diff --git a/.github/bot_config.yml b/.github/bot_config.yml
index 952afc316e7..9d90ab82b33 100644
--- a/.github/bot_config.yml
+++ b/.github/bot_config.yml
@@ -21,4 +21,6 @@
# A list of assignees
assignees:
- - saikumarchalla
+ - laxmareddyp
+ - jettibharat
+ - lakshmikala
diff --git a/.github/stale.yml b/.github/stale.yml
deleted file mode 100644
index 7eef5309ecd..00000000000
--- a/.github/stale.yml
+++ /dev/null
@@ -1,39 +0,0 @@
- # Copyright 2019 The TensorFlow Authors. All Rights Reserved.
- #
- # Licensed under the Apache License, Version 2.0 (the "License");
- # you may not use this file except in compliance with the License.
- # You may obtain a copy of the License at
- #
- # http://www.apache.org/licenses/LICENSE-2.0
- #
- # Unless required by applicable law or agreed to in writing, software
- # distributed under the License is distributed on an "AS IS" BASIS,
- # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
- # See the License for the specific language governing permissions and
- # limitations under the License.
- # ============================================================================
- #
- # THIS IS A GENERATED DOCKERFILE.
- #
- # This file was assembled from multiple pieces, whose use is documented
- # throughout. Please refer to the TensorFlow dockerfiles documentation
- # for more information.
-
-# Number of days of inactivity before an Issue or Pull Request becomes stale
-daysUntilStale: 7
-# Number of days of inactivity before a stale Issue or Pull Request is closed
-daysUntilClose: 7
-# Only issues or pull requests with all of these labels are checked if stale. Defaults to `[]` (disabled)
-onlyLabels:
- - stat:awaiting response
-# Comment to post when marking as stale. Set to `false` to disable
-markComment: >
- This issue has been automatically marked as stale because it has not had
- recent activity. It will be closed if no further activity occurs. Thank you.
-# Comment to post when removing the stale label. Set to `false` to disable
-unmarkComment: false
-closeComment: >
- Closing as stale. Please reopen if you'd like to work on this further.
-limitPerRun: 30
-# Limit to only `issues` or `pulls`
-only: issues
diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml
index 744f440b053..eb321f8a943 100644
--- a/.github/workflows/ci.yml
+++ b/.github/workflows/ci.yml
@@ -1,6 +1,9 @@
name: CI
on: pull_request
+permissions:
+ contents: read
+
jobs:
pylint:
runs-on: ubuntu-latest
diff --git a/.github/workflows/stale.yaml b/.github/workflows/stale.yaml
new file mode 100644
index 00000000000..2fab200edac
--- /dev/null
+++ b/.github/workflows/stale.yaml
@@ -0,0 +1,67 @@
+# Copyright 2023 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+# ==============================================================================
+
+# This workflow alerts and then closes the stale issues/PRs after specific time
+# You can adjust the behavior by modifying this file.
+# For more information, see:
+# https://github.com/actions/stale
+
+name: 'Close stale issues and PRs'
+"on":
+ schedule:
+ - cron: "30 1 * * *"
+permissions:
+ contents: read
+ issues: write
+ pull-requests: write
+
+jobs:
+ stale:
+ runs-on: ubuntu-latest
+ steps:
+ - uses: 'actions/stale@v7'
+ with:
+ #Comma separated list of labels that can be assigned to issues to exclude them from being marked as stale
+ exempt-issue-labels: 'override-stale'
+ #Comma separated list of labels that can be assigned to PRs to exclude them from being marked as stale
+ exempt-pr-labels: "override-stale"
+ #Limit the No. of API calls in one run default value is 30.
+ operations-per-run: 1000
+ #Prevent to remove stale label when PRs or issues are updated.
+ remove-stale-when-updated: false
+ # comment on issue if not active for more then 7 days.
+ stale-issue-message: 'This issue has been marked stale because it has no recent activity since 7 days. It will be closed if no further activity occurs. Thank you.'
+ # comment on PR if not active for more then 14 days.
+ stale-pr-message: 'This PR has been marked stale because it has no recent activity since 14 days. It will be closed if no further activity occurs. Thank you.'
+ # comment on issue if stale for more then 7 days.
+ close-issue-message: This issue was closed due to lack of activity after being marked stale for past 7 days.
+ # comment on PR if stale for more then 14 days.
+ close-pr-message: This PR was closed due to lack of activity after being marked stale for past 14 days.
+ # Number of days of inactivity before an Issue Request becomes stale
+ days-before-issue-stale: 7
+ # Number of days of inactivity before a stale Issue is closed
+ days-before-issue-close: 7
+ # reason for closed the issue default value is not_planned
+ close-issue-reason: completed
+ # Number of days of inactivity before a stale PR is closed
+ days-before-pr-close: 14
+ # Number of days of inactivity before an PR Request becomes stale
+ days-before-pr-stale: 14
+ # Check for label to stale or close the issue/PR
+ any-of-labels: 'stat:awaiting response'
+ # override stale to stalled for PR
+ stale-pr-label: 'stale'
+ # override stale to stalled for Issue
+ stale-issue-label: "stale"
diff --git a/CODEOWNERS b/CODEOWNERS
index a46d472ca35..0a9c6abb62f 100644
--- a/CODEOWNERS
+++ b/CODEOWNERS
@@ -1,13 +1,13 @@
* @tensorflow/tf-model-garden-team
-/official/ @rachellj218 @saberkun @jaeyounkim
-/official/nlp/ @saberkun @lehougoogle @rachellj218 @jaeyounkim
+/official/ @rachellj218
+/official/nlp/ @lehougoogle @rachellj218
/official/recommendation/ranking/ @gagika
-/official/vision/ @xianzhidu @yeqingli @arashwan @saberkun @rachellj218 @jaeyounkim
-/official/vision/beta/projects/assemblenet/ @mryoo @yeqingli
-/official/vision/beta/projects/deepmac_maskrcnn/ @vighneshbirodkar
-/official/vision/beta/projects/movinet/ @hyperparticle @yuanliangzhe @yeqingli
-/official/vision/beta/projects/simclr/ @luotigerlsx @chentingpc @saxenasaurabh
-/official/vision/beta/projects/video_ssl/ @richardaecn @yeqingli
+/official/vision/ @yeqingli @arashwan @rachellj218
+/official/vision/projects/assemblenet/ @yeqingli
+/official/vision/projects/deepmac_maskrcnn/ @vighneshbirodkar
+/official/vision/projects/movinet/ @hyperparticle @yuanliangzhe @yeqingli
+/official/vision/projects/simclr/ @luotigerlsx @saxenasaurabh
+/official/vision/projects/video_ssl/ @yeqingli
/research/adversarial_text/ @rsepassi @a-dai
/research/attention_ocr/ @xavigibert
/research/audioset/ @plakal @dpwe
@@ -24,6 +24,6 @@
/research/object_detection/ @jch1 @tombstone @pkulzc
/research/pcl_rl/ @ofirnachum
/research/rebar/ @gjtucker
-/research/seq_flow_lite/ @thunderfyc
+/research/seq_flow_lite/ @thunderfyc @karunreddy30
/research/slim/ @sguada @marksandler2
/research/vid2depth/ @rezama
diff --git a/README.md b/README.md
index e9c928a4e4f..db37fce8ad1 100644
--- a/README.md
+++ b/README.md
@@ -3,7 +3,8 @@
[](https://badge.fury.io/py/tensorflow)
-[](https://badge.fury.io/py/tensorflow)
+[](https://badge.fury.io/py/tf-models-official)
+
# Welcome to the Model Garden for TensorFlow
@@ -19,7 +20,7 @@ extent possible though not all models are suitable.
| Directory | Description |
|-----------|-------------|
-| [official](official) | • A collection of example implementations for SOTA models using the latest TensorFlow 2's high-level APIs • Officially maintained, supported, and kept up to date with the latest TensorFlow 2 APIs by TensorFlow • Reasonably optimized for fast performance while still being easy to read |
+| [official](official) | • A collection of example implementations for SOTA models using the latest TensorFlow 2's high-level APIs • Officially maintained, supported, and kept up to date with the latest TensorFlow 2 APIs by TensorFlow • Reasonably optimized for fast performance while still being easy to read For more details on the capabilities, check the guide on the [Model-garden](https://www.tensorflow.org/tfmodels)|
| [research](research) | • A collection of research model implementations in TensorFlow 1 or 2 by researchers • Maintained and supported by researchers |
| [community](community) | • A curated list of the GitHub repositories with machine learning models and implementations powered by TensorFlow 2 |
| [orbit](orbit) | • A flexible and lightweight library that users can easily use or fork when writing customized training loop code in TensorFlow 2.x. It seamlessly integrates with `tf.distribute` and supports running on different device types (CPU, GPU, and TPU). |
@@ -32,23 +33,20 @@ To install the current release of tensorflow-models, please follow any one of th
-**tf-models-official** is the stable Model Garden package.
-pip will install all models and dependencies automatically.
+**tf-models-official** is the stable Model Garden package. Please check out the [releases](https://github.com/tensorflow/models/releases) to see what are available modules.
-```shell
-pip3 install tf-models-official
-```
-
-If you are using nlp packages, please also install **tensorflow-text**:
+pip3 will install all models and dependencies automatically.
```shell
-pip3 install tensorflow-text
+pip3 install tf-models-official
```
-Please check out our [example](https://github.com/tensorflow/text/blob/master/docs/tutorials/fine_tune_bert.ipynb)
+Please check out our examples:
+ - [basic library import](https://github.com/tensorflow/models/blob/master/tensorflow_models/tensorflow_models_pypi.ipynb)
+ - [nlp model building](https://github.com/tensorflow/models/blob/master/docs/nlp/index.ipynb)
to learn how to use a PIP package.
-Note that **tf-models-official** may not include the latest changes in this
+Note that **tf-models-official** may not include the latest changes in the master branch of this
github repo. To include latest changes, you may install **tf-models-nightly**,
which is the nightly Model Garden package created daily automatically.
@@ -56,11 +54,6 @@ which is the nightly Model Garden package created daily automatically.
pip3 install tf-models-nightly
```
-If you are using `nlp` packages, please also install tensorflow-text-nightly
-
-```shell
-pip3 install tensorflow-text-nightly
-```
@@ -80,6 +73,11 @@ git clone https://github.com/tensorflow/models.git
export PYTHONPATH=$PYTHONPATH:/path/to/models
```
+If you are using in a Windows environment, you may need to use the following command with PowerShell:
+```shell
+$env:PYTHONPATH += ":\path\to\models"
+```
+
If you are using a Colab notebook, please set the Python path with os.environ.
```python
@@ -90,7 +88,7 @@ os.environ['PYTHONPATH'] += ":/path/to/models"
3. Install other dependencies
```shell
-pip3 install --user -r official/requirements.txt
+pip3 install --user -r models/official/requirements.txt
```
Finally, if you are using nlp packages, please also install
@@ -123,8 +121,8 @@ If you use TensorFlow Model Garden in your research, please cite this repository
```
@misc{tensorflowmodelgarden2020,
- author = {Hongkun Yu, Chen Chen, Xianzhi Du, Yeqing Li, Abdullah Rashwan, Le Hou, Pengchong Jin, Fan Yang,
- Frederick Liu, Jaeyoun Kim, and Jing Li},
+ author = {Hongkun Yu and Chen Chen and Xianzhi Du and Yeqing Li and Abdullah Rashwan and Le Hou and Pengchong Jin and Fan Yang
+ and Frederick Liu and Jaeyoun Kim and Jing Li},
title = {{TensorFlow Model Garden}},
howpublished = {\url{https://github.com/tensorflow/models}},
year = {2020}
diff --git a/docs/README.md b/docs/README.md
new file mode 100644
index 00000000000..dcf37f88464
--- /dev/null
+++ b/docs/README.md
@@ -0,0 +1,17 @@
+# Public docs for TensorFlow Models
+
+This directory contains the top-level public documentation for
+[TensorFlow Models](https://github.com/tensorflow/models).
+
+This directory is mirrored to https://tensorflow.org/tfmodels, and is mainly
+concerned with documenting the tools provided in the `tensorflow_models` pip
+package (including `orbit`).
+
+Api-reference pages are
+[available on the site](https://www.tensorflow.org/api_docs/more).
+
+The
+[Official Models](https://github.com/tensorflow/models/blob/master/official/projects)
+and [Research Models](https://github.com/tensorflow/models/blob/master/research)
+directories are not described in detail here, refer to the individual project
+directories for more information.
diff --git a/docs/index.md b/docs/index.md
new file mode 100644
index 00000000000..2b1535a71fe
--- /dev/null
+++ b/docs/index.md
@@ -0,0 +1,140 @@
+# Model Garden overview
+
+The TensorFlow Model Garden provides implementations of many state-of-the-art
+machine learning (ML) models for vision and natural language processing (NLP),
+as well as workflow tools to let you quickly configure and run those models on
+standard datasets. Whether you are looking to benchmark performance for a
+well-known model, verify the results of recently released research, or extend
+existing models, the Model Garden can help you drive your ML research and
+applications forward.
+
+The Model Garden includes the following resources for machine learning
+developers:
+
+- [**Official models**](#official) for vision and NLP, maintained by Google
+ engineers
+- [**Research models**](#research) published as part of ML research papers
+- [**Training experiment framework**](#training_framework) for fast,
+ declarative training configuration of official models
+- [**Specialized ML operations**](#ops) for vision and natural language
+ processing (NLP)
+- [**Model training loop**](#orbit) management with Orbit
+
+These resources are built to be used with the TensorFlow Core framework and
+integrate with your existing TensorFlow development projects. Model
+Garden resources are also provided under an [open
+source](https://github.com/tensorflow/models/blob/master/LICENSE) license, so
+you can freely extend and distribute the models and tools.
+
+Practical ML models are computationally intensive to train and run, and may
+require accelerators such as Graphical Processing Units (GPUs) and Tensor
+Processing Units (TPUs). Most of the models in Model Garden were trained on
+large datasets using TPUs. However, you can also train and run these models on
+GPU and CPU processors.
+
+## Model Garden models
+
+The machine learning models in the Model Garden include full code so you can
+test, train, or re-train them for research and experimentation. The Model Garden
+includes two primary categories of models: *official models* and *research
+models*.
+
+### Official models {:#official}
+
+The [Official Models](https://github.com/tensorflow/models/tree/master/official)
+repository is a collection of state-of-the-art models, with a focus on
+vision and natural language processing (NLP).
+These models are implemented using current TensorFlow 2.x high-level
+APIs. Model libraries in this repository are optimized for fast performance and
+actively maintained by Google engineers. The official models include additional
+metadata you can use to quickly configure experiments using the Model Garden
+[training experiment framework](#training_framework).
+
+### Research models {:#research}
+
+The [Research Models](https://github.com/tensorflow/models/tree/master/research)
+repository is a collection of models published as code resources for research
+papers. These models are implemented using both TensorFlow 1.x and 2.x. Model
+libraries in the research folder are supported by the code owners and the
+research community.
+
+## Training experiment framework {:#training_framework}
+
+The Model Garden training experiment framework lets you quickly assemble and run
+training experiments using its official models and standard datasets. The
+training framework uses additional metadata included with the Model Garden's
+official models to allow you to configure models quickly using a declarative
+programming model. You can define a training experiment using Python commands in
+the
+[TensorFlow Model library](https://www.tensorflow.org/api_docs/python/tfm/core)
+or configure training using a YAML configuration file, like this
+[example](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml).
+
+The training framework uses
+[`tfm.core.base_trainer.ExperimentConfig`](https://www.tensorflow.org/api_docs/python/tfm/core/base_trainer/ExperimentConfig)
+as the configuration object, which contains the following top-level
+configuration objects:
+
+- [`runtime`](https://www.tensorflow.org/api_docs/python/tfm/core/base_task/RuntimeConfig):
+ Defines the processing hardware, distribution strategy, and other
+ performance optimizations
+- [`task`](https://www.tensorflow.org/api_docs/python/tfm/core/config_definitions/TaskConfig):
+ Defines the model, training data, losses, and initialization
+- [`trainer`](https://www.tensorflow.org/api_docs/python/tfm/core/base_trainer/TrainerConfig):
+ Defines the optimizer, training loops, evaluation loops, summaries, and
+ checkpoints
+
+For a complete example using the Model Garden training experiment framework, see
+the [Image classification with Model Garden](vision/image_classification.ipynb)
+tutorial. For information on the training experiment framework, check out the
+[TensorFlow Models API documentation](https://tensorflow.org/api_docs/python/tfm/core).
+If you are looking for a solution to manage training loops for your model
+training experiments, check out [Orbit](#orbit).
+
+## Specialized ML operations {:#ops}
+
+The Model Garden contains many vision and NLP operations specifically designed
+to execute state-of-the-art models that run efficiently on GPUs and TPUs. Review
+the TensorFlow Models Vision library API docs for a list of specialized
+[vision operations](https://www.tensorflow.org/api_docs/python/tfm/vision).
+Review the TensorFlow Models NLP Library API docs for a list of
+[NLP operations](https://www.tensorflow.org/api_docs/python/tfm/nlp). These
+libraries also include additional utility functions used for vision and NLP data
+processing, training, and model execution.
+
+## Training loops with Orbit {:#orbit}
+
+There are two default options for training TensorFlow models:
+
+* Use the high-level Keras
+[Model.fit](https://www.tensorflow.org/api_docs/python/tf/keras/Model#fit)
+function. If your model and training procedure fit the assumptions of Keras'
+`Model.fit` (incremental gradient descent on batches of data) method this can
+be very convenient.
+* Write a custom training loop
+[with keras](https://www.tensorflow.org/guide/keras/writing_a_training_loop_from_scratch),
+or [without](https://www.tensorflow.org/guide/core/logistic_regression_core).
+You can write a custom training loop with low-level TensorFlow methods such as
+`tf.GradientTape` or `tf.function`. However, this approach requires a lot of
+boilerplate code, and doesn't do anything to simplify distributed training.
+
+Orbit tries to provide a third option in between these two extremes.
+
+Orbit is a flexible, lightweight library designed to make it easier to
+write custom training loops in TensorFlow 2.x, and works well with the Model
+Garden [training experiment framework](#training_framework). Orbit handles
+common model training tasks such as saving checkpoints, running model
+evaluations, and setting up summary writing. It seamlessly integrates with
+`tf.distribute` and supports running on different device types, including CPU,
+GPU, and TPU hardware. The Orbit tool is also [open
+source](https://github.com/tensorflow/models/blob/master/orbit/LICENSE), so you
+can extend and adapt to your model training needs.
+
+The Orbit guide is available [here](orbit/index.ipynb).
+
+Note: You can customize how the Keras API executes training. Mainly you must
+override the `Model.train_step` method or use `keras.callbacks` like
+`callbacks.ModelCheckpoint` or `callbacks.TensorBoard`. For more information
+about modifying the behavior of `train_step`, check out the
+[Customize what happens in Model.fit](https://www.tensorflow.org/guide/keras/customizing_what_happens_in_fit)
+page.
diff --git a/docs/nlp/_guide_toc.yaml b/docs/nlp/_guide_toc.yaml
new file mode 100644
index 00000000000..84b866d8628
--- /dev/null
+++ b/docs/nlp/_guide_toc.yaml
@@ -0,0 +1,9 @@
+toc:
+- heading: TensorFlow Models - NLP
+ style: divider
+- title: "Overview"
+ path: /tfmodels/nlp
+- title: "Customize a transformer encoder"
+ path: /tfmodels/nlp/customize_encoder
+- title: "Load LM checkpoints"
+ path: /tfmodels/nlp/load_lm_ckpts
diff --git a/official/colab/nlp/customize_encoder.ipynb b/docs/nlp/customize_encoder.ipynb
similarity index 69%
rename from official/colab/nlp/customize_encoder.ipynb
rename to docs/nlp/customize_encoder.ipynb
index aeddb29f963..7e6fd0f32af 100644
--- a/official/colab/nlp/customize_encoder.ipynb
+++ b/docs/nlp/customize_encoder.ipynb
@@ -1,19 +1,4 @@
{
- "nbformat": 4,
- "nbformat_minor": 0,
- "metadata": {
- "colab": {
- "name": "Customizing a Transformer Encoder",
- "private_outputs": true,
- "provenance": [],
- "collapsed_sections": [],
- "toc_visible": true
- },
- "kernelspec": {
- "display_name": "Python 3",
- "name": "python3"
- }
- },
"cells": [
{
"cell_type": "markdown",
@@ -21,15 +6,17 @@
"id": "Bp8t2AI8i7uP"
},
"source": [
- "##### Copyright 2020 The TensorFlow Authors."
+ "##### Copyright 2022 The TensorFlow Authors."
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"cellView": "form",
"id": "rxPj2Lsni9O4"
},
+ "outputs": [],
"source": [
"#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n",
"# you may not use this file except in compliance with the License.\n",
@@ -42,9 +29,7 @@
"# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
"# See the License for the specific language governing permissions and\n",
"# limitations under the License."
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -63,16 +48,16 @@
"source": [
"
"
]
@@ -87,7 +72,7 @@
"\n",
"The [TensorFlow Models NLP library](https://github.com/tensorflow/models/tree/master/official/nlp/modeling) is a collection of tools for building and training modern high performance natural language models.\n",
"\n",
- "The [TransformEncoder](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/networks/encoder_scaffold.py) is the core of this library, and lots of new network architectures are proposed to improve the encoder. In this Colab notebook, we will learn how to customize the encoder to employ new network architectures."
+ "The `tfm.nlp.networks.EncoderScaffold` is the core of this library, and lots of new network architectures are proposed to improve the encoder. In this Colab notebook, we will learn how to customize the encoder to employ new network architectures."
]
},
{
@@ -114,14 +99,25 @@
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
- "id": "thsKZDjhswhR"
+ "id": "mfHI5JyuJ1y9"
},
+ "outputs": [],
"source": [
- "!pip install -q tf-models-official==2.4.0"
- ],
+ "!pip install -q opencv-python"
+ ]
+ },
+ {
+ "cell_type": "code",
"execution_count": null,
- "outputs": []
+ "metadata": {
+ "id": "thsKZDjhswhR"
+ },
+ "outputs": [],
+ "source": [
+ "!pip install -q tf-models-official"
+ ]
},
{
"cell_type": "markdown",
@@ -134,19 +130,18 @@
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "my4dp-RMssQe"
},
+ "outputs": [],
"source": [
"import numpy as np\n",
"import tensorflow as tf\n",
"\n",
- "from official.modeling import activations\n",
- "from official.nlp import modeling\n",
- "from official.nlp.modeling import layers, losses, models, networks"
- ],
- "execution_count": null,
- "outputs": []
+ "import tensorflow_models as tfm\n",
+ "nlp = tfm.nlp"
+ ]
},
{
"cell_type": "markdown",
@@ -156,14 +151,16 @@
"source": [
"## Canonical BERT encoder\n",
"\n",
- "Before learning how to customize the encoder, let's firstly create a canonical BERT enoder and use it to instantiate a `BertClassifier` for classification task."
+ "Before learning how to customize the encoder, let's firstly create a canonical BERT enoder and use it to instantiate a `bert_classifier.BertClassifier` for classification task."
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "Oav8sbgstWc-"
},
+ "outputs": [],
"source": [
"cfg = {\n",
" \"vocab_size\": 100,\n",
@@ -171,22 +168,20 @@
" \"num_layers\": 3,\n",
" \"num_attention_heads\": 4,\n",
" \"intermediate_size\": 64,\n",
- " \"activation\": activations.gelu,\n",
+ " \"activation\": tfm.utils.activations.gelu,\n",
" \"dropout_rate\": 0.1,\n",
" \"attention_dropout_rate\": 0.1,\n",
" \"max_sequence_length\": 16,\n",
" \"type_vocab_size\": 2,\n",
" \"initializer\": tf.keras.initializers.TruncatedNormal(stddev=0.02),\n",
"}\n",
- "bert_encoder = modeling.networks.BertEncoder(**cfg)\n",
+ "bert_encoder = nlp.networks.BertEncoder(**cfg)\n",
"\n",
"def build_classifier(bert_encoder):\n",
- " return modeling.models.BertClassifier(bert_encoder, num_classes=2)\n",
+ " return nlp.models.BertClassifier(bert_encoder, num_classes=2)\n",
"\n",
"canonical_classifier_model = build_classifier(bert_encoder)"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -194,16 +189,18 @@
"id": "Qe2UWI6_tsHo"
},
"source": [
- "`canonical_classifier_model` can be trained using the training data. For details about how to train the model, please see the colab [fine_tuning_bert.ipynb](https://github.com/tensorflow/models/blob/master/official/colab/fine_tuning_bert.ipynb). We skip the code that trains the model here.\n",
+ "`canonical_classifier_model` can be trained using the training data. For details about how to train the model, please see the [Fine tuning bert](https://www.tensorflow.org/text/tutorials/fine_tune_bert) notebook. We skip the code that trains the model here.\n",
"\n",
"After training, we can apply the model to do prediction.\n"
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "csED2d-Yt5h6"
},
+ "outputs": [],
"source": [
"def predict(model):\n",
" batch_size = 3\n",
@@ -216,9 +213,7 @@
" print(model([word_ids, mask, type_ids], training=False))\n",
"\n",
"predict(canonical_classifier_model)"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -249,7 +244,7 @@
"source": [
"### Use EncoderScaffold\n",
"\n",
- "`EncoderScaffold` allows users to provide a custom embedding subnetwork\n",
+ "`networks.EncoderScaffold` allows users to provide a custom embedding subnetwork\n",
" (which will replace the standard embedding logic) and/or a custom hidden layer class (which will replace the `Transformer` instantiation in the encoder)."
]
},
@@ -261,30 +256,32 @@
"source": [
"#### Without Customization\n",
"\n",
- "Without any customization, `EncoderScaffold` behaves the same the canonical `BertEncoder`.\n",
+ "Without any customization, `networks.EncoderScaffold` behaves the same the canonical `networks.BertEncoder`.\n",
"\n",
- "As shown in the following example, `EncoderScaffold` can load `BertEncoder`'s weights and output the same values:"
+ "As shown in the following example, `networks.EncoderScaffold` can load `networks.BertEncoder`'s weights and output the same values:"
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "ktNzKuVByZQf"
},
+ "outputs": [],
"source": [
"default_hidden_cfg = dict(\n",
" num_attention_heads=cfg[\"num_attention_heads\"],\n",
" intermediate_size=cfg[\"intermediate_size\"],\n",
- " intermediate_activation=activations.gelu,\n",
+ " intermediate_activation=cfg[\"activation\"],\n",
" dropout_rate=cfg[\"dropout_rate\"],\n",
" attention_dropout_rate=cfg[\"attention_dropout_rate\"],\n",
- " kernel_initializer=tf.keras.initializers.TruncatedNormal(0.02),\n",
+ " kernel_initializer=cfg[\"initializer\"],\n",
")\n",
"default_embedding_cfg = dict(\n",
" vocab_size=cfg[\"vocab_size\"],\n",
" type_vocab_size=cfg[\"type_vocab_size\"],\n",
" hidden_size=cfg[\"hidden_size\"],\n",
- " initializer=tf.keras.initializers.TruncatedNormal(0.02),\n",
+ " initializer=cfg[\"initializer\"],\n",
" dropout_rate=cfg[\"dropout_rate\"],\n",
" max_seq_length=cfg[\"max_sequence_length\"]\n",
")\n",
@@ -294,17 +291,15 @@
" num_hidden_instances=cfg[\"num_layers\"],\n",
" pooled_output_dim=cfg[\"hidden_size\"],\n",
" return_all_layer_outputs=True,\n",
- " pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(0.02),\n",
+ " pooler_layer_initializer=cfg[\"initializer\"],\n",
")\n",
"\n",
- "encoder_scaffold = modeling.networks.EncoderScaffold(**default_kwargs)\n",
+ "encoder_scaffold = nlp.networks.EncoderScaffold(**default_kwargs)\n",
"classifier_model_from_encoder_scaffold = build_classifier(encoder_scaffold)\n",
"classifier_model_from_encoder_scaffold.set_weights(\n",
" canonical_classifier_model.get_weights())\n",
"predict(classifier_model_from_encoder_scaffold)"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -316,31 +311,31 @@
"\n",
"Next, we show how to use a customized embedding network.\n",
"\n",
- "We firstly build an embedding network that will replace the default network. This one will have 2 inputs (`mask` and `word_ids`) instead of 3, and won't use positional embeddings."
+ "We first build an embedding network that would replace the default network. This one will have 2 inputs (`mask` and `word_ids`) instead of 3, and won't use positional embeddings."
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "LTinnaG6vcsw"
},
+ "outputs": [],
"source": [
"word_ids = tf.keras.layers.Input(\n",
" shape=(cfg['max_sequence_length'],), dtype=tf.int32, name=\"input_word_ids\")\n",
"mask = tf.keras.layers.Input(\n",
" shape=(cfg['max_sequence_length'],), dtype=tf.int32, name=\"input_mask\")\n",
- "embedding_layer = modeling.layers.OnDeviceEmbedding(\n",
+ "embedding_layer = nlp.layers.OnDeviceEmbedding(\n",
" vocab_size=cfg['vocab_size'],\n",
" embedding_width=cfg['hidden_size'],\n",
- " initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),\n",
+ " initializer=cfg[\"initializer\"],\n",
" name=\"word_embeddings\")\n",
"word_embeddings = embedding_layer(word_ids)\n",
- "attention_mask = layers.SelfAttentionMask()([word_embeddings, mask])\n",
+ "attention_mask = nlp.layers.SelfAttentionMask()([word_embeddings, mask])\n",
"new_embedding_network = tf.keras.Model([word_ids, mask],\n",
" [word_embeddings, attention_mask])"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -354,14 +349,14 @@
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "fO9zKFE4OpHp"
},
+ "outputs": [],
"source": [
"tf.keras.utils.plot_model(new_embedding_network, show_shapes=True, dpi=48)"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -369,14 +364,16 @@
"id": "9cOaGQHLv12W"
},
"source": [
- "We then can build a new encoder using the above `new_embedding_network`."
+ "We can then build a new encoder using the above `new_embedding_network`."
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "mtFDMNf2vIl9"
},
+ "outputs": [],
"source": [
"kwargs = dict(default_kwargs)\n",
"\n",
@@ -384,16 +381,14 @@
"kwargs['embedding_cls'] = new_embedding_network\n",
"kwargs['embedding_data'] = embedding_layer.embeddings\n",
"\n",
- "encoder_with_customized_embedding = modeling.networks.EncoderScaffold(**kwargs)\n",
+ "encoder_with_customized_embedding = nlp.networks.EncoderScaffold(**kwargs)\n",
"classifier_model = build_classifier(encoder_with_customized_embedding)\n",
"# ... Train the model ...\n",
"print(classifier_model.inputs)\n",
"\n",
"# Assert that there are only two inputs.\n",
"assert len(classifier_model.inputs) == 2"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -403,34 +398,34 @@
"source": [
"#### Customized Transformer\n",
"\n",
- "User can also override the [hidden_cls](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/networks/encoder_scaffold.py#L103) argument in `EncoderScaffold`'s constructor to employ a customized Transformer layer.\n",
+ "Users can also override the `hidden_cls` argument in `networks.EncoderScaffold`'s constructor employ a customized Transformer layer.\n",
"\n",
- "See [ReZeroTransformer](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/rezero_transformer.py) for how to implement a customized Transformer layer.\n",
+ "See [the source of `nlp.layers.ReZeroTransformer`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/rezero_transformer.py) for how to implement a customized Transformer layer.\n",
"\n",
- "Following is an example of using `ReZeroTransformer`:\n"
+ "The following is an example of using `nlp.layers.ReZeroTransformer`:\n"
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "uAIarLZgw6pA"
},
+ "outputs": [],
"source": [
"kwargs = dict(default_kwargs)\n",
"\n",
"# Use ReZeroTransformer.\n",
- "kwargs['hidden_cls'] = modeling.layers.ReZeroTransformer\n",
+ "kwargs['hidden_cls'] = nlp.layers.ReZeroTransformer\n",
"\n",
- "encoder_with_rezero_transformer = modeling.networks.EncoderScaffold(**kwargs)\n",
+ "encoder_with_rezero_transformer = nlp.networks.EncoderScaffold(**kwargs)\n",
"classifier_model = build_classifier(encoder_with_rezero_transformer)\n",
"# ... Train the model ...\n",
"predict(classifier_model)\n",
"\n",
"# Assert that the variable `rezero_alpha` from ReZeroTransformer exists.\n",
"assert 'rezero_alpha' in ''.join([x.name for x in classifier_model.trainable_weights])"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -438,10 +433,9 @@
"id": "6PMHFdvnxvR0"
},
"source": [
- "### Use [TransformerScaffold](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/transformer_scaffold.py)\n",
+ "### Use `nlp.layers.TransformerScaffold`\n",
"\n",
- "The above method of customizing `Transformer` requires rewriting the whole `Transformer` layer, while sometimes you may only want to customize either attention layer or feedforward block. In this case, [TransformerScaffold](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/transformer_scaffold.py) can be used.\n",
- "\n"
+ "The above method of customizing the model requires rewriting the whole `nlp.layers.Transformer` layer, while sometimes you may only want to customize either attention layer or feedforward block. In this case, `nlp.layers.TransformerScaffold` can be used.\n"
]
},
{
@@ -452,37 +446,48 @@
"source": [
"#### Customize Attention Layer\n",
"\n",
- "User can also override the [attention_cls](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/transformer_scaffold.py#L45) argument in `TransformerScaffold`'s constructor to employ a customized Attention layer.\n",
+ "User can also override the `attention_cls` argument in `layers.TransformerScaffold`'s constructor to employ a customized Attention layer.\n",
"\n",
- "See [TalkingHeadsAttention](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/talking_heads_attention.py) for how to implement a customized `Attention` layer.\n",
+ "See [the source of `nlp.layers.TalkingHeadsAttention`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/talking_heads_attention.py) for how to implement a customized `Attention` layer.\n",
"\n",
- "Following is an example of using [TalkingHeadsAttention](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/talking_heads_attention.py):"
+ "Following is an example of using `nlp.layers.TalkingHeadsAttention`:"
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "nFrSMrZuyNeQ"
},
+ "outputs": [],
"source": [
"# Use TalkingHeadsAttention\n",
"hidden_cfg = dict(default_hidden_cfg)\n",
- "hidden_cfg['attention_cls'] = modeling.layers.TalkingHeadsAttention\n",
+ "hidden_cfg['attention_cls'] = nlp.layers.TalkingHeadsAttention\n",
"\n",
"kwargs = dict(default_kwargs)\n",
- "kwargs['hidden_cls'] = modeling.layers.TransformerScaffold\n",
+ "kwargs['hidden_cls'] = nlp.layers.TransformerScaffold\n",
"kwargs['hidden_cfg'] = hidden_cfg\n",
"\n",
- "encoder = modeling.networks.EncoderScaffold(**kwargs)\n",
+ "encoder = nlp.networks.EncoderScaffold(**kwargs)\n",
"classifier_model = build_classifier(encoder)\n",
"# ... Train the model ...\n",
"predict(classifier_model)\n",
"\n",
"# Assert that the variable `pre_softmax_weight` from TalkingHeadsAttention exists.\n",
"assert 'pre_softmax_weight' in ''.join([x.name for x in classifier_model.trainable_weights])"
- ],
+ ]
+ },
+ {
+ "cell_type": "code",
"execution_count": null,
- "outputs": []
+ "metadata": {
+ "id": "tKkZ8spzYmpc"
+ },
+ "outputs": [],
+ "source": [
+ "tf.keras.utils.plot_model(encoder_with_rezero_transformer, show_shapes=True, dpi=48)"
+ ]
},
{
"cell_type": "markdown",
@@ -494,35 +499,35 @@
"\n",
"Similiarly, one could also customize the feedforward layer.\n",
"\n",
- "See [GatedFeedforward](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/gated_feedforward.py) for how to implement a customized feedforward layer.\n",
+ "See [the source of `nlp.layers.GatedFeedforward`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/gated_feedforward.py) for how to implement a customized feedforward layer.\n",
"\n",
- "Following is an example of using [GatedFeedforward](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/layers/gated_feedforward.py)."
+ "Following is an example of using `nlp.layers.GatedFeedforward`:"
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "XAbKy_l4y_-i"
},
+ "outputs": [],
"source": [
- "# Use TalkingHeadsAttention\n",
+ "# Use GatedFeedforward\n",
"hidden_cfg = dict(default_hidden_cfg)\n",
- "hidden_cfg['feedforward_cls'] = modeling.layers.GatedFeedforward\n",
+ "hidden_cfg['feedforward_cls'] = nlp.layers.GatedFeedforward\n",
"\n",
"kwargs = dict(default_kwargs)\n",
- "kwargs['hidden_cls'] = modeling.layers.TransformerScaffold\n",
+ "kwargs['hidden_cls'] = nlp.layers.TransformerScaffold\n",
"kwargs['hidden_cfg'] = hidden_cfg\n",
"\n",
- "encoder_with_gated_feedforward = modeling.networks.EncoderScaffold(**kwargs)\n",
+ "encoder_with_gated_feedforward = nlp.networks.EncoderScaffold(**kwargs)\n",
"classifier_model = build_classifier(encoder_with_gated_feedforward)\n",
"# ... Train the model ...\n",
"predict(classifier_model)\n",
"\n",
"# Assert that the variable `gate` from GatedFeedforward exists.\n",
"assert 'gate' in ''.join([x.name for x in classifier_model.trainable_weights])"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -530,26 +535,28 @@
"id": "a_8NWUhkzeAq"
},
"source": [
- "### Build a new Encoder using building blocks from KerasBERT.\n",
+ "### Build a new Encoder\n",
"\n",
"Finally, you could also build a new encoder using building blocks in the modeling library.\n",
"\n",
- "See [AlbertEncoder](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/networks/albert_encoder.py) as an example:\n"
+ "See [the source for `nlp.networks.AlbertEncoder`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/networks/albert_encoder.py) as an example of how to do this. \n",
+ "\n",
+ "Here is an example using `nlp.networks.AlbertEncoder`:\n"
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "xsiA3RzUzmUM"
},
+ "outputs": [],
"source": [
- "albert_encoder = modeling.networks.AlbertEncoder(**cfg)\n",
+ "albert_encoder = nlp.networks.AlbertEncoder(**cfg)\n",
"classifier_model = build_classifier(albert_encoder)\n",
"# ... Train the model ...\n",
"predict(classifier_model)"
- ],
- "execution_count": null,
- "outputs": []
+ ]
},
{
"cell_type": "markdown",
@@ -557,19 +564,33 @@
"id": "MeidDfhlHKSO"
},
"source": [
- "Inspecting the `albert_encoder`, we see it stacks the same `Transformer` layer multiple times."
+ "Inspecting the `albert_encoder`, we see it stacks the same `Transformer` layer multiple times (note the loop-back on the \"Transformer\" block below.."
]
},
{
"cell_type": "code",
+ "execution_count": null,
"metadata": {
"id": "Uv_juT22HERW"
},
+ "outputs": [],
"source": [
"tf.keras.utils.plot_model(albert_encoder, show_shapes=True, dpi=48)"
- ],
- "execution_count": null,
- "outputs": []
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "collapsed_sections": [],
+ "name": "customize_encoder.ipynb",
+ "provenance": [],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
}
- ]
-}
\ No newline at end of file
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/colab/decoding_api_in_tf_nlp.ipynb b/docs/nlp/decoding_api.ipynb
similarity index 74%
rename from official/colab/decoding_api_in_tf_nlp.ipynb
rename to docs/nlp/decoding_api.ipynb
index 726b382e228..0846c40263e 100644
--- a/official/colab/decoding_api_in_tf_nlp.ipynb
+++ b/docs/nlp/decoding_api.ipynb
@@ -31,6 +31,37 @@
"# limitations under the License."
]
},
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "2X-XaMSVcLua"
+ },
+ "source": [
+ "# Decoding API"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hYEwGTeCXnnX"
+ },
+ "source": [
+ "
"
+ ]
+ },
{
"cell_type": "markdown",
"metadata": {
@@ -45,25 +76,14 @@
]
},
{
- "cell_type": "markdown",
+ "cell_type": "code",
+ "execution_count": null,
"metadata": {
- "id": "hYEwGTeCXnnX"
+ "id": "G4BhAu01HZcM"
},
+ "outputs": [],
"source": [
- "\u003ctable class=\"tfo-notebook-buttons\" align=\"left\"\u003e\n",
- " \u003ctd\u003e\n",
- " \u003ca target=\"_blank\" href=\"https://www.tensorflow.org/official_models/tutorials/decoding_api_in_tf_nlp.ipynb\"\u003e\u003cimg src=\"https://www.tensorflow.org/images/tf_logo_32px.png\" /\u003eView on TensorFlow.org\u003c/a\u003e\n",
- " \u003c/td\u003e\n",
- " \u003ctd\u003e\n",
- " \u003ca target=\"_blank\" href=\"https://colab.research.google.com/github/tensorflow/models/blob/master/official/colab/decoding_api_in_tf_nlp.ipynb\"\u003e\u003cimg src=\"https://www.tensorflow.org/images/colab_logo_32px.png\" /\u003eRun in Google Colab\u003c/a\u003e\n",
- " \u003c/td\u003e\n",
- " \u003ctd\u003e\n",
- " \u003ca target=\"_blank\" href=\"https://github.com/tensorflow/models/blob/master/official/colab/decoding_api_in_tf_nlp.ipynb\"\u003e\u003cimg src=\"https://www.tensorflow.org/images/GitHub-Mark-32px.png\" /\u003eView source on GitHub\u003c/a\u003e\n",
- " \u003c/td\u003e\n",
- " \u003ctd\u003e\n",
- " \u003ca href=\"https://storage.googleapis.com/tensorflow_docs/models/official/colab/decoding_api_in_tf_nlp.ipynb\"\u003e\u003cimg src=\"https://www.tensorflow.org/images/download_logo_32px.png\" /\u003eDownload notebook\u003c/a\u003e\n",
- " \u003c/td\u003e\n",
- "\u003c/table\u003e"
+ "!pip uninstall -y opencv-python"
]
},
{
@@ -74,7 +94,7 @@
},
"outputs": [],
"source": [
- "pip install tf-models-nightly"
+ "!pip install tf-models-official"
]
},
{
@@ -92,9 +112,20 @@
"\n",
"import tensorflow as tf\n",
"\n",
- "from official import nlp\n",
- "from official.nlp.modeling.ops import sampling_module\n",
- "from official.nlp.modeling.ops import beam_search"
+ "from tensorflow_models import nlp"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "T92ccAzlnGqh"
+ },
+ "outputs": [],
+ "source": [
+ "def length_norm(length, dtype):\n",
+ " \"\"\"Return length normalization factor.\"\"\"\n",
+ " return tf.pow(((5. + tf.cast(length, dtype)) / 6.), 0.0)"
]
},
{
@@ -103,13 +134,14 @@
"id": "0AWgyo-IQ5sP"
},
"source": [
- "# Decoding API\n",
+ "## Overview\n",
+ "\n",
"This API provides an interface to experiment with different decoding strategies used for auto-regressive models.\n",
"\n",
"1. The following sampling strategies are provided in sampling_module.py, which inherits from the base Decoding class:\n",
" * [top_p](https://arxiv.org/abs/1904.09751) : [github](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/ops/sampling_module.py#L65) \n",
"\n",
- " This implementation chooses most probable logits with cumulative probabilities upto top_p.\n",
+ " This implementation chooses the most probable logits with cumulative probabilities up to top_p.\n",
"\n",
" * [top_k](https://arxiv.org/pdf/1805.04833.pdf) : [github](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/ops/sampling_module.py#L48)\n",
"\n",
@@ -133,7 +165,7 @@
"## Initialize Sampling Module in TF-NLP.\n",
"\n",
"\n",
- "\u003e **symbols_to_logits_fn** : This is a closure implemented by the users of the API. The input to this closure will be \n",
+ "> **symbols_to_logits_fn** : This is a closure implemented by the users of the API. The input to this closure will be \n",
"```\n",
"Args:\n",
" 1] ids [batch_size, .. (index + 1 or 1 if padded_decode is True)],\n",
@@ -147,7 +179,7 @@
"Here is a [reference](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/ops/beam_search_test.py#L88) implementation for the above closure.\n",
"\n",
"\n",
- "\u003e **length_normalization_fn** : Closure for returning length normalization parameter.\n",
+ "> **length_normalization_fn** : Closure for returning length normalization parameter.\n",
"```\n",
"Args: \n",
" 1] length : scalar for decoded step index.\n",
@@ -159,21 +191,21 @@
" return tf.pow(((5. + tf.cast(length, dtype)) / 6.), 0.0)\n",
"```\n",
"\n",
- "\u003e **vocab_size** : Output vocabulary size.\n",
+ "> **vocab_size** : Output vocabulary size.\n",
"\n",
- "\u003e **max_decode_length** : Scalar for total number of decoding steps.\n",
+ "> **max_decode_length** : Scalar for total number of decoding steps.\n",
"\n",
- "\u003e **eos_id** : Decoding will stop if all output decoded ids in the batch have this ID.\n",
+ "> **eos_id** : Decoding will stop if all output decoded ids in the batch have this ID.\n",
"\n",
- "\u003e **padded_decode** : Set this to True if running on TPU. Tensors are padded to max_decoding_length if this is True.\n",
+ "> **padded_decode** : Set this to True if running on TPU. Tensors are padded to max_decoding_length if this is True.\n",
"\n",
- "\u003e **top_k** : top_k is enabled if this value is \u003e 1.\n",
+ "> **top_k** : top_k is enabled if this value is > 1.\n",
"\n",
- "\u003e **top_p** : top_p is enabled if this value is \u003e 0 and \u003c 1.0\n",
+ "> **top_p** : top_p is enabled if this value is > 0 and < 1.0\n",
"\n",
- "\u003e **sampling_temperature** : This is used to re-estimate the softmax output. Temperature skews the distribution towards high probability tokens and lowers the mass in tail distribution. Value has to be positive. Low temperature is equivalent to greedy and makes the distribution sharper, while high temperature makes it more flat.\n",
+ "> **sampling_temperature** : This is used to re-estimate the softmax output. Temperature skews the distribution towards high-probability tokens and lowers the mass in the tail distribution. Value has to be positive. Low temperature is equivalent to greedy and makes the distribution sharper, while high temperature makes it flatter.\n",
"\n",
- "\u003e **enable_greedy** : By default, this is true and greedy decoding is enabled.\n"
+ "> **enable_greedy** : By default, this is true and greedy decoding is enabled.\n"
]
},
{
@@ -182,7 +214,7 @@
"id": "lV1RRp6ihnGX"
},
"source": [
- "# Initialize the Model Hyper-parameters"
+ "## Initialize the Model Hyper-parameters"
]
},
{
@@ -193,44 +225,32 @@
},
"outputs": [],
"source": [
- "params = {}\n",
- "params['num_heads'] = 2\n",
- "params['num_layers'] = 2\n",
- "params['batch_size'] = 2\n",
- "params['n_dims'] = 256\n",
- "params['max_decode_length'] = 4"
+ "params = {\n",
+ " 'num_heads': 2,\n",
+ " 'num_layers': 2,\n",
+ " 'batch_size': 2,\n",
+ " 'n_dims': 256,\n",
+ " 'max_decode_length': 4}"
]
},
{
"cell_type": "markdown",
"metadata": {
- "id": "UGvmd0_dRFYI"
+ "id": "CYXkoplAij01"
},
"source": [
- "## What is a Cache?\n",
- "In auto-regressive architectures like Transformer based [Encoder-Decoder](https://arxiv.org/abs/1706.03762) models, \n",
- "Cache is used for fast sequential decoding.\n",
- "It is a nested dictionary storing pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention blocks) for every layer.\n",
- "\n",
- "```\n",
- "{\n",
- " 'layer_%d' % layer: {\n",
- " 'k': tf.zeros([params['batch_size'], params['max_decode_length'], params['num_heads'], params['n_dims']/params['num_heads']], dtype=tf.float32),\n",
- " 'v': tf.zeros([params['batch_size'], params['max_decode_length'], params['num_heads'], params['n_dims']/params['num_heads']], dtype=tf.float32)\n",
- " } for layer in range(params['num_layers']),\n",
- " 'model_specific_item' : Model specific tensor shape,\n",
- "}\n",
- "\n",
- "```"
+ "## Initialize cache. "
]
},
{
"cell_type": "markdown",
"metadata": {
- "id": "CYXkoplAij01"
+ "id": "UGvmd0_dRFYI"
},
"source": [
- "# Initialize cache. "
+ "In auto-regressive architectures like Transformer based [Encoder-Decoder](https://arxiv.org/abs/1706.03762) models, \n",
+ "Cache is used for fast sequential decoding.\n",
+ "It is a nested dictionary storing pre-computed hidden-states (key and values in the self-attention blocks and the cross-attention blocks) for every layer."
]
},
{
@@ -243,35 +263,15 @@
"source": [
"cache = {\n",
" 'layer_%d' % layer: {\n",
- " 'k': tf.zeros([params['batch_size'], params['max_decode_length'], params['num_heads'], params['n_dims']/params['num_heads']], dtype=tf.float32),\n",
- " 'v': tf.zeros([params['batch_size'], params['max_decode_length'], params['num_heads'], params['n_dims']/params['num_heads']], dtype=tf.float32)\n",
+ " 'k': tf.zeros(\n",
+ " shape=[params['batch_size'], params['max_decode_length'], params['num_heads'], params['n_dims'] // params['num_heads']],\n",
+ " dtype=tf.float32),\n",
+ " 'v': tf.zeros(\n",
+ " shape=[params['batch_size'], params['max_decode_length'], params['num_heads'], params['n_dims'] // params['num_heads']],\n",
+ " dtype=tf.float32)\n",
" } for layer in range(params['num_layers'])\n",
" }\n",
- "print(\"cache key shape for layer 1 :\", cache['layer_1']['k'].shape)"
- ]
- },
- {
- "cell_type": "markdown",
- "metadata": {
- "id": "nNY3Xn8SiblP"
- },
- "source": [
- "# Define closure for length normalization. **optional.**\n",
- "\n",
- "\n"
- ]
- },
- {
- "cell_type": "code",
- "execution_count": null,
- "metadata": {
- "id": "T92ccAzlnGqh"
- },
- "outputs": [],
- "source": [
- "def length_norm(length, dtype):\n",
- " \"\"\"Return length normalization factor.\"\"\"\n",
- " return tf.pow(((5. + tf.cast(length, dtype)) / 6.), 0.0)"
+ "print(\"cache value shape for layer 1 :\", cache['layer_1']['k'].shape)"
]
},
{
@@ -280,15 +280,14 @@
"id": "syl7I5nURPgW"
},
"source": [
- "# Create model_fn\n",
+ "### Create model_fn\n",
" In practice, this will be replaced by an actual model implementation such as [here](https://github.com/tensorflow/models/blob/master/official/nlp/transformer/transformer.py#L236)\n",
"```\n",
"Args:\n",
"i : Step that is being decoded.\n",
"Returns:\n",
" logit probabilities of size [batch_size, 1, vocab_size]\n",
- "```\n",
- "\n"
+ "```\n"
]
},
{
@@ -307,15 +306,6 @@
" return probabilities[:, i, :]"
]
},
- {
- "cell_type": "markdown",
- "metadata": {
- "id": "DBMUkaVmVZBg"
- },
- "source": [
- "# Initialize symbols_to_logits_fn\n"
- ]
- },
{
"cell_type": "code",
"execution_count": null,
@@ -339,7 +329,7 @@
"id": "R_tV3jyWVL47"
},
"source": [
- "# Greedy \n",
+ "## Greedy \n",
"Greedy decoding selects the token id with the highest probability as its next id: $id_t = argmax_{w}P(id | id_{1:t-1})$ at each timestep $t$. The following sketch shows greedy decoding. "
]
},
@@ -370,7 +360,7 @@
"id": "s4pTTsQXVz5O"
},
"source": [
- "# top_k sampling\n",
+ "## top_k sampling\n",
"In *Top-K* sampling, the *K* most likely next token ids are filtered and the probability mass is redistributed among only those *K* ids. "
]
},
@@ -404,7 +394,7 @@
"id": "Jp3G-eE_WI4Y"
},
"source": [
- "# top_p sampling\n",
+ "## top_p sampling\n",
"Instead of sampling only from the most likely *K* token ids, in *Top-p* sampling chooses from the smallest possible set of ids whose cumulative probability exceeds the probability *p*."
]
},
@@ -438,7 +428,7 @@
"id": "2hcuyJ2VWjDz"
},
"source": [
- "# Beam search decoding\n",
+ "## Beam search decoding\n",
"Beam search reduces the risk of missing hidden high probability token ids by keeping the most likely num_beams of hypotheses at each time step and eventually choosing the hypothesis that has the overall highest probability. "
]
},
diff --git a/docs/nlp/fine_tune_bert.ipynb b/docs/nlp/fine_tune_bert.ipynb
new file mode 100644
index 00000000000..24c58c35bcb
--- /dev/null
+++ b/docs/nlp/fine_tune_bert.ipynb
@@ -0,0 +1,1550 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "vXLA5InzXydn"
+ },
+ "source": [
+ "##### Copyright 2019 The TensorFlow Authors."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "RuRlpLL-X0R_"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n",
+ "# you may not use this file except in compliance with the License.\n",
+ "# You may obtain a copy of the License at\n",
+ "#\n",
+ "# https://www.apache.org/licenses/LICENSE-2.0\n",
+ "#\n",
+ "# Unless required by applicable law or agreed to in writing, software\n",
+ "# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
+ "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
+ "# See the License for the specific language governing permissions and\n",
+ "# limitations under the License."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "1mLJmVotXs64"
+ },
+ "source": [
+ "# Fine-tuning a BERT model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hYEwGTeCXnnX"
+ },
+ "source": [
+ "
"
]
},
{
@@ -103,7 +103,7 @@
},
"outputs": [],
"source": [
- "!pip install -q tf-models-official==2.4.0"
+ "!pip install tf-models-official"
]
},
{
@@ -126,8 +126,7 @@
"import numpy as np\n",
"import tensorflow as tf\n",
"\n",
- "from official.nlp import modeling\n",
- "from official.nlp.modeling import layers, losses, models, networks"
+ "from tensorflow_models import nlp"
]
},
{
@@ -151,9 +150,9 @@
"source": [
"### Build a `BertPretrainer` model wrapping `BertEncoder`\n",
"\n",
- "The [BertEncoder](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/networks/bert_encoder.py) implements the Transformer-based encoder as described in [BERT paper](https://arxiv.org/abs/1810.04805). It includes the embedding lookups and transformer layers, but not the masked language model or classification task networks.\n",
+ "The `nlp.networks.BertEncoder` class implements the Transformer-based encoder as described in [BERT paper](https://arxiv.org/abs/1810.04805). It includes the embedding lookups and transformer layers (`nlp.layers.TransformerEncoderBlock`), but not the masked language model or classification task networks.\n",
"\n",
- "The [BertPretrainer](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/models/bert_pretrainer.py) allows a user to pass in a transformer stack, and instantiates the masked language model and classification networks that are used to create the training objectives."
+ "The `nlp.models.BertPretrainer` class allows a user to pass in a transformer stack, and instantiates the masked language model and classification networks that are used to create the training objectives."
]
},
{
@@ -166,9 +165,10 @@
"source": [
"# Build a small transformer network.\n",
"vocab_size = 100\n",
- "sequence_length = 16\n",
- "network = modeling.networks.BertEncoder(\n",
- " vocab_size=vocab_size, num_layers=2, sequence_length=16)"
+ "network = nlp.networks.BertEncoder(\n",
+ " vocab_size=vocab_size, \n",
+ " # The number of TransformerEncoderBlock layers\n",
+ " num_layers=3)"
]
},
{
@@ -177,7 +177,7 @@
"id": "0NH5irV5KTMS"
},
"source": [
- "Inspecting the encoder, we see it contains few embedding layers, stacked `Transformer` layers and are connected to three input layers:\n",
+ "Inspecting the encoder, we see it contains few embedding layers, stacked `nlp.layers.TransformerEncoderBlock` layers and are connected to three input layers:\n",
"\n",
"`input_word_ids`, `input_type_ids` and `input_mask`.\n"
]
@@ -190,7 +190,7 @@
},
"outputs": [],
"source": [
- "tf.keras.utils.plot_model(network, show_shapes=True, dpi=48)"
+ "tf.keras.utils.plot_model(network, show_shapes=True, expand_nested=True, dpi=48)"
]
},
{
@@ -203,7 +203,7 @@
"source": [
"# Create a BERT pretrainer with the created network.\n",
"num_token_predictions = 8\n",
- "bert_pretrainer = modeling.models.BertPretrainer(\n",
+ "bert_pretrainer = nlp.models.BertPretrainer(\n",
" network, num_classes=2, num_token_predictions=num_token_predictions, output='predictions')"
]
},
@@ -213,7 +213,7 @@
"id": "d5h5HT7gNHx_"
},
"source": [
- "Inspecting the `bert_pretrainer`, we see it wraps the `encoder` with additional `MaskedLM` and `Classification` heads."
+ "Inspecting the `bert_pretrainer`, we see it wraps the `encoder` with additional `MaskedLM` and `nlp.layers.ClassificationHead` heads."
]
},
{
@@ -224,7 +224,7 @@
},
"outputs": [],
"source": [
- "tf.keras.utils.plot_model(bert_pretrainer, show_shapes=True, dpi=48)"
+ "tf.keras.utils.plot_model(bert_pretrainer, show_shapes=True, expand_nested=True, dpi=48)"
]
},
{
@@ -236,7 +236,9 @@
"outputs": [],
"source": [
"# We can feed some dummy data to get masked language model and sentence output.\n",
+ "sequence_length = 16\n",
"batch_size = 2\n",
+ "\n",
"word_id_data = np.random.randint(vocab_size, size=(batch_size, sequence_length))\n",
"mask_data = np.random.randint(2, size=(batch_size, sequence_length))\n",
"type_id_data = np.random.randint(2, size=(batch_size, sequence_length))\n",
@@ -246,8 +248,8 @@
" [word_id_data, mask_data, type_id_data, masked_lm_positions_data])\n",
"lm_output = outputs[\"masked_lm\"]\n",
"sentence_output = outputs[\"classification\"]\n",
- "print(lm_output)\n",
- "print(sentence_output)"
+ "print(f'lm_output: shape={lm_output.shape}, dtype={lm_output.dtype!r}')\n",
+ "print(f'sentence_output: shape={sentence_output.shape}, dtype={sentence_output.dtype!r}')"
]
},
{
@@ -272,14 +274,15 @@
"masked_lm_weights_data = np.random.randint(2, size=(batch_size, num_token_predictions))\n",
"next_sentence_labels_data = np.random.randint(2, size=(batch_size))\n",
"\n",
- "mlm_loss = modeling.losses.weighted_sparse_categorical_crossentropy_loss(\n",
+ "mlm_loss = nlp.losses.weighted_sparse_categorical_crossentropy_loss(\n",
" labels=masked_lm_ids_data,\n",
" predictions=lm_output,\n",
" weights=masked_lm_weights_data)\n",
- "sentence_loss = modeling.losses.weighted_sparse_categorical_crossentropy_loss(\n",
+ "sentence_loss = nlp.losses.weighted_sparse_categorical_crossentropy_loss(\n",
" labels=next_sentence_labels_data,\n",
" predictions=sentence_output)\n",
"loss = mlm_loss + sentence_loss\n",
+ "\n",
"print(loss)"
]
},
@@ -290,8 +293,7 @@
},
"source": [
"With the loss, you can optimize the model.\n",
- "After training, we can save the weights of TransformerEncoder for the downstream fine-tuning tasks. Please see [run_pretraining.py](https://github.com/tensorflow/models/blob/master/official/nlp/bert/run_pretraining.py) for the full example.\n",
- "\n"
+ "After training, we can save the weights of TransformerEncoder for the downstream fine-tuning tasks. Please see [run_pretraining.py](https://github.com/tensorflow/models/blob/master/official/legacy/bert/run_pretraining.py) for the full example.\n"
]
},
{
@@ -315,9 +317,9 @@
"source": [
"### Build a BertSpanLabeler wrapping BertEncoder\n",
"\n",
- "[BertSpanLabeler](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/models/bert_span_labeler.py) implements a simple single-span start-end predictor (that is, a model that predicts two values: a start token index and an end token index), suitable for SQuAD-style tasks.\n",
+ "The `nlp.models.BertSpanLabeler` class implements a simple single-span start-end predictor (that is, a model that predicts two values: a start token index and an end token index), suitable for SQuAD-style tasks.\n",
"\n",
- "Note that `BertSpanLabeler` wraps a `BertEncoder`, the weights of which can be restored from the above pretraining model.\n"
+ "Note that `nlp.models.BertSpanLabeler` wraps a `nlp.networks.BertEncoder`, the weights of which can be restored from the above pretraining model.\n"
]
},
{
@@ -328,11 +330,11 @@
},
"outputs": [],
"source": [
- "network = modeling.networks.BertEncoder(\n",
- " vocab_size=vocab_size, num_layers=2, sequence_length=sequence_length)\n",
+ "network = nlp.networks.BertEncoder(\n",
+ " vocab_size=vocab_size, num_layers=2)\n",
"\n",
"# Create a BERT trainer with the created network.\n",
- "bert_span_labeler = modeling.models.BertSpanLabeler(network)"
+ "bert_span_labeler = nlp.models.BertSpanLabeler(network)"
]
},
{
@@ -341,7 +343,7 @@
"id": "QpB9pgj4PpMg"
},
"source": [
- "Inspecting the `bert_span_labeler`, we see it wraps the encoder with additional `SpanLabeling` that outputs `start_position` and `end_postion`."
+ "Inspecting the `bert_span_labeler`, we see it wraps the encoder with additional `SpanLabeling` that outputs `start_position` and `end_position`."
]
},
{
@@ -352,7 +354,7 @@
},
"outputs": [],
"source": [
- "tf.keras.utils.plot_model(bert_span_labeler, show_shapes=True, dpi=48)"
+ "tf.keras.utils.plot_model(bert_span_labeler, show_shapes=True, expand_nested=True, dpi=48)"
]
},
{
@@ -370,8 +372,9 @@
"\n",
"# Feed the data to the model.\n",
"start_logits, end_logits = bert_span_labeler([word_id_data, mask_data, type_id_data])\n",
- "print(start_logits)\n",
- "print(end_logits)"
+ "\n",
+ "print(f'start_logits: shape={start_logits.shape}, dtype={start_logits.dtype!r}')\n",
+ "print(f'end_logits: shape={end_logits.shape}, dtype={end_logits.dtype!r}')"
]
},
{
@@ -410,7 +413,7 @@
"id": "Zdf03YtZmd_d"
},
"source": [
- "With the `loss`, you can optimize the model. Please see [run_squad.py](https://github.com/tensorflow/models/blob/master/official/nlp/bert/run_squad.py) for the full example."
+ "With the `loss`, you can optimize the model. Please see [run_squad.py](https://github.com/tensorflow/models/blob/master/official/legacy/bert/run_squad.py) for the full example."
]
},
{
@@ -432,7 +435,7 @@
"source": [
"### Build a BertClassifier model wrapping BertEncoder\n",
"\n",
- "[BertClassifier](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/models/bert_classifier.py) implements a [CLS] token classification model containing a single classification head."
+ "`nlp.models.BertClassifier` implements a [CLS] token classification model containing a single classification head."
]
},
{
@@ -443,12 +446,12 @@
},
"outputs": [],
"source": [
- "network = modeling.networks.BertEncoder(\n",
- " vocab_size=vocab_size, num_layers=2, sequence_length=sequence_length)\n",
+ "network = nlp.networks.BertEncoder(\n",
+ " vocab_size=vocab_size, num_layers=2)\n",
"\n",
"# Create a BERT trainer with the created network.\n",
"num_classes = 2\n",
- "bert_classifier = modeling.models.BertClassifier(\n",
+ "bert_classifier = nlp.models.BertClassifier(\n",
" network, num_classes=num_classes)"
]
},
@@ -469,7 +472,7 @@
},
"outputs": [],
"source": [
- "tf.keras.utils.plot_model(bert_classifier, show_shapes=True, dpi=48)"
+ "tf.keras.utils.plot_model(bert_classifier, show_shapes=True, expand_nested=True, dpi=48)"
]
},
{
@@ -487,7 +490,7 @@
"\n",
"# Feed the data to the model.\n",
"logits = bert_classifier([word_id_data, mask_data, type_id_data])\n",
- "print(logits)"
+ "print(f'logits: shape={logits.shape}, dtype={logits.dtype!r}')"
]
},
{
@@ -522,15 +525,13 @@
"id": "mzBqOylZo3og"
},
"source": [
- "With the `loss`, you can optimize the model. Please see [run_classifier.py](https://github.com/tensorflow/models/blob/master/official/nlp/bert/run_classifier.py) or the colab [fine_tuning_bert.ipynb](https://github.com/tensorflow/models/blob/master/official/colab/fine_tuning_bert.ipynb) for the full example."
+ "With the `loss`, you can optimize the model. Please see the [Fine tune_bert](https://www.tensorflow.org/text/tutorials/fine_tune_bert) notebook or the [model training documentation](https://github.com/tensorflow/models/blob/master/official/nlp/docs/train.md) for the full example."
]
}
],
"metadata": {
"colab": {
- "collapsed_sections": [],
- "name": "Introduction to the TensorFlow Models NLP library",
- "private_outputs": true,
+ "name": "nlp_modeling_library_intro.ipynb",
"provenance": [],
"toc_visible": true
},
diff --git a/docs/nlp/load_lm_ckpts.ipynb b/docs/nlp/load_lm_ckpts.ipynb
new file mode 100644
index 00000000000..e4e6083a6e4
--- /dev/null
+++ b/docs/nlp/load_lm_ckpts.ipynb
@@ -0,0 +1,692 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "30155835fc9f"
+ },
+ "source": [
+ "##### Copyright 2022 The TensorFlow Authors."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "906e07f6e562"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n",
+ "# you may not use this file except in compliance with the License.\n",
+ "# You may obtain a copy of the License at\n",
+ "#\n",
+ "# https://www.apache.org/licenses/LICENSE-2.0\n",
+ "#\n",
+ "# Unless required by applicable law or agreed to in writing, software\n",
+ "# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
+ "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
+ "# See the License for the specific language governing permissions and\n",
+ "# limitations under the License."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "5hrbPTziJK15"
+ },
+ "source": [
+ "# Load LM Checkpoints using Model Garden"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "-PYqCW1II75I"
+ },
+ "source": [
+ "
"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "c6J4IoNfN9jp"
+ },
+ "source": [
+ "This tutorial trains a [DeepLabV3](https://arxiv.org/pdf/1706.05587.pdf) with [Mobilenet V2](https://arxiv.org/abs/1801.04381) as backbone model from the [TensorFlow Model Garden](https://pypi.org/project/tf-models-official/) package (tensorflow-models).\n",
+ "\n",
+ "\n",
+ "[Model Garden](https://www.tensorflow.org/tfmodels) contains a collection of state-of-the-art models, implemented with TensorFlow's high-level APIs. The implementations demonstrate the best practices for modeling, letting users to take full advantage of TensorFlow for their research and product development.\n",
+ "\n",
+ "**Dataset**: [Oxford-IIIT Pets](https://www.tensorflow.org/datasets/catalog/oxford_iiit_pet)\n",
+ "\n",
+ "* The Oxford-IIIT pet dataset is a 37 category pet image dataset with roughly 200 images for each class. The images have large variations in scale, pose and lighting. All images have an associated ground truth annotation of breed.\n",
+ "\n",
+ "\n",
+ "**This tutorial demonstrates how to:**\n",
+ "\n",
+ "1. Use models from the TensorFlow Models package.\n",
+ "2. Train/Fine-tune a pre-built DeepLabV3 with mobilenet as backbone for Semantic Segmentation.\n",
+ "3. Export the trained/tuned DeepLabV3 model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "AlxYhP0XFnDn"
+ },
+ "source": [
+ "## Install necessary dependencies"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "pXWAySwgaWpN"
+ },
+ "outputs": [],
+ "source": [
+ "!pip install -U -q \"tf-models-official\""
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "uExUsXlgaPD6"
+ },
+ "source": [
+ "## Import required libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "mOmKZ3Vky5t9"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "import pprint\n",
+ "import numpy as np\n",
+ "import matplotlib.pyplot as plt\n",
+ "\n",
+ "from IPython import display"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nF8IHrXua_0b"
+ },
+ "outputs": [],
+ "source": [
+ "import tensorflow as tf\n",
+ "import tensorflow_datasets as tfds\n",
+ "\n",
+ "\n",
+ "import orbit\n",
+ "import tensorflow_models as tfm\n",
+ "from official.vision.data import tfrecord_lib\n",
+ "from official.vision.utils import summary_manager\n",
+ "from official.vision.serving import export_saved_model_lib\n",
+ "from official.vision.utils.object_detection import visualization_utils\n",
+ "\n",
+ "pp = pprint.PrettyPrinter(indent=4) # Set Pretty Print Indentation\n",
+ "print(tf.__version__) # Check the version of tensorflow used\n",
+ "\n",
+ "%matplotlib inline"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "gMs4l2dpaTd3"
+ },
+ "source": [
+ "## Custom dataset preparation for semantic segmentation\n",
+ "Models in Official repository (of model-garden) require models in a TFRecords dataformat.\n",
+ "\n",
+ "Please check [this resource](https://www.tensorflow.org/tutorials/load_data/tfrecord) to learn more about TFRecords data format.\n",
+ "\n",
+ "[Oxford_IIIT_pet:3](https://www.tensorflow.org/datasets/catalog/oxford_iiit_pet) dataset is taken from [Tensorflow Datasets](https://www.tensorflow.org/datasets/catalog/overview)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JpWK1Z-N3fHh"
+ },
+ "outputs": [],
+ "source": [
+ "(train_ds, val_ds, test_ds), info = tfds.data_source(\n",
+ " 'oxford_iiit_pet:4.*.*',\n",
+ " split=['train+test[:50%]', 'test[50%:80%]', 'test[80%:100%]'],\n",
+ " with_info=True)\n",
+ "info"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Sq6s11E1bMJB"
+ },
+ "source": [
+ "### Helper function to encode dataset as tfrecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "NlEf_C-DjDHG"
+ },
+ "outputs": [],
+ "source": [
+ "def process_record(record):\n",
+ " keys_to_features = {\n",
+ " 'image/encoded': tfrecord_lib.convert_to_feature(\n",
+ " tf.io.encode_jpeg(record['image']).numpy()),\n",
+ " 'image/height': tfrecord_lib.convert_to_feature(record['image'].shape[0]),\n",
+ " 'image/width': tfrecord_lib.convert_to_feature(record['image'].shape[1]),\n",
+ " 'image/segmentation/class/encoded':tfrecord_lib.convert_to_feature(\n",
+ " tf.io.encode_png(record['segmentation_mask'] - 1).numpy())\n",
+ " }\n",
+ " example = tf.train.Example(\n",
+ " features=tf.train.Features(feature=keys_to_features))\n",
+ " return example"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "FoapGlIebP9r"
+ },
+ "source": [
+ "### Write TFRecords to a folder"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "dDbMn5q551LQ"
+ },
+ "outputs": [],
+ "source": [
+ "output_dir = './oxford_iiit_pet_tfrecords/'\n",
+ "LOG_EVERY = 100\n",
+ "if not os.path.exists(output_dir):\n",
+ " os.mkdir(output_dir)\n",
+ "\n",
+ "def write_tfrecords(dataset, output_path, num_shards=1):\n",
+ " writers = [\n",
+ " tf.io.TFRecordWriter(\n",
+ " output_path + '-%05d-of-%05d.tfrecord' % (i, num_shards))\n",
+ " for i in range(num_shards)\n",
+ " ]\n",
+ " for idx, record in enumerate(dataset):\n",
+ " if idx % LOG_EVERY == 0:\n",
+ " print('On image %d', idx)\n",
+ " tf_example = process_record(record)\n",
+ " writers[idx % num_shards].write(tf_example.SerializeToString())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "QHDD-D7rbZj7"
+ },
+ "source": [
+ "### Write training data as TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qxJnVUfT0qBJ"
+ },
+ "outputs": [],
+ "source": [
+ "output_train_tfrecs = output_dir + 'train'\n",
+ "write_tfrecords(train_ds, output_train_tfrecs, num_shards=10)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ap55RwVFbhtu"
+ },
+ "source": [
+ "### Write validation data as TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Fgq-VxF79ucR"
+ },
+ "outputs": [],
+ "source": [
+ "output_val_tfrecs = output_dir + 'val'\n",
+ "write_tfrecords(val_ds, output_val_tfrecs, num_shards=5)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0AZoIEzRbxZu"
+ },
+ "source": [
+ "### Write test data as TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "QmwFmbP69t0U"
+ },
+ "outputs": [],
+ "source": [
+ "output_test_tfrecs = output_dir + 'test'\n",
+ "write_tfrecords(test_ds, output_test_tfrecs, num_shards=5)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "uEFzV-6ZfBZW"
+ },
+ "source": [
+ "## Configure the DeepLabV3 Mobilenet model for custom dataset"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "_LPEIvLsqSaG"
+ },
+ "outputs": [],
+ "source": [
+ "train_data_tfrecords = './oxford_iiit_pet_tfrecords/train*'\n",
+ "val_data_tfrecords = './oxford_iiit_pet_tfrecords/val*'\n",
+ "test_data_tfrecords = './oxford_iiit_pet_tfrecords/test*'\n",
+ "trained_model = './trained_model/'\n",
+ "export_dir = './exported_model/'"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "1ZlSiSRyb1Q6"
+ },
+ "source": [
+ "In Model Garden, the collections of parameters that define a model are called *configs*. Model Garden can create a config based on a known set of parameters via a [factory](https://en.wikipedia.org/wiki/Factory_method_pattern).\n",
+ "\n",
+ "\n",
+ "Use the `mnv2_deeplabv3_pascal` experiment configuration, as defined by `tfm.vision.configs.semantic_segmentation.mnv2_deeplabv3_pascal`.\n",
+ "\n",
+ "Please find all the registered experiements [here](https://www.tensorflow.org/api_docs/python/tfm/core/exp_factory/get_exp_config)\n",
+ "\n",
+ "The configuration defines an experiment to train a [DeepLabV3](https://arxiv.org/pdf/1706.05587.pdf) model with MobilenetV2 as backbone and [ASPP](https://arxiv.org/pdf/1606.00915v2.pdf) as decoder.\n",
+ "\n",
+ "There are also other alternative experiments available such as\n",
+ "\n",
+ "* `seg_deeplabv3_pascal`\n",
+ "* `seg_deeplabv3plus_pascal`\n",
+ "* `seg_resnetfpn_pascal`\n",
+ "* `mnv2_deeplabv3plus_cityscapes`\n",
+ "\n",
+ "and more. One can switch to them by changing the experiment name argument to the `get_exp_config` function.\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "bj5UZ6BkfJCX"
+ },
+ "outputs": [],
+ "source": [
+ "exp_config = tfm.core.exp_factory.get_exp_config('mnv2_deeplabv3_pascal')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "B8jyG-jGIdFs"
+ },
+ "outputs": [],
+ "source": [
+ "model_ckpt_path = './model_ckpt/'\n",
+ "if not os.path.exists(model_ckpt_path):\n",
+ " os.mkdir(model_ckpt_path)\n",
+ "\n",
+ "!gcloud storage cp gs://tf_model_garden/cloud/vision-2.0/deeplab/deeplabv3_mobilenetv2_coco/best_ckpt-63.data-00000-of-00001 './model_ckpt/'\n",
+ "!gcloud storage cp gs://tf_model_garden/cloud/vision-2.0/deeplab/deeplabv3_mobilenetv2_coco/best_ckpt-63.index './model_ckpt/'"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "QBYvVFZXhSGQ"
+ },
+ "source": [
+ "### Adjust the model and dataset configurations so that it works with custom dataset."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "o_Z_vWW9-5Sy"
+ },
+ "outputs": [],
+ "source": [
+ "num_classes = 3\n",
+ "WIDTH, HEIGHT = 128, 128\n",
+ "input_size = [HEIGHT, WIDTH, 3]\n",
+ "BATCH_SIZE = 16\n",
+ "\n",
+ "# Backbone Config\n",
+ "exp_config.task.init_checkpoint = model_ckpt_path + 'best_ckpt-63'\n",
+ "exp_config.task.freeze_backbone = True\n",
+ "\n",
+ "# Model Config\n",
+ "exp_config.task.model.num_classes = num_classes\n",
+ "exp_config.task.model.input_size = input_size\n",
+ "\n",
+ "# Training Data Config\n",
+ "exp_config.task.train_data.aug_scale_min = 1.0\n",
+ "exp_config.task.train_data.aug_scale_max = 1.0\n",
+ "exp_config.task.train_data.input_path = train_data_tfrecords\n",
+ "exp_config.task.train_data.global_batch_size = BATCH_SIZE\n",
+ "exp_config.task.train_data.dtype = 'float32'\n",
+ "exp_config.task.train_data.output_size = [HEIGHT, WIDTH]\n",
+ "exp_config.task.train_data.preserve_aspect_ratio = False\n",
+ "exp_config.task.train_data.seed = 21 # Reproducable Training Data\n",
+ "\n",
+ "# Validation Data Config\n",
+ "exp_config.task.validation_data.input_path = val_data_tfrecords\n",
+ "exp_config.task.validation_data.global_batch_size = BATCH_SIZE\n",
+ "exp_config.task.validation_data.dtype = 'float32'\n",
+ "exp_config.task.validation_data.output_size = [HEIGHT, WIDTH]\n",
+ "exp_config.task.validation_data.preserve_aspect_ratio = False\n",
+ "exp_config.task.validation_data.groundtruth_padded_size = [HEIGHT, WIDTH]\n",
+ "exp_config.task.validation_data.seed = 21 # Reproducable Validation Data\n",
+ "exp_config.task.validation_data.resize_eval_groundtruth = True # To enable validation loss"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0HDg5eKniMGJ"
+ },
+ "source": [
+ "### Adjust the trainer configuration."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WASJZ3gUH8ni"
+ },
+ "outputs": [],
+ "source": [
+ "logical_device_names = [logical_device.name\n",
+ " for logical_device in tf.config.list_logical_devices()]\n",
+ "\n",
+ "if 'GPU' in ''.join(logical_device_names):\n",
+ " print('This may be broken in Colab.')\n",
+ " device = 'GPU'\n",
+ "elif 'TPU' in ''.join(logical_device_names):\n",
+ " print('This may be broken in Colab.')\n",
+ " device = 'TPU'\n",
+ "else:\n",
+ " print('Running on CPU is slow, so only train for a few steps.')\n",
+ " device = 'CPU'\n",
+ "\n",
+ "\n",
+ "train_steps = 2000\n",
+ "exp_config.trainer.steps_per_loop = int(train_ds.__len__() // BATCH_SIZE)\n",
+ "\n",
+ "exp_config.trainer.summary_interval = exp_config.trainer.steps_per_loop # steps_per_loop = num_of_validation_examples // eval_batch_size\n",
+ "exp_config.trainer.checkpoint_interval = exp_config.trainer.steps_per_loop\n",
+ "exp_config.trainer.validation_interval = exp_config.trainer.steps_per_loop\n",
+ "exp_config.trainer.validation_steps = int(train_ds.__len__() // BATCH_SIZE) # validation_steps = num_of_validation_examples // eval_batch_size\n",
+ "exp_config.trainer.train_steps = train_steps\n",
+ "exp_config.trainer.optimizer_config.warmup.linear.warmup_steps = exp_config.trainer.steps_per_loop\n",
+ "exp_config.trainer.optimizer_config.learning_rate.type = 'cosine'\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.decay_steps = train_steps\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.initial_learning_rate = 0.1\n",
+ "exp_config.trainer.optimizer_config.warmup.linear.warmup_learning_rate = 0.05"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "R66w5MwkiO8Z"
+ },
+ "source": [
+ "### Print the modified configuration."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ckpjzrqfhoSn"
+ },
+ "outputs": [],
+ "source": [
+ "pp.pprint(exp_config.as_dict())\n",
+ "display.Javascript('google.colab.output.setIframeHeight(\"500px\");')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "FYwzdGKAiSOV"
+ },
+ "source": [
+ "### Set up the distribution strategy."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iwiOuYRRqdBi"
+ },
+ "outputs": [],
+ "source": [
+ "# Setting up the Strategy\n",
+ "if exp_config.runtime.mixed_precision_dtype == tf.float16:\n",
+ " tf.keras.mixed_precision.set_global_policy('mixed_float16')\n",
+ "\n",
+ "if 'GPU' in ''.join(logical_device_names):\n",
+ " distribution_strategy = tf.distribute.MirroredStrategy()\n",
+ "elif 'TPU' in ''.join(logical_device_names):\n",
+ " tf.tpu.experimental.initialize_tpu_system()\n",
+ " tpu = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='/device:TPU_SYSTEM:0')\n",
+ " distribution_strategy = tf.distribute.experimental.TPUStrategy(tpu)\n",
+ "else:\n",
+ " print('Warning: this will be really slow.')\n",
+ " distribution_strategy = tf.distribute.OneDeviceStrategy(logical_device_names[0])\n",
+ "\n",
+ "print(\"Done\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZLtk1GIIiVR2"
+ },
+ "source": [
+ "## Create the `Task` object (`tfm.core.base_task.Task`) from the `config_definitions.TaskConfig`.\n",
+ "\n",
+ "The `Task` object has all the methods necessary for building the dataset, building the model, and running training & evaluation. These methods are driven by `tfm.core.train_lib.run_experiment`."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ASTB5D2UISSr"
+ },
+ "outputs": [],
+ "source": [
+ "model_dir = './trained_model/'\n",
+ "\n",
+ "with distribution_strategy.scope():\n",
+ " task = tfm.core.task_factory.get_task(exp_config.task, logging_dir=model_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "YIQ26TW-ihzA"
+ },
+ "source": [
+ "## Visualize a batch of the data."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "412WyIUAIdCr"
+ },
+ "outputs": [],
+ "source": [
+ "for images, masks in task.build_inputs(exp_config.task.train_data).take(1):\n",
+ " print()\n",
+ " print(f'images.shape: {str(images.shape):16} images.dtype: {images.dtype!r}')\n",
+ " print(f'masks.shape: {str(masks[\"masks\"].shape):16} images.dtype: {masks[\"masks\"].dtype!r}')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "3GgluDVJixMd"
+ },
+ "source": [
+ "### Helper function for visualizing the results from TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "1kueMMfERvLx"
+ },
+ "outputs": [],
+ "source": [
+ "def plot_masks(display_list):\n",
+ " plt.figure(figsize=(15, 15))\n",
+ "\n",
+ " title = ['Input Image', 'True Mask', 'Predicted Mask']\n",
+ "\n",
+ " for i in range(len(display_list)):\n",
+ " plt.subplot(1, len(display_list), i+1)\n",
+ " plt.title(title[i])\n",
+ " plt.imshow(tf.keras.utils.array_to_img(display_list[i]))\n",
+ "\n",
+ "\n",
+ " plt.axis('off')\n",
+ " plt.show()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZCtt09G7i3dq"
+ },
+ "source": [
+ "### Visualization of training data\n",
+ "\n",
+ "Image Title represents what is depicted from the image.\n",
+ "\n",
+ "Same helper function can be used while visualizing predicted mask"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "YwUPf9V2B6SR"
+ },
+ "outputs": [],
+ "source": [
+ "num_examples = 3\n",
+ "\n",
+ "for images, masks in task.build_inputs(exp_config.task.train_data).take(num_examples):\n",
+ " plot_masks([images[0], masks['masks'][0]])"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "MeJ5w8KfjMmP"
+ },
+ "source": [
+ "## Train and evaluate\n",
+ "**IoU**: is defined as the area of the intersection divided by the area of the union of a predicted mask and ground truth mask."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ru3aHTCySHoH"
+ },
+ "outputs": [],
+ "source": [
+ "\n",
+ "model, eval_logs = tfm.core.train_lib.run_experiment(\n",
+ " distribution_strategy=distribution_strategy,\n",
+ " task=task,\n",
+ " mode='train_and_eval',\n",
+ " params=exp_config,\n",
+ " model_dir=model_dir,\n",
+ " eval_summary_manager=summary_manager.maybe_build_eval_summary_manager(\n",
+ " params=exp_config, model_dir=model_dir),\n",
+ " run_post_eval=True)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "vt3WmtxhjfGe"
+ },
+ "source": [
+ "## Load logs in tensorboard"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "A9rct_7BoJFb"
+ },
+ "outputs": [],
+ "source": [
+ "%load_ext tensorboard\n",
+ "%tensorboard --logdir './trained_model'"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "v6XaGoUuji7P"
+ },
+ "source": [
+ "## Saving and exporting the trained model\n",
+ "\n",
+ "The `keras.Model` object returned by `train_lib.run_experiment` expects the data to be normalized by the dataset loader using the same mean and variance statiscics in `preprocess_ops.normalize_image(image, offset=MEAN_RGB, scale=STDDEV_RGB)`. This export function handles those details, so you can pass `tf.uint8` images and get the correct results."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "GVsnyqzdnxHd"
+ },
+ "outputs": [],
+ "source": [
+ "export_saved_model_lib.export_inference_graph(\n",
+ " input_type='image_tensor',\n",
+ " batch_size=1,\n",
+ " input_image_size=[HEIGHT, WIDTH],\n",
+ " params=exp_config,\n",
+ " checkpoint_path=tf.train.latest_checkpoint(model_dir),\n",
+ " export_dir=export_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "nM1S-tjIjvAr"
+ },
+ "source": [
+ "## Importing SavedModel"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Nxi9pEwluUcT"
+ },
+ "outputs": [],
+ "source": [
+ "imported = tf.saved_model.load(export_dir)\n",
+ "model_fn = imported.signatures['serving_default']"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LbBfl6AUj_My"
+ },
+ "source": [
+ "## Visualize predictions"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qifGt_ohpFhn"
+ },
+ "outputs": [],
+ "source": [
+ "def create_mask(pred_mask):\n",
+ " pred_mask = tf.math.argmax(pred_mask, axis=-1)\n",
+ " pred_mask = pred_mask[..., tf.newaxis]\n",
+ " return pred_mask[0]\n",
+ "\n",
+ "\n",
+ "for record in test_ds.take(15):\n",
+ " image = tf.image.resize(record['image'], size=[HEIGHT, WIDTH])\n",
+ " image = tf.cast(image, dtype=tf.uint8)\n",
+ " mask = tf.image.resize(record['segmentation_mask'], size=[HEIGHT, WIDTH])\n",
+ " predicted_mask = model_fn(tf.expand_dims(record['image'], axis=0))\n",
+ " plot_masks([image, mask, create_mask(predicted_mask['logits'])])"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "name": "semantic_segmentation.ipynb",
+ "provenance": [],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/README-TPU.md b/official/README-TPU.md
index a6031c44f03..d1c0d4ee2b5 100644
--- a/official/README-TPU.md
+++ b/official/README-TPU.md
@@ -2,28 +2,28 @@
## Natural Language Processing
-* [bert](nlp/bert): A powerful pre-trained language representation model:
+* [bert](https://arxiv.org/abs/1810.04805): A powerful pre-trained language representation model:
BERT, which stands for Bidirectional Encoder Representations from
Transformers.
- [BERT FineTuning with Cloud TPU](https://cloud.google.com/tpu/docs/tutorials/bert-2.x) provides step by step instructions on Cloud TPU training. You can look [Bert MNLI Tensorboard.dev metrics](https://tensorboard.dev/experiment/LijZ1IrERxKALQfr76gndA) for MNLI fine tuning task.
+ [BERT FineTuning with Cloud TPU](https://cloud.google.com/ai-platform/training/docs/algorithms/bert-start) provides step by step instructions on Cloud TPU training. You can look [Bert MNLI Tensorboard.dev metrics](https://tensorboard.dev/experiment/LijZ1IrERxKALQfr76gndA) for MNLI fine tuning task.
* [transformer](nlp/transformer): A transformer model to translate the WMT
English to German dataset.
[Training transformer on Cloud TPU](https://cloud.google.com/tpu/docs/tutorials/transformer-2.x) for step by step instructions on Cloud TPU training.
## Computer Vision
-* [efficientnet](vision/image_classification): A family of convolutional
+* [efficientnet](https://github.com/tensorflow/models/blob/master/official/vision/modeling/backbones/efficientnet.py): A family of convolutional
neural networks that scale by balancing network depth, width, and
resolution and can be used to classify ImageNet's dataset of 1000 classes.
See [Tensorboard.dev training metrics](https://tensorboard.dev/experiment/KnaWjrq5TXGfv0NW5m7rpg/#scalars).
-* [mnist](vision/image_classification): A basic model to classify digits
+* [mnist](https://www.tensorflow.org/datasets/catalog/mnist): A basic model to classify digits
from the MNIST dataset. See [Running MNIST on Cloud TPU](https://cloud.google.com/tpu/docs/tutorials/mnist-2.x) tutorial and [Tensorboard.dev metrics](https://tensorboard.dev/experiment/mIah5lppTASvrHqWrdr6NA).
-* [mask-rcnn](vision/detection): An object detection and instance segmentation model. See [Tensorboard.dev training metrics](https://tensorboard.dev/experiment/LH7k0fMsRwqUAcE09o9kPA).
-* [resnet](vision/image_classification): A deep residual network that can
+* [mask-rcnn](https://www.tensorflow.org/api_docs/python/tfm/vision/configs/maskrcnn/MaskRCNN): An object detection and instance segmentation model. See [Tensorboard.dev training metrics](https://tensorboard.dev/experiment/LH7k0fMsRwqUAcE09o9kPA).
+* [resnet]((https://www.tensorflow.org/api_docs/python/tfm/vision/configs/image_classification/image_classification_imagenet)): A deep residual network that can
be used to classify ImageNet's dataset of 1000 classes.
See [Training ResNet on Cloud TPU](https://cloud.google.com/tpu/docs/tutorials/resnet-2.x) tutorial and [Tensorboard.dev metrics](https://tensorboard.dev/experiment/CxlDK8YMRrSpYEGtBRpOhg).
-* [retinanet](vision/detection): A fast and powerful object detector. See [Tensorboard.dev training metrics](https://tensorboard.dev/experiment/b8NRnWU3TqG6Rw0UxueU6Q).
-* [shapemask](vision/detection): An object detection and instance segmentation model using shape priors. See [Tensorboard.dev training metrics](https://tensorboard.dev/experiment/ZbXgVoc6Rf6mBRlPj0JpLA).
+* [retinanet](https://www.tensorflow.org/api_docs/python/tfm/vision/retinanet): A fast and powerful object detector. See [Tensorboard.dev training metrics](https://tensorboard.dev/experiment/b8NRnWU3TqG6Rw0UxueU6Q).
+* [shapemask](https://cloud.google.com/tpu/docs/tutorials/shapemask-2.x): An object detection and instance segmentation model using shape priors. See [Tensorboard.dev training metrics](https://tensorboard.dev/experiment/ZbXgVoc6Rf6mBRlPj0JpLA).
## Recommendation
* [dlrm](recommendation/ranking): [Deep Learning Recommendation Model for
diff --git a/official/README.md b/official/README.md
index 3502d3b5843..6b25fec103d 100644
--- a/official/README.md
+++ b/official/README.md
@@ -14,6 +14,9 @@ being easy to read.
These models are used as end-to-end tests, ensuring that the models run
with the same or improved speed and performance with each new TensorFlow build.
+The API documentation of the latest stable release is published to
+[tensorflow.org](https://www.tensorflow.org/api_docs/python/tfm).
+
## More models to come!
The team is actively developing new models.
@@ -55,6 +58,7 @@ In the near future, we will add:
|-------|-------------------|
| [RetinaNet](vision/MODEL_GARDEN.md) | [Focal Loss for Dense Object Detection](https://arxiv.org/abs/1708.02002) |
| [Mask R-CNN](vision/MODEL_GARDEN.md) | [Mask R-CNN](https://arxiv.org/abs/1703.06870) |
+| [YOLO](projects/yolo/README.md) | [YOLOv7: Trainable bag-of-freebies sets new state-of-the-art for real-time object detectors](https://arxiv.org/abs/2207.02696) |
| [SpineNet](vision/MODEL_GARDEN.md) | [SpineNet: Learning Scale-Permuted Backbone for Recognition and Localization](https://arxiv.org/abs/1912.05027) |
| [Cascade RCNN-RS and RetinaNet-RS](vision/MODEL_GARDEN.md) | [Simple Training Strategies and Model Scaling for Object Detection](https://arxiv.org/abs/2107.00057)|
@@ -66,12 +70,32 @@ In the near future, we will add:
### [Natural Language Processing](nlp/README.md)
+#### Pre-trained Language Model
+
+| Model | Reference (Paper) |
+|-------|-------------------|
+| [ALBERT](nlp/MODEL_GARDEN.md#available-model-configs) | [ALBERT: A Lite BERT for Self-supervised Learning of Language Representations](https://arxiv.org/abs/1909.11942) |
+| [BERT](nlp/MODEL_GARDEN.md#available-model-configs) | [BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding](https://arxiv.org/abs/1810.04805) |
+| [ELECTRA](nlp/tasks/electra_task.py) | [ELECTRA: Pre-training Text Encoders as Discriminators Rather Than Generators](https://arxiv.org/abs/2003.10555) |
+
+
+#### Neural Machine Translation
+
| Model | Reference (Paper) |
|-------|-------------------|
-| [ALBERT (A Lite BERT)](nlp/MODEL_GARDEN.md#available-model-configs) | [ALBERT: A Lite BERT for Self-supervised Learning of Language Representations](https://arxiv.org/abs/1909.11942) |
-| [BERT (Bidirectional Encoder Representations from Transformers)](nlp/MODEL_GARDEN.md#available-model-configs) | [BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding](https://arxiv.org/abs/1810.04805) |
-| [NHNet (News Headline generation model)](projects/nhnet) | [Generating Representative Headlines for News Stories](https://arxiv.org/abs/2001.09386) |
| [Transformer](nlp/MODEL_GARDEN.md#available-model-configs) | [Attention Is All You Need](https://arxiv.org/abs/1706.03762) |
+
+#### Natural Language Generation
+
+| Model | Reference (Paper) |
+|-------|-------------------|
+| [NHNet (News Headline generation model)](projects/nhnet) | [Generating Representative Headlines for News Stories](https://arxiv.org/abs/2001.09386) |
+
+
+#### Knowledge Distillation
+
+| Model | Reference (Paper) |
+|-------|-------------------|
| [MobileBERT](projects/mobilebert) | [MobileBERT: a Compact Task-Agnostic BERT for Resource-Limited Devices](https://arxiv.org/abs/2004.02984) |
### Recommendation
@@ -97,17 +121,21 @@ pip3 install tensorflow-text-nightly # when model uses `nlp` packages
* Incase of stable versions, targeting a specific release, Tensorflow-models
repository version numbers match with the target TensorFlow release. For
-example, [TensorFlow-models v2.5.0]
-(https://github.com/tensorflow/models/releases/tag/v2.5.0)
-is compatible with [TensorFlow v2.5.0]
-(https://github.com/tensorflow/tensorflow/releases/tag/v2.5.0).
-This is equivalent to the following.
+example, [TensorFlow-models v2.8.x](https://github.com/tensorflow/models/releases/tag/v2.8.0)
+is compatible with [TensorFlow v2.8.x](https://github.com/tensorflow/tensorflow/releases/tag/v2.8.0).
+This is equivalent to the following:
```shell
-pip3 install tf-models-official==2.5.0
-pip3 install tensorflow-text==2.5.0 # when model uses `nlp` packages
+pip3 install tf-models-official==2.8.0
+pip3 install tensorflow-text==2.8.0 # when models in uses `nlp` packages
```
+Starting from 2.9.x release, we release the modeling library as
+`tensorflow_models` package and users can `import tensorflow_models` directly to
+access to the exported symbols. If you are
+using the latest nightly version or github code directly, please follow the
+docstrings in the github.
+
Please follow the below steps before running models in this repository.
### Requirements
@@ -123,7 +151,26 @@ don't recommend earlier versions.
### Installation
Please check [here](https://github.com/tensorflow/models#Installation) for the
-instructions
+instructions.
+
+Available pypi packages:
+
+* [tf-models-official](https://pypi.org/project/tf-models-official/)
+* [tf-models-nightly](https://pypi.org/project/tf-models-nightly/): nightly
+release with the latest changes.
+* [tf-models-no-deps](https://pypi.org/project/tf-models-no-deps/): without
+`tensorflow` and `tensorflow-text` in the `install_requires` list.
+
+### Examples and Tutorials
+
+Get started with TensorFlow Model Garden by exploring the provided examples and tutorials:
+
+* [NLP](https://www.tensorflow.org/tfmodels/nlp)
+* [Image classification](https://www.tensorflow.org/tfmodels/vision/image_classification)
+* [Object detection](https://www.tensorflow.org/tfmodels/vision/object_detection)
+* [Semantic Segmentation](https://www.tensorflow.org/tfmodels/vision/semantic_segmentation)
+* [Instance Segmentation](https://www.tensorflow.org/tfmodels/vision/instance_segmentation)
+
## Contributions
diff --git a/official/__init__.py b/official/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/__init__.py
+++ b/official/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/common/__init__.py b/official/common/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/common/__init__.py
+++ b/official/common/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/common/dataset_fn.py b/official/common/dataset_fn.py
index 52138d717a0..06067b7190a 100644
--- a/official/common/dataset_fn.py
+++ b/official/common/dataset_fn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -31,7 +31,7 @@
import functools
from typing import Any, Callable, Type, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
PossibleDatasetType = Union[Type[tf.data.Dataset], Callable[[tf.Tensor], Any]]
diff --git a/official/common/distribute_utils.py b/official/common/distribute_utils.py
index 480bbf8c772..b5c1f17b3fa 100644
--- a/official/common/distribute_utils.py
+++ b/official/common/distribute_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import json
import os
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _collective_communication(all_reduce_alg):
@@ -96,7 +96,7 @@ def get_distribution_strategy(distribution_strategy="mirrored",
num_packs=1,
tpu_address=None,
**kwargs):
- """Return a DistributionStrategy for running the model.
+ """Return a Strategy for running the model.
Args:
distribution_strategy: a string specifying which distribution strategy to
@@ -119,7 +119,7 @@ def get_distribution_strategy(distribution_strategy="mirrored",
**kwargs: Additional kwargs for internal usages.
Returns:
- tf.distribute.DistibutionStrategy object.
+ tf.distribute.Strategy object.
Raises:
ValueError: if `distribution_strategy` is "off" or "one_device" and
`num_gpus` is larger than 1; or `num_gpus` is negative or if
@@ -148,8 +148,32 @@ def get_distribution_strategy(distribution_strategy="mirrored",
if distribution_strategy == "tpu":
# When tpu_address is an empty string, we communicate with local TPUs.
- cluster_resolver = tpu_initialize(tpu_address)
- return tf.distribute.TPUStrategy(cluster_resolver)
+ # Bug workaround that in v5p we need to explicitly specify the device
+ # assignment when using tpu strategy, adding device assignment to the
+ # strategy.
+ cluster_resolver = tf.distribute.cluster_resolver.TPUClusterResolver(
+ tpu=tpu_address
+ )
+ if tpu_address not in ("", "local"):
+ tf.config.experimental_connect_to_cluster(cluster_resolver)
+ topology = tf.tpu.experimental.initialize_tpu_system(cluster_resolver)
+
+ device_assignment = None
+ if hasattr(tf.tpu.experimental, "HardWareFeature"):
+ hardware_feature = tf.tpu.experimental.HardWareFeature(
+ cluster_resolver.tpu_hardware_feature
+ )
+ if (
+ hardware_feature.embedding_feature
+ == tf.tpu.experimental.HardwareFeature.EmbeddingFeature.V2 # pyrefly: ignore[missing-attribute]
+ ):
+ tpu_metadata = cluster_resolver.get_tpu_system_metadata()
+ device_assignment = tf.tpu.experimental.DeviceAssignment.build(
+ topology, num_replicas=tpu_metadata.num_cores
+ )
+
+ return tf.distribute.TPUStrategy(
+ cluster_resolver, experimental_device_assignment=device_assignment)
if distribution_strategy == "multi_worker_mirrored":
return tf.distribute.experimental.MultiWorkerMirroredStrategy(
diff --git a/official/common/distribute_utils_test.py b/official/common/distribute_utils_test.py
index f06ee3ba628..be1de09d1c8 100644
--- a/official/common/distribute_utils_test.py
+++ b/official/common/distribute_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for distribution util functions."""
import sys
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
@@ -43,14 +43,14 @@ def test_invalid_args(self):
def test_one_device_strategy_cpu(self):
ds = distribute_utils.get_distribution_strategy('one_device', num_gpus=0)
- self.assertEquals(ds.num_replicas_in_sync, 1)
- self.assertEquals(len(ds.extended.worker_devices), 1)
+ self.assertEqual(ds.num_replicas_in_sync, 1)
+ self.assertEqual(len(ds.extended.worker_devices), 1)
self.assertIn('CPU', ds.extended.worker_devices[0])
def test_one_device_strategy_gpu(self):
ds = distribute_utils.get_distribution_strategy('one_device', num_gpus=1)
- self.assertEquals(ds.num_replicas_in_sync, 1)
- self.assertEquals(len(ds.extended.worker_devices), 1)
+ self.assertEqual(ds.num_replicas_in_sync, 1)
+ self.assertEqual(len(ds.extended.worker_devices), 1)
self.assertIn('GPU', ds.extended.worker_devices[0])
def test_mirrored_strategy(self):
@@ -58,8 +58,8 @@ def test_mirrored_strategy(self):
_ = distribute_utils.get_distribution_strategy(num_gpus=0)
# 5 GPUs.
ds = distribute_utils.get_distribution_strategy(num_gpus=5)
- self.assertEquals(ds.num_replicas_in_sync, 5)
- self.assertEquals(len(ds.extended.worker_devices), 5)
+ self.assertEqual(ds.num_replicas_in_sync, 5)
+ self.assertEqual(len(ds.extended.worker_devices), 5)
for device in ds.extended.worker_devices:
self.assertIn('GPU', device)
@@ -105,12 +105,13 @@ def test_tpu_strategy(self):
ds, tf.distribute.TPUStrategy)
def test_invalid_strategy(self):
- with self.assertRaisesRegexp(
- ValueError,
- 'distribution_strategy must be a string but got: False. If'):
+ with self.assertRaisesRegex(
+ ValueError, 'distribution_strategy must be a string but got: False. If'
+ ):
distribute_utils.get_distribution_strategy(False)
- with self.assertRaisesRegexp(
- ValueError, 'distribution_strategy must be a string but got: 1'):
+ with self.assertRaisesRegex(
+ ValueError, 'distribution_strategy must be a string but got: 1'
+ ):
distribute_utils.get_distribution_strategy(1)
def test_get_strategy_scope(self):
diff --git a/official/common/flags.py b/official/common/flags.py
index 245769d8f40..549fedf1f10 100644
--- a/official/common/flags.py
+++ b/official/common/flags.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -45,7 +45,8 @@ def define_flags():
default=None,
enum_values=[
'train', 'eval', 'train_and_eval', 'continuous_eval',
- 'continuous_train_and_eval', 'train_and_validate'
+ 'continuous_train_and_eval', 'train_and_validate',
+ 'train_and_post_eval'
],
help='Mode to run: `train`, `eval`, `train_and_eval`, '
'`continuous_eval`, `continuous_train_and_eval` and '
@@ -108,3 +109,33 @@ def define_flags():
flags.DEFINE_string(
'tf_data_service', default=None, help='The tf.data service address')
+
+ flags.DEFINE_string(
+ 'tpu_platform', default=None, help='TPU platform type.')
+
+ flags.DEFINE_string(
+ 'tfhub_handle',
+ default=None,
+ help=(
+ 'TFHub handle for publishing the model to TFHub. When this flag '
+ 'is set your train.py should use xm_tfhub to publish to TFHub. '
+ 'If using TFleX, prefer tflex_output_uri + TfHubPusher component.'
+ ),
+ )
+
+ # To use TFleX components that consume a model, set this flag by configuring
+ # a TFleX tfx.borg.types.standard_artifacts.Model output artifact.
+ # (go/tflex-xm#connect-your-component-to-other-components)
+ # Meanwhile have your train.py save a serializable model to this destination:
+ # model.save(os.path.join(FLAGS.tflex_output_uri, 'Format-Servo'),
+ # save_format='tf')
+ # (https://www.tensorflow.org/guide/keras/serialization_and_saving)
+ flags.DEFINE_string(
+ 'tflex_output_uri',
+ default=None,
+ help=(
+ 'When running in TFleX, you can configure an XManagerLauncher output'
+ ' to set this flag'
+ ' (go/tflex-xm#connect-your-component-to-other-components)'
+ ),
+ )
diff --git a/official/common/registry_imports.py b/official/common/registry_imports.py
index eb9af692a4a..e76513e57da 100644
--- a/official/common/registry_imports.py
+++ b/official/common/registry_imports.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/common/streamz_counters.py b/official/common/streamz_counters.py
index 5def620ec75..1289e770ac8 100644
--- a/official/common/streamz_counters.py
+++ b/official/common/streamz_counters.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/core/__init__.py b/official/core/__init__.py
index 30d59d59e03..743aee16b9e 100644
--- a/official/core/__init__.py
+++ b/official/core/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,14 +13,19 @@
# limitations under the License.
"""Core is shared by both `nlp` and `vision`."""
+
from official.core import actions
from official.core import base_task
from official.core import base_trainer
from official.core import config_definitions
from official.core import exp_factory
from official.core import export_base
+from official.core import file_writers
from official.core import input_reader
from official.core import registry
+from official.core import savedmodel_checkpoint_manager
from official.core import task_factory
+from official.core import tf_example_builder
+from official.core import tf_example_feature_key
from official.core import train_lib
from official.core import train_utils
diff --git a/official/core/actions.py b/official/core/actions.py
index 4d51d309436..333c7ece049 100644
--- a/official/core/actions.py
+++ b/official/core/actions.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import gin
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_trainer
from official.core import config_definitions
@@ -39,16 +39,16 @@ class PruningAction:
def __init__(
self,
export_dir: str,
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
):
"""Initializes the instance.
Args:
export_dir: `str` for the export directory of the pruning summaries.
- model: `tf.keras.Model` model instance used for training. This will be
+ model: `tf_keras.Model` model instance used for training. This will be
used to assign a pruning step to each prunable weight.
- optimizer: `tf.keras.optimizers.Optimizer` optimizer instance used for
+ optimizer: `tf_keras.optimizers.Optimizer` optimizer instance used for
training. This will be used to find the current training steps.
"""
# TODO(b/221490190): Avoid local import when the bug is fixed.
@@ -84,14 +84,14 @@ class EMACheckpointing:
def __init__(self,
export_dir: str,
- optimizer: tf.keras.optimizers.Optimizer,
+ optimizer: tf_keras.optimizers.Optimizer,
checkpoint: tf.train.Checkpoint,
max_to_keep: int = 1):
"""Initializes the instance.
Args:
export_dir: `str` for the export directory of the EMA average weights.
- optimizer: `tf.keras.optimizers.Optimizer` optimizer instance used for
+ optimizer: `tf_keras.optimizers.Optimizer` optimizer instance used for
training. This will be used to swap the model weights with the average
weigths.
checkpoint: `tf.train.Checkpoint` instance.
@@ -165,7 +165,7 @@ def __call__(self, outputs: orbit.runner.Output):
% self.recover_counter)
return True
if (self.global_step >= self.recovery_begin_steps and
- loss_value > self.loss_upper_bound):
+ loss_value > self.loss_upper_bound): # pyrefly: ignore[unsupported-operation]
self.recover_counter += 1
if self.recover_counter > self.recovery_max_trials:
raise RuntimeError(
@@ -221,4 +221,16 @@ def get_train_actions(
action=RecoveryAction(checkpoint_manager),
)
train_actions.append(recover_action)
+
+ if (
+ params.trainer.preemption_on_demand_checkpoint
+ and trainer.strategy.cluster_resolver
+ ):
+ on_demand_checkpoint_action = orbit.actions.SaveCheckpointIfPreempted(
+ trainer.strategy.cluster_resolver,
+ checkpoint_manager,
+ trainer.global_step,
+ keep_running_after_save=True,
+ )
+ train_actions.append(on_demand_checkpoint_action)
return train_actions
diff --git a/official/core/actions_test.py b/official/core/actions_test.py
index a42360b66ad..c9fce8273eb 100644
--- a/official/core/actions_test.py
+++ b/official/core/actions_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl.testing import parameterized
import numpy as np
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -27,12 +27,12 @@
from official.modeling import optimization
-class TestModel(tf.keras.Model):
+class TestModel(tf_keras.Model):
def __init__(self):
super().__init__()
self.value = tf.Variable(0.0)
- self.dense = tf.keras.layers.Dense(2)
+ self.dense = tf_keras.layers.Dense(2)
_ = self.dense(tf.zeros((2, 2), tf.float32))
def call(self, x, training=None):
@@ -51,7 +51,7 @@ def test_ema_checkpointing(self, distribution):
with distribution.scope():
directory = self.create_tempdir()
model = TestModel()
- optimizer = tf.keras.optimizers.SGD()
+ optimizer = tf_keras.optimizers.SGD()
optimizer = optimization.ExponentialMovingAverage(
optimizer, trainable_weights_only=False)
@@ -81,7 +81,7 @@ def test_ema_checkpointing(self, distribution):
# Raises an error for a normal optimizer.
with self.assertRaisesRegex(ValueError,
'Optimizer has to be instance of.*'):
- _ = actions.EMACheckpointing(directory, tf.keras.optimizers.SGD(),
+ _ = actions.EMACheckpointing(directory, tf_keras.optimizers.SGD(),
checkpoint)
@combinations.generate(
@@ -121,7 +121,7 @@ def test_pruning(self, distribution):
with distribution.scope():
directory = self.get_temp_dir()
model = TestModel()
- optimizer = tf.keras.optimizers.SGD()
+ optimizer = tf_keras.optimizers.SGD()
pruning = actions.PruningAction(directory, model, optimizer)
pruning({})
diff --git a/official/core/base_task.py b/official/core/base_task.py
index 56b9bc4392e..110f6cb7f10 100644
--- a/official/core/base_task.py
+++ b/official/core/base_task.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from typing import Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions
from official.modeling import optimization
@@ -57,7 +57,9 @@ def __init__(self,
"""
super().__init__(name=name)
self._task_config = params
- self._logging_dir = logging_dir
+ self._logging_dir = (
+ logging_dir or ""
+ ) # Empty directory hints current working dir.
@property
def task_config(self):
@@ -108,7 +110,7 @@ def create_optimizer(cls, optimizer_config: OptimizationConfig,
return optimizer
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""[Optional] A callback function used as CheckpointManager's init_fn.
This function will be called when no checkpoint is found for the model.
@@ -139,7 +141,7 @@ def initialize(self, model: tf.keras.Model):
logging.info("Finished loading pretrained checkpoint from %s",
ckpt_dir_or_file)
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model: # pyrefly: ignore[bad-return]
"""[Optional] Creates model architecture.
Returns:
@@ -220,8 +222,8 @@ def process_compiled_metrics(self, compiled_metrics, labels, model_outputs):
def train_step(self,
inputs,
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
metrics=None):
"""Does forward and backward.
@@ -258,14 +260,14 @@ def train_step(self,
# For mixed precision, when a LossScaleOptimizer is used, the loss is
# scaled to avoid numeric underflow.
if isinstance(optimizer,
- tf.keras.mixed_precision.LossScaleOptimizer):
+ tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
if isinstance(optimizer,
- tf.keras.mixed_precision.LossScaleOptimizer):
+ tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
logs = {self.loss: loss}
@@ -277,7 +279,7 @@ def train_step(self,
logs.update({m.name: m.result() for m in model.metrics})
return logs
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
"""Validation step.
With distribution strategies, this method runs on devices.
@@ -306,7 +308,7 @@ def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
logs.update({m.name: m.result() for m in model.metrics})
return logs
- def inference_step(self, inputs, model: tf.keras.Model):
+ def inference_step(self, inputs, model: tf_keras.Model):
"""Performs the forward step.
With distribution strategies, this method runs on devices.
diff --git a/official/core/base_trainer.py b/official/core/base_trainer.py
index 4ac35b47da2..fb774700da2 100644
--- a/official/core/base_trainer.py
+++ b/official/core/base_trainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from absl import logging
import gin
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions
@@ -47,10 +47,19 @@ def init_async(self):
tf.distribute.experimental.coordinator.ClusterCoordinator(
self._strategy))
+ def coordinator_for_async(
+ self,
+ ) -> tf.distribute.experimental.coordinator.ClusterCoordinator:
+ if not self._coordinator:
+ raise ValueError(
+ "Coordinator uninitialized for async run. Call init_async() first."
+ )
+ return self._coordinator
+
def join(self):
"""Join all async steps. Only useful in aysnc training."""
if getattr(self, "_is_async", False):
- self._coordinator.join()
+ self.coordinator_for_async().join()
def create_train_loop_fn(self):
"""Creates a eval loop from the given step function and options."""
@@ -58,7 +67,9 @@ def create_train_loop_fn(self):
if getattr(self, "_is_async", False):
def _async_loop_fn(iterator, num_steps):
- self._coordinator.schedule(train_loop_fn, args=(iterator, num_steps))
+ self.coordinator_for_async().schedule(
+ train_loop_fn, args=(iterator, num_steps)
+ )
return _async_loop_fn
else:
@@ -76,7 +87,9 @@ def create_eval_loop_fn(self, has_state: bool):
def _async_loop_fn(iterator, num_steps, state=None, reduce_fn=None):
assert state is None
assert reduce_fn is None
- self._coordinator.schedule(eval_loop_fn, args=(iterator, num_steps))
+ self.coordinator_for_async().schedule(
+ eval_loop_fn, args=(iterator, num_steps)
+ )
return _async_loop_fn
else:
@@ -102,7 +115,9 @@ def distribute_dataset(self, dataset_or_fn, *args, **kwargs):
*args, **kwargs)
per_worker_dataset_fn = tf.function(per_worker_dataset_fn)
- return self._coordinator.create_per_worker_dataset(per_worker_dataset_fn)
+ return self.coordinator_for_async().create_per_worker_dataset(
+ per_worker_dataset_fn
+ )
else:
return orbit.utils.make_distributed_dataset(self._strategy, dataset_or_fn,
*args, **kwargs)
@@ -127,7 +142,7 @@ def __init__(
self,
config: ExperimentConfig,
task: base_task.Task,
- model: tf.keras.Model,
+ model: tf_keras.Model,
optimizer: tf.optimizers.Optimizer,
train: bool = True,
evaluate: bool = True,
@@ -141,7 +156,7 @@ def __init__(
Args:
config: An `ExperimentConfig` instance specifying experiment config.
task: A base_task.Task instance.
- model: The model instance, e.g. a tf.keras.Model instance.
+ model: The model instance, e.g. a tf_keras.Model instance.
optimizer: tf.optimizers.Optimizer instance.
train: bool, whether or not this trainer will be used for training.
default to True.
@@ -192,8 +207,8 @@ def __init__(
optimizer=self.optimizer,
**checkpoint_items)
- self._train_loss = tf.keras.metrics.Mean("training_loss", dtype=tf.float32)
- self._validation_loss = tf.keras.metrics.Mean(
+ self._train_loss = tf_keras.metrics.Mean("training_loss", dtype=tf.float32)
+ self._validation_loss = tf_keras.metrics.Mean(
"validation_loss", dtype=tf.float32)
model_metrics = model.metrics if hasattr(model, "metrics") else []
@@ -343,6 +358,28 @@ def train_loop_end(self):
logs["learning_rate"] = self.optimizer.learning_rate
return logs
+ def next_train_inputs(self, iterator):
+ """Fetches the next inputs for the model during train.
+
+ This method consumes the input iterator and returns the next inputs for the
+ model.
+
+ This method provides a way to control how to fetch the next model input, and
+ what data to send to the model.
+
+ Note: This function runs on the host side when accelerators are used.
+
+ Note: Depending on the training setup this may or may not run in eager mode.
+ In most cases it will be run in graph mode.
+
+ Args:
+ iterator: Dataset iterator to generate the next inputs from.
+
+ Returns:
+ The inputs to the model.
+ """
+ return next(iterator)
+
def train_step(self, iterator):
"""See base class."""
@@ -359,8 +396,8 @@ def step_fn(inputs):
self._train_loss.update_state(logs[self.task.loss])
self.global_step.assign_add(1)
- self.strategy.run(
- step_fn, args=(next(iterator),), options=self._runtime_options)
+ inputs = self.next_train_inputs(iterator)
+ self.strategy.run(step_fn, args=(inputs,), options=self._runtime_options)
def eval_begin(self):
"""Sets up metrics."""
@@ -371,6 +408,31 @@ def eval_begin(self):
optimization.ExponentialMovingAverage):
self.optimizer.swap_weights()
+ def next_eval_inputs(self, iterator):
+ """Fetches the next inputs for the model during eval.
+
+ This method consumes the input iterator and returns the next inputs for the
+ model and an additional logs dict. The output dict remains in the host (not
+ sent to GPUs/TPUs) and is merged with the model outputs which will be
+ processed later in `aggregate_logs`. This is useful for sending extra logs
+ downstream that are not compatible with the accelerators.
+
+ Note: This function runs on the host side when accelerators are used.
+
+ Note: Depending on the training setup this may or may not run in eager mode.
+ In most cases it will be run in graph mode.
+
+ Args:
+ iterator: Dataset iterator to generate the next inputs from.
+
+ Returns:
+ The inputs to the model, and an additional logs dictionnary. The logs
+ are not passed to the model, instead they are merged with model output
+ logs.
+ """
+ passthrough_logs = dict()
+ return next(iterator), passthrough_logs
+
def eval_step(self, iterator):
"""See base class."""
@@ -381,11 +443,26 @@ def step_fn(inputs):
self._validation_loss.update_state(logs[self.task.loss])
return logs
- distributed_outputs = self.strategy.run(step_fn, args=(next(iterator),))
- return tf.nest.map_structure(self.strategy.experimental_local_results,
- distributed_outputs)
-
- def eval_end(self, aggregated_logs=None):
+ inputs, passthrough_logs = self.next_eval_inputs(iterator)
+ distributed_outputs = self.strategy.run(step_fn, args=(inputs,))
+ logs = tf.nest.map_structure(
+ self.strategy.experimental_local_results, distributed_outputs
+ )
+
+ if set(logs.keys()) & set(passthrough_logs.keys()):
+ logging.warning(
+ (
+ "Conflict between the pasthrough log keys and the returned model"
+ " log keys. Found %r keys in the passthrough logs and %r keys in"
+ " the model logs. Model log keys takes precedence."
+ ),
+ logs.keys(),
+ passthrough_logs.keys(),
+ )
+
+ return {**passthrough_logs, **logs}
+
+ def eval_end(self, aggregated_logs=None): # pyrefly: ignore[bad-override]
"""Processes evaluation results."""
self.join()
logs = {}
diff --git a/official/core/base_trainer_test.py b/official/core/base_trainer_test.py
index b4aacb100a7..22f5a36c68c 100644
--- a/official/core/base_trainer_test.py
+++ b/official/core/base_trainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,7 +22,7 @@
from absl.testing import parameterized
import orbit
import portpicker
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -312,14 +312,16 @@ def test_configure_optimizer(self, mixed_precision_dtype, loss_scale):
trainer = self.create_test_trainer(config)
if mixed_precision_dtype == 'float16':
self.assertIsInstance(trainer.optimizer,
- tf.keras.mixed_precision.LossScaleOptimizer)
+ tf_keras.mixed_precision.LossScaleOptimizer)
if loss_scale in (None, 'dynamic'):
self.assertTrue(trainer.optimizer.dynamic)
else:
self.assertFalse(trainer.optimizer.dynamic)
self.assertEqual(trainer.optimizer.initial_scale, loss_scale)
else:
- self.assertIsInstance(trainer.optimizer, tf.keras.optimizers.SGD)
+ self.assertIsInstance(
+ trainer.optimizer,
+ (tf_keras.optimizers.SGD, tf_keras.optimizers.legacy.SGD))
metrics = trainer.train(tf.convert_to_tensor(5, dtype=tf.int32))
self.assertIn('training_loss', metrics)
@@ -347,7 +349,7 @@ def test_export_best_ckpt(self):
def test_model_with_compiled_loss(self):
task = mock_task.MockTask()
model = task.build_model()
- model.compile(loss=tf.keras.losses.CategoricalCrossentropy())
+ model.compile(loss=tf_keras.losses.CategoricalCrossentropy())
trainer = trainer_lib.Trainer(
self._config,
task,
diff --git a/official/core/config_definitions.py b/official/core/config_definitions.py
index c57bd21d8f8..685d20fd8fe 100644
--- a/official/core/config_definitions.py
+++ b/official/core/config_definitions.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -55,6 +55,8 @@ class DataConfig(base_config.Config):
interleaving files.
block_length: The number of consecutive elements to produce from each input
element before cycling to another input element when interleaving files.
+ ram_budget: RAM budget for tf.data service in GB. If None, tf.data will use
+ 50% of the available host RAM.
deterministic: A boolean controlling whether determinism should be enforced.
sharding: Whether sharding is used in the input pipeline.
enable_tf_data_service: A boolean indicating whether to enable tf.data
@@ -62,7 +64,7 @@ class DataConfig(base_config.Config):
tf_data_service_address: The URI of a tf.data service to offload
preprocessing onto during training. The URI should be in the format
"protocol://address", e.g. "grpc://tf-data-service:5050". It can be
- overridden by `FLAGS.tf_data_service` flag in the binary.
+ overridden by `FLAGS.tf_data_service` flag in the binary.
tf_data_service_job_name: The name of the tf.data service job. This argument
makes it possible for multiple datasets to share the same job. The default
behavior is that the dataset creates anonymous, exclusively owned jobs.
@@ -75,22 +77,49 @@ class DataConfig(base_config.Config):
decoding when loading dataset from TFDS. Use comma to separate multiple
features. The main use case is to skip the image/video decoding for better
performance.
+ enable_shared_tf_data_service_between_parallel_trainers: A bool. When set to
+ true, only a single tf.data service will be started, and it will be shared
+ between all the trainer run simultaneously, e.g. using vizier to tune
+ hyperparameters. This will save CPU and RAM resources compared to running
+ separate tf.data service for each trainer. Notice that if batch size is
+ different for different trainers, the field
+ apply_tf_data_service_before_batching also needs to be true so that only a
+ single tf.data service instance will be created. In this case, tf.data
+ service will be applied before batching operation. So make sure to not
+ apply any processing steps after batching (e.g. in postprocess_fn) since
+ they wouldn't be paralleled by tf.data service and may slow down your
+ tf.data pipeline. When using shared tf.data service, the tf.data dataset
+ must be infinite, and slow trainer may skip certain training examples.
+ More details about shared tf.data service can be found at:
+ https://www.tensorflow.org/api_docs/python/tf/data/experimental/service#sharing_tfdata_service_with_concurrent_trainers.
+ apply_tf_data_service_before_batching: A bool. If set to True, tf.data
+ service will be applied before batching operation. This is useful to make
+ sure only a single tf.data service instance is created when
+ enable_shared_tf_data_service_between_parallel_trainers is true and batch
+ size is changing between parallel trainers.
+ trainer_id: A string. The id of the trainer if there are multiple parallel
+ trainer running at the same time, e.g. in vizier tuning case. It will be
+ automatically set if this field is needed. Users does not need to set it
+ when creating experiment configs.
seed: An optional seed to use for deterministic shuffling/preprocessing.
prefetch_buffer_size: An int specifying the buffer size of prefetch
datasets. If None, the buffer size is autotuned. Specifying this is useful
in case autotuning uses up too much memory by making the buffer size too
high.
+ autotune_algorithm: If specified, use this algorithm for AUTOTUNE. See:
+ https://www.tensorflow.org/api_docs/python/tf/data/experimental/AutotuneAlgorithm
"""
input_path: Union[Sequence[str], str, base_config.Config] = ""
- tfds_name: str = ""
+ tfds_name: Union[str, base_config.Config] = ""
tfds_split: str = ""
global_batch_size: int = 0
- is_training: bool = None
+ is_training: Optional[bool] = None
drop_remainder: bool = True
shuffle_buffer_size: int = 100
cache: bool = False
cycle_length: Optional[int] = None
block_length: int = 1
+ ram_budget: Optional[int] = None
deterministic: Optional[bool] = None
sharding: bool = True
enable_tf_data_service: bool = False
@@ -99,8 +128,12 @@ class DataConfig(base_config.Config):
tfds_data_dir: str = ""
tfds_as_supervised: bool = False
tfds_skip_decoding_feature: str = ""
+ enable_shared_tf_data_service_between_parallel_trainers: bool = False
+ apply_tf_data_service_before_batching: bool = False
+ trainer_id: Optional[str] = None
seed: Optional[int] = None
prefetch_buffer_size: Optional[int] = None
+ autotune_algorithm: Optional[str] = None
@dataclasses.dataclass
@@ -166,6 +199,7 @@ class RuntimeConfig(base_config.Config):
# Global model parallelism configurations.
num_cores_per_replica: int = 1
default_shard_dim: int = -1
+ use_tpu_mp_strategy: bool = False
def model_parallelism(self):
return dict(
@@ -183,6 +217,7 @@ class TrainerConfig(base_config.Config):
train_tf_while_loop: whether or not to use tf while loop.
train_tf_function: whether or not to use tf_function for training loop.
eval_tf_function: whether or not to use tf_function for eval.
+ eval_tf_while_loop: whether or not to use tf while loop for eval.
allow_tpu_summary: Whether to allow summary happen inside the XLA program
runs on TPU through automatic outside compilation.
steps_per_loop: number of steps per loop to report training metrics. This
@@ -210,8 +245,12 @@ class TrainerConfig(base_config.Config):
trainer should compare the evaluation metrics. This can be either `higher`
(higher the better) or `lower` (lower the better).
validation_summary_subdir: A 'str', sub directory for saving eval summary.
+ preemption_on_demand_checkpoint: whether or not to save on-demand
+ checkpoints after a preemption.
"""
- optimizer_config: OptimizationConfig = OptimizationConfig()
+ optimizer_config: OptimizationConfig = dataclasses.field(
+ default_factory=OptimizationConfig
+ )
# Orbit settings.
train_tf_while_loop: bool = True
train_tf_function: bool = True
@@ -242,6 +281,8 @@ class TrainerConfig(base_config.Config):
# we will retore the model states.
recovery_max_trials: int = 0
validation_summary_subdir: str = "validation"
+ # Preemption on-demand checkpoint.
+ preemption_on_demand_checkpoint: bool = True # copybara-replace
@dataclasses.dataclass
@@ -249,19 +290,23 @@ class TaskConfig(base_config.Config):
"""Config passed to task."""
init_checkpoint: str = ""
model: Optional[base_config.Config] = None
- train_data: DataConfig = DataConfig()
- validation_data: DataConfig = DataConfig()
+ train_data: DataConfig = dataclasses.field(default_factory=DataConfig)
+ validation_data: DataConfig = dataclasses.field(default_factory=DataConfig)
name: Optional[str] = None
# Configs for differential privacy
# These configs are only effective if you use create_optimizer in
# tensorflow_models/official/core/base_task.py
+ # DEPRECATED b/264611883
differential_privacy_config: Optional[
dp_configs.DifferentialPrivacyConfig] = None
+ # Whether to show image summary. Useful to visualize model predictions. Only
+ # work for vision tasks.
+ allow_image_summary: bool = False
@dataclasses.dataclass
class ExperimentConfig(base_config.Config):
"""Top-level configuration."""
- task: TaskConfig = TaskConfig()
- trainer: TrainerConfig = TrainerConfig()
- runtime: RuntimeConfig = RuntimeConfig()
+ task: TaskConfig = dataclasses.field(default_factory=TaskConfig)
+ trainer: TrainerConfig = dataclasses.field(default_factory=TrainerConfig)
+ runtime: RuntimeConfig = dataclasses.field(default_factory=RuntimeConfig)
diff --git a/official/core/exp_factory.py b/official/core/exp_factory.py
index fef74449877..f713437b40c 100644
--- a/official/core/exp_factory.py
+++ b/official/core/exp_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/core/export_base.py b/official/core/export_base.py
index 0ee9163d725..1300cb07915 100644
--- a/official/core/export_base.py
+++ b/official/core/export_base.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from typing import Any, Callable, Dict, Mapping, List, Optional, Text, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
MAX_DIRECTORY_CREATION_ATTEMPTS = 10
@@ -30,7 +30,7 @@ class ExportModule(tf.Module, metaclass=abc.ABCMeta):
def __init__(self,
params,
- model: Union[tf.Module, tf.keras.Model],
+ model: Union[tf.Module, tf_keras.Model],
inference_step: Optional[Callable[..., Any]] = None,
*,
preprocessor: Optional[Callable[..., Any]] = None,
@@ -68,7 +68,7 @@ def _inference_step(inputs, model=None):
if inference_step is not None:
self.inference_step = functools.partial(inference_step, model=self.model)
else:
- if issubclass(type(model), tf.keras.Model):
+ if issubclass(type(model), tf_keras.Model):
# Default to self.model.call instead of self.model.__call__ to avoid
# keras tracing logic designed for training.
# Since most of Model Garden's call doesn't not have training kwargs
diff --git a/official/core/export_base_test.py b/official/core/export_base_test.py
index e08a4a420f9..38cf3ffe43b 100644
--- a/official/core/export_base_test.py
+++ b/official/core/export_base_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import os
from typing import Any, Dict, Mapping, Text
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import export_base
@@ -41,7 +41,7 @@ class ExportBaseTest(tf.test.TestCase):
def test_export_module(self):
tmp_dir = self.get_temp_dir()
- model = tf.keras.layers.Dense(2)
+ model = tf_keras.layers.Dense(2)
inputs = tf.ones([2, 4], tf.float32)
expected_output = model(inputs, training=False)
module = TestModule(params=None, model=model)
@@ -67,7 +67,7 @@ def test_export_module(self):
def test_custom_inference_step(self):
tmp_dir = self.get_temp_dir()
- model = tf.keras.layers.Dense(2)
+ model = tf_keras.layers.Dense(2)
inputs = tf.ones([2, 4], tf.float32)
def _inference_step(inputs, model):
diff --git a/official/core/file_writers.py b/official/core/file_writers.py
new file mode 100644
index 00000000000..fa776714530
--- /dev/null
+++ b/official/core/file_writers.py
@@ -0,0 +1,80 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""File writer functions for dataset preparation, infra validation, and unit tests."""
+
+import io
+from typing import Optional, Sequence, Union
+
+import tensorflow as tf, tf_keras
+
+
+def write_small_dataset(examples: Sequence[Union[tf.train.Example,
+ tf.train.SequenceExample]],
+ output_path: str,
+ file_type: str = 'tfrecord') -> None:
+ """Writes `examples` to a file at `output_path` with type `file_type`.
+
+ CAVEAT: This function is not recommended for writing large datasets, since it
+ will loop through `examples` and perform write operation sequentially.
+
+ Args:
+ examples: List of tf.train.Example or tf.train.SequenceExample.
+ output_path: Output path for the dataset.
+ file_type: A string indicating the file format, could be: 'tfrecord',
+ 'tfrecords', 'tfrecord_compressed', 'tfrecords_gzip', 'riegeli'. The
+ string is case insensitive.
+ """
+ file_type = file_type.lower()
+
+ if file_type == 'tfrecord' or file_type == 'tfrecords':
+ _write_tfrecord(examples, output_path)
+ elif file_type == 'tfrecord_compressed' or file_type == 'tfrecords_gzip':
+ _write_tfrecord(examples, output_path,
+ tf.io.TFRecordOptions(compression_type='GZIP'))
+ elif file_type == 'riegeli':
+ _write_riegeli(examples, output_path)
+ else:
+ raise ValueError(f'Unknown file_type: {file_type}')
+
+
+def _write_tfrecord(examples: Sequence[Union[tf.train.Example,
+ tf.train.SequenceExample]],
+ output_path: str,
+ options: Optional[tf.io.TFRecordOptions] = None) -> None:
+ """Writes `examples` to a TFRecord file at `output_path`.
+
+ Args:
+ examples: A list of tf.train.Example.
+ output_path: Output path for the dataset.
+ options: Options used for manipulating TFRecord files.
+ """
+ with tf.io.TFRecordWriter(output_path, options) as writer:
+ for example in examples:
+ writer.write(example.SerializeToString())
+
+
+def _write_riegeli(examples: Sequence[Union[tf.train.Example,
+ tf.train.SequenceExample]],
+ output_path: str) -> None:
+ """Writes `examples` to a Riegeli file at `output_path`.
+
+ Args:
+ examples: A list of tf.train.Example.
+ output_path: Output path for the dataset.
+ """
+ with io.FileIO(output_path, 'wb') as fileio:
+ import riegeli # pylint: disable=g-import-not-at-top
+ with riegeli.RecordWriter(fileio) as writer:
+ writer.write_messages(examples)
diff --git a/official/core/file_writers_test.py b/official/core/file_writers_test.py
new file mode 100644
index 00000000000..8230e91a76b
--- /dev/null
+++ b/official/core/file_writers_test.py
@@ -0,0 +1,53 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for file_writers."""
+
+import os
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.core import file_writers
+from official.core import tf_example_builder
+
+
+class FileWritersTest(tf.test.TestCase, parameterized.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_bytes_feature('foo', 'Hello World!')
+ self._example = example_builder.example
+
+ @parameterized.parameters('tfrecord', 'TFRecord', 'tfrecords',
+ 'tfrecord_compressed', 'TFRecord_Compressed',
+ 'tfrecords_gzip')
+ def test_write_small_dataset_success(self, file_type):
+ temp_dir = self.create_tempdir()
+ temp_dataset_file = os.path.join(temp_dir.full_path, 'train')
+ file_writers.write_small_dataset([self._example], temp_dataset_file,
+ file_type)
+ self.assertTrue(os.path.exists(temp_dataset_file))
+
+ def test_write_small_dataset_unrecognized_format(self):
+ file_type = 'bar'
+ temp_dir = self.create_tempdir()
+ temp_dataset_file = os.path.join(temp_dir.full_path, 'train')
+ with self.assertRaises(ValueError):
+ file_writers.write_small_dataset([self._example], temp_dataset_file,
+ file_type)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/core/input_reader.py b/official/core/input_reader.py
index 72285ea3b05..20ab91e7800 100644
--- a/official/core/input_reader.py
+++ b/official/core/input_reader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,11 +13,12 @@
# limitations under the License.
"""A common dataset reader."""
+import dataclasses
import random
from typing import Any, Callable, Dict, List, Optional, Sequence, Text, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.core import config_definitions as cfg
@@ -147,7 +148,8 @@ def _shard_files_then_read(matched_files: List[str],
return dataset
-def _read_tfds(tfds_builder: tfds.core.DatasetBuilder,
+def _read_tfds(tfds_name: Text,
+ tfds_data_dir: Text,
tfds_split: Text,
tfds_skip_decoding_feature: Text,
tfds_as_supervised: bool,
@@ -158,49 +160,54 @@ def _read_tfds(tfds_builder: tfds.core.DatasetBuilder,
cycle_length: Optional[int] = None,
block_length: Optional[int] = None) -> tf.data.Dataset:
"""Reads a dataset from tfds."""
- # No op if exist.
- tfds_builder.download_and_prepare()
+ repeat_filenames = is_training and not cache
+ read_config = tfds.ReadConfig(
+ interleave_cycle_length=cycle_length,
+ interleave_block_length=block_length,
+ input_context=input_context,
+ shuffle_seed=seed, # pyrefly: ignore[bad-argument-type]
+ repeat_filenames=repeat_filenames,
+ # Only assert cardinality when we have a finite dataset.
+ assert_cardinality=not repeat_filenames,
+ skip_prefetch=True)
+
decoders = {}
if tfds_skip_decoding_feature:
for skip_feature in tfds_skip_decoding_feature.split(','):
decoders[skip_feature.strip()] = tfds.decode.SkipDecoding()
- if tfds_builder.info.splits:
- num_shards = len(tfds_builder.info.splits[tfds_split].file_instructions)
- else:
- # The tfds mock path often does not provide splits.
- num_shards = 1
- if input_context and num_shards < input_context.num_input_pipelines:
- # The number of files in the dataset split is smaller than the number of
- # input pipelines. We read the entire dataset first and then shard in the
- # host memory.
- read_config = tfds.ReadConfig(
- interleave_cycle_length=cycle_length,
- interleave_block_length=block_length,
- input_context=None,
- shuffle_seed=seed)
- dataset = tfds_builder.as_dataset(
- split=tfds_split,
- shuffle_files=is_training,
- as_supervised=tfds_as_supervised,
- decoders=decoders,
- read_config=read_config)
- dataset = dataset.shard(input_context.num_input_pipelines,
- input_context.input_pipeline_id)
- else:
- read_config = tfds.ReadConfig(
- interleave_cycle_length=cycle_length,
- interleave_block_length=block_length,
- input_context=input_context,
- shuffle_seed=seed)
- dataset = tfds_builder.as_dataset(
- split=tfds_split,
- shuffle_files=is_training,
- as_supervised=tfds_as_supervised,
- decoders=decoders,
- read_config=read_config)
- if is_training and not cache:
- dataset = dataset.repeat()
+ if tfds_name.startswith('mldataset.'):
+ dataset = tfds.load(name=tfds_name,
+ split=tfds_split,
+ as_supervised=tfds_as_supervised,
+ decoders=decoders if decoders else None,
+ read_config=read_config)
+ else:
+ builder = tfds.builder(tfds_name, data_dir=tfds_data_dir)
+ if builder.info.splits:
+ num_shards = len(builder.info.splits[tfds_split].file_instructions)
+ else:
+ # The tfds mock path often does not provide splits.
+ num_shards = 1
+ load_kwargs = dict(
+ name=tfds_name, download=True, split=tfds_split,
+ shuffle_files=is_training, as_supervised=tfds_as_supervised,
+ decoders=decoders if decoders else None)
+ if tfds_data_dir:
+ load_kwargs.update({'data_dir': tfds_data_dir})
+
+ if input_context and num_shards < input_context.num_input_pipelines:
+ # The number of files in the dataset split is smaller than the number of
+ # input pipelines. We read the entire dataset first and then shard in the
+ # host memory.
+ read_config = dataclasses.replace(read_config, input_context=None)
+ load_kwargs.update({'read_config': read_config}) # pyrefly: ignore[no-matching-overload]
+ dataset = tfds.load(**load_kwargs)
+ dataset = dataset.shard(input_context.num_input_pipelines,
+ input_context.input_pipeline_id)
+ else:
+ load_kwargs.update({'read_config': read_config}) # pyrefly: ignore[no-matching-overload]
+ dataset = tfds.load(**load_kwargs)
return dataset
@@ -211,17 +218,23 @@ class InputReader:
# instances.
static_randnum = _get_random_integer()
- def __init__(self,
- params: cfg.DataConfig,
- dataset_fn=tf.data.TFRecordDataset,
- decoder_fn: Optional[Callable[..., Any]] = None,
- combine_fn: Optional[Callable[..., Any]] = None,
- sample_fn: Optional[Callable[..., Any]] = None,
- parser_fn: Optional[Callable[..., Any]] = None,
- transform_and_batch_fn: Optional[Callable[
- [tf.data.Dataset, Optional[tf.distribute.InputContext]],
- tf.data.Dataset]] = None,
- postprocess_fn: Optional[Callable[..., Any]] = None):
+ def __init__(
+ self,
+ params: cfg.DataConfig,
+ dataset_fn=tf.data.TFRecordDataset,
+ decoder_fn: Optional[Callable[..., Any]] = None,
+ combine_fn: Optional[Callable[..., Any]] = None,
+ sample_fn: Optional[Callable[..., Any]] = None,
+ parser_fn: Optional[Callable[..., Any]] = None,
+ filter_fn: Optional[Callable[..., tf.Tensor]] = None,
+ transform_and_batch_fn: Optional[
+ Callable[
+ [tf.data.Dataset, Optional[tf.distribute.InputContext]],
+ tf.data.Dataset,
+ ]
+ ] = None,
+ postprocess_fn: Optional[Callable[..., Any]] = None,
+ ):
"""Initializes an InputReader instance.
Args:
@@ -239,6 +252,8 @@ def __init__(self,
parser_fn: An optional `callable` that takes the decoded raw tensors dict
and parse them into a dictionary of tensors that can be consumed by the
model. It will be executed after decoder_fn.
+ filter_fn: An optional `callable` mapping a dataset element to a boolean.
+ It will be executed after parser_fn.
transform_and_batch_fn: An optional `callable` that takes a
`tf.data.Dataset` object and an optional `tf.distribute.InputContext` as
input, and returns a `tf.data.Dataset` object. It will be executed after
@@ -253,12 +268,14 @@ def __init__(self,
'specified, but got %s and %s.' %
(params.input_path, params.tfds_name))
- if isinstance(params.input_path,
- cfg.base_config.Config) and combine_fn is None:
+ if (isinstance(params.input_path, cfg.base_config.Config) or
+ isinstance(params.tfds_name, cfg.base_config.Config)
+ ) and combine_fn is None:
raise ValueError(
- 'A `combine_fn` is required if the `input_path` is a dictionary.')
+ 'A combine_fn is required if `input_path` or `tfds_name` is a dict.')
- self._tfds_builder = None
+ self._tfds_name = params.tfds_name
+ self._tfds_data_dir = params.tfds_data_dir
self._matched_files = None
if not params.input_path:
# Read dataset from TFDS.
@@ -266,8 +283,6 @@ def __init__(self,
raise ValueError(
'`tfds_name` is %s, but `tfds_split` is not specified.' %
params.tfds_name)
- self._tfds_builder = tfds.builder(
- params.tfds_name, data_dir=params.tfds_data_dir)
else:
self._matched_files = self.get_files(params.input_path)
@@ -291,9 +306,12 @@ def __init__(self,
self._parser_fn = parser_fn
self._transform_and_batch_fn = transform_and_batch_fn
self._postprocess_fn = postprocess_fn
+ self._filter_fn = filter_fn
self._seed = params.seed
- self._prefetch_buffer_size = (params.prefetch_buffer_size or
- tf.data.experimental.AUTOTUNE)
+ self._prefetch_buffer_size = (
+ params.prefetch_buffer_size or tf.data.experimental.AUTOTUNE)
+ self._autotune_algorithm = params.autotune_algorithm
+ self._ram_budget = params.ram_budget
# When tf.data service is enabled, each data service worker should get
# different random seeds. Thus, we set `seed` to None.
@@ -306,6 +324,11 @@ def __init__(self,
self._enable_tf_data_service = (
params.enable_tf_data_service and params.tf_data_service_address)
self._tf_data_service_address = params.tf_data_service_address
+ self._enable_shared_tf_data_service_between_parallel_trainers = (
+ params.enable_shared_tf_data_service_between_parallel_trainers)
+ self._apply_tf_data_service_before_batching = (
+ params.apply_tf_data_service_before_batching)
+ self._trainer_id = params.trainer_id
if self._enable_tf_data_service:
# Add a random seed as the tf.data service job name suffix, so tf.data
# service doesn't reuse the previous state if TPU worker gets preempted.
@@ -322,15 +345,15 @@ def __init__(self,
f'{self.static_randnum}')
self._enable_round_robin_tf_data_service = params.get(
'enable_round_robin_tf_data_service', False)
-
- @property
- def tfds_info(self) -> tfds.core.DatasetInfo:
- """Returns TFDS dataset info, if available."""
- if self._tfds_builder:
- return self._tfds_builder.info
- else:
- raise ValueError('tfds_info is not available, because the dataset '
- 'is not loaded from tfds.')
+ if self._enable_shared_tf_data_service_between_parallel_trainers:
+ # When shared tf.data service is enabled, only a single tf.data service
+ # instance should be created and shared between parallel trainers. If
+ # the global batch size is different across trainers,
+ # params.apply_tf_data_service_before_batching should be set to true
+ # because tf.data service with different batch sizes will be considered
+ # separate tf.data service instances.
+ self._tf_data_service_job_name = (
+ f'{params.tf_data_service_job_name}_{self.static_randnum}')
def get_files(self, input_path):
"""Gets matched files. Can be overridden by subclasses."""
@@ -351,17 +374,21 @@ def _read_data_source(
matched_files: Union[Dict[str, List[str]], List[str]],
dataset_fn,
input_context: Optional[tf.distribute.InputContext] = None,
- tfds_builder: Optional[tfds.core.DatasetBuilder] = None):
+ ):
"""Reads the data source (files/tfds) to a dataset."""
def _files_to_dataset(files: List[str]) -> tf.data.Dataset:
if len(files) > 1:
if input_context and (len(files) < input_context.num_input_pipelines):
logging.warn(
- 'The number of files %d is less than the number of input pipelines '
- '%d. We will send all input files to every worker. '
- 'Please consider sharding your data into more files.', len(files),
- input_context.num_input_pipelines)
+ (
+ 'The number of files %d is less than the number of input '
+ 'pipelines %d. We will send all input files to every worker. '
+ 'Please consider sharding your data into more files.'
+ ),
+ len(files),
+ input_context.num_input_pipelines,
+ )
return _read_files_then_shard(
files,
dataset_fn,
@@ -391,18 +418,35 @@ def _files_to_dataset(files: List[str]) -> tf.data.Dataset:
raise ValueError('It is unexpected that `tfds_builder` is None and '
'there is also no `files`.')
- if tfds_builder:
- dataset = _read_tfds(
- tfds_builder=self._tfds_builder,
- tfds_split=self._tfds_split,
- tfds_skip_decoding_feature=self._tfds_skip_decoding_feature,
- tfds_as_supervised=self._tfds_as_supervised,
- input_context=input_context,
- seed=self._seed,
- is_training=self._is_training,
- cache=self._cache,
- cycle_length=self._cycle_length,
- block_length=self._block_length)
+ if self._tfds_name:
+ if isinstance(self._tfds_name, cfg.base_config.Config):
+ dataset = {}
+ for k, tfds_name in self._tfds_name.as_dict().items():
+ dataset[k] = _read_tfds(
+ tfds_name=tfds_name,
+ tfds_data_dir=self._tfds_data_dir,
+ tfds_split=self._tfds_split,
+ tfds_skip_decoding_feature=self._tfds_skip_decoding_feature,
+ tfds_as_supervised=self._tfds_as_supervised,
+ input_context=input_context,
+ seed=self._seed,
+ is_training=self._is_training,
+ cache=self._cache,
+ cycle_length=self._cycle_length,
+ block_length=self._block_length)
+ else:
+ dataset = _read_tfds(
+ tfds_name=self._tfds_name,
+ tfds_data_dir=self._tfds_data_dir,
+ tfds_split=self._tfds_split,
+ tfds_skip_decoding_feature=self._tfds_skip_decoding_feature,
+ tfds_as_supervised=self._tfds_as_supervised,
+ input_context=input_context,
+ seed=self._seed,
+ is_training=self._is_training,
+ cache=self._cache,
+ cycle_length=self._cycle_length,
+ block_length=self._block_length)
elif isinstance(matched_files, (list, tuple)):
dataset = _files_to_dataset(matched_files)
elif isinstance(matched_files, dict):
@@ -432,11 +476,14 @@ def _shuffle_and_decode(ds):
dataset = tf.nest.map_structure(_shuffle_and_decode, dataset)
if tf.nest.is_nested(dataset):
- dataset = self._combine_fn(dataset)
+ dataset = self._combine_fn(dataset) # pyrefly: ignore[not-callable]
if self._sample_fn is not None:
- dataset = dataset.apply(self._sample_fn)
- dataset = _maybe_map_fn(dataset, self._parser_fn)
+ dataset = dataset.apply(self._sample_fn) # pyrefly: ignore[missing-attribute]
+ dataset = _maybe_map_fn(dataset, self._parser_fn) # pyrefly: ignore[bad-argument-type]
+
+ if self._filter_fn is not None:
+ dataset = dataset.filter(self._filter_fn)
if self._cache:
dataset = dataset.cache()
@@ -444,6 +491,19 @@ def _shuffle_and_decode(ds):
dataset = dataset.repeat()
dataset = dataset.shuffle(self._shuffle_buffer_size, seed=self._seed)
+ # Applies tf.data service before batching operations. This is useful when
+ # tf.data service is shared between parallel trainers, and batch size is
+ # changing between parallel trainers. Then batch size is changing, tf.data
+ # services will be considered different instances if applied after batching
+ # operations, which make it difficult to share between parallel trainers.
+ # However, if there are additional expensive operations in
+ # self._transform_and_batch_fn and self._postprocess_fn, the entire tf.data
+ # pipeline could be slowed down. In this case, try to move these dataset
+ # operations into early stages if possible.
+ if (self._enable_shared_tf_data_service_between_parallel_trainers and
+ self._apply_tf_data_service_before_batching):
+ dataset = self._maybe_apply_data_service(dataset, input_context)
+
if self._transform_and_batch_fn is not None:
dataset = self._transform_and_batch_fn(dataset, input_context)
else:
@@ -469,13 +529,18 @@ def _maybe_apply_data_service(
num_consumers = input_context.num_input_pipelines * (
replicas_per_input_pipeline)
range_dataset = tf.data.Dataset.range(replicas_per_input_pipeline)
+ tfds_kwargs = {
+ 'processing_mode': 'parallel_epochs',
+ 'service': self._tf_data_service_address,
+ 'job_name': self._tf_data_service_job_name,
+ 'num_consumers': num_consumers
+ }
+ if self._enable_shared_tf_data_service_between_parallel_trainers:
+ raise ValueError('Shared tf.data service does not support round-robin'
+ ' tf.data service.')
dataset = range_dataset.map(lambda i: dataset.apply( # pylint: disable=g-long-lambda
tf.data.experimental.service.distribute(
- processing_mode='parallel_epochs',
- service=self._tf_data_service_address,
- job_name=self._tf_data_service_job_name,
- consumer_index=base_consumer_index + i,
- num_consumers=num_consumers)))
+ consumer_index=base_consumer_index + i, **tfds_kwargs)))
# Use parallel interleave to read multiple batches from a tf.data
# service worker in parallel.
dataset = dataset.interleave(
@@ -484,11 +549,21 @@ def _maybe_apply_data_service(
num_parallel_calls=replicas_per_input_pipeline,
deterministic=True)
else:
+ tfds_kwargs = {
+ 'processing_mode': 'parallel_epochs',
+ 'service': self._tf_data_service_address,
+ 'job_name': self._tf_data_service_job_name,
+ }
+ if self._enable_shared_tf_data_service_between_parallel_trainers:
+ tfds_kwargs.update({
+ 'processing_mode':
+ tf.data.experimental.service.ShardingPolicy.OFF,
+ 'cross_trainer_cache':
+ tf.data.experimental.service.CrossTrainerCache(
+ trainer_id=self._trainer_id)
+ })
dataset = dataset.apply(
- tf.data.experimental.service.distribute(
- processing_mode='parallel_epochs',
- service=self._tf_data_service_address,
- job_name=self._tf_data_service_job_name))
+ tf.data.experimental.service.distribute(**tfds_kwargs))
return dataset
def read(self,
@@ -496,15 +571,29 @@ def read(self,
dataset: Optional[tf.data.Dataset] = None) -> tf.data.Dataset:
"""Generates a tf.data.Dataset object."""
if dataset is None:
- dataset = self._read_data_source(self._matched_files, self._dataset_fn,
- input_context, self._tfds_builder)
+ dataset = self._read_data_source(self._matched_files, self._dataset_fn, # pyrefly: ignore[bad-argument-type]
+ input_context)
dataset = self._decode_and_parse_dataset(dataset, self._global_batch_size,
input_context)
dataset = _maybe_map_fn(dataset, self._postprocess_fn)
- dataset = self._maybe_apply_data_service(dataset, input_context)
+ if not (self._enable_shared_tf_data_service_between_parallel_trainers and
+ self._apply_tf_data_service_before_batching):
+ dataset = self._maybe_apply_data_service(dataset, input_context)
if self._deterministic is not None:
options = tf.data.Options()
- options.experimental_deterministic = self._deterministic
+ options.deterministic = self._deterministic
dataset = dataset.with_options(options)
+ if self._autotune_algorithm:
+ options = tf.data.Options()
+ options.autotune.autotune_algorithm = (
+ tf.data.experimental.AutotuneAlgorithm[self._autotune_algorithm]
+ )
+ dataset = dataset.with_options(options)
+
+ if self._ram_budget:
+ options = tf.data.Options()
+ options.autotune.ram_budget = self._ram_budget * 1024 * 1024 * 1024
+ dataset = dataset.with_options(options)
+
return dataset.prefetch(self._prefetch_buffer_size)
diff --git a/official/core/registry.py b/official/core/registry.py
index 5fdaf48ad87..3bce3524d33 100644
--- a/official/core/registry.py
+++ b/official/core/registry.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,6 @@
# limitations under the License.
"""Registry utility."""
-from absl import logging
def register(registered_collection, reg_key):
@@ -55,16 +54,8 @@ def decorator(fn_or_cls):
leaf_reg_key = reg_key
if leaf_reg_key in collection:
- if "beta" in fn_or_cls.__module__:
- # TODO(yeqing): Clean this temporary branch for beta.
- logging.warn(
- "Duplicate registeration of beta module "
- "name %r new %r old %r", reg_key, collection[leaf_reg_key],
- fn_or_cls.__module__)
- return fn_or_cls
- else:
- raise KeyError("Function or class {} registered multiple times.".format(
- leaf_reg_key))
+ raise KeyError("Function or class {} registered multiple times.".format(
+ leaf_reg_key))
collection[leaf_reg_key] = fn_or_cls
return fn_or_cls
diff --git a/official/core/registry_test.py b/official/core/registry_test.py
index 559b918e1e2..643a0412eeb 100644
--- a/official/core/registry_test.py
+++ b/official/core/registry_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for registry."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import registry
diff --git a/official/core/savedmodel_checkpoint_manager.py b/official/core/savedmodel_checkpoint_manager.py
new file mode 100644
index 00000000000..07aaf410094
--- /dev/null
+++ b/official/core/savedmodel_checkpoint_manager.py
@@ -0,0 +1,258 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Custom checkpoint manager that also exports saved models."""
+
+import os
+import re
+import time
+from typing import Callable, List, Mapping, Optional, Union
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+SAVED_MODULES_PATH_SUFFIX = 'saved_modules'
+
+
+def make_saved_modules_directory_name(checkpoint_name: str) -> str:
+ return f'{checkpoint_name}_{SAVED_MODULES_PATH_SUFFIX}'
+
+
+class SavedModelCheckpointManager(tf.train.CheckpointManager):
+ """A CheckpointManager that also exports `SavedModel`s."""
+
+ def __init__(self,
+ checkpoint: tf.train.Checkpoint,
+ directory: str,
+ max_to_keep: int,
+ modules_to_export: Optional[Mapping[str, tf.Module]] = None,
+ keep_checkpoint_every_n_hours: Optional[int] = None,
+ checkpoint_name: str = 'ckpt',
+ step_counter: Optional[tf.Variable] = None,
+ checkpoint_interval: Optional[int] = None,
+ init_fn: Optional[Callable[[], None]] = None):
+ """See base class."""
+ super().__init__(
+ checkpoint=checkpoint,
+ directory=directory,
+ max_to_keep=max_to_keep,
+ keep_checkpoint_every_n_hours=keep_checkpoint_every_n_hours,
+ checkpoint_name=checkpoint_name,
+ step_counter=step_counter,
+ checkpoint_interval=checkpoint_interval,
+ init_fn=init_fn)
+ self._modules_to_export = modules_to_export
+ self._savedmodels = self.get_existing_savedmodels()
+
+ def save(self,
+ checkpoint_number: Optional[int] = None,
+ check_interval: bool = True,
+ options: Optional[tf.train.CheckpointOptions] = None):
+ """See base class."""
+ checkpoint_path = super().save(
+ checkpoint_number=checkpoint_number,
+ check_interval=check_interval,
+ options=options)
+ if not checkpoint_path: # Nothing got written.
+ return
+ if not self._modules_to_export: # No modules to export.
+ logging.info('Skip saving SavedModel due to empty modules_to_export.')
+ return checkpoint_path
+
+ # Save the models for the checkpoint that just got written.
+ saved_modules_directory = make_saved_modules_directory_name(checkpoint_path)
+ # Atomic export of SavedModel. Write into a temporary direcotory and then
+ # rename as the final direcotory after finishing the writing.
+ # This can avoid trying to read an unfinished savedmodel.
+ saved_modules_directory_tmp = saved_modules_directory + '_temp'
+ for model_name, model in self._modules_to_export.items():
+ signatures = getattr(model, 'saved_model_signatures', None)
+ if signatures is not None:
+ tf.saved_model.save(
+ obj=model,
+ export_dir=os.path.join(saved_modules_directory_tmp, model_name),
+ signatures=signatures)
+ if tf.io.gfile.exists(saved_modules_directory_tmp):
+ tf.io.gfile.rename(saved_modules_directory_tmp, saved_modules_directory)
+
+ saved_modules_directories_to_keep = [
+ make_saved_modules_directory_name(ckpt) for ckpt in self.checkpoints
+ ]
+ existing_saved_modules_dirs = self.get_existing_savedmodels()
+
+ self._savedmodels = []
+ # Keep savedmodels in the same order as checkpoints (from oldest to newest).
+ for saved_modules_dir_to_keep in saved_modules_directories_to_keep:
+ if saved_modules_dir_to_keep in existing_saved_modules_dirs:
+ self._savedmodels.append(saved_modules_dir_to_keep)
+
+ for existing_saved_modules_dir in existing_saved_modules_dirs:
+ if existing_saved_modules_dir not in self._savedmodels:
+ tf.io.gfile.rmtree(existing_saved_modules_dir)
+
+ return checkpoint_path
+
+ def get_existing_savedmodels(self) -> List[str]:
+ """Gets a list of all existing SavedModel paths in `directory`.
+
+ Returns:
+ A list of all existing SavedModel paths.
+ """
+ saved_modules_glob = make_saved_modules_directory_name(
+ self._checkpoint_prefix + '-*')
+ savedmodels = tf.io.gfile.glob(saved_modules_glob)
+ # Filter out temporary savedmodel.
+ savedmodels = [
+ savedmodel
+ for savedmodel in savedmodels
+ if savedmodel.endswith(SAVED_MODULES_PATH_SUFFIX)
+ ]
+ return savedmodels
+
+ @property
+ def latest_savedmodel(self) -> Union[str, None]:
+ """The path of the most recent SavedModel in `directory`.
+
+ Returns:
+ The latest SavedModel path. If there are no SavedModels, returns `None`.
+ """
+ if self._savedmodels:
+ return self._savedmodels[-1]
+ return None
+
+ @property
+ def savedmodels(self) -> List[str]:
+ """A list of managed SavedModels.
+
+ Returns:
+ A list of SavedModel paths, sorted from oldest to newest.
+ """
+ return self._savedmodels
+
+ @property
+ def modules_to_export(self) -> Union[Mapping[str, tf.Module], None]:
+ return self._modules_to_export
+
+ def get_savedmodel_number_from_path(self,
+ savedmodel_path: str) -> Union[int, None]:
+ """Gets the savedmodel_number/checkpoint_number from savedmodel filepath.
+
+ The savedmodel_number is global step when using with orbit controller.
+
+ Args:
+ savedmodel_path: savedmodel directory path.
+
+ Returns:
+ Savedmodel number or None if no matched pattern found in savedmodel path.
+ """
+ pattern = rf'\d+_{SAVED_MODULES_PATH_SUFFIX}$'
+ savedmodel_number = re.search(pattern, savedmodel_path)
+ if savedmodel_number:
+ savedmodel_number = savedmodel_number.group()
+ return int(savedmodel_number[:-len(SAVED_MODULES_PATH_SUFFIX) - 1])
+ return None
+
+ def savedmodels_iterator(self,
+ min_interval_secs: float = 0,
+ timeout: Optional[float] = None,
+ timeout_fn: Optional[Callable[[], bool]] = None):
+ """Continuously yield new SavedModel files as they appear.
+
+ The iterator only checks for new savedmodels when control flow has been
+ reverted to it. The logic is same to the `train.checkpoints_iterator`.
+
+ Args:
+ min_interval_secs: The minimum number of seconds between yielding
+ savedmodels.
+ timeout: The maximum number of seconds to wait between savedmodels. If
+ left as `None`, then the process will wait indefinitely.
+ timeout_fn: Optional function to call after a timeout. If the function
+ returns True, then it means that no new savedmodels will be generated
+ and the iterator will exit. The function is called with no arguments.
+
+ Yields:
+ String paths to latest SavedModel files as they arrive.
+ """
+ savedmodel_path = None
+ while True:
+ new_savedmodel_path = self.wait_for_new_savedmodel(
+ savedmodel_path, timeout=timeout)
+ if new_savedmodel_path is None:
+ if not timeout_fn:
+ # timed out
+ logging.info('Timed-out waiting for a savedmodel.')
+ return
+ if timeout_fn():
+ # The timeout_fn indicated that we are truly done.
+ return
+ else:
+ # The timeout_fn indicated that more savedmodels may come.
+ continue
+ start = time.time()
+ savedmodel_path = new_savedmodel_path
+ yield savedmodel_path
+ time_to_next_eval = start + min_interval_secs - time.time()
+ if time_to_next_eval > 0:
+ time.sleep(time_to_next_eval)
+
+ def wait_for_new_savedmodel(
+ self,
+ last_savedmodel: Optional[str] = None,
+ seconds_to_sleep: float = 1.0,
+ timeout: Optional[float] = None) -> Union[str, None]:
+ """Waits until a new savedmodel file is found.
+
+ Args:
+ last_savedmodel: The last savedmodel path used or `None` if we're
+ expecting a savedmodel for the first time.
+ seconds_to_sleep: The number of seconds to sleep for before looking for a
+ new savedmodel.
+ timeout: The maximum number of seconds to wait. If left as `None`, then
+ the process will wait indefinitely.
+
+ Returns:
+ A new savedmodel path, or None if the timeout was reached.
+ """
+ logging.info('Waiting for new savedmodel at %s', self._directory)
+ stop_time = time.time() + timeout if timeout is not None else None
+
+ last_savedmodel_number = -1
+ if last_savedmodel:
+ last_savedmodel_number = self.get_savedmodel_number_from_path(
+ last_savedmodel)
+
+ while True:
+ if stop_time is not None and time.time() + seconds_to_sleep > stop_time:
+ return None
+
+ existing_savedmodels = {}
+ for savedmodel_path in self.get_existing_savedmodels():
+ savedmodel_number = self.get_savedmodel_number_from_path(
+ savedmodel_path)
+ if savedmodel_number is not None:
+ existing_savedmodels[savedmodel_number] = savedmodel_path
+
+ # Find the first savedmodel with larger step number as next savedmodel.
+ savedmodel_path = None
+ existing_savedmodels = dict(sorted(existing_savedmodels.items()))
+ for savedmodel_number in existing_savedmodels:
+ if savedmodel_number > last_savedmodel_number:
+ savedmodel_path = existing_savedmodels[savedmodel_number]
+ break
+
+ if savedmodel_path:
+ logging.info('Found new savedmodel at %s', savedmodel_path)
+ return savedmodel_path
+ else:
+ time.sleep(seconds_to_sleep)
diff --git a/official/core/savedmodel_checkpoint_manager_test.py b/official/core/savedmodel_checkpoint_manager_test.py
new file mode 100644
index 00000000000..299560801d3
--- /dev/null
+++ b/official/core/savedmodel_checkpoint_manager_test.py
@@ -0,0 +1,125 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import os
+import time
+from typing import Iterable
+
+import tensorflow as tf, tf_keras
+
+from official.core import savedmodel_checkpoint_manager
+
+
+def _models_exist(checkpoint_path: str, models: Iterable[str]) -> bool:
+ for model_name in models:
+ if not tf.io.gfile.isdir(
+ os.path.join(
+ savedmodel_checkpoint_manager.make_saved_modules_directory_name(
+ checkpoint_path), model_name)):
+ return False
+ return True
+
+
+class _ModelForTest(tf_keras.Model):
+ def __init__(self, hidden_size: int = 8):
+ super().__init__()
+ self.dense = tf_keras.layers.Dense(hidden_size)
+
+ @tf.function(input_signature=[tf.TensorSpec([None, 16])])
+ def call(self, inputs):
+ return self.dense(inputs)
+
+ @property
+ def saved_model_signatures(self):
+ # Build SavedModel signatures.
+ return dict(serving_default=self.call)
+
+
+class CheckpointManagerTest(tf.test.TestCase):
+
+ def _create_manager(self, max_to_keep: int = 1) -> tf.train.CheckpointManager:
+ """Sets up SavedModelCheckpointManager object.
+
+ Args:
+ max_to_keep: max number of savedmodels to keep.
+
+ Returns:
+ created savedmodel manager.
+ """
+ models = {
+ 'model_1': _ModelForTest(12),
+ 'model_2': _ModelForTest(14),
+ }
+ checkpoint = tf.train.Checkpoint()
+ manager = savedmodel_checkpoint_manager.SavedModelCheckpointManager(
+ checkpoint=checkpoint,
+ directory=self.get_temp_dir(),
+ max_to_keep=max_to_keep,
+ modules_to_export=models)
+ return manager
+
+ def test_max_to_keep(self):
+ manager = self._create_manager()
+ models = manager.modules_to_export
+ first_path = manager.save()
+ second_path = manager.save()
+
+ savedmodel = savedmodel_checkpoint_manager.make_saved_modules_directory_name(
+ manager.latest_checkpoint)
+ self.assertEqual(savedmodel, manager.latest_savedmodel)
+ self.assertTrue(_models_exist(second_path, models.keys()))
+ self.assertFalse(_models_exist(first_path, models.keys()))
+
+ def test_returns_none_after_timeout(self):
+ manager = self._create_manager()
+ start = time.time()
+ ret = manager.wait_for_new_savedmodel(
+ None, timeout=1.0, seconds_to_sleep=0.5)
+ end = time.time()
+ self.assertIsNone(ret)
+ # We've waited 0.5 second.
+ self.assertGreater(end, start + 0.5)
+ # The timeout kicked in.
+ self.assertLess(end, start + 0.6)
+
+ def test_saved_model_iterator(self):
+ manager = self._create_manager(max_to_keep=2)
+ self.assertIsNotNone(manager.save(checkpoint_number=1))
+ self.assertIsNotNone(manager.save(checkpoint_number=2))
+ self.assertIsNotNone(manager.save(checkpoint_number=3))
+
+ # Savedmodels are in time order.
+ expected_savedmodels = manager.savedmodels
+ # Order not guaranteed.
+ existing_savedmodels = manager.get_existing_savedmodels()
+ savedmodels = list(manager.savedmodels_iterator(timeout=3.0))
+ self.assertEqual(savedmodels, expected_savedmodels)
+ self.assertEqual(set(savedmodels), set(existing_savedmodels))
+
+ def test_saved_model_iterator_timeout_fn(self):
+ manager = self._create_manager()
+ timeout_fn_calls = [0]
+
+ def timeout_fn():
+ timeout_fn_calls[0] += 1
+ return timeout_fn_calls[0] > 3
+
+ results = list(
+ manager.savedmodels_iterator(timeout=0.1, timeout_fn=timeout_fn))
+ self.assertEqual([], results)
+ self.assertEqual(4, timeout_fn_calls[0])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/core/task_factory.py b/official/core/task_factory.py
index 4dee1fe2e21..020853afd76 100644
--- a/official/core/task_factory.py
+++ b/official/core/task_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/core/test_utils.py b/official/core/test_utils.py
index 7edeff7c632..96ad10204cd 100644
--- a/official/core/test_utils.py
+++ b/official/core/test_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,18 +14,18 @@
"""Utils for testing."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class FakeKerasModel(tf.keras.Model):
+class FakeKerasModel(tf_keras.Model):
"""Fake keras model for testing."""
def __init__(self):
super().__init__()
- self.dense = tf.keras.layers.Dense(4, activation=tf.nn.relu)
- self.dense2 = tf.keras.layers.Dense(4, activation=tf.nn.relu)
+ self.dense = tf_keras.layers.Dense(4, activation=tf.nn.relu)
+ self.dense2 = tf_keras.layers.Dense(4, activation=tf.nn.relu)
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
return self.dense2(self.dense(inputs))
diff --git a/official/core/tf_example_builder.py b/official/core/tf_example_builder.py
new file mode 100644
index 00000000000..34ab3fd52c2
--- /dev/null
+++ b/official/core/tf_example_builder.py
@@ -0,0 +1,144 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builder class for preparing tf.train.Example."""
+
+# https://www.python.org/dev/peps/pep-0563/#enabling-the-future-behavior-in-python-3-7
+from __future__ import annotations
+
+from typing import Mapping, Sequence, Union
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+BytesValueType = Union[bytes, Sequence[bytes], str, Sequence[str]]
+
+_to_array = lambda v: [v] if not isinstance(v, (list, np.ndarray)) else v
+_to_bytes = lambda v: v.encode() if isinstance(v, str) else v
+_to_bytes_array = lambda v: list(map(_to_bytes, _to_array(v)))
+
+
+class TfExampleBuilder(object):
+ """Builder class for preparing tf.train.Example.
+
+ Read API doc at https://www.tensorflow.org/api_docs/python/tf/train/Example.
+
+ Example usage:
+ >>> example_builder = TfExampleBuilder()
+ >>> example = (
+ example_builder.add_bytes_feature('feature_a', 'foobarbaz')
+ .add_ints_feature('feature_b', [1, 2, 3])
+ .example)
+ """
+
+ def __init__(self) -> None:
+ self._example = tf.train.Example()
+
+ @property
+ def example(self) -> tf.train.Example:
+ """Returns a copy of the generated tf.train.Example proto."""
+ return self._example
+
+ @property
+ def serialized_example(self) -> str:
+ """Returns a serialized string of the generated tf.train.Example proto."""
+ return self._example.SerializeToString()
+
+ def set(self, example: tf.train.Example) -> TfExampleBuilder:
+ """Sets the example."""
+ self._example = example
+ return self
+
+ def reset(self) -> TfExampleBuilder:
+ """Resets the example to an empty proto."""
+ self._example = tf.train.Example()
+ return self
+
+ ###### Basic APIs for primitive data types ######
+ def add_feature_dict(
+ self, feature_dict: Mapping[str, tf.train.Feature]) -> TfExampleBuilder:
+ """Adds the predefined `feature_dict` to the example.
+
+ Note: Please prefer to using feature-type-specific methods.
+
+ Args:
+ feature_dict: A dictionary from tf.Example feature key to
+ tf.train.Feature.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ for k, v in feature_dict.items():
+ self._example.features.feature[k].CopyFrom(v)
+ return self
+
+ def add_feature(self, key: str,
+ feature: tf.train.Feature) -> TfExampleBuilder:
+ """Adds predefined `feature` with `key` to the example.
+
+ Args:
+ key: String key of the feature.
+ feature: The feature to be added to the example.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ self._example.features.feature[key].CopyFrom(feature)
+ return self
+
+ def add_bytes_feature(self, key: str,
+ value: BytesValueType) -> TfExampleBuilder:
+ """Adds byte(s) or string(s) with `key` to the example.
+
+ Args:
+ key: String key of the feature.
+ value: The byte(s) or string(s) to be added to the example.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ return self.add_feature(
+ key,
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=_to_bytes_array(value))))
+
+ def add_ints_feature(self, key: str,
+ value: Union[int, Sequence[int]]) -> TfExampleBuilder:
+ """Adds integer(s) with `key` to the example.
+
+ Args:
+ key: String key of the feature.
+ value: The integer(s) to be added to the example.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ return self.add_feature(
+ key,
+ tf.train.Feature(int64_list=tf.train.Int64List(value=_to_array(value))))
+
+ def add_floats_feature(
+ self, key: str, value: Union[float, Sequence[float]]) -> TfExampleBuilder:
+ """Adds float(s) with `key` to the example.
+
+ Args:
+ key: String key of the feature.
+ value: The float(s) to be added to the example.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ return self.add_feature(
+ key,
+ tf.train.Feature(float_list=tf.train.FloatList(value=_to_array(value))))
diff --git a/official/core/tf_example_builder_test.py b/official/core/tf_example_builder_test.py
new file mode 100644
index 00000000000..1def20877bf
--- /dev/null
+++ b/official/core/tf_example_builder_test.py
@@ -0,0 +1,165 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for tf_example_builder.
+
+See `test_add_image_matrix_feature_with_fake_image` for the typical structure of
+a unit test.
+"""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+from official.core import tf_example_builder
+
+
+class TfExampleBuilderTest(tf.test.TestCase, parameterized.TestCase):
+
+ def test_init_an_empty_example(self):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example = example_builder.example
+ self.assertProtoEquals('', example)
+
+ def test_init_an_empty_serialized_example(self):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example = example_builder.serialized_example
+ self.assertProtoEquals('', example)
+
+ def test_add_feature(self):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_feature(
+ 'foo',
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[b'Hello World!'])))
+ example = example_builder.example
+ # Use proto text to show how the entire proto would look like.
+ self.assertProtoEquals(
+ """
+ features: {
+ feature: {
+ key: "foo"
+ value: {
+ bytes_list: {
+ value: "Hello World!"
+ }
+ }
+ }
+ }""", example)
+
+ def test_add_feature_dict(self):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_feature_dict({
+ 'foo':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[b'Hello World!'])),
+ 'bar':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[299, 792, 458]))
+ })
+ example = example_builder.example
+ # Use proto text to show how the entire proto would look like.
+ self.assertProtoEquals(
+ """
+ features: {
+ feature: {
+ key: "foo"
+ value: {
+ bytes_list: {
+ value: "Hello World!"
+ }
+ }
+ }
+ feature: {
+ key: "bar"
+ value: {
+ int64_list: {
+ value: 299
+ value: 792
+ value: 458
+ }
+ }
+ }
+ }""", example)
+
+ @parameterized.named_parameters(
+ ('single_bytes', b'Hello World!', b'Hello World!'),
+ ('single_string', 'Hello World!', b'Hello World!'))
+ def test_add_single_byte_feature(self, value, expected_value):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_bytes_feature('foo', value)
+ example = example_builder.example
+ # Use constructor to easily work with test parameters.
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'foo':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[expected_value]))
+ })), example)
+
+ @parameterized.named_parameters(
+ ('multiple_bytes', [b'Hello World!', b'Good Morning!'
+ ], [b'Hello World!', b'Good Morning!']),
+ ('multiple_sring', ['Hello World!', 'Good Morning!'
+ ], [b'Hello World!', b'Good Morning!']))
+ def test_add_multiple_bytes_feature(self, values, expected_values):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_bytes_feature('foo', values)
+ example = example_builder.example
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'foo':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=expected_values))
+ })), example)
+
+ @parameterized.named_parameters(
+ ('single_integer', 123, [123]),
+ ('multiple_integers', [123, 456, 789], [123, 456, 789]))
+ def test_add_ints_feature(self, value, expected_value):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_ints_feature('bar', value)
+ example = example_builder.example
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'bar':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=expected_value))
+ })), example)
+
+ @parameterized.named_parameters(
+ ('single_float', 3.14, [3.14]),
+ ('multiple_floats', [3.14, 1.57, 6.28], [3.14, 1.57, 6.28]))
+ def test_add_floats_feature(self, value, expected_value):
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_floats_feature('baz', value)
+ example = example_builder.example
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'baz':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=expected_value))
+ })), example)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/core/tf_example_feature_key.py b/official/core/tf_example_feature_key.py
new file mode 100644
index 00000000000..feb2ae4e1b8
--- /dev/null
+++ b/official/core/tf_example_feature_key.py
@@ -0,0 +1,62 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data classes for tf.Example proto feature keys.
+
+Feature keys are grouped by feature types. Key names follow conventions in
+go/tf-example.
+"""
+import dataclasses
+import functools
+from typing import Optional
+
+# Disable init function to use the one defined in base class.
+dataclass = functools.partial(dataclasses.dataclass(init=False)) # pyrefly: ignore[bad-argument-type]
+
+
+@dataclass
+class TfExampleFeatureKeyBase:
+ """Base dataclass for defining tf.Example proto feature keys.
+
+ This class defines the logic of adding prefix to feature keys. Subclasses
+ will define feature keys for a specific feature type in data fields.
+
+ NOTE: Please follow subclass examples in this module to define feature keys
+ for a new feature type.
+ """
+
+ def __init__(self, prefix: Optional[str] = None):
+ """Instantiates the feature key class.
+
+ Adds a string prefix to all fields of a feature key instance if `prefix` is
+ not None nor empty.
+
+ Example usage:
+
+ >>> test_key = EncodedImageFeatureKey()
+ >>> test_key.encoded
+ image/encoded
+ >>> test_key = EncodedImageFeatureKey('prefix')
+ >>> test_key.encoded
+ prefix/image/encoded
+
+ Args:
+ prefix: A prefix string that will be added before the feature key string
+ with a trailing slash '/'.
+ """
+ if prefix:
+ for field in dataclasses.fields(self): # pytype: disable=wrong-arg-types # re-none
+ key_name = field.name
+ key_value = getattr(self, key_name)
+ setattr(self, key_name, f'{prefix}/{key_value}')
diff --git a/official/core/tf_example_feature_key_test.py b/official/core/tf_example_feature_key_test.py
new file mode 100644
index 00000000000..a53aa4fd1c9
--- /dev/null
+++ b/official/core/tf_example_feature_key_test.py
@@ -0,0 +1,49 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for tf_example_feature_key."""
+import dataclasses
+import inspect
+from absl.testing import absltest
+from absl.testing import parameterized
+
+from official.core import tf_example_feature_key
+
+
+@tf_example_feature_key.dataclass
+class TestFeatureKey(tf_example_feature_key.TfExampleFeatureKeyBase):
+ test: str = 'foo/bar'
+
+
+class TfExampleFeatureKeyTest(parameterized.TestCase):
+
+ def test_add_prefix_success(self):
+ test_key = TestFeatureKey('prefix')
+ self.assertEqual(test_key.test, 'prefix/foo/bar')
+
+ @parameterized.parameters(None, '')
+ def test_add_prefix_skip_success(self, prefix):
+ test_key = TestFeatureKey(prefix)
+ self.assertEqual(test_key.test, 'foo/bar')
+
+ def test_all_feature_key_classes_are_valid(self):
+ for _, obj in inspect.getmembers(tf_example_feature_key):
+ if inspect.isclass(obj):
+ self.assertTrue(dataclasses.is_dataclass(obj))
+ self.assertTrue(
+ issubclass(obj, tf_example_feature_key.TfExampleFeatureKeyBase))
+
+
+if __name__ == '__main__':
+ absltest.main()
diff --git a/official/core/train_lib.py b/official/core/train_lib.py
index 1c9dbcfbde4..f5067f82f43 100644
--- a/official/core/train_lib.py
+++ b/official/core/train_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,12 @@
"""TFM common training driver library."""
# pytype: disable=attribute-error
import os
-from typing import Any, Mapping, Optional, Tuple
-
-# Import libraries
+import tempfile
+from typing import Any, List, Mapping, Optional, Tuple
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import actions
from official.core import base_task
@@ -32,6 +31,279 @@
maybe_create_best_ckpt_exporter = train_utils.maybe_create_best_ckpt_exporter
+class OrbitExperimentRunner:
+ """Runs experiment with Orbit training loop.
+
+ The default experiment runner for model garden experiments. User can
+ customize the experiment pipeline by subclassing this class and replacing
+ components or functions.
+
+ For example, an experiment runner with customized checkpoint manager:
+
+ ```python
+ class MyExpRunnerWithExporter(OrbitExperimentRunner):
+ def _maybe_build_checkpoint_manager(sefl):
+ # Replaces the default CheckpointManger with a customized one.
+ return MyCheckpointManager(*args)
+
+ # In user code, instead of the orginal
+ # `OrbitExperimentRunner(..).run(mode)`, now user can do:
+ MyExpRunnerWithExporter(**needed_kwargs).run(mode)
+ ```
+
+ Similar override can be done to other components.
+ """
+
+ def __init__(
+ self,
+ distribution_strategy: tf.distribute.Strategy,
+ task: base_task.Task,
+ mode: str,
+ params: config_definitions.ExperimentConfig,
+ model_dir: str,
+ run_post_eval: bool = False,
+ save_summary: bool = True,
+ train_actions: Optional[List[orbit.Action]] = None,
+ eval_actions: Optional[List[orbit.Action]] = None,
+ trainer: Optional[base_trainer.Trainer] = None,
+ controller_cls=orbit.Controller,
+ summary_manager: Optional[orbit.utils.SummaryManager] = None,
+ eval_summary_manager: Optional[orbit.utils.SummaryManager] = None,
+ enable_async_checkpointing: bool = False,
+ ):
+ """Constructor.
+
+ Args:
+ distribution_strategy: A distribution strategy.
+ task: A Task instance.
+ mode: A 'str', specifying the mode. Can be 'train', 'eval',
+ 'train_and_eval' or 'continuous_eval'.
+ params: ExperimentConfig instance.
+ model_dir: A 'str', a path to store model checkpoints and summaries.
+ run_post_eval: Whether to run post eval once after training, metrics logs
+ are returned.
+ save_summary: Whether to save train and validation summary.
+ train_actions: Optional list of Orbit train actions.
+ eval_actions: Optional list of Orbit eval actions.
+ trainer: the base_trainer.Trainer instance. It should be created within
+ the strategy.scope().
+ controller_cls: The controller class to manage the train and eval process.
+ Must be a orbit.Controller subclass.
+ summary_manager: Instance of the summary manager to override default
+ summary manager.
+ eval_summary_manager: Instance of the eval summary manager to override
+ default eval summary manager.
+ enable_async_checkpointing: Optional boolean indicating whether to enable
+ async checkpoint saving.
+ """
+ self.strategy = distribution_strategy or tf.distribute.get_strategy()
+ self._params = params
+ self._model_dir = model_dir
+ self._mode = mode
+ self._run_post_eval = run_post_eval
+
+ self._trainer = trainer or self._build_trainer(
+ task,
+ train='train' in mode,
+ evaluate=('eval' in mode) or run_post_eval)
+ assert self.trainer is not None
+ self._checkpoint_manager = self._maybe_build_checkpoint_manager()
+ self._summary_manager = summary_manager
+ self._eval_summary_manager = eval_summary_manager
+ self._controller = self._build_controller(
+ trainer=self.trainer if 'train' in mode else None,
+ evaluator=self.trainer,
+ save_summary=save_summary,
+ train_actions=train_actions,
+ eval_actions=eval_actions,
+ controller_cls=controller_cls,
+ enable_async_checkpointing=enable_async_checkpointing)
+
+ @property
+ def params(self) -> config_definitions.ExperimentConfig:
+ """The whole experiment parameters object."""
+ return self._params
+
+ @property
+ def model_dir(self) -> str:
+ """Path to the model folder, which stores checkpoints, params, log, etc."""
+ return self._model_dir
+
+ @property
+ def trainer(self) -> base_trainer.Trainer:
+ """The underlying Orbit Trainer object."""
+ return self._trainer
+
+ @property
+ def checkpoint_manager(self) -> Optional[tf.train.CheckpointManager]:
+ """The CheckpointManager that stores the checkpoints in a train job."""
+ return self._checkpoint_manager
+
+ @property
+ def controller(self) -> orbit.Controller:
+ """The Orbit controller object."""
+ return self._controller
+
+ def _build_trainer(self, task: base_task.Task, train: bool,
+ evaluate: bool) -> base_trainer.Trainer:
+ """Create trainer."""
+ with self.strategy.scope():
+ trainer = train_utils.create_trainer(
+ self.params,
+ task,
+ train=train,
+ evaluate=evaluate,
+ checkpoint_exporter=self._build_best_checkpoint_exporter())
+ return trainer
+
+ def _build_best_checkpoint_exporter(self):
+ return maybe_create_best_ckpt_exporter(self.params, self.model_dir)
+
+ def _maybe_build_checkpoint_manager(
+ self) -> Optional[tf.train.CheckpointManager]:
+ """Maybe create a CheckpointManager."""
+ assert self.trainer is not None
+ if self.trainer.checkpoint:
+ if self.model_dir is None:
+ raise ValueError('model_dir must be specified, but got None')
+
+ if (not self.strategy) or self.strategy.extended.should_checkpoint:
+ ckpt_path = self.model_dir
+ max_to_keep = self.params.trainer.max_to_keep
+ else:
+ # In multi worker training we need every worker to save checkpoint,
+ # because variables can trigger synchronization on read and
+ # synchronization needs all workers to participate. To avoid workers
+ # overriding each other we save to a temporary directory on non-chief
+ # workers.
+ ckpt_path = tempfile.mkdtemp()
+ max_to_keep = 1
+
+ checkpoint_manager = tf.train.CheckpointManager(
+ self.trainer.checkpoint,
+ directory=ckpt_path,
+ max_to_keep=max_to_keep,
+ step_counter=self.trainer.global_step,
+ checkpoint_interval=self.params.trainer.checkpoint_interval,
+ init_fn=self.trainer.initialize)
+ else:
+ checkpoint_manager = None
+ return checkpoint_manager
+
+ def _build_controller(
+ self,
+ trainer,
+ evaluator,
+ save_summary: bool = True,
+ train_actions: Optional[List[orbit.Action]] = None,
+ eval_actions: Optional[List[orbit.Action]] = None,
+ controller_cls=orbit.Controller,
+ enable_async_checkpointing: bool = False,
+ ) -> orbit.Controller:
+ """Builds a Orbit controler."""
+ train_actions = [] if not train_actions else train_actions
+ if trainer:
+ checkpoint_manager = self.checkpoint_manager
+ assert checkpoint_manager, 'Checkpoint manager required but undefined.'
+ train_actions += actions.get_train_actions(
+ self.params,
+ trainer,
+ self.model_dir,
+ checkpoint_manager=checkpoint_manager,
+ )
+
+ eval_actions = [] if not eval_actions else eval_actions
+ if evaluator:
+ eval_actions += actions.get_eval_actions(self.params, evaluator,
+ self.model_dir)
+
+ if save_summary:
+ eval_summary_dir = os.path.join(
+ self.model_dir, self.params.trainer.validation_summary_subdir
+ )
+ else:
+ eval_summary_dir = None
+
+ controller = controller_cls(
+ strategy=self.strategy,
+ trainer=trainer,
+ evaluator=evaluator,
+ global_step=self.trainer.global_step,
+ steps_per_loop=self.params.trainer.steps_per_loop,
+ checkpoint_manager=self.checkpoint_manager,
+ enable_async_checkpointing=enable_async_checkpointing,
+ summary_dir=os.path.join(self.model_dir, 'train')
+ if (save_summary)
+ else None,
+ eval_summary_dir=eval_summary_dir,
+ summary_interval=self.params.trainer.summary_interval
+ if (save_summary)
+ else None,
+ train_actions=train_actions,
+ eval_actions=eval_actions,
+ summary_manager=self._summary_manager
+ if hasattr(self, '_summary_manager')
+ else None,
+ eval_summary_manager=self._eval_summary_manager
+ if hasattr(self, '_eval_summary_manager')
+ else None,
+ )
+ return controller
+
+ def run(self) -> Tuple[tf_keras.Model, Mapping[str, Any]]:
+ """Run experiments by mode.
+
+ Returns:
+ A 2-tuple of (model, eval_logs).
+ model: `tf_keras.Model` instance.
+ eval_logs: returns eval metrics logs when run_post_eval is set to True,
+ otherwise, returns {}.
+ """
+ mode = self._mode
+ params = self.params
+ logging.info('Starts to execute mode: %s', mode)
+ with self.strategy.scope():
+ if mode == 'train' or mode == 'train_and_post_eval':
+ self.controller.train(steps=params.trainer.train_steps)
+ elif mode == 'train_and_eval':
+ self.controller.train_and_evaluate(
+ train_steps=params.trainer.train_steps,
+ eval_steps=params.trainer.validation_steps,
+ eval_interval=params.trainer.validation_interval)
+ elif mode == 'eval':
+ self.controller.evaluate(steps=params.trainer.validation_steps)
+ elif mode == 'continuous_eval':
+
+ def timeout_fn():
+ if self.trainer.global_step.numpy() >= params.trainer.train_steps:
+ return True
+ return False
+
+ self.controller.evaluate_continuously(
+ steps=params.trainer.validation_steps,
+ timeout=params.trainer.continuous_eval_timeout,
+ timeout_fn=timeout_fn)
+ else:
+ raise NotImplementedError('The mode is not implemented: %s' % mode)
+
+ num_params = train_utils.try_count_params(self.trainer.model)
+ if num_params is not None:
+ logging.info('Number of trainable params in model: %f Millions.',
+ num_params / 10.**6)
+
+ flops = train_utils.try_count_flops(self.trainer.model)
+ if flops is not None:
+ logging.info('FLOPs (multi-adds) in model: %f Billions.',
+ flops / 10.**9 / 2)
+
+ if self._run_post_eval or mode == 'train_and_post_eval':
+ with self.strategy.scope():
+ return self.trainer.model, self.controller.evaluate( # pytype: disable=bad-return-type # always-use-property-annotation
+ steps=params.trainer.validation_steps)
+ else:
+ return self.trainer.model, {}
+
+
def run_experiment(
distribution_strategy: tf.distribute.Strategy,
task: base_task.Task,
@@ -40,9 +312,14 @@ def run_experiment(
model_dir: str,
run_post_eval: bool = False,
save_summary: bool = True,
+ train_actions: Optional[List[orbit.Action]] = None,
+ eval_actions: Optional[List[orbit.Action]] = None,
trainer: Optional[base_trainer.Trainer] = None,
- controller_cls=orbit.Controller
-) -> Tuple[tf.keras.Model, Mapping[str, Any]]:
+ controller_cls=orbit.Controller,
+ summary_manager: Optional[orbit.utils.SummaryManager] = None,
+ eval_summary_manager: Optional[orbit.utils.SummaryManager] = None,
+ enable_async_checkpointing: bool = False,
+) -> Tuple[tf_keras.Model, Mapping[str, Any]]:
"""Runs train/eval configured by the experiment params.
Args:
@@ -55,96 +332,39 @@ def run_experiment(
run_post_eval: Whether to run post eval once after training, metrics logs
are returned.
save_summary: Whether to save train and validation summary.
+ train_actions: Optional list of Orbit train actions.
+ eval_actions: Optional list of Orbit eval actions.
trainer: the base_trainer.Trainer instance. It should be created within the
strategy.scope().
controller_cls: The controller class to manage the train and eval process.
Must be a orbit.Controller subclass.
+ summary_manager: Instance of the summary manager to override default summary
+ manager.
+ eval_summary_manager: Instance of the eval summary manager to override
+ default eval summary manager.
+ enable_async_checkpointing: Optional boolean indicating whether to enable
+ async checkpoint saving.
Returns:
A 2-tuple of (model, eval_logs).
- model: `tf.keras.Model` instance.
+ model: `tf_keras.Model` instance.
eval_logs: returns eval metrics logs when run_post_eval is set to True,
otherwise, returns {}.
"""
-
- with distribution_strategy.scope():
- if not trainer:
- trainer = train_utils.create_trainer(
- params,
- task,
- train='train' in mode,
- evaluate=('eval' in mode) or run_post_eval,
- checkpoint_exporter=maybe_create_best_ckpt_exporter(
- params, model_dir))
-
- if trainer.checkpoint:
- if model_dir is None:
- raise ValueError('model_dir must be specified, but got None')
- checkpoint_manager = tf.train.CheckpointManager(
- trainer.checkpoint,
- directory=model_dir,
- max_to_keep=params.trainer.max_to_keep,
- step_counter=trainer.global_step,
- checkpoint_interval=params.trainer.checkpoint_interval,
- init_fn=trainer.initialize)
- else:
- checkpoint_manager = None
-
- controller = controller_cls(
- strategy=distribution_strategy,
- trainer=trainer if 'train' in mode else None,
- evaluator=trainer,
- global_step=trainer.global_step,
- steps_per_loop=params.trainer.steps_per_loop,
- checkpoint_manager=checkpoint_manager,
- summary_dir=os.path.join(model_dir, 'train') if (save_summary) else None,
- eval_summary_dir=os.path.join(model_dir,
- params.trainer.validation_summary_subdir) if
- (save_summary) else None,
- summary_interval=params.trainer.summary_interval if
- (save_summary) else None,
- train_actions=actions.get_train_actions(
- params, trainer, model_dir, checkpoint_manager=checkpoint_manager),
- eval_actions=actions.get_eval_actions(params, trainer, model_dir))
-
- logging.info('Starts to execute mode: %s', mode)
- with distribution_strategy.scope():
- if mode == 'train':
- controller.train(steps=params.trainer.train_steps)
- elif mode == 'train_and_eval':
- controller.train_and_evaluate(
- train_steps=params.trainer.train_steps,
- eval_steps=params.trainer.validation_steps,
- eval_interval=params.trainer.validation_interval)
- elif mode == 'eval':
- controller.evaluate(steps=params.trainer.validation_steps)
- elif mode == 'continuous_eval':
-
- def timeout_fn():
- if trainer.global_step.numpy() >= params.trainer.train_steps:
- return True
- return False
-
- controller.evaluate_continuously(
- steps=params.trainer.validation_steps,
- timeout=params.trainer.continuous_eval_timeout,
- timeout_fn=timeout_fn)
- else:
- raise NotImplementedError('The mode is not implemented: %s' % mode)
-
- num_params = train_utils.try_count_params(trainer.model)
- if num_params is not None:
- logging.info('Number of trainable params in model: %f Millions.',
- num_params / 10.**6)
-
- flops = train_utils.try_count_flops(trainer.model)
- if flops is not None:
- logging.info('FLOPs (multi-adds) in model: %f Billions.',
- flops / 10.**9 / 2)
-
- if run_post_eval:
- with distribution_strategy.scope():
- return trainer.model, trainer.evaluate(
- tf.convert_to_tensor(params.trainer.validation_steps))
- else:
- return trainer.model, {}
+ runner = OrbitExperimentRunner(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=mode,
+ params=params,
+ model_dir=model_dir,
+ run_post_eval=run_post_eval,
+ save_summary=save_summary,
+ train_actions=train_actions,
+ eval_actions=eval_actions,
+ trainer=trainer,
+ controller_cls=controller_cls,
+ summary_manager=summary_manager,
+ eval_summary_manager=eval_summary_manager,
+ enable_async_checkpointing=enable_async_checkpointing,
+ )
+ return runner.run()
diff --git a/official/core/train_lib_test.py b/official/core/train_lib_test.py
index 61a4ccee13e..0c526b44218 100644
--- a/official/core/train_lib_test.py
+++ b/official/core/train_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl.testing import flagsaver
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -117,6 +117,61 @@ def test_end_to_end(self, distribution_strategy, flag_mode, run_post_eval):
model_dir=model_dir,
run_post_eval=run_post_eval)
+ @combinations.generate(
+ combinations.combine(
+ distribution_strategy=[
+ strategy_combinations.default_strategy,
+ strategy_combinations.cloud_tpu_strategy,
+ strategy_combinations.one_device_strategy_gpu,
+ ],
+ flag_mode=['train', 'eval', 'train_and_eval'],
+ run_post_eval=[True, False]))
+ def test_end_to_end_class(self, distribution_strategy, flag_mode,
+ run_post_eval):
+ model_dir = self.get_temp_dir()
+ flags_dict = dict(
+ experiment='mock',
+ mode=flag_mode,
+ model_dir=model_dir,
+ params_override=json.dumps(self._test_config))
+ with flagsaver.flagsaver(**flags_dict):
+ params = train_utils.parse_configuration(flags.FLAGS)
+ train_utils.serialize_config(params, model_dir)
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ _, logs = train_lib.OrbitExperimentRunner(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=flag_mode,
+ params=params,
+ model_dir=model_dir,
+ run_post_eval=run_post_eval).run()
+
+ if 'eval' in flag_mode:
+ self.assertTrue(
+ tf.io.gfile.exists(
+ os.path.join(model_dir,
+ params.trainer.validation_summary_subdir)))
+ if run_post_eval:
+ self.assertNotEmpty(logs)
+ else:
+ self.assertEmpty(logs)
+ self.assertNotEmpty(
+ tf.io.gfile.glob(os.path.join(model_dir, 'params.yaml')))
+ if flag_mode == 'eval':
+ return
+ self.assertNotEmpty(
+ tf.io.gfile.glob(os.path.join(model_dir, 'checkpoint')))
+ # Tests continuous evaluation.
+ _, logs = train_lib.OrbitExperimentRunner(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode='continuous_eval',
+ params=params,
+ model_dir=model_dir,
+ run_post_eval=run_post_eval).run()
+
@combinations.generate(
combinations.combine(
distribution_strategy=[
@@ -148,12 +203,12 @@ def build_losses(labels, model_outputs, aux_losses=None):
task.build_losses = build_losses
with self.assertRaises(RuntimeError):
- train_lib.run_experiment(
+ train_lib.OrbitExperimentRunner(
distribution_strategy=distribution_strategy,
task=task,
mode=flag_mode,
params=params,
- model_dir=model_dir)
+ model_dir=model_dir).run()
@combinations.generate(
combinations.combine(
@@ -194,12 +249,12 @@ def build_losses(labels, model_outputs, aux_losses=None):
task.build_losses = build_losses
- model, _ = train_lib.run_experiment(
+ model, _ = train_lib.OrbitExperimentRunner(
distribution_strategy=distribution_strategy,
task=task,
mode=flag_mode,
params=params,
- model_dir=model_dir)
+ model_dir=model_dir).run()
after_weights = model.get_weights()
for left, right in zip(before_weights, after_weights):
self.assertAllEqual(left, right)
diff --git a/official/core/train_utils.py b/official/core/train_utils.py
index 45aab1af405..40cb5163c30 100644
--- a/official/core/train_utils.py
+++ b/official/core/train_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Training utils."""
-import copy
+
import dataclasses
import inspect
import json
@@ -23,10 +23,12 @@
from absl import logging
import gin
+import numpy as np
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=g-direct-tensorflow-import
+from tensorflow.python.framework import ops
from tensorflow.python.framework.convert_to_constants import convert_variables_to_constants_v2_as_graph
# pylint: enable=g-direct-tensorflow-import
from official.core import base_task
@@ -36,6 +38,9 @@
from official.modeling import hyperparams
+BEST_CHECKPOINT_NAME = 'best_ckpt'
+
+
def get_leaf_nested_dict(d: Dict[str, Any], keys: List[str]) -> Dict[str, Any]:
"""Get leaf from a dictionary with arbitrary depth with a list of keys.
@@ -82,6 +87,29 @@ def cast_leaf_nested_dict(d: Dict[str, Any],
return d
+def _filter_leaf_nested_dict(
+ d: Dict[str, Any], predicate: Callable[[Any], bool]
+) -> Dict[str, Any]:
+ """Filters the leaves of a dictionary with arbitrary depth in place.
+
+ Args:
+ d: The dictionary to extract value from.
+ predicate: A function that will be called on every leave item. When the
+ function returns True the leave will be kept. Otherwise the leave will be
+ dropped.
+
+ Returns:
+ A new dictionray with filtered result.
+ """
+ result = {}
+ for key, value in d.items():
+ if isinstance(value, dict):
+ result[key] = _filter_leaf_nested_dict(value, predicate)
+ elif predicate(value):
+ result[key] = value
+ return result
+
+
def maybe_create_best_ckpt_exporter(params: config_definitions.ExperimentConfig,
data_dir: str) -> Any:
"""Maybe create a BestCheckpointExporter object, according to the config."""
@@ -138,7 +166,7 @@ def _get_checkpoint_manager(self, checkpoint):
checkpoint,
directory=self._export_dir,
max_to_keep=1,
- checkpoint_name='best_ckpt')
+ checkpoint_name=BEST_CHECKPOINT_NAME)
return self._checkpoint_manager
@@ -187,7 +215,11 @@ def _new_metric_is_better(self, old_logs, new_logs):
def export_best_eval_metric(self, eval_logs, global_step):
"""Export evaluation results of the best checkpoint into a json file."""
- eval_logs_ext = copy.copy(eval_logs)
+ # eval_log_ext may contains non-scalar tensors, such as image data when
+ # `allow_image_summary` is True. Here we only keep scalar tensors.
+ eval_logs_ext = _filter_leaf_nested_dict(
+ eval_logs, lambda x: tf.rank(x) <= 1
+ )
eval_logs_ext['best_ckpt_global_step'] = global_step
eval_logs_ext = cast_leaf_nested_dict(
eval_logs_ext, lambda x: float(orbit.utils.get_value(x)))
@@ -211,7 +243,7 @@ def best_ckpt_path(self):
def create_optimizer(task: base_task.Task,
params: config_definitions.ExperimentConfig
- ) -> tf.keras.optimizers.Optimizer:
+ ) -> tf_keras.optimizers.Optimizer:
"""A create optimizer util to be backward compatability with new args."""
if 'dp_config' in inspect.signature(task.create_optimizer).parameters:
dp_config = None
@@ -260,6 +292,7 @@ class ParseConfigOptions:
tpu: str = ''
tf_data_service: str = ''
params_override: str = ''
+ strict_override: bool = True
def __contains__(self, name):
return name in dataclasses.asdict(self)
@@ -298,9 +331,13 @@ def base_experiment(self):
def parse_config_file(self, params):
"""Override the configs of params from the config_file."""
+ is_strict = True
+ if isinstance(self._flags_obj, ParseConfigOptions):
+ is_strict = self._flags_obj.strict_override
for config_file in self._flags_obj.config_file or []:
params = hyperparams.override_params_dict(
- params, config_file, is_strict=True)
+ params, config_file, is_strict=is_strict
+ )
return params
def parse_runtime(self, params):
@@ -331,14 +368,28 @@ def parse_data_service(self, params):
return params
def parse_params_override(self, params):
+ """Overrides params from the --params_override flag.
+
+ Args:
+ params: A ParamsDict object to be overridden.
+
+ Returns:
+ The overridden ParamsDict object.
+ """
# Get the second level of override from `--params_override`.
# `--params_override` is typically used as a further override over the
# template. For example, one may define a particular template for training
# ResNet50 on ImageNet in a config file and pass it via `--config_file`,
# then define different learning rates and pass it via `--params_override`.
if self._flags_obj.params_override:
+ is_strict = True
+ if isinstance(self._flags_obj, ParseConfigOptions):
+ is_strict = self._flags_obj.strict_override
params = hyperparams.override_params_dict(
- params, self._flags_obj.params_override, is_strict=True)
+ params,
+ self._flags_obj.params_override,
+ is_strict=is_strict,
+ )
return params
@@ -439,7 +490,7 @@ def remove_ckpts(model_dir):
tf.io.gfile.remove(file_to_remove)
-def write_model_params(model: Union[tf.Module, tf.keras.Model],
+def write_model_params(model: Union[tf.Module, tf_keras.Model],
output_path: str) -> None:
"""Writes the model parameters and shapes to a file.
@@ -457,7 +508,7 @@ def write_model_params(model: Union[tf.Module, tf.keras.Model],
def try_count_params(
- model: Union[tf.Module, tf.keras.Model],
+ model: Union[tf.Module, tf_keras.Model],
trainable_only: bool = False):
"""Count the number of parameters if model is possible.
@@ -487,7 +538,7 @@ def try_count_params(
return total_params
-def try_count_flops(model: Union[tf.Module, tf.keras.Model],
+def try_count_flops(model: Union[tf.Module, tf_keras.Model],
inputs_kwargs: Optional[Dict[str, Any]] = None,
output_path: Optional[str] = None):
"""Counts and returns model FLOPs.
@@ -535,3 +586,44 @@ def try_count_flops(model: Union[tf.Module, tf.keras.Model],
'reached before this run.', e)
return None
return None
+
+
+@ops.RegisterStatistics('Einsum', 'flops')
+def _einsum_flops(graph, node):
+ """Calculates the compute resources needed for Einsum."""
+ assert len(node.input) == 2
+ x_shape = tf.compat.v1.graph_util.tensor_shape_from_node_def_name(
+ graph, node.input[0])
+ y_shape = tf.compat.v1.graph_util.tensor_shape_from_node_def_name(
+ graph, node.input[1])
+ x_shape.assert_is_fully_defined()
+ y_shape.assert_is_fully_defined()
+ x_shape = x_shape.as_list()
+ y_shape = y_shape.as_list()
+ equation = str(node.attr['equation'])
+ equation = (
+ equation.replace('s:', '')
+ .replace('"', '')
+ .replace(' ', '')
+ .replace('\n', '')
+ )
+ x_str = equation.split(',')[0]
+ y_r_str = equation.split(',')[1]
+ y_str = y_r_str.split('->')[0]
+ r_str = y_r_str.split('->')[1]
+ shape_dic = {}
+ contracted = set()
+ for indice in x_str + y_str:
+ if indice in x_str:
+ indice_dim = x_shape[x_str.find(indice)]
+ elif indice in y_str:
+ indice_dim = y_shape[y_str.find(indice)]
+ else:
+ raise ValueError('indice {} not found in inputs'.format(indice))
+ shape_dic[indice] = indice_dim
+ if indice not in r_str:
+ contracted.add(indice)
+ madds = np.prod([shape_dic[indice] for indice in r_str]) * (
+ np.prod([shape_dic[indice] for indice in contracted]))
+ flops = 2 * madds
+ return ops.OpStats('flops', flops)
diff --git a/official/core/train_utils_test.py b/official/core/train_utils_test.py
index dbc49d2b7d5..8149de928aa 100644
--- a/official/core/train_utils_test.py
+++ b/official/core/train_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,8 +17,10 @@
import os
import pprint
+from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+import yaml
from official.core import exp_factory
from official.core import test_utils
@@ -47,7 +49,7 @@ def foo():
return experiment_config
-class TrainUtilsTest(tf.test.TestCase):
+class TrainUtilsTest(tf.test.TestCase, parameterized.TestCase):
def test_get_leaf_nested_dict(self):
d = {'a': {'i': {'x': 5}}}
@@ -138,6 +140,87 @@ def test_construct_experiment_from_flags(self):
self.assertEqual(params_from_obj.trainer.train_steps, 10)
self.assertEqual(params_from_obj.trainer.validation_steps, 11)
+ @parameterized.named_parameters(
+ dict(
+ testcase_name='strict_with_extra',
+ strict_override=True,
+ has_extra=True,
+ expect_error=True,
+ ),
+ dict(
+ testcase_name='non_strict_with_extra',
+ strict_override=False,
+ has_extra=True,
+ expect_error=False,
+ ),
+ dict(
+ testcase_name='strict_no_extra',
+ strict_override=True,
+ has_extra=False,
+ expect_error=False,
+ ),
+ dict(
+ testcase_name='non_strict_no_extra',
+ strict_override=False,
+ has_extra=False,
+ expect_error=False,
+ ),
+ )
+ def test_parse_configuration_strict_override_file(
+ self, strict_override, has_extra, expect_error
+ ):
+ tempdir = self.create_tempdir().full_path
+ config_path = os.path.join(tempdir, 'config.yaml')
+ override_config = {
+ 'task': {
+ 'model': {
+ 'model_id': 'override',
+ },
+ },
+ 'trainer': {
+ 'train_steps': 500,
+ },
+ }
+ if has_extra:
+ override_config['task']['extra_key'] = 'extra_value'
+
+ with open(config_path, 'w') as f:
+ yaml.dump(override_config, f)
+
+ options = train_utils.ParseConfigOptions(
+ experiment='foo',
+ config_file=[config_path],
+ strict_override=strict_override,
+ )
+
+ if expect_error:
+ with self.assertRaises(KeyError):
+ train_utils.parse_configuration(options)
+ else:
+ params = train_utils.parse_configuration(options)
+ self.assertEqual(params.task.model.model_id, 'override')
+ self.assertEqual(params.trainer.train_steps, 500)
+ if has_extra:
+ self.assertTrue(hasattr(params.task, 'extra_key'))
+ self.assertEqual(params.task.extra_key, 'extra_value')
+ else:
+ self.assertFalse(hasattr(params.task, 'extra_key'))
+
+ def test_parse_configuration_strict_override_string(self):
+ # Test non-strict loading with params_override string
+ options_non_strict_override = train_utils.ParseConfigOptions(
+ experiment='foo',
+ config_file=[],
+ params_override='task.another_extra=test,trainer.train_steps=100',
+ strict_override=False,
+ )
+ params_override = train_utils.parse_configuration(
+ options_non_strict_override
+ )
+ self.assertEqual(params_override.trainer.train_steps, 100)
+ self.assertTrue(hasattr(params_override.task, 'another_extra'))
+ self.assertEqual(params_override.task.another_extra, 'test')
+
class BestCheckpointExporterTest(tf.test.TestCase):
@@ -193,6 +276,23 @@ def test_export_best_eval_metric(self):
metric,
{'test_metric': {'metric_1': 5.0}, 'best_ckpt_global_step': 100.0})
+ def test_export_best_eval_metric_skips_non_scalar_values(self):
+ model_dir = self.create_tempdir().full_path
+ metric_name = 'test_metric|metric_1'
+ exporter = train_utils.BestCheckpointExporter(model_dir, metric_name,
+ 'higher')
+ image = tf.zeros(shape=[16, 8, 1])
+ eval_logs = {'test_metric': {'metric_1': 5.0, 'image': image}}
+
+ exporter.export_best_eval_metric(eval_logs, 100)
+
+ with tf.io.gfile.GFile(os.path.join(model_dir, 'info.json'),
+ 'rb') as reader:
+ metric = json.loads(reader.read())
+ self.assertAllEqual(
+ metric,
+ {'test_metric': {'metric_1': 5.0}, 'best_ckpt_global_step': 100.0})
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/legacy/__init__.py b/official/legacy/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/__init__.py
+++ b/official/legacy/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/albert/__init__.py b/official/legacy/albert/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/albert/__init__.py
+++ b/official/legacy/albert/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/albert/configs.py b/official/legacy/albert/configs.py
index 7baf693aee8..b45243b264a 100644
--- a/official/legacy/albert/configs.py
+++ b/official/legacy/albert/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/bert/README.md b/official/legacy/bert/README.md
index cf4062a6dc3..82bd5e9d9c5 100644
--- a/official/legacy/bert/README.md
+++ b/official/legacy/bert/README.md
@@ -1,6 +1,6 @@
# BERT (Bidirectional Encoder Representations from Transformers)
-**WARNING**: We are on the way to deprecate most of the code in this directory.
+**WARNING**: We are on the way to deprecating most of the code in this directory.
Please see
[this link](../g3doc/tutorials/bert_new.md)
for the new tutorial and use the new code in `nlp/modeling`. This README is
@@ -55,7 +55,7 @@ The new checkpoints are:**
* **[`BERT-Base, Multilingual Cased`](https://storage.googleapis.com/cloud-tpu-checkpoints/bert/keras_bert/multi_cased_L-12_H-768_A-12.tar.gz)**:
104 languages, 12-layer, 768-hidden, 12-heads, 110M parameters
-We recommend to host checkpoints on Google Cloud storage buckets when you use
+We recommend to host checkpoints on Google Cloud Storage buckets when you use
Cloud GPU/TPU.
### Restoring from Checkpoints
@@ -127,7 +127,7 @@ pip install tf-nightly
There is no change to generate pre-training data. Please use the script
[`../data/create_pretraining_data.py`](../data/create_pretraining_data.py)
-which is essentially branched from [BERT research repo](https://github.com/google-research/bert)
+which is essentially branched from the [BERT research repo](https://github.com/google-research/bert)
to get processed pre-training data and it adapts to TF2 symbols and python3
compatibility.
@@ -155,7 +155,7 @@ To prepare the fine-tuning data for final model training, use the
[`../data/create_finetuning_data.py`](../data/create_finetuning_data.py) script.
Resulting datasets in `tf_record` format and training meta data should be later
passed to training or evaluation scripts. The task-specific arguments are
-described in following sections:
+described in the following sections:
* GLUE
@@ -163,7 +163,7 @@ Users can download the
[GLUE data](https://gluebenchmark.com/tasks) by running
[this script](https://gist.github.com/W4ngatang/60c2bdb54d156a41194446737ce03e2e)
and unpack it to some directory `$GLUE_DIR`.
-Also, users can download [Pretrained Checkpoint](#access-to-pretrained-checkpoints) and locate on some directory `$BERT_DIR` instead of using checkpoints on Google Cloud Storage.
+Also, users can download [Pretrained Checkpoint](#access-to-pretrained-checkpoints) and locate it on some directory `$BERT_DIR` instead of using checkpoints on Google Cloud Storage.
```shell
export GLUE_DIR=~/glue
@@ -271,12 +271,12 @@ python run_classifier.py \
```
Alternatively, instead of specifying `init_checkpoint`, you can specify
-`hub_module_url` to employ a pretraind BERT hub module, e.g.,
+`hub_module_url` to employ a pre-trained BERT hub module, e.g.,
` --hub_module_url=https://tfhub.dev/tensorflow/bert_en_uncased_L-24_H-1024_A-16/1`.
After training a model, to get predictions from the classifier, you can set the
`--mode=predict` and offer the test set tfrecords to `--eval_data_path`.
-Output will be created in file called test_results.tsv in the output folder.
+The output will be created in file called test_results.tsv in the output folder.
Each line will contain output for each sample, columns are the class
probabilities.
@@ -291,7 +291,7 @@ python run_classifier.py \
--distribution_strategy=mirrored
```
-To use TPU, you only need to switch distribution strategy type to `tpu` with TPU
+To use TPU, you only need to switch the distribution strategy type to `tpu` with TPU
information and use remote storage for model checkpoints.
```shell
@@ -325,7 +325,7 @@ and callbacks will not be called inside the loop.
### SQuAD 1.1
The Stanford Question Answering Dataset (SQuAD) is a popular question answering
-benchmark dataset. See more in [SQuAD website](https://rajpurkar.github.io/SQuAD-explorer/).
+benchmark dataset. See more on [SQuAD website](https://rajpurkar.github.io/SQuAD-explorer/).
We use the `BERT-Large` (uncased_L-24_H-1024_A-16) as an example throughout the
workflow.
@@ -353,14 +353,14 @@ python run_squad.py \
--distribution_strategy=mirrored
```
-Similarily, you can replace `init_checkpoint` FLAG with `hub_module_url` to
+Similarly, you can replace `init_checkpoint` FLAG with `hub_module_url` to
specify a hub module path.
`run_squad.py` writes the prediction for `--predict_file` by default. If you set
the `--model=predict` and offer the SQuAD test data, the scripts will generate
the prediction json file.
-To use TPU, you need switch distribution strategy type to `tpu` with TPU
+To use TPU, you need to switch the distribution strategy type to `tpu` with TPU
information.
```shell
diff --git a/official/legacy/bert/__init__.py b/official/legacy/bert/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/legacy/bert/__init__.py
+++ b/official/legacy/bert/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/bert/bert_models.py b/official/legacy/bert/bert_models.py
index 21d095174cc..1c7dc5c6df4 100644
--- a/official/legacy/bert/bert_models.py
+++ b/official/legacy/bert/bert_models.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""BERT models that are compatible with TF 2.0."""
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_hub as hub
from official.legacy.albert import configs as albert_configs
from official.legacy.bert import configs
@@ -24,7 +24,7 @@
from official.nlp.modeling import networks
-class BertPretrainLossAndMetricLayer(tf.keras.layers.Layer):
+class BertPretrainLossAndMetricLayer(tf_keras.layers.Layer):
"""Returns layer that computes custom loss and metrics for pretraining."""
def __init__(self, vocab_size, **kwargs):
@@ -38,7 +38,7 @@ def _add_metrics(self, lm_output, lm_labels, lm_label_weights,
lm_example_loss, sentence_output, sentence_labels,
next_sentence_loss):
"""Adds metrics."""
- masked_lm_accuracy = tf.keras.metrics.sparse_categorical_accuracy(
+ masked_lm_accuracy = tf_keras.metrics.sparse_categorical_accuracy(
lm_labels, lm_output)
numerator = tf.reduce_sum(masked_lm_accuracy * lm_label_weights)
denominator = tf.reduce_sum(lm_label_weights) + 1e-5
@@ -49,7 +49,7 @@ def _add_metrics(self, lm_output, lm_labels, lm_label_weights,
self.add_metric(lm_example_loss, name='lm_example_loss', aggregation='mean')
if sentence_labels is not None:
- next_sentence_accuracy = tf.keras.metrics.sparse_categorical_accuracy(
+ next_sentence_accuracy = tf_keras.metrics.sparse_categorical_accuracy(
sentence_labels, sentence_output)
self.add_metric(
next_sentence_accuracy,
@@ -70,7 +70,7 @@ def call(self,
lm_label_weights = tf.cast(lm_label_weights, tf.float32)
lm_output_logits = tf.cast(lm_output_logits, tf.float32)
- lm_prediction_losses = tf.keras.losses.sparse_categorical_crossentropy(
+ lm_prediction_losses = tf_keras.losses.sparse_categorical_crossentropy(
lm_label_ids, lm_output_logits, from_logits=True)
lm_numerator_loss = tf.reduce_sum(lm_prediction_losses * lm_label_weights)
lm_denominator_loss = tf.reduce_sum(lm_label_weights)
@@ -79,7 +79,7 @@ def call(self,
if sentence_labels is not None:
sentence_output_logits = tf.cast(sentence_output_logits, tf.float32)
- sentence_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ sentence_loss = tf_keras.losses.sparse_categorical_crossentropy(
sentence_labels, sentence_output_logits, from_logits=True)
sentence_loss = tf.reduce_mean(sentence_loss)
loss = mask_label_loss + sentence_loss
@@ -123,7 +123,7 @@ def get_transformer_encoder(bert_config,
type_vocab_size=bert_config.type_vocab_size,
hidden_size=bert_config.hidden_size,
max_seq_length=bert_config.max_position_embeddings,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
dropout_rate=bert_config.hidden_dropout_prob,
)
@@ -133,7 +133,7 @@ def get_transformer_encoder(bert_config,
intermediate_activation=tf_utils.get_activation(bert_config.hidden_act),
dropout_rate=bert_config.hidden_dropout_prob,
attention_dropout_rate=bert_config.attention_probs_dropout_prob,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
)
kwargs = dict(
@@ -141,7 +141,7 @@ def get_transformer_encoder(bert_config,
hidden_cfg=hidden_cfg,
num_hidden_instances=bert_config.num_hidden_layers,
pooled_output_dim=bert_config.hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range))
# Relies on gin configuration to define the Transformer encoder arguments.
@@ -159,7 +159,7 @@ def get_transformer_encoder(bert_config,
max_sequence_length=bert_config.max_position_embeddings,
type_vocab_size=bert_config.type_vocab_size,
embedding_width=bert_config.embedding_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range))
if isinstance(bert_config, albert_configs.AlbertConfig):
return networks.AlbertEncoder(**kwargs)
@@ -192,32 +192,32 @@ def pretrain_model(bert_config,
save weights after pretraining, and (3) optional core `BertPretrainer`
object if argument `return_core_pretrainer_model` is True.
"""
- input_word_ids = tf.keras.layers.Input(
+ input_word_ids = tf_keras.layers.Input(
shape=(seq_length,), name='input_word_ids', dtype=tf.int32)
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(seq_length,), name='input_mask', dtype=tf.int32)
- input_type_ids = tf.keras.layers.Input(
+ input_type_ids = tf_keras.layers.Input(
shape=(seq_length,), name='input_type_ids', dtype=tf.int32)
- masked_lm_positions = tf.keras.layers.Input(
+ masked_lm_positions = tf_keras.layers.Input(
shape=(max_predictions_per_seq,),
name='masked_lm_positions',
dtype=tf.int32)
- masked_lm_ids = tf.keras.layers.Input(
+ masked_lm_ids = tf_keras.layers.Input(
shape=(max_predictions_per_seq,), name='masked_lm_ids', dtype=tf.int32)
- masked_lm_weights = tf.keras.layers.Input(
+ masked_lm_weights = tf_keras.layers.Input(
shape=(max_predictions_per_seq,),
name='masked_lm_weights',
dtype=tf.int32)
if use_next_sentence_label:
- next_sentence_labels = tf.keras.layers.Input(
+ next_sentence_labels = tf_keras.layers.Input(
shape=(1,), name='next_sentence_labels', dtype=tf.int32)
else:
next_sentence_labels = None
transformer_encoder = get_transformer_encoder(bert_config, seq_length)
if initializer is None:
- initializer = tf.keras.initializers.TruncatedNormal(
+ initializer = tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range)
pretrainer_model = models.BertPretrainer(
network=transformer_encoder,
@@ -247,7 +247,7 @@ def pretrain_model(bert_config,
if use_next_sentence_label:
inputs['next_sentence_labels'] = next_sentence_labels
- keras_model = tf.keras.Model(inputs=inputs, outputs=output_loss)
+ keras_model = tf_keras.Model(inputs=inputs, outputs=output_loss)
if return_core_pretrainer_model:
return keras_model, transformer_encoder, pretrainer_model
else:
@@ -274,23 +274,23 @@ def squad_model(bert_config,
(2) the core BERT transformer encoder.
"""
if initializer is None:
- initializer = tf.keras.initializers.TruncatedNormal(
+ initializer = tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range)
if not hub_module_url:
bert_encoder = get_transformer_encoder(bert_config, max_seq_length)
return models.BertSpanLabeler(
network=bert_encoder, initializer=initializer), bert_encoder
- input_word_ids = tf.keras.layers.Input(
+ input_word_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_word_ids')
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_mask')
- input_type_ids = tf.keras.layers.Input(
+ input_type_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_type_ids')
core_model = hub.KerasLayer(hub_module_url, trainable=hub_module_trainable)
pooled_output, sequence_output = core_model(
[input_word_ids, input_mask, input_type_ids])
- bert_encoder = tf.keras.Model(
+ bert_encoder = tf_keras.Model(
inputs={
'input_word_ids': input_word_ids,
'input_mask': input_mask,
@@ -330,7 +330,7 @@ def classifier_model(bert_config,
if final_layer_initializer is not None:
initializer = final_layer_initializer
else:
- initializer = tf.keras.initializers.TruncatedNormal(
+ initializer = tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range)
if not hub_module_url:
@@ -342,21 +342,21 @@ def classifier_model(bert_config,
dropout_rate=bert_config.hidden_dropout_prob,
initializer=initializer), bert_encoder
- input_word_ids = tf.keras.layers.Input(
+ input_word_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_word_ids')
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_mask')
- input_type_ids = tf.keras.layers.Input(
+ input_type_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_type_ids')
bert_model = hub.KerasLayer(hub_module_url, trainable=hub_module_trainable)
pooled_output, _ = bert_model([input_word_ids, input_mask, input_type_ids])
- output = tf.keras.layers.Dropout(rate=bert_config.hidden_dropout_prob)(
+ output = tf_keras.layers.Dropout(rate=bert_config.hidden_dropout_prob)(
pooled_output)
- output = tf.keras.layers.Dense(
+ output = tf_keras.layers.Dense(
num_labels, kernel_initializer=initializer, name='output')(
output)
- return tf.keras.Model(
+ return tf_keras.Model(
inputs={
'input_word_ids': input_word_ids,
'input_mask': input_mask,
diff --git a/official/legacy/bert/bert_models_test.py b/official/legacy/bert/bert_models_test.py
index e64c013c40d..a28ad048bd6 100644
--- a/official/legacy/bert/bert_models_test.py
+++ b/official/legacy/bert/bert_models_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -12,7 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import bert_models
from official.legacy.bert import configs as bert_configs
@@ -43,7 +43,7 @@ def test_pretrain_model(self):
max_predictions_per_seq=2,
initializer=None,
use_next_sentence_label=True)
- self.assertIsInstance(model, tf.keras.Model)
+ self.assertIsInstance(model, tf_keras.Model)
self.assertIsInstance(encoder, networks.BertEncoder)
# model has one scalar output: loss value.
@@ -64,8 +64,8 @@ def test_squad_model(self):
initializer=None,
hub_module_url=None,
hub_module_trainable=None)
- self.assertIsInstance(model, tf.keras.Model)
- self.assertIsInstance(core_model, tf.keras.Model)
+ self.assertIsInstance(model, tf_keras.Model)
+ self.assertIsInstance(core_model, tf_keras.Model)
# Expect two output from model: start positions and end positions
self.assertIsInstance(model.output, list)
@@ -87,8 +87,8 @@ def test_classifier_model(self):
final_layer_initializer=None,
hub_module_url=None,
hub_module_trainable=None)
- self.assertIsInstance(model, tf.keras.Model)
- self.assertIsInstance(core_model, tf.keras.Model)
+ self.assertIsInstance(model, tf_keras.Model)
+ self.assertIsInstance(core_model, tf_keras.Model)
# model has one classification output with num_labels=3.
self.assertEqual(model.output.shape.as_list(), [None, 3])
diff --git a/official/legacy/bert/common_flags.py b/official/legacy/bert/common_flags.py
index 32ad7059f04..014df62fea0 100644
--- a/official/legacy/bert/common_flags.py
+++ b/official/legacy/bert/common_flags.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Defining common flags used across all BERT models/applications."""
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.utils import hyperparams_flags
from official.utils.flags import core as flags_core
diff --git a/official/legacy/bert/configs.py b/official/legacy/bert/configs.py
index bbded193265..8c4f71d5bf2 100644
--- a/official/legacy/bert/configs.py
+++ b/official/legacy/bert/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import json
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class BertConfig(object):
diff --git a/official/legacy/bert/export_tfhub.py b/official/legacy/bert/export_tfhub.py
index 69dd49865e2..7c6c1e6fcb8 100644
--- a/official/legacy/bert/export_tfhub.py
+++ b/official/legacy/bert/export_tfhub.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,11 +20,10 @@
from typing import Text
-# Import libraries
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import bert_models
from official.legacy.bert import configs
@@ -45,7 +44,7 @@
"What kind of BERT model to export.")
-def create_bert_model(bert_config: configs.BertConfig) -> tf.keras.Model:
+def create_bert_model(bert_config: configs.BertConfig) -> tf_keras.Model:
"""Creates a BERT keras core model from BERT configuration.
Args:
@@ -55,11 +54,11 @@ def create_bert_model(bert_config: configs.BertConfig) -> tf.keras.Model:
A keras model.
"""
# Adds input layers just as placeholders.
- input_word_ids = tf.keras.layers.Input(
+ input_word_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name="input_word_ids")
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name="input_mask")
- input_type_ids = tf.keras.layers.Input(
+ input_type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name="input_type_ids")
transformer_encoder = bert_models.get_transformer_encoder(
bert_config, sequence_length=None)
@@ -67,7 +66,7 @@ def create_bert_model(bert_config: configs.BertConfig) -> tf.keras.Model:
[input_word_ids, input_mask, input_type_ids])
# To keep consistent with legacy hub modules, the outputs are
# "pooled_output" and "sequence_output".
- return tf.keras.Model(
+ return tf_keras.Model(
inputs=[input_word_ids, input_mask, input_type_ids],
outputs=[pooled_output, sequence_output]), transformer_encoder
@@ -77,7 +76,7 @@ def export_bert_tfhub(bert_config: configs.BertConfig,
hub_destination: Text,
vocab_file: Text,
do_lower_case: bool = None):
- """Restores a tf.keras.Model and saves for TF-Hub."""
+ """Restores a tf_keras.Model and saves for TF-Hub."""
# If do_lower_case is not explicit, default to checking whether "uncased" is
# in the vocab file name
if do_lower_case is None:
@@ -99,7 +98,7 @@ def export_bert_squad_tfhub(bert_config: configs.BertConfig,
hub_destination: Text,
vocab_file: Text,
do_lower_case: bool = None):
- """Restores a tf.keras.Model for BERT with SQuAD and saves for TF-Hub."""
+ """Restores a tf_keras.Model for BERT with SQuAD and saves for TF-Hub."""
# If do_lower_case is not explicit, default to checking whether "uncased" is
# in the vocab file name
if do_lower_case is None:
diff --git a/official/legacy/bert/export_tfhub_test.py b/official/legacy/bert/export_tfhub_test.py
index 68146fb5814..45c449c8c43 100644
--- a/official/legacy/bert/export_tfhub_test.py
+++ b/official/legacy/bert/export_tfhub_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_hub as hub
from official.legacy.bert import configs
@@ -94,9 +94,9 @@ def _dropout_mean_stddev(training, num_runs=20):
self.assertGreater(_dropout_mean_stddev(training=True), 1e-3)
# Test propagation of seq_length in shape inference.
- input_word_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_mask = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_type_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_word_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_mask = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_type_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
pooled_output, sequence_output = hub_layer(
[input_word_ids, input_mask, input_type_ids])
self.assertEqual(pooled_output.shape.as_list(), [None, hidden_size])
diff --git a/official/legacy/bert/input_pipeline.py b/official/legacy/bert/input_pipeline.py
index 045f16ce76b..f5064e20cec 100644
--- a/official/legacy/bert/input_pipeline.py
+++ b/official/legacy/bert/input_pipeline.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""BERT model input pipelines."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def decode_record(record, name_to_features):
diff --git a/official/legacy/bert/model_saving_utils.py b/official/legacy/bert/model_saving_utils.py
index 6a0d7074972..abd22789e20 100644
--- a/official/legacy/bert/model_saving_utils.py
+++ b/official/legacy/bert/model_saving_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,11 +17,11 @@
import os
import typing
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def export_bert_model(model_export_path: typing.Text,
- model: tf.keras.Model,
+ model: tf_keras.Model,
checkpoint_dir: typing.Optional[typing.Text] = None,
restore_model_using_load_weights: bool = False) -> None:
"""Export BERT model for serving which does not include the optimizer.
@@ -45,8 +45,8 @@ def export_bert_model(model_export_path: typing.Text,
"""
if not model_export_path:
raise ValueError('model_export_path must be specified.')
- if not isinstance(model, tf.keras.Model):
- raise ValueError('model must be a tf.keras.Model object.')
+ if not isinstance(model, tf_keras.Model):
+ raise ValueError('model must be a tf_keras.Model object.')
if checkpoint_dir:
if restore_model_using_load_weights:
diff --git a/official/legacy/bert/model_training_utils.py b/official/legacy/bert/model_training_utils.py
index f7c8e443be3..987368c9eac 100644
--- a/official/legacy/bert/model_training_utils.py
+++ b/official/legacy/bert/model_training_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
import tempfile
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.util import deprecation
from official.common import distribute_utils
from official.modeling import grad_utils
@@ -222,7 +222,7 @@ def run_customized_training_loop(
strategy, model_fn, loss_fn, model_dir, steps_per_epoch, train_input_fn
]
- steps_between_evals = int(steps_per_epoch / num_eval_per_epoch)
+ steps_between_evals = int(steps_per_epoch / num_eval_per_epoch) # pyrefly: ignore[unsupported-operation]
if [arg for arg in required_arguments if arg is None]:
raise ValueError('`strategy`, `model_fn`, `loss_fn`, `model_dir`, '
'`steps_per_epoch` and `train_input_fn` are required '
@@ -259,14 +259,14 @@ def run_customized_training_loop(
raise ValueError(
'if `metric_fn` is specified, metric_fn must be a callable.')
- total_training_steps = steps_per_epoch * epochs
+ total_training_steps = steps_per_epoch * epochs # pyrefly: ignore[unsupported-operation]
train_iterator = _get_input_iterator(train_input_fn, strategy)
- eval_loss_metric = tf.keras.metrics.Mean('training_loss', dtype=tf.float32)
+ eval_loss_metric = tf_keras.metrics.Mean('training_loss', dtype=tf.float32)
with distribute_utils.get_strategy_scope(strategy):
# To correctly place the model weights on accelerators,
# model and optimizer should be created in scope.
- model, sub_model = model_fn()
+ model, sub_model = model_fn() # pyrefly: ignore[not-callable]
if not hasattr(model, 'optimizer'):
raise ValueError('User should set optimizer attribute to model '
'inside `model_fn`.')
@@ -274,7 +274,7 @@ def run_customized_training_loop(
raise ValueError('sub_model_export_name is specified as %s, but '
'sub_model is None.' % sub_model_export_name)
- callback_list = tf.keras.callbacks.CallbackList(
+ callback_list = tf_keras.callbacks.CallbackList(
callbacks=custom_callbacks, model=model)
optimizer = model.optimizer
@@ -287,20 +287,20 @@ def run_customized_training_loop(
checkpoint.read(init_checkpoint).assert_existing_objects_matched()
logging.info('Loading from checkpoint file completed')
- train_loss_metric = tf.keras.metrics.Mean('training_loss', dtype=tf.float32)
+ train_loss_metric = tf_keras.metrics.Mean('training_loss', dtype=tf.float32)
eval_metrics = metric_fn() if metric_fn else []
if not isinstance(eval_metrics, list):
eval_metrics = [eval_metrics]
# If evaluation is required, make a copy of metric as it will be used by
# both train and evaluation.
train_metrics = [
- metric.__class__.from_config(metric.get_config())
+ metric.__class__.from_config(metric.get_config()) # pyrefly: ignore[missing-attribute]
for metric in eval_metrics
]
# Create summary writers
if _should_export_summary(strategy):
- summary_dir = os.path.join(model_dir, 'summaries')
+ summary_dir = os.path.join(model_dir, 'summaries') # pyrefly: ignore[no-matching-overload]
else:
# In multi worker training we need every worker to write summary, because
# variables can trigger synchronization on read and synchronization needs
@@ -326,12 +326,12 @@ def _replicated_step(inputs):
inputs, labels = inputs
with tf.GradientTape() as tape:
model_outputs = model(inputs, training=True)
- loss = loss_fn(labels, model_outputs)
+ loss = loss_fn(labels, model_outputs) # pyrefly: ignore[not-callable]
# Raw loss is used for reporting in metrics/logs.
raw_loss = loss
if scale_loss:
# Scales down the loss for gradients to be invariant from replicas.
- loss = loss / strategy.num_replicas_in_sync
+ loss = loss / strategy.num_replicas_in_sync # pyrefly: ignore[missing-attribute]
if explicit_allreduce:
grad_utils.minimize_using_explicit_allreduce(tape, optimizer, loss,
@@ -340,7 +340,7 @@ def _replicated_step(inputs):
post_allreduce_callbacks,
allreduce_bytes_per_pack)
else:
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
with tape:
scaled_loss = optimizer.get_scaled_loss(loss)
scaled_grads = tape.gradient(scaled_loss, training_vars)
@@ -370,7 +370,7 @@ def train_steps(iterator, steps):
'retracing.')
for _ in tf.range(steps):
- strategy.run(_replicated_step, args=(next(iterator),))
+ strategy.run(_replicated_step, args=(next(iterator),)) # pyrefly: ignore[missing-attribute]
def train_single_step(iterator):
"""Performs a distributed training step.
@@ -381,7 +381,7 @@ def train_single_step(iterator):
Raises:
ValueError: Any of the arguments or tensor shapes are invalid.
"""
- strategy.run(_replicated_step, args=(next(iterator),))
+ strategy.run(_replicated_step, args=(next(iterator),)) # pyrefly: ignore[missing-attribute]
def test_step(iterator):
"""Calculates evaluation metrics on distributed devices."""
@@ -392,13 +392,13 @@ def _test_step_fn(inputs):
inputs, labels = inputs
model_outputs = model(inputs, training=False)
for metric in eval_metrics:
- metric.update_state(labels, model_outputs)
+ metric.update_state(labels, model_outputs) # pyrefly: ignore[missing-attribute]
return model_outputs, labels
- outputs, labels = strategy.run(_test_step_fn, args=(next(iterator),))
- outputs = tf.nest.map_structure(strategy.experimental_local_results,
+ outputs, labels = strategy.run(_test_step_fn, args=(next(iterator),)) # pyrefly: ignore[missing-attribute]
+ outputs = tf.nest.map_structure(strategy.experimental_local_results, # pyrefly: ignore[missing-attribute]
outputs)
- labels = tf.nest.map_structure(strategy.experimental_local_results,
+ labels = tf.nest.map_structure(strategy.experimental_local_results, # pyrefly: ignore[missing-attribute]
labels)
return outputs, labels
@@ -422,14 +422,14 @@ def _run_evaluation(current_training_step, test_iterator):
# gather all the logits and labels here to calculate the evaluation loss
# outside.
loss_list, loss_weights = list(), list()
- for _ in range(eval_steps):
+ for _ in range(eval_steps): # pyrefly: ignore[bad-argument-type]
outputs, labels = test_step(test_iterator)
for cur_logits, cur_labels in zip(outputs, labels):
# This is to handle cases when cur_labels is not a single tensor,
# but a dict of tensors.
cur_weight = tf.shape(tf.nest.flatten(cur_labels)[0])[0]
if cur_weight != 0:
- loss_list.append(loss_fn(cur_labels, cur_logits).numpy())
+ loss_list.append(loss_fn(cur_labels, cur_logits).numpy()) # pyrefly: ignore[not-callable]
loss_weights.append(cur_weight)
# The sample_weights are the actual number of examples in each batch,
# a summation of numbers of examples in each replica if using
diff --git a/official/legacy/bert/model_training_utils_test.py b/official/legacy/bert/model_training_utils_test.py
index 298c9282c85..baa40ba5632 100644
--- a/official/legacy/bert/model_training_utils_test.py
+++ b/official/legacy/bert/model_training_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl.testing import parameterized
from absl.testing.absltest import mock
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -94,16 +94,16 @@ def create_model_fn(input_shape, num_classes, use_float16=False):
def _model_fn():
"""A one-layer softmax model suitable for testing."""
- input_layer = tf.keras.layers.Input(shape=input_shape)
- x = tf.keras.layers.Dense(num_classes, activation='relu')(input_layer)
- output_layer = tf.keras.layers.Dense(num_classes, activation='softmax')(x)
- sub_model = tf.keras.models.Model(input_layer, x, name='sub_model')
- model = tf.keras.models.Model(input_layer, output_layer, name='model')
+ input_layer = tf_keras.layers.Input(shape=input_shape)
+ x = tf_keras.layers.Dense(num_classes, activation='relu')(input_layer)
+ output_layer = tf_keras.layers.Dense(num_classes, activation='softmax')(x)
+ sub_model = tf_keras.models.Model(input_layer, x, name='sub_model')
+ model = tf_keras.models.Model(input_layer, output_layer, name='model')
model.add_metric(
tf.reduce_mean(input_layer), name='mean_input', aggregation='mean')
- model.optimizer = tf.keras.optimizers.SGD(learning_rate=0.1, momentum=0.9)
+ model.optimizer = tf_keras.optimizers.SGD(learning_rate=0.1, momentum=0.9)
if use_float16:
- model.optimizer = tf.keras.mixed_precision.LossScaleOptimizer(
+ model.optimizer = tf_keras.mixed_precision.LossScaleOptimizer(
model.optimizer)
return model, sub_model
@@ -112,7 +112,7 @@ def _model_fn():
def metric_fn():
"""Gets a tf.keras metric object."""
- return tf.keras.metrics.CategoricalAccuracy(name='accuracy', dtype=tf.float32)
+ return tf_keras.metrics.CategoricalAccuracy(name='accuracy', dtype=tf.float32)
def summaries_with_matching_keyword(keyword, summary_dir):
@@ -131,7 +131,7 @@ def check_eventfile_for_keyword(keyword, summary_dir):
return any(summaries_with_matching_keyword(keyword, summary_dir))
-class RecordingCallback(tf.keras.callbacks.Callback):
+class RecordingCallback(tf_keras.callbacks.Callback):
def __init__(self):
self.batch_begin = [] # (batch, logs)
@@ -165,7 +165,7 @@ def run_training(self, strategy, model_dir, steps_per_loop, run_eagerly):
model_training_utils.run_customized_training_loop(
strategy=strategy,
model_fn=self._model_fn,
- loss_fn=tf.keras.losses.categorical_crossentropy,
+ loss_fn=tf_keras.losses.categorical_crossentropy,
model_dir=model_dir,
steps_per_epoch=20,
steps_per_loop=steps_per_loop,
@@ -195,7 +195,7 @@ def test_train_eager_single_step(self, distribution):
@combinations.generate(eager_gpu_strategy_combinations())
def test_train_eager_mixed_precision(self, distribution):
model_dir = self.create_tempdir().full_path
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
self._model_fn = create_model_fn(
input_shape=[128], num_classes=3, use_float16=True)
self.run_training(
@@ -255,7 +255,7 @@ def test_train_check_callbacks(self, distribution):
model_training_utils.run_customized_training_loop(
strategy=distribution,
model_fn=self._model_fn,
- loss_fn=tf.keras.losses.categorical_crossentropy,
+ loss_fn=tf_keras.losses.categorical_crossentropy,
model_dir=model_dir,
steps_per_epoch=20,
num_eval_per_epoch=4,
diff --git a/official/legacy/bert/run_classifier.py b/official/legacy/bert/run_classifier.py
index 6e9ea466bce..5dcd6888485 100644
--- a/official/legacy/bert/run_classifier.py
+++ b/official/legacy/bert/run_classifier.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,12 +19,11 @@
import math
import os
-# Import libraries
from absl import app
from absl import flags
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.bert import bert_models
from official.legacy.bert import common_flags
@@ -153,11 +152,11 @@ def _get_classifier_model():
use_float16=common_flags.use_float16())
return classifier_model, core_model
- # tf.keras.losses objects accept optional sample_weight arguments (eg. coming
+ # tf_keras.losses objects accept optional sample_weight arguments (eg. coming
# from the dataset) to compute weighted loss, as used for the regression
# tasks. The classification tasks, using the custom get_loss_fn don't accept
# sample weights though.
- loss_fn = (tf.keras.losses.MeanSquaredError() if is_regression
+ loss_fn = (tf_keras.losses.MeanSquaredError() if is_regression
else get_loss_fn(num_classes))
# Defines evaluation metrics function, which will create metrics in the
@@ -166,12 +165,12 @@ def _get_classifier_model():
metric_fn = custom_metrics
elif is_regression:
metric_fn = functools.partial(
- tf.keras.metrics.MeanSquaredError,
+ tf_keras.metrics.MeanSquaredError,
'mean_squared_error',
dtype=tf.float32)
else:
metric_fn = functools.partial(
- tf.keras.metrics.SparseCategoricalAccuracy,
+ tf_keras.metrics.SparseCategoricalAccuracy,
'accuracy',
dtype=tf.float32)
@@ -230,7 +229,7 @@ def run_keras_compile_fit(model_dir,
steps_per_execution=steps_per_loop)
summary_dir = os.path.join(model_dir, 'summaries')
- summary_callback = tf.keras.callbacks.TensorBoard(summary_dir)
+ summary_callback = tf_keras.callbacks.TensorBoard(summary_dir)
checkpoint = tf.train.Checkpoint(model=bert_model, optimizer=optimizer)
checkpoint_manager = tf.train.CheckpointManager(
checkpoint,
@@ -347,7 +346,7 @@ def export_classifier(model_export_path, input_meta_data, bert_config,
raise ValueError('Export path is not specified: %s' % model_dir)
# Export uses float32 for now, even if training uses mixed precision.
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
classifier_model = bert_models.classifier_model(
bert_config,
input_meta_data.get('num_labels', 1),
@@ -422,7 +421,7 @@ def custom_main(custom_callbacks=None, custom_metrics=None):
"""Run classification or regression.
Args:
- custom_callbacks: list of tf.keras.Callbacks passed to training loop.
+ custom_callbacks: list of tf_keras.Callbacks passed to training loop.
custom_metrics: list of metrics passed to the training loop.
"""
gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_param)
diff --git a/official/legacy/bert/run_pretraining.py b/official/legacy/bert/run_pretraining.py
index 6a1b1d7a59b..afef9e35b30 100644
--- a/official/legacy/bert/run_pretraining.py
+++ b/official/legacy/bert/run_pretraining.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,12 +14,11 @@
"""Run masked LM/next sentence pre-training for BERT in TF 2.x."""
-# Import libraries
from absl import app
from absl import flags
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.bert import bert_models
from official.legacy.bert import common_flags
@@ -63,7 +62,7 @@ def get_pretrain_dataset_fn(input_file_pattern, seq_length,
def _dataset_fn(ctx=None):
"""Returns tf.data.Dataset for distributed BERT pretraining."""
input_patterns = input_file_pattern.split(',')
- batch_size = ctx.get_per_replica_batch_size(global_batch_size)
+ batch_size = ctx.get_per_replica_batch_size(global_batch_size) # pyrefly: ignore[missing-attribute]
train_dataset = input_pipeline.create_pretrain_dataset(
input_patterns,
seq_length,
diff --git a/official/legacy/bert/run_squad.py b/official/legacy/bert/run_squad.py
index ee63bc96f73..9bcc14531a9 100644
--- a/official/legacy/bert/run_squad.py
+++ b/official/legacy/bert/run_squad.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,12 +18,11 @@
import os
import time
-# Import libraries
from absl import app
from absl import flags
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.bert import configs as bert_configs
from official.legacy.bert import run_squad_helper
diff --git a/official/legacy/bert/run_squad_helper.py b/official/legacy/bert/run_squad_helper.py
index be2e97dac5f..e5f2947dc7f 100644
--- a/official/legacy/bert/run_squad_helper.py
+++ b/official/legacy/bert/run_squad_helper.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import bert_models
from official.legacy.bert import common_flags
from official.legacy.bert import input_pipeline
@@ -97,9 +97,9 @@ def define_common_squad_flags():
def squad_loss_fn(start_positions, end_positions, start_logits, end_logits):
"""Returns sparse categorical crossentropy for start/end logits."""
- start_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ start_loss = tf_keras.losses.sparse_categorical_crossentropy(
start_positions, start_logits, from_logits=True)
- end_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ end_loss = tf_keras.losses.sparse_categorical_crossentropy(
end_positions, end_logits, from_logits=True)
total_loss = (tf.reduce_mean(start_loss) + tf.reduce_mean(end_loss)) / 2
@@ -160,7 +160,7 @@ def get_squad_model_to_predict(strategy, bert_config, checkpoint_path,
"""Gets a squad model to make predictions."""
with strategy.scope():
# Prediction always uses float32, even if training uses mixed precision.
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
squad_model, _ = bert_models.squad_model(
bert_config,
input_meta_data['max_seq_length'],
@@ -464,7 +464,7 @@ def export_squad(model_export_path, input_meta_data, bert_config):
if not model_export_path:
raise ValueError('Export path is not specified: %s' % model_export_path)
# Export uses float32 for now, even if training uses mixed precision.
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
squad_model, _ = bert_models.squad_model(bert_config,
input_meta_data['max_seq_length'])
model_saving_utils.export_bert_model(
diff --git a/official/legacy/bert/serving.py b/official/legacy/bert/serving.py
index 1666435aa8f..cb409b6a5b1 100644
--- a/official/legacy/bert/serving.py
+++ b/official/legacy/bert/serving.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import bert_models
from official.legacy.bert import configs
@@ -36,7 +36,7 @@
FLAGS = flags.FLAGS
-class BertServing(tf.keras.Model):
+class BertServing(tf_keras.Model):
"""Bert transformer encoder model for serving."""
def __init__(self, bert_config, name_to_features=None, name="serving_model"):
diff --git a/official/legacy/detection/__init__.py b/official/legacy/detection/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/__init__.py
+++ b/official/legacy/detection/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/configs/__init__.py b/official/legacy/detection/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/configs/__init__.py
+++ b/official/legacy/detection/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/configs/base_config.py b/official/legacy/detection/configs/base_config.py
index e274d91adc0..c942fc07af2 100644
--- a/official/legacy/detection/configs/base_config.py
+++ b/official/legacy/detection/configs/base_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/configs/factory.py b/official/legacy/detection/configs/factory.py
index d14f4b4e766..25f46be53af 100644
--- a/official/legacy/detection/configs/factory.py
+++ b/official/legacy/detection/configs/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/configs/maskrcnn_config.py b/official/legacy/detection/configs/maskrcnn_config.py
index 275cbf5e608..bf1ff4279b2 100644
--- a/official/legacy/detection/configs/maskrcnn_config.py
+++ b/official/legacy/detection/configs/maskrcnn_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/configs/olnmask_config.py b/official/legacy/detection/configs/olnmask_config.py
index 74e786c1fef..4b10575494f 100644
--- a/official/legacy/detection/configs/olnmask_config.py
+++ b/official/legacy/detection/configs/olnmask_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/configs/retinanet_config.py b/official/legacy/detection/configs/retinanet_config.py
index d3bd1ef19eb..d3a9b914a65 100644
--- a/official/legacy/detection/configs/retinanet_config.py
+++ b/official/legacy/detection/configs/retinanet_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/configs/shapemask_config.py b/official/legacy/detection/configs/shapemask_config.py
index 321a364f624..c3df0e4f674 100644
--- a/official/legacy/detection/configs/shapemask_config.py
+++ b/official/legacy/detection/configs/shapemask_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/dataloader/__init__.py b/official/legacy/detection/dataloader/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/dataloader/__init__.py
+++ b/official/legacy/detection/dataloader/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/dataloader/anchor.py b/official/legacy/detection/dataloader/anchor.py
index a5d90ed6c1d..42f895ccbea 100644
--- a/official/legacy/detection/dataloader/anchor.py
+++ b/official/legacy/detection/dataloader/anchor.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.utils import box_utils
from official.vision.ops import iou_similarity
from official.vision.utils.object_detection import argmax_matcher
diff --git a/official/legacy/detection/dataloader/factory.py b/official/legacy/detection/dataloader/factory.py
index 3bc8985eb43..21fd8c7fcf7 100644
--- a/official/legacy/detection/dataloader/factory.py
+++ b/official/legacy/detection/dataloader/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/dataloader/input_reader.py b/official/legacy/detection/dataloader/input_reader.py
index 4ffa729eda1..41121bddd41 100644
--- a/official/legacy/detection/dataloader/input_reader.py
+++ b/official/legacy/detection/dataloader/input_reader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
from typing import Optional, Text
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import factory
from official.legacy.detection.dataloader import mode_keys as ModeKeys
from official.modeling.hyperparams import params_dict
diff --git a/official/legacy/detection/dataloader/maskrcnn_parser.py b/official/legacy/detection/dataloader/maskrcnn_parser.py
index f69fa3260f0..62366f960bd 100644
--- a/official/legacy/detection/dataloader/maskrcnn_parser.py
+++ b/official/legacy/detection/dataloader/maskrcnn_parser.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Data parser and processing for Mask R-CNN."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import anchor
from official.legacy.detection.dataloader import mode_keys as ModeKeys
diff --git a/official/legacy/detection/dataloader/mode_keys.py b/official/legacy/detection/dataloader/mode_keys.py
index 93eb7d3ad9e..a6cda1f37fc 100644
--- a/official/legacy/detection/dataloader/mode_keys.py
+++ b/official/legacy/detection/dataloader/mode_keys.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/dataloader/olnmask_parser.py b/official/legacy/detection/dataloader/olnmask_parser.py
index b569d66be72..4cce8de5c3b 100644
--- a/official/legacy/detection/dataloader/olnmask_parser.py
+++ b/official/legacy/detection/dataloader/olnmask_parser.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Data parser and processing for Mask R-CNN."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import anchor
from official.legacy.detection.dataloader.maskrcnn_parser import Parser as MaskrcnnParser
diff --git a/official/legacy/detection/dataloader/retinanet_parser.py b/official/legacy/detection/dataloader/retinanet_parser.py
index 55058af79dd..8fba2307dfa 100644
--- a/official/legacy/detection/dataloader/retinanet_parser.py
+++ b/official/legacy/detection/dataloader/retinanet_parser.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
Focal Loss for Dense Object Detection. arXiv:1708.02002
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import anchor
from official.legacy.detection.dataloader import mode_keys as ModeKeys
diff --git a/official/legacy/detection/dataloader/shapemask_parser.py b/official/legacy/detection/dataloader/shapemask_parser.py
index 5feeb21d430..c3c4600c74f 100644
--- a/official/legacy/detection/dataloader/shapemask_parser.py
+++ b/official/legacy/detection/dataloader/shapemask_parser.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
ShapeMask: Learning to Segment Novel Objects by Refining Shape Priors.
arXiv:1904.03239.
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import anchor
from official.legacy.detection.dataloader import mode_keys as ModeKeys
diff --git a/official/legacy/detection/dataloader/tf_example_decoder.py b/official/legacy/detection/dataloader/tf_example_decoder.py
index 9e65509ce15..e5504bd122c 100644
--- a/official/legacy/detection/dataloader/tf_example_decoder.py
+++ b/official/legacy/detection/dataloader/tf_example_decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
A decoder to decode string tensors containing serialized tensorflow.Example
protos for object detection.
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class TfExampleDecoder(object):
diff --git a/official/legacy/detection/evaluation/__init__.py b/official/legacy/detection/evaluation/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/evaluation/__init__.py
+++ b/official/legacy/detection/evaluation/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/evaluation/coco_evaluator.py b/official/legacy/detection/evaluation/coco_evaluator.py
index 222763b5e40..badf1884179 100644
--- a/official/legacy/detection/evaluation/coco_evaluator.py
+++ b/official/legacy/detection/evaluation/coco_evaluator.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -38,16 +38,246 @@
import numpy as np
from pycocotools import cocoeval
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.evaluation import coco_utils
from official.legacy.detection.utils import class_utils
+class OlnCOCOevalWrapper(cocoeval.COCOeval):
+ """COCOeval wrapper class.
+
+ Rewritten based on cocoapi: (pycocotools/cocoeval.py)
+
+ This class wraps COCOEVAL API object, which provides the following additional
+ functionalities:
+ 1. summarze 'all', 'seen', and 'novel' split output print-out, e.g., AR at
+ different K proposals, AR and AP resutls for 'seen' and 'novel' class
+ splits.
+ """
+
+ def __init__(self, coco_gt, coco_dt, iou_type='box'):
+ super(OlnCOCOevalWrapper, self).__init__(
+ cocoGt=coco_gt, cocoDt=coco_dt, iouType=iou_type)
+
+ def summarize(self):
+ """Compute and display summary metrics for evaluation results.
+
+ Delta to the standard cocoapi function:
+ More Averate Recall metrics are produced with different top-K proposals.
+ Note this function can *only* be applied on the default parameter
+ setting.
+ Raises:
+ Exception: Please run accumulate() first.
+ """
+
+ def _summarize(ap=1, iou_thr=None, area_rng='all', max_dets=100):
+ p = self.params
+ i_str = (' {:<18} {} @[ IoU={:<9} | area={:>6s} | maxDets={:>3d} ] = '
+ '{:0.3f}')
+ title_str = 'Average Precision' if ap == 1 else 'Average Recall'
+ type_str = '(AP)' if ap == 1 else '(AR)'
+ iou_str = '{:0.2f}:{:0.2f}'.format(
+ p.iouThrs[0],
+ p.iouThrs[-1]) if iou_thr is None else '{:0.2f}'.format(iou_thr)
+
+ aind = [i for i, a_rng in enumerate(p.areaRngLbl) if a_rng == area_rng]
+ mind = [i for i, m_det in enumerate(p.maxDets) if m_det == max_dets]
+ if ap == 1:
+ # dimension of precision: [TxRxKxAxM]
+ s = self.eval['precision']
+ # IoU
+ if iou_thr is not None:
+ t = np.where(iou_thr == p.iouThrs)[0]
+ s = s[t]
+ s = s[:, :, :, aind, mind]
+ else:
+ # dimension of recall: [TxKxAxM]
+ s = self.eval['recall']
+ if iou_thr is not None:
+ t = np.where(iou_thr == p.iouThrs)[0]
+ s = s[t]
+ s = s[:, :, aind, mind]
+
+ if not (s[s > -1]).any():
+ mean_s = -1
+ else:
+ mean_s = np.mean(s[s > -1])
+ print(
+ i_str.format(title_str, type_str, iou_str, area_rng, max_dets,
+ mean_s))
+ return mean_s
+
+ def _summarize_dets():
+ stats = np.zeros((14,))
+ stats[0] = _summarize(1)
+ stats[1] = _summarize(
+ 1,
+ iou_thr=.5,
+ )
+ stats[2] = _summarize(
+ 1,
+ iou_thr=.75,
+ )
+ stats[3] = _summarize(
+ 1,
+ area_rng='small',
+ )
+ stats[4] = _summarize(
+ 1,
+ area_rng='medium',
+ )
+ stats[5] = _summarize(
+ 1,
+ area_rng='large',
+ )
+
+ stats[6] = _summarize(0, max_dets=self.params.maxDets[0]) # 10
+ stats[7] = _summarize(0, max_dets=self.params.maxDets[1]) # 20
+ stats[8] = _summarize(0, max_dets=self.params.maxDets[2]) # 50
+ stats[9] = _summarize(0, max_dets=self.params.maxDets[3]) # 100
+ stats[10] = _summarize(0, max_dets=self.params.maxDets[4]) # 200
+
+ stats[11] = _summarize(0, area_rng='small', max_dets=10)
+ stats[12] = _summarize(0, area_rng='medium', max_dets=10)
+ stats[13] = _summarize(0, area_rng='large', max_dets=10)
+ return stats
+
+ if not self.eval:
+ raise Exception('Please run accumulate() first')
+ summarize = _summarize_dets
+ self.stats = summarize()
+
+
+class OlnCOCOevalXclassWrapper(OlnCOCOevalWrapper):
+ """COCOeval wrapper class.
+
+ Rewritten based on cocoapi: (pycocotools/cocoeval.py)
+ Delta to the standard cocoapi:
+ Detections that hit the 'seen' class objects are ignored in top-K proposals.
+
+ This class wraps COCOEVAL API object, which provides the following additional
+ functionalities:
+ 1. Include ignore-class split (e.g., 'voc' or 'nonvoc').
+ 2. Do not count (or ignore) box proposals hitting ignore-class when
+ evaluating Average Recall at top-K proposals.
+ """
+
+ def __init__(self, coco_gt, coco_dt, iou_type='box'):
+ super(OlnCOCOevalXclassWrapper, self).__init__(
+ coco_gt=coco_gt, coco_dt=coco_dt, iou_type=iou_type)
+
+ def evaluateImg(self, img_id, cat_id, a_rng, max_det):
+ p = self.params
+ if p.useCats:
+ gt = self._gts[img_id, cat_id]
+ dt = self._dts[img_id, cat_id]
+ else:
+ gt, dt = [], []
+ for c_id in p.catIds:
+ gt.extend(self._gts[img_id, c_id])
+ dt.extend(self._dts[img_id, c_id])
+
+ if not gt and not dt:
+ return None
+
+ for g in gt:
+ if g['ignore'] or (g['area'] < a_rng[0] or g['area'] > a_rng[1]):
+ g['_ignore'] = 1
+ else:
+ g['_ignore'] = 0
+ # Class manipulation: ignore the 'ignored_split'.
+ if 'ignored_split' in g and g['ignored_split'] == 1:
+ g['_ignore'] = 1
+
+ # sort dt highest score first, sort gt ignore last
+ gtind = np.argsort([g['_ignore'] for g in gt], kind='mergesort')
+ gt = [gt[i] for i in gtind]
+ dtind = np.argsort([-d['score'] for d in dt], kind='mergesort')
+ dt = [dt[i] for i in dtind[0:max_det]]
+ iscrowd = [int(o['iscrowd']) for o in gt]
+ # load computed ious
+ # ious = self.ious[img_id, cat_id][:, gtind] if len(
+ # self.ious[img_id, cat_id]) > 0 else self.ious[img_id, cat_id]
+ if self.ious[img_id, cat_id].any():
+ ious = self.ious[img_id, cat_id][:, gtind]
+ else:
+ ious = self.ious[img_id, cat_id]
+
+ tt = len(p.iouThrs)
+ gg = len(gt)
+ dd = len(dt)
+ gtm = np.zeros((tt, gg))
+ dtm = np.zeros((tt, dd))
+ gt_ig = np.array([g['_ignore'] for g in gt])
+ dt_ig = np.zeros((tt, dd))
+ # indicator of whether the gt object class is of ignored_split or not.
+ gt_ig_split = np.array([g['ignored_split'] for g in gt])
+ dt_ig_split = np.zeros((dd))
+
+ if ious.any():
+ for tind, t in enumerate(p.iouThrs):
+ for dind, d in enumerate(dt):
+ # information about best match so far (m=-1 -> unmatched)
+ iou = min([t, 1 - 1e-10])
+ m = -1
+ for gind, g in enumerate(gt):
+ # if this gt already matched, and not a crowd, continue
+ if gtm[tind, gind] > 0 and not iscrowd[gind]:
+ continue
+ # if dt matched to reg gt, and on ignore gt, stop
+ if m > -1 and gt_ig[m] == 0 and gt_ig[gind] == 1:
+ break
+ # continue to next gt unless better match made
+ if ious[dind, gind] < iou:
+ continue
+ # if match successful and best so far, store appropriately
+ iou = ious[dind, gind]
+ m = gind
+ # if match made store id of match for both dt and gt
+ if m == -1:
+ continue
+ dt_ig[tind, dind] = gt_ig[m]
+ dtm[tind, dind] = gt[m]['id']
+ gtm[tind, m] = d['id']
+
+ # Activate to ignore the seen-class detections.
+ if tind == 0: # Register just only once: tind > 0 is also fine.
+ dt_ig_split[dind] = gt_ig_split[m]
+
+ # set unmatched detections outside of area range to ignore
+ a = np.array([d['area'] < a_rng[0] or d['area'] > a_rng[1] for d in dt
+ ]).reshape((1, len(dt)))
+ dt_ig = np.logical_or(dt_ig, np.logical_and(dtm == 0, np.repeat(a, tt, 0)))
+
+ # Activate to ignore the seen-class detections.
+ # Take only eval_split (eg, nonvoc) and ignore seen_split (eg, voc).
+ if dt_ig_split.sum() > 0:
+ dtm = dtm[:, dt_ig_split == 0]
+ dt_ig = dt_ig[:, dt_ig_split == 0]
+ len_dt = min(max_det, len(dt))
+ dt = [dt[i] for i in range(len_dt) if dt_ig_split[i] == 0]
+
+ # store results for given image and category
+ return {
+ 'image_id': img_id,
+ 'category_id': cat_id,
+ 'aRng': a_rng,
+ 'maxDet': max_det,
+ 'dtIds': [d['id'] for d in dt],
+ 'gtIds': [g['id'] for g in gt],
+ 'dtMatches': dtm,
+ 'gtMatches': gtm,
+ 'dtScores': [d['score'] for d in dt],
+ 'gtIgnore': gt_ig,
+ 'dtIgnore': dt_ig,
+ }
+
+
class MetricWrapper(object):
"""Metric Wrapper of the COCO evaluator."""
# This is only a wrapper for COCO metric and works on for numpy array. So it
- # doesn't inherit from tf.keras.layers.Layer or tf.keras.metrics.Metric.
+ # doesn't inherit from tf_keras.layers.Layer or tf_keras.metrics.Metric.
def __init__(self, evaluator):
self._evaluator = evaluator
diff --git a/official/legacy/detection/evaluation/coco_utils.py b/official/legacy/detection/evaluation/coco_utils.py
index 6c3692d011a..6d863c5e79f 100644
--- a/official/legacy/detection/evaluation/coco_utils.py
+++ b/official/legacy/detection/evaluation/coco_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -27,7 +27,7 @@
from pycocotools import coco
from pycocotools import mask as mask_api
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import tf_example_decoder
from official.legacy.detection.utils import box_utils
@@ -238,9 +238,7 @@ def convert_groundtruths_to_coco_dataset(groundtruths, label_map=None):
(boxes[j, k, 2] - boxes[j, k, 0]))
if 'masks' in groundtruths:
mask = Image.open(six.BytesIO(groundtruths['masks'][i][j, k]))
- width, height = mask.size
- np_mask = (
- np.array(mask.getdata()).reshape(height, width).astype(np.uint8))
+ np_mask = np.array(mask, dtype=np.uint8)
np_mask[np_mask > 0] = 255
encoded_mask = mask_api.encode(np.asfortranarray(np_mask))
ann['segmentation'] = encoded_mask
diff --git a/official/legacy/detection/evaluation/factory.py b/official/legacy/detection/evaluation/factory.py
index b47de01f9e1..ee98fbc48d3 100644
--- a/official/legacy/detection/evaluation/factory.py
+++ b/official/legacy/detection/evaluation/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/executor/__init__.py b/official/legacy/detection/executor/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/executor/__init__.py
+++ b/official/legacy/detection/executor/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/executor/detection_executor.py b/official/legacy/detection/executor/detection_executor.py
index 396de52cd62..0919f6bc9de 100644
--- a/official/legacy/detection/executor/detection_executor.py
+++ b/official/legacy/detection/executor/detection_executor.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.executor import distributed_executor as executor
from official.vision.utils.object_detection import visualization_utils
@@ -63,11 +63,11 @@ def _create_replicated_step(self,
logging.info('Filter trainable variables from %d to %d',
len(model.trainable_variables), len(trainable_variables))
update_state_fn = lambda labels, outputs: None
- if isinstance(metric, tf.keras.metrics.Metric):
+ if isinstance(metric, tf_keras.metrics.Metric):
update_state_fn = metric.update_state
else:
logging.error('Detection: train metric is not an instance of '
- 'tf.keras.metrics.Metric.')
+ 'tf_keras.metrics.Metric.')
def _replicated_step(inputs):
"""Replicated training step."""
@@ -151,7 +151,7 @@ def _run_evaluation(self, test_step, current_training_step, metric,
break
metric_result = metric.result()
- if isinstance(metric, tf.keras.metrics.Metric):
+ if isinstance(metric, tf_keras.metrics.Metric):
metric_result = tf.nest.map_structure(lambda x: x.numpy().astype(float),
metric_result)
logging.info('Step: [%d] Validation metric = %s', current_training_step,
diff --git a/official/legacy/detection/executor/distributed_executor.py b/official/legacy/detection/executor/distributed_executor.py
index ad6f5a22758..8efe2f8cd8f 100644
--- a/official/legacy/detection/executor/distributed_executor.py
+++ b/official/legacy/detection/executor/distributed_executor.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,7 +25,7 @@
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import,g-import-not-at-top,redefined-outer-name,reimported
from official.common import distribute_utils
@@ -63,12 +63,12 @@ def metrics_as_dict(metric):
Args:
metric: metric(s) to be put into the list. `metric` could be an object, a
- list, or a dict of tf.keras.metrics.Metric or has the `required_method`.
+ list, or a dict of tf_keras.metrics.Metric or has the `required_method`.
Returns:
A dictionary of valid metrics.
"""
- if isinstance(metric, tf.keras.metrics.Metric):
+ if isinstance(metric, tf_keras.metrics.Metric):
metrics = {metric.name: metric}
elif isinstance(metric, list):
metrics = {m.name: m for m in metric}
@@ -141,7 +141,7 @@ def __init__(self, strategy, params, model_fn, loss_fn, is_multi_host=False):
strategy: an instance of tf.distribute.Strategy.
params: Model configuration needed to run distribution strategy.
model_fn: Keras model function. Signature:
- (params: ParamsDict) -> tf.keras.models.Model.
+ (params: ParamsDict) -> tf_keras.models.Model.
loss_fn: loss function. Signature:
(y_true: Tensor, y_pred: Tensor) -> Tensor
is_multi_host: Set to True when using multi hosts for training, like multi
@@ -223,8 +223,8 @@ def _create_replicated_step(self,
strategy: an instance of tf.distribute.Strategy.
model: (Tensor, bool) -> Tensor. model function.
loss_fn: (y_true: Tensor, y_pred: Tensor) -> Tensor.
- optimizer: tf.keras.optimizers.Optimizer.
- metric: tf.keras.metrics.Metric subclass.
+ optimizer: tf_keras.optimizers.Optimizer.
+ metric: tf_keras.metrics.Metric subclass.
Returns:
The training step callable.
@@ -261,8 +261,8 @@ def _create_train_step(self,
strategy: an instance of tf.distribute.Strategy.
model: (Tensor, bool) -> Tensor. model function.
loss_fn: (y_true: Tensor, y_pred: Tensor) -> Tensor.
- optimizer: tf.keras.optimizers.Optimizer.
- metric: tf.keras.metrics.Metric subclass.
+ optimizer: tf_keras.optimizers.Optimizer.
+ metric: tf_keras.metrics.Metric subclass.
Returns:
The training step callable.
@@ -332,8 +332,8 @@ def train(
train_metric_fn: Optional[Callable[[], Any]] = None,
eval_metric_fn: Optional[Callable[[], Any]] = None,
summary_writer_fn: Callable[[Text, Text], SummaryWriter] = SummaryWriter,
- init_checkpoint: Optional[Callable[[tf.keras.Model], Any]] = None,
- custom_callbacks: Optional[List[tf.keras.callbacks.Callback]] = None,
+ init_checkpoint: Optional[Callable[[tf_keras.Model], Any]] = None,
+ custom_callbacks: Optional[List[tf_keras.callbacks.Callback]] = None,
continuous_eval: bool = False,
save_config: bool = True):
"""Runs distributed training.
@@ -412,7 +412,7 @@ def _run_callbacks_on_batch_end(batch):
train_loss = None
train_metric_result = None
eval_metric_result = None
- tf.keras.backend.set_learning_phase(1)
+ tf_keras.backend.set_learning_phase(1)
with strategy.scope():
# To correctly place the model weights on accelerators,
# model and optimizer should be created in scope.
@@ -444,10 +444,10 @@ def _run_callbacks_on_batch_end(batch):
eval_metric = eval_metric_fn()
train_metric = train_metric_fn()
- train_summary_writer = summary_writer_fn(model_dir, 'eval_train')
+ train_summary_writer = summary_writer_fn(model_dir, 'eval_train') # pyrefly: ignore[bad-argument-type]
self.train_summary_writer = train_summary_writer.writer
- test_summary_writer = summary_writer_fn(model_dir, 'eval_test')
+ test_summary_writer = summary_writer_fn(model_dir, 'eval_test') # pyrefly: ignore[bad-argument-type]
self.eval_summary_writer = test_summary_writer.writer
# Use training summary writer in TimeHistory if it's in use
@@ -472,7 +472,7 @@ def _run_callbacks_on_batch_end(batch):
_save_checkpoint(checkpoint, model_dir,
checkpoint_name.format(step=current_step))
if test_step:
- eval_iterator = self._get_input_iterator(eval_input_fn, strategy)
+ eval_iterator = self._get_input_iterator(eval_input_fn, strategy) # pyrefly: ignore[bad-argument-type]
eval_metric_result = self._run_evaluation(test_step, current_step,
eval_metric, eval_iterator)
logging.info('Step: %s evalation metric = %s.', current_step,
@@ -506,7 +506,7 @@ def _run_callbacks_on_batch_end(batch):
train_metric_result = train_loss
if callable(optimizer.lr):
train_metric_result.update(
- {'learning_rate': optimizer.lr(current_step).numpy()})
+ {'learning_rate': optimizer.lr(current_step).numpy()}) # pyrefly: ignore[missing-attribute]
else:
train_metric_result.update({'learning_rate': optimizer.lr.numpy()})
logging.info('Train Step: %d/%d / loss = %s / training metric = %s',
@@ -526,7 +526,7 @@ def _run_callbacks_on_batch_end(batch):
last_save_checkpoint_step = current_step
if continuous_eval and current_step < total_steps and test_step:
- eval_iterator = self._get_input_iterator(eval_input_fn, strategy)
+ eval_iterator = self._get_input_iterator(eval_input_fn, strategy) # pyrefly: ignore[bad-argument-type]
eval_metric_result = self._run_evaluation(test_step, current_step,
eval_metric, eval_iterator)
logging.info('Step: %s evalation metric = %s.', current_step,
@@ -547,7 +547,7 @@ def _run_callbacks_on_batch_end(batch):
if test_step:
logging.info('Running final evaluation after training is complete.')
- eval_iterator = self._get_input_iterator(eval_input_fn, strategy)
+ eval_iterator = self._get_input_iterator(eval_input_fn, strategy) # pyrefly: ignore[bad-argument-type]
eval_metric_result = self._run_evaluation(test_step, current_step,
eval_metric, eval_iterator)
logging.info('Final evaluation metric = %s.', eval_metric_result)
@@ -662,8 +662,8 @@ def evaluate_checkpoint(self,
raise ValueError('if `eval_metric_fn` is specified, '
'eval_metric_fn must be a callable.')
- old_phase = tf.keras.backend.learning_phase()
- tf.keras.backend.set_learning_phase(0)
+ old_phase = tf_keras.backend.learning_phase()
+ tf_keras.backend.set_learning_phase(0)
params = self._params
strategy = self._strategy
# To reduce unnecessary send/receive input pipeline operation, we place
@@ -683,8 +683,14 @@ def evaluate_checkpoint(self,
if not checkpoint_path:
raise ValueError('checkpoint path is empty')
reader = tf.compat.v1.train.NewCheckpointReader(checkpoint_path)
- current_step = reader.get_tensor(
- 'optimizer/iter/.ATTRIBUTES/VARIABLE_VALUE')
+ if reader.has_tensor('optimizer/iter/.ATTRIBUTES/VARIABLE_VALUE'):
+ # Legacy keras optimizer iteration.
+ current_step = reader.get_tensor(
+ 'optimizer/iter/.ATTRIBUTES/VARIABLE_VALUE')
+ else:
+ # New keras optimizer iteration.
+ current_step = reader.get_tensor(
+ 'optimizer/_iterations/.ATTRIBUTES/VARIABLE_VALUE')
logging.info('Checkpoint file %s found and restoring from '
'checkpoint', checkpoint_path)
status = checkpoint.restore(checkpoint_path)
@@ -696,10 +702,10 @@ def evaluate_checkpoint(self,
eval_metric, eval_iterator)
logging.info('Step: %s evalation metric = %s.', current_step,
eval_metric_result)
- summary_writer(metrics=eval_metric_result, step=current_step)
+ summary_writer(metrics=eval_metric_result, step=current_step) # pyrefly: ignore[not-callable]
reset_states(eval_metric)
- tf.keras.backend.set_learning_phase(old_phase)
+ tf_keras.backend.set_learning_phase(old_phase)
return eval_metric_result, current_step
def predict(self):
@@ -743,8 +749,8 @@ class MyDistributedExecutor(DistributedExecutor):
"""
def __init__(self, strategy_type=None, strategy_config=None):
- _ = distribute_utils.configure_cluster(strategy_config.worker_hosts,
- strategy_config.task_index)
+ _ = distribute_utils.configure_cluster(strategy_config.worker_hosts, # pyrefly: ignore[missing-attribute]
+ strategy_config.task_index) # pyrefly: ignore[missing-attribute]
"""Constructor.
Args:
@@ -756,10 +762,10 @@ def __init__(self, strategy_type=None, strategy_config=None):
"""
self._strategy = distribute_utils.get_distribution_strategy(
distribution_strategy=strategy_type,
- num_gpus=strategy_config.num_gpus,
- all_reduce_alg=strategy_config.all_reduce_alg,
- num_packs=strategy_config.num_packs,
- tpu_address=strategy_config.tpu)
+ num_gpus=strategy_config.num_gpus, # pyrefly: ignore[missing-attribute]
+ all_reduce_alg=strategy_config.all_reduce_alg, # pyrefly: ignore[missing-attribute]
+ num_packs=strategy_config.num_packs, # pyrefly: ignore[missing-attribute]
+ tpu_address=strategy_config.tpu) # pyrefly: ignore[missing-attribute]
@property
def strategy(self):
diff --git a/official/legacy/detection/main.py b/official/legacy/detection/main.py
index 9071e7c990c..55ded1d7411 100644
--- a/official/legacy/detection/main.py
+++ b/official/legacy/detection/main.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.detection.configs import factory as config_factory
diff --git a/official/legacy/detection/modeling/__init__.py b/official/legacy/detection/modeling/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/modeling/__init__.py
+++ b/official/legacy/detection/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/modeling/architecture/__init__.py b/official/legacy/detection/modeling/architecture/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/modeling/architecture/__init__.py
+++ b/official/legacy/detection/modeling/architecture/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/modeling/architecture/factory.py b/official/legacy/detection/modeling/architecture/factory.py
index 4b755200e11..173630030b4 100644
--- a/official/legacy/detection/modeling/architecture/factory.py
+++ b/official/legacy/detection/modeling/architecture/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/modeling/architecture/fpn.py b/official/legacy/detection/modeling/architecture/fpn.py
index 6b9edf6dfe3..f40e926fb94 100644
--- a/official/legacy/detection/modeling/architecture/fpn.py
+++ b/official/legacy/detection/modeling/architecture/fpn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,7 +26,7 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.modeling.architecture import nn_ops
from official.legacy.detection.ops import spatial_transform_ops
@@ -62,9 +62,9 @@ def __init__(self,
self._fpn_feat_dims = fpn_feat_dims
if use_separable_conv:
self._conv2d_op = functools.partial(
- tf.keras.layers.SeparableConv2D, depth_multiplier=1)
+ tf_keras.layers.SeparableConv2D, depth_multiplier=1)
else:
- self._conv2d_op = tf.keras.layers.Conv2D
+ self._conv2d_op = tf_keras.layers.Conv2D
if activation == 'relu':
self._activation_op = tf.nn.relu
elif activation == 'swish':
diff --git a/official/legacy/detection/modeling/architecture/heads.py b/official/legacy/detection/modeling/architecture/heads.py
index 430cb01d79d..3011701b28b 100644
--- a/official/legacy/detection/modeling/architecture/heads.py
+++ b/official/legacy/detection/modeling/architecture/heads.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,13 +21,13 @@
import functools
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.modeling.architecture import nn_ops
from official.legacy.detection.ops import spatial_transform_ops
-class RpnHead(tf.keras.layers.Layer):
+class RpnHead(tf_keras.layers.Layer):
"""Region Proposal Network head."""
def __init__(
@@ -74,13 +74,13 @@ def __init__(
if use_separable_conv:
self._conv2d_op = functools.partial(
- tf.keras.layers.SeparableConv2D,
+ tf_keras.layers.SeparableConv2D,
depth_multiplier=1,
bias_initializer=tf.zeros_initializer())
else:
self._conv2d_op = functools.partial(
- tf.keras.layers.Conv2D,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ tf_keras.layers.Conv2D,
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
bias_initializer=tf.zeros_initializer())
self._rpn_conv = self._conv2d_op(
@@ -138,7 +138,7 @@ def call(self, features, is_training=None):
return scores_outputs, box_outputs
-class OlnRpnHead(tf.keras.layers.Layer):
+class OlnRpnHead(tf_keras.layers.Layer):
"""Region Proposal Network for Object Localization Network (OLN)."""
def __init__(
@@ -183,13 +183,13 @@ def __init__(
if use_separable_conv:
self._conv2d_op = functools.partial(
- tf.keras.layers.SeparableConv2D,
+ tf_keras.layers.SeparableConv2D,
depth_multiplier=1,
bias_initializer=tf.zeros_initializer())
else:
self._conv2d_op = functools.partial(
- tf.keras.layers.Conv2D,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ tf_keras.layers.Conv2D,
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
bias_initializer=tf.zeros_initializer())
self._rpn_conv = self._conv2d_op(
@@ -262,7 +262,7 @@ def __call__(self, features, is_training=None):
return scores_outputs, box_outputs, center_outputs
-class FastrcnnHead(tf.keras.layers.Layer):
+class FastrcnnHead(tf_keras.layers.Layer):
"""Fast R-CNN box head."""
def __init__(
@@ -303,13 +303,13 @@ def __init__(
self._num_filters = num_filters
if use_separable_conv:
self._conv2d_op = functools.partial(
- tf.keras.layers.SeparableConv2D,
+ tf_keras.layers.SeparableConv2D,
depth_multiplier=1,
bias_initializer=tf.zeros_initializer())
else:
self._conv2d_op = functools.partial(
- tf.keras.layers.Conv2D,
- kernel_initializer=tf.keras.initializers.VarianceScaling(
+ tf_keras.layers.Conv2D,
+ kernel_initializer=tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
bias_initializer=tf.zeros_initializer())
@@ -344,7 +344,7 @@ def __init__(
self._fc_bn_ops = []
for i in range(self._num_fcs):
self._fc_ops.append(
- tf.keras.layers.Dense(
+ tf_keras.layers.Dense(
units=self._fc_dims,
activation=(None
if self._use_batch_norm else self._activation_op),
@@ -352,14 +352,14 @@ def __init__(
if self._use_batch_norm:
self._fc_bn_ops.append(self._norm_activation(fused=False))
- self._class_predict = tf.keras.layers.Dense(
+ self._class_predict = tf_keras.layers.Dense(
self._num_classes,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
bias_initializer=tf.zeros_initializer(),
name='class-predict')
- self._box_predict = tf.keras.layers.Dense(
+ self._box_predict = tf_keras.layers.Dense(
self._num_classes * 4,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.001),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.001),
bias_initializer=tf.zeros_initializer(),
name='box-predict')
@@ -403,7 +403,7 @@ def call(self, roi_features, is_training=None):
return class_outputs, box_outputs
-class OlnBoxScoreHead(tf.keras.layers.Layer):
+class OlnBoxScoreHead(tf_keras.layers.Layer):
"""Box head of Object Localization Network (OLN)."""
def __init__(
@@ -442,13 +442,13 @@ def __init__(
self._num_filters = num_filters
if use_separable_conv:
self._conv2d_op = functools.partial(
- tf.keras.layers.SeparableConv2D,
+ tf_keras.layers.SeparableConv2D,
depth_multiplier=1,
bias_initializer=tf.zeros_initializer())
else:
self._conv2d_op = functools.partial(
- tf.keras.layers.Conv2D,
- kernel_initializer=tf.keras.initializers.VarianceScaling(
+ tf_keras.layers.Conv2D,
+ kernel_initializer=tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
bias_initializer=tf.zeros_initializer())
@@ -483,7 +483,7 @@ def __init__(
self._fc_bn_ops = []
for i in range(self._num_fcs):
self._fc_ops.append(
- tf.keras.layers.Dense(
+ tf_keras.layers.Dense(
units=self._fc_dims,
activation=(None
if self._use_batch_norm else self._activation_op),
@@ -491,19 +491,19 @@ def __init__(
if self._use_batch_norm:
self._fc_bn_ops.append(self._norm_activation(fused=False))
- self._class_predict = tf.keras.layers.Dense(
+ self._class_predict = tf_keras.layers.Dense(
self._num_classes,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
bias_initializer=tf.zeros_initializer(),
name='class-predict')
- self._box_predict = tf.keras.layers.Dense(
+ self._box_predict = tf_keras.layers.Dense(
self._num_classes * 4,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.001),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.001),
bias_initializer=tf.zeros_initializer(),
name='box-predict')
- self._score_predict = tf.keras.layers.Dense(
+ self._score_predict = tf_keras.layers.Dense(
1,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
bias_initializer=tf.zeros_initializer(),
name='score-predict')
@@ -547,7 +547,7 @@ def __call__(self, roi_features, is_training=None):
return class_outputs, box_outputs, score_outputs
-class MaskrcnnHead(tf.keras.layers.Layer):
+class MaskrcnnHead(tf_keras.layers.Layer):
"""Mask R-CNN head."""
def __init__(
@@ -584,13 +584,13 @@ def __init__(
self._num_filters = num_filters
if use_separable_conv:
self._conv2d_op = functools.partial(
- tf.keras.layers.SeparableConv2D,
+ tf_keras.layers.SeparableConv2D,
depth_multiplier=1,
bias_initializer=tf.zeros_initializer())
else:
self._conv2d_op = functools.partial(
- tf.keras.layers.Conv2D,
- kernel_initializer=tf.keras.initializers.VarianceScaling(
+ tf_keras.layers.Conv2D,
+ kernel_initializer=tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
bias_initializer=tf.zeros_initializer())
if activation == 'relu':
@@ -613,13 +613,13 @@ def __init__(
activation=(None
if self._use_batch_norm else self._activation_op),
name='mask-conv-l%d' % i))
- self._mask_conv_transpose = tf.keras.layers.Conv2DTranspose(
+ self._mask_conv_transpose = tf_keras.layers.Conv2DTranspose(
self._num_filters,
kernel_size=(2, 2),
strides=(2, 2),
padding='valid',
activation=(None if self._use_batch_norm else self._activation_op),
- kernel_initializer=tf.keras.initializers.VarianceScaling(
+ kernel_initializer=tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
bias_initializer=tf.zeros_initializer(),
name='conv5-mask')
@@ -673,18 +673,12 @@ class the ROI is.
])
with tf.name_scope('masks_post_processing'):
- # TODO(pengchong): Figure out the way not to use the static inferred
- # batch size.
- batch_size, num_masks = class_indices.get_shape().as_list()
- mask_outputs = tf.transpose(a=mask_outputs, perm=[0, 1, 4, 2, 3])
- # Constructs indices for gather.
- batch_indices = tf.tile(
- tf.expand_dims(tf.range(batch_size), axis=1), [1, num_masks])
- mask_indices = tf.tile(
- tf.expand_dims(tf.range(num_masks), axis=0), [batch_size, 1])
- gather_indices = tf.stack(
- [batch_indices, mask_indices, class_indices], axis=2)
- mask_outputs = tf.gather_nd(mask_outputs, gather_indices)
+ mask_outputs = tf.gather(
+ mask_outputs,
+ tf.cast(class_indices, tf.int32),
+ axis=-1,
+ batch_dims=2,
+ )
return mask_outputs
@@ -741,18 +735,18 @@ def _box_net_batch_norm_name(self, i, level):
def _build_class_net_layers(self, norm_activation):
"""Build re-usable layers for class prediction network."""
if self._use_separable_conv:
- self._class_predict = tf.keras.layers.SeparableConv2D(
+ self._class_predict = tf_keras.layers.SeparableConv2D(
self._num_classes * self._anchors_per_location,
kernel_size=(3, 3),
bias_initializer=tf.constant_initializer(-np.log((1 - 0.01) / 0.01)),
padding='same',
name='class-predict')
else:
- self._class_predict = tf.keras.layers.Conv2D(
+ self._class_predict = tf_keras.layers.Conv2D(
self._num_classes * self._anchors_per_location,
kernel_size=(3, 3),
bias_initializer=tf.constant_initializer(-np.log((1 - 0.01) / 0.01)),
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=1e-5),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=1e-5),
padding='same',
name='class-predict')
self._class_conv = []
@@ -760,7 +754,7 @@ def _build_class_net_layers(self, norm_activation):
for i in range(self._num_convs):
if self._use_separable_conv:
self._class_conv.append(
- tf.keras.layers.SeparableConv2D(
+ tf_keras.layers.SeparableConv2D(
self._num_filters,
kernel_size=(3, 3),
bias_initializer=tf.zeros_initializer(),
@@ -769,11 +763,11 @@ def _build_class_net_layers(self, norm_activation):
name='class-' + str(i)))
else:
self._class_conv.append(
- tf.keras.layers.Conv2D(
+ tf_keras.layers.Conv2D(
self._num_filters,
kernel_size=(3, 3),
bias_initializer=tf.zeros_initializer(),
- kernel_initializer=tf.keras.initializers.RandomNormal(
+ kernel_initializer=tf_keras.initializers.RandomNormal(
stddev=0.01),
activation=None,
padding='same',
@@ -785,18 +779,18 @@ def _build_class_net_layers(self, norm_activation):
def _build_box_net_layers(self, norm_activation):
"""Build re-usable layers for box prediction network."""
if self._use_separable_conv:
- self._box_predict = tf.keras.layers.SeparableConv2D(
+ self._box_predict = tf_keras.layers.SeparableConv2D(
4 * self._anchors_per_location,
kernel_size=(3, 3),
bias_initializer=tf.zeros_initializer(),
padding='same',
name='box-predict')
else:
- self._box_predict = tf.keras.layers.Conv2D(
+ self._box_predict = tf_keras.layers.Conv2D(
4 * self._anchors_per_location,
kernel_size=(3, 3),
bias_initializer=tf.zeros_initializer(),
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=1e-5),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=1e-5),
padding='same',
name='box-predict')
self._box_conv = []
@@ -804,7 +798,7 @@ def _build_box_net_layers(self, norm_activation):
for i in range(self._num_convs):
if self._use_separable_conv:
self._box_conv.append(
- tf.keras.layers.SeparableConv2D(
+ tf_keras.layers.SeparableConv2D(
self._num_filters,
kernel_size=(3, 3),
activation=None,
@@ -813,12 +807,12 @@ def _build_box_net_layers(self, norm_activation):
name='box-' + str(i)))
else:
self._box_conv.append(
- tf.keras.layers.Conv2D(
+ tf_keras.layers.Conv2D(
self._num_filters,
kernel_size=(3, 3),
activation=None,
bias_initializer=tf.zeros_initializer(),
- kernel_initializer=tf.keras.initializers.RandomNormal(
+ kernel_initializer=tf_keras.initializers.RandomNormal(
stddev=0.01),
padding='same',
name='box-' + str(i)))
@@ -892,7 +886,7 @@ def __init__(self, num_classes, num_downsample_channels, mask_crop_size,
self._shape_prior_path = shape_prior_path
self._use_category_for_mask = use_category_for_mask
- self._shape_prior_fc = tf.keras.layers.Dense(
+ self._shape_prior_fc = tf_keras.layers.Dense(
self._num_downsample_channels, name='shape-prior-fc')
def __call__(self, fpn_features, boxes, outer_boxes, classes, is_training):
@@ -992,7 +986,7 @@ class ids.
# Reduce spatial dimension of features. The features have shape
# [batch_size, num_instances, num_channels].
features = tf.reduce_mean(features, axis=(2, 3))
- logits = tf.keras.layers.Dense(
+ logits = tf_keras.layers.Dense(
self._mask_num_classes * self._num_clusters,
kernel_initializer=tf.random_normal_initializer(stddev=0.01),
name='classify-shape-prior-fc')(features)
@@ -1038,7 +1032,7 @@ def __init__(self,
self._num_convs = num_convs
self._norm_activation = norm_activation
- self._coarse_mask_fc = tf.keras.layers.Dense(
+ self._coarse_mask_fc = tf_keras.layers.Dense(
self._num_downsample_channels, name='coarse-mask-fc')
self._class_conv = []
@@ -1046,11 +1040,11 @@ def __init__(self,
for i in range(self._num_convs):
self._class_conv.append(
- tf.keras.layers.Conv2D(
+ tf_keras.layers.Conv2D(
self._num_downsample_channels,
kernel_size=(3, 3),
bias_initializer=tf.zeros_initializer(),
- kernel_initializer=tf.keras.initializers.RandomNormal(
+ kernel_initializer=tf_keras.initializers.RandomNormal(
stddev=0.01),
padding='same',
name='coarse-mask-class-%d' % i))
@@ -1058,12 +1052,12 @@ def __init__(self,
self._class_norm_activation.append(
norm_activation(name='coarse-mask-class-%d-bn' % i))
- self._class_predict = tf.keras.layers.Conv2D(
+ self._class_predict = tf_keras.layers.Conv2D(
self._mask_num_classes,
kernel_size=(1, 1),
# Focal loss bias initialization to have foreground 0.01 probability.
bias_initializer=tf.constant_initializer(-np.log((1 - 0.01) / 0.01)),
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
padding='same',
name='coarse-mask-class-predict')
@@ -1164,10 +1158,10 @@ def __init__(self,
self._num_convs = num_convs
self.up_sample_factor = upsample_factor
- self._fine_mask_fc = tf.keras.layers.Dense(
+ self._fine_mask_fc = tf_keras.layers.Dense(
self._num_downsample_channels, name='fine-mask-fc')
- self._upsample_conv = tf.keras.layers.Conv2DTranspose(
+ self._upsample_conv = tf_keras.layers.Conv2DTranspose(
self._num_downsample_channels,
(self.up_sample_factor, self.up_sample_factor),
(self.up_sample_factor, self.up_sample_factor),
@@ -1177,11 +1171,11 @@ def __init__(self,
self._fine_class_bn = []
for i in range(self._num_convs):
self._fine_class_conv.append(
- tf.keras.layers.Conv2D(
+ tf_keras.layers.Conv2D(
self._num_downsample_channels,
kernel_size=(3, 3),
bias_initializer=tf.zeros_initializer(),
- kernel_initializer=tf.keras.initializers.RandomNormal(
+ kernel_initializer=tf_keras.initializers.RandomNormal(
stddev=0.01),
activation=None,
padding='same',
@@ -1189,12 +1183,12 @@ def __init__(self,
self._fine_class_bn.append(
norm_activation(name='fine-mask-class-%d-bn' % i))
- self._class_predict_conv = tf.keras.layers.Conv2D(
+ self._class_predict_conv = tf_keras.layers.Conv2D(
self._mask_num_classes,
kernel_size=(1, 1),
# Focal loss bias initialization to have foreground 0.01 probability.
bias_initializer=tf.constant_initializer(-np.log((1 - 0.01) / 0.01)),
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
padding='same',
name='fine-mask-class-predict')
diff --git a/official/legacy/detection/modeling/architecture/identity.py b/official/legacy/detection/modeling/architecture/identity.py
index 7d3280dbd5e..eae3105c607 100644
--- a/official/legacy/detection/modeling/architecture/identity.py
+++ b/official/legacy/detection/modeling/architecture/identity.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/modeling/architecture/nn_blocks.py b/official/legacy/detection/modeling/architecture/nn_blocks.py
index ab61d3239a9..9ac546de702 100644
--- a/official/legacy/detection/modeling/architecture/nn_blocks.py
+++ b/official/legacy/detection/modeling/architecture/nn_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,12 +18,12 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
-class ResidualBlock(tf.keras.layers.Layer):
+class ResidualBlock(tf_keras.layers.Layer):
"""A residual block."""
def __init__(self,
@@ -50,9 +50,9 @@ def __init__(self,
for the first block of a block group, which may change the number of
filters and the resolution.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
activation: `str` name of the activation function.
use_sync_bn: if True, use synchronized batch normalization.
@@ -75,10 +75,10 @@ def __init__(self,
self._bias_regularizer = bias_regularizer
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -86,7 +86,7 @@ def __init__(self,
def build(self, input_shape):
if self._use_projection:
- self._shortcut = tf.keras.layers.Conv2D(
+ self._shortcut = tf_keras.layers.Conv2D(
filters=self._filters,
kernel_size=1,
strides=self._strides,
@@ -99,7 +99,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv1 = tf.keras.layers.Conv2D(
+ self._conv1 = tf_keras.layers.Conv2D(
filters=self._filters,
kernel_size=3,
strides=self._strides,
@@ -113,7 +113,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv2 = tf.keras.layers.Conv2D(
+ self._conv2 = tf_keras.layers.Conv2D(
filters=self._filters,
kernel_size=3,
strides=1,
@@ -162,7 +162,7 @@ def call(self, inputs):
return self._activation_fn(x + shortcut)
-class BottleneckBlock(tf.keras.layers.Layer):
+class BottleneckBlock(tf_keras.layers.Layer):
"""A standard bottleneck block."""
def __init__(self,
@@ -189,9 +189,9 @@ def __init__(self,
for the first block of a block group, which may change the number of
filters and the resolution.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
activation: `str` name of the activation function.
use_sync_bn: if True, use synchronized batch normalization.
@@ -213,10 +213,10 @@ def __init__(self,
self._kernel_regularizer = kernel_regularizer
self._bias_regularizer = bias_regularizer
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -224,7 +224,7 @@ def __init__(self,
def build(self, input_shape):
if self._use_projection:
- self._shortcut = tf.keras.layers.Conv2D(
+ self._shortcut = tf_keras.layers.Conv2D(
filters=self._filters * 4,
kernel_size=1,
strides=self._strides,
@@ -237,7 +237,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv1 = tf.keras.layers.Conv2D(
+ self._conv1 = tf_keras.layers.Conv2D(
filters=self._filters,
kernel_size=1,
strides=1,
@@ -250,7 +250,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv2 = tf.keras.layers.Conv2D(
+ self._conv2 = tf_keras.layers.Conv2D(
filters=self._filters,
kernel_size=3,
strides=self._strides,
@@ -264,7 +264,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv3 = tf.keras.layers.Conv2D(
+ self._conv3 = tf_keras.layers.Conv2D(
filters=self._filters * 4,
kernel_size=1,
strides=1,
diff --git a/official/legacy/detection/modeling/architecture/nn_ops.py b/official/legacy/detection/modeling/architecture/nn_ops.py
index 70f47c9af0b..40b3f45341c 100644
--- a/official/legacy/detection/modeling/architecture/nn_ops.py
+++ b/official/legacy/detection/modeling/architecture/nn_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,10 +20,10 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class NormActivation(tf.keras.layers.Layer):
+class NormActivation(tf_keras.layers.Layer):
"""Combined Normalization and Activation layers."""
def __init__(self,
@@ -54,10 +54,10 @@ def __init__(self,
"""
super(NormActivation, self).__init__(trainable=trainable)
if init_zero:
- gamma_initializer = tf.keras.initializers.Zeros()
+ gamma_initializer = tf_keras.initializers.Zeros()
else:
- gamma_initializer = tf.keras.initializers.Ones()
- self._normalization_op = tf.keras.layers.BatchNormalization(
+ gamma_initializer = tf_keras.initializers.Ones()
+ self._normalization_op = tf_keras.layers.BatchNormalization(
momentum=momentum,
epsilon=epsilon,
center=True,
diff --git a/official/legacy/detection/modeling/architecture/resnet.py b/official/legacy/detection/modeling/architecture/resnet.py
index 0a8182bfe4a..40180b1c0a6 100644
--- a/official/legacy/detection/modeling/architecture/resnet.py
+++ b/official/legacy/detection/modeling/architecture/resnet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.modeling.architecture import nn_ops
@@ -159,7 +159,7 @@ def conv2d_fixed_padding(self, inputs, filters, kernel_size, strides):
if strides > 1:
inputs = self.fixed_padding(inputs, kernel_size)
- return tf.keras.layers.Conv2D(
+ return tf_keras.layers.Conv2D(
filters=filters,
kernel_size=kernel_size,
strides=strides,
@@ -309,7 +309,7 @@ def model(inputs, is_training=None):
inputs = tf.identity(inputs, 'initial_conv')
inputs = self._norm_activation()(inputs, is_training=is_training)
- inputs = tf.keras.layers.MaxPool2D(
+ inputs = tf_keras.layers.MaxPool2D(
pool_size=3, strides=2, padding='SAME',
data_format=self._data_format)(
inputs)
diff --git a/official/legacy/detection/modeling/architecture/spinenet.py b/official/legacy/detection/modeling/architecture/spinenet.py
index ea86a70f28d..c36021082d9 100644
--- a/official/legacy/detection/modeling/architecture/spinenet.py
+++ b/official/legacy/detection/modeling/architecture/spinenet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,11 +22,11 @@
import math
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.modeling.architecture import nn_blocks
from official.modeling import tf_utils
-layers = tf.keras.layers
+layers = tf_keras.layers
FILTER_SIZE_MAP = {
1: 32,
@@ -111,14 +111,14 @@ def build_block_specs(block_specs=None):
return [BlockSpec(*b) for b in block_specs]
-class SpineNet(tf.keras.Model):
+class SpineNet(tf_keras.Model):
"""Class to build SpineNet models."""
def __init__(self,
- input_specs=tf.keras.layers.InputSpec(shape=[None, 640, 640, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[None, 640, 640, 3]),
min_level=3,
max_level=7,
- block_specs=build_block_specs(),
+ block_specs=None,
endpoints_num_filters=256,
resample_alpha=0.5,
block_repeats=1,
@@ -134,7 +134,9 @@ def __init__(self,
"""SpineNet model."""
self._min_level = min_level
self._max_level = max_level
- self._block_specs = block_specs
+ self._block_specs = (
+ build_block_specs() if block_specs is None else block_specs
+ )
self._endpoints_num_filters = endpoints_num_filters
self._resample_alpha = resample_alpha
self._block_repeats = block_repeats
@@ -159,13 +161,13 @@ def __init__(self,
else:
self._norm = layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
# Build SpineNet.
- inputs = tf.keras.Input(shape=input_specs.shape[1:])
+ inputs = tf_keras.Input(shape=input_specs.shape[1:])
net = self._build_stem(inputs=inputs)
net = self._build_scale_permuted_network(
@@ -451,10 +453,10 @@ class SpineNetBuilder(object):
def __init__(self,
model_id,
- input_specs=tf.keras.layers.InputSpec(shape=[None, 640, 640, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[None, 640, 640, 3]),
min_level=3,
max_level=7,
- block_specs=build_block_specs(),
+ block_specs=None,
kernel_initializer='VarianceScaling',
kernel_regularizer=None,
bias_regularizer=None,
@@ -469,7 +471,7 @@ def __init__(self,
self._input_specs = input_specs
self._min_level = min_level
self._max_level = max_level
- self._block_specs = block_specs
+ self._block_specs = block_specs or build_block_specs()
self._endpoints_num_filters = scaling_params['endpoints_num_filters']
self._resample_alpha = scaling_params['resample_alpha']
self._block_repeats = scaling_params['block_repeats']
diff --git a/official/legacy/detection/modeling/base_model.py b/official/legacy/detection/modeling/base_model.py
index aa84f468263..92ac9407b8f 100644
--- a/official/legacy/detection/modeling/base_model.py
+++ b/official/legacy/detection/modeling/base_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
import abc
import re
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.modeling import checkpoint_utils
from official.legacy.detection.modeling import learning_rates
from official.legacy.detection.modeling import optimizers
diff --git a/official/legacy/detection/modeling/checkpoint_utils.py b/official/legacy/detection/modeling/checkpoint_utils.py
index 1765a059c30..d0fe22cfc7d 100644
--- a/official/legacy/detection/modeling/checkpoint_utils.py
+++ b/official/legacy/detection/modeling/checkpoint_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,7 +26,7 @@
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _build_assignment_map(keras_model,
@@ -40,7 +40,7 @@ def _build_assignment_map(keras_model,
the new Keras name.
Args:
- keras_model: tf.keras.Model object to provide variables to assign.
+ keras_model: tf_keras.Model object to provide variables to assign.
prefix: prefix in the variable name to be remove for alignment with names in
the checkpoint.
skip_variables_regex: regular expression to math the names of variables that
diff --git a/official/legacy/detection/modeling/factory.py b/official/legacy/detection/modeling/factory.py
index 3d852b8d040..58175ade16a 100644
--- a/official/legacy/detection/modeling/factory.py
+++ b/official/legacy/detection/modeling/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/modeling/learning_rates.py b/official/legacy/detection/modeling/learning_rates.py
index bbd34873981..9d0f6db8980 100644
--- a/official/legacy/detection/modeling/learning_rates.py
+++ b/official/legacy/detection/modeling/learning_rates.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,12 +19,12 @@
from __future__ import print_function
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.hyperparams import params_dict
class StepLearningRateWithLinearWarmup(
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
"""Class to generate learning rate tensor."""
def __init__(self, total_steps, params):
@@ -57,7 +57,7 @@ def get_config(self):
class CosineLearningRateWithLinearWarmup(
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
"""Class to generate learning rate tensor."""
def __init__(self, total_steps, params):
diff --git a/official/legacy/detection/modeling/losses.py b/official/legacy/detection/modeling/losses.py
index f3423993390..144d450978c 100644
--- a/official/legacy/detection/modeling/losses.py
+++ b/official/legacy/detection/modeling/losses.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from __future__ import print_function
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def focal_loss(logits, targets, alpha, gamma, normalizer):
@@ -90,8 +90,8 @@ class RpnScoreLoss(object):
def __init__(self, params):
self._rpn_batch_size_per_im = params.rpn_batch_size_per_im
- self._binary_crossentropy = tf.keras.losses.BinaryCrossentropy(
- reduction=tf.keras.losses.Reduction.SUM, from_logits=True)
+ self._binary_crossentropy = tf_keras.losses.BinaryCrossentropy(
+ reduction=tf_keras.losses.Reduction.SUM, from_logits=True)
def __call__(self, score_outputs, labels):
"""Computes total RPN detection loss.
@@ -153,8 +153,8 @@ def __init__(self, params):
# The delta is typically around the mean value of regression target.
# for instances, the regression targets of 512x512 input with 6 anchors on
# P2-P6 pyramid is about [0.1, 0.1, 0.2, 0.2].
- self._huber_loss = tf.keras.losses.Huber(
- delta=params.huber_loss_delta, reduction=tf.keras.losses.Reduction.SUM)
+ self._huber_loss = tf_keras.losses.Huber(
+ delta=params.huber_loss_delta, reduction=tf_keras.losses.Reduction.SUM)
def __call__(self, box_outputs, labels):
"""Computes total RPN detection loss.
@@ -199,8 +199,8 @@ class OlnRpnCenterLoss(object):
"""Object Localization Network RPN centerness regression loss function."""
def __init__(self):
- self._l1_loss = tf.keras.losses.MeanAbsoluteError(
- reduction=tf.keras.losses.Reduction.SUM)
+ self._l1_loss = tf_keras.losses.MeanAbsoluteError(
+ reduction=tf_keras.losses.Reduction.SUM)
def __call__(self, center_outputs, labels):
"""Computes total RPN centerness regression loss.
@@ -342,8 +342,8 @@ class FastrcnnClassLoss(object):
"""Fast R-CNN classification loss function."""
def __init__(self):
- self._categorical_crossentropy = tf.keras.losses.CategoricalCrossentropy(
- reduction=tf.keras.losses.Reduction.SUM, from_logits=True)
+ self._categorical_crossentropy = tf_keras.losses.CategoricalCrossentropy(
+ reduction=tf_keras.losses.Reduction.SUM, from_logits=True)
def __call__(self, class_outputs, class_targets):
"""Computes the class loss (Fast-RCNN branch) of Mask-RCNN.
@@ -388,8 +388,8 @@ def __init__(self, params):
# The delta is typically around the mean value of regression target.
# for instances, the regression targets of 512x512 input with 6 anchors on
# P2-P6 pyramid is about [0.1, 0.1, 0.2, 0.2].
- self._huber_loss = tf.keras.losses.Huber(
- delta=params.huber_loss_delta, reduction=tf.keras.losses.Reduction.SUM)
+ self._huber_loss = tf_keras.losses.Huber(
+ delta=params.huber_loss_delta, reduction=tf_keras.losses.Reduction.SUM)
def __call__(self, box_outputs, class_targets, box_targets):
"""Computes the box loss (Fast-RCNN branch) of Mask-RCNN.
@@ -465,8 +465,8 @@ class OlnBoxScoreLoss(object):
def __init__(self, params):
self._ignore_threshold = params.ignore_threshold
- self._l1_loss = tf.keras.losses.MeanAbsoluteError(
- reduction=tf.keras.losses.Reduction.SUM)
+ self._l1_loss = tf_keras.losses.MeanAbsoluteError(
+ reduction=tf_keras.losses.Reduction.SUM)
def __call__(self, score_outputs, score_targets):
"""Computes the class loss (Fast-RCNN branch) of Mask-RCNN.
@@ -505,8 +505,8 @@ class MaskrcnnLoss(object):
"""Mask R-CNN instance segmentation mask loss function."""
def __init__(self):
- self._binary_crossentropy = tf.keras.losses.BinaryCrossentropy(
- reduction=tf.keras.losses.Reduction.SUM, from_logits=True)
+ self._binary_crossentropy = tf_keras.losses.BinaryCrossentropy(
+ reduction=tf_keras.losses.Reduction.SUM, from_logits=True)
def __call__(self, mask_outputs, mask_targets, select_class_targets):
"""Computes the mask loss of Mask-RCNN.
@@ -616,8 +616,8 @@ class RetinanetBoxLoss(object):
"""RetinaNet box loss."""
def __init__(self, params):
- self._huber_loss = tf.keras.losses.Huber(
- delta=params.huber_loss_delta, reduction=tf.keras.losses.Reduction.SUM)
+ self._huber_loss = tf_keras.losses.Huber(
+ delta=params.huber_loss_delta, reduction=tf_keras.losses.Reduction.SUM)
def __call__(self, box_outputs, labels, num_positives):
"""Computes box detection loss.
@@ -695,8 +695,8 @@ class ShapemaskLoss(object):
"""ShapeMask mask loss function wrapper."""
def __init__(self):
- self._binary_crossentropy = tf.keras.losses.BinaryCrossentropy(
- reduction=tf.keras.losses.Reduction.SUM, from_logits=True)
+ self._binary_crossentropy = tf_keras.losses.BinaryCrossentropy(
+ reduction=tf_keras.losses.Reduction.SUM, from_logits=True)
def __call__(self, logits, labels, valid_mask):
"""ShapeMask mask cross entropy loss function wrapper.
diff --git a/official/legacy/detection/modeling/maskrcnn_model.py b/official/legacy/detection/modeling/maskrcnn_model.py
index 576457b6122..a6d30906dc1 100644
--- a/official/legacy/detection/modeling/maskrcnn_model.py
+++ b/official/legacy/detection/modeling/maskrcnn_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import anchor
from official.legacy.detection.dataloader import mode_keys
@@ -239,31 +239,31 @@ def build_input_layers(self, params, mode):
batch_size = params.train.batch_size
input_layer = {
'image':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=input_shape,
batch_size=batch_size,
name='image',
dtype=tf.bfloat16 if self._use_bfloat16 else tf.float32),
'image_info':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[4, 2],
batch_size=batch_size,
name='image_info',
),
'gt_boxes':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[params.maskrcnn_parser.max_num_instances, 4],
batch_size=batch_size,
name='gt_boxes'),
'gt_classes':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[params.maskrcnn_parser.max_num_instances],
batch_size=batch_size,
name='gt_classes',
dtype=tf.int64),
}
if self._include_mask:
- input_layer['gt_masks'] = tf.keras.layers.Input(
+ input_layer['gt_masks'] = tf_keras.layers.Input(
shape=[
params.maskrcnn_parser.max_num_instances,
params.maskrcnn_parser.mask_crop_size,
@@ -275,13 +275,13 @@ def build_input_layers(self, params, mode):
batch_size = params.eval.batch_size
input_layer = {
'image':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=input_shape,
batch_size=batch_size,
name='image',
dtype=tf.bfloat16 if self._use_bfloat16 else tf.float32),
'image_info':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[4, 2],
batch_size=batch_size,
name='image_info',
@@ -294,9 +294,9 @@ def build_model(self, params, mode):
input_layers = self.build_input_layers(self._params, mode)
outputs = self.model_outputs(input_layers, mode)
- model = tf.keras.models.Model(
+ model = tf_keras.models.Model(
inputs=input_layers, outputs=outputs, name='maskrcnn')
- assert model is not None, 'Fail to build tf.keras.Model.'
+ assert model is not None, 'Fail to build tf_keras.Model.'
model.optimizer = self.build_optimizer()
self._keras_model = model
diff --git a/official/legacy/detection/modeling/olnmask_model.py b/official/legacy/detection/modeling/olnmask_model.py
index 255ff86e4f6..7feca2b1b38 100644
--- a/official/legacy/detection/modeling/olnmask_model.py
+++ b/official/legacy/detection/modeling/olnmask_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import anchor
from official.legacy.detection.dataloader import mode_keys
@@ -368,31 +368,31 @@ def build_input_layers(self, params, mode):
batch_size = params.train.batch_size
input_layer = {
'image':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=input_shape,
batch_size=batch_size,
name='image',
dtype=tf.bfloat16 if self._use_bfloat16 else tf.float32),
'image_info':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[4, 2],
batch_size=batch_size,
name='image_info',
),
'gt_boxes':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[params.olnmask_parser.max_num_instances, 4],
batch_size=batch_size,
name='gt_boxes'),
'gt_classes':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[params.olnmask_parser.max_num_instances],
batch_size=batch_size,
name='gt_classes',
dtype=tf.int64),
}
if self._include_mask:
- input_layer['gt_masks'] = tf.keras.layers.Input(
+ input_layer['gt_masks'] = tf_keras.layers.Input(
shape=[
params.olnmask_parser.max_num_instances,
params.olnmask_parser.mask_crop_size,
@@ -404,13 +404,13 @@ def build_input_layers(self, params, mode):
batch_size = params.eval.batch_size
input_layer = {
'image':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=input_shape,
batch_size=batch_size,
name='image',
dtype=tf.bfloat16 if self._use_bfloat16 else tf.float32),
'image_info':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[4, 2],
batch_size=batch_size,
name='image_info',
@@ -423,9 +423,9 @@ def build_model(self, params, mode):
input_layers = self.build_input_layers(self._params, mode)
outputs = self.model_outputs(input_layers, mode)
- model = tf.keras.models.Model(
+ model = tf_keras.models.Model(
inputs=input_layers, outputs=outputs, name='olnmask')
- assert model is not None, 'Fail to build tf.keras.Model.'
+ assert model is not None, 'Fail to build tf_keras.Model.'
model.optimizer = self.build_optimizer()
self._keras_model = model
diff --git a/official/legacy/detection/modeling/optimizers.py b/official/legacy/detection/modeling/optimizers.py
index d8ff456aa77..7c676f5e61b 100644
--- a/official/legacy/detection/modeling/optimizers.py
+++ b/official/legacy/detection/modeling/optimizers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class OptimizerFactory(object):
@@ -30,18 +30,18 @@ def __init__(self, params):
"""Creates optimized based on the specified flags."""
if params.type == 'momentum':
self._optimizer = functools.partial(
- tf.keras.optimizers.SGD,
+ tf_keras.optimizers.SGD,
momentum=params.momentum,
nesterov=params.nesterov)
elif params.type == 'adam':
- self._optimizer = tf.keras.optimizers.Adam
+ self._optimizer = tf_keras.optimizers.Adam
elif params.type == 'adadelta':
- self._optimizer = tf.keras.optimizers.Adadelta
+ self._optimizer = tf_keras.optimizers.Adadelta
elif params.type == 'adagrad':
- self._optimizer = tf.keras.optimizers.Adagrad
+ self._optimizer = tf_keras.optimizers.Adagrad
elif params.type == 'rmsprop':
self._optimizer = functools.partial(
- tf.keras.optimizers.RMSprop, momentum=params.momentum)
+ tf_keras.optimizers.RMSprop, momentum=params.momentum)
else:
raise ValueError('Unsupported optimizer type `{}`.'.format(params.type))
diff --git a/official/legacy/detection/modeling/retinanet_model.py b/official/legacy/detection/modeling/retinanet_model.py
index 7e87717cc9e..86ea936615b 100644
--- a/official/legacy/detection/modeling/retinanet_model.py
+++ b/official/legacy/detection/modeling/retinanet_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import mode_keys
from official.legacy.detection.evaluation import factory as eval_factory
@@ -57,7 +57,7 @@ def __init__(self, params):
self._transpose_input = params.train.transpose_input
assert not self._transpose_input, 'Transpose input is not supported.'
# Input layer.
- self._input_layer = tf.keras.layers.Input(
+ self._input_layer = tf_keras.layers.Input(
shape=(None, None, params.retinanet_parser.num_channels),
name='',
dtype=tf.bfloat16 if self._use_bfloat16 else tf.float32)
@@ -118,9 +118,9 @@ def build_model(self, params, mode=None):
if self._keras_model is None:
outputs = self.model_outputs(self._input_layer, mode)
- model = tf.keras.models.Model(
+ model = tf_keras.models.Model(
inputs=self._input_layer, outputs=outputs, name='retinanet')
- assert model is not None, 'Fail to build tf.keras.Model.'
+ assert model is not None, 'Fail to build tf_keras.Model.'
model.optimizer = self.build_optimizer()
self._keras_model = model
diff --git a/official/legacy/detection/modeling/shapemask_model.py b/official/legacy/detection/modeling/shapemask_model.py
index 6d01e122b0b..3443e346f1e 100644
--- a/official/legacy/detection/modeling/shapemask_model.py
+++ b/official/legacy/detection/modeling/shapemask_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.dataloader import anchor
from official.legacy.detection.dataloader import mode_keys
@@ -103,7 +103,7 @@ def build_outputs(self, inputs, mode):
# Wrapping if else code paths into a layer to make the checkpoint loadable
# in prediction mode.
- class SampledBoxesLayer(tf.keras.layers.Layer):
+ class SampledBoxesLayer(tf_keras.layers.Layer):
"""ShapeMask model function."""
def call(self, inputs, val_boxes, val_classes, val_outer_boxes, training):
@@ -210,28 +210,28 @@ def build_input_layers(self, params, mode):
batch_size = params.train.batch_size
input_layer = {
'image':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=input_shape,
batch_size=batch_size,
name='image',
dtype=tf.bfloat16 if self._use_bfloat16 else tf.float32),
'image_info':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[4, 2], batch_size=batch_size, name='image_info'),
'mask_classes':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[params.shapemask_parser.num_sampled_masks],
batch_size=batch_size,
name='mask_classes',
dtype=tf.int64),
'mask_outer_boxes':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[params.shapemask_parser.num_sampled_masks, 4],
batch_size=batch_size,
name='mask_outer_boxes',
dtype=tf.float32),
'mask_boxes':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[params.shapemask_parser.num_sampled_masks, 4],
batch_size=batch_size,
name='mask_boxes',
@@ -241,13 +241,13 @@ def build_input_layers(self, params, mode):
batch_size = params.eval.batch_size
input_layer = {
'image':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=input_shape,
batch_size=batch_size,
name='image',
dtype=tf.bfloat16 if self._use_bfloat16 else tf.float32),
'image_info':
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=[4, 2], batch_size=batch_size, name='image_info'),
}
return input_layer
@@ -257,9 +257,9 @@ def build_model(self, params, mode):
input_layers = self.build_input_layers(self._params, mode)
outputs = self.model_outputs(input_layers, mode)
- model = tf.keras.models.Model(
+ model = tf_keras.models.Model(
inputs=input_layers, outputs=outputs, name='shapemask')
- assert model is not None, 'Fail to build tf.keras.Model.'
+ assert model is not None, 'Fail to build tf_keras.Model.'
model.optimizer = self.build_optimizer()
self._keras_model = model
diff --git a/official/legacy/detection/ops/__init__.py b/official/legacy/detection/ops/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/ops/__init__.py
+++ b/official/legacy/detection/ops/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/ops/nms.py b/official/legacy/detection/ops/nms.py
index 24fdcef87bc..1f9ca465eb6 100644
--- a/official/legacy/detection/ops/nms.py
+++ b/official/legacy/detection/ops/nms.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.utils import box_utils
@@ -67,8 +67,9 @@ def _suppression_loop_body(boxes, iou_threshold, output_size, idx):
output_size: the updated output_size.
idx: the updated induction variable.
"""
- num_tiles = tf.shape(boxes)[1] // NMS_TILE_SIZE
- batch_size = tf.shape(boxes)[0]
+ boxes_shape = tf.shape(boxes)
+ num_tiles = boxes_shape[1] // NMS_TILE_SIZE
+ batch_size = boxes_shape[0]
# Iterates over tiles that can possibly suppress the current tile.
box_slice = tf.slice(boxes, [0, idx * NMS_TILE_SIZE, 0],
@@ -97,7 +98,7 @@ def _suppression_loop_body(boxes, iou_threshold, output_size, idx):
boxes = tf.tile(tf.expand_dims(
box_slice, [1]), [1, num_tiles, 1, 1]) * mask + tf.reshape(
boxes, [batch_size, num_tiles, NMS_TILE_SIZE, 4]) * (1 - mask)
- boxes = tf.reshape(boxes, [batch_size, -1, 4])
+ boxes = tf.reshape(boxes, boxes_shape)
# Updates output_size.
output_size += tf.reduce_sum(
diff --git a/official/legacy/detection/ops/postprocess_ops.py b/official/legacy/detection/ops/postprocess_ops.py
index 8b4a8b6d9f0..b755e8277f9 100644
--- a/official/legacy/detection/ops/postprocess_ops.py
+++ b/official/legacy/detection/ops/postprocess_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.ops import nms
from official.legacy.detection.utils import box_utils
@@ -291,7 +291,7 @@ def _generate_detections_batched(boxes, scores, max_total_size,
return nmsed_boxes, nmsed_scores, nmsed_classes, valid_detections
-class MultilevelDetectionGenerator(tf.keras.layers.Layer):
+class MultilevelDetectionGenerator(tf_keras.layers.Layer):
"""Generates detected boxes with scores and classes for one-stage detector."""
def __init__(self, min_level, max_level, params):
@@ -338,7 +338,7 @@ def call(self, box_outputs, class_outputs, anchor_boxes, image_shape):
return nmsed_boxes, nmsed_scores, nmsed_classes, valid_detections
-class GenericDetectionGenerator(tf.keras.layers.Layer):
+class GenericDetectionGenerator(tf_keras.layers.Layer):
"""Generates the final detected boxes with scores and classes."""
def __init__(self, params):
diff --git a/official/legacy/detection/ops/roi_ops.py b/official/legacy/detection/ops/roi_ops.py
index 7aeb1a91b1f..f17cf54644b 100644
--- a/official/legacy/detection/ops/roi_ops.py
+++ b/official/legacy/detection/ops/roi_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.ops import nms
from official.legacy.detection.utils import box_utils
@@ -170,7 +170,7 @@ def multilevel_propose_rois(rpn_boxes,
return selected_rois, selected_roi_scores
-class ROIGenerator(tf.keras.layers.Layer):
+class ROIGenerator(tf_keras.layers.Layer):
"""Proposes RoIs for the second stage processing."""
def __init__(self, params):
diff --git a/official/legacy/detection/ops/spatial_transform_ops.py b/official/legacy/detection/ops/spatial_transform_ops.py
index db9cf98fb80..19c8510afad 100644
--- a/official/legacy/detection/ops/spatial_transform_ops.py
+++ b/official/legacy/detection/ops/spatial_transform_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
_EPSILON = 1e-8
diff --git a/official/legacy/detection/ops/target_ops.py b/official/legacy/detection/ops/target_ops.py
index 7b8e208b99b..9f00c93bab4 100644
--- a/official/legacy/detection/ops/target_ops.py
+++ b/official/legacy/detection/ops/target_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.ops import spatial_transform_ops
from official.legacy.detection.utils import box_utils
@@ -292,7 +292,7 @@ def sample_and_crop_foreground_masks(candidate_rois,
return foreground_rois, foreground_classes, cropped_foreground_masks
-class ROISampler(tf.keras.layers.Layer):
+class ROISampler(tf_keras.layers.Layer):
"""Samples RoIs and creates training targets."""
def __init__(self, params):
@@ -517,7 +517,7 @@ def assign_and_sample_proposals_and_scores(self,
sampled_gt_classes, sampled_gt_indices)
-class MaskSampler(tf.keras.layers.Layer):
+class MaskSampler(tf_keras.layers.Layer):
"""Samples and creates mask training targets."""
def __init__(self, mask_target_size, num_mask_samples_per_image):
diff --git a/official/legacy/detection/utils/__init__.py b/official/legacy/detection/utils/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/detection/utils/__init__.py
+++ b/official/legacy/detection/utils/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/utils/box_utils.py b/official/legacy/detection/utils/box_utils.py
index f52b4d52c12..8449844826c 100644
--- a/official/legacy/detection/utils/box_utils.py
+++ b/official/legacy/detection/utils/box_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from __future__ import print_function
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
EPSILON = 1e-8
BBOX_XFORM_CLIP = np.log(1000. / 16.)
diff --git a/official/legacy/detection/utils/class_utils.py b/official/legacy/detection/utils/class_utils.py
index fe5525c6926..adfee9d1e53 100644
--- a/official/legacy/detection/utils/class_utils.py
+++ b/official/legacy/detection/utils/class_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/detection/utils/dataloader_utils.py b/official/legacy/detection/utils/dataloader_utils.py
index a3a34eb658b..99d500e8eee 100644
--- a/official/legacy/detection/utils/dataloader_utils.py
+++ b/official/legacy/detection/utils/dataloader_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Utility functions for dataloader."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.utils import input_utils
diff --git a/official/legacy/detection/utils/input_utils.py b/official/legacy/detection/utils/input_utils.py
index 12b7c0be168..5df094cbba7 100644
--- a/official/legacy/detection/utils/input_utils.py
+++ b/official/legacy/detection/utils/input_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import math
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.detection.utils import box_utils
from official.vision.utils.object_detection import preprocessor
diff --git a/official/legacy/detection/utils/mask_utils.py b/official/legacy/detection/utils/mask_utils.py
index deb86a51605..d25eb9a9f96 100644
--- a/official/legacy/detection/utils/mask_utils.py
+++ b/official/legacy/detection/utils/mask_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,8 +19,8 @@
import math
-import numpy as np
import cv2
+import numpy as np
def paste_instance_masks(masks, detected_boxes, image_height, image_width):
diff --git a/official/legacy/image_classification/__init__.py b/official/legacy/image_classification/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/image_classification/__init__.py
+++ b/official/legacy/image_classification/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/image_classification/augment.py b/official/legacy/image_classification/augment.py
index add7ed631ca..a5fac371182 100644
--- a/official/legacy/image_classification/augment.py
+++ b/official/legacy/image_classification/augment.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,8 +25,7 @@
import math
from typing import Any, Dict, List, Optional, Text, Tuple
-from keras.layers.preprocessing import image_preprocessing as image_ops
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# This signifies the max integer that the controller RNN could predict for the
@@ -163,6 +162,86 @@ def _convert_angles_to_transform(angles: tf.Tensor, image_width: tf.Tensor,
)
+def apply_transform_to_images(
+ images,
+ transforms,
+ fill_mode='reflect',
+ fill_value=0.0,
+ interpolation='bilinear',
+ output_shape=None,
+ name=None,
+):
+ """Applies the given transform(s) to the image(s).
+
+ Args:
+ images: A tensor of shape `(num_images, num_rows, num_columns,
+ num_channels)` (NHWC). The rank must be statically known (the shape is
+ not `TensorShape(None)`).
+ transforms: Projective transform matrix/matrices. A vector of length 8 or
+ tensor of size N x 8. If one row of transforms is [a0, a1, a2, b0, b1,
+ b2, c0, c1], then it maps the *output* point `(x, y)` to a transformed
+ *input* point `(x', y') = ((a0 x + a1 y + a2) / k, (b0 x + b1 y + b2) /
+ k)`, where `k = c0 x + c1 y + 1`. The transforms are *inverted* compared
+ to the transform mapping input points to output points. Note that
+ gradients are not backpropagated into transformation parameters.
+ fill_mode: Points outside the boundaries of the input are filled according
+ to the given mode (one of `{"constant", "reflect", "wrap", "nearest"}`).
+ fill_value: a float represents the value to be filled outside the
+ boundaries when `fill_mode="constant"`.
+ interpolation: Interpolation mode. Supported values: `"nearest"`,
+ `"bilinear"`.
+ output_shape: Output dimension after the transform, `[height, width]`. If
+ `None`, output is the same size as input image.
+ name: The name of the op. Fill mode behavior for each valid value is as
+ follows
+ - `"reflect"`: `(d c b a | a b c d | d c b a)` The input is extended by
+ reflecting about the edge of the last pixel.
+ - `"constant"`: `(k k k k | a b c d | k k k k)` The input is extended by
+ filling all values beyond the edge with the same constant value k = 0.
+ - `"wrap"`: `(a b c d | a b c d | a b c d)` The input is extended by
+ wrapping around to the opposite edge.
+ - `"nearest"`: `(a a a a | a b c d | d d d d)` The input is extended by
+ the nearest pixel. Input shape: 4D tensor with shape:
+ `(samples, height, width, channels)`, in `"channels_last"` format.
+ Output shape: 4D tensor with shape: `(samples, height, width, channels)`,
+ in `"channels_last"` format.
+
+ Returns:
+ Image(s) with the same type and shape as `images`, with the given
+ transform(s) applied. Transformed coordinates outside of the input image
+ will be filled with zeros.
+ """
+ with tf.name_scope(name or 'transform'):
+ if output_shape is None:
+ output_shape = tf.shape(images)[1:3]
+ if not tf.executing_eagerly():
+ output_shape_value = tf.get_static_value(output_shape)
+ if output_shape_value is not None:
+ output_shape = output_shape_value
+
+ output_shape = tf.convert_to_tensor(
+ output_shape, tf.int32, name='output_shape'
+ )
+
+ if not output_shape.get_shape().is_compatible_with([2]):
+ raise ValueError(
+ 'output_shape must be a 1-D Tensor of 2 elements: '
+ 'new_height, new_width, instead got '
+ f'output_shape={output_shape}'
+ )
+
+ fill_value = tf.convert_to_tensor(fill_value, tf.float32, name='fill_value')
+
+ return tf.raw_ops.ImageProjectiveTransformV3(
+ images=images,
+ output_shape=output_shape,
+ fill_value=fill_value,
+ transforms=transforms,
+ fill_mode=fill_mode.upper(),
+ interpolation=interpolation.upper(),
+ )
+
+
def transform(image: tf.Tensor, transforms) -> tf.Tensor:
"""Prepares input data for `image_ops.transform`."""
original_ndims = tf.rank(image)
@@ -170,8 +249,9 @@ def transform(image: tf.Tensor, transforms) -> tf.Tensor:
if transforms.shape.rank == 1:
transforms = transforms[None]
image = to_4d(image)
- image = image_ops.transform(
- images=image, transforms=transforms, interpolation='nearest')
+ image = apply_transform_to_images(
+ images=image, transforms=transforms, interpolation='nearest'
+ )
return from_4d(image, original_ndims)
@@ -186,7 +266,7 @@ def translate(image: tf.Tensor, translations) -> tf.Tensor:
The translated version of the image.
"""
- transforms = _convert_translation_to_transform(translations)
+ transforms = _convert_translation_to_transform(translations) # pytype: disable=wrong-arg-types # always-use-return-annotations
return transform(image, transforms=transforms)
@@ -839,10 +919,6 @@ def policy_v0():
the policy.
"""
- # TODO(dankondratyuk): tensorflow_addons defines custom ops, which
- # for some reason are not included when building/linking
- # This results in the error, "Op type not registered
- # 'Addons>ImageProjectiveTransformV2' in binary" when running on borg TPUs
policy = [
[('Equalize', 0.8, 1), ('ShearY', 0.8, 4)],
[('Color', 0.4, 9), ('Equalize', 0.6, 3)],
diff --git a/official/legacy/image_classification/augment_test.py b/official/legacy/image_classification/augment_test.py
index 139e10195b4..1ba8a43c86d 100644
--- a/official/legacy/image_classification/augment_test.py
+++ b/official/legacy/image_classification/augment_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification import augment
diff --git a/official/legacy/image_classification/callbacks.py b/official/legacy/image_classification/callbacks.py
index 061826dbd05..43a32095528 100644
--- a/official/legacy/image_classification/callbacks.py
+++ b/official/legacy/image_classification/callbacks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from typing import Any, List, MutableMapping, Optional, Text
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import optimization
from official.utils.misc import keras_utils
@@ -38,19 +38,19 @@ def get_callbacks(
batch_size: int = 0,
log_steps: int = 0,
model_dir: Optional[str] = None,
- backup_and_restore: bool = False) -> List[tf.keras.callbacks.Callback]:
+ backup_and_restore: bool = False) -> List[tf_keras.callbacks.Callback]:
"""Get all callbacks."""
model_dir = model_dir or ''
callbacks = []
if model_checkpoint:
ckpt_full_path = os.path.join(model_dir, 'model.ckpt-{epoch:04d}')
callbacks.append(
- tf.keras.callbacks.ModelCheckpoint(
+ tf_keras.callbacks.ModelCheckpoint(
ckpt_full_path, save_weights_only=True, verbose=1))
if backup_and_restore:
backup_dir = os.path.join(model_dir, 'tmp')
callbacks.append(
- tf.keras.callbacks.experimental.BackupAndRestore(backup_dir))
+ tf_keras.callbacks.experimental.BackupAndRestore(backup_dir))
if include_tensorboard:
callbacks.append(
CustomTensorBoard(
@@ -82,14 +82,14 @@ def get_callbacks(
def get_scalar_from_tensor(t: tf.Tensor) -> int:
"""Utility function to convert a Tensor to a scalar."""
- t = tf.keras.backend.get_value(t)
+ t = tf_keras.backend.get_value(t)
if callable(t):
- return t()
+ return t() # pyrefly: ignore[bad-return]
else:
- return t
+ return t # pyrefly: ignore[bad-return]
-class CustomTensorBoard(tf.keras.callbacks.TensorBoard):
+class CustomTensorBoard(tf_keras.callbacks.TensorBoard):
"""A customized TensorBoard callback that tracks additional datapoints.
Metrics tracked:
@@ -157,7 +157,7 @@ def _calculate_lr(self) -> int:
return get_scalar_from_tensor(
self._get_base_optimizer()._decayed_lr(var_dtype=tf.float32)) # pylint:disable=protected-access
- def _get_base_optimizer(self) -> tf.keras.optimizers.Optimizer:
+ def _get_base_optimizer(self) -> tf_keras.optimizers.Optimizer:
"""Get the base optimizer used by the current model."""
optimizer = self.model.optimizer
@@ -169,7 +169,7 @@ def _get_base_optimizer(self) -> tf.keras.optimizers.Optimizer:
return optimizer
-class MovingAverageCallback(tf.keras.callbacks.Callback):
+class MovingAverageCallback(tf_keras.callbacks.Callback):
"""A Callback to be used with a `ExponentialMovingAverage` optimizer.
Applies moving average weights to the model during validation time to test
@@ -187,7 +187,7 @@ def __init__(self, overwrite_weights_on_train_end: bool = False, **kwargs):
super(MovingAverageCallback, self).__init__(**kwargs)
self.overwrite_weights_on_train_end = overwrite_weights_on_train_end
- def set_model(self, model: tf.keras.Model):
+ def set_model(self, model: tf_keras.Model):
super(MovingAverageCallback, self).set_model(model)
assert isinstance(self.model.optimizer,
optimization.ExponentialMovingAverage)
@@ -204,7 +204,7 @@ def on_train_end(self, logs: Optional[MutableMapping[Text, Any]] = None):
self.model.optimizer.assign_average_vars(self.model.variables)
-class AverageModelCheckpoint(tf.keras.callbacks.ModelCheckpoint):
+class AverageModelCheckpoint(tf_keras.callbacks.ModelCheckpoint):
"""Saves and, optionally, assigns the averaged weights.
Taken from tfa.callbacks.AverageModelCheckpoint.
@@ -212,7 +212,7 @@ class AverageModelCheckpoint(tf.keras.callbacks.ModelCheckpoint):
Attributes:
update_weights: If True, assign the moving average weights to the model, and
save them. If False, keep the old non-averaged weights, but the saved
- model uses the average weights. See `tf.keras.callbacks.ModelCheckpoint`
+ model uses the average weights. See `tf_keras.callbacks.ModelCheckpoint`
for the other args.
"""
diff --git a/official/legacy/image_classification/classifier_trainer.py b/official/legacy/image_classification/classifier_trainer.py
index 66577f6079e..a39b7597b42 100644
--- a/official/legacy/image_classification/classifier_trainer.py
+++ b/official/legacy/image_classification/classifier_trainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.image_classification import callbacks as custom_callbacks
from official.legacy.image_classification import dataset_factory
@@ -38,7 +38,7 @@
from official.utils.misc import keras_utils
-def get_models() -> Mapping[str, tf.keras.Model]:
+def get_models() -> Mapping[str, tf_keras.Model]:
"""Returns the mapping from model type name to Keras model."""
return {
'efficientnet': efficientnet_model.EfficientNet.from_name,
@@ -64,26 +64,26 @@ def _get_metrics(one_hot: bool) -> Mapping[Text, Any]:
return {
# (name, metric_fn)
'acc':
- tf.keras.metrics.CategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
'accuracy':
- tf.keras.metrics.CategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
'top_1':
- tf.keras.metrics.CategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
'top_5':
- tf.keras.metrics.TopKCategoricalAccuracy(
+ tf_keras.metrics.TopKCategoricalAccuracy(
k=5, name='top_5_accuracy'),
}
else:
return {
# (name, metric_fn)
'acc':
- tf.keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
'accuracy':
- tf.keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
'top_1':
- tf.keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
'top_5':
- tf.keras.metrics.SparseTopKCategoricalAccuracy(
+ tf_keras.metrics.SparseTopKCategoricalAccuracy(
k=5, name='top_5_accuracy'),
}
@@ -192,7 +192,7 @@ def _get_params_from_flags(flags_obj: flags.FlagValues):
return params
-def resume_from_checkpoint(model: tf.keras.Model, model_dir: str,
+def resume_from_checkpoint(model: tf_keras.Model, model_dir: str,
train_steps: int) -> int:
"""Resumes from the latest checkpoint, if possible.
@@ -233,7 +233,7 @@ def initialize(params: base_configs.ExperimentConfig,
data_format = 'channels_first'
else:
data_format = 'channels_last'
- tf.keras.backend.set_image_data_format(data_format)
+ tf_keras.backend.set_image_data_format(data_format)
if params.runtime.run_eagerly:
# Enable eager execution to allow step-by-step debugging
tf.config.experimental_run_functions_eagerly(True)
@@ -348,10 +348,10 @@ def train_and_eval(
steps_per_loop = train_steps if params.train.set_epoch_loop else 1
if one_hot:
- loss_obj = tf.keras.losses.CategoricalCrossentropy(
+ loss_obj = tf_keras.losses.CategoricalCrossentropy(
label_smoothing=params.model.loss.label_smoothing)
else:
- loss_obj = tf.keras.losses.SparseCategoricalCrossentropy()
+ loss_obj = tf_keras.losses.SparseCategoricalCrossentropy()
model.compile(
optimizer=optimizer,
loss=loss_obj,
diff --git a/official/legacy/image_classification/classifier_trainer_test.py b/official/legacy/image_classification/classifier_trainer_test.py
index 2be5d85727f..c8c896afae6 100644
--- a/official/legacy/image_classification/classifier_trainer_test.py
+++ b/official/legacy/image_classification/classifier_trainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,7 +25,7 @@
from absl import flags
from absl.testing import flagsaver
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
diff --git a/official/legacy/image_classification/classifier_trainer_util_test.py b/official/legacy/image_classification/classifier_trainer_util_test.py
index 19a05fa678b..2c6c7cb58eb 100644
--- a/official/legacy/image_classification/classifier_trainer_util_test.py
+++ b/official/legacy/image_classification/classifier_trainer_util_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,7 +22,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification import classifier_trainer
from official.legacy.image_classification import dataset_factory
@@ -30,12 +30,12 @@
from official.legacy.image_classification.configs import base_configs
-def get_trivial_model(num_classes: int) -> tf.keras.Model:
+def get_trivial_model(num_classes: int) -> tf_keras.Model:
"""Creates and compiles trivial model for ImageNet dataset."""
model = test_utils.trivial_model(num_classes=num_classes)
lr = 0.01
- optimizer = tf.keras.optimizers.SGD(learning_rate=lr)
- loss_obj = tf.keras.losses.SparseCategoricalCrossentropy()
+ optimizer = tf_keras.optimizers.SGD(learning_rate=lr)
+ loss_obj = tf_keras.losses.SparseCategoricalCrossentropy()
model.compile(optimizer=optimizer, loss=loss_obj, run_eagerly=True)
return model
@@ -120,7 +120,7 @@ class EmptyClass:
def test_resume_from_checkpoint(self):
"""Tests functionality for resuming from checkpoint."""
# Set the keras policy
- tf.keras.mixed_precision.set_global_policy('mixed_bfloat16')
+ tf_keras.mixed_precision.set_global_policy('mixed_bfloat16')
# Get the model, datasets, and compile it.
model = get_trivial_model(10)
@@ -131,7 +131,7 @@ def test_resume_from_checkpoint(self):
train_steps = 10
ds = get_trivial_data()
callbacks = [
- tf.keras.callbacks.ModelCheckpoint(
+ tf_keras.callbacks.ModelCheckpoint(
os.path.join(model_dir, 'model.ckpt-{epoch:04d}'),
save_weights_only=True)
]
diff --git a/official/legacy/image_classification/configs/__init__.py b/official/legacy/image_classification/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/image_classification/configs/__init__.py
+++ b/official/legacy/image_classification/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/image_classification/configs/base_configs.py b/official/legacy/image_classification/configs/base_configs.py
index 7fd230b418e..4237153d14b 100644
--- a/official/legacy/image_classification/configs/base_configs.py
+++ b/official/legacy/image_classification/configs/base_configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -112,10 +112,16 @@ class TrainConfig(hyperparams.Config):
resume_checkpoint: bool = None
epochs: int = None
steps: int = None
- callbacks: CallbacksConfig = CallbacksConfig()
+ callbacks: CallbacksConfig = dataclasses.field(
+ default_factory=CallbacksConfig
+ )
metrics: MetricsConfig = None
- tensorboard: TensorBoardConfig = TensorBoardConfig()
- time_history: TimeHistoryConfig = TimeHistoryConfig()
+ tensorboard: TensorBoardConfig = dataclasses.field(
+ default_factory=TensorBoardConfig
+ )
+ time_history: TimeHistoryConfig = dataclasses.field(
+ default_factory=TimeHistoryConfig
+ )
set_epoch_loop: bool = False
diff --git a/official/legacy/image_classification/configs/configs.py b/official/legacy/image_classification/configs/configs.py
index 87fb5df5b6f..175a0f99fbb 100644
--- a/official/legacy/image_classification/configs/configs.py
+++ b/official/legacy/image_classification/configs/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -35,83 +35,138 @@ class EfficientNetImageNetConfig(base_configs.ExperimentConfig):
evaluation: An `EvalConfig` instance.
model: A `ModelConfig` instance.
"""
- export: base_configs.ExportConfig = base_configs.ExportConfig()
- runtime: base_configs.RuntimeConfig = base_configs.RuntimeConfig()
- train_dataset: dataset_factory.DatasetConfig = dataset_factory.ImageNetConfig(
- split='train')
- validation_dataset: dataset_factory.DatasetConfig = dataset_factory.ImageNetConfig(
- split='validation')
- train: base_configs.TrainConfig = base_configs.TrainConfig(
- resume_checkpoint=True,
- epochs=500,
- steps=None,
- callbacks=base_configs.CallbacksConfig(
- enable_checkpoint_and_export=True, enable_tensorboard=True),
- metrics=['accuracy', 'top_5'],
- time_history=base_configs.TimeHistoryConfig(log_steps=100),
- tensorboard=base_configs.TensorBoardConfig(
- track_lr=True, write_model_weights=False),
- set_epoch_loop=False)
- evaluation: base_configs.EvalConfig = base_configs.EvalConfig(
- epochs_between_evals=1, steps=None)
- model: base_configs.ModelConfig = efficientnet_config.EfficientNetModelConfig(
+ export: base_configs.ExportConfig = dataclasses.field(
+ default_factory=base_configs.ExportConfig
+ )
+ runtime: base_configs.RuntimeConfig = dataclasses.field(
+ default_factory=base_configs.RuntimeConfig
+ )
+ train_dataset: dataset_factory.DatasetConfig = dataclasses.field(
+ default_factory=lambda: dataset_factory.ImageNetConfig(split='train')
+ )
+ validation_dataset: dataset_factory.DatasetConfig = dataclasses.field(
+ default_factory=lambda: dataset_factory.ImageNetConfig(split='validation')
+ )
+ train: base_configs.TrainConfig = dataclasses.field(
+ default_factory=lambda: base_configs.TrainConfig( # pylint: disable=g-long-lambda
+ resume_checkpoint=True,
+ epochs=500,
+ steps=None,
+ callbacks=base_configs.CallbacksConfig(
+ enable_checkpoint_and_export=True, enable_tensorboard=True
+ ),
+ metrics=['accuracy', 'top_5'],
+ time_history=base_configs.TimeHistoryConfig(log_steps=100),
+ tensorboard=base_configs.TensorBoardConfig(
+ track_lr=True, write_model_weights=False
+ ),
+ set_epoch_loop=False,
+ )
+ )
+ evaluation: base_configs.EvalConfig = dataclasses.field(
+ default_factory=lambda: base_configs.EvalConfig( # pylint: disable=g-long-lambda
+ epochs_between_evals=1, steps=None
+ )
+ )
+ model: base_configs.ModelConfig = dataclasses.field(
+ default_factory=efficientnet_config.EfficientNetModelConfig
)
@dataclasses.dataclass
class ResNetImagenetConfig(base_configs.ExperimentConfig):
"""Base configuration to train resnet-50 on ImageNet."""
- export: base_configs.ExportConfig = base_configs.ExportConfig()
- runtime: base_configs.RuntimeConfig = base_configs.RuntimeConfig()
- train_dataset: dataset_factory.DatasetConfig = \
- dataset_factory.ImageNetConfig(split='train',
- one_hot=False,
- mean_subtract=True,
- standardize=True)
- validation_dataset: dataset_factory.DatasetConfig = \
- dataset_factory.ImageNetConfig(split='validation',
- one_hot=False,
- mean_subtract=True,
- standardize=True)
- train: base_configs.TrainConfig = base_configs.TrainConfig(
- resume_checkpoint=True,
- epochs=90,
- steps=None,
- callbacks=base_configs.CallbacksConfig(
- enable_checkpoint_and_export=True, enable_tensorboard=True),
- metrics=['accuracy', 'top_5'],
- time_history=base_configs.TimeHistoryConfig(log_steps=100),
- tensorboard=base_configs.TensorBoardConfig(
- track_lr=True, write_model_weights=False),
- set_epoch_loop=False)
- evaluation: base_configs.EvalConfig = base_configs.EvalConfig(
- epochs_between_evals=1, steps=None)
- model: base_configs.ModelConfig = resnet_config.ResNetModelConfig()
+ export: base_configs.ExportConfig = dataclasses.field(
+ default_factory=base_configs.ExportConfig
+ )
+ runtime: base_configs.RuntimeConfig = dataclasses.field(
+ default_factory=base_configs.RuntimeConfig
+ )
+ train_dataset: dataset_factory.DatasetConfig = dataclasses.field(
+ default_factory=lambda: dataset_factory.ImageNetConfig( # pylint: disable=g-long-lambda
+ split='train', one_hot=False, mean_subtract=True, standardize=True
+ )
+ )
+ validation_dataset: dataset_factory.DatasetConfig = dataclasses.field(
+ default_factory=lambda: dataset_factory.ImageNetConfig( # pylint: disable=g-long-lambda
+ split='validation',
+ one_hot=False,
+ mean_subtract=True,
+ standardize=True,
+ )
+ )
+ train: base_configs.TrainConfig = dataclasses.field(
+ default_factory=lambda: base_configs.TrainConfig( # pylint: disable=g-long-lambda
+ resume_checkpoint=True,
+ epochs=90,
+ steps=None,
+ callbacks=base_configs.CallbacksConfig(
+ enable_checkpoint_and_export=True, enable_tensorboard=True
+ ),
+ metrics=['accuracy', 'top_5'],
+ time_history=base_configs.TimeHistoryConfig(log_steps=100),
+ tensorboard=base_configs.TensorBoardConfig(
+ track_lr=True, write_model_weights=False
+ ),
+ set_epoch_loop=False,
+ )
+ )
+ evaluation: base_configs.EvalConfig = dataclasses.field(
+ default_factory=lambda: base_configs.EvalConfig( # pylint: disable=g-long-lambda
+ epochs_between_evals=1, steps=None
+ )
+ )
+ model: base_configs.ModelConfig = dataclasses.field(
+ default_factory=resnet_config.ResNetModelConfig
+ )
@dataclasses.dataclass
class VGGImagenetConfig(base_configs.ExperimentConfig):
"""Base configuration to train vgg-16 on ImageNet."""
- export: base_configs.ExportConfig = base_configs.ExportConfig()
- runtime: base_configs.RuntimeConfig = base_configs.RuntimeConfig()
- train_dataset: dataset_factory.DatasetConfig = dataset_factory.ImageNetConfig(
- split='train', one_hot=False, mean_subtract=True, standardize=True)
- validation_dataset: dataset_factory.DatasetConfig = dataset_factory.ImageNetConfig(
- split='validation', one_hot=False, mean_subtract=True, standardize=True)
- train: base_configs.TrainConfig = base_configs.TrainConfig(
- resume_checkpoint=True,
- epochs=90,
- steps=None,
- callbacks=base_configs.CallbacksConfig(
- enable_checkpoint_and_export=True, enable_tensorboard=True),
- metrics=['accuracy', 'top_5'],
- time_history=base_configs.TimeHistoryConfig(log_steps=100),
- tensorboard=base_configs.TensorBoardConfig(
- track_lr=True, write_model_weights=False),
- set_epoch_loop=False)
- evaluation: base_configs.EvalConfig = base_configs.EvalConfig(
- epochs_between_evals=1, steps=None)
- model: base_configs.ModelConfig = vgg_config.VGGModelConfig()
+ export: base_configs.ExportConfig = dataclasses.field(
+ default_factory=base_configs.ExportConfig
+ )
+ runtime: base_configs.RuntimeConfig = dataclasses.field(
+ default_factory=base_configs.RuntimeConfig
+ )
+ train_dataset: dataset_factory.DatasetConfig = dataclasses.field(
+ default_factory=lambda: dataset_factory.ImageNetConfig( # pylint: disable=g-long-lambda
+ split='train', one_hot=False, mean_subtract=True, standardize=True
+ )
+ )
+ validation_dataset: dataset_factory.DatasetConfig = dataclasses.field(
+ default_factory=lambda: dataset_factory.ImageNetConfig( # pylint: disable=g-long-lambda
+ split='validation',
+ one_hot=False,
+ mean_subtract=True,
+ standardize=True,
+ )
+ )
+ train: base_configs.TrainConfig = dataclasses.field(
+ default_factory=lambda: base_configs.TrainConfig( # pylint: disable=g-long-lambda
+ resume_checkpoint=True,
+ epochs=90,
+ steps=None,
+ callbacks=base_configs.CallbacksConfig(
+ enable_checkpoint_and_export=True, enable_tensorboard=True
+ ),
+ metrics=['accuracy', 'top_5'],
+ time_history=base_configs.TimeHistoryConfig(log_steps=100),
+ tensorboard=base_configs.TensorBoardConfig(
+ track_lr=True, write_model_weights=False
+ ),
+ set_epoch_loop=False,
+ )
+ )
+ evaluation: base_configs.EvalConfig = dataclasses.field(
+ default_factory=lambda: base_configs.EvalConfig( # pylint: disable=g-long-lambda
+ epochs_between_evals=1, steps=None
+ )
+ )
+ model: base_configs.ModelConfig = dataclasses.field(
+ default_factory=vgg_config.VGGModelConfig
+ )
def get_config(model: str, dataset: str) -> base_configs.ExperimentConfig:
diff --git a/official/legacy/image_classification/dataset_factory.py b/official/legacy/image_classification/dataset_factory.py
index 19a757046b3..6b00e96f693 100644
--- a/official/legacy/image_classification/dataset_factory.py
+++ b/official/legacy/image_classification/dataset_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from typing import Any, List, Mapping, Optional, Tuple, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.legacy.image_classification import augment
from official.legacy.image_classification import preprocessing
@@ -48,7 +48,7 @@ class AugmentConfig(base_config.Config):
def build(self) -> augment.ImageAugment:
"""Build the augmenter using this config."""
params = self.params or {}
- augmenter = AUGMENTERS.get(self.name, None)
+ augmenter = AUGMENTERS.get(self.name, None) # pyrefly: ignore[no-matching-overload]
return augmenter(**params) if augmenter is not None else None
@@ -116,7 +116,7 @@ class DatasetConfig(base_config.Config):
num_devices: int = 1
dtype: str = 'float32'
one_hot: bool = True
- augmenter: AugmentConfig = AugmentConfig()
+ augmenter: AugmentConfig = dataclasses.field(default_factory=AugmentConfig)
download: bool = False
shuffle_buffer_size: int = 10000
file_shuffle_buffer_size: int = 1024
diff --git a/official/legacy/image_classification/efficientnet/__init__.py b/official/legacy/image_classification/efficientnet/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/image_classification/efficientnet/__init__.py
+++ b/official/legacy/image_classification/efficientnet/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/image_classification/efficientnet/common_modules.py b/official/legacy/image_classification/efficientnet/common_modules.py
index 28be6962047..6f8dd187b4e 100644
--- a/official/legacy/image_classification/efficientnet/common_modules.py
+++ b/official/legacy/image_classification/efficientnet/common_modules.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,13 +18,13 @@
from __future__ import print_function
from typing import Optional, Text
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow.compat.v1 as tf1
from tensorflow.python.tpu import tpu_function
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class TpuBatchNormalization(tf.keras.layers.BatchNormalization):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class TpuBatchNormalization(tf_keras.layers.BatchNormalization):
"""Cross replica batch normalization."""
def __init__(self, fused: Optional[bool] = False, **kwargs):
@@ -48,7 +48,7 @@ def _cross_replica_average(self, t: tf.Tensor, num_shards_per_group: int):
return tf1.tpu.cross_replica_sum(t, group_assignment) / tf.cast(
num_shards_per_group, t.dtype)
- def _moments(self, inputs: tf.Tensor, reduction_axes: int, keep_dims: int):
+ def _moments(self, inputs: tf.Tensor, reduction_axes: int, keep_dims: int): # pyrefly: ignore[bad-override]
"""Compute the mean and variance: it overrides the original _moments."""
shard_mean, shard_variance = super(TpuBatchNormalization, self)._moments(
inputs, reduction_axes, keep_dims=keep_dims)
@@ -71,7 +71,7 @@ def _moments(self, inputs: tf.Tensor, reduction_axes: int, keep_dims: int):
return (shard_mean, shard_variance)
-def get_batch_norm(batch_norm_type: Text) -> tf.keras.layers.BatchNormalization:
+def get_batch_norm(batch_norm_type: Text) -> tf_keras.layers.BatchNormalization:
"""A helper to create a batch normalization getter.
Args:
@@ -79,12 +79,12 @@ def get_batch_norm(batch_norm_type: Text) -> tf.keras.layers.BatchNormalization:
will use `TpuBatchNormalization`.
Returns:
- An instance of `tf.keras.layers.BatchNormalization`.
+ An instance of `tf_keras.layers.BatchNormalization`.
"""
if batch_norm_type == 'tpu':
return TpuBatchNormalization
- return tf.keras.layers.BatchNormalization # pytype: disable=bad-return-type # typed-keras
+ return tf_keras.layers.BatchNormalization # pytype: disable=bad-return-type # typed-keras
def count_params(model, trainable_only=True):
@@ -94,11 +94,11 @@ def count_params(model, trainable_only=True):
else:
return int(
np.sum([
- tf.keras.backend.count_params(p) for p in model.trainable_weights
+ tf_keras.backend.count_params(p) for p in model.trainable_weights
]))
-def load_weights(model: tf.keras.Model,
+def load_weights(model: tf_keras.Model,
model_weights_path: Text,
weights_format: Text = 'saved_model'):
"""Load model weights from the given file path.
@@ -110,7 +110,7 @@ def load_weights(model: tf.keras.Model,
'checkpoint'.
"""
if weights_format == 'saved_model':
- loaded_model = tf.keras.models.load_model(model_weights_path)
+ loaded_model = tf_keras.models.load_model(model_weights_path)
model.set_weights(loaded_model.get_weights())
else:
model.load_weights(model_weights_path)
diff --git a/official/legacy/image_classification/efficientnet/efficientnet_config.py b/official/legacy/image_classification/efficientnet/efficientnet_config.py
index 148851cf687..cddbac87cde 100644
--- a/official/legacy/image_classification/efficientnet/efficientnet_config.py
+++ b/official/legacy/image_classification/efficientnet/efficientnet_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -55,19 +55,28 @@ class EfficientNetModelConfig(base_configs.ModelConfig):
'dtype': 'float32',
}
})
- loss: base_configs.LossConfig = base_configs.LossConfig(
- name='categorical_crossentropy', label_smoothing=0.1)
- optimizer: base_configs.OptimizerConfig = base_configs.OptimizerConfig(
- name='rmsprop',
- decay=0.9,
- epsilon=0.001,
- momentum=0.9,
- moving_average_decay=None)
- learning_rate: base_configs.LearningRateConfig = base_configs.LearningRateConfig( # pylint: disable=line-too-long
- name='exponential',
- initial_lr=0.008,
- decay_epochs=2.4,
- decay_rate=0.97,
- warmup_epochs=5,
- scale_by_batch_size=1. / 128.,
- staircase=True)
+ loss: base_configs.LossConfig = dataclasses.field(
+ default_factory=lambda: base_configs.LossConfig( # pylint: disable=g-long-lambda
+ name='categorical_crossentropy', label_smoothing=0.1
+ )
+ )
+ optimizer: base_configs.OptimizerConfig = dataclasses.field(
+ default_factory=lambda: base_configs.OptimizerConfig( # pylint: disable=g-long-lambda
+ name='rmsprop',
+ decay=0.9,
+ epsilon=0.001,
+ momentum=0.9,
+ moving_average_decay=None,
+ )
+ )
+ learning_rate: base_configs.LearningRateConfig = dataclasses.field(
+ default_factory=lambda: base_configs.LearningRateConfig( # pylint: disable=g-long-lambda
+ name='exponential',
+ initial_lr=0.008,
+ decay_epochs=2.4,
+ decay_rate=0.97,
+ warmup_epochs=5,
+ scale_by_batch_size=1.0 / 128.0,
+ staircase=True,
+ )
+ )
diff --git a/official/legacy/image_classification/efficientnet/efficientnet_model.py b/official/legacy/image_classification/efficientnet/efficientnet_model.py
index a9aa243b037..e4ec052b028 100644
--- a/official/legacy/image_classification/efficientnet/efficientnet_model.py
+++ b/official/legacy/image_classification/efficientnet/efficientnet_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -27,7 +27,7 @@
from typing import Any, Dict, Optional, Text, Tuple
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification import preprocessing
from official.legacy.image_classification.efficientnet import common_modules
from official.modeling import tf_utils
@@ -134,7 +134,7 @@ def round_filters(filters: int, config: ModelConfig) -> int:
if not width_coefficient:
return filters
- filters *= width_coefficient
+ filters *= width_coefficient # pyrefly: ignore[bad-assignment]
min_depth = min_depth or divisor
new_filters = max(min_depth, int(filters + divisor / 2) // divisor * divisor)
# Make sure that round down does not go down by more than 10%.
@@ -163,7 +163,7 @@ def conv2d_block(inputs: tf.Tensor,
batch_norm = common_modules.get_batch_norm(config.batch_norm)
bn_momentum = config.bn_momentum
bn_epsilon = config.bn_epsilon
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
weight_decay = config.weight_decay
name = name or ''
@@ -175,21 +175,21 @@ def conv2d_block(inputs: tf.Tensor,
'use_bias': use_bias,
'padding': 'same',
'name': name + '_conv2d',
- 'kernel_regularizer': tf.keras.regularizers.l2(weight_decay),
- 'bias_regularizer': tf.keras.regularizers.l2(weight_decay),
+ 'kernel_regularizer': tf_keras.regularizers.l2(weight_decay),
+ 'bias_regularizer': tf_keras.regularizers.l2(weight_decay),
}
if depthwise:
- conv2d = tf.keras.layers.DepthwiseConv2D
+ conv2d = tf_keras.layers.DepthwiseConv2D
init_kwargs.update({'depthwise_initializer': CONV_KERNEL_INITIALIZER})
else:
- conv2d = tf.keras.layers.Conv2D
+ conv2d = tf_keras.layers.Conv2D
init_kwargs.update({
'filters': conv_filters,
'kernel_initializer': CONV_KERNEL_INITIALIZER
})
- x = conv2d(**init_kwargs)(inputs)
+ x = conv2d(**init_kwargs)(inputs) # pyrefly: ignore[missing-argument]
if use_batch_norm:
bn_axis = 1 if data_format == 'channels_first' else -1
@@ -201,7 +201,7 @@ def conv2d_block(inputs: tf.Tensor,
x)
if activation is not None:
- x = tf.keras.layers.Activation(activation, name=name + '_activation')(x)
+ x = tf_keras.layers.Activation(activation, name=name + '_activation')(x)
return x
@@ -223,7 +223,7 @@ def mb_conv_block(inputs: tf.Tensor,
use_se = config.use_se
activation = tf_utils.get_activation(config.activation)
drop_connect_rate = config.drop_connect_rate
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
use_depthwise = block.conv_type != 'no_depthwise'
prefix = prefix or ''
@@ -276,8 +276,8 @@ def mb_conv_block(inputs: tf.Tensor,
else:
se_shape = (1, 1, filters)
- se = tf.keras.layers.GlobalAveragePooling2D(name=prefix + 'se_squeeze')(x)
- se = tf.keras.layers.Reshape(se_shape, name=prefix + 'se_reshape')(se)
+ se = tf_keras.layers.GlobalAveragePooling2D(name=prefix + 'se_squeeze')(x)
+ se = tf_keras.layers.Reshape(se_shape, name=prefix + 'se_reshape')(se)
se = conv2d_block(
se,
@@ -295,7 +295,7 @@ def mb_conv_block(inputs: tf.Tensor,
use_batch_norm=False,
activation='sigmoid',
name=prefix + 'se_expand')
- x = tf.keras.layers.multiply([x, se], name=prefix + 'se_excite')
+ x = tf_keras.layers.multiply([x, se], name=prefix + 'se_excite')
# Output phase
x = conv2d_block(
@@ -303,7 +303,7 @@ def mb_conv_block(inputs: tf.Tensor,
# Add identity so that quantization-aware training can insert quantization
# ops correctly.
- x = tf.keras.layers.Activation(
+ x = tf_keras.layers.Activation(
tf_utils.get_activation('identity'), name=prefix + 'id')(
x)
@@ -314,19 +314,19 @@ def mb_conv_block(inputs: tf.Tensor,
# The only difference between dropout and dropconnect in TF is scaling by
# drop_connect_rate during training. See:
# https://github.com/keras-team/keras/pull/9898#issuecomment-380577612
- x = tf.keras.layers.Dropout(
+ x = tf_keras.layers.Dropout(
drop_connect_rate, noise_shape=(None, 1, 1, 1), name=prefix + 'drop')(
x)
- x = tf.keras.layers.add([x, inputs], name=prefix + 'add')
+ x = tf_keras.layers.add([x, inputs], name=prefix + 'add')
return x
-def efficientnet(image_input: tf.keras.layers.Input, config: ModelConfig): # pytype: disable=invalid-annotation # typed-keras
+def efficientnet(image_input: tf_keras.layers.Input, config: ModelConfig): # pytype: disable=invalid-annotation # typed-keras
"""Creates an EfficientNet graph given the model parameters.
- This function is wrapped by the `EfficientNet` class to make a tf.keras.Model.
+ This function is wrapped by the `EfficientNet` class to make a tf_keras.Model.
Args:
image_input: the input batch of images
@@ -345,14 +345,14 @@ def efficientnet(image_input: tf.keras.layers.Input, config: ModelConfig): # py
num_classes = config.num_classes
input_channels = config.input_channels
rescale_input = config.rescale_input
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
dtype = config.dtype
weight_decay = config.weight_decay
x = image_input
if data_format == 'channels_first':
# Happens on GPU/TPU if available.
- x = tf.keras.layers.Permute((3, 1, 2))(x)
+ x = tf_keras.layers.Permute((3, 1, 2))(x)
if rescale_input:
x = preprocessing.normalize_images(
x, num_channels=input_channels, dtype=dtype, data_format=data_format)
@@ -405,22 +405,22 @@ def efficientnet(image_input: tf.keras.layers.Input, config: ModelConfig): # py
name='top')
# Build classifier
- x = tf.keras.layers.GlobalAveragePooling2D(name='top_pool')(x)
+ x = tf_keras.layers.GlobalAveragePooling2D(name='top_pool')(x)
if dropout_rate and dropout_rate > 0:
- x = tf.keras.layers.Dropout(dropout_rate, name='top_dropout')(x)
- x = tf.keras.layers.Dense(
+ x = tf_keras.layers.Dropout(dropout_rate, name='top_dropout')(x)
+ x = tf_keras.layers.Dense(
num_classes,
kernel_initializer=DENSE_KERNEL_INITIALIZER,
- kernel_regularizer=tf.keras.regularizers.l2(weight_decay),
- bias_regularizer=tf.keras.regularizers.l2(weight_decay),
+ kernel_regularizer=tf_keras.regularizers.l2(weight_decay),
+ bias_regularizer=tf_keras.regularizers.l2(weight_decay),
name='logits')(
x)
- x = tf.keras.layers.Activation('softmax', name='probs')(x)
+ x = tf_keras.layers.Activation('softmax', name='probs')(x)
return x
-class EfficientNet(tf.keras.Model):
+class EfficientNet(tf_keras.Model):
"""Wrapper class for an EfficientNet Keras model.
Contains helper methods to build, manage, and save metadata about the model.
@@ -443,7 +443,7 @@ def __init__(self,
input_channels = self.config.input_channels
model_name = self.config.model_name
input_shape = (None, None, input_channels) # Should handle any size image
- image_input = tf.keras.layers.Input(shape=input_shape)
+ image_input = tf_keras.layers.Input(shape=input_shape)
output = efficientnet(image_input, self.config)
diff --git a/official/legacy/image_classification/efficientnet/tfhub_export.py b/official/legacy/image_classification/efficientnet/tfhub_export.py
index b82e587fdd8..a851a8c1526 100644
--- a/official/legacy/image_classification/efficientnet/tfhub_export.py
+++ b/official/legacy/image_classification/efficientnet/tfhub_export.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification.efficientnet import efficientnet_model
@@ -36,22 +36,22 @@
def export_tfhub(model_path, hub_destination, model_name):
- """Restores a tf.keras.Model and saves for TF-Hub."""
+ """Restores a tf_keras.Model and saves for TF-Hub."""
model_configs = dict(efficientnet_model.MODEL_CONFIGS)
config = model_configs[model_name]
- image_input = tf.keras.layers.Input(
+ image_input = tf_keras.layers.Input(
shape=(None, None, 3), name="image_input", dtype=tf.float32)
x = image_input * 255.0
- ouputs = efficientnet_model.efficientnet(x, config)
- hub_model = tf.keras.Model(image_input, ouputs)
+ outputs = efficientnet_model.efficientnet(x, config)
+ hub_model = tf_keras.Model(image_input, outputs)
ckpt = tf.train.Checkpoint(model=hub_model)
ckpt.restore(model_path).assert_existing_objects_matched()
hub_model.save(
os.path.join(hub_destination, "classification"), include_optimizer=False)
feature_vector_output = hub_model.get_layer(name="top_pool").get_output_at(0)
- hub_model2 = tf.keras.Model(image_input, feature_vector_output)
+ hub_model2 = tf_keras.Model(image_input, feature_vector_output)
hub_model2.save(
os.path.join(hub_destination, "feature-vector"), include_optimizer=False)
diff --git a/official/legacy/image_classification/learning_rate.py b/official/legacy/image_classification/learning_rate.py
index 248cc8472e1..90bcc46547e 100644
--- a/official/legacy/image_classification/learning_rate.py
+++ b/official/legacy/image_classification/learning_rate.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,16 +20,16 @@
from typing import Any, Mapping, Optional
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
BASE_LEARNING_RATE = 0.1
-class WarmupDecaySchedule(tf.keras.optimizers.schedules.LearningRateSchedule):
+class WarmupDecaySchedule(tf_keras.optimizers.schedules.LearningRateSchedule):
"""A wrapper for LearningRateSchedule that includes warmup steps."""
def __init__(self,
- lr_schedule: tf.keras.optimizers.schedules.LearningRateSchedule,
+ lr_schedule: tf_keras.optimizers.schedules.LearningRateSchedule,
warmup_steps: int,
warmup_lr: Optional[float] = None):
"""Add warmup decay to a learning rate schedule.
@@ -73,7 +73,7 @@ def get_config(self) -> Mapping[str, Any]:
return config
-class CosineDecayWithWarmup(tf.keras.optimizers.schedules.LearningRateSchedule):
+class CosineDecayWithWarmup(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Class to generate learning rate tensor."""
def __init__(self, batch_size: int, total_steps: int, warmup_steps: int):
diff --git a/official/legacy/image_classification/learning_rate_test.py b/official/legacy/image_classification/learning_rate_test.py
index 77dc65c571f..a337679b951 100644
--- a/official/legacy/image_classification/learning_rate_test.py
+++ b/official/legacy/image_classification/learning_rate_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification import learning_rate
@@ -32,7 +32,7 @@ def test_warmup_decay(self):
decay_rate = 0.01
warmup_steps = 10
- base_lr = tf.keras.optimizers.schedules.ExponentialDecay(
+ base_lr = tf_keras.optimizers.schedules.ExponentialDecay(
initial_learning_rate=initial_lr,
decay_steps=decay_steps,
decay_rate=decay_rate)
diff --git a/official/legacy/image_classification/mnist_main.py b/official/legacy/image_classification/mnist_main.py
index cf60631444e..1fb2b5a5b4a 100644
--- a/official/legacy/image_classification/mnist_main.py
+++ b/official/legacy/image_classification/mnist_main.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,11 +19,10 @@
import os
-# Import libraries
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.common import distribute_utils
from official.legacy.image_classification.resnet import common
@@ -36,29 +35,29 @@
def build_model():
"""Constructs the ML model used to predict handwritten digits."""
- image = tf.keras.layers.Input(shape=(28, 28, 1))
+ image = tf_keras.layers.Input(shape=(28, 28, 1))
- y = tf.keras.layers.Conv2D(filters=32,
+ y = tf_keras.layers.Conv2D(filters=32,
kernel_size=5,
padding='same',
activation='relu')(image)
- y = tf.keras.layers.MaxPooling2D(pool_size=(2, 2),
+ y = tf_keras.layers.MaxPooling2D(pool_size=(2, 2),
strides=(2, 2),
padding='same')(y)
- y = tf.keras.layers.Conv2D(filters=32,
+ y = tf_keras.layers.Conv2D(filters=32,
kernel_size=5,
padding='same',
activation='relu')(y)
- y = tf.keras.layers.MaxPooling2D(pool_size=(2, 2),
+ y = tf_keras.layers.MaxPooling2D(pool_size=(2, 2),
strides=(2, 2),
padding='same')(y)
- y = tf.keras.layers.Flatten()(y)
- y = tf.keras.layers.Dense(1024, activation='relu')(y)
- y = tf.keras.layers.Dropout(0.4)(y)
+ y = tf_keras.layers.Flatten()(y)
+ y = tf_keras.layers.Dense(1024, activation='relu')(y)
+ y = tf_keras.layers.Dropout(0.4)(y)
- probs = tf.keras.layers.Dense(10, activation='softmax')(y)
+ probs = tf_keras.layers.Dense(10, activation='softmax')(y)
- model = tf.keras.models.Model(image, probs, name='mnist')
+ model = tf_keras.models.Model(image, probs, name='mnist')
return model
@@ -104,9 +103,9 @@ def run(flags_obj, datasets_override=None, strategy_override=None):
eval_input_dataset = mnist_test.cache().repeat().batch(flags_obj.batch_size)
with strategy_scope:
- lr_schedule = tf.keras.optimizers.schedules.ExponentialDecay(
+ lr_schedule = tf_keras.optimizers.schedules.ExponentialDecay(
0.05, decay_steps=100000, decay_rate=0.96)
- optimizer = tf.keras.optimizers.SGD(learning_rate=lr_schedule)
+ optimizer = tf_keras.optimizers.SGD(learning_rate=lr_schedule)
model = build_model()
model.compile(
@@ -120,9 +119,9 @@ def run(flags_obj, datasets_override=None, strategy_override=None):
ckpt_full_path = os.path.join(flags_obj.model_dir, 'model.ckpt-{epoch:04d}')
callbacks = [
- tf.keras.callbacks.ModelCheckpoint(
+ tf_keras.callbacks.ModelCheckpoint(
ckpt_full_path, save_weights_only=True),
- tf.keras.callbacks.TensorBoard(log_dir=flags_obj.model_dir),
+ tf_keras.callbacks.TensorBoard(log_dir=flags_obj.model_dir),
]
num_eval_examples = mnist.info.splits['test'].num_examples
diff --git a/official/legacy/image_classification/mnist_test.py b/official/legacy/image_classification/mnist_test.py
index 384a6a9abb3..dabfd1192ca 100644
--- a/official/legacy/image_classification/mnist_test.py
+++ b/official/legacy/image_classification/mnist_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
import functools
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
diff --git a/official/legacy/image_classification/optimizer_factory.py b/official/legacy/image_classification/optimizer_factory.py
index 42be845837f..557a5283820 100644
--- a/official/legacy/image_classification/optimizer_factory.py
+++ b/official/legacy/image_classification/optimizer_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,38 +17,192 @@
from __future__ import division
from __future__ import print_function
-from typing import Any, Dict, Optional, Text
+from typing import Any, Dict, Optional, Text, Union
from absl import logging
-import tensorflow as tf
-import tensorflow_addons as tfa
+import numpy as np
+import tensorflow as tf, tf_keras
+
from official.legacy.image_classification import learning_rate
from official.legacy.image_classification.configs import base_configs
from official.modeling import optimization
+from official.modeling.optimization import legacy_adamw
# pylint: disable=protected-access
+FloatTensorLike = Union[tf.Tensor, float, np.float16, np.float32, np.float64]
+
+
+class Lookahead(tf_keras.optimizers.legacy.Optimizer):
+ """This class allows to extend optimizers with the lookahead mechanism.
+
+ The mechanism is proposed by Michael R. Zhang et.al in the paper [Lookahead
+ Optimizer: k steps forward, 1 step back] (https://arxiv.org/abs/1907.08610v1).
+ The optimizer iteratively updates two sets of weights: the search directions
+ for weights are chosen by the inner optimizer, while the "slow weights" are
+ updated each `k` steps based on the directions of the "fast weights" and the
+ two sets of weights are synchronized. This method improves the learning
+ stability and lowers the variance of its inner optimizer.
+
+ Example of usage:
+
+ ```python
+ opt = tf_keras.optimizers.SGD(learning_rate) opt =
+ tfa.optimizers.Lookahead(opt)
+ ```
+ """
+
+ def __init__(
+ self,
+ optimizer: tf_keras.optimizers.Optimizer,
+ sync_period: int = 6,
+ slow_step_size: FloatTensorLike = 0.5,
+ name: str = 'Lookahead',
+ **kwargs,
+ ):
+ """Wrap optimizer with the lookahead mechanism.
+
+ Args:
+ optimizer: The original optimizer that will be used to compute and apply
+ the gradients.
+ sync_period: An integer. The synchronization period of lookahead. Enable
+ lookahead mechanism by setting it with a positive value.
+ slow_step_size: A floating point value. The ratio for updating the slow
+ weights.
+ name: Optional name for the operations created when applying gradients.
+ Defaults to "Lookahead".
+ **kwargs: keyword arguments. Allowed to be {`clipnorm`, `clipvalue`, `lr`,
+ `decay`}. `clipnorm` is clip gradients by norm; `clipvalue` is clip
+ gradients by value, `decay` is included for backward compatibility to
+ allow time inverse decay of learning rate. `lr` is included for backward
+ compatibility, recommended to use `learning_rate` instead.
+ """
+ super().__init__(name, **kwargs)
+
+ if isinstance(optimizer, str):
+ optimizer = tf_keras.optimizers.get(optimizer)
+ if not isinstance(
+ optimizer,
+ (tf_keras.optimizers.Optimizer, tf_keras.optimizers.legacy.Optimizer),
+ ):
+ raise TypeError(
+ 'optimizer is not an object of tf_keras.optimizers.Optimizer'
+ )
+
+ self._optimizer = optimizer
+ self._set_hyper('sync_period', sync_period)
+ self._set_hyper('slow_step_size', slow_step_size)
+ self._initialized = False
+ self._track_trackable(self._optimizer, 'lh_base_optimizer')
+
+ def _create_slots(self, var_list):
+ self._optimizer._create_slots(var_list=var_list) # pylint: disable=protected-access
+ for var in var_list:
+ self.add_slot(var, 'slow', initializer=var)
+
+ def _create_hypers(self):
+ self._optimizer._create_hypers() # pylint: disable=protected-access
+
+ def _prepare(self, var_list):
+ return self._optimizer._prepare(var_list=var_list) # pylint: disable=protected-access
+
+ def apply_gradients(
+ self, grads_and_vars, name=None, skip_gradients_aggregation=None, **kwargs
+ ):
+ self._optimizer._iterations = self.iterations # pylint: disable=protected-access
+ return super().apply_gradients(grads_and_vars, name, **kwargs)
+
+ def _look_ahead_op(self, var):
+ var_dtype = var.dtype.base_dtype
+ slow_var = self.get_slot(var, 'slow')
+ local_step = tf.cast(self.iterations + 1, tf.dtypes.int64)
+ sync_period = self._get_hyper('sync_period', tf.dtypes.int64)
+ slow_step_size = self._get_hyper('slow_step_size', var_dtype)
+ step_back = slow_var + slow_step_size * (var - slow_var)
+ sync_cond = tf.equal(
+ tf.math.floordiv(local_step, sync_period) * sync_period, local_step
+ )
+ with tf.control_dependencies([step_back]):
+ slow_update = slow_var.assign(
+ tf.where(sync_cond, step_back, slow_var),
+ use_locking=self._use_locking,
+ )
+ var_update = var.assign(
+ tf.where(sync_cond, step_back, var), use_locking=self._use_locking
+ )
+ return tf.group(slow_update, var_update)
+
+ @property
+ def weights(self):
+ return self._weights + self._optimizer.weights
+
+ def _resource_apply_dense(self, grad, var):
+ train_op = self._optimizer._resource_apply_dense(grad, var) # pylint: disable=protected-access
+ with tf.control_dependencies([train_op]):
+ look_ahead_op = self._look_ahead_op(var)
+ return tf.group(train_op, look_ahead_op)
+
+ def _resource_apply_sparse(self, grad, var, indices):
+ train_op = self._optimizer._resource_apply_sparse( # pylint: disable=protected-access
+ grad, var, indices
+ )
+ with tf.control_dependencies([train_op]):
+ look_ahead_op = self._look_ahead_op(var)
+ return tf.group(train_op, look_ahead_op)
+
+ def get_config(self):
+ config = {
+ 'optimizer': tf_keras.optimizers.serialize(self._optimizer),
+ 'sync_period': self._serialize_hyperparameter('sync_period'),
+ 'slow_step_size': self._serialize_hyperparameter('slow_step_size'),
+ }
+ base_config = super().get_config()
+ return {**base_config, **config}
+
+ @property
+ def learning_rate(self):
+ return self._optimizer._get_hyper('learning_rate')
+
+ @learning_rate.setter
+ def learning_rate(self, value):
+ self._optimizer._set_hyper('learning_rate', value)
+
+ @property
+ def lr(self):
+ return self.learning_rate
+
+ @lr.setter
+ def lr(self, lr):
+ self.learning_rate = lr
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ optimizer = tf_keras.optimizers.deserialize(
+ config.pop('optimizer'), custom_objects=custom_objects
+ )
+ return cls(optimizer, **config)
+
def build_optimizer(
optimizer_name: Text,
- base_learning_rate: tf.keras.optimizers.schedules.LearningRateSchedule,
+ base_learning_rate: tf_keras.optimizers.schedules.LearningRateSchedule,
params: Dict[Text, Any],
- model: Optional[tf.keras.Model] = None):
+ model: Optional[tf_keras.Model] = None):
"""Build the optimizer based on name.
Args:
optimizer_name: String representation of the optimizer name. Examples: sgd,
momentum, rmsprop.
- base_learning_rate: `tf.keras.optimizers.schedules.LearningRateSchedule`
+ base_learning_rate: `tf_keras.optimizers.schedules.LearningRateSchedule`
base learning rate.
params: String -> Any dictionary representing the optimizer params. This
should contain optimizer specific parameters such as `base_learning_rate`,
`decay`, etc.
- model: The `tf.keras.Model`. This is used for the shadow copy if using
+ model: The `tf_keras.Model`. This is used for the shadow copy if using
`ExponentialMovingAverage`.
Returns:
- A tf.keras.Optimizer.
+ A tf_keras.optimizers.legacy.Optimizer.
Raises:
ValueError if the provided optimizer_name is not supported.
@@ -60,12 +214,12 @@ def build_optimizer(
if optimizer_name == 'sgd':
logging.info('Using SGD optimizer')
nesterov = params.get('nesterov', False)
- optimizer = tf.keras.optimizers.SGD(
+ optimizer = tf_keras.optimizers.legacy.SGD(
learning_rate=base_learning_rate, nesterov=nesterov)
elif optimizer_name == 'momentum':
logging.info('Using momentum optimizer')
nesterov = params.get('nesterov', False)
- optimizer = tf.keras.optimizers.SGD(
+ optimizer = tf_keras.optimizers.legacy.SGD(
learning_rate=base_learning_rate,
momentum=params['momentum'],
nesterov=nesterov)
@@ -74,7 +228,7 @@ def build_optimizer(
rho = params.get('decay', None) or params.get('rho', 0.9)
momentum = params.get('momentum', 0.9)
epsilon = params.get('epsilon', 1e-07)
- optimizer = tf.keras.optimizers.RMSprop(
+ optimizer = tf_keras.optimizers.legacy.RMSprop(
learning_rate=base_learning_rate,
rho=rho,
momentum=momentum,
@@ -84,7 +238,7 @@ def build_optimizer(
beta_1 = params.get('beta_1', 0.9)
beta_2 = params.get('beta_2', 0.999)
epsilon = params.get('epsilon', 1e-07)
- optimizer = tf.keras.optimizers.Adam(
+ optimizer = tf_keras.optimizers.legacy.Adam(
learning_rate=base_learning_rate,
beta_1=beta_1,
beta_2=beta_2,
@@ -95,18 +249,19 @@ def build_optimizer(
beta_1 = params.get('beta_1', 0.9)
beta_2 = params.get('beta_2', 0.999)
epsilon = params.get('epsilon', 1e-07)
- optimizer = tfa.optimizers.AdamW(
- weight_decay=weight_decay,
+ optimizer = legacy_adamw.AdamWeightDecay(
learning_rate=base_learning_rate,
+ weight_decay_rate=weight_decay,
beta_1=beta_1,
beta_2=beta_2,
- epsilon=epsilon)
+ epsilon=epsilon,
+ )
else:
raise ValueError('Unknown optimizer %s' % optimizer_name)
if params.get('lookahead', None):
logging.info('Using lookahead optimizer.')
- optimizer = tfa.optimizers.Lookahead(optimizer)
+ optimizer = Lookahead(optimizer)
# Moving average should be applied last, as it's applied at test time
moving_average_decay = params.get('moving_average_decay', 0.)
@@ -152,7 +307,7 @@ def build_learning_rate(params: base_configs.LearningRateConfig,
'Using exponential learning rate with: '
'initial_learning_rate: %f, decay_steps: %d, '
'decay_rate: %f', base_lr, decay_steps, decay_rate)
- lr = tf.keras.optimizers.schedules.ExponentialDecay(
+ lr = tf_keras.optimizers.schedules.ExponentialDecay(
initial_learning_rate=base_lr,
decay_steps=decay_steps,
decay_rate=decay_rate,
@@ -164,17 +319,17 @@ def build_learning_rate(params: base_configs.LearningRateConfig,
logging.info(
'Using stepwise learning rate. Parameters: '
'boundaries: %s, values: %s', boundaries, multipliers)
- lr = tf.keras.optimizers.schedules.PiecewiseConstantDecay(
+ lr = tf_keras.optimizers.schedules.PiecewiseConstantDecay(
boundaries=boundaries, values=multipliers)
elif decay_type == 'cosine_with_warmup':
lr = learning_rate.CosineDecayWithWarmup(
batch_size=batch_size,
- total_steps=train_epochs * train_steps,
+ total_steps=train_epochs * train_steps, # pyrefly: ignore[unsupported-operation]
warmup_steps=warmup_steps)
if warmup_steps > 0:
if decay_type not in ['cosine_with_warmup']:
logging.info('Applying %d warmup steps to the learning rate',
warmup_steps)
lr = learning_rate.WarmupDecaySchedule(
- lr, warmup_steps, warmup_lr=base_lr)
- return lr
+ lr, warmup_steps, warmup_lr=base_lr) # pyrefly: ignore[unbound-name]
+ return lr # pyrefly: ignore[unbound-name]
diff --git a/official/legacy/image_classification/optimizer_factory_test.py b/official/legacy/image_classification/optimizer_factory_test.py
index e0974505790..3b51b5e3214 100644
--- a/official/legacy/image_classification/optimizer_factory_test.py
+++ b/official/legacy/image_classification/optimizer_factory_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,17 +20,17 @@
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification import optimizer_factory
from official.legacy.image_classification.configs import base_configs
class OptimizerFactoryTest(tf.test.TestCase, parameterized.TestCase):
- def build_toy_model(self) -> tf.keras.Model:
+ def build_toy_model(self) -> tf_keras.Model:
"""Creates a toy `tf.Keras.Model`."""
- model = tf.keras.Sequential()
- model.add(tf.keras.layers.Dense(1, input_shape=(1,)))
+ model = tf_keras.Sequential()
+ model.add(tf_keras.layers.Dense(1, input_shape=(1,)))
return model
@parameterized.named_parameters(
@@ -57,7 +57,9 @@ def test_optimizer(self, optimizer_name, moving_average_decay, lookahead):
base_learning_rate=params['learning_rate'],
params=params,
model=model)
- self.assertTrue(issubclass(type(optimizer), tf.keras.optimizers.Optimizer))
+ self.assertTrue(
+ issubclass(type(optimizer), tf_keras.optimizers.legacy.Optimizer)
+ )
def test_unknown_optimizer(self):
with self.assertRaises(ValueError):
@@ -84,7 +86,7 @@ def test_learning_rate_without_decay_or_warmups(self):
params=params, batch_size=batch_size, train_steps=train_steps)
self.assertTrue(
issubclass(
- type(lr), tf.keras.optimizers.schedules.LearningRateSchedule))
+ type(lr), tf_keras.optimizers.schedules.LearningRateSchedule))
@parameterized.named_parameters(('exponential', 'exponential'),
('cosine_with_warmup', 'cosine_with_warmup'))
@@ -111,7 +113,7 @@ def test_learning_rate_with_decay_and_warmup(self, lr_decay_type):
train_steps=train_steps)
self.assertTrue(
issubclass(
- type(lr), tf.keras.optimizers.schedules.LearningRateSchedule))
+ type(lr), tf_keras.optimizers.schedules.LearningRateSchedule))
if __name__ == '__main__':
diff --git a/official/legacy/image_classification/preprocessing.py b/official/legacy/image_classification/preprocessing.py
index 78b58243afb..cc0e260a307 100644
--- a/official/legacy/image_classification/preprocessing.py
+++ b/official/legacy/image_classification/preprocessing.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from __future__ import division
from __future__ import print_function
from typing import List, Optional, Text, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification import augment
diff --git a/official/legacy/image_classification/resnet/__init__.py b/official/legacy/image_classification/resnet/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/image_classification/resnet/__init__.py
+++ b/official/legacy/image_classification/resnet/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/image_classification/resnet/common.py b/official/legacy/image_classification/resnet/common.py
index 59786136e29..53bd325fc08 100644
--- a/official/legacy/image_classification/resnet/common.py
+++ b/official/legacy/image_classification/resnet/common.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import os
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.utils.flags import core as flags_core
@@ -35,7 +35,7 @@
class PiecewiseConstantDecayWithWarmup(
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
"""Piecewise constant decay with warmup schedule."""
def __init__(self,
@@ -81,17 +81,18 @@ def __call__(self, step):
def _get_learning_rate(self, step):
"""Compute learning rate at given step."""
+ step = tf.cast(step, dtype=tf.float32)
+ warmup_steps = tf.cast(self.warmup_steps, dtype=tf.float32)
with tf.name_scope('PiecewiseConstantDecayWithWarmup'):
def warmup_lr(step):
- return self.rescaled_lr * (
- tf.cast(step, tf.float32) / tf.cast(self.warmup_steps, tf.float32))
+ return self.rescaled_lr * (step / warmup_steps)
def piecewise_lr(step):
return tf.compat.v1.train.piecewise_constant(step, self.step_boundaries,
self.lr_values)
- return tf.cond(step < self.warmup_steps, lambda: warmup_lr(step),
+ return tf.cond(step < warmup_steps, lambda: warmup_lr(step),
lambda: piecewise_lr(step))
def get_config(self):
@@ -105,10 +106,14 @@ def get_config(self):
}
-def get_optimizer(learning_rate=0.1):
+def get_optimizer(learning_rate=0.1, use_legacy_optimizer=True):
"""Returns optimizer to use."""
# The learning_rate is overwritten at the beginning of each step by callback.
- return tf.keras.optimizers.SGD(learning_rate=learning_rate, momentum=0.9)
+ if use_legacy_optimizer:
+ return tf_keras.optimizers.legacy.SGD(
+ learning_rate=learning_rate, momentum=0.9)
+ else:
+ return tf_keras.optimizers.SGD(learning_rate=learning_rate, momentum=0.9)
def get_callbacks(pruning_method=None,
@@ -122,7 +127,7 @@ def get_callbacks(pruning_method=None,
callbacks = [time_callback]
if FLAGS.enable_tensorboard:
- tensorboard_callback = tf.keras.callbacks.TensorBoard(
+ tensorboard_callback = tf_keras.callbacks.TensorBoard(
log_dir=FLAGS.model_dir, profile_batch=FLAGS.profile_steps)
callbacks.append(tensorboard_callback)
@@ -138,7 +143,7 @@ def get_callbacks(pruning_method=None,
if model_dir is not None:
ckpt_full_path = os.path.join(model_dir, 'model.ckpt-{epoch:04d}')
callbacks.append(
- tf.keras.callbacks.ModelCheckpoint(
+ tf_keras.callbacks.ModelCheckpoint(
ckpt_full_path, save_weights_only=True))
return callbacks
diff --git a/official/legacy/image_classification/resnet/imagenet_preprocessing.py b/official/legacy/image_classification/resnet/imagenet_preprocessing.py
index d60107035da..68cc0af01db 100644
--- a/official/legacy/image_classification/resnet/imagenet_preprocessing.py
+++ b/official/legacy/image_classification/resnet/imagenet_preprocessing.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -38,7 +38,7 @@
import os
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
DEFAULT_IMAGE_SIZE = 224
NUM_CHANNELS = 3
@@ -257,13 +257,13 @@ def get_parse_record_fn(use_keras_image_data_format=False):
This is useful by handling different types of Keras models. For instance,
the current resnet_model.resnet50 input format is always channel-last,
whereas the keras_applications mobilenet input format depends on
- tf.keras.backend.image_data_format(). We should set
+ tf_keras.backend.image_data_format(). We should set
use_keras_image_data_format=False for the former and True for the latter.
Args:
use_keras_image_data_format: A boolean denoting whether data format is keras
backend image data format. If False, the image format is channel-last. If
- True, the image format matches tf.keras.backend.image_data_format().
+ True, the image format matches tf_keras.backend.image_data_format().
Returns:
Function to use for parsing the records.
@@ -272,7 +272,7 @@ def get_parse_record_fn(use_keras_image_data_format=False):
def parse_record_fn(raw_record, is_training, dtype):
image, label = parse_record(raw_record, is_training, dtype)
if use_keras_image_data_format:
- if tf.keras.backend.image_data_format() == 'channels_first':
+ if tf_keras.backend.image_data_format() == 'channels_first':
image = tf.transpose(image, perm=[2, 0, 1])
return image, label
diff --git a/official/legacy/image_classification/resnet/resnet_config.py b/official/legacy/image_classification/resnet/resnet_config.py
index 9c406282162..08a073fa2c2 100644
--- a/official/legacy/image_classification/resnet/resnet_config.py
+++ b/official/legacy/image_classification/resnet/resnet_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -36,20 +36,28 @@ class ResNetModelConfig(base_configs.ModelConfig):
'rescale_inputs': False,
})
# pylint: enable=g-long-lambda
- loss: base_configs.LossConfig = base_configs.LossConfig(
- name='sparse_categorical_crossentropy')
- optimizer: base_configs.OptimizerConfig = base_configs.OptimizerConfig(
- name='momentum',
- decay=0.9,
- epsilon=0.001,
- momentum=0.9,
- moving_average_decay=None)
- learning_rate: base_configs.LearningRateConfig = (
- base_configs.LearningRateConfig(
+ loss: base_configs.LossConfig = dataclasses.field(
+ default_factory=lambda: base_configs.LossConfig( # pylint: disable=g-long-lambda
+ name='sparse_categorical_crossentropy'
+ )
+ )
+ optimizer: base_configs.OptimizerConfig = dataclasses.field(
+ default_factory=lambda: base_configs.OptimizerConfig( # pylint: disable=g-long-lambda
+ name='momentum',
+ decay=0.9,
+ epsilon=0.001,
+ momentum=0.9,
+ moving_average_decay=None,
+ )
+ )
+ learning_rate: base_configs.LearningRateConfig = dataclasses.field(
+ default_factory=lambda: base_configs.LearningRateConfig( # pylint: disable=g-long-lambda
name='stepwise',
initial_lr=0.1,
examples_per_epoch=1281167,
boundaries=[30, 60, 80],
warmup_epochs=5,
- scale_by_batch_size=1. / 256.,
- multipliers=[0.1 / 256, 0.01 / 256, 0.001 / 256, 0.0001 / 256]))
+ scale_by_batch_size=1.0 / 256.0,
+ multipliers=[0.1 / 256, 0.01 / 256, 0.001 / 256, 0.0001 / 256],
+ )
+ )
diff --git a/official/legacy/image_classification/resnet/resnet_ctl_imagenet_main.py b/official/legacy/image_classification/resnet/resnet_ctl_imagenet_main.py
index 328d2890a73..8de0c8f7b6f 100644
--- a/official/legacy/image_classification/resnet/resnet_ctl_imagenet_main.py
+++ b/official/legacy/image_classification/resnet/resnet_ctl_imagenet_main.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,12 +17,11 @@
import math
import os
-# Import libraries
from absl import app
from absl import flags
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.image_classification.resnet import common
from official.legacy.image_classification.resnet import imagenet_preprocessing
@@ -113,7 +112,7 @@ def run(flags_obj):
if data_format is None:
data_format = ('channels_first' if tf.config.list_physical_devices('GPU')
else 'channels_last')
- tf.keras.backend.set_image_data_format(data_format)
+ tf_keras.backend.set_image_data_format(data_format)
strategy = distribute_utils.get_distribution_strategy(
distribution_strategy=flags_obj.distribution_strategy,
diff --git a/official/legacy/image_classification/resnet/resnet_model.py b/official/legacy/image_classification/resnet/resnet_model.py
index 545d06ecc9a..e89f9fd1432 100644
--- a/official/legacy/image_classification/resnet/resnet_model.py
+++ b/official/legacy/image_classification/resnet/resnet_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""ResNet50 model for Keras.
-Adapted from tf.keras.applications.resnet50.ResNet50().
+Adapted from tf_keras.applications.resnet50.ResNet50().
This is ResNet model version 1.5.
Related papers/blogs:
@@ -27,14 +27,14 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification.resnet import imagenet_preprocessing
-layers = tf.keras.layers
+layers = tf_keras.layers
def _gen_l2_regularizer(use_l2_regularizer=True, l2_weight_decay=1e-4):
- return tf.keras.regularizers.L2(
+ return tf_keras.regularizers.L2(
l2_weight_decay) if use_l2_regularizer else None
@@ -62,7 +62,7 @@ def identity_block(input_tensor,
Output tensor for the block.
"""
filters1, filters2, filters3 = filters
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
bn_axis = 3
else:
bn_axis = 1
@@ -150,7 +150,7 @@ def conv_block(input_tensor,
Output tensor for the block.
"""
filters1, filters2, filters3 = filters
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
bn_axis = 3
else:
bn_axis = 1
@@ -249,7 +249,7 @@ def resnet50(num_classes,
# Hub image modules expect inputs in the range [0, 1]. This rescales these
# inputs to the range expected by the trained model.
x = layers.Lambda(
- lambda x: x * 255.0 - tf.keras.backend.constant( # pylint: disable=g-long-lambda
+ lambda x: x * 255.0 - tf_keras.backend.constant( # pylint: disable=g-long-lambda
imagenet_preprocessing.CHANNEL_MEANS,
shape=[1, 1, 3],
dtype=x.dtype),
@@ -258,7 +258,7 @@ def resnet50(num_classes,
else:
x = img_input
- if tf.keras.backend.image_data_format() == 'channels_first':
+ if tf_keras.backend.image_data_format() == 'channels_first':
x = layers.Permute((3, 1, 2))(x)
bn_axis = 1
else: # channels_last
@@ -322,4 +322,4 @@ def resnet50(num_classes,
x = layers.Activation('softmax', dtype='float32')(x)
# Create model.
- return tf.keras.Model(img_input, x, name='resnet50')
+ return tf_keras.Model(img_input, x, name='resnet50')
diff --git a/official/legacy/image_classification/resnet/resnet_runnable.py b/official/legacy/image_classification/resnet/resnet_runnable.py
index c8f9ade935f..7b6e8786d1b 100644
--- a/official/legacy/image_classification/resnet/resnet_runnable.py
+++ b/official/legacy/image_classification/resnet/resnet_runnable.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Runs a ResNet model on the ImageNet dataset using custom training loops."""
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification.resnet import common
from official.legacy.image_classification.resnet import imagenet_preprocessing
from official.legacy.image_classification.resnet import resnet_model
@@ -76,11 +76,11 @@ def __init__(self, flags_obj, time_callback, epoch_steps):
use_float16=self.dtype == tf.float16,
loss_scale=flags_core.get_loss_scale(flags_obj, default_for_fp16=128))
- self.train_loss = tf.keras.metrics.Mean('train_loss', dtype=tf.float32)
- self.train_accuracy = tf.keras.metrics.SparseCategoricalAccuracy(
+ self.train_loss = tf_keras.metrics.Mean('train_loss', dtype=tf.float32)
+ self.train_accuracy = tf_keras.metrics.SparseCategoricalAccuracy(
'train_accuracy', dtype=tf.float32)
- self.test_loss = tf.keras.metrics.Mean('test_loss', dtype=tf.float32)
- self.test_accuracy = tf.keras.metrics.SparseCategoricalAccuracy(
+ self.test_loss = tf_keras.metrics.Mean('test_loss', dtype=tf.float32)
+ self.test_accuracy = tf_keras.metrics.SparseCategoricalAccuracy(
'test_accuracy', dtype=tf.float32)
self.checkpoint = tf.train.Checkpoint(
@@ -140,7 +140,7 @@ def step_fn(inputs):
with tf.GradientTape() as tape:
logits = self.model(images, training=True)
- prediction_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ prediction_loss = tf_keras.losses.sparse_categorical_crossentropy(
labels, logits)
loss = tf.reduce_sum(prediction_loss) * (1.0 /
self.flags_obj.batch_size)
@@ -187,7 +187,7 @@ def step_fn(inputs):
"""Function to run on the device."""
images, labels = inputs
logits = self.model(images, training=False)
- loss = tf.keras.losses.sparse_categorical_crossentropy(labels, logits)
+ loss = tf_keras.losses.sparse_categorical_crossentropy(labels, logits)
loss = tf.reduce_sum(loss) * (1.0 / self.flags_obj.batch_size)
self.test_loss.update_state(loss)
self.test_accuracy.update_state(labels, logits)
diff --git a/official/legacy/image_classification/resnet/tfhub_export.py b/official/legacy/image_classification/resnet/tfhub_export.py
index 1d7d743ddeb..f1ebade900b 100644
--- a/official/legacy/image_classification/resnet/tfhub_export.py
+++ b/official/legacy/image_classification/resnet/tfhub_export.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,11 +20,10 @@
import os
-# Import libraries
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification.resnet import imagenet_preprocessing
from official.legacy.image_classification.resnet import resnet_model
@@ -38,7 +37,7 @@
def export_tfhub(model_path, hub_destination):
- """Restores a tf.keras.Model and saves for TF-Hub."""
+ """Restores a tf_keras.Model and saves for TF-Hub."""
model = resnet_model.resnet50(
num_classes=imagenet_preprocessing.NUM_CLASSES, rescale_inputs=True)
model.load_weights(model_path)
@@ -48,7 +47,7 @@ def export_tfhub(model_path, hub_destination):
# Extracts a sub-model to use pooling feature vector as model output.
image_input = model.get_layer(index=0).get_output_at(0)
feature_vector_output = model.get_layer(name="reduce_mean").get_output_at(0)
- hub_model = tf.keras.Model(image_input, feature_vector_output)
+ hub_model = tf_keras.Model(image_input, feature_vector_output)
# Exports a SavedModel.
hub_model.save(
diff --git a/official/legacy/image_classification/test_utils.py b/official/legacy/image_classification/test_utils.py
index 871ac7e30f0..e6b2b887840 100644
--- a/official/legacy/image_classification/test_utils.py
+++ b/official/legacy/image_classification/test_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,20 +18,20 @@
from __future__ import division
from __future__ import print_function
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def trivial_model(num_classes):
"""Trivial model for ImageNet dataset."""
input_shape = (224, 224, 3)
- img_input = tf.keras.layers.Input(shape=input_shape)
+ img_input = tf_keras.layers.Input(shape=input_shape)
- x = tf.keras.layers.Lambda(
- lambda x: tf.keras.backend.reshape(x, [-1, 224 * 224 * 3]),
+ x = tf_keras.layers.Lambda(
+ lambda x: tf_keras.backend.reshape(x, [-1, 224 * 224 * 3]),
name='reshape')(img_input)
- x = tf.keras.layers.Dense(1, name='fc1')(x)
- x = tf.keras.layers.Dense(num_classes, name='fc1000')(x)
- x = tf.keras.layers.Activation('softmax', dtype='float32')(x)
+ x = tf_keras.layers.Dense(1, name='fc1')(x)
+ x = tf_keras.layers.Dense(num_classes, name='fc1000')(x)
+ x = tf_keras.layers.Activation('softmax', dtype='float32')(x)
- return tf.keras.models.Model(img_input, x, name='trivial')
+ return tf_keras.models.Model(img_input, x, name='trivial')
diff --git a/official/legacy/image_classification/vgg/__init__.py b/official/legacy/image_classification/vgg/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/legacy/image_classification/vgg/__init__.py
+++ b/official/legacy/image_classification/vgg/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/image_classification/vgg/vgg_config.py b/official/legacy/image_classification/vgg/vgg_config.py
index 0bf936744fa..639c3a3a0ef 100644
--- a/official/legacy/image_classification/vgg/vgg_config.py
+++ b/official/legacy/image_classification/vgg/vgg_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -29,16 +29,27 @@ class VGGModelConfig(base_configs.ModelConfig):
'batch_size': None,
'use_l2_regularizer': True
})
- loss: base_configs.LossConfig = base_configs.LossConfig(
- name='sparse_categorical_crossentropy')
- optimizer: base_configs.OptimizerConfig = base_configs.OptimizerConfig(
- name='momentum', epsilon=0.001, momentum=0.9, moving_average_decay=None)
- learning_rate: base_configs.LearningRateConfig = (
- base_configs.LearningRateConfig(
+ loss: base_configs.LossConfig = dataclasses.field(
+ default_factory=lambda: base_configs.LossConfig( # pylint: disable=g-long-lambda
+ name='sparse_categorical_crossentropy'
+ )
+ )
+ optimizer: base_configs.OptimizerConfig = dataclasses.field(
+ default_factory=lambda: base_configs.OptimizerConfig( # pylint: disable=g-long-lambda
+ name='momentum',
+ epsilon=0.001,
+ momentum=0.9,
+ moving_average_decay=None,
+ )
+ )
+ learning_rate: base_configs.LearningRateConfig = dataclasses.field(
+ default_factory=lambda: base_configs.LearningRateConfig( # pylint: disable=g-long-lambda
name='stepwise',
initial_lr=0.01,
examples_per_epoch=1281167,
boundaries=[30, 60],
warmup_epochs=0,
- scale_by_batch_size=1. / 256.,
- multipliers=[0.01 / 256, 0.001 / 256, 0.0001 / 256]))
+ scale_by_batch_size=1.0 / 256.0,
+ multipliers=[0.01 / 256, 0.001 / 256, 0.0001 / 256],
+ )
+ )
diff --git a/official/legacy/image_classification/vgg/vgg_model.py b/official/legacy/image_classification/vgg/vgg_model.py
index b93e22555c5..7bd745fa4e0 100644
--- a/official/legacy/image_classification/vgg/vgg_model.py
+++ b/official/legacy/image_classification/vgg/vgg_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,19 +14,19 @@
"""VGG16 model for Keras.
-Adapted from tf.keras.applications.vgg16.VGG16().
+Adapted from tf_keras.applications.vgg16.VGG16().
Related papers/blogs:
- https://arxiv.org/abs/1409.1556
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-layers = tf.keras.layers
+layers = tf_keras.layers
def _gen_l2_regularizer(use_l2_regularizer=True, l2_weight_decay=1e-4):
- return tf.keras.regularizers.L2(
+ return tf_keras.regularizers.L2(
l2_weight_decay) if use_l2_regularizer else None
@@ -53,7 +53,7 @@ def vgg16(num_classes,
x = img_input
- if tf.keras.backend.image_data_format() == 'channels_first':
+ if tf_keras.backend.image_data_format() == 'channels_first':
x = layers.Permute((3, 1, 2))(x)
bn_axis = 1
else: # channels_last
@@ -266,4 +266,4 @@ def vgg16(num_classes,
x = layers.Activation('softmax', dtype='float32')(x)
# Create model.
- return tf.keras.Model(img_input, x, name='vgg16')
+ return tf_keras.Model(img_input, x, name='vgg16')
diff --git a/official/legacy/transformer/__init__.py b/official/legacy/transformer/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/transformer/__init__.py
+++ b/official/legacy/transformer/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/transformer/attention_layer.py b/official/legacy/transformer/attention_layer.py
index e759b462cb1..1209ce784d4 100644
--- a/official/legacy/transformer/attention_layer.py
+++ b/official/legacy/transformer/attention_layer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,12 @@
"""Implementation of multiheaded attention and self-attention layers."""
import math
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
-class Attention(tf.keras.layers.Layer):
+
+class Attention(tf_keras.layers.Layer):
"""Multi-headed attention layer."""
def __init__(self, hidden_size, num_heads, attention_dropout):
@@ -46,31 +48,31 @@ def build(self, input_shape):
def _glorot_initializer(fan_in, fan_out):
limit = math.sqrt(6.0 / (fan_in + fan_out))
- return tf.keras.initializers.RandomUniform(minval=-limit, maxval=limit)
+ return tf_keras.initializers.RandomUniform(minval=-limit, maxval=limit)
attention_initializer = _glorot_initializer(input_shape.as_list()[-1],
self.hidden_size)
- self.query_dense_layer = tf.keras.layers.experimental.EinsumDense(
+ self.query_dense_layer = tf_keras.layers.EinsumDense(
"BTE,ENH->BTNH",
output_shape=(None, self.num_heads, size_per_head),
- kernel_initializer=attention_initializer,
+ kernel_initializer=tf_utils.clone_initializer(attention_initializer),
bias_axes=None,
name="query")
- self.key_dense_layer = tf.keras.layers.experimental.EinsumDense(
+ self.key_dense_layer = tf_keras.layers.EinsumDense(
"BTE,ENH->BTNH",
output_shape=(None, self.num_heads, size_per_head),
- kernel_initializer=attention_initializer,
+ kernel_initializer=tf_utils.clone_initializer(attention_initializer),
bias_axes=None,
name="key")
- self.value_dense_layer = tf.keras.layers.experimental.EinsumDense(
+ self.value_dense_layer = tf_keras.layers.EinsumDense(
"BTE,ENH->BTNH",
output_shape=(None, self.num_heads, size_per_head),
- kernel_initializer=attention_initializer,
+ kernel_initializer=tf_utils.clone_initializer(attention_initializer),
bias_axes=None,
name="value")
output_initializer = _glorot_initializer(self.hidden_size, self.hidden_size)
- self.output_dense_layer = tf.keras.layers.experimental.EinsumDense(
+ self.output_dense_layer = tf_keras.layers.EinsumDense(
"BTNH,NHE->BTE",
output_shape=(None, self.hidden_size),
kernel_initializer=output_initializer,
diff --git a/official/legacy/transformer/beam_search_v1.py b/official/legacy/transformer/beam_search_v1.py
index 533cc01b211..93f5e93c96a 100644
--- a/official/legacy/transformer/beam_search_v1.py
+++ b/official/legacy/transformer/beam_search_v1.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/transformer/compute_bleu.py b/official/legacy/transformer/compute_bleu.py
index c1b01e11a97..254e8a10633 100644
--- a/official/legacy/transformer/compute_bleu.py
+++ b/official/legacy/transformer/compute_bleu.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -27,7 +27,7 @@
from absl import logging
import six
from six.moves import range
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer.utils import metrics
from official.legacy.transformer.utils import tokenizer
diff --git a/official/legacy/transformer/compute_bleu_test.py b/official/legacy/transformer/compute_bleu_test.py
index 24159248eb2..3f1e8e0ee93 100644
--- a/official/legacy/transformer/compute_bleu_test.py
+++ b/official/legacy/transformer/compute_bleu_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import tempfile
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import compute_bleu
diff --git a/official/legacy/transformer/data_download.py b/official/legacy/transformer/data_download.py
index ca3c65c3f69..f9eb887e568 100644
--- a/official/legacy/transformer/data_download.py
+++ b/official/legacy/transformer/data_download.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -188,7 +188,7 @@ def download_and_extract(path, url, input_filename, target_filename):
Full paths to extracted input and target files.
Raises:
- OSError: if the the download/extraction fails.
+ OSError: if the download/extraction fails.
"""
# Check if extracted files already exist in path
input_file = find_file(path, input_filename)
diff --git a/official/legacy/transformer/data_pipeline.py b/official/legacy/transformer/data_pipeline.py
index 484c8e97a59..1cd2acc2ea4 100644
--- a/official/legacy/transformer/data_pipeline.py
+++ b/official/legacy/transformer/data_pipeline.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -50,7 +50,7 @@
import os
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.utils.misc import model_helpers
diff --git a/official/legacy/transformer/embedding_layer.py b/official/legacy/transformer/embedding_layer.py
index 398a950df2b..282ce7653f9 100644
--- a/official/legacy/transformer/embedding_layer.py
+++ b/official/legacy/transformer/embedding_layer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Implementation of embedding layer with shared weights."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class EmbeddingSharedWeights(tf.keras.layers.Layer):
+class EmbeddingSharedWeights(tf_keras.layers.Layer):
"""Calculates input embeddings and pre-softmax linear with shared weights."""
def __init__(self, vocab_size, hidden_size):
diff --git a/official/legacy/transformer/ffn_layer.py b/official/legacy/transformer/ffn_layer.py
index 8e24a1e8428..f2d0c41bdd4 100644
--- a/official/legacy/transformer/ffn_layer.py
+++ b/official/legacy/transformer/ffn_layer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Implementation of fully connected network."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class FeedForwardNetwork(tf.keras.layers.Layer):
+class FeedForwardNetwork(tf_keras.layers.Layer):
"""Fully connected feedforward network."""
def __init__(self, hidden_size, filter_size, relu_dropout):
@@ -34,12 +34,12 @@ def __init__(self, hidden_size, filter_size, relu_dropout):
self.relu_dropout = relu_dropout
def build(self, input_shape):
- self.filter_dense_layer = tf.keras.layers.Dense(
+ self.filter_dense_layer = tf_keras.layers.Dense(
self.filter_size,
use_bias=True,
activation=tf.nn.relu,
name="filter_layer")
- self.output_dense_layer = tf.keras.layers.Dense(
+ self.output_dense_layer = tf_keras.layers.Dense(
self.hidden_size, use_bias=True, name="output_layer")
super(FeedForwardNetwork, self).build(input_shape)
diff --git a/official/legacy/transformer/metrics.py b/official/legacy/transformer/metrics.py
index b469e6c6f67..f761bf3436f 100644
--- a/official/legacy/transformer/metrics.py
+++ b/official/legacy/transformer/metrics.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,7 +25,7 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _pad_tensors_to_same_length(x, y):
@@ -131,7 +131,7 @@ def padded_neg_log_perplexity(logits, labels, vocab_size):
return -num, den
-class MetricLayer(tf.keras.layers.Layer):
+class MetricLayer(tf_keras.layers.Layer):
"""Custom a layer of metrics for Transformer model."""
def __init__(self, vocab_size):
@@ -144,11 +144,11 @@ def build(self, input_shape):
neg_log_perplexity = functools.partial(
padded_neg_log_perplexity, vocab_size=self.vocab_size)
self.metric_mean_fns = [
- (tf.keras.metrics.Mean("accuracy"), padded_accuracy),
- (tf.keras.metrics.Mean("accuracy_top5"), padded_accuracy_top5),
- (tf.keras.metrics.Mean("accuracy_per_sequence"),
+ (tf_keras.metrics.Mean("accuracy"), padded_accuracy),
+ (tf_keras.metrics.Mean("accuracy_top5"), padded_accuracy_top5),
+ (tf_keras.metrics.Mean("accuracy_per_sequence"),
padded_sequence_accuracy),
- (tf.keras.metrics.Mean("neg_log_perplexity"), neg_log_perplexity),
+ (tf_keras.metrics.Mean("neg_log_perplexity"), neg_log_perplexity),
]
super(MetricLayer, self).build(input_shape)
diff --git a/official/legacy/transformer/misc.py b/official/legacy/transformer/misc.py
index ff8930a6601..72ccfe4da58 100644
--- a/official/legacy/transformer/misc.py
+++ b/official/legacy/transformer/misc.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
# pylint: disable=g-bad-import-order
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import model_params
from official.utils.flags import core as flags_core
@@ -250,7 +250,7 @@ def get_callbacks():
callbacks.append(time_callback)
if FLAGS.enable_tensorboard:
- tensorboard_callback = tf.keras.callbacks.TensorBoard(
+ tensorboard_callback = tf_keras.callbacks.TensorBoard(
log_dir=FLAGS.model_dir)
callbacks.append(tensorboard_callback)
diff --git a/official/legacy/transformer/model_params.py b/official/legacy/transformer/model_params.py
index 70e464be20a..8e6bc1cf814 100644
--- a/official/legacy/transformer/model_params.py
+++ b/official/legacy/transformer/model_params.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/transformer/model_utils.py b/official/legacy/transformer/model_utils.py
index 36095238822..0913b98bc16 100644
--- a/official/legacy/transformer/model_utils.py
+++ b/official/legacy/transformer/model_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import math
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# Very low numbers to represent -infinity. We do not actually use -Inf, since we
# want to be able to multiply these values by zero to get zero. (-Inf * 0 = NaN)
diff --git a/official/legacy/transformer/model_utils_test.py b/official/legacy/transformer/model_utils_test.py
index 0758caa1870..71276781049 100644
--- a/official/legacy/transformer/model_utils_test.py
+++ b/official/legacy/transformer/model_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Test Transformer model helper methods."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import model_utils
diff --git a/official/legacy/transformer/optimizer.py b/official/legacy/transformer/optimizer.py
index 70e96ab6bad..a20016efc05 100644
--- a/official/legacy/transformer/optimizer.py
+++ b/official/legacy/transformer/optimizer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Optimizer from addons and learning rate scheduler."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class LearningRateSchedule(tf.keras.optimizers.schedules.LearningRateSchedule):
+class LearningRateSchedule(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Learning rate schedule."""
def __init__(self, initial_learning_rate, hidden_size, warmup_steps):
diff --git a/official/legacy/transformer/transformer.py b/official/legacy/transformer/transformer.py
index ed5d874900d..752aac30701 100644
--- a/official/legacy/transformer/transformer.py
+++ b/official/legacy/transformer/transformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
Transformer model code source: https://github.com/tensorflow/tensor2tensor
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import attention_layer
from official.legacy.transformer import embedding_layer
@@ -38,32 +38,32 @@ def create_model(params, is_train):
"""Creates transformer model."""
with tf.name_scope("model"):
if is_train:
- inputs = tf.keras.layers.Input((None,), dtype="int64", name="inputs")
- targets = tf.keras.layers.Input((None,), dtype="int64", name="targets")
+ inputs = tf_keras.layers.Input((None,), dtype="int64", name="inputs")
+ targets = tf_keras.layers.Input((None,), dtype="int64", name="targets")
internal_model = Transformer(params, name="transformer_v2")
logits = internal_model([inputs, targets], training=is_train)
vocab_size = params["vocab_size"]
label_smoothing = params["label_smoothing"]
if params["enable_metrics_in_training"]:
logits = metrics.MetricLayer(vocab_size)([logits, targets])
- logits = tf.keras.layers.Lambda(
+ logits = tf_keras.layers.Lambda(
lambda x: x, name="logits", dtype=tf.float32)(
logits)
- model = tf.keras.Model([inputs, targets], logits)
+ model = tf_keras.Model([inputs, targets], logits)
loss = metrics.transformer_loss(logits, targets, label_smoothing,
vocab_size)
model.add_loss(loss)
return model
else:
- inputs = tf.keras.layers.Input((None,), dtype="int64", name="inputs")
+ inputs = tf_keras.layers.Input((None,), dtype="int64", name="inputs")
internal_model = Transformer(params, name="transformer_v2")
ret = internal_model([inputs], training=is_train)
outputs, scores = ret["outputs"], ret["scores"]
- return tf.keras.Model(inputs, [outputs, scores])
+ return tf_keras.Model(inputs, [outputs, scores])
-class Transformer(tf.keras.Model):
+class Transformer(tf_keras.Model):
"""Transformer model with Keras.
Implemented as described in: https://arxiv.org/pdf/1706.03762.pdf
@@ -339,7 +339,7 @@ def predict(self, encoder_outputs, encoder_decoder_attention_bias, training):
return {"outputs": top_decoded_ids, "scores": top_scores}
-class PrePostProcessingWrapper(tf.keras.layers.Layer):
+class PrePostProcessingWrapper(tf_keras.layers.Layer):
"""Wrapper class that applies layer pre-processing and post-processing."""
def __init__(self, layer, params):
@@ -350,7 +350,7 @@ def __init__(self, layer, params):
def build(self, input_shape):
# Create normalization layer
- self.layer_norm = tf.keras.layers.LayerNormalization(
+ self.layer_norm = tf_keras.layers.LayerNormalization(
epsilon=1e-6, dtype="float32")
super(PrePostProcessingWrapper, self).build(input_shape)
@@ -375,7 +375,7 @@ def call(self, x, *args, **kwargs):
return x + y
-class EncoderStack(tf.keras.layers.Layer):
+class EncoderStack(tf_keras.layers.Layer):
"""Transformer encoder stack.
The encoder stack is made up of N identical layers. Each layer is composed
@@ -406,7 +406,7 @@ def build(self, input_shape):
])
# Create final layer normalization layer.
- self.output_normalization = tf.keras.layers.LayerNormalization(
+ self.output_normalization = tf_keras.layers.LayerNormalization(
epsilon=1e-6, dtype="float32")
super(EncoderStack, self).build(input_shape)
@@ -446,7 +446,7 @@ def call(self, encoder_inputs, attention_bias, inputs_padding, training):
return self.output_normalization(encoder_inputs)
-class DecoderStack(tf.keras.layers.Layer):
+class DecoderStack(tf_keras.layers.Layer):
"""Transformer decoder stack.
Like the encoder stack, the decoder stack is made up of N identical layers.
@@ -480,7 +480,7 @@ def build(self, input_shape):
PrePostProcessingWrapper(enc_dec_attention_layer, params),
PrePostProcessingWrapper(feed_forward_network, params)
])
- self.output_normalization = tf.keras.layers.LayerNormalization(
+ self.output_normalization = tf_keras.layers.LayerNormalization(
epsilon=1e-6, dtype="float32")
super(DecoderStack, self).build(input_shape)
diff --git a/official/legacy/transformer/transformer_forward_test.py b/official/legacy/transformer/transformer_forward_test.py
index 5efdc4178f4..9024c36971c 100644
--- a/official/legacy/transformer/transformer_forward_test.py
+++ b/official/legacy/transformer/transformer_forward_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Forward pass test for Transformer model refactoring."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import metrics
from official.legacy.transformer import model_params
@@ -30,7 +30,7 @@ def _count_params(layer, trainable_only=True):
else:
return int(
np.sum([
- tf.keras.backend.count_params(p) for p in layer.trainable_weights
+ tf_keras.backend.count_params(p) for p in layer.trainable_weights
]))
@@ -66,8 +66,8 @@ def _create_model(params, is_train):
name="transformer_v2")
if is_train:
- inputs = tf.keras.layers.Input((None,), dtype="int64", name="inputs")
- targets = tf.keras.layers.Input((None,), dtype="int64", name="targets")
+ inputs = tf_keras.layers.Input((None,), dtype="int64", name="inputs")
+ targets = tf_keras.layers.Input((None,), dtype="int64", name="targets")
internal_model = models.Seq2SeqTransformer(**model_kwargs)
logits = internal_model(
dict(inputs=inputs, targets=targets), training=is_train)
@@ -75,24 +75,24 @@ def _create_model(params, is_train):
label_smoothing = params["label_smoothing"]
if params["enable_metrics_in_training"]:
logits = metrics.MetricLayer(vocab_size)([logits, targets])
- logits = tf.keras.layers.Lambda(
+ logits = tf_keras.layers.Lambda(
lambda x: x, name="logits", dtype=tf.float32)(
logits)
- model = tf.keras.Model([inputs, targets], logits)
+ model = tf_keras.Model([inputs, targets], logits)
loss = metrics.transformer_loss(logits, targets, label_smoothing,
vocab_size)
model.add_loss(loss)
return model
batch_size = params["decode_batch_size"] if params["padded_decode"] else None
- inputs = tf.keras.layers.Input((None,),
+ inputs = tf_keras.layers.Input((None,),
batch_size=batch_size,
dtype="int64",
name="inputs")
internal_model = models.Seq2SeqTransformer(**model_kwargs)
ret = internal_model(dict(inputs=inputs), training=is_train)
outputs, scores = ret["outputs"], ret["scores"]
- return tf.keras.Model(inputs, [outputs, scores])
+ return tf_keras.Model(inputs, [outputs, scores])
class TransformerForwardTest(tf.test.TestCase):
diff --git a/official/legacy/transformer/transformer_layers_test.py b/official/legacy/transformer/transformer_layers_test.py
index c2080443965..a3996abde02 100644
--- a/official/legacy/transformer/transformer_layers_test.py
+++ b/official/legacy/transformer/transformer_layers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for layers in Transformer."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import attention_layer
from official.legacy.transformer import embedding_layer
@@ -109,10 +109,10 @@ def test_feed_forward_network(self):
def test_metric_layer(self):
vocab_size = 50
- logits = tf.keras.layers.Input((None, vocab_size),
+ logits = tf_keras.layers.Input((None, vocab_size),
dtype="float32",
name="logits")
- targets = tf.keras.layers.Input((None,), dtype="int64", name="targets")
+ targets = tf_keras.layers.Input((None,), dtype="int64", name="targets")
output_logits = metrics.MetricLayer(vocab_size)([logits, targets])
self.assertEqual(output_logits.shape.as_list(), [
None,
diff --git a/official/legacy/transformer/transformer_main.py b/official/legacy/transformer/transformer_main.py
index ec1e7634045..84507284c25 100644
--- a/official/legacy/transformer/transformer_main.py
+++ b/official/legacy/transformer/transformer_main.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,12 +21,10 @@
import os
import tempfile
-# Import libraries
-
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.transformer import compute_bleu
@@ -208,7 +206,7 @@ def train(self):
current_step = opt.iterations.numpy()
if params["use_ctl"]:
- train_loss_metric = tf.keras.metrics.Mean(
+ train_loss_metric = tf_keras.metrics.Mean(
"training_loss", dtype=tf.float32)
if params["enable_tensorboard"]:
summary_writer = tf.summary.create_file_writer(
@@ -411,7 +409,7 @@ def _create_callbacks(self, cur_log_dir, params):
if params["enable_checkpointing"]:
ckpt_full_path = os.path.join(cur_log_dir, "cp-{epoch:04d}.ckpt")
callbacks.append(
- tf.keras.callbacks.ModelCheckpoint(
+ tf_keras.callbacks.ModelCheckpoint(
ckpt_full_path, save_weights_only=params["save_weights_only"]))
return callbacks
@@ -434,7 +432,7 @@ def _create_optimizer(self):
lr_schedule = optimizer.LearningRateSchedule(
params["learning_rate"], params["hidden_size"],
params["learning_rate_warmup_steps"])
- opt = tf.keras.optimizers.Adam(
+ opt = tf_keras.optimizers.Adam(
lr_schedule,
params["optimizer_adam_beta1"],
params["optimizer_adam_beta2"],
@@ -457,8 +455,6 @@ def _ensure_dir(log_dir):
def main(_):
flags_obj = flags.FLAGS
- if flags_obj.enable_mlir_bridge:
- tf.config.experimental.enable_mlir_bridge()
task = TransformerTask(flags_obj)
# Execute flag override logic for better model performance
diff --git a/official/legacy/transformer/transformer_main_test.py b/official/legacy/transformer/transformer_main_test.py
index 82077858102..0e8e1dae920 100644
--- a/official/legacy/transformer/transformer_main_test.py
+++ b/official/legacy/transformer/transformer_main_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl import flags
from absl.testing import flagsaver
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.eager import context # pylint: disable=ungrouped-imports
from official.legacy.transformer import misc
from official.legacy.transformer import transformer_main
diff --git a/official/legacy/transformer/transformer_test.py b/official/legacy/transformer/transformer_test.py
index a6cedb48e1d..7e5eda50d7d 100644
--- a/official/legacy/transformer/transformer_test.py
+++ b/official/legacy/transformer/transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Test Transformer model."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import model_params
from official.legacy.transformer import transformer
diff --git a/official/legacy/transformer/translate.py b/official/legacy/transformer/translate.py
index abbf82f5f1a..1da9b832b92 100644
--- a/official/legacy/transformer/translate.py
+++ b/official/legacy/transformer/translate.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,9 @@
"""Translate text or files using trained transformer model."""
-# Import libraries
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer.utils import tokenizer
@@ -110,7 +109,7 @@ def input_generator():
if distribution_strategy:
for j in range(batch_size - len(lines)):
lines.append([tokenizer.EOS_ID])
- batch = tf.keras.preprocessing.sequence.pad_sequences(
+ batch = tf_keras.preprocessing.sequence.pad_sequences(
lines,
maxlen=params["decode_max_length"],
dtype="int32",
diff --git a/official/legacy/transformer/utils/__init__.py b/official/legacy/transformer/utils/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/legacy/transformer/utils/__init__.py
+++ b/official/legacy/transformer/utils/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/transformer/utils/metrics.py b/official/legacy/transformer/utils/metrics.py
index 23261ac474a..bcc683a1eb5 100644
--- a/official/legacy/transformer/utils/metrics.py
+++ b/official/legacy/transformer/utils/metrics.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -471,7 +471,7 @@ def _lcs(x, y):
def _f_lcs(llcs, m, n):
"""Computes the LCS-based F-measure score.
- Source: http://research.microsoft.com/en-us/um/people/cyl/download/papers/
+ Source: https://research.microsoft.com/en-us/um/people/cyl/download/papers/
rouge-working-note-v1.3.1.pdf
Args:
diff --git a/official/legacy/transformer/utils/tokenizer.py b/official/legacy/transformer/utils/tokenizer.py
index 9533846d2fc..5d71399a94b 100644
--- a/official/legacy/transformer/utils/tokenizer.py
+++ b/official/legacy/transformer/utils/tokenizer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -28,7 +28,7 @@
import numpy as np
import six
from six.moves import xrange # pylint: disable=redefined-builtin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=g-complex-comprehension
PAD = ""
diff --git a/official/legacy/transformer/utils/tokenizer_test.py b/official/legacy/transformer/utils/tokenizer_test.py
index 2b582b99c6f..e81a07295e3 100644
--- a/official/legacy/transformer/utils/tokenizer_test.py
+++ b/official/legacy/transformer/utils/tokenizer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import collections
import tempfile
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer.utils import tokenizer
diff --git a/official/legacy/xlnet/__init__.py b/official/legacy/xlnet/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/legacy/xlnet/__init__.py
+++ b/official/legacy/xlnet/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/xlnet/classifier_utils.py b/official/legacy/xlnet/classifier_utils.py
index 27aaf4ade18..102ec16f1e2 100644
--- a/official/legacy/xlnet/classifier_utils.py
+++ b/official/legacy/xlnet/classifier_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/xlnet/common_flags.py b/official/legacy/xlnet/common_flags.py
index b1ee5c3e86d..4c2078ba0ba 100644
--- a/official/legacy/xlnet/common_flags.py
+++ b/official/legacy/xlnet/common_flags.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/xlnet/data_utils.py b/official/legacy/xlnet/data_utils.py
index 0048832d6ea..e5cf5020a8b 100644
--- a/official/legacy/xlnet/data_utils.py
+++ b/official/legacy/xlnet/data_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
special_symbols = {
"": 0,
@@ -529,7 +529,7 @@ def parser(record):
for key in list(example.keys()):
val = example[key]
- if tf.keras.backend.is_sparse(val):
+ if tf_keras.backend.is_sparse(val):
val = tf.sparse.to_dense(val)
if val.dtype == tf.int64:
val = tf.cast(val, tf.int32)
diff --git a/official/legacy/xlnet/optimization.py b/official/legacy/xlnet/optimization.py
index 28be940e1c6..c891178874a 100644
--- a/official/legacy/xlnet/optimization.py
+++ b/official/legacy/xlnet/optimization.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,11 +15,11 @@
"""Functions and classes related to optimization (weight updates)."""
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp import optimization
-class WarmUp(tf.keras.optimizers.schedules.LearningRateSchedule):
+class WarmUp(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Applys a warmup schedule on a given learning rate decay schedule."""
def __init__(self,
@@ -69,7 +69,7 @@ def create_optimizer(init_lr,
weight_decay_rate=0.0):
"""Creates an optimizer with learning rate schedule."""
# Implements linear decay of the learning rate.
- learning_rate_fn = tf.keras.optimizers.schedules.PolynomialDecay(
+ learning_rate_fn = tf_keras.optimizers.schedules.PolynomialDecay(
initial_learning_rate=init_lr,
decay_steps=num_train_steps - num_warmup_steps,
end_learning_rate=init_lr * min_lr_ratio)
@@ -92,7 +92,7 @@ def create_optimizer(init_lr,
include_in_weight_decay=["r_s_bias", "r_r_bias", "r_w_bias"])
else:
logging.info("Using Adam with adam_epsilon=%.9f", (adam_epsilon))
- optimizer = tf.keras.optimizers.Adam(
+ optimizer = tf_keras.optimizers.legacy.Adam(
learning_rate=learning_rate_fn, epsilon=adam_epsilon)
return optimizer, learning_rate_fn
diff --git a/official/legacy/xlnet/preprocess_classification_data.py b/official/legacy/xlnet/preprocess_classification_data.py
index d517e486b03..7768e47b778 100644
--- a/official/legacy/xlnet/preprocess_classification_data.py
+++ b/official/legacy/xlnet/preprocess_classification_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,12 +18,11 @@
import csv
import os
-# Import libraries
from absl import app
from absl import flags
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import sentencepiece as spm
from official.legacy.xlnet import classifier_utils
diff --git a/official/legacy/xlnet/preprocess_pretrain_data.py b/official/legacy/xlnet/preprocess_pretrain_data.py
index aaf60ba5e4a..197334d8823 100644
--- a/official/legacy/xlnet/preprocess_pretrain_data.py
+++ b/official/legacy/xlnet/preprocess_pretrain_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,6 @@
import os
import random
-# Import libraries
from absl import app
from absl import flags
from absl import logging
@@ -129,7 +128,7 @@ def _create_data(idx, input_paths):
continue
input_data = np.array(input_data, dtype=np.int64)
- sent_ids = np.array(sent_ids, dtype=np.bool)
+ sent_ids = np.array(sent_ids, dtype=bool)
total_line_cnt += line_cnt
input_shards.append((input_data, sent_ids))
@@ -348,7 +347,7 @@ def _is_start_piece(piece):
def _sample_mask(sp, seg, reverse=False, max_gram=5, goal_num_predict=None):
"""Samples `goal_num_predict` tokens for partial prediction."""
seg_len = len(seg)
- mask = np.array([False] * seg_len, dtype=np.bool)
+ mask = np.array([False] * seg_len, dtype=bool)
num_predict = 0
@@ -411,7 +410,7 @@ def _sample_mask_ngram(sp, seg, reverse=False, max_gram=5,
"""Sample `goal_num_predict` tokens for partial prediction."""
seg_len = len(seg)
- mask = np.array([False] * seg_len, dtype=np.bool)
+ mask = np.array([False] * seg_len, dtype=bool)
num_predict = 0
@@ -614,7 +613,7 @@ def _convert_example(example, use_bfloat16):
"""Cast int64 into int32 and float32 to bfloat16 if use_bfloat16."""
for key in list(example.keys()):
val = example[key]
- if tf.keras.backend.is_sparse(val):
+ if tf_keras.backend.is_sparse(val):
val = tf.sparse.to_dense(val)
if val.dtype == tf.int64:
val = tf.cast(val, tf.int32)
diff --git a/official/legacy/xlnet/preprocess_squad_data.py b/official/legacy/xlnet/preprocess_squad_data.py
index e99177c838e..7e563019d43 100644
--- a/official/legacy/xlnet/preprocess_squad_data.py
+++ b/official/legacy/xlnet/preprocess_squad_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,11 +18,10 @@
import os
import random
-# Import libraries
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import sentencepiece as spm
from official.legacy.xlnet import squad_utils
diff --git a/official/legacy/xlnet/preprocess_utils.py b/official/legacy/xlnet/preprocess_utils.py
index 19cae9174c2..7594ba40f75 100644
--- a/official/legacy/xlnet/preprocess_utils.py
+++ b/official/legacy/xlnet/preprocess_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/legacy/xlnet/run_classifier.py b/official/legacy/xlnet/run_classifier.py
index 258e6116ab3..0688a27c1b2 100644
--- a/official/legacy/xlnet/run_classifier.py
+++ b/official/legacy/xlnet/run_classifier.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,12 @@
"""XLNet classification finetuning runner in tf2.0."""
import functools
-# Import libraries
from absl import app
from absl import flags
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.common import distribute_utils
from official.legacy.xlnet import common_flags
@@ -123,7 +122,7 @@ def _run_evaluation(test_iterator):
def get_metric_fn():
- train_acc_metric = tf.keras.metrics.SparseCategoricalAccuracy(
+ train_acc_metric = tf_keras.metrics.SparseCategoricalAccuracy(
"acc", dtype=tf.float32)
return train_acc_metric
diff --git a/official/legacy/xlnet/run_pretrain.py b/official/legacy/xlnet/run_pretrain.py
index 311f283a9cb..1eef1bd5b68 100644
--- a/official/legacy/xlnet/run_pretrain.py
+++ b/official/legacy/xlnet/run_pretrain.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,11 +17,10 @@
import functools
import os
-# Import libraries
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.common import distribute_utils
from official.legacy.xlnet import common_flags
diff --git a/official/legacy/xlnet/run_squad.py b/official/legacy/xlnet/run_squad.py
index 29a5c5c451c..3c09a2ea0cf 100644
--- a/official/legacy/xlnet/run_squad.py
+++ b/official/legacy/xlnet/run_squad.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,12 +19,11 @@
import os
import pickle
-# Import libraries
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
import sentencepiece as spm
from official.common import distribute_utils
diff --git a/official/legacy/xlnet/squad_utils.py b/official/legacy/xlnet/squad_utils.py
index 641e8818f48..9828a0fe4cb 100644
--- a/official/legacy/xlnet/squad_utils.py
+++ b/official/legacy/xlnet/squad_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -30,7 +30,7 @@
from absl import logging
import numpy as np
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.xlnet import data_utils
from official.legacy.xlnet import preprocess_utils
diff --git a/official/legacy/xlnet/training_utils.py b/official/legacy/xlnet/training_utils.py
index 5fd924e8bbf..dd4fb8c8858 100644
--- a/official/legacy/xlnet/training_utils.py
+++ b/official/legacy/xlnet/training_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from typing import Any, Callable, Dict, Optional, Text
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import model_training_utils
from official.legacy.xlnet import data_utils
@@ -51,11 +51,11 @@ def train(
train_input_fn: Callable,
total_training_steps: int,
steps_per_loop: int,
- optimizer: tf.keras.optimizers.Optimizer,
- learning_rate_fn: tf.keras.optimizers.schedules.LearningRateSchedule,
- eval_fn: Optional[Callable[[tf.keras.Model, int, tf.summary.SummaryWriter],
+ optimizer: tf_keras.optimizers.Optimizer,
+ learning_rate_fn: tf_keras.optimizers.schedules.LearningRateSchedule,
+ eval_fn: Optional[Callable[[tf_keras.Model, int, tf.summary.SummaryWriter],
Any]] = None,
- metric_fn: Optional[Callable[[], tf.keras.metrics.Metric]] = None,
+ metric_fn: Optional[Callable[[], tf_keras.metrics.Metric]] = None,
init_checkpoint: Optional[Text] = None,
init_from_transformerxl: Optional[bool] = False,
model_dir: Optional[Text] = None,
@@ -140,7 +140,7 @@ def train(
if not hasattr(model, "optimizer"):
raise ValueError("User should set optimizer attribute to model.")
- train_loss_metric = tf.keras.metrics.Mean("training_loss", dtype=tf.float32)
+ train_loss_metric = tf_keras.metrics.Mean("training_loss", dtype=tf.float32)
train_metric = None
if metric_fn:
train_metric = metric_fn()
diff --git a/official/legacy/xlnet/xlnet_config.py b/official/legacy/xlnet/xlnet_config.py
index d8ee7e6a07f..76f638e1b92 100644
--- a/official/legacy/xlnet/xlnet_config.py
+++ b/official/legacy/xlnet/xlnet_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import json
import os
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def create_run_config(is_training, is_finetune, flags):
diff --git a/official/legacy/xlnet/xlnet_modeling.py b/official/legacy/xlnet/xlnet_modeling.py
index f03354f62ab..fa2b2c2e129 100644
--- a/official/legacy/xlnet/xlnet_modeling.py
+++ b/official/legacy/xlnet/xlnet_modeling.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,22 +17,22 @@
import copy
import warnings
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.xlnet import data_utils
from official.nlp.modeling import networks
def gelu(x):
- return tf.keras.activations.gelu(x, approximate=True)
+ return tf_keras.activations.gelu(x, approximate=True)
def _get_initializer(flags):
"""Get variable initializer."""
if flags.init_method == "uniform":
- initializer = tf.keras.initializers.RandomUniform(
+ initializer = tf_keras.initializers.RandomUniform(
minval=-flags.init_range, maxval=flags.init_range)
elif flags.init_method == "normal":
- initializer = tf.keras.initializers.RandomNormal(stddev=flags.init_std)
+ initializer = tf_keras.initializers.RandomNormal(stddev=flags.init_std)
else:
raise ValueError("Initializer {} not supported".format(flags.init_method))
return initializer
@@ -78,7 +78,7 @@ def _cache_mem(curr_out, prev_mem, mem_len, reuse_len=None):
else:
new_mem = tf.concat([prev_mem, curr_out], 0)[-mem_len:]
- return tf.keras.backend.stop_gradient(new_mem)
+ return tf_keras.backend.stop_gradient(new_mem)
def is_special_none_tensor(tensor):
@@ -86,8 +86,8 @@ def is_special_none_tensor(tensor):
return tensor.shape.ndims == 0 and tensor.dtype == tf.int32
-@tf.keras.utils.register_keras_serializable(package="Text")
-class RelativePositionEncoding(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class RelativePositionEncoding(tf_keras.layers.Layer):
"""Creates a relative positional encoding.
This layer creates a relative positional encoding as described in
@@ -132,7 +132,7 @@ def call(self, pos_seq, batch_size=None):
return pos_emb
-class RelativeAttention(tf.keras.layers.Layer):
+class RelativeAttention(tf_keras.layers.Layer):
"""Core calculations for relative attention."""
def __init__(self, dropout_att, scale):
@@ -143,7 +143,7 @@ def __init__(self, dropout_att, scale):
def build(self, unused_input_shapes):
"""Implements build() for the layer."""
- self.attention_probs_dropout = tf.keras.layers.Dropout(
+ self.attention_probs_dropout = tf_keras.layers.Dropout(
rate=self.dropout_att)
super(RelativeAttention, self).build(unused_input_shapes)
@@ -185,7 +185,7 @@ def call(self, q_head, k_head_h, v_head_h, k_head_r, seg_embed, seg_mat,
return attn_vec
-class PositionwiseFF(tf.keras.layers.Layer):
+class PositionwiseFF(tf_keras.layers.Layer):
"""Positionwise feed-forward layer."""
def __init__(self, d_model, d_inner, dropout, kernel_initializer,
@@ -207,20 +207,20 @@ def build(self, unused_input_shapes):
raise (ValueError("Unsupported activation type {}".format(
self.activation_type)))
self.inner_projection_layer = (
- tf.keras.layers.Dense(
+ tf_keras.layers.Dense(
units=self.d_inner,
activation=activation,
kernel_initializer=self.kernel_initializer,
name="layer_1"))
self.output_projection_layer = (
- tf.keras.layers.Dense(
+ tf_keras.layers.Dense(
units=self.d_model,
kernel_initializer=self.kernel_initializer,
name="layer_2"))
- self.output_dropout = tf.keras.layers.Dropout(
+ self.output_dropout = tf_keras.layers.Dropout(
rate=self.dropout, name="drop_2")
self.output_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="LayerNorm", axis=-1, epsilon=1e-12))
super(PositionwiseFF, self).build(unused_input_shapes)
@@ -234,7 +234,7 @@ def call(self, inp):
return output
-class EmbeddingLookup(tf.keras.layers.Layer):
+class EmbeddingLookup(tf_keras.layers.Layer):
"""Looks up words embeddings for id tensor."""
def __init__(self, n_token, d_embed, initializer, **kwargs):
@@ -257,7 +257,7 @@ def call(self, inputs):
return tf.nn.embedding_lookup(self.lookup_table, inputs)
-class RelativeMultiheadAttention(tf.keras.layers.Layer):
+class RelativeMultiheadAttention(tf_keras.layers.Layer):
"""Multi-head attention with relative embedding."""
def __init__(self, d_model, n_head, d_head, dropout, dropout_att,
@@ -274,7 +274,7 @@ def build(self, unused_input_shapes):
"""Implements build() for the layer."""
self.scale = 1.0 / (self.d_head**0.5)
- self.output_layer_norm = tf.keras.layers.LayerNormalization(
+ self.output_layer_norm = tf_keras.layers.LayerNormalization(
name="LayerNorm", axis=-1, epsilon=1e-12)
self.kh_projection_layer = self.add_weight(
@@ -302,7 +302,7 @@ def build(self, unused_input_shapes):
shape=[self.d_model, self.n_head, self.d_head],
initializer=self.initializer)
- self.attention_dropout = tf.keras.layers.Dropout(rate=self.dropout)
+ self.attention_dropout = tf_keras.layers.Dropout(rate=self.dropout)
super(RelativeMultiheadAttention, self).build(unused_input_shapes)
@@ -360,7 +360,7 @@ def call(self, h, g, r, r_w_bias, r_r_bias, seg_mat, r_s_bias, seg_embed,
return (output_h, output_g)
-class TransformerXLModel(tf.keras.layers.Layer):
+class TransformerXLModel(tf_keras.layers.Layer):
"""Defines a Transformer-XL computation graph with additional support for XLNet."""
def __init__(self,
@@ -452,8 +452,8 @@ def build(self, unused_input_shapes):
dtype=self.tf_float,
name="word_embedding")
- self.h_dropout = tf.keras.layers.Dropout(rate=self.dropout)
- self.g_dropout = tf.keras.layers.Dropout(rate=self.dropout)
+ self.h_dropout = tf_keras.layers.Dropout(rate=self.dropout)
+ self.g_dropout = tf_keras.layers.Dropout(rate=self.dropout)
if self.untie_r:
self.r_w_bias = (
@@ -501,7 +501,7 @@ def build(self, unused_input_shapes):
self.mask_emb = self.add_weight(
"mask_emb/mask_emb", shape=[1, 1, self.d_model], dtype=self.tf_float)
- self.emb_dropout = tf.keras.layers.Dropout(rate=self.dropout)
+ self.emb_dropout = tf_keras.layers.Dropout(rate=self.dropout)
self.fwd_position_embedding = RelativePositionEncoding(self.d_model)
self.bwd_position_embedding = RelativePositionEncoding(self.d_model)
@@ -526,7 +526,7 @@ def build(self, unused_input_shapes):
activation_type=self.ff_activation,
name="layer_%d/ff" % (i)))
- self.output_dropout = tf.keras.layers.Dropout(rate=self.dropout)
+ self.output_dropout = tf_keras.layers.Dropout(rate=self.dropout)
super(TransformerXLModel, self).build(unused_input_shapes)
@@ -741,7 +741,7 @@ def call(self, inputs):
return output, new_mems, None
-class PretrainingXLNetModel(tf.keras.Model):
+class PretrainingXLNetModel(tf_keras.Model):
"""XLNet keras model combined with pretraining LM loss layer.
See the original paper: https://arxiv.org/pdf/1906.08237.pdf
@@ -826,7 +826,7 @@ def call(self, features):
return self.new_mems, model_output
-class ClassificationXLNetModel(tf.keras.Model):
+class ClassificationXLNetModel(tf_keras.Model):
"""XLNet keras model combined with classification loss layer.
See the original paper: https://arxiv.org/pdf/1906.08237.pdf
@@ -901,11 +901,11 @@ def call(self, features):
summary = self.summarization_layer(attention_output)
per_example_loss, logits = self.cl_loss_layer(hidden=summary, labels=label)
- self.add_loss(tf.keras.backend.mean(per_example_loss))
+ self.add_loss(tf_keras.backend.mean(per_example_loss))
return new_mems, logits
-class LMLossLayer(tf.keras.layers.Layer):
+class LMLossLayer(tf_keras.layers.Layer):
"""Layer computing cross entropy loss for language modeling."""
def __init__(self,
@@ -945,12 +945,12 @@ def __init__(self,
def build(self, unused_input_shapes):
"""Implements build() for the layer."""
if self.use_proj:
- self.proj_layer = tf.keras.layers.Dense(
+ self.proj_layer = tf_keras.layers.Dense(
units=self.hidden_size,
kernel_initializer=self.initializer,
activation=gelu,
name="lm_projection/dense")
- self.proj_layer_norm = tf.keras.layers.LayerNormalization(
+ self.proj_layer_norm = tf_keras.layers.LayerNormalization(
axis=-1, epsilon=1e-12, name="lm_projection/LayerNorm")
if not self.tie_weight:
self.softmax_w = self.add_weight(
@@ -984,7 +984,7 @@ def call(self, hidden, target, lookup_table, target_mask):
return total_loss, logits
-class Summarization(tf.keras.layers.Layer):
+class Summarization(tf_keras.layers.Layer):
"""The layer to pool the output from XLNet model into a vector."""
def __init__(self,
@@ -1024,12 +1024,12 @@ def __init__(self,
def build(self, unused_input_shapes):
"""Implements build() for the layer."""
if self.use_proj:
- self.proj_layer = tf.keras.layers.Dense(
+ self.proj_layer = tf_keras.layers.Dense(
units=self.hidden_size,
kernel_initializer=self.initializer,
activation=tf.nn.tanh,
name="summary")
- self.dropout_layer = tf.keras.layers.Dropout(rate=self.dropout_rate)
+ self.dropout_layer = tf_keras.layers.Dropout(rate=self.dropout_rate)
super(Summarization, self).build(unused_input_shapes)
@@ -1047,7 +1047,7 @@ def call(self, inputs):
return summary
-class ClassificationLossLayer(tf.keras.layers.Layer):
+class ClassificationLossLayer(tf_keras.layers.Layer):
"""Layer computing cross entropy loss for classification task."""
def __init__(self, n_class, initializer, **kwargs):
@@ -1065,7 +1065,7 @@ def __init__(self, n_class, initializer, **kwargs):
def build(self, unused_input_shapes):
"""Implements build() for the layer."""
- self.proj_layer = tf.keras.layers.Dense(
+ self.proj_layer = tf_keras.layers.Dense(
units=self.n_class, kernel_initializer=self.initializer, name="logit")
super(ClassificationLossLayer, self).build(unused_input_shapes)
@@ -1080,7 +1080,7 @@ def call(self, hidden, labels):
return loss, logits
-class QAXLNetModel(tf.keras.Model):
+class QAXLNetModel(tf_keras.Model):
"""XLNet keras model combined with question answering loss layer.
See the original paper: https://arxiv.org/pdf/1906.08237.pdf
@@ -1161,7 +1161,7 @@ def call(self, features, training=False):
return results
-class QALossLayer(tf.keras.layers.Layer):
+class QALossLayer(tf_keras.layers.Layer):
"""Layer computing position and regression loss for question answering task."""
def __init__(self, hidden_size, start_n_top, end_n_top, initializer,
@@ -1185,28 +1185,28 @@ def __init__(self, hidden_size, start_n_top, end_n_top, initializer,
def build(self, unused_input_shapes):
"""Implements build() for the layer."""
- self.start_logits_proj_layer = tf.keras.layers.Dense(
+ self.start_logits_proj_layer = tf_keras.layers.Dense(
units=1, kernel_initializer=self.initializer, name="start_logits/dense")
- self.end_logits_proj_layer0 = tf.keras.layers.Dense(
+ self.end_logits_proj_layer0 = tf_keras.layers.Dense(
units=self.hidden_size,
kernel_initializer=self.initializer,
activation=tf.nn.tanh,
name="end_logits/dense_0")
- self.end_logits_proj_layer1 = tf.keras.layers.Dense(
+ self.end_logits_proj_layer1 = tf_keras.layers.Dense(
units=1, kernel_initializer=self.initializer, name="end_logits/dense_1")
- self.end_logits_layer_norm = tf.keras.layers.LayerNormalization(
+ self.end_logits_layer_norm = tf_keras.layers.LayerNormalization(
axis=-1, epsilon=1e-12, name="end_logits/LayerNorm")
- self.answer_class_proj_layer0 = tf.keras.layers.Dense(
+ self.answer_class_proj_layer0 = tf_keras.layers.Dense(
units=self.hidden_size,
kernel_initializer=self.initializer,
activation=tf.nn.tanh,
name="answer_class/dense_0")
- self.answer_class_proj_layer1 = tf.keras.layers.Dense(
+ self.answer_class_proj_layer1 = tf_keras.layers.Dense(
units=1,
kernel_initializer=self.initializer,
use_bias=False,
name="answer_class/dense_1")
- self.ans_feature_dropout = tf.keras.layers.Dropout(rate=self.dropout_rate)
+ self.ans_feature_dropout = tf_keras.layers.Dropout(rate=self.dropout_rate)
super(QALossLayer, self).build(unused_input_shapes)
def __call__(self, hidden, p_mask, cls_index, **kwargs):
diff --git a/official/modeling/__init__.py b/official/modeling/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/modeling/__init__.py
+++ b/official/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/activations/__init__.py b/official/modeling/activations/__init__.py
index af39ee619c8..fffc2c2ffac 100644
--- a/official/modeling/activations/__init__.py
+++ b/official/modeling/activations/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,8 +14,10 @@
"""Activations package definition."""
from official.modeling.activations.gelu import gelu
+from official.modeling.activations.mish import mish
from official.modeling.activations.relu import relu6
from official.modeling.activations.sigmoid import hard_sigmoid
+from official.modeling.activations.squared_relu import squared_relu
from official.modeling.activations.swish import hard_swish
from official.modeling.activations.swish import identity
from official.modeling.activations.swish import simple_swish
diff --git a/official/modeling/activations/gelu.py b/official/modeling/activations/gelu.py
index 1ca79ebb662..60f17ade9ae 100644
--- a/official/modeling/activations/gelu.py
+++ b/official/modeling/activations/gelu.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Gaussian error linear unit."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
def gelu(x):
"""Gaussian Error Linear Unit.
@@ -29,4 +29,4 @@ def gelu(x):
Returns:
`x` with the GELU activation applied.
"""
- return tf.keras.activations.gelu(x, approximate=True)
+ return tf_keras.activations.gelu(x, approximate=True)
diff --git a/official/modeling/activations/gelu_test.py b/official/modeling/activations/gelu_test.py
index 727a714e38b..7b373edb3f5 100644
--- a/official/modeling/activations/gelu_test.py
+++ b/official/modeling/activations/gelu_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,14 +14,12 @@
"""Tests for the Gaussian error linear unit."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.modeling import activations
-@keras_parameterized.run_all_keras_modes
-class GeluTest(keras_parameterized.TestCase):
+class GeluTest(tf.test.TestCase):
def test_gelu(self):
expected_data = [[0.14967535, 0., -0.10032465],
diff --git a/official/modeling/activations/mish.py b/official/modeling/activations/mish.py
new file mode 100644
index 00000000000..de6eb29ac96
--- /dev/null
+++ b/official/modeling/activations/mish.py
@@ -0,0 +1,36 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Self Regularized Non-Monotonic Activation Function."""
+
+import tensorflow as tf, tf_keras
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+def mish(x) -> tf.Tensor:
+ """Mish activation function.
+
+ Mish: A Self Regularized Non-Monotonic Activation Function
+ https://arxiv.org/pdf/1908.08681.pdf
+
+ Mish(x) = x * tanh(ln(1+e^x))
+
+ Args:
+ x: A `Tensor` representing preactivation values.
+
+ Returns:
+ The activation value.
+ """
+ x = tf.convert_to_tensor(x)
+ return x * tf.tanh(tf.nn.softplus(x))
diff --git a/official/modeling/activations/mish_test.py b/official/modeling/activations/mish_test.py
new file mode 100644
index 00000000000..f4f713f1d70
--- /dev/null
+++ b/official/modeling/activations/mish_test.py
@@ -0,0 +1,30 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for the customized Mish activation."""
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import activations
+
+
+class MishTest(tf.test.TestCase):
+
+ def test_mish(self):
+ x = tf.constant([1.0, 0.0])
+ self.assertAllClose([0.86509839, 0.0], activations.mish(x))
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/modeling/activations/relu.py b/official/modeling/activations/relu.py
index 410be29d266..ce16af14b8a 100644
--- a/official/modeling/activations/relu.py
+++ b/official/modeling/activations/relu.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Customized Relu activation."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
def relu6(features):
"""Computes the Relu6 activation function.
diff --git a/official/modeling/activations/relu_test.py b/official/modeling/activations/relu_test.py
index 45a8339e2a2..9b0592ee877 100644
--- a/official/modeling/activations/relu_test.py
+++ b/official/modeling/activations/relu_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,15 +14,12 @@
"""Tests for the customized Relu activation."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import \
- keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.modeling import activations
-@keras_parameterized.run_all_keras_modes
-class CustomizedReluTest(keras_parameterized.TestCase):
+class CustomizedReluTest(tf.test.TestCase):
def test_relu6(self):
features = [[.25, 0, -.25], [-1, -2, 3]]
diff --git a/official/modeling/activations/sigmoid.py b/official/modeling/activations/sigmoid.py
index a3fc77fa5ea..f60d8d0d0f8 100644
--- a/official/modeling/activations/sigmoid.py
+++ b/official/modeling/activations/sigmoid.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Customized Sigmoid activation."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
def hard_sigmoid(features):
"""Computes the hard sigmoid activation function.
diff --git a/official/modeling/activations/sigmoid_test.py b/official/modeling/activations/sigmoid_test.py
index e5a1a61f97f..788a0fad7fc 100644
--- a/official/modeling/activations/sigmoid_test.py
+++ b/official/modeling/activations/sigmoid_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,15 +15,12 @@
"""Tests for the customized Sigmoid activation."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import \
- keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.modeling import activations
-@keras_parameterized.run_all_keras_modes
-class CustomizedSigmoidTest(keras_parameterized.TestCase):
+class CustomizedSigmoidTest(tf.test.TestCase):
def _hard_sigmoid_nn(self, x):
x = np.float32(x)
diff --git a/official/modeling/activations/squared_relu.py b/official/modeling/activations/squared_relu.py
new file mode 100644
index 00000000000..a619f5885df
--- /dev/null
+++ b/official/modeling/activations/squared_relu.py
@@ -0,0 +1,31 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Customized Squared ReLU activation."""
+
+import tensorflow as tf, tf_keras
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+def squared_relu(features: tf.Tensor) -> tf.Tensor:
+ """Computes the Squared ReLU activation function.
+
+ Args:
+ features: A `Tensor` representing preactivation values.
+
+ Returns:
+ The activation value.
+ """
+ features_tensor = tf.convert_to_tensor(features)
+ return tf.math.square(tf.nn.relu(features_tensor))
diff --git a/official/modeling/activations/squared_relu_test.py b/official/modeling/activations/squared_relu_test.py
new file mode 100644
index 00000000000..599d9e1068d
--- /dev/null
+++ b/official/modeling/activations/squared_relu_test.py
@@ -0,0 +1,37 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for the customized Squared ReLU activation."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.modeling import activations
+
+
+class CustomizedSquaredReluTest(tf.test.TestCase):
+
+ def _squared_relu_nn(self, x):
+ x = np.float32(x)
+ return tf.math.square(tf.nn.relu(x))
+
+ def test_squared_relu(self):
+ features = [[0.25, 0, -0.25], [-1, -2, 3]]
+ customized_squared_relu_data = activations.squared_relu(features)
+ squared_relu_data = self._squared_relu_nn(features)
+ self.assertAllClose(customized_squared_relu_data, squared_relu_data)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/modeling/activations/swish.py b/official/modeling/activations/swish.py
index 3d9372370ce..610173974d2 100644
--- a/official/modeling/activations/swish.py
+++ b/official/modeling/activations/swish.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Customized Swish activation."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
def simple_swish(features):
"""Computes the Swish activation function.
@@ -38,7 +38,7 @@ def simple_swish(features):
return features * tf.nn.sigmoid(features)
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
def hard_swish(features):
"""Computes a hard version of the swish function.
@@ -56,7 +56,7 @@ def hard_swish(features):
return features * tf.nn.relu6(features + tf.cast(3., fdtype)) * (1. / 6.)
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
def identity(features):
"""Computes the identity function.
diff --git a/official/modeling/activations/swish_test.py b/official/modeling/activations/swish_test.py
index 1eb5fa2a94f..0c8c9c76a66 100644
--- a/official/modeling/activations/swish_test.py
+++ b/official/modeling/activations/swish_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,14 +14,12 @@
"""Tests for the customized Swish activation."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.modeling import activations
-@keras_parameterized.run_all_keras_modes
-class CustomizedSwishTest(keras_parameterized.TestCase):
+class CustomizedSwishTest(tf.test.TestCase):
def _hard_swish_np(self, x):
x = np.float32(x)
diff --git a/official/modeling/fast_training/experimental/tf2_utils_2x_wide.py b/official/modeling/fast_training/experimental/tf2_utils_2x_wide.py
index af0760277ae..0f1e02ddad8 100644
--- a/official/modeling/fast_training/experimental/tf2_utils_2x_wide.py
+++ b/official/modeling/fast_training/experimental/tf2_utils_2x_wide.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def expand_vector(v: np.ndarray) -> np.ndarray:
@@ -36,13 +36,14 @@ def expand_vector(v: np.ndarray) -> np.ndarray:
def expand_1_axis(w: np.ndarray,
epsilon: float,
axis: int) -> np.ndarray:
- """Expands either the first dimension or the last dimension of w.
+ """Expands either the first or last dimension of w.
- If `axis = 0`, the following constraint will be satisfied:
+ If `axis = 0`, the following expression will be satisfied:
matmul(x, w) ==
- matmul(expand_vector(x), expand_1_axis(w, epsilon=0.1, axis=0))
+ matmul(expand_vector(x), expand_1_axis(w, axis=0))
- If `axis = -1`, the following constraint will be satisfied if `epsilon = 0.0`:
+ If `axis = -1` and `epsilon = 0.0`, the following constraint will be
+ satisfied:
expand_vector(matmul(x, w)) ==
2 * matmul(x, expand_1_axis(w, epsilon=0.0, axis=-1))
@@ -54,9 +55,12 @@ def expand_1_axis(w: np.ndarray,
Returns:
Expanded numpy array.
"""
- assert axis in (0, -1), (
- "Only support expanding the first or the last dimension. "
- "Got: {}".format(axis))
+
+ if axis not in (0, -1):
+ raise ValueError(
+ "Only support expanding the first or the last dimension. "
+ "Got: {}".format(axis)
+ )
rank = len(w.shape)
@@ -65,7 +69,7 @@ def expand_1_axis(w: np.ndarray,
sign_flip = np.array([1, -1])
for _ in range(rank - 1):
- sign_flip = np.expand_dims(sign_flip, axis=-1 if axis == 0 else 0)
+ sign_flip = np.expand_dims(sign_flip, axis=axis - 1)
sign_flip = np.tile(sign_flip,
[w.shape[0]] + [1] * (rank - 2) + [w.shape[-1]])
@@ -76,9 +80,9 @@ def expand_1_axis(w: np.ndarray,
def expand_2_axes(w: np.ndarray,
epsilon: float) -> np.ndarray:
- """Expands the first dimension and the last dimension of w.
+ """Expands the first and last dimension of w.
- The following constraint will be satisfied:
+ This operation satisfies the following expression:
expand_vector(matmul(x, w)) == matmul(expand_vector(x), expand_2_axes(w))
Args:
@@ -109,8 +113,8 @@ def var_to_var(var_from: tf.Variable,
epsilon: float):
"""Expands a variable to another variable.
- Assume the shape of `var_from` is (a, b, ..., y, z), the shape of `var_to`
- can be (a, ..., z * 2), (a * 2, ..., z * 2), (a * 2, ..., z)
+ Assuming the shape of `var_from` is (a, b, ..., y, z), then shape of `var_to`
+ must be one of (a, ..., z * 2), (a * 2, ..., z * 2), or (a * 2, ..., z).
If the shape of `var_to` is (a, ..., 2 * z):
For any x, tf.matmul(x, var_to) ~= expand_vector(tf.matmul(x, var_from)) / 2
@@ -131,21 +135,30 @@ def var_to_var(var_from: tf.Variable,
if shape_from == shape_to:
var_to.assign(var_from)
+ return
+
+ var_from_np = var_from.numpy()
+
+ if len(shape_from) == len(shape_to) == 1:
+ var_to.assign(expand_vector(var_from_np))
+ return
- elif len(shape_from) == 1 and len(shape_to) == 1:
- var_to.assign(expand_vector(var_from.numpy()))
+ a_from, z_from = shape_from[0], shape_from[-1]
+ a_to, z_to = shape_to[0], shape_to[-1]
- elif shape_from[0] * 2 == shape_to[0] and shape_from[-1] == shape_to[-1]:
- var_to.assign(expand_1_axis(var_from.numpy(), epsilon=epsilon, axis=0))
+ if a_to == 2 * a_from and z_to == z_from:
+ var_to.assign(expand_1_axis(var_from_np, epsilon=epsilon, axis=0))
+ return
- elif shape_from[0] == shape_to[0] and shape_from[-1] * 2 == shape_to[-1]:
- var_to.assign(expand_1_axis(var_from.numpy(), epsilon=epsilon, axis=-1))
+ if a_to == a_from and z_to == 2 * z_from:
+ var_to.assign(expand_1_axis(var_from_np, epsilon=epsilon, axis=-1))
+ return
- elif shape_from[0] * 2 == shape_to[0] and shape_from[-1] * 2 == shape_to[-1]:
- var_to.assign(expand_2_axes(var_from.numpy(), epsilon=epsilon))
+ if a_to == 2 * a_from and z_to == 2 * z_from:
+ var_to.assign(expand_2_axes(var_from_np, epsilon=epsilon))
+ return
- else:
- raise ValueError("Shape not supported, {}, {}".format(shape_from, shape_to))
+ raise ValueError("Shape not supported, {}, {}".format(shape_from, shape_to))
def model_to_model_2x_wide(model_from: tf.Module,
@@ -156,22 +169,23 @@ def model_to_model_2x_wide(model_from: tf.Module,
Also makes sure that the output of the model is not changed after expanding.
For example:
```
- model_narrow = tf.keras.Sequential()
- model_narrow.add(tf.keras.Input(shape=(3,)))
- model_narrow.add(tf.keras.layers.Dense(4))
- model_narrow.add(tf.keras.layers.Dense(1))
+ model_narrow = tf_keras.Sequential()
+ model_narrow.add(tf_keras.Input(shape=(3,)))
+ model_narrow.add(tf_keras.layers.Dense(4))
+ model_narrow.add(tf_keras.layers.Dense(1))
- model_wide = tf.keras.Sequential()
- model_wide.add(tf.keras.Input(shape=(6,)))
- model_wide.add(tf.keras.layers.Dense(8))
- model_wide.add(tf.keras.layers.Dense(1))
+ model_wide = tf_keras.Sequential()
+ model_wide.add(tf_keras.Input(shape=(6,)))
+ model_wide.add(tf_keras.layers.Dense(8))
+ model_wide.add(tf_keras.layers.Dense(1))
model_to_model_2x_wide(model_narrow, model_wide)
assert model_narrow([[1, 2, 3]]) == model_wide([[1, 1, 2, 2, 3, 3]])
```
- We assume that `model_from` and `model_to` has the same architecture and only
- widths of them differ.
+ We assume that `model_from` and `model_to` have the same architecture and
+ differ
+ only in widths.
Args:
model_from: input model to expand.
diff --git a/official/modeling/fast_training/experimental/tf2_utils_2x_wide_test.py b/official/modeling/fast_training/experimental/tf2_utils_2x_wide_test.py
index 2b95110b606..a0640769e71 100644
--- a/official/modeling/fast_training/experimental/tf2_utils_2x_wide_test.py
+++ b/official/modeling/fast_training/experimental/tf2_utils_2x_wide_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for tf2_utils_2x_wide."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.fast_training.experimental import tf2_utils_2x_wide
@@ -71,19 +71,44 @@ def test_expand_3d_tensor_axis_2(self):
o1 = np.matmul(x, w1)
self.assertAllClose(o0, np.sum(o1.reshape(2, 2), axis=-1))
+ def test_relations(self):
+ x = np.array([10, 11])
+ w = np.random.rand(2, 2)
+ # matmul(x, w) == matmul(expand_vector(x), expand_1_axis(w, axis=0))
+ lhs = np.matmul(x, w)
+ rhs = np.matmul(
+ tf2_utils_2x_wide.expand_vector(x),
+ tf2_utils_2x_wide.expand_1_axis(w, epsilon=0.1, axis=0),
+ )
+ self.assertAllClose(lhs, rhs)
+ # expand_vector(matmul(x, w)) ==
+ # 2 * matmul(x, expand_1_axis(w, epsilon=0.0, axis=-1))
+ lhs = tf2_utils_2x_wide.expand_vector(np.matmul(x, w))
+ rhs = 2 * np.matmul(
+ x, tf2_utils_2x_wide.expand_1_axis(w, epsilon=0.0, axis=-1)
+ )
+ self.assertAllClose(lhs, rhs)
+ # expand_vector(matmul(x, w)) == matmul(expand_vector(x), expand_2_axes(w))
+ lhs = tf2_utils_2x_wide.expand_vector(np.matmul(x, w))
+ rhs = np.matmul(
+ tf2_utils_2x_wide.expand_vector(x),
+ tf2_utils_2x_wide.expand_2_axes(w, epsilon=0.1),
+ )
+ self.assertAllClose(lhs, rhs)
+
def test_end_to_end(self):
"""Covers expand_vector, expand_2_axes, and expand_1_axis."""
- model_narrow = tf.keras.Sequential()
- model_narrow.add(tf.keras.Input(shape=(3,)))
- model_narrow.add(tf.keras.layers.Dense(4))
- model_narrow.add(tf.keras.layers.Dense(4))
- model_narrow.add(tf.keras.layers.Dense(1))
-
- model_wide = tf.keras.Sequential()
- model_wide.add(tf.keras.Input(shape=(6,)))
- model_wide.add(tf.keras.layers.Dense(8))
- model_wide.add(tf.keras.layers.Dense(8))
- model_wide.add(tf.keras.layers.Dense(1))
+ model_narrow = tf_keras.Sequential()
+ model_narrow.add(tf_keras.Input(shape=(3,)))
+ model_narrow.add(tf_keras.layers.Dense(4))
+ model_narrow.add(tf_keras.layers.Dense(4))
+ model_narrow.add(tf_keras.layers.Dense(1))
+
+ model_wide = tf_keras.Sequential()
+ model_wide.add(tf_keras.Input(shape=(6,)))
+ model_wide.add(tf_keras.layers.Dense(8))
+ model_wide.add(tf_keras.layers.Dense(8))
+ model_wide.add(tf_keras.layers.Dense(1))
x0 = np.array([[1, 2, 3]])
x1 = np.array([[1, 1, 2, 2, 3, 3]])
diff --git a/official/modeling/fast_training/progressive/policies.py b/official/modeling/fast_training/progressive/policies.py
index 52c3e73b486..e5c58198318 100644
--- a/official/modeling/fast_training/progressive/policies.py
+++ b/official/modeling/fast_training/progressive/policies.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from typing import Any, Mapping
from absl import logging
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import streamz_counters
from official.modeling.fast_training.progressive import utils
@@ -96,12 +96,12 @@ def num_steps(self, stage_id: int) -> int:
@abc.abstractmethod
def get_model(self,
stage_id: int,
- old_model: tf.keras.Model = None) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ old_model: tf_keras.Model = None) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Return model for this stage. For initialization, `old_model` = None."""
pass
@abc.abstractmethod
- def get_optimizer(self, stage_id: int) -> tf.keras.optimizers.Optimizer:
+ def get_optimizer(self, stage_id: int) -> tf_keras.optimizers.Optimizer:
"""Return optimizer for this stage."""
pass
@@ -116,7 +116,7 @@ def get_eval_dataset(self, stage_id: int) -> tf.data.Dataset:
pass
@property
- def cur_model(self) -> tf.keras.Model:
+ def cur_model(self) -> tf_keras.Model:
return self._volatiles.model
@property
@@ -132,7 +132,7 @@ def cur_eval_dataset(self) -> tf.data.Dataset:
return self._cur_eval_dataset
@property
- def cur_optimizer(self) -> tf.keras.optimizers.Optimizer:
+ def cur_optimizer(self) -> tf_keras.optimizers.Optimizer:
return self._volatiles.optimizer
@property
@@ -174,5 +174,5 @@ def update_pt_stage(self, global_step: int, pass_old_model=True) -> None:
new_optimizer = self.get_optimizer(new_stage_id)
self._volatiles.reassign_trackable(optimizer=new_optimizer)
new_model = self.get_model(
- new_stage_id, old_model=self.cur_model if pass_old_model else None)
+ new_stage_id, old_model=self.cur_model if pass_old_model else None) # pyrefly: ignore[bad-argument-type]
self._volatiles.reassign_trackable(model=new_model)
diff --git a/official/modeling/fast_training/progressive/train.py b/official/modeling/fast_training/progressive/train.py
index 612a485c6b4..134dd2852ce 100644
--- a/official/modeling/fast_training/progressive/train.py
+++ b/official/modeling/fast_training/progressive/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/fast_training/progressive/train_lib.py b/official/modeling/fast_training/progressive/train_lib.py
index 1fdb1d1c23c..a520cfb0944 100644
--- a/official/modeling/fast_training/progressive/train_lib.py
+++ b/official/modeling/fast_training/progressive/train_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,10 +22,9 @@
import os
from typing import Any, Mapping, Tuple
-# Import libraries
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions
from official.core import train_lib as base_train_lib
@@ -39,7 +38,7 @@ def run_experiment(distribution_strategy: tf.distribute.Strategy,
model_dir: str,
run_post_eval: bool = False,
save_summary: bool = True) \
--> Tuple[tf.keras.Model, Mapping[str, Any]]:
+-> Tuple[tf_keras.Model, Mapping[str, Any]]:
"""Runs train/eval configured by the experiment params.
Args:
@@ -55,7 +54,7 @@ def run_experiment(distribution_strategy: tf.distribute.Strategy,
Returns:
A 2-tuple of (model, eval_logs).
- model: `tf.keras.Model` instance.
+ model: `tf_keras.Model` instance.
eval_logs: returns eval metrics logs when run_post_eval is set to True,
otherwise, returns {}.
"""
diff --git a/official/modeling/fast_training/progressive/train_lib_test.py b/official/modeling/fast_training/progressive/train_lib_test.py
index fdc35b2e823..597f7004bab 100644
--- a/official/modeling/fast_training/progressive/train_lib_test.py
+++ b/official/modeling/fast_training/progressive/train_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl.testing import parameterized
import dataclasses
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
diff --git a/official/modeling/fast_training/progressive/trainer.py b/official/modeling/fast_training/progressive/trainer.py
index af24af52787..f29ff1045ef 100644
--- a/official/modeling/fast_training/progressive/trainer.py
+++ b/official/modeling/fast_training/progressive/trainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,11 +23,10 @@
import os
from typing import Any, Optional
-# Import libraries
from absl import logging
import gin
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import base_trainer as trainer_lib
from official.core import config_definitions
@@ -116,8 +115,8 @@ def __init__(
global_step=self.global_step,
**self._task.cur_checkpoint_items)
- self._train_loss = tf.keras.metrics.Mean('training_loss', dtype=tf.float32)
- self._validation_loss = tf.keras.metrics.Mean(
+ self._train_loss = tf_keras.metrics.Mean('training_loss', dtype=tf.float32)
+ self._validation_loss = tf_keras.metrics.Mean(
'validation_loss', dtype=tf.float32)
self._train_metrics = self.task.build_metrics(
training=True) + self.model.metrics
diff --git a/official/modeling/fast_training/progressive/trainer_test.py b/official/modeling/fast_training/progressive/trainer_test.py
index 47303551282..fa611f9176a 100644
--- a/official/modeling/fast_training/progressive/trainer_test.py
+++ b/official/modeling/fast_training/progressive/trainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl.testing import parameterized
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -67,11 +67,11 @@ def num_steps(self, stage_id: int) -> int:
def get_model(self,
stage_id: int,
- old_model: tf.keras.Model) -> tf.keras.Model:
+ old_model: tf_keras.Model) -> tf_keras.Model:
del stage_id, old_model
return self.build_model()
- def get_optimizer(self, stage_id: int) -> tf.keras.optimizers.Optimizer:
+ def get_optimizer(self, stage_id: int) -> tf_keras.optimizers.Optimizer:
optimizer_type = 'sgd' if stage_id == 0 else 'adamw'
optimizer_config = cfg.OptimizationConfig({
'optimizer': {'type': optimizer_type},
@@ -226,9 +226,13 @@ def test_configure_optimizer(self, mixed_precision_dtype, loss_scale):
task = TestPolicy(None, config.task)
trainer = trainer_lib.ProgressiveTrainer(config, task, self.get_temp_dir())
if mixed_precision_dtype != 'float16':
- self.assertIsInstance(trainer.optimizer, tf.keras.optimizers.SGD)
+ self.assertIsInstance(
+ trainer.optimizer,
+ (tf_keras.optimizers.SGD, tf_keras.optimizers.legacy.SGD))
elif mixed_precision_dtype == 'float16' and loss_scale is None:
- self.assertIsInstance(trainer.optimizer, tf.keras.optimizers.SGD)
+ self.assertIsInstance(
+ trainer.optimizer,
+ (tf_keras.optimizers.SGD, tf_keras.optimizers.legacy.SGD))
metrics = trainer.train(tf.convert_to_tensor(5, dtype=tf.int32))
self.assertIn('training_loss', metrics)
diff --git a/official/modeling/fast_training/progressive/utils.py b/official/modeling/fast_training/progressive/utils.py
index 73418d0ab8f..970cd37d1bf 100644
--- a/official/modeling/fast_training/progressive/utils.py
+++ b/official/modeling/fast_training/progressive/utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,13 @@
"""Util classes and functions."""
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=g-direct-tensorflow-import
-from tensorflow.python.training.tracking import tracking
+from tensorflow.python.trackable import autotrackable
-class VolatileTrackable(tracking.AutoTrackable):
+class VolatileTrackable(autotrackable.AutoTrackable):
"""A util class to keep Trackables that might change instances."""
def __init__(self, **kwargs):
diff --git a/official/modeling/grad_utils.py b/official/modeling/grad_utils.py
index 22479e6ff3b..85f708206f7 100644
--- a/official/modeling/grad_utils.py
+++ b/official/modeling/grad_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _filter_grads(grads_and_vars):
@@ -98,7 +98,7 @@ def minimize_using_explicit_allreduce(tape,
Args:
tape: An instance of `tf.GradientTape`.
- optimizer: An instance of `tf.keras.optimizers.Optimizer`.
+ optimizer: An instance of `tf_keras.optimizers.Optimizer`.
loss: the loss tensor.
trainable_variables: A list of model Variables.
pre_allreduce_callbacks: A list of callback functions that takes gradients
@@ -117,7 +117,7 @@ def minimize_using_explicit_allreduce(tape,
in one pack.
"""
if isinstance(optimizer,
- tf.keras.mixed_precision.LossScaleOptimizer):
+ tf_keras.mixed_precision.LossScaleOptimizer):
# FP16 GPU code path
with tape:
scaled_loss = optimizer.get_scaled_loss(loss)
diff --git a/official/modeling/grad_utils_test.py b/official/modeling/grad_utils_test.py
index ded7794ab58..18006173f5c 100644
--- a/official/modeling/grad_utils_test.py
+++ b/official/modeling/grad_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for grad_utils."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import grad_utils
from official.modeling import performance
@@ -23,9 +23,9 @@ class GradUtilsTest(tf.test.TestCase):
def test_minimize(self):
- optimizer = tf.keras.optimizers.SGD(0.1)
+ optimizer = tf_keras.optimizers.SGD(0.1)
with tf.GradientTape() as tape:
- model = tf.keras.layers.Dense(2)
+ model = tf_keras.layers.Dense(2)
outputs = model(tf.zeros((2, 2), tf.float32))
loss = tf.reduce_mean(outputs)
@@ -35,10 +35,10 @@ def test_minimize(self):
def test_minimize_fp16(self):
optimizer = performance.configure_optimizer(
- tf.keras.optimizers.SGD(0.1), use_float16=True)
+ tf_keras.optimizers.SGD(0.1), use_float16=True)
performance.set_mixed_precision_policy(tf.float16)
with tf.GradientTape() as tape:
- model = tf.keras.layers.Dense(2)
+ model = tf_keras.layers.Dense(2)
outputs = model(tf.zeros((2, 2), tf.float16))
loss = tf.reduce_mean(outputs)
@@ -51,11 +51,11 @@ def _clip_by_global_norm(grads_and_vars):
(grads, _) = tf.clip_by_global_norm(grads, clip_norm=1.0)
return zip(grads, tvars)
with tf.GradientTape() as tape:
- model = tf.keras.layers.Dense(2)
+ model = tf_keras.layers.Dense(2)
outputs = model(tf.zeros((2, 2), tf.float16))
loss = tf.reduce_mean(outputs)
optimizer = performance.configure_optimizer(
- tf.keras.optimizers.SGD(0.1), use_float16=True, loss_scale=128)
+ tf_keras.optimizers.SGD(0.1), use_float16=True, loss_scale=128)
grad_utils.minimize_using_explicit_allreduce(
tape,
optimizer,
diff --git a/official/modeling/hyperparams/__init__.py b/official/modeling/hyperparams/__init__.py
index 5503ad8e478..3a14dc44f09 100644
--- a/official/modeling/hyperparams/__init__.py
+++ b/official/modeling/hyperparams/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/hyperparams/base_config.py b/official/modeling/hyperparams/base_config.py
index f68b16b3645..fbb6d77604f 100644
--- a/official/modeling/hyperparams/base_config.py
+++ b/official/modeling/hyperparams/base_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,18 +13,22 @@
# limitations under the License.
"""Base configurations to standardize experiments."""
+
import copy
import dataclasses
import functools
import inspect
-from typing import Any, List, Mapping, Optional, Type
+import types
+import typing
+from typing import Any, List, Mapping, Optional, Type, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import yaml
from official.modeling.hyperparams import params_dict
+
_BOUND = set()
@@ -54,6 +58,16 @@ def _wrapper(self, *args, **kwargs): # pylint: disable=unused-argument
return decorator
+def _is_optional(field):
+ # Two styles of annotating optional fields:
+ # Optional[T] -> typing.Union
+ # T | None -> types.UnionType
+ is_union = typing.get_origin(field) in (Union, types.UnionType)
+ # An optional field is a union of a type, and NoneType.
+ args = typing.get_args(field)
+ return is_union and len(args) == 2 and type(None) in args
+
+
@dataclasses.dataclass
class Config(params_dict.ParamsDict):
"""The base configuration class that supports YAML/JSON based overrides.
@@ -87,6 +101,20 @@ def __post_init__(self, default_params, restrictions):
def BUILDER(self):
return self._BUILDER
+ @classmethod
+ def _get_annotations(cls):
+ """Returns valid annotations.
+
+ Note: this is similar to dataclasses.__annotations__ except it also includes
+ annotations from its parent classes.
+ """
+ all_annotations = typing.get_type_hints(cls)
+ # Removes Config class annotation from the value, e.g., default_params,
+ # restrictions, etc.
+ for k in Config.__annotations__:
+ del all_annotations[k]
+ return all_annotations
+
@classmethod
def _isvalidsequence(cls, v):
"""Check if the input values are valid sequences.
@@ -148,11 +176,15 @@ def _export_config(cls, v):
raise TypeError('Unknown type: {!r}'.format(type(v)))
@classmethod
- def _get_subconfig_type(cls, k) -> Type[params_dict.ParamsDict]:
+ def _get_subconfig_type(
+ cls, k, subconfig_type=None
+ ) -> Type[params_dict.ParamsDict]:
"""Get element type by the field name.
Args:
k: the key/name of the field.
+ subconfig_type: default subconfig_type. If None, it is set to
+ Config.
Returns:
Config as default. If a type annotation is found for `k`,
@@ -160,22 +192,34 @@ def _get_subconfig_type(cls, k) -> Type[params_dict.ParamsDict]:
2) returns the element type if the annotation of `k` is List[SubType]
or Tuple[SubType].
"""
- subconfig_type = Config
- if k in cls.__annotations__:
+ if not subconfig_type:
+ subconfig_type = Config
+
+ def _is_subtype(x, target) -> bool:
+ return isinstance(x, type) and issubclass(x, target)
+
+ annotations = cls._get_annotations()
+ if k in annotations:
# Directly Config subtype.
- type_annotation = cls.__annotations__[k] # pytype: disable=invalid-annotation
- if (isinstance(type_annotation, type) and
- issubclass(type_annotation, Config)):
- subconfig_type = cls.__annotations__[k] # pytype: disable=invalid-annotation
- else:
- # Check if the field is a sequence of subtypes.
- field_type = getattr(type_annotation, '__origin__', type(None))
- if (isinstance(field_type, type) and
- issubclass(field_type, cls.SEQUENCE_TYPES)):
- element_type = getattr(type_annotation, '__args__', [type(None)])[0]
- subconfig_type = (
- element_type if issubclass(element_type, params_dict.ParamsDict)
- else subconfig_type)
+ type_annotation = annotations[k]
+ i = 0
+ # Loop for striping the Optional annotation.
+ traverse_in = True
+ while traverse_in:
+ i += 1
+ if _is_subtype(type_annotation, Config):
+ subconfig_type = type_annotation
+ break
+ else:
+ # If the field is a sequence of sub-config types or an Optional
+ # sub-config, then strip the container and process the sub-config.
+ is_sequence = _is_subtype(
+ typing.get_origin(type_annotation), cls.SEQUENCE_TYPES
+ )
+ if is_sequence or _is_optional(type_annotation):
+ type_annotation = typing.get_args(type_annotation)[0]
+ continue
+ traverse_in = False
return subconfig_type
def _set(self, k, v):
@@ -202,8 +246,11 @@ def is_null(k):
# If the key not exist or the value is None, a new Config-family object
# sould be created for the key.
self.__dict__[k] = subconfig_type(v)
- else:
+ elif hasattr(self.__dict__[k], 'override'):
self.__dict__[k].override(v)
+ else:
+ # The key exists but it cannot be overridden. For example, it's a str.
+ self.__dict__[k] = subconfig_type(v)
elif not is_null(k) and isinstance(v, self.SEQUENCE_TYPES) and all(
[not isinstance(e, self.IMMUTABLE_TYPES) for e in v]):
if len(self.__dict__[k]) == len(v):
@@ -256,9 +303,11 @@ def _override(self, override_dict, is_strict=True):
else:
self._set(k, v)
else:
- if isinstance(v, dict) and self.__dict__[k]:
+ if isinstance(v, dict) and hasattr(self.__dict__[k], '_override'):
self.__dict__[k]._override(v, is_strict) # pylint: disable=protected-access
- elif isinstance(v, params_dict.ParamsDict) and self.__dict__[k]:
+ elif isinstance(v, params_dict.ParamsDict) and hasattr(
+ self.__dict__[k], '_override'
+ ):
self.__dict__[k]._override(v.as_dict(), is_strict) # pylint: disable=protected-access
else:
self._set(k, v)
@@ -300,6 +349,9 @@ def from_json(cls, file_path: str):
@classmethod
def from_args(cls, *args, **kwargs):
"""Builds a config from the given list of arguments."""
+ # Note we intend to keep `__annotations__` instead of `_get_annotations`.
+ # Assuming a parent class of (a, b) with the sub-class of (c, d), the
+ # sub-class will take (c, d) for args, rather than starting from (a, b).
attributes = list(cls.__annotations__.keys())
default_params = {a: p for a, p in zip(attributes, args)}
default_params.update(kwargs)
diff --git a/official/modeling/hyperparams/base_config_test.py b/official/modeling/hyperparams/base_config_test.py
index b27352af895..eab97cb492b 100644
--- a/official/modeling/hyperparams/base_config_test.py
+++ b/official/modeling/hyperparams/base_config_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,11 +13,12 @@
# limitations under the License.
import pprint
-from typing import List, Tuple
+import dataclasses
+from typing import List, Optional, Tuple
from absl.testing import parameterized
-import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
from official.modeling.hyperparams import base_config
@@ -31,7 +32,8 @@ class DumpConfig1(base_config.Config):
class DumpConfig2(base_config.Config):
c: int = 2
d: str = 'text'
- e: DumpConfig1 = DumpConfig1()
+ e: DumpConfig1 = dataclasses.field(default_factory=DumpConfig1)
+ optional_e: Optional[DumpConfig1] = None
@dataclasses.dataclass
@@ -54,6 +56,17 @@ class DummyConfig5(base_config.Config):
z: Tuple[str] = ('a',)
+@dataclasses.dataclass
+class DumpConfig6(base_config.Config):
+ test_config1: Optional[DumpConfig1] = None
+
+
+@dataclasses.dataclass
+class ModernOptionalConfig(base_config.Config):
+ leaf: DumpConfig1 | None = None
+ leaves: tuple[DumpConfig1 | None, ...] = tuple()
+
+
class BaseConfigTest(parameterized.TestCase, tf.test.TestCase):
def assertHasSameTypes(self, c, d, msg=''):
@@ -321,6 +334,17 @@ def test_override(self):
self.assertEqual(type(params.a[1].c), base_config.Config)
self.assertEqual(pprint.pformat(params.a[1].c.d), '5')
+ def test_override_scalar_with_dict_replaces_the_whole_value(self):
+ params = base_config.Config({'a': 1})
+ params.override({'a': {'b': 2}}, is_strict=False)
+ self.assertEqual(type(params.a), base_config.Config)
+ self.assertEqual(params.a.b, 2)
+
+ def test_override_dict_with_scalar_replaces_the_whole_value(self):
+ params = base_config.Config({'a': {'b': 2}})
+ params.override({'a': 1}, is_strict=False)
+ self.assertEqual(params.a, 1)
+
@parameterized.parameters(
([{}],),
(({},),),
@@ -342,6 +366,34 @@ def test_ppformat(self):
]),
"['s', 1, 1.0, True, None, {}, [], (), {8: 9, (2,): (3, [4], {6: 7})}]")
+ def test_with_superclass_override(self):
+ config = DumpConfig2()
+ config.override({'optional_e': {'a': 2}})
+ self.assertEqual(
+ config.optional_e.as_dict(),
+ {
+ 'a': 2,
+ 'b': 'text',
+ },
+ )
+
+ # Previously, the following will fail. See b/274696969 for context.
+ config = DumpConfig3()
+ config.override({'optional_e': {'a': 2}})
+ self.assertEqual(
+ config.optional_e.as_dict(),
+ {
+ 'a': 2,
+ 'b': 'text',
+ },
+ )
+
+ def test_get_annotations_without_base_config_leak(self):
+ with self.assertRaisesRegex(
+ KeyError, "The key 'restrictions' does not exist"
+ ):
+ DumpConfig3().override({'restrictions': None})
+
def test_with_restrictions(self):
restrictions = ['e.a[a-zA-Z][\w\.]*) # variable name: "var" or "x"
+ (?P[a-zA-Z][\w\.]*)(?P\[?[0-9]*\]?) # variable name: "var" or "x" followed by optional index: "[0]" or "[23]"
\s*=\s*
((?P\'(.*?)\' # single quote
|
@@ -41,15 +42,15 @@
_CONST_VALUE_RE = re.compile(r'(\d.*|-\d.*|None)')
-# Yaml loader with an implicit resolver to parse float decimal and exponential
+# Yaml LOADER with an implicit resolver to parse float decimal and exponential
# format. The regular experission parse the following cases:
# 1- Decimal number with an optional exponential term.
# 2- Integer number with an exponential term.
# 3- Decimal number with an optional exponential term.
# 4- Decimal number.
-LOADER = yaml.SafeLoader
-LOADER.add_implicit_resolver(
+_LOADER = yaml.FullLoader
+_LOADER.add_implicit_resolver(
'tag:yaml.org,2002:float',
re.compile(r'''
^(?:[-+]?(?:[0-9][0-9_]*)\\.[0-9_]*(?:[eE][-+]?[0-9]+)?
@@ -192,6 +193,9 @@ def _override(self, override_dict, is_strict=True):
'To extend the existing keys, use '
'`override` with `is_strict` = False.'.format(k))
else:
+ logging.warning(
+ 'Adding new key `%s` to ParamsDict because is_strict=False.', k
+ )
self._set(k, v)
else:
if isinstance(v, dict):
@@ -223,7 +227,7 @@ def validate(self):
"""Validate the parameters consistency based on the restrictions.
This method validates the internal consistency using the pre-defined list of
- restrictions. A restriction is defined as a string which specfiies a binary
+ restrictions. A restriction is defined as a string which specifies a binary
operation. The supported binary operations are {'==', '!=', '<', '<=', '>',
'>='}. Note that the meaning of these operators are consistent with the
underlying Python immplementation. Users should make sure the define
@@ -288,42 +292,42 @@ def _get_kvs(tokens, params_dict):
_, left_v, _, right_v = _get_kvs(tokens, params_dict)
if left_v != right_v:
raise KeyError(
- 'Found inconsistncy between key `{}` and key `{}`.'.format(
+ 'Found inconsistency between key `{}` and key `{}`.'.format(
tokens[0], tokens[1]))
elif '!=' in restriction:
tokens = restriction.split('!=')
_, left_v, _, right_v = _get_kvs(tokens, params_dict)
if left_v == right_v:
raise KeyError(
- 'Found inconsistncy between key `{}` and key `{}`.'.format(
- tokens[0], tokens[1]))
- elif '<' in restriction:
- tokens = restriction.split('<')
- _, left_v, _, right_v = _get_kvs(tokens, params_dict)
- if left_v >= right_v:
- raise KeyError(
- 'Found inconsistncy between key `{}` and key `{}`.'.format(
+ 'Found inconsistency between key `{}` and key `{}`.'.format(
tokens[0], tokens[1]))
elif '<=' in restriction:
tokens = restriction.split('<=')
_, left_v, _, right_v = _get_kvs(tokens, params_dict)
if left_v > right_v:
raise KeyError(
- 'Found inconsistncy between key `{}` and key `{}`.'.format(
+ 'Found inconsistency between key `{}` and key `{}`.'.format(
tokens[0], tokens[1]))
- elif '>' in restriction:
- tokens = restriction.split('>')
+ elif '<' in restriction:
+ tokens = restriction.split('<')
_, left_v, _, right_v = _get_kvs(tokens, params_dict)
- if left_v <= right_v:
+ if left_v >= right_v:
raise KeyError(
- 'Found inconsistncy between key `{}` and key `{}`.'.format(
+ 'Found inconsistency between key `{}` and key `{}`.'.format(
tokens[0], tokens[1]))
elif '>=' in restriction:
tokens = restriction.split('>=')
_, left_v, _, right_v = _get_kvs(tokens, params_dict)
if left_v < right_v:
raise KeyError(
- 'Found inconsistncy between key `{}` and key `{}`.'.format(
+ 'Found inconsistency between key `{}` and key `{}`.'.format(
+ tokens[0], tokens[1]))
+ elif '>' in restriction:
+ tokens = restriction.split('>')
+ _, left_v, _, right_v = _get_kvs(tokens, params_dict)
+ if left_v <= right_v:
+ raise KeyError(
+ 'Found inconsistency between key `{}` and key `{}`.'.format(
tokens[0], tokens[1]))
else:
raise ValueError('Unsupported relation in restriction.')
@@ -332,7 +336,7 @@ def _get_kvs(tokens, params_dict):
def read_yaml_to_params_dict(file_path: str):
"""Reads a YAML file to a ParamsDict."""
with tf.io.gfile.GFile(file_path, 'r') as f:
- params_dict = yaml.load(f, Loader=LOADER)
+ params_dict = yaml.load(f, Loader=_LOADER)
return ParamsDict(params_dict)
@@ -385,6 +389,8 @@ def nested_csv_str_to_json_str(csv_str):
if not csv_str:
return ''
+ array_param_map = collections.defaultdict(str)
+ max_index_map = collections.defaultdict(str)
formatted_entries = []
nested_map = collections.defaultdict(list)
pos = 0
@@ -398,6 +404,27 @@ def nested_csv_str_to_json_str(csv_str):
m_dict = m.groupdict()
name = m_dict['name']
v = m_dict['val']
+ bracketed_index = m_dict['bracketed_index']
+ # If we reach the name of the array.
+ if bracketed_index and '.' not in name:
+ # Extract the array's index by removing '[' and ']'
+ index = int(bracketed_index[1:-1])
+ if '.' in v:
+ numeric_val = float(v)
+ else:
+ numeric_val = int(v)
+ # Add the value to the array.
+ if name not in array_param_map:
+ max_index_map[name] = index # pyrefly: ignore[unsupported-operation]
+ array_param_map[name] = [None] * (index + 1) # pyrefly: ignore[unsupported-operation]
+ array_param_map[name][index] = numeric_val # pyrefly: ignore[unsupported-operation]
+ elif index < max_index_map[name]: # pyrefly: ignore[unsupported-operation]
+ array_param_map[name][index] = numeric_val # pyrefly: ignore[unsupported-operation]
+ else:
+ array_param_map[name] += [None] * (index - max_index_map[name]) # pyrefly: ignore[unsupported-operation]
+ array_param_map[name][index] = numeric_val # pyrefly: ignore[unsupported-operation]
+ max_index_map[name] = index # pyrefly: ignore[unsupported-operation]
+ continue
# If a GCS path (e.g. gs://...) is provided, wrap this in quotes
# as yaml.load would otherwise throw an exception
@@ -407,7 +434,10 @@ def nested_csv_str_to_json_str(csv_str):
name_nested = name.split('.')
if len(name_nested) > 1:
grouping = name_nested[0]
- value = '.'.join(name_nested[1:]) + '=' + v
+ if bracketed_index:
+ value = '.'.join(name_nested[1:]) + bracketed_index + '=' + v
+ else:
+ value = '.'.join(name_nested[1:]) + '=' + v
nested_map[grouping].append(value)
else:
formatted_entries.append('%s : %s' % (name, v))
@@ -416,6 +446,13 @@ def nested_csv_str_to_json_str(csv_str):
value = ','.join(value)
value = nested_csv_str_to_json_str(value)
formatted_entries.append('%s : %s' % (grouping, value))
+
+ # Add array parameters and check that the array is fully initialized.
+ for name in array_param_map:
+ if any(v is None for v in array_param_map[name]):
+ raise ValueError('Did not pass all values of array: %s' % name)
+ formatted_entries.append('%s : %s' % (name, array_param_map[name]))
+
return '{' + ', '.join(formatted_entries) + '}'
@@ -453,12 +490,12 @@ def override_params_dict(params, dict_or_string_or_yaml_file, is_strict):
nested_csv_str_to_json_str(dict_or_string_or_yaml_file))
except ValueError:
pass
- params_dict = yaml.load(dict_or_string_or_yaml_file, Loader=LOADER)
+ params_dict = yaml.load(dict_or_string_or_yaml_file, Loader=_LOADER)
if isinstance(params_dict, dict):
params.override(params_dict, is_strict)
else:
with tf.io.gfile.GFile(dict_or_string_or_yaml_file) as f:
- params.override(yaml.load(f, Loader=yaml.FullLoader), is_strict)
+ params.override(yaml.load(f, Loader=_LOADER), is_strict)
else:
raise ValueError('Unknown input type to parse.')
return params
diff --git a/official/modeling/hyperparams/params_dict_test.py b/official/modeling/hyperparams/params_dict_test.py
index 145590a4c2c..80efb7b9316 100644
--- a/official/modeling/hyperparams/params_dict_test.py
+++ b/official/modeling/hyperparams/params_dict_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import os
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import yaml
from official.modeling.hyperparams import params_dict
@@ -167,6 +167,7 @@ def test_validate(self):
'a': 10
}
}, ['b == c'])
+ params.validate()
# Raise error due to inconsistency
with self.assertRaises(KeyError):
@@ -176,7 +177,7 @@ def test_validate(self):
# Valid rule.
params = params_dict.ParamsDict({'a': 1, 'c': {'a': 1}}, ['a == c.a'])
- # Overridding violates the existing rule, raise error upon validate.
+ # Overriding violates the existing rule, raise error upon validate.
params.override({'a': 11})
with self.assertRaises(KeyError):
params.validate()
@@ -198,6 +199,10 @@ def test_validate(self):
}, ['a == None', 'c.a == 1'])
params.validate()
+ # Valid restrictions with inequality.
+ params = params_dict.ParamsDict({'a': 1}, ['a >= 1'])
+ params.validate()
+
class ParamsDictIOTest(tf.test.TestCase):
@@ -220,7 +225,7 @@ def test_save_params_dict_to_yaml(self):
params_dict.save_params_dict_to_yaml(params, output_yaml_file)
with tf.io.gfile.GFile(output_yaml_file, 'r') as f:
- params_d = yaml.load(f)
+ params_d = yaml.load(f, Loader=yaml.Loader)
self.assertEqual(params.a, params_d['a'])
self.assertEqual(params.b, params_d['b'])
self.assertEqual(params.c.c1, params_d['c']['c1'])
@@ -364,7 +369,7 @@ def test_basic_csv_str_load(self):
csv_str = 'a=1,b=2,c=3'
expected_output = {'a': 1, 'b': 2, 'c': 3}
converted_csv_str = params_dict.nested_csv_str_to_json_str(csv_str)
- converted_dict = yaml.load(converted_csv_str)
+ converted_dict = yaml.load(converted_csv_str, Loader=yaml.Loader)
self.assertDictEqual(converted_dict, expected_output)
def test_basic_nested_csv_str_to_json_str(self):
@@ -377,7 +382,7 @@ def test_basic_nested_csv_str_load(self):
csv_str = 'a=1,b.b1=2,c.c1=3'
expected_output = {'a': 1, 'b': {'b1': 2}, 'c': {'c1': 3}}
converted_csv_str = params_dict.nested_csv_str_to_json_str(csv_str)
- converted_dict = yaml.load(converted_csv_str)
+ converted_dict = yaml.load(converted_csv_str, Loader=yaml.Loader)
self.assertDictEqual(converted_dict, expected_output)
def test_complex_nested_csv_str_to_json_str(self):
@@ -390,13 +395,30 @@ def test_complex_nested_csv_str_load(self):
csv_str = 'a.aa.aaa.aaaaa.a=1,a.a=2'
expected_output = {'a': {'aa': {'aaa': {'aaaaa': {'a': 1}}}, 'a': 2}}
converted_csv_str = params_dict.nested_csv_str_to_json_str(csv_str)
- converted_dict = yaml.load(converted_csv_str)
+ converted_dict = yaml.load(converted_csv_str, Loader=yaml.Loader)
self.assertDictEqual(converted_dict, expected_output)
+ def test_int_array_param_nested_csv_str_to_json_str(self):
+ csv_str = 'a.b[2]=3,a.b[0]=1,a.b[1]=2'
+ json_str = '{a : {b : [1, 2, 3]}}'
+ converted_csv_str = params_dict.nested_csv_str_to_json_str(csv_str)
+ self.assertEqual(converted_csv_str, json_str)
+
+ def test_float_array_param_nested_csv_str_to_json_str(self):
+ csv_str = 'a.b[1]=3.45,a.b[2]=1.32,a.b[0]=2.232'
+ json_str = '{a : {b : [2.232, 3.45, 1.32]}}'
+ converted_csv_str = params_dict.nested_csv_str_to_json_str(csv_str)
+ self.assertEqual(converted_csv_str, json_str)
+
+ def test_incomplete_array_param_nested_csv_str_to_json_str(self):
+ csv_str = 'a.b[0]=1,a.b[2]=2'
+ self.assertRaises(ValueError, params_dict.nested_csv_str_to_json_str,
+ csv_str)
+
def test_csv_str_load_supported_datatypes(self):
csv_str = 'a=1,b=2.,c=[1,2,3],d=\'hello, there\',e=\"Hi.\"'
converted_csv_str = params_dict.nested_csv_str_to_json_str(csv_str)
- converted_dict = yaml.load(converted_csv_str)
+ converted_dict = yaml.load(converted_csv_str, Loader=yaml.Loader)
self.assertEqual(converted_dict['a'], 1)
self.assertEqual(converted_dict['b'], 2.)
self.assertEqual(converted_dict['c'], [1, 2, 3])
diff --git a/official/modeling/multitask/__init__.py b/official/modeling/multitask/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/modeling/multitask/__init__.py
+++ b/official/modeling/multitask/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/multitask/base_model.py b/official/modeling/multitask/base_model.py
index 573bdc75399..93ce0f70092 100644
--- a/official/modeling/multitask/base_model.py
+++ b/official/modeling/multitask/base_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Abstraction of multi-task model."""
from typing import Text, Dict
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class MultiTaskBaseModel(tf.Module):
@@ -25,11 +25,11 @@ def __init__(self, **kwargs):
super().__init__(**kwargs)
self._sub_tasks = self._instantiate_sub_tasks()
- def _instantiate_sub_tasks(self) -> Dict[Text, tf.keras.Model]:
+ def _instantiate_sub_tasks(self) -> Dict[Text, tf_keras.Model]:
"""Abstract function that sets up the computation for each sub-task.
Returns:
- A map from task name (as string) to a tf.keras.Model object that
+ A map from task name (as string) to a tf_keras.Model object that
represents the sub-task in the multi-task pool.
"""
raise NotImplementedError(
@@ -37,9 +37,18 @@ def _instantiate_sub_tasks(self) -> Dict[Text, tf.keras.Model]:
@property
def sub_tasks(self):
- """Fetch a map of task name (string) to task model (tf.keras.Model)."""
+ """Fetch a map of task name (string) to task model (tf_keras.Model)."""
return self._sub_tasks
def initialize(self):
"""Optional function that loads a pre-train checkpoint."""
return
+
+ def build(self):
+ """Builds the networks for tasks to make sure variables are created."""
+ # Try to build all sub tasks.
+ for task_model in self._sub_tasks.values():
+ # Assumes all the tf.Module models are built because we don't have any
+ # way to check them.
+ if isinstance(task_model, tf_keras.Model) and not task_model.built:
+ _ = task_model(task_model.inputs)
diff --git a/official/modeling/multitask/base_trainer.py b/official/modeling/multitask/base_trainer.py
index e3bf18718ed..0a62f60bc06 100644
--- a/official/modeling/multitask/base_trainer.py
+++ b/official/modeling/multitask/base_trainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import gin
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import optimization
from official.modeling.multitask import base_model
@@ -33,7 +33,7 @@ class MultiTaskBaseTrainer(orbit.StandardTrainer):
def __init__(self,
multi_task: multitask.MultiTask,
- multi_task_model: Union[tf.keras.Model,
+ multi_task_model: Union[tf_keras.Model,
base_model.MultiTaskBaseModel],
optimizer: tf.optimizers.Optimizer,
trainer_options=None,
@@ -113,9 +113,9 @@ def training_losses(self):
# Builds the per-task metrics and losses.
# This the total summed training loss of tasks in the joint training.
self._training_losses = dict(
- total_loss=tf.keras.metrics.Mean("training_loss", dtype=tf.float32))
+ total_loss=tf_keras.metrics.Mean("training_loss", dtype=tf.float32))
for name in self.multi_task.tasks:
- self._training_losses[name] = tf.keras.metrics.Mean(
+ self._training_losses[name] = tf_keras.metrics.Mean(
"training_loss", dtype=tf.float32)
return self._training_losses
@@ -164,7 +164,7 @@ def step_fn(inputs):
task_metrics=self.training_metrics)
for key, loss in losses.items():
self.training_losses[key].update_state(loss)
+ self.global_step.assign_add(1)
self.strategy.run(
step_fn, args=(tf.nest.map_structure(next, iterator_map),))
- self.global_step.assign_add(1)
diff --git a/official/modeling/multitask/base_trainer_test.py b/official/modeling/multitask/base_trainer_test.py
index 2eb5acd252f..fe54099f5d1 100644
--- a/official/modeling/multitask/base_trainer_test.py
+++ b/official/modeling/multitask/base_trainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for multitask.base_trainer."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -47,7 +47,7 @@ def test_multitask_joint_trainer(self, distribution):
task_weights = {"foo": 1.0, "bar": 1.0}
test_multitask = multitask.MultiTask(
tasks=tasks, task_weights=task_weights)
- test_optimizer = tf.keras.optimizers.SGD(0.1)
+ test_optimizer = tf_keras.optimizers.SGD(0.1)
model = test_utils.MockMultiTaskModel()
test_trainer = base_trainer.MultiTaskBaseTrainer(
multi_task=test_multitask,
@@ -70,7 +70,7 @@ def test_trainer_with_configs(self):
task_config=test_utils.BarConfig(),
task_weight=0.5)))
test_multitask = multitask.MultiTask.from_config(config)
- test_optimizer = tf.keras.optimizers.SGD(0.1)
+ test_optimizer = tf_keras.optimizers.SGD(0.1)
model = test_utils.MockMultiTaskModel()
test_trainer = base_trainer.MultiTaskBaseTrainer(
multi_task=test_multitask,
diff --git a/official/modeling/multitask/configs.py b/official/modeling/multitask/configs.py
index a77d2c09560..afcd8bef9af 100644
--- a/official/modeling/multitask/configs.py
+++ b/official/modeling/multitask/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,9 +13,8 @@
# limitations under the License.
"""Configuration definitions for multi-task training."""
-from typing import Optional, Tuple
-
import dataclasses
+from typing import Optional, Tuple
from official.core import config_definitions as cfg
from official.modeling import hyperparams
@@ -26,7 +25,7 @@
class TaskRoutine(hyperparams.Config):
# TODO(hongkuny): deprecate the task_name once we migrated client code.
task_name: str = ""
- task_config: cfg.TaskConfig = None
+ task_config: cfg.TaskConfig | None = None
eval_steps: Optional[int] = None
task_weight: Optional[float] = 1.0
@@ -34,8 +33,12 @@ class TaskRoutine(hyperparams.Config):
@dataclasses.dataclass
class MultiTaskConfig(hyperparams.Config):
init_checkpoint: str = ""
- model: hyperparams.Config = None
+ model: hyperparams.Config | None = None
task_routines: Tuple[TaskRoutine, ...] = ()
+ # Configs for differential privacy
+ # These configs are only effective if you use create_optimizer in
+ # tensorflow_models/official/core/base_task.py
+ # DEPRECATED b/264611883
differential_privacy_config: Optional[
dp_configs.DifferentialPrivacyConfig] = None
@@ -54,23 +57,35 @@ class AnnealingSampleConfig(hyperparams.Config):
@dataclasses.dataclass
class TaskSamplingConfig(hyperparams.OneOfConfig):
type: str = ""
- uniform: hyperparams.Config = hyperparams.Config()
- proportional: ProportionalSampleConfig = ProportionalSampleConfig()
- annealing: AnnealingSampleConfig = AnnealingSampleConfig()
+ uniform: hyperparams.Config = dataclasses.field(
+ default_factory=hyperparams.Config
+ )
+ proportional: ProportionalSampleConfig = dataclasses.field(
+ default_factory=ProportionalSampleConfig
+ )
+ annealing: AnnealingSampleConfig = dataclasses.field(
+ default_factory=AnnealingSampleConfig
+ )
@dataclasses.dataclass
class MultiTaskTrainerConfig(cfg.TrainerConfig):
trainer_type: str = "interleaving"
- task_sampler: TaskSamplingConfig = TaskSamplingConfig(type="proportional")
+ task_sampler: TaskSamplingConfig = dataclasses.field(
+ default_factory=lambda: TaskSamplingConfig(type="proportional")
+ )
@dataclasses.dataclass
class MultiTaskExperimentConfig(hyperparams.Config):
"""An experiment config for multi-task training and multi-task evaluation."""
- task: MultiTaskConfig = MultiTaskConfig()
- trainer: MultiTaskTrainerConfig = MultiTaskTrainerConfig()
- runtime: cfg.RuntimeConfig = cfg.RuntimeConfig()
+ task: MultiTaskConfig = dataclasses.field(default_factory=MultiTaskConfig)
+ trainer: MultiTaskTrainerConfig = dataclasses.field(
+ default_factory=MultiTaskTrainerConfig
+ )
+ runtime: cfg.RuntimeConfig = dataclasses.field(
+ default_factory=cfg.RuntimeConfig
+ )
@dataclasses.dataclass
diff --git a/official/modeling/multitask/evaluator.py b/official/modeling/multitask/evaluator.py
index 9433a318afb..9708ed38909 100644
--- a/official/modeling/multitask/evaluator.py
+++ b/official/modeling/multitask/evaluator.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from typing import Dict, List, Optional, Union
import gin
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import train_utils
@@ -33,7 +33,7 @@ class MultiTaskEvaluator(orbit.AbstractEvaluator):
def __init__(
self,
eval_tasks: List[base_task.Task],
- model: Union[tf.keras.Model, base_model.MultiTaskBaseModel],
+ model: Union[tf_keras.Model, base_model.MultiTaskBaseModel],
global_step: Optional[tf.Variable] = None,
eval_steps: Optional[Dict[str, int]] = None,
checkpoint_exporter: Optional[train_utils.BestCheckpointExporter] = None):
@@ -41,7 +41,7 @@ def __init__(
Args:
eval_tasks: A list of tasks to evaluate.
- model: tf.keras.Model instance.
+ model: tf_keras.Model instance.
global_step: the global step variable.
eval_steps: a dictionary of steps to run eval keyed by task names.
checkpoint_exporter: an object that has the `maybe_export_checkpoint`
@@ -124,7 +124,7 @@ def validation_losses(self):
# Builds the per-task metrics and losses.
self._validation_losses = {}
for task in self.tasks:
- self._validation_losses[task.name] = tf.keras.metrics.Mean(
+ self._validation_losses[task.name] = tf_keras.metrics.Mean(
"validation_loss", dtype=tf.float32)
return self._validation_losses
diff --git a/official/modeling/multitask/evaluator_test.py b/official/modeling/multitask/evaluator_test.py
index 660adcfc34f..9f404b7d678 100644
--- a/official/modeling/multitask/evaluator_test.py
+++ b/official/modeling/multitask/evaluator_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for multitask.evaluator."""
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -35,11 +35,11 @@ def all_strategy_combinations():
)
-class MockModel(tf.keras.Model):
+class MockModel(tf_keras.Model):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
- self.dense = tf.keras.layers.Dense(1)
+ self.dense = tf_keras.layers.Dense(1)
def call(self, inputs):
print(inputs, type(inputs))
@@ -55,7 +55,7 @@ class MockTask(base_task.Task):
def build_metrics(self, training: bool = True):
del training
- return [tf.keras.metrics.Accuracy(name="acc")]
+ return [tf_keras.metrics.Accuracy(name="acc")]
def build_inputs(self, params):
@@ -73,7 +73,7 @@ def generate_data(_):
generate_data, num_parallel_calls=tf.data.experimental.AUTOTUNE)
return dataset.prefetch(buffer_size=1).batch(2, drop_remainder=True)
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
logs = super().validation_step(inputs, model, metrics)
logs["counter"] = tf.ones((1,), dtype=tf.float32)
return logs
diff --git a/official/modeling/multitask/interleaving_trainer.py b/official/modeling/multitask/interleaving_trainer.py
index 25fd20d1c01..b25b17e4d6e 100644
--- a/official/modeling/multitask/interleaving_trainer.py
+++ b/official/modeling/multitask/interleaving_trainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Union
import gin
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.multitask import base_model
from official.modeling.multitask import base_trainer
from official.modeling.multitask import multitask
@@ -29,9 +29,11 @@ class MultiTaskInterleavingTrainer(base_trainer.MultiTaskBaseTrainer):
def __init__(self,
multi_task: multitask.MultiTask,
- multi_task_model: Union[tf.keras.Model,
+ multi_task_model: Union[tf_keras.Model,
base_model.MultiTaskBaseModel],
- optimizer: tf.optimizers.Optimizer,
+ optimizer: Union[tf.optimizers.Optimizer,
+ tf_keras.optimizers.experimental.Optimizer,
+ tf_keras.optimizers.legacy.Optimizer],
task_sampler: sampler.TaskSampler,
trainer_options=None):
super().__init__(
@@ -69,9 +71,25 @@ def step_fn(inputs):
name: orbit.utils.create_global_step() for name in self.multi_task.tasks
}
+ # If the new Keras optimizer is used, we require all model variables are
+ # created before the training and let the optimizer to create the slot
+ # variable all together.
+ if isinstance(optimizer, tf_keras.optimizers.experimental.Optimizer):
+ multi_task_model.build() # pyrefly: ignore[missing-argument]
+ optimizer.build(multi_task_model.trainable_variables)
+
def task_step_counter(self, name):
return self._task_step_counters[name]
+ def _task_train_step(self, name):
+ """Runs one training step and updates counters."""
+ def _step_fn(inputs):
+ self._task_train_step_map[name](inputs)
+ self.global_step.assign_add(1)
+ self.task_step_counter(name).assign_add(1)
+
+ return _step_fn
+
def train_step(self, iterator_map):
# Sample one task to train according to a multinomial distribution
rn = tf.random.stateless_uniform(shape=[], seed=(0, self.global_step))
@@ -87,9 +105,7 @@ def train_step(self, iterator_map):
end = cumulative_sample_distribution[idx + 1]
if rn >= begin and rn < end:
self._strategy.run(
- self._task_train_step_map[name], args=(next(iterator_map[name]),))
- self.global_step.assign_add(1)
- self.task_step_counter(name).assign_add(1)
+ self._task_train_step(name), args=(next(iterator_map[name]),))
def train_loop_end(self):
"""Record loss and metric values per task."""
diff --git a/official/modeling/multitask/interleaving_trainer_test.py b/official/modeling/multitask/interleaving_trainer_test.py
index 6f871713ca7..93bedcc7fb2 100644
--- a/official/modeling/multitask/interleaving_trainer_test.py
+++ b/official/modeling/multitask/interleaving_trainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for multitask.interleaving_trainer."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -46,7 +46,7 @@ def test_multitask_interleaving_trainer(self, distribution):
test_utils.MockBarTask(params=test_utils.BarConfig(), name="bar")
]
test_multitask = multitask.MultiTask(tasks=tasks)
- test_optimizer = tf.keras.optimizers.SGD(0.1)
+ test_optimizer = tf_keras.optimizers.SGD(0.1)
model = test_utils.MockMultiTaskModel()
sampler = task_sampler.UniformTaskSampler(
task_weights=test_multitask.task_weights)
@@ -75,7 +75,7 @@ def test_trainer_with_configs(self, distribution):
task_weight=1.0)))
with distribution.scope():
test_multitask = multitask.MultiTask.from_config(config)
- test_optimizer = tf.keras.optimizers.SGD(0.1)
+ test_optimizer = tf_keras.optimizers.SGD(0.1)
model = test_utils.MockMultiTaskModel()
num_step = 1000
sampler = task_sampler.AnnealingTaskSampler(
diff --git a/official/modeling/multitask/multitask.py b/official/modeling/multitask/multitask.py
index 23e85afe879..d7531ccf894 100644
--- a/official/modeling/multitask/multitask.py
+++ b/official/modeling/multitask/multitask.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import abc
from typing import Dict, List, Optional, Text, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions
from official.core import task_factory
@@ -103,7 +103,7 @@ def create_optimizer(cls,
def joint_train_step(self, task_inputs,
multi_task_model: base_model.MultiTaskBaseModel,
- optimizer: tf.keras.optimizers.Optimizer, task_metrics,
+ optimizer: tf_keras.optimizers.Optimizer, task_metrics,
**kwargs):
"""The joint train step.
@@ -138,10 +138,10 @@ def joint_train_step(self, task_inputs,
self.tasks[name].process_metrics(task_metrics[name], labels, outputs,
**kwargs)
- # Scales loss as the default gradients allreduce performs sum inside
- # the optimizer.
- scaled_loss = total_loss / tf.distribute.get_strategy(
- ).num_replicas_in_sync
+ # Scales loss as the default gradients allreduce performs sum inside
+ # the optimizer.
+ scaled_loss = total_loss / tf.distribute.get_strategy(
+ ).num_replicas_in_sync
tvars = multi_task_model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
optimizer.apply_gradients(list(zip(grads, tvars)))
diff --git a/official/modeling/multitask/task_sampler.py b/official/modeling/multitask/task_sampler.py
index 5e062bd45b5..2b563428f92 100644
--- a/official/modeling/multitask/task_sampler.py
+++ b/official/modeling/multitask/task_sampler.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Utils to sample tasks for interleaved optimization."""
import abc
from typing import Union, Dict, Text
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.multitask import configs
diff --git a/official/modeling/multitask/task_sampler_test.py b/official/modeling/multitask/task_sampler_test.py
index 8b3d95ff462..eec141d81f8 100644
--- a/official/modeling/multitask/task_sampler_test.py
+++ b/official/modeling/multitask/task_sampler_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Tests for multitask.task_sampler."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.multitask import configs
from official.modeling.multitask import task_sampler as sampler
diff --git a/official/modeling/multitask/test_utils.py b/official/modeling/multitask/test_utils.py
index d9f32c28c0f..4bdefe1e1ba 100644
--- a/official/modeling/multitask/test_utils.py
+++ b/official/modeling/multitask/test_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,22 +14,24 @@
"""Testing utils for mock models and tasks."""
from typing import Dict, Text
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
from official.core import task_factory
from official.modeling.multitask import base_model
-class MockFooModel(tf.keras.Model):
+class MockFooModel(tf_keras.Model):
"""A mock model can consume 'foo' and 'bar' inputs."""
def __init__(self, shared_layer, *args, **kwargs):
super().__init__(*args, **kwargs)
self._share_layer = shared_layer
- self._foo_specific_layer = tf.keras.layers.Dense(1)
+ self._foo_specific_layer = tf_keras.layers.Dense(1)
+ self.inputs = {"foo": tf_keras.Input(shape=(2,), dtype=tf.float32),
+ "bar": tf_keras.Input(shape=(2,), dtype=tf.float32)}
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
self.add_loss(tf.zeros((1,), dtype=tf.float32))
if "foo" in inputs:
input_tensor = inputs["foo"]
@@ -38,14 +40,16 @@ def call(self, inputs):
return self._foo_specific_layer(self._share_layer(input_tensor))
-class MockBarModel(tf.keras.Model):
+class MockBarModel(tf_keras.Model):
+ """A mock model can only consume 'bar' inputs."""
def __init__(self, shared_layer, *args, **kwargs):
super().__init__(*args, **kwargs)
self._share_layer = shared_layer
- self._bar_specific_layer = tf.keras.layers.Dense(1)
+ self._bar_specific_layer = tf_keras.layers.Dense(1)
+ self.inputs = {"bar": tf_keras.Input(shape=(2,), dtype=tf.float32)}
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
self.add_loss(tf.zeros((2,), dtype=tf.float32))
return self._bar_specific_layer(self._share_layer(inputs["bar"]))
@@ -53,10 +57,10 @@ def call(self, inputs):
class MockMultiTaskModel(base_model.MultiTaskBaseModel):
def __init__(self, *args, **kwargs):
- self._shared_dense = tf.keras.layers.Dense(1)
+ self._shared_dense = tf_keras.layers.Dense(1)
super().__init__(*args, **kwargs)
- def _instantiate_sub_tasks(self) -> Dict[Text, tf.keras.Model]:
+ def _instantiate_sub_tasks(self) -> Dict[Text, tf_keras.Model]:
return {
"foo": MockFooModel(self._shared_dense),
"bar": MockBarModel(self._shared_dense)
@@ -92,16 +96,16 @@ class MockFooTask(base_task.Task):
def build_metrics(self, training: bool = True):
del training
- return [tf.keras.metrics.Accuracy(name="foo_acc")]
+ return [tf_keras.metrics.Accuracy(name="foo_acc")]
def build_inputs(self, params):
return mock_data("foo")
- def build_model(self) -> tf.keras.Model:
- return MockFooModel(shared_layer=tf.keras.layers.Dense(1))
+ def build_model(self) -> tf_keras.Model:
+ return MockFooModel(shared_layer=tf_keras.layers.Dense(1))
def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
- loss = tf.keras.losses.mean_squared_error(labels, model_outputs)
+ loss = tf_keras.losses.mean_squared_error(labels, model_outputs)
if aux_losses:
loss += tf.add_n(aux_losses)
return tf.reduce_mean(loss)
@@ -113,13 +117,13 @@ class MockBarTask(base_task.Task):
def build_metrics(self, training: bool = True):
del training
- return [tf.keras.metrics.Accuracy(name="bar_acc")]
+ return [tf_keras.metrics.Accuracy(name="bar_acc")]
def build_inputs(self, params):
return mock_data("bar")
def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
- loss = tf.keras.losses.mean_squared_error(labels, model_outputs)
+ loss = tf_keras.losses.mean_squared_error(labels, model_outputs)
if aux_losses:
loss += tf.add_n(aux_losses)
return tf.reduce_mean(loss)
diff --git a/official/modeling/multitask/train_lib.py b/official/modeling/multitask/train_lib.py
index a730bb160e8..b54ff364cce 100644
--- a/official/modeling/multitask/train_lib.py
+++ b/official/modeling/multitask/train_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,10 @@
"""Multitask training driver library."""
# pytype: disable=attribute-error
import os
-from typing import Any, List, Optional, Tuple
+from typing import Any, List, Mapping, Optional, Tuple, Union, Callable
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import base_trainer as core_lib
from official.core import train_utils
@@ -40,12 +40,17 @@ def run_experiment(
*,
distribution_strategy: tf.distribute.Strategy,
task: multitask.MultiTask,
- model: base_model.MultiTaskBaseModel,
+ model: base_model.MultiTaskBaseModel | tf_keras.Model,
mode: str,
params: configs.MultiTaskExperimentConfig,
model_dir: str,
- trainer: base_trainer.MultiTaskBaseTrainer = None
-) -> base_model.MultiTaskBaseModel:
+ run_post_eval: bool = False,
+ trainer: base_trainer.MultiTaskBaseTrainer = None,
+ eval_summary_manager: Optional[orbit.utils.SummaryManagerInterface] = None,
+ best_ckpt_exporter_creator: Optional[Any] = train_utils
+ .maybe_create_best_ckpt_exporter
+) -> Union[base_model.MultiTaskBaseModel, Tuple[base_model.MultiTaskBaseModel,
+ Mapping[Any, Any]]]:
"""Runs train/eval configured by the experiment params.
Args:
@@ -56,8 +61,15 @@ def run_experiment(
or 'continuous_eval'.
params: ExperimentConfig instance.
model_dir: A 'str', a path to store model checkpoints and summaries.
+ run_post_eval: Whether to run post eval once after training, metrics logs
+ are returned.
trainer: (optional) A multi-task trainer to use. If none is provided, a
default one will be created based on `params`.
+ eval_summary_manager: Instance of the eval summary manager. If set, the
+ `eval_summary_dir` will be ignored. Otherwise the eval summary manager
+ will be created internally for TensorBoard summaries by default from the
+ `eval_summary_dir`.
+ best_ckpt_exporter_creator: A functor for creating best checkpoint exporter.
Returns:
model: `base_model.MultiTaskBaseModel` instance.
@@ -66,15 +78,8 @@ def run_experiment(
is_training = 'train' in mode
is_eval = 'eval' in mode
with distribution_strategy.scope():
- optimizer = train_utils.create_optimizer(task, params)
- kwargs = dict(multi_task=task, multi_task_model=model, optimizer=optimizer)
- if params.trainer.trainer_type == 'interleaving':
- sampler = task_sampler.get_task_sampler(params.trainer.task_sampler,
- task.task_weights)
- kwargs.update(dict(task_sampler=sampler))
- if trainer is None:
- trainer = TRAINERS[params.trainer.trainer_type](
- **kwargs) if is_training else None
+ if is_training and trainer is None:
+ trainer = get_trainer(distribution_strategy, params, task, model)
if is_eval:
eval_steps = task.task_eval_steps
evaluator = evaluator_lib.MultiTaskEvaluator(
@@ -82,8 +87,7 @@ def run_experiment(
model=model,
eval_steps=eval_steps,
global_step=trainer.global_step if is_training else None,
- checkpoint_exporter=train_utils.maybe_create_best_ckpt_exporter(
- params, model_dir))
+ checkpoint_exporter=best_ckpt_exporter_creator(params, model_dir))
else:
evaluator = None
@@ -94,7 +98,6 @@ def run_experiment(
checkpoint = evaluator.checkpoint
global_step = evaluator.global_step
- # TODO(hongkuny,haozhangthu): Revisit initialization method.
checkpoint_manager = tf.train.CheckpointManager(
checkpoint,
directory=model_dir,
@@ -112,6 +115,7 @@ def run_experiment(
checkpoint_manager=checkpoint_manager,
summary_dir=os.path.join(model_dir, 'train'),
eval_summary_dir=os.path.join(model_dir, 'validation'),
+ eval_summary_manager=eval_summary_manager,
summary_interval=params.trainer.summary_interval)
logging.info('Starts to execute mode: %s', mode)
@@ -139,7 +143,62 @@ def timeout_fn():
else:
raise NotImplementedError('The mode is not implemented: %s' % mode)
- return model
+ if run_post_eval:
+ return model, evaluator.evaluate(
+ tf.convert_to_tensor(params.trainer.validation_steps)) # pytype: disable=bad-return-type # typed-keras
+ else:
+ return model
+
+
+def get_trainer(
+ distribution_strategy: tf.distribute.Strategy,
+ params: configs.MultiEvalExperimentConfig,
+ task: multitask.MultiTask,
+ model: base_model.MultiTaskBaseModel | tf_keras.Model,
+) -> orbit.StandardTrainer:
+ """Creates a multi-task trainer for the given task.
+
+ Args:
+ distribution_strategy: A distribution strategy.
+ params: ExperimentConfig instance.
+ task: A MultiTaskTask instance.
+ model: A MultiTaskBaseModel instance.
+
+ Returns:
+ An Orbit trainer instance.
+ """
+ with distribution_strategy.scope():
+ kwargs = dict(
+ multi_task=task,
+ multi_task_model=model,
+ optimizer=train_utils.create_optimizer(task, params),
+ )
+ if params.trainer.trainer_type == 'interleaving':
+ kwargs.update(
+ task_sampler=task_sampler.get_task_sampler(
+ params.trainer.task_sampler, task.task_weights
+ )
+ )
+ return TRAINERS[params.trainer.trainer_type](**kwargs)
+
+
+TrainActionsFactoryType = Callable[
+ [
+ configs.MultiEvalExperimentConfig,
+ orbit.StandardTrainer,
+ str,
+ tf.train.CheckpointManager,
+ ],
+ List[orbit.Action],
+]
+EvalActionsFactoryType = Callable[
+ [
+ configs.MultiEvalExperimentConfig,
+ orbit.AbstractEvaluator,
+ str,
+ ],
+ List[orbit.Action],
+]
def run_experiment_with_multitask_eval(
@@ -152,7 +211,13 @@ def run_experiment_with_multitask_eval(
model_dir: str,
run_post_eval: bool = False,
save_summary: bool = True,
- trainer: Optional[core_lib.Trainer] = None) -> Tuple[Any, Any]:
+ trainer: Optional[core_lib.Trainer] = None,
+ eval_summary_manager: Optional[orbit.utils.SummaryManagerInterface] = None,
+ best_ckpt_exporter_creator: Optional[Any] = train_utils
+ .maybe_create_best_ckpt_exporter,
+ train_actions_factory: Optional[TrainActionsFactoryType] = None,
+ eval_actions_factory: Optional[EvalActionsFactoryType] = None,
+) -> Tuple[Any, Any]:
"""Runs train/eval configured by the experiment params.
Args:
@@ -169,9 +234,16 @@ def run_experiment_with_multitask_eval(
trainer: the core_lib.Trainer instance. It should be created within the
strategy.scope(). If not provided, an instance will be created by default
if `mode` contains 'train'.
+ eval_summary_manager: Instance of the eval summary manager. If set, the
+ `eval_summary_dir` will be ignored. Otherwise the eval summary manager
+ will be created internally for TensorBoard summaries by default from the
+ `eval_summary_dir`.
+ best_ckpt_exporter_creator: A functor for creating best checkpoint exporter.
+ train_actions_factory: Optional factory function to create train actions.
+ eval_actions_factory: Optional factory function to create eval actions.
Returns:
- model: `tf.keras.Model` instance.
+ model: `tf_keras.Model` instance.
"""
is_training = 'train' in mode
@@ -187,7 +259,19 @@ def run_experiment_with_multitask_eval(
evaluate=False)
else:
trainer = None
- model = trainer.model if trainer else train_task.build_model()
+
+ # Build the model or fetch the pre-cached one (which could be either
+ # multi-task model or single task model).
+ if trainer is None:
+ if isinstance(train_task, multitask.MultiTask):
+ model = train_task.build_multitask_model()
+ else:
+ model = train_task.build_model()
+ else:
+ if isinstance(trainer, base_trainer.MultiTaskBaseTrainer):
+ model = trainer.multi_task_model
+ else:
+ model = trainer.model
if is_eval:
eval_steps = dict([(task_routine.task_config.name,
@@ -198,8 +282,7 @@ def run_experiment_with_multitask_eval(
model=model,
global_step=trainer.global_step if is_training else None,
eval_steps=eval_steps,
- checkpoint_exporter=train_utils.maybe_create_best_ckpt_exporter(
- params, model_dir))
+ checkpoint_exporter=best_ckpt_exporter_creator(params, model_dir))
else:
evaluator = None
@@ -218,6 +301,23 @@ def run_experiment_with_multitask_eval(
checkpoint_interval=params.trainer.checkpoint_interval,
init_fn=trainer.initialize if trainer else None)
+ if trainer and train_actions_factory:
+ # pytype: disable=wrong-keyword-args
+ train_actions = train_actions_factory(
+ params=params,
+ trainer=trainer,
+ model_dir=model_dir,
+ checkpoint_manager=checkpoint_manager,
+ )
+ # pytype: enable=wrong-keyword-args
+ else:
+ train_actions = None
+
+ if evaluator and eval_actions_factory:
+ eval_actions = eval_actions_factory(params, evaluator, model_dir)
+ else:
+ eval_actions = None
+
controller = orbit.Controller(
strategy=distribution_strategy,
trainer=trainer,
@@ -228,8 +328,12 @@ def run_experiment_with_multitask_eval(
summary_dir=os.path.join(model_dir, 'train') if save_summary else None,
eval_summary_dir=os.path.join(model_dir, 'validation') if
(save_summary) else None,
+ eval_summary_manager=eval_summary_manager,
summary_interval=params.trainer.summary_interval if
- (save_summary) else None)
+ (save_summary) else None,
+ train_actions=train_actions,
+ eval_actions=eval_actions,
+ )
logging.info('Starts to execute mode: %s', mode)
with distribution_strategy.scope():
diff --git a/official/modeling/multitask/train_lib_test.py b/official/modeling/multitask/train_lib_test.py
index 4c5fad2eb14..fd9e1c1c474 100644
--- a/official/modeling/multitask/train_lib_test.py
+++ b/official/modeling/multitask/train_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for multitask.train_lib."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -58,8 +58,9 @@ def setUp(self):
strategy_combinations.one_device_strategy_gpu,
],
mode='eager',
+ optimizer=['sgd_experimental', 'sgd'],
flag_mode=['train', 'eval', 'train_and_eval']))
- def test_end_to_end(self, distribution_strategy, flag_mode):
+ def test_end_to_end(self, distribution_strategy, optimizer, flag_mode):
model_dir = self.get_temp_dir()
experiment_config = configs.MultiTaskExperimentConfig(
task=configs.MultiTaskConfig(
@@ -70,6 +71,7 @@ def test_end_to_end(self, distribution_strategy, flag_mode):
task_name='bar', task_config=test_utils.BarConfig()))))
experiment_config = params_dict.override_params_dict(
experiment_config, self._test_config, is_strict=False)
+ experiment_config.trainer.optimizer_config.optimizer.type = optimizer
with distribution_strategy.scope():
test_multitask = multitask.MultiTask.from_config(experiment_config.task)
model = test_utils.MockMultiTaskModel()
diff --git a/official/modeling/optimization/__init__.py b/official/modeling/optimization/__init__.py
index c02b2b9a913..3b87f7fc68d 100644
--- a/official/modeling/optimization/__init__.py
+++ b/official/modeling/optimization/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/optimization/adafactor_optimizer.py b/official/modeling/optimization/adafactor_optimizer.py
index b7f1944e61e..6082c66ada1 100644
--- a/official/modeling/optimization/adafactor_optimizer.py
+++ b/official/modeling/optimization/adafactor_optimizer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/optimization/configs/__init__.py b/official/modeling/optimization/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/modeling/optimization/configs/__init__.py
+++ b/official/modeling/optimization/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/optimization/configs/learning_rate_config.py b/official/modeling/optimization/configs/learning_rate_config.py
index 9af3cb673f8..d4fb0437214 100644
--- a/official/modeling/optimization/configs/learning_rate_config.py
+++ b/official/modeling/optimization/configs/learning_rate_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,9 +13,9 @@
# limitations under the License.
"""Dataclasses for learning rate schedule config."""
+import dataclasses
from typing import List, Optional
-import dataclasses
from official.modeling.hyperparams import base_config
@@ -23,7 +23,7 @@
class ConstantLrConfig(base_config.Config):
"""Configuration for constant learning rate.
- This class is a containers for the constant learning rate decay configs.
+ This class is a container for the constant learning rate decay configs.
Attributes:
name: The name of the learning rate schedule. Defaults to Constant.
@@ -68,7 +68,7 @@ class StepwiseLrConfig(base_config.Config):
class ExponentialLrConfig(base_config.Config):
"""Configuration for exponential learning rate decay.
- This class is a containers for the exponential learning rate decay configs.
+ This class is a container for the exponential learning rate decay configs.
Attributes:
name: The name of the learning rate schedule. Defaults to ExponentialDecay.
@@ -92,7 +92,7 @@ class ExponentialLrConfig(base_config.Config):
class PolynomialLrConfig(base_config.Config):
"""Configuration for polynomial learning rate decay.
- This class is a containers for the polynomial learning rate decay configs.
+ This class is a container for the polynomial learning rate decay configs.
Attributes:
name: The name of the learning rate schedule. Defaults to PolynomialDecay.
@@ -118,8 +118,8 @@ class PolynomialLrConfig(base_config.Config):
class CosineLrConfig(base_config.Config):
"""Configuration for Cosine learning rate decay.
- This class is a containers for the cosine learning rate decay configs,
- tf.keras.experimental.CosineDecay.
+ This class is a container for the cosine learning rate decay configs,
+ tf_keras.experimental.CosineDecay.
Attributes:
name: The name of the learning rate schedule. Defaults to CosineDecay.
@@ -137,6 +137,35 @@ class CosineLrConfig(base_config.Config):
offset: int = 0
+@dataclasses.dataclass
+class CosineRestartsLrConfig(base_config.Config):
+ """Configuration for Cosine restarts learning rate decay.
+
+ This class is a container for the cosine restarts learning rate decay configs,
+ tf_keras.experimental.CosineDecayRestarts.
+
+ Attributes:
+ name: Name of the learning rate schedule. Defaults to CosineDecayRestarts.
+ initial_learning_rate: A float. The initial learning rate. Defaults to None.
+ first_decay_steps: A scalar `int32` or `int64` `Tensor` or a Python
+ number. Number of steps to decay over.
+ t_mul: A scalar `float32` or `float64` `Tensor` or a Python number.
+ Used to derive the number of iterations in the i-th period.
+ m_mul: A scalar `float32` or `float64` `Tensor` or a Python number.
+ Used to derive the initial learning rate of the i-th period.
+ alpha: A float. Minimum learning rate value as a fraction of
+ initial_learning_rate.
+ offset: An int. The offset applied to steps. Defaults to 0.
+ """
+ name: str = 'CosineDecayRestarts'
+ initial_learning_rate: Optional[float] = None
+ first_decay_steps: Optional[int] = None
+ t_mul: float = 2.0
+ m_mul: float = 1.0
+ alpha: float = 0.0
+ offset: int = 0
+
+
@dataclasses.dataclass
class DirectPowerLrConfig(base_config.Config):
"""Configuration for DirectPower learning rate decay.
diff --git a/official/modeling/optimization/configs/optimization_config.py b/official/modeling/optimization/configs/optimization_config.py
index a237eb4d96c..623d0f3e26c 100644
--- a/official/modeling/optimization/configs/optimization_config.py
+++ b/official/modeling/optimization/configs/optimization_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,9 +18,8 @@
It also has two helper functions get_optimizer_config, and get_lr_config from
an OptimizationConfig class.
"""
-from typing import Optional
-
import dataclasses
+from typing import Optional
from official.modeling.hyperparams import base_config
from official.modeling.hyperparams import oneof
@@ -42,21 +41,49 @@ class OptimizerConfig(oneof.OneOfConfig):
lars: lars optimizer.
adagrad: adagrad optimizer.
slide: slide optimizer.
+ adafactor: adafactor optimizer.
+ adafactor_keras: adafactor optimizer.
"""
type: Optional[str] = None
- sgd: opt_cfg.SGDConfig = opt_cfg.SGDConfig()
- sgd_experimental: opt_cfg.SGDExperimentalConfig = (
- opt_cfg.SGDExperimentalConfig())
- adam: opt_cfg.AdamConfig = opt_cfg.AdamConfig()
- adam_experimental: opt_cfg.AdamExperimentalConfig = (
- opt_cfg.AdamExperimentalConfig())
- adamw: opt_cfg.AdamWeightDecayConfig = opt_cfg.AdamWeightDecayConfig()
- lamb: opt_cfg.LAMBConfig = opt_cfg.LAMBConfig()
- rmsprop: opt_cfg.RMSPropConfig = opt_cfg.RMSPropConfig()
- lars: opt_cfg.LARSConfig = opt_cfg.LARSConfig()
- adagrad: opt_cfg.AdagradConfig = opt_cfg.AdagradConfig()
- slide: opt_cfg.SLIDEConfig = opt_cfg.SLIDEConfig()
- adafactor: opt_cfg.AdafactorConfig = opt_cfg.AdafactorConfig()
+ sgd: opt_cfg.SGDConfig = dataclasses.field(default_factory=opt_cfg.SGDConfig)
+ sgd_experimental: opt_cfg.SGDExperimentalConfig = dataclasses.field(
+ default_factory=opt_cfg.SGDExperimentalConfig
+ )
+ adam: opt_cfg.AdamConfig = dataclasses.field(
+ default_factory=opt_cfg.AdamConfig
+ )
+ adam_experimental: opt_cfg.AdamExperimentalConfig = dataclasses.field(
+ default_factory=opt_cfg.AdamExperimentalConfig
+ )
+ adamw: opt_cfg.AdamWeightDecayConfig = dataclasses.field(
+ default_factory=opt_cfg.AdamWeightDecayConfig
+ )
+ adamw_experimental: opt_cfg.AdamWeightDecayExperimentalConfig = (
+ dataclasses.field(
+ default_factory=opt_cfg.AdamWeightDecayExperimentalConfig
+ )
+ )
+ lamb: opt_cfg.LAMBConfig = dataclasses.field(
+ default_factory=opt_cfg.LAMBConfig
+ )
+ rmsprop: opt_cfg.RMSPropConfig = dataclasses.field(
+ default_factory=opt_cfg.RMSPropConfig
+ )
+ lars: opt_cfg.LARSConfig = dataclasses.field(
+ default_factory=opt_cfg.LARSConfig
+ )
+ adagrad: opt_cfg.AdagradConfig = dataclasses.field(
+ default_factory=opt_cfg.AdagradConfig
+ )
+ slide: opt_cfg.SLIDEConfig = dataclasses.field(
+ default_factory=opt_cfg.SLIDEConfig
+ )
+ adafactor: opt_cfg.AdafactorConfig = dataclasses.field(
+ default_factory=opt_cfg.AdafactorConfig
+ )
+ adafactor_keras: opt_cfg.AdafactorKerasConfig = dataclasses.field(
+ default_factory=opt_cfg.AdafactorKerasConfig
+ )
@dataclasses.dataclass
@@ -77,18 +104,36 @@ class LrConfig(oneof.OneOfConfig):
step_cosine_with_offset: Step cosine with a step offset.
"""
type: Optional[str] = None
- constant: lr_cfg.ConstantLrConfig = lr_cfg.ConstantLrConfig()
- stepwise: lr_cfg.StepwiseLrConfig = lr_cfg.StepwiseLrConfig()
- exponential: lr_cfg.ExponentialLrConfig = lr_cfg.ExponentialLrConfig()
- polynomial: lr_cfg.PolynomialLrConfig = lr_cfg.PolynomialLrConfig()
- cosine: lr_cfg.CosineLrConfig = lr_cfg.CosineLrConfig()
- power: lr_cfg.DirectPowerLrConfig = lr_cfg.DirectPowerLrConfig()
- power_linear: lr_cfg.PowerAndLinearDecayLrConfig = (
- lr_cfg.PowerAndLinearDecayLrConfig())
- power_with_offset: lr_cfg.PowerDecayWithOffsetLrConfig = (
- lr_cfg.PowerDecayWithOffsetLrConfig())
- step_cosine_with_offset: lr_cfg.StepCosineLrConfig = (
- lr_cfg.StepCosineLrConfig())
+ constant: lr_cfg.ConstantLrConfig = dataclasses.field(
+ default_factory=lr_cfg.ConstantLrConfig
+ )
+ stepwise: lr_cfg.StepwiseLrConfig = dataclasses.field(
+ default_factory=lr_cfg.StepwiseLrConfig
+ )
+ exponential: lr_cfg.ExponentialLrConfig = dataclasses.field(
+ default_factory=lr_cfg.ExponentialLrConfig
+ )
+ polynomial: lr_cfg.PolynomialLrConfig = dataclasses.field(
+ default_factory=lr_cfg.PolynomialLrConfig
+ )
+ cosine: lr_cfg.CosineLrConfig = dataclasses.field(
+ default_factory=lr_cfg.CosineLrConfig
+ )
+ cosine_restarts: lr_cfg.CosineRestartsLrConfig = dataclasses.field(
+ default_factory=lr_cfg.CosineRestartsLrConfig
+ )
+ power: lr_cfg.DirectPowerLrConfig = dataclasses.field(
+ default_factory=lr_cfg.DirectPowerLrConfig
+ )
+ power_linear: lr_cfg.PowerAndLinearDecayLrConfig = dataclasses.field(
+ default_factory=lr_cfg.PowerAndLinearDecayLrConfig
+ )
+ power_with_offset: lr_cfg.PowerDecayWithOffsetLrConfig = dataclasses.field(
+ default_factory=lr_cfg.PowerDecayWithOffsetLrConfig
+ )
+ step_cosine_with_offset: lr_cfg.StepCosineLrConfig = dataclasses.field(
+ default_factory=lr_cfg.StepCosineLrConfig
+ )
@dataclasses.dataclass
@@ -101,8 +146,12 @@ class WarmupConfig(oneof.OneOfConfig):
polynomial: polynomial warmup config.
"""
type: Optional[str] = None
- linear: lr_cfg.LinearWarmupConfig = lr_cfg.LinearWarmupConfig()
- polynomial: lr_cfg.PolynomialWarmupConfig = lr_cfg.PolynomialWarmupConfig()
+ linear: lr_cfg.LinearWarmupConfig = dataclasses.field(
+ default_factory=lr_cfg.LinearWarmupConfig
+ )
+ polynomial: lr_cfg.PolynomialWarmupConfig = dataclasses.field(
+ default_factory=lr_cfg.PolynomialWarmupConfig
+ )
@dataclasses.dataclass
@@ -116,7 +165,9 @@ class OptimizationConfig(base_config.Config):
learning_rate: learning rate oneof config.
warmup: warmup oneof config.
"""
- optimizer: OptimizerConfig = OptimizerConfig()
+ optimizer: OptimizerConfig = dataclasses.field(
+ default_factory=OptimizerConfig
+ )
ema: Optional[opt_cfg.EMAConfig] = None
- learning_rate: LrConfig = LrConfig()
- warmup: WarmupConfig = WarmupConfig()
+ learning_rate: LrConfig = dataclasses.field(default_factory=LrConfig)
+ warmup: WarmupConfig = dataclasses.field(default_factory=WarmupConfig)
diff --git a/official/modeling/optimization/configs/optimization_config_test.py b/official/modeling/optimization/configs/optimization_config_test.py
index 6fc11fea022..e55a17a5e05 100644
--- a/official/modeling/optimization/configs/optimization_config_test.py
+++ b/official/modeling/optimization/configs/optimization_config_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for optimization_config.py."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.optimization.configs import learning_rate_config as lr_cfg
from official.modeling.optimization.configs import optimization_config
diff --git a/official/modeling/optimization/configs/optimizer_config.py b/official/modeling/optimization/configs/optimizer_config.py
index 977ffa497c9..4aae11e3364 100644
--- a/official/modeling/optimization/configs/optimizer_config.py
+++ b/official/modeling/optimization/configs/optimizer_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -40,7 +40,7 @@ class BaseOptimizerConfig(base_config.Config):
class SGDConfig(BaseOptimizerConfig):
"""Configuration for SGD optimizer.
- The attributes for this class matches the arguments of tf.keras.optimizer.SGD.
+ The attributes for this class matches the arguments of tf_keras.optimizer.SGD.
Attributes:
name: name of the optimizer.
@@ -61,7 +61,7 @@ class SGDExperimentalConfig(BaseOptimizerConfig):
"""Configuration for SGD optimizer.
The attributes for this class matches the arguments of
- `tf.keras.optimizer.experimental.SGD`.
+ `tf_keras.optimizer.experimental.SGD`.
Attributes:
name: name of the optimizer.
@@ -80,7 +80,7 @@ class RMSPropConfig(BaseOptimizerConfig):
"""Configuration for RMSProp optimizer.
The attributes for this class matches the arguments of
- tf.keras.optimizers.RMSprop.
+ tf_keras.optimizers.RMSprop.
Attributes:
name: name of the optimizer.
@@ -101,7 +101,7 @@ class AdagradConfig(BaseOptimizerConfig):
"""Configuration for Adagrad optimizer.
The attributes of this class match the arguments of
- tf.keras.optimizer.Adagrad.
+ tf_keras.optimizer.Adagrad.
Attributes:
name: name of the optimizer.
@@ -119,7 +119,7 @@ class AdamConfig(BaseOptimizerConfig):
"""Configuration for Adam optimizer.
The attributes for this class matches the arguments of
- tf.keras.optimizer.Adam.
+ tf_keras.optimizer.Adam.
Attributes:
name: name of the optimizer.
@@ -141,7 +141,7 @@ class AdamExperimentalConfig(BaseOptimizerConfig):
"""Configuration for experimental Adam optimizer.
The attributes for this class matches the arguments of
- `tf.keras.optimizer.experimental.Adam`.
+ `tf_keras.optimizer.experimental.Adam`.
Attributes:
name: name of the optimizer.
@@ -190,12 +190,37 @@ class AdamWeightDecayConfig(BaseOptimizerConfig):
gradient_clip_norm: float = 1.0
+@dataclasses.dataclass
+class AdamWeightDecayExperimentalConfig(BaseOptimizerConfig):
+ """Configuration for Adam optimizer with weight decay.
+
+ Attributes:
+ name: name of the optimizer.
+ beta_1: decay rate for 1st order moments.
+ beta_2: decay rate for 2st order moments.
+ epsilon: epsilon value used for numerical stability in the optimizer.
+ amsgrad: boolean. Whether to apply AMSGrad variant of this algorithm from
+ the paper "On the Convergence of Adam and beyond".
+ weight_decay: float. Weight decay rate. Default to 0.
+ global_clipnorm: A positive float. Clips the gradients to this maximum
+ L2-norm. Default to 1.0.
+ jit_compile: if True, jit compile will be used.
+ """
+ name: str = "AdamWeightDecayExperimental"
+ beta_1: float = 0.9
+ beta_2: float = 0.999
+ epsilon: float = 1e-07
+ amsgrad: bool = False
+ weight_decay: float = 0.0
+ global_clipnorm: float = 1.0
+ jit_compile: bool = False
+
+
@dataclasses.dataclass
class LAMBConfig(BaseOptimizerConfig):
"""Configuration for LAMB optimizer.
- The attributes for this class matches the arguments of
- tensorflow_addons.optimizers.LAMB.
+ The attributes for this class matches the arguments of LAMB optimizer.
Attributes:
name: name of the optimizer.
@@ -311,3 +336,39 @@ class AdafactorConfig(BaseOptimizerConfig):
min_dim_size_to_factor: int = 128
epsilon1: float = 1e-30
epsilon2: float = 1e-3
+ weight_decay: Optional[float] = None
+ include_in_weight_decay: Optional[str] = None
+
+
+@dataclasses.dataclass
+class AdafactorKerasConfig(BaseOptimizerConfig):
+ """Configuration for AdafactorKeras optimizer.
+
+ The attributes for this class matches the arguments of the Adafactor
+ implementation provided by keras.
+
+ Attributes:
+ learning_rate: Initial value for the learning rate: either a floating
+ point value, or a
+ `tf_keras.optimizers.schedules.LearningRateSchedule` instance.
+ Defaults to 0.001.
+ beta_2_decay: float, defaults to -0.8. The decay rate of `beta_2`.
+ epsilon_1: float, defaults to 1e-30. A small offset to keep denominator
+ away from 0.
+ epsilon_2: float, defaults to 1e-3. A small offset to avoid learning
+ rate becoming too small by time.
+ clip_threshold: float, defaults to 1.0. Clipping threshold. This is a
+ part of Adafactor algorithm, independent from `clipnorm`, `clipvalue`
+ and `global_clipnorm`.
+ relative_step: bool, defaults to True. If `learning_rate` is a constant
+ and `relative_step=True`, learning rate will be adjusted based on
+ current iterations. This is a default learning rate decay in
+ Adafactor.
+ """
+ name: str = "Adafactor"
+ learning_rate: float = 0.001
+ beta_2_decay: float = -0.8
+ epsilon_1: float = 1e-30
+ epsilon_2: float = 1e-3
+ clip_threshold: float = 1.0
+ relative_step: bool = True
diff --git a/official/modeling/optimization/ema_optimizer.py b/official/modeling/optimization/ema_optimizer.py
index f1094d12092..aa9229d6476 100644
--- a/official/modeling/optimization/ema_optimizer.py
+++ b/official/modeling/optimization/ema_optimizer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,14 +14,41 @@
"""Exponential moving average optimizer."""
-from typing import List, Optional, Text
+from typing import List, Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=protected-access
-class ExponentialMovingAverage(tf.keras.optimizers.Optimizer):
+def maybe_merge_call(fn, strategy, *args, **kwargs):
+ """Maybe invoke `fn` via `merge_call` which may or may not be fulfilled.
+
+ The caller of this utility function requests to invoke `fn` via `merge_call`
+ at `tf.distribute.Strategy`'s best efforts. It is `tf.distribute`'s internal
+ whether the request is honored, depending on the `Strategy`. See
+ `tf.distribute.ReplicaContext.merge_call()` for more information.
+
+ This is adapted from tensorflow/python/distribute/merge_call_interim.py.
+
+ Args:
+ fn: the function to be invoked.
+ strategy: the `tf.distribute.Strategy` to call `fn` with.
+ *args: the positional arguments to be passed in to `fn`.
+ **kwargs: the keyword arguments to be passed in to `fn`.
+
+ Returns:
+ The return value of the `fn` call.
+ """
+ if strategy.extended._use_merge_call():
+ return tf.distribute.get_replica_context().merge_call(
+ fn, args=args, kwargs=kwargs
+ )
+ else:
+ return fn(strategy, *args, **kwargs)
+
+
+class ExponentialMovingAverage(tf_keras.optimizers.legacy.Optimizer):
"""Optimizer that computes an exponential moving average of the variables.
Empirically it has been found that using the moving average of the trained
@@ -32,7 +59,7 @@ class ExponentialMovingAverage(tf.keras.optimizers.Optimizer):
Example of usage for training:
```python
- opt = tf.keras.optimizers.SGD(learning_rate)
+ opt = tf_keras.optimizers.SGD(learning_rate)
opt = ExponentialMovingAverage(opt)
opt.shadow_copy(model)
@@ -47,17 +74,17 @@ class ExponentialMovingAverage(tf.keras.optimizers.Optimizer):
"""
def __init__(self,
- optimizer: tf.keras.optimizers.Optimizer,
+ optimizer: tf_keras.optimizers.Optimizer,
trainable_weights_only: bool = True,
average_decay: float = 0.99,
start_step: int = 0,
dynamic_decay: bool = True,
- name: Text = 'ExponentialMovingAverage',
+ name: str = 'ExponentialMovingAverage',
**kwargs):
"""Construct a new ExponentialMovingAverage optimizer.
Args:
- optimizer: `tf.keras.optimizers.Optimizer` that will be
+ optimizer: `tf_keras.optimizers.Optimizer` that will be
used to compute and apply gradients.
trainable_weights_only: 'bool', if True, only model trainable weights will
be updated. Otherwise, all model weights will be updated. This mainly
@@ -80,11 +107,11 @@ def __init__(self,
self._start_step = tf.constant(start_step, tf.float32)
self._dynamic_decay = dynamic_decay
self._optimizer = optimizer
- self._track_trackable(self._optimizer, 'base_optimizer')
+ self._track_trackable(self._optimizer, 'ema_base_optimizer')
self._average_weights = None
self._model_weights = None
- def shadow_copy(self, model: tf.keras.Model):
+ def shadow_copy(self, model: tf_keras.Model):
"""Creates shadow variables for the given model weights."""
if self._trainable_weights_only:
@@ -106,14 +133,15 @@ def has_shadow_copy(self):
def _create_slots(self, var_list):
self._optimizer._create_slots(var_list=var_list) # pylint: disable=protected-access
- def apply_gradients(self, grads_and_vars, name: Optional[Text] = None):
+ def apply_gradients(self, grads_and_vars, name: Optional[str] = None):
result = self._optimizer.apply_gradients(grads_and_vars, name)
- self.update_average(self.iterations)
+ maybe_merge_call(self.update_average, tf.distribute.get_strategy())
return result
@tf.function
- def update_average(self, step: tf.Tensor):
- step = tf.cast(step, tf.float32)
+ def update_average(self, strategy):
+ # Compute current decay value.
+ step = tf.cast(self.iterations, tf.float32)
if step < self._start_step:
decay = tf.constant(0., tf.float32)
elif self._dynamic_decay:
@@ -122,18 +150,16 @@ def update_average(self, step: tf.Tensor):
else:
decay = self._average_decay
- def _apply_moving(v_moving, v_normal):
- diff = v_moving - v_normal
- v_moving.assign_sub(tf.cast(1. - decay, v_moving.dtype) * diff)
- return v_moving
+ def _apply_moving(average, normal):
+ diff = average - normal
+ average.assign_sub(tf.cast(1.0 - decay, average.dtype) * diff)
+ return average
- def _update(strategy, v_moving_and_v_normal):
- for v_moving, v_normal in v_moving_and_v_normal:
- strategy.extended.update(v_moving, _apply_moving, args=(v_normal,))
-
- ctx = tf.distribute.get_replica_context()
- return ctx.merge_call(_update, args=(zip(self._average_weights,
- self._model_weights),))
+ # Update moving average with the latest value.
+ for average, normal in zip(self._average_weights, self._model_weights): # pyrefly: ignore[bad-argument-type]
+ strategy.extended.update(
+ average, _apply_moving, args=(normal,), group=False
+ )
def swap_weights(self):
"""Swap the average and moving weights.
@@ -147,8 +173,9 @@ def swap_weights(self):
strategy = tf.distribute.get_strategy()
strategy.run(self._swap_weights, args=())
else:
- raise ValueError('Swapping weights must occur under a '
- 'tf.distribute.Strategy')
+ raise ValueError(
+ 'Swapping weights must occur under a tf.distribute.Strategy.'
+ )
@tf.function
def _swap_weights(self):
@@ -162,16 +189,30 @@ def fn_2(a, b):
a.assign_sub(b)
return a
- def swap(strategy, a_and_b):
+ def _swap(strategy, a_and_b):
"""Swap `a` and `b` and mirror to all devices."""
for a, b in a_and_b:
strategy.extended.update(a, fn_0, args=(b,)) # a = a + b
strategy.extended.update(b, fn_1, args=(a,)) # b = a - b
strategy.extended.update(a, fn_2, args=(b,)) # a = a - b
- ctx = tf.distribute.get_replica_context()
- return ctx.merge_call(
- swap, args=(zip(self._average_weights, self._model_weights),))
+ # Use merge_call if requested by strategy and always for TPUStrategy as
+ # the use of merge_call is not recommended and deprecated for other
+ # strategies such as mirrored strategy (MS) and multi-worker mirrored
+ # strategy (MWMS) if nccl/collective_ops are used, which can operate in
+ # pure replica context.
+ strategy = tf.distribute.get_strategy()
+ if isinstance(strategy, tf.distribute.TPUStrategy):
+ maybe_merge_call(
+ _swap,
+ strategy,
+ zip(self._average_weights, self._model_weights), # pyrefly: ignore[bad-argument-type]
+ )
+ else:
+ _swap(
+ strategy,
+ zip(self._average_weights, self._model_weights), # pyrefly: ignore[bad-argument-type]
+ )
def assign_average_vars(self, var_list: List[tf.Variable]):
"""Assign variables in var_list with their respective averages.
@@ -238,7 +279,7 @@ def _resource_apply_sparse_duplicate_indices(self, grad, var, indices):
def get_config(self):
config = {
- 'optimizer': tf.keras.optimizers.serialize(self._optimizer),
+ 'optimizer': tf_keras.optimizers.serialize(self._optimizer),
'average_decay': self._average_decay,
'start_step': self._start_step,
'dynamic_decay': self._dynamic_decay,
@@ -248,7 +289,7 @@ def get_config(self):
@classmethod
def from_config(cls, config, custom_objects=None):
- optimizer = tf.keras.optimizers.deserialize(
+ optimizer = tf_keras.optimizers.deserialize(
config.pop('optimizer'),
custom_objects=custom_objects,
)
diff --git a/official/modeling/optimization/lamb.py b/official/modeling/optimization/lamb.py
new file mode 100644
index 00000000000..7339c616c92
--- /dev/null
+++ b/official/modeling/optimization/lamb.py
@@ -0,0 +1,252 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Layer-wise Adaptive Moments (LAMB) optimizer.
+
+See paper [Large Batch Optimization for Deep Learning: Training BERT in
+76 minutes](https://arxiv.org/abs/1904.00962).
+"""
+import re
+from typing import Optional, Union, Callable, List
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+FloatTensorLike = Union[tf.Tensor, float, np.float16, np.float32]
+
+
+@tf_keras.utils.register_keras_serializable(package="Addons")
+class LAMB(tf_keras.optimizers.legacy.Optimizer):
+ """Optimizer that implements the Layer-wise Adaptive Moments (LAMB).
+
+ See paper [Large Batch Optimization for Deep Learning: Training BERT
+ in 76 minutes](https://arxiv.org/abs/1904.00962).
+ """
+
+ def __init__(
+ self,
+ learning_rate: Union[FloatTensorLike, Callable] = 0.001,
+ beta_1: FloatTensorLike = 0.9,
+ beta_2: FloatTensorLike = 0.999,
+ epsilon: FloatTensorLike = 1e-6,
+ weight_decay_rate: FloatTensorLike = 0.0,
+ exclude_from_weight_decay: Optional[List[str]] = None,
+ exclude_from_layer_adaptation: Optional[List[str]] = None,
+ name: str = "LAMB",
+ **kwargs,
+ ):
+ """Construct a new LAMB optimizer.
+
+ Args:
+ learning_rate: A `Tensor` or a floating point value. or a schedule that
+ is a `tf_keras.optimizers.schedules.LearningRateSchedule` The learning
+ rate.
+ beta_1: A `float` value or a constant `float` tensor. The exponential
+ decay rate for the 1st moment estimates.
+ beta_2: A `float` value or a constant `float` tensor. The exponential
+ decay rate for the 2nd moment estimates.
+ epsilon: A small constant for numerical stability.
+ weight_decay_rate: weight decay rate.
+ exclude_from_weight_decay: List of regex patterns of variables excluded
+ from weight decay. Variables whose name contain a substring matching
+ the pattern will be excluded.
+ exclude_from_layer_adaptation: List of regex patterns of variables
+ excluded from layer adaptation. Variables whose name contain a
+ substring matching the pattern will be excluded.
+ name: Optional name for the operations created when applying gradients.
+ Defaults to "LAMB".
+ **kwargs: keyword arguments. Allowed to be {`clipnorm`, `clipvalue`,
+ `lr`, `decay`}. `clipnorm` is clip gradients by norm; `clipvalue` is
+ clip gradients by value, `decay` is included for backward
+ compatibility to allow time inverse decay of learning rate. `lr` is
+ included for backward compatibility, recommended to use
+ `learning_rate` instead.
+ """
+ super().__init__(name, **kwargs)
+
+ # Just adding the square of the weights to the loss function is *not*
+ # the correct way of using L2 regularization/weight decay with Adam,
+ # since that will interact with the m and v parameters in strange ways.
+ #
+ # Instead we want to decay the weights in a manner that doesn't interact
+ # with the m/v parameters.
+ self._set_hyper("weight_decay_rate", weight_decay_rate)
+ self._set_hyper("learning_rate", kwargs.get("lr", learning_rate))
+
+ # This is learning rate decay for using keras learning rate schedule.
+ self._set_hyper("decay", self._initial_decay)
+ self._set_hyper("beta_1", beta_1)
+ self._set_hyper("beta_2", beta_2)
+ self.epsilon = epsilon or tf.backend_config.epsilon()
+ self.exclude_from_weight_decay = exclude_from_weight_decay
+ # exclude_from_layer_adaptation is set to exclude_from_weight_decay if
+ # the arg is None.
+ if exclude_from_layer_adaptation:
+ self.exclude_from_layer_adaptation = exclude_from_layer_adaptation
+ else:
+ self.exclude_from_layer_adaptation = exclude_from_weight_decay
+
+ def _create_slots(self, var_list):
+ # Create slots for the first and second moments.
+ # Separate for-loops to respect the ordering of slot variables from v1.
+ for var in var_list:
+ self.add_slot(var, "m")
+ for var in var_list:
+ self.add_slot(var, "v")
+
+ def _prepare_local(self, var_device, var_dtype, apply_state):
+ super()._prepare_local(var_device, var_dtype, apply_state)
+
+ local_step = tf.cast(self.iterations + 1, var_dtype)
+ beta_1_t = tf.identity(self._get_hyper("beta_1", var_dtype))
+ beta_2_t = tf.identity(self._get_hyper("beta_2", var_dtype))
+ weight_decay_rate = tf.identity(
+ self._get_hyper("weight_decay_rate", var_dtype)
+ )
+ beta_1_power = tf.pow(beta_1_t, local_step)
+ beta_2_power = tf.pow(beta_2_t, local_step)
+ apply_state[(var_device, var_dtype)].update(
+ dict(
+ weight_decay_rate=weight_decay_rate,
+ epsilon=tf.convert_to_tensor(self.epsilon, var_dtype),
+ beta_1_t=beta_1_t,
+ beta_1_power=beta_1_power,
+ one_minus_beta_1_t=1 - beta_1_t,
+ beta_2_t=beta_2_t,
+ beta_2_power=beta_2_power,
+ one_minus_beta_2_t=1 - beta_2_t,
+ )
+ )
+
+ def _resource_apply_dense(self, grad, var, apply_state=None):
+ var_device, var_dtype = var.device, var.dtype.base_dtype
+ coefficients = (apply_state or {}).get(
+ (var_device, var_dtype)
+ ) or self._fallback_apply_state(var_device, var_dtype)
+
+ # m_t = beta1 * m + (1 - beta1) * g_t
+ m = self.get_slot(var, "m")
+ m_scaled_g_values = grad * coefficients["one_minus_beta_1_t"]
+ m_t = m * coefficients["beta_1_t"] + m_scaled_g_values
+ m_t = m.assign(m_t, use_locking=self._use_locking)
+ # v_t = beta2 * v + (1 - beta2) * (g_t * g_t)
+ v = self.get_slot(var, "v")
+ v_scaled_g_values = (grad * grad) * coefficients["one_minus_beta_2_t"]
+ v_t = v * coefficients["beta_2_t"] + v_scaled_g_values
+ v_t = v.assign(v_t, use_locking=self._use_locking)
+
+ m_t_hat = m_t / (1.0 - coefficients["beta_1_power"])
+ v_t_hat = v_t / (1.0 - coefficients["beta_2_power"])
+
+ v_sqrt = tf.sqrt(v_t_hat)
+ update = m_t_hat / (v_sqrt + coefficients["epsilon"])
+
+ var_name = self._get_variable_name(var.name)
+ if self._do_use_weight_decay(var_name):
+ update += coefficients["weight_decay_rate"] * var
+
+ ratio = 1.0
+ if self._do_layer_adaptation(var_name):
+ w_norm = tf.norm(var, ord=2)
+ g_norm = tf.norm(update, ord=2)
+ ratio = tf.where(
+ tf.greater(w_norm, 0),
+ tf.where(tf.greater(g_norm, 0), (w_norm / g_norm), 1.0),
+ 1.0,
+ )
+
+ var_update = var - ratio * coefficients["lr_t"] * update
+ return var.assign(var_update, use_locking=self._use_locking)
+
+ def _resource_apply_sparse(self, grad, var, indices, apply_state=None):
+ var_device, var_dtype = var.device, var.dtype.base_dtype
+ coefficients = (apply_state or {}).get(
+ (var_device, var_dtype)
+ ) or self._fallback_apply_state(var_device, var_dtype)
+
+ # m_t = beta1 * m + (1 - beta1) * g_t
+ m = self.get_slot(var, "m")
+ m_scaled_g_values = grad * coefficients["one_minus_beta_1_t"]
+ m_t = m.assign(m * coefficients["beta_1_t"], use_locking=self._use_locking)
+ with tf.control_dependencies([m_t]):
+ m_t = self._resource_scatter_add(m, indices, m_scaled_g_values)
+
+ # v_t = beta2 * v + (1 - beta2) * (g_t * g_t)
+ v = self.get_slot(var, "v")
+ v_scaled_g_values = (grad * grad) * coefficients["one_minus_beta_2_t"]
+ v_t = v.assign(v * coefficients["beta_2_t"], use_locking=self._use_locking)
+ with tf.control_dependencies([v_t]):
+ v_t = self._resource_scatter_add(v, indices, v_scaled_g_values)
+
+ m_t_hat = m_t / (1.0 - coefficients["beta_1_power"])
+ v_t_hat = v_t / (1.0 - coefficients["beta_2_power"])
+
+ v_sqrt = tf.sqrt(v_t_hat)
+ update = m_t_hat / (v_sqrt + coefficients["epsilon"])
+
+ var_name = self._get_variable_name(var.name)
+ if self._do_use_weight_decay(var_name):
+ update += coefficients["weight_decay_rate"] * var
+
+ ratio = 1.0
+ if self._do_layer_adaptation(var_name):
+ w_norm = tf.norm(var, ord=2)
+ g_norm = tf.norm(update, ord=2)
+ ratio = tf.where(
+ tf.greater(w_norm, 0),
+ tf.where(tf.greater(g_norm, 0), (w_norm / g_norm), 1.0),
+ 1.0,
+ )
+
+ var_update = var.assign_sub(
+ ratio * coefficients["lr_t"] * update, use_locking=self._use_locking
+ )
+ return tf.group(*[var_update, m_t, v_t])
+
+ def get_config(self):
+ config = super().get_config()
+ config.update({
+ "learning_rate": self._serialize_hyperparameter("learning_rate"),
+ "weight_decay_rate": self._serialize_hyperparameter(
+ "weight_decay_rate"
+ ),
+ "decay": self._serialize_hyperparameter("decay"),
+ "beta_1": self._serialize_hyperparameter("beta_1"),
+ "beta_2": self._serialize_hyperparameter("beta_2"),
+ "epsilon": self.epsilon,
+ })
+ return config
+
+ def _do_use_weight_decay(self, param_name):
+ """Whether to use L2 weight decay for `param_name`."""
+ if self.exclude_from_weight_decay:
+ for r in self.exclude_from_weight_decay:
+ if re.search(r, param_name) is not None:
+ return False
+ return True
+
+ def _do_layer_adaptation(self, param_name):
+ """Whether to do layer-wise learning rate adaptation for `param_name`."""
+ if self.exclude_from_layer_adaptation:
+ for r in self.exclude_from_layer_adaptation:
+ if re.search(r, param_name) is not None:
+ return False
+ return True
+
+ def _get_variable_name(self, param_name):
+ """Get the variable name from the tensor name."""
+ m = re.match("^(.*):\\d+$", param_name)
+ if m is not None:
+ param_name = m.group(1)
+ return param_name
diff --git a/official/modeling/optimization/lamb_test.py b/official/modeling/optimization/lamb_test.py
new file mode 100644
index 00000000000..24781a9a8e6
--- /dev/null
+++ b/official/modeling/optimization/lamb_test.py
@@ -0,0 +1,177 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for LAMB Optimizer."""
+import numpy as np
+from numpy import linalg
+
+import tensorflow as tf, tf_keras
+
+from official.modeling.optimization import lamb
+
+
+def lamb_update_numpy(param,
+ g_t,
+ t,
+ m,
+ v,
+ lr=0.001,
+ lamb_wd=0.0,
+ beta1=0.9,
+ beta2=0.999,
+ epsilon=1e-6):
+
+ m_t = beta1 * m + (1 - beta1) * g_t
+ v_t = beta2 * v + (1 - beta2) * g_t * g_t
+
+ m_t_hat = m_t / (1 - beta1**(t + 1))
+ v_t_hat = v_t / (1 - beta2**(t + 1))
+ update = m_t_hat / (np.sqrt(v_t_hat) + epsilon)
+
+ update += lamb_wd * param
+
+ w_norm = linalg.norm(param, ord=2)
+ g_norm = linalg.norm(update, ord=2)
+ ratio = np.where(w_norm > 0, np.where(g_norm > 0, (w_norm / g_norm), 1.0),
+ 1.0)
+
+ param_t = param - ratio * lr * update
+ return param_t, m_t, v_t
+
+
+def get_beta_accumulators(opt, dtype):
+ local_step = tf.cast(opt.iterations + 1, dtype)
+ beta_1_t = tf.cast(opt._get_hyper("beta_1"), dtype)
+ beta_1_power = tf.math.pow(beta_1_t, local_step)
+ beta_2_t = tf.cast(opt._get_hyper("beta_2"), dtype)
+ beta_2_power = tf.math.pow(beta_2_t, local_step)
+ return (beta_1_power, beta_2_power)
+
+
+class LAMBTest(tf.test.TestCase):
+
+ def test_sparse(self):
+ dtype = tf.float32
+ # Initialize tf for numpy implementation.
+ m0, v0, m1, v1 = 0.0, 0.0, 0.0, 0.0
+ var0_np = np.array([1.0, 1.0, 2.0], dtype=dtype.as_numpy_dtype)
+ grads0_np = np.array([0.1, 0.0, 0.1], dtype=dtype.as_numpy_dtype)
+ var1_np = np.array([3.0, 3.0, 4.0], dtype=dtype.as_numpy_dtype)
+ grads1_np = np.array([0.01, 0.0, 0.01], dtype=dtype.as_numpy_dtype)
+
+ var0 = tf.Variable(var0_np)
+ var1 = tf.Variable(var1_np)
+ grads0_np_indices = np.array([0, 2], dtype=np.int32)
+ grads0 = tf.IndexedSlices(
+ tf.constant(grads0_np[grads0_np_indices]),
+ tf.constant(grads0_np_indices),
+ tf.constant([3]),
+ )
+ grads1_np_indices = np.array([0, 2], dtype=np.int32)
+ grads1 = tf.IndexedSlices(
+ tf.constant(grads1_np[grads1_np_indices]),
+ tf.constant(grads1_np_indices),
+ tf.constant([3]),
+ )
+ opt = lamb.LAMB()
+
+ # Fetch params to validate initial values
+ np.testing.assert_allclose(np.asanyarray([1.0, 1.0, 2.0]), var0.numpy())
+ np.testing.assert_allclose(np.asanyarray([3.0, 3.0, 4.0]), var1.numpy())
+
+ # Run 3 steps of LAMB
+ for t in range(3):
+ beta_1_power, beta_2_power = get_beta_accumulators(opt, dtype)
+ self.assertAllClose(0.9 ** (t + 1), beta_1_power)
+ self.assertAllClose(0.999 ** (t + 1), beta_2_power)
+
+ opt.apply_gradients(zip([grads0, grads1], [var0, var1]))
+
+ var0_np, m0, v0 = lamb_update_numpy(var0_np, grads0_np, t, m0, v0)
+ var1_np, m1, v1 = lamb_update_numpy(var1_np, grads1_np, t, m1, v1)
+
+ # Validate updated params
+ self.assertAllClose(var0_np, var0.numpy())
+ self.assertAllClose(var1_np, var1.numpy())
+
+ def test_basic_with_learning_rate_decay(self):
+ dtype = tf.float32
+ # Initialize variables for numpy implementation.
+ m0, v0, m1, v1 = 0.0, 0.0, 0.0, 0.0
+ var0_np = np.array([1.0, 2.0], dtype=dtype.as_numpy_dtype)
+ grads0_np = np.array([0.1, 0.1], dtype=dtype.as_numpy_dtype)
+ var1_np = np.array([3.0, 4.0], dtype=dtype.as_numpy_dtype)
+ grads1_np = np.array([0.01, 0.01], dtype=dtype.as_numpy_dtype)
+
+ var0 = tf.Variable(var0_np, name="var0")
+ var1 = tf.Variable(var1_np, name="var1")
+ grads0 = tf.constant(grads0_np)
+ grads1 = tf.constant(grads1_np)
+
+ learning_rate = 0.001
+ beta_1 = 0.9
+ beta_2 = 0.999
+ epsilon = 1e-7
+ decay = 0.5
+ lamb_wd = 0.01
+
+ opt = lamb.LAMB(
+ learning_rate=learning_rate,
+ beta_1=beta_1,
+ beta_2=beta_2,
+ epsilon=epsilon,
+ weight_decay_rate=lamb_wd,
+ decay=decay,
+ )
+
+ # Run 3 steps of LAMB
+ for t in range(3):
+ opt.apply_gradients(zip([grads0, grads1], [var0, var1]))
+
+ lr_np = learning_rate / (1 + decay * t)
+
+ var0_np, m0, v0 = lamb_update_numpy(
+ var0_np, grads0_np, t, m0, v0, lr=lr_np, lamb_wd=lamb_wd)
+ var1_np, m1, v1 = lamb_update_numpy(
+ var1_np, grads1_np, t, m1, v1, lr=lr_np, lamb_wd=lamb_wd)
+
+ # Validate updated params
+ self.assertAllClose(var0_np, var0.numpy())
+ self.assertAllClose(var1_np, var1.numpy())
+
+ def test_exclude_weight_decay(self):
+ opt = lamb.LAMB(
+ 0.01, weight_decay_rate=0.01, exclude_from_weight_decay=["var1"]
+ )
+ assert opt._do_use_weight_decay("var0")
+ assert not opt._do_use_weight_decay("var1")
+ assert not opt._do_use_weight_decay("var1_weight")
+
+ def test_exclude_layer_adaptation(self):
+ opt = lamb.LAMB(0.01, exclude_from_layer_adaptation=["var1"])
+ assert opt._do_layer_adaptation("var0")
+ assert not opt._do_layer_adaptation("var1")
+ assert not opt._do_layer_adaptation("var1_weight")
+
+ def test_serialization(self):
+ optimizer = lamb.LAMB(1e-4)
+ config = tf_keras.optimizers.serialize(optimizer, use_legacy_format=True)
+ new_optimizer = tf_keras.optimizers.deserialize(
+ config, use_legacy_format=True
+ )
+ assert new_optimizer.get_config() == optimizer.get_config()
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/modeling/optimization/lars_optimizer.py b/official/modeling/optimization/lars.py
similarity index 98%
rename from official/modeling/optimization/lars_optimizer.py
rename to official/modeling/optimization/lars.py
index b2e14132ca5..37dc02cfd85 100644
--- a/official/modeling/optimization/lars_optimizer.py
+++ b/official/modeling/optimization/lars.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,13 +16,13 @@
import re
from typing import Text, List, Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=protected-access
-class LARS(tf.keras.optimizers.Optimizer):
+class LARS(tf_keras.optimizers.legacy.Optimizer):
"""Layer-wise Adaptive Rate Scaling for large batch training.
Introduced by "Large Batch Training of Convolutional Networks" by Y. You,
diff --git a/official/modeling/optimization/legacy_adamw.py b/official/modeling/optimization/legacy_adamw.py
index 55abc6859c8..9e4ae6d4cec 100644
--- a/official/modeling/optimization/legacy_adamw.py
+++ b/official/modeling/optimization/legacy_adamw.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,10 +17,10 @@
import re
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class AdamWeightDecay(tf.keras.optimizers.Adam):
+class AdamWeightDecay(tf_keras.optimizers.legacy.Adam):
"""Adam enables L2 weight decay and clip_by_global_norm on gradients.
[Warning!]: Keras optimizer supports gradient clipping and has an AdamW
diff --git a/official/modeling/optimization/lr_schedule.py b/official/modeling/optimization/lr_schedule.py
index 4b79d9e3e79..2ea29094f92 100644
--- a/official/modeling/optimization/lr_schedule.py
+++ b/official/modeling/optimization/lr_schedule.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,35 +17,36 @@
import math
from typing import Mapping, Any, Union, Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _make_offset_wrapper(new_class_name: str, base_lr_class):
"""Generates a offset wrapper of learning rate schedule.
- It will returns a subclass of the the `base_lr_class`, the subclass takes an
+ It will returns a subclass of the `base_lr_class`, the subclass takes an
`offset` argument in the constructor. When the new class instance is called,
the behavior is:
new_class_object(step) = base_lr_class_object(step - offset)
Example:
CosineDecayWithOffset = _make_offset_wrapper(
- 'CosineDecayWithOffset', tf.keras.experimental.CosineDecay)
+ 'CosineDecayWithOffset',
+ tf_keras.optimizers.schedules.CosineDecay)
# Use the lr:
lr = CosineDecayWithOffset(offset=100, initial_learning_rate=0.1,
decay_steps=1000)
- lr(101) # equals to tf.keras.experimental.CosineDecay(...)(101-100)
+ lr(101) # equals to keras.optimizers.schedules.CosineDecay(...)(101-100)
Args:
new_class_name: the name of the new class.
base_lr_class: the base learning rate schedule class. Should be subclass of
- tf.keras.optimizers.schedules.LearningRateSchedule
+ tf_keras.optimizers.schedules.LearningRateSchedule
Returns:
A new class (subclass of the base_lr_class) that can take an offset.
"""
assert issubclass(base_lr_class,
- tf.keras.optimizers.schedules.LearningRateSchedule), (
+ tf_keras.optimizers.schedules.LearningRateSchedule), (
"base_lr_class should be subclass of keras "
f"LearningRateSchedule, got {base_lr_class}")
@@ -79,22 +80,28 @@ def offset_learning_rate_call(self, step):
PiecewiseConstantDecayWithOffset = _make_offset_wrapper(
"PiecewiseConstantDecayWithOffset",
- tf.keras.optimizers.schedules.PiecewiseConstantDecay)
+ tf_keras.optimizers.schedules.PiecewiseConstantDecay)
PolynomialDecayWithOffset = _make_offset_wrapper(
- "PolynomialDecayWithOffset", tf.keras.optimizers.schedules.PolynomialDecay)
+ "PolynomialDecayWithOffset", tf_keras.optimizers.schedules.PolynomialDecay)
ExponentialDecayWithOffset = _make_offset_wrapper(
"ExponentialDecayWithOffset",
- tf.keras.optimizers.schedules.ExponentialDecay)
-CosineDecayWithOffset = _make_offset_wrapper("CosineDecayWithOffset",
- tf.keras.experimental.CosineDecay)
-
-
-class LinearWarmup(tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.ExponentialDecay)
+CosineDecayWithOffset = _make_offset_wrapper(
+ "CosineDecayWithOffset",
+ tf_keras.optimizers.schedules.CosineDecay,
+)
+CosineDecayRestartsWithOffset = _make_offset_wrapper(
+ "CosineDecayRestartsWithOffset",
+ tf_keras.optimizers.schedules.CosineDecayRestarts,
+)
+
+
+class LinearWarmup(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Linear warmup schedule."""
def __init__(self,
after_warmup_lr_sched: Union[
- tf.keras.optimizers.schedules.LearningRateSchedule, float],
+ tf_keras.optimizers.schedules.LearningRateSchedule, float],
warmup_steps: int,
warmup_learning_rate: float,
name: Optional[str] = None):
@@ -110,7 +117,7 @@ def __init__(self,
steps.
Args:
- after_warmup_lr_sched: tf.keras.optimizers.schedules .LearningRateSchedule
+ after_warmup_lr_sched: tf_keras.optimizers.schedules .LearningRateSchedule
or a constant.
warmup_steps: Number of the warmup steps.
warmup_learning_rate: Initial learning rate for the warmup.
@@ -122,7 +129,7 @@ def __init__(self,
self._warmup_steps = warmup_steps
self._init_warmup_lr = warmup_learning_rate
if isinstance(after_warmup_lr_sched,
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
self._final_warmup_lr = after_warmup_lr_sched(warmup_steps)
else:
self._final_warmup_lr = tf.cast(after_warmup_lr_sched, dtype=tf.float32)
@@ -136,7 +143,7 @@ def __call__(self, step: int):
(self._final_warmup_lr - self._init_warmup_lr))
if isinstance(self._after_warmup_lr_sched,
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
after_warmup_lr = self._after_warmup_lr_sched(step)
else:
after_warmup_lr = tf.cast(self._after_warmup_lr_sched, dtype=tf.float32)
@@ -148,13 +155,13 @@ def __call__(self, step: int):
def get_config(self) -> Mapping[str, Any]:
if isinstance(self._after_warmup_lr_sched,
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
config = {
"after_warmup_lr_sched": self._after_warmup_lr_sched.get_config()} # pytype: disable=attribute-error
else:
config = {"after_warmup_lr_sched": self._after_warmup_lr_sched} # pytype: disable=attribute-error
- config.update({
+ config.update({ # pyrefly: ignore[no-matching-overload]
"warmup_steps": self._warmup_steps,
"warmup_learning_rate": self._init_warmup_lr,
"name": self._name
@@ -162,18 +169,18 @@ def get_config(self) -> Mapping[str, Any]:
return config
-class PolynomialWarmUp(tf.keras.optimizers.schedules.LearningRateSchedule):
+class PolynomialWarmUp(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Applies polynomial warmup schedule on a given learning rate decay schedule."""
def __init__(self,
after_warmup_lr_sched: Union[
- tf.keras.optimizers.schedules.LearningRateSchedule, float],
+ tf_keras.optimizers.schedules.LearningRateSchedule, float],
warmup_steps: int,
power: float = 1.0,
name: str = "PolynomialWarmup"):
super().__init__()
if isinstance(after_warmup_lr_sched,
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
self._initial_learning_rate = after_warmup_lr_sched(warmup_steps)
else:
self._initial_learning_rate = tf.cast(
@@ -203,7 +210,7 @@ def __call__(self, step):
tf.math.pow(warmup_percent_done, self._power))
if isinstance(self._after_warmup_lr_sched,
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
after_warmup_lr = self._after_warmup_lr_sched(step)
else:
after_warmup_lr = tf.cast(self._after_warmup_lr_sched, dtype=tf.float32)
@@ -216,13 +223,13 @@ def __call__(self, step):
def get_config(self) -> Mapping[str, Any]:
if isinstance(self._after_warmup_lr_sched,
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
config = {
"after_warmup_lr_sched": self._after_warmup_lr_sched.get_config()} # pytype: disable=attribute-error
else:
config = {"after_warmup_lr_sched": self._after_warmup_lr_sched} # pytype: disable=attribute-error
- config.update({
+ config.update({ # pyrefly: ignore[no-matching-overload]
"warmup_steps": self._warmup_steps,
"power": self._power,
"name": self._name
@@ -230,7 +237,7 @@ def get_config(self) -> Mapping[str, Any]:
return config
-class DirectPowerDecay(tf.keras.optimizers.schedules.LearningRateSchedule):
+class DirectPowerDecay(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Learning rate schedule follows lr * (step)^power."""
def __init__(self,
@@ -267,7 +274,7 @@ def get_config(self):
}
-class PowerAndLinearDecay(tf.keras.optimizers.schedules.LearningRateSchedule):
+class PowerAndLinearDecay(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Learning rate schedule with multiplied by linear decay at the end.
The schedule has the following behavoir.
@@ -334,7 +341,7 @@ def get_config(self):
}
-class PowerDecayWithOffset(tf.keras.optimizers.schedules.LearningRateSchedule):
+class PowerDecayWithOffset(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Power learning rate decay with offset.
Learning rate equals to `pre_offset_learning_rate` if `step` < `offset`.
@@ -387,7 +394,7 @@ def get_config(self):
class StepCosineDecayWithOffset(
- tf.keras.optimizers.schedules.LearningRateSchedule):
+ tf_keras.optimizers.schedules.LearningRateSchedule):
"""Stepwise cosine learning rate decay with offset.
Learning rate is equivalent to one or more cosine decay(s) starting and
@@ -460,10 +467,6 @@ def __call__(self, global_step):
tf.constant(math.pi) * (global_step) /
(init_total_steps)) + 1.0) / 2.0 + next_init_lr)
learning_rate = cosine_learning_rate
- tf.compat.v1.logging.info("DEBUG lr %r next lr %r", learning_rate,
- cosine_learning_rate)
- tf.compat.v1.logging.info("DEBUG lr %r next lr %r inittotalstep %r",
- init_lr, next_init_lr, init_total_steps)
for i in range(1, num_levels):
next_init_lr = lr_levels[i]
@@ -471,9 +474,6 @@ def __call__(self, global_step):
next_total_steps = level_total_steps[i]
next_next_init_lr = lr_levels[i + 1] if num_levels > i + 1 else 0.
- tf.compat.v1.logging.info(
- "DEBUG step %r nilr %r nss %r nts %r nnilr %r", global_step,
- next_init_lr, next_start_step, next_total_steps, next_next_init_lr)
next_cosine_learning_rate = ((next_init_lr - next_next_init_lr) *
(tf.cos(
tf.constant(math.pi) *
@@ -482,8 +482,6 @@ def __call__(self, global_step):
next_next_init_lr)
learning_rate = tf.where(global_step >= next_start_step,
next_cosine_learning_rate, learning_rate)
- tf.compat.v1.logging.info("DEBUG lr %r next lr %r", learning_rate,
- next_cosine_learning_rate)
return learning_rate
diff --git a/official/modeling/optimization/lr_schedule_test.py b/official/modeling/optimization/lr_schedule_test.py
index df74db692eb..aa996d1e818 100644
--- a/official/modeling/optimization/lr_schedule_test.py
+++ b/official/modeling/optimization/lr_schedule_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for lr_schedule."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.optimization import lr_schedule
@@ -77,6 +77,7 @@ class OffsetLearningRateTest(tf.test.TestCase, parameterized.TestCase):
dict(class_name=lr_schedule.PolynomialDecayWithOffset),
dict(class_name=lr_schedule.ExponentialDecayWithOffset),
dict(class_name=lr_schedule.CosineDecayWithOffset),
+ dict(class_name=lr_schedule.CosineDecayRestartsWithOffset),
)
def test_generated_docstring(self, class_name):
self.assertNotEmpty(class_name.__init__.__doc__)
@@ -104,6 +105,38 @@ def test_offset(self, class_name, kwarg):
for step in range(10, 101, 10):
self.assertEqual(offset_lr(step), base_lr(step - offset))
+ @parameterized.named_parameters(
+ dict(testcase_name='period_1', start=5, end=15, offset_shift=0),
+ dict(testcase_name='period_2', start=15, end=25, offset_shift=10),
+ dict(testcase_name='period_3', start=25, end=35, offset_shift=20),
+ )
+ def test_cosine_decay_restarts_with_offset(self, start, end, offset_shift):
+ offset = 5
+ m_mul = 0.5
+ first_decay_steps = 10
+
+ restarts_lr = lr_schedule.CosineDecayRestartsWithOffset(
+ offset=offset,
+ initial_learning_rate=1.0,
+ first_decay_steps=first_decay_steps,
+ t_mul=1.0,
+ m_mul=m_mul,
+ )
+
+ # The equivalent cosine decay with offset.
+ period_start = offset + offset_shift
+ initial_learning_rate = m_mul ** (period_start // first_decay_steps)
+ decay_lr = lr_schedule.CosineDecayWithOffset(
+ offset=period_start,
+ initial_learning_rate=initial_learning_rate,
+ decay_steps=first_decay_steps,
+ )
+
+ for step in range(start, end):
+ self.assertAlmostEqual(
+ restarts_lr(step).numpy(), decay_lr(step).numpy(), places=5
+ )
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/modeling/optimization/optimizer_factory.py b/official/modeling/optimization/optimizer_factory.py
index 1b57d10b83f..48705cdfcdf 100644
--- a/official/modeling/optimization/optimizer_factory.py
+++ b/official/modeling/optimization/optimizer_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,39 +13,55 @@
# limitations under the License.
"""Optimizer factory class."""
-from typing import Callable, Optional, Union, List, Tuple
+from typing import Callable, List, Optional, Tuple, Union
import gin
-import tensorflow as tf
-import tensorflow_addons.optimizers as tfa_optimizers
+import tensorflow as tf, tf_keras
from official.modeling.optimization import slide_optimizer
from official.modeling.optimization import adafactor_optimizer
from official.modeling.optimization import ema_optimizer
-from official.modeling.optimization import lars_optimizer
+from official.modeling.optimization import lamb
+from official.modeling.optimization import lars
from official.modeling.optimization import legacy_adamw
from official.modeling.optimization import lr_schedule
from official.modeling.optimization.configs import optimization_config as opt_cfg
-OPTIMIZERS_CLS = {
- 'sgd': tf.keras.optimizers.SGD,
- # TODO(chenmoneygithub): experimental.SGD
- 'adam': tf.keras.optimizers.Adam,
- # TODO(chenmoneygithub): experimental.Adam
+# Optimizer CLS to be used in both legacy and new path.
+SHARED_OPTIMIZERS = {
+ 'sgd_experimental': tf_keras.optimizers.experimental.SGD,
+ 'adam_experimental': tf_keras.optimizers.experimental.Adam,
'adamw': legacy_adamw.AdamWeightDecay,
- 'lamb': tfa_optimizers.LAMB,
- 'rmsprop': tf.keras.optimizers.RMSprop,
- 'lars': lars_optimizer.LARS,
- 'adagrad': tf.keras.optimizers.Adagrad,
+ 'adamw_experimental': tf_keras.optimizers.experimental.AdamW,
+ 'lamb': lamb.LAMB,
+ 'lars': lars.LARS,
'slide': slide_optimizer.SLIDE,
'adafactor': adafactor_optimizer.Adafactor,
+ 'adafactor_keras': tf_keras.optimizers.Adafactor,
}
+LEGACY_OPTIMIZERS_CLS = {
+ 'sgd': tf_keras.optimizers.legacy.SGD,
+ 'adam': tf_keras.optimizers.legacy.Adam,
+ 'rmsprop': tf_keras.optimizers.legacy.RMSprop,
+ 'adagrad': tf_keras.optimizers.legacy.Adagrad,
+}
+LEGACY_OPTIMIZERS_CLS.update(SHARED_OPTIMIZERS)
+
+NEW_OPTIMIZERS_CLS = {
+ 'sgd': tf_keras.optimizers.experimental.SGD,
+ 'adam': tf_keras.optimizers.experimental.Adam,
+ 'rmsprop': tf_keras.optimizers.experimental.RMSprop,
+ 'adagrad': tf_keras.optimizers.experimental.Adagrad,
+}
+NEW_OPTIMIZERS_CLS.update(SHARED_OPTIMIZERS) # pyrefly: ignore[no-matching-overload]
+
LR_CLS = {
'stepwise': lr_schedule.PiecewiseConstantDecayWithOffset,
'polynomial': lr_schedule.PolynomialDecayWithOffset,
'exponential': lr_schedule.ExponentialDecayWithOffset,
'cosine': lr_schedule.CosineDecayWithOffset,
+ 'cosine_restarts': lr_schedule.CosineDecayRestartsWithOffset,
'power': lr_schedule.DirectPowerDecay,
'power_linear': lr_schedule.PowerAndLinearDecay,
'power_with_offset': lr_schedule.PowerDecayWithOffset,
@@ -59,7 +75,12 @@
def register_optimizer_cls(key: str,
- optimizer_config_cls: tf.keras.optimizers.Optimizer):
+ optimizer_config_cls: Union[
+ tf_keras.optimizers.Optimizer,
+ tf_keras.optimizers.legacy.Optimizer,
+ tf_keras.optimizers.experimental.Optimizer
+ ],
+ use_legacy_optimizer: bool = True):
"""Register customize optimizer cls.
The user will still need to subclass data classes in
@@ -67,11 +88,17 @@ def register_optimizer_cls(key: str,
Args:
key: A string to that the optimizer_config_cls is registered with.
- optimizer_config_cls: A class which inherits tf.keras.optimizers.Optimizer.
+ optimizer_config_cls: A class which inherits tf_keras.optimizers.Optimizer.
+ use_legacy_optimizer: A boolean that indicates if using legacy optimizers.
"""
- if key in OPTIMIZERS_CLS:
- raise ValueError('%s already registered in OPTIMIZER_CLS.' % key)
- OPTIMIZERS_CLS[key] = optimizer_config_cls
+ if use_legacy_optimizer:
+ if key in LEGACY_OPTIMIZERS_CLS:
+ raise ValueError('%s already registered in LEGACY_OPTIMIZERS_CLS.' % key)
+ LEGACY_OPTIMIZERS_CLS[key] = optimizer_config_cls
+ else:
+ if key in NEW_OPTIMIZERS_CLS:
+ raise ValueError('%s already registered in NEW_OPTIMIZERS_CLS.' % key)
+ NEW_OPTIMIZERS_CLS[key] = optimizer_config_cls
class OptimizerFactory:
@@ -86,6 +113,8 @@ class OptimizerFactory:
(4) Build optimizer.
This is a typical example for using this class:
+
+ ```
params = {
'optimizer': {
'type': 'sgd',
@@ -105,6 +134,7 @@ class OptimizerFactory:
opt_factory = OptimizerFactory(opt_config)
lr = opt_factory.build_learning_rate()
optimizer = opt_factory.build_optimizer(lr)
+ ```
"""
def __init__(self, config: opt_cfg.OptimizationConfig):
@@ -140,7 +170,7 @@ def build_learning_rate(self):
lr_config.learning_rate is returned.
Returns:
- tf.keras.optimizers.schedules.LearningRateSchedule instance. If
+ tf_keras.optimizers.schedules.LearningRateSchedule instance. If
learning rate type is consant, lr_config.learning_rate is returned.
"""
if self._lr_type == 'constant':
@@ -156,15 +186,16 @@ def build_learning_rate(self):
@gin.configurable
def build_optimizer(
self,
- lr: Union[tf.keras.optimizers.schedules.LearningRateSchedule, float],
+ lr: Union[tf_keras.optimizers.schedules.LearningRateSchedule, float],
gradient_aggregator: Optional[Callable[
[List[Tuple[tf.Tensor, tf.Tensor]]], List[Tuple[tf.Tensor,
tf.Tensor]]]] = None,
gradient_transformers: Optional[List[Callable[
[List[Tuple[tf.Tensor, tf.Tensor]]], List[Tuple[tf.Tensor,
tf.Tensor]]]]] = None,
- postprocessor: Optional[Callable[[tf.keras.optimizers.Optimizer],
- tf.keras.optimizers.Optimizer]] = None):
+ postprocessor: Optional[Callable[[tf_keras.optimizers.Optimizer],
+ tf_keras.optimizers.Optimizer]] = None,
+ use_legacy_optimizer: bool = True):
"""Build optimizer.
Builds optimizer from config. It takes learning rate as input, and builds
@@ -173,7 +204,7 @@ def build_optimizer(
Args:
lr: A floating point value, or a
- tf.keras.optimizers.schedules.LearningRateSchedule instance.
+ tf_keras.optimizers.schedules.LearningRateSchedule instance.
gradient_aggregator: Optional function to overwrite gradient aggregation.
gradient_transformers: Optional list of functions to use to transform
gradients before applying updates to Variables. The functions are
@@ -182,10 +213,11 @@ def build_optimizer(
global_clipnorm should not be set when gradient_transformers is passed.
postprocessor: An optional function for postprocessing the optimizer. It
takes an optimizer and returns an optimizer.
+ use_legacy_optimizer: A boolean that indicates if using legacy optimizers.
Returns:
- `tf.keras.optimizers.Optimizer` or
- `tf.keras.optimizers.experimental.Optimizer` instance.
+ `tf_keras.optimizers.legacy.Optimizer` or
+ `tf_keras.optimizers.experimental.Optimizer` instance.
"""
optimizer_dict = self._optimizer_config.as_dict()
@@ -203,17 +235,34 @@ def build_optimizer(
if gradient_transformers is not None:
optimizer_dict['gradient_transformers'] = gradient_transformers
- optimizer = OPTIMIZERS_CLS[self._optimizer_type](**optimizer_dict)
+ if use_legacy_optimizer:
+ optimizer = LEGACY_OPTIMIZERS_CLS[self._optimizer_type](**optimizer_dict)
+ else:
+ if 'decay' in optimizer_dict:
+ raise ValueError(
+ '`decay` is deprecated in new Keras optimizer, please reflect the '
+ 'decay logic in `lr` or set `use_legacy_optimizer=True` to use the '
+ 'legacy optimizer.')
+ optimizer = NEW_OPTIMIZERS_CLS[self._optimizer_type](**optimizer_dict)
if self._use_ema:
+ if not use_legacy_optimizer:
+ raise ValueError(
+ 'EMA can only work with the legacy optimizer, please set '
+ '`use_legacy_optimizer=True`.')
optimizer = ema_optimizer.ExponentialMovingAverage(
optimizer, **self._ema_config.as_dict())
if postprocessor:
optimizer = postprocessor(optimizer)
- assert isinstance(
- optimizer, (tf.keras.optimizers.Optimizer,
- tf.keras.optimizers.experimental.Optimizer)
- ), ('OptimizerFactory.build_optimizer returning a non-optimizer object: '
- '{}'.format(optimizer))
-
- return optimizer
+ if isinstance(optimizer, tf_keras.optimizers.Optimizer):
+ return optimizer
+ # The following check makes sure the function won't break in older TF
+ # version because of missing the experimental/legacy package.
+ if hasattr(tf_keras.optimizers, 'experimental'):
+ if isinstance(optimizer, tf_keras.optimizers.experimental.Optimizer):
+ return optimizer
+ if hasattr(tf_keras.optimizers, 'legacy'):
+ if isinstance(optimizer, tf_keras.optimizers.legacy.Optimizer):
+ return optimizer
+ raise TypeError('OptimizerFactory.build_optimizer returning a '
+ 'non-optimizer object: {}'.format(optimizer))
diff --git a/official/modeling/optimization/optimizer_factory_test.py b/official/modeling/optimization/optimizer_factory_test.py
index d513a15aba9..a7c0d5a8d16 100644
--- a/official/modeling/optimization/optimizer_factory_test.py
+++ b/official/modeling/optimization/optimizer_factory_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,9 +13,10 @@
# limitations under the License.
"""Tests for optimizer_factory.py."""
+import math
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.optimization import optimizer_factory
from official.modeling.optimization.configs import optimization_config
@@ -37,7 +38,7 @@ def test_optimizers(self, optimizer_type):
}
}
}
- optimizer_cls = optimizer_factory.OPTIMIZERS_CLS[optimizer_type]
+ optimizer_cls = optimizer_factory.LEGACY_OPTIMIZERS_CLS[optimizer_type]
expected_optimizer_config = optimizer_cls().get_config()
expected_optimizer_config['learning_rate'] = 0.1
@@ -49,6 +50,36 @@ def test_optimizers(self, optimizer_type):
self.assertIsInstance(optimizer, optimizer_cls)
self.assertEqual(expected_optimizer_config, optimizer.get_config())
+ @parameterized.parameters(('sgd'), ('rmsprop'), ('adam'), ('adamw'), ('lamb'),
+ ('lars'), ('adagrad'))
+ def test_new_optimizers(self, optimizer_type):
+ params = {
+ 'optimizer': {
+ 'type': optimizer_type
+ },
+ 'learning_rate': {
+ 'type': 'constant',
+ 'constant': {
+ 'learning_rate': 0.1
+ }
+ }
+ }
+ optimizer_cls = optimizer_factory.NEW_OPTIMIZERS_CLS[optimizer_type]
+ expected_optimizer_config = optimizer_cls().get_config()
+ expected_optimizer_config['learning_rate'] = 0.1
+
+ opt_config = optimization_config.OptimizationConfig(params)
+ if optimizer_type == 'sgd':
+ # Delete unsupported arg `decay` from SGDConfig.
+ delattr(opt_config.optimizer.sgd, 'decay')
+ opt_factory = optimizer_factory.OptimizerFactory(opt_config)
+ lr = opt_factory.build_learning_rate()
+ optimizer = opt_factory.build_optimizer(
+ lr, postprocessor=lambda x: x, use_legacy_optimizer=False)
+
+ self.assertIsInstance(optimizer, optimizer_cls)
+ self.assertEqual(expected_optimizer_config, optimizer.get_config())
+
def test_gradient_aggregator(self):
params = {
'optimizer': {
@@ -69,6 +100,9 @@ def test_gradient_aggregator(self):
zero_grads = lambda gv: [(tf.zeros_like(g), v) for g, v in gv]
optimizer = opt_factory.build_optimizer(lr, gradient_aggregator=zero_grads)
+ if isinstance(optimizer, tf_keras.optimizers.experimental.Optimizer):
+ self.skipTest('New Keras optimizer does not support '
+ '`gradient_aggregator` arg.')
var0 = tf.Variable([1.0, 2.0])
var1 = tf.Variable([3.0, 4.0])
@@ -140,6 +174,25 @@ def test_missing_types(self):
optimizer_factory.OptimizerFactory(
optimization_config.OptimizationConfig(params))
+ def test_wrong_return_type(self):
+ optimizer_type = 'sgd'
+ params = {
+ 'optimizer': {
+ 'type': optimizer_type
+ },
+ 'learning_rate': {
+ 'type': 'constant',
+ 'constant': {
+ 'learning_rate': 0.1
+ }
+ }
+ }
+
+ opt_config = optimization_config.OptimizationConfig(params)
+ opt_factory = optimizer_factory.OptimizerFactory(opt_config)
+ with self.assertRaises(TypeError):
+ _ = opt_factory.build_optimizer(0.1, postprocessor=lambda x: None)
+
# TODO(b/187559334) refactor lr_schedule tests into `lr_schedule_test.py`.
@@ -284,6 +337,41 @@ def test_cosine_lr_schedule(self):
for step, value in expected_lr_step_values:
self.assertAlmostEqual(lr(step).numpy(), value)
+ def test_cosine_restarts_lr_schedule(self):
+ params = {
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine_restarts',
+ 'cosine_restarts': {
+ 'initial_learning_rate': 0.1,
+ 'first_decay_steps': 1000,
+ 't_mul': 0.75,
+ 'm_mul': 0.5
+ }
+ }
+ }
+ expected_lr_step_values = [
+ # Period 1: length 1000.
+ [0, 0.1],
+ [250, 0.1 * (1.0 + math.cos(math.pi * 250.0 / 1000.0)) / 2.0],
+ [500, 0.1 * (1.0 + math.cos(math.pi * 500.0 / 1000.0)) / 2.0],
+ [750, 0.1 * (1.0 + math.cos(math.pi * 750.0 / 1000.0)) / 2.0],
+ # Period 2: length 750, starts at 1000 with init_lr = 0.1 * 0.5 = 0.05.
+ [1000, 0.05],
+ [1250, 0.05 * (1.0 + math.cos(math.pi * 250.0 / 750.0)) / 2.0],
+ ]
+ opt_config = optimization_config.OptimizationConfig(params)
+ opt_factory = optimizer_factory.OptimizerFactory(opt_config)
+ lr = opt_factory.build_learning_rate()
+
+ for step, value in expected_lr_step_values:
+ self.assertAlmostEqual(lr(step).numpy(), value)
+
def test_constant_lr_with_warmup_schedule(self):
params = {
'optimizer': {
@@ -469,7 +557,7 @@ class MyClass():
pass
optimizer_factory.register_optimizer_cls('test', MyClass)
- self.assertIn('test', optimizer_factory.OPTIMIZERS_CLS)
+ self.assertIn('test', optimizer_factory.LEGACY_OPTIMIZERS_CLS)
with self.assertRaisesRegex(ValueError, 'test already registered.*'):
optimizer_factory.register_optimizer_cls('test', MyClass)
diff --git a/official/modeling/optimization/slide_optimizer.py b/official/modeling/optimization/slide_optimizer.py
index 8bbd4687461..500f2875eed 100644
--- a/official/modeling/optimization/slide_optimizer.py
+++ b/official/modeling/optimization/slide_optimizer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/performance.py b/official/modeling/performance.py
index 3c6f6d15a56..17f8485a11f 100644
--- a/official/modeling/performance.py
+++ b/official/modeling/performance.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Functions and classes related to training performance."""
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def configure_optimizer(optimizer,
@@ -29,25 +29,25 @@ def configure_optimizer(optimizer,
del use_graph_rewrite
if use_float16:
if loss_scale in (None, 'dynamic'):
- optimizer = tf.keras.mixed_precision.LossScaleOptimizer(optimizer)
+ optimizer = tf_keras.mixed_precision.LossScaleOptimizer(optimizer)
else:
# loss_scale is a number. We interpret that as a fixed loss scale.
- optimizer = tf.keras.mixed_precision.LossScaleOptimizer(
+ optimizer = tf_keras.mixed_precision.LossScaleOptimizer(
optimizer, dynamic=False, initial_scale=loss_scale)
return optimizer
def set_mixed_precision_policy(dtype, loss_scale=None):
- """Sets the global `tf.keras.mixed_precision.Policy`."""
+ """Sets the global `tf_keras.mixed_precision.Policy`."""
# TODO(b/191894773): Remove loss_scale argument
assert loss_scale is None, (
'The loss_scale argument must be None. The argument exists for '
'historical reasons and will be removed soon.')
if dtype == tf.float16:
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
elif dtype == tf.bfloat16:
- tf.keras.mixed_precision.set_global_policy('mixed_bfloat16')
+ tf_keras.mixed_precision.set_global_policy('mixed_bfloat16')
elif dtype == tf.float32:
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
else:
raise ValueError('Unexpected dtype: %s' % dtype)
diff --git a/official/modeling/privacy/__init__.py b/official/modeling/privacy/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/modeling/privacy/__init__.py
+++ b/official/modeling/privacy/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/privacy/configs.py b/official/modeling/privacy/configs.py
index c8d4692a563..cc5502fd301 100644
--- a/official/modeling/privacy/configs.py
+++ b/official/modeling/privacy/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/modeling/privacy/configs_test.py b/official/modeling/privacy/configs_test.py
index 485e4e5c4ac..6c882aafe08 100644
--- a/official/modeling/privacy/configs_test.py
+++ b/official/modeling/privacy/configs_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for configs."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.privacy import configs
diff --git a/official/modeling/privacy/ops.py b/official/modeling/privacy/ops.py
index 8b024702085..631bb354068 100644
--- a/official/modeling/privacy/ops.py
+++ b/official/modeling/privacy/ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,12 +15,24 @@
"""Ops for differential privacy (gradient) transforms."""
from typing import List, Tuple
-import tensorflow as tf
+import warnings
+
+import tensorflow as tf, tf_keras
def clip_l2_norm(grads_vars: List[Tuple[tf.Tensor, tf.Tensor]],
l2_norm_clip: float) -> List[Tuple[tf.Tensor, tf.Tensor]]:
- """Clip gradients by global norm."""
+ """DEPRECATED Clip gradients by global norm.
+
+ Args:
+ grads_vars: List of tuple of gradient and its corresponding variables
+ l2_norm_clip: Float for differential privacy norm
+
+ Returns:
+ List of clipped gradients and its corresponding variables
+ """
+ warnings.warn("`clip_l2_norm` deprecated.",
+ DeprecationWarning)
gradients = []
variables = []
@@ -33,10 +45,19 @@ def clip_l2_norm(grads_vars: List[Tuple[tf.Tensor, tf.Tensor]],
def add_noise(grads_vars: List[Tuple[tf.Tensor, tf.Tensor]],
noise_stddev: float) -> List[Tuple[tf.Tensor, tf.Tensor]]:
- """Add noise to gradients."""
+ """DEPRECATED Add noise to gradients.
+
+ Args:
+ grads_vars: List of tuple of gradient and its corresponding variables
+ noise_stddev: Noise multiplier
+
+ Returns:
+ List of noised gradients and its corresponding variables
+ """
+ warnings.warn("`add_noise` deprecated.", DeprecationWarning)
+
ret = []
for (g, v) in grads_vars:
noise = tf.random.normal(tf.shape(g), stddev=noise_stddev)
ret.append((g + noise, v))
return ret
-
diff --git a/official/modeling/privacy/ops_test.py b/official/modeling/privacy/ops_test.py
index 4f5d580c75f..3d1369da9e4 100644
--- a/official/modeling/privacy/ops_test.py
+++ b/official/modeling/privacy/ops_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from unittest import mock
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.privacy import ops
diff --git a/official/modeling/tf_utils.py b/official/modeling/tf_utils.py
index bff4d52c14b..1c91457f859 100644
--- a/official/modeling/tf_utils.py
+++ b/official/modeling/tf_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,8 +14,10 @@
"""Common TF utilities."""
+import functools
+import inspect
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.util import deprecation
from official.modeling import activations
@@ -23,7 +25,7 @@
@deprecation.deprecated(
None,
- "tf.keras.layers.Layer supports multiple positional args and kwargs as "
+ "tf_keras.layers.Layer supports multiple positional args and kwargs as "
"input tensors. pack/unpack inputs to override __call__ is no longer "
"needed.")
def pack_inputs(inputs):
@@ -48,7 +50,7 @@ def pack_inputs(inputs):
@deprecation.deprecated(
None,
- "tf.keras.layers.Layer supports multiple positional args and kwargs as "
+ "tf_keras.layers.Layer supports multiple positional args and kwargs as "
"input tensors. pack/unpack inputs to override __call__ is no longer "
"needed.")
def unpack_inputs(inputs):
@@ -82,19 +84,22 @@ def is_special_none_tensor(tensor):
return tensor.shape.ndims == 0 and tensor.dtype == tf.int32
-def get_activation(identifier, use_keras_layer=False):
- """Maps a identifier to a Python function, e.g., "relu" => `tf.nn.relu`.
+def get_activation(identifier, use_keras_layer=False, **kwargs):
+ """Maps an identifier to a Python function, e.g., "relu" => `tf.nn.relu`.
It checks string first and if it is one of customized activation not in TF,
the corresponding activation will be returned. For non-customized activation
- names and callable identifiers, always fallback to tf.keras.activations.get.
+ names and callable identifiers, always fallback to tf_keras.activations.get.
Prefers using keras layers when use_keras_layer=True. Now it only supports
- 'relu', 'linear', 'identity', 'swish'.
+ 'relu', 'linear', 'identity', 'swish', 'mish', 'leaky_relu', and 'gelu'.
Args:
identifier: String name of the activation function or callable.
use_keras_layer: If True, use keras layer if identifier is allow-listed.
+ **kwargs: Keyword arguments to use to instantiate an activation function.
+ Available only for 'leaky_relu' and 'gelu' when using keras layers.
+ For example: get_activation('leaky_relu', use_keras_layer=True, alpha=0.1)
Returns:
A Python function corresponding to the activation function or a keras
@@ -110,11 +115,15 @@ def get_activation(identifier, use_keras_layer=False):
"swish": "swish",
"sigmoid": "sigmoid",
"relu6": tf.nn.relu6,
+ "leaky_relu": functools.partial(tf.nn.leaky_relu, **kwargs),
"hard_swish": activations.hard_swish,
"hard_sigmoid": activations.hard_sigmoid,
+ "mish": activations.mish,
+ "gelu": functools.partial(tf.nn.gelu, **kwargs),
+ "squared_relu": activations.squared_relu,
}
if identifier in keras_layer_allowlist:
- return tf.keras.layers.Activation(keras_layer_allowlist[identifier])
+ return tf_keras.layers.Activation(keras_layer_allowlist[identifier])
name_to_fn = {
"gelu": activations.gelu,
"simple_swish": activations.simple_swish,
@@ -122,10 +131,12 @@ def get_activation(identifier, use_keras_layer=False):
"relu6": activations.relu6,
"hard_sigmoid": activations.hard_sigmoid,
"identity": activations.identity,
+ "mish": activations.mish,
+ "squared_relu": activations.squared_relu,
}
if identifier in name_to_fn:
- return tf.keras.activations.get(name_to_fn[identifier])
- return tf.keras.activations.get(identifier)
+ return tf_keras.activations.get(name_to_fn[identifier])
+ return tf_keras.activations.get(identifier)
def get_shape_list(tensor, expected_rank=None, name=None):
@@ -272,3 +283,92 @@ def cross_replica_concat(value, axis, name="cross_replica_concat"):
if value.shape.as_list()[0] is None:
raise RuntimeError(f"{value} has unknown batch.")
return context.all_gather(value, axis=axis)
+
+
+def clone_initializer(initializer):
+ # Keras initializer is going to be stateless, which mean reusing the same
+ # initializer will produce same init value when the shapes are the same.
+ if isinstance(initializer, tf_keras.initializers.Initializer):
+ return initializer.__class__.from_config(initializer.get_config())
+ # When the input is string/dict or other serialized configs, caller will
+ # create a new keras Initializer instance based on that, and we don't need to
+ # do anything
+ return initializer
+
+
+def serialize_keras_object(obj):
+ if hasattr(tf_keras.utils, "legacy"):
+ return tf_keras.utils.legacy.serialize_keras_object(obj)
+ else:
+ return tf_keras.utils.serialize_keras_object(obj)
+
+
+def deserialize_keras_object(
+ config, module_objects=None, custom_objects=None, printable_module_name=None
+):
+ if hasattr(tf_keras.utils, "legacy"):
+ return tf_keras.utils.legacy.deserialize_keras_object(
+ config, custom_objects, module_objects, printable_module_name
+ )
+ else:
+ return tf_keras.utils.deserialize_keras_object(
+ config, custom_objects, module_objects, printable_module_name # pyrefly: ignore[bad-argument-count]
+ )
+
+
+def serialize_layer(layer, use_legacy_format=False):
+ if (
+ "use_legacy_format"
+ in inspect.getfullargspec(tf_keras.layers.serialize).args
+ ):
+ return tf_keras.layers.serialize(layer, use_legacy_format=use_legacy_format)
+ else:
+ return tf_keras.layers.serialize(layer)
+
+
+def serialize_initializer(initializer, use_legacy_format=False):
+ if (
+ "use_legacy_format"
+ in inspect.getfullargspec(tf_keras.initializers.serialize).args
+ ):
+ return tf_keras.initializers.serialize(
+ initializer, use_legacy_format=use_legacy_format
+ )
+ else:
+ return tf_keras.initializers.serialize(initializer)
+
+
+def serialize_regularizer(regularizer, use_legacy_format=False):
+ if (
+ "use_legacy_format"
+ in inspect.getfullargspec(tf_keras.regularizers.serialize).args
+ ):
+ return tf_keras.regularizers.serialize(
+ regularizer, use_legacy_format=use_legacy_format
+ )
+ else:
+ return tf_keras.regularizers.serialize(regularizer)
+
+
+def serialize_constraint(constraint, use_legacy_format=False):
+ if (
+ "use_legacy_format"
+ in inspect.getfullargspec(tf_keras.constraints.serialize).args
+ ):
+ return tf_keras.constraints.serialize(
+ constraint, use_legacy_format=use_legacy_format
+ )
+ else:
+ return tf_keras.constraints.serialize(constraint)
+
+
+def serialize_activation(activation, use_legacy_format=False):
+ if (
+ "use_legacy_format"
+ in inspect.getfullargspec(tf_keras.activations.serialize).args
+ ):
+ return tf_keras.activations.serialize(
+ activation, use_legacy_format=use_legacy_format
+ )
+ else:
+ return tf_keras.activations.serialize(activation)
diff --git a/official/modeling/tf_utils_test.py b/official/modeling/tf_utils_test.py
index 0b6a1662cd9..88c3a7c4d60 100644
--- a/official/modeling/tf_utils_test.py
+++ b/official/modeling/tf_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for tf_utils."""
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -26,7 +26,8 @@ def all_strategy_combinations():
return combinations.combine(
strategy=[
strategy_combinations.cloud_tpu_strategy,
- strategy_combinations.mirrored_strategy_with_two_gpus,
+ # TODO(b/285797201):disable multi-gpu tests due to hanging.
+ # strategy_combinations.mirrored_strategy_with_two_gpus,
],
mode='eager',
)
@@ -84,6 +85,24 @@ def function():
for gradient in per_replica_gradients.values:
self.assertAllClose(gradient, num_cores * tf.ones(shape))
+ @parameterized.parameters(('relu', True), ('relu', False),
+ ('leaky_relu', False), ('leaky_relu', True),
+ ('mish', True), ('mish', False), ('gelu', True))
+ def test_get_activations(self, name, use_keras_layer):
+ fn = tf_utils.get_activation(name, use_keras_layer)
+ self.assertIsNotNone(fn)
+
+ @combinations.generate(all_strategy_combinations())
+ def test_get_leaky_relu_layer(self, strategy):
+ @tf.function
+ def forward(x):
+ fn = tf_utils.get_activation(
+ 'leaky_relu', use_keras_layer=True, alpha=0.1)
+ return strategy.run(fn, args=(x,)).values[0]
+
+ got = forward(tf.constant([-1]))
+ self.assertAllClose(got, tf.constant([-0.1]))
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/nightly_requirements.txt b/official/nightly_requirements.txt
new file mode 100644
index 00000000000..aac74f52c7c
--- /dev/null
+++ b/official/nightly_requirements.txt
@@ -0,0 +1,31 @@
+six
+google-api-python-client>=1.6.7
+kaggle>=1.3.9
+numpy>=1.20
+oauth2client
+pandas>=0.22.0
+psutil>=5.4.3
+py-cpuinfo>=3.3.0
+scipy>=0.19.1
+tf-hub-nightly
+tensorflow-model-optimization>=0.4.1
+tensorflow-datasets
+tf-keras-nightly
+gin-config
+tf_slim>=1.1.0
+Cython
+matplotlib
+pyyaml
+# CV related dependencies
+opencv-python-headless
+Pillow
+pycocotools
+# NLP related dependencies
+seqeval
+sentencepiece
+sacrebleu
+# Projects/vit dependencies
+immutabledict
+# Fix CI
+wrapt>=1.15
+
diff --git a/official/nlp/MODEL_GARDEN.md b/official/nlp/MODEL_GARDEN.md
index 09294e942e8..c5f55bd6bc8 100644
--- a/official/nlp/MODEL_GARDEN.md
+++ b/official/nlp/MODEL_GARDEN.md
@@ -1,5 +1,4 @@
# TF-NLP Model Garden
-
## Introduction
The TF-NLP library provides a collection of scripts for training and
@@ -7,6 +6,15 @@ evaluating transformer-based models, on various tasks such as sentence
classification, question answering, and translation. Additionally, we provide
checkpoints of pretrained models which can be finetuned on downstream tasks.
+⚠️ Disclaimer: Checkpoints are based on training with publicly available datasets.
+Some datasets contain limitations, including non-commercial use limitations. Please review the terms and conditions made available by third parties before using
+the datasets provided. Checkpoints are licensed under
+[Apache 2.0](https://github.com/tensorflow/models/blob/master/LICENSE).
+
+⚠️ Disclaimer: Datasets hyperlinked from this page are not owned or distributed
+by Google. Such datasets are made available by third parties. Please review the
+terms and conditions made available by the third parties before using the data.
+
### How to Train Models
Model Garden can be easily installed with
@@ -18,7 +26,7 @@ on how to train models with this codebase.
By default, the experiment runs on GPUs. To run on TPUs, one should overwrite
`runtime.distribution_strategy` and set the tpu address. See [RuntimeConfig](https://github.com/tensorflow/models/blob/master/official/core/config_definitions.py) for details.
-In general, the experiments can run with the folloing command by setting the
+In general, the experiments can run with the following command by setting the
corresponding `${TASK}`, `${TASK_CONFIG}`, `${MODEL_CONFIG}`.
```
EXPERIMENT=???
@@ -64,7 +72,7 @@ Note that
[How to Train Models](https://github.com/tensorflow/models/blob/master/official/nlp/docs/train.md)
-[List of Pretrained Models for finetuning](https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md)
+[List of Pre-trained Models for finetuning](https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md)
[How to Publish Models](https://github.com/tensorflow/models/blob/master/official/nlp/docs/tfhub.md)
diff --git a/official/nlp/README.md b/official/nlp/README.md
index fcd1d2f77c5..e3425e8c075 100644
--- a/official/nlp/README.md
+++ b/official/nlp/README.md
@@ -1,64 +1,126 @@
# TF-NLP Model Garden
-⚠️ Disclaimer: All datasets hyperlinked from this page are not owned or
-distributed by Google. The dataset is made available by third parties.
-Please review the terms and conditions made available by the third parties
-before using the data.
+⚠️ Disclaimer: Datasets hyperlinked from this page are not owned or distributed
+by Google. Such datasets are made available by third parties. Please review the
+terms and conditions made available by the third parties before using the data.
-This codebase provides a Natrual Language Processing modeling toolkit written in
+This codebase provides a Natural Language Processing modeling toolkit written in
[TF2](https://www.tensorflow.org/guide/effective_tf2). It allows researchers and
developers to reproduce state-of-the-art model results and train custom models
to experiment new research ideas.
## Features
-* Reusable and modularized modeling building blocks
-* State-of-the-art reproducible
-* Easy to customize and extend
-* End-to-end training
-* Distributed trainable on both GPUs and TPUs
+* Reusable and modularized modeling building blocks
+* State-of-the-art reproducible
+* Easy to customize and extend
+* End-to-end training
+* Distributed trainable on both GPUs and TPUs
## Major components
### Libraries
We provide modeling library to allow users to train custom models for new
-research ideas. Detailed intructions can be found in READMEs in each folder.
+research ideas. Detailed instructions can be found in READMEs in each folder.
* [modeling/](modeling): modeling library that provides building blocks
(e.g.,Layers, Networks, and Models) that can be assembled into
- transformer-based achitectures .
+ transformer-based architectures.
* [data/](data): binaries and utils for input preprocessing, tokenization,
etc.
+### Layers
+
+Layers are the fundamental building blocks for NLP models. They can be used to
+assemble new `tf.keras` layers or models.
+
+| Layers |
+| ------------ |
+| [BertPackInputs](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/BertPackInputs) \| [BertTokenizer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/BertTokenizer) \| [BigBirdAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/BigBirdAttention) \| [BigBirdMasks](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/BigBirdMasks) \| [BlockDiagFeedforward](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/BlockDiagFeedforward) \| [CachedAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/CachedAttention) |
+| [ClassificationHead](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/ClassificationHead) \| [ExpertsChooseMaskedRouter](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/ExpertsChooseMaskedRouter) \| [FactorizedEmbedding](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/FactorizedEmbedding) \| [FastWordpieceBertTokenizer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/FastWordpieceBertTokenizer) |
+| [FeedForwardExperts](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/FeedForwardExperts) \| [FourierTransformLayer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/FourierTransformLayer) \| [GatedFeedforward](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/GatedFeedforward) \| [GaussianProcessClassificationHead](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/GaussianProcessClassificationHead) |
+| [HartleyTransformLayer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/HartleyTransformLayer) \| [KernelAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/KernelAttention) \| [KernelMask](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/KernelMask) \| [LinearTransformLayer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/LinearTransformLayer) \| [MaskedLM](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MaskedLM) \| [MaskedSoftmax](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MaskedSoftmax) |
+| [MatMulWithMargin](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MatMulWithMargin) \| [MixingMechanism](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MixingMechanism) \| [MobileBertEmbedding](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MobileBertEmbedding) \| [MobileBertMaskedLM](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MobileBertMaskedLM) |
+| [MobileBertTransformer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MobileBertTransformer) \| [MoeLayer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MoeLayer) \| [MoeLayerWithBackbone](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MoeLayerWithBackbone) \| [MultiChannelAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MultiChannelAttention) \| [MultiClsHeads](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MultiClsHeads) |
+| [MultiHeadRelativeAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/MultiHeadRelativeAttention) \| [OnDeviceEmbedding](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/OnDeviceEmbedding) \| [PackBertEmbeddings](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/PackBertEmbeddings) \| [PerDimScaleAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/PerDimScaleAttention) |
+| [PerQueryDenseHead](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/PerQueryDenseHead) \| [PositionEmbedding](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/PositionEmbedding) \| [RandomFeatureGaussianProcess](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/RandomFeatureGaussianProcess) \| [ReZeroTransformer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/ReZeroTransformer) |
+| [RelativePositionBias](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/RelativePositionBias) \| [RelativePositionEmbedding](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/RelativePositionEmbedding) \| [ReuseMultiHeadAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/ReuseMultiHeadAttention) \| [ReuseTransformer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/ReuseTransformer) |
+| [SelectTopK](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/SelectTopK) \| [SelfAttentionMask](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/SelfAttentionMask) \| [SentencepieceTokenizer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/SentencepieceTokenizer) \| [SpectralNormalization](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/SpectralNormalization) |
+| [SpectralNormalizationConv2D](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/SpectralNormalizationConv2D) \| [StridedTransformerEncoderBlock](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/StridedTransformerEncoderBlock) \| [StridedTransformerScaffold](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/StridedTransformerScaffold) |
+| [TNTransformerExpandCondense](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TNTransformerExpandCondense) \| [TalkingHeadsAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TalkingHeadsAttention) \| [TokenImportanceWithMovingAvg](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TokenImportanceWithMovingAvg) \| [Transformer](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/Transformer) |
+| [TransformerDecoderBlock](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TransformerDecoderBlock) \| [TransformerEncoderBlock](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TransformerEncoderBlock) \| [TransformerScaffold](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TransformerScaffold) \| [TransformerXL](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TransformerXL) |
+| [TransformerXLBlock](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TransformerXLBlock) \|[get_mask](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/get_mask) \|[TwoStreamRelativeAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/TwoStreamRelativeAttention) \| [VotingAttention](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/VotingAttention) \| [extract_gp_layer_kwargs](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/extract_gp_layer_kwargs) |
+| [extract_spec_norm_kwargs](https://www.tensorflow.org/api_docs/python/tfm/nlp/layers/extract_spec_norm_kwargs) |
+
+### Networks
+
+Networks are combinations of `tf.keras` layers (and possibly other networks).
+They are `tf.keras` models that would not be trained alone. It encapsulates
+common network structures like a transformer encoder into an easily handled
+object with a standardized configuration.
+
+| Networks |
+| -------------- |
+| [AlbertEncoder](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/AlbertEncoder) \| [BertEncoder](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/BertEncoder) \| [BertEncoderV2](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/BertEncoderV2) \| [Classification](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/Classification) \| [EncoderScaffold](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/EncoderScaffold) \| [FNet](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/FNet) \| [MobileBERTEncoder](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/MobileBERTEncoder) |
+| [FunnelTransformerEncoder](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/FunnelTransformerEncoder) \| [PackedSequenceEmbedding](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/PackedSequenceEmbedding) \| [SpanLabeling](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/SpanLabeling) \| [SparseMixer](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/SparseMixer) \| [XLNetBase](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/XLNetBase) |
+| [XLNetSpanLabeling](https://www.tensorflow.org/api_docs/python/tfm/nlp/networks/XLNetSpanLabeling) |
+
+### Models
+
+Models are combinations of `tf.keras` layers and models that can be trained.
+Several pre-built canned models are provided to train encoder networks. These
+models are intended as both convenience functions and canonical examples.
+
+| Models |
+| ------------ |
+| [BertClassifier](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/BertClassifier) \| [BertPretrainer](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/BertPretrainer) \| [BertPretrainerV2](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/BertPretrainerV2) \| [BertSpanLabeler](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/BertSpanLabeler) \| [BertTokenClassifier](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/BertTokenClassifier) \| [DualEncoder](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/DualEncoder) |
+| [ElectraPretrainer](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/ElectraPretrainer) \| [Seq2SeqTransformer](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/Seq2SeqTransformer) \| [T5Transformer](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/T5Transformer) \| [T5TransformerParams](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/T5TransformerParams) \| [TransformerDecoder](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/TransformerDecoder) |
+| [TransformerEncoder](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/TransformerEncoder) \| [XLNetClassifier](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/XLNetClassifier) \| [XLNetPretrainer](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/XLNetPretrainer) \| [XLNetSpanLabeler](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/XLNetSpanLabeler) \| [attention_initializer](https://www.tensorflow.org/api_docs/python/tfm/nlp/models/attention_initializer) |
+
+### Losses
+
+Losses contains common loss computation used in NLP tasks.
+
+| Losses |
+| ------------ |
+| [weighted_sparse_categorical_crossentropy_loss](https://www.tensorflow.org/api_docs/python/tfm/nlp/losses/weighted_sparse_categorical_crossentropy_loss) |
+
### State-of-the-Art models and examples
We provide SoTA model implementations, pre-trained models, training and
evaluation examples, and command lines. Detail instructions can be found in the
-READMEs for specific papers.
-
-1. [BERT](MODEL_GARDEN.md#available-model-configs): [BERT: Pre-training of Deep Bidirectional Transformers for
- Language Understanding](https://arxiv.org/abs/1810.04805) by Devlin et al.,
- 2018
+READMEs for specific papers. Below are some papers implemented in the repository
+and more NLP projects can be found in the
+[`projects`](https://github.com/tensorflow/models/tree/master/official/projects)
+folder:
+
+1. [BERT](MODEL_GARDEN.md#available-model-configs): [BERT: Pre-training of Deep
+ Bidirectional Transformers for Language
+ Understanding](https://arxiv.org/abs/1810.04805) by Devlin et al., 2018
2. [ALBERT](MODEL_GARDEN.md#available-model-configs):
[A Lite BERT for Self-supervised Learning of Language Representations](https://arxiv.org/abs/1909.11942)
by Lan et al., 2019
-3. [XLNet](xlnet):
+3. [XLNet](MODEL_GARDEN.md):
[XLNet: Generalized Autoregressive Pretraining for Language Understanding](https://arxiv.org/abs/1906.08237)
by Yang et al., 2019
-4. [Transformer for translation](transformer):
+4. [Transformer for translation](MODEL_GARDEN.md#available-model-configs):
[Attention Is All You Need](https://arxiv.org/abs/1706.03762) by Vaswani et
al., 2017
### Common Training Driver
We provide a single common driver [train.py](train.py) to train above SoTA
-models on popluar tasks. Please see [docs/train.md](docs/train.md) for
-more details.
-
+models on popular tasks. Please see [docs/train.md](docs/train.md) for more
+details.
### Pre-trained models with checkpoints and TF-Hub
We provide a large collection of baselines and checkpoints for NLP pre-trained
models. Please see [docs/pretrained_models.md](docs/pretrained_models.md) for
more details.
+
+## More Documentations
+
+Please read through the model training tutorials and references in the
+[docs/ folder](docs/README.md).
diff --git a/official/nlp/__init__.py b/official/nlp/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/nlp/__init__.py
+++ b/official/nlp/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/configs/__init__.py b/official/nlp/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/nlp/configs/__init__.py
+++ b/official/nlp/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/configs/bert.py b/official/nlp/configs/bert.py
index dc3b42d512b..1e7fa139a07 100644
--- a/official/nlp/configs/bert.py
+++ b/official/nlp/configs/bert.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -37,7 +37,11 @@ class ClsHeadConfig(base_config.Config):
@dataclasses.dataclass
class PretrainerConfig(base_config.Config):
"""Pretrainer configuration."""
- encoder: encoders.EncoderConfig = encoders.EncoderConfig()
+ encoder: encoders.EncoderConfig = dataclasses.field(
+ default_factory=encoders.EncoderConfig
+ )
cls_heads: List[ClsHeadConfig] = dataclasses.field(default_factory=list)
mlm_activation: str = "gelu"
mlm_initializer_range: float = 0.02
+ # Currently only used for mobile bert.
+ mlm_output_weights_use_proj: bool = False
diff --git a/official/nlp/configs/electra.py b/official/nlp/configs/electra.py
index 0c55e50e5e8..c7f33ec2405 100644
--- a/official/nlp/configs/electra.py
+++ b/official/nlp/configs/electra.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -31,6 +31,10 @@ class ElectraPretrainerConfig(base_config.Config):
discriminator_loss_weight: float = 50.0
tie_embeddings: bool = True
disallow_correct: bool = False
- generator_encoder: encoders.EncoderConfig = encoders.EncoderConfig()
- discriminator_encoder: encoders.EncoderConfig = encoders.EncoderConfig()
+ generator_encoder: encoders.EncoderConfig = dataclasses.field(
+ default_factory=encoders.EncoderConfig
+ )
+ discriminator_encoder: encoders.EncoderConfig = dataclasses.field(
+ default_factory=encoders.EncoderConfig
+ )
cls_heads: List[bert.ClsHeadConfig] = dataclasses.field(default_factory=list)
diff --git a/official/nlp/configs/encoders.py b/official/nlp/configs/encoders.py
index 0b182f18b9a..bc8f85e018a 100644
--- a/official/nlp/configs/encoders.py
+++ b/official/nlp/configs/encoders.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,10 +17,10 @@
Includes configurations and factory methods.
"""
import dataclasses
-from typing import Optional
+from typing import Optional, Sequence, Union
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.modeling import tf_utils
@@ -46,10 +46,39 @@ class BertEncoderConfig(hyperparams.Config):
embedding_size: Optional[int] = None
output_range: Optional[int] = None
return_all_encoder_outputs: bool = False
+ return_attention_scores: bool = False
+ return_word_embeddings: bool = False
# Pre/Post-LN Transformer
norm_first: bool = False
+@dataclasses.dataclass
+class FunnelEncoderConfig(hyperparams.Config):
+ """Funnel encoder configuration."""
+ vocab_size: int = 30522
+ hidden_size: int = 768
+ num_layers: int = 12
+ num_attention_heads: int = 12
+ max_position_embeddings: int = 512
+ type_vocab_size: int = 16
+ inner_dim: int = 3072
+ hidden_activation: str = "gelu"
+ approx_gelu: bool = True
+ dropout_rate: float = 0.1
+ attention_dropout_rate: float = 0.1
+ pool_type: str = "max"
+ pool_stride: Union[int, Sequence[Union[int, float]]] = 2
+ unpool_length: int = 0
+ initializer_range: float = 0.02
+ output_range: Optional[int] = None
+ embedding_width: Optional[int] = None
+ embedding_layer: Optional[tf_keras.layers.Layer] = None
+ norm_first: bool = False
+ share_rezero: bool = False
+ append_dense_inputs: bool = False
+ transformer_cls: str = "TransformerEncoderBlock"
+
+
@dataclasses.dataclass
class MobileBertEncoderConfig(hyperparams.Config):
"""MobileBERT encoder configuration.
@@ -221,25 +250,120 @@ class XLNetEncoderConfig(hyperparams.Config):
two_stream: bool = False
+@dataclasses.dataclass
+class QueryBertConfig(hyperparams.Config):
+ """Query BERT encoder configuration."""
+ vocab_size: int = 30522
+ hidden_size: int = 768
+ num_layers: int = 12
+ num_attention_heads: int = 12
+ hidden_activation: str = "gelu"
+ intermediate_size: int = 3072
+ dropout_rate: float = 0.1
+ attention_dropout_rate: float = 0.1
+ max_position_embeddings: int = 512
+ type_vocab_size: int = 2
+ initializer_range: float = 0.02
+ embedding_size: Optional[int] = None
+ output_range: Optional[int] = None
+ return_all_encoder_outputs: bool = False
+ return_attention_scores: bool = False
+ # Pre/Post-LN Transformer
+ norm_first: bool = False
+
+
+@dataclasses.dataclass
+class FNetEncoderConfig(hyperparams.Config):
+ """FNet encoder configuration."""
+ vocab_size: int = 30522
+ hidden_size: int = 768
+ num_layers: int = 12
+ num_attention_heads: int = 12
+ inner_activation: str = "gelu"
+ inner_dim: int = 3072
+ output_dropout: float = 0.1
+ attention_dropout: float = 0.1
+ max_sequence_length: int = 512
+ type_vocab_size: int = 2
+ initializer_range: float = 0.02
+ embedding_width: Optional[int] = None
+ output_range: Optional[int] = None
+ norm_first: bool = False
+ use_fft: bool = False
+ attention_layers: Sequence[int] = ()
+
+
+@dataclasses.dataclass
+class SparseMixerEncoderConfig(hyperparams.Config):
+ """SparseMixer encoder configuration."""
+ vocab_size: int = 30522
+ hidden_size: int = 768
+ num_layers: int = 14
+ moe_layers: Sequence[int] = (5, 6, 7, 8)
+ attention_layers: Sequence[int] = (10, 11, 12, 13)
+ num_experts: int = 16
+ train_capacity_factor: float = 1.
+ eval_capacity_factor: float = 1.
+ examples_per_group: float = 1.
+ use_fft: bool = False
+ num_attention_heads: int = 8
+ max_sequence_length: int = 512
+ type_vocab_size: int = 2
+ inner_dim: int = 3072
+ inner_activation: str = "gelu"
+ output_dropout: float = 0.1
+ attention_dropout: float = 0.1
+ initializer_range: float = 0.02
+ output_range: Optional[int] = None
+ embedding_width: Optional[int] = None
+ norm_first: bool = False
+
+
@dataclasses.dataclass
class EncoderConfig(hyperparams.OneOfConfig):
"""Encoder configuration."""
type: Optional[str] = "bert"
- albert: AlbertEncoderConfig = AlbertEncoderConfig()
- bert: BertEncoderConfig = BertEncoderConfig()
- bert_v2: BertEncoderConfig = BertEncoderConfig()
- bigbird: BigBirdEncoderConfig = BigBirdEncoderConfig()
- kernel: KernelEncoderConfig = KernelEncoderConfig()
- mobilebert: MobileBertEncoderConfig = MobileBertEncoderConfig()
- reuse: ReuseEncoderConfig = ReuseEncoderConfig()
- xlnet: XLNetEncoderConfig = XLNetEncoderConfig()
+ albert: AlbertEncoderConfig = dataclasses.field(
+ default_factory=AlbertEncoderConfig
+ )
+ bert: BertEncoderConfig = dataclasses.field(default_factory=BertEncoderConfig)
+ bert_v2: BertEncoderConfig = dataclasses.field(
+ default_factory=BertEncoderConfig
+ )
+ bigbird: BigBirdEncoderConfig = dataclasses.field(
+ default_factory=BigBirdEncoderConfig
+ )
+ funnel: FunnelEncoderConfig = dataclasses.field(
+ default_factory=FunnelEncoderConfig
+ )
+ kernel: KernelEncoderConfig = dataclasses.field(
+ default_factory=KernelEncoderConfig
+ )
+ mobilebert: MobileBertEncoderConfig = dataclasses.field(
+ default_factory=MobileBertEncoderConfig
+ )
+ reuse: ReuseEncoderConfig = dataclasses.field(
+ default_factory=ReuseEncoderConfig
+ )
+ xlnet: XLNetEncoderConfig = dataclasses.field(
+ default_factory=XLNetEncoderConfig
+ )
+ query_bert: QueryBertConfig = dataclasses.field(
+ default_factory=QueryBertConfig
+ )
+ fnet: FNetEncoderConfig = dataclasses.field(default_factory=FNetEncoderConfig)
+ sparse_mixer: SparseMixerEncoderConfig = dataclasses.field(
+ default_factory=SparseMixerEncoderConfig
+ )
# If `any` is used, the encoder building relies on any.BUILDER.
- any: hyperparams.Config = hyperparams.Config()
+ any: hyperparams.Config = dataclasses.field(
+ default_factory=hyperparams.Config
+ )
@gin.configurable
def build_encoder(config: EncoderConfig,
- embedding_layer: Optional[tf.keras.layers.Layer] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
encoder_cls=None,
bypass_config: bool = False):
"""Instantiate a Transformer encoder network from EncoderConfig.
@@ -257,7 +381,7 @@ def build_encoder(config: EncoderConfig,
An encoder instance.
"""
if bypass_config:
- return encoder_cls()
+ return encoder_cls() # pyrefly: ignore[not-callable]
encoder_type = config.type
encoder_cfg = config.get()
if encoder_cls and encoder_cls.__name__ == "EncoderScaffold":
@@ -266,7 +390,7 @@ def build_encoder(config: EncoderConfig,
type_vocab_size=encoder_cfg.type_vocab_size,
hidden_size=encoder_cfg.hidden_size,
max_seq_length=encoder_cfg.max_position_embeddings,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
dropout_rate=encoder_cfg.dropout_rate,
)
@@ -277,7 +401,7 @@ def build_encoder(config: EncoderConfig,
encoder_cfg.hidden_activation),
dropout_rate=encoder_cfg.dropout_rate,
attention_dropout_rate=encoder_cfg.attention_dropout_rate,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
)
kwargs = dict(
@@ -285,7 +409,7 @@ def build_encoder(config: EncoderConfig,
hidden_cfg=hidden_cfg,
num_hidden_instances=encoder_cfg.num_layers,
pooled_output_dim=encoder_cfg.hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
return_all_layer_outputs=encoder_cfg.return_all_encoder_outputs,
dict_outputs=True)
@@ -294,10 +418,10 @@ def build_encoder(config: EncoderConfig,
if encoder_type == "any":
encoder = encoder_cfg.BUILDER(encoder_cfg)
if not isinstance(encoder,
- (tf.Module, tf.keras.Model, tf.keras.layers.Layer)):
+ (tf.Module, tf_keras.Model, tf_keras.layers.Layer)):
raise ValueError("The BUILDER returns an unexpected instance. The "
"`build_encoder` should returns a tf.Module, "
- "tf.keras.Model or tf.keras.layers.Layer. However, "
+ "tf_keras.Model or tf_keras.layers.Layer. However, "
f"we get {encoder.__class__}")
return encoder
@@ -336,7 +460,7 @@ def build_encoder(config: EncoderConfig,
activation=tf_utils.get_activation(encoder_cfg.hidden_activation),
dropout_rate=encoder_cfg.dropout_rate,
attention_dropout_rate=encoder_cfg.attention_dropout_rate,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
dict_outputs=True)
@@ -357,7 +481,7 @@ def build_encoder(config: EncoderConfig,
block_size=encoder_cfg.block_size,
max_position_embeddings=encoder_cfg.max_position_embeddings,
type_vocab_size=encoder_cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
embedding_width=encoder_cfg.embedding_width,
use_gradient_checkpointing=encoder_cfg.use_gradient_checkpointing)
@@ -366,13 +490,13 @@ def build_encoder(config: EncoderConfig,
type_vocab_size=encoder_cfg.type_vocab_size,
hidden_size=encoder_cfg.hidden_size,
max_seq_length=encoder_cfg.max_position_embeddings,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
dropout_rate=encoder_cfg.dropout_rate)
attention_cfg = dict(
num_heads=encoder_cfg.num_attention_heads,
key_dim=int(encoder_cfg.hidden_size // encoder_cfg.num_attention_heads),
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
max_rand_mask_length=encoder_cfg.max_position_embeddings,
num_rand_blocks=encoder_cfg.num_rand_blocks,
@@ -387,7 +511,7 @@ def build_encoder(config: EncoderConfig,
dropout_rate=encoder_cfg.dropout_rate,
attention_dropout_rate=encoder_cfg.attention_dropout_rate,
norm_first=encoder_cfg.norm_first,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
attention_cls=layers.BigBirdAttention,
attention_cfg=attention_cfg)
@@ -399,26 +523,60 @@ def build_encoder(config: EncoderConfig,
mask_cls=layers.BigBirdMasks,
mask_cfg=dict(block_size=encoder_cfg.block_size),
pooled_output_dim=encoder_cfg.hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
return_all_layer_outputs=False,
dict_outputs=True,
layer_idx_as_attention_seed=True)
return networks.EncoderScaffold(**kwargs)
+ if encoder_type == "funnel":
+
+ if encoder_cfg.hidden_activation == "gelu":
+ activation = tf_utils.get_activation(
+ encoder_cfg.hidden_activation,
+ approximate=encoder_cfg.approx_gelu)
+ else:
+ activation = tf_utils.get_activation(encoder_cfg.hidden_activation)
+
+ return networks.FunnelTransformerEncoder(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ num_attention_heads=encoder_cfg.num_attention_heads,
+ max_sequence_length=encoder_cfg.max_position_embeddings,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ inner_dim=encoder_cfg.inner_dim,
+ inner_activation=activation,
+ output_dropout=encoder_cfg.dropout_rate,
+ attention_dropout=encoder_cfg.attention_dropout_rate,
+ pool_type=encoder_cfg.pool_type,
+ pool_stride=encoder_cfg.pool_stride,
+ unpool_length=encoder_cfg.unpool_length,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ output_range=encoder_cfg.output_range,
+ embedding_width=encoder_cfg.embedding_width,
+ embedding_layer=embedding_layer,
+ norm_first=encoder_cfg.norm_first,
+ share_rezero=encoder_cfg.share_rezero,
+ append_dense_inputs=encoder_cfg.append_dense_inputs,
+ transformer_cls=encoder_cfg.transformer_cls,
+ )
+
if encoder_type == "kernel":
embedding_cfg = dict(
vocab_size=encoder_cfg.vocab_size,
type_vocab_size=encoder_cfg.type_vocab_size,
hidden_size=encoder_cfg.hidden_size,
max_seq_length=encoder_cfg.max_position_embeddings,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
dropout_rate=encoder_cfg.dropout_rate)
attention_cfg = dict(
num_heads=encoder_cfg.num_attention_heads,
key_dim=int(encoder_cfg.hidden_size // encoder_cfg.num_attention_heads),
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
feature_transform=encoder_cfg.feature_transform,
num_random_features=encoder_cfg.num_random_features,
@@ -435,7 +593,7 @@ def build_encoder(config: EncoderConfig,
dropout_rate=encoder_cfg.dropout_rate,
attention_dropout_rate=encoder_cfg.attention_dropout_rate,
norm_first=encoder_cfg.norm_first,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
attention_cls=layers.KernelAttention,
attention_cfg=attention_cfg)
@@ -446,7 +604,7 @@ def build_encoder(config: EncoderConfig,
num_hidden_instances=encoder_cfg.num_layers,
mask_cls=layers.KernelMask,
pooled_output_dim=encoder_cfg.hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
return_all_layer_outputs=False,
dict_outputs=True,
@@ -473,7 +631,7 @@ def build_encoder(config: EncoderConfig,
inner_activation=encoder_cfg.inner_activation,
use_cls_mask=encoder_cfg.use_cls_mask,
embedding_width=encoder_cfg.embedding_width,
- initializer=tf.keras.initializers.RandomNormal(
+ initializer=tf_keras.initializers.RandomNormal(
stddev=encoder_cfg.initializer_range))
if encoder_type == "reuse":
@@ -482,7 +640,7 @@ def build_encoder(config: EncoderConfig,
type_vocab_size=encoder_cfg.type_vocab_size,
hidden_size=encoder_cfg.hidden_size,
max_seq_length=encoder_cfg.max_position_embeddings,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
dropout_rate=encoder_cfg.dropout_rate)
hidden_cfg = dict(
@@ -493,7 +651,7 @@ def build_encoder(config: EncoderConfig,
output_dropout=encoder_cfg.dropout_rate,
attention_dropout=encoder_cfg.attention_dropout_rate,
norm_first=encoder_cfg.norm_first,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
reuse_attention=encoder_cfg.reuse_attention,
use_relative_pe=encoder_cfg.use_relative_pe,
@@ -505,7 +663,7 @@ def build_encoder(config: EncoderConfig,
hidden_cfg=hidden_cfg,
num_hidden_instances=encoder_cfg.num_layers,
pooled_output_dim=encoder_cfg.hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
return_all_layer_outputs=False,
dict_outputs=True,
@@ -513,6 +671,81 @@ def build_encoder(config: EncoderConfig,
recursive=True)
return networks.EncoderScaffold(**kwargs)
+ if encoder_type == "query_bert":
+ embedding_layer = layers.FactorizedEmbedding(
+ vocab_size=encoder_cfg.vocab_size,
+ embedding_width=encoder_cfg.embedding_size,
+ output_dim=encoder_cfg.hidden_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ name="word_embeddings")
+ return networks.BertEncoderV2(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ num_attention_heads=encoder_cfg.num_attention_heads,
+ intermediate_size=encoder_cfg.intermediate_size,
+ activation=tf_utils.get_activation(encoder_cfg.hidden_activation),
+ dropout_rate=encoder_cfg.dropout_rate,
+ attention_dropout_rate=encoder_cfg.attention_dropout_rate,
+ max_sequence_length=encoder_cfg.max_position_embeddings,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ output_range=encoder_cfg.output_range,
+ embedding_layer=embedding_layer,
+ return_all_encoder_outputs=encoder_cfg.return_all_encoder_outputs,
+ return_attention_scores=encoder_cfg.return_attention_scores,
+ dict_outputs=True,
+ norm_first=encoder_cfg.norm_first)
+
+ if encoder_type == "fnet":
+ return networks.FNet(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ num_attention_heads=encoder_cfg.num_attention_heads,
+ inner_dim=encoder_cfg.inner_dim,
+ inner_activation=tf_utils.get_activation(encoder_cfg.inner_activation),
+ output_dropout=encoder_cfg.output_dropout,
+ attention_dropout=encoder_cfg.attention_dropout,
+ max_sequence_length=encoder_cfg.max_sequence_length,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ output_range=encoder_cfg.output_range,
+ embedding_width=encoder_cfg.embedding_width,
+ embedding_layer=embedding_layer,
+ norm_first=encoder_cfg.norm_first,
+ use_fft=encoder_cfg.use_fft,
+ attention_layers=encoder_cfg.attention_layers)
+
+ if encoder_type == "sparse_mixer":
+ return networks.SparseMixer(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ moe_layers=encoder_cfg.moe_layers,
+ attention_layers=encoder_cfg.attention_layers,
+ num_experts=encoder_cfg.num_experts,
+ train_capacity_factor=encoder_cfg.train_capacity_factor,
+ eval_capacity_factor=encoder_cfg.eval_capacity_factor,
+ examples_per_group=encoder_cfg.examples_per_group,
+ use_fft=encoder_cfg.use_fft,
+ num_attention_heads=encoder_cfg.num_attention_heads,
+ max_sequence_length=encoder_cfg.max_sequence_length,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ inner_dim=encoder_cfg.inner_dim,
+ inner_activation=tf_utils.get_activation(encoder_cfg.inner_activation),
+ output_dropout=encoder_cfg.output_dropout,
+ attention_dropout=encoder_cfg.attention_dropout,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ output_range=encoder_cfg.output_range,
+ embedding_width=encoder_cfg.embedding_width,
+ norm_first=encoder_cfg.norm_first,
+ embedding_layer=embedding_layer)
+
bert_encoder_cls = networks.BertEncoder
if encoder_type == "bert_v2":
bert_encoder_cls = networks.BertEncoderV2
@@ -530,11 +763,13 @@ def build_encoder(config: EncoderConfig,
attention_dropout_rate=encoder_cfg.attention_dropout_rate,
max_sequence_length=encoder_cfg.max_position_embeddings,
type_vocab_size=encoder_cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
output_range=encoder_cfg.output_range,
embedding_width=encoder_cfg.embedding_size,
embedding_layer=embedding_layer,
return_all_encoder_outputs=encoder_cfg.return_all_encoder_outputs,
+ return_attention_scores=encoder_cfg.return_attention_scores,
+ return_word_embeddings=encoder_cfg.return_word_embeddings,
dict_outputs=True,
norm_first=encoder_cfg.norm_first)
diff --git a/official/nlp/configs/encoders_test.py b/official/nlp/configs/encoders_test.py
index 6012c55fe9b..418817c1d4b 100644
--- a/official/nlp/configs/encoders_test.py
+++ b/official/nlp/configs/encoders_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for official.nlp.configs.encoders."""
import os
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.nlp.configs import encoders
diff --git a/official/nlp/configs/experiment_configs.py b/official/nlp/configs/experiment_configs.py
index 006d6d7d558..9a81f574e94 100644
--- a/official/nlp/configs/experiment_configs.py
+++ b/official/nlp/configs/experiment_configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/configs/experiments/wiki_books_pretrain.yaml b/official/nlp/configs/experiments/wiki_books_pretrain.yaml
new file mode 100644
index 00000000000..be126e0d0bf
--- /dev/null
+++ b/official/nlp/configs/experiments/wiki_books_pretrain.yaml
@@ -0,0 +1,48 @@
+task:
+ init_checkpoint: ''
+ model:
+ cls_heads: [{activation: tanh, cls_token_idx: 0, dropout_rate: 0.1, inner_dim: 768, name: next_sentence, num_classes: 2}]
+ train_data:
+ drop_remainder: true
+ global_batch_size: 512
+ input_path: '[Your processed wiki data path]*,[Your processed books data path]*'
+ is_training: true
+ max_predictions_per_seq: 76
+ seq_length: 512
+ use_next_sentence_label: true
+ use_position_id: false
+ use_v2_feature_names: true
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 512
+ input_path: '[Your processed wiki data path]-00000-of-00500,[Your processed books data path]-00000-of-00500'
+ is_training: false
+ max_predictions_per_seq: 76
+ seq_length: 512
+ use_next_sentence_label: true
+ use_position_id: false
+ use_v2_feature_names: true
+trainer:
+ checkpoint_interval: 20000
+ max_to_keep: 5
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ cycle: false
+ decay_steps: 1000000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0001
+ power: 1.0
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 10000
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ train_steps: 1000000
+ validation_interval: 1000
+ validation_steps: 64
diff --git a/official/nlp/configs/experiments/wiki_tfds_pretrain.yaml b/official/nlp/configs/experiments/wiki_tfds_pretrain.yaml
new file mode 100644
index 00000000000..388bce3210a
--- /dev/null
+++ b/official/nlp/configs/experiments/wiki_tfds_pretrain.yaml
@@ -0,0 +1,50 @@
+task:
+ init_checkpoint: ''
+ model:
+ cls_heads: [{activation: tanh, cls_token_idx: 0, dropout_rate: 0.1, inner_dim: 768, name: next_sentence, num_classes: 2}]
+ train_data:
+ drop_remainder: true
+ global_batch_size: 512
+ is_training: true
+ max_predictions_per_seq: 76
+ seq_length: 512
+ use_next_sentence_label: false
+ use_whole_word_masking: true
+ tfds_name: wikipedia/20201201.en
+ tfds_split: train
+ vocab_file_path: 'Please provide the vocab file path.'
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 32
+ is_training: false
+ max_predictions_per_seq: 76
+ seq_length: 512
+ use_next_sentence_label: false
+ use_whole_word_masking: true
+ tfds_name: wikipedia/20201201.en
+ tfds_split: train
+ vocab_file_path: 'Please provide the vocab file path.'
+trainer:
+ checkpoint_interval: 20000
+ max_to_keep: 5
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ cycle: false
+ decay_steps: 1000000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0001
+ power: 1.0
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 10000
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ train_steps: 1000000
+ validation_interval: 1000
+ validation_steps: 64
diff --git a/official/nlp/configs/finetuning_experiments.py b/official/nlp/configs/finetuning_experiments.py
index 23833d4cf49..c5d6f63b5d6 100644
--- a/official/nlp/configs/finetuning_experiments.py
+++ b/official/nlp/configs/finetuning_experiments.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/configs/pretraining_experiments.py b/official/nlp/configs/pretraining_experiments.py
index 1eedb878280..9dacdc2147a 100644
--- a/official/nlp/configs/pretraining_experiments.py
+++ b/official/nlp/configs/pretraining_experiments.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,8 +19,11 @@
from official.modeling import optimization
from official.nlp.data import pretrain_dataloader
from official.nlp.data import pretrain_dynamic_dataloader
+from official.nlp.data import pretrain_text_dataloader
+from official.nlp.tasks import electra_task
from official.nlp.tasks import masked_lm
+
_TRAINER = cfg.TrainerConfig(
train_steps=1000000,
optimizer_config=optimization.OptimizationConfig({
@@ -82,3 +85,51 @@ def bert_dynamic() -> cfg.ExperimentConfig:
'task.validation_data.is_training != None'
])
return config
+
+
+@exp_factory.register_config_factory('bert/text_wiki_pretraining')
+def bert_text_wiki_pretraining() -> cfg.ExperimentConfig:
+ r"""BERT with wiki text tfds.
+
+ Note that: only wikipedia english corpus is used. It cannot exactly reproduce
+ BERT training setup because the next sentence sampling is hard to match the
+ implementation with tf ops.
+ """
+ config = cfg.ExperimentConfig(
+ task=masked_lm.MaskedLMConfig(
+ train_data=pretrain_text_dataloader.BertPretrainTextDataConfig(
+ tfds_name='wikipedia/20201201.en',
+ tfds_split='train',
+ vocab_file_path='TODO for users',
+ ),
+ validation_data=pretrain_text_dataloader.BertPretrainTextDataConfig(
+ tfds_name='wikipedia/20201201.en',
+ tfds_split='train',
+ vocab_file_path='TODO for users',
+ is_training=False)),
+ trainer=_TRAINER,
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
+
+
+@exp_factory.register_config_factory('electra/pretraining')
+def electra_pretrain() -> cfg.ExperimentConfig:
+ """ELECTRA pretraining experiment."""
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(enable_xla=True),
+ task=electra_task.ElectraPretrainConfig(
+ train_data=pretrain_dataloader.BertPretrainDataConfig(),
+ validation_data=pretrain_dataloader.BertPretrainDataConfig(
+ is_training=False
+ ),
+ ),
+ trainer=_TRAINER,
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ],
+ )
+ return config
diff --git a/official/nlp/configs/wmt_transformer_experiments.py b/official/nlp/configs/wmt_transformer_experiments.py
index bdef599fa42..2d6455a4da1 100644
--- a/official/nlp/configs/wmt_transformer_experiments.py
+++ b/official/nlp/configs/wmt_transformer_experiments.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/continuous_finetune_lib.py b/official/nlp/continuous_finetune_lib.py
index 988b62c6032..e28ae7688ec 100644
--- a/official/nlp/continuous_finetune_lib.py
+++ b/official/nlp/continuous_finetune_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from typing import Any, Mapping, Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.core import config_definitions
diff --git a/official/nlp/continuous_finetune_lib_test.py b/official/nlp/continuous_finetune_lib_test.py
index 6ed727d73e2..6ca92a2156a 100644
--- a/official/nlp/continuous_finetune_lib_test.py
+++ b/official/nlp/continuous_finetune_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from absl import flags
from absl.testing import flagsaver
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.common import registry_imports
diff --git a/official/nlp/data/__init__.py b/official/nlp/data/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/nlp/data/__init__.py
+++ b/official/nlp/data/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/data/classifier_data_lib.py b/official/nlp/data/classifier_data_lib.py
index 3e95caf719b..445a6cea3a8 100644
--- a/official/nlp/data/classifier_data_lib.py
+++ b/official/nlp/data/classifier_data_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
import os
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.nlp.tools import tokenization
@@ -599,7 +599,7 @@ def _create_examples_tfds(self, set_type):
text_a = self.process_text_fn(example["sentence1"])
text_b = self.process_text_fn(example["sentence2"])
if set_type != "test":
- label = self.label_type(example["label"])
+ label = self.label_type(example["label"]) # pyrefly: ignore[not-callable]
examples.append(
InputExample(
guid=guid, text_a=text_a, text_b=text_b, label=label,
@@ -668,7 +668,7 @@ def __init__(self,
self._labels = list(range(info.features[self.label_key].num_classes))
def _process_tfds_params_str(self, params_str):
- """Extracts TFDS parameters from a comma-separated assignements string."""
+ """Extracts TFDS parameters from a comma-separated assignments string."""
dtype_map = {"int": int, "float": float}
cast_str_to_bool = lambda s: s.lower() not in ["false", "0"]
@@ -729,7 +729,7 @@ def _create_examples(self, split_name, set_type):
text_a = self.process_text_fn(example[self.text_key])
if self.text_b_key:
text_b = self.process_text_fn(example[self.text_b_key])
- label = self.label_type(example[self.label_key])
+ label = self.label_type(example[self.label_key]) # pyrefly: ignore[not-callable]
if self.skip_label is not None and label == self.skip_label:
continue
if self.weight_key:
@@ -1599,7 +1599,7 @@ def generate_tf_record_from_data_file(processor,
if is_regression:
meta_data["task_type"] = "bert_regression"
- meta_data["label_type"] = {int: "int", float: "float"}[label_type]
+ meta_data["label_type"] = {int: "int", float: "float"}[label_type] # pyrefly: ignore[bad-index]
else:
meta_data["task_type"] = "bert_classification"
meta_data["num_labels"] = len(processor.get_labels())
@@ -1607,6 +1607,6 @@ def generate_tf_record_from_data_file(processor,
meta_data["has_sample_weights"] = True
if eval_data_output_path:
- meta_data["eval_data_size"] = len(eval_input_data_examples)
+ meta_data["eval_data_size"] = len(eval_input_data_examples) # pyrefly: ignore[unbound-name]
return meta_data
diff --git a/official/nlp/data/classifier_data_lib_test.py b/official/nlp/data/classifier_data_lib_test.py
index f7a517da0a2..6825d3ebb0a 100644
--- a/official/nlp/data/classifier_data_lib_test.py
+++ b/official/nlp/data/classifier_data_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import tempfile
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.nlp.data import classifier_data_lib
diff --git a/official/nlp/data/create_finetuning_data.py b/official/nlp/data/create_finetuning_data.py
index be1b6b444a0..6e3d0e5d097 100644
--- a/official/nlp/data/create_finetuning_data.py
+++ b/official/nlp/data/create_finetuning_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,10 +18,9 @@
import json
import os
-# Import libraries
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import classifier_data_lib
from official.nlp.data import sentence_retrieval_lib
# word-piece tokenizer based squad_lib
@@ -165,7 +164,7 @@
"while ALBERT uses SentencePiece tokenizer.")
flags.DEFINE_string(
- "tfds_params", "", "Comma-separated list of TFDS parameter assigments for "
+ "tfds_params", "", "Comma-separated list of TFDS parameter assignments for "
"generic classfication data import (for more details "
"see the TfdsProcessor class documentation).")
@@ -270,7 +269,7 @@ def generate_classifier_dataset():
}
task_name = FLAGS.classification_task_name.lower()
if task_name not in processors:
- raise ValueError("Task not found: %s" % (task_name))
+ raise ValueError("Task not found: %s" % (task_name,))
processor = processors[task_name](process_text_fn=processor_text_fn)
return classifier_data_lib.generate_tf_record_from_data_file(
diff --git a/official/nlp/data/create_pretraining_data.py b/official/nlp/data/create_pretraining_data.py
index 4d5eae4de05..e3b979225cf 100644
--- a/official/nlp/data/create_pretraining_data.py
+++ b/official/nlp/data/create_pretraining_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,11 +18,10 @@
import itertools
import random
-# Import libraries
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.tools import tokenization
@@ -35,8 +34,26 @@
"output_file", None,
"Output TF example file (or comma-separated list of files).")
-flags.DEFINE_string("vocab_file", None,
- "The vocabulary file that the BERT model was trained on.")
+flags.DEFINE_enum(
+ "tokenization",
+ "WordPiece",
+ ["WordPiece", "SentencePiece"],
+ "Specifies the tokenizer implementation, i.e., whether to use WordPiece "
+ "or SentencePiece tokenizer. Canonical BERT uses WordPiece tokenizer, "
+ "while ALBERT uses SentencePiece tokenizer.",
+)
+
+flags.DEFINE_string(
+ "vocab_file",
+ None,
+ "For WordPiece tokenization, the vocabulary file of the tokenizer.",
+)
+
+flags.DEFINE_string(
+ "sp_model_file",
+ "",
+ "For SentencePiece tokenization, the path to the model of the tokenizer.",
+)
flags.DEFINE_bool(
"do_lower_case", True,
@@ -44,8 +61,10 @@
"models and False for cased models.")
flags.DEFINE_bool(
- "do_whole_word_mask", False,
- "Whether to use whole word masking rather than per-WordPiece masking.")
+ "do_whole_word_mask",
+ False,
+ "Whether to use whole word masking rather than per-token masking.",
+)
flags.DEFINE_integer(
"max_ngram_size", None,
@@ -198,16 +217,19 @@ def create_float_feature(values):
return feature
-def create_training_instances(input_files,
- tokenizer,
- max_seq_length,
- dupe_factor,
- short_seq_prob,
- masked_lm_prob,
- max_predictions_per_seq,
- rng,
- do_whole_word_mask=False,
- max_ngram_size=None):
+def create_training_instances(
+ input_files,
+ tokenizer,
+ processor_text_fn,
+ max_seq_length,
+ dupe_factor,
+ short_seq_prob,
+ masked_lm_prob,
+ max_predictions_per_seq,
+ rng,
+ do_whole_word_mask=False,
+ max_ngram_size=None,
+):
"""Create `TrainingInstance`s from raw text."""
all_documents = [[]]
@@ -219,11 +241,8 @@ def create_training_instances(input_files,
# that the "next sentence prediction" task doesn't span between documents.
for input_file in input_files:
with tf.io.gfile.GFile(input_file, "rb") as reader:
- while True:
- line = tokenization.convert_to_unicode(reader.readline())
- if not line:
- break
- line = line.strip()
+ for line in reader:
+ line = processor_text_fn(line)
# Empty lines are used as document delimiters
if not line:
@@ -432,7 +451,7 @@ def _contiguous(sorted_grams):
def _masking_ngrams(grams, max_ngram_size, max_masked_tokens, rng):
"""Create a list of masking {1, ..., n}-grams from a list of one-grams.
- This is an extention of 'whole word masking' to mask multiple, contiguous
+ This is an extension of 'whole word masking' to mask multiple, contiguous
words such as (e.g., "the red boat").
Each input gram represents the token indices of a single word,
@@ -488,7 +507,7 @@ def _masking_ngrams(grams, max_ngram_size, max_masked_tokens, rng):
rng.shuffle(v)
# Create the weighting for n-gram length selection.
- # Stored cummulatively for `random.choices` below.
+ # Stored cumulatively for `random.choices` below.
cummulative_weights = list(
itertools.accumulate([1./n for n in range(1, max_ngram_size+1)]))
@@ -498,7 +517,7 @@ def _masking_ngrams(grams, max_ngram_size, max_masked_tokens, rng):
# Loop until we have enough masked tokens or there are no more candidate
# n-grams of any length.
# Each code path should ensure one or more elements from `ngrams` are removed
- # to guarentee this loop terminates.
+ # to guarantee this loop terminates.
while (sum(masked_tokens) < max_masked_tokens and
sum(len(s) for s in ngrams.values())):
# Pick an n-gram size based on our weights.
@@ -535,7 +554,7 @@ def _masking_ngrams(grams, max_ngram_size, max_masked_tokens, rng):
return output_ngrams
-def _wordpieces_to_grams(tokens):
+def _tokens_to_grams(tokens):
"""Reconstitue grams (words) from `tokens`.
E.g.,
@@ -543,7 +562,8 @@ def _wordpieces_to_grams(tokens):
grams: [ [1,2), [2, 4), [4,5) , [5, 6)]
Args:
- tokens: list of wordpieces
+ tokens: list of tokens (word pieces or sentence pieces).
+
Returns:
List of _Grams representing spans of whole words
(without "[CLS]" and "[SEP]").
@@ -570,7 +590,7 @@ def create_masked_lm_predictions(tokens, masked_lm_prob,
max_ngram_size=None):
"""Creates the predictions for the masked LM objective."""
if do_whole_word_mask:
- grams = _wordpieces_to_grams(tokens)
+ grams = _tokens_to_grams(tokens)
else:
# Here we consider each token to be a word to allow for sub-word masking.
if max_ngram_size:
@@ -633,9 +653,28 @@ def truncate_seq_pair(tokens_a, tokens_b, max_num_tokens, rng):
trunc_tokens.pop()
+def get_processor_text_fn(is_sentence_piece, do_lower_case):
+ def processor_text_fn(text):
+ text = tokenization.convert_to_unicode(text)
+ if is_sentence_piece:
+ # Additional preprocessing specific to the SentencePiece tokenizer.
+ text = tokenization.preprocess_text(text, lower=do_lower_case)
+
+ return text.strip()
+
+ return processor_text_fn
+
+
def main(_):
- tokenizer = tokenization.FullTokenizer(
- vocab_file=FLAGS.vocab_file, do_lower_case=FLAGS.do_lower_case)
+ if FLAGS.tokenization == "WordPiece":
+ tokenizer = tokenization.FullTokenizer(
+ vocab_file=FLAGS.vocab_file, do_lower_case=FLAGS.do_lower_case
+ )
+ processor_text_fn = get_processor_text_fn(False, FLAGS.do_lower_case)
+ else:
+ assert FLAGS.tokenization == "SentencePiece"
+ tokenizer = tokenization.FullSentencePieceTokenizer(FLAGS.sp_model_file)
+ processor_text_fn = get_processor_text_fn(True, FLAGS.do_lower_case)
input_files = []
for input_pattern in FLAGS.input_file.split(","):
@@ -647,9 +686,18 @@ def main(_):
rng = random.Random(FLAGS.random_seed)
instances = create_training_instances(
- input_files, tokenizer, FLAGS.max_seq_length, FLAGS.dupe_factor,
- FLAGS.short_seq_prob, FLAGS.masked_lm_prob, FLAGS.max_predictions_per_seq,
- rng, FLAGS.do_whole_word_mask, FLAGS.max_ngram_size)
+ input_files,
+ tokenizer,
+ processor_text_fn,
+ FLAGS.max_seq_length,
+ FLAGS.dupe_factor,
+ FLAGS.short_seq_prob,
+ FLAGS.masked_lm_prob,
+ FLAGS.max_predictions_per_seq,
+ rng,
+ FLAGS.do_whole_word_mask,
+ FLAGS.max_ngram_size,
+ )
output_files = FLAGS.output_file.split(",")
logging.info("*** Writing to output files ***")
@@ -665,5 +713,4 @@ def main(_):
if __name__ == "__main__":
flags.mark_flag_as_required("input_file")
flags.mark_flag_as_required("output_file")
- flags.mark_flag_as_required("vocab_file")
app.run(main)
diff --git a/official/nlp/data/create_pretraining_data_test.py b/official/nlp/data/create_pretraining_data_test.py
index da50d5479e4..ea83b3aa551 100644
--- a/official/nlp/data/create_pretraining_data_test.py
+++ b/official/nlp/data/create_pretraining_data_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for official.nlp.data.create_pretraining_data."""
import random
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import create_pretraining_data as cpd
@@ -43,7 +43,7 @@ def assertTokens(self, input_tokens, output_tokens, masked_positions,
continue
self.fail("invalid mask value: {}".format(output_token))
- def test_wordpieces_to_grams(self):
+ def test_tokens_to_grams(self):
tests = [
(["That", "cone"], [(0, 1), (1, 2)]),
(["That", "cone", "##s"], [(0, 1), (1, 3)]),
@@ -52,7 +52,7 @@ def test_wordpieces_to_grams(self):
(["[CLS]", "Up", "##dog", "[SEP]", "Down"], [(1, 3), (4, 5)]),
]
for inp, expected in tests:
- output = cpd._wordpieces_to_grams(inp)
+ output = cpd._tokens_to_grams(inp)
self.assertEqual(expected, output)
def test_window(self):
@@ -81,8 +81,8 @@ def test_create_masked_lm_predictions(self):
rng=rng,
do_whole_word_mask=False,
max_ngram_size=None))
- self.assertEqual(len(masked_positions), 3)
- self.assertEqual(len(masked_labels), 3)
+ self.assertLen(masked_positions, 3)
+ self.assertLen(masked_labels, 3)
self.assertTokens(tokens, output_tokens, masked_positions, masked_labels)
def test_create_masked_lm_predictions_whole_word(self):
@@ -100,8 +100,8 @@ def test_create_masked_lm_predictions_whole_word(self):
max_ngram_size=None))
# since we can't get exactly three tokens without breaking a word we
# only take two.
- self.assertEqual(len(masked_positions), 2)
- self.assertEqual(len(masked_labels), 2)
+ self.assertLen(masked_positions, 2)
+ self.assertLen(masked_labels, 2)
self.assertTokens(tokens, output_tokens, masked_positions, masked_labels)
# ensure that we took an entire word.
self.assertIn(masked_labels, [["a", "##a"], ["b", "##b"], ["c", "##c"]])
@@ -119,8 +119,8 @@ def test_create_masked_lm_predictions_ngram(self):
rng=rng,
do_whole_word_mask=True,
max_ngram_size=3))
- self.assertEqual(len(masked_positions), 76)
- self.assertEqual(len(masked_labels), 76)
+ self.assertLen(masked_positions, 76)
+ self.assertLen(masked_labels, 76)
self.assertTokens(tokens, output_tokens, masked_positions, masked_labels)
diff --git a/official/nlp/data/create_xlnet_pretraining_data.py b/official/nlp/data/create_xlnet_pretraining_data.py
index 3657962fd19..97b7a9d8f3b 100644
--- a/official/nlp/data/create_xlnet_pretraining_data.py
+++ b/official/nlp/data/create_xlnet_pretraining_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,14 +23,12 @@
from typing import Iterable, Mapping, List, Optional, Tuple
import unicodedata
-# Import libraries
-
from absl import app
from absl import flags
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.tools import tokenization
@@ -225,7 +223,7 @@ def preprocess_and_tokenize_input_files(
continue
total_number_of_lines += line_count
all_tokens = np.array(all_tokens, dtype=np.int64)
- all_sentence_ids = np.array(all_sentence_ids, dtype=np.bool)
+ all_sentence_ids = np.array(all_sentence_ids, dtype=bool)
all_data.append((all_tokens, all_sentence_ids))
logging.info("Completed text preprocessing. Total number of lines: %d",
@@ -271,7 +269,7 @@ def _create_a_and_b_segments(
Args:
tokens: The 1D input token ids. This represents an individual entry within a
batch.
- sentence_ids: The 1D input sentence ids. This represents an indivdual entry
+ sentence_ids: The 1D input sentence ids. This represents an individual entry
within a batch. This should be the same length as `tokens`.
begin_index: The reference beginning index to split data.
total_length: The target combined length of segments A and B.
diff --git a/official/nlp/data/create_xlnet_pretraining_data_test.py b/official/nlp/data/create_xlnet_pretraining_data_test.py
index 6a3b96833ed..521d4673634 100644
--- a/official/nlp/data/create_xlnet_pretraining_data_test.py
+++ b/official/nlp/data/create_xlnet_pretraining_data_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import create_xlnet_pretraining_data as cpd
diff --git a/official/nlp/data/data_loader.py b/official/nlp/data/data_loader.py
index b962d5f9740..d54ea92f287 100644
--- a/official/nlp/data/data_loader.py
+++ b/official/nlp/data/data_loader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import abc
from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class DataLoader(metaclass=abc.ABCMeta):
diff --git a/official/nlp/data/data_loader_factory.py b/official/nlp/data/data_loader_factory.py
index f3a2decb8c5..18badd6960c 100644
--- a/official/nlp/data/data_loader_factory.py
+++ b/official/nlp/data/data_loader_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/data/data_loader_factory_test.py b/official/nlp/data/data_loader_factory_test.py
index 518717a3f37..10af357c5ac 100644
--- a/official/nlp/data/data_loader_factory_test.py
+++ b/official/nlp/data/data_loader_factory_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for official.nlp.data.data_loader_factory."""
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.nlp.data import data_loader_factory
diff --git a/official/nlp/data/dual_encoder_dataloader.py b/official/nlp/data/dual_encoder_dataloader.py
index 1818d07f0bc..8f5954de07b 100644
--- a/official/nlp/data/dual_encoder_dataloader.py
+++ b/official/nlp/data/dual_encoder_dataloader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,14 +13,15 @@
# limitations under the License.
"""Loads dataset for the dual encoder (retrieval) task."""
+import dataclasses
import functools
import itertools
from typing import Iterable, Mapping, Optional, Tuple
-import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_hub as hub
+from official.common import dataset_fn
from official.core import config_definitions as cfg
from official.core import input_reader
from official.nlp.data import data_loader
@@ -47,6 +48,7 @@ class DualEncoderDataConfig(cfg.DataConfig):
right_text_fields: Tuple[str] = ('right_input',)
is_training: bool = True
seq_length: int = 128
+ file_type: str = 'tfrecord'
@data_loader_factory.register_data_loader_cls(DualEncoderDataConfig)
@@ -140,6 +142,6 @@ def load(self, input_context: Optional[tf.distribute.InputContext] = None):
params=self._params,
# Skip `decoder_fn` for tfds input.
decoder_fn=self._decode if self._params.input_path else None,
- dataset_fn=tf.data.TFRecordDataset,
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
postprocess_fn=self._bert_preprocess)
return reader.read(input_context)
diff --git a/official/nlp/data/dual_encoder_dataloader_test.py b/official/nlp/data/dual_encoder_dataloader_test.py
index bebdc1531ef..63a717e5b79 100644
--- a/official/nlp/data/dual_encoder_dataloader_test.py
+++ b/official/nlp/data/dual_encoder_dataloader_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import dual_encoder_dataloader
diff --git a/official/nlp/data/pretrain_dataloader.py b/official/nlp/data/pretrain_dataloader.py
index f2a33cd4277..0e0b330b36f 100644
--- a/official/nlp/data/pretrain_dataloader.py
+++ b/official/nlp/data/pretrain_dataloader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,13 +13,14 @@
# limitations under the License.
"""Loads dataset for the BERT pretraining task."""
+import dataclasses
from typing import Mapping, Optional
from absl import logging
-import dataclasses
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.common import dataset_fn
from official.core import config_definitions as cfg
from official.core import input_reader
from official.nlp.data import data_loader
@@ -42,6 +43,7 @@ class BertPretrainDataConfig(cfg.DataConfig):
# v2_feature_names is True, the data loader assumes the tf.Examples use
# `input_word_ids` and `input_type_ids` as keys.
use_v2_feature_names: bool = False
+ file_type: str = 'tfrecord'
@data_loader_factory.register_data_loader_cls(BertPretrainDataConfig)
@@ -128,7 +130,10 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
def load(self, input_context: Optional[tf.distribute.InputContext] = None):
"""Returns a tf.dataset.Dataset."""
reader = input_reader.InputReader(
- params=self._params, decoder_fn=self._decode, parser_fn=self._parse)
+ params=self._params,
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
+ decoder_fn=self._decode,
+ parser_fn=self._parse)
return reader.read(input_context)
@@ -138,7 +143,7 @@ class XLNetPretrainDataConfig(cfg.DataConfig):
Attributes:
input_path: See base class.
- global_batch_size: See base calss.
+ global_batch_size: See base class.
is_training: See base class.
seq_length: The length of each sequence.
max_predictions_per_seq: The number of predictions per sequence.
@@ -146,23 +151,20 @@ class XLNetPretrainDataConfig(cfg.DataConfig):
should be the same value used during pretrain data creation.
sample_strategy: The strategy used to sample factorization permutations.
Possible values: 'single_token', 'whole_word', 'token_span', 'word_span'.
- min_num_tokens: The minimum number of tokens to sample in a span.
- This is used when `sample_strategy` is 'token_span'.
- max_num_tokens: The maximum number of tokens to sample in a span.
- This is used when `sample_strategy` is 'token_span'.
- min_num_words: The minimum number of words to sample in a span.
- This is used when `sample_strategy` is 'word_span'.
- max_num_words: The maximum number of words to sample in a span.
- This is used when `sample_strategy` is 'word_span'.
- permutation_size: The length of the longest permutation. This can be set
- to `reuse_length`. This should NOT be greater than `reuse_length`,
- otherwise this may introduce data leaks.
+ min_num_tokens: The minimum number of tokens to sample in a span. This is
+ used when `sample_strategy` is 'token_span'.
+ max_num_tokens: The maximum number of tokens to sample in a span. This is
+ used when `sample_strategy` is 'token_span'.
+ min_num_words: The minimum number of words to sample in a span. This is used
+ when `sample_strategy` is 'word_span'.
+ max_num_words: The maximum number of words to sample in a span. This is used
+ when `sample_strategy` is 'word_span'.
+ permutation_size: The length of the longest permutation. This can be set to
+ `reuse_length`. This should NOT be greater than `reuse_length`, otherwise
+ this may introduce data leaks.
leak_ratio: The percentage of masked tokens that are leaked.
- segment_sep_id: The ID of the SEP token used when preprocessing
- the dataset.
- segment_cls_id: The ID of the CLS token used when preprocessing
- the dataset.
-
+ segment_sep_id: The ID of the SEP token used when preprocessing the dataset.
+ segment_cls_id: The ID of the CLS token used when preprocessing the dataset.
"""
input_path: str = ''
global_batch_size: int = 512
@@ -205,12 +207,9 @@ def __init__(self, params: XLNetPretrainDataConfig):
def _decode(self, record: tf.Tensor):
"""Decodes a serialized tf.Example."""
name_to_features = {
- 'input_word_ids':
- tf.io.FixedLenFeature([self._seq_length], tf.int64),
- 'input_type_ids':
- tf.io.FixedLenFeature([self._seq_length], tf.int64),
- 'boundary_indices':
- tf.io.VarLenFeature(tf.int64),
+ 'input_word_ids': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'input_type_ids': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'boundary_indices': tf.io.VarLenFeature(tf.int64),
}
example = tf.io.parse_single_example(record, name_to_features)
@@ -236,20 +235,20 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
else:
boundary = None
- input_mask = self._online_sample_mask(inputs=inputs, boundary=boundary)
+ input_mask = self._online_sample_mask(inputs=inputs, boundary=boundary) # pyrefly: ignore[bad-argument-type]
if self._reuse_length > 0:
if self._permutation_size > self._reuse_length:
logging.warning(
'`permutation_size` is greater than `reuse_length` (%d > %d).'
- 'This may introduce data leakage.',
- self._permutation_size, self._reuse_length)
+ 'This may introduce data leakage.', self._permutation_size,
+ self._reuse_length)
# Enable the memory mechanism.
# Permute the reuse and non-reuse segments separately.
non_reuse_len = self._seq_length - self._reuse_length
- if not (self._reuse_length % self._permutation_size == 0
- and non_reuse_len % self._permutation_size == 0):
+ if not (self._reuse_length % self._permutation_size == 0 and
+ non_reuse_len % self._permutation_size == 0):
raise ValueError('`reuse_length` and `seq_length` should both be '
'a multiple of `permutation_size`.')
@@ -260,17 +259,20 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
input_mask=input_mask[:self._reuse_length])
# Creates permutation mask and target mask for the rest of tokens in
- # current example, which are concatentation of two new segments.
+ # current example, which are concatenation of two new segments.
perm_mask_1, target_mask_1, tokens_1, masked_1 = self._get_factorization(
inputs[self._reuse_length:], input_mask[self._reuse_length:])
- perm_mask_0 = tf.concat(
- [perm_mask_0,
- tf.zeros([self._reuse_length, non_reuse_len], dtype=tf.int32)],
- axis=1)
- perm_mask_1 = tf.concat(
- [tf.ones([non_reuse_len, self._reuse_length], dtype=tf.int32),
- perm_mask_1], axis=1)
+ perm_mask_0 = tf.concat([
+ perm_mask_0,
+ tf.zeros([self._reuse_length, non_reuse_len], dtype=tf.int32)
+ ],
+ axis=1)
+ perm_mask_1 = tf.concat([
+ tf.ones([non_reuse_len, self._reuse_length], dtype=tf.int32),
+ perm_mask_1
+ ],
+ axis=1)
perm_mask = tf.concat([perm_mask_0, perm_mask_1], axis=0)
target_mask = tf.concat([target_mask_0, target_mask_1], axis=0)
tokens = tf.concat([tokens_0, tokens_1], axis=0)
@@ -283,8 +285,8 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
# Permute the entire sequence together
perm_mask, target_mask, tokens, masked_tokens = self._get_factorization(
inputs=inputs, input_mask=input_mask)
- x['permutation_mask'] = tf.reshape(
- perm_mask, [self._seq_length, self._seq_length])
+ x['permutation_mask'] = tf.reshape(perm_mask,
+ [self._seq_length, self._seq_length])
x['input_word_ids'] = tokens
x['masked_tokens'] = masked_tokens
@@ -313,7 +315,8 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
target_mask = tf.concat([
tf.ones([actual_num_predict], dtype=tf.int32),
tf.zeros([pad_len], dtype=tf.int32)
- ], axis=0)
+ ],
+ axis=0)
x['target_mask'] = tf.reshape(target_mask,
[self._max_predictions_per_seq])
else:
@@ -321,16 +324,14 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
x['target_mask'] = tf.reshape(target_mask, [self._seq_length])
return x
- def _index_pair_to_mask(self,
- begin_indices: tf.Tensor,
+ def _index_pair_to_mask(self, begin_indices: tf.Tensor,
end_indices: tf.Tensor,
inputs: tf.Tensor) -> tf.Tensor:
"""Converts beginning and end indices into an actual mask."""
non_func_mask = tf.logical_and(
tf.not_equal(inputs, self._sep_id), tf.not_equal(inputs, self._cls_id))
all_indices = tf.where(
- non_func_mask,
- tf.range(self._seq_length, dtype=tf.int32),
+ non_func_mask, tf.range(self._seq_length, dtype=tf.int32),
tf.constant(-1, shape=[self._seq_length], dtype=tf.int32))
candidate_matrix = tf.cast(
tf.logical_and(all_indices[None, :] >= begin_indices[:, None],
@@ -352,8 +353,7 @@ def _single_token_mask(self, inputs: tf.Tensor) -> tf.Tensor:
masked_pos = tf.random.shuffle(non_func_indices)
masked_pos = tf.sort(masked_pos[:self._max_predictions_per_seq])
- sparse_indices = tf.stack(
- [tf.zeros_like(masked_pos), masked_pos], axis=-1)
+ sparse_indices = tf.stack([tf.zeros_like(masked_pos), masked_pos], axis=-1)
sparse_indices = tf.cast(sparse_indices, tf.int64)
sparse_indices = tf.sparse.SparseTensor(
@@ -361,14 +361,11 @@ def _single_token_mask(self, inputs: tf.Tensor) -> tf.Tensor:
values=tf.ones_like(masked_pos),
dense_shape=(1, self._seq_length))
- target_mask = tf.sparse.to_dense(
- sp_input=sparse_indices,
- default_value=0)
+ target_mask = tf.sparse.to_dense(sp_input=sparse_indices, default_value=0)
return tf.squeeze(tf.cast(target_mask, tf.bool))
- def _whole_word_mask(self,
- inputs: tf.Tensor,
+ def _whole_word_mask(self, inputs: tf.Tensor,
boundary: tf.Tensor) -> tf.Tensor:
"""Samples whole words as prediction targets."""
pair_indices = tf.concat([boundary[:-1, None], boundary[1:, None]], axis=1)
@@ -378,9 +375,7 @@ def _whole_word_mask(self,
end_indices = cand_pair_indices[:, 1]
return self._index_pair_to_mask(
- begin_indices=begin_indices,
- end_indices=end_indices,
- inputs=inputs)
+ begin_indices=begin_indices, end_indices=end_indices, inputs=inputs)
def _token_span_mask(self, inputs: tf.Tensor) -> tf.Tensor:
"""Samples token spans as prediction targets."""
@@ -429,13 +424,9 @@ def _token_span_mask(self, inputs: tf.Tensor) -> tf.Tensor:
end_indices = tf.gather(end_indices, order)
return self._index_pair_to_mask(
- begin_indices=begin_indices,
- end_indices=end_indices,
- inputs=inputs)
+ begin_indices=begin_indices, end_indices=end_indices, inputs=inputs)
- def _word_span_mask(self,
- inputs: tf.Tensor,
- boundary: tf.Tensor):
+ def _word_span_mask(self, inputs: tf.Tensor, boundary: tf.Tensor):
"""Sample whole word spans as prediction targets."""
min_num_words = self._params.min_num_words
max_num_words = self._params.max_num_words
@@ -486,12 +477,9 @@ def _word_span_mask(self,
end_indices = tf.gather(end_indices, order)
return self._index_pair_to_mask(
- begin_indices=begin_indices,
- end_indices=end_indices,
- inputs=inputs)
+ begin_indices=begin_indices, end_indices=end_indices, inputs=inputs)
- def _online_sample_mask(self,
- inputs: tf.Tensor,
+ def _online_sample_mask(self, inputs: tf.Tensor,
boundary: tf.Tensor) -> tf.Tensor:
"""Samples target positions for predictions.
@@ -531,15 +519,13 @@ def _online_sample_mask(self,
else:
raise NotImplementedError('Invalid sample strategy.')
- def _get_factorization(self,
- inputs: tf.Tensor,
- input_mask: tf.Tensor):
+ def _get_factorization(self, inputs: tf.Tensor, input_mask: tf.Tensor):
"""Samples a permutation of the factorization order.
Args:
inputs: the input tokens.
- input_mask: the `bool` Tensor of the same shape as `inputs`.
- If `True`, then this means select for partial prediction.
+ input_mask: the `bool` Tensor of the same shape as `inputs`. If `True`,
+ then this means select for partial prediction.
Returns:
perm_mask: An `int32` Tensor of shape [seq_length, seq_length] consisting
@@ -552,7 +538,6 @@ def _get_factorization(self,
input. This token will not be included in the loss.
tokens: int32 Tensor of shape [seq_length].
masked_tokens: int32 Tensor of shape [seq_length].
-
"""
factorization_length = tf.shape(inputs)[0]
# Generate permutation indices
@@ -576,8 +561,8 @@ def _get_factorization(self,
if self._leak_ratio > 0:
leak_tokens = tf.logical_and(
masked_tokens,
- tf.random.uniform([factorization_length],
- maxval=1.0) < self._leak_ratio)
+ tf.random.uniform([factorization_length], maxval=1.0) <
+ self._leak_ratio)
can_attend_self = tf.logical_or(non_masked_or_func_tokens, leak_tokens)
else:
can_attend_self = non_masked_or_func_tokens
diff --git a/official/nlp/data/pretrain_dataloader_test.py b/official/nlp/data/pretrain_dataloader_test.py
index ce7f216f9af..2bdaee01bb8 100644
--- a/official/nlp/data/pretrain_dataloader_test.py
+++ b/official/nlp/data/pretrain_dataloader_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import pretrain_dataloader
diff --git a/official/nlp/data/pretrain_dynamic_dataloader.py b/official/nlp/data/pretrain_dynamic_dataloader.py
index ab614454680..9cba8b71128 100644
--- a/official/nlp/data/pretrain_dynamic_dataloader.py
+++ b/official/nlp/data/pretrain_dynamic_dataloader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Optional, Tuple
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import input_reader
diff --git a/official/nlp/data/pretrain_dynamic_dataloader_test.py b/official/nlp/data/pretrain_dynamic_dataloader_test.py
index 3927b799355..c1eae4cd3a5 100644
--- a/official/nlp/data/pretrain_dynamic_dataloader_test.py
+++ b/official/nlp/data/pretrain_dynamic_dataloader_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl.testing import parameterized
import numpy as np
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
diff --git a/official/nlp/data/pretrain_text_dataloader.py b/official/nlp/data/pretrain_text_dataloader.py
new file mode 100644
index 00000000000..0b05d22a06b
--- /dev/null
+++ b/official/nlp/data/pretrain_text_dataloader.py
@@ -0,0 +1,226 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Loads text dataset for the BERT pretraining task."""
+import dataclasses
+from typing import List, Mapping, Optional, Text
+
+import tensorflow as tf, tf_keras
+import tensorflow_text as tf_text
+
+from official.common import dataset_fn
+from official.core import config_definitions as cfg
+from official.core import input_reader
+from official.nlp.data import data_loader
+from official.nlp.data import data_loader_factory
+from official.nlp.modeling.ops import segment_extractor
+
+
+@dataclasses.dataclass
+class BertPretrainTextDataConfig(cfg.DataConfig):
+ """Data config for BERT pretraining task (tasks/masked_lm) from text."""
+ input_path: str = ""
+ doc_batch_size: int = 8
+ global_batch_size: int = 512
+ is_training: bool = True
+ seq_length: int = 512
+ max_predictions_per_seq: int = 76
+ use_next_sentence_label: bool = True
+ # The name of the text feature fields. The text features will be
+ # concatenated in order.
+ # Note: More than 1 field name is not compatible with NSP.
+ text_field_names: Optional[List[str]] = dataclasses.field(
+ default_factory=lambda: ["text"])
+ vocab_file_path: str = ""
+ masking_rate: float = 0.15
+ use_whole_word_masking: bool = False
+ file_type: str = "tfrecord"
+
+
+_CLS_TOKEN = b"[CLS]"
+_SEP_TOKEN = b"[SEP]"
+_MASK_TOKEN = b"[MASK]"
+_NUM_OOV_BUCKETS = 1
+# Accounts for [CLS] and 2 x [SEP] tokens
+_NUM_SPECIAL_TOKENS = 3
+
+
+@data_loader_factory.register_data_loader_cls(BertPretrainTextDataConfig)
+class BertPretrainTextDataLoader(data_loader.DataLoader):
+ """A class to load text dataset for BERT pretraining task."""
+
+ def __init__(self, params):
+ """Inits `BertPretrainTextDataLoader` class.
+
+ Args:
+ params: A `BertPretrainTextDataConfig` object.
+ """
+ if len(params.text_field_names) > 1 and params.use_next_sentence_label:
+ raise ValueError("Currently there is no support for more than text field "
+ "while generating next sentence labels.")
+
+ self._params = params
+ self._seq_length = params.seq_length
+ self._max_predictions_per_seq = params.max_predictions_per_seq
+ self._use_next_sentence_label = params.use_next_sentence_label
+ self._masking_rate = params.masking_rate
+ self._use_whole_word_masking = params.use_whole_word_masking
+
+ lookup_table_init = tf.lookup.TextFileInitializer(
+ params.vocab_file_path,
+ key_dtype=tf.string,
+ key_index=tf.lookup.TextFileIndex.WHOLE_LINE,
+ value_dtype=tf.int64,
+ value_index=tf.lookup.TextFileIndex.LINE_NUMBER)
+ self._vocab_lookup_table = tf.lookup.StaticVocabularyTable(
+ lookup_table_init,
+ num_oov_buckets=_NUM_OOV_BUCKETS,
+ lookup_key_dtype=tf.string)
+
+ self._cls_token = self._vocab_lookup_table.lookup(tf.constant(_CLS_TOKEN))
+ self._sep_token = self._vocab_lookup_table.lookup(tf.constant(_SEP_TOKEN))
+ self._mask_token = self._vocab_lookup_table.lookup(tf.constant(_MASK_TOKEN))
+
+ # -_NUM_OOV_BUCKETS to offset unused OOV bucket.
+ self._vocab_size = self._vocab_lookup_table.size() - _NUM_OOV_BUCKETS
+
+ def _decode(self, record: tf.Tensor) -> Mapping[Text, tf.Tensor]:
+ """Decodes a serialized tf.Example."""
+ name_to_features = {}
+ for text_field_name in self._params.text_field_names:
+ name_to_features[text_field_name] = tf.io.FixedLenFeature([], tf.string)
+ return tf.io.parse_single_example(record, name_to_features)
+
+ def _tokenize(self, segments):
+ """Tokenize the input segments."""
+ # Tokenize segments
+ tokenizer = tf_text.BertTokenizer(
+ self._vocab_lookup_table, token_out_type=tf.int64)
+
+ if self._use_whole_word_masking:
+ # tokenize the segments which should have the shape:
+ # [num_sentence, (num_words), (num_wordpieces)]
+ segments = [tokenizer.tokenize(s) for s in segments]
+ else:
+ # tokenize the segments and merge out the token dimension so that each
+ # segment has the shape: [num_sentence, (num_wordpieces)]
+ segments = [tokenizer.tokenize(s).merge_dims(-2, -1) for s in segments]
+
+ # Truncate inputs
+ trimmer = tf_text.WaterfallTrimmer(
+ self._seq_length - _NUM_SPECIAL_TOKENS, axis=-1)
+ truncated_segments = trimmer.trim(segments)
+
+ # Combine segments, get segment ids and add special tokens
+ return tf_text.combine_segments(
+ truncated_segments,
+ start_of_sequence_id=self._cls_token,
+ end_of_segment_id=self._sep_token)
+
+ def _bert_preprocess(self, record: Mapping[str, tf.Tensor]):
+ """Parses raw tensors into a dict of tensors to be consumed by the model."""
+ if self._use_next_sentence_label:
+ input_text = record[self._params.text_field_names[0]]
+ # Split sentences
+ sentence_breaker = tf_text.RegexSplitter()
+ sentences = sentence_breaker.split(input_text)
+
+ # Extract next-sentence-prediction labels and segments
+ next_or_random_segment, is_next = (
+ segment_extractor.get_next_sentence_labels(sentences))
+ # merge dims to change shape from [num_docs, (num_segments)] to
+ # [total_num_segments]
+ is_next = is_next.merge_dims(-2, -1)
+
+ # construct segments with shape [(num_sentence)]
+ segments = [
+ sentences.merge_dims(-2, -1),
+ next_or_random_segment.merge_dims(-2, -1)
+ ]
+ else:
+ segments = [record[name] for name in self._params.text_field_names]
+
+ segments_combined, segment_ids = self._tokenize(segments)
+
+ # Dynamic masking
+ item_selector = tf_text.RandomItemSelector(
+ self._max_predictions_per_seq,
+ selection_rate=self._masking_rate,
+ unselectable_ids=[self._cls_token, self._sep_token],
+ shuffle_fn=(tf.identity if self._params.deterministic else None))
+ values_chooser = tf_text.MaskValuesChooser(
+ vocab_size=self._vocab_size, mask_token=self._mask_token)
+ masked_input_ids, masked_lm_positions, masked_lm_ids = (
+ tf_text.mask_language_model(
+ segments_combined,
+ item_selector=item_selector,
+ mask_values_chooser=values_chooser,
+ ))
+
+ # Pad out to fixed shape and get input mask.
+ seq_lengths = {
+ "input_word_ids": self._seq_length,
+ "input_type_ids": self._seq_length,
+ "masked_lm_positions": self._max_predictions_per_seq,
+ "masked_lm_ids": self._max_predictions_per_seq,
+ }
+ model_inputs = {
+ "input_word_ids": masked_input_ids,
+ "input_type_ids": segment_ids,
+ "masked_lm_positions": masked_lm_positions,
+ "masked_lm_ids": masked_lm_ids,
+ }
+ padded_inputs_and_mask = tf.nest.map_structure(tf_text.pad_model_inputs,
+ model_inputs, seq_lengths)
+ model_inputs = {
+ k: padded_inputs_and_mask[k][0] for k in padded_inputs_and_mask
+ }
+ model_inputs["masked_lm_weights"] = tf.cast(
+ padded_inputs_and_mask["masked_lm_ids"][1], tf.float32)
+ model_inputs["input_mask"] = padded_inputs_and_mask["input_word_ids"][1]
+
+ if self._use_next_sentence_label:
+ model_inputs["next_sentence_labels"] = is_next # pyrefly: ignore[unbound-name]
+
+ for name in model_inputs:
+ t = model_inputs[name]
+ if t.dtype == tf.int64:
+ t = tf.cast(t, tf.int32)
+ model_inputs[name] = t
+
+ return model_inputs
+
+ def load(self, input_context: Optional[tf.distribute.InputContext] = None):
+ """Returns a tf.dataset.Dataset."""
+
+ def _batch_docs(dataset, input_context):
+ per_core_doc_batch_size = (
+ input_context.get_per_replica_batch_size(self._params.doc_batch_size)
+ if input_context else self._params.doc_batch_size)
+ return dataset.batch(per_core_doc_batch_size)
+
+ reader = input_reader.InputReader(
+ params=self._params,
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
+ decoder_fn=self._decode if self._params.input_path else None,
+ transform_and_batch_fn=_batch_docs
+ if self._use_next_sentence_label else None,
+ postprocess_fn=self._bert_preprocess)
+ transformed_inputs = reader.read(input_context)
+ per_core_example_batch_size = (
+ input_context.get_per_replica_batch_size(self._params.global_batch_size)
+ if input_context else self._params.global_batch_size)
+ batched_inputs = transformed_inputs.unbatch().batch(
+ per_core_example_batch_size, self._params.drop_remainder)
+ return batched_inputs.prefetch(tf.data.experimental.AUTOTUNE)
diff --git a/official/nlp/data/question_answering_dataloader.py b/official/nlp/data/question_answering_dataloader.py
index 171c0d3b228..7dc7803fc43 100644
--- a/official/nlp/data/question_answering_dataloader.py
+++ b/official/nlp/data/question_answering_dataloader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,10 +13,11 @@
# limitations under the License.
"""Loads dataset for the question answering (e.g, SQuAD) task."""
+import dataclasses
from typing import Mapping, Optional
-import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.common import dataset_fn
from official.core import config_definitions as cfg
from official.core import input_reader
from official.nlp.data import data_loader
@@ -44,6 +45,7 @@ class QADataConfig(cfg.DataConfig):
tokenization: str = 'WordPiece' # WordPiece or SentencePiece
do_lower_case: bool = True
xlnet_format: bool = False
+ file_type: str = 'tfrecord'
@data_loader_factory.register_data_loader_cls(QADataConfig)
@@ -106,5 +108,8 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
def load(self, input_context: Optional[tf.distribute.InputContext] = None):
"""Returns a tf.dataset.Dataset."""
reader = input_reader.InputReader(
- params=self._params, decoder_fn=self._decode, parser_fn=self._parse)
+ params=self._params,
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
+ decoder_fn=self._decode,
+ parser_fn=self._parse)
return reader.read(input_context)
diff --git a/official/nlp/data/question_answering_dataloader_test.py b/official/nlp/data/question_answering_dataloader_test.py
index 9767ef0a7c1..cc0e3ff415b 100644
--- a/official/nlp/data/question_answering_dataloader_test.py
+++ b/official/nlp/data/question_answering_dataloader_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import os
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import question_answering_dataloader
diff --git a/official/nlp/data/sentence_prediction_dataloader.py b/official/nlp/data/sentence_prediction_dataloader.py
index c601b9d72d5..17c207fa123 100644
--- a/official/nlp/data/sentence_prediction_dataloader.py
+++ b/official/nlp/data/sentence_prediction_dataloader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,11 +13,11 @@
# limitations under the License.
"""Loads dataset for the sentence prediction (classification) task."""
+import dataclasses
import functools
from typing import List, Mapping, Optional, Tuple
-import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_hub as hub
from official.common import dataset_fn
diff --git a/official/nlp/data/sentence_prediction_dataloader_test.py b/official/nlp/data/sentence_prediction_dataloader_test.py
index d4f0d8559b1..7b79d0e14a2 100644
--- a/official/nlp/data/sentence_prediction_dataloader_test.py
+++ b/official/nlp/data/sentence_prediction_dataloader_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from sentencepiece import SentencePieceTrainer
from official.nlp.data import sentence_prediction_dataloader as loader
diff --git a/official/nlp/data/sentence_retrieval_lib.py b/official/nlp/data/sentence_retrieval_lib.py
index 947dbb77949..bd16038f47d 100644
--- a/official/nlp/data/sentence_retrieval_lib.py
+++ b/official/nlp/data/sentence_retrieval_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -135,7 +135,7 @@ def generate_sentence_retrevial_tf_record(processor,
for lang_b in [lang_a, "en"]:
if eval_data_output_path:
eval_input_data_examples = processor.get_dev_examples(
- data_dir, os.path.join(path_pattern.format(lang_a, lang_b)))
+ data_dir, os.path.join(path_pattern.format(lang_a, lang_b))) # pyrefly: ignore[unbound-name]
num_eval_data = len(eval_input_data_examples)
logging.info("Processing %d dev examples of %s-en.%s", num_eval_data,
diff --git a/official/nlp/data/squad_lib.py b/official/nlp/data/squad_lib.py
index 2d198e6c1b4..5ebc7a11777 100644
--- a/official/nlp/data/squad_lib.py
+++ b/official/nlp/data/squad_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
import six
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.tools import tokenization
@@ -492,7 +492,7 @@ def _improve_answer_span(doc_tokens, input_start, input_end, tokenizer,
#
# However, this is not always possible. Consider the following:
#
- # Question: What country is the top exporter of electornics?
+ # Question: What country is the top exporter of electronics?
# Context: The Japanese electronics industry is the lagest in the world.
# Answer: Japan
#
@@ -720,7 +720,7 @@ def postprocess_output(all_examples,
start_logit=pred.start_logit,
end_logit=pred.end_logit))
- # if we didn't inlude the empty option in the n-best, inlcude it
+ # if we didn't include the empty option in the n-best, include it
if version_2_with_negative and not xlnet_format:
if "" not in seen_predictions:
nbest.append(
@@ -815,7 +815,7 @@ def get_final_text(pred_text, orig_text, do_lower_case, verbose=False):
# What we really want to return is "Steve Smith".
#
# Therefore, we have to apply a semi-complicated alignment heruistic between
- # `pred_text` and `orig_text` to get a character-to-charcter alignment. This
+ # `pred_text` and `orig_text` to get a character-to-character alignment. This
# can fail in certain cases in which case we just return `orig_text`.
def _strip_spaces(text):
diff --git a/official/nlp/data/squad_lib_sp.py b/official/nlp/data/squad_lib_sp.py
index abd4abfbc09..14127ff4e07 100644
--- a/official/nlp/data/squad_lib_sp.py
+++ b/official/nlp/data/squad_lib_sp.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,7 +26,7 @@
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.tools import tokenization
@@ -63,7 +63,7 @@ def __repr__(self):
tokenization.printable_text(self.question_text))
s += ", paragraph_text: [%s]" % (" ".join(self.paragraph_text))
if self.start_position:
- s += ", start_position: %d" % (self.start_position)
+ s += ", start_position: %d" % (self.start_position,)
if self.start_position:
s += ", end_position: %d" % (self.end_position)
if self.start_position:
@@ -307,8 +307,8 @@ def _lcs_match(max_dist, n=n, m=m):
if (i, j) not in g:
break
if g[(i, j)] == 2:
- orig_to_chartok_index[i] = j
- chartok_to_orig_index[j] = i
+ orig_to_chartok_index[i] = j # pyrefly: ignore[unsupported-operation]
+ chartok_to_orig_index[j] = i # pyrefly: ignore[unsupported-operation]
i, j = i - 1, j - 1
elif g[(i, j)] == 1:
j = j - 1
@@ -466,8 +466,8 @@ def process_class(seg_class):
doc_start = doc_span.start
doc_end = doc_span.start + doc_span.length - 1
out_of_span = False
- if not (tok_start_position >= doc_start and
- tok_end_position <= doc_end):
+ if not (tok_start_position >= doc_start and # pyrefly: ignore[unbound-name]
+ tok_end_position <= doc_end): # pyrefly: ignore[unbound-name]
out_of_span = True
if out_of_span:
# continue
@@ -511,7 +511,7 @@ def process_class(seg_class):
if is_training and not span_is_impossible:
pieces = [
tokenizer.sp_model.IdToPiece(token)
- for token in tokens[start_position:(end_position + 1)]
+ for token in tokens[start_position:(end_position + 1)] # pyrefly: ignore[unsupported-operation]
]
answer_text = tokenizer.sp_model.DecodePieces(pieces)
logging.info("start_position: %d", (start_position))
@@ -557,7 +557,7 @@ def process_class(seg_class):
else:
cnt_pos += 1
- if not is_training and feature:
+ if not is_training and feature: # pyrefly: ignore[unbound-name]
assert batch_size
num_padding = 0
num_examples = unique_id - base_id
@@ -776,7 +776,7 @@ def postprocess_output(all_examples,
start_logit=pred.start_logit,
end_logit=pred.end_logit))
- # if we didn't inlude the empty option in the n-best, include it
+ # if we didn't include the empty option in the n-best, include it
if version_2_with_negative and not xlnet_format:
if "" not in seen_predictions:
nbest.append(
diff --git a/official/nlp/data/tagging_data_lib.py b/official/nlp/data/tagging_data_lib.py
index c73d7108a9b..621daf72a99 100644
--- a/official/nlp/data/tagging_data_lib.py
+++ b/official/nlp/data/tagging_data_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import classifier_data_lib
from official.nlp.tools import tokenization
diff --git a/official/nlp/data/tagging_data_lib_test.py b/official/nlp/data/tagging_data_lib_test.py
index 6a1679f5e28..c6905ae879e 100644
--- a/official/nlp/data/tagging_data_lib_test.py
+++ b/official/nlp/data/tagging_data_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import random
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import tagging_data_lib
from official.nlp.tools import tokenization
diff --git a/official/nlp/data/tagging_dataloader.py b/official/nlp/data/tagging_dataloader.py
index f02d49ab94b..62030c68943 100644
--- a/official/nlp/data/tagging_dataloader.py
+++ b/official/nlp/data/tagging_dataloader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,10 +13,11 @@
# limitations under the License.
"""Loads dataset for the tagging (e.g., NER/POS) task."""
+import dataclasses
from typing import Mapping, Optional
-import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.common import dataset_fn
from official.core import config_definitions as cfg
from official.core import input_reader
from official.nlp.data import data_loader
@@ -29,6 +30,7 @@ class TaggingDataConfig(cfg.DataConfig):
is_training: bool = True
seq_length: int = 128
include_sentence_id: bool = False
+ file_type: str = 'tfrecord'
@data_loader_factory.register_data_loader_cls(TaggingDataConfig)
@@ -81,5 +83,8 @@ def _parse(self, record: Mapping[str, tf.Tensor]):
def load(self, input_context: Optional[tf.distribute.InputContext] = None):
"""Returns a tf.dataset.Dataset."""
reader = input_reader.InputReader(
- params=self._params, decoder_fn=self._decode, parser_fn=self._parse)
+ params=self._params,
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
+ decoder_fn=self._decode,
+ parser_fn=self._parse)
return reader.read(input_context)
diff --git a/official/nlp/data/tagging_dataloader_test.py b/official/nlp/data/tagging_dataloader_test.py
index 3d2be5e97c3..72d5a887fd9 100644
--- a/official/nlp/data/tagging_dataloader_test.py
+++ b/official/nlp/data/tagging_dataloader_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import tagging_dataloader
diff --git a/official/nlp/data/train_sentencepiece.py b/official/nlp/data/train_sentencepiece.py
index c019bc1677d..838c2c9920a 100644
--- a/official/nlp/data/train_sentencepiece.py
+++ b/official/nlp/data/train_sentencepiece.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -28,7 +28,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from sentencepiece import SentencePieceTrainer
@@ -36,7 +36,7 @@
FLAGS = flags.FLAGS
flags.DEFINE_string("output_model_path", None,
- "Path to save the the sentencepiece model.")
+ "Path to save the sentencepiece model.")
flags.mark_flag_as_required("output_model_path")
flags.DEFINE_string("tfds_dir", None, "Directory of the tfds.")
diff --git a/official/nlp/data/wmt_dataloader.py b/official/nlp/data/wmt_dataloader.py
index e801e9d7433..52c683112e8 100644
--- a/official/nlp/data/wmt_dataloader.py
+++ b/official/nlp/data/wmt_dataloader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -34,7 +34,7 @@
from typing import Dict, Optional
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_text as tftxt
from official.core import config_definitions as cfg
from official.core import input_reader
@@ -117,7 +117,7 @@ def _batch_examples(dataset, batch_size, max_length):
# Validates bucket batch sizes.
if any([batch_size <= 0 for batch_size in bucket_batch_sizes]):
raise ValueError(
- 'The token budget, global batch size, is too small to yeild 0 bucket '
+ 'The token budget, global batch size, is too small to yield 0 bucket '
'window: %s' % str(bucket_batch_sizes))
# bucket_id will be a tensor, so convert this list to a tensor as well.
diff --git a/official/nlp/data/wmt_dataloader_test.py b/official/nlp/data/wmt_dataloader_test.py
index 82e56f599d2..2d410ed0b47 100644
--- a/official/nlp/data/wmt_dataloader_test.py
+++ b/official/nlp/data/wmt_dataloader_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from sentencepiece import SentencePieceTrainer
from official.nlp.data import wmt_dataloader
@@ -41,7 +41,7 @@ def _generate_record_file(filepath, src_lines, tgt_lines, unique_id=False):
}
if unique_id:
features['unique_id'] = tf.train.Feature(
- int64_list=tf.train.Int64List(value=[i])),
+ int64_list=tf.train.Int64List(value=[i]))
example = tf.train.Example(
features=tf.train.Features(
feature=features))
diff --git a/official/nlp/docs/data_processing.md b/official/nlp/docs/data_processing.md
new file mode 100644
index 00000000000..049b0c9f6a0
--- /dev/null
+++ b/official/nlp/docs/data_processing.md
@@ -0,0 +1,100 @@
+# TF-NLP Data Processing
+
+## Code locations
+
+Open sourced data processing libraries:
+[tensorflow_models/official/nlp/data/](https://github.com/tensorflow/models/tree/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/nlp/data)
+
+## Preprocess data offline v.s. TFDS
+
+Inside TF-NLP, there are flexible ways to provide training data to the input
+pipeline: 1) using python scripts/beam/flume to process/tokenize the data
+offline; 2) reading the text data directly from
+[TFDS](https://www.tensorflow.org/datasets/api_docs/python/tfds) and using
+[TF.Text](https://www.tensorflow.org/tutorials/tensorflow_text/intro) for
+tokenization and preprocessing inside the tf.data input pipeline.
+
+### Preprocessing scripts
+
+We have implemented data preprocessing for multiple datasets in the following
+python scripts:
+
+* [create_pretraining_data.py](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/nlp/data/create_pretraining_data.py)
+
+* [create_finetuning_data.py](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/nlp/data/create_finetuning_data.py)
+
+Then, the processed files with `tf.Example` protos inside should be specified to
+the `input_path` argument in
+[`DataConfig`](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/core/config_definitions.py#L28).
+
+### TFDS usages
+
+For convenience and consolidation, we built a common
+[input_reader.py](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/core/input_reader.py)
+library to standardize input reading, which has built-in pass for TFDS.
+Specifying the arguments in the
+[`DataConfig`](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/core/config_definitions.py#L28),
+`tfds_name`, `tfds_data_dir` and `tfds_split`, will let the tf.data pipeline
+read from the corresponding dataset inside TFDS.
+
+## DataLoaders
+
+To manage multiple datasets and processing functions, we defined the
+[DataLoader](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/nlp/data/data_loader.py)
+class to work with the
+[data loader factory](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/nlp/data/data_loader_factory.py).
+
+Each dataloader defines the tf.data input pipeline inside the `load` method.
+
+```python
+@abc.abstractmethod
+def load(
+ self,
+ input_context: Optional[tf.distribute.InputContext] = None
+) -> tf.data.Dataset:
+```
+
+Then, the `load` method is called inside each NLP task's `build_input` method
+and the trainer wrap that to create distributed datasets.
+
+```python
+def build_inputs(self, params, input_context=None):
+ """Returns tf.data.Dataset for pretraining."""
+ data_loader = YourDataLoader(params)
+ return data_loader.load(input_context)
+```
+
+By default, in the example above, `params` is the `train_data` or
+`validation_data` field of the `task` field of the experiment config. `params`
+is a type of `DataConfig`.
+
+It is important to note that, for TPU training, the entire `load` method will
+run on the TPU workers and it requires that the function does not access
+resources outside, e.g. the task attributes.
+
+To work with raw text features, we need to use the `DataLoader`s handling the
+text data with TF.Text. You can take the following dataloaders as references:
+
+* [sentence_prediction_dataloader.py](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/nlp/data/sentence_prediction_dataloader.py)
+ for BERT GLUE fine tuning using TFDS with raw text features.
+
+## Speed up training using TF.data service and dynamic sequence length on TPUs
+
+With TF 2.x, we can enable some types of dynamic shapes on TPUs, thanks to TF
+2.x programing model and TPUStrategy/XLA works.
+
+Depending on the data distribution, we are seeing 50% to 90% speed up on typical
+text data for BERT pretraining applications relative to padded static shape
+inputs.
+
+To enable dynamic sequence, we need to use
+`tf data service` for the global bucketizing over
+sequences. To enable it, you can simply add `--enable_tf_data_service` when you
+start experiments.
+
+To pair with tf data service, we need to use the dataloaders that has the
+bucketizing function implemented. You can take the following dataloaders as
+references:
+
+* [pretrain_dynamic_dataloader.py](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/nlp/data/pretrain_dynamic_dataloader.py)
+ for BERT pretraining on the tokenized datasets.
diff --git a/official/nlp/docs/faq.md b/official/nlp/docs/faq.md
new file mode 100644
index 00000000000..b509a9c0902
--- /dev/null
+++ b/official/nlp/docs/faq.md
@@ -0,0 +1,360 @@
+# Frequently Asked Questions
+
+
+
+
+## Introduction
+
+Goal of this document is to capture Frequently Asked Questions (FAQs) related to
+TensorFlow-Models-NLP (TF-NLP). The source of these questions is limited to
+external resources (GitHub, StackOverflow,Google groups etc).
+
+## FAQs of TF-NLP
+
+--------------------------------------------------------------------------------
+
+**Q1: How to cite TF-NLP as the libraries are used for research code bases
+externally?**
+
+If you use TensorFlow Model Garden in your research github repos, please cite
+this repository in your publication. The citation is at the following
+[location](https://github.com/tensorflow/models#citing-tensorflow-model-garden).
+
+--------------------------------------------------------------------------------
+
+**Q2: How to Load NLP Pretrained Models ?**
+
+* [**How to Initialize from Checkpoint:**](https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md#how-to-load-pretrained-models)
+ If you use the TF-NLP training library, you can specify the checkpoint path
+ link directly when launching your job. For example, follow the BERT
+ [fine-tuning command](https://github.com/tensorflow/models/blob/master/official/nlp/docs/train.md#fine-tuning-squad-with-a-pre-trained-bert-checkpoint),
+ to initialize the model from the checkpoint specified by \
+ `--params_override=task.init_checkpoint=PATH_TO_INIT_CKPT`
+
+* [**How to load TF-HUB SavedModel:**](https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md#how-to-load-tf-hub-savedmodel)
+ TF NLP's fine-tuning tasks such as question answering (SQuAD) and sentence
+ prediction (GLUE) support loading a model from TF-HUB. These built-in tasks
+ support a specific task.hub_module_url parameter. To set this parameter,
+ follow the BERT
+ [fine-tuning command](https://github.com/tensorflow/models/blob/master/official/nlp/docs/train.md#fine-tuning-sentence-classification-with-bert-from-tf-hub),
+ and replace `--params_override=task.init_checkpoint=...` with \
+ `--params_override=task.hub_module_url=TF_HUB_URL`.
+
+--------------------------------------------------------------------------------
+
+**Q3: How do I go about changing the pretraining loss functions for BERT ?**
+
+You can change the loss function for the pretraining in the
+[code](https://github.com/tensorflow/models/blob/d93c7e932de27522b2fa3b115f58d06d6f640537/official/nlp/tasks/masked_lm.py#L76)
+here.
+
+--------------------------------------------------------------------------------
+
+**Q4: The
+[transformer code](https://github.com/tensorflow/models/blob/d93c7e932de27522b2fa3b115f58d06d6f640537/official/nlp/modeling/models/seq2seq_transformer.py#L31)
+extends keras.Model. Can I use the constructs like model.fit() for training as
+we do for any tf2/keras model? Are there any tutorials and starting points to
+set up the training and evaluation of a transformer model using TF-NLP?**
+
+Keras Model native `fit()` and `predict()` do not work for the seq2seq
+transformer model. TF model garden uses the workflow defined
+[here](https://github.com/tensorflow/models/blob/d93c7e932de27522b2fa3b115f58d06d6f640537/official/nlp/docs/train.md#model-garden-nlp-common-training-driver).
+\
+The
+[code](https://github.com/tensorflow/models/blob/91d543a1a976e513822f03e63cf7e7d2dc0d92e1/official/nlp/tasks/translation.py)
+defines the translation task.
+
+--------------------------------------------------------------------------------
+
+**Q5: Is there an easy way to set up a model server from a checkpoint (as
+opposed to an exported saved_model)?**
+
+Model server requires saved_model. If you just want to inspect the outputs, this
+[colab](https://www.tensorflow.org/tfmodels/nlp/customize_encoder)
+can help.
+
+--------------------------------------------------------------------------------
+
+**Q6: Training with global batch size (4096) and local batch size (128) on 4x4
+TPUs is very slow. Will the quality change by increasing TPUs to 8x8 with fixed
+local batch size (128) and global batch size (16392)?**
+
+Experiment configuration can be overridden by `--params_override`
+[FLAG](https://github.com/tensorflow/models/blob/master/official/nlp/docs/train.md#overriding-configuration-via-yaml-and-flags)
+through the command line. It only supports scalars. Please find the
+[implementation](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/modeling/hyperparams/params_dict.py#L339)
+here.
+
+--------------------------------------------------------------------------------
+
+**Q7: Training with global batch size (4096) and local batch size (128) on 4x4
+TPUs is very slow. Will the quality change by increasing TPUs to 8x8 with fixed
+local batch size (128) and global batch size (16392)?**
+
+The global batch size should be the key factor. As you increase the batch size,
+you may need to tune the Learning Rate to match the quality of the smaller batch
+size. If the task is retrieval it is recommended using the global softmax. An
+example can be found
+[here](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/modeling/tf_utils.py#L225).
+
+--------------------------------------------------------------------------------
+
+**Q8: In some TF NLP
+[examples](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/tasks/question_answering.py#L15),
+the model output logits are casted into float32: Isn't logits already in the
+format of float?**
+
+For mixed precision training, the activations inside the model could be
+bfloat16/float16 format. The model output logits are casted into float32 to make
+sure the softmax and losses are calculated in float32. This is done to avoid any
+numeric issues that may occur if the intermediate tensor flowing from the
+softmax to the loss is float16 or bfloat16. You can also refer to the
+[mixed precision guide](https://www.tensorflow.org/guide/mixed_precision#building_the_model)
+for more information.
+
+--------------------------------------------------------------------------------
+
+**Q9: Is it possible to use gradient clipping in the optimizer used in the Bert
+encoder? If yes, Is there any sample on its usage ?**
+
+We have the `gradient_clip_norm` argument in AdamW. Also new
+[Keras optimizers](https://www.tensorflow.org/api_docs/python/tf/keras/optimizers/Optimizer)
+offer `global_clipnorm`, `clipnorm` and `clipvalue` as kwargs.
+
+Please refer to the
+[Example](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/configs/experiments/glue_mnli_matched.yaml#L23)
+below:
+
+```
+optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ weight_decay_rate: 0.05
+ gradient_clip_norm: 0.0
+ type: adamw
+```
+
+Please find the bert paper using legacy
+[implementation](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/modeling/optimization/legacy_adamw.py#L78)
+here[[ref]](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/projects/detr/configs/detr.py#L88).
+
+--------------------------------------------------------------------------------
+
+**Q10: I am trying to create an embedding table with 4.7 million rows and 512
+dimensions. However, the `nlp.modeling.layers.OnDeviceEmbedding` fails with the
+following error: UnknownError: Attempting to allocate 4.54G. That was not
+possible. There are 2.94G free.; \
+Is there a way to increase this capacity or alternatives to OnDeviceEmbedding
+that can work in the same framework?**
+
+The embedding with 4.7 million rows and 512 dimensions looks very big. This will
+be placed on the TPU tensor core. \
+Below tips might help:
+
+* Try to reduce the number of rows
+* Consider
+ [mixed_precision_dtype](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/core/config_definitions.py#L147):
+ 'bfloat16' training to reduce memory cost.
+
+--------------------------------------------------------------------------------
+
+**Q11: What is the difference between seq_length in glue_mnli_matched.yaml and
+max_position_embeddings in bert_en_uncased_base.yaml ? Why are they not the
+same?**
+
+`seq_length` is the padded input length and `max_position_embeddings` is the
+size of learned position embeddings. Seq_length value should be always less or
+equal to max_position_embeddings value (seq_length <= max_position_embeddings).
+
+--------------------------------------------------------------------------------
+
+**Q12: While running a model using the tf-nlp framework, it is noticed that when
+the number of validation steps (even by 10) is increased, the experiments get
+much slower. Is that expected?**
+
+This is not expected for 10 validation steps. Recommended tips below:
+
+* Increase the validation interval
+* Use `--add_eval` to start a side-car job for eval
+* Collect xprof for the eval job. It is known that tf2 eager execution is
+ slow.
+
+--------------------------------------------------------------------------------
+
+**Q13: How to load checkpoints for the BERT model? Any recommendations on how to
+deal with the variables mismatch error?**
+
+We recommend using `tf.train.Checkpoint` and manage the objects (including inner
+layers) directly. The details on restoring the encoder weights can be found
+[here](https://www.tensorflow.org/tfmodels/nlp/fine_tune_bert#restore_the_encoder_weights).
+More on TF-NLP checkpoint tutorial is
+[here](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/docs/pretrained_models.md#pre-trained-models)
+\
+The variable mismatch error is due to the classifier_model not equal to the
+threephil model. The recommendation is using the same code and class of
+threephil model to read the checkpoint. The keras functional model cannot
+guarantee the python objects are matched if the model creation code is
+different. \
+More to read as: https://www.tensorflow.org/guide/checkpoint
+
+--------------------------------------------------------------------------------
+
+**Q14: Fail to save Bert2Bert model instance without passing the label input
+i.e. target_id ?**
+
+Bert2Bert needs input_ids, input_mask, segment_ids and target_ids to train. You
+should save the model with all features provided.
+
+If you care about inference and there is no target_id, you should not use Keras
+model.save(). Keras does not support None as inputs. Instead, we directly define
+a tf.Module including the bert2bert core model and save the tf.function using
+tf.saved_model.save() API. Refer
+[example](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/serving/serving_modules.py#L414)
+for the translation task. Usually, the seq2seq model is not friendly to Keras
+assumptions.
+
+--------------------------------------------------------------------------------
+
+**Q15: How to fix the TPU Inference error with the Transformer?**
+
+The potential causes for the error may be having different inputs, and the batch
+size of one of them differs from the rest.
+
+Here are some explanations and troubleshooting tips :
+
+* Resolve the batching issue by implementing
+ signature batching
+* Address the dynamic dimension problem by setting `max_batch_size` and
+ `allowed_batch_sizes` to 1.
+
+--------------------------------------------------------------------------------
+
+**Q16: Are there any models/methods that can improve the latency of the
+feed-forward neural network portion of the transformer encoder block (on CPU and
+GPU)?**
+
+There are `sparsemixture` and `Conditional computation` blocks to speed up. The
+`Block sparse feedforward`
+[layer](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/layers/block_diag_feedforward.py#L15)
+might be promising for performance purposes. This would work nicely on CPU and
+GPU since reshaping ops in this layer are free on CPU/GPUs. It offers speed-up
+for models of similar sizes (a caveat is we observed some quality drop with
+block sparse feedforward in the past).
+
+Refer to
+[Sparse Mixer encoder](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/networks/sparse_mixer.py)
+network and
+[FNet encoder](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/networks/fnet.py)
+network for some more sparsemixture references.
+
+Conditional computation
+is an AI model architecture where specific sections of the computational graph
+are activated based on input conditions. Models following this paradigm
+demonstrate efficiency, especially with increased model capacity or reduced
+inference latency.
+
+Refer
+[ExpandCondense tensor network layer](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/layers/tn_expand_condense.py)
+and
+[Gated linear feedforward layer](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/layers/gated_feedforward.py)
+for FFN blocks. The above mentioned techniques work really well with long
+sequence length.
+
+Please refer to the additional notes below based on your specific use cases.
+
+* For small student models, we used only 1 expert and route much fewer tokens
+ to the FFN expert.
+* We need to set routing_group_size so each routing combines the tokens in
+ multiple sequences and selects for example 1/4 of the tokens.
+* This will work well in the case of distillation or when we can pretrain the
+ model. There will be a quality gap because a lot of tokens skip the FFN
+ computation.
+
+--------------------------------------------------------------------------------
+
+**Q17: How to obtain final layer embeddings from a model? Is there an example?**
+
+Refer to the `call`
+[method](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/networks/bert_encoder.py#L280)
+of the Transformer-based BERT encoder network. The `sequence_output` is the last
+layer embeddings [batch_size, seq len, hidden size].
+
+--------------------------------------------------------------------------------
+
+**Q18: Is it possible to convert public TF hub models like
+[sentence-t5](https://tfhub.dev/google/collections/sentence-t5/1) for TPU use?**
+
+The Inference Converter V2 deploys user-provided function(s) on the XLA device
+(TPU or XLA GPU) and optimizes them.
+
+--------------------------------------------------------------------------------
+
+**Q19: Is it possible to have a dynamic batch size for `edit5` models using
+[sampling modules](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/ops/sampling_module.py)?**
+
+This may depend on the
+[decoding algorithm](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/ops/decoding_module.py#L136)
+for `beam_search`, the source of the issue is at the sample initial time it
+needs to allocate the [batch_size, beam_size, ...] buffer so that batch size is
+fixed. However, note that it may not be easily achievable.
+
+Users can also see that, in
+AutoMUM distillation
+[sampling module](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/ops/sampling_module.py)
+which makes the batch size static.
+
+Possibly, for greedy decoding, it can be done since it doesn't require the
+`beam_size`.
+
+--------------------------------------------------------------------------------
+
+**Q20: Is multi-label tagging distillation supported by text tagging
+distillation?**
+
+Currently the template is just doing basic things of per token binary
+classification. If you intend to perform multi-label classification for each
+token, it shouldn't be overly challenging. It mainly involves adjusting the
+number of classes and switching to a multi-label loss.
+
+--------------------------------------------------------------------------------
+
+**Q21: The TFM
+[Bert](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/networks/bert_encoder.py#L132)
+intentionally utilizes an `OnDeviceEmbedding`. Is it possible to incorporate an
+option to implement `CPU-forced` embedding table ideas by putting embeddings for
+transformer models on CPU to save HBM memory?**
+
+For the optimization, users can just place word embeddings on cpu. Just
+utilizing the `input_word_embeddings` path in
+[BertEncoderV2](https://github.com/tensorflow/models/blob/12cfda05b3fd34a3dd7b3271cd922cd00d0d0c41/official/nlp/modeling/networks/bert_encoder.py#L238)
+class for optimizing HBM usage during serving is sufficient.
+
+--------------------------------------------------------------------------------
+
+**Q22: Is there a possibility of getting TF2 versions of Gemini/MUM? Basically,
+a checkpoint converter and a TF2-variant of instantiating the corresponding
+Transformer?**
+
+[JAX](https://github.com/jax-ml/jax) is the way forward at the moment for
+Gemini.
+
+--------------------------------------------------------------------------------
+
+**Q23:Is it possible to
+perform MLM pretraining in text tagging as well??**
+
+ The MLM functionality in `text_tagging` is
+currently not available.
+
+--------------------------------------------------------------------------------
+
+## Glossary
+
+Acronym | Meaning
+------- | --------------------------
+TFM | Tensorflow Models
+FAQs | Frequently Asked Questions
+TF | TensorFlow
diff --git a/official/nlp/docs/multi_task.md b/official/nlp/docs/multi_task.md
new file mode 100644
index 00000000000..05796ac3263
--- /dev/null
+++ b/official/nlp/docs/multi_task.md
@@ -0,0 +1,97 @@
+# Multi-task library
+
+## Overview
+
+The multi-task library offers lite-weight interfaces and common components to support multi-task
+training and evaluation. It makes no assumption on task types and specific model
+structure details, instead, it is designed to be a scaffold that effectively
+compose single tasks together. Common training scheduling are implemented in the
+default module and it saves possibility for further extension on customized use
+cases.
+
+The multi-task library support:
+
+- *joint* training: individual tasks perform forward passes to get a joint
+ loss and one backward pass happens.
+- *alternative* training: individual tasks perform independent forward and
+ backward pass. The mixture of tasks is controlled by sampling different
+ tasks for train steps.
+
+## Library components
+
+### Interfaces
+
+* [multitask.py](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/multitask.py#L15)
+ serves as a stakeholder of multiple
+ [`Task`](https://github.com/tensorflow/models/blob/master/official/core/base_task.py#L34)
+ instances as well as holding information about multi-task scheduling, such
+ as task weight.
+
+* [base_model.py](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/base_model.py)
+ offers access to each single task's forward computation, where each task is
+ represented as a `tf.keras.Model` instance. Parameter sharing between tasks
+ is left to concrete implementation.
+
+### Common components
+
+* [base_trainer.py](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/base_trainer.py)
+ provides an abstraction to optimize a multi-task model that involves with
+ heterogeneous datasets. By default it conducts joint backward step. Task can
+ be balanced through setting different task weight on corresponding task
+ loss.
+
+* [interleaving_trainer.py](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/interleaving_trainer.py)
+ derives from base trainer and hence shares its housekeeping logic such as
+ loss, metric aggregation and reporting. Unlike the base trainer which
+ conducts joint backward step, interleaving trainer alternates between tasks
+ and effectively mixes single task training step on heterogeneous data sets.
+ Task sampling with respect to a probabilistic distribution will be supported
+ to facilitate task balancing.
+
+* [evaluator.py](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/evaluator.py)
+ conducts a combination of evaluation of each single task. It simply loops
+ through specified tasks and conducts evaluation with corresponding data
+ sets.
+
+* [train_lib.py](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/train_lib.py)
+ puts together model, tasks then trainer and triggers training evaluation
+ execution.
+
+* [configs.py](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/configs.py)
+ provides a top level view on the entire system. Configuration objects are
+ mimicked or composed from corresponding single task components to reuse
+ whenever possible and maintain consistency. For example,
+ [`TaskRoutine`](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/configs.py#L25)
+ effectively reuses
+ [`Task`](https://github.com/tensorflow/models/blob/master/official/core/base_task.py#L34);
+ and
+ [`MultiTaskConfig`](https://github.com/tensorflow/models/blob/master/official/modeling/multitask/configs.py#L34)
+ serves as a similar role of
+ [`TaskConfig`](https://github.com/tensorflow/models/blob/master/official/core/config_definitions.py#L211)
+
+### Notes on single task composability
+
+The library is designed to be able to put together multi-task model by composing
+single task implementations. This is reflected in many aspects:
+
+* Base model interface allows single task's `tf.keras.Model` implementation to
+ be reused, given the shared parts in a potential multi-task case are passed
+ in through constructor. A good example of this is
+ [`BertClassifier`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/models/bert_classifier.py#L24)
+ and
+ [`BertSpanLabeler`](https://github.com/tensorflow/models/blob/master/official/nlp/modeling/models/bert_span_labeler.py),
+ where the backbone network is initialized out of classifier object. Hence a
+ multi-task model that conducts both classification + sequence labeling using
+ a shared backbone encoder could be easily created from existing code.
+
+* Multi-task interface holds a set of Task objects, hence completely reuse the
+ input functions, loss functions, metrics with corresponding aggregation and
+ reduction logic. **Note, under multi-task training situation, the
+ [`build_model()`](https://github.com/tensorflow/models/blob/master/official/core/base_task.py#L144)
+ are not used**, given partially shared structure cannot be specified with
+ only one single task.
+
+* Interleaving trainer works on top of each single task's
+ [`train_step()`](https://github.com/tensorflow/models/blob/master/official/core/base_task.py#L223).
+ This hides the optimization details from each single task and focuses on
+ optimization scheduling and task balancing.
\ No newline at end of file
diff --git a/official/nlp/docs/optimization.md b/official/nlp/docs/optimization.md
new file mode 100644
index 00000000000..aa14eb14410
--- /dev/null
+++ b/official/nlp/docs/optimization.md
@@ -0,0 +1,174 @@
+# Optimizer and Learning Rate Scheduler
+
+This page describes the
+[optimization package](https://github.com/tensorflow/models/tree/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/optimization)
+for Tensorflow Official Models (TFM) which includes optimizers, and learning
+rate schedulers.
+
+## Building Optimizer and LR Scheduler
+
+We use an Optimizer factory class to manage optimizer and learning rate
+creation. Optimizer factory takes a config as an input, and it has member
+functions that are used to build optimizer and learning rate schedule. To create
+an optimizer and a LR schedule through OptimizerFactory, you need to do the
+following:
+
+1. Define optimization config, this includes optimizer, and learning rate
+ schedule.
+2. Initialize the OptimizerFactory instance using the optimization config.
+3. Build the learning rate, and the optimizer using the class member functions.
+
+The following is an example for creating an SGD optimizer with stepwise LR
+scheduler with linear warmup:
+
+```python
+params = {'optimizer': { 'type': 'sgd',
+ 'sgd': {'momentum': 0.9}},
+ 'learning_rate': {'type': 'stepwise',
+ 'stepwise': {
+ 'boundaries': [10000, 20000],
+ 'values': [0.1, 0.01, 0.001]}},
+ 'warmup': {'type': 'linear',
+ 'linear': {'warmup_steps': 500,
+ 'warmup_learning_rate': 0.01}}}
+# Defines optimization config from a dictionary.
+opt_config = optimization.OptimizationConfig(params)
+# Initializes an optimization factory from optimization config.
+opt_factory = optimization.OptimizerFactory(opt_config)
+# Builds the desired learning rate scheduling instance.
+lr = opt_factory.build_learning_rate()
+# Builds the optimizer instance with the desired learning rate schedule.
+optimizer = opt_factory.build_optimizer(lr)
+```
+
+To initialize an OptimizerFactory, `optimizer` and `learning_rate` fields must
+be defined, while `warmup` is an optional field. The field `type` is used to
+define the type of each optimization component. The set of available types are
+explained in details in the following sections.
+
+In the following sections, we explain how to create different optimizers,
+learning rate, and warmup schedulers. We also explain how to add new optimizers,
+or learning rate schedulers.
+
+## Optimizers
+
+The list of supported optimizers can be found
+[here](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/optimization/optimizer_factory.py#L31).
+
+```python
+OPTIMIZERS_CLS = {
+ 'sgd': tf.keras.optimizers.SGD,
+ 'adam': tf.keras.optimizers.Adam,
+ 'adamw': nlp_optimization.AdamWeightDecay,
+ 'lamb': tfa_optimizers.LAMB,
+ 'rmsprop': tf.keras.optimizers.RMSprop
+}
+```
+
+You can specify the type of optimizer to be one of the above using
+[oneof](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/hyperparams/oneof.py)
+config. The available config fields can be found
+[here](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/optimization/configs/optimizer_config.py).
+
+All optimizers support gradient clipping methods: clip by value, clip by norm,
+clip by global norm. To specify which method to use, you need to specify the
+appropriate field list
+[here](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/optimization/configs/optimizer_config.py#L34).
+
+### Example
+
+We will specify an rmsprop optimizer with discounting factor (rho) of 0.9, and
+global norm gradient clipping of 10.0. Below is the config to be used.
+
+```python
+params = {'optimizer': { 'type': 'rmsprop',
+ 'rmsprop': {'rho': 0.9,
+ 'global_clipnorm': 10.0}}}
+```
+
+### Adding a New Optimizer
+
+To add a new optimizer, you need to do the following:
+
+1. Create a
+ [custom](https://www.tensorflow.org/api_docs/python/tf/keras/optimizers/Optimizer#creating_a_custom_optimizer_2)
+ of tf.keras.optimizers.Optimizer.
+2. Add the required config fields under
+ [optimization/configs/optimizer_config.py](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/optimization/configs/optimizer_config.py).
+3. Add the optimizer class to the list of available optimizer classes in
+ (optimizer_factor)[https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/optimization/optimizer_factory.py]
+
+## Learning Rate and Warmup Schedules
+
+Learning rate with an optional warmup can be configured by specifying
+`learning_rate`, and `warmup` fields in optimization config. `learning_rate` is
+a required field, while `warmup` is an optional one. The list of supported
+`learning_rate` and `warmup` schedules can be found
+[here](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/modeling/optimization/optimizer_factory.py#L59).
+
+```python
+LR_CLS = {
+ 'stepwise': tf.keras.optimizers.schedules.PiecewiseConstantDecay,
+ 'polynomial': tf.keras.optimizers.schedules.PolynomialDecay,
+ 'exponential': tf.keras.optimizers.schedules.ExponentialDecay,
+ 'cosine': tf.keras.experimental.CosineDecay,
+ 'power': lr_schedule.DirectPowerDecay,
+}
+
+WARMUP_CLS = {
+ 'linear': lr_schedule.LinearWarmup,
+ 'polynomial': lr_schedule.PolynomialWarmUp
+}
+```
+
+In addition, a `constant` learning rate can be specified.
+
+## How Learning Rate Works
+
+Learning rate takes `step` as an input, and it returns the learning rate value.
+As the training progresses, usually learning rate value decays. Warmup schedule
+is often used to stablize the training. Warmup schedule starts from a low
+learning rate value, and it gradually increases until it reaches the initial
+value for the regular learning rate decay schedule. We combine `learning_rate`
+(lr) with `warmup` (warmup) schedules as follows
+
+* Steps [0, warmup_steps): `learning_rate = warmup(step)`
+* Steps [warmup_steps, train_steps): `learning_rate = lr(step)`
+* We designed the warmup schedule such that final warmup learning rate is
+ inferred from the learning rate schedule (i.e.
+ `learning_rate(warmup_steps) = warmup(warmup_steps)`). Note that, warmup
+ schedule doesn't delay the regular learning rate decay by warmup_steps,
+ instead it replaces it.
+
+Learning rate value is logged every
+[summary_interval](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/core/config_definitions.py#L262).
+If warmup_steps are less that the `summary_interval`, you won't be able to see
+warmup values.
+
+### Example
+
+We want to specify a cosine learning rate decay with decay_steps of 20000, with
+a linear warmup schedule for the first 500 steps.
+
+```python
+params = {'learning_rate': {'type': 'cosine',
+ 'cosine': {'decay_steps': 20000}},
+ 'warmup': {'type': 'linear',
+ 'linear': {'warmup_steps': 500}}}
+```
+
+## Customizing Optimizer Inside Task
+
+Optimizer and learning rate are created inside the
+[task](https://github.com/tensorflow/models/blob/28d972a0b30b628cbb7f67a090ea564c3eda99ea/official/core/base_task.py#L99).
+If different optimizers/learning rate schedulers are needed, they can be defined
+by overriding the class method.
+
+## Important Factors To Consider
+
+* Batch size: Changing batch size usually requires scaling learning rate
+ values, and number of training steps. Make sure that you change appropriate
+ values as batch size changes.
+* Train steps: Train steps is highly correlated with fields such as
+ `decay_steps` for cosine learning rate decay. Changing one without changing
+ the other might result in undesired behavior.
diff --git a/official/nlp/docs/pretrain.md b/official/nlp/docs/pretrain.md
new file mode 100644
index 00000000000..c66de6261bc
--- /dev/null
+++ b/official/nlp/docs/pretrain.md
@@ -0,0 +1,102 @@
+# Model Garden NLP Pre-training Experiments
+
+This user guide describes experiments on pre-training models in TF-NLP. Here we demonstrate how to run the pre-training experiments on TPU/GPU environment. Please refer to the corresponding experiment to get more detailed instructions.
+
+## Pre-train a BERT from scratch
+
+
+
+This example pre-trains a BERT model with Wikipedia and Books datasets used by
+the original BERT paper.
+The [BERT repo](https://github.com/tensorflow/models/blob/master/official/nlp/data/create_pretraining_data.py)
+contains detailed information about the Wikipedia dump and
+[BookCorpus](https://yknzhu.wixsite.com/mbweb). Of course, the pre-training
+recipe is generic and you can apply the same recipe to your own corpus.
+
+Please use the script
+[`create_pretraining_data.py`](https://github.com/tensorflow/models/blob/master/official/nlp/data/create_pretraining_data.py)
+which is essentially branched from [BERT research repo](https://github.com/google-research/bert)
+to get processed pre-training data and it adapts to TF2 symbols and python3
+compatibility.
+
+Running the pre-training script requires an input and output directory, as well
+as a vocab file. Note that `max_seq_length` will need to match the sequence
+length parameter you specify when you run pre-training.
+
+```shell
+export WORKING_DIR='local disk or cloud location'
+export BERT_DIR='local disk or cloud location'
+python models/official/nlp/data/create_pretraining_data.py \
+ --input_file=$WORKING_DIR/input/input.txt \
+ --output_file=$WORKING_DIR/output/tf_examples.tfrecord \
+ --vocab_file=$BERT_DIR/wwm_uncased_L-24_H-1024_A-16/vocab.txt \
+ --do_lower_case=True \
+ --max_seq_length=512 \
+ --max_predictions_per_seq=76 \
+ --masked_lm_prob=0.15 \
+ --random_seed=12345 \
+ --dupe_factor=5
+```
+
+Then, you can update the yaml configuration file, e.g.
+`configs/experiments/wiki_books_pretrain.yaml` to specify your data paths and
+update masking-related hyper parameters to match with your specification for
+the pretraining data. When your data have multiple shards, you can
+use `*` to include multiple files.
+
+To train different BERT sizes, you need to adjust:
+
+```
+model:
+ cls_heads: [{activation: tanh, cls_token_idx: 0, dropout_rate: 0.1, inner_dim: 768, name: next_sentence, num_classes: 2}]
+```
+
+to match the hidden dimensions.
+
+Then, you can start the training and evaluation jobs, which runs the
+[`bert/pretraining`](https://github.com/tensorflow/models/blob/master/official/nlp/configs/pretraining_experiments.py#L51)
+experiment:
+
+```shell
+export OUTPUT_DIR=gs://some_bucket/my_output_dir
+export PARAMS=$PARAMS,runtime.distribution_strategy=tpu
+
+python3 train.py \
+ --experiment=bert/pretraining \
+ --mode=train_and_eval \
+ --model_dir=$OUTPUT_DIR \
+ --config_file=configs/models/bert_en_uncased_base.yaml \
+ --config_file=configs/experiments/wiki_books_pretrain.yaml \
+ --tpu=${TPU_NAME} \
+ --params_override=$PARAMS
+```
+
+## Pre-train BERT MLM with TFDS datasets
+
+This example pre-trains a BERT MLM model with tensorflow_datasets (TFDS) and use tf.text for pre-processing using TPUs. Note that: only wikipedia english corpus is used.
+
+You can start the training and evaluation jobs, which runs the
+[`bert/text_wiki_pretraining`](https://github.com/tensorflow/models/blob/master/official/nlp/configs/pretraining_experiments.py#L88)
+experiment:
+
+```shell
+export OUTPUT_DIR=gs://some_bucket/my_output_dir
+
+# See the following link for more pre-trained checkpoints:
+# https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md
+export BERT_DIR=~/cased_L-12_H-768_A-12
+
+# Override the configurations by FLAGS. Alternatively, you can directly edit
+# `configs/experiments/wiki_tfds_pretrain.yaml` to specify corresponding fields.
+export PARAMS=$PARAMS,task.validation_data.vocab_file_path=$BERT_DIR/vocab.txt
+export PARAMS=$PARAMS,task.train_data.vocab_file_path=$BERT_DIR/vocab.txt
+export PARAMS=$PARAMS,runtime.distribution_strategy=tpu
+
+python3 train.py \
+ --experiment=bert/text_wiki_pretraining \
+ --mode=train_and_eval \
+ --model_dir=$OUTPUT_DIR \
+ --config_file=configs/experiments/wiki_tfds_pretrain.yaml \
+ --tpu=${TPU_NAME} \
+ --params_override=$PARAMS
+```
\ No newline at end of file
diff --git a/official/nlp/docs/pretrained_models.md b/official/nlp/docs/pretrained_models.md
index 581548ff5c8..570a4c88e8a 100644
--- a/official/nlp/docs/pretrained_models.md
+++ b/official/nlp/docs/pretrained_models.md
@@ -1,5 +1,15 @@
# Pre-trained Models
+⚠️ Disclaimer: Checkpoints are based on training with publicly available datasets.
+Some datasets contain limitations, including non-commercial use limitations.
+Please review the terms and conditions made available by third parties before
+using the datasets provided. Checkpoints are licensed under
+[Apache 2.0](https://github.com/tensorflow/models/blob/master/LICENSE).
+
+⚠️ Disclaimer: Datasets hyperlinked from this page are not owned or distributed
+by Google. Such datasets are made available by third parties. Please review the
+terms and conditions made available by the third parties before using the data.
+
We provide a large collection of baselines and checkpoints for NLP pre-trained
models.
@@ -7,8 +17,9 @@ models.
### How to Initialize from Checkpoint
-**Note:** TF-HUB/Savedmodel is the preferred way to distribute models as it is
-self-contained. Please consider using TF-HUB for finetuning tasks first.
+**Note:** TF-HUB/Kaggle-Savedmodel is the preferred way to distribute models as
+it is self-contained. Please consider using TF-HUB/Kaggle for finetuning tasks
+first.
If you use the [NLP training library](train.md),
you can specify the checkpoint path link directly when launching your job. For
@@ -20,11 +31,11 @@ python3 train.py \
--params_override=task.init_checkpoint=PATH_TO_INIT_CKPT
```
-### How to load TF-HUB SavedModel
+### How to load TF-HUB/Kaggle SavedModel
Finetuning tasks such as question answering (SQuAD) and sentence
-prediction (GLUE) support loading a model from TF-HUB. These built-in tasks
-support a specific `task.hub_module_url` parameter. To set this parameter,
+prediction (GLUE) support loading a model from TF-HUB/Kaggle. These built-in
+tasks support a specific `task.hub_module_url` parameter. To set this parameter,
replace `--params_override=task.init_checkpoint=...` with
`--params_override=task.hub_module_url=TF_HUB_URL`, like below:
@@ -45,7 +56,7 @@ in order to keep consistent with BERT paper.
### Checkpoints
-Model | Configuration | Training Data | Checkpoint & Vocabulary | TF-HUB SavedModels
+Model | Configuration | Training Data | Checkpoint & Vocabulary | Kaggle SavedModels
---------------------------------------- | :--------------------------: | ------------: | ----------------------: | ------:
BERT-base uncased English | uncased_L-12_H-768_A-12 | Wiki + Books | [uncased_L-12_H-768_A-12](https://storage.googleapis.com/tf_model_garden/nlp/bert/v3/uncased_L-12_H-768_A-12.tar.gz) | [`BERT-Base, Uncased`](https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/)
BERT-base cased English | cased_L-12_H-768_A-12 | Wiki + Books | [cased_L-12_H-768_A-12](https://storage.googleapis.com/tf_model_garden/nlp/bert/v3/cased_L-12_H-768_A-12.tar.gz) | [`BERT-Base, Cased`](https://tfhub.dev/tensorflow/bert_en_cased_L-12_H-768_A-12/)
@@ -65,7 +76,7 @@ We also have pretrained BERT models with variants in both network architecture
and training methodologies. These models achieve higher downstream accuracy
scores.
-Model | Configuration | Training Data | TF-HUB SavedModels | Comment
+Model | Configuration | Training Data | Kaggle SavedModels | Comment
-------------------------------- | :----------------------: | -----------------------: | ------------------------------------------------------------------------------------: | ------:
BERT-base talking heads + ggelu | uncased_L-12_H-768_A-12 | Wiki + Books | [talkheads_ggelu_base](https://tfhub.dev/tensorflow/talkheads_ggelu_bert_en_base/1) | BERT-base trained with [talking heads attention](https://arxiv.org/abs/2003.02436) and [gated GeLU](https://arxiv.org/abs/2002.05202).
BERT-large talking heads + ggelu | uncased_L-24_H-1024_A-16 | Wiki + Books | [talkheads_ggelu_large](https://tfhub.dev/tensorflow/talkheads_ggelu_bert_en_large/1) | BERT-large trained with [talking heads attention](https://arxiv.org/abs/2003.02436) and [gated GeLU](https://arxiv.org/abs/2002.05202).
@@ -87,13 +98,12 @@ ALBERT repository.
### Checkpoints
-Model | Training Data | Checkpoint & Vocabulary | TF-HUB SavedModels
+Model | Training Data | Checkpoint & Vocabulary | Kaggle SavedModels
---------------------------------------- | ------------: | ----------------------: | ------:
-ALBERT-base English | Wiki + Books | [`ALBERT Base`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_base.tar.gz) | https://tfhub.dev/tensorflow/albert_en_base/3
-ALBERT-large English | Wiki + Books | [`ALBERT Large`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_large.tar.gz) | https://tfhub.dev/tensorflow/albert_en_large/3
-ALBERT-xlarge English | Wiki + Books | [`ALBERT XLarge`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_xlarge.tar.gz) | https://tfhub.dev/tensorflow/albert_en_xlarge/3
-ALBERT-xxlarge English | Wiki + Books | [`ALBERT XXLarge`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_xxlarge.tar.gz) | https://tfhub.dev/tensorflow/albert_en_xxlarge/3
-
+ALBERT-base English | Wiki + Books | [`ALBERT Base`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_base.tar.gz) | [albert_en_base](https://tfhub.dev/tensorflow/albert_en_base/3)
+ALBERT-large English | Wiki + Books | [`ALBERT Large`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_large.tar.gz) | [albert_en_large](https://tfhub.dev/tensorflow/albert_en_large/3)
+ALBERT-xlarge English | Wiki + Books | [`ALBERT XLarge`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_xlarge.tar.gz) | [albert_en_xlarge](https://tfhub.dev/tensorflow/albert_en_xlarge/3)
+ALBERT-xxlarge English | Wiki + Books | [`ALBERT XXLarge`](https://storage.googleapis.com/tf_model_garden/nlp/albert/albert_xxlarge.tar.gz) | [albert_en_xxlarge](https://tfhub.dev/tensorflow/albert_en_xxlarge/3)
## ELECTRA
diff --git a/official/nlp/docs/tfhub.md b/official/nlp/docs/tfhub.md
index c6fe9a2f8f4..fb00ce3d150 100644
--- a/official/nlp/docs/tfhub.md
+++ b/official/nlp/docs/tfhub.md
@@ -72,10 +72,10 @@ encoder_inputs = dict(
)
encoder_outputs = encoder(encoder_inputs)
assert encoder_outputs.keys() == {
- "pooled_output", # Shape [batch_size, width], dtype=float32
- "default", # Alias for "pooled_output" (aligns with other models)
- "sequence_output", # Shape [batch_size, seq_length, width], dtype=float32
- "encoder_outputs", # List of Tensors with outputs of all transformer layers
+ "pooled_output", # Shape [batch_size, width], dtype=float32
+ "default", # Alias for "pooled_output" (aligns with other models)
+ "sequence_output", # Shape [batch_size, seq_length, width], dtype=float32
+ "encoder_outputs", # List of Tensors with outputs of all transformer layers
}
```
@@ -170,10 +170,10 @@ mlm_inputs = dict(
)
mlm_outputs = encoder.mlm(mlm_inputs)
assert mlm_outputs.keys() == {
- "pooled_output", # Shape [batch, width], dtype=float32
- "sequence_output", # Shape [batch, seq_length, width], dtype=float32
- "encoder_outputs", # List of Tensors with outputs of all transformer layers
- "mlm_logits" # Shape [batch, num_predictions, vocab_size], dtype=float32
+ "pooled_output", # Shape [batch, width], dtype=float32
+ "sequence_output", # Shape [batch, seq_length, width], dtype=float32
+ "encoder_outputs", # List of Tensors with outputs of all transformer layers
+ "mlm_logits" # Shape [batch, num_predictions, vocab_size], dtype=float32
}
```
@@ -246,9 +246,9 @@ preprocessor = hub.load(...)
text_input = ... # Shape [batch_size], dtype=tf.string
encoder_inputs = preprocessor(text_input, seq_length=seq_length)
assert encoder_inputs.keys() == {
- "input_word_ids", # Shape [batch_size, seq_length], dtype=int32
- "input_mask", # Shape [batch_size, seq_length], dtype=int32
- "input_type_ids" # Shape [batch_size, seq_length], dtype=int32
+ "input_word_ids", # Shape [batch_size, seq_length], dtype=int32
+ "input_mask", # Shape [batch_size, seq_length], dtype=int32
+ "input_type_ids" # Shape [batch_size, seq_length], dtype=int32
}
```
diff --git a/official/nlp/docs/train.md b/official/nlp/docs/train.md
index 051666de9e6..d53fa30dca5 100644
--- a/official/nlp/docs/train.md
+++ b/official/nlp/docs/train.md
@@ -46,7 +46,7 @@ trains a BERT-base model on GLUE/MNLI-matched which is a sentence prediction
task.
```shell
-PARAMS=runtime.distribution_strategy=mirrored # Train no GPU
+PARAMS=runtime.distribution_strategy=mirrored # Train on GPU
PARAMS=${PARAMS},task.train_data.input_path=/path-to-your-training-data/
python3 train.py \
@@ -101,11 +101,67 @@ export PYTHONPATH=$PYTHONPATH:/path/to/models
pip3 install --user -r official/requirements.txt
```
+### Fine-tuning SQuAD with a pre-trained BERT checkpoint
+
+This example fine-tunes a pre-trained BERT checkpoint on the
+Stanford Question Answering Dataset (SQuAD) using TPUs.
+The [SQuAD website](https://rajpurkar.github.io/SQuAD-explorer/) contains
+detailed information about the SQuAD datasets and evaluation. After downloading
+the SQuAD datasets and the [pre-trained BERT checkpoints](https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md),
+you can run the following command to prepare the `tf_record` files:
+
+```shell
+export SQUAD_DIR=~/squad
+export BERT_DIR=~/uncased_L-12_H-768_A-12
+export OUTPUT_DATA_DIR=gs://some_bucket/datasets
+
+python3 create_finetuning_data.py \
+ --squad_data_file=${SQUAD_DIR}/train-v1.1.json \
+ --vocab_file=${BERT_DIR}/vocab.txt \
+ --train_data_output_path=${OUTPUT_DATA_DIR}/train.tf_record \
+ --meta_data_file_path=${OUTPUT_DATA_DIR}/squad_meta_data \
+ --fine_tuning_task_type=squad --max_seq_length=384
+```
+
+Note: To create fine-tuning data with SQuAD 2.0, you need to add flag `--version_2_with_negative=True`.
+
+Then, you can start the training and evaluation jobs:
+
+```shell
+export SQUAD_DIR=~/squad
+export INPUT_DATA_DIR=gs://some_bucket/datasets
+export OUTPUT_DIR=gs://some_bucket/my_output_dir
+
+# See the following link for more pre-trained checkpoints:
+# https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md
+export BERT_DIR=~/uncased_L-12_H-768_A-12
+
+# Override the configurations by FLAGS. Alternatively, you can directly edit
+# `configs/experiments/squad_v1.1.yaml` to specify corresponding fields.
+# Also note that the training data is the pre-processed tf_record file, while
+# the validation file is the raw json file.
+export PARAMS=task.train_data.input_path=$INPUT_DATA_DIR/train.tf_record
+export PARAMS=$PARAMS,task.validation_data.input_path=$SQUAD_DIR/dev-v1.1.json
+export PARAMS=$PARAMS,task.validation_data.vocab_file=$BERT_DIR/vocab.txt
+export PARAMS=$PARAMS,task.init_checkpoint=$BERT_DIR/bert_model.ckpt
+export PARAMS=$PARAMS,runtime.distribution_strategy=tpu
+
+python3 train.py \
+ --experiment=bert/squad \
+ --mode=train_and_eval \
+ --model_dir=$OUTPUT_DIR \
+ --config_file=configs/models/bert_en_uncased_base.yaml \
+ --config_file=configs/experiments/squad_v1.1.yaml \
+ --tpu=${TPU_NAME} \
+ --params_override=$PARAMS
+
+```
+
### Fine-tuning Sentence Classification with BERT from TF-Hub
-
-This example fine-tunes BERT-base from TF-Hub on the the Multi-Genre Natural
+
+This example fine-tunes BERT-base from TF-Hub on the Multi-Genre Natural
Language Inference (MultiNLI) corpus using TPUs.
Firstly, you can prepare the fine-tuning data using
@@ -133,8 +189,7 @@ python3 data/create_finetuning_data.py \
```
Resulting training and evaluation datasets in `tf_record` format will be later
-passed to [train.py](train.py). We will support to read dataset from
-tensorflow_datasets (TFDS) and use tf.text for pre-processing soon.
+passed to [train.py](train.py).
Then you can execute the following commands to start the training and evaluation
job.
@@ -168,67 +223,3 @@ python3 train.py \
You can monitor the training progress in the console and find the output
models in `$OUTPUT_DIR`.
-
-
-
-### Fine-tuning SQuAD with a pre-trained BERT checkpoint
-
-
-
-This example fine-tunes a pre-trained BERT checkpoint on the
-Stanford Question Answering Dataset (SQuAD) using TPUs.
-The [SQuAD website](https://rajpurkar.github.io/SQuAD-explorer/) contains
-detailed information about the SQuAD datasets and evaluation. After downloading
-the SQuAD datasets and the [pre-trained BERT checkpoints](https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md),
-you can run the following command to prepare the `tf_record` files:
-
-```shell
-export SQUAD_DIR=~/squad
-export BERT_DIR=~/uncased_L-12_H-768_A-12
-export OUTPUT_DATA_DIR=gs://some_bucket/datasets
-
-python3 create_finetuning_data.py \
- --squad_data_file=${SQUAD_DIR}/train-v1.1.json \
- --vocab_file=${BERT_DIR}/vocab.txt \
- --train_data_output_path=${OUTPUT_DATA_DIR}/train.tf_record \
- --meta_data_file_path=${OUTPUT_DATA_DIR}/squad_meta_data \
- --fine_tuning_task_type=squad --max_seq_length=384
-```
-
-Note: To create fine-tuning data with SQuAD 2.0, you need to add flag `--version_2_with_negative=True`.
-
-Then, you can start the training and evaluation jobs:
-
-```shell
-export SQUAD_DIR=~/squad
-export INPUT_DATA_DIR=gs://some_bucket/datasets
-export OUTPUT_DIR=gs://some_bucket/my_output_dir
-
-# See the following link for more pre-trained checkpoints:
-# https://github.com/tensorflow/models/blob/master/official/nlp/docs/pretrained_models.md
-export BERT_DIR=~/uncased_L-12_H-768_A-12
-
-# Override the configurations by FLAGS. Alternatively, you can directly edit
-# `configs/experiments/squad_v1.1.yaml` to specify corresponding fields.
-# Also note that the training data is the pre-processed tf_record file, while
-# the validation file is the raw json file.
-export PARAMS=task.train_data.input_path=$INPUT_DATA_DIR/train.tf_record
-export PARAMS=$PARAMS,task.validation_data.input_path=$SQUAD_DIR/dev-v1.1.json
-export PARAMS=$PARAMS,task.validation_data.vocab_file=$BERT_DIR/vocab.txt
-export PARAMS=$PARAMS,task.init_checkpoint=$BERT_DIR/bert_model.ckpt
-export PARAMS=$PARAMS,runtime.distribution_strategy=tpu
-
-python3 train.py \
- --experiment=bert/squad \
- --mode=train_and_eval \
- --model_dir=$OUTPUT_DIR \
- --config_file=configs/models/bert_en_uncased_base.yaml \
- --config_file=configs/experiments/squad_v1.1.yaml \
- --tpu=${TPU_NAME} \
- --params_override=$PARAMS
-
-```
-
-
-
-Note: More examples about pre-training will come soon.
diff --git a/official/nlp/finetuning/binary_helper.py b/official/nlp/finetuning/binary_helper.py
index 10ad91e9377..1a12712c4bc 100644
--- a/official/nlp/finetuning/binary_helper.py
+++ b/official/nlp/finetuning/binary_helper.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from typing import Any, Dict, List, Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.modeling import hyperparams
diff --git a/official/nlp/finetuning/glue/flags.py b/official/nlp/finetuning/glue/flags.py
index 0ad161bc662..c7279d4ef50 100644
--- a/official/nlp/finetuning/glue/flags.py
+++ b/official/nlp/finetuning/glue/flags.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,16 +25,6 @@ def define_flags():
# ===========================================================================
# Glue binary flags.
# ===========================================================================
- flags.DEFINE_enum(
- 'mode', 'train_eval_and_predict',
- ['train_eval_and_predict', 'train_eval', 'predict'],
- 'The mode to run the binary. If `train_eval_and_predict` '
- 'it will (1) train on the training data and (2) evaluate on '
- 'the validation data and (3) finally generate predictions '
- 'on the prediction data; if `train_eval`, it will only '
- 'run training and evaluation; if `predict`, it will only '
- 'run prediction using the model in `model_dir`.')
-
flags.DEFINE_enum('task_name', None, [
'AX', 'COLA', 'MNLI', 'MRPC', 'QNLI', 'QQP', 'RTE', 'SST-2', 'STS-B',
'WNLI'
diff --git a/official/nlp/finetuning/glue/run_glue.py b/official/nlp/finetuning/glue/run_glue.py
index 54d45150b4d..cbf0eb414d0 100644
--- a/official/nlp/finetuning/glue/run_glue.py
+++ b/official/nlp/finetuning/glue/run_glue.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
# Imports registered experiment configs.
@@ -46,6 +46,22 @@
'used when creating the Cloud TPU, or a grpc://ip.address.of.tpu:8470 url.')
flags.DEFINE_integer('num_gpus', 1, 'The number of GPUs to use at each worker.')
+_MODE = flags.DEFINE_enum(
+ 'mode', 'train_eval_and_predict',
+ ['train_eval_and_predict', 'train_eval', 'predict'],
+ 'The mode to run the binary. If `train_eval_and_predict` '
+ 'it will (1) train on the training data and (2) evaluate on '
+ 'the validation data and (3) finally generate predictions '
+ 'on the prediction data; if `train_eval`, it will only '
+ 'run training and evaluation; if `predict`, it will only '
+ 'run prediction using the model in `model_dir`.')
+
+# TODO(kitsing) The `params_override` flag is currently not being used.
+# Only declared to make xm_job_3p.XMTPUJob happy.
+_PARAMS_OVERRIDE = flags.DEFINE_string(
+ 'params_override', '', 'Overridden parameters.'
+)
+
FLAGS = flags.FLAGS
EXPERIMENT_TYPE = 'bert/sentence_prediction'
@@ -55,9 +71,9 @@
'AX': 'matthews_corrcoef',
'COLA': 'matthews_corrcoef',
'MNLI': 'cls_accuracy',
- 'MRPC': 'cls_accuracy',
+ 'MRPC': 'f1',
'QNLI': 'cls_accuracy',
- 'QQP': 'cls_accuracy',
+ 'QQP': 'f1',
'RTE': 'cls_accuracy',
'SST-2': 'cls_accuracy',
'STS-B': 'pearson_spearman_corr',
@@ -93,11 +109,16 @@ def _override_exp_config_by_flags(exp_config, input_meta_data):
binary_helper.override_sentence_prediction_task_config,
num_classes=input_meta_data['num_labels'],
metric_type='matthews_corrcoef')
- elif FLAGS.task_name in ('MNLI', 'MRPC', 'QNLI', 'QQP', 'RTE', 'SST-2',
+ elif FLAGS.task_name in ('MNLI', 'QNLI', 'RTE', 'SST-2',
'WNLI'):
override_task_cfg_fn = functools.partial(
binary_helper.override_sentence_prediction_task_config,
num_classes=input_meta_data['num_labels'])
+ elif FLAGS.task_name in ('QQP', 'MRPC'):
+ override_task_cfg_fn = functools.partial(
+ binary_helper.override_sentence_prediction_task_config,
+ metric_type='f1',
+ num_classes=input_meta_data['num_labels'])
elif FLAGS.task_name in ('STS-B',):
override_task_cfg_fn = functools.partial(
binary_helper.override_sentence_prediction_task_config,
@@ -229,7 +250,7 @@ def main(argv):
with distribution_strategy.scope():
task = None
- if 'train_eval' in FLAGS.mode:
+ if 'train_eval' in _MODE.value:
logging.info('Starting training and eval...')
logging.info('Model dir: %s', FLAGS.model_dir)
@@ -245,7 +266,7 @@ def main(argv):
params=exp_config,
model_dir=FLAGS.model_dir)
- if 'predict' in FLAGS.mode:
+ if 'predict' in _MODE.value:
logging.info('Starting predict...')
# When mode is `predict`, `task` will be None.
if task is None:
diff --git a/official/nlp/finetuning/superglue/flags.py b/official/nlp/finetuning/superglue/flags.py
index 68457ea379b..f1203dc2acc 100644
--- a/official/nlp/finetuning/superglue/flags.py
+++ b/official/nlp/finetuning/superglue/flags.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/finetuning/superglue/run_superglue.py b/official/nlp/finetuning/superglue/run_superglue.py
index 773abbd0315..faee582b09c 100644
--- a/official/nlp/finetuning/superglue/run_superglue.py
+++ b/official/nlp/finetuning/superglue/run_superglue.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
# Imports registered experiment configs.
diff --git a/official/nlp/metrics/__init__.py b/official/nlp/metrics/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/nlp/metrics/__init__.py
+++ b/official/nlp/metrics/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/metrics/bleu.py b/official/nlp/metrics/bleu.py
index 01c6ae5faa2..003f2cdb9ed 100644
--- a/official/nlp/metrics/bleu.py
+++ b/official/nlp/metrics/bleu.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,7 +25,7 @@
import unicodedata
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class UnicodeRegex(object):
@@ -149,15 +149,15 @@ def compute_bleu(reference_corpus,
for i in range(0, max_order):
if possible_matches_by_order[i] > 0:
- precisions[i] = float(matches_by_order[i]) / possible_matches_by_order[i]
+ precisions[i] = float(matches_by_order[i]) / possible_matches_by_order[i] # pyrefly: ignore[unsupported-operation]
if matches_by_order[i] > 0:
- precisions[i] = float(
+ precisions[i] = float( # pyrefly: ignore[unsupported-operation]
matches_by_order[i]) / possible_matches_by_order[i]
else:
smooth *= 2
- precisions[i] = 1.0 / (smooth * possible_matches_by_order[i])
+ precisions[i] = 1.0 / (smooth * possible_matches_by_order[i]) # pyrefly: ignore[unsupported-operation]
else:
- precisions[i] = 0.0
+ precisions[i] = 0.0 # pyrefly: ignore[unsupported-operation]
if max(precisions) > 0:
p_log_sum = sum(math.log(p) for p in precisions if p)
diff --git a/official/nlp/metrics/bleu_test.py b/official/nlp/metrics/bleu_test.py
index 9097ad8fd2c..f39ea95f363 100644
--- a/official/nlp/metrics/bleu_test.py
+++ b/official/nlp/metrics/bleu_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import tempfile
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.metrics import bleu
diff --git a/official/nlp/modeling/README.md b/official/nlp/modeling/README.md
index 6a31447d577..05a90248b6d 100644
--- a/official/nlp/modeling/README.md
+++ b/official/nlp/modeling/README.md
@@ -20,7 +20,7 @@ examples.
* [`losses`](losses) contains common loss computation used in NLP tasks.
Please see the colab
-[nlp_modeling_library_intro.ipynb](https://colab.sandbox.google.com/github/tensorflow/models/blob/master/official/colab/nlp/nlp_modeling_library_intro.ipynb)
+[NLP modeling library intro.ipynb](https://colab.sandbox.google.com/github/tensorflow/models/blob/master/docs/nlp/index.ipynb)
for how to build transformer-based NLP models using above primitives.
Besides the pre-defined primitives, it also provides scaffold classes to allow
@@ -43,7 +43,7 @@ custom hidden layer (which will replace the Transformer instantiation in the
encoder).
Please see the colab
-[customize_encoder.ipynb](https://colab.sandbox.google.com/github/tensorflow/models/blob/master/official/colab/nlp/customize_encoder.ipynb)
+[customize_encoder.ipynb](https://colab.sandbox.google.com/github/tensorflow/models/blob/master/docs/nlp/customize_encoder.ipynb)
for how to use scaffold classes to build noval achitectures.
BERT and ALBERT models in this repo are implemented using this library.
diff --git a/official/nlp/modeling/__init__.py b/official/nlp/modeling/__init__.py
index 6159d986d39..2c5dee2ba02 100644
--- a/official/nlp/modeling/__init__.py
+++ b/official/nlp/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,8 +14,8 @@
"""NLP Modeling Library.
-This library provides a set of Keras primitives (`tf.keras.Layer` and
-`tf.keras.Model`) that can be assembled into transformer-based models.
+This library provides a set of Keras primitives (`tf_keras.Layer` and
+`tf_keras.Model`) that can be assembled into transformer-based models.
They are flexible, validated, interoperable, and both TF1 and TF2 compatible.
"""
from official.nlp.modeling import layers
diff --git a/official/nlp/modeling/layers/README.md b/official/nlp/modeling/layers/README.md
index fb4069abd5b..54c00661ffd 100644
--- a/official/nlp/modeling/layers/README.md
+++ b/official/nlp/modeling/layers/README.md
@@ -13,7 +13,7 @@ assemble new `tf.keras` layers or models.
["Big Bird: Transformers for Longer Sequences"](https://arxiv.org/abs/2007.14062).
* [CachedAttention](attention.py) implements an attention layer with cache
- used for auto-agressive decoding.
+ used for auto-aggressive decoding.
* [KernelAttention](kernel_attention.py) implements a group of attention
mechansim that express the self-attention as a linear dot-product of
diff --git a/official/nlp/modeling/layers/__init__.py b/official/nlp/modeling/layers/__init__.py
index 268e1708d72..92829a99dff 100644
--- a/official/nlp/modeling/layers/__init__.py
+++ b/official/nlp/modeling/layers/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,11 +17,15 @@
They can be used to assemble new `tf.keras` layers or models.
"""
# pylint: disable=wildcard-import
+
+from official.nlp.modeling.layers import util
from official.nlp.modeling.layers.attention import *
from official.nlp.modeling.layers.bigbird_attention import BigBirdAttention
from official.nlp.modeling.layers.bigbird_attention import BigBirdMasks
from official.nlp.modeling.layers.block_diag_feedforward import BlockDiagFeedforward
+from official.nlp.modeling.layers.block_sparse_attention import MultiHeadAttention as BlockSparseAttention
from official.nlp.modeling.layers.cls_head import *
+from official.nlp.modeling.layers.factorized_embedding import FactorizedEmbedding
from official.nlp.modeling.layers.gated_feedforward import GatedFeedforward
from official.nlp.modeling.layers.gaussian_process import RandomFeatureGaussianProcess
from official.nlp.modeling.layers.kernel_attention import KernelAttention
@@ -29,11 +33,25 @@
from official.nlp.modeling.layers.masked_lm import MaskedLM
from official.nlp.modeling.layers.masked_softmax import MaskedSoftmax
from official.nlp.modeling.layers.mat_mul_with_margin import MatMulWithMargin
+from official.nlp.modeling.layers.mixing import FourierTransformLayer
+from official.nlp.modeling.layers.mixing import HartleyTransformLayer
+from official.nlp.modeling.layers.mixing import LinearTransformLayer
+from official.nlp.modeling.layers.mixing import MixingMechanism
from official.nlp.modeling.layers.mobile_bert_layers import MobileBertEmbedding
from official.nlp.modeling.layers.mobile_bert_layers import MobileBertMaskedLM
from official.nlp.modeling.layers.mobile_bert_layers import MobileBertTransformer
+from official.nlp.modeling.layers.moe import ExpertsChooseMaskedRouter
+from official.nlp.modeling.layers.moe import FeedForwardExperts
+from official.nlp.modeling.layers.moe import MoeLayer
+from official.nlp.modeling.layers.moe import MoeLayerWithBackbone
from official.nlp.modeling.layers.multi_channel_attention import *
+from official.nlp.modeling.layers.multi_query_attention import MultiHeadAttention as MultiQueryAttention
from official.nlp.modeling.layers.on_device_embedding import OnDeviceEmbedding
+from official.nlp.modeling.layers.pack_optimization import PackBertEmbeddings
+from official.nlp.modeling.layers.pack_optimization import StridedReZeroTransformer
+from official.nlp.modeling.layers.pack_optimization import StridedTransformerEncoderBlock
+from official.nlp.modeling.layers.pack_optimization import StridedTransformerScaffold
+from official.nlp.modeling.layers.per_dim_scale_attention import PerDimScaleAttention
from official.nlp.modeling.layers.position_embedding import PositionEmbedding
from official.nlp.modeling.layers.position_embedding import RelativePositionBias
from official.nlp.modeling.layers.position_embedding import RelativePositionEmbedding
@@ -43,7 +61,7 @@
from official.nlp.modeling.layers.reuse_transformer import ReuseTransformer
from official.nlp.modeling.layers.rezero_transformer import ReZeroTransformer
from official.nlp.modeling.layers.routing import *
-from official.nlp.modeling.layers.self_attention_mask import SelfAttentionMask
+from official.nlp.modeling.layers.self_attention_mask import *
from official.nlp.modeling.layers.spectral_normalization import *
from official.nlp.modeling.layers.talking_heads_attention import TalkingHeadsAttention
from official.nlp.modeling.layers.text_layers import BertPackInputs
@@ -51,7 +69,8 @@
from official.nlp.modeling.layers.text_layers import FastWordpieceBertTokenizer
from official.nlp.modeling.layers.text_layers import SentencepieceTokenizer
from official.nlp.modeling.layers.tn_transformer_expand_condense import TNTransformerExpandCondense
-from official.nlp.modeling.layers.transformer import *
+from official.nlp.modeling.layers.transformer import Transformer
+from official.nlp.modeling.layers.transformer import TransformerDecoderBlock
from official.nlp.modeling.layers.transformer_encoder_block import TransformerEncoderBlock
from official.nlp.modeling.layers.transformer_scaffold import TransformerScaffold
from official.nlp.modeling.layers.transformer_xl import TransformerXL
diff --git a/official/nlp/modeling/layers/attention.py b/official/nlp/modeling/layers/attention.py
index 35dc1eed53b..7ff337ab120 100644
--- a/official/nlp/modeling/layers/attention.py
+++ b/official/nlp/modeling/layers/attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,17 +16,17 @@
# pylint: disable=g-classes-have-attributes
import math
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-EinsumDense = tf.keras.layers.experimental.EinsumDense
-MultiHeadAttention = tf.keras.layers.MultiHeadAttention
+EinsumDense = tf_keras.layers.EinsumDense
+MultiHeadAttention = tf_keras.layers.MultiHeadAttention
-@tf.keras.utils.register_keras_serializable(package="Text")
-class CachedAttention(tf.keras.layers.MultiHeadAttention):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class CachedAttention(tf_keras.layers.MultiHeadAttention):
"""Attention layer with cache used for autoregressive decoding.
- Arguments are the same as `tf.keras.layers.MultiHeadAttention` layer.
+ Arguments are the same as `tf_keras.layers.MultiHeadAttention` layer.
"""
def _update_cache(self, key, value, cache, decode_loop_step):
diff --git a/official/nlp/modeling/layers/attention_test.py b/official/nlp/modeling/layers/attention_test.py
index 1f3d73d164a..288fb947c44 100644
--- a/official/nlp/modeling/layers/attention_test.py
+++ b/official/nlp/modeling/layers/attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,8 @@
"""Tests for the attention layer."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import attention
@@ -32,8 +31,7 @@ def _create_cache(batch_size, init_decode_length, num_heads, head_size):
}
-@keras_parameterized.run_all_keras_modes
-class CachedAttentionTest(keras_parameterized.TestCase):
+class CachedAttentionTest(tf.test.TestCase):
def test_masked_attention(self):
"""Test with a mask tensor."""
diff --git a/official/nlp/modeling/layers/bigbird_attention.py b/official/nlp/modeling/layers/bigbird_attention.py
index 8f6f3d61471..9094564a2d1 100644
--- a/official/nlp/modeling/layers/bigbird_attention.py
+++ b/official/nlp/modeling/layers/bigbird_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Keras-based bigbird attention layer."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
MAX_SEQ_LEN = 4096
@@ -368,7 +368,7 @@ def bigbird_block_sparse_attention(
return context_layer
-class BigBirdMasks(tf.keras.layers.Layer):
+class BigBirdMasks(tf_keras.layers.Layer):
"""Creates bigbird attention masks."""
def __init__(self, block_size, **kwargs):
@@ -390,8 +390,8 @@ def call(self, inputs, mask):
return [band_mask, encoder_from_mask, encoder_to_mask, blocked_encoder_mask]
-@tf.keras.utils.register_keras_serializable(package="Text")
-class BigBirdAttention(tf.keras.layers.MultiHeadAttention):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class BigBirdAttention(tf_keras.layers.MultiHeadAttention):
"""BigBird, a sparse attention mechanism.
This layer follows the paper "Big Bird: Transformers for Longer Sequences"
@@ -432,7 +432,7 @@ def __init__(self,
self.rand_attn = tf.constant(rand_attn, dtype=tf.int32)
def _compute_attention(self, query, key, value, attention_mask=None):
- (band_mask, encoder_from_mask, encoder_to_mask,
+ (band_mask, encoder_from_mask, encoder_to_mask, # pyrefly: ignore[not-iterable]
blocked_encoder_mask) = attention_mask
query_shape = tf.shape(query)
from_seq_length = query_shape[1]
@@ -458,7 +458,7 @@ def _compute_attention(self, query, key, value, attention_mask=None):
to_block_size=self._to_block_size,
rand_attn=rand_attn)
- def call(self, query, value, key=None, attention_mask=None, **kwargs):
+ def call(self, query, value, key=None, attention_mask=None, **kwargs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
if not self._built_from_signature:
self._build_from_signature(query=query, value=value, key=key)
if key is None:
diff --git a/official/nlp/modeling/layers/bigbird_attention_test.py b/official/nlp/modeling/layers/bigbird_attention_test.py
index 3764ce49db6..339b024c798 100644
--- a/official/nlp/modeling/layers/bigbird_attention_test.py
+++ b/official/nlp/modeling/layers/bigbird_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for official.nlp.projects.bigbird.attention."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import bigbird_attention as attention
diff --git a/official/nlp/modeling/layers/block_diag_feedforward.py b/official/nlp/modeling/layers/block_diag_feedforward.py
index 21fe144a15f..a6188fbd21d 100644
--- a/official/nlp/modeling/layers/block_diag_feedforward.py
+++ b/official/nlp/modeling/layers/block_diag_feedforward.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,10 +16,12 @@
# pylint: disable=g-classes-have-attributes
from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
-class BlockDiagFeedforward(tf.keras.layers.Layer):
+
+class BlockDiagFeedforward(tf_keras.layers.Layer):
"""Block diagonal feedforward layer.
This layer replaces the weight matrix of the output_dense layer with a block
@@ -51,13 +53,13 @@ def __init__(
apply_mixing: bool = True,
kernel_initializer: str = "glorot_uniform",
bias_initializer: str = "zeros",
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- activity_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- kernel_constraint: Optional[tf.keras.constraints.Constraint] = None,
- bias_constraint: Optional[tf.keras.constraints.Constraint] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ activity_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ kernel_constraint: Optional[tf_keras.constraints.Constraint] = None,
+ bias_constraint: Optional[tf_keras.constraints.Constraint] = None,
**kwargs): # pylint: disable=g-doc-args
- super(BlockDiagFeedforward, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._intermediate_size = intermediate_size
self._intermediate_activation = intermediate_activation
self._dropout = dropout
@@ -68,60 +70,64 @@ def __init__(
raise ValueError("Intermediate_size (%d) isn't a multiple of num_blocks "
"(%d)." % (intermediate_size, num_blocks))
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
def build(self, input_shape):
hidden_size = input_shape.as_list()[-1]
common_kwargs = dict(
- kernel_initializer=self._kernel_initializer,
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
kernel_constraint=self._kernel_constraint,
bias_constraint=self._bias_constraint)
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
"abc,cde->abde",
output_shape=(None, self._num_blocks,
self._intermediate_size // self._num_blocks),
bias_axes="de",
name="intermediate",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
policy = tf.float32
- self._intermediate_activation_layer = tf.keras.layers.Activation(
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
self._intermediate_activation, dtype=policy)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._output_dense = tf_keras.layers.EinsumDense(
"abde,deo->abdo",
- output_shape=(None, self._num_blocks,
- hidden_size // self._num_blocks),
+ output_shape=(None, self._num_blocks, hidden_size // self._num_blocks),
bias_axes="do",
name="output",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
if self._apply_mixing:
- self._output_mixing = tf.keras.layers.experimental.EinsumDense(
+ self._output_mixing = tf_keras.layers.EinsumDense(
"abdo,de->abeo",
output_shape=(None, self._num_blocks,
hidden_size // self._num_blocks),
name="output_mixing",
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- self._output_reshape = tf.keras.layers.Reshape((-1, hidden_size))
+ self._output_reshape = tf_keras.layers.Reshape((-1, hidden_size))
- self._output_dropout = tf.keras.layers.Dropout(rate=self._dropout)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._dropout)
def get_config(self):
config = {
@@ -136,21 +142,21 @@ def get_config(self):
"apply_mixing":
self._apply_mixing,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint)
+ tf_keras.constraints.serialize(self._bias_constraint)
}
- base_config = super(BlockDiagFeedforward, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def call(self, inputs):
diff --git a/official/nlp/modeling/layers/block_diag_feedforward_test.py b/official/nlp/modeling/layers/block_diag_feedforward_test.py
index e9b5b4e5e48..431b6b91091 100644
--- a/official/nlp/modeling/layers/block_diag_feedforward_test.py
+++ b/official/nlp/modeling/layers/block_diag_feedforward_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,20 +16,18 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import block_diag_feedforward
# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class BlockDiagFeedforwardTest(keras_parameterized.TestCase):
+class BlockDiagFeedforwardTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(BlockDiagFeedforwardTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
@parameterized.parameters(
(1, True, "float32"),
@@ -42,7 +40,7 @@ def tearDown(self):
(2, False, "mixed_float16"),
)
def test_layer_creation(self, num_blocks, apply_mixing, dtype):
- tf.keras.mixed_precision.set_global_policy(dtype)
+ tf_keras.mixed_precision.set_global_policy(dtype)
kwargs = dict(
intermediate_size=128,
intermediate_activation="relu",
@@ -56,7 +54,7 @@ def test_layer_creation(self, num_blocks, apply_mixing, dtype):
sequence_length = 64
width = 128
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -72,7 +70,7 @@ def test_layer_creation(self, num_blocks, apply_mixing, dtype):
(2, False, "mixed_float16"),
)
def test_layer_invocation(self, num_blocks, apply_mixing, dtype):
- tf.keras.mixed_precision.set_global_policy(dtype)
+ tf_keras.mixed_precision.set_global_policy(dtype)
kwargs = dict(
intermediate_size=16,
intermediate_activation="relu",
@@ -86,11 +84,11 @@ def test_layer_invocation(self, num_blocks, apply_mixing, dtype):
sequence_length = 16
width = 32
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(data_tensor, output_tensor)
+ model = tf_keras.Model(data_tensor, output_tensor)
# Invoke the model on test data.
batch_size = 6
diff --git a/official/nlp/modeling/layers/block_sparse_attention.py b/official/nlp/modeling/layers/block_sparse_attention.py
new file mode 100644
index 00000000000..20e65b0a392
--- /dev/null
+++ b/official/nlp/modeling/layers/block_sparse_attention.py
@@ -0,0 +1,359 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Block sparse attention converts query/key/value into blocks and performs diagonal block sparse attention."""
+import collections
+import logging
+
+import tensorflow as tf, tf_keras
+
+
+def _large_compatible_negative(tensor_type):
+ """Large negative number as Tensor.
+
+ This function is necessary because the standard value for epsilon
+ in this module (-1e9) cannot be represented using tf.float16
+
+ Args:
+ tensor_type: a dtype to determine the type.
+
+ Returns:
+ a large negative number.
+ """
+ # In case of dtype=float16 (e.g., for mixed-precision), the largest
+ # negative number (dtypes.float16.min) is divided by 2, in order to
+ # avoid overflows when summing negative inputs.
+ if tensor_type == tf.float16:
+ return tf.float16.min / 2.0
+ return -1e9
+
+
+class MultiHeadAttention(tf_keras.layers.MultiHeadAttention):
+ """Multi-head block sparse attention layer."""
+
+ def __init__(
+ self,
+ src_block_size=None,
+ tgt_block_size=None,
+ use_sigmoid_attn=False,
+ sigmoid_attn_bias=None,
+ num_kv_heads=None,
+ **kwargs
+ ):
+ """Initializes the block sparse attention layer.
+
+ Args:
+ src_block_size: The block size of the query. An integer that divides the
+ sequence length into blocks.
+ tgt_block_size: The block size of the key/value. An integer that divides
+ the sequence length into blocks. The number of blocks in the source and
+ target must be the same.
+ use_sigmoid_attn: If enabled, uses sigmoid instead of softmax to compute
+ attn probs. https://arxiv.org/pdf/2409.04431
+ sigmoid_attn_bias: Bias for sigmoid attn. Suggested value -ln(seq_len).
+ num_kv_heads: Number of key/value heads in the multi-head self attention.
+ Refer to multi_query_attention.py for more details.
+ **kwargs: Args passed to the base class.
+ """
+ super().__init__(**kwargs)
+ if src_block_size is None or src_block_size <= 0:
+ raise ValueError("src_block_size must be specified.")
+ self._src_block_size = src_block_size
+ self._tgt_block_size = tgt_block_size or self._src_block_size
+ self._num_kv_heads = num_kv_heads
+ if num_kv_heads is not None and num_kv_heads != 1:
+ raise ValueError(
+ "num_kv_heads must be 1. Grouped-query attention is not supported."
+ )
+ self._use_sigmoid_attn = use_sigmoid_attn
+ self._sigmoid_attn_bias = sigmoid_attn_bias
+ if self._use_sigmoid_attn:
+ if self._sigmoid_attn_bias is None:
+ raise ValueError(
+ "sigmoid_attn_bias must be specified for sigmoid attn."
+ )
+
+ def get_config(self):
+ config = super().get_config()
+ config.update({
+ "src_block_size": self._src_block_size,
+ "tgt_block_size": self._tgt_block_size,
+ "use_sigmoid_attn": self._use_sigmoid_attn,
+ "sigmoid_attn_bias": self._sigmoid_attn_bias,
+ "num_kv_heads": self._num_kv_heads,
+ })
+ return config
+
+ def _build_from_signature(self, query, value, key=None):
+ # pytype: disable=attribute-error
+ super()._build_from_signature(query, value, key)
+ # pytype: enable=attribute-error
+ # If block sizes are same as sequence lengths, we defer to default attn.
+ if (
+ self._query_shape[-2] == self._src_block_size
+ and self._key_shape[-2] == self._tgt_block_size
+ ):
+ return
+ # The following capital letters are used to denote the tensor dimension
+ # parameters:
+ # B = batch size
+ # S = length of the key/value (target)
+ # D = model dimension.
+ # T = length of the query (source)
+ # t = block size of the source.
+ # s = block size of the target.
+ # L = number of blocks in the source/target.
+ # N = number of attention heads
+ # H = dimensions of each attention head.
+ with tf.init_scope():
+ proj_einsum_eqn = "BTD,DNH->BNTH"
+ bias_axes = "NH"
+ qk_output_shape = [
+ self._num_heads,
+ None,
+ self._key_dim,
+ ]
+ v_output_shape = [
+ self._num_heads,
+ None,
+ self._value_dim,
+ ]
+ self._query_dense = tf_keras.layers.EinsumDense(
+ proj_einsum_eqn,
+ output_shape=qk_output_shape,
+ bias_axes=bias_axes if self._use_bias else None,
+ name="query",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+ if self._num_kv_heads == 1:
+ self._key_dense = tf_keras.layers.EinsumDense(
+ "BTD,DH->BTH",
+ output_shape=[None, self._key_dim],
+ bias_axes="H" if self._use_bias else None,
+ name="key",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+ self._value_dense = tf_keras.layers.EinsumDense(
+ "BTD,DH->BTH",
+ output_shape=[None, self._value_dim],
+ bias_axes="H" if self._use_bias else None,
+ name="value",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+ else:
+ self._key_dense = tf_keras.layers.EinsumDense(
+ proj_einsum_eqn,
+ output_shape=qk_output_shape,
+ bias_axes=bias_axes if self._use_bias else None,
+ name="key",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+ self._value_dense = tf_keras.layers.EinsumDense(
+ proj_einsum_eqn,
+ output_shape=v_output_shape,
+ bias_axes=bias_axes if self._use_bias else None,
+ name="value",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+ if self._key_shape[-2] == self._tgt_block_size:
+ if self._num_kv_heads == 1:
+ self._dot_product_equation = "BsH,BNLtH->BNLts"
+ self._combine_equation = "BNLts,BsH->BNLtH"
+ else:
+ self._dot_product_equation = "BNsH,BNLtH->BNLts"
+ self._combine_equation = "BNLts,BNsH->BNLtH"
+ else:
+ if self._num_kv_heads == 1:
+ self._dot_product_equation = "BLsH,BNLtH->BNLts"
+ self._combine_equation = "BNLts,BLsH->BNLtH"
+ else:
+ self._dot_product_equation = "BNLsH,BNLtH->BNLts"
+ self._combine_equation = "BNLts,BNLsH->BNLtH"
+ if self._output_shape:
+ if not isinstance(self._output_shape, collections.abc.Sized):
+ output_shape = [self._output_shape]
+ else:
+ output_shape = self._output_shape
+ else:
+ output_shape = [self._query_shape[-1]]
+ output_shape = [None] + output_shape # pyrefly: ignore[unsupported-operation]
+ self._output_dense = tf_keras.layers.EinsumDense(
+ "BNTH,DNH->BTD",
+ output_shape=output_shape,
+ bias_axes="D" if self._use_bias else None,
+ name="attention_output",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+
+ def _block_diagonal_mask(self, attention_mask, dtype=None):
+ """Converts the attention mask to block diagonal."""
+ # Uses the same key mask for the entire query sequence since softmax
+ # is applied only on the key axis.
+ tgt_num_blocks = self._key_shape[-2] // self._tgt_block_size
+ if tgt_num_blocks == 1:
+ src_num_blocks = self._query_shape[-2] // self._src_block_size
+ result = tf.reshape(
+ attention_mask,
+ [-1, src_num_blocks, self._src_block_size, self._tgt_block_size],
+ )
+ else:
+ attention_mask = tf.cast(attention_mask[:, 0, :], dtype=dtype)
+ attention_mask = tf.reshape(
+ attention_mask,
+ [
+ -1,
+ tgt_num_blocks,
+ self._tgt_block_size,
+ ],
+ )
+ result = tf.einsum("BLQ,BLK->BLQK", attention_mask, attention_mask)
+ return result
+
+ def _masked_softmax(self, attention_scores, attention_mask=None):
+ # Normalize the attention scores to probabilities.
+ # `attention_scores` = [B, N, L, T, S]
+ if attention_mask is not None:
+ # `attention_mask` = [B, 1, L, T, S]
+ attention_mask = tf.expand_dims(attention_mask, axis=1)
+ if self._use_sigmoid_attn:
+ if attention_mask is not None:
+ adder = (1.0 - tf.cast(attention_mask, attention_scores.dtype)) * (
+ _large_compatible_negative(attention_scores.dtype)
+ )
+ attention_scores += adder
+ attention_scores += self._sigmoid_attn_bias
+ return tf_keras.activations.sigmoid(attention_scores)
+ else:
+ return self._softmax(attention_scores, attention_mask)
+
+ def _compute_attention(
+ self, query, key, value, attention_mask=None, training=None
+ ):
+ # If block sizes are same as sequence lengths, we defer to default attn.
+ if (
+ self._query_shape[-2] == self._src_block_size
+ and self._key_shape[-2] == self._tgt_block_size
+ ):
+ logging.info(
+ "Computing default attention as block sizes are equal to sequence"
+ " lengths."
+ )
+ # pytype: disable=attribute-error
+ return super()._compute_attention(
+ query,
+ key,
+ value,
+ attention_mask=attention_mask,
+ training=training,
+ )
+ # pytype: enable=attribute-error
+ # src_num_blocks and tgt_num_blocks are the number of blocks in the source
+ # and target. Care should be taken to ensure that the number of blocks in
+ # the source and target are the same.
+ if self._query_shape[-2] % self._src_block_size != 0:
+ raise ValueError(
+ "query_shape[-2] must be divisible by src_block_size."
+ )
+ if self._key_shape[-2] % self._tgt_block_size != 0:
+ raise ValueError(
+ "key_shape[-2] must be divisible by tgt_block_size."
+ )
+ src_num_blocks = self._query_shape[-2] // self._src_block_size
+ tgt_num_blocks = self._key_shape[-2] // self._tgt_block_size
+
+ if src_num_blocks != tgt_num_blocks and tgt_num_blocks != 1:
+ raise ValueError(
+ "src_num_blocks must be equal to tgt_num_blocks."
+ )
+ # Convert the query/key/value into blocks to perform block diagonal
+ # attention.
+ query_blocks = tf.reshape(query, [
+ -1,
+ self._num_heads,
+ src_num_blocks,
+ self._src_block_size,
+ self._key_dim,
+ ])
+ if tgt_num_blocks != 1 and self._num_kv_heads != 1:
+ key_blocks = tf.reshape(key, [
+ -1,
+ self._num_heads,
+ tgt_num_blocks,
+ self._tgt_block_size,
+ self._key_dim,
+ ])
+ value_blocks = tf.reshape(value, [
+ -1,
+ self._num_heads,
+ tgt_num_blocks,
+ self._tgt_block_size,
+ self._value_dim,
+ ])
+ elif tgt_num_blocks != 1 and self._num_kv_heads == 1:
+ key_blocks = tf.reshape(key, [
+ -1,
+ tgt_num_blocks,
+ self._tgt_block_size,
+ self._key_dim,
+ ])
+ value_blocks = tf.reshape(value, [
+ -1,
+ tgt_num_blocks,
+ self._tgt_block_size,
+ self._value_dim,
+ ])
+ else:
+ key_blocks = key
+ value_blocks = value
+ if attention_mask is not None:
+ attention_mask = self._block_diagonal_mask(attention_mask, key.dtype)
+ # pytype: disable=attribute-error
+ attention_output, attention_scores = super()._compute_attention(
+ query_blocks,
+ key_blocks,
+ value_blocks,
+ attention_mask=attention_mask,
+ training=training,
+ )
+ # pytype: enable=attribute-error
+ # Reshape the attention output to the original shape.
+ attention_output = tf.reshape(attention_output, [
+ -1,
+ self._num_heads,
+ self._query_shape[1],
+ self._value_dim,
+ ])
+ return attention_output, attention_scores
+
+ def call(
+ self,
+ query,
+ value,
+ key=None,
+ attention_mask=None,
+ return_attention_scores=False,
+ training=None,
+ use_causal_mask=False,
+ ):
+ if use_causal_mask:
+ raise ValueError("use_causal_mask is not supported.")
+ return super().call(
+ query,
+ value,
+ key=key,
+ attention_mask=attention_mask,
+ return_attention_scores=return_attention_scores,
+ training=training,
+ use_causal_mask=use_causal_mask,
+ )
diff --git a/official/nlp/modeling/layers/block_sparse_attention_test.py b/official/nlp/modeling/layers/block_sparse_attention_test.py
new file mode 100644
index 00000000000..3b21c2eef7e
--- /dev/null
+++ b/official/nlp/modeling/layers/block_sparse_attention_test.py
@@ -0,0 +1,433 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for block sparse attention layer."""
+
+import math
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling.layers import block_sparse_attention
+
+
+class BlockSparseAttentionTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ ("key_value_same_proj", None, None, [40, 80]),
+ ("key_value_same_proj_mqa", None, None, [40, 80], False, 1),
+ ("key_value_same_proj_multi_query_blocks", None, None, [40, 80], True),
+ (
+ "key_value_same_proj_multi_query_blocks_mqa",
+ None,
+ None,
+ [40, 80],
+ True,
+ 1,
+ ),
+ ("key_value_different_proj", 32, 60, [40, 60]),
+ ("key_value_different_proj_mqa", 32, 60, [40, 60], False, 1),
+ ("key_value_different_proj_multi_query_blocks", 32, 60, [40, 60], True),
+ (
+ "key_value_different_proj_multi_query_blocks_mqa",
+ 32,
+ 60,
+ [40, 60],
+ True,
+ 1,
+ ),
+ )
+ def test_non_masked_attention(
+ self,
+ value_dim,
+ output_shape,
+ output_dims,
+ multi_query_blocks=False,
+ num_kv_heads=None,
+ ):
+ """Test that the attention layer can be created without a mask tensor."""
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=12,
+ key_dim=64,
+ value_dim=value_dim,
+ output_shape=output_shape,
+ src_block_size=10,
+ tgt_block_size=20 if multi_query_blocks else 5,
+ num_kv_heads=num_kv_heads,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(20, 80))
+ output = test_layer(query=query, value=value)
+ self.assertEqual(output.shape.as_list(), [None] + output_dims)
+
+ def test_non_masked_self_attention(self):
+ """Test with one input (self-attenntion) and no mask tensor."""
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=12, key_dim=64, src_block_size=10, tgt_block_size=10
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ output = test_layer(query, query)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+
+ @parameterized.named_parameters(
+ ("with_bias", True),
+ ("with_bias_mqa", True, False, False, 1),
+ ("with_bias_multi_query_blocks", True, False, True),
+ ("with_bias_multi_query_blocks_mqa", True, False, True, 1),
+ ("no_bias", False),
+ ("no_bias_mqa", False, False, False, 1),
+ ("no_bias_multi_query_blocks", False, False, True),
+ ("no_bias_multi_query_blocks_mqa", False, False, True, 1),
+ ("with_sigmoid_attn", True, True),
+ ("with_sigmoid_attn_mqa", True, True, False, 1),
+ ("with_sigmoid_attn_multi_query_blocks", True, True, True),
+ ("with_sigmoid_attn_multi_query_blocks_mqa", True, True, True, 1),
+ )
+ def test_masked_attention(
+ self,
+ use_bias,
+ use_sigmoid_attn=False,
+ multi_query_blocks=False,
+ num_kv_heads=None,
+ ):
+ """Test with a mask tensor."""
+ if use_sigmoid_attn:
+ sigmoid_attn_bias = -math.log(2)
+ else:
+ sigmoid_attn_bias = None
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=4,
+ key_dim=2,
+ use_bias=use_bias,
+ src_block_size=2,
+ tgt_block_size=2 if multi_query_blocks else 1,
+ use_sigmoid_attn=use_sigmoid_attn,
+ sigmoid_attn_bias=sigmoid_attn_bias,
+ num_kv_heads=num_kv_heads,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ batch_size = 3
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
+ output = test_layer(query=query, value=value, attention_mask=mask_tensor)
+
+ # Create a model containing the test layer.
+ model = tf_keras.Model([query, value, mask_tensor], output)
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = 10 * np.random.random_sample((batch_size, 4, 8))
+ to_data = 10 * np.random.random_sample((batch_size, 2, 8))
+
+ # Invoke the data with a random set of mask data. This should mask at
+ # least one element.
+ mask_data = np.random.randint(2, size=(batch_size, 4, 2))
+ masked_output_data = model.predict([from_data, to_data, mask_data])
+
+ # Invoke the same data, but with a null mask (where no elements are
+ # masked).
+ null_mask_data = np.ones((batch_size, 4, 2))
+ unmasked_output_data = model.predict([from_data, to_data, null_mask_data])
+
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ # Tests the layer with three inputs: Q, K, V.
+ key = tf_keras.Input(shape=(2, 8))
+ output = test_layer(
+ query, value=value, key=key, attention_mask=mask_tensor
+ )
+ model = tf_keras.Model([query, value, key, mask_tensor], output)
+
+ masked_output_data = model.predict(
+ [from_data, to_data, to_data, mask_data]
+ )
+ unmasked_output_data = model.predict(
+ [from_data, to_data, to_data, null_mask_data]
+ )
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ if use_bias:
+ self.assertLen(test_layer._query_dense.trainable_variables, 2)
+ self.assertLen(test_layer._output_dense.trainable_variables, 2)
+ else:
+ self.assertLen(test_layer._query_dense.trainable_variables, 1)
+ self.assertLen(test_layer._output_dense.trainable_variables, 1)
+
+ @parameterized.named_parameters(
+ ("default_with_softmax", False),
+ ("default_with_sigmoid", True),
+ )
+ def test_default_masked_attention(
+ self,
+ use_sigmoid_attn=False,
+ ):
+ """Test with a mask tensor."""
+ seq_len = 8
+ if use_sigmoid_attn:
+ sigmoid_attn_bias = -math.log(seq_len)
+ else:
+ sigmoid_attn_bias = None
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=4,
+ key_dim=2,
+ use_bias=True,
+ src_block_size=seq_len,
+ tgt_block_size=seq_len,
+ use_sigmoid_attn=use_sigmoid_attn,
+ sigmoid_attn_bias=sigmoid_attn_bias,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ batch_size = 3
+ query = tf_keras.Input(shape=(seq_len, 8))
+ value = tf_keras.Input(shape=(seq_len, 8))
+ mask_tensor = tf_keras.Input(shape=(seq_len, seq_len))
+ output = test_layer(query=query, value=value, attention_mask=mask_tensor)
+
+ # Create a model containing the test layer.
+ model = tf_keras.Model([query, value, mask_tensor], output)
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = 10 * np.random.random_sample((batch_size, seq_len, 8))
+ to_data = 10 * np.random.random_sample((batch_size, seq_len, 8))
+
+ # Invoke the data with a random set of mask data. This should mask at
+ # least one element.
+ mask_data = np.random.randint(2, size=(batch_size, seq_len, seq_len))
+ masked_output_data = model.predict([from_data, to_data, mask_data])
+
+ # Invoke the same data, but with a null mask (where no elements are
+ # masked).
+ null_mask_data = np.ones((batch_size, seq_len, seq_len))
+ unmasked_output_data = model.predict([from_data, to_data, null_mask_data])
+
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ # Tests the layer with three inputs: Q, K, V.
+ key = tf_keras.Input(shape=(seq_len, 8))
+ output = test_layer(
+ query, value=value, key=key, attention_mask=mask_tensor
+ )
+ model = tf_keras.Model([query, value, key, mask_tensor], output)
+
+ masked_output_data = model.predict(
+ [from_data, to_data, to_data, mask_data]
+ )
+ unmasked_output_data = model.predict(
+ [from_data, to_data, to_data, null_mask_data]
+ )
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ self.assertLen(test_layer._query_dense.trainable_variables, 2)
+ self.assertLen(test_layer._output_dense.trainable_variables, 2)
+
+ def test_masked_attention_with_scores(self):
+ """Test with a mask tensor."""
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=4, key_dim=2, src_block_size=2, tgt_block_size=1,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ batch_size = 3
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
+ output = test_layer(query=query, value=value, attention_mask=mask_tensor)
+
+ # Create a model containing the test layer.
+ model = tf_keras.Model([query, value, mask_tensor], output)
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = 10 * np.random.random_sample((batch_size, 4, 8))
+ to_data = 10 * np.random.random_sample((batch_size, 2, 8))
+
+ # Invoke the data with a random set of mask data. This should mask at
+ # least one element.
+ mask_data = np.random.randint(2, size=(batch_size, 4, 2))
+ masked_output_data = model.predict([from_data, to_data, mask_data])
+
+ # Invoke the same data, but with a null mask (where no elements are
+ # masked).
+ null_mask_data = np.ones((batch_size, 4, 2))
+ unmasked_output_data = model.predict([from_data, to_data, null_mask_data])
+
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ # Create a model containing attention scores.
+ output, scores = test_layer(
+ query=query,
+ value=value,
+ attention_mask=mask_tensor,
+ return_attention_scores=True,
+ )
+ model = tf_keras.Model([query, value, mask_tensor], [output, scores])
+ masked_output_data_score, masked_score = model.predict(
+ [from_data, to_data, mask_data]
+ )
+ unmasked_output_data_score, unmasked_score = model.predict(
+ [from_data, to_data, null_mask_data]
+ )
+ self.assertNotAllClose(masked_output_data_score, unmasked_output_data_score)
+ self.assertAllClose(masked_output_data, masked_output_data_score)
+ self.assertAllClose(unmasked_output_data, unmasked_output_data_score)
+ self.assertNotAllClose(masked_score, unmasked_score)
+
+ def test_initializer(self):
+ """Test with a specified initializer."""
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=12,
+ key_dim=64,
+ src_block_size=10,
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ output = test_layer(query, query)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+
+ # Make sure the sub layers have different kernel init value, and not
+ # reusing the initializers.
+ self.assertNotAllClose(
+ tf_keras.backend.eval(test_layer._query_dense.kernel),
+ tf_keras.backend.eval(test_layer._key_dense.kernel),
+ )
+ self.assertNotAllClose(
+ tf_keras.backend.eval(test_layer._query_dense.kernel),
+ tf_keras.backend.eval(test_layer._value_dense.kernel),
+ )
+ self.assertNotAllClose(
+ tf_keras.backend.eval(test_layer._query_dense.kernel),
+ tf_keras.backend.eval(test_layer._output_dense.kernel),
+ )
+
+ @parameterized.named_parameters(
+ ("bfloat16", tf.bfloat16),
+ ("float16", tf.float16),
+ ("float32", tf.float32),
+ ("float64", tf.float64),
+ )
+ def test_sublayer_dtypes(self, dtype):
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=12, key_dim=64, src_block_size=10, dtype=dtype
+ )
+
+ query = tf_keras.Input(shape=(40, 80), dtype=dtype)
+ # Build the layer
+ test_layer(query=query, value=query)
+
+ self.assertEqual(test_layer._query_dense.dtype, dtype)
+ self.assertEqual(test_layer._key_dense.dtype, dtype)
+ self.assertEqual(test_layer._value_dense.dtype, dtype)
+ self.assertEqual(test_layer._output_dense.dtype, dtype)
+
+ def test_dropout(self):
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=2, key_dim=2, dropout=0.5, src_block_size=2, tgt_block_size=1,
+ )
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = tf_keras.backend.ones(shape=(32, 4, 8))
+ to_data = tf_keras.backend.ones(shape=(32, 2, 8))
+ train_out = test_layer(from_data, to_data, None, None, None, True)
+ test_out = test_layer(from_data, to_data, None, None, None, False)
+
+ # Output should be close when not in training mode,
+ # and should not be close when enabling dropout in training mode.
+ self.assertNotAllClose(
+ tf_keras.backend.eval(train_out), tf_keras.backend.eval(test_out)
+ )
+
+ def test_query_mask_progagation(self):
+ """Test automatic propagation of the query's mask."""
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=2,
+ key_dim=2,
+ src_block_size=2,
+ tgt_block_size=1,
+ )
+ self.assertTrue(test_layer.supports_masking)
+ query = tf.constant(
+ [[1, 2, 3, 0, 0, 0], [3, 3, 1, 1, 2, 0], [1, 1, 0, 0, 0, 0]]
+ )
+ masked_query = tf_keras.layers.Embedding(4, 8, mask_zero=True)(query)
+ value = tf.random.normal((3, 3, 8))
+ output = test_layer(query=masked_query, value=value)
+ self.assertTrue(hasattr(output, "_keras_mask"))
+ self.assertAllEqual(masked_query._keras_mask, output._keras_mask)
+
+ def test_value_mask(self):
+ """Test that the value mask is taken into account."""
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=2,
+ key_dim=2,
+ src_block_size=2,
+ tgt_block_size=1,
+ )
+ query = tf.constant(
+ [[1, 2, 3, 0, 0, 0], [3, 3, 1, 1, 2, 0], [1, 1, 0, 0, 0, 0]]
+ )
+ masked_query = tf_keras.layers.Embedding(4, 8, mask_zero=True)(query)
+ value = tf.constant([[5, 4, 0], [3, 0, 0], [2, 1, 1]])
+ masked_value = tf_keras.layers.Embedding(6, 8, mask_zero=True)(value)
+ output = test_layer(
+ query=masked_query,
+ value=masked_value,
+ )
+ mask = tf.constant(
+ [[[True, True, False]] * 3 + [[False, False, False]] * 2]
+ + [[[True, False, False]] * 5]
+ + [[[True, True, True]] + [[False, False, False]] * 4]
+ )
+ del masked_query._keras_mask
+ del masked_value._keras_mask
+ output_with_manual_mask = test_layer(
+ query=masked_query, value=masked_value, attention_mask=mask
+ )
+ self.assertAllClose(output, output_with_manual_mask)
+
+ def test_masks_are_cast_to_bool(self):
+ """Test that the implicit and explicit masks are cast to bool."""
+ test_layer = block_sparse_attention.MultiHeadAttention(
+ num_heads=2, key_dim=2, src_block_size=2, tgt_block_size=1,
+ )
+ query = np.array(
+ [[1, 2, 3, 0, 0, 0], [3, 3, 1, 1, 2, 0], [1, 1, 0, 0, 0, 0]]
+ )
+ masked_query = tf_keras.layers.Embedding(4, 8, mask_zero=True)(query)
+ masked_query._keras_mask = tf.cast(masked_query._keras_mask, tf.float32)
+ value = np.array([[5, 4, 0], [3, 0, 0], [2, 1, 1]])
+ masked_value = tf_keras.layers.Embedding(6, 8, mask_zero=True)(value)
+ masked_value._keras_mask = tf.cast(masked_value._keras_mask, tf.float32)
+ float_mask = tf.constant([[[1.0]]])
+ # if all works well, the following should not raise any exception:
+ _ = test_layer(
+ query=masked_query,
+ value=masked_value,
+ attention_mask=float_mask,
+ )
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/nlp/modeling/layers/cls_head.py b/official/nlp/modeling/layers/cls_head.py
index 82156e4d004..d0fd8e7f539 100644
--- a/official/nlp/modeling/layers/cls_head.py
+++ b/official/nlp/modeling/layers/cls_head.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""A Classification head layer which is common used with sequence encoders."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
@@ -22,7 +22,7 @@
from official.nlp.modeling.layers import spectral_normalization
-class ClassificationHead(tf.keras.layers.Layer):
+class ClassificationHead(tf_keras.layers.Layer):
"""Pooling head for sentence-level classification tasks."""
def __init__(self,
@@ -50,19 +50,21 @@ def __init__(self,
self.inner_dim = inner_dim
self.num_classes = num_classes
self.activation = tf_utils.get_activation(activation)
- self.initializer = tf.keras.initializers.get(initializer)
+ self.initializer = tf_keras.initializers.get(initializer)
self.cls_token_idx = cls_token_idx
if self.inner_dim:
- self.dense = tf.keras.layers.Dense(
+ self.dense = tf_keras.layers.Dense(
units=self.inner_dim,
activation=self.activation,
- kernel_initializer=self.initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name="pooler_dense")
- self.dropout = tf.keras.layers.Dropout(rate=self.dropout_rate)
+ self.dropout = tf_keras.layers.Dropout(rate=self.dropout_rate)
- self.out_proj = tf.keras.layers.Dense(
- units=num_classes, kernel_initializer=self.initializer, name="logits")
+ self.out_proj = tf_keras.layers.Dense(
+ units=num_classes,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
+ name="logits")
def call(self, features: tf.Tensor, only_project: bool = False):
"""Implements call().
@@ -95,8 +97,8 @@ def get_config(self):
"dropout_rate": self.dropout_rate,
"num_classes": self.num_classes,
"inner_dim": self.inner_dim,
- "activation": tf.keras.activations.serialize(self.activation),
- "initializer": tf.keras.initializers.serialize(self.initializer),
+ "activation": tf_keras.activations.serialize(self.activation),
+ "initializer": tf_keras.initializers.serialize(self.initializer),
}
config.update(super(ClassificationHead, self).get_config())
return config
@@ -110,7 +112,7 @@ def checkpoint_items(self):
return {self.dense.name: self.dense}
-class MultiClsHeads(tf.keras.layers.Layer):
+class MultiClsHeads(tf_keras.layers.Layer):
"""Pooling heads sharing the same pooling stem."""
def __init__(self,
@@ -139,21 +141,22 @@ def __init__(self,
self.inner_dim = inner_dim
self.cls_list = cls_list
self.activation = tf_utils.get_activation(activation)
- self.initializer = tf.keras.initializers.get(initializer)
+ self.initializer = tf_keras.initializers.get(initializer)
self.cls_token_idx = cls_token_idx
if self.inner_dim:
- self.dense = tf.keras.layers.Dense(
+ self.dense = tf_keras.layers.Dense(
units=inner_dim,
activation=self.activation,
- kernel_initializer=self.initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name="pooler_dense")
- self.dropout = tf.keras.layers.Dropout(rate=self.dropout_rate)
+ self.dropout = tf_keras.layers.Dropout(rate=self.dropout_rate)
self.out_projs = []
for name, num_classes in cls_list:
self.out_projs.append(
- tf.keras.layers.Dense(
- units=num_classes, kernel_initializer=self.initializer,
+ tf_keras.layers.Dense(
+ units=num_classes,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name=name))
def call(self, features: tf.Tensor, only_project: bool = False):
@@ -190,8 +193,8 @@ def get_config(self):
"cls_token_idx": self.cls_token_idx,
"cls_list": self.cls_list,
"inner_dim": self.inner_dim,
- "activation": tf.keras.activations.serialize(self.activation),
- "initializer": tf.keras.initializers.serialize(self.initializer),
+ "activation": tf_keras.activations.serialize(self.activation),
+ "initializer": tf_keras.initializers.serialize(self.initializer),
}
config.update(super().get_config())
return config
@@ -277,11 +280,11 @@ def __init__(self,
if use_gp_layer:
self.out_proj = gaussian_process.RandomFeatureGaussianProcess(
self.num_classes,
- kernel_initializer=self.initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name="logits",
**self.gp_layer_kwargs)
- def call(self, features, training=False, return_covmat=False):
+ def call(self, features, training=False, return_covmat=False): # pyrefly: ignore[bad-override]
"""Returns model output.
Dring training, the model returns raw logits. During evaluation, the model
@@ -361,3 +364,97 @@ def extract_spec_norm_kwargs(kwargs):
return dict(
iteration=kwargs.pop("iteration", 1),
norm_multiplier=kwargs.pop("norm_multiplier", .99))
+
+
+class PerQueryDenseHead(tf_keras.layers.Layer):
+ """Pooling head used for EncT5 style models.
+
+ This module projects each query to use a different projection.
+
+ For a input shape= [bs, num_queries, hidden_size], it projects each query to
+ (features). Ending up with shape= [bs, num_queries, features].
+
+ For example, for classification with a few classes, one may use num_queries
+ as 1 and features as number of classes. For multilabel classification, one
+ may use num_queries as number of classes and features as 2. So each query
+ represents a binary classification of one label.
+ """
+
+ def __init__(self,
+ num_queries: int,
+ features: int,
+ use_bias: bool = False,
+ kernel_initializer: str = "glorot_uniform",
+ **kwargs):
+ """Initializes the `PerQueryDenseHead`.
+
+ Args:
+ num_queries: number of queries (the learnable embeddings in the input
+ sequences) from the decoder.
+ features: int with numbers of output features. Each query with be
+ projected to this number with a different projection.
+ use_bias: whether to add a bias to the output.
+ kernel_initializer: Initializer for dense layer kernels.
+ **kwargs: Keyword arguments.
+ """
+ super().__init__(**kwargs)
+ self.num_queries = num_queries
+ self.features = features
+
+ self.use_bias = use_bias
+ self.kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+
+ def build(self, input_shape):
+ input_shape = tf.TensorShape(input_shape)
+ # Hidden size.
+ last_dim = tf.compat.dimension_value(input_shape[-1])
+
+ self.hidden_size = last_dim
+ self.kernel = self.add_weight(
+ "kernel",
+ shape=[self.num_queries, last_dim, self.features],
+ initializer=self.kernel_initializer,
+ dtype=self.dtype,
+ trainable=True)
+ if self.use_bias:
+ self.bias = self.add_weight(
+ "bias",
+ shape=[
+ self.num_queries,
+ self.features,
+ ],
+ dtype=self.dtype,
+ trainable=True)
+ else:
+ self.bias = None
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ """Implements call().
+
+ Args:
+ inputs: a rank-3 Tensor of shape= [bs, num_queries, hidden_size].
+
+ Returns:
+ A Tensor, shape= [batch size, num_queries, features].
+ """
+
+ outputs = tf.einsum("bqh,qhf->bqf", inputs, self.kernel)
+ if self.use_bias:
+ outputs += self.bias
+ return outputs
+
+ def get_config(self):
+ config = {
+ "num_queries":
+ self.num_queries,
+ "features":
+ self.features,
+ "kernel_initializer":
+ tf_keras.activations.serialize(self.kernel_initializer),
+ }
+ config.update(super(PerQueryDenseHead, self).get_config())
+ return config
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
diff --git a/official/nlp/modeling/layers/cls_head_test.py b/official/nlp/modeling/layers/cls_head_test.py
index 8bcfb0bba37..c0003d251de 100644
--- a/official/nlp/modeling/layers/cls_head_test.py
+++ b/official/nlp/modeling/layers/cls_head_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for cls_head."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import cls_head
@@ -199,5 +199,29 @@ def test_sngp_kwargs_serialization(self):
self.assertEqual(layer_config["norm_multiplier"], 1.)
self.assertEqual(layer_config["num_inducing"], 512)
+
+class PerQueryDenseHeadTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(("single_query", 1, 3, False),
+ ("multi_queries", 10, 2, False),
+ ("with_bias", 10, 2, True))
+ def test_layer_invocation(self, num_queries, features, use_bias):
+ batch_size = 5
+ hidden_size = 10
+ layer = cls_head.PerQueryDenseHead(
+ num_queries=num_queries, features=features, use_bias=use_bias)
+ inputs = tf.zeros(
+ shape=(batch_size, num_queries, hidden_size), dtype=tf.float32)
+ outputs = layer(inputs)
+ self.assertEqual(outputs.shape, [batch_size, num_queries, features])
+
+ def test_layer_serialization(self):
+ layer = cls_head.PerQueryDenseHead(
+ num_queries=10, features=2, use_bias=True)
+ new_layer = cls_head.PerQueryDenseHead.from_config(layer.get_config())
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(layer.get_config(), new_layer.get_config())
+
if __name__ == "__main__":
tf.test.main()
diff --git a/official/nlp/modeling/layers/factorized_embedding.py b/official/nlp/modeling/layers/factorized_embedding.py
new file mode 100644
index 00000000000..7f4003ba6ca
--- /dev/null
+++ b/official/nlp/modeling/layers/factorized_embedding.py
@@ -0,0 +1,76 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""A factorized embedding layer."""
+# pylint: disable=g-classes-have-attributes
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.nlp.modeling.layers import on_device_embedding
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class FactorizedEmbedding(on_device_embedding.OnDeviceEmbedding):
+ """A factorized embeddings layer for supporting larger embeddings.
+
+ Arguments:
+ vocab_size: Number of elements in the vocabulary.
+ embedding_width: Width of word embeddings.
+ output_dim: The output dimension of this layer.
+ initializer: The initializer to use for the embedding weights. Defaults to
+ "glorot_uniform".
+ use_one_hot: Whether to use tf.one_hot over tf.gather for the embedding
+ lookup. Defaults to False (that is, using tf.gather). Setting this option
+ to True may improve performance, especially on small vocabulary sizes, but
+ will generally require more memory.
+ scale_factor: Whether to scale the output embeddings. Defaults to None (that
+ is, not to scale). Setting this option to a float will let values in
+ output embeddings multiplied by scale_factor.
+ """
+
+ def __init__(self,
+ vocab_size: int,
+ embedding_width: int,
+ output_dim: int,
+ initializer='glorot_uniform',
+ use_one_hot=False,
+ scale_factor=None,
+ **kwargs):
+ super().__init__(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ initializer=initializer,
+ use_one_hot=use_one_hot,
+ scale_factor=scale_factor,
+ **kwargs)
+ self._output_dim = output_dim
+
+ def get_config(self):
+ config = {'output_dim': self._output_dim}
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(config.items()))
+
+ def build(self, input_shape):
+ self._embedding_projection = tf_keras.layers.EinsumDense(
+ '...x,xy->...y',
+ output_shape=self._output_dim,
+ bias_axes=None,
+ kernel_initializer=tf_utils.clone_initializer(self._initializer),
+ name='embedding_projection')
+ super().build(input_shape)
+
+ def call(self, inputs):
+ output = super().call(inputs)
+ return self._embedding_projection(output)
diff --git a/official/nlp/modeling/layers/factorized_embedding_test.py b/official/nlp/modeling/layers/factorized_embedding_test.py
new file mode 100644
index 00000000000..bf6cdfa3432
--- /dev/null
+++ b/official/nlp/modeling/layers/factorized_embedding_test.py
@@ -0,0 +1,70 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for FactorizedEmbedding layer."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling.layers import factorized_embedding
+
+
+class FactorizedEmbeddingTest(tf.test.TestCase):
+
+ def test_layer_creation(self):
+ vocab_size = 31
+ embedding_width = 27
+ output_dim = 45
+ test_layer = factorized_embedding.FactorizedEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ output_dim=output_dim)
+ # Create a 2-dimensional input (the first dimension is implicit).
+ sequence_length = 23
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
+ output_tensor = test_layer(input_tensor)
+
+ # The output should be the same as the input, save that it has an extra
+ # embedding_width dimension on the end.
+ expected_output_shape = [None, sequence_length, output_dim]
+ self.assertEqual(expected_output_shape, output_tensor.shape.as_list())
+ self.assertEqual(output_tensor.dtype, tf.float32)
+
+ def test_layer_invocation(self):
+ vocab_size = 31
+ embedding_width = 27
+ output_dim = 45
+ test_layer = factorized_embedding.FactorizedEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ output_dim=output_dim)
+ # Create a 2-dimensional input (the first dimension is implicit).
+ sequence_length = 23
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
+ output_tensor = test_layer(input_tensor)
+
+ # Create a model from the test layer.
+ model = tf_keras.Model(input_tensor, output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 3
+ input_data = np.random.randint(
+ vocab_size, size=(batch_size, sequence_length))
+ output = model.predict(input_data)
+ self.assertEqual(tf.float32, output.dtype)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/nlp/modeling/layers/gated_feedforward.py b/official/nlp/modeling/layers/gated_feedforward.py
index 54db81eab0b..4430fc53193 100644
--- a/official/nlp/modeling/layers/gated_feedforward.py
+++ b/official/nlp/modeling/layers/gated_feedforward.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,12 +16,15 @@
# pylint: disable=g-classes-have-attributes
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
+from official.nlp.modeling.layers import util
-@tf.keras.utils.register_keras_serializable(package="Text")
+
+@tf_keras.utils.register_keras_serializable(package="Text")
@gin.configurable
-class GatedFeedforward(tf.keras.layers.Layer):
+class GatedFeedforward(tf_keras.layers.Layer):
"""Gated linear feedforward layer.
This layer follows the paper "GLU Variants Improve Transformer"
@@ -34,7 +37,7 @@ class GatedFeedforward(tf.keras.layers.Layer):
dropout: Dropout probability for the output dropout.
use_gate: Whether to use gated linear units. If True, assuming `GELU` as the
activation and omitting bias, will apply
- `GEGLU(x, W, V, W_2) = (GEGLU(xW) * xV)W2`; if False, will follow
+ `GEGLU(x, W, V, W_2) = (GELU(xW) * xV)W2`; if False, will follow
"Attention Is All You Need" (https://arxiv.org/abs/1706.03762) paper and
apply `FFN(x, W, W_2) = GELU(xW_1)W_2.`
num_blocks: The number of feedforward blocks to stack. Each block contains a
@@ -55,9 +58,9 @@ class GatedFeedforward(tf.keras.layers.Layer):
"""
def __init__(self,
- intermediate_size,
- intermediate_activation,
- dropout,
+ inner_dim=768,
+ inner_activation=tf_utils.get_activation("gelu"),
+ dropout=0.0,
use_gate=True,
apply_output_layer_norm=True,
num_blocks=1,
@@ -70,9 +73,12 @@ def __init__(self,
kernel_constraint=None,
bias_constraint=None,
**kwargs):
- super(GatedFeedforward, self).__init__(**kwargs)
- self._intermediate_size = intermediate_size
- self._intermediate_activation = intermediate_activation
+ inner_dim = kwargs.pop("intermediate_size", inner_dim)
+ inner_activation = kwargs.pop("intermediate_activation", inner_activation)
+ util.filter_kwargs(kwargs)
+ super().__init__(**kwargs)
+ self._inner_dim = inner_dim
+ self._inner_activation = inner_activation
self._dropout = dropout
self._use_gate = use_gate
self._num_blocks = num_blocks
@@ -83,32 +89,30 @@ def __init__(self,
"The dropout_position should be either `before_residual` or"
"`after_residual`, got: %s" % self._dropout_position)
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
def build(self, input_shape):
hidden_size = input_shape.as_list()[-1]
common_kwargs = dict(
- kernel_initializer=self._kernel_initializer,
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
kernel_constraint=self._kernel_constraint,
bias_constraint=self._bias_constraint)
self._intermediate_dense = []
- self._intermediate_activation_layers = []
+ self._inner_activation_layers = []
self._gate_dense = []
self._output_dense = []
self._output_dropout = []
self._output_layer_norm = []
- activation_policy = tf.keras.mixed_precision.global_policy()
+ activation_policy = tf_keras.mixed_precision.global_policy()
if activation_policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
@@ -116,35 +120,47 @@ def build(self, input_shape):
activation_policy = tf.float32
for i in range(self._num_blocks):
self._intermediate_dense.append(
- tf.keras.layers.experimental.EinsumDense(
+ tf_keras.layers.EinsumDense(
"abc,cd->abd",
- output_shape=(None, self._intermediate_size),
+ output_shape=(None, self._inner_dim),
bias_axes="d",
name="intermediate_%d" % i,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(
+ self._bias_initializer),
**common_kwargs))
- self._intermediate_activation_layers.append(
- tf.keras.layers.Activation(
- self._intermediate_activation, dtype=activation_policy))
+ self._inner_activation_layers.append(
+ tf_keras.layers.Activation(
+ self._inner_activation, dtype=activation_policy))
if self._use_gate:
self._gate_dense.append(
- tf.keras.layers.experimental.EinsumDense(
+ tf_keras.layers.EinsumDense(
"abc,cd->abd",
- output_shape=(None, self._intermediate_size),
+ output_shape=(None, self._inner_dim),
bias_axes="d",
name="gate_%d" % i,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(
+ self._bias_initializer),
**common_kwargs))
self._output_dense.append(
- tf.keras.layers.experimental.EinsumDense(
+ tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
name="output_%d" % i,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(
+ self._bias_initializer),
**common_kwargs))
- self._output_dropout.append(tf.keras.layers.Dropout(rate=self._dropout))
+ self._output_dropout.append(tf_keras.layers.Dropout(rate=self._dropout))
# Use float32 in layernorm for numeric stability.
if self._apply_output_layer_norm:
self._output_layer_norm.append(
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="output_layer_norm_%d" % i,
axis=-1,
epsilon=1e-12,
@@ -152,10 +168,10 @@ def build(self, input_shape):
def get_config(self):
config = {
- "intermediate_size":
- self._intermediate_size,
- "intermediate_activation":
- self._intermediate_activation,
+ "inner_dim":
+ self._inner_dim,
+ "inner_activation":
+ self._inner_activation,
"dropout":
self._dropout,
"use_gate":
@@ -165,21 +181,21 @@ def get_config(self):
"dropout_position":
self._dropout_position,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint)
+ tf_keras.constraints.serialize(self._bias_constraint)
}
- base_config = super(GatedFeedforward, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def call(self, inputs):
@@ -187,7 +203,7 @@ def call(self, inputs):
for i in range(self._num_blocks):
layer_input = layer_output
intermediate_output = self._intermediate_dense[i](layer_input)
- intermediate_output = self._intermediate_activation_layers[i](
+ intermediate_output = self._inner_activation_layers[i](
intermediate_output)
if self._use_gate:
gated_linear = self._gate_dense[i](layer_input)
diff --git a/official/nlp/modeling/layers/gated_feedforward_test.py b/official/nlp/modeling/layers/gated_feedforward_test.py
index 8f69cd4fa58..757a55633aa 100644
--- a/official/nlp/modeling/layers/gated_feedforward_test.py
+++ b/official/nlp/modeling/layers/gated_feedforward_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,20 +16,16 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import gated_feedforward
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class GatedFeedforwardTest(keras_parameterized.TestCase):
+class GatedFeedforwardTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(GatedFeedforwardTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
@parameterized.parameters(
(True, 1, "after_residual", "float32"),
@@ -42,10 +38,10 @@ def tearDown(self):
(False, 1, "before_residual", "mixed_float16"),
)
def test_layer_creation(self, use_gate, num_blocks, dropout_position, dtype):
- tf.keras.mixed_precision.set_global_policy(dtype)
+ tf_keras.mixed_precision.set_global_policy(dtype)
kwargs = dict(
- intermediate_size=128,
- intermediate_activation="relu",
+ inner_dim=128,
+ inner_activation="relu",
dropout=0.1,
use_gate=use_gate,
num_blocks=num_blocks,
@@ -57,7 +53,7 @@ def test_layer_creation(self, use_gate, num_blocks, dropout_position, dtype):
sequence_length = 64
width = 128
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -74,10 +70,10 @@ def test_layer_creation(self, use_gate, num_blocks, dropout_position, dtype):
)
def test_layer_invocation(self, use_gate, num_blocks, dropout_position,
dtype):
- tf.keras.mixed_precision.set_global_policy(dtype)
+ tf_keras.mixed_precision.set_global_policy(dtype)
kwargs = dict(
- intermediate_size=16,
- intermediate_activation="relu",
+ inner_dim=16,
+ inner_activation="relu",
dropout=0.1,
use_gate=use_gate,
num_blocks=num_blocks,
@@ -89,11 +85,11 @@ def test_layer_invocation(self, use_gate, num_blocks, dropout_position,
sequence_length = 16
width = 32
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(data_tensor, output_tensor)
+ model = tf_keras.Model(data_tensor, output_tensor)
# Invoke the model on test data.
batch_size = 6
@@ -104,8 +100,8 @@ def test_layer_invocation(self, use_gate, num_blocks, dropout_position,
def test_serialize_deserialize(self):
kwargs = dict(
- intermediate_size=16,
- intermediate_activation="relu",
+ inner_dim=16,
+ inner_activation="relu",
dropout=0.1,
use_gate=False,
num_blocks=4,
diff --git a/official/nlp/modeling/layers/gaussian_process.py b/official/nlp/modeling/layers/gaussian_process.py
index ac47eaf2f02..e23b145b496 100644
--- a/official/nlp/modeling/layers/gaussian_process.py
+++ b/official/nlp/modeling/layers/gaussian_process.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,13 @@
"""Definitions for random feature Gaussian process layer."""
import math
-import tensorflow as tf
+import tensorflow as tf, tf_keras
_SUPPORTED_LIKELIHOOD = ('binary_logistic', 'poisson', 'gaussian')
-class RandomFeatureGaussianProcess(tf.keras.layers.Layer):
+class RandomFeatureGaussianProcess(tf_keras.layers.Layer):
"""Gaussian process layer with random feature approximation [1].
During training, the model updates the maximum a posteriori (MAP) logits
@@ -98,7 +98,7 @@ def __init__(self,
scale_random_features: (bool) Whether to scale the random feature
by sqrt(2. / num_inducing).
use_custom_random_features: (bool) Whether to use custom random
- features implemented using tf.keras.layers.Dense.
+ features implemented using tf_keras.layers.Dense.
custom_random_features_initializer: (str or callable) Initializer for
the random features. Default to random normal which approximates a RBF
kernel function if activation function is cos.
@@ -116,7 +116,7 @@ def __init__(self,
name: (string) Layer name.
**gp_output_kwargs: Additional keyword arguments to dense output layer.
"""
- super(RandomFeatureGaussianProcess, self).__init__(name=name, dtype=dtype)
+ super().__init__(name=name, dtype=dtype)
self.units = units
self.num_inducing = num_inducing
@@ -151,14 +151,14 @@ def __init__(self,
minval=0., maxval=2. * math.pi)
if self.custom_random_features_initializer is None:
self.custom_random_features_initializer = (
- tf.keras.initializers.RandomNormal(stddev=1.))
+ tf_keras.initializers.RandomNormal(stddev=1.))
if self.custom_random_features_activation is None:
self.custom_random_features_activation = tf.math.cos
def build(self, input_shape):
# Defines model layers.
if self.normalize_input:
- self._input_norm_layer = tf.keras.layers.LayerNormalization(
+ self._input_norm_layer = tf_keras.layers.LayerNormalization(
name='gp_input_normalization')
self._input_norm_layer.build(input_shape)
input_shape = self._input_norm_layer.compute_output_shape(input_shape)
@@ -177,10 +177,10 @@ def build(self, input_shape):
name='gp_covariance')
self._gp_cov_layer.build(input_shape)
- self._gp_output_layer = tf.keras.layers.Dense(
+ self._gp_output_layer = tf_keras.layers.Dense(
units=self.units,
use_bias=False,
- kernel_regularizer=tf.keras.regularizers.l2(self.l2_regularization),
+ kernel_regularizer=tf_keras.regularizers.l2(self.l2_regularization),
dtype=self.dtype,
name='gp_output_weights',
**self.gp_output_kwargs)
@@ -197,8 +197,8 @@ def build(self, input_shape):
def _make_random_feature_layer(self, name):
"""Defines random feature layer depending on kernel type."""
if not self.use_custom_random_features:
- # Use default RandomFourierFeatures layer from tf.keras.
- return tf.keras.layers.experimental.RandomFourierFeatures(
+ # Use default RandomFourierFeatures layer from tf_keras.
+ return tf_keras.layers.experimental.RandomFourierFeatures(
output_dim=self.num_inducing,
kernel_initializer=self.gp_kernel_type,
scale=self.gp_kernel_scale,
@@ -207,11 +207,11 @@ def _make_random_feature_layer(self, name):
name=name)
if self.gp_kernel_type.lower() == 'linear':
- custom_random_feature_layer = tf.keras.layers.Lambda(
+ custom_random_feature_layer = tf_keras.layers.Lambda(
lambda x: x, name=name)
else:
# Use user-supplied configurations.
- custom_random_feature_layer = tf.keras.layers.Dense(
+ custom_random_feature_layer = tf_keras.layers.Dense(
units=self.num_inducing,
use_bias=True,
activation=self.custom_random_features_activation,
@@ -226,7 +226,7 @@ def reset_covariance_matrix(self):
"""Resets covariance matrix of the GP layer.
This function is useful for reseting the model's covariance matrix at the
- begining of a new epoch.
+ beginning of a new epoch.
"""
self._gp_cov_layer.reset_precision_matrix()
@@ -260,14 +260,14 @@ def call(self, inputs, global_step=None, training=None):
# Assembles model output.
model_output = [gp_output,]
if self.return_gp_cov:
- model_output.append(gp_covmat)
+ model_output.append(gp_covmat) # pyrefly: ignore[unbound-name]
if self.return_random_features:
model_output.append(gp_feature)
return model_output
-class LaplaceRandomFeatureCovariance(tf.keras.layers.Layer):
+class LaplaceRandomFeatureCovariance(tf_keras.layers.Layer):
"""Computes the Gaussian Process covariance using Laplace method.
At training time, this layer updates the Gaussian process posterior using
@@ -324,7 +324,7 @@ def build(self, input_shape):
name='gp_precision_matrix',
shape=(gp_feature_dim, gp_feature_dim),
dtype=self.dtype,
- initializer=tf.keras.initializers.Identity(self.ridge_penalty),
+ initializer=tf_keras.initializers.Identity(self.ridge_penalty),
trainable=False,
aggregation=tf.VariableAggregation.ONLY_FIRST_REPLICA))
self.built = True
@@ -380,7 +380,7 @@ def reset_precision_matrix(self):
"""Resets precision matrix to its initial value.
This function is useful for reseting the model's covariance matrix at the
- begining of a new epoch.
+ beginning of a new epoch.
"""
precision_matrix_reset_op = self.precision_matrix.assign(
self.initial_precision_matrix)
@@ -417,7 +417,7 @@ def compute_predictive_covariance(self, gp_feature):
def _get_training_value(self, training=None):
if training is None:
- training = tf.keras.backend.learning_phase()
+ training = tf_keras.backend.learning_phase()
if isinstance(training, int):
training = bool(training)
diff --git a/official/nlp/modeling/layers/gaussian_process_test.py b/official/nlp/modeling/layers/gaussian_process_test.py
index 7a9a56fe452..043e1b26704 100644
--- a/official/nlp/modeling/layers/gaussian_process_test.py
+++ b/official/nlp/modeling/layers/gaussian_process_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import gaussian_process
@@ -205,9 +205,9 @@ def test_state_saving_and_loading(self):
input_data = np.random.random((1, 2))
rfgp_model = gaussian_process.RandomFeatureGaussianProcess(units=1)
- inputs = tf.keras.Input((2,), batch_size=1)
+ inputs = tf_keras.Input((2,), batch_size=1)
outputs = rfgp_model(inputs)
- model = tf.keras.Model(inputs, outputs)
+ model = tf_keras.Model(inputs, outputs)
gp_output, gp_covmat = model.predict(input_data)
# Save and then load the model.
@@ -215,7 +215,7 @@ def test_state_saving_and_loading(self):
self.addCleanup(shutil.rmtree, temp_dir)
saved_model_dir = os.path.join(temp_dir, 'rfgp_model')
model.save(saved_model_dir)
- new_model = tf.keras.models.load_model(saved_model_dir)
+ new_model = tf_keras.models.load_model(saved_model_dir)
gp_output_new, gp_covmat_new = new_model.predict(input_data)
self.assertAllClose(gp_output, gp_output_new, atol=1e-4)
diff --git a/official/nlp/modeling/layers/kernel_attention.py b/official/nlp/modeling/layers/kernel_attention.py
index cce2ce9f96b..53767ebb8d6 100644
--- a/official/nlp/modeling/layers/kernel_attention.py
+++ b/official/nlp/modeling/layers/kernel_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,12 +16,14 @@
import functools
import math
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
_NUMERIC_STABLER = 1e-6
-class KernelMask(tf.keras.layers.Layer):
+class KernelMask(tf_keras.layers.Layer):
"""Creates kernel attention mask.
inputs: from_tensor: 2D or 3D Tensor of shape
@@ -39,6 +41,242 @@ def call(self, inputs, mask):
return mask
+def pad_to_chunk_length(tensor, axis, chunk_length, padding=None):
+ """Pads a tensor so that shape[axis] is divisible by chunk_length.
+
+ Args:
+ tensor: Input tensor to pad.
+ axis: Axis to pad along.
+ chunk_length: The output tensor will have shape[axis] divisible by
+ chunk_length.
+ padding: Pad the input tensor across the axis from either left or right if
+ padding is set to "left" or "right"; applies no padding if padding is set
+ to None. In the latter case, the axis dimension of the input tensor must
+ be divisible by the chunk_length.
+
+ Returns:
+ Padded tensor with shape[axis] divisible by chunk_length.
+ """
+ if padding is None:
+ return tensor
+ shape = tf.shape(tensor)
+ rank = tf.rank(tensor)
+ if axis < 0:
+ axis += rank
+ axis_length = shape[axis]
+ pad_length = -axis_length % chunk_length
+ if padding == "right":
+ axis_paddings = [[0, pad_length]]
+ elif padding == "left":
+ axis_paddings = [[pad_length, 0]]
+ else:
+ raise ValueError(
+ "Illegal padding value; must be one of \"left\", \"right\" or None.")
+ paddings = tf.concat([
+ tf.zeros([axis, 2], dtype=tf.int32), axis_paddings,
+ tf.zeros([rank - axis - 1, 2], dtype=tf.int32)
+ ],
+ axis=0)
+ return tf.pad(tensor, paddings)
+
+
+def split_tensor_into_chunks(tensor, axis, chunk_length):
+ """Reshape tensor along given axis using chunk_length.
+
+ Args:
+ tensor: Input tensor.
+ axis: Reshape tensor along this axis.
+ chunk_length: Split the axis into [axis/chunk_length, chunk_length]
+
+ Returns:
+ Reshaped tensor.
+ """
+ shape = tf.shape(tensor)
+ num_chunks = shape[axis] // chunk_length
+ new_shape = tf.concat(
+ [shape[:axis], [num_chunks, chunk_length], shape[(axis + 1):]], axis=0)
+ return tf.reshape(tensor, new_shape)
+
+
+def rectangular_window_sum(tensor, window_length):
+ """Summarizes tensor elements over a sliding rectangular window.
+
+ Sums elements of the input tensor of shape [B, T', C', H, dim]
+ across a rectangular window sliding along the dimension T'.
+
+ Args:
+ tensor: Tensor of shape `[B, T', C', H, dim]`.
+ window_length: The length of the rectangular window.
+
+ Returns:
+ A tensor of shape [B, T', C', H, dim] containing sums over the
+ window.
+ """
+ tensor_cumsum = tf.cumsum(tensor, axis=-4)
+ tensor_winsum = tensor_cumsum - tf.pad(
+ tensor_cumsum,
+ [[0, 0], [window_length, 0], [0, 0], [0, 0], [0, 0]])[:, :-window_length]
+ return tensor_winsum
+
+
+def weighted_window_sum(tensor, window_length, window_weights):
+ """Summarizes tensor elements over a sliding weighted window.
+
+ Computes a weighted sum of elements of the input tensor of shape [B,
+ T', C', H, dim] across a window sliding along the dimension T'.
+
+ Args:
+ tensor: Tensor of shape `[B, T', C', H, dim]`.
+ window_length: The length of the window.
+ window_weights: Tensor of shape [window_length] containing window weights.
+
+ Returns:
+ A tensor of shape [B, T', C', H, dim] containing sums over the
+ window.
+ """
+ # Flatten the last three dimensions of the [B, T', C', H, dim] shape
+ # into a single channels dimension.
+ tensor_shape = tf.shape(tensor)
+ tensor_2d = tf.reshape(tensor, [tensor_shape[0], tensor_shape[1], 1, -1])
+
+ # Apply the same weights to all channels.
+ conv_filter = tf.tile(
+ tf.reshape(window_weights, [-1, 1, 1, 1]),
+ multiples=[1, 1, tf.shape(tensor_2d)[-1], 1])
+ tensor_winsum_2d = tf.nn.depthwise_conv2d(
+ tensor_2d,
+ conv_filter,
+ strides=[1, 1, 1, 1],
+ padding=[[0, 0], [window_length - 1, 0], [0, 0], [0, 0]])
+
+ # Unflatten the channels dimension into the original shape.
+ tensor_winsum = tf.reshape(tensor_winsum_2d, tensor_shape)
+ return tensor_winsum
+
+
+def causal_windowed_performer_attention(query_matrix,
+ key_matrix,
+ value_matrix,
+ chunk_length,
+ window_length,
+ window_decay=None,
+ padding=None,
+ cache=None):
+ """Applies windowed causal kernel attention with query, key, value tensors.
+
+ We partition the T-length input sequence into N chunks, each of
+ chunk_length tokens (thus: T = N * chunk_length). Within each chunk,
+ we apply bidirectional (non-causal) Performers’ implicit attention
+ and we model relationships between different chunks using
+ Performers’ causal attention. We consider windowed causal variant of
+ performer, where the current chunk attends only to the window of
+ window_length of the most recent chunks.
+
+ Below is an example with T=9, chunk_length=3, window_length=2. In
+ this example 1 indicates attention is computed between the pair
+ while 0 indicates attention is not computed between the pairs:
+
+ 111000000
+ 111000000
+ 111000000
+ 111111000
+ 111111000
+ 111111000
+ 000111111
+ 000111111
+ 000111111
+
+ User can ensure sequence_length is divisible by chunk_length or use
+ padding="left"/"right" to pad the sequence length either at the left
+ or right respectively and make it divisible by chunk_length.
+
+ Args:
+ query_matrix: Kernel query `Tensor` of shape `[B, T, H, dim]`.
+ key_matrix: Kernel key `Tensor` of shape `[B, T, H, dim]`.
+ value_matrix: Value `Tensor` of shape `[B, T, H, out_dim]`.
+ chunk_length: Length of each chunk in tokens.
+ window_length: Length of attention window in chunks.
+ window_decay: Float window decay factor or `None`. If set, exponentially
+ decay past attention window values by this factor before summation.
+ padding: Pad the query, value and key input tensors across the axis from
+ either left or right if padding is set to "left" or "right"; apply no
+ padding if padding is set to None. In the latter case, the axis dimension
+ of the query, value and key input tensors must be divisible by the
+ chunk_length.
+ cache: Cache to accumulate history in memory. Used at inferecne time
+ (streaming, decoding) for causal attention.
+
+ Returns:
+ Window causal performer attention of shape `[B, T, H, out_dim]`.
+ """
+ if cache is None: # Training
+ old_shape = tf.shape(value_matrix)
+
+ query_matrix = pad_to_chunk_length(query_matrix, -3, chunk_length, padding)
+ key_matrix = pad_to_chunk_length(key_matrix, -3, chunk_length, padding)
+ value_matrix = pad_to_chunk_length(value_matrix, -3, chunk_length, padding)
+
+ new_shape = tf.shape(value_matrix)
+ chunked_query_matrix = split_tensor_into_chunks(
+ query_matrix, -3,
+ chunk_length) # [-1, T//chunk_length, chunk_length, N, dim]
+ chunked_key_matrix = split_tensor_into_chunks(
+ key_matrix, -3,
+ chunk_length) # [-1, T//chunk_length, chunk_length, N, dim]
+ chunked_value_matrix = split_tensor_into_chunks(
+ value_matrix, -3,
+ chunk_length) # [-1, T//chunk_length, chunk_length, N, out_dim]
+
+ kp_v = tf.einsum("BTCHD,BTCHO->BTHDO", chunked_key_matrix,
+ chunked_value_matrix)
+
+ k_sum = tf.math.reduce_sum(chunked_key_matrix, axis=-3, keepdims=True)
+
+ if window_decay is None:
+ kp_v_winsum = rectangular_window_sum(kp_v, window_length)
+ k_winsum = rectangular_window_sum(k_sum, window_length)
+ else:
+ # Compute exponentially decaying weights.
+ decaying_weights = tf.math.pow(
+ tf.convert_to_tensor(window_decay, dtype=value_matrix.dtype),
+ tf.range(window_length - 1, -1, delta=-1, dtype=value_matrix.dtype))
+ kp_v_winsum = weighted_window_sum(kp_v, window_length, decaying_weights)
+ k_winsum = weighted_window_sum(k_sum, window_length, decaying_weights)
+
+ numerator = tf.einsum(
+ "BTCHD,BTHDO->BTCHO", chunked_query_matrix, kp_v_winsum)
+
+ k_winsum = tf.squeeze(k_winsum, -3)
+ denominator = tf.einsum("BTCHD,BTHD->BTCH", chunked_query_matrix, k_winsum)
+ denominator = tf.expand_dims(denominator, -1) + _NUMERIC_STABLER
+ attention = numerator / denominator
+ attention = tf.reshape(attention, new_shape)
+
+ start = tf.zeros([old_shape.shape[0]], dtype=old_shape.dtype)
+ attention = tf.slice(attention, start, old_shape)
+
+ # Queued window cache (drop instead of decay) not yet supported.
+ else: # Streaming
+
+ if window_decay is None or window_decay > 1.0 or window_decay < 0.0:
+ raise ValueError("window_decay should be in (0.0, 1.0) and not None.")
+ kv = window_decay * cache["kv"] + tf.einsum(
+ "BTHD,BTHO->BHOD", key_matrix, value_matrix)
+ cache["kv"] = kv
+ k_sum = window_decay * cache["k_sum"] + tf.reduce_sum(key_matrix, axis=1)
+ cache["k_sum"] = k_sum
+ denominator = tf.einsum("BTHD,BHD->BTH", query_matrix, k_sum)
+ # The below is equivalent to but converts to TF Lite better than:
+ # tf.einsum("BTHD,BTH->BTHD",
+ # query_matrix, 1.0 / (denominator + _NUMERIC_STABLER))
+ inverse_denominator = 1.0 / (denominator + _NUMERIC_STABLER)
+ # Add another dimension to align for the broadcast multiplication.
+ fused_query_denominator = query_matrix * tf.expand_dims(inverse_denominator,
+ -1)
+ attention = tf.einsum("BTHD,BHOD->BTHO", fused_query_denominator, kv)
+ return attention
+
+
def create_projection_matrix(m, d, seed=None):
r"""Constructs the matrix of random projections.
@@ -56,8 +294,8 @@ def create_projection_matrix(m, d, seed=None):
The matrix of random projections of the shape [m, d].
"""
nb_full_blocks = math.ceil(m / d)
- block_list = tf.TensorArray(tf.float32,
- size=tf.cast(nb_full_blocks, dtype=tf.int32))
+ block_list = tf.TensorArray(
+ tf.float32, size=tf.cast(nb_full_blocks, dtype=tf.int32))
stateful = False
if seed is None:
stateful = True
@@ -85,11 +323,13 @@ def create_projection_matrix(m, d, seed=None):
return tf.linalg.matmul(tf.linalg.diag(multiplier), final_matrix)
-def _generalized_kernel(x, projection_matrix, f, h):
+def _generalized_kernel(x, y, is_query, projection_matrix, f, h):
"""Generalized kernel in RETHINKING ATTENTION WITH PERFORMERS.
Args:
x: The feature being transformed with shape [B, T, N ,H].
+ y: The extra stats-tensor of shape [B, T, N ,H].
+ is_query: True if x is a query-tensor.
projection_matrix: The matrix with shape [M, H] that we projecct x to, where
M is the number of projections.
f: A non-linear function applied on x or projected x.
@@ -99,7 +339,8 @@ def _generalized_kernel(x, projection_matrix, f, h):
Returns:
Transformed feature.
"""
-
+ del y
+ del is_query
if projection_matrix is None:
return h(x) * f(x)
else:
@@ -108,28 +349,147 @@ def _generalized_kernel(x, projection_matrix, f, h):
tf.cast(tf.shape(projection_matrix)[0], tf.float32))
+def expplus(data_orig,
+ other_data,
+ is_query,
+ projection_matrix=None,
+ numerical_stabilizer=0.000001,
+ normalize_data=True,
+ numerical_renormalizer=True,
+ extra_renormalize_exp_fun=False):
+ """FAVOR++ mechanism from the CRT paper: https://arxiv.org/abs/2205.15317 .
+
+ Args:
+ data_orig: data tensor of shape [B,T,H,D] for which random features aree to
+ be computed
+ other_data: additional tensor of the shape [B,F,H,D] used to collect stats
+ to determine the exact instantiation of the random feature mechanism
+ is_query: boolean indicating whether tensor is a query tensor
+ projection_matrix: tensor of the shape [M,D] encoding random projections for
+ random features (M stands for the number of random features)
+ numerical_stabilizer: numerical stabilizer for the kernel features
+ normalize_data: whether to sqrt-d-normalize queries/keys as in the regular
+ attention
+ numerical_renormalizer: whether to apply additional renormalization for
+ numerical stability
+ extra_renormalize_exp_fun: extra renormalizer for the exponential mapping
+ applied to construct random features
+
+ Returns:
+ Random feature map tensor for the unbiased softmax-kernel estimation.
+ """
+
+ data = data_orig
+ if projection_matrix is None:
+ return data_orig
+ projection_matrix = tf.cast(projection_matrix, data.dtype)
+ if normalize_data:
+ data_normalizer = 1.0 / tf.math.sqrt(
+ (tf.math.sqrt(tf.dtypes.cast(data.shape[-1], data.dtype))))
+ else:
+ data_normalizer = 1.0
+ lengths = tf.math.square(data)
+ lengths = tf.reduce_sum(lengths, axis=tf_keras.backend.ndim(data) - 1)
+ lengths = tf.expand_dims(lengths, axis=tf_keras.backend.ndim(data) - 1)
+ lengths = tf.math.sqrt(lengths)
+ data /= lengths
+ ratio = 1.0 / tf.math.sqrt(
+ tf.dtypes.cast(projection_matrix.shape[0], data.dtype))
+ data_dash = tf.einsum("blhd,md->blhm", data_normalizer * data,
+ projection_matrix)
+ diag_data = tf.math.square(data)
+ diag_data = tf.math.reduce_sum(
+ diag_data, axis=tf_keras.backend.ndim(data) - 1)
+ diag_data = (diag_data / 2.0) * data_normalizer * data_normalizer
+ diag_data = tf.expand_dims(diag_data, axis=tf_keras.backend.ndim(data) - 1)
+
+ # Calculating coefficients A, B of the FAVOR++ mechanism:
+ _, l, _, _ = tf_utils.get_shape_list(data_orig)
+
+ l = tf.cast(l, dtype=tf.float32)
+ first_sum_of_squares = tf.math.square(data)
+ first_sum_of_squares = tf.math.reduce_sum(
+ first_sum_of_squares, axis=(1, -1), keepdims=True)
+ first_sum_of_squares *= (data_normalizer * data_normalizer)
+ first_sum_of_squares /= l # data.shape[1]
+ second_sum_of_squares = tf.math.square(other_data)
+ second_sum_of_squares = tf.math.reduce_sum(
+ second_sum_of_squares, axis=(1, -1), keepdims=True)
+ second_sum_of_squares *= (data_normalizer * data_normalizer)
+ second_sum_of_squares /= l # other_data.shape[1]
+ data_sum = tf.math.reduce_sum(data, axis=(1,), keepdims=True)
+ other_data_sum = tf.math.reduce_sum(other_data, axis=(1,), keepdims=True)
+ d_prod = tf.einsum("blhd,blhd->blh", data_sum, other_data_sum)
+ d_prod = tf.expand_dims(d_prod, axis=-1)
+ d_prod *= (data_normalizer * data_normalizer)
+ d_prod *= (2.0 / (l * l))
+ ave = first_sum_of_squares + second_sum_of_squares + d_prod
+ dim = projection_matrix.shape[-1]
+ a_coeff = (1.0 / (4.0 * ave)) * (
+ tf.math.sqrt((2.0 * ave + dim) *
+ (2.0 * ave + dim) + 8.0 * dim * ave) - 2.0 * ave - dim)
+ a_coeff = (1.0 - 1.0 / a_coeff) / 8.0
+ b_coeff = tf.math.sqrt(1.0 - 4.0 * a_coeff)
+ d_coeff = tf.math.pow(1.0 - 4.0 * a_coeff, dim / 4.0)
+ a_coeff = tf.stop_gradient(a_coeff)
+ b_coeff = tf.stop_gradient(b_coeff)
+ d_coeff = tf.stop_gradient(d_coeff)
+
+ # Calculating diag_omega for the FAVOR++ mechanism:
+ diag_omega = tf.math.square(projection_matrix)
+ diag_omega = tf.math.reduce_sum(
+ diag_omega, axis=tf_keras.backend.ndim(projection_matrix) - 1)
+ diag_omega = tf.expand_dims(diag_omega, axis=0)
+ diag_omega = tf.expand_dims(diag_omega, axis=0)
+ diag_omega = tf.expand_dims(diag_omega, axis=0)
+ diag_omega = a_coeff * diag_omega
+
+ if numerical_renormalizer:
+ if is_query:
+ last_dims_t = (len(data_dash.shape) - 1,)
+ stab = b_coeff * tf.math.reduce_max(
+ data_dash, axis=last_dims_t, keepdims=True)
+ else:
+ stab = b_coeff * tf.math.reduce_max(data_dash, keepdims=True)
+ if extra_renormalize_exp_fun:
+ extra_stab = tf.reduce_max(diag_data, axis=1, keepdims=True)
+ stab = tf.math.maximum(stab, extra_stab)
+ data_dash = ratio * d_coeff * (
+ tf.math.exp(b_coeff * data_dash - stab - diag_data + diag_omega) +
+ numerical_stabilizer)
+ else:
+ data_dash = ratio * d_coeff * (
+ tf.math.exp(b_coeff * data_dash - diag_data + diag_omega) +
+ numerical_stabilizer)
+
+ return data_dash
+
+
# pylint: disable=g-long-lambda
-_TRANSFORM_MAP = {
+_CAUSAL_SUPPORT_TRANSFORM_MAP = {
"elu":
functools.partial(
_generalized_kernel,
- f=lambda x: tf.keras.activations.elu(x) + 1,
+ f=lambda x: tf_keras.activations.elu(x) + 1,
h=lambda x: 1),
"relu":
functools.partial(
- _generalized_kernel, f=tf.keras.activations.relu, h=lambda x: 1),
+ _generalized_kernel,
+ # Improve numerical stability and avoid NaNs in some cases by adding
+ # a tiny epsilon.
+ f=lambda x: tf_keras.activations.relu(x) + 1e-3,
+ h=lambda x: 1),
"square":
- functools.partial(
- _generalized_kernel, f=tf.math.square, h=lambda x: 1),
+ functools.partial(_generalized_kernel, f=tf.math.square, h=lambda x: 1),
"exp":
functools.partial(
_generalized_kernel,
# Avoid exp explosion by shifting.
- f=lambda x: tf.math.exp(
- x - tf.math.reduce_max(x, axis=[1, 2, 3], keepdims=True)),
- h=lambda x: tf.math.exp(
- -0.5 * tf.math.reduce_sum(
- tf.math.square(x), axis=-1, keepdims=True)),),
+ f=lambda x: tf.math.exp(x - tf.math.reduce_max(
+ x, axis=[1, 2, 3], keepdims=True)),
+ h=lambda x: tf.math.exp(-0.5 * tf.math.reduce_sum(
+ tf.math.square(x), axis=-1, keepdims=True)),
+ ),
"expmod":
functools.partial(
_generalized_kernel,
@@ -142,10 +502,20 @@ def _generalized_kernel(x, projection_matrix, f, h):
"identity":
functools.partial(_generalized_kernel, f=lambda x: x, h=lambda x: 1)
}
+
+_NON_CAUSAL_SUPPORT_TRANSFORM_MAP = {
+ "expplus": expplus,
+}
+
+_TRANSFORM_MAP = {
+ **_CAUSAL_SUPPORT_TRANSFORM_MAP,
+ **_NON_CAUSAL_SUPPORT_TRANSFORM_MAP
+}
+
# pylint: enable=g-long-lambda
-class KernelAttention(tf.keras.layers.MultiHeadAttention):
+class KernelAttention(tf_keras.layers.MultiHeadAttention):
"""A variant of efficient transformers which replaces softmax with kernels.
This module combines ideas from the two following papers:
@@ -154,6 +524,9 @@ class KernelAttention(tf.keras.layers.MultiHeadAttention):
(https://arxiv.org/abs/2009.14794)
- exp (Lemma 1, positive), relu
- random/deterministic projection
+ Chefs' Random Tables: Non-Trigonometric Random Features
+ (https://arxiv.org/abs/2205.15317)
+ - expplus (OPRF mechanism)
Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention
(https://arxiv.org/abs/2006.16236)
@@ -179,12 +552,18 @@ def __init__(self,
begin_kernel=0,
scale=None,
scale_by_length=False,
+ use_causal_windowed=False,
+ causal_chunk_length=1,
+ causal_window_length=3,
+ causal_window_decay=None,
+ causal_padding=None,
**kwargs):
r"""Constructor of KernelAttention.
Args:
- feature_transform: A non-linear transform of the keys and quries. Possible
- transforms are "elu", "relu", "square", "exp", "expmod", "identity".
+ feature_transform: A non-linear transform of the keys and queries.
+ Possible transforms are "elu", "relu", "square", "exp", "expplus",
+ "expmod", "identity".
num_random_features: Number of random features to be used for projection.
if num_random_features <= 0, no production is used before transform.
seed: The seed to begin drawing random features. Once the seed is set, the
@@ -204,6 +583,18 @@ def __init__(self,
the dot product based on key length. Set as log_512^(n) to stablize
attention entropy against length. Refer to
https://kexue.fm/archives/8823 for details.
+ use_causal_windowed: If true perform windowed causal attention. See
+ causal_windowed_performer_attention function docstring for more details.
+ causal_chunk_length: Length of each chunk in tokens.
+ causal_window_length: Length of attention window in chunks.
+ causal_window_decay: Float window decay factor or `None`. If set,
+ exponentially decay past attention window values by this factor before
+ summation.
+ causal_padding: Pad the query, value and key input tensors across the axis
+ from either left or right if padding is set to "left" or "right"; apply
+ no padding if padding is set to None. In the latter case, the axis
+ dimension of the query, value and key input tensors must be divisible by
+ the chunk_length.
**kwargs: The same arguments `MultiHeadAttention` layer.
"""
if feature_transform not in _TRANSFORM_MAP:
@@ -233,6 +624,14 @@ def __init__(self,
self._projection_matrix = create_projection_matrix(
self._num_random_features, self._key_dim,
tf.constant([self._seed, self._seed + 1]))
+ self.use_causal_windowed = use_causal_windowed
+ self.causal_chunk_length = causal_chunk_length
+ self.causal_window_length = causal_window_length
+ self.causal_window_decay = causal_window_decay
+ self.causal_padding = causal_padding
+ if self.use_causal_windowed and self._is_short_seq:
+ raise ValueError(
+ "use_causal_windowed and short_seq methods are mutually exclusive")
def _compute_attention(self,
query,
@@ -241,6 +640,7 @@ def _compute_attention(self,
feature_transform,
is_short_seq,
attention_mask=None,
+ cache=None,
training=False,
numeric_stabler=_NUMERIC_STABLER):
"""Applies kernel attention with query, key, value tensors.
@@ -260,6 +660,8 @@ def _compute_attention(self,
attention_mask: a boolean mask of shape `[B, S]`, that prevents attenting
to masked positions. Note that the mask is only appied to the keys. User
may want to mask the output if query contains pads.
+ cache: Cache to accumulate history in memory. Used at inferecne time
+ (streaming, decoding) for causal attention.
training: Python boolean indicating whether the layer should behave in
training mode (adding dropout) or in inference mode (doing nothing).
numeric_stabler: A scalar value added to avoid divide by 0.
@@ -268,6 +670,7 @@ def _compute_attention(self,
attention_output: Multi-headed outputs of attention computation.
"""
projection_matrix = None
+
if self._num_random_features > 0:
if self._redraw and training:
projection_matrix = create_projection_matrix(self._num_random_features,
@@ -293,23 +696,35 @@ def _compute_attention(self,
key *= tf.math.sqrt(scale)
query *= tf.math.sqrt(scale)
- key = _TRANSFORM_MAP[feature_transform](key, projection_matrix)
- query = _TRANSFORM_MAP[feature_transform](query, projection_matrix)
+ key_prime = _TRANSFORM_MAP[feature_transform](key, query, False,
+ projection_matrix)
+ query_prime = _TRANSFORM_MAP[feature_transform](query, key, True,
+ projection_matrix)
if attention_mask is not None:
- key = tf.einsum("BSNH,BS->BSNH", key, attention_mask)
+ key_prime = tf.einsum("BSNH,BS->BSNH", key_prime, attention_mask)
if is_short_seq:
- attention_scores = tf.einsum("BTNH,BSNH->BTSN", query, key)
+ attention_scores = tf.einsum("BTNH,BSNH->BTSN", query_prime, key_prime)
attention_scores = tf.nn.softmax(attention_scores, axis=2)
attention_output = tf.einsum("BTSN,BSNH->BTNH", attention_scores, value)
+ elif self.use_causal_windowed:
+ attention_output = causal_windowed_performer_attention(
+ query_prime,
+ key_prime,
+ value,
+ chunk_length=self.causal_chunk_length,
+ window_length=self.causal_window_length,
+ window_decay=self.causal_window_decay,
+ padding=self.causal_padding,
+ cache=cache)
else:
- kv = tf.einsum("BSNH,BSND->BNDH", key, value)
+ kv = tf.einsum("BSNH,BSND->BNDH", key_prime, value)
denominator = 1.0 / (
- tf.einsum("BTNH,BNH->BTN", query, tf.reduce_sum(key, axis=1)) +
- _NUMERIC_STABLER)
- attention_output = tf.einsum(
- "BTNH,BNDH,BTN->BTND", query, kv, denominator)
+ tf.einsum("BTNH,BNH->BTN", query_prime,
+ tf.reduce_sum(key_prime, axis=1)) + _NUMERIC_STABLER)
+ attention_output = tf.einsum("BTNH,BNDH,BTN->BTND", query_prime, kv,
+ denominator)
return attention_output
def _build_from_signature(self, query, value, key=None):
@@ -324,15 +739,12 @@ def _build_from_signature(self, query, value, key=None):
kernel_constraint=self._kernel_constraint,
bias_constraint=self._bias_constraint)
self._output_dense_softmax = self._make_output_dense(
- self._query_shape.rank - 1, common_kwargs,
+ self._query_shape.rank - 1,
+ common_kwargs,
name="attention_output_softmax")
- self._dropout_softmax = tf.keras.layers.Dropout(rate=self._dropout)
+ self._dropout_softmax = tf_keras.layers.Dropout(rate=self._dropout)
- def call(self,
- query,
- value,
- key=None,
- attention_mask=None,
+ def call(self, query, value, key=None, attention_mask=None, cache=None, # pyrefly: ignore[bad-override]
training=False):
"""Compute attention with kernel mechanism.
@@ -344,12 +756,29 @@ def call(self,
attention_mask: a boolean mask of shape `[B, S]`, that prevents attenting
to masked positions. Note that the mask is only appied to the keys. User
may want to mask the output if query contains pads.
+ cache: Cache to accumulate history in memory. Used at inferecne time
+ (streaming, decoding) for causal attention.
training: Python boolean indicating whether the layer should behave in
training mode (adding dropout) or in inference mode (doing nothing).
Returns:
Multi-headed outputs of attention computation.
"""
+ if cache is not None:
+ if training:
+ raise ValueError(
+ "Cache is not supported when training is True.")
+ if not self.use_causal_windowed:
+ raise ValueError(
+ "Cache is not supported for non use_causal_windowed case.")
+ if self._begin_kernel:
+ raise ValueError(
+ "Cache is not supported when begin_kernel is set since the bahvior "
+ "is too complicated.")
+ if self._feature_transform in _NON_CAUSAL_SUPPORT_TRANSFORM_MAP:
+ raise ValueError("Cache is not supported for feature_transform %s" %
+ (self._feature_transform))
+
if not self._built_from_signature:
self._build_from_signature(query=query, value=value, key=key)
if key is None:
@@ -368,26 +797,26 @@ def call(self,
if self._begin_kernel > 0:
attention_output_softmax = self._compute_attention(
- query[:, :self._begin_kernel],
- key, value, "identity", True, attention_mask, training)
+ query[:, :self._begin_kernel], key, value, "identity", True,
+ attention_mask, training)
attention_output_softmax = self._dropout_softmax(attention_output_softmax)
attention_output_softmax = self._output_dense_softmax(
attention_output_softmax)
attention_output_kernel = self._compute_attention(
- query[:, self._begin_kernel:],
- key, value, self._feature_transform, self._is_short_seq,
- attention_mask, training)
+ query[:, self._begin_kernel:], key, value, self._feature_transform,
+ self._is_short_seq, attention_mask, training)
attention_output_kernel = self._dropout_layer(attention_output_kernel)
- attention_output_kernel = self._output_dense(
- attention_output_kernel)
+ attention_output_kernel = self._output_dense(attention_output_kernel)
attention_output = tf.concat(
[attention_output_softmax, attention_output_kernel], axis=1)
else:
attention_output = self._compute_attention(query, key, value,
self._feature_transform,
self._is_short_seq,
- attention_mask, training)
+ attention_mask,
+ cache,
+ training)
# This is actually dropping out entire tokens to attend to, which might
# seem a bit unusual, but is taken from the original Transformer paper.
attention_output = self._dropout_layer(attention_output)
@@ -403,6 +832,12 @@ def get_config(self):
"is_short_seq": self._is_short_seq,
"begin_kernel": self._begin_kernel,
"scale": self._scale,
+ "scale_by_length": self._scale_by_length,
+ "use_causal_windowed": self.use_causal_windowed,
+ "causal_chunk_length": self.causal_chunk_length,
+ "causal_window_length": self.causal_window_length,
+ "causal_window_decay": self.causal_window_decay,
+ "causal_padding": self.causal_padding,
}
base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
diff --git a/official/nlp/modeling/layers/kernel_attention_test.py b/official/nlp/modeling/layers/kernel_attention_test.py
index 4225cd79408..1b372bfae8b 100644
--- a/official/nlp/modeling/layers/kernel_attention_test.py
+++ b/official/nlp/modeling/layers/kernel_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,12 +16,12 @@
import itertools
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import kernel_attention as attention
-_FEATURE_TRANSFORM = ['relu', 'elu', 'exp']
+_FEATURE_TRANSFORM = ["relu", "elu", "exp", "expplus"]
_REDRAW = [True, False]
_TRAINING = [True, False]
_IS_SHORT_SEQ = [True, False]
@@ -30,6 +30,64 @@
class KernelAttentionTest(tf.test.TestCase, parameterized.TestCase):
+ # expplus is only designed for bi-directional use case.
+ # exp can be numeric unstable.
+ @parameterized.parameters(itertools.product(
+ ["relu", "elu"], [1, 4], [0.9]))
+ def test_causal_windowed_attention_projection_streaming(
+ self, feature_transform, causal_chunk_length, causal_weight_decay):
+ num_heads = 12
+ key_dim = 64
+ seq_length = 16
+ num_chunks = seq_length // causal_chunk_length
+ causal_window_length = num_chunks
+ batch_size = 2
+ training = False
+ num_random_features = 0
+ test_layer = attention.KernelAttention(
+ num_heads=num_heads,
+ key_dim=key_dim,
+ feature_transform=feature_transform,
+ num_random_features=num_random_features,
+ redraw=False,
+ is_short_seq=False,
+ begin_kernel=False,
+ use_causal_windowed=True,
+ causal_chunk_length=causal_chunk_length,
+ causal_window_length=causal_window_length,
+ causal_window_decay=causal_weight_decay,
+ causal_padding=None,
+ )
+ query = tf.random.normal(
+ shape=(batch_size, seq_length, key_dim), seed=2)
+ value = query
+ encoder_inputs_mask = tf.ones((batch_size, seq_length), dtype=tf.int32)
+ masks = tf.cast(encoder_inputs_mask, dtype=tf.float32)
+ output = test_layer(
+ query=query,
+ value=value,
+ attention_mask=masks,
+ training=training)
+ dim = num_random_features if num_random_features > 0 else key_dim
+ kv_cache = tf.zeros(
+ (batch_size, num_heads, dim, dim))
+ k_sum_cache = tf.zeros((batch_size, num_heads, dim))
+ stream_output = []
+ cache = {"kv": kv_cache, "k_sum": k_sum_cache}
+ for i in range(num_chunks):
+ stream_output.append(
+ test_layer(
+ query=query[:, i * causal_chunk_length:(i + 1) *
+ causal_chunk_length, :],
+ value=value[:, i * causal_chunk_length:(i + 1) *
+ causal_chunk_length, :],
+ attention_mask=masks[:, i * causal_chunk_length:(i + 1) *
+ causal_chunk_length],
+ cache=cache,
+ training=training))
+ stream_output = tf.concat(stream_output, axis=1)
+ self.assertAllClose(output, stream_output)
+
@parameterized.parameters(
itertools.product(_FEATURE_TRANSFORM, [127], _TRAINING, [True, False],
_IS_SHORT_SEQ, _BEGIN_KERNEL))
@@ -60,6 +118,41 @@ def test_attention_projection(
training=training)
self.assertEqual(output.shape, [batch_size, seq_length, key_dim])
+ @parameterized.parameters(
+ itertools.product(["relu", "exp"], [127], _TRAINING, [True, False],
+ [0], [None, 0.97], [None, "left", "right"]))
+ def test_causal_windowed_attention_projection(
+ self, feature_transform, num_random_features, training, redraw,
+ begin_kernel, causal_window_decay, causal_padding):
+ num_heads = 12
+ key_dim = 64
+ seq_length = 1024
+ batch_size = 2
+ test_layer = attention.KernelAttention(
+ num_heads=num_heads,
+ key_dim=key_dim,
+ feature_transform=feature_transform,
+ num_random_features=num_random_features,
+ redraw=redraw,
+ is_short_seq=False,
+ begin_kernel=begin_kernel,
+ use_causal_windowed=True,
+ causal_chunk_length=8,
+ causal_window_length=3,
+ causal_window_decay=causal_window_decay,
+ causal_padding=causal_padding)
+ query = tf.random.normal(
+ shape=(batch_size, seq_length, key_dim))
+ value = query
+ encoder_inputs_mask = tf.zeros((batch_size, seq_length), dtype=tf.int32)
+ masks = tf.cast(encoder_inputs_mask, dtype=tf.float32)
+ output = test_layer(
+ query=query,
+ value=value,
+ attention_mask=masks,
+ training=training)
+ self.assertEqual(output.shape, [batch_size, seq_length, key_dim])
+
@parameterized.parameters(itertools.product(
_FEATURE_TRANSFORM, [0], _TRAINING, [False],
_IS_SHORT_SEQ, _BEGIN_KERNEL))
@@ -117,14 +210,14 @@ def test_attention_scale_by_length(self, seq_length):
self.assertNotAllClose(output_scale_by_length, output_no_scale_by_length)
def test_unsupported_feature_transform(self):
- with self.assertRaisesRegex(ValueError, 'Unsupported feature_transform.*'):
- _ = attention.KernelAttention(feature_transform='test')
+ with self.assertRaisesRegex(ValueError, "Unsupported feature_transform.*"):
+ _ = attention.KernelAttention(feature_transform="test")
def test_redraw_true_no_projection(self):
with self.assertRaisesRegex(
- ValueError, 'There is nothing to redraw when num_random_features.*'):
+ ValueError, "There is nothing to redraw when num_random_features.*"):
_ = attention.KernelAttention(
- num_heads=2, key_dim=64, feature_transform='elu',
+ num_heads=2, key_dim=64, feature_transform="elu",
num_random_features=0, redraw=True)
def test_config(self):
@@ -133,7 +226,7 @@ def test_config(self):
test_layer = attention.KernelAttention(
num_heads=num_heads,
key_dim=key_dim,
- feature_transform='exp',
+ feature_transform="exp",
num_random_features=128,
is_short_seq=True)
new_layer = attention.KernelAttention.from_config(
@@ -141,5 +234,25 @@ def test_config(self):
# If the serialization was successful, the new config should match the old.
self.assertAllEqual(test_layer.get_config(), new_layer.get_config())
-if __name__ == '__main__':
+ def test_rectangular_window_sum(self):
+ x = tf.ones([2, 5, 2, 2, 2])
+ winsum = attention.rectangular_window_sum(x, 3)
+ self.assertEqual(winsum.shape, x.shape)
+ self.assertAllClose(
+ tf.tile(
+ tf.reshape([1., 2., 3., 3., 3.], [1, -1, 1, 1, 1]),
+ [2, 1, 2, 2, 2]),
+ winsum)
+
+ def test_weighted_window_sum(self):
+ x = tf.ones([2, 5, 2, 2, 2])
+ winsum = attention.weighted_window_sum(x, 3, [0.01, 0.1, 1.])
+ self.assertEqual(winsum.shape, x.shape)
+ self.assertAllClose(
+ tf.tile(
+ tf.reshape([1., 1.1, 1.11, 1.11, 1.11], [1, -1, 1, 1, 1]),
+ [2, 1, 2, 2, 2]),
+ winsum)
+
+if __name__ == "__main__":
tf.test.main()
diff --git a/official/nlp/modeling/layers/masked_lm.py b/official/nlp/modeling/layers/masked_lm.py
index c622d91b752..7cd6d881288 100644
--- a/official/nlp/modeling/layers/masked_lm.py
+++ b/official/nlp/modeling/layers/masked_lm.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,11 @@
"""Masked language model network."""
# pylint: disable=g-classes-have-attributes
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MaskedLM(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MaskedLM(tf_keras.layers.Layer):
"""Masked language model network head for BERT modeling.
This layer implements a masked language model based on the provided
@@ -47,10 +47,10 @@ def __init__(self,
output='logits',
name=None,
**kwargs):
- super(MaskedLM, self).__init__(name=name, **kwargs)
+ super().__init__(name=name, **kwargs)
self.embedding_table = embedding_table
self.activation = activation
- self.initializer = tf.keras.initializers.get(initializer)
+ self.initializer = tf_keras.initializers.get(initializer)
if output not in ('predictions', 'logits'):
raise ValueError(
@@ -60,12 +60,12 @@ def __init__(self,
def build(self, input_shape):
self._vocab_size, hidden_size = self.embedding_table.shape
- self.dense = tf.keras.layers.Dense(
+ self.dense = tf_keras.layers.Dense(
hidden_size,
activation=self.activation,
kernel_initializer=self.initializer,
name='transform/dense')
- self.layer_norm = tf.keras.layers.LayerNormalization(
+ self.layer_norm = tf_keras.layers.LayerNormalization(
axis=-1, epsilon=1e-12, name='transform/LayerNorm')
self.bias = self.add_weight(
'output_bias/bias',
@@ -73,7 +73,7 @@ def build(self, input_shape):
initializer='zeros',
trainable=True)
- super(MaskedLM, self).build(input_shape)
+ super().build(input_shape)
def call(self, sequence_data, masked_positions):
masked_lm_input = self._gather_indexes(sequence_data, masked_positions)
@@ -81,10 +81,16 @@ def call(self, sequence_data, masked_positions):
lm_data = self.layer_norm(lm_data)
lm_data = tf.matmul(lm_data, self.embedding_table, transpose_b=True)
logits = tf.nn.bias_add(lm_data, self.bias)
- masked_positions_length = masked_positions.shape.as_list()[1] or tf.shape(
- masked_positions)[1]
- logits = tf.reshape(logits,
- [-1, masked_positions_length, self._vocab_size])
+ masked_positions_length = (
+ masked_positions.shape.as_list()[1] or tf.shape(masked_positions)[1]
+ )
+ batch_size = (
+ masked_positions.shape.as_list()[0] or tf.shape(masked_positions)[0]
+ )
+ logits = tf.reshape(
+ logits,
+ [batch_size, masked_positions_length, self._vocab_size],
+ )
if self._output_type == 'logits':
return logits
return tf.nn.log_softmax(logits)
diff --git a/official/nlp/modeling/layers/masked_lm_test.py b/official/nlp/modeling/layers/masked_lm_test.py
index 0cd3ce0721a..47ceb04a604 100644
--- a/official/nlp/modeling/layers/masked_lm_test.py
+++ b/official/nlp/modeling/layers/masked_lm_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,20 +13,15 @@
# limitations under the License.
"""Tests for masked language model network."""
-
+from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
-
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import masked_lm
from official.nlp.modeling.networks import bert_encoder
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class MaskedLMTest(keras_parameterized.TestCase):
+class MaskedLMTest(tf.test.TestCase, parameterized.TestCase):
def create_layer(self,
vocab_size,
@@ -57,8 +52,8 @@ def test_layer_creation(self):
vocab_size=vocab_size, hidden_size=hidden_size)
# Make sure that the output tensor of the masked LM is the right shape.
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_positions = tf.keras.Input(shape=(num_predictions,), dtype=tf.int32)
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_positions = tf_keras.Input(shape=(num_predictions,), dtype=tf.int32)
output = test_layer(lm_input_tensor, masked_positions=masked_positions)
expected_output_shape = [None, num_predictions, vocab_size]
@@ -87,14 +82,14 @@ def test_layer_invocation_with_external_logits(self):
output='logits')
# Create a model from the masked LM layer.
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_positions = tf.keras.Input(shape=(num_predictions,), dtype=tf.int32)
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_positions = tf_keras.Input(shape=(num_predictions,), dtype=tf.int32)
output = test_layer(lm_input_tensor, masked_positions)
logit_output = logit_layer(lm_input_tensor, masked_positions)
- logit_output = tf.keras.layers.Activation(tf.nn.log_softmax)(logit_output)
+ logit_output = tf_keras.layers.Activation(tf.nn.log_softmax)(logit_output)
logit_layer.set_weights(test_layer.get_weights())
- model = tf.keras.Model([lm_input_tensor, masked_positions], output)
- logits_model = tf.keras.Model(([lm_input_tensor, masked_positions]),
+ model = tf_keras.Model([lm_input_tensor, masked_positions], output)
+ logits_model = tf_keras.Model(([lm_input_tensor, masked_positions]),
logit_output)
# Invoke the masked LM on some fake data to make sure there are no runtime
@@ -115,19 +110,28 @@ def test_layer_invocation_with_external_logits(self):
self.assertEqual(expected_output_shape, outputs.shape)
self.assertAllClose(ref_outputs, outputs)
- def test_layer_invocation(self):
+ @parameterized.named_parameters(
+ dict(
+ testcase_name='default',
+ num_predictions=21,
+ ),
+ dict(
+ testcase_name='zero_predictions',
+ num_predictions=0,
+ ),
+ )
+ def test_layer_invocation(self, num_predictions):
vocab_size = 100
sequence_length = 32
hidden_size = 64
- num_predictions = 21
test_layer = self.create_layer(
vocab_size=vocab_size, hidden_size=hidden_size)
# Create a model from the masked LM layer.
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_positions = tf.keras.Input(shape=(num_predictions,), dtype=tf.int32)
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_positions = tf_keras.Input(shape=(num_predictions,), dtype=tf.int32)
output = test_layer(lm_input_tensor, masked_positions)
- model = tf.keras.Model([lm_input_tensor, masked_positions], output)
+ model = tf_keras.Model([lm_input_tensor, masked_positions], output)
# Invoke the masked LM on some fake data to make sure there are no runtime
# errors in the code.
@@ -136,7 +140,9 @@ def test_layer_invocation(self):
(batch_size, sequence_length, hidden_size))
masked_position_data = np.random.randint(
2, size=(batch_size, num_predictions))
- _ = model.predict([lm_input_data, masked_position_data])
+ res = model.predict([lm_input_data, masked_position_data])
+ expected_shape = (batch_size, num_predictions, vocab_size)
+ self.assertEqual(expected_shape, res.shape)
def test_unknown_output_type_fails(self):
with self.assertRaisesRegex(ValueError, 'Unknown `output` value "bad".*'):
diff --git a/official/nlp/modeling/layers/masked_softmax.py b/official/nlp/modeling/layers/masked_softmax.py
index db0a0fcaaf1..06e5125189e 100644
--- a/official/nlp/modeling/layers/masked_softmax.py
+++ b/official/nlp/modeling/layers/masked_softmax.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Keras-based softmax layer with optional masking."""
# pylint: disable=g-classes-have-attributes
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _large_compatible_negative(tensor_type):
@@ -35,8 +35,8 @@ def _large_compatible_negative(tensor_type):
return -1e9
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MaskedSoftmax(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MaskedSoftmax(tf_keras.layers.Layer):
"""Performs a softmax with optional masking on a tensor.
Args:
@@ -53,7 +53,7 @@ def __init__(self,
self._normalization_axes = (-1,)
else:
self._normalization_axes = normalization_axes
- super(MaskedSoftmax, self).__init__(**kwargs)
+ super().__init__(**kwargs)
def call(self, scores, mask=None):
@@ -81,5 +81,5 @@ def get_config(self):
'mask_expansion_axes': self._mask_expansion_axes,
'normalization_axes': self._normalization_axes
}
- base_config = super(MaskedSoftmax, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
diff --git a/official/nlp/modeling/layers/masked_softmax_test.py b/official/nlp/modeling/layers/masked_softmax_test.py
index d6fe410b164..51ff6ef9b87 100644
--- a/official/nlp/modeling/layers/masked_softmax_test.py
+++ b/official/nlp/modeling/layers/masked_softmax_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,22 +15,18 @@
"""Tests for Keras-based masked softmax layer."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import masked_softmax
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class MaskedSoftmaxLayerTest(keras_parameterized.TestCase):
+class MaskedSoftmaxLayerTest(tf.test.TestCase):
def test_non_masked_softmax(self):
test_layer = masked_softmax.MaskedSoftmax()
- input_tensor = tf.keras.Input(shape=(4, 8))
+ input_tensor = tf_keras.Input(shape=(4, 8))
output = test_layer(input_tensor)
- model = tf.keras.Model(input_tensor, output)
+ model = tf_keras.Model(input_tensor, output)
input_data = 10 * np.random.random_sample((3, 4, 8))
output_data = model.predict(input_data)
@@ -39,10 +35,10 @@ def test_non_masked_softmax(self):
def test_masked_softmax(self):
test_layer = masked_softmax.MaskedSoftmax()
- input_tensor = tf.keras.Input(shape=(4, 8))
- mask_tensor = tf.keras.Input(shape=(4, 8))
+ input_tensor = tf_keras.Input(shape=(4, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 8))
output = test_layer(input_tensor, mask_tensor)
- model = tf.keras.Model([input_tensor, mask_tensor], output)
+ model = tf_keras.Model([input_tensor, mask_tensor], output)
input_data = 10 * np.random.random_sample((3, 4, 8))
mask_data = np.random.randint(2, size=(3, 4, 8))
@@ -54,9 +50,9 @@ def test_masked_softmax(self):
def test_masked_softmax_with_none_mask(self):
test_layer = masked_softmax.MaskedSoftmax()
- input_tensor = tf.keras.Input(shape=(4, 8))
+ input_tensor = tf_keras.Input(shape=(4, 8))
output = test_layer(input_tensor, None)
- model = tf.keras.Model(input_tensor, output)
+ model = tf_keras.Model(input_tensor, output)
input_data = 10 * np.random.random_sample((3, 4, 8))
output_data = model.predict(input_data)
@@ -65,10 +61,10 @@ def test_masked_softmax_with_none_mask(self):
def test_softmax_with_axes_expansion(self):
test_layer = masked_softmax.MaskedSoftmax(mask_expansion_axes=[1])
- input_tensor = tf.keras.Input(shape=(4, 8))
- mask_tensor = tf.keras.Input(shape=(8))
+ input_tensor = tf_keras.Input(shape=(4, 8))
+ mask_tensor = tf_keras.Input(shape=(8))
output = test_layer(input_tensor, mask_tensor)
- model = tf.keras.Model([input_tensor, mask_tensor], output)
+ model = tf_keras.Model([input_tensor, mask_tensor], output)
input_data = 10 * np.random.random_sample((3, 4, 8))
mask_data = np.random.randint(2, size=(3, 8))
@@ -84,10 +80,10 @@ def test_masked_softmax_high_dims(self):
mask_expansion_axes=[1], normalization_axes=[6, 7])
input_shape = [2, 3, 4, 5, 6, 7, 8]
mask_shape = [5, 6, 7, 8]
- input_tensor = tf.keras.Input(shape=input_shape)
- mask_tensor = tf.keras.Input(shape=mask_shape)
+ input_tensor = tf_keras.Input(shape=input_shape)
+ mask_tensor = tf_keras.Input(shape=mask_shape)
output = test_layer(input_tensor, mask_tensor)
- model = tf.keras.Model([input_tensor, mask_tensor], output)
+ model = tf_keras.Model([input_tensor, mask_tensor], output)
input_data = 10 * np.random.random_sample([3] + input_shape)
mask_data = np.random.randint(2, size=[3] + mask_shape)
diff --git a/official/nlp/modeling/layers/mat_mul_with_margin.py b/official/nlp/modeling/layers/mat_mul_with_margin.py
index 9bc8721d20f..933998baae2 100644
--- a/official/nlp/modeling/layers/mat_mul_with_margin.py
+++ b/official/nlp/modeling/layers/mat_mul_with_margin.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,14 +16,13 @@
# pylint: disable=g-classes-have-attributes
from typing import Tuple
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MatMulWithMargin(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MatMulWithMargin(tf_keras.layers.Layer):
"""This layer computs a dot product matrix given two encoded inputs.
Args:
@@ -36,7 +35,7 @@ def __init__(self,
logit_scale=1.0,
logit_margin=0.0,
**kwargs):
- super(MatMulWithMargin, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self.logit_scale = logit_scale
self.logit_margin = logit_margin
@@ -61,7 +60,7 @@ def get_config(self):
config = {
'logit_scale': self.logit_scale,
'logit_margin': self.logit_margin}
- config.update(super(MatMulWithMargin, self).get_config())
+ config.update(super().get_config())
return config
@classmethod
diff --git a/official/nlp/modeling/layers/mat_mul_with_margin_test.py b/official/nlp/modeling/layers/mat_mul_with_margin_test.py
index 4a02d51362e..8da35d0e60e 100644
--- a/official/nlp/modeling/layers/mat_mul_with_margin_test.py
+++ b/official/nlp/modeling/layers/mat_mul_with_margin_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,21 +14,20 @@
"""Tests for mat_mul_with_margin layer."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import mat_mul_with_margin
-class MatMulWithMarginTest(keras_parameterized.TestCase):
+class MatMulWithMarginTest(tf.test.TestCase):
def test_layer_invocation(self):
"""Validate that the Keras object can be created and invoked."""
input_width = 512
test_layer = mat_mul_with_margin.MatMulWithMargin()
# Create a 2-dimensional input (the first dimension is implicit).
- left_encoded = tf.keras.Input(shape=(input_width,), dtype=tf.float32)
- right_encoded = tf.keras.Input(shape=(input_width,), dtype=tf.float32)
+ left_encoded = tf_keras.Input(shape=(input_width,), dtype=tf.float32)
+ right_encoded = tf_keras.Input(shape=(input_width,), dtype=tf.float32)
left_logits, right_logits = test_layer(left_encoded, right_encoded)
# Validate that the outputs are of the expected shape.
diff --git a/official/nlp/modeling/layers/mixing.py b/official/nlp/modeling/layers/mixing.py
new file mode 100644
index 00000000000..090ec3162ec
--- /dev/null
+++ b/official/nlp/modeling/layers/mixing.py
@@ -0,0 +1,285 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras-based mixing layers.
+
+Based on the mixing layers use by FNet
+(https://aclanthology.org/2022.naacl-main.319/) and Sparse Mixers
+(https://arxiv.org/abs/2205.12399).
+
+Mixing layers can be used as drop in replacements for self-attention layers. For
+interoperability with attention layers, we use the same `query` and `value` call
+signature.
+
+Note: These mixing layers currently only support encoder stacks. Decoder stacks
+can be supported in the future by utilizing the `value` inputs.
+"""
+
+import enum
+import functools
+from typing import Callable, Tuple, Union
+
+import gin
+import numpy as np
+from scipy import linalg
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+
+default_kernel_initializer = tf_keras.initializers.TruncatedNormal(stddev=2e-2)
+
+
+@gin.constants_from_enum
+class MixingMechanism(enum.Enum):
+ """Determines the type of mixing layer.
+
+ Possible options:
+ FOURIER: Fourier Transform mixing.
+ LINEAR: Mixing using dense matrix multiplications with learnable weights.
+ HARTLEY: Hartley Transform mixing.
+ """
+ FOURIER = "fourier"
+ HARTLEY = "hartley"
+ LINEAR = "linear"
+
+
+class MixingLayer(tf_keras.layers.Layer):
+ """Mixing layer base class.
+
+ This class cannot be used directly. It just specifies the API for mixing
+ layer subclasses. For interoperability with attention layers, we use the same
+ `query` and `value` call signature.
+
+ Based on the mixing layers use by FNet
+ (https://aclanthology.org/2022.naacl-main.319/) and Sparse Mixers
+ (https://arxiv.org/abs/2205.12399).
+ """
+
+ def __init__(self, name: str = "mixing", **kwargs):
+ """Initializes layer.
+
+ Args:
+ name: Name for layer.
+ **kwargs: Keyword arguments.
+ """
+ super().__init__(name=name, **kwargs)
+
+ def call(self, query: tf.Tensor, value: tf.Tensor, **kwargs) -> tf.Tensor:
+ """Calls the layer.
+
+ Subclasses should return tensors of shape
+ [batch_size, max_seq_length, hidden_dim].
+
+ Args:
+ query: Batch of input embeddings, typically of shape [batch_size,
+ max_seq_length, hidden_dim].
+ value: Unused. Included to match attention layer API.
+ **kwargs: Optional arguments to catch unused attention keyword arguments.
+
+ Raises:
+ NotImplementedError. This class should not be called directly.
+ """
+ raise NotImplementedError("Abstract method")
+
+
+class FourierTransformLayer(MixingLayer):
+ """Fourier Transform layer.
+
+ Applies 2D Fourier Transform over final two dimensions of `query` inputs -
+ typically the sequence and hidden dimensions.
+ """
+
+ def __init__(self,
+ use_fft: bool = False,
+ name: str = "fourier_transform",
+ **kwargs):
+ """Initializes layer.
+
+ Args:
+ use_fft: Whether to use Fast Fourier Transform (True) or the Discrete
+ Fourier Transform (DFT) matrix (False) to compute the Fourier Transform.
+ See _pick_fourier_transform() for recommendations on when to use FFT or
+ DFT.
+ name: Name for layer.
+ **kwargs: Keyword arguments.
+ """
+ super().__init__(name=name, **kwargs)
+ self.use_fft = use_fft
+
+ def build(self, input_shape: Tuple[int, ...]):
+ """Picks the Fourier Transform implementation."""
+ self.fourier_transform = _pick_fourier_transform(
+ self.use_fft,
+ max_seq_length=input_shape[-2],
+ hidden_dim=input_shape[-1])
+
+ def call(self, query: tf.Tensor, value: tf.Tensor, **kwargs) -> tf.Tensor:
+ """Applies layer to `query`.
+
+ Args:
+ query: Batch of input embeddings, typically of shape [batch_size,
+ max_seq_length, hidden_dim].
+ value: Unused. Included to match attention layer API.
+ **kwargs: Optional arguments to catch unused attention keyword arguments.
+
+ Returns:
+ Real part of discrete Fourier Transform of `query` inputs with shape
+ [batch_size, max_seq_length, hidden_dim].
+ """
+ del value # Ignored by encoder-only mixing layers
+ query = tf.cast(query, tf.complex64)
+ return tf.math.real(self.fourier_transform(query))
+
+
+class HartleyTransformLayer(MixingLayer):
+ """Hartley Transform layer.
+
+ Applies 2D Hartley Transform over final two dimensions of `query` inputs -
+ typically the sequence and hidden dimensions.
+ """
+
+ def __init__(self,
+ use_fft: bool = False,
+ name: str = "hartley_transform",
+ **kwargs):
+ """Initializes layer.
+
+ Args:
+ use_fft: Whether to use Fast Fourier Transform (True) or the Discrete
+ Fourier Transform (DFT) matrix (False) to compute the Hartley Transform.
+ See _pick_fourier_transform() for recommendations on when to use FFT or
+ DFT.
+ name: Name for layer.
+ **kwargs: Keyword arguments.
+ """
+ super().__init__(name=name, **kwargs)
+ self.use_fft = use_fft
+
+ def build(self, input_shape: Tuple[int, ...]):
+ """Picks the Fourier Transform implementation."""
+ self.fourier_transform = _pick_fourier_transform(
+ self.use_fft,
+ max_seq_length=input_shape[-2],
+ hidden_dim=input_shape[-1])
+
+ def call(self, query: tf.Tensor, value: tf.Tensor, **kwargs) -> tf.Tensor:
+ """Applies layer to `query`.
+
+ Args:
+ query: Batch of input embeddings, typically of shape [batch_size,
+ max_seq_length, hidden_dim].
+ value: Unused. Included to match attention layer API.
+ **kwargs: Optional arguments to catch unused attention keyword arguments.
+
+ Returns:
+ Real part of discrete Hartley Transform of `query` inputs with shape
+ [batch_size, max_seq_length, hidden_dim].
+ """
+ del value # Ignored by encoder-only mixing layers
+ query = tf.cast(query, tf.complex64)
+ frequencies = self.fourier_transform(query)
+ return tf.math.real(frequencies) - tf.math.imag(frequencies)
+
+
+class LinearTransformLayer(MixingLayer):
+ """Dense, linear transformation layer.
+
+ Applies matrix multiplications over sequence and hidden dimensions.
+ """
+
+ def __init__(self,
+ kernel_initializer: _Initializer = default_kernel_initializer,
+ name: str = "linear_transform",
+ **kwargs):
+ """Initializes layer.
+
+ Args:
+ kernel_initializer: Initialization scheme for kernel.
+ name: Name for layer.
+ **kwargs: Keyword arguments.
+ """
+ super().__init__(name=name, **kwargs)
+ self.kernel_initializer = kernel_initializer
+
+ def build(self, input_shape: Tuple[int, ...]):
+ """Creates the hidden and sequence matrix variables of the layer."""
+ self.mat_hidden = self.add_weight(
+ shape=(input_shape[-1], input_shape[-1]),
+ initializer=tf_utils.clone_initializer(self.kernel_initializer),
+ trainable=True,
+ name="hidden_kernel")
+ self.mat_seq = self.add_weight(
+ shape=(input_shape[-2], input_shape[-2]),
+ initializer=tf_utils.clone_initializer(self.kernel_initializer),
+ trainable=True,
+ name="seq_kernel")
+
+ def call(self, query: tf.Tensor, value: tf.Tensor, **kwargs) -> tf.Tensor:
+ """Applies layer to `query`.
+
+ Args:
+ query: Batch of input embeddings, typically of shape [batch_size,
+ max_seq_length, hidden_dim].
+ value: Unused. Included to match attention layer API.
+ **kwargs: Optional arguments to catch unused attention keyword arguments.
+
+ Returns:
+ Linearly transformed `query` inputs with shape
+ [batch_size, max_seq_length, hidden_dim].
+ """
+ del value # Ignored by encoder-only mixing layers
+
+ return tf.einsum("bij,jk,ni->bnk", query, self.mat_hidden, self.mat_seq)
+
+
+def _pick_fourier_transform(
+ use_fft: bool, max_seq_length: int,
+ hidden_dim: int) -> Callable[[tf.Tensor], tf.Tensor]:
+ """Returns FFT or DFT Fourier Transform implementation.
+
+ On TPUs, we recommend using the Discrete Fourier Transform (DFT) matrix
+ (use_fft=False), except for very long sequence lengths. On GPUs and CPUs, the
+ Fast Fourier Transform (use_fft=True) is generally optimal for all sequence
+ lengths.
+
+ Note: When using the FFT it is recommended to use a sequence length that is a
+ power of 2.
+
+ Args:
+ use_fft: If True, return FFT. Otherwise, return DFT matrix.
+ max_seq_length: Maximum sequence length of inputs. Only used if
+ use_fft=False.
+ hidden_dim: Size of hidden dimension of inputs. Only used if use_fft=False.
+
+ Returns:
+ Fourier Transform.
+ """
+ if use_fft:
+ return tf.signal.fft2d
+ else:
+ dft_mat_seq = linalg.dft(max_seq_length).astype(np.complex64)
+ dft_mat_hidden = linalg.dft(hidden_dim).astype(np.complex64)
+
+ def two_dim_matmul(x: tf.Tensor, matrix_dim_one: tf.Tensor,
+ matrix_dim_two: tf.Tensor) -> tf.Tensor:
+ """Applies 2D matrix multiplication to input tensors of rank >= 2."""
+ return tf.einsum("...ij,jk,ni->...nk", tf.cast(x, tf.complex64),
+ matrix_dim_two, matrix_dim_one)
+
+ return functools.partial(
+ two_dim_matmul,
+ matrix_dim_one=dft_mat_seq,
+ matrix_dim_two=dft_mat_hidden)
diff --git a/official/nlp/modeling/layers/mixing_test.py b/official/nlp/modeling/layers/mixing_test.py
new file mode 100644
index 00000000000..716d9006f44
--- /dev/null
+++ b/official/nlp/modeling/layers/mixing_test.py
@@ -0,0 +1,109 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for mixing.py."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling.layers import mixing
+
+
+class MixingTest(tf.test.TestCase):
+
+ def test_base_mixing_layer(self):
+ inputs = tf.random.uniform((3, 8, 16),
+ minval=0,
+ maxval=10,
+ dtype=tf.float32)
+
+ with self.assertRaisesRegex(NotImplementedError, "Abstract method"):
+ _ = mixing.MixingLayer()(query=inputs, value=inputs)
+
+ def test_fourier_layer(self):
+ batch_size = 4
+ max_seq_length = 8
+ hidden_dim = 16
+
+ inputs = tf.random.uniform((batch_size, max_seq_length, hidden_dim),
+ minval=0,
+ maxval=10,
+ dtype=tf.float32)
+ outputs = mixing.FourierTransformLayer(use_fft=True)(
+ query=inputs, value=inputs)
+ self.assertEqual(outputs.shape, (batch_size, max_seq_length, hidden_dim))
+
+ def test_hartley_layer(self):
+ batch_size = 3
+ max_seq_length = 16
+ hidden_dim = 4
+
+ inputs = tf.random.uniform((batch_size, max_seq_length, hidden_dim),
+ minval=0,
+ maxval=12,
+ dtype=tf.float32)
+ outputs = mixing.HartleyTransformLayer(use_fft=True)(
+ query=inputs, value=inputs)
+ self.assertEqual(outputs.shape, (batch_size, max_seq_length, hidden_dim))
+
+ def test_linear_mixing_layer(self):
+ batch_size = 2
+ max_seq_length = 4
+ hidden_dim = 3
+
+ inputs = tf.ones((batch_size, max_seq_length, hidden_dim), dtype=tf.float32)
+ outputs = mixing.LinearTransformLayer(
+ kernel_initializer=tf_keras.initializers.Ones())(
+ query=inputs, value=inputs)
+
+ # hidden_dim * (max_seq_length * 1) = 12.
+ expected_outputs = [
+ [
+ [12., 12., 12.],
+ [12., 12., 12.],
+ [12., 12., 12.],
+ [12., 12., 12.],
+ ],
+ [
+ [12., 12., 12.],
+ [12., 12., 12.],
+ [12., 12., 12.],
+ [12., 12., 12.],
+ ],
+ ]
+ np.testing.assert_allclose(outputs, expected_outputs, rtol=1e-6, atol=1e-6)
+
+ def test_pick_fourier_transform(self):
+ # Ensure we don't hit an edge case which exceeds the fixed numerical error.
+ tf.random.set_seed(1)
+ np.random.seed(1)
+
+ batch_size = 3
+ max_seq_length = 4
+ hidden_dim = 8
+
+ fft = mixing._pick_fourier_transform(
+ use_fft=True, max_seq_length=max_seq_length, hidden_dim=hidden_dim)
+ dft_matmul = mixing._pick_fourier_transform(
+ use_fft=False, max_seq_length=max_seq_length, hidden_dim=hidden_dim)
+
+ inputs = tf.random.uniform([batch_size, max_seq_length, hidden_dim])
+ inputs = tf.cast(inputs, tf.complex64)
+
+ np.testing.assert_allclose(
+ fft(inputs), dft_matmul(inputs), rtol=1e-6, atol=1e-6)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/nlp/modeling/layers/mobile_bert_layers.py b/official/nlp/modeling/layers/mobile_bert_layers.py
index ef66adb897c..ea66f95725f 100644
--- a/official/nlp/modeling/layers/mobile_bert_layers.py
+++ b/official/nlp/modeling/layers/mobile_bert_layers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,18 +13,19 @@
# limitations under the License.
"""MobileBERT embedding and transformer layers."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling.layers import on_device_embedding
from official.nlp.modeling.layers import position_embedding
-@tf.keras.utils.register_keras_serializable(package='Text')
-class NoNorm(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class NoNorm(tf_keras.layers.Layer):
"""Apply element-wise linear transformation to the last dimension."""
def __init__(self, name=None):
- super(NoNorm, self).__init__(name=name)
+ super().__init__(name=name)
def build(self, shape):
kernal_size = shape[-1]
@@ -54,7 +55,7 @@ def _get_norm_layer(normalization_type='no_norm', name=None):
if normalization_type == 'no_norm':
layer = NoNorm(name=name)
elif normalization_type == 'layer_norm':
- layer = tf.keras.layers.LayerNormalization(
+ layer = tf_keras.layers.LayerNormalization(
name=name,
axis=-1,
epsilon=1e-12,
@@ -64,8 +65,8 @@ def _get_norm_layer(normalization_type='no_norm', name=None):
return layer
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MobileBertEmbedding(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MobileBertEmbedding(tf_keras.layers.Layer):
"""Performs an embedding lookup for MobileBERT.
This layer includes word embedding, token type embedding, position embedding.
@@ -78,7 +79,7 @@ def __init__(self,
output_embed_size,
max_sequence_length=512,
normalization_type='no_norm',
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
dropout_rate=0.1,
**kwargs):
"""Class initialization.
@@ -96,38 +97,38 @@ def __init__(self,
dropout_rate: Dropout rate.
**kwargs: keyword arguments.
"""
- super(MobileBertEmbedding, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self.word_vocab_size = word_vocab_size
self.word_embed_size = word_embed_size
self.type_vocab_size = type_vocab_size
self.output_embed_size = output_embed_size
self.max_sequence_length = max_sequence_length
self.normalization_type = normalization_type
- self.initializer = tf.keras.initializers.get(initializer)
+ self.initializer = tf_keras.initializers.get(initializer)
self.dropout_rate = dropout_rate
self.word_embedding = on_device_embedding.OnDeviceEmbedding(
self.word_vocab_size,
self.word_embed_size,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(self.initializer),
name='word_embedding')
self.type_embedding = on_device_embedding.OnDeviceEmbedding(
self.type_vocab_size,
self.output_embed_size,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(self.initializer),
name='type_embedding')
self.pos_embedding = position_embedding.PositionEmbedding(
max_length=max_sequence_length,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(self.initializer),
name='position_embedding')
- self.word_embedding_proj = tf.keras.layers.experimental.EinsumDense(
+ self.word_embedding_proj = tf_keras.layers.EinsumDense(
'abc,cd->abd',
output_shape=[None, self.output_embed_size],
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
bias_axes='d',
name='embedding_projection')
self.layer_norm = _get_norm_layer(normalization_type, 'embedding_norm')
- self.dropout_layer = tf.keras.layers.Dropout(
+ self.dropout_layer = tf_keras.layers.Dropout(
self.dropout_rate,
name='embedding_dropout')
@@ -139,7 +140,7 @@ def get_config(self):
'output_embed_size': self.output_embed_size,
'max_sequence_length': self.max_sequence_length,
'normalization_type': self.normalization_type,
- 'initializer': tf.keras.initializers.serialize(self.initializer),
+ 'initializer': tf_keras.initializers.serialize(self.initializer),
'dropout_rate': self.dropout_rate
}
base_config = super(MobileBertEmbedding, self).get_config()
@@ -165,8 +166,8 @@ def call(self, input_ids, token_type_ids=None):
return embedding_out
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MobileBertTransformer(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MobileBertTransformer(tf_keras.layers.Layer):
"""Transformer block for MobileBERT.
An implementation of one layer (block) of Transformer with bottleneck and
@@ -188,7 +189,7 @@ def __init__(self,
key_query_shared_bottleneck=True,
num_feedforward_networks=4,
normalization_type='no_norm',
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
**kwargs):
"""Class initialization.
@@ -220,7 +221,7 @@ def __init__(self,
Raises:
ValueError: A Tensor shape or parameter is invalid.
"""
- super(MobileBertTransformer, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
self.intermediate_size = intermediate_size
@@ -232,7 +233,7 @@ def __init__(self,
self.key_query_shared_bottleneck = key_query_shared_bottleneck
self.num_feedforward_networks = num_feedforward_networks
self.normalization_type = normalization_type
- self.initializer = tf.keras.initializers.get(initializer)
+ self.initializer = tf_keras.initializers.get(initializer)
if intra_bottleneck_size % num_attention_heads != 0:
raise ValueError(
@@ -242,11 +243,11 @@ def __init__(self,
self.block_layers = {}
# add input bottleneck
- dense_layer_2d = tf.keras.layers.experimental.EinsumDense(
+ dense_layer_2d = tf_keras.layers.EinsumDense(
'abc,cd->abd',
output_shape=[None, self.intra_bottleneck_size],
bias_axes='d',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='bottleneck_input/dense')
layer_norm = _get_norm_layer(self.normalization_type,
name='bottleneck_input/norm')
@@ -254,11 +255,11 @@ def __init__(self,
layer_norm]
if self.key_query_shared_bottleneck:
- dense_layer_2d = tf.keras.layers.experimental.EinsumDense(
+ dense_layer_2d = tf_keras.layers.EinsumDense(
'abc,cd->abd',
output_shape=[None, self.intra_bottleneck_size],
bias_axes='d',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='kq_shared_bottleneck/dense')
layer_norm = _get_norm_layer(self.normalization_type,
name='kq_shared_bottleneck/norm')
@@ -266,13 +267,13 @@ def __init__(self,
layer_norm]
# add attention layer
- attention_layer = tf.keras.layers.MultiHeadAttention(
+ attention_layer = tf_keras.layers.MultiHeadAttention(
num_heads=self.num_attention_heads,
key_dim=attention_head_size,
value_dim=attention_head_size,
dropout=self.attention_probs_dropout_prob,
output_shape=self.intra_bottleneck_size,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='attention')
layer_norm = _get_norm_layer(self.normalization_type,
name='attention/norm')
@@ -284,19 +285,20 @@ def __init__(self,
for ffn_layer_idx in range(self.num_feedforward_networks):
layer_prefix = f'ffn_layer_{ffn_layer_idx}'
layer_name = layer_prefix + '/intermediate_dense'
- intermediate_layer = tf.keras.layers.experimental.EinsumDense(
+ intermediate_layer = tf_keras.layers.EinsumDense(
'abc,cd->abd',
- activation=self.intermediate_act_fn,
+ activation=tf_utils.get_activation(self.intermediate_act_fn),
output_shape=[None, self.intermediate_size],
bias_axes='d',
- kernel_initializer=initializer,
- name=layer_name)
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
+ name=layer_name,
+ )
layer_name = layer_prefix + '/output_dense'
- output_layer = tf.keras.layers.experimental.EinsumDense(
+ output_layer = tf_keras.layers.EinsumDense(
'abc,cd->abd',
output_shape=[None, self.intra_bottleneck_size],
bias_axes='d',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name=layer_name)
layer_name = layer_prefix + '/norm'
layer_norm = _get_norm_layer(self.normalization_type,
@@ -306,14 +308,14 @@ def __init__(self,
layer_norm])
# add output bottleneck
- bottleneck = tf.keras.layers.experimental.EinsumDense(
+ bottleneck = tf_keras.layers.EinsumDense(
'abc,cd->abd',
output_shape=[None, self.hidden_size],
activation=None,
bias_axes='d',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='bottleneck_output/dense')
- dropout_layer = tf.keras.layers.Dropout(
+ dropout_layer = tf_keras.layers.Dropout(
self.hidden_dropout_prob,
name='bottleneck_output/dropout')
layer_norm = _get_norm_layer(self.normalization_type,
@@ -335,7 +337,7 @@ def get_config(self):
'key_query_shared_bottleneck': self.key_query_shared_bottleneck,
'num_feedforward_networks': self.num_feedforward_networks,
'normalization_type': self.normalization_type,
- 'initializer': tf.keras.initializers.serialize(self.initializer),
+ 'initializer': tf_keras.initializers.serialize(self.initializer),
}
base_config = super(MobileBertTransformer, self).get_config()
return dict(list(base_config.items()) + list(config.items()))
@@ -419,7 +421,7 @@ def call(self,
bottleneck = self.block_layers['bottleneck_output'][0]
dropout_layer = self.block_layers['bottleneck_output'][1]
layer_norm = self.block_layers['bottleneck_output'][2]
- layer_output = bottleneck(layer_output)
+ layer_output = bottleneck(layer_output) # pyrefly: ignore[unbound-name]
layer_output = dropout_layer(layer_output)
layer_output = layer_norm(layer_output + prev_output)
@@ -429,8 +431,8 @@ def call(self,
return layer_output
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MobileBertMaskedLM(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MobileBertMaskedLM(tf_keras.layers.Layer):
"""Masked language model network head for BERT modeling.
This layer implements a masked language model based on the provided
@@ -445,6 +447,7 @@ def __init__(self,
activation=None,
initializer='glorot_uniform',
output='logits',
+ output_weights_use_proj=False,
**kwargs):
"""Class initialization.
@@ -455,34 +458,45 @@ def __init__(self,
uniform initializer.
output: The output style for this layer. Can be either `logits` or
`predictions`.
+ output_weights_use_proj: Use projection instead of concating extra output
+ weights, this may reduce the MLM task accuracy but will reduce the model
+ params as well.
**kwargs: keyword arguments.
"""
- super(MobileBertMaskedLM, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self.embedding_table = embedding_table
self.activation = activation
- self.initializer = tf.keras.initializers.get(initializer)
+ self.initializer = tf_keras.initializers.get(initializer)
if output not in ('predictions', 'logits'):
raise ValueError(
('Unknown `output` value "%s". `output` can be either "logits" or '
'"predictions"') % output)
self._output_type = output
+ self._output_weights_use_proj = output_weights_use_proj
def build(self, input_shape):
self._vocab_size, embedding_width = self.embedding_table.shape
hidden_size = input_shape[-1]
- self.dense = tf.keras.layers.Dense(
+ self.dense = tf_keras.layers.Dense(
hidden_size,
activation=self.activation,
- kernel_initializer=self.initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='transform/dense')
if hidden_size > embedding_width:
- self.extra_output_weights = self.add_weight(
- 'extra_output_weights',
- shape=(self._vocab_size, hidden_size - embedding_width),
- initializer=self.initializer,
- trainable=True)
+ if self._output_weights_use_proj:
+ self.extra_output_weights = self.add_weight(
+ 'output_weights_proj',
+ shape=(embedding_width, hidden_size),
+ initializer=tf_utils.clone_initializer(self.initializer),
+ trainable=True)
+ else:
+ self.extra_output_weights = self.add_weight(
+ 'extra_output_weights',
+ shape=(self._vocab_size, hidden_size - embedding_width),
+ initializer=tf_utils.clone_initializer(self.initializer),
+ trainable=True)
elif hidden_size == embedding_width:
self.extra_output_weights = None
else:
@@ -490,7 +504,7 @@ def build(self, input_shape):
'hidden size %d cannot be smaller than embedding width %d.' %
(hidden_size, embedding_width))
- self.layer_norm = tf.keras.layers.LayerNormalization(
+ self.layer_norm = tf_keras.layers.LayerNormalization(
axis=-1, epsilon=1e-12, name='transform/LayerNorm')
self.bias = self.add_weight(
'output_bias/bias',
@@ -507,10 +521,16 @@ def call(self, sequence_data, masked_positions):
if self.extra_output_weights is None:
lm_data = tf.matmul(lm_data, self.embedding_table, transpose_b=True)
else:
- lm_data = tf.matmul(
- lm_data,
- tf.concat([self.embedding_table, self.extra_output_weights], axis=1),
- transpose_b=True)
+ if self._output_weights_use_proj:
+ lm_data = tf.matmul(
+ lm_data, self.extra_output_weights, transpose_b=True)
+ lm_data = tf.matmul(lm_data, self.embedding_table, transpose_b=True)
+ else:
+ lm_data = tf.matmul(
+ lm_data,
+ tf.concat([self.embedding_table, self.extra_output_weights],
+ axis=1),
+ transpose_b=True)
logits = tf.nn.bias_add(lm_data, self.bias)
masked_positions_length = masked_positions.shape.as_list()[1] or tf.shape(
diff --git a/official/nlp/modeling/layers/mobile_bert_layers_test.py b/official/nlp/modeling/layers/mobile_bert_layers_test.py
index b5c3c5e3fd3..b7fb67a4c96 100644
--- a/official/nlp/modeling/layers/mobile_bert_layers_test.py
+++ b/official/nlp/modeling/layers/mobile_bert_layers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import mobile_bert_layers
from official.nlp.modeling.networks import mobile_bert_encoder
@@ -42,7 +42,7 @@ def test_embedding_layer_with_token_type(self):
output = layer(input_seq, token_type)
output_shape = output.shape.as_list()
expected_shape = [1, 4, 16]
- self.assertListEqual(output_shape, expected_shape, msg=None)
+ self.assertListEqual(output_shape, expected_shape)
def test_embedding_layer_without_token_type(self):
layer = mobile_bert_layers.MobileBertEmbedding(10, 8, 2, 16)
@@ -50,7 +50,7 @@ def test_embedding_layer_without_token_type(self):
output = layer(input_seq)
output_shape = output.shape.as_list()
expected_shape = [1, 4, 16]
- self.assertListEqual(output_shape, expected_shape, msg=None)
+ self.assertListEqual(output_shape, expected_shape)
def test_embedding_layer_get_config(self):
layer = mobile_bert_layers.MobileBertEmbedding(
@@ -60,7 +60,7 @@ def test_embedding_layer_get_config(self):
output_embed_size=32,
max_sequence_length=32,
normalization_type='layer_norm',
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.01),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01),
dropout_rate=0.5)
layer_config = layer.get_config()
new_layer = mobile_bert_layers.MobileBertEmbedding.from_config(layer_config)
@@ -72,7 +72,7 @@ def test_no_norm(self):
output = layer(feature)
output_shape = output.shape.as_list()
expected_shape = [2, 3, 4]
- self.assertListEqual(output_shape, expected_shape, msg=None)
+ self.assertListEqual(output_shape, expected_shape)
@parameterized.named_parameters(('with_kq_shared_bottleneck', False),
('without_kq_shared_bottleneck', True))
@@ -83,7 +83,17 @@ def test_transfomer_kq_shared_bottleneck(self, is_kq_shared):
output = layer(feature)
output_shape = output.shape.as_list()
expected_shape = [2, 3, 512]
- self.assertListEqual(output_shape, expected_shape, msg=None)
+ self.assertListEqual(output_shape, expected_shape)
+
+ def test_transformer_with_squared_relu(self):
+ feature = tf.random.uniform([2, 3, 512])
+ layer = mobile_bert_layers.MobileBertTransformer(
+ intermediate_act_fn='squared_relu'
+ )
+ output = layer(feature)
+ output_shape = output.shape.as_list()
+ expected_shape = [2, 3, 512]
+ self.assertListEqual(output_shape, expected_shape)
def test_transfomer_with_mask(self):
feature = tf.random.uniform([2, 3, 512])
@@ -94,7 +104,7 @@ def test_transfomer_with_mask(self):
output = layer(feature, input_mask)
output_shape = output.shape.as_list()
expected_shape = [2, 3, 512]
- self.assertListEqual(output_shape, expected_shape, msg=None)
+ self.assertListEqual(output_shape, expected_shape)
def test_transfomer_return_attention_score(self):
sequence_length = 5
@@ -104,8 +114,7 @@ def test_transfomer_return_attention_score(self):
num_attention_heads=num_attention_heads)
_, attention_score = layer(feature, return_attention_scores=True)
expected_shape = [2, num_attention_heads, sequence_length, sequence_length]
- self.assertListEqual(
- attention_score.shape.as_list(), expected_shape, msg=None)
+ self.assertListEqual(attention_score.shape.as_list(), expected_shape)
def test_transformer_get_config(self):
layer = mobile_bert_layers.MobileBertTransformer(
@@ -120,7 +129,7 @@ def test_transformer_get_config(self):
key_query_shared_bottleneck=False,
num_feedforward_networks=2,
normalization_type='layer_norm',
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.01),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01),
name='block')
layer_config = layer.get_config()
new_layer = mobile_bert_layers.MobileBertTransformer.from_config(
@@ -163,8 +172,8 @@ def test_layer_creation(self):
embedding_width=embedding_width)
# Make sure that the output tensor of the masked LM is the right shape.
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_positions = tf.keras.Input(shape=(num_predictions,), dtype=tf.int32)
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_positions = tf_keras.Input(shape=(num_predictions,), dtype=tf.int32)
output = test_layer(lm_input_tensor, masked_positions=masked_positions)
expected_output_shape = [None, num_predictions, vocab_size]
@@ -196,14 +205,14 @@ def test_layer_invocation_with_external_logits(self):
output='logits')
# Create a model from the masked LM layer.
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_positions = tf.keras.Input(shape=(num_predictions,), dtype=tf.int32)
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_positions = tf_keras.Input(shape=(num_predictions,), dtype=tf.int32)
output = test_layer(lm_input_tensor, masked_positions)
logit_output = logit_layer(lm_input_tensor, masked_positions)
- logit_output = tf.keras.layers.Activation(tf.nn.log_softmax)(logit_output)
+ logit_output = tf_keras.layers.Activation(tf.nn.log_softmax)(logit_output)
logit_layer.set_weights(test_layer.get_weights())
- model = tf.keras.Model([lm_input_tensor, masked_positions], output)
- logits_model = tf.keras.Model(([lm_input_tensor, masked_positions]),
+ model = tf_keras.Model([lm_input_tensor, masked_positions], output)
+ logits_model = tf_keras.Model(([lm_input_tensor, masked_positions]),
logit_output)
# Invoke the masked LM on some fake data to make sure there are no runtime
@@ -236,10 +245,10 @@ def test_layer_invocation(self):
embedding_width=embedding_width)
# Create a model from the masked LM layer.
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_positions = tf.keras.Input(shape=(num_predictions,), dtype=tf.int32)
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_positions = tf_keras.Input(shape=(num_predictions,), dtype=tf.int32)
output = test_layer(lm_input_tensor, masked_positions)
- model = tf.keras.Model([lm_input_tensor, masked_positions], output)
+ model = tf_keras.Model([lm_input_tensor, masked_positions], output)
# Invoke the masked LM on some fake data to make sure there are no runtime
# errors in the code.
@@ -263,8 +272,8 @@ def test_hidden_size_smaller_than_embedding_width(self):
ValueError, 'hidden size 8 cannot be smaller than embedding width 16.'):
test_layer = self.create_layer(
vocab_size=8, hidden_size=8, embedding_width=16)
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_positions = tf.keras.Input(
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_positions = tf_keras.Input(
shape=(num_predictions,), dtype=tf.int32)
_ = test_layer(lm_input_tensor, masked_positions)
diff --git a/official/nlp/modeling/layers/moe.py b/official/nlp/modeling/layers/moe.py
new file mode 100644
index 00000000000..8ac49bba836
--- /dev/null
+++ b/official/nlp/modeling/layers/moe.py
@@ -0,0 +1,721 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Mixture of Experts layers and their routing mechanisms."""
+
+import dataclasses
+from typing import Callable, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+
+
+_InitializerType = tf_keras.initializers.Initializer
+
+
+_DEFAULT_KERNEL_INITIALIZER = tf_keras.initializers.TruncatedNormal(stddev=2e-2)
+_DEFAULT_BIAS_INITIALIZER = tf_keras.initializers.Zeros()
+
+
+################## Routers (gating functions) ##################
+
+
+def _router_z_loss(router_logits: tf.Tensor) -> float:
+ """Computes router z-loss.
+
+ The router z-loss was introduced in Designing Effective Sparse Expert Models
+ (https://arxiv.org/abs/2202.08906). It encourages router logits to remain
+ small in an effort to improve stability.
+
+ Args:
+ router_logits: [num_groups, tokens_per_group, num_experts] router
+ logits.
+
+ Returns:
+ Scalar router z-loss .
+ """
+ num_groups = tf.shape(router_logits)[0]
+ tokens_per_group = router_logits.shape[1]
+
+ log_z = tf.math.reduce_logsumexp(router_logits, axis=-1)
+ z_loss = log_z**2
+ return tf.math.reduce_sum(z_loss) / tf.cast(
+ num_groups * tokens_per_group, tf.float32)
+
+
+@dataclasses.dataclass
+class RouterMask:
+ """Dispatch and combine arrays for expert routing with masked matmuls.
+
+ Attributes:
+ dispatch_mask:
+ [num_groups, tokens_per_group, num_experts, expert_capacity]
+ dispatch array that is 1 if the token gets routed to the
+ corresponding expert, and 0 otherwise.
+ combine_array:
+ [num_groups, tokens_per_group, num_experts, expert_capacity]
+ combine array used for combining expert outputs and
+ scaling with router probability.
+ """
+ dispatch_mask: tf.Tensor
+ combine_array: tf.Tensor
+
+RouterOutput = RouterMask
+
+
+class Router(tf_keras.layers.Layer):
+ """Abstract base router class, defining router API and inner workings.
+
+ Computations are performed in float32 for stability, and returned after
+ conversion according to the precision policy. See the discussion of
+ "selective precision" in https://arxiv.org/abs/2101.03961.
+
+ Uses Keras add_loss() and add_metric() APIs.
+
+ Attributes:
+ num_experts: Number of experts, used to check consistency with
+ FeedForwardExperts.
+ jitter_noise: Amplitude of jitter noise applied to router logits.
+ router_weights: Dense layer that computes logits for all tokens, which are
+ then used as expert or token weights.
+ """
+
+ def __init__(
+ self,
+ num_experts: int,
+ *,
+ jitter_noise: float = 0.0,
+ use_bias: bool = True,
+ kernel_initializer: _InitializerType = _DEFAULT_KERNEL_INITIALIZER,
+ bias_initializer: _InitializerType = _DEFAULT_BIAS_INITIALIZER,
+ router_z_loss_weight: float = 0.0,
+ export_metrics: bool = True,
+ name: str = "router",
+ **kwargs):
+ """Init.
+
+ Args:
+ num_experts: Number of experts.
+ jitter_noise: Amplitude of jitter noise applied to router logits.
+ use_bias: Whether or not to use the bias term in computing the router
+ weights.
+ kernel_initializer: Kernel initializer for router weights.
+ bias_initializer: Bias initializer for router weights.
+ router_z_loss_weight: Weight for router_z_loss. Use non-zero values if
+ running into training instability (esp. with dtype 'bfloat16' or lower).
+ export_metrics: Whether to export metrics using Keras add_metric API.
+ name: Layer name.
+ **kwargs: Forwarded to super.
+ """
+ super().__init__(name=name, **kwargs)
+
+ self.num_experts = num_experts # Used to check consistency with
+ # FeedForwardExperts.
+ self.jitter_noise = jitter_noise
+ self.router_z_loss_weight = router_z_loss_weight
+ self._export_metrics = export_metrics
+
+ self.router_weights = tf_keras.layers.Dense(
+ num_experts,
+ use_bias=use_bias,
+ kernel_initializer=tf_utils.clone_initializer(kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(bias_initializer),
+ name="router_weights",
+ dtype=tf.float32)
+
+ def call(self,
+ inputs: tf.Tensor,
+ *,
+ expert_capacity: int,
+ training: Optional[bool] = None) -> RouterOutput:
+ """Computes dispatch and combine arrays for routing to experts.
+
+ Args:
+ inputs: Inputs to send to experts of shape
+ [num_groups, tokens_per_group, hidden_dim].
+ expert_capacity: Each group will send this many tokens to each expert.
+ training: If true, apply jitter noise during routing. If not provided
+ taken from tf_keras.backend.
+
+ Returns:
+ Router indices or mask arrays (depending on router type).
+ """
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+
+ # inputs shape [num_groups, tokens_per_group, hidden_dim]
+ router_probs, router_logits = self._compute_router_probabilities(
+ inputs, apply_jitter=training)
+ # router_probs [num_groups, tokens_per_group, num_experts]
+ # router_logits [num_groups, tokens_per_group, num_experts]
+ unscaled_router_z_loss = _router_z_loss(router_logits)
+ router_z_loss = self.router_z_loss_weight * unscaled_router_z_loss
+ self.add_loss(router_z_loss)
+ if self._export_metrics:
+ self.add_metric(unscaled_router_z_loss, name="unscaled_router_z_loss")
+ self.add_metric(router_z_loss, name="router_z_loss")
+
+ routing_instructions = self._compute_routing_instructions(
+ router_probs, expert_capacity)
+ return routing_instructions
+
+ def _compute_router_probabilities(
+ self, inputs: tf.Tensor,
+ apply_jitter: bool) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Computes router probabilities from input tokens.
+
+ Args:
+ inputs: Inputs from which router probabilities are computed, shape
+ [num_groups, tokens_per_group, hidden_dim].
+ apply_jitter: If true, apply jitter noise.
+
+ Returns:
+ - [num_groups, tokens_per_group, num_experts] probabilities for
+ each token and expert. Used for routing tokens to experts.
+ - [num_groups, tokens_per_group, num_experts] raw router logits.
+ Used for computing router z-loss.
+ """
+ if apply_jitter and self.jitter_noise > 0:
+ inputs *= tf.random.uniform(
+ tf.shape(inputs),
+ minval=1.0 - self.jitter_noise,
+ maxval=1.0 + self.jitter_noise,
+ dtype=inputs.dtype)
+ # inputs , router_logits
+ router_logits = self.router_weights(inputs)
+ router_probs = tf_keras.activations.softmax(router_logits, axis=-1)
+ return router_probs, router_logits
+
+ def _compute_routing_instructions(self, router_probs: tf.Tensor,
+ expert_capacity: int) -> RouterOutput:
+ """Computes instructions for routing inputs to experts."""
+ raise NotImplementedError(
+ "Router is an abstract class that should be subclassed.")
+
+
+class MaskedRouter(Router):
+ """Abstract base router class for masked matmul dispatch routers.
+
+ MaskedRouter(s) return RouterMask(s) containing a dispatch mask and combine
+ array for sending and receiving (via masked matmuls) inputs and outputs to and
+ from experts.
+
+ Routing using masked matmuls is generally faster than scatter-based routing on
+ TPUs.
+
+ Uses Keras add_loss() and add_metric() APIs.
+ """
+
+ def _compute_routing_instructions(self, router_probs: tf.Tensor,
+ expert_capacity: int) -> RouterMask:
+ """Computes masks for the top-k experts per token.
+
+ Args:
+ router_probs: [num_groups, tokens_per_group, num_experts]
+ probabilities used to determine the routing of tokens to the experts.
+ expert_capacity: Each group will send this many tokens to each expert.
+
+ Returns:
+ Router mask arrays.
+ """
+ raise NotImplementedError(
+ "MaskedRouter is an abstract class that should be subclassed.")
+
+
+class ExpertsChooseMaskedRouter(MaskedRouter):
+ """Masked matmul router using experts choose tokens assignment.
+
+ This router uses the same mechanism as in Mixture-of-Experts with Expert
+ Choice (https://arxiv.org/abs/2202.09368): each expert selects its top
+ expert_capacity tokens. An individual token may be processed by multiple
+ experts or none at all.
+
+ Note: "experts choose routing" should not be used in decoder blocks because it
+ breaks the autoregressive behavior, leading to a mismatch between training
+ (teacher forcing) and inference (autoregressive decoding).
+
+ Uses Keras add_loss() and add_metric() APIs.
+ """
+
+ def _compute_routing_instructions(self, router_probs: tf.Tensor,
+ expert_capacity: int) -> RouterMask:
+ """Computes masks for the highest probability token per expert.
+
+ Args:
+ router_probs: [num_groups, tokens_per_group, num_experts]
+ probabilities used to determine the routing of tokens to the experts.
+ expert_capacity: Each group will send this many tokens to each expert.
+
+ Returns:
+ Dispatch and combine arrays for routing with masked matmuls.
+ """
+ num_groups = tf.shape(router_probs)[0]
+ tokens_per_group = router_probs.shape[1]
+
+ router_probs_t = tf.transpose(router_probs, perm=[0, 2, 1])
+ # router_probs_t: [num_groups, num_experts, tokens_per_group]
+ # Top expert_capacity router probability and corresponding token indices for
+ # each expert.
+ # Shapes [num_groups, num_experts, expert_capacity]
+ _, expert_index = tf.math.top_k(
+ router_probs_t, k=expert_capacity, sorted=False)
+
+ # Convert to one-hot mask of expert indices for each token in each group.
+ # Shape: [num_groups, tokens_per_group, num_experts, expert_capacity].
+ dispatch_mask = tf.one_hot(
+ expert_index, tokens_per_group, axis=1, dtype=router_probs.dtype)
+
+ # The combine array will be used for combining expert outputs, scaled by the
+ # router probabilities.
+ # Shape: [num_groups, num_experts, tokens_per_group, expert_capacity]
+ combine_array = tf.expand_dims(router_probs, axis=3) * dispatch_mask
+
+ # Add load balancing loss.
+ # Each expert is choosing tokens until it reaches full capacity, so we don't
+ # need an auxiliary loading balancing loss for expert choice routing.
+ if self._export_metrics:
+ self.add_metric(0.0, name="load_balancing_loss")
+
+ # Gather expert metrics.
+ # Number of tokens that were dispatched to at least one expert.
+ num_tokens = num_groups * tokens_per_group
+ num_tokens_dispatched_somewhere = tf.math.reduce_sum(tf.math.reduce_max(
+ dispatch_mask, axis=(-1, -2)))
+ fraction_tokens_left_behind = 1.0 - tf.cast(
+ num_tokens_dispatched_somewhere, tf.float32) / tf.cast(
+ num_tokens, tf.float32)
+
+ # Total number of tokens that were dispatched (one token could be
+ # dispatched to multiple experts).
+ num_tokens_dispatched = tf.math.reduce_sum(dispatch_mask)
+ # Of the tokens dispatched, how confident was the router in its routing?
+ router_confidence = tf.math.reduce_sum(
+ combine_array) / num_tokens_dispatched
+
+ expert_usage = 1.0 # Experts fully utilized when "expert choose tokens"
+
+ self.add_metric(fraction_tokens_left_behind,
+ name="fraction_tokens_left_behind")
+ self.add_metric(router_confidence, name="router_confidence")
+ self.add_metric(expert_usage, name="expert_usage")
+
+ # Return to default dtype now that router computation is complete.
+ dispatch_mask = tf.cast(dispatch_mask, self.compute_dtype)
+ combine_array = tf.cast(combine_array, self.compute_dtype)
+ output = RouterMask(dispatch_mask, combine_array)
+ return output
+
+
+################## Model layers ##################
+
+
+class FeedForward(tf_keras.layers.Layer):
+ """Feed-forward layer - position independent, dense, nonlinear transformation.
+
+ Typically used in an MLP Transformer block.
+ """
+
+ def __init__(
+ self,
+ d_ff: int,
+ *,
+ inner_dropout: float = 0.0,
+ output_dropout: float = 0.0,
+ activation: Callable[[tf.Tensor], tf.Tensor] = tf_keras.activations.gelu,
+ kernel_initializer: _InitializerType = _DEFAULT_KERNEL_INITIALIZER,
+ bias_initializer: _InitializerType = _DEFAULT_BIAS_INITIALIZER,
+ name: str = "feed_forward",
+ **kwargs):
+ """Initializes layer.
+
+ Args:
+ d_ff: Dimension of feed-forward layer.
+ inner_dropout: The dropout probability to be applied after intermediate
+ activations.
+ output_dropout: The dropout probability to be applied after output layer.
+ activation: (Nonlinear) transform applied in layer.
+ kernel_initializer: Initialization scheme for kernel.
+ bias_initializer: Initialization scheme for bias.
+ name: Layer name.
+ **kwargs: Forwarded to super.
+ """
+ super().__init__(name=name, **kwargs)
+ self.activation = activation
+ self.kernel_initializer = kernel_initializer
+ self.bias_initializer = bias_initializer
+
+ self.intermediate_layer = tf_keras.layers.Dense(
+ d_ff,
+ kernel_initializer=tf_utils.clone_initializer(self.kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self.bias_initializer),
+ name="intermediate")
+ self.inner_dropout_layer = tf_keras.layers.Dropout(
+ inner_dropout)
+ self.output_dropout_layer = tf_keras.layers.Dropout(output_dropout)
+
+ def build(self, input_shape: Tuple[int, int, int]):
+ """Creates the input shape dependent output weight variables."""
+ self.output_layer = tf_keras.layers.Dense(
+ input_shape[-1],
+ kernel_initializer=tf_utils.clone_initializer(self.kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self.bias_initializer),
+ name="output")
+
+ def call(self,
+ inputs: tf.Tensor,
+ *,
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Applies layer to inputs.
+
+ Args:
+ inputs: Batch of input embeddings, of shape
+ [batch_size, seq_len, hidden_dim].
+ training: Only apply dropout during training.
+
+ Returns:
+ Transformed inputs with the same shape as inputs
+ [batch_size, seq_len, hidden_dim].
+ """
+ x = self.intermediate_layer(inputs)
+ x = self.activation(x)
+ x = self.inner_dropout_layer(x, training=training)
+ x = self.output_layer(x)
+ x = self.output_dropout_layer(x, training=training)
+ return x
+
+
+class FeedForwardExperts(tf_keras.layers.Layer):
+ """Feed-forward layer with multiple experts.
+
+ Note that call() takes inputs with shape
+ [num_groups, num_experts, expert_capacity, hidden_dim]
+ which is different from the usual [batch_size, seq_len, hidden_dim] used by
+ the FeedForward layer.
+
+ The experts are independent FeedForward layers of the
+ same shape, i.e. the kernel doesn't have shape [hidden_dim, out_dim], but
+ [num_experts, hidden_dim, out_dim].
+ """
+
+ def __init__(
+ self,
+ num_experts: int,
+ d_ff: int,
+ *,
+ inner_dropout: float = 0.0,
+ output_dropout: float = 0.0,
+ activation: Callable[[tf.Tensor], tf.Tensor] = tf_keras.activations.gelu,
+ kernel_initializer: _InitializerType = _DEFAULT_KERNEL_INITIALIZER,
+ bias_initializer: _InitializerType = _DEFAULT_BIAS_INITIALIZER,
+ name: str = "experts",
+ **kwargs):
+ """Initializes layer.
+
+ Args:
+ num_experts: Number of experts (i.e. number of independent feed-forward
+ blocks).
+ d_ff: Dimension of feed-forward layer of each expert.
+ inner_dropout: The dropout probability to be applied after intermediate
+ activations.
+ output_dropout: The dropout probability to be applied after output layer.
+ activation: (Nonlinear) transform applied in layer.
+ kernel_initializer: Initialization scheme for kernel.
+ bias_initializer: Initialization scheme for bias.
+ name: Layer name.
+ **kwargs: Forwarded to super.
+ """
+ super().__init__(name=name, **kwargs)
+ self.num_experts = num_experts
+ self.activation = activation
+ self.kernel_initializer = kernel_initializer
+ self.bias_initializer = bias_initializer
+
+ self.intermediate_layer = tf_keras.layers.EinsumDense(
+ "gech,ehf->gecf",
+ output_shape=(self.num_experts, None, d_ff),
+ bias_axes="ef",
+ kernel_initializer=tf_utils.clone_initializer(self.kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self.bias_initializer),
+ name="intermediate")
+ self.inner_dropout_layer = tf_keras.layers.Dropout(
+ inner_dropout)
+ self.output_dropout_layer = tf_keras.layers.Dropout(output_dropout)
+
+ def build(self, input_shape: Tuple[int, int, int, int]):
+ """Creates the input shape dependent output weight variables."""
+ if input_shape[1] != self.num_experts:
+ raise ValueError(
+ f"Input shape {input_shape} is inconsistent with num_experts "
+ f"{self.num_experts}.")
+
+ self.output_layer = tf_keras.layers.EinsumDense(
+ "gecf,efh->gech",
+ output_shape=(self.num_experts, None, input_shape[-1]),
+ bias_axes="eh",
+ kernel_initializer=tf_utils.clone_initializer(self.kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self.bias_initializer),
+ name="output")
+
+ def call(self,
+ inputs: tf.Tensor,
+ *,
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Applies layer to inputs.
+
+ Args:
+ inputs: Inputs of shape
+ [num_groups, num_experts, expert_capacity, hidden_dim].
+ training: Only apply dropout during training.
+
+ Returns:
+ Transformed inputs with the same shape as inputs
+ [num_groups, num_experts, expert_capacity, hidden_dim].
+ """
+ x = self.intermediate_layer(inputs)
+ x = self.activation(x)
+ x = self.inner_dropout_layer(x, training=training)
+ x = self.output_layer(x)
+ x = self.output_dropout_layer(x, training=training)
+ return x
+
+
+class MoeLayer(tf_keras.layers.Layer):
+ """Sparse MoE layer with per-token routing.
+
+ In this TF implementation, all experts need to fit onto a single device
+ allowing for batch parallelism only.
+
+ Uses Keras add_loss() and add_metric() APIs.
+
+ Attributes:
+ num_experts: Number of experts (i.e. number of independent feed-forward
+ blocks).
+ """
+
+ def __init__(
+ self,
+ experts: FeedForwardExperts,
+ router: MaskedRouter,
+ *,
+ train_capacity_factor: float = 1.0,
+ eval_capacity_factor: float = 1.0,
+ examples_per_group: float = 1.0,
+ name: str = "moe",
+ **kwargs):
+ """Init.
+
+ Args:
+ experts: Instance of FeedForwardExperts. Needs to have the same
+ num_experts as the router.
+ router: Instance of MaskedRouter to route the tokens to
+ the different experts.
+ train_capacity_factor: Scaling factor to increase the expert token
+ capacity during training. This factor plays an analogous, but slightly
+ different, role depending on the routing assignment algorithm:
+ - For "tokens choose" routing, the capacity factor only affects the
+ maximum number of tokens that an expert will process. It does not
+ affect how many experts a given token is routed to; see the
+ num_selected_experts attributes of "tokens choose" routers.
+ - For "experts choose" routing, because experts always fill their
+ buffer, increasing the capacity factor will increase the number of
+ tokens that an expert will process AND will indirectly increase the
+ number of experts that a given token is routed to.
+ eval_capacity_factor: As above, but used during evaluation.
+ examples_per_group: Number of examples to form a group. Router then
+ performs top_k token selection for each expert on a per group basis.
+ E.g. when `examples_per_group=4.0`, tokens are assigned to experts in
+ groups formed from 4 examples. When `examples_per_group=0.5`,
+ each example is split into 2 groups.
+ `examples_per_group` must divide the local batch size.
+ A larger group size will result in slower but more accurate top-k and
+ sorting computations, whereas a smaller group size will result in faster
+ but more approximate (and potentially less stable) routing choices.
+ In practice, we find that imperfect routing choices are tolerable and
+ recommend choosing a group size on the order of 4096 tokens, although
+ this number will vary based on model configuration and size.
+ name: Layer name.
+ **kwargs: Forwarded to super.
+ """
+ super().__init__(name=name, **kwargs)
+ self._experts = experts
+ self._router = router
+
+ self.num_experts = experts.num_experts
+ assert experts.num_experts == router.num_experts
+
+ self._train_capacity_factor = train_capacity_factor
+ self._eval_capacity_factor = eval_capacity_factor
+ self._examples_per_group = examples_per_group
+
+ def call(self,
+ inputs: tf.Tensor,
+ *,
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Applies MoeLayer.
+
+ Args:
+ inputs: Batch of input embeddings of shape
+ [batch_size, seq_length, hidden_dim].
+ training: Only apply dropout and jitter noise during training. If not
+ provided taken from tf_keras.backend.
+
+ Returns:
+ Transformed inputs with same shape as inputs:
+ [batch_size, seq_length, hidden_dim].
+
+ Raises:
+ ValueError if we cannot find a group_size satisfying given requirements.
+ """
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+
+ # inputs shape [batch_size, seq_length, hidden_dim]
+ batch_size, seq_length, hidden_dim = inputs.shape
+ if batch_size is not None:
+ if self._examples_per_group > batch_size:
+ raise ValueError(
+ f"examples_per_group={self._examples_per_group} is larger than the "
+ "number of examples available in the local (per-device) batch_size="
+ f"{batch_size}. Either decrease examples_per_group or increase the "
+ "batch_size.")
+ tokens_per_group = int(seq_length * self._examples_per_group)
+
+ if training:
+ capacity_factor = self._train_capacity_factor
+ else:
+ capacity_factor = self._eval_capacity_factor
+ # Each group will send expert_capacity tokens to each expert.
+ expert_capacity = int(
+ round(capacity_factor * tokens_per_group / self.num_experts))
+
+ # Reshape batch and sequence/token dimensions for expert routing.
+ x = tf.reshape(inputs, (-1, tokens_per_group, hidden_dim))
+
+ x = self._mask_and_dispatch_to_experts(x, expert_capacity, training)
+
+ # Return to original input shape.
+ x = tf.reshape(x, (-1, seq_length, hidden_dim))
+ return x
+
+ def _mask_and_dispatch_to_experts(self, inputs: tf.Tensor,
+ expert_capacity: int,
+ training: bool) -> tf.Tensor:
+ """Wraps expert masked routing and dispatching algorithm.
+
+ This algorithm takes the following steps:
+ (1) Compute dispatch mask and combine array using self._router.
+ (2) Dispatch inputs to experts based on dispatch mask.
+ (3) Recombine individual expert outputs using combine array.
+
+ Args:
+ inputs: [num_groups, tokens_per_group, hidden_dim] inputs to
+ send to experts.
+ expert_capacity: Each group will send this many tokens to each expert.
+ training: If true, apply jitter noise during routing and dropout
+ during expert computation.
+
+ Returns:
+ [num_groups, num_tokens_per_group, hidden_dim] outputs from
+ experts.
+ """
+ # Shape [num_groups, tokens_per_group, num_experts, expert_capacity]
+ router_mask = self._router(
+ inputs,
+ expert_capacity=expert_capacity,
+ training=training)
+
+ # Shape [num_groups, num_experts, expert_capacity, hidden_dim]
+ expert_inputs = tf.einsum(
+ "gtec,gth->gech",
+ router_mask.dispatch_mask,
+ inputs)
+
+ expert_outputs = self._experts(expert_inputs, training=training)
+
+ # Shape [num_groups, tokens_per_group, hidden_dim]
+ combined_outputs = tf.einsum(
+ "gtec,gech->gth",
+ router_mask.combine_array,
+ expert_outputs)
+
+ return combined_outputs
+
+
+class MoeLayerWithBackbone(tf_keras.layers.Layer):
+ """Sparse MoE layer plus a FeedForward layer evaluated for all tokens.
+
+ Uses Keras add_loss() and add_metric() APIs.
+ """
+
+ def __init__(
+ self,
+ moe: MoeLayer,
+ backbone_d_ff: int,
+ *,
+ inner_dropout: float = 0.0,
+ output_dropout: float = 0.0,
+ activation: Callable[[tf.Tensor],
+ tf.Tensor] = tf_keras.activations.gelu,
+ kernel_initializer: _InitializerType = _DEFAULT_KERNEL_INITIALIZER,
+ bias_initializer: _InitializerType = _DEFAULT_BIAS_INITIALIZER,
+ name: str = "moe_with_backbone",
+ **kwargs):
+ """Init.
+
+ Args:
+ moe: Instance of MoeLayer with experts and router.
+ backbone_d_ff: Dimension of feed-forward layer of a lightweight backbone,
+ which is evaluated for all tokens.
+ inner_dropout: The dropout probability to be applied after intermediate
+ activations for the backbone.
+ output_dropout: The dropout probability to be applied after the output
+ of the backbone.
+ activation: (Nonlinear) transform applied in the backbone.
+ kernel_initializer: Initialization scheme for kernels in the backbone.
+ bias_initializer: Initialization scheme for biases in the backbone.
+ name: Layer name.
+ **kwargs: Forwarded to super.
+ """
+ super().__init__(name=name, **kwargs)
+ self._moe = moe
+
+ self._backbone = FeedForward(
+ backbone_d_ff,
+ inner_dropout=inner_dropout,
+ output_dropout=output_dropout,
+ activation=activation,
+ kernel_initializer=tf_utils.clone_initializer(kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(bias_initializer),
+ name="backbone")
+
+ def call(self,
+ inputs: tf.Tensor,
+ *,
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Applies MoeLayerWithBackbone layer.
+
+ Args:
+ inputs: Batch of input embeddings of shape
+ [batch_size, seq_length, hidden_dim].
+ training: Only apply dropout and jitter noise during training. If not
+ provided taken from tf_keras.backend.
+
+ Returns:
+ Transformed inputs with same shape as inputs:
+ [batch_size, seq_length, hidden_dim].
+ """
+ return self._backbone(
+ inputs, training=training) + self._moe(
+ inputs, training=training)
diff --git a/official/nlp/modeling/layers/moe_test.py b/official/nlp/modeling/layers/moe_test.py
new file mode 100644
index 00000000000..17f80789b8a
--- /dev/null
+++ b/official/nlp/modeling/layers/moe_test.py
@@ -0,0 +1,238 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for moe.py."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling.layers import moe
+
+
+def small_config():
+ """Creates a small model config that can be used by all tests."""
+ config = {}
+ config['d_ff'] = 32
+ config['output_dropout'] = 0.1
+
+ config['num_experts'] = 2
+ config['expert_d_ff'] = 33
+ config['expert_dropout_rate'] = 0.1
+ config['jitter_noise'] = 0.1
+ config['train_capacity_factor'] = 1.0
+ config['eval_capacity_factor'] = 1.0
+ config['examples_per_group'] = 2.0
+
+ config['backbone_d_ff'] = 13
+ return config
+
+
+def make_input_ones(batch_size: int = 4,
+ seq_length: int = 10,
+ hidden_dim: int = 7) -> tf.Tensor:
+ return tf.ones((batch_size, seq_length, hidden_dim))
+
+
+def make_experts_input_ones(num_groups: int = 1,
+ num_experts: int = 2,
+ expert_capacity: int = 5,
+ hidden_dim: int = 7) -> tf.Tensor:
+ return tf.ones((num_groups, num_experts, expert_capacity, hidden_dim))
+
+
+class MoeTest(tf.test.TestCase):
+
+ def tearDown(self):
+ super().tearDown()
+ tf_keras.mixed_precision.set_global_policy('float32')
+
+ def test_router_z_loss_dtype(self):
+ x = tf.constant([[[10.0, 5.0]]], dtype=tf.float32)
+ y = moe._router_z_loss(x)
+ expected = (5 + np.log(np.exp(5) + 1))**2
+ self.assertAllClose(expected, y, atol=1e-7)
+ self.assertDTypeEqual(y, tf.float32)
+
+ def test_router_z_loss_shape(self):
+ x = make_input_ones(2, 5, 7)
+ y = moe._router_z_loss(x)
+ expected = (np.log(7) + 1)**2
+ self.assertAllClose(expected, y, atol=1e-7)
+
+ def test_experts_choose_masked_router_dtype_shape(self):
+ tf_keras.mixed_precision.set_global_policy('mixed_bfloat16')
+ num_groups = 2
+ tokens_per_group = 3
+ hidden_dim = tokens_per_group
+ num_experts = tokens_per_group
+ expert_capacity = 2
+ x = np.zeros([num_groups, tokens_per_group, hidden_dim])
+ x[0, 0, 0] += 1
+ x[0, :2, :2] += 1
+ x[1, 1:, 1:] += 1
+ x[1, -1, -1] += 1
+
+ router = moe.ExpertsChooseMaskedRouter(
+ num_experts=num_experts,
+ jitter_noise=0.1,
+ use_bias=True,
+ kernel_initializer=tf_keras.initializers.get('identity'),
+ bias_initializer=tf_keras.initializers.get('ones'))
+ router_mask = router(x, expert_capacity=expert_capacity, training=False)
+
+ self.assertDTypeEqual(router_mask.dispatch_mask, tf.bfloat16)
+ self.assertDTypeEqual(router_mask.combine_array, tf.bfloat16)
+
+ expect_shape = [num_groups, tokens_per_group, num_experts, expert_capacity]
+ self.assertEqual(expect_shape, router_mask.dispatch_mask.shape)
+ self.assertEqual(expect_shape, router_mask.combine_array.shape)
+
+ # top_k call may not be sorted, so can't compare the output directly
+ # Check that the output contains only 0s and 1s
+ out_dm = router_mask.dispatch_mask.numpy()
+ self.assertSetEqual({0, 1}, set(out_dm.flatten().astype(np.int32)))
+ # Check that the right tokens for selected
+ out_dm_indices = np.dot(
+ out_dm.transpose((0, 2, 3, 1)), np.arange(tokens_per_group))
+ # Shape [num_groups, num_experts, expert_capacity]
+ self.assertSetEqual({0, 1}, set(out_dm_indices[0, 0, :].astype(np.int32)))
+ self.assertSetEqual({1, 2}, set(out_dm_indices[0, 1, :].astype(np.int32)))
+ self.assertSetEqual({1, 2}, set(out_dm_indices[0, 2, :].astype(np.int32)))
+ self.assertSetEqual({0, 1}, set(out_dm_indices[1, 0, :].astype(np.int32)))
+ self.assertSetEqual({0, 1}, set(out_dm_indices[1, 1, :].astype(np.int32)))
+ self.assertSetEqual({1, 2}, set(out_dm_indices[1, 2, :].astype(np.int32)))
+
+ out_ca = router_mask.combine_array.numpy()
+ out_ca = np.dot(out_ca, np.ones((expert_capacity,)))
+
+ expected_combine_array = np.array([[[0.66, 0.0, 0.0], [0.42, 0.42, 0.16],
+ [0.0, 0.33, 0.33]],
+ [[0.33, 0.33, 0.0], [0.16, 0.42, 0.42],
+ [0.0, 0.0, 0.66]]])
+ self.assertAllClose(expected_combine_array, out_ca, atol=1e-2)
+
+ def test_feed_forward_shape_and_vars(self):
+ config = small_config()
+ layer = moe.FeedForward(
+ d_ff=config['d_ff'], output_dropout=config['output_dropout'])
+ inputs = make_input_ones()
+ outputs = layer(inputs)
+ self.assertAllEqual(tf.shape(inputs), tf.shape(outputs))
+ var_names = sorted([v.name for v in layer.trainable_variables])
+ self.assertAllEqual([
+ 'feed_forward/intermediate/bias:0',
+ 'feed_forward/intermediate/kernel:0', 'feed_forward/output/bias:0',
+ 'feed_forward/output/kernel:0'
+ ], var_names)
+
+ def test_feed_forward_manual(self):
+ config = small_config()
+ layer = moe.FeedForward(
+ d_ff=config['d_ff'],
+ output_dropout=config['output_dropout'],
+ activation=tf_keras.activations.relu,
+ kernel_initializer=tf_keras.initializers.get('ones'),
+ bias_initializer=tf_keras.initializers.get('ones'))
+ inputs = make_input_ones(1, 2, 3)
+ outputs = layer(inputs, training=False)
+ manual_outputs = tf.constant([[[129.0, 129.0, 129.0], [129.0, 129.0,
+ 129.0]]])
+ self.assertAllClose(manual_outputs, outputs, atol=1e-7)
+
+ def test_feed_forward_experts_shape_and_vars(self):
+ config = small_config()
+ layer = moe.FeedForwardExperts(
+ num_experts=config['num_experts'],
+ d_ff=config['expert_d_ff'],
+ output_dropout=config['expert_dropout_rate'])
+ inputs = make_experts_input_ones()
+ outputs = layer(inputs)
+ self.assertAllEqual(tf.shape(inputs), tf.shape(outputs))
+ var_names = sorted([v.name for v in layer.trainable_variables])
+ self.assertAllEqual([
+ 'experts/intermediate/bias:0', 'experts/intermediate/kernel:0',
+ 'experts/output/bias:0', 'experts/output/kernel:0'
+ ], var_names)
+
+ def test_feed_forward_experts_manual(self):
+ config = small_config()
+ layer = moe.FeedForwardExperts(
+ num_experts=1,
+ d_ff=config['expert_d_ff'],
+ output_dropout=config['expert_dropout_rate'],
+ activation=tf_keras.activations.relu,
+ kernel_initializer=tf_keras.initializers.get('ones'),
+ bias_initializer=tf_keras.initializers.get('ones'))
+ inputs = make_experts_input_ones(1, 1, 2, 3)
+ outputs = layer(inputs, training=False)
+ manual_outputs = tf.constant([[[[133.0, 133.0, 133.0],
+ [133.0, 133.0, 133.0]]]])
+ self.assertAllClose(manual_outputs, outputs, atol=1e-7)
+
+ def test_moe_layer(self):
+ config = small_config()
+ experts = moe.FeedForwardExperts(
+ num_experts=config['num_experts'],
+ d_ff=config['expert_d_ff'],
+ output_dropout=config['expert_dropout_rate'])
+ router = moe.ExpertsChooseMaskedRouter(
+ config['num_experts'], jitter_noise=config['jitter_noise'])
+ moe_layer = moe.MoeLayer(
+ experts,
+ router,
+ train_capacity_factor=config['train_capacity_factor'],
+ eval_capacity_factor=config['eval_capacity_factor'],
+ examples_per_group=config['examples_per_group'])
+
+ inputs = make_input_ones()
+ outputs = moe_layer(inputs, training=True)
+ self.assertAllEqual(tf.shape(inputs), tf.shape(outputs))
+
+ var_names = sorted([v.name for v in moe_layer.trainable_variables])
+ self.assertAllEqual([
+ 'moe/experts/intermediate/bias:0', 'moe/experts/intermediate/kernel:0',
+ 'moe/experts/output/bias:0', 'moe/experts/output/kernel:0',
+ 'moe/router/router_weights/bias:0', 'moe/router/router_weights/kernel:0'
+ ], var_names)
+ self.assertLen(moe_layer.losses, 1)
+ metrics = [metric.name for metric in moe_layer.metrics]
+ self.assertSetEqual(
+ {
+ 'router_z_loss', 'unscaled_router_z_loss', 'load_balancing_loss',
+ 'fraction_tokens_left_behind', 'router_confidence', 'expert_usage'
+ }, set(metrics))
+
+ def test_moe_layer_with_backbone(self):
+ config = small_config()
+ experts = moe.FeedForwardExperts(
+ num_experts=config['num_experts'],
+ d_ff=config['expert_d_ff'],
+ output_dropout=config['expert_dropout_rate'])
+ router = moe.ExpertsChooseMaskedRouter(
+ config['num_experts'], jitter_noise=config['jitter_noise'])
+ moe_layer = moe.MoeLayer(
+ experts,
+ router,
+ train_capacity_factor=config['train_capacity_factor'],
+ eval_capacity_factor=config['eval_capacity_factor'],
+ examples_per_group=config['examples_per_group'])
+ layer = moe.MoeLayerWithBackbone(moe_layer, config['backbone_d_ff'])
+
+ inputs = make_input_ones()
+ outputs = layer(inputs)
+ self.assertAllEqual(tf.shape(inputs), tf.shape(outputs))
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/nlp/modeling/layers/multi_channel_attention.py b/official/nlp/modeling/layers/multi_channel_attention.py
index f6ba0e007cc..a6ec6dc5413 100644
--- a/official/nlp/modeling/layers/multi_channel_attention.py
+++ b/official/nlp/modeling/layers/multi_channel_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,12 +17,13 @@
import math
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
from official.modeling import tf_utils
from official.nlp.modeling.layers import masked_softmax
-class VotingAttention(tf.keras.layers.Layer):
+class VotingAttention(tf_keras.layers.Layer):
"""Voting Attention layer.
Args:
@@ -48,38 +49,40 @@ def __init__(self,
kernel_constraint=None,
bias_constraint=None,
**kwargs):
- super(VotingAttention, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._num_heads = num_heads
self._head_size = head_size
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
def build(self, unused_input_shapes):
common_kwargs = dict(
- kernel_initializer=self._kernel_initializer,
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
kernel_constraint=self._kernel_constraint,
bias_constraint=self._bias_constraint)
- self._query_dense = tf.keras.layers.experimental.EinsumDense(
+ self._query_dense = tf_keras.layers.EinsumDense(
"BAE,ENH->BANH",
output_shape=(None, self._num_heads, self._head_size),
bias_axes="NH",
name="query",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- self._key_dense = tf.keras.layers.experimental.EinsumDense(
+ self._key_dense = tf_keras.layers.EinsumDense(
"BAE,ENH->BANH",
output_shape=(None, self._num_heads, self._head_size),
bias_axes="NH",
name="key",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- super(VotingAttention, self).build(unused_input_shapes)
+ super().build(unused_input_shapes)
def call(self, encoder_outputs, doc_attention_mask):
num_docs = tf_utils.get_shape_list(encoder_outputs, expected_rank=[4])[1]
@@ -100,7 +103,7 @@ def call(self, encoder_outputs, doc_attention_mask):
return tf.nn.softmax(doc_attention_probs + infadder)
-class MultiChannelAttention(tf.keras.layers.MultiHeadAttention):
+class MultiChannelAttention(tf_keras.layers.MultiHeadAttention):
"""Multi-channel Attention layer.
Introduced in, [Generating Representative Headlines for News Stories
@@ -120,10 +123,10 @@ class MultiChannelAttention(tf.keras.layers.MultiHeadAttention):
"""
def _build_attention(self, rank):
- super(MultiChannelAttention, self)._build_attention(rank) # pytype: disable=attribute-error # typed-keras
+ super()._build_attention(rank) # pytype: disable=attribute-error # typed-keras
self._masked_softmax = masked_softmax.MaskedSoftmax(mask_expansion_axes=[2])
- def call(self,
+ def call(self, # pyrefly: ignore[bad-override]
query,
value,
key=None,
diff --git a/official/nlp/modeling/layers/multi_channel_attention_test.py b/official/nlp/modeling/layers/multi_channel_attention_test.py
index 8c022046756..0ff152a845d 100644
--- a/official/nlp/modeling/layers/multi_channel_attention_test.py
+++ b/official/nlp/modeling/layers/multi_channel_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for projects.nhnet.multi_channel_attention."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import multi_channel_attention
diff --git a/official/nlp/modeling/layers/multi_query_attention.py b/official/nlp/modeling/layers/multi_query_attention.py
new file mode 100644
index 00000000000..ad2437921f4
--- /dev/null
+++ b/official/nlp/modeling/layers/multi_query_attention.py
@@ -0,0 +1,423 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras-based attention layers to support multi-query attention.
+
+Based on https://arxiv.org/pdf/1911.02150.pdf and
+https://arxiv.org/pdf/2305.13245.pdf.
+"""
+
+import math
+import string
+from typing import Optional, Sequence, Union
+
+import gin
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
+
+_CHR_IDX = string.ascii_lowercase
+
+
+def _build_proj_equation(free_dims: int, bound_dims: int, output_dims: int):
+ """Builds an einsum equation for projections inside attention layer.
+
+ Args:
+ free_dims: The number of free dimensions which are copied from input to
+ output.
+ bound_dims: The number of bound dimensions part of input which are combined
+ with the kernel to produce output.
+ output_dims: The number of output dimensions.
+
+ Returns:
+ A tuple of einsum equation, bias axes and output rank.
+ """
+
+ input_str = ""
+ kernel_str = ""
+ output_str = ""
+ bias_axes = ""
+ letter_offset = 0
+ for i in range(free_dims):
+ char = _CHR_IDX[i + letter_offset]
+ input_str += char
+ output_str += char
+
+ letter_offset += free_dims
+ for i in range(bound_dims):
+ char = _CHR_IDX[i + letter_offset]
+ input_str += char
+ kernel_str += char
+
+ letter_offset += bound_dims
+ for i in range(output_dims):
+ char = _CHR_IDX[i + letter_offset]
+ kernel_str += char
+ output_str += char
+ bias_axes += char
+ equation = f"{input_str},{kernel_str}->{output_str}"
+
+ return equation, bias_axes, len(output_str)
+
+
+def _get_output_shape(
+ output_rank: int, known_last_dims: Sequence[int]
+) -> list[Optional[int]]:
+ return [None] * (output_rank - len(known_last_dims)) + list(known_last_dims)
+
+
+class MultiHeadAttention(tf_keras.layers.MultiHeadAttention):
+ """Multi-query attention layer."""
+
+ def __init__(
+ self, num_kv_heads=None, enable_gqa_optimization=False, **kwargs
+ ):
+ # num_kv_heads defines the number of key/value heads. A value of 1 means
+ # that the key/value heads are shared across all query heads. Any other
+ # value must be less than num_heads and must divide num_heads exactly. If
+ # num_kv_heads is greater than 1, query heads are split into groups of
+ # num_kv_heads.
+ super().__init__(**kwargs)
+ self._num_kv_heads = num_kv_heads or self._num_heads
+ # TODO(akandoor): Remove this flag once the GQA optimization is rolled out.
+ # This flag is used to enable order of K,G in the einsum equations.
+ # This optimization is only used in GQA, and is disabled by default.
+ # If enabled, the einsum equations are:
+ # 1. Dot product: "...SKH,...TKGH->...KGTS"
+ # 2. Combine: "...KGTS,...SKH->...TKGH"
+ # If disabled, the einsum equations are:
+ # 1. Dot product: "...SKH,...TKnH->...nKTS"
+ # 2. Combine: "...nKTS,...SKH->...TnKH"
+ self._enable_gqa_optimization = enable_gqa_optimization
+ assert (
+ self._num_kv_heads < self._num_heads
+ ), "num_kv_heads must be less than num_heads."
+ assert (
+ self._num_heads % self._num_kv_heads == 0
+ ), "num_kv_heads needs to divide num_heads exactly."
+
+ def get_config(self):
+ config = super().get_config()
+ config.update({"num_kv_heads": self._num_kv_heads})
+ return config
+
+ def _build_from_signature(
+ self,
+ query: Union[tf.Tensor, tf.TensorShape],
+ value: Union[tf.Tensor, tf.TensorShape],
+ key: Optional[Union[tf.Tensor, tf.TensorShape]] = None,
+ ):
+ """Builds layers and variables.
+
+ Once the method is called, self._built_from_signature will be set to
+ True.
+
+ Args:
+ query: Query tensor or TensorShape.
+ value: Value tensor or TensorShape.
+ key: Key tensor or TensorShape.
+ """
+ # pytype: disable=attribute-error
+ super()._build_from_signature(query=query, value=value, key=key)
+ # pytype: enable=attribute-error
+
+ with tf.init_scope():
+ # Key, value are shared across heads in multi-query attention.
+ # Overwrite the K, V projections, logits & attend einsum equations to
+ # remove the number of attention head dimension in K, V related tensors.
+ #
+ # The following capital letters are used to denote the tensor dimension
+ # parameters:
+ # B = batch size
+ # S = length of the key/value (source)
+ # T = length of the query (target)
+ # N = number of query attention heads
+ # K = number of key/value heads
+ # n = N // K
+ # H = dimensions of each attention head.
+ #
+ if self._num_kv_heads == 1:
+ output_dims = 1
+ key_last_dims = [self._key_dim]
+ value_last_dims = [self._value_dim]
+ self._dot_product_equation = "...SH,...TNH->...NTS"
+ self._combine_equation = "...NTS,...SH->...TNH"
+ else:
+ output_dims = 2
+ key_last_dims = [self._num_kv_heads, self._key_dim]
+ value_last_dims = [self._num_kv_heads, self._value_dim]
+ if self._enable_gqa_optimization:
+ self._dot_product_equation = "...SKH,...TKGH->...KGTS"
+ self._combine_equation = "...KGTS,...SKH->...TKGH"
+ else:
+ self._dot_product_equation = "...SKH,...TKnH->...nKTS"
+ self._combine_equation = "...nKTS,...SKH->...TnKH"
+
+ einsum_equation, bias_axes, output_rank = _build_proj_equation(
+ free_dims=self._key_shape.rank - 1,
+ bound_dims=1,
+ output_dims=output_dims,
+ )
+ self._key_dense = tf_keras.layers.EinsumDense(
+ einsum_equation,
+ output_shape=_get_output_shape(output_rank - 1, key_last_dims),
+ bias_axes=bias_axes if self._use_bias else None,
+ name="key",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+ einsum_equation, bias_axes, output_rank = _build_proj_equation(
+ free_dims=self._value_shape.rank - 1,
+ bound_dims=1,
+ output_dims=output_dims,
+ )
+ self._value_dense = tf_keras.layers.EinsumDense(
+ einsum_equation,
+ output_shape=_get_output_shape(output_rank - 1, value_last_dims),
+ bias_axes=bias_axes if self._use_bias else None,
+ name="value",
+ **self._get_common_kwargs_for_sublayer(),
+ )
+ self._qkv_rank = (
+ output_rank if self._num_kv_heads > 1 else output_rank + 1
+ )
+
+ def _compute_attention(
+ self, query, key, value, attention_mask=None, training=None
+ ):
+ if self._num_kv_heads > 1:
+ query = tf.reshape(
+ query,
+ [
+ tf.shape(query)[0],
+ tf.shape(query)[1],
+ self._num_kv_heads,
+ self._num_heads // self._num_kv_heads,
+ tf.shape(query)[-1],
+ ],
+ )
+
+ # pytype: disable=attribute-error
+ attention_output, attention_scores = super()._compute_attention(
+ query, key, value, attention_mask=attention_mask, training=training
+ )
+ # pytype: enable=attribute-error
+ if self._num_kv_heads != 1:
+ attention_output = tf.reshape(
+ attention_output,
+ [
+ tf.shape(attention_output)[0],
+ tf.shape(attention_output)[1],
+ self._num_heads,
+ tf.shape(attention_output)[-1],
+ ],
+ )
+ attention_scores = tf.reshape(
+ attention_scores,
+ [
+ tf.shape(attention_scores)[0],
+ self._num_heads,
+ tf.shape(attention_scores)[-2],
+ tf.shape(attention_scores)[-1],
+ ],
+ )
+ return attention_output, attention_scores
+
+
+@tf_keras.utils.register_keras_serializable(package="Text")
+@gin.configurable
+class TalkingHeadsMultiQueryAttention(MultiHeadAttention):
+ """Implements Talking-Heads Attention combined with Multi-Query Attention.
+
+ See https://arxiv.org/pdf/2003.02436 for more details.
+ TODO(akandoor): Make num talking heads configurable. Currently, num talking
+ heads is fixed to num query heads.
+
+ This class inherits from MultiQueryAttention to get the MQA-specific
+ logic for __init__, get_config.
+
+ It then overrides _build_from_signature to add the talking-heads weights
+ and overrides _compute_attention to merge the MQA wrapper (for
+ reshaping) with the THA computation (for pre/post-softmax projections).
+ """
+
+ def _build_from_signature(
+ self,
+ query: Union[tf.Tensor, tf.TensorShape],
+ value: Union[tf.Tensor, tf.TensorShape],
+ key: Optional[Union[tf.Tensor, tf.TensorShape]] = None,
+ ):
+ """Builds layers and variables."""
+ # Call the parent (MultiQueryAttention) _build_from_signature.
+ super()._build_from_signature(query=query, value=value, key=key)
+ # Now, *after* all MQA setup is done, we add the THA setup logic.
+ qkv_rank = self._qkv_rank
+ # TalkingHeadsAttention logic to the MQA build logic.
+ num_batch_dims = qkv_rank - len(self._attention_axes) - 2
+ attn_scores_rank = num_batch_dims + 1 + len(self._attention_axes) * 2
+ scores_notation = _CHR_IDX[:attn_scores_rank]
+ projection_notation = scores_notation[num_batch_dims] + (
+ _CHR_IDX[attn_scores_rank])
+ projected_scores_notation = scores_notation[:num_batch_dims] + (
+ _CHR_IDX[attn_scores_rank] + scores_notation[num_batch_dims + 1:])
+ self._talking_heads_equation = "%s,%s->%s" % (
+ scores_notation, projection_notation, projected_scores_notation)
+
+ with tf.init_scope():
+ self._pre_softmax_weight = self.add_weight(
+ "pre_softmax_weight",
+ shape=(self._num_heads, self._num_heads),
+ initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ regularizer=self._kernel_regularizer,
+ constraint=self._kernel_constraint,
+ dtype=self.dtype,
+ trainable=True)
+ self._post_softmax_weight = self.add_weight(
+ "post_softmax_weight",
+ shape=(self._num_heads, self._num_heads),
+ initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ regularizer=self._kernel_regularizer,
+ constraint=self._kernel_constraint,
+ dtype=self.dtype,
+ trainable=True)
+
+ def _compute_attention(
+ self, query, key, value, attention_mask=None, training=None
+ ):
+ """Applies Dot-product attention, merging MQA wrapper and THA computation.
+
+ Args:
+ query: Projected query `Tensor` of shape `[B, T, N, key_dim]`.
+ key: Projected key `Tensor` of shape `[B, T, N, key_dim]`.
+ value: Projected value `Tensor` of shape `[B, T, N, value_dim]`.
+ attention_mask: a boolean mask of shape `[B, T, S]`, that prevents
+ attention to certain positions.
+ training: Python boolean indicating whether the layer should behave in
+ training mode (adding dropout) or in inference mode (doing nothing).
+
+ Returns:
+ attention_output: Multi-headed outputs of attention computation.
+ attention_scores: Multi-headed attention weights.
+ """
+ # This is the MQA "wrapper" logic for grouped queries
+ query_shape = tf.shape(query)
+ if self._num_kv_heads > 1:
+ query = tf.reshape(
+ query,
+ [
+ query_shape[0],
+ query_shape[1],
+ self._num_kv_heads,
+ self._num_heads // self._num_kv_heads,
+ query_shape[-1],
+ ],
+ )
+
+ # This is the THA "computation" logic
+ # Note: Applying scalar multiply at the smaller end of einsum improves
+ # XLA performance, but may introduce slight numeric differences in
+ # the Transformer attention head.
+ query = tf.multiply(
+ query, 1.0 / math.sqrt(float(self._key_dim))
+ )
+
+ # Note: self._dot_product_equation was set by _build_from_signature
+ # (from MQA) to be MQA-compatible.
+ attention_scores = tf.einsum(self._dot_product_equation, key, query)
+
+ # --- Talking-Heads modification for MQA ---
+ # The THA _talking_heads_equation expects scores of shape [B, N, T, S].
+ # The MQA _dot_product_equation produces [B, K, G, T, S].
+ # We must reshape before and after applying TH logic.
+ scores_shape = tf.shape(attention_scores)
+ if self._num_kv_heads > 1:
+ # Reshape from [B, K, G, T, S] to [B, N, T, S]
+ attention_scores = tf.reshape(
+ attention_scores,
+ [
+ scores_shape[0], # Batch
+ self._num_heads, # N = K * G
+ scores_shape[-2], # T
+ scores_shape[-1] # S
+ ]
+ )
+
+ # Apply linear projection before softmax
+ attention_scores = tf.einsum(self._talking_heads_equation, attention_scores,
+ self._pre_softmax_weight)
+
+ # Normalize the attention scores to probabilities.
+ attention_scores = self._masked_softmax(attention_scores, attention_mask)
+
+ # Apply linear projection after softmax
+ attention_scores = tf.einsum(self._talking_heads_equation, attention_scores,
+ self._post_softmax_weight)
+
+ # Reshape back to MQA-compatible shape [B, K, G, T, S]
+ # before the final combine_equation
+ if self._num_kv_heads > 1:
+ if self._enable_gqa_optimization:
+ attention_scores = tf.reshape(
+ attention_scores,
+ [
+ scores_shape[0], # B
+ self._num_kv_heads, # K
+ self._num_heads // self._num_kv_heads, # G
+ scores_shape[-2], # T
+ scores_shape[-1] # S
+ ]
+ )
+ else:
+ attention_scores = tf.reshape(
+ attention_scores,
+ [
+ scores_shape[0], # B
+ self._num_heads // self._num_kv_heads, # G
+ self._num_kv_heads, # K
+ scores_shape[-2], # T
+ scores_shape[-1] # S
+ ]
+ )
+
+ # This is actually dropping out entire tokens to attend to.
+ attention_scores_dropout = self._dropout_layer(
+ attention_scores, training=training)
+
+ # Note: self._combine_equation was set by _build_from_signature
+ # (from MQA) to be MQA-compatible.
+ attention_output = tf.einsum(self._combine_equation,
+ attention_scores_dropout, value)
+
+ # This is the MQA "wrapper" logic for grouped queries
+ if self._num_kv_heads > 1:
+ attention_output = tf.reshape(
+ attention_output,
+ [
+ query_shape[0],
+ query_shape[1],
+ self._num_heads,
+ tf.shape(attention_output)[-1],
+ ],
+ )
+ # We also need to reshape the final scores back to [B, N, T, S]
+ # for the return value.
+ attention_scores = tf.reshape(
+ attention_scores,
+ [
+ query_shape[0],
+ self._num_heads,
+ tf.shape(attention_scores)[-2],
+ tf.shape(attention_scores)[-1],
+ ],
+ )
+
+ return attention_output, attention_scores
diff --git a/official/nlp/modeling/layers/multi_query_attention_test.py b/official/nlp/modeling/layers/multi_query_attention_test.py
new file mode 100644
index 00000000000..a29c8907fe9
--- /dev/null
+++ b/official/nlp/modeling/layers/multi_query_attention_test.py
@@ -0,0 +1,415 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for multi-query attention layer."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling.layers import multi_query_attention
+
+
+class MultiQueryAttentionTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ ("key_value_same_proj_mqa", 1, None, None, [40, 80]),
+ ("key_value_different_proj_mqa", 1, 32, 60, [40, 60]),
+ ("key_value_same_proj_gqa", 3, None, None, [40, 80]),
+ ("key_value_different_proj_gqa", 3, 32, 60, [40, 60]),
+ )
+ def test_non_masked_attention(
+ self, num_kv_heads, value_dim, output_shape, output_dims
+ ):
+ """Test that the attention layer can be created without a mask tensor."""
+ test_layer = multi_query_attention.MultiHeadAttention(
+ num_heads=12,
+ num_kv_heads=num_kv_heads,
+ key_dim=64,
+ value_dim=value_dim,
+ output_shape=output_shape,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(20, 80))
+ output = test_layer(query=query, value=value)
+ self.assertEqual(output.shape.as_list(), [None] + output_dims)
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 3),
+ )
+ def test_non_masked_self_attention(self, num_kv_heads):
+ """Test with one input (self-attenntion) and no mask tensor."""
+ test_layer = multi_query_attention.MultiHeadAttention(
+ num_heads=12, num_kv_heads=num_kv_heads, key_dim=64
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ output = test_layer(query, query)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 3),
+ )
+ def test_attention_scores(self, num_kv_heads):
+ """Test attention outputs with coefficients."""
+ test_layer = multi_query_attention.MultiHeadAttention(
+ num_heads=12, num_kv_heads=num_kv_heads, key_dim=64
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ output, coef = test_layer(query, query, return_attention_scores=True)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+ self.assertEqual(coef.shape.as_list(), [None, 12, 40, 40])
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 3),
+ )
+ def test_attention_scores_with_values(self, num_kv_heads):
+ """Test attention outputs with coefficients."""
+ test_layer = multi_query_attention.MultiHeadAttention(
+ num_heads=12, num_kv_heads=num_kv_heads, key_dim=64
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(60, 80))
+ output, coef = test_layer(query, value, return_attention_scores=True)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+ self.assertEqual(coef.shape.as_list(), [None, 12, 40, 60])
+
+ @parameterized.named_parameters(
+ ("with_bias_mqa", 1, True),
+ ("no_bias_mqa", 1, False),
+ ("with_bias_gqa", 2, True),
+ ("no_bias_gqa", 2, False),
+ )
+ def test_masked_attention(self, num_kv_heads, use_bias):
+ """Test with a mask tensor."""
+ test_layer = multi_query_attention.MultiHeadAttention(
+ num_heads=4, num_kv_heads=num_kv_heads, key_dim=2, use_bias=use_bias
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ batch_size = 3
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
+ output = test_layer(query=query, value=value, attention_mask=mask_tensor)
+
+ # Create a model containing the test layer.
+ model = tf_keras.Model([query, value, mask_tensor], output)
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = 10 * np.random.random_sample((batch_size, 4, 8))
+ to_data = 10 * np.random.random_sample((batch_size, 2, 8))
+
+ # Invoke the data with a random set of mask data. This should mask at
+ # least one element.
+ mask_data = np.random.randint(2, size=(batch_size, 4, 2))
+ masked_output_data = model.predict([from_data, to_data, mask_data])
+
+ # Invoke the same data, but with a null mask (where no elements are
+ # masked).
+ null_mask_data = np.ones((batch_size, 4, 2))
+ unmasked_output_data = model.predict([from_data, to_data, null_mask_data])
+
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ # Tests the layer with three inputs: Q, K, V.
+ key = tf_keras.Input(shape=(2, 8))
+ output = test_layer(
+ query, value=value, key=key, attention_mask=mask_tensor
+ )
+ model = tf_keras.Model([query, value, key, mask_tensor], output)
+
+ masked_output_data = model.predict(
+ [from_data, to_data, to_data, mask_data]
+ )
+ unmasked_output_data = model.predict(
+ [from_data, to_data, to_data, null_mask_data]
+ )
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ if use_bias:
+ self.assertLen(test_layer._query_dense.trainable_variables, 2)
+ self.assertLen(test_layer._output_dense.trainable_variables, 2)
+ else:
+ self.assertLen(test_layer._query_dense.trainable_variables, 1)
+ self.assertLen(test_layer._output_dense.trainable_variables, 1)
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 2),
+ )
+ def test_masked_attention_with_scores(self, num_kv_heads):
+ """Test with a mask tensor."""
+ test_layer = multi_query_attention.MultiHeadAttention(
+ num_heads=4, num_kv_heads=num_kv_heads, key_dim=2
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ batch_size = 3
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
+ output = test_layer(query=query, value=value, attention_mask=mask_tensor)
+
+ # Create a model containing the test layer.
+ model = tf_keras.Model([query, value, mask_tensor], output)
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = 10 * np.random.random_sample((batch_size, 4, 8))
+ to_data = 10 * np.random.random_sample((batch_size, 2, 8))
+
+ # Invoke the data with a random set of mask data. This should mask at
+ # least one element.
+ mask_data = np.random.randint(2, size=(batch_size, 4, 2))
+ masked_output_data = model.predict([from_data, to_data, mask_data])
+
+ # Invoke the same data, but with a null mask (where no elements are
+ # masked).
+ null_mask_data = np.ones((batch_size, 4, 2))
+ unmasked_output_data = model.predict([from_data, to_data, null_mask_data])
+
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ # Create a model containing attention scores.
+ output, scores = test_layer(
+ query=query,
+ value=value,
+ attention_mask=mask_tensor,
+ return_attention_scores=True,
+ )
+ model = tf_keras.Model([query, value, mask_tensor], [output, scores])
+ masked_output_data_score, masked_score = model.predict(
+ [from_data, to_data, mask_data]
+ )
+ unmasked_output_data_score, unmasked_score = model.predict(
+ [from_data, to_data, null_mask_data]
+ )
+ self.assertNotAllClose(masked_output_data_score, unmasked_output_data_score)
+ self.assertAllClose(masked_output_data, masked_output_data_score)
+ self.assertAllClose(unmasked_output_data, unmasked_output_data_score)
+ self.assertNotAllClose(masked_score, unmasked_score)
+
+
+class TalkingHeadsMultiQueryAttentionTest(
+ tf.test.TestCase, parameterized.TestCase
+):
+
+ @parameterized.named_parameters(
+ ("key_value_same_proj_mqa", 1, None, None, [40, 80]),
+ ("key_value_different_proj_mqa", 1, 32, 60, [40, 60]),
+ ("key_value_same_proj_gqa", 3, None, None, [40, 80]),
+ ("key_value_different_proj_gqa", 3, 32, 60, [40, 60]),
+ )
+ def test_non_masked_attention(
+ self, num_kv_heads, value_dim, output_shape, output_dims
+ ):
+ """Test that the attention layer can be created without a mask tensor."""
+ test_layer = multi_query_attention.TalkingHeadsMultiQueryAttention(
+ num_heads=12,
+ num_kv_heads=num_kv_heads,
+ enable_gqa_optimization=num_kv_heads > 1,
+ key_dim=64,
+ value_dim=value_dim,
+ output_shape=output_shape,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(20, 80))
+ output = test_layer(query=query, value=value)
+ self.assertEqual(output.shape.as_list(), [None] + output_dims)
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 3),
+ )
+ def test_non_masked_self_attention(self, num_kv_heads):
+ """Test with one input (self-attenntion) and no mask tensor."""
+ test_layer = multi_query_attention.TalkingHeadsMultiQueryAttention(
+ num_heads=12,
+ num_kv_heads=num_kv_heads,
+ enable_gqa_optimization=num_kv_heads > 1,
+ key_dim=64,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ output = test_layer(query, query)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 3),
+ )
+ def test_attention_scores(self, num_kv_heads):
+ """Test attention outputs with coefficients."""
+ test_layer = multi_query_attention.TalkingHeadsMultiQueryAttention(
+ num_heads=12, num_kv_heads=num_kv_heads, key_dim=64
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ output, coef = test_layer(query, query, return_attention_scores=True)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+ self.assertEqual(coef.shape.as_list(), [None, 12, 40, 40])
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 3),
+ )
+ def test_attention_scores_with_values(self, num_kv_heads):
+ """Test attention outputs with coefficients."""
+ test_layer = multi_query_attention.TalkingHeadsMultiQueryAttention(
+ num_heads=12, num_kv_heads=num_kv_heads, key_dim=64
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(60, 80))
+ output, coef = test_layer(query, value, return_attention_scores=True)
+ self.assertEqual(output.shape.as_list(), [None, 40, 80])
+ self.assertEqual(coef.shape.as_list(), [None, 12, 40, 60])
+
+ @parameterized.named_parameters(
+ ("with_bias_mqa", 1, True),
+ ("no_bias_mqa", 1, False),
+ ("with_bias_gqa", 2, True),
+ ("no_bias_gqa", 2, False),
+ )
+ def test_masked_attention(self, num_kv_heads, use_bias):
+ """Test with a mask tensor."""
+ test_layer = multi_query_attention.TalkingHeadsMultiQueryAttention(
+ num_heads=4,
+ num_kv_heads=num_kv_heads,
+ enable_gqa_optimization=num_kv_heads > 1,
+ key_dim=2,
+ use_bias=use_bias,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ batch_size = 3
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
+ output = test_layer(query=query, value=value, attention_mask=mask_tensor)
+
+ # Create a model containing the test layer.
+ model = tf_keras.Model([query, value, mask_tensor], output)
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = 10 * np.random.random_sample((batch_size, 4, 8))
+ to_data = 10 * np.random.random_sample((batch_size, 2, 8))
+
+ # Invoke the data with a random set of mask data. This should mask at
+ # least one element.
+ mask_data = np.random.randint(2, size=(batch_size, 4, 2))
+ masked_output_data = model.predict([from_data, to_data, mask_data])
+
+ # Invoke the same data, but with a null mask (where no elements are
+ # masked).
+ null_mask_data = np.ones((batch_size, 4, 2))
+ unmasked_output_data = model.predict([from_data, to_data, null_mask_data])
+
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ # Tests the layer with three inputs: Q, K, V.
+ key = tf_keras.Input(shape=(2, 8))
+ output = test_layer(
+ query, value=value, key=key, attention_mask=mask_tensor
+ )
+ model = tf_keras.Model([query, value, key, mask_tensor], output)
+
+ masked_output_data = model.predict(
+ [from_data, to_data, to_data, mask_data]
+ )
+ unmasked_output_data = model.predict(
+ [from_data, to_data, to_data, null_mask_data]
+ )
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ if use_bias:
+ self.assertLen(test_layer._query_dense.trainable_variables, 2)
+ self.assertLen(test_layer._output_dense.trainable_variables, 2)
+ else:
+ self.assertLen(test_layer._query_dense.trainable_variables, 1)
+ self.assertLen(test_layer._output_dense.trainable_variables, 1)
+
+ @parameterized.named_parameters(
+ ("_mqa", 1),
+ ("_gqa", 2),
+ )
+ def test_masked_attention_with_scores(self, num_kv_heads):
+ """Test with a mask tensor."""
+ test_layer = multi_query_attention.TalkingHeadsMultiQueryAttention(
+ num_heads=4, num_kv_heads=num_kv_heads, key_dim=2
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ batch_size = 3
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
+ output = test_layer(query=query, value=value, attention_mask=mask_tensor)
+
+ # Create a model containing the test layer.
+ model = tf_keras.Model([query, value, mask_tensor], output)
+
+ # Generate data for the input (non-mask) tensors.
+ from_data = 10 * np.random.random_sample((batch_size, 4, 8))
+ to_data = 10 * np.random.random_sample((batch_size, 2, 8))
+
+ # Invoke the data with a random set of mask data. This should mask at
+ # least one element.
+ mask_data = np.random.randint(2, size=(batch_size, 4, 2))
+ masked_output_data = model.predict([from_data, to_data, mask_data])
+
+ # Invoke the same data, but with a null mask (where no elements are
+ # masked).
+ null_mask_data = np.ones((batch_size, 4, 2))
+ unmasked_output_data = model.predict([from_data, to_data, null_mask_data])
+
+ # Because one data is masked and one is not, the outputs should not be
+ # the same.
+ self.assertNotAllClose(masked_output_data, unmasked_output_data)
+
+ # Create a model containing attention scores.
+ output, scores = test_layer(
+ query=query,
+ value=value,
+ attention_mask=mask_tensor,
+ return_attention_scores=True,
+ )
+ model = tf_keras.Model([query, value, mask_tensor], [output, scores])
+ masked_output_data_score, masked_score = model.predict(
+ [from_data, to_data, mask_data]
+ )
+ unmasked_output_data_score, unmasked_score = model.predict(
+ [from_data, to_data, null_mask_data]
+ )
+ self.assertNotAllClose(masked_output_data_score, unmasked_output_data_score)
+ self.assertAllClose(masked_output_data, masked_output_data_score)
+ self.assertAllClose(unmasked_output_data, unmasked_output_data_score)
+ self.assertNotAllClose(masked_score, unmasked_score)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/nlp/modeling/layers/on_device_embedding.py b/official/nlp/modeling/layers/on_device_embedding.py
index be000427f4e..13c48d39eff 100644
--- a/official/nlp/modeling/layers/on_device_embedding.py
+++ b/official/nlp/modeling/layers/on_device_embedding.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,11 +15,11 @@
"""Keras-based one-hot embedding layer."""
# pylint: disable=g-classes-have-attributes
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package="Text")
-class OnDeviceEmbedding(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class OnDeviceEmbedding(tf_keras.layers.Layer):
"""Performs an embedding lookup suitable for accelerator devices.
This layer uses either tf.gather or tf.one_hot to translate integer indices to
@@ -37,6 +37,9 @@ class OnDeviceEmbedding(tf.keras.layers.Layer):
scale_factor: Whether to scale the output embeddings. Defaults to None (that
is, not to scale). Setting this option to a float will let values in
output embeddings multiplied by scale_factor.
+ weight_fallback_dtype: When keras mix precision inferred wrong dtype for
+ variables, `weight_fallback_dtype` will be used to define the dtype of
+ weights.
"""
def __init__(self,
@@ -45,14 +48,19 @@ def __init__(self,
initializer="glorot_uniform",
use_one_hot=False,
scale_factor=None,
+ weight_fallback_dtype=tf.float32,
**kwargs):
- super(OnDeviceEmbedding, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._vocab_size = vocab_size
self._embedding_width = embedding_width
self._initializer = initializer
self._use_one_hot = use_one_hot
self._scale_factor = scale_factor
+ # Backup control of the weight dtype because Keras mix precision sometimes
+ # depends on the input to infer the compute dtype, but the inputs of
+ # this layer are int type.
+ self._weight_fallback_dtype = weight_fallback_dtype
def get_config(self):
config = {
@@ -61,28 +69,37 @@ def get_config(self):
"initializer": self._initializer,
"use_one_hot": self._use_one_hot,
"scale_factor": self._scale_factor,
+ "weight_fallback_dtype": self._weight_fallback_dtype,
}
- base_config = super(OnDeviceEmbedding, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def build(self, input_shape):
+ if (
+ self.dtype is not None
+ and not tf.dtypes.as_dtype(self.dtype).is_floating
+ ):
+ # Keras failed to infer the right dtype.
+ dtype = self._weight_fallback_dtype
+ else:
+ dtype = self.dtype
self.embeddings = self.add_weight(
"embeddings",
shape=[self._vocab_size, self._embedding_width],
initializer=self._initializer,
- dtype=tf.float32)
+ dtype=dtype)
- super(OnDeviceEmbedding, self).build(input_shape)
+ super().build(input_shape)
def call(self, inputs):
flat_inputs = tf.reshape(inputs, [-1])
if self._use_one_hot:
- dtype = self._compute_dtype
+ dtype = self.compute_dtype
if not tf.dtypes.as_dtype(dtype).is_floating:
- # TensorFlow 1 compatibility. In TF1, self._compute_dtype is int32
+ # TensorFlow 1 compatibility. In TF1, self.compute_dtype is int32
# instead of a floating-point dtype, as the dtype is inferred from the
# dtype of the inputs
- dtype = tf.float32
+ dtype = self._weight_fallback_dtype
one_hot_data = tf.one_hot(
flat_inputs, depth=self._vocab_size, dtype=dtype)
embeddings = tf.matmul(one_hot_data, self.embeddings)
@@ -90,7 +107,6 @@ def call(self, inputs):
embeddings = tf.gather(self.embeddings, flat_inputs)
embeddings = tf.reshape(
embeddings,
- # Work around b/142213824: prefer concat to shape over a Python list.
tf.concat([tf.shape(inputs), [self._embedding_width]], axis=0))
embeddings.set_shape(inputs.shape.as_list() + [self._embedding_width])
if self._scale_factor:
diff --git a/official/nlp/modeling/layers/on_device_embedding_test.py b/official/nlp/modeling/layers/on_device_embedding_test.py
index 373cfdb6d3d..e19ca0ca5a9 100644
--- a/official/nlp/modeling/layers/on_device_embedding_test.py
+++ b/official/nlp/modeling/layers/on_device_embedding_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,16 +15,12 @@
"""Tests for Keras-based one-hot embedding layer."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import on_device_embedding
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class OnDeviceEmbeddingTest(keras_parameterized.TestCase):
+class OnDeviceEmbeddingTest(tf.test.TestCase):
def test_layer_creation(self):
vocab_size = 31
@@ -33,7 +29,7 @@ def test_layer_creation(self):
vocab_size=vocab_size, embedding_width=embedding_width)
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# The output should be the same as the input, save that it has an extra
@@ -50,7 +46,7 @@ def test_layer_creation_with_mixed_precision(self):
dtype="mixed_float16")
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# The output should be the same as the input, save that it has an extra
@@ -66,11 +62,11 @@ def test_layer_invocation(self):
vocab_size=vocab_size, embedding_width=embedding_width)
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -88,11 +84,11 @@ def test_layer_invocation_with_mixed_precision(self):
dtype="mixed_float16")
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -111,7 +107,7 @@ def test_one_hot_layer_creation(self):
use_one_hot=True)
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# The output should be the same as the input, save that it has an extra
@@ -130,7 +126,7 @@ def test_one_hot_layer_creation_with_mixed_precision(self):
use_one_hot=True)
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# The output should be the same as the input, save that it has an extra
@@ -148,11 +144,11 @@ def test_one_hot_layer_invocation(self):
use_one_hot=True)
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -172,11 +168,11 @@ def test_one_hot_layer_invocation_with_mixed_precision(self):
use_one_hot=True)
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -194,11 +190,11 @@ def test_use_scale_layer_invocation(self):
scale_factor=embedding_width**0.5)
# Create a 2-dimensional input (the first dimension is implicit).
sequence_length = 23
- input_tensor = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
+ input_tensor = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
output_tensor = test_layer(input_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
diff --git a/official/nlp/modeling/layers/pack_optimization.py b/official/nlp/modeling/layers/pack_optimization.py
new file mode 100644
index 00000000000..bb58f7b59e8
--- /dev/null
+++ b/official/nlp/modeling/layers/pack_optimization.py
@@ -0,0 +1,258 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pack sequence optimization on accelerators."""
+from typing import Dict
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
+from official.nlp.modeling.layers import rezero_transformer
+from official.nlp.modeling.layers import self_attention_mask
+from official.nlp.modeling.layers import transformer_encoder_block
+from official.nlp.modeling.layers import transformer_scaffold
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class PackBertEmbeddings(tf_keras.layers.Layer):
+ """Performs packing tricks for BERT inputs to improve TPU utilization."""
+
+ def __init__(self, pack_sequences: int, **kwargs):
+ super().__init__(**kwargs)
+ self.pack_sequences = pack_sequences
+
+ def call(self, input_embeddings: tf.Tensor,
+ input_mask: tf.Tensor) -> Dict[str, tf.Tensor]:
+ batch_size, seq_len, embedding_dim = tf_utils.get_shape_list(
+ input_embeddings, expected_rank=3)
+ reduced_batch_size = batch_size // self.pack_sequences
+ packed_seq_len = self.pack_sequences * seq_len
+ packed_embeddings = tf.reshape(
+ input_embeddings, [reduced_batch_size, packed_seq_len, embedding_dim])
+ input_mask = tf.reshape(input_mask, [reduced_batch_size, packed_seq_len])
+ example_ids = 1 + tf.range(self.pack_sequences)
+ # Shape: [batch_size, seq_len, pack_sequences].
+ example_ids = tf.tile(example_ids[None, :, None],
+ [reduced_batch_size, 1, seq_len])
+ example_ids = tf.reshape(example_ids, [reduced_batch_size, packed_seq_len])
+ example_ids = tf.where(
+ tf.math.equal(input_mask, 0), tf.zeros_like(example_ids), example_ids)
+ packing_mask = tf.cast(
+ tf.equal(
+ tf.expand_dims(example_ids, 2), tf.expand_dims(example_ids, 1)),
+ dtype=tf.bool)
+
+ attention_mask = self_attention_mask.get_mask(
+ packed_embeddings, input_mask, dtype=tf.bool)
+
+ combined_attention_mask = tf.cast(
+ tf.math.logical_and(attention_mask, packing_mask), tf.float32)
+
+ return dict(
+ packed_embeddings=packed_embeddings,
+ combined_attention_mask=combined_attention_mask)
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class StridedTransformerEncoderBlock(
+ transformer_encoder_block.TransformerEncoderBlock):
+ """Transformer layer for packing optimization to stride over inputs."""
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ if self._output_range is not None:
+ raise ValueError('StridedTransformerEncoderBlock does not '
+ 'support `output_range` argument.')
+ # TODO(b/337888023): Support block sparse attention with strided inputs.
+ if self._src_block_size is not None:
+ raise ValueError('StridedTransformerEncoderBlock does not '
+ 'support block sparse attention.')
+
+ def call(self, inputs, stride: tf.Tensor):
+ if isinstance(inputs, (list, tuple)):
+ if len(inputs) == 2:
+ input_tensor, attention_mask = inputs
+ key_value = None
+ elif len(inputs) == 3:
+ input_tensor, key_value, attention_mask = inputs
+ else:
+ raise ValueError('Unexpected inputs to %s with length at %d' %
+ (self.__class__, len(inputs)))
+ else:
+ input_tensor, key_value, attention_mask = (inputs, None, None)
+
+ if self._norm_first:
+ source_tensor = input_tensor[:, ::stride, :]
+ input_tensor = self._attention_layer_norm(input_tensor)
+ if key_value is not None:
+ key_value = self._attention_layer_norm_kv(key_value)
+ target_tensor = input_tensor[:, ::stride, :]
+ if attention_mask is not None:
+ attention_mask = attention_mask[:, ::stride, :]
+
+ if key_value is None:
+ key_value = input_tensor
+ attention_output = self._attention_layer(
+ query=target_tensor, value=key_value, attention_mask=attention_mask)
+ attention_output = self._attention_dropout(attention_output)
+
+ if self._norm_first:
+ # Important to not combine `self._norm_first` and
+ # `self._use_query_residual` into one if clause because else is only for
+ # `_norm_first == False`.
+ if self._use_query_residual:
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
+ else:
+ if self._use_query_residual:
+ attention_output = target_tensor + attention_output
+ attention_output = self._attention_layer_norm(attention_output)
+
+ if self._norm_first:
+ source_attention_output = attention_output
+ attention_output = self._output_layer_norm(attention_output)
+ inner_output = self._intermediate_dense(attention_output)
+ inner_output = self._intermediate_activation_layer(inner_output)
+ inner_output = self._inner_dropout_layer(inner_output)
+ layer_output = self._output_dense(inner_output)
+ layer_output = self._output_dropout(layer_output)
+
+ if self._norm_first:
+ return source_attention_output + layer_output # pyrefly: ignore[unbound-name]
+
+ layer_output = tf.cast(layer_output, tf.float32)
+ return self._output_layer_norm(layer_output + attention_output)
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class StridedReZeroTransformer(rezero_transformer.ReZeroTransformer):
+ """ReZeroTransformer for packing optimization to stride over inputs."""
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ if self._output_range is not None:
+ raise ValueError(f'{self.__class__} does not '
+ 'support `output_range` argument.')
+ # TODO(b/337888023): Support block sparse attention with strided inputs.
+ if self._src_block_size is not None:
+ raise ValueError(f'{self.__class__} does not '
+ 'support block sparse attention.')
+
+ def call(self, inputs, stride: tf.Tensor):
+ if isinstance(inputs, (list, tuple)):
+ if len(inputs) == 2:
+ input_tensor, attention_mask = inputs
+ key_value = None
+ elif len(inputs) == 3:
+ input_tensor, key_value, attention_mask = inputs
+ else:
+ raise ValueError(f'Unexpected inputs to {self.__class__} with '
+ f'length at {len(inputs)}.')
+ else:
+ input_tensor, key_value, attention_mask = (inputs, None, None)
+
+ target_tensor = input_tensor[:, ::stride, :]
+ if attention_mask is not None:
+ attention_mask = attention_mask[:, ::stride, :]
+
+ if key_value is None:
+ key_value = input_tensor
+
+ attention_output = self._attention_layer(
+ query=target_tensor, value=key_value, attention_mask=attention_mask)
+ attention_output = self._attention_dropout(attention_output)
+ attention_output = target_tensor + self._rezero_a * attention_output
+ if self._use_layer_norm:
+ attention_output = self._attention_layer_norm(attention_output)
+ else:
+ attention_output = tf.cast(attention_output, tf.float32)
+
+ intermediate_output = self._intermediate_dense(attention_output)
+ intermediate_output = self._inner_activation_layer(intermediate_output)
+ layer_output = self._output_dense(intermediate_output)
+ layer_output = self._output_dropout(layer_output)
+ layer_output = attention_output + tf.cast(self._rezero_a_ffn * layer_output,
+ tf.float32)
+ if self._use_layer_norm:
+ layer_output = self._output_layer_norm(layer_output)
+
+ return layer_output
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class StridedTransformerScaffold(transformer_scaffold.TransformerScaffold):
+ """TransformerScaffold for packing optimization to stride over inputs."""
+
+ def call(self, inputs, stride: tf.Tensor, training=None):
+ if isinstance(inputs, (list, tuple)):
+ if len(inputs) == 2:
+ input_tensor, attention_mask = inputs
+ key_value = None
+ elif len(inputs) == 3:
+ input_tensor, key_value, attention_mask = inputs
+ else:
+ raise ValueError('Unexpected inputs to %s with length at %d' %
+ (self.__class__, len(inputs)))
+ else:
+ input_tensor, key_value, attention_mask = (inputs, None, None)
+
+ if key_value is None:
+ key_value = input_tensor
+
+ if self._norm_first:
+ source_tensor = input_tensor[:, ::stride, :]
+ input_tensor = self._attention_layer_norm(input_tensor, training=training)
+ if attention_mask is not None:
+ attention_mask = attention_mask[:, ::stride, :]
+ target_tensor = input_tensor[:, ::stride, :]
+
+ attention_output = self._attention_layer(
+ query=target_tensor,
+ value=key_value,
+ attention_mask=attention_mask,
+ training=training)
+ attention_output = self._attention_dropout(
+ attention_output, training=training)
+
+ if self._norm_first:
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
+ else:
+ attention_output = self._attention_layer_norm(
+ target_tensor + attention_output, training=training)
+ if self._norm_first:
+ source_attention_output = attention_output
+ attention_output = self._output_layer_norm(
+ attention_output, training=training)
+
+ if self._feedforward_block is None:
+ intermediate_output = self._intermediate_dense(attention_output)
+ intermediate_output = self._intermediate_activation_layer(
+ intermediate_output)
+ layer_output = self._output_dense(intermediate_output, training=training)
+ layer_output = self._output_dropout(layer_output, training=training)
+ layer_output = tf.cast(layer_output, tf.float32)
+ if self._norm_first:
+ layer_output = source_attention_output + layer_output # pyrefly: ignore[unbound-name]
+ else:
+ layer_output = self._output_layer_norm(
+ layer_output + attention_output, training=training)
+ else:
+ if self._norm_first:
+ # if norm_first, assume the feedforward block will not apply layer norm
+ layer_output = self._feedforward_block(
+ attention_output, training=training)
+ layer_output += source_attention_output # pyrefly: ignore[unbound-name]
+ else:
+ # if not norm_first, assume that the feedforwad does apply layer norm
+ layer_output = self._feedforward_block(
+ attention_output, training=training)
+
+ return layer_output
diff --git a/official/nlp/modeling/layers/pack_optimization_test.py b/official/nlp/modeling/layers/pack_optimization_test.py
new file mode 100644
index 00000000000..64fecc7912c
--- /dev/null
+++ b/official/nlp/modeling/layers/pack_optimization_test.py
@@ -0,0 +1,66 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for pack_optimization."""
+
+import tensorflow as tf, tf_keras
+from official.nlp.modeling.layers import pack_optimization
+
+
+class PackOptimizationTest(tf.test.TestCase):
+
+ def test_bert_embedding_packing(self):
+ batch_size, seq_len, embed_dim = 2, 4, 8
+ pack_sequences = 2
+ token_and_position_embed = tf.ones((batch_size, seq_len, embed_dim),
+ dtype=tf.float32)
+ input_mask = tf.ones((batch_size, seq_len), dtype=tf.int32)
+
+ layer = pack_optimization.PackBertEmbeddings(pack_sequences=pack_sequences)
+ outputs = layer(token_and_position_embed, input_mask)
+ self.assertEqual(outputs["packed_embeddings"].shape, (1, 8, embed_dim))
+ self.assertEqual(outputs["combined_attention_mask"].shape, (1, 8, 8))
+
+ def test_strided_transformer_encoder_block(self):
+ inputs = tf.zeros((2, 4, 8), dtype=tf.float32)
+ attention_mask = tf.ones((2, 4, 4), dtype=tf.float32)
+ transformer = pack_optimization.StridedTransformerEncoderBlock(
+ num_attention_heads=2, inner_dim=4, inner_activation="relu")
+ outputs = transformer([inputs, attention_mask],
+ stride=tf.constant(2, dtype=tf.int32))
+ self.assertEqual(outputs.shape, (2, 2, 8))
+
+ def test_strided_rezero_transformer(self):
+ inputs = tf.zeros((2, 4, 8), dtype=tf.float32)
+ attention_mask = tf.ones((2, 4, 4), dtype=tf.float32)
+ transformer = pack_optimization.StridedReZeroTransformer(
+ num_attention_heads=2, inner_dim=4, inner_activation="relu")
+ outputs = transformer([inputs, attention_mask],
+ stride=tf.constant(2, dtype=tf.int32))
+ self.assertEqual(outputs.shape, (2, 2, 8))
+
+ def test_strided_scaffold(self):
+ inputs = tf.zeros((2, 4, 8), dtype=tf.float32)
+ attention_mask = tf.ones((2, 4, 4), dtype=tf.float32)
+ test_layer = pack_optimization.StridedTransformerScaffold(
+ num_attention_heads=2,
+ inner_dim=128,
+ inner_activation="relu")
+ outputs = test_layer([inputs, attention_mask],
+ stride=tf.constant(2, dtype=tf.int32))
+ self.assertEqual(outputs.shape, (2, 2, 8))
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/nlp/modeling/layers/per_dim_scale_attention.py b/official/nlp/modeling/layers/per_dim_scale_attention.py
new file mode 100644
index 00000000000..7a3c3c9ee20
--- /dev/null
+++ b/official/nlp/modeling/layers/per_dim_scale_attention.py
@@ -0,0 +1,101 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras-based attention layer with learnable per dim scaling."""
+import gin
+import numpy as np
+import tensorflow as tf, tf_keras
+
+
+@gin.configurable
+@tf_keras.utils.register_keras_serializable(package='Text')
+class PerDimScaleAttention(tf_keras.layers.MultiHeadAttention):
+ """Learn scales for individual dims.
+
+ It can improve quality but might hurt training stability.
+ """
+
+ def _build_from_signature(self, query, value, key=None):
+ super()._build_from_signature(query=query, value=value, key=key) # pytype: disable=attribute-error
+ self._scale_dim = self._key_dim
+ with tf.init_scope():
+ self.per_dim_scale = self.add_weight(
+ name='per_dim_scale',
+ shape=(self._scale_dim,),
+ initializer='zeros',
+ dtype=self.dtype,
+ trainable=True)
+
+ def _scale_query(self, query):
+ # 1.0/tf.nn.softplus(0.0) = 1.442695041. Hard code this number so that we
+ # can avoid unnecessary XLA op fusion mess on TPU.
+ r_softplus_0 = 1.442695041
+ scale = tf.constant(
+ r_softplus_0 / np.sqrt(float(self._scale_dim)), dtype=query.dtype)
+
+ scale *= tf.nn.softplus(self.per_dim_scale)
+ return query * scale
+
+ def _compute_attention(self,
+ query,
+ key,
+ value,
+ attention_mask=None,
+ training=None):
+ query = self._scale_query(query)
+
+ attention_scores = tf.einsum(self._dot_product_equation, key, query)
+
+ attention_scores = self._masked_softmax(attention_scores, attention_mask)
+
+ attention_scores_dropout = self._dropout_layer(
+ attention_scores, training=training)
+
+ # `context_layer` = [B, T, N, H]
+ attention_output = tf.einsum(self._combine_equation,
+ attention_scores_dropout, value)
+ return attention_output, attention_scores
+
+ def call( # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ self,
+ query,
+ value,
+ key=None,
+ attention_mask=None,
+ return_attention_scores=False,
+ training=None,
+ ):
+ if not self._built_from_signature:
+ self._build_from_signature(query=query, value=value, key=key)
+ if key is None:
+ key = value
+
+ # N = `num_attention_heads`
+ # H = `size_per_head`
+ # `query` = [B, T, N ,H]
+ query = self._query_dense(query)
+
+ # `key` = [B, S, N, H]
+ key = self._key_dense(key)
+
+ # `value` = [B, S, N, H]
+ value = self._value_dense(value)
+
+ attention_output, attention_scores = self._compute_attention(
+ query, key, value, attention_mask, training)
+ attention_output = self._output_dense(attention_output)
+
+ if return_attention_scores:
+ return attention_output, attention_scores
+ return attention_output
diff --git a/official/nlp/modeling/layers/per_dim_scale_attention_test.py b/official/nlp/modeling/layers/per_dim_scale_attention_test.py
new file mode 100644
index 00000000000..4bb0accab17
--- /dev/null
+++ b/official/nlp/modeling/layers/per_dim_scale_attention_test.py
@@ -0,0 +1,52 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for PerDimScaleAttention."""
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling.layers import per_dim_scale_attention as attention
+
+
+class PerDimScaleAttentionTest(tf.test.TestCase):
+
+ def test_attention(self):
+ num_heads = 12
+ key_dim = 64
+ seq_length = 1024
+ batch_size = 2
+ test_layer = attention.PerDimScaleAttention(
+ num_heads=num_heads, key_dim=key_dim)
+ query = tf.random.normal(
+ shape=(batch_size, seq_length, key_dim * num_heads))
+ value = query
+ output = test_layer(query=query, value=value)
+ self.assertEqual(output.shape,
+ [batch_size, seq_length, key_dim * num_heads])
+
+ def test_config(self):
+ num_heads = 12
+ key_dim = 64
+ test_layer = attention.PerDimScaleAttention(
+ num_heads=num_heads, key_dim=key_dim)
+ print(test_layer.get_config())
+ new_layer = attention.PerDimScaleAttention.from_config(
+ test_layer.get_config())
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(test_layer.get_config(), new_layer.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/nlp/modeling/layers/position_embedding.py b/official/nlp/modeling/layers/position_embedding.py
index 86ee2fc6e99..4ba87ec21e3 100644
--- a/official/nlp/modeling/layers/position_embedding.py
+++ b/official/nlp/modeling/layers/position_embedding.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,21 +17,21 @@
import math
from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
-Initializer = tf.keras.initializers.Initializer
+Initializer = tf_keras.initializers.Initializer
-@tf.keras.utils.register_keras_serializable(package="Text")
-class PositionEmbedding(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class PositionEmbedding(tf_keras.layers.Layer):
"""Creates a positional embedding.
Example:
```python
position_embedding = PositionEmbedding(max_length=100)
- inputs = tf.keras.Input((100, 32), dtype=tf.float32)
+ inputs = tf_keras.Input((100, 32), dtype=tf.float32)
outputs = position_embedding(inputs)
```
@@ -53,26 +53,26 @@ def __init__(self,
seq_axis=1,
**kwargs):
- super(PositionEmbedding, self).__init__(**kwargs)
+ super().__init__(**kwargs)
if max_length is None:
raise ValueError(
"`max_length` must be an Integer, not `None`."
)
self._max_length = max_length
- self._initializer = tf.keras.initializers.get(initializer)
+ self._initializer = tf_keras.initializers.get(initializer)
self._seq_axis = seq_axis
def get_config(self):
config = {
"max_length": self._max_length,
- "initializer": tf.keras.initializers.serialize(self._initializer),
+ "initializer": tf_keras.initializers.serialize(self._initializer),
"seq_axis": self._seq_axis,
}
base_config = super(PositionEmbedding, self).get_config()
return dict(list(base_config.items()) + list(config.items()))
def build(self, input_shape):
- dimension_list = input_shape.as_list()
+ dimension_list = input_shape
width = dimension_list[-1]
weight_sequence_length = self._max_length
@@ -81,7 +81,7 @@ def build(self, input_shape):
shape=[weight_sequence_length, width],
initializer=self._initializer)
- super(PositionEmbedding, self).build(input_shape)
+ super().build(input_shape)
def call(self, inputs):
input_shape = tf.shape(inputs)
@@ -94,8 +94,8 @@ def call(self, inputs):
return tf.broadcast_to(position_embeddings, input_shape)
-@tf.keras.utils.register_keras_serializable(package="Text")
-class RelativePositionEmbedding(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class RelativePositionEmbedding(tf_keras.layers.Layer):
"""Creates a positional embedding.
This layer calculates the position encoding as a mix of sine and cosine
@@ -227,8 +227,8 @@ def _relative_position_bucket(relative_position,
return ret
-@tf.keras.utils.register_keras_serializable(package="Text")
-class RelativePositionBias(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class RelativePositionBias(tf_keras.layers.Layer):
"""Relative position embedding via per-head bias in T5 style.
Reference implementation in MeshTF:
@@ -254,7 +254,7 @@ def __init__(self,
if embeddings_initializer:
self._embed_init = embeddings_initializer
else:
- self._embed_init = tf.keras.initializers.TruncatedNormal(stddev=1.0)
+ self._embed_init = tf_keras.initializers.TruncatedNormal(stddev=1.0)
with tf.name_scope(self.name):
self._relative_attention_bias = self.add_weight(
"rel_embedding",
@@ -274,7 +274,7 @@ def get_config(self):
"bidirectional":
self.bidirectional,
"embeddings_initializer":
- tf.keras.initializers.serialize(self._embed_init),
+ tf_keras.initializers.serialize(self._embed_init),
}
base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
diff --git a/official/nlp/modeling/layers/position_embedding_test.py b/official/nlp/modeling/layers/position_embedding_test.py
index f9f17085421..809a04d71fb 100644
--- a/official/nlp/modeling/layers/position_embedding_test.py
+++ b/official/nlp/modeling/layers/position_embedding_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,16 +16,12 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import position_embedding
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class PositionEmbeddingLayerTest(keras_parameterized.TestCase):
+class PositionEmbeddingLayerTest(tf.test.TestCase):
def test_static_layer_output_shape(self):
# Create a 3-dimensional input (the first dimension is implicit).
@@ -33,7 +29,7 @@ def test_static_layer_output_shape(self):
test_layer = position_embedding.PositionEmbedding(
max_length=sequence_length)
width = 30
- input_tensor = tf.keras.Input(shape=(sequence_length, width))
+ input_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(input_tensor)
# When using static positional embedding shapes, the output is expected
@@ -49,7 +45,7 @@ def test_non_default_axis_static(self):
test_layer = position_embedding.PositionEmbedding(
max_length=sequence_length, seq_axis=2)
width = 30
- input_tensor = tf.keras.Input(shape=(width, sequence_length, width))
+ input_tensor = tf_keras.Input(shape=(width, sequence_length, width))
output_tensor = test_layer(input_tensor)
# When using static positional embedding shapes, the output is expected
@@ -65,7 +61,7 @@ def test_float16_dtype(self):
test_layer = position_embedding.PositionEmbedding(
max_length=sequence_length, dtype="float16")
width = 30
- input_tensor = tf.keras.Input(shape=(sequence_length, width))
+ input_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(input_tensor)
# When using static positional embedding shapes, the output is expected
@@ -81,7 +77,7 @@ def test_dynamic_layer_output_shape(self):
max_length=max_sequence_length)
# Create a 3-dimensional input (the first dimension is implicit).
width = 30
- input_tensor = tf.keras.Input(shape=(None, width))
+ input_tensor = tf_keras.Input(shape=(None, width))
output_tensor = test_layer(input_tensor)
# When using dynamic positional embedding shapes, the output is expected
@@ -96,7 +92,7 @@ def test_non_default_axis_dynamic(self):
max_length=max_sequence_length, seq_axis=2)
# Create a 3-dimensional input (the first dimension is implicit).
width = 30
- input_tensor = tf.keras.Input(shape=(None, None, width))
+ input_tensor = tf_keras.Input(shape=(None, None, width))
output_tensor = test_layer(input_tensor)
# When using dynamic positional embedding shapes, the output is expected
@@ -111,10 +107,10 @@ def test_dynamic_layer_slicing(self):
max_length=max_sequence_length)
# Create a 3-dimensional input (the first dimension is implicit).
width = 30
- input_tensor = tf.keras.Input(shape=(None, width))
+ input_tensor = tf_keras.Input(shape=(None, width))
output_tensor = test_layer(input_tensor)
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
# Create input data that is shorter than max_sequence_length, which should
# trigger a down-slice.
@@ -129,10 +125,7 @@ def test_dynamic_layer_slicing(self):
self.assertAllEqual([1, input_length, width], output_data.shape)
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class RelativePositionEmbeddingLayerTest(keras_parameterized.TestCase):
+class RelativePositionEmbeddingLayerTest(tf.test.TestCase):
def test_relative_tensor_input(self):
hidden_size = 8
@@ -164,8 +157,7 @@ def test_relative_length_input(self):
self.assertAllEqual(output_tensor, expected_output_tensor)
-@keras_parameterized.run_all_keras_modes
-class RelativePositionBiasTest(keras_parameterized.TestCase):
+class RelativePositionBiasTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.named_parameters(("bidirectional", True),
("unidirectional", False))
diff --git a/official/nlp/modeling/layers/relative_attention.py b/official/nlp/modeling/layers/relative_attention.py
index 7312cbf8614..6e357da6289 100644
--- a/official/nlp/modeling/layers/relative_attention.py
+++ b/official/nlp/modeling/layers/relative_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Keras-based relative attention layers."""
import math
import string
-import tensorflow as tf
+import tensorflow as tf, tf_keras
_CHR_IDX = string.ascii_lowercase
@@ -69,12 +69,12 @@ def _rel_shift(x, klen=-1):
return x
-@tf.keras.utils.register_keras_serializable(package="Text")
-class MultiHeadRelativeAttention(tf.keras.layers.MultiHeadAttention):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class MultiHeadRelativeAttention(tf_keras.layers.MultiHeadAttention):
"""A multi-head attention layer with relative attention + position encoding.
This layer shares the same input/output projections as the common
- `tf.keras.layers.MultiHeadAttention` layer.
+ `tf_keras.layers.MultiHeadAttention` layer.
When it calculates attention logits, position encoding is projected to form
relative keys. The logits are composed by shifted relative logits and content
@@ -98,14 +98,14 @@ class MultiHeadRelativeAttention(tf.keras.layers.MultiHeadAttention):
`[B, L, dim]`.
segment_matrix: Optional `Tensor` representing segmentation IDs used in
XLNet of shape `[B, S, S + M]`.
- segment_encoding: Optional `Tensor` representing the segmentation
- encoding as used in XLNet of shape `[2, num_heads, dim]`.
- segment_attention_bias: Optional trainable bias parameter added to the
- query had when calculating the segment-based attention score used in
- XLNet of shape `[num_heads, dim]`.
+ segment_encoding: Optional `Tensor` representing the segmentation encoding
+ as used in XLNet of shape `[2, num_heads, dim]`.
+ segment_attention_bias: Optional trainable bias parameter added to the query
+ had when calculating the segment-based attention score used in XLNet of
+ shape `[num_heads, dim]`.
state: Optional `Tensor` of shape `[B, M, E]` where M is the length of the
- state or memory.
- If passed, this is also attended over as in Transformer XL.
+ state or memory. If passed, this is also attended over as in Transformer
+ XL.
attention_mask: A boolean mask of shape `[B, T, S]` that prevents attention
to certain positions.
"""
@@ -144,7 +144,7 @@ def _build_from_signature(self, query, value, key=None):
with tf.init_scope():
einsum_equation, _, output_rank = _build_proj_equation(
key_shape.rank - 1, bound_dims=1, output_dims=2)
- self._encoding_dense = tf.keras.layers.experimental.EinsumDense(
+ self._encoding_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1,
[self._num_heads, self._key_dim]),
@@ -228,7 +228,7 @@ def compute_attention(self,
value)
return attention_output
- def call(self,
+ def call(self, # pytype: disable=signature-mismatch # overriding-parameter-count-checks
query,
value,
content_attention_bias,
@@ -255,8 +255,8 @@ def call(self,
Args:
query: attention input.
value: attention input.
- content_attention_bias: A trainable bias parameter added to the query
- head when calculating the content-based attention score.
+ content_attention_bias: A trainable bias parameter added to the query head
+ when calculating the content-based attention score.
positional_attention_bias: A trainable bias parameter added to the query
head when calculating the position-based attention score.
key: attention input.
@@ -264,8 +264,8 @@ def call(self,
value.
segment_matrix: Optional `Tensor` representing segmentation IDs used in
XLNet.
- segment_encoding: Optional `Tensor` representing the segmentation
- encoding as used in XLNet.
+ segment_encoding: Optional `Tensor` representing the segmentation encoding
+ as used in XLNet.
segment_attention_bias: Optional trainable bias parameter added to the
query had when calculating the segment-based attention score used in
XLNet.
@@ -319,7 +319,7 @@ def call(self,
return attention_output
-@tf.keras.utils.register_keras_serializable(package="Text")
+@tf_keras.utils.register_keras_serializable(package="Text")
class TwoStreamRelativeAttention(MultiHeadRelativeAttention):
"""Two-stream relative self-attention for XLNet.
@@ -333,7 +333,7 @@ class TwoStreamRelativeAttention(MultiHeadRelativeAttention):
but not the content.
This layer shares the same build signature as
- `tf.keras.layers.MultiHeadAttention` but has different input/output
+ `tf_keras.layers.MultiHeadAttention` but has different input/output
projections.
**Note: This layer is currently experimental.
@@ -364,7 +364,7 @@ class TwoStreamRelativeAttention(MultiHeadRelativeAttention):
prevents attention to certain position for query attention computation.
"""
- def call(self,
+ def call(self, # pyrefly: ignore[bad-override]
content_stream,
content_attention_bias,
positional_attention_bias,
@@ -394,22 +394,22 @@ def call(self,
content_stream: The content representation, commonly referred to as h.
This serves a similar role to the standard hidden states in
Transformer-XL.
- content_attention_bias: A trainable bias parameter added to the query
- head when calculating the content-based attention score.
+ content_attention_bias: A trainable bias parameter added to the query head
+ when calculating the content-based attention score.
positional_attention_bias: A trainable bias parameter added to the query
head when calculating the position-based attention score.
- query_stream: The query representation, commonly referred to as g.
- This only has access to contextual information and position, but not
- content. If not provided, then this is MultiHeadRelativeAttention with
+ query_stream: The query representation, commonly referred to as g. This
+ only has access to contextual information and position, but not content.
+ If not provided, then this is MultiHeadRelativeAttention with
self-attention.
relative_position_encoding: relative positional encoding for key and
value.
- target_mapping: Optional `Tensor` representing the target mapping used
- in partial prediction.
+ target_mapping: Optional `Tensor` representing the target mapping used in
+ partial prediction.
segment_matrix: Optional `Tensor` representing segmentation IDs used in
XLNet.
- segment_encoding: Optional `Tensor` representing the segmentation
- encoding as used in XLNet.
+ segment_encoding: Optional `Tensor` representing the segmentation encoding
+ as used in XLNet.
segment_attention_bias: Optional trainable bias parameter added to the
query head when calculating the segment-based attention score.
state: (default None) optional state. If passed, this is also attended
@@ -417,8 +417,8 @@ def call(self,
content_attention_mask: (default None) Optional mask that is added to
content attention logits. If state is not None, the mask source sequence
dimension should extend M.
- query_attention_mask: (default None) Optional mask that is added to
- query attention logits. If state is not None, the mask source sequence
+ query_attention_mask: (default None) Optional mask that is added to query
+ attention logits. If state is not None, the mask source sequence
dimension should extend M.
Returns:
@@ -496,4 +496,3 @@ def call(self,
query_attention_output = self._output_dense(query_attention_output)
return content_attention_output, query_attention_output
-
diff --git a/official/nlp/modeling/layers/relative_attention_test.py b/official/nlp/modeling/layers/relative_attention_test.py
index d07093f72af..6734c845b1d 100644
--- a/official/nlp/modeling/layers/relative_attention_test.py
+++ b/official/nlp/modeling/layers/relative_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,11 @@
"""Tests for the attention layer."""
+from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import relative_attention
@@ -111,8 +111,7 @@ def _create_mock_attention_data(
return data
-@keras_parameterized.run_all_keras_modes
-class MultiHeadRelativeAttentionTest(keras_parameterized.TestCase):
+class MultiHeadRelativeAttentionTest(tf.test.TestCase, parameterized.TestCase):
@combinations.generate(combinations.combine(
value_dim=[32, 64],
@@ -147,8 +146,7 @@ def test_attention_scores(self,
self.assertEqual(output.shape, [batch_size, seq_length, key_dim])
-@keras_parameterized.run_all_keras_modes
-class TwoStreamRelativeAttentionTest(keras_parameterized.TestCase):
+class TwoStreamRelativeAttentionTest(tf.test.TestCase, parameterized.TestCase):
@combinations.generate(combinations.combine(
num_predictions=[2, 10],
diff --git a/official/nlp/modeling/layers/reuse_attention.py b/official/nlp/modeling/layers/reuse_attention.py
index 419cf136052..2ebc4a022c1 100644
--- a/official/nlp/modeling/layers/reuse_attention.py
+++ b/official/nlp/modeling/layers/reuse_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,9 @@
import string
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
_CHR_IDX = string.ascii_lowercase
@@ -107,7 +109,7 @@ def _get_output_shape(output_rank, known_last_dims):
return [None] * (output_rank - len(known_last_dims)) + list(known_last_dims)
-class ReuseMultiHeadAttention(tf.keras.layers.Layer):
+class ReuseMultiHeadAttention(tf_keras.layers.Layer):
"""MultiHeadAttention layer.
This is an implementation of multi-headed attention as described in the paper
@@ -136,8 +138,8 @@ class ReuseMultiHeadAttention(tf.keras.layers.Layer):
Returns the additional attention weights over heads.
>>> layer = MultiHeadAttention(num_heads=2, key_dim=2)
- >>> target = tf.keras.Input(shape=[8, 16])
- >>> source = tf.keras.Input(shape=[4, 16])
+ >>> target = tf_keras.Input(shape=[8, 16])
+ >>> source = tf_keras.Input(shape=[4, 16])
>>> output_tensor, weights = layer(target, source,
... return_attention_scores=True)
>>> print(output_tensor.shape)
@@ -148,7 +150,7 @@ class ReuseMultiHeadAttention(tf.keras.layers.Layer):
Performs 2D self-attention over a 5D input tensor on axes 2 and 3.
>>> layer = MultiHeadAttention(num_heads=2, key_dim=2, attention_axes=(2, 3))
- >>> input_tensor = tf.keras.Input(shape=[5, 3, 4, 16])
+ >>> input_tensor = tf_keras.Input(shape=[5, 3, 4, 16])
>>> output_tensor = layer(input_tensor, input_tensor)
>>> print(output_tensor.shape)
(None, 5, 3, 4, 16)
@@ -221,7 +223,7 @@ def __init__(self,
kernel_constraint=None,
bias_constraint=None,
**kwargs):
- super(ReuseMultiHeadAttention, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._num_heads = num_heads
self._key_dim = key_dim
self._value_dim = value_dim if value_dim else key_dim
@@ -237,12 +239,12 @@ def __init__(self,
self._pe_max_seq_length = pe_max_seq_length
self._use_bias = use_bias
self._output_shape = output_shape
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
if attention_axes is not None and not isinstance(attention_axes,
collections.abc.Sized):
self._attention_axes = (attention_axes,)
@@ -253,7 +255,7 @@ def __init__(self,
# Use relative PE only if reuse_heads < num_heads.
if self._use_relative_pe and self._reuse_heads < self._num_heads:
# Determine the dtype from global policy.
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
policy = tf.bfloat16
elif policy.name == "mixed_float16":
@@ -282,24 +284,24 @@ def get_config(self):
"use_relative_pe": self._use_relative_pe,
"pe_max_seq_length": self._pe_max_seq_length,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
+ tf_keras.constraints.serialize(self._bias_constraint),
"query_shape": self._query_shape,
"key_shape": self._key_shape,
"value_shape": self._value_shape,
}
- base_config = super(ReuseMultiHeadAttention, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
@classmethod
@@ -347,8 +349,6 @@ def _build_from_signature(self, query, value, key=None):
self._key_shape = tf.TensorShape(key)
common_kwargs = dict(
- kernel_initializer=self._kernel_initializer,
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
@@ -362,42 +362,61 @@ def _build_from_signature(self, query, value, key=None):
if self._reuse_heads < self._num_heads:
einsum_equation, bias_axes, output_rank = _build_proj_equation(
free_dims, bound_dims=1, output_dims=2)
- self._query_dense = tf.keras.layers.experimental.EinsumDense(
+ self._query_dense = tf_keras.layers.EinsumDense(
einsum_equation,
- output_shape=_get_output_shape(output_rank - 1, [
- self._num_heads - self._reuse_heads, self._key_dim]),
+ output_shape=_get_output_shape(
+ output_rank - 1,
+ [self._num_heads - self._reuse_heads, self._key_dim]),
bias_axes=bias_axes if self._use_bias else None,
name="query",
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
einsum_equation, bias_axes, output_rank = _build_proj_equation(
self._key_shape.rank - 1, bound_dims=1, output_dims=2)
- self._key_dense = tf.keras.layers.experimental.EinsumDense(
+ self._key_dense = tf_keras.layers.EinsumDense(
einsum_equation,
- output_shape=_get_output_shape(output_rank - 1, [
- self._num_heads - self._reuse_heads, self._key_dim]),
+ output_shape=_get_output_shape(
+ output_rank - 1,
+ [self._num_heads - self._reuse_heads, self._key_dim]),
bias_axes=bias_axes if self._use_bias else None,
name="key",
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
einsum_equation, bias_axes, output_rank = _build_proj_equation(
self._value_shape.rank - 1, bound_dims=1, output_dims=2)
self._value_dense = []
if self._reuse_heads > 0:
- self._value_dense.append(tf.keras.layers.experimental.EinsumDense(
- einsum_equation,
- output_shape=_get_output_shape(
- output_rank - 1, [self._reuse_heads, self._value_dim]),
- bias_axes=bias_axes if self._use_bias else None,
- name="value_reuse",
- **common_kwargs))
+ self._value_dense.append(
+ tf_keras.layers.EinsumDense(
+ einsum_equation,
+ output_shape=_get_output_shape(
+ output_rank - 1, [self._reuse_heads, self._value_dim]),
+ bias_axes=bias_axes if self._use_bias else None,
+ name="value_reuse",
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(
+ self._bias_initializer),
+ **common_kwargs))
if self._reuse_heads < self._num_heads:
- self._value_dense.append(tf.keras.layers.experimental.EinsumDense(
- einsum_equation,
- output_shape=_get_output_shape(output_rank - 1, [
- self._num_heads - self._reuse_heads, self._value_dim]),
- bias_axes=bias_axes if self._use_bias else None,
- name="value_new",
- **common_kwargs))
+ self._value_dense.append(
+ tf_keras.layers.EinsumDense(
+ einsum_equation,
+ output_shape=_get_output_shape(
+ output_rank - 1,
+ [self._num_heads - self._reuse_heads, self._value_dim]),
+ bias_axes=bias_axes if self._use_bias else None,
+ name="value_new",
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(
+ self._bias_initializer),
+ **common_kwargs))
# Builds the attention computations for multi-head dot product attention.
# These computations could be wrapped into the keras attention layer once
@@ -431,21 +450,23 @@ def _make_output_dense(self, free_dims, common_kwargs, name=None,
else:
output_shape = self._output_shape
else:
- output_shape = [self._query_shape[-1]]
+ output_shape = [self._query_shape[-1]] # pyrefly: ignore[unsupported-operation]
einsum_equation, bias_axes, output_rank = _build_proj_equation(
free_dims, bound_dims=2, output_dims=len(output_shape))
- return tf.keras.layers.experimental.EinsumDense(
+ return tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1, output_shape),
bias_axes=bias_axes if (use_bias and self._use_bias) else None,
name=name,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
def _build_attention(self, rank):
"""Builds multi-head dot-product attention computations.
This function builds attributes necessary for `_compute_attention` to
- costomize attention computation to replace the default dot-product
+ customize attention computation to replace the default dot-product
attention.
Args:
@@ -454,13 +475,13 @@ def _build_attention(self, rank):
if self._attention_axes is None:
self._attention_axes = tuple(range(1, rank - 2))
else:
- self._attention_axes = tuple(self._attention_axes)
+ self._attention_axes = tuple(self._attention_axes) # pyrefly: ignore[bad-argument-type]
self._dot_product_equation, self._combine_equation, attn_scores_rank = (
_build_attention_equation(rank, attn_axes=self._attention_axes))
norm_axes = tuple(
range(attn_scores_rank - len(self._attention_axes), attn_scores_rank))
- self._softmax = tf.keras.layers.Softmax(axis=norm_axes)
- self._dropout_layer = tf.keras.layers.Dropout(rate=self._dropout)
+ self._softmax = tf_keras.layers.Softmax(axis=norm_axes)
+ self._dropout_layer = tf_keras.layers.Dropout(rate=self._dropout)
def _masked_softmax(self, attention_scores, attention_mask=None):
# Normalize the attention scores to probabilities.
@@ -468,7 +489,7 @@ def _masked_softmax(self, attention_scores, attention_mask=None):
if attention_mask is not None:
# The expand dim happens starting from the `num_heads` dimension,
# (, num_heads, )
- mask_expansion_axes = [-len(self._attention_axes) * 2 - 1]
+ mask_expansion_axes = [-len(self._attention_axes) * 2 - 1] # pyrefly: ignore[bad-argument-type]
for _ in range(len(attention_scores.shape) - len(attention_mask.shape)):
attention_mask = tf.expand_dims(
attention_mask, axis=mask_expansion_axes)
@@ -522,7 +543,7 @@ def _compute_attention(self,
tf.shape(query)[1], tf.shape(key)[1])
new_scores = self._masked_softmax(new_scores, attention_mask)
if self._reuse_heads > 0: # Partial reuse
- reuse_scores = reuse_scores[:, :self._reuse_heads, :, :]
+ reuse_scores = reuse_scores[:, :self._reuse_heads, :, :] # pyrefly: ignore[unsupported-operation]
attention_scores = tf.concat([new_scores, reuse_scores], 1)
else: # No reuse
attention_scores = new_scores
diff --git a/official/nlp/modeling/layers/reuse_attention_test.py b/official/nlp/modeling/layers/reuse_attention_test.py
index fe9e71d2f06..225bdbc8e36 100644
--- a/official/nlp/modeling/layers/reuse_attention_test.py
+++ b/official/nlp/modeling/layers/reuse_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import reuse_attention as attention
@@ -36,8 +36,8 @@ def test_non_masked_attention(self, value_dim, output_shape, output_dims):
value_dim=value_dim,
output_shape=output_shape)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
- value = tf.keras.Input(shape=(20, 80))
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(20, 80))
output = test_layer(query=query, value=value)
self.assertEqual(output.shape.as_list(), [None] + output_dims)
@@ -46,7 +46,7 @@ def test_non_masked_self_attention(self):
test_layer = attention.ReuseMultiHeadAttention(
num_heads=12, key_dim=64)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
+ query = tf_keras.Input(shape=(40, 80))
output = test_layer(query, query)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
@@ -55,7 +55,7 @@ def test_attention_scores(self):
test_layer = attention.ReuseMultiHeadAttention(
num_heads=12, key_dim=64)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
+ query = tf_keras.Input(shape=(40, 80))
output, coef = test_layer(query, query, return_attention_scores=True)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
self.assertEqual(coef.shape.as_list(), [None, 12, 40, 40])
@@ -65,8 +65,8 @@ def test_attention_scores_with_values(self):
test_layer = attention.ReuseMultiHeadAttention(
num_heads=12, key_dim=64)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
- value = tf.keras.Input(shape=(60, 80))
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(60, 80))
output, coef = test_layer(query, value, return_attention_scores=True)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
self.assertEqual(coef.shape.as_list(), [None, 12, 40, 60])
@@ -83,14 +83,14 @@ def test_masked_attention(self, use_bias, reuse_attention):
reuse_attention=reuse_attention)
# Create a 3-dimensional input (the first dimension is implicit).
batch_size = 3
- query = tf.keras.Input(shape=(4, 8))
- value = tf.keras.Input(shape=(2, 8))
- mask_tensor = tf.keras.Input(shape=(4, 2))
- reuse_attention_scores = tf.keras.Input(shape=(2, 4, 2))
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
+ reuse_attention_scores = tf_keras.Input(shape=(2, 4, 2))
output = test_layer(query=query, value=value, attention_mask=mask_tensor,
reuse_attention_scores=reuse_attention_scores)
# Create a model containing the test layer.
- model = tf.keras.Model(
+ model = tf_keras.Model(
[query, value, mask_tensor, reuse_attention_scores], output)
# Generate data for the input (non-mask) tensors.
@@ -116,10 +116,10 @@ def test_masked_attention(self, use_bias, reuse_attention):
self.assertNotAllClose(masked_output_data, unmasked_output_data)
# Tests the layer with three inputs: Q, K, V.
- key = tf.keras.Input(shape=(2, 8))
+ key = tf_keras.Input(shape=(2, 8))
output = test_layer(query, value=value, key=key, attention_mask=mask_tensor,
reuse_attention_scores=reuse_attention_scores)
- model = tf.keras.Model(
+ model = tf_keras.Model(
[query, value, key, mask_tensor, reuse_attention_scores], output)
masked_output_data = model.predict(
@@ -152,9 +152,9 @@ def test_initializer(self):
test_layer = attention.ReuseMultiHeadAttention(
num_heads=12,
key_dim=64,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
+ query = tf_keras.Input(shape=(40, 80))
output = test_layer(query, query)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
@@ -164,13 +164,13 @@ def test_masked_attention_with_scores(self):
num_heads=2, key_dim=2)
# Create a 3-dimensional input (the first dimension is implicit).
batch_size = 3
- query = tf.keras.Input(shape=(4, 8))
- value = tf.keras.Input(shape=(2, 8))
- mask_tensor = tf.keras.Input(shape=(4, 2))
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
output = test_layer(query=query, value=value, attention_mask=mask_tensor)
# Create a model containing the test layer.
- model = tf.keras.Model([query, value, mask_tensor], output)
+ model = tf_keras.Model([query, value, mask_tensor], output)
# Generate data for the input (non-mask) tensors.
from_data = 10 * np.random.random_sample((batch_size, 4, 8))
@@ -193,7 +193,7 @@ def test_masked_attention_with_scores(self):
output, scores = test_layer(
query=query, value=value, attention_mask=mask_tensor,
return_attention_scores=True)
- model = tf.keras.Model([query, value, mask_tensor], [output, scores])
+ model = tf_keras.Model([query, value, mask_tensor], [output, scores])
masked_output_data_score, masked_score = model.predict(
[from_data, to_data, mask_data])
unmasked_output_data_score, unmasked_score = model.predict(
@@ -230,12 +230,12 @@ def test_high_dim_attention(self, q_dims, v_dims, mask_dims, attention_axes):
null_mask_data = np.ones(mask_shape)
# Because one data is masked and one is not, the outputs should not be the
# same.
- query_tensor = tf.keras.Input(query_shape[1:], name="query")
- value_tensor = tf.keras.Input(value_shape[1:], name="value")
- mask_tensor = tf.keras.Input(mask_shape[1:], name="mask")
+ query_tensor = tf_keras.Input(query_shape[1:], name="query")
+ value_tensor = tf_keras.Input(value_shape[1:], name="value")
+ mask_tensor = tf_keras.Input(mask_shape[1:], name="mask")
output = test_layer(query=query_tensor, value=value_tensor,
attention_mask=mask_tensor)
- model = tf.keras.Model([query_tensor, value_tensor, mask_tensor], output)
+ model = tf_keras.Model([query_tensor, value_tensor, mask_tensor], output)
self.assertNotAllClose(
model.predict([query, value, mask_data]),
@@ -246,24 +246,24 @@ def test_dropout(self):
num_heads=2, key_dim=2, dropout=0.5)
# Generate data for the input (non-mask) tensors.
- from_data = tf.keras.backend.ones(shape=(32, 4, 8))
- to_data = tf.keras.backend.ones(shape=(32, 2, 8))
+ from_data = tf_keras.backend.ones(shape=(32, 4, 8))
+ to_data = tf_keras.backend.ones(shape=(32, 2, 8))
train_out = test_layer(from_data, to_data, None, None, None, True)
test_out = test_layer(from_data, to_data, None, None, None, False)
# Output should be close when not in training mode,
# and should not be close when enabling dropout in training mode.
self.assertNotAllClose(
- tf.keras.backend.eval(train_out),
- tf.keras.backend.eval(test_out))
+ tf_keras.backend.eval(train_out),
+ tf_keras.backend.eval(test_out))
def test_non_masked_self_attention_with_reuse(self):
"""Test with one input (self-attenntion) and no mask tensor."""
test_layer = attention.ReuseMultiHeadAttention(
num_heads=12, key_dim=64, reuse_attention=True)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
- reuse_scores = tf.keras.Input(shape=(12, 40, 40))
+ query = tf_keras.Input(shape=(40, 80))
+ reuse_scores = tf_keras.Input(shape=(12, 40, 40))
output = test_layer(query, query, reuse_attention_scores=reuse_scores)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
@@ -281,22 +281,22 @@ def test_non_masked_self_attention_with_relative_pe(self, reuse_attention,
num_heads=12, key_dim=64, reuse_attention=reuse_attention,
use_relative_pe=True, pe_max_seq_length=pe_max_seq_length)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
- reuse_scores = tf.keras.Input(shape=(12, 40, 40))
+ query = tf_keras.Input(shape=(40, 80))
+ reuse_scores = tf_keras.Input(shape=(12, 40, 40))
output = test_layer(query, query, reuse_attention_scores=reuse_scores)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
- query = tf.keras.Input(shape=(30, 80))
- reuse_scores = tf.keras.Input(shape=(12, 30, 30))
+ query = tf_keras.Input(shape=(30, 80))
+ reuse_scores = tf_keras.Input(shape=(12, 30, 30))
output = test_layer(query, query, reuse_attention_scores=reuse_scores)
self.assertEqual(output.shape.as_list(), [None, 30, 80])
- query = tf.keras.Input(shape=(30, 80))
- key = tf.keras.Input(shape=(20, 80))
- reuse_scores = tf.keras.Input(shape=(12, 30, 20))
+ query = tf_keras.Input(shape=(30, 80))
+ key = tf_keras.Input(shape=(20, 80))
+ reuse_scores = tf_keras.Input(shape=(12, 30, 20))
output = test_layer(query, key, reuse_attention_scores=reuse_scores)
self.assertEqual(output.shape.as_list(), [None, 30, 80])
- query = tf.keras.Input(shape=(50, 80))
- key = tf.keras.Input(shape=(60, 80))
- reuse_scores = tf.keras.Input(shape=(12, 50, 60))
+ query = tf_keras.Input(shape=(50, 80))
+ key = tf_keras.Input(shape=(60, 80))
+ reuse_scores = tf_keras.Input(shape=(12, 50, 60))
output = test_layer(query, key, reuse_attention_scores=reuse_scores)
self.assertEqual(output.shape.as_list(), [None, 50, 80])
diff --git a/official/nlp/modeling/layers/reuse_transformer.py b/official/nlp/modeling/layers/reuse_transformer.py
index fba459573cd..54b9571259f 100644
--- a/official/nlp/modeling/layers/reuse_transformer.py
+++ b/official/nlp/modeling/layers/reuse_transformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,11 +13,13 @@
# limitations under the License.
"""Keras-based TransformerEncoder block layer."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
from official.nlp.modeling.layers import reuse_attention as attention
-class ReuseTransformer(tf.keras.layers.Layer):
+class ReuseTransformer(tf_keras.layers.Layer):
"""Transformer layer.
This layer implements the ReuseTransformer Encoder from
@@ -106,13 +108,13 @@ def __init__(self,
self._output_dropout = output_dropout
self._output_dropout_rate = output_dropout
self._output_range = output_range
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
@@ -128,10 +130,11 @@ def __init__(self,
self._max_reuse_layer_idx < self._layer_idx)):
self._reuse_attention = 0
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
+ self._attention_initializer = tf_keras.initializers.get(
attention_initializer)
else:
- self._attention_initializer = self._kernel_initializer
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer)
self._attention_axes = attention_axes
def build(self, input_shape):
@@ -156,7 +159,6 @@ def build(self, input_shape):
else:
self._attention_head_size = self._head_size
common_kwargs = dict(
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
@@ -168,49 +170,52 @@ def build(self, input_shape):
dropout=self._attention_dropout,
use_bias=self._use_bias,
kernel_initializer=self._attention_initializer,
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
attention_axes=self._attention_axes,
reuse_attention=self._reuse_attention,
use_relative_pe=self._use_relative_pe,
pe_max_seq_length=self._pe_max_seq_length,
name="self_attention",
**common_kwargs)
- self._attention_dropout = tf.keras.layers.Dropout(
+ self._attention_dropout = tf_keras.layers.Dropout(
rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype=tf.float32))
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, self._inner_dim),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="intermediate",
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
# TODO(b/154538392): Investigate this.
policy = tf.float32
- self._intermediate_activation_layer = tf.keras.layers.Activation(
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
self._inner_activation, dtype=policy)
- self._inner_dropout_layer = tf.keras.layers.Dropout(
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
rate=self._inner_dropout)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._output_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, hidden_size),
bias_axes="d",
name="output",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -240,19 +245,19 @@ def get_config(self):
"pe_max_seq_length": self._pe_max_seq_length,
"max_reuse_layer_idx": self._max_reuse_layer_idx,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
+ tf_keras.constraints.serialize(self._bias_constraint),
"use_bias":
self._use_bias,
"norm_first":
@@ -262,7 +267,7 @@ def get_config(self):
"inner_dropout":
self._inner_dropout,
"attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer),
+ tf_keras.initializers.serialize(self._attention_initializer),
"attention_axes": self._attention_axes,
}
base_config = super(ReuseTransformer, self).get_config()
@@ -329,9 +334,9 @@ def call(self, inputs):
reuse_attention_scores=reuse_attention_scores,
return_attention_scores=True)
attention_output, attention_scores = attention_output
- attention_output = self._attention_dropout(attention_output)
+ attention_output = self._attention_dropout(attention_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- attention_output = source_tensor + attention_output
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self._attention_layer_norm(target_tensor +
attention_output)
@@ -343,10 +348,10 @@ def call(self, inputs):
inner_output = self._intermediate_activation_layer(inner_output)
inner_output = self._inner_dropout_layer(inner_output)
layer_output = self._output_dense(inner_output)
- layer_output = self._output_dropout(layer_output)
+ layer_output = self._output_dropout(layer_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- return source_attention_output + layer_output, attention_scores
+ return source_attention_output + layer_output, attention_scores # pyrefly: ignore[unbound-name]
# During mixed precision training, layer norm output is always fp32 for now.
# Casts fp32 for the subsequent add.
diff --git a/official/nlp/modeling/layers/reuse_transformer_test.py b/official/nlp/modeling/layers/reuse_transformer_test.py
index 0376906e909..f0b6d22546b 100644
--- a/official/nlp/modeling/layers/reuse_transformer_test.py
+++ b/official/nlp/modeling/layers/reuse_transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import reuse_transformer
@@ -27,7 +27,7 @@ class ReuseTransformerLayerTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(ReuseTransformerLayerTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
def test_layer_creation(self, transformer_cls):
test_layer = transformer_cls(
@@ -35,7 +35,7 @@ def test_layer_creation(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor, _ = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -46,9 +46,9 @@ def test_layer_creation_with_mask(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor, _ = test_layer([data_tensor, mask_tensor])
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -59,11 +59,11 @@ def test_layer_invocation(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(data_tensor, output_tensor)
+ model = tf_keras.Model(data_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -78,13 +78,13 @@ def test_layer_invocation_with_mask(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -206,19 +206,19 @@ def test_layer_output_range_with_pre_norm(self, transformer_cls):
new_output_tensor, output_tensor[:, 0:1, :], atol=0.002, rtol=0.01)
def test_layer_invocation_with_float16_dtype(self, transformer_cls):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
test_layer = transformer_cls(
num_attention_heads=10, inner_dim=2048, inner_activation='relu')
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -236,11 +236,11 @@ def test_transform_with_initializer(self, transformer_cls):
num_attention_heads=10,
inner_dim=2048,
inner_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output, _ = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output.shape.as_list())
@@ -250,12 +250,12 @@ def test_dynamic_layer_sequence(self, transformer_cls):
num_attention_heads=10,
inner_dim=2048,
inner_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Create a 3-dimensional input (the first dimension is implicit).
width = 30
- input_tensor = tf.keras.Input(shape=(None, width))
+ input_tensor = tf_keras.Input(shape=(None, width))
output_tensor, _ = test_layer(input_tensor)
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
input_length = 17
input_data = np.ones((1, input_length, width))
@@ -279,7 +279,7 @@ def test_use_bias_norm_first(self):
norm_first=True,
norm_epsilon=1e-6,
inner_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
# Forward path.
dummy_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
@@ -300,7 +300,7 @@ def test_get_config(self):
norm_first=True,
norm_epsilon=1e-6,
inner_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
encoder_block_config = encoder_block.get_config()
new_encoder_block = reuse_transformer.ReuseTransformer.from_config(
@@ -325,23 +325,20 @@ def test_several_attention_axes(self, attention_axes):
num_cols = 13
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(num_rows, num_cols, width))
+ data_tensor = tf_keras.Input(shape=(num_rows, num_cols, width))
output_tensor, _ = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@parameterized.named_parameters(
- ('plain', False, False, False),
- ('plain_returnscore', False, True, False),
- ('plain_with_relative_pe', False, False, True),
- ('reuse_all', True, False, False),
- ('reuse_all_returnscore', True, True, False),
- ('reuse_all_with_relative_pe', True, False, True),
- ('reuse_5', 5, False, False),
- ('reuse_5_returnscore', 5, True, False),
- ('reuse_5_with_relative_pe', 5, False, True),)
- def test_layer_invocation_with_mask(self, reuse_attention,
- return_attention_scores, use_relative_pe):
+ ('plain_returnscore', False, False),
+ ('plain_with_relative_pe', False, True),
+ ('reuse_all_returnscore', True, False),
+ ('reuse_all_with_relative_pe', True, True),
+ ('reuse_5_returnscore', 5, False),
+ ('reuse_5_with_relative_pe', 5, True),
+ )
+ def test_layer_invocation_with_mask(self, reuse_attention, use_relative_pe):
test_layer = reuse_transformer.ReuseTransformer(
num_attention_heads=10,
inner_dim=2048,
@@ -351,19 +348,23 @@ def test_layer_invocation_with_mask(self, reuse_attention,
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
- return_scores_tensor = tf.keras.Input(shape=(1,))
- reuse_attention_scores = tf.keras.Input(
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
+ reuse_attention_scores = tf_keras.Input(
shape=(10, sequence_length, sequence_length))
output_tensor, _ = test_layer(
[data_tensor, mask_tensor, reuse_attention_scores])
# Create a model from the test layer.
- model = tf.keras.Model(
- ([data_tensor, mask_tensor, reuse_attention_scores],
- return_scores_tensor), output_tensor)
+ model = tf_keras.Model(
+ [
+ data_tensor,
+ mask_tensor,
+ reuse_attention_scores,
+ ],
+ output_tensor,
+ )
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -376,8 +377,7 @@ def test_layer_invocation_with_mask(self, reuse_attention,
2, size=(batch_size, sequence_length, sequence_length))
reuse_scores = np.random.rand(
batch_size, 10, sequence_length, sequence_length)
- _ = model.predict([input_data, mask_data, reuse_scores],
- return_attention_scores)
+ _ = model.predict([input_data, mask_data, reuse_scores])
@parameterized.named_parameters(
('without_relative_pe_with_pe_max_seq_length_10', False, 10),
@@ -386,20 +386,20 @@ def test_layer_invocation_with_mask(self, reuse_attention,
('with_relative_pe_with_pe_max_seq_length_100', True, 100))
def test_layer_invocation_with_float16_with_relative_pe(
self, use_relative_pe, pe_max_seq_length):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
test_layer = reuse_transformer.ReuseTransformer(
num_attention_heads=10, inner_dim=2048, inner_activation='relu',
use_relative_pe=use_relative_pe, pe_max_seq_length=pe_max_seq_length)
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
diff --git a/official/nlp/modeling/layers/rezero_transformer.py b/official/nlp/modeling/layers/rezero_transformer.py
index 626753a8680..5885d2edf06 100644
--- a/official/nlp/modeling/layers/rezero_transformer.py
+++ b/official/nlp/modeling/layers/rezero_transformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,16 +14,21 @@
"""Keras-based rezero-transformer block layer (Transformer with ReZero)."""
# pylint: disable=g-classes-have-attributes
+from typing import Optional
+from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
+from official.nlp.modeling.layers import block_sparse_attention
+from official.nlp.modeling.layers import multi_query_attention
from official.nlp.modeling.layers import util
-@tf.keras.utils.register_keras_serializable(package="Text")
+@tf_keras.utils.register_keras_serializable(package="Text")
@gin.configurable
-class ReZeroTransformer(tf.keras.layers.Layer):
+class ReZeroTransformer(tf_keras.layers.Layer):
"""Transformer layer with ReZero.
This layer implements the Transformer from "Attention Is All You Need".
@@ -33,8 +38,10 @@ class ReZeroTransformer(tf.keras.layers.Layer):
Args:
num_attention_heads: Number of attention heads.
- intermediate_size: Size of the intermediate layer.
- intermediate_activation: Activation for the intermediate layer.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network.
dropout_rate: Dropout probability for the post-attention and output dropout.
attention_dropout_rate: Dropout probability for within the attention layer.
output_range: the sequence output range, [0, output_range) by slicing the
@@ -48,12 +55,18 @@ class ReZeroTransformer(tf.keras.layers.Layer):
bias_constraint: Constraint for dense layer kernels.
use_layer_norm: If add layer_norm on top of the ReZero.
share_rezero: If attention layer and FFN layer share the same alpha.
+ num_kv_heads: Number of key-value heads for multi-query attention. Refer to
+ `multi_query_attention.MultiHeadAttention` for more details.
+ src_block_size: Source block size. Refer to
+ `block_sparse_attention.MultiHeadAttention` for more details.
+ tgt_block_size: Target block size. Refer to
+ `block_sparse_attention.MultiHeadAttention` for more details.
"""
def __init__(self,
num_attention_heads,
- intermediate_size,
- intermediate_activation,
+ inner_dim=768,
+ inner_activation=tf_utils.get_activation("gelu"),
dropout_rate=0.0,
attention_dropout_rate=0.0,
output_range=None,
@@ -66,29 +79,60 @@ def __init__(self,
bias_constraint=None,
use_layer_norm=False,
share_rezero=True,
+ num_kv_heads=None,
+ src_block_size=None,
+ tgt_block_size=None,
+ linformer_dim=None,
+ linformer_shared_kv_projection=True,
+ use_sigmoid_attn=False,
+ sigmoid_attn_bias=None,
**kwargs):
# attention_dropout will override attention_dropout_rate.
# This is to unify the input params with TransformerEncoderBlock.
attention_dropout_rate = kwargs.pop("attention_dropout",
attention_dropout_rate)
dropout_rate = kwargs.pop("output_dropout", dropout_rate)
+ inner_dim = kwargs.pop("intermediate_size", inner_dim)
+ inner_activation = kwargs.pop("intermediate_activation", inner_activation)
util.filter_kwargs(kwargs)
- super(ReZeroTransformer, self).__init__(**kwargs)
+ super().__init__(**kwargs)
+
+ # Deprecation warning.
+ if output_range is not None:
+ logging.warning("`output_range` is avaliable as an argument for `call()`."
+ "The `output_range` as __init__ argument is deprecated.")
self._num_heads = num_attention_heads
- self._intermediate_size = intermediate_size
- self._intermediate_activation = intermediate_activation
+ self._inner_dim = inner_dim
+ self._inner_activation = inner_activation
self._attention_dropout_rate = attention_dropout_rate
self._dropout_rate = dropout_rate
self._output_range = output_range
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_layer_norm = use_layer_norm
self._share_rezero = share_rezero
+ self._num_kv_heads = num_kv_heads
+ self._src_block_size = src_block_size
+ self._tgt_block_size = tgt_block_size
+ self._linformer_dim = linformer_dim
+ self._linformer_shared_kv_projection = linformer_shared_kv_projection
+ self._use_sigmoid_attn = use_sigmoid_attn
+ self._sigmoid_attn_bias = sigmoid_attn_bias
+ if self._linformer_dim is not None or self._use_sigmoid_attn:
+ raise ValueError(
+ "Linformer and Sigmoid attention are not supported in ReZero"
+ " Transformer."
+ )
+ if self._num_kv_heads is not None and self._src_block_size is not None:
+ raise ValueError(
+ "Block sparse attention does not support Multi-query attention."
+ " Specify only one of them."
+ )
def build(self, input_shape):
if isinstance(input_shape, tf.TensorShape):
@@ -97,82 +141,110 @@ def build(self, input_shape):
input_tensor_shape = tf.TensorShape(input_shape[0])
else:
raise ValueError(
- "The type of input shape argument is not supported, got: %s" %
- type(input_shape))
+ "The type of input shape argument is not supported, got: %s"
+ % type(input_shape)
+ )
if len(input_tensor_shape.as_list()) != 3:
- raise ValueError("TransformerLayer expects a three-dimensional input of "
- "shape [batch, sequence, width].")
+ raise ValueError(
+ "TransformerLayer expects a three-dimensional input of "
+ "shape [batch, sequence, width]."
+ )
batch_size, sequence_length, hidden_size = input_tensor_shape
if len(input_shape) == 2:
mask_tensor_shape = tf.TensorShape(input_shape[1])
expected_mask_tensor_shape = tf.TensorShape(
- [batch_size, sequence_length, sequence_length])
+ [batch_size, sequence_length, sequence_length]
+ )
if not expected_mask_tensor_shape.is_compatible_with(mask_tensor_shape):
- raise ValueError("When passing a mask tensor to TransformerLayer, the "
- "mask tensor must be of shape [batch, "
- "sequence_length, sequence_length] (here %s). Got a "
- "mask tensor of shape %s." %
- (expected_mask_tensor_shape, mask_tensor_shape))
+ raise ValueError(
+ "When passing a mask tensor to TransformerLayer, the "
+ "mask tensor must be of shape [batch, "
+ "sequence_length, sequence_length] (here %s). Got a "
+ "mask tensor of shape %s."
+ % (expected_mask_tensor_shape, mask_tensor_shape)
+ )
if hidden_size % self._num_heads != 0:
raise ValueError(
"The input size (%d) is not a multiple of the number of attention "
- "heads (%d)" % (hidden_size, self._num_heads))
+ "heads (%d)" % (hidden_size, self._num_heads)
+ )
self._attention_head_size = int(hidden_size // self._num_heads)
common_kwargs = dict(
- kernel_initializer=self._kernel_initializer,
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
kernel_constraint=self._kernel_constraint,
- bias_constraint=self._bias_constraint)
- self._attention_layer = tf.keras.layers.MultiHeadAttention(
+ bias_constraint=self._bias_constraint,
+ )
+ attention_kwargs = dict(
num_heads=self._num_heads,
key_dim=self._attention_head_size,
dropout=self._attention_dropout_rate,
name="self_attention",
- **common_kwargs)
- self._attention_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ )
+ if self._src_block_size is not None:
+ attention_kwargs.update(
+ src_block_size=self._src_block_size,
+ tgt_block_size=self._tgt_block_size,
+ name="block_sparse_attention",
+ )
+ attention_fn = block_sparse_attention.MultiHeadAttention
+ elif self._num_kv_heads is not None:
+ attention_kwargs.update(
+ num_kv_heads=self._num_kv_heads,
+ name="multi_query_attention",
+ )
+ attention_fn = multi_query_attention.MultiHeadAttention
+ else:
+ attention_fn = tf_keras.layers.MultiHeadAttention
+ self._attention_layer = attention_fn(**attention_kwargs, **common_kwargs)
+ self._attention_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
if self._use_layer_norm:
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
- self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
- name="self_attention_layer_norm",
- axis=-1,
- epsilon=1e-12,
- dtype=tf.float32))
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self._attention_layer_norm = tf_keras.layers.LayerNormalization(
+ name="self_attention_layer_norm",
+ axis=-1,
+ epsilon=1e-12,
+ dtype=tf.float32,
+ )
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
- output_shape=(None, self._intermediate_size),
+ output_shape=(None, self._inner_dim),
bias_axes="d",
name="intermediate",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
# TODO(b/154538392): Investigate this.
policy = tf.float32
- self._intermediate_activation_layer = tf.keras.layers.Activation(
- self._intermediate_activation, dtype=policy)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._inner_activation_layer = tf_keras.layers.Activation(
+ self._inner_activation, dtype=policy)
+ self._output_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
name="output",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
if self._use_layer_norm:
# Use float32 in layernorm for numeric stability.
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm", axis=-1, epsilon=1e-12, dtype=tf.float32)
self._rezero_a = self.add_weight(
name="rezero_alpha",
- initializer=tf.keras.initializers.Zeros(),
+ initializer=tf_keras.initializers.Zeros(),
trainable=True,
dtype=tf.float32)
@@ -181,20 +253,20 @@ def build(self, input_shape):
else:
self._rezero_a_ffn = self.add_weight(
name="rezero_alpha_ffn",
- initializer=tf.keras.initializers.Zeros(),
+ initializer=tf_keras.initializers.Zeros(),
trainable=True,
dtype=tf.float32)
- super(ReZeroTransformer, self).build(input_shape)
+ super().build(input_shape)
def get_config(self):
config = {
"num_attention_heads":
self._num_heads,
- "intermediate_size":
- self._intermediate_size,
- "intermediate_activation":
- self._intermediate_activation,
+ "inner_dim":
+ self._inner_dim,
+ "inner_activation":
+ self._inner_activation,
"dropout_rate":
self._dropout_rate,
"attention_dropout_rate":
@@ -205,22 +277,34 @@ def get_config(self):
self._use_layer_norm,
"share_rezero":
self._share_rezero,
+ "num_kv_heads":
+ self._num_kv_heads,
+ "src_block_size":
+ self._src_block_size,
+ "tgt_block_size":
+ self._tgt_block_size,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
+ tf_keras.constraints.serialize(self._bias_constraint),
+ "linformer_dim": self._linformer_dim,
+ "linformer_shared_kv_projection": (
+ self._linformer_shared_kv_projection
+ ),
+ "use_sigmoid_attn": self._use_sigmoid_attn,
+ "sigmoid_attn_bias": self._sigmoid_attn_bias,
}
- base_config = super(ReZeroTransformer, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def reset_rezero(self):
@@ -228,7 +312,7 @@ def reset_rezero(self):
if not self._share_rezero:
self._rezero_a_ffn.assign(0.)
- def call(self, inputs):
+ def call(self, inputs, output_range: Optional[tf.Tensor] = None) -> tf.Tensor:
if isinstance(inputs, (list, tuple)):
if len(inputs) == 2:
input_tensor, attention_mask = inputs
@@ -241,10 +325,12 @@ def call(self, inputs):
else:
input_tensor, key_value, attention_mask = (inputs, None, None)
- if self._output_range:
- target_tensor = input_tensor[:, 0:self._output_range, :]
+ if output_range is None:
+ output_range = self._output_range
+ if output_range:
+ target_tensor = input_tensor[:, 0:output_range, :]
if attention_mask is not None:
- attention_mask = attention_mask[:, 0:self._output_range, :]
+ attention_mask = attention_mask[:, 0:output_range, :]
else:
target_tensor = input_tensor
@@ -261,8 +347,7 @@ def call(self, inputs):
attention_output = tf.cast(attention_output, tf.float32)
intermediate_output = self._intermediate_dense(attention_output)
- intermediate_output = self._intermediate_activation_layer(
- intermediate_output)
+ intermediate_output = self._inner_activation_layer(intermediate_output)
layer_output = self._output_dense(intermediate_output)
layer_output = self._output_dropout(layer_output)
# During mixed precision training, attention_output is from layer norm and
diff --git a/official/nlp/modeling/layers/rezero_transformer_test.py b/official/nlp/modeling/layers/rezero_transformer_test.py
index ed940626c89..67788782de5 100644
--- a/official/nlp/modeling/layers/rezero_transformer_test.py
+++ b/official/nlp/modeling/layers/rezero_transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,25 +16,21 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import rezero_transformer
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class TransformerWithReZeroLayerTest(keras_parameterized.TestCase):
+class TransformerWithReZeroLayerTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(TransformerWithReZeroLayerTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
@parameterized.named_parameters(('no_share_attn_ffn', False),
('share_attn_ffn', True))
def test_layer_invocation_with_float16_dtype(self, share_rezero):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
test_layer = rezero_transformer.ReZeroTransformer(
num_attention_heads=10,
intermediate_size=2048,
@@ -43,13 +39,13 @@ def test_layer_invocation_with_float16_dtype(self, share_rezero):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -70,9 +66,9 @@ def test_rezero_without_layer_norm(self):
use_layer_norm=False)
input_length, width = 16, 30
- input_tensor = tf.keras.Input(shape=(input_length, width))
+ input_tensor = tf_keras.Input(shape=(input_length, width))
output_tensor = test_layer(input_tensor)
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
input_data = np.random.rand(2, input_length, width)
test_layer._rezero_a.assign(1.0)
@@ -89,9 +85,9 @@ def test_rezero_with_layer_norm(self):
use_layer_norm=True)
input_length, width = 16, 30
- input_tensor = tf.keras.Input(shape=(input_length, width))
+ input_tensor = tf_keras.Input(shape=(input_length, width))
output_tensor = test_layer(input_tensor)
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
input_data = np.random.rand(2, input_length, width) + 2.0
output_data = model.predict(input_data)
@@ -128,12 +124,15 @@ def test_layer_output_range(self):
new_output_tensor = new_layer([input_data, mask_data])
self.assertAllClose(new_output_tensor, output_tensor[:, 0:1, :])
+ output_tensor = test_layer([input_data, mask_data], output_range=1)
+ self.assertAllClose(new_output_tensor, output_tensor, atol=5e-5, rtol=0.003)
+
def test_separate_qkv(self):
test_layer = rezero_transformer.ReZeroTransformer(
num_attention_heads=2,
intermediate_size=128,
intermediate_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Forward path.
q_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
kv_tensor = tf.zeros([2, 8, 16], dtype=tf.float32)
@@ -142,6 +141,69 @@ def test_separate_qkv(self):
output = test_layer(inputs)
self.assertEqual(output.shape, q_tensor.shape)
+ @parameterized.named_parameters(('_mqa', 1),
+ ('_gqa', 5))
+ def test_rezero_with_kv_heads(self, num_kv_heads):
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
+ test_layer = rezero_transformer.ReZeroTransformer(
+ num_attention_heads=10,
+ intermediate_size=2048,
+ intermediate_activation='relu',
+ num_kv_heads=num_kv_heads,
+ )
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ # Create a 2-dimensional input (the first dimension is implicit).
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
+ output_tensor = test_layer([data_tensor, mask_tensor])
+
+ # Create a model from the test layer.
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 6
+ input_data = (10 * np.random.random_sample(
+ (batch_size, sequence_length, width)))
+ # The attention mask should be of shape (batch, from_seq_len, to_seq_len),
+ # which here is (batch, sequence_length, sequence_length)
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length, sequence_length))
+ _ = model.predict([input_data, mask_data])
+
+ def test_rezero_with_block_sparse_attention(self):
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
+ test_layer = rezero_transformer.ReZeroTransformer(
+ num_attention_heads=10,
+ intermediate_size=2048,
+ intermediate_activation='relu',
+ src_block_size=3,
+ tgt_block_size=3,
+ )
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ # Create a 2-dimensional input (the first dimension is implicit).
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
+ output_tensor = test_layer([data_tensor, mask_tensor])
+
+ # Create a model from the test layer.
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 6
+ input_data = (10 * np.random.random_sample(
+ (batch_size, sequence_length, width)))
+ # The attention mask should be of shape (batch, from_seq_len, to_seq_len),
+ # which here is (batch, sequence_length, sequence_length)
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length, sequence_length))
+ _ = model.predict([input_data, mask_data])
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/nlp/modeling/layers/routing.py b/official/nlp/modeling/layers/routing.py
index 6eb42f1e4b9..c5d355f605a 100644
--- a/official/nlp/modeling/layers/routing.py
+++ b/official/nlp/modeling/layers/routing.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,11 +18,11 @@
Later on, different sets of tokens can potentially go to different experts.
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package="Text")
-class TokenImportanceWithMovingAvg(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class TokenImportanceWithMovingAvg(tf_keras.layers.Layer):
"""Routing based on per-token importance value."""
def __init__(self,
@@ -33,13 +33,13 @@ def __init__(self,
self._vocab_size = vocab_size
self._init_importance = init_importance
self._moving_average_beta = moving_average_beta
- super(TokenImportanceWithMovingAvg, self).__init__(**kwargs)
+ super().__init__(**kwargs)
def build(self, input_shape):
self._importance_embedding = self.add_weight(
name="importance_embed",
shape=(self._vocab_size),
- initializer=tf.keras.initializers.Constant(self._init_importance),
+ initializer=tf_keras.initializers.Constant(self._init_importance),
trainable=False)
def get_config(self):
@@ -51,7 +51,7 @@ def get_config(self):
"moving_average_beta":
self._moving_average_beta,
}
- base_config = super(TokenImportanceWithMovingAvg, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def update_token_importance(self, token_ids, importance):
@@ -70,8 +70,8 @@ def call(self, inputs):
return tf.gather(self._importance_embedding, inputs)
-@tf.keras.utils.register_keras_serializable(package="Text")
-class SelectTopK(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class SelectTopK(tf_keras.layers.Layer):
"""Select top-k + random-k tokens according to importance."""
def __init__(self,
@@ -80,7 +80,7 @@ def __init__(self,
**kwargs):
self._top_k = top_k
self._random_k = random_k
- super(SelectTopK, self).__init__(**kwargs)
+ super().__init__(**kwargs)
def get_config(self):
config = {
@@ -89,7 +89,7 @@ def get_config(self):
"random_k":
self._random_k,
}
- base_config = super(SelectTopK, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def call(self, inputs):
diff --git a/official/nlp/modeling/layers/routing_test.py b/official/nlp/modeling/layers/routing_test.py
index 8d124187f3c..ae0b3b916e5 100644
--- a/official/nlp/modeling/layers/routing_test.py
+++ b/official/nlp/modeling/layers/routing_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import routing
diff --git a/official/nlp/modeling/layers/self_attention_mask.py b/official/nlp/modeling/layers/self_attention_mask.py
index ed9783adaea..f63b60b6995 100644
--- a/official/nlp/modeling/layers/self_attention_mask.py
+++ b/official/nlp/modeling/layers/self_attention_mask.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,12 +13,40 @@
# limitations under the License.
"""Keras layer that creates a self-attention mask."""
+from typing import Optional
+import tensorflow as tf, tf_keras
-import tensorflow as tf
+def get_mask(inputs: tf.Tensor,
+ to_mask: tf.Tensor,
+ dtype: Optional[tf.DType] = None) -> tf.Tensor:
+ """Gets a 3D self-attention mask.
-@tf.keras.utils.register_keras_serializable(package='Text')
-class SelfAttentionMask(tf.keras.layers.Layer):
+ Args:
+ inputs: from_tensor: 2D or 3D Tensor of shape [batch_size, from_seq_length,
+ ...].
+ to_mask: int32 Tensor of shape [batch_size, to_seq_length].
+ dtype: the output Tensor dtype.
+
+ Returns:
+ float Tensor of shape [batch_size, from_seq_length, to_seq_length].
+ """
+ from_shape = tf.shape(inputs)
+ batch_size = from_shape[0]
+ from_seq_length = from_shape[1]
+ dtype = inputs.dtype if dtype is None else dtype
+
+ to_shape = tf.shape(to_mask)
+ to_seq_length = to_shape[1]
+
+ to_mask = tf.cast(
+ tf.reshape(to_mask, [batch_size, 1, to_seq_length]), dtype=dtype)
+
+ return tf.broadcast_to(to_mask, [batch_size, from_seq_length, to_seq_length])
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class SelfAttentionMask(tf_keras.layers.Layer):
"""Create 3D attention mask from a 2D tensor mask.
inputs[0]: from_tensor: 2D or 3D Tensor of shape
@@ -33,26 +61,4 @@ def call(self, inputs, to_mask=None):
if isinstance(inputs, list) and to_mask is None:
to_mask = inputs[1]
inputs = inputs[0]
- from_shape = tf.shape(inputs)
- batch_size = from_shape[0]
- from_seq_length = from_shape[1]
-
- to_shape = tf.shape(to_mask)
- to_seq_length = to_shape[1]
-
- to_mask = tf.cast(
- tf.reshape(to_mask, [batch_size, 1, to_seq_length]),
- dtype=inputs.dtype)
-
- # We don't assume that `from_tensor` is a mask (although it could be). We
- # don't actually care if we attend *from* padding tokens (only *to* padding)
- # tokens so we create a tensor of all ones.
- #
- # `broadcast_ones` = [batch_size, from_seq_length, 1]
- broadcast_ones = tf.ones(
- shape=[batch_size, from_seq_length, 1], dtype=inputs.dtype)
-
- # Here we broadcast along two dimensions to create the mask.
- mask = broadcast_ones * to_mask
-
- return mask
+ return get_mask(inputs, to_mask) # pyrefly: ignore[bad-argument-type]
diff --git a/official/nlp/modeling/layers/spectral_normalization.py b/official/nlp/modeling/layers/spectral_normalization.py
index 00f07ce4ddb..6a0390bc48d 100644
--- a/official/nlp/modeling/layers/spectral_normalization.py
+++ b/official/nlp/modeling/layers/spectral_normalization.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -30,10 +30,10 @@
"""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class SpectralNormalization(tf.keras.layers.Wrapper):
+class SpectralNormalization(tf_keras.layers.Wrapper):
"""Implements spectral normalization for Dense layer."""
def __init__(self,
@@ -47,7 +47,7 @@ def __init__(self,
"""Initializer.
Args:
- layer: (tf.keras.layers.Layer) A TF Keras layer to apply normalization to.
+ layer: (tf_keras.layers.Layer) A TF Keras layer to apply normalization to.
iteration: (int) The number of power iteration to perform to estimate
weight matrix's singular value.
norm_multiplier: (float) Multiplicative constant to threshold the
@@ -71,24 +71,23 @@ def __init__(self,
if inhere_layer_name:
wrapper_name = layer.name
- if not isinstance(layer, tf.keras.layers.Layer):
- raise ValueError('`layer` must be a `tf.keras.layer.Layer`. '
+ if not isinstance(layer, tf_keras.layers.Layer):
+ raise ValueError('`layer` must be a `tf_keras.layer.Layer`. '
'Observed `{}`'.format(layer))
- super(SpectralNormalization, self).__init__(
+ super().__init__(
layer, name=wrapper_name, **kwargs)
- def build(self, input_shape):
- super(SpectralNormalization, self).build(input_shape)
+ def build(self, input_shape): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ super().build(input_shape)
self.layer.kernel._aggregation = self.aggregation # pylint: disable=protected-access
self._dtype = self.layer.kernel.dtype
self.w = self.layer.kernel
self.w_shape = self.w.shape.as_list()
- self.uv_initializer = tf.initializers.random_normal()
self.v = self.add_weight(
shape=(1, np.prod(self.w_shape[:-1])),
- initializer=self.uv_initializer,
+ initializer=tf.initializers.random_normal(),
trainable=False,
name='v',
dtype=self.dtype,
@@ -96,7 +95,7 @@ def build(self, input_shape):
self.u = self.add_weight(
shape=(1, self.w_shape[-1]),
- initializer=self.uv_initializer,
+ initializer=tf.initializers.random_normal(),
trainable=False,
name='u',
dtype=self.dtype,
@@ -151,7 +150,7 @@ def restore_weights(self):
return self.layer.kernel.assign(self.w)
-class SpectralNormalizationConv2D(tf.keras.layers.Wrapper):
+class SpectralNormalizationConv2D(tf_keras.layers.Wrapper):
"""Implements spectral normalization for Conv2D layer based on [3]."""
def __init__(self,
@@ -165,7 +164,7 @@ def __init__(self,
"""Initializer.
Args:
- layer: (tf.keras.layers.Layer) A TF Keras layer to apply normalization to.
+ layer: (tf_keras.layers.Layer) A TF Keras layer to apply normalization to.
iteration: (int) The number of power iteration to perform to estimate
weight matrix's singular value.
norm_multiplier: (float) Multiplicative constant to threshold the
@@ -190,14 +189,15 @@ def __init__(self,
# Set layer attributes.
layer._name += '_spec_norm'
- if not isinstance(layer, tf.keras.layers.Conv2D):
+ if not isinstance(layer, tf_keras.layers.Conv2D):
raise ValueError(
- 'layer must be a `tf.keras.layer.Conv2D` instance. You passed: {input}'
+ 'layer must be a `tf_keras.layer.Conv2D` instance. You passed: {input}'
.format(input=layer))
- super(SpectralNormalizationConv2D, self).__init__(layer, **kwargs)
+ super().__init__(layer, **kwargs)
- def build(self, input_shape):
- self.layer.build(input_shape)
+ def build(self, input_shape): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ if not self.layer.built:
+ self.layer.build(input_shape)
self.layer.kernel._aggregation = self.aggregation # pylint: disable=protected-access
self._dtype = self.layer.kernel.dtype
@@ -221,11 +221,10 @@ def build(self, input_shape):
self.in_shape = (uv_dim, in_height, in_width, in_channel)
self.out_shape = (uv_dim, out_height, out_width, out_channel)
- self.uv_initializer = tf.initializers.random_normal()
self.v = self.add_weight(
shape=self.in_shape,
- initializer=self.uv_initializer,
+ initializer=tf.initializers.random_normal(),
trainable=False,
name='v',
dtype=self.dtype,
@@ -233,13 +232,13 @@ def build(self, input_shape):
self.u = self.add_weight(
shape=self.out_shape,
- initializer=self.uv_initializer,
+ initializer=tf.initializers.random_normal(),
trainable=False,
name='u',
dtype=self.dtype,
aggregation=self.aggregation)
- super(SpectralNormalizationConv2D, self).build()
+ super().build()
def call(self, inputs):
u_update_op, v_update_op, w_update_op = self.update_weights()
diff --git a/official/nlp/modeling/layers/spectral_normalization_test.py b/official/nlp/modeling/layers/spectral_normalization_test.py
index 2600acd89be..45da78d8ff9 100644
--- a/official/nlp/modeling/layers/spectral_normalization_test.py
+++ b/official/nlp/modeling/layers/spectral_normalization_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,12 +23,12 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import spectral_normalization
-DenseLayer = tf.keras.layers.Dense(10)
-Conv2DLayer = tf.keras.layers.Conv2D(filters=64, kernel_size=3, padding='valid')
+DenseLayer = tf_keras.layers.Dense(10)
+Conv2DLayer = tf_keras.layers.Conv2D(filters=64, kernel_size=3, padding='valid')
def _compute_spectral_norm(weight):
@@ -66,7 +66,7 @@ def test_spec_norm_magnitude(self, input_shape, layer, norm_wrapper):
spectral_norm_computed = _compute_spectral_norm(normalized_kernel)
spectral_norm_expected = self.norm_multiplier
self.assertAllClose(
- spectral_norm_computed, spectral_norm_expected, atol=5e-2)
+ spectral_norm_computed, spectral_norm_expected, atol=1e-1)
# Test that the normalized layer is K-Lipschitz. In particular, if the layer
# is a function f, then ||f(x1) - f(x2)||_2 <= K * ||(x1 - x2)||_2, where K
diff --git a/official/nlp/modeling/layers/talking_heads_attention.py b/official/nlp/modeling/layers/talking_heads_attention.py
index 69605bc8f22..f1ccd943a36 100644
--- a/official/nlp/modeling/layers/talking_heads_attention.py
+++ b/official/nlp/modeling/layers/talking_heads_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,14 +18,16 @@
import string
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
_CHR_IDX = string.ascii_lowercase
-@tf.keras.utils.register_keras_serializable(package="Text")
+@tf_keras.utils.register_keras_serializable(package="Text")
@gin.configurable
-class TalkingHeadsAttention(tf.keras.layers.MultiHeadAttention):
+class TalkingHeadsAttention(tf_keras.layers.MultiHeadAttention):
"""Implements Talking-Heads Attention.
This is an implementation of Talking-Heads Attention based on the paper
@@ -33,7 +35,7 @@ class TalkingHeadsAttention(tf.keras.layers.MultiHeadAttention):
multi-head attention by including linearprojections across the attention-heads
dimension, immediately before and after the softmax operation.
- See the base class `tf.keras.layers.MultiHeadAttention` for more details.
+ See the base class `tf_keras.layers.MultiHeadAttention` for more details.
Args:
num_heads: Number of attention heads.
@@ -87,7 +89,7 @@ def _build_attention(self, qkv_rank):
self._pre_softmax_weight = self.add_weight(
"pre_softmax_weight",
shape=(self._num_heads, self._num_heads),
- initializer=self._kernel_initializer,
+ initializer=tf_utils.clone_initializer(self._kernel_initializer),
regularizer=self._kernel_regularizer,
constraint=self._kernel_constraint,
dtype=self.dtype,
@@ -95,7 +97,7 @@ def _build_attention(self, qkv_rank):
self._post_softmax_weight = self.add_weight(
"post_softmax_weight",
shape=(self._num_heads, self._num_heads),
- initializer=self._kernel_initializer,
+ initializer=tf_utils.clone_initializer(self._kernel_initializer),
regularizer=self._kernel_regularizer,
constraint=self._kernel_constraint,
dtype=self.dtype,
diff --git a/official/nlp/modeling/layers/talking_heads_attention_test.py b/official/nlp/modeling/layers/talking_heads_attention_test.py
index 6f14e2023c2..0bba2d8d401 100644
--- a/official/nlp/modeling/layers/talking_heads_attention_test.py
+++ b/official/nlp/modeling/layers/talking_heads_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,17 +16,13 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import talking_heads_attention
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
# This test is revised base on attention.MultiHeadAttentionTest.
-@keras_parameterized.run_all_keras_modes
-class TalkingHeadsAttentionTest(keras_parameterized.TestCase):
+class TalkingHeadsAttentionTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.named_parameters(
("key_value_same_proj", None, None, [40, 80]),
@@ -40,8 +36,8 @@ def test_non_masked_attention(self, value_dim, output_shape, output_dims):
value_dim=value_dim,
output_shape=output_shape)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
- value = tf.keras.Input(shape=(20, 80))
+ query = tf_keras.Input(shape=(40, 80))
+ value = tf_keras.Input(shape=(20, 80))
output = test_layer(query=query, value=value)
self.assertEqual(output.shape.as_list(), [None] + output_dims)
@@ -50,7 +46,7 @@ def test_non_masked_self_attention(self):
test_layer = talking_heads_attention.TalkingHeadsAttention(
num_heads=12, key_dim=64)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
+ query = tf_keras.Input(shape=(40, 80))
output = test_layer(query=query, value=query)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
@@ -59,7 +55,7 @@ def test_attention_scores(self):
test_layer = talking_heads_attention.TalkingHeadsAttention(
num_heads=12, key_dim=64)
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
+ query = tf_keras.Input(shape=(40, 80))
output, coef = test_layer(query=query, value=query,
return_attention_scores=True)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
@@ -72,13 +68,13 @@ def test_masked_attention(self, use_bias):
num_heads=12, key_dim=2, use_bias=use_bias)
# Create a 3-dimensional input (the first dimension is implicit).
batch_size = 3
- query = tf.keras.Input(shape=(4, 8))
- value = tf.keras.Input(shape=(2, 8))
- mask_tensor = tf.keras.Input(shape=(4, 2))
+ query = tf_keras.Input(shape=(4, 8))
+ value = tf_keras.Input(shape=(2, 8))
+ mask_tensor = tf_keras.Input(shape=(4, 2))
output = test_layer(query=query, value=value, attention_mask=mask_tensor)
# Create a model containing the test layer.
- model = tf.keras.Model([query, value, mask_tensor], output)
+ model = tf_keras.Model([query, value, mask_tensor], output)
# Generate data for the input (non-mask) tensors.
from_data = 10 * np.random.random_sample((batch_size, 4, 8))
@@ -98,10 +94,10 @@ def test_masked_attention(self, use_bias):
self.assertNotAllClose(masked_output_data, unmasked_output_data)
# Tests the layer with three inputs: Q, K, V.
- key = tf.keras.Input(shape=(2, 8))
+ key = tf_keras.Input(shape=(2, 8))
output = test_layer(
query=query, value=value, key=key, attention_mask=mask_tensor)
- model = tf.keras.Model([query, value, key, mask_tensor], output)
+ model = tf_keras.Model([query, value, key, mask_tensor], output)
masked_output_data = model.predict([from_data, to_data, to_data, mask_data])
unmasked_output_data = model.predict(
@@ -122,9 +118,9 @@ def test_initializer(self):
test_layer = talking_heads_attention.TalkingHeadsAttention(
num_heads=12,
key_dim=64,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Create a 3-dimensional input (the first dimension is implicit).
- query = tf.keras.Input(shape=(40, 80))
+ query = tf_keras.Input(shape=(40, 80))
output = test_layer(query=query, value=query)
self.assertEqual(output.shape.as_list(), [None, 40, 80])
diff --git a/official/nlp/modeling/layers/text_layers.py b/official/nlp/modeling/layers/text_layers.py
index 60b2f11a7a6..8de009c6e4e 100644
--- a/official/nlp/modeling/layers/text_layers.py
+++ b/official/nlp/modeling/layers/text_layers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from typing import Any, Dict, List, Mapping, Optional, Text, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
try:
# pytype: disable=import-error
@@ -58,7 +58,7 @@ def fn(x):
return result
-class BertTokenizer(tf.keras.layers.Layer):
+class BertTokenizer(tf_keras.layers.Layer):
"""Wraps TF.Text's BertTokenizer with pre-defined vocab as a Keras Layer.
Attributes:
@@ -120,7 +120,7 @@ def __init__(self, *,
tokenizer_kwargs = dict(tokenizer_kwargs or {})
if lower_case is not None:
tokenizer_kwargs["lower_case"] = lower_case
- self._bert_tokenizer = text.BertTokenizer(self._vocab_table,
+ self._bert_tokenizer = text.BertTokenizer(self._vocab_table, # pyrefly: ignore[missing-attribute]
**tokenizer_kwargs)
@property
@@ -235,7 +235,7 @@ def _create_special_tokens_dict(self, vocab_table, vocab_file):
return result
-class SentencepieceTokenizer(tf.keras.layers.Layer):
+class SentencepieceTokenizer(tf_keras.layers.Layer):
"""Wraps `tf_text.SentencepieceTokenizer` as a Keras Layer.
Attributes:
@@ -317,7 +317,7 @@ def __init__(self,
self._special_tokens_dict = self._create_special_tokens_dict()
def _create_tokenizer(self):
- return text.SentencepieceTokenizer(
+ return text.SentencepieceTokenizer( # pyrefly: ignore[missing-attribute]
model=self._model_serialized_proto,
out_type=tf.int32,
nbest_size=self._nbest_size,
@@ -347,11 +347,11 @@ def call(self, inputs: tf.Tensor):
if self.tokenize_with_offsets:
raise ValueError("`tokenize_with_offsets` is not supported yet when "
"`strip_diacritics` is set to True (b/181866850).")
- inputs = text.normalize_utf8(inputs, "NFD")
+ inputs = text.normalize_utf8(inputs, "NFD") # pyrefly: ignore[missing-attribute]
inputs = tf.strings.regex_replace(inputs, r"\p{Mn}", "")
if self._lower_case:
- inputs = text.case_fold_utf8(inputs)
+ inputs = text.case_fold_utf8(inputs) # pyrefly: ignore[missing-attribute]
# Prepare to reshape the result to work around broken shape inference.
batch_size = tf.shape(inputs)[0]
@@ -437,7 +437,7 @@ def _create_special_tokens_dict(self):
return result
-class BertPackInputs(tf.keras.layers.Layer):
+class BertPackInputs(tf_keras.layers.Layer):
"""Packs tokens into model inputs for BERT."""
def __init__(self,
@@ -573,22 +573,22 @@ def bert_pack_inputs(inputs: Union[tf.RaggedTensor, List[tf.RaggedTensor]],
# fall back to some ad-hoc truncation.
num_special_tokens = len(inputs) + 1
if truncator == "round_robin":
- trimmed_segments = text.RoundRobinTrimmer(seq_length -
+ trimmed_segments = text.RoundRobinTrimmer(seq_length - # pyrefly: ignore[missing-attribute]
num_special_tokens).trim(inputs)
elif truncator == "waterfall":
- trimmed_segments = text.WaterfallTrimmer(
+ trimmed_segments = text.WaterfallTrimmer( # pyrefly: ignore[missing-attribute]
seq_length - num_special_tokens).trim(inputs)
else:
raise ValueError("Unsupported truncator: %s" % truncator)
# Combine segments.
- segments_combined, segment_ids = text.combine_segments(
+ segments_combined, segment_ids = text.combine_segments( # pyrefly: ignore[missing-attribute]
trimmed_segments,
start_of_sequence_id=start_of_sequence_id,
end_of_segment_id=end_of_segment_id)
# Pad to dense Tensors.
- input_word_ids, _ = text.pad_model_inputs(segments_combined, seq_length,
+ input_word_ids, _ = text.pad_model_inputs(segments_combined, seq_length, # pyrefly: ignore[missing-attribute]
pad_value=padding_id)
- input_type_ids, input_mask = text.pad_model_inputs(segment_ids, seq_length,
+ input_type_ids, input_mask = text.pad_model_inputs(segment_ids, seq_length, # pyrefly: ignore[missing-attribute]
pad_value=0)
# Work around broken shape inference.
output_shape = tf.stack([
@@ -602,7 +602,7 @@ def _reshape(t):
input_type_ids=_reshape(input_type_ids))
-class FastWordpieceBertTokenizer(tf.keras.layers.Layer):
+class FastWordpieceBertTokenizer(tf_keras.layers.Layer):
"""A bert tokenizer keras layer using text.FastWordpieceTokenizer.
See details: "Fast WordPiece Tokenization" (https://arxiv.org/abs/2012.15524)
@@ -633,11 +633,11 @@ def __init__(self,
super().__init__(**kwargs)
logging.info("Initialize a FastWordpieceBertTokenizer.")
self.tokenize_with_offsets = tokenize_with_offsets
- self._basic_tokenizer = bert_tokenizer.BasicTokenizer(lower_case=lower_case)
+ self._basic_tokenizer = bert_tokenizer.BasicTokenizer(lower_case=lower_case) # pyrefly: ignore[missing-attribute]
# Read the vocab file into a list of tokens to create `fast_wp_tokenizer`.
self._vocab = [line.rstrip() for line in tf.io.gfile.GFile(vocab_file)]
- self._fast_wp_tokenizer = text.FastWordpieceTokenizer(
+ self._fast_wp_tokenizer = text.FastWordpieceTokenizer( # pyrefly: ignore[missing-attribute]
vocab=self._vocab, token_out_type=tf.int32, no_pretokenization=True)
self._special_tokens_dict = self._create_special_tokens_dict()
diff --git a/official/nlp/modeling/layers/text_layers_test.py b/official/nlp/modeling/layers/text_layers_test.py
index d3bc63352c1..b2b7b081768 100644
--- a/official/nlp/modeling/layers/text_layers_test.py
+++ b/official/nlp/modeling/layers/text_layers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import tempfile
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow import estimator as tf_estimator
from sentencepiece import SentencePieceTrainer
@@ -103,7 +103,7 @@ def input_fn():
with tf.init_scope():
self.assertFalse(tf.executing_eagerly())
# Build a preprocessing Model.
- sentences = tf.keras.layers.Input(shape=[], dtype=tf.string)
+ sentences = tf_keras.layers.Input(shape=[], dtype=tf.string)
bert_tokenizer = text_layers.BertTokenizer(
vocab_file=vocab_file, lower_case=True)
special_tokens_dict = bert_tokenizer.get_special_tokens_dict()
@@ -112,7 +112,7 @@ def input_fn():
tokens = bert_tokenizer(sentences)
packed_inputs = text_layers.BertPackInputs(
4, special_tokens_dict=special_tokens_dict)(tokens)
- preprocessing = tf.keras.Model(sentences, packed_inputs)
+ preprocessing = tf_keras.Model(sentences, packed_inputs)
# Map the dataset.
ds = tf.data.Dataset.from_tensors(
(tf.constant(["abc", "DEF"]), tf.constant([0, 1])))
@@ -214,7 +214,7 @@ def input_fn():
with tf.init_scope():
self.assertFalse(tf.executing_eagerly())
# Build a preprocessing Model.
- sentences = tf.keras.layers.Input(shape=[], dtype=tf.string)
+ sentences = tf_keras.layers.Input(shape=[], dtype=tf.string)
sentencepiece_tokenizer = text_layers.SentencepieceTokenizer(
model_file_path=self._spm_path, lower_case=True, nbest_size=0)
special_tokens_dict = sentencepiece_tokenizer.get_special_tokens_dict()
@@ -223,7 +223,7 @@ def input_fn():
tokens = sentencepiece_tokenizer(sentences)
packed_inputs = text_layers.BertPackInputs(
4, special_tokens_dict=special_tokens_dict)(tokens)
- preprocessing = tf.keras.Model(sentences, packed_inputs)
+ preprocessing = tf_keras.Model(sentences, packed_inputs)
# Map the dataset.
ds = tf.data.Dataset.from_tensors(
(tf.constant(["abc", "DEF"]), tf.constant([0, 1])))
@@ -294,9 +294,9 @@ def test_serialize_deserialize(self):
def test_saving(self):
sentencepiece_tokenizer = text_layers.SentencepieceTokenizer(
model_file_path=self._spm_path, lower_case=True, nbest_size=0)
- inputs = tf.keras.layers.Input([], dtype=tf.string)
+ inputs = tf_keras.layers.Input([], dtype=tf.string)
outputs = sentencepiece_tokenizer(inputs)
- model = tf.keras.Model(inputs, outputs)
+ model = tf_keras.Model(inputs, outputs)
export_path = tempfile.mkdtemp(dir=self.get_temp_dir())
model.save(export_path, signatures={})
@@ -520,7 +520,7 @@ def input_fn():
with tf.init_scope():
self.assertFalse(tf.executing_eagerly())
# Build a preprocessing Model.
- sentences = tf.keras.layers.Input(shape=[], dtype=tf.string)
+ sentences = tf_keras.layers.Input(shape=[], dtype=tf.string)
bert_tokenizer = text_layers.FastWordpieceBertTokenizer(
vocab_file=vocab_file, lower_case=True)
special_tokens_dict = bert_tokenizer.get_special_tokens_dict()
@@ -529,7 +529,7 @@ def input_fn():
tokens = bert_tokenizer(sentences)
packed_inputs = text_layers.BertPackInputs(
4, special_tokens_dict=special_tokens_dict)(tokens)
- preprocessing = tf.keras.Model(sentences, packed_inputs)
+ preprocessing = tf_keras.Model(sentences, packed_inputs)
# Map the dataset.
ds = tf.data.Dataset.from_tensors(
(tf.constant(["abc", "DEF"]), tf.constant([0, 1])))
diff --git a/official/nlp/modeling/layers/tn_expand_condense.py b/official/nlp/modeling/layers/tn_expand_condense.py
index 33247d3689a..e6cb34bf0e7 100644
--- a/official/nlp/modeling/layers/tn_expand_condense.py
+++ b/official/nlp/modeling/layers/tn_expand_condense.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,14 +15,16 @@
"""ExpandCondense tensor network layer used in TN-BERT."""
# pylint: disable=g-classes-have-attributes
from typing import List, Optional, Text, Any, Dict
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-Layer = tf.keras.layers.Layer
-activations = tf.keras.activations
-initializers = tf.keras.initializers
+from official.modeling import tf_utils
+Layer = tf_keras.layers.Layer
+activations = tf_keras.activations
+initializers = tf_keras.initializers
-@tf.keras.utils.register_keras_serializable(package='Text')
+
+@tf_keras.utils.register_keras_serializable(package='Text')
class TNExpandCondense(Layer):
"""A TPU-optimized TensorNetwork layer.
@@ -64,7 +66,7 @@ def __init__(self,
if 'input_shape' not in kwargs and 'input_dim' in kwargs:
kwargs['input_shape'] = (kwargs.pop('input_dim'),)
- super(TNExpandCondense, self).__init__(**kwargs)
+ super().__init__(**kwargs)
assert proj_multiplier in [
2, 4, 6, 8, 10, 12
@@ -84,7 +86,7 @@ def build(self, input_shape: List[int]) -> None:
'The last dimension of the inputs to `TNExpandCondense` '
'should be defined. Found `None`.')
- super(TNExpandCondense, self).build(input_shape)
+ super().build(input_shape)
self.proj_size = self.proj_multiplier * input_shape[-1]
@@ -98,24 +100,24 @@ def build(self, input_shape: List[int]) -> None:
name='w1',
shape=(input_shape[-1], input_shape[-1]),
trainable=True,
- initializer=self.kernel_initializer)
+ initializer=tf_utils.clone_initializer(self.kernel_initializer))
self.w2 = self.add_weight(
name='w2',
shape=(128, (128 * (self.proj_size // input_shape[-1]))),
trainable=True,
- initializer=self.kernel_initializer)
+ initializer=tf_utils.clone_initializer(self.kernel_initializer))
self.w3 = self.add_weight(
name='w3',
shape=(128 * (self.proj_size // input_shape[-1]), 128),
trainable=True,
- initializer=self.kernel_initializer)
+ initializer=tf_utils.clone_initializer(self.kernel_initializer))
self.w4 = self.add_weight(
name='w4',
shape=(input_shape[-1] // 128, 128, input_shape[-1]),
trainable=True,
- initializer=self.kernel_initializer)
+ initializer=tf_utils.clone_initializer(self.kernel_initializer))
if self.use_bias:
self.bias = self.add_weight(
@@ -176,5 +178,5 @@ def get_config(self) -> Dict[Any, Any]:
getattr(self, initializer_arg))
# Get base config
- base_config = super(TNExpandCondense, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
diff --git a/official/nlp/modeling/layers/tn_expand_condense_test.py b/official/nlp/modeling/layers/tn_expand_condense_test.py
index 04f211d119b..6a79f8e4490 100644
--- a/official/nlp/modeling/layers/tn_expand_condense_test.py
+++ b/official/nlp/modeling/layers/tn_expand_condense_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,9 +18,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
-# pylint: disable=g-direct-tensorflow-import
-from tensorflow.python.keras.testing_utils import layer_test
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers.tn_expand_condense import TNExpandCondense
@@ -29,44 +27,25 @@ class TNLayerTest(tf.test.TestCase, parameterized.TestCase):
"""
def setUp(self):
- super(TNLayerTest, self).setUp()
+ super().setUp()
self.labels = np.concatenate((np.ones((50, 1)), np.zeros((50, 1))), axis=0)
def _build_model(self, data, proj_multiple=2):
- model = tf.keras.models.Sequential()
+ model = tf_keras.models.Sequential()
model.add(
TNExpandCondense(
proj_multiplier=proj_multiple,
use_bias=True,
activation='relu',
input_shape=(data.shape[-1],)))
- model.add(tf.keras.layers.Dense(1, activation='sigmoid'))
+ model.add(tf_keras.layers.Dense(1, activation='sigmoid'))
return model
- @parameterized.parameters((768, 6), (1024, 2))
- def test_keras_layer(self, input_dim, proj_multiple):
- self.skipTest('Disable the test for now since it imports '
- 'keras.testing_utils, will reenable this test after we '
- 'fix the b/184578869')
- # TODO(scottzhu): Reenable after fix b/184578869
- data = np.random.normal(size=(100, input_dim))
- data = data.astype(np.float32)
- layer_test(
- TNExpandCondense,
- kwargs={
- 'proj_multiplier': proj_multiple,
- 'input_shape': data.shape
- },
- input_shape=data.shape,
- input_data=data,
- expected_output_shape=(None, data.shape[-1]),
- expected_output_dtype=data.dtype)
-
@parameterized.parameters((768, 6), (1024, 2))
def test_train(self, input_dim, proj_multiple):
+ tf_keras.utils.set_random_seed(0)
data = np.random.randint(10, size=(100, input_dim))
model = self._build_model(data, proj_multiple)
- tf.random.set_seed(0)
model.compile(
optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])
@@ -81,7 +60,7 @@ def test_train(self, input_dim, proj_multiple):
@parameterized.parameters((768, 6), (1024, 2))
def test_weights_change(self, input_dim, proj_multiple):
- tf.random.set_seed(0)
+ tf_keras.utils.set_random_seed(0)
data = np.random.randint(10, size=(100, input_dim))
model = self._build_model(data, proj_multiple)
model.compile(
@@ -111,7 +90,7 @@ def test_output_shape(self, input_dim, proj_multiple):
def test_expandcondense_num_parameters(self, input_dim, proj_multiple):
data = np.random.randint(10, size=(100, input_dim))
proj_size = proj_multiple * data.shape[-1]
- model = tf.keras.models.Sequential()
+ model = tf_keras.models.Sequential()
model.add(
TNExpandCondense(
proj_multiplier=proj_multiple,
@@ -171,7 +150,7 @@ def test_model_save(self, input_dim, proj_multiple):
save_path = os.path.join(self.get_temp_dir(), 'test_model')
model.save(save_path)
- loaded_model = tf.keras.models.load_model(save_path)
+ loaded_model = tf_keras.models.load_model(save_path)
# Compare model predictions and loaded_model predictions
self.assertAllEqual(model.predict(data), loaded_model.predict(data))
diff --git a/official/nlp/modeling/layers/tn_transformer_expand_condense.py b/official/nlp/modeling/layers/tn_transformer_expand_condense.py
index d52525064b8..35dfb927a90 100644
--- a/official/nlp/modeling/layers/tn_transformer_expand_condense.py
+++ b/official/nlp/modeling/layers/tn_transformer_expand_condense.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,17 +14,16 @@
"""TN-BERT TNTransformerExpandCondense employing Expand-Condense layer instead of Dense."""
# pylint: disable=g-classes-have-attributes
-# Import libraries
-
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling.layers.tn_expand_condense import TNExpandCondense
-@tf.keras.utils.register_keras_serializable(package="Text")
+@tf_keras.utils.register_keras_serializable(package="Text")
@gin.configurable
-class TNTransformerExpandCondense(tf.keras.layers.Layer):
+class TNTransformerExpandCondense(tf_keras.layers.Layer):
"""Transformer layer using tensor network Expand-Condense layer.
This layer implements the Transformer from transformer.py, with a single
@@ -77,7 +76,7 @@ def __init__(self,
intermediate_dropout=0.0,
attention_initializer=None,
**kwargs):
- super(TNTransformerExpandCondense, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._num_heads = num_attention_heads
self._intermediate_size = intermediate_size
@@ -85,22 +84,23 @@ def __init__(self,
self._attention_dropout_rate = attention_dropout_rate
self._dropout_rate = dropout_rate
self._output_range = output_range
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
self._intermediate_dropout = intermediate_dropout
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
+ self._attention_initializer = tf_keras.initializers.get(
attention_initializer)
else:
- self._attention_initializer = self._kernel_initializer
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer)
def build(self, input_shape):
input_tensor = input_shape[0] if len(input_shape) == 2 else input_shape
@@ -128,25 +128,25 @@ def build(self, input_shape):
"heads (%d)" % (hidden_size, self._num_heads))
self._attention_head_size = int(hidden_size // self._num_heads)
common_kwargs = dict(
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
kernel_constraint=self._kernel_constraint,
bias_constraint=self._bias_constraint)
- self._attention_layer = tf.keras.layers.MultiHeadAttention(
+ self._attention_layer = tf_keras.layers.MultiHeadAttention(
num_heads=self._num_heads,
key_dim=self._attention_head_size,
dropout=self._attention_dropout_rate,
use_bias=self._use_bias,
kernel_initializer=self._attention_initializer,
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="self_attention",
**common_kwargs)
- self._attention_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self._attention_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -160,15 +160,15 @@ def build(self, input_shape):
kernel_initializer=self._kernel_initializer,
bias_initializer=self._bias_initializer)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
# Use float32 in layernorm for numeric stability.
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype=tf.float32)
- super(TNTransformerExpandCondense, self).build(input_shape)
+ super().build(input_shape)
def get_config(self):
config = {
@@ -185,19 +185,19 @@ def get_config(self):
"output_range":
self._output_range,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
+ tf_keras.constraints.serialize(self._bias_constraint),
"use_bias":
self._use_bias,
"norm_first":
@@ -207,9 +207,9 @@ def get_config(self):
"intermediate_dropout":
self._intermediate_dropout,
"attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer)
+ tf_keras.initializers.serialize(self._attention_initializer)
}
- base_config = super(TNTransformerExpandCondense, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def call(self, inputs):
@@ -220,7 +220,7 @@ def call(self, inputs):
if self._output_range:
target_tensor = input_tensor[:, 0:self._output_range, :]
- attention_mask = attention_mask[:, 0:self._output_range, :]
+ attention_mask = attention_mask[:, 0:self._output_range, :] # pyrefly: ignore[unsupported-operation]
else:
if self._norm_first:
source_tensor = input_tensor
@@ -231,7 +231,7 @@ def call(self, inputs):
query=target_tensor, value=input_tensor, attention_mask=attention_mask)
attention_output = self._attention_dropout(attention_output)
if self._norm_first:
- attention_output = source_tensor + attention_output
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self._attention_layer_norm(target_tensor +
attention_output)
@@ -246,7 +246,7 @@ def call(self, inputs):
# add.
layer_output = tf.cast(layer_output, tf.float32)
if self._norm_first:
- layer_output = source_attention_output + layer_output
+ layer_output = source_attention_output + layer_output # pyrefly: ignore[unbound-name]
else:
layer_output = self._output_layer_norm(layer_output + attention_output)
diff --git a/official/nlp/modeling/layers/tn_transformer_test.py b/official/nlp/modeling/layers/tn_transformer_test.py
index af52661a99b..6dfbca1e64e 100644
--- a/official/nlp/modeling/layers/tn_transformer_test.py
+++ b/official/nlp/modeling/layers/tn_transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,21 +16,17 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers.tn_transformer_expand_condense import TNTransformerExpandCondense
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
@parameterized.named_parameters(('tn', TNTransformerExpandCondense))
-class TransformerLayerTest(keras_parameterized.TestCase):
+class TransformerLayerTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(TransformerLayerTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
def test_layer_creation(self, transformer_cls):
test_layer = transformer_cls(
@@ -40,7 +36,7 @@ def test_layer_creation(self, transformer_cls):
sequence_length = 21
width = 256
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -53,9 +49,9 @@ def test_layer_creation_with_mask(self, transformer_cls):
sequence_length = 21
width = 256
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -68,9 +64,9 @@ def test_layer_creation_with_incorrect_mask_fails(self, transformer_cls):
sequence_length = 21
width = 256
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length - 3))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length - 3))
with self.assertRaisesRegex(ValueError, 'When passing a mask tensor.*'):
_ = test_layer([data_tensor, mask_tensor])
@@ -82,11 +78,11 @@ def test_layer_invocation(self, transformer_cls):
sequence_length = 21
width = 256
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(data_tensor, output_tensor)
+ model = tf_keras.Model(data_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -103,13 +99,13 @@ def test_layer_invocation_with_mask(self, transformer_cls):
sequence_length = 21
width = 256
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -151,7 +147,7 @@ def test_layer_output_range(self, transformer_cls):
new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
def test_layer_invocation_with_float16_dtype(self, transformer_cls):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
test_layer = transformer_cls(
num_attention_heads=16,
intermediate_size=2048,
@@ -159,13 +155,13 @@ def test_layer_invocation_with_float16_dtype(self, transformer_cls):
sequence_length = 21
width = 256
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -183,11 +179,11 @@ def test_transform_with_initializer(self, transformer_cls):
num_attention_heads=16,
intermediate_size=2048,
intermediate_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
sequence_length = 21
width = 256
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output.shape.as_list())
@@ -197,12 +193,12 @@ def test_dynamic_layer_sequence(self, transformer_cls):
num_attention_heads=16,
intermediate_size=2048,
intermediate_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Create a 3-dimensional input (the first dimension is implicit).
width = 256
- input_tensor = tf.keras.Input(shape=(None, width))
+ input_tensor = tf_keras.Input(shape=(None, width))
output_tensor = test_layer(input_tensor)
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
input_length = 17
input_data = np.ones((1, input_length, width))
diff --git a/official/nlp/modeling/layers/transformer.py b/official/nlp/modeling/layers/transformer.py
index 8982a4d3ed0..c9696350408 100644
--- a/official/nlp/modeling/layers/transformer.py
+++ b/official/nlp/modeling/layers/transformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,22 +15,27 @@
"""Keras-based transformer block layer."""
# pylint: disable=g-classes-have-attributes
+from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling.layers import attention
from official.nlp.modeling.layers import multi_channel_attention
from official.nlp.modeling.layers import transformer_encoder_block
from official.nlp.modeling.layers.util import tf_function_if_eager
-@tf.keras.utils.register_keras_serializable(package="Text")
+@tf_keras.utils.register_keras_serializable(package="Text")
class Transformer(transformer_encoder_block.TransformerEncoderBlock):
"""Transformer layer.
This layer implements the Transformer from "Attention Is All You Need".
(https://arxiv.org/abs/1706.03762).
+ **Warning: this layer is deprecated. Please don't use it. Use the
+ `TransformerEncoderBlock` layer instead.**
+
Args:
num_attention_heads: Number of attention heads.
intermediate_size: Size of the intermediate layer.
@@ -97,49 +102,49 @@ def __init__(self,
inner_dropout=intermediate_dropout,
attention_initializer=attention_initializer,
**kwargs)
+ logging.warning("The `Transformer` layer is deprecated. Please directly "
+ "use `TransformerEncoderBlock`.")
def get_config(self):
return {
- "num_attention_heads":
- self._num_heads,
- "intermediate_size":
- self._inner_dim,
- "intermediate_activation":
- self._inner_activation,
- "dropout_rate":
- self._output_dropout_rate,
- "attention_dropout_rate":
- self._attention_dropout_rate,
- "output_range":
- self._output_range,
- "kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
- "bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
- "kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
- "bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
- "activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
- "kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
- "bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
- "use_bias":
- self._use_bias,
- "norm_first":
- self._norm_first,
- "norm_epsilon":
- self._norm_epsilon,
- "intermediate_dropout":
- self._inner_dropout,
- "attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer)
+ "num_attention_heads": self._num_heads,
+ "intermediate_size": self._inner_dim,
+ "intermediate_activation": self._inner_activation,
+ "dropout_rate": self._output_dropout_rate,
+ "attention_dropout_rate": self._attention_dropout_rate,
+ "output_range": self._output_range,
+ "kernel_initializer": tf_utils.serialize_initializer(
+ self._kernel_initializer, use_legacy_format=True
+ ),
+ "bias_initializer": tf_utils.serialize_initializer(
+ self._bias_initializer, use_legacy_format=True
+ ),
+ "kernel_regularizer": tf_utils.serialize_regularizer(
+ self._kernel_regularizer, use_legacy_format=True
+ ),
+ "bias_regularizer": tf_utils.serialize_regularizer(
+ self._bias_regularizer, use_legacy_format=True
+ ),
+ "activity_regularizer": tf_utils.serialize_regularizer(
+ self._activity_regularizer, use_legacy_format=True
+ ),
+ "kernel_constraint": tf_utils.serialize_constraint(
+ self._kernel_constraint, use_legacy_format=True
+ ),
+ "bias_constraint": tf_utils.serialize_constraint(
+ self._bias_constraint, use_legacy_format=True
+ ),
+ "use_bias": self._use_bias,
+ "norm_first": self._norm_first,
+ "norm_epsilon": self._norm_epsilon,
+ "intermediate_dropout": self._inner_dropout,
+ "attention_initializer": tf_utils.serialize_initializer(
+ self._attention_initializer, use_legacy_format=True
+ ),
}
-@tf.keras.utils.register_keras_serializable(package="Text")
+@tf_keras.utils.register_keras_serializable(package="Text")
@gin.configurable
class CompiledTransformer(Transformer):
@@ -148,8 +153,8 @@ def call(self, inputs):
return super().call(inputs)
-@tf.keras.utils.register_keras_serializable(package="Text")
-class TransformerDecoderBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class TransformerDecoderBlock(tf_keras.layers.Layer):
"""Single transformer layer for decoder.
It has three sub-layers:
@@ -181,6 +186,8 @@ class TransformerDecoderBlock(tf.keras.layers.Layer):
intermediate_dropout: Dropout probability for intermediate_dropout_layer.
attention_initializer: Initializer for kernels of attention layers. If set
`None`, attention layers use kernel_initializer as initializer for kernel.
+ self_attention_cls: An optional class to use for self attention.
+ cross_attention_cls: An optional class to use for cross attention.
"""
def __init__(self,
@@ -202,32 +209,44 @@ def __init__(self,
norm_epsilon=1e-12,
intermediate_dropout=0.0,
attention_initializer=None,
+ self_attention_cls=None,
+ cross_attention_cls=None,
**kwargs):
super().__init__(**kwargs)
self.num_attention_heads = num_attention_heads
self.intermediate_size = intermediate_size
- self.intermediate_activation = tf.keras.activations.get(
+ self.intermediate_activation = tf_keras.activations.get(
intermediate_activation)
self.dropout_rate = dropout_rate
self.attention_dropout_rate = attention_dropout_rate
self.multi_channel_cross_attention = multi_channel_cross_attention
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
self._intermediate_dropout = intermediate_dropout
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
+ self._attention_initializer = tf_keras.initializers.get(
attention_initializer)
else:
- self._attention_initializer = self._kernel_initializer
- if self.multi_channel_cross_attention:
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer)
+
+ self._self_attention_cls = self_attention_cls or attention.CachedAttention
+
+ if cross_attention_cls is not None:
+ self._cross_attention_cls = cross_attention_cls
+ if self.multi_channel_cross_attention:
+ logging.warning(
+ "%s will be used for cross attention", cross_attention_cls
+ )
+ elif self.multi_channel_cross_attention:
self._cross_attention_cls = multi_channel_attention.MultiChannelAttention
else:
self._cross_attention_cls = attention.MultiHeadAttention
@@ -244,32 +263,34 @@ def build(self, input_shape):
"heads (%d)" % (hidden_size, self.num_attention_heads))
self.attention_head_size = int(hidden_size) // self.num_attention_heads
common_kwargs = dict(
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
kernel_constraint=self._kernel_constraint,
bias_constraint=self._bias_constraint)
# Self attention.
- self.self_attention = attention.CachedAttention(
+ self.self_attention = self._self_attention_cls(
num_heads=self.num_attention_heads,
key_dim=self.attention_head_size,
dropout=self.attention_dropout_rate,
use_bias=self._use_bias,
- kernel_initializer=self._attention_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._attention_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="self_attention",
**common_kwargs)
- self.self_attention_output_dense = tf.keras.layers.experimental.EinsumDense(
+ self.self_attention_output_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="output",
**common_kwargs)
- self.self_attention_dropout = tf.keras.layers.Dropout(
+ self.self_attention_dropout = tf_keras.layers.Dropout(
rate=self.dropout_rate)
self.self_attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -281,40 +302,44 @@ def build(self, input_shape):
dropout=self.attention_dropout_rate,
output_shape=hidden_size,
use_bias=self._use_bias,
- kernel_initializer=self._attention_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._attention_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="attention/encdec",
**common_kwargs)
- self.encdec_attention_dropout = tf.keras.layers.Dropout(
+ self.encdec_attention_dropout = tf_keras.layers.Dropout(
rate=self.dropout_rate)
self.encdec_attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="attention/encdec_output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype="float32"))
# Feed-forward projection.
- self.intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self.intermediate_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, self.intermediate_size),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="intermediate",
**common_kwargs)
- self.intermediate_activation_layer = tf.keras.layers.Activation(
+ self.intermediate_activation_layer = tf_keras.layers.Activation(
self.intermediate_activation)
- self._intermediate_dropout_layer = tf.keras.layers.Dropout(
+ self._intermediate_dropout_layer = tf_keras.layers.Dropout(
rate=self._intermediate_dropout)
- self.output_dense = tf.keras.layers.experimental.EinsumDense(
+ self.output_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="output",
**common_kwargs)
- self.output_dropout = tf.keras.layers.Dropout(rate=self.dropout_rate)
- self.output_layer_norm = tf.keras.layers.LayerNormalization(
+ self.output_dropout = tf_keras.layers.Dropout(rate=self.dropout_rate)
+ self.output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -323,42 +348,42 @@ def build(self, input_shape):
def get_config(self):
config = {
- "num_attention_heads":
- self.num_attention_heads,
- "intermediate_size":
- self.intermediate_size,
- "intermediate_activation":
- self.intermediate_activation,
- "dropout_rate":
- self.dropout_rate,
- "attention_dropout_rate":
- self.attention_dropout_rate,
- "multi_channel_cross_attention":
- self.multi_channel_cross_attention,
- "kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
- "bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
- "kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
- "bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
- "activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
- "kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
- "bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
- "use_bias":
- self._use_bias,
- "norm_first":
- self._norm_first,
- "norm_epsilon":
- self._norm_epsilon,
- "intermediate_dropout":
- self._intermediate_dropout,
- "attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer)
+ "num_attention_heads": self.num_attention_heads,
+ "intermediate_size": self.intermediate_size,
+ "intermediate_activation": self.intermediate_activation,
+ "dropout_rate": self.dropout_rate,
+ "attention_dropout_rate": self.attention_dropout_rate,
+ "multi_channel_cross_attention": self.multi_channel_cross_attention,
+ "kernel_initializer": tf_utils.serialize_initializer(
+ self._kernel_initializer, use_legacy_format=True
+ ),
+ "bias_initializer": tf_utils.serialize_initializer(
+ self._bias_initializer, use_legacy_format=True
+ ),
+ "kernel_regularizer": tf_utils.serialize_regularizer(
+ self._kernel_regularizer, use_legacy_format=True
+ ),
+ "bias_regularizer": tf_utils.serialize_regularizer(
+ self._bias_regularizer, use_legacy_format=True
+ ),
+ "activity_regularizer": tf_utils.serialize_regularizer(
+ self._activity_regularizer, use_legacy_format=True
+ ),
+ "kernel_constraint": tf_utils.serialize_constraint(
+ self._kernel_constraint, use_legacy_format=True
+ ),
+ "bias_constraint": tf_utils.serialize_constraint(
+ self._bias_constraint, use_legacy_format=True
+ ),
+ "use_bias": self._use_bias,
+ "norm_first": self._norm_first,
+ "norm_epsilon": self._norm_epsilon,
+ "intermediate_dropout": self._intermediate_dropout,
+ "attention_initializer": tf_utils.serialize_initializer(
+ self._attention_initializer, use_legacy_format=True
+ ),
+ "self_attention_cls": self._self_attention_cls,
+ "cross_attention_cls": self._cross_attention_cls,
}
base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
@@ -408,9 +433,10 @@ def call(self, inputs, cache=None, decode_loop_step=None):
# Accesses the 5-th input tensor for the doc-attention probabilities.
cross_attn_inputs["context_attention_weights"] = inputs[-1]
attention_output = self.encdec_attention(**cross_attn_inputs)
+
attention_output = self.encdec_attention_dropout(attention_output)
if self._norm_first:
- attention_output = source_self_attention_output + attention_output
+ attention_output = source_self_attention_output + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self.encdec_attention_layer_norm(
self_attention_output + attention_output)
@@ -425,7 +451,7 @@ def call(self, inputs, cache=None, decode_loop_step=None):
layer_output = self.output_dense(intermediate_output)
layer_output = self.output_dropout(layer_output)
if self._norm_first:
- layer_output = source_attention_output + layer_output
+ layer_output = source_attention_output + layer_output # pyrefly: ignore[unbound-name]
else:
layer_output = self.output_layer_norm(layer_output + attention_output)
return layer_output, cache
diff --git a/official/nlp/modeling/layers/transformer_encoder_block.py b/official/nlp/modeling/layers/transformer_encoder_block.py
index b44d5ab3ac0..54056b44f88 100644
--- a/official/nlp/modeling/layers/transformer_encoder_block.py
+++ b/official/nlp/modeling/layers/transformer_encoder_block.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,19 +13,68 @@
# limitations under the License.
"""Keras-based TransformerEncoder block layer."""
+from typing import Any, Optional, Sequence
+from absl import logging
+import tensorflow as tf, tf_keras
-import tensorflow as tf
-
+from official.modeling import tf_utils
+from official.nlp.modeling.layers import block_sparse_attention
+from official.nlp.modeling.layers import multi_query_attention
+from official.nlp.modeling.layers import talking_heads_attention
from official.nlp.modeling.layers import util
-@tf.keras.utils.register_keras_serializable(package="Text")
-class TransformerEncoderBlock(tf.keras.layers.Layer):
+class RMSNorm(tf_keras.layers.Layer):
+ """Root mean square layer normalization layer."""
+
+ def __init__(
+ self,
+ axis: int | Sequence[int] = -1,
+ epsilon: float = 1e-6,
+ **kwargs,
+ ):
+ """Initializes RMSNorm.
+
+ Args:
+ axis: The axis that the input is normalized over.
+ epsilon: A small value added to the mean square for numerical stability.
+ **kwargs: Keyword arguments passed to the base layer.
+ """
+ super().__init__(**kwargs)
+ self.axis = [axis] if isinstance(axis, int) else axis
+ self.epsilon = epsilon
+
+ def build(self, input_shape: tf.TensorShape | Sequence[int | None]):
+ input_shape = tf.TensorShape(input_shape)
+ scale_shape = [1] * input_shape.rank
+ for dim in self.axis:
+ scale_shape[dim] = input_shape[dim]
+ with tf.name_scope(self.name):
+ self.scale = self.add_weight(
+ name="scale",
+ shape=scale_shape,
+ initializer="ones",
+ experimental_autocast=False,
+ )
+ super().build(input_shape)
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ input_dtype = inputs.dtype
+ inputs = tf.cast(inputs, tf.float32)
+ var = tf.math.reduce_mean(
+ tf.math.square(inputs), axis=self.axis, keepdims=True
+ )
+ outputs = inputs * tf.math.rsqrt(var + self.epsilon) * self.scale
+ return tf.cast(outputs, input_dtype)
+
+
+@tf_keras.utils.register_keras_serializable(package="Text")
+class TransformerEncoderBlock(tf_keras.layers.Layer):
"""TransformerEncoderBlock layer.
This layer implements the Transformer Encoder from
"Attention Is All You Need". (https://arxiv.org/abs/1706.03762),
- which combines a `tf.keras.layers.MultiHeadAttention` layer with a
+ which combines a `tf_keras.layers.MultiHeadAttention` layer with a
two-layer feedforward network.
References:
@@ -34,32 +83,47 @@ class TransformerEncoderBlock(tf.keras.layers.Layer):
Understanding](https://arxiv.org/abs/1810.04805)
"""
- def __init__(self,
- num_attention_heads,
- inner_dim,
- inner_activation,
- output_range=None,
- kernel_initializer="glorot_uniform",
- bias_initializer="zeros",
- kernel_regularizer=None,
- bias_regularizer=None,
- activity_regularizer=None,
- kernel_constraint=None,
- bias_constraint=None,
- use_bias=True,
- norm_first=False,
- norm_epsilon=1e-12,
- output_dropout=0.0,
- attention_dropout=0.0,
- inner_dropout=0.0,
- attention_initializer=None,
- attention_axes=None,
- use_query_residual=True,
- key_dim=None,
- value_dim=None,
- output_last_dim=None,
- diff_q_kv_att_layer_norm=False,
- **kwargs):
+ def __init__(
+ self,
+ num_attention_heads,
+ inner_dim,
+ inner_activation,
+ output_range=None,
+ kernel_initializer="glorot_uniform",
+ bias_initializer="zeros",
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ activity_regularizer=None,
+ kernel_constraint=None,
+ bias_constraint=None,
+ use_bias=True,
+ norm_first=False,
+ norm_epsilon=1e-12,
+ use_rms_norm=False,
+ output_dropout=0.0,
+ attention_dropout=0.0,
+ inner_dropout=0.0,
+ attention_initializer=None,
+ attention_axes=None,
+ use_query_residual=True,
+ key_dim=None,
+ value_dim=None,
+ output_last_dim=None,
+ diff_q_kv_att_layer_norm=False,
+ return_attention_scores=False,
+ num_kv_heads=None,
+ src_block_size=None,
+ tgt_block_size=None,
+ use_sigmoid_attn=False,
+ sigmoid_attn_bias=None,
+ linformer_dim=None,
+ linformer_shared_kv_projection=True,
+ lowrank_query_seq_proj_dim=None,
+ enable_talking_heads=False,
+ enable_gqa_optimization=False,
+ softmax_robust_masking=False,
+ **kwargs,
+ ):
"""Initializes `TransformerEncoderBlock`.
Note: If `output_last_dim` is used and `use_query_residual` is `True`, the
@@ -73,7 +137,7 @@ def __init__(self,
E.g. let's say input dims are `[batch_size, seq_dim, input_last_dim]`.
Scenario 1: If `output_last_dim` is not `None`, then the output dims of this
module would be `[batch_size, seq_dim, output_last_dim]`. Note `key_dim` is
- is overriden by `output_last_dim`.
+ overridden by `output_last_dim`.
Scenario 2: If `output_last_dim` is `None` and `key_dim` is not `None`, then
the output dims of this module would be `[batch_size, seq_dim, key_dim]`.
Scenario 3: If the `output_last_dim` and `key_dim` are both `None`, the
@@ -100,6 +164,7 @@ def __init__(self,
dense layers. If set False, output of attention and intermediate dense
layers is normalized.
norm_epsilon: Epsilon value to initialize normalization layers.
+ use_rms_norm: Whether to use RMSNorm instead of LayerNorm.
output_dropout: Dropout probability for the post-attention and output
dropout.
attention_dropout: Dropout probability for within the attention layer.
@@ -111,62 +176,119 @@ def __init__(self,
attention_axes: axes over which the attention is applied. `None` means
attention over all axes, but batch, heads, and features.
use_query_residual: Toggle to execute residual connection after attention.
- key_dim: `key_dim` for the `tf.keras.layers.MultiHeadAttention`. If
+ key_dim: `key_dim` for the `tf_keras.layers.MultiHeadAttention`. If
`None`, we use the first `input_shape`'s last dim.
- value_dim: `value_dim` for the `tf.keras.layers.MultiHeadAttention`.
+ value_dim: `value_dim` for the `tf_keras.layers.MultiHeadAttention`.
output_last_dim: Final dimension of the output of this module. This also
- dictates the value for the final dimension of the
- multi-head-attention. When it's `None`, we use, in order of decreasing
- precedence, `key_dim` * `num_heads` or the first `input_shape`'s last
- dim as the output's last dim.
+ dictates the value for the final dimension of the multi-head-attention.
+ When it's `None`, we use, in order of decreasing precedence, `key_dim` *
+ `num_heads` or the first `input_shape`'s last dim as the output's last
+ dim.
diff_q_kv_att_layer_norm: If `True`, create a separate attention layer
- norm layer for query and key-value if `norm_first` is `True`. Invalid
- to set to `True` if `norm_first` is `False`.
+ norm layer for query and key-value if `norm_first` is `True`. Invalid to
+ set to `True` if `norm_first` is `False`.
+ return_attention_scores: If `True`, the output of this layer will be a
+ tuple and additionally contain the attention scores in the shape of
+ `[batch_size, num_attention_heads, seq_dim, seq_dim]`.
+ num_kv_heads: Number of key-value heads for multi-query attention. Refer
+ to `multi_query_attention.MultiHeadAttention` for more details.
+ src_block_size: Source block size. Refer to
+ `block_sparse_attention.MultiHeadAttention` for more details.
+ tgt_block_size: Target block size. Refer to
+ `block_sparse_attention.MultiHeadAttention` for more details.
+ use_sigmoid_attn: This param is only used in
+ `block_sparse_attention.MultiHeadAttention`
+ sigmoid_attn_bias: This param is only used in
+ `block_sparse_attention.MultiHeadAttention`
+ linformer_dim: Applies low-rank factorization on keys/values as in
+ https://arxiv.org/pdf/2006.04768.
+ linformer_shared_kv_projection: If set, projection layer is shared for
+ keys and values.
+ lowrank_query_seq_proj_dim: If set, applies a projection layer on query
+ sequence to the given dimension. go/constformer-doc
+ enable_talking_heads: Enable talking heads as in
+ https://arxiv.org/pdf/2003.02436.
+ enable_gqa_optimization: Enable GQA optimization in multi-query attention.
+ This flag is valid only when num_kv_heads is set for GQA.
+ softmax_robust_masking: If true, will use a more numerically robust
+ masking impl for softmax.
**kwargs: keyword arguments.
"""
util.filter_kwargs(kwargs)
super().__init__(**kwargs)
+ # Deprecation warning.
+ if output_range is not None:
+ logging.warning("`output_range` is available as an argument for `call()`."
+ "The `output_range` as __init__ argument is deprecated.")
+
self._num_heads = num_attention_heads
self._inner_dim = inner_dim
self._inner_activation = inner_activation
- self._attention_dropout = attention_dropout
self._attention_dropout_rate = attention_dropout
- self._output_dropout = output_dropout
self._output_dropout_rate = output_dropout
self._output_range = output_range
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
+ self._use_rms_norm = use_rms_norm
self._inner_dropout = inner_dropout
self._use_query_residual = use_query_residual
self._key_dim = key_dim
self._value_dim = value_dim
self._output_last_dim = output_last_dim
self._diff_q_kv_att_layer_norm = diff_q_kv_att_layer_norm
+ self._return_attention_scores = return_attention_scores
+ self._num_kv_heads = num_kv_heads
+ self._src_block_size = src_block_size
+ self._tgt_block_size = tgt_block_size
+ self._use_sigmoid_attn = use_sigmoid_attn
+ self._sigmoid_attn_bias = sigmoid_attn_bias
+ self._linformer_dim = linformer_dim
+ self._linformer_shared_kv_projection = linformer_shared_kv_projection
+ self._lowrank_query_seq_proj_dim = lowrank_query_seq_proj_dim
+ self._enable_talking_heads = enable_talking_heads
+ self._enable_gqa_optimization = enable_gqa_optimization
+ self._softmax_robust_masking = softmax_robust_masking
+ if (
+ self._src_block_size is not None
+ and self._num_kv_heads is not None
+ and self._num_kv_heads != 1
+ ):
+ raise ValueError(
+ "Block sparse attention only supports Multi-query attention.Please"
+ " set num_kv_heads to 1 to enable MQA with block sparse attention."
+ )
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
- attention_initializer)
+ self._attention_initializer = tf_keras.initializers.get(
+ attention_initializer
+ )
else:
- self._attention_initializer = self._kernel_initializer
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer
+ )
self._attention_axes = attention_axes
if self._diff_q_kv_att_layer_norm and not self._norm_first:
- raise ValueError("Setting `diff_q_and_kv_attention_layer_norm` to True"
- "when `norm_first` is False is invalid.")
+ raise ValueError(
+ "Setting `diff_q_and_kv_attention_layer_norm` to True"
+ "when `norm_first` is False is invalid."
+ )
def build(self, input_shape):
if isinstance(input_shape, tf.TensorShape):
input_tensor_shape = input_shape
elif isinstance(input_shape, (list, tuple)):
input_tensor_shape = tf.TensorShape(input_shape[0])
+ elif isinstance(input_shape, dict):
+ input_tensor_shape = tf.TensorShape(input_shape["input_tensor"])
else:
raise ValueError(
"The type of input shape argument is not supported, got: %s" %
@@ -176,9 +298,9 @@ def build(self, input_shape):
einsum_equation = "...bc,cd->...bd"
hidden_size = input_tensor_shape[-1]
if hidden_size % self._num_heads != 0:
- raise ValueError(
+ logging.warning(
"The input size (%d) is not a multiple of the number of attention "
- "heads (%d)" % (hidden_size, self._num_heads))
+ "heads (%d)", hidden_size, self._num_heads)
if self._key_dim is None:
self._key_dim = int(hidden_size // self._num_heads)
if self._output_last_dim is None:
@@ -186,140 +308,290 @@ def build(self, input_shape):
else:
last_output_shape = self._output_last_dim
- common_kwargs = dict(
- bias_initializer=self._bias_initializer,
- kernel_regularizer=self._kernel_regularizer,
- bias_regularizer=self._bias_regularizer,
- activity_regularizer=self._activity_regularizer,
- kernel_constraint=self._kernel_constraint,
- bias_constraint=self._bias_constraint)
- self._attention_layer = tf.keras.layers.MultiHeadAttention(
+ attention_layer_kwargs = dict(
num_heads=self._num_heads,
key_dim=self._key_dim,
value_dim=self._value_dim,
- dropout=self._attention_dropout,
+ dropout=self._attention_dropout_rate,
use_bias=self._use_bias,
kernel_initializer=self._attention_initializer,
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
attention_axes=self._attention_axes,
output_shape=self._output_last_dim,
+ softmax_robust_masking=self._softmax_robust_masking,
name="self_attention",
- **common_kwargs)
- self._attention_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ )
+ common_kwargs = dict(
+ bias_regularizer=self._bias_regularizer,
+ activity_regularizer=self._activity_regularizer,
+ kernel_constraint=self._kernel_constraint,
+ bias_constraint=self._bias_constraint,
+ )
+ if self._src_block_size is not None:
+ if self._enable_talking_heads:
+ raise ValueError(
+ "Block sparse attention does not support talking heads. Please"
+ " set enable_talking_heads to False."
+ )
+ attention_layer_kwargs.update(
+ src_block_size=self._src_block_size,
+ tgt_block_size=self._tgt_block_size,
+ use_sigmoid_attn=self._use_sigmoid_attn,
+ sigmoid_attn_bias=self._sigmoid_attn_bias,
+ num_kv_heads=self._num_kv_heads,
+ name="block_sparse_attention",
+ )
+ attention_fn = block_sparse_attention.MultiHeadAttention
+ elif self._num_kv_heads is not None:
+ attention_layer_kwargs.update(
+ num_kv_heads=self._num_kv_heads,
+ enable_gqa_optimization=self._enable_gqa_optimization,
+ name="multi_query_attention",
+ )
+ if self._enable_talking_heads:
+ attention_fn = (
+ multi_query_attention.TalkingHeadsMultiQueryAttention
+ )
+ else:
+ attention_fn = multi_query_attention.MultiHeadAttention
+ elif self._enable_talking_heads:
+ attention_layer_kwargs.update(
+ name="talking_heads_attention",
+ )
+ attention_fn = (
+ talking_heads_attention.TalkingHeadsAttention
+ )
+ else:
+ attention_fn = tf_keras.layers.MultiHeadAttention
+ self._attention_layer = attention_fn(
+ **attention_layer_kwargs, **common_kwargs
+ )
+ self._attention_dropout = tf_keras.layers.Dropout(
+ rate=self._attention_dropout_rate
+ )
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
- self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
- name="self_attention_layer_norm",
- axis=-1,
- epsilon=self._norm_epsilon,
- dtype=tf.float32))
+ if self._use_rms_norm:
+ self._attention_layer_norm = RMSNorm(
+ epsilon=self._norm_epsilon,
+ name="self_attention_layer_norm",
+ )
+ else:
+ self._attention_layer_norm = tf_keras.layers.LayerNormalization(
+ name="self_attention_layer_norm",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32,
+ )
self._attention_layer_norm_kv = self._attention_layer_norm
if self._diff_q_kv_att_layer_norm:
- self._attention_layer_norm_kv = (
- tf.keras.layers.LayerNormalization(
- name="self_attention_layer_norm_kv",
- axis=-1,
- epsilon=self._norm_epsilon,
- dtype=tf.float32))
-
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ if self._use_rms_norm:
+ self._attention_layer_norm_kv = RMSNorm(
+ epsilon=self._norm_epsilon,
+ name="self_attention_layer_norm_kv",
+ )
+ else:
+ self._attention_layer_norm_kv = tf_keras.layers.LayerNormalization(
+ name="self_attention_layer_norm_kv",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32,
+ )
+
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, self._inner_dim),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
name="intermediate",
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
# TODO(b/154538392): Investigate this.
policy = tf.float32
- self._intermediate_activation_layer = tf.keras.layers.Activation(
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
self._inner_activation, dtype=policy)
- self._inner_dropout_layer = tf.keras.layers.Dropout(
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
rate=self._inner_dropout)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._output_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, last_output_shape),
bias_axes="d",
name="output",
- kernel_initializer=self._kernel_initializer,
- **common_kwargs)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ **common_kwargs,
+ )
+ self._output_dropout = tf_keras.layers.Dropout(
+ rate=self._output_dropout_rate
+ )
# Use float32 in layernorm for numeric stability.
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
- dtype=tf.float32)
-
- super(TransformerEncoderBlock, self).build(input_shape)
+ dtype=tf.float32,
+ )
+ if self._linformer_dim is not None:
+ if self._linformer_shared_kv_projection:
+ low_rank_dim = self._linformer_dim
+ else:
+ low_rank_dim = 2 * self._linformer_dim
+ self._lowrank_kv_projection = tf_keras.layers.EinsumDense(
+ "...bc,cd->...bd",
+ output_shape=(None, low_rank_dim),
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer
+ ),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ name="lowrank_kv_projection",
+ **common_kwargs,
+ )
+ if self._lowrank_query_seq_proj_dim is not None:
+ self._lowrank_query_seq_projection = tf_keras.layers.EinsumDense(
+ # Squash the sequence-length dimension; keep embedding as is.
+ "...ij,ik->...kj",
+ output_shape=(
+ self._lowrank_query_seq_proj_dim,
+ hidden_size,
+ ),
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer
+ ),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ name="constformer_projection",
+ **common_kwargs,
+ )
+ super().build(input_shape)
def get_config(self):
config = {
- "num_attention_heads":
- self._num_heads,
- "inner_dim":
- self._inner_dim,
- "inner_activation":
- self._inner_activation,
- "output_dropout":
- self._output_dropout_rate,
- "attention_dropout":
- self._attention_dropout_rate,
- "output_range":
- self._output_range,
- "kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
- "bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
- "kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
- "bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
- "activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
- "kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
- "bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
- "use_bias":
- self._use_bias,
- "norm_first":
- self._norm_first,
- "norm_epsilon":
- self._norm_epsilon,
- "inner_dropout":
- self._inner_dropout,
- "attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer),
+ "num_attention_heads": self._num_heads,
+ "inner_dim": self._inner_dim,
+ "inner_activation": self._inner_activation,
+ "output_dropout": self._output_dropout_rate,
+ "attention_dropout": self._attention_dropout_rate,
+ "output_range": self._output_range,
+ "kernel_initializer": tf_utils.serialize_initializer(
+ self._kernel_initializer, use_legacy_format=True
+ ),
+ "bias_initializer": tf_utils.serialize_initializer(
+ self._bias_initializer, use_legacy_format=True
+ ),
+ "kernel_regularizer": tf_utils.serialize_regularizer(
+ self._kernel_regularizer, use_legacy_format=True
+ ),
+ "bias_regularizer": tf_utils.serialize_regularizer(
+ self._bias_regularizer, use_legacy_format=True
+ ),
+ "activity_regularizer": tf_utils.serialize_regularizer(
+ self._activity_regularizer, use_legacy_format=True
+ ),
+ "kernel_constraint": tf_utils.serialize_constraint(
+ self._kernel_constraint, use_legacy_format=True
+ ),
+ "bias_constraint": tf_utils.serialize_constraint(
+ self._bias_constraint, use_legacy_format=True
+ ),
+ "use_bias": self._use_bias,
+ "norm_first": self._norm_first,
+ "norm_epsilon": self._norm_epsilon,
+ "inner_dropout": self._inner_dropout,
+ "attention_initializer": tf_utils.serialize_initializer(
+ self._attention_initializer, use_legacy_format=True
+ ),
"attention_axes": self._attention_axes,
- "use_query_residual":
- self._use_query_residual,
- "key_dim":
- self._key_dim,
- "value_dim":
- self._value_dim,
- "output_last_dim":
- self._output_last_dim,
- "diff_q_kv_att_layer_norm":
- self._diff_q_kv_att_layer_norm,
+ "use_query_residual": self._use_query_residual,
+ "key_dim": self._key_dim,
+ "value_dim": self._value_dim,
+ "output_last_dim": self._output_last_dim,
+ "diff_q_kv_att_layer_norm": self._diff_q_kv_att_layer_norm,
+ "num_kv_heads": self._num_kv_heads,
+ "src_block_size": self._src_block_size,
+ "tgt_block_size": self._tgt_block_size,
+ "use_sigmoid_attn": self._use_sigmoid_attn,
+ "sigmoid_attn_bias": self._sigmoid_attn_bias,
+ "linformer_dim": self._linformer_dim,
+ "linformer_shared_kv_projection": self._linformer_shared_kv_projection,
+ "lowrank_query_seq_proj_dim": self._lowrank_query_seq_proj_dim,
+ "softmax_robust_masking": self._softmax_robust_masking
}
- base_config = super(TransformerEncoderBlock, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
- def call(self, inputs):
+ def _apply_lowrank_query_projection(
+ self,
+ query: tf.Tensor,
+ attention_mask: tf.Tensor | None,
+ ):
+ """Applies constformer projection to the source tensor."""
+
+ # Don't project the source tensor if the `lowrank_query_seq_projection`
+ # (constformer) dimension is the same as the input
+ # sequence dimension.
+ if (
+ self._lowrank_query_seq_proj_dim is None
+ or query.shape[1] == self._lowrank_query_seq_proj_dim
+ ):
+ return query
+ # Don't overwrite the attention mask.
+ query = self._apply_query_mask(attention_mask, query)
+ dtype = query.dtype
+ query = self._lowrank_query_seq_projection(query)
+ query = tf.cast(query, dtype)
+ return query
+
+ def _apply_query_mask(
+ self,
+ attention_mask: tf.Tensor | None,
+ query: tf.Tensor,
+ ):
+ """Applying mask before the low rank factorization so that padding is accounted for.
+
+ Applies mask to query only if the dimension of query matches the mask. This
+ is to avoid the projection from happening multiple times while stacking
+ the transformer layers.
+
+ Args:
+ attention_mask: The attention_mask tensor.
+ query: The query tensor.
+
+ Returns:
+ query: The query tensor after applying the mask.
+ """
+ if attention_mask is None:
+ return query
+ if attention_mask.shape[1] != query.shape[1]:
+ # Skip the mask application for query.
+ logging.info(
+ "Skipping mask application on query. Shape mismatch: %s vs %s",
+ attention_mask.shape,
+ query.shape,
+ )
+ return query
+
+ query_mask = tf.cast(attention_mask[:, :, 0], dtype=query.dtype)
+ query = query * tf.expand_dims(query_mask, axis=-1)
+ return query
+
+ def call(self, inputs: Any, output_range: Optional[tf.Tensor] = None) -> Any:
"""Transformer self-attention encoder block call.
Args:
- inputs: a single tensor or a list of tensors.
- `input tensor` as the single sequence of embeddings.
- [`input tensor`, `attention mask`] to have the additional attention
- mask.
- [`query tensor`, `key value tensor`, `attention mask`] to have separate
- input streams for the query, and key/value to the multi-head
- attention.
+ inputs: a single tensor or a list of tensors, or a dictionary. `input
+ tensor` as the single sequence of embeddings. [`input tensor`,
+ `attention mask`] to have the additional attention mask. [`query
+ tensor`, `key value tensor`, `attention mask`] to have separate input
+ streams for the query, and key/value to the multi-head attention. If
+ dictionary is provided, it must contain the following keys:
+ `input_tensor`, `attention_mask`, `key_value_tensor`.
+ output_range: the sequence output range, [0, output_range) for slicing the
+ target sequence. `None` means the target sequence is not sliced. If you
+ would like to have no change to the model training, it is better to only
+ set the `output_range` for serving.
Returns:
An output tensor with the same dimensions as input/query tensor.
@@ -333,30 +605,99 @@ def call(self, inputs):
else:
raise ValueError("Unexpected inputs to %s with length at %d" %
(self.__class__, len(inputs)))
+ elif isinstance(inputs, dict):
+ if not set(inputs.keys()).issubset(
+ set(["input_tensor", "key_value_tensor", "attention_mask"])
+ ):
+ raise ValueError(
+ f"Unexpected keys in input dictionary to: {inputs.keys()}"
+ )
+ try:
+ input_tensor = inputs["input_tensor"]
+ except KeyError as e:
+ raise ValueError(
+ "Missing required key `input_tensor` in input dictionary."
+ ) from e
+ key_value = inputs.get("key_value_tensor", None)
+ attention_mask = inputs.get("attention_mask", None)
else:
input_tensor, key_value, attention_mask = (inputs, None, None)
- if self._output_range:
+ if output_range is None:
+ output_range = self._output_range
+ if output_range:
if self._norm_first:
- source_tensor = input_tensor[:, 0:self._output_range, :]
+ source_tensor = input_tensor[:, 0:output_range, :]
+ if self._use_query_residual:
+ # `source_tensor` is only used for the residual connection.
+ source_tensor = self._apply_lowrank_query_projection(
+ source_tensor, attention_mask
+ )
+
input_tensor = self._attention_layer_norm(input_tensor)
if key_value is not None:
key_value = self._attention_layer_norm_kv(key_value)
- target_tensor = input_tensor[:, 0:self._output_range, :]
+ target_tensor = input_tensor[:, 0:output_range, :]
if attention_mask is not None:
- attention_mask = attention_mask[:, 0:self._output_range, :]
+ attention_mask = attention_mask[:, 0:output_range, :]
else:
if self._norm_first:
source_tensor = input_tensor
+ if self._use_query_residual:
+ # `source_tensor` is only used for the residual connection.
+ source_tensor = self._apply_lowrank_query_projection(
+ source_tensor, attention_mask
+ )
input_tensor = self._attention_layer_norm(input_tensor)
if key_value is not None:
key_value = self._attention_layer_norm_kv(key_value)
target_tensor = input_tensor
+ # Project the query to the constformer dimension.
+ target_tensor = self._apply_lowrank_query_projection(
+ target_tensor, attention_mask
+ )
+
if key_value is None:
key_value = input_tensor
- attention_output = self._attention_layer(
- query=target_tensor, value=key_value, attention_mask=attention_mask)
+
+ key = key_value
+ value = key_value
+ if self._linformer_dim is not None:
+ if attention_mask is not None:
+ # Applying mask before the low rank factorization so that padding is
+ # accounted for.
+ query_mask = tf.cast(attention_mask[:, :, 0], dtype=target_tensor.dtype)
+ if self._lowrank_query_seq_proj_dim is None:
+ target_tensor = target_tensor * tf.expand_dims(query_mask, axis=-1)
+ key_mask = tf.cast(attention_mask[:, 0, :], dtype=target_tensor.dtype)
+ key_value = key_value * tf.expand_dims(key_mask, axis=-1)
+ attention_mask = None
+ key_value = tf.transpose(key_value, [0, 2, 1])
+ key_value = self._lowrank_kv_projection(key_value)
+ if self._linformer_shared_kv_projection:
+ key_value = tf.transpose(key_value, [0, 2, 1])
+ key = key_value
+ value = key_value
+ else:
+ key = tf.transpose(key_value[:, :, : self._linformer_dim], [0, 2, 1])
+ value = tf.transpose(key_value[:, :, self._linformer_dim :], [0, 2, 1])
+
+ if self._return_attention_scores:
+ attention_output, attention_scores = self._attention_layer(
+ query=target_tensor,
+ key=key,
+ value=value,
+ attention_mask=attention_mask,
+ return_attention_scores=True,
+ )
+ else:
+ attention_output = self._attention_layer(
+ query=target_tensor,
+ key=key,
+ value=value,
+ attention_mask=attention_mask,
+ )
attention_output = self._attention_dropout(attention_output)
if self._norm_first:
@@ -364,7 +705,7 @@ def call(self, inputs):
# `self._use_query_residual` into one if clause because else is only for
# `_norm_first == False`.
if self._use_query_residual:
- attention_output = source_tensor + attention_output
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
else:
if self._use_query_residual:
attention_output = target_tensor + attention_output
@@ -380,9 +721,14 @@ def call(self, inputs):
layer_output = self._output_dropout(layer_output)
if self._norm_first:
- return source_attention_output + layer_output
+ layer_output = source_attention_output + layer_output # pyrefly: ignore[unbound-name]
+ else:
+ # During mixed precision training, layer norm output is always fp32 for
+ # now. Casts fp32 for the subsequent add.
+ layer_output = tf.cast(layer_output, tf.float32)
+ layer_output = self._output_layer_norm(layer_output + attention_output)
- # During mixed precision training, layer norm output is always fp32 for now.
- # Casts fp32 for the subsequent add.
- layer_output = tf.cast(layer_output, tf.float32)
- return self._output_layer_norm(layer_output + attention_output)
+ if self._return_attention_scores:
+ return layer_output, attention_scores # pyrefly: ignore[unbound-name]
+ else:
+ return layer_output
diff --git a/official/nlp/modeling/layers/transformer_encoder_block_test.py b/official/nlp/modeling/layers/transformer_encoder_block_test.py
index e8c5ec8a52b..d0ef9fe72dd 100644
--- a/official/nlp/modeling/layers/transformer_encoder_block_test.py
+++ b/official/nlp/modeling/layers/transformer_encoder_block_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,22 +14,22 @@
"""Tests for Keras-based transformer block layer."""
+import math
+
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers.transformer_encoder_block import TransformerEncoderBlock
-@keras_parameterized.run_all_keras_modes
-@parameterized.named_parameters(
- ('base', TransformerEncoderBlock))
-class TransformerEncoderBlockLayerTest(keras_parameterized.TestCase):
+@parameterized.named_parameters(('base', TransformerEncoderBlock))
+class TransformerEncoderBlockLayerTest(
+ tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(TransformerEncoderBlockLayerTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
def test_layer_creation(self, transformer_cls):
test_layer = transformer_cls(
@@ -37,7 +37,7 @@ def test_layer_creation(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -48,24 +48,43 @@ def test_layer_creation_with_mask(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
+ def test_layer_creation_with_dict_inputs(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048, inner_activation='relu'
+ )
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ # Create a 2-dimensional input (the first dimension is implicit).
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
+ inputs = {
+ 'input_tensor': data_tensor,
+ 'key_value_tensor': data_tensor,
+ 'attention_mask': mask_tensor,
+ }
+ output_tensor = test_layer(inputs)
+ # The default output of a transformer layer should be the same as the input.
+ self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
+
def test_layer_invocation(self, transformer_cls):
test_layer = transformer_cls(
num_attention_heads=10, inner_dim=2048, inner_activation='relu')
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(data_tensor, output_tensor)
+ model = tf_keras.Model(data_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -80,13 +99,47 @@ def test_layer_invocation_with_mask(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 6
+ input_data = 10 * np.random.random_sample(
+ (batch_size, sequence_length, width)
+ )
+ # The attention mask should be of shape (batch, from_seq_len, to_seq_len),
+ # which here is (batch, sequence_length, sequence_length)
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length, sequence_length)
+ )
+ _ = model.predict([input_data, mask_data])
+
+ def test_layer_invocation_with_dict_inputs(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048, inner_activation='relu'
+ )
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ # Create a 2-dimensional input (the first dimension is implicit).
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
+ inputs = {
+ 'input_tensor': data_tensor,
+ 'key_value_tensor': data_tensor,
+ 'attention_mask': mask_tensor,
+ }
+
+ output_tensor = test_layer(inputs)
+
+ # Create a model from the test layer.
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -117,18 +170,22 @@ def test_layer_output_range(self, transformer_cls):
new_layer = transformer_cls(
num_attention_heads=10,
inner_dim=2048,
- inner_activation='relu',
- output_range=1)
- _ = new_layer([input_data, mask_data])
+ inner_activation='relu')
+ _ = new_layer([input_data, mask_data], output_range=1)
new_layer.set_weights(test_layer.get_weights())
- new_output_tensor = new_layer([input_data, mask_data])
+ new_output_tensor = new_layer([input_data, mask_data], output_range=1)
self.assertAllClose(
new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
+ output_tensor = test_layer([input_data, mask_data], output_range=1)
+ self.assertAllClose(new_output_tensor, output_tensor, atol=5e-5, rtol=0.003)
+
def test_layer_output_range_without_mask(self, transformer_cls):
test_layer = transformer_cls(
- num_attention_heads=10, inner_dim=2048,
- inner_activation='relu', norm_first=True)
+ num_attention_heads=10,
+ inner_dim=2048,
+ inner_activation='relu',
+ norm_first=True)
sequence_length = 21
width = 80
@@ -143,18 +200,19 @@ def test_layer_output_range_without_mask(self, transformer_cls):
num_attention_heads=10,
inner_dim=2048,
inner_activation='relu',
- output_range=1,
norm_first=True)
- _ = new_layer(input_data)
+ _ = new_layer(input_data, output_range=1)
new_layer.set_weights(test_layer.get_weights())
- new_output_tensor = new_layer(input_data)
+ new_output_tensor = new_layer(input_data, output_range=1)
self.assertAllClose(
new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
def test_layer_output_range_with_pre_norm(self, transformer_cls):
test_layer = transformer_cls(
- num_attention_heads=10, inner_dim=2048,
- inner_activation='relu', norm_first=True)
+ num_attention_heads=10,
+ inner_dim=2048,
+ inner_activation='relu',
+ norm_first=True)
sequence_length = 21
width = 80
@@ -171,28 +229,30 @@ def test_layer_output_range_with_pre_norm(self, transformer_cls):
num_attention_heads=10,
inner_dim=2048,
inner_activation='relu',
- output_range=1,
norm_first=True)
- _ = new_layer([input_data, mask_data])
+ _ = new_layer([input_data, mask_data], output_range=1)
new_layer.set_weights(test_layer.get_weights())
- new_output_tensor = new_layer([input_data, mask_data])
+ new_output_tensor = new_layer([input_data, mask_data], output_range=1)
self.assertAllClose(
new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
+ output_tensor = test_layer([input_data, mask_data], output_range=1)
+ self.assertAllClose(new_output_tensor, output_tensor, atol=5e-5, rtol=0.003)
+
def test_layer_invocation_with_float16_dtype(self, transformer_cls):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
test_layer = transformer_cls(
num_attention_heads=10, inner_dim=2048, inner_activation='relu')
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -210,11 +270,11 @@ def test_transform_with_initializer(self, transformer_cls):
num_attention_heads=10,
inner_dim=2048,
inner_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output.shape.as_list())
@@ -224,12 +284,12 @@ def test_dynamic_layer_sequence(self, transformer_cls):
num_attention_heads=10,
inner_dim=2048,
inner_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Create a 3-dimensional input (the first dimension is implicit).
width = 30
- input_tensor = tf.keras.Input(shape=(None, width))
+ input_tensor = tf_keras.Input(shape=(None, width))
output_tensor = test_layer(input_tensor)
- model = tf.keras.Model(input_tensor, output_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
input_length = 17
input_data = np.ones((1, input_length, width))
@@ -242,7 +302,7 @@ def test_separate_qkv(self, transformer_cls):
num_attention_heads=2,
inner_dim=128,
inner_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Forward path.
q_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
kv_tensor = tf.zeros([2, 8, 16], dtype=tf.float32)
@@ -252,13 +312,12 @@ def test_separate_qkv(self, transformer_cls):
self.assertEqual(output.shape, q_tensor.shape)
-@keras_parameterized.run_all_keras_modes
class TransformerEncoderBlockLayerTestWithoutParams(
- keras_parameterized.TestCase):
+ tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(TransformerEncoderBlockLayerTestWithoutParams, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
def test_raises_invalid_arg_error_when_q_kv_dims_are_different(self):
test_layer = TransformerEncoderBlock(
@@ -274,17 +333,14 @@ def test_raises_invalid_arg_error_when_q_kv_dims_are_different(self):
with self.assertRaises(tf.errors.InvalidArgumentError):
test_layer(inputs)
- @parameterized.named_parameters(
- ('output_range_not_none', 2),
- ('output_range_none', None))
+ @parameterized.named_parameters(('output_range_not_none', 2),
+ ('output_range_none', None))
def test_needs_diff_q_kv_att_layer_norm_to_be_true_for_diff_q_and_kv_dims(
- self,
- output_range):
+ self, output_range):
test_layer = TransformerEncoderBlock(
num_attention_heads=2,
inner_dim=128,
inner_activation='relu',
- output_range=output_range,
norm_first=True)
# Forward path.
q_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
@@ -292,7 +348,7 @@ def test_needs_diff_q_kv_att_layer_norm_to_be_true_for_diff_q_and_kv_dims(
dummy_mask = tf.zeros([2, 4, 8], dtype=tf.float32)
inputs = [q_tensor, kv_tensor, dummy_mask]
with self.assertRaises(tf.errors.InvalidArgumentError):
- test_layer(inputs)
+ test_layer(inputs, output_range=output_range)
test_layer = TransformerEncoderBlock(
num_attention_heads=2,
@@ -303,9 +359,8 @@ def test_needs_diff_q_kv_att_layer_norm_to_be_true_for_diff_q_and_kv_dims(
# Forward path.
test_layer(inputs)
- @parameterized.named_parameters(
- ('norm_first_is_true', True),
- ('norm_first_is_false', False))
+ @parameterized.named_parameters(('norm_first_is_true', True),
+ ('norm_first_is_false', False))
def test_use_query_residual_false_removes_add_op(self, norm_first):
graph_with_res = tf.Graph()
with graph_with_res.as_default():
@@ -314,9 +369,9 @@ def test_use_query_residual_false_removes_add_op(self, norm_first):
inner_dim=128,
inner_activation='relu',
norm_first=norm_first)
- inputs = tf.keras.Input(shape=(None, None, 2))
+ inputs = tf_keras.Input(shape=(None, None, 2))
outputs = layer(inputs)
- tf.keras.Model(inputs=inputs, outputs=outputs)
+ tf_keras.Model(inputs=inputs, outputs=outputs)
graph_without_res = tf.Graph()
with graph_without_res.as_default():
@@ -326,9 +381,9 @@ def test_use_query_residual_false_removes_add_op(self, norm_first):
inner_activation='relu',
norm_first=norm_first,
use_query_residual=False)
- inputs = tf.keras.Input(shape=(None, None, 2))
+ inputs = tf_keras.Input(shape=(None, None, 2))
outputs = layer(inputs)
- tf.keras.Model(inputs=inputs, outputs=outputs)
+ tf_keras.Model(inputs=inputs, outputs=outputs)
graph_with_res_names = {x.name for x in graph_with_res.get_operations()}
graph_without_res_names = {
x.name for x in graph_without_res.get_operations()
@@ -338,15 +393,10 @@ def test_use_query_residual_false_removes_add_op(self, norm_first):
list(graph_with_res_names - graph_without_res_names)[0])
self.assertEmpty(graph_without_res_names - graph_with_res_names)
- @parameterized.named_parameters(
- ('key_dim_is_none', None, 128, 2, 128 // 2),
- ('key_dim_is_not_none', 30, 128, 2, 30))
- def test_key_dim(
- self,
- key_dim,
- q_tensor_last_dim,
- some_num_attention_heads,
- expected):
+ @parameterized.named_parameters(('key_dim_is_none', None, 128, 2, 128 // 2),
+ ('key_dim_is_not_none', 30, 128, 2, 30))
+ def test_key_dim(self, key_dim, q_tensor_last_dim, some_num_attention_heads,
+ expected):
some_inner_dim = 32
some_inner_activation = 'relu'
test_layer = TransformerEncoderBlock(
@@ -360,28 +410,16 @@ def test_key_dim(
dummy_mask = tf.zeros([2, 4, 8], dtype=tf.float32)
test_layer([q_tensor, kv_tensor, dummy_mask])
- self.assertEqual(
- expected,
- test_layer._attention_layer.get_config()['key_dim'])
+ self.assertEqual(expected,
+ test_layer._attention_layer.get_config()['key_dim'])
@parameterized.named_parameters(
- ('output_last_dim_is_none_use_query_residual_false',
- False,
- None,
- 128,
- 128),
- ('output_last_dim_is_none_use_query_residual_true',
- True,
- None,
- 128,
+ ('output_last_dim_is_none_use_query_residual_false', False, None, 128,
128),
+ ('output_last_dim_is_none_use_query_residual_true', True, None, 128, 128),
('output_last_dim_is_not_none', False, 30, 128, 30))
- def test_output_last_dim(
- self,
- use_query_residual,
- output_last_dim,
- q_tensor_last_dim,
- expected):
+ def test_output_last_dim(self, use_query_residual, output_last_dim,
+ q_tensor_last_dim, expected):
some_num_attention_heads = 2
some_inner_dim = 32
some_inner_activation = 'relu'
@@ -401,15 +439,10 @@ def test_output_last_dim(
self.assertEqual(output.numpy().shape[-1], expected)
- @parameterized.named_parameters(
- ('value_dim_is_none', None, 128, 2, 128 // 2),
- ('value_dim_is_not_none', 30, 128, 2, 30))
- def test_value_dim(
- self,
- value_dim,
- q_tensor_last_dim,
- some_num_attention_heads,
- expected):
+ @parameterized.named_parameters(('value_dim_is_none', None, 128, 2, 128 // 2),
+ ('value_dim_is_not_none', 30, 128, 2, 30))
+ def test_value_dim(self, value_dim, q_tensor_last_dim,
+ some_num_attention_heads, expected):
some_inner_dim = 32
some_inner_activation = 'relu'
test_layer = TransformerEncoderBlock(
@@ -423,13 +456,11 @@ def test_value_dim(
dummy_mask = tf.zeros([2, 4, 8], dtype=tf.float32)
test_layer([q_tensor, kv_tensor, dummy_mask])
- self.assertEqual(
- expected,
- test_layer._attention_layer.get_config()['value_dim'])
+ self.assertEqual(expected,
+ test_layer._attention_layer.get_config()['value_dim'])
-@keras_parameterized.run_all_keras_modes
-class TransformerArgumentTest(keras_parameterized.TestCase):
+class TransformerArgumentTest(tf.test.TestCase, parameterized.TestCase):
def test_use_bias_norm_first(self):
num_attention_heads = 2
@@ -444,7 +475,29 @@ def test_use_bias_norm_first(self):
norm_first=True,
norm_epsilon=1e-6,
inner_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
+ minval=0., maxval=1.))
+ # Forward path.
+ dummy_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
+ dummy_mask = tf.zeros([2, 4, 4], dtype=tf.float32)
+ inputs = [dummy_tensor, dummy_mask]
+ output = encoder_block(inputs)
+ self.assertEqual(output.shape, (2, 4, hidden_size))
+
+ def test_use_rms_norm(self):
+ num_attention_heads = 2
+ hidden_size = 16
+ encoder_block = TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=32,
+ inner_activation='relu',
+ output_dropout=0.1,
+ attention_dropout=0.1,
+ use_bias=False,
+ use_rms_norm=True,
+ norm_epsilon=1e-6,
+ inner_dropout=0.1,
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
# Forward path.
dummy_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
@@ -597,7 +650,7 @@ def test_get_config(self):
norm_first=True,
norm_epsilon=1e-6,
inner_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.),
use_query_residual=False,
key_dim=20,
@@ -627,11 +680,320 @@ def test_several_attention_axes(self, attention_axes):
num_cols = 13
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(num_rows, num_cols, width))
+ data_tensor = tf_keras.Input(shape=(num_rows, num_cols, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
+ @parameterized.parameters(
+ {
+ 'output_dropout': 0.1,
+ 'attention_dropout': 0.2,
+ 'inner_dropout': 0.3
+ }, {
+ 'output_dropout': 0.0,
+ 'attention_dropout': 0.2,
+ 'inner_dropout': 0.3
+ }, {
+ 'output_dropout': 0.1,
+ 'attention_dropout': 0.0,
+ 'inner_dropout': 0.3
+ }, {
+ 'output_dropout': 0.1,
+ 'attention_dropout': 0.2,
+ 'inner_dropout': 0.0
+ })
+ def test_dropout_config(self, output_dropout, attention_dropout,
+ inner_dropout):
+ test_layer = TransformerEncoderBlock(
+ num_attention_heads=2,
+ inner_dim=32,
+ inner_activation='relu',
+ output_dropout=output_dropout,
+ attention_dropout=attention_dropout,
+ inner_dropout=inner_dropout)
+ seq_len = 21
+ hidden_size = 512
+ input_tensor = tf_keras.Input(shape=(seq_len, hidden_size))
+ _ = test_layer(input_tensor)
+
+ true_output_dropout = test_layer._output_dropout.get_config()['rate']
+ true_attention_dropout = test_layer._attention_dropout.get_config()['rate']
+ true_inner_dropout = test_layer._inner_dropout_layer.get_config()['rate']
+ self.assertEqual(true_output_dropout, output_dropout)
+ self.assertEqual(true_attention_dropout, attention_dropout)
+ self.assertEqual(true_inner_dropout, inner_dropout)
+
+ @parameterized.named_parameters(
+ (
+ 'return_attention_scores_is_false',
+ False,
+ ),
+ (
+ 'return_attention_scores_is_true',
+ True,
+ ),
+ )
+ def test_return_attention_scores(self, return_attention_scores):
+ num_attention_heads = 7
+ sequence_length = 21
+ width = 80
+
+ test_layer = TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=2048,
+ inner_activation='relu',
+ return_attention_scores=return_attention_scores,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+
+ expected_layer_output_shape = [None, sequence_length, width]
+ expected_attention_scores_shape = [
+ None,
+ num_attention_heads,
+ sequence_length,
+ sequence_length,
+ ]
+
+ if return_attention_scores:
+ self.assertIsInstance(output_tensor, tuple)
+ self.assertLen(output_tensor, 2)
+ # First is the standard output.
+ self.assertEqual(
+ output_tensor[0].shape.as_list(), expected_layer_output_shape
+ )
+ # Second is the attention scores.
+ self.assertEqual(
+ output_tensor[1].shape.as_list(), expected_attention_scores_shape
+ )
+ else:
+ # Only the standard layer output.
+ self.assertEqual(
+ output_tensor.shape.as_list(), expected_layer_output_shape
+ )
+
+ @parameterized.named_parameters(
+ ('mqa', 1),
+ ('gqa', 4),
+ ('talking_heads_mqa', 1, True),
+ ('talking_heads_gqa', 4, True),
+ )
+ def test_attention_with_kv_heads(
+ self, num_kv_heads, enable_talking_heads=False
+ ):
+ num_attention_heads = 8
+ sequence_length = 21
+ width = 80
+
+ test_layer = TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=2048,
+ inner_activation='relu',
+ return_attention_scores=True,
+ num_kv_heads=num_kv_heads,
+ enable_talking_heads=enable_talking_heads,
+ enable_gqa_optimization=True,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+
+ expected_layer_output_shape = [None, sequence_length, width]
+ expected_attention_scores_shape = [
+ None,
+ num_attention_heads,
+ sequence_length,
+ sequence_length,
+ ]
+
+ self.assertIsInstance(output_tensor, tuple)
+ self.assertLen(output_tensor, 2)
+ # First is the standard output.
+ self.assertEqual(
+ output_tensor[0].shape.as_list(), expected_layer_output_shape
+ )
+ # Second is the attention scores.
+ self.assertEqual(
+ output_tensor[1].shape.as_list(), expected_attention_scores_shape
+ )
+
+ @parameterized.named_parameters(
+ ('use_softmax_attn', False),
+ ('use_softmax_attn_mqa', False, 1),
+ ('use_sigmoid_attn', True),
+ ('use_sigmoid_attn_mqa', True, 1),
+ )
+ def test_block_sparse_attention(self, use_sigmoid_attn, num_kv_heads=None):
+ num_attention_heads = 8
+ sequence_length = 21
+ width = 80
+ src_block_size = 7
+ tgt_block_size = 7
+
+ test_layer = TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=2048,
+ inner_activation='relu',
+ return_attention_scores=True,
+ src_block_size=src_block_size,
+ tgt_block_size=tgt_block_size,
+ num_kv_heads=num_kv_heads,
+ use_sigmoid_attn=use_sigmoid_attn,
+ sigmoid_attn_bias=-math.log(sequence_length)
+ if use_sigmoid_attn
+ else None,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+
+ expected_layer_output_shape = [None, sequence_length, width]
+ expected_attention_scores_shape = [
+ None,
+ num_attention_heads,
+ sequence_length//src_block_size,
+ src_block_size,
+ tgt_block_size,
+ ]
+
+ self.assertIsInstance(output_tensor, tuple)
+ self.assertLen(output_tensor, 2)
+ # First is the standard output.
+ self.assertEqual(
+ output_tensor[0].shape.as_list(), expected_layer_output_shape
+ )
+ # Second is the attention scores.
+ self.assertEqual(
+ output_tensor[1].shape.as_list(), expected_attention_scores_shape
+ )
+
+ @parameterized.named_parameters(
+ ('unshared_kv_projection', False),
+ ('shared_kv_projection', True),
+ )
+ def test_low_rank_attention(self, shared_kv_projection):
+ num_attention_heads = 8
+ sequence_length = 21
+ linformer_dim = 7
+ width = 80
+
+ test_layer = TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=2048,
+ inner_activation='relu',
+ return_attention_scores=True,
+ linformer_dim=linformer_dim,
+ linformer_shared_kv_projection=shared_kv_projection,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+
+ expected_layer_output_shape = [None, sequence_length, width]
+ expected_attention_scores_shape = [
+ None,
+ num_attention_heads,
+ sequence_length,
+ linformer_dim,
+ ]
+
+ self.assertIsInstance(output_tensor, tuple)
+ self.assertLen(output_tensor, 2)
+ # First is the standard output.
+ self.assertEqual(
+ output_tensor[0].shape.as_list(), expected_layer_output_shape
+ )
+ # Second is the attention scores.
+ self.assertEqual(
+ output_tensor[1].shape.as_list(), expected_attention_scores_shape
+ )
+
+ def test_low_rank_attention_with_constformer(self):
+ num_attention_heads = 8
+ sequence_length = 21
+ linformer_dim = 7
+ lowrank_query_seq_proj_dim = 10
+ width = 80
+ shared_kv_projection = False
+
+ test_layer = TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=2048,
+ inner_activation='relu',
+ return_attention_scores=True,
+ linformer_dim=linformer_dim,
+ linformer_shared_kv_projection=shared_kv_projection,
+ lowrank_query_seq_proj_dim=lowrank_query_seq_proj_dim,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+
+ # The output from constformer has bottlenecked sequence length.
+ expected_layer_output_shape = [None, lowrank_query_seq_proj_dim, width]
+ # Note that attentions scores with Constformer don't have same
+ # interpretation as the original attention scores, since the sequence
+ # length is squashed.
+ expected_attention_scores_shape = [
+ None,
+ num_attention_heads,
+ lowrank_query_seq_proj_dim,
+ linformer_dim,
+ ]
+
+ self.assertIsInstance(output_tensor, tuple)
+ self.assertLen(output_tensor, 2)
+ # First is the standard output.
+ self.assertEqual(
+ output_tensor[0].shape.as_list(), expected_layer_output_shape
+ )
+ # Second is the attention scores.
+ self.assertEqual(
+ output_tensor[1].shape.as_list(), expected_attention_scores_shape
+ )
+
+ def test_low_rank_attention_with_constformer_no_linformer(self):
+ num_attention_heads = 8
+ sequence_length = 21
+ lowrank_query_seq_proj_dim = 10
+ width = 80
+
+ test_layer = TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=2048,
+ inner_activation='relu',
+ return_attention_scores=True,
+ lowrank_query_seq_proj_dim=lowrank_query_seq_proj_dim,
+ )
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+
+ # The output from constformer has bottlenecked sequence length.
+ expected_layer_output_shape = [None, lowrank_query_seq_proj_dim, width]
+ # Note that attentions scores with Constformer don't have same
+ # interpretation as the original attention scores, since the sequence
+ # length is squashed.
+ expected_attention_scores_shape = [
+ None,
+ num_attention_heads,
+ lowrank_query_seq_proj_dim,
+ sequence_length,
+ ]
+
+ self.assertIsInstance(output_tensor, tuple)
+ self.assertLen(output_tensor, 2)
+ # First is the standard output.
+ self.assertEqual(
+ output_tensor[0].shape.as_list(), expected_layer_output_shape
+ )
+ # Second is the attention scores.
+ self.assertEqual(
+ output_tensor[1].shape.as_list(), expected_attention_scores_shape
+ )
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/nlp/modeling/layers/transformer_scaffold.py b/official/nlp/modeling/layers/transformer_scaffold.py
index c836624b880..c0d178c56c8 100644
--- a/official/nlp/modeling/layers/transformer_scaffold.py
+++ b/official/nlp/modeling/layers/transformer_scaffold.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,14 +17,16 @@
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling.layers import attention
+from official.nlp.modeling.layers import util
-@tf.keras.utils.register_keras_serializable(package="Text")
+@tf_keras.utils.register_keras_serializable(package="Text")
@gin.configurable
-class TransformerScaffold(tf.keras.layers.Layer):
+class TransformerScaffold(tf_keras.layers.Layer):
"""Transformer scaffold layer.
This layer implements the Transformer from "Attention Is All You Need".
@@ -37,8 +39,10 @@ class TransformerScaffold(tf.keras.layers.Layer):
Args:
num_attention_heads: Number of attention heads.
- intermediate_size: Size of the intermediate layer.
- intermediate_activation: Activation for the intermediate layer.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network.
attention_cls: A class to instantiate attention layer, or a layer instance.
attention_cfg: The config with which to instantiate `attention_cls`. Ignored
if attention_cls is a layer instance or None. If `attention_cls` is a
@@ -58,12 +62,16 @@ class TransformerScaffold(tf.keras.layers.Layer):
Ignored if feedforward_cls is a layer instance or is None. If
`feedforward_cls` is a class, but `feedforward_cfg` is None, following
kwargs will be used to instantiate the feedforward instance: {
- "intermediate_size": intermediate_size,
- "intermediate_activation": intermediate_activation,
+ "inner_dim": inner_dim,
+ "inner_activation": inner_activation,
"dropout": dropout_rate,
"name": "feedforward" }.
dropout_rate: Dropout probability for the post-attention and output dropout.
attention_dropout_rate: Dropout probability for within the attention layer.
+ norm_first: Whether to normalize inputs to attention and intermediate
+ dense layers. If set False, output of attention and intermediate dense
+ layers is normalized.
+ norm_epsilon: Epsilon value to initialize normalization layers.
kernel_initializer: Initializer for dense layer kernels.
bias_initializer: Initializer for dense layer biases.
kernel_regularizer: Regularizer for dense layer kernels.
@@ -75,8 +83,8 @@ class TransformerScaffold(tf.keras.layers.Layer):
def __init__(self,
num_attention_heads,
- intermediate_size,
- intermediate_activation,
+ inner_dim=768,
+ inner_activation=tf_utils.get_activation("gelu"),
attention_cls=attention.MultiHeadAttention,
attention_cfg=None,
feedforward_cls=None,
@@ -84,6 +92,7 @@ def __init__(self,
dropout_rate=0.0,
attention_dropout_rate=0.0,
norm_first=False,
+ norm_epsilon=1e-12,
kernel_initializer="glorot_uniform",
bias_initializer="zeros",
kernel_regularizer=None,
@@ -92,29 +101,39 @@ def __init__(self,
kernel_constraint=None,
bias_constraint=None,
**kwargs):
- super(TransformerScaffold, self).__init__(**kwargs)
+ inner_dim = kwargs.pop("intermediate_size", inner_dim)
+ inner_activation = kwargs.pop("inner_activation", inner_activation)
+ util.filter_kwargs(kwargs)
+ super().__init__(**kwargs)
self._attention_cfg = attention_cfg
self._attention_cls = attention_cls
self._feedforward_cls = feedforward_cls
self._feedforward_cfg = feedforward_cfg
self._norm_first = norm_first
+ self._norm_epsilon = norm_epsilon
self._num_heads = num_attention_heads
- self._intermediate_size = intermediate_size
- self._intermediate_activation = intermediate_activation
+ self._inner_dim = inner_dim
+ self._inner_activation = inner_activation
self._attention_dropout_rate = attention_dropout_rate
self._dropout_rate = dropout_rate
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
def build(self, input_shape):
- input_tensor_shape = input_shape[0] if (
- len(input_shape) == 2) else input_shape
- input_tensor_shape = tf.TensorShape(input_tensor_shape)
+ if isinstance(input_shape, tf.TensorShape):
+ input_tensor_shape = input_shape
+ elif isinstance(input_shape, (list, tuple)):
+ input_tensor_shape = tf.TensorShape(input_shape[0])
+ else:
+ raise ValueError(
+ "The type of input shape argument is not supported, got: %s" %
+ type(input_shape))
+
if len(input_tensor_shape.as_list()) != 3:
raise ValueError(
"TransformerScaffold expects a three-dimensional input of "
@@ -127,8 +146,6 @@ def build(self, input_shape):
self._attention_head_size = int(hidden_size // self._num_heads)
common_kwargs = dict(
- kernel_initializer=self._kernel_initializer,
- bias_initializer=self._bias_initializer,
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activity_regularizer=self._activity_regularizer,
@@ -136,8 +153,14 @@ def build(self, input_shape):
bias_constraint=self._bias_constraint)
def get_layer_instance(instance_or_cls, config, default_config):
- if isinstance(instance_or_cls, tf.keras.layers.Layer):
+ if isinstance(instance_or_cls, tf_keras.layers.Layer):
return instance_or_cls
+ elif isinstance(instance_or_cls, dict):
+ return get_layer_instance(
+ tf_keras.utils.deserialize_keras_object(instance_or_cls),
+ config,
+ default_config,
+ )
else:
if config is None:
return instance_or_cls(**default_config)
@@ -145,6 +168,9 @@ def get_layer_instance(instance_or_cls, config, default_config):
return instance_or_cls(**config)
default_attention_cfg = {
+ "kernel_initializer": tf_utils.clone_initializer(
+ self._kernel_initializer),
+ "bias_initializer": tf_utils.clone_initializer(self._bias_initializer),
"num_heads": self._num_heads,
"key_dim": self._attention_head_size,
"dropout": self._attention_dropout_rate,
@@ -158,8 +184,15 @@ def get_layer_instance(instance_or_cls, config, default_config):
if self._feedforward_cls is not None:
default_feedforward_cfg = {
- "intermediate_size": self._intermediate_size,
- "intermediate_activation": self._intermediate_activation,
+ "kernel_initializer": tf_utils.clone_initializer(
+ self._kernel_initializer),
+ "bias_initializer": tf_utils.clone_initializer(
+ self._bias_initializer),
+ "inner_dim": self._inner_dim,
+ "inner_activation": self._inner_activation,
+ # TODO(hongkuny): try to update all ffn block args.
+ "intermediate_size": self._inner_dim,
+ "intermediate_activation": self._inner_activation,
"dropout": self._dropout_rate,
"name": "feedforward",
}
@@ -173,100 +206,119 @@ def get_layer_instance(instance_or_cls, config, default_config):
# self._dropout_rate controls dropout rates at two places:
# after attention, and after FFN.
- self._attention_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self._attention_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
- epsilon=1e-12,
+ epsilon=self._norm_epsilon,
dtype=tf.float32))
if self._feedforward_block is None:
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
- output_shape=(None, self._intermediate_size),
+ output_shape=(None, self._inner_dim),
bias_axes="d",
name="intermediate",
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
# TODO(b/154538392): Investigate this.
policy = tf.float32
- self._intermediate_activation_layer = tf.keras.layers.Activation(
- self._intermediate_activation, dtype=policy)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
+ self._inner_activation, dtype=policy)
+ self._output_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
name="output",
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
**common_kwargs)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
# Use float32 in layernorm for numeric stability.
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
- name="output_layer_norm", axis=-1, epsilon=1e-12, dtype=tf.float32)
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
+ name="output_layer_norm",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32)
- super(TransformerScaffold, self).build(input_shape)
+ super().build(input_shape)
logging.info("%s configs: %s", self.__class__.__name__, self.get_config())
def get_config(self):
config = {
- "attention_cls":
- self._attention_layer,
- "feedforward_cls":
- self._feedforward_block,
- "num_attention_heads":
- self._num_heads,
- "intermediate_size":
- self._intermediate_size,
- "intermediate_activation":
- self._intermediate_activation,
- "dropout_rate":
- self._dropout_rate,
- "attention_dropout_rate":
- self._attention_dropout_rate,
- "norm_first":
- self._norm_first,
- "kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
- "bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
- "kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
- "bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
- "activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
- "kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
- "bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint)
+ "attention_cls": self._attention_layer,
+ "feedforward_cls": self._feedforward_block,
+ "num_attention_heads": self._num_heads,
+ "inner_dim": self._inner_dim,
+ "inner_activation": self._inner_activation,
+ "dropout_rate": self._dropout_rate,
+ "attention_dropout_rate": self._attention_dropout_rate,
+ "norm_first": self._norm_first,
+ "norm_epsilon": self._norm_epsilon,
+ "kernel_initializer": tf_utils.serialize_initializer(
+ self._kernel_initializer, use_legacy_format=True
+ ),
+ "bias_initializer": tf_utils.serialize_initializer(
+ self._bias_initializer, use_legacy_format=True
+ ),
+ "kernel_regularizer": tf_utils.serialize_regularizer(
+ self._kernel_regularizer, use_legacy_format=True
+ ),
+ "bias_regularizer": tf_utils.serialize_regularizer(
+ self._bias_regularizer, use_legacy_format=True
+ ),
+ "activity_regularizer": tf_utils.serialize_regularizer(
+ self._activity_regularizer, use_legacy_format=True
+ ),
+ "kernel_constraint": tf_utils.serialize_constraint(
+ self._kernel_constraint, use_legacy_format=True
+ ),
+ "bias_constraint": tf_utils.serialize_constraint(
+ self._bias_constraint, use_legacy_format=True
+ ),
}
- base_config = super(TransformerScaffold, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def call(self, inputs, training=None):
- if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
- input_tensor, attention_mask = inputs
+ if isinstance(inputs, (list, tuple)):
+ if len(inputs) == 2:
+ input_tensor, attention_mask = inputs
+ key_value = None
+ elif len(inputs) == 3:
+ input_tensor, key_value, attention_mask = inputs
+ else:
+ raise ValueError("Unexpected inputs to %s with length at %d" %
+ (self.__class__, len(inputs)))
else:
- input_tensor, attention_mask = (inputs, None)
+ input_tensor, key_value, attention_mask = (inputs, None, None)
+
+ if key_value is None:
+ key_value = input_tensor
if self._norm_first:
source_tensor = input_tensor
input_tensor = self._attention_layer_norm(input_tensor, training=training)
attention_output = self._attention_layer(
- query=input_tensor, value=input_tensor, attention_mask=attention_mask,
+ query=input_tensor, value=key_value, attention_mask=attention_mask,
training=training)
attention_output = self._attention_dropout(attention_output,
training=training)
if self._norm_first:
- attention_output = source_tensor + attention_output
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self._attention_layer_norm(input_tensor +
attention_output,
@@ -287,7 +339,7 @@ def call(self, inputs, training=None):
# add.
layer_output = tf.cast(layer_output, tf.float32)
if self._norm_first:
- layer_output = source_attention_output + layer_output
+ layer_output = source_attention_output + layer_output # pyrefly: ignore[unbound-name]
else:
layer_output = self._output_layer_norm(layer_output + attention_output,
training=training)
@@ -296,9 +348,11 @@ def call(self, inputs, training=None):
# if norm_first, assume the feedforward block will not apply layer norm
layer_output = self._feedforward_block(attention_output,
training=training)
- layer_output += source_attention_output
+ layer_output += source_attention_output # pyrefly: ignore[unbound-name]
else:
- # if not norm_first, assume that the feedforwad does apply layer norm
+ # Attention: if not norm_first, assume that the feedforwad does apply
+ # layer norm. The feedford also apply residual connection. Please
+ # read the `GatedFeedforward` as a concrete example.
layer_output = self._feedforward_block(attention_output,
training=training)
diff --git a/official/nlp/modeling/layers/transformer_scaffold_test.py b/official/nlp/modeling/layers/transformer_scaffold_test.py
index cf19331b317..186776892ae 100644
--- a/official/nlp/modeling/layers/transformer_scaffold_test.py
+++ b/official/nlp/modeling/layers/transformer_scaffold_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,8 @@
"""Tests for Keras-based transformer block layer."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import attention
from official.nlp.modeling.layers import transformer_scaffold
@@ -26,7 +25,7 @@
# at any point, the list passed to the config object will be filled with a
# boolean 'True'. We register this class as a Keras serializable so we can
# test serialization below.
-@tf.keras.utils.register_keras_serializable(package='TestOnlyAttention')
+@tf_keras.utils.register_keras_serializable(package='TestOnlyAttention')
class ValidatedAttentionLayer(attention.MultiHeadAttention):
def __init__(self, call_list, **kwargs):
@@ -48,8 +47,8 @@ def get_config(self):
# at any point, the list passed to the config object will be filled with a
# boolean 'True'. We register this class as a Keras serializable so we can
# test serialization below.
-@tf.keras.utils.register_keras_serializable(package='TestOnlyFeedforward')
-class ValidatedFeedforwardLayer(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='TestOnlyFeedforward')
+class ValidatedFeedforwardLayer(tf_keras.layers.Layer):
def __init__(self, call_list, activation, **kwargs):
super(ValidatedFeedforwardLayer, self).__init__(**kwargs)
@@ -57,8 +56,8 @@ def __init__(self, call_list, activation, **kwargs):
self.activation = activation
def build(self, input_shape):
- hidden_size = input_shape.as_list()[-1]
- self._feedforward_dense = tf.keras.layers.experimental.EinsumDense(
+ hidden_size = input_shape[-1]
+ self._feedforward_dense = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
@@ -76,14 +75,11 @@ def get_config(self):
return config
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class TransformerLayerTest(keras_parameterized.TestCase):
+class TransformerLayerTest(tf.test.TestCase):
def tearDown(self):
super(TransformerLayerTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
def test_layer_creation(self):
sequence_length = 21
@@ -99,11 +95,11 @@ def test_layer_creation(self):
attention_cls=ValidatedAttentionLayer,
attention_cfg=attention_layer_cfg,
num_attention_heads=10,
- intermediate_size=2048,
- intermediate_activation='relu')
+ inner_dim=2048,
+ inner_activation='relu')
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -134,11 +130,11 @@ def test_layer_creation_with_feedforward_cls(self):
feedforward_cls=ValidatedFeedforwardLayer,
feedforward_cfg=feedforward_layer_cfg,
num_attention_heads=10,
- intermediate_size=None,
- intermediate_activation=None)
+ inner_dim=None,
+ inner_activation=None)
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -165,13 +161,13 @@ def test_layer_creation_with_mask(self):
attention_cls=ValidatedAttentionLayer,
attention_cfg=attention_layer_cfg,
num_attention_heads=10,
- intermediate_size=2048,
- intermediate_activation='relu')
+ inner_dim=2048,
+ inner_activation='relu')
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -194,15 +190,15 @@ def test_layer_invocation(self):
attention_cls=ValidatedAttentionLayer,
attention_cfg=attention_layer_cfg,
num_attention_heads=10,
- intermediate_size=2048,
- intermediate_activation='relu')
+ inner_dim=2048,
+ inner_activation='relu')
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(data_tensor, output_tensor)
+ model = tf_keras.Model(data_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -236,17 +232,17 @@ def test_layer_invocation_with_feedforward_cls(self):
attention_cfg=attention_layer_cfg,
feedforward_cls=feedforward_layer,
num_attention_heads=10,
- intermediate_size=None,
- intermediate_activation=None)
+ inner_dim=None,
+ inner_activation=None)
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -280,17 +276,17 @@ def test_layer_invocation_with_mask(self):
attention_cls=ValidatedAttentionLayer,
attention_cfg=attention_layer_cfg,
num_attention_heads=10,
- intermediate_size=2048,
- intermediate_activation='relu')
+ inner_dim=2048,
+ inner_activation='relu')
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -308,7 +304,7 @@ def test_layer_invocation_with_mask(self):
self.assertTrue(call_list[0], "The passed layer class wasn't instantiated.")
def test_layer_invocation_with_float16_dtype(self):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
sequence_length = 21
width = 80
@@ -322,17 +318,17 @@ def test_layer_invocation_with_float16_dtype(self):
attention_cls=ValidatedAttentionLayer,
attention_cfg=attention_layer_cfg,
num_attention_heads=10,
- intermediate_size=2048,
- intermediate_activation='relu')
+ inner_dim=2048,
+ inner_activation='relu')
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -363,12 +359,12 @@ def test_transform_with_initializer(self):
attention_cls=ValidatedAttentionLayer,
attention_cfg=attention_layer_cfg,
num_attention_heads=10,
- intermediate_size=2048,
- intermediate_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ inner_dim=2048,
+ inner_activation='relu',
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output.shape.as_list())
@@ -392,17 +388,17 @@ def test_layer_restoration_from_config(self):
attention_cls=ValidatedAttentionLayer,
attention_cfg=attention_layer_cfg,
num_attention_heads=10,
- intermediate_size=2048,
- intermediate_activation='relu')
+ inner_dim=2048,
+ inner_activation='relu')
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -421,7 +417,7 @@ def test_layer_restoration_from_config(self):
# Create a new model from the old config, and copy the weights. These models
# should have identical outputs.
- new_model = tf.keras.Model.from_config(serialized_data)
+ new_model = tf_keras.Model.from_config(serialized_data)
new_model.set_weights(model.get_weights())
output = new_model.predict([input_data, mask_data])
@@ -458,17 +454,17 @@ def test_layer_with_feedforward_cls_restoration_from_config(self):
feedforward_cls=ValidatedFeedforwardLayer,
feedforward_cfg=feedforward_layer_cfg,
num_attention_heads=10,
- intermediate_size=None,
- intermediate_activation=None)
+ inner_dim=None,
+ inner_activation=None)
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -484,7 +480,7 @@ def test_layer_with_feedforward_cls_restoration_from_config(self):
serialized_data = model.get_config()
# Create a new model from the old config, and copy the weights. These models
# should have identical outputs.
- new_model = tf.keras.Model.from_config(serialized_data)
+ new_model = tf_keras.Model.from_config(serialized_data)
new_model.set_weights(model.get_weights())
output = new_model.predict([input_data, mask_data])
diff --git a/official/nlp/modeling/layers/transformer_test.py b/official/nlp/modeling/layers/transformer_test.py
index 8ee11b9196f..b35cef9061f 100644
--- a/official/nlp/modeling/layers/transformer_test.py
+++ b/official/nlp/modeling/layers/transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,8 @@
"""Tests for Keras-based transformer block layer."""
-import tensorflow as tf
-
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
from official.nlp.modeling.layers import transformer
@@ -31,8 +30,7 @@ def _create_cache(batch_size, init_decode_length, num_heads, head_size):
}
-@keras_parameterized.run_all_keras_modes
-class TransformerDecoderBlockTest(keras_parameterized.TestCase):
+class TransformerDecoderBlockTest(parameterized.TestCase):
def test_decoder_block_with_cache(self):
num_attention_heads = 2
@@ -66,7 +64,7 @@ def test_use_bias_norm_first(self):
norm_first=True,
norm_epsilon=1e-6,
intermediate_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
# Forward path.
dummy_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
@@ -87,13 +85,67 @@ def test_get_config(self):
norm_first=True,
norm_epsilon=1e-6,
intermediate_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
decoder_block_config = decoder_block.get_config()
new_decoder_block = transformer.TransformerDecoderBlock.from_config(
decoder_block_config)
self.assertEqual(decoder_block_config, new_decoder_block.get_config())
+ @parameterized.named_parameters(
+ ('default', False, False),
+ ('custom_self_attention', True, False),
+ ('custom_cross_attention', False, True),
+ ('custom_self_and_cross_attention', True, True),
+ )
+ def test_decoder_block_with_self_attention_override(
+ self, custom_self_attention, custom_cross_attention
+ ):
+ self_attention_called = False
+ cross_attention_called = False
+
+ class SelfAttention:
+ """Dummy implementation of custom attention."""
+
+ def __init__(self, *args, **kwargs):
+ pass
+
+ def __call__(self, query, value, attention_mask, cache, decode_loop_step):
+ nonlocal self_attention_called
+ self_attention_called = True
+ return query, cache
+
+ class CrossAttention:
+ """Dummy implementation of custom attention."""
+
+ def __init__(self, *args, **kwargs):
+ pass
+
+ def __call__(self, query, value, attention_mask):
+ nonlocal cross_attention_called
+ cross_attention_called = True
+ return query
+
+ num_attention_heads = 2
+ hidden_size = 16
+ decoder_block = transformer.TransformerDecoderBlock(
+ num_attention_heads=num_attention_heads,
+ intermediate_size=32,
+ intermediate_activation='relu',
+ dropout_rate=0.1,
+ attention_dropout_rate=0.1,
+ self_attention_cls=SelfAttention if custom_self_attention else None,
+ cross_attention_cls=CrossAttention if custom_cross_attention else None,
+ )
+ # Forward path.
+ dummy_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
+ dummy_mask = tf.zeros([2, 4, 4], dtype=tf.float32)
+ inputs = [dummy_tensor, dummy_tensor, dummy_mask, dummy_mask]
+ output, _ = decoder_block(inputs)
+ self.assertEqual(output.shape, (2, 4, hidden_size))
+ self.assertEqual(self_attention_called, custom_self_attention)
+ self.assertEqual(cross_attention_called, custom_cross_attention)
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/nlp/modeling/layers/transformer_xl.py b/official/nlp/modeling/layers/transformer_xl.py
index 25355d650eb..330a1cc636e 100644
--- a/official/nlp/modeling/layers/transformer_xl.py
+++ b/official/nlp/modeling/layers/transformer_xl.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,8 +16,9 @@
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling.layers import relative_attention
@@ -50,8 +51,8 @@ def _cache_memory(current_state, previous_state, memory_length, reuse_length=0):
return tf.stop_gradient(new_mem)
-@tf.keras.utils.register_keras_serializable(package="Text")
-class TransformerXLBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class TransformerXLBlock(tf_keras.layers.Layer):
"""Transformer XL block.
This implements a Transformer XL block from "Transformer-XL: Attentive
@@ -102,7 +103,7 @@ def __init__(self,
**kwargs):
"""Initializes TransformerXLBlock layer."""
- super(TransformerXLBlock, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._vocab_size = vocab_size
self._num_heads = num_attention_heads
self._head_size = head_size
@@ -148,39 +149,39 @@ def build(self, input_shape):
value_dim=self._head_size,
dropout=self._attention_dropout_rate,
use_bias=False,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
name="rel_attn")
- self._attention_dropout = tf.keras.layers.Dropout(
+ self._attention_dropout = tf_keras.layers.Dropout(
rate=self._attention_dropout_rate)
- self._attention_layer_norm = tf.keras.layers.LayerNormalization(
+ self._attention_layer_norm = tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype=tf.float32)
- self._inner_dense = tf.keras.layers.experimental.EinsumDense(
+ self._inner_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, self._inner_size),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
name="inner")
- self._inner_activation_layer = tf.keras.layers.Activation(
+ self._inner_activation_layer = tf_keras.layers.Activation(
self._inner_activation)
- self._inner_dropout_layer = tf.keras.layers.Dropout(
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
rate=self._inner_dropout)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._output_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
name="output",
- kernel_initializer=self._kernel_initializer)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer))
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon)
- super(TransformerXLBlock, self).build(input_shape)
+ super().build(input_shape)
def get_config(self):
config = {
@@ -209,7 +210,7 @@ def get_config(self):
"inner_dropout":
self._inner_dropout,
}
- base_config = super(TransformerXLBlock, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def call(self,
@@ -318,7 +319,7 @@ def call(self,
return attention_output
-class TransformerXL(tf.keras.layers.Layer):
+class TransformerXL(tf_keras.layers.Layer):
"""Transformer XL.
This layer combines multiple Transformer XL blocks from "Transformer-XL:
@@ -370,7 +371,7 @@ def __init__(self,
inner_activation="relu",
**kwargs):
"""Initializes TransformerXL."""
- super(TransformerXL, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._vocab_size = vocab_size
self._initializer = initializer
@@ -398,17 +399,17 @@ def __init__(self,
"content_attention_bias",
shape=attention_bias_shape,
dtype=tf.float32,
- initializer=self._initializer)
+ initializer=tf_utils.clone_initializer(self._initializer))
self.positional_attention_bias = self.add_weight(
"positional_attention_bias",
shape=attention_bias_shape,
dtype=tf.float32,
- initializer=self._initializer)
+ initializer=tf_utils.clone_initializer(self._initializer))
self.segment_attention_bias = self.add_weight(
"segment_attention_bias",
shape=attention_bias_shape,
dtype=tf.float32,
- initializer=self._initializer)
+ initializer=tf_utils.clone_initializer(self._initializer))
self.transformer_xl_layers = []
for i in range(self._num_layers):
@@ -427,7 +428,7 @@ def __init__(self,
kernel_initializer="variance_scaling",
name="layer_%d" % i))
- self.output_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self.output_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
def get_config(self):
config = {
@@ -460,7 +461,7 @@ def get_config(self):
"inner_activation":
self._inner_activation,
}
- base_config = super(TransformerXL, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def call(self,
@@ -523,7 +524,7 @@ def call(self,
segment_attention_bias = (self.segment_attention_bias
if self._tie_attention_biases
else self.segment_attention_bias[i])
- segment_encoding = segment_embedding[i]
+ segment_encoding = segment_embedding[i] # pyrefly: ignore[unsupported-operation]
content_attention_bias = (self.content_attention_bias
if self._tie_attention_biases
diff --git a/official/nlp/modeling/layers/transformer_xl_test.py b/official/nlp/modeling/layers/transformer_xl_test.py
index 375d96ec8ff..acf026d2712 100644
--- a/official/nlp/modeling/layers/transformer_xl_test.py
+++ b/official/nlp/modeling/layers/transformer_xl_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,11 @@
"""Tests for Transformer XL."""
+from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.layers import transformer_xl
@@ -115,8 +115,7 @@ def create_mock_transformer_xl_data(
return data
-@keras_parameterized.run_all_keras_modes
-class TransformerXLBlockTest(keras_parameterized.TestCase):
+class TransformerXLBlockTest(tf.test.TestCase, parameterized.TestCase):
@combinations.generate(combinations.combine(
memory_length=[0, 4],
@@ -186,8 +185,7 @@ def test_get_config(self):
self.assertEqual(transformer_xl_block_config, new_block.get_config())
-@keras_parameterized.run_all_keras_modes
-class TransformerXLTest(keras_parameterized.TestCase):
+class TransformerXLTest(tf.test.TestCase, parameterized.TestCase):
@combinations.generate(combinations.combine(
two_stream=[True, False],
@@ -233,7 +231,7 @@ def test_transformer_xl(
inner_size=inner_size,
dropout_rate=0.,
attention_dropout_rate=0.,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
two_stream=two_stream,
tie_attention_biases=tie_attention_biases,
memory_length=memory_length,
@@ -246,7 +244,7 @@ def test_transformer_xl(
else:
self.assertEqual(attention_output.shape,
[batch_size, seq_length, hidden_size])
- self.assertEqual(len(cached_memory_states), num_layers)
+ self.assertLen(cached_memory_states, num_layers)
def test_get_config(self):
transformer_xl_layer = transformer_xl.TransformerXL(
@@ -258,7 +256,7 @@ def test_get_config(self):
inner_size=12,
dropout_rate=0.,
attention_dropout_rate=0.,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
two_stream=False,
tie_attention_biases=True,
memory_length=0,
diff --git a/official/nlp/modeling/layers/util.py b/official/nlp/modeling/layers/util.py
index a3a7820712a..f686ff670d5 100644
--- a/official/nlp/modeling/layers/util.py
+++ b/official/nlp/modeling/layers/util.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class TfFunctionIfEagerDecorator(object):
diff --git a/official/nlp/modeling/losses/__init__.py b/official/nlp/modeling/losses/__init__.py
index 2cb70ee5e5a..0bcb5924c61 100644
--- a/official/nlp/modeling/losses/__init__.py
+++ b/official/nlp/modeling/losses/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy.py b/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy.py
index 81c9b38c544..fcd3cf7209f 100644
--- a/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy.py
+++ b/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Weighted sparse categorical cross-entropy losses."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _adjust_labels(labels, predictions):
@@ -61,7 +61,7 @@ def loss(labels, predictions, weights=None, from_logits=False):
labels, predictions = _adjust_labels(labels, predictions)
_validate_rank(labels, predictions, weights)
- example_losses = tf.keras.losses.sparse_categorical_crossentropy(
+ example_losses = tf_keras.losses.sparse_categorical_crossentropy(
labels, predictions, from_logits=from_logits)
if weights is None:
diff --git a/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy_test.py b/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy_test.py
index 3acab53394b..4214b6b9c28 100644
--- a/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy_test.py
+++ b/official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,16 +15,14 @@
"""Tests for masked LM loss."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import layers
from official.nlp.modeling import networks
from official.nlp.modeling.losses import weighted_sparse_categorical_crossentropy
-@keras_parameterized.run_all_keras_modes
-class ClassificationLossTest(keras_parameterized.TestCase):
+class ClassificationLossTest(tf.test.TestCase):
def create_lm_model(self,
vocab_size,
@@ -41,9 +39,9 @@ def create_lm_model(self,
hidden_size=hidden_size,
num_attention_heads=4,
)
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
_ = xformer_stack([word_ids, mask, type_ids])
# Create a maskedLM from the transformer stack.
@@ -51,11 +49,11 @@ def create_lm_model(self,
embedding_table=xformer_stack.get_embedding_table(), output=output)
# Create a model from the masked LM layer.
- lm_input_tensor = tf.keras.Input(shape=(sequence_length, hidden_size))
- masked_lm_positions = tf.keras.Input(
+ lm_input_tensor = tf_keras.Input(shape=(sequence_length, hidden_size))
+ masked_lm_positions = tf_keras.Input(
shape=(num_predictions,), dtype=tf.int32)
output = test_layer(lm_input_tensor, masked_positions=masked_lm_positions)
- return tf.keras.Model([lm_input_tensor, masked_lm_positions], output)
+ return tf_keras.Model([lm_input_tensor, masked_lm_positions], output)
def test_loss_3d_input(self):
"""Test overall loss with a 3-dimensional input, from a masked LM."""
diff --git a/official/nlp/modeling/models/__init__.py b/official/nlp/modeling/models/__init__.py
index afe28858b0b..7ade8e4a8dc 100644
--- a/official/nlp/modeling/models/__init__.py
+++ b/official/nlp/modeling/models/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/modeling/models/bert_classifier.py b/official/nlp/modeling/models/bert_classifier.py
index 105f00497a0..6b80cbdcf1c 100644
--- a/official/nlp/modeling/models/bert_classifier.py
+++ b/official/nlp/modeling/models/bert_classifier.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,13 @@
"""BERT cls-token classifier."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
-@tf.keras.utils.register_keras_serializable(package='Text')
-class BertClassifier(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BertClassifier(tf_keras.Model):
"""Classifier model based on a BERT-style transformer-based encoder.
This is an implementation of the network structure surrounding a transformer
@@ -79,7 +79,7 @@ def __init__(self,
cls_inputs = outputs[1]
else:
cls_inputs = outputs['pooled_output']
- cls_inputs = tf.keras.layers.Dropout(rate=dropout_rate)(cls_inputs)
+ cls_inputs = tf_keras.layers.Dropout(rate=dropout_rate)(cls_inputs)
else:
outputs = network(inputs)
if isinstance(outputs, list):
@@ -117,7 +117,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self.classifier = classifier
diff --git a/official/nlp/modeling/models/bert_classifier_test.py b/official/nlp/modeling/models/bert_classifier_test.py
index 98c4d8287f2..ca7b23e88b2 100644
--- a/official/nlp/modeling/models/bert_classifier_test.py
+++ b/official/nlp/modeling/models/bert_classifier_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,18 +15,14 @@
"""Tests for BERT trainer network."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import layers
from official.nlp.modeling import networks
from official.nlp.modeling.models import bert_classifier
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class BertClassifierTest(keras_parameterized.TestCase):
+class BertClassifierTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.named_parameters(('single_cls', 1, False), ('3_cls', 3, False),
('3_cls_dictoutputs', 3, True))
@@ -43,9 +39,9 @@ def test_bert_trainer(self, num_classes, dict_outputs):
test_network, num_classes=num_classes)
# Create a set of 2-dimensional inputs (the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
# Invoke the trainer model on the inputs. This causes the layer to be built.
cls_outs = bert_trainer_model([word_ids, mask, type_ids])
diff --git a/official/nlp/modeling/models/bert_pretrainer.py b/official/nlp/modeling/models/bert_pretrainer.py
index 4c8f98a94d8..6ff2e71a016 100644
--- a/official/nlp/modeling/models/bert_pretrainer.py
+++ b/official/nlp/modeling/models/bert_pretrainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,14 +20,15 @@
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
from official.nlp.modeling import networks
-@tf.keras.utils.register_keras_serializable(package='Text')
-class BertPretrainer(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BertPretrainer(tf_keras.Model):
"""BERT pretraining model.
[Note] Please use the new `BertPretrainerV2` for your projects.
@@ -91,7 +92,7 @@ def __init__(self,
'requested num_token_predictions %s.' %
(sequence_output_length, num_token_predictions))
- masked_lm_positions = tf.keras.layers.Input(
+ masked_lm_positions = tf_keras.layers.Input(
shape=(num_token_predictions,),
name='masked_lm_positions',
dtype=tf.int32)
@@ -102,7 +103,7 @@ def __init__(self,
masked_lm = layers.MaskedLM(
embedding_table=embedding_table,
activation=activation,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
output=output,
name='cls/predictions')
lm_outputs = masked_lm(
@@ -111,7 +112,7 @@ def __init__(self,
classification = networks.Classification(
input_width=cls_output.shape[-1],
num_classes=num_classes,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
output=output,
name='classification')
sentence_outputs = classification(cls_output)
@@ -142,7 +143,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self.encoder = network
@@ -157,9 +158,9 @@ def from_config(cls, config, custom_objects=None):
return cls(**config)
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
@gin.configurable
-class BertPretrainerV2(tf.keras.Model):
+class BertPretrainerV2(tf_keras.Model):
"""BERT pretraining model V2.
Adds the masked language model head and optional classification heads upon the
@@ -188,11 +189,11 @@ class BertPretrainerV2(tf.keras.Model):
def __init__(
self,
- encoder_network: tf.keras.Model,
+ encoder_network: tf_keras.Model,
mlm_activation=None,
mlm_initializer='glorot_uniform',
- classification_heads: Optional[List[tf.keras.layers.Layer]] = None,
- customized_masked_lm: Optional[tf.keras.layers.Layer] = None,
+ classification_heads: Optional[List[tf_keras.layers.Layer]] = None,
+ customized_masked_lm: Optional[tf_keras.layers.Layer] = None,
name: str = 'bert',
**kwargs):
super().__init__(self, name=name, **kwargs)
@@ -217,7 +218,7 @@ def __init__(
activation=mlm_activation,
initializer=mlm_initializer,
name='cls/predictions')
- masked_lm_positions = tf.keras.layers.Input(
+ masked_lm_positions = tf_keras.layers.Input(
shape=(None,), name='masked_lm_positions', dtype=tf.int32)
if isinstance(inputs, dict):
inputs['masked_lm_positions'] = masked_lm_positions
@@ -225,7 +226,7 @@ def __init__(
inputs.append(masked_lm_positions)
self.inputs = inputs
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
if isinstance(inputs, list):
logging.warning('List inputs to BertPretrainer are discouraged.')
inputs = dict([
diff --git a/official/nlp/modeling/models/bert_pretrainer_test.py b/official/nlp/modeling/models/bert_pretrainer_test.py
index 86977737221..4fc90ccfbd5 100644
--- a/official/nlp/modeling/models/bert_pretrainer_test.py
+++ b/official/nlp/modeling/models/bert_pretrainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,18 +16,14 @@
import itertools
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import layers
from official.nlp.modeling import networks
from official.nlp.modeling.models import bert_pretrainer
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class BertPretrainerTest(keras_parameterized.TestCase):
+class BertPretrainerTest(tf.test.TestCase, parameterized.TestCase):
def test_bert_pretrainer(self):
"""Validate that the Keras object can be created."""
@@ -48,10 +44,10 @@ def test_bert_pretrainer(self):
num_token_predictions=num_token_predictions)
# Create a set of 2-dimensional inputs (the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- masked_lm_positions = tf.keras.Input(
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ masked_lm_positions = tf_keras.Input(
shape=(num_token_predictions,), dtype=tf.int32)
# Invoke the trainer model on the inputs. This causes the layer to be built.
@@ -109,7 +105,7 @@ def test_serialize_deserialize(self):
new_bert_trainer_model.get_config())
-class BertPretrainerV2Test(keras_parameterized.TestCase):
+class BertPretrainerV2Test(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters(itertools.product(
(False, True),
@@ -121,6 +117,7 @@ def test_bert_pretrainerv2(self, dict_outputs, return_all_encoder_outputs,
use_customized_masked_lm, has_masked_lm_positions):
"""Validate that the Keras object can be created."""
# Build a transformer network to use within the BERT trainer.
+ del dict_outputs, return_all_encoder_outputs
vocab_size = 100
sequence_length = 512
hidden_size = 48
@@ -144,11 +141,11 @@ def test_bert_pretrainerv2(self, dict_outputs, return_all_encoder_outputs,
num_token_predictions = 20
# Create a set of 2-dimensional inputs (the first dimension is implicit).
inputs = dict(
- input_word_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32))
+ input_word_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32))
if has_masked_lm_positions:
- inputs['masked_lm_positions'] = tf.keras.Input(
+ inputs['masked_lm_positions'] = tf_keras.Input(
shape=(num_token_predictions,), dtype=tf.int32)
# Invoke the trainer model on the inputs. This causes the layer to be built.
@@ -196,10 +193,10 @@ def test_multiple_cls_outputs(self):
num_token_predictions = 20
# Create a set of 2-dimensional inputs (the first dimension is implicit).
inputs = dict(
- input_word_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- masked_lm_positions=tf.keras.Input(
+ input_word_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ masked_lm_positions=tf_keras.Input(
shape=(num_token_predictions,), dtype=tf.int32))
# Invoke the trainer model on the inputs. This causes the layer to be built.
diff --git a/official/nlp/modeling/models/bert_span_labeler.py b/official/nlp/modeling/models/bert_span_labeler.py
index 5edc62967b8..f105b5548b9 100644
--- a/official/nlp/modeling/models/bert_span_labeler.py
+++ b/official/nlp/modeling/models/bert_span_labeler.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,13 @@
"""BERT Question Answering model."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import networks
-@tf.keras.utils.register_keras_serializable(package='Text')
-class BertSpanLabeler(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BertSpanLabeler(tf_keras.Model):
"""Span labeler model based on a BERT-style transformer-based encoder.
This is an implementation of the network structure surrounding a transformer
@@ -80,10 +80,10 @@ def __init__(self,
# Use identity layers wrapped in lambdas to explicitly name the output
# tensors. This allows us to use string-keyed dicts in Keras fit/predict/
# evaluate calls.
- start_logits = tf.keras.layers.Lambda(
+ start_logits = tf_keras.layers.Lambda(
tf.identity, name='start_positions')(
start_logits)
- end_logits = tf.keras.layers.Lambda(
+ end_logits = tf_keras.layers.Lambda(
tf.identity, name='end_positions')(
end_logits)
@@ -109,7 +109,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self.span_labeling = span_labeling
diff --git a/official/nlp/modeling/models/bert_span_labeler_test.py b/official/nlp/modeling/models/bert_span_labeler_test.py
index 9f9da14c30f..23f991a0437 100644
--- a/official/nlp/modeling/models/bert_span_labeler_test.py
+++ b/official/nlp/modeling/models/bert_span_labeler_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,17 +15,13 @@
"""Tests for BERT trainer network."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import networks
from official.nlp.modeling.models import bert_span_labeler
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class BertSpanLabelerTest(keras_parameterized.TestCase):
+class BertSpanLabelerTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters(True, False)
def test_bert_trainer(self, dict_outputs):
@@ -40,15 +36,15 @@ def test_bert_trainer(self, dict_outputs):
bert_trainer_model = bert_span_labeler.BertSpanLabeler(test_network)
# Create a set of 2-dimensional inputs (the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
# Invoke the trainer model on the inputs. This causes the layer to be built.
cls_outs = bert_trainer_model([word_ids, mask, type_ids])
# Validate that there are 2 outputs are of the expected shape.
- self.assertEqual(2, len(cls_outs))
+ self.assertLen(cls_outs, 2)
expected_shape = [None, sequence_length]
for out in cls_outs:
self.assertAllEqual(expected_shape, out.shape.as_list())
diff --git a/official/nlp/modeling/models/bert_token_classifier.py b/official/nlp/modeling/models/bert_token_classifier.py
index 6375aa4b61c..c98b17d91cb 100644
--- a/official/nlp/modeling/models/bert_token_classifier.py
+++ b/official/nlp/modeling/models/bert_token_classifier.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,11 +15,11 @@
"""BERT token classifier."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-@tf.keras.utils.register_keras_serializable(package='Text')
-class BertTokenClassifier(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BertTokenClassifier(tf_keras.Model):
"""Token classifier model based on a BERT-style transformer-based encoder.
This is an implementation of the network structure surrounding a transformer
@@ -68,10 +68,10 @@ def __init__(self,
sequence_output = outputs[0]
else:
sequence_output = outputs['sequence_output']
- sequence_output = tf.keras.layers.Dropout(rate=dropout_rate)(
+ sequence_output = tf_keras.layers.Dropout(rate=dropout_rate)(
sequence_output)
- classifier = tf.keras.layers.Dense(
+ classifier = tf_keras.layers.Dense(
num_classes,
activation=None,
kernel_initializer=initializer,
@@ -81,7 +81,7 @@ def __init__(self,
output_tensors = {'logits': logits}
elif output == 'predictions':
output_tensors = {
- 'predictions': tf.keras.layers.Activation(tf.nn.log_softmax)(logits)
+ 'predictions': tf_keras.layers.Activation(tf.nn.log_softmax)(logits)
}
else:
raise ValueError(
@@ -115,7 +115,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self.classifier = classifier
diff --git a/official/nlp/modeling/models/bert_token_classifier_test.py b/official/nlp/modeling/models/bert_token_classifier_test.py
index 83765f5fed5..ffa193d3d0b 100644
--- a/official/nlp/modeling/models/bert_token_classifier_test.py
+++ b/official/nlp/modeling/models/bert_token_classifier_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,17 +15,13 @@
"""Tests for BERT token classifier."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import networks
from official.nlp.modeling.models import bert_token_classifier
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class BertTokenClassifierTest(keras_parameterized.TestCase):
+class BertTokenClassifierTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters((True, True), (False, False))
def test_bert_trainer(self, dict_outputs, output_encoder_outputs):
@@ -49,9 +45,9 @@ def test_bert_trainer(self, dict_outputs, output_encoder_outputs):
output_encoder_outputs=output_encoder_outputs)
# Create a set of 2-dimensional inputs (the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
# Invoke the trainer model on the inputs. This causes the layer to be built.
outputs = bert_trainer_model([word_ids, mask, type_ids])
diff --git a/official/nlp/modeling/models/dual_encoder.py b/official/nlp/modeling/models/dual_encoder.py
index b5b948c11a1..43d19c89541 100644
--- a/official/nlp/modeling/models/dual_encoder.py
+++ b/official/nlp/modeling/models/dual_encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,13 @@
"""Trainer network for dual encoder style models."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
-@tf.keras.utils.register_keras_serializable(package='Text')
-class DualEncoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class DualEncoder(tf_keras.Model):
"""A dual encoder model based on a transformer-based encoder.
This is an implementation of the dual encoder network structure based on the
@@ -43,7 +43,7 @@ class DualEncoder(tf.keras.Model):
"""
def __init__(self,
- network: tf.keras.Model,
+ network: tf_keras.Model,
max_seq_length: int = 32,
normalize: bool = True,
logit_scale: float = 1.0,
@@ -52,19 +52,19 @@ def __init__(self,
**kwargs) -> None:
if output == 'logits':
- left_word_ids = tf.keras.layers.Input(
+ left_word_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='left_word_ids')
- left_mask = tf.keras.layers.Input(
+ left_mask = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='left_mask')
- left_type_ids = tf.keras.layers.Input(
+ left_type_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='left_type_ids')
else:
# Keep the consistant with legacy BERT hub module input names.
- left_word_ids = tf.keras.layers.Input(
+ left_word_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_word_ids')
- left_mask = tf.keras.layers.Input(
+ left_mask = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_mask')
- left_type_ids = tf.keras.layers.Input(
+ left_type_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='input_type_ids')
left_inputs = [left_word_ids, left_mask, left_type_ids]
@@ -75,16 +75,16 @@ def __init__(self,
left_sequence_output = left_outputs['sequence_output']
left_encoded = left_outputs['pooled_output']
if normalize:
- left_encoded = tf.keras.layers.Lambda(
+ left_encoded = tf_keras.layers.Lambda(
lambda x: tf.nn.l2_normalize(x, axis=1))(
left_encoded)
if output == 'logits':
- right_word_ids = tf.keras.layers.Input(
+ right_word_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='right_word_ids')
- right_mask = tf.keras.layers.Input(
+ right_mask = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='right_mask')
- right_type_ids = tf.keras.layers.Input(
+ right_type_ids = tf_keras.layers.Input(
shape=(max_seq_length,), dtype=tf.int32, name='right_type_ids')
right_inputs = [right_word_ids, right_mask, right_type_ids]
@@ -94,7 +94,7 @@ def __init__(self,
else:
right_encoded = right_outputs['pooled_output']
if normalize:
- right_encoded = tf.keras.layers.Lambda(
+ right_encoded = tf_keras.layers.Lambda(
lambda x: tf.nn.l2_normalize(x, axis=1))(
right_encoded)
@@ -143,7 +143,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self.network = network
diff --git a/official/nlp/modeling/models/dual_encoder_test.py b/official/nlp/modeling/models/dual_encoder_test.py
index 699277966d2..6ba72e08119 100644
--- a/official/nlp/modeling/models/dual_encoder_test.py
+++ b/official/nlp/modeling/models/dual_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,17 +15,13 @@
"""Tests for dual encoder network."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import networks
from official.nlp.modeling.models import dual_encoder
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class DualEncoderTest(keras_parameterized.TestCase):
+class DualEncoderTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters((192, 'logits'), (768, 'predictions'))
def test_dual_encoder(self, hidden_size, output):
@@ -44,13 +40,13 @@ def test_dual_encoder(self, hidden_size, output):
test_network, max_seq_length=sequence_length, output=output)
# Create a set of 2-dimensional inputs (the first dimension is implicit).
- left_word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- left_mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- left_type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ left_word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ left_mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ left_type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
- right_word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- right_mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- right_type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ right_word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ right_mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ right_type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
if output == 'logits':
outputs = dual_encoder_model([
@@ -72,6 +68,7 @@ def test_dual_encoder(self, hidden_size, output):
def test_dual_encoder_tensor_call(self, hidden_size, output):
"""Validate that the Keras object can be invoked."""
# Build a transformer network to use within the dual encoder model.
+ del hidden_size
sequence_length = 2
test_network = networks.BertEncoder(vocab_size=100, num_layers=2)
diff --git a/official/nlp/modeling/models/electra_pretrainer.py b/official/nlp/modeling/models/electra_pretrainer.py
index 96ab689e502..ed00f683195 100644
--- a/official/nlp/modeling/models/electra_pretrainer.py
+++ b/official/nlp/modeling/models/electra_pretrainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,14 +17,14 @@
import copy
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.nlp.modeling import layers
-@tf.keras.utils.register_keras_serializable(package='Text')
-class ElectraPretrainer(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class ElectraPretrainer(tf_keras.Model):
"""ELECTRA network training model.
This is an implementation of the network structure described in "ELECTRA:
@@ -96,23 +96,24 @@ def __init__(self,
self.masked_lm = layers.MaskedLM(
embedding_table=generator_network.get_embedding_table(),
activation=mlm_activation,
- initializer=mlm_initializer,
+ initializer=tf_utils.clone_initializer(mlm_initializer),
output=output_type,
name='generator_masked_lm')
self.classification = layers.ClassificationHead(
inner_dim=generator_network.get_config()['hidden_size'],
num_classes=num_classes,
- initializer=mlm_initializer,
+ initializer=tf_utils.clone_initializer(mlm_initializer),
name='generator_classification_head')
- self.discriminator_projection = tf.keras.layers.Dense(
+ self.discriminator_projection = tf_keras.layers.Dense(
units=discriminator_network.get_config()['hidden_size'],
activation=mlm_activation,
- kernel_initializer=mlm_initializer,
+ kernel_initializer=tf_utils.clone_initializer(mlm_initializer),
name='discriminator_projection_head')
- self.discriminator_head = tf.keras.layers.Dense(
- units=1, kernel_initializer=mlm_initializer)
+ self.discriminator_head = tf_keras.layers.Dense(
+ units=1,
+ kernel_initializer=tf_utils.clone_initializer(mlm_initializer))
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
"""ELECTRA forward pass.
Args:
diff --git a/official/nlp/modeling/models/electra_pretrainer_test.py b/official/nlp/modeling/models/electra_pretrainer_test.py
index 23864934993..f138e4ce525 100644
--- a/official/nlp/modeling/models/electra_pretrainer_test.py
+++ b/official/nlp/modeling/models/electra_pretrainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,17 +14,13 @@
"""Tests for ELECTRA pre trainer network."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import networks
from official.nlp.modeling.models import electra_pretrainer
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class ElectraPretrainerTest(keras_parameterized.TestCase):
+class ElectraPretrainerTest(tf.test.TestCase):
def test_electra_pretrainer(self):
"""Validate that the Keras object can be created."""
@@ -54,12 +50,12 @@ def test_electra_pretrainer(self):
disallow_correct=True)
# Create a set of 2-dimensional inputs (the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- lm_positions = tf.keras.Input(
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ lm_positions = tf_keras.Input(
shape=(num_token_predictions,), dtype=tf.int32)
- lm_ids = tf.keras.Input(shape=(num_token_predictions,), dtype=tf.int32)
+ lm_ids = tf_keras.Input(shape=(num_token_predictions,), dtype=tf.int32)
inputs = {
'input_word_ids': word_ids,
'input_mask': mask,
diff --git a/official/nlp/modeling/models/seq2seq_transformer.py b/official/nlp/modeling/models/seq2seq_transformer.py
index d33e690250e..82c354a267b 100644
--- a/official/nlp/modeling/models/seq2seq_transformer.py
+++ b/official/nlp/modeling/models/seq2seq_transformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,11 @@
Model paper: https://arxiv.org/pdf/1706.03762.pdf
"""
+import inspect
import math
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
from official.modeling import tf_utils
from official.nlp.modeling import layers
from official.nlp.modeling.ops import beam_search
@@ -26,7 +28,7 @@
EOS_ID = 1
-class Seq2SeqTransformer(tf.keras.Model):
+class Seq2SeqTransformer(tf_keras.Model):
"""Transformer model with Keras.
Implemented as described in: https://arxiv.org/pdf/1706.03762.pdf
@@ -87,8 +89,8 @@ def __init__(self,
self.decoder_layer = decoder_layer
self.position_embedding = layers.RelativePositionEmbedding(
hidden_size=self._embedding_width)
- self.encoder_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
- self.decoder_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self.encoder_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
+ self.decoder_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
def get_config(self):
config = {
@@ -144,7 +146,7 @@ def _parse_inputs(self, inputs):
return embedded_inputs, boolean_mask, input_shape, source_dtype
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
"""Calculate target logits or inferred target sequences.
Args:
@@ -197,7 +199,7 @@ def call(self, inputs):
encoder_inputs = self.encoder_dropout(encoder_inputs)
- encoder_outputs = self.encoder_layer(
+ encoder_outputs = self.encoder_layer( # pyrefly: ignore[not-callable]
encoder_inputs, attention_mask=attention_mask)
if targets is None:
@@ -214,7 +216,7 @@ def call(self, inputs):
# Create cache storing decoder attention values for each layer.
init_decode_length = (max_decode_length if self._padded_decode else 0)
- num_heads = self.decoder_layer.num_attention_heads
+ num_heads = self.decoder_layer.num_attention_heads # pyrefly: ignore[missing-attribute]
dim_per_head = self._embedding_width // num_heads
# Cache dtype needs to match beam_search dtype.
@@ -229,7 +231,7 @@ def call(self, inputs):
tf.zeros(
[batch_size, init_decode_length, num_heads, dim_per_head],
dtype=self.compute_dtype)
- } for layer in range(self.decoder_layer.num_layers)
+ } for layer in range(self.decoder_layer.num_layers) # pyrefly: ignore[missing-attribute]
}
# pylint: enable=g-complex-comprehension
@@ -282,7 +284,7 @@ def call(self, inputs):
tf.expand_dims(boolean_mask, axis=1), dtype=source_dtype)
attention_mask = tf.tile(attention_mask, [1, decoder_length, 1])
- outputs = self.decoder_layer(
+ outputs = self.decoder_layer( # pyrefly: ignore[not-callable]
decoder_inputs,
encoder_outputs,
self_attention_mask=self_attention_mask,
@@ -340,7 +342,7 @@ def symbols_to_logits_fn(ids, i, cache):
attention_mask = cache.get("encoder_decoder_attention_mask")
attention_mask = tf.tile(attention_mask, [1, decoder_length, 1])
- decoder_outputs = self.decoder_layer(
+ decoder_outputs = self.decoder_layer( # pyrefly: ignore[not-callable]
decoder_input,
cache.get("encoder_outputs"),
self_attention_mask=self_attention_mask,
@@ -357,7 +359,7 @@ def symbols_to_logits_fn(ids, i, cache):
return symbols_to_logits_fn
-class TransformerEncoder(tf.keras.layers.Layer):
+class TransformerEncoder(tf_keras.layers.Layer):
"""Transformer encoder.
Transformer encoder is made up of N identical layers. Each layer is composed
@@ -394,7 +396,7 @@ def __init__(self,
layers is normalized.
norm_epsilon: Epsilon value to initialize normalization layers.
intermediate_dropout: Dropout probability for intermediate_dropout_layer.
- **kwargs: key word arguemnts passed to tf.keras.layers.Layer.
+ **kwargs: key word arguemnts passed to tf_keras.layers.Layer.
"""
super(TransformerEncoder, self).__init__(**kwargs)
@@ -426,7 +428,7 @@ def build(self, input_shape):
inner_dropout=self._intermediate_dropout,
attention_initializer=attention_initializer(input_shape[2]),
name=("layer_%d" % i)))
- self.output_normalization = tf.keras.layers.LayerNormalization(
+ self.output_normalization = tf_keras.layers.LayerNormalization(
epsilon=self._norm_epsilon, dtype="float32")
super(TransformerEncoder, self).build(input_shape)
@@ -469,7 +471,7 @@ def call(self, encoder_inputs, attention_mask=None):
return output_tensor
-class TransformerDecoder(tf.keras.layers.Layer):
+class TransformerDecoder(tf_keras.layers.Layer):
"""Transformer decoder.
Like the encoder, the decoder is made up of N identical layers.
@@ -491,6 +493,8 @@ def __init__(self,
norm_first=True,
norm_epsilon=1e-6,
intermediate_dropout=0.0,
+ self_attention_cls=None,
+ cross_attention_cls=None,
**kwargs):
"""Initialize a Transformer decoder.
@@ -508,7 +512,11 @@ def __init__(self,
layers is normalized.
norm_epsilon: Epsilon value to initialize normalization layers.
intermediate_dropout: Dropout probability for intermediate_dropout_layer.
- **kwargs: key word arguemnts passed to tf.keras.layers.Layer.
+ self_attention_cls: An optional class to use for self attention
+ or a function that provides the class per layer.
+ cross_attention_cls: An optional class to use for cross attention
+ or a function that provides the class per layer.
+ **kwargs: key word arguemnts passed to tf_keras.layers.Layer.
"""
super(TransformerDecoder, self).__init__(**kwargs)
self.num_layers = num_layers
@@ -521,11 +529,26 @@ def __init__(self,
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
self._intermediate_dropout = intermediate_dropout
+ self._self_attention_cls = self_attention_cls
+ self._cross_attention_cls = cross_attention_cls
def build(self, input_shape):
"""Implements build() for the layer."""
+
+ def _select_attention_cls(attention_cls, index):
+ cls = None
+ if attention_cls is not None:
+ cls = (
+ attention_cls(index)
+ if inspect.isfunction(attention_cls)
+ else attention_cls
+ )
+ return cls
+
self.decoder_layers = []
for i in range(self.num_layers):
+ self_attention_cls = _select_attention_cls(self._self_attention_cls, i)
+ cross_attention_cls = _select_attention_cls(self._cross_attention_cls, i)
self.decoder_layers.append(
layers.TransformerDecoderBlock(
num_attention_heads=self.num_attention_heads,
@@ -538,8 +561,10 @@ def build(self, input_shape):
norm_epsilon=self._norm_epsilon,
intermediate_dropout=self._intermediate_dropout,
attention_initializer=attention_initializer(input_shape[2]),
- name=("layer_%d" % i)))
- self.output_normalization = tf.keras.layers.LayerNormalization(
+ name=("layer_%d" % i),
+ self_attention_cls=self_attention_cls,
+ cross_attention_cls=cross_attention_cls))
+ self.output_normalization = tf_keras.layers.LayerNormalization(
epsilon=1e-6, dtype="float32")
super(TransformerDecoder, self).build(input_shape)
@@ -554,7 +579,9 @@ def get_config(self):
"use_bias": self._use_bias,
"norm_first": self._norm_first,
"norm_epsilon": self._norm_epsilon,
- "intermediate_dropout": self._intermediate_dropout
+ "intermediate_dropout": self._intermediate_dropout,
+ "self_attention_cls": self._self_attention_cls,
+ "cross_attention_cls": self._cross_attention_cls,
}
base_config = super(TransformerDecoder, self).get_config()
return dict(list(base_config.items()) + list(config.items()))
@@ -620,4 +647,4 @@ def attention_initializer(hidden_size):
"""Initializer for attention layers in Seq2SeqTransformer."""
hidden_size = int(hidden_size)
limit = math.sqrt(6.0 / (hidden_size + hidden_size))
- return tf.keras.initializers.RandomUniform(minval=-limit, maxval=limit)
+ return tf_keras.initializers.RandomUniform(minval=-limit, maxval=limit)
diff --git a/official/nlp/modeling/models/seq2seq_transformer_test.py b/official/nlp/modeling/models/seq2seq_transformer_test.py
index f45f5f3cef6..a7c6a9eb373 100644
--- a/official/nlp/modeling/models/seq2seq_transformer_test.py
+++ b/official/nlp/modeling/models/seq2seq_transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,16 +17,24 @@
from absl import logging
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
+from official.nlp.modeling.layers import attention
from official.nlp.modeling.models import seq2seq_transformer
class Seq2SeqTransformerTest(tf.test.TestCase, parameterized.TestCase):
- def _build_model(self, padded_decode, decode_max_length, embedding_width):
+ def _build_model(
+ self,
+ padded_decode,
+ decode_max_length,
+ embedding_width,
+ self_attention_cls=None,
+ cross_attention_cls=None,
+ ):
num_layers = 1
num_attention_heads = 2
intermediate_size = 32
@@ -43,7 +51,11 @@ def _build_model(self, padded_decode, decode_max_length, embedding_width):
norm_epsilon=1e-6,
intermediate_dropout=0.01)
encoder_layer = seq2seq_transformer.TransformerEncoder(**encdec_kwargs)
- decoder_layer = seq2seq_transformer.TransformerDecoder(**encdec_kwargs)
+ decoder_layer = seq2seq_transformer.TransformerDecoder(
+ **encdec_kwargs,
+ self_attention_cls=self_attention_cls,
+ cross_attention_cls=cross_attention_cls
+ )
return seq2seq_transformer.Seq2SeqTransformer(
vocab_size=vocab_size,
@@ -64,8 +76,41 @@ def _build_model(self, padded_decode, decode_max_length, embedding_width):
],
embed=[True, False],
is_training=[True, False],
+ custom_self_attention=[False, True],
+ custom_cross_attention=[False, True],
mode="eager"))
- def test_create_model_with_ds(self, distribution, embed, is_training):
+ def test_create_model_with_ds(
+ self,
+ distribution,
+ embed,
+ is_training,
+ custom_self_attention,
+ custom_cross_attention,
+ ):
+ self_attention_called = False
+ cross_attention_called = False
+
+ class SelfAttention(attention.CachedAttention):
+ """Dummy implementation of custom attention."""
+
+ def __call__(
+ self, *args, **kwargs
+ ):
+ nonlocal self_attention_called
+ self_attention_called = True
+ return super().__call__(*args, **kwargs)
+
+ class CrossAttention:
+ """Dummy implementation of custom attention."""
+
+ def __init__(self, *args, **kwargs):
+ pass
+
+ def __call__(self, query, value, attention_mask, **kwargs):
+ nonlocal cross_attention_called
+ cross_attention_called = True
+ return query
+
with distribution.scope():
padded_decode = isinstance(
distribution,
@@ -74,7 +119,14 @@ def test_create_model_with_ds(self, distribution, embed, is_training):
batch_size = 4
embedding_width = 16
model = self._build_model(
- padded_decode, decode_max_length, embedding_width)
+ padded_decode,
+ decode_max_length,
+ embedding_width,
+ self_attention_cls=SelfAttention if custom_self_attention else None,
+ cross_attention_cls=CrossAttention
+ if custom_cross_attention
+ else None,
+ )
@tf.function
def step(inputs):
@@ -91,7 +143,7 @@ def _step_fn(inputs):
embedded_inputs=np.zeros(
(batch_size, decode_max_length, embedding_width),
dtype=np.float32),
- input_masks=np.ones((batch_size, decode_max_length), dtype=np.bool))
+ input_masks=np.ones((batch_size, decode_max_length), dtype=bool))
else:
fake_inputs = dict(
inputs=np.zeros((batch_size, decode_max_length), dtype=np.int32))
@@ -105,6 +157,8 @@ def _step_fn(inputs):
local_outputs = step(fake_inputs)
logging.info("local_outputs=%s", local_outputs)
self.assertEqual(local_outputs["outputs"][0].shape, (4, 10))
+ self.assertEqual(self_attention_called, custom_self_attention)
+ self.assertEqual(cross_attention_called, custom_cross_attention)
@parameterized.parameters(True, False)
def test_create_savedmodel(self, padded_decode):
diff --git a/official/nlp/modeling/models/t5.py b/official/nlp/modeling/models/t5.py
index c7bc2e1c806..a87d2132164 100644
--- a/official/nlp/modeling/models/t5.py
+++ b/official/nlp/modeling/models/t5.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -28,7 +28,7 @@
from typing import Callable, Dict, Optional, Sequence, Text, Union
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
@@ -55,6 +55,7 @@ def create_variable(self,
initializer: Initializer,
dtype: tf.DType = tf.float32,
**kwargs):
+ initializer = tf_utils.clone_initializer(initializer)
return tf.Variable(initializer(shape, dtype=dtype, **kwargs), name=name)
def read_variable(self,
@@ -156,7 +157,7 @@ def __init__(self,
if embeddings_initializer:
self.embed_init = embeddings_initializer
else:
- self.embed_init = tf.keras.initializers.TruncatedNormal(stddev=1.0)
+ self.embed_init = tf_keras.initializers.TruncatedNormal(stddev=1.0)
with self.name_scope:
self.embeddings = self.create_variable(
"embedding", [self.vocab_size, self.features],
@@ -223,7 +224,7 @@ def __init__(self, hidden_size: int, epsilon: float = 1e-6, **kwargs):
self.weight = self.create_variable(
"scale", [hidden_size],
dtype=self.dtype,
- initializer=tf.keras.initializers.Ones())
+ initializer=tf_keras.initializers.Ones())
@tf.Module.with_name_scope
def __call__(self, x):
@@ -253,14 +254,14 @@ def __init__(self,
self.use_bias = use_bias
self.w_init = w_init
if self.use_bias:
- self.b_init = b_init if b_init else tf.keras.initializers.Zeros()
+ self.b_init = b_init if b_init else tf_keras.initializers.Zeros()
elif b_init is not None:
raise ValueError("When not using a bias the b_init must be None.")
with self.name_scope:
if self.w_init is None:
stddev = 1 / math.sqrt(self.in_features)
- self.w_init = tf.keras.initializers.HeNormal()
+ self.w_init = tf_keras.initializers.HeNormal()
self.w = self.create_variable(
"kernel", [self.in_features, self.out_features],
@@ -321,13 +322,13 @@ def __init__(self,
self.bias_shape = (self.out_features,)
bias_rank = 1
if self.use_bias:
- self.b_init = b_init or tf.keras.initializers.Zeros()
+ self.b_init = b_init or tf_keras.initializers.Zeros()
elif b_init is not None:
raise ValueError("When not using a bias the b_init must be None.")
with self.name_scope:
if self.w_init is None:
- self.w_init = tf.keras.initializers.HeNormal()
+ self.w_init = tf_keras.initializers.HeNormal()
self.w = self.create_variable(
"kernel",
@@ -444,7 +445,7 @@ def __init__(self,
b_init=bias_initializer,
dtype=self.dtype,
name="wo")
- self.dropout = Dropout(rate=dropout_rate)
+ self.dropout = Dropout(rate=dropout_rate) # pyrefly: ignore[bad-argument-type]
@tf.Module.with_name_scope
def __call__(self,
@@ -588,7 +589,8 @@ def __init__(self,
init_std_rescaling = tf.math.sqrt(tf.cast(self.d_kv, dtype=self.dtype))
query_w_init = (
lambda *args, **kwargs: ( # pylint: disable=g-long-lambda
- weight_initializer(*args, **kwargs) / init_std_rescaling))
+ tf_utils.clone_initializer(weight_initializer)
+ (*args, **kwargs) / init_std_rescaling))
self.q = Linear3D(
self.d_model,
self.d_kv,
@@ -626,7 +628,7 @@ def __init__(self,
b_init=bias_initializer,
dtype=self.dtype,
name="o")
- self.dropout = Dropout(dropout_rate)
+ self.dropout = Dropout(dropout_rate) # pyrefly: ignore[bad-argument-type]
def _update_cache(self, key, value, cache, decode_position):
"""Updates cache states and gets full-length key/value tensors."""
@@ -705,6 +707,7 @@ def __call__(self,
if mask is not None:
scores += mask # (bs, n_heads, qlen, klen)
weights = tf.nn.softmax(tf.cast(scores, tf.float32), axis=-1)
+ output_scores = weights
# weights shape = (bs, n_heads, qlen, klen)
weights = tf.cast(weights, scores.dtype)
weight_shape = tf_utils.get_shape_list(weights)
@@ -718,6 +721,7 @@ def __call__(self,
c = self.o(c)
outputs = dict(context=c)
+ outputs["attention_scores"] = output_scores
if cache:
outputs["cache"] = cache
return outputs
@@ -750,10 +754,10 @@ def __init__(self,
name="attention")
self.layer_norm = RMSNorm(
hidden_size=d_model,
- epsilon=layer_norm_epsilon,
+ epsilon=layer_norm_epsilon, # pyrefly: ignore[bad-argument-type]
dtype=self.dtype,
name="layer_norm")
- self.dropout = Dropout(dropout_rate)
+ self.dropout = Dropout(dropout_rate) # pyrefly: ignore[bad-argument-type]
@tf.Module.with_name_scope
def __call__(self,
@@ -807,10 +811,10 @@ def __init__(self,
name="attention")
self.layer_norm = RMSNorm(
hidden_size=d_model,
- epsilon=layer_norm_epsilon,
+ epsilon=layer_norm_epsilon, # pyrefly: ignore[bad-argument-type]
dtype=self.dtype,
name="layer_norm")
- self.dropout = Dropout(dropout_rate)
+ self.dropout = Dropout(dropout_rate) # pyrefly: ignore[bad-argument-type]
@tf.Module.with_name_scope
def __call__(self,
@@ -851,6 +855,7 @@ def __init__(self,
rescale_query: bool = False,
weight_initializer: Optional[Initializer] = None,
bias_initializer: Optional[Initializer] = None,
+ return_attention_scores: bool = False,
**kwargs):
super().__init__(**kwargs)
with self.name_scope:
@@ -866,7 +871,7 @@ def __init__(self,
name="self_attention")
self.ffn_layer_norm = RMSNorm(
hidden_size=d_model,
- epsilon=layer_norm_epsilon,
+ epsilon=layer_norm_epsilon, # pyrefly: ignore[bad-argument-type]
dtype=self.dtype,
name="ffn_layer_norm")
self.ffn = FFN(
@@ -878,7 +883,8 @@ def __init__(self,
bias_initializer=bias_initializer,
dtype=self.dtype,
name="ffn")
- self.ffn_output_dropout = Dropout(dropout_rate)
+ self.ffn_output_dropout = Dropout(dropout_rate) # pyrefly: ignore[bad-argument-type]
+ self.return_attention_scores = return_attention_scores
@tf.Module.with_name_scope
def __call__(self,
@@ -900,7 +906,8 @@ def __call__(self,
ffn_output = self.ffn_output_dropout(
ffn_output, noise_shape=tensor_shape, training=training)
ffn_output = attn_output + ffn_output
-
+ if self.return_attention_scores:
+ return ffn_output, attention_outputs["attention_scores"]
return ffn_output
@@ -943,7 +950,7 @@ def __init__(self,
name="cross_attention")
self.ffn_layer_norm = RMSNorm(
hidden_size=d_model,
- epsilon=layer_norm_epsilon,
+ epsilon=layer_norm_epsilon, # pyrefly: ignore[bad-argument-type]
dtype=self.dtype,
name="ffn_layer_norm")
self.ffn = FFN(
@@ -955,7 +962,7 @@ def __init__(self,
bias_initializer=bias_initializer,
dtype=self.dtype,
name="ffn")
- self.ffn_output_dropout = Dropout(dropout_rate,)
+ self.ffn_output_dropout = Dropout(dropout_rate,) # pyrefly: ignore[bad-argument-type]
@tf.Module.with_name_scope
def __call__(self,
@@ -1012,7 +1019,7 @@ class T5TransformerParams:
relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128
relative_embeddings_initializer: Optional[Initializer] = None
- weight_initializer: Optional[Initializer] = (tf.keras.initializers.HeNormal())
+ weight_initializer: Optional[Initializer] = (tf_keras.initializers.HeNormal())
bias_initializer: Optional[Initializer] = None
rescale_query: bool = False
bidirectional: bool = True
@@ -1021,6 +1028,10 @@ class T5TransformerParams:
num_decoder_layers: Optional[int] = None
one_hot_embedding: bool = True
layer_sharing: bool = False
+ # If true, uses one relative embedding for all encoder layers and one for all
+ # decoder layers. Otherwise, have relative embedding for each layer.
+ use_shared_relative_position_bias: bool = True
+ return_attention_scores: bool = False
class Encoder(Module):
@@ -1049,17 +1060,34 @@ def __init__(self,
self.input_embed = shared_embedding
# Creates an alias to the input embed for encoder-only models.
self.word_embed = self.input_embed
- self.relative_embedding = RelativePositionEmbedding(
- num_heads=self.config.num_heads,
- relative_attention_num_buckets=self.config
- .relative_attention_num_buckets,
- relative_attention_max_distance=self.config
- .relative_attention_max_distance,
- bidirectional=self.config.bidirectional,
- embeddings_initializer=self.config.relative_embeddings_initializer,
- dtype=self.dtype,
- compute_dtype=self.compute_dtype,
- name="relative_posemb")
+ if config.use_shared_relative_position_bias:
+ self.relative_embedding = RelativePositionEmbedding(
+ num_heads=self.config.num_heads,
+ relative_attention_num_buckets=self.config
+ .relative_attention_num_buckets,
+ relative_attention_max_distance=self.config
+ .relative_attention_max_distance,
+ bidirectional=self.config.bidirectional,
+ embeddings_initializer=self.config.relative_embeddings_initializer,
+ dtype=self.dtype,
+ compute_dtype=self.compute_dtype,
+ name="relative_posemb")
+ else:
+ self.relative_embeddings = []
+ for layer_idx in range(self.config.num_layers):
+ relative_embedding = RelativePositionEmbedding(
+ num_heads=self.config.num_heads,
+ relative_attention_num_buckets=self.config
+ .relative_attention_num_buckets,
+ relative_attention_max_distance=self.config
+ .relative_attention_max_distance,
+ bidirectional=self.config.bidirectional,
+ embeddings_initializer=self.config
+ .relative_embeddings_initializer,
+ dtype=self.dtype,
+ compute_dtype=self.compute_dtype,
+ name=f"relative_posemb_{layer_idx}")
+ self.relative_embeddings.append(relative_embedding)
self.input_dropout = Dropout(self.config.dropout_rate,)
self.encoder_layers = []
for layer_idx in range(self.config.num_layers):
@@ -1077,6 +1105,7 @@ def __init__(self,
rescale_query=self.config.rescale_query,
weight_initializer=self.config.weight_initializer,
bias_initializer=self.config.bias_initializer,
+ return_attention_scores=self.config.return_attention_scores,
dtype=self.dtype,
name="encoder_block_%d" % layer_idx))
self.output_norm = RMSNorm(
@@ -1086,6 +1115,26 @@ def __init__(self,
name="final_layer_norm")
self.output_dropout = Dropout(self.config.dropout_rate,)
+ @tf.Module.with_name_scope
+ def get_relpos_bias(self,
+ input_length: int,
+ dense_inputs: tf.Tensor,
+ layer_idx: Optional[int] = None) -> tf.Tensor:
+ if self.config.use_shared_relative_position_bias:
+ position_bias = self.relative_embedding(input_length, input_length)
+ else:
+ position_bias = self.relative_embeddings[layer_idx](input_length, # pyrefly: ignore[bad-index]
+ input_length)
+ if dense_inputs is not None:
+ # Here we ignore relative position bias for dense embeddings.
+ # TODO(yejiayu): If we proceed to video use cases, rework this part.
+ dense_input_length = tf_utils.get_shape_list(dense_inputs)[1]
+ # Position bias shape: [batch, 1, len, len]
+ paddings = tf.constant([[0, 0], [0, 0], [0, dense_input_length],
+ [0, dense_input_length]])
+ position_bias = tf.pad(position_bias, paddings, "CONSTANT")
+ return position_bias
+
@tf.Module.with_name_scope
def __call__(self,
inputs=None,
@@ -1097,8 +1146,8 @@ def __call__(self,
Args:
inputs: input word ids. Optional if dense data are provided.
encoder_mask: the encoder self-attention mask.
- dense_inputs: dense input data. Concat after the embedding if word ids
- are provided.
+ dense_inputs: dense input data. Concat after the embedding if word ids are
+ provided.
training: whether it is training pass, affecting dropouts.
Returns:
@@ -1111,7 +1160,7 @@ def __call__(self,
inputs_array = []
if inputs is not None:
inputs_array.append(
- self.input_embed(inputs, one_hot=cfg.one_hot_embedding))
+ self.input_embed(inputs, one_hot=cfg.one_hot_embedding)) # pyrefly: ignore[not-callable]
if dense_inputs is not None:
inputs_array.append(dense_inputs)
if not inputs_array:
@@ -1125,26 +1174,25 @@ def __call__(self,
input_length = tf_utils.get_shape_list(inputs)[1]
else:
input_length = 0
- position_bias = self.relative_embedding(input_length, input_length)
- if dense_inputs is not None:
- # Here we ignore relative position bias for dense embeddings.
- # TODO(yejiayu): If we proceed to video use cases, rework this part.
- dense_input_length = tf_utils.get_shape_list(dense_inputs)[1]
- # Position bias shape: [batch, 1, len, len]
- paddings = tf.constant([[0, 0], [0, 0], [0, dense_input_length],
- [0, dense_input_length]])
- position_bias = tf.pad(position_bias, paddings, "CONSTANT")
+ attention_outputs = []
for i in range(cfg.num_layers):
+ position_bias = self.get_relpos_bias(input_length, dense_inputs, i)
x = self.encoder_layers[i](
x,
attention_mask=encoder_mask,
position_bias=position_bias,
training=training)
+ if self.config.return_attention_scores:
+ x, attention_scores = x
+ attention_outputs.append(attention_scores)
encoded = self.output_norm(x)
encoded = self.output_dropout(encoded, training=training)
- return encoded
+ if self.config.return_attention_scores:
+ return encoded, attention_outputs
+ else:
+ return encoded
class Decoder(Module):
@@ -1178,17 +1226,34 @@ def __init__(self,
self.target_embed = shared_embedding
self.target_dropout = Dropout(self.config.dropout_rate,)
# Position bias for the target self attention.
- self.relative_embedding = RelativePositionEmbedding(
- num_heads=self.config.num_heads,
- relative_attention_num_buckets=self.config
- .relative_attention_num_buckets,
- relative_attention_max_distance=self.config
- .relative_attention_max_distance,
- bidirectional=self.config.bidirectional,
- embeddings_initializer=self.config.relative_embeddings_initializer,
- dtype=self.dtype,
- compute_dtype=self.compute_dtype,
- name="relative_posemb")
+ if config.use_shared_relative_position_bias:
+ self.relative_embedding = RelativePositionEmbedding(
+ num_heads=self.config.num_heads,
+ relative_attention_num_buckets=self.config
+ .relative_attention_num_buckets,
+ relative_attention_max_distance=self.config
+ .relative_attention_max_distance,
+ bidirectional=self.config.bidirectional,
+ embeddings_initializer=self.config.relative_embeddings_initializer,
+ dtype=self.dtype,
+ compute_dtype=self.compute_dtype,
+ name="relative_posemb")
+ else:
+ self.relative_embeddings = []
+ for layer_idx in range(self.config.num_decoder_layers):
+ relative_embedding = RelativePositionEmbedding(
+ num_heads=self.config.num_heads,
+ relative_attention_num_buckets=self.config
+ .relative_attention_num_buckets,
+ relative_attention_max_distance=self.config
+ .relative_attention_max_distance,
+ bidirectional=self.config.bidirectional,
+ embeddings_initializer=self.config
+ .relative_embeddings_initializer,
+ dtype=self.dtype,
+ compute_dtype=self.compute_dtype,
+ name=f"relative_posemb_{layer_idx}")
+ self.relative_embeddings.append(relative_embedding)
self.decoder_layers = []
for layer_idx in range(self.config.num_decoder_layers):
if self.config.layer_sharing and layer_idx > 0:
@@ -1221,6 +1286,13 @@ def __init__(self,
dtype=self.dtype,
name="logits")
+ @tf.Module.with_name_scope
+ def get_relpos_bias(self, input_length: int, layer_idx: int) -> tf.Tensor:
+ if self.config.use_shared_relative_position_bias:
+ return self.relative_embedding(input_length, input_length)
+ else:
+ return self.relative_embeddings[layer_idx](input_length, input_length)
+
@tf.Module.with_name_scope
def __call__(self,
decoder_input_tokens,
@@ -1248,7 +1320,10 @@ def __call__(self,
training: Whether it is training pass, affecting dropouts.
Returns:
- output of a transformer encoder.
+ output of a transformer encoder including
+ 1. logits: Logits for each word in the vocab.
+ 2. raw_logits: Logits along the moded dimension.
+ 3. cache: Used for decoding in inference mode.
"""
cfg = self.config
# Casts inputs to the dtype.
@@ -1257,16 +1332,18 @@ def __call__(self,
decoder_mask = tf.cast(decoder_mask, self.compute_dtype)
if encoder_decoder_mask is not None:
encoder_decoder_mask = tf.cast(encoder_decoder_mask, self.compute_dtype)
- x = self.target_embed(decoder_input_tokens, one_hot=cfg.one_hot_embedding)
+ x = self.target_embed(decoder_input_tokens, one_hot=cfg.one_hot_embedding) # pyrefly: ignore[not-callable]
tensor_shape = tf_utils.get_shape_list(x)
tensor_shape[-2] = 1
x = self.target_dropout(x, noise_shape=tensor_shape, training=training)
- if cache is not None:
- position_bias = self.relative_embedding(max_decode_len, max_decode_len)
- else:
- input_length = tf_utils.get_shape_list(decoder_input_tokens)[1]
- position_bias = self.relative_embedding(input_length, input_length)
- for i in range(cfg.num_decoder_layers):
+
+ for i in range(cfg.num_decoder_layers): # pyrefly: ignore[bad-argument-type]
+ if cache is not None:
+ position_bias = self.get_relpos_bias(max_decode_len, i)
+ else:
+ input_length = tf_utils.get_shape_list(decoder_input_tokens)[1]
+ position_bias = self.get_relpos_bias(input_length, i)
+
if cache is None:
x, _ = self.decoder_layers[i](
x,
@@ -1296,7 +1373,7 @@ def __call__(self,
logits = logits / math.sqrt(cfg.d_model)
else:
logits = self.logits_dense(output)
- return logits, cache
+ return dict(logits=logits, cache=cache, raw_logits=output)
class T5Transformer(Module):
@@ -1310,6 +1387,7 @@ def __init__(self,
# Builds the model components.
shared_embedding = config.shared_embedding
self.compute_dtype = compute_dtype
+ self.config = config
self.decoder_cfg = dataclasses.replace(config, bidirectional=False)
if self.decoder_cfg.num_decoder_layers is None:
self.decoder_cfg.num_decoder_layers = self.decoder_cfg.num_layers
@@ -1327,12 +1405,12 @@ def __init__(self,
self.shared_embedding = None
self.encoder = Encoder(
self.encoder_cfg,
- self.shared_embedding,
+ self.shared_embedding, # pyrefly: ignore[bad-argument-type]
dtype=self.dtype,
compute_dtype=self.compute_dtype)
self.decoder = Decoder(
self.decoder_cfg,
- self.shared_embedding,
+ self.shared_embedding, # pyrefly: ignore[bad-argument-type]
dtype=self.dtype,
compute_dtype=self.compute_dtype)
@@ -1390,7 +1468,7 @@ def decode(
cache=None,
max_decode_len=None,
decode=False,
- training=False):
+ training=False) -> Dict[str, tf.Tensor]:
eligible_inputs_array = []
if encoder_input_tokens is not None:
eligible_inputs = tf.cast(
@@ -1447,7 +1525,7 @@ def decode(
decoder_mask = (1.0 - tf.cast(decoder_mask, self.compute_dtype)) * -1e9
encoder_decoder_mask = (
1.0 - tf.cast(encoder_decoder_mask, self.compute_dtype)) * -1e9
- logits, cache = self.decoder(
+ outputs = self.decoder(
decoder_input_tokens,
encoded,
decode_position=decode_position,
@@ -1457,7 +1535,8 @@ def decode(
max_decode_len=max_decode_len,
decode=decode,
training=training)
- return dict(logits=logits, encoded=encoded, cache=cache)
+ outputs["encoded"] = encoded
+ return outputs
@tf.Module.with_name_scope
def __call__(self,
@@ -1492,6 +1571,8 @@ def __call__(self,
encoder_dense_inputs=encoder_dense_inputs,
encoder_dense_segment_ids=encoder_dense_segment_ids,
training=training)
+ if self.config.return_attention_scores:
+ encoded, attn_scores = encoded
outputs = self.decode(
encoded=encoded,
decoder_target_tokens=decoder_target_tokens,
@@ -1503,6 +1584,8 @@ def __call__(self,
decoder_segment_ids=decoder_segment_ids,
training=training)
outputs["encoded"] = encoded
+ if self.config.return_attention_scores:
+ outputs["attention_scores"] = attn_scores # pyrefly: ignore[unbound-name]
return outputs
@property
diff --git a/official/nlp/modeling/models/t5_test.py b/official/nlp/modeling/models/t5_test.py
index 53e04911949..1cec42e6d54 100644
--- a/official/nlp/modeling/models/t5_test.py
+++ b/official/nlp/modeling/models/t5_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -62,7 +62,7 @@ def test_embed(self, dtype):
features=4,
compute_dtype=dtype,
name="foo",
- embeddings_initializer=tf.keras.initializers.Zeros())
+ embeddings_initializer=tf_keras.initializers.Zeros())
self.assertAllClose(l(inputs), tf.zeros((2, 2, 4), dtype))
@parameterized.named_parameters(("bfloat16", tf.bfloat16),
@@ -82,7 +82,7 @@ def test_linear(self, dtype):
l = t5.Linear(
in_features=4,
out_features=4,
- w_init=tf.keras.initializers.Ones(),
+ w_init=tf_keras.initializers.Ones(),
name="foo")
inputs = tf.ones((2, 4), dtype=dtype)
outputs = l(inputs)
@@ -97,7 +97,7 @@ def test_linear3d(self):
out_features=4,
num_heads=2,
to_3d=True,
- w_init=tf.keras.initializers.Ones(),
+ w_init=tf_keras.initializers.Ones(),
name="foo")
inputs = np.ones((batch_size, 2, 4), dtype=np.float32)
self.assertEqual(l(inputs).shape, (batch_size, 2, 2, 4))
@@ -107,7 +107,7 @@ def test_linear3d(self):
out_features=4,
num_heads=2,
to_3d=False,
- w_init=tf.keras.initializers.Ones(),
+ w_init=tf_keras.initializers.Ones(),
name="foo")
inputs = np.ones((batch_size, 2, 2, 2), dtype=np.float32)
self.assertEqual(l(inputs).shape, (batch_size, 2, 4))
@@ -140,14 +140,14 @@ def test_relative_position(self, dtype):
l = t5.RelativePositionEmbedding(
num_heads=4,
bidirectional=False,
- embeddings_initializer=tf.keras.initializers.Ones(),
+ embeddings_initializer=tf_keras.initializers.Ones(),
compute_dtype=dtype,
name="foo")
self.assertEqual(l(4, 2).shape, (1, 4, 4, 2))
l = t5.RelativePositionEmbedding(
num_heads=4,
bidirectional=True,
- embeddings_initializer=tf.keras.initializers.Ones(),
+ embeddings_initializer=tf_keras.initializers.Ones(),
compute_dtype=dtype,
name="bar")
outputs = l(4, 2)
@@ -172,7 +172,7 @@ def test_attention(self, distribution):
pos_embed = t5.RelativePositionEmbedding(
num_heads=4,
bidirectional=False,
- embeddings_initializer=tf.keras.initializers.Ones(),
+ embeddings_initializer=tf_keras.initializers.Ones(),
name="pos_embed")
position_bias = pos_embed(from_seq_length, from_seq_length)
l = t5.MultiHeadAttention(d_model=4, d_kv=2, num_heads=4, dropout_rate=0.1)
@@ -234,7 +234,7 @@ def test_attention_layers(self, distribution):
pos_embed = t5.RelativePositionEmbedding(
num_heads=head_size,
bidirectional=False,
- embeddings_initializer=tf.keras.initializers.Ones(),
+ embeddings_initializer=tf_keras.initializers.Ones(),
name="pos_embed")
l = t5.SelfAttention(
d_model=4, d_kv=head_size, num_heads=num_heads, dropout_rate=0.1)
@@ -302,7 +302,7 @@ def test_encoder_block(self):
pos_embed = t5.RelativePositionEmbedding(
num_heads=2,
bidirectional=True,
- embeddings_initializer=tf.keras.initializers.Ones(),
+ embeddings_initializer=tf_keras.initializers.Ones(),
name="bar")
attention_mask = t5.make_attention_mask(
tf.ones((batch_size, from_seq_length)),
@@ -322,7 +322,7 @@ def test_encdec_block(self):
pos_embed = t5.RelativePositionEmbedding(
num_heads=2,
bidirectional=True,
- embeddings_initializer=tf.keras.initializers.Ones(),
+ embeddings_initializer=tf_keras.initializers.Ones(),
name="bar")
encoder_decoder_mask = t5.make_attention_mask(
tf.ones((batch_size, from_seq_length)),
@@ -348,12 +348,36 @@ def test_encoder(self, dtype):
num_heads=4,
d_ff=16,
vocab_size=10,
- vocab_embeddings_initializer=tf.keras.initializers.Ones(),
- relative_embeddings_initializer=tf.keras.initializers.Ones())
+ vocab_embeddings_initializer=tf_keras.initializers.Ones(),
+ relative_embeddings_initializer=tf_keras.initializers.Ones())
encoder = t5.Encoder(config, compute_dtype=dtype)
encoded = encoder(tf.zeros((4, 8), dtype=tf.int32))
self.assertEqual(encoded.shape, (4, 8, config.d_model))
+ @parameterized.named_parameters(("return_score", True),
+ ("not_return_score", False))
+ def test_encoder_att_scores(self, return_attention_scores):
+ config = t5.T5TransformerParams(
+ num_layers=2,
+ d_model=4,
+ d_kv=3,
+ num_heads=4,
+ d_ff=16,
+ vocab_size=10,
+ vocab_embeddings_initializer=tf_keras.initializers.Ones(),
+ relative_embeddings_initializer=tf_keras.initializers.Ones(),
+ return_attention_scores=return_attention_scores)
+ encoder = t5.Encoder(config, compute_dtype=tf.float32)
+ encoded = encoder(tf.zeros((4, 8), dtype=tf.int32))
+ if return_attention_scores:
+ encoded, scores = encoded
+ self.assertEqual(encoded.shape, (4, 8, config.d_model))
+ self.assertIsNotNone(scores)
+ self.assertLen(scores, 2)
+ self.assertEqual(scores[0].shape, (4, 4, 8, 8))
+ else:
+ self.assertEqual(encoded.shape, (4, 8, config.d_model))
+
@parameterized.named_parameters(("bfloat16", tf.bfloat16),
("float32", tf.float32))
def test_encoder_with_dense(self, dtype):
@@ -364,8 +388,8 @@ def test_encoder_with_dense(self, dtype):
num_heads=4,
d_ff=16,
vocab_size=10,
- vocab_embeddings_initializer=tf.keras.initializers.Ones(),
- relative_embeddings_initializer=tf.keras.initializers.Ones())
+ vocab_embeddings_initializer=tf_keras.initializers.Ones(),
+ relative_embeddings_initializer=tf_keras.initializers.Ones())
encoder = t5.Encoder(config, compute_dtype=dtype)
encoded = encoder(
tf.zeros((4, 8), dtype=tf.int32),
@@ -382,8 +406,8 @@ def test_encoder_only_dense(self, dtype):
num_heads=4,
d_ff=16,
vocab_size=10,
- vocab_embeddings_initializer=tf.keras.initializers.Ones(),
- relative_embeddings_initializer=tf.keras.initializers.Ones())
+ vocab_embeddings_initializer=tf_keras.initializers.Ones(),
+ relative_embeddings_initializer=tf_keras.initializers.Ones())
encoder = t5.Encoder(config, compute_dtype=dtype)
encoded = encoder(dense_inputs=tf.ones((4, 2, 4), dtype=dtype))
self.assertEqual(encoded.shape, (4, 2, config.d_model))
@@ -397,13 +421,14 @@ def test_decoder(self):
num_heads=4,
d_ff=16,
vocab_size=10,
- vocab_embeddings_initializer=tf.keras.initializers.Ones(),
- relative_embeddings_initializer=tf.keras.initializers.Ones())
+ vocab_embeddings_initializer=tf_keras.initializers.Ones(),
+ relative_embeddings_initializer=tf_keras.initializers.Ones())
decoder = t5.Decoder(config)
batch_size = 4
targets = tf.zeros((4, 8), dtype=tf.int32)
encoded = tf.zeros((4, 8, config.d_model), dtype=tf.float32)
- logits, cache = decoder(targets, encoded)
+ outputs = decoder(targets, encoded)
+ logits = outputs["logits"]
self.assertEqual(logits.shape, (4, 8, config.vocab_size))
cache = {}
@@ -412,13 +437,15 @@ def test_decoder(self):
cache[1] = _create_cache(batch_size, max_decode_len, config.num_heads,
config.d_kv)
targets = tf.zeros((4, 1), dtype=tf.int32)
- logits, cache = decoder(
+ outputs = decoder(
targets,
encoded,
decode_position=2,
cache=cache,
decode=True,
max_decode_len=max_decode_len)
+ logits = outputs["logits"]
+ cache = outputs["cache"]
self.assertEqual(logits.shape, (batch_size, 1, config.vocab_size))
for entry in cache.values():
for tensor in entry.values():
@@ -476,11 +503,66 @@ def test_transformer(self, ffn_activations, logits_via_embedding,
self.assertEqual(outputs["logits"].shape,
(batch_size, 1, config.vocab_size))
for v in transformer.trainable_variables:
- print(v.name, v.shape)
+ self.assertEqual(v.dtype, tf.float32)
+
+ def test_transformer_return_attn_scores(self):
+ max_decode_len = 10
+ config = t5.T5TransformerParams(
+ num_layers=1,
+ d_model=8,
+ d_kv=4,
+ num_heads=4,
+ d_ff=32,
+ vocab_size=10,
+ shared_embedding=True,
+ layer_sharing=False,
+ ffn_activations=("relu",),
+ logits_via_embedding=True,
+ return_attention_scores=True,
+ )
+ transformer = t5.T5Transformer(config, compute_dtype=tf.float32)
+ self.assertLen(transformer.trainable_variables, 26)
+ inputs = tf.convert_to_tensor(
+ np.array([[2, 2, 1, 3, 1, 0], [3, 3, 1, 2, 2, 1]])
+ )
+ segments = tf.convert_to_tensor(
+ np.array([[1, 1, 1, 2, 2, 0], [1, 1, 1, 2, 2, 2]])
+ )
+
+ outputs = transformer(
+ encoder_input_tokens=inputs,
+ decoder_input_tokens=inputs,
+ decoder_target_tokens=inputs,
+ encoder_segment_ids=segments,
+ decoder_segment_ids=segments,
+ )
+ self.assertIn("attention_scores", outputs)
+ self.assertLen(outputs["attention_scores"], 1)
+ self.assertEqual(outputs["attention_scores"][0].shape, (2, 4, 6, 6))
+ cache = {}
+ batch_size = 2
+ cache[0] = _create_cache(
+ batch_size,
+ max_decode_len,
+ config.num_heads,
+ config.d_kv,
+ dtype=tf.float32,
+ )
+ outputs = transformer.decode(
+ encoder_input_tokens=inputs,
+ encoded=outputs["encoded"],
+ decoder_target_tokens=tf.ones((batch_size, 1), dtype=tf.int32),
+ decode_position=1,
+ decode=True,
+ max_decode_len=max_decode_len,
+ cache=cache)
+ self.assertEqual(outputs["logits"].shape,
+ (batch_size, 1, config.vocab_size))
+ for v in transformer.trainable_variables:
self.assertEqual(v.dtype, tf.float32)
@parameterized.named_parameters(
- ("t5_10", ("relu",), True, 26, False, tf.float32),)
+ ("t5_10_dense", ("relu",), True, 26, False, tf.float32),)
def test_transformer_with_dense(self, ffn_activations, logits_via_embedding,
expect_num_variables, layer_sharing, dtype):
max_decode_len = 10
@@ -496,6 +578,7 @@ def test_transformer_with_dense(self, ffn_activations, logits_via_embedding,
ffn_activations=ffn_activations,
logits_via_embedding=logits_via_embedding)
transformer = t5.T5Transformer(config, compute_dtype=dtype)
+
self.assertLen(transformer.trainable_variables, expect_num_variables)
inputs = tf.convert_to_tensor(
np.array([[2, 2, 1, 3, 1, 0], [3, 3, 1, 2, 2, 1]]))
@@ -528,7 +611,74 @@ def test_transformer_with_dense(self, ffn_activations, logits_via_embedding,
self.assertEqual(outputs["logits"].shape,
(batch_size, 1, config.vocab_size))
for v in transformer.trainable_variables:
- print(v.name, v.shape)
+ self.assertEqual(v.dtype, tf.float32)
+
+ @parameterized.named_parameters(
+ ("t5_10_dense_layerwise_relpos",
+ ("relu",), True, 26, False, tf.float32, False, 1),
+ ("t5_10_dense_shared_relpos_d2",
+ ("relu",), True, 39, False, tf.float32, True, 2),
+ ("t5_10_dense_layerwise_relpos_d2",
+ ("relu",), True, 40, False, tf.float32, False, 2),
+ )
+ def test_transformer_with_lw_relpos(self, ffn_activations,
+ logits_via_embedding,
+ expect_num_variables, layer_sharing,
+ dtype, use_shared_relpos,
+ num_decoder_layers):
+ max_decode_len = 10
+ config = t5.T5TransformerParams(
+ num_layers=1,
+ num_decoder_layers=num_decoder_layers,
+ d_model=8,
+ d_kv=4,
+ num_heads=4,
+ d_ff=32,
+ vocab_size=10,
+ shared_embedding=True,
+ layer_sharing=layer_sharing,
+ ffn_activations=ffn_activations,
+ logits_via_embedding=logits_via_embedding,
+ use_shared_relative_position_bias=use_shared_relpos)
+ transformer = t5.T5Transformer(config, compute_dtype=dtype)
+
+ self.assertLen(transformer.trainable_variables, expect_num_variables)
+ inputs = tf.convert_to_tensor(
+ np.array([[2, 2, 1, 3, 1, 0], [3, 3, 1, 2, 2, 1]]))
+ segments = tf.convert_to_tensor(
+ np.array([[1, 1, 1, 2, 2, 0], [1, 1, 1, 2, 2, 2]]))
+
+ dense_inputs = tf.convert_to_tensor(np.random.randn(2, 2, 8), dtype=dtype)
+ dense_segments = tf.convert_to_tensor(np.array([[1, 2], [1, 2]]))
+ outputs = transformer(
+ encoder_input_tokens=inputs,
+ encoder_dense_inputs=dense_inputs,
+ decoder_input_tokens=inputs,
+ decoder_target_tokens=inputs,
+ encoder_segment_ids=segments,
+ encoder_dense_segment_ids=dense_segments,
+ decoder_segment_ids=segments)
+ cache = {}
+ batch_size = 2
+ for i in range(num_decoder_layers):
+ cache[i] = _create_cache(
+ batch_size,
+ max_decode_len,
+ config.num_heads,
+ config.d_kv,
+ dtype=dtype)
+ outputs = transformer.decode(
+ encoder_input_tokens=inputs,
+ encoder_dense_inputs=dense_inputs,
+ encoded=outputs["encoded"],
+ decoder_target_tokens=tf.ones((batch_size, 1), dtype=tf.int32),
+ decode_position=1,
+ decode=True,
+ max_decode_len=max_decode_len,
+ cache=cache)
+ self.assertEqual(outputs["logits"].shape,
+ (batch_size, 1, config.vocab_size))
+ for v in transformer.trainable_variables:
self.assertEqual(v.dtype, tf.float32)
@parameterized.named_parameters(
@@ -580,7 +730,6 @@ def test_transformer_with_dense_only(self, ffn_activations,
self.assertEqual(outputs["logits"].shape,
(batch_size, 1, config.vocab_size))
for v in transformer.trainable_variables:
- print(v.name, v.shape)
self.assertEqual(v.dtype, tf.float32)
@parameterized.named_parameters(
@@ -635,7 +784,6 @@ def test_transformer_different_num_decoder_layers(self, ffn_activations,
self.assertEqual(outputs["logits"].shape,
(batch_size, 1, config.vocab_size))
for v in transformer.trainable_variables:
- print(v.name, v.shape)
self.assertEqual(v.dtype, tf.float32)
diff --git a/official/nlp/modeling/models/xlnet.py b/official/nlp/modeling/models/xlnet.py
index eea637e0316..5c3c5f2fa3a 100644
--- a/official/nlp/modeling/models/xlnet.py
+++ b/official/nlp/modeling/models/xlnet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,13 +17,13 @@
from typing import Any, Mapping, Optional, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
from official.nlp.modeling import networks
-class XLNetMaskedLM(tf.keras.layers.Layer):
+class XLNetMaskedLM(tf_keras.layers.Layer):
"""XLNet pretraining head."""
def __init__(self,
@@ -40,12 +40,12 @@ def __init__(self,
self._activation = activation
def build(self, input_shape):
- self.dense = tf.keras.layers.Dense(
+ self.dense = tf_keras.layers.Dense(
units=self._hidden_size,
activation=self._activation,
kernel_initializer=self._initializer,
name='transform/dense')
- self.layer_norm = tf.keras.layers.LayerNormalization(
+ self.layer_norm = tf_keras.layers.LayerNormalization(
axis=-1, epsilon=1e-12, name='transform/LayerNorm')
self.bias = self.add_weight(
'output_bias/bias',
@@ -76,8 +76,8 @@ def get_config(self) -> Mapping[str, Any]:
return dict(list(base_config.items()) + list(config.items()))
-@tf.keras.utils.register_keras_serializable(package='Text')
-class XLNetPretrainer(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class XLNetPretrainer(tf_keras.Model):
"""XLNet-based pretrainer.
This is an implementation of the network structure surrounding a
@@ -96,7 +96,7 @@ class XLNetPretrainer(tf.keras.Model):
def __init__(
self,
- network: Union[tf.keras.layers.Layer, tf.keras.Model],
+ network: Union[tf_keras.layers.Layer, tf_keras.Model],
mlm_activation=None,
mlm_initializer='glorot_uniform',
name: Optional[str] = None,
@@ -117,7 +117,7 @@ def __init__(
hidden_size=self._hidden_size,
initializer=self._initializer)
- def call(self, inputs: Mapping[str, Any]):
+ def call(self, inputs: Mapping[str, Any]): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
input_word_ids = inputs['input_word_ids']
input_type_ids = inputs['input_type_ids']
masked_tokens = inputs['masked_tokens']
@@ -152,8 +152,8 @@ def checkpoint_items(self):
return dict(encoder=self._network)
-@tf.keras.utils.register_keras_serializable(package='Text')
-class XLNetClassifier(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class XLNetClassifier(tf_keras.Model):
"""Classifier model based on XLNet.
This is an implementation of the network structure surrounding a
@@ -176,9 +176,9 @@ class XLNetClassifier(tf.keras.Model):
def __init__(
self,
- network: Union[tf.keras.layers.Layer, tf.keras.Model],
+ network: Union[tf_keras.layers.Layer, tf_keras.Model],
num_classes: int,
- initializer: tf.keras.initializers.Initializer = 'random_normal',
+ initializer: tf_keras.initializers.Initializer = 'random_normal', # pyrefly: ignore[bad-function-definition]
summary_type: str = 'last',
dropout_rate: float = 0.1,
head_name: str = 'sentence_prediction', # pytype: disable=annotation-type-mismatch # typed-keras
@@ -212,7 +212,7 @@ def __init__(
cls_token_idx=cls_token_idx,
name=head_name)
- def call(self, inputs: Mapping[str, Any]):
+ def call(self, inputs: Mapping[str, Any]): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
input_ids = inputs['input_word_ids']
segment_ids = inputs['input_type_ids']
input_mask = tf.cast(inputs['input_mask'], tf.float32)
@@ -244,8 +244,8 @@ def checkpoint_items(self):
return items
-@tf.keras.utils.register_keras_serializable(package='Text')
-class XLNetSpanLabeler(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class XLNetSpanLabeler(tf_keras.Model):
"""Span labeler model based on XLNet.
This is an implementation of the network structure surrounding a
@@ -266,12 +266,12 @@ class XLNetSpanLabeler(tf.keras.Model):
def __init__(
self,
- network: Union[tf.keras.layers.Layer, tf.keras.Model],
+ network: Union[tf_keras.layers.Layer, tf_keras.Model],
start_n_top: int = 5,
end_n_top: int = 5,
dropout_rate: float = 0.1,
- span_labeling_activation: tf.keras.initializers.Initializer = 'tanh',
- initializer: tf.keras.initializers.Initializer = 'glorot_uniform', # pytype: disable=annotation-type-mismatch # typed-keras
+ span_labeling_activation: tf_keras.initializers.Initializer = 'tanh', # pyrefly: ignore[bad-function-definition]
+ initializer: tf_keras.initializers.Initializer = 'glorot_uniform', # pytype: disable=annotation-type-mismatch # typed-keras
**kwargs):
super().__init__(**kwargs)
self._config = {
@@ -305,7 +305,7 @@ def __init__(
dropout_rate=self._dropout_rate,
initializer=self._initializer)
- def call(self, inputs: Mapping[str, Any]):
+ def call(self, inputs: Mapping[str, Any]): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
input_word_ids = inputs['input_word_ids']
input_type_ids = inputs['input_type_ids']
input_mask = inputs['input_mask']
diff --git a/official/nlp/modeling/models/xlnet_test.py b/official/nlp/modeling/models/xlnet_test.py
index e22883508da..c89e2ba9ea4 100644
--- a/official/nlp/modeling/models/xlnet_test.py
+++ b/official/nlp/modeling/models/xlnet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,14 +17,13 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling import networks
from official.nlp.modeling.models import xlnet
-def _get_xlnet_base() -> tf.keras.layers.Layer:
+def _get_xlnet_base() -> tf_keras.layers.Layer:
"""Returns a trivial base XLNet model."""
return networks.XLNetBase(
vocab_size=100,
@@ -37,17 +36,14 @@ def _get_xlnet_base() -> tf.keras.layers.Layer:
attention_dropout_rate=0.,
attention_type='bi',
bi_data=True,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
two_stream=False,
tie_attention_biases=True,
reuse_length=0,
inner_activation='relu')
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class XLNetMaskedLMTest(keras_parameterized.TestCase):
+class XLNetMaskedLMTest(tf.test.TestCase):
def test_xlnet_masked_lm_head(self):
hidden_size = 10
@@ -62,8 +58,7 @@ def test_xlnet_masked_lm_head(self):
self.assertAllClose(mlm_output.shape, (batch_size, hidden_size))
-@keras_parameterized.run_all_keras_modes
-class XLNetPretrainerTest(keras_parameterized.TestCase):
+class XLNetPretrainerTest(tf.test.TestCase):
def test_xlnet_trainer(self):
"""Validates that the Keras object can be created."""
@@ -75,19 +70,19 @@ def test_xlnet_trainer(self):
# Create an XLNet trainer with the created network.
xlnet_trainer_model = xlnet.XLNetPretrainer(network=xlnet_base)
inputs = dict(
- input_word_ids=tf.keras.layers.Input(
+ input_word_ids=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_word_ids'),
- input_type_ids=tf.keras.layers.Input(
+ input_type_ids=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_type_ids'),
- input_mask=tf.keras.layers.Input(
+ input_mask=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_mask'),
- permutation_mask=tf.keras.layers.Input(
+ permutation_mask=tf_keras.layers.Input(
shape=(seq_length, seq_length,), dtype=tf.int32,
name='permutation_mask'),
- target_mapping=tf.keras.layers.Input(
+ target_mapping=tf_keras.layers.Input(
shape=(num_predictions, seq_length), dtype=tf.int32,
name='target_mapping'),
- masked_tokens=tf.keras.layers.Input(
+ masked_tokens=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='masked_tokens'))
logits, _ = xlnet_trainer_model(inputs)
@@ -144,8 +139,7 @@ def test_serialize_deserialize(self):
new_xlnet_trainer_model.get_config())
-@keras_parameterized.run_all_keras_modes
-class XLNetClassifierTest(keras_parameterized.TestCase):
+class XLNetClassifierTest(tf.test.TestCase, parameterized.TestCase):
def test_xlnet_trainer(self):
"""Validate that the Keras object can be created."""
@@ -158,20 +152,20 @@ def test_xlnet_trainer(self):
xlnet_trainer_model = xlnet.XLNetClassifier(
network=xlnet_base,
num_classes=num_classes,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
summary_type='last',
dropout_rate=0.1)
inputs = dict(
- input_word_ids=tf.keras.layers.Input(
+ input_word_ids=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_word_ids'),
- input_type_ids=tf.keras.layers.Input(
+ input_type_ids=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_type_ids'),
- input_mask=tf.keras.layers.Input(
+ input_mask=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_mask'),
- permutation_mask=tf.keras.layers.Input(
+ permutation_mask=tf_keras.layers.Input(
shape=(seq_length, seq_length,), dtype=tf.int32,
name='permutation_mask'),
- masked_tokens=tf.keras.layers.Input(
+ masked_tokens=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='masked_tokens'))
logits = xlnet_trainer_model(inputs)
@@ -190,7 +184,7 @@ def test_xlnet_tensor_call(self, num_classes):
xlnet_trainer_model = xlnet.XLNetClassifier(
network=xlnet_base,
num_classes=num_classes,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
summary_type='last',
dropout_rate=0.1)
@@ -215,7 +209,7 @@ def test_serialize_deserialize(self):
xlnet_trainer_model = xlnet.XLNetClassifier(
network=xlnet_base,
num_classes=2,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
summary_type='last',
dropout_rate=0.1)
@@ -232,8 +226,7 @@ def test_serialize_deserialize(self):
new_xlnet_trainer_model.get_config())
-@keras_parameterized.run_all_keras_modes
-class XLNetSpanLabelerTest(keras_parameterized.TestCase):
+class XLNetSpanLabelerTest(tf.test.TestCase):
def test_xlnet_trainer(self):
"""Validate that the Keras object can be created."""
@@ -247,21 +240,21 @@ def test_xlnet_trainer(self):
network=xlnet_base,
start_n_top=top_n,
end_n_top=top_n,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
span_labeling_activation='tanh',
dropout_rate=0.1)
inputs = dict(
- input_word_ids=tf.keras.layers.Input(
+ input_word_ids=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_word_ids'),
- input_type_ids=tf.keras.layers.Input(
+ input_type_ids=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_type_ids'),
- input_mask=tf.keras.layers.Input(
+ input_mask=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_mask'),
- paragraph_mask=tf.keras.layers.Input(
+ paragraph_mask=tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='paragraph_mask'),
- class_index=tf.keras.layers.Input(
+ class_index=tf_keras.layers.Input(
shape=(), dtype=tf.int32, name='class_index'),
- start_positions=tf.keras.layers.Input(
+ start_positions=tf_keras.layers.Input(
shape=(), dtype=tf.int32, name='start_positions'))
outputs = xlnet_trainer_model(inputs)
self.assertIsInstance(outputs, dict)
@@ -307,7 +300,7 @@ def test_serialize_deserialize(self):
network=xlnet_base,
start_n_top=2,
end_n_top=2,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
span_labeling_activation='tanh',
dropout_rate=0.1)
diff --git a/official/nlp/modeling/networks/README.md b/official/nlp/modeling/networks/README.md
index b192399a727..87cc571e84a 100644
--- a/official/nlp/modeling/networks/README.md
+++ b/official/nlp/modeling/networks/README.md
@@ -2,38 +2,50 @@
Networks are combinations of `tf.keras` layers (and possibly other networks).
They are `tf.keras` models that would not be trained alone. It encapsulates
-common network structures like a transformer encoder into an easily
-handled object with a standardized configuration.
-
-* [`BertEncoder`](bert_encoder.py) implements a bi-directional
-Transformer-based encoder as described in ["BERT: Pre-training of Deep
-Bidirectional Transformers for Language Understanding"](https://arxiv.org/abs/1810.04805).
-It includes the embedding lookups, transformer layers and pooling layer.
-
-* [`AlbertEncoder`](albert_encoder.py) implements a
-Transformer-encoder described in the paper ["ALBERT: A Lite BERT for
-Self-supervised Learning of Language Representations"]
-(https://arxiv.org/abs/1909.11942). Compared with [BERT](https://arxiv.org/abs/1810.04805),
-ALBERT refactorizes embedding parameters into two smaller matrices and shares
-parameters across layers.
-
-* [`MobileBERTEncoder`](mobile_bert_encoder.py) implements the
-MobileBERT network described in the paper ["MobileBERT: a Compact Task-Agnostic
-BERT for Resource-Limited Devices"](https://arxiv.org/abs/2004.02984).
-
-* [`Classification`](classification.py) contains a single hidden layer, and is
-intended for use as a classification or regression (if number of classes is set
-to 1) head.
-
-* [`PackedSequenceEmbedding`](packed_sequence_embedding.py) implements an
-embedding network that supports packed sequences and position ids.
-
-* [`SpanLabeling`](span_labeling.py) implements a single-span labeler
-(that is, a prediction head that can predict one start and end index per batch
-item) based on a single dense hidden layer. It can be used in the SQuAD task.
-
-* [`XLNetBase`](xlnet_base.py) implements the base network used in "XLNet:
-Generalized Autoregressive Pretraining for Language Understanding"
-(https://arxiv.org/abs/1906.08237). It includes embedding lookups,
-relative position encodings, mask computations, segment matrix computations and
-Transformer XL layers using one or two stream relative self-attention.
+common network structures like a transformer encoder into an easily handled
+object with a standardized configuration.
+
+* [`BertEncoder`](bert_encoder.py) implements a bi-directional
+ Transformer-based encoder as described in ["BERT: Pre-training of Deep
+ Bidirectional Transformers for Language
+ Understanding"](https://arxiv.org/abs/1810.04805). It includes the embedding
+ lookups, transformer layers and pooling layer.
+
+* [`AlbertEncoder`](albert_encoder.py) implements a Transformer-encoder
+ described in the paper ["ALBERT: A Lite BERT for Self-supervised Learning of
+ Language Representations"](https://arxiv.org/abs/1909.11942). Compared with
+ [BERT](https://arxiv.org/abs/1810.04805), ALBERT refactorizes embedding
+ parameters into two smaller matrices and shares parameters across layers.
+
+* [`MobileBERTEncoder`](mobile_bert_encoder.py) implements the MobileBERT
+ network described in the paper
+ ["MobileBERT: a Compact Task-Agnostic BERT for Resource-Limited Devices"](https://arxiv.org/abs/2004.02984).
+
+* [`Classification`](classification.py) contains a single hidden layer, and is
+ intended for use as a classification or regression (if number of classes is
+ set to 1) head.
+
+* [`PackedSequenceEmbedding`](packed_sequence_embedding.py) implements an
+ embedding network that supports packed sequences and position ids.
+
+* [`SpanLabeling`](span_labeling.py) implements a single-span labeler (that
+ is, a prediction head that can predict one start and end index per batch
+ item) based on a single dense hidden layer. It can be used in the SQuAD
+ task.
+
+* [`XLNetBase`](xlnet_base.py) implements the base network used in "XLNet:
+ Generalized Autoregressive Pretraining for Language Understanding"
+ (https://arxiv.org/abs/1906.08237). It includes embedding lookups, relative
+ position encodings, mask computations, segment matrix computations and
+ Transformer XL layers using one or two stream relative self-attention.
+
+* [`FNet`](fnet.py) implements the encoder model from
+ ["FNet: Mixing Tokens with Fourier Transforms"](https://aclanthology.org/2022.naacl-main.319/).
+ FNet has the same structure as a Transformer encoder, except that all or
+ most of the self-attention sublayers are replaced with Fourier sublayers.
+
+* [`Sparse Mixer`](sparse_mixer.py) implements the encoder model from
+ ["Sparse Mixers: Combining MoE and Mixing to build a more efficient BERT "](https://arxiv.org/abs/2205.12399/).
+ Sparse Mixer consists of layers of heterogeneous encoder blocks. Each
+ encoder block contains a linear mixing or an attention sublayer together
+ with a (dense) MLP or sparsely activated Mixture-of-Experts sublayer.
diff --git a/official/nlp/modeling/networks/__init__.py b/official/nlp/modeling/networks/__init__.py
index b9e766e2cfc..2883449a9c9 100644
--- a/official/nlp/modeling/networks/__init__.py
+++ b/official/nlp/modeling/networks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,9 +23,11 @@
from official.nlp.modeling.networks.bert_encoder import BertEncoderV2
from official.nlp.modeling.networks.classification import Classification
from official.nlp.modeling.networks.encoder_scaffold import EncoderScaffold
+from official.nlp.modeling.networks.fnet import FNet
from official.nlp.modeling.networks.funnel_transformer import FunnelTransformerEncoder
from official.nlp.modeling.networks.mobile_bert_encoder import MobileBERTEncoder
from official.nlp.modeling.networks.packed_sequence_embedding import PackedSequenceEmbedding
from official.nlp.modeling.networks.span_labeling import SpanLabeling
from official.nlp.modeling.networks.span_labeling import XLNetSpanLabeling
+from official.nlp.modeling.networks.sparse_mixer import SparseMixer
from official.nlp.modeling.networks.xlnet_base import XLNetBase
diff --git a/official/nlp/modeling/networks/albert_encoder.py b/official/nlp/modeling/networks/albert_encoder.py
index bdca7575143..92d4f7da114 100644
--- a/official/nlp/modeling/networks/albert_encoder.py
+++ b/official/nlp/modeling/networks/albert_encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,14 +15,15 @@
"""ALBERT (https://arxiv.org/abs/1810.04805) text encoder network."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import activations
+from official.modeling import tf_utils
from official.nlp.modeling import layers
-@tf.keras.utils.register_keras_serializable(package='Text')
-class AlbertEncoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class AlbertEncoder(tf_keras.Model):
"""ALBERT (https://arxiv.org/abs/1810.04805) text encoder network.
This network implements the encoder described in the paper "ALBERT: A Lite
@@ -74,17 +75,17 @@ def __init__(self,
activation=activations.gelu,
dropout_rate=0.1,
attention_dropout_rate=0.1,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
dict_outputs=False,
**kwargs):
- activation = tf.keras.activations.get(activation)
- initializer = tf.keras.initializers.get(initializer)
+ activation = tf_keras.activations.get(activation)
+ initializer = tf_keras.initializers.get(initializer)
- word_ids = tf.keras.layers.Input(
+ word_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- mask = tf.keras.layers.Input(
+ mask = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_mask')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
if embedding_width is None:
@@ -92,13 +93,13 @@ def __init__(self,
embedding_layer = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
name='word_embeddings')
word_embeddings = embedding_layer(word_ids)
# Always uses dynamic slicing for simplicity.
position_embedding_layer = layers.PositionEmbedding(
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
max_length=max_sequence_length,
name='position_embedding')
position_embeddings = position_embedding_layer(word_embeddings)
@@ -107,27 +108,27 @@ def __init__(self,
layers.OnDeviceEmbedding(
vocab_size=type_vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_one_hot=True,
name='type_embeddings')(type_ids))
- embeddings = tf.keras.layers.Add()(
+ embeddings = tf_keras.layers.Add()(
[word_embeddings, position_embeddings, type_embeddings])
embeddings = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm',
axis=-1,
epsilon=1e-12,
dtype=tf.float32)(embeddings))
- embeddings = (tf.keras.layers.Dropout(rate=dropout_rate)(embeddings))
+ embeddings = (tf_keras.layers.Dropout(rate=dropout_rate)(embeddings))
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
if embedding_width != hidden_size:
- embeddings = tf.keras.layers.experimental.EinsumDense(
+ embeddings = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='embedding_projection')(
embeddings)
@@ -139,7 +140,7 @@ def __init__(self,
inner_activation=activation,
output_dropout=dropout_rate,
attention_dropout=attention_dropout_rate,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='transformer')
encoder_outputs = []
for _ in range(num_layers):
@@ -150,10 +151,10 @@ def __init__(self,
# like this will create a SliceOpLambda layer. This is better than a Lambda
# layer with Python code, because that is fundamentally less portable.
first_token_tensor = data[:, 0, :]
- cls_output = tf.keras.layers.Dense(
+ cls_output = tf_keras.layers.Dense(
units=hidden_size,
activation='tanh',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='pooler_transform')(
first_token_tensor)
if dict_outputs:
@@ -172,7 +173,7 @@ def __init__(self,
# created using the Functional API. Once super().__init__ is called, we
# can assign attributes to `self` - note that all `self` assignments are
# below this line.
- super(AlbertEncoder, self).__init__(
+ super().__init__(
inputs=[word_ids, mask, type_ids], outputs=outputs, **kwargs)
config_dict = {
'vocab_size': vocab_size,
@@ -183,10 +184,10 @@ def __init__(self,
'max_sequence_length': max_sequence_length,
'type_vocab_size': type_vocab_size,
'intermediate_size': intermediate_size,
- 'activation': tf.keras.activations.serialize(activation),
+ 'activation': tf_keras.activations.serialize(activation),
'dropout_rate': dropout_rate,
'attention_dropout_rate': attention_dropout_rate,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_keras.initializers.serialize(initializer),
}
# We are storing the config dict as a namedtuple here to ensure checkpoint
@@ -194,7 +195,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self._embedding_layer = embedding_layer
self._position_embedding_layer = position_embedding_layer
@@ -206,5 +207,5 @@ def get_config(self):
return dict(self._config._asdict())
@classmethod
- def from_config(cls, config):
+ def from_config(cls, config): # pyrefly: ignore[bad-override]
return cls(**config)
diff --git a/official/nlp/modeling/networks/albert_encoder_test.py b/official/nlp/modeling/networks/albert_encoder_test.py
index f7116afc915..7eed061fee9 100644
--- a/official/nlp/modeling/networks/albert_encoder_test.py
+++ b/official/nlp/modeling/networks/albert_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,26 +14,18 @@
"""Tests for ALBERT transformer-based text encoder network."""
-from __future__ import absolute_import
-from __future__ import division
-from __future__ import print_function
-
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.networks import albert_encoder
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class AlbertEncoderTest(keras_parameterized.TestCase):
+class AlbertEncoderTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(AlbertEncoderTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
@parameterized.named_parameters(
dict(testcase_name="default", expected_dtype=tf.float32),
@@ -49,15 +41,15 @@ def test_network_creation(self, expected_dtype):
num_attention_heads=2,
num_layers=3)
if expected_dtype == tf.float16:
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a small TransformerEncoder for testing.
test_network = albert_encoder.AlbertEncoder(**kwargs)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask, type_ids])
expected_data_shape = [None, sequence_length, hidden_size]
@@ -94,13 +86,13 @@ def test_network_invocation(self):
num_layers=num_layers,
type_vocab_size=num_types)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask, type_ids])
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -122,7 +114,7 @@ def test_network_invocation(self):
num_attention_heads=2,
num_layers=num_layers,
type_vocab_size=num_types)
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
_ = model.predict([word_id_data, mask_data, type_id_data])
# Tests dictionary outputs.
@@ -148,7 +140,7 @@ def test_network_invocation(self):
self.assertLen(dict_outputs["pooled_output"], num_layers)
def test_serialize_deserialize(self):
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a network object that sets all of its config options.
kwargs = dict(
vocab_size=100,
@@ -166,10 +158,10 @@ def test_serialize_deserialize(self):
network = albert_encoder.AlbertEncoder(**kwargs)
expected_config = dict(kwargs)
- expected_config["activation"] = tf.keras.activations.serialize(
- tf.keras.activations.get(expected_config["activation"]))
- expected_config["initializer"] = tf.keras.initializers.serialize(
- tf.keras.initializers.get(expected_config["initializer"]))
+ expected_config["activation"] = tf_keras.activations.serialize(
+ tf_keras.activations.get(expected_config["activation"]))
+ expected_config["initializer"] = tf_keras.initializers.serialize(
+ tf_keras.initializers.get(expected_config["initializer"]))
self.assertEqual(network.get_config(), expected_config)
# Create another network object from the first object's config.
diff --git a/official/nlp/modeling/networks/bert_dense_encoder_test.py b/official/nlp/modeling/networks/bert_dense_encoder_test.py
index 3f884e9a343..59f6341c3f3 100644
--- a/official/nlp/modeling/networks/bert_dense_encoder_test.py
+++ b/official/nlp/modeling/networks/bert_dense_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,23 +14,18 @@
"""Tests for transformer-based bert encoder network with dense features as inputs."""
-# Import libraries
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.networks import bert_encoder
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class BertEncoderV2Test(keras_parameterized.TestCase):
+class BertEncoderV2Test(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(BertEncoderV2Test, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
def test_dict_outputs_network_creation(self):
hidden_size = 32
@@ -46,14 +41,14 @@ def test_dict_outputs_network_creation(self):
with_dense_inputs=True,
**kwargs)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
- dense_inputs = tf.keras.Input(
+ dense_inputs = tf_keras.Input(
shape=(dense_sequence_length, hidden_size), dtype=tf.float32)
- dense_mask = tf.keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
- dense_type_ids = tf.keras.Input(
+ dense_mask = tf_keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
+ dense_type_ids = tf_keras.Input(
shape=(dense_sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
@@ -69,7 +64,7 @@ def test_dict_outputs_network_creation(self):
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, 3)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
expected_data_shape = [
None, sequence_length + dense_sequence_length, hidden_size
@@ -95,14 +90,14 @@ def test_dict_outputs_all_encoder_outputs_network_creation(self):
dict_outputs=True,
with_dense_inputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
- dense_inputs = tf.keras.Input(
+ dense_inputs = tf_keras.Input(
shape=(dense_sequence_length, hidden_size), dtype=tf.float32)
- dense_mask = tf.keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
- dense_type_ids = tf.keras.Input(
+ dense_mask = tf_keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
+ dense_type_ids = tf_keras.Input(
shape=(dense_sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
@@ -134,7 +129,7 @@ def test_dict_outputs_network_creation_with_float16_dtype(self):
hidden_size = 32
sequence_length = 21
dense_sequence_length = 20
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a small BertEncoder for testing.
test_network = bert_encoder.BertEncoderV2(
vocab_size=100,
@@ -144,14 +139,14 @@ def test_dict_outputs_network_creation_with_float16_dtype(self):
dict_outputs=True,
with_dense_inputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
- dense_inputs = tf.keras.Input(
+ dense_inputs = tf_keras.Input(
shape=(dense_sequence_length, hidden_size), dtype=tf.float32)
- dense_mask = tf.keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
- dense_type_ids = tf.keras.Input(
+ dense_mask = tf_keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
+ dense_type_ids = tf_keras.Input(
shape=(dense_sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
@@ -196,17 +191,17 @@ def test_dict_outputs_network_invocation(
num_attention_heads=2,
num_layers=3,
type_vocab_size=num_types,
- output_range=output_range,
dict_outputs=True,
- with_dense_inputs=True)
+ with_dense_inputs=True,
+ output_range=output_range)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- dense_inputs = tf.keras.Input(
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dense_inputs = tf_keras.Input(
shape=(dense_sequence_length, hidden_size), dtype=tf.float32)
- dense_mask = tf.keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
- dense_type_ids = tf.keras.Input(
+ dense_mask = tf_keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
+ dense_type_ids = tf_keras.Input(
shape=(dense_sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
@@ -221,7 +216,7 @@ def test_dict_outputs_network_invocation(
pooled = dict_outputs["pooled_output"]
# Create a model based off of this network:
- model = tf.keras.Model(
+ model = tf_keras.Model(
[word_ids, mask, type_ids, dense_inputs, dense_mask, dense_type_ids],
[data, pooled])
@@ -267,7 +262,7 @@ def test_dict_outputs_network_invocation(
dense_type_ids=dense_type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model(
+ model = tf_keras.Model(
[word_ids, mask, type_ids, dense_inputs, dense_mask, dense_type_ids],
[data, pooled])
outputs = model.predict([
@@ -289,7 +284,7 @@ def test_dict_outputs_network_invocation(
embedding_width=embedding_width,
dict_outputs=True)
- dense_inputs = tf.keras.Input(
+ dense_inputs = tf_keras.Input(
shape=(dense_sequence_length, embedding_width), dtype=tf.float32)
dense_input_data = np.zeros(
(batch_size, dense_sequence_length, embedding_width), dtype=float)
@@ -304,7 +299,7 @@ def test_dict_outputs_network_invocation(
dense_type_ids=dense_type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model(
+ model = tf_keras.Model(
[word_ids, mask, type_ids, dense_inputs, dense_mask, dense_type_ids],
[data, pooled])
outputs = model.predict([
@@ -326,14 +321,14 @@ def test_embeddings_as_inputs(self):
num_layers=3,
with_dense_inputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
- dense_inputs = tf.keras.Input(
+ dense_inputs = tf_keras.Input(
shape=(dense_sequence_length, hidden_size), dtype=tf.float32)
- dense_mask = tf.keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
- dense_type_ids = tf.keras.Input(
+ dense_mask = tf_keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
+ dense_type_ids = tf_keras.Input(
shape=(dense_sequence_length,), dtype=tf.int32)
test_network.build(
diff --git a/official/nlp/modeling/networks/bert_encoder.py b/official/nlp/modeling/networks/bert_encoder.py
index ed2f3dd6c8b..ae8f07018bf 100644
--- a/official/nlp/modeling/networks/bert_encoder.py
+++ b/official/nlp/modeling/networks/bert_encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,18 +17,19 @@
from typing import Any, Callable, Optional, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
-
-_Initializer = Union[str, tf.keras.initializers.Initializer]
+_Initializer = Union[str, tf_keras.initializers.Initializer]
_Activation = Union[str, Callable[..., Any]]
-_approx_gelu = lambda x: tf.keras.activations.gelu(x, approximate=True)
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
-class BertEncoderV2(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BertEncoderV2(tf_keras.layers.Layer):
"""Bi-directional Transformer-based encoder network.
This network implements a bi-directional Transformer-based encoder as
@@ -48,8 +49,7 @@ class BertEncoderV2(tf.keras.layers.Layer):
num_attention_heads: The number of attention heads for each transformer. The
hidden size must be divisible by the number of attention heads.
max_sequence_length: The maximum sequence length that this encoder can
- consume. If None, max_sequence_length uses the value from sequence length.
- This determines the variable shape for positional embeddings.
+ consume. This determines the variable shape for positional embeddings.
type_vocab_size: The number of types that the 'type_ids' input can take.
inner_dim: The output dimension of the first Dense layer in a two-layer
feedforward network for each transformer.
@@ -75,6 +75,12 @@ class BertEncoderV2(tf.keras.layers.Layer):
layers. If set False, output of attention and intermediate dense layers is
normalized.
with_dense_inputs: Whether to accept dense embeddings as the input.
+ return_attention_scores: Whether to add an additional output containing the
+ attention scores of all transformer layers. This will be a list of length
+ `num_layers`, and each element will be in the shape [batch_size,
+ num_attention_heads, seq_dim, seq_dim].
+ return_word_embeddings: If true, also return the input word embedding
+ sequence in the bert inference output.
"""
def __init__(
@@ -89,13 +95,15 @@ def __init__(
inner_activation: _Activation = _approx_gelu,
output_dropout: float = 0.1,
attention_dropout: float = 0.1,
- initializer: _Initializer = tf.keras.initializers.TruncatedNormal(
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
stddev=0.02),
output_range: Optional[int] = None,
embedding_width: Optional[int] = None,
- embedding_layer: Optional[tf.keras.layers.Layer] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
norm_first: bool = False,
with_dense_inputs: bool = False,
+ return_attention_scores: bool = False,
+ return_word_embeddings: bool = False,
**kwargs):
# Pops kwargs that are used in V1 implementation.
if 'dict_outputs' in kwargs:
@@ -112,8 +120,10 @@ def __init__(
attention_dropout = kwargs.pop('attention_dropout_rate')
super().__init__(**kwargs)
- activation = tf.keras.activations.get(inner_activation)
- initializer = tf.keras.initializers.get(initializer)
+ self._output_range = output_range
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
if embedding_width is None:
embedding_width = hidden_size
@@ -122,43 +132,44 @@ def __init__(
self._embedding_layer = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
name='word_embeddings')
else:
self._embedding_layer = embedding_layer
self._position_embedding_layer = layers.PositionEmbedding(
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
max_length=max_sequence_length,
name='position_embedding')
self._type_embedding_layer = layers.OnDeviceEmbedding(
vocab_size=type_vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_one_hot=True,
name='type_embeddings')
- self._embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
- self._embedding_dropout = tf.keras.layers.Dropout(
+ self._embedding_dropout = tf_keras.layers.Dropout(
rate=output_dropout, name='embedding_dropout')
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
self._embedding_projection = None
if embedding_width != hidden_size:
- self._embedding_projection = tf.keras.layers.experimental.EinsumDense(
+ self._embedding_projection = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='embedding_projection')
self._transformer_layers = []
self._attention_mask_layer = layers.SelfAttentionMask(
name='self_attention_mask')
+ self._num_layers = num_layers
for i in range(num_layers):
layer = layers.TransformerEncoderBlock(
num_attention_heads=num_attention_heads,
@@ -167,15 +178,15 @@ def __init__(
output_dropout=output_dropout,
attention_dropout=attention_dropout,
norm_first=norm_first,
- output_range=output_range if i == num_layers - 1 else None,
- kernel_initializer=initializer,
+ return_attention_scores=return_attention_scores,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='transformer/layer_%d' % i)
self._transformer_layers.append(layer)
- self._pooler_layer = tf.keras.layers.Dense(
+ self._pooler_layer = tf_keras.layers.Dense(
units=hidden_size,
activation='tanh',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='pooler_transform')
self._config = {
@@ -186,31 +197,37 @@ def __init__(
'max_sequence_length': max_sequence_length,
'type_vocab_size': type_vocab_size,
'inner_dim': inner_dim,
- 'inner_activation': tf.keras.activations.serialize(activation),
+ 'inner_activation': tf_utils.serialize_activation(
+ activation, use_legacy_format=True
+ ),
'output_dropout': output_dropout,
'attention_dropout': attention_dropout,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_utils.serialize_initializer(
+ initializer, use_legacy_format=True
+ ),
'output_range': output_range,
'embedding_width': embedding_width,
'embedding_layer': embedding_layer,
'norm_first': norm_first,
'with_dense_inputs': with_dense_inputs,
+ 'return_attention_scores': return_attention_scores,
+ 'return_word_embeddings': return_word_embeddings,
}
if with_dense_inputs:
self.inputs = dict(
- input_word_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- dense_inputs=tf.keras.Input(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_inputs=tf_keras.Input(
shape=(None, embedding_width), dtype=tf.float32),
- dense_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- dense_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
+ dense_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
)
else:
self.inputs = dict(
- input_word_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32))
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32))
def call(self, inputs):
word_embeddings = None
@@ -230,16 +247,10 @@ def call(self, inputs):
word_embeddings = self._embedding_layer(word_ids)
if dense_inputs is not None:
- # Concat the dense embeddings at sequence end.
- word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
- type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
mask = tf.concat([mask, dense_mask], axis=1)
- # absolute position embeddings.
- position_embeddings = self._position_embedding_layer(word_embeddings)
- type_embeddings = self._type_embedding_layer(type_ids)
-
- embeddings = word_embeddings + position_embeddings + type_embeddings
+ embeddings = self._get_embeddings(word_ids, type_ids, word_embeddings, # pyrefly: ignore[bad-argument-type]
+ dense_inputs, dense_type_ids)
embeddings = self._embedding_norm_layer(embeddings)
embeddings = self._embedding_dropout(embeddings)
@@ -249,19 +260,33 @@ def call(self, inputs):
attention_mask = self._attention_mask_layer(embeddings, mask)
encoder_outputs = []
+ attention_outputs = []
x = embeddings
- for layer in self._transformer_layers:
- x = layer([x, attention_mask])
+ for i, layer in enumerate(self._transformer_layers):
+ transformer_output_range = None
+ if i == self._num_layers - 1:
+ transformer_output_range = self._output_range
+ x = layer([x, attention_mask], output_range=transformer_output_range)
+ if self._config['return_attention_scores']:
+ x, attention_scores = x
+ attention_outputs.append(attention_scores)
encoder_outputs.append(x)
last_encoder_output = encoder_outputs[-1]
first_token_tensor = last_encoder_output[:, 0, :]
pooled_output = self._pooler_layer(first_token_tensor)
- return dict(
+ output = dict(
sequence_output=encoder_outputs[-1],
pooled_output=pooled_output,
encoder_outputs=encoder_outputs)
+ if self._config['return_attention_scores']:
+ output['attention_scores'] = attention_outputs
+
+ if self._config['return_word_embeddings']:
+ output['word_embeddings'] = embeddings
+
+ return output
def get_embedding_table(self):
return self._embedding_layer.embeddings
@@ -295,9 +320,27 @@ def from_config(cls, config, custom_objects=None):
return cls(**config)
+ def _get_embeddings(self, word_ids: tf.Tensor, type_ids: tf.Tensor,
+ word_embeddings: Optional[tf.Tensor],
+ dense_inputs: Optional[tf.Tensor],
+ dense_type_ids: Optional[tf.Tensor]) -> tf.Tensor:
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ # Concat the dense embeddings at sequence end.
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+
+ type_embeddings = self._type_embedding_layer(type_ids)
+
+ # absolute position embeddings.
+ position_embeddings = self._position_embedding_layer(word_embeddings)
+ return word_embeddings + position_embeddings + type_embeddings
+
-@tf.keras.utils.register_keras_serializable(package='Text')
-class BertEncoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BertEncoder(tf_keras.Model):
"""Bi-directional Transformer-based encoder network.
This network implements a bi-directional Transformer-based encoder as
@@ -324,13 +367,13 @@ class BertEncoder(tf.keras.Model):
This determines the variable shape for positional embeddings.
type_vocab_size: The number of types that the 'type_ids' input can take.
inner_dim: The output dimension of the first Dense layer in a two-layer
- feedforward network for each transformer.
+ feedforward network for each transformer.
inner_activation: The activation for the first Dense layer in a two-layer
- feedforward network for each transformer.
+ feedforward network for each transformer.
output_dropout: Dropout probability for the post-attention and output
- dropout.
- attention_dropout: The dropout rate to use for the attention layers
- within the transformer layers.
+ dropout.
+ attention_dropout: The dropout rate to use for the attention layers within
+ the transformer layers.
initializer: The initialzer to use for all weights in this encoder.
output_range: The sequence output range, [0, output_range), by slicing the
target sequence of the last transformer layer. `None` means the entire
@@ -341,16 +384,22 @@ class BertEncoder(tf.keras.Model):
matrices in the shape of ['vocab_size', 'embedding_width'] and
['embedding_width', 'hidden_size'] ('embedding_width' is usually much
smaller than 'hidden_size').
- embedding_layer: An optional Layer instance which will be called to
- generate embeddings for the input word IDs.
- norm_first: Whether to normalize inputs to attention and intermediate
- dense layers. If set False, output of attention and intermediate dense
- layers is normalized.
+ embedding_layer: An optional Layer instance which will be called to generate
+ embeddings for the input word IDs.
+ norm_first: Whether to normalize inputs to attention and intermediate dense
+ layers. If set False, output of attention and intermediate dense layers is
+ normalized.
dict_outputs: Whether to use a dictionary as the model outputs.
return_all_encoder_outputs: Whether to output sequence embedding outputs of
all encoder transformer layers. Note: when the following `dict_outputs`
argument is True, all encoder outputs are always returned in the dict,
keyed by `encoder_outputs`.
+ return_attention_scores: Whether to add an additional output containing the
+ attention scores of all transformer layers. This will be a list of length
+ `num_layers`, and each element will be in the shape [batch_size,
+ num_attention_heads, seq_dim, seq_dim].
+ return_word_embeddings: If true, also return the input word embedding
+ sequence in the bert inference output.
"""
def __init__(
@@ -362,16 +411,18 @@ def __init__(
max_sequence_length=512,
type_vocab_size=16,
inner_dim=3072,
- inner_activation=lambda x: tf.keras.activations.gelu(x, approximate=True),
+ inner_activation=lambda x: tf_keras.activations.gelu(x, approximate=True),
output_dropout=0.1,
attention_dropout=0.1,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
output_range=None,
embedding_width=None,
embedding_layer=None,
norm_first=False,
dict_outputs=False,
return_all_encoder_outputs=False,
+ return_attention_scores: bool = False,
+ return_word_embeddings: bool = False,
**kwargs):
if 'sequence_length' in kwargs:
kwargs.pop('sequence_length')
@@ -392,14 +443,14 @@ def __init__(
if 'attention_dropout_rate' in kwargs:
attention_dropout = kwargs.pop('attention_dropout_rate')
- activation = tf.keras.activations.get(inner_activation)
- initializer = tf.keras.initializers.get(initializer)
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
- word_ids = tf.keras.layers.Input(
+ word_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- mask = tf.keras.layers.Input(
+ mask = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_mask')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
if embedding_width is None:
@@ -409,7 +460,7 @@ def __init__(
embedding_layer_inst = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
name='word_embeddings')
else:
embedding_layer_inst = embedding_layer
@@ -417,35 +468,35 @@ def __init__(
# Always uses dynamic slicing for simplicity.
position_embedding_layer = layers.PositionEmbedding(
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
max_length=max_sequence_length,
name='position_embedding')
position_embeddings = position_embedding_layer(word_embeddings)
type_embedding_layer = layers.OnDeviceEmbedding(
vocab_size=type_vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_one_hot=True,
name='type_embeddings')
type_embeddings = type_embedding_layer(type_ids)
- embeddings = tf.keras.layers.Add()(
+ embeddings = tf_keras.layers.Add()(
[word_embeddings, position_embeddings, type_embeddings])
- embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
embeddings = embedding_norm_layer(embeddings)
- embeddings = (tf.keras.layers.Dropout(rate=output_dropout)(embeddings))
+ embeddings = (tf_keras.layers.Dropout(rate=output_dropout)(embeddings))
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
if embedding_width != hidden_size:
- embedding_projection = tf.keras.layers.experimental.EinsumDense(
+ embedding_projection = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='embedding_projection')
embeddings = embedding_projection(embeddings)
else:
@@ -455,11 +506,11 @@ def __init__(
data = embeddings
attention_mask = layers.SelfAttentionMask()(data, mask)
encoder_outputs = []
+ attention_outputs = []
for i in range(num_layers):
- if i == num_layers - 1 and output_range is not None:
+ transformer_output_range = None
+ if i == num_layers - 1:
transformer_output_range = output_range
- else:
- transformer_output_range = None
layer = layers.TransformerEncoderBlock(
num_attention_heads=num_attention_heads,
inner_dim=inner_dim,
@@ -467,11 +518,15 @@ def __init__(
output_dropout=output_dropout,
attention_dropout=attention_dropout,
norm_first=norm_first,
- output_range=transformer_output_range,
- kernel_initializer=initializer,
+ return_attention_scores=return_attention_scores,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='transformer/layer_%d' % i)
transformer_layers.append(layer)
- data = layer([data, attention_mask])
+ data = layer([data, attention_mask],
+ output_range=transformer_output_range)
+ if return_attention_scores:
+ data, attention_scores = data
+ attention_outputs.append(attention_scores)
encoder_outputs.append(data)
last_encoder_output = encoder_outputs[-1]
@@ -479,10 +534,10 @@ def __init__(
# like this will create a SliceOpLambda layer. This is better than a Lambda
# layer with Python code, because that is fundamentally less portable.
first_token_tensor = last_encoder_output[:, 0, :]
- pooler_layer = tf.keras.layers.Dense(
+ pooler_layer = tf_keras.layers.Dense(
units=hidden_size,
activation='tanh',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='pooler_transform')
cls_output = pooler_layer(first_token_tensor)
@@ -491,6 +546,11 @@ def __init__(
pooled_output=cls_output,
encoder_outputs=encoder_outputs,
)
+ if return_attention_scores:
+ outputs['attention_scores'] = attention_outputs
+
+ if return_word_embeddings:
+ outputs['word_embeddings'] = embeddings
if dict_outputs:
super().__init__(
@@ -503,6 +563,8 @@ def __init__(
else:
sequence_output = outputs['sequence_output']
outputs = [sequence_output, cls_output]
+ if return_attention_scores:
+ outputs.append(attention_outputs)
super().__init__( # pylint: disable=bad-super-call
inputs=[word_ids, mask, type_ids],
outputs=outputs,
@@ -525,15 +587,21 @@ def __init__(
'max_sequence_length': max_sequence_length,
'type_vocab_size': type_vocab_size,
'inner_dim': inner_dim,
- 'inner_activation': tf.keras.activations.serialize(activation),
+ 'inner_activation': tf_utils.serialize_activation(
+ activation, use_legacy_format=True
+ ),
'output_dropout': output_dropout,
'attention_dropout': attention_dropout,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_utils.serialize_initializer(
+ initializer, use_legacy_format=True
+ ),
'output_range': output_range,
'embedding_width': embedding_width,
'embedding_layer': embedding_layer,
'norm_first': norm_first,
'dict_outputs': dict_outputs,
+ 'return_attention_scores': return_attention_scores,
+ 'return_word_embeddings': return_word_embeddings,
}
# pylint: disable=protected-access
self._setattr_tracking = False
diff --git a/official/nlp/modeling/networks/bert_encoder_test.py b/official/nlp/modeling/networks/bert_encoder_test.py
index acf773af7e2..ad63da36eec 100644
--- a/official/nlp/modeling/networks/bert_encoder_test.py
+++ b/official/nlp/modeling/networks/bert_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,23 +14,18 @@
"""Tests for transformer-based bert encoder network."""
-# Import libraries
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.networks import bert_encoder
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class BertEncoderTest(keras_parameterized.TestCase):
+class BertEncoderTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(BertEncoderTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
@parameterized.named_parameters(
("encoder_v2", bert_encoder.BertEncoderV2),
@@ -51,9 +46,9 @@ def test_dict_outputs_network_creation(self, encoder_cls):
num_layers=3,
**kwargs)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
@@ -61,7 +56,7 @@ def test_dict_outputs_network_creation(self, encoder_cls):
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, 3)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
expected_data_shape = [None, sequence_length, hidden_size]
expected_pooled_shape = [None, hidden_size]
@@ -87,9 +82,9 @@ def test_dict_outputs_all_encoder_outputs_network_creation(self, encoder_cls):
num_layers=3,
dict_outputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
all_encoder_outputs = dict_outputs["encoder_outputs"]
@@ -106,6 +101,74 @@ def test_dict_outputs_all_encoder_outputs_network_creation(self, encoder_cls):
self.assertAllEqual(tf.float32, all_encoder_outputs[-1].dtype)
self.assertAllEqual(tf.float32, pooled.dtype)
+ @parameterized.named_parameters(
+ ("encoder_v2", bert_encoder.BertEncoderV2),
+ ("encoder_v1", bert_encoder.BertEncoder),
+ )
+ def test_dict_outputs_network_creation_return_attention_scores(
+ self, encoder_cls):
+ hidden_size = 32
+ sequence_length = 21
+ num_attention_heads = 5
+ num_layers = 3
+ # Create a small BertEncoder for testing.
+ test_network = encoder_cls(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=num_attention_heads,
+ num_layers=num_layers,
+ return_attention_scores=True,
+ dict_outputs=True)
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = test_network(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ all_attention_outputs = dict_outputs["attention_scores"]
+
+ expected_data_shape = [
+ None, num_attention_heads, sequence_length, sequence_length
+ ]
+ self.assertLen(all_attention_outputs, num_layers)
+ for data in all_attention_outputs:
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, all_attention_outputs[-1].dtype)
+
+ @parameterized.named_parameters(
+ ("encoder_v2", bert_encoder.BertEncoderV2),
+ ("encoder_v1", bert_encoder.BertEncoder),
+ )
+ def test_dict_outputs_network_creation_return_word_embeddings(
+ self, encoder_cls):
+ hidden_size = 32
+ sequence_length = 21
+ num_attention_heads = 5
+ num_layers = 3
+ # Create a small BertEncoder for testing.
+ test_network = encoder_cls(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=num_attention_heads,
+ num_layers=num_layers,
+ return_word_embeddings=True,
+ dict_outputs=True)
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = test_network(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ word_embeddings = dict_outputs["word_embeddings"]
+
+ expected_data_shape = [None, sequence_length, hidden_size]
+ self.assertAllEqual(expected_data_shape, word_embeddings.shape)
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, word_embeddings[-1].dtype)
+
@parameterized.named_parameters(
("encoder_v2", bert_encoder.BertEncoderV2),
("encoder_v1", bert_encoder.BertEncoder),
@@ -113,7 +176,7 @@ def test_dict_outputs_all_encoder_outputs_network_creation(self, encoder_cls):
def test_dict_outputs_network_creation_with_float16_dtype(self, encoder_cls):
hidden_size = 32
sequence_length = 21
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a small BertEncoder for testing.
test_network = encoder_cls(
vocab_size=100,
@@ -122,9 +185,9 @@ def test_dict_outputs_network_creation_with_float16_dtype(self, encoder_cls):
num_layers=3,
dict_outputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
@@ -162,16 +225,16 @@ def test_dict_outputs_network_invocation(
output_range=output_range,
dict_outputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -198,7 +261,7 @@ def test_dict_outputs_network_invocation(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[1], sequence_length)
@@ -216,7 +279,7 @@ def test_dict_outputs_network_invocation(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[-1], hidden_size)
self.assertTrue(hasattr(test_network, "_embedding_projection"))
@@ -231,9 +294,9 @@ def test_embeddings_as_inputs(self):
num_attention_heads=2,
num_layers=3)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
test_network.build(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
embeddings = test_network.get_embedding_layer()(word_ids)
@@ -275,15 +338,41 @@ def test_serialize_deserialize(self):
embedding_width=16,
embedding_layer=None,
norm_first=False)
- network = bert_encoder.BertEncoder(**kwargs)
- # Validate that the config can be forced to JSON.
- _ = network.to_json()
+ with self.subTest("BertEncoder"):
+ network = bert_encoder.BertEncoder(**kwargs)
+
+ # Validate that the config can be forced to JSON.
+ _ = network.to_json()
- # Tests model saving/loading.
- model_path = self.get_temp_dir() + "/model"
- network.save(model_path)
- _ = tf.keras.models.load_model(model_path)
+ # Tests model saving/loading with SavedModel.
+ model_path = self.get_temp_dir() + "/model"
+ network.save(model_path)
+ _ = tf_keras.models.load_model(model_path)
+
+ # Test model saving/loading with Keras V3.
+ keras_path = self.get_temp_dir() + "/model.keras"
+ network.save(keras_path)
+ _ = tf_keras.models.load_model(keras_path)
+
+ with self.subTest("BertEncoderV2"):
+ new_net = bert_encoder.BertEncoderV2(**kwargs)
+ inputs = new_net.inputs
+ outputs = new_net(inputs)
+ network_v2 = tf_keras.Model(inputs=inputs, outputs=outputs)
+
+ # Validate that the config can be forced to JSON.
+ _ = network_v2.to_json()
+
+ # Tests model saving/loading with SavedModel.
+ model_path = self.get_temp_dir() + "/v2_model"
+ network_v2.save(model_path)
+ _ = tf_keras.models.load_model(model_path)
+
+ # Test model saving/loading with Keras V3.
+ keras_path = self.get_temp_dir() + "/v2_model.keras"
+ network_v2.save(keras_path)
+ _ = tf_keras.models.load_model(keras_path)
def test_network_creation(self):
hidden_size = 32
@@ -295,14 +384,14 @@ def test_network_creation(self):
num_attention_heads=2,
num_layers=3)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask, type_ids])
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, 3)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
expected_data_shape = [None, sequence_length, hidden_size]
expected_pooled_shape = [None, hidden_size]
@@ -353,9 +442,9 @@ def test_all_encoder_outputs_network_creation(self):
num_layers=3,
return_all_encoder_outputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
all_encoder_outputs, pooled = test_network([word_ids, mask, type_ids])
expected_data_shape = [None, sequence_length, hidden_size]
@@ -369,10 +458,38 @@ def test_all_encoder_outputs_network_creation(self):
self.assertAllEqual(tf.float32, all_encoder_outputs[-1].dtype)
self.assertAllEqual(tf.float32, pooled.dtype)
+ def test_attention_scores_output_network_creation(self):
+ hidden_size = 32
+ sequence_length = 21
+ num_attention_heads = 5
+ num_layers = 3
+ # Create a small BertEncoder for testing.
+ test_network = bert_encoder.BertEncoder(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=num_attention_heads,
+ num_layers=num_layers,
+ return_attention_scores=True)
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ _, _, all_attention_outputs = test_network([word_ids, mask, type_ids])
+
+ expected_data_shape = [
+ None, num_attention_heads, sequence_length, sequence_length
+ ]
+ self.assertLen(all_attention_outputs, num_layers)
+ for data in all_attention_outputs:
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, all_attention_outputs[-1].dtype)
+
def test_network_creation_with_float16_dtype(self):
hidden_size = 32
sequence_length = 21
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a small BertEncoder for testing.
test_network = bert_encoder.BertEncoder(
vocab_size=100,
@@ -380,9 +497,9 @@ def test_network_creation_with_float16_dtype(self):
num_attention_heads=2,
num_layers=3)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask, type_ids])
expected_data_shape = [None, sequence_length, hidden_size]
@@ -413,13 +530,13 @@ def test_network_invocation(self, output_range, out_seq_len):
type_vocab_size=num_types,
output_range=output_range)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask, type_ids])
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -442,7 +559,7 @@ def test_network_invocation(self, output_range, out_seq_len):
num_layers=3,
type_vocab_size=num_types)
data, pooled = test_network([word_ids, mask, type_ids])
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[1], sequence_length)
@@ -456,7 +573,7 @@ def test_network_invocation(self, output_range, out_seq_len):
type_vocab_size=num_types,
embedding_width=16)
data, pooled = test_network([word_ids, mask, type_ids])
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[-1], hidden_size)
self.assertTrue(hasattr(test_network, "_embedding_projection"))
@@ -466,7 +583,7 @@ class BertEncoderV2CompatibilityTest(tf.test.TestCase):
def tearDown(self):
super().tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
def test_weights_forward_compatible(self):
batch_size = 3
@@ -481,8 +598,7 @@ def test_weights_forward_compatible(self):
hidden_size=hidden_size,
num_attention_heads=2,
num_layers=3,
- type_vocab_size=num_types,
- output_range=None)
+ type_vocab_size=num_types)
word_id_data = np.random.randint(
vocab_size, size=(batch_size, sequence_length))
@@ -541,8 +657,7 @@ def test_checkpoint_forward_compatible(self):
hidden_size=hidden_size,
num_attention_heads=2,
num_layers=3,
- type_vocab_size=num_types,
- output_range=None)
+ type_vocab_size=num_types)
word_id_data = np.random.randint(
vocab_size, size=(batch_size, sequence_length))
@@ -601,7 +716,7 @@ def test_keras_model_checkpoint_forward_compatible(self):
old_net = bert_encoder.BertEncoder(**kwargs)
inputs = old_net.inputs
outputs = old_net(inputs)
- old_model = tf.keras.Model(inputs=inputs, outputs=outputs)
+ old_model = tf_keras.Model(inputs=inputs, outputs=outputs)
old_model_outputs = old_model(data)
ckpt = tf.train.Checkpoint(net=old_model)
path = ckpt.save(self.get_temp_dir())
@@ -609,7 +724,7 @@ def test_keras_model_checkpoint_forward_compatible(self):
new_net = bert_encoder.BertEncoderV2(**kwargs)
inputs = new_net.inputs
outputs = new_net(inputs)
- new_model = tf.keras.Model(inputs=inputs, outputs=outputs)
+ new_model = tf_keras.Model(inputs=inputs, outputs=outputs)
new_ckpt = tf.train.Checkpoint(net=new_model)
status = new_ckpt.restore(path)
diff --git a/official/nlp/modeling/networks/classification.py b/official/nlp/modeling/networks/classification.py
index 67fa0dd2608..c4d560f0f62 100644
--- a/official/nlp/modeling/networks/classification.py
+++ b/official/nlp/modeling/networks/classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,12 +15,12 @@
"""Classification and regression network."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.util import deprecation
-@tf.keras.utils.register_keras_serializable(package='Text')
-class Classification(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class Classification(tf_keras.Model):
"""Classification network head for BERT modeling.
This network implements a simple classifier head based on a dense layer. If
@@ -49,10 +49,10 @@ def __init__(self,
output='logits',
**kwargs):
- cls_output = tf.keras.layers.Input(
+ cls_output = tf_keras.layers.Input(
shape=(input_width,), name='cls_output', dtype=tf.float32)
- logits = tf.keras.layers.Dense(
+ logits = tf_keras.layers.Dense(
num_classes,
activation=None,
kernel_initializer=initializer,
@@ -62,11 +62,11 @@ def __init__(self,
if output == 'logits':
output_tensors = logits
elif output == 'predictions':
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == 'mixed_bfloat16':
# b/158514794: bf16 is not stable with post-softmax cross-entropy.
policy = tf.float32
- output_tensors = tf.keras.layers.Activation(
+ output_tensors = tf_keras.layers.Activation(
tf.nn.log_softmax, dtype=policy)(
logits)
else:
@@ -74,7 +74,7 @@ def __init__(self,
('Unknown `output` value "%s". `output` can be either "logits" or '
'"predictions"') % output)
- super(Classification, self).__init__(
+ super().__init__(
inputs=[cls_output], outputs=output_tensors, **kwargs)
# b/164516224
@@ -95,7 +95,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self.logits = logits
diff --git a/official/nlp/modeling/networks/classification_test.py b/official/nlp/modeling/networks/classification_test.py
index 3f055181327..ba7512f511c 100644
--- a/official/nlp/modeling/networks/classification_test.py
+++ b/official/nlp/modeling/networks/classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,22 +14,14 @@
"""Tests for classification network."""
-from __future__ import absolute_import
-from __future__ import division
-from __future__ import print_function
-
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.networks import classification
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class ClassificationTest(keras_parameterized.TestCase):
+class ClassificationTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters(1, 10)
def test_network_creation(self, num_classes):
@@ -38,7 +30,7 @@ def test_network_creation(self, num_classes):
test_object = classification.Classification(
input_width=input_width, num_classes=num_classes)
# Create a 2-dimensional input (the first dimension is implicit).
- cls_data = tf.keras.Input(shape=(input_width,), dtype=tf.float32)
+ cls_data = tf_keras.Input(shape=(input_width,), dtype=tf.float32)
output = test_object(cls_data)
# Validate that the outputs are of the expected shape.
@@ -52,11 +44,11 @@ def test_network_invocation(self, num_classes):
test_object = classification.Classification(
input_width=input_width, num_classes=num_classes, output='predictions')
# Create a 2-dimensional input (the first dimension is implicit).
- cls_data = tf.keras.Input(shape=(input_width,), dtype=tf.float32)
+ cls_data = tf_keras.Input(shape=(input_width,), dtype=tf.float32)
output = test_object(cls_data)
# Invoke the network as part of a Model.
- model = tf.keras.Model(cls_data, output)
+ model = tf_keras.Model(cls_data, output)
input_data = 10 * np.random.random_sample((3, input_width))
_ = model.predict(input_data)
@@ -68,10 +60,10 @@ def test_network_invocation_with_internal_logits(self):
input_width=input_width, num_classes=num_classes, output='predictions')
# Create a 2-dimensional input (the first dimension is implicit).
- cls_data = tf.keras.Input(shape=(input_width,), dtype=tf.float32)
+ cls_data = tf_keras.Input(shape=(input_width,), dtype=tf.float32)
output = test_object(cls_data)
- model = tf.keras.Model(cls_data, output)
- logits_model = tf.keras.Model(test_object.inputs, test_object.logits)
+ model = tf_keras.Model(cls_data, output)
+ logits_model = tf_keras.Model(test_object.inputs, test_object.logits)
batch_size = 3
input_data = 10 * np.random.random_sample((batch_size, input_width))
@@ -84,9 +76,9 @@ def test_network_invocation_with_internal_logits(self):
self.assertEqual(expected_output_shape, logits.shape)
# Ensure that the logits, when softmaxed, create the outputs.
- input_tensor = tf.keras.Input(expected_output_shape[1:])
- output_tensor = tf.keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
- softmax_model = tf.keras.Model(input_tensor, output_tensor)
+ input_tensor = tf_keras.Input(expected_output_shape[1:])
+ output_tensor = tf_keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
+ softmax_model = tf_keras.Model(input_tensor, output_tensor)
calculated_softmax = softmax_model.predict(logits)
self.assertAllClose(outputs, calculated_softmax)
@@ -100,10 +92,10 @@ def test_network_invocation_with_internal_and_external_logits(
input_width=input_width, num_classes=num_classes, output='logits')
# Create a 2-dimensional input (the first dimension is implicit).
- cls_data = tf.keras.Input(shape=(input_width,), dtype=tf.float32)
+ cls_data = tf_keras.Input(shape=(input_width,), dtype=tf.float32)
output = test_object(cls_data)
- model = tf.keras.Model(cls_data, output)
- logits_model = tf.keras.Model(test_object.inputs, test_object.logits)
+ model = tf_keras.Model(cls_data, output)
+ logits_model = tf_keras.Model(test_object.inputs, test_object.logits)
batch_size = 3
input_data = 10 * np.random.random_sample((batch_size, input_width))
@@ -128,12 +120,12 @@ def test_network_invocation_with_logit_output(self):
logit_object.set_weights(test_object.get_weights())
# Create a 2-dimensional input (the first dimension is implicit).
- cls_data = tf.keras.Input(shape=(input_width,), dtype=tf.float32)
+ cls_data = tf_keras.Input(shape=(input_width,), dtype=tf.float32)
output = test_object(cls_data)
logit_output = logit_object(cls_data)
- model = tf.keras.Model(cls_data, output)
- logits_model = tf.keras.Model(cls_data, logit_output)
+ model = tf_keras.Model(cls_data, output)
+ logits_model = tf_keras.Model(cls_data, logit_output)
batch_size = 3
input_data = 10 * np.random.random_sample((batch_size, input_width))
@@ -146,9 +138,9 @@ def test_network_invocation_with_logit_output(self):
self.assertEqual(expected_output_shape, logits.shape)
# Ensure that the logits, when softmaxed, create the outputs.
- input_tensor = tf.keras.Input(expected_output_shape[1:])
- output_tensor = tf.keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
- softmax_model = tf.keras.Model(input_tensor, output_tensor)
+ input_tensor = tf_keras.Input(expected_output_shape[1:])
+ output_tensor = tf_keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
+ softmax_model = tf_keras.Model(input_tensor, output_tensor)
calculated_softmax = softmax_model.predict(logits)
self.assertAllClose(outputs, calculated_softmax)
diff --git a/official/nlp/modeling/networks/encoder_scaffold.py b/official/nlp/modeling/networks/encoder_scaffold.py
index 4f546f93214..72f125f431a 100644
--- a/official/nlp/modeling/networks/encoder_scaffold.py
+++ b/official/nlp/modeling/networks/encoder_scaffold.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,14 +19,15 @@
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
-@tf.keras.utils.register_keras_serializable(package='Text')
+@tf_keras.utils.register_keras_serializable(package='Text')
@gin.configurable
-class EncoderScaffold(tf.keras.Model):
+class EncoderScaffold(tf_keras.Model):
"""Bi-directional Transformer-based encoder network scaffold.
This network allows users to flexibly implement an encoder similar to the one
@@ -109,7 +110,7 @@ class or instance defines the inputs to this encoder and outputs (1)
def __init__(self,
pooled_output_dim,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
embedding_cls=None,
embedding_cfg=None,
@@ -133,50 +134,50 @@ def __init__(self,
**embedding_cfg) if embedding_cfg else embedding_cls()
else:
embedding_network = embedding_cls
- inputs = embedding_network.inputs
- embeddings, attention_mask = embedding_network(inputs)
+ inputs = embedding_network.inputs # pyrefly: ignore[missing-attribute]
+ embeddings, attention_mask = embedding_network(inputs) # pyrefly: ignore[not-callable]
embedding_layer = None
position_embedding_layer = None
type_embedding_layer = None
embedding_norm_layer = None
else:
embedding_network = None
- seq_length = embedding_cfg.get('seq_length', None)
- word_ids = tf.keras.layers.Input(
+ seq_length = embedding_cfg.get('seq_length', None) # pyrefly: ignore[missing-attribute]
+ word_ids = tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_word_ids')
- mask = tf.keras.layers.Input(
+ mask = tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_mask')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(seq_length,), dtype=tf.int32, name='input_type_ids')
inputs = [word_ids, mask, type_ids]
embedding_layer = layers.OnDeviceEmbedding(
- vocab_size=embedding_cfg['vocab_size'],
- embedding_width=embedding_cfg['hidden_size'],
- initializer=embedding_cfg['initializer'],
+ vocab_size=embedding_cfg['vocab_size'], # pyrefly: ignore[unsupported-operation]
+ embedding_width=embedding_cfg['hidden_size'], # pyrefly: ignore[unsupported-operation]
+ initializer=tf_utils.clone_initializer(embedding_cfg['initializer']), # pyrefly: ignore[unsupported-operation]
name='word_embeddings')
word_embeddings = embedding_layer(word_ids)
# Always uses dynamic slicing for simplicity.
position_embedding_layer = layers.PositionEmbedding(
- initializer=embedding_cfg['initializer'],
- max_length=embedding_cfg['max_seq_length'],
+ initializer=tf_utils.clone_initializer(embedding_cfg['initializer']), # pyrefly: ignore[unsupported-operation]
+ max_length=embedding_cfg['max_seq_length'], # pyrefly: ignore[unsupported-operation]
name='position_embedding')
position_embeddings = position_embedding_layer(word_embeddings)
type_embedding_layer = layers.OnDeviceEmbedding(
- vocab_size=embedding_cfg['type_vocab_size'],
- embedding_width=embedding_cfg['hidden_size'],
- initializer=embedding_cfg['initializer'],
+ vocab_size=embedding_cfg['type_vocab_size'], # pyrefly: ignore[unsupported-operation]
+ embedding_width=embedding_cfg['hidden_size'], # pyrefly: ignore[unsupported-operation]
+ initializer=tf_utils.clone_initializer(embedding_cfg['initializer']), # pyrefly: ignore[unsupported-operation]
use_one_hot=True,
name='type_embeddings')
type_embeddings = type_embedding_layer(type_ids)
- embeddings = tf.keras.layers.Add()(
+ embeddings = tf_keras.layers.Add()(
[word_embeddings, position_embeddings, type_embeddings])
- embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm',
axis=-1,
epsilon=1e-12,
@@ -184,15 +185,15 @@ def __init__(self,
embeddings = embedding_norm_layer(embeddings)
embeddings = (
- tf.keras.layers.Dropout(
- rate=embedding_cfg['dropout_rate'])(embeddings))
+ tf_keras.layers.Dropout(
+ rate=embedding_cfg['dropout_rate'])(embeddings)) # pyrefly: ignore[unsupported-operation]
mask_cfg = {} if mask_cfg is None else mask_cfg
if inspect.isclass(mask_cls):
mask_layer = mask_cls(**mask_cfg)
else:
mask_layer = mask_cls
- attention_mask = mask_layer(embeddings, mask)
+ attention_mask = mask_layer(embeddings, mask) # pyrefly: ignore[not-callable]
data = embeddings
@@ -224,15 +225,15 @@ def __init__(self,
else:
layer = cur_hidden_cls
if recursive:
- data, recursive_states = layer([data, attention_mask, recursive_states])
+ data, recursive_states = layer([data, attention_mask, recursive_states]) # pyrefly: ignore[not-callable]
else:
- data = layer([data, attention_mask])
+ data = layer([data, attention_mask]) # pyrefly: ignore[not-callable]
layer_output_data.append(data)
hidden_layers.append(layer)
if layer_norm_before_pooling:
# Normalize the final output.
- output_layer_norm = tf.keras.layers.LayerNormalization(
+ output_layer_norm = tf_keras.layers.LayerNormalization(
name='final_layer_norm',
axis=-1,
epsilon=1e-12)
@@ -243,7 +244,9 @@ def __init__(self,
# like this will create a SliceOpLambda layer. This is better than a Lambda
# layer with Python code, because that is fundamentally less portable.
first_token_tensor = last_layer_output[:, 0, :]
- pooler_layer = tf.keras.layers.Dense(
+ pooler_layer_initializer = tf_keras.initializers.get(
+ pooler_layer_initializer)
+ pooler_layer = tf_keras.layers.Dense(
units=pooled_output_dim,
activation='tanh',
kernel_initializer=pooler_layer_initializer,
@@ -268,7 +271,7 @@ def __init__(self,
# created using the Functional API. Once super().__init__ is called, we
# can assign attributes to `self` - note that all `self` assignments are
# below this line.
- super(EncoderScaffold, self).__init__(
+ super().__init__(
inputs=inputs, outputs=outputs, **kwargs)
self._hidden_cls = hidden_cls
@@ -293,7 +296,7 @@ def __init__(self,
self._embedding_norm_layer = embedding_norm_layer
self._hidden_layers = hidden_layers
if self._layer_norm_before_pooling:
- self._output_layer_norm = output_layer_norm
+ self._output_layer_norm = output_layer_norm # pyrefly: ignore[unbound-name]
self._pooler_layer = pooler_layer
self._layer_idx_as_attention_seed = layer_idx_as_attention_seed
@@ -303,7 +306,8 @@ def get_config(self):
config_dict = {
'num_hidden_instances': self._num_hidden_instances,
'pooled_output_dim': self._pooled_output_dim,
- 'pooler_layer_initializer': self._pooler_layer_initializer,
+ 'pooler_layer_initializer': tf_keras.initializers.serialize(
+ self._pooler_layer_initializer),
'embedding_cls': self._embedding_network,
'embedding_cfg': self._embedding_cfg,
'layer_norm_before_pooling': self._layer_norm_before_pooling,
@@ -323,9 +327,9 @@ def get_config(self):
# `self._hidden_cfg` may contain `class`, e.g., when `hidden_cfg` is
# `TransformerScaffold`, `attention_cls` argument can be a `class`.
if inspect.isclass(v):
- config_dict[cfg_name][k] = tf.keras.utils.get_registered_name(v)
+ config_dict[cfg_name][k] = tf_keras.utils.get_registered_name(v) # pyrefly: ignore[unsupported-operation]
else:
- config_dict[cfg_name][k] = v
+ config_dict[cfg_name][k] = v # pyrefly: ignore[unsupported-operation]
clss = {
'hidden_cls': self._hidden_cls,
@@ -335,7 +339,7 @@ def get_config(self):
for cls_name, cls in clss.items():
if inspect.isclass(cls):
key = '{}_string'.format(cls_name)
- config_dict[key] = tf.keras.utils.get_registered_name(cls)
+ config_dict[key] = tf_keras.utils.get_registered_name(cls)
else:
config_dict[cls_name] = cls
@@ -348,7 +352,7 @@ def from_config(cls, config, custom_objects=None):
for cls_name in cls_names:
cls_string = '{}_string'.format(cls_name)
if cls_string in config:
- config[cls_name] = tf.keras.utils.get_registered_object(
+ config[cls_name] = tf_keras.utils.get_registered_object(
config[cls_string], custom_objects=custom_objects)
del config[cls_string]
return cls(**config)
@@ -357,7 +361,7 @@ def get_embedding_table(self):
if self._embedding_network is None:
# In this case, we don't have a custom embedding network and can return
# the standard embedding data.
- return self._embedding_layer.embeddings
+ return self._embedding_layer.embeddings # pyrefly: ignore[missing-attribute]
if self._embedding_data is None:
raise RuntimeError(('The EncoderScaffold %s does not have a reference '
diff --git a/official/nlp/modeling/networks/encoder_scaffold_test.py b/official/nlp/modeling/networks/encoder_scaffold_test.py
index bc0b02e3cf0..b037c409f90 100644
--- a/official/nlp/modeling/networks/encoder_scaffold_test.py
+++ b/official/nlp/modeling/networks/encoder_scaffold_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,8 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.modeling import activations
from official.nlp.modeling import layers
from official.nlp.modeling.networks import encoder_scaffold
@@ -28,7 +27,7 @@
# at any point, the list passed to the config object will be filled with a
# boolean 'True'. We register this class as a Keras serializable so we can
# test serialization below.
-@tf.keras.utils.register_keras_serializable(package="TestOnly")
+@tf_keras.utils.register_keras_serializable(package="TestOnly")
class ValidatedTransformerLayer(layers.Transformer):
def __init__(self, call_list, call_class=None, **kwargs):
@@ -43,7 +42,7 @@ def call(self, inputs):
def get_config(self):
config = super(ValidatedTransformerLayer, self).get_config()
config["call_list"] = self.list
- config["call_class"] = tf.keras.utils.get_registered_name(self.call_class)
+ config["call_class"] = tf_keras.utils.get_registered_name(self.call_class)
return config
@@ -52,7 +51,7 @@ def get_config(self):
# object will be filled with a
# boolean 'True'. We register this class as a Keras serializable so we can
# test serialization below.
-@tf.keras.utils.register_keras_serializable(package="TestOnly")
+@tf_keras.utils.register_keras_serializable(package="TestOnly")
class ValidatedMaskLayer(layers.SelfAttentionMask):
def __init__(self, call_list, call_class=None, **kwargs):
@@ -67,23 +66,20 @@ def call(self, inputs, mask):
def get_config(self):
config = super(ValidatedMaskLayer, self).get_config()
config["call_list"] = self.list
- config["call_class"] = tf.keras.utils.get_registered_name(self.call_class)
+ config["call_class"] = tf_keras.utils.get_registered_name(self.call_class)
return config
-@tf.keras.utils.register_keras_serializable(package="TestLayerOnly")
-class TestLayer(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="TestLayerOnly")
+class TestLayer(tf_keras.layers.Layer):
pass
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class EncoderScaffoldLayerClassTest(keras_parameterized.TestCase):
+class EncoderScaffoldLayerClassTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(EncoderScaffoldLayerClassTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
@parameterized.named_parameters(
dict(testcase_name="only_final_output", return_all_layer_outputs=False),
@@ -98,7 +94,7 @@ def test_network_creation(self, return_all_layer_outputs):
"hidden_size": hidden_size,
"seq_length": sequence_length,
"max_seq_length": sequence_length,
- "initializer": tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ "initializer": tf_keras.initializers.TruncatedNormal(stddev=0.02),
"dropout_rate": 0.1,
}
@@ -115,7 +111,7 @@ def test_network_creation(self, return_all_layer_outputs):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
"call_list":
call_list
}
@@ -128,7 +124,7 @@ def test_network_creation(self, return_all_layer_outputs):
test_network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=num_hidden_instances,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cls=ValidatedTransformerLayer,
hidden_cfg=hidden_cfg,
@@ -138,9 +134,9 @@ def test_network_creation(self, return_all_layer_outputs):
layer_norm_before_pooling=True,
return_all_layer_outputs=return_all_layer_outputs)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
output_data, pooled = test_network([word_ids, mask, type_ids])
if return_all_layer_outputs:
@@ -151,7 +147,7 @@ def test_network_creation(self, return_all_layer_outputs):
data = output_data
self.assertIsInstance(test_network.hidden_layers, list)
self.assertLen(test_network.hidden_layers, num_hidden_instances)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
expected_data_shape = [None, sequence_length, hidden_size]
expected_pooled_shape = [None, hidden_size]
@@ -170,7 +166,7 @@ def test_network_creation(self, return_all_layer_outputs):
self.assertTrue(hasattr(test_network, "_output_layer_norm"))
def test_network_creation_with_float16_dtype(self):
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
hidden_size = 32
sequence_length = 21
embedding_cfg = {
@@ -179,7 +175,7 @@ def test_network_creation_with_float16_dtype(self):
"hidden_size": hidden_size,
"seq_length": sequence_length,
"max_seq_length": sequence_length,
- "initializer": tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ "initializer": tf_keras.initializers.TruncatedNormal(stddev=0.02),
"dropout_rate": 0.1,
}
hidden_cfg = {
@@ -194,20 +190,20 @@ def test_network_creation_with_float16_dtype(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
}
# Create a small EncoderScaffold for testing.
test_network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cfg=hidden_cfg,
embedding_cfg=embedding_cfg)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask, type_ids])
expected_data_shape = [None, sequence_length, hidden_size]
@@ -231,7 +227,7 @@ def test_network_invocation(self):
"hidden_size": hidden_size,
"seq_length": sequence_length,
"max_seq_length": sequence_length,
- "initializer": tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ "initializer": tf_keras.initializers.TruncatedNormal(stddev=0.02),
"dropout_rate": 0.1,
}
hidden_cfg = {
@@ -246,26 +242,26 @@ def test_network_invocation(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
}
# Create a small EncoderScaffold for testing.
test_network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cfg=hidden_cfg,
embedding_cfg=embedding_cfg,
dict_outputs=True)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
outputs = test_network([word_ids, mask, type_ids])
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], outputs)
+ model = tf_keras.Model([word_ids, mask, type_ids], outputs)
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -286,7 +282,7 @@ def test_network_invocation(self):
"hidden_size": hidden_size,
"seq_length": sequence_length,
"max_seq_length": sequence_length * 2,
- "initializer": tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ "initializer": tf_keras.initializers.TruncatedNormal(stddev=0.02),
"dropout_rate": 0.1,
}
hidden_cfg = {
@@ -301,18 +297,18 @@ def test_network_invocation(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
}
# Create a small EncoderScaffold for testing.
test_network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cfg=hidden_cfg,
embedding_cfg=embedding_cfg)
outputs = test_network([word_ids, mask, type_ids])
- model = tf.keras.Model([word_ids, mask, type_ids], outputs)
+ model = tf_keras.Model([word_ids, mask, type_ids], outputs)
_ = model.predict([word_id_data, mask_data, type_id_data])
def test_serialize_deserialize(self):
@@ -325,7 +321,7 @@ def test_serialize_deserialize(self):
"hidden_size": hidden_size,
"seq_length": sequence_length,
"max_seq_length": sequence_length,
- "initializer": tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ "initializer": tf_keras.initializers.TruncatedNormal(stddev=0.02),
"dropout_rate": 0.1,
}
hidden_cfg = {
@@ -340,13 +336,13 @@ def test_serialize_deserialize(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
}
# Create a small EncoderScaffold for testing.
network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cfg=hidden_cfg,
embedding_cfg=embedding_cfg)
@@ -362,20 +358,20 @@ def test_serialize_deserialize(self):
self.assertAllEqual(network.get_config(), new_network.get_config())
-class Embeddings(tf.keras.Model):
+class Embeddings(tf_keras.Model):
def __init__(self, vocab_size, hidden_size):
super().__init__()
self.inputs = [
- tf.keras.layers.Input(
+ tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name="input_word_ids"),
- tf.keras.layers.Input(shape=(None,), dtype=tf.int32, name="input_mask")
+ tf_keras.layers.Input(shape=(None,), dtype=tf.int32, name="input_mask")
]
self.attention_mask = layers.SelfAttentionMask()
self.embedding_layer = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=hidden_size,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
name="word_embeddings")
def call(self, inputs):
@@ -384,8 +380,7 @@ def call(self, inputs):
return word_embeddings, self.attention_mask([word_embeddings, mask])
-@keras_parameterized.run_all_keras_modes
-class EncoderScaffoldEmbeddingNetworkTest(keras_parameterized.TestCase):
+class EncoderScaffoldEmbeddingNetworkTest(tf.test.TestCase):
def test_network_invocation(self):
hidden_size = 32
@@ -409,25 +404,25 @@ def test_network_invocation(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
}
# Create a small EncoderScaffold for testing.
test_network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cfg=hidden_cfg,
embedding_cls=network)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask])
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask], [data, pooled])
+ model = tf_keras.Model([word_ids, mask], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -446,18 +441,18 @@ def test_serialize_deserialize(self):
# will have 2 inputs (mask and word_ids) instead of 3, and won't use
# positional embeddings.
- word_ids = tf.keras.layers.Input(
+ word_ids = tf_keras.layers.Input(
shape=(sequence_length,), dtype=tf.int32, name="input_word_ids")
- mask = tf.keras.layers.Input(
+ mask = tf_keras.layers.Input(
shape=(sequence_length,), dtype=tf.int32, name="input_mask")
embedding_layer = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=hidden_size,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
name="word_embeddings")
word_embeddings = embedding_layer(word_ids)
attention_mask = layers.SelfAttentionMask()([word_embeddings, mask])
- network = tf.keras.Model([word_ids, mask],
+ network = tf_keras.Model([word_ids, mask],
[word_embeddings, attention_mask])
hidden_cfg = {
@@ -472,14 +467,14 @@ def test_serialize_deserialize(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
}
# Create a small EncoderScaffold for testing.
test_network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cfg=hidden_cfg,
embedding_cls=network,
@@ -496,14 +491,14 @@ def test_serialize_deserialize(self):
self.assertAllEqual(test_network.get_config(), new_network.get_config())
# Create a model based off of the old and new networks:
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = new_network([word_ids, mask])
- new_model = tf.keras.Model([word_ids, mask], [data, pooled])
+ new_model = tf_keras.Model([word_ids, mask], [data, pooled])
data, pooled = test_network([word_ids, mask])
- model = tf.keras.Model([word_ids, mask], [data, pooled])
+ model = tf_keras.Model([word_ids, mask], [data, pooled])
# Copy the weights between models.
new_model.set_weights(model.get_weights())
@@ -525,8 +520,8 @@ def test_serialize_deserialize(self):
new_network.get_embedding_table()
-@keras_parameterized.run_all_keras_modes
-class EncoderScaffoldHiddenInstanceTest(keras_parameterized.TestCase):
+class EncoderScaffoldHiddenInstanceTest(
+ tf.test.TestCase, parameterized.TestCase):
def test_network_invocation(self):
hidden_size = 32
@@ -540,7 +535,7 @@ def test_network_invocation(self):
"hidden_size": hidden_size,
"seq_length": sequence_length,
"max_seq_length": sequence_length,
- "initializer": tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ "initializer": tf_keras.initializers.TruncatedNormal(stddev=0.02),
"dropout_rate": 0.1,
}
@@ -557,7 +552,7 @@ def test_network_invocation(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
"call_list":
call_list
}
@@ -574,20 +569,20 @@ def test_network_invocation(self):
test_network = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cls=xformer,
mask_cls=xmask,
embedding_cfg=embedding_cfg)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = test_network([word_ids, mask, type_ids])
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -624,7 +619,7 @@ def test_hidden_cls_list(self):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
"call_list":
call_list
}
@@ -640,7 +635,7 @@ def test_hidden_cls_list(self):
test_network_a = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
hidden_cls=xformer,
mask_cls=xmask,
@@ -649,7 +644,7 @@ def test_hidden_cls_list(self):
test_network_b = encoder_scaffold.EncoderScaffold(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
mask_cls=xmask,
embedding_cls=test_network_a.embedding_network,
@@ -661,25 +656,25 @@ def test_hidden_cls_list(self):
test_network_c = encoder_scaffold.EncoderScaffold(
num_hidden_instances=2,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
mask_cls=xmask,
embedding_cls=test_network_a.embedding_network,
hidden_cls=hidden_layers)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
# Create model based off of network a:
data_a, pooled_a = test_network_a([word_ids, mask])
- model_a = tf.keras.Model([word_ids, mask], [data_a, pooled_a])
+ model_a = tf_keras.Model([word_ids, mask], [data_a, pooled_a])
# Create model based off of network b:
data_b, pooled_b = test_network_b([word_ids, mask])
- model_b = tf.keras.Model([word_ids, mask], [data_b, pooled_b])
+ model_b = tf_keras.Model([word_ids, mask], [data_b, pooled_b])
# Create model based off of network b:
data_c, pooled_c = test_network_c([word_ids, mask])
- model_c = tf.keras.Model([word_ids, mask], [data_c, pooled_c])
+ model_c = tf_keras.Model([word_ids, mask], [data_c, pooled_c])
batch_size = 3
word_id_data = np.random.randint(
@@ -709,7 +704,7 @@ def test_serialize_deserialize(self, use_hidden_cls_instance):
"hidden_size": hidden_size,
"seq_length": sequence_length,
"max_seq_length": sequence_length,
- "initializer": tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ "initializer": tf_keras.initializers.TruncatedNormal(stddev=0.02),
"dropout_rate": 0.1,
}
@@ -726,7 +721,7 @@ def test_serialize_deserialize(self, use_hidden_cls_instance):
"attention_dropout_rate":
0.1,
"kernel_initializer":
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
"call_list":
call_list,
"call_class":
@@ -739,7 +734,7 @@ def test_serialize_deserialize(self, use_hidden_cls_instance):
kwargs = dict(
num_hidden_instances=3,
pooled_output_dim=hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=0.02),
embedding_cfg=embedding_cfg)
@@ -767,15 +762,15 @@ def test_serialize_deserialize(self, use_hidden_cls_instance):
self.assertAllEqual(test_network.get_config(), new_network.get_config())
# Create a model based off of the old and new networks:
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
data, pooled = new_network([word_ids, mask, type_ids])
- new_model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ new_model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
data, pooled = test_network([word_ids, mask, type_ids])
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Copy the weights between models.
new_model.set_weights(model.get_weights())
diff --git a/official/nlp/modeling/networks/fnet.py b/official/nlp/modeling/networks/fnet.py
new file mode 100644
index 00000000000..82ef4e03a1f
--- /dev/null
+++ b/official/nlp/modeling/networks/fnet.py
@@ -0,0 +1,351 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""FNet encoder network.
+
+Based on ["FNet: Mixing Tokens with Fourier Transforms"]
+(https://aclanthology.org/2022.naacl-main.319/).
+"""
+# pylint: disable=g-classes-have-attributes
+
+from typing import Any, Callable, Optional, Sequence, Union
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.nlp.modeling import layers
+
+_Activation = Union[str, Callable[..., Any]]
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
+
+
+class FNet(tf_keras.layers.Layer):
+ """FNet encoder network.
+
+ Based on ["FNet: Mixing Tokens with Fourier Transforms"]
+ (https://aclanthology.org/2022.naacl-main.319/). FNet is an efficient
+ Transformer-like encoder network that replaces self-attention sublayers with
+ Fourier sublayers.
+
+ This implementation defaults to the canonical FNet Base model, but the network
+ also supports more general mixing models (e.g. 'Linear', 'HNet') and hybrid
+ models (e.g. 'FNet-Hybrid') models that use both mixing and self-attention
+ layers. The input length is fixed to 'max_sequence_length'.
+
+ Args:
+ vocab_size: The size of the token vocabulary.
+ hidden_size: The size of the transformer hidden layers.
+ num_layers: The number of transformer layers.
+ mixing_mechanism: Type of mixing mechanism used in place of self-attention
+ layers. Defaults to FNet ('Fourier') mixing.
+ use_fft: Only used for spectral mixing mechanisms. Determines whether to use
+ Fast Fourier Transform (True) or the Discrete Fourier Transform (DFT)
+ matrix (False; default) to compute the Fourier Transform. See
+ layers.FourierTransformLayer or layers.HartleyTransformLayer for advice.
+ attention_layers: Specifies which layers, if any, should be attention layers
+ in the encoder. The remaining [0, num_layers) setminus attention_layers
+ will use the specified `mixing_mechanism`. If using attention layers, a
+ good rule of thumb is to place them in the final few layers.
+ num_attention_heads: The number of attention heads for each transformer. The
+ hidden size must be divisible by the number of attention heads.
+ max_sequence_length: The only sequence length that this encoder can
+ consume. This determines the variable shape for positional embeddings and
+ the size of the mixing matrices.
+ type_vocab_size: The number of types that the 'type_ids' input can take.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ output_dropout: Dropout probability for the post-attention and output
+ dropout.
+ attention_dropout: The dropout rate to use for the attention layers within
+ the transformer layers.
+ initializer: The initializer to use for all weights in this encoder.
+ output_range: The sequence output range, [0, output_range), by slicing the
+ target sequence of the last transformer layer. `None` means the entire
+ target sequence will attend to the source sequence, which yields the full
+ output.
+ embedding_width: The width of the word embeddings. If the embedding width is
+ not equal to hidden size, embedding parameters will be factorized into two
+ matrices in the shape of ['vocab_size', 'embedding_width'] and
+ ['embedding_width', 'hidden_size'] ('embedding_width' is usually much
+ smaller than 'hidden_size').
+ embedding_layer: An optional Layer instance which will be called to generate
+ embeddings for the input word IDs.
+ norm_first: Whether to normalize inputs to attention and intermediate dense
+ layers. If set False, output of attention and intermediate dense layers is
+ normalized.
+ with_dense_inputs: Whether to accept dense embeddings as the input.
+ """
+
+ def __init__(
+ self,
+ vocab_size: int,
+ hidden_size: int = 768,
+ num_layers: int = 12,
+ mixing_mechanism: layers.MixingMechanism = layers.MixingMechanism.FOURIER,
+ use_fft: bool = False,
+ attention_layers: Sequence[int] = (),
+ num_attention_heads: int = 12,
+ max_sequence_length: int = 512,
+ type_vocab_size: int = 16,
+ inner_dim: int = 3072,
+ inner_activation: _Activation = _approx_gelu,
+ output_dropout: float = 0.1,
+ attention_dropout: float = 0.1,
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
+ stddev=0.02),
+ output_range: Optional[int] = None,
+ embedding_width: Optional[int] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
+ norm_first: bool = False,
+ with_dense_inputs: bool = False,
+ **kwargs):
+ super().__init__(**kwargs)
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
+
+ if embedding_width is None:
+ embedding_width = hidden_size
+
+ self._config = {
+ 'vocab_size': vocab_size,
+ 'hidden_size': hidden_size,
+ 'num_layers': num_layers,
+ 'mixing_mechanism': mixing_mechanism,
+ 'use_fft': use_fft,
+ 'attention_layers': attention_layers,
+ 'num_attention_heads': num_attention_heads,
+ 'max_sequence_length': max_sequence_length,
+ 'type_vocab_size': type_vocab_size,
+ 'inner_dim': inner_dim,
+ 'inner_activation': tf_keras.activations.serialize(activation),
+ 'output_dropout': output_dropout,
+ 'attention_dropout': attention_dropout,
+ 'initializer': tf_keras.initializers.serialize(initializer),
+ 'output_range': output_range,
+ 'embedding_width': embedding_width,
+ 'embedding_layer': embedding_layer,
+ 'norm_first': norm_first,
+ 'with_dense_inputs': with_dense_inputs,
+ }
+
+ if embedding_layer is None:
+ self._embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ initializer=tf_utils.clone_initializer(initializer),
+ name='word_embeddings')
+ else:
+ self._embedding_layer = embedding_layer
+
+ self._position_embedding_layer = layers.PositionEmbedding(
+ initializer=tf_utils.clone_initializer(initializer),
+ max_length=max_sequence_length,
+ name='position_embedding')
+
+ self._type_embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=type_vocab_size,
+ embedding_width=embedding_width,
+ initializer=tf_utils.clone_initializer(initializer),
+ use_one_hot=True,
+ name='type_embeddings')
+
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
+ name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
+
+ self._embedding_dropout = tf_keras.layers.Dropout(
+ rate=output_dropout, name='embedding_dropout')
+
+ # We project the 'embedding' output to 'hidden_size' if it is not already
+ # 'hidden_size'.
+ self._embedding_projection = None
+ if embedding_width != hidden_size:
+ self._embedding_projection = tf_keras.layers.EinsumDense(
+ '...x,xy->...y',
+ output_shape=hidden_size,
+ bias_axes='y',
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='embedding_projection')
+
+ self._transformer_layers = []
+ for layer in range(num_layers):
+ if layer in attention_layers:
+ mixing_layer = layers.MultiHeadAttention(
+ num_heads=num_attention_heads,
+ key_dim=int(hidden_size // num_attention_heads),
+ dropout=attention_dropout,
+ use_bias=True,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='self_attention',
+ )
+ else:
+ mixing_layer = self._init_mixing_sublayer(layer)
+
+ block = layers.TransformerScaffold(
+ num_attention_heads=num_attention_heads,
+ inner_dim=inner_dim,
+ inner_activation=inner_activation,
+ attention_cls=mixing_layer,
+ feedforward_cls=None, # Fallback to default FeedForward class
+ output_dropout=output_dropout,
+ attention_dropout=attention_dropout,
+ norm_first=norm_first,
+ output_range=output_range if layer == num_layers - 1 else None,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='transformer/layer_%d' % layer)
+ self._transformer_layers.append(block)
+
+ self._attention_mask_layer = layers.SelfAttentionMask(
+ name='self_attention_mask')
+
+ self._pooler_layer = tf_keras.layers.Dense(
+ units=hidden_size,
+ activation='tanh',
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='pooler_transform')
+
+ if with_dense_inputs:
+ self.inputs = dict(
+ # The total length of token ids and dense inputs still has to be
+ # max_sequence_length. It is checked in call().
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_inputs=tf_keras.Input(
+ shape=(None, embedding_width), dtype=tf.float32),
+ dense_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ )
+
+ else:
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(
+ shape=(max_sequence_length,), dtype=tf.int32),
+ input_mask=tf_keras.Input(
+ shape=(max_sequence_length,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(
+ shape=(max_sequence_length,), dtype=tf.int32))
+ self._max_sequence_length = max_sequence_length
+
+ def call(self, inputs):
+ word_embeddings = None
+ if isinstance(inputs, dict):
+ word_ids = inputs.get('input_word_ids')
+ mask = inputs.get('input_mask')
+ type_ids = inputs.get('input_type_ids')
+ word_embeddings = inputs.get('input_word_embeddings', None)
+
+ dense_inputs = inputs.get('dense_inputs', None)
+ dense_mask = inputs.get('dense_mask', None)
+ dense_type_ids = inputs.get('dense_type_ids', None)
+ else:
+ raise ValueError('Unexpected inputs type (%s) to %s.' %
+ (type(inputs), self.__class__))
+
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ # Concat the dense embeddings at sequence end.
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+ mask = tf.concat([mask, dense_mask], axis=1)
+
+ # FNet: Sequence length must be the same as `max_sequence_length`.
+ word_embeddings = tf.ensure_shape(word_embeddings,
+ [None, self._max_sequence_length, None])
+
+ # Absolute position embeddings.
+ position_embeddings = self._position_embedding_layer(word_embeddings)
+ type_embeddings = self._type_embedding_layer(type_ids)
+
+ embeddings = word_embeddings + position_embeddings + type_embeddings
+ embeddings = self._embedding_norm_layer(embeddings)
+ embeddings = self._embedding_dropout(embeddings)
+
+ if self._embedding_projection is not None:
+ embeddings = self._embedding_projection(embeddings)
+
+ attention_mask = self._attention_mask_layer(embeddings, mask)
+
+ encoder_outputs = []
+ x = embeddings
+ for layer in self._transformer_layers:
+ x = layer([x, attention_mask])
+ encoder_outputs.append(x)
+
+ last_encoder_output = encoder_outputs[-1]
+ first_token_tensor = last_encoder_output[:, 0, :]
+ pooled_output = self._pooler_layer(first_token_tensor)
+
+ output = dict(
+ sequence_output=encoder_outputs[-1],
+ pooled_output=pooled_output,
+ encoder_outputs=encoder_outputs)
+ return output
+
+ def get_embedding_table(self):
+ return self._embedding_layer.embeddings
+
+ def get_embedding_layer(self):
+ return self._embedding_layer
+
+ def get_config(self):
+ return dict(self._config)
+
+ @property
+ def transformer_layers(self):
+ """List of Transformer layers in the encoder."""
+ return self._transformer_layers
+
+ @property
+ def pooler_layer(self):
+ """The pooler dense layer after the transformer layers."""
+ return self._pooler_layer
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ if 'embedding_layer' in config and config['embedding_layer'] is not None:
+ warn_string = (
+ 'You are reloading a model that was saved with a '
+ 'potentially-shared embedding layer object. If you contine to '
+ 'train this model, the embedding layer will no longer be shared. '
+ 'To work around this, load the model outside of the Keras API.')
+ print('WARNING: ' + warn_string)
+ logging.warn(warn_string)
+
+ return cls(**config)
+
+ def _init_mixing_sublayer(self, layer: int):
+ """Initializes config-dependent mixing sublayer."""
+ if self._config['mixing_mechanism'] == layers.MixingMechanism.FOURIER:
+ mixing_sublayer = layers.FourierTransformLayer(
+ use_fft=self._config['use_fft'], name='fourier_transform')
+ elif self._config['mixing_mechanism'] == layers.MixingMechanism.HARTLEY:
+ mixing_sublayer = layers.HartleyTransformLayer(
+ use_fft=self._config['use_fft'], name='hartley_transform')
+ elif self._config['mixing_mechanism'] == layers.MixingMechanism.LINEAR:
+ mixing_sublayer = layers.LinearTransformLayer(
+ kernel_initializer=tf_utils.clone_initializer(
+ self._config['initializer']),
+ name='linear_transform')
+ else:
+ raise ValueError('Unsupported mixing mechanism: %s' %
+ self._config['mixing_mechanism'])
+
+ return mixing_sublayer
diff --git a/official/nlp/modeling/networks/fnet_test.py b/official/nlp/modeling/networks/fnet_test.py
new file mode 100644
index 00000000000..fecd1a2a79b
--- /dev/null
+++ b/official/nlp/modeling/networks/fnet_test.py
@@ -0,0 +1,119 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for FNet encoder network."""
+
+from typing import Sequence
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+from official.nlp.modeling.networks import fnet
+
+
+class FNetTest(parameterized.TestCase, tf.test.TestCase):
+
+ def tearDown(self):
+ super(FNetTest, self).tearDown()
+ tf_keras.mixed_precision.set_global_policy("float32")
+
+ @parameterized.named_parameters(
+ ("fnet", layers.MixingMechanism.FOURIER, ()),
+ ("fnet_hybrid", layers.MixingMechanism.FOURIER, (1, 2)),
+ ("hnet", layers.MixingMechanism.HARTLEY, ()),
+ ("hnet_hybrid", layers.MixingMechanism.HARTLEY, (1, 2)),
+ ("linear", layers.MixingMechanism.LINEAR, ()),
+ ("linear_hybrid", layers.MixingMechanism.LINEAR, (0,)),
+ ("bert", layers.MixingMechanism.FOURIER, (0, 1, 2)),
+ )
+ def test_network(self, mixing_mechanism: layers.MixingMechanism,
+ attention_layers: Sequence[int]):
+ num_layers = 3
+ hidden_size = 32
+ sequence_length = 21
+ test_network = fnet.FNet(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=2,
+ max_sequence_length=sequence_length,
+ num_layers=num_layers,
+ mixing_mechanism=mixing_mechanism,
+ attention_layers=attention_layers)
+
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+
+ dict_outputs = test_network(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ data = dict_outputs["sequence_output"]
+ pooled = dict_outputs["pooled_output"]
+
+ self.assertIsInstance(test_network.transformer_layers, list)
+ self.assertLen(test_network.transformer_layers, 3)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
+
+ expected_data_shape = [None, sequence_length, hidden_size]
+ expected_pooled_shape = [None, hidden_size]
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+ self.assertAllEqual(expected_pooled_shape, pooled.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, data.dtype)
+ self.assertAllEqual(tf.float32, pooled.dtype)
+
+ def test_embeddings_as_inputs(self):
+ hidden_size = 32
+ sequence_length = 21
+ test_network = fnet.FNet(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=2,
+ max_sequence_length=sequence_length,
+ num_layers=3)
+
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+
+ test_network.build(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ embeddings = test_network.get_embedding_layer()(word_ids)
+
+ # Calls with the embeddings.
+ dict_outputs = test_network(
+ dict(
+ input_word_embeddings=embeddings,
+ input_mask=mask,
+ input_type_ids=type_ids))
+ all_encoder_outputs = dict_outputs["encoder_outputs"]
+ pooled = dict_outputs["pooled_output"]
+
+ expected_data_shape = [None, sequence_length, hidden_size]
+ expected_pooled_shape = [None, hidden_size]
+ self.assertLen(all_encoder_outputs, 3)
+ for data in all_encoder_outputs:
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+ self.assertAllEqual(expected_pooled_shape, pooled.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, all_encoder_outputs[-1].dtype)
+ self.assertAllEqual(tf.float32, pooled.dtype)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/nlp/modeling/networks/funnel_transformer.py b/official/nlp/modeling/networks/funnel_transformer.py
index 5be3e8308a7..af3057ab435 100644
--- a/official/nlp/modeling/networks/funnel_transformer.py
+++ b/official/nlp/modeling/networks/funnel_transformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,14 +15,17 @@
"""Funnel Transformer network."""
# pylint: disable=g-classes-have-attributes
-from typing import Any, Callable, Optional, Union, Sequence
+import math
+from typing import Any, Callable, Optional, Sequence, Union
+
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
-_Initializer = Union[str, tf.keras.initializers.Initializer]
+_Initializer = Union[str, tf_keras.initializers.Initializer]
_Activation = Union[str, Callable[..., Any]]
_MAX = 'max'
@@ -39,12 +42,12 @@
'ReZeroTransformer': layers.ReZeroTransformer
}
-_approx_gelu = lambda x: tf.keras.activations.gelu(x, approximate=True)
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
def _get_policy_dtype():
try:
- return tf.keras.mixed_precision.global_policy().compute_dtype or tf.float32
+ return tf_keras.mixed_precision.global_policy().compute_dtype or tf.float32
except AttributeError: # tf1 has no attribute 'global_policy'
return tf.float32
@@ -90,8 +93,39 @@ def _pool_and_concat(mask, unpool_length: int, strides: Union[Sequence[int],
return mask
-def _create_truncated_avg_transforms(seq_length: int,
- pool_strides: Sequence[int]):
+def _create_fractional_pool_transform(sl: int, pool_factor: float):
+ """Create pooling transform for fractional pooling factor."""
+
+ assert pool_factor > 1.0, '`pool_factor` should be > 1.0.'
+
+ psl = int(sl / pool_factor)
+ gcd_ = math.gcd(sl, psl)
+ # It is expected chunk_sl and chunk_psl are small integers.
+ # The transform is built by tiling a [chunk_sl, chunk_psl] submatrix
+ # gcd_ times. The submatrix sums to chunk_psl.
+ chunk_sl = sl // gcd_
+ chunk_psl = psl // gcd_
+ num_one_entries = chunk_psl - 1
+ num_frac_entries = chunk_sl - (chunk_psl - 1)
+
+ # The transform is of shape [sl, psl].
+ transform = np.zeros((sl, psl))
+ for i in range(sl // chunk_sl):
+ row_start = chunk_sl * i
+ col_start = chunk_psl * i
+ for idx in range(num_one_entries):
+ transform[row_start + idx][col_start + idx] = 1.0
+ for idx in range(num_frac_entries):
+ transform[row_start + num_one_entries + idx][
+ col_start + num_one_entries
+ ] = (1.0 / num_frac_entries)
+
+ return tf.constant(transform, dtype=_get_policy_dtype())
+
+
+def _create_truncated_avg_transforms(
+ seq_length: int, pool_strides: Sequence[int]
+):
"""Computes pooling transforms.
The pooling_transform is of shape [seq_length,
@@ -121,15 +155,20 @@ def _create_truncated_avg_transforms(seq_length: int,
if pool_stride == 1:
pooling_transforms.append(None)
else:
- pooled_seq_length = seq_length // pool_stride
-
- pfac, sl, psl = pool_stride, seq_length, pooled_seq_length
- transform = [[1.0 if (i // pfac) == j else 0.0
- for j in range(psl)]
- for i in range(sl)]
- transform = tf.constant(transform, dtype=_get_policy_dtype())
-
- pooling_transforms.append(transform / pool_stride)
+ pooled_seq_length = int(seq_length / pool_stride)
+ if (1.0 * pool_stride).is_integer():
+ pfac, sl, psl = pool_stride, seq_length, pooled_seq_length
+
+ transform = [
+ [1.0 if (i // pfac) == j else 0.0 for j in range(psl)]
+ for i in range(sl)
+ ]
+ transform = (
+ tf.constant(transform, dtype=_get_policy_dtype()) / pool_stride
+ )
+ else:
+ transform = _create_fractional_pool_transform(seq_length, pool_stride)
+ pooling_transforms.append(transform)
seq_length = pooled_seq_length
return pooling_transforms
@@ -163,20 +202,24 @@ def create_2d_mask(from_length, mask):
if pool_stride == 1:
attention_masks.append(create_2d_mask(seq_length, layer_mask))
else:
- pooled_seq_length = seq_length // pool_stride
+ pooled_seq_length = tf.cast(
+ tf.cast(seq_length, tf.float32) / tf.cast(pool_stride, tf.float32),
+ tf.int32,
+ )
attention_masks.append(create_2d_mask(pooled_seq_length, layer_mask))
layer_mask = tf.cast(
tf.einsum('BF,FT->BT', layer_mask, transform) > 0.0,
- dtype=layer_mask.dtype)
+ dtype=layer_mask.dtype,
+ )
seq_length = pooled_seq_length
del seq_length
return attention_masks
-@tf.keras.utils.register_keras_serializable(package='Text')
-class FunnelTransformerEncoder(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class FunnelTransformerEncoder(tf_keras.layers.Layer):
"""Funnel Transformer-based encoder network.
Funnel Transformer Implementation of https://arxiv.org/abs/2006.03236.
@@ -242,21 +285,30 @@ def __init__(
output_dropout: float = 0.1,
attention_dropout: float = 0.1,
pool_type: str = _MAX,
- pool_stride: int = 2,
+ pool_stride: Union[int, Sequence[Union[int, float]]] = 2,
unpool_length: int = 0,
- initializer: _Initializer = tf.keras.initializers.TruncatedNormal(
- stddev=0.02),
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
+ stddev=0.02
+ ),
output_range: Optional[int] = None,
embedding_width: Optional[int] = None,
- embedding_layer: Optional[tf.keras.layers.Layer] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
norm_first: bool = False,
transformer_cls: Union[
- str, tf.keras.layers.Layer] = layers.TransformerEncoderBlock,
- share_rezero: bool = True,
- **kwargs):
+ str, tf_keras.layers.Layer
+ ] = layers.TransformerEncoderBlock,
+ share_rezero: bool = False,
+ append_dense_inputs: bool = False,
+ **kwargs
+ ):
super().__init__(**kwargs)
- activation = tf.keras.activations.get(inner_activation)
- initializer = tf.keras.initializers.get(initializer)
+
+ if output_range is not None:
+ logging.warning('`output_range` is available as an argument for `call()`.'
+ 'The `output_range` as __init__ argument is deprecated.')
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
if embedding_width is None:
embedding_width = hidden_size
@@ -265,38 +317,38 @@ def __init__(
self._embedding_layer = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
name='word_embeddings')
else:
self._embedding_layer = embedding_layer
self._position_embedding_layer = layers.PositionEmbedding(
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
max_length=max_sequence_length,
name='position_embedding')
self._type_embedding_layer = layers.OnDeviceEmbedding(
vocab_size=type_vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_one_hot=True,
name='type_embeddings')
- self._embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
- self._embedding_dropout = tf.keras.layers.Dropout(
+ self._embedding_dropout = tf_keras.layers.Dropout(
rate=output_dropout, name='embedding_dropout')
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
self._embedding_projection = None
if embedding_width != hidden_size:
- self._embedding_projection = tf.keras.layers.experimental.EinsumDense(
+ self._embedding_projection = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='embedding_projection')
self._transformer_layers = []
@@ -305,6 +357,7 @@ def __init__(
# Will raise an error if the string is not supported.
if isinstance(transformer_cls, str):
transformer_cls = _str2transformer_cls[transformer_cls]
+ self._num_layers = num_layers
for i in range(num_layers):
layer = transformer_cls(
num_attention_heads=num_attention_heads,
@@ -315,16 +368,15 @@ def __init__(
output_dropout=output_dropout,
attention_dropout=attention_dropout,
norm_first=norm_first,
- output_range=output_range if i == num_layers - 1 else None,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
share_rezero=share_rezero,
name='transformer/layer_%d' % i)
self._transformer_layers.append(layer)
- self._pooler_layer = tf.keras.layers.Dense(
+ self._pooler_layer = tf_keras.layers.Dense(
units=hidden_size,
activation='tanh',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='pooler_transform')
if isinstance(pool_stride, int):
# TODO(b/197133196): Pooling layer can be shared.
@@ -333,83 +385,76 @@ def __init__(
if len(pool_stride) != num_layers:
raise ValueError('Lengths of pool_stride and num_layers are not equal.')
pool_strides = pool_stride
- # TODO(crickwu): explore tf.keras.layers.serialize method.
+
+ is_fractional_pooling = False in [
+ (1.0 * pool_stride).is_integer() for pool_stride in pool_strides
+ ]
+ if is_fractional_pooling and pool_type in [_MAX, _AVG]:
+ raise ValueError(
+ 'Fractional pooling is only supported for'
+ ' `pool_type`=`truncated_average`'
+ )
+
+ # TODO(crickwu): explore tf_keras.layers.serialize method.
if pool_type == _MAX:
- pool_cls = tf.keras.layers.MaxPooling1D
+ pool_cls = tf_keras.layers.MaxPooling1D
elif pool_type == _AVG:
- pool_cls = tf.keras.layers.AveragePooling1D
+ pool_cls = tf_keras.layers.AveragePooling1D
elif pool_type == _TRUNCATED_AVG:
# TODO(b/203665205): unpool_length should be implemented.
if unpool_length != 0:
raise ValueError('unpool_length is not supported by truncated_avg now.')
- # Compute the attention masks and pooling transforms.
- self._pooling_transforms = _create_truncated_avg_transforms(
- max_sequence_length, pool_strides)
else:
raise ValueError('pool_type not supported.')
if pool_type in (_MAX, _AVG):
self._att_input_pool_layers = []
for layer_pool_stride in pool_strides:
- att_input_pool_layer = pool_cls(
+ att_input_pool_layer = pool_cls( # pyrefly: ignore[unbound-name]
pool_size=layer_pool_stride,
strides=layer_pool_stride,
padding='same',
name='att_input_pool_layer')
self._att_input_pool_layers.append(att_input_pool_layer)
+ self._max_sequence_length = max_sequence_length
self._pool_strides = pool_strides # This is a list here.
self._unpool_length = unpool_length
self._pool_type = pool_type
+ self._append_dense_inputs = append_dense_inputs
self._config = {
- 'vocab_size':
- vocab_size,
- 'hidden_size':
- hidden_size,
- 'num_layers':
- num_layers,
- 'num_attention_heads':
- num_attention_heads,
- 'max_sequence_length':
- max_sequence_length,
- 'type_vocab_size':
- type_vocab_size,
- 'inner_dim':
- inner_dim,
- 'inner_activation':
- tf.keras.activations.serialize(activation),
- 'output_dropout':
- output_dropout,
- 'attention_dropout':
- attention_dropout,
- 'initializer':
- tf.keras.initializers.serialize(initializer),
- 'output_range':
- output_range,
- 'embedding_width':
- embedding_width,
- 'embedding_layer':
- embedding_layer,
- 'norm_first':
- norm_first,
- 'pool_type':
- pool_type,
- 'pool_stride':
- pool_stride,
- 'unpool_length':
- unpool_length,
- 'transformer_cls':
- _transformer_cls2str.get(transformer_cls, str(transformer_cls))
+ 'vocab_size': vocab_size,
+ 'hidden_size': hidden_size,
+ 'num_layers': num_layers,
+ 'num_attention_heads': num_attention_heads,
+ 'max_sequence_length': max_sequence_length,
+ 'type_vocab_size': type_vocab_size,
+ 'inner_dim': inner_dim,
+ 'inner_activation': tf_keras.activations.serialize(activation),
+ 'output_dropout': output_dropout,
+ 'attention_dropout': attention_dropout,
+ 'initializer': tf_keras.initializers.serialize(initializer),
+ 'output_range': output_range,
+ 'embedding_width': embedding_width,
+ 'embedding_layer': embedding_layer,
+ 'norm_first': norm_first,
+ 'pool_type': pool_type,
+ 'pool_stride': pool_stride,
+ 'unpool_length': unpool_length,
+ 'transformer_cls': _transformer_cls2str.get(
+ transformer_cls, str(transformer_cls)
+ ),
}
self.inputs = dict(
- input_word_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32))
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32))
- def call(self, inputs):
+ def call(self, inputs, output_range: Optional[tf.Tensor] = None):
# inputs are [word_ids, mask, type_ids]
+ word_embeddings = None
if isinstance(inputs, (list, tuple)):
logging.warning('List inputs to %s are discouraged.', self.__class__)
if len(inputs) == 3:
@@ -418,14 +463,19 @@ def call(self, inputs):
dense_mask = None
dense_type_ids = None
elif len(inputs) == 6:
- word_ids, mask, type_ids, dense_inputs, dense_mask, dense_type_ids = inputs
+ word_ids, mask, type_ids, dense_inputs, dense_mask, dense_type_ids = (
+ inputs
+ )
else:
- raise ValueError('Unexpected inputs to %s with length at %d.' %
- (self.__class__, len(inputs)))
+ raise ValueError(
+ 'Unexpected inputs to %s with length at %d.'
+ % (self.__class__, len(inputs))
+ )
elif isinstance(inputs, dict):
word_ids = inputs.get('input_word_ids')
mask = inputs.get('input_mask')
type_ids = inputs.get('input_type_ids')
+ word_embeddings = inputs.get('input_word_embeddings', None)
dense_inputs = inputs.get('dense_inputs', None)
dense_mask = inputs.get('dense_mask', None)
@@ -433,19 +483,31 @@ def call(self, inputs):
else:
raise ValueError('Unexpected inputs type to %s.' % self.__class__)
- word_embeddings = self._embedding_layer(word_ids)
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
if dense_inputs is not None:
- # Concat the dense embeddings at sequence begin so unpool_len can control
- # embedding not being pooled.
- word_embeddings = tf.concat([dense_inputs, word_embeddings], axis=1)
- type_ids = tf.concat([dense_type_ids, type_ids], axis=1)
- mask = tf.concat([dense_mask, mask], axis=1)
+ # Allow concatenation of the dense embeddings at sequence end if requested
+ # and `unpool_length`` is set as zero
+ if self._append_dense_inputs:
+ if self._unpool_length != 0:
+ raise ValueError(
+ 'unpool_length is not supported by append_dense_inputs now.'
+ )
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+ mask = tf.concat([mask, dense_mask], axis=1)
+ else:
+ # Concat the dense embeddings at sequence begin so unpool_len can
+ # control embedding not being pooled.
+ word_embeddings = tf.concat([dense_inputs, word_embeddings], axis=1)
+ type_ids = tf.concat([dense_type_ids, type_ids], axis=1)
+ mask = tf.concat([dense_mask, mask], axis=1)
# absolute position embeddings
position_embeddings = self._position_embedding_layer(word_embeddings)
type_embeddings = self._type_embedding_layer(type_ids)
- embeddings = tf.keras.layers.add(
+ embeddings = tf_keras.layers.add(
[word_embeddings, position_embeddings, type_embeddings])
embeddings = self._embedding_norm_layer(embeddings)
embeddings = self._embedding_dropout(embeddings)
@@ -462,13 +524,19 @@ def call(self, inputs):
attention_mask = _pool_and_concat(
attention_mask,
unpool_length=self._unpool_length,
- strides=self._pool_strides[0],
+ strides=self._pool_strides[0], # pyrefly: ignore[bad-argument-type]
axes=[1])
for i, layer in enumerate(self._transformer_layers):
+ transformer_output_range = None
+ if i == self._num_layers - 1:
+ transformer_output_range = output_range
+
# Bypass no pooling cases.
if self._pool_strides[i] == 1:
- x = layer([x, x, attention_mask])
+ x = layer(
+ [x, x, attention_mask], output_range=transformer_output_range
+ )
else:
# Pools layer for compressing the query length.
pooled_inputs = self._att_input_pool_layers[i](
@@ -478,35 +546,46 @@ def call(self, inputs):
x[:, :self._unpool_length, :],
dtype=pooled_inputs.dtype), pooled_inputs),
axis=1)
- x = layer([query_inputs, x, attention_mask])
+ x = layer([query_inputs, x, attention_mask],
+ output_range=transformer_output_range)
# Pools the corresponding attention_mask.
if i < len(self._transformer_layers) - 1:
attention_mask = _pool_and_concat(
attention_mask,
unpool_length=self._unpool_length,
- strides=[self._pool_strides[i + 1], self._pool_strides[i]],
+ strides=[self._pool_strides[i + 1], self._pool_strides[i]], # pyrefly: ignore[bad-argument-type]
axes=[1, 2])
encoder_outputs.append(x)
elif self._pool_type == _TRUNCATED_AVG:
- attention_masks = _create_truncated_avg_masks(mask, self._pool_strides,
- self._pooling_transforms)
+ # Compute the attention masks and pooling transforms.
+ # Note we do not compute this in __init__ due to inference converter issue
+ # b/215659399.
+ pooling_transforms = _create_truncated_avg_transforms(
+ self._max_sequence_length, self._pool_strides) # pyrefly: ignore[bad-argument-type]
+ attention_masks = _create_truncated_avg_masks(mask, self._pool_strides, # pyrefly: ignore[bad-argument-type]
+ pooling_transforms)
for i, layer in enumerate(self._transformer_layers):
attention_mask = attention_masks[i]
+ transformer_output_range = None
+ if i == self._num_layers - 1:
+ transformer_output_range = output_range
# Bypass no pooling cases.
if self._pool_strides[i] == 1:
- x = layer([x, x, attention_mask])
+ x = layer([x, x, attention_mask],
+ output_range=transformer_output_range)
else:
pooled_inputs = tf.einsum(
'BFD,FT->BTD',
tf.cast(x[:, self._unpool_length:, :], _get_policy_dtype()
), # extra casting for faster mixed computation.
- self._pooling_transforms[i])
+ pooling_transforms[i])
query_inputs = tf.concat(
values=(tf.cast(
x[:, :self._unpool_length, :],
dtype=pooled_inputs.dtype), pooled_inputs),
axis=1)
- x = layer([query_inputs, x, attention_mask])
+ x = layer([query_inputs, x, attention_mask],
+ output_range=transformer_output_range)
encoder_outputs.append(x)
last_encoder_output = encoder_outputs[-1]
diff --git a/official/nlp/modeling/networks/funnel_transformer_test.py b/official/nlp/modeling/networks/funnel_transformer_test.py
index 98dc3e1a971..5512d1a9b1a 100644
--- a/official/nlp/modeling/networks/funnel_transformer_test.py
+++ b/official/nlp/modeling/networks/funnel_transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,12 +16,12 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.networks import funnel_transformer
-class SingleLayerModel(tf.keras.Model):
+class SingleLayerModel(tf_keras.Model):
def __init__(self, layer):
super().__init__()
@@ -35,7 +35,7 @@ class FunnelTransformerEncoderTest(parameterized.TestCase, tf.test.TestCase):
def tearDown(self):
super(FunnelTransformerEncoderTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
@parameterized.named_parameters(
("mix_truncated_avg_rezero", "mixed_float16", tf.float16, "truncated_avg",
@@ -52,7 +52,7 @@ def tearDown(self):
("float32_avg", "float32", tf.float32, "avg", "TransformerEncoderBlock"))
def test_network_creation(self, policy, pooled_dtype, pool_type,
transformer_cls):
- tf.keras.mixed_precision.set_global_policy(policy)
+ tf_keras.mixed_precision.set_global_policy(policy)
hidden_size = 32
sequence_length = 21
@@ -70,16 +70,16 @@ def test_network_creation(self, policy, pooled_dtype, pool_type,
unpool_length=0,
transformer_cls=transformer_cls)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, num_layers)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
# Stride=2 compresses sequence length to half the size at each layer.
# For pool_type = max or avg,
@@ -101,8 +101,12 @@ def test_network_creation(self, policy, pooled_dtype, pool_type,
self.assertAllEqual(tf.float32, data.dtype)
self.assertAllEqual(pooled_dtype, pooled.dtype)
- def test_network_creation_dense(self):
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ @parameterized.named_parameters(
+ ("append_dense_inputs", True),
+ ("dense_inputs_at_sequence_begin", False),
+ )
+ def test_network_creation_dense(self, append_dense_inputs):
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
pool_type = "avg"
hidden_size = 32
@@ -120,16 +124,17 @@ def test_network_creation_dense(self):
pool_type=pool_type,
max_sequence_length=sequence_length + dense_sequence_length,
unpool_length=0,
- transformer_cls="TransformerEncoderBlock")
+ transformer_cls="TransformerEncoderBlock",
+ append_dense_inputs=append_dense_inputs)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
- dense_inputs = tf.keras.Input(
+ dense_inputs = tf_keras.Input(
shape=(dense_sequence_length, hidden_size), dtype=tf.float32)
- dense_mask = tf.keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
- dense_type_ids = tf.keras.Input(
+ dense_mask = tf_keras.Input(shape=(dense_sequence_length,), dtype=tf.int32)
+ dense_type_ids = tf_keras.Input(
shape=(dense_sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
@@ -139,7 +144,7 @@ def test_network_creation_dense(self):
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, num_layers)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
# Stride=2 compresses sequence length to half the size at each layer.
# For pool_type = max or avg,
@@ -150,6 +155,36 @@ def test_network_creation_dense(self):
self.assertAllEqual(expected_data_shape, data.shape.as_list())
self.assertAllEqual(expected_pooled_shape, pooled.shape.as_list())
+ @parameterized.named_parameters(
+ ("frac_pool_rezero", "ReZeroTransformer"),
+ ("frac_pool_vanilla", "TransformerEncoderBlock"),
+ )
+ def test_fractional_pooling(self, transformer_cls):
+ hidden_size = 16
+ sequence_length = 32
+ pool_strides = [1.33333, 3, 2, 1]
+ num_layers = 4
+ pool_type = "truncated_avg"
+ test_network = funnel_transformer.FunnelTransformerEncoder(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=2,
+ num_layers=num_layers,
+ pool_stride=pool_strides,
+ pool_type=pool_type,
+ max_sequence_length=sequence_length,
+ unpool_length=0,
+ transformer_cls=transformer_cls)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = test_network([word_ids, mask, type_ids])
+ data = dict_outputs["sequence_output"]
+
+ expected_data_shape = [None, 4, hidden_size]
+
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+
def test_invalid_stride_and_num_layers(self):
hidden_size = 32
num_layers = 3
@@ -186,9 +221,9 @@ def test_all_encoder_outputs_network_creation(self, pool_stride,
pool_stride=pool_stride,
unpool_length=unpool_length)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network([word_ids, mask, type_ids])
all_encoder_outputs = dict_outputs["encoder_outputs"]
pooled = dict_outputs["pooled_output"]
@@ -210,38 +245,46 @@ def test_all_encoder_outputs_network_creation(self, pool_stride,
self.assertAllEqual(tf.float32, pooled.dtype)
@parameterized.named_parameters(
- ("all_sequence", None, 3, 0),
- ("output_range", 1, 1, 0),
- ("all_sequence_wit_unpool", None, 4, 1),
- ("output_range_with_unpool", 1, 1, 1),
- ("output_range_with_large_unpool", 1, 1, 2),
+ ("all_sequence", None, 3, 0, 2),
+ ("output_range", 1, 1, 0, 2),
+ ("all_sequence_with_unpool", None, 4, 1, 2),
+ ("output_range_with_unpool", 1, 1, 1, 2),
+ ("output_range_with_large_unpool", 1, 1, 2, 2),
+ ("output_range_with_no_pooling", 1, 1, 0, 1),
+ ("output_range_with_unpool_and_no_pooling", 1, 1, 1, 1),
)
- def test_network_invocation(self, output_range, out_seq_len, unpool_length):
+ def test_network_invocation(
+ self,
+ output_range,
+ out_seq_len,
+ unpool_length,
+ pool_stride,
+ ):
hidden_size = 32
sequence_length = 21
vocab_size = 57
num_types = 7
- pool_stride = 2
+ num_layers = 3
# Create a small FunnelTransformerEncoder for testing.
test_network = funnel_transformer.FunnelTransformerEncoder(
vocab_size=vocab_size,
hidden_size=hidden_size,
num_attention_heads=2,
- num_layers=3,
+ num_layers=num_layers,
type_vocab_size=num_types,
- output_range=output_range,
pool_stride=pool_stride,
unpool_length=unpool_length)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- dict_outputs = test_network([word_ids, mask, type_ids])
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = test_network([word_ids, mask, type_ids],
+ output_range=output_range)
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -262,15 +305,18 @@ def test_network_invocation(self, output_range, out_seq_len, unpool_length):
hidden_size=hidden_size,
max_sequence_length=max_sequence_length,
num_attention_heads=2,
- num_layers=3,
+ num_layers=num_layers,
type_vocab_size=num_types,
pool_stride=pool_stride)
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
- self.assertEqual(outputs[0].shape[1], 3)
+ expected_sequence_length = float(sequence_length)
+ for _ in range(num_layers):
+ expected_sequence_length = np.ceil(expected_sequence_length / pool_stride)
+ self.assertEqual(outputs[0].shape[1], expected_sequence_length)
# Creates a FunnelTransformerEncoder with embedding_width != hidden_size
test_network = funnel_transformer.FunnelTransformerEncoder(
@@ -285,11 +331,48 @@ def test_network_invocation(self, output_range, out_seq_len, unpool_length):
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[-1], hidden_size)
self.assertTrue(hasattr(test_network, "_embedding_projection"))
+ def test_embeddings_as_inputs(self):
+ hidden_size = 32
+ sequence_length = 21
+ # Create a small BertEncoder for testing.
+ test_network = funnel_transformer.FunnelTransformerEncoder(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=2,
+ num_layers=3,
+ pool_stride=2,
+ )
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ test_network.build(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids)
+ )
+ embeddings = test_network.get_embedding_layer()(word_ids)
+ # Calls with the embeddings.
+ dict_outputs = test_network(
+ dict(
+ input_word_embeddings=embeddings,
+ input_mask=mask,
+ input_type_ids=type_ids,
+ )
+ )
+ all_encoder_outputs = dict_outputs["encoder_outputs"]
+ pooled = dict_outputs["pooled_output"]
+
+ expected_pooled_shape = [None, hidden_size]
+ self.assertAllEqual(expected_pooled_shape, pooled.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, all_encoder_outputs[-1].dtype)
+ self.assertAllEqual(tf.float32, pooled.dtype)
+
def test_serialize_deserialize(self):
# Create a network object that sets all of its config options.
kwargs = dict(
@@ -314,10 +397,10 @@ def test_serialize_deserialize(self):
transformer_cls="TransformerEncoderBlock")
network = funnel_transformer.FunnelTransformerEncoder(**kwargs)
expected_config = dict(kwargs)
- expected_config["inner_activation"] = tf.keras.activations.serialize(
- tf.keras.activations.get(expected_config["inner_activation"]))
- expected_config["initializer"] = tf.keras.initializers.serialize(
- tf.keras.initializers.get(expected_config["initializer"]))
+ expected_config["inner_activation"] = tf_keras.activations.serialize(
+ tf_keras.activations.get(expected_config["inner_activation"]))
+ expected_config["initializer"] = tf_keras.initializers.serialize(
+ tf_keras.initializers.get(expected_config["initializer"]))
self.assertEqual(network.get_config(), expected_config)
# Create another network object from the first object's config.
new_network = funnel_transformer.FunnelTransformerEncoder.from_config(
@@ -342,7 +425,7 @@ def test_serialize_deserialize(self):
_ = network_wrapper.predict([word_id_data, mask_data, type_id_data])
network_wrapper.save(model_path)
- _ = tf.keras.models.load_model(model_path)
+ _ = tf_keras.models.load_model(model_path)
if __name__ == "__main__":
diff --git a/official/nlp/modeling/networks/mobile_bert_encoder.py b/official/nlp/modeling/networks/mobile_bert_encoder.py
index 710691a70fb..65af4d5de31 100644
--- a/official/nlp/modeling/networks/mobile_bert_encoder.py
+++ b/official/nlp/modeling/networks/mobile_bert_encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,13 @@
"""MobileBERT text encoder network."""
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
@gin.configurable
-class MobileBERTEncoder(tf.keras.Model):
+class MobileBERTEncoder(tf_keras.Model):
"""A Keras functional API implementation for MobileBERT encoder."""
def __init__(self,
@@ -84,7 +84,7 @@ def __init__(self,
**kwargs: Other keyworded and arguments.
"""
self._self_setattr_tracking = False
- initializer = tf.keras.initializers.TruncatedNormal(
+ initializer = tf_keras.initializers.TruncatedNormal(
stddev=initializer_range)
# layer instantiation
@@ -117,11 +117,11 @@ def __init__(self,
self._transformer_layers.append(transformer)
# input tensor
- input_ids = tf.keras.layers.Input(
+ input_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(None,), dtype=input_mask_dtype, name='input_mask')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
self.inputs = [input_ids, input_mask, type_ids]
@@ -146,7 +146,7 @@ def __init__(self,
first_token = tf.squeeze(prev_output[:, 0:1, :], axis=1)
if classifier_activation:
- self._pooler_layer = tf.keras.layers.experimental.EinsumDense(
+ self._pooler_layer = tf_keras.layers.EinsumDense(
'ab,bc->ac',
output_shape=hidden_size,
activation=tf.tanh,
@@ -162,9 +162,39 @@ def __init__(self,
pooled_output=first_token,
encoder_outputs=all_layer_outputs,
attention_scores=all_attention_scores)
-
- super(MobileBERTEncoder, self).__init__(
+ super().__init__(
inputs=self.inputs, outputs=outputs, **kwargs)
+ self._config = dict(
+ word_vocab_size=word_vocab_size,
+ word_embed_size=word_embed_size,
+ type_vocab_size=type_vocab_size,
+ max_sequence_length=max_sequence_length,
+ num_blocks=num_blocks,
+ hidden_size=hidden_size,
+ num_attention_heads=num_attention_heads,
+ intermediate_size=intermediate_size,
+ intermediate_act_fn=intermediate_act_fn,
+ hidden_dropout_prob=hidden_dropout_prob,
+ attention_probs_dropout_prob=attention_probs_dropout_prob,
+ intra_bottleneck_size=intra_bottleneck_size,
+ initializer_range=initializer_range,
+ use_bottleneck_attention=use_bottleneck_attention,
+ key_query_shared_bottleneck=key_query_shared_bottleneck,
+ num_feedforward_networks=num_feedforward_networks,
+ normalization_type=normalization_type,
+ classifier_activation=classifier_activation,
+ input_mask_dtype=input_mask_dtype,
+ **kwargs,
+ )
+ if 'name' not in self._config:
+ self._config['name'] = self.name
+
+ def get_config(self):
+ return dict(self._config)
+
+ @classmethod
+ def from_config(cls, config): # pyrefly: ignore[bad-override]
+ return cls(**config)
def get_embedding_table(self):
return self.embedding_layer.word_embedding.embeddings
diff --git a/official/nlp/modeling/networks/mobile_bert_encoder_test.py b/official/nlp/modeling/networks/mobile_bert_encoder_test.py
index 1b119005b32..6c8ed75d21f 100644
--- a/official/nlp/modeling/networks/mobile_bert_encoder_test.py
+++ b/official/nlp/modeling/networks/mobile_bert_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import models
from official.nlp.modeling.networks import mobile_bert_encoder
@@ -55,9 +55,9 @@ def test_mobilebert_encoder(self, act_fn, kq_shared_bottleneck,
normalization_type=normalization_type,
classifier_activation=use_pooler)
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
outputs = test_network([word_ids, mask, type_ids])
layer_output, pooler_output = outputs['sequence_output'], outputs[
'pooled_output']
@@ -80,9 +80,9 @@ def test_mobilebert_encoder_return_all_layer_output(self):
hidden_size=hidden_size,
num_blocks=num_blocks)
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
outputs = test_network([word_ids, mask, type_ids])
all_layer_output = outputs['encoder_outputs']
@@ -101,11 +101,11 @@ def test_mobilebert_encoder_invocation(self, input_mask_dtype):
num_blocks=num_blocks,
input_mask_dtype=input_mask_dtype)
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=input_mask_dtype)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=input_mask_dtype)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
outputs = test_network([word_ids, mask, type_ids])
- model = tf.keras.Model([word_ids, mask, type_ids], outputs)
+ model = tf_keras.Model([word_ids, mask, type_ids], outputs)
input_seq = generate_fake_input(
batch_size=1, seq_len=sequence_length, vocab_size=vocab_size)
@@ -130,11 +130,11 @@ def test_mobilebert_encoder_invocation_with_attention_score(self):
hidden_size=hidden_size,
num_blocks=num_blocks)
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
outputs = test_network([word_ids, mask, type_ids])
- model = tf.keras.Model([word_ids, mask, type_ids], outputs)
+ model = tf_keras.Model([word_ids, mask, type_ids], outputs)
input_seq = generate_fake_input(
batch_size=1, seq_len=sequence_length, vocab_size=vocab_size)
@@ -156,9 +156,9 @@ def test_mobilebert_encoder_for_downstream_task(self, task, prediction_shape):
num_classes = 5
classifier = task(network=mobilebert_encoder, num_classes=num_classes)
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
prediction = classifier([word_ids, mask, type_ids])
if task == models.BertTokenClassifier:
prediction = prediction['logits']
diff --git a/official/nlp/modeling/networks/packed_sequence_embedding.py b/official/nlp/modeling/networks/packed_sequence_embedding.py
index 18ff59f4f42..7cd8ef7bff1 100644
--- a/official/nlp/modeling/networks/packed_sequence_embedding.py
+++ b/official/nlp/modeling/networks/packed_sequence_embedding.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,14 +15,14 @@
"""An embedding network supporting packed sequences and position ids."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.nlp.modeling import layers
-@tf.keras.utils.register_keras_serializable(package='Text')
-class PackedSequenceEmbedding(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class PackedSequenceEmbedding(tf_keras.Model):
"""An embedding network supporting packed sequences and position ids.
This network implements an embedding layer similar to the one described in
@@ -60,7 +60,7 @@ def __init__(self,
use_position_id=False,
pack_multiple_sequences=False,
**kwargs):
- initializer = tf.keras.initializers.get(initializer)
+ initializer = tf_keras.initializers.get(initializer)
if embedding_width is None:
embedding_width = hidden_size
config_dict = {
@@ -69,21 +69,21 @@ def __init__(self,
'embedding_width': embedding_width,
'hidden_size': hidden_size,
'max_seq_length': max_seq_length,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_keras.initializers.serialize(initializer),
'dropout_rate': dropout_rate,
'use_position_id': use_position_id,
'pack_multiple_sequences': pack_multiple_sequences,
}
- word_ids = tf.keras.layers.Input(
+ word_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- mask = tf.keras.layers.Input(
+ mask = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_mask')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
inputs = [word_ids, mask, type_ids]
if use_position_id:
- position_ids = tf.keras.layers.Input(
+ position_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='position_ids')
inputs.append(position_ids)
else:
@@ -97,13 +97,13 @@ def __init__(self,
embedding_layer = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
name='word_embeddings')
word_embeddings = embedding_layer(word_ids)
# Always uses dynamic slicing for simplicity.
position_embedding_layer = PositionEmbeddingWithSubSeqMask(
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_dynamic_slicing=True,
max_sequence_length=max_seq_length,
name='position_embedding')
@@ -114,40 +114,40 @@ def __init__(self,
layers.OnDeviceEmbedding(
vocab_size=type_vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_one_hot=True,
name='type_embeddings')(type_ids))
- embeddings = tf.keras.layers.Add()(
+ embeddings = tf_keras.layers.Add()(
[word_embeddings, position_embeddings, type_embeddings])
- embeddings = tf.keras.layers.LayerNormalization(
+ embeddings = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)(
embeddings)
- embeddings = tf.keras.layers.Dropout(
+ embeddings = tf_keras.layers.Dropout(
rate=dropout_rate, dtype=tf.float32)(
embeddings)
if embedding_width != hidden_size:
- embeddings = tf.keras.layers.experimental.EinsumDense(
+ embeddings = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes=None,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='embedding_projection')(
embeddings)
attention_mask = layers.SelfAttentionMask()(embeddings, mask)
if sub_seq_mask is not None:
- attention_mask = tf.keras.layers.Lambda(
+ attention_mask = tf_keras.layers.Lambda(
lambda x: x[0] * tf.cast(x[1], x[0].dtype))(
[attention_mask, sub_seq_mask])
outputs = [embeddings, attention_mask]
- super(PackedSequenceEmbedding, self).__init__(
+ super().__init__(
inputs=inputs, outputs=outputs, **kwargs)
# TF does not track immutable attrs which do not contain Trackables,
# so by creating a config namedtuple instead of a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self._embedding_layer = embedding_layer
self._position_embedding_layer = position_embedding_layer
@@ -163,8 +163,8 @@ def from_config(cls, config, custom_objects=None):
return cls(**config)
-@tf.keras.utils.register_keras_serializable(package='Text')
-class PackedSequenceMask(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class PackedSequenceMask(tf_keras.layers.Layer):
"""A layer to create a mask to indicate multiple sub sequences."""
def call(self, input_ids):
@@ -189,8 +189,8 @@ def call(self, input_ids):
return tf.equal(seq_ids, tf.transpose(seq_ids, [0, 2, 1]))
-@tf.keras.utils.register_keras_serializable(package='Text')
-class PositionEmbeddingWithSubSeqMask(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class PositionEmbeddingWithSubSeqMask(tf_keras.layers.Layer):
"""Creates a positional embedding with sub-sequence masking.
This layer creates a positional embedding as described in "BERT: Pre-training
@@ -221,22 +221,22 @@ def __init__(self,
if 'dtype' not in kwargs:
kwargs['dtype'] = 'float32'
- super(PositionEmbeddingWithSubSeqMask, self).__init__(**kwargs)
+ super().__init__(**kwargs)
if use_dynamic_slicing and max_sequence_length is None:
raise ValueError(
'If `use_dynamic_slicing` is True, `max_sequence_length` must be set.'
)
self._max_sequence_length = max_sequence_length
- self._initializer = tf.keras.initializers.get(initializer)
+ self._initializer = tf_keras.initializers.get(initializer)
self._use_dynamic_slicing = use_dynamic_slicing
def get_config(self):
config = {
'max_sequence_length': self._max_sequence_length,
- 'initializer': tf.keras.initializers.serialize(self._initializer),
+ 'initializer': tf_keras.initializers.serialize(self._initializer),
'use_dynamic_slicing': self._use_dynamic_slicing,
}
- base_config = super(PositionEmbeddingWithSubSeqMask, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def build(self, input_shape):
@@ -273,7 +273,7 @@ def build(self, input_shape):
shape=[weight_sequence_length, width],
initializer=self._initializer)
- super(PositionEmbeddingWithSubSeqMask, self).build(input_shape)
+ super().build(input_shape)
def call(self, inputs, position_ids=None, sub_sequence_mask=None):
"""Implements call() for the layer.
diff --git a/official/nlp/modeling/networks/packed_sequence_embedding_test.py b/official/nlp/modeling/networks/packed_sequence_embedding_test.py
index 64080f3c8f2..3b489e1fd1e 100644
--- a/official/nlp/modeling/networks/packed_sequence_embedding_test.py
+++ b/official/nlp/modeling/networks/packed_sequence_embedding_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,9 @@
"""Tests for official.nlp.modeling.networks.packed_sequence_embedding."""
-# Import libraries
-
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.networks import packed_sequence_embedding
@@ -27,7 +25,7 @@ class PackedSequenceEmbeddingTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(PackedSequenceEmbeddingTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
@parameterized.parameters([
(True, True, True),
@@ -39,7 +37,7 @@ def test_network_creation(self, use_position_id, pack_multiple_sequences,
use_float16):
"""Validate that the Keras object can be created."""
if use_float16:
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
seq_length = 16
vocab_size = 100
max_position_embeddings = 32
@@ -52,7 +50,7 @@ def test_network_creation(self, use_position_id, pack_multiple_sequences,
embedding_width=embedding_width,
hidden_size=hidden_size,
max_seq_length=max_position_embeddings,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
dropout_rate=0.1,
use_position_id=use_position_id,
pack_multiple_sequences=pack_multiple_sequences,
@@ -60,22 +58,22 @@ def test_network_creation(self, use_position_id, pack_multiple_sequences,
test_object = packed_sequence_embedding.PackedSequenceEmbedding(
**embedding_cfg)
- input_word_ids = tf.keras.Input(shape=(seq_length,), dtype=tf.int32)
- input_mask = tf.keras.Input(shape=(seq_length,), dtype=tf.int32)
- input_type_ids = tf.keras.Input(shape=(seq_length,), dtype=tf.int32)
+ input_word_ids = tf_keras.Input(shape=(seq_length,), dtype=tf.int32)
+ input_mask = tf_keras.Input(shape=(seq_length,), dtype=tf.int32)
+ input_type_ids = tf_keras.Input(shape=(seq_length,), dtype=tf.int32)
network_inputs = {
'input_word_ids': input_word_ids,
'input_mask': input_mask,
'input_type_ids': input_type_ids,
}
if use_position_id:
- network_inputs['position_ids'] = tf.keras.Input(
+ network_inputs['position_ids'] = tf_keras.Input(
shape=(seq_length,), dtype=tf.int32)
embedding, mask = test_object(network_inputs)
# Create a model based off of this network:
- model = tf.keras.Model(network_inputs, [embedding, mask])
+ model = tf_keras.Model(network_inputs, [embedding, mask])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -99,7 +97,7 @@ def test_network_creation(self, use_position_id, pack_multiple_sequences,
self.assertAllEqual(expected_attention_mask_shape, attention_mask.shape)
def test_serialize_deserialize(self):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
# Create a network object that sets all of its config options.
embedding_cfg = dict(
vocab_size=100,
@@ -107,7 +105,7 @@ def test_serialize_deserialize(self):
embedding_width=64,
hidden_size=64,
max_seq_length=32,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
dropout_rate=0.1,
use_position_id=True,
pack_multiple_sequences=False,
@@ -115,8 +113,8 @@ def test_serialize_deserialize(self):
network = packed_sequence_embedding.PackedSequenceEmbedding(**embedding_cfg)
expected_config = dict(embedding_cfg)
- expected_config['initializer'] = tf.keras.initializers.serialize(
- tf.keras.initializers.get(expected_config['initializer']))
+ expected_config['initializer'] = tf_keras.initializers.serialize(
+ tf_keras.initializers.get(expected_config['initializer']))
self.assertEqual(network.get_config(), expected_config)
# Create another network object from the first object's config.
diff --git a/official/nlp/modeling/networks/span_labeling.py b/official/nlp/modeling/networks/span_labeling.py
index 6dc73d3abf3..4adebbd89d5 100644
--- a/official/nlp/modeling/networks/span_labeling.py
+++ b/official/nlp/modeling/networks/span_labeling.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,9 @@
"""Span labeling network."""
# pylint: disable=g-classes-have-attributes
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
def _apply_paragraph_mask(logits, paragraph_mask):
@@ -24,8 +26,8 @@ def _apply_paragraph_mask(logits, paragraph_mask):
return tf.nn.log_softmax(masked_logits, -1), masked_logits
-@tf.keras.utils.register_keras_serializable(package='Text')
-class SpanLabeling(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class SpanLabeling(tf_keras.Model):
"""Span labeling network head for BERT modeling.
This network implements a simple single-span labeler based on a dense layer.
@@ -48,10 +50,10 @@ def __init__(self,
output='logits',
**kwargs):
- sequence_data = tf.keras.layers.Input(
+ sequence_data = tf_keras.layers.Input(
shape=(None, input_width), name='sequence_data', dtype=tf.float32)
- intermediate_logits = tf.keras.layers.Dense(
+ intermediate_logits = tf_keras.layers.Dense(
2, # This layer predicts start location and end location.
activation=activation,
kernel_initializer=initializer,
@@ -59,9 +61,9 @@ def __init__(self,
sequence_data)
start_logits, end_logits = self._split_output_tensor(intermediate_logits)
- start_predictions = tf.keras.layers.Activation(tf.nn.log_softmax)(
+ start_predictions = tf_keras.layers.Activation(tf.nn.log_softmax)(
start_logits)
- end_predictions = tf.keras.layers.Activation(tf.nn.log_softmax)(end_logits)
+ end_predictions = tf_keras.layers.Activation(tf.nn.log_softmax)(end_logits)
if output == 'logits':
output_tensors = [start_logits, end_logits]
@@ -79,7 +81,7 @@ def __init__(self,
# created using the Functional API. Once super().__init__ is called, we
# can assign attributes to `self` - note that all `self` assignments are
# below this line.
- super(SpanLabeling, self).__init__(
+ super().__init__(
inputs=[sequence_data], outputs=output_tensors, **kwargs)
config_dict = {
'input_width': input_width,
@@ -92,7 +94,7 @@ def __init__(self,
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self.start_logits = start_logits
self.end_logits = end_logits
@@ -109,7 +111,7 @@ def from_config(cls, config, custom_objects=None):
return cls(**config)
-class XLNetSpanLabeling(tf.keras.layers.Layer):
+class XLNetSpanLabeling(tf_keras.layers.Layer):
"""Span labeling network head for XLNet on SQuAD2.0.
This networks implements a span-labeler based on dense layers and question
@@ -154,33 +156,33 @@ def __init__(self,
raise ValueError('`start_n_top` must be greater than 1.')
self._start_n_top = start_n_top
self._end_n_top = end_n_top
- self.start_logits_dense = tf.keras.layers.Dense(
+ self.start_logits_dense = tf_keras.layers.Dense(
units=1,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='predictions/transform/start_logits')
- self.end_logits_inner_dense = tf.keras.layers.Dense(
+ self.end_logits_inner_dense = tf_keras.layers.Dense(
units=input_width,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
activation=activation,
name='predictions/transform/end_logits/inner')
- self.end_logits_layer_norm = tf.keras.layers.LayerNormalization(
+ self.end_logits_layer_norm = tf_keras.layers.LayerNormalization(
axis=-1, epsilon=1e-12,
name='predictions/transform/end_logits/layernorm')
- self.end_logits_output_dense = tf.keras.layers.Dense(
+ self.end_logits_output_dense = tf_keras.layers.Dense(
units=1,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='predictions/transform/end_logits/output')
- self.answer_logits_inner = tf.keras.layers.Dense(
+ self.answer_logits_inner = tf_keras.layers.Dense(
units=input_width,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
activation=activation,
name='predictions/transform/answer_logits/inner')
- self.answer_logits_dropout = tf.keras.layers.Dropout(rate=dropout_rate)
- self.answer_logits_output = tf.keras.layers.Dense(
+ self.answer_logits_dropout = tf_keras.layers.Dropout(rate=dropout_rate)
+ self.answer_logits_output = tf_keras.layers.Dense(
units=1,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
use_bias=False,
name='predictions/transform/answer_logits/output')
@@ -309,8 +311,8 @@ def call(self,
end_top_index = tf.reshape(
end_top_index,
[-1, self._start_n_top * self._end_n_top])
- output_dict['start_top_predictions'] = start_top_predictions
- output_dict['start_top_index'] = start_top_index
+ output_dict['start_top_predictions'] = start_top_predictions # pyrefly: ignore[unbound-name]
+ output_dict['start_top_index'] = start_top_index # pyrefly: ignore[unbound-name]
output_dict['end_top_predictions'] = end_top_predictions
output_dict['end_top_index'] = end_top_index
diff --git a/official/nlp/modeling/networks/span_labeling_test.py b/official/nlp/modeling/networks/span_labeling_test.py
index a51a0a7c6ec..9bf2fd7ef4a 100644
--- a/official/nlp/modeling/networks/span_labeling_test.py
+++ b/official/nlp/modeling/networks/span_labeling_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,16 +14,12 @@
"""Tests for span_labeling network."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.networks import span_labeling
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class SpanLabelingTest(keras_parameterized.TestCase):
+class SpanLabelingTest(tf.test.TestCase):
def test_network_creation(self):
"""Validate that the Keras object can be created."""
@@ -32,7 +28,7 @@ def test_network_creation(self):
test_network = span_labeling.SpanLabeling(
input_width=input_width, output='predictions')
# Create a 3-dimensional input (the first dimension is implicit).
- sequence_data = tf.keras.Input(
+ sequence_data = tf_keras.Input(
shape=(sequence_length, input_width), dtype=tf.float32)
start_outputs, end_outputs = test_network(sequence_data)
@@ -48,10 +44,10 @@ def test_network_invocation(self):
test_network = span_labeling.SpanLabeling(input_width=input_width)
# Create a 3-dimensional input (the first dimension is implicit).
- sequence_data = tf.keras.Input(
+ sequence_data = tf_keras.Input(
shape=(sequence_length, input_width), dtype=tf.float32)
outputs = test_network(sequence_data)
- model = tf.keras.Model(sequence_data, outputs)
+ model = tf_keras.Model(sequence_data, outputs)
# Invoke the network as part of a Model.
batch_size = 3
@@ -71,11 +67,11 @@ def test_network_invocation_with_internal_logit_output(self):
test_network = span_labeling.SpanLabeling(
input_width=input_width, output='predictions')
# Create a 3-dimensional input (the first dimension is implicit).
- sequence_data = tf.keras.Input(
+ sequence_data = tf_keras.Input(
shape=(sequence_length, input_width), dtype=tf.float32)
output = test_network(sequence_data)
- model = tf.keras.Model(sequence_data, output)
- logit_model = tf.keras.Model(
+ model = tf_keras.Model(sequence_data, output)
+ logit_model = tf_keras.Model(
test_network.inputs,
[test_network.start_logits, test_network.end_logits])
@@ -93,9 +89,9 @@ def test_network_invocation_with_internal_logit_output(self):
self.assertEqual(expected_output_shape, end_logits.shape)
# Ensure that the logits, when softmaxed, create the outputs.
- input_tensor = tf.keras.Input(expected_output_shape[1:])
- output_tensor = tf.keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
- softmax_model = tf.keras.Model(input_tensor, output_tensor)
+ input_tensor = tf_keras.Input(expected_output_shape[1:])
+ output_tensor = tf_keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
+ softmax_model = tf_keras.Model(input_tensor, output_tensor)
start_softmax = softmax_model.predict(start_logits)
self.assertAllClose(start_outputs, start_softmax)
@@ -113,12 +109,12 @@ def test_network_invocation_with_external_logit_output(self):
logit_network.set_weights(test_network.get_weights())
# Create a 3-dimensional input (the first dimension is implicit).
- sequence_data = tf.keras.Input(
+ sequence_data = tf_keras.Input(
shape=(sequence_length, input_width), dtype=tf.float32)
output = test_network(sequence_data)
logit_output = logit_network(sequence_data)
- model = tf.keras.Model(sequence_data, output)
- logit_model = tf.keras.Model(sequence_data, logit_output)
+ model = tf_keras.Model(sequence_data, output)
+ logit_model = tf_keras.Model(sequence_data, logit_output)
batch_size = 3
input_data = 10 * np.random.random_sample(
@@ -134,9 +130,9 @@ def test_network_invocation_with_external_logit_output(self):
self.assertEqual(expected_output_shape, end_logits.shape)
# Ensure that the logits, when softmaxed, create the outputs.
- input_tensor = tf.keras.Input(expected_output_shape[1:])
- output_tensor = tf.keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
- softmax_model = tf.keras.Model(input_tensor, output_tensor)
+ input_tensor = tf_keras.Input(expected_output_shape[1:])
+ output_tensor = tf_keras.layers.Activation(tf.nn.log_softmax)(input_tensor)
+ softmax_model = tf_keras.Model(input_tensor, output_tensor)
start_softmax = softmax_model.predict(start_logits)
self.assertAllClose(start_outputs, start_softmax)
@@ -165,8 +161,7 @@ def test_unknown_output_type_fails(self):
_ = span_labeling.SpanLabeling(input_width=10, output='bad')
-@keras_parameterized.run_all_keras_modes
-class XLNetSpanLabelingTest(keras_parameterized.TestCase):
+class XLNetSpanLabelingTest(tf.test.TestCase):
def test_basic_invocation_train(self):
batch_size = 2
@@ -233,11 +228,11 @@ def test_subclass_invocation(self):
hidden_size = 4
batch_size = 2
- sequence_data = tf.keras.Input(shape=(seq_length, hidden_size),
+ sequence_data = tf_keras.Input(shape=(seq_length, hidden_size),
dtype=tf.float32)
- class_index = tf.keras.Input(shape=(), dtype=tf.uint8)
- paragraph_mask = tf.keras.Input(shape=(seq_length), dtype=tf.float32)
- start_positions = tf.keras.Input(shape=(), dtype=tf.int32)
+ class_index = tf_keras.Input(shape=(), dtype=tf.uint8)
+ paragraph_mask = tf_keras.Input(shape=(seq_length), dtype=tf.float32)
+ start_positions = tf_keras.Input(shape=(), dtype=tf.int32)
layer = span_labeling.XLNetSpanLabeling(
input_width=hidden_size,
@@ -251,7 +246,7 @@ def test_subclass_invocation(self):
class_index=class_index,
paragraph_mask=paragraph_mask,
start_positions=start_positions)
- model = tf.keras.Model(
+ model = tf_keras.Model(
inputs={
'sequence_data': sequence_data,
'class_index': class_index,
@@ -282,8 +277,8 @@ def test_subclass_invocation(self):
# Test `call` with training flag.
# Note: this fails due to incompatibility with the functional API.
- with self.assertRaisesRegexp(AssertionError,
- 'Could not compute output KerasTensor'):
+ with self.assertRaisesRegex(AssertionError,
+ 'Could not compute output KerasTensor'):
model(inputs, training=True)
def test_serialize_deserialize(self):
diff --git a/official/nlp/modeling/networks/sparse_mixer.py b/official/nlp/modeling/networks/sparse_mixer.py
new file mode 100644
index 00000000000..35818bd0c41
--- /dev/null
+++ b/official/nlp/modeling/networks/sparse_mixer.py
@@ -0,0 +1,406 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Sparse Mixer encoder network.
+
+Based on ["Sparse Mixers: Combining MoE and Mixing to build a more efficient
+BERT"](https://arxiv.org/abs/2205.12399).
+"""
+# pylint: disable=g-classes-have-attributes
+
+from typing import Any, Callable, Optional, Sequence, Union
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.nlp.modeling import layers
+
+_Activation = Union[str, Callable[..., Any]]
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
+
+
+class SparseMixer(tf_keras.layers.Layer):
+ """Sparse Mixer encoder network.
+
+ Based on ["Sparse Mixers: Combining MoE and Mixing to build a more efficient
+ BERT"](https://arxiv.org/abs/2205.12399). Sparse Mixer is an efficient
+ encoder network that replaces typical Transformer encoder blocks with a
+ combination of linear mixing and sparsely activated Mixture-of-Experts (MoE)
+ sublayers.
+
+ This implementation defaults to the canonical Sparse Mixer Base model. To use
+ the "Fast Sparse Mixer" configuration, set `*_capacity_factor`=0.5. This
+ yields a sparser and faster variant of the canonical Sparse Mixer model, in
+ which each expert processes roughly 50% less tokens.
+
+ Notes:
+ - The underlying MoeLayer uses the Keras add_loss() and add_metric() APIs to
+ propagate auxiliary MoE losses and metrics. Any model using this network,
+ should collect these losses and, if desired, metrics.
+ - The input length is fixed to 'max_sequence_length' to accomodate the mixing
+ mechanisms.
+
+ Args:
+ vocab_size: The size of the token vocabulary.
+ hidden_size: The size of the transformer hidden layers.
+ num_layers: The number of transformer layers.
+ moe_layers: Specifies which layers, if any, should be sparsely activated
+ Mixture-of-Experts (MoE) layers. The remaining [0, num_layers) setminus
+ moe_layers will use the vanilla MLP sublayers. Defaults to placing MoE
+ layers in the middle of the model.
+ attention_layers: Specifies which layers, if any, should be attention layers
+ in the encoder. The remaining [0, num_layers) setminus attention_layers
+ will use the specified `mixing_mechanism`. If using attention layers, a
+ good rule of thumb is to place them in the final few layers.
+ num_experts: Number of experts. Experts are themselves MLP modules, with the
+ same `inner_dim` and `inner_activation` as the vanilla MLP sublayers.
+ train_capacity_factor: Scaling factor to increase the expert token capacity
+ during training. See layers.MoeLayer for further details. The "Fast Sparse
+ Mixer" increases model sparsity (and speed) by using a capacity factor of
+ 0.5.
+ eval_capacity_factor: As above, but used during evaluation.
+ max_group_size: The total number of tokens on each device is subdivided into
+ groups of this size. Router computations are then performed on a per-group
+ basis. See layers.MoeLayer for further details.
+ mixing_mechanism: Type of mixing mechanism used in place of self-attention
+ layers. Defaults to 'Linear' mixing.
+ use_fft: Only used for spectral mixing mechanisms. Determines whether to use
+ Fast Fourier Transform (True) or the Discrete Fourier Transform (DFT)
+ matrix (False; default) to compute the Fourier Transform. See
+ layers.FourierTransformLayer or layers.HartleyTransformLayer for advice.
+ num_attention_heads: The number of attention heads for each transformer. The
+ hidden size must be divisible by the number of attention heads.
+ max_sequence_length: The only sequence length that this encoder can consume.
+ This determines the variable shape for positional embeddings and the size
+ of the mixing matrices.
+ type_vocab_size: The number of types that the 'type_ids' input can take.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ output_dropout: Dropout probability for the post-attention and output
+ dropout.
+ attention_dropout: The dropout rate to use for the attention layers within
+ the transformer layers.
+ initializer: The initializer to use for all weights in this encoder.
+ output_range: The sequence output range, [0, output_range), by slicing the
+ target sequence of the last transformer layer. `None` means the entire
+ target sequence will attend to the source sequence, which yields the full
+ output.
+ embedding_width: The width of the word embeddings. If the embedding width is
+ not equal to hidden size, embedding parameters will be factorized into two
+ matrices in the shape of ['vocab_size', 'embedding_width'] and
+ ['embedding_width', 'hidden_size'] ('embedding_width' is usually much
+ smaller than 'hidden_size').
+ embedding_layer: An optional Layer instance which will be called to generate
+ embeddings for the input word IDs.
+ norm_first: Whether to normalize inputs to attention and intermediate dense
+ layers. If set False, output of attention and intermediate dense layers is
+ normalized.
+ with_dense_inputs: Whether to accept dense embeddings as the input.
+ export_metrics: Whether to export metrics using Keras add_metric API.
+ """
+
+ def __init__(
+ self,
+ vocab_size: int,
+ hidden_size: int = 512,
+ num_layers: int = 14,
+ moe_layers: Sequence[int] = (5, 6, 7, 8),
+ attention_layers: Sequence[int] = (10, 11, 12, 13),
+ num_experts: int = 16,
+ train_capacity_factor: float = 1.,
+ eval_capacity_factor: float = 1.,
+ examples_per_group: float = 1.,
+ mixing_mechanism: layers.MixingMechanism = layers.MixingMechanism.LINEAR,
+ use_fft: bool = False,
+ num_attention_heads: int = 8,
+ max_sequence_length: int = 512,
+ type_vocab_size: int = 16,
+ inner_dim: int = 2048,
+ inner_activation: _Activation = _approx_gelu,
+ output_dropout: float = 0.1,
+ attention_dropout: float = 0.1,
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
+ stddev=0.02),
+ output_range: Optional[int] = None,
+ embedding_width: Optional[int] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
+ norm_first: bool = False,
+ with_dense_inputs: bool = False,
+ export_metrics: bool = True,
+ **kwargs):
+ super().__init__(**kwargs)
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
+
+ if embedding_width is None:
+ embedding_width = hidden_size
+
+ self._config = {
+ 'vocab_size': vocab_size,
+ 'hidden_size': hidden_size,
+ 'num_layers': num_layers,
+ 'moe_layers': moe_layers,
+ 'num_experts': num_experts,
+ 'train_capacity_factor': train_capacity_factor,
+ 'eval_capacity_factor': eval_capacity_factor,
+ 'examples_per_group': examples_per_group,
+ 'mixing_mechanism': mixing_mechanism,
+ 'use_fft': use_fft,
+ 'attention_layers': attention_layers,
+ 'num_attention_heads': num_attention_heads,
+ 'max_sequence_length': max_sequence_length,
+ 'type_vocab_size': type_vocab_size,
+ 'inner_dim': inner_dim,
+ 'inner_activation': tf_keras.activations.serialize(activation),
+ 'output_dropout': output_dropout,
+ 'attention_dropout': attention_dropout,
+ 'initializer': tf_keras.initializers.serialize(initializer),
+ 'output_range': output_range,
+ 'embedding_width': embedding_width,
+ 'embedding_layer': embedding_layer,
+ 'norm_first': norm_first,
+ 'with_dense_inputs': with_dense_inputs,
+ 'export_metrics': export_metrics,
+ }
+
+ if embedding_layer is None:
+ self._embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ initializer=tf_utils.clone_initializer(initializer),
+ name='word_embeddings')
+ else:
+ self._embedding_layer = embedding_layer
+
+ self._position_embedding_layer = layers.PositionEmbedding(
+ initializer=tf_utils.clone_initializer(initializer),
+ max_length=max_sequence_length,
+ name='position_embedding')
+
+ self._type_embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=type_vocab_size,
+ embedding_width=embedding_width,
+ initializer=tf_utils.clone_initializer(initializer),
+ use_one_hot=True,
+ name='type_embeddings')
+
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
+ name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
+
+ self._embedding_dropout = tf_keras.layers.Dropout(
+ rate=output_dropout, name='embedding_dropout')
+
+ # We project the 'embedding' output to 'hidden_size' if it is not already
+ # 'hidden_size'.
+ self._embedding_projection = None
+ if embedding_width != hidden_size:
+ self._embedding_projection = tf_keras.layers.EinsumDense(
+ '...x,xy->...y',
+ output_shape=hidden_size,
+ bias_axes='y',
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='embedding_projection')
+
+ self._transformer_layers = []
+ for layer in range(num_layers):
+ if layer in attention_layers:
+ mixing_layer = layers.MultiHeadAttention(
+ num_heads=num_attention_heads,
+ key_dim=int(hidden_size // num_attention_heads),
+ dropout=attention_dropout,
+ use_bias=True,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='self_attention',
+ )
+ else:
+ mixing_layer = self._init_mixing_sublayer(layer)
+
+ if layer in moe_layers:
+ feedforward_layer = layers.MoeLayer(
+ experts=layers.FeedForwardExperts(
+ num_experts=num_experts,
+ d_ff=inner_dim,
+ output_dropout=output_dropout,
+ activation=inner_activation,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='experts'),
+ router=layers.ExpertsChooseMaskedRouter(
+ num_experts=num_experts,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ export_metrics=export_metrics,
+ name='router'),
+ train_capacity_factor=train_capacity_factor,
+ eval_capacity_factor=eval_capacity_factor,
+ examples_per_group=examples_per_group,
+ name='moe')
+ else:
+ feedforward_layer = None # Fallback to default (dense) MLP class
+
+ block = layers.TransformerScaffold(
+ num_attention_heads=num_attention_heads,
+ inner_dim=inner_dim,
+ inner_activation=inner_activation,
+ attention_cls=mixing_layer,
+ feedforward_cls=feedforward_layer,
+ output_dropout=output_dropout,
+ attention_dropout=attention_dropout,
+ norm_first=norm_first,
+ output_range=output_range if layer == num_layers - 1 else None,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='transformer/layer_%d' % layer)
+ self._transformer_layers.append(block)
+
+ self._attention_mask_layer = layers.SelfAttentionMask(
+ name='self_attention_mask')
+
+ self._pooler_layer = tf_keras.layers.Dense(
+ units=hidden_size,
+ activation='tanh',
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='pooler_transform')
+
+ if with_dense_inputs:
+ self.inputs = dict(
+ # The total length of token ids and dense inputs still has to be
+ # max_sequence_length. It is checked in call().
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_inputs=tf_keras.Input(
+ shape=(None, embedding_width), dtype=tf.float32),
+ dense_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ )
+ else:
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(
+ shape=(max_sequence_length,), dtype=tf.int32),
+ input_mask=tf_keras.Input(
+ shape=(max_sequence_length,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(
+ shape=(max_sequence_length,), dtype=tf.int32))
+ self._max_sequence_length = max_sequence_length
+
+ def call(self, inputs):
+ word_embeddings = None
+ if isinstance(inputs, dict):
+ word_ids = inputs.get('input_word_ids')
+ mask = inputs.get('input_mask')
+ type_ids = inputs.get('input_type_ids')
+ word_embeddings = inputs.get('input_word_embeddings', None)
+
+ dense_inputs = inputs.get('dense_inputs', None)
+ dense_mask = inputs.get('dense_mask', None)
+ dense_type_ids = inputs.get('dense_type_ids', None)
+ else:
+ raise ValueError('Unexpected inputs type (%s) to %s.' %
+ (type(inputs), self.__class__))
+
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ # Concat the dense embeddings at sequence end.
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+ mask = tf.concat([mask, dense_mask], axis=1)
+
+ # SparseMixer: Sequence length must be the same as `max_sequence_length`.
+ word_embeddings = tf.ensure_shape(word_embeddings,
+ [None, self._max_sequence_length, None])
+
+ # Absolute position embeddings.
+ position_embeddings = self._position_embedding_layer(word_embeddings)
+ type_embeddings = self._type_embedding_layer(type_ids)
+
+ embeddings = word_embeddings + position_embeddings + type_embeddings
+ embeddings = self._embedding_norm_layer(embeddings)
+ embeddings = self._embedding_dropout(embeddings)
+
+ if self._embedding_projection is not None:
+ embeddings = self._embedding_projection(embeddings)
+
+ attention_mask = self._attention_mask_layer(embeddings, mask)
+
+ encoder_outputs = []
+ x = embeddings
+ for layer in self._transformer_layers:
+ x = layer([x, attention_mask])
+ encoder_outputs.append(x)
+
+ last_encoder_output = encoder_outputs[-1]
+ first_token_tensor = last_encoder_output[:, 0, :]
+ pooled_output = self._pooler_layer(first_token_tensor)
+
+ output = dict(
+ sequence_output=encoder_outputs[-1],
+ pooled_output=pooled_output,
+ encoder_outputs=encoder_outputs)
+ return output
+
+ def get_embedding_table(self):
+ return self._embedding_layer.embeddings
+
+ def get_embedding_layer(self):
+ return self._embedding_layer
+
+ def get_config(self):
+ return dict(self._config)
+
+ @property
+ def transformer_layers(self):
+ """List of Transformer layers in the encoder."""
+ return self._transformer_layers
+
+ @property
+ def pooler_layer(self):
+ """The pooler dense layer after the transformer layers."""
+ return self._pooler_layer
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ if 'embedding_layer' in config and config['embedding_layer'] is not None:
+ warn_string = (
+ 'You are reloading a model that was saved with a '
+ 'potentially-shared embedding layer object. If you contine to '
+ 'train this model, the embedding layer will no longer be shared. '
+ 'To work around this, load the model outside of the Keras API.')
+ print('WARNING: ' + warn_string)
+ logging.warn(warn_string)
+
+ return cls(**config)
+
+ def _init_mixing_sublayer(self, layer: int):
+ """Initializes config-dependent mixing sublayer."""
+ if self._config['mixing_mechanism'] == layers.MixingMechanism.FOURIER:
+ mixing_sublayer = layers.FourierTransformLayer(
+ use_fft=self._config['use_fft'], name='fourier_transform')
+ elif self._config['mixing_mechanism'] == layers.MixingMechanism.HARTLEY:
+ mixing_sublayer = layers.HartleyTransformLayer(
+ use_fft=self._config['use_fft'], name='hartley_transform')
+ elif self._config['mixing_mechanism'] == layers.MixingMechanism.LINEAR:
+ mixing_sublayer = layers.LinearTransformLayer(
+ kernel_initializer=tf_utils.clone_initializer(
+ self._config['initializer']),
+ name='linear_transform')
+ else:
+ raise ValueError('Unsupported mixing mechanism: %s' %
+ self._config['mixing_mechanism'])
+
+ return mixing_sublayer
diff --git a/official/nlp/modeling/networks/sparse_mixer_test.py b/official/nlp/modeling/networks/sparse_mixer_test.py
new file mode 100644
index 00000000000..250e63ebc71
--- /dev/null
+++ b/official/nlp/modeling/networks/sparse_mixer_test.py
@@ -0,0 +1,143 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for Sparse Mixer encoder network."""
+
+from typing import Sequence
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+from official.nlp.modeling.networks import sparse_mixer
+
+
+class SparseMixerTest(parameterized.TestCase, tf.test.TestCase):
+
+ def tearDown(self):
+ super().tearDown()
+ tf_keras.mixed_precision.set_global_policy("float32")
+
+ @parameterized.named_parameters(
+ dict(
+ testcase_name="sparse_mixer",
+ mixing_mechanism=layers.MixingMechanism.LINEAR,
+ moe_layers=(1,),
+ attention_layers=(2,)),
+ dict(
+ testcase_name="fnet",
+ mixing_mechanism=layers.MixingMechanism.FOURIER,
+ moe_layers=(),
+ attention_layers=()),
+ dict(
+ testcase_name="sparse_hnet",
+ mixing_mechanism=layers.MixingMechanism.HARTLEY,
+ moe_layers=(0, 1, 2),
+ attention_layers=(1, 2)),
+ dict(
+ testcase_name="sparse_bert",
+ mixing_mechanism=layers.MixingMechanism.LINEAR,
+ moe_layers=(0, 1, 2), # All layers use MoE
+ attention_layers=(0, 1, 2)), # All layers use attention
+ )
+ def test_network(self, mixing_mechanism: layers.MixingMechanism,
+ attention_layers: Sequence[int], moe_layers: Sequence[int]):
+ num_layers = 3
+ hidden_size = 16
+ sequence_length = 32
+ test_network = sparse_mixer.SparseMixer(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=2,
+ max_sequence_length=sequence_length,
+ num_layers=num_layers,
+ moe_layers=moe_layers,
+ num_experts=8,
+ mixing_mechanism=mixing_mechanism,
+ attention_layers=attention_layers)
+
+ batch_size = 4
+ word_ids = tf_keras.Input(
+ shape=(sequence_length,), batch_size=batch_size, dtype=tf.int32)
+ mask = tf_keras.Input(
+ shape=(sequence_length,), batch_size=batch_size, dtype=tf.int32)
+ type_ids = tf_keras.Input(
+ shape=(sequence_length,), batch_size=batch_size, dtype=tf.int32)
+
+ dict_outputs = test_network(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ data = dict_outputs["sequence_output"]
+ pooled = dict_outputs["pooled_output"]
+
+ self.assertIsInstance(test_network.transformer_layers, list)
+ self.assertLen(test_network.transformer_layers, 3)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
+
+ expected_data_shape = [batch_size, sequence_length, hidden_size]
+ expected_pooled_shape = [batch_size, hidden_size]
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+ self.assertAllEqual(expected_pooled_shape, pooled.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, data.dtype)
+ self.assertAllEqual(tf.float32, pooled.dtype)
+
+ def test_embeddings_as_inputs(self):
+ hidden_size = 32
+ sequence_length = 8
+ test_network = sparse_mixer.SparseMixer(
+ vocab_size=100,
+ hidden_size=hidden_size,
+ num_attention_heads=2,
+ max_sequence_length=sequence_length,
+ num_layers=3,
+ moe_layers=(1,),
+ num_experts=4,
+ attention_layers=(2,))
+
+ batch_size = 2
+ word_ids = tf_keras.Input(
+ shape=(sequence_length), batch_size=batch_size, dtype=tf.int32)
+ mask = tf_keras.Input(
+ shape=(sequence_length,), batch_size=batch_size, dtype=tf.int32)
+ type_ids = tf_keras.Input(
+ shape=(sequence_length,), batch_size=batch_size, dtype=tf.int32)
+
+ test_network.build(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ embeddings = test_network.get_embedding_layer()(word_ids)
+
+ # Calls with the embeddings.
+ dict_outputs = test_network(
+ dict(
+ input_word_embeddings=embeddings,
+ input_mask=mask,
+ input_type_ids=type_ids))
+ all_encoder_outputs = dict_outputs["encoder_outputs"]
+ pooled = dict_outputs["pooled_output"]
+
+ expected_data_shape = [batch_size, sequence_length, hidden_size]
+ expected_pooled_shape = [batch_size, hidden_size]
+ self.assertLen(all_encoder_outputs, 3)
+ for data in all_encoder_outputs:
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+ self.assertAllEqual(expected_pooled_shape, pooled.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, all_encoder_outputs[-1].dtype)
+ self.assertAllEqual(tf.float32, pooled.dtype)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/nlp/modeling/networks/xlnet_base.py b/official/nlp/modeling/networks/xlnet_base.py
index fbb276a7071..1ca0c6a1c18 100644
--- a/official/nlp/modeling/networks/xlnet_base.py
+++ b/official/nlp/modeling/networks/xlnet_base.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,8 +16,9 @@
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
from official.nlp.modeling.layers import transformer_xl
@@ -365,7 +366,7 @@ def _compute_positional_encoding(
return relative_position_encoding
-class RelativePositionEncoding(tf.keras.layers.Layer):
+class RelativePositionEncoding(tf_keras.layers.Layer):
"""Creates a relative positional encoding.
This layer creates a relative positional encoding as described in
@@ -383,7 +384,7 @@ class RelativePositionEncoding(tf.keras.layers.Layer):
"""
def __init__(self, hidden_size, **kwargs):
- super(RelativePositionEncoding, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._hidden_size = hidden_size
self._inv_freq = 1.0 / (10000.0**(
tf.range(0, self._hidden_size, 2.0) / self._hidden_size))
@@ -411,8 +412,8 @@ def call(self, pos_seq, batch_size=None):
return relative_position_encoding
-@tf.keras.utils.register_keras_serializable(package="Text")
-class XLNetBase(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class XLNetBase(tf_keras.layers.Layer):
"""Base XLNet model.
Attributes:
@@ -475,7 +476,7 @@ def __init__(self,
use_cls_mask=False,
embedding_width=None,
**kwargs):
- super(XLNetBase, self).__init__(**kwargs)
+ super().__init__(**kwargs)
self._vocab_size = vocab_size
self._initializer = initializer
@@ -507,12 +508,12 @@ def __init__(self,
self._embedding_layer = layers.OnDeviceEmbedding(
vocab_size=self._vocab_size,
embedding_width=embedding_width,
- initializer=self._initializer,
+ initializer=tf_utils.clone_initializer(self._initializer),
dtype=tf.float32,
name="word_embedding")
- self._dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self._dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
- self.embedding_dropout = tf.keras.layers.Dropout(rate=self._dropout_rate)
+ self.embedding_dropout = tf_keras.layers.Dropout(rate=self._dropout_rate)
self.position_encoding = RelativePositionEncoding(self._hidden_size)
self._transformer_xl = transformer_xl.TransformerXL(
@@ -573,7 +574,7 @@ def get_config(self):
"embedding_width":
self._embedding_width,
}
- base_config = super(XLNetBase, self).get_config()
+ base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
def get_embedding_lookup_table(self):
@@ -600,7 +601,7 @@ def __call__(self,
"target_mapping": target_mapping,
"masked_tokens": masked_tokens
}
- return super(XLNetBase, self).__call__(inputs, **kwargs)
+ return super().__call__(inputs, **kwargs)
def call(self, inputs):
"""Implements call() for the layer."""
@@ -666,7 +667,7 @@ def call(self, inputs):
shape=[self._num_layers, 2, self._num_attention_heads,
self._head_size],
dtype=tf.float32,
- initializer=self._initializer)
+ initializer=tf_utils.clone_initializer(self._initializer))
segment_embedding = self._segment_embedding
segment_matrix = _compute_segment_matrix(
diff --git a/official/nlp/modeling/networks/xlnet_base_test.py b/official/nlp/modeling/networks/xlnet_base_test.py
index c2abda38711..83c115e0bbe 100644
--- a/official/nlp/modeling/networks/xlnet_base_test.py
+++ b/official/nlp/modeling/networks/xlnet_base_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,16 +13,16 @@
# limitations under the License.
"""Tests for Keras based XLNet model."""
+
+from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.networks import xlnet_base
-@keras_parameterized.run_all_keras_modes
-class RelativePositionEncodingTest(keras_parameterized.TestCase):
+class RelativePositionEncodingTest(tf.test.TestCase):
def test_positional_embedding(self):
"""A low-dimensional example is tested.
@@ -47,7 +47,7 @@ def test_positional_embedding(self):
self.assertAllClose(encoding, target)
-class ComputePositionEncodingTest(keras_parameterized.TestCase):
+class ComputePositionEncodingTest(tf.test.TestCase, parameterized.TestCase):
@combinations.generate(combinations.combine(
attention_type=["uni", "bi"],
@@ -111,7 +111,7 @@ def test_causal_attention_mask_with_same_length(self):
self.assertAllClose(causal_attention_mask, expected_output)
-class MaskComputationTests(keras_parameterized.TestCase):
+class MaskComputationTests(tf.test.TestCase, parameterized.TestCase):
@combinations.generate(combinations.combine(
use_input_mask=[False, True],
@@ -411,7 +411,7 @@ def test_xlnet_model(self):
attention_dropout_rate=0.,
attention_type="bi",
bi_data=True,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
two_stream=False,
tie_attention_biases=True,
reuse_length=0,
@@ -435,7 +435,7 @@ def test_get_config(self):
attention_dropout_rate=0.,
attention_type="bi",
bi_data=True,
- initializer=tf.keras.initializers.RandomNormal(stddev=0.1),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.1),
two_stream=False,
tie_attention_biases=True,
memory_length=0,
diff --git a/official/nlp/modeling/ops/__init__.py b/official/nlp/modeling/ops/__init__.py
index 3e1c9d52d23..3918cc5bfac 100644
--- a/official/nlp/modeling/ops/__init__.py
+++ b/official/nlp/modeling/ops/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,5 +14,7 @@
"""Ops package definition."""
from official.nlp.modeling.ops.beam_search import sequence_beam_search
+from official.nlp.modeling.ops.beam_search import SequenceBeamSearch
+from official.nlp.modeling.ops.sampling_module import SamplingModule
from official.nlp.modeling.ops.segment_extractor import get_next_sentence_labels
from official.nlp.modeling.ops.segment_extractor import get_sentence_order_labels
diff --git a/official/nlp/modeling/ops/beam_search.py b/official/nlp/modeling/ops/beam_search.py
index fc406819f81..831eb8797e3 100644
--- a/official/nlp/modeling/ops/beam_search.py
+++ b/official/nlp/modeling/ops/beam_search.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Beam search to find the translated sequence with the highest probability."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def inf(dtype):
@@ -69,6 +69,9 @@ class _StateKeys(object):
# At the beginning, all of the sequences in FINISHED_SEQ are filler values.
# True -> finished sequence, False -> filler. Shape [batch_size, beam_size]
FINISHED_FLAGS = "FINISHED_FLAGS"
+ # for prefix matching hack. The BS will only constraint the next token to
+ # where the mask is 1.
+ CONSTRAINT_MASK = "CONSTRAINT_MASK"
def _expand_to_same_rank(tensor, target):
@@ -99,48 +102,58 @@ def _expand_to_same_rank(tensor, target):
class SequenceBeamSearch(tf.Module):
"""Implementation of beam search loop."""
- def __init__(self,
- symbols_to_logits_fn,
- vocab_size,
- beam_size,
- alpha,
- max_decode_length,
- eos_id,
- padded_decode,
- dtype=tf.float32):
+ def __init__(
+ self,
+ symbols_to_logits_fn,
+ vocab_size,
+ beam_size,
+ alpha,
+ max_decode_length,
+ eos_id,
+ padded_decode,
+ dtype=tf.float32,
+ noise_multiplier: float = 0.0,
+ decoding_name=None,
+ ):
"""Initialize sequence beam search.
Args:
- symbols_to_logits_fn: A function to provide logits, which is the
- interface to the Transformer model. The passed in arguments are: ids ->
- A tensor with shape [batch_size * beam_size, index]. index -> A
- scalar. cache -> A nested dictionary of tensors [batch_size *
- beam_size, ...].
- The function must return a tuple of logits and the updated cache: logits
- -> A tensor with shape [batch * beam_size, vocab_size]. updated cache
- -> A nested dictionary with the same structure as the input cache.
+ symbols_to_logits_fn: A function to provide logits, which is the interface
+ to the Transformer model. The passed in arguments are: ids -> A tensor
+ with shape [batch_size * beam_size, index]. index -> A scalar. cache ->
+ A nested dictionary of tensors [batch_size * beam_size, ...]. The
+ function must return a tuple of logits and the updated cache: logits ->
+ A tensor with shape [batch * beam_size, vocab_size]. updated cache -> A
+ nested dictionary with the same structure as the input cache.
vocab_size: An integer, the size of the vocabulary, used for topk
computation.
beam_size: An integer, number of beams for beam search.
alpha: A float, defining the strength of length normalization.
max_decode_length: An integer, the maximum number of steps to decode a
sequence.
- eos_id: An integer. ID of end of sentence token.
+ eos_id: An integer or a list. ID of end of sentence token.
padded_decode: A bool, indicating if max_sequence_length padding is used
for beam search.
dtype: A tensorflow data type used for score computation. The default is
tf.float32.
+ noise_multiplier: The amount of noise.
+ decoding_name: an optional name for the decoding loop tensors.
"""
self.symbols_to_logits_fn = symbols_to_logits_fn
self.vocab_size = vocab_size
self.beam_size = beam_size
self.alpha = alpha
self.max_decode_length = max_decode_length
- self.eos_id = eos_id
+ if isinstance(eos_id, list):
+ self.eos_id = eos_id
+ else:
+ self.eos_id = [eos_id]
self.padded_decode = padded_decode
self.dtype = tf.as_dtype(dtype)
+ self.decoding_name = decoding_name
+ self.noise_multiplier = noise_multiplier
- def search(self, initial_ids, initial_cache):
+ def search(self, initial_ids, initial_cache, constraint_mask=None):
"""Beam search for sequences with highest scores.
Args:
@@ -148,6 +161,9 @@ def search(self, initial_ids, initial_cache):
with shape [batch_size, 1]
initial_cache: dictionary storing values to be passed into the
symbols_to_logits_fn.
+ constraint_mask: a [vocab_size] tensor, with 1 represent prefix. During
+ autoregressive decoding, the first token should be among where the
+ constraint_mask is 1.
Returns:
finished_seq and finished_scores.
@@ -155,8 +171,9 @@ def search(self, initial_ids, initial_cache):
batch_size = (
initial_ids.shape.as_list()[0]
if self.padded_decode else tf.shape(initial_ids)[0])
- state, state_shapes = self._create_initial_state(initial_ids, initial_cache,
- batch_size)
+ state, state_shapes = self._create_initial_state(
+ initial_ids, initial_cache, batch_size, constraint_mask=constraint_mask
+ )
def _grow_alive_seq(state):
"""Grow alive sequences by one token, collect top 2*beam_size sequences.
@@ -194,6 +211,28 @@ def _grow_alive_seq(state):
flat_logits, flat_cache = self.symbols_to_logits_fn(
flat_ids, i, flat_cache)
+ if _StateKeys.CONSTRAINT_MASK in state:
+ constraint_mask = state[_StateKeys.CONSTRAINT_MASK]
+ constraint_mask = tf.cond(
+ tf.equal(i, 0),
+ lambda: constraint_mask,
+ lambda: tf.ones_like(constraint_mask),
+ )
+ penalty = tf.cast(
+ tf.cast(constraint_mask != 1, tf.int32) * 999_999_999,
+ flat_logits.dtype,
+ )
+ flat_logits = flat_logits - penalty[tf.newaxis, :]
+ else:
+ constraint_mask = None
+
+ if self.noise_multiplier > 0:
+ noise = tf.random.uniform(flat_logits.shape, dtype=flat_logits.dtype)
+ # Generates standard Gumbel(0, 1) noise, GSE Tensors
+ noise = -tf.math.log(-tf.math.log(noise))
+ # NOMUTANTS -- may not impact final result.
+ flat_logits = flat_logits + noise * self.noise_multiplier
+
# Unflatten logits to shape [batch_size, beam_size, vocab_size]
logits = _unflatten_beam_dim(flat_logits, batch_size, self.beam_size)
new_cache = tf.nest.map_structure(
@@ -204,7 +243,7 @@ def _grow_alive_seq(state):
candidate_log_probs = _log_prob_from_logits(logits)
# Calculate new log probabilities if each of the alive sequences were
- # extended # by the the candidate IDs.
+ # extended # by the candidate IDs.
# Shape [batch_size, beam_size, vocab_size]
log_probs = candidate_log_probs + tf.expand_dims(alive_log_probs, axis=2)
@@ -233,7 +272,7 @@ def _grow_alive_seq(state):
else:
topk_seq = tf.concat(
[topk_seq, tf.expand_dims(topk_ids, axis=2)], axis=2)
- return topk_seq, topk_log_probs, topk_ids, new_cache
+ return topk_seq, topk_log_probs, topk_ids, new_cache, constraint_mask
def _get_new_alive_state(new_seq, new_log_probs, new_finished_flags,
new_cache):
@@ -346,8 +385,15 @@ def _search_step(state):
new state dictionary.
"""
# Grow alive sequences by one token.
- new_seq, new_log_probs, topk_ids, new_cache = _grow_alive_seq(state)
- new_finished_flags = tf.equal(topk_ids, self.eos_id)
+ new_seq, new_log_probs, topk_ids, new_cache, constraint_mask = (
+ _grow_alive_seq(state)
+ )
+ new_finished_flags = tf.equal(topk_ids, self.eos_id[0])
+ for eos_id in self.eos_id[1:]:
+ one_finished_flags = tf.equal(topk_ids, eos_id)
+ new_finished_flags = tf.logical_or(
+ new_finished_flags, one_finished_flags
+ )
# Collect top beam_size alive sequences
alive_state = _get_new_alive_state(new_seq, new_log_probs,
new_finished_flags, new_cache)
@@ -361,6 +407,8 @@ def _search_step(state):
new_state = {_StateKeys.CUR_INDEX: state[_StateKeys.CUR_INDEX] + 1}
new_state.update(alive_state)
new_state.update(finished_state)
+ if constraint_mask is not None:
+ new_state[_StateKeys.CONSTRAINT_MASK] = constraint_mask
return [new_state]
finished_state = tf.nest.map_structure(
@@ -370,7 +418,8 @@ def _search_step(state):
_search_step,
loop_vars=[state],
shape_invariants=[state_shapes],
- parallel_iterations=1))
+ parallel_iterations=1,
+ name=self.decoding_name))
finished_state = finished_state[0]
return self._process_finished_state(finished_state)
@@ -392,7 +441,9 @@ def _process_finished_state(self, finished_state):
finished_scores = tf.where(score_cond, finished_scores, alive_log_probs)
return finished_seq, finished_scores
- def _create_initial_state(self, initial_ids, initial_cache, batch_size):
+ def _create_initial_state(
+ self, initial_ids, initial_cache, batch_size, constraint_mask=None
+ ):
"""Return initial state dictionary and its shape invariants."""
for key, value in initial_cache.items():
for inner_value in tf.nest.flatten(value):
@@ -443,6 +494,8 @@ def _create_initial_state(self, initial_ids, initial_cache, batch_size):
_StateKeys.FINISHED_SCORES: finished_scores,
_StateKeys.FINISHED_FLAGS: finished_flags
}
+ if constraint_mask is not None:
+ state[_StateKeys.CONSTRAINT_MASK] = constraint_mask
# Create state invariants for each value in the state dictionary. Each
# dimension must be a constant or None. A None dimension means either:
@@ -486,6 +539,10 @@ def _create_initial_state(self, initial_ids, initial_cache, batch_size):
_StateKeys.FINISHED_FLAGS:
tf.TensorShape([None, self.beam_size])
}
+ if constraint_mask is not None:
+ state_shape_invariants[_StateKeys.CONSTRAINT_MASK] = tf.TensorShape(
+ [self.vocab_size]
+ )
return state, state_shape_invariants
@@ -564,7 +621,7 @@ def _gather_beams(nested, beam_indices, batch_size, new_beam_size):
Nested structure containing tensors with shape
[batch_size, new_beam_size, ...]
"""
- # Computes the i'th coodinate that contains the batch index for gather_nd.
+ # Computes the i'th coordinate that contains the batch index for gather_nd.
# Batch pos is a tensor like [[0,0,0,0,],[1,1,1,1],..].
batch_pos = tf.range(batch_size * new_beam_size) // new_beam_size
batch_pos = tf.reshape(batch_pos, [batch_size, new_beam_size])
@@ -578,26 +635,31 @@ def _gather_beams(nested, beam_indices, batch_size, new_beam_size):
nested)
-def sequence_beam_search(symbols_to_logits_fn,
- initial_ids,
- initial_cache,
- vocab_size,
- beam_size,
- alpha,
- max_decode_length,
- eos_id,
- padded_decode=False,
- dtype="float32"):
+def sequence_beam_search(
+ symbols_to_logits_fn,
+ initial_ids,
+ initial_cache,
+ vocab_size,
+ beam_size,
+ alpha,
+ max_decode_length,
+ eos_id,
+ padded_decode=False,
+ dtype="float32",
+ noise_multiplier: float = 0.0,
+ decoding_name=None,
+ constraint_mask=None,
+):
"""Search for sequence of subtoken ids with the largest probability.
Args:
symbols_to_logits_fn: A function that takes in ids, index, and cache as
arguments. The passed in arguments will have shape: ids -> A tensor with
- shape [batch_size * beam_size, index]. index -> A scalar. cache -> A
- nested dictionary of tensors [batch_size * beam_size, ...].
- The function must return a tuple of logits and new cache: logits -> A
- tensor with shape [batch * beam_size, vocab_size]. new cache -> A nested
- dictionary with the same shape/structure as the inputted cache.
+ shape [batch_size * beam_size, index]. index -> A scalar. cache -> A
+ nested dictionary of tensors [batch_size * beam_size, ...]. The function
+ must return a tuple of logits and new cache: logits -> A tensor with shape
+ [batch * beam_size, vocab_size]. new cache -> A nested dictionary with the
+ same shape/structure as the inputted cache.
initial_ids: An int32 tensor with shape [batch_size]. Starting ids for each
batch item.
initial_cache: A dictionary, containing starting decoder variables
@@ -612,14 +674,28 @@ def sequence_beam_search(symbols_to_logits_fn,
beam search.
dtype: A tensorflow data type used for score computation. The default is
tf.float32.
+ noise_multiplier: The amount of noise.
+ decoding_name: an optional name for the decoding loop tensors.
+ constraint_mask: The BS will only constraint the next token to where the
+ mask is 1.
Returns:
Top decoded sequences [batch_size, beam_size, max_decode_length]
sequence scores [batch_size, beam_size]
"""
- sbs = SequenceBeamSearch(symbols_to_logits_fn, vocab_size, beam_size, alpha,
- max_decode_length, eos_id, padded_decode, dtype)
- return sbs.search(initial_ids, initial_cache)
+ sbs = SequenceBeamSearch(
+ symbols_to_logits_fn,
+ vocab_size,
+ beam_size,
+ alpha,
+ max_decode_length,
+ eos_id,
+ padded_decode,
+ dtype,
+ noise_multiplier,
+ decoding_name,
+ )
+ return sbs.search(initial_ids, initial_cache, constraint_mask=constraint_mask)
def _log_prob_from_logits(logits):
diff --git a/official/nlp/modeling/ops/beam_search_test.py b/official/nlp/modeling/ops/beam_search_test.py
index 6a541a12891..a4b5295f920 100644
--- a/official/nlp/modeling/ops/beam_search_test.py
+++ b/official/nlp/modeling/ops/beam_search_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Test beam search helper methods."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.ops import beam_search
@@ -60,10 +60,13 @@ def test_gather_beams(self):
y)
@parameterized.named_parameters([
- ('padded_decode_true', True),
- ('padded_decode_false', False),
+ ('padded_decode_true_with_name', True, 0.0, 'decoding'),
+ ('padded_decode_false_with_name', False, 0.0, 'decoding'),
+ ('padded_decode_true_without_name', True, 0.0, None),
+ ('padded_decode_false_without_name', False, 0.0, None),
+ ('padded_decode_false_with_noise', False, 0.5, 'decoding'),
])
- def test_sequence_beam_search(self, padded_decode):
+ def test_sequence_beam_search(self, padded_decode, noise_multiplier, name):
# batch_size*beam_size, max_decode_length, vocab_size
probabilities = tf.constant([[[0.2, 0.7, 0.1], [0.5, 0.3, 0.2],
[0.1, 0.8, 0.1]],
@@ -91,8 +94,109 @@ def symbols_to_logits_fn(_, i, cache):
max_decode_length=3,
eos_id=9,
padded_decode=padded_decode,
- dtype=tf.float32)
- self.assertAllEqual([[[0, 1, 0, 1], [0, 1, 1, 2]]], predictions)
+ dtype=tf.float32,
+ noise_multiplier=noise_multiplier,
+ decoding_name=name,
+ )
+ if noise_multiplier > 0:
+ self.assertAllEqual([[[0, 1, 0, 1], [0, 0, 2, 2]]], predictions)
+ else:
+ self.assertAllEqual([[[0, 1, 0, 1], [0, 1, 1, 2]]], predictions)
+
+ @parameterized.named_parameters([
+ ('padded_decode_true_with_name', True, 0.0, 'decoding'),
+ ('padded_decode_false_with_name', False, 0.0, 'decoding'),
+ ('padded_decode_true_without_name', True, 0.0, None),
+ ('padded_decode_false_without_name', False, 0.0, None),
+ ('padded_decode_false_with_noise', False, 0.5, 'decoding'),
+ ])
+ def test_sequence_beam_search_multi_eos(
+ self, padded_decode, noise_multiplier, name
+ ):
+ # batch_size*beam_size, max_decode_length, vocab_size
+ probabilities = tf.constant([
+ [[0.2, 0.7, 0.1], [0.5, 0.3, 0.2], [0.1, 0.8, 0.1]],
+ [[0.1, 0.8, 0.1], [0.3, 0.4, 0.3], [0.2, 0.1, 0.7]],
+ ])
+ # batch_size, max_decode_length, num_heads, embed_size per head
+ x = tf.zeros([1, 3, 2, 32], dtype=tf.float32)
+ cache = {'layer_%d' % layer: {'k': x, 'v': x} for layer in range(2)}
+
+ def _get_test_symbols_to_logits_fn():
+ """Test function that returns logits for next token."""
+
+ def symbols_to_logits_fn(_, i, cache):
+ logits = tf.cast(probabilities[:, i, :], tf.float32)
+ return logits, cache
+
+ return symbols_to_logits_fn
+
+ predictions, _ = beam_search.sequence_beam_search(
+ symbols_to_logits_fn=_get_test_symbols_to_logits_fn(),
+ initial_ids=tf.zeros([1], dtype=tf.int32),
+ initial_cache=cache,
+ vocab_size=3,
+ beam_size=2,
+ alpha=0.6,
+ max_decode_length=3,
+ eos_id=[9, 10],
+ padded_decode=padded_decode,
+ dtype=tf.float32,
+ noise_multiplier=noise_multiplier,
+ decoding_name=name,
+ )
+ if noise_multiplier > 0:
+ self.assertAllEqual([[[0, 1, 0, 1], [0, 0, 2, 2]]], predictions)
+ else:
+ self.assertAllEqual([[[0, 1, 0, 1], [0, 1, 1, 2]]], predictions)
+
+ @parameterized.named_parameters([
+ ('padded_decode_true_with_name', True, 0.0, 'decoding'),
+ ('padded_decode_false_with_name', False, 0.0, 'decoding'),
+ ('padded_decode_true_without_name', True, 0.0, None),
+ ('padded_decode_false_without_name', False, 0.0, None),
+ ('padded_decode_false_with_noise', False, 0.5, 'decoding'),
+ ])
+ def test_sequence_beam_search_with_prefix_constraint(
+ self, padded_decode, noise_multiplier, name
+ ):
+ # batch_size*beam_size, max_decode_length, vocab_size
+ probabilities = tf.constant([
+ [[0.2, 0.7, 0.1], [0.5, 0.3, 0.2], [0.1, 0.8, 0.1]],
+ [[0.1, 0.8, 0.1], [0.3, 0.4, 0.3], [0.2, 0.1, 0.7]],
+ ])
+ # batch_size, max_decode_length, num_heads, embed_size per head
+ x = tf.zeros([1, 3, 2, 32], dtype=tf.float32)
+ cache = {'layer_%d' % layer: {'k': x, 'v': x} for layer in range(2)}
+
+ def _get_test_symbols_to_logits_fn():
+ """Test function that returns logits for next token."""
+
+ def symbols_to_logits_fn(_, i, cache):
+ logits = tf.cast(probabilities[:, i, :], tf.float32)
+ return logits, cache
+
+ return symbols_to_logits_fn
+
+ predictions, _ = beam_search.sequence_beam_search(
+ symbols_to_logits_fn=_get_test_symbols_to_logits_fn(),
+ initial_ids=tf.zeros([1], dtype=tf.int32),
+ initial_cache=cache,
+ vocab_size=3,
+ beam_size=2,
+ alpha=0.6,
+ max_decode_length=3,
+ eos_id=[9, 10],
+ padded_decode=padded_decode,
+ dtype=tf.float32,
+ noise_multiplier=noise_multiplier,
+ decoding_name=name,
+ constraint_mask=tf.constant([1, 0, 0]),
+ )
+ if noise_multiplier > 0:
+ self.assertAllEqual([[[0, 0, 0, 1], [0, 0, 0, 2]]], predictions)
+ else:
+ self.assertAllEqual([[[0, 0, 0, 1], [0, 0, 1, 2]]], predictions)
if __name__ == '__main__':
diff --git a/official/nlp/modeling/ops/decoding_module.py b/official/nlp/modeling/ops/decoding_module.py
index 40a02bb4bf6..fe4bd78e884 100644
--- a/official/nlp/modeling/ops/decoding_module.py
+++ b/official/nlp/modeling/ops/decoding_module.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,14 +15,14 @@
"""Base class for Decoding Strategies (beam_search, top_k, top_p and greedy)."""
import abc
-from typing import Any, Callable, Dict, Tuple
+from typing import Any, Callable, Dict, Optional, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.framework import dtypes
from official.modeling import tf_utils
-Output = Tuple[tf.Tensor, tf.Tensor]
+Output = Tuple[tf.Tensor, tf.Tensor, Optional[tf.Tensor]]
InternalState = Tuple[tf.Tensor, tf.Tensor, tf.Tensor, Dict]
InitialState = Tuple[Dict[str, Any], Dict[str, Any]]
@@ -46,6 +46,10 @@ class StateKeys:
# the previous iteration.
ALIVE_CACHE = "ALIVE_CACHE"
+ # The initial model state/cache after model processing the initial token.
+ # The cache will be filled if extra_cache_output is true.
+ INITIAL_OUTPUT_CACHE = "INITIAL_OUTPUT_CACHE"
+
# Top finished sequences for each batch item.
# Has shape [batch_size, beam_size, CUR_INDEX + 1]. Sequences that are
# shorter than CUR_INDEX + 1 are padded with 0s.
@@ -108,7 +112,9 @@ class DecodingModule(tf.Module, metaclass=abc.ABCMeta):
def __init__(self,
length_normalization_fn: Callable[[int, tf.DType], float],
- dtype: tf.DType = tf.float32):
+ dtype: tf.DType = tf.float32,
+ decoding_name: Optional[str] = None,
+ extra_cache_output: bool = False):
"""Initialize the Decoding Module.
Args:
@@ -116,31 +122,39 @@ def __init__(self,
parameter. Function accepts input as length, dtype and returns float.
dtype: A tensorflow data type used for score computation. The default is
tf.float32.
+ decoding_name: an optional name for the decoding loop tensors.
+ extra_cache_output: If true, the first cache will be in the states.
"""
self.length_normalization_fn = length_normalization_fn
self.dtype = tf.as_dtype(dtype)
+ self.decoding_name = decoding_name
def generate(self,
initial_ids: tf.Tensor,
- initial_cache: Dict[str, tf.Tensor]) -> Output:
+ initial_cache: Dict[str, tf.Tensor],
+ initial_log_probs: Optional[tf.Tensor] = None) -> Output:
"""Implements the decoding strategy (beam_search or sampling).
Args:
- initial_ids: initial ids to pass into the symbols_to_logits_fn.
- int tensor with shape [batch_size, 1]
+ initial_ids: initial ids to pass into the symbols_to_logits_fn. int tensor
+ with shape [batch_size, 1]
initial_cache: dictionary for caching model outputs from previous step.
+ initial_log_probs: Optionally initial log probs if there is a prefix
+ sequence we want to start to decode from.
+
Returns:
Tuple of tensors representing
finished_sequence: shape [batch, max_seq_length]
finished_scores: [batch]
+ first_cache: The cache after init token
"""
batch_size = (
initial_ids.shape.as_list()[0]
if self.padded_decode else tf.shape(initial_ids)[0])
- state, state_shapes = self._create_initial_state(initial_ids,
- initial_cache,
- batch_size)
+ state, state_shapes = self._create_initial_state(initial_ids, initial_cache,
+ batch_size,
+ initial_log_probs)
def _generate_step(state):
topk_seq, topk_log_probs, topk_ids, new_cache = self._grow_alive_seq(
@@ -160,6 +174,17 @@ def _generate_step(state):
}
new_state.update(alive_state)
new_state.update(finished_state)
+ if self.extra_cache_output:
+ i = state[StateKeys.CUR_INDEX]
+ old_cache = state[StateKeys.INITIAL_OUTPUT_CACHE]
+
+ def update_with_cache(new_state, cache):
+ """Updates new_state with cache."""
+ new_state.update({StateKeys.INITIAL_OUTPUT_CACHE: cache})
+
+ tf.cond(
+ tf.equal(i, 0), lambda: update_with_cache(new_state, new_cache),
+ lambda: update_with_cache(new_state, old_cache))
return [new_state]
finished_state = tf.nest.map_structure(
@@ -169,15 +194,18 @@ def _generate_step(state):
_generate_step,
loop_vars=[state],
shape_invariants=[state_shapes],
- parallel_iterations=1))
+ parallel_iterations=1,
+ name=self.decoding_name))
final_state = self._process_finished_state(finished_state[0])
return final_state
@abc.abstractmethod
- def _create_initial_state(self,
- initial_ids: tf.Tensor,
- initial_cache: Dict[str, tf.Tensor],
- batch_size: int) -> InitialState:
+ def _create_initial_state(
+ self,
+ initial_ids: tf.Tensor,
+ initial_cache: Dict[str, tf.Tensor],
+ batch_size: int,
+ initial_log_probs: Optional[tf.Tensor] = None) -> InitialState:
"""Return initial state dictionary and its shape invariants."""
pass
@@ -277,6 +305,3 @@ def inf(self):
return dtypes.float16.max
else:
raise AssertionError("Invalid dtype: %s" % self.dtype)
-
-
-
diff --git a/official/nlp/modeling/ops/decoding_module_test.py b/official/nlp/modeling/ops/decoding_module_test.py
index f29b5d6720a..fcec5f96588 100644
--- a/official/nlp/modeling/ops/decoding_module_test.py
+++ b/official/nlp/modeling/ops/decoding_module_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Test decoding utility methods."""
import abc
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.ops import decoding_module
@@ -29,6 +29,7 @@ class TestSubclass(decoding_module.DecodingModule, metaclass=abc.ABCMeta):
def __init__(self,
length_normalization_fn=length_normalization,
+ extra_cache_output=True,
dtype=tf.float32):
super(TestSubclass, self).__init__(
length_normalization_fn=length_normalization, dtype=dtype)
diff --git a/official/nlp/modeling/ops/sampling_module.py b/official/nlp/modeling/ops/sampling_module.py
index 4278ef2c03b..ae2b1fb1925 100644
--- a/official/nlp/modeling/ops/sampling_module.py
+++ b/official/nlp/modeling/ops/sampling_module.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from typing import Any, Callable, Dict, Optional
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.ops import decoding_module
@@ -55,10 +55,12 @@ def sample_top_k(logits, top_k):
Returns:
Logits with top_k filtering applied.
"""
+ top_k = tf.clip_by_value(
+ top_k, clip_value_min=1, clip_value_max=tf.shape(logits)[-1])
top_k_logits = tf.math.top_k(logits, k=top_k)
indices_to_remove = logits < tf.expand_dims(top_k_logits[0][..., -1], -1)
top_k_logits = set_tensor_by_indices_to_value(logits, indices_to_remove,
- np.NINF)
+ -np.inf)
return top_k_logits
@@ -101,7 +103,7 @@ def sample_top_p(logits, top_p):
indices_to_remove = scatter_values_on_batch_indices(sorted_indices_to_remove,
sorted_indices)
top_p_logits = set_tensor_by_indices_to_value(logits, indices_to_remove,
- np.NINF)
+ -np.inf)
return top_p_logits
@@ -160,7 +162,9 @@ def __init__(self,
top_p=1.0,
sample_temperature=0.0,
enable_greedy: bool = True,
- dtype: tf.DType = tf.float32):
+ dtype: tf.DType = tf.float32,
+ decoding_name: Optional[str] = None,
+ extra_cache_output: bool = False):
"""Initialize sampling module."""
self.symbols_to_logits_fn = symbols_to_logits_fn
self.length_normalization_fn = length_normalization_fn
@@ -174,8 +178,13 @@ def __init__(self,
self.sample_temperature = tf.convert_to_tensor(
sample_temperature, dtype=tf.float32)
self.enable_greedy = enable_greedy
+ self.decoding_name = decoding_name
+ self.extra_cache_output = extra_cache_output
super(SamplingModule, self).__init__(
- length_normalization_fn=length_normalization_fn, dtype=dtype)
+ length_normalization_fn=length_normalization_fn,
+ dtype=dtype,
+ decoding_name=decoding_name,
+ extra_cache_output=extra_cache_output)
def _grow_alive_seq(self,
state: Dict[str, Any],
@@ -241,10 +250,13 @@ def _grow_alive_seq(self,
topk_seq = tf.concat([alive_seq, topk_ids], axis=-1)
return topk_seq, topk_log_probs, topk_ids, new_cache
- def _create_initial_state(self,
- initial_ids: tf.Tensor,
- initial_cache: Dict[str, tf.Tensor],
- batch_size: int) -> decoding_module.InitialState:
+ def _create_initial_state(
+ self,
+ initial_ids: tf.Tensor,
+ initial_cache: Dict[str, tf.Tensor],
+ batch_size: int,
+ initial_log_probs: Optional[tf.Tensor] = None
+ ) -> decoding_module.InitialState:
"""Return initial state dictionary and its shape invariants."""
for key, value in initial_cache.items():
for inner_value in tf.nest.flatten(value):
@@ -264,8 +276,11 @@ def _create_initial_state(self,
alive_seq = tf.tile(alive_seq, [1, self.max_decode_length + 1])
# Initial log probabilities with shape [batch_size, 1].
- initial_log_probs = tf.constant([[0.]], dtype=self.dtype)
- alive_log_probs = tf.tile(initial_log_probs, [batch_size, 1])
+ if initial_log_probs is None:
+ initial_log_probs = tf.constant([[0.]], dtype=self.dtype)
+ alive_log_probs = tf.tile(initial_log_probs, [batch_size, 1])
+ else:
+ alive_log_probs = initial_log_probs
alive_cache = initial_cache
@@ -294,16 +309,14 @@ def _create_initial_state(self,
decoding_module.StateKeys.CUR_INDEX:
tf.TensorShape([]),
decoding_module.StateKeys.ALIVE_SEQ:
- tf.TensorShape(
- [batch_size, self.max_decode_length + 1]),
+ tf.TensorShape([batch_size, self.max_decode_length + 1]),
decoding_module.StateKeys.ALIVE_LOG_PROBS:
tf.TensorShape([batch_size, 1]),
decoding_module.StateKeys.ALIVE_CACHE:
tf.nest.map_structure(lambda state: state.get_shape(),
alive_cache),
decoding_module.StateKeys.FINISHED_SEQ:
- tf.TensorShape(
- [batch_size, self.max_decode_length + 1]),
+ tf.TensorShape([batch_size, self.max_decode_length + 1]),
decoding_module.StateKeys.FINISHED_SCORES:
tf.TensorShape([batch_size, 1]),
decoding_module.StateKeys.FINISHED_FLAGS:
@@ -318,9 +331,8 @@ def _create_initial_state(self,
decoding_module.StateKeys.ALIVE_LOG_PROBS:
tf.TensorShape([None, 1]),
decoding_module.StateKeys.ALIVE_CACHE:
- tf.nest.map_structure(
- decoding_module.get_shape_keep_last_dim,
- alive_cache),
+ tf.nest.map_structure(decoding_module.get_shape_keep_last_dim,
+ alive_cache),
decoding_module.StateKeys.FINISHED_SEQ:
tf.TensorShape([None, None]),
decoding_module.StateKeys.FINISHED_SCORES:
@@ -329,6 +341,22 @@ def _create_initial_state(self,
tf.TensorShape([None, 1])
}
+ if self.extra_cache_output:
+ state.update(
+ {decoding_module.StateKeys.INITIAL_OUTPUT_CACHE: alive_cache})
+ if self.padded_decode:
+ state_shape_invariants.update({
+ decoding_module.StateKeys.INITIAL_OUTPUT_CACHE:
+ tf.nest.map_structure(lambda state: state.get_shape(),
+ alive_cache)
+ })
+ else:
+ state_shape_invariants.update({
+ decoding_module.StateKeys.INITIAL_OUTPUT_CACHE:
+ tf.nest.map_structure(decoding_module.get_shape_keep_last_dim,
+ alive_cache),
+ })
+
return state, state_shape_invariants
def _get_new_alive_state(self, new_seq: tf.Tensor, new_log_probs: tf.Tensor,
@@ -422,6 +450,9 @@ def _process_finished_state(
finished_scores)
finished_seq = tf.where(seq_cond, finished_seq, alive_seq)
finished_scores = tf.where(score_cond, finished_scores, alive_log_probs)
+ if self.extra_cache_output:
+ return finished_seq, finished_scores, finished_state[
+ decoding_module.StateKeys.INITIAL_OUTPUT_CACHE]
return finished_seq, finished_scores
def _continue_search(self, state) -> tf.Tensor:
diff --git a/official/nlp/modeling/ops/segment_extractor.py b/official/nlp/modeling/ops/segment_extractor.py
index 8016b65cdd4..02f405f4987 100644
--- a/official/nlp/modeling/ops/segment_extractor.py
+++ b/official/nlp/modeling/ops/segment_extractor.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Module for extracting segments from sentences in documents."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# Get a random tensor like `positions` and make some decisions
@@ -27,7 +27,7 @@ def _get_random(positions, random_fn):
return positions.with_flat_values(flat_random)
-# For every position j in a row, sample a position preceeding j or
+# For every position j in a row, sample a position preceding j or
# a position which is [0, j-1]
def _random_int_up_to(maxval, random_fn):
# Need to cast because the int kernel for uniform doesn't support bcast.
@@ -87,22 +87,22 @@ def get_sentence_order_labels(sentences,
dtype.
random_threshold: (optional) A float threshold between 0 and 1, used to
determine whether to extract a random, out-of-batch sentence or a
- suceeding sentence. Higher value favors succeeding sentence.
+ succeeding sentence. Higher value favors succeeding sentence.
random_next_threshold: (optional) A float threshold between 0 and 1, used to
determine whether to extract either a random, out-of-batch, or succeeding
- sentence or a preceeding sentence. Higher value favors preceeding
+ sentence or a preceding sentence. Higher value favors preceding
sentences.
random_fn: (optional) An op used to generate random float values.
Returns:
- a tuple of (preceeding_or_random_next, is_suceeding_or_random) where:
- preceeding_or_random_next: a `RaggedTensor` of strings with the same shape
- as `sentences` and contains either a preceeding, suceeding, or random
+ a tuple of (preceding_or_random_next, is_succeeding_or_random) where:
+ preceding_or_random_next: a `RaggedTensor` of strings with the same shape
+ as `sentences` and contains either a preceding, succeeding, or random
out-of-batch sentence respective to its counterpart in `sentences` and
- dependent on its label in `is_preceeding_or_random_next`.
- is_suceeding_or_random: a `RaggedTensor` of bool values with the
+ dependent on its label in `is_preceding_or_random_next`.
+ is_succeeding_or_random: a `RaggedTensor` of bool values with the
same shape as `sentences` and is True if it's corresponding sentence in
- `preceeding_or_random_next` is a random or suceeding sentence, False
+ `preceding_or_random_next` is a random or succeeding sentence, False
otherwise.
"""
# Create a RaggedTensor in the same shape as sentences ([doc, (sentences)])
diff --git a/official/nlp/modeling/ops/segment_extractor_test.py b/official/nlp/modeling/ops/segment_extractor_test.py
index 3fb6f566731..4fae3d0791b 100644
--- a/official/nlp/modeling/ops/segment_extractor_test.py
+++ b/official/nlp/modeling/ops/segment_extractor_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import functools
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling.ops import segment_extractor
diff --git a/official/nlp/optimization.py b/official/nlp/optimization.py
index 68351ee1aa9..0854ea51a95 100644
--- a/official/nlp/optimization.py
+++ b/official/nlp/optimization.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,14 +16,16 @@
from absl import logging
import gin
-import tensorflow as tf
-import tensorflow_addons.optimizers as tfa_optimizers
+import tensorflow as tf, tf_keras
+
+from official.modeling.optimization import lamb
from official.modeling.optimization import legacy_adamw
AdamWeightDecay = legacy_adamw.AdamWeightDecay
+LAMB = lamb.LAMB
-class WarmUp(tf.keras.optimizers.schedules.LearningRateSchedule):
+class WarmUp(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Applies a warmup schedule on a given learning rate decay schedule."""
def __init__(self,
@@ -71,13 +73,15 @@ def create_optimizer(init_lr,
num_warmup_steps,
end_lr=0.0,
optimizer_type='adamw',
- beta_1=0.9):
+ beta_1=0.9,
+ poly_power=1.0):
"""Creates an optimizer with learning rate schedule."""
# Implements linear decay of the learning rate.
- lr_schedule = tf.keras.optimizers.schedules.PolynomialDecay(
+ lr_schedule = tf_keras.optimizers.schedules.PolynomialDecay(
initial_learning_rate=init_lr,
decay_steps=num_train_steps,
- end_learning_rate=end_lr)
+ end_learning_rate=end_lr,
+ power=poly_power)
if num_warmup_steps:
lr_schedule = WarmUp(
initial_learning_rate=init_lr,
@@ -95,13 +99,14 @@ def create_optimizer(init_lr,
exclude_from_weight_decay=['LayerNorm', 'layer_norm', 'bias'])
elif optimizer_type == 'lamb':
logging.info('using Lamb optimizer')
- optimizer = tfa_optimizers.LAMB(
+ optimizer = LAMB(
learning_rate=lr_schedule,
weight_decay_rate=0.01,
beta_1=beta_1,
beta_2=0.999,
epsilon=1e-6,
- exclude_from_weight_decay=['LayerNorm', 'layer_norm', 'bias'])
+ exclude_from_weight_decay=['LayerNorm', 'layer_norm', 'bias'],
+ )
else:
raise ValueError('Unsupported optimizer type: ', optimizer_type)
diff --git a/official/nlp/serving/__init__.py b/official/nlp/serving/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/nlp/serving/__init__.py
+++ b/official/nlp/serving/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/serving/export_savedmodel.py b/official/nlp/serving/export_savedmodel.py
index 81300efba2d..c3391042d48 100644
--- a/official/nlp/serving/export_savedmodel.py
+++ b/official/nlp/serving/export_savedmodel.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -72,6 +72,12 @@ def define_flags():
flags.DEFINE_bool("convert_tpu", False, "")
flags.DEFINE_multi_integer("allowed_batch_size", None,
"Allowed batch sizes for batching ops.")
+ flags.DEFINE_integer("num_batch_threads", 4,
+ "Number of threads to do TPU batching.")
+ flags.DEFINE_integer("batch_timeout_micros", 100000,
+ "TPU batch function timeout in microseconds.")
+ flags.DEFINE_integer("max_enqueued_batches", 1000,
+ "Max number of batches in queue for TPU batching.")
def lookup_export_module(task: base_task.Task):
@@ -100,7 +106,7 @@ def create_export_module(*, task_name: Text, config_file: Text,
@dataclasses.dataclass
class Dummy(base_config.Config):
- task: task_config_cls = task_config_cls()
+ task: task_config_cls = dataclasses.field(default_factory=task_config_cls)
dummy_exp = Dummy()
dummy_exp = hyperparams.override_params_dict(
@@ -130,21 +136,30 @@ def main(_):
if FLAGS.convert_tpu:
# pylint: disable=g-import-not-at-top
- from cloud_tpu.inference_converter import converter_cli
- from cloud_tpu.inference_converter import converter_options_pb2
+ from cloud_tpu.inference_converter_v2 import converter_options_v2_pb2
+ from cloud_tpu.inference_converter_v2.python import converter
+
tpu_dir = os.path.join(export_dir, "tpu")
- options = converter_options_pb2.ConverterOptions()
+ batch_options = []
if FLAGS.allowed_batch_size is not None:
allowed_batch_sizes = sorted(FLAGS.allowed_batch_size)
- options.batch_options.num_batch_threads = 4
- options.batch_options.max_batch_size = allowed_batch_sizes[-1]
- options.batch_options.batch_timeout_micros = 100000
- options.batch_options.allowed_batch_sizes[:] = allowed_batch_sizes
- options.batch_options.max_enqueued_batches = 1000
- converter_cli.ConvertSavedModel(
- export_dir, tpu_dir, function_alias="tpu_candidate", options=options,
- graph_rewrite_only=True)
-
+ batch_option = converter_options_v2_pb2.BatchOptionsV2(
+ num_batch_threads=FLAGS.num_batch_threads,
+ max_batch_size=allowed_batch_sizes[-1],
+ batch_timeout_micros=FLAGS.batch_timeout_micros,
+ allowed_batch_sizes=allowed_batch_sizes,
+ max_enqueued_batches=FLAGS.max_enqueued_batches
+ )
+ batch_options.append(batch_option)
+
+ converter_options = converter_options_v2_pb2.ConverterOptionsV2(
+ tpu_functions=[
+ converter_options_v2_pb2.TpuFunction(function_alias="tpu_candidate")
+ ],
+ batch_options=batch_options,
+ )
+
+ converter.ConvertSavedModel(export_dir, tpu_dir, converter_options)
if __name__ == "__main__":
define_flags()
diff --git a/official/nlp/serving/export_savedmodel_test.py b/official/nlp/serving/export_savedmodel_test.py
index 1f1a82a90d2..f4d6019861e 100644
--- a/official/nlp/serving/export_savedmodel_test.py
+++ b/official/nlp/serving/export_savedmodel_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import bert
from official.nlp.configs import encoders
from official.nlp.serving import export_savedmodel
diff --git a/official/nlp/serving/export_savedmodel_util.py b/official/nlp/serving/export_savedmodel_util.py
index 8fe163c72e9..06a077b64bb 100644
--- a/official/nlp/serving/export_savedmodel_util.py
+++ b/official/nlp/serving/export_savedmodel_util.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,9 +13,9 @@
# limitations under the License.
"""Common library to export a SavedModel from the export module."""
-from typing import Dict, List, Optional, Text, Union
+from typing import Dict, List, Optional, Union, Any
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import export_base
@@ -23,11 +23,12 @@
def export(export_module: export_base.ExportModule,
- function_keys: Union[List[Text], Dict[Text, Text]],
- export_savedmodel_dir: Text,
- checkpoint_path: Optional[Text] = None,
+ function_keys: Union[List[str], Dict[str, str]],
+ export_savedmodel_dir: str,
+ checkpoint_path: Optional[str] = None,
timestamped: bool = True,
- module_key: Optional[Text] = None) -> Text:
+ module_key: Optional[str] = None,
+ checkpoint_kwargs: Optional[Dict[str, Any]] = None) -> str:
"""Exports to SavedModel format.
Args:
@@ -40,6 +41,8 @@ def export(export_module: export_base.ExportModule,
timestamped: Whether to export the savedmodel to a timestamped directory.
module_key: Optional string to identify a checkpoint object to load for the
model in the export module.
+ checkpoint_kwargs: Optional dict used as keyword args to create the
+ checkpoint object. Not used if module_key is present.
Returns:
The savedmodel directory path.
@@ -50,6 +53,8 @@ def export(export_module: export_base.ExportModule,
if module_key:
kwargs = {module_key: export_module.model}
checkpoint = tf.train.Checkpoint(**kwargs)
+ elif checkpoint_kwargs:
+ checkpoint = tf.train.Checkpoint(**checkpoint_kwargs)
else:
checkpoint = None
return export_base.export(
diff --git a/official/nlp/serving/serving_modules.py b/official/nlp/serving/serving_modules.py
index 1621e0de543..439be84acb9 100644
--- a/official/nlp/serving/serving_modules.py
+++ b/official/nlp/serving/serving_modules.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import dataclasses
from typing import Dict, List, Optional, Text
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_text as tf_text
from official.core import export_base
@@ -65,7 +65,7 @@ class Params(base_config.Config):
# ...or load preprocessing from a SavedModel at this location.
preprocessing_hub_module_url: str = ""
- def __init__(self, params, model: tf.keras.Model, inference_step=None):
+ def __init__(self, params, model: tf_keras.Model, inference_step=None):
super().__init__(params, model, inference_step)
if params.use_v2_feature_names:
self.input_word_ids_field = "input_word_ids"
@@ -188,7 +188,7 @@ def get_inference_signatures(self, function_keys: Dict[Text, Text]):
class MaskedLM(export_base.ExportModule):
"""The export module for the Bert Pretrain (MaskedLM) task."""
- def __init__(self, params, model: tf.keras.Model, inference_step=None):
+ def __init__(self, params, model: tf_keras.Model, inference_step=None):
super().__init__(params, model, inference_step)
if params.use_v2_feature_names:
self.input_word_ids_field = "input_word_ids"
@@ -269,7 +269,7 @@ class Params(base_config.Config):
parse_sequence_length: Optional[int] = None
use_v2_feature_names: bool = True
- def __init__(self, params, model: tf.keras.Model, inference_step=None):
+ def __init__(self, params, model: tf_keras.Model, inference_step=None):
super().__init__(params, model, inference_step)
if params.use_v2_feature_names:
self.input_word_ids_field = "input_word_ids"
@@ -344,7 +344,7 @@ class Params(base_config.Config):
use_v2_feature_names: bool = True
output_encoder_outputs: bool = False
- def __init__(self, params, model: tf.keras.Model, inference_step=None):
+ def __init__(self, params, model: tf_keras.Model, inference_step=None):
super().__init__(params, model, inference_step)
if params.use_v2_feature_names:
self.input_word_ids_field = "input_word_ids"
@@ -420,7 +420,7 @@ class Params(base_config.Config):
# Needs to be specified if padded_decode is True/on TPUs.
batch_size: Optional[int] = None
- def __init__(self, params, model: tf.keras.Model, inference_step=None):
+ def __init__(self, params, model: tf_keras.Model, inference_step=None):
super().__init__(params, model, inference_step)
self._sp_tokenizer = tf_text.SentencepieceTokenizer(
model=tf.io.gfile.GFile(params.sentencepiece_model_path, "rb").read(),
diff --git a/official/nlp/serving/serving_modules_test.py b/official/nlp/serving/serving_modules_test.py
index e967c606629..e0f4c6c4879 100644
--- a/official/nlp/serving/serving_modules_test.py
+++ b/official/nlp/serving/serving_modules_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from sentencepiece import SentencePieceTrainer
from official.core import export_base
diff --git a/official/nlp/tasks/__init__.py b/official/nlp/tasks/__init__.py
index cec41dff173..fa39a45d9a0 100644
--- a/official/nlp/tasks/__init__.py
+++ b/official/nlp/tasks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/tasks/dual_encoder.py b/official/nlp/tasks/dual_encoder.py
index 116456b5909..c8b320dba95 100644
--- a/official/nlp/tasks/dual_encoder.py
+++ b/official/nlp/tasks/dual_encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,9 @@
"""Dual encoder (retrieval) task."""
from typing import Mapping, Tuple
-# Import libraries
from absl import logging
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -48,8 +47,9 @@ class ModelConfig(base_config.Config):
# Defining k for calculating metrics recall@k.
eval_top_k: Tuple[int, ...] = (1, 3, 10)
- encoder: encoders.EncoderConfig = (
- encoders.EncoderConfig())
+ encoder: encoders.EncoderConfig = dataclasses.field(
+ default_factory=encoders.EncoderConfig
+ )
@dataclasses.dataclass
@@ -60,9 +60,11 @@ class DualEncoderConfig(cfg.TaskConfig):
init_checkpoint: str = ''
hub_module_url: str = ''
# Defines the concrete model config at instantiation time.
- model: ModelConfig = ModelConfig()
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ model: ModelConfig = dataclasses.field(default_factory=ModelConfig)
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
@task_factory.register_task_cls(DualEncoderConfig)
@@ -138,12 +140,12 @@ def dummy_data(_):
def build_metrics(self, training=None):
del training
- metrics = [tf.keras.metrics.Mean(name='batch_size_per_core')]
+ metrics = [tf_keras.metrics.Mean(name='batch_size_per_core')]
for k in self.task_config.model.eval_top_k:
- metrics.append(tf.keras.metrics.SparseTopKCategoricalAccuracy(
+ metrics.append(tf_keras.metrics.SparseTopKCategoricalAccuracy(
k=k, name=f'left_recall_at_{k}'))
if self.task_config.model.bidirectional:
- metrics.append(tf.keras.metrics.SparseTopKCategoricalAccuracy(
+ metrics.append(tf_keras.metrics.SparseTopKCategoricalAccuracy(
k=k, name=f'right_recall_at_{k}'))
return metrics
@@ -168,7 +170,7 @@ def process_metrics(self, metrics, labels, model_outputs):
def validation_step(self,
inputs,
- model: tf.keras.Model,
+ model: tf_keras.Model,
metrics=None) -> Mapping[str, tf.Tensor]:
outputs = model(inputs)
loss = self.build_losses(
diff --git a/official/nlp/tasks/dual_encoder_test.py b/official/nlp/tasks/dual_encoder_test.py
index 3e1a72605ae..18595837a13 100644
--- a/official/nlp/tasks/dual_encoder_test.py
+++ b/official/nlp/tasks/dual_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import configs
from official.nlp.configs import bert
@@ -53,7 +53,7 @@ def _run_task(self, config):
dataset.batch(10)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
model.save(os.path.join(self.get_temp_dir(), "saved_model"))
@@ -69,7 +69,7 @@ def test_task(self):
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
diff --git a/official/nlp/tasks/electra_task.py b/official/nlp/tasks/electra_task.py
index 9473c0d4e06..3aeed6f7df0 100644
--- a/official/nlp/tasks/electra_task.py
+++ b/official/nlp/tasks/electra_task.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""ELECTRA pretraining task (Joint Masked LM and Replaced Token Detection)."""
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -32,16 +32,22 @@
@dataclasses.dataclass
class ElectraPretrainConfig(cfg.TaskConfig):
"""The model config."""
- model: electra.ElectraPretrainerConfig = electra.ElectraPretrainerConfig(
- cls_heads=[
- bert.ClsHeadConfig(
- inner_dim=768,
- num_classes=2,
- dropout_rate=0.1,
- name='next_sentence')
- ])
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ model: electra.ElectraPretrainerConfig = dataclasses.field(
+ default_factory=lambda: electra.ElectraPretrainerConfig( # pylint: disable=g-long-lambda
+ cls_heads=[
+ bert.ClsHeadConfig(
+ inner_dim=768,
+ num_classes=2,
+ dropout_rate=0.1,
+ name='next_sentence',
+ )
+ ]
+ )
+ )
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
def _build_pretrainer(
@@ -68,7 +74,7 @@ def _build_pretrainer(
num_token_predictions=config.num_masked_tokens,
mlm_activation=tf_utils.get_activation(
generator_encoder_cfg.hidden_activation),
- mlm_initializer=tf.keras.initializers.TruncatedNormal(
+ mlm_initializer=tf_keras.initializers.TruncatedNormal(
stddev=generator_encoder_cfg.initializer_range),
classification_heads=[
layers.ClassificationHead(**cfg.as_dict()) for cfg in config.cls_heads
@@ -91,7 +97,7 @@ def build_losses(self,
metrics = dict([(metric.name, metric) for metric in metrics])
# generator lm and (optional) nsp loss.
- lm_prediction_losses = tf.keras.losses.sparse_categorical_crossentropy(
+ lm_prediction_losses = tf_keras.losses.sparse_categorical_crossentropy(
labels['masked_lm_ids'],
tf.cast(model_outputs['lm_outputs'], tf.float32),
from_logits=True)
@@ -104,7 +110,7 @@ def build_losses(self,
sentence_labels = labels['next_sentence_labels']
sentence_outputs = tf.cast(
model_outputs['sentence_outputs'], dtype=tf.float32)
- sentence_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ sentence_loss = tf_keras.losses.sparse_categorical_crossentropy(
sentence_labels, sentence_outputs, from_logits=True)
metrics['next_sentence_loss'].update_state(sentence_loss)
total_loss = mlm_loss + sentence_loss
@@ -158,19 +164,19 @@ def dummy_data(_):
def build_metrics(self, training=None):
del training
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
- tf.keras.metrics.Mean(name='lm_example_loss'),
- tf.keras.metrics.SparseCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
+ tf_keras.metrics.Mean(name='lm_example_loss'),
+ tf_keras.metrics.SparseCategoricalAccuracy(
name='discriminator_accuracy'),
]
if self.task_config.train_data.use_next_sentence_label:
metrics.append(
- tf.keras.metrics.SparseCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(
name='next_sentence_accuracy'))
- metrics.append(tf.keras.metrics.Mean(name='next_sentence_loss'))
+ metrics.append(tf_keras.metrics.Mean(name='next_sentence_loss'))
- metrics.append(tf.keras.metrics.Mean(name='discriminator_loss'))
- metrics.append(tf.keras.metrics.Mean(name='total_loss'))
+ metrics.append(tf_keras.metrics.Mean(name='discriminator_loss'))
+ metrics.append(tf_keras.metrics.Mean(name='total_loss'))
return metrics
@@ -191,8 +197,8 @@ def process_metrics(self, metrics, labels, model_outputs):
model_outputs['disc_label'], discrim_full_logits,
labels['input_mask'])
- def train_step(self, inputs, model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer, metrics):
+ def train_step(self, inputs, model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer, metrics):
"""Does forward and backward.
Args:
@@ -221,7 +227,7 @@ def train_step(self, inputs, model: tf.keras.Model,
self.process_metrics(metrics, inputs, outputs)
return {self.loss: loss}
- def validation_step(self, inputs, model: tf.keras.Model, metrics):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics):
"""Validatation step.
Args:
diff --git a/official/nlp/tasks/electra_task_test.py b/official/nlp/tasks/electra_task_test.py
index 4018c9220ac..f012e6087ba 100644
--- a/official/nlp/tasks/electra_task_test.py
+++ b/official/nlp/tasks/electra_task_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for official.nlp.tasks.electra_task."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import bert
from official.nlp.configs import electra
@@ -51,7 +51,7 @@ def test_task(self):
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
diff --git a/official/nlp/tasks/masked_lm.py b/official/nlp/tasks/masked_lm.py
index f784b141676..923896a9ed3 100644
--- a/official/nlp/tasks/masked_lm.py
+++ b/official/nlp/tasks/masked_lm.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Masked language task."""
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -31,15 +31,25 @@
@dataclasses.dataclass
class MaskedLMConfig(cfg.TaskConfig):
"""The model config."""
- model: bert.PretrainerConfig = bert.PretrainerConfig(cls_heads=[
- bert.ClsHeadConfig(
- inner_dim=768, num_classes=2, dropout_rate=0.1, name='next_sentence')
- ])
+ model: bert.PretrainerConfig = dataclasses.field(
+ default_factory=lambda: bert.PretrainerConfig( # pylint: disable=g-long-lambda
+ cls_heads=[
+ bert.ClsHeadConfig(
+ inner_dim=768,
+ num_classes=2,
+ dropout_rate=0.1,
+ name='next_sentence',
+ )
+ ]
+ )
+ )
# TODO(b/154564893): Mathematically, scale_loss should be True.
# However, it works better with scale_loss being False.
scale_loss: bool = False
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
@task_factory.register_task_cls(MaskedLMConfig)
@@ -58,7 +68,7 @@ def build_model(self, params=None):
] if config.cls_heads else []
return models.BertPretrainerV2(
mlm_activation=tf_utils.get_activation(config.mlm_activation),
- mlm_initializer=tf.keras.initializers.TruncatedNormal(
+ mlm_initializer=tf_keras.initializers.TruncatedNormal(
stddev=config.mlm_initializer_range),
encoder_network=encoder_network,
classification_heads=cls_heads)
@@ -70,7 +80,7 @@ def build_losses(self,
aux_losses=None) -> tf.Tensor:
with tf.name_scope('MaskedLMTask/losses'):
metrics = dict([(metric.name, metric) for metric in metrics])
- lm_prediction_losses = tf.keras.losses.sparse_categorical_crossentropy(
+ lm_prediction_losses = tf_keras.losses.sparse_categorical_crossentropy(
labels['masked_lm_ids'],
tf.cast(model_outputs['mlm_logits'], tf.float32),
from_logits=True)
@@ -85,7 +95,7 @@ def build_losses(self,
sentence_outputs = tf.cast(
model_outputs['next_sentence'], dtype=tf.float32)
sentence_loss = tf.reduce_mean(
- tf.keras.losses.sparse_categorical_crossentropy(
+ tf_keras.losses.sparse_categorical_crossentropy(
sentence_labels, sentence_outputs, from_logits=True))
metrics['next_sentence_loss'].update_state(sentence_loss)
total_loss = mlm_loss + sentence_loss
@@ -123,15 +133,15 @@ def dummy_data(_):
def build_metrics(self, training=None):
del training
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
- tf.keras.metrics.Mean(name='lm_example_loss')
+ tf_keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
+ tf_keras.metrics.Mean(name='lm_example_loss')
]
# TODO(hongkuny): rethink how to manage metrics creation with heads.
if self.task_config.train_data.use_next_sentence_label:
metrics.append(
- tf.keras.metrics.SparseCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(
name='next_sentence_accuracy'))
- metrics.append(tf.keras.metrics.Mean(name='next_sentence_loss'))
+ metrics.append(tf_keras.metrics.Mean(name='next_sentence_loss'))
return metrics
def process_metrics(self, metrics, labels, model_outputs):
@@ -145,8 +155,8 @@ def process_metrics(self, metrics, labels, model_outputs):
metrics['next_sentence_accuracy'].update_state(
labels['next_sentence_labels'], model_outputs['next_sentence'])
- def train_step(self, inputs, model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer, metrics):
+ def train_step(self, inputs, model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer, metrics):
"""Does forward and backward.
Args:
@@ -172,14 +182,14 @@ def train_step(self, inputs, model: tf.keras.Model,
scaled_loss = loss / tf.distribute.get_strategy().num_replicas_in_sync
tvars = model.trainable_variables
if self.task_config.scale_loss:
- grads = tape.gradient(scaled_loss, tvars)
+ grads = tape.gradient(scaled_loss, tvars) # pyrefly: ignore[unbound-name]
else:
grads = tape.gradient(loss, tvars)
optimizer.apply_gradients(list(zip(grads, tvars)))
self.process_metrics(metrics, inputs, outputs)
return {self.loss: loss}
- def validation_step(self, inputs, model: tf.keras.Model, metrics):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics):
"""Validatation step.
Args:
diff --git a/official/nlp/tasks/masked_lm_determinism_test.py b/official/nlp/tasks/masked_lm_determinism_test.py
new file mode 100644
index 00000000000..86c124b36bf
--- /dev/null
+++ b/official/nlp/tasks/masked_lm_determinism_test.py
@@ -0,0 +1,103 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests that masked LM models are deterministic when determinism is enabled."""
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.configs import bert
+from official.nlp.configs import encoders
+from official.nlp.data import pretrain_dataloader
+from official.nlp.tasks import masked_lm
+
+
+class MLMTaskTest(tf.test.TestCase):
+
+ def _build_dataset(self, params, vocab_size):
+ def dummy_data(_):
+ dummy_ids = tf.random.uniform((1, params.seq_length), maxval=vocab_size,
+ dtype=tf.int32)
+ dummy_mask = tf.ones((1, params.seq_length), dtype=tf.int32)
+ dummy_type_ids = tf.zeros((1, params.seq_length), dtype=tf.int32)
+ dummy_lm = tf.zeros((1, params.max_predictions_per_seq), dtype=tf.int32)
+ return dict(
+ input_word_ids=dummy_ids,
+ input_mask=dummy_mask,
+ input_type_ids=dummy_type_ids,
+ masked_lm_positions=dummy_lm,
+ masked_lm_ids=dummy_lm,
+ masked_lm_weights=tf.cast(dummy_lm, dtype=tf.float32),
+ next_sentence_labels=tf.zeros((1, 1), dtype=tf.int32))
+
+ dataset = tf.data.Dataset.range(1)
+ dataset = dataset.repeat()
+ dataset = dataset.map(
+ dummy_data, num_parallel_calls=tf.data.experimental.AUTOTUNE)
+ return dataset
+
+ def _build_and_run_model(self, config, num_steps=5):
+ task = masked_lm.MaskedLMTask(config)
+ model = task.build_model()
+ metrics = task.build_metrics()
+ dataset = self._build_dataset(config.train_data,
+ config.model.encoder.get().vocab_size)
+
+ iterator = iter(dataset)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
+
+ # Run training
+ for _ in range(num_steps):
+ logs = task.train_step(next(iterator), model, optimizer, metrics=metrics)
+ for metric in metrics:
+ logs[metric.name] = metric.result()
+
+ # Run validation
+ validation_logs = task.validation_step(next(iterator), model,
+ metrics=metrics)
+ for metric in metrics:
+ validation_logs[metric.name] = metric.result()
+
+ return logs, validation_logs, model.weights
+
+ def test_task_determinism(self):
+ config = masked_lm.MaskedLMConfig(
+ init_checkpoint=self.get_temp_dir(),
+ scale_loss=True,
+ model=bert.PretrainerConfig(
+ encoder=encoders.EncoderConfig(
+ bert=encoders.BertEncoderConfig(vocab_size=30522,
+ num_layers=1)),
+ cls_heads=[
+ bert.ClsHeadConfig(
+ inner_dim=10, num_classes=2, name="next_sentence")
+ ]),
+ train_data=pretrain_dataloader.BertPretrainDataConfig(
+ max_predictions_per_seq=20,
+ seq_length=128,
+ global_batch_size=1))
+
+ tf_keras.utils.set_random_seed(1)
+ logs1, validation_logs1, weights1 = self._build_and_run_model(config)
+ tf_keras.utils.set_random_seed(1)
+ logs2, validation_logs2, weights2 = self._build_and_run_model(config)
+
+ self.assertEqual(logs1["loss"], logs2["loss"])
+ self.assertEqual(validation_logs1["loss"], validation_logs2["loss"])
+ for weight1, weight2 in zip(weights1, weights2):
+ self.assertAllEqual(weight1, weight2)
+
+
+if __name__ == "__main__":
+ tf.config.experimental.enable_op_determinism()
+ tf.test.main()
diff --git a/official/nlp/tasks/masked_lm_test.py b/official/nlp/tasks/masked_lm_test.py
index 221fa6c0978..ac4418d2678 100644
--- a/official/nlp/tasks/masked_lm_test.py
+++ b/official/nlp/tasks/masked_lm_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for official.nlp.tasks.masked_lm."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import bert
from official.nlp.configs import encoders
@@ -47,7 +47,7 @@ def test_task(self):
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
diff --git a/official/nlp/tasks/question_answering.py b/official/nlp/tasks/question_answering.py
index d9c7508fe86..e5647741403 100644
--- a/official/nlp/tasks/question_answering.py
+++ b/official/nlp/tasks/question_answering.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -41,7 +41,9 @@
@dataclasses.dataclass
class ModelConfig(base_config.Config):
"""A base span labeler configuration."""
- encoder: encoders.EncoderConfig = encoders.EncoderConfig()
+ encoder: encoders.EncoderConfig = dataclasses.field(
+ default_factory=encoders.EncoderConfig
+ )
@dataclasses.dataclass
@@ -53,9 +55,11 @@ class QuestionAnsweringConfig(cfg.TaskConfig):
n_best_size: int = 20
max_answer_length: int = 30
null_score_diff_threshold: float = 0.0
- model: ModelConfig = ModelConfig()
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ model: ModelConfig = dataclasses.field(default_factory=ModelConfig)
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
@dataclasses.dataclass
@@ -107,7 +111,7 @@ def build_model(self):
encoder_cfg = self.task_config.model.encoder.get()
return models.BertSpanLabeler(
network=encoder_network,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range))
def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
@@ -115,11 +119,11 @@ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
end_positions = labels['end_positions']
start_logits, end_logits = model_outputs
- start_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ start_loss = tf_keras.losses.sparse_categorical_crossentropy(
start_positions,
tf.cast(start_logits, dtype=tf.float32),
from_logits=True)
- end_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ end_loss = tf_keras.losses.sparse_categorical_crossentropy(
end_positions, tf.cast(end_logits, dtype=tf.float32), from_logits=True)
loss = (tf.reduce_mean(start_loss) + tf.reduce_mean(end_loss)) / 2
@@ -220,9 +224,9 @@ def build_metrics(self, training=None):
return []
# TODO(lehou): a list of metrics doesn't work the same as in compile/fit.
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(
name='start_position_accuracy'),
- tf.keras.metrics.SparseCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(
name='end_position_accuracy'),
]
return metrics
@@ -244,7 +248,7 @@ def process_compiled_metrics(self, compiled_metrics, labels, model_outputs):
'end_positions': end_logits
})
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
features, _ = inputs
unique_ids = features.pop('unique_ids')
model_outputs = self.inference_step(features, model)
@@ -347,7 +351,7 @@ def build_model(self):
network=encoder_network,
start_n_top=self.task_config.n_best_size,
end_n_top=self.task_config.n_best_size,
- initializer=tf.keras.initializers.RandomNormal(
+ initializer=tf_keras.initializers.RandomNormal(
stddev=encoder_cfg.initializer_range))
def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
@@ -364,7 +368,7 @@ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
start_positions, start_logits)
end_loss = tf.nn.sparse_softmax_cross_entropy_with_logits(
end_positions, end_logits)
- is_impossible_loss = tf.keras.losses.binary_crossentropy(
+ is_impossible_loss = tf_keras.losses.binary_crossentropy(
is_impossible, class_logits, from_logits=True)
loss = (tf.reduce_mean(start_loss) + tf.reduce_mean(end_loss)) / 2
@@ -408,7 +412,7 @@ def _dummy_data(self, params, _):
is_impossible=zero)
return x, y
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
features, _ = inputs
unique_ids = features.pop('unique_ids')
model_outputs = self.inference_step(features, model)
@@ -455,7 +459,7 @@ def aggregate_logs(self, state=None, step_outputs=None):
def predict(task: QuestionAnsweringTask, params: cfg.DataConfig,
- model: tf.keras.Model):
+ model: tf_keras.Model):
"""Predicts on the input data.
Args:
diff --git a/official/nlp/tasks/question_answering_test.py b/official/nlp/tasks/question_answering_test.py
index cc50592a829..7aca34beb8f 100644
--- a/official/nlp/tasks/question_answering_test.py
+++ b/official/nlp/tasks/question_answering_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import bert
from official.nlp.configs import encoders
@@ -90,7 +90,7 @@ def _run_task(self, config):
train_dataset = task.build_inputs(config.train_data)
train_iterator = iter(train_dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(train_iterator), model, optimizer, metrics=metrics)
val_dataset = task.build_inputs(config.validation_data)
@@ -101,7 +101,8 @@ def _run_task(self, config):
logs = task.aggregate_logs(step_outputs=logs)
metrics = task.reduce_aggregated_logs(logs)
self.assertIn("final_f1", metrics)
- model.save(os.path.join(self.get_temp_dir(), "saved_model"))
+ model.save(os.path.join(self.get_temp_dir(), "saved_model.keras"),
+ save_format="keras")
@parameterized.parameters(
itertools.product(
@@ -109,6 +110,7 @@ def _run_task(self, config):
("WordPiece", "SentencePiece"),
))
def test_task(self, version_2_with_negative, tokenization):
+ del tokenization
# Saves a checkpoint.
pretrain_cfg = bert.PretrainerConfig(
encoder=self._encoder_config,
@@ -135,7 +137,7 @@ def _export_bert_tfhub(self):
bert=encoders.BertEncoderConfig(vocab_size=30522, num_layers=1)))
encoder_inputs_dict = {x.name: x for x in encoder.inputs}
encoder_output_dict = encoder(encoder_inputs_dict)
- core_model = tf.keras.Model(
+ core_model = tf_keras.Model(
inputs=encoder_inputs_dict, outputs=encoder_output_dict)
hub_destination = os.path.join(self.get_temp_dir(), "hub")
core_model.save(hub_destination, include_optimizer=False, save_format="tf")
@@ -238,7 +240,7 @@ def _run_task(self, config):
train_dataset = task.build_inputs(config.train_data)
train_iterator = iter(train_dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(train_iterator), model, optimizer, metrics=metrics)
val_dataset = task.build_inputs(config.validation_data)
diff --git a/official/nlp/tasks/sentence_prediction.py b/official/nlp/tasks/sentence_prediction.py
index dd6c5514441..efcb32535cd 100644
--- a/official/nlp/tasks/sentence_prediction.py
+++ b/official/nlp/tasks/sentence_prediction.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
import orbit
from scipy import stats
from sklearn import metrics as sklearn_metrics
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -34,7 +34,7 @@
from official.nlp.tasks import utils
METRIC_TYPES = frozenset(
- ['accuracy', 'matthews_corrcoef', 'pearson_spearman_corr'])
+ ['accuracy', 'f1', 'matthews_corrcoef', 'pearson_spearman_corr'])
@dataclasses.dataclass
@@ -42,7 +42,7 @@ class ModelConfig(base_config.Config):
"""A classifier/regressor configuration."""
num_classes: int = 0
use_encoder_pooler: bool = False
- encoder: encoders.EncoderConfig = encoders.EncoderConfig()
+ encoder: encoders.EncoderConfig = dataclasses.field(default_factory=encoders.EncoderConfig)
@dataclasses.dataclass
@@ -55,9 +55,9 @@ class SentencePredictionConfig(cfg.TaskConfig):
hub_module_url: str = ''
metric_type: str = 'accuracy'
# Defines the concrete model config at instantiation time.
- model: ModelConfig = ModelConfig()
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ model: ModelConfig = dataclasses.field(default_factory=ModelConfig)
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
@task_factory.register_task_cls(SentencePredictionConfig)
@@ -88,22 +88,25 @@ def build_model(self):
return models.XLNetClassifier(
network=encoder_network,
num_classes=self.task_config.model.num_classes,
- initializer=tf.keras.initializers.RandomNormal(
+ initializer=tf_keras.initializers.RandomNormal(
stddev=encoder_cfg.initializer_range))
else:
return models.BertClassifier(
network=encoder_network,
num_classes=self.task_config.model.num_classes,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
use_encoder_pooler=self.task_config.model.use_encoder_pooler)
def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
label_ids = labels[self.label_field]
if self.task_config.model.num_classes == 1:
- loss = tf.keras.losses.mean_squared_error(label_ids, model_outputs)
+ # Reshape to avoid silent broadcasting bugs in Keras.
+ loss = tf_keras.losses.mean_squared_error(
+ tf.reshape(label_ids, [-1]), tf.reshape(model_outputs, [-1])
+ )
else:
- loss = tf.keras.losses.sparse_categorical_crossentropy(
+ loss = tf_keras.losses.sparse_categorical_crossentropy(
label_ids, tf.cast(model_outputs, tf.float32), from_logits=True)
if aux_losses:
@@ -139,15 +142,15 @@ def dummy_data(_):
def build_metrics(self, training=None):
del training
if self.task_config.model.num_classes == 1:
- metrics = [tf.keras.metrics.MeanSquaredError()]
+ metrics = [tf_keras.metrics.MeanSquaredError()]
elif self.task_config.model.num_classes == 2:
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
- tf.keras.metrics.AUC(name='auc', curve='PR'),
+ tf_keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
+ tf_keras.metrics.AUC(name='auc', curve='PR'),
]
else:
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
+ tf_keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
]
return metrics
@@ -164,15 +167,18 @@ def process_metrics(self, metrics, labels, model_outputs):
def process_compiled_metrics(self, compiled_metrics, labels, model_outputs):
compiled_metrics.update_state(labels[self.label_field], model_outputs)
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
- if self.metric_type == 'accuracy':
- return super(SentencePredictionTask,
- self).validation_step(inputs, model, metrics)
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
features, labels = inputs, inputs
outputs = self.inference_step(features, model)
loss = self.build_losses(
labels=labels, model_outputs=outputs, aux_losses=model.losses)
logs = {self.loss: loss}
+ if metrics:
+ self.process_metrics(metrics, labels, outputs)
+ if model.compiled_metrics:
+ self.process_compiled_metrics(model.compiled_metrics, labels, outputs)
+ logs.update({m.name: m.result() for m in metrics or []})
+ logs.update({m.name: m.result() for m in model.metrics})
if self.metric_type == 'matthews_corrcoef':
logs.update({
'sentence_prediction': # Ensure one prediction along batch dimension.
@@ -180,7 +186,7 @@ def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
'labels':
labels[self.label_field],
})
- if self.metric_type == 'pearson_spearman_corr':
+ else:
logs.update({
'sentence_prediction': outputs,
'labels': labels[self.label_field],
@@ -193,27 +199,29 @@ def aggregate_logs(self, state=None, step_outputs=None):
if state is None:
state = {'sentence_prediction': [], 'labels': []}
state['sentence_prediction'].append(
- np.concatenate([v.numpy() for v in step_outputs['sentence_prediction']],
+ np.concatenate([v.numpy() for v in step_outputs['sentence_prediction']], # pyrefly: ignore[unsupported-operation]
axis=0))
state['labels'].append(
- np.concatenate([v.numpy() for v in step_outputs['labels']], axis=0))
+ np.concatenate([v.numpy() for v in step_outputs['labels']], axis=0)) # pyrefly: ignore[unsupported-operation]
return state
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
if self.metric_type == 'accuracy':
return None
+
+ preds = np.concatenate(aggregated_logs['sentence_prediction'], axis=0)
+ labels = np.concatenate(aggregated_logs['labels'], axis=0)
+ if self.metric_type == 'f1':
+ preds = np.argmax(preds, axis=1)
+ return {self.metric_type: sklearn_metrics.f1_score(labels, preds)}
elif self.metric_type == 'matthews_corrcoef':
- preds = np.concatenate(aggregated_logs['sentence_prediction'], axis=0)
preds = np.reshape(preds, -1)
- labels = np.concatenate(aggregated_logs['labels'], axis=0)
labels = np.reshape(labels, -1)
return {
self.metric_type: sklearn_metrics.matthews_corrcoef(preds, labels)
}
elif self.metric_type == 'pearson_spearman_corr':
- preds = np.concatenate(aggregated_logs['sentence_prediction'], axis=0)
preds = np.reshape(preds, -1)
- labels = np.concatenate(aggregated_logs['labels'], axis=0)
labels = np.reshape(labels, -1)
pearson_corr = stats.pearsonr(preds, labels)[0]
spearman_corr = stats.spearmanr(preds, labels)[0]
@@ -249,7 +257,7 @@ def initialize(self, model):
def predict(task: SentencePredictionTask,
params: cfg.DataConfig,
- model: tf.keras.Model,
+ model: tf_keras.Model,
params_aug: Optional[cfg.DataConfig] = None,
test_time_aug_wgt: float = 0.3) -> List[Union[int, float]]:
"""Predicts on the input data.
diff --git a/official/nlp/tasks/sentence_prediction_test.py b/official/nlp/tasks/sentence_prediction_test.py
index f3fb8b4e08d..67d1f66670c 100644
--- a/official/nlp/tasks/sentence_prediction_test.py
+++ b/official/nlp/tasks/sentence_prediction_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import bert
from official.nlp.configs import encoders
@@ -83,7 +83,7 @@ def _run_task(self, config):
functools.partial(task.build_inputs, config.train_data))
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(learning_rate=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
model.save(os.path.join(self.get_temp_dir(), "saved_model"))
return task.validation_step(next(iterator), model, metrics=metrics)
@@ -120,7 +120,7 @@ def test_task(self, init_cls_pooler):
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(learning_rate=0.1)
task.initialize(model)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
@@ -144,14 +144,14 @@ def test_metrics_and_losses(self, num_classes):
model = task.build_model()
metrics = task.build_metrics()
if num_classes == 1:
- self.assertIsInstance(metrics[0], tf.keras.metrics.MeanSquaredError)
+ self.assertIsInstance(metrics[0], tf_keras.metrics.MeanSquaredError)
else:
self.assertIsInstance(metrics[0],
- tf.keras.metrics.SparseCategoricalAccuracy)
+ tf_keras.metrics.SparseCategoricalAccuracy)
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(learning_rate=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
logs = task.validation_step(next(iterator), model, metrics=metrics)
@@ -161,8 +161,26 @@ def test_metrics_and_losses(self, num_classes):
else:
self.assertLess(loss, 1.0)
+ def test_build_losses_broadcasting_regression(self):
+ config = sentence_prediction.SentencePredictionConfig(
+ init_checkpoint=self.get_temp_dir(),
+ model=self.get_model_config(num_classes=1),
+ train_data=self._train_data_config,
+ )
+ task = sentence_prediction.SentencePredictionTask(config)
+
+ labels = {"label_ids": tf.constant([1.0, 2.0], dtype=tf.float32)}
+ model_outputs = tf.constant([[1.0], [4.0]], dtype=tf.float32)
+
+ loss = task.build_losses(labels, model_outputs)
+ # True MSE between [1.0, 2.0] and [[1.0], [4.0]] is (0^2 + 2^2) / 2 = 2.0.
+ # Without the reshape fix, Keras broadcasting results in an incorrect loss
+ # of 3.5.
+ self.assertAllClose(loss, 2.0)
+
@parameterized.parameters(("matthews_corrcoef", 2),
- ("pearson_spearman_corr", 1))
+ ("pearson_spearman_corr", 1),
+ ("f1", 2))
def test_np_metrics(self, metric_type, num_classes):
config = sentence_prediction.SentencePredictionConfig(
metric_type=metric_type,
@@ -218,7 +236,7 @@ def _export_bert_tfhub(self):
bert=encoders.BertEncoderConfig(vocab_size=30522, num_layers=1)))
encoder_inputs_dict = {x.name: x for x in encoder.inputs}
encoder_output_dict = encoder(encoder_inputs_dict)
- core_model = tf.keras.Model(
+ core_model = tf_keras.Model(
inputs=encoder_inputs_dict, outputs=encoder_output_dict)
hub_destination = os.path.join(self.get_temp_dir(), "hub")
core_model.save(hub_destination, include_optimizer=False, save_format="tf")
diff --git a/official/nlp/tasks/tagging.py b/official/nlp/tasks/tagging.py
index 5f2a3f64fc2..4ce34a6e98b 100644
--- a/official/nlp/tasks/tagging.py
+++ b/official/nlp/tasks/tagging.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from seqeval import metrics as seqeval_metrics
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -35,7 +35,7 @@
@dataclasses.dataclass
class ModelConfig(base_config.Config):
"""A base span labeler configuration."""
- encoder: encoders.EncoderConfig = encoders.EncoderConfig()
+ encoder: encoders.EncoderConfig = dataclasses.field(default_factory=encoders.EncoderConfig)
head_dropout: float = 0.1
head_initializer_range: float = 0.02
@@ -46,7 +46,7 @@ class TaggingConfig(cfg.TaskConfig):
# At most one of `init_checkpoint` and `hub_module_url` can be specified.
init_checkpoint: str = ''
hub_module_url: str = ''
- model: ModelConfig = ModelConfig()
+ model: ModelConfig = dataclasses.field(default_factory=ModelConfig)
# The real class names, the order of which should match real label id.
# Note that a word may be tokenized into multiple word_pieces tokens, and
@@ -54,8 +54,8 @@ class TaggingConfig(cfg.TaskConfig):
# of the word, and a negative label id is assigned to the remaining tokens.
# The negative label id will not contribute to loss and metrics.
class_names: Optional[List[str]] = None
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
def _masked_labels_and_weights(y_true):
@@ -95,7 +95,7 @@ def build_model(self):
return models.BertTokenClassifier(
network=encoder_network,
num_classes=len(self.task_config.class_names),
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=self.task_config.model.head_initializer_range),
dropout_rate=self.task_config.model.head_dropout,
output='logits',
@@ -104,7 +104,7 @@ def build_model(self):
def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
logits = tf.cast(model_outputs['logits'], tf.float32)
masked_labels, masked_weights = _masked_labels_and_weights(labels)
- loss = tf.keras.losses.sparse_categorical_crossentropy(
+ loss = tf_keras.losses.sparse_categorical_crossentropy(
masked_labels, logits, from_logits=True)
numerator_loss = tf.reduce_sum(loss * masked_weights)
denominator_loss = tf.reduce_sum(masked_weights)
@@ -138,13 +138,13 @@ def dummy_data(_):
return data_loader_factory.get_data_loader(params).load(input_context)
- def inference_step(self, inputs, model: tf.keras.Model):
+ def inference_step(self, inputs, model: tf_keras.Model):
"""Performs the forward step."""
logits = model(inputs, training=False)['logits']
return {'logits': logits,
'predict_ids': tf.argmax(logits, axis=-1, output_type=tf.int32)}
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
"""Validatation step.
Args:
@@ -185,8 +185,8 @@ def id_to_class_name(batched_ids):
# Convert id to class names, because `seqeval_metrics` relies on the class
# name to decide IOB tags.
- state['predict_class'].extend(id_to_class_name(step_outputs['predict_ids']))
- state['label_class'].extend(id_to_class_name(step_outputs['label_ids']))
+ state['predict_class'].extend(id_to_class_name(step_outputs['predict_ids'])) # pyrefly: ignore[unsupported-operation]
+ state['label_class'].extend(id_to_class_name(step_outputs['label_ids'])) # pyrefly: ignore[unsupported-operation]
return state
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
@@ -207,7 +207,7 @@ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
def predict(task: TaggingTask,
params: cfg.DataConfig,
- model: tf.keras.Model) -> List[Tuple[int, int, List[int]]]:
+ model: tf_keras.Model) -> List[Tuple[int, int, List[int]]]:
"""Predicts on the input data.
Args:
diff --git a/official/nlp/tasks/tagging_test.py b/official/nlp/tasks/tagging_test.py
index e888abb5614..9652acf4d32 100644
--- a/official/nlp/tasks/tagging_test.py
+++ b/official/nlp/tasks/tagging_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import encoders
from official.nlp.data import tagging_dataloader
@@ -67,7 +67,7 @@ def _run_task(self, config):
functools.partial(task.build_inputs, config.train_data))
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
model.save(os.path.join(self.get_temp_dir(), "saved_model"))
@@ -89,7 +89,7 @@ def test_task(self):
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
task.initialize(model)
@@ -100,7 +100,7 @@ def _export_bert_tfhub(self):
bert=encoders.BertEncoderConfig(vocab_size=30522, num_layers=1)))
encoder_inputs_dict = {x.name: x for x in encoder.inputs}
encoder_output_dict = encoder(encoder_inputs_dict)
- core_model = tf.keras.Model(
+ core_model = tf_keras.Model(
inputs=encoder_inputs_dict, outputs=encoder_output_dict)
hub_destination = os.path.join(self.get_temp_dir(), "hub")
core_model.save(hub_destination, include_optimizer=False, save_format="tf")
diff --git a/official/nlp/tasks/translation.py b/official/nlp/tasks/translation.py
index bb9591d4617..e963122a1d9 100644
--- a/official/nlp/tasks/translation.py
+++ b/official/nlp/tasks/translation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl import logging
import sacrebleu
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_text as tftxt
from official.core import base_task
@@ -98,8 +98,8 @@ class EncDecoder(base_config.Config):
@dataclasses.dataclass
class ModelConfig(base_config.Config):
"""A base Seq2Seq model configuration."""
- encoder: EncDecoder = EncDecoder()
- decoder: EncDecoder = EncDecoder()
+ encoder: EncDecoder = dataclasses.field(default_factory=EncDecoder)
+ decoder: EncDecoder = dataclasses.field(default_factory=EncDecoder)
embedding_width: int = 512
dropout_rate: float = 0.1
@@ -117,9 +117,11 @@ class ModelConfig(base_config.Config):
@dataclasses.dataclass
class TranslationConfig(cfg.TaskConfig):
"""The translation task config."""
- model: ModelConfig = ModelConfig()
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ model: ModelConfig = dataclasses.field(default_factory=ModelConfig)
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
# Tokenization
sentencepiece_model_path: str = ""
# Evaluation.
@@ -201,7 +203,7 @@ def __init__(self, params: cfg.TaskConfig, logging_dir=None, name=None):
self._references, self._tf_record_input_path = write_test_record(
params.validation_data, self.logging_dir)
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
"""Creates model architecture.
Returns:
@@ -264,8 +266,8 @@ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
def train_step(self,
inputs,
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
metrics=None):
"""Does forward and backward.
@@ -290,13 +292,13 @@ def train_step(self,
# For mixed precision, when a LossScaleOptimizer is used, the loss is
# scaled to avoid numeric underflow.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
logs = {self.loss: loss}
@@ -304,7 +306,7 @@ def train_step(self,
self.process_metrics(metrics, inputs["targets"], outputs)
return logs
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
unique_ids = inputs.pop("unique_id")
# Validation loss
outputs = model(inputs, training=False)
@@ -328,9 +330,9 @@ def aggregate_logs(self, state=None, step_outputs=None):
state = {}
for in_token_ids, out_token_ids, unique_ids in zip(
- step_outputs["inputs"],
- step_outputs["outputs"],
- step_outputs["unique_ids"]):
+ step_outputs["inputs"], # pyrefly: ignore[unsupported-operation]
+ step_outputs["outputs"], # pyrefly: ignore[unsupported-operation]
+ step_outputs["unique_ids"]): # pyrefly: ignore[unsupported-operation]
for in_ids, out_ids, u_id in zip(
in_token_ids.numpy(), out_token_ids.numpy(), unique_ids.numpy()):
state[u_id] = (in_ids, out_ids)
diff --git a/official/nlp/tasks/translation_test.py b/official/nlp/tasks/translation_test.py
index 30cd8b7f352..2252bc7e20b 100644
--- a/official/nlp/tasks/translation_test.py
+++ b/official/nlp/tasks/translation_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from sentencepiece import SentencePieceTrainer
from official.nlp.data import wmt_dataloader
@@ -97,7 +97,7 @@ def test_task(self):
model = task.build_model()
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer)
def test_no_sentencepiece_path(self):
diff --git a/official/nlp/tasks/utils.py b/official/nlp/tasks/utils.py
index 44295e6590b..4f911e6d6b8 100644
--- a/official/nlp/tasks/utils.py
+++ b/official/nlp/tasks/utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,24 +16,24 @@
from typing import Any, Callable
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_hub as hub
-def get_encoder_from_hub(hub_model_path: str) -> tf.keras.Model:
+def get_encoder_from_hub(hub_model_path: str) -> tf_keras.Model:
"""Gets an encoder from hub.
Args:
hub_model_path: The path to the tfhub model.
Returns:
- A tf.keras.Model.
+ A tf_keras.Model.
"""
- input_word_ids = tf.keras.layers.Input(
+ input_word_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_mask')
- input_type_ids = tf.keras.layers.Input(
+ input_type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
hub_layer = hub.KerasLayer(hub_model_path, trainable=True)
output_dict = {}
@@ -43,7 +43,7 @@ def get_encoder_from_hub(hub_model_path: str) -> tf.keras.Model:
input_type_ids=input_type_ids)
output_dict = hub_layer(dict_input)
- return tf.keras.Model(inputs=dict_input, outputs=output_dict)
+ return tf_keras.Model(inputs=dict_input, outputs=output_dict)
def predict(predict_step_fn: Callable[[Any], Any],
diff --git a/official/nlp/tools/__init__.py b/official/nlp/tools/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/nlp/tools/__init__.py
+++ b/official/nlp/tools/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/tools/export_tfhub.py b/official/nlp/tools/export_tfhub.py
index e81dabd32fd..ab6ea74e9ae 100644
--- a/official/nlp/tools/export_tfhub.py
+++ b/official/nlp/tools/export_tfhub.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/tools/export_tfhub_lib.py b/official/nlp/tools/export_tfhub_lib.py
index e1a428b67bd..87539af5282 100644
--- a/official/nlp/tools/export_tfhub_lib.py
+++ b/official/nlp/tools/export_tfhub_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,12 +21,11 @@
from typing import Optional, Text, Tuple
-# Import libraries
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=g-direct-tensorflow-import TODO(b/175369555): Remove these.
from tensorflow.core.protobuf import saved_model_pb2
-from tensorflow.python.ops import control_flow_ops
+from tensorflow.python.ops import control_flow_assert
# pylint: enable=g-direct-tensorflow-import
from official.legacy.bert import configs
from official.modeling import tf_utils
@@ -49,7 +48,7 @@ def get_bert_encoder(bert_config):
attention_dropout_rate=bert_config.attention_probs_dropout_prob,
max_sequence_length=bert_config.max_position_embeddings,
type_vocab_size=bert_config.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
embedding_width=bert_config.embedding_size,
dict_outputs=True)
@@ -80,7 +79,7 @@ def _create_model(
bert_config: Optional[configs.BertConfig] = None,
encoder_config: Optional[encoders.EncoderConfig] = None,
with_mlm: bool,
-) -> Tuple[tf.keras.Model, tf.keras.Model]:
+) -> Tuple[tf_keras.Model, tf_keras.Model]:
"""Creates the model to export and the model to restore the checkpoint.
Args:
@@ -119,7 +118,7 @@ def _create_model(
# For interchangeability with other text representations,
# add "default" as an alias for BERT's whole-input reptesentations.
encoder_output_dict["default"] = encoder_output_dict["pooled_output"]
- core_model = tf.keras.Model(
+ core_model = tf_keras.Model(
inputs=encoder_inputs_dict, outputs=encoder_output_dict)
if with_mlm:
@@ -133,9 +132,12 @@ def _create_model(
encoder_network=encoder,
mlm_activation=tf_utils.get_activation(hidden_act))
- pretrainer_inputs_dict = {x.name: x for x in pretrainer.inputs}
+ if isinstance(pretrainer.inputs, dict):
+ pretrainer_inputs_dict = pretrainer.inputs
+ else:
+ pretrainer_inputs_dict = {x.name: x for x in pretrainer.inputs}
pretrainer_output_dict = pretrainer(pretrainer_inputs_dict)
- mlm_model = tf.keras.Model(
+ mlm_model = tf_keras.Model(
inputs=pretrainer_inputs_dict, outputs=pretrainer_output_dict)
# Set `_auto_track_sub_layers` to False, so that the additional weights
# from `mlm` sub-object will not be included in the core model.
@@ -312,7 +314,7 @@ def create_preprocessing(*,
sp_model_file: Optional[str] = None,
do_lower_case: bool,
tokenize_with_offsets: bool,
- default_seq_length: int) -> tf.keras.Model:
+ default_seq_length: int) -> tf_keras.Model:
"""Returns a preprocessing Model for given tokenization parameters.
This function builds a Keras Model with attached subobjects suitable for
@@ -333,7 +335,7 @@ def create_preprocessing(*,
bert_pack_inputs subobject.
Returns:
- A tf.keras.Model object with several attached subobjects, suitable for
+ A tf_keras.Model object with several attached subobjects, suitable for
saving as a preprocessing SavedModel.
"""
# Select tokenizer.
@@ -353,7 +355,7 @@ def create_preprocessing(*,
# The root object of the preprocessing model can be called to do
# one-shot preprocessing for users with single-sentence inputs.
- sentences = tf.keras.layers.Input(shape=(), dtype=tf.string, name="sentences")
+ sentences = tf_keras.layers.Input(shape=(), dtype=tf.string, name="sentences")
if tokenize_with_offsets:
tokens, start_offsets, limit_offsets = tokenize(sentences)
else:
@@ -362,24 +364,24 @@ def create_preprocessing(*,
seq_length=default_seq_length,
special_tokens_dict=tokenize.get_special_tokens_dict())
model_inputs = pack(tokens)
- preprocessing = tf.keras.Model(sentences, model_inputs)
+ preprocessing = tf_keras.Model(sentences, model_inputs)
# Individual steps of preprocessing are made available as named subobjects
# to enable more general preprocessing. For saving, they need to be Models
# in their own right.
- preprocessing.tokenize = tf.keras.Model(sentences, tokens)
+ preprocessing.tokenize = tf_keras.Model(sentences, tokens)
# Provide an equivalent to tokenize.get_special_tokens_dict().
preprocessing.tokenize.get_special_tokens_dict = tf.train.Checkpoint()
preprocessing.tokenize.get_special_tokens_dict.__call__ = tf.function(
lambda: tokenize.get_special_tokens_dict(), # pylint: disable=[unnecessary-lambda]
input_signature=[])
if tokenize_with_offsets:
- preprocessing.tokenize_with_offsets = tf.keras.Model(
- sentences, [tokens, start_offsets, limit_offsets])
+ preprocessing.tokenize_with_offsets = tf_keras.Model(
+ sentences, [tokens, start_offsets, limit_offsets]) # pyrefly: ignore[unbound-name]
preprocessing.tokenize_with_offsets.get_special_tokens_dict = (
preprocessing.tokenize.get_special_tokens_dict)
# Conceptually, this should be
- # preprocessing.bert_pack_inputs = tf.keras.Model(tokens, model_inputs)
+ # preprocessing.bert_pack_inputs = tf_keras.Model(tokens, model_inputs)
# but technicalities require us to use a wrapper (see comments there).
# In particular, seq_length can be overridden when calling this.
preprocessing.bert_pack_inputs = BertPackInputsSavedModelWrapper(pack)
@@ -453,15 +455,15 @@ def _dont_assert(condition, data, summarize=None, name="Assert"):
@contextlib.contextmanager
def _maybe_disable_assert(disable_assert):
- """Scoped monkey patch of control_flow_ops.Assert to a no-op."""
+ """Scoped monkey patch of control_flow_assert.Assert to a no-op."""
if not disable_assert:
yield
return
- original_assert = control_flow_ops.Assert
- control_flow_ops.Assert = _dont_assert
+ original_assert = control_flow_assert.Assert
+ control_flow_assert.Assert = _dont_assert
yield
- control_flow_ops.Assert = original_assert
+ control_flow_assert.Assert = original_assert
def _check_no_assert(saved_model_path):
diff --git a/official/nlp/tools/export_tfhub_lib_test.py b/official/nlp/tools/export_tfhub_lib_test.py
index 51bb87319d7..4d733227b3f 100644
--- a/official/nlp/tools/export_tfhub_lib_test.py
+++ b/official/nlp/tools/export_tfhub_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow import estimator as tf_estimator
import tensorflow_hub as hub
import tensorflow_text as text
@@ -111,7 +111,7 @@ def _read_asset(asset: tf.saved_model.Asset):
def _find_lambda_layers(layer):
"""Returns list of all Lambda layers in a Keras model."""
- if isinstance(layer, tf.keras.layers.Lambda):
+ if isinstance(layer, tf_keras.layers.Lambda):
return [layer]
elif hasattr(layer, "layers"): # It's nested, like a Model.
result = []
@@ -233,9 +233,9 @@ def _dropout_mean_stddev(training, num_runs=20):
self.assertGreater(_dropout_mean_stddev(training=True), 1e-3)
# Test propagation of seq_length in shape inference.
- input_word_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_mask = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_type_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_word_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_mask = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_type_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
input_dict = dict(
input_word_ids=input_word_ids,
input_mask=input_mask,
@@ -469,9 +469,9 @@ def _dropout_mean_stddev_mlm(training, num_runs=20):
self.assertGreater(_dropout_mean_stddev_mlm(training=True), 1e-3)
# Test propagation of seq_length in shape inference.
- input_word_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_mask = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_type_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_word_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_mask = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_type_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
input_dict = dict(
input_word_ids=input_word_ids,
input_mask=input_mask,
@@ -1006,7 +1006,7 @@ def _get_special_tokens_dict(obj):
def input_fn():
self.assertFalse(tf.executing_eagerly())
# Build a preprocessing Model.
- sentences = tf.keras.layers.Input(shape=[], dtype=tf.string)
+ sentences = tf_keras.layers.Input(shape=[], dtype=tf.string)
preprocess = tf.saved_model.load(preprocess_export_path)
tokenize = hub.KerasLayer(preprocess.tokenize)
special_tokens_dict = _get_special_tokens_dict(tokenize.resolved_object)
@@ -1016,7 +1016,7 @@ def input_fn():
packed_inputs = layers.BertPackInputs(
4, special_tokens_dict=special_tokens_dict)(
tokens)
- preprocessing = tf.keras.Model(sentences, packed_inputs)
+ preprocessing = tf_keras.Model(sentences, packed_inputs)
# Map the dataset.
ds = tf.data.Dataset.from_tensors(
(tf.constant(["abc", "D EF"]), tf.constant([0, 1])))
diff --git a/official/nlp/tools/squad_evaluate_v1_1.py b/official/nlp/tools/squad_evaluate_v1_1.py
index 795fa471e3d..84e4b3b8117 100644
--- a/official/nlp/tools/squad_evaluate_v1_1.py
+++ b/official/nlp/tools/squad_evaluate_v1_1.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/tools/squad_evaluate_v2_0.py b/official/nlp/tools/squad_evaluate_v2_0.py
index ac02f72bec5..02799f6402c 100644
--- a/official/nlp/tools/squad_evaluate_v2_0.py
+++ b/official/nlp/tools/squad_evaluate_v2_0.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/tools/tf1_bert_checkpoint_converter_lib.py b/official/nlp/tools/tf1_bert_checkpoint_converter_lib.py
index b34bd00088f..a15747347d7 100644
--- a/official/nlp/tools/tf1_bert_checkpoint_converter_lib.py
+++ b/official/nlp/tools/tf1_bert_checkpoint_converter_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/nlp/tools/tf2_albert_encoder_checkpoint_converter.py b/official/nlp/tools/tf2_albert_encoder_checkpoint_converter.py
index 4583e4c4c65..7b774f0eee9 100644
--- a/official/nlp/tools/tf2_albert_encoder_checkpoint_converter.py
+++ b/official/nlp/tools/tf2_albert_encoder_checkpoint_converter.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,7 +22,7 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.albert import configs
from official.modeling import tf_utils
from official.nlp.modeling import models
@@ -94,7 +94,7 @@ def _create_albert_model(cfg):
attention_dropout_rate=cfg.attention_probs_dropout_prob,
max_sequence_length=cfg.max_position_embeddings,
type_vocab_size=cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=cfg.initializer_range))
return albert_encoder
@@ -112,7 +112,7 @@ def _create_pretrainer_model(cfg):
pretrainer = models.BertPretrainerV2(
encoder_network=albert_encoder,
mlm_activation=tf_utils.get_activation(cfg.hidden_act),
- mlm_initializer=tf.keras.initializers.TruncatedNormal(
+ mlm_initializer=tf_keras.initializers.TruncatedNormal(
stddev=cfg.initializer_range))
# Makes sure masked_lm layer's variables in pretrainer are created.
_ = pretrainer(pretrainer.inputs)
diff --git a/official/nlp/tools/tf2_bert_encoder_checkpoint_converter.py b/official/nlp/tools/tf2_bert_encoder_checkpoint_converter.py
index ddbff775faf..2987ce257b2 100644
--- a/official/nlp/tools/tf2_bert_encoder_checkpoint_converter.py
+++ b/official/nlp/tools/tf2_bert_encoder_checkpoint_converter.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -24,7 +24,7 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import configs
from official.modeling import tf_utils
from official.nlp.modeling import models
@@ -71,7 +71,7 @@ def _create_bert_model(cfg):
attention_dropout_rate=cfg.attention_probs_dropout_prob,
max_sequence_length=cfg.max_position_embeddings,
type_vocab_size=cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=cfg.initializer_range),
embedding_width=cfg.embedding_size)
@@ -91,7 +91,7 @@ def _create_bert_pretrainer_model(cfg):
pretrainer = models.BertPretrainerV2(
encoder_network=bert_encoder,
mlm_activation=tf_utils.get_activation(cfg.hidden_act),
- mlm_initializer=tf.keras.initializers.TruncatedNormal(
+ mlm_initializer=tf_keras.initializers.TruncatedNormal(
stddev=cfg.initializer_range))
# Makes sure the pretrainer variables are created.
_ = pretrainer(pretrainer.inputs)
diff --git a/official/nlp/tools/tokenization.py b/official/nlp/tools/tokenization.py
index 65d2b7717b1..2ee53dcfd37 100644
--- a/official/nlp/tools/tokenization.py
+++ b/official/nlp/tools/tokenization.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -24,7 +24,7 @@
import unicodedata
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import sentencepiece as spm
@@ -78,7 +78,7 @@ def validate_case_matches_checkpoint(do_lower_case, init_checkpoint):
"should pass in `--do_lower_case=%s` so that the fine-tuning matches "
"how the model was pre-training. If this error is wrong, please "
"just comment out this check." %
- (actual_flag, init_checkpoint, model_name, case_name, opposite_flag))
+ (actual_flag, init_checkpoint, model_name, case_name, opposite_flag)) # pyrefly: ignore[unbound-name]
def convert_to_unicode(text):
@@ -92,8 +92,8 @@ def convert_to_unicode(text):
raise ValueError("Unsupported string type: %s" % (type(text)))
elif six.PY2:
if isinstance(text, str):
- return text.decode("utf-8", "ignore")
- elif isinstance(text, unicode):
+ return text.decode("utf-8", "ignore") # pyrefly: ignore[missing-attribute]
+ elif isinstance(text, unicode): # pyrefly: ignore[unknown-name]
return text
else:
raise ValueError("Unsupported string type: %s" % (type(text)))
@@ -116,7 +116,7 @@ def printable_text(text):
elif six.PY2:
if isinstance(text, str):
return text
- elif isinstance(text, unicode):
+ elif isinstance(text, unicode): # pyrefly: ignore[unknown-name]
return text.encode("utf-8")
else:
raise ValueError("Unsupported string type: %s" % (type(text)))
@@ -428,16 +428,18 @@ def preprocess_text(inputs, remove_space=True, lower=False):
The preprocessed text.
"""
+ # Byte strings need to be explicitly decoded to unicode text,
+ # typically using UTF-8. A latin-1 fallback is included for
+ # backward compatibility with legacy sentence piece models.
+ if isinstance(inputs, six.binary_type):
+ try:
+ inputs = six.ensure_text(inputs, "utf-8")
+ except UnicodeDecodeError:
+ inputs = six.ensure_text(inputs, "latin-1")
outputs = inputs
if remove_space:
outputs = " ".join(inputs.strip().split())
- if six.PY2 and isinstance(outputs, str):
- try:
- outputs = six.ensure_text(outputs, "utf-8")
- except UnicodeDecodeError:
- outputs = six.ensure_text(outputs, "latin-1")
-
outputs = unicodedata.normalize("NFKD", outputs)
outputs = "".join([c for c in outputs if not unicodedata.combining(c)])
if lower:
diff --git a/official/nlp/tools/tokenization_test.py b/official/nlp/tools/tokenization_test.py
index c67a7e53d44..58e808784c6 100644
--- a/official/nlp/tools/tokenization_test.py
+++ b/official/nlp/tools/tokenization_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import tempfile
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.tools import tokenization
@@ -151,6 +151,18 @@ def test_is_punctuation(self):
self.assertFalse(tokenization._is_punctuation(u"A"))
self.assertFalse(tokenization._is_punctuation(u" "))
+ def test_preprocess_text(self):
+ self.assertEqual(tokenization.preprocess_text("hello world"), "hello world")
+ self.assertEqual(tokenization.preprocess_text(b"hello \xc3\xa9"), "hello e")
+ self.assertEqual(tokenization.preprocess_text(b"hello \xe9"), "hello e")
+ self.assertEqual(
+ tokenization.preprocess_text(b"hello world", remove_space=True),
+ "hello world",
+ )
+ self.assertEqual(
+ tokenization.preprocess_text("Hello World", lower=True), "hello world"
+ )
+
if __name__ == "__main__":
tf.test.main()
diff --git a/official/nlp/train.py b/official/nlp/train.py
index feef3d54ea5..8c32ef42632 100644
--- a/official/nlp/train.py
+++ b/official/nlp/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,9 @@
from absl import app
from absl import flags
+from absl import logging
import gin
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
# pylint: disable=unused-import
@@ -36,6 +38,50 @@
default=None,
help='The number of total training steps for the pretraining job.')
+flags.DEFINE_bool(
+ 'enable_async_checkpointing',
+ default=True,
+ help='A boolean indicating whether to enable async checkpoint saving')
+
+
+def _run_experiment_with_preemption_recovery(params, model_dir):
+ """Runs experiment and tries to reconnect when encounting a preemption."""
+ keep_training = True
+ while keep_training:
+ preemption_watcher = None
+ try:
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu,
+ **params.runtime.model_parallelism())
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+ # pylint: disable=line-too-long
+ preemption_watcher = None # copybara-replace
+ # pylint: enable=line-too-long
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir,
+ enable_async_checkpointing=FLAGS.enable_async_checkpointing)
+
+ keep_training = False
+ except tf.errors.OpError as e:
+ if preemption_watcher and preemption_watcher.preemption_message:
+ preemption_watcher.block_until_worker_exit()
+ logging.info(
+ 'Some TPU workers had been preempted (message: %s), '
+ 'retarting training from the last checkpoint...',
+ preemption_watcher.preemption_message)
+ keep_training = True
+ else:
+ raise e from None
+
def main(_):
gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
@@ -58,21 +104,7 @@ def main(_):
if params.runtime.mixed_precision_dtype:
performance.set_mixed_precision_policy(
params.runtime.mixed_precision_dtype)
- distribution_strategy = distribute_utils.get_distribution_strategy(
- distribution_strategy=params.runtime.distribution_strategy,
- all_reduce_alg=params.runtime.all_reduce_alg,
- num_gpus=params.runtime.num_gpus,
- tpu_address=params.runtime.tpu,
- **params.runtime.model_parallelism())
- with distribution_strategy.scope():
- task = task_factory.get_task(params.task, logging_dir=model_dir)
-
- train_lib.run_experiment(
- distribution_strategy=distribution_strategy,
- task=task,
- mode=FLAGS.mode,
- params=params,
- model_dir=model_dir)
+ _run_experiment_with_preemption_recovery(params, model_dir)
train_utils.save_gin_config(FLAGS.mode, model_dir)
diff --git a/official/pip_package/setup.py b/official/pip_package/setup.py
index 8399275ab73..2b97707514c 100644
--- a/official/pip_package/setup.py
+++ b/official/pip_package/setup.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,8 +20,8 @@
from setuptools import find_packages
from setuptools import setup
-version = '2.8.0'
-tf_version = '2.8.0' # Major version.
+version = '2.20.0'
+tf_version = '2.20.0' # Major version.
project_name = 'tf-models-official'
@@ -38,12 +38,16 @@
sys.argv.pop(project_name_idx)
-def _get_requirements():
+def _get_requirements(is_nightly=False):
"""Parses requirements.txt file."""
install_requires_tmp = []
dependency_links_tmp = []
+ if is_nightly:
+ file_name = '../nightly_requirements.txt'
+ else:
+ file_name = '../requirements.txt'
with open(
- os.path.join(os.path.dirname(__file__), '../requirements.txt'), 'r') as f:
+ os.path.join(os.path.dirname(__file__), file_name), 'r') as f:
for line in f:
package_name = line.strip()
# Skip empty line or comments starting with "#".
@@ -55,13 +59,16 @@ def _get_requirements():
install_requires_tmp.append(package_name)
return install_requires_tmp, dependency_links_tmp
-install_requires, dependency_links = _get_requirements()
-
if project_name == 'tf-models-nightly':
+ install_requires, dependency_links = _get_requirements(is_nightly=True)
+ version_split = version.split('.')
+ version_split[1] = str(int(version_split[1]) + 1)
+ version = '.'.join(version_split)
version += '.dev' + datetime.datetime.now().strftime('%Y%m%d')
install_requires.append('tf-nightly')
install_requires.append('tensorflow-text-nightly')
else:
+ install_requires, dependency_links = _get_requirements()
install_requires.append(f'tensorflow~={tf_version}')
install_requires.append(f'tensorflow-text~={tf_version}')
diff --git a/official/projects/README.md b/official/projects/README.md
index 9c94fdd1106..69b76dd4332 100644
--- a/official/projects/README.md
+++ b/official/projects/README.md
@@ -1,11 +1,43 @@
# TensorFlow Model Garden Modeling Projects
-This directory contains projects using TensorFlow Model Garden Modeling
-libraries.
+This directory contains projects using Modeling libraries of TensorFlow Model
+Garden. More details about each project can be found in the individual
+project folders listed below.
-## Projects
+⚠️ Disclaimer: Checkpoints included in the project folders listed below are
+based on training with publicly available datasets. Some datasets contain
+limitations, including non-commercial use limitations. Please review the terms
+and conditions made available by third parties before using the datasets
+provided. Checkpoints are licensed under
+[Apache 2.0](https://github.com/tensorflow/models/blob/master/LICENSE).
+
+⚠️ Disclaimer: Datasets hyperlinked from the project folders listed below are
+not owned or distributed by Google. Such datasets are made available by third
+parties. Please review the terms and conditions made available by the third
+parties before using the data.
-* [NHNet](nhnet):
- [Generating Representative Headlines for News Stories](https://arxiv.org/abs/2001.09386)
- by Gu et al, 2020
+## Projects
+* [AssembleNet](./assemblenet/README.md)
+* [BASNet](./basnet/README.md)
+* [BigBird](./bigbird/README.md)
+* [DeepMAC Mask-RCNN](./deepmac_maskrcnn/README.md)
+* [DETR](./detr/README.md)
+* [Edge-TPU for Vision and NLP](./edgetpu/README.md)
+* [Language-agnostic BERT Sentence Embedding](./labse/README.md)
+* [Long-Document Transformer](./longformer/README.md)
+* [MobileBERT](./mobilebert/README.md)
+* [MoViNets](./movinet/README.md)
+* [News Headline Generation Model: NHNet](./nhnet/README.md)
+* [Training with Pruning](./pruning/README.md)
+* [QAT for Computer Vision](./qat/vision/README.md)
+* [Roformer Project](./roformer/README.md)
+* [Training ELECTRA Augmented with Multi-word Selection](./teams/README.md)
+* [NLP example project](./text_classification_example/README.md)
+* [TensorNetwork BERT](./tn_bert/README.md)
+* [Token Dropping for Efficient BERT Pretraining](./token_dropping/README.md)
+* [Spatiotemporal Contrastive Video Representation Learning](./video_ssl/README.md)
+* [Vision Transformer (ViT)](./vit/README.md)
+* [Data-Efficient Image Transformer (DEIT)](./vit/README.md)
+* [Volumetric Models](./volumetric_models/README.md)
+* [YouTube-8M Tensorflow Starter Code](./yt8m/README.md)
diff --git a/official/projects/__init__.py b/official/projects/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/__init__.py
+++ b/official/projects/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/assemblenet/README.md b/official/projects/assemblenet/README.md
index 74a97a536db..014830fe2b9 100644
--- a/official/projects/assemblenet/README.md
+++ b/official/projects/assemblenet/README.md
@@ -117,10 +117,10 @@ provided
Example of training AssembleNet with UCF101 TF Datasets.
```bash
-python -m official.vision.beta.projects.assemblenet.trian \
+python -m official.projects.assemblenet.trian \
--mode=train_and_eval --experiment=assemblenet_ucf101 \
--model_dir='YOUR_GS_BUCKET_TO_SAVE_MODEL' \
---config_file=./official/vision/beta/projects/assemblenet/\
+--config_file=./official/projects/assemblenet/\
--ucf101_assemblenet_tpu.yaml \
--tpu=TPU_NAME
```
@@ -128,10 +128,10 @@ python -m official.vision.beta.projects.assemblenet.trian \
Example of training AssembleNet++ with UCF101 TF Datasets.
```bash
-python -m official.vision.beta.projects.assemblenet.trian \
+python -m official.projects.assemblenet.trian \
--mode=train_and_eval --experiment=assemblenetplus_ucf101 \
--model_dir='YOUR_GS_BUCKET_TO_SAVE_MODEL' \
---config_file=./official/vision/beta/projects/assemblenet/\
+--config_file=./official/projects/assemblenet/\
--ucf101_assemblenet_plus_tpu.yaml \
--tpu=TPU_NAME
```
diff --git a/official/projects/assemblenet/configs/assemblenet.py b/official/projects/assemblenet/configs/assemblenet.py
index 08301dc276d..60ced7f6d6e 100644
--- a/official/projects/assemblenet/configs/assemblenet.py
+++ b/official/projects/assemblenet/configs/assemblenet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -61,7 +61,7 @@ def flat_lists_to_blocks(model_structures, model_edge_weights):
if node[0] < 0:
block = BlockSpec(level=node[0], temporal_dilation=node[1])
else:
- block = BlockSpec(
+ block = BlockSpec( # pytype: disable=wrong-arg-types
level=node[0],
input_blocks=node[1],
num_filters=node[2],
@@ -195,17 +195,24 @@ class Backbone3D(backbones_3d.Backbone3D):
assemblenet_plus : AssembleNetPlus backbone config.
"""
type: Optional[str] = None
- assemblenet: AssembleNet = AssembleNet()
- assemblenet_plus: AssembleNetPlus = AssembleNetPlus()
+ assemblenet: AssembleNet = dataclasses.field(default_factory=AssembleNet)
+ assemblenet_plus: AssembleNetPlus = dataclasses.field(
+ default_factory=AssembleNetPlus
+ )
@dataclasses.dataclass
class AssembleNetModel(video_classification.VideoClassificationModel):
"""The AssembleNet model config."""
model_type: str = 'assemblenet'
- backbone: Backbone3D = Backbone3D(type='assemblenet')
- norm_activation: common.NormActivation = common.NormActivation(
- norm_momentum=0.99, norm_epsilon=1e-5, use_sync_bn=True)
+ backbone: Backbone3D = dataclasses.field(
+ default_factory=lambda: Backbone3D(type='assemblenet')
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ norm_momentum=0.99, norm_epsilon=1e-5, use_sync_bn=True
+ )
+ )
max_pool_predictions: bool = False
@@ -213,9 +220,14 @@ class AssembleNetModel(video_classification.VideoClassificationModel):
class AssembleNetPlusModel(video_classification.VideoClassificationModel):
"""The AssembleNet model config."""
model_type: str = 'assemblenet_plus'
- backbone: Backbone3D = Backbone3D(type='assemblenet_plus')
- norm_activation: common.NormActivation = common.NormActivation(
- norm_momentum=0.99, norm_epsilon=1e-5, use_sync_bn=True)
+ backbone: Backbone3D = dataclasses.field(
+ default_factory=lambda: Backbone3D(type='assemblenet_plus')
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ norm_momentum=0.99, norm_epsilon=1e-5, use_sync_bn=True
+ )
+ )
max_pool_predictions: bool = False
diff --git a/official/projects/assemblenet/configs/assemblenet_test.py b/official/projects/assemblenet/configs/assemblenet_test.py
index f11c21135c0..cada250f471 100644
--- a/official/projects/assemblenet/configs/assemblenet_test.py
+++ b/official/projects/assemblenet/configs/assemblenet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.projects.assemblenet.configs import assemblenet
diff --git a/official/projects/assemblenet/modeling/assemblenet.py b/official/projects/assemblenet/modeling/assemblenet.py
index 3c2417a94e8..c435ef9cb61 100644
--- a/official/projects/assemblenet/modeling/assemblenet.py
+++ b/official/projects/assemblenet/modeling/assemblenet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -51,7 +51,7 @@
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.assemblenet.configs import assemblenet as cfg
@@ -59,7 +59,7 @@
from official.vision.modeling import factory_3d as model_factory
from official.vision.modeling.backbones import factory as backbone_factory
-layers = tf.keras.layers
+layers = tf_keras.layers
intermediate_channel_size = [64, 128, 256, 512]
@@ -76,7 +76,7 @@ def fixed_padding(inputs, kernel_size):
A padded `Tensor` of the same `data_format` with size either intact
(if `kernel_size == 1`) or padded (if `kernel_size > 1`).
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
pad_total = kernel_size - 1
pad_beg = pad_total // 2
pad_end = pad_total - pad_beg
@@ -118,7 +118,7 @@ def reshape_temporal_conv1d_bn(inputs: tf.Tensor,
A padded `Tensor` of the same `data_format` with size either intact
(if `kernel_size == 1`) or padded (if `kernel_size > 1`).
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
feature_shape = inputs.shape
@@ -128,23 +128,23 @@ def reshape_temporal_conv1d_bn(inputs: tf.Tensor,
[-1, num_frames, feature_shape[1] * feature_shape[2], feature_shape[3]])
if temporal_dilation == 1:
- inputs = tf.keras.layers.Conv2D(
+ inputs = tf_keras.layers.Conv2D(
filters=filters,
kernel_size=(kernel_size, 1),
strides=1,
padding='SAME',
use_bias=False,
- kernel_initializer=tf.keras.initializers.VarianceScaling())(
+ kernel_initializer=tf_keras.initializers.VarianceScaling())(
inputs=inputs)
else:
- inputs = tf.keras.layers.Conv2D(
+ inputs = tf_keras.layers.Conv2D(
filters=filters,
kernel_size=(kernel_size, 1),
strides=1,
padding='SAME',
dilation_rate=(temporal_dilation, 1),
use_bias=False,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=math.sqrt(2.0 / (kernel_size * feature_shape[3]))))(
inputs=inputs)
@@ -164,7 +164,7 @@ def conv2d_fixed_padding(inputs: tf.Tensor, filters: int, kernel_size: int,
"""Strided 2-D convolution with explicit padding.
The padding is consistent and is based only on `kernel_size`, not on the
- dimensions of `inputs` (as opposed to using `tf.keras.layers.Conv2D` alone).
+ dimensions of `inputs` (as opposed to using `tf_keras.layers.Conv2D` alone).
Args:
inputs: `Tensor` of size `[batch, channels, height_in, width_in]`.
@@ -178,13 +178,13 @@ def conv2d_fixed_padding(inputs: tf.Tensor, filters: int, kernel_size: int,
if strides > 1:
inputs = fixed_padding(inputs, kernel_size)
- return tf.keras.layers.Conv2D(
+ return tf_keras.layers.Conv2D(
filters=filters,
kernel_size=kernel_size,
strides=strides,
padding=('SAME' if strides == 1 else 'VALID'),
use_bias=False,
- kernel_initializer=tf.keras.initializers.VarianceScaling())(
+ kernel_initializer=tf_keras.initializers.VarianceScaling())(
inputs=inputs)
@@ -211,18 +211,18 @@ def conv3d_same_padding(inputs: tf.Tensor,
"""
if isinstance(kernel_size, int):
if do_2d_conv:
- kernel_size = [1, kernel_size, kernel_size]
+ kernel_size = [1, kernel_size, kernel_size] # pyrefly: ignore[bad-assignment]
else:
- kernel_size = [kernel_size, kernel_size, kernel_size]
+ kernel_size = [kernel_size, kernel_size, kernel_size] # pyrefly: ignore[bad-assignment]
- return tf.keras.layers.Conv3D(
+ return tf_keras.layers.Conv3D(
filters=filters,
kernel_size=kernel_size,
strides=[1, strides, strides],
padding='SAME',
dilation_rate=[temporal_dilation, 1, 1],
use_bias=False,
- kernel_initializer=tf.keras.initializers.VarianceScaling())(
+ kernel_initializer=tf_keras.initializers.VarianceScaling())(
inputs=inputs)
@@ -374,7 +374,7 @@ def spatial_resize_and_concat(inputs):
Returns:
The output `Tensor` after concatenation.
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
# Do nothing if only 1 input
@@ -393,7 +393,7 @@ def spatial_resize_and_concat(inputs):
for i in range(len(inputs)):
if inputs[i].shape[1] != sm_size[0] or inputs[i].shape[2] != sm_size[1]:
ratio = (inputs[i].shape[1] + 1) // sm_size[0]
- inputs[i] = tf.keras.layers.MaxPool2D([ratio, ratio],
+ inputs[i] = tf_keras.layers.MaxPool2D([ratio, ratio],
ratio,
padding='same')(
inputs[i])
@@ -431,7 +431,7 @@ def __init__(self,
self._index = index
self._use_5d_mode = use_5d_mode
self._model_edge_weights = model_edge_weights
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
def get_config(self):
@@ -452,7 +452,7 @@ def build(self, input_shape: tf.TensorShape):
if self._index is None or not self._model_edge_weights:
self._edge_weights = self.add_weight(
shape=self._weights_shape,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
mean=0.0, stddev=0.01),
trainable=True,
name='agg_weights')
@@ -499,11 +499,11 @@ def call(self,
assert sm_size[0] != 0
ratio = (inp.shape[h_channel_loc] + 1) // sm_size[0]
if use_5d_mode:
- inp = tf.keras.layers.MaxPool3D([1, ratio, ratio], [1, ratio, ratio],
+ inp = tf_keras.layers.MaxPool3D([1, ratio, ratio], [1, ratio, ratio],
padding='same')(
inp)
else:
- inp = tf.keras.layers.MaxPool2D([ratio, ratio], ratio,
+ inp = tf_keras.layers.MaxPool2D([ratio, ratio], ratio,
padding='same')(
inp)
@@ -610,7 +610,7 @@ def rgb_conv_stem(inputs,
Returns:
The output `Tensor`.
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
if temporal_dilation < 1:
@@ -634,7 +634,7 @@ def rgb_conv_stem(inputs,
bn_epsilon=bn_epsilon,
use_sync_bn=use_sync_bn)
- inputs = tf.keras.layers.MaxPool2D(
+ inputs = tf_keras.layers.MaxPool2D(
pool_size=3, strides=2, padding='SAME')(
inputs=inputs)
inputs = tf.identity(inputs, 'initial_max_pool')
@@ -673,7 +673,7 @@ def flow_conv_stem(inputs,
inputs)
inputs = tf.nn.relu(inputs)
- inputs = tf.keras.layers.MaxPool2D(
+ inputs = tf_keras.layers.MaxPool2D(
pool_size=2, strides=2, padding='SAME')(
inputs=inputs)
inputs = tf.identity(inputs, 'initial_max_pool')
@@ -704,7 +704,7 @@ def multi_stream_heads(streams,
def _pool_and_reshape(net):
# The activation is 7x7 so this is a global average pool.
- net = tf.keras.layers.GlobalAveragePooling2D()(inputs=net)
+ net = tf_keras.layers.GlobalAveragePooling2D()(inputs=net)
net = tf.identity(net, 'final_avg_pool0')
net = tf.reshape(net, [-1, num_frames, num_channels])
@@ -724,7 +724,7 @@ def _pool_and_reshape(net):
if len(final_nodes) > 1:
outputs = outputs / len(final_nodes)
- outputs = tf.keras.layers.Dense(
+ outputs = tf_keras.layers.Dense(
units=num_classes,
kernel_initializer=tf.random_normal_initializer(stddev=.01))(
inputs=outputs)
@@ -739,7 +739,7 @@ def _pool_and_reshape(net):
return outputs
-class AssembleNet(tf.keras.Model):
+class AssembleNet(tf_keras.Model):
"""AssembleNet backbone."""
def __init__(
@@ -766,7 +766,7 @@ def __init__(
inputs of the same resolution.
num_frames: the number of frames in the input tensor.
model_structure: AssembleNet model structure in the string format.
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
Dimension should be `[batch*time, height, width, channels]`.
model_edge_weights: AssembleNet model structure connection weights in the
string format.
@@ -776,8 +776,8 @@ def __init__(
combine_method: 'str' for the weighted summation to fuse different blocks.
**kwargs: pass through arguments.
"""
- inputs = tf.keras.Input(shape=input_specs.shape[1:])
- data_format = tf.keras.backend.image_data_format()
+ inputs = tf_keras.Input(shape=input_specs.shape[1:])
+ data_format = tf_keras.backend.image_data_format()
# Creation of the model graph.
logging.info('model_structure=%r', model_structure)
@@ -836,7 +836,7 @@ def __init__(
streams.append(inputs)
elif structure[i][0] == -2:
inputs = flow_conv_stem(
- flow_inputs,
+ flow_inputs, # pyrefly: ignore[unbound-name]
stem_filters,
temporal_dilation=structure[i][1],
bn_decay=bn_decay,
@@ -883,7 +883,7 @@ def __init__(
inputs=original_inputs, outputs=streams, **kwargs)
-class AssembleNetModel(tf.keras.Model):
+class AssembleNetModel(tf_keras.Model):
"""An AssembleNet model builder."""
def __init__(self,
@@ -892,7 +892,7 @@ def __init__(self,
num_frames: int,
model_structure: List[Any],
input_specs: Optional[Mapping[str,
- tf.keras.layers.InputSpec]] = None,
+ tf_keras.layers.InputSpec]] = None,
max_pool_predictions: bool = False,
**kwargs):
if not input_specs:
@@ -914,7 +914,7 @@ def __init__(self,
grouping[model_structure[i][0]].append(i)
inputs = {
- k: tf.keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
+ k: tf_keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
}
streams = self._backbone(inputs['image'])
@@ -985,7 +985,7 @@ def assemblenet_v1(assemblenet_depth: int,
**kwargs):
"""Returns the AssembleNet model for a given size and number of output classes."""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
if assemblenet_depth not in ASSEMBLENET_SPECS:
@@ -995,7 +995,7 @@ def assemblenet_v1(assemblenet_depth: int,
params = ASSEMBLENET_SPECS[assemblenet_depth]
backbone = AssembleNet(
block_fn=params['block'],
- num_blocks=params['num_blocks'],
+ num_blocks=params['num_blocks'], # pyrefly: ignore[bad-argument-type]
num_frames=num_frames,
model_structure=model_structure,
input_specs=input_specs,
@@ -1014,11 +1014,11 @@ def assemblenet_v1(assemblenet_depth: int,
@backbone_factory.register_backbone_builder('assemblenet')
def build_assemblenet_v1(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
norm_activation_config: hyperparams.Config,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None
-) -> tf.keras.Model:
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> tf_keras.Model:
"""Builds assemblenet backbone."""
del l2_regularizer
@@ -1033,13 +1033,13 @@ def build_assemblenet_v1(
backbone_cfg.blocks)
params = ASSEMBLENET_SPECS[assemblenet_depth]
block_fn = functools.partial(
- params['block'],
+ params['block'], # pyrefly: ignore[bad-argument-type, not-callable]
use_sync_bn=norm_activation_config.use_sync_bn,
bn_decay=norm_activation_config.norm_momentum,
bn_epsilon=norm_activation_config.norm_epsilon)
backbone = AssembleNet(
block_fn=block_fn,
- num_blocks=params['num_blocks'],
+ num_blocks=params['num_blocks'], # pyrefly: ignore[bad-argument-type]
num_frames=backbone_cfg.num_frames,
model_structure=model_structure,
input_specs=input_specs,
@@ -1055,10 +1055,10 @@ def build_assemblenet_v1(
@model_factory.register_model_builder('assemblenet')
def build_assemblenet_model(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
model_config: cfg.AssembleNetModel,
num_classes: int,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None):
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None):
"""Builds assemblenet model."""
input_specs_dict = {'image': input_specs}
backbone = build_assemblenet_v1(input_specs, model_config.backbone,
diff --git a/official/projects/assemblenet/modeling/assemblenet_plus.py b/official/projects/assemblenet/modeling/assemblenet_plus.py
index c07657bdf1f..99df7aa1db5 100644
--- a/official/projects/assemblenet/modeling/assemblenet_plus.py
+++ b/official/projects/assemblenet/modeling/assemblenet_plus.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -58,7 +58,7 @@
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.assemblenet.configs import assemblenet as cfg
@@ -67,7 +67,7 @@
from official.vision.modeling import factory_3d as model_factory
from official.vision.modeling.backbones import factory as backbone_factory
-layers = tf.keras.layers
+layers = tf_keras.layers
def softmax_merge_peer_attentions(peers):
@@ -80,11 +80,11 @@ def softmax_merge_peer_attentions(peers):
Returns:
The output `Tensor` of size `[batch*time, channels].
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
dtype = peers[0].dtype
assert data_format == 'channels_last'
- initial_attn_weights = tf.keras.initializers.TruncatedNormal(stddev=0.01)(
+ initial_attn_weights = tf_keras.initializers.TruncatedNormal(stddev=0.01)(
[len(peers)])
attn_weights = tf.cast(tf.nn.softmax(initial_attn_weights), dtype)
weighted_peers = []
@@ -116,7 +116,7 @@ def apply_attention(inputs,
Returns:
The output `Tensor` after concatenation.
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
if use_5d_mode:
@@ -128,7 +128,7 @@ def apply_attention(inputs,
attn = softmax_merge_peer_attentions(attention_in)
else:
attn = tf.math.reduce_mean(inputs, [h_channel_loc, h_channel_loc + 1])
- attn = tf.keras.layers.Dense(
+ attn = tf_keras.layers.Dense(
units=inputs.shape[-1],
kernel_initializer=tf.random_normal_initializer(stddev=.01))(
inputs=attn)
@@ -180,7 +180,7 @@ def __init__(self,
self._use_5d_mode = use_5d_mode
self._model_edge_weights = model_edge_weights
self._num_object_classes = num_object_classes
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
def get_config(self):
@@ -202,7 +202,7 @@ def build(self, input_shape: tf.TensorShape):
if self._index is None or not self._model_edge_weights:
self._edge_weights = self.add_weight(
shape=self._weights_shape,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
mean=0.0, stddev=0.01),
trainable=True,
name='agg_weights')
@@ -253,11 +253,11 @@ def call(self,
assert sm_size[0] != 0
ratio = (inp.shape[h_channel_loc] + 1) // sm_size[0]
if use_5d_mode:
- inp = tf.keras.layers.MaxPool3D([1, ratio, ratio], [1, ratio, ratio],
+ inp = tf_keras.layers.MaxPool3D([1, ratio, ratio], [1, ratio, ratio],
padding='same')(
inp)
else:
- inp = tf.keras.layers.MaxPool2D([ratio, ratio], ratio,
+ inp = tf_keras.layers.MaxPool2D([ratio, ratio], ratio,
padding='same')(
inp)
@@ -375,7 +375,7 @@ def object_conv_stem(inputs):
Returns:
The output `Tensor`.
"""
- inputs = tf.keras.layers.MaxPool2D(
+ inputs = tf_keras.layers.MaxPool2D(
pool_size=4, strides=4, padding='SAME')(
inputs=inputs)
inputs = tf.identity(inputs, 'initial_max_pool')
@@ -383,7 +383,7 @@ def object_conv_stem(inputs):
return inputs
-class AssembleNetPlus(tf.keras.Model):
+class AssembleNetPlus(tf_keras.Model):
"""AssembleNet++ backbone."""
def __init__(self,
@@ -410,7 +410,7 @@ def __init__(self,
inputs of the same resolution.
num_frames: the number of frames in the input tensor.
model_structure: AssembleNetPlus model structure in the string format.
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
Dimension should be `[batch*time, height, width, channels]`.
model_edge_weights: AssembleNet model structure connection weight in the
string format.
@@ -425,7 +425,7 @@ def __init__(self,
Model `function` that takes in `inputs` and `is_training` and returns the
output `Tensor` of the AssembleNetPlus model.
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
# Creation of the model graph.
logging.info('model_structure=%r', model_structure)
@@ -434,11 +434,11 @@ def __init__(self,
structure = model_structure
if use_object_input:
- original_inputs = tf.keras.Input(shape=input_specs[0].shape[1:])
- object_inputs = tf.keras.Input(shape=input_specs[1].shape[1:])
+ original_inputs = tf_keras.Input(shape=input_specs[0].shape[1:])
+ object_inputs = tf_keras.Input(shape=input_specs[1].shape[1:])
input_specs = input_specs[0]
else:
- original_inputs = tf.keras.Input(shape=input_specs.shape[1:])
+ original_inputs = tf_keras.Input(shape=input_specs.shape[1:])
object_inputs = None
original_num_frames = num_frames
@@ -493,7 +493,7 @@ def __init__(self,
streams.append(inputs)
elif structure[i][0] == -2:
inputs = asn.flow_conv_stem(
- flow_inputs,
+ flow_inputs, # pyrefly: ignore[unbound-name]
stem_filters,
temporal_dilation=structure[i][1],
bn_decay=bn_decay,
@@ -529,7 +529,7 @@ def __init__(self,
for node_index in nodes_below:
attn = tf.reduce_mean(streams[node_index], [1, 2])
- attn = tf.keras.layers.Dense(
+ attn = tf_keras.layers.Dense(
units=lg_channel,
kernel_initializer=tf.random_normal_initializer(stddev=.01))(
inputs=attn)
@@ -564,8 +564,8 @@ def __init__(self,
inputs=inputs, outputs=streams, **kwargs)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class AssembleNetPlusModel(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class AssembleNetPlusModel(tf_keras.Model):
"""An AssembleNet++ model builder."""
def __init__(self,
@@ -574,7 +574,7 @@ def __init__(self,
num_frames: int,
model_structure: List[Any],
input_specs: Optional[Dict[str,
- tf.keras.layers.InputSpec]] = None,
+ tf_keras.layers.InputSpec]] = None,
max_pool_predictions: bool = False,
use_object_input: bool = False,
**kwargs):
@@ -602,7 +602,7 @@ def __init__(self,
grouping[model_structure[i][0]].append(i)
inputs = {
- k: tf.keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
+ k: tf_keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
}
if use_object_input:
@@ -650,7 +650,7 @@ def assemblenet_plus(assemblenet_depth: int,
**kwargs):
"""Returns the AssembleNet++ model for a given size and number of output classes."""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
if assemblenet_depth not in asn.ASSEMBLENET_SPECS:
@@ -671,7 +671,7 @@ def assemblenet_plus(assemblenet_depth: int,
input_specs=input_specs,
model_edge_weights=model_edge_weights,
use_object_input=use_object_input,
- attention_mode=attention_mode,
+ attention_mode=attention_mode, # pyrefly: ignore[bad-argument-type]
**kwargs)
return AssembleNetPlusModel(
backbone,
@@ -686,11 +686,11 @@ def assemblenet_plus(assemblenet_depth: int,
@backbone_factory.register_backbone_builder('assemblenet_plus')
def build_assemblenet_plus(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
norm_activation_config: hyperparams.Config,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None
-) -> tf.keras.Model:
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> tf_keras.Model:
"""Builds assemblenet++ backbone."""
del l2_regularizer
@@ -728,10 +728,10 @@ def build_assemblenet_plus(
@model_factory.register_model_builder('assemblenet_plus')
def build_assemblenet_plus_model(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
model_config: cfg.AssembleNetPlusModel,
num_classes: int,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None):
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None):
"""Builds assemblenet++ model."""
input_specs_dict = {'image': input_specs}
backbone = build_assemblenet_plus(input_specs, model_config.backbone,
diff --git a/official/projects/assemblenet/modeling/assemblenet_plus_test.py b/official/projects/assemblenet/modeling/assemblenet_plus_test.py
index a2799c0b045..f1775acd039 100644
--- a/official/projects/assemblenet/modeling/assemblenet_plus_test.py
+++ b/official/projects/assemblenet/modeling/assemblenet_plus_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.assemblenet.configs import assemblenet as asn_config
from official.projects.assemblenet.modeling import assemblenet_plus as asnp
@@ -39,8 +39,8 @@ def test_network_creation(self, depth, use_object_input, attention_mode):
vid_input = (batch_size * num_frames, img_size, img_size, 3)
obj_input = (batch_size * num_frames, img_size, img_size,
num_object_classes)
- input_specs = (tf.keras.layers.InputSpec(shape=(vid_input)),
- tf.keras.layers.InputSpec(shape=(obj_input)))
+ input_specs = (tf_keras.layers.InputSpec(shape=(vid_input)),
+ tf_keras.layers.InputSpec(shape=(obj_input)))
vid_inputs = np.random.rand(batch_size * num_frames, img_size, img_size,
3)
obj_inputs = np.random.rand(batch_size * num_frames, img_size, img_size,
@@ -52,7 +52,7 @@ def test_network_creation(self, depth, use_object_input, attention_mode):
edge_weights = asn_config.full_asnp_structure_weights
else:
# video input: (batch_size, FLAGS.num_frames, image_size, image_size, 3)
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=(batch_size, num_frames, img_size, img_size, 3))
inputs = np.random.rand(batch_size, num_frames, img_size, img_size, 3)
diff --git a/official/projects/assemblenet/modeling/rep_flow_2d_layer.py b/official/projects/assemblenet/modeling/rep_flow_2d_layer.py
index 2b6439342ed..75108cbb242 100644
--- a/official/projects/assemblenet/modeling/rep_flow_2d_layer.py
+++ b/official/projects/assemblenet/modeling/rep_flow_2d_layer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -24,9 +24,9 @@
"""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-layers = tf.keras.layers
+layers = tf_keras.layers
BATCH_NORM_DECAY = 0.99
BATCH_NORM_EPSILON = 1e-5
@@ -53,7 +53,7 @@ def build_batch_norm(init_zero: bool = False,
else:
gamma_initializer = tf.ones_initializer()
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
assert data_format == 'channels_last'
if data_format == 'channels_first':
@@ -91,7 +91,7 @@ def divergence(p1, p2, f_grad_x, f_grad_y, name):
Returns:
A `Tensor` with the same `data_format` and shape as input.
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
df = 'NHWC' if data_format == 'channels_last' else 'NCHW'
with tf.name_scope('divergence_' + name):
@@ -108,7 +108,7 @@ def divergence(p1, p2, f_grad_x, f_grad_y, name):
def forward_grad(x, f_grad_x, f_grad_y, name):
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
with tf.name_scope('forward_grad_' + name):
df = 'NHWC' if data_format == 'channels_last' else 'NCHW'
grad_x = tf.nn.conv2d(x, f_grad_x, [1, 1, 1, 1], 'SAME', data_format=df)
@@ -257,7 +257,7 @@ def build(self, input_shape: tf.TensorShape):
strides=1,
padding='same',
use_bias=False,
- kernel_initializer=tf.keras.initializers.VarianceScaling(),
+ kernel_initializer=tf_keras.initializers.VarianceScaling(),
name='rf/bottleneck1')
self._bottleneck_conv2 = layers.Conv2D(
filters=self._depth,
@@ -265,11 +265,11 @@ def build(self, input_shape: tf.TensorShape):
strides=1,
padding='same',
use_bias=False,
- kernel_initializer=tf.keras.initializers.VarianceScaling(),
+ kernel_initializer=tf_keras.initializers.VarianceScaling(),
name='rf/bottleneck2')
self._batch_norm = build_batch_norm(init_zero=True)
- def call(self, inputs: tf.Tensor, training: bool = None) -> tf.Tensor:
+ def call(self, inputs: tf.Tensor, training: bool = None) -> tf.Tensor: # pytype: disable=annotation-type-mismatch
"""Perform representation flows.
Args:
@@ -280,7 +280,7 @@ def call(self, inputs: tf.Tensor, training: bool = None) -> tf.Tensor:
Returns:
A tensor of the same shape as the inputs.
"""
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
df = 'NHWC' if data_format == 'channels_last' else 'NCHW'
axis = 3 if data_format == 'channels_last' else 1 # channel axis
dtype = inputs.dtype
@@ -397,7 +397,7 @@ def call(self, inputs: tf.Tensor, training: bool = None) -> tf.Tensor:
flow = tf.ensure_shape(flow, output_shape)
return flow
else:
- flow = self._bottleneck_conv2(flow)
+ flow = self._bottleneck_conv2(flow) # pyrefly: ignore[not-callable]
flow = self._batch_norm(flow)
flow = tf.ensure_shape(flow, residual.shape)
diff --git a/official/projects/assemblenet/train.py b/official/projects/assemblenet/train.py
index 54b682ef059..c32df15685e 100644
--- a/official/projects/assemblenet/train.py
+++ b/official/projects/assemblenet/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/assemblenet/train_test.py b/official/projects/assemblenet/train_test.py
index b3fda067904..e7b5ac557f5 100644
--- a/official/projects/assemblenet/train_test.py
+++ b/official/projects/assemblenet/train_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl import flags
from absl import logging
from absl.testing import flagsaver
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.assemblenet import train as train_lib
from official.vision.dataloaders import tfexample_utils
diff --git a/official/projects/backbone_reuse/README.md b/official/projects/backbone_reuse/README.md
index 371de30f696..0981c9e2260 100644
--- a/official/projects/backbone_reuse/README.md
+++ b/official/projects/backbone_reuse/README.md
@@ -1,5 +1,41 @@
# Proper Reuse of Image Classification Features Improves Object Detection
-Coming soon
-1. CVPR 2022 paper
-2. Table of results
\ No newline at end of file
+This project brings the backbone freezing training approach into the Mask-RCNN
+architecture. Please see the paper for more details
+\([arxiv](https://arxiv.org/abs/2204.00484) - selected for oral presentation at
+CVPR 2022\).
+
+### Training Mask-Rcnn Models with backbone frozen.
+
+#### Freezing Resnet-RS-101 checkpoint (ImageNet pretrained).
+
+1. Download the ResNet-RS-101 pretrained checkpoint from
+ [TF-Vision Model Garden](https://github.com/tensorflow/models/tree/master/official/vision#resnet-rs-models-trained-with-various-settings),
+ \([checkpoint](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i192.tar.gz)\)
+
+2. Config files used in our Resnet-101 ablations are included in the
+ [configs folder](https://github.com/tensorflow/models/tree/master/official/projects/backbone_reuse/configs/experiments/faster_rcnn).
+ Select one according to the target architecture (FPN, NASFPN, NASFPN +
+ Cascades) and training schedule preference (shorter--72 epochs, or longer
+ --600 epochs).
+
+3. Change the config flag `init_checkpoint` to point to the downloaded file.
+
+You are all set. Follow the standard TFVision Mask-Rcnn training pipeline to
+complete the training.
+
+#### How does it work?
+
+The config files set the task's flag `freeze_backbone: true`. This flag prevents
+the pretrained backbone weights from being updated during the downstream model
+training.
+
+## Citation
+
+```
+@inproceedings{vasconcelos2022backbonefreeze,
+ title = {Proper Reuse of Image Classification Features Improves Object Detection},
+ author = {Cristina Vasconcelos and Vighnesh Birodkar and Vincent Dumoulin},
+ booktitle={CVPR}
+ year={2022},
+```
diff --git a/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_fpn_600epochs.yaml b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_fpn_600epochs.yaml
new file mode 100644
index 00000000000..7b4f285b890
--- /dev/null
+++ b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_fpn_600epochs.yaml
@@ -0,0 +1,34 @@
+task:
+ # init_checkpoint: 'a_pretrained_backbone_checkpoint'
+ init_checkpoint_modules: backbone
+ freeze_backbone: true
+ model:
+ backbone:
+ resnet:
+ model_id: 101
+ replace_stem_max_pool: true
+ resnetd_shortcut: true
+ scale_stem: true
+ se_ratio: 0.25
+ stem_type: v1
+ type: resnet
+ decoder:
+ type: fpn
+ train_data:
+ global_batch_size: 256
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ stepwise:
+ boundaries: [265684, 272615]
+ name: PiecewiseConstantDecay
+ offset: 0
+ values: [0.32, 0.032, 0.0032]
+ type: stepwise
+ steps_per_loop: 462
+ summary_interval: 462
+ train_steps: 277235
diff --git a/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_fpn_72epochs.yaml b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_fpn_72epochs.yaml
new file mode 100644
index 00000000000..436f7a0c7f4
--- /dev/null
+++ b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_fpn_72epochs.yaml
@@ -0,0 +1,34 @@
+task:
+ # init_checkpoint: 'a_pretrained_backbone_checkpoint'
+ init_checkpoint_modules: backbone
+ freeze_backbone: true
+ model:
+ backbone:
+ resnet:
+ model_id: 101
+ replace_stem_max_pool: true
+ resnetd_shortcut: true
+ scale_stem: true
+ se_ratio: 0.25
+ stem_type: v1
+ type: resnet
+ decoder:
+ type: fpn
+ train_data:
+ global_batch_size: 256
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ stepwise:
+ boundaries: [22176, 31416]
+ name: PiecewiseConstantDecay
+ offset: 0
+ values: [0.32, 0.032, 0.0032]
+ type: stepwise
+ steps_per_loop: 462
+ summary_interval: 462
+ train_steps: 33264
diff --git a/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_nasfpn_600epochs.yaml b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_nasfpn_600epochs.yaml
new file mode 100644
index 00000000000..13db989d407
--- /dev/null
+++ b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_nasfpn_600epochs.yaml
@@ -0,0 +1,34 @@
+task:
+ # init_checkpoint: 'a_pretrained_backbone_checkpoint'
+ init_checkpoint_modules: backbone
+ freeze_backbone: true
+ model:
+ backbone:
+ resnet:
+ model_id: 101
+ replace_stem_max_pool: true
+ resnetd_shortcut: true
+ scale_stem: true
+ se_ratio: 0.25
+ stem_type: v1
+ type: resnet
+ decoder:
+ type: nasfpn
+ train_data:
+ global_batch_size: 256
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ stepwise:
+ boundaries: [265684, 272615]
+ name: PiecewiseConstantDecay
+ offset: 0
+ values: [0.32, 0.032, 0.0032]
+ type: stepwise
+ steps_per_loop: 462
+ summary_interval: 462
+ train_steps: 277235
diff --git a/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_nasfpn_72epochs.yaml b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_nasfpn_72epochs.yaml
new file mode 100644
index 00000000000..74a35ee9a53
--- /dev/null
+++ b/official/projects/backbone_reuse/configs/experiments/retinanet/retinanet_resnet101_nasfpn_72epochs.yaml
@@ -0,0 +1,34 @@
+task:
+ # init_checkpoint: 'a_pretrained_backbone_checkpoint'
+ init_checkpoint_modules: backbone
+ freeze_backbone: true
+ model:
+ backbone:
+ resnet:
+ model_id: 101
+ replace_stem_max_pool: true
+ resnetd_shortcut: true
+ scale_stem: true
+ se_ratio: 0.25
+ stem_type: v1
+ type: resnet
+ decoder:
+ type: nasfpn
+ train_data:
+ global_batch_size: 256
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ stepwise:
+ boundaries: [22176, 31416]
+ name: PiecewiseConstantDecay
+ offset: 0
+ values: [0.32, 0.032, 0.0032]
+ type: stepwise
+ steps_per_loop: 462
+ summary_interval: 462
+ train_steps: 33264
diff --git a/official/projects/basnet/configs/basnet.py b/official/projects/basnet/configs/basnet.py
index 3c971d3ca66..e238ea828c2 100644
--- a/official/projects/basnet/configs/basnet.py
+++ b/official/projects/basnet/configs/basnet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -47,7 +47,9 @@ class BASNetModel(hyperparams.Config):
"""BASNet model config."""
input_size: List[int] = dataclasses.field(default_factory=list)
use_bias: bool = False
- norm_activation: common.NormActivation = common.NormActivation()
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
@dataclasses.dataclass
@@ -61,10 +63,14 @@ class Losses(hyperparams.Config):
@dataclasses.dataclass
class BASNetTask(cfg.TaskConfig):
"""The model config."""
- model: BASNetModel = BASNetModel()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- losses: Losses = Losses()
+ model: BASNetModel = dataclasses.field(default_factory=BASNetModel)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
gradient_clip_norm: float = 0.0
init_checkpoint: Optional[str] = None
init_checkpoint_modules: Union[
@@ -99,7 +105,7 @@ def basnet_duts() -> cfg.ExperimentConfig:
config = cfg.ExperimentConfig(
task=BASNetTask(
model=BASNetModel(
- input_size=[None, None, 3],
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
use_bias=True,
norm_activation=common.NormActivation(
activation='relu',
diff --git a/official/projects/basnet/configs/basnet_test.py b/official/projects/basnet/configs/basnet_test.py
index 3e474ab098d..1cf300f7ebc 100644
--- a/official/projects/basnet/configs/basnet_test.py
+++ b/official/projects/basnet/configs/basnet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import exp_factory
diff --git a/official/projects/basnet/evaluation/metrics.py b/official/projects/basnet/evaluation/metrics.py
index 88fb5907222..cad3c0d7b6d 100644
--- a/official/projects/basnet/evaluation/metrics.py
+++ b/official/projects/basnet/evaluation/metrics.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/basnet/evaluation/metrics_test.py b/official/projects/basnet/evaluation/metrics_test.py
index e37b0185ff5..c51f9a9db23 100644
--- a/official/projects/basnet/evaluation/metrics_test.py
+++ b/official/projects/basnet/evaluation/metrics_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for metrics.py."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.basnet.evaluation import metrics
@@ -32,7 +32,7 @@ def test_mae(self):
mae_obj.update_state(labels, inputs)
output = mae_obj.result()
- mae_tf = tf.keras.metrics.MeanAbsoluteError()
+ mae_tf = tf_keras.metrics.MeanAbsoluteError()
mae_tf.reset_state()
mae_tf.update_state(labels[0], inputs[0])
compare = mae_tf.result().numpy()
@@ -51,8 +51,8 @@ def test_max_f(self):
max_f_obj.update_state(labels, inputs)
output = max_f_obj.result()
- pre_tf = tf.keras.metrics.Precision(thresholds=0.78)
- rec_tf = tf.keras.metrics.Recall(thresholds=0.78)
+ pre_tf = tf_keras.metrics.Precision(thresholds=0.78)
+ rec_tf = tf_keras.metrics.Recall(thresholds=0.78)
pre_tf.reset_state()
rec_tf.reset_state()
pre_tf.update_state(labels[0], inputs[0])
diff --git a/official/projects/basnet/losses/basnet_losses.py b/official/projects/basnet/losses/basnet_losses.py
index 023d3c6358f..c58ebe758fd 100644
--- a/official/projects/basnet/losses/basnet_losses.py
+++ b/official/projects/basnet/losses/basnet_losses.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Losses used for BASNet models."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
EPSILON = 1e-5
@@ -22,8 +22,8 @@ class BASNetLoss:
"""BASNet hybrid loss."""
def __init__(self):
- self._binary_crossentropy = tf.keras.losses.BinaryCrossentropy(
- reduction=tf.keras.losses.Reduction.SUM, from_logits=False)
+ self._binary_crossentropy = tf_keras.losses.BinaryCrossentropy(
+ reduction=tf_keras.losses.Reduction.SUM, from_logits=False)
self._ssim = tf.image.ssim
def __call__(self, sigmoids, labels):
diff --git a/official/projects/basnet/modeling/basnet_model.py b/official/projects/basnet/modeling/basnet_model.py
index cef6d456d64..0b08159ccb6 100644
--- a/official/projects/basnet/modeling/basnet_model.py
+++ b/official/projects/basnet/modeling/basnet_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Mapping
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.projects.basnet.modeling import nn_blocks
@@ -52,8 +52,8 @@
]
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class BASNetModel(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class BASNetModel(tf_keras.Model):
"""A BASNet model.
Boundary-Awar network (BASNet) were proposed in:
@@ -88,7 +88,7 @@ def __init__(self,
self.decoder = decoder
self.refinement = refinement
- def call(self, inputs, training=None):
+ def call(self, inputs, training=None): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
features = self.backbone(inputs)
if self.decoder:
@@ -119,13 +119,13 @@ def from_config(cls, config, custom_objects=None):
return cls(**config)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class BASNetEncoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class BASNetEncoder(tf_keras.Model):
"""BASNet encoder."""
def __init__(
self,
- input_specs=tf.keras.layers.InputSpec(shape=[None, None, None, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[None, None, None, 3]),
activation='relu',
use_sync_bn=False,
use_bias=True,
@@ -138,7 +138,7 @@ def __init__(
"""BASNet encoder initialization function.
Args:
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
activation: `str` name of the activation function.
use_sync_bn: if True, use synchronized batch normalization.
use_bias: if True, use bias in conv2d.
@@ -146,9 +146,9 @@ def __init__(
norm_epsilon: `float` small float added to variance to avoid dividing by
zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
**kwargs: keyword arguments to be passed.
"""
@@ -159,22 +159,22 @@ def __init__(
self._norm_momentum = norm_momentum
self._norm_epsilon = norm_epsilon
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
+ self._norm = tf_keras.layers.BatchNormalization
self._kernel_initializer = kernel_initializer
self._kernel_regularizer = kernel_regularizer
self._bias_regularizer = bias_regularizer
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
bn_axis = -1
else:
bn_axis = 1
# Build BASNet Encoder.
- inputs = tf.keras.Input(shape=input_specs.shape[1:])
+ inputs = tf_keras.Input(shape=input_specs.shape[1:])
- x = tf.keras.layers.Conv2D(
+ x = tf_keras.layers.Conv2D(
filters=64, kernel_size=3, strides=1,
use_bias=self._use_bias, padding='same',
kernel_initializer=self._kernel_initializer,
@@ -196,7 +196,7 @@ def __init__(
name='block_group_l{}'.format(i + 2))
endpoints[str(i)] = x
if spec[3]:
- x = tf.keras.layers.MaxPool2D(pool_size=2, strides=2, padding='same')(x)
+ x = tf_keras.layers.MaxPool2D(pool_size=2, strides=2, padding='same')(x)
self._output_specs = {l: endpoints[l].get_shape() for l in endpoints}
super(BASNetEncoder, self).__init__(
inputs=inputs, outputs=endpoints, **kwargs)
@@ -263,9 +263,9 @@ def output_specs(self):
@factory.register_backbone_builder('basnet_encoder')
def build_basnet_encoder(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
model_config,
- l2_regularizer: tf.keras.regularizers.Regularizer = None) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds BASNet Encoder backbone from a config."""
backbone_type = model_config.backbone.type
norm_activation_config = model_config.norm_activation
@@ -281,8 +281,8 @@ def build_basnet_encoder(
kernel_regularizer=l2_regularizer)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class BASNetDecoder(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class BASNetDecoder(tf_keras.layers.Layer):
"""BASNet decoder."""
def __init__(self,
@@ -305,8 +305,8 @@ def __init__(self,
norm_epsilon: `float` small float added to variance to avoid dividing by
zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
**kwargs: keyword arguments to be passed.
"""
super(BASNetDecoder, self).__init__(**kwargs)
@@ -322,12 +322,12 @@ def __init__(self,
}
self._activation = tf_utils.get_activation(activation)
- self._concat = tf.keras.layers.Concatenate(axis=-1)
- self._sigmoid = tf.keras.layers.Activation(activation='sigmoid')
+ self._concat = tf_keras.layers.Concatenate(axis=-1)
+ self._sigmoid = tf_keras.layers.Activation(activation='sigmoid')
def build(self, input_shape):
"""Creates the variables of the BASNet decoder."""
- conv_op = tf.keras.layers.Conv2D
+ conv_op = tf_keras.layers.Conv2D
conv_kwargs = {
'kernel_size': 3,
'strides': 1,
@@ -358,7 +358,7 @@ def build(self, input_shape):
filters=1,
padding='same',
**conv_kwargs))
- self._out_usmps.append(tf.keras.layers.UpSampling2D(
+ self._out_usmps.append(tf_keras.layers.UpSampling2D(
size=spec[6],
interpolation='bilinear'
))
@@ -381,7 +381,7 @@ def build(self, input_shape):
filters=1,
padding='same',
**conv_kwargs))
- self._out_usmps.append(tf.keras.layers.UpSampling2D(
+ self._out_usmps.append(tf_keras.layers.UpSampling2D(
size=spec[6],
interpolation='bilinear'
))
@@ -415,7 +415,7 @@ def call(self, backbone_output: Mapping[str, tf.Tensor]):
for block in blocks:
x = block(x)
sup[str(i+1)] = x
- x = tf.keras.layers.UpSampling2D(
+ x = tf_keras.layers.UpSampling2D(
size=2,
interpolation='bilinear'
)(x)
diff --git a/official/projects/basnet/modeling/basnet_model_test.py b/official/projects/basnet/modeling/basnet_model_test.py
index 8f59904e5d1..a15db6bdec2 100644
--- a/official/projects/basnet/modeling/basnet_model_test.py
+++ b/official/projects/basnet/modeling/basnet_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.basnet.modeling import basnet_model
from official.projects.basnet.modeling import refunet
@@ -32,7 +32,7 @@ def test_basnet_network_creation(
self, input_size):
"""Test for creation of a segmentation network."""
inputs = np.random.rand(2, input_size, input_size, 3)
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = basnet_model.BASNetEncoder()
decoder = basnet_model.BASNetDecoder()
diff --git a/official/projects/basnet/modeling/nn_blocks.py b/official/projects/basnet/modeling/nn_blocks.py
index 1254c9c78d8..6dc6350c057 100644
--- a/official/projects/basnet/modeling/nn_blocks.py
+++ b/official/projects/basnet/modeling/nn_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,13 @@
"""Contains common building blocks for BasNet model."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class ConvBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class ConvBlock(tf_keras.layers.Layer):
"""A (Conv+BN+Activation) block."""
def __init__(self,
@@ -47,9 +47,9 @@ def __init__(self,
dilation_rate: `int`, dilation rate for conv layers.
kernel_size: `int`, kernel size of conv layers.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
activation: `str` name of the activation function.
use_bias: `bool`, whether or not use bias in conv layers.
@@ -75,10 +75,10 @@ def __init__(self,
'norm_epsilon': norm_epsilon
}
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -93,7 +93,7 @@ def build(self, input_shape):
'bias_regularizer': self._config_dict['bias_regularizer'],
}
- self._conv0 = tf.keras.layers.Conv2D(
+ self._conv0 = tf_keras.layers.Conv2D(
filters=self._config_dict['filters'],
kernel_size=self._config_dict['kernel_size'],
strides=self._config_dict['strides'],
@@ -117,8 +117,8 @@ def call(self, inputs, training=None):
return x
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class ResBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class ResBlock(tf_keras.layers.Layer):
"""A residual block."""
def __init__(self,
@@ -147,9 +147,9 @@ def __init__(self,
filters and the resolution.
kernel_initializer: A `str` of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2d.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2d.
Default to None.
activation: A `str` name of the activation function.
use_sync_bn: A `bool`. If True, use synchronized batch normalization.
@@ -173,10 +173,10 @@ def __init__(self,
'norm_epsilon': norm_epsilon
}
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -193,7 +193,7 @@ def build(self, input_shape):
}
if self._config_dict['use_projection']:
- self._shortcut = tf.keras.layers.Conv2D(
+ self._shortcut = tf_keras.layers.Conv2D(
filters=self._config_dict['filters'],
kernel_size=1,
strides=self._config_dict['strides'],
@@ -206,7 +206,7 @@ def build(self, input_shape):
momentum=self._config_dict['norm_momentum'],
epsilon=self._config_dict['norm_epsilon'])
- self._conv1 = tf.keras.layers.Conv2D(
+ self._conv1 = tf_keras.layers.Conv2D(
kernel_size=3,
strides=self._config_dict['strides'],
**conv_kwargs)
@@ -215,7 +215,7 @@ def build(self, input_shape):
momentum=self._config_dict['norm_momentum'],
epsilon=self._config_dict['norm_epsilon'])
- self._conv2 = tf.keras.layers.Conv2D(
+ self._conv2 = tf_keras.layers.Conv2D(
kernel_size=3,
strides=1,
**conv_kwargs)
diff --git a/official/projects/basnet/modeling/refunet.py b/official/projects/basnet/modeling/refunet.py
index a052adc9ef4..3e85cfeaafe 100644
--- a/official/projects/basnet/modeling/refunet.py
+++ b/official/projects/basnet/modeling/refunet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,12 +13,12 @@
# limitations under the License.
"""RefUNet model."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.basnet.modeling import nn_blocks
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class RefUnet(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class RefUnet(tf_keras.layers.Layer):
"""Residual Refinement Module of BASNet.
Boundary-Aware network (BASNet) were proposed in:
@@ -46,9 +46,9 @@ def __init__(self,
norm_epsilon: `float` small float added to variance to avoid dividing by
zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
**kwargs: keyword arguments to be passed.
"""
@@ -63,19 +63,19 @@ def __init__(self,
'kernel_regularizer': kernel_regularizer,
'bias_regularizer': bias_regularizer,
}
- self._concat = tf.keras.layers.Concatenate(axis=-1)
- self._sigmoid = tf.keras.layers.Activation(activation='sigmoid')
- self._maxpool = tf.keras.layers.MaxPool2D(
+ self._concat = tf_keras.layers.Concatenate(axis=-1)
+ self._sigmoid = tf_keras.layers.Activation(activation='sigmoid')
+ self._maxpool = tf_keras.layers.MaxPool2D(
pool_size=2,
strides=2,
padding='valid')
- self._upsample = tf.keras.layers.UpSampling2D(
+ self._upsample = tf_keras.layers.UpSampling2D(
size=2,
interpolation='bilinear')
def build(self, input_shape):
"""Creates the variables of the BASNet decoder."""
- conv_op = tf.keras.layers.Conv2D
+ conv_op = tf_keras.layers.Conv2D
conv_kwargs = {
'kernel_size': 3,
'strides': 1,
diff --git a/official/projects/basnet/serving/basnet.py b/official/projects/basnet/serving/basnet.py
index c9f5cb9a10c..42b1c0358b4 100644
--- a/official/projects/basnet/serving/basnet.py
+++ b/official/projects/basnet/serving/basnet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,21 +14,17 @@
"""Export module for BASNet."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.basnet.tasks import basnet
from official.vision.serving import semantic_segmentation
-MEAN_RGB = (0.485 * 255, 0.456 * 255, 0.406 * 255)
-STDDEV_RGB = (0.229 * 255, 0.224 * 255, 0.225 * 255)
-
-
class BASNetModule(semantic_segmentation.SegmentationModule):
"""BASNet Module."""
def _build_model(self):
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[self._batch_size] + self._input_image_size + [3])
return basnet.build_basnet_model(
diff --git a/official/projects/basnet/serving/export_saved_model.py b/official/projects/basnet/serving/export_saved_model.py
index 417beac57fd..29c9789c7d9 100644
--- a/official/projects/basnet/serving/export_saved_model.py
+++ b/official/projects/basnet/serving/export_saved_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/basnet/tasks/basnet.py b/official/projects/basnet/tasks/basnet.py
index a68463abf7b..42eae56e977 100644
--- a/official/projects/basnet/tasks/basnet.py
+++ b/official/projects/basnet/tasks/basnet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import dataset_fn
from official.core import base_task
@@ -31,9 +31,9 @@
def build_basnet_model(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
model_config: exp_cfg.BASNetModel,
- l2_regularizer: tf.keras.regularizers.Regularizer = None):
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None):
"""Builds BASNet model."""
norm_activation_config = model_config.norm_activation
backbone = basnet_model.BASNetEncoder(
@@ -71,14 +71,14 @@ class BASNetTask(base_task.Task):
def build_model(self):
"""Builds basnet model."""
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None] + self.task_config.model.input_size)
l2_weight_decay = self.task_config.losses.l2_weight_decay
# Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
- l2_regularizer = (tf.keras.regularizers.l2(
+ l2_regularizer = (tf_keras.regularizers.l2(
l2_weight_decay / 2.0) if l2_weight_decay else None)
model = build_basnet_model(
@@ -87,7 +87,7 @@ def build_model(self):
l2_regularizer=l2_regularizer)
return model
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loads pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -203,7 +203,7 @@ def train_step(self, inputs, model, optimizer, metrics=None):
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
@@ -211,7 +211,7 @@ def train_step(self, inputs, model, optimizer, metrics=None):
# Scales back gradient before apply_gradients when LossScaleOptimizer is
# used.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
# Apply gradient clipping.
@@ -261,14 +261,14 @@ def aggregate_logs(self, state=None, step_outputs=None):
self.relaxf_metric.reset_states()
state = self.mae_metric
self.mae_metric.update_state(
- step_outputs[self.mae_metric.name][0],
- step_outputs[self.mae_metric.name][1])
+ step_outputs[self.mae_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ step_outputs[self.mae_metric.name][1]) # pyrefly: ignore[unsupported-operation]
self.maxf_metric.update_state(
- step_outputs[self.maxf_metric.name][0],
- step_outputs[self.maxf_metric.name][1])
+ step_outputs[self.maxf_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ step_outputs[self.maxf_metric.name][1]) # pyrefly: ignore[unsupported-operation]
self.relaxf_metric.update_state(
- step_outputs[self.relaxf_metric.name][0],
- step_outputs[self.relaxf_metric.name][1])
+ step_outputs[self.relaxf_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ step_outputs[self.relaxf_metric.name][1]) # pyrefly: ignore[unsupported-operation]
return state
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
diff --git a/official/projects/basnet/train.py b/official/projects/basnet/train.py
index d30321ac373..7f4959d12ec 100644
--- a/official/projects/basnet/train.py
+++ b/official/projects/basnet/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/bigbird/__init__.py b/official/projects/bigbird/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/bigbird/__init__.py
+++ b/official/projects/bigbird/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/bigbird/encoder.py b/official/projects/bigbird/encoder.py
index 991a7de12bb..32a6ae9efbb 100644
--- a/official/projects/bigbird/encoder.py
+++ b/official/projects/bigbird/encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,10 @@
"""Transformer-based text encoder network."""
# pylint: disable=g-classes-have-attributes
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import activations
+from official.modeling import tf_utils
from official.nlp import modeling
from official.nlp.modeling import layers
from official.projects.bigbird import recompute_grad
@@ -50,8 +51,8 @@ def f(*args):
return f(emb, *mask)
-@tf.keras.utils.register_keras_serializable(package='Text')
-class BigBirdEncoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BigBirdEncoder(tf_keras.Model):
"""Transformer-based encoder network with BigBird attentions.
*Note* that the network is constructed by
@@ -100,15 +101,15 @@ def __init__(self,
activation=activations.gelu,
dropout_rate=0.1,
attention_dropout_rate=0.1,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
embedding_width=None,
use_gradient_checkpointing=False,
**kwargs):
- activation = tf.keras.activations.get(activation)
- initializer = tf.keras.initializers.get(initializer)
+ activation = tf_keras.activations.get(activation)
+ initializer = tf_keras.initializers.get(initializer)
if use_gradient_checkpointing:
- tf.keras.layers.Dropout = recomputing_dropout.RecomputingDropout
+ tf_keras.layers.Dropout = recomputing_dropout.RecomputingDropout
layer_cls = RecomputeTransformerLayer
else:
layer_cls = layers.TransformerScaffold
@@ -124,18 +125,22 @@ def __init__(self,
'intermediate_size': intermediate_size,
'block_size': block_size,
'num_rand_blocks': num_rand_blocks,
- 'activation': tf.keras.activations.serialize(activation),
+ 'activation': tf_utils.serialize_activation(
+ activation, use_legacy_format=True
+ ),
'dropout_rate': dropout_rate,
'attention_dropout_rate': attention_dropout_rate,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_utils.serialize_initializer(
+ initializer, use_legacy_format=True
+ ),
'embedding_width': embedding_width,
}
- word_ids = tf.keras.layers.Input(
+ word_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- mask = tf.keras.layers.Input(
+ mask = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_mask')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
if embedding_width is None:
@@ -161,19 +166,19 @@ def __init__(self,
name='type_embeddings')
type_embeddings = self._type_embedding_layer(type_ids)
- embeddings = tf.keras.layers.Add()(
+ embeddings = tf_keras.layers.Add()(
[word_embeddings, position_embeddings, type_embeddings])
- self._embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
embeddings = self._embedding_norm_layer(embeddings)
- embeddings = tf.keras.layers.Dropout(rate=dropout_rate)(embeddings)
+ embeddings = tf_keras.layers.Dropout(rate=dropout_rate)(embeddings)
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
if embedding_width != hidden_size:
- self._embedding_projection = tf.keras.layers.experimental.EinsumDense(
+ self._embedding_projection = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
diff --git a/official/projects/bigbird/encoder_test.py b/official/projects/bigbird/encoder_test.py
index 9b683372050..93a6f4f5c0f 100644
--- a/official/projects/bigbird/encoder_test.py
+++ b/official/projects/bigbird/encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for official.nlp.projects.bigbird.encoder."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.bigbird import encoder
@@ -53,7 +53,7 @@ def test_save_restore(self):
ref_outputs = network(inputs)
model_path = self.get_temp_dir() + "/model"
network.save(model_path)
- loaded = tf.keras.models.load_model(model_path)
+ loaded = tf_keras.models.load_model(model_path)
outputs = loaded(inputs)
self.assertAllClose(outputs["sequence_output"],
ref_outputs["sequence_output"])
diff --git a/official/projects/bigbird/experiment_configs.py b/official/projects/bigbird/experiment_configs.py
index 0ad3e4e5820..2e70975e11b 100644
--- a/official/projects/bigbird/experiment_configs.py
+++ b/official/projects/bigbird/experiment_configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/bigbird/recompute_grad.py b/official/projects/bigbird/recompute_grad.py
index be9424d6031..143c7138319 100644
--- a/official/projects/bigbird/recompute_grad.py
+++ b/official/projects/bigbird/recompute_grad.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class RecomputeContext(
@@ -138,7 +138,7 @@ def _first_element(x):
if not dtype.is_floating:
raise ValueError('_force_data_dependency only supports floating dtypes.')
zero = np.finfo(dtype.as_numpy_dtype).tiny * first_compute_sum
- return [
+ return [ # pyrefly: ignore[bad-return]
x + tf.cast(zero, x.dtype) if x is not None else None
for x in then_compute
]
diff --git a/official/projects/bigbird/recomputing_dropout.py b/official/projects/bigbird/recomputing_dropout.py
index fb3e565b966..85459841101 100644
--- a/official/projects/bigbird/recomputing_dropout.py
+++ b/official/projects/bigbird/recomputing_dropout.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Keras dropout layer that is aware of `RecomputeContext`."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.bigbird import recompute_grad as recompute_grad_lib
from official.projects.bigbird import stateless_dropout as stateless_dropout_lib
@@ -57,8 +57,8 @@ def smart_cond(pred, true_fn=None, false_fn=None, name=None):
# See https://www.tensorflow.org/api_docs/python/tf/keras/layers/Dropout.
-class RecomputingDropout(tf.keras.layers.Layer):
- """`tf.keras.layers.Dropout` that supports `recompute_grad`."""
+class RecomputingDropout(tf_keras.layers.Layer):
+ """`tf_keras.layers.Dropout` that supports `recompute_grad`."""
def __init__(self,
rate,
@@ -78,7 +78,7 @@ def __init__(self,
seed: A Python integer to use as random seed.
force_recomputation: If `True`, then raises an error if called outside a
recompute context.
- **kwargs: Keyword arguments for `tf.keras.layers.Layer`.
+ **kwargs: Keyword arguments for `tf_keras.layers.Layer`.
"""
super(RecomputingDropout, self).__init__(**kwargs)
@@ -121,7 +121,7 @@ def call(self, inputs, training=None):
a recompute context.
"""
if training is None:
- training = tf.keras.backend.learning_phase()
+ training = tf_keras.backend.learning_phase()
def dropped_inputs():
"""Randomly drops elements of `inputs` when `training=True`."""
diff --git a/official/projects/bigbird/stateless_dropout.py b/official/projects/bigbird/stateless_dropout.py
index 49941253c64..55d84f2cbc6 100644
--- a/official/projects/bigbird/stateless_dropout.py
+++ b/official/projects/bigbird/stateless_dropout.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from typing import Optional, Sequence, Text, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _as_shape(shape: Union[Sequence[int], tf.TensorShape]) -> tf.TensorShape:
@@ -104,7 +104,7 @@ def stateless_dropout(x: tf.Tensor,
return x
rate = tf.convert_to_tensor(rate, dtype=x.dtype, name='rate')
- rate.shape.assert_has_rank(0)
+ rate.shape.assert_has_rank(0) # pyrefly: ignore[missing-attribute]
noise_shape = _get_noise_shape(x, noise_shape)
# Sample a uniform distribution on [0.0, 1.0) and select values larger than
# rate.
diff --git a/official/vision/beta/projects/centernet/README.md b/official/projects/centernet/README.md
similarity index 97%
rename from official/vision/beta/projects/centernet/README.md
rename to official/projects/centernet/README.md
index 6c80d8e371e..d61e22f3798 100644
--- a/official/vision/beta/projects/centernet/README.md
+++ b/official/projects/centernet/README.md
@@ -22,7 +22,7 @@ heatmaps (one heatmap for each class) is needed to predict the object. CenterNet
proves that this can be done without a significant difference in accuracy.
-## Enviroment setup
+## Environment setup
The code can be run on multiple GPUs or TPUs with different distribution
strategies. See the TensorFlow distributed training
@@ -37,7 +37,7 @@ install -r ./official/requirements.txt`
To train the model on Coco, try the following command:
```
-python3 -m official.vision.beta.projects.centernet.train \
+python3 -m official.projects.centernet.train \
--mode=train_and_eval \
--experiment=centernet_hourglass_coco \
--model_dir={MODEL_DIR} \
diff --git a/official/vision/beta/projects/centernet/configs/__init__.py b/official/projects/centernet/__init__.py
similarity index 89%
rename from official/vision/beta/projects/centernet/configs/__init__.py
rename to official/projects/centernet/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/vision/beta/projects/centernet/configs/__init__.py
+++ b/official/projects/centernet/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/beta/projects/centernet/utils/checkpoints/__init__.py b/official/projects/centernet/common/__init__.py
similarity index 89%
rename from official/vision/beta/projects/centernet/utils/checkpoints/__init__.py
rename to official/projects/centernet/common/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/vision/beta/projects/centernet/utils/checkpoints/__init__.py
+++ b/official/projects/centernet/common/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/beta/projects/simclr/common/registry_imports.py b/official/projects/centernet/common/registry_imports.py
similarity index 65%
rename from official/vision/beta/projects/simclr/common/registry_imports.py
rename to official/projects/centernet/common/registry_imports.py
index c674017fd47..ab58628de3a 100644
--- a/official/vision/beta/projects/simclr/common/registry_imports.py
+++ b/official/projects/centernet/common/registry_imports.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,8 +15,8 @@
"""All necessary imports for registration."""
# pylint: disable=unused-import
+from official.projects.centernet.configs import centernet
+from official.projects.centernet.modeling import centernet_model
+from official.projects.centernet.modeling.backbones import hourglass
+from official.projects.centernet.tasks import centernet as centernet_task
from official.vision import registry_imports
-from official.vision.beta.projects.simclr.configs import simclr
-from official.vision.beta.projects.simclr.losses import contrastive_losses
-from official.vision.beta.projects.simclr.modeling import simclr_model
-from official.vision.beta.projects.simclr.tasks import simclr as simclr_task
diff --git a/official/vision/beta/projects/centernet/ops/__init__.py b/official/projects/centernet/configs/__init__.py
similarity index 89%
rename from official/vision/beta/projects/centernet/ops/__init__.py
rename to official/projects/centernet/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/vision/beta/projects/centernet/ops/__init__.py
+++ b/official/projects/centernet/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/beta/projects/centernet/configs/backbones.py b/official/projects/centernet/configs/backbones.py
similarity index 88%
rename from official/vision/beta/projects/centernet/configs/backbones.py
rename to official/projects/centernet/configs/backbones.py
index 170aa496932..680fcbe0ec1 100644
--- a/official/vision/beta/projects/centernet/configs/backbones.py
+++ b/official/projects/centernet/configs/backbones.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -32,4 +32,4 @@ class Hourglass(hyperparams.Config):
@dataclasses.dataclass
class Backbone(backbones.Backbone):
- hourglass: Hourglass = Hourglass()
+ hourglass: Hourglass = dataclasses.field(default_factory=Hourglass)
diff --git a/official/vision/beta/projects/centernet/configs/centernet.py b/official/projects/centernet/configs/centernet.py
similarity index 78%
rename from official/vision/beta/projects/centernet/configs/centernet.py
rename to official/projects/centernet/configs/centernet.py
index 1fd026f4a49..0a136f33205 100644
--- a/official/vision/beta/projects/centernet/configs/centernet.py
+++ b/official/projects/centernet/configs/centernet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,7 +22,7 @@
from official.core import exp_factory
from official.modeling import hyperparams
from official.modeling import optimization
-from official.vision.beta.projects.centernet.configs import backbones
+from official.projects.centernet.configs import backbones
from official.vision.configs import common
@@ -37,8 +37,12 @@ class TfExampleDecoder(hyperparams.Config):
@dataclasses.dataclass
class DataDecoder(hyperparams.OneOfConfig):
type: Optional[str] = 'simple_decoder'
- simple_decoder: TfExampleDecoder = TfExampleDecoder()
- label_map_decoder: TfExampleDecoderLabelMap = TfExampleDecoderLabelMap()
+ simple_decoder: TfExampleDecoder = dataclasses.field(
+ default_factory=TfExampleDecoder
+ )
+ label_map_decoder: TfExampleDecoderLabelMap = dataclasses.field(
+ default_factory=TfExampleDecoderLabelMap
+ )
@dataclasses.dataclass
@@ -66,8 +70,8 @@ class DataConfig(cfg.DataConfig):
global_batch_size: int = 32
is_training: bool = True
dtype: str = 'float16'
- decoder: DataDecoder = DataDecoder()
- parser: Parser = Parser()
+ decoder: DataDecoder = dataclasses.field(default_factory=DataDecoder)
+ parser: Parser = dataclasses.field(default_factory=Parser)
shuffle_buffer_size: int = 10000
file_type: str = 'tfrecord'
drop_remainder: bool = True
@@ -82,7 +86,7 @@ class DetectionLoss(hyperparams.Config):
@dataclasses.dataclass
class Losses(hyperparams.Config):
- detection: DetectionLoss = DetectionLoss()
+ detection: DetectionLoss = dataclasses.field(default_factory=DetectionLoss)
gaussian_iou: float = 0.7
class_offset: int = 1
@@ -112,13 +116,21 @@ class CenterNetModel(hyperparams.Config):
num_classes: int = 90
max_num_instances: int = 128
input_size: List[int] = dataclasses.field(default_factory=list)
- backbone: backbones.Backbone = backbones.Backbone(
- type='hourglass', hourglass=backbones.Hourglass(model_id=52))
- head: CenterNetHead = CenterNetHead()
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='hourglass', hourglass=backbones.Hourglass(model_id=52)
+ )
+ )
+ head: CenterNetHead = dataclasses.field(default_factory=CenterNetHead)
# pylint: disable=line-too-long
- detection_generator: CenterNetDetectionGenerator = CenterNetDetectionGenerator()
- norm_activation: common.NormActivation = common.NormActivation(
- norm_momentum=0.1, norm_epsilon=1e-5, use_sync_bn=True)
+ detection_generator: CenterNetDetectionGenerator = dataclasses.field(
+ default_factory=CenterNetDetectionGenerator
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ norm_momentum=0.1, norm_epsilon=1e-5, use_sync_bn=True
+ )
+ )
@dataclasses.dataclass
@@ -129,17 +141,25 @@ class CenterNetDetection(hyperparams.Config):
@dataclasses.dataclass
class CenterNetSubTasks(hyperparams.Config):
- detection: CenterNetDetection = CenterNetDetection()
+ detection: CenterNetDetection = dataclasses.field(
+ default_factory=CenterNetDetection
+ )
@dataclasses.dataclass
class CenterNetTask(cfg.TaskConfig):
"""Config for centernet task."""
- model: CenterNetModel = CenterNetModel()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- subtasks: CenterNetSubTasks = CenterNetSubTasks()
- losses: Losses = Losses()
+ model: CenterNetModel = dataclasses.field(default_factory=CenterNetModel)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ subtasks: CenterNetSubTasks = dataclasses.field(
+ default_factory=CenterNetSubTasks
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
gradient_clip_norm: float = 10.0
per_category_metrics: bool = False
weight_decay: float = 5e-4
diff --git a/official/vision/beta/projects/centernet/configs/centernet_test.py b/official/projects/centernet/configs/centernet_test.py
similarity index 83%
rename from official/vision/beta/projects/centernet/configs/centernet_test.py
rename to official/projects/centernet/configs/centernet_test.py
index 50c4bd0296a..7e44682ea9e 100644
--- a/official/vision/beta/projects/centernet/configs/centernet_test.py
+++ b/official/projects/centernet/configs/centernet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,12 +15,12 @@
"""Tests for centernet."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import exp_factory
-from official.vision.beta.projects.centernet.common import registry_imports # pylint: disable=unused-import
-from official.vision.beta.projects.centernet.configs import centernet as exp_cfg
+from official.projects.centernet.common import registry_imports # pylint: disable=unused-import
+from official.projects.centernet.configs import centernet as exp_cfg
class CenterNetConfigTest(tf.test.TestCase, parameterized.TestCase):
diff --git a/official/vision/beta/projects/centernet/configs/experiments/coco-centernet-hourglass-gpu.yaml b/official/projects/centernet/configs/experiments/coco-centernet-hourglass-gpu.yaml
similarity index 100%
rename from official/vision/beta/projects/centernet/configs/experiments/coco-centernet-hourglass-gpu.yaml
rename to official/projects/centernet/configs/experiments/coco-centernet-hourglass-gpu.yaml
diff --git a/official/vision/beta/projects/centernet/configs/experiments/coco-centernet-hourglass-tpu.yaml b/official/projects/centernet/configs/experiments/coco-centernet-hourglass-tpu.yaml
similarity index 100%
rename from official/vision/beta/projects/centernet/configs/experiments/coco-centernet-hourglass-tpu.yaml
rename to official/projects/centernet/configs/experiments/coco-centernet-hourglass-tpu.yaml
diff --git a/official/vision/beta/projects/__init__.py b/official/projects/centernet/dataloaders/__init__.py
similarity index 89%
rename from official/vision/beta/projects/__init__.py
rename to official/projects/centernet/dataloaders/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/vision/beta/projects/__init__.py
+++ b/official/projects/centernet/dataloaders/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/beta/projects/centernet/dataloaders/centernet_input.py b/official/projects/centernet/dataloaders/centernet_input.py
similarity index 97%
rename from official/vision/beta/projects/centernet/dataloaders/centernet_input.py
rename to official/projects/centernet/dataloaders/centernet_input.py
index e337ced1aaf..a51ca364f72 100644
--- a/official/vision/beta/projects/centernet/dataloaders/centernet_input.py
+++ b/official/projects/centernet/dataloaders/centernet_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,11 +16,11 @@
from typing import Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.ops import box_list
-from official.vision.beta.projects.centernet.ops import box_list_ops
-from official.vision.beta.projects.centernet.ops import preprocess_ops as cn_prep_ops
+from official.projects.centernet.ops import box_list
+from official.projects.centernet.ops import box_list_ops
+from official.projects.centernet.ops import preprocess_ops as cn_prep_ops
from official.vision.dataloaders import parser
from official.vision.dataloaders import utils
from official.vision.ops import box_ops
diff --git a/official/projects/centernet/losses/__init__.py b/official/projects/centernet/losses/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/losses/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/centernet/losses/centernet_losses.py b/official/projects/centernet/losses/centernet_losses.py
similarity index 97%
rename from official/vision/beta/projects/centernet/losses/centernet_losses.py
rename to official/projects/centernet/losses/centernet_losses.py
index 4cb7b0fe8d0..3c3802b07ec 100644
--- a/official/vision/beta/projects/centernet/losses/centernet_losses.py
+++ b/official/projects/centernet/losses/centernet_losses.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Losses for centernet model."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class PenaltyReducedLogisticFocalLoss(object):
diff --git a/official/vision/beta/projects/centernet/losses/centernet_losses_test.py b/official/projects/centernet/losses/centernet_losses_test.py
similarity index 95%
rename from official/vision/beta/projects/centernet/losses/centernet_losses_test.py
rename to official/projects/centernet/losses/centernet_losses_test.py
index 903f28aed83..f6fbe0206b7 100644
--- a/official/vision/beta/projects/centernet/losses/centernet_losses_test.py
+++ b/official/projects/centernet/losses/centernet_losses_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
"""Tests for losses of centernet model."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.losses import centernet_losses
+from official.projects.centernet.losses import centernet_losses
LOG_2 = np.log(2)
LOG_3 = np.log(3)
diff --git a/official/projects/centernet/modeling/__init__.py b/official/projects/centernet/modeling/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/modeling/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/centernet/modeling/backbones/__init__.py b/official/projects/centernet/modeling/backbones/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/modeling/backbones/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/centernet/modeling/backbones/hourglass.py b/official/projects/centernet/modeling/backbones/hourglass.py
similarity index 91%
rename from official/vision/beta/projects/centernet/modeling/backbones/hourglass.py
rename to official/projects/centernet/modeling/backbones/hourglass.py
index d96ca5ccb6d..b33a0f48ae1 100644
--- a/official/vision/beta/projects/centernet/modeling/backbones/hourglass.py
+++ b/official/projects/centernet/modeling/backbones/hourglass.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,10 +16,10 @@
from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
-from official.vision.beta.projects.centernet.modeling.layers import cn_nn_blocks
+from official.projects.centernet.modeling.layers import cn_nn_blocks
from official.vision.modeling.backbones import factory
from official.vision.modeling.backbones import mobilenet
from official.vision.modeling.layers import nn_blocks
@@ -48,14 +48,14 @@
}
-class Hourglass(tf.keras.Model):
+class Hourglass(tf_keras.Model):
"""CenterNet Hourglass backbone."""
def __init__(
self,
model_id: int,
input_channel_dims: int,
- input_specs=tf.keras.layers.InputSpec(shape=[None, None, None, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[None, None, None, 3]),
num_hourglasses: int = 1,
initial_downsample: bool = True,
activation: str = 'relu',
@@ -63,8 +63,8 @@ def __init__(
norm_momentum=0.1,
norm_epsilon=1e-5,
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
**kwargs):
"""Initialize Hourglass backbone.
@@ -72,7 +72,7 @@ def __init__(
model_id: An `int` of the scale of Hourglass backbone model.
input_channel_dims: `int`, number of filters used to downsample the
input image.
- input_specs: A `tf.keras.layers.InputSpec` of specs of the input tensor.
+ input_specs: A `tf_keras.layers.InputSpec` of specs of the input tensor.
num_hourglasses: `int``, number of hourglass blocks in backbone. For
example, hourglass-104 has two hourglass-52 modules.
initial_downsample: `bool`, whether or not to downsample the input.
@@ -81,9 +81,9 @@ def __init__(
norm_momentum: `float`, momentum for the batch normalization layers.
norm_epsilon: `float`, epsilon for the batch normalization layers.
kernel_initializer: A `str` for kernel initializer of conv layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2D.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
Default to None.
**kwargs: Additional keyword arguments to be passed.
"""
@@ -104,7 +104,7 @@ def __init__(
self._channel_dims_per_stage = [item * self._input_channel_dims
for item in specs['channel_dims_per_stage']]
- inputs = tf.keras.layers.Input(shape=input_specs.shape[1:])
+ inputs = tf_keras.layers.Input(shape=input_specs.shape[1:])
inp_filters = self._channel_dims_per_stage[0]
@@ -203,8 +203,8 @@ def __init__(
norm_epsilon=self._norm_epsilon
)(x_hg)
- x_downsampled = tf.keras.layers.Add()([inter_hg_conv1, inter_hg_conv2])
- x_downsampled = tf.keras.layers.ReLU()(x_downsampled)
+ x_downsampled = tf_keras.layers.Add()([inter_hg_conv1, inter_hg_conv2])
+ x_downsampled = tf_keras.layers.ReLU()(x_downsampled)
x_downsampled = nn_blocks.ResidualBlock(
filters=inp_filters,
@@ -250,11 +250,11 @@ def output_specs(self):
@factory.register_backbone_builder('hourglass')
def build_hourglass(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
norm_activation_config: hyperparams.Config,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None
- ) -> tf.keras.Model:
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+ ) -> tf_keras.Model:
"""Builds Hourglass backbone from a configuration."""
backbone_type = backbone_config.type
backbone_cfg = backbone_config.get()
diff --git a/official/vision/beta/projects/centernet/modeling/backbones/hourglass_test.py b/official/projects/centernet/modeling/backbones/hourglass_test.py
similarity index 74%
rename from official/vision/beta/projects/centernet/modeling/backbones/hourglass_test.py
rename to official/projects/centernet/modeling/backbones/hourglass_test.py
index e845044bec8..6f372ff997f 100644
--- a/official/vision/beta/projects/centernet/modeling/backbones/hourglass_test.py
+++ b/official/projects/centernet/modeling/backbones/hourglass_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,11 +16,11 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.common import registry_imports # pylint: disable=unused-import
-from official.vision.beta.projects.centernet.configs import backbones
-from official.vision.beta.projects.centernet.modeling.backbones import hourglass
+from official.projects.centernet.common import registry_imports # pylint: disable=unused-import
+from official.projects.centernet.configs import backbones
+from official.projects.centernet.modeling.backbones import hourglass
from official.vision.configs import common
@@ -28,7 +28,7 @@ class HourglassTest(tf.test.TestCase, parameterized.TestCase):
def test_hourglass(self):
backbone = hourglass.build_hourglass(
- input_specs=tf.keras.layers.InputSpec(shape=[None, 512, 512, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[None, 512, 512, 3]),
backbone_config=backbones.Backbone(type='hourglass'),
norm_activation_config=common.NormActivation(use_sync_bn=True)
)
diff --git a/official/vision/beta/projects/centernet/modeling/centernet_model.py b/official/projects/centernet/modeling/centernet_model.py
similarity index 81%
rename from official/vision/beta/projects/centernet/modeling/centernet_model.py
rename to official/projects/centernet/modeling/centernet_model.py
index 3b8de7534b2..2ab16f5f079 100644
--- a/official/vision/beta/projects/centernet/modeling/centernet_model.py
+++ b/official/projects/centernet/modeling/centernet_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,16 +16,16 @@
from typing import Mapping, Union, Any
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class CenterNetModel(tf.keras.Model):
+class CenterNetModel(tf_keras.Model):
"""CenterNet Model."""
def __init__(self,
- backbone: tf.keras.Model,
- head: tf.keras.Model,
- detection_generator: tf.keras.layers.Layer,
+ backbone: tf_keras.Model,
+ head: tf_keras.Model,
+ detection_generator: tf_keras.layers.Layer,
**kwargs):
"""CenterNet Model.
@@ -41,9 +41,9 @@ def __init__(self,
self._detection_generator = detection_generator
self._head = head
- def call(self,
+ def call(self, # pytype: disable=annotation-type-mismatch,signature-mismatch
inputs: tf.Tensor,
- training: bool = None,
+ training: bool = None, # pyrefly: ignore[bad-function-definition]
**kwargs) -> Mapping[str, tf.Tensor]:
features = self._backbone(inputs)
raw_outputs = self._head(features)
@@ -55,7 +55,7 @@ def call(self,
@property
def checkpoint_items(
- self) -> Mapping[str, Union[tf.keras.Model, tf.keras.layers.Layer]]:
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
"""Returns a dictionary of items to be additionally checkpointed."""
items = dict(backbone=self.backbone, head=self.head)
diff --git a/official/vision/beta/projects/centernet/modeling/centernet_model_test.py b/official/projects/centernet/modeling/centernet_model_test.py
similarity index 79%
rename from official/vision/beta/projects/centernet/modeling/centernet_model_test.py
rename to official/projects/centernet/modeling/centernet_model_test.py
index 08cb7fca3e5..f98b0703232 100644
--- a/official/vision/beta/projects/centernet/modeling/centernet_model_test.py
+++ b/official/projects/centernet/modeling/centernet_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,13 @@
"""Test for centernet detection model."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.configs import backbones
-from official.vision.beta.projects.centernet.modeling import centernet_model
-from official.vision.beta.projects.centernet.modeling.backbones import hourglass
-from official.vision.beta.projects.centernet.modeling.heads import centernet_head
-from official.vision.beta.projects.centernet.modeling.layers import detection_generator
+from official.projects.centernet.configs import backbones
+from official.projects.centernet.modeling import centernet_model
+from official.projects.centernet.modeling.backbones import hourglass
+from official.projects.centernet.modeling.heads import centernet_head
+from official.projects.centernet.modeling.layers import detection_generator
from official.vision.configs import common
@@ -29,7 +29,7 @@ class CenterNetTest(parameterized.TestCase, tf.test.TestCase):
def testBuildCenterNet(self):
backbone = hourglass.build_hourglass(
- input_specs=tf.keras.layers.InputSpec(shape=[None, 512, 512, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[None, 512, 512, 3]),
backbone_config=backbones.Backbone(type='hourglass'),
norm_activation_config=common.NormActivation(use_sync_bn=True)
)
diff --git a/official/projects/centernet/modeling/heads/__init__.py b/official/projects/centernet/modeling/heads/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/modeling/heads/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/centernet/modeling/heads/centernet_head.py b/official/projects/centernet/modeling/heads/centernet_head.py
similarity index 90%
rename from official/vision/beta/projects/centernet/modeling/heads/centernet_head.py
rename to official/projects/centernet/modeling/heads/centernet_head.py
index 754ef8d5469..6f201ca40ea 100644
--- a/official/vision/beta/projects/centernet/modeling/heads/centernet_head.py
+++ b/official/projects/centernet/modeling/heads/centernet_head.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,14 +14,14 @@
"""Contains the definitions of head for CenterNet."""
-from typing import Any, Mapping, Dict, List
+from typing import Any, Dict, List, Mapping
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.modeling.layers import cn_nn_blocks
+from official.projects.centernet.modeling.layers import cn_nn_blocks
-class CenterNetHead(tf.keras.Model):
+class CenterNetHead(tf_keras.Model):
"""CenterNet Head."""
def __init__(self,
@@ -61,7 +61,7 @@ def __init__(self,
self._heatmap_bias = heatmap_bias
self._num_inputs = len(input_levels)
- inputs = {level: tf.keras.layers.Input(shape=self._input_specs[level][1:])
+ inputs = {level: tf_keras.layers.Input(shape=self._input_specs[level][1:])
for level in input_levels}
outputs = {}
@@ -102,4 +102,4 @@ def from_config(cls, config, custom_objects=None):
@property
def output_specs(self) -> Mapping[str, tf.TensorShape]:
"""A dict of {level: TensorShape} pairs for the model output."""
- return self._output_specs
+ return self._output_specs # pytype: disable=bad-return-type
diff --git a/official/vision/beta/projects/centernet/modeling/heads/centernet_head_test.py b/official/projects/centernet/modeling/heads/centernet_head_test.py
similarity index 87%
rename from official/vision/beta/projects/centernet/modeling/heads/centernet_head_test.py
rename to official/projects/centernet/modeling/heads/centernet_head_test.py
index bcaad046e51..921db906123 100644
--- a/official/vision/beta/projects/centernet/modeling/heads/centernet_head_test.py
+++ b/official/projects/centernet/modeling/heads/centernet_head_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,9 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.modeling.heads import centernet_head
+from official.projects.centernet.modeling.heads import centernet_head
class CenterNetHeadTest(tf.test.TestCase, parameterized.TestCase):
@@ -30,8 +30,8 @@ def test_decoder_shape(self):
'ct_size': 2,
}
input_specs = {
- '2_0': tf.keras.layers.InputSpec(shape=(None, 128, 128, 256)).shape,
- '2': tf.keras.layers.InputSpec(shape=(None, 128, 128, 256)).shape,
+ '2_0': tf_keras.layers.InputSpec(shape=(None, 128, 128, 256)).shape,
+ '2': tf_keras.layers.InputSpec(shape=(None, 128, 128, 256)).shape,
}
input_levels = ['2', '2_0']
diff --git a/official/projects/centernet/modeling/layers/__init__.py b/official/projects/centernet/modeling/layers/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/modeling/layers/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/centernet/modeling/layers/cn_nn_blocks.py b/official/projects/centernet/modeling/layers/cn_nn_blocks.py
similarity index 92%
rename from official/vision/beta/projects/centernet/modeling/layers/cn_nn_blocks.py
rename to official/projects/centernet/modeling/layers/cn_nn_blocks.py
index eba920e4283..66de09c57ab 100644
--- a/official/vision/beta/projects/centernet/modeling/layers/cn_nn_blocks.py
+++ b/official/projects/centernet/modeling/layers/cn_nn_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import List, Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.modeling.layers import nn_blocks
@@ -41,8 +41,8 @@ def _make_repeated_residual_blocks(
initial_stride: int = 1,
initial_skip_conv: bool = False,
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
):
"""Stack Residual blocks one after the other.
@@ -61,9 +61,9 @@ def _make_repeated_residual_blocks(
convolution. This is useful when the number of channels in the input
are not the same as residual_channels.
kernel_initializer: A `str` for kernel initializer of convolutional layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2D.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
Default to None.
Returns:
@@ -120,10 +120,10 @@ def _make_repeated_residual_blocks(
kernel_regularizer=kernel_regularizer,
bias_regularizer=bias_regularizer))
- return tf.keras.Sequential(blocks)
+ return tf_keras.Sequential(blocks)
-class HourglassBlock(tf.keras.layers.Layer):
+class HourglassBlock(tf_keras.layers.Layer):
"""Hourglass module: an encoder-decoder block."""
def __init__(
@@ -135,8 +135,8 @@ def __init__(
norm_momentum: float = 0.1,
norm_epsilon: float = 1e-5,
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
**kwargs):
"""Initialize Hourglass module.
@@ -158,9 +158,9 @@ def __init__(
norm_momentum: `float`, momentum for the batch normalization layers.
norm_epsilon: `float`, epsilon for the batch normalization layers.
kernel_initializer: A `str` for kernel initializer of conv layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2D.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
Default to None.
**kwargs: Additional keyword arguments to be passed.
"""
@@ -241,7 +241,7 @@ def build(self, input_shape):
kernel_initializer=self._kernel_initializer,
kernel_regularizer=self._kernel_regularizer)
- self.upsample_layer = tf.keras.layers.UpSampling2D(
+ self.upsample_layer = tf_keras.layers.UpSampling2D(
size=2,
interpolation='nearest')
@@ -273,7 +273,7 @@ def get_config(self):
return config
-class CenterNetHeadConv(tf.keras.layers.Layer):
+class CenterNetHeadConv(tf_keras.layers.Layer):
"""Convolution block for the CenterNet head."""
def __init__(self,
@@ -297,15 +297,15 @@ def __init__(self,
def build(self, input_shape):
n_channels = input_shape[-1]
- self.conv1 = tf.keras.layers.Conv2D(
+ self.conv1 = tf_keras.layers.Conv2D(
filters=n_channels,
kernel_size=(3, 3),
padding='same')
- self.relu = tf.keras.layers.ReLU()
+ self.relu = tf_keras.layers.ReLU()
# Initialize bias to the last Conv2D Layer
- self.conv2 = tf.keras.layers.Conv2D(
+ self.conv2 = tf_keras.layers.Conv2D(
filters=self._output_filters,
kernel_size=(1, 1),
padding='valid',
diff --git a/official/vision/beta/projects/centernet/modeling/layers/cn_nn_blocks_test.py b/official/projects/centernet/modeling/layers/cn_nn_blocks_test.py
similarity index 90%
rename from official/vision/beta/projects/centernet/modeling/layers/cn_nn_blocks_test.py
rename to official/projects/centernet/modeling/layers/cn_nn_blocks_test.py
index 324d3c4bb66..afedebd1d61 100644
--- a/official/vision/beta/projects/centernet/modeling/layers/cn_nn_blocks_test.py
+++ b/official/projects/centernet/modeling/layers/cn_nn_blocks_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,13 +19,13 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.modeling.layers import cn_nn_blocks
+from official.projects.centernet.modeling.layers import cn_nn_blocks
from official.vision.modeling.layers import nn_blocks
-class HourglassBlockPyTorch(tf.keras.layers.Layer):
+class HourglassBlockPyTorch(tf_keras.layers.Layer):
"""An CornerNet-style implementation of the hourglass block."""
def __init__(self, dims, modules, k=0, **kwargs):
@@ -63,7 +63,7 @@ def build(self, input_shape):
next_dim = dims[k + 1]
self.up1 = self.make_up_layer(3, curr_dim, curr_dim, curr_mod, **kwargs)
- self.max1 = tf.keras.layers.MaxPool2D(strides=2)
+ self.max1 = tf_keras.layers.MaxPool2D(strides=2)
self.low1 = self.make_hg_layer(3, curr_dim, next_dim, curr_mod, **kwargs)
if self.n - k > 1:
self.low2 = type(self)(dims, modules, k=k + 1, **kwargs)
@@ -72,8 +72,8 @@ def build(self, input_shape):
3, next_dim, next_dim, next_mod, **kwargs)
self.low3 = self.make_hg_layer_revr(
3, next_dim, curr_dim, curr_mod, **kwargs)
- self.up2 = tf.keras.layers.UpSampling2D(2)
- self.merge = tf.keras.layers.Add()
+ self.up2 = tf_keras.layers.UpSampling2D(2)
+ self.merge = tf_keras.layers.Add()
super(HourglassBlockPyTorch, self).build(input_shape)
@@ -91,7 +91,7 @@ def make_layer(self, k, inp_dim, out_dim, modules, **kwargs):
nn_blocks.ResidualBlock(out_dim, 1, use_projection=True, **kwargs)]
for _ in range(1, modules):
layers.append(nn_blocks.ResidualBlock(out_dim, 1, **kwargs))
- return tf.keras.Sequential(layers)
+ return tf_keras.Sequential(layers)
def make_layer_revr(self, k, inp_dim, out_dim, modules, **kwargs):
layers = []
@@ -100,7 +100,7 @@ def make_layer_revr(self, k, inp_dim, out_dim, modules, **kwargs):
nn_blocks.ResidualBlock(inp_dim, 1, **kwargs))
layers.append(
nn_blocks.ResidualBlock(out_dim, 1, use_projection=True, **kwargs))
- return tf.keras.Sequential(layers)
+ return tf_keras.Sequential(layers)
def make_up_layer(self, k, inp_dim, out_dim, modules, **kwargs):
return self.make_layer(k, inp_dim, out_dim, modules, **kwargs)
@@ -121,7 +121,7 @@ def test_hourglass_block(self):
dims = [256, 256, 384, 384, 384, 512]
modules = [2, 2, 2, 2, 2, 4]
model = cn_nn_blocks.HourglassBlock(dims, modules)
- test_input = tf.keras.Input((512, 512, 256))
+ test_input = tf_keras.Input((512, 512, 256))
_ = model(test_input)
filter_sizes = [256, 256, 384, 384, 384, 512]
diff --git a/official/vision/beta/projects/centernet/modeling/layers/detection_generator.py b/official/projects/centernet/modeling/layers/detection_generator.py
similarity index 89%
rename from official/vision/beta/projects/centernet/modeling/layers/detection_generator.py
rename to official/projects/centernet/modeling/layers/detection_generator.py
index e1686ed7bc7..c53cc1d1d9f 100644
--- a/official/vision/beta/projects/centernet/modeling/layers/detection_generator.py
+++ b/official/projects/centernet/modeling/layers/detection_generator.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,18 +23,18 @@
from typing import Any, Mapping
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.ops import loss_ops
-from official.vision.beta.projects.centernet.ops import nms_ops
+from official.projects.centernet.ops import loss_ops
+from official.projects.centernet.ops import nms_ops
from official.vision.ops import box_ops
-class CenterNetDetectionGenerator(tf.keras.layers.Layer):
+class CenterNetDetectionGenerator(tf_keras.layers.Layer):
"""CenterNet Detection Generator."""
def __init__(self,
- input_image_dims: int = 512,
+ input_image_dims: tuple[int, int] | int = 512,
net_down_scale: int = 4,
max_detections: int = 100,
peak_error: float = 1e-6,
@@ -47,7 +47,10 @@ def __init__(self,
"""Initialize CenterNet Detection Generator.
Args:
- input_image_dims: An `int` that specifies the input image size.
+ input_image_dims: The input image size. If it is a tuple of two `int`s, it
+ is the size (height, width) of the input images. If it is an `int`, the
+ input images are supposed to be squared images whose height and width
+ are equal.
net_down_scale: An `int` that specifies stride of the output.
max_detections: An `int` specifying the maximum number of bounding
boxes generated. This is an upper bound, so the number of generated
@@ -67,6 +70,9 @@ def __init__(self,
"""
super(CenterNetDetectionGenerator, self).__init__(**kwargs)
+ if isinstance(input_image_dims, int):
+ input_image_dims = (input_image_dims, input_image_dims)
+
# Object center selection parameters
self._max_detections = max_detections
self._peak_error = peak_error
@@ -246,10 +252,28 @@ def get_boxes(self,
return boxes, detection_classes
def convert_strided_predictions_to_normalized_boxes(self, boxes: tf.Tensor):
+ """Converts strided predictions to normalized boxes.
+
+ Args:
+ boxes: A tf.Tensor of shape [batch_size, num_predictions, 4], representing
+ the strided predictions of the detected objects.
+
+ Returns:
+ A tf.Tensor of shape [batch_size, num_predictions, 4], representing
+ the normalized boxes of the detected objects.
+ """
boxes = boxes * tf.cast(self._net_down_scale, boxes.dtype)
- boxes = boxes / tf.cast(self._input_image_dims, boxes.dtype)
- boxes = tf.clip_by_value(boxes, 0.0, 1.0)
- return boxes
+
+ height = tf.cast(self._input_image_dims[0], boxes.dtype)
+ width = tf.cast(self._input_image_dims[1], boxes.dtype)
+ ymin = boxes[..., 0:1] / height
+ xmin = boxes[..., 1:2] / width
+ ymax = boxes[..., 2:3] / height
+ xmax = boxes[..., 3:4] / width
+
+ normalized_boxes = tf.concat([ymin, xmin, ymax, xmax], axis=-1)
+ normalized_boxes = tf.clip_by_value(normalized_boxes, 0.0, 1.0)
+ return normalized_boxes
def __call__(self, inputs):
# Get heatmaps from decoded outputs via final hourglass stack output
@@ -308,8 +332,7 @@ def __call__(self, inputs):
nms_thresh=0.4)
num_det = tf.reduce_sum(tf.cast(scores > 0, dtype=tf.int32), axis=1)
- boxes = box_ops.denormalize_boxes(
- boxes, [self._input_image_dims, self._input_image_dims])
+ boxes = box_ops.denormalize_boxes(boxes, self._input_image_dims)
return {
'boxes': boxes,
diff --git a/official/projects/centernet/modeling/layers/detection_generator_test.py b/official/projects/centernet/modeling/layers/detection_generator_test.py
new file mode 100644
index 00000000000..cac5cae7be1
--- /dev/null
+++ b/official/projects/centernet/modeling/layers/detection_generator_test.py
@@ -0,0 +1,152 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for Centernet detection_generator."""
+
+from collections.abc import Mapping, Sequence
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.centernet.modeling.layers import detection_generator
+
+
+def _build_input_example(
+ batch_size: int, height: int, width: int, num_classes: int, num_outputs: int
+) -> Mapping[str, Sequence[tf.Tensor]]:
+ """Builds a random input example for CenterNetDetectionGenerator.
+
+ Args:
+ batch_size: The batch size.
+ height: The height of the feature_map.
+ width: The width of the feature_map.
+ num_classes: The number of classes to detect.
+ num_outputs: The number of output heatmaps, which corresponds to the length
+ of CenterNetHead's input_levels.
+
+ Returns:
+ A dictionary, mapping from feature names to sequences of tensors.
+ """
+ return {
+ 'ct_heatmaps': [
+ tf.random.normal([batch_size, height, width, num_classes])
+ for _ in range(num_outputs)
+ ],
+ 'ct_size': [
+ tf.random.normal([batch_size, height, width, 2])
+ for _ in range(num_outputs)
+ ],
+ 'ct_offset': [
+ tf.random.normal([batch_size, height, width, 2])
+ for _ in range(num_outputs)
+ ],
+ }
+
+
+class CenterNetDetectionGeneratorTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (1, 256),
+ (1, 512),
+ (2, 256),
+ (2, 512),
+ )
+ def test_squered_image_forward(self, batch_size, input_image_dims):
+ max_detections = 128
+ num_classes = 80
+ generator = detection_generator.CenterNetDetectionGenerator(
+ input_image_dims=input_image_dims, max_detections=max_detections
+ )
+ test_input = _build_input_example(
+ batch_size=batch_size,
+ height=input_image_dims,
+ width=input_image_dims,
+ num_classes=num_classes,
+ num_outputs=2,
+ )
+
+ output = generator(test_input)
+
+ self.assert_detection_generator_output_shapes(
+ output, batch_size, max_detections
+ )
+
+ @parameterized.parameters(
+ (1, (256, 512)),
+ (1, (512, 256)),
+ (2, (256, 512)),
+ (2, (512, 256)),
+ )
+ def test_rectangular_image_forward(self, batch_size, input_image_dims):
+ max_detections = 128
+ num_classes = 80
+ generator = detection_generator.CenterNetDetectionGenerator(
+ input_image_dims=input_image_dims, max_detections=max_detections
+ )
+ test_input = _build_input_example(
+ batch_size=batch_size,
+ height=input_image_dims[0],
+ width=input_image_dims[1],
+ num_classes=num_classes,
+ num_outputs=2,
+ )
+
+ output = generator(test_input)
+
+ self.assert_detection_generator_output_shapes(
+ output, batch_size, max_detections
+ )
+
+ def assert_detection_generator_output_shapes(
+ self,
+ output: Mapping[str, tf.Tensor],
+ batch_size: int,
+ max_detections: int,
+ ):
+ self.assertAllEqual(output['boxes'].shape, (batch_size, max_detections, 4))
+ self.assertAllEqual(output['classes'].shape, (batch_size, max_detections))
+ self.assertAllEqual(
+ output['confidence'].shape, (batch_size, max_detections)
+ )
+ self.assertAllEqual(output['num_detections'].shape, (batch_size,))
+
+ @parameterized.parameters(
+ (256,),
+ (512,),
+ ((256, 512),),
+ ((512, 256),),
+ )
+ def test_serialize_deserialize(self, input_image_dims):
+ kwargs = {
+ 'input_image_dims': input_image_dims,
+ 'net_down_scale': 4,
+ 'max_detections': 128,
+ 'peak_error': 1e-6,
+ 'peak_extract_kernel_size': 3,
+ 'class_offset': 1,
+ 'use_nms': False,
+ 'nms_pre_thresh': 0.1,
+ 'nms_thresh': 0.5,
+ }
+
+ generator = detection_generator.CenterNetDetectionGenerator(**kwargs)
+ new_generator = detection_generator.CenterNetDetectionGenerator.from_config(
+ generator.get_config()
+ )
+
+ self.assertAllEqual(generator.get_config(), new_generator.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/centernet/ops/__init__.py b/official/projects/centernet/ops/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/ops/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/centernet/ops/box_list.py b/official/projects/centernet/ops/box_list.py
similarity index 98%
rename from official/vision/beta/projects/centernet/ops/box_list.py
rename to official/projects/centernet/ops/box_list.py
index 4e93b9fd631..6d5d1e34e18 100644
--- a/official/vision/beta/projects/centernet/ops/box_list.py
+++ b/official/projects/centernet/ops/box_list.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -33,7 +33,7 @@
* Tensors are always provided as (flat) [N, 4] tensors.
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def _get_dim_as_int(dim):
diff --git a/official/vision/beta/projects/centernet/ops/box_list_ops.py b/official/projects/centernet/ops/box_list_ops.py
similarity index 98%
rename from official/vision/beta/projects/centernet/ops/box_list_ops.py
rename to official/projects/centernet/ops/box_list_ops.py
index c811371f0ca..be79ac57e62 100644
--- a/official/vision/beta/projects/centernet/ops/box_list_ops.py
+++ b/official/projects/centernet/ops/box_list_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,9 @@
"""Bounding Box List operations."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.ops import box_list
+from official.projects.centernet.ops import box_list
from official.vision.ops import sampling_ops
diff --git a/official/vision/beta/projects/centernet/ops/loss_ops.py b/official/projects/centernet/ops/loss_ops.py
similarity index 98%
rename from official/vision/beta/projects/centernet/ops/loss_ops.py
rename to official/projects/centernet/ops/loss_ops.py
index db7875c110e..dc64c2e0833 100644
--- a/official/vision/beta/projects/centernet/ops/loss_ops.py
+++ b/official/projects/centernet/ops/loss_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Operations for compute losses for centernet."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.ops import sampling_ops
diff --git a/official/vision/beta/projects/centernet/ops/nms_ops.py b/official/projects/centernet/ops/nms_ops.py
similarity index 96%
rename from official/vision/beta/projects/centernet/ops/nms_ops.py
rename to official/projects/centernet/ops/nms_ops.py
index adc0891d274..bc91e9d509a 100644
--- a/official/vision/beta/projects/centernet/ops/nms_ops.py
+++ b/official/projects/centernet/ops/nms_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,9 @@
"""nms computation."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import box_ops
+from official.projects.yolo.ops import box_ops
NMS_TILE_SIZE = 512
diff --git a/official/vision/beta/projects/centernet/ops/preprocess_ops.py b/official/projects/centernet/ops/preprocess_ops.py
similarity index 98%
rename from official/vision/beta/projects/centernet/ops/preprocess_ops.py
rename to official/projects/centernet/ops/preprocess_ops.py
index c7fba31461e..5b81e9cd126 100644
--- a/official/vision/beta/projects/centernet/ops/preprocess_ops.py
+++ b/official/projects/centernet/ops/preprocess_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,10 +16,10 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.ops import box_list
-from official.vision.beta.projects.centernet.ops import box_list_ops
+from official.projects.centernet.ops import box_list
+from official.projects.centernet.ops import box_list_ops
def _get_or_create_preprocess_rand_vars(generator_func,
diff --git a/official/vision/beta/projects/centernet/ops/target_assigner.py b/official/projects/centernet/ops/target_assigner.py
similarity index 99%
rename from official/vision/beta/projects/centernet/ops/target_assigner.py
rename to official/projects/centernet/ops/target_assigner.py
index 0bbe39dffe6..37a01ff3d6e 100644
--- a/official/vision/beta/projects/centernet/ops/target_assigner.py
+++ b/official/projects/centernet/ops/target_assigner.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Dict, List
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.ops import sampling_ops
diff --git a/official/vision/beta/projects/centernet/ops/target_assigner_test.py b/official/projects/centernet/ops/target_assigner_test.py
similarity index 97%
rename from official/vision/beta/projects/centernet/ops/target_assigner_test.py
rename to official/projects/centernet/ops/target_assigner_test.py
index 0886ae66088..1499779d8b7 100644
--- a/official/vision/beta/projects/centernet/ops/target_assigner_test.py
+++ b/official/projects/centernet/ops/target_assigner_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
"""Tests for targets generations of centernet."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.centernet.ops import target_assigner
+from official.projects.centernet.ops import target_assigner
from official.vision.ops import preprocess_ops
diff --git a/official/projects/centernet/tasks/__init__.py b/official/projects/centernet/tasks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/tasks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/centernet/tasks/centernet.py b/official/projects/centernet/tasks/centernet.py
similarity index 92%
rename from official/vision/beta/projects/centernet/tasks/centernet.py
rename to official/projects/centernet/tasks/centernet.py
index 4230c2059b9..d872810e30f 100644
--- a/official/vision/beta/projects/centernet/tasks/centernet.py
+++ b/official/projects/centernet/tasks/centernet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,19 +17,19 @@
from typing import Any, List, Optional, Tuple
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import input_reader
from official.core import task_factory
-from official.vision.beta.projects.centernet.configs import centernet as exp_cfg
-from official.vision.beta.projects.centernet.dataloaders import centernet_input
-from official.vision.beta.projects.centernet.losses import centernet_losses
-from official.vision.beta.projects.centernet.modeling import centernet_model
-from official.vision.beta.projects.centernet.modeling.heads import centernet_head
-from official.vision.beta.projects.centernet.modeling.layers import detection_generator
-from official.vision.beta.projects.centernet.ops import loss_ops
-from official.vision.beta.projects.centernet.ops import target_assigner
+from official.projects.centernet.configs import centernet as exp_cfg
+from official.projects.centernet.dataloaders import centernet_input
+from official.projects.centernet.losses import centernet_losses
+from official.projects.centernet.modeling import centernet_model
+from official.projects.centernet.modeling.heads import centernet_head
+from official.projects.centernet.modeling.layers import detection_generator
+from official.projects.centernet.ops import loss_ops
+from official.projects.centernet.ops import target_assigner
from official.vision.dataloaders import tf_example_decoder
from official.vision.dataloaders import tfds_factory
from official.vision.dataloaders import tf_example_label_map_decoder
@@ -90,14 +90,14 @@ def build_inputs(self,
def build_model(self):
"""get an instance of CenterNet."""
model_config = self.task_config.model
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None] + model_config.input_size)
l2_weight_decay = self.task_config.weight_decay
# Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
- l2_regularizer = (tf.keras.regularizers.l2(
+ l2_regularizer = (tf_keras.regularizers.l2(
l2_weight_decay / 2.0) if l2_weight_decay else None)
backbone = factory.build_backbone(
@@ -130,7 +130,10 @@ def build_model(self):
peak_extract_kernel_size=dg_config.peak_extract_kernel_size,
class_offset=dg_config.class_offset,
net_down_scale=self._net_down_scale,
- input_image_dims=model_config.input_size[0],
+ input_image_dims=(
+ model_config.input_size[0],
+ model_config.input_size[1],
+ ),
use_nms=dg_config.use_nms,
nms_pre_thresh=dg_config.nms_pre_thresh,
nms_thresh=dg_config.nms_thresh)
@@ -142,7 +145,7 @@ def build_model(self):
return model
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loading pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -302,7 +305,7 @@ def build_metrics(self, training=True):
metrics = []
metric_names = ['total_loss', 'ct_loss', 'scale_loss', 'ct_offset_loss']
for name in metric_names:
- metrics.append(tf.keras.metrics.Mean(name, dtype=tf.float32))
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
if not training:
if (self.task_config.validation_data.tfds_name
@@ -318,8 +321,8 @@ def build_metrics(self, training=True):
def train_step(self,
inputs: Tuple[Any, Any],
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
metrics: Optional[List[Any]] = None):
"""Does forward and backward.
@@ -346,7 +349,7 @@ def train_step(self,
scaled_loss = losses['total_loss'] / num_replicas
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
# compute the gradient
@@ -354,7 +357,7 @@ def train_step(self,
gradients = tape.gradient(scaled_loss, tvars)
# get unscaled loss if the scaled loss was used
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
gradients = optimizer.get_unscaled_gradients(gradients)
if self.task_config.gradient_clip_norm > 0.0:
@@ -374,7 +377,7 @@ def train_step(self,
def validation_step(self,
inputs: Tuple[Any, Any],
- model: tf.keras.Model,
+ model: tf_keras.Model,
metrics: Optional[List[Any]] = None):
"""Validation step.
@@ -417,8 +420,8 @@ def aggregate_logs(self, state=None, step_outputs=None):
if state is None:
self.coco_metric.reset_states()
state = self.coco_metric
- self.coco_metric.update_state(step_outputs[self.coco_metric.name][0],
- step_outputs[self.coco_metric.name][1])
+ self.coco_metric.update_state(step_outputs[self.coco_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ step_outputs[self.coco_metric.name][1]) # pyrefly: ignore[unsupported-operation]
return state
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
diff --git a/official/vision/beta/projects/centernet/train.py b/official/projects/centernet/train.py
similarity index 93%
rename from official/vision/beta/projects/centernet/train.py
rename to official/projects/centernet/train.py
index 8488e6084a1..7adb14bb293 100644
--- a/official/vision/beta/projects/centernet/train.py
+++ b/official/projects/centernet/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from official.core import train_lib
from official.core import train_utils
from official.modeling import performance
-from official.vision.beta.projects.centernet.common import registry_imports # pylint: disable=unused-import
+from official.projects.centernet.common import registry_imports # pylint: disable=unused-import
FLAGS = flags.FLAGS
diff --git a/official/projects/centernet/utils/__init__.py b/official/projects/centernet/utils/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/utils/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/centernet/utils/checkpoints/__init__.py b/official/projects/centernet/utils/checkpoints/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/centernet/utils/checkpoints/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/centernet/utils/checkpoints/config_classes.py b/official/projects/centernet/utils/checkpoints/config_classes.py
similarity index 86%
rename from official/vision/beta/projects/centernet/utils/checkpoints/config_classes.py
rename to official/projects/centernet/utils/checkpoints/config_classes.py
index 5c67085f9f6..bcd0708c62a 100644
--- a/official/vision/beta/projects/centernet/utils/checkpoints/config_classes.py
+++ b/official/projects/centernet/utils/checkpoints/config_classes.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,7 +26,7 @@
from typing import Dict, Optional
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class Config(abc.ABC):
@@ -36,7 +36,7 @@ def get_weights(self):
"""Generates the weights needed to be loaded into the layer."""
raise NotImplementedError
- def load_weights(self, layer: tf.keras.layers.Layer) -> int:
+ def load_weights(self, layer: tf_keras.layers.Layer) -> int:
"""Assign weights to layer.
Given a layer, this function retrieves the weights for that layer in an
@@ -47,7 +47,7 @@ def load_weights(self, layer: tf.keras.layers.Layer) -> int:
will be raised by set_weights().
Args:
- layer: A `tf.keras.layers.Layer`.
+ layer: A `tf_keras.layers.Layer`.
Returns:
@@ -77,8 +77,8 @@ class Conv2DBNCFG(Config):
repr=False, default=None)
def __post_init__(self):
- conv_weights_dict = self.weights_dict['conv']
- norm_weights_dict = self.weights_dict['norm']
+ conv_weights_dict = self.weights_dict['conv'] # pyrefly: ignore[unsupported-operation]
+ norm_weights_dict = self.weights_dict['norm'] # pyrefly: ignore[unsupported-operation]
self.weights = conv_weights_dict['kernel']
@@ -138,12 +138,12 @@ class ResidualBlockCFG(Config):
repr=False, default=None)
def __post_init__(self):
- conv_weights_dict = self.weights_dict['conv']
- norm_weights_dict = self.weights_dict['norm']
- conv_block_weights_dict = self.weights_dict['conv_block']
+ conv_weights_dict = self.weights_dict['conv'] # pyrefly: ignore[unsupported-operation]
+ norm_weights_dict = self.weights_dict['norm'] # pyrefly: ignore[unsupported-operation]
+ conv_block_weights_dict = self.weights_dict['conv_block'] # pyrefly: ignore[unsupported-operation]
- if 'skip' in self.weights_dict:
- skip_weights_dict = self.weights_dict['skip']
+ if 'skip' in self.weights_dict: # pyrefly: ignore[not-iterable]
+ skip_weights_dict = self.weights_dict['skip'] # pyrefly: ignore[unsupported-operation]
self.skip_weights = skip_weights_dict['conv']['kernel']
self.skip_beta = skip_weights_dict['norm']['beta']
self.skip_gamma = skip_weights_dict['norm']['gamma']
@@ -207,8 +207,8 @@ class HeadConvCFG(Config):
repr=False, default=None)
def __post_init__(self):
- conv_1_weights_dict = self.weights_dict['layer_with_weights-0']
- conv_2_weights_dict = self.weights_dict['layer_with_weights-1']
+ conv_1_weights_dict = self.weights_dict['layer_with_weights-0'] # pyrefly: ignore[unsupported-operation]
+ conv_2_weights_dict = self.weights_dict['layer_with_weights-1'] # pyrefly: ignore[unsupported-operation]
self.conv_1_weights = conv_1_weights_dict['kernel']
self.conv_1_bias = conv_1_weights_dict['bias']
@@ -230,10 +230,10 @@ class HourglassCFG(Config):
weights_dict: Optional[Dict[str, np.ndarray]] = dataclasses.field(
repr=False, default=None)
- is_last_stage: bool = dataclasses.field(repr=False, default=None)
+ is_last_stage: bool = dataclasses.field(repr=False, default=None) # pyrefly: ignore[bad-assignment]
def __post_init__(self):
- self.is_last_stage = False if 'inner_block' in self.weights_dict else True
+ self.is_last_stage = False if 'inner_block' in self.weights_dict else True # pyrefly: ignore[not-iterable]
def get_weights(self):
"""It is not used in this class."""
@@ -271,15 +271,15 @@ def load_weights(self, layer):
layer.submodules[3]
]
enc_dec_weight_dicts = [
- self.weights_dict['encoder_block1'],
- self.weights_dict['encoder_block2'],
- self.weights_dict['decoder_block']
+ self.weights_dict['encoder_block1'], # pyrefly: ignore[unsupported-operation]
+ self.weights_dict['encoder_block2'], # pyrefly: ignore[unsupported-operation]
+ self.weights_dict['decoder_block'] # pyrefly: ignore[unsupported-operation]
]
for l, weights_dict in zip(enc_dec_layers, enc_dec_weight_dicts):
n_weights += self.load_block_weights(l, weights_dict)
- if len(self.weights_dict['inner_block']) == 1:
+ if len(self.weights_dict['inner_block']) == 1: # pyrefly: ignore[unsupported-operation]
# still in an outer hourglass
inner_weights_dict = self.weights_dict['inner_block']['0']
else:
@@ -287,7 +287,7 @@ def load_weights(self, layer):
inner_weights_dict = self.weights_dict['inner_block']
inner_hg_layer = layer.submodules[2]
- inner_hg_cfg = type(self)(weights_dict=inner_weights_dict)
+ inner_hg_cfg = type(self)(weights_dict=inner_weights_dict) # pyrefly: ignore[bad-argument-type]
n_weights += inner_hg_cfg.load_weights(inner_hg_layer)
else:
diff --git a/official/vision/beta/projects/centernet/utils/checkpoints/config_data.py b/official/projects/centernet/utils/checkpoints/config_data.py
similarity index 73%
rename from official/vision/beta/projects/centernet/utils/checkpoints/config_data.py
rename to official/projects/centernet/utils/checkpoints/config_data.py
index 76de9dfd2d1..a603418e88f 100644
--- a/official/vision/beta/projects/centernet/utils/checkpoints/config_data.py
+++ b/official/projects/centernet/utils/checkpoints/config_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
import numpy as np
-from official.vision.beta.projects.centernet.utils.checkpoints import config_classes
+from official.projects.centernet.utils.checkpoints import config_classes
Conv2DBNCFG = config_classes.Conv2DBNCFG
HeadConvCFG = config_classes.HeadConvCFG
@@ -41,54 +41,54 @@ def get_cfg_list(self, name):
return [
# Downsampling Layers
Conv2DBNCFG(
- weights_dict=self.weights_dict['downsample_input']['conv_block']),
+ weights_dict=self.weights_dict['downsample_input']['conv_block']), # pyrefly: ignore[unsupported-operation]
ResidualBlockCFG(
- weights_dict=self.weights_dict['downsample_input'][
+ weights_dict=self.weights_dict['downsample_input'][ # pyrefly: ignore[unsupported-operation]
'residual_block']),
# Hourglass
HourglassCFG(
- weights_dict=self.weights_dict['hourglass_network']['0']),
+ weights_dict=self.weights_dict['hourglass_network']['0']), # pyrefly: ignore[unsupported-operation]
Conv2DBNCFG(
- weights_dict=self.weights_dict['output_conv']['0']),
+ weights_dict=self.weights_dict['output_conv']['0']), # pyrefly: ignore[unsupported-operation]
# Intermediate
Conv2DBNCFG(
- weights_dict=self.weights_dict['intermediate_conv1']['0']),
+ weights_dict=self.weights_dict['intermediate_conv1']['0']), # pyrefly: ignore[unsupported-operation]
Conv2DBNCFG(
- weights_dict=self.weights_dict['intermediate_conv2']['0']),
+ weights_dict=self.weights_dict['intermediate_conv2']['0']), # pyrefly: ignore[unsupported-operation]
ResidualBlockCFG(
- weights_dict=self.weights_dict['intermediate_residual']['0']),
+ weights_dict=self.weights_dict['intermediate_residual']['0']), # pyrefly: ignore[unsupported-operation]
# Hourglass
HourglassCFG(
- weights_dict=self.weights_dict['hourglass_network']['1']),
+ weights_dict=self.weights_dict['hourglass_network']['1']), # pyrefly: ignore[unsupported-operation]
Conv2DBNCFG(
- weights_dict=self.weights_dict['output_conv']['1']),
+ weights_dict=self.weights_dict['output_conv']['1']), # pyrefly: ignore[unsupported-operation]
]
elif name == 'extremenet':
return [
# Downsampling Layers
Conv2DBNCFG(
- weights_dict=self.weights_dict['downsample_input']['conv_block']),
+ weights_dict=self.weights_dict['downsample_input']['conv_block']), # pyrefly: ignore[unsupported-operation]
ResidualBlockCFG(
- weights_dict=self.weights_dict['downsample_input'][
+ weights_dict=self.weights_dict['downsample_input'][ # pyrefly: ignore[unsupported-operation]
'residual_block']),
# Hourglass
HourglassCFG(
- weights_dict=self.weights_dict['hourglass_network']['0']),
+ weights_dict=self.weights_dict['hourglass_network']['0']), # pyrefly: ignore[unsupported-operation]
Conv2DBNCFG(
- weights_dict=self.weights_dict['output_conv']['0']),
+ weights_dict=self.weights_dict['output_conv']['0']), # pyrefly: ignore[unsupported-operation]
# Intermediate
Conv2DBNCFG(
- weights_dict=self.weights_dict['intermediate_conv1']['0']),
+ weights_dict=self.weights_dict['intermediate_conv1']['0']), # pyrefly: ignore[unsupported-operation]
Conv2DBNCFG(
- weights_dict=self.weights_dict['intermediate_conv2']['0']),
+ weights_dict=self.weights_dict['intermediate_conv2']['0']), # pyrefly: ignore[unsupported-operation]
ResidualBlockCFG(
- weights_dict=self.weights_dict['intermediate_residual']['0']),
+ weights_dict=self.weights_dict['intermediate_residual']['0']), # pyrefly: ignore[unsupported-operation]
# Hourglass
HourglassCFG(
- weights_dict=self.weights_dict['hourglass_network']['1']),
+ weights_dict=self.weights_dict['hourglass_network']['1']), # pyrefly: ignore[unsupported-operation]
Conv2DBNCFG(
- weights_dict=self.weights_dict['output_conv']['1']),
+ weights_dict=self.weights_dict['output_conv']['1']), # pyrefly: ignore[unsupported-operation]
]
@@ -102,10 +102,10 @@ class HeadConfigData:
def get_cfg_list(self, name):
if name == 'detection_2d':
return [
- HeadConvCFG(weights_dict=self.weights_dict['object_center']['0']),
- HeadConvCFG(weights_dict=self.weights_dict['object_center']['1']),
- HeadConvCFG(weights_dict=self.weights_dict['box.Soffset']['0']),
- HeadConvCFG(weights_dict=self.weights_dict['box.Soffset']['1']),
- HeadConvCFG(weights_dict=self.weights_dict['box.Sscale']['0']),
- HeadConvCFG(weights_dict=self.weights_dict['box.Sscale']['1'])
+ HeadConvCFG(weights_dict=self.weights_dict['object_center']['0']), # pyrefly: ignore[unsupported-operation]
+ HeadConvCFG(weights_dict=self.weights_dict['object_center']['1']), # pyrefly: ignore[unsupported-operation]
+ HeadConvCFG(weights_dict=self.weights_dict['box.Soffset']['0']), # pyrefly: ignore[unsupported-operation]
+ HeadConvCFG(weights_dict=self.weights_dict['box.Soffset']['1']), # pyrefly: ignore[unsupported-operation]
+ HeadConvCFG(weights_dict=self.weights_dict['box.Sscale']['0']), # pyrefly: ignore[unsupported-operation]
+ HeadConvCFG(weights_dict=self.weights_dict['box.Sscale']['1']) # pyrefly: ignore[unsupported-operation]
]
diff --git a/official/vision/beta/projects/centernet/utils/checkpoints/load_weights.py b/official/projects/centernet/utils/checkpoints/load_weights.py
similarity index 94%
rename from official/vision/beta/projects/centernet/utils/checkpoints/load_weights.py
rename to official/projects/centernet/utils/checkpoints/load_weights.py
index cddd4b868b8..5ad057149fa 100644
--- a/official/vision/beta/projects/centernet/utils/checkpoints/load_weights.py
+++ b/official/projects/centernet/utils/checkpoints/load_weights.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,9 @@
"""Functions used to load the ODAPI CenterNet checkpoint."""
-from official.vision.beta.projects.centernet.modeling.layers import cn_nn_blocks
-from official.vision.beta.projects.centernet.utils.checkpoints import config_classes
-from official.vision.beta.projects.centernet.utils.checkpoints import config_data
+from official.projects.centernet.modeling.layers import cn_nn_blocks
+from official.projects.centernet.utils.checkpoints import config_classes
+from official.projects.centernet.utils.checkpoints import config_data
from official.vision.modeling.backbones import mobilenet
from official.vision.modeling.layers import nn_blocks
diff --git a/official/vision/beta/projects/centernet/utils/checkpoints/read_checkpoints.py b/official/projects/centernet/utils/checkpoints/read_checkpoints.py
similarity index 97%
rename from official/vision/beta/projects/centernet/utils/checkpoints/read_checkpoints.py
rename to official/projects/centernet/utils/checkpoints/read_checkpoints.py
index 4128f404600..60061dcf6d7 100644
--- a/official/vision/beta/projects/centernet/utils/checkpoints/read_checkpoints.py
+++ b/official/projects/centernet/utils/checkpoints/read_checkpoints.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Functions used to convert a TF checkpoint into a dictionary."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def update_weights_dict(weights_dict, variable_key, value):
diff --git a/official/vision/beta/projects/centernet/utils/tf2_centernet_checkpoint_converter.py b/official/projects/centernet/utils/tf2_centernet_checkpoint_converter.py
similarity index 84%
rename from official/vision/beta/projects/centernet/utils/tf2_centernet_checkpoint_converter.py
rename to official/projects/centernet/utils/tf2_centernet_checkpoint_converter.py
index 576a2cbd9a6..5c8d59390b8 100644
--- a/official/vision/beta/projects/centernet/utils/tf2_centernet_checkpoint_converter.py
+++ b/official/projects/centernet/utils/tf2_centernet_checkpoint_converter.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,16 +17,16 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
-
-from official.vision.beta.projects.centernet.common import registry_imports # pylint: disable=unused-import
-from official.vision.beta.projects.centernet.configs import backbones
-from official.vision.beta.projects.centernet.configs import centernet
-from official.vision.beta.projects.centernet.modeling import centernet_model
-from official.vision.beta.projects.centernet.modeling.heads import centernet_head
-from official.vision.beta.projects.centernet.modeling.layers import detection_generator
-from official.vision.beta.projects.centernet.utils.checkpoints import load_weights
-from official.vision.beta.projects.centernet.utils.checkpoints import read_checkpoints
+import tensorflow as tf, tf_keras
+
+from official.projects.centernet.common import registry_imports # pylint: disable=unused-import
+from official.projects.centernet.configs import backbones
+from official.projects.centernet.configs import centernet
+from official.projects.centernet.modeling import centernet_model
+from official.projects.centernet.modeling.heads import centernet_head
+from official.projects.centernet.modeling.layers import detection_generator
+from official.projects.centernet.utils.checkpoints import load_weights
+from official.projects.centernet.utils.checkpoints import read_checkpoints
from official.vision.modeling.backbones import factory
FLAGS = flags.FLAGS
@@ -58,7 +58,7 @@ def _create_centernet_model(model_id: int = 52,
model_config = task_config.model
backbone = factory.build_backbone(
- input_specs=tf.keras.layers.InputSpec(shape=[1, 512, 512, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[1, 512, 512, 3]),
backbone_config=model_config.backbone,
norm_activation_config=model_config.norm_activation)
diff --git a/official/projects/const_cl/README.md b/official/projects/const_cl/README.md
index 488b455fda9..b8dce2ba825 100644
--- a/official/projects/const_cl/README.md
+++ b/official/projects/const_cl/README.md
@@ -1,5 +1,60 @@
# Contextualized Spatial-Temporal Contrastive Learning with Self-Supervision
-(WIP) This repository contains the official implementation of
+[](https://arxiv.org/abs/2112.05181)
+
+
+This repository contains the official implementation of
[Contextualized Spatio-Temporal Contrastive Learning with Self-Supervision](https://arxiv.org/abs/2112.05181)
in TF2.
+
+
+
+
+
+## Description
+
+Most of existing video-language pre-training methods focus on instance-level
+alignment between video clips and captions via global contrastive learning but
+neglect rich fine-grained local information, which is of importance to
+downstream tasks requiring temporal localization and semantic reasoning. In this
+work, we propose a simple yet effective video-language pre-training framework,
+namely G-ViLM, to learn discriminative spatiotemporal features. Two novel
+designs involving spatiotemporal grounding and temporal grouping promote
+learning local region-noun alignment and temporal-aware features simultaneously.
+Specifically, spatiotemporal grounding aggregates semantically similar video
+tokens and aligns them with noun phrases extracted from the caption to promote
+local region-noun correspondences. Moreover, temporal grouping leverages
+cut-and-paste to manually create temporal scene changes and then learns
+distinguishable features from different scenes. Comprehensive evaluations
+demonstrate that G-ViLM performs favorably against existing approaches on four
+representative downstream tasks, covering text-video retrieval, video question
+answering, video action recognition and temporal action localization. G-ViLM
+performs competitively on all evaluated tasks and in particular achieves R@10 of
+65.1 on zero-shot MSR-VTT retrieval, over 9% higher than the state-of-the-art
+method.
+
+## Pre-trained Model Performance
+
+All models are pre-trained from scratch with `region_generator = RANDOM` and `context_length = 5` as described in the paper.
+
+We report the mean average
+precision on AVA v2.2 and AVA-Kinetics validation set and precision/success rate
+on Object Tracking Benchmark 2015.
+
+| Method | Parameters | Dataset | Pretrain Steps | AVA(mAP) | AVAK(mAP) | OTB(P/S) |
+| :--------------: | :----: | :--: | :--: |:----: |:-----------: | :----------: |
+| CVRL | 31.7M | Kinetics-400 | 200k | 18.4% | 24.1% | 75.4/53.7 |
+| ConST-CL | 31.7M | Kinetics-400 | 100k | 22.1% | 28.0% | 77.4/54.3 |
+| ConST-CL | 31.7M | Kinetics-400 | 200k | 24.1% | 30.5% | 78.1/55.2 |
+
+
+## Citation
+
+```
+@inproceedings{yuan2022constcl,
+ title={Contexualized Spatio-Temporal Contrastive Learning with Self-Supervision},
+ author={Yuan, Liangzhe and Qian, Rui and Cui, Yin and Gong, Boqing and Schroff, Florian and Yang, Ming-Hsuan and Adam, Hartwig and Liu, Ting},
+ journal={CVPR},
+ year={2022}
+}
+```
diff --git a/official/projects/const_cl/configs/backbones_3d.py b/official/projects/const_cl/configs/backbones_3d.py
new file mode 100644
index 00000000000..2bd8df0a050
--- /dev/null
+++ b/official/projects/const_cl/configs/backbones_3d.py
@@ -0,0 +1,59 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""3D Backbones configurations."""
+import dataclasses
+from typing import Tuple
+
+from official.vision.configs import backbones_3d
+
+
+ResNet3DBlock = backbones_3d.ResNet3DBlock
+
+
+@dataclasses.dataclass
+class ResNet3DY(backbones_3d.ResNet3D):
+ pass
+
+
+@dataclasses.dataclass
+class ResNet3DY50(ResNet3DY):
+ """Block specifications of the Resnet50 (3DY) model."""
+ model_id: int = 50
+ block_specs: Tuple[
+ ResNet3DBlock, ResNet3DBlock, ResNet3DBlock, ResNet3DBlock] = (
+ ResNet3DBlock(temporal_strides=1,
+ temporal_kernel_sizes=(3, 3, 3),
+ use_self_gating=True),
+ ResNet3DBlock(temporal_strides=1,
+ temporal_kernel_sizes=(3, 1, 3, 1),
+ use_self_gating=True),
+ ResNet3DBlock(temporal_strides=1,
+ temporal_kernel_sizes=(3, 1, 3, 1, 3, 1),
+ use_self_gating=True),
+ ResNet3DBlock(temporal_strides=1,
+ temporal_kernel_sizes=(1, 3, 1),
+ use_self_gating=True))
+
+
+@dataclasses.dataclass
+class Backbone3D(backbones_3d.Backbone3D):
+ """Configuration for backbones.
+
+ Attributes:
+ type: type of backbone be used, one of the fields below.
+ resnet_3dy: resnet_3dy backbone config.
+ """
+ type: str = 'resnet_3dy'
+ resnet_3dy: ResNet3DY = dataclasses.field(default_factory=ResNet3DY50)
diff --git a/official/projects/const_cl/configs/backbones_3d_test.py b/official/projects/const_cl/configs/backbones_3d_test.py
new file mode 100644
index 00000000000..cf584564e11
--- /dev/null
+++ b/official/projects/const_cl/configs/backbones_3d_test.py
@@ -0,0 +1,32 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for backbones_3d."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.configs import backbones_3d
+
+
+class Backbones3DTest(tf.test.TestCase):
+
+ def test_conv3dy_config(self):
+ config = backbones_3d.Backbone3D(
+ type='resnet_3dy',
+ resnet_3d=backbones_3d.ResNet3DY50())
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/configs/const_cl.py b/official/projects/const_cl/configs/const_cl.py
new file mode 100644
index 00000000000..85502ffbe35
--- /dev/null
+++ b/official/projects/const_cl/configs/const_cl.py
@@ -0,0 +1,107 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Video classification configuration definition."""
+
+import dataclasses
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.const_cl.configs import backbones_3d as backbones_3d_cfg
+from official.projects.const_cl.configs import head as head_cfg
+from official.vision.configs import common
+from official.vision.configs import video_classification
+
+
+VideoClassificationTask = video_classification.VideoClassificationTask
+
+
+@dataclasses.dataclass
+class ConstCLPretrainTask(VideoClassificationTask):
+ pass
+
+
+@dataclasses.dataclass
+class DataConfig(video_classification.DataConfig):
+ """The base configuration for building datasets."""
+ zero_centering_image: bool = True
+ is_ssl: bool = False
+ num_instances: int = 8
+
+
+@dataclasses.dataclass
+class ConstCLModel(hyperparams.Config):
+ """The model config."""
+ model_type: str = 'video_classification'
+ backbone: backbones_3d_cfg.Backbone3D = dataclasses.field(
+ default_factory=lambda: backbones_3d_cfg.Backbone3D( # pylint: disable=g-long-lambda
+ type='resnet_3dy', resnet_3dy=backbones_3d_cfg.ResNet3DY50()
+ )
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ use_sync_bn=False, norm_momentum=0.9, norm_epsilon=1e-5
+ )
+ )
+ global_head: head_cfg.MLP = dataclasses.field(
+ default_factory=lambda: head_cfg.MLP( # pylint: disable=g-long-lambda
+ use_sync_bn=False, normalize_inputs=False, norm_momentum=0.9
+ )
+ )
+ local_head: head_cfg.InstanceReconstructor = dataclasses.field(
+ default_factory=head_cfg.InstanceReconstructor
+ )
+
+
+@dataclasses.dataclass
+class ConstCLLosses(hyperparams.Config):
+ """The config for ConST-CL losses."""
+ normalize_inputs: bool = True
+ global_temperature: float = 0.1
+ local_temperature: float = 0.2
+ global_weight: float = 1.0
+ local_weight: float = 0.001
+ l2_weight_decay: float = 0.0
+
+
+@exp_factory.register_config_factory('const_cl_pretrain_kinetics400')
+def const_cl_pretrain_kinetics400() -> cfg.ExperimentConfig:
+ """Pretrains SSL Video classification on Kinectics 400 with ResNet."""
+ exp = video_classification.video_classification_kinetics400()
+ exp.task = ConstCLPretrainTask(**exp.task.as_dict())
+ exp.task.train_data.zero_centering_image = True
+ exp.task.train_data = DataConfig(is_ssl=True, **exp.task.train_data.as_dict())
+ exp.task.train_data.feature_shape = (16, 224, 224, 3)
+ exp.task.train_data.temporal_stride = 2
+ exp.task.train_data.aug_min_area_ratio = 0.3
+ exp.task.model = ConstCLModel()
+ exp.task.model.model_type = 'const_cl_model'
+ exp.task.losses = ConstCLLosses()
+ return exp
+
+
+@exp_factory.register_config_factory('const_cl_pretrain_kinetics600')
+def const_cl_pretrain_kinetics600() -> cfg.ExperimentConfig:
+ """Pretrains SSL Video classification on Kinectics 400 with ResNet."""
+ exp = video_classification.video_classification_kinetics600()
+ exp.task = ConstCLPretrainTask(**exp.task.as_dict())
+ exp.task.train_data.zero_centering_image = True
+ exp.task.train_data = DataConfig(is_ssl=True, **exp.task.train_data.as_dict())
+ exp.task.train_data.feature_shape = (16, 224, 224, 3)
+ exp.task.train_data.temporal_stride = 2
+ exp.task.train_data.aug_min_area_ratio = 0.3
+ exp.task.model = ConstCLModel()
+ exp.task.model.model_type = 'const_cl_model'
+ exp.task.losses = ConstCLLosses()
+ return exp
diff --git a/official/projects/const_cl/configs/const_cl_test.py b/official/projects/const_cl/configs/const_cl_test.py
new file mode 100644
index 00000000000..c097a5e8cfb
--- /dev/null
+++ b/official/projects/const_cl/configs/const_cl_test.py
@@ -0,0 +1,42 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# pylint: disable=unused-import
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official import vision
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.const_cl.configs import const_cl as exp_cfg
+
+
+class VideoClassificationConfigTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.parameters(('const_cl_pretrain_kinetics400',),
+ ('const_cl_pretrain_kinetics600',))
+ def test_const_cl_pretrain_configs(self, config_name):
+ config = exp_factory.get_exp_config(config_name)
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task, exp_cfg.ConstCLPretrainTask)
+ self.assertIsInstance(config.task.model, exp_cfg.ConstCLModel)
+ self.assertIsInstance(config.task.losses, exp_cfg.ConstCLLosses)
+ self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
+ config.task.train_data.is_training = None
+ with self.assertRaises(KeyError):
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/configs/head.py b/official/projects/const_cl/configs/head.py
new file mode 100644
index 00000000000..5d7469707f3
--- /dev/null
+++ b/official/projects/const_cl/configs/head.py
@@ -0,0 +1,76 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Configs for different model heads."""
+import dataclasses
+
+from typing import Optional
+
+from official.modeling import hyperparams
+
+
+@dataclasses.dataclass
+class MLP(hyperparams.Config):
+ """Config for the MLP head."""
+ normalize_inputs: bool = True
+ num_hidden_channels: int = 2048
+ num_hidden_layers: int = 3
+ num_output_channels: int = 128
+ use_sync_bn: bool = False
+ norm_momentum: float = 0.997
+ norm_epsilon: float = 1e-5
+ activation: Optional[str] = 'relu'
+
+
+@dataclasses.dataclass
+class InstanceReconstructor(hyperparams.Config):
+ """Config for the instance reconstructor head."""
+ normalize_inputs: bool = True
+ context_level: int = 5
+ # parameters for projector
+ num_output_channels: int = 2048
+ # parameters for RoiAligner
+ crop_size: int = 4
+ sample_offset: float = 0.5
+ # parameters for TxDecoder
+ num_tx_channels: int = 128
+ num_tx_layers: int = 3
+ num_tx_heads: int = 3
+ use_bias: bool = True
+ activation: Optional[str] = 'gelu'
+ dropout_rate: float = 0.0
+ layer_norm_epsilon: float = 1e-12
+ use_positional_embedding: bool = True
+
+
+@dataclasses.dataclass
+class ActionTransformer(hyperparams.Config):
+ """Config for the action transformer head."""
+ # parameters for classifier
+ num_hidden_layers: int = 0
+ num_hidden_channels: int = 0
+ use_sync_bn: bool = True
+ activation: str = 'relu'
+ # parameters for RoiAligner
+ crop_size: int = 4
+ sample_offset: float = 0.5
+ # parameters for TxDecoder
+ num_tx_channels: int = 128
+ num_tx_layers: int = 3
+ num_tx_heads: int = 3
+ use_bias: bool = True
+ tx_activation: Optional[str] = 'gelu'
+ dropout_rate: float = 0.0
+ layer_norm_epsilon: float = 1e-12
+ use_positional_embedding: bool = True
diff --git a/official/projects/const_cl/configs/head_test.py b/official/projects/const_cl/configs/head_test.py
new file mode 100644
index 00000000000..ae1136114ae
--- /dev/null
+++ b/official/projects/const_cl/configs/head_test.py
@@ -0,0 +1,49 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for head."""
+
+import tensorflow as tf, tf_keras
+from official.projects.const_cl.configs import head as head_cfg
+
+
+class HeadTest(tf.test.TestCase):
+
+ def test_mlp_head_valid(self):
+ config = head_cfg.MLP(
+ num_hidden_channels=128,
+ num_hidden_layers=4,
+ num_output_channels=1280,
+ use_sync_bn=True,
+ norm_momentum=0.99,
+ norm_epsilon=1e-5,
+ activation='relu')
+ config.validate()
+
+ def test_instance_reconstructor_head_valid(self):
+ config = head_cfg.InstanceReconstructor(
+ num_output_channels=1280,
+ layer_norm_epsilon=1e-12,
+ activation='relu')
+ config.validate()
+
+ def test_action_transformer_head_valid(self):
+ config = head_cfg.ActionTransformer(
+ activation='relu',
+ tx_activation='relu')
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/configs/yaml/const_cl_pretrain_k400_100k.yaml b/official/projects/const_cl/configs/yaml/const_cl_pretrain_k400_100k.yaml
new file mode 100644
index 00000000000..6e1dffbe297
--- /dev/null
+++ b/official/projects/const_cl/configs/yaml/const_cl_pretrain_k400_100k.yaml
@@ -0,0 +1,77 @@
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ dropout_rate: 1.0
+ norm_activation:
+ use_sync_bn: true
+ global_head:
+ normalize_inputs: false
+ use_sync_bn: true
+ backbone:
+ type: 'resnet_3dy'
+ resnet_3dy:
+ block_specs: !!python/tuple
+ - temporal_kernel_sizes: !!python/tuple
+ - 1
+ - 1
+ - 1
+ temporal_strides: 1
+ use_self_gating: false
+ - temporal_kernel_sizes: !!python/tuple
+ - 1
+ - 1
+ - 1
+ - 1
+ temporal_strides: 1
+ use_self_gating: false
+ - temporal_kernel_sizes: !!python/tuple
+ - 3
+ - 3
+ - 3
+ - 3
+ - 3
+ - 3
+ temporal_strides: 1
+ use_self_gating: false
+ - temporal_kernel_sizes: !!python/tuple
+ - 3
+ - 3
+ - 3
+ temporal_strides: 1
+ use_self_gating: false
+ model_id: 50
+ stem_conv_temporal_kernel_size: 5
+ stem_conv_temporal_stride: 2
+ stem_pool_temporal_stride: 1
+ train_data:
+ name: kinetics400
+ num_examples: 235693
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 2
+ global_batch_size: 1024
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 1024
+ prefetch_buffer_size: 1024
+ losses:
+ l2_weight_decay: 0.000001
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 40.96
+ decay_steps: 100000
+ optimizer:
+ sgd:
+ nesterov: false
+ warmup:
+ linear:
+ warmup_steps: 2500
+ train_steps: 100000
+ steps_per_loop: 100
+ summary_interval: 100
diff --git a/official/projects/const_cl/configs/yaml/const_cl_pretrain_k400_200k.yaml b/official/projects/const_cl/configs/yaml/const_cl_pretrain_k400_200k.yaml
new file mode 100644
index 00000000000..565ea896528
--- /dev/null
+++ b/official/projects/const_cl/configs/yaml/const_cl_pretrain_k400_200k.yaml
@@ -0,0 +1,77 @@
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ dropout_rate: 1.0
+ norm_activation:
+ use_sync_bn: true
+ global_head:
+ normalize_inputs: false
+ use_sync_bn: true
+ backbone:
+ type: 'resnet_3dy'
+ resnet_3dy:
+ block_specs: !!python/tuple
+ - temporal_kernel_sizes: !!python/tuple
+ - 1
+ - 1
+ - 1
+ temporal_strides: 1
+ use_self_gating: false
+ - temporal_kernel_sizes: !!python/tuple
+ - 1
+ - 1
+ - 1
+ - 1
+ temporal_strides: 1
+ use_self_gating: false
+ - temporal_kernel_sizes: !!python/tuple
+ - 3
+ - 3
+ - 3
+ - 3
+ - 3
+ - 3
+ temporal_strides: 1
+ use_self_gating: false
+ - temporal_kernel_sizes: !!python/tuple
+ - 3
+ - 3
+ - 3
+ temporal_strides: 1
+ use_self_gating: false
+ model_id: 50
+ stem_conv_temporal_kernel_size: 5
+ stem_conv_temporal_stride: 2
+ stem_pool_temporal_stride: 1
+ train_data:
+ name: kinetics400
+ num_examples: 235693
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 2
+ global_batch_size: 1024
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 1024
+ prefetch_buffer_size: 1024
+ losses:
+ l2_weight_decay: 0.000001
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 40.96
+ decay_steps: 200000
+ optimizer:
+ sgd:
+ nesterov: false
+ warmup:
+ linear:
+ warmup_steps: 2500
+ train_steps: 200000
+ steps_per_loop: 100
+ summary_interval: 100
diff --git a/official/projects/const_cl/datasets/video_ssl_inputs.py b/official/projects/const_cl/datasets/video_ssl_inputs.py
new file mode 100644
index 00000000000..f63401fbefd
--- /dev/null
+++ b/official/projects/const_cl/datasets/video_ssl_inputs.py
@@ -0,0 +1,105 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Video SSL datasets."""
+
+from typing import Dict, Tuple, Optional
+import tensorflow as tf, tf_keras
+
+from official.projects.video_ssl.dataloaders import video_ssl_input
+
+
+IMAGE_KEY = video_ssl_input.IMAGE_KEY
+LABEL_KEY = video_ssl_input.LABEL_KEY
+Decoder = video_ssl_input.Decoder
+
+
+class Parser(video_ssl_input.Parser):
+ """Parses a video dataset for SSL."""
+
+ def __init__(self,
+ input_params,
+ image_key: str = IMAGE_KEY,
+ label_key: str = LABEL_KEY):
+ super().__init__(
+ input_params=input_params,
+ image_key=image_key,
+ label_key=label_key)
+
+ self._num_instances = input_params.num_instances
+ self._num_frames = input_params.feature_shape[0]
+
+ def _generate_random_positions(
+ self, seed: Optional[int] = None) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Generates random instance positions in videos."""
+
+ num_frames = self._num_frames * 2 if self._is_ssl else self._num_frames
+ shape = [num_frames, self._num_instances, 4]
+ xmin = tf.random.uniform(shape=shape[:-1],
+ minval=0.0,
+ maxval=1.0,
+ dtype=tf.float32,
+ seed=seed)
+ ymin = tf.random.uniform(shape=shape[:-1],
+ minval=0.0,
+ maxval=1.0,
+ dtype=tf.float32,
+ seed=seed)
+ xdelta = tf.random.uniform(shape=shape[:-1],
+ minval=0.1,
+ maxval=0.5,
+ dtype=tf.float32,
+ seed=seed)
+ aspect_ratio = tf.random.uniform(shape=shape[:-1],
+ minval=0.5,
+ maxval=2.0,
+ dtype=tf.float32,
+ seed=seed)
+ ydelta = xdelta * aspect_ratio
+ xmax = tf.math.minimum(xmin + xdelta, 1.0 - 1e-3)
+ ymax = tf.math.minimum(ymin + ydelta, 1.0 - 1e-3)
+ random_positions = tf.stack([ymin, xmin, ymax, xmax], axis=-1)
+ random_positions = tf.cast(random_positions, dtype=self._dtype)
+ instances_mask = tf.ones(shape[:-1], dtype=tf.bool)
+ return random_positions, instances_mask
+
+ def _parse_train_data(
+ self, decoded_tensors: Dict[str, tf.Tensor]
+ ) -> Tuple[Dict[str, tf.Tensor], tf.Tensor]:
+ """Parses data for training."""
+ features, label = super()._parse_train_data(decoded_tensors=decoded_tensors)
+ instances_position, instances_mask = self._generate_random_positions(
+ seed=1234)
+ features.update({
+ 'instances_position': instances_position,
+ 'instances_mask': instances_mask,
+ })
+ return features, label
+
+
+class PostBatchProcessor(video_ssl_input.PostBatchProcessor):
+ """Processes a video and label dataset which is batched."""
+
+ def __call__(self,
+ features: Dict[str, tf.Tensor],
+ label: tf.Tensor) -> Tuple[Dict[str, tf.Tensor], tf.Tensor]:
+ """Postprocesses features and label tensors."""
+ features, label = super().__call__(features, label)
+
+ for key in ['instances_position', 'instances_mask']:
+ if key in features and self._is_ssl and self._is_training:
+ features[key] = tf.concat(
+ tf.split(features[key], num_or_size_splits=2, axis=1), axis=0)
+
+ return features, label
diff --git a/official/projects/const_cl/datasets/video_ssl_inputs_test.py b/official/projects/const_cl/datasets/video_ssl_inputs_test.py
new file mode 100644
index 00000000000..0ab6bc02fb9
--- /dev/null
+++ b/official/projects/const_cl/datasets/video_ssl_inputs_test.py
@@ -0,0 +1,81 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for video_ssl_inputs."""
+
+import io
+import numpy as np
+from PIL import Image
+
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.configs import const_cl as exp_cfg
+from official.projects.const_cl.datasets import video_ssl_inputs
+
+
+AUDIO_KEY = 'features/audio'
+
+
+def fake_seq_example():
+ """Creates fake data."""
+ random_image = np.random.randint(0, 256, size=(263, 320, 3), dtype=np.uint8)
+ random_image = Image.fromarray(random_image)
+ label = 42
+ with io.BytesIO() as buffer:
+ random_image.save(buffer, format='JPEG')
+ raw_image_bytes = buffer.getvalue()
+
+ seq_example = tf.train.SequenceExample()
+ seq_example.feature_lists.feature_list.get_or_create(
+ video_ssl_inputs.IMAGE_KEY).feature.add().bytes_list.value[:] = [
+ raw_image_bytes
+ ]
+ seq_example.feature_lists.feature_list.get_or_create(
+ video_ssl_inputs.IMAGE_KEY).feature.add().bytes_list.value[:] = [
+ raw_image_bytes
+ ]
+ seq_example.context.feature[
+ video_ssl_inputs.LABEL_KEY].int64_list.value[:] = [label]
+
+ random_audio = np.random.normal(size=(10, 256)).tolist()
+ for s in random_audio:
+ seq_example.feature_lists.feature_list.get_or_create(
+ AUDIO_KEY).feature.add().float_list.value[:] = s
+ return seq_example, label
+
+
+class VideoSslInputsTest(tf.test.TestCase):
+
+ def test_video_ssl_input_pretrain(self):
+ params = exp_cfg.const_cl_pretrain_kinetics400().task.train_data
+
+ decoder = video_ssl_inputs.Decoder()
+ parser = video_ssl_inputs.Parser(params).parse_fn(params.is_training)
+ seq_example, _ = fake_seq_example()
+
+ input_tensor = tf.constant(seq_example.SerializeToString())
+ decoded_tensors = decoder.decode(input_tensor)
+ output_tensor = parser(decoded_tensors)
+ features, _ = output_tensor
+ image = features['image']
+ instances_position = features['instances_position']
+ instances_mask = features['instances_mask']
+
+ self.assertAllEqual(image.shape, (32, 224, 224, 3))
+ self.assertAllEqual(instances_position.shape, (32, 8, 4))
+ self.assertAllEqual(instances_mask.shape, (32, 8))
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/losses/losses.py b/official/projects/const_cl/losses/losses.py
new file mode 100644
index 00000000000..1213d87f19d
--- /dev/null
+++ b/official/projects/const_cl/losses/losses.py
@@ -0,0 +1,361 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""The losses for ConST-CL."""
+
+from typing import Mapping
+
+import tensorflow as tf, tf_keras
+
+from tensorflow.compiler.tf2xla.python import xla # pylint: disable=g-direct-tensorflow-import
+from official.projects.video_ssl.losses import losses as video_ssl_losses
+
+tpu_cross_replica_concat = video_ssl_losses.tpu_cross_replica_concat
+
+
+_LARGE_NUM = 1e9
+
+
+class ContrastiveLoss(object):
+ """InfoNCE loss.
+
+ Reference: Oord et al. "Representation learning with contrastive
+ predictive coding" NeurIPS 2019.
+ """
+
+ def __init__(self,
+ normalize_inputs: bool,
+ temperature: float):
+ """Computes contrastive loss.
+
+ Args:
+ normalize_inputs: whether or not to l2 normalize the inputs vector.
+ temperature: temperature in the InfoNCE contrastive loss.
+ """
+ self._normalize_inputs = normalize_inputs
+ self._temperature = temperature
+
+ def __call__(self,
+ inputs: tf.Tensor,
+ num_replicas: int = 1) -> Mapping[str, tf.Tensor]:
+ """Calculates the loss.
+
+ Args:
+ inputs: the embeddings (in shape [2*B, C]) from video clips after the
+ projection head.
+ num_replicas: the number of TPU replicas.
+
+ Returns:
+ a dictionary contains calculated loss and statistics.
+ """
+ inputs1, inputs2 = tf.split(inputs, num_or_size_splits=2, axis=0)
+ if self._normalize_inputs:
+ inputs1 = tf.math.l2_normalize(inputs1, -1)
+ inputs2 = tf.math.l2_normalize(inputs2, -1)
+ batch_size = tf.shape(inputs1)[0]
+
+ if num_replicas == 1:
+ # This is the local version.
+ inputs1_large = inputs1
+ inputs2_large = inputs2
+ labels = tf.one_hot(tf.range(batch_size), batch_size * 2)
+ masks = tf.one_hot(tf.range(batch_size), batch_size)
+ else:
+ # This is the cross-tpu version.
+ inputs1_large = tpu_cross_replica_concat(inputs1, num_replicas)
+ inputs2_large = tpu_cross_replica_concat(inputs2, num_replicas)
+ enlarged_batch_size = tf.shape(inputs1_large)[0]
+ replica_id = tf.cast(tf.cast(xla.replica_id(), tf.uint32), tf.int32)
+ labels_idx = tf.range(batch_size) + replica_id * batch_size
+ labels = tf.one_hot(labels_idx, enlarged_batch_size * 2)
+ masks = tf.one_hot(labels_idx, enlarged_batch_size)
+
+ logits_aa = tf.matmul(
+ inputs1, inputs1_large, transpose_b=True) / self._temperature
+ logits_aa = logits_aa - tf.cast(masks, logits_aa.dtype) * _LARGE_NUM
+ logits_bb = tf.matmul(
+ inputs2, inputs2_large, transpose_b=True) / self._temperature
+ logits_bb = logits_bb - tf.cast(masks, logits_bb.dtype) * _LARGE_NUM
+ logits_ab = tf.matmul(
+ inputs1, inputs2_large, transpose_b=True) / self._temperature
+ logits_ba = tf.matmul(
+ inputs2, inputs1_large, transpose_b=True) / self._temperature
+
+ loss_a = tf.reduce_mean(tf.nn.softmax_cross_entropy_with_logits(
+ labels, tf.concat([logits_ab, logits_aa], 1)))
+ loss_b = tf.reduce_mean(tf.nn.softmax_cross_entropy_with_logits(
+ labels, tf.concat([logits_ba, logits_bb], 1)))
+ loss = loss_a + loss_b
+
+ contrast_prob = tf.nn.softmax(logits_ab)
+ contrast_entropy = - tf.reduce_mean(
+ tf.reduce_sum(contrast_prob * tf.math.log(contrast_prob + 1e-8), -1))
+
+ contrast_acc = tf.equal(tf.argmax(labels, 1), tf.argmax(logits_ab, axis=1))
+ contrast_acc = tf.reduce_mean(tf.cast(contrast_acc, tf.float32))
+
+ return {
+ 'loss': loss,
+ 'contrastive_accuracy': contrast_acc,
+ 'contrastive_entropy': contrast_entropy,
+ }
+
+
+class InstanceContrastiveLoss(object):
+ """Instance Contrastive Loss.
+
+ Reference: Yuan et al. "Contextualized Spatio-Temporal Contrastive Learning
+ with Self-Supervision" CVPR 2022.
+ """
+
+ def __init__(self,
+ normalize_inputs: bool,
+ temperature: float):
+ self._normalize_inputs = normalize_inputs
+ self._temperature = temperature
+
+ def __call__(self,
+ predictions: Mapping[str, tf.Tensor],
+ num_replicas: int = 1) -> Mapping[str, tf.Tensor]:
+ """Computes contrastive loss for spatio-temporal instance embeddings.
+
+ Args:
+ predictions: a dictionary of the model outputs, contains
+ 'instances_a2b': the reconstructed instance features from view a -> b.
+ In shape [B, N, C].
+ 'instances_b2a': the reconstructed instance features from view b -> a.
+ In shape [B, N, C].
+ 'instances_a': the target instance features in view a. In shape
+ [B, N, C].
+ 'instances_b': the target instance features in view b. In shape
+ [B, N, C].
+ 'masks_a': the vaidity boolean mask for instances in view a. In shape
+ [B, N].
+ 'masks_b': the vaidity boolean mask for instances in view b. In shape
+ [B, N].
+ num_replicas: the number of TPU replicas.
+
+ Returns:
+ A loss scalar.
+ The staticstics for positive examples.
+ The staticstics for negative examples.
+ """
+
+ inst_a2b = predictions['instances_a2b']
+ inst_b2a = predictions['instances_b2a']
+ inst_a = predictions['instances_a']
+ inst_b = predictions['instances_b']
+ masks_a = tf.cast(predictions['masks_a'][..., None], dtype=inst_a.dtype)
+ masks_b = tf.cast(predictions['masks_b'][..., None], dtype=inst_b.dtype)
+
+ if self._normalize_inputs:
+ inst_a2b = tf.math.l2_normalize(inst_a2b, axis=-1)
+ inst_b2a = tf.math.l2_normalize(inst_b2a, axis=-1)
+ inst_a = tf.math.l2_normalize(inst_a, axis=-1)
+ inst_b = tf.math.l2_normalize(inst_b, axis=-1)
+
+ b, n = inst_a.shape.as_list()[:2]
+ batch_index = tf.range(b)
+
+ # Computes similarity based on raw features in view a and b.
+ similarity_ab = tf.einsum('ijc,ikc->ijk', inst_a, inst_b)
+
+ # Loss on translated_a2b.
+ similarity_ab_index = tf.argmax(similarity_ab, axis=2, output_type=tf.int32)
+ lookup_a2b_index = tf.stack(
+ [tf.tile(batch_index[:, None], [1, n]), similarity_ab_index], axis=-1)
+ loss_and_stats_a = self._compute_constrastive_loss(
+ positive_lookup_index=lookup_a2b_index,
+ inst_translated=inst_a2b,
+ inst_target=inst_b,
+ inst_mask=masks_a,
+ num_replicas=num_replicas)
+
+ # Loss on translated_b2a.
+ similarity_ba_index = tf.argmax(similarity_ab, axis=1, output_type=tf.int32)
+ lookup_b2a_index = tf.stack(
+ [tf.tile(batch_index[:, None], [1, n]), similarity_ba_index], axis=-1)
+ loss_and_stats_b = self._compute_constrastive_loss(
+ positive_lookup_index=lookup_b2a_index,
+ inst_translated=inst_b2a,
+ inst_target=inst_a,
+ inst_mask=masks_b,
+ num_replicas=num_replicas)
+
+ loss_and_stats = {}
+ for key in loss_and_stats_a:
+ loss_and_stats[key] = 0.5 * (
+ loss_and_stats_a[key] + loss_and_stats_b[key])
+ return loss_and_stats
+
+ def _get_negative_similarity_statistics(
+ self,
+ logits: tf.Tensor,
+ batch_masks: tf.Tensor,
+ inst_mask: tf.Tensor) -> Mapping[str, tf.Tensor]:
+ """Gets negative examples similarity statistics.
+
+ Args:
+ logits: the logits matrix.
+ batch_masks: the batch validity mask.
+ inst_mask: the instance validity mask.
+
+ Returns:
+ logs: a dictionary of logs.
+ """
+ # logits = [b, n, bl, n]
+ # batch_masks = [b, n, bl, n]
+ # inst_mask = [b, n, 1]
+ inst_mask = tf.cast(inst_mask, logits.dtype)
+ batch_masks = tf.cast(batch_masks, logits.dtype)
+ batch_masks = tf.ones_like(batch_masks) - batch_masks
+ masks = batch_masks * inst_mask[..., None]
+ # Recover the raw similarity and mask self-similarity, which will be
+ # removed from negative samples.
+ similarity = logits * masks * self._temperature
+ similarity_mean = tf.reduce_sum(similarity) / tf.reduce_sum(masks)
+
+ similarity_masks = tf.squeeze(inst_mask, axis=-1)
+ similarity_max = similarity - (1.0 - masks) * _LARGE_NUM
+ similarity_max = tf.reduce_max(similarity_max, axis=[-1, -2])
+ similarity_max = tf.reduce_sum(
+ similarity_max * similarity_masks) / tf.reduce_sum(similarity_masks)
+
+ similarity_min = similarity + (1.0 - masks) * _LARGE_NUM
+ similarity_min = tf.reduce_min(similarity_min, axis=[-1, -2])
+ similarity_min = tf.reduce_sum(
+ similarity_min * similarity_masks) / tf.reduce_sum(similarity_masks)
+ logs = {
+ 'negative_similarity_mean': similarity_mean,
+ 'negative_similarity_min': similarity_min,
+ 'negative_similarity_max': similarity_max,
+ }
+ return logs
+
+ def _get_positive_similarity_statistics(
+ self,
+ logits: tf.Tensor,
+ inst_mask: tf.Tensor) -> Mapping[str, tf.Tensor]:
+ """Gets positive examples similarity statistics.
+
+ Args:
+ logits: the logits matrix.
+ inst_mask: the instance validity mask.
+
+ Returns:
+ logs: a dictionary of logs.
+ """
+ # logits in shape [b, n]
+ # inst_mask in shape [b, n, 1]
+ inst_mask = tf.squeeze(inst_mask, axis=-1)
+ inst_mask = tf.cast(inst_mask, dtype=logits.dtype)
+ similarity = logits * inst_mask * self._temperature
+
+ num_instances = tf.reduce_sum(inst_mask)
+ similarity_mean = tf.reduce_sum(similarity) / num_instances
+
+ similarity_max = similarity - (1.0 - inst_mask) * _LARGE_NUM
+ similarity_max = tf.reduce_max(similarity_max)
+
+ similarity_min = similarity + (1.0 - inst_mask) * _LARGE_NUM
+ similarity_min = tf.reduce_min(similarity_min)
+
+ logs = {
+ 'positive_similarity_mean': similarity_mean,
+ 'positive_similarity_min': similarity_min,
+ 'positive_similarity_max': similarity_max,
+ }
+ return logs
+
+ def _compute_constrastive_loss(
+ self,
+ positive_lookup_index: tf.Tensor,
+ inst_translated: tf.Tensor,
+ inst_target: tf.Tensor,
+ inst_mask: tf.Tensor,
+ num_replicas: int = 1) -> Mapping[str, tf.Tensor]:
+ """Computes constrastive loss.
+
+ Args:
+ positive_lookup_index: the index tensor to look-up the corresponding
+ features in inst_target. In shape [B, N].
+ inst_translated: a float tensor of shape [B, N, C] of translated instance
+ features by the transformer head.
+ inst_target: a float tensor of shape [B, N, C] of instance features on the
+ target domain. Note that the order of inst_target is not necessarily
+ matched to inst_translated.
+ inst_mask: a boolean tensor of shape [B, N, 1] suggesting valid instances
+ in inst_translated.
+ num_replicas: the number of TPU replicas.
+
+ Returns:
+ loss_and_stats: a dictionary of loss and intermediate statistics.
+ """
+ b, n = inst_translated.shape.as_list()[:2]
+
+ if num_replicas == 1:
+ inst_target_large = inst_target
+ b_large = tf.shape(inst_target_large)[0]
+ labels_idx = tf.range(b)
+ else:
+ inst_target_large = tpu_cross_replica_concat(
+ inst_target,
+ num_replicas)
+ b_large = tf.shape(inst_target_large)[0]
+ # NOTE: make sure to use xla.replica_id() here and in
+ # tpu_cross_replica_concat to consistently align the replica_id.
+ # replicator.replica_id != xla.replica_id()
+ replica_id = tf.cast(tf.cast(xla.replica_id(), tf.uint32), tf.int32)
+ labels_idx = tf.range(b) + replica_id * b
+
+ # [B, BL], 1 indicates positive batches.
+ batch_masks = tf.one_hot(labels_idx, b_large)
+ # [B, N, BL, N]
+ batch_masks = tf.tile(batch_masks[:, None, :, None], [1, n, 1, n])
+
+ # Construct negative examples.
+ logits_negative = tf.einsum(
+ 'ijc,pqc->ijpq',
+ inst_translated, inst_target_large) / self._temperature
+ # Get negative statistics.
+ negative_stats = self._get_negative_similarity_statistics(
+ logits_negative, batch_masks, inst_mask)
+ logits_negative = logits_negative - tf.cast(
+ batch_masks, logits_negative.dtype) * _LARGE_NUM
+ logits_negative = tf.reshape(logits_negative, [b * n, b_large * n])
+
+ # Construct positive examples.
+ inst_matched = tf.gather_nd(
+ inst_target, positive_lookup_index, name='matched_inst')
+ logits_positive = tf.einsum(
+ 'ijc,ijc->ij',
+ inst_translated, inst_matched) / self._temperature
+ # Get positive statistics.
+ positive_stats = self._get_positive_similarity_statistics(
+ logits_positive, inst_mask)
+ logits_positive = tf.reshape(logits_positive, [b * n, 1])
+
+ logits_all = tf.concat([logits_positive, logits_negative], axis=1)
+ loss_pos = tf.reduce_logsumexp(logits_positive, 1)
+ loss_all = tf.reduce_logsumexp(logits_all, 1)
+ loss = (loss_all - loss_pos) * tf.reshape(inst_mask, [b * n])
+
+ # Average across instances.
+ loss = tf.math.divide_no_nan(
+ tf.reduce_sum(loss), tf.reduce_sum(inst_mask))
+
+ loss_and_stats = {'loss': loss}
+ loss_and_stats.update(negative_stats)
+ loss_and_stats.update(positive_stats)
+ return loss_and_stats
diff --git a/official/projects/const_cl/losses/losses_test.py b/official/projects/const_cl/losses/losses_test.py
new file mode 100644
index 00000000000..641124e0305
--- /dev/null
+++ b/official/projects/const_cl/losses/losses_test.py
@@ -0,0 +1,85 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for losses."""
+
+import tensorflow as tf, tf_keras
+from official.projects.const_cl.losses import losses
+
+
+class LossesTest(tf.test.TestCase):
+
+ def test_constrative_loss(self):
+ contrastive_loss = losses.ContrastiveLoss(normalize_inputs=True,
+ temperature=0.1)
+ inputs1 = tf.constant(
+ [[1, 2, 3, 4], [5, 6, 7, 8], [4, 3, 2, 1], [8, 7, 6, 5]],
+ dtype=tf.float32)
+ inputs2 = tf.constant(
+ [[1, 2, 3, 4], [4, 3, 2, 1], [5, 6, 7, 8], [8, 7, 6, 5]],
+ dtype=tf.float32)
+ inputs = tf.concat([inputs1, inputs2], axis=0)
+ contrastive_loss_dict = contrastive_loss(inputs)
+
+ self.assertAlmostEqual(contrastive_loss_dict['contrastive_accuracy'], 0.5)
+ self.assertAlmostEqual(
+ contrastive_loss_dict['loss'].numpy(), 4.136947, places=4)
+
+ def test_instance_constrative_loss(self):
+ instance_contrastive_loss = losses.InstanceContrastiveLoss(
+ normalize_inputs=True, temperature=0.1)
+ inst_a = tf.constant(
+ [[[1, 2, 3, 4], [5, 6, 7, 8], [-1, -1, -1, -1], [-1, -1, -1, -1]],
+ [[-1, -1, -1, -1], [-1, -1, -1, -1], [4, 3, 2, 1], [8, 7, 6, 5]],
+ [[1, 2, 3, 4], [5, 6, 7, 8], [4, 3, 2, 1], [8, 7, 6, 5]]],
+ dtype=tf.float32)
+ inst_b = tf.constant([[[1.5, 2.5, 3.5, 4.5], [5.5, 6.5, 7.5, 8.5],
+ [-1, -1, -1, -1], [-1, -1, -1, -1]],
+ [[-1, -1, -1, -1], [-1, -1, -1, -1],
+ [5.5, 6.5, 7.5, 8.5], [8.5, 7.5, 6.5, 5.5]],
+ [[1.5, 2.5, 3.5, 4.5], [4.5, 3.5, 2.5, 1.5],
+ [5.5, 6.5, 7.5, 8.5], [8.5, 7.5, 6.5, 5.5]]],
+ dtype=tf.float32)
+
+ inst_a2b = inst_b
+ inst_b2a = inst_a
+
+ masks_a = tf.constant(
+ [[True, True, False, False],
+ [False, False, True, True],
+ [True, True, True, True]], dtype=tf.bool)
+ masks_b = tf.constant(
+ [[True, True, False, False],
+ [False, False, True, True],
+ [True, True, True, True]], dtype=tf.bool)
+
+ predictions = {
+ 'instances_a': inst_a,
+ 'instances_b': inst_b,
+ 'instances_a2b': inst_a2b,
+ 'instances_b2a': inst_b2a,
+ 'masks_a': masks_a,
+ 'masks_b': masks_b}
+ contrastive_loss_dict = instance_contrastive_loss(
+ predictions=predictions)
+
+ self.assertContainsSubset(
+ list(contrastive_loss_dict.keys()), [
+ 'loss', 'positive_similarity_mean', 'positive_similarity_min',
+ 'positive_similarity_max', 'negative_similarity_mean',
+ 'negative_similarity_min', 'negative_similarity_max'
+ ])
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/modeling/backbones/nn_blocks_3d.py b/official/projects/const_cl/modeling/backbones/nn_blocks_3d.py
new file mode 100644
index 00000000000..3ca84724e45
--- /dev/null
+++ b/official/projects/const_cl/modeling/backbones/nn_blocks_3d.py
@@ -0,0 +1,122 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains common building blocks for 3D networks."""
+import tensorflow as tf, tf_keras
+
+from official.vision.modeling.layers import nn_blocks_3d
+from official.vision.modeling.layers import nn_layers
+
+SelfGating = nn_blocks_3d.SelfGating
+
+
+class BottleneckBlock3D(nn_blocks_3d.BottleneckBlock3D):
+ """Creates a 3D bottleneck block."""
+
+ def build(self, input_shape):
+ self._shortcut_maxpool = tf_keras.layers.MaxPool3D(
+ pool_size=[1, 1, 1],
+ strides=[
+ self._temporal_strides, self._spatial_strides, self._spatial_strides
+ ])
+
+ self._shortcut_conv = tf_keras.layers.Conv3D(
+ filters=4 * self._filters,
+ kernel_size=1,
+ strides=[
+ self._temporal_strides, self._spatial_strides, self._spatial_strides
+ ],
+ use_bias=False,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='shortcut_conv')
+ self._norm0 = self._norm(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ name='shortcut_conv/batch_norm')
+
+ self._temporal_conv = tf_keras.layers.Conv3D(
+ filters=self._filters,
+ kernel_size=[self._temporal_kernel_size, 1, 1],
+ strides=[self._temporal_strides, 1, 1],
+ padding='same',
+ use_bias=False,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='temporal_conv')
+ self._norm1 = self._norm(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ name='temporal_conv/batch_norm')
+
+ self._spatial_conv = tf_keras.layers.Conv3D(
+ filters=self._filters,
+ kernel_size=[1, 3, 3],
+ strides=[1, self._spatial_strides, self._spatial_strides],
+ padding='same',
+ use_bias=False,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='spatial_conv')
+ self._norm2 = self._norm(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ name='spatial_conv/batch_norm')
+
+ self._expand_conv = tf_keras.layers.Conv3D(
+ filters=4 * self._filters,
+ kernel_size=[1, 1, 1],
+ strides=[1, 1, 1],
+ padding='same',
+ use_bias=False,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='expand_conv')
+ self._norm3 = self._norm(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ name='expand_conv/batch_norm/')
+
+ if self._se_ratio and self._se_ratio > 0 and self._se_ratio <= 1:
+ self._squeeze_excitation = nn_layers.SqueezeExcitation(
+ in_filters=self._filters * 4,
+ out_filters=self._filters * 4,
+ se_ratio=self._se_ratio,
+ use_3d_input=True,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='se_layer')
+ else:
+ self._squeeze_excitation = None
+
+ if self._stochastic_depth_drop_rate:
+ self._stochastic_depth = nn_layers.StochasticDepth(
+ self._stochastic_depth_drop_rate)
+ else:
+ self._stochastic_depth = None
+
+ if self._use_self_gating:
+ self._self_gating = SelfGating(filters=4 * self._filters,
+ name='self_gating')
+ else:
+ self._self_gating = None
diff --git a/official/projects/const_cl/modeling/backbones/nn_blocks_3d_test.py b/official/projects/const_cl/modeling/backbones/nn_blocks_3d_test.py
new file mode 100644
index 00000000000..72978bc592b
--- /dev/null
+++ b/official/projects/const_cl/modeling/backbones/nn_blocks_3d_test.py
@@ -0,0 +1,70 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for resnet."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.modeling.backbones import nn_blocks_3d
+
+
+class NNBlocksTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (nn_blocks_3d.BottleneckBlock3D, 1, 1, 2, True, 0.2, 0.1),
+ (nn_blocks_3d.BottleneckBlock3D, 3, 2, 1, False, 0.0, 0.0),
+ )
+ def test_bottleneck_block_creation(self, block_fn, temporal_kernel_size,
+ temporal_strides, spatial_strides,
+ use_self_gating, se_ratio,
+ stochastic_depth):
+ temporal_size = 16
+ spatial_size = 128
+ filters = 256
+ inputs = tf_keras.Input(
+ shape=(temporal_size, spatial_size, spatial_size, filters * 4),
+ batch_size=1)
+ block = block_fn(
+ filters=filters,
+ temporal_kernel_size=temporal_kernel_size,
+ temporal_strides=temporal_strides,
+ spatial_strides=spatial_strides,
+ use_self_gating=use_self_gating,
+ se_ratio=se_ratio,
+ stochastic_depth_drop_rate=stochastic_depth)
+
+ features = block(inputs)
+
+ self.assertAllEqual([
+ 1, temporal_size // temporal_strides, spatial_size // spatial_strides,
+ spatial_size // spatial_strides, filters * 4
+ ], features.shape.as_list())
+ vnames = [v.name for v in block.trainable_variables]
+ expected_names = [
+ 'bottleneck_block3d/temporal_conv/kernel:0',
+ 'bottleneck_block3d/temporal_conv/batch_norm/gamma:0',
+ 'bottleneck_block3d/temporal_conv/batch_norm/beta:0',
+ 'bottleneck_block3d/spatial_conv/kernel:0',
+ 'bottleneck_block3d/spatial_conv/batch_norm/gamma:0',
+ 'bottleneck_block3d/spatial_conv/batch_norm/beta:0',
+ 'bottleneck_block3d/expand_conv/kernel:0',
+ 'bottleneck_block3d/expand_conv/batch_norm/gamma:0',
+ 'bottleneck_block3d/expand_conv/batch_norm/beta:0'
+ ]
+ self.assertContainsSubset(expected_names, vnames)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/modeling/backbones/resnet_3d.py b/official/projects/const_cl/modeling/backbones/resnet_3d.py
new file mode 100644
index 00000000000..e3be5217164
--- /dev/null
+++ b/official/projects/const_cl/modeling/backbones/resnet_3d.py
@@ -0,0 +1,391 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definitions of 3D Residual Networks."""
+from typing import Any, Callable, List, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import hyperparams
+from official.modeling import tf_utils
+from official.projects.const_cl.modeling.backbones import nn_blocks_3d
+from official.vision.modeling.backbones import factory
+from official.vision.modeling.backbones import resnet_3d
+from official.vision.modeling.layers import nn_layers
+
+layers = tf_keras.layers
+
+RESNET_SPECS = resnet_3d.RESNET_SPECS
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class ResNet3DY(tf_keras.Model):
+ """Creates a 3D ResNet family model with branched res5 block."""
+
+ def __init__(
+ self,
+ model_id: int,
+ temporal_strides: List[int],
+ temporal_kernel_sizes: List[Tuple[int]],
+ use_self_gating: Optional[List[int]] = None,
+ input_specs: tf_keras.layers.InputSpec = layers.InputSpec(
+ shape=[None, None, None, None, 3]),
+ stem_type: str = 'v0',
+ stem_conv_temporal_kernel_size: int = 5,
+ stem_conv_temporal_stride: int = 2,
+ stem_pool_temporal_stride: int = 2,
+ init_stochastic_depth_rate: float = 0.0,
+ activation: str = 'relu',
+ se_ratio: Optional[float] = None,
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ kernel_initializer: str = 'VarianceScaling',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initializes a 3D ResNet model.
+
+ Args:
+ model_id: An `int` of depth of ResNet backbone model.
+ temporal_strides: A list of integers that specifies the temporal strides
+ for all 3d blocks.
+ temporal_kernel_sizes: A list of tuples that specifies the temporal kernel
+ sizes for all 3d blocks in different block groups.
+ use_self_gating: A list of booleans to specify applying self-gating module
+ or not in each block group. If None, self-gating is not applied.
+ input_specs: A `tf_keras.layers.InputSpec` of the input tensor.
+ stem_type: A `str` of stem type of ResNet. Default to `v0`. If set to
+ `v1`, use ResNet-D type stem (https://arxiv.org/abs/1812.01187).
+ stem_conv_temporal_kernel_size: An `int` of temporal kernel size for the
+ first conv layer.
+ stem_conv_temporal_stride: An `int` of temporal stride for the first conv
+ layer.
+ stem_pool_temporal_stride: An `int` of temporal stride for the first pool
+ layer.
+ init_stochastic_depth_rate: A `float` of initial stochastic depth rate.
+ activation: A `str` of name of the activation function.
+ se_ratio: A `float` or None. Ratio of the Squeeze-and-Excitation layer.
+ use_sync_bn: If True, use synchronized batch normalization.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_initializer: A str for kernel initializer of convolutional layers.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default to None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ Default to None.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+
+ self._model_id = model_id
+ self._temporal_strides = temporal_strides
+ self._temporal_kernel_sizes = temporal_kernel_sizes
+ self._input_specs = input_specs
+ self._stem_type = stem_type
+ self._stem_conv_temporal_kernel_size = stem_conv_temporal_kernel_size
+ self._stem_conv_temporal_stride = stem_conv_temporal_stride
+ self._stem_pool_temporal_stride = stem_pool_temporal_stride
+ self._use_self_gating = use_self_gating
+ self._se_ratio = se_ratio
+ self._init_stochastic_depth_rate = init_stochastic_depth_rate
+ self._use_sync_bn = use_sync_bn
+ self._activation = activation
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ if use_sync_bn:
+ self._norm = layers.experimental.SyncBatchNormalization
+ else:
+ self._norm = layers.BatchNormalization
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ else:
+ self._bn_axis = 1
+
+ # Build ResNet3D backbone.
+ inputs = tf_keras.Input(shape=input_specs.shape[1:])
+ self._build_model(inputs)
+
+ def _build_model(self, inputs):
+ """Builds model architecture.
+
+ Args:
+ inputs: the Keras input spec.
+
+ Returns:
+ endpoints: A dictionary of backbone endpoint features.
+ """
+ # Build stem.
+ self._build_stem(inputs, stem_type=self._stem_type)
+
+ temporal_kernel_size = 1 if self._stem_pool_temporal_stride == 1 else 3
+ self._max_pool = layers.MaxPool3D(
+ pool_size=[temporal_kernel_size, 3, 3],
+ strides=[self._stem_pool_temporal_stride, 2, 2],
+ padding='same')
+
+ # Build intermediate blocks and endpoints.
+ resnet_specs = RESNET_SPECS[self._model_id]
+ if len(self._temporal_strides) != len(resnet_specs) or len(
+ self._temporal_kernel_sizes) != len(resnet_specs):
+ raise ValueError(
+ 'Number of blocks in temporal specs should equal to resnet_specs.')
+
+ self._blocks = {}
+ for i, resnet_spec in enumerate(resnet_specs):
+ if resnet_spec[0] == 'bottleneck3d':
+ block_fn = nn_blocks_3d.BottleneckBlock3D
+ else:
+ raise ValueError('Block fn `{}` is not supported.'.format(
+ resnet_spec[0]))
+
+ use_self_gating = (
+ self._use_self_gating[i] if self._use_self_gating else False)
+ self._blocks[f'res_{i+2}'] = self._build_block_group(
+ inputs=inputs,
+ filters=resnet_spec[1],
+ temporal_kernel_sizes=self._temporal_kernel_sizes[i],
+ temporal_strides=self._temporal_strides[i],
+ spatial_strides=(1 if i == 0 else 2),
+ block_fn=block_fn,
+ block_repeats=resnet_spec[2],
+ stochastic_depth_drop_rate=nn_layers.get_stochastic_depth_rate(
+ self._init_stochastic_depth_rate, i + 2, 5),
+ use_self_gating=use_self_gating, # pyrefly: ignore[bad-argument-type]
+ name='res_{}'.format(i + 2))
+
+ # Duplicate res5 block.
+ resnet_specs = RESNET_SPECS[self._model_id]
+ resnet_spec = resnet_specs[-1]
+ i = len(resnet_specs) - 1
+
+ if resnet_spec[0] == 'bottleneck3d':
+ block_fn = nn_blocks_3d.BottleneckBlock3D
+ else:
+ raise ValueError('Block fn `{}` is not supported.'.format(
+ resnet_spec[0]))
+
+ use_self_gating = (
+ self._use_self_gating[i] if self._use_self_gating else False)
+ block_layers = self._build_block_group(
+ inputs=inputs,
+ filters=resnet_spec[1],
+ temporal_kernel_sizes=self._temporal_kernel_sizes[i],
+ temporal_strides=self._temporal_strides[i],
+ spatial_strides=(1 if i == 0 else 2),
+ block_fn=block_fn,
+ block_repeats=resnet_spec[2],
+ stochastic_depth_drop_rate=nn_layers.get_stochastic_depth_rate(
+ self._init_stochastic_depth_rate, i + 2, 5),
+ use_self_gating=use_self_gating, # pyrefly: ignore[bad-argument-type]
+ name='res_{}_1'.format(i + 2))
+ self._res_5_1_layers = block_layers
+
+ def _build_stem(self, inputs, stem_type):
+ """Builds stem layer."""
+ del inputs
+ # Build stem.
+ if stem_type == 'v0':
+ self._stem_conv = layers.Conv3D(
+ filters=64,
+ kernel_size=[self._stem_conv_temporal_kernel_size, 7, 7],
+ strides=[self._stem_conv_temporal_stride, 2, 2],
+ use_bias=False,
+ padding='same',
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='stem')
+ self._stem_bn = self._norm(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ name='stem/batch_norm')
+ self._stem_activation = tf_utils.get_activation(self._activation)
+ else:
+ raise ValueError(f'Stem type {stem_type} not supported.')
+
+ def _build_block_group(
+ self,
+ inputs: tf.Tensor,
+ filters: int,
+ temporal_kernel_sizes: Tuple[int],
+ temporal_strides: int,
+ spatial_strides: int,
+ block_fn: Callable[
+ ..., tf_keras.layers.Layer] = nn_blocks_3d.BottleneckBlock3D,
+ block_repeats: int = 1,
+ stochastic_depth_drop_rate: float = 0.0,
+ use_self_gating: bool = False,
+ name: str = 'block_group'):
+ """Creates one group of blocks for the ResNet3D model.
+
+ Args:
+ inputs: A `tf.Tensor` of size `[batch, channels, height, width]`.
+ filters: An `int` of number of filters for the first convolution of the
+ layer.
+ temporal_kernel_sizes: A tuple that specifies the temporal kernel sizes
+ for each block in the current group.
+ temporal_strides: An `int` of temporal strides for the first convolution
+ in this group.
+ spatial_strides: An `int` stride to use for the first convolution of the
+ layer. If greater than 1, this layer will downsample the input.
+ block_fn: Either `nn_blocks.ResidualBlock` or `nn_blocks.BottleneckBlock`.
+ block_repeats: An `int` of number of blocks contained in the layer.
+ stochastic_depth_drop_rate: A `float` of drop rate of the current block
+ group.
+ use_self_gating: A `bool` that specifies whether to apply self-gating
+ module or not.
+ name: A `str` name for the block.
+
+ Returns:
+ The output `tf.Tensor` of the block layer.
+ """
+ del inputs
+ if len(temporal_kernel_sizes) != block_repeats:
+ raise ValueError(
+ 'Number of elements in temporal_kernel_sizes must equal to '
+ 'block_repeats.')
+
+ # Only apply self-gating module in the last block.
+ use_self_gating_list = [False] * (block_repeats - 1) + [use_self_gating]
+
+ name = 'cell'
+ block_layers = {}
+ block_layers[f'{name}_0'] = block_fn(
+ filters=filters,
+ temporal_kernel_size=temporal_kernel_sizes[0],
+ temporal_strides=temporal_strides,
+ spatial_strides=spatial_strides,
+ stochastic_depth_drop_rate=stochastic_depth_drop_rate,
+ use_self_gating=use_self_gating_list[0],
+ se_ratio=self._se_ratio,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ activation=self._activation,
+ use_sync_bn=self._use_sync_bn,
+ norm_momentum=self._norm_momentum,
+ norm_epsilon=self._norm_epsilon,
+ name=f'{name}_0')
+
+ for i in range(1, block_repeats):
+ block_layers[f'{name}_{i}'] = block_fn(
+ filters=filters,
+ temporal_kernel_size=temporal_kernel_sizes[i],
+ temporal_strides=1,
+ spatial_strides=1,
+ stochastic_depth_drop_rate=stochastic_depth_drop_rate,
+ use_self_gating=use_self_gating_list[i],
+ se_ratio=self._se_ratio,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ activation=self._activation,
+ use_sync_bn=self._use_sync_bn,
+ norm_momentum=self._norm_momentum,
+ norm_epsilon=self._norm_epsilon,
+ name=f'{name}_{i}')
+
+ return block_layers
+
+ def call(self, inputs: tf.Tensor, training: bool = False, mask: Any = None):
+ """Calls ResNet3DY model."""
+ del mask
+ x = self._stem_conv(inputs, training=training)
+ x = self._stem_bn(x, training=training)
+ x = self._stem_activation(x)
+ x = self._max_pool(x)
+
+ res4 = None
+ endpoints = {}
+ for i, block_layers in enumerate(self._blocks.values()):
+ for block_fn in block_layers.values():
+ x = block_fn(x, training=training)
+ endpoints[f'{i + 2}'] = x
+ if i + 2 == 4:
+ res4 = x
+
+ for block_fn in self._res_5_1_layers.values():
+ res4 = block_fn(res4, training=training)
+ endpoints['5_1'] = res4
+ return endpoints
+
+ def get_config(self):
+ config_dict = {
+ 'model_id': self._model_id,
+ 'temporal_strides': self._temporal_strides,
+ 'temporal_kernel_sizes': self._temporal_kernel_sizes,
+ 'stem_type': self._stem_type,
+ 'stem_conv_temporal_kernel_size': self._stem_conv_temporal_kernel_size,
+ 'stem_conv_temporal_stride': self._stem_conv_temporal_stride,
+ 'stem_pool_temporal_stride': self._stem_pool_temporal_stride,
+ 'use_self_gating': self._use_self_gating,
+ 'se_ratio': self._se_ratio,
+ 'init_stochastic_depth_rate': self._init_stochastic_depth_rate,
+ 'activation': self._activation,
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_regularizer': self._bias_regularizer,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+
+@factory.register_backbone_builder('resnet_3dy')
+def build_resnet3dy(
+ input_specs: tf_keras.layers.InputSpec,
+ backbone_config: hyperparams.Config,
+ norm_activation_config: hyperparams.Config,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> tf_keras.Model:
+ """Builds ResNet 3d-Y backbone from a config."""
+ backbone_cfg = backbone_config.get()
+
+ # Flatten configs before passing to the backbone.
+ temporal_strides = []
+ temporal_kernel_sizes = []
+ use_self_gating = []
+ for block_spec in backbone_cfg.block_specs:
+ temporal_strides.append(block_spec.temporal_strides)
+ temporal_kernel_sizes.append(block_spec.temporal_kernel_sizes)
+ use_self_gating.append(block_spec.use_self_gating)
+
+ return ResNet3DY(
+ model_id=backbone_cfg.model_id,
+ temporal_strides=temporal_strides,
+ temporal_kernel_sizes=temporal_kernel_sizes,
+ use_self_gating=use_self_gating,
+ input_specs=input_specs,
+ stem_type=backbone_cfg.stem_type,
+ stem_conv_temporal_kernel_size=backbone_cfg
+ .stem_conv_temporal_kernel_size,
+ stem_conv_temporal_stride=backbone_cfg.stem_conv_temporal_stride,
+ stem_pool_temporal_stride=backbone_cfg.stem_pool_temporal_stride,
+ init_stochastic_depth_rate=backbone_cfg.stochastic_depth_drop_rate,
+ se_ratio=backbone_cfg.se_ratio,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer)
diff --git a/official/projects/const_cl/modeling/backbones/resnet_3d_test.py b/official/projects/const_cl/modeling/backbones/resnet_3d_test.py
new file mode 100644
index 00000000000..fc28f86d1d8
--- /dev/null
+++ b/official/projects/const_cl/modeling/backbones/resnet_3d_test.py
@@ -0,0 +1,104 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for resnet."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.modeling.backbones import resnet_3d
+
+
+class ResNet3DTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (128, 50, 4, 'v0', False, 0.0),
+ (128, 50, 4, 'v0', False, 0.2),
+ (256, 50, 4, 'v0', True, 0.2),
+ )
+ def test_network_creation(self, input_size, model_id, endpoint_filter_scale,
+ stem_type, se_ratio, init_stochastic_depth_rate):
+ """Test creation of ResNet3D family models."""
+ tf_keras.backend.set_image_data_format('channels_last')
+ temporal_strides = [1, 1, 1, 1]
+ temporal_kernel_sizes = [(3, 3, 3), (3, 1, 3, 1), (3, 1, 3, 1, 3, 1),
+ (1, 3, 1)]
+ use_self_gating = [True, False, True, False]
+
+ network = resnet_3d.ResNet3DY(
+ model_id=model_id,
+ temporal_strides=temporal_strides,
+ temporal_kernel_sizes=temporal_kernel_sizes,
+ use_self_gating=use_self_gating,
+ stem_type=stem_type,
+ se_ratio=se_ratio,
+ init_stochastic_depth_rate=init_stochastic_depth_rate)
+ inputs = tf_keras.Input(shape=(8, input_size, input_size, 3), batch_size=1)
+ endpoints = network(inputs)
+
+ self.assertAllEqual([
+ 1, 2, input_size / 2**2, input_size / 2**2, 64 * endpoint_filter_scale
+ ], endpoints['2'].shape.as_list())
+ self.assertAllEqual([
+ 1, 2, input_size / 2**3, input_size / 2**3, 128 * endpoint_filter_scale
+ ], endpoints['3'].shape.as_list())
+ self.assertAllEqual([
+ 1, 2, input_size / 2**4, input_size / 2**4, 256 * endpoint_filter_scale
+ ], endpoints['4'].shape.as_list())
+ self.assertAllEqual([
+ 1, 2, input_size / 2**5, input_size / 2**5, 512 * endpoint_filter_scale
+ ], endpoints['5'].shape.as_list())
+ self.assertAllEqual([
+ 1, 2, input_size / 2**5, input_size / 2**5, 512 * endpoint_filter_scale
+ ], endpoints['5_1'].shape.as_list())
+
+ def test_serialize_deserialize(self):
+ # Create a network object that sets all of its config options.
+ kwargs = dict(
+ model_id=50,
+ temporal_strides=[1, 1, 1, 1],
+ temporal_kernel_sizes=[(3, 3, 3), (3, 1, 3, 1), (3, 1, 3, 1, 3, 1),
+ (1, 3, 1)],
+ stem_type='v0',
+ stem_conv_temporal_kernel_size=5,
+ stem_conv_temporal_stride=2,
+ stem_pool_temporal_stride=2,
+ se_ratio=0.0,
+ use_self_gating=None,
+ init_stochastic_depth_rate=0.0,
+ use_sync_bn=False,
+ activation='relu',
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ )
+ network = resnet_3d.ResNet3DY(**kwargs)
+
+ expected_config = dict(kwargs)
+ self.assertEqual(network.get_config(), expected_config)
+
+ # Create another network object from the first object's config.
+ new_network = resnet_3d.ResNet3DY.from_config(network.get_config())
+
+ # Validate that the config can be forced to JSON.
+ _ = new_network.to_json()
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(network.get_config(), new_network.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/modeling/const_cl_model.py b/official/projects/const_cl/modeling/const_cl_model.py
new file mode 100644
index 00000000000..a7b3391c1a8
--- /dev/null
+++ b/official/projects/const_cl/modeling/const_cl_model.py
@@ -0,0 +1,232 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builds ConST-CL SSL models."""
+from typing import Mapping, Optional
+
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.configs import const_cl as const_cl_cfg
+from official.projects.const_cl.modeling.heads import instance_reconstructor
+from official.projects.const_cl.modeling.heads import simple
+
+from official.vision.modeling import backbones
+from official.vision.modeling import factory_3d as model_factory
+
+layers = tf_keras.layers
+
+
+class ConstCLModel(tf_keras.Model):
+ """A ConST-CL SSL model class builder."""
+
+ def __init__(
+ self,
+ backbone,
+ input_specs: Optional[Mapping[str, tf_keras.layers.InputSpec]] = None,
+ # global_head
+ num_hidden_layers: int = 3,
+ num_hidden_channels: int = 1024,
+ num_output_channels: int = 128,
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 1e-5,
+ activation: Optional[str] = None,
+ normalize_global_features: bool = False,
+ # local_head
+ context_level: int = 1,
+ num_tx_output_channels: int = 1024,
+ crop_size: int = 4,
+ sample_offset: float = 0.5,
+ num_tx_channels: int = 128,
+ num_tx_layers: int = 3,
+ num_tx_heads: int = 3,
+ use_bias: bool = True,
+ tx_activation: str = 'gelu',
+ dropout_rate: float = 0.0,
+ layer_norm_epsilon: float = 1e-6,
+ use_positional_embedding: bool = True,
+ normalize_local_features: bool = True,
+ **kwargs):
+ """Video Classification initialization function.
+
+ Args:
+ backbone: a 3d backbone network.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
+ num_hidden_layers: the number of hidden layers in the MLP.
+ num_hidden_channels: the number of hidden nodes in the MLP.
+ num_output_channels: the number of final output nodes in the MLP.
+ use_sync_bn: whether to use sync batch norm in the MLP.
+ norm_momentum: the MLP batch norm momentum.
+ norm_epsilon: the MLP batch norm epsilon.
+ activation: the MLP activation function.
+ normalize_global_features: whether to normalize inputs to the MLP.
+ context_level: the number of context frame to use.
+ num_tx_output_channels: the number of final output channels for instance
+ reconstrcutor.
+ crop_size: the ROI aligner crop size.
+ sample_offset: the ROI aligner sample offset.
+ num_tx_channels: the Transformer decoder head channels.
+ num_tx_layers: the number of Transformer decoder layers.
+ num_tx_heads: the number of Transformer decoder heads per layer.
+ use_bias: whether to use bias in the Transformer.
+ tx_activation: the activation function to use in the Transformer.
+ dropout_rate: the dropout rate for Transformer.
+ layer_norm_epsilon: the layer norm epsilon.
+ use_positional_embedding: whether to use positional embedding.
+ normalize_local_features: whether to normalize input embeddings.
+ **kwargs: keyword arguments to be passed.
+ """
+ if not input_specs:
+ input_specs = {
+ 'image': layers.InputSpec(shape=[None, None, None, None, 3]),
+ 'instances_position': layers.InputSpec(shape=[None, None, None, 4]),
+ 'instances_mask': layers.InputSpec(shape=[None, None, None]),
+ }
+ self._self_setattr_tracking = False
+ self._config_dict = {
+ 'backbone': backbone,
+ 'num_hidden_layers': num_hidden_layers,
+ 'num_hidden_channels': num_hidden_channels,
+ 'num_output_channels': num_output_channels,
+ 'use_sync_bn': use_sync_bn,
+ 'norm_momentum': norm_momentum,
+ 'norm_epsilon': norm_epsilon,
+ 'activation': activation,
+ 'normalize_global_features': normalize_global_features,
+ 'context_level': context_level,
+ 'num_tx_output_channels': num_tx_output_channels,
+ 'crop_size': crop_size,
+ 'sample_offset': sample_offset,
+ 'num_tx_channels': num_tx_channels,
+ 'num_tx_layers': num_tx_layers,
+ 'num_tx_heads': num_tx_heads,
+ 'use_bias': use_bias,
+ 'tx_activation': tx_activation,
+ 'dropout_rate': dropout_rate,
+ 'layer_norm_epsilon': layer_norm_epsilon,
+ 'use_positional_embedding': use_positional_embedding,
+ 'normalize_local_features': normalize_local_features,
+ }
+ self._input_specs = input_specs
+ self._backbone = backbone
+
+ inputs = {
+ k: tf_keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
+ }
+ endpoints = backbone(inputs['image'])
+
+ res5 = endpoints['5']
+ res5 = tf_keras.layers.GlobalAveragePooling3D()(res5)
+ res5_1 = endpoints['5_1']
+
+ global_embeddings = simple.MLP(
+ num_hidden_layers=num_hidden_layers,
+ num_hidden_channels=num_hidden_channels,
+ num_output_channels=num_output_channels,
+ use_sync_bn=use_sync_bn,
+ norm_momentum=norm_momentum,
+ norm_epsilon=norm_epsilon,
+ activation=activation,
+ normalize_inputs=normalize_global_features)(res5)
+
+ instance_inputs = {
+ 'features': res5_1,
+ 'instances_position': inputs['instances_position'],
+ 'instances_mask': inputs['instances_mask'],
+ }
+ instances_outputs = instance_reconstructor.InstanceReconstructor(
+ context_level=context_level,
+ # parameters for projector
+ num_output_channels=num_tx_output_channels,
+ # parameters for RoiAligner
+ crop_size=crop_size,
+ sample_offset=sample_offset,
+ # parameters for TxDecoder
+ num_tx_channels=num_tx_channels,
+ num_tx_layers=num_tx_layers,
+ num_tx_heads=num_tx_heads,
+ use_bias=use_bias,
+ activation=tx_activation,
+ dropout_rate=dropout_rate,
+ layer_norm_epsilon=layer_norm_epsilon,
+ use_positional_embedding=use_positional_embedding,
+ normalize_inputs=normalize_local_features)(instance_inputs)
+
+ outputs = instances_outputs
+ outputs['global_embeddings'] = global_embeddings
+ super().__init__(inputs=inputs, outputs=outputs, **kwargs)
+
+ @property
+ def checkpoint_items(self):
+ """Returns a dictionary of items to be additionally checkpointed."""
+ return dict(backbone=self.backbone)
+
+ @property
+ def backbone(self):
+ return self._backbone
+
+ def get_config(self):
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+
+@model_factory.register_model_builder('const_cl_model')
+def build_const_cl_pretrain_model(
+ input_specs_dict: Mapping[str, tf_keras.layers.InputSpec],
+ model_config: const_cl_cfg.ConstCLModel,
+ num_classes: int,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> ConstCLModel:
+ """Builds the ConST-CL video ssl model."""
+ del num_classes
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs_dict['image'],
+ backbone_config=model_config.backbone,
+ norm_activation_config=model_config.norm_activation,
+ l2_regularizer=l2_regularizer)
+
+ # Norm layer type in the MLP head should same with backbone
+ if (model_config.norm_activation.use_sync_bn
+ != model_config.global_head.use_sync_bn):
+ raise ValueError('Should use the same batch normalization type.')
+
+ return ConstCLModel(
+ backbone=backbone,
+ input_specs=input_specs_dict,
+ # global_head
+ num_hidden_channels=model_config.global_head.num_hidden_channels,
+ num_hidden_layers=model_config.global_head.num_hidden_layers,
+ num_output_channels=model_config.global_head.num_output_channels,
+ use_sync_bn=model_config.global_head.use_sync_bn,
+ norm_momentum=model_config.global_head.norm_momentum,
+ norm_epsilon=model_config.global_head.norm_epsilon,
+ activation=model_config.global_head.activation,
+ normalize_global_features=model_config.global_head.normalize_inputs,
+ # local_head
+ context_level=model_config.local_head.context_level,
+ num_tx_output_channels=model_config.local_head.num_output_channels,
+ crop_size=model_config.local_head.crop_size,
+ sample_offset=model_config.local_head.sample_offset,
+ num_tx_channels=model_config.local_head.num_tx_channels,
+ num_tx_layers=model_config.local_head.num_tx_layers,
+ num_tx_heads=model_config.local_head.num_tx_heads,
+ use_bias=model_config.local_head.use_bias,
+ tx_activation=model_config.local_head.activation,
+ dropout_rate=model_config.local_head.dropout_rate,
+ layer_norm_epsilon=model_config.local_head.layer_norm_epsilon,
+ use_positional_embedding=model_config.local_head.use_positional_embedding,
+ normalize_local_features=model_config.local_head.normalize_inputs)
diff --git a/official/projects/const_cl/modeling/const_cl_model_test.py b/official/projects/const_cl/modeling/const_cl_model_test.py
new file mode 100644
index 00000000000..215b1467c10
--- /dev/null
+++ b/official/projects/const_cl/modeling/const_cl_model_test.py
@@ -0,0 +1,47 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for const_cl_model."""
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.configs import const_cl as const_cl_cfg
+from official.projects.const_cl.modeling import const_cl_model
+# pylint: disable=unused-import
+from official.projects.const_cl.modeling.backbones import resnet_3d
+# pylint: enable=unused-import
+
+
+class ConstClModelTest(tf.test.TestCase):
+
+ def test_build_const_cl_pretrain_model(self):
+ model_config = const_cl_cfg.ConstCLModel()
+ images_input_specs = tf_keras.layers.InputSpec(
+ shape=[None, 16, 224, 224, 4])
+ boxes_input_specs = tf_keras.layers.InputSpec(shape=[None, 16, 8, 4])
+ masks_input_specs = tf_keras.layers.InputSpec(shape=[None, 16, 8])
+
+ input_specs_dict = {
+ 'image': images_input_specs,
+ 'instances_position': boxes_input_specs,
+ 'instances_mask': masks_input_specs,
+ }
+ model = const_cl_model.build_const_cl_pretrain_model(
+ input_specs_dict=input_specs_dict,
+ model_config=model_config,
+ num_classes=500)
+ self.assertIsInstance(model, const_cl_model.ConstCLModel)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/modeling/heads/instance_reconstructor.py b/official/projects/const_cl/modeling/heads/instance_reconstructor.py
new file mode 100644
index 00000000000..6d807a249c4
--- /dev/null
+++ b/official/projects/const_cl/modeling/heads/instance_reconstructor.py
@@ -0,0 +1,261 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""The instance feature reconstructor head."""
+
+from typing import Mapping
+
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.modeling.heads import transformer_decoder
+from official.vision.modeling.layers import roi_aligner
+
+
+def _get_shape(x):
+ """Helper function to return shape of a given tensor."""
+ static = x.shape.as_list()
+ dynamic = tf.shape(x)
+ return [dynamic[i] if s is None else s for i, s in enumerate(static)]
+
+
+class InstanceReconstructor(tf_keras.layers.Layer):
+ """The SSL head for reconstructing contextualized instance representations."""
+
+ def __init__(self,
+ context_level: int = 1,
+ # parameters for projector
+ num_output_channels: int = 1024,
+ # parameters for RoiAligner
+ crop_size: int = 4,
+ sample_offset: float = 0.5,
+ # parameters for TxDecoder
+ num_tx_channels: int = 128,
+ num_tx_layers: int = 3,
+ num_tx_heads: int = 3,
+ use_bias: bool = True,
+ activation: str = 'gelu',
+ dropout_rate: float = 0.0,
+ layer_norm_epsilon: float = 1e-6,
+ use_positional_embedding: bool = True,
+ normalize_inputs: bool = True,
+ **kwargs):
+ """InstanceReconstructor SSL head initializer.
+
+ Args:
+ context_level: the number of context frame to use.
+ num_output_channels: the number of final output channels.
+ crop_size: the ROI aligner crop size.
+ sample_offset: the ROI aligner sample offset.
+ num_tx_channels: the Transformer decoder head channels.
+ num_tx_layers: the number of Transformer decoder layers.
+ num_tx_heads: the number of Transformer decoder heads per layer.
+ use_bias: whether to use bias.
+ activation: the activation function to use.
+ dropout_rate: the dropout rate.
+ layer_norm_epsilon: the layer norm epsilon.
+ use_positional_embedding: whether to use positional embedding.
+ normalize_inputs: whether to normalize input embeddings.
+ **kwargs: the kwargs.
+ """
+
+ super().__init__(**kwargs)
+ self._normalize_inputs = normalize_inputs
+ self._context_level = context_level
+ self._num_output_channels = num_output_channels
+ self._crop_size = crop_size
+ self._sample_offset = sample_offset
+ self._num_tx_channels = num_tx_channels
+ self._num_tx_layers = num_tx_layers
+ self._num_tx_heads = num_tx_heads
+ self._use_bias = use_bias
+ self._activation = activation
+ self._dropout_rate = dropout_rate
+ self._layer_norm_epsilon = layer_norm_epsilon
+ self._use_positional_embedding = use_positional_embedding
+
+ self._roi_aligner = roi_aligner.MultilevelROIAligner(
+ crop_size=crop_size,
+ sample_offset=sample_offset)
+
+ if self._use_positional_embedding:
+ self._spatial_mlp = [
+ tf_keras.layers.Dense(
+ 4, use_bias=True, activation='relu', name='spatial_mlp_l1'),
+ tf_keras.layers.Dense(
+ 8, use_bias=True, name='spatial_mlp_l2')]
+ self._temporal_mlp = [
+ tf_keras.layers.Dense(
+ 4, use_bias=True, activation='relu', name='temporal_mlp_l1'),
+ tf_keras.layers.Dense(
+ 8, use_bias=True, name='temporal_mlp_l2')]
+
+ self._attention_decoder = transformer_decoder.TransformerDecoder(
+ num_channels=num_tx_channels,
+ num_layers=num_tx_layers,
+ num_heads=num_tx_heads,
+ use_bias=use_bias,
+ activation=activation,
+ dropout_rate=dropout_rate,
+ layer_norm_epsilon=layer_norm_epsilon)
+
+ self._projection_layer = tf_keras.layers.Dense(num_output_channels)
+
+ def _get_memory_embeddings(self, inputs: tf.Tensor) -> tf.Tensor:
+ """Uniformly samples frames to construct memory embeddings."""
+ if self._context_level % 2 == 0:
+ raise ValueError('context_level should be specified as odd number.')
+
+ num_frames = tf.shape(inputs)[1]
+ keyframe_index = num_frames // 2
+ stride = num_frames // self._context_level
+ start = self._context_level // 2 * -1
+ stop = self._context_level // 2 + 1 # exclusive
+
+ memories = []
+ for idx in range(start, stop):
+ idx = idx * stride + keyframe_index
+ memories.append(inputs[:, idx, ...])
+
+ memories = tf.stack(memories, axis=1)
+ return memories
+
+ def _add_positional_embedding(self, inputs: tf.Tensor) -> tf.Tensor:
+ """Adds positional embeddings to the inputs tensor."""
+ # Compute the locations using meshgrid.
+ b, t, h, w = _get_shape(inputs)[:4]
+ mesh = tf.meshgrid(tf.range(t), tf.range(h), tf.range(w), indexing='ij')
+ position = tf.cast(
+ tf.tile(
+ tf.expand_dims(tf.stack(mesh, axis=-1), axis=0), [b, 1, 1, 1, 1]),
+ tf.float32)
+
+ # Make the positions relative to center point.
+ # The mean of all position coordinates would be the center point anyway
+ center_position = tf.reduce_mean(position, axis=[1, 2, 3], keepdims=True)
+ position -= center_position
+
+ # Apply learneable layers.
+ temporal_position = position[..., :1]
+ for mlp in self._temporal_mlp:
+ temporal_position = mlp(temporal_position)
+ spatial_position = position[..., 1:]
+ for mlp in self._spatial_mlp:
+ spatial_position = mlp(spatial_position)
+
+ return tf.concat([inputs, temporal_position, spatial_position], axis=-1)
+
+ def _keyframe_roi_pooling(self,
+ features: tf.Tensor,
+ boxes: tf.Tensor,
+ training: bool = True) -> tf.Tensor:
+ """Pools ROI features on the keyframe.
+
+ Args:
+ features: a 5D tensor in shape [B, T, H, W, C].
+ boxes: normalized box coordinates, a 4D tensor in shape [B, T', N, 4].
+ training: whether in training mode.
+
+ Returns:
+ roi_feature: pooled ROI-features in shape [B, N, C].
+ """
+ if features.shape.ndims != 5:
+ raise ValueError('Expected features is a rank-5 tensor. Got shape %s' %
+ features.shape)
+
+ keyframe_index = tf.shape(boxes)[1] // 2
+ t, h, w = _get_shape(features)[1:4]
+ roi_features = {'0': features[:, t // 2, ...]}
+ keyframe_boxes = boxes[:, keyframe_index, ...]
+ unnormalized_boxes = keyframe_boxes * tf.convert_to_tensor(
+ [h, w, h, w], keyframe_boxes.dtype)
+ # roi_features in shape [B, N, h, w, C]
+ roi_features = self._roi_aligner(
+ roi_features, unnormalized_boxes, training=training)
+
+ roi_shape = _get_shape(roi_features)
+ # Perform average_pooling on ROI-pooled features.
+ roi_features = tf.reshape(roi_features, [-1] + roi_shape[2:])
+ roi_features = tf.reduce_mean(roi_features, axis=[1, 2])
+ roi_features = tf.reshape(roi_features, roi_shape[:2] + roi_shape[-1:])
+ return roi_features
+
+ def call(self,
+ inputs: Mapping[str, tf.Tensor],
+ training: bool = False) -> Mapping[str, tf.Tensor]:
+ """Forward calls.
+
+ Args:
+ inputs: the inputs dictionary contains
+ 'features': the instance embeddings in shape [2*B, T', H, W, C].
+ 'instances_positions': the instance boxes in shape [2*B, T, N, 4].
+ 'instances_mask': the validity mask for each instance position, in
+ [2*B, T, N].
+ training: whether in training mode.
+
+ Returns:
+ the context-guided reconstructed instance representations.
+ """
+
+ dense_embeddings_raw = inputs['features']
+ instances_position = inputs['instances_position']
+ instances_mask = inputs['instances_mask']
+
+ if self._normalize_inputs:
+ dense_embeddings_raw = tf.math.l2_normalize(dense_embeddings_raw, axis=-1)
+
+ def _keyframe_temporal_pooling(inputs):
+ t = tf.shape(inputs)[1] // 2
+ return inputs[:, t:t+1, ...]
+
+ dense_embeddings = _keyframe_temporal_pooling(dense_embeddings_raw)
+ instances_position = _keyframe_temporal_pooling(instances_position)
+ instances_mask = _keyframe_temporal_pooling(instances_mask)
+ instances_mask_a, instances_mask_b = tf.split(
+ tf.squeeze(instances_mask, axis=1), num_or_size_splits=2, axis=0)
+
+ inst_embeddings = self._keyframe_roi_pooling(
+ features=dense_embeddings,
+ boxes=instances_position,
+ training=training)
+
+ inst_embeddings_a, inst_embeddings_b = tf.split(inst_embeddings, 2, axis=0)
+ memory = self._get_memory_embeddings(dense_embeddings_raw)
+
+ # Add the positional embeddings before roi_pooling and tx_decoder.
+ if self._use_positional_embedding:
+ memory = self._add_positional_embedding(memory)
+ memory_a, memory_b = tf.split(memory, 2, axis=0)
+
+ # Reconstruct inst_a2b by querying in memory_b.
+ inst_embeddings_a2b = self._attention_decoder(
+ inputs=inst_embeddings_a, memory=memory_b, training=training)
+ inst_embeddings_a2b = inst_embeddings_a2b['hidden_states'][-1]
+ inst_embeddings_a2b = self._projection_layer(
+ inst_embeddings_a2b, training=training)
+ # Reconstruct inst_b2a by querying in memory_a.
+ inst_embeddings_b2a = self._attention_decoder(
+ inputs=inst_embeddings_b, memory=memory_a, training=training)
+ inst_embeddings_b2a = inst_embeddings_b2a['hidden_states'][-1]
+ inst_embeddings_b2a = self._projection_layer(
+ inst_embeddings_b2a, training=training)
+
+ outputs = {
+ 'inst_a2b': inst_embeddings_a2b,
+ 'inst_b2a': inst_embeddings_b2a,
+ 'inst_a': inst_embeddings_a,
+ 'inst_b': inst_embeddings_b,
+ 'masks_a': instances_mask_a,
+ 'masks_b': instances_mask_b,
+ }
+ return outputs
diff --git a/official/projects/const_cl/modeling/heads/instance_reconstructor_test.py b/official/projects/const_cl/modeling/heads/instance_reconstructor_test.py
new file mode 100644
index 00000000000..1290b2706a7
--- /dev/null
+++ b/official/projects/const_cl/modeling/heads/instance_reconstructor_test.py
@@ -0,0 +1,46 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for instance_reconstructor."""
+
+import tensorflow as tf, tf_keras
+from official.projects.const_cl.modeling.heads import instance_reconstructor
+
+
+class InstanceReconstructorTest(tf.test.TestCase):
+
+ def test_instance_reconstructor_return_shapes(self):
+ decoder = instance_reconstructor.InstanceReconstructor()
+
+ inputs = {
+ 'features': tf.ones([12, 5, 7, 7, 128]),
+ 'instances_position': tf.random.uniform([12, 5, 16, 4]),
+ 'instances_mask': tf.ones([12, 5, 16], tf.bool)
+ }
+
+ outputs = decoder(inputs, training=True)
+ self.assertContainsSubset(
+ list(outputs.keys()),
+ ['inst_a2b', 'inst_b2a', 'inst_a', 'inst_b', 'masks_a', 'masks_b'])
+
+ self.assertAllEqual(outputs['inst_a2b'].shape, [6, 16, 1024])
+ self.assertAllEqual(outputs['inst_a'].shape, [6, 16, 128])
+ self.assertAllEqual(outputs['inst_b2a'].shape, [6, 16, 1024])
+ self.assertAllEqual(outputs['inst_b'].shape, [6, 16, 128])
+ self.assertAllEqual(outputs['masks_b'].shape, [6, 16])
+ self.assertAllEqual(outputs['masks_a'].shape, [6, 16])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/modeling/heads/simple.py b/official/projects/const_cl/modeling/heads/simple.py
new file mode 100644
index 00000000000..09a3d181276
--- /dev/null
+++ b/official/projects/const_cl/modeling/heads/simple.py
@@ -0,0 +1,109 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Constructs simple heads."""
+
+from typing import Any, Mapping, Optional
+
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
+
+
+class MLP(tf_keras.layers.Layer):
+ """Constructs the Multi-Layer Perceptron head."""
+
+ def __init__(self,
+ num_hidden_layers: int,
+ num_hidden_channels: int,
+ num_output_channels: int,
+ use_sync_bn: bool,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 1e-5,
+ activation: Optional[str] = None,
+ normalize_inputs: bool = False,
+ **kwargs):
+ """Multi-Layer Perceptron initialization.
+
+ Args:
+ num_hidden_layers: the number of hidden layers in the MLP.
+ num_hidden_channels: the number of hidden nodes.
+ num_output_channels: the number of final output nodes.
+ use_sync_bn: whether to use sync batch norm.
+ norm_momentum: the batch norm momentum.
+ norm_epsilon: the batch norm epsilon.
+ activation: the activation function.
+ normalize_inputs: whether to normalize inputs.
+ **kwargs: keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+
+ self._num_hidden_layers = num_hidden_layers
+ self._num_hidden_channels = num_hidden_channels
+ self._num_output_channels = num_output_channels
+ self._use_sync_bn = use_sync_bn
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._activation = activation
+ self._normalize_inputs = normalize_inputs
+
+ self._layers = []
+ # MLP hidden layers
+ for _ in range(num_hidden_layers):
+ self._layers.append(
+ tf_keras.layers.Dense(num_hidden_channels, use_bias=False))
+ if use_sync_bn:
+ self._layers.append(
+ tf_keras.layers.experimental.SyncBatchNormalization(
+ momentum=norm_momentum,
+ epsilon=norm_epsilon))
+ else:
+ self._layers.append(
+ tf_keras.layers.BatchNormalization(
+ momentum=norm_momentum,
+ epsilon=norm_epsilon))
+ if activation is not None:
+ self._layers.append(tf_utils.get_activation(activation))
+
+ # Projection head
+ self._layers.append(tf_keras.layers.Dense(num_output_channels))
+
+ def call(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ """Forward calls with N-D inputs tensor."""
+ if self._normalize_inputs:
+ inputs = tf.nn.l2_normalize(inputs, axis=-1)
+
+ for layer in self._layers:
+ if isinstance(layer, tf_keras.layers.Layer):
+ inputs = layer(inputs, training=training)
+ else: # activation
+ inputs = layer(inputs)
+ return inputs
+
+ def get_config(self) -> Mapping[str, Any]:
+ """Gets class config parameters."""
+ config_dict = {
+ 'num_hidden_layer': self._num_hidden_layer,
+ 'num_hidden_channels': self._num_hidden_channels,
+ 'num_output_channels': self._num_output_channels,
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'activation': self._activation,
+ 'normalize_inputs': self._normalize_inputs}
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]):
+ """Factory constructor from config."""
+ return cls(**config)
diff --git a/official/projects/const_cl/modeling/heads/simple_test.py b/official/projects/const_cl/modeling/heads/simple_test.py
new file mode 100644
index 00000000000..57c785b4887
--- /dev/null
+++ b/official/projects/const_cl/modeling/heads/simple_test.py
@@ -0,0 +1,41 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for simple."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.projects.const_cl.modeling.heads import simple
+
+
+class SimpleTest(tf.test.TestCase):
+
+ def test_mlp_construction(self):
+ mlp_head = simple.MLP(
+ num_hidden_layers=3,
+ num_hidden_channels=128,
+ num_output_channels=56,
+ use_sync_bn=False,
+ activation='relu')
+ inputs = tf.zeros([2, 512])
+ outputs = mlp_head(inputs, training=False)
+
+ num_params = np.sum(
+ [np.prod(v.get_shape()) for v in mlp_head.trainable_weights])
+ self.assertEqual(num_params, 106296)
+ self.assertAllEqual(outputs.shape, [2, 56])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/modeling/heads/transformer_decoder.py b/official/projects/const_cl/modeling/heads/transformer_decoder.py
new file mode 100644
index 00000000000..b8026e488b8
--- /dev/null
+++ b/official/projects/const_cl/modeling/heads/transformer_decoder.py
@@ -0,0 +1,343 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Definition for Transformer heads."""
+
+from typing import Any, Mapping, Optional, Union, List, Sequence
+from absl import logging
+
+import tensorflow as tf, tf_keras
+
+
+def _get_shape(x: tf.Tensor):
+ """Helper function to return shape of a given tensor."""
+ static = x.shape.as_list()
+ dynamic = tf.shape(x)
+ return [dynamic[i] if s is None else s for i, s in enumerate(static)]
+
+
+class DecoderUnit(tf_keras.layers.Layer):
+ """Constructs the decoder MHA module used in Transformer layers."""
+
+ def __init__(self,
+ num_channels: int,
+ use_bias: bool,
+ dropout_rate: float,
+ activation: str,
+ layer_norm_epsilon: float,
+ **kwargs):
+
+ super().__init__(**kwargs)
+ self._num_channels = num_channels
+ self._use_bias = use_bias
+ self._dropout_rate = dropout_rate
+ self._activation = activation
+ self._layer_norm_epsilon = layer_norm_epsilon
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Builds the layer.
+
+ Args:
+ input_shape: the input shape for the keras tensor.
+ """
+ # Query, key, and value mapping.
+ self.layer_q = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ name='query')
+ self.layer_k = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ name='key')
+ self.layer_v = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ name='value')
+
+ self.dropout = tf_keras.layers.Dropout(self._dropout_rate)
+ # Note here is a different behavior for contrib_layers.layer_norm and
+ # tf_keras.layers.LayerNormalization, where by default, the former
+ # calculates mean/variance across all axes except the first one
+ # (batch axis), while the latter one computes statistics only on the last
+ # axis.
+ self.layer_norm = tf_keras.layers.LayerNormalization(
+ epsilon=self._layer_norm_epsilon,
+ name='layer_norm')
+
+ self.ffn1 = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=self._activation,
+ name='ffn1')
+ self.ffn2 = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ name='ffn2')
+
+ super().build(input_shape)
+
+ def call(self,
+ query: tf.Tensor,
+ memory: Optional[tf.Tensor],
+ training: bool = False) -> Mapping[str, tf.Tensor]:
+ """Forward pass of the Transformer decoder unit.
+
+ Args:
+ query: the input query tensor.
+ memory: the input memory tensor for key/value pairs. If None,
+ self-attention will be performed.
+ training: whether in training mode.
+
+ Returns:
+ outputs: the output dictionary contains 'hidden_states' and
+ 'attention weights' matrix.
+ """
+ if memory is None:
+ memory = query
+
+ tensor_q = self.layer_q(query) # (bs, qlen, inner_dim)
+ tensor_k = self.layer_k(memory) # (bs, klen, inner_dim)
+ tensor_v = self.layer_v(memory) # (bs, klen, inner_dim)
+
+ scores = tf.matmul(tensor_q, tensor_k, transpose_b=True)
+ # Scales attention_scores.
+ dk = tf.cast(_get_shape(tensor_k)[-1], dtype=scores.dtype)
+ scores = scores / tf.math.sqrt(dk)
+
+ # Shape: (bs, seq_len, seq_len)
+ attention_weights = tf.nn.softmax(scores, axis=-1)
+ # Shape: (bs, seq_len, dim_per_head)
+ attention_features = tf.matmul(attention_weights, tensor_v)
+ # Shape: (bs, seq_len, seq_len)
+ attention_features = self.dropout(attention_features, training=training)
+
+ hidden_states = attention_features + tensor_q
+ hidden_states = self.layer_norm(hidden_states)
+
+ # Shape: (bs, seq_len, out_dim)
+ hidden_states = self.ffn1(hidden_states)
+ hidden_states = self.ffn2(hidden_states)
+
+ outputs = {
+ 'hidden_states': hidden_states,
+ 'attention_weights': attention_weights,
+ }
+ return outputs
+
+ def get_config(self) -> Mapping[str, Any]:
+ """Gets class config parameters."""
+ config_dict = {
+ 'num_channels': self._num_channels,
+ 'use_bias': self._use_bias,
+ 'dropout_rate': self._dropout_rate,
+ 'activation': self._activation,
+ 'layer_norm_epsilon': self._layer_norm_epsilon,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]):
+ """Factory constructor from config."""
+ return cls(**config)
+
+
+class TransformerDecoderLayer(tf_keras.layers.Layer):
+ """Constructs the main Transformer decoder module which includes MHA + FFN."""
+
+ def __init__(self,
+ num_channels: int,
+ num_heads: int,
+ use_bias: bool,
+ activation: str,
+ dropout_rate: float,
+ layer_norm_epsilon: float,
+ name: str = 'decoder_layer',
+ **kwargs):
+ super().__init__(name=name)
+
+ self._num_channels = num_channels
+ self._num_heads = num_heads
+ self._use_bias = use_bias
+ self._activation = activation
+ self._dropout_rate = dropout_rate
+ self._layer_norm_epsilon = layer_norm_epsilon
+ self._name = name
+
+ self._mha_units = []
+ for i in range(num_heads):
+ self._mha_units.append(
+ DecoderUnit(
+ num_channels=num_channels,
+ use_bias=use_bias,
+ dropout_rate=dropout_rate,
+ activation=activation,
+ layer_norm_epsilon=layer_norm_epsilon,
+ name='mha_{}'.format(i)))
+
+ def call(
+ self,
+ inputs: tf.Tensor,
+ memory: Optional[tf.Tensor] = None,
+ training: bool = False
+ ) -> Mapping[str, Union[tf.Tensor, Sequence[tf.Tensor]]]:
+ """Forward pass of the Transformer decoder layer.
+
+ Args:
+ inputs: the input query tensor.
+ memory: the input memory tensor for key/value pairs. If None,
+ self-attention will be performed.
+ training: whether in training mode.
+
+ Returns:
+ outputs: the output dictionary contains 'hidden_states' and
+ 'attention weights' matrix.
+ """
+
+ if memory is None:
+ logging.info('No memory tokens are provided. Performing self-attention '
+ 'on input tokens in TransfomerDecoder.')
+
+ all_head_feats = []
+ all_head_attentions = []
+ for i in range(self._num_heads):
+ outputs = self._mha_units[i](
+ query=inputs, memory=memory, training=training)
+ all_head_feats.append(outputs['hidden_states'])
+ all_head_attentions.append(outputs['attention_weights'])
+
+ outputs = {
+ 'hidden_states': tf.concat(all_head_feats, axis=-1),
+ 'attention_weights': all_head_attentions,
+ }
+ return outputs
+
+ def get_config(self) -> Mapping[str, Any]:
+ """Gets class config parameters."""
+ config_dict = {
+ 'num_channels': self._num_channels,
+ 'num_heads': self._num_heads,
+ 'use_bias': self._use_bias,
+ 'activation': self._activation,
+ 'dropout_rate': self._dropout_rate,
+ 'layer_norm_epsilon': self._layer_norm_epsilon,
+ 'name': self._name,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]):
+ """Factory constructor from config."""
+ return cls(**config)
+
+
+class TransformerDecoder(tf_keras.layers.Layer):
+ """Constructs the final Transformer decoder stack."""
+
+ def __init__(self,
+ num_channels: int,
+ num_layers: int,
+ num_heads: int,
+ use_bias: bool,
+ activation: str,
+ dropout_rate: float,
+ layer_norm_epsilon: float,
+ name: str = 'transformer_decoder',
+ **kwargs):
+ super().__init__(name=name)
+
+ self._num_channels = num_channels
+ self._num_layers = num_layers
+ self._num_heads = num_heads
+ self._use_bias = use_bias
+ self._activation = activation
+ self._dropout_rate = dropout_rate
+ self._layer_norm_epsilon = layer_norm_epsilon
+
+ self._layers = []
+ for n in range(self._num_layers):
+ self._layers.append(
+ TransformerDecoderLayer(
+ num_channels=num_channels,
+ num_heads=num_heads,
+ use_bias=use_bias,
+ activation=activation,
+ dropout_rate=dropout_rate,
+ layer_norm_epsilon=layer_norm_epsilon,
+ name='layer_{}'.format(n)))
+
+ def call(self,
+ inputs: tf.Tensor,
+ memory: Optional[tf.Tensor] = None,
+ training: bool = False) -> Mapping[str, Sequence[tf.Tensor]]:
+ """Forward pass of the Transformer decoder.
+
+ Args:
+ inputs: the input query tensor.
+ memory: the input memory tensor for key/value pairs. If None,
+ self-attention will be performed.
+ training: whether in training mode.
+
+ Returns:
+ outputs: the output dictionary contains 'hidden_states' and
+ 'attention weights' matrix.
+ """
+
+ all_hidden_states = ()
+ all_attentions = ()
+
+ memory_shape = _get_shape(memory) # pyrefly: ignore[bad-argument-type]
+ memory = tf.reshape(memory, [memory_shape[0], -1, memory_shape[-1]])
+ hidden_states = inputs
+
+ for layer in self._layers:
+ layer_outputs = layer(inputs=hidden_states,
+ memory=memory,
+ training=training)
+
+ # layer_outputs is a dictionary with the following keys:
+ # hidden_states, self_attention_weights
+ hidden_states = layer_outputs['hidden_states']
+ all_attentions += (layer_outputs['attention_weights'],)
+
+ # Add last layer
+ all_hidden_states += (hidden_states,)
+
+ outputs = {
+ 'hidden_states': all_hidden_states,
+ 'attention_weights': all_attentions,
+ }
+
+ return outputs
+
+ def get_config(self) -> Mapping[str, Any]:
+ """Gets class config parameters."""
+ config_dict = {
+ 'num_channels': self._num_channels,
+ 'num_layers': self._num_layers,
+ 'num_heads': self._num_heads,
+ 'use_bias': self._use_bias,
+ 'activation': self._activation,
+ 'dropout_rate': self._dropout_rate,
+ 'layer_norm_epsilon': self._layer_norm_epsilon,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]):
+ """Factory constructor from config."""
+ return cls(**config)
diff --git a/official/projects/const_cl/modeling/heads/transformer_decoder_test.py b/official/projects/const_cl/modeling/heads/transformer_decoder_test.py
new file mode 100644
index 00000000000..e29791fec03
--- /dev/null
+++ b/official/projects/const_cl/modeling/heads/transformer_decoder_test.py
@@ -0,0 +1,126 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for TransformerDecoder."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.const_cl.modeling.heads import transformer_decoder
+
+
+class TransformerTest(tf.test.TestCase):
+
+ def test_decoder_unit_return_shape(self):
+ decoder_unit = transformer_decoder.DecoderUnit(
+ num_channels=128,
+ use_bias=True,
+ dropout_rate=0.5,
+ activation='relu',
+ layer_norm_epsilon=1e-7)
+ batch_size = 16
+ num_inputs = 128
+ num_channels = 256
+ input_tensor = tf.zeros([batch_size, num_inputs, num_channels])
+ memory_tensor = tf.ones([batch_size, num_inputs * 4, num_channels])
+ outputs = decoder_unit(input_tensor, memory_tensor, training=False)
+ self.assertAllEqual(outputs['hidden_states'].shape,
+ [batch_size, num_inputs, num_inputs])
+
+ self.assertAllEqual(outputs['attention_weights'].shape,
+ [batch_size, num_inputs, 4 * num_inputs])
+
+ def test_decoder_unit_serialize_deserialize(self):
+ decoder_unit = transformer_decoder.DecoderUnit(
+ num_channels=128,
+ use_bias=True,
+ dropout_rate=0.5,
+ activation='relu',
+ layer_norm_epsilon=1e-7)
+ config = decoder_unit.get_config()
+ new_decoder_unit = (
+ transformer_decoder.DecoderUnit.from_config(config))
+ self.assertAllEqual(
+ decoder_unit.get_config(), new_decoder_unit.get_config())
+
+ def test_decoder_layer_return_shape(self):
+ decoder_layer = transformer_decoder.TransformerDecoderLayer(
+ num_channels=128,
+ num_heads=3,
+ use_bias=True,
+ dropout_rate=0.5,
+ activation='relu',
+ layer_norm_epsilon=1e-7)
+ batch_size = 16
+ num_inputs = 128
+ num_channels = 256
+ input_tensor = tf.zeros([batch_size, num_inputs, num_channels])
+ memory_tensor = tf.ones([batch_size, num_inputs * 4, num_channels])
+ outputs = decoder_layer(input_tensor, memory_tensor, training=False)
+ self.assertAllEqual(outputs['hidden_states'].shape,
+ [batch_size, num_inputs, num_inputs * 3])
+
+ self.assertAllEqual(outputs['attention_weights'][-1].shape,
+ [batch_size, num_inputs, 4 * num_inputs])
+
+ def test_decoder_layer_serialize_deserialize(self):
+ decoder_layer = transformer_decoder.TransformerDecoderLayer(
+ num_channels=128,
+ num_heads=3,
+ use_bias=True,
+ dropout_rate=0.5,
+ activation='relu',
+ layer_norm_epsilon=1e-7)
+ config = decoder_layer.get_config()
+ new_decoder_layer = (
+ transformer_decoder.TransformerDecoderLayer.from_config(config))
+ self.assertAllEqual(
+ decoder_layer.get_config(), new_decoder_layer.get_config())
+
+ def test_decoder_return_shape(self):
+ decoder = transformer_decoder.TransformerDecoder(
+ num_channels=128,
+ num_layers=5,
+ num_heads=3,
+ use_bias=True,
+ dropout_rate=0.5,
+ activation='relu',
+ layer_norm_epsilon=1e-7)
+ batch_size = 16
+ num_inputs = 128
+ num_channels = 256
+ input_tensor = tf.zeros([batch_size, num_inputs, num_channels])
+ memory_tensor = tf.ones([batch_size, num_inputs * 4, num_channels])
+ outputs = decoder(input_tensor, memory_tensor, training=False)
+ self.assertLen(outputs['attention_weights'], 5)
+ self.assertAllEqual(outputs['hidden_states'][-1].shape,
+ [batch_size, num_inputs, num_inputs * 3])
+ self.assertAllEqual(outputs['attention_weights'][-1][-1].shape,
+ [batch_size, num_inputs, 4 * num_inputs])
+
+ def test_decoder_serialize_deserialize(self):
+ decoder = transformer_decoder.TransformerDecoder(
+ num_channels=128,
+ num_layers=5,
+ num_heads=3,
+ use_bias=True,
+ dropout_rate=0.5,
+ activation='relu',
+ layer_norm_epsilon=1e-7)
+ config = decoder.get_config()
+ new_decoder = transformer_decoder.TransformerDecoder.from_config(config)
+ self.assertAllEqual(
+ decoder.get_config(), new_decoder.get_config())
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/tasks/const_cl.py b/official/projects/const_cl/tasks/const_cl.py
new file mode 100644
index 00000000000..9a1114157f4
--- /dev/null
+++ b/official/projects/const_cl/tasks/const_cl.py
@@ -0,0 +1,224 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Video ssl pretrain task definition."""
+from typing import Any, Optional
+
+from absl import logging
+import tensorflow as tf, tf_keras
+from official.core import input_reader
+from official.core import task_factory
+from official.projects.const_cl.configs import const_cl as exp_cfg
+from official.projects.const_cl.datasets import video_ssl_inputs
+from official.projects.const_cl.losses import losses
+from official.projects.video_ssl.tasks import pretrain as video_ssl_pretrain
+from official.vision.modeling import factory_3d
+
+
+@task_factory.register_task_cls(exp_cfg.ConstCLPretrainTask)
+class ConstCLPretrainTask(video_ssl_pretrain.VideoSSLPretrainTask):
+ """A task for video contextualized ssl pretraining."""
+
+ def build_model(self):
+ """Builds video ssl pretraining model."""
+ common_input_shape = [
+ d1 if d1 == d2 else None
+ for d1, d2 in zip(self.task_config.train_data.feature_shape,
+ self.task_config.validation_data.feature_shape)
+ ]
+
+ num_frames = common_input_shape[0]
+ num_instances = self.task_config.train_data.num_instances
+ input_specs_dict = {
+ 'image':
+ tf_keras.layers.InputSpec(shape=[None] + common_input_shape),
+ 'instances_position':
+ tf_keras.layers.InputSpec(
+ shape=[None, num_frames, num_instances, 4]),
+ 'instances_mask':
+ tf_keras.layers.InputSpec(shape=[None, num_frames, num_instances]),
+ }
+
+ logging.info('Build model input %r', common_input_shape)
+
+ model = factory_3d.build_model(
+ self.task_config.model.model_type,
+ input_specs=input_specs_dict,
+ model_config=self.task_config.model,
+ num_classes=self.task_config.train_data.num_classes)
+ return model
+
+ def build_inputs(self,
+ params: exp_cfg.DataConfig,
+ input_context: Optional[Any] = None) -> tf.data.Dataset:
+ """Builds ConST-CL SSL input."""
+
+ parser = video_ssl_inputs.Parser(input_params=params)
+ postprocess_fn = video_ssl_inputs.PostBatchProcessor(params)
+
+ reader = input_reader.InputReader(
+ params,
+ dataset_fn=self._get_dataset_fn(params),
+ decoder_fn=self._get_decoder_fn(params),
+ parser_fn=parser.parse_fn(params.is_training),
+ postprocess_fn=postprocess_fn)
+
+ dataset = reader.read(input_context=input_context)
+ return dataset
+
+ def build_losses(self, model_outputs, num_replicas, model):
+ """Sparse categorical cross entropy loss.
+
+ Args:
+ model_outputs: Output logits of the model.
+ num_replicas: distributed replica number.
+ model: keras model for calculating weight decay.
+
+ Returns:
+ The total loss tensor.
+ """
+ all_losses = {}
+ logging_metrics = {}
+ losses_config = self.task_config.losses
+ total_loss = None
+
+ global_loss = losses.ContrastiveLoss(
+ normalize_inputs=losses_config.normalize_inputs,
+ temperature=losses_config.global_temperature)
+ local_loss = losses.InstanceContrastiveLoss(
+ normalize_inputs=losses_config.normalize_inputs,
+ temperature=losses_config.local_temperature)
+ # Compute global loss.
+ global_inputs = model_outputs['global_embeddings']
+ global_loss_dict = global_loss(inputs=global_inputs,
+ num_replicas=num_replicas)
+ # Compute local loss.
+ local_inputs = {
+ 'instances_a2b': model_outputs['inst_a2b'],
+ 'instances_b2a': model_outputs['inst_b2a'],
+ 'instances_a': model_outputs['inst_a'],
+ 'instances_b': model_outputs['inst_b'],
+ 'masks_a': model_outputs['masks_a'],
+ 'masks_b': model_outputs['masks_b'],
+ }
+ local_loss_dict = local_loss(predictions=local_inputs,
+ num_replicas=num_replicas)
+ # Compute regularization loss.
+ reg_loss = losses_config.l2_weight_decay * tf.add_n([
+ tf.nn.l2_loss(v) for v in model.trainable_variables
+ if 'kernel' in v.name])
+
+ total_loss = (global_loss_dict['loss'] * losses_config.global_weight +
+ local_loss_dict['loss'] * losses_config.local_weight +
+ reg_loss)
+ all_losses.update({
+ 'total_loss': total_loss
+ })
+ all_losses[self.loss] = total_loss
+
+ logging_metrics['regularization_loss'] = reg_loss
+ for k, v in global_loss_dict.items():
+ logging_metrics['global_loss/' + k] = v
+ for k, v in local_loss_dict.items():
+ logging_metrics['local_loss/' + k] = v
+ return all_losses, logging_metrics
+
+ def build_metrics(self, training=True):
+ """Gets streaming metrics for training/validation."""
+ metrics = [
+ tf_keras.metrics.Mean(name='regularization_loss'),
+
+ tf_keras.metrics.Mean(name='global_loss/loss'),
+ tf_keras.metrics.Mean(name='global_loss/contrastive_accuracy'),
+ tf_keras.metrics.Mean(name='global_loss/contrastive_entropy'),
+
+ tf_keras.metrics.Mean(name='local_loss/loss'),
+ tf_keras.metrics.Mean(name='local_loss/positive_similarity_mean'),
+ tf_keras.metrics.Mean(name='local_loss/positive_similarity_max'),
+ tf_keras.metrics.Mean(name='local_loss/positive_similarity_min'),
+ tf_keras.metrics.Mean(name='local_loss/negative_similarity_mean'),
+ tf_keras.metrics.Mean(name='local_loss/negative_similarity_max'),
+ tf_keras.metrics.Mean(name='local_loss/negative_similarity_min'),
+ ]
+ return metrics
+
+ def process_metrics(self, metrics, contrastive_metrics):
+ """Processes and updates metrics."""
+ for metric in metrics:
+ v = contrastive_metrics[metric.name]
+ metric.update_state(v)
+
+ def train_step(self, inputs, model, optimizer, metrics=None):
+ """Forward and backward pass.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ features, _ = inputs
+
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+ with tf.GradientTape() as tape:
+ outputs = model(features, training=True)
+ # Casting output layer as float32 is necessary when mixed_precision is
+ # mixed_float16 or mixed_bfloat16 to ensure output is casted as float32.
+ outputs = tf.nest.map_structure(
+ lambda x: tf.cast(x, tf.float32), outputs)
+
+ all_losses, contrastive_metrics = self.build_losses(
+ model_outputs=outputs, num_replicas=num_replicas,
+ model=model)
+ scaled_loss = all_losses[self.loss]
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(
+ optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scales back gradient before apply_gradients when LossScaleOptimizer is
+ # used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ logs = all_losses
+ if metrics:
+ self.process_metrics(metrics, contrastive_metrics)
+ logs.update({m.name: m.result() for m in metrics})
+ return logs
+
+ def validation_step(self, inputs, model, metrics=None):
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ raise NotImplementedError
+
+ def inference_step(self, features, model):
+ """Performs the forward step."""
+ raise NotImplementedError
diff --git a/official/projects/const_cl/tasks/const_cl_test.py b/official/projects/const_cl/tasks/const_cl_test.py
new file mode 100644
index 00000000000..4b6d85b94fa
--- /dev/null
+++ b/official/projects/const_cl/tasks/const_cl_test.py
@@ -0,0 +1,85 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for ConST-CL pretrain task definition."""
+import functools
+import os
+import random
+
+import orbit
+import tensorflow as tf, tf_keras
+
+# pylint: disable=unused-import
+from official.core import exp_factory
+from official.core import task_factory
+from official.modeling import optimization
+from official.projects.const_cl.modeling import const_cl_model
+from official.projects.const_cl.modeling.backbones import resnet_3d
+from official.projects.const_cl.tasks import const_cl
+from official.vision.dataloaders import tfexample_utils
+# pylint: enable=unused-import
+
+
+class ConstCLPretrainTaskTest(tf.test.TestCase):
+
+ def setUp(self):
+ super(ConstCLPretrainTaskTest, self).setUp()
+ data_dir = os.path.join(self.get_temp_dir(), 'data')
+ tf.io.gfile.makedirs(data_dir)
+ self._data_path = os.path.join(data_dir, 'data.tfrecord')
+ # pylint: disable=g-complex-comprehension
+ examples = [
+ tfexample_utils.make_video_test_example(
+ image_shape=(36, 36, 3),
+ audio_shape=(20, 128),
+ label=random.randint(0, 100)) for _ in range(2)
+ ]
+ # pylint: enable=g-complex-comprehension
+ tfexample_utils.dump_to_tfrecord(self._data_path, tf_examples=examples)
+
+ def test_task(self):
+ config = exp_factory.get_exp_config('const_cl_pretrain_kinetics400')
+ config.task.train_data.global_batch_size = 2
+ config.task.train_data.input_path = self._data_path
+
+ task = const_cl.ConstCLPretrainTask(
+ config.task)
+ model = task.build_model()
+ metrics = task.build_metrics()
+ strategy = tf.distribute.get_strategy()
+
+ dataset = orbit.utils.make_distributed_dataset(
+ strategy,
+ functools.partial(task.build_inputs),
+ config.task.train_data)
+
+ iterator = iter(dataset)
+ opt_factory = optimization.OptimizerFactory(config.trainer.optimizer_config)
+ optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
+ logs = task.train_step(next(iterator), model, optimizer, metrics=metrics)
+ self.assertIn('total_loss', logs)
+ self.assertIn('regularization_loss', logs)
+ self.assertIn('global_loss/loss', logs)
+ self.assertIn('global_loss/contrastive_accuracy', logs)
+ self.assertIn('global_loss/contrastive_entropy', logs)
+ self.assertIn('local_loss/loss', logs)
+
+ def test_task_factory(self):
+ config = exp_factory.get_exp_config('const_cl_pretrain_kinetics400')
+ task = task_factory.get_task(config.task)
+ self.assertIs(type(task), const_cl.ConstCLPretrainTask)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/const_cl/train.py b/official/projects/const_cl/train.py
new file mode 100644
index 00000000000..ea53d8fda04
--- /dev/null
+++ b/official/projects/const_cl/train.py
@@ -0,0 +1,77 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Training driver."""
+
+from absl import app
+from absl import flags
+import gin
+
+# pylint: disable=unused-import
+from official.common import distribute_utils
+from official.common import flags as tfm_flags
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+from official.projects.const_cl.modeling import const_cl_model
+from official.projects.const_cl.modeling.backbones import resnet_3d
+from official.projects.const_cl.tasks.google import const_cl
+# pylint: enable=unused-import
+
+
+FLAGS = flags.FLAGS
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ if 'train_and_eval' in FLAGS.mode:
+ assert (params.task.train_data.feature_shape ==
+ params.task.validation_data.feature_shape), (
+ f'train {params.task.train_data.feature_shape} != validate '
+ f'{params.task.validation_data.feature_shape}')
+
+ # Set mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
+ # dtype is float16
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu)
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir)
+
+ train_utils.save_gin_config(FLAGS.mode, model_dir)
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(main)
diff --git a/official/projects/cots_detector/README.md b/official/projects/cots_detector/README.md
new file mode 100644
index 00000000000..550382bbced
--- /dev/null
+++ b/official/projects/cots_detector/README.md
@@ -0,0 +1,32 @@
+# Crown-of-Thorns Starfish Detection Pipeline
+
+[](https://colab.research.google.com/github/tensorflow/models/blob/master/official/projects/cots_detector/crown_of_thorns_starfish_detection_pipeline.ipynb?force_crab_mode=1)
+
+This repository shows how to detect crown-of-thorns starfish (COTS) using a
+pre-trained COTS detector implemented in TensorFlow.
+
+
+
+## Description
+
+Coral reefs are some of the most diverse and important ecosystems in the world,
+however they face a number of rising threats that have resulted in massive
+global declines. In Australia, outbreaks of the coral-eating crown-of-thorns
+starfish (COTS) have been shown to cause major coral loss, with just 15 starfish
+in a hectare being able to strip a reef of 90% of its coral tissue. While COTS
+naturally exist in the Indo-Pacific, overfishing and excess run-off nutrients
+have led to massive outbreaks that are devastating already vulnerable coral
+communities.
+
+Controlling COTS populations is critical to promoting coral growth and
+resilience, so Google teamed up with Australia’s national science agency,
+[CSIRO](https://www.csiro.au/en/), to tackle this problem. We trained ML object
+detection models to help scale underwater surveys, enabling the monitoring and
+mapping out these harmful invertebrates with the ultimate goal of helping
+control teams to address and prioritize outbreaks.
+
+## Get started
+
+[Open the notebook in Colab](https://colab.research.google.com/github/tensorflow/models/blob/master/official/projects/cots_detector/crown_of_thorns_starfish_detection_pipeline.ipynb?force_crab_mode=1)
+to run the COTS detection pipeline.
diff --git a/official/projects/cots_detector/crown_of_thorns_starfish_detection_pipeline.ipynb b/official/projects/cots_detector/crown_of_thorns_starfish_detection_pipeline.ipynb
new file mode 100644
index 00000000000..c0b9fa65142
--- /dev/null
+++ b/official/projects/cots_detector/crown_of_thorns_starfish_detection_pipeline.ipynb
@@ -0,0 +1,1473 @@
+{
+ "cells": [
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "xBH8CcrkV3IU"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n",
+ "# you may not use this file except in compliance with the License.\n",
+ "# You may obtain a copy of the License at\n",
+ "#\n",
+ "# https://www.apache.org/licenses/LICENSE-2.0\n",
+ "#\n",
+ "# Unless required by applicable law or agreed to in writing, software\n",
+ "# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
+ "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
+ "# See the License for the specific language governing permissions and\n",
+ "# limitations under the License."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "9CzbXNRovpbc"
+ },
+ "source": [
+ "# Crown-of-Thorns Starfish Detection Pipeline"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Lpb0yoNjiWhw"
+ },
+ "source": [
+ "\u003ctable class=\"tfo-notebook-buttons\" align=\"left\"\u003e\n",
+ " \u003ctd\u003e\n",
+ " \u003ca target=\"_blank\" href=\"https://colab.research.google.com/github/tensorflow/models/blob/master/official/projects/cots_detector/crown_of_thorns_starfish_detection_pipeline.ipynb?force_crab_mode=1\"\u003e\u003cimg src=\"https://www.tensorflow.org/images/colab_logo_32px.png\" /\u003eRun in Google Colab\u003c/a\u003e\n",
+ " \u003c/td\u003e\n",
+ " \u003ctd\u003e\n",
+ " \u003ca target=\"_blank\" href=\"https://github.com/tensorflow/models/blob/master/official/projects/cots_detector/crown_of_thorns_starfish_detection_pipeline.ipynb\"\u003e\u003cimg src=\"https://www.tensorflow.org/images/GitHub-Mark-32px.png\" /\u003eView on GitHub\u003c/a\u003e\n",
+ " \u003c/td\u003e\n",
+ "\u003c/table\u003e"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "GUQ1x137ysLD"
+ },
+ "source": [
+ "Coral reefs are some of the most diverse and important ecosystems in the world , however they face a number of rising threats that have resulted in massive global declines. In Australia, outbreaks of the coral-eating crown-of-thorns starfish (COTS) have been shown to cause major coral loss, with just 15 starfish in a hectare being able to strip a reef of 90% of its coral tissue. While COTS naturally exist in the Indo-Pacific, overfishing and excess run-off nutrients have led to massive outbreaks that are devastating already vulnerable coral communities.\n",
+ "\n",
+ "Controlling COTS populations is critical to promoting coral growth and resilience, so Google teamed up with Australia’s national science agency, [CSIRO](https://www.csiro.au/en/), to tackle this problem. We trained ML object detection models to help scale underwater surveys, enabling the monitoring and mapping out these harmful invertebrates with the ultimate goal of helping control teams to address and prioritize outbreaks."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "jDiIX2xawkJw"
+ },
+ "source": [
+ "## About this notebook\n",
+ "\n",
+ "This notebook tutorial shows how to detect COTS using a pre-trained COTS detector implemented in TensorFlow. On top of just running the model on each frame of the video, the tracking code in this notebook aligns detections from frame to frame creating a consistent track for each COTS. Each track is given an id and frame count. Here is an example image from a video of a reef showing labeled COTS starfish.\n",
+ "\n",
+ "\u003cimg src=\"https://storage.googleapis.com/download.tensorflow.org/data/cots_detection/COTS_detected_sample.png\"\u003e"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "YxCF1t-Skag8"
+ },
+ "source": [
+ "It is recommended to enable GPU to accelerate the inference. On CPU, this runs for about 40 minutes, but on GPU it takes only 10 minutes. (In Colab it should already be set to GPU in the Runtime menu: *Runtime \u003e Change runtime type \u003e Hardware accelerator \u003e select \"GPU\"*)."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "a4R2T97u442o"
+ },
+ "source": [
+ "## Setup \n",
+ "\n",
+ "Install all needed packages."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "5Gs7XvCGlwlj"
+ },
+ "outputs": [],
+ "source": [
+ "# remove the existing datascience package to avoid package conflicts in the colab environment\n",
+ "!pip3 uninstall -y datascience\n",
+ "!pip3 install -q opencv-python\n",
+ "!pip3 install PILLOW"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "w-UQ87240x5R"
+ },
+ "outputs": [],
+ "source": [
+ "# Imports\n",
+ "import base64\n",
+ "import copy\n",
+ "import dataclasses\n",
+ "import glob\n",
+ "import logging\n",
+ "import mimetypes\n",
+ "import os\n",
+ "import pathlib\n",
+ "import subprocess\n",
+ "import time\n",
+ "import textwrap\n",
+ "from typing import Dict, Iterable, List, Optional, Tuple\n",
+ "\n",
+ "from absl import logging as absl_logging\n",
+ "from IPython import display\n",
+ "import cv2\n",
+ "import matplotlib.pyplot as plt\n",
+ "import numpy as np\n",
+ "import PIL.Image\n",
+ "import tensorflow as tf\n",
+ "from tqdm import tqdm"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "gsSclJg4sJbX"
+ },
+ "source": [
+ "Define all needed variables."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iKMCvnZEXBBT"
+ },
+ "outputs": [],
+ "source": [
+ "model_name = \"cots_1080_v1\" #@param [\"cots_1080_v1\", \"cots_720_v1\"]\n",
+ "test_sequence_name = \"test3\" #@param [\"test1\", \"test2\", \"test3\"]"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ORLJSdLq4-gd"
+ },
+ "outputs": [],
+ "source": [
+ "cots_model = f\"https://storage.googleapis.com/download.tensorflow.org/models/cots_detection/{model_name}.zip\"\n",
+ "\n",
+ "# Alternatively, this dataset can be downloaded through CSIRO's Data Access Portal at https://data.csiro.au/collection/csiro:54830v2\n",
+ "sample_data_link = f\"https://storage.googleapis.com/download.tensorflow.org/data/cots_detection/sample_images.zip\"\n",
+ "\n",
+ "preview_video_path = \"preview.mp4\"\n",
+ "detection_small_video_path = \"COTS_detection.mp4\"\n",
+ "detection_csv_path = \"detections.csv\""
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "FNwP3s-5xgaF"
+ },
+ "source": [
+ "You also need to retrieve the sample data. This sample data is made up of a series of chronological images."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "DF_c_ZMXdPRN"
+ },
+ "outputs": [],
+ "source": [
+ "sample_data_path = tf.keras.utils.get_file(origin=sample_data_link)\n",
+ "# Unzip data\n",
+ "!mkdir sample_images\n",
+ "!unzip -o -q {sample_data_path} -d sample_images"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Ghf-4E5-ZiJn"
+ },
+ "source": [
+ "Convert the images to a video file:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "kCdWsbO1afIJ"
+ },
+ "outputs": [],
+ "source": [
+ "tmp_video_path = \"tmp_preview.mp4\"\n",
+ "\n",
+ "filenames = sorted(glob.glob(f\"sample_images/{test_sequence_name}/*.jpg\"))\n",
+ "img = cv2.imread(filenames[0])\n",
+ "height, width, layers = img.shape\n",
+ "size = (width, height)\n",
+ "\n",
+ "video_writer = cv2.VideoWriter(\n",
+ " filename=tmp_video_path,\n",
+ " fourcc=cv2.VideoWriter_fourcc(*\"MP4V\"), \n",
+ " fps=15, \n",
+ " frameSize=size)\n",
+ " \n",
+ "for filename in tqdm(filenames):\n",
+ " img = cv2.imread(filename)\n",
+ " video_writer.write(img)\n",
+ "cv2.destroyAllWindows()\n",
+ "video_writer.release()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cHsKpPyviWmF"
+ },
+ "source": [
+ "Re-encode the video, and reduce its size (Colab crashes if you try to embed the full size video)."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "_li0qe-gh1iT"
+ },
+ "outputs": [],
+ "source": [
+ "subprocess.check_call([\n",
+ " \"ffmpeg\", \"-y\", \"-i\", tmp_video_path,\n",
+ " \"-vf\",\"scale=800:-1\",\n",
+ " \"-crf\", \"18\",\n",
+ " \"-preset\", \"veryfast\",\n",
+ " \"-vcodec\", \"libx264\", preview_video_path])"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "2ItoiHyYQGya"
+ },
+ "source": [
+ "The images you downloaded are frames of a movie showing a top view of a coral reef with crown-of-thorns starfish. Use the `base64` data-URL trick to embed the video in this notebook:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "u0fqXQUzdZCu"
+ },
+ "outputs": [],
+ "source": [
+ "def embed_video_file(path: os.PathLike) -\u003e display.HTML:\n",
+ " \"\"\"Embeds a file in the notebook as an html tag with a data-url.\"\"\"\n",
+ " path = pathlib.Path(path)\n",
+ " mime, unused_encoding = mimetypes.guess_type(str(path))\n",
+ " data = path.read_bytes()\n",
+ "\n",
+ " b64 = base64.b64encode(data).decode()\n",
+ " return display.HTML(\n",
+ " textwrap.dedent(\"\"\"\n",
+ " \u003cvideo width=\"640\" height=\"480\" controls\u003e\n",
+ " \u003csource src=\"data:{mime};base64,{b64}\" type=\"{mime}\"\u003e\n",
+ " Your browser does not support the video tag.\n",
+ " \u003c/video\u003e\n",
+ " \"\"\").format(mime=mime, b64=b64))\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "SiOsbr8xePkg"
+ },
+ "outputs": [],
+ "source": [
+ "embed_video_file(preview_video_path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "9Z0DTbWrZMZ-"
+ },
+ "source": [
+ "Can you see them? There are a lot. The goal of the model is to put boxes around all of the starfish. Each starfish will get its own ID, and that ID will be stable as the camera passes over it."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "d0iALUwM0g2p"
+ },
+ "source": [
+ "## Load the model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "fVq6vNBTxM62"
+ },
+ "source": [
+ "Download the trained COTS detection model that matches your preferences from earlier."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "No5jRA1TxXj0"
+ },
+ "outputs": [],
+ "source": [
+ "model_path = tf.keras.utils.get_file(origin=cots_model)\n",
+ "# Unzip model\n",
+ "!mkdir {model_name}\n",
+ "!unzip -o -q {model_path} -d {model_name}"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ezyuSHK5ap__"
+ },
+ "source": [
+ "Load trained model from disk and create the inference function `model_fn()`. This might take a little while."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "HXQnNjwl8Beu"
+ },
+ "outputs": [],
+ "source": [
+ "absl_logging.set_verbosity(absl_logging.ERROR)\n",
+ "\n",
+ "tf.config.optimizer.set_experimental_options({'auto_mixed_precision': True})\n",
+ "tf.config.optimizer.set_jit(True)\n",
+ "\n",
+ "model_fn = tf.saved_model.load(model_name).signatures['serving_default']"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "OvLuznhUa7uG"
+ },
+ "source": [
+ "Here's one test image. How many COTS can you see?"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XmQF_2L_a7Hu"
+ },
+ "outputs": [],
+ "source": [
+ "example_frame_number = 52\n",
+ "image = tf.io.read_file(filenames[example_frame_number])\n",
+ "image = tf.io.decode_jpeg(image)\n",
+ "\n",
+ "# Caution PIL and tf use \"RGB\" color order, while cv2 uses \"BGR\".\n",
+ "PIL.Image.fromarray(image.numpy())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "KSOf4V8WhTHF"
+ },
+ "source": [
+ "## Raw model outputs\n",
+ "\n",
+ "Try running the model on the image. The model expects a batch of images so add an outer `batch` dimension before calling the model.\n",
+ "\n",
+ "Note: The model only runs correctly with a batch size of 1.\n",
+ "\n",
+ "The result is a dictionary with a number of fields. For all fields the first dimension of the shape is the `batch` dimension, "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iqLHo8h0c2pW"
+ },
+ "outputs": [],
+ "source": [
+ "image_batch = image[tf.newaxis, ...]\n",
+ "result = model_fn(image_batch)\n",
+ "\n",
+ "print(f\"{'image_batch':20s}- shape: {image_batch.shape}\")\n",
+ "\n",
+ "for key, value in result.items():\n",
+ " print(f\"{key:20s}- shape: {value.shape}\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0xuNoKLCjyDz"
+ },
+ "source": [
+ "The `num_detections` field gives the number of valid detections, but this is always 100. There are always 100 locations that _could_ be a COTS."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nGCDZJQvkIOL"
+ },
+ "outputs": [],
+ "source": [
+ "print('\\nnum_detections: ', result['num_detections'].numpy())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cSd7JJYqkPz7"
+ },
+ "source": [
+ "Similarly the `detection_classes` field is always `0`, since the model only detects 1 class: COTS."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JoY8bJrfkcuS"
+ },
+ "outputs": [],
+ "source": [
+ "print('detection_classes: \\n', result['detection_classes'].numpy())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "X2nVLSOokyog"
+ },
+ "source": [
+ "What actually matters here is the detection scores, indicating the quality of each detection: "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iepEgCc2jsRD"
+ },
+ "outputs": [],
+ "source": [
+ "result['detection_scores'].numpy()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Fn2B0nbplAFy"
+ },
+ "source": [
+ "You need to choose a threshold that determines what counts as a good detection. This frame has a few good detections:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "a30Uyc0WlK2a"
+ },
+ "outputs": [],
+ "source": [
+ "good_detections = result['detection_scores'] \u003e 0.4\n",
+ "good_detections.numpy()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Y_xrbQiAlWrK"
+ },
+ "source": [
+ "## Bounding boxes and detections\n",
+ "\n",
+ "Build a class to handle the detection boxes:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "S5inzqu-4JhT"
+ },
+ "outputs": [],
+ "source": [
+ "@dataclasses.dataclass(frozen=True)\n",
+ "class BBox:\n",
+ " x0: float\n",
+ " y0: float\n",
+ " x1: float\n",
+ " y1: float\n",
+ "\n",
+ " def replace(self, **kwargs):\n",
+ " d = self.__dict__.copy()\n",
+ " d.update(kwargs)\n",
+ " return type(self)(**d)\n",
+ "\n",
+ " @property\n",
+ " def center(self)-\u003e Tuple[float, float]:\n",
+ " return ((self.x0+self.x1)/2, (self.y0+self.y1)/2)\n",
+ " \n",
+ " @property\n",
+ " def width(self) -\u003e float:\n",
+ " return self.x1 - self.x0\n",
+ "\n",
+ " @property\n",
+ " def height(self) -\u003e float:\n",
+ " return self.y1 - self.y0\n",
+ "\n",
+ " @property\n",
+ " def area(self)-\u003e float:\n",
+ " return (self.x1 - self.x0 + 1) * (self.y1 - self.y0 + 1)\n",
+ " \n",
+ " def intersection(self, other)-\u003e Optional['BBox']:\n",
+ " x0 = max(self.x0, other.x0)\n",
+ " y0 = max(self.y0, other.y0)\n",
+ " x1 = min(self.x1, other.x1)\n",
+ " y1 = min(self.y1, other.y1)\n",
+ " if x0 \u003e x1 or y0 \u003e y1:\n",
+ " return None\n",
+ " return BBox(x0, y0, x1, y1)\n",
+ "\n",
+ " def iou(self, other):\n",
+ " intersection = self.intersection(other)\n",
+ " if intersection is None:\n",
+ " return 0\n",
+ " \n",
+ " ia = intersection.area\n",
+ "\n",
+ " return ia/(self.area + other.area - ia)\n",
+ " \n",
+ " def draw(self, image, label=None, color=(0, 140, 255)):\n",
+ " image = np.asarray(image)\n",
+ " cv2.rectangle(image, \n",
+ " (int(self.x0), int(self.y0)),\n",
+ " (int(self.x1), int(self.y1)),\n",
+ " color,\n",
+ " thickness=2)\n",
+ " if label is not None:\n",
+ " cv2.putText(image, str(label), \n",
+ " (int(self.x0), int(self.y0-10)),\n",
+ " cv2.FONT_HERSHEY_SIMPLEX,\n",
+ " 0.9, color, thickness=2)\n",
+ " return image"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "2izYMR9Q6Dn0"
+ },
+ "source": [
+ "And a class to represent a `Detection`, with a method to create a list of detections from the model's output:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "tybwY3eaY803"
+ },
+ "outputs": [],
+ "source": [
+ "@dataclasses.dataclass(frozen=True)\n",
+ "class Detection:\n",
+ " \"\"\"Detection dataclass.\"\"\"\n",
+ " class_id: int\n",
+ " score: float\n",
+ " bbox: BBox\n",
+ " threshold:float = 0.4\n",
+ "\n",
+ " def replace(self, **kwargs):\n",
+ " d = self.__dict__.copy()\n",
+ " d.update(kwargs)\n",
+ " return type(self)(**d)\n",
+ "\n",
+ " @classmethod\n",
+ " def process_model_output(\n",
+ " cls, image, detections: Dict[str, tf.Tensor]\n",
+ " ) -\u003e Iterable['Detection']:\n",
+ " \n",
+ " # The model only works on a batch size of 1.\n",
+ " detection_boxes = detections['detection_boxes'].numpy()[0]\n",
+ " detection_classes = detections['detection_classes'].numpy()[0].astype(np.int32)\n",
+ " detection_scores = detections['detection_scores'].numpy()[0]\n",
+ "\n",
+ " img_h, img_w = image.shape[0:2]\n",
+ "\n",
+ " valid_indices = detection_scores \u003e= cls.threshold\n",
+ " classes = detection_classes[valid_indices]\n",
+ " scores = detection_scores[valid_indices]\n",
+ " boxes = detection_boxes[valid_indices, :]\n",
+ " detections = []\n",
+ "\n",
+ " for class_id, score, box in zip(classes, scores, boxes):\n",
+ " detections.append(\n",
+ " Detection(\n",
+ " class_id=class_id,\n",
+ " score=score,\n",
+ " bbox=BBox(\n",
+ " x0=box[1] * img_w,\n",
+ " y0=box[0] * img_h,\n",
+ " x1=box[3] * img_w,\n",
+ " y1=box[2] * img_h,)))\n",
+ "\n",
+ " return detections"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "QRZ9Q5meHl84"
+ },
+ "source": [
+ "## Preview some detections\n",
+ "\n",
+ "Now you can preview the model's output:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Px7AoFCn-psx"
+ },
+ "outputs": [],
+ "source": [
+ "detections = Detection.process_model_output(image, result)\n",
+ "\n",
+ "for n, det in enumerate(detections):\n",
+ " det.bbox.draw(image, label=n+1, color=(255, 140, 0))\n",
+ "\n",
+ "PIL.Image.fromarray(image.numpy())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "B1q_n1xJLm60"
+ },
+ "source": [
+ "That works well for one frame, but to count the number of COTS in a video you'll need to track the detections from frame to frame. The raw detection indices are not stable, they're just sorted by the detection score. Below both sets of detections are overlaid on the second image with the first frame's detections in white and the second frame's in orange, the indices are not aligned. The positions are shifted because of camera motion between the two frames:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "PLtxJFPuLma0"
+ },
+ "outputs": [],
+ "source": [
+ "image2 = tf.io.read_file(filenames[example_frame_number+5]) # five frames later\n",
+ "image2 = tf.io.decode_jpeg(image2)\n",
+ "result2 = model_fn(image2[tf.newaxis, ...])\n",
+ "detections2 = Detection.process_model_output(image2, result2)\n",
+ "\n",
+ "for n, det in enumerate(detections):\n",
+ " det.bbox.draw(image2, label=n+1, color=(255, 255, 255))\n",
+ "\n",
+ "for n, det in enumerate(detections2):\n",
+ " det.bbox.draw(image2, label=n+1, color=(255, 140, 0))\n",
+ "\n",
+ "PIL.Image.fromarray(image2.numpy())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "CoRxLon5MZ35"
+ },
+ "source": [
+ "## Use optical flow to align detections\n",
+ "\n",
+ "The two sets of bounding boxes above don't line up because of camera movement. \n",
+ "To see in more detail how tracks are aligned, initialize the tracker with the first image, and then run the optical flow step, `propagate_tracks`. "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "wb_nkcPJJx2t"
+ },
+ "outputs": [],
+ "source": [
+ "def default_of_params():\n",
+ " its=20\n",
+ " eps=0.03\n",
+ " return {\n",
+ " 'winSize': (64,64),\n",
+ " 'maxLevel': 3,\n",
+ " 'criteria': (cv2.TermCriteria_COUNT + cv2.TermCriteria_EPS, its, eps)\n",
+ " }"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "mHVPymG8F2ke"
+ },
+ "outputs": [],
+ "source": [
+ "def propagate_detections(detections, image1, image2, of_params=None):\n",
+ " if of_params is None:\n",
+ " of_params = default_of_params()\n",
+ "\n",
+ " bboxes = [det.bbox for det in detections]\n",
+ " centers = np.float32([[bbox.center for bbox in bboxes]])\n",
+ " widths = np.float32([[bbox.width for bbox in bboxes]])\n",
+ " heights = np.float32([[bbox.height for bbox in bboxes]])\n",
+ "\n",
+ "\n",
+ " new_centers, status, error = cv2.calcOpticalFlowPyrLK(\n",
+ " image1, image2, centers, None, **of_params)\n",
+ "\n",
+ " x0s = new_centers[...,0] - widths/2\n",
+ " x1s = new_centers[...,0] + widths/2\n",
+ " y0s = new_centers[...,1] - heights/2\n",
+ " y1s = new_centers[...,1] + heights/2\n",
+ "\n",
+ " updated_detections = []\n",
+ " for i, det in enumerate(detections):\n",
+ " det = det.replace(\n",
+ " bbox = BBox(x0s[0,i], y0s[0,i], x1s[0,i], y1s[0,i]))\n",
+ " updated_detections.append(det)\n",
+ " return updated_detections"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "dCjgvoZnOcBu"
+ },
+ "source": [
+ "Now keep the white boxes for the initial detections, and the orange boxes for the new set of detections. But add the optical-flow propagated tracks in green. You can see that by using optical-flow to propagate the old detections to the new frame the alignment is quite good. It's this alignment between the old and new detections (between the green and orange boxes) that allows the tracker to make a persistent track for each COTS. "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "aeTny8YnHwTw"
+ },
+ "outputs": [],
+ "source": [
+ "image = tf.io.read_file(filenames[example_frame_number])\n",
+ "image = tf.io.decode_jpeg(image).numpy()\n",
+ "image_gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n",
+ "\n",
+ "image2 = tf.io.read_file(filenames[example_frame_number+5]) # five frames later\n",
+ "image2 = tf.io.decode_jpeg(image2).numpy()\n",
+ "image2_gray = cv2.cvtColor(image2, cv2.COLOR_BGR2GRAY)\n",
+ "\n",
+ "updated_detections = propagate_detections(detections, image_gray, image2_gray)\n",
+ "\n",
+ "\n",
+ "for det in detections:\n",
+ " det.bbox.draw(image2, color=(255, 255, 255))\n",
+ "\n",
+ "for det in updated_detections:\n",
+ " det.bbox.draw(image2, color=(0, 255, 0))\n",
+ "\n",
+ "for det in detections2:\n",
+ " det.bbox.draw(image2, color=(255, 140, 0))\n",
+ "\n",
+ "PIL.Image.fromarray(image2)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "jbZ-7ICCENWG"
+ },
+ "source": [
+ "## Define **OpticalFlowTracker** class\n",
+ "\n",
+ "These help track the movement of each COTS object across the video frames.\n",
+ "\n",
+ "The tracker collects related detections into `Track` objects. \n",
+ "\n",
+ "The class's init is defined below, it's methods are defined in the following cells.\n",
+ "\n",
+ "The `__init__` method just initializes the track counter (`track_id`), and sets some default values for the tracking and optical flow configurations. "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "3j2Ka1uGEoz4"
+ },
+ "outputs": [],
+ "source": [
+ "class OpticalFlowTracker:\n",
+ " \"\"\"Optical flow tracker.\"\"\"\n",
+ "\n",
+ " @classmethod\n",
+ " def add_method(cls, fun):\n",
+ " \"\"\"Attach a new method to the class.\"\"\"\n",
+ " setattr(cls, fun.__name__, fun)\n",
+ "\n",
+ "\n",
+ " def __init__(self, tid=1, ft=3.0, iou=0.5, tt=2.0, bb=32, of_params=None):\n",
+ " # Bookkeeping for the tracks.\n",
+ " # The running track count, incremented for each new track.\n",
+ " self.track_id = tid\n",
+ " self.tracks = []\n",
+ " self.prev_image = None\n",
+ " self.prev_time = None\n",
+ "\n",
+ " # Configuration for the track cleanup logic.\n",
+ " # How long to apply optical flow tracking without getting positive \n",
+ " # detections (sec).\n",
+ " self.track_flow_time = ft * 1000\n",
+ " # Required IoU overlap to link a detection to a track.\n",
+ " self.overlap_threshold = iou\n",
+ " # Used to detect if detector needs to be reset.\n",
+ " self.time_threshold = tt * 1000\n",
+ " self.border = bb\n",
+ "\n",
+ " if of_params is None:\n",
+ " of_params = default_of_params()\n",
+ " self.of_params = of_params\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "yBLSv0Fi_JJD"
+ },
+ "source": [
+ "Internally the tracker will use small `Track` and `Tracklet` classes to organize the data. The `Tracklet` class is just a `Detection` with a timestamp, while a `Track` is a track ID, the most recent detection and a list of `Tracklet` objects forming the history of the track."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "gCQFfAkaY_WN"
+ },
+ "outputs": [],
+ "source": [
+ "@dataclasses.dataclass(frozen=True)\n",
+ "class Tracklet:\n",
+ " timestamp:float\n",
+ " detection:Detection\n",
+ "\n",
+ " def replace(self, **kwargs):\n",
+ " d = self.__dict__.copy()\n",
+ " d.update(kwargs)\n",
+ " return type(self)(**d)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "7qVW1a_YZBgL"
+ },
+ "outputs": [],
+ "source": [
+ "@dataclasses.dataclass(frozen=True)\n",
+ "class Track:\n",
+ " \"\"\"Tracker entries.\"\"\"\n",
+ " id:int\n",
+ " det: Detection\n",
+ " linked_dets:List[Tracklet] = dataclasses.field(default_factory=list)\n",
+ "\n",
+ " def replace(self, **kwargs):\n",
+ " d = self.__dict__.copy()\n",
+ " d.update(kwargs)\n",
+ " return type(self)(**d)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Ntl_4oUp_1nD"
+ },
+ "source": [
+ "The tracker keeps a list of active `Track` objects.\n",
+ "\n",
+ "The main `update` method takes an image, along with the list of detections and the timestamp for that image. On each frame step it performs the following sub-tasks:\n",
+ "\n",
+ "* The tracker uses optical flow to calculate where each `Track` expects to see a new `Detection`.\n",
+ "* The tracker matches up the actual detections for the frame to the expected detections for each Track.\n",
+ "* If a detection doesn't get matched to an existing track, a new track is created for the detection.\n",
+ "* If a track stops getting assigned new detections, it is eventually deactivated. "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "koZ0mjFTpiTv"
+ },
+ "outputs": [],
+ "source": [
+ "@OpticalFlowTracker.add_method\n",
+ "def update(self, image_bgr, detections, timestamp):\n",
+ " start = time.time()\n",
+ "\n",
+ " image = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2GRAY)\n",
+ "\n",
+ " # Remove dead tracks.\n",
+ " self.tracks = self.cleanup_tracks(image, timestamp)\n",
+ "\n",
+ " # Run optical flow to update existing tracks.\n",
+ " if self.prev_time is not None:\n",
+ " self.tracks = self.propagate_tracks(image)\n",
+ "\n",
+ " # Update the track list based on the new detections\n",
+ " self.apply_detections_to_tracks(image, detections, timestamp)\n",
+ "\n",
+ " self.prev_image = image\n",
+ " self.prev_time = timestamp\n",
+ "\n",
+ " return self.tracks"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "U-6__zF2CHFS"
+ },
+ "source": [
+ "The `cleanup_tracks` method clears tracks that are too old or are too close to the edge of the image."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "HQBj8GihjF3-"
+ },
+ "outputs": [],
+ "source": [
+ "@OpticalFlowTracker.add_method\n",
+ "def cleanup_tracks(self, image, timestamp) -\u003e List[Track]:\n",
+ " image_w = image.shape[1]\n",
+ " image_h = image.shape[0]\n",
+ "\n",
+ " # Assume tracker is invalid if too much time has passed!\n",
+ " if (self.prev_time is not None and\n",
+ " timestamp - self.prev_time \u003e self.time_threshold):\n",
+ " logging.info(\n",
+ " 'Too much time since last update, resetting tracker.')\n",
+ " return []\n",
+ "\n",
+ " # Remove tracks which are:\n",
+ " # - Touching the image edge.\n",
+ " # - Have existed for a long time without linking a real detection.\n",
+ " active_tracks = []\n",
+ " for track in self.tracks:\n",
+ " bbox = track.det.bbox\n",
+ " if (bbox.x0 \u003c self.border or bbox.y0 \u003c self.border or\n",
+ " bbox.x1 \u003e= (image_w - self.border) or\n",
+ " bbox.y1 \u003e= (image_h - self.border)):\n",
+ " logging.info(f'Removing track {track.id} because it\\'s near the border')\n",
+ " continue\n",
+ "\n",
+ " time_since_last_detection = timestamp - track.linked_dets[-1].timestamp\n",
+ " if (time_since_last_detection \u003e self.track_flow_time):\n",
+ " logging.info(f'Removing track {track.id} because it\\'s too old '\n",
+ " f'({time_since_last_detection:.02f}s)')\n",
+ " continue\n",
+ "\n",
+ " active_tracks.append(track)\n",
+ "\n",
+ " return active_tracks"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "DVzNcESxC6vY"
+ },
+ "source": [
+ "The `propagate_tracks` method uses optical flow to update each track's bounding box's position to predict their location in the new image: "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "0GycdAflCs6v"
+ },
+ "outputs": [],
+ "source": [
+ "@OpticalFlowTracker.add_method\n",
+ "def propagate_tracks(self, image):\n",
+ " if not self.tracks:\n",
+ " return self.tracks[:]\n",
+ "\n",
+ " detections = [track.det for track in self.tracks]\n",
+ " detections = propagate_detections(detections, self.prev_image, image, self.of_params)\n",
+ "\n",
+ " return [track.replace(det=det) \n",
+ " for track, det in zip(self.tracks, detections)]\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "uLbVeetwD0ph"
+ },
+ "source": [
+ "The `apply_detections_to_tracks` method compares each detection to the updated bounding box for each track. The detection is added to the track that matches best, if the match is better than the `overlap_threshold`. If no track is better than the threshold, the detection is used to create a new track. \n",
+ "\n",
+ "If a track has no new detection assigned to it the predicted detection is used."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "j6pfRhDRlApe"
+ },
+ "outputs": [],
+ "source": [
+ "@OpticalFlowTracker.add_method\n",
+ "def apply_detections_to_tracks(self, image, detections, timestamp):\n",
+ " image_w = image.shape[1]\n",
+ " image_h = image.shape[0]\n",
+ "\n",
+ " # Insert new detections.\n",
+ " detected_obj_track_ids = set()\n",
+ "\n",
+ " for detection in detections:\n",
+ " bbox = detection.bbox\n",
+ " if (bbox.x0 \u003c self.border or bbox.y0 \u003c self.border or\n",
+ " bbox.x1 \u003e= image_w - self.border or\n",
+ " bbox.y1 \u003e= image_h - self.border):\n",
+ " logging.debug('Skipping detection because it\\'s close to the border.')\n",
+ " continue\n",
+ "\n",
+ " # See if detection can be linked to an existing track.\n",
+ " linked = False\n",
+ " overlap_index = 0\n",
+ " overlap_max = -1000\n",
+ " for track_index, track in enumerate(self.tracks):\n",
+ " logging.debug('Testing track %d', track_index)\n",
+ " if track.det.class_id != detection.class_id:\n",
+ " continue\n",
+ " overlap = detection.bbox.iou(track.det.bbox)\n",
+ " if overlap \u003e overlap_max:\n",
+ " overlap_index = track_index\n",
+ " overlap_max = overlap\n",
+ "\n",
+ " # Link to existing track with maximal IoU.\n",
+ " if overlap_max \u003e self.overlap_threshold:\n",
+ " track = self.tracks[overlap_index]\n",
+ " self.tracks[overlap_index] = track.replace(det=detection)\n",
+ " track.linked_dets.append(Tracklet(timestamp, detection))\n",
+ " detected_obj_track_ids.add(track.id)\n",
+ " linked = True\n",
+ "\n",
+ " if not linked:\n",
+ " logging.info(f'Creating new track with ID {self.track_id}')\n",
+ " new_track = Track(self.track_id, detection)\n",
+ " new_track.linked_dets.append(Tracklet(timestamp, detection))\n",
+ " detected_obj_track_ids.add(self.track_id)\n",
+ " self.tracks.append(new_track)\n",
+ " self.track_id += 1\n",
+ "\n",
+ " for track in self.tracks:\n",
+ " # If the detector does not find the obj but estimated in the tracker, \n",
+ " # add the estimated one to that tracker's linked_dets\n",
+ " if track.id not in detected_obj_track_ids:\n",
+ " track.linked_dets.append(Tracklet(timestamp, track.det))"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "gY0AH-KUHPlC"
+ },
+ "source": [
+ "## Test run the tracker\n",
+ "\n",
+ "So reload the test images, and run the detections to test out the tracker.\n",
+ "\n",
+ "On the first frame it creates and returns one track per detection:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "7Ekkj_XFGdfq"
+ },
+ "outputs": [],
+ "source": [
+ "example_frame_number = 52\n",
+ "image = tf.io.read_file(filenames[example_frame_number])\n",
+ "image = tf.io.decode_jpeg(image)\n",
+ "result = model_fn(image[tf.newaxis, ...])\n",
+ "detections = Detection.process_model_output(image, result)\n",
+ "\n",
+ "tracker = OpticalFlowTracker()\n",
+ "tracks = tracker.update(image.numpy(), detections, timestamp = 0)\n",
+ "\n",
+ "print(f'detections : {len(detections)}') \n",
+ "print(f'tracks : {len(tracks)}')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "WovDYdNMII-n"
+ },
+ "source": [
+ "On the second frame many of the detections get assigned to existing tracks:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "7iFEKwgMGi5n"
+ },
+ "outputs": [],
+ "source": [
+ "image2 = tf.io.read_file(filenames[example_frame_number+5]) # five frames later\n",
+ "image2 = tf.io.decode_jpeg(image2)\n",
+ "result2 = model_fn(image2[tf.newaxis, ...])\n",
+ "detections2 = Detection.process_model_output(image2, result2)\n",
+ "\n",
+ "new_tracks = tracker.update(image2.numpy(), detections2, timestamp = 1000)\n",
+ "\n",
+ "print(f'detections : {len(detections2)}') \n",
+ "print(f'tracks : {len(new_tracks)}')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "dbkedwiVrxnQ"
+ },
+ "source": [
+ "Now the track IDs should be consistent:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "QexJR5gerw6q"
+ },
+ "outputs": [],
+ "source": [
+ "test_img = image2.numpy()\n",
+ "for n,track in enumerate(tracks):\n",
+ " track.det.bbox.draw(test_img, label=n, color=(255, 255, 255))\n",
+ "\n",
+ "for n,track in enumerate(new_tracks):\n",
+ " track.det.bbox.draw(test_img, label=n, color=(255, 140, 0))\n",
+ "\n",
+ "PIL.Image.fromarray(test_img)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "OW5gGixy1osE"
+ },
+ "source": [
+ "## Perform the COTS detection inference and tracking."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "f21596933d08"
+ },
+ "source": [
+ "The main tracking loop will perform the following: \n",
+ "\n",
+ "1. Load the images in order.\n",
+ "2. Run the model on the image.\n",
+ "3. Update the tracker with the new images and detections.\n",
+ "4. Keep information about each track (id, current index and length) analysis or display. \n",
+ "\n",
+ "The `TrackAnnotation` class, below, will collect the data about each track:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "lESJE0qXxubm"
+ },
+ "outputs": [],
+ "source": [
+ "@dataclasses.dataclass(frozen=True)\n",
+ "class TrackAnnotation:\n",
+ " det: Detection\n",
+ " seq_id: int\n",
+ " seq_idx: int\n",
+ " seq_length: Optional[int] = None\n",
+ "\n",
+ " def replace(self, **kwargs):\n",
+ " d = self.__dict__.copy()\n",
+ " d.update(kwargs)\n",
+ " return type(self)(**d)\n",
+ "\n",
+ " def annotation_str(self):\n",
+ " return f\"{self.seq_id} ({self.seq_idx}/{self.seq_length})\"\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "3863fb28cd34"
+ },
+ "source": [
+ "The `parse_image` function, below, will take `(index, filename)` pairs load the images as tensors and return `(timestamp_ms, filename, image)` triples, assuming 30fps"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Dn7efhr0GBGz"
+ },
+ "outputs": [],
+ "source": [
+ "# Read a jpg image and decode it to a uint8 tf tensor.\n",
+ "def parse_image(index, filename):\n",
+ " image = tf.io.read_file(filename)\n",
+ " image = tf.io.decode_jpeg(image)\n",
+ " timestamp_ms = 1000*index/30 # assuming 30fps\n",
+ " return (timestamp_ms, filename, image)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "8f878e4b0852"
+ },
+ "source": [
+ "Here is the main tracker loop. Note that initially the saved `TrackAnnotations` don't contain the track lengths. The lengths are collected in the `track_length_for_id` dict."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "cqN8RGBgVbr4"
+ },
+ "outputs": [],
+ "source": [
+ "# Create a tracker object\n",
+ "tracker = OpticalFlowTracker(tid=1)\n",
+ "# Record tracking responses from the tracker\n",
+ "detection_result = []\n",
+ "# Record the length of each tracking sequence\n",
+ "track_length_for_id = {}\n",
+ "\n",
+ "# Create a data loader\n",
+ "file_list = sorted(glob.glob(f\"sample_images/{test_sequence_name}/*.jpg\"))\n",
+ "list_ds = tf.data.Dataset.from_tensor_slices(file_list).enumerate()\n",
+ "images_ds = list_ds.map(parse_image)\n",
+ "\n",
+ "# Traverse the dataset with batch size = 1, you cannot change the batch size\n",
+ "for timestamp_ms, file_path, images in tqdm(images_ds.batch(1, drop_remainder=True)):\n",
+ " # get detection result\n",
+ " detections = Detection.process_model_output(images[0], model_fn(images))\n",
+ "\n",
+ " # Feed detection results and the corresponding timestamp to the tracker, and then get tracker response\n",
+ " tracks = tracker.update(images[0].numpy(), detections, timestamp_ms[0])\n",
+ " annotations = []\n",
+ " for track in tracks:\n",
+ " anno = TrackAnnotation(\n",
+ " det=track.det,\n",
+ " seq_id = track.id,\n",
+ " seq_idx = len(track.linked_dets)\n",
+ " )\n",
+ " annotations.append(anno)\n",
+ " track_length_for_id[track.id] = len(track.linked_dets)\n",
+ " \n",
+ " detection_result.append((file_path.numpy()[0].decode(), annotations))"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "29306d7f32df"
+ },
+ "source": [
+ "Once the tracking loop has completed you can update the track length (`seq_length`) for each annotation from the `track_length_for_id` dict:"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "oPSfnQ1o04Rx"
+ },
+ "outputs": [],
+ "source": [
+ "def update_annotation_lengths(detection_result, track_length_for_id):\n",
+ " new_result = []\n",
+ " for file_path, annotations in detection_result:\n",
+ " new_annotations = []\n",
+ " for anno in annotations:\n",
+ " anno = anno.replace(seq_length=track_length_for_id[anno.seq_id])\n",
+ " new_annotations.append(anno)\n",
+ " new_result.append((file_path, new_annotations))\n",
+ " return new_result"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "zda914lv1o_v"
+ },
+ "outputs": [],
+ "source": [
+ "detection_result = update_annotation_lengths(detection_result, track_length_for_id)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "QkpmYRyFAMlM"
+ },
+ "source": [
+ "## Output the detection results and play the result video\n",
+ "\n",
+ "Once the inference is done, we draw the bounding boxes and track information onto each frame's image. Finally, we combine all frames into a video for visualisation."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "gWMJG7g95MGk"
+ },
+ "outputs": [],
+ "source": [
+ "detection_full_video_path = \"COTS_detection_full_size.mp4\"\n",
+ "detect_video_writer = cv2.VideoWriter(\n",
+ " filename=detection_full_video_path,\n",
+ " fourcc=cv2.VideoWriter_fourcc(*\"MP4V\"), \n",
+ " fps=15, \n",
+ " frameSize=size)\n",
+ "\n",
+ "for file_path, annotations in tqdm(detection_result):\n",
+ " image = cv2.imread(file_path)\n",
+ " for anno in annotations:\n",
+ " anno.det.bbox.draw(image, label=anno.annotation_str(), color=(0, 140, 255))\n",
+ " detect_video_writer.write(image)\n",
+ "cv2.destroyAllWindows()\n",
+ "\n",
+ "detect_video_writer.release()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "9s1myz67jcV8"
+ },
+ "outputs": [],
+ "source": [
+ "subprocess.check_call([\n",
+ " \"ffmpeg\",\"-y\", \"-i\", detection_full_video_path,\n",
+ " \"-vf\",\"scale=800:-1\",\n",
+ " \"-crf\", \"18\",\n",
+ " \"-preset\", \"veryfast\",\n",
+ " \"-vcodec\", \"libx264\", detection_small_video_path])"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "wsK5cvX5jkL7"
+ },
+ "outputs": [],
+ "source": [
+ "embed_video_file(detection_small_video_path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "n1oOgMR2zzIl"
+ },
+ "source": [
+ "The output video is now saved as movie at `detection_full_video_path`. You can download your video by uncommenting the following code."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "tyHucK8lbGXk"
+ },
+ "outputs": [],
+ "source": [
+ "#try:\n",
+ "# from google.colab import files\n",
+ "# files.download(detection_full_video_path)\n",
+ "#except ImportError:\n",
+ "# pass"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "collapsed_sections": [],
+ "name": "crown_of_thorns_starfish_detection_pipeline.ipynb",
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/deepmac_maskrcnn/README.md b/official/projects/deepmac_maskrcnn/README.md
index a0e449f74c7..e0dc3cafa5b 100644
--- a/official/projects/deepmac_maskrcnn/README.md
+++ b/official/projects/deepmac_maskrcnn/README.md
@@ -19,14 +19,14 @@ for more details.
agnostic mode and `task.allowed_mask_class_ids` controls which classes are
allowed to have masks during training.
* Majority of experiments and ablations from the paper are perfomed with the
- [DeepMAC model](../../../../../research/object_detection/g3doc/deepmac.md)
+ [DeepMAC model](https://github.com/tensorflow/models/blob/master/research/object_detection/g3doc/deepmac.md)
in the Object Detection API code base.
## Prerequisites
### Prepare dataset
-Use [create_coco_tf_record.py](../../data/create_coco_tf_record.py) to create
+Use [create_coco_tf_record.py](https://github.com/tensorflow/models/blob/master/official/vision/data/create_coco_tf_record.py) to create
the COCO dataset. The data needs to be store in a
[Google cloud storage bucket](https://cloud.google.com/storage/docs/creating-buckets)
so that it can be accessed by the TPU.
@@ -103,11 +103,17 @@ ResNet-50 | Hourglass-52 | `deep_mask_head_rcnn_voc_r50_hg52.yaml` |
ResNet-101 | Hourglass-52 | `deep_mask_head_rcnn_voc_r101_hg52.yaml` | 34.4
SpienNet-143 | Hourglass-52 | `deep_mask_head_rcnn_voc_spinenet143_hg52.yaml` | 38.7
+## Checkpoints
+This model takes Image + boxes as input and produces per-box instance
+masks as output.
+
+* [Mask-RCNN SpineNet backbone](https://storage.googleapis.com/tf_model_garden/vision/deepmac_maskrcnn/deepmarc_spinenet.zip)
+
## See also
* [DeepMAC model](https://github.com/tensorflow/models/blob/master/research/object_detection/g3doc/deepmac.md)
in the Object Detection API code base.
-* Project website - [git.io/deepmac](https://git.io/deepmac)
+* Project website - [git.io/deepmac](https://google.github.io/deepmac/)
## Citation
diff --git a/official/projects/deepmac_maskrcnn/__init__.py b/official/projects/deepmac_maskrcnn/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/deepmac_maskrcnn/__init__.py
+++ b/official/projects/deepmac_maskrcnn/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/common/__init__.py b/official/projects/deepmac_maskrcnn/common/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/deepmac_maskrcnn/common/__init__.py
+++ b/official/projects/deepmac_maskrcnn/common/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/common/registry_imports.py b/official/projects/deepmac_maskrcnn/common/registry_imports.py
index 018e01f61c1..0100c703787 100644
--- a/official/projects/deepmac_maskrcnn/common/registry_imports.py
+++ b/official/projects/deepmac_maskrcnn/common/registry_imports.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/configs/__init__.py b/official/projects/deepmac_maskrcnn/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/deepmac_maskrcnn/configs/__init__.py
+++ b/official/projects/deepmac_maskrcnn/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn.py b/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn.py
index 932e76dc883..a3c52014b75 100644
--- a/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn.py
+++ b/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -35,14 +35,16 @@ class DeepMaskHead(maskrcnn_config.MaskHead):
@dataclasses.dataclass
class DeepMaskHeadRCNN(maskrcnn_config.MaskRCNN):
- mask_head: Optional[DeepMaskHead] = DeepMaskHead()
+ mask_head: Optional[DeepMaskHead] = dataclasses.field(
+ default_factory=DeepMaskHead
+ )
use_gt_boxes_for_masks: bool = False
@dataclasses.dataclass
class DeepMaskHeadRCNNTask(maskrcnn_config.MaskRCNNTask):
"""Configuration for the deep mask head R-CNN task."""
- model: DeepMaskHeadRCNN = DeepMaskHeadRCNN()
+ model: DeepMaskHeadRCNN = dataclasses.field(default_factory=DeepMaskHeadRCNN)
@exp_factory.register_config_factory('deep_mask_head_rcnn_resnetfpn_coco')
@@ -56,21 +58,21 @@ def deep_mask_head_rcnn_resnetfpn_coco() -> cfg.ExperimentConfig:
config = cfg.ExperimentConfig(
runtime=cfg.RuntimeConfig(mixed_precision_dtype='bfloat16'),
task=DeepMaskHeadRCNNTask(
- init_checkpoint='gs://cloud-tpu-checkpoints/vision-2.0/resnet50_imagenet/ckpt-28080',
- init_checkpoint_modules='backbone',
- annotation_file=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE,
+ init_checkpoint='gs://cloud-tpu-checkpoints/vision-2.0/resnet50_imagenet/ckpt-28080', # pyrefly: ignore[unexpected-keyword]
+ init_checkpoint_modules='backbone', # pyrefly: ignore[unexpected-keyword]
+ annotation_file=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE, # pyrefly: ignore[unexpected-keyword]
'instances_val2017.json'),
model=DeepMaskHeadRCNN(
num_classes=91, input_size=[1024, 1024, 3], include_mask=True), # pytype: disable=wrong-keyword-args
- losses=maskrcnn_config.Losses(l2_weight_decay=0.00004),
- train_data=maskrcnn_config.DataConfig(
+ losses=maskrcnn_config.Losses(l2_weight_decay=0.00004), # pyrefly: ignore[unexpected-keyword]
+ train_data=maskrcnn_config.DataConfig( # pyrefly: ignore[unexpected-keyword]
input_path=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE,
'train*'),
is_training=True,
global_batch_size=global_batch_size,
parser=maskrcnn_config.Parser(
aug_rand_hflip=True, aug_scale_min=0.8, aug_scale_max=1.25)),
- validation_data=maskrcnn_config.DataConfig(
+ validation_data=maskrcnn_config.DataConfig( # pyrefly: ignore[unexpected-keyword]
input_path=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE,
'val*'),
is_training=False,
@@ -123,34 +125,34 @@ def deep_mask_head_rcnn_spinenet_coco() -> cfg.ExperimentConfig:
config = cfg.ExperimentConfig(
runtime=cfg.RuntimeConfig(mixed_precision_dtype='bfloat16'),
task=DeepMaskHeadRCNNTask(
- annotation_file=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE,
+ annotation_file=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE, # pyrefly: ignore[unexpected-keyword]
'instances_val2017.json'), # pytype: disable=wrong-keyword-args
model=DeepMaskHeadRCNN(
- backbone=backbones.Backbone(
+ backbone=backbones.Backbone( # pyrefly: ignore[unexpected-keyword]
type='spinenet',
spinenet=backbones.SpineNet(
model_id='49',
min_level=3,
max_level=7,
)),
- decoder=decoders.Decoder(
+ decoder=decoders.Decoder( # pyrefly: ignore[unexpected-keyword]
type='identity', identity=decoders.Identity()),
- anchor=maskrcnn_config.Anchor(anchor_size=3),
- norm_activation=common.NormActivation(use_sync_bn=True),
- num_classes=91,
- input_size=[640, 640, 3],
- min_level=3,
- max_level=7,
+ anchor=maskrcnn_config.Anchor(anchor_size=3), # pyrefly: ignore[unexpected-keyword]
+ norm_activation=common.NormActivation(use_sync_bn=True), # pyrefly: ignore[unexpected-keyword]
+ num_classes=91, # pyrefly: ignore[unexpected-keyword]
+ input_size=[640, 640, 3], # pyrefly: ignore[unexpected-keyword]
+ min_level=3, # pyrefly: ignore[unexpected-keyword]
+ max_level=7, # pyrefly: ignore[unexpected-keyword]
include_mask=True), # pytype: disable=wrong-keyword-args
- losses=maskrcnn_config.Losses(l2_weight_decay=0.00004),
- train_data=maskrcnn_config.DataConfig(
+ losses=maskrcnn_config.Losses(l2_weight_decay=0.00004), # pyrefly: ignore[unexpected-keyword]
+ train_data=maskrcnn_config.DataConfig( # pyrefly: ignore[unexpected-keyword]
input_path=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE,
'train*'),
is_training=True,
global_batch_size=train_batch_size,
parser=maskrcnn_config.Parser(
aug_rand_hflip=True, aug_scale_min=0.5, aug_scale_max=2.0)),
- validation_data=maskrcnn_config.DataConfig(
+ validation_data=maskrcnn_config.DataConfig( # pyrefly: ignore[unexpected-keyword]
input_path=os.path.join(maskrcnn_config.COCO_INPUT_PATH_BASE,
'val*'),
is_training=False,
diff --git a/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn_config_test.py b/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn_config_test.py
index 03a7d527473..75a7ff070e5 100644
--- a/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn_config_test.py
+++ b/official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn_config_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Check that the config is set correctly."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.deepmac_maskrcnn.configs import deep_mask_head_rcnn
diff --git a/official/projects/deepmac_maskrcnn/modeling/__init__.py b/official/projects/deepmac_maskrcnn/modeling/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/deepmac_maskrcnn/modeling/__init__.py
+++ b/official/projects/deepmac_maskrcnn/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/modeling/heads/__init__.py b/official/projects/deepmac_maskrcnn/modeling/heads/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/deepmac_maskrcnn/modeling/heads/__init__.py
+++ b/official/projects/deepmac_maskrcnn/modeling/heads/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/modeling/heads/hourglass_network.py b/official/projects/deepmac_maskrcnn/modeling/heads/hourglass_network.py
index b6f3cac996d..6f0eb169ee4 100644
--- a/official/projects/deepmac_maskrcnn/modeling/heads/hourglass_network.py
+++ b/official/projects/deepmac_maskrcnn/modeling/heads/hourglass_network.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -32,14 +32,14 @@
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
BATCH_NORM_EPSILON = 1e-5
BATCH_NORM_MOMENTUM = 0.1
BATCH_NORM_FUSED = True
-class IdentityLayer(tf.keras.layers.Layer):
+class IdentityLayer(tf_keras.layers.Layer):
"""A layer which passes through the input as it is."""
def call(self, inputs):
@@ -58,14 +58,14 @@ def _get_padding_for_kernel_size(kernel_size):
def batchnorm():
try:
- return tf.keras.layers.experimental.SyncBatchNormalization(
+ return tf_keras.layers.experimental.SyncBatchNormalization(
name='batchnorm', epsilon=1e-5, momentum=0.1)
except AttributeError:
- return tf.keras.layers.BatchNormalization(
+ return tf_keras.layers.BatchNormalization(
name='batchnorm', epsilon=1e-5, momentum=0.1, fused=BATCH_NORM_FUSED)
-class ConvolutionalBlock(tf.keras.layers.Layer):
+class ConvolutionalBlock(tf_keras.layers.Layer):
"""Block that aggregates Convolution + Norm layer + ReLU."""
def __init__(self, kernel_size, out_channels, stride=1, relu=True,
@@ -87,18 +87,18 @@ def __init__(self, kernel_size, out_channels, stride=1, relu=True,
# TODO(vighneshb) Explore if removing and using padding option in conv
# layer works.
- self.pad = tf.keras.layers.ZeroPadding2D(padding_size)
+ self.pad = tf_keras.layers.ZeroPadding2D(padding_size)
else:
self.pad = IdentityLayer()
- self.conv = tf.keras.layers.Conv2D(
+ self.conv = tf_keras.layers.Conv2D(
filters=out_channels, kernel_size=kernel_size, use_bias=False,
strides=stride, padding=padding)
self.norm = batchnorm()
if relu:
- self.relu = tf.keras.layers.ReLU()
+ self.relu = tf_keras.layers.ReLU()
else:
self.relu = IdentityLayer()
@@ -123,7 +123,7 @@ def __init__(self, out_channels, stride):
out_channels=out_channels, kernel_size=1, stride=stride, relu=False)
-class ResidualBlock(tf.keras.layers.Layer):
+class ResidualBlock(tf_keras.layers.Layer):
"""A Residual block."""
def __init__(self, out_channels, skip_conv=False, kernel_size=3, stride=1,
@@ -142,7 +142,7 @@ def __init__(self, out_channels, skip_conv=False, kernel_size=3, stride=1,
self.conv_block = ConvolutionalBlock(
kernel_size=kernel_size, out_channels=out_channels, stride=stride)
- self.conv = tf.keras.layers.Conv2D(
+ self.conv = tf_keras.layers.Conv2D(
filters=out_channels, kernel_size=kernel_size, use_bias=False,
strides=1, padding=padding)
self.norm = batchnorm()
@@ -153,7 +153,7 @@ def __init__(self, out_channels, skip_conv=False, kernel_size=3, stride=1,
else:
self.skip = IdentityLayer()
- self.relu = tf.keras.layers.ReLU()
+ self.relu = tf_keras.layers.ReLU()
def call(self, inputs):
net = self.conv_block(inputs)
@@ -163,7 +163,7 @@ def call(self, inputs):
return self.relu(net + net_skip)
-class InputDownsampleBlock(tf.keras.layers.Layer):
+class InputDownsampleBlock(tf_keras.layers.Layer):
"""Block for the initial feature downsampling."""
def __init__(self, out_channels_initial_conv, out_channels_residual_block):
@@ -187,7 +187,7 @@ def call(self, inputs):
return self.residual_block(self.conv_block(inputs))
-class InputConvBlock(tf.keras.layers.Layer):
+class InputConvBlock(tf_keras.layers.Layer):
"""Block for the initial feature convolution.
This block is used in the hourglass network when we don't want to downsample
@@ -283,7 +283,7 @@ def _apply_blocks(inputs, blocks):
return net
-class EncoderDecoderBlock(tf.keras.layers.Layer):
+class EncoderDecoderBlock(tf_keras.layers.Layer):
"""An encoder-decoder block which recursively defines the hourglass network."""
def __init__(self, num_stages, channel_dims, blocks_per_stage,
@@ -316,7 +316,7 @@ def __init__(self, num_stages, channel_dims, blocks_per_stage,
self.encoder_decoder_shortcut = encoder_decoder_shortcut
if encoder_decoder_shortcut:
- self.merge_features = tf.keras.layers.Add()
+ self.merge_features = tf_keras.layers.Add()
self.encoder_block1 = _make_repeated_residual_blocks(
out_channels=out_channels, num_blocks=blocks_per_stage[0],
initial_stride=1)
@@ -343,7 +343,7 @@ def __init__(self, num_stages, channel_dims, blocks_per_stage,
residual_channels=out_channels_downsampled,
out_channels=out_channels, num_blocks=blocks_per_stage[0])
- self.upsample = tf.keras.layers.UpSampling2D(initial_stride)
+ self.upsample = tf_keras.layers.UpSampling2D(initial_stride)
def call(self, inputs):
@@ -357,12 +357,12 @@ def call(self, inputs):
upsampled_outputs = self.upsample(decoded_outputs)
if self.encoder_decoder_shortcut:
- return self.merge_features([encoded_outputs, upsampled_outputs])
+ return self.merge_features([encoded_outputs, upsampled_outputs]) # pyrefly: ignore[unbound-name]
else:
return upsampled_outputs
-class HourglassNetwork(tf.keras.Model):
+class HourglassNetwork(tf_keras.Model):
"""The hourglass network."""
def __init__(self, num_stages, input_channel_dims, channel_dims_per_stage,
@@ -437,9 +437,9 @@ def __init__(self, num_stages, input_channel_dims, channel_dims_per_stage,
ResidualBlock(out_channels=channel_dims_per_stage[0])
)
- self.intermediate_relu = tf.keras.layers.ReLU()
+ self.intermediate_relu = tf_keras.layers.ReLU()
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
if self.initial_downsample:
inputs = self.downsample_input(inputs)
diff --git a/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads.py b/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads.py
index 96a3eb579d9..6b2747f1834 100644
--- a/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads.py
+++ b/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,16 +14,14 @@
"""Instance prediction heads."""
-# Import libraries
-
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.projects.deepmac_maskrcnn.modeling.heads import hourglass_network
-class DeepMaskHead(tf.keras.layers.Layer):
+class DeepMaskHead(tf_keras.layers.Layer):
"""Creates a mask head."""
def __init__(self,
@@ -59,9 +57,9 @@ def __init__(self,
normalization across different replicas.
norm_momentum: A `float` of normalization momentum for the moving average.
norm_epsilon: A `float` added to variance to avoid dividing by zero.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default is None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2D.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
class_agnostic: A `bool`. If set, we use a single channel mask head that
is shared between all classes.
convnet_variant: A `str` denoting the architecture of network used in the
@@ -86,16 +84,16 @@ def __init__(self,
'convnet_variant': convnet_variant,
}
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
self._activation = tf_utils.get_activation(activation)
def _get_conv_op_and_kwargs(self):
- conv_op = (tf.keras.layers.SeparableConv2D
+ conv_op = (tf_keras.layers.SeparableConv2D
if self._config_dict['use_separable_conv']
- else tf.keras.layers.Conv2D)
+ else tf_keras.layers.Conv2D)
conv_kwargs = {
'filters': self._config_dict['num_filters'],
'kernel_size': 3,
@@ -103,9 +101,9 @@ def _get_conv_op_and_kwargs(self):
}
if self._config_dict['use_separable_conv']:
conv_kwargs.update({
- 'depthwise_initializer': tf.keras.initializers.VarianceScaling(
+ 'depthwise_initializer': tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
- 'pointwise_initializer': tf.keras.initializers.VarianceScaling(
+ 'pointwise_initializer': tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
'bias_initializer': tf.zeros_initializer(),
'depthwise_regularizer': self._config_dict['kernel_regularizer'],
@@ -114,7 +112,7 @@ def _get_conv_op_and_kwargs(self):
})
else:
conv_kwargs.update({
- 'kernel_initializer': tf.keras.initializers.VarianceScaling(
+ 'kernel_initializer': tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
'bias_initializer': tf.zeros_initializer(),
'kernel_regularizer': self._config_dict['kernel_regularizer'],
@@ -125,9 +123,9 @@ def _get_conv_op_and_kwargs(self):
def _get_bn_op_and_kwargs(self):
- bn_op = (tf.keras.layers.experimental.SyncBatchNormalization
+ bn_op = (tf_keras.layers.experimental.SyncBatchNormalization
if self._config_dict['use_sync_bn']
- else tf.keras.layers.BatchNormalization)
+ else tf_keras.layers.BatchNormalization)
bn_kwargs = {
'axis': self._bn_axis,
'momentum': self._config_dict['norm_momentum'],
@@ -143,12 +141,12 @@ def build(self, input_shape):
self._build_convnet_variant()
- self._deconv = tf.keras.layers.Conv2DTranspose(
+ self._deconv = tf_keras.layers.Conv2DTranspose(
filters=self._config_dict['num_filters'],
kernel_size=self._config_dict['upsample_factor'],
strides=self._config_dict['upsample_factor'],
padding='valid',
- kernel_initializer=tf.keras.initializers.VarianceScaling(
+ kernel_initializer=tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
bias_initializer=tf.zeros_initializer(),
kernel_regularizer=self._config_dict['kernel_regularizer'],
@@ -170,9 +168,9 @@ def build(self, input_shape):
}
if self._config_dict['use_separable_conv']:
conv_kwargs.update({
- 'depthwise_initializer': tf.keras.initializers.VarianceScaling(
+ 'depthwise_initializer': tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
- 'pointwise_initializer': tf.keras.initializers.VarianceScaling(
+ 'pointwise_initializer': tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
'bias_initializer': tf.zeros_initializer(),
'depthwise_regularizer': self._config_dict['kernel_regularizer'],
@@ -181,7 +179,7 @@ def build(self, input_shape):
})
else:
conv_kwargs.update({
- 'kernel_initializer': tf.keras.initializers.VarianceScaling(
+ 'kernel_initializer': tf_keras.initializers.VarianceScaling(
scale=2, mode='fan_out', distribution='untruncated_normal'),
'bias_initializer': tf.zeros_initializer(),
'kernel_regularizer': self._config_dict['kernel_regularizer'],
@@ -209,11 +207,12 @@ def call(self, inputs, training=None):
"""
roi_features, roi_classes = inputs
features_shape = tf.shape(roi_features)
- batch_size, num_rois, height, width, filters = (
- features_shape[0], features_shape[1], features_shape[2],
- features_shape[3], features_shape[4])
- if batch_size is None:
- batch_size = tf.shape(roi_features)[0]
+ num_rois, height, width, filters = (
+ features_shape[1],
+ features_shape[2],
+ features_shape[3],
+ features_shape[4],
+ )
x = tf.reshape(roi_features, [-1, height, width, filters])
@@ -229,40 +228,26 @@ def call(self, inputs, training=None):
mask_width = width * self._config_dict['upsample_factor']
if self._config_dict['class_agnostic']:
- logits = tf.reshape(logits, [-1, num_rois, mask_height, mask_width, 1])
+ return tf.reshape(logits, [-1, num_rois, mask_height, mask_width])
else:
logits = tf.reshape(
logits,
[-1, num_rois, mask_height, mask_width,
self._config_dict['num_classes']])
-
- batch_indices = tf.tile(
- tf.expand_dims(tf.range(batch_size), axis=1), [1, num_rois])
- mask_indices = tf.tile(
- tf.expand_dims(tf.range(num_rois), axis=0), [batch_size, 1])
-
- if self._config_dict['class_agnostic']:
- class_gather_indices = tf.zeros_like(roi_classes, dtype=tf.int32)
- else:
- class_gather_indices = tf.cast(roi_classes, dtype=tf.int32)
-
- gather_indices = tf.stack(
- [batch_indices, mask_indices, class_gather_indices],
- axis=2)
- mask_outputs = tf.gather_nd(
- tf.transpose(logits, [0, 1, 4, 2, 3]), gather_indices)
- return mask_outputs
+ return tf.gather(
+ logits, tf.cast(roi_classes, dtype=tf.int32), axis=-1, batch_dims=2
+ )
def _build_convnet_variant(self):
variant = self._config_dict['convnet_variant']
if variant == 'default':
- conv_op, conv_kwargs = self._get_conv_op_and_kwargs()
bn_op, bn_kwargs = self._get_bn_op_and_kwargs()
self._convs = []
self._conv_norms = []
for i in range(self._config_dict['num_convs']):
conv_name = 'mask-conv_{}'.format(i)
+ conv_op, conv_kwargs = self._get_conv_op_and_kwargs()
self._convs.append(conv_op(name=conv_name, **conv_kwargs))
bn_name = 'mask-conv-bn_{}'.format(i)
self._conv_norms.append(bn_op(name=bn_name, **bn_kwargs))
diff --git a/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads_test.py b/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads_test.py
index 20cdc0fcab6..41a88db0a2b 100644
--- a/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads_test.py
+++ b/official/projects/deepmac_maskrcnn/modeling/heads/instance_heads_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,9 @@
"""Tests for instance_heads.py."""
-# Import libraries
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.deepmac_maskrcnn.modeling.heads import instance_heads as deep_instance_heads
diff --git a/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model.py b/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model.py
index 488485e2881..a2ab95cf26b 100644
--- a/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model.py
+++ b/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,12 +16,11 @@
from typing import List, Mapping, Optional, Union
-# Import libraries
-
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.modeling import maskrcnn_model
+from official.vision.ops import box_ops
def resize_as(source, size):
@@ -35,19 +34,19 @@ class DeepMaskRCNNModel(maskrcnn_model.MaskRCNNModel):
"""The Mask R-CNN model."""
def __init__(self,
- backbone: tf.keras.Model,
- decoder: tf.keras.Model,
- rpn_head: tf.keras.layers.Layer,
- detection_head: Union[tf.keras.layers.Layer,
- List[tf.keras.layers.Layer]],
- roi_generator: tf.keras.layers.Layer,
- roi_sampler: Union[tf.keras.layers.Layer,
- List[tf.keras.layers.Layer]],
- roi_aligner: tf.keras.layers.Layer,
- detection_generator: tf.keras.layers.Layer,
- mask_head: Optional[tf.keras.layers.Layer] = None,
- mask_sampler: Optional[tf.keras.layers.Layer] = None,
- mask_roi_aligner: Optional[tf.keras.layers.Layer] = None,
+ backbone: tf_keras.Model,
+ decoder: tf_keras.Model,
+ rpn_head: tf_keras.layers.Layer,
+ detection_head: Union[tf_keras.layers.Layer,
+ List[tf_keras.layers.Layer]],
+ roi_generator: tf_keras.layers.Layer,
+ roi_sampler: Union[tf_keras.layers.Layer,
+ List[tf_keras.layers.Layer]],
+ roi_aligner: tf_keras.layers.Layer,
+ detection_generator: tf_keras.layers.Layer,
+ mask_head: Optional[tf_keras.layers.Layer] = None,
+ mask_sampler: Optional[tf_keras.layers.Layer] = None,
+ mask_roi_aligner: Optional[tf_keras.layers.Layer] = None,
class_agnostic_bbox_pred: bool = False,
cascade_class_ensemble: bool = False,
min_level: Optional[int] = None,
@@ -55,13 +54,14 @@ def __init__(self,
num_scales: Optional[int] = None,
aspect_ratios: Optional[List[float]] = None,
anchor_size: Optional[float] = None,
+ outer_boxes_scale: float = 1.0,
use_gt_boxes_for_masks=False,
**kwargs):
"""Initializes the Mask R-CNN model.
Args:
- backbone: `tf.keras.Model`, the backbone network.
- decoder: `tf.keras.Model`, the decoder network.
+ backbone: `tf_keras.Model`, the backbone network.
+ decoder: `tf_keras.Model`, the decoder network.
rpn_head: the RPN head.
detection_head: the detection head or a list of heads.
roi_generator: the ROI generator.
@@ -86,11 +86,13 @@ def __init__(self,
aspect_ratios=[1.0, 2.0, 0.5] adds three anchors on each scale level.
anchor_size: A number representing the scale of size of the base anchor to
the feature stride 2^level.
+ outer_boxes_scale: a float to scale up the bounding boxes to generate
+ more inclusive masks. The scale is expected to be >=1.0.
use_gt_boxes_for_masks: bool, if set, crop using groundtruth boxes instead
of proposals for training mask head
**kwargs: keyword arguments to be passed.
"""
- super(DeepMaskRCNNModel, self).__init__(
+ super().__init__(
backbone=backbone,
decoder=decoder,
rpn_head=rpn_head,
@@ -109,6 +111,7 @@ def __init__(self,
num_scales=num_scales,
aspect_ratios=aspect_ratios,
anchor_size=anchor_size,
+ outer_boxes_scale=outer_boxes_scale,
**kwargs)
self._config_dict['use_gt_boxes_for_masks'] = use_gt_boxes_for_masks
@@ -120,24 +123,44 @@ def call(self,
gt_boxes: Optional[tf.Tensor] = None,
gt_classes: Optional[tf.Tensor] = None,
gt_masks: Optional[tf.Tensor] = None,
+ gt_outer_boxes: Optional[tf.Tensor] = None,
training: Optional[bool] = None) -> Mapping[str, tf.Tensor]:
-
+ call_box_outputs_kwargs = {
+ 'images': images,
+ 'image_shape': image_shape,
+ 'anchor_boxes': anchor_boxes,
+ 'gt_boxes': gt_boxes,
+ 'gt_classes': gt_classes,
+ 'training': training
+ }
+ if self.outer_boxes_scale > 1.0:
+ call_box_outputs_kwargs['gt_outer_boxes'] = gt_outer_boxes
model_outputs, intermediate_outputs = self._call_box_outputs(
- images=images, image_shape=image_shape, anchor_boxes=anchor_boxes,
- gt_boxes=gt_boxes, gt_classes=gt_classes, training=training)
+ **call_box_outputs_kwargs)
if not self._include_mask:
return model_outputs
+ if self.outer_boxes_scale == 1.0:
+ current_rois = intermediate_outputs['current_rois']
+ matched_gt_boxes = intermediate_outputs['matched_gt_boxes']
+ mask_head_gt_boxes = gt_boxes
+ else:
+ current_rois = box_ops.compute_outer_boxes(
+ intermediate_outputs['current_rois'],
+ tf.expand_dims(image_shape, axis=1), self.outer_boxes_scale)
+ matched_gt_boxes = intermediate_outputs['matched_gt_outer_boxes']
+ mask_head_gt_boxes = gt_outer_boxes
+
model_mask_outputs = self._call_mask_outputs(
model_box_outputs=model_outputs,
features=model_outputs['decoder_features'],
- current_rois=intermediate_outputs['current_rois'],
+ current_rois=current_rois,
matched_gt_indices=intermediate_outputs['matched_gt_indices'],
- matched_gt_boxes=intermediate_outputs['matched_gt_boxes'],
+ matched_gt_boxes=matched_gt_boxes,
matched_gt_classes=intermediate_outputs['matched_gt_classes'],
- gt_masks=gt_masks,
- gt_classes=gt_classes,
- gt_boxes=gt_boxes,
+ gt_masks=gt_masks, # pyrefly: ignore[bad-argument-type]
+ gt_classes=gt_classes, # pyrefly: ignore[bad-argument-type]
+ gt_boxes=mask_head_gt_boxes, # pyrefly: ignore[bad-argument-type]
training=training)
model_outputs.update(model_mask_outputs)
return model_outputs
@@ -194,7 +217,10 @@ def _call_mask_outputs(
})
else:
- rois = model_outputs['detection_boxes']
+ if self.outer_boxes_scale == 1.0:
+ rois = model_outputs['detection_boxes']
+ else:
+ rois = model_outputs['detection_outer_boxes']
roi_classes = model_outputs['detection_classes']
# Mask RoI align.
@@ -204,8 +230,8 @@ def _call_mask_outputs(
mask_head_classes = gt_classes
else:
- roi_aligner_boxes = rois
- mask_head_classes = roi_classes
+ roi_aligner_boxes = rois # pyrefly: ignore[unbound-name]
+ mask_head_classes = roi_classes # pyrefly: ignore[unbound-name]
mask_logits, mask_probs = self._features_to_mask_outputs(
features, roi_aligner_boxes, mask_head_classes)
diff --git a/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model_test.py b/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model_test.py
index 08e9ab5376f..53d7b845d44 100644
--- a/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model_test.py
+++ b/official/projects/deepmac_maskrcnn/modeling/maskrcnn_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,9 @@
"""Tests for maskrcnn_model.py."""
-# Import libraries
-
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.deepmac_maskrcnn.modeling import maskrcnn_model
from official.projects.deepmac_maskrcnn.modeling.heads import instance_heads as deep_instance_heads
@@ -50,7 +48,7 @@ def construct_model_and_anchors(image_size, use_gt_boxes_for_masks):
image_size=image_size).multilevel_boxes
num_anchors_per_location = len(aspect_ratios) * num_scales
- input_specs = tf.keras.layers.InputSpec(shape=[None, None, None, 3])
+ input_specs = tf_keras.layers.InputSpec(shape=[None, None, None, 3])
backbone = resnet.ResNet(model_id=50, input_specs=input_specs)
decoder = fpn.FPN(
min_level=min_level,
@@ -92,12 +90,14 @@ def construct_model_and_anchors(image_size, use_gt_boxes_for_masks):
class MaskRCNNModelTest(parameterized.TestCase, tf.test.TestCase):
@parameterized.parameters(
- (False, False,),
- (False, True,),
- (True, False,),
- (True, True,),
+ (False, False, False),
+ (False, True, False),
+ (True, False, True),
+ (True, False, False),
+ (True, True, True),
+ (True, True, False),
)
- def test_forward(self, use_gt_boxes_for_masks, training):
+ def test_forward(self, use_gt_boxes_for_masks, training, use_outer_boxes):
image_size = (256, 256)
images = np.random.rand(2, image_size[0], image_size[1], 3)
image_shape = np.array([[224, 100], [100, 224]])
@@ -105,6 +105,9 @@ def test_forward(self, use_gt_boxes_for_masks, training):
image_size, use_gt_boxes_for_masks)
gt_boxes = tf.zeros((2, 16, 4), dtype=tf.float32)
+ gt_outer_boxes = None
+ if use_outer_boxes:
+ gt_outer_boxes = tf.zeros((2, 16, 4), dtype=tf.float32)
gt_masks = tf.zeros((2, 16, 32, 32))
gt_classes = tf.zeros((2, 16), dtype=tf.int32)
results = model(images.astype(np.uint8),
@@ -113,6 +116,7 @@ def test_forward(self, use_gt_boxes_for_masks, training):
gt_boxes,
gt_classes,
gt_masks,
+ gt_outer_boxes,
training=training)
self.assertIn('rpn_boxes', results)
@@ -137,12 +141,12 @@ def test_forward(self, use_gt_boxes_for_masks, training):
)
def test_image_and_boxes(self, batch_size, num_boxes):
image_size = (640, 640)
- images = np.random.rand(1, image_size[0], image_size[1], 3).astype(
+ images = np.random.rand(batch_size, image_size[0], image_size[1], 3).astype(
np.float32)
model, _ = construct_model_and_anchors(
image_size, use_gt_boxes_for_masks=True)
- boxes = np.zeros((1, num_boxes, 4), dtype=np.float32)
+ boxes = np.zeros((batch_size, num_boxes, 4), dtype=np.float32)
boxes[:, :, [2, 3]] = 1.0
boxes = tf.constant(boxes)
results = model.call_images_and_boxes(images, boxes)
diff --git a/official/projects/deepmac_maskrcnn/serving/__init__.py b/official/projects/deepmac_maskrcnn/serving/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/deepmac_maskrcnn/serving/__init__.py
+++ b/official/projects/deepmac_maskrcnn/serving/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/serving/detection.py b/official/projects/deepmac_maskrcnn/serving/detection.py
index 9e3bbfd2889..ee0a76a2cb4 100644
--- a/official/projects/deepmac_maskrcnn/serving/detection.py
+++ b/official/projects/deepmac_maskrcnn/serving/detection.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Dict, Mapping, Text
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.deepmac_maskrcnn.configs import deep_mask_head_rcnn as cfg
from official.projects.deepmac_maskrcnn.modeling import maskrcnn_model
@@ -61,7 +61,7 @@ def _build_model(self):
ValueError("batch_size can't be None for detection models")
if self.params.task.model.detection_generator.nms_version != 'batched':
ValueError('Only batched_nms is supported.')
- input_specs = tf.keras.layers.InputSpec(shape=[self._batch_size] +
+ input_specs = tf_keras.layers.InputSpec(shape=[self._batch_size] +
self._input_image_size + [3])
if isinstance(self.params.task.model, cfg.DeepMaskHeadRCNN):
diff --git a/official/projects/deepmac_maskrcnn/serving/detection_test.py b/official/projects/deepmac_maskrcnn/serving/detection_test.py
index cd832e2821f..a2078a20cd6 100644
--- a/official/projects/deepmac_maskrcnn/serving/detection_test.py
+++ b/official/projects/deepmac_maskrcnn/serving/detection_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl.testing import parameterized
import numpy as np
from PIL import Image
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import exp_factory
from official.projects.deepmac_maskrcnn.serving import detection
diff --git a/official/projects/deepmac_maskrcnn/serving/export_saved_model.py b/official/projects/deepmac_maskrcnn/serving/export_saved_model.py
index b88aa70eb8a..b1ad4e43bfb 100644
--- a/official/projects/deepmac_maskrcnn/serving/export_saved_model.py
+++ b/official/projects/deepmac_maskrcnn/serving/export_saved_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/tasks/__init__.py b/official/projects/deepmac_maskrcnn/tasks/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/deepmac_maskrcnn/tasks/__init__.py
+++ b/official/projects/deepmac_maskrcnn/tasks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/deepmac_maskrcnn/tasks/deep_mask_head_rcnn.py b/official/projects/deepmac_maskrcnn/tasks/deep_mask_head_rcnn.py
index 4f4682837ee..9d446f52095 100644
--- a/official/projects/deepmac_maskrcnn/tasks/deep_mask_head_rcnn.py
+++ b/official/projects/deepmac_maskrcnn/tasks/deep_mask_head_rcnn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Mask R-CNN variant with support for deep mask heads."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import task_factory
from official.projects.deepmac_maskrcnn.configs import deep_mask_head_rcnn as deep_mask_head_rcnn_config
@@ -33,9 +33,9 @@
# Taken from modeling/factory.py
-def build_maskrcnn(input_specs: tf.keras.layers.InputSpec,
+def build_maskrcnn(input_specs: tf_keras.layers.InputSpec,
model_config: deep_mask_head_rcnn_config.DeepMaskHeadRCNN,
- l2_regularizer: tf.keras.regularizers.Regularizer = None): # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None): # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds Mask R-CNN model."""
norm_activation_config = model_config.norm_activation
backbone = backbones.factory.build_backbone(
@@ -120,7 +120,8 @@ def build_maskrcnn(input_specs: tf.keras.layers.InputSpec,
pre_nms_score_threshold=generator_config.pre_nms_score_threshold,
nms_iou_threshold=generator_config.nms_iou_threshold,
max_num_detections=generator_config.max_num_detections,
- nms_version=generator_config.nms_version)
+ nms_version=generator_config.nms_version,
+ use_sigmoid_probability=generator_config.use_sigmoid_probability)
if model_config.include_mask:
mask_head = deep_instance_heads.DeepMaskHead(
@@ -162,6 +163,14 @@ def build_maskrcnn(input_specs: tf.keras.layers.InputSpec,
mask_head=mask_head,
mask_sampler=mask_sampler_obj,
mask_roi_aligner=mask_roi_aligner_obj,
+ class_agnostic_bbox_pred=detection_head_config.class_agnostic_bbox_pred,
+ cascade_class_ensemble=detection_head_config.cascade_class_ensemble,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ num_scales=model_config.anchor.num_scales,
+ aspect_ratios=model_config.anchor.aspect_ratios,
+ anchor_size=model_config.anchor.anchor_size,
+ outer_boxes_scale=model_config.outer_boxes_scale,
use_gt_boxes_for_masks=model_config.use_gt_boxes_for_masks)
return model
@@ -171,20 +180,29 @@ class DeepMaskHeadRCNNTask(maskrcnn.MaskRCNNTask):
"""Mask R-CNN with support for deep mask heads."""
def build_model(self):
- """Build Mask R-CNN model."""
+ """Builds Mask R-CNN model."""
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None] + self.task_config.model.input_size)
l2_weight_decay = self.task_config.losses.l2_weight_decay
# Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
- l2_regularizer = (tf.keras.regularizers.l2(
+ l2_regularizer = (tf_keras.regularizers.l2(
l2_weight_decay / 2.0) if l2_weight_decay else None)
model = build_maskrcnn(
input_specs=input_specs,
model_config=self.task_config.model,
- l2_regularizer=l2_regularizer)
+ l2_regularizer=l2_regularizer) # pyrefly: ignore[bad-argument-type]
+
+ if self.task_config.freeze_backbone:
+ model.backbone.trainable = False
+
+ # Builds the model through warm-up call.
+ dummy_images = tf_keras.Input(self.task_config.model.input_size)
+ dummy_image_shape = tf_keras.layers.Input([2])
+ _ = model(dummy_images, image_shape=dummy_image_shape, training=False)
+
return model
diff --git a/official/projects/deepmac_maskrcnn/train.py b/official/projects/deepmac_maskrcnn/train.py
index ac866f51ded..9328c9e9e59 100644
--- a/official/projects/deepmac_maskrcnn/train.py
+++ b/official/projects/deepmac_maskrcnn/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/detr/README.md b/official/projects/detr/README.md
index e8860f5e1eb..8bab41dba81 100644
--- a/official/projects/detr/README.md
+++ b/official/projects/detr/README.md
@@ -19,8 +19,8 @@ detr/experiments.
| Model | Resolution | Batch size | Epochs | Decay@ | Params (M) | Box AP | Dashboard | Checkpoint | Experiment |
| --------- | :--------: | ----------:| ------:| -----: | ---------: | -----: | --------: | ---------: | ---------: |
-| DETR-ResNet-50 | 1333x1333 |64|300| 200 |41 | 40.6 | [tensorboard](https://tensorboard.dev/experiment/o2IEZnniRYu6pqViBeopIg/#scalars) | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/detr/detr_resnet_50_300.tar.gz) | detr_r50_300epochs.sh |
-| DETR-ResNet-50 | 1333x1333 |64|500| 400 |41 | 42.0| [tensorboard](https://tensorboard.dev/experiment/YFMDKpESR4yjocPh5HgfRw/) | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/detr/detr_resnet_50_500.tar.gz) | detr_r50_500epochs.sh |
+| DETR-ResNet-50 | 1333x1333 |64|300| 200 |41 | 40.6 | [tensorboard](https://tensorboard.dev/experiment/o2IEZnniRYu6pqViBeopIg/#scalars) | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/detr/detr_resnet_50_300.tar.gz) | [detr_r50_300epochs.sh](https://github.com/tensorflow/models/blob/master/official/projects/detr/experiments/detr_r50_300epochs.sh) |
+| DETR-ResNet-50 | 1333x1333 |64|500| 400 |41 | 42.0| [tensorboard](https://tensorboard.dev/experiment/YFMDKpESR4yjocPh5HgfRw/) | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/detr/detr_resnet_50_500.tar.gz) | [detr_r50_500epochs.sh](https://github.com/tensorflow/models/blob/master/official/projects/detr/experiments/detr_r50_500epochs.sh) |
| DETR-ResNet-50 | 1333x1333 |64|300| 200 |41 | 40.6 | paper | NA | NA |
| DETR-ResNet-50 | 1333x1333 |64|500| 400 |41 | 42.0 | paper | NA | NA |
| DETR-DC5-ResNet-50 | 1333x1333 |64|500| 400 |41 | 43.3 | paper | NA | NA |
diff --git a/official/projects/detr/__init__.py b/official/projects/detr/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/configs/__init__.py b/official/projects/detr/configs/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/configs/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/configs/detr.py b/official/projects/detr/configs/detr.py
index dfab9400367..bf24fdbe4b4 100644
--- a/official/projects/detr/configs/detr.py
+++ b/official/projects/detr/configs/detr.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,44 +15,93 @@
"""DETR configurations."""
import dataclasses
+import os
+from typing import List, Optional, Union
+
from official.core import config_definitions as cfg
from official.core import exp_factory
+from official.modeling import hyperparams
from official.projects.detr import optimization
from official.projects.detr.dataloaders import coco
+from official.vision.configs import backbones
+from official.vision.configs import common
+
+
+@dataclasses.dataclass
+class DataConfig(cfg.DataConfig):
+ """Input config for training."""
+ input_path: str = ''
+ tfds_name: str = ''
+ tfds_split: str = 'train'
+ global_batch_size: int = 0
+ is_training: bool = False
+ dtype: str = 'bfloat16'
+ decoder: common.DataDecoder = dataclasses.field(default_factory=common.DataDecoder)
+ shuffle_buffer_size: int = 10000
+ file_type: str = 'tfrecord'
+ drop_remainder: bool = True
@dataclasses.dataclass
-class DetectionConfig(cfg.TaskConfig):
- """The translation task config."""
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+class Losses(hyperparams.Config):
+ class_offset: int = 0
lambda_cls: float = 1.0
lambda_box: float = 5.0
lambda_giou: float = 2.0
-
- init_ckpt: str = ''
- num_classes: int = 81 # 0: background
background_cls_weight: float = 0.1
+ l2_weight_decay: float = 1e-4
+
+
+@dataclasses.dataclass
+class Detr(hyperparams.Config):
+ """Detr model definations."""
+ num_queries: int = 100
+ hidden_size: int = 256
+ num_classes: int = 91 # 0: background
num_encoder_layers: int = 6
num_decoder_layers: int = 6
+ input_size: List[int] = dataclasses.field(default_factory=list)
+ backbone: backbones.Backbone = dataclasses.field(default_factory=lambda:backbones.Backbone(
+ type='resnet', resnet=backbones.ResNet(model_id=50, bn_trainable=False)))
+ norm_activation: common.NormActivation = dataclasses.field(default_factory=common.NormActivation)
+ backbone_endpoint_name: str = '5'
- # Make DETRConfig.
- num_queries: int = 100
- num_hidden: int = 256
+
+@dataclasses.dataclass
+class DetrTask(cfg.TaskConfig):
+ model: Detr = dataclasses.field(default_factory=Detr)
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+ init_checkpoint_modules: Union[str, List[str]] = 'all' # all, backbone
+ annotation_file: Optional[str] = None
per_category_metrics: bool = False
+COCO_INPUT_PATH_BASE = 'coco'
+COCO_TRAIN_EXAMPLES = 118287
+COCO_VAL_EXAMPLES = 5000
+
+
@exp_factory.register_config_factory('detr_coco')
def detr_coco() -> cfg.ExperimentConfig:
"""Config to get results that matches the paper."""
train_batch_size = 64
eval_batch_size = 64
- num_train_data = 118287
+ num_train_data = COCO_TRAIN_EXAMPLES
num_steps_per_epoch = num_train_data // train_batch_size
train_steps = 500 * num_steps_per_epoch # 500 epochs
decay_at = train_steps - 100 * num_steps_per_epoch # 400 epochs
config = cfg.ExperimentConfig(
- task=DetectionConfig(
+ task=DetrTask(
+ init_checkpoint='',
+ init_checkpoint_modules='backbone',
+ model=Detr(
+ num_classes=81,
+ input_size=[1333, 1333, 3],
+ norm_activation=common.NormActivation()),
+ losses=Losses(),
train_data=coco.COCODataConfig(
tfds_name='coco/2017',
tfds_split='train',
@@ -65,9 +114,7 @@ def detr_coco() -> cfg.ExperimentConfig:
tfds_split='validation',
is_training=False,
global_batch_size=eval_batch_size,
- drop_remainder=False
- )
- ),
+ drop_remainder=False)),
trainer=cfg.TrainerConfig(
train_steps=train_steps,
validation_steps=-1,
@@ -95,8 +142,135 @@ def detr_coco() -> cfg.ExperimentConfig:
'values': [0.0001, 1.0e-05]
}
},
- })
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ])
+ return config
+
+
+@exp_factory.register_config_factory('detr_coco_tfrecord')
+def detr_coco_tfrecord() -> cfg.ExperimentConfig:
+ """Config to get results that matches the paper."""
+ train_batch_size = 64
+ eval_batch_size = 64
+ steps_per_epoch = COCO_TRAIN_EXAMPLES // train_batch_size
+ train_steps = 300 * steps_per_epoch # 300 epochs
+ decay_at = train_steps - 100 * steps_per_epoch # 200 epochs
+ config = cfg.ExperimentConfig(
+ task=DetrTask(
+ init_checkpoint='',
+ init_checkpoint_modules='backbone',
+ annotation_file=os.path.join(COCO_INPUT_PATH_BASE,
+ 'instances_val2017.json'),
+ model=Detr(
+ input_size=[1333, 1333, 3],
+ norm_activation=common.NormActivation()),
+ losses=Losses(),
+ train_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=1000,
),
+ validation_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=False,
+ )),
+ trainer=cfg.TrainerConfig(
+ train_steps=train_steps,
+ validation_steps=COCO_VAL_EXAMPLES // eval_batch_size,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ validation_interval=5 * steps_per_epoch,
+ max_to_keep=1,
+ best_checkpoint_export_subdir='best_ckpt',
+ best_checkpoint_eval_metric='AP',
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'detr_adamw',
+ 'detr_adamw': {
+ 'weight_decay_rate': 1e-4,
+ 'global_clipnorm': 0.1,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0
+ }
+ },
+ 'learning_rate': {
+ 'type': 'stepwise',
+ 'stepwise': {
+ 'boundaries': [decay_at],
+ 'values': [0.0001, 1.0e-05]
+ }
+ },
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ])
+ return config
+
+
+@exp_factory.register_config_factory('detr_coco_tfds')
+def detr_coco_tfds() -> cfg.ExperimentConfig:
+ """Config to get results that matches the paper."""
+ train_batch_size = 64
+ eval_batch_size = 64
+ steps_per_epoch = COCO_TRAIN_EXAMPLES // train_batch_size
+ train_steps = 300 * steps_per_epoch # 300 epochs
+ decay_at = train_steps - 100 * steps_per_epoch # 200 epochs
+ config = cfg.ExperimentConfig(
+ task=DetrTask(
+ init_checkpoint='',
+ init_checkpoint_modules='backbone',
+ model=Detr(
+ num_classes=81,
+ input_size=[1333, 1333, 3],
+ norm_activation=common.NormActivation()),
+ losses=Losses(class_offset=1),
+ train_data=DataConfig(
+ tfds_name='coco/2017',
+ tfds_split='train',
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=1000,
+ ),
+ validation_data=DataConfig(
+ tfds_name='coco/2017',
+ tfds_split='validation',
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=False)),
+ trainer=cfg.TrainerConfig(
+ train_steps=train_steps,
+ validation_steps=COCO_VAL_EXAMPLES // eval_batch_size,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ validation_interval=5 * steps_per_epoch,
+ max_to_keep=1,
+ best_checkpoint_export_subdir='best_ckpt',
+ best_checkpoint_eval_metric='AP',
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'detr_adamw',
+ 'detr_adamw': {
+ 'weight_decay_rate': 1e-4,
+ 'global_clipnorm': 0.1,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0
+ }
+ },
+ 'learning_rate': {
+ 'type': 'stepwise',
+ 'stepwise': {
+ 'boundaries': [decay_at],
+ 'values': [0.0001, 1.0e-05]
+ }
+ },
+ })),
restrictions=[
'task.train_data.is_training != None',
])
diff --git a/official/projects/detr/configs/detr_test.py b/official/projects/detr/configs/detr_test.py
index bb96f62ca2a..b9159a03022 100644
--- a/official/projects/detr/configs/detr_test.py
+++ b/official/projects/detr/configs/detr_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import exp_factory
@@ -27,15 +27,25 @@
class DetrTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters(('detr_coco',))
- def test_detr_configs(self, config_name):
+ def test_detr_configs_tfds(self, config_name):
config = exp_factory.get_exp_config(config_name)
self.assertIsInstance(config, cfg.ExperimentConfig)
- self.assertIsInstance(config.task, exp_cfg.DetectionConfig)
+ self.assertIsInstance(config.task, exp_cfg.DetrTask)
self.assertIsInstance(config.task.train_data, coco.COCODataConfig)
config.task.train_data.is_training = None
with self.assertRaises(KeyError):
config.validate()
+ @parameterized.parameters(('detr_coco_tfrecord'), ('detr_coco_tfds'))
+ def test_detr_configs(self, config_name):
+ config = exp_factory.get_exp_config(config_name)
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task, exp_cfg.DetrTask)
+ self.assertIsInstance(config.task.train_data, cfg.DataConfig)
+ config.task.train_data.is_training = None
+ with self.assertRaises(KeyError):
+ config.validate()
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/projects/detr/dataloaders/__init__.py b/official/projects/detr/dataloaders/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/dataloaders/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/dataloaders/coco.py b/official/projects/detr/dataloaders/coco.py
index 4f28a53c208..25eb3fd531b 100644
--- a/official/projects/detr/dataloaders/coco.py
+++ b/official/projects/detr/dataloaders/coco.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import dataclasses
from typing import Optional, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import input_reader
@@ -145,7 +145,7 @@ def _transform_and_batch_fn(
self._params.global_batch_size
) if input_context else self._params.global_batch_size
dataset = dataset.batch(
- per_replica_batch_size, drop_remainder=self._params.is_training)
+ per_replica_batch_size, drop_remainder=self._params.drop_remainder)
return dataset
def load(self, input_context: Optional[tf.distribute.InputContext] = None):
diff --git a/official/projects/detr/dataloaders/coco_test.py b/official/projects/detr/dataloaders/coco_test.py
index cad38e18c00..1b496f4366a 100644
--- a/official/projects/detr/dataloaders/coco_test.py
+++ b/official/projects/detr/dataloaders/coco_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.projects.detr.dataloaders import coco
@@ -31,7 +31,7 @@ def _gen_fn():
'image/id': np.random.randint(0, 100),
'image/filename': 'test',
'objects': {
- 'is_crowd': np.ones(shape=(num_boxes), dtype=np.bool),
+ 'is_crowd': np.ones(shape=(num_boxes), dtype=bool),
'bbox': np.ones(shape=(num_boxes, 4), dtype=np.float32),
'label': np.ones(shape=(num_boxes), dtype=np.int64),
'id': np.ones(shape=(num_boxes), dtype=np.int64),
diff --git a/official/projects/detr/dataloaders/detr_input.py b/official/projects/detr/dataloaders/detr_input.py
new file mode 100644
index 00000000000..eb04f9580f9
--- /dev/null
+++ b/official/projects/detr/dataloaders/detr_input.py
@@ -0,0 +1,175 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""COCO data loader for DETR."""
+
+from typing import Tuple
+import tensorflow as tf, tf_keras
+
+from official.vision.dataloaders import parser
+
+from official.vision.ops import box_ops
+from official.vision.ops import preprocess_ops
+
+RESIZE_SCALES = (480, 512, 544, 576, 608, 640, 672, 704, 736, 768, 800)
+
+
+class Parser(parser.Parser):
+ """Parse an image and its annotations into a dictionary of tensors."""
+
+ def __init__(self,
+ class_offset: int = 0,
+ output_size: Tuple[int, int] = (1333, 1333),
+ max_num_boxes: int = 100,
+ resize_scales: Tuple[int, ...] = RESIZE_SCALES,
+ aug_rand_hflip=True):
+ self._class_offset = class_offset
+ self._output_size = output_size
+ self._max_num_boxes = max_num_boxes
+ self._resize_scales = resize_scales
+ self._aug_rand_hflip = aug_rand_hflip
+
+ def _parse_train_data(self, data):
+ """Parses data for training and evaluation."""
+ classes = data['groundtruth_classes'] + self._class_offset
+ boxes = data['groundtruth_boxes']
+ is_crowd = data['groundtruth_is_crowd']
+
+ # Gets original image.
+ image = data['image']
+
+ # Normalizes image with mean and std pixel values.
+ image = preprocess_ops.normalize_image(image)
+ image, boxes, _ = preprocess_ops.random_horizontal_flip(image, boxes)
+
+ do_crop = tf.greater(tf.random.uniform([]), 0.5)
+ if do_crop:
+ # Rescale
+ boxes = box_ops.denormalize_boxes(boxes, tf.shape(image)[:2])
+ index = tf.random.categorical(tf.zeros([1, 3]), 1)[0]
+ scales = tf.gather([400.0, 500.0, 600.0], index, axis=0)
+ short_side = scales[0]
+ image, image_info = preprocess_ops.resize_image(image, short_side)
+ boxes = preprocess_ops.resize_and_crop_boxes(boxes, image_info[2, :],
+ image_info[1, :],
+ image_info[3, :])
+ boxes = box_ops.normalize_boxes(boxes, image_info[1, :])
+
+ # Do croping
+ shape = tf.cast(image_info[1], dtype=tf.int32)
+ h = tf.random.uniform([],
+ 384,
+ tf.math.minimum(shape[0], 600),
+ dtype=tf.int32)
+ w = tf.random.uniform([],
+ 384,
+ tf.math.minimum(shape[1], 600),
+ dtype=tf.int32)
+ i = tf.random.uniform([], 0, shape[0] - h + 1, dtype=tf.int32)
+ j = tf.random.uniform([], 0, shape[1] - w + 1, dtype=tf.int32)
+ image = tf.image.crop_to_bounding_box(image, i, j, h, w)
+ boxes = tf.clip_by_value(
+ (boxes[..., :] * tf.cast(
+ tf.stack([shape[0], shape[1], shape[0], shape[1]]),
+ dtype=tf.float32) -
+ tf.cast(tf.stack([i, j, i, j]), dtype=tf.float32)) /
+ tf.cast(tf.stack([h, w, h, w]), dtype=tf.float32), 0.0, 1.0)
+ scales = tf.constant(self._resize_scales, dtype=tf.float32)
+ index = tf.random.categorical(tf.zeros([1, 11]), 1)[0]
+ scales = tf.gather(scales, index, axis=0)
+
+ image_shape = tf.shape(image)[:2]
+ boxes = box_ops.denormalize_boxes(boxes, image_shape)
+ short_side = scales[0]
+ image, image_info = preprocess_ops.resize_image(image, short_side,
+ max(self._output_size))
+ boxes = preprocess_ops.resize_and_crop_boxes(boxes, image_info[2, :],
+ image_info[1, :],
+ image_info[3, :])
+ boxes = box_ops.normalize_boxes(boxes, image_info[1, :])
+
+ # Filters out ground truth boxes that are all zeros.
+ indices = box_ops.get_non_empty_box_indices(boxes)
+ boxes = tf.gather(boxes, indices)
+ classes = tf.gather(classes, indices)
+ is_crowd = tf.gather(is_crowd, indices)
+ boxes = box_ops.yxyx_to_cycxhw(boxes)
+
+ image = tf.image.pad_to_bounding_box(image, 0, 0, self._output_size[0],
+ self._output_size[1])
+ labels = {
+ 'classes':
+ preprocess_ops.clip_or_pad_to_fixed_size(classes,
+ self._max_num_boxes),
+ 'boxes':
+ preprocess_ops.clip_or_pad_to_fixed_size(boxes, self._max_num_boxes)
+ }
+
+ return image, labels
+
+ def _parse_eval_data(self, data):
+ """Parses data for training and evaluation."""
+ classes = data['groundtruth_classes']
+ boxes = data['groundtruth_boxes']
+ is_crowd = data['groundtruth_is_crowd']
+
+ # Gets original image and its size.
+ image = data['image']
+
+ # Normalizes image with mean and std pixel values.
+ image = preprocess_ops.normalize_image(image)
+
+ scales = tf.constant([self._resize_scales[-1]], tf.float32)
+
+ image_shape = tf.shape(image)[:2]
+ boxes = box_ops.denormalize_boxes(boxes, image_shape)
+ gt_boxes = boxes
+ short_side = scales[0]
+ image, image_info = preprocess_ops.resize_image(image, short_side,
+ max(self._output_size))
+ boxes = preprocess_ops.resize_and_crop_boxes(boxes, image_info[2, :],
+ image_info[1, :],
+ image_info[3, :])
+ boxes = box_ops.normalize_boxes(boxes, image_info[1, :])
+
+ # Filters out ground truth boxes that are all zeros.
+ indices = box_ops.get_non_empty_box_indices(boxes)
+ boxes = tf.gather(boxes, indices)
+ classes = tf.gather(classes, indices)
+ is_crowd = tf.gather(is_crowd, indices)
+ boxes = box_ops.yxyx_to_cycxhw(boxes)
+
+ image = tf.image.pad_to_bounding_box(image, 0, 0, self._output_size[0],
+ self._output_size[1])
+ labels = {
+ 'classes':
+ preprocess_ops.clip_or_pad_to_fixed_size(classes,
+ self._max_num_boxes),
+ 'boxes':
+ preprocess_ops.clip_or_pad_to_fixed_size(boxes, self._max_num_boxes)
+ }
+ labels.update({
+ 'id':
+ int(data['source_id']),
+ 'image_info':
+ image_info,
+ 'is_crowd':
+ preprocess_ops.clip_or_pad_to_fixed_size(is_crowd,
+ self._max_num_boxes),
+ 'gt_boxes':
+ preprocess_ops.clip_or_pad_to_fixed_size(gt_boxes,
+ self._max_num_boxes),
+ })
+
+ return image, labels
diff --git a/official/projects/detr/experiments/__init__.py b/official/projects/detr/experiments/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/experiments/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/experiments/detr_r50_300epochs.sh b/official/projects/detr/experiments/detr_r50_300epochs.sh
index 98c52ce3c44..162f974306b 100644
--- a/official/projects/detr/experiments/detr_r50_300epochs.sh
+++ b/official/projects/detr/experiments/detr_r50_300epochs.sh
@@ -3,4 +3,4 @@ python3 official/projects/detr/train.py \
--experiment=detr_coco \
--mode=train_and_eval \
--model_dir=/tmp/logging_dir/ \
- --params_override=task.init_ckpt='gs://tf_model_garden/vision/resnet50_imagenet/ckpt-62400',trainer.train_steps=554400
+ --params_override=task.init_checkpoint='gs://tf_model_garden/vision/resnet50_imagenet/ckpt-62400',trainer.train_steps=554400,trainer.optimizer_config.learning_rate.stepwise.boundaries="[369600]"
diff --git a/official/projects/detr/experiments/detr_r50_500epochs.sh b/official/projects/detr/experiments/detr_r50_500epochs.sh
index f3febd6ecdd..58036040578 100644
--- a/official/projects/detr/experiments/detr_r50_500epochs.sh
+++ b/official/projects/detr/experiments/detr_r50_500epochs.sh
@@ -3,4 +3,4 @@ python3 official/projects/detr/train.py \
--experiment=detr_coco \
--mode=train_and_eval \
--model_dir=/tmp/logging_dir/ \
- --params_override=task.init_ckpt='gs://tf_model_garden/vision/resnet50_imagenet/ckpt-62400'
+ --params_override=task.init_checkpoint='gs://tf_model_garden/vision/resnet50_imagenet/ckpt-62400'
diff --git a/official/projects/detr/modeling/__init__.py b/official/projects/detr/modeling/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/modeling/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/modeling/detr.py b/official/projects/detr/modeling/detr.py
index 63af00c2db3..224b8d620df 100644
--- a/official/projects/detr/modeling/detr.py
+++ b/official/projects/detr/modeling/detr.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,11 +20,13 @@
for graph serializaiton.
"""
import math
-import tensorflow as tf
+from typing import Any, List
+
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.projects.detr.modeling import transformer
-from official.vision.modeling.backbones import resnet
+from official.vision.ops import box_ops
def position_embedding_sine(attention_mask,
@@ -93,14 +95,48 @@ def position_embedding_sine(attention_mask,
return embeddings
-class DETR(tf.keras.Model):
+def postprocess(outputs: dict[str, tf.Tensor]) -> dict[str, tf.Tensor]:
+ """Performs post-processing on model output.
+
+ Args:
+ outputs: The raw model output.
+
+ Returns:
+ Postprocessed model output.
+ """
+ predictions = {
+ "detection_boxes": # Box coordinates are relative values here.
+ box_ops.cycxhw_to_yxyx(outputs["box_outputs"]),
+ "detection_scores":
+ tf.math.reduce_max(
+ tf.nn.softmax(outputs["cls_outputs"])[:, :, 1:], axis=-1),
+ "detection_classes":
+ tf.math.argmax(outputs["cls_outputs"][:, :, 1:], axis=-1) + 1,
+ # Fix this. It's not being used at the moment.
+ "num_detections":
+ tf.reduce_sum(
+ tf.cast(
+ tf.math.greater(
+ tf.math.reduce_max(outputs["cls_outputs"], axis=-1), 0),
+ tf.int32),
+ axis=-1)
+ }
+ return predictions
+
+
+class DETR(tf_keras.Model):
"""DETR model with Keras.
DETR consists of backbone, query embedding, DETRTransformer,
class and box heads.
"""
- def __init__(self, num_queries, hidden_size, num_classes,
+ def __init__(self,
+ backbone,
+ backbone_endpoint_name,
+ num_queries,
+ hidden_size,
+ num_classes,
num_encoder_layers=6,
num_decoder_layers=6,
dropout_rate=0.1,
@@ -114,13 +150,17 @@ def __init__(self, num_queries, hidden_size, num_classes,
self._dropout_rate = dropout_rate
if hidden_size % 2 != 0:
raise ValueError("hidden_size must be a multiple of 2.")
- # TODO(frederickliu): Consider using the backbone factory.
- # TODO(frederickliu): Add to factory once we get skeleton code in.
- self._backbone = resnet.ResNet(50, bn_trainable=False)
+ self._backbone = backbone
+ self._backbone_endpoint_name = backbone_endpoint_name
def build(self, input_shape=None):
- self._input_proj = tf.keras.layers.Conv2D(
+ self._input_proj = tf_keras.layers.Conv2D(
self._hidden_size, 1, name="detr/conv2d")
+ self._build_detection_decoder()
+ super().build(input_shape)
+
+ def _build_detection_decoder(self):
+ """Builds detection decoder."""
self._transformer = DETRTransformer(
num_encoder_layers=self._num_encoder_layers,
num_decoder_layers=self._num_decoder_layers,
@@ -128,37 +168,38 @@ def build(self, input_shape=None):
self._query_embeddings = self.add_weight(
"detr/query_embeddings",
shape=[self._num_queries, self._hidden_size],
- initializer=tf.keras.initializers.RandomNormal(mean=0., stddev=1.),
+ initializer=tf_keras.initializers.RandomNormal(mean=0., stddev=1.),
dtype=tf.float32)
sqrt_k = math.sqrt(1.0 / self._hidden_size)
- self._class_embed = tf.keras.layers.Dense(
+ self._class_embed = tf_keras.layers.Dense(
self._num_classes,
- kernel_initializer=tf.keras.initializers.RandomUniform(-sqrt_k, sqrt_k),
+ kernel_initializer=tf_keras.initializers.RandomUniform(-sqrt_k, sqrt_k),
name="detr/cls_dense")
self._bbox_embed = [
- tf.keras.layers.Dense(
+ tf_keras.layers.Dense(
self._hidden_size, activation="relu",
- kernel_initializer=tf.keras.initializers.RandomUniform(
+ kernel_initializer=tf_keras.initializers.RandomUniform(
-sqrt_k, sqrt_k),
name="detr/box_dense_0"),
- tf.keras.layers.Dense(
+ tf_keras.layers.Dense(
self._hidden_size, activation="relu",
- kernel_initializer=tf.keras.initializers.RandomUniform(
+ kernel_initializer=tf_keras.initializers.RandomUniform(
-sqrt_k, sqrt_k),
name="detr/box_dense_1"),
- tf.keras.layers.Dense(
- 4, kernel_initializer=tf.keras.initializers.RandomUniform(
+ tf_keras.layers.Dense(
+ 4, kernel_initializer=tf_keras.initializers.RandomUniform(
-sqrt_k, sqrt_k),
name="detr/box_dense_2")]
- self._sigmoid = tf.keras.layers.Activation("sigmoid")
- super().build(input_shape)
+ self._sigmoid = tf_keras.layers.Activation("sigmoid")
@property
- def backbone(self) -> tf.keras.Model:
+ def backbone(self) -> tf_keras.Model:
return self._backbone
def get_config(self):
return {
+ "backbone": self._backbone,
+ "backbone_endpoint_name": self._backbone_endpoint_name,
"num_queries": self._num_queries,
"hidden_size": self._hidden_size,
"num_classes": self._num_classes,
@@ -168,18 +209,24 @@ def get_config(self):
}
@classmethod
- def from_config(cls, config):
+ def from_config(cls, config): # pyrefly: ignore[bad-override]
return cls(**config)
- def call(self, inputs):
- batch_size = tf.shape(inputs)[0]
+ def _generate_image_mask(self, inputs: tf.Tensor,
+ target_shape: tf.Tensor) -> tf.Tensor:
+ """Generates image mask from input image."""
mask = tf.expand_dims(
tf.cast(tf.not_equal(tf.reduce_sum(inputs, axis=-1), 0), inputs.dtype),
axis=-1)
- features = self._backbone(inputs)["5"]
- shape = tf.shape(features)
mask = tf.image.resize(
- mask, shape[1:3], method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)
+ mask, target_shape, method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)
+ return mask
+
+ def call(self, inputs: tf.Tensor, training: bool = None) -> List[Any]: # pytype: disable=annotation-type-mismatch,signature-mismatch
+ batch_size = tf.shape(inputs)[0]
+ features = self._backbone(inputs)[self._backbone_endpoint_name]
+ shape = tf.shape(features)
+ mask = self._generate_image_mask(inputs, shape[1: 3])
pos_embed = position_embedding_sine(
mask[:, :, :, 0], num_pos_features=self._hidden_size)
@@ -208,34 +255,55 @@ def call(self, inputs):
box_out = layer(box_out)
output_coord = self._sigmoid(box_out)
out = {"cls_outputs": output_class, "box_outputs": output_coord}
+ if not training:
+ out.update(postprocess(out))
out_list.append(out)
+
return out_list
-class DETRTransformer(tf.keras.layers.Layer):
+class DETRTransformer(tf_keras.layers.Layer):
"""Encoder and Decoder of DETR."""
- def __init__(self, num_encoder_layers=6, num_decoder_layers=6,
- dropout_rate=0.1, **kwargs):
+ def __init__(
+ self,
+ num_encoder_layers=6,
+ num_decoder_layers=6,
+ num_attention_heads=8,
+ intermediate_size=2048,
+ dropout_rate=0.1,
+ **kwargs
+ ):
super().__init__(**kwargs)
self._dropout_rate = dropout_rate
self._num_encoder_layers = num_encoder_layers
self._num_decoder_layers = num_decoder_layers
+ self._num_attention_heads = num_attention_heads
+ self._intermediate_size = intermediate_size
def build(self, input_shape=None):
- self._encoder = transformer.TransformerEncoder(
- attention_dropout_rate=self._dropout_rate,
- dropout_rate=self._dropout_rate,
- intermediate_dropout=self._dropout_rate,
- norm_first=False,
- num_layers=self._num_encoder_layers,
- )
+ if self._num_encoder_layers > 0:
+ self._encoder = transformer.TransformerEncoder(
+ attention_dropout_rate=self._dropout_rate,
+ dropout_rate=self._dropout_rate,
+ intermediate_dropout=self._dropout_rate,
+ norm_first=False,
+ num_layers=self._num_encoder_layers,
+ num_attention_heads=self._num_attention_heads,
+ intermediate_size=self._intermediate_size,
+ )
+ else:
+ self._encoder = None
+
self._decoder = transformer.TransformerDecoder(
attention_dropout_rate=self._dropout_rate,
dropout_rate=self._dropout_rate,
intermediate_dropout=self._dropout_rate,
norm_first=False,
- num_layers=self._num_decoder_layers)
+ num_layers=self._num_decoder_layers,
+ num_attention_heads=self._num_attention_heads,
+ intermediate_size=self._intermediate_size,
+ )
super().build(input_shape)
def get_config(self):
@@ -253,8 +321,12 @@ def call(self, inputs):
input_shape = tf_utils.get_shape_list(sources)
source_attention_mask = tf.tile(
tf.expand_dims(mask, axis=1), [1, input_shape[1], 1])
- memory = self._encoder(
- sources, attention_mask=source_attention_mask, pos_embed=pos_embed)
+ if self._encoder is not None:
+ memory = self._encoder(
+ sources, attention_mask=source_attention_mask, pos_embed=pos_embed)
+ else:
+ memory = sources
+
target_shape = tf_utils.get_shape_list(targets)
cross_attention_mask = tf.tile(
tf.expand_dims(mask, axis=1), [1, target_shape[1], 1])
diff --git a/official/projects/detr/modeling/detr_test.py b/official/projects/detr/modeling/detr_test.py
index 31a18ee85d7..2da6a02e157 100644
--- a/official/projects/detr/modeling/detr_test.py
+++ b/official/projects/detr/modeling/detr_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,8 +13,9 @@
# limitations under the License.
"""Tests for tensorflow_models.official.projects.detr.detr."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.detr.modeling import detr
+from official.vision.modeling.backbones import resnet
class DetrTest(tf.test.TestCase):
@@ -25,7 +26,10 @@ def test_forward(self):
num_classes = 10
image_size = 640
batch_size = 2
- model = detr.DETR(num_queries, hidden_size, num_classes)
+ backbone = resnet.ResNet(50, bn_trainable=False)
+ backbone_endpoint_name = '5'
+ model = detr.DETR(backbone, backbone_endpoint_name, num_queries,
+ hidden_size, num_classes)
outs = model(tf.ones((batch_size, image_size, image_size, 3)))
self.assertLen(outs, 6) # intermediate decoded outputs.
for out in outs:
@@ -47,6 +51,8 @@ def test_get_from_config_detr_transformer(self):
def test_get_from_config_detr(self):
config = {
+ 'backbone': resnet.ResNet(50, bn_trainable=False),
+ 'backbone_endpoint_name': '5',
'num_queries': 2,
'hidden_size': 4,
'num_classes': 10,
diff --git a/official/projects/detr/modeling/transformer.py b/official/projects/detr/modeling/transformer.py
index 06a419874a5..4f74c13b01d 100644
--- a/official/projects/detr/modeling/transformer.py
+++ b/official/projects/detr/modeling/transformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,13 +18,14 @@
cross-attention layer.
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
from official.nlp.modeling import models
-class TransformerEncoder(tf.keras.layers.Layer):
+class TransformerEncoder(tf_keras.layers.Layer):
"""Transformer encoder.
Transformer encoder is made up of N identical layers. Each layer is composed
@@ -61,7 +62,7 @@ def __init__(self,
layers is normalized.
norm_epsilon: Epsilon value to initialize normalization layers.
intermediate_dropout: Dropout probability for intermediate_dropout_layer.
- **kwargs: key word arguemnts passed to tf.keras.layers.Layer.
+ **kwargs: key word arguemnts passed to tf_keras.layers.Layer.
"""
super(TransformerEncoder, self).__init__(**kwargs)
@@ -91,10 +92,11 @@ def build(self, input_shape):
norm_first=self._norm_first,
norm_epsilon=self._norm_epsilon,
inner_dropout=self._intermediate_dropout,
- attention_initializer=models.seq2seq_transformer
- .attention_initializer(input_shape[2]),
+ attention_initializer=tf_utils.clone_initializer(
+ models.seq2seq_transformer.attention_initializer(
+ input_shape[2])),
name=("layer_%d" % i)))
- self.output_normalization = tf.keras.layers.LayerNormalization(
+ self.output_normalization = tf_keras.layers.LayerNormalization(
epsilon=self._norm_epsilon, dtype="float32")
super(TransformerEncoder, self).build(input_shape)
@@ -138,12 +140,12 @@ def call(self, encoder_inputs, attention_mask=None, pos_embed=None):
return output_tensor
-class TransformerEncoderBlock(tf.keras.layers.Layer):
+class TransformerEncoderBlock(tf_keras.layers.Layer):
"""TransformerEncoderBlock layer.
This layer implements the Transformer Encoder from
"Attention Is All You Need". (https://arxiv.org/abs/1706.03762),
- which combines a `tf.keras.layers.MultiHeadAttention` layer with a
+ which combines a `tf_keras.layers.MultiHeadAttention` layer with a
two-layer feedforward network. The only difference: position embedding is
added to the query and key of self-attention.
@@ -219,22 +221,23 @@ def __init__(self,
self._output_dropout = output_dropout
self._output_dropout_rate = output_dropout
self._output_range = output_range
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
self._inner_dropout = inner_dropout
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
+ self._attention_initializer = tf_keras.initializers.get(
attention_initializer)
else:
- self._attention_initializer = self._kernel_initializer
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer)
self._attention_axes = attention_axes
def build(self, input_shape):
@@ -262,7 +265,7 @@ def build(self, input_shape):
activity_regularizer=self._activity_regularizer,
kernel_constraint=self._kernel_constraint,
bias_constraint=self._bias_constraint)
- self._attention_layer = tf.keras.layers.MultiHeadAttention(
+ self._attention_layer = tf_keras.layers.MultiHeadAttention(
num_heads=self._num_heads,
key_dim=self._attention_head_size,
dropout=self._attention_dropout,
@@ -271,42 +274,42 @@ def build(self, input_shape):
attention_axes=self._attention_axes,
name="self_attention",
**common_kwargs)
- self._attention_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ self._attention_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype=tf.float32))
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, self._inner_dim),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
name="intermediate",
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
# TODO(b/154538392): Investigate this.
policy = tf.float32
- self._intermediate_activation_layer = tf.keras.layers.Activation(
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
self._inner_activation, dtype=policy)
- self._inner_dropout_layer = tf.keras.layers.Dropout(
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
rate=self._inner_dropout)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._output_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, hidden_size),
bias_axes="d",
name="output",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
**common_kwargs)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -316,44 +319,41 @@ def build(self, input_shape):
def get_config(self):
config = {
- "num_attention_heads":
- self._num_heads,
- "inner_dim":
- self._inner_dim,
- "inner_activation":
- self._inner_activation,
- "output_dropout":
- self._output_dropout_rate,
- "attention_dropout":
- self._attention_dropout_rate,
- "output_range":
- self._output_range,
- "kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
- "bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
- "kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
- "bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
- "activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
- "kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
- "bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
- "use_bias":
- self._use_bias,
- "norm_first":
- self._norm_first,
- "norm_epsilon":
- self._norm_epsilon,
- "inner_dropout":
- self._inner_dropout,
- "attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer),
- "attention_axes":
- self._attention_axes,
+ "num_attention_heads": self._num_heads,
+ "inner_dim": self._inner_dim,
+ "inner_activation": self._inner_activation,
+ "output_dropout": self._output_dropout_rate,
+ "attention_dropout": self._attention_dropout_rate,
+ "output_range": self._output_range,
+ "kernel_initializer": tf_utils.serialize_initializer(
+ self._kernel_initializer, use_legacy_format=True
+ ),
+ "bias_initializer": tf_utils.serialize_initializer(
+ self._bias_initializer, use_legacy_format=True
+ ),
+ "kernel_regularizer": tf_utils.serialize_regularizer(
+ self._kernel_regularizer, use_legacy_format=True
+ ),
+ "bias_regularizer": tf_utils.serialize_regularizer(
+ self._bias_regularizer, use_legacy_format=True
+ ),
+ "activity_regularizer": tf_utils.serialize_regularizer(
+ self._activity_regularizer, use_legacy_format=True
+ ),
+ "kernel_constraint": tf_utils.serialize_constraint(
+ self._kernel_constraint, use_legacy_format=True
+ ),
+ "bias_constraint": tf_utils.serialize_constraint(
+ self._bias_constraint, use_legacy_format=True
+ ),
+ "use_bias": self._use_bias,
+ "norm_first": self._norm_first,
+ "norm_epsilon": self._norm_epsilon,
+ "inner_dropout": self._inner_dropout,
+ "attention_initializer": tf_utils.serialize_initializer(
+ self._attention_initializer, use_legacy_format=True
+ ),
+ "attention_axes": self._attention_axes,
}
base_config = super(TransformerEncoderBlock, self).get_config()
return dict(list(base_config.items()) + list(config.items()))
@@ -398,9 +398,9 @@ def call(self, inputs):
key=key_value + pos_embed,
value=key_value,
attention_mask=attention_mask)
- attention_output = self._attention_dropout(attention_output)
+ attention_output = self._attention_dropout(attention_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- attention_output = source_tensor + attention_output
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self._attention_layer_norm(target_tensor +
attention_output)
@@ -411,10 +411,10 @@ def call(self, inputs):
inner_output = self._intermediate_activation_layer(inner_output)
inner_output = self._inner_dropout_layer(inner_output)
layer_output = self._output_dense(inner_output)
- layer_output = self._output_dropout(layer_output)
+ layer_output = self._output_dropout(layer_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- return source_attention_output + layer_output
+ return source_attention_output + layer_output # pyrefly: ignore[unbound-name]
# During mixed precision training, layer norm output is always fp32 for now.
# Casts fp32 for the subsequent add.
@@ -422,7 +422,7 @@ def call(self, inputs):
return self._output_layer_norm(layer_output + attention_output)
-class TransformerDecoder(tf.keras.layers.Layer):
+class TransformerDecoder(tf_keras.layers.Layer):
"""Transformer decoder.
Like the encoder, the decoder is made up of N identical layers.
@@ -461,7 +461,7 @@ def __init__(self,
layers is normalized.
norm_epsilon: Epsilon value to initialize normalization layers.
intermediate_dropout: Dropout probability for intermediate_dropout_layer.
- **kwargs: key word arguemnts passed to tf.keras.layers.Layer.
+ **kwargs: key word arguemnts passed to tf_keras.layers.Layer.
"""
super(TransformerDecoder, self).__init__(**kwargs)
self.num_layers = num_layers
@@ -490,10 +490,11 @@ def build(self, input_shape):
norm_first=self._norm_first,
norm_epsilon=self._norm_epsilon,
intermediate_dropout=self._intermediate_dropout,
- attention_initializer=models.seq2seq_transformer
- .attention_initializer(input_shape[2]),
+ attention_initializer=tf_utils.clone_initializer(
+ models.seq2seq_transformer.attention_initializer(
+ input_shape[2])),
name=("layer_%d" % i)))
- self.output_normalization = tf.keras.layers.LayerNormalization(
+ self.output_normalization = tf_keras.layers.LayerNormalization(
epsilon=self._norm_epsilon, dtype="float32")
super(TransformerDecoder, self).build(input_shape)
@@ -577,7 +578,7 @@ def call(self,
return self.output_normalization(output_tensor)
-class TransformerDecoderBlock(tf.keras.layers.Layer):
+class TransformerDecoderBlock(tf_keras.layers.Layer):
"""Single transformer layer for decoder.
It has three sub-layers:
@@ -632,31 +633,32 @@ def __init__(self,
attention_initializer: Initializer for kernels of attention layers. If set
`None`, attention layers use kernel_initializer as initializer for
kernel.
- **kwargs: key word arguemnts passed to tf.keras.layers.Layer.
+ **kwargs: key word arguemnts passed to tf_keras.layers.Layer.
"""
super().__init__(**kwargs)
self.num_attention_heads = num_attention_heads
self.intermediate_size = intermediate_size
- self.intermediate_activation = tf.keras.activations.get(
+ self.intermediate_activation = tf_keras.activations.get(
intermediate_activation)
self.dropout_rate = dropout_rate
self.attention_dropout_rate = attention_dropout_rate
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
self._intermediate_dropout = intermediate_dropout
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
+ self._attention_initializer = tf_keras.initializers.get(
attention_initializer)
else:
- self._attention_initializer = self._kernel_initializer
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer)
self._cross_attention_cls = layers.attention.MultiHeadAttention
def build(self, input_shape):
@@ -686,17 +688,17 @@ def build(self, input_shape):
kernel_initializer=self._attention_initializer,
name="self_attention",
**common_kwargs)
- self.self_attention_output_dense = tf.keras.layers.experimental.EinsumDense(
+ self.self_attention_output_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
name="output",
**common_kwargs)
- self.self_attention_dropout = tf.keras.layers.Dropout(
+ self.self_attention_dropout = tf_keras.layers.Dropout(
rate=self.dropout_rate)
self.self_attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -712,36 +714,36 @@ def build(self, input_shape):
name="attention/encdec",
**common_kwargs)
- self.encdec_attention_dropout = tf.keras.layers.Dropout(
+ self.encdec_attention_dropout = tf_keras.layers.Dropout(
rate=self.dropout_rate)
self.encdec_attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="attention/encdec_output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype="float32"))
# Feed-forward projection.
- self.intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self.intermediate_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, self.intermediate_size),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
name="intermediate",
**common_kwargs)
- self.intermediate_activation_layer = tf.keras.layers.Activation(
+ self.intermediate_activation_layer = tf_keras.layers.Activation(
self.intermediate_activation)
- self._intermediate_dropout_layer = tf.keras.layers.Dropout(
+ self._intermediate_dropout_layer = tf_keras.layers.Dropout(
rate=self._intermediate_dropout)
- self.output_dense = tf.keras.layers.experimental.EinsumDense(
+ self.output_dense = tf_keras.layers.EinsumDense(
"abc,cd->abd",
output_shape=(None, hidden_size),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
name="output",
**common_kwargs)
- self.output_dropout = tf.keras.layers.Dropout(rate=self.dropout_rate)
- self.output_layer_norm = tf.keras.layers.LayerNormalization(
+ self.output_dropout = tf_keras.layers.Dropout(rate=self.dropout_rate)
+ self.output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -750,40 +752,41 @@ def build(self, input_shape):
def get_config(self):
config = {
- "num_attention_heads":
- self.num_attention_heads,
- "intermediate_size":
- self.intermediate_size,
- "intermediate_activation":
- tf.keras.activations.serialize(self.intermediate_activation),
- "dropout_rate":
- self.dropout_rate,
- "attention_dropout_rate":
- self.attention_dropout_rate,
- "kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
- "bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
- "kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
- "bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
- "activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
- "kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
- "bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
- "use_bias":
- self._use_bias,
- "norm_first":
- self._norm_first,
- "norm_epsilon":
- self._norm_epsilon,
- "intermediate_dropout":
- self._intermediate_dropout,
- "attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer)
+ "num_attention_heads": self.num_attention_heads,
+ "intermediate_size": self.intermediate_size,
+ "intermediate_activation": tf_utils.serialize_activation(
+ self.intermediate_activation, use_legacy_format=True
+ ),
+ "dropout_rate": self.dropout_rate,
+ "attention_dropout_rate": self.attention_dropout_rate,
+ "kernel_initializer": tf_utils.serialize_initializer(
+ self._kernel_initializer, use_legacy_format=True
+ ),
+ "bias_initializer": tf_utils.serialize_initializer(
+ self._bias_initializer, use_legacy_format=True
+ ),
+ "kernel_regularizer": tf_utils.serialize_regularizer(
+ self._kernel_regularizer, use_legacy_format=True
+ ),
+ "bias_regularizer": tf_utils.serialize_regularizer(
+ self._bias_regularizer, use_legacy_format=True
+ ),
+ "activity_regularizer": tf_utils.serialize_regularizer(
+ self._activity_regularizer, use_legacy_format=True
+ ),
+ "kernel_constraint": tf_utils.serialize_constraint(
+ self._kernel_constraint, use_legacy_format=True
+ ),
+ "bias_constraint": tf_utils.serialize_constraint(
+ self._bias_constraint, use_legacy_format=True
+ ),
+ "use_bias": self._use_bias,
+ "norm_first": self._norm_first,
+ "norm_epsilon": self._norm_epsilon,
+ "intermediate_dropout": self._intermediate_dropout,
+ "attention_initializer": tf_utils.serialize_initializer(
+ self._attention_initializer, use_legacy_format=True
+ ),
}
base_config = super().get_config()
return dict(list(base_config.items()) + list(config.items()))
@@ -825,7 +828,7 @@ def call(self, inputs, cache=None, decode_loop_step=None):
attention_output = self.encdec_attention(**cross_attn_inputs)
attention_output = self.encdec_attention_dropout(attention_output)
if self._norm_first:
- attention_output = source_self_attention_output + attention_output
+ attention_output = source_self_attention_output + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self.encdec_attention_layer_norm(
self_attention_output + attention_output)
@@ -840,7 +843,7 @@ def call(self, inputs, cache=None, decode_loop_step=None):
layer_output = self.output_dense(intermediate_output)
layer_output = self.output_dropout(layer_output)
if self._norm_first:
- layer_output = source_attention_output + layer_output
+ layer_output = source_attention_output + layer_output # pyrefly: ignore[unbound-name]
else:
layer_output = self.output_layer_norm(layer_output + attention_output)
return layer_output, cache
diff --git a/official/projects/detr/modeling/transformer_test.py b/official/projects/detr/modeling/transformer_test.py
index 0752403a2a8..83755b371d8 100644
--- a/official/projects/detr/modeling/transformer_test.py
+++ b/official/projects/detr/modeling/transformer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for transformer."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.detr.modeling import transformer
diff --git a/official/projects/detr/ops/__init__.py b/official/projects/detr/ops/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/ops/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/ops/matchers.py b/official/projects/detr/ops/matchers.py
index 56f25585ed2..5ba810d099f 100644
--- a/official/projects/detr/ops/matchers.py
+++ b/official/projects/detr/ops/matchers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,7 +26,7 @@
Based on the original implementation by Jiquan Ngiam .
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
diff --git a/official/projects/detr/ops/matchers_test.py b/official/projects/detr/ops/matchers_test.py
index 71c607123a3..8ae16c0668e 100644
--- a/official/projects/detr/ops/matchers_test.py
+++ b/official/projects/detr/ops/matchers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import numpy as np
from scipy import optimize
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.detr.ops import matchers
diff --git a/official/projects/detr/optimization.py b/official/projects/detr/optimization.py
index a9da1740d11..31199c24863 100644
--- a/official/projects/detr/optimization.py
+++ b/official/projects/detr/optimization.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Customized optimizer to match paper results."""
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import optimization
from official.nlp import optimization as nlp_optimization
@@ -27,7 +27,9 @@ class DETRAdamWConfig(optimization.AdamWeightDecayConfig):
@dataclasses.dataclass
class OptimizerConfig(optimization.OptimizerConfig):
- detr_adamw: DETRAdamWConfig = DETRAdamWConfig()
+ detr_adamw: DETRAdamWConfig = dataclasses.field(
+ default_factory=DETRAdamWConfig
+ )
@dataclasses.dataclass
@@ -41,7 +43,9 @@ class OptimizationConfig(optimization.OptimizationConfig):
learning_rate: learning rate oneof config.
warmup: warmup oneof config.
"""
- optimizer: OptimizerConfig = OptimizerConfig()
+ optimizer: OptimizerConfig = dataclasses.field(
+ default_factory=OptimizerConfig
+ )
# TODO(frederickliu): figure out how to make this configuable.
diff --git a/official/projects/detr/serving/__init__.py b/official/projects/detr/serving/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/serving/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/serving/export_module.py b/official/projects/detr/serving/export_module.py
new file mode 100644
index 00000000000..d0227da2f48
--- /dev/null
+++ b/official/projects/detr/serving/export_module.py
@@ -0,0 +1,103 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Export module for DETR model."""
+import tensorflow as tf, tf_keras
+
+from official.projects.detr.modeling import detr
+from official.vision.modeling import backbones
+from official.vision.ops import preprocess_ops
+from official.vision.serving import detection
+
+
+class DETRModule(detection.DetectionModule):
+ """DETR detection module."""
+
+ def _build_model(self) -> tf_keras.Model:
+ input_specs = tf_keras.layers.InputSpec(shape=[self._batch_size] +
+ self._input_image_size +
+ [self._num_channels])
+
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=self.params.task.model.backbone,
+ norm_activation_config=self.params.task.model.norm_activation)
+
+ model = detr.DETR(backbone, self.params.task.model.backbone_endpoint_name,
+ self.params.task.model.num_queries,
+ self.params.task.model.hidden_size,
+ self.params.task.model.num_classes,
+ self.params.task.model.num_encoder_layers,
+ self.params.task.model.num_decoder_layers)
+ model(tf_keras.Input(input_specs.shape[1:]))
+ return model
+
+ def _build_inputs(self, image: tf.Tensor) -> tuple[tf.Tensor, tf.Tensor]:
+ """Builds detection model inputs for serving."""
+ # Normalizes image with mean and std pixel values.
+ image = preprocess_ops.normalize_image(
+ image, offset=preprocess_ops.MEAN_RGB, scale=preprocess_ops.STDDEV_RGB)
+
+ image, image_info = preprocess_ops.resize_image(
+ image, size=self._input_image_size)
+
+ return image, image_info
+
+ def serve(self, images: tf.Tensor) -> dict[str, tf.Tensor]:
+ """Cast image to float and run inference.
+
+ Args:
+ images: uint8 Tensor of shape [batch_size, None, None, 3]
+
+ Returns:
+ Tensor holding classification output logits.
+ """
+ # Skip image preprocessing when input_type is tflite so it is compatible
+ # with TFLite quantization.
+ image_info = None
+ if self._input_type != 'tflite':
+ with tf.device('cpu:0'):
+ images = tf.cast(images, dtype=tf.float32)
+
+ images_spec = tf.TensorSpec(
+ shape=self._input_image_size + [3], dtype=tf.float32)
+ image_info_spec = tf.TensorSpec(shape=[4, 2], dtype=tf.float32)
+
+ images, image_info = tf.nest.map_structure(
+ tf.identity,
+ tf.map_fn(
+ self._build_inputs,
+ elems=images,
+ fn_output_signature=(images_spec, image_info_spec),
+ parallel_iterations=32))
+
+ outputs = self.inference_step(images)[-1]
+ outputs = {
+ 'detection_boxes': outputs['detection_boxes'],
+ 'detection_scores': outputs['detection_scores'],
+ 'detection_classes': outputs['detection_classes'],
+ 'num_detections': outputs['num_detections']
+ }
+ if image_info is not None:
+ outputs['detection_boxes'] = outputs['detection_boxes'] * tf.expand_dims(
+ tf.concat([
+ image_info[:, 1:2, 0], image_info[:, 1:2, 1],
+ image_info[:, 1:2, 0], image_info[:, 1:2, 1]
+ ],
+ axis=1),
+ axis=1)
+
+ outputs.update({'image_info': image_info})
+
+ return outputs
diff --git a/official/projects/detr/serving/export_module_test.py b/official/projects/detr/serving/export_module_test.py
new file mode 100644
index 00000000000..56105d83574
--- /dev/null
+++ b/official/projects/detr/serving/export_module_test.py
@@ -0,0 +1,98 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Test for DETR export module."""
+
+import io
+import os
+
+from absl.testing import parameterized
+import numpy as np
+from PIL import Image
+import tensorflow as tf, tf_keras
+
+from official.core import exp_factory
+from official.projects.detr.configs import detr as exp_cfg # pylint: disable=unused-import
+from official.projects.detr.serving import export_module
+
+
+class ExportModuleTest(tf.test.TestCase, parameterized.TestCase):
+
+ def _get_module(self, input_type):
+ params = exp_factory.get_exp_config('detr_coco')
+ return export_module.DETRModule(
+ params,
+ batch_size=1,
+ input_image_size=[384, 384],
+ input_type=input_type)
+
+ def _export_from_module(self, module, input_type, save_directory):
+ signatures = module.get_inference_signatures(
+ {input_type: 'serving_default'})
+ tf.saved_model.save(module, save_directory, signatures=signatures)
+
+ def _get_dummy_input(self, input_type):
+ """Gets dummy input for the given input type."""
+
+ if input_type == 'image_tensor':
+ return tf.zeros((1, 384, 384, 3), dtype=np.uint8)
+ elif input_type == 'image_bytes':
+ image = Image.fromarray(np.zeros((384, 384, 3), dtype=np.uint8))
+ byte_io = io.BytesIO()
+ image.save(byte_io, 'PNG')
+ return [byte_io.getvalue()]
+ elif input_type == 'tf_example':
+ image_tensor = tf.zeros((384, 384, 3), dtype=tf.uint8)
+ encoded_jpeg = tf.image.encode_jpeg(tf.constant(image_tensor)).numpy()
+ example = tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[encoded_jpeg])),
+ })).SerializeToString()
+ return [example]
+
+ @parameterized.parameters(
+ {'input_type': 'image_tensor'},
+ {'input_type': 'image_bytes'},
+ {'input_type': 'tf_example'},
+ )
+ def test_export(self, input_type='image_tensor'):
+ tmp_dir = self.get_temp_dir()
+ module = self._get_module(input_type)
+ self._export_from_module(module, input_type, tmp_dir)
+
+ self.assertTrue(os.path.exists(os.path.join(tmp_dir, 'saved_model.pb')))
+ self.assertTrue(
+ os.path.exists(os.path.join(tmp_dir, 'variables', 'variables.index')))
+ self.assertTrue(
+ os.path.exists(
+ os.path.join(tmp_dir, 'variables',
+ 'variables.data-00000-of-00001')))
+
+ imported = tf.saved_model.load(tmp_dir)
+ predict_fn = imported.signatures['serving_default']
+
+ images = self._get_dummy_input(input_type)
+ outputs = predict_fn(tf.constant(images))
+
+ self.assertNotEmpty(outputs['detection_boxes'])
+ self.assertNotEmpty(outputs['detection_classes'])
+ self.assertNotEmpty(outputs['detection_scores'])
+ self.assertNotEmpty(outputs['num_detections'])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/detr/serving/export_saved_model.py b/official/projects/detr/serving/export_saved_model.py
new file mode 100644
index 00000000000..f21290c08df
--- /dev/null
+++ b/official/projects/detr/serving/export_saved_model.py
@@ -0,0 +1,109 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+r"""Vision models export binary for serving/inference.
+
+To export a trained checkpoint in saved_model format (shell script):
+
+EXPERIMENT_TYPE = XX
+CHECKPOINT_PATH = XX
+EXPORT_DIR_PATH = XX
+export_saved_model --experiment=${EXPERIMENT_TYPE} \
+ --export_dir=${EXPORT_DIR_PATH}/ \
+ --checkpoint_path=${CHECKPOINT_PATH} \
+ --batch_size=2 \
+ --input_image_size=224,224
+
+To serve (python):
+
+export_dir_path = XX
+input_type = XX
+input_images = XX
+imported = tf.saved_model.load(export_dir_path)
+model_fn = imported.signatures['serving_default']
+output = model_fn(input_images)
+"""
+
+from absl import app
+from absl import flags
+
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.detr.configs import detr as exp_cfg # pylint: disable=unused-import
+from official.projects.detr.serving import export_module
+from official.vision.serving import export_saved_model_lib
+
+FLAGS = flags.FLAGS
+
+_EXPERIMENT = flags.DEFINE_string('experiment', None,
+ 'experiment type, e.g. detr_coco')
+_EXPORT_DIR = flags.DEFINE_string('export_dir', None, 'The export directory.')
+_CHECKPOINT_PATH = flags.DEFINE_string('checkpoint_path', None,
+ 'Checkpoint path.')
+_CONFIG_FILE = flags.DEFINE_multi_string(
+ 'config_file',
+ default=None,
+ help='YAML/JSON files which specifies overrides. The override order '
+ 'follows the order of args. Note that each file '
+ 'can be used as an override template to override the default parameters '
+ 'specified in Python. If the same parameter is specified in both '
+ '`--config_file` and `--params_override`, `config_file` will be used '
+ 'first, followed by params_override.')
+_PARAMS_OVERRIDE = flags.DEFINE_string(
+ 'params_override', '',
+ 'The JSON/YAML file or string which specifies the parameter to be overriden'
+ ' on top of `config_file` template.')
+_BATCH_SIZE = flags.DEFINE_integer('batch_size', None, 'The batch size.')
+_IMAGE_TYPE = flags.DEFINE_string(
+ 'input_type', 'image_tensor',
+ 'One of `image_tensor`, `image_bytes`, `tf_example` and `tflite`.')
+_INPUT_IMAGE_SIZE = flags.DEFINE_string(
+ 'input_image_size', '224,224',
+ 'The comma-separated string of two integers representing the height,width '
+ 'of the input to the model.')
+
+
+def main(_):
+
+ params = exp_factory.get_exp_config(_EXPERIMENT.value)
+ for config_file in _CONFIG_FILE.value or []:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=False)
+ if _PARAMS_OVERRIDE.value:
+ params = hyperparams.override_params_dict(
+ params, _PARAMS_OVERRIDE.value, is_strict=False)
+
+ params.validate()
+ params.lock()
+
+ input_image_size = [int(x) for x in _INPUT_IMAGE_SIZE.value.split(',')]
+ module = export_module.DETRModule(
+ params=params,
+ batch_size=_BATCH_SIZE.value,
+ input_image_size=input_image_size,
+ input_type=_IMAGE_TYPE.value,
+ num_channels=3)
+
+ export_saved_model_lib.export_inference_graph(
+ input_type=_IMAGE_TYPE.value,
+ batch_size=_BATCH_SIZE.value,
+ input_image_size=input_image_size,
+ params=params,
+ checkpoint_path=_CHECKPOINT_PATH.value,
+ export_dir=_EXPORT_DIR.value,
+ export_module=module)
+
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/detr/tasks/__init__.py b/official/projects/detr/tasks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/detr/tasks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/detr/tasks/detection.py b/official/projects/detr/tasks/detection.py
index 37a11d8b41a..27590fa33f5 100644
--- a/official/projects/detr/tasks/detection.py
+++ b/official/projects/detr/tasks/detection.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,21 +13,30 @@
# limitations under the License.
"""DETR detection task definition."""
+from typing import Optional
-import tensorflow as tf
+from absl import logging
+import tensorflow as tf, tf_keras
+from official.common import dataset_fn
from official.core import base_task
from official.core import task_factory
from official.projects.detr.configs import detr as detr_cfg
from official.projects.detr.dataloaders import coco
+from official.projects.detr.dataloaders import detr_input
from official.projects.detr.modeling import detr
from official.projects.detr.ops import matchers
+from official.vision.dataloaders import input_reader_factory
+from official.vision.dataloaders import tf_example_decoder
+from official.vision.dataloaders import tfds_factory
+from official.vision.dataloaders import tf_example_label_map_decoder
from official.vision.evaluation import coco_evaluator
+from official.vision.modeling import backbones
from official.vision.ops import box_ops
-@task_factory.register_task_cls(detr_cfg.DetectionConfig)
-class DectectionTask(base_task.Task):
+@task_factory.register_task_cls(detr_cfg.DetrTask)
+class DetectionTask(base_task.Task):
"""A single-replica view of training procedure.
DETR task provides artifacts for training/evalution procedures, including
@@ -37,62 +46,131 @@ class DectectionTask(base_task.Task):
def build_model(self):
"""Build DETR model."""
- model = detr.DETR(
- self._task_config.num_queries,
- self._task_config.num_hidden,
- self._task_config.num_classes,
- self._task_config.num_encoder_layers,
- self._task_config.num_decoder_layers)
+
+ input_specs = tf_keras.layers.InputSpec(shape=[None] +
+ self._task_config.model.input_size)
+
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=self._task_config.model.backbone,
+ norm_activation_config=self._task_config.model.norm_activation)
+
+ model = detr.DETR(backbone,
+ self._task_config.model.backbone_endpoint_name,
+ self._task_config.model.num_queries,
+ self._task_config.model.hidden_size,
+ self._task_config.model.num_classes,
+ self._task_config.model.num_encoder_layers,
+ self._task_config.model.num_decoder_layers)
return model
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loading pretrained checkpoint."""
- ckpt = tf.train.Checkpoint(backbone=model.backbone)
- status = ckpt.read(self._task_config.init_ckpt)
- status.expect_partial().assert_existing_objects_matched()
-
- def build_inputs(self, params, input_context=None):
+ if not self._task_config.init_checkpoint:
+ return
+
+ ckpt_dir_or_file = self._task_config.init_checkpoint
+
+ # Restoring checkpoint.
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+
+ if self._task_config.init_checkpoint_modules == 'all':
+ ckpt = tf.train.Checkpoint(**model.checkpoint_items)
+ status = ckpt.restore(ckpt_dir_or_file)
+ status.assert_consumed()
+ elif self._task_config.init_checkpoint_modules == 'backbone':
+ ckpt = tf.train.Checkpoint(backbone=model.backbone)
+ status = ckpt.restore(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
+
+ def build_inputs(self,
+ params,
+ input_context: Optional[tf.distribute.InputContext] = None):
"""Build input dataset."""
- return coco.COCODataLoader(params).load(input_context)
+ if isinstance(params, coco.COCODataConfig):
+ dataset = coco.COCODataLoader(params).load(input_context)
+ else:
+ if params.tfds_name:
+ decoder = tfds_factory.get_detection_decoder(params.tfds_name)
+ else:
+ decoder_cfg = params.decoder.get()
+ if params.decoder.type == 'simple_decoder':
+ decoder = tf_example_decoder.TfExampleDecoder(
+ regenerate_source_id=decoder_cfg.regenerate_source_id)
+ elif params.decoder.type == 'label_map_decoder':
+ decoder = tf_example_label_map_decoder.TfExampleDecoderLabelMap(
+ label_map=decoder_cfg.label_map,
+ regenerate_source_id=decoder_cfg.regenerate_source_id)
+ else:
+ raise ValueError('Unknown decoder type: {}!'.format(
+ params.decoder.type))
+
+ parser = detr_input.Parser(
+ class_offset=self._task_config.losses.class_offset,
+ output_size=self._task_config.model.input_size[:2],
+ )
+
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training))
+ dataset = reader.read(input_context=input_context)
+
+ return dataset
def _compute_cost(self, cls_outputs, box_outputs, cls_targets, box_targets):
# Approximate classification cost with 1 - prob[target class].
# The 1 is a constant that doesn't change the matching, it can be ommitted.
# background: 0
- cls_cost = self._task_config.lambda_cls * tf.gather(
- -tf.nn.softmax(cls_outputs), cls_targets, batch_dims=1, axis=-1)
+ cls_cost = self._task_config.losses.lambda_cls * tf.gather(
+ -tf.nn.softmax(cls_outputs), cls_targets, batch_dims=1, axis=-1
+ )
# Compute the L1 cost between boxes,
- paired_differences = self._task_config.lambda_box * tf.abs(
- tf.expand_dims(box_outputs, 2) - tf.expand_dims(box_targets, 1))
+ paired_differences = self._task_config.losses.lambda_box * tf.abs(
+ tf.expand_dims(box_outputs, 2) - tf.expand_dims(box_targets, 1)
+ )
box_cost = tf.reduce_sum(paired_differences, axis=-1)
# Compute the giou cost betwen boxes
- giou_cost = self._task_config.lambda_giou * -box_ops.bbox_generalized_overlap(
- box_ops.cycxhw_to_yxyx(box_outputs),
- box_ops.cycxhw_to_yxyx(box_targets))
+ giou_cost = (
+ self._task_config.losses.lambda_giou
+ * -box_ops.bbox_generalized_overlap(
+ box_ops.cycxhw_to_yxyx(box_outputs),
+ box_ops.cycxhw_to_yxyx(box_targets),
+ )
+ )
total_cost = cls_cost + box_cost + giou_cost
max_cost = (
- self._task_config.lambda_cls * 0.0 + self._task_config.lambda_box * 4. +
- self._task_config.lambda_giou * 0.0)
+ self._task_config.losses.lambda_cls * 1.0
+ + self._task_config.losses.lambda_box * 4.0
+ + self._task_config.losses.lambda_giou * 1.0
+ )
# Set pads to large constant
valid = tf.expand_dims(
- tf.cast(tf.not_equal(cls_targets, 0), dtype=total_cost.dtype), axis=1)
+ tf.cast(tf.not_equal(cls_targets, 0), dtype=total_cost.dtype), axis=1
+ )
total_cost = (1 - valid) * max_cost + valid * total_cost
# Set inf of nan to large constant
total_cost = tf.where(
tf.logical_or(tf.math.is_nan(total_cost), tf.math.is_inf(total_cost)),
max_cost * tf.ones_like(total_cost, dtype=total_cost.dtype),
- total_cost)
+ total_cost,
+ )
return total_cost
def build_losses(self, outputs, labels, aux_losses=None):
- """Build DETR losses."""
+ """Builds DETR losses."""
cls_outputs = outputs['cls_outputs']
box_outputs = outputs['box_outputs']
cls_targets = labels['classes']
@@ -115,35 +193,26 @@ def build_losses(self, outputs, labels, aux_losses=None):
# Down-weight background to account for class imbalance.
xentropy = tf.nn.sparse_softmax_cross_entropy_with_logits(
labels=cls_targets, logits=cls_assigned)
- cls_loss = self._task_config.lambda_cls * tf.where(
- background,
- self._task_config.background_cls_weight * xentropy,
- xentropy
- )
+ cls_loss = self._task_config.losses.lambda_cls * tf.where(
+ background, self._task_config.losses.background_cls_weight * xentropy,
+ xentropy)
cls_weights = tf.where(
background,
- self._task_config.background_cls_weight * tf.ones_like(cls_loss),
- tf.ones_like(cls_loss)
- )
+ self._task_config.losses.background_cls_weight * tf.ones_like(cls_loss),
+ tf.ones_like(cls_loss))
# Box loss is only calculated on non-background class.
l_1 = tf.reduce_sum(tf.abs(box_assigned - box_targets), axis=-1)
- box_loss = self._task_config.lambda_box * tf.where(
- background,
- tf.zeros_like(l_1),
- l_1
- )
+ box_loss = self._task_config.losses.lambda_box * tf.where(
+ background, tf.zeros_like(l_1), l_1)
# Giou loss is only calculated on non-background class.
giou = tf.linalg.diag_part(1.0 - box_ops.bbox_generalized_overlap(
box_ops.cycxhw_to_yxyx(box_assigned),
box_ops.cycxhw_to_yxyx(box_targets)
))
- giou_loss = self._task_config.lambda_giou * tf.where(
- background,
- tf.zeros_like(giou),
- giou
- )
+ giou_loss = self._task_config.losses.lambda_giou * tf.where(
+ background, tf.zeros_like(giou), giou)
# Consider doing all reduce once in train_step to speed up.
num_boxes_per_replica = tf.reduce_sum(num_boxes)
@@ -160,19 +229,20 @@ def build_losses(self, outputs, labels, aux_losses=None):
tf.reduce_sum(giou_loss), num_boxes_sum)
aux_losses = tf.add_n(aux_losses) if aux_losses else 0.0
+
total_loss = cls_loss + box_loss + giou_loss + aux_losses
return total_loss, cls_loss, box_loss, giou_loss
def build_metrics(self, training=True):
- """Build detection metrics."""
+ """Builds detection metrics."""
metrics = []
metric_names = ['cls_loss', 'box_loss', 'giou_loss']
for name in metric_names:
- metrics.append(tf.keras.metrics.Mean(name, dtype=tf.float32))
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
if not training:
self.coco_metric = coco_evaluator.COCOEvaluator(
- annotation_file='',
+ annotation_file=self._task_config.annotation_file,
include_mask=False,
need_rescale_bboxes=True,
per_category_metrics=self._task_config.per_category_metrics)
@@ -201,8 +271,11 @@ def train_step(self, inputs, model, optimizer, metrics=None):
for output in outputs:
# Computes per-replica loss.
- layer_loss, layer_cls_loss, layer_box_loss, layer_giou_loss = self.build_losses(
- outputs=output, labels=labels, aux_losses=model.losses)
+ layer_loss, layer_cls_loss, layer_box_loss, layer_giou_loss = (
+ self.build_losses(
+ outputs=output, labels=labels, aux_losses=model.losses
+ )
+ )
loss += layer_loss
cls_loss += layer_cls_loss
box_loss += layer_box_loss
@@ -212,13 +285,13 @@ def train_step(self, inputs, model, optimizer, metrics=None):
scaled_loss = loss
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
# Scales back gradient when LossScaleOptimizer is used.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
@@ -277,31 +350,50 @@ def validation_step(self, inputs, model, metrics=None):
# Evaluator class handles loss metric for you.
logs = {self.loss: loss}
+ # This is for backward compatibility.
+ if 'detection_boxes' not in outputs:
+ detection_boxes = box_ops.cycxhw_to_yxyx(
+ outputs['box_outputs']) * tf.expand_dims(
+ tf.concat([
+ labels['image_info'][:, 1:2, 0], labels['image_info'][:, 1:2,
+ 1],
+ labels['image_info'][:, 1:2, 0], labels['image_info'][:, 1:2,
+ 1]
+ ],
+ axis=1),
+ axis=1)
+ else:
+ detection_boxes = outputs['detection_boxes']
+
+ detection_scores = tf.math.reduce_max(
+ tf.nn.softmax(outputs['cls_outputs'])[:, :, 1:], axis=-1
+ ) if 'detection_scores' not in outputs else outputs['detection_scores']
+
+ if 'detection_classes' not in outputs:
+ detection_classes = tf.math.argmax(
+ outputs['cls_outputs'][:, :, 1:], axis=-1) + 1
+ else:
+ detection_classes = outputs['detection_classes']
+
+ if 'num_detections' not in outputs:
+ num_detections = tf.reduce_sum(
+ tf.cast(
+ tf.math.greater(
+ tf.math.reduce_max(outputs['cls_outputs'], axis=-1), 0),
+ tf.int32),
+ axis=-1)
+ else:
+ num_detections = outputs['num_detections']
+
predictions = {
- 'detection_boxes':
- box_ops.cycxhw_to_yxyx(outputs['box_outputs'])
- * tf.expand_dims(
- tf.concat([
- labels['image_info'][:, 1:2, 0],
- labels['image_info'][:, 1:2, 1],
- labels['image_info'][:, 1:2, 0],
- labels['image_info'][:, 1:2, 1]
- ],
- axis=1),
- axis=1),
- 'detection_scores':
- tf.math.reduce_max(
- tf.nn.softmax(outputs['cls_outputs'])[:, :, 1:], axis=-1),
- 'detection_classes':
- tf.math.argmax(outputs['cls_outputs'][:, :, 1:], axis=-1) + 1,
- # Fix this. It's not being used at the moment.
- 'num_detections': tf.reduce_sum(
- tf.cast(
- tf.math.greater(tf.math.reduce_max(
- outputs['cls_outputs'], axis=-1), 0), tf.int32), axis=-1),
+ 'detection_boxes': detection_boxes,
+ 'detection_scores': detection_scores,
+ 'detection_classes': detection_classes,
+ 'num_detections': num_detections,
'source_id': labels['id'],
'image_info': labels['image_info']
}
+
ground_truths = {
'source_id': labels['id'],
'height': labels['image_info'][:, 0:1, 0],
@@ -333,8 +425,8 @@ def aggregate_logs(self, state=None, step_outputs=None):
state = self.coco_metric
state.update_state(
- step_outputs['ground_truths'],
- step_outputs['predictions'])
+ step_outputs['ground_truths'], # pyrefly: ignore[unsupported-operation]
+ step_outputs['predictions']) # pyrefly: ignore[unsupported-operation]
return state
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
diff --git a/official/projects/detr/tasks/detection_test.py b/official/projects/detr/tasks/detection_test.py
index 65766c98ad5..5759108586f 100644
--- a/official/projects/detr/tasks/detection_test.py
+++ b/official/projects/detr/tasks/detection_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,14 @@
"""Tests for detection."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.projects.detr import optimization
from official.projects.detr.configs import detr as detr_cfg
from official.projects.detr.dataloaders import coco
from official.projects.detr.tasks import detection
+from official.vision.configs import backbones
_NUM_EXAMPLES = 10
@@ -36,7 +37,7 @@ def _gen_fn():
'image/id': np.random.randint(0, 100),
'image/filename': 'test',
'objects': {
- 'is_crowd': np.ones(shape=(num_boxes), dtype=np.bool),
+ 'is_crowd': np.ones(shape=(num_boxes), dtype=bool),
'bbox': np.ones(shape=(num_boxes, 4), dtype=np.float32),
'label': np.ones(shape=(num_boxes), dtype=np.int64),
'id': np.ones(shape=(num_boxes), dtype=np.int64),
@@ -58,9 +59,16 @@ def _as_dataset(self, *args, **kwargs):
class DetectionTest(tf.test.TestCase):
def test_train_step(self):
- config = detr_cfg.DetectionConfig(
- num_encoder_layers=1,
- num_decoder_layers=1,
+ config = detr_cfg.DetrTask(
+ model=detr_cfg.Detr(
+ input_size=[1333, 1333, 3],
+ num_encoder_layers=1,
+ num_decoder_layers=1,
+ num_classes=81,
+ backbone=backbones.Backbone(
+ type='resnet',
+ resnet=backbones.ResNet(model_id=10, bn_trainable=False))
+ ),
train_data=coco.COCODataConfig(
tfds_name='coco/2017',
tfds_split='validation',
@@ -68,7 +76,7 @@ def test_train_step(self):
global_batch_size=2,
))
with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
- task = detection.DectectionTask(config)
+ task = detection.DetectionTask(config)
model = task.build_model()
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
@@ -88,13 +96,20 @@ def test_train_step(self):
}
},
})
- optimizer = detection.DectectionTask.create_optimizer(opt_cfg)
+ optimizer = detection.DetectionTask.create_optimizer(opt_cfg)
task.train_step(next(iterator), model, optimizer)
def test_validation_step(self):
- config = detr_cfg.DetectionConfig(
- num_encoder_layers=1,
- num_decoder_layers=1,
+ config = detr_cfg.DetrTask(
+ model=detr_cfg.Detr(
+ input_size=[1333, 1333, 3],
+ num_encoder_layers=1,
+ num_decoder_layers=1,
+ num_classes=81,
+ backbone=backbones.Backbone(
+ type='resnet',
+ resnet=backbones.ResNet(model_id=10, bn_trainable=False))
+ ),
validation_data=coco.COCODataConfig(
tfds_name='coco/2017',
tfds_split='validation',
@@ -103,7 +118,79 @@ def test_validation_step(self):
))
with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
- task = detection.DectectionTask(config)
+ task = detection.DetectionTask(config)
+ model = task.build_model()
+ metrics = task.build_metrics(training=False)
+ dataset = task.build_inputs(config.validation_data)
+ iterator = iter(dataset)
+ logs = task.validation_step(next(iterator), model, metrics)
+ state = task.aggregate_logs(step_outputs=logs)
+ task.reduce_aggregated_logs(state)
+
+
+class DetectionTFDSTest(tf.test.TestCase):
+
+ def test_train_step(self):
+ config = detr_cfg.DetrTask(
+ model=detr_cfg.Detr(
+ input_size=[1333, 1333, 3],
+ num_encoder_layers=1,
+ num_decoder_layers=1,
+ backbone=backbones.Backbone(
+ type='resnet',
+ resnet=backbones.ResNet(model_id=10, bn_trainable=False))
+ ),
+ losses=detr_cfg.Losses(class_offset=1),
+ train_data=detr_cfg.DataConfig(
+ tfds_name='coco/2017',
+ tfds_split='validation',
+ is_training=True,
+ global_batch_size=2,
+ ))
+ with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
+ task = detection.DetectionTask(config)
+ model = task.build_model()
+ dataset = task.build_inputs(config.train_data)
+ iterator = iter(dataset)
+ opt_cfg = optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'detr_adamw',
+ 'detr_adamw': {
+ 'weight_decay_rate': 1e-4,
+ 'global_clipnorm': 0.1,
+ }
+ },
+ 'learning_rate': {
+ 'type': 'stepwise',
+ 'stepwise': {
+ 'boundaries': [120000],
+ 'values': [0.0001, 1.0e-05]
+ }
+ },
+ })
+ optimizer = detection.DetectionTask.create_optimizer(opt_cfg)
+ task.train_step(next(iterator), model, optimizer)
+
+ def test_validation_step(self):
+ config = detr_cfg.DetrTask(
+ model=detr_cfg.Detr(
+ input_size=[1333, 1333, 3],
+ num_encoder_layers=1,
+ num_decoder_layers=1,
+ backbone=backbones.Backbone(
+ type='resnet',
+ resnet=backbones.ResNet(model_id=10, bn_trainable=False))
+ ),
+ losses=detr_cfg.Losses(class_offset=1),
+ validation_data=detr_cfg.DataConfig(
+ tfds_name='coco/2017',
+ tfds_split='validation',
+ is_training=False,
+ global_batch_size=2,
+ ))
+
+ with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
+ task = detection.DetectionTask(config)
model = task.build_model()
metrics = task.build_metrics(training=False)
dataset = task.build_inputs(config.validation_data)
diff --git a/official/projects/detr/train.py b/official/projects/detr/train.py
index a34da6843b4..99ee98c06ca 100644
--- a/official/projects/detr/train.py
+++ b/official/projects/detr/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/nlp/__init__.py b/official/projects/edgetpu/nlp/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/nlp/__init__.py
+++ b/official/projects/edgetpu/nlp/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/nlp/configs/__init__.py b/official/projects/edgetpu/nlp/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/nlp/configs/__init__.py
+++ b/official/projects/edgetpu/nlp/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/nlp/configs/params.py b/official/projects/edgetpu/nlp/configs/params.py
index 39a83ba26eb..cb3307e0cfb 100644
--- a/official/projects/edgetpu/nlp/configs/params.py
+++ b/official/projects/edgetpu/nlp/configs/params.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -52,20 +52,31 @@ class OrbitParams(base_config.Config):
@dataclasses.dataclass
class OptimizerParams(optimization.OptimizationConfig):
"""Optimizer parameters for MobileBERT-EdgeTPU."""
- optimizer: optimization.OptimizerConfig = optimization.OptimizerConfig(
- type='adamw',
- adamw=optimization.AdamWeightDecayConfig(
- weight_decay_rate=0.01,
- exclude_from_weight_decay=['LayerNorm', 'layer_norm', 'bias']))
- learning_rate: optimization.LrConfig = optimization.LrConfig(
- type='polynomial',
- polynomial=optimization.PolynomialLrConfig(
- initial_learning_rate=1e-4,
- decay_steps=1000000,
- end_learning_rate=0.0))
- warmup: optimization.WarmupConfig = optimization.WarmupConfig(
- type='polynomial',
- polynomial=optimization.PolynomialWarmupConfig(warmup_steps=10000))
+ optimizer: optimization.OptimizerConfig = dataclasses.field(
+ default_factory=lambda: optimization.OptimizerConfig( # pylint: disable=g-long-lambda
+ type='adamw',
+ adamw=optimization.AdamWeightDecayConfig(
+ weight_decay_rate=0.01,
+ exclude_from_weight_decay=['LayerNorm', 'layer_norm', 'bias'],
+ ),
+ )
+ )
+ learning_rate: optimization.LrConfig = dataclasses.field(
+ default_factory=lambda: optimization.LrConfig( # pylint: disable=g-long-lambda
+ type='polynomial',
+ polynomial=optimization.PolynomialLrConfig(
+ initial_learning_rate=1e-4,
+ decay_steps=1000000,
+ end_learning_rate=0.0,
+ ),
+ )
+ )
+ warmup: optimization.WarmupConfig = dataclasses.field(
+ default_factory=lambda: optimization.WarmupConfig( # pylint: disable=g-long-lambda
+ type='polynomial',
+ polynomial=optimization.PolynomialWarmupConfig(warmup_steps=10000),
+ )
+ )
@dataclasses.dataclass
@@ -144,16 +155,26 @@ class EdgeTPUBERTCustomParams(base_config.Config):
distill_ground_truth_ratio: A float number representing the ratio between
distillation output and ground truth.
"""
- train_datasest: DatasetParams = DatasetParams()
- eval_dataset: DatasetParams = DatasetParams()
- teacher_model: Optional[PretrainerModelParams] = PretrainerModelParams()
- student_model: PretrainerModelParams = PretrainerModelParams()
+ train_datasest: DatasetParams = dataclasses.field(
+ default_factory=DatasetParams
+ )
+ eval_dataset: DatasetParams = dataclasses.field(default_factory=DatasetParams)
+ teacher_model: Optional[PretrainerModelParams] = dataclasses.field(
+ default_factory=PretrainerModelParams
+ )
+ student_model: PretrainerModelParams = dataclasses.field(
+ default_factory=PretrainerModelParams
+ )
teacher_model_init_checkpoint: str = ''
student_model_init_checkpoint: str = ''
- layer_wise_distillation: LayerWiseDistillationParams = (
- LayerWiseDistillationParams())
- end_to_end_distillation: EndToEndDistillationParams = (
- EndToEndDistillationParams())
- optimizer: OptimizerParams = OptimizerParams()
- runtime: RuntimeParams = RuntimeParams()
- orbit_config: OrbitParams = OrbitParams()
+ layer_wise_distillation: LayerWiseDistillationParams = dataclasses.field(
+ default_factory=LayerWiseDistillationParams
+ )
+ end_to_end_distillation: EndToEndDistillationParams = dataclasses.field(
+ default_factory=EndToEndDistillationParams
+ )
+ optimizer: OptimizerParams = dataclasses.field(
+ default_factory=OptimizerParams
+ )
+ runtime: RuntimeParams = dataclasses.field(default_factory=RuntimeParams)
+ orbit_config: OrbitParams = dataclasses.field(default_factory=OrbitParams)
diff --git a/official/projects/edgetpu/nlp/experiments/mobilebert_edgetpu_xxs.yaml b/official/projects/edgetpu/nlp/experiments/mobilebert_edgetpu_xxs.yaml
index 18fbaf7a2f4..86c26569339 100644
--- a/official/projects/edgetpu/nlp/experiments/mobilebert_edgetpu_xxs.yaml
+++ b/official/projects/edgetpu/nlp/experiments/mobilebert_edgetpu_xxs.yaml
@@ -63,6 +63,7 @@ student_model:
type: mobilebert
mlm_activation: relu
mlm_initializer_range: 0.02
+ mlm_output_weights_use_proj: true
teacher_model:
cls_heads: []
encoder:
diff --git a/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer.py b/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer.py
index 2adeb246bf0..cc7fee7a5be 100644
--- a/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer.py
+++ b/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import optimization
from official.nlp import modeling
@@ -65,9 +65,9 @@ def _get_distribution_losses(teacher, student):
def _get_attention_loss(teacher_score, student_score):
"""Function to calculate attention loss for transformer layers."""
# Note that the definition of KLDivergence here is a little different from
- # the original one (tf.keras.losses.KLDivergence). We adopt this approach
+ # the original one (tf_keras.losses.KLDivergence). We adopt this approach
# to stay consistent with the TF1 implementation.
- teacher_weight = tf.keras.activations.softmax(teacher_score, axis=-1)
+ teacher_weight = tf_keras.activations.softmax(teacher_score, axis=-1)
student_log_weight = tf.nn.log_softmax(student_score, axis=-1)
kl_divergence = -(teacher_weight * student_log_weight)
kl_divergence = tf.math.reduce_sum(kl_divergence, axis=-1, keepdims=True)
@@ -91,7 +91,7 @@ def _build_sub_encoder(encoder, stage_number):
layer_output, attention_score = encoder.transformer_layers[layer_idx](
layer_output, attention_mask, return_attention_scores=True)
- return tf.keras.Model(
+ return tf_keras.Model(
inputs=[input_ids, input_mask, type_ids],
outputs=[layer_output, attention_score])
@@ -145,7 +145,7 @@ def __init__(self,
self.current_optimizer = self.layer_wise_optimizer
# A non-trainable layer for feature normalization for transfer loss.
- self._layer_norm = tf.keras.layers.LayerNormalization(
+ self._layer_norm = tf_keras.layers.LayerNormalization(
axis=-1,
beta_initializer='zeros',
gamma_initializer='ones',
@@ -204,7 +204,7 @@ def build_model(self):
stage_number=int(self.stage * self.ratio))
teacher_output_feature, teacher_attention_score = teacher_sub_encoder(
inputs)
- return tf.keras.Model(
+ return tf_keras.Model(
inputs=inputs,
outputs=dict(
student_output_feature=student_output_feature,
@@ -216,7 +216,7 @@ def build_model(self):
inputs = self.student_model.inputs
student_pretrainer_outputs = self.student_model(inputs)
teacher_pretrainer_outputs = self.teacher_model(inputs)
- model = tf.keras.Model(
+ model = tf_keras.Model(
inputs=inputs,
outputs=dict(
student_pretrainer_outputs=student_pretrainer_outputs,
@@ -253,19 +253,19 @@ def build_optimizer(self, config):
def build_metrics(self):
"""Creates metrics functions for the training."""
self.train_metrics = {
- 'feature_transfer_mse': tf.keras.metrics.Mean(),
- 'beta_transfer_loss': tf.keras.metrics.Mean(),
- 'gamma_transfer_loss': tf.keras.metrics.Mean(),
- 'attention_transfer_loss': tf.keras.metrics.Mean(),
- 'masked_lm_accuracy': tf.keras.metrics.SparseCategoricalAccuracy(),
- 'lm_example_loss': tf.keras.metrics.Mean(),
- 'total_loss': tf.keras.metrics.Mean(),
- 'next_sentence_accuracy': tf.keras.metrics.SparseCategoricalAccuracy(),
- 'next_sentence_loss': tf.keras.metrics.Mean(),
+ 'feature_transfer_mse': tf_keras.metrics.Mean(),
+ 'beta_transfer_loss': tf_keras.metrics.Mean(),
+ 'gamma_transfer_loss': tf_keras.metrics.Mean(),
+ 'attention_transfer_loss': tf_keras.metrics.Mean(),
+ 'masked_lm_accuracy': tf_keras.metrics.SparseCategoricalAccuracy(),
+ 'lm_example_loss': tf_keras.metrics.Mean(),
+ 'total_loss': tf_keras.metrics.Mean(),
+ 'next_sentence_accuracy': tf_keras.metrics.SparseCategoricalAccuracy(),
+ 'next_sentence_loss': tf_keras.metrics.Mean(),
}
self.eval_metrics = {
- 'masked_lm_accuracy': tf.keras.metrics.SparseCategoricalAccuracy(),
- 'next_sentence_accuracy': tf.keras.metrics.SparseCategoricalAccuracy(),
+ 'masked_lm_accuracy': tf_keras.metrics.SparseCategoricalAccuracy(),
+ 'next_sentence_accuracy': tf_keras.metrics.SparseCategoricalAccuracy(),
}
def build_exported_ckpt_manager(self):
@@ -298,7 +298,7 @@ def calculate_loss_metrics(self, labels, outputs):
if self.mode == DistillationMode.LAYER_WISE:
teacher_feature = outputs['teacher_output_feature']
student_feature = outputs['student_output_feature']
- feature_transfer_loss = tf.keras.losses.mean_squared_error(
+ feature_transfer_loss = tf_keras.losses.mean_squared_error(
self._layer_norm(teacher_feature), self._layer_norm(student_feature))
# feature_transfer_loss = tf.reduce_mean(feature_transfer_loss)
feature_transfer_loss *= self.layer_wise_distill_config.hidden_distill_factor
@@ -349,7 +349,7 @@ def calculate_loss_metrics(self, labels, outputs):
sentence_outputs = tf.cast(
student_pretrainer_output['next_sentence'], dtype=tf.float32)
sentence_loss = tf.reduce_mean(
- tf.keras.losses.sparse_categorical_crossentropy(
+ tf_keras.losses.sparse_categorical_crossentropy(
sentence_labels, sentence_outputs, from_logits=True))
total_loss += sentence_loss
else:
@@ -357,13 +357,13 @@ def calculate_loss_metrics(self, labels, outputs):
if self.mode == DistillationMode.LAYER_WISE:
self.train_metrics['feature_transfer_mse'].update_state(
- feature_transfer_loss)
- self.train_metrics['beta_transfer_loss'].update_state(beta_loss)
- self.train_metrics['gamma_transfer_loss'].update_state(gamma_loss)
- self.train_metrics['attention_transfer_loss'].update_state(attention_loss)
+ feature_transfer_loss) # pyrefly: ignore[unbound-name]
+ self.train_metrics['beta_transfer_loss'].update_state(beta_loss) # pyrefly: ignore[unbound-name]
+ self.train_metrics['gamma_transfer_loss'].update_state(gamma_loss) # pyrefly: ignore[unbound-name]
+ self.train_metrics['attention_transfer_loss'].update_state(attention_loss) # pyrefly: ignore[unbound-name]
elif self.mode == DistillationMode.END2END:
- self.train_metrics['lm_example_loss'].update_state(mlm_loss)
- self.train_metrics['next_sentence_loss'].update_state(sentence_loss)
+ self.train_metrics['lm_example_loss'].update_state(mlm_loss) # pyrefly: ignore[unbound-name]
+ self.train_metrics['next_sentence_loss'].update_state(sentence_loss) # pyrefly: ignore[unbound-name]
self.train_metrics['total_loss'].update_state(total_loss)
return total_loss
@@ -467,7 +467,7 @@ def train_loop_end(self):
# e2e distillation training stage.
if self.exported_ckpt_manager is None:
self.build_exported_ckpt_manager()
- self.exported_ckpt_manager.save(
+ self.exported_ckpt_manager.save( # pyrefly: ignore[missing-attribute]
checkpoint_number=self.current_step.numpy(),
check_interval=True)
@@ -512,7 +512,7 @@ def step_fn(inputs):
self.strategy.run(step_fn, args=(next(iterator),))
- def eval_end(self):
+ def eval_end(self): # pyrefly: ignore[bad-override]
return {'masked_lm_accuracy':
self.eval_metrics['masked_lm_accuracy'].result(),
'next_sentence_accuracy':
diff --git a/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer_test.py b/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer_test.py
index b411c4946f3..5fff6be2019 100644
--- a/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer_test.py
+++ b/official/projects/edgetpu/nlp/mobilebert_edgetpu_trainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for mobilebert_edgetpu_trainer.py."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.nlp import mobilebert_edgetpu_trainer
from official.projects.edgetpu.nlp.configs import params
diff --git a/official/projects/edgetpu/nlp/modeling/__init__.py b/official/projects/edgetpu/nlp/modeling/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/nlp/modeling/__init__.py
+++ b/official/projects/edgetpu/nlp/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/nlp/modeling/edgetpu_layers.py b/official/projects/edgetpu/nlp/modeling/edgetpu_layers.py
index 95c08ec73ac..cc1683f6b8c 100644
--- a/official/projects/edgetpu/nlp/modeling/edgetpu_layers.py
+++ b/official/projects/edgetpu/nlp/modeling/edgetpu_layers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -24,14 +24,14 @@
import string
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
_CHR_IDX = string.ascii_lowercase
-# This function is directly copied from the tf.keras.layers.MultiHeadAttention
+# This function is directly copied from the tf_keras.layers.MultiHeadAttention
# implementation.
def _build_attention_equation(rank, attn_axes):
"""Builds einsum equations for the attention computation.
@@ -81,8 +81,8 @@ def _build_attention_equation(rank, attn_axes):
return dot_product_equation, combine_equation, attn_scores_rank
-@tf.keras.utils.register_keras_serializable(package='Text')
-class EdgeTPUSoftmax(tf.keras.layers.Softmax):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class EdgeTPUSoftmax(tf_keras.layers.Softmax):
"""EdgeTPU/Quantization friendly implementation for the SoftMax.
When export quant model, use -120 mask value.
@@ -111,26 +111,26 @@ def call(self, inputs, mask=None):
return tf.exp(inputs - tf.reduce_logsumexp(
inputs, axis=self.axis, keepdims=True))
else:
- return tf.keras.backend.softmax(inputs, axis=self.axis[0])
- return tf.keras.backend.softmax(inputs, axis=self.axis)
+ return tf_keras.backend.softmax(inputs, axis=self.axis[0])
+ return tf_keras.backend.softmax(inputs, axis=self.axis)
-@tf.keras.utils.register_keras_serializable(package='Text')
-class EdgeTPUMultiHeadAttention(tf.keras.layers.MultiHeadAttention):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class EdgeTPUMultiHeadAttention(tf_keras.layers.MultiHeadAttention):
"""Quantization friendly implementation for the MultiHeadAttention."""
def _build_attention(self, rank):
"""Builds multi-head dot-product attention computations.
This function builds attributes necessary for `_compute_attention` to
- costomize attention computation to replace the default dot-product
+ customize attention computation to replace the default dot-product
attention.
Args:
rank: the rank of query, key, value tensors.
"""
if self._attention_axes is None:
- self._attention_axes = tuple(range(1, rank - 2))
+ self._attention_axes = tuple(range(1, rank - 2)) # pyrefly: ignore[bad-assignment]
else:
self._attention_axes = tuple(self._attention_axes)
self._dot_product_equation, self._combine_equation, attn_scores_rank = (
@@ -139,7 +139,7 @@ def _build_attention(self, rank):
norm_axes = tuple(
range(attn_scores_rank - len(self._attention_axes), attn_scores_rank))
self._softmax = EdgeTPUSoftmax(axis=norm_axes)
- self._dropout_layer = tf.keras.layers.Dropout(rate=self._dropout)
+ self._dropout_layer = tf_keras.layers.Dropout(rate=self._dropout)
class EdgetpuMobileBertTransformer(layers.MobileBertTransformer):
diff --git a/official/projects/edgetpu/nlp/modeling/edgetpu_layers_test.py b/official/projects/edgetpu/nlp/modeling/edgetpu_layers_test.py
index 1ed5570d2d1..6ae96c8afe4 100644
--- a/official/projects/edgetpu/nlp/modeling/edgetpu_layers_test.py
+++ b/official/projects/edgetpu/nlp/modeling/edgetpu_layers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for custom layers used by MobileBERT-EdgeTPU."""
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.nlp.modeling import edgetpu_layers
diff --git a/official/projects/edgetpu/nlp/modeling/encoder.py b/official/projects/edgetpu/nlp/modeling/encoder.py
index a3790da7f01..fbbc61cd2c2 100644
--- a/official/projects/edgetpu/nlp/modeling/encoder.py
+++ b/official/projects/edgetpu/nlp/modeling/encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,15 +14,15 @@
"""MobileBERT text encoder network."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp import modeling
from official.nlp.modeling import layers
from official.projects.edgetpu.nlp.modeling import edgetpu_layers
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MobileBERTEncoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MobileBERTEncoder(tf_keras.Model):
"""A Keras functional API implementation for MobileBERT encoder."""
def __init__(self,
@@ -91,7 +91,7 @@ def __init__(self,
**kwargs: Other keyworded and arguments.
"""
self._self_setattr_tracking = False
- initializer = tf.keras.initializers.TruncatedNormal(
+ initializer = tf_keras.initializers.TruncatedNormal(
stddev=initializer_range)
# layer instantiation
@@ -132,11 +132,11 @@ def __init__(self,
self._transformer_layers.append(transformer)
# input tensor
- input_ids = tf.keras.layers.Input(
+ input_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(None,), dtype=input_mask_dtype, name='input_mask')
self.inputs = [input_ids, input_mask, type_ids]
@@ -161,7 +161,7 @@ def __init__(self,
first_token = tf.squeeze(prev_output[:, 0:1, :], axis=1)
if classifier_activation:
- self._pooler_layer = tf.keras.layers.experimental.EinsumDense(
+ self._pooler_layer = tf_keras.layers.EinsumDense(
'ab,bc->ac',
output_shape=hidden_size,
activation=tf.tanh,
diff --git a/official/projects/edgetpu/nlp/modeling/model_builder.py b/official/projects/edgetpu/nlp/modeling/model_builder.py
index 09ab3a3dee2..b82df986379 100644
--- a/official/projects/edgetpu/nlp/modeling/model_builder.py
+++ b/official/projects/edgetpu/nlp/modeling/model_builder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Build MobileBERT-EdgeTPU model."""
from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.nlp import modeling
@@ -25,10 +25,10 @@
def build_bert_pretrainer(pretrainer_cfg: params.PretrainerModelParams,
- encoder: Optional[tf.keras.Model] = None,
- masked_lm: Optional[tf.keras.Model] = None,
+ encoder: Optional[tf_keras.Model] = None,
+ masked_lm: Optional[tf_keras.Model] = None,
quantization_friendly: Optional[bool] = False,
- name: Optional[str] = None) -> tf.keras.Model:
+ name: Optional[str] = None) -> tf_keras.Model:
"""Builds pretrainer.
Args:
@@ -83,8 +83,9 @@ def _get_embedding_table(encoder):
masked_lm = masked_lm or modeling.layers.MobileBertMaskedLM(
embedding_table=_get_embedding_table(encoder),
activation=tf_utils.get_activation(pretrainer_cfg.mlm_activation),
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=pretrainer_cfg.mlm_initializer_range),
+ output_weights_use_proj=pretrainer_cfg.mlm_output_weights_use_proj,
name='cls/predictions')
pretrainer = edgetpu_pretrainer.MobileBERTEdgeTPUPretrainer(
diff --git a/official/projects/edgetpu/nlp/modeling/model_builder_test.py b/official/projects/edgetpu/nlp/modeling/model_builder_test.py
index 159dd2d7b44..1e0d75f1dfb 100644
--- a/official/projects/edgetpu/nlp/modeling/model_builder_test.py
+++ b/official/projects/edgetpu/nlp/modeling/model_builder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for mobilebert_edgetpu.model_builder.py."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp import modeling
from official.nlp.configs import encoders
@@ -62,7 +62,7 @@ def test_initialization_with_mlm(self):
word_embed_size=128,
type_vocab_size=2,
output_embed_size=encoders.MobileBertEncoderConfig().hidden_size)
- dummy_input = tf.keras.layers.Input(
+ dummy_input = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32)
_ = embedding(dummy_input)
embedding_table = embedding.word_embedding.embeddings
diff --git a/official/projects/edgetpu/nlp/modeling/pretrainer.py b/official/projects/edgetpu/nlp/modeling/pretrainer.py
index 8607f3e817c..f2d6a88513d 100644
--- a/official/projects/edgetpu/nlp/modeling/pretrainer.py
+++ b/official/projects/edgetpu/nlp/modeling/pretrainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,13 +17,13 @@
import copy
from typing import List, Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
-@tf.keras.utils.register_keras_serializable(package='Text')
-class MobileBERTEdgeTPUPretrainer(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MobileBERTEdgeTPUPretrainer(tf_keras.Model):
"""BERT pretraining model V2.
Adds the masked language model head and optional classification heads upon the
@@ -52,11 +52,11 @@ class MobileBERTEdgeTPUPretrainer(tf.keras.Model):
def __init__(
self,
- encoder_network: tf.keras.Model,
+ encoder_network: tf_keras.Model,
mlm_activation=None,
mlm_initializer='glorot_uniform',
- classification_heads: Optional[List[tf.keras.layers.Layer]] = None,
- customized_masked_lm: Optional[tf.keras.layers.Layer] = None,
+ classification_heads: Optional[List[tf_keras.layers.Layer]] = None,
+ customized_masked_lm: Optional[tf_keras.layers.Layer] = None,
name: str = 'bert',
**kwargs):
@@ -76,7 +76,7 @@ def __init__(
raise ValueError('encoder_network\'s output should be either a list '
'or a dict, but got %s' % encoder_network_outputs)
- masked_lm_positions = tf.keras.layers.Input(
+ masked_lm_positions = tf_keras.layers.Input(
shape=(None,), name='masked_lm_positions', dtype=tf.int32)
inputs.append(masked_lm_positions)
masked_lm_layer = customized_masked_lm or layers.MaskedLM(
diff --git a/official/projects/edgetpu/nlp/modeling/pretrainer_test.py b/official/projects/edgetpu/nlp/modeling/pretrainer_test.py
index e896d0da1a9..23228e3cbce 100644
--- a/official/projects/edgetpu/nlp/modeling/pretrainer_test.py
+++ b/official/projects/edgetpu/nlp/modeling/pretrainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import itertools
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
from official.nlp.modeling import networks
@@ -59,10 +59,10 @@ def test_mobilebert_edgetpu_pretrainer(
num_token_predictions = 20
# Create a set of 2-dimensional inputs (the first dimension is implicit).
inputs = dict(
- input_word_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32))
- inputs['masked_lm_positions'] = tf.keras.Input(
+ input_word_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32))
+ inputs['masked_lm_positions'] = tf_keras.Input(
shape=(num_token_predictions,), dtype=tf.int32)
# Invoke the trainer model on the inputs. This causes the layer to be built.
@@ -109,10 +109,10 @@ def test_multiple_cls_outputs(self):
num_token_predictions = 20
# Create a set of 2-dimensional inputs (the first dimension is implicit).
inputs = dict(
- input_word_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(sequence_length,), dtype=tf.int32),
- masked_lm_positions=tf.keras.Input(
+ input_word_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ masked_lm_positions=tf_keras.Input(
shape=(num_token_predictions,), dtype=tf.int32))
# Invoke the trainer model on the inputs. This causes the layer to be built.
diff --git a/official/projects/edgetpu/nlp/run_mobilebert_edgetpu_train.py b/official/projects/edgetpu/nlp/run_mobilebert_edgetpu_train.py
index 812a0d051e6..86a7386440c 100644
--- a/official/projects/edgetpu/nlp/run_mobilebert_edgetpu_train.py
+++ b/official/projects/edgetpu/nlp/run_mobilebert_edgetpu_train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl import flags
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.common import flags as tfm_flags
diff --git a/official/projects/edgetpu/nlp/serving/__init__.py b/official/projects/edgetpu/nlp/serving/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/nlp/serving/__init__.py
+++ b/official/projects/edgetpu/nlp/serving/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/nlp/serving/export_tflite_squad.py b/official/projects/edgetpu/nlp/serving/export_tflite_squad.py
index b66c54a84d4..e03b55fbcd1 100644
--- a/official/projects/edgetpu/nlp/serving/export_tflite_squad.py
+++ b/official/projects/edgetpu/nlp/serving/export_tflite_squad.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -31,7 +31,7 @@
from absl import flags
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import flags as tfm_flags
from official.nlp.data import data_loader_factory
@@ -58,9 +58,9 @@
'with random weights if path is None.')
-def build_model_for_serving(model: tf.keras.Model,
+def build_model_for_serving(model: tf_keras.Model,
sequence_length: int = 384,
- batch_size: int = 1) -> tf.keras.Model:
+ batch_size: int = 1) -> tf_keras.Model:
"""Builds MLPerf evaluation compatible models.
To run the model on device, the model input/output datatype and node names
@@ -73,27 +73,27 @@ def build_model_for_serving(model: tf.keras.Model,
Returns:
Keras model with new input/output nodes.
"""
- word_ids = tf.keras.Input(shape=(sequence_length,),
+ word_ids = tf_keras.Input(shape=(sequence_length,),
batch_size=batch_size,
dtype=tf.int32,
name='input_word_ids')
- mask = tf.keras.Input(shape=(sequence_length,),
+ mask = tf_keras.Input(shape=(sequence_length,),
batch_size=batch_size,
dtype=tf.int32, name='input_mask')
- type_ids = tf.keras.Input(shape=(sequence_length,),
+ type_ids = tf_keras.Input(shape=(sequence_length,),
batch_size=batch_size,
dtype=tf.int32, name='input_type_ids')
model_output = model([word_ids, type_ids, mask])
# Use identity layers wrapped in lambdas to explicitly name the output
# tensors.
- start_logits = tf.keras.layers.Lambda(
+ start_logits = tf_keras.layers.Lambda(
tf.identity, name='start_positions')(
model_output[0])
- end_logits = tf.keras.layers.Lambda(
+ end_logits = tf_keras.layers.Lambda(
tf.identity, name='end_positions')(
model_output[1])
- model = tf.keras.Model(
+ model = tf_keras.Model(
inputs=[word_ids, type_ids, mask],
outputs=[start_logits, end_logits])
@@ -127,7 +127,7 @@ def main(argv: Sequence[str]) -> None:
encoder_network = pretrainer_model.encoder_network
model = models.BertSpanLabeler(
network=encoder_network,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.01))
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01))
# Load model weights.
if FLAGS.model_checkpoint is not None:
diff --git a/official/projects/edgetpu/nlp/serving/export_tflite_squad_test.py b/official/projects/edgetpu/nlp/serving/export_tflite_squad_test.py
index 10c1b0d51a8..3dbfe662717 100644
--- a/official/projects/edgetpu/nlp/serving/export_tflite_squad_test.py
+++ b/official/projects/edgetpu/nlp/serving/export_tflite_squad_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for export_tflite_squad."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import models
from official.projects.edgetpu.nlp.configs import params
@@ -32,7 +32,7 @@ def setUp(self):
encoder_network = pretrainer_model.encoder_network
self.span_labeler = models.BertSpanLabeler(
network=encoder_network,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.01))
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01))
def test_model_input_output(self):
test_model = export_tflite_squad.build_model_for_serving(self.span_labeler)
diff --git a/official/projects/edgetpu/nlp/utils/__init__.py b/official/projects/edgetpu/nlp/utils/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/nlp/utils/__init__.py
+++ b/official/projects/edgetpu/nlp/utils/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/nlp/utils/utils.py b/official/projects/edgetpu/nlp/utils/utils.py
index ea0594a160e..8054448728d 100644
--- a/official/projects/edgetpu/nlp/utils/utils.py
+++ b/official/projects/edgetpu/nlp/utils/utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import pprint
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.edgetpu.nlp.configs import params
@@ -86,7 +86,7 @@ def get_model_dir(experiment_params, flags_obj):
return flags_obj.model_dir
-def load_checkpoint(model: tf.keras.Model, ckpt_path: str):
+def load_checkpoint(model: tf_keras.Model, ckpt_path: str):
"""Initializes model with the checkpoint."""
ckpt_dir_or_file = ckpt_path
diff --git a/official/projects/edgetpu/nlp/utils/utils_test.py b/official/projects/edgetpu/nlp/utils/utils_test.py
index 82131baab74..e7bd9381713 100644
--- a/official/projects/edgetpu/nlp/utils/utils_test.py
+++ b/official/projects/edgetpu/nlp/utils/utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for utils.py."""
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.nlp.configs import params
from official.projects.edgetpu.nlp.modeling import model_builder
diff --git a/official/projects/edgetpu/vision/__init__.py b/official/projects/edgetpu/vision/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/__init__.py
+++ b/official/projects/edgetpu/vision/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/configs/__init__.py b/official/projects/edgetpu/vision/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/configs/__init__.py
+++ b/official/projects/edgetpu/vision/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/configs/mobilenet_edgetpu_config.py b/official/projects/edgetpu/vision/configs/mobilenet_edgetpu_config.py
index 5970e533c5e..9dc4821dd98 100644
--- a/official/projects/edgetpu/vision/configs/mobilenet_edgetpu_config.py
+++ b/official/projects/edgetpu/vision/configs/mobilenet_edgetpu_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,8 +19,6 @@
import os
from typing import Any, Mapping, Optional
-# Import libraries
-
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.modeling import optimization
@@ -60,7 +58,9 @@ class MobilenetEdgeTPUTaskConfig(base_config.ImageClassificationTask):
saved_model_path: Instead of initializing a model from the model config,
the model can be loaded from a file path.
"""
- model: MobilenetEdgeTPUModelConfig = MobilenetEdgeTPUModelConfig()
+ model: MobilenetEdgeTPUModelConfig = dataclasses.field(
+ default_factory=MobilenetEdgeTPUModelConfig
+ )
saved_model_path: Optional[str] = None
@@ -83,19 +83,19 @@ def mobilenet_edgetpu_base_experiment_config(
eval_batch_size = 4096
steps_per_epoch = IMAGENET_TRAIN_EXAMPLES // train_batch_size
mobilenet_edgetpu_config = MobilenetEdgeTPUModelConfig(
- num_classes=1001, input_size=[224, 224, 3])
- mobilenet_edgetpu_config.model_params.model_name = model_name
+ num_classes=1001, input_size=[224, 224, 3]) # pyrefly: ignore[unexpected-keyword]
+ mobilenet_edgetpu_config.model_params.model_name = model_name # pyrefly: ignore[missing-attribute]
config = cfg.ExperimentConfig(
task=MobilenetEdgeTPUTaskConfig(
model=mobilenet_edgetpu_config,
- losses=base_config.Losses(label_smoothing=0.1),
- train_data=base_config.DataConfig(
+ losses=base_config.Losses(label_smoothing=0.1), # pyrefly: ignore[unexpected-keyword]
+ train_data=base_config.DataConfig( # pyrefly: ignore[unexpected-keyword]
input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'train*'),
is_training=True,
global_batch_size=train_batch_size,
dtype='bfloat16',
aug_type=common.Augmentation(type='autoaug')),
- validation_data=base_config.DataConfig(
+ validation_data=base_config.DataConfig( # pyrefly: ignore[unexpected-keyword]
input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'valid*'),
is_training=False,
dtype='bfloat16',
diff --git a/official/projects/edgetpu/vision/configs/semantic_segmentation_config.py b/official/projects/edgetpu/vision/configs/semantic_segmentation_config.py
index 10012436d96..d5f3808ff6d 100644
--- a/official/projects/edgetpu/vision/configs/semantic_segmentation_config.py
+++ b/official/projects/edgetpu/vision/configs/semantic_segmentation_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -47,7 +47,9 @@ class Backbone(backbones.Backbone):
spinenet_seg: spinenet-seg backbone config.
"""
type: Optional[str] = None
- mobilenet_edgetpu: MobileNetEdgeTPU = MobileNetEdgeTPU()
+ mobilenet_edgetpu: MobileNetEdgeTPU = dataclasses.field(
+ default_factory=MobileNetEdgeTPU
+ )
@dataclasses.dataclass
@@ -55,9 +57,16 @@ class CustomSemanticSegmentationTaskConfig(base_cfg.SemanticSegmentationTask):
"""Same config for custom taks."""
model: Optional[base_cfg.SemanticSegmentationModel] = None
- train_data: base_cfg.DataConfig = base_cfg.DataConfig(is_training=True)
- validation_data: base_cfg.DataConfig = base_cfg.DataConfig(is_training=False)
- evaluation: base_cfg.Evaluation = base_cfg.Evaluation()
+ train_data: base_cfg.DataConfig = dataclasses.field(
+ default_factory=lambda: base_cfg.DataConfig(is_training=True)
+ )
+ validation_data: base_cfg.DataConfig = dataclasses.field(
+ default_factory=lambda: base_cfg.DataConfig(is_training=False)
+ )
+ evaluation: base_cfg.Evaluation = dataclasses.field(
+ default_factory=base_cfg.Evaluation
+ )
+ allow_image_summary: bool = False
# ADE 20K Dataset
@@ -68,18 +77,26 @@ class CustomSemanticSegmentationTaskConfig(base_cfg.SemanticSegmentationTask):
PRETRAINED_CKPT_PATH_BASE = 'gs://**/placeholder_for_edgetpu_models'
BACKBONE_PRETRAINED_CHECKPOINT = {
- 'mobilenet_edgetpu_v2_l':
- PRETRAINED_CKPT_PATH_BASE +
- '/pretrained_checkpoints/mobilenet_edgetpu_v2_l/ckpt-171600',
- 'mobilenet_edgetpu_v2_m':
- PRETRAINED_CKPT_PATH_BASE +
- '/pretrained_checkpoints/mobilenet_edgetpu_v2_m/ckpt-171600',
- 'mobilenet_edgetpu_v2_s':
- PRETRAINED_CKPT_PATH_BASE +
- '/pretrained_checkpoints/mobilenet_edgetpu_v2_s/ckpt-171600',
- 'mobilenet_edgetpu_v2_xs':
- PRETRAINED_CKPT_PATH_BASE +
- '/pretrained_checkpoints/mobilenet_edgetpu_v2_xs/ckpt-171600',
+ 'mobilenet_edgetpu_v2_l': (
+ PRETRAINED_CKPT_PATH_BASE
+ + '/pretrained_checkpoints/mobilenet_edgetpu_v2_l/ckpt-171600'
+ ),
+ 'mobilenet_edgetpu_v2_m': (
+ PRETRAINED_CKPT_PATH_BASE
+ + '/pretrained_checkpoints/mobilenet_edgetpu_v2_m/ckpt-171600'
+ ),
+ 'mobilenet_edgetpu_v2_s': (
+ PRETRAINED_CKPT_PATH_BASE
+ + '/pretrained_checkpoints/mobilenet_edgetpu_v2_s/ckpt-171600'
+ ),
+ 'mobilenet_edgetpu_v2_xs': (
+ PRETRAINED_CKPT_PATH_BASE
+ + '/pretrained_checkpoints/mobilenet_edgetpu_v2_xs/ckpt-171600'
+ ),
+ 'mobilenet_edgetpu_v2_tiny': (
+ PRETRAINED_CKPT_PATH_BASE
+ + '/pretrained_checkpoints/mobilenet_edgetpu_v2_tiny/ckpt-171600'
+ ),
}
BACKBONE_HEADPOINT = {
@@ -87,6 +104,7 @@ class CustomSemanticSegmentationTaskConfig(base_cfg.SemanticSegmentationTask):
'mobilenet_edgetpu_v2_m': 4,
'mobilenet_edgetpu_v2_s': 4,
'mobilenet_edgetpu_v2_xs': 4,
+ 'mobilenet_edgetpu_v2_tiny': 4,
}
BACKBONE_LOWER_FEATURES = {
@@ -94,6 +112,7 @@ class CustomSemanticSegmentationTaskConfig(base_cfg.SemanticSegmentationTask):
'mobilenet_edgetpu_v2_m': 3,
'mobilenet_edgetpu_v2_s': 3,
'mobilenet_edgetpu_v2_xs': 3,
+ 'mobilenet_edgetpu_v2_tiny': 3,
}
@@ -207,7 +226,15 @@ def seg_deeplabv3plus_ade20k(backbone: str):
# Experiment configs for 32 output classes
@exp_factory.register_config_factory(
- 'deeplabv3plus_mobilenet_edgetpuv2_m_ade20k_32')
+ 'deeplabv3plus_mobilenet_edgetpuv2_l_ade20k_32'
+)
+def deeplabv3plus_mobilenet_edgetpuv2_l_ade20k_32() -> cfg.ExperimentConfig:
+ return seg_deeplabv3plus_ade20k_32('mobilenet_edgetpu_v2_l')
+
+
+@exp_factory.register_config_factory(
+ 'deeplabv3plus_mobilenet_edgetpuv2_m_ade20k_32'
+)
def deeplabv3plus_mobilenet_edgetpuv2_m_ade20k_32() -> cfg.ExperimentConfig:
return seg_deeplabv3plus_ade20k_32('mobilenet_edgetpu_v2_m')
@@ -224,23 +251,42 @@ def deeplabv3plus_mobilenet_edgetpuv2_xs_ade20k_32() -> cfg.ExperimentConfig:
return seg_deeplabv3plus_ade20k_32('mobilenet_edgetpu_v2_xs')
+@exp_factory.register_config_factory(
+ 'deeplabv3plus_mobilenet_edgetpuv2_tiny_ade20k_32'
+)
+def deeplabv3plus_mobilenet_edgetpuv2_tiny_ade20k_32() -> cfg.ExperimentConfig:
+ return seg_deeplabv3plus_ade20k_32('mobilenet_edgetpu_v2_tiny')
+
+
# Experiment configs for 151 output classes
@exp_factory.register_config_factory(
- 'deeplabv3plus_mobilenet_edgetpuv2_m_ade20k')
+ 'deeplabv3plus_mobilenet_edgetpuv2_l_ade20k'
+)
+def deeplabv3plus_mobilenet_edgetpuv2_l_ade20k() -> cfg.ExperimentConfig:
+ return seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_l')
+
+
+@exp_factory.register_config_factory(
+ 'deeplabv3plus_mobilenet_edgetpuv2_m_ade20k'
+)
def deeplabv3plus_mobilenet_edgetpuv2_m_ade20k() -> cfg.ExperimentConfig:
- config = seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_m')
- return config
+ return seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_m')
@exp_factory.register_config_factory(
'deeplabv3plus_mobilenet_edgetpuv2_s_ade20k')
def deeplabv3plus_mobilenet_edgetpuv2_s_ade20k() -> cfg.ExperimentConfig:
- config = seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_s')
- return config
+ return seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_s')
@exp_factory.register_config_factory(
'deeplabv3plus_mobilenet_edgetpuv2_xs_ade20k')
def deeplabv3plus_mobilenet_edgetpuv2_xs_ade20k() -> cfg.ExperimentConfig:
- config = seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_xs')
- return config
+ return seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_xs')
+
+
+@exp_factory.register_config_factory(
+ 'deeplabv3plus_mobilenet_edgetpuv2_tiny_ade20k'
+)
+def deeplabv3plus_mobilenet_edgetpuv2_tiny_ade20k() -> cfg.ExperimentConfig:
+ return seg_deeplabv3plus_ade20k('mobilenet_edgetpu_v2_tiny')
diff --git a/official/projects/edgetpu/vision/configs/semantic_segmentation_searched_config.py b/official/projects/edgetpu/vision/configs/semantic_segmentation_searched_config.py
index 87213ff6dc7..126492d7064 100644
--- a/official/projects/edgetpu/vision/configs/semantic_segmentation_searched_config.py
+++ b/official/projects/edgetpu/vision/configs/semantic_segmentation_searched_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,8 +20,6 @@
import os
from typing import Any, List, Mapping, Optional
-# Import libraries
-
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.modeling import hyperparams
@@ -72,8 +70,10 @@ class AutosegEdgeTPUModelConfig(hyperparams.Config):
"""Autoseg-EdgeTPU segmentation model config."""
num_classes: int = 0
input_size: List[int] = dataclasses.field(default_factory=list)
- backbone: backbones.Backbone = backbones.Backbone()
- head: BiFPNHeadConfig = BiFPNHeadConfig()
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=backbones.Backbone
+ )
+ head: BiFPNHeadConfig = dataclasses.field(default_factory=BiFPNHeadConfig)
model_params: Mapping[str, Any] = dataclasses.field(
default_factory=lambda: { # pylint: disable=g-long-lambda
'model_name': 'autoseg_edgetpu_backbone_s',
@@ -91,10 +91,16 @@ class AutosegEdgeTPUModelConfig(hyperparams.Config):
class AutosegEdgeTPUTaskConfig(base_cfg.SemanticSegmentationTask):
"""The task config inherited from the base segmentation task."""
- model: AutosegEdgeTPUModelConfig = AutosegEdgeTPUModelConfig()
- train_data: base_cfg.DataConfig = base_cfg.DataConfig(is_training=True)
- validation_data: base_cfg.DataConfig = base_cfg.DataConfig(is_training=False)
- losses: Losses = Losses()
+ model: AutosegEdgeTPUModelConfig = dataclasses.field(
+ default_factory=AutosegEdgeTPUModelConfig
+ )
+ train_data: base_cfg.DataConfig = dataclasses.field(
+ default_factory=lambda: base_cfg.DataConfig(is_training=True)
+ )
+ validation_data: base_cfg.DataConfig = dataclasses.field(
+ default_factory=lambda: base_cfg.DataConfig(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
init_checkpoint: Optional[str] = None
init_checkpoint_modules: str = 'backbone' # all or backbone
model_output_keys: Optional[List[int]] = dataclasses.field(
diff --git a/official/projects/edgetpu/vision/dataloaders/__init__.py b/official/projects/edgetpu/vision/dataloaders/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/dataloaders/__init__.py
+++ b/official/projects/edgetpu/vision/dataloaders/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/dataloaders/classification_input.py b/official/projects/edgetpu/vision/dataloaders/classification_input.py
index 1c7f532d93b..b44297dc2a4 100644
--- a/official/projects/edgetpu/vision/dataloaders/classification_input.py
+++ b/official/projects/edgetpu/vision/dataloaders/classification_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,8 +13,7 @@
# limitations under the License.
"""Classification decoder and parser."""
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import classification_input
from official.vision.ops import preprocess_ops
diff --git a/official/projects/edgetpu/vision/dataloaders/classification_input_test.py b/official/projects/edgetpu/vision/dataloaders/classification_input_test.py
index ecd552b072d..12931af71ff 100644
--- a/official/projects/edgetpu/vision/dataloaders/classification_input_test.py
+++ b/official/projects/edgetpu/vision/dataloaders/classification_input_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests classification_input.py."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.dataloaders import classification_input
from official.vision.configs import common
from official.vision.dataloaders import tfexample_utils
diff --git a/official/projects/edgetpu/vision/modeling/__init__.py b/official/projects/edgetpu/vision/modeling/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/modeling/__init__.py
+++ b/official/projects/edgetpu/vision/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/modeling/backbones/__init__.py b/official/projects/edgetpu/vision/modeling/backbones/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/modeling/backbones/__init__.py
+++ b/official/projects/edgetpu/vision/modeling/backbones/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu.py b/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu.py
index 2fa8d10e597..91ea3542467 100644
--- a/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu.py
+++ b/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,17 +14,15 @@
"""Contains definitions of mobilenet_edgetpu_v2 Networks."""
-# Import libraries
-
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.edgetpu.vision.modeling.mobilenet_edgetpu_v1_model import MobilenetEdgeTPU
from official.projects.edgetpu.vision.modeling.mobilenet_edgetpu_v2_model import MobilenetEdgeTPUV2
from official.vision.modeling.backbones import factory
-layers = tf.keras.layers
+layers = tf_keras.layers
# MobileNet-EdgeTPU-V2 configs.
MOBILENET_EDGETPU_V2_CONFIGS = frozenset([
@@ -48,7 +46,7 @@
])
-def freeze_large_filters(model: tf.keras.Model, threshold: int):
+def freeze_large_filters(model: tf_keras.Model, threshold: int):
"""Freezes layer with large number of filters."""
for layer in model.layers:
if isinstance(layer.output_shape, tuple):
@@ -59,9 +57,9 @@ def freeze_large_filters(model: tf.keras.Model, threshold: int):
@factory.register_backbone_builder('mobilenet_edgetpu')
-def build_mobilenet_edgetpu(input_specs: tf.keras.layers.InputSpec,
+def build_mobilenet_edgetpu(input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
- **unused_kwargs) -> tf.keras.Model:
+ **unused_kwargs) -> tf_keras.Model:
"""Builds MobileNetEdgeTpu backbone from a config."""
backbone_type = backbone_config.type
backbone_cfg = backbone_config.get()
@@ -75,6 +73,7 @@ def build_mobilenet_edgetpu(input_specs: tf.keras.layers.InputSpec,
'batch_norm': 'tpu',
'rescale_input': False,
'resolution': input_specs.shape[1:3],
+ 'input_channels': input_specs.shape[3],
'backbone_only': True,
'features_as_dict': True,
'dtype': 'bfloat16'
@@ -90,6 +89,7 @@ def build_mobilenet_edgetpu(input_specs: tf.keras.layers.InputSpec,
'batch_norm': 'tpu',
'rescale_input': False,
'resolution': input_specs.shape[1:3],
+ 'input_channels': input_specs.shape[3],
'backbone_only': True,
'dtype': 'bfloat16'
},
diff --git a/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu_test.py b/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu_test.py
index 9043aeb0608..712a9af177b 100644
--- a/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu_test.py
+++ b/official/projects/edgetpu/vision/modeling/backbones/mobilenet_edgetpu_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,8 @@
"""Tests for MobileNet."""
-# Import libraries
-
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling.backbones import mobilenet_edgetpu
@@ -47,14 +45,17 @@ class MobileNetEdgeTPUTest(parameterized.TestCase, tf.test.TestCase):
('mobilenet_edgetpu_v2_l', (1, None, None, 3)),
('mobilenet_edgetpu', (1, 512, 512, 3)),
('mobilenet_edgetpu_dm1p25', (1, None, None, 3)),
+ ('mobilenet_edgetpu', (1, None, None, 6)),
+ ('mobilenet_edgetpu_v2_tiny', (1, None, None, 6)),
)
def test_mobilenet_creation(self, model_id, input_shape):
"""Test creation of MobileNet family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
test_model = mobilenet_edgetpu.build_mobilenet_edgetpu(
input_specs=TestInputSpec(input_shape),
backbone_config=TestBackboneConfig(model_id))
+
self.assertGreater(len(test_model.outputs), 1)
diff --git a/official/projects/edgetpu/vision/modeling/common_modules.py b/official/projects/edgetpu/vision/modeling/common_modules.py
index 284a2e8e46f..13e437a9c06 100644
--- a/official/projects/edgetpu/vision/modeling/common_modules.py
+++ b/official/projects/edgetpu/vision/modeling/common_modules.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,8 @@
"""Common modeling utilities."""
from typing import Optional, Tuple
-# Import libraries
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow.compat.v1 as tf1
from tensorflow.python.tpu import tpu_function # pylint: disable=g-direct-tensorflow-import
@@ -26,8 +25,8 @@
STDDEV_RGB = (0.5 * 255, 0.5 * 255, 0.5 * 255)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class TpuBatchNormalization(tf.keras.layers.BatchNormalization):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class TpuBatchNormalization(tf_keras.layers.BatchNormalization):
"""Cross replica batch normalization."""
def __init__(self, fused: Optional[bool] = False, **kwargs):
@@ -51,10 +50,14 @@ def _cross_replica_average(self, t: tf.Tensor, num_shards_per_group: int):
return tf1.tpu.cross_replica_sum(t, group_assignment) / tf.cast(
num_shards_per_group, t.dtype)
- def _moments(self, inputs: tf.Tensor, reduction_axes: int, keep_dims: int):
+ def _moments(self,
+ inputs: tf.Tensor,
+ reduction_axes: int,
+ keep_dims: int,
+ mask: Optional[tf.Tensor] = None):
"""Compute the mean and variance: it overrides the original _moments."""
shard_mean, shard_variance = super(TpuBatchNormalization, self)._moments(
- inputs, reduction_axes, keep_dims=keep_dims)
+ inputs, reduction_axes, keep_dims=keep_dims, mask=mask)
num_shards = tpu_function.get_tpu_context().number_of_shards or 1
if num_shards <= 8: # Skip cross_replica for 2x2 or smaller slices.
@@ -74,7 +77,7 @@ def _moments(self, inputs: tf.Tensor, reduction_axes: int, keep_dims: int):
return (shard_mean, shard_variance)
-def get_batch_norm(batch_norm_type: str) -> tf.keras.layers.BatchNormalization:
+def get_batch_norm(batch_norm_type: str) -> tf_keras.layers.BatchNormalization:
"""A helper to create a batch normalization getter.
Args:
@@ -82,12 +85,12 @@ def get_batch_norm(batch_norm_type: str) -> tf.keras.layers.BatchNormalization:
will use `TpuBatchNormalization`.
Returns:
- An instance of `tf.keras.layers.BatchNormalization`.
+ An instance of `tf_keras.layers.BatchNormalization`.
"""
if batch_norm_type == 'tpu':
return TpuBatchNormalization
- return tf.keras.layers.BatchNormalization # pytype: disable=bad-return-type # typed-keras
+ return tf_keras.layers.BatchNormalization # pytype: disable=bad-return-type # typed-keras
def count_params(model, trainable_only=True):
@@ -95,11 +98,11 @@ def count_params(model, trainable_only=True):
if not trainable_only:
return model.count_params()
else:
- return int(np.sum([tf.keras.backend.count_params(p)
+ return int(np.sum([tf_keras.backend.count_params(p)
for p in model.trainable_weights]))
-def load_weights(model: tf.keras.Model,
+def load_weights(model: tf_keras.Model,
model_weights_path: str,
checkpoint_format: str = 'tf_checkpoint'):
"""Load model weights from the given file path.
diff --git a/official/projects/edgetpu/vision/modeling/custom_layers.py b/official/projects/edgetpu/vision/modeling/custom_layers.py
index 8a78f6b6464..3eda96ac028 100644
--- a/official/projects/edgetpu/vision/modeling/custom_layers.py
+++ b/official/projects/edgetpu/vision/modeling/custom_layers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,22 +14,25 @@
"""Customized keras layers used in the EdgeTPU models."""
+from collections.abc import MutableMapping
import inspect
-from typing import Any, MutableMapping, Optional, Union, Tuple
-import tensorflow as tf
+from typing import Any, Optional, Union
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
-class GroupConv2D(tf.keras.layers.Conv2D):
+
+class GroupConv2D(tf_keras.layers.Conv2D):
"""2D group convolution as a Keras Layer."""
def __init__(self,
filters: int,
- kernel_size: Union[int, Tuple[int, int]],
+ kernel_size: Union[int, tuple[int, int]],
groups: int,
- strides: Tuple[int, int] = (1, 1),
+ strides: tuple[int, int] = (1, 1),
padding: str = 'valid',
data_format: str = 'channels_last',
- dilation_rate: Tuple[int, int] = (1, 1),
+ dilation_rate: tuple[int, int] = (1, 1),
activation: Any = None,
use_bias: bool = True,
kernel_initializer: Any = 'glorot_uniform',
@@ -39,10 +42,10 @@ def __init__(self,
activity_regularizer: Any = None,
kernel_constraint: Any = None,
bias_constraint: Any = None,
- batch_norm_layer: Optional[tf.keras.layers.Layer] = None,
+ batch_norm_layer: Optional[tf_keras.layers.Layer] = None,
bn_epsilon: float = 1e-3,
bn_momentum: float = 0.99,
- **kwargs: Any) -> tf.keras.layers.Layer:
+ **kwargs: Any) -> tf_keras.layers.Layer: # pyrefly: ignore[bad-return]
"""Creates a 2D group convolution keras layer.
Args:
@@ -81,7 +84,7 @@ def __init__(self,
bias_constraint: Constraint function applied to the bias vector ( see
`keras.constraints`).
batch_norm_layer: The batch normalization layer to use. This is typically
- tf.keras.layer.BatchNormalization or a derived class.
+ tf_keras.layer.BatchNormalization or a derived class.
bn_epsilon: Batch normalization epsilon.
bn_momentum: Momentum used for moving average in batch normalization.
**kwargs: Additional keyword arguments.
@@ -102,9 +105,9 @@ def __init__(self,
ValueError: if `batch_norm_layer` is not a callable when provided.
ValueError: when both `strides` > 1 and `dilation_rate` > 1.
"""
- if groups <= 1 or groups >= filters:
- raise ValueError('Number of groups should be greater than 1 and less '
- 'than the output filters.')
+ if groups <= 1 or groups > filters:
+ raise ValueError(f'Number of groups {groups} should be greater than 1 and'
+ f' less or equal than the output filters {filters}.')
self._groups = groups
if data_format != 'channels_last':
raise ValueError(
@@ -123,8 +126,8 @@ def __init__(self,
self.batch_norm_layer = []
if self.use_batch_norm:
self.batch_norm_layer = [
- batch_norm_layer(
- axis=-1, momentum=self.bn_momentum, epsilon=self.bn_epsilon)
+ batch_norm_layer( # pyrefly: ignore[not-callable]
+ axis=-1, momentum=self.bn_momentum, epsilon=self.bn_epsilon) # pyrefly: ignore[unexpected-keyword]
for i in range(self._groups)
]
@@ -147,7 +150,7 @@ def __init__(self,
groups=1,
**kwargs) # pytype: disable=bad-return-type # typed-keras
- def build(self, input_shape: Tuple[int, ...]) -> None:
+ def build(self, input_shape: tuple[int, ...]) -> None:
"""Builds GroupConv2D layer as a collection of smaller Conv2D layers."""
input_shape = tf.TensorShape(input_shape)
input_channel = self._get_input_channel(input_shape)
@@ -157,7 +160,7 @@ def build(self, input_shape: Tuple[int, ...]) -> None:
f'by number of groups: {self._groups}.')
self.group_input_channel = int(input_channel / self._groups)
- self.group_output_channel = int(self.filters / self._groups)
+ self.group_output_channel = int(self.filters / self._groups) # pyrefly: ignore[unsupported-operation]
self.group_kernel_shape = self.kernel_size + (self.group_input_channel,
self.group_output_channel)
@@ -168,7 +171,7 @@ def build(self, input_shape: Tuple[int, ...]) -> None:
self.add_weight(
name='kernel_{}'.format(g),
shape=self.group_kernel_shape,
- initializer=self.kernel_initializer,
+ initializer=tf_utils.clone_initializer(self.kernel_initializer),
regularizer=self.kernel_regularizer,
constraint=self.kernel_constraint,
trainable=True,
@@ -178,13 +181,13 @@ def build(self, input_shape: Tuple[int, ...]) -> None:
self.add_weight(
name='bias_{}'.format(g),
shape=(self.group_output_channel,),
- initializer=self.bias_initializer,
+ initializer=tf_utils.clone_initializer(self.bias_initializer),
regularizer=self.bias_regularizer,
constraint=self.bias_constraint,
trainable=True,
dtype=self.dtype))
channel_axis = self._get_channel_axis()
- self.input_spec = tf.keras.layers.InputSpec(
+ self.input_spec = tf_keras.layers.InputSpec(
ndim=self.rank + 2, axes={channel_axis: input_channel})
self._build_conv_op_data_shape = input_shape[-(self.rank + 1):]
@@ -216,7 +219,7 @@ def call(self, inputs: Any, training: Optional[bool] = None) -> Any:
if self.use_bias:
output_slice = tf.nn.bias_add(
- output_slice, self.bias[i], data_format='NHWC')
+ output_slice, self.bias[i], data_format='NHWC') # pyrefly: ignore[unsupported-operation]
# Apply batch norm after bias addition.
if self.use_batch_norm:
@@ -264,19 +267,19 @@ def from_config(cls, config):
return cls(**config)
-class GroupConv2DKerasModel(tf.keras.Model):
+class GroupConv2DKerasModel(tf_keras.Model):
"""2D group convolution as a keras model."""
def __init__(self,
filters: int,
- kernel_size: Tuple[int, int],
+ kernel_size: tuple[int, int],
groups: int,
- batch_norm_layer: Optional[tf.keras.layers.Layer] = None,
+ batch_norm_layer: Optional[tf_keras.layers.Layer] = None,
bn_epsilon: float = 1e-3,
bn_momentum: float = 0.99,
data_format: str = 'channels_last',
padding: str = 'valid',
- **kwargs: Any) -> tf.keras.Model:
+ **kwargs: Any) -> tf_keras.Model: # pyrefly: ignore[bad-return]
"""Creates a 2D group convolution layer as a keras model.
Args:
@@ -287,7 +290,7 @@ def __init__(self,
specify the same value for all spatial dimensions.
groups: The number of input/output channel groups.
batch_norm_layer: The batch normalization layer to use. This is typically
- tf.keras.layer.BatchNormalization or a derived class.
+ tf_keras.layer.BatchNormalization or a derived class.
bn_epsilon: Batch normalization epsilon.
bn_momentum: Momentum used for moving average in batch normalization.
data_format: The ordering of the dimensions in the inputs. `channels_last`
@@ -314,12 +317,12 @@ def __init__(self,
self.batch_norm_layer = batch_norm_layer
self.use_batch_norm = False
if self.batch_norm_layer is not None:
- if not inspect.isclass(self.batch_norm_layer):
+ if not inspect.isclass(self.batch_norm_layer): # pytype: disable=not-supported-yet
raise ValueError('batch_norm_layer is not a class.')
self.use_batch_norm = True
if 'activation' in kwargs.keys():
- self.activation = tf.keras.activations.get(kwargs['activation'])
+ self.activation = tf_keras.activations.get(kwargs['activation'])
kwargs.pop('activation')
else:
self.activation = None
@@ -335,15 +338,15 @@ def __init__(self,
for _ in range(self._groups):
# Override the activation so that batchnorm can be applied after the conv.
self.conv_layers.append(
- tf.keras.layers.Conv2D(per_conv_filter_size, kernel_size, **kwargs))
+ tf_keras.layers.Conv2D(per_conv_filter_size, kernel_size, **kwargs))
if self.use_batch_norm:
for _ in range(self._groups):
self.bn_layers.append(
- self.batch_norm_layer(
+ self.batch_norm_layer( # pyrefly: ignore[not-callable]
axis=-1, momentum=bn_momentum, epsilon=bn_epsilon)) # pytype: disable=bad-return-type # typed-keras
- def call(self, inputs: Any) -> Any:
+ def call(self, inputs: Any) -> Any: # pytype: disable=signature-mismatch # overriding-parameter-count-checks
"""Applies 2d group convolution on the inputs."""
input_shape = inputs.get_shape().as_list()
if input_shape[-1] % self._groups != 0:
@@ -356,7 +359,7 @@ def call(self, inputs: Any) -> Any:
output_slice = self.conv_layers[g](input_slices[g])
if self.use_batch_norm:
output_slice = self.bn_layers[g](output_slice)
- output_slice = self.activation(output_slice)
+ output_slice = self.activation(output_slice) # pyrefly: ignore[not-callable]
output_slices.append(output_slice)
outputs = tf.concat(output_slices, axis=-1)
@@ -392,7 +395,7 @@ def argmax(input_tensor,
which axis of the input Tensor to reduce across. For vectors, use axis =
0.
output_type: An optional tf.DType. Note that default is different from
- tflite (int64) to make default behavior compatible with darwinn.
+ tflite (int64) to make default behavior compatible with Edge TPU.
name: Optional name for operations.
keepdims: If true, retains reduced dimensions with length 1.
epsilon: Optional small number which is intended to be always below
@@ -443,14 +446,14 @@ def argmax(input_tensor,
name=name))
-class ArgmaxKerasLayer(tf.keras.layers.Layer):
+class ArgmaxKerasLayer(tf_keras.layers.Layer):
"""Implements argmax as a keras model."""
def __init__(self,
axis=-1,
name=None,
output_type=tf.dtypes.int32,
- **kwargs: Any) -> tf.keras.Model:
+ **kwargs: Any) -> tf_keras.Model: # pyrefly: ignore[bad-return]
"""Implements argmax as a keras model.
Args:
diff --git a/official/projects/edgetpu/vision/modeling/custom_layers_test.py b/official/projects/edgetpu/vision/modeling/custom_layers_test.py
index c07ce224ee3..2d8d6ab9684 100644
--- a/official/projects/edgetpu/vision/modeling/custom_layers_test.py
+++ b/official/projects/edgetpu/vision/modeling/custom_layers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import itertools
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling import custom_layers
GROUPS = [2, 4]
@@ -25,7 +25,7 @@
OUTPUT_CHANNEL = [8, 16]
USE_BATCH_NORM = [True, False]
ACTIVATION = ['relu', 'linear']
-BATCH_NORM_LAYER = tf.keras.layers.BatchNormalization
+BATCH_NORM_LAYER = tf_keras.layers.BatchNormalization
# 2 functionally identical group conv implementations.
GROUP_CONV_IMPL = {
@@ -121,7 +121,7 @@ def test_equivalence(self, groups, input_channel, output_channel,
use_bias=False,
batch_norm_layer=batch_norm_layer,
activation=activation)
- gc_layer = tf.keras.Sequential([custom_layers.GroupConv2D(**kwargs)])
+ gc_layer = tf_keras.Sequential([custom_layers.GroupConv2D(**kwargs)])
gc_model = custom_layers.GroupConv2DKerasModel(**kwargs)
gc_layer.build(input_shape=(None, 3, 3, input_channel))
gc_model.build(input_shape=(None, 3, 3, input_channel))
@@ -184,7 +184,3 @@ def test_reference_match(self, shape, input_type, output_type):
test_output = custom_layers.argmax(
random_inputs, axis=axis, output_type=output_type)
self.assertAllEqual(control_output, test_output)
-
-
-if __name__ == '__main__':
- tf.test.main()
diff --git a/official/projects/edgetpu/vision/modeling/heads/__init__.py b/official/projects/edgetpu/vision/modeling/heads/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/modeling/heads/__init__.py
+++ b/official/projects/edgetpu/vision/modeling/heads/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/modeling/heads/bifpn_head.py b/official/projects/edgetpu/vision/modeling/heads/bifpn_head.py
index 6fcd3d48093..e474d104427 100644
--- a/official/projects/edgetpu/vision/modeling/heads/bifpn_head.py
+++ b/official/projects/edgetpu/vision/modeling/heads/bifpn_head.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,11 +17,9 @@
import itertools
from typing import Text, Optional
-# Import libraries
-
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling import common_modules
@@ -71,7 +69,7 @@ def build_batch_norm(is_training_bn: bool,
if is_training_bn:
batch_norm_class = common_modules.get_batch_norm(strategy)
else:
- batch_norm_class = tf.keras.layers.BatchNormalization
+ batch_norm_class = tf_keras.layers.BatchNormalization
bn_layer = batch_norm_class(
axis=axis,
@@ -141,12 +139,12 @@ def get_conv_op(conv_type):
kernel_size = int(conv_type.split('_')[-1])
if conv_type.startswith('sep'):
conv_op = functools.partial(
- tf.keras.layers.SeparableConv2D,
+ tf_keras.layers.SeparableConv2D,
depth_multiplier=1,
kernel_size=(kernel_size, kernel_size))
elif conv_type.startswith('conv'):
conv_op = functools.partial(
- tf.keras.layers.Conv2D, kernel_size=(kernel_size, kernel_size))
+ tf_keras.layers.Conv2D, kernel_size=(kernel_size, kernel_size))
else:
raise ValueError('Unknown conv type: {}'.format(conv_type))
return conv_op
@@ -218,7 +216,7 @@ def resize(feat,
return tf.cast(feat, dtype)
-class ResampleFeatureMap(tf.keras.layers.Layer):
+class ResampleFeatureMap(tf_keras.layers.Layer):
"""Resamples feature map for downsampling or upsampling."""
def __init__(self,
@@ -248,14 +246,14 @@ def _pool2d(self, inputs, height, width, target_height, target_width):
height_stride_size = int((height - 1) // target_height + 1)
width_stride_size = int((width - 1) // target_width + 1)
if self.pooling_type == 'max':
- return tf.keras.layers.MaxPooling2D(
+ return tf_keras.layers.MaxPooling2D(
pool_size=[height_stride_size + 1, width_stride_size + 1],
strides=[height_stride_size, width_stride_size],
padding='SAME',
data_format=self.data_format)(
inputs)
if self.pooling_type == 'avg':
- return tf.keras.layers.AveragePooling2D(
+ return tf_keras.layers.AveragePooling2D(
pool_size=[height_stride_size + 1, width_stride_size + 1],
strides=[height_stride_size, width_stride_size],
padding='SAME',
@@ -278,14 +276,14 @@ def _maybe_apply_1x1(self, feat, training, num_channels):
def build(self, feat_shape):
num_channels = self.target_num_channels or feat_shape[-1]
- self.conv2d = tf.keras.layers.Conv2D(
+ self.conv2d = tf_keras.layers.Conv2D(
num_channels, (1, 1),
padding='same',
data_format=self.data_format,
name='conv2d')
self.bn = build_batch_norm(
- is_training_bn=self.is_training_bn,
- data_format=self.data_format,
+ is_training_bn=self.is_training_bn, # pyrefly: ignore[bad-argument-type]
+ data_format=self.data_format, # pyrefly: ignore[bad-argument-type]
strategy=self.strategy,
name='bn')
self.built = True
@@ -321,7 +319,7 @@ def call(self, feat, training, all_feats):
return feat
-class FNode(tf.keras.layers.Layer):
+class FNode(tf_keras.layers.Layer):
"""A Keras Layer implementing BiFPN Node."""
def __init__(self,
@@ -460,7 +458,7 @@ def call(self, feats, training):
return feats + [new_node]
-class OpAfterCombine(tf.keras.layers.Layer):
+class OpAfterCombine(tf_keras.layers.Layer):
"""Operation after combining input features during feature fusiong."""
def __init__(self,
@@ -501,7 +499,7 @@ def call(self, new_node, training):
return new_node
-class FPNCells(tf.keras.layers.Layer):
+class FPNCells(tf_keras.layers.Layer):
"""FPN cells."""
def __init__(self,
@@ -565,7 +563,7 @@ def call(self, feats, training):
return feats
-class FPNCell(tf.keras.layers.Layer):
+class FPNCell(tf_keras.layers.Layer):
"""A single FPN cell."""
def __init__(self,
@@ -619,7 +617,7 @@ def _call(feats):
return _call(feats)
-class SegClassNet(tf.keras.layers.Layer):
+class SegClassNet(tf_keras.layers.Layer):
"""Segmentation class prediction network."""
def __init__(self,
@@ -670,7 +668,7 @@ def __init__(self,
self.min_level = min_level
self.max_level = max_level
self.fullres_output = fullres_output
- self.fullres_conv_transpose = fullres_skip_connections
+ self.fullres_skip_connections = fullres_skip_connections
self.fnode = FNode(
0, # Always use the first level with highest resolution.
@@ -703,7 +701,7 @@ def __init__(self,
padding='same',
activation=act_type,
name='fullres_conv_%d' % i)
- self.fullres_conv_transpose[str(i)] = tf.keras.layers.Conv2DTranspose(
+ self.fullres_conv_transpose[str(i)] = tf_keras.layers.Conv2DTranspose(
filters=num_filters,
data_format=data_format,
kernel_size=3,
@@ -726,8 +724,8 @@ def call(self, inputs, backbone_feats, training):
if self.fullres_output:
for i in reversed(range(self.min_level)):
- if self.config.fullres_skip_connections:
- net = tf.keras.layers.Concatenate()([net, backbone_feats[i + 1]])
+ if self.fullres_skip_connections:
+ net = tf_keras.layers.Concatenate()([net, backbone_feats[i + 1]])
net = self.fullres_conv[str(i)](net)
net = self.fullres_conv_transpose[str(i)](net)
diff --git a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model.py b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model.py
index fa3f36cc55a..1223b6bda72 100644
--- a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model.py
+++ b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,8 @@
"""Contains definitions for MobilenetEdgeTPU image classification models."""
from typing import Any, Dict, Optional, Text
-# Import libraries
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling import common_modules
from official.projects.edgetpu.vision.modeling import mobilenet_edgetpu_v1_model_blocks
@@ -34,8 +33,8 @@
}
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class MobilenetEdgeTPU(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MobilenetEdgeTPU(tf_keras.Model):
"""Wrapper class for a MobilenetEdgeTPU Keras model.
Contains helper methods to build, manage, and save metadata about the model.
@@ -63,7 +62,7 @@ def __init__(self,
else:
input_shape = (self.config.resolution, self.config.resolution,
input_channels)
- image_input = tf.keras.layers.Input(shape=input_shape)
+ image_input = tf_keras.layers.Input(shape=input_shape)
output = mobilenet_edgetpu_v1_model_blocks.mobilenet_edgetpu(
image_input, self.config)
diff --git a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_blocks.py b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_blocks.py
index 29d93d3d92b..937abd1126d 100644
--- a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_blocks.py
+++ b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,9 +17,8 @@
import math
from typing import Any, Optional, Tuple, Union
-# Import libraries
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.modeling.hyperparams import base_config
@@ -121,7 +120,7 @@ def round_filters(filters: int,
if not width_coefficient:
return filters
- filters *= width_coefficient
+ filters *= width_coefficient # pyrefly: ignore[bad-assignment]
min_depth = min_depth or divisor
new_filters = max(min_depth, int(filters + divisor / 2) // divisor * divisor)
# Make sure that round down does not go down by more than 10%.
@@ -150,7 +149,7 @@ def conv2d_block(inputs: tf.Tensor,
batch_norm = common_modules.get_batch_norm(config.batch_norm)
bn_momentum = config.bn_momentum
bn_epsilon = config.bn_epsilon
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
weight_decay = config.weight_decay
name = name or ''
@@ -162,19 +161,19 @@ def conv2d_block(inputs: tf.Tensor,
'use_bias': use_bias,
'padding': 'same',
'name': name + '_conv2d',
- 'kernel_regularizer': tf.keras.regularizers.l2(weight_decay),
- 'bias_regularizer': tf.keras.regularizers.l2(weight_decay),
+ 'kernel_regularizer': tf_keras.regularizers.l2(weight_decay),
+ 'bias_regularizer': tf_keras.regularizers.l2(weight_decay),
}
if depthwise:
- conv2d = tf.keras.layers.DepthwiseConv2D
+ conv2d = tf_keras.layers.DepthwiseConv2D
init_kwargs.update({'depthwise_initializer': CONV_KERNEL_INITIALIZER})
else:
- conv2d = tf.keras.layers.Conv2D
+ conv2d = tf_keras.layers.Conv2D
init_kwargs.update({'filters': conv_filters,
'kernel_initializer': CONV_KERNEL_INITIALIZER})
- x = conv2d(**init_kwargs)(inputs)
+ x = conv2d(**init_kwargs)(inputs) # pyrefly: ignore[missing-argument]
if use_batch_norm:
bn_axis = 1 if data_format == 'channels_first' else -1
@@ -184,7 +183,7 @@ def conv2d_block(inputs: tf.Tensor,
name=name + '_bn')(x)
if activation is not None:
- x = tf.keras.layers.Activation(activation,
+ x = tf_keras.layers.Activation(activation,
name=name + '_activation')(x)
return x
@@ -207,7 +206,7 @@ def mb_conv_block(inputs: tf.Tensor,
use_se = config.use_se
activation = tf_utils.get_activation(config.activation)
drop_connect_rate = config.drop_connect_rate
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
use_depthwise = block.conv_type == 'depthwise'
prefix = prefix or ''
@@ -259,8 +258,8 @@ def mb_conv_block(inputs: tf.Tensor,
else:
se_shape = (1, 1, filters)
- se = tf.keras.layers.GlobalAveragePooling2D(name=prefix + 'se_squeeze')(x)
- se = tf.keras.layers.Reshape(se_shape, name=prefix + 'se_reshape')(se)
+ se = tf_keras.layers.GlobalAveragePooling2D(name=prefix + 'se_squeeze')(x)
+ se = tf_keras.layers.Reshape(se_shape, name=prefix + 'se_reshape')(se)
se = conv2d_block(se,
num_reduced_filters,
@@ -276,7 +275,7 @@ def mb_conv_block(inputs: tf.Tensor,
use_batch_norm=False,
activation='sigmoid',
name=prefix + 'se_expand')
- x = tf.keras.layers.multiply([x, se], name=prefix + 'se_excite')
+ x = tf_keras.layers.multiply([x, se], name=prefix + 'se_excite')
# Output phase
x = conv2d_block(x,
@@ -287,7 +286,7 @@ def mb_conv_block(inputs: tf.Tensor,
# Add identity so that quantization-aware training can insert quantization
# ops correctly.
- x = tf.keras.layers.Activation('linear', name=prefix + 'id')(x)
+ x = tf_keras.layers.Activation('linear', name=prefix + 'id')(x)
if (block.id_skip
and all(s == 1 for s in block.strides)
@@ -297,20 +296,20 @@ def mb_conv_block(inputs: tf.Tensor,
# The only difference between dropout and dropconnect in TF is scaling by
# drop_connect_rate during training. See:
# https://github.com/keras-team/keras/pull/9898#issuecomment-380577612
- x = tf.keras.layers.Dropout(drop_connect_rate,
+ x = tf_keras.layers.Dropout(drop_connect_rate,
noise_shape=(None, 1, 1, 1),
name=prefix + 'drop')(x)
- x = tf.keras.layers.add([x, inputs], name=prefix + 'add')
+ x = tf_keras.layers.add([x, inputs], name=prefix + 'add')
return x
-def mobilenet_edgetpu(image_input: tf.keras.layers.Input, config: ModelConfig): # pytype: disable=invalid-annotation # typed-keras
+def mobilenet_edgetpu(image_input: tf_keras.layers.Input, config: ModelConfig): # pytype: disable=invalid-annotation # typed-keras
"""Creates a MobilenetEdgeTPU graph given the model parameters.
This function is wrapped by the `MobilenetEdgeTPU` class to make a
- tf.keras.Model.
+ tf_keras.Model.
Args:
image_input: the input batch of images
@@ -330,14 +329,14 @@ def mobilenet_edgetpu(image_input: tf.keras.layers.Input, config: ModelConfig):
num_classes = config.num_classes
input_channels = config.input_channels
rescale_input = config.rescale_input
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
dtype = config.dtype
weight_decay = config.weight_decay
x = image_input
if data_format == 'channels_first':
# Happens on GPU/TPU if available.
- x = tf.keras.layers.Permute((3, 1, 2))(x)
+ x = tf_keras.layers.Permute((3, 1, 2))(x)
if rescale_input:
x = common_modules.normalize_images(
x, num_channels=input_channels, dtype=dtype, data_format=data_format)
@@ -396,18 +395,18 @@ def mobilenet_edgetpu(image_input: tf.keras.layers.Input, config: ModelConfig):
# Build classifier
pool_size = (x.shape.as_list()[1], x.shape.as_list()[2])
- x = tf.keras.layers.AveragePooling2D(pool_size, name='top_pool')(x)
+ x = tf_keras.layers.AveragePooling2D(pool_size, name='top_pool')(x)
if dropout_rate and dropout_rate > 0:
- x = tf.keras.layers.Dropout(dropout_rate, name='top_dropout')(x)
- x = tf.keras.layers.Conv2D(
+ x = tf_keras.layers.Dropout(dropout_rate, name='top_dropout')(x)
+ x = tf_keras.layers.Conv2D(
num_classes,
1,
kernel_initializer=DENSE_KERNEL_INITIALIZER,
- kernel_regularizer=tf.keras.regularizers.l2(weight_decay),
- bias_regularizer=tf.keras.regularizers.l2(weight_decay),
+ kernel_regularizer=tf_keras.regularizers.l2(weight_decay),
+ bias_regularizer=tf_keras.regularizers.l2(weight_decay),
name='logits')(
x)
- x = tf.keras.layers.Activation('softmax', name='probs')(x)
+ x = tf_keras.layers.Activation('softmax', name='probs')(x)
x = tf.squeeze(x, axis=[1, 2])
return x
diff --git a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_test.py b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_test.py
index a4ca070a908..4c31abdc30e 100644
--- a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_test.py
+++ b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v1_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import os
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.image_classification import preprocessing
from official.projects.edgetpu.vision.modeling import common_modules
from official.projects.edgetpu.vision.modeling import mobilenet_edgetpu_v1_model
@@ -47,13 +47,13 @@ class MobilenetEdgeTPUBlocksTest(tf.test.TestCase):
def setUp(self):
super(tf.test.TestCase, self).setUp()
# Ensure no model duplicates
- tf.keras.backend.clear_session()
+ tf_keras.backend.clear_session()
def test_bottleneck_block(self):
"""Test for creating a model with bottleneck block arguments."""
images = tf.zeros((4, 224, 224, 3), dtype=tf.float32)
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
blocks = [
mobilenet_edgetpu_v1_model_blocks.BlockConfig.from_args(
@@ -121,7 +121,7 @@ def test_fused_bottleneck_block(self):
"""Test for creating a model with fused bottleneck block arguments."""
images = tf.zeros((4, 224, 224, 3), dtype=tf.float32)
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
blocks = [
mobilenet_edgetpu_v1_model_blocks.BlockConfig.from_args(
@@ -160,7 +160,7 @@ def test_variables(self):
"""Test for variables in blocks to be included in `model.variables`."""
images = tf.zeros((4, 224, 224, 3), dtype=tf.float32)
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
blocks = [
mobilenet_edgetpu_v1_model_blocks.BlockConfig.from_args(
@@ -195,7 +195,7 @@ class MobilenetEdgeTPUBuildTest(tf.test.TestCase):
def setUp(self):
super(tf.test.TestCase, self).setUp()
# Ensure no model duplicates
- tf.keras.backend.clear_session()
+ tf_keras.backend.clear_session()
def test_create_mobilenet_edgetpu(self):
model = mobilenet_edgetpu_v1_model.MobilenetEdgeTPU()
@@ -207,7 +207,7 @@ class MobilenetEdgeTPUPredictTest(tf.test.TestCase):
def setUp(self):
super(tf.test.TestCase, self).setUp()
# Ensure no model duplicates
- tf.keras.backend.clear_session()
+ tf_keras.backend.clear_session()
def _copy_saved_model_to_local(self, model_ckpt):
# Copy saved model to local first for speed
diff --git a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model.py b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model.py
index 9321cb47ec6..5bcbe78edfe 100644
--- a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model.py
+++ b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Any, Mapping, Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling import common_modules
from official.projects.edgetpu.vision.modeling import mobilenet_edgetpu_v2_model_blocks
@@ -45,8 +45,8 @@
}
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class MobilenetEdgeTPUV2(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MobilenetEdgeTPUV2(tf_keras.Model):
"""Wrapper class for a MobilenetEdgeTPUV2 Keras model.
Contains helper methods to build, manage, and save metadata about the model.
@@ -86,7 +86,7 @@ def __init__(self,
else:
input_shape = (self.config.resolution, self.config.resolution,
input_channels)
- image_input = tf.keras.layers.Input(shape=input_shape)
+ image_input = tf_keras.layers.Input(shape=input_shape)
output = mobilenet_edgetpu_v2_model_blocks.mobilenet_edgetpu_v2(
image_input, self.config)
diff --git a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks.py b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks.py
index a66c72a7c1f..3371ab7b9d6 100644
--- a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks.py
+++ b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,8 @@
import dataclasses
import math
from typing import Any, Dict, List, Optional, Tuple, Union
-# Import libraries
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.modeling.hyperparams import base_config
@@ -26,7 +25,7 @@
from official.projects.edgetpu.vision.modeling import common_modules
from official.projects.edgetpu.vision.modeling import custom_layers
-InitializerType = Optional[Union[str, tf.keras.initializers.Initializer]]
+InitializerType = Optional[Union[str, tf_keras.initializers.Initializer]]
@dataclasses.dataclass
@@ -43,7 +42,7 @@ class BlockType(oneof.OneOfConfig):
@dataclasses.dataclass
class BlockSearchConfig(base_config.Config):
"""Config for searchable BlockConfig parameters."""
- op_type: BlockType = BlockType()
+ op_type: BlockType = dataclasses.field(default_factory=BlockType)
kernel_size: Optional[int] = None
expand_ratio: Optional[int] = None
stride: Optional[int] = None
@@ -363,7 +362,7 @@ def autoseg_edgetpu_backbone_base(
'drop_connect_rate': drop_connect_rate,
}
if blocks_overrides:
- param_overrides['blocks'] = blocks_overrides
+ param_overrides['blocks'] = blocks_overrides # pyrefly: ignore[bad-assignment]
config = config.replace(**param_overrides)
return config
@@ -661,7 +660,7 @@ def round_filters(filters: int,
if not width_coefficient:
return filters
- filters *= width_coefficient
+ filters *= width_coefficient # pyrefly: ignore[bad-assignment]
min_depth = min_depth or divisor
new_filters = max(min_depth, int(filters + divisor / 2) // divisor * divisor)
# Make sure that round down does not go down by more than 10%.
@@ -684,22 +683,22 @@ def groupconv2d_block(conv_filters: Optional[int],
use_batch_norm: bool = True,
use_bias: bool = False,
activation: Any = None,
- name: Optional[str] = None) -> tf.keras.layers.Layer:
+ name: Optional[str] = None) -> tf_keras.layers.Layer:
"""2D group convolution with batchnorm and activation."""
batch_norm = common_modules.get_batch_norm(config.batch_norm)
bn_momentum = config.bn_momentum
bn_epsilon = config.bn_epsilon
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
weight_decay = config.weight_decay
if group_size is None:
group_size = config.group_base_size
name = name or ''
# Compute the # of groups
- if conv_filters % group_size != 0:
+ if conv_filters % group_size != 0: # pyrefly: ignore[unsupported-operation]
raise ValueError(f'Number of filters: {conv_filters} is not divisible by '
f'size of the groups: {group_size}')
- groups = int(conv_filters / group_size)
+ groups = int(conv_filters / group_size) # pyrefly: ignore[unsupported-operation]
# Collect args based on what kind of groupconv2d block is desired
init_kwargs = {
'kernel_size': kernel_size,
@@ -707,8 +706,8 @@ def groupconv2d_block(conv_filters: Optional[int],
'use_bias': use_bias,
'padding': 'same',
'name': name + '_groupconv2d',
- 'kernel_regularizer': tf.keras.regularizers.l2(weight_decay),
- 'bias_regularizer': tf.keras.regularizers.l2(weight_decay),
+ 'kernel_regularizer': tf_keras.regularizers.l2(weight_decay),
+ 'bias_regularizer': tf_keras.regularizers.l2(weight_decay),
'filters': conv_filters,
'groups': groups,
'batch_norm_layer': batch_norm if use_batch_norm else None,
@@ -730,12 +729,12 @@ def conv2d_block_as_layers(
activation: Any = None,
depthwise: bool = False,
kernel_initializer: InitializerType = None,
- name: Optional[str] = None) -> List[tf.keras.layers.Layer]:
+ name: Optional[str] = None) -> List[tf_keras.layers.Layer]:
"""A conv2d followed by batch norm and an activation."""
batch_norm = common_modules.get_batch_norm(config.batch_norm)
bn_momentum = config.bn_momentum
bn_epsilon = config.bn_epsilon
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
weight_decay = config.weight_decay
name = name or ''
@@ -747,22 +746,22 @@ def conv2d_block_as_layers(
'use_bias': use_bias,
'padding': 'same',
'name': name + '_conv2d',
- 'kernel_regularizer': tf.keras.regularizers.l2(weight_decay),
- 'bias_regularizer': tf.keras.regularizers.l2(weight_decay),
+ 'kernel_regularizer': tf_keras.regularizers.l2(weight_decay),
+ 'bias_regularizer': tf_keras.regularizers.l2(weight_decay),
}
- sequential_layers: List[tf.keras.layers.Layer] = []
+ sequential_layers: List[tf_keras.layers.Layer] = []
if depthwise:
- conv2d = tf.keras.layers.DepthwiseConv2D
+ conv2d = tf_keras.layers.DepthwiseConv2D
init_kwargs.update({'depthwise_initializer': kernel_initializer})
else:
- conv2d = tf.keras.layers.Conv2D
+ conv2d = tf_keras.layers.Conv2D
init_kwargs.update({
'filters': conv_filters,
'kernel_initializer': kernel_initializer
})
- sequential_layers.append(conv2d(**init_kwargs))
+ sequential_layers.append(conv2d(**init_kwargs)) # pyrefly: ignore[missing-argument]
if use_batch_norm:
bn_axis = 1 if data_format == 'channels_first' else -1
@@ -775,7 +774,7 @@ def conv2d_block_as_layers(
if activation is not None:
sequential_layers.append(
- tf.keras.layers.Activation(activation, name=name + '_activation'))
+ tf_keras.layers.Activation(activation, name=name + '_activation'))
return sequential_layers
@@ -807,7 +806,7 @@ def conv2d_block(inputs: tf.Tensor,
return x
-# Do not inherit from (tf.keras.layers.Layer), will break weights loading.
+# Do not inherit from (tf_keras.layers.Layer), will break weights loading.
class _MbConvBlock:
"""Mobile Inverted Residual Bottleneck composite layer."""
@@ -819,11 +818,11 @@ def __call__(self, inputs: tf.Tensor, training=False):
se = x
for layer in self.squeeze_excitation:
se = layer(se)
- x = tf.keras.layers.multiply([x, se], name=self.name + 'se_excite')
+ x = tf_keras.layers.multiply([x, se], name=self.name + 'se_excite')
for layer in self.project_block:
x = layer(x)
if self.has_skip_add:
- x = tf.keras.layers.add([x, inputs], name=self.name + 'add')
+ x = tf_keras.layers.add([x, inputs], name=self.name + 'add')
return x
def __init__(self,
@@ -840,7 +839,7 @@ def __init__(self,
use_se = config.use_se
activation = tf_utils.get_activation(config.activation)
drop_connect_rate = config.drop_connect_rate
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
use_depthwise = block.conv_type == 'depthwise'
use_groupconv = block.conv_type == 'group'
prefix = prefix or ''
@@ -851,9 +850,9 @@ def __init__(self,
filters = block.input_filters * block.expand_ratio
- self.expand_block: List[tf.keras.layers.Layer] = []
- self.squeeze_excitation: List[tf.keras.layers.Layer] = []
- self.project_block: List[tf.keras.layers.Layer] = []
+ self.expand_block: List[tf_keras.layers.Layer] = []
+ self.squeeze_excitation: List[tf_keras.layers.Layer] = []
+ self.project_block: List[tf_keras.layers.Layer] = []
if block.fused_project:
raise NotImplementedError('Fused projection is not supported.')
@@ -878,7 +877,7 @@ def __init__(self,
kernel_size=block.kernel_size,
strides=block.strides,
activation=activation,
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
name=prefix + 'fused'))
else:
if block.expand_ratio != 1:
@@ -889,7 +888,7 @@ def __init__(self,
config=config,
kernel_size=(1, 1),
activation=activation,
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
name=prefix + 'expand'))
# Main kernel, after the expansion (if applicable, i.e. not fused).
@@ -900,7 +899,7 @@ def __init__(self,
kernel_size=block.kernel_size,
strides=block.strides,
activation=activation,
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
depthwise=True,
name=prefix + 'depthwise'))
elif use_groupconv:
@@ -927,9 +926,9 @@ def __init__(self,
se_shape = (1, 1, filters)
self.squeeze_excitation.append(
- tf.keras.layers.GlobalAveragePooling2D(name=prefix + 'se_squeeze'))
+ tf_keras.layers.GlobalAveragePooling2D(name=prefix + 'se_squeeze'))
self.squeeze_excitation.append(
- tf.keras.layers.Reshape(se_shape, name=prefix + 'se_reshape'))
+ tf_keras.layers.Reshape(se_shape, name=prefix + 'se_reshape'))
self.squeeze_excitation.extend(
conv2d_block_as_layers(
conv_filters=num_reduced_filters,
@@ -937,7 +936,7 @@ def __init__(self,
use_bias=True,
use_batch_norm=False,
activation=activation,
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
name=prefix + 'se_reduce'))
self.squeeze_excitation.extend(
conv2d_block_as_layers(
@@ -946,7 +945,7 @@ def __init__(self,
use_bias=True,
use_batch_norm=False,
activation='sigmoid',
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
name=prefix + 'se_expand'))
# Output phase
@@ -955,13 +954,13 @@ def __init__(self,
conv_filters=block.output_filters,
config=config,
activation=None,
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
name=prefix + 'project'))
# Add identity so that quantization-aware training can insert quantization
# ops correctly.
self.project_block.append(
- tf.keras.layers.Activation('linear', name=prefix + 'id'))
+ tf_keras.layers.Activation('linear', name=prefix + 'id'))
self.has_skip_add = False
if (block.id_skip
@@ -974,7 +973,7 @@ def __init__(self,
# by drop_connect_rate during training. See:
# https://github.com/keras-team/keras/pull/9898#issuecomment-380577612
self.project_block.append(
- tf.keras.layers.Dropout(
+ tf_keras.layers.Dropout(
drop_connect_rate,
noise_shape=(None, 1, 1, 1),
name=prefix + 'drop'))
@@ -998,12 +997,12 @@ def mb_conv_block(inputs: tf.Tensor,
return _MbConvBlock(block, config, prefix)(inputs)
-def mobilenet_edgetpu_v2(image_input: tf.keras.layers.Input,
+def mobilenet_edgetpu_v2(image_input: tf_keras.layers.Input, # pyrefly: ignore[not-a-type]
config: ModelConfig): # pytype: disable=invalid-annotation # typed-keras
"""Creates a MobilenetEdgeTPUV2 graph given the model parameters.
This function is wrapped by the `MobilenetEdgeTPUV2` class to make a
- tf.keras.Model.
+ tf_keras.Model.
Args:
image_input: the input batch of images
@@ -1030,14 +1029,14 @@ def mobilenet_edgetpu_v2(image_input: tf.keras.layers.Input,
num_classes = config.num_classes
input_channels = config.input_channels
rescale_input = config.rescale_input
- data_format = tf.keras.backend.image_data_format()
+ data_format = tf_keras.backend.image_data_format()
dtype = config.dtype
weight_decay = config.weight_decay
x = image_input
if data_format == 'channels_first':
# Happens on GPU/TPU if available.
- x = tf.keras.layers.Permute((3, 1, 2))(x)
+ x = tf_keras.layers.Permute((3, 1, 2))(x)
if rescale_input:
x = common_modules.normalize_images(
x, num_channels=input_channels, dtype=dtype, data_format=data_format)
@@ -1050,7 +1049,7 @@ def mobilenet_edgetpu_v2(image_input: tf.keras.layers.Input,
kernel_size=[stem_kernel_size, stem_kernel_size],
strides=[2, 2],
activation=activation,
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
name='stem')
# Build blocks
@@ -1101,23 +1100,23 @@ def mobilenet_edgetpu_v2(image_input: tf.keras.layers.Input,
conv_filters=round_filters(top_base_filters, config),
config=config,
activation=activation,
- kernel_initializer=conv_kernel_initializer,
+ kernel_initializer=conv_kernel_initializer, # pyrefly: ignore[bad-argument-type]
name='top')
# Build classifier
pool_size = (x.shape.as_list()[1], x.shape.as_list()[2])
- x = tf.keras.layers.AveragePooling2D(pool_size, name='top_pool')(x)
+ x = tf_keras.layers.AveragePooling2D(pool_size, name='top_pool')(x)
if dropout_rate and dropout_rate > 0:
- x = tf.keras.layers.Dropout(dropout_rate, name='top_dropout')(x)
- x = tf.keras.layers.Conv2D(
+ x = tf_keras.layers.Dropout(dropout_rate, name='top_dropout')(x)
+ x = tf_keras.layers.Conv2D(
num_classes,
1,
kernel_initializer=dense_kernel_initializer,
- kernel_regularizer=tf.keras.regularizers.l2(weight_decay),
- bias_regularizer=tf.keras.regularizers.l2(weight_decay),
+ kernel_regularizer=tf_keras.regularizers.l2(weight_decay),
+ bias_regularizer=tf_keras.regularizers.l2(weight_decay),
name='logits')(
x)
- x = tf.keras.layers.Activation('softmax', name='probs')(x)
+ x = tf_keras.layers.Activation('softmax', name='probs')(x)
x = tf.squeeze(x, axis=[1, 2])
return x
diff --git a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks_test.py b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks_test.py
index 1ad600399d1..b2b3a79613f 100644
--- a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks_test.py
+++ b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_blocks_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for mobilenet_edgetpu_v2_model_blocks."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling import custom_layers
from official.projects.edgetpu.vision.modeling import mobilenet_edgetpu_v2_model_blocks
@@ -27,28 +27,28 @@ def setUp(self):
self.model_config = mobilenet_edgetpu_v2_model_blocks.ModelConfig()
def test_model_creatation(self):
- model_input = tf.keras.layers.Input(shape=(224, 224, 1))
+ model_input = tf_keras.layers.Input(shape=(224, 224, 1))
model_output = mobilenet_edgetpu_v2_model_blocks.mobilenet_edgetpu_v2(
image_input=model_input,
config=self.model_config)
- test_model = tf.keras.Model(inputs=model_input, outputs=model_output)
- self.assertIsInstance(test_model, tf.keras.Model)
+ test_model = tf_keras.Model(inputs=model_input, outputs=model_output)
+ self.assertIsInstance(test_model, tf_keras.Model)
self.assertEqual(test_model.input.shape, (None, 224, 224, 1))
self.assertEqual(test_model.output.shape, (None, 1001))
def test_model_with_customized_kernel_initializer(self):
self.model_config.conv_kernel_initializer = 'he_uniform'
self.model_config.dense_kernel_initializer = 'glorot_normal'
- model_input = tf.keras.layers.Input(shape=(224, 224, 1))
+ model_input = tf_keras.layers.Input(shape=(224, 224, 1))
model_output = mobilenet_edgetpu_v2_model_blocks.mobilenet_edgetpu_v2(
image_input=model_input,
config=self.model_config)
- test_model = tf.keras.Model(inputs=model_input, outputs=model_output)
+ test_model = tf_keras.Model(inputs=model_input, outputs=model_output)
conv_layer_stack = []
for layer in test_model.layers:
- if (isinstance(layer, tf.keras.layers.Conv2D) or
- isinstance(layer, tf.keras.layers.DepthwiseConv2D) or
+ if (isinstance(layer, tf_keras.layers.Conv2D) or
+ isinstance(layer, tf_keras.layers.DepthwiseConv2D) or
isinstance(layer, custom_layers.GroupConv2D)):
conv_layer_stack.append(layer)
self.assertGreater(len(conv_layer_stack), 2)
@@ -56,16 +56,16 @@ def test_model_with_customized_kernel_initializer(self):
for layer in conv_layer_stack[:-1]:
if isinstance(layer, custom_layers.GroupConv2D):
self.assertIsInstance(layer.kernel_initializer,
- tf.keras.initializers.GlorotUniform)
- elif isinstance(layer, tf.keras.layers.Conv2D):
+ tf_keras.initializers.GlorotUniform)
+ elif isinstance(layer, tf_keras.layers.Conv2D):
self.assertIsInstance(layer.kernel_initializer,
- tf.keras.initializers.HeUniform)
- elif isinstance(layer, tf.keras.layers.DepthwiseConv2D):
+ tf_keras.initializers.HeUniform)
+ elif isinstance(layer, tf_keras.layers.DepthwiseConv2D):
self.assertIsInstance(layer.depthwise_initializer,
- tf.keras.initializers.HeUniform)
+ tf_keras.initializers.HeUniform)
self.assertIsInstance(conv_layer_stack[-1].kernel_initializer,
- tf.keras.initializers.GlorotNormal)
+ tf_keras.initializers.GlorotNormal)
if __name__ == '__main__':
diff --git a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_test.py b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_test.py
index 7044d7d93e5..140ddc6f1a8 100644
--- a/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_test.py
+++ b/official/projects/edgetpu/vision/modeling/mobilenet_edgetpu_v2_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling import common_modules
from official.projects.edgetpu.vision.modeling import mobilenet_edgetpu_v2_model
@@ -28,7 +28,7 @@ class MobilenetEdgeTPUV2BuildTest(tf.test.TestCase, parameterized.TestCase):
def setUp(self):
super(tf.test.TestCase, self).setUp()
# Ensure no model duplicates
- tf.keras.backend.clear_session()
+ tf_keras.backend.clear_session()
def test_create_mobilenet_edgetpu(self):
model = mobilenet_edgetpu_v2_model.MobilenetEdgeTPUV2()
@@ -54,7 +54,7 @@ def test_model_save_load(self):
first_conv_layer = model.get_layer('stem_conv2d')
kernel_tensor = first_conv_layer.trainable_weights[0].numpy()
model.save('/tmp/test_model')
- loaded_model = tf.keras.models.load_model('/tmp/test_model')
+ loaded_model = tf_keras.models.load_model('/tmp/test_model')
loaded_first_conv_layer = loaded_model.get_layer('stem_conv2d')
loaded_kernel_tensor = loaded_first_conv_layer.trainable_weights[0].numpy()
diff --git a/official/projects/edgetpu/vision/modeling/optimized_multiheadattention_layer.py b/official/projects/edgetpu/vision/modeling/optimized_multiheadattention_layer.py
new file mode 100644
index 00000000000..30a087fead6
--- /dev/null
+++ b/official/projects/edgetpu/vision/modeling/optimized_multiheadattention_layer.py
@@ -0,0 +1,164 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""MultiHeadAttention layer optimized for EdgeTPU.
+
+Compared to tf_keras.layers.MultiHeadAttention, this layer performs query-key
+multiplication instead of key-query multiplication to remove an unnecessary
+transpose.
+"""
+import math
+import string
+from typing import Optional, Tuple
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+_CHR_IDX = string.ascii_lowercase
+
+
+def _build_attention_equation(
+ rank: int, attn_axes: Tuple[int, ...]) -> Tuple[str, str, int]:
+ """Builds einsum equations for the attention computation.
+
+ Query, key, value inputs after projection are expected to have the shape as:
+ `(bs, , , num_heads, channels)`.
+ `bs` and `` are treated as ``.
+
+ The attention operations can be generalized:
+ (1) Query-key dot product:
+ `(, , num_heads, channels), (,
+ , num_heads, channels) -> (,
+ num_heads, , )`
+ (2) Combination:
+ `(, num_heads, , ),
+ (, , num_heads, channels) -> (, , num_heads, channels)`
+
+ Args:
+ rank: Rank of query, key, value tensors.
+ attn_axes: List/tuple of axes, `[-1, rank)`, that attention will be
+ applied to.
+
+ Returns:
+ Einsum equations.
+ """
+ target_notation = _CHR_IDX[:rank]
+ # `batch_dims` includes the head dim.
+ batch_dims = tuple(np.delete(range(rank), attn_axes + (rank - 1,)))
+ letter_offset = rank
+ source_notation = ""
+ for i in range(rank):
+ if i in batch_dims or i == rank - 1:
+ source_notation += target_notation[i]
+ else:
+ source_notation += _CHR_IDX[letter_offset]
+ letter_offset += 1
+
+ product_notation = "".join([target_notation[i] for i in batch_dims] +
+ [target_notation[i] for i in attn_axes] +
+ [source_notation[i] for i in attn_axes])
+ dot_product_equation = "%s,%s->%s" % (
+ target_notation,
+ source_notation,
+ product_notation,
+ )
+ attn_scores_rank = len(product_notation)
+ combine_equation = "%s,%s->%s" % (
+ product_notation,
+ source_notation,
+ target_notation,
+ )
+ return dot_product_equation, combine_equation, attn_scores_rank
+
+
+class OptimizedMultiHeadAttention(tf_keras.layers.MultiHeadAttention):
+ """MultiHeadAttention with query-key multiplication.
+
+ Currently, this layer only works for self-attention but not for
+ cross-attention. TODO(b/243166060).
+ """
+
+ def _build_attention(self, rank: int) -> None:
+ """Builds multi-head dot-product attention computations.
+
+ This function builds attributes necessary for `_compute_attention` to
+ customize attention computation to replace the default dot-product
+ attention.
+
+ Args:
+ rank: the rank of query, key, value tensors.
+ """
+ if self._attention_axes is None:
+ self._attention_axes = tuple(range(1, rank - 2)) # pyrefly: ignore[bad-assignment]
+ else:
+ self._attention_axes = tuple(self._attention_axes)
+ (
+ self._dot_product_equation,
+ self._combine_equation,
+ attn_scores_rank,
+ ) = _build_attention_equation(
+ rank, attn_axes=self._attention_axes)
+ norm_axes = tuple(
+ range(attn_scores_rank - len(self._attention_axes), attn_scores_rank))
+ self._softmax = tf_keras.layers.Softmax(axis=norm_axes)
+ self._dropout_layer = tf_keras.layers.Dropout(rate=self._dropout)
+
+ def _compute_attention(
+ self,
+ query: tf.Tensor,
+ key: tf.Tensor,
+ value: tf.Tensor,
+ attention_mask: Optional[tf.Tensor] = None,
+ training: Optional[bool] = None) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Applies Dot-product attention with query, key, value tensors.
+
+ This function defines the computation inside `call` with projected
+ multi-head Q, K, V inputs. Users can override this function for
+ customized attention implementation.
+
+ Args:
+ query: Projected query `Tensor` of shape `(B, T, N, key_dim)`.
+ key: Projected key `Tensor` of shape `(B, S, N, key_dim)`.
+ value: Projected value `Tensor` of shape `(B, S, N, value_dim)`.
+ attention_mask: a boolean mask of shape `(B, T, S)`, that prevents
+ attention to certain positions. It is generally not needed if the
+ `query` and `value` (and/or `key`) are masked.
+ training: Python boolean indicating whether the layer should behave in
+ training mode (adding dropout) or in inference mode (doing nothing).
+
+ Returns:
+ attention_output: Multi-headed outputs of attention computation.
+ attention_scores: Multi-headed attention weights.
+ """
+ # Note: Applying scalar multiply at the smaller end of einsum improves
+ # XLA performance, but may introduce slight numeric differences in
+ # the Transformer attention head.
+ query = tf.multiply(query, 1.0 / math.sqrt(float(self._key_dim)))
+
+ # Take the dot product between "query" and "key" to get the raw
+ # attention scores.
+ attention_scores = tf.einsum(self._dot_product_equation, query, key)
+
+ attention_scores = self._masked_softmax(attention_scores, attention_mask)
+
+ # This is actually dropping out entire tokens to attend to, which might
+ # seem a bit unusual, but is taken from the original Transformer paper.
+ attention_scores_dropout = self._dropout_layer(
+ attention_scores, training=training)
+
+ # `context_layer` = [B, T, N, H]
+ attention_output = tf.einsum(self._combine_equation,
+ attention_scores_dropout, value)
+ return attention_output, attention_scores
diff --git a/official/projects/edgetpu/vision/modeling/optimized_multiheadattention_layer_test.py b/official/projects/edgetpu/vision/modeling/optimized_multiheadattention_layer_test.py
new file mode 100644
index 00000000000..505ebf0c299
--- /dev/null
+++ b/official/projects/edgetpu/vision/modeling/optimized_multiheadattention_layer_test.py
@@ -0,0 +1,81 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for optimized_multiheadattention_layer."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.edgetpu.vision.modeling import optimized_multiheadattention_layer
+
+_BATCH_SIZE = 32
+_SEQ_LEN = 4
+_EMBEDDING_SIZE = 8
+_NUM_HEADS = 2
+_KEY_DIM = 2
+
+
+class OptimizedMultiheadattentionLayerTest(tf.test.TestCase):
+
+ def test_same_output(self):
+ """Tests that OptimizedMultiHeadAttention returns the expected outputs."""
+
+ input_tensor_1 = tf.random.uniform((_BATCH_SIZE, _SEQ_LEN, _EMBEDDING_SIZE))
+ input_tensor_2 = tf.random.uniform((_BATCH_SIZE, _SEQ_LEN, _EMBEDDING_SIZE))
+
+ # Instantiate layer and call with inputs to build.
+ orig_layer = tf_keras.layers.MultiHeadAttention(
+ num_heads=_NUM_HEADS, key_dim=_KEY_DIM)
+ _ = orig_layer(input_tensor_1, input_tensor_2)
+ opt_layer = optimized_multiheadattention_layer.OptimizedMultiHeadAttention(
+ num_heads=_NUM_HEADS, key_dim=_KEY_DIM)
+ _ = opt_layer(input_tensor_1, input_tensor_2)
+
+ # Set the weights of the two layers to be the same.
+ query_dense_weights = np.random.uniform(
+ size=(_EMBEDDING_SIZE, _NUM_HEADS, _KEY_DIM))
+ query_dense_bias = np.random.uniform(size=(_NUM_HEADS, _KEY_DIM))
+ key_dense_weights = np.random.uniform(
+ size=(_EMBEDDING_SIZE, _NUM_HEADS, _KEY_DIM))
+ key_dense_bias = np.random.uniform(size=(_NUM_HEADS, _KEY_DIM))
+ value_dense_weights = np.random.uniform(
+ size=(_EMBEDDING_SIZE, _NUM_HEADS, _KEY_DIM))
+ value_dense_bias = np.random.uniform(size=(_NUM_HEADS, _KEY_DIM))
+ attention_output_dense_weights = np.random.uniform(
+ size=(_NUM_HEADS, _KEY_DIM, _EMBEDDING_SIZE))
+ attention_output_dense_bias = np.random.uniform(size=(_EMBEDDING_SIZE,))
+
+ orig_layer._query_dense.set_weights([query_dense_weights, query_dense_bias])
+ orig_layer._key_dense.set_weights([key_dense_weights, key_dense_bias])
+ orig_layer._value_dense.set_weights([value_dense_weights, value_dense_bias])
+ orig_layer._output_dense.set_weights(
+ [attention_output_dense_weights, attention_output_dense_bias])
+
+ opt_layer._query_dense.set_weights([query_dense_weights, query_dense_bias])
+ opt_layer._key_dense.set_weights([key_dense_weights, key_dense_bias])
+ opt_layer._value_dense.set_weights([value_dense_weights, value_dense_bias])
+ opt_layer._output_dense.set_weights(
+ [attention_output_dense_weights, attention_output_dense_bias])
+
+ # Calculate two sets of attention outputs and scores and compare.
+ orig_attn_output, orig_attn_score = orig_layer(
+ input_tensor_1, input_tensor_2, return_attention_scores=True)
+ opt_attn_output, opt_attn_score = opt_layer(
+ input_tensor_1, input_tensor_2, return_attention_scores=True)
+ self.assertAllClose(orig_attn_output, opt_attn_output)
+ self.assertAllClose(orig_attn_score, opt_attn_score)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/edgetpu/vision/serving/__init__.py b/official/projects/edgetpu/vision/serving/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/serving/__init__.py
+++ b/official/projects/edgetpu/vision/serving/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/serving/export_tflite.py b/official/projects/edgetpu/vision/serving/export_tflite.py
index 725edd793a3..2db35fa2dd6 100644
--- a/official/projects/edgetpu/vision/serving/export_tflite.py
+++ b/official/projects/edgetpu/vision/serving/export_tflite.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -35,7 +35,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.modeling import common_modules
from official.projects.edgetpu.vision.serving import export_util
@@ -53,13 +53,15 @@
flags.DEFINE_bool(
'export_keras_model', False,
'Export SavedModel format: if False, export TF SavedModel with'
- 'tf.saved_model API; if True, export Keras SavedModel with tf.keras.Model'
+ 'tf.saved_model API; if True, export Keras SavedModel with tf_keras.Model'
'API.')
flags.DEFINE_string('output_dir', None, 'Directory to output exported files.')
flags.DEFINE_integer(
'image_size', 224,
'Size of the input image. Ideally should be the same as the image_size used '
'in training config.')
+flags.DEFINE_bool(
+ 'fix_batch_size', True, 'Whether to export model with fixed batch size.')
flags.DEFINE_string(
'output_layer', None,
'Layer name to take the output from. Can be used to take the output from '
@@ -144,13 +146,15 @@ def run_export():
'chose an output layer.', export_config.output_layer)
return
output_layer = model.get_layer(export_config.output_layer)
- model = tf.keras.Model(model.input, output_layer.output)
+ model = tf_keras.Model(model.input, output_layer.output)
+
+ batch_size = 1 if FLAGS.fix_batch_size else None
- model_input = tf.keras.Input(
+ model_input = tf_keras.Input(
shape=(export_config.image_size, export_config.image_size, 3),
- batch_size=1)
+ batch_size=batch_size)
model_output = export_util.finalize_serving(model(model_input), export_config)
- model_for_inference = tf.keras.Model(model_input, model_output)
+ model_for_inference = tf_keras.Model(model_input, model_output)
# Convert to tflite. Quantize if quantization parameters are specified.
converter = tf.lite.TFLiteConverter.from_keras_model(model_for_inference)
diff --git a/official/projects/edgetpu/vision/serving/export_tflite_test.py b/official/projects/edgetpu/vision/serving/export_tflite_test.py
index 6a0ae90629c..ef8e26e3f65 100644
--- a/official/projects/edgetpu/vision/serving/export_tflite_test.py
+++ b/official/projects/edgetpu/vision/serving/export_tflite_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import exp_factory
from official.core import task_factory
@@ -49,10 +49,10 @@ def _build_experiment_model(experiment_type):
def _build_model(config):
model = _build_experiment_model(config.model_name)
- model_input = tf.keras.Input(
+ model_input = tf_keras.Input(
shape=(config.image_size, config.image_size, 3), batch_size=1)
model_output = export_util.finalize_serving(model(model_input), config)
- model_for_inference = tf.keras.Model(model_input, model_output)
+ model_for_inference = tf_keras.Model(model_input, model_output)
return model_for_inference
diff --git a/official/projects/edgetpu/vision/serving/export_util.py b/official/projects/edgetpu/vision/serving/export_util.py
index 71b688272d8..472acbcd4d2 100644
--- a/official/projects/edgetpu/vision/serving/export_util.py
+++ b/official/projects/edgetpu/vision/serving/export_util.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import dataclasses
from typing import List, Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.core import exp_factory
@@ -93,7 +93,9 @@ class ExportConfig(base_config.Config):
For example: --finalize_method=resize128,argmax,resize512,squeeze will do
resize bilinear to 128x128, then argmax then resize nn to 512x512
"""
- quantization_config: QuantizationConfig = QuantizationConfig()
+ quantization_config: QuantizationConfig = dataclasses.field(
+ default_factory=QuantizationConfig
+ )
model_name: Optional[str] = None
output_layer: Optional[str] = None
ckpt_path: Optional[str] = None
@@ -107,6 +109,12 @@ class ExportConfig(base_config.Config):
def finalize_serving(model_output, export_config):
"""Adds extra layers based on the provided configuration."""
+ if isinstance(model_output, dict):
+ return {
+ key: finalize_serving(model_output[key], export_config)
+ for key in model_output
+ }
+
finalize_method = export_config.finalize_method
output_layer = model_output
if not finalize_method or finalize_method[0] == 'none':
@@ -183,8 +191,7 @@ def representative_dataset_gen(export_config):
"""Gets a python generator of numpy arrays for the given dataset."""
quantization_config = export_config.quantization_config
dataset = tfds.builder(
- quantization_config.dataset_name,
- data_dir=quantization_config.dataset_dir)
+ quantization_config.dataset_name, try_gcs=True)
dataset.download_and_prepare()
data = dataset.as_dataset()[quantization_config.dataset_split]
iterator = data.as_numpy_iterator()
@@ -201,7 +208,8 @@ def configure_tflite_converter(export_config, converter):
"""Common code for picking up quantization parameters."""
quantization_config = export_config.quantization_config
if quantization_config.quantize:
- if quantization_config.dataset_dir is None:
+ if (quantization_config.dataset_dir is
+ None) and (quantization_config.dataset_name is None):
raise ValueError(
'Must provide a representative dataset when quantizing the model.')
converter.optimizations = [tf.lite.Optimize.DEFAULT]
diff --git a/official/projects/edgetpu/vision/serving/inference_visualization_tool.ipynb b/official/projects/edgetpu/vision/serving/inference_visualization_tool.ipynb
index e163df3052d..f0795773c03 100644
--- a/official/projects/edgetpu/vision/serving/inference_visualization_tool.ipynb
+++ b/official/projects/edgetpu/vision/serving/inference_visualization_tool.ipynb
@@ -57,7 +57,7 @@
" sandbox_path = web_path.split('/')[-1]\n",
" !rm -f {sandbox_path}\n",
" if web_path[:2] == \"gs\":\n",
- " !gsutil cp {web_path} {sandbox_path}\n",
+ " !gcloud storage cp {web_path} {sandbox_path}\n",
" else:\n",
" !wget -v {web_path} --no-check-certificate\n",
" return sandbox_path\n"
@@ -122,7 +122,7 @@
},
"source": [
"MODEL_HOME='gs://tf_model_garden/models/edgetpu/checkpoint_and_tflite/vision/segmentation-edgetpu/tflite/default_argmax'\n",
- "!gsutil ls {MODEL_HOME}"
+ "!gcloud storage ls {MODEL_HOME}"
],
"execution_count": null,
"outputs": []
diff --git a/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator.py b/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator.py
index c4afb000b01..4dafc1831f6 100644
--- a/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator.py
+++ b/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,11 @@
from typing import Tuple
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
+# pylint: disable=g-direct-tensorflow-import
+from tensorflow.lite.python import interpreter as tfl_interpreter
+# pylint: enable=g-direct-tensorflow-import
@dataclasses.dataclass
@@ -58,8 +62,9 @@ def evaluate_single_image(self, eval_input: EvaluationInput) -> bool:
Returns:
Whether the estimation is correct.
"""
- interpreter = tf.lite.Interpreter(
- model_content=self._model_content, num_threads=1)
+ interpreter = tfl_interpreter.Interpreter(
+ model_content=self._model_content, num_threads=1
+ )
interpreter.allocate_tensors()
# Get input and output tensors and quantization details.
input_details = interpreter.get_input_details()
diff --git a/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_run.py b/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_run.py
index f74f90ac2fb..52296329180 100644
--- a/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_run.py
+++ b/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_run.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from typing import Sequence
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import exp_factory
from official.projects.edgetpu.vision.serving import tflite_imagenet_evaluator
diff --git a/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_test.py b/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_test.py
index 3fcaffa4537..9fbab81bcc1 100644
--- a/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_test.py
+++ b/official/projects/edgetpu/vision/serving/tflite_imagenet_evaluator_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for tflite_imagenet_evaluator."""
from unittest import mock
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.edgetpu.vision.serving import tflite_imagenet_evaluator
diff --git a/official/projects/edgetpu/vision/tasks/__init__.py b/official/projects/edgetpu/vision/tasks/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/edgetpu/vision/tasks/__init__.py
+++ b/official/projects/edgetpu/vision/tasks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/edgetpu/vision/tasks/image_classification.py b/official/projects/edgetpu/vision/tasks/image_classification.py
index b3e1a22e603..c7347539a08 100644
--- a/official/projects/edgetpu/vision/tasks/image_classification.py
+++ b/official/projects/edgetpu/vision/tasks/image_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from typing import Any, List, Mapping, Optional, Tuple
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import dataset_fn
from official.core import base_task
@@ -45,7 +45,7 @@ def _copy_recursively(src: str, dst: str) -> None:
overwrite=True)
-def get_models() -> Mapping[str, tf.keras.Model]:
+def get_models() -> Mapping[str, tf_keras.Model]:
"""Returns the mapping from model type name to Keras model."""
model_mapping = {}
@@ -63,7 +63,7 @@ def add_models(name: str, constructor: Any):
return model_mapping
-def load_searched_model(saved_model_path: str) -> tf.keras.Model:
+def load_searched_model(saved_model_path: str) -> tf_keras.Model:
"""Loads saved model from file.
Excepting loading MobileNet-EdgeTPU-V1/V2 models, we can also load searched
@@ -83,7 +83,7 @@ def load_searched_model(saved_model_path: str) -> tf.keras.Model:
raise ValueError('Saved model path is invalid.')
load_options = tf.saved_model.LoadOptions(
experimental_io_device='/job:localhost')
- model = tf.keras.models.load_model(load_path, options=load_options)
+ model = tf_keras.models.load_model(load_path, options=load_options)
return model
@@ -115,7 +115,7 @@ def build_model(self):
return model
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loads pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -190,7 +190,7 @@ def build_losses(self,
Args:
labels: Input groundtruth labels.
model_outputs: Output logits of the classifier.
- aux_losses: The auxiliarly loss tensors, i.e. `losses` in tf.keras.Model.
+ aux_losses: The auxiliarly loss tensors, i.e. `losses` in tf_keras.Model.
Returns:
The total loss tensor.
@@ -200,13 +200,13 @@ def build_losses(self,
if not is_multilabel:
if losses_config.one_hot:
- total_loss = tf.keras.losses.categorical_crossentropy(
+ total_loss = tf_keras.losses.categorical_crossentropy(
labels,
model_outputs,
from_logits=False,
label_smoothing=losses_config.label_smoothing)
else:
- total_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ total_loss = tf_keras.losses.sparse_categorical_crossentropy(
labels, model_outputs, from_logits=True)
else:
# Multi-label weighted binary cross entropy loss.
@@ -221,20 +221,20 @@ def build_losses(self,
return total_loss
def build_metrics(self,
- training: bool = True) -> List[tf.keras.metrics.Metric]:
+ training: bool = True) -> List[tf_keras.metrics.Metric]:
"""Gets streaming metrics for training/validation."""
is_multilabel = self.task_config.train_data.is_multilabel
if not is_multilabel:
k = self.task_config.evaluation.top_k
if self.task_config.losses.one_hot:
metrics = [
- tf.keras.metrics.CategoricalAccuracy(name='accuracy'),
- tf.keras.metrics.TopKCategoricalAccuracy(
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.TopKCategoricalAccuracy(
k=k, name='top_{}_accuracy'.format(k))]
else:
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
- tf.keras.metrics.SparseTopKCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.SparseTopKCategoricalAccuracy(
k=k, name='top_{}_accuracy'.format(k))]
else:
metrics = []
@@ -243,30 +243,30 @@ def build_metrics(self,
# TODO(arashwan): Investigate adding following metric to train.
if not training:
metrics = [
- tf.keras.metrics.AUC(
+ tf_keras.metrics.AUC(
name='globalPR-AUC',
curve='PR',
multi_label=False,
from_logits=True),
- tf.keras.metrics.AUC(
+ tf_keras.metrics.AUC(
name='meanPR-AUC',
curve='PR',
multi_label=True,
num_labels=self.task_config.model.num_classes,
from_logits=True),
]
- return metrics
+ return metrics # pyrefly: ignore[bad-return]
def train_step(self,
inputs: Tuple[Any, Any],
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
metrics: Optional[List[Any]] = None):
"""Does forward and backward.
Args:
- inputs: A tuple of of input tensors of (features, labels).
- model: A tf.keras.Model instance.
+ inputs: A tuple of input tensors of (features, labels).
+ model: A tf_keras.Model instance.
optimizer: The optimizer for this training step.
metrics: A nested structure of metrics objects.
@@ -292,7 +292,7 @@ def train_step(self,
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
if isinstance(
- optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
@@ -300,7 +300,7 @@ def train_step(self,
# Scales back gradient before apply_gradients when LossScaleOptimizer is
# used.
if isinstance(
- optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
@@ -314,13 +314,13 @@ def train_step(self,
def validation_step(self,
inputs: Tuple[Any, Any],
- model: tf.keras.Model,
+ model: tf_keras.Model,
metrics: Optional[List[Any]] = None):
- """Runs validatation step.
+ """Runs validation step.
Args:
- inputs: A tuple of of input tensors of (features, labels).
- model: A tf.keras.Model instance.
+ inputs: A tuple of input tensors of (features, labels).
+ model: A tf_keras.Model instance.
metrics: A nested structure of metrics objects.
Returns:
@@ -344,6 +344,6 @@ def validation_step(self,
logs.update({m.name: m.result() for m in model.metrics})
return logs
- def inference_step(self, inputs: tf.Tensor, model: tf.keras.Model):
+ def inference_step(self, inputs: tf.Tensor, model: tf_keras.Model):
"""Performs the forward step."""
return model(inputs, training=False)
diff --git a/official/projects/edgetpu/vision/tasks/image_classification_test.py b/official/projects/edgetpu/vision/tasks/image_classification_test.py
index be250d9d405..a9626f454c3 100644
--- a/official/projects/edgetpu/vision/tasks/image_classification_test.py
+++ b/official/projects/edgetpu/vision/tasks/image_classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import exp_factory
from official.modeling import optimization
diff --git a/official/projects/edgetpu/vision/tasks/semantic_segmentation.py b/official/projects/edgetpu/vision/tasks/semantic_segmentation.py
index d5cac8120fa..cdb62744dc7 100644
--- a/official/projects/edgetpu/vision/tasks/semantic_segmentation.py
+++ b/official/projects/edgetpu/vision/tasks/semantic_segmentation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Any, Mapping, Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import dataset_fn
from official.core import config_definitions as cfg
@@ -99,7 +99,7 @@ def build_inputs(self,
return dataset
-class AutosegEdgeTPU(tf.keras.Model):
+class AutosegEdgeTPU(tf_keras.Model):
"""Segmentation keras network without pre/post-processing."""
def __init__(self,
@@ -208,7 +208,7 @@ def __init__(self,
fullres_output=fullres_output,
num_classes=num_classes)
- def call(self, inputs, training):
+ def call(self, inputs, training): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
# call backbone network.
all_feats = self.backbone(inputs, training=training)
if self.use_original_backbone_features:
@@ -232,7 +232,7 @@ def call(self, inputs, training):
return class_outputs
-def get_models() -> Mapping[str, tf.keras.Model]:
+def get_models() -> Mapping[str, tf_keras.Model]:
"""Returns the mapping from model type name to Keras model."""
model_mapping = {}
diff --git a/official/projects/edgetpu/vision/tasks/semantic_segmentation_test.py b/official/projects/edgetpu/vision/tasks/semantic_segmentation_test.py
index d12eb8dcdcd..89b58ee3f03 100644
--- a/official/projects/edgetpu/vision/tasks/semantic_segmentation_test.py
+++ b/official/projects/edgetpu/vision/tasks/semantic_segmentation_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import exp_factory
@@ -90,11 +90,14 @@ def test_task(self, config_name):
class AutosegEdgeTPUTaskTest(tf.test.TestCase, parameterized.TestCase):
- @parameterized.parameters(('autoseg_edgetpu_xs',))
+ @parameterized.parameters(
+ ('autoseg_edgetpu_xs',), ('autoseg_edgetpu_s',), ('autoseg_edgetpu_m',)
+ )
def test_task(self, config_name):
config_to_backbone_mapping = {
'autoseg_edgetpu_xs': 'autoseg_edgetpu_backbone_xs',
- 'autoseg_edgetpu_s': 'autoseg_edgetpu_backone_s'
+ 'autoseg_edgetpu_s': 'autoseg_edgetpu_backbone_s',
+ 'autoseg_edgetpu_m': 'autoseg_edgetpu_backbone_m',
}
config = autoseg_cfg.autoseg_edgetpu_experiment_config(
config_to_backbone_mapping[config_name], init_backbone=False)
diff --git a/official/projects/edgetpu/vision/train.py b/official/projects/edgetpu/vision/train.py
index d08da93810d..913ae12436c 100644
--- a/official/projects/edgetpu/vision/train.py
+++ b/official/projects/edgetpu/vision/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/fffner/README.md b/official/projects/fffner/README.md
new file mode 100644
index 00000000000..6cb88581a43
--- /dev/null
+++ b/official/projects/fffner/README.md
@@ -0,0 +1,94 @@
+# Formulating Few-shot Fine-tuning Towards Language Model Pre-training: A Pilot Study on Named Entity Recognition
+
+**DISCLAIMER**: This implementation is still under development.
+
+[](https://arxiv.org/abs/2205.11799)
+
+This repository is the official implementation of the following
+paper.
+
+* [Formulating Few-shot Fine-tuning Towards Language Model Pre-training: A Pilot Study on Named Entity Recognition](https://arxiv.org/abs/2205.11799)
+
+## Description
+
+FFF-NER is a training task for effective Few-shot Named Entity Recognition.
+You can also refer to the author's [GitHub repository](https://github.com/ZihanWangKi/fffner).
+
+## Maintainers
+
+* Zihan Wang ([zihanwangki](https://github.com/zihanwangki))
+
+## Requirements
+
+[](https://badge.fury.io/py/tensorflow)
+[](https://badge.fury.io/py/tf-models-official)
+
+## Training & Evaluation
+
+It can run on Google Cloud Platform using Cloud TPU.
+[Here](https://cloud.google.com/tpu/docs/how-to) is the instruction of using
+Cloud TPU.
+
+### Setup
+You will need to first convert a pre-trained language model to the encoder
+format we are using. The following command by default converts a base size
+bert uncased model.
+```shell
+python3 utils/convert_checkpoint_tensorflow.py
+```
+Then, you will need to convert the dataset into a tf_record for training.
+`utils/create_data.py` contains the script to do so. Example dataset and
+dataset format can be found in the
+[official repo](https://github.com/ZihanWangKi/fffner/tree/main/dataset)
+Suppose the dataset is stored as
+`/data/fffner_datasets/conll2003/few_shot_5_0.words` and
+`/data/fffner_datasets/conll2003/few_shot_5_0.ner`,
+where `/data/fffner_datasets/conll2003/` also contains the dataset
+configuration and testing data,
+then,
+```
+export PATH_TO_DATA_FOLDER=/data/fffner_datasets/
+export DATASET_NAME=conll2003
+export TRAINING_FOLD=few_shot_5_0
+```
+and
+```shell
+python3 utils/create_data.py $PATH_TO_DATA_FOLDER $DATASET_NAME $TRAINING_FOLD
+```
+creates the training fold.
+
+### Training
+
+```shell
+PATH_TO_TRAINING_RECORD=conll2003_few_shot_5_0.tf_record # path to the training record
+PATH_TO_TESTING_RECORD=conll2003_test.tf_record # path to the evaluation record
+TPU_NAME="" # The name assigned while creating a Cloud TPU
+MODEL_DIR=/tmp/conll2003_ew_shot_5_0 # directory to store the experiment
+# Now launch the experiment.
+python3 -m official.projects.mosaic.train \
+ --experiment=fffner/ner \
+ --config_file=experiments/base_conll2003.yaml \
+ --params_override="task.train_data.input_path=${PATH_TO_TRAINING_RECORD},task.validation_data.input_path=${PATH_TO_TESTING_RECORD},runtime.distribution_strategy=tpu"
+ --mode=train_and_eval \
+ --tpu=$TPU_NAME \
+ --model_dir=$MODEL_DIR
+```
+## License
+
+[](https://opensource.org/licenses/Apache-2.0)
+
+This project is licensed under the terms of the **Apache License 2.0**.
+
+## Citation
+
+If you want to cite this repository in your work, please consider citing the
+paper.
+
+```
+@article{wang2022formulating,
+ title={Formulating Few-shot Fine-tuning Towards Language Model Pre-training: A Pilot Study on Named Entity Recognition},
+ author={Wang, Zihan and Zhao, Kewen and Wang, Zilong and Shang, Jingbo},
+ journal={arXiv preprint arXiv:2205.11799},
+ year={2022}
+}
+```
diff --git a/official/projects/fffner/experiments/base_conll2003.yaml b/official/projects/fffner/experiments/base_conll2003.yaml
new file mode 100644
index 00000000000..60de3eb09f6
--- /dev/null
+++ b/official/projects/fffner/experiments/base_conll2003.yaml
@@ -0,0 +1,46 @@
+task:
+ init_checkpoint: 'tf-bert-uncased'
+ model:
+ num_classes_is_entity: 2
+ num_classes_entity_type: 4
+ encoder:
+ type: any
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: true
+ seq_length: 128
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 1024
+ input_path: TODO
+ is_training: false
+ seq_length: 128
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 2189
+ end_learning_rate: 0.0
+ initial_learning_rate: 2.0e-05
+ power: 1.0
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 0
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ # 2335 * 30 / 32
+ train_steps: 2189
+ validation_interval: 500
+ best_checkpoint_export_subdir: "best_ckpts"
+ best_checkpoint_eval_metric: "overall_f1"
+ best_checkpoint_metric_comp: "higher"
diff --git a/official/projects/fffner/experiments/base_restaurants.yaml b/official/projects/fffner/experiments/base_restaurants.yaml
new file mode 100644
index 00000000000..1617d1d25c8
--- /dev/null
+++ b/official/projects/fffner/experiments/base_restaurants.yaml
@@ -0,0 +1,46 @@
+task:
+ init_checkpoint: 'tf-bert-uncased'
+ model:
+ num_classes_is_entity: 2
+ num_classes_entity_type: 8
+ encoder:
+ type: any
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: true
+ seq_length: 128
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 1024
+ input_path: TODO
+ is_training: false
+ seq_length: 128
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 2518
+ end_learning_rate: 0.0
+ initial_learning_rate: 2.0e-05
+ power: 1.0
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 0
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ # 2686 * 30 / 32
+ train_steps: 2518
+ validation_interval: 500
+ best_checkpoint_export_subdir: "best_ckpts"
+ best_checkpoint_eval_metric: "overall_f1"
+ best_checkpoint_metric_comp: "higher"
diff --git a/official/projects/fffner/fffner.py b/official/projects/fffner/fffner.py
new file mode 100644
index 00000000000..524f5052622
--- /dev/null
+++ b/official/projects/fffner/fffner.py
@@ -0,0 +1,45 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""The encoder used for FFFNER."""
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.modeling.hyperparams import base_config
+from official.nlp.configs import encoders
+from official.projects.fffner import fffner_encoder
+
+FFFNerEncoderConfig = encoders.BertEncoderConfig
+
+
+@base_config.bind(FFFNerEncoderConfig)
+def get_encoder(encoder_cfg: FFFNerEncoderConfig):
+ """Gets the FFNerEncoder from the configurations."""
+ encoder = fffner_encoder.FFFNerEncoder(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ num_attention_heads=encoder_cfg.num_attention_heads,
+ inner_dim=encoder_cfg.intermediate_size,
+ inner_activation=tf_utils.get_activation(encoder_cfg.hidden_activation),
+ output_dropout=encoder_cfg.dropout_rate,
+ attention_dropout=encoder_cfg.attention_dropout_rate,
+ max_sequence_length=encoder_cfg.max_position_embeddings,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ output_range=encoder_cfg.output_range,
+ embedding_width=encoder_cfg.embedding_size,
+ norm_first=encoder_cfg.norm_first)
+ return encoder
diff --git a/official/projects/fffner/fffner_classifier.py b/official/projects/fffner/fffner_classifier.py
new file mode 100644
index 00000000000..1a0134a60af
--- /dev/null
+++ b/official/projects/fffner/fffner_classifier.py
@@ -0,0 +1,145 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""FFF-NER special token classifier."""
+# pylint: disable=g-classes-have-attributes
+import collections
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class FFFNerClassifier(tf_keras.Model):
+ """Classifier model based on a BERT-style transformer-based encoder.
+
+ This is an implementation of the network structure surrounding a transformer
+ encoder as described in "BERT: Pre-training of Deep Bidirectional Transformers
+ for Language Understanding" (https://arxiv.org/abs/1810.04805).
+
+ The BertClassifier allows a user to pass in a transformer stack, and
+ instantiates a classification network based on the passed `num_classes`
+ argument. If `num_classes` is set to 1, a regression network is instantiated.
+
+ *Note* that the model is constructed by
+ [Keras Functional API](https://keras.io/guides/functional_api/).
+
+ Args:
+ network: A transformer network. This network should output a sequence output
+ and a classification output. Furthermore, it should expose its embedding
+ table via a "get_embedding_table" method.
+ num_classes: Number of classes to predict from the classification network.
+ initializer: The initializer (if any) to use in the classification networks.
+ Defaults to a Glorot uniform initializer.
+ dropout_rate: The dropout probability of the cls head.
+ use_encoder_pooler: Whether to use the pooler layer pre-defined inside the
+ encoder.
+ head_name_is_entity: Name of the classification head.
+ head_name_entity_type: Name of the classification head.
+ """
+
+ def __init__(self,
+ network,
+ num_classes_is_entity,
+ num_classes_entity_type,
+ initializer='glorot_uniform',
+ dropout_rate=0.1,
+ use_encoder_pooler=True,
+ head_name_is_entity='fffner_prediction_is_entity',
+ head_name_entity_type='fffner_prediction_entity_type',
+ cls_head=None,
+ **kwargs):
+ self.num_classes_is_entity = num_classes_is_entity
+ self.num_classes_entity_type = num_classes_entity_type
+ self.head_name_is_entity = head_name_is_entity
+ self.head_name_entity_type = head_name_entity_type
+ self.initializer = initializer
+ self.use_encoder_pooler = use_encoder_pooler
+ assert use_encoder_pooler, ('Customized pooling & classification function '
+ 'is used')
+
+ # We want to use the inputs of the passed network as the inputs to this
+ # Model. To do this, we need to keep a handle to the network inputs for use
+ # when we construct the Model object at the end of init.
+ inputs = network.inputs
+
+ outputs = network(inputs)
+ if isinstance(outputs, list):
+ cls_inputs = outputs[1]
+ else:
+ cls_inputs = outputs['pooled_output']
+ cls_inputs = tf_keras.layers.Dropout(rate=dropout_rate)(cls_inputs)
+
+ classifier_is_entity = layers.ClassificationHead(
+ inner_dim=0 if use_encoder_pooler else cls_inputs.shape[-1],
+ num_classes=num_classes_is_entity,
+ initializer=initializer,
+ dropout_rate=dropout_rate,
+ name=head_name_is_entity)
+ classifier_entity_type = layers.ClassificationHead(
+ inner_dim=0 if use_encoder_pooler else cls_inputs.shape[-1],
+ num_classes=num_classes_entity_type,
+ initializer=initializer,
+ dropout_rate=dropout_rate,
+ name=head_name_entity_type)
+
+ predictions_is_entity = classifier_is_entity(cls_inputs[:, 0, :])
+ predictions_entity_type = classifier_entity_type(cls_inputs[:, 1, :])
+
+ super().__init__(
+ inputs=inputs,
+ outputs=[predictions_is_entity, predictions_entity_type],
+ **kwargs)
+ self._network = network
+ self._cls_head = cls_head
+
+ config_dict = self._make_config_dict()
+ # We are storing the config dict as a namedtuple here to ensure checkpoint
+ # compatibility with an earlier version of this model which did not track
+ # the config dict attribute. TF does not track immutable attrs which
+ # do not contain Trackables, so by creating a config namedtuple instead of
+ # a dict we avoid tracking it.
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
+ self._config = config_cls(**config_dict)
+ self.classifier_is_entity = classifier_is_entity
+ self.classifier_entity_type = classifier_entity_type
+
+ @property
+ def checkpoint_items(self):
+ items = dict(encoder=self._network)
+ if hasattr(self.classifier_is_entity, 'checkpoint_items'):
+ for key, item in self.classifier_is_entity.checkpoint_items.items():
+ items['.'.join([self.classifier_is_entity.name, key])] = item
+ if hasattr(self.classifier_entity_type, 'checkpoint_items'):
+ for key, item in self.classifier_entity_type.checkpoint_items.items():
+ items['.'.join([self.classifier_entity_type.name, key])] = item
+ return items
+
+ def get_config(self):
+ return dict(self._config._asdict())
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+ def _make_config_dict(self):
+ return {
+ 'network': self._network,
+ 'num_classes_is_entity': self.num_classes_is_entity,
+ 'num_classes_entity_type': self.num_classes_entity_type,
+ 'head_name_is_entity': self.head_name_is_entity,
+ 'head_name_entity_type': self.head_name_entity_type,
+ 'initializer': self.initializer,
+ 'use_encoder_pooler': self.use_encoder_pooler,
+ }
diff --git a/official/projects/fffner/fffner_dataloader.py b/official/projects/fffner/fffner_dataloader.py
new file mode 100644
index 00000000000..eff2886271e
--- /dev/null
+++ b/official/projects/fffner/fffner_dataloader.py
@@ -0,0 +1,129 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Loads dataset for the FFFNER task."""
+import dataclasses
+from typing import Mapping, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import config_definitions as cfg
+from official.core import input_reader
+from official.nlp.data import data_loader
+from official.nlp.data import data_loader_factory
+
+LABEL_TYPES_MAP = {'int': tf.int64, 'float': tf.float32}
+
+
+@dataclasses.dataclass
+class FFFNerDataConfig(cfg.DataConfig):
+ """Data config for sentence prediction task (tasks/sentence_prediction)."""
+ input_path: str = ''
+ global_batch_size: int = 32
+ is_training: bool = True
+ seq_length: int = 128
+ label_type: str = 'int'
+ # Whether to include the example id number.
+ include_example_id: bool = False
+ label_field_is_entity: str = 'is_entity_label'
+ label_field_entity_type: str = 'entity_type_label'
+ # Maps the key in TfExample to feature name.
+ # E.g 'label_ids' to 'next_sentence_labels'
+ label_name: Optional[Tuple[str, str]] = None
+ # Either tfrecord, sstable, or recordio.
+ file_type: str = 'tfrecord'
+
+
+@data_loader_factory.register_data_loader_cls(FFFNerDataConfig)
+class FFFNerDataLoader(data_loader.DataLoader):
+ """A class to load dataset for sentence prediction (classification) task."""
+
+ def __init__(self, params):
+ self._params = params
+ self._seq_length = params.seq_length
+ self._include_example_id = params.include_example_id
+ self._label_field_is_entity = params.label_field_is_entity
+ self._label_field_entity_type = params.label_field_entity_type
+ if params.label_name:
+ self._label_name_mapping = dict(
+ [params.label_name_is_entity, params.label_name_entity_type])
+ else:
+ self._label_name_mapping = dict()
+
+ def name_to_features_spec(self):
+ """Defines features to decode. Subclass may override to append features."""
+ label_type = LABEL_TYPES_MAP[self._params.label_type]
+ name_to_features = {
+ 'input_ids': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'input_mask': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'segment_ids': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'is_entity_token_pos': tf.io.FixedLenFeature([1], tf.int64),
+ 'entity_type_token_pos': tf.io.FixedLenFeature([1], tf.int64),
+ self._label_field_is_entity: tf.io.FixedLenFeature([], label_type),
+ self._label_field_entity_type: tf.io.FixedLenFeature([], label_type),
+ 'sentence_id': tf.io.FixedLenFeature([1], tf.int64),
+ 'span_start': tf.io.FixedLenFeature([1], tf.int64),
+ 'span_end': tf.io.FixedLenFeature([1], tf.int64),
+ }
+ if self._include_example_id:
+ name_to_features['example_id'] = tf.io.FixedLenFeature([], tf.int64)
+
+ return name_to_features
+
+ def _decode(self, record: tf.Tensor):
+ """Decodes a serialized tf.Example."""
+ example = tf.io.parse_single_example(record, self.name_to_features_spec())
+
+ # tf.Example only supports tf.int64, but the TPU only supports tf.int32.
+ # So cast all int64 to int32.
+ for name in example:
+ t = example[name]
+ if t.dtype == tf.int64:
+ t = tf.cast(t, tf.int32)
+ example[name] = t
+
+ return example
+
+ def _parse(self, record: Mapping[str, tf.Tensor]):
+ """Parses raw tensors into a dict of tensors to be consumed by the model."""
+ key_mapping = {
+ 'input_ids': 'input_word_ids',
+ 'input_mask': 'input_mask',
+ 'segment_ids': 'input_type_ids',
+ 'is_entity_token_pos': 'is_entity_token_pos',
+ 'entity_type_token_pos': 'entity_type_token_pos',
+ 'is_entity_label': 'is_entity_label',
+ 'entity_type_label': 'entity_type_label',
+ 'sentence_id': 'sentence_id',
+ 'span_start': 'span_start',
+ 'span_end': 'span_end',
+ }
+ ret = {}
+ for record_key in record:
+ if record_key in key_mapping:
+ ret[key_mapping[record_key]] = record[record_key]
+ else:
+ ret[record_key] = record[record_key]
+
+ return ret
+
+ def load(self, input_context: Optional[tf.distribute.InputContext] = None):
+ """Returns a tf.dataset.Dataset."""
+ reader = input_reader.InputReader(
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
+ params=self._params,
+ decoder_fn=self._decode,
+ parser_fn=self._parse)
+ return reader.read(input_context)
diff --git a/official/projects/fffner/fffner_encoder.py b/official/projects/fffner/fffner_encoder.py
new file mode 100644
index 00000000000..f92045c8244
--- /dev/null
+++ b/official/projects/fffner/fffner_encoder.py
@@ -0,0 +1,346 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Transformer-based encoder network for FFFNER."""
+# pylint: disable=g-classes-have-attributes
+
+from typing import Any, Callable, Optional, Union
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.nlp.modeling import layers
+
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+_Activation = Union[str, Callable[..., Any]]
+
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
+
+
+class FFFNerEncoder(tf_keras.layers.Layer):
+ """Transformer-based encoder network for FFFNER.
+
+ The main difference is that it takes in additional positional arguments and
+ returns last layer representations at those positions.
+ Args:
+ vocab_size: The size of the token vocabulary.
+ hidden_size: The size of the transformer hidden layers.
+ num_layers: The number of transformer layers.
+ num_attention_heads: The number of attention heads for each transformer. The
+ hidden size must be divisible by the number of attention heads.
+ max_sequence_length: The maximum sequence length that this encoder can
+ consume. This determines the variable shape for positional embeddings.
+ type_vocab_size: The number of types that the 'type_ids' input can take.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ output_dropout: Dropout probability for the post-attention and output
+ dropout.
+ attention_dropout: The dropout rate to use for the attention layers within
+ the transformer layers.
+ initializer: The initialzer to use for all weights in this encoder.
+ output_range: The sequence output range, [0, output_range), by slicing the
+ target sequence of the last transformer layer. `None` means the entire
+ target sequence will attend to the source sequence, which yields the full
+ output.
+ embedding_width: The width of the word embeddings. If the embedding width is
+ not equal to hidden size, embedding parameters will be factorized into two
+ matrices in the shape of ['vocab_size', 'embedding_width'] and
+ ['embedding_width', 'hidden_size'] ('embedding_width' is usually much
+ smaller than 'hidden_size').
+ embedding_layer: An optional Layer instance which will be called to generate
+ embeddings for the input word IDs.
+ norm_first: Whether to normalize inputs to attention and intermediate dense
+ layers. If set False, output of attention and intermediate dense layers is
+ normalized.
+ with_dense_inputs: Whether to accept dense embeddings as the input.
+ return_attention_scores: Whether to add an additional output containing the
+ attention scores of all transformer layers. This will be a list of length
+ `num_layers`, and each element will be in the shape [batch_size,
+ num_attention_heads, seq_dim, seq_dim].
+ """
+
+ def __init__(
+ self,
+ vocab_size: int,
+ hidden_size: int = 768,
+ num_layers: int = 12,
+ num_attention_heads: int = 12,
+ max_sequence_length: int = 512,
+ type_vocab_size: int = 16,
+ inner_dim: int = 3072,
+ inner_activation: _Activation = _approx_gelu,
+ output_dropout: float = 0.1,
+ attention_dropout: float = 0.1,
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
+ stddev=0.02),
+ output_range: Optional[int] = None,
+ embedding_width: Optional[int] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
+ norm_first: bool = False,
+ with_dense_inputs: bool = False,
+ return_attention_scores: bool = False,
+ **kwargs):
+ if 'dict_outputs' in kwargs:
+ kwargs.pop('dict_outputs')
+ if 'return_all_encoder_outputs' in kwargs:
+ kwargs.pop('return_all_encoder_outputs')
+ if 'intermediate_size' in kwargs:
+ inner_dim = kwargs.pop('intermediate_size')
+ if 'activation' in kwargs:
+ inner_activation = kwargs.pop('activation')
+ if 'dropout_rate' in kwargs:
+ output_dropout = kwargs.pop('dropout_rate')
+ if 'attention_dropout_rate' in kwargs:
+ attention_dropout = kwargs.pop('attention_dropout_rate')
+ super().__init__(**kwargs)
+
+ self._output_range = output_range
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
+
+ if embedding_width is None:
+ embedding_width = hidden_size
+
+ if embedding_layer is None:
+ self._embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ initializer=tf_utils.clone_initializer(initializer),
+ name='word_embeddings')
+ else:
+ self._embedding_layer = embedding_layer
+
+ self._position_embedding_layer = layers.PositionEmbedding(
+ initializer=tf_utils.clone_initializer(initializer),
+ max_length=max_sequence_length,
+ name='position_embedding')
+
+ self._type_embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=type_vocab_size,
+ embedding_width=embedding_width,
+ initializer=tf_utils.clone_initializer(initializer),
+ use_one_hot=True,
+ name='type_embeddings')
+
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
+ name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
+
+ self._embedding_dropout = tf_keras.layers.Dropout(
+ rate=output_dropout, name='embedding_dropout')
+
+ # We project the 'embedding' output to 'hidden_size' if it is not already
+ # 'hidden_size'.
+ self._embedding_projection = None
+ if embedding_width != hidden_size:
+ self._embedding_projection = tf_keras.layers.EinsumDense(
+ '...x,xy->...y',
+ output_shape=hidden_size,
+ bias_axes='y',
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='embedding_projection')
+
+ self._transformer_layers = []
+ self._attention_mask_layer = layers.SelfAttentionMask(
+ name='self_attention_mask')
+ self._num_layers = num_layers
+ for i in range(num_layers):
+ layer = layers.TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=inner_dim,
+ inner_activation=inner_activation,
+ output_dropout=output_dropout,
+ attention_dropout=attention_dropout,
+ norm_first=norm_first,
+ return_attention_scores=return_attention_scores,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='transformer/layer_%d' % i)
+ self._transformer_layers.append(layer)
+
+ self._pooler_layer_is_entity = tf_keras.layers.Dense(
+ units=hidden_size,
+ activation='tanh',
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='pooler_transform_is_entity')
+ self._pooler_layer_entity_type = tf_keras.layers.Dense(
+ units=hidden_size,
+ activation='tanh',
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='pooler_transform_entity_type')
+
+ self._config = {
+ 'vocab_size': vocab_size,
+ 'hidden_size': hidden_size,
+ 'num_layers': num_layers,
+ 'num_attention_heads': num_attention_heads,
+ 'max_sequence_length': max_sequence_length,
+ 'type_vocab_size': type_vocab_size,
+ 'inner_dim': inner_dim,
+ 'inner_activation': tf_keras.activations.serialize(activation),
+ 'output_dropout': output_dropout,
+ 'attention_dropout': attention_dropout,
+ 'initializer': tf_keras.initializers.serialize(initializer),
+ 'output_range': output_range,
+ 'embedding_width': embedding_width,
+ 'embedding_layer': embedding_layer,
+ 'norm_first': norm_first,
+ 'with_dense_inputs': with_dense_inputs,
+ 'return_attention_scores': return_attention_scores,
+ }
+ if with_dense_inputs:
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_inputs=tf_keras.Input(
+ shape=(None, embedding_width), dtype=tf.float32),
+ dense_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ is_entity_token_pos=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ entity_type_token_pos=tf_keras.Input(shape=(None,), dtype=tf.int32))
+ else:
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ is_entity_token_pos=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ entity_type_token_pos=tf_keras.Input(shape=(None,), dtype=tf.int32))
+
+ def call(self, inputs):
+ word_embeddings = None
+ if isinstance(inputs, dict):
+ word_ids = inputs.get('input_word_ids')
+ mask = inputs.get('input_mask')
+ type_ids = inputs.get('input_type_ids')
+ word_embeddings = inputs.get('input_word_embeddings', None)
+
+ dense_inputs = inputs.get('dense_inputs', None)
+ dense_mask = inputs.get('dense_mask', None)
+ dense_type_ids = inputs.get('dense_type_ids', None)
+
+ is_entity_token_pos = inputs.get('is_entity_token_pos', None)
+ entity_type_token_pos = inputs.get('entity_type_token_pos', None)
+ else:
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ mask = tf.concat([mask, dense_mask], axis=1)
+
+ embeddings = self._get_embeddings(word_ids, type_ids, word_embeddings, # pyrefly: ignore[bad-argument-type]
+ dense_inputs, dense_type_ids)
+ embeddings = self._embedding_norm_layer(embeddings)
+ embeddings = self._embedding_dropout(embeddings)
+
+ if self._embedding_projection is not None:
+ embeddings = self._embedding_projection(embeddings)
+
+ attention_mask = self._attention_mask_layer(embeddings, mask)
+
+ encoder_outputs = []
+ attention_outputs = []
+ x = embeddings
+ for i, layer in enumerate(self._transformer_layers):
+ transformer_output_range = None
+ if i == self._num_layers - 1:
+ transformer_output_range = self._output_range
+ x = layer([x, attention_mask], output_range=transformer_output_range)
+ if self._config['return_attention_scores']:
+ x, attention_scores = x
+ attention_outputs.append(attention_scores)
+ encoder_outputs.append(x)
+
+ last_encoder_output = encoder_outputs[-1]
+ encoder_output_is_entity = tf.gather(
+ last_encoder_output, indices=is_entity_token_pos, axis=1, batch_dims=1)
+ encoder_output_entity_type = tf.gather(
+ last_encoder_output,
+ indices=entity_type_token_pos,
+ axis=1,
+ batch_dims=1)
+ cls_output_is_entity = self._pooler_layer_is_entity(
+ encoder_output_is_entity)
+ cls_output_entity_type = self._pooler_layer_entity_type(
+ encoder_output_entity_type)
+
+ pooled_output = tf.concat([cls_output_is_entity, cls_output_entity_type], 1)
+
+ output = dict(
+ sequence_output=encoder_outputs[-1],
+ pooled_output=pooled_output,
+ encoder_outputs=encoder_outputs)
+ if self._config['return_attention_scores']:
+ output['attention_scores'] = attention_outputs
+ return output
+
+ def get_embedding_table(self):
+ return self._embedding_layer.embeddings
+
+ def get_embedding_layer(self):
+ return self._embedding_layer
+
+ def get_config(self):
+ return dict(self._config)
+
+ @property
+ def transformer_layers(self):
+ """List of Transformer layers in the encoder."""
+ return self._transformer_layers
+
+ @property
+ def pooler_layer_is_entity(self):
+ """The pooler dense layer for is entity classification after the transformer layers.
+ """
+ return self._pooler_layer_is_entity
+
+ @property
+ def pooler_layer_entity_type(self):
+ """The pooler dense layer for entity type classification after the transformer layers.
+ """
+ return self._pooler_layer_entity_type
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ if 'embedding_layer' in config and config['embedding_layer'] is not None:
+ warn_string = (
+ 'You are reloading a model that was saved with a '
+ 'potentially-shared embedding layer object. If you contine to '
+ 'train this model, the embedding layer will no longer be shared. '
+ 'To work around this, load the model outside of the Keras API.')
+ print('WARNING: ' + warn_string)
+ logging.warn(warn_string)
+
+ return cls(**config)
+
+ def _get_embeddings(self, word_ids: tf.Tensor, type_ids: tf.Tensor,
+ word_embeddings: Optional[tf.Tensor],
+ dense_inputs: Optional[tf.Tensor],
+ dense_type_ids: Optional[tf.Tensor]) -> tf.Tensor:
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ # Concat the dense embeddings at sequence end.
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+
+ type_embeddings = self._type_embedding_layer(type_ids)
+
+ # absolute position embeddings.
+ position_embeddings = self._position_embedding_layer(word_embeddings)
+ return word_embeddings + position_embeddings + type_embeddings
diff --git a/official/projects/fffner/fffner_encoder_test.py b/official/projects/fffner/fffner_encoder_test.py
new file mode 100644
index 00000000000..249b8ee4dd8
--- /dev/null
+++ b/official/projects/fffner/fffner_encoder_test.py
@@ -0,0 +1,68 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for official.nlp.projects.fffner.fffner_encoder."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.fffner import fffner_encoder
+
+
+class FFFNerEncoderTest(tf.test.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ np.random.seed(0)
+ tf.random.set_seed(0)
+
+ def test_encoder(self):
+ sequence_length = 128
+ batch_size = 2
+ vocab_size = 1024
+ hidden_size = 256
+ network = fffner_encoder.FFFNerEncoder(
+ vocab_size=vocab_size,
+ hidden_size=hidden_size,
+ num_layers=1,
+ num_attention_heads=4,
+ max_sequence_length=512,
+ dict_outputs=True)
+ word_id_data = np.random.randint(
+ vocab_size, size=(batch_size, sequence_length), dtype=np.int32)
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length), dtype=np.int32)
+ type_id_data = np.random.randint(
+ 2, size=(batch_size, sequence_length), dtype=np.int32)
+ is_entity_token_pos = np.random.randint(
+ sequence_length, size=(batch_size,), dtype=np.int32)
+ entity_type_token_pos = np.random.randint(
+ sequence_length, size=(batch_size,), dtype=np.int32)
+ inputs = {
+ 'input_word_ids': word_id_data,
+ 'input_mask': mask_data,
+ 'input_type_ids': type_id_data,
+ 'is_entity_token_pos': is_entity_token_pos,
+ 'entity_type_token_pos': entity_type_token_pos
+ }
+ outputs = network(inputs)
+ self.assertEqual(outputs['sequence_output'].shape,
+ (batch_size, sequence_length, hidden_size))
+
+ self.assertEqual(outputs['pooled_output'].shape,
+ (batch_size, 2 * hidden_size))
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/fffner/fffner_experiments.py b/official/projects/fffner/fffner_experiments.py
new file mode 100644
index 00000000000..9fd185a7f54
--- /dev/null
+++ b/official/projects/fffner/fffner_experiments.py
@@ -0,0 +1,68 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""FFFNER experiment configurations."""
+# pylint: disable=g-doc-return-or-yield,line-too-long
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.nlp.configs import encoders
+from official.projects.fffner import fffner
+from official.projects.fffner import fffner_dataloader
+from official.projects.fffner import fffner_prediction
+
+AdamWeightDecay = optimization.AdamWeightDecayConfig
+PolynomialLr = optimization.PolynomialLrConfig
+PolynomialWarmupConfig = optimization.PolynomialWarmupConfig
+
+
+@exp_factory.register_config_factory('fffner/ner')
+def fffner_ner() -> cfg.ExperimentConfig:
+ """Defines fffner experiments."""
+ config = cfg.ExperimentConfig(
+ task=fffner_prediction.FFFNerPredictionConfig(
+ model=fffner_prediction.FFFNerModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=fffner.FFFNerEncoderConfig())),
+ train_data=fffner_dataloader.FFFNerDataConfig(),
+ validation_data=fffner_dataloader.FFFNerDataConfig(
+ is_training=False, drop_remainder=False,
+ include_example_id=True)),
+ trainer=cfg.TrainerConfig(
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate':
+ 0.01,
+ 'exclude_from_weight_decay':
+ ['LayerNorm', 'layer_norm', 'bias'],
+ }
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 2e-5,
+ 'end_learning_rate': 0.0,
+ }
+ },
+ 'warmup': {
+ 'type': 'polynomial'
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
diff --git a/official/projects/fffner/fffner_prediction.py b/official/projects/fffner/fffner_prediction.py
new file mode 100644
index 00000000000..2d2d984b6d7
--- /dev/null
+++ b/official/projects/fffner/fffner_prediction.py
@@ -0,0 +1,372 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""FFFNER prediction task."""
+import collections
+import dataclasses
+
+from absl import logging
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import config_definitions as cfg
+from official.core import task_factory
+from official.modeling import tf_utils
+from official.modeling.hyperparams import base_config
+from official.nlp.configs import encoders
+from official.nlp.data import data_loader_factory
+from official.nlp.tasks import utils
+from official.projects.fffner import fffner_classifier
+
+METRIC_TYPES = frozenset(
+ ['accuracy', 'matthews_corrcoef', 'pearson_spearman_corr'])
+
+
+@dataclasses.dataclass
+class FFFNerModelConfig(base_config.Config):
+ """A classifier/regressor configuration."""
+ num_classes_is_entity: int = 0
+ num_classes_entity_type: int = 0
+ use_encoder_pooler: bool = True
+ encoder: encoders.EncoderConfig = dataclasses.field(
+ default_factory=encoders.EncoderConfig
+ )
+
+
+@dataclasses.dataclass
+class FFFNerPredictionConfig(cfg.TaskConfig):
+ """The model config."""
+ # At most one of `init_checkpoint` and `hub_module_url` can
+ # be specified.
+ init_checkpoint: str = ''
+ init_cls_pooler: bool = False
+ hub_module_url: str = ''
+ metric_type: str = 'accuracy'
+ # Defines the concrete model config at instantiation time.
+ model: FFFNerModelConfig = dataclasses.field(
+ default_factory=FFFNerModelConfig
+ )
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
+
+
+@task_factory.register_task_cls(FFFNerPredictionConfig)
+class FFFNerTask(base_task.Task):
+ """Task object for FFFNer."""
+
+ def __init__(self, params: cfg.TaskConfig, logging_dir=None, name=None):
+ super().__init__(params, logging_dir, name=name)
+ if params.metric_type not in METRIC_TYPES:
+ raise ValueError('Invalid metric_type: {}'.format(params.metric_type))
+ self.metric_type = params.metric_type
+ self.label_field_is_entity = 'is_entity_label'
+ self.label_field_entity_type = 'entity_type_label'
+
+ def build_model(self):
+ if self.task_config.hub_module_url and self.task_config.init_checkpoint:
+ raise ValueError('At most one of `hub_module_url` and '
+ '`init_checkpoint` can be specified.')
+ if self.task_config.hub_module_url:
+ encoder_network = utils.get_encoder_from_hub(
+ self.task_config.hub_module_url)
+ else:
+ encoder_network = encoders.build_encoder(self.task_config.model.encoder)
+ encoder_cfg = self.task_config.model.encoder.get()
+ if self.task_config.model.encoder.type == 'xlnet':
+ assert False, 'Not supported yet'
+ else:
+ return fffner_classifier.FFFNerClassifier(
+ # encoder_network.inputs
+ network=encoder_network,
+ num_classes_is_entity=self.task_config.model.num_classes_is_entity,
+ num_classes_entity_type=self.task_config.model
+ .num_classes_entity_type,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ use_encoder_pooler=self.task_config.model.use_encoder_pooler)
+
+ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
+ label_ids_is_entity = labels[self.label_field_is_entity]
+ label_ids_entity_type = labels[self.label_field_entity_type]
+ loss_is_entity = tf_keras.losses.sparse_categorical_crossentropy(
+ label_ids_is_entity,
+ tf.cast(model_outputs[0], tf.float32),
+ from_logits=True)
+ loss_entity_type = tf_keras.losses.sparse_categorical_crossentropy(
+ label_ids_entity_type,
+ tf.cast(model_outputs[1], tf.float32),
+ from_logits=True)
+ loss = loss_is_entity + loss_entity_type
+
+ if aux_losses:
+ loss += tf.add_n(aux_losses)
+ return tf_utils.safe_mean(loss)
+
+ def build_inputs(self, params, input_context=None):
+ """Returns tf.data.Dataset for sentence_prediction task."""
+ if params.input_path == 'dummy':
+
+ def dummy_data(_):
+ dummy_ids = tf.zeros((1, params.seq_length), dtype=tf.int32)
+ x = dict(
+ input_word_ids=dummy_ids,
+ input_mask=dummy_ids,
+ input_type_ids=dummy_ids,
+ is_entity_token_pos=tf.zeros((1, 1), dtype=tf.int32),
+ entity_type_token_pos=tf.ones((1, 1), dtype=tf.int32))
+
+ x[self.label_field_is_entity] = tf.zeros((1, 1), dtype=tf.int32)
+ x[self.label_field_entity_type] = tf.zeros((1, 1), dtype=tf.int32)
+ return x
+
+ dataset = tf.data.Dataset.range(1)
+ dataset = dataset.repeat()
+ dataset = dataset.map(
+ dummy_data, num_parallel_calls=tf.data.experimental.AUTOTUNE)
+ return dataset
+
+ return data_loader_factory.get_data_loader(params).load(input_context)
+
+ def build_metrics(self, training=None):
+ del training
+ metrics = [
+ tf_keras.metrics.SparseCategoricalAccuracy(
+ name='cls_accuracy_is_entity'),
+ tf_keras.metrics.SparseCategoricalAccuracy(
+ name='cls_accuracy_entity_type'),
+ ]
+ return metrics
+
+ def process_metrics(self, metrics, labels, model_outputs):
+ for metric in metrics:
+ if metric.name == 'cls_accuracy_is_entity':
+ metric.update_state(labels[self.label_field_is_entity],
+ model_outputs[0])
+ if metric.name == 'cls_accuracy_entity_type':
+ metric.update_state(labels[self.label_field_entity_type],
+ model_outputs[1])
+
+ def process_compiled_metrics(self, compiled_metrics, labels, model_outputs):
+ compiled_metrics.update_state(labels[self.label_field_is_entity],
+ model_outputs[0])
+ compiled_metrics.update_state(labels[self.label_field_entity_type],
+ model_outputs[1])
+
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
+ features, labels = inputs, inputs
+ outputs = self.inference_step(features, model)
+ loss = self.build_losses(
+ labels=labels, model_outputs=outputs, aux_losses=model.losses)
+ logs = {self.loss: loss}
+ if metrics:
+ self.process_metrics(metrics, labels, outputs)
+ if model.compiled_metrics:
+ self.process_compiled_metrics(model.compiled_metrics, labels, outputs)
+ logs.update({m.name: m.result() for m in metrics or []})
+ logs.update({m.name: m.result() for m in model.metrics})
+ logs.update({
+ 'sentence_prediction_is_entity': outputs[0],
+ 'sentence_prediction_entity_type': outputs[1],
+ 'labels_is_entity': labels[self.label_field_is_entity],
+ 'labels_entity_type': labels[self.label_field_entity_type],
+ 'id': labels['example_id'],
+ 'sentence_id': labels['sentence_id'],
+ 'span_start': labels['span_start'],
+ 'span_end': labels['span_end']
+ })
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ if state is None:
+ state = {
+ 'sentence_prediction_is_entity': [],
+ 'sentence_prediction_entity_type': [],
+ 'labels_is_entity': [],
+ 'labels_entity_type': [],
+ 'ids': [],
+ 'sentence_id': [],
+ 'span_start': [],
+ 'span_end': []
+ }
+ state['sentence_prediction_is_entity'].append(
+ np.concatenate(
+ [v.numpy() for v in step_outputs['sentence_prediction_is_entity']], # pyrefly: ignore[unsupported-operation]
+ axis=0))
+ state['sentence_prediction_entity_type'].append(
+ np.concatenate([
+ v.numpy() for v in step_outputs['sentence_prediction_entity_type'] # pyrefly: ignore[unsupported-operation]
+ ],
+ axis=0))
+ state['labels_is_entity'].append(
+ np.concatenate([v.numpy() for v in step_outputs['labels_is_entity']], # pyrefly: ignore[unsupported-operation]
+ axis=0))
+ state['labels_entity_type'].append(
+ np.concatenate([v.numpy() for v in step_outputs['labels_entity_type']], # pyrefly: ignore[unsupported-operation]
+ axis=0))
+ state['ids'].append(
+ np.concatenate([v.numpy() for v in step_outputs['id']], axis=0)) # pyrefly: ignore[unsupported-operation]
+ state['sentence_id'].append(
+ np.concatenate([v.numpy() for v in step_outputs['sentence_id']], # pyrefly: ignore[unsupported-operation]
+ axis=0))
+ state['span_start'].append(
+ np.concatenate([v.numpy() for v in step_outputs['span_start']], axis=0)) # pyrefly: ignore[unsupported-operation]
+ state['span_end'].append(
+ np.concatenate([v.numpy() for v in step_outputs['span_end']], axis=0)) # pyrefly: ignore[unsupported-operation]
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ sentence_prediction_is_entity = np.concatenate(
+ aggregated_logs['sentence_prediction_is_entity'], axis=0)
+ sentence_prediction_is_entity = np.reshape(
+ sentence_prediction_is_entity,
+ (-1, self.task_config.model.num_classes_is_entity))
+ sentence_prediction_entity_type = np.concatenate(
+ aggregated_logs['sentence_prediction_entity_type'], axis=0)
+ sentence_prediction_entity_type = np.reshape(
+ sentence_prediction_entity_type,
+ (-1, self.task_config.model.num_classes_entity_type))
+ labels_is_entity = np.concatenate(
+ aggregated_logs['labels_is_entity'], axis=0)
+ labels_is_entity = np.reshape(labels_is_entity, -1)
+ labels_entity_type = np.concatenate(
+ aggregated_logs['labels_entity_type'], axis=0)
+ labels_entity_type = np.reshape(labels_entity_type, -1)
+
+ ids = np.concatenate(aggregated_logs['ids'], axis=0)
+ ids = np.reshape(ids, -1)
+ sentence_id = np.concatenate(aggregated_logs['sentence_id'], axis=0)
+ sentence_id = np.reshape(sentence_id, -1)
+ span_start = np.concatenate(aggregated_logs['span_start'], axis=0)
+ span_start = np.reshape(span_start, -1)
+ span_end = np.concatenate(aggregated_logs['span_end'], axis=0)
+ span_end = np.reshape(span_end, -1)
+
+ def resolve(length, spans, prediction_confidence):
+ used = [False] * length
+ spans = sorted(
+ spans,
+ key=lambda x: prediction_confidence[(x[0], x[1])],
+ reverse=True)
+ real_spans = []
+ for span_start, span_end, ent_type in spans:
+ fill = False
+ for s in range(span_start, span_end + 1):
+ if used[s]:
+ fill = True
+ break
+ if not fill:
+ real_spans.append((span_start, span_end, ent_type))
+ for s in range(span_start, span_end + 1):
+ used[s] = True
+ return real_spans
+
+ def get_p_r_f(truth, pred):
+ n_pred = len(pred)
+ n_truth = len(truth)
+ n_correct = len(set(pred) & set(truth))
+ precision = 1. * n_correct / n_pred if n_pred != 0 else 0.0
+ recall = 1. * n_correct / n_truth if n_truth != 0 else 0.0
+ f1 = 2 * precision * recall / (
+ precision + recall) if precision + recall != 0.0 else 0.0
+ return {
+ 'n_pred': n_pred,
+ 'n_truth': n_truth,
+ 'n_correct': n_correct,
+ 'precision': precision,
+ 'recall': recall,
+ 'f1': f1,
+ }
+
+ def softmax(x):
+ x = np.array(x)
+ e_x = np.exp(x - np.max(x))
+ return e_x / e_x.sum(axis=0)
+
+ per_sid_results = collections.defaultdict(list)
+ for _, sent_id, sp_start, sp_end, is_entity_label, is_entity_logit, entity_type_label, entity_type_logit in zip(
+ ids, sentence_id, span_start, span_end, labels_is_entity,
+ sentence_prediction_is_entity, labels_entity_type,
+ sentence_prediction_entity_type):
+ if sent_id > 0:
+ per_sid_results[sent_id].append(
+ (sp_start, sp_end, is_entity_label, is_entity_logit,
+ entity_type_label, entity_type_logit))
+ ground_truth = []
+ prediction_is_entity = []
+ prediction_entity_type = []
+ for key in sorted(list(per_sid_results.keys())):
+ results = per_sid_results[key]
+ gt_entities = []
+ predictied_entities = []
+ prediction_confidence = {}
+ prediction_confidence_type = {}
+ length = 0
+ for span_start, span_end, ground_truth_span, prediction_span, ground_truth_type, prediction_type in results:
+ if ground_truth_span == 1:
+ gt_entities.append((span_start, span_end, ground_truth_type))
+ if prediction_span[1] > prediction_span[0]:
+ predictied_entities.append(
+ (span_start, span_end, np.argmax(prediction_type).item()))
+ prediction_confidence[(span_start,
+ span_end)] = max(softmax(prediction_span))
+ prediction_confidence_type[(span_start,
+ span_end)] = max(softmax(prediction_type))
+ length = max(length, span_end)
+ length += 1
+ ground_truth.extend([(key, *x) for x in gt_entities])
+ prediction_is_entity.extend([(key, *x) for x in predictied_entities])
+ resolved_predicted = resolve(length, predictied_entities,
+ prediction_confidence)
+ prediction_entity_type.extend([(key, *x) for x in resolved_predicted])
+
+ raw = get_p_r_f(ground_truth, prediction_is_entity)
+ resolved = get_p_r_f(ground_truth, prediction_entity_type)
+ return {
+ 'raw_f1': raw['f1'],
+ 'raw_precision': raw['precision'],
+ 'raw_recall': raw['recall'],
+ 'resolved_f1': resolved['f1'],
+ 'resolved_precision': resolved['precision'],
+ 'resolved_recall': resolved['recall'],
+ 'overall_f1': raw['f1'] + resolved['f1'],
+ }
+
+ def initialize(self, model):
+ """Load a pretrained checkpoint (if exists) and then train from iter 0."""
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ logging.info('Trying to load pretrained checkpoint from %s',
+ ckpt_dir_or_file)
+ if ckpt_dir_or_file and tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+ if not ckpt_dir_or_file:
+ logging.info('No checkpoint file found from %s. Will not load.',
+ ckpt_dir_or_file)
+ return
+
+ pretrain2finetune_mapping = {
+ 'encoder': model.checkpoint_items['encoder'],
+ }
+ if self.task_config.init_cls_pooler:
+ # This option is valid when use_encoder_pooler is false.
+ pretrain2finetune_mapping[
+ 'next_sentence.pooler_dense'] = model.checkpoint_items[
+ 'sentence_prediction.pooler_dense']
+ ckpt = tf.train.Checkpoint(**pretrain2finetune_mapping)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
diff --git a/official/projects/fffner/train.py b/official/projects/fffner/train.py
new file mode 100644
index 00000000000..1239ef94822
--- /dev/null
+++ b/official/projects/fffner/train.py
@@ -0,0 +1,27 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""FFFNER trainer."""
+
+from absl import app
+
+from official.common import flags as tfm_flags
+# pylint: disable=unused-import
+from official.nlp import train
+from official.projects.fffner import fffner_experiments
+# pylint: enable=unused-import
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(train.main)
diff --git a/official/projects/fffner/utils/convert_checkpoint_huggingface.py b/official/projects/fffner/utils/convert_checkpoint_huggingface.py
new file mode 100644
index 00000000000..e9a2a57ec75
--- /dev/null
+++ b/official/projects/fffner/utils/convert_checkpoint_huggingface.py
@@ -0,0 +1,160 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Converts pre-trained pytorch checkpoint into a tf encoder checkpoint."""
+
+import os
+
+from absl import app
+import numpy as np
+import tensorflow as tf, tf_keras
+import transformers
+
+from official.modeling import tf_utils
+from official.projects.fffner.fffner import FFFNerEncoderConfig
+from official.projects.fffner.fffner_encoder import FFFNerEncoder
+
+
+def _get_huggingface_bert_model_and_config(huggingface_model_name_or_path):
+ model = transformers.AutoModel.from_pretrained(huggingface_model_name_or_path)
+
+ return {n: p.data.numpy() for n, p in model.named_parameters()}, model.config
+
+
+def _create_fffner_model(huggingface_bert_config):
+ """Creates a Longformer model."""
+ encoder_cfg = FFFNerEncoderConfig()
+ encoder = FFFNerEncoder(
+ vocab_size=huggingface_bert_config.vocab_size,
+ hidden_size=huggingface_bert_config.hidden_size,
+ num_layers=huggingface_bert_config.num_hidden_layers,
+ num_attention_heads=huggingface_bert_config.num_attention_heads,
+ inner_dim=huggingface_bert_config.intermediate_size,
+ inner_activation=tf_utils.get_activation(
+ huggingface_bert_config.hidden_act),
+ output_dropout=huggingface_bert_config.hidden_dropout_prob,
+ attention_dropout=huggingface_bert_config.attention_probs_dropout_prob,
+ max_sequence_length=huggingface_bert_config.max_position_embeddings,
+ type_vocab_size=huggingface_bert_config.type_vocab_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ output_range=encoder_cfg.output_range,
+ embedding_width=huggingface_bert_config.hidden_size,
+ norm_first=encoder_cfg.norm_first)
+ return encoder
+
+
+# pylint: disable=protected-access
+def convert(encoder, bert_model):
+ """Convert a Huggingface transformers bert encoder to the one in the codebase.
+ """
+ num_layers = encoder._config["num_layers"]
+ num_attention_heads = encoder._config["num_attention_heads"]
+ hidden_size = encoder._config["hidden_size"]
+ head_size = hidden_size // num_attention_heads
+ assert head_size * num_attention_heads == hidden_size
+ encoder._embedding_layer.set_weights(
+ [bert_model["embeddings.word_embeddings.weight"]])
+ encoder._embedding_norm_layer.set_weights([
+ bert_model["embeddings.LayerNorm.weight"],
+ bert_model["embeddings.LayerNorm.bias"]
+ ])
+ encoder._type_embedding_layer.set_weights(
+ [bert_model["embeddings.token_type_embeddings.weight"]])
+ encoder._position_embedding_layer.set_weights(
+ [bert_model["embeddings.position_embeddings.weight"]])
+ for layer_num in range(num_layers):
+ encoder._transformer_layers[
+ layer_num]._attention_layer._key_dense.set_weights([
+ bert_model[f"encoder.layer.{layer_num}.attention.self.key.weight"].T
+ .reshape((hidden_size, num_attention_heads, head_size)),
+ bert_model[f"encoder.layer.{layer_num}.attention.self.key.bias"]
+ .reshape((num_attention_heads, head_size))
+ ])
+ encoder._transformer_layers[
+ layer_num]._attention_layer._query_dense.set_weights([
+ bert_model[f"encoder.layer.{layer_num}.attention.self.query.weight"]
+ .T.reshape((hidden_size, num_attention_heads, head_size)),
+ bert_model[f"encoder.layer.{layer_num}.attention.self.query.bias"]
+ .reshape((num_attention_heads, head_size))
+ ])
+ encoder._transformer_layers[
+ layer_num]._attention_layer._value_dense.set_weights([
+ bert_model[f"encoder.layer.{layer_num}.attention.self.value.weight"]
+ .T.reshape((hidden_size, num_attention_heads, head_size)),
+ bert_model[f"encoder.layer.{layer_num}.attention.self.value.bias"]
+ .reshape((num_attention_heads, head_size))
+ ])
+ encoder._transformer_layers[
+ layer_num]._attention_layer._output_dense.set_weights([
+ bert_model[
+ f"encoder.layer.{layer_num}.attention.output.dense.weight"].T
+ .reshape((num_attention_heads, head_size, hidden_size)),
+ bert_model[f"encoder.layer.{layer_num}.attention.output.dense.bias"]
+ ])
+ encoder._transformer_layers[layer_num]._attention_layer_norm.set_weights([
+ bert_model[
+ f"encoder.layer.{layer_num}.attention.output.LayerNorm.weight"],
+ bert_model[f"encoder.layer.{layer_num}.attention.output.LayerNorm.bias"]
+ ])
+ encoder._transformer_layers[layer_num]._intermediate_dense.set_weights([
+ bert_model[f"encoder.layer.{layer_num}.intermediate.dense.weight"].T,
+ bert_model[f"encoder.layer.{layer_num}.intermediate.dense.bias"]
+ ])
+ encoder._transformer_layers[layer_num]._output_dense.set_weights([
+ bert_model[f"encoder.layer.{layer_num}.output.dense.weight"].T,
+ bert_model[f"encoder.layer.{layer_num}.output.dense.bias"]
+ ])
+ encoder._transformer_layers[layer_num]._output_layer_norm.set_weights([
+ bert_model[f"encoder.layer.{layer_num}.output.LayerNorm.weight"],
+ bert_model[f"encoder.layer.{layer_num}.output.LayerNorm.bias"]
+ ])
+
+
+def convert_checkpoint(huggingface_model_name_or_path, output_path):
+ """Converts and save the checkpoint."""
+ output_dir, _ = os.path.split(output_path)
+ tf.io.gfile.makedirs(output_dir)
+
+ huggingface_bert_model, huggingface_bert_config = _get_huggingface_bert_model_and_config(
+ huggingface_model_name_or_path)
+ encoder = _create_fffner_model(huggingface_bert_config)
+ sequence_length = 128
+ batch_size = 2
+ word_id_data = np.random.randint(
+ 10, size=(batch_size, sequence_length), dtype=np.int32)
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length), dtype=np.int32)
+ type_id_data = np.random.randint(
+ 2, size=(batch_size, sequence_length), dtype=np.int32)
+ is_entity_token_pos = np.zeros((batch_size, 1), dtype=np.int32)
+ entity_type_token_pos = np.ones((batch_size, 1), dtype=np.int32)
+ inputs = {
+ "input_word_ids": word_id_data,
+ "input_mask": mask_data,
+ "input_type_ids": type_id_data,
+ "is_entity_token_pos": is_entity_token_pos,
+ "entity_type_token_pos": entity_type_token_pos,
+ }
+ encoder(inputs)
+ convert(encoder, huggingface_bert_model)
+ tf.train.Checkpoint(encoder=encoder).write(output_path)
+
+
+def main(_):
+ convert_checkpoint("bert-base-uncased", "bert-uncased")
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/official/projects/fffner/utils/convert_checkpoint_tensorflow.py b/official/projects/fffner/utils/convert_checkpoint_tensorflow.py
new file mode 100644
index 00000000000..5bf08384ec8
--- /dev/null
+++ b/official/projects/fffner/utils/convert_checkpoint_tensorflow.py
@@ -0,0 +1,179 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Converts pre-trained encoder into a fffner encoder checkpoint."""
+import os
+
+from absl import app
+import numpy as np
+import tensorflow as tf, tf_keras
+import tensorflow_hub as hub
+
+from official.projects.fffner.fffner import FFFNerEncoderConfig
+from official.projects.fffner.fffner_encoder import FFFNerEncoder
+
+
+def _get_tensorflow_bert_model_and_config(tfhub_handle_encoder):
+ """Gets the BERT model name-parameters pairs and configurations."""
+ bert_model = hub.KerasLayer(tfhub_handle_encoder)
+ bert_model_weights_name = [w.name for w in bert_model.weights]
+ bert_model_weights = bert_model.get_weights()
+ named_parameters = {
+ n: p for n, p in zip(bert_model_weights_name, bert_model_weights)
+ }
+ config = {}
+ config["num_attention_heads"], _, config["hidden_size"] = named_parameters[
+ "transformer/layer_0/self_attention/attention_output/kernel:0"].shape
+ _, config["intermediate_size"] = named_parameters[
+ "transformer/layer_0/intermediate/kernel:0"].shape
+ num_hidden_layers = 0
+ while f"transformer/layer_{num_hidden_layers}/self_attention/query/kernel:0" in named_parameters:
+ num_hidden_layers += 1
+ config["num_hidden_layers"] = num_hidden_layers
+ config["vocab_size"], _ = named_parameters[
+ "word_embeddings/embeddings:0"].shape
+ config["max_position_embeddings"], _ = named_parameters[
+ "position_embedding/embeddings:0"].shape
+ config["type_vocab_size"], _ = named_parameters[
+ "type_embeddings/embeddings:0"].shape
+ return named_parameters, config
+
+
+def _create_fffner_model(bert_config):
+ """Creates a Longformer model."""
+ encoder_cfg = FFFNerEncoderConfig()
+ encoder = FFFNerEncoder(
+ vocab_size=bert_config["vocab_size"],
+ hidden_size=bert_config["hidden_size"],
+ num_layers=bert_config["num_hidden_layers"],
+ num_attention_heads=bert_config["num_attention_heads"],
+ inner_dim=bert_config["intermediate_size"],
+ max_sequence_length=bert_config["max_position_embeddings"],
+ type_vocab_size=bert_config["type_vocab_size"],
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range),
+ output_range=encoder_cfg.output_range,
+ embedding_width=bert_config["hidden_size"],
+ norm_first=encoder_cfg.norm_first)
+ return encoder
+
+
+# pylint: disable=protected-access
+def convert(encoder, bert_model):
+ """Convert a Tensorflow transformers bert encoder to the one in the codebase.
+ """
+ num_layers = encoder._config["num_layers"]
+ num_attention_heads = encoder._config["num_attention_heads"]
+ hidden_size = encoder._config["hidden_size"]
+ head_size = hidden_size // num_attention_heads
+ assert head_size * num_attention_heads == hidden_size
+ encoder._embedding_layer.set_weights(
+ [bert_model["word_embeddings/embeddings:0"]])
+ encoder._embedding_norm_layer.set_weights([
+ bert_model["embeddings/layer_norm/gamma:0"],
+ bert_model["embeddings/layer_norm/beta:0"]
+ ])
+ encoder._type_embedding_layer.set_weights(
+ [bert_model["type_embeddings/embeddings:0"]])
+ encoder._position_embedding_layer.set_weights(
+ [bert_model["position_embedding/embeddings:0"]])
+ for layer_num in range(num_layers):
+ encoder._transformer_layers[
+ layer_num]._attention_layer._key_dense.set_weights([
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/key/kernel:0"],
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/key/bias:0"]
+ ])
+ encoder._transformer_layers[
+ layer_num]._attention_layer._query_dense.set_weights([
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/query/kernel:0"],
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/query/bias:0"]
+ ])
+ encoder._transformer_layers[
+ layer_num]._attention_layer._value_dense.set_weights([
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/value/kernel:0"],
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/value/bias:0"]
+ ])
+
+ encoder._transformer_layers[layer_num]._attention_layer._output_dense.set_weights([
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/attention_output/kernel:0"],
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention/attention_output/bias:0"]
+ ])
+ encoder._transformer_layers[layer_num]._attention_layer_norm.set_weights([
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention_layer_norm/gamma:0"],
+ bert_model[
+ f"transformer/layer_{layer_num}/self_attention_layer_norm/beta:0"]
+ ])
+
+ encoder._transformer_layers[layer_num]._intermediate_dense.set_weights([
+ bert_model[f"transformer/layer_{layer_num}/intermediate/kernel:0"],
+ bert_model[f"transformer/layer_{layer_num}/intermediate/bias:0"]
+ ])
+ encoder._transformer_layers[layer_num]._output_dense.set_weights([
+ bert_model[f"transformer/layer_{layer_num}/output/kernel:0"],
+ bert_model[f"transformer/layer_{layer_num}/output/bias:0"]
+ ])
+ encoder._transformer_layers[layer_num]._output_layer_norm.set_weights([
+ bert_model[f"transformer/layer_{layer_num}/output_layer_norm/gamma:0"],
+ bert_model[f"transformer/layer_{layer_num}/output_layer_norm/beta:0"]
+ ])
+
+
+def convert_checkpoint(output_path, tfhub_handle_encoder):
+ """Converts and save the checkpoint."""
+ output_dir, _ = os.path.split(output_path)
+ tf.io.gfile.makedirs(output_dir)
+
+ bert_model, bert_config = _get_tensorflow_bert_model_and_config(
+ tfhub_handle_encoder)
+ encoder = _create_fffner_model(bert_config)
+ sequence_length = 128
+ batch_size = 2
+ word_id_data = np.random.randint(
+ 10, size=(batch_size, sequence_length), dtype=np.int32)
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length), dtype=np.int32)
+ type_id_data = np.random.randint(
+ 2, size=(batch_size, sequence_length), dtype=np.int32)
+ is_entity_token_pos = np.zeros((batch_size, 1), dtype=np.int32)
+ entity_type_token_pos = np.ones((batch_size, 1), dtype=np.int32)
+ inputs = {
+ "input_word_ids": word_id_data,
+ "input_mask": mask_data,
+ "input_type_ids": type_id_data,
+ "is_entity_token_pos": is_entity_token_pos,
+ "entity_type_token_pos": entity_type_token_pos,
+ }
+ encoder(inputs)
+ convert(encoder, bert_model)
+ tf.train.Checkpoint(encoder=encoder).write(output_path)
+
+
+def main(_):
+ convert_checkpoint(
+ output_path="tf-bert-uncased",
+ tfhub_handle_encoder="https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/3"
+ )
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/official/projects/fffner/utils/create_data.py b/official/projects/fffner/utils/create_data.py
new file mode 100644
index 00000000000..8cd29f3073b
--- /dev/null
+++ b/official/projects/fffner/utils/create_data.py
@@ -0,0 +1,347 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Creates the datasets for FFF-NER model."""
+import collections
+import json
+import math
+import os
+import sys
+
+import numpy as np
+import tensorflow as tf, tf_keras
+from tqdm import tqdm
+import transformers
+
+
+class NERDataset:
+ """A Named Entity Recognition dataset for FFF-NER model."""
+
+ def __init__(self, words_path, labels_path, tokenizer, is_train,
+ label_to_entity_type_index, ablation_not_mask,
+ ablation_no_brackets, ablation_span_type_together):
+ """Instantiates the class.
+
+ Args:
+ words_path: Path to the .words file that contains the text.
+ labels_path: Path to the .ner file that contains NER labels for the text.
+ tokenizer: A huggingface tokenizer.
+ is_train: If creating a dataset for training, otherwise testing.
+ label_to_entity_type_index: A mapping of NER labels to indices.
+ ablation_not_mask: An ablation experiment that does not use mask tokens.
+ ablation_no_brackets: An ablation experiment that does not use brackets.
+ ablation_span_type_together: An ablation experiment that does span and
+ type prediction together at a single token.
+ """
+ self.words_path = words_path
+ self.labels_path = labels_path
+ self.tokenizer = tokenizer
+ self.is_train = is_train
+ self.label_to_entity_type_index = label_to_entity_type_index
+ self.ablation_no_brackets = ablation_no_brackets
+ self.ablation_span_type_together = ablation_span_type_together
+ self.ablation_not_mask = ablation_not_mask
+
+ self.left_bracket = self.tokenize_word(" [")[0]
+ self.right_bracket = self.tokenize_word(" ]")[0]
+ self.mask_id = self.tokenizer.mask_token_id
+ self.cls_token_id = self.tokenizer.cls_token_id
+ self.sep_token_id = self.tokenizer.sep_token_id
+
+ self.data = []
+ self.id_to_sentence_infos = dict()
+ self.id_counter = 0
+ self.all_tokens = []
+ self.all_labels = []
+ self.max_seq_len_in_data = 0
+ self.max_len = 128
+
+ def read_file(self):
+ """Reads the input files from words_path and labels_paths."""
+ with open(self.words_path) as f1, open(self.labels_path) as f2:
+ for _, (l1, l2) in enumerate(zip(f1, f2)):
+ tokens = l1.strip().split(" ")
+ labels = l2.strip().split(" ")
+ # since we are use [ and ], we replace all [, ] in the text with (, )
+ tokens = ["(" if token == "[" else token for token in tokens]
+ tokens = [")" if token == "]" else token for token in tokens]
+ yield tokens, labels
+
+ def tokenize_word(self, word):
+ """Calls the tokenizer to produce word ids from text."""
+ result = self.tokenizer(word, add_special_tokens=False)
+ return result["input_ids"]
+
+ def tokenize_word_list(self, word_list):
+ return [self.tokenize_word(word) for word in word_list]
+
+ def process_to_input(self, input_ids, is_entity_token_pos,
+ entity_type_token_pos, is_entity_label,
+ entity_type_label, sid, span_start, span_end):
+ """Process and store sentence and span id information."""
+ self.id_counter += 1
+ self.id_to_sentence_infos[self.id_counter] = {
+ "sid": sid, # sentence id
+ "span_start": span_start,
+ "span_end": span_end,
+ }
+ seqlen = len(input_ids)
+ self.max_seq_len_in_data = max(self.max_seq_len_in_data, seqlen)
+ return {
+ "input_ids": input_ids,
+ "attention_mask": [1] * seqlen,
+ "is_entity_token_pos": is_entity_token_pos,
+ "entity_type_token_pos": entity_type_token_pos,
+ "is_entity_label": 1 if is_entity_label else 0,
+ "entity_type_label": entity_type_label,
+ "sentence_id": sid,
+ "span_start": span_start,
+ "span_end": span_end,
+ "id": self.id_counter,
+ }
+
+ def process_word_list_and_spans_to_inputs(self, sid, word_list, spans):
+ """Constructs the fffner input with spans and types."""
+ tokenized_word_list = self.tokenize_word_list(word_list)
+ final_len = sum(len(x) for x in tokenized_word_list)
+ final_len = 2 + 3 + 2 + 3 + final_len # account for mask and brackets
+ if final_len > self.max_len:
+ print(f"final_len {final_len} too long, skipping")
+ return
+ for span_start, span_end, span_type, span_label in spans:
+ assert span_type == "mask"
+ input_ids = []
+ input_ids.append(self.cls_token_id)
+ for ids in tokenized_word_list[:span_start]:
+ input_ids.extend(ids)
+
+ if not self.ablation_span_type_together:
+ if not self.ablation_no_brackets:
+ input_ids.append(self.left_bracket)
+ is_entity_token_pos = len(input_ids)
+ input_ids.append(self.mask_id if not self.ablation_not_mask else 8487)
+ if not self.ablation_no_brackets:
+ input_ids.append(self.right_bracket)
+
+ if not self.ablation_no_brackets:
+ input_ids.append(self.left_bracket)
+ for ids in tokenized_word_list[span_start:span_end + 1]:
+ input_ids.extend(ids)
+ if not self.ablation_no_brackets:
+ input_ids.append(self.right_bracket)
+
+ if not self.ablation_no_brackets:
+ input_ids.append(self.left_bracket)
+
+ entity_type_token_pos = len(input_ids)
+ if self.ablation_span_type_together:
+ is_entity_token_pos = len(input_ids)
+
+ input_ids.append(self.mask_id if not self.ablation_not_mask else 2828)
+ if not self.ablation_no_brackets:
+ input_ids.append(self.right_bracket)
+
+ for ids in tokenized_word_list[span_end + 1:]:
+ input_ids.extend(ids)
+ input_ids.append(self.sep_token_id)
+ is_entity_label = span_label in self.label_to_entity_type_index
+ entity_type_label = self.label_to_entity_type_index.get(span_label, 0)
+ yield self.process_to_input(input_ids, is_entity_token_pos,
+ entity_type_token_pos, is_entity_label,
+ entity_type_label, sid, span_start, span_end)
+
+ def bio_labels_to_spans(self, bio_labels):
+ """Gets labels to spans."""
+ spans = []
+ for i, label in enumerate(bio_labels):
+ if label.startswith("B-"):
+ spans.append([i, i, label[2:]])
+ elif label.startswith("I-"):
+ if spans:
+ print("Error... I-tag should not start a span")
+ spans.append([i, i, label[2:]])
+ elif spans[-1][1] != i - 1 or spans[-1][2] != label[2:]:
+ print("Error... I-tag not consistent with previous tag")
+ spans.append([i, i, label[2:]])
+ else:
+ spans[-1][1] = i
+ elif label.startswith("O"):
+ pass
+ else:
+ assert False, bio_labels
+ spans = list(
+ filter(lambda x: x[2] in self.label_to_entity_type_index.keys(), spans))
+ return spans
+
+ def collate_fn(self, batch):
+ batch = self.tokenizer.pad(
+ batch,
+ padding="max_length",
+ max_length=self.max_len,
+ )
+ return batch
+
+ def prepare(self, negative_multiplier=3.):
+ """Constructs negative sampling and handling train/test differences."""
+ desc = ("prepare data for training"
+ if self.is_train else "prepare data for testing")
+ total_missed_entities = 0
+ total_entities = 0
+ for sid, (tokens, labels) in tqdm(enumerate(self.read_file()), desc=desc):
+ self.all_tokens.append(tokens)
+ self.all_labels.append(labels)
+ entity_spans = self.bio_labels_to_spans(labels)
+ entity_spans_dict = {
+ (start, end): ent_type for start, end, ent_type in entity_spans
+ }
+ num_entities = len(entity_spans_dict)
+ num_negatives = int(
+ (len(tokens) + num_entities * 10) * negative_multiplier)
+ num_negatives = min(num_negatives, len(tokens) * (len(tokens) + 1) // 2)
+ min_words = 1
+ max_words = len(tokens)
+ total_entities += len(entity_spans)
+
+ spans = []
+ if self.is_train:
+ is_token_entity_prefix = [0] * (len(tokens) + 1)
+ for start, end, _ in entity_spans:
+ for i in range(start, end + 1):
+ is_token_entity_prefix[i + 1] = 1
+ for i in range(len(tokens)):
+ is_token_entity_prefix[i + 1] += is_token_entity_prefix[i]
+
+ negative_spans = []
+ negative_spans_probs = []
+ for n_words in range(min_words, max_words + 1):
+ for i in range(len(tokens) - n_words + 1):
+ j = i + n_words - 1
+ ent_type = entity_spans_dict.get((i, j), "O")
+ if not self.is_train or ent_type != "O":
+ spans.append((i, j, "mask", ent_type))
+ else:
+ negative_spans.append((i, j, "mask", ent_type))
+ intersection_size = (is_token_entity_prefix[j + 1] -
+ is_token_entity_prefix[i] + 1) / (
+ j + 1 - i)
+ negative_spans_probs.append(math.e**intersection_size)
+
+ if negative_spans and num_negatives > 0:
+ negative_spans_probs = np.array(negative_spans_probs) / np.sum(
+ negative_spans_probs)
+ negative_span_indices = np.random.choice(
+ len(negative_spans),
+ num_negatives,
+ replace=True,
+ p=negative_spans_probs)
+ spans.extend([negative_spans[x] for x in negative_span_indices])
+ else:
+ for n_words in range(min_words, max_words + 1):
+ for i in range(len(tokens) - n_words + 1):
+ j = i + n_words - 1
+ ent_type = entity_spans_dict.get((i, j), "O")
+ spans.append((i, j, "mask", ent_type))
+
+ for instance in self.process_word_list_and_spans_to_inputs(
+ sid, tokens, spans):
+ self.data.append(instance)
+ print(f"{total_missed_entities}/{total_entities} are ignored due to length")
+ print(f"Total {self.__len__()} instances")
+
+ def __len__(self):
+ return len(self.data)
+
+ def __getitem__(self, idx):
+ return self.data[idx]
+
+
+if __name__ == "__main__":
+ path_to_data_folder = sys.argv[1]
+ dataset_name = sys.argv[2]
+ train_file = sys.argv[3]
+ dataset = os.path.join(path_to_data_folder, dataset_name)
+ test_file = "test"
+ _tokenizer = transformers.AutoTokenizer.from_pretrained("bert-base-uncased")
+ entity_map = json.load(open(os.path.join(dataset, "entity_map.json")))
+ _label_to_entity_type_index = {
+ k: i for i, k in enumerate(list(entity_map.keys()))
+ }
+ train_ds = NERDataset(
+ words_path=os.path.join(dataset, train_file + ".words"),
+ labels_path=os.path.join(dataset, train_file + ".ner"),
+ tokenizer=_tokenizer,
+ is_train=True,
+ ablation_not_mask=False,
+ ablation_no_brackets=False,
+ ablation_span_type_together=False,
+ label_to_entity_type_index=_label_to_entity_type_index)
+ eval_ds = NERDataset(
+ words_path=os.path.join(dataset, test_file + ".words"),
+ labels_path=os.path.join(dataset, test_file + ".ner"),
+ tokenizer=_tokenizer,
+ is_train=False,
+ ablation_not_mask=False,
+ ablation_no_brackets=False,
+ ablation_span_type_together=False,
+ label_to_entity_type_index=_label_to_entity_type_index)
+ train_ds.prepare(negative_multiplier=3)
+ train_data = train_ds.collate_fn(train_ds.data)
+ eval_ds.prepare(negative_multiplier=3)
+ eval_data = eval_ds.collate_fn(eval_ds.data)
+
+ def file_based_convert_examples_to_features(examples, output_file):
+ """Convert a set of `InputExample`s to a TFRecord file."""
+ tf.io.gfile.makedirs(os.path.dirname(output_file))
+ writer = tf.io.TFRecordWriter(output_file)
+
+ for ex_index in range(len(examples["input_ids"])):
+ if ex_index % 10000 == 0:
+ print(f"Writing example {ex_index} of {len(examples['input_ids'])}")
+ print(examples["input_ids"][ex_index])
+
+ def create_int_feature(values):
+ f = tf.train.Feature(int64_list=tf.train.Int64List(value=list(values)))
+ return f
+
+ features = collections.OrderedDict()
+ features["input_ids"] = create_int_feature(
+ examples["input_ids"][ex_index])
+ features["input_mask"] = create_int_feature(
+ examples["attention_mask"][ex_index])
+ features["segment_ids"] = create_int_feature(
+ [0] * len(examples["attention_mask"][ex_index]))
+ features["is_entity_token_pos"] = create_int_feature(
+ [examples["is_entity_token_pos"][ex_index]])
+ features["entity_type_token_pos"] = create_int_feature(
+ [examples["entity_type_token_pos"][ex_index]])
+ features["is_entity_label"] = create_int_feature(
+ [examples["is_entity_label"][ex_index]])
+ features["entity_type_label"] = create_int_feature(
+ [examples["entity_type_label"][ex_index]])
+ features["example_id"] = create_int_feature([examples["id"][ex_index]])
+ features["sentence_id"] = create_int_feature(
+ [examples["sentence_id"][ex_index]])
+ features["span_start"] = create_int_feature(
+ [examples["span_start"][ex_index]])
+ features["span_end"] = create_int_feature(
+ [examples["span_end"][ex_index]])
+ tf_example = tf.train.Example(
+ features=tf.train.Features(feature=features))
+ writer.write(tf_example.SerializeToString())
+ writer.close()
+
+ file_based_convert_examples_to_features(
+ train_data, f"{dataset_name}_{train_file}.tf_record")
+ file_based_convert_examples_to_features(
+ eval_data, f"{dataset_name}_{test_file}.tf_record")
diff --git a/official/projects/labse/README.md b/official/projects/labse/README.md
index bb6abcca16d..7d4f8bc8330 100644
--- a/official/projects/labse/README.md
+++ b/official/projects/labse/README.md
@@ -71,7 +71,7 @@ pretraining:
TPU=local
VOCAB=???
INIT_CHECKPOINT=???
-PARAMS="task.train_data.input_data=/path/to/train/data"
+PARAMS="task.train_data.input_path=/path/to/train/data"
PARAMS="${PARAMS},task.train_data.vocab_file=${VOCAB}"
PARAMS="${PARAMS},task.validation_data.input_path=/path/to/validation/data"
PARAMS="${PARAMS},task.validation_data.vocab_file=${VOCAB}"
diff --git a/official/projects/labse/config_labse.py b/official/projects/labse/config_labse.py
index 4dba0e32a03..ddb1d9f4555 100644
--- a/official/projects/labse/config_labse.py
+++ b/official/projects/labse/config_labse.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -29,16 +29,27 @@
@dataclasses.dataclass
class LaBSEOptimizationConfig(optimization.OptimizationConfig):
"""Bert optimization config."""
- optimizer: optimization.OptimizerConfig = optimization.OptimizerConfig(
- type="adamw", adamw=AdamWeightDecay())
- learning_rate: optimization.LrConfig = optimization.LrConfig(
- type="polynomial",
- polynomial=PolynomialLr(
- initial_learning_rate=1e-4,
- decay_steps=1000000,
- end_learning_rate=0.0))
- warmup: optimization.WarmupConfig = optimization.WarmupConfig(
- type="polynomial", polynomial=PolynomialWarmupConfig(warmup_steps=10000))
+ optimizer: optimization.OptimizerConfig = dataclasses.field(
+ default_factory=lambda: optimization.OptimizerConfig( # pylint: disable=g-long-lambda
+ type="adamw", adamw=AdamWeightDecay()
+ )
+ )
+ learning_rate: optimization.LrConfig = dataclasses.field(
+ default_factory=lambda: optimization.LrConfig( # pylint: disable=g-long-lambda
+ type="polynomial",
+ polynomial=PolynomialLr(
+ initial_learning_rate=1e-4,
+ decay_steps=1000000,
+ end_learning_rate=0.0,
+ ),
+ )
+ )
+ warmup: optimization.WarmupConfig = dataclasses.field(
+ default_factory=lambda: optimization.WarmupConfig( # pylint: disable=g-long-lambda
+ type="polynomial",
+ polynomial=PolynomialWarmupConfig(warmup_steps=10000),
+ )
+ )
@exp_factory.register_config_factory("labse/train")
diff --git a/official/projects/labse/export_tfhub.py b/official/projects/labse/export_tfhub.py
index 6adb53c7984..b0b28c91318 100644
--- a/official/projects/labse/export_tfhub.py
+++ b/official/projects/labse/export_tfhub.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -35,7 +35,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import bert_models
from official.legacy.bert import configs
@@ -96,7 +96,7 @@ def _get_do_lower_case(do_lower_case, vocab_file):
def create_labse_model(bert_tfhub_module: Text,
bert_config: configs.BertConfig,
- normalize: bool) -> tf.keras.Model:
+ normalize: bool) -> tf_keras.Model:
"""Creates a LaBSE keras core model from BERT configuration.
Args:
@@ -127,7 +127,7 @@ def create_labse_model(bert_tfhub_module: Text,
def export_labse_model(bert_tfhub_module: Text, bert_config: configs.BertConfig,
model_checkpoint_path: Text, hub_destination: Text,
vocab_file: Text, do_lower_case: bool, normalize: bool):
- """Restores a tf.keras.Model and saves for TF-Hub."""
+ """Restores a tf_keras.Model and saves for TF-Hub."""
core_model, encoder = create_labse_model(
bert_tfhub_module, bert_config, normalize)
checkpoint = tf.train.Checkpoint(encoder=encoder)
diff --git a/official/projects/labse/export_tfhub_test.py b/official/projects/labse/export_tfhub_test.py
index f45c200441c..550f8fd89c6 100644
--- a/official/projects/labse/export_tfhub_test.py
+++ b/official/projects/labse/export_tfhub_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,8 @@
import os
-# Import libraries
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_hub as hub
from official.legacy.bert import configs
from official.projects.labse import export_tfhub
@@ -97,9 +96,9 @@ def _dropout_mean_stddev(training, num_runs=20):
self.assertGreater(_dropout_mean_stddev(training=True), 1e-3)
# Test propagation of seq_length in shape inference.
- input_word_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_mask = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
- input_type_ids = tf.keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_word_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_mask = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
+ input_type_ids = tf_keras.layers.Input(shape=(seq_length,), dtype=tf.int32)
outputs = hub_layer([input_word_ids, input_mask, input_type_ids])
self.assertEqual(outputs["pooled_output"].shape.as_list(),
[None, hidden_size])
diff --git a/official/projects/labse/train.py b/official/projects/labse/train.py
index 7e9cc7d11c3..1eb709b0cfd 100644
--- a/official/projects/labse/train.py
+++ b/official/projects/labse/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/longformer/longformer.py b/official/projects/longformer/longformer.py
index 76a491ccb5f..0bd11e50c68 100644
--- a/official/projects/longformer/longformer.py
+++ b/official/projects/longformer/longformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import dataclasses
from typing import List
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.modeling.hyperparams import base_config
@@ -61,7 +61,7 @@ def get_encoder(encoder_cfg: LongformerEncoderConfig):
attention_dropout=encoder_cfg.attention_dropout_rate,
max_sequence_length=encoder_cfg.max_position_embeddings,
type_vocab_size=encoder_cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
output_range=encoder_cfg.output_range,
embedding_width=encoder_cfg.embedding_size,
diff --git a/official/projects/longformer/longformer_attention.py b/official/projects/longformer/longformer_attention.py
index 3f3980e81a4..bcfe6f4f5c1 100644
--- a/official/projects/longformer/longformer_attention.py
+++ b/official/projects/longformer/longformer_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import string
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.tf_utils import get_shape_list
@@ -104,8 +104,8 @@ def _get_output_shape(output_rank, known_last_dims):
return [None] * (output_rank - len(known_last_dims)) + list(known_last_dims)
-@tf.keras.utils.register_keras_serializable(package="Text")
-class LongformerAttention(tf.keras.layers.MultiHeadAttention):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class LongformerAttention(tf_keras.layers.MultiHeadAttention):
"""LongformerAttention.
Args:
@@ -170,14 +170,14 @@ def _build_from_signature(self, query, value, key=None):
free_dims = self._query_shape.rank - 1
einsum_equation, bias_axes, output_rank = _build_proj_equation(
free_dims, bound_dims=1, output_dims=2)
- self._query_dense = tf.keras.layers.experimental.EinsumDense(
+ self._query_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1,
[self._num_heads, self._key_dim]),
bias_axes=bias_axes if self._use_bias else None,
name="query",
**common_kwargs)
- self._global_query_dense = tf.keras.layers.experimental.EinsumDense(
+ self._global_query_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1,
[self._num_heads, self._key_dim]),
@@ -186,14 +186,14 @@ def _build_from_signature(self, query, value, key=None):
**common_kwargs)
einsum_equation, bias_axes, output_rank = _build_proj_equation(
self._key_shape.rank - 1, bound_dims=1, output_dims=2)
- self._key_dense = tf.keras.layers.experimental.EinsumDense(
+ self._key_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1,
[self._num_heads, self._key_dim]),
bias_axes=bias_axes if self._use_bias else None,
name="key",
**common_kwargs)
- self._global_key_dense = tf.keras.layers.experimental.EinsumDense(
+ self._global_key_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1,
[self._num_heads, self._key_dim]),
@@ -202,14 +202,14 @@ def _build_from_signature(self, query, value, key=None):
**common_kwargs)
einsum_equation, bias_axes, output_rank = _build_proj_equation(
self._value_shape.rank - 1, bound_dims=1, output_dims=2)
- self._value_dense = tf.keras.layers.experimental.EinsumDense(
+ self._value_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1,
[self._num_heads, self._value_dim]),
bias_axes=bias_axes if self._use_bias else None,
name="value",
**common_kwargs)
- self._global_value_dense = tf.keras.layers.experimental.EinsumDense(
+ self._global_value_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=_get_output_shape(output_rank - 1,
[self._num_heads, self._value_dim]),
@@ -221,13 +221,13 @@ def _build_from_signature(self, query, value, key=None):
# These computations could be wrapped into the keras attention layer once
# it support mult-head einsum computations.
self._build_attention(output_rank)
- self._global_dropout_layer = tf.keras.layers.Dropout(rate=self._dropout)
+ self._global_dropout_layer = tf_keras.layers.Dropout(rate=self._dropout)
# self._output_dense = self._make_output_dense(
# free_dims, common_kwargs, "attention_output")
- self._output_dense = tf.keras.layers.Dense(
+ self._output_dense = tf_keras.layers.Dense(
units=self._num_heads * self._key_dim, name="dense", **common_kwargs)
- def call(self,
+ def call(self, # pyrefly: ignore[bad-override]
hidden_states,
attention_mask=None,
is_index_masked=None,
@@ -332,13 +332,13 @@ def call(self,
# self.one_sided_attn_window_size * 2 + 1]
if self.global_attention_size > 0:
masked_index = tf.tile(
- is_index_masked[:, :, None, None],
+ is_index_masked[:, :, None, None], # pyrefly: ignore[unsupported-operation]
(1, 1, self._num_heads, self._one_sided_attn_window_size * 2 +
max_num_global_attn_indices + 1),
)
else:
masked_index = tf.tile(
- is_index_masked[:, :, None, None],
+ is_index_masked[:, :, None, None], # pyrefly: ignore[unsupported-operation]
(1, 1, self._num_heads, self._one_sided_attn_window_size * 2 + 1),
)
@@ -412,13 +412,13 @@ def call(self,
# global attn
if self.global_attention_size > 0:
masked_global_attn_index = tf.tile(
- is_index_global_attn[:, :, None, None],
+ is_index_global_attn[:, :, None, None], # pyrefly: ignore[unsupported-operation]
(1, 1, self._num_heads, self._one_sided_attn_window_size * 2 +
max_num_global_attn_indices + 1),
)
else:
masked_global_attn_index = tf.tile(
- is_index_global_attn[:, :, None, None],
+ is_index_global_attn[:, :, None, None], # pyrefly: ignore[unsupported-operation]
(1, 1, self._num_heads, self._one_sided_attn_window_size * 2 + 1),
)
diff --git a/official/projects/longformer/longformer_attention_test.py b/official/projects/longformer/longformer_attention_test.py
index 9211987e62a..462f8952a11 100644
--- a/official/projects/longformer/longformer_attention_test.py
+++ b/official/projects/longformer/longformer_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for official.nlp.projects.longformer.longformer_attention."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.tf_utils import get_shape_list
from official.projects.longformer import longformer_attention
diff --git a/official/projects/longformer/longformer_encoder.py b/official/projects/longformer/longformer_encoder.py
index 65a9c1de13a..6a1a06c32c6 100644
--- a/official/projects/longformer/longformer_encoder.py
+++ b/official/projects/longformer/longformer_encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,18 +19,18 @@
from typing import Any, Callable, List, Optional, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.tf_utils import get_shape_list
from official.nlp.modeling import layers
from official.projects.longformer.longformer_encoder_block import LongformerEncoderBlock
-_Initializer = Union[str, tf.keras.initializers.Initializer]
-_approx_gelu = lambda x: tf.keras.activations.gelu(x, approximate=True)
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
-class LongformerEncoder(tf.keras.layers.Layer):
+class LongformerEncoder(tf_keras.layers.Layer):
"""LongformerEncoder.
Args:
@@ -86,11 +86,11 @@ def __init__(
inner_activation: Callable[..., Any] = _approx_gelu,
output_dropout: float = 0.1,
attention_dropout: float = 0.1,
- initializer: _Initializer = tf.keras.initializers.TruncatedNormal(
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
stddev=0.02),
output_range: Optional[int] = None,
embedding_width: Optional[int] = None,
- embedding_layer: Optional[tf.keras.layers.Layer] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
norm_first: bool = False,
**kwargs):
super().__init__(**kwargs)
@@ -99,8 +99,8 @@ def __init__(
self._global_attention_size = global_attention_size
self._pad_token_id = pad_token_id
- activation = tf.keras.activations.get(inner_activation)
- initializer = tf.keras.initializers.get(initializer)
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
if embedding_width is None:
embedding_width = hidden_size
@@ -126,17 +126,17 @@ def __init__(
use_one_hot=True,
name='type_embeddings')
- self._embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
- self._embedding_dropout = tf.keras.layers.Dropout(
+ self._embedding_dropout = tf_keras.layers.Dropout(
rate=output_dropout, name='embedding_dropout')
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
self._embedding_projection = None
if embedding_width != hidden_size:
- self._embedding_projection = tf.keras.layers.experimental.EinsumDense(
+ self._embedding_projection = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
@@ -152,7 +152,7 @@ def __init__(
num_attention_heads=num_attention_heads,
inner_dim=inner_dim,
inner_activation=inner_activation,
- attention_window=attention_window[i],
+ attention_window=attention_window[i], # pyrefly: ignore[bad-index]
layer_id=i,
output_dropout=output_dropout,
attention_dropout=attention_dropout,
@@ -162,7 +162,7 @@ def __init__(
name=f'transformer/layer_{i}')
self._transformer_layers.append(layer)
- self._pooler_layer = tf.keras.layers.Dense(
+ self._pooler_layer = tf_keras.layers.Dense(
units=hidden_size,
activation='tanh',
kernel_initializer=initializer,
@@ -176,10 +176,10 @@ def __init__(
'max_sequence_length': max_sequence_length,
'type_vocab_size': type_vocab_size,
'inner_dim': inner_dim,
- 'inner_activation': tf.keras.activations.serialize(activation),
+ 'inner_activation': tf_keras.activations.serialize(activation),
'output_dropout': output_dropout,
'attention_dropout': attention_dropout,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_keras.initializers.serialize(initializer),
'output_range': output_range,
'embedding_width': embedding_width,
'embedding_layer': embedding_layer,
@@ -189,9 +189,9 @@ def __init__(
'pad_token_id': pad_token_id,
}
self.inputs = dict(
- input_word_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32))
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32))
def call(self, inputs):
word_embeddings = None
@@ -318,7 +318,7 @@ def _pad_to_window_size(
pad_token_id,
):
# padding
- attention_window = max(self._attention_window)
+ attention_window = max(self._attention_window) # pyrefly: ignore[bad-argument-type]
assert (attention_window %
2 == 0), ('`attention_window` should be an even value.'
diff --git a/official/projects/longformer/longformer_encoder_block.py b/official/projects/longformer/longformer_encoder_block.py
index 84999efaf3a..c5ce0b29bfd 100644
--- a/official/projects/longformer/longformer_encoder_block.py
+++ b/official/projects/longformer/longformer_encoder_block.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,12 +14,12 @@
"""Longformer attention layer. Modified From huggingface/transformers."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.longformer.longformer_attention import LongformerAttention
-@tf.keras.utils.register_keras_serializable(package="Text")
-class LongformerEncoderBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class LongformerEncoderBlock(tf_keras.layers.Layer):
"""LongformerEncoderBlock.
Args:
@@ -96,19 +96,19 @@ def __init__(
self._output_dropout = output_dropout
self._output_dropout_rate = output_dropout
self._output_range = output_range
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
self._inner_dropout = inner_dropout
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
+ self._attention_initializer = tf_keras.initializers.get(
attention_initializer)
else:
self._attention_initializer = self._kernel_initializer
@@ -154,39 +154,39 @@ def build(self, input_shape):
name="self_attention",
**common_kwargs)
# TFLongformerSelfOutput.dropout
- self._attention_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ self._attention_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
# TFLongformerSelfOutput.Layernorm
self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype=tf.float32))
# TFLongformerIntermediate
# TFLongformerIntermediate.dense
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, self._inner_dim),
bias_axes="d",
kernel_initializer=self._kernel_initializer,
name="intermediate",
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
# TODO(b/154538392): Investigate this.
policy = tf.float32
# TFLongformerIntermediate.intermediate_act_fn
- self._intermediate_activation_layer = tf.keras.layers.Activation(
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
self._inner_activation, dtype=policy)
- self._inner_dropout_layer = tf.keras.layers.Dropout(
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
rate=self._inner_dropout)
# TFLongformerOutput
# TFLongformerOutput.dense
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._output_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, hidden_size),
bias_axes="d",
@@ -194,10 +194,10 @@ def build(self, input_shape):
kernel_initializer=self._kernel_initializer,
**common_kwargs)
# TFLongformerOutput.dropout
- self._output_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
# TFLongformerOutput.layernorm
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -220,19 +220,19 @@ def get_config(self):
"output_range":
self._output_range,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
+ tf_keras.constraints.serialize(self._bias_constraint),
"use_bias":
self._use_bias,
"norm_first":
@@ -242,7 +242,7 @@ def get_config(self):
"inner_dropout":
self._inner_dropout,
"attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer),
+ tf_keras.initializers.serialize(self._attention_initializer),
"attention_axes":
self._attention_axes,
}
@@ -314,9 +314,9 @@ def call(self, inputs):
is_index_global_attn=is_index_global_attn,
)
# TFLongformerAttention.TFLongformerSelfOutput.* - {.dense}
- attention_output = self._attention_dropout(attention_output)
+ attention_output = self._attention_dropout(attention_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- attention_output = source_tensor + attention_output
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self._attention_layer_norm(target_tensor +
attention_output)
@@ -329,10 +329,10 @@ def call(self, inputs):
inner_output = self._inner_dropout_layer(inner_output)
# TFLongformerOutput
layer_output = self._output_dense(inner_output)
- layer_output = self._output_dropout(layer_output)
+ layer_output = self._output_dropout(layer_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- return source_attention_output + layer_output
+ return source_attention_output + layer_output # pyrefly: ignore[unbound-name]
# During mixed precision training, layer norm output is always fp32 for now.
# Casts fp32 for the subsequent add.
diff --git a/official/projects/longformer/longformer_encoder_test.py b/official/projects/longformer/longformer_encoder_test.py
index cf24d7c926b..49a478b353f 100644
--- a/official/projects/longformer/longformer_encoder_test.py
+++ b/official/projects/longformer/longformer_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from official.projects.longformer.longformer_encoder import LongformerEncoder
diff --git a/official/projects/longformer/longformer_experiments.py b/official/projects/longformer/longformer_experiments.py
index e93672806d8..3c52e84c74f 100644
--- a/official/projects/longformer/longformer_experiments.py
+++ b/official/projects/longformer/longformer_experiments.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/longformer/train.py b/official/projects/longformer/train.py
index 5486c1902d8..c2041f168a2 100644
--- a/official/projects/longformer/train.py
+++ b/official/projects/longformer/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/longformer/utils/convert_pretrained_pytorch_checkpoint_to_tf.py b/official/projects/longformer/utils/convert_pretrained_pytorch_checkpoint_to_tf.py
index 38fcc84e077..98c52fe04a0 100644
--- a/official/projects/longformer/utils/convert_pretrained_pytorch_checkpoint_to_tf.py
+++ b/official/projects/longformer/utils/convert_pretrained_pytorch_checkpoint_to_tf.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl import app
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import transformers
from official.modeling import tf_utils
@@ -54,7 +54,7 @@ def _create_longformer_model():
attention_dropout=encoder_cfg.attention_dropout_rate,
max_sequence_length=encoder_cfg.max_position_embeddings,
type_vocab_size=encoder_cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
output_range=encoder_cfg.output_range,
embedding_width=encoder_cfg.embedding_size,
diff --git a/official/projects/longformer/utils/longformer_tokenizer_to_tfrecord.py b/official/projects/longformer/utils/longformer_tokenizer_to_tfrecord.py
index 9fc85a391ed..10bf9dfd7ca 100644
--- a/official/projects/longformer/utils/longformer_tokenizer_to_tfrecord.py
+++ b/official/projects/longformer/utils/longformer_tokenizer_to_tfrecord.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
import datasets
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import transformers
pretrained_lm = "allenai/longformer-base-4096"
diff --git a/official/projects/lra/README.md b/official/projects/lra/README.md
new file mode 100644
index 00000000000..28bd78555c9
--- /dev/null
+++ b/official/projects/lra/README.md
@@ -0,0 +1,67 @@
+# Long Range Arena
+
+This repository contains TensorFlow 2.x implementation for Long Range Arena
+Tasks, including baseline for MEGA, Transformer, and Linformer. The codebase is
+adapted from (https://github.com/google-research/long-range-arena).
+
+## Training on LRA Tasks
+Example job script to train Transformer on ListOPs task:
+
+```bash
+TRAIN_DATA=task.train_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_train.tf_record,task.validation_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_eval.tf_record
+
+PYTHONPATH=[/PATH/TO/MODEL_GARDEN] \
+ python3 train.py \
+ --experiment=transformer/lra_listops \
+ --config_file=../experiments/lra_listops.yaml \
+ --params_override="${TRAIN_DATA},runtime.distribution_strategy=tpu" \
+ --tpu=local \
+ --model_dir=[OUTPUT_DIR] \
+ --mode=train_and_eval
+```
+
+To train Linformer on ListOPs task:
+
+```bash
+TRAIN_DATA=task.train_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_train.tf_record,task.validation_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_eval.tf_record
+
+PYTHONPATH=[/PATH/TO/MODEL_GARDEN] \
+ python3 train.py \
+ --experiment=linformer/lra_listops \
+ --config_file=../experiments/lra_listops_linformer.yaml \
+ --params_override="${TRAIN_DATA},runtime.distribution_strategy=tpu" \
+ --tpu=local \
+ --model_dir=[OUTPUT_DIR] \
+ --mode=train_and_eval
+```
+
+To train MEGA on Text task (Reproduced Acc = 87.55):
+
+```bash
+TRAIN_DATA=task.train_data.input_path=gs://model-garden-ucsd-zihan/lra_imdb_4096_train.tf_record,task.validation_data.input_path=gs://model-garden-ucsd-zihan/lra_imdb_4096_eval.tf_record
+
+PYTHONPATH=[/PATH/TO/MODEL_GARDEN] \
+ python3 train.py \
+ --experiment=mega/lra_imdb \
+ --config_file=../experiments/lra_imdb_mega.yaml \
+ --params_override="${TRAIN_DATA},runtime.distribution_strategy=tpu" \
+ --tpu=local \
+ --model_dir=[OUTPUT_DIR] \
+ --mode=train_and_eval
+```
+
+## Data Paths and Experiment Configs
+Dataset Paths are listed below:
+
+| | Path |
+|------------|-------------------------------------------------------------------------|
+| ListOps | gs://model-garden-ucsd-zihan/lra_listops_[train/eval/test].tf_record |
+| IMDB | gs://model-garden-ucsd-zihan/lra_imdb_[train/eval/test].tf_record |
+| IMDB-4096 | gs://model-garden-ucsd-zihan/lra_imdb_4096_[train/eval/test].tf_record |
+| AAN | gs://model-garden-ucsd-zihan/lra_aan_[train/eval/test].tf_record |
+| CIFAR10 | gs://model-garden-ucsd-zihan/lra_cifar_[train/eval/test].tf_record |
+| Pathfinder | gs://model-garden-ucsd-zihan/lra_pathfinder_[train/eval/test].tf_record |
+
+Experiment Configs can be found in the `experiments` subfolder.
+
+
diff --git a/official/projects/lra/experiments/lra_aan.yaml b/official/projects/lra/experiments/lra_aan.yaml
new file mode 100644
index 00000000000..43b8d9dd000
--- /dev/null
+++ b/official/projects/lra/experiments/lra_aan.yaml
@@ -0,0 +1,59 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ max_seq_length: 4000
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 128
+ hidden_activation: gelu
+ hidden_size: 128
+ initializer_range: 0.02
+ intermediate_size: 512
+ max_position_embeddings: 4000
+ num_attention_heads: 4
+ num_layers: 4
+ type_vocab_size: 2
+ vocab_size: 258
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: true
+ seq_length: 4000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: false
+ seq_length: 4000
+trainer:
+ checkpoint_interval: 500
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 5000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0005
+ power: 0.5
+ type: polynomial
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.1
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 800
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 5000
+ validation_interval: 500
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_aan_linformer.yaml b/official/projects/lra/experiments/lra_aan_linformer.yaml
new file mode 100644
index 00000000000..2e6902653a6
--- /dev/null
+++ b/official/projects/lra/experiments/lra_aan_linformer.yaml
@@ -0,0 +1,60 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ max_seq_length: 4000
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 128
+ hidden_activation: gelu
+ hidden_size: 128
+ initializer_range: 0.02
+ intermediate_size: 512
+ max_position_embeddings: 4000
+ num_attention_heads: 4
+ num_layers: 4
+ low_rank_features: 32
+ type_vocab_size: 2
+ vocab_size: 258
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: true
+ seq_length: 4000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: false
+ seq_length: 4000
+trainer:
+ checkpoint_interval: 500
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 5000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0005
+ power: 0.5
+ type: polynomial
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.1
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 800
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 5000
+ validation_interval: 500
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_aan_mega.yaml b/official/projects/lra/experiments/lra_aan_mega.yaml
new file mode 100644
index 00000000000..158388bc470
--- /dev/null
+++ b/official/projects/lra/experiments/lra_aan_mega.yaml
@@ -0,0 +1,69 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ max_seq_length: 4000
+ use_encoder_pooler: true
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 128
+ hidden_activation: gelu
+ initializer_range: 0.02
+ intermediate_size: 256
+ max_position_embeddings: 4000
+ num_attention_heads: 1
+ num_layers: 6
+ type_vocab_size: 2
+ vocab_size: 258
+ zdim: 64
+ hdim: 256
+ ndim: 16
+ activation: 'silu'
+ bidirectional: true
+ norm_first: true
+ dropout: 0.0
+ hidden_dropout: 0.1
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: true
+ seq_length: 4000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: false
+ seq_length: 4000
+trainer:
+ checkpoint_interval: 500
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 20000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.003
+ power: 0.5
+ type: polynomial
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.98
+ weight_decay_rate: 0.04
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 2000
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 20000
+ validation_interval: 500
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_cifar.yaml b/official/projects/lra/experiments/lra_cifar.yaml
new file mode 100644
index 00000000000..4b4b9c8699d
--- /dev/null
+++ b/official/projects/lra/experiments/lra_cifar.yaml
@@ -0,0 +1,61 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 10
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.2
+ dropout_rate: 0.3
+ embedding_size: 128
+ hidden_activation: gelu
+ hidden_size: 64
+ initializer_range: 0.02
+ intermediate_size: 128
+ max_position_embeddings: 1024
+ num_attention_heads: 8
+ num_layers: 1
+ type_vocab_size: 2
+ vocab_size: 257
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 256
+ input_path: TODO
+ is_training: true
+ seq_length: 1024
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 256
+ input_path: TODO
+ is_training: false
+ seq_length: 1024
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 35200
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0005
+ power: 1
+ type: polynomial
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.98
+ weight_decay_rate: 0.0
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 176
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 35200
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_cifar_linformer.yaml b/official/projects/lra/experiments/lra_cifar_linformer.yaml
new file mode 100644
index 00000000000..adc170c9bc2
--- /dev/null
+++ b/official/projects/lra/experiments/lra_cifar_linformer.yaml
@@ -0,0 +1,62 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 10
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.2
+ dropout_rate: 0.3
+ embedding_size: 128
+ hidden_activation: gelu
+ hidden_size: 64
+ initializer_range: 0.02
+ intermediate_size: 128
+ max_position_embeddings: 1024
+ num_attention_heads: 8
+ num_layers: 1
+ low_rank_features: 32
+ type_vocab_size: 2
+ vocab_size: 257
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 256
+ input_path: TODO
+ is_training: true
+ seq_length: 1024
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 256
+ input_path: TODO
+ is_training: false
+ seq_length: 1024
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 35200
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0005
+ power: 1
+ type: polynomial
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.98
+ weight_decay_rate: 0.0
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 176
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 35200
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_cifar_mega.yaml b/official/projects/lra/experiments/lra_cifar_mega.yaml
new file mode 100644
index 00000000000..238477197d3
--- /dev/null
+++ b/official/projects/lra/experiments/lra_cifar_mega.yaml
@@ -0,0 +1,70 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 10
+ max_seq_length: 1024
+ use_encoder_pooler: true
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 160
+ hidden_activation: gelu
+ initializer_range: 0.02
+ intermediate_size: 320
+ max_position_embeddings: 1024
+ num_attention_heads: 1
+ num_layers: 8
+ type_vocab_size: 2
+ vocab_size: 100
+ zdim: 96
+ hdim: 320
+ ndim: 16
+ activation: 'silu'
+ bidirectional: true
+ norm_first: true
+ dropout: 0.0
+ hidden_dropout: 0.0
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: true
+ seq_length: 1024
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: false
+ seq_length: 1024
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 156250
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.001
+ power: 1
+ type: polynomial
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.98
+ weight_decay_rate: 0.02
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 2000
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 156250
+ validation_interval: 1000
+ validation_steps: 9999999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_imdb.yaml b/official/projects/lra/experiments/lra_imdb.yaml
new file mode 100644
index 00000000000..00017eeb237
--- /dev/null
+++ b/official/projects/lra/experiments/lra_imdb.yaml
@@ -0,0 +1,58 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 256
+ hidden_activation: gelu
+ hidden_size: 256
+ initializer_range: 0.02
+ intermediate_size: 1024
+ max_position_embeddings: 1000
+ num_attention_heads: 4
+ num_layers: 4
+ type_vocab_size: 2
+ vocab_size: 258
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: true
+ seq_length: 1000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: false
+ seq_length: 1000
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 20000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.00005
+ power: 0.5
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 8000
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 20000
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_imdb_linformer.yaml b/official/projects/lra/experiments/lra_imdb_linformer.yaml
new file mode 100644
index 00000000000..9bd0a6356d7
--- /dev/null
+++ b/official/projects/lra/experiments/lra_imdb_linformer.yaml
@@ -0,0 +1,59 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 256
+ hidden_activation: gelu
+ hidden_size: 256
+ initializer_range: 0.02
+ intermediate_size: 1024
+ max_position_embeddings: 1000
+ num_attention_heads: 4
+ num_layers: 4
+ low_rank_features: 32
+ type_vocab_size: 2
+ vocab_size: 258
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: true
+ seq_length: 1000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: TODO
+ is_training: false
+ seq_length: 1000
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 20000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.00005
+ power: 0.5
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 8000
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 20000
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_imdb_mega.yaml b/official/projects/lra/experiments/lra_imdb_mega.yaml
new file mode 100644
index 00000000000..25f1a6277bd
--- /dev/null
+++ b/official/projects/lra/experiments/lra_imdb_mega.yaml
@@ -0,0 +1,66 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ use_encoder_pooler: true
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 128
+ hidden_activation: gelu
+ initializer_range: 0.02
+ intermediate_size: 256
+ max_position_embeddings: 1000
+ num_attention_heads: 1
+ num_layers: 4
+ type_vocab_size: 2
+ vocab_size: 256
+ zdim: 64
+ hdim: 256
+ ndim: 16
+ activation: 'silu'
+ bidirectional: true
+ norm_first: false
+ dropout: 0.1
+ hidden_dropout: 0.1
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: true
+ seq_length: 1000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: false
+ seq_length: 1000
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 25000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.004
+ power: 1
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 10000
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 50000
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_listops.yaml b/official/projects/lra/experiments/lra_listops.yaml
new file mode 100644
index 00000000000..87a3dbf2aac
--- /dev/null
+++ b/official/projects/lra/experiments/lra_listops.yaml
@@ -0,0 +1,58 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 10
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 512
+ hidden_activation: gelu
+ hidden_size: 512
+ initializer_range: 0.02
+ intermediate_size: 1024
+ max_position_embeddings: 2000
+ num_attention_heads: 8
+ num_layers: 4
+ type_vocab_size: 2
+ vocab_size: 100
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: true
+ seq_length: 2000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: false
+ seq_length: 2000
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 5000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.00005
+ power: 0.5
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 1000
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 5000
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_listops_linformer.yaml b/official/projects/lra/experiments/lra_listops_linformer.yaml
new file mode 100644
index 00000000000..fe2361ed409
--- /dev/null
+++ b/official/projects/lra/experiments/lra_listops_linformer.yaml
@@ -0,0 +1,58 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 10
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 512
+ hidden_activation: gelu
+ hidden_size: 512
+ initializer_range: 0.02
+ intermediate_size: 1024
+ max_position_embeddings: 2000
+ num_attention_heads: 8
+ num_layers: 4
+ low_rank_features: 32
+ type_vocab_size: 2
+ vocab_size: 100
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: true
+ seq_length: 2000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: false
+ seq_length: 2000
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 5000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.00005
+ power: 0.5
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 1000
+ type: polynomial
+ summary_interval: 100
+ train_steps: 5000
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_listops_mega.yaml b/official/projects/lra/experiments/lra_listops_mega.yaml
new file mode 100644
index 00000000000..0c297a4459a
--- /dev/null
+++ b/official/projects/lra/experiments/lra_listops_mega.yaml
@@ -0,0 +1,66 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 10
+ use_encoder_pooler: true
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.0
+ dropout_rate: 0.1
+ embedding_size: 80
+ hidden_activation: gelu
+ initializer_range: 0.02
+ intermediate_size: 160
+ max_position_embeddings: 2000
+ num_attention_heads: 1
+ num_layers: 6
+ type_vocab_size: 2
+ vocab_size: 100
+ zdim: 64
+ hdim: 160
+ ndim: 16
+ activation: 'silu'
+ bidirectional: true
+ norm_first: false
+ dropout: 0.1
+ hidden_dropout: 0.0
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: true
+ seq_length: 2000
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 64
+ input_path: TODO
+ is_training: false
+ seq_length: 2000
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 90000
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0001
+ power: 1
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 3000
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 90000
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_pathfinder.yaml b/official/projects/lra/experiments/lra_pathfinder.yaml
new file mode 100644
index 00000000000..be6eac8d953
--- /dev/null
+++ b/official/projects/lra/experiments/lra_pathfinder.yaml
@@ -0,0 +1,62 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ use_encoder_pooler: true
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 128
+ hidden_activation: gelu
+ hidden_size: 64
+ initializer_range: 0.02
+ intermediate_size: 128
+ max_position_embeddings: 1024
+ num_attention_heads: 8
+ num_layers: 1
+ type_vocab_size: 2
+ vocab_size: 257
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 512
+ input_path: TODO
+ is_training: true
+ seq_length: 1024
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 512
+ input_path: TODO
+ is_training: false
+ seq_length: 1024
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 62500
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.001
+ power: 0.5
+ type: polynomial
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.98
+ weight_decay_rate: 0.0
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 313
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 62500
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_pathfinder_linformer.yaml b/official/projects/lra/experiments/lra_pathfinder_linformer.yaml
new file mode 100644
index 00000000000..c14c0ab7d0c
--- /dev/null
+++ b/official/projects/lra/experiments/lra_pathfinder_linformer.yaml
@@ -0,0 +1,62 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ use_encoder_pooler: true
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 256
+ hidden_activation: gelu
+ hidden_size: 128
+ initializer_range: 0.02
+ intermediate_size: 512
+ max_position_embeddings: 1024
+ num_attention_heads: 4
+ num_layers: 2
+ low_rank_features: 128
+ type_vocab_size: 2
+ vocab_size: 257
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 512
+ input_path: TODO
+ is_training: true
+ seq_length: 1024
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 512
+ input_path: TODO
+ is_training: false
+ seq_length: 1024
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 62500
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.0001
+ power: 0.5
+ type: polynomial
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.98
+ weight_decay_rate: 0.0
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 313
+ type: polynomial
+ summary_interval: 100
+ train_steps: 62500
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/experiments/lra_pathfinder_mega.yaml b/official/projects/lra/experiments/lra_pathfinder_mega.yaml
new file mode 100644
index 00000000000..3153fcecdd9
--- /dev/null
+++ b/official/projects/lra/experiments/lra_pathfinder_mega.yaml
@@ -0,0 +1,69 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ use_encoder_pooler: true
+ encoder:
+ type: any
+ any:
+ attention_dropout_rate: 0.1
+ dropout_rate: 0.1
+ embedding_size: 128
+ hidden_activation: gelu
+ initializer_range: 0.02
+ intermediate_size: 256
+ max_position_embeddings: 1024
+ num_attention_heads: 1
+ num_layers: 6
+ type_vocab_size: 2
+ vocab_size: 257
+ zdim: 64
+ hdim: 256
+ ndim: 16
+ activation: 'silu'
+ bidirectional: true
+ norm_first: true
+ dropout: 0.0
+ hidden_dropout: 0.0
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 512
+ input_path: TODO
+ is_training: true
+ seq_length: 1024
+ validation_data:
+ drop_remainder: true
+ global_batch_size: 512
+ input_path: TODO
+ is_training: false
+ seq_length: 1024
+trainer:
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 7200
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 62500
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.001
+ power: 0.5
+ type: polynomial
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.98
+ weight_decay_rate: 0.01
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 313
+ type: polynomial
+ steps_per_loop: 100
+ summary_interval: 100
+ train_steps: 62500
+ validation_interval: 1000
+ validation_steps: 99999
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/lra/exponential_moving_average.py b/official/projects/lra/exponential_moving_average.py
new file mode 100644
index 00000000000..0aabe76b205
--- /dev/null
+++ b/official/projects/lra/exponential_moving_average.py
@@ -0,0 +1,177 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras-based MegaEncoder block layer."""
+
+from typing import Optional
+import tensorflow as tf, tf_keras
+
+
+class MultiHeadEMA(tf_keras.layers.Layer):
+ """Exponential Moving Average Layer.
+
+ See "https://arxiv.org/abs/2209.10655" for more details.
+ """
+
+ def __init__(
+ self, embed_dim, ndim=2, bidirectional=False, truncation=None, **kwargs
+ ):
+ super().__init__(**kwargs)
+
+ self.embed_dim = embed_dim
+ self.ndim = ndim
+ self.bidirectional = bidirectional
+ self.truncation = truncation
+ self.scale = tf.math.sqrt(1.0 / self.ndim)
+
+ self.kernel_dim = 2 * embed_dim if self.bidirectional else embed_dim
+
+ self._kernel = None
+ self._coeffs = None
+
+ def build(self, input_shape):
+ self.damping_factor = self.add_weight(
+ shape=(self.kernel_dim, self.ndim, 1),
+ initializer="random_normal",
+ trainable=True,
+ name="damping_factor",
+ dtype=tf.float32,
+ )
+ self.decay_factor = self.add_weight(
+ shape=(self.kernel_dim, self.ndim, 1),
+ initializer="random_normal",
+ trainable=True,
+ name="decay_factor",
+ dtype=tf.float32,
+ )
+ self.ema_expansion_matrix = self.add_weight(
+ shape=(self.kernel_dim, self.ndim, 1),
+ initializer="random_normal",
+ trainable=True,
+ name="ema_expansion_matrix",
+ dtype=tf.float32,
+ )
+ self.kernel_projection_matrix = self.add_weight(
+ shape=(self.kernel_dim, self.ndim),
+ initializer="random_normal",
+ trainable=True,
+ name="kernel_projection_matrix",
+ dtype=tf.float32,
+ )
+ self.residual_weight = self.add_weight(
+ shape=(self.embed_dim,),
+ initializer="ones",
+ trainable=True,
+ name="residual_weight",
+ dtype=tf.float32,
+ )
+
+ super().build(input_shape)
+
+ def _calc_coeffs(self):
+ self._coeffs = None
+ # D x N x 1
+ damping_factor = tf.math.sigmoid(self.damping_factor)
+ decay_factor = tf.math.sigmoid(self.decay_factor)
+ previous_timestep_weight = 1.0 - damping_factor * decay_factor
+ return damping_factor, previous_timestep_weight
+
+ def _compute_kernel(self, length: int):
+ self._kernel = None
+ # D x N x 1
+ damping_factor, previous_timestep_weight = self._calc_coeffs()
+ # D x N x L
+ vander = tf.cast(
+ tf.reshape(tf.range(length), shape=(1, 1, length)),
+ dtype=damping_factor.dtype,
+ ) * tf.math.log(previous_timestep_weight)
+ kernel = (damping_factor * self.ema_expansion_matrix) * tf.math.exp(vander)
+ # D x L
+ return tf.einsum(
+ "dnl,dn->dl", kernel, self.kernel_projection_matrix * self.scale
+ )
+
+ def coeffs(self):
+ if self.training:
+ return self._calc_coeffs()
+ else:
+ if self._coeffs is None:
+ self._coeffs = self._calc_coeffs()
+ return self._coeffs
+
+ def kernel(self, length: int):
+ assert self.truncation is None, "WEIRD!"
+ kernel_size = (
+ length if self.truncation is None else min(self.truncation, length)
+ )
+ return self._compute_kernel(kernel_size)
+
+ def call(self, x, padding_mask: Optional[tf.Tensor] = None) -> tf.Tensor:
+ """Input shape: Time x Batch x Channel.
+
+ Args:
+ x: Tensor input.
+ padding_mask (ByteTensor, optional): mask to exclude keys that are pads,
+ of shape `(batch, src_len)`, where padding elements are indicated by
+ 1s.
+ Returns:
+ transformed: transformed Tensor.
+ """
+
+ seq_len, _, embed_dim = x.shape
+ assert embed_dim == self.embed_dim
+ if seq_len is None:
+ seq_len = 1
+
+ # L x B x D
+ residual = x * self.residual_weight
+
+ # L x B x D -> B x D x L
+ x = tf.transpose(x, perm=(1, 2, 0))
+
+ # Masking of the tensor
+ if padding_mask is not None:
+ x = x * tf.cast(tf.expand_dims(padding_mask, axis=1), x.dtype)
+
+ k = self.kernel(seq_len)
+
+ kernel_size = k.shape[1]
+ fft_len = seq_len
+ s = 0
+
+ if self.bidirectional:
+ k1, k2 = tf.split(k, [self.embed_dim, self.embed_dim], axis=0)
+ # D x 2*L-1
+ padding_l = tf.constant([[0, 0], [kernel_size - 1, 0]])
+ padding_r = tf.constant([[0, 0], [0, kernel_size - 1]])
+ padding_x = tf.constant([[0, 0], [0, 0], [kernel_size - 1, 0]])
+ k = tf.pad(k1, padding_l) + tf.pad(tf.reverse(k2, axis=[-1]), padding_r)
+ x = tf.pad(x, padding_x)
+ fft_len = fft_len + kernel_size - 1
+ s = 2 * kernel_size - 2
+
+ k_f = tf.signal.rfft(
+ k, fft_length=tf.constant([2 * fft_len], dtype=tf.int32)
+ )
+ x_f = tf.signal.rfft(
+ x, fft_length=tf.constant([2 * fft_len], dtype=tf.int32)
+ )
+ # B x D x L
+ out = tf.signal.irfft(
+ x_f * k_f, fft_length=tf.constant([2 * fft_len], dtype=tf.int32)
+ )[..., s : s + seq_len]
+
+ # B x D x L -> L x B x D
+ out = tf.nn.silu(tf.transpose(out, perm=(2, 0, 1)) + residual)
+ return out
diff --git a/official/projects/lra/linformer.py b/official/projects/lra/linformer.py
new file mode 100644
index 00000000000..c7cbfab4478
--- /dev/null
+++ b/official/projects/lra/linformer.py
@@ -0,0 +1,68 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Linformer model configurations and instantiation methods."""
+import dataclasses
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.modeling.hyperparams import base_config
+from official.nlp.configs import encoders
+from official.projects.lra.linformer_encoder import LinformerEncoder
+
+
+@dataclasses.dataclass
+class LinformerEncoderConfig(encoders.BertEncoderConfig):
+ """Extra paramerters for Linformer configs.
+
+ Attributes:
+ pad_token_id: the token id for the pad token
+ low_rank_features: number of dimensions for low-rank projection
+ """
+
+ pad_token_id: int = 0
+ low_rank_features: int = 256
+
+
+@base_config.bind(LinformerEncoderConfig)
+def get_encoder(encoder_cfg: LinformerEncoderConfig):
+ """Gets a 'LinformerEncoder' object.
+
+ Args:
+ encoder_cfg: A 'LinformerEncoderConfig'.
+
+ Returns:
+ A encoder object.
+ """
+ encoder = LinformerEncoder(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ num_attention_heads=encoder_cfg.num_attention_heads,
+ low_rank_features=encoder_cfg.low_rank_features,
+ inner_dim=encoder_cfg.intermediate_size,
+ inner_activation=tf_utils.get_activation(encoder_cfg.hidden_activation),
+ output_dropout=encoder_cfg.dropout_rate,
+ attention_dropout=encoder_cfg.attention_dropout_rate,
+ max_sequence_length=encoder_cfg.max_position_embeddings,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range
+ ),
+ output_range=encoder_cfg.output_range,
+ embedding_width=encoder_cfg.embedding_size,
+ norm_first=encoder_cfg.norm_first,
+ )
+ return encoder
diff --git a/official/projects/lra/linformer_encoder.py b/official/projects/lra/linformer_encoder.py
new file mode 100644
index 00000000000..62c88f92f63
--- /dev/null
+++ b/official/projects/lra/linformer_encoder.py
@@ -0,0 +1,306 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Linformer encoder. Modified From huggingface/transformers."""
+
+# pylint: disable=g-classes-have-attributes
+
+from typing import Any, Callable, Optional, Union
+
+from absl import logging
+import tensorflow as tf, tf_keras
+import tensorflow_models as tfm
+
+from official.modeling import tf_utils
+from official.projects.lra.linformer_encoder_block import LinformerEncoderBlock
+
+layers = tfm.nlp.layers
+
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
+
+
+class LinformerEncoder(tf_keras.layers.Layer):
+ """LinformerEncoder.
+
+ Args:
+ vocab_size: The size of the token vocabulary.
+ hidden_size: The size of the transformer hidden layers.
+ num_layers: The number of transformer layers.
+ num_attention_heads: The number of attention heads for each transformer. The
+ hidden size must be divisible by the number of attention heads.
+ low_rank_features: The number of dimensions for low rank projection.
+ max_sequence_length: The maximum sequence length that this encoder can
+ consume. If None, max_sequence_length uses the value from sequence length.
+ This determines the variable shape for positional embeddings.
+ type_vocab_size: The number of types that the 'type_ids' input can take.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ output_dropout: Dropout probability for the post-attention and output
+ dropout.
+ attention_dropout: The dropout rate to use for the attention layers within
+ the transformer layers.
+ initializer: The initialzer to use for all weights in this encoder.
+ output_range: The sequence output range, [0, output_range), by slicing the
+ target sequence of the last transformer layer. `None` means the entire
+ target sequence will attend to the source sequence, which yields the full
+ output.
+ embedding_width: The width of the word embeddings. If the embedding width is
+ not equal to hidden size, embedding parameters will be factorized into two
+ matrices in the shape of ['vocab_size', 'embedding_width'] and
+ ['embedding_width', 'hidden_size'] ('embedding_width' is usually much
+ smaller than 'hidden_size').
+ embedding_layer: An optional Layer instance which will be called to generate
+ embeddings for the input word IDs.
+ norm_first: Whether to normalize inputs to attention and intermediate dense
+ layers. If set False, output of attention and intermediate dense layers is
+ normalized.
+ """
+
+ def __init__(
+ self,
+ vocab_size: int,
+ hidden_size: int = 768,
+ num_layers: int = 12,
+ num_attention_heads: int = 12,
+ low_rank_features: int = 32,
+ max_sequence_length: int = 512,
+ type_vocab_size: int = 16,
+ inner_dim: int = 3072,
+ inner_activation: Callable[..., Any] = _approx_gelu,
+ output_dropout: float = 0.1,
+ attention_dropout: float = 0.1,
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
+ stddev=0.02
+ ),
+ output_range: Optional[int] = None,
+ embedding_width: Optional[int] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
+ norm_first: bool = False,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ # Linformer args
+ self._low_rank_features = low_rank_features
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
+
+ if embedding_width is None:
+ embedding_width = hidden_size
+
+ if embedding_layer is None:
+ self._embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ initializer=initializer,
+ name='word_embeddings',
+ )
+ else:
+ self._embedding_layer = embedding_layer
+
+ self._position_embedding_layer = layers.PositionEmbedding(
+ initializer=initializer,
+ max_length=max_sequence_length,
+ name='position_embedding',
+ )
+
+ self._type_embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=type_vocab_size,
+ embedding_width=embedding_width,
+ initializer=initializer,
+ use_one_hot=True,
+ name='type_embeddings',
+ )
+
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
+ name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32
+ )
+
+ self._embedding_dropout = tf_keras.layers.Dropout(
+ rate=output_dropout, name='embedding_dropout'
+ )
+
+ # We project the 'embedding' output to 'hidden_size' if it is not already
+ # 'hidden_size'.
+ self._embedding_projection = None
+ if embedding_width != hidden_size:
+ self._embedding_projection = tf_keras.layers.EinsumDense(
+ '...x,xy->...y',
+ output_shape=hidden_size,
+ bias_axes='y',
+ kernel_initializer=initializer,
+ name='embedding_projection',
+ )
+
+ self._transformer_layers = []
+ self._attention_mask_layer = layers.SelfAttentionMask(
+ name='self_attention_mask'
+ )
+ for i in range(num_layers):
+ layer = LinformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ low_rank_features=low_rank_features,
+ inner_dim=inner_dim,
+ inner_activation=inner_activation,
+ output_dropout=output_dropout,
+ attention_dropout=attention_dropout,
+ norm_first=norm_first,
+ return_attention_scores=False,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='transformer/layer_%d' % i,
+ )
+ self._transformer_layers.append(layer)
+ self._num_layers = num_layers
+
+ self._pooler_layer = tf_keras.layers.Dense(
+ units=hidden_size,
+ activation='tanh',
+ kernel_initializer=initializer,
+ name='pooler_transform',
+ )
+
+ self._config = {
+ 'vocab_size': vocab_size,
+ 'hidden_size': hidden_size,
+ 'num_layers': num_layers,
+ 'low_rank_features': low_rank_features,
+ 'num_attention_heads': num_attention_heads,
+ 'max_sequence_length': max_sequence_length,
+ 'type_vocab_size': type_vocab_size,
+ 'inner_dim': inner_dim,
+ 'inner_activation': tf_keras.activations.serialize(activation),
+ 'output_dropout': output_dropout,
+ 'attention_dropout': attention_dropout,
+ 'initializer': tf_keras.initializers.serialize(initializer),
+ 'output_range': output_range,
+ 'embedding_width': embedding_width,
+ 'embedding_layer': embedding_layer,
+ 'norm_first': norm_first,
+ }
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ )
+
+ def call(self, inputs):
+ if isinstance(inputs, dict):
+ word_embeddings = inputs.get('input_word_embeddings', None)
+ type_ids = inputs.get('input_type_ids', None)
+ if 'input_word_ids' in inputs.keys():
+ word_ids = inputs.get('input_word_ids')
+ mask = inputs.get('input_mask')
+ elif 'left_word_ids' in inputs.keys():
+ word_ids = inputs.get('left_word_ids')
+ mask = inputs.get('left_mask')
+ elif 'right_word_ids' in inputs.keys():
+ word_ids = inputs.get('right_word_ids')
+ mask = inputs.get('right_mask')
+ dense_inputs = inputs.get('dense_inputs', None)
+ dense_mask = inputs.get('dense_mask', None)
+ dense_type_ids = inputs.get('dense_type_ids', None)
+ elif isinstance(inputs, list):
+ ## Dual Encoder Tasks
+ word_ids, mask = inputs
+ word_embeddings = None
+ type_ids = None
+ dense_inputs, dense_mask, dense_type_ids = None, None, None
+ else:
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+
+ if type_ids is None:
+ type_ids = tf.zeros_like(mask) # pyrefly: ignore[unbound-name]
+
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids) # pyrefly: ignore[unbound-name]
+
+ if dense_inputs is not None:
+ mask = tf.concat([mask, dense_mask], axis=1)
+
+ embeddings = self._get_embeddings(
+ word_ids, type_ids, word_embeddings, dense_inputs, dense_type_ids # pyrefly: ignore[bad-argument-type]
+ )
+ embeddings = self._embedding_norm_layer(embeddings)
+ embeddings = self._embedding_dropout(embeddings)
+
+ if self._embedding_projection is not None:
+ embeddings = self._embedding_projection(embeddings)
+
+ attention_mask = self._attention_mask_layer(embeddings, mask)
+
+ encoder_outputs = []
+ x = embeddings
+ for layer in self._transformer_layers:
+ x = layer([x, attention_mask])
+ encoder_outputs.append(x)
+
+ last_encoder_output = encoder_outputs[-1]
+ first_token_tensor = last_encoder_output[:, 0, :]
+ pooled_output = self._pooler_layer(first_token_tensor)
+
+ output = dict(
+ sequence_output=encoder_outputs[-1],
+ pooled_output=pooled_output,
+ encoder_outputs=encoder_outputs,
+ )
+
+ return output
+
+ def get_embedding_table(self):
+ return self._embedding_layer.embeddings
+
+ def get_embedding_layer(self):
+ return self._embedding_layer
+
+ def get_config(self):
+ return dict(self._config)
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ if 'embedding_layer' in config and config['embedding_layer'] is not None:
+ warn_string = (
+ 'You are reloading a model that was saved with a '
+ 'potentially-shared embedding layer object. If you contine to '
+ 'train this model, the embedding layer will no longer be shared. '
+ 'To work around this, load the model outside of the Keras API.'
+ )
+ print('WARNING: ' + warn_string)
+ logging.warn(warn_string)
+
+ return cls(**config)
+
+ def _get_embeddings(
+ self,
+ word_ids: tf.Tensor,
+ type_ids: tf.Tensor,
+ word_embeddings: Optional[tf.Tensor],
+ dense_inputs: Optional[tf.Tensor],
+ dense_type_ids: Optional[tf.Tensor],
+ ) -> tf.Tensor:
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ # Concat the dense embeddings at sequence end.
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+
+ type_embeddings = self._type_embedding_layer(type_ids)
+
+ # absolute position embeddings.
+ position_embeddings = self._position_embedding_layer(word_embeddings)
+ return word_embeddings + position_embeddings + type_embeddings
diff --git a/official/projects/lra/linformer_encoder_block.py b/official/projects/lra/linformer_encoder_block.py
new file mode 100644
index 00000000000..aab8a5e03f8
--- /dev/null
+++ b/official/projects/lra/linformer_encoder_block.py
@@ -0,0 +1,458 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras-based LinformerEncoder block layer."""
+
+from typing import Any, Optional
+
+from absl import logging
+import tensorflow as tf, tf_keras
+import tensorflow_models as tfm
+
+from official.modeling import tf_utils
+
+
+@tf_keras.utils.register_keras_serializable(package="Text")
+class LinformerEncoderBlock(tf_keras.layers.Layer):
+ """LinformerEncoderBlock layer.
+
+ This layer implements the Linformer Encoder from
+ "Linformer: Self-Attention with Linear Complexity".
+ (https://arxiv.org/abs/2006.04768)
+
+ References:
+ [Linformer: Self-Attention with Linear Complexity]
+ (https://arxiv.org/abs/2006.04768)
+ [Long Range Arena: A Benchmark for Efficient Transformers]
+ (https://arxiv.org/abs/2011.04006)
+ """
+
+ def __init__(
+ self,
+ num_attention_heads,
+ inner_dim,
+ inner_activation,
+ low_rank_features,
+ kernel_initializer="glorot_uniform",
+ bias_initializer="zeros",
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ activity_regularizer=None,
+ kernel_constraint=None,
+ bias_constraint=None,
+ use_bias=True,
+ norm_first=False,
+ norm_epsilon=1e-12,
+ output_dropout=0.0,
+ attention_dropout=0.0,
+ inner_dropout=0.0,
+ attention_initializer=None,
+ attention_axes=None,
+ use_query_residual=True,
+ key_dim=None,
+ value_dim=None,
+ output_last_dim=None,
+ diff_q_kv_att_layer_norm=False,
+ return_attention_scores=False,
+ **kwargs
+ ):
+ """Initializes `LinformerEncoder`.
+
+ Note: If `output_last_dim` is used and `use_query_residual` is `True`, the
+ `output_last_dim`'s value must equal the first input's last dimension for
+ the query residual connection to work. This is because the residual
+ connection after the multi-head-attention requires their dimensions to
+ match. If `use_query_residual` is `False`, the `output_last_dim` dictactes
+ the last dimension of the output of this module and the
+ multi-head-attention.
+
+ E.g. let's say input dims are `[batch_size, seq_dim, input_last_dim]`.
+ Scenario 1: If `output_last_dim` is not `None`, then the output dims of this
+ module would be `[batch_size, seq_dim, output_last_dim]`. Note `key_dim` is
+ overriden by `output_last_dim`.
+ Scenario 2: If `output_last_dim` is `None` and `key_dim` is not `None`, then
+ the output dims of this module would be `[batch_size, seq_dim, key_dim]`.
+ Scenario 3: If the `output_last_dim` and `key_dim` are both `None`, the
+ output dims would be `[batch_size, seq_dim, input_last_dim]`.
+
+ Args:
+ num_attention_heads: Number of attention heads.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network.
+ low_rank_features: The number of dimensions for low-rank projection.
+ kernel_initializer: Initializer for dense layer kernels.
+ bias_initializer: Initializer for dense layer biases.
+ kernel_regularizer: Regularizer for dense layer kernels.
+ bias_regularizer: Regularizer for dense layer biases.
+ activity_regularizer: Regularizer for dense layer activity.
+ kernel_constraint: Constraint for dense layer kernels.
+ bias_constraint: Constraint for dense layer kernels.
+ use_bias: Whether to enable use_bias in attention layer. If set False,
+ use_bias in attention layer is disabled.
+ norm_first: Whether to normalize inputs to attention and intermediate
+ dense layers. If set False, output of attention and intermediate dense
+ layers is normalized.
+ norm_epsilon: Epsilon value to initialize normalization layers.
+ output_dropout: Dropout probability for the post-attention and output
+ dropout.
+ attention_dropout: Dropout probability for within the attention layer.
+ inner_dropout: Dropout probability for the first Dense layer in a
+ two-layer feedforward network.
+ attention_initializer: Initializer for kernels of attention layers. If set
+ `None`, attention layers use kernel_initializer as initializer for
+ kernel.
+ attention_axes: axes over which the attention is applied. `None` means
+ attention over all axes, but batch, heads, and features.
+ use_query_residual: Toggle to execute residual connection after attention.
+ key_dim: `key_dim` for the `tf_keras.layers.MultiHeadAttention`. If
+ `None`, we use the first `input_shape`'s last dim.
+ value_dim: `value_dim` for the `tf_keras.layers.MultiHeadAttention`.
+ output_last_dim: Final dimension of the output of this module. This also
+ dictates the value for the final dimension of the multi-head-attention.
+ When it's `None`, we use, in order of decreasing precedence, `key_dim` *
+ `num_heads` or the first `input_shape`'s last dim as the output's last
+ dim.
+ diff_q_kv_att_layer_norm: If `True`, create a separate attention layer
+ norm layer for query and key-value if `norm_first` is `True`. Invalid to
+ set to `True` if `norm_first` is `False`.
+ return_attention_scores: If `True`, the output of this layer will be a
+ tuple and additionally contain the attention scores in the shape of
+ `[batch_size, num_attention_heads, seq_dim, seq_dim]`.
+ **kwargs: keyword arguments.
+ """
+ tfm.nlp.layers.util.filter_kwargs(kwargs)
+ super().__init__(**kwargs)
+
+ self._num_heads = num_attention_heads
+ self._low_rank_features = low_rank_features
+ self._inner_dim = inner_dim
+ self._inner_activation = inner_activation
+ self._attention_dropout_rate = attention_dropout
+ self._output_dropout_rate = output_dropout
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
+ self._use_bias = use_bias
+ self._norm_first = norm_first
+ self._norm_epsilon = norm_epsilon
+ self._inner_dropout = inner_dropout
+ self._use_query_residual = use_query_residual
+ self._key_dim = key_dim
+ self._value_dim = value_dim
+ self._output_last_dim = output_last_dim
+ self._diff_q_kv_att_layer_norm = diff_q_kv_att_layer_norm
+ self._return_attention_scores = return_attention_scores
+ if attention_initializer:
+ self._attention_initializer = tf_keras.initializers.get(
+ attention_initializer
+ )
+ else:
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer
+ )
+ self._attention_axes = attention_axes
+
+ if self._diff_q_kv_att_layer_norm and not self._norm_first:
+ raise ValueError(
+ "Setting `diff_q_and_kv_attention_layer_norm` to True"
+ "when `norm_first` is False is invalid."
+ )
+
+ def build(self, input_shape):
+ if isinstance(input_shape, tf.TensorShape):
+ input_tensor_shape = input_shape
+ elif isinstance(input_shape, (list, tuple)):
+ input_tensor_shape = tf.TensorShape(input_shape[0])
+ else:
+ raise ValueError(
+ "The type of input shape argument is not supported, got: %s"
+ % type(input_shape)
+ )
+ einsum_equation = "abc,cd->abd"
+ if len(input_tensor_shape.as_list()) > 3:
+ einsum_equation = "...bc,cd->...bd"
+ hidden_size = input_tensor_shape[-1]
+ if hidden_size % self._num_heads != 0:
+ logging.warning(
+ (
+ "The input size (%d) is not a multiple of the number of attention"
+ " heads (%d)"
+ ),
+ hidden_size,
+ self._num_heads,
+ )
+ if self._key_dim is None:
+ self._key_dim = int(hidden_size // self._num_heads)
+ if self._output_last_dim is None:
+ last_output_shape = hidden_size
+ else:
+ last_output_shape = self._output_last_dim
+
+ common_kwargs = dict(
+ bias_regularizer=self._bias_regularizer,
+ activity_regularizer=self._activity_regularizer,
+ kernel_constraint=self._kernel_constraint,
+ bias_constraint=self._bias_constraint,
+ )
+ self._key_projection = tf_keras.layers.Dense(
+ self._low_rank_features,
+ activation=None,
+ use_bias=False,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ name="key_low_rank_projection",
+ **common_kwargs
+ )
+ self._value_projection = tf_keras.layers.Dense(
+ self._low_rank_features,
+ activation=None,
+ use_bias=False,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ name="value_low_rank_projection",
+ **common_kwargs
+ )
+ self._attention_layer = tf_keras.layers.MultiHeadAttention(
+ num_heads=self._num_heads,
+ key_dim=self._low_rank_features,
+ value_dim=self._low_rank_features,
+ dropout=self._attention_dropout_rate,
+ use_bias=self._use_bias,
+ kernel_initializer=self._attention_initializer,
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ attention_axes=self._attention_axes,
+ output_shape=self._output_last_dim,
+ name="self_attention",
+ **common_kwargs
+ )
+ self._attention_dropout = tf_keras.layers.Dropout(
+ rate=self._attention_dropout_rate
+ )
+ # Use float32 in layernorm for numeric stability.
+ # It is probably safe in mixed_float16, but we haven't validated this yet.
+ self._attention_layer_norm = tf_keras.layers.LayerNormalization(
+ name="self_attention_layer_norm",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32,
+ )
+ self._attention_layer_norm_kv = self._attention_layer_norm
+ if self._diff_q_kv_att_layer_norm:
+ self._attention_layer_norm_kv = tf_keras.layers.LayerNormalization(
+ name="self_attention_layer_norm_kv",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32,
+ )
+
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
+ einsum_equation,
+ output_shape=(None, self._inner_dim),
+ bias_axes="d",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ name="intermediate",
+ **common_kwargs
+ )
+ policy = tf_keras.mixed_precision.global_policy()
+ if policy.name == "mixed_bfloat16":
+ # bfloat16 causes BERT with the LAMB optimizer to not converge
+ # as well, so we use float32.
+ # TODO(b/154538392): Investigate this.
+ policy = tf.float32
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
+ self._inner_activation, dtype=policy
+ )
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
+ rate=self._inner_dropout
+ )
+ self._output_dense = tf_keras.layers.EinsumDense(
+ einsum_equation,
+ output_shape=(None, last_output_shape),
+ bias_axes="d",
+ name="output",
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ bias_initializer=tf_utils.clone_initializer(self._bias_initializer),
+ **common_kwargs
+ )
+ self._output_dropout = tf_keras.layers.Dropout(
+ rate=self._output_dropout_rate
+ )
+ # Use float32 in layernorm for numeric stability.
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
+ name="output_layer_norm",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32,
+ )
+
+ super().build(input_shape)
+
+ def get_config(self):
+ config = {
+ "num_attention_heads": self._num_heads,
+ "low_rank_features": self._low_rank_features,
+ "inner_dim": self._inner_dim,
+ "inner_activation": self._inner_activation,
+ "output_dropout": self._output_dropout_rate,
+ "attention_dropout": self._attention_dropout_rate,
+ "kernel_initializer": tf_keras.initializers.serialize(
+ self._kernel_initializer
+ ),
+ "bias_initializer": tf_keras.initializers.serialize(
+ self._bias_initializer
+ ),
+ "kernel_regularizer": tf_keras.regularizers.serialize(
+ self._kernel_regularizer
+ ),
+ "bias_regularizer": tf_keras.regularizers.serialize(
+ self._bias_regularizer
+ ),
+ "activity_regularizer": tf_keras.regularizers.serialize(
+ self._activity_regularizer
+ ),
+ "kernel_constraint": tf_keras.constraints.serialize(
+ self._kernel_constraint
+ ),
+ "bias_constraint": tf_keras.constraints.serialize(
+ self._bias_constraint
+ ),
+ "use_bias": self._use_bias,
+ "norm_first": self._norm_first,
+ "norm_epsilon": self._norm_epsilon,
+ "inner_dropout": self._inner_dropout,
+ "attention_initializer": tf_keras.initializers.serialize(
+ self._attention_initializer
+ ),
+ "attention_axes": self._attention_axes,
+ "use_query_residual": self._use_query_residual,
+ "key_dim": self._key_dim,
+ "value_dim": self._value_dim,
+ "output_last_dim": self._output_last_dim,
+ "diff_q_kv_att_layer_norm": self._diff_q_kv_att_layer_norm,
+ }
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(config.items()))
+
+ def call(self, inputs: Any, output_range: Optional[tf.Tensor] = None) -> Any:
+ """Transformer self-attention encoder block call.
+
+ Args:
+ inputs: a single tensor or a list of tensors. `input tensor` as the single
+ sequence of embeddings. [`input tensor`, `attention mask`] to have the
+ additional attention mask. [`query tensor`, `key value tensor`,
+ `attention mask`] to have separate input streams for the query, and
+ key/value to the multi-head attention.
+ output_range: the sequence output range, [0, output_range) for slicing the
+ target sequence. `None` means the target sequence is not sliced. If you
+ would like to have no change to the model training, it is better to only
+ set the `output_range` for serving.
+
+ Returns:
+ An output tensor with the same dimensions as input/query tensor.
+ """
+ if isinstance(inputs, (list, tuple)):
+ if len(inputs) == 2:
+ input_tensor, attention_mask = inputs
+ key_value = None
+ elif len(inputs) == 3:
+ input_tensor, key_value, attention_mask = inputs
+ else:
+ raise ValueError(
+ "Unexpected inputs to %s with length at %d"
+ % (self.__class__, len(inputs))
+ )
+ else:
+ input_tensor, key_value, attention_mask = (inputs, None, None)
+
+ if output_range:
+ if self._norm_first:
+ source_tensor = input_tensor[:, 0:output_range, :]
+ input_tensor = self._attention_layer_norm(input_tensor)
+ if key_value is not None:
+ key_value = self._attention_layer_norm_kv(key_value)
+ target_tensor = input_tensor[:, 0:output_range, :]
+ if attention_mask is not None:
+ attention_mask = attention_mask[:, 0:output_range, :]
+ else:
+ if self._norm_first:
+ source_tensor = input_tensor
+ input_tensor = self._attention_layer_norm(input_tensor)
+ if key_value is not None:
+ key_value = self._attention_layer_norm_kv(key_value)
+ target_tensor = input_tensor
+
+ if key_value is None:
+ key_value = input_tensor
+
+ ## Low Rank Projection Here
+ key = self._key_projection(key_value)
+ value = self._value_projection(input_tensor)
+ ## Low Rank Projection Done
+
+ if self._return_attention_scores:
+ attention_output, attention_scores = self._attention_layer(
+ query=target_tensor,
+ key=key,
+ value=value,
+ attention_mask=attention_mask,
+ return_attention_scores=True,
+ )
+ else:
+ attention_output = self._attention_layer(
+ query=target_tensor,
+ key=key,
+ value=value,
+ attention_mask=attention_mask,
+ )
+ attention_output = self._attention_dropout(attention_output)
+
+ if self._norm_first:
+ # Important to not combine `self._norm_first` and
+ # `self._use_query_residual` into one if clause because else is only for
+ # `_norm_first == False`.
+ if self._use_query_residual:
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
+ else:
+ if self._use_query_residual:
+ attention_output = target_tensor + attention_output
+ attention_output = self._attention_layer_norm(attention_output)
+
+ if self._norm_first:
+ source_attention_output = attention_output
+ attention_output = self._output_layer_norm(attention_output)
+ inner_output = self._intermediate_dense(attention_output)
+ inner_output = self._intermediate_activation_layer(inner_output)
+ inner_output = self._inner_dropout_layer(inner_output)
+ layer_output = self._output_dense(inner_output)
+ layer_output = self._output_dropout(layer_output)
+
+ if self._norm_first:
+ layer_output = source_attention_output + layer_output # pyrefly: ignore[unbound-name]
+ else:
+ # During mixed precision training, layer norm output is always fp32 for
+ # now. Casts fp32 for the subsequent add.
+ layer_output = tf.cast(layer_output, tf.float32)
+ layer_output = self._output_layer_norm(layer_output + attention_output)
+
+ if self._return_attention_scores:
+ return layer_output, attention_scores # pyrefly: ignore[unbound-name]
+ else:
+ return layer_output
diff --git a/official/projects/lra/linformer_experiments.py b/official/projects/lra/linformer_experiments.py
new file mode 100644
index 00000000000..7e7ad880030
--- /dev/null
+++ b/official/projects/lra/linformer_experiments.py
@@ -0,0 +1,155 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Linformer experiments."""
+# pylint: disable=g-doc-return-or-yield,line-too-long
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.nlp.configs import encoders
+from official.nlp.data import sentence_prediction_dataloader
+from official.nlp.tasks import sentence_prediction
+from official.projects.lra import lra_dual_encoder_dataloader
+from official.projects.lra import lra_dual_encoder_task
+from official.projects.lra.linformer import LinformerEncoderConfig
+
+
+AdamWeightDecay = optimization.AdamWeightDecayConfig
+PolynomialLr = optimization.PolynomialLrConfig
+PolynomialWarmupConfig = optimization.PolynomialWarmupConfig
+
+_TRAINER = cfg.TrainerConfig(
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.01,
+ 'exclude_from_weight_decay': [
+ 'LayerNorm',
+ 'layer_norm',
+ 'bias',
+ ],
+ },
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 3e-5,
+ 'end_learning_rate': 0.0,
+ },
+ },
+ 'warmup': {'type': 'polynomial'},
+ })
+)
+
+
+@exp_factory.register_config_factory('linformer/lra_listops')
+def linformer_listops() -> cfg.ExperimentConfig:
+ """Linformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=LinformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('linformer/lra_imdb')
+def linformer_imdb() -> cfg.ExperimentConfig:
+ """Linformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=LinformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('linformer/lra_cifar')
+def linformer_cifar() -> cfg.ExperimentConfig:
+ """Linformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=LinformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('linformer/lra_pathfinder')
+def linformer_pathfinder() -> cfg.ExperimentConfig:
+ """Linformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=LinformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('linformer/lra_aan')
+def linformer_aan() -> cfg.ExperimentConfig:
+ """Linformer LRA Task."""
+ config = cfg.ExperimentConfig(
+ task=lra_dual_encoder_task.DualEncoderConfig(
+ model=lra_dual_encoder_task.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=LinformerEncoderConfig()
+ )
+ ),
+ train_data=lra_dual_encoder_dataloader.DualEncoderDataConfig(),
+ validation_data=lra_dual_encoder_dataloader.DualEncoderDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
diff --git a/official/projects/lra/lra_dual_encoder.py b/official/projects/lra/lra_dual_encoder.py
new file mode 100644
index 00000000000..b9dd8cc61d2
--- /dev/null
+++ b/official/projects/lra/lra_dual_encoder.py
@@ -0,0 +1,135 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Trainer network for dual encoder style models."""
+# pylint: disable=g-classes-have-attributes
+import collections
+import tensorflow as tf, tf_keras
+
+import tensorflow_models as tfm
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class LRADualEncoder(tf_keras.layers.Layer):
+ """A dual encoder model based on a transformer-based encoder.
+
+ This is an implementation of the dual encoder network structure based on the
+ transfomer stack, as described in ["Language-agnostic BERT Sentence
+ Embedding"](https://arxiv.org/abs/2007.01852)
+
+ The DualEncoder allows a user to pass in a transformer stack, and build a dual
+ encoder model based on the transformer stack.
+
+ Args:
+ network: A transformer network which should output an encoding output.
+ max_seq_length: The maximum allowed sequence length for transformer.
+ normalize: If set to True, normalize the encoding produced by transfomer.
+ logit_scale: The scaling factor of dot products when doing training.
+ logit_margin: The margin between positive and negative when doing training.
+ output: The output style for this network. Can be either `logits` or
+ `predictions`. If set to `predictions`, it will output the embedding
+ producted by transformer network.
+ """
+
+ def __init__(
+ self,
+ network,
+ num_classes,
+ max_seq_length,
+ dropout_rate=0.1,
+ initializer='glorot_uniform',
+ use_encoder_pooler=True,
+ inner_dim=None,
+ head_name='dual_encode',
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+
+ config_dict = {
+ 'network': network,
+ 'num_classes': num_classes,
+ 'head_name': head_name,
+ 'max_seq_length': max_seq_length,
+ 'initializer': initializer,
+ 'use_encoder_pooler': use_encoder_pooler,
+ 'inner_dim': inner_dim,
+ }
+ # We are storing the config dict as a namedtuple here to ensure checkpoint
+ # compatibility with an earlier version of this model which did not track
+ # the config dict attribute. TF does not track immutable attrs which
+ # do not contain Trackables, so by creating a config namedtuple instead of
+ # a dict we avoid tracking it.
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
+ self._config = config_cls(**config_dict)
+ self._use_encoder_pooler = use_encoder_pooler
+
+ self.network = network
+ self.classifier = tfm.nlp.layers.ClassificationHead(
+ inner_dim=0 if use_encoder_pooler else inner_dim,
+ num_classes=num_classes,
+ initializer=initializer,
+ dropout_rate=dropout_rate,
+ name=head_name,
+ )
+
+ def call(self, inputs):
+ if isinstance(inputs, dict):
+ left_word_ids = inputs.get('left_word_ids')
+ left_mask = inputs.get('left_mask')
+
+ right_word_ids = inputs.get('right_word_ids')
+ right_mask = inputs.get('right_mask')
+ else:
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+
+ inputs = [left_word_ids, left_mask, right_word_ids, right_mask]
+
+ left_inputs = [left_word_ids, left_mask]
+ left_outputs = self.network(left_inputs)
+ right_inputs = [right_word_ids, right_mask]
+ right_outputs = self.network(right_inputs)
+
+ if self._use_encoder_pooler:
+ # Because we have a copy of inputs to create this Model object, we can
+ # invoke the Network object with its own input tensors to start the Model.
+ if isinstance(left_outputs, list):
+ left_cls_inputs = left_outputs[1]
+ right_cls_inputs = right_outputs[1]
+ else:
+ left_cls_inputs = left_outputs['pooled_output']
+ right_cls_inputs = right_outputs['pooled_output']
+ else:
+ if isinstance(left_outputs, list):
+ left_cls_inputs = left_outputs[0]
+ right_cls_inputs = right_outputs[0]
+ else:
+ left_cls_inputs = left_outputs['sequence_output']
+ right_cls_inputs = right_outputs['sequence_output']
+
+ cls_inputs = tf.concat([left_cls_inputs, right_cls_inputs], -1)
+ predictions = self.classifier(cls_inputs)
+ return predictions
+
+ def get_config(self):
+ return dict(self._config._asdict())
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+ @property
+ def checkpoint_items(self):
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(encoder=self.network)
+ return items
diff --git a/official/projects/lra/lra_dual_encoder_dataloader.py b/official/projects/lra/lra_dual_encoder_dataloader.py
new file mode 100644
index 00000000000..1894ab8fe1a
--- /dev/null
+++ b/official/projects/lra/lra_dual_encoder_dataloader.py
@@ -0,0 +1,124 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Loads dataset for the similarity comparison (classification) task."""
+
+import dataclasses
+from typing import Mapping, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import config_definitions as cfg
+from official.core import input_reader
+from official.nlp.data import data_loader
+from official.nlp.data import data_loader_factory
+
+
+LABEL_TYPES_MAP = {'int': tf.int64, 'float': tf.float32}
+
+
+@dataclasses.dataclass
+class DualEncoderDataConfig(cfg.DataConfig):
+ """Data config for similarity comparison task."""
+
+ input_path: str = ''
+ global_batch_size: int = 32
+ is_training: bool = True
+ seq_length: int = 128
+ label_type: str = 'int'
+ # Whether to include the example id number.
+ include_example_id: bool = False
+ label_field: str = 'label_ids'
+ # Maps the key in TfExample to feature name.
+ # E.g 'label_ids' to 'next_sentence_labels'
+ label_name: Optional[Tuple[str, str]] = None
+ # Either tfrecord, sstable, or recordio.
+ file_type: str = 'tfrecord'
+
+
+@data_loader_factory.register_data_loader_cls(DualEncoderDataConfig)
+class DualEncoderDataLoader(data_loader.DataLoader):
+ """A class to load dataset for similarity comparison (classification) task."""
+
+ def __init__(self, params):
+ self._params = params
+ self._seq_length = params.seq_length
+ self._include_example_id = params.include_example_id
+ self._label_field = params.label_field
+ if params.label_name:
+ self._label_name_mapping = dict([params.label_name])
+ else:
+ self._label_name_mapping = dict()
+
+ def name_to_features_spec(self):
+ """Defines features to decode. Subclass may override to append features."""
+ label_type = LABEL_TYPES_MAP[self._params.label_type]
+ name_to_features = {
+ 'left_word_ids': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'left_mask': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'right_word_ids': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ 'right_mask': tf.io.FixedLenFeature([self._seq_length], tf.int64),
+ self._label_field: tf.io.FixedLenFeature([], label_type),
+ }
+ if self._include_example_id:
+ name_to_features['example_id'] = tf.io.FixedLenFeature([], tf.int64)
+
+ return name_to_features
+
+ def _decode(self, record: tf.Tensor):
+ """Decodes a serialized tf.Example."""
+ example = tf.io.parse_single_example(record, self.name_to_features_spec())
+
+ # tf.Example only supports tf.int64, but the TPU only supports tf.int32.
+ # So cast all int64 to int32.
+ for name in example:
+ t = example[name]
+ if t.dtype == tf.int64:
+ t = tf.cast(t, tf.int32)
+ example[name] = t
+
+ return example
+
+ def _parse(self, record: Mapping[str, tf.Tensor]):
+ """Parses raw tensors into a dict of tensors to be consumed by the model."""
+ key_mapping = {
+ 'left_ids': 'left_word_ids',
+ 'left_mask': 'left_mask',
+ 'right_ids': 'right_word_ids',
+ 'right_mask': 'right_mask',
+ }
+ ret = {}
+ for record_key in record:
+ if record_key in key_mapping:
+ ret[key_mapping[record_key]] = record[record_key]
+ else:
+ ret[record_key] = record[record_key]
+
+ if self._label_field in self._label_name_mapping:
+ ret[self._label_name_mapping[self._label_field]] = record[
+ self._label_field
+ ]
+
+ return ret
+
+ def load(self, input_context: Optional[tf.distribute.InputContext] = None):
+ """Returns a tf.dataset.Dataset."""
+ reader = input_reader.InputReader(
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
+ params=self._params,
+ decoder_fn=self._decode,
+ parser_fn=self._parse,
+ )
+ return reader.read(input_context)
diff --git a/official/projects/lra/lra_dual_encoder_task.py b/official/projects/lra/lra_dual_encoder_task.py
new file mode 100644
index 00000000000..9eaa17ff9d4
--- /dev/null
+++ b/official/projects/lra/lra_dual_encoder_task.py
@@ -0,0 +1,349 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Trainer network for dual encoder style models."""
+# pylint: disable=g-classes-have-attributes
+import dataclasses
+from typing import List, Union, Optional
+
+from absl import logging
+import numpy as np
+import orbit
+from scipy import stats
+from sklearn import metrics as sklearn_metrics
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import config_definitions as cfg
+from official.core import task_factory
+from official.modeling import tf_utils
+from official.modeling.hyperparams import base_config
+from official.nlp.configs import encoders
+from official.nlp.data import data_loader_factory
+from official.nlp.tasks import utils
+
+from official.projects.lra import lra_dual_encoder
+
+METRIC_TYPES = frozenset(
+ ['accuracy', 'f1', 'matthews_corrcoef', 'pearson_spearman_corr']
+)
+
+
+@dataclasses.dataclass
+class ModelConfig(base_config.Config):
+ """A classifier/regressor configuration."""
+
+ num_classes: int = 2
+ use_encoder_pooler: bool = False
+ encoder: encoders.EncoderConfig = encoders.EncoderConfig()
+ max_seq_length: int = 512
+
+
+@dataclasses.dataclass
+class DualEncoderConfig(cfg.TaskConfig):
+ """The model config."""
+
+ # At most one of `init_checkpoint` and `hub_module_url` can
+ # be specified.
+ init_checkpoint: str = ''
+ init_cls_pooler: bool = False
+ hub_module_url: str = ''
+ metric_type: str = 'accuracy'
+ # Defines the concrete model config at instantiation time.
+ model: ModelConfig = ModelConfig()
+ train_data: cfg.DataConfig = cfg.DataConfig()
+ validation_data: cfg.DataConfig = cfg.DataConfig()
+
+
+@task_factory.register_task_cls(DualEncoderConfig)
+class DualEncoderTask(base_task.Task):
+ """Task object for DualEncoderTask."""
+
+ def __init__(self, params: cfg.TaskConfig, logging_dir=None, name=None):
+ super().__init__(params, logging_dir, name=name)
+ if params.metric_type not in METRIC_TYPES:
+ raise ValueError('Invalid metric_type: {}'.format(params.metric_type))
+ self.metric_type = params.metric_type
+ if hasattr(params.train_data, 'label_field'):
+ self.label_field = params.train_data.label_field
+ else:
+ self.label_field = 'label_ids'
+
+ def build_model(self):
+ if self.task_config.hub_module_url and self.task_config.init_checkpoint:
+ raise ValueError(
+ 'At most one of `hub_module_url` and '
+ '`init_checkpoint` can be specified.'
+ )
+ if self.task_config.hub_module_url:
+ encoder_network = utils.get_encoder_from_hub(
+ self.task_config.hub_module_url
+ )
+ else:
+ encoder_network = encoders.build_encoder(self.task_config.model.encoder)
+ encoder_cfg = self.task_config.model.encoder.get()
+
+ return lra_dual_encoder.LRADualEncoder(
+ network=encoder_network,
+ max_seq_length=self.task_config.model.max_seq_length,
+ num_classes=self.task_config.model.num_classes,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range
+ ),
+ use_encoder_pooler=self.task_config.model.use_encoder_pooler,
+ inner_dim=encoder_cfg.hidden_size * 2,
+ )
+
+ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
+ label_ids = labels[self.label_field]
+ if self.task_config.model.num_classes == 1:
+ loss = tf_keras.losses.mean_squared_error(label_ids, model_outputs)
+ else:
+ loss = tf_keras.losses.sparse_categorical_crossentropy(
+ label_ids, tf.cast(model_outputs, tf.float32), from_logits=True
+ )
+
+ if aux_losses:
+ loss += tf.add_n(aux_losses)
+ return tf_utils.safe_mean(loss)
+
+ def build_inputs(self, params, input_context=None):
+ """Returns tf.data.Dataset for sentence_prediction task."""
+ if params.input_path == 'dummy':
+
+ def dummy_data(_):
+ dummy_ids = tf.zeros((1, params.seq_length), dtype=tf.int32)
+ x = dict(
+ left_word_ids=dummy_ids,
+ left_mask=dummy_ids,
+ right_word_ids=dummy_ids,
+ right_mask=dummy_ids,
+ )
+
+ if self.task_config.model.num_classes == 1:
+ y = tf.zeros((1,), dtype=tf.float32)
+ else:
+ y = tf.zeros((1, 1), dtype=tf.int32)
+ x[self.label_field] = y
+ return x
+
+ dataset = tf.data.Dataset.range(1)
+ dataset = dataset.repeat()
+ dataset = dataset.map(
+ dummy_data, num_parallel_calls=tf.data.experimental.AUTOTUNE
+ )
+ return dataset
+
+ return data_loader_factory.get_data_loader(params).load(input_context)
+
+ def build_metrics(self, training=None):
+ del training
+ if self.task_config.model.num_classes == 1:
+ metrics = [tf_keras.metrics.MeanSquaredError()]
+ elif self.task_config.model.num_classes == 2:
+ metrics = [
+ tf_keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
+ tf_keras.metrics.AUC(name='auc', curve='PR'),
+ ]
+ else:
+ metrics = [
+ tf_keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
+ ]
+ return metrics
+
+ def process_metrics(self, metrics, labels, model_outputs):
+ for metric in metrics:
+ if metric.name == 'auc':
+ # Convert the logit to probability and extract the probability of True..
+ metric.update_state(
+ labels[self.label_field],
+ tf.expand_dims(tf.nn.softmax(model_outputs)[:, 1], axis=1),
+ )
+ if metric.name == 'cls_accuracy':
+ metric.update_state(labels[self.label_field], model_outputs)
+
+ def process_compiled_metrics(self, compiled_metrics, labels, model_outputs):
+ compiled_metrics.update_state(labels[self.label_field], model_outputs)
+
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
+ features, labels = inputs, inputs
+ outputs = self.inference_step(features, model)
+ loss = self.build_losses(
+ labels=labels, model_outputs=outputs, aux_losses=model.losses
+ )
+ logs = {self.loss: loss}
+ if metrics:
+ self.process_metrics(metrics, labels, outputs)
+ if model.compiled_metrics:
+ self.process_compiled_metrics(model.compiled_metrics, labels, outputs)
+ logs.update({m.name: m.result() for m in metrics or []})
+ logs.update({m.name: m.result() for m in model.metrics})
+ if self.metric_type == 'matthews_corrcoef':
+ logs.update({
+ 'sentence_prediction': (
+ tf.expand_dims( # Ensure one prediction along batch dimension.
+ tf.math.argmax(outputs, axis=1), axis=1
+ )
+ ),
+ 'labels': labels[self.label_field],
+ })
+ else:
+ logs.update({
+ 'sentence_prediction': outputs,
+ 'labels': labels[self.label_field],
+ })
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ if self.metric_type == 'accuracy':
+ return None
+ if state is None:
+ state = {'sentence_prediction': [], 'labels': []}
+ state['sentence_prediction'].append(
+ np.concatenate(
+ [v.numpy() for v in step_outputs['sentence_prediction']], axis=0 # pyrefly: ignore[unsupported-operation]
+ )
+ )
+ state['labels'].append(
+ np.concatenate([v.numpy() for v in step_outputs['labels']], axis=0) # pyrefly: ignore[unsupported-operation]
+ )
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ if self.metric_type == 'accuracy':
+ return None
+
+ preds = np.concatenate(aggregated_logs['sentence_prediction'], axis=0)
+ labels = np.concatenate(aggregated_logs['labels'], axis=0)
+ if self.metric_type == 'f1':
+ preds = np.argmax(preds, axis=1)
+ return {self.metric_type: sklearn_metrics.f1_score(labels, preds)}
+ elif self.metric_type == 'matthews_corrcoef':
+ preds = np.reshape(preds, -1)
+ labels = np.reshape(labels, -1)
+ return {
+ self.metric_type: sklearn_metrics.matthews_corrcoef(preds, labels)
+ }
+ elif self.metric_type == 'pearson_spearman_corr':
+ preds = np.reshape(preds, -1)
+ labels = np.reshape(labels, -1)
+ pearson_corr = stats.pearsonr(preds, labels)[0]
+ spearman_corr = stats.spearmanr(preds, labels)[0]
+ corr_metric = (pearson_corr + spearman_corr) / 2
+ return {self.metric_type: corr_metric}
+
+ def initialize(self, model):
+ """Load a pretrained checkpoint (if exists) and then train from iter 0."""
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ logging.info(
+ 'Trying to load pretrained checkpoint from %s', ckpt_dir_or_file
+ )
+ if ckpt_dir_or_file and tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+ if not ckpt_dir_or_file:
+ logging.info(
+ 'No checkpoint file found from %s. Will not load.', ckpt_dir_or_file
+ )
+ return
+
+ pretrain2finetune_mapping = {
+ 'encoder': model.checkpoint_items['encoder'],
+ }
+ if self.task_config.init_cls_pooler:
+ # This option is valid when use_encoder_pooler is false.
+ pretrain2finetune_mapping['next_sentence.pooler_dense'] = (
+ model.checkpoint_items['sentence_prediction.pooler_dense']
+ )
+ ckpt = tf.train.Checkpoint(**pretrain2finetune_mapping)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ logging.info(
+ 'Finished loading pretrained checkpoint from %s', ckpt_dir_or_file
+ )
+
+
+def predict(
+ task: DualEncoderTask,
+ params: cfg.DataConfig,
+ model: tf_keras.Model,
+ params_aug: Optional[cfg.DataConfig] = None,
+ test_time_aug_wgt: float = 0.3,
+) -> List[Union[int, float]]:
+ """Predicts on the input data.
+
+ Args:
+ task: A `DualEncoderTask` object.
+ params: A `cfg.DataConfig` object.
+ model: A keras.Model.
+ params_aug: A `cfg.DataConfig` object for augmented data.
+ test_time_aug_wgt: Test time augmentation weight. The prediction score will
+ use (1. - test_time_aug_wgt) original prediction plus test_time_aug_wgt
+ augmented prediction.
+
+ Returns:
+ A list of predictions with length of `num_examples`. For regression task,
+ each element in the list is the predicted score; for classification task,
+ each element is the predicted class id.
+ """
+
+ def predict_step(inputs):
+ """Replicated prediction calculation."""
+ x = inputs
+ example_id = x.pop('example_id')
+ outputs = task.inference_step(x, model)
+ return dict(example_id=example_id, predictions=outputs)
+
+ def aggregate_fn(state, outputs):
+ """Concatenates model's outputs."""
+ if state is None:
+ state = []
+
+ for per_replica_example_id, per_replica_batch_predictions in zip(
+ outputs['example_id'], outputs['predictions']
+ ):
+ state.extend(zip(per_replica_example_id, per_replica_batch_predictions))
+ return state
+
+ dataset = orbit.utils.make_distributed_dataset(
+ tf.distribute.get_strategy(), task.build_inputs, params
+ )
+ outputs = utils.predict(predict_step, aggregate_fn, dataset)
+
+ # When running on TPU POD, the order of output cannot be maintained,
+ # so we need to sort by example_id.
+ outputs = sorted(outputs, key=lambda x: x[0])
+ is_regression = task.task_config.model.num_classes == 1
+ if params_aug is not None:
+ dataset_aug = orbit.utils.make_distributed_dataset(
+ tf.distribute.get_strategy(), task.build_inputs, params_aug
+ )
+ outputs_aug = utils.predict(predict_step, aggregate_fn, dataset_aug)
+ outputs_aug = sorted(outputs_aug, key=lambda x: x[0])
+ if is_regression:
+ return [
+ (1.0 - test_time_aug_wgt) * x[1] + test_time_aug_wgt * y[1]
+ for x, y in zip(outputs, outputs_aug)
+ ]
+ else:
+ return [
+ tf.argmax(
+ (1.0 - test_time_aug_wgt) * x[1] + test_time_aug_wgt * y[1],
+ axis=-1,
+ )
+ for x, y in zip(outputs, outputs_aug)
+ ]
+ if is_regression:
+ return [x[1] for x in outputs]
+ else:
+ return [tf.argmax(x[1], axis=-1) for x in outputs]
diff --git a/official/projects/lra/mega.py b/official/projects/lra/mega.py
new file mode 100644
index 00000000000..a60718253a1
--- /dev/null
+++ b/official/projects/lra/mega.py
@@ -0,0 +1,76 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Mega model configurations and instantiation methods."""
+import dataclasses
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.modeling.hyperparams import base_config
+from official.nlp.configs import encoders
+from official.projects.lra.mega_encoder import MegaEncoder
+
+
+@dataclasses.dataclass
+class MegaEncoderConfig(encoders.BertEncoderConfig):
+ """Extra paramerters for Mega configs.
+
+ Attributes:
+ pad_token_id: the token id for the pad token
+ low_rank_features: number of dimensions for low-rank projection
+ """
+
+ zdim: int = 64
+ hdim: int = 256
+ ndim: int = 16
+ activation: str = 'silu'
+ bidirectional: bool = False
+ dropout: float = 0.0
+ hidden_dropout: float = 0.0
+
+
+@base_config.bind(MegaEncoderConfig)
+def get_encoder(encoder_cfg: MegaEncoderConfig):
+ """Gets a 'MegaEncoder' object.
+
+ Args:
+ encoder_cfg: A 'MegaEncoderConfig'.
+
+ Returns:
+ A encoder object.
+ """
+ encoder = MegaEncoder(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ zdim=encoder_cfg.zdim,
+ hdim=encoder_cfg.hdim,
+ ndim=encoder_cfg.ndim,
+ activation=encoder_cfg.activation,
+ bidirectional=encoder_cfg.bidirectional,
+ dropout=encoder_cfg.dropout,
+ hidden_dropout=encoder_cfg.hidden_dropout,
+ inner_activation=tf_utils.get_activation(encoder_cfg.hidden_activation),
+ attention_dropout=encoder_cfg.attention_dropout_rate,
+ max_sequence_length=encoder_cfg.max_position_embeddings,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range
+ ),
+ output_range=encoder_cfg.output_range,
+ embedding_width=encoder_cfg.embedding_size,
+ norm_first=encoder_cfg.norm_first,
+ )
+ return encoder
diff --git a/official/projects/lra/mega_encoder.py b/official/projects/lra/mega_encoder.py
new file mode 100644
index 00000000000..0d8d9123fed
--- /dev/null
+++ b/official/projects/lra/mega_encoder.py
@@ -0,0 +1,303 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Mega encoder. Modified From huggingface/transformers."""
+
+# pylint: disable=g-classes-have-attributes
+
+from typing import Any, Callable, Optional, Union
+
+from absl import logging
+import tensorflow as tf, tf_keras
+import tensorflow_models as tfm
+
+from official.modeling import tf_utils
+from official.projects.lra.moving_average_gated_attention import MovingAverageGatedAttention
+
+
+layers = tfm.nlp.layers
+
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class MegaEncoder(tf_keras.layers.Layer):
+ """MegaEncoder.
+
+ Args:
+ vocab_size: The size of the token vocabulary.
+ embedding_width: The number of embedding dimensions.
+ intermediate_size: The number of dimension for MLP layers.
+ num_layers: The number of transformer layers.
+ max_sequence_length: The maximum sequence length that this encoder can
+ consume. If None, max_sequence_length uses the value from sequence length.
+ This determines the variable shape for positional embeddings.
+ type_vocab_size: The number of types that the 'type_ids' input can take.
+ zdim: hidden dimension for gates used in MEGA Layer.
+ hdim: hidden dimension used in MEGA Layer.
+ ndim: number of EMA used in MEGA layer.
+ activation: The activation for the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ bidirectional: Whether to use bidirectional EMA.
+ dropout: Dropout probability for the post-attention and output dropout.
+ attention_dropout: The dropout rate to use for the attention layers within
+ the transformer layers.
+ hidden_dropout: The dropout rate to use for hidden states in MEGA.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ initializer: The initialzer to use for all weights in this encoder.
+ output_range: The sequence output range, [0, output_range), by slicing the
+ target sequence of the last transformer layer. `None` means the entire
+ target sequence will attend to the source sequence, which yields the full
+ output.
+ embedding_layer: An optional Layer instance which will be called to generate
+ embeddings for the input word IDs.
+ norm_first: Whether to normalize inputs to attention and intermediate dense
+ layers. If set False, output of attention and intermediate dense layers is
+ normalized.
+ """
+
+ def __init__(
+ self,
+ vocab_size: int,
+ embedding_width: int = 128,
+ intermediate_size: int = 256,
+ num_layers: int = 12,
+ max_sequence_length: int = 512,
+ type_vocab_size: int = 16,
+ zdim: int = 64,
+ hdim: int = 256,
+ ndim: int = 16,
+ activation='silu',
+ bidirectional=False,
+ dropout: float = 0.0,
+ attention_dropout: float = 0.0,
+ hidden_dropout: float = 0.0,
+ inner_activation: Callable[..., Any] = _approx_gelu,
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
+ stddev=0.02
+ ),
+ output_range: Optional[int] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
+ norm_first: bool = False,
+ hidden_size: Optional[int] = None,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ # Mega args
+ initializer = tf_keras.initializers.get(initializer)
+
+ if embedding_layer is None:
+ self._embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ initializer=initializer,
+ name='word_embeddings',
+ )
+ else:
+ self._embedding_layer = embedding_layer
+
+ self._position_embedding_layer = layers.PositionEmbedding(
+ initializer=initializer,
+ max_length=max_sequence_length,
+ name='position_embedding',
+ )
+
+ self._type_embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=type_vocab_size,
+ embedding_width=embedding_width,
+ initializer=initializer,
+ use_one_hot=True,
+ name='type_embeddings',
+ )
+
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
+ name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32
+ )
+
+ self._embedding_dropout = tf_keras.layers.Dropout(
+ rate=dropout, name='embedding_dropout'
+ )
+
+ self._transformer_layers = []
+ self._attention_mask_layer = layers.SelfAttentionMask(
+ name='self_attention_mask'
+ )
+ for _ in range(num_layers):
+ layer = MovingAverageGatedAttention(
+ embed_dim=embedding_width,
+ zdim=zdim,
+ hdim=hdim,
+ ndim=ndim,
+ intermediate_size=intermediate_size,
+ inner_activation=inner_activation,
+ dropout=dropout,
+ attention_dropout=attention_dropout,
+ hidden_dropout=hidden_dropout,
+ activation=activation,
+ bidirectional=bidirectional,
+ prenorm=norm_first,
+ max_positions=max_sequence_length,
+ use_bias=True,
+ return_attention_scores=False,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ )
+ self._transformer_layers.append(layer)
+ self._num_layers = num_layers
+ self._pooler_layer = tf_keras.layers.Dense(
+ units=embedding_width,
+ activation='silu',
+ kernel_initializer=initializer,
+ name='pooler_transform',
+ )
+ self._config = {
+ 'vocab_size': vocab_size,
+ 'num_layers': num_layers,
+ 'max_sequence_length': max_sequence_length,
+ 'type_vocab_size': type_vocab_size,
+ 'zdim': zdim,
+ 'hdim': hdim,
+ 'ndim': ndim,
+ 'activation': activation,
+ 'bidirectional': bidirectional,
+ 'dropout': dropout,
+ 'attention_dropout': attention_dropout,
+ 'hidden_dropout': hidden_dropout,
+ 'inner_activation': tf_keras.activations.serialize(inner_activation),
+ 'initializer': tf_keras.initializers.serialize(initializer),
+ 'output_range': output_range,
+ 'embedding_width': embedding_width,
+ 'embedding_layer': embedding_layer,
+ 'norm_first': norm_first,
+ }
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ )
+
+ def call(self, inputs):
+ word_embeddings = None
+
+ if isinstance(inputs, dict):
+ if 'input_word_ids' in inputs.keys():
+ word_ids = inputs.get('input_word_ids')
+ mask = inputs.get('input_mask')
+ type_ids = inputs.get('input_type_ids', None)
+ word_embeddings = inputs.get('input_word_embeddings', None)
+ elif 'left_word_ids' in inputs.keys():
+ word_ids = inputs.get('left_word_ids')
+ mask = inputs.get('left_mask')
+ elif 'right_word_ids' in inputs.keys():
+ word_ids = inputs.get('right_word_ids')
+ mask = inputs.get('right_mask')
+ dense_inputs = inputs.get('dense_inputs', None)
+ dense_mask = inputs.get('dense_mask', None)
+ elif isinstance(inputs, list):
+ ## Dual Encoder Tasks
+ word_ids, mask = inputs
+ type_ids = None
+ dense_inputs, dense_mask = None, None
+ else:
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+
+ if type_ids is None: # pyrefly: ignore[unbound-name]
+ type_ids = tf.zeros_like(mask) # pyrefly: ignore[unbound-name]
+
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids) # pyrefly: ignore[unbound-name]
+
+ if dense_inputs is not None:
+ mask = tf.concat([mask, dense_mask], axis=1)
+
+ embeddings = self._embedding_norm_layer(word_embeddings)
+ embeddings = self._embedding_dropout(embeddings)
+
+ encoder_outputs = []
+ x = embeddings
+
+ for l in range(self._num_layers):
+ if x.shape[0] is None:
+ pass
+ else:
+ x = self._transformer_layers[l]([x, mask])
+ encoder_outputs.append(x)
+
+ last_encoder_output = encoder_outputs[-1]
+ avg_token_tensor = tf.math.reduce_mean(last_encoder_output, axis=1)
+ pooled_output = self._pooler_layer(avg_token_tensor)
+
+ output = dict(
+ sequence_output=encoder_outputs[-1],
+ pooled_output=pooled_output,
+ encoder_outputs=encoder_outputs,
+ )
+
+ return output
+
+ def get_embedding_table(self):
+ return self._embedding_layer.embeddings
+
+ def get_embedding_layer(self):
+ return self._embedding_layer
+
+ def get_config(self):
+ return dict(self._config)
+
+ @property
+ def transformer_layers(self):
+ """List of Transformer layers in the encoder."""
+ return self._transformer_layers
+
+ @property
+ def pooler_layer(self):
+ """The pooler dense layer after the transformer layers."""
+ return self._pooler_layer
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ if 'embedding_layer' in config and config['embedding_layer'] is not None:
+ warn_string = (
+ 'You are reloading a model that was saved with a '
+ 'potentially-shared embedding layer object. If you contine to '
+ 'train this model, the embedding layer will no longer be shared. '
+ 'To work around this, load the model outside of the Keras API.'
+ )
+ print('WARNING: ' + warn_string)
+ logging.warn(warn_string)
+
+ return cls(**config)
+
+ def _get_embeddings(
+ self,
+ word_ids: tf.Tensor,
+ type_ids: tf.Tensor,
+ word_embeddings: Optional[tf.Tensor],
+ dense_inputs: Optional[tf.Tensor],
+ dense_type_ids: Optional[tf.Tensor],
+ ) -> tf.Tensor:
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ # Concat the dense embeddings at sequence end.
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+
+ type_embeddings = self._type_embedding_layer(type_ids)
+
+ # absolute position embeddings.
+ position_embeddings = self._position_embedding_layer(word_embeddings)
+ return word_embeddings + position_embeddings + type_embeddings
diff --git a/official/projects/lra/mega_encoder_test.py b/official/projects/lra/mega_encoder_test.py
new file mode 100644
index 00000000000..52b327750ee
--- /dev/null
+++ b/official/projects/lra/mega_encoder_test.py
@@ -0,0 +1,51 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for official.nlp.projects.lra.mega_encoder."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.lra import mega_encoder
+
+
+class MegaEncoderTest(tf.test.TestCase):
+
+ def test_encoder(self):
+ sequence_length = 1024
+ batch_size = 2
+ vocab_size = 1024
+ network = mega_encoder.MegaEncoder(
+ num_layers=1,
+ vocab_size=1024,
+ max_sequence_length=4096,
+ )
+ word_id_data = np.random.randint(
+ vocab_size, size=(batch_size, sequence_length)
+ )
+ mask_data = np.random.randint(2, size=(batch_size, sequence_length))
+ type_id_data = np.random.randint(2, size=(batch_size, sequence_length))
+ outputs = network({
+ "input_word_ids": word_id_data,
+ "input_mask": mask_data,
+ "input_type_ids": type_id_data,
+ })
+ self.assertEqual(
+ outputs["sequence_output"].shape,
+ (batch_size, sequence_length, 128),
+ )
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/lra/mega_experiments.py b/official/projects/lra/mega_experiments.py
new file mode 100644
index 00000000000..6bf0fa970d7
--- /dev/null
+++ b/official/projects/lra/mega_experiments.py
@@ -0,0 +1,153 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Mega experiments."""
+# pylint: disable=g-doc-return-or-yield,line-too-long
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.nlp.configs import encoders
+from official.nlp.data import sentence_prediction_dataloader
+from official.nlp.tasks import sentence_prediction
+from official.projects.lra import lra_dual_encoder_dataloader
+from official.projects.lra import lra_dual_encoder_task
+from official.projects.lra.mega import MegaEncoderConfig
+
+AdamWeightDecay = optimization.AdamWeightDecayConfig
+PolynomialLr = optimization.PolynomialLrConfig
+PolynomialWarmupConfig = optimization.PolynomialWarmupConfig
+
+_TRAINER = cfg.TrainerConfig(
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.01,
+ 'exclude_from_weight_decay': [
+ 'LayerNorm',
+ 'layer_norm',
+ 'bias',
+ ],
+ },
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 1e-7,
+ 'end_learning_rate': 0.0,
+ },
+ },
+ 'warmup': {'type': 'polynomial'},
+ })
+)
+
+
+@exp_factory.register_config_factory('mega/lra_listops')
+def mega_listops() -> cfg.ExperimentConfig:
+ """Mega lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=MegaEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('mega/lra_imdb')
+def mega_imdb() -> cfg.ExperimentConfig:
+ """Mega lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=MegaEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('mega/lra_cifar')
+def mega_cifar() -> cfg.ExperimentConfig:
+ """Mega lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=MegaEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('mega/lra_pathfinder')
+def mega_pathfinder() -> cfg.ExperimentConfig:
+ """Mega lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=MegaEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('mega/lra_aan')
+def mega_aan() -> cfg.ExperimentConfig:
+ """Mega LRA task."""
+ config = cfg.ExperimentConfig(
+ task=lra_dual_encoder_task.DualEncoderConfig(
+ model=lra_dual_encoder_task.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=MegaEncoderConfig()
+ )
+ ),
+ train_data=lra_dual_encoder_dataloader.DualEncoderDataConfig(),
+ validation_data=lra_dual_encoder_dataloader.DualEncoderDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
diff --git a/official/projects/lra/moving_average_gated_attention.py b/official/projects/lra/moving_average_gated_attention.py
new file mode 100644
index 00000000000..0935b2da3b0
--- /dev/null
+++ b/official/projects/lra/moving_average_gated_attention.py
@@ -0,0 +1,350 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras-based MegaEncoder block layer."""
+
+from typing import Any
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.projects.lra.exponential_moving_average import MultiHeadEMA
+
+
+def get_activation_fn(activation):
+ ## Helper Function for Activation
+ if activation == "silu":
+ return tf.nn.silu
+ elif activation == "softmax":
+ return tf.nn.softmax
+ else:
+ raise NotImplementedError
+ return
+
+
+class RelativePositionBias(tf_keras.layers.Layer):
+ """Relative position embedding layer with bias."""
+
+ def __init__(self, max_positions):
+ super().__init__()
+ self.max_positions = max_positions
+
+ def build(self, input_shape):
+ gauss_init = tf_keras.initializers.RandomNormal(mean=0.0, stddev=0.02)
+ self.rel_pos_bias = tf.Variable(
+ gauss_init(shape=[2 * self.max_positions - 1], dtype=tf.float32),
+ trainable=True,
+ )
+
+ def call(self, seq_len):
+ if seq_len is None:
+ seq_len = self.max_positions
+ seq_len = tf.get_static_value(seq_len)
+ # seq_len * 2 -1
+ b = self.rel_pos_bias[
+ (self.max_positions - seq_len) : (self.max_positions + seq_len - 1)
+ ]
+ # seq_len * 3 - 1
+ t = tf.pad(b, paddings=tf.constant([[0, seq_len]]))
+ # (seq_len * 3 - 1) * seq_len
+ t = tf.tile(t, (seq_len,))
+ t = t[:-seq_len]
+ # seq_len x (3 * seq_len - 2)
+ t = tf.reshape(t, shape=(seq_len, 3 * seq_len - 2))
+ r = (2 * seq_len - 1) // 2
+ start = r
+ end = t.shape[1] - r
+ t = t[:, start:end]
+ return t
+
+
+class MovingAverageGatedAttention(tf_keras.layers.Layer):
+ """MegaEncoderBlock layer.
+
+ This layer implements the Mega Encoder from
+ "Mega: Moving Average Equipped Gated Attention".
+ (https://arxiv.org/abs/2209.10655)
+ """
+
+ def __init__(
+ self,
+ embed_dim,
+ zdim,
+ hdim,
+ ndim,
+ intermediate_size,
+ inner_activation=None,
+ dropout=0.0,
+ attention_dropout=0.0,
+ hidden_dropout=0.0,
+ activation="silu",
+ bidirectional=False,
+ truncation=None,
+ prenorm=True,
+ max_positions=1024,
+ use_bias=True,
+ kernel_initializer="glorot_uniform",
+ bias_initializer="zeros",
+ attention_initializer=None,
+ attention_axes=None,
+ return_attention_scores=False,
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ activity_regularizer=None,
+ kernel_constraint=None,
+ bias_constraint=None,
+ ):
+ self.embed_dim = embed_dim
+ self.hdim = hdim
+ self.zdim = zdim
+ self.ndim = ndim
+ self.inner_dim = intermediate_size
+ self.activation = get_activation_fn(activation=activation)
+ self.inner_activation = inner_activation
+ self.scaling = self.zdim**-0.5
+
+ self.dropout = tf_keras.layers.Dropout(rate=dropout)
+ self.hidden_dropout = tf_keras.layers.Dropout(rate=hidden_dropout)
+ self.attention_dropout_rate = attention_dropout
+ self.attention_dropout = tf_keras.layers.Dropout(rate=attention_dropout)
+
+ self.ffn_intermediate_dropout = tf_keras.layers.Dropout(rate=hidden_dropout)
+ self.output_dropout = tf_keras.layers.Dropout(rate=hidden_dropout)
+
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
+
+ if attention_initializer:
+ self._attention_initializer = tf_keras.initializers.get(
+ attention_initializer
+ )
+ else:
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer
+ )
+ self._attention_axes = attention_axes
+ self._use_bias = use_bias
+ self.return_attention_scores = return_attention_scores
+
+ self.prenorm = prenorm
+ self.norm = tf_keras.layers.LayerNormalization(axis=-1)
+ self.ffn_norm = tf_keras.layers.LayerNormalization(axis=-1)
+
+ self.move = MultiHeadEMA(
+ embed_dim, ndim=ndim, bidirectional=bidirectional, truncation=truncation
+ )
+
+ self.max_positions = max_positions
+ super().__init__()
+
+ def build(self, input_shape):
+ gauss_init = tf_keras.initializers.RandomNormal(mean=0.0, stddev=0.02)
+ zero_init = tf_keras.initializers.Zeros()
+
+ self.v_proj = tf_keras.layers.Dense(
+ self.hdim,
+ activation=None,
+ use_bias=True,
+ kernel_initializer=tf_utils.clone_initializer(gauss_init),
+ bias_initializer=tf_utils.clone_initializer(zero_init),
+ name="v_proj",
+ )
+
+ self.mx_proj = tf_keras.layers.Dense(
+ self.zdim + self.hdim + 2 * self.embed_dim,
+ activation=None,
+ use_bias=True,
+ kernel_initializer=tf_utils.clone_initializer(gauss_init),
+ bias_initializer=tf_utils.clone_initializer(zero_init),
+ name="mx_proj",
+ )
+
+ self.h_proj = tf_keras.layers.Dense(
+ self.embed_dim,
+ activation=None,
+ use_bias=True,
+ kernel_initializer=tf_utils.clone_initializer(gauss_init),
+ bias_initializer=tf_utils.clone_initializer(zero_init),
+ name="h_proj",
+ )
+
+ self._intermediate_dense = tf_keras.layers.Dense(
+ self.inner_dim, use_bias=True
+ )
+
+ self._output_dense = tf_keras.layers.Dense(self.embed_dim, use_bias=True)
+
+ policy = tf_keras.mixed_precision.global_policy()
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
+ self.inner_activation, dtype=policy
+ )
+
+ self.gamma = tf.Variable(
+ gauss_init(shape=[2, self.zdim], dtype=tf.float32), trainable=True
+ )
+ self.beta = tf.Variable(
+ zero_init(shape=[2, self.zdim], dtype=tf.float32), trainable=True
+ )
+
+ self.rel_pos_bias = RelativePositionBias(max_positions=self.max_positions)
+
+ super().build(input_shape)
+
+ def get_config(self):
+ base_config = super().get_config()
+ base_config.update({
+ "embed_dim": self.embed_dim,
+ "zdim": self.zdim,
+ "hdim": self.hdim,
+ "dropout": self.dropout,
+ "attention_dropout": self.attention_dropout_rate,
+ "kernel_initializer": tf_keras.initializers.serialize(
+ self._kernel_initializer
+ ),
+ "bias_initializer": tf_keras.initializers.serialize(
+ self._bias_initializer
+ ),
+ "use_bias": self._use_bias,
+ "prenorm": self.prenorm,
+ "max_positions": self.max_positions,
+ "attention_initializer": tf_keras.initializers.serialize(
+ self._attention_initializer
+ ),
+ "attention_axes": self._attention_axes,
+ "return_attention_scores": self.return_attention_scores,
+ })
+ return base_config
+
+ def _softmax_attention(self, q, k):
+ slen = k.shape[1]
+ # C x C
+ if slen is None:
+ slen = 2
+ bias = self.rel_pos_bias(slen)
+
+ # scaled attention
+ q = q * self.scaling
+ # B x K x C x C
+ qk = tf.matmul(q, tf.transpose(k, perm=(0, 2, 1))) + bias
+
+ attn_weights = tf.nn.softmax(qk, axis=-1)
+ return attn_weights
+
+ def call(self, inputs: Any) -> Any:
+ """MEGA encoder block call.
+
+ Args:
+ inputs: a single tensor or a list of tensors. `input tensor`
+ as the single sequence of embeddings. [`input tensor`,
+ `attention mask`] to have the
+ additional attention mask. [`query tensor`, `key value tensor`,
+ `attention mask`] to have separate input streams for the query, and
+ key/value to the multi-head attention.
+ Returns:
+ An output tensor with the same dimensions as input/query tensor.
+ """
+ if isinstance(inputs, (list, tuple)):
+ if len(inputs) == 2:
+ (input_tensor, attention_mask) = inputs
+ key_value = None
+ elif len(inputs) == 3:
+ (input_tensor, key_value, attention_mask) = inputs
+ else:
+ raise ValueError(
+ "Unexpected inputs to %s with length at %d"
+ % (self.__class__, len(inputs))
+ )
+ else:
+ (input_tensor, key_value, attention_mask) = (inputs, None, None)
+
+ if self.prenorm:
+ input_tensor = self.norm(input_tensor)
+ if key_value is not None:
+ key_value = self.norm(key_value)
+
+ ## B*L*D -> L*B*D
+ ## Multi-Dimensional Damped EMA
+ x = tf.transpose(input_tensor, perm=[1, 0, 2])
+ residual = x
+
+ seq_len, bsz, _ = x.shape
+
+ # L x B x E
+ v = self.activation(self.v_proj(x))
+
+ # L x B x D
+ mx = self.move(x, attention_mask)
+ mx = self.dropout(mx)
+
+ # L x B x D -> L x B x (2*D+S+E)
+ base = self.mx_proj(mx)
+
+ u, zr, hx = tf.split(
+ base, [self.embed_dim, self.zdim + self.hdim, self.embed_dim], axis=-1
+ )
+ # L x B x D
+ u = tf.math.sigmoid(u)
+ # L x B x (E+S)
+ z, r = tf.split(tf.nn.silu(zr), [self.zdim, self.hdim], axis=-1)
+ # L x B x S -> L x B x 1 x S -> L x B x 2 x S
+ z = tf.expand_dims(z, axis=2) * self.gamma + self.beta
+ # L x B x 2 x S -> L x B x S
+ q, k = tf.unstack(z, axis=2)
+
+ # L x B x D -> B x L x D
+ q = tf.transpose(q, perm=(1, 0, 2))
+ k = tf.transpose(k, perm=(1, 0, 2))
+ # L x B x E -> B x L x E
+ v = tf.transpose(v, perm=(1, 0, 2))
+
+ attn_weights = self._softmax_attention(q, k)
+ v = self.hidden_dropout(v)
+ kernel = tf.squeeze(self.attention_dropout(attn_weights))
+ # B x K x C x E -> B x L x E -> L x B x E
+ h = tf.transpose(
+ tf.reshape(
+ tf.linalg.matmul(kernel, v), shape=(bsz, seq_len, self.hdim)
+ ),
+ perm=(1, 0, 2),
+ )
+
+ # L x B x E -> L x B x D
+ h = self.activation(hx + self.h_proj(h * r))
+ h = self.dropout(h)
+ # L x B x D
+ out = residual + tf.math.multiply(u, h - residual)
+
+ if not self.prenorm:
+ out = self.norm(out)
+
+ out = tf.transpose(out, perm=(1, 0, 2))
+
+ if self.prenorm:
+ out = self.ffn_norm(out)
+
+ inner_output = self._intermediate_dense(out)
+ inner_output = self._intermediate_activation_layer(inner_output)
+ inner_output = self.ffn_intermediate_dropout(inner_output)
+ layer_output = self._output_dense(inner_output)
+ layer_output = self.output_dropout(layer_output) + out
+
+ if not self.prenorm:
+ layer_output = self.ffn_norm(layer_output)
+
+ return layer_output
diff --git a/official/projects/lra/train.py b/official/projects/lra/train.py
new file mode 100644
index 00000000000..01166b60584
--- /dev/null
+++ b/official/projects/lra/train.py
@@ -0,0 +1,74 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""A customized training library for the specific task."""
+
+from absl import app
+from absl import flags
+import gin
+
+from official.common import distribute_utils
+from official.common import flags as tfm_flags
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+from official.projects.lra import linformer_experiments # pylint:disable=unused-import
+from official.projects.lra import mega_experiments # pylint:disable=unused-import
+from official.projects.lra import transformer_experiments # pylint:disable=unused-import
+
+
+FLAGS = flags.FLAGS
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ # Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
+ # dtype is float16
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu,
+ **params.runtime.model_parallelism()
+ )
+
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir,
+ )
+
+ train_utils.save_gin_config(FLAGS.mode, model_dir)
+
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(main)
diff --git a/official/projects/lra/transformer.py b/official/projects/lra/transformer.py
new file mode 100644
index 00000000000..05884cda36e
--- /dev/null
+++ b/official/projects/lra/transformer.py
@@ -0,0 +1,61 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Longformer model configurations and instantiation methods."""
+import dataclasses
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.modeling.hyperparams import base_config
+from official.nlp.configs import encoders
+from official.projects.lra.transformer_encoder import TransformerEncoder
+
+
+@dataclasses.dataclass
+class TransformerEncoderConfig(encoders.BertEncoderConfig):
+ """Extra paramerters for Transformer configs.
+
+ Attributes: For in-place usage only
+ """
+
+
+@base_config.bind(TransformerEncoderConfig)
+def get_encoder(encoder_cfg: TransformerEncoderConfig):
+ """Gets a 'TransformerEncoder' object.
+
+ Args:
+ encoder_cfg: A 'TransformerEncoderConfig'.
+
+ Returns:
+ A encoder object.
+ """
+ encoder = TransformerEncoder(
+ vocab_size=encoder_cfg.vocab_size,
+ hidden_size=encoder_cfg.hidden_size,
+ num_layers=encoder_cfg.num_layers,
+ num_attention_heads=encoder_cfg.num_attention_heads,
+ inner_dim=encoder_cfg.intermediate_size,
+ inner_activation=tf_utils.get_activation(encoder_cfg.hidden_activation),
+ output_dropout=encoder_cfg.dropout_rate,
+ attention_dropout=encoder_cfg.attention_dropout_rate,
+ max_sequence_length=encoder_cfg.max_position_embeddings,
+ type_vocab_size=encoder_cfg.type_vocab_size,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range
+ ),
+ output_range=encoder_cfg.output_range,
+ embedding_width=encoder_cfg.embedding_size,
+ norm_first=encoder_cfg.norm_first,
+ )
+ return encoder
diff --git a/official/projects/lra/transformer_encoder.py b/official/projects/lra/transformer_encoder.py
new file mode 100644
index 00000000000..ac1218e76ef
--- /dev/null
+++ b/official/projects/lra/transformer_encoder.py
@@ -0,0 +1,310 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Transformer encoder."""
+
+# pylint: disable=g-classes-have-attributes
+
+from typing import Any, Callable, Optional, Union
+
+from absl import logging
+import tensorflow as tf, tf_keras
+import tensorflow_models as tfm
+
+from official.modeling import tf_utils
+
+layers = tfm.nlp.layers
+
+_Initializer = Union[str, tf_keras.initializers.Initializer]
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
+
+
+class TransformerEncoder(tf_keras.layers.Layer):
+ """TransformerEncoder.
+
+ Args:
+ vocab_size: The size of the token vocabulary.
+ pad_token_id: the token id for the pad token
+ hidden_size: The size of the transformer hidden layers.
+ num_layers: The number of transformer layers.
+ num_attention_heads: The number of attention heads for each transformer. The
+ hidden size must be divisible by the number of attention heads.
+ max_sequence_length: The maximum sequence length that this encoder can
+ consume. If None, max_sequence_length uses the value from sequence length.
+ This determines the variable shape for positional embeddings.
+ type_vocab_size: The number of types that the 'type_ids' input can take.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network for each transformer.
+ output_dropout: Dropout probability for the post-attention and output
+ dropout.
+ attention_dropout: The dropout rate to use for the attention layers within
+ the transformer layers.
+ initializer: The initialzer to use for all weights in this encoder.
+ output_range: The sequence output range, [0, output_range), by slicing the
+ target sequence of the last transformer layer. `None` means the entire
+ target sequence will attend to the source sequence, which yields the full
+ output.
+ embedding_width: The width of the word embeddings. If the embedding width is
+ not equal to hidden size, embedding parameters will be factorized into two
+ matrices in the shape of ['vocab_size', 'embedding_width'] and
+ ['embedding_width', 'hidden_size'] ('embedding_width' is usually much
+ smaller than 'hidden_size').
+ embedding_layer: An optional Layer instance which will be called to generate
+ embeddings for the input word IDs.
+ norm_first: Whether to normalize inputs to attention and intermediate dense
+ layers. If set False, output of attention and intermediate dense layers is
+ normalized.
+ """
+
+ def __init__(
+ self,
+ vocab_size: int,
+ hidden_size: int = 768,
+ num_layers: int = 12,
+ num_attention_heads: int = 12,
+ max_sequence_length: int = 512,
+ type_vocab_size: int = 16,
+ inner_dim: int = 3072,
+ inner_activation: Callable[..., Any] = _approx_gelu,
+ output_dropout: float = 0.1,
+ attention_dropout: float = 0.1,
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
+ stddev=0.02
+ ),
+ output_range: Optional[int] = None,
+ embedding_width: Optional[int] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
+ norm_first: bool = False,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
+
+ if embedding_width is None:
+ embedding_width = hidden_size
+
+ if embedding_layer is None:
+ self._embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=embedding_width,
+ initializer=initializer,
+ name='word_embeddings',
+ )
+ else:
+ self._embedding_layer = embedding_layer
+
+ self._position_embedding_layer = layers.PositionEmbedding(
+ initializer=initializer,
+ max_length=max_sequence_length,
+ name='position_embedding',
+ )
+
+ self._type_embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=type_vocab_size,
+ embedding_width=embedding_width,
+ initializer=initializer,
+ use_one_hot=True,
+ name='type_embeddings',
+ )
+
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
+ name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32
+ )
+
+ self._embedding_dropout = tf_keras.layers.Dropout(
+ rate=output_dropout, name='embedding_dropout'
+ )
+
+ # We project the 'embedding' output to 'hidden_size' if it is not already
+ # 'hidden_size'.
+ self._embedding_projection = None
+ if embedding_width != hidden_size:
+ self._embedding_projection = tf_keras.layers.EinsumDense(
+ '...x,xy->...y',
+ output_shape=hidden_size,
+ bias_axes='y',
+ kernel_initializer=initializer,
+ name='embedding_projection',
+ )
+
+ self._transformer_layers = []
+ self._attention_mask_layer = layers.SelfAttentionMask(
+ name='self_attention_mask'
+ )
+ for i in range(num_layers):
+ layer = layers.TransformerEncoderBlock(
+ num_attention_heads=num_attention_heads,
+ inner_dim=inner_dim,
+ inner_activation=inner_activation,
+ output_dropout=output_dropout,
+ attention_dropout=attention_dropout,
+ norm_first=norm_first,
+ return_attention_scores=False,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
+ name='transformer/layer_%d' % i,
+ )
+ self._transformer_layers.append(layer)
+ self._num_layers = num_layers
+
+ self._pooler_layer = tf_keras.layers.Dense(
+ units=hidden_size,
+ activation='tanh',
+ kernel_initializer=initializer,
+ name='pooler_transform',
+ )
+ self._config = {
+ 'vocab_size': vocab_size,
+ 'hidden_size': hidden_size,
+ 'num_layers': num_layers,
+ 'num_attention_heads': num_attention_heads,
+ 'max_sequence_length': max_sequence_length,
+ 'type_vocab_size': type_vocab_size,
+ 'inner_dim': inner_dim,
+ 'inner_activation': tf_keras.activations.serialize(activation),
+ 'output_dropout': output_dropout,
+ 'attention_dropout': attention_dropout,
+ 'initializer': tf_keras.initializers.serialize(initializer),
+ 'output_range': output_range,
+ 'embedding_width': embedding_width,
+ 'embedding_layer': embedding_layer,
+ 'norm_first': norm_first,
+ }
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ )
+
+ def call(self, inputs):
+ word_embeddings = None
+
+ if isinstance(inputs, dict):
+ if 'input_word_ids' in inputs.keys():
+ word_ids = inputs.get('input_word_ids')
+ mask = inputs.get('input_mask')
+ type_ids = inputs.get('input_type_ids', None)
+ word_embeddings = inputs.get('input_word_embeddings', None)
+ elif 'left_word_ids' in inputs.keys():
+ word_ids = inputs.get('left_word_ids')
+ mask = inputs.get('left_mask')
+ elif 'right_word_ids' in inputs.keys():
+ word_ids = inputs.get('right_word_ids')
+ mask = inputs.get('right_mask')
+ dense_inputs = inputs.get('dense_inputs', None)
+ dense_mask = inputs.get('dense_mask', None)
+ dense_type_ids = inputs.get('dense_type_ids', None)
+ elif isinstance(inputs, list):
+ ## Dual Encoder Tasks
+ word_ids, mask = inputs
+ type_ids = None
+ dense_inputs, dense_mask, dense_type_ids = None, None, None
+ else:
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+
+ if type_ids is None: # pyrefly: ignore[unbound-name]
+ type_ids = tf.zeros_like(mask) # pyrefly: ignore[unbound-name]
+
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids) # pyrefly: ignore[unbound-name]
+
+ if dense_inputs is not None:
+ mask = tf.concat([mask, dense_mask], axis=1)
+
+ embeddings = self._get_embeddings(
+ word_ids, type_ids, word_embeddings, dense_inputs, dense_type_ids # pyrefly: ignore[bad-argument-type]
+ )
+ embeddings = self._embedding_norm_layer(embeddings)
+ embeddings = self._embedding_dropout(embeddings)
+
+ if self._embedding_projection is not None:
+ embeddings = self._embedding_projection(embeddings)
+
+ attention_mask = self._attention_mask_layer(embeddings, mask)
+
+ encoder_outputs = []
+ x = embeddings
+ for layer in self._transformer_layers:
+ x = layer([x, attention_mask])
+ encoder_outputs.append(x)
+
+ last_encoder_output = encoder_outputs[-1]
+ first_token_tensor = last_encoder_output[:, 0, :]
+ pooled_output = self._pooler_layer(first_token_tensor)
+
+ output = dict(
+ sequence_output=encoder_outputs[-1],
+ pooled_output=pooled_output,
+ encoder_outputs=encoder_outputs,
+ )
+
+ return output
+
+ def get_embedding_table(self):
+ return self._embedding_layer.embeddings
+
+ def get_embedding_layer(self):
+ return self._embedding_layer
+
+ def get_config(self):
+ return dict(self._config)
+
+ @property
+ def transformer_layers(self):
+ """List of Transformer layers in the encoder."""
+ return self._transformer_layers
+
+ @property
+ def pooler_layer(self):
+ """The pooler dense layer after the transformer layers."""
+ return self._pooler_layer
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ if 'embedding_layer' in config and config['embedding_layer'] is not None:
+ warn_string = (
+ 'You are reloading a model that was saved with a '
+ 'potentially-shared embedding layer object. If you contine to '
+ 'train this model, the embedding layer will no longer be shared. '
+ 'To work around this, load the model outside of the Keras API.'
+ )
+ print('WARNING: ' + warn_string)
+ logging.warn(warn_string)
+
+ return cls(**config)
+
+ def _get_embeddings(
+ self,
+ word_ids: tf.Tensor,
+ type_ids: tf.Tensor,
+ word_embeddings: Optional[tf.Tensor],
+ dense_inputs: Optional[tf.Tensor],
+ dense_type_ids: Optional[tf.Tensor],
+ ) -> tf.Tensor:
+ if word_embeddings is None:
+ word_embeddings = self._embedding_layer(word_ids)
+
+ if dense_inputs is not None:
+ # Concat the dense embeddings at sequence end.
+ word_embeddings = tf.concat([word_embeddings, dense_inputs], axis=1)
+ type_ids = tf.concat([type_ids, dense_type_ids], axis=1)
+
+ type_embeddings = self._type_embedding_layer(type_ids)
+
+ # absolute position embeddings.
+ position_embeddings = self._position_embedding_layer(word_embeddings)
+ return word_embeddings + position_embeddings + type_embeddings
diff --git a/official/projects/lra/transformer_experiments.py b/official/projects/lra/transformer_experiments.py
new file mode 100644
index 00000000000..99f0183538b
--- /dev/null
+++ b/official/projects/lra/transformer_experiments.py
@@ -0,0 +1,154 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Transformer experiments."""
+# pylint: disable=g-doc-return-or-yield,line-too-long
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.nlp.configs import encoders
+from official.nlp.data import sentence_prediction_dataloader
+from official.nlp.tasks import sentence_prediction
+
+from official.projects.lra import lra_dual_encoder_dataloader
+from official.projects.lra import lra_dual_encoder_task
+from official.projects.lra.transformer import TransformerEncoderConfig
+
+AdamWeightDecay = optimization.AdamWeightDecayConfig
+PolynomialLr = optimization.PolynomialLrConfig
+PolynomialWarmupConfig = optimization.PolynomialWarmupConfig
+
+_TRAINER = cfg.TrainerConfig(
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.01,
+ 'exclude_from_weight_decay': [
+ 'LayerNorm',
+ 'layer_norm',
+ 'bias',
+ ],
+ },
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 3e-5,
+ 'end_learning_rate': 0.0,
+ },
+ },
+ 'warmup': {'type': 'polynomial'},
+ })
+)
+
+
+@exp_factory.register_config_factory('transformer/lra_listops')
+def transformer_listops() -> cfg.ExperimentConfig:
+ """Transformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=TransformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('transformer/lra_imdb')
+def transformer_imdb() -> cfg.ExperimentConfig:
+ """Transformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=TransformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('transformer/lra_cifar')
+def transformer_cifar() -> cfg.ExperimentConfig:
+ """Transformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=TransformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('transformer/lra_pathfinder')
+def transformer_pathfinder() -> cfg.ExperimentConfig:
+ """Transformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=sentence_prediction.SentencePredictionConfig(
+ model=sentence_prediction.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=TransformerEncoderConfig()
+ )
+ ),
+ train_data=sentence_prediction_dataloader.SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
+
+
+@exp_factory.register_config_factory('transformer/lra_aan')
+def transformer_aan() -> cfg.ExperimentConfig:
+ """Transformer lra fine-tuning."""
+ config = cfg.ExperimentConfig(
+ task=lra_dual_encoder_task.DualEncoderConfig(
+ model=lra_dual_encoder_task.ModelConfig(
+ encoder=encoders.EncoderConfig(
+ type='any', any=TransformerEncoderConfig()
+ )
+ ),
+ train_data=lra_dual_encoder_dataloader.DualEncoderDataConfig(),
+ validation_data=lra_dual_encoder_dataloader.DualEncoderDataConfig(
+ is_training=False, drop_remainder=False
+ ),
+ ),
+ trainer=_TRAINER,
+ )
+ return config
diff --git a/official/projects/mae/README.md b/official/projects/mae/README.md
new file mode 100644
index 00000000000..424d41328b7
--- /dev/null
+++ b/official/projects/mae/README.md
@@ -0,0 +1,42 @@
+# Masked Autoencoders Are Scalable Vision Learners (MAE)
+
+TF2 implementation of [MAE](https://arxiv.org/abs/2111.06377).
+
+## Imagenet pretrain
+
+Model | reolution | pathch size | batch size | epochs | target pixel norm | val MSE
+------------ | ------: | ------: | -----:| -----:| -----: | --------:
+(a) ViT-L14 | 224x224 | 14 | 4096 | 800 | no | 0.2456
+(b) ViT-L14 | 224x224 | 14 | 4096 | 800 | yes | 0.3630
+(c) ViT-L16 | 224x224 | 16 | 4096 | 800 | yes | 0.3866
+
+
+## ImageNet linear probing
+
+Model | resolution | pathch size | base learning rate | batch size | init checkpoint | epochs | top1 Acc | dashboard
+------------ | :--------: | -----:| -----:| -----:| -----:| -----:| -----: | --------:
+ViT-L14 | 224x224 | 14 | 0.1 | 16384 | (b) | 90 | 72.8 | -
+ViT-L16 | 224x224 | 16 | 0.1 | 16384 | (c) | 90 | 73.0 | -
+ViT-L16 | 224x224 | 16 | 0.1 | 16384 | norm | 90 | 73.9 | Table 1 (d)
+
+## ImageNet finetune
+
+Model | resolution | pathch size | base learning rate | batch size | init checkpoint | epochs | top1 Acc | dashboard
+------------ | :--------: | -----:| -----:| -----:| -----:| -----:| -----: | -----:
+ViT-L14 | 224x224 | 14 | 0.001 | 1024 | (a) | 50 | 84.4 | -
+ViT-L14 | 224x224 | 14 | 0.001 | 1024 | (b) | 50 | 85.3 | -
+ViT-L14 | 224x224 | 14 | 0.00075 | 1024 |(b) | 50 | 85.4 | -
+ViT-L14 | 224x224 | 14 | 0.0001 | 4096| scratch | 200 | 82.4 | -
+ViT-L16 | 224x224 | 16 | 0.001 | 1024 | (c)| 50 | 84.9 | -
+ViT-L16 | 224x224 | 16 | 0.001 | 1024| no-norm | 50 | 84.9 | Table 1(d)
+ViT-L16 | 224x224 | 16 | 0.001 | 1024| norm | 50 | 85.4 | paper section 4.
+ViT-L16 | 224x224 | 16 | 0.0001 | 4096| scratch | 200 | 82.5 | paper section 4.
+
+
+## Known discrepancy with the paper:
+
+* ~-0.9 linear probing top1 acc (w/ norm) compared to paper results with patch
+ size 16.
+
+* ~-0.5 finetune top1 acc (w/ norm) compared to paper results with patch
+ size 16.
diff --git a/official/projects/mae/configs/linear_probe.py b/official/projects/mae/configs/linear_probe.py
new file mode 100644
index 00000000000..14b7f9d93aa
--- /dev/null
+++ b/official/projects/mae/configs/linear_probe.py
@@ -0,0 +1,87 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""ViT linear probing configurations."""
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.projects.mae.tasks import linear_probe
+from official.vision.configs import image_classification
+
+
+@exp_factory.register_config_factory('vit_imagenet_mae_linear_probe')
+def vit_imagenet_mae_linear_probe() -> cfg.ExperimentConfig:
+ """Config to get results that matches the paper."""
+ train_batch_size = 16384
+ eval_batch_size = 1024
+ imagenet_size = 1281167
+ steps_per_epoch = imagenet_size // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=linear_probe.ViTLinearProbeConfig( # pylint: disable=unexpected-keyword-arg
+ train_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='train',
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=20000,
+ ),
+ validation_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='validation',
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=False,
+ aug_rand_hflip=False,
+ ),
+ init_stochastic_depth_rate=0.0,
+ init_checkpoint='Please provide',
+ ),
+ trainer=cfg.TrainerConfig(
+ train_steps=90 * steps_per_epoch,
+ validation_steps=48,
+ steps_per_loop=100,
+ summary_interval=100,
+ checkpoint_interval=100,
+ validation_interval=100,
+ max_to_keep=1,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'lars',
+ 'lars': {
+ 'weight_decay_rate': 0.0,
+ 'momentum': 0.9,
+ },
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.1 * train_batch_size / 256,
+ 'decay_steps': 90 * steps_per_epoch,
+ },
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 10 * steps_per_epoch,
+ 'warmup_learning_rate': 0,
+ },
+ },
+ }),
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ],
+ )
+ return config
diff --git a/official/projects/mae/configs/mae.py b/official/projects/mae/configs/mae.py
new file mode 100644
index 00000000000..243cb860e75
--- /dev/null
+++ b/official/projects/mae/configs/mae.py
@@ -0,0 +1,107 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""MAE configurations."""
+
+import dataclasses
+from typing import Tuple
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.vision.configs import image_classification
+
+
+@dataclasses.dataclass
+class MAEConfig(cfg.TaskConfig):
+ """The translation task config."""
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
+ masking_ratio: float = 0.75
+ patch_h: int = 14
+ patch_w: int = 14
+ num_classes: int = 1000
+ input_size: Tuple[int, int] = (224, 224)
+ norm_target: bool = False
+
+
+@exp_factory.register_config_factory('mae_imagenet')
+def mae_imagenet() -> cfg.ExperimentConfig:
+ """Config to get results that matches the paper."""
+ train_batch_size = 4096
+ eval_batch_size = 4096
+ imagenet_size = 1281167
+ steps_per_epoch = imagenet_size // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=MAEConfig(
+ train_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='train',
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=10000,
+ crop_area_range=(0.2, 1.0),
+ ),
+ validation_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='validation',
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=False,
+ )
+ ),
+ trainer=cfg.TrainerConfig(
+ train_steps=800 * steps_per_epoch,
+ validation_steps=24,
+ steps_per_loop=1000,
+ summary_interval=1000,
+ checkpoint_interval=1000,
+ validation_interval=1000,
+ max_to_keep=5,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'beta_2': 0.95,
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm':
+ 0.0,
+ 'exclude_from_weight_decay': [
+ 'LayerNorm', 'layer_norm', 'bias']
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate':
+ 1.5 * 1e-4 * train_batch_size / 256,
+ 'decay_steps': 800 * steps_per_epoch
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 40 * steps_per_epoch,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ])
+ return config
diff --git a/official/projects/mae/configs/vit.py b/official/projects/mae/configs/vit.py
new file mode 100644
index 00000000000..c618fa56b5d
--- /dev/null
+++ b/official/projects/mae/configs/vit.py
@@ -0,0 +1,191 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""ViT configurations."""
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.mae import optimization
+from official.projects.mae.tasks import image_classification as vit
+from official.vision.configs import common
+from official.vision.configs import image_classification
+
+
+vars_substr = [
+ 'token_layer/cls', 'dense_1/kernel', 'vi_t_classifier/dense',
+ 'encoder/layer_normalization', 'encoder/transformer_encoder_block/'
+]
+
+layers_idx = [0, 0, 25, 24, 1]
+
+for i in range(1, 24):
+ vars_substr.append('encoder/transformer_encoder_block_%s/' % str(i))
+ layers_idx.append(i + 1)
+
+
+@exp_factory.register_config_factory('vit_imagenet_mae_finetune')
+def vit_imagenet_mae_finetune() -> cfg.ExperimentConfig:
+ """Config to get results that matches the paper."""
+ train_batch_size = 1024
+ eval_batch_size = 1024
+ imagenet_size = 1281167
+ steps_per_epoch = imagenet_size // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=vit.ViTConfig(
+ train_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='train',
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=10000,
+ aug_type=common.Augmentation(
+ type='randaug',
+ randaug=common.RandAugment(
+ magnitude=9,
+ magnitude_std=0.5,
+ exclude_ops=['Cutout', 'Invert'],
+ ),
+ ),
+ ),
+ validation_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='validation',
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=False,
+ aug_rand_hflip=False,
+ ),
+ init_stochastic_depth_rate=0.1,
+ init_checkpoint='Please provide',
+ ),
+ trainer=cfg.TrainerConfig(
+ train_steps=50 * steps_per_epoch,
+ validation_steps=48,
+ steps_per_loop=2000,
+ summary_interval=2000,
+ checkpoint_interval=2000,
+ validation_interval=2000,
+ max_to_keep=1,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'vit_adamw',
+ 'vit_adamw': {
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0,
+ 'beta_2': 0.999,
+ 'layer_decay': 0.75,
+ 'vars_substr': vars_substr,
+ 'layers_idx': layers_idx,
+ 'exclude_from_weight_decay': ['cls'],
+ },
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 1e-3 * train_batch_size / 256,
+ 'decay_steps': 50 * steps_per_epoch,
+ },
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 5 * steps_per_epoch,
+ 'warmup_learning_rate': 0,
+ },
+ },
+ }),
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ],
+ )
+ return config
+
+
+@exp_factory.register_config_factory('vit_imagenet_scratch')
+def vit_imagenet_scratch() -> cfg.ExperimentConfig:
+ """Config to get results that matches the paper."""
+ train_batch_size = 4096
+ eval_batch_size = 1024
+ imagenet_size = 1281167
+ steps_per_epoch = imagenet_size // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=vit.ViTConfig(
+ train_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='train',
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=10000,
+ aug_type=common.Augmentation(
+ type='randaug',
+ randaug=common.RandAugment(
+ magnitude=9,
+ magnitude_std=0.5,
+ exclude_ops=['Cutout', 'Invert'])
+ )
+ ),
+ validation_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='validation',
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=False,
+ aug_rand_hflip=False,
+ )
+ ),
+ trainer=cfg.TrainerConfig(
+ train_steps=200 * steps_per_epoch,
+ validation_steps=48,
+ steps_per_loop=1000,
+ summary_interval=1000,
+ checkpoint_interval=1000,
+ validation_interval=1000,
+ max_to_keep=1,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'vit_adamw',
+ 'vit_adamw': {
+ 'weight_decay_rate': 0.3,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0,
+ 'beta_2': 0.95,
+ 'exclude_from_weight_decay': ['cls']
+ }
+ },
+ 'ema': {
+ 'average_decay': 0.9999,
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate':
+ 1e-4 * train_batch_size / 256,
+ 'decay_steps': 200 * steps_per_epoch
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 20 * steps_per_epoch,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ])
+ return config
diff --git a/official/projects/mae/modeling/masked_ae.py b/official/projects/mae/modeling/masked_ae.py
new file mode 100644
index 00000000000..ba0dee9d0e2
--- /dev/null
+++ b/official/projects/mae/modeling/masked_ae.py
@@ -0,0 +1,122 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Models for MAE."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.mae.modeling import utils
+from official.vision.modeling.backbones import vit
+
+
+class MaskedAE(tf_keras.Model):
+ """MAE model."""
+
+ def __init__(self,
+ encoder,
+ name=None,
+ **kwargs):
+ super(MaskedAE, self).__init__(name=name, **kwargs)
+ self.encoder = encoder
+ self.pixels_per_patch = self.encoder.patch_h * self.encoder.patch_w * 3
+
+ def build(self, input_shape):
+ self.decoder = vit.Encoder(
+ num_layers=8,
+ mlp_dim=2048,
+ num_heads=16,
+ dropout_rate=0.0,
+ attention_dropout_rate=0.0,
+ add_pos_embed=False
+ )
+ self.mask = self.add_weight(
+ 'mask', (1, 1, 512),
+ initializer=tf_keras.initializers.RandomNormal(stddev=0.02))
+ self.to_pixels = tf_keras.layers.Dense(self.pixels_per_patch)
+ self.linear = tf_keras.layers.Dense(512)
+ super().build(input_shape)
+
+ def add_position_embed(self, patch_embeds, num_rows, num_cols):
+ # patch_embeds is 1d (N, 1+H*W, D) with cls token.
+ shape = tf.shape(patch_embeds)
+ position_embedding = utils.position_embedding_sine(
+ tf.ones((shape[0], num_rows, num_cols), dtype=patch_embeds.dtype),
+ 512, normalize=False)
+ position_embedding = tf.reshape(
+ position_embedding, (shape[0], num_rows * num_cols, -1))
+ return patch_embeds + tf.concat(
+ [tf.zeros((shape[0], 1, shape[2]), dtype=patch_embeds.dtype),
+ position_embedding
+ ], axis=1)
+
+ def call(self, inputs, training=None, masking=None):
+ patches = inputs['patches']
+ masked_indices = tf.cast(inputs['masked_indices'], tf.int32)
+ unmasked_indices = tf.cast(inputs['unmasked_indices'], tf.int32)
+ batch_size = tf.shape(patches)[0]
+ num_h_patches = tf.shape(patches)[1]
+ num_w_patches = tf.shape(patches)[2]
+ num_patches = num_h_patches * num_w_patches
+ num_masks = tf.shape(masked_indices)[1]
+ patch_embeds = self.encoder.to_embed(patches)
+ patch_embeds = self.encoder.add_position_embed(patch_embeds)
+ patch_embeds = tf.reshape(
+ patch_embeds,
+ (batch_size, num_patches, -1))
+ patch_embeds = self.encoder.insert_cls(patch_embeds)
+
+ unmasked_indices = tf.concat(
+ [tf.zeros((batch_size, 1), unmasked_indices.dtype),
+ unmasked_indices + 1],
+ axis=1)
+ masked_indices = masked_indices + 1
+ unmasked_patch_embeds = tf.gather(
+ patch_embeds, unmasked_indices, batch_dims=1)
+ encoded = self.encoder({'embeddings': unmasked_patch_embeds})
+ encoded = self.linear(encoded)
+
+ zeros = tf.zeros((batch_size, num_patches + 1, 512))
+
+ unmasked_embed = tf.tensor_scatter_nd_add(
+ zeros,
+ tf.stack([
+ tf.tile(
+ tf.expand_dims(tf.range(batch_size), axis=1),
+ [1, num_patches + 1 - num_masks]), unmasked_indices
+ ],
+ axis=-1),
+ encoded)
+ mask_embeds = tf.tile(self.mask, [batch_size, num_masks, 1])
+ full_embed = tf.tensor_scatter_nd_add(
+ unmasked_embed,
+ tf.stack([
+ tf.tile(
+ tf.expand_dims(tf.range(batch_size), axis=1),
+ [1, num_masks]), masked_indices
+ ],
+ axis=-1),
+ mask_embeds)
+ full_embed = self.add_position_embed(
+ full_embed, num_h_patches, num_w_patches)
+
+ decoded = self.decoder(full_embed)
+ pred_pixel_values = self.to_pixels(
+ tf.gather(decoded, masked_indices, batch_dims=1))
+ return pred_pixel_values
+
+ @property
+ def checkpoint_items(self):
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(encoder=self.encoder)
+ return items
diff --git a/official/projects/mae/modeling/utils.py b/official/projects/mae/modeling/utils.py
new file mode 100644
index 00000000000..24276ca0648
--- /dev/null
+++ b/official/projects/mae/modeling/utils.py
@@ -0,0 +1,86 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utils for MAE."""
+
+import math
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
+
+
+# TODO(frederickliu): Move this to vision ops and add tests.
+def position_embedding_sine(attention_mask,
+ num_pos_features=256,
+ temperature=10000.,
+ normalize=True,
+ scale=2 * math.pi):
+ """Sine-based positional embeddings for 2D images.
+
+ Args:
+ attention_mask: a `bool` Tensor specifying the size of the input image to
+ the Transformer and which elements are padded, of size [batch_size,
+ height, width]
+ num_pos_features: a `int` specifying the number of positional features,
+ should be equal to the hidden size of the Transformer network
+ temperature: a `float` specifying the temperature of the positional
+ embedding. Any type that is converted to a `float` can also be accepted.
+ normalize: a `bool` determining whether the positional embeddings should be
+ normalized between [0, scale] before application of the sine and cos
+ functions.
+ scale: a `float` if normalize is True specifying the scale embeddings before
+ application of the embedding function.
+
+ Returns:
+ embeddings: a `float` tensor of the same shape as input_tensor specifying
+ the positional embeddings based on sine features.
+ """
+ if num_pos_features % 2 != 0:
+ raise ValueError(
+ "Number of embedding features (num_pos_features) must be even when "
+ "column and row embeddings are concatenated.")
+ num_pos_features = num_pos_features // 2
+
+ # Produce row and column embeddings based on total size of the image
+ # [batch_size, height, width]
+ attention_mask = tf.cast(attention_mask, tf.float32)
+ row_embedding = tf.cumsum(attention_mask, 1)
+ col_embedding = tf.cumsum(attention_mask, 2)
+
+ if normalize:
+ eps = 1e-6
+ row_embedding = row_embedding / (row_embedding[:, -1:, :] + eps) * scale
+ col_embedding = col_embedding / (col_embedding[:, :, -1:] + eps) * scale
+
+ dim_t = tf.range(num_pos_features, dtype=row_embedding.dtype)
+ dim_t = tf.pow(temperature, 2 * (dim_t // 2) / num_pos_features)
+
+ # Creates positional embeddings for each row and column position
+ # [batch_size, height, width, num_pos_features]
+ pos_row = tf.expand_dims(row_embedding, -1) / dim_t
+ pos_col = tf.expand_dims(col_embedding, -1) / dim_t
+ pos_row = tf.stack(
+ [tf.sin(pos_row[:, :, :, 0::2]),
+ tf.cos(pos_row[:, :, :, 1::2])], axis=4)
+ pos_col = tf.stack(
+ [tf.sin(pos_col[:, :, :, 0::2]),
+ tf.cos(pos_col[:, :, :, 1::2])], axis=4)
+
+ # final_shape = pos_row.shape.as_list()[:3] + [-1]
+ final_shape = tf_utils.get_shape_list(pos_row)[:3] + [-1]
+ pos_row = tf.reshape(pos_row, final_shape)
+ pos_col = tf.reshape(pos_col, final_shape)
+ output = tf.concat([pos_row, pos_col], -1)
+
+ embeddings = tf.cast(output, tf.float32)
+ return embeddings
diff --git a/official/projects/mae/modeling/vit.py b/official/projects/mae/modeling/vit.py
new file mode 100644
index 00000000000..8e91d13f858
--- /dev/null
+++ b/official/projects/mae/modeling/vit.py
@@ -0,0 +1,125 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Models for ViT."""
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.projects.mae.modeling import utils
+from official.vision.modeling.backbones import vit
+
+
+def to_patch(images, patch_height, patch_width):
+ """Image (NHWC) to patches (N(H' W')(patch_height patch_width c))."""
+ batch_size, h, w, c = tf_utils.get_shape_list(images)
+ num_h = h // patch_height
+ num_w = w // patch_width
+ x = tf.reshape(images,
+ (batch_size, num_h, patch_height, num_w, patch_width, c))
+ x = tf.einsum('nhpwqc->nhwpqc', x)
+ x = tf.reshape(x, (batch_size, num_h, num_w, patch_height * patch_width * c))
+ return x
+
+
+class ViTClassifier(tf_keras.Model):
+ """ViT classifier for finetune."""
+
+ def __init__(self, encoder, num_classes, **kwargs):
+ super().__init__(**kwargs)
+ self.encoder = encoder
+ self.linear = tf_keras.layers.Dense(
+ num_classes,
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=2e-5))
+
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ encoded = self.encoder({'images': inputs})
+ return self.linear(encoded[:, 0])
+
+
+class ViTLinearClassifier(tf_keras.Model):
+ """ViT classifier for linear probing."""
+
+ def __init__(self, encoder, num_classes, use_sync_bn=True, **kwargs):
+ super().__init__(**kwargs)
+ self.encoder = encoder
+ self.linear = tf_keras.layers.Dense(
+ num_classes,
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01))
+ if use_sync_bn:
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ self._norm = tf_keras.layers.BatchNormalization
+ self.batch_norm = self._norm(
+ axis=-1, epsilon=1e-6, center=False, scale=False, momentum=0.9)
+
+ def call(self, inputs, training=False): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ encoded = self.encoder({'images': inputs})
+ features = self.batch_norm(encoded[:, 0], training=training)
+ return self.linear(features)
+
+
+class VisionTransformer(tf_keras.Model):
+ """ViT backbone."""
+
+ def __init__(self,
+ patch_h,
+ patch_w,
+ init_stochastic_depth_rate=0.0,
+ **kwargs):
+ super().__init__(**kwargs)
+ self.patch_h = patch_h
+ self.patch_w = patch_w
+ self.init_stochastic_depth_rate = init_stochastic_depth_rate
+
+ def build(self, input_shape):
+ self.patch_to_embed = tf_keras.layers.Dense(1024)
+ # ViT-L
+ self.encoder = vit.Encoder(
+ num_layers=24,
+ mlp_dim=4096,
+ num_heads=16,
+ dropout_rate=0.0,
+ attention_dropout_rate=0.0,
+ init_stochastic_depth_rate=self.init_stochastic_depth_rate,
+ add_pos_embed=False,
+ )
+ self.token_cls = vit.TokenLayer()
+ super().build(input_shape)
+
+ def to_embed(self, patches):
+ return self.patch_to_embed(patches)
+
+ def insert_cls(self, patch_embeds):
+ return self.token_cls(patch_embeds)
+
+ def add_position_embed(self, patch_embeds):
+ return patch_embeds + utils.position_embedding_sine(
+ tf.ones_like(patch_embeds[..., 0]), 1024, normalize=False)
+
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ if isinstance(inputs, dict):
+ images = inputs.get('images', None)
+ patch_embeds = inputs.get('embeddings', None)
+ else:
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+ if images is not None:
+ patches = to_patch(images, self.patch_h, self.patch_w)
+ patch_embeds = self.to_embed(patches)
+ patch_shape = tf.shape(patch_embeds)
+ patch_embeds = self.add_position_embed(patch_embeds)
+ patch_embeds = tf.reshape(patch_embeds,
+ (patch_shape[0], -1, patch_shape[-1]))
+ patch_embeds = self.insert_cls(patch_embeds)
+ return self.encoder(patch_embeds)
diff --git a/official/projects/mae/optimization.py b/official/projects/mae/optimization.py
new file mode 100644
index 00000000000..8e7f6f3e574
--- /dev/null
+++ b/official/projects/mae/optimization.py
@@ -0,0 +1,225 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Customized optimizer to match paper results."""
+
+import dataclasses
+from typing import List, Optional
+
+from absl import logging
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import optimization
+from official.nlp import optimization as nlp_optimization
+
+
+@dataclasses.dataclass
+class ViTAdamWConfig(optimization.AdamWeightDecayConfig):
+ layer_decay: Optional[float] = 1.0
+ vars_substr: Optional[List[str]] = None
+ layers_idx: Optional[List[int]] = None
+
+
+@dataclasses.dataclass
+class OptimizerConfig(optimization.OptimizerConfig):
+ vit_adamw: ViTAdamWConfig = dataclasses.field(default_factory=ViTAdamWConfig)
+
+
+@dataclasses.dataclass
+class OptimizationConfig(optimization.OptimizationConfig):
+ """Configuration for optimizer and learning rate schedule.
+
+ Attributes:
+ optimizer: optimizer oneof config.
+ ema: optional exponential moving average optimizer config, if specified, ema
+ optimizer will be used.
+ learning_rate: learning rate oneof config.
+ warmup: warmup oneof config.
+ """
+ optimizer: OptimizerConfig = dataclasses.field(
+ default_factory=OptimizerConfig
+ )
+
+
+# TODO(frederickliu): figure out how to make this configuable.
+# TODO(frederickliu): Study if this is needed.
+class _ViTAdamW(nlp_optimization.AdamWeightDecay):
+ """Custom AdamW to support different lr scaling for backbone.
+
+ The code is copied from AdamWeightDecay and Adam with learning scaling.
+ """
+
+ def __init__(self,
+ learning_rate=0.001,
+ beta_1=0.9,
+ beta_2=0.999,
+ epsilon=1e-7,
+ amsgrad=False,
+ weight_decay_rate=0.0,
+ include_in_weight_decay=None,
+ exclude_from_weight_decay=None,
+ gradient_clip_norm=1.0,
+ layer_decay=1.0,
+ vars_substr=None,
+ layers_idx=None,
+ name='ViTAdamWeightDecay',
+ **kwargs):
+ super(_ViTAdamW,
+ self).__init__(learning_rate, beta_1, beta_2, epsilon, amsgrad,
+ weight_decay_rate, include_in_weight_decay,
+ exclude_from_weight_decay, gradient_clip_norm, name,
+ **kwargs)
+ self._layer_decay = layer_decay
+ self._vars_substr = vars_substr
+ self._layers_idx = layers_idx
+ self._max_idx = max(layers_idx) + 1 if layers_idx is not None else 1
+
+ def _resource_apply_dense(self, grad, var, apply_state=None):
+ lr_t, kwargs = self._get_lr(var.device, var.dtype.base_dtype, apply_state)
+ apply_state = kwargs['apply_state']
+ if (
+ self._layer_decay != 1.0
+ and self._vars_substr is not None
+ and self._layers_idx is not None
+ ):
+ is_decayed = False
+ for var_substr, idx in zip(self._vars_substr, self._layers_idx):
+ if var_substr in var.name:
+ decay_factor = self._layer_decay ** (self._max_idx - idx)
+ lr_t = lr_t * decay_factor
+ is_decayed = True
+ logging.debug(
+ 'Applying layer-wise lr decay: %s: %f', var.name, decay_factor)
+ break
+ if not is_decayed:
+ logging.debug('Ignore layer-wise lr decay: %s', var.name)
+ decay = self._decay_weights_op(var, lr_t, apply_state)
+ with tf.control_dependencies([decay]):
+ var_device, var_dtype = var.device, var.dtype.base_dtype
+ coefficients = ((apply_state or {}).get((var_device, var_dtype))
+ or self._fallback_apply_state(var_device, var_dtype))
+
+ m = self.get_slot(var, 'm')
+ v = self.get_slot(var, 'v')
+ lr = coefficients['lr_t']
+ if (
+ self._layer_decay != 1.0
+ and self._vars_substr is not None
+ and self._layers_idx is not None
+ ):
+ for var_substr, idx in zip(self._vars_substr, self._layers_idx):
+ if var_substr in var.name:
+ lr = lr * (self._layer_decay ** (self._max_idx - idx))
+ break
+
+ if not self.amsgrad:
+ return tf.raw_ops.ResourceApplyAdam(
+ var=var.handle,
+ m=m.handle,
+ v=v.handle,
+ beta1_power=coefficients['beta_1_power'],
+ beta2_power=coefficients['beta_2_power'],
+ lr=lr,
+ beta1=coefficients['beta_1_t'],
+ beta2=coefficients['beta_2_t'],
+ epsilon=coefficients['epsilon'],
+ grad=grad,
+ use_locking=self._use_locking)
+ else:
+ vhat = self.get_slot(var, 'vhat')
+ return tf.raw_ops.ResourceApplyAdamWithAmsgrad(
+ var=var.handle,
+ m=m.handle,
+ v=v.handle,
+ vhat=vhat.handle,
+ beta1_power=coefficients['beta_1_power'],
+ beta2_power=coefficients['beta_2_power'],
+ lr=lr,
+ beta1=coefficients['beta_1_t'],
+ beta2=coefficients['beta_2_t'],
+ epsilon=coefficients['epsilon'],
+ grad=grad,
+ use_locking=self._use_locking)
+
+ def _resource_apply_sparse(self, grad, var, indices, apply_state=None):
+ lr_t, kwargs = self._get_lr(var.device, var.dtype.base_dtype, apply_state)
+ apply_state = kwargs['apply_state']
+ if (
+ self._layer_decay != 1.0
+ and self._vars_substr is not None
+ and self._layers_idx is not None
+ ):
+ is_decayed = False
+ for var_substr, idx in zip(self._vars_substr, self._layers_idx):
+ if var_substr in var.name:
+ decay_factor = self._layer_decay ** (self._max_idx - idx)
+ lr_t = lr_t * decay_factor
+ is_decayed = True
+ logging.debug(
+ 'Applying layer-wise lr decay: %s: %f', var.name, decay_factor)
+ break
+ if not is_decayed:
+ logging.debug('Ignore layer-wise lr decay: %s', var.name)
+ decay = self._decay_weights_op(var, lr_t, apply_state)
+ with tf.control_dependencies([decay]):
+ var_device, var_dtype = var.device, var.dtype.base_dtype
+ coefficients = ((apply_state or {}).get((var_device, var_dtype))
+ or self._fallback_apply_state(var_device, var_dtype))
+
+ # m_t = beta1 * m + (1 - beta1) * g_t
+ m = self.get_slot(var, 'm')
+ m_scaled_g_values = grad * coefficients['one_minus_beta_1_t']
+ m_t = tf.compat.v1.assign(m, m * coefficients['beta_1_t'],
+ use_locking=self._use_locking)
+ with tf.control_dependencies([m_t]):
+ m_t = self._resource_scatter_add(m, indices, m_scaled_g_values)
+
+ # v_t = beta2 * v + (1 - beta2) * (g_t * g_t)
+ v = self.get_slot(var, 'v')
+ v_scaled_g_values = (grad * grad) * coefficients['one_minus_beta_2_t']
+ v_t = tf.compat.v1.assign(v, v * coefficients['beta_2_t'],
+ use_locking=self._use_locking)
+ with tf.control_dependencies([v_t]):
+ v_t = self._resource_scatter_add(v, indices, v_scaled_g_values)
+ lr = coefficients['lr_t']
+ if (
+ self._layer_decay != 1.0
+ and self._vars_substr is not None
+ and self._layers_idx is not None
+ ):
+ for var_substr, idx in zip(self._vars_substr, self._layers_idx):
+ if var_substr in var.name:
+ lr = lr * (self._layer_decay ** (self._max_idx - idx))
+ break
+ if not self.amsgrad:
+ v_sqrt = tf.sqrt(v_t)
+ var_update = tf.compat.v1.assign_sub(
+ var, lr * m_t / (v_sqrt + coefficients['epsilon']),
+ use_locking=self._use_locking)
+ return tf.group(*[var_update, m_t, v_t])
+ else:
+ v_hat = self.get_slot(var, 'vhat')
+ v_hat_t = tf.maximum(v_hat, v_t)
+ with tf.control_dependencies([v_hat_t]):
+ v_hat_t = tf.compat.v1.assign(
+ v_hat, v_hat_t, use_locking=self._use_locking)
+ v_hat_sqrt = tf.sqrt(v_hat_t)
+ var_update = tf.compat.v1.assign_sub(
+ var,
+ lr* m_t / (v_hat_sqrt + coefficients['epsilon']),
+ use_locking=self._use_locking)
+ return tf.group(*[var_update, m_t, v_t, v_hat_t])
+
+optimization.register_optimizer_cls('vit_adamw', _ViTAdamW)
diff --git a/official/projects/mae/tasks/image_classification.py b/official/projects/mae/tasks/image_classification.py
new file mode 100644
index 00000000000..e0cbf3236bc
--- /dev/null
+++ b/official/projects/mae/tasks/image_classification.py
@@ -0,0 +1,128 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Image classification task with ViT."""
+
+import dataclasses
+from typing import Optional, Tuple
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import config_definitions as cfg
+from official.core import input_reader
+from official.core import task_factory
+from official.projects.mae.modeling import vit
+from official.vision.dataloaders import classification_input
+from official.vision.dataloaders import tfds_factory
+from official.vision.ops import augment
+
+
+@dataclasses.dataclass
+class ViTConfig(cfg.TaskConfig):
+ """The translation task config."""
+
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
+ patch_h: int = 14
+ patch_w: int = 14
+ num_classes: int = 1000
+ input_size: Tuple[int, int] = (224, 224)
+ init_stochastic_depth_rate: float = 0.2
+
+
+@task_factory.register_task_cls(ViTConfig)
+class ViTClassificationTask(base_task.Task):
+ """Image classificaiton with ViT and load checkpoint if exists."""
+
+ def build_model(self) -> tf_keras.Model:
+ encoder = vit.VisionTransformer(
+ self.task_config.patch_h,
+ self.task_config.patch_w,
+ self.task_config.init_stochastic_depth_rate)
+ model = vit.ViTClassifier(encoder, self.task_config.num_classes)
+ model(tf.ones((1, 224, 224, 3)))
+ return model
+
+ def build_inputs(self,
+ params,
+ input_context: Optional[tf.distribute.InputContext] = None):
+ num_classes = self.task_config.num_classes
+ input_size = self.task_config.input_size
+ image_field_key = self.task_config.train_data.image_field_key
+ label_field_key = self.task_config.train_data.label_field_key
+
+ decoder = tfds_factory.get_classification_decoder(params.tfds_name)
+ parser = classification_input.Parser(
+ output_size=input_size[:2],
+ num_classes=num_classes,
+ image_field_key=image_field_key,
+ label_field_key=label_field_key,
+ decode_jpeg_only=params.decode_jpeg_only,
+ aug_rand_hflip=params.aug_rand_hflip,
+ aug_type=params.aug_type,
+ color_jitter=params.color_jitter,
+ random_erasing=params.random_erasing,
+ dtype=params.dtype)
+
+ if params.is_training:
+ postprocess_fn = augment.MixupAndCutmix(
+ mixup_alpha=0.8,
+ cutmix_alpha=1.0,
+ prob=1.0 if params.is_training else 0.0,
+ label_smoothing=0.1,
+ num_classes=num_classes)
+ else:
+ postprocess_fn = lambda images, labels: ( # pylint:disable=g-long-lambda
+ images, tf.one_hot(labels, num_classes))
+
+ reader = input_reader.InputReader(
+ params=params,
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training),
+ postprocess_fn=postprocess_fn)
+
+ dataset = reader.read(input_context=input_context)
+ return dataset
+
+ def initialize(self, model: tf_keras.Model):
+ """Load encoder if checkpoint exists.
+
+ Args:
+ model: The keras.Model built or used by this task.
+ """
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+ if not ckpt_dir_or_file:
+ return
+
+ checkpoint_items = dict(encoder=model.encoder)
+ ckpt = tf.train.Checkpoint(**checkpoint_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ def build_metrics(self, training=None):
+ del training
+ metrics = [
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
+ ]
+ return metrics
+
+ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
+ return tf_keras.losses.categorical_crossentropy(
+ labels,
+ model_outputs,
+ from_logits=True)
diff --git a/official/projects/mae/tasks/image_classification_test.py b/official/projects/mae/tasks/image_classification_test.py
new file mode 100644
index 00000000000..e66d353acfc
--- /dev/null
+++ b/official/projects/mae/tasks/image_classification_test.py
@@ -0,0 +1,95 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for image_classification."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+import tensorflow_datasets as tfds
+
+from official.modeling import optimization
+from official.projects.mae.tasks import image_classification as vit_cls
+from official.vision.configs import image_classification
+
+
+_NUM_EXAMPLES = 10
+
+
+def _gen_fn():
+ h = np.random.randint(0, 300)
+ w = np.random.randint(0, 300)
+ return {
+ 'image': np.ones(shape=(h, w, 3), dtype=np.uint8),
+ 'label': np.random.randint(0, 100),
+ 'file_name': 'test',
+ }
+
+
+def _as_dataset(self, *args, **kwargs):
+ del args
+ del kwargs
+ return tf.data.Dataset.from_generator(
+ lambda: (_gen_fn() for i in range(_NUM_EXAMPLES)),
+ output_types=self.info.features.dtype,
+ output_shapes=self.info.features.shape,
+ )
+
+
+class ImageClassificationTest(tf.test.TestCase):
+
+ def test_train_step(self):
+ config = vit_cls.ViTConfig(
+ num_classes=1000,
+ train_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='validation',
+ is_training=True,
+ global_batch_size=2,
+ ),
+ )
+ with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
+ task = vit_cls.ViTClassificationTask(config)
+ model = task.build_model()
+ dataset = task.build_inputs(config.train_data)
+ iterator = iter(dataset)
+ opt_cfg = optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 1.5 * 1e-4,
+ 'decay_steps': 5
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 1,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })
+ optimizer = vit_cls.ViTClassificationTask.create_optimizer(opt_cfg)
+ task.train_step(next(iterator), model, optimizer)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mae/tasks/linear_probe.py b/official/projects/mae/tasks/linear_probe.py
new file mode 100644
index 00000000000..e1de9e93dc2
--- /dev/null
+++ b/official/projects/mae/tasks/linear_probe.py
@@ -0,0 +1,114 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Image classification task with ViT and linear probe."""
+import dataclasses
+from typing import Optional
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import input_reader
+from official.core import task_factory
+from official.projects.mae.modeling import vit
+from official.projects.mae.tasks import image_classification
+from official.vision.dataloaders import classification_input
+from official.vision.dataloaders import tfds_factory
+
+
+@dataclasses.dataclass
+class ViTLinearProbeConfig(image_classification.ViTConfig):
+ """The LinearProbe task config."""
+
+
+@task_factory.register_task_cls(ViTLinearProbeConfig)
+class ViTLinearProbeTask(base_task.Task):
+ """Image classificaiton with ViT and load checkpoint if exists."""
+
+ def build_model(self) -> tf_keras.Model:
+ encoder = vit.VisionTransformer(
+ self.task_config.patch_h,
+ self.task_config.patch_w,
+ self.task_config.init_stochastic_depth_rate,
+ )
+ # Freeze backbone.
+ encoder.trainable = False
+ model = vit.ViTLinearClassifier(encoder, self.task_config.num_classes)
+ model(tf.ones((1, 224, 224, 3)))
+ return model
+
+ def build_inputs(
+ self, params, input_context: Optional[tf.distribute.InputContext] = None
+ ):
+ num_classes = self.task_config.num_classes
+ input_size = self.task_config.input_size
+ image_field_key = self.task_config.train_data.image_field_key
+ label_field_key = self.task_config.train_data.label_field_key
+
+ decoder = tfds_factory.get_classification_decoder(params.tfds_name)
+ parser = classification_input.Parser(
+ output_size=input_size[:2],
+ num_classes=num_classes,
+ image_field_key=image_field_key,
+ label_field_key=label_field_key,
+ decode_jpeg_only=params.decode_jpeg_only,
+ aug_rand_hflip=params.aug_rand_hflip,
+ aug_type=params.aug_type,
+ color_jitter=params.color_jitter,
+ random_erasing=params.random_erasing,
+ dtype=params.dtype,
+ )
+
+ postprocess_fn = lambda images, labels: ( # pylint:disable=g-long-lambda
+ images,
+ tf.one_hot(labels, num_classes),
+ )
+
+ reader = input_reader.InputReader(
+ params=params,
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training),
+ postprocess_fn=postprocess_fn,
+ )
+
+ dataset = reader.read(input_context=input_context)
+ return dataset
+
+ def initialize(self, model: tf_keras.Model):
+ """Load encoder if checkpoint exists.
+
+ Args:
+ model: The keras.Model built or used by this task.
+ """
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+ if not ckpt_dir_or_file:
+ return
+
+ checkpoint_items = dict(encoder=model.encoder)
+ ckpt = tf.train.Checkpoint(**checkpoint_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ def build_metrics(self, training=None):
+ del training
+ metrics = [
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
+ ]
+ return metrics
+
+ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
+ return tf_keras.losses.categorical_crossentropy(
+ labels, model_outputs, from_logits=True
+ )
diff --git a/official/projects/mae/tasks/linear_probe_test.py b/official/projects/mae/tasks/linear_probe_test.py
new file mode 100644
index 00000000000..47e443ebeac
--- /dev/null
+++ b/official/projects/mae/tasks/linear_probe_test.py
@@ -0,0 +1,94 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for image_classification."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+import tensorflow_datasets as tfds
+
+from official.modeling import optimization
+from official.projects.mae.tasks import linear_probe
+from official.vision.configs import image_classification
+
+
+_NUM_EXAMPLES = 10
+
+
+def _gen_fn():
+ h = np.random.randint(0, 300)
+ w = np.random.randint(0, 300)
+ return {
+ 'image': np.ones(shape=(h, w, 3), dtype=np.uint8),
+ 'label': np.random.randint(0, 100),
+ 'file_name': 'test',
+ }
+
+
+def _as_dataset(self, *args, **kwargs):
+ del args
+ del kwargs
+ return tf.data.Dataset.from_generator(
+ lambda: (_gen_fn() for i in range(_NUM_EXAMPLES)),
+ output_types=self.info.features.dtype,
+ output_shapes=self.info.features.shape,
+ )
+
+
+class ImageClassificationTest(tf.test.TestCase):
+
+ def test_train_step(self):
+ config = linear_probe.ViTLinearProbeConfig(
+ num_classes=1000,
+ train_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='validation',
+ is_training=True,
+ global_batch_size=2,
+ ))
+ with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
+ task = linear_probe.ViTLinearProbeTask(config)
+ model = task.build_model()
+ dataset = task.build_inputs(config.train_data)
+ iterator = iter(dataset)
+ opt_cfg = optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 1.5 * 1e-4,
+ 'decay_steps': 5
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 1,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })
+ optimizer = linear_probe.ViTLinearProbeTask.create_optimizer(opt_cfg)
+ task.train_step(next(iterator), model, optimizer)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mae/tasks/masked_ae.py b/official/projects/mae/tasks/masked_ae.py
new file mode 100644
index 00000000000..9700a61896a
--- /dev/null
+++ b/official/projects/mae/tasks/masked_ae.py
@@ -0,0 +1,106 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Task for masked autoencoder pretraining."""
+
+from typing import Optional
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import input_reader
+from official.core import task_factory
+from official.modeling import tf_utils
+from official.projects.mae.configs import mae as mae_cfg
+from official.projects.mae.modeling import masked_ae
+from official.projects.mae.modeling import vit
+from official.vision.dataloaders import classification_input
+from official.vision.dataloaders import tfds_factory
+
+
+@task_factory.register_task_cls(mae_cfg.MAEConfig)
+class MaskedAETask(base_task.Task):
+ """Task for masked autoencoder training."""
+
+ def build_model(self) -> tf_keras.Model:
+ encoder = vit.VisionTransformer(
+ self.task_config.patch_h,
+ self.task_config.patch_w,
+ 0.0)
+ # trigger build to be called.
+ input_size = self.task_config.input_size
+ encoder({'images': tf.ones((1, input_size[0], input_size[1], 3))})
+ model = masked_ae.MaskedAE(encoder)
+ return model
+
+ def build_inputs(self,
+ params,
+ input_context: Optional[tf.distribute.InputContext] = None):
+ num_classes = self.task_config.num_classes
+ input_size = self.task_config.input_size
+ image_field_key = self.task_config.train_data.image_field_key
+ label_field_key = self.task_config.train_data.label_field_key
+
+ decoder = tfds_factory.get_classification_decoder(params.tfds_name)
+ parser = classification_input.Parser(
+ output_size=input_size[:2],
+ num_classes=num_classes,
+ image_field_key=image_field_key,
+ label_field_key=label_field_key,
+ decode_jpeg_only=params.decode_jpeg_only,
+ aug_rand_hflip=params.aug_rand_hflip,
+ aug_type=params.aug_type,
+ color_jitter=params.color_jitter,
+ random_erasing=params.random_erasing,
+ dtype=params.dtype,
+ crop_area_range=params.crop_area_range)
+
+ def patch_and_mask(images, labels):
+ del labels
+ patches = vit.to_patch(
+ images, self.task_config.patch_h, self.task_config.patch_w)
+ batch_size, num_h_patches, num_w_patches = tf_utils.get_shape_list(
+ patches)[:3]
+ num_patches = num_h_patches * num_w_patches
+ num_masked = tf.cast(
+ self.task_config.masking_ratio * num_patches, dtype=tf.int32)
+ r = tf.random.uniform((batch_size, num_patches))
+ rand_indices = tf.argsort(r)
+
+ masked_indices = rand_indices[:, :num_masked]
+ unmasked_indices = rand_indices[:, num_masked:]
+ patches_1d = tf.reshape(patches, (batch_size, num_patches, -1))
+ masked_patches = tf.gather(patches_1d, masked_indices, batch_dims=1)
+
+ if self.task_config.norm_target:
+ mean = tf.reduce_mean(masked_patches, axis=-1, keepdims=True)
+ var = tf.math.reduce_variance(masked_patches, axis=-1, keepdims=True)
+ std = (var + 1.e-6)**.5
+ masked_patches = (masked_patches - mean) / std
+
+ return {'patches': patches,
+ 'masked_indices': masked_indices,
+ 'unmasked_indices': unmasked_indices}, masked_patches
+
+ reader = input_reader.InputReader(
+ params=params,
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training),
+ postprocess_fn=patch_and_mask)
+
+ dataset = reader.read(input_context=input_context)
+ return dataset
+
+ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
+ return tf_keras.metrics.mean_squared_error(
+ labels, model_outputs)
diff --git a/official/projects/mae/tasks/masked_ae_test.py b/official/projects/mae/tasks/masked_ae_test.py
new file mode 100644
index 00000000000..88080d7d116
--- /dev/null
+++ b/official/projects/mae/tasks/masked_ae_test.py
@@ -0,0 +1,95 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for masked_ae."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+import tensorflow_datasets as tfds
+
+from official.modeling import optimization
+from official.projects.mae.configs import mae as mae_cfg
+from official.projects.mae.tasks import masked_ae
+from official.vision.configs import image_classification
+
+
+_NUM_EXAMPLES = 10
+
+
+def _gen_fn():
+ np.random.seed(0) # Some seed may cause jpeg decoding to fail.
+ h = np.random.randint(0, 300)
+ w = np.random.randint(0, 300)
+ return {
+ 'image': np.ones(shape=(h, w, 3), dtype=np.uint8),
+ 'label': np.random.randint(0, 100),
+ 'file_name': 'test',
+ }
+
+
+def _as_dataset(self, *args, **kwargs):
+ del args
+ del kwargs
+ return tf.data.Dataset.from_generator(
+ lambda: (_gen_fn() for i in range(_NUM_EXAMPLES)),
+ output_types=self.info.features.dtype,
+ output_shapes=self.info.features.shape,
+ )
+
+
+class MAETest(tf.test.TestCase):
+
+ def test_train_step(self):
+ config = mae_cfg.MAEConfig(
+ train_data=image_classification.DataConfig(
+ tfds_name='imagenet2012',
+ tfds_split='validation',
+ is_training=True,
+ global_batch_size=2,
+ ))
+ with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
+ task = masked_ae.MaskedAETask(config)
+ model = task.build_model()
+ dataset = task.build_inputs(config.train_data)
+ iterator = iter(dataset)
+ opt_cfg = optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 1.5 * 1e-4,
+ 'decay_steps': 5
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 1,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })
+ optimizer = masked_ae.MaskedAETask.create_optimizer(opt_cfg)
+ task.train_step(next(iterator), model, optimizer)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mae/train.py b/official/projects/mae/train.py
new file mode 100644
index 00000000000..a1d105f8764
--- /dev/null
+++ b/official/projects/mae/train.py
@@ -0,0 +1,32 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""TensorFlow Model Garden Vision training driver, register MAE configs."""
+
+from absl import app
+
+from official.common import flags as tfm_flags
+# pylint: disable=unused-import
+from official.projects.mae.configs import linear_probe
+from official.projects.mae.configs import mae
+from official.projects.mae.configs import vit
+from official.projects.mae.tasks import image_classification
+from official.projects.mae.tasks import linear_probe as linear_probe_task
+from official.projects.mae.tasks import masked_ae
+# pylint: enable=unused-import
+from official.vision import train
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(train.main)
diff --git a/official/projects/maskconver/README.md b/official/projects/maskconver/README.md
new file mode 100644
index 00000000000..47d93685fb5
--- /dev/null
+++ b/official/projects/maskconver/README.md
@@ -0,0 +1,82 @@
+# MaskConver: Revisiting Pure Convolution Model for Panoptic Segmentation (WACV 2024)
+
+[](https://arxiv.org/abs/2312.06052)
+
+
+[MaskConver](https://arxiv.org/abs/2312.06052) is a pure convolutional panoptic
+architecture. MaskConver proposes to fully unify things and stuff representation
+by predicting their centers. To that extent, it creates a lightweight class
+embedding module that can break the ties when multiple centers co-exist in the
+same location. Furthermore, our study shows that the decoder design is critical
+in ensuring that the model has sufficient context for accurate detection and
+segmentation. We introduce a powerful ConvNeXt-UNet decoder that closes the
+performance gap between convolution- and transformer based models. With ResNet50
+backbone, our MaskConver achieves 53.6% PQ on the COCO panoptic val set,
+outperforming the modern convolution-based model, Panoptic FCN, by 9.3% as well
+as transformer-based models such as Mask2Former (+1.7% PQ) and kMaX-DeepLab
+(+0.6% PQ). Additionally, MaskConver with a MobileNet backbone reaches 37.2% PQ,
+improving over Panoptic-DeepLab by +6.4% under the same FLOPs/latency
+constraints. A further optimized version of MaskConver achieves 29.7% PQ, while
+running in real-time on mobile devices.
+
+
+MaskConver meta-architecture:
+
+
+
+
+
+The meta architecture of MaskConver contains four components: backbone (gray),
+pixel decoder (pink), prediction heads (light blue), and mask embedding
+generator (green). The backbone is any commonly deployed neural network, e.g.,
+ResNet50. We propose a novel ConvNeXt-UNet for the pixel decoder, which
+effectively captures long-range context and high-level semantics by stacking
+many ConvNeXt blocks at the highest level of backbone. We propose three
+prediction heads: Center Heatmap Head (for predicting center point heatmaps),
+Center Embedding Head (for predicting the embeddings for center points), and
+Mask Feature Head (for generating mask features). The Mask Embedding Generator
+first produces the class embeddings via a lookup table (Class Embedding Lookup
+Table module) by taking the predicted semantic classes from the top-K center
+points. The output mask embeddings are obtained by modulating the class
+embeddings with the center embeddings (via addition and MLP) to mitigate the
+center point collision between instances of different classes. In the end, the
+mask features are multiplied with the mask embeddings to generate the final
+binary masks. Unlike transformer-based methods, MaskConver only exploits
+convolutions without any self- or cross-attentions.
+
+
+
+
+
+
+
+
+
+
+## Performance Reference
+
+
+| Backbone | Image Size | Params | FLOPS | PQ | Latency | Link |
+| ------------ | ---------- | ------ | ----- | ---- | -------- | ---- |
+| MobileNet-MH | 256x256 | 3.4M | 1.5B | 29.7 | 17.12 ms | maskconver_mobilenetv3p5_rf_256_coco.yaml |
+| MobileNet-MH | 640x640 | 3.4M | 9.58B | 37.2 | 24.93 ms | maskconver_mobilenetv3p5_rf_640_coco.yaml|
+
+
+### Citation
+
+Should you find this repository useful, please consider citing:
+
+```
+@inproceedings{rashwan2024maskconver,
+ title={MaskConver: Revisiting Pure Convolution Model for Panoptic Segmentation},
+ author={Abdullah Rashwan and Jiageng Zhang and Ali Taalimi and Fan Yang and Xingyi Zhou and Chaochao Yan and Liang-Chieh Chen and Yeqing Li},
+ year={2024},
+ booktitle={2024 IEEE Winter Conference on Applications of Computer Vision (WACV)},
+ organization={IEEE}
+}
+```
+
+
+
+
+
diff --git a/official/projects/maskconver/__init__.py b/official/projects/maskconver/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/maskconver/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/maskconver/configs/__init__.py b/official/projects/maskconver/configs/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/maskconver/configs/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/maskconver/configs/backbones.py b/official/projects/maskconver/configs/backbones.py
new file mode 100644
index 00000000000..04a60882208
--- /dev/null
+++ b/official/projects/maskconver/configs/backbones.py
@@ -0,0 +1,43 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Backbones configurations."""
+import dataclasses
+
+from typing import List, Optional
+from official.modeling import hyperparams
+from official.vision.configs.google import backbones
+
+
+@dataclasses.dataclass
+class ResNetUNet(hyperparams.Config):
+ """ResNetUNet config."""
+ model_id: int = 50
+ depth_multiplier: float = 1.0
+ stem_type: str = 'v0'
+ se_ratio: float = 0.0
+ stochastic_depth_drop_rate: float = 0.0
+ scale_stem: bool = True
+ resnetd_shortcut: bool = False
+ replace_stem_max_pool: bool = False
+ bn_trainable: bool = True
+ classification_output: bool = False
+ upsample_kernel_sizes: Optional[List[int]] = None
+ upsample_repeats: Optional[List[int]] = None
+ upsample_filters: Optional[List[int]] = None
+
+
+@dataclasses.dataclass
+class Backbone(backbones.Backbone):
+ resnet_unet: ResNetUNet = dataclasses.field(default_factory=ResNetUNet)
diff --git a/official/projects/maskconver/configs/decoders.py b/official/projects/maskconver/configs/decoders.py
new file mode 100644
index 00000000000..a6bc60876eb
--- /dev/null
+++ b/official/projects/maskconver/configs/decoders.py
@@ -0,0 +1,36 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Decoders configurations."""
+import dataclasses
+
+from official.modeling import hyperparams
+from official.vision.configs import decoders
+
+
+@dataclasses.dataclass
+class MaskConverFPN(hyperparams.Config):
+ """FPN config."""
+ num_filters: int = 256
+ fusion_type: str = 'sum'
+ use_separable_conv: bool = False
+ use_keras_layer: bool = False
+ use_layer_norm: bool = True
+ depthwise_kernel_size: int = 7
+
+
+@dataclasses.dataclass
+class Decoder(decoders.Decoder):
+ maskconver_fpn: MaskConverFPN = dataclasses.field(
+ default_factory=MaskConverFPN)
diff --git a/official/projects/maskconver/configs/experiments/imagenet_resnetunet_tpu.yaml b/official/projects/maskconver/configs/experiments/imagenet_resnetunet_tpu.yaml
new file mode 100644
index 00000000000..b2285f752e2
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/imagenet_resnetunet_tpu.yaml
@@ -0,0 +1,40 @@
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ num_classes: 1001
+ input_size: [384, 384, 3]
+ backbone:
+ type: 'resnet_unet'
+ resnet_unet:
+ model_id: 50
+ stochastic_depth_drop_rate: 0.1
+ classification_output: true
+ upsample_kernel_sizes: [7, 7, 7]
+ upsample_repeats: [18, 1, 1]
+ upsample_filters: [384, 384, 384]
+ norm_activation:
+ activation: 'gelu'
+ norm_momentum: 0.0
+ norm_epsilon: 0.00001
+ use_sync_bn: true
+ dropout_rate: 0.0
+ validation_data:
+ global_batch_size: 1024
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ initial_learning_rate: 0.004
+ name: CosineDecay
+ offset: 0
+ type: cosine
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.05
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
diff --git a/official/projects/maskconver/configs/experiments/maskconver_mobilenetv3p5_rf_256_coco.yaml b/official/projects/maskconver/configs/experiments/maskconver_mobilenetv3p5_rf_256_coco.yaml
new file mode 100644
index 00000000000..ab53f2d30e6
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_mobilenetv3p5_rf_256_coco.yaml
@@ -0,0 +1,111 @@
+# Train on 4x8 TPU and eval on GPU. PQ: 29.66
+# http://tb/7684098868201894928
+# Note: Get PQ 29.7 with official evaluation.
+# Note: We pad the model output to 640x640 for better eval.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ init_checkpoint: 'maskconver_seg_mnv3p5rf_coco_200k/43437096'
+ init_checkpoint_modules: ['backbone', 'decoder']
+ losses:
+ l2_weight_decay: 0.00001
+ mask_weight: 5.0
+ model:
+ input_size: [256, 256, 3]
+ level: 3
+ embedding_size: 256
+ padded_output_size: [640, 640]
+ num_instances: 50
+ norm_activation:
+ activation: 'relu'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ backbone:
+ mobilenet:
+ filter_size_scale: 1.0
+ model_id: MobileNetMultiAVGSeg
+ stochastic_depth_drop_rate: 0.0
+ output_stride: 16
+ type: mobilenet
+ decoder:
+ aspp:
+ dilation_rates: [6, 12, 18]
+ dropout_rate: 0.0
+ level: 4
+ num_filters: 256
+ spp_layer_version: v1
+ use_depthwise_convolution: true
+ type: 'aspp'
+ class_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ num_convs: 2
+ per_pixel_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ num_convs: 2
+ mask_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ num_convs: 2
+ panoptic_generator:
+ object_mask_threshold: 0.01
+ overlap_threshold: 0.7
+ rescale_predictions: true
+ small_area_threshold: 256
+ train_data:
+ global_batch_size: 64
+ parser:
+ max_num_stuff_centers: 1
+ gaussaian_iou: 0.7
+ aug_scale_max: 1.9
+ aug_scale_min: 0.1
+ aug_type: null
+ validation_data:
+ global_batch_size: 1
+ parser:
+ segmentation_resize_eval_groundtruth: false
+ segmentation_groundtruth_padded_size: [640, 640]
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 500000
+ initial_learning_rate: 0.04
+ type: cosine
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 2000
+ type: linear
+ steps_per_loop: 100
+ summary_interval: 1000
+ train_steps: 500000
+ validation_interval: 1000
+ validation_steps: 5000
+ checkpoint_interval: 1000
diff --git a/official/projects/maskconver/configs/experiments/maskconver_mobilenetv3p5_rf_640_coco.yaml b/official/projects/maskconver/configs/experiments/maskconver_mobilenetv3p5_rf_640_coco.yaml
new file mode 100644
index 00000000000..fd6bc6589bb
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_mobilenetv3p5_rf_640_coco.yaml
@@ -0,0 +1,110 @@
+# Train on 4x8 TPU and eval on GPU. PQ: 37.02
+# http://tb/3805525975368328422
+# Note: Get PQ 37.2 with official evaluation.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ init_checkpoint: 'maskconver_seg_mnv3p5rf_coco_200k/43437096'
+ init_checkpoint_modules: ['backbone', 'decoder']
+ losses:
+ l2_weight_decay: 0.00001
+ mask_weight: 5.0
+ model:
+ input_size: [640, 640, 3]
+ level: 3
+ embedding_size: 256
+ padded_output_size: [640, 640]
+ num_instances: 50
+ norm_activation:
+ activation: 'relu'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ backbone:
+ mobilenet:
+ filter_size_scale: 1.0
+ model_id: MobileNetMultiAVGSeg
+ stochastic_depth_drop_rate: 0.0
+ output_stride: 16
+ type: mobilenet
+ decoder:
+ aspp:
+ dilation_rates: [6, 12, 18]
+ dropout_rate: 0.0
+ level: 4
+ num_filters: 256
+ spp_layer_version: v1
+ use_depthwise_convolution: true
+ type: 'aspp'
+ class_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ num_convs: 2
+ per_pixel_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ num_convs: 2
+ mask_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ num_convs: 2
+ panoptic_generator:
+ object_mask_threshold: 0.01
+ overlap_threshold: 0.7
+ rescale_predictions: true
+ small_area_threshold: 256
+ train_data:
+ global_batch_size: 64
+ parser:
+ max_num_stuff_centers: 1
+ gaussaian_iou: 0.7
+ aug_scale_max: 1.9
+ aug_scale_min: 0.1
+ aug_type: null
+ validation_data:
+ global_batch_size: 1
+ parser:
+ segmentation_resize_eval_groundtruth: false
+ segmentation_groundtruth_padded_size: [640, 640]
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 500000
+ initial_learning_rate: 0.04
+ type: cosine
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 2000
+ type: linear
+ steps_per_loop: 100
+ summary_interval: 1000
+ train_steps: 500000
+ validation_interval: 1000
+ validation_steps: 5000
+ checkpoint_interval: 1000
diff --git a/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_coco.yaml b/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_coco.yaml
new file mode 100644
index 00000000000..df56fc3ddb1
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_coco.yaml
@@ -0,0 +1,104 @@
+# For pre-training.
+# Train on 4x4 TPU and eval on GPU.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ init_checkpoint: 'mnv_avg_imagenet/32040105'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 0.00001
+ mask_weight: 10.0
+ model:
+ input_size: [640, 640, 3]
+ num_classes: 91
+ level: 3
+ embedding_size: 256
+ padded_output_size: [640, 640]
+ norm_activation:
+ activation: 'relu'
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
+ use_sync_bn: true
+ backbone:
+ mobilenet:
+ filter_size_scale: 1.0
+ model_id: MobileNetMultiAVGSeg
+ output_stride: 16
+ type: mobilenet
+ decoder:
+ aspp:
+ dilation_rates: [6, 12, 18]
+ level: 4
+ num_filters: 256
+ use_depthwise_convolution: true
+ type: 'aspp'
+ class_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ use_depthwise_convolution: true
+ num_convs: 2
+ mask_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ use_depthwise_convolution: true
+ num_convs: 2
+ per_pixel_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 3
+ low_level_num_filters: 256
+ num_filters: 256
+ use_depthwise_convolution: true
+ num_convs: 2
+ panoptic_generator:
+ object_mask_threshold: 0.01
+ small_area_threshold: 256
+ overlap_threshold: 0.7
+ rescale_predictions: true
+ train_data:
+ input_path: 'image_segmentation/dataset/mscoco/mscoco_alltasks_trainvalminusminival2014'
+ output_size: [640, 640]
+ global_batch_size: 32
+ gaussaian_iou: 0.7
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ aug_type: null
+ validation_data:
+ input_path: 'image_segmentation/dataset/mscoco/mscoco_alltasks_minival2014''
+ output_size: [640, 640]
+ global_batch_size: 1
+ resize_eval_groundtruth: false
+ groundtruth_padded_size: [640, 640]
+trainer:
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 200000
+ initial_learning_rate: 0.01
+ type: cosine
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 500
+ type: linear
+ train_steps: 200000
+ validation_steps: 5000
+ steps_per_loop: 100
+ validation_interval: 5000
+ checkpoint_interval: 5000
+ summary_interval: 100
diff --git a/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_dlv3psummerge_cityscapes.yaml b/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_dlv3psummerge_cityscapes.yaml
new file mode 100644
index 00000000000..4615938f915
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_dlv3psummerge_cityscapes.yaml
@@ -0,0 +1,116 @@
+# Train on 4x8 TPU and eval on 4x4 TPU. Best mIoU: 75.54%
+# http://tb/2097099367278640078
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: float32
+task:
+ init_checkpoint: 'maskconver_seg_mnv3p5rf_coco_200k/43437096'
+ init_checkpoint_modules: [backbone]
+ losses:
+ l2_weight_decay: 0.0001
+ mask_weight: 5.0
+ model:
+ input_size: [1024, 2048, 3]
+ num_classes: 19
+ level: 4
+ embedding_size: 128
+ padded_output_size: [1024, 2048]
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ backbone:
+ mobilenet:
+ model_id: MobileNetMultiAVGSeg
+ output_intermediate_endpoints: true
+ output_stride: 16
+ type: mobilenet
+ decoder:
+ aspp:
+ level: 4
+ dilation_rates: [6, 12, 18]
+ num_filters: 128
+ pool_kernel_size: [512, 1024]
+ use_depthwise_convolution: true
+ type: aspp
+ class_head:
+ feature_fusion: deeplabv3plus
+ level: 4
+ low_level: '4'
+ low_level_num_filters: 48
+ num_convs: 2
+ num_filters: 128
+ use_depthwise_convolution: true
+ per_pixel_embedding_head:
+ feature_fusion: deeplabv3plus
+ level: 4
+ low_level: '4'
+ low_level_num_filters: 48
+ num_convs: 2
+ num_filters: 128
+ use_depthwise_convolution: true
+ mask_embedding_head:
+ feature_fusion: deeplabv3plus
+ level: 4
+ low_level: '4'
+ low_level_num_filters: 48
+ num_convs: 2
+ num_filters: 128
+ use_depthwise_convolution: true
+ panoptic_generator:
+ object_mask_threshold: 0.01
+ small_area_threshold: 256
+ overlap_threshold: 0.0
+ rescale_predictions: false
+ train_data:
+ input_path: 'dataset/cityscapes/tfrecord/train_fine'
+ is_training: true
+ global_batch_size: 64
+ output_size: [1024, 2048]
+ gaussaian_iou: 0.7
+ aug_scale_max: 1.9
+ aug_scale_min: 0.1
+ aug_type:
+ autoaug:
+ augmentation_name: panoptic_deeplab_policy
+ cutout_const: 100
+ translate_const: 250
+ type: autoaug
+ max_num_stuff_centers: 5
+ validation_data:
+ input_path: 'dataset/cityscapes/tfrecord/val_fine'
+ is_training: false
+ drop_remainder: false
+ global_batch_size: 16
+ output_size: [1024, 2048]
+ resize_eval_groundtruth: true
+ groundtruth_padded_size: [1024, 2048]
+trainer:
+ best_checkpoint_eval_metric: mean_iou
+ best_checkpoint_export_subdir: best_ckpt
+ best_checkpoint_metric_comp: higher
+ optimizer_config:
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ decay_steps: 100000
+ initial_learning_rate: 0.1
+ type: cosine
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 925
+ type: linear
+ train_steps: 100000
+ validation_steps: 31
+ steps_per_loop: 185
+ validation_interval: 185
+ checkpoint_interval: 185
+ summary_interval: 185
diff --git a/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_dlv3psummerge_pascal.yaml b/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_dlv3psummerge_pascal.yaml
new file mode 100644
index 00000000000..517ef7fae12
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_seg_mobilenetv3p5_rf_dlv3psummerge_pascal.yaml
@@ -0,0 +1,112 @@
+# Train and eval on 4x4 TPU. Best mIoU 76.4%
+# http://tb/1941403217464639568
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: float32
+task:
+ init_checkpoint: 'maskconver_seg_mnetv3p5rf_coco_conv_bn997/44073905'
+ init_checkpoint_modules: [backbone, decoder]
+ losses:
+ l2_weight_decay: 1.0e-04
+ mask_weight: 5.0
+ model:
+ input_size: [512, 512, 3]
+ num_classes: 21
+ level: 4
+ embedding_size: 256
+ padded_output_size: [512, 512]
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
+ use_sync_bn: true
+ backbone:
+ mobilenet:
+ model_id: MobileNetMultiAVGSeg
+ output_stride: 16
+ type: mobilenet
+ decoder:
+ aspp:
+ level: 4
+ dilation_rates: [6, 12, 18]
+ num_filters: 256
+ use_depthwise_convolution: true
+ type: aspp
+ class_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 4
+ low_level_num_filters: 256
+ num_convs: 2
+ use_depthwise_convolution: true
+ per_pixel_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 4
+ low_level_num_filters: 256
+ num_convs: 1
+ use_depthwise_convolution: true
+ mask_embedding_head:
+ feature_fusion: deeplabv3plus_sum_to_merge
+ level: 4
+ low_level: 4
+ low_level_num_filters: 256
+ num_convs: 1
+ use_depthwise_convolution: true
+ panoptic_generator:
+ object_mask_threshold: 0.1
+ small_area_threshold: 256
+ overlap_threshold: 0.0
+ rescale_predictions: false
+ train_data:
+ input_path: gs://**/pascal_voc_seg/train_aug*
+ is_training: true
+ global_batch_size: 32
+ output_size: [512, 512]
+ gaussaian_iou: 0.7
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ aug_type:
+ autoaug:
+ augmentation_name: panoptic_deeplab_policy
+ cutout_const: 100
+ translate_const: 250
+ type: autoaug
+ validation_data:
+ input_path: gs://**/pascal_voc_seg/val*
+ is_training: false
+ drop_remainder: false
+ global_batch_size: 32
+ output_size: [512, 512]
+ resize_eval_groundtruth: true
+ groundtruth_padded_size: [512, 512]
+trainer:
+ best_checkpoint_eval_metric: mean_iou
+ best_checkpoint_export_subdir: best_ckpt
+ best_checkpoint_metric_comp: higher
+ optimizer_config:
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ polynomial:
+ decay_steps: 33000
+ initial_learning_rate: 0.007
+ power: 0.9
+ type: polynomial
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 660
+ type: linear
+ train_steps: 33000
+ validation_steps: 22
+ steps_per_loop: 165
+ validation_interval: 165
+ checkpoint_interval: 165
+ summary_interval: 165
diff --git a/official/projects/maskconver/configs/experiments/maskconver_seg_spinenetseg_se_coco.yaml b/official/projects/maskconver/configs/experiments/maskconver_seg_spinenetseg_se_coco.yaml
new file mode 100644
index 00000000000..8f589d5d9e3
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_seg_spinenetseg_se_coco.yaml
@@ -0,0 +1,108 @@
+# For pre-training.
+# Train on 4x4 TPU and eval on GPU.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ init_checkpoint: 'maskconver_spinenet49se_imgnet/39653711/ema_checkpoints'
+ init_checkpoint_modules: 'backbone'
+ losses:
+ l2_weight_decay: 0.00001
+ mask_weight: 10.0
+ ignore_label: 255
+ model:
+ input_size: [640, 640, 3]
+ num_classes: 91
+ level: 3
+ embedding_size: 256
+ padded_output_size: [640, 640]
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ backbone:
+ type: 'spinenet_seg'
+ spinenet_seg:
+ model_id: '49'
+ stem_type: 'v1'
+ resnetd_shortcut: true
+ se_ratio: 0.25
+ replace_stem_max_pool: true
+ stochastic_depth_drop_rate: 0.1
+ decoder:
+ aspp:
+ level: 3
+ dilation_rates: [12, 24, 36]
+ num_filters: 256
+ use_depthwise_convolution: false
+ type: aspp
+ class_head:
+ feature_fusion: 'deeplabv3'
+ level: 3
+ num_convs: 2
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ per_pixel_embedding_head:
+ feature_fusion: 'deeplabv3'
+ level: 3
+ num_convs: 2
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ mask_embedding_head:
+ feature_fusion: 'deeplabv3'
+ level: 3
+ num_convs: 2
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ panoptic_generator:
+ object_mask_threshold: 0.1
+ small_area_threshold: 256
+ overlap_threshold: 0.0
+ rescale_predictions: true
+ train_data:
+ input_path: 'image_segmentation/dataset/mscoco/mscoco_alltasks_trainvalminusminival2014'
+ global_batch_size: 32
+ gaussaian_iou: 0.7
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ aug_type:
+ autoaug:
+ augmentation_name: panoptic_deeplab_policy
+ cutout_const: 100
+ translate_const: 250
+ type: autoaug
+ validation_data:
+ input_path: 'image_segmentation/dataset/mscoco/mscoco_alltasks_minival2014''
+ global_batch_size: 1
+ resize_eval_groundtruth: false
+ groundtruth_padded_size: [640, 640]
+trainer:
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 200000
+ initial_learning_rate: 0.01
+ type: cosine
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 500
+ type: linear
+ train_steps: 200000
+ validation_steps: 5000
+ steps_per_loop: 100
+ validation_interval: 5000
+ checkpoint_interval: 5000
+ summary_interval: 100
diff --git a/official/projects/maskconver/configs/experiments/maskconver_seg_spinenetseg_se_pascal.yaml b/official/projects/maskconver/configs/experiments/maskconver_seg_spinenetseg_se_pascal.yaml
new file mode 100644
index 00000000000..47c3a035325
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_seg_spinenetseg_se_pascal.yaml
@@ -0,0 +1,107 @@
+# Train and eval on 4x4 TPU. Best mIoU: 83.95%
+# http://tb/6094566970421137074
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ init_checkpoint: 'maskconver_seg_sn49_coco_200k/40849575' # maskconver seg coco things
+ init_checkpoint_modules: ['backbone', 'decoder']
+ losses:
+ l2_weight_decay: 0.00001
+ mask_weight: 10.0
+ ignore_label: 255
+ model:
+ input_size: [512, 512, 3]
+ embedding_size: 256
+ padded_output_size: [512, 512]
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ backbone:
+ type: 'spinenet_seg'
+ spinenet_seg:
+ model_id: '49'
+ stem_type: 'v1'
+ resnetd_shortcut: true
+ se_ratio: 0.25
+ replace_stem_max_pool: true
+ stochastic_depth_drop_rate: 0.1
+ decoder:
+ aspp:
+ level: 3
+ pool_kernel_size: null
+ dilation_rates: [12, 24, 36]
+ use_depthwise_convolution: false
+ type: aspp
+ class_head:
+ level: 3
+ num_convs: 2
+ feature_fusion: 'deeplabv3'
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ per_pixel_embedding_head:
+ level: 3
+ num_convs: 2
+ feature_fusion: 'deeplabv3'
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ mask_embedding_head:
+ level: 3
+ num_convs: 2
+ feature_fusion: 'deeplabv3'
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ panoptic_generator:
+ object_mask_threshold: 0.1
+ overlap_threshold: 0.0
+ rescale_predictions: false
+ small_area_threshold: 256
+ train_data:
+ global_batch_size: 32
+ gaussaian_iou: 0.7
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ aug_type:
+ autoaug:
+ augmentation_name: panoptic_deeplab_policy
+ cutout_const: 100
+ translate_const: 250
+ type: autoaug
+ validation_data:
+ global_batch_size: 32
+ resize_eval_groundtruth: true
+ groundtruth_padded_size: [512, 512]
+trainer:
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ optimizer_config:
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ stepwise:
+ boundaries: [25000]
+ values: [0.001, 0.0001]
+ type: stepwise
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 500
+ type: linear
+ train_steps: 30000
+ validation_steps: 45
+ steps_per_loop: 100
+ validation_interval: 1000
+ checkpoint_interval: 1000
+ summary_interval: 100
diff --git a/official/projects/maskconver/configs/experiments/maskconver_spinenetseg_coco.yaml b/official/projects/maskconver/configs/experiments/maskconver_spinenetseg_coco.yaml
new file mode 100644
index 00000000000..d93b1612500
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/maskconver_spinenetseg_coco.yaml
@@ -0,0 +1,81 @@
+# Train on 4x8 TPU and eval on GPU.
+# http://tb/6152990635147738742 - PQ 47.
+# Note: Above tensorboard has the best number, but using a different config compared to this file.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ init_checkpoint: 'spinenetseg/28194596/ema_checkpoints'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 0.00001
+ mask_weight: 10.0
+ model:
+ input_size: [640, 640, 3]
+ embedding_size: 256
+ padded_output_size: [640, 640]
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ backbone:
+ type: 'spinenet_seg'
+ spinenet_seg:
+ model_id: '49'
+ stochastic_depth_drop_rate: 0.0
+ decoder:
+ aspp:
+ level: 3
+ pool_kernel_size: null
+ dilation_rates: [12, 24, 36]
+ type: aspp
+ class_head:
+ level: 3
+ num_convs: 2
+ feature_fusion: 'deeplabv3'
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ per_pixel_embedding_head:
+ level: 3
+ num_convs: 2
+ feature_fusion: 'deeplabv3'
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ mask_embedding_head:
+ level: 3
+ num_convs: 2
+ feature_fusion: 'deeplabv3'
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ panoptic_generator:
+ object_mask_threshold: 0.2
+ train_data:
+ global_batch_size: 64
+ parser:
+ gaussaian_iou: 0.6
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 200000
+ initial_learning_rate: 0.08
+ type: cosine
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 2000
+ type: linear
+ train_steps: 200000
+ steps_per_loop: 100
+ validation_interval: 1000
+ checkpoint_interval: 1000
+ summary_interval: 1000
diff --git a/official/projects/maskconver/configs/experiments/multiscale_maskconver_resnet50fpn_coco.yaml b/official/projects/maskconver/configs/experiments/multiscale_maskconver_resnet50fpn_coco.yaml
new file mode 100644
index 00000000000..dd1368e1503
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/multiscale_maskconver_resnet50fpn_coco.yaml
@@ -0,0 +1,88 @@
+# Train on 4x8 TPU and eval on GPU.
+# http://tb/6152990635147738742 - PQ 47.
+# Note: Above tensorboard has the best number, but using a different config compared to this file.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ init_checkpoint: 'gs://cloud-tpu-checkpoints/vision-2.0/resnet50_imagenet/ckpt-28080'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 1.0e-5
+ mask_weight: 5.0
+ model:
+ input_size: [640, 640, 3]
+ embedding_size: 256
+ padded_output_size: [640, 640]
+ min_level: 3
+ max_level: 7
+ num_instances: 100
+ norm_activation:
+ activation: swish
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ backbone:
+ type: resnet
+ decoder:
+ fpn:
+ num_filters: 256
+ type: fpn
+ class_head:
+ num_convs: 2
+ num_filters: 256
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ per_pixel_embedding_head:
+ decoder_min_level: 3
+ decoder_max_level: 7
+ level: 2
+ num_convs: 2
+ num_filters: 256
+ feature_fusion: panoptic_fpn_fusion
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ mask_embedding_head:
+ num_convs: 2
+ num_filters: 256
+ prediction_kernel_size: 3
+ use_depthwise_convolution: false
+ upsample_factor: 1
+ panoptic_generator:
+ object_mask_threshold: 0.01
+ overlap_threshold: 0.7
+ small_area_threshold: 4
+ train_data:
+ global_batch_size: 64
+ parser:
+ gaussaian_iou: 0.7
+ aug_scale_max: 1.9
+ aug_scale_min: 0.1
+ max_num_stuff_centers: 1
+ validation_data:
+ global_batch_size: 16
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 200000
+ initial_learning_rate: 0.04
+ type: cosine
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 2000
+ type: linear
+ train_steps: 200000
+ validation_steps: 312
+ steps_per_loop: 154
+ validation_interval: 308
+ checkpoint_interval: 308
+ summary_interval: 308
diff --git a/official/projects/maskconver/configs/experiments/multiscale_maskconver_resnetufpn_coco.yaml b/official/projects/maskconver/configs/experiments/multiscale_maskconver_resnetufpn_coco.yaml
new file mode 100644
index 00000000000..85e0cbed61e
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/multiscale_maskconver_resnetufpn_coco.yaml
@@ -0,0 +1,240 @@
+runtime:
+ mixed_precision_dtype: float32
+task:
+ differential_privacy_config: null
+ init_checkpoint: 'resnetubn_fpn_seg_coco1024_1811se/48546144/ema_checkpoints'
+ init_checkpoint_modules: ["backbone"]
+ losses:
+ alpha: 2.0
+ beta: 4.0
+ ignore_label: 0
+ l2_weight_decay: 0.0
+ loss_weight: 1.0
+ mask_weight: 10.0
+ top_k_percent_pixels_category: 1.0
+ top_k_percent_pixels_instance: 1.0
+ use_groundtruth_dimension: true
+ model:
+ anchor:
+ anchor_size: 8.0
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 1
+ backbone:
+ type: 'resnet_unet'
+ resnet_unet:
+ model_id: 50
+ stochastic_depth_drop_rate: 0.1
+ classification_output: false
+ bn_trainable: false
+ upsample_kernel_sizes: [7, 7, 7]
+ upsample_repeats: [18, 1, 1]
+ upsample_filters: [384, 384, 384]
+ class_head:
+ num_convs: 4
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ depthwise_kernel_size: 7
+ decoder:
+ maskconver_fpn:
+ num_filters: 256
+ use_separable_conv: true
+ fusion_type: concat
+ type: maskconver_fpn
+ mask_decoder: null
+ embedding_size: 256
+ input_size: [1280, 1280, 3]
+ level: 3
+ mask_embedding_head:
+ num_convs: 4
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ depthwise_kernel_size: 7
+ use_depthwise_convolution: true
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: gelu
+ norm_epsilon: 0.001
+ norm_momentum: 0.0
+ use_sync_bn: true
+ num_anchors: 100
+ num_classes: 201
+ num_instances: 256
+ num_thing_classes: 91
+ padded_output_size: [640, 640]
+ panoptic_fusion_num_filters: 256
+ panoptic_generator:
+ object_mask_threshold: 0.2
+ overlap_threshold: 0.75
+ rescale_predictions: true
+ small_area_threshold: 16
+ use_hardware_optimization: false
+ per_pixel_embedding_head:
+ decoder_max_level: 7
+ decoder_min_level: 3
+ feature_fusion: pyramid_fusion
+ level: 2
+ low_level: 2
+ low_level_num_filters: 48
+ num_convs: 4
+ num_filters: 256
+ prediction_kernel_size: 1
+ upsample_factor: 1
+ use_depthwise_convolution: true
+ depthwise_kernel_size: 7
+ name: null
+ panoptic_quality_evaluator:
+ ignored_label: 0
+ max_instances_per_category: 256
+ num_categories: 201
+ offset: 16777216
+ report_per_class_metrics: false
+ rescale_predictions: true
+ train_data:
+ apply_tf_data_service_before_batching: false
+ block_length: 1
+ cache: false
+ cycle_length: null
+ decoder:
+ simple_decoder:
+ include_panoptic_masks: true
+ mask_binarize_threshold: null
+ panoptic_category_mask_key: image/panoptic/category_mask
+ panoptic_instance_mask_key: image/panoptic/instance_mask
+ regenerate_source_id: false
+ type: simple_decoder
+ deterministic: null
+ drop_remainder: true
+ dtype: float32
+ enable_shared_tf_data_service_between_parallel_trainers: false
+ enable_tf_data_service: false
+ file_type: tfrecord
+ global_batch_size: 128
+ input_path: 'datasets/coco/tfrecords/train-nocrowd'
+ is_training: true
+ num_examples: -1
+ parser:
+ copypaste:
+ copypaste_aug_scale_max: 1.5
+ copypaste_aug_scale_min: 0.025
+ aug_scale_max: 2.9
+ aug_scale_min: 0.05
+ copypaste_frequency: 0.5
+ aug_rand_hflip: true
+ aug_scale_max: 1.0
+ aug_scale_min: 1.0
+ aug_type:
+ autoaug:
+ augmentation_name: panoptic_deeplab_policy
+ cutout_const: 100
+ translate_const: 250
+ type: autoaug
+ fpn_low_range: [0, 100, 200, 400, 800]
+ fpn_high_range: [100, 200, 400, 800, 12500000]
+ gaussaian_iou: 0.7
+ panoptic_ignore_label: 0
+ segmentation_groundtruth_padded_size: []
+ segmentation_ignore_label: 0
+ segmentation_resize_eval_groundtruth: true
+ prefetch_buffer_size: 2
+ seed: null
+ sharding: true
+ shuffle_buffer_size: 10000
+ tf_data_service_address: null
+ tf_data_service_job_name: null
+ tfds_as_supervised: false
+ tfds_data_dir: ''
+ tfds_name: ''
+ tfds_skip_decoding_feature: ''
+ tfds_split: ''
+ trainer_id: null
+ validation_data:
+ apply_tf_data_service_before_batching: false
+ block_length: 1
+ cache: false
+ cycle_length: null
+ decoder:
+ simple_decoder:
+ include_panoptic_masks: true
+ mask_binarize_threshold: null
+ panoptic_category_mask_key: image/panoptic/category_mask
+ panoptic_instance_mask_key: image/panoptic/instance_mask
+ regenerate_source_id: false
+ type: simple_decoder
+ deterministic: null
+ drop_remainder: false
+ dtype: float32
+ enable_shared_tf_data_service_between_parallel_trainers: false
+ enable_tf_data_service: false
+ file_type: tfrecord
+ global_batch_size: 16
+ input_path: 'datasets/coco/tfrecords/val-nocrowd'
+ is_training: false
+ num_examples: -1
+ parser:
+ aug_rand_hflip: false
+ aug_scale_max: 1.0
+ aug_scale_min: 1.0
+ aug_type:
+ type: null
+ fpn_high_range: []
+ fpn_low_range: []
+ panoptic_ignore_label: 0
+ segmentation_groundtruth_padded_size: [640, 640]
+ segmentation_ignore_label: 0
+ segmentation_resize_eval_groundtruth: false
+ prefetch_buffer_size: 8
+trainer:
+ best_checkpoint_eval_metric: 'panoptic_quality/All_pq'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ allow_tpu_summary: false
+ checkpoint_interval: 1000
+ continuous_eval_timeout: 3600
+ eval_tf_function: true
+ eval_tf_while_loop: false
+ loss_upper_bound: 1000000.0
+ max_to_keep: 5
+ optimizer_config:
+ ema:
+ average_decay: 0.99996
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 270000
+ initial_learning_rate: 0.001
+ name: CosineDecay
+ offset: 0
+ type: cosine
+ optimizer:
+ adamw:
+ amsgrad: false
+ beta_1: 0.9
+ beta_2: 0.999
+ clipnorm: null
+ clipvalue: null
+ epsilon: 1.0e-07
+ exclude_from_weight_decay: null
+ global_clipnorm: null
+ gradient_clip_norm: 1.0
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0.0
+ warmup_steps: 500
+ type: linear
+ steps_per_loop: 25
+ summary_interval: 100
+ train_steps: 270000
+ train_tf_function: true
+ train_tf_while_loop: true
+ validation_interval: 308
+ validation_steps: 312
+ validation_summary_subdir: validation
diff --git a/official/projects/maskconver/configs/experiments/resnetu_cocoseg_tpu.yaml b/official/projects/maskconver/configs/experiments/resnetu_cocoseg_tpu.yaml
new file mode 100644
index 00000000000..fcd9295a542
--- /dev/null
+++ b/official/projects/maskconver/configs/experiments/resnetu_cocoseg_tpu.yaml
@@ -0,0 +1,111 @@
+# --experiment_type=seg_deeplabv3_pascal
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ min_level: 2
+ max_level: 7
+ num_classes: 91
+ input_size: [1024, 1024, 3]
+ backbone:
+ type: 'resnet_unet'
+ resnet_unet:
+ model_id: 50
+ stochastic_depth_drop_rate: 0.1
+ classification_output: false
+ bn_trainable: false
+ upsample_filters: [384, 384, 384]
+ upsample_kernel_sizes: [7, 7, 7]
+ upsample_repeats: [18, 1, 1]
+ decoder:
+ fpn:
+ num_filters: 256
+ use_separable_conv: true
+ type: fpn
+ head:
+ feature_fusion: pyramid_fusion
+ level: 2
+ num_convs: 4
+ use_depthwise_convolution: true
+ prediction_kernel_size: 1
+ norm_activation:
+ activation: 'gelu'
+ norm_epsilon: 0.001
+ norm_momentum: 0.0
+ use_sync_bn: true
+ init_checkpoint: 'res50unet_imgnet_1811_se0625/78934910/ema_checkpoints'
+ init_checkpoint_modules: 'backbone'
+ losses:
+ l2_weight_decay: 0.0
+ top_k_percent_pixels: 1.0 # only backpropagate loss for the topk 100% pixels.
+ train_data:
+ output_size: [1024, 1024]
+ input_path: 'image_segmentation/dataset/mscoco/mscoco_alltasks_trainvalminusminival2014'
+ is_training: true
+ global_batch_size: 256
+ dtype: 'float32'
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ output_size: [1024, 1024]
+ input_path: 'image_segmentation/dataset/mscoco/mscoco_alltasks_minival2014''
+ is_training: false
+ global_batch_size: 16
+ dtype: 'float32'
+ drop_remainder: false
+ resize_eval_groundtruth: true
+trainer:
+ allow_tpu_summary: false
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: ''
+ best_checkpoint_metric_comp: higher
+ checkpoint_interval: 200
+ continuous_eval_timeout: 3600
+ eval_tf_function: true
+ eval_tf_while_loop: false
+ loss_upper_bound: 1000000.0
+ max_to_keep: 5
+ optimizer_config:
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ alpha: 0.03
+ decay_steps: 64000
+ initial_learning_rate: 0.0003
+ name: CosineDecay
+ offset: 0
+ type: cosine
+ optimizer:
+ adamw:
+ amsgrad: false
+ beta_1: 0.9
+ beta_2: 0.999
+ clipnorm: null
+ clipvalue: null
+ epsilon: 1.0e-07
+ exclude_from_weight_decay: null
+ global_clipnorm: null
+ gradient_clip_norm: 1.0
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0.000001
+ warmup_steps: 2000
+ type: linear
+ recovery_begin_steps: 0
+ recovery_max_trials: 0
+ steps_per_loop: 50
+ summary_interval: 100
+ train_steps: 64000
+ train_tf_function: true
+ train_tf_while_loop: true
+ validation_interval: 308
+ validation_steps: 312
+ validation_summary_subdir: validation
diff --git a/official/projects/maskconver/configs/maskconver.py b/official/projects/maskconver/configs/maskconver.py
new file mode 100644
index 00000000000..77d3e4b3fe5
--- /dev/null
+++ b/official/projects/maskconver/configs/maskconver.py
@@ -0,0 +1,523 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic Mask R-CNN configuration definition."""
+
+import dataclasses
+import os
+from typing import List, Optional
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.modeling import optimization
+from official.projects.maskconver.configs import backbones
+from official.projects.maskconver.configs import decoders
+from official.vision.configs import common
+from official.vision.configs import maskrcnn
+from official.vision.configs import semantic_segmentation
+
+
+_COCO_INPUT_PATH_BASE = 'coco/tfrecords'
+_COCO_TRAIN_EXAMPLES = 118287
+_COCO_VAL_EXAMPLES = 5000
+
+# PASCAL VOC 2012 Dataset
+_PASCAL_TRAIN_EXAMPLES = 10582
+_PASCAL_VAL_EXAMPLES = 1449
+_PASCAL_INPUT_PATH_BASE = 'gs://**/pascal_voc_seg'
+
+# Cityscapes Dataset
+_CITYSCAPES_TRAIN_EXAMPLES = 2975
+_CITYSCAPES_VAL_EXAMPLES = 500
+_CITYSCAPES_INPUT_PATH_BASE = 'cityscapes/tfrecord'
+
+
+# pytype: disable=wrong-keyword-args
+# pylint: disable=unexpected-keyword-arg
+
+
+@dataclasses.dataclass
+class CopyPaste(hyperparams.Config):
+ copypaste_frequency: float = 0.5
+ aug_scale_min: float = 0.1
+ aug_scale_max: float = 1.9
+ copypaste_aug_scale_max: float = 1.0
+ copypaste_aug_scale_min: float = 0.05
+
+
+@dataclasses.dataclass
+class Parser(maskrcnn.Parser):
+ """MaskConver parser config."""
+ # If segmentation_resize_eval_groundtruth is set to False, original image
+ # sizes are used for eval. In that case,
+ # segmentation_groundtruth_padded_size has to be specified too to allow for
+ # batching the variable input sizes of images.
+ segmentation_resize_eval_groundtruth: bool = True
+ segmentation_groundtruth_padded_size: List[int] = dataclasses.field(
+ default_factory=list)
+ segmentation_ignore_label: int = 0
+ panoptic_ignore_label: int = 0
+ # Setting this to true will enable parsing category_mask and instance_mask.
+ include_panoptic_masks: bool = True
+ gaussaian_iou: float = 0.7
+ max_num_stuff_centers: int = 3
+ aug_type: common.Augmentation = dataclasses.field(
+ default_factory=common.Augmentation
+ )
+ copypaste: CopyPaste = dataclasses.field(default_factory=CopyPaste)
+
+
+@dataclasses.dataclass
+class TfExampleDecoder(common.TfExampleDecoder):
+ """A simple TF Example decoder config."""
+ # Setting this to true will enable decoding category_mask and instance_mask.
+ include_panoptic_masks: bool = True
+ panoptic_category_mask_key: str = 'image/panoptic/category_mask'
+ panoptic_instance_mask_key: str = 'image/panoptic/instance_mask'
+
+
+@dataclasses.dataclass
+class DataDecoder(common.DataDecoder):
+ """Data decoder config."""
+ simple_decoder: TfExampleDecoder = dataclasses.field(
+ default_factory=TfExampleDecoder
+ )
+
+
+@dataclasses.dataclass
+class DataConfig(maskrcnn.DataConfig):
+ """Input config for training."""
+ decoder: DataDecoder = dataclasses.field(default_factory=DataDecoder)
+ parser: Parser = dataclasses.field(default_factory=Parser)
+ dtype: str = 'float32'
+ prefetch_buffer_size: int = 8
+
+
+@dataclasses.dataclass
+class Anchor(hyperparams.Config):
+ num_scales: int = 1
+ aspect_ratios: List[float] = dataclasses.field(
+ default_factory=lambda: [0.5, 1.0, 2.0])
+ anchor_size: float = 8.0
+
+
+@dataclasses.dataclass
+class PanopticGenerator(hyperparams.Config):
+ """MaskConver panoptic generator."""
+ object_mask_threshold: float = 0.001
+ small_area_threshold: int = 0
+ overlap_threshold: float = 0.5
+ rescale_predictions: bool = True
+ use_hardware_optimization: bool = False
+
+
+@dataclasses.dataclass
+class SegmentationHead(semantic_segmentation.SegmentationHead):
+ """Segmentation head config."""
+ depthwise_kernel_size: int = 7
+ use_layer_norm: bool = False
+
+
+@dataclasses.dataclass
+class MaskConver(hyperparams.Config):
+ """MaskConver model config."""
+ num_classes: int = 0
+ num_thing_classes: int = 0
+ num_instances: int = 100
+ embedding_size: int = 512
+ padded_output_size: List[int] = dataclasses.field(default_factory=list)
+ input_size: List[int] = dataclasses.field(default_factory=list)
+ min_level: int = 2
+ max_level: int = 6
+ num_anchors: int = 100
+ panoptic_fusion_num_filters: int = 256
+ anchor: Anchor = dataclasses.field(default_factory=Anchor)
+ level: int = 3
+ class_head: SegmentationHead = dataclasses.field(
+ default_factory=SegmentationHead
+ )
+ mask_embedding_head: SegmentationHead = dataclasses.field(
+ default_factory=SegmentationHead
+ )
+ per_pixel_embedding_head: SegmentationHead = dataclasses.field(
+ default_factory=SegmentationHead
+ )
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=backbones.Backbone
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder(type='identity')
+ )
+ mask_decoder: Optional[decoders.Decoder] = dataclasses.field(
+ default_factory=lambda: decoders.Decoder(type='identity')
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
+ panoptic_generator: PanopticGenerator = dataclasses.field(
+ default_factory=PanopticGenerator
+ )
+
+
+@dataclasses.dataclass
+class Losses(hyperparams.Config):
+ """maskconver loss config."""
+ l2_weight_decay: float = 0.0
+ ignore_label: int = 0
+ use_groundtruth_dimension: bool = True
+ top_k_percent_pixels_category: float = 1.0
+ top_k_percent_pixels_instance: float = 1.0
+ loss_weight: float = 1.0
+ mask_weight: float = 10.0
+ beta: float = 4.0
+ alpha: float = 2.0
+
+
+@dataclasses.dataclass
+class PanopticQualityEvaluator(hyperparams.Config):
+ """Panoptic Quality Evaluator config."""
+ num_categories: int = 2
+ ignored_label: int = 0
+ max_instances_per_category: int = 256
+ offset: int = 256 * 256 * 256
+ is_thing: List[float] = dataclasses.field(
+ default_factory=list)
+ rescale_predictions: bool = True
+ report_per_class_metrics: bool = False
+
+###################################
+###### PANOPTIC SEGMENTATION ######
+###################################
+
+
+@dataclasses.dataclass
+class MaskConverTask(cfg.TaskConfig):
+ """MaskConverTask task config."""
+ model: MaskConver = dataclasses.field(default_factory=MaskConver)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ # pylint: disable=g-long-lambda
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(
+ is_training=False, drop_remainder=False
+ )
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+
+ init_checkpoint_modules: Optional[List[str]] = dataclasses.field(
+ default_factory=list)
+ panoptic_quality_evaluator: PanopticQualityEvaluator = dataclasses.field(
+ default_factory=PanopticQualityEvaluator
+ )
+ # pylint: enable=g-long-lambda
+
+
+@exp_factory.register_config_factory('maskconver_coco')
+def maskconver_coco() -> cfg.ExperimentConfig:
+ """COCO panoptic segmentation with MaskConver."""
+ train_batch_size = 64
+ eval_batch_size = 8
+ # steps_per_epoch = _COCO_TRAIN_EXAMPLES // train_batch_size
+ validation_steps = _COCO_VAL_EXAMPLES // eval_batch_size
+
+ # coco panoptic dataset has category ids ranging from [0-200] inclusive.
+ # 0 is not used and represents the background class
+ # ids 1-91 represent thing categories (91)
+ # ids 92-200 represent stuff categories (109)
+ # for the segmentation task, we continue using id=0 for the background
+ # and map all thing categories to id=1, the remaining 109 stuff categories
+ # are shifted by an offset=90 given by num_thing classes - 1. This shifting
+ # will make all the stuff categories begin from id=2 and end at id=110
+ num_panoptic_categories = 201
+ num_thing_categories = 91
+ # num_semantic_segmentation_classes = 111
+
+ is_thing = [False]
+ for idx in range(1, num_panoptic_categories):
+ is_thing.append(True if idx < num_thing_categories else False)
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(
+ mixed_precision_dtype='float32', enable_xla=False),
+ task=MaskConverTask(
+ init_checkpoint='gs://cloud-tpu-checkpoints/vision-2.0/resnet50_imagenet/ckpt-28080', # pylint: disable=line-too-long
+ init_checkpoint_modules=['backbone'],
+ model=MaskConver(
+ num_classes=201, num_thing_classes=91, input_size=[512, 512, 3],
+ padded_output_size=[512, 512]),
+ losses=Losses(l2_weight_decay=0.0),
+ train_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'train-nocrowd*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ parser=Parser(
+ aug_rand_hflip=True, aug_scale_min=0.5, aug_scale_max=1.5,
+ aug_type=common.Augmentation(
+ type='autoaug',
+ autoaug=common.AutoAugment(
+ augmentation_name='panoptic_deeplab_policy')))),
+ validation_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'val-nocrowd*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ parser=Parser(
+ segmentation_resize_eval_groundtruth=True,
+ segmentation_groundtruth_padded_size=[640, 640]),
+ drop_remainder=False),
+ panoptic_quality_evaluator=PanopticQualityEvaluator(
+ num_categories=num_panoptic_categories,
+ ignored_label=0,
+ is_thing=is_thing,
+ rescale_predictions=True)),
+ trainer=cfg.TrainerConfig(
+ train_steps=200000,
+ validation_steps=validation_steps,
+ validation_interval=1000,
+ steps_per_loop=1000,
+ summary_interval=1000,
+ checkpoint_interval=1000,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.08,
+ 'decay_steps': 200000,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 2000,
+ 'warmup_learning_rate': 0.0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
+
+###################################
+###### SEMANTIC SEGMENTATION ######
+###################################
+
+
+@dataclasses.dataclass
+class SegDataConfig(cfg.DataConfig):
+ """Input config for training."""
+ output_size: List[int] = dataclasses.field(default_factory=list)
+ # If crop_size is specified, image will be resized first to
+ # output_size, then crop of size crop_size will be cropped.
+ crop_size: List[int] = dataclasses.field(default_factory=list)
+ input_path: str = ''
+ global_batch_size: int = 0
+ is_training: bool = True
+ dtype: str = 'float32'
+ shuffle_buffer_size: int = 1000
+ prefetch_buffer_size: int = 8
+ cycle_length: int = 10
+ # If resize_eval_groundtruth is set to False, original image sizes are used
+ # for eval. In that case, groundtruth_padded_size has to be specified too to
+ # allow for batching the variable input sizes of images.
+ resize_eval_groundtruth: bool = True
+ groundtruth_padded_size: List[int] = dataclasses.field(default_factory=list)
+ aug_scale_min: float = 1.0
+ aug_scale_max: float = 1.0
+ aug_rand_hflip: bool = True
+ preserve_aspect_ratio: bool = True
+ aug_policy: Optional[str] = None
+ drop_remainder: bool = True
+ file_type: str = 'tfrecord'
+ gaussaian_iou: float = 0.7
+ max_num_stuff_centers: int = 3
+ max_num_instances: int = 100
+ aug_type: common.Augmentation = dataclasses.field(
+ default_factory=common.Augmentation)
+
+
+@dataclasses.dataclass
+class MaskConverSegTask(cfg.TaskConfig):
+ """MaskConverTask task config."""
+ model: MaskConver = dataclasses.field(default_factory=MaskConver)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: SegDataConfig(is_training=True)
+ )
+ # pylint: disable=g-long-lambda
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: SegDataConfig(
+ is_training=False, drop_remainder=False
+ )
+ )
+ # pylint: enable=g-long-lambda
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+
+ init_checkpoint_modules: Optional[List[str]] = dataclasses.field(
+ default_factory=list)
+
+
+@exp_factory.register_config_factory('maskconver_seg_pascal')
+def maskconver_seg_pascal() -> cfg.ExperimentConfig:
+ """COCO panoptic segmentation with MaskConver."""
+ train_batch_size = 64
+ eval_batch_size = 8
+ validation_steps = _PASCAL_VAL_EXAMPLES // eval_batch_size
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(
+ mixed_precision_dtype='float32', enable_xla=False),
+ task=MaskConverSegTask(
+ init_checkpoint='gs://cloud-tpu-checkpoints/vision-2.0/resnet50_imagenet/ckpt-28080', # pylint: disable=line-too-long
+ init_checkpoint_modules=['backbone'],
+ model=MaskConver(
+ num_classes=21, num_thing_classes=91, input_size=[512, 512, 3],
+ padded_output_size=[512, 512]),
+ losses=Losses(l2_weight_decay=0.00004),
+ train_data=SegDataConfig(
+ input_path=os.path.join(_PASCAL_INPUT_PATH_BASE, 'train_aug*'),
+ output_size=[512, 512],
+ is_training=True,
+ global_batch_size=train_batch_size,
+ aug_scale_min=0.5,
+ aug_scale_max=2.0,
+ aug_type=common.Augmentation(
+ type='autoaug',
+ autoaug=common.AutoAugment(
+ augmentation_name='panoptic_deeplab_policy'))),
+ validation_data=SegDataConfig(
+ input_path=os.path.join(_PASCAL_INPUT_PATH_BASE, 'val*'),
+ output_size=[512, 512],
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ resize_eval_groundtruth=False,
+ groundtruth_padded_size=[512, 512],
+ drop_remainder=False)),
+ trainer=cfg.TrainerConfig(
+ train_steps=200000,
+ validation_steps=validation_steps,
+ validation_interval=1000,
+ steps_per_loop=1000,
+ summary_interval=1000,
+ checkpoint_interval=1000,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.08,
+ 'decay_steps': 200000,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 2000,
+ 'warmup_learning_rate': 0.0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
+
+
+@exp_factory.register_config_factory('maskconver_seg_cityscapes')
+def maskconver_seg_cityscapes() -> cfg.ExperimentConfig:
+ """Cityscapes semantic segmentation with MaskConver."""
+ train_batch_size = 32
+ eval_batch_size = 8
+ validation_steps = _CITYSCAPES_VAL_EXAMPLES // eval_batch_size
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(
+ mixed_precision_dtype='float32', enable_xla=False),
+ task=MaskConverSegTask(
+ init_checkpoint='maskconver_seg_mnv3p5rf_coco_200k/43437096', # pylint: disable=line-too-long
+ init_checkpoint_modules=['backbone'],
+ model=MaskConver(
+ num_classes=19, input_size=[None, None, 3],
+ padded_output_size=[1024, 2048]),
+ losses=Losses(l2_weight_decay=0.00004),
+ train_data=SegDataConfig(
+ input_path=os.path.join(_CITYSCAPES_INPUT_PATH_BASE,
+ 'train_fine*'),
+ output_size=[1024, 2048],
+ crop_size=[512, 1024],
+ is_training=True,
+ global_batch_size=train_batch_size,
+ aug_scale_min=0.5,
+ aug_scale_max=2.0,
+ aug_type=common.Augmentation(
+ type='autoaug',
+ autoaug=common.AutoAugment(
+ augmentation_name='panoptic_deeplab_policy'))),
+ validation_data=SegDataConfig(
+ input_path=os.path.join(_CITYSCAPES_INPUT_PATH_BASE, 'val_fine*'),
+ output_size=[1024, 2048],
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ resize_eval_groundtruth=False,
+ groundtruth_padded_size=[1024, 2048],
+ drop_remainder=False)),
+ trainer=cfg.TrainerConfig(
+ train_steps=100000,
+ validation_steps=validation_steps,
+ validation_interval=185,
+ steps_per_loop=185,
+ summary_interval=185,
+ checkpoint_interval=185,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 0.01,
+ 'decay_steps': 100000,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 925,
+ 'warmup_learning_rate': 0.0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
diff --git a/official/projects/maskconver/configs/multiscale_maskconver.py b/official/projects/maskconver/configs/multiscale_maskconver.py
new file mode 100644
index 00000000000..93bc6de69e7
--- /dev/null
+++ b/official/projects/maskconver/configs/multiscale_maskconver.py
@@ -0,0 +1,215 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Multi-scale Maskconver configuration definition."""
+
+import dataclasses
+import os
+from typing import List, Optional
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.modeling import optimization
+from official.projects.maskconver.configs import maskconver
+from official.vision.configs import common
+
+# pylint: disable=unused-private-name
+# pytype: disable=wrong-keyword-args
+# pylint: disable=unexpected-keyword-arg
+
+_COCO_INPUT_PATH_BASE = 'coco/tfrecords'
+_COCO_TRAIN_EXAMPLES = 118287
+_COCO_VAL_EXAMPLES = 5000
+
+TfExampleDecoder = maskconver.TfExampleDecoder
+DataDecoder = maskconver.DataDecoder
+DataConfig = maskconver.DataConfig
+Losses = maskconver.Losses
+PanopticGenerator = maskconver.PanopticGenerator
+PanopticQualityEvaluator = maskconver.PanopticQualityEvaluator
+
+
+@dataclasses.dataclass
+class CopyPaste(hyperparams.Config):
+ copypaste_frequency: float = 1.0
+ aug_scale_min: float = 0.1
+ aug_scale_max: float = 1.9
+ copypaste_aug_scale_max: float = 1.0
+ copypaste_aug_scale_min: float = 0.05
+
+
+@dataclasses.dataclass
+class Parser(hyperparams.Config):
+ """MaskConver parser config."""
+ aug_rand_hflip: bool = False
+ aug_scale_min: float = 1.0
+ aug_scale_max: float = 1.0
+ # If segmentation_resize_eval_groundtruth is set to False, original image
+ # sizes are used for eval. In that case,
+ # segmentation_groundtruth_padded_size has to be specified too to allow for
+ # batching the variable input sizes of images.
+ segmentation_resize_eval_groundtruth: bool = True
+ segmentation_groundtruth_padded_size: List[int] = dataclasses.field(
+ default_factory=list)
+ segmentation_ignore_label: int = 0
+ panoptic_ignore_label: int = 0
+ # Setting this to true will enable parsing category_mask and instance_mask.
+ include_panoptic_masks: bool = True
+ gaussaian_iou: float = 0.7
+ max_num_instances: int = 256
+ aug_type: common.Augmentation = dataclasses.field(
+ default_factory=common.Augmentation)
+ fpn_low_range: List[int] = dataclasses.field(default_factory=list)
+ fpn_high_range: List[int] = dataclasses.field(default_factory=list)
+ mask_target_level: int = 1
+ copypaste: CopyPaste = dataclasses.field(default_factory=CopyPaste)
+
+
+@dataclasses.dataclass
+class MultiScaleMaskConverHead(hyperparams.Config):
+ """Segmentation head config."""
+ num_convs: int = 4
+ num_filters: int = 256
+ use_depthwise_convolution: bool = False
+ prediction_kernel_size: int = 3
+ upsample_factor: int = 1
+ depthwise_kernel_size: int = 7
+ use_layer_norm: bool = True
+
+
+@dataclasses.dataclass
+class MultiScaleMaskConver(maskconver.MaskConver):
+ """Multi-scale MaskConver model config."""
+ min_level: int = 3
+ max_level: int = 7
+ num_instances: int = 100
+ class_head: MultiScaleMaskConverHead = dataclasses.field(
+ default_factory=MultiScaleMaskConverHead
+ )
+ mask_embedding_head: MultiScaleMaskConverHead = dataclasses.field(
+ default_factory=MultiScaleMaskConverHead
+ )
+ per_pixel_embedding_head: maskconver.SegmentationHead = dataclasses.field(
+ default_factory=lambda: maskconver.SegmentationHead(use_layer_norm=True)
+ )
+
+
+###################################
+###### PANOPTIC SEGMENTATION ######
+###################################
+
+
+@dataclasses.dataclass
+class MultiScaleMaskConverTask(cfg.TaskConfig):
+ """MaskConverTask task config."""
+ model: MultiScaleMaskConver = dataclasses.field(
+ default_factory=MultiScaleMaskConver
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ # pylint: disable=g-long-lambda
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(
+ is_training=False, drop_remainder=False
+ )
+ )
+ # pylint: enable=g-long-lambda
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+
+ init_checkpoint_modules: Optional[List[str]] = dataclasses.field(
+ default_factory=list)
+ panoptic_quality_evaluator: PanopticQualityEvaluator = dataclasses.field(
+ default_factory=PanopticQualityEvaluator
+ )
+
+
+@exp_factory.register_config_factory('multiscale_maskconver_coco')
+def multiscale_maskconver_coco() -> cfg.ExperimentConfig:
+ """COCO panoptic segmentation with MaskConver."""
+ train_batch_size = 128
+ eval_batch_size = 1
+ validation_steps = _COCO_VAL_EXAMPLES // eval_batch_size
+
+ # coco panoptic dataset has category ids ranging from [0-200] inclusive.
+ # 0 is not used and represents the background class
+ # ids 1-91 represent thing categories (91)
+ # ids 92-200 represent stuff categories (109)
+ # for the segmentation task, we continue using id=0 for the background
+ # and map all thing categories to id=1, the remaining 109 stuff categories
+ # are shifted by an offset=90 given by num_thing classes - 1. This shifting
+ # will make all the stuff categories begin from id=2 and end at id=110
+ num_panoptic_categories = 201
+ num_thing_categories = 91
+ # num_semantic_segmentation_classes = 111
+
+ is_thing = [False]
+ for idx in range(1, num_panoptic_categories):
+ is_thing.append(True if idx < num_thing_categories else False)
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(
+ mixed_precision_dtype='float32', enable_xla=False),
+ task=MultiScaleMaskConverTask(
+ init_checkpoint='gs://cloud-tpu-checkpoints/vision-2.0/resnet50_imagenet/ckpt-28080', # pylint: disable=line-too-long
+ init_checkpoint_modules=['backbone'],
+ model=MultiScaleMaskConver(
+ num_classes=201,
+ num_thing_classes=91,
+ input_size=[640, 640, 3],
+ padded_output_size=[640, 640]),
+ losses=Losses(l2_weight_decay=1e-4),
+ train_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ parser=Parser(
+ aug_rand_hflip=True,
+ aug_scale_min=0.1,
+ aug_scale_max=1.9,
+ fpn_low_range=[0, 40, 80, 160, 320],
+ fpn_high_range=[64, 128, 256, 512, 10000000],
+ aug_type=common.Augmentation(
+ type='autoaug',
+ autoaug=common.AutoAugment(
+ augmentation_name='panoptic_deeplab_policy')))),
+ validation_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ parser=Parser(
+ segmentation_resize_eval_groundtruth=False,
+ segmentation_groundtruth_padded_size=[640, 640]),
+ drop_remainder=False),
+ panoptic_quality_evaluator=PanopticQualityEvaluator(
+ num_categories=num_panoptic_categories,
+ ignored_label=0,
+ is_thing=is_thing,
+ rescale_predictions=True)),
+ trainer=cfg.TrainerConfig(
+ train_steps=200000,
+ validation_steps=validation_steps,
+ validation_interval=1000,
+ steps_per_loop=1000,
+ summary_interval=1000,
+ checkpoint_interval=1000,
+ optimizer_config=optimization.OptimizationConfig()),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
+
diff --git a/official/projects/maskconver/dataloaders/maskconver_segmentation_input.py b/official/projects/maskconver/dataloaders/maskconver_segmentation_input.py
new file mode 100644
index 00000000000..8d3b4b994df
--- /dev/null
+++ b/official/projects/maskconver/dataloaders/maskconver_segmentation_input.py
@@ -0,0 +1,382 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data parser and processing for maskconver segmentation dataloader."""
+from typing import Optional
+
+import tensorflow as tf, tf_keras
+from tensorflow_addons import image as tfa_image
+
+from official.projects.centernet.ops import target_assigner
+from official.vision.configs import common
+from official.vision.dataloaders import segmentation_input
+from official.vision.ops import augment
+from official.vision.ops import preprocess_ops
+
+
+class Parser(segmentation_input.Parser):
+ """Parser to parse an image and its annotations into a dictionary of tensors."""
+
+ def __init__(self,
+ output_size,
+ num_classes,
+ crop_size=None,
+ level=3,
+ resize_eval_groundtruth=True,
+ groundtruth_padded_size=None,
+ ignore_label=255,
+ aug_rand_hflip=False,
+ preserve_aspect_ratio=True,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ gaussian_iou=0.7,
+ aug_type: Optional[common.Augmentation] = None,
+ max_num_instances: int = 100,
+ max_num_stuff_centers: int = 3,
+ dtype='float32'):
+ """Initializes parameters for parsing annotations in the dataset.
+
+ Args:
+ output_size: `Tensor` or `list` for [height, width] of output image. The
+ output_size should be divided by the largest feature stride 2^max_level.
+ num_classes: `int` for number of classes.
+ crop_size: `Tensor` or `list` for [height, width] of the crop. If
+ specified a training crop of size crop_size is returned. This is useful
+ for cropping original images during training while evaluating on
+ original image sizes.
+ level: level for output masks.
+ resize_eval_groundtruth: `bool`, if True, eval groundtruth masks are
+ resized to output_size.
+ groundtruth_padded_size: `Tensor` or `list` for [height, width]. When
+ resize_eval_groundtruth is set to False, the groundtruth masks are
+ padded to this size.
+ ignore_label: `int` the pixel with ignore label will not used for training
+ and evaluation.
+ aug_rand_hflip: `bool`, if True, augment training with random
+ horizontal flip.
+ preserve_aspect_ratio: `bool`, if True, the aspect ratio is preserved,
+ otherwise, the image is resized to output_size.
+ aug_scale_min: `float`, the minimum scale applied to `output_size` for
+ data augmentation during training.
+ aug_scale_max: `float`, the maximum scale applied to `output_size` for
+ data augmentation during training.
+ gaussian_iou: `float`, used for generating the center heatmaps.
+ aug_type: Optional[common.Augmentation].
+ max_num_instances: max number of instances.
+ max_num_stuff_centers: max number of stuff centers.
+ dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
+ """
+ super(Parser, self).__init__(
+ output_size=output_size,
+ crop_size=crop_size,
+ resize_eval_groundtruth=resize_eval_groundtruth,
+ groundtruth_padded_size=groundtruth_padded_size,
+ ignore_label=ignore_label,
+ aug_rand_hflip=aug_rand_hflip,
+ preserve_aspect_ratio=preserve_aspect_ratio,
+ aug_scale_min=aug_scale_min,
+ aug_scale_max=aug_scale_max)
+
+ self._num_classes = num_classes
+ self.level = level
+ self._gaussian_iou = gaussian_iou
+ self._max_num_stuff_centers = max_num_stuff_centers
+ self._max_num_instances = max_num_instances
+
+ if aug_type and aug_type.type:
+ if aug_type.type == 'autoaug':
+ self._augmenter = augment.AutoAugment(
+ augmentation_name=aug_type.autoaug.augmentation_name,
+ cutout_const=aug_type.autoaug.cutout_const,
+ translate_const=aug_type.autoaug.translate_const)
+ else:
+ raise ValueError('Augmentation policy {} not supported.'.format(
+ aug_type.type))
+ else:
+ self._augmenter = None
+
+ def _prepare_image_and_label(self, data, use_augment=False):
+ """Prepare normalized image and label."""
+ image = tf.io.decode_image(data['image/encoded'], channels=3)
+ label = tf.io.decode_image(data['image/segmentation/class/encoded'],
+ channels=1)
+ height = data['image/height']
+ width = data['image/width']
+ image = tf.reshape(image, (height, width, 3))
+
+ label = tf.reshape(label, (1, height, width))
+ label = tf.cast(label, tf.float32)
+
+ if use_augment and self._augmenter is not None:
+ image = self._augmenter.distort(image)
+ # Normalizes image with mean and std pixel values.
+ image = preprocess_ops.normalize_image(image)
+
+ if not self._preserve_aspect_ratio:
+ label = tf.reshape(label, [data['image/height'], data['image/width'], 1])
+ image = tf.image.resize(image, self._output_size, method='bilinear')
+ label = tf.image.resize(label, self._output_size, method='nearest')
+ label = tf.reshape(label[:, :, -1], [1] + self._output_size)
+
+ return image, label
+
+ def _parse_train_data(self, data):
+ """Parses data for training and evaluation."""
+ image, label = self._prepare_image_and_label(data, use_augment=True)
+
+ if self._crop_size:
+ label = tf.reshape(label, [data['image/height'], data['image/width'], 1])
+ # If output_size is specified, resize image, and label to desired
+ # output_size.
+ if self._output_size:
+ image = tf.image.resize(image, self._output_size, method='bilinear')
+ label = tf.image.resize(label, self._output_size, method='nearest')
+
+ image_mask = tf.concat([image, label], axis=2)
+ image_mask_crop = tf.image.random_crop(image_mask,
+ self._crop_size + [4])
+ image = image_mask_crop[:, :, :-1]
+ label = tf.reshape(image_mask_crop[:, :, -1], [1] + self._crop_size)
+
+ # Flips image randomly during training.
+ if self._aug_rand_hflip:
+ image, _, label = preprocess_ops.random_horizontal_flip(
+ image, masks=label)
+
+ train_image_size = self._crop_size if self._crop_size else self._output_size
+ # Resizes and crops image.
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ train_image_size,
+ train_image_size,
+ aug_scale_min=self._aug_scale_min,
+ aug_scale_max=self._aug_scale_max)
+
+ # Resizes and crops boxes.
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+
+ # Pad label and make sure the padded region assigned to the ignore label.
+ # The label is first offset by +1 and then padded with 0.
+ label += 1
+ label = tf.expand_dims(label, axis=3)
+ label = preprocess_ops.resize_and_crop_masks(
+ label, image_scale, train_image_size, offset)
+ label -= 1
+ label = tf.where(tf.equal(label, -1),
+ self._ignore_label * tf.ones_like(label), label)
+ label = tf.squeeze(label, axis=0)
+ label = tf.squeeze(label, axis=-1)
+
+ valid_mask = tf.not_equal(label, self._ignore_label)
+
+ stuff_classes, _ = tf.unique(tf.reshape(label, [-1]))
+ stuff_classes = tf.boolean_mask(
+ stuff_classes, stuff_classes != self._ignore_label)
+
+ # Compute bboxes for stuff classes
+ def get_bbox(mask):
+ rows = tf.math.count_nonzero(mask, axis=0, keepdims=None, dtype=tf.bool)
+ columns = tf.math.count_nonzero(
+ mask, axis=1, keepdims=None, dtype=tf.bool)
+ indices = tf.where(tf.equal(columns, True))
+ y_min = indices[0]
+ y_max = indices[-1]
+ indices = tf.where(tf.equal(rows, True))
+ x_min = indices[0]
+ x_max = indices[-1]
+ return tf.stack([y_min, x_min, y_max, x_max])
+
+ def get_stuff_class_label(stuff_class, k=self._max_num_stuff_centers):
+ mask = tf.cast(
+ label == stuff_class, tf.int32)
+ smoothed_mask = tf.nn.max_pool(mask[:, :, None], 21, 1, padding='SAME')
+ smoothed_mask = tf.nn.max_pool(smoothed_mask, 21, 1, padding='SAME')
+ smoothed_mask = tf.nn.max_pool(smoothed_mask, 21, 1, padding='SAME')
+
+ smoothed_mask = tf.cast(tf.squeeze(smoothed_mask, axis=2), tf.int32)
+ connected_components = tfa_image.connected_components( # pyrefly: ignore[not-callable]
+ images=smoothed_mask)
+ counts = tf.math.bincount(connected_components)
+ ids = tf.argsort(counts[1:], axis=-1, direction='DESCENDING') + 1
+
+ masks = tf.cast(tf.repeat(mask[None, :, :], k, axis=0), tf.float32)
+ stuff_classes = tf.cast(stuff_class, tf.int64) * tf.ones([k], tf.int64)
+
+ # Pad or clip ids
+ ids_length = tf.shape(ids)[0]
+ ids_length = tf.clip_by_value(ids_length, 0, k)
+ batch_weights = 1.0 * tf.ones([1, k], tf.float32)
+ ids = ids[:ids_length]
+ padding_length = tf.maximum(0, k - ids_length)
+ paddings = tf.cast(-1 * tf.ones([padding_length]), ids.dtype)
+ ids = tf.concat([ids, paddings], axis=0)
+
+ def get_bbox_and_batch_mask(island_id):
+ if island_id == -1:
+ bbox = tf.zeros([4], tf.float32)
+ batch_mask = -1 * tf.ones([1], tf.int32)
+ smask = tf.cast(tf.zeros_like(connected_components), tf.float32)
+ else:
+ smask = tf.cast(connected_components, tf.int32) == island_id
+ smask = tf.cast(
+ tf.logical_and(tf.cast(smask, tf.bool), tf.cast(mask, tf.bool)),
+ tf.float32)
+ if tf.reduce_sum(smask) < 1024:
+ batch_mask = -1 * tf.ones([1], tf.int32)
+ else:
+ batch_mask = tf.ones([1], tf.int32)
+ bbox = tf.cast(tf.reshape(get_bbox(smask), [-1]), tf.float32)
+ return smask, bbox, batch_mask
+
+ stuff_center_masks, bboxes, batch_mask = tf.map_fn(
+ get_bbox_and_batch_mask,
+ ids,
+ fn_output_signature=((tf.TensorSpec(train_image_size, tf.float32),
+ tf.TensorSpec([4,], tf.float32),
+ tf.TensorSpec([1], tf.int32))))
+
+ stuff_classes = tf.reshape(
+ stuff_classes, [1, k])
+ batch_mask = tf.reshape(batch_mask, [1, k])
+ stuff_center_masks = tf.reshape(
+ stuff_center_masks, [1, k] + train_image_size)
+ return masks[None, :, :, :], bboxes[
+ None, :, :], stuff_classes, batch_mask, batch_weights
+
+ centers = self._max_num_stuff_centers
+ stuff_masks, stuff_bboxes, stuff_classes, stuff_batch_mask, stuff_batch_weigths = tf.map_fn(
+ get_stuff_class_label,
+ stuff_classes,
+ fn_output_signature=(tf.TensorSpec([1, centers] + train_image_size),
+ tf.TensorSpec([1, centers, 4,]),
+ tf.TensorSpec([1, centers], tf.int64),
+ tf.TensorSpec([1, centers], tf.int32),
+ tf.TensorSpec([1, centers], tf.float32)))
+ stuff_masks = tf.reshape(stuff_masks, [-1] + train_image_size)
+ stuff_bboxes = tf.reshape(stuff_bboxes, [-1, 4])
+ stuff_classes = tf.reshape(stuff_classes, [-1])
+ stuff_batch_mask = tf.reshape(stuff_batch_mask, [-1])
+ stuff_batch_weigths = tf.reshape(stuff_batch_weigths, [-1])
+ masks = stuff_masks[stuff_batch_mask == 1]
+ bboxes = stuff_bboxes[stuff_batch_mask == 1]
+ classes = stuff_classes[stuff_batch_mask == 1]
+
+ width_ratio = 1 / float(2 ** self.level)
+ height_ratio = 1 / float(2 ** self.level)
+
+ # Original box coordinates
+ # [max_num_instances, ]
+ ytl, ybr = bboxes[..., 0], bboxes[..., 2]
+ xtl, xbr = bboxes[..., 1], bboxes[..., 3]
+ yct = (ytl + ybr) / 2
+ xct = (xtl + xbr) / 2
+
+ # Scaled box coordinates (could be floating point)
+ # [max_num_instances, ]
+ scale_xct = xct * width_ratio
+ scale_yct = yct * height_ratio
+
+ # Floor the scaled box coordinates to be placed on heatmaps
+ # [max_num_instances, ]
+ scale_xct_floor = tf.math.floor(scale_xct)
+ scale_yct_floor = tf.math.floor(scale_yct)
+
+ # Get the scaled box dimensions for computing the gaussian radius
+ # [max_num_instances, ]
+ box_widths = bboxes[..., 3] - bboxes[..., 1]
+ box_heights = bboxes[..., 2] - bboxes[..., 0]
+
+ box_widths = box_widths * width_ratio
+ box_heights = box_heights * height_ratio
+
+ ct_heatmap = target_assigner.assign_center_targets(
+ out_height=int(train_image_size[0] * height_ratio),
+ out_width=int(train_image_size[1] * width_ratio),
+ y_center=scale_yct,
+ x_center=scale_xct,
+ boxes_height=box_heights,
+ boxes_width=box_widths,
+ channel_onehot=tf.one_hot(
+ tf.cast(classes, tf.int32),
+ self._num_classes, off_value=0.),
+ gaussian_iou=self._gaussian_iou)
+ box_indices = tf.cast(
+ tf.stack([scale_yct_floor, scale_xct_floor], axis=-1), dtype=tf.int32)
+
+ seg_classes = preprocess_ops.clip_or_pad_to_fixed_size(
+ classes, self._max_num_instances, -1)
+ seg_masks = preprocess_ops.clip_or_pad_to_fixed_size(
+ masks, self._max_num_instances, -1.0)
+ seg_masks = tf.transpose(seg_masks, [1, 2, 0])
+ seg_boxes = preprocess_ops.clip_or_pad_to_fixed_size(
+ bboxes, self._max_num_instances, -1)
+
+ box_indices = preprocess_ops.clip_or_pad_to_fixed_size(
+ box_indices, self._max_num_instances, 0)
+
+ labels = {
+ 'image': image,
+ 'seg_ct_heatmaps': ct_heatmap,
+ 'seg_classes': seg_classes,
+ 'seg_masks': seg_masks,
+ 'seg_mask_weights': tf.cast(seg_classes >= 0, tf.float32),
+ 'seg_valid_mask': tf.cast(valid_mask, tf.float32),
+ 'seg_boxes': seg_boxes,
+ 'seg_box_indices': box_indices,
+ 'num_instances': tf.reduce_sum(
+ tf.cast(seg_classes >= 0, tf.float32)),
+ 'image_info': image_info,
+ }
+ return image, labels
+
+ def _parse_eval_data(self, data):
+ """Parses data for training and evaluation."""
+ image, label = self._prepare_image_and_label(data)
+ # The label is first offset by +1 and then padded with 0.
+ label += 1
+ label = tf.expand_dims(label, axis=3)
+
+ # Resizes and crops image.
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image, self._output_size, self._output_size)
+
+ if self._resize_eval_groundtruth:
+ # Resizes eval masks to match input image sizes. In that case, mean IoU
+ # is computed on output_size not the original size of the images.
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ label = preprocess_ops.resize_and_crop_masks(label, image_scale,
+ self._output_size, offset)
+ else:
+ label = tf.image.pad_to_bounding_box(
+ label, 0, 0, self._groundtruth_padded_size[0],
+ self._groundtruth_padded_size[1])
+
+ label -= 1
+ label = tf.where(tf.equal(label, -1),
+ self._ignore_label * tf.ones_like(label), label)
+ label = tf.squeeze(label, axis=0)
+
+ valid_mask = tf.not_equal(label, self._ignore_label)
+ labels = {
+ 'masks': label,
+ 'valid_masks': valid_mask,
+ 'image_info': image_info
+ }
+
+ # Cast image as self._dtype
+ image = tf.cast(image, dtype=self._dtype)
+
+ return image, labels
diff --git a/official/projects/maskconver/dataloaders/multiscale_maskconver_input.py b/official/projects/maskconver/dataloaders/multiscale_maskconver_input.py
new file mode 100644
index 00000000000..2e10502c5d8
--- /dev/null
+++ b/official/projects/maskconver/dataloaders/multiscale_maskconver_input.py
@@ -0,0 +1,528 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data parser and processing for Panoptic MaskConver."""
+from typing import Optional
+
+import tensorflow as tf, tf_keras
+
+from official.projects.centernet.ops import target_assigner
+from official.vision.configs import common
+from official.vision.dataloaders import parser
+from official.vision.dataloaders import tf_example_decoder
+from official.vision.dataloaders import utils
+from official.vision.ops import augment
+from official.vision.ops import preprocess_ops
+
+
+class TfExampleDecoder(tf_example_decoder.TfExampleDecoder):
+ """Tensorflow Example proto decoder."""
+
+ def __init__(
+ self,
+ regenerate_source_id: bool,
+ mask_binarize_threshold: float,
+ include_panoptic_masks: bool,
+ panoptic_category_mask_key: str = 'image/panoptic/category_mask',
+ panoptic_instance_mask_key: str = 'image/panoptic/instance_mask'):
+
+ super().__init__(
+ include_mask=True, regenerate_source_id=regenerate_source_id,
+ mask_binarize_threshold=None)
+
+ self._include_panoptic_masks = include_panoptic_masks
+ self._panoptic_category_mask_key = panoptic_category_mask_key
+ self._panoptic_instance_mask_key = panoptic_instance_mask_key
+
+ if include_panoptic_masks:
+ self._segmentation_keys_to_features = {
+ panoptic_category_mask_key:
+ tf.io.FixedLenFeature((), tf.string, default_value=''),
+ panoptic_instance_mask_key:
+ tf.io.FixedLenFeature((), tf.string, default_value='')}
+
+ def decode(self, serialized_example):
+ decoded_tensors = super().decode(serialized_example)
+
+ if self._include_panoptic_masks:
+ parsed_tensors = tf.io.parse_single_example(
+ serialized_example, self._segmentation_keys_to_features)
+ category_mask = tf.io.decode_image(
+ parsed_tensors[self._panoptic_category_mask_key],
+ channels=1)
+ instance_mask = tf.io.decode_image(
+ parsed_tensors[self._panoptic_instance_mask_key],
+ channels=1)
+ category_mask.set_shape([None, None, 1])
+ instance_mask.set_shape([None, None, 1])
+
+ decoded_tensors.update({
+ 'groundtruth_panoptic_category_mask':
+ category_mask,
+ 'groundtruth_panoptic_instance_mask':
+ instance_mask})
+ return decoded_tensors
+
+
+class Parser(parser.Parser):
+ """Parser to parse an image and its annotations into a dictionary of tensors."""
+
+ def __init__(self,
+ output_size,
+ min_level,
+ max_level,
+ fpn_low_range,
+ fpn_high_range,
+ aug_rand_hflip=False,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ max_num_instances=100,
+ segmentation_resize_eval_groundtruth=True,
+ segmentation_groundtruth_padded_size=None,
+ segmentation_ignore_label=255,
+ panoptic_ignore_label=0,
+ level=3,
+ mask_target_level=1,
+ num_panoptic_categories=201,
+ num_thing_categories=91,
+ gaussian_iou=0.7,
+ aug_type: Optional[common.Augmentation] = None,
+ dtype='float32'):
+ """Initializes parameters for parsing annotations in the dataset.
+
+ Args:
+ output_size: `Tensor` or `list` for [height, width] of output image. The
+ output_size should be divided by the largest feature stride 2^max_level.
+ min_level: `int` number of minimum level of the output feature pyramid.
+ max_level: `int` number of maximum level of the output feature pyramid.
+ fpn_low_range: List of `int`.
+ fpn_high_range: List of `int`.
+ aug_rand_hflip: `bool`, if True, augment training with random
+ horizontal flip.
+ aug_scale_min: `float`, the minimum scale applied to `output_size` for
+ data augmentation during training.
+ aug_scale_max: `float`, the maximum scale applied to `output_size` for
+ data augmentation during training.
+ max_num_instances: `int` number of maximum number of instances in an
+ image. The groundtruth data will be padded to `max_num_instances`.
+ segmentation_resize_eval_groundtruth: `bool`, for whether or not to
+ resize the eval groundtruth masks.
+ segmentation_groundtruth_padded_size: `Tensor` or `list` for [height,
+ width]. When resize_eval_groundtruth is set to False, the groundtruth
+ masks are padded to this size.
+ segmentation_ignore_label: `int` the pixels with ignore label will not be
+ used for training and evaluation.
+ panoptic_ignore_label: `int` the pixels with ignore label will not be used
+ by the PQ evaluator.
+ level: `int`, output level, used to generate target assignments.
+ mask_target_level: `int`, target level for panoptic masks, default is 1.
+ num_panoptic_categories: `int`, number of panoptic categories.
+ num_thing_categories: `int1, number of thing categories.
+ gaussian_iou: `float`, used for generating the center heatmaps.
+ aug_type: An optional Augmentation object with params for AutoAugment.
+ dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
+ """
+ super().__init__()
+
+ self.aug_rand_hflip = aug_rand_hflip
+ self._segmentation_resize_eval_groundtruth = (
+ segmentation_resize_eval_groundtruth
+ )
+ if (not segmentation_resize_eval_groundtruth) and (
+ segmentation_groundtruth_padded_size is None
+ ):
+ raise ValueError(
+ 'segmentation_groundtruth_padded_size ([height, width]) needs to be'
+ 'specified when segmentation_resize_eval_groundtruth is False.'
+ )
+ self._segmentation_groundtruth_padded_size = (
+ segmentation_groundtruth_padded_size
+ )
+ self._dtype = dtype
+ self._max_num_instances = max_num_instances
+ self._output_size = output_size
+ self._min_level = min_level
+ self._max_level = max_level
+ self._aug_scale_min = aug_scale_min
+ self._aug_scale_max = aug_scale_max
+ self._segmentation_ignore_label = segmentation_ignore_label
+ self._panoptic_ignore_label = panoptic_ignore_label
+ self.level = level
+ self._mask_target_level = mask_target_level
+ self._num_panoptic_categories = num_panoptic_categories
+ self._num_thing_categories = num_thing_categories
+ self._gaussian_iou = gaussian_iou
+ self.fpn_low_range = fpn_low_range
+ self.fpn_high_range = fpn_high_range
+
+ if aug_type and aug_type.type:
+ if aug_type.type == 'autoaug':
+ self._augmenter = augment.AutoAugment(
+ augmentation_name=aug_type.autoaug.augmentation_name,
+ cutout_const=aug_type.autoaug.cutout_const,
+ translate_const=aug_type.autoaug.translate_const)
+ else:
+ raise ValueError('Augmentation policy {} not supported.'.format(
+ aug_type.type))
+ else:
+ self._augmenter = None
+
+ def _parse_train_data(self, data):
+ """Parses data for training.
+
+ Args:
+ data: the decoded tensor dictionary from TfExampleDecoder.
+
+ Returns:
+ image: image tensor that is preproessed to have normalized value and
+ dimension [output_size[0], output_size[1], 3]
+ labels: a dictionary of tensors used for training. The following describes
+ {key: value} pairs in the dictionary.
+ """
+ # Flips image randomly during training.
+ if self.aug_rand_hflip:
+ instance_mask = data['groundtruth_panoptic_instance_mask']
+ category_mask = data['groundtruth_panoptic_category_mask']
+ image_mask = tf.concat([
+ tf.cast(data['image'], category_mask.dtype), category_mask,
+ instance_mask
+ ], axis=2)
+
+ image_mask, _, _ = preprocess_ops.random_horizontal_flip(image_mask)
+
+ instance_mask = image_mask[:, :, -1:]
+ category_mask = image_mask[:, :, -2:-1]
+ image = tf.cast(image_mask[:, :, :-2], tf.uint8)
+
+ if self._augmenter is not None:
+ image = self._augmenter.distort(image)
+
+ image = preprocess_ops.normalize_image(image) # pyrefly: ignore[unbound-name]
+
+ # Resizes and crops image.
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ self._output_size,
+ padded_size=preprocess_ops.compute_padded_size(self._output_size,
+ 2**self._max_level),
+ aug_scale_min=self._aug_scale_min,
+ aug_scale_max=self._aug_scale_max)
+
+ def _process_mask(mask, ignore_label, image_info):
+ mask = tf.cast(mask, dtype=tf.float32)
+ mask = tf.reshape(mask, shape=[1, data['height'], data['width'], 1])
+ mask += 1
+
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ mask = preprocess_ops.resize_and_crop_masks(
+ mask, image_scale, self._output_size, offset)
+
+ mask -= 1
+ # Assigns ignore label to the padded region.
+ mask = tf.where(
+ tf.equal(mask, -1),
+ ignore_label * tf.ones_like(mask),
+ mask)
+ mask = tf.squeeze(mask, axis=0)
+ return mask
+
+ panoptic_category_mask = _process_mask(
+ category_mask, # pyrefly: ignore[unbound-name]
+ -1, image_info)
+ panoptic_instance_mask = _process_mask(
+ instance_mask, # pyrefly: ignore[unbound-name]
+ self._panoptic_ignore_label, image_info)
+
+ mask_level = self._mask_target_level
+ mask_size = [
+ int(self._output_size[0] / 2**mask_level),
+ int(self._output_size[1] / 2**mask_level)]
+ panoptic_category_mask = tf.image.resize(
+ panoptic_category_mask, mask_size, method='nearest')
+ panoptic_instance_mask = tf.image.resize(
+ panoptic_instance_mask, mask_size, method='nearest')
+
+ panoptic_category_mask = tf.cast(panoptic_category_mask[..., 0], tf.int32)
+ panoptic_instance_mask = tf.cast(panoptic_instance_mask[..., 0], tf.int32)
+
+ padding_mask = tf.cast(panoptic_category_mask == -1, tf.float32)
+
+ def get_bbox(mask):
+ """Gets the bounding box of a mask."""
+ rows = tf.math.count_nonzero(mask, axis=0, keepdims=None, dtype=tf.bool)
+ columns = tf.math.count_nonzero(
+ mask, axis=1, keepdims=None, dtype=tf.bool)
+ indices = tf.where(tf.equal(columns, True))
+ y_min = indices[0]
+ y_max = indices[-1]
+ indices = tf.where(tf.equal(rows, True))
+ x_min = indices[0]
+ x_max = indices[-1]
+ return tf.stack([y_min, x_min, y_max, x_max])
+
+ def get_thing_class_label(instance_id):
+ """Gets bounding box, center, and labels for a mask."""
+ mask = tf.cast(panoptic_instance_mask == instance_id, tf.int32)
+ bbox = tf.cast(tf.reshape(get_bbox(mask), [-1]), tf.float32)
+ # Computing the mask centers.
+ center = tf.reduce_mean(tf.cast(tf.where(mask > 0), tf.float32), axis=0)
+ inds = tf.where(panoptic_instance_mask == instance_id)
+ thing_classes = tf.gather_nd(panoptic_category_mask, inds)
+ thing_classes, _, counts = tf.unique_with_counts(thing_classes)
+ max_index = tf.argmax(counts)
+ thing_class = tf.cast(thing_classes[max_index], tf.int64)
+ return tf.cast(mask, tf.float32), bbox, center, thing_class
+
+ instance_ids, _ = tf.unique(tf.reshape(panoptic_instance_mask, [-1]))
+ instance_ids = tf.boolean_mask(instance_ids, instance_ids > 0)
+ things_masks, things_bboxes, things_centers, things_classes = tf.map_fn(
+ get_thing_class_label,
+ instance_ids,
+ fn_output_signature=(
+ tf.TensorSpec(mask_size),
+ tf.TensorSpec([4]),
+ tf.TensorSpec([2]),
+ tf.int64,
+ ),
+ )
+
+ stuff_classes, _ = tf.unique(tf.reshape(panoptic_category_mask, [-1]))
+ stuff_classes = tf.boolean_mask(
+ stuff_classes, stuff_classes > self._num_thing_categories - 1
+ )
+
+ def get_stuff_class_label(stuff_class):
+ """Gets bounding box, center, and labels for a mask."""
+ mask = tf.cast(panoptic_category_mask == stuff_class, tf.int32)
+ bbox = tf.cast(tf.reshape(get_bbox(mask), [-1]), tf.float32)
+ center = tf.cast(tf.reduce_mean(tf.where(mask > 0), axis=0), tf.float32)
+ stuff_class = tf.cast(stuff_class, tf.int64)
+ return tf.cast(mask, tf.float32), bbox, center, stuff_class
+
+ stuff_masks, stuff_bboxes, stuff_centers, stuff_classes = tf.map_fn(
+ get_stuff_class_label,
+ stuff_classes,
+ fn_output_signature=(
+ tf.TensorSpec(mask_size),
+ tf.TensorSpec([4]),
+ tf.TensorSpec([2]),
+ tf.int64,
+ ),
+ )
+
+ # Start to generate panoptic heatmap.
+ bboxes = tf.concat([things_bboxes, stuff_bboxes], axis=0)
+ centers = tf.concat([things_centers, stuff_centers], axis=0)
+ classes = tf.concat([things_classes, stuff_classes], axis=0)
+ panoptic_masks = tf.concat([things_masks, stuff_masks], axis=0)
+ groundtruth_ids = tf.range(tf.shape(panoptic_masks)[0], dtype=tf.int32)
+
+ box_widths = bboxes[..., 3] - bboxes[..., 1]
+ box_heights = bboxes[..., 2] - bboxes[..., 0]
+ diag_lengths = (tf.sqrt(box_widths**2 + box_heights**2) / 2)[None]
+
+ # Compute according levels for each bounding box.
+ fpn_low_range = tf.constant(self.fpn_low_range, tf.float32)[:, None] / (
+ 2**mask_level
+ )
+ fpn_high_range = tf.constant(self.fpn_high_range, tf.float32)[:, None] / (
+ 2**mask_level
+ )
+ levels = tf.logical_and(
+ diag_lengths >= fpn_low_range, diag_lengths <= fpn_high_range
+ )
+
+ ct_heatmaps, box_indices = [], []
+ offset = 0
+ level_ids = []
+ level_bboxes, level_classes = [], []
+ for level in range(self._min_level, self._max_level + 1):
+ level_idx = level - self._min_level
+ level_ids.append(tf.where(levels[level_idx]))
+ level_bboxes.append(tf.boolean_mask(bboxes, levels[level_idx]))
+ level_classes.append(tf.boolean_mask(classes, levels[level_idx]))
+
+ level_center = tf.boolean_mask(centers, levels[level_idx])
+
+ width_ratio = 1 / float(2 ** (level - mask_level))
+ height_ratio = 1 / float(2 ** (level - mask_level))
+
+ out_height = int(self._output_size[0] / float(2 ** (level)))
+ out_width = int(self._output_size[1] / float(2 ** (level)))
+
+ # Original box coordinates
+ # [max_num_instances, ]
+ ytl, ybr = level_bboxes[-1][..., 0], level_bboxes[-1][..., 2]
+ xtl, xbr = level_bboxes[-1][..., 1], level_bboxes[-1][..., 3]
+
+ # Scaled centers (could be floating point)
+ # [max_num_instances, ]
+ scale_xct = level_center[:, 1] * width_ratio
+ scale_yct = level_center[:, 0] * height_ratio
+
+ # Floor the scaled box coordinates to be placed on heatmaps
+ # [max_num_instances, ]
+ scale_xct_floor = tf.math.floor(scale_xct)
+ scale_yct_floor = tf.math.floor(scale_yct)
+ scale_indices_floor = scale_yct_floor * out_width + scale_xct_floor
+
+ # Get the scaled box dimensions for computing the gaussian radius
+ # [max_num_instances, ]
+ box_widths = (xbr - xtl) * width_ratio
+ box_heights = (ybr - ytl) * height_ratio
+
+ ct_heatmap = target_assigner.assign_center_targets(
+ out_height=out_height,
+ out_width=out_width,
+ y_center=scale_yct,
+ x_center=scale_xct,
+ boxes_height=box_heights,
+ boxes_width=box_widths,
+ channel_onehot=tf.one_hot(
+ tf.cast(level_classes[-1], tf.int32),
+ self._num_panoptic_categories, off_value=0.),
+ gaussian_iou=self._gaussian_iou)
+ ct_heatmaps.append(tf.reshape(
+ ct_heatmap, [-1, self._num_panoptic_categories]))
+ box_indices.append(
+ tf.cast(scale_indices_floor + offset, dtype=tf.int32))
+ offset += out_width * out_height
+
+ ct_heatmaps = tf.concat(ct_heatmaps, axis=0)
+ box_indices = tf.concat(box_indices, axis=0)
+ ids = tf.concat(level_ids, axis=0)[:, 0]
+ bboxes = tf.concat(level_bboxes, axis=0)
+ classes = tf.concat(level_classes, axis=0)
+ panoptic_masks = tf.gather(panoptic_masks, ids)
+ groundtruth_ids = tf.gather(groundtruth_ids, ids)
+
+ panoptic_mask_weights = tf.cast(classes >= 0, tf.float32)
+ panoptic_mask_weights = preprocess_ops.clip_or_pad_to_fixed_size(
+ panoptic_mask_weights, self._max_num_instances, 0
+ )
+ panoptic_classes = preprocess_ops.clip_or_pad_to_fixed_size(
+ classes, self._max_num_instances, -1
+ )
+ panoptic_masks = preprocess_ops.clip_or_pad_to_fixed_size(
+ panoptic_masks, self._max_num_instances, -1.0
+ )
+ panoptic_masks = tf.transpose(panoptic_masks, [1, 2, 0])
+ panoptic_boxes = preprocess_ops.clip_or_pad_to_fixed_size(
+ bboxes, self._max_num_instances, -1
+ )
+ box_indices = preprocess_ops.clip_or_pad_to_fixed_size(
+ box_indices, self._max_num_instances, 0
+ )
+ groundtruth_ids = preprocess_ops.clip_or_pad_to_fixed_size(
+ groundtruth_ids, self._max_num_instances, -1
+ )
+
+ labels = {}
+ labels.update({
+ 'panoptic_heatmaps': ct_heatmaps,
+ 'panoptic_classes': panoptic_classes,
+ 'panoptic_masks': panoptic_masks,
+ 'panoptic_mask_weights': tf.cast(panoptic_classes >= 0, tf.float32),
+ 'panoptic_padding_mask': padding_mask,
+ 'panoptic_boxes': panoptic_boxes,
+ 'panoptic_box_indices': box_indices,
+ 'num_instances': tf.reduce_sum(
+ tf.cast(panoptic_classes >= 0, tf.float32)
+ ),
+ 'groundtruth_ids': groundtruth_ids,
+ })
+ return image, labels
+
+ def _parse_eval_data(self, data):
+ """Parses data for evaluation.
+
+ Args:
+ data: the decoded tensor dictionary from TfExampleDecoder.
+
+ Returns:
+ A dictionary of {'images': image, 'labels': labels} where
+ image: image tensor that is preproessed to have normalized value and
+ dimension [output_size[0], output_size[1], 3]
+ labels: a dictionary of tensors used for training.
+ """
+ def _process_mask(mask, ignore_label, image_info):
+ mask = tf.cast(mask, dtype=tf.float32)
+ mask = tf.reshape(mask, shape=[1, data['height'], data['width'], 1])
+ mask += 1
+
+ if self._segmentation_resize_eval_groundtruth:
+ # Resizes eval masks to match input image sizes. In that case, mean IoU
+ # is computed on output_size not the original size of the images.
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ mask = preprocess_ops.resize_and_crop_masks(
+ mask, image_scale, self._output_size, offset)
+ else:
+ mask = tf.image.pad_to_bounding_box(
+ mask, 0, 0,
+ self._segmentation_groundtruth_padded_size[0], # pyrefly: ignore[unsupported-operation]
+ self._segmentation_groundtruth_padded_size[1]) # pyrefly: ignore[unsupported-operation]
+ mask -= 1
+ # Assign ignore label to the padded region.
+ mask = tf.where(
+ tf.equal(mask, -1),
+ ignore_label * tf.ones_like(mask),
+ mask)
+ mask = tf.squeeze(mask, axis=0)
+ return mask
+
+ image = data['image']
+ # Normalizes image with mean and std pixel values.
+ image = preprocess_ops.normalize_image(image)
+
+ # Resizes and crops image.
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ self._output_size,
+ padded_size=preprocess_ops.compute_padded_size(
+ self._output_size, 2 ** self._max_level),
+ aug_scale_min=1.0,
+ aug_scale_max=1.0)
+
+ # Casts input image to self._dtype
+ image = tf.cast(image, dtype=self._dtype)
+ groundtruths = {
+ 'source_id': utils.process_source_id(data['source_id']),
+ 'height': data['height'],
+ 'width': data['width'],
+ }
+ labels = {'image_info': image_info, 'groundtruths': groundtruths}
+
+ panoptic_category_mask = _process_mask(
+ data['groundtruth_panoptic_category_mask'],
+ self._panoptic_ignore_label,
+ image_info,
+ )
+ panoptic_instance_mask = _process_mask(
+ data['groundtruth_panoptic_instance_mask'], 0, image_info
+ )
+
+ panoptic_category_mask = panoptic_category_mask[:, :, 0]
+ panoptic_instance_mask = panoptic_instance_mask[:, :, 0]
+
+ labels['groundtruths'].update({
+ 'gt_panoptic_category_mask': tf.cast(
+ panoptic_category_mask, dtype=tf.int32
+ ),
+ 'gt_panoptic_instance_mask': tf.cast(
+ panoptic_instance_mask, dtype=tf.int32
+ ),
+ })
+ return image, labels
diff --git a/official/projects/maskconver/dataloaders/panoptic_maskrcnn_input.py b/official/projects/maskconver/dataloaders/panoptic_maskrcnn_input.py
new file mode 100644
index 00000000000..b6368da0ab1
--- /dev/null
+++ b/official/projects/maskconver/dataloaders/panoptic_maskrcnn_input.py
@@ -0,0 +1,591 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data parser and processing for Panoptic Mask R-CNN."""
+from typing import Optional
+
+import tensorflow as tf, tf_keras
+from tensorflow_addons import image as tfa_image
+
+from official.projects.centernet.ops import target_assigner
+from official.vision.configs import common
+from official.vision.dataloaders import maskrcnn_input
+from official.vision.dataloaders import tf_example_decoder
+from official.vision.ops import augment
+from official.vision.ops import preprocess_ops
+
+
+class TfExampleDecoder(tf_example_decoder.TfExampleDecoder):
+ """Tensorflow Example proto decoder."""
+
+ def __init__(
+ self,
+ regenerate_source_id: bool,
+ mask_binarize_threshold: float,
+ include_panoptic_masks: bool,
+ panoptic_category_mask_key: str = 'image/panoptic/category_mask',
+ panoptic_instance_mask_key: str = 'image/panoptic/instance_mask'):
+ super(TfExampleDecoder, self).__init__(
+ include_mask=True,
+ regenerate_source_id=regenerate_source_id,
+ mask_binarize_threshold=None)
+
+ self._include_panoptic_masks = include_panoptic_masks
+ self._panoptic_category_mask_key = panoptic_category_mask_key
+ self._panoptic_instance_mask_key = panoptic_instance_mask_key
+ keys_to_features = {
+ 'image/segmentation/class/encoded':
+ tf.io.FixedLenFeature((), tf.string, default_value='')}
+
+ if include_panoptic_masks:
+ keys_to_features.update({
+ panoptic_category_mask_key:
+ tf.io.FixedLenFeature((), tf.string, default_value=''),
+ panoptic_instance_mask_key:
+ tf.io.FixedLenFeature((), tf.string, default_value='')})
+ self._segmentation_keys_to_features = keys_to_features
+
+ def decode(self, serialized_example):
+ decoded_tensors = super(TfExampleDecoder, self).decode(serialized_example)
+ parsed_tensors = tf.io.parse_single_example(
+ serialized_example, self._segmentation_keys_to_features)
+ segmentation_mask = tf.io.decode_image(
+ parsed_tensors['image/segmentation/class/encoded'],
+ channels=1)
+ segmentation_mask.set_shape([None, None, 1])
+ decoded_tensors.update({'groundtruth_segmentation_mask': segmentation_mask})
+
+ if self._include_panoptic_masks:
+ category_mask = tf.io.decode_image(
+ parsed_tensors[self._panoptic_category_mask_key],
+ channels=1)
+ instance_mask = tf.io.decode_image(
+ parsed_tensors[self._panoptic_instance_mask_key],
+ channels=1)
+ category_mask.set_shape([None, None, 1])
+ instance_mask.set_shape([None, None, 1])
+
+ decoded_tensors.update({
+ 'groundtruth_panoptic_category_mask':
+ category_mask,
+ 'groundtruth_panoptic_instance_mask':
+ instance_mask})
+ return decoded_tensors
+
+
+class Parser(maskrcnn_input.Parser):
+ """Parser to parse an image and its annotations into a dictionary of tensors."""
+
+ def __init__(self,
+ output_size,
+ min_level,
+ max_level,
+ num_scales,
+ aspect_ratios,
+ anchor_size,
+ rpn_match_threshold=0.7,
+ rpn_unmatched_threshold=0.3,
+ rpn_batch_size_per_im=256,
+ rpn_fg_fraction=0.5,
+ aug_rand_hflip=False,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ skip_crowd_during_training=True,
+ max_num_instances=100,
+ mask_crop_size=112,
+ segmentation_resize_eval_groundtruth=True,
+ segmentation_groundtruth_padded_size=None,
+ segmentation_ignore_label=255,
+ panoptic_ignore_label=0,
+ include_panoptic_masks=True,
+ level=3,
+ num_panoptic_categories=201,
+ num_thing_categories=91,
+ gaussian_iou=0.7,
+ aug_type: Optional[common.Augmentation] = None,
+ max_num_stuff_centers=1,
+ dtype='float32'):
+ """Initializes parameters for parsing annotations in the dataset.
+
+ Args:
+ output_size: `Tensor` or `list` for [height, width] of output image. The
+ output_size should be divided by the largest feature stride 2^max_level.
+ min_level: `int` number of minimum level of the output feature pyramid.
+ max_level: `int` number of maximum level of the output feature pyramid.
+ num_scales: `int` number representing intermediate scales added
+ on each level. For instance, num_scales=2 adds one additional
+ intermediate anchor scales [2^0, 2^0.5] on each level.
+ aspect_ratios: `list` of float numbers representing the aspect raito
+ anchors added on each level. The number indicates the ratio of width to
+ height. For instance, aspect_ratios=[1.0, 2.0, 0.5] adds three anchors
+ on each scale level.
+ anchor_size: `float` number representing the scale of size of the base
+ anchor to the feature stride 2^level.
+ rpn_match_threshold: `float`, match threshold for anchors in RPN.
+ rpn_unmatched_threshold: `float`, unmatched threshold for anchors in RPN.
+ rpn_batch_size_per_im: `int` for batch size per image in RPN.
+ rpn_fg_fraction: `float` for forground fraction per batch in RPN.
+ aug_rand_hflip: `bool`, if True, augment training with random
+ horizontal flip.
+ aug_scale_min: `float`, the minimum scale applied to `output_size` for
+ data augmentation during training.
+ aug_scale_max: `float`, the maximum scale applied to `output_size` for
+ data augmentation during training.
+ skip_crowd_during_training: `bool`, if True, skip annotations labeled with
+ `is_crowd` equals to 1.
+ max_num_instances: `int` number of maximum number of instances in an
+ image. The groundtruth data will be padded to `max_num_instances`.
+ mask_crop_size: the size which groundtruth mask is cropped to.
+ segmentation_resize_eval_groundtruth: `bool`, if True, eval groundtruth
+ masks are resized to output_size.
+ segmentation_groundtruth_padded_size: `Tensor` or `list` for [height,
+ width]. When resize_eval_groundtruth is set to False, the groundtruth
+ masks are padded to this size.
+ segmentation_ignore_label: `int` the pixels with ignore label will not be
+ used for training and evaluation.
+ panoptic_ignore_label: `int` the pixels with ignore label will not be used
+ by the PQ evaluator.
+ include_panoptic_masks: `bool`, if True, category_mask and instance_mask
+ will be parsed. Set this to true if PQ evaluator is enabled.
+ level: `int`, output level, used to generate target assignments.
+ num_panoptic_categories: `int`, number of panoptic categories.
+ num_thing_categories: `int1, number of thing categories.
+ gaussian_iou: `float`, used for generating the center heatmaps.
+ aug_type: An optional Augmentation object with params for AutoAugment.
+ max_num_stuff_centers: `int`, max number of stuff centers.
+ dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
+ """
+ super(Parser, self).__init__(
+ output_size=output_size,
+ min_level=min_level,
+ max_level=max_level,
+ num_scales=num_scales,
+ aspect_ratios=aspect_ratios,
+ anchor_size=anchor_size,
+ rpn_match_threshold=rpn_match_threshold,
+ rpn_unmatched_threshold=rpn_unmatched_threshold,
+ rpn_batch_size_per_im=rpn_batch_size_per_im,
+ rpn_fg_fraction=rpn_fg_fraction,
+ aug_rand_hflip=False,
+ aug_scale_min=aug_scale_min,
+ aug_scale_max=aug_scale_max,
+ skip_crowd_during_training=skip_crowd_during_training,
+ max_num_instances=max_num_instances,
+ include_mask=True,
+ mask_crop_size=mask_crop_size,
+ dtype=dtype)
+
+ self.aug_rand_hflip = aug_rand_hflip
+ self._segmentation_resize_eval_groundtruth = segmentation_resize_eval_groundtruth
+ if (not segmentation_resize_eval_groundtruth) and (
+ segmentation_groundtruth_padded_size is None):
+ raise ValueError(
+ 'segmentation_groundtruth_padded_size ([height, width]) needs to be'
+ 'specified when segmentation_resize_eval_groundtruth is False.')
+ self._segmentation_groundtruth_padded_size = segmentation_groundtruth_padded_size
+ self._segmentation_ignore_label = segmentation_ignore_label
+ self._panoptic_ignore_label = panoptic_ignore_label
+ self._include_panoptic_masks = include_panoptic_masks
+ self.level = level
+ self._num_panoptic_categories = num_panoptic_categories
+ self._num_thing_categories = num_thing_categories
+ self._gaussian_iou = gaussian_iou
+ self._max_num_stuff_centers = max_num_stuff_centers
+
+ if aug_type and aug_type.type:
+ if aug_type.type == 'autoaug':
+ self._augmenter = augment.AutoAugment(
+ augmentation_name=aug_type.autoaug.augmentation_name,
+ cutout_const=aug_type.autoaug.cutout_const,
+ translate_const=aug_type.autoaug.translate_const)
+ else:
+ raise ValueError('Augmentation policy {} not supported.'.format(
+ aug_type.type))
+ else:
+ self._augmenter = None
+
+ def _parse_train_data(self, data):
+ """Parses data for training.
+
+ Args:
+ data: the decoded tensor dictionary from TfExampleDecoder.
+
+ Returns:
+ image: image tensor that is preproessed to have normalized value and
+ dimension [output_size[0], output_size[1], 3]
+ labels: a dictionary of tensors used for training. The following describes
+ {key: value} pairs in the dictionary.
+ image_info: a 2D `Tensor` that encodes the information of the image and
+ the applied preprocessing. It is in the format of
+ [[original_height, original_width], [scaled_height, scaled_width]],
+ anchor_boxes: ordered dictionary with keys
+ [min_level, min_level+1, ..., max_level]. The values are tensor with
+ shape [height_l, width_l, 4] representing anchor boxes at each level.
+ rpn_score_targets: ordered dictionary with keys
+ [min_level, min_level+1, ..., max_level]. The values are tensor with
+ shape [height_l, width_l, anchors_per_location]. The height_l and
+ width_l represent the dimension of class logits at l-th level.
+ rpn_box_targets: ordered dictionary with keys
+ [min_level, min_level+1, ..., max_level]. The values are tensor with
+ shape [height_l, width_l, anchors_per_location * 4]. The height_l and
+ width_l represent the dimension of bounding box regression output at
+ l-th level.
+ gt_boxes: Groundtruth bounding box annotations. The box is represented
+ in [y1, x1, y2, x2] format. The coordinates are w.r.t the scaled
+ image that is fed to the network. The tennsor is padded with -1 to
+ the fixed dimension [self._max_num_instances, 4].
+ gt_classes: Groundtruth classes annotations. The tennsor is padded
+ with -1 to the fixed dimension [self._max_num_instances].
+ gt_masks: Groundtruth masks cropped by the bounding box and
+ resized to a fixed size determined by mask_crop_size.
+ gt_segmentation_mask: Groundtruth mask for segmentation head, this is
+ resized to a fixed size determined by output_size.
+ gt_segmentation_valid_mask: Binary mask that marks the pixels that
+ are supposed to be used in computing the segmentation loss while
+ training.
+ """
+ segmentation_mask = data['groundtruth_segmentation_mask']
+
+ # Flips image randomly during training.
+ if self.aug_rand_hflip:
+ masks = data['groundtruth_instance_masks']
+ instance_mask = data['groundtruth_panoptic_instance_mask']
+ category_mask = data['groundtruth_panoptic_category_mask']
+ image_mask = tf.concat(
+ [data['image'], segmentation_mask, category_mask, instance_mask],
+ axis=2)
+
+ image_mask, boxes, masks = preprocess_ops.random_horizontal_flip(
+ image_mask, data['groundtruth_boxes'], masks)
+
+ segmentation_mask = image_mask[:, :, -3:-2]
+ category_mask = image_mask[:, :, -2:-1]
+ instance_mask = image_mask[:, :, -1:]
+ image = image_mask[:, :, :-3]
+
+ data['image'] = image
+ if self._augmenter is not None:
+ data['image'] = self._augmenter.distort(image)
+
+ data['groundtruth_boxes'] = boxes
+ data['groundtruth_instance_masks'] = masks
+ data['groundtruth_panoptic_instance_mask'] = instance_mask
+ data['groundtruth_panoptic_category_mask'] = category_mask
+
+ image, labels = super(Parser, self)._parse_train_data(data)
+
+ image_info = labels['image_info']
+
+ def _process_mask(mask, ignore_label, image_info):
+ mask = tf.cast(mask, dtype=tf.float32)
+ mask = tf.reshape(mask, shape=[1, data['height'], data['width'], 1])
+ mask += 1
+
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ mask = preprocess_ops.resize_and_crop_masks(
+ mask, image_scale, self._output_size, offset)
+
+ mask -= 1
+ # Assign ignore label to the padded region.
+ mask = tf.where(
+ tf.equal(mask, -1),
+ ignore_label * tf.ones_like(mask),
+ mask)
+ mask = tf.squeeze(mask, axis=0)
+ return mask
+
+ if self._include_panoptic_masks:
+ panoptic_category_mask = _process_mask(
+ data['groundtruth_panoptic_category_mask'],
+ -1, image_info)
+ panoptic_instance_mask = _process_mask(
+ data['groundtruth_panoptic_instance_mask'],
+ self._panoptic_ignore_label, image_info)
+
+ panoptic_category_mask = panoptic_category_mask[:, :, 0]
+ panoptic_instance_mask = panoptic_instance_mask[:, :, 0]
+
+ padding_mask = tf.cast(panoptic_category_mask == -1, tf.float32)
+ labels.update({
+ 'gt_panoptic_category_mask':
+ tf.cast(panoptic_category_mask, dtype=tf.int32),
+ 'gt_panoptic_instance_mask':
+ tf.cast(panoptic_instance_mask, dtype=tf.int32)})
+
+ def get_bbox(mask):
+ rows = tf.math.count_nonzero(mask, axis=0, keepdims=None, dtype=tf.bool)
+ columns = tf.math.count_nonzero(
+ mask, axis=1, keepdims=None, dtype=tf.bool)
+ indices = tf.where(tf.equal(columns, True))
+ y_min = indices[0]
+ y_max = indices[-1]
+ indices = tf.where(tf.equal(rows, True))
+ x_min = indices[0]
+ x_max = indices[-1]
+ return tf.stack([y_min, x_min, y_max, x_max])
+
+ def get_thing_class_label(instance_id):
+ mask = tf.cast(labels['gt_panoptic_instance_mask'] == instance_id,
+ tf.int32)
+ bbox = tf.cast(tf.reshape(get_bbox(mask), [-1]), tf.float32)
+ inds = tf.where(labels['gt_panoptic_instance_mask'] == instance_id)
+ thing_classes = tf.gather_nd(labels['gt_panoptic_category_mask'], inds)
+ thing_classes, _, counts = tf.unique_with_counts(thing_classes)
+ max_index = tf.argmax(counts)
+ thing_class = tf.cast(thing_classes[max_index], tf.int64)
+ return tf.cast(mask, tf.float32), bbox, thing_class
+
+ instance_ids, _ = tf.unique(
+ tf.reshape(labels['gt_panoptic_instance_mask'], [-1]))
+ instance_ids = tf.boolean_mask(instance_ids, instance_ids > 0)
+ things_masks, things_bboxes, things_classes = tf.map_fn(
+ get_thing_class_label,
+ instance_ids,
+ fn_output_signature=(tf.TensorSpec(self._output_size), tf.TensorSpec([
+ 4,
+ ]), tf.int64))
+
+ stuff_classes, _ = tf.unique(
+ tf.reshape(labels['gt_panoptic_category_mask'], [-1]))
+ stuff_classes = tf.boolean_mask(
+ stuff_classes, stuff_classes > self._num_thing_categories - 1)
+
+ # Compute classes and bboxes for stuff classes
+ def get_stuff_class_label(stuff_class, k=self._max_num_stuff_centers):
+ mask = tf.cast(
+ labels['gt_panoptic_category_mask'] == stuff_class, tf.int32)
+ smoothed_mask = tf.nn.max_pool(mask[:, :, None], 21, 1, padding='SAME')
+ smoothed_mask = tf.nn.max_pool(smoothed_mask, 21, 1, padding='SAME')
+ smoothed_mask = tf.nn.max_pool(smoothed_mask, 21, 1, padding='SAME')
+
+ smoothed_mask = tf.cast(tf.squeeze(smoothed_mask, axis=2), tf.int32)
+ connected_components = tfa_image.connected_components( # pyrefly: ignore[not-callable]
+ images=smoothed_mask)
+ counts = tf.math.bincount(connected_components)
+ ids = tf.argsort(counts[1:], axis=-1, direction='DESCENDING') + 1
+
+ masks = tf.cast(tf.repeat(mask[None, :, :], k, axis=0), tf.float32)
+ stuff_classes = tf.cast(stuff_class, tf.int64) * tf.ones([k], tf.int64)
+
+ # Pad or clip ids
+ ids_length = tf.shape(ids)[0]
+ ids_length = tf.clip_by_value(ids_length, 0, k)
+ # batch_weights = 1 / tf.cast(ids_length, tf.float32) * tf.ones([1, k])
+ batch_weights = 1.0 * tf.ones([1, k], tf.float32)
+ ids = ids[:ids_length]
+ padding_length = tf.maximum(0, k - ids_length)
+ paddings = tf.cast(-1 * tf.ones([padding_length]), ids.dtype)
+ ids = tf.concat([ids, paddings], axis=0)
+
+ def get_bbox_and_batch_mask(island_id):
+ if island_id == -1:
+ bbox = tf.zeros([4], tf.float32)
+ batch_mask = -1 * tf.ones([1], tf.int32)
+ smask = tf.cast(tf.zeros_like(connected_components), tf.float32)
+ else:
+ smask = tf.cast(connected_components, tf.int32) == island_id
+ smask = tf.cast(
+ tf.logical_and(tf.cast(smask, tf.bool), tf.cast(mask, tf.bool)),
+ tf.float32)
+ if tf.reduce_sum(smask) < 1024:
+ batch_mask = -1 * tf.ones([1], tf.int32)
+ else:
+ batch_mask = tf.ones([1], tf.int32)
+ bbox = tf.cast(tf.reshape(get_bbox(smask), [-1]), tf.float32)
+ return smask, bbox, batch_mask
+
+ stuff_center_masks, bboxes, batch_mask = tf.map_fn(
+ get_bbox_and_batch_mask,
+ ids,
+ fn_output_signature=((tf.TensorSpec(self._output_size, tf.float32),
+ tf.TensorSpec([4,], tf.float32),
+ tf.TensorSpec([1], tf.int32))))
+
+ stuff_classes = tf.reshape(
+ stuff_classes, [1, k])
+ batch_mask = tf.reshape(batch_mask, [1, k])
+ stuff_center_masks = tf.reshape(
+ stuff_center_masks, [1, k] + self._output_size)
+ return masks[None, :, :, :], bboxes[
+ None, :, :], stuff_classes, batch_mask, batch_weights
+
+ centers = self._max_num_stuff_centers
+ stuff_masks, stuff_bboxes, stuff_classes, stuff_batch_mask, stuff_batch_weigths = tf.map_fn(
+ get_stuff_class_label,
+ stuff_classes,
+ fn_output_signature=(tf.TensorSpec([1, centers] + self._output_size),
+ tf.TensorSpec([1, centers, 4,]),
+ tf.TensorSpec([1, centers], tf.int64),
+ tf.TensorSpec([1, centers], tf.int32),
+ tf.TensorSpec([1, centers], tf.float32)))
+ stuff_masks = tf.reshape(stuff_masks, [-1] + self._output_size)
+ stuff_bboxes = tf.reshape(stuff_bboxes, [-1, 4])
+ stuff_classes = tf.reshape(stuff_classes, [-1])
+ stuff_batch_mask = tf.reshape(stuff_batch_mask, [-1])
+ stuff_batch_weigths = tf.reshape(stuff_batch_weigths, [-1])
+ stuff_masks = stuff_masks[stuff_batch_mask == 1]
+ stuff_bboxes = stuff_bboxes[stuff_batch_mask == 1]
+ stuff_classes = stuff_classes[stuff_batch_mask == 1]
+
+ # Start to generate panoptic heatmap.
+ bboxes = tf.concat([things_bboxes, stuff_bboxes], axis=0)
+ classes = tf.concat([things_classes, stuff_classes], axis=0)
+ panoptic_masks = tf.concat([things_masks, stuff_masks], axis=0)
+
+ width_ratio = 1 / float(2 ** self.level)
+ height_ratio = 1 / float(2 ** self.level)
+
+ # Original box coordinates
+ # [max_num_instances, ]
+ ytl, ybr = bboxes[..., 0], bboxes[..., 2]
+ xtl, xbr = bboxes[..., 1], bboxes[..., 3]
+ yct = (ytl + ybr) / 2
+ xct = (xtl + xbr) / 2
+
+ # Scaled box coordinates (could be floating point)
+ # [max_num_instances, ]
+ scale_xct = xct * width_ratio
+ scale_yct = yct * height_ratio
+
+ # Floor the scaled box coordinates to be placed on heatmaps
+ # [max_num_instances, ]
+ scale_xct_floor = tf.math.floor(scale_xct)
+ scale_yct_floor = tf.math.floor(scale_yct)
+
+ # Get the scaled box dimensions for computing the gaussian radius
+ # [max_num_instances, ]
+ box_widths = bboxes[..., 3] - bboxes[..., 1]
+ box_heights = bboxes[..., 2] - bboxes[..., 0]
+
+ box_widths = box_widths * width_ratio
+ box_heights = box_heights * height_ratio
+
+ ct_heatmap = target_assigner.assign_center_targets(
+ out_height=int(self._output_size[0] * height_ratio),
+ out_width=int(self._output_size[1] * width_ratio),
+ y_center=scale_yct,
+ x_center=scale_xct,
+ boxes_height=box_heights,
+ boxes_width=box_widths,
+ channel_onehot=tf.one_hot(
+ tf.cast(classes, tf.int32),
+ self._num_panoptic_categories, off_value=0.),
+ gaussian_iou=self._gaussian_iou)
+ box_indices = tf.cast(
+ tf.stack([scale_yct_floor, scale_xct_floor], axis=-1), dtype=tf.int32)
+
+ panoptic_classes = preprocess_ops.clip_or_pad_to_fixed_size(
+ classes, self._max_num_instances, -1)
+ panoptic_masks = preprocess_ops.clip_or_pad_to_fixed_size(
+ panoptic_masks, self._max_num_instances, -1.0)
+ panoptic_masks = tf.transpose(panoptic_masks, [1, 2, 0])
+ panoptic_boxes = preprocess_ops.clip_or_pad_to_fixed_size(
+ bboxes, self._max_num_instances, -1)
+
+ box_indices = preprocess_ops.clip_or_pad_to_fixed_size(
+ box_indices, self._max_num_instances, 0)
+
+ labels = {}
+ labels.update({
+ 'panoptic_heatmaps': ct_heatmap,
+ 'panoptic_classes': panoptic_classes,
+ 'panoptic_masks': panoptic_masks,
+ 'panoptic_mask_weights': tf.cast(panoptic_classes >= 0, tf.float32),
+ 'panoptic_padding_mask': padding_mask,
+ 'panoptic_boxes': panoptic_boxes,
+ 'panoptic_box_indices': box_indices,
+ 'num_instances': tf.reduce_sum(
+ tf.cast(panoptic_classes >= 0, tf.float32))
+ })
+ return image, labels
+
+ def _parse_eval_data(self, data):
+ """Parses data for evaluation.
+
+ Args:
+ data: the decoded tensor dictionary from TfExampleDecoder.
+
+ Returns:
+ A dictionary of {'images': image, 'labels': labels} where
+ image: image tensor that is preproessed to have normalized value and
+ dimension [output_size[0], output_size[1], 3]
+ labels: a dictionary of tensors used for training. The following
+ describes {key: value} pairs in the dictionary.
+ source_ids: Source image id. Default value -1 if the source id is
+ empty in the groundtruth annotation.
+ image_info: a 2D `Tensor` that encodes the information of the image
+ and the applied preprocessing. It is in the format of
+ [[original_height, original_width], [scaled_height, scaled_width]],
+ anchor_boxes: ordered dictionary with keys
+ [min_level, min_level+1, ..., max_level]. The values are tensor with
+ shape [height_l, width_l, 4] representing anchor boxes at each
+ level.
+ """
+ def _process_mask(mask, ignore_label, image_info):
+ mask = tf.cast(mask, dtype=tf.float32)
+ mask = tf.reshape(mask, shape=[1, data['height'], data['width'], 1])
+ mask += 1
+
+ if self._segmentation_resize_eval_groundtruth:
+ # Resizes eval masks to match input image sizes. In that case, mean IoU
+ # is computed on output_size not the original size of the images.
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ mask = preprocess_ops.resize_and_crop_masks(
+ mask, image_scale, self._output_size, offset)
+ else:
+ mask = tf.image.pad_to_bounding_box(
+ mask, 0, 0,
+ self._segmentation_groundtruth_padded_size[0], # pyrefly: ignore[unsupported-operation]
+ self._segmentation_groundtruth_padded_size[1]) # pyrefly: ignore[unsupported-operation]
+ mask -= 1
+ # Assign ignore label to the padded region.
+ mask = tf.where(
+ tf.equal(mask, -1),
+ ignore_label * tf.ones_like(mask),
+ mask)
+ mask = tf.squeeze(mask, axis=0)
+ return mask
+
+ image, labels = super(Parser, self)._parse_eval_data(data)
+ image_info = labels['image_info']
+
+ segmentation_mask = _process_mask(
+ data['groundtruth_segmentation_mask'],
+ self._segmentation_ignore_label, image_info)
+ segmentation_valid_mask = tf.not_equal(
+ segmentation_mask, self._segmentation_ignore_label)
+ labels['groundtruths'].update({
+ 'gt_segmentation_mask': segmentation_mask,
+ 'gt_segmentation_valid_mask': segmentation_valid_mask})
+
+ if self._include_panoptic_masks:
+ panoptic_category_mask = _process_mask(
+ data['groundtruth_panoptic_category_mask'],
+ self._panoptic_ignore_label, image_info)
+ panoptic_instance_mask = _process_mask(
+ data['groundtruth_panoptic_instance_mask'],
+ 0, image_info)
+
+ panoptic_category_mask = panoptic_category_mask[:, :, 0]
+ panoptic_instance_mask = panoptic_instance_mask[:, :, 0]
+
+ labels['groundtruths'].update({
+ 'gt_panoptic_category_mask':
+ tf.cast(panoptic_category_mask, dtype=tf.int32),
+ 'gt_panoptic_instance_mask':
+ tf.cast(panoptic_instance_mask, dtype=tf.int32)})
+
+ return image, labels
diff --git a/official/projects/maskconver/docs/Table1.png b/official/projects/maskconver/docs/Table1.png
new file mode 100644
index 00000000000..c0d6f065adb
Binary files /dev/null and b/official/projects/maskconver/docs/Table1.png differ
diff --git a/official/projects/maskconver/docs/maskconver_architecture.pdf b/official/projects/maskconver/docs/maskconver_architecture.pdf
new file mode 100644
index 00000000000..f939ad5ced1
Binary files /dev/null and b/official/projects/maskconver/docs/maskconver_architecture.pdf differ
diff --git a/official/projects/maskconver/docs/maskconver_architecture.png b/official/projects/maskconver/docs/maskconver_architecture.png
new file mode 100644
index 00000000000..7b5f982d1af
Binary files /dev/null and b/official/projects/maskconver/docs/maskconver_architecture.png differ
diff --git a/official/projects/maskconver/docs/maskconver_plot.png b/official/projects/maskconver/docs/maskconver_plot.png
new file mode 100644
index 00000000000..7e000000365
Binary files /dev/null and b/official/projects/maskconver/docs/maskconver_plot.png differ
diff --git a/official/projects/maskconver/docs/multiscale_maskconver.png b/official/projects/maskconver/docs/multiscale_maskconver.png
new file mode 100644
index 00000000000..4539db11004
Binary files /dev/null and b/official/projects/maskconver/docs/multiscale_maskconver.png differ
diff --git a/official/projects/maskconver/losses/maskconver_losses.py b/official/projects/maskconver/losses/maskconver_losses.py
new file mode 100644
index 00000000000..160db62c5bc
--- /dev/null
+++ b/official/projects/maskconver/losses/maskconver_losses.py
@@ -0,0 +1,174 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Losses for maskconver model."""
+
+import functools
+import tensorflow as tf, tf_keras
+
+LARGE_NUM = 1e9
+
+
+class PenaltyReducedLogisticFocalLoss(object):
+ """Penalty-reduced pixelwise logistic regression with focal loss."""
+
+ def __init__(self, alpha=2.0, beta=4.0, sigmoid_clip_value=1e-4):
+ """Constructor.
+
+ The loss is defined in Equation (1) of the Objects as Points[1] paper.
+ Although the loss is defined per-pixel in the output space, this class
+ assumes that each pixel is an anchor to be compatible with the base class.
+
+ [1]: https://arxiv.org/abs/1904.07850
+
+ Args:
+ alpha: Focussing parameter of the focal loss. Increasing this will
+ decrease the loss contribution of the well classified examples.
+ beta: The local penalty reduction factor. Increasing this will decrease
+ the contribution of loss due to negative pixels near the keypoint.
+ sigmoid_clip_value: The sigmoid operation used internally will be clipped
+ between [sigmoid_clip_value, 1 - sigmoid_clip_value)
+ """
+ self._alpha = alpha
+ self._beta = beta
+ self._sigmoid_clip_value = sigmoid_clip_value
+ super().__init__()
+
+ def __call__(self, prediction_tensor, target_tensor, weights=1.0):
+ """Compute loss function.
+
+ In all input tensors, `num_anchors` is the total number of pixels in the
+ the output space.
+
+ Args:
+ prediction_tensor: A float tensor of shape [batch_size, num_anchors,
+ num_classes] representing the predicted unscaled logits for each class.
+ The function will compute sigmoid on this tensor internally.
+ target_tensor: A float tensor of shape [batch_size, num_anchors,
+ num_classes] representing a tensor with the 'splatted' keypoints,
+ possibly using a gaussian kernel. This function assumes that
+ the target is bounded between [0, 1].
+ weights: a float tensor of shape, either [batch_size, num_anchors,
+ num_classes] or [batch_size, num_anchors, 1]. If the shape is
+ [batch_size, num_anchors, 1], all the classses are equally weighted.
+
+ Returns:
+ loss: a float tensor of shape [batch_size, num_anchors, num_classes]
+ representing the value of the loss function.
+ """
+ with tf.name_scope('prlf_loss'):
+ ignore = tf.cast(tf.math.equal(target_tensor, -1.0), tf.float32)
+ is_present_tensor = tf.math.equal(target_tensor, 1.0)
+ prediction_tensor = tf.clip_by_value(
+ tf.sigmoid(prediction_tensor),
+ self._sigmoid_clip_value, 1 - self._sigmoid_clip_value)
+
+ positive_loss = (tf.math.pow((1 - prediction_tensor), self._alpha) *
+ tf.math.log(prediction_tensor))
+ negative_loss = (tf.math.pow((1 - target_tensor), self._beta) *
+ tf.math.pow(prediction_tensor, self._alpha) *
+ tf.math.log(1 - prediction_tensor))
+
+ loss = -tf.where(is_present_tensor, positive_loss, negative_loss)
+ return loss * weights * (1 - ignore)
+
+
+class EmbedLoss:
+ """Embedding loss class."""
+
+ def __init__(
+ self, projection_norm: bool = True,
+ temperature: float = 0.1,
+ max_num_detections: int = 100,
+ num_classes: int = 91,
+ class_agnostic: bool = False):
+ """Initializes `ContrastiveLoss`."""
+ self._projection_norm = projection_norm
+ self._temperature = temperature
+ self._max_num_detections = max_num_detections
+ self._num_classes = num_classes
+ self._class_agnostic = class_agnostic
+
+ def __call__(self, embed_outputs, matched_gt_indices, matched_gt_classes):
+ """Computes the embedding loss."""
+ with tf.name_scope('feature_nms_loss'):
+ # Get (normalized) hidden1 and hidden2.
+ if self._projection_norm:
+ embed_outputs = tf.math.l2_normalize(embed_outputs, -1)
+ batch_size = tf.shape(embed_outputs)[0]
+ num_boxes = tf.shape(embed_outputs)[1]
+ matched_gt_indices = tf.cast(matched_gt_indices, tf.int32)
+ if not self._class_agnostic:
+ matched_gt_classes = tf.cast(matched_gt_classes, tf.int32)
+
+ # labels = [batch_size, num_boxes, max_num_detections]
+ labels = tf.one_hot(matched_gt_indices, self._max_num_detections)
+ ta_matmul = functools.partial(tf.matmul, transpose_a=True)
+ # num_boxes_per_index = [batch_size, max_num_detections]
+ num_boxes_per_index = tf.reduce_sum(labels, axis=1)
+
+ # mean_emb = [batch_size, max_num_detections, embed_dim]
+ mean_emb = tf.stop_gradient(tf.math.divide_no_nan(
+ ta_matmul(labels, embed_outputs),
+ tf.expand_dims(num_boxes_per_index, -1)))
+
+ if self._projection_norm:
+ mean_emb = tf.math.l2_normalize(mean_emb, -1)
+
+ tb_matmul = functools.partial(tf.matmul, transpose_b=True)
+ # logits = [batch_size, num_boxes, max_num_detections]
+ logits = tb_matmul(embed_outputs, mean_emb) / self._temperature
+
+ mask = tf.ones([batch_size, num_boxes, self._max_num_detections],
+ tf.int32)
+ if not self._class_agnostic:
+ # Force value "-1"(negative sample)->"0" to be used for indices.
+ # These wrong indices will be ignored since there is no proposals with
+ # gt_classes = 0 (background)
+ classes = tf.one_hot(matched_gt_classes, self._num_classes)
+ matched_gt_classes = tf.nn.relu(matched_gt_classes)
+ matched_gt_indices = tf.nn.relu(matched_gt_indices)
+ batch_indices = tf.ones([batch_size, num_boxes],
+ tf.int32) * tf.expand_dims(
+ range(batch_size), -1)
+ indices = tf.stack(
+ [batch_indices, matched_gt_classes, matched_gt_indices], -1)
+ # gt_class_indices = [batch_size, num_classes, max_num_detections]
+ gt_class_indices = tf.scatter_nd(
+ indices, tf.ones([batch_size, num_boxes]),
+ [batch_size, self._num_classes, self._max_num_detections])
+ # A class which has multiple objects in an image is updated more than
+ # once.
+ gt_class_indices = tf.clip_by_value(
+ gt_class_indices, clip_value_min=0, clip_value_max=1)
+ # mask = [batch_size, num_boxes, max_num_detections]
+ mask = tf.cast(tf.matmul(classes, gt_class_indices), tf.int32)
+
+ valid_gt_indices = tf.math.count_nonzero(
+ labels, axis=1, dtype=tf.int32)
+ mask *= tf.expand_dims(valid_gt_indices, 1)
+
+ mask = tf.where(tf.equal(mask, 0), LARGE_NUM, 0.)
+
+ loss_mask = tf.math.count_nonzero(labels, axis=2)
+ loss_mask = tf.where(tf.equal(loss_mask, 0), 0., 1.)
+
+ logits = logits - mask
+
+ loss = tf.nn.softmax_cross_entropy_with_logits(
+ labels,
+ logits)
+ loss = tf.reduce_sum(loss * loss_mask) / (tf.reduce_sum(loss_mask) + 1.0)
+
+ return loss
diff --git a/official/projects/maskconver/modeling/factory.py b/official/projects/maskconver/modeling/factory.py
new file mode 100644
index 00000000000..5c66460915b
--- /dev/null
+++ b/official/projects/maskconver/modeling/factory.py
@@ -0,0 +1,316 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Factory method to build panoptic segmentation model."""
+from typing import Optional, Union
+
+import tensorflow as tf, tf_keras
+
+from official.projects.maskconver.configs import maskconver as maskconver_cfg
+from official.projects.maskconver.configs import multiscale_maskconver as multiscale_maskconver_cfg
+from official.projects.maskconver.modeling import maskconver_model
+from official.projects.maskconver.modeling import multiscale_maskconver_model
+from official.projects.maskconver.modeling.layers import maskconver_head
+from official.projects.maskconver.modeling.layers import multiscale_maskconver_head
+from official.projects.maskconver.modeling.layers import panoptic_segmentation_generator
+from official.vision.modeling import backbones
+from official.vision.modeling import decoders
+
+
+def build_maskconver_model(
+ input_specs: tf_keras.layers.InputSpec,
+ model_config: maskconver_cfg.MaskConver,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ backbone: Optional[tf_keras.Model] = None,
+ decoder: Optional[Union[tf_keras.Model, tf_keras.layers.Layer]] = None,
+ segmentation_inference: bool = False,
+) -> tf_keras.Model:
+ """Builds Segmentation model."""
+ norm_activation_config = model_config.norm_activation
+ if not backbone:
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=model_config.backbone,
+ norm_activation_config=norm_activation_config,
+ l2_regularizer=l2_regularizer)
+
+ if not decoder:
+ decoder = decoders.factory.build_decoder(
+ input_specs=backbone.output_specs,
+ model_config=model_config,
+ l2_regularizer=l2_regularizer)
+
+ class_head_config = model_config.class_head
+ mask_embedding_head_config = model_config.mask_embedding_head
+ per_pixel_embedding_head_config = model_config.per_pixel_embedding_head
+
+ # pylint: disable=line-too-long
+ class_head = maskconver_head.MaskConverHead(
+ num_classes=model_config.num_classes,
+ level=class_head_config.level,
+ num_convs=class_head_config.num_convs,
+ num_filters=class_head_config.num_filters,
+ use_layer_norm=class_head_config.use_layer_norm,
+ depthwise_kernel_size=class_head_config.depthwise_kernel_size,
+ use_depthwise_convolution=class_head_config.use_depthwise_convolution,
+ prediction_kernel_size=class_head_config.prediction_kernel_size,
+ upsample_factor=class_head_config.upsample_factor,
+ feature_fusion=class_head_config.feature_fusion,
+ decoder_min_level=class_head_config.decoder_min_level,
+ decoder_max_level=class_head_config.decoder_max_level,
+ low_level=class_head_config.low_level,
+ low_level_num_filters=class_head_config.low_level_num_filters,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer,
+ bias_initializer=tf.constant_initializer(-2.19))
+
+ mask_embedding_head = maskconver_head.MaskConverHead(
+ num_classes=model_config.embedding_size,
+ level=mask_embedding_head_config.level,
+ num_convs=mask_embedding_head_config.num_convs,
+ num_filters=mask_embedding_head_config.num_filters,
+ use_depthwise_convolution=mask_embedding_head_config.use_depthwise_convolution,
+ use_layer_norm=mask_embedding_head_config.use_layer_norm,
+ depthwise_kernel_size=mask_embedding_head_config.depthwise_kernel_size,
+ prediction_kernel_size=mask_embedding_head_config.prediction_kernel_size,
+ upsample_factor=mask_embedding_head_config.upsample_factor,
+ feature_fusion=mask_embedding_head_config.feature_fusion,
+ decoder_min_level=mask_embedding_head_config.decoder_min_level,
+ decoder_max_level=mask_embedding_head_config.decoder_max_level,
+ low_level=mask_embedding_head_config.low_level,
+ low_level_num_filters=mask_embedding_head_config.low_level_num_filters,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer,
+ bias_initializer=tf.constant_initializer(0.0))
+
+ per_pixel_embedding_head = maskconver_head.MaskConverHead(
+ num_classes=model_config.embedding_size,
+ level=per_pixel_embedding_head_config.level,
+ num_convs=per_pixel_embedding_head_config.num_convs,
+ num_filters=per_pixel_embedding_head_config.num_filters,
+ use_depthwise_convolution=per_pixel_embedding_head_config.use_depthwise_convolution,
+ depthwise_kernel_size=per_pixel_embedding_head_config.depthwise_kernel_size,
+ use_layer_norm=per_pixel_embedding_head_config.use_layer_norm,
+ prediction_kernel_size=per_pixel_embedding_head_config.prediction_kernel_size,
+ upsample_factor=per_pixel_embedding_head_config.upsample_factor,
+ feature_fusion=per_pixel_embedding_head_config.feature_fusion,
+ decoder_min_level=per_pixel_embedding_head_config.decoder_min_level,
+ decoder_max_level=per_pixel_embedding_head_config.decoder_max_level,
+ low_level=per_pixel_embedding_head_config.low_level,
+ low_level_num_filters=per_pixel_embedding_head_config.low_level_num_filters,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer,
+ bias_initializer=tf.constant_initializer(0.0))
+
+ proposal_generator = panoptic_segmentation_generator.MaskConverProposalGenerator(
+ max_proposals=model_config.num_instances,
+ peak_error=1e-6,
+ peak_extract_kernel_size=3)
+
+ postprocessing_config = model_config.panoptic_generator
+ if segmentation_inference:
+ is_thing = [False] * model_config.num_classes
+ else:
+ is_thing = ([False] + [True] * (model_config.num_thing_classes - 1) +
+ [False] *
+ (model_config.num_classes - model_config.num_thing_classes))
+ panoptic_generator = None
+ if postprocessing_config:
+ panoptic_generator = panoptic_segmentation_generator.MaskConverPanopticGenerator(
+ output_size=model_config.padded_output_size,
+ num_classes=model_config.num_classes,
+ is_thing=is_thing,
+ num_instances=model_config.num_instances,
+ object_mask_threshold=postprocessing_config.object_mask_threshold,
+ small_area_threshold=postprocessing_config.small_area_threshold,
+ overlap_threshold=postprocessing_config.overlap_threshold,
+ rescale_predictions=postprocessing_config.rescale_predictions,
+ use_hardware_optimization=postprocessing_config.use_hardware_optimization,
+ )
+ # pylint: enable=line-too-long
+ mlp_embedding_head = maskconver_head.MLP(
+ hidden_dim=model_config.embedding_size,
+ output_dim=model_config.embedding_size,
+ num_layers=2,
+ activation=norm_activation_config.activation,
+ l2_regularizer=l2_regularizer)
+
+ model = maskconver_model.MaskConverModel(
+ backbone,
+ decoder,
+ embedding_head=mask_embedding_head,
+ class_head=class_head,
+ per_pixel_embeddings_head=per_pixel_embedding_head,
+ mlp_embedding_head=mlp_embedding_head,
+ proposal_generator=proposal_generator,
+ panoptic_generator=panoptic_generator,
+ level=model_config.level,
+ padded_output_size=model_config.padded_output_size,
+ l2_regularizer=l2_regularizer,
+ embedding_size=model_config.embedding_size,
+ num_classes=model_config.num_classes)
+ return model
+
+
+def build_multiscale_maskconver_model(
+ input_specs: tf_keras.layers.InputSpec,
+ model_config: multiscale_maskconver_cfg.MultiScaleMaskConver,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ backbone: Optional[tf_keras.regularizers.Regularizer] = None,
+ decoder: Optional[tf_keras.regularizers.Regularizer] = None,
+ segmentation_inference: bool = False,
+) -> tf_keras.Model:
+ """Builds multiscale MaskConver model."""
+ norm_activation_config = model_config.norm_activation
+ if not backbone:
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=model_config.backbone,
+ norm_activation_config=norm_activation_config,
+ l2_regularizer=l2_regularizer)
+
+ if not decoder:
+ decoder = decoders.factory.build_decoder(
+ input_specs=backbone.output_specs,
+ model_config=model_config,
+ l2_regularizer=l2_regularizer)
+
+ mask_decoder = None
+ if model_config.mask_decoder:
+ temp_model_config = multiscale_maskconver_cfg.MultiScaleMaskConver()
+ temp_model_config.override(model_config)
+ temp_model_config.decoder.override(model_config.mask_decoder)
+ mask_decoder = decoders.factory.build_decoder(
+ input_specs=backbone.output_specs,
+ model_config=temp_model_config,
+ l2_regularizer=l2_regularizer)
+
+ class_head_config = model_config.class_head
+ mask_embedding_head_config = model_config.mask_embedding_head
+ per_pixel_embedding_head_config = model_config.per_pixel_embedding_head
+
+ # pylint: disable=line-too-long
+ class_head = multiscale_maskconver_head.MultiScaleMaskConverHead(
+ num_classes=model_config.num_classes,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ num_convs=class_head_config.num_convs,
+ depthwise_kernel_size=class_head_config.depthwise_kernel_size,
+ use_layer_norm=class_head_config.use_layer_norm,
+ num_filters=class_head_config.num_filters,
+ use_depthwise_convolution=class_head_config.use_depthwise_convolution,
+ prediction_kernel_size=class_head_config.prediction_kernel_size,
+ upsample_factor=class_head_config.upsample_factor,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer,
+ bias_initializer=tf.constant_initializer(-2.19),)
+
+ mask_embedding_head = multiscale_maskconver_head.MultiScaleMaskConverHead(
+ num_classes=model_config.embedding_size,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ num_convs=mask_embedding_head_config.num_convs,
+ num_filters=mask_embedding_head_config.num_filters,
+ use_depthwise_convolution=mask_embedding_head_config
+ .use_depthwise_convolution,
+ use_layer_norm=mask_embedding_head_config.use_layer_norm,
+ depthwise_kernel_size=mask_embedding_head_config.depthwise_kernel_size,
+ prediction_kernel_size=mask_embedding_head_config.prediction_kernel_size,
+ upsample_factor=mask_embedding_head_config.upsample_factor,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer,
+ bias_initializer=tf.constant_initializer(0.0))
+
+ per_pixel_embedding_head = maskconver_head.MaskConverHead(
+ num_classes=model_config.embedding_size,
+ level=per_pixel_embedding_head_config.level,
+ num_convs=per_pixel_embedding_head_config.num_convs,
+ num_filters=per_pixel_embedding_head_config.num_filters,
+ use_layer_norm=per_pixel_embedding_head_config.use_layer_norm,
+ use_depthwise_convolution=per_pixel_embedding_head_config.use_depthwise_convolution,
+ depthwise_kernel_size=per_pixel_embedding_head_config.depthwise_kernel_size,
+ prediction_kernel_size=per_pixel_embedding_head_config.prediction_kernel_size,
+ upsample_factor=per_pixel_embedding_head_config.upsample_factor,
+ feature_fusion=per_pixel_embedding_head_config.feature_fusion,
+ decoder_min_level=per_pixel_embedding_head_config.decoder_min_level,
+ decoder_max_level=per_pixel_embedding_head_config.decoder_max_level,
+ low_level=per_pixel_embedding_head_config.low_level,
+ low_level_num_filters=per_pixel_embedding_head_config.low_level_num_filters,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer,
+ bias_initializer=tf.constant_initializer(0.0))
+
+ postprocessing_config = model_config.panoptic_generator
+ if segmentation_inference:
+ is_thing = [False] * model_config.num_classes
+ else:
+ is_thing = ([False] + [True] * (model_config.num_thing_classes - 1) +
+ [False] *
+ (model_config.num_classes - model_config.num_thing_classes))
+ panoptic_generator = None
+ if postprocessing_config:
+ panoptic_generator = panoptic_segmentation_generator.MaskConverPanopticGenerator(
+ output_size=model_config.padded_output_size,
+ num_classes=model_config.num_classes,
+ is_thing=is_thing,
+ num_instances=model_config.num_instances,
+ object_mask_threshold=postprocessing_config.object_mask_threshold,
+ small_area_threshold=postprocessing_config.small_area_threshold,
+ overlap_threshold=postprocessing_config.overlap_threshold,
+ rescale_predictions=postprocessing_config.rescale_predictions,
+ use_hardware_optimization=postprocessing_config.use_hardware_optimization,
+ )
+ # pylint: enable=line-too-long
+ mlp_embedding_head = maskconver_head.MLP(
+ hidden_dim=1024,
+ output_dim=model_config.embedding_size,
+ num_layers=2,
+ activation=norm_activation_config.activation,
+ l2_regularizer=l2_regularizer)
+
+ model = multiscale_maskconver_model.MultiScaleMaskConverModel(
+ backbone,
+ decoder,
+ mask_decoder=mask_decoder, # pyrefly: ignore[unbound-name]
+ embedding_head=mask_embedding_head,
+ class_head=class_head,
+ per_pixel_embeddings_head=per_pixel_embedding_head,
+ mlp_embedding_head=mlp_embedding_head,
+ panoptic_generator=panoptic_generator,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ max_proposals=model_config.num_instances,
+ padded_output_size=model_config.padded_output_size,
+ l2_regularizer=l2_regularizer,
+ embedding_size=model_config.embedding_size,
+ num_classes=model_config.num_classes)
+ return model
diff --git a/official/projects/maskconver/modeling/fpn.py b/official/projects/maskconver/modeling/fpn.py
new file mode 100644
index 00000000000..8eaaab7ba6f
--- /dev/null
+++ b/official/projects/maskconver/modeling/fpn.py
@@ -0,0 +1,275 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains the definitions of a custom Feature Pyramid Networks (FPN).
+
+We allow using layer norm, and 7x7 convs.
+"""
+
+from typing import Any, Mapping, Optional
+
+from absl import logging
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import hyperparams
+from official.modeling import tf_utils
+from official.vision.modeling.decoders import factory
+from official.vision.ops import spatial_transform_ops
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MaskConverFPN(tf_keras.Model):
+ """Creates a custom Feature Pyramid Network (FPN).
+
+ This implements the paper:
+ Tsung-Yi Lin, Piotr Dollar, Ross Girshick, Kaiming He, Bharath Hariharan, and
+ Serge Belongie.
+ Feature Pyramid Networks for Object Detection.
+ (https://arxiv.org/pdf/1612.03144)
+
+ We allow using layer norm and customize the kernel size for the depthwise
+ convs.
+ """
+
+ def __init__(
+ self,
+ input_specs: Mapping[str, tf.TensorShape],
+ min_level: int = 3,
+ max_level: int = 7,
+ num_filters: int = 256,
+ fusion_type: str = 'sum',
+ use_separable_conv: bool = False,
+ depthwise_kernel_size: int = 3,
+ use_keras_layer: bool = False,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ use_layer_norm: bool = False,
+ kernel_initializer: str = 'VarianceScaling',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initializes a Feature Pyramid Network (FPN).
+
+ Args:
+ input_specs: A `dict` of input specifications. A dictionary consists of
+ {level: TensorShape} from a backbone.
+ min_level: An `int` of minimum level in FPN output feature maps.
+ max_level: An `int` of maximum level in FPN output feature maps.
+ num_filters: An `int` number of filters in FPN layers.
+ fusion_type: A `str` of `sum` or `concat`. Whether performing sum or
+ concat for feature fusion.
+ use_separable_conv: A `bool`. If True use separable convolution for
+ convolution in FPN layers.
+ depthwise_kernel_size: An `int` for kernel size of the depthwise kernel
+ size.
+ use_keras_layer: A `bool`. If Ture use keras layers as many as possible.
+ activation: A `str` name of the activation function.
+ use_sync_bn: A `bool`. If True, use synchronized batch normalization.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ use_layer_norm: A `bool` of whether or not using layer norm.
+ kernel_initializer: A `str` name of kernel_initializer for convolutional
+ layers.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ self._config_dict = {
+ 'input_specs': input_specs,
+ 'min_level': min_level,
+ 'max_level': max_level,
+ 'num_filters': num_filters,
+ 'fusion_type': fusion_type,
+ 'use_separable_conv': use_separable_conv,
+ 'use_keras_layer': use_keras_layer,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'norm_momentum': norm_momentum,
+ 'norm_epsilon': norm_epsilon,
+ 'kernel_initializer': kernel_initializer,
+ 'kernel_regularizer': kernel_regularizer,
+ 'bias_regularizer': bias_regularizer,
+ 'use_layer_norm': use_layer_norm,
+ 'depthwise_kernel_size': depthwise_kernel_size,
+ }
+ if use_separable_conv:
+ conv2d = tf_keras.layers.SeparableConv2D
+ else:
+ conv2d = tf_keras.layers.Conv2D
+ if use_sync_bn:
+ norm = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ norm = tf_keras.layers.BatchNormalization
+
+ activation_fn = tf_utils.get_activation(activation, use_keras_layer=True)
+
+ # Build input feature pyramid.
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ bn_axis = -1
+ else:
+ bn_axis = 1
+
+ if use_layer_norm:
+ norm_layer = lambda: tf_keras.layers.LayerNormalization(epsilon=1e-6)
+ else:
+ norm_layer = lambda: norm( # pylint: disable=g-long-lambda
+ axis=bn_axis, momentum=norm_momentum, epsilon=norm_epsilon)
+
+ # Get input feature pyramid from backbone.
+ logging.info('FPN input_specs: %s', input_specs)
+ inputs = self._build_input_pyramid(input_specs, min_level)
+ backbone_max_level = min(int(max(inputs.keys())), max_level)
+
+ # Build lateral connections.
+ feats_lateral = {}
+ for level in range(min_level, backbone_max_level + 1):
+ feats_lateral[str(level)] = conv2d(
+ filters=num_filters,
+ kernel_size=1,
+ padding='same',
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer)(
+ inputs[str(level)])
+
+ # Build top-down path.
+ feats = {str(backbone_max_level): feats_lateral[str(backbone_max_level)]}
+ for level in range(backbone_max_level - 1, min_level - 1, -1):
+ feat_a = spatial_transform_ops.nearest_upsampling(
+ feats[str(level + 1)], 2, use_keras_layer=use_keras_layer)
+ feat_b = feats_lateral[str(level)]
+
+ if fusion_type == 'sum':
+ if use_keras_layer:
+ feats[str(level)] = tf_keras.layers.Add()([feat_a, feat_b])
+ else:
+ feats[str(level)] = feat_a + feat_b
+ elif fusion_type == 'concat':
+ if use_keras_layer:
+ feats[str(level)] = tf_keras.layers.Concatenate(axis=-1)(
+ [feat_a, feat_b])
+ else:
+ feats[str(level)] = tf.concat([feat_a, feat_b], axis=-1)
+ else:
+ raise ValueError('Fusion type {} not supported.'.format(fusion_type))
+
+ # Build post-hoc 3x3 convolution kernel.
+ for level in range(min_level, backbone_max_level + 1):
+ feats[str(level)] = conv2d(
+ filters=num_filters,
+ strides=1,
+ kernel_size=depthwise_kernel_size if use_separable_conv else 3,
+ padding='same',
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer)(
+ feats[str(level)])
+
+ # Build coarser FPN levels introduced for RetinaNet.
+ for level in range(backbone_max_level + 1, max_level + 1):
+ feats_in = feats[str(level - 1)]
+ if level > backbone_max_level + 1:
+ feats_in = activation_fn(feats_in)
+ feats[str(level)] = conv2d(
+ filters=num_filters,
+ strides=2,
+ kernel_size=3,
+ padding='same',
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ )(feats_in)
+
+ # Apply batch norm layers.
+ for level in range(min_level, max_level + 1):
+ feats[str(level)] = norm_layer()(feats[str(level)])
+
+ self._output_specs = {
+ str(level): feats[str(level)].get_shape()
+ for level in range(min_level, max_level + 1)
+ }
+
+ super().__init__(inputs=inputs, outputs=feats, **kwargs)
+
+ def _build_input_pyramid(self, input_specs: Mapping[str, tf.TensorShape],
+ min_level: int):
+ assert isinstance(input_specs, dict)
+ if min(input_specs.keys()) > str(min_level):
+ raise ValueError(
+ 'Backbone min level should be less or equal to FPN min level')
+
+ inputs = {}
+ for level, spec in input_specs.items():
+ inputs[level] = tf_keras.Input(shape=spec[1:])
+ return inputs
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+ @property
+ def output_specs(self) -> Mapping[str, tf.TensorShape]:
+ """A dict of {level: TensorShape} pairs for the model output."""
+ return self._output_specs
+
+
+@factory.register_decoder_builder('maskconver_fpn')
+def build_fpn_decoder(
+ input_specs: Mapping[str, tf.TensorShape],
+ model_config: hyperparams.Config,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> tf_keras.Model:
+ """Builds FPN decoder from a config.
+
+ Args:
+ input_specs: A `dict` of input specifications. A dictionary consists of
+ {level: TensorShape} from a backbone.
+ model_config: A OneOfConfig. Model config.
+ l2_regularizer: A `tf_keras.regularizers.Regularizer` instance. Default to
+ None.
+
+ Returns:
+ A `tf_keras.Model` instance of the MaskConverFPN decoder.
+
+ Raises:
+ ValueError: If the model_config.decoder.type is not `maskconver_fpn`.
+ """
+ decoder_type = model_config.decoder.type
+ decoder_cfg = model_config.decoder.get()
+ if decoder_type != 'maskconver_fpn':
+ raise ValueError(f'Inconsistent decoder type {decoder_type}. '
+ 'Need to be `maskconver_fpn`.')
+ norm_activation_config = model_config.norm_activation
+ return MaskConverFPN(
+ input_specs=input_specs,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ num_filters=decoder_cfg.num_filters,
+ fusion_type=decoder_cfg.fusion_type,
+ use_separable_conv=decoder_cfg.use_separable_conv,
+ use_keras_layer=decoder_cfg.use_keras_layer,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ use_layer_norm=decoder_cfg.use_layer_norm,
+ depthwise_kernel_size=decoder_cfg.depthwise_kernel_size,
+ kernel_regularizer=l2_regularizer)
diff --git a/official/projects/maskconver/modeling/layers/copypaste.py b/official/projects/maskconver/modeling/layers/copypaste.py
new file mode 100644
index 00000000000..5ee14014415
--- /dev/null
+++ b/official/projects/maskconver/modeling/layers/copypaste.py
@@ -0,0 +1,250 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic CopyPaste."""
+import random
+
+import tensorflow as tf, tf_keras
+
+from official.vision.ops import preprocess_ops
+
+
+GLOBAL_SEED_SET = False
+PAD_VALUE = 0
+
+
+def random_uniform_strong(minval,
+ maxval,
+ dtype=tf.float32,
+ seed=None,
+ shape=None):
+ """A unified function for consistent random number generation.
+
+ Equivalent to tf.random.uniform, except that minval and maxval are flipped if
+ minval is greater than maxval. Seed Safe random number generator.
+
+ Args:
+ minval: An `int` for a lower or upper endpoint of the interval from which to
+ choose the random number.
+ maxval: An `int` for the other endpoint.
+ dtype: The output type of the tensor.
+ seed: An `int` used to set the seed.
+ shape: List or 1D tf.Tensor, output shape of the random generator.
+
+ Returns:
+ A random tensor of type `dtype` that falls between `minval` and `maxval`
+ excluding the larger one.
+ """
+ if GLOBAL_SEED_SET:
+ seed = None
+
+ if minval > maxval:
+ minval, maxval = maxval, minval
+ return tf.random.uniform(
+ shape=shape or [], minval=minval, maxval=maxval, seed=seed, dtype=dtype)
+
+
+class CopyPaste:
+ """Panoptic CopyPaste."""
+
+ def __init__(self,
+ output_size,
+ copypaste_frequency=0.5,
+ stuff_mask_drop_rate=0.0,
+ copypaste_aug_scale_max=1.0,
+ copypaste_aug_scale_min=1.0,
+ num_thing_classes=91,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ random_flip=False,
+ pad_value=PAD_VALUE,
+ seed=None):
+ """Initializes parameters for Copy Paste.
+
+ Args:
+ output_size: `Tensor` or `List` for [height, width] of output image.
+ copypaste_frequency: `float` indicating how often to apply copypaste.
+ stuff_mask_drop_rate: `float` indicating drop rate for stuff masks.
+ copypaste_aug_scale_max: `float`, how much to scale the copypaste
+ image.
+ copypaste_aug_scale_min: `float`, how much to scale the copypaste
+ image.
+ num_thing_classes: `int`, number of thing classes.
+ aug_scale_min: `float` indicating the minimum scaling value for image
+ scale jitter.
+ aug_scale_max: `float` indicating the maximum scaling value for image
+ scale jitter.
+ random_flip: `bool` whether or not to random flip the image.
+ pad_value: `int` padding value.
+ seed: `int` the seed for random number generation.
+ """
+
+ self._output_size = output_size
+ self._aug_scale_min = aug_scale_min
+ self._aug_scale_max = aug_scale_max
+ self._copypaste_aug_scale_min = copypaste_aug_scale_min
+ self._copypaste_aug_scale_max = copypaste_aug_scale_max
+ self._random_flip = random_flip
+ self._pad_value = pad_value
+ self._copypaste_frequency = copypaste_frequency
+ self._stuff_mask_drop_rate = stuff_mask_drop_rate
+ self._num_thing_classes = num_thing_classes
+
+ self._deterministic = seed is not None
+ self._seed = seed if seed is not None else random.randint(0, 2**30)
+
+ def _process_image(self, sample, aug_min, aug_max, seed=None):
+ """Process and augment each image."""
+ if self._random_flip:
+ instance_mask = sample['groundtruth_panoptic_instance_mask']
+ category_mask = sample['groundtruth_panoptic_category_mask']
+ image_mask = tf.concat(
+ [tf.cast(sample['image'], tf.uint8), category_mask, instance_mask],
+ axis=2)
+
+ image_mask, _, _ = preprocess_ops.random_horizontal_flip(
+ image_mask)
+
+ instance_mask = image_mask[:, :, -1:]
+ category_mask = image_mask[:, :, -2:-1]
+ image = tf.cast(image_mask[:, :, :-2], tf.uint8)
+
+ # Resizes and crops image.
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ self._output_size,
+ padded_size=self._output_size,
+ aug_scale_min=aug_min,
+ aug_scale_max=aug_max,
+ seed=seed)
+
+ def _process_mask(mask, ignore_label, image_info):
+ mask = tf.cast(mask, dtype=tf.float32)
+ mask = tf.reshape(mask, shape=[1, sample['height'], sample['width'], 1])
+ mask += 1
+
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ mask = preprocess_ops.resize_and_crop_masks(
+ mask, image_scale, self._output_size, offset)
+
+ mask -= 1
+ # Assign ignore label to the padded region.
+ mask = tf.where(
+ tf.equal(mask, -1),
+ ignore_label * tf.ones_like(mask),
+ mask)
+ mask = tf.squeeze(mask, axis=0)
+ return mask
+
+ panoptic_category_mask = _process_mask(
+ category_mask,
+ 0, image_info)
+ panoptic_instance_mask = _process_mask(
+ instance_mask,
+ 0, image_info)
+ sample['image'] = image
+ sample['height'] = tf.cast(self._output_size[0], tf.int32)
+ sample['width'] = tf.cast(self._output_size[1], tf.int32)
+ sample['groundtruth_panoptic_category_mask'] = panoptic_category_mask
+ sample['groundtruth_panoptic_instance_mask'] = panoptic_instance_mask
+
+ return sample
+
+ def _patch(self, one, two):
+ """Stitch together 2 images in totality."""
+ sample = one
+ unique_instance_ids, _ = tf.unique(
+ tf.reshape(two['groundtruth_panoptic_instance_mask'], [-1]))
+ first_image = one['image']
+ first_instance_mask = one['groundtruth_panoptic_instance_mask']
+ second_instance_mask = two['groundtruth_panoptic_instance_mask']
+ first_category_mask = one['groundtruth_panoptic_category_mask']
+ second_category_mask = two['groundtruth_panoptic_category_mask']
+ max_id = tf.reduce_max(one['groundtruth_panoptic_instance_mask'])
+
+ for inst_id in unique_instance_ids:
+ num = random_uniform_strong(
+ 0.0, 1.0, dtype=tf.float32, seed=self._seed)
+ if tf.logical_and(inst_id > 0, num < self._copypaste_frequency):
+ first_instance_mask = tf.where(second_instance_mask == inst_id,
+ second_instance_mask + max_id + 1,
+ first_instance_mask)
+ first_image = tf.where(second_instance_mask == inst_id, two['image'],
+ first_image)
+ first_category_mask = tf.where(second_instance_mask == inst_id,
+ second_category_mask,
+ first_category_mask)
+ stuff_classes, _ = tf.unique(
+ tf.reshape(two['groundtruth_panoptic_category_mask'], [-1]))
+ stuff_classes = tf.boolean_mask(
+ stuff_classes, stuff_classes >= self._num_thing_classes)
+
+ for stuff_class in stuff_classes:
+ num = random_uniform_strong(
+ 0.0, 1.0, dtype=tf.float32, seed=self._seed)
+ if num < self._copypaste_frequency:
+ random_tensor = tf.random.uniform(
+ self._output_size + [1], minval=0.0, maxval=1.0, seed=self._seed)
+ stuff_mask_to_copy = tf.logical_and(
+ second_category_mask == stuff_class,
+ random_tensor > self._stuff_mask_drop_rate)
+ first_image = tf.where(stuff_mask_to_copy, two['image'],
+ first_image)
+ first_category_mask = tf.where(stuff_mask_to_copy,
+ second_category_mask,
+ first_category_mask)
+ first_instance_mask = tf.where(stuff_mask_to_copy,
+ tf.zeros_like(first_instance_mask),
+ first_instance_mask)
+
+ sample['image'] = first_image
+ sample['groundtruth_panoptic_instance_mask'] = first_instance_mask
+ sample['groundtruth_panoptic_category_mask'] = first_category_mask
+
+ sample['image'] = tf.cast(first_image, tf.uint8)
+ return sample
+
+ def _copypaste(self, one, two):
+ """Apply copypaste on 2 images."""
+ one = self._process_image(one, self._aug_scale_min, self._aug_scale_max,
+ self._seed)
+ two = self._process_image(
+ two, self._copypaste_aug_scale_min, self._copypaste_aug_scale_max,
+ self._seed + 1)
+ copypasted = self._patch(one, two)
+ return copypasted
+
+ def _apply(self, dataset):
+ """Apply copypaste to an input dataset."""
+ determ = self._deterministic
+ dataset = dataset.prefetch(tf.data.AUTOTUNE)
+ one = dataset.shuffle(1000, seed=self._seed, reshuffle_each_iteration=True)
+ two = dataset.shuffle(
+ 1000, seed=self._seed + 1, reshuffle_each_iteration=True)
+
+ dataset = tf.data.Dataset.zip((one, two))
+ dataset = dataset.map(
+ self._copypaste,
+ num_parallel_calls=tf.data.AUTOTUNE,
+ deterministic=determ)
+
+ return dataset
+
+ def copypaste_fn(self, is_training=True):
+ """Determine which function to apply based on whether model is training."""
+ if is_training:
+ return self._apply
+ else:
+ return lambda dataset: dataset
diff --git a/official/projects/maskconver/modeling/layers/maskconver_head.py b/official/projects/maskconver/modeling/layers/maskconver_head.py
new file mode 100644
index 00000000000..bcbe3ff868c
--- /dev/null
+++ b/official/projects/maskconver/modeling/layers/maskconver_head.py
@@ -0,0 +1,321 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definition for postprocessing layer to genrate panoptic segmentations."""
+
+from typing import Any, List, Optional, Union
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.vision.modeling.layers import nn_layers
+from official.vision.ops import spatial_transform_ops
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MaskConverHead(tf_keras.layers.Layer):
+ """Creates a MaskConver head."""
+
+ def __init__(
+ self,
+ num_classes: int,
+ level: Union[int, str],
+ num_convs: int = 2,
+ num_filters: int = 256,
+ use_depthwise_convolution: bool = False,
+ depthwise_kernel_size: int = 3,
+ prediction_kernel_size: int = 1,
+ upsample_factor: int = 1,
+ feature_fusion: Optional[str] = None,
+ decoder_min_level: Optional[int] = None,
+ decoder_max_level: Optional[int] = None,
+ low_level: int = 2,
+ low_level_num_filters: int = 48,
+ num_decoder_filters: int = 256,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ use_layer_norm: bool = False,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_initializer: Optional[Any] = tf.constant_initializer(0.0),
+ **kwargs):
+ """Initializes a maskconver head.
+
+ Args:
+ num_classes: An `int` number of mask classification categories. The number
+ of classes does not include background class.
+ level: An `int` or `str`, level to use to build maskconver head.
+ num_convs: An `int` number of stacked convolution before the last
+ prediction layer.
+ num_filters: An `int` number to specify the number of filters used.
+ Default is 256.
+ use_depthwise_convolution: A bool to specify if use depthwise separable
+ convolutions.
+ depthwise_kernel_size: An `int` for depthwise kernel size.
+ prediction_kernel_size: An `int` number to specify the kernel size of the
+ prediction layer.
+ upsample_factor: An `int` number to specify the upsampling factor to
+ generate finer mask. Default 1 means no upsampling is applied.
+ feature_fusion: One of the constants in nn_layers.FeatureFusion, namely
+ `deeplabv3plus`, `pyramid_fusion`, `panoptic_fpn_fusion`,
+ `deeplabv3plus_sum_to_merge`, or None. If `deeplabv3plus`, features from
+ decoder_features[level] will be fused with low level feature maps from
+ backbone. If `pyramid_fusion`, multiscale features will be resized and
+ fused at the target level.
+ decoder_min_level: An `int` of minimum level from decoder to use in
+ feature fusion. It is only used when feature_fusion is set to
+ `panoptic_fpn_fusion`.
+ decoder_max_level: An `int` of maximum level from decoder to use in
+ feature fusion. It is only used when feature_fusion is set to
+ `panoptic_fpn_fusion`.
+ low_level: An `int` of backbone level to be used for feature fusion. It is
+ used when feature_fusion is set to `deeplabv3plus` or
+ `deeplabv3plus_sum_to_merge`.
+ low_level_num_filters: An `int` of reduced number of filters for the low
+ level features before fusing it with higher level features. It is only
+ used when feature_fusion is set to `deeplabv3plus` or
+ `deeplabv3plus_sum_to_merge`.
+ num_decoder_filters: An `int` of number of filters in the decoder outputs.
+ It is only used when feature_fusion is set to `panoptic_fpn_fusion`.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ use_layer_norm: A `bool` for whether to use layer norm.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ bias_initializer: Bias initializer for the classification layer.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+
+ self._config_dict = {
+ 'num_classes': num_classes,
+ 'level': level,
+ 'num_convs': num_convs,
+ 'num_filters': num_filters,
+ 'use_depthwise_convolution': use_depthwise_convolution,
+ 'depthwise_kernel_size': depthwise_kernel_size,
+ 'prediction_kernel_size': prediction_kernel_size,
+ 'upsample_factor': upsample_factor,
+ 'feature_fusion': feature_fusion,
+ 'decoder_min_level': decoder_min_level,
+ 'decoder_max_level': decoder_max_level,
+ 'low_level': low_level,
+ 'low_level_num_filters': low_level_num_filters,
+ 'num_decoder_filters': num_decoder_filters,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'norm_momentum': norm_momentum,
+ 'norm_epsilon': norm_epsilon,
+ 'kernel_regularizer': kernel_regularizer,
+ 'bias_regularizer': bias_regularizer,
+ 'bias_initializer': bias_initializer,
+ 'use_layer_norm': use_layer_norm,
+ }
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ else:
+ self._bn_axis = 1
+ self._activation = tf_utils.get_activation(activation)
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the variables of the segmentation head."""
+ use_depthwise_convolution = self._config_dict['use_depthwise_convolution']
+ conv_op = tf_keras.layers.Conv2D
+ if self._config_dict['use_layer_norm']:
+ bn_layer = lambda: tf_keras.layers.LayerNormalization(epsilon=1e-6)
+ else:
+ bn_kwargs = {
+ 'axis': self._bn_axis,
+ 'momentum': self._config_dict['norm_momentum'],
+ 'epsilon': self._config_dict['norm_epsilon'],
+ }
+ if self._config_dict['use_sync_bn']:
+ bn_layer = lambda: tf_keras.layers.experimental.SyncBatchNormalization( # pylint: disable=g-long-lambda
+ **bn_kwargs)
+ else:
+ bn_layer = lambda: tf_keras.layers.BatchNormalization(**bn_kwargs)
+
+ if self._config_dict['feature_fusion'] in {'deeplabv3plus',
+ 'deeplabv3plus_sum_to_merge'}:
+ # Deeplabv3+ feature fusion layers.
+ self._dlv3p_conv = conv_op(
+ kernel_size=1,
+ padding='same',
+ use_bias=False,
+ kernel_initializer=tf_keras.initializers.he_normal(),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ name='segmentation_head_deeplabv3p_fusion_conv',
+ filters=self._config_dict['low_level_num_filters'])
+
+ self._dlv3p_norm = bn_layer()
+
+ elif self._config_dict['feature_fusion'] == 'panoptic_fpn_fusion':
+ self._panoptic_fpn_fusion = nn_layers.PanopticFPNFusion(
+ min_level=self._config_dict['decoder_min_level'],
+ max_level=self._config_dict['decoder_max_level'],
+ target_level=self._config_dict['level'],
+ num_filters=self._config_dict['num_filters'],
+ num_fpn_filters=self._config_dict['num_decoder_filters'],
+ activation=self._config_dict['activation'],
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'])
+
+ # Segmentation head layers.
+ self._convs = []
+ self._norms = []
+ for i in range(self._config_dict['num_convs']): # pyrefly: ignore[bad-argument-type]
+ if use_depthwise_convolution:
+ self._convs.append(
+ tf_keras.layers.DepthwiseConv2D(
+ name='segmentation_head_depthwise_conv_{}'.format(i),
+ kernel_size=self._config_dict['depthwise_kernel_size'],
+ padding='same',
+ use_bias=False,
+ depth_multiplier=1))
+ self._norms.append(bn_layer())
+ conv_name = 'segmentation_head_conv_{}'.format(i)
+ self._convs.append(
+ conv_op(
+ name=conv_name,
+ filters=self._config_dict['num_filters'],
+ kernel_size=3 if not use_depthwise_convolution else 1,
+ padding='same',
+ use_bias=False,
+ kernel_initializer=tf_keras.initializers.he_normal(),
+ kernel_regularizer=self._config_dict['kernel_regularizer']))
+ self._norms.append(bn_layer())
+
+ self._classifier = conv_op(
+ name='segmentation_output',
+ filters=self._config_dict['num_classes'],
+ kernel_size=self._config_dict['prediction_kernel_size'],
+ padding='same',
+ bias_initializer=self._config_dict['bias_initializer'],
+ kernel_initializer=tf_keras.initializers.truncated_normal(stddev=0.01),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'])
+
+ super().build(input_shape)
+
+ def call(self, inputs):
+ """Forward pass of the segmentation head.
+
+ It supports both a tuple of 2 tensors or 2 dictionaries. The first is
+ backbone endpoints, and the second is decoder endpoints. When inputs are
+ tensors, they are from a single level of feature maps. When inputs are
+ dictionaries, they contain multiple levels of feature maps, where the key
+ is the index of feature map.
+
+ Args:
+ inputs: A tuple of 2 feature map tensors of shape
+ [batch, height_l, width_l, channels] or 2 dictionaries of tensors:
+ - key: A `str` of the level of the multilevel features.
+ - values: A `tf.Tensor` of the feature map tensors, whose shape is
+ [batch, height_l, width_l, channels].
+ The first is backbone endpoints, and the second is decoder endpoints.
+ Returns:
+ segmentation prediction mask: A `tf.Tensor` of the segmentation mask
+ scores predicted from input features.
+ """
+
+ backbone_output = inputs[0]
+ decoder_output = inputs[1]
+ if self._config_dict['feature_fusion'] in {'deeplabv3plus',
+ 'deeplabv3plus_sum_to_merge'}:
+ # deeplabv3+ feature fusion
+ x = decoder_output[str(self._config_dict['level'])] if isinstance(
+ decoder_output, dict) else decoder_output
+ y = backbone_output[str(self._config_dict['low_level'])] if isinstance(
+ backbone_output, dict) else backbone_output
+ y = self._dlv3p_norm(self._dlv3p_conv(y))
+ y = self._activation(y)
+
+ x = tf.image.resize(
+ x, tf.shape(y)[1:3], method=tf.image.ResizeMethod.BILINEAR)
+ x = tf.cast(x, dtype=y.dtype)
+ if self._config_dict['feature_fusion'] == 'deeplabv3plus':
+ x = tf.concat([x, y], axis=self._bn_axis)
+ else:
+ x = tf_keras.layers.Add()([x, y])
+ elif self._config_dict['feature_fusion'] == 'pyramid_fusion':
+ if not isinstance(decoder_output, dict):
+ raise ValueError('Only support dictionary decoder_output.')
+ x = nn_layers.pyramid_feature_fusion(decoder_output,
+ self._config_dict['level'])
+ elif self._config_dict['feature_fusion'] == 'panoptic_fpn_fusion':
+ x = self._panoptic_fpn_fusion(decoder_output)
+ else:
+ x = decoder_output[str(self._config_dict['level'])] if isinstance(
+ decoder_output, dict) else decoder_output
+
+ for conv, norm in zip(self._convs, self._norms):
+ x = conv(x)
+ x = norm(x)
+ x = self._activation(x)
+ if self._config_dict['upsample_factor'] > 1: # pyrefly: ignore[unsupported-operation]
+ x = spatial_transform_ops.nearest_upsampling(
+ x, scale=self._config_dict['upsample_factor'])
+
+ return self._classifier(x)
+
+ def get_config(self):
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(self._config_dict.items()))
+
+ @classmethod
+ def from_config(cls, config):
+ return cls(**config)
+
+
+class MLP(tf_keras.Model):
+ """MLP."""
+
+ def __init__(
+ self,
+ hidden_dim: int = 256,
+ output_dim: int = 256,
+ num_layers: int = 2,
+ activation: str = 'swish',
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self.num_layers = num_layers
+ dims = [hidden_dim] * (num_layers - 1)
+ # pylint: disable=g-complex-comprehension
+ bn_layer = lambda: tf_keras.layers.LayerNormalization(epsilon=1e-6)
+ self.dense = [
+ tf_keras.layers.Dense(d, kernel_regularizer=l2_regularizer)
+ for d in dims + [output_dim]
+ ]
+ self.norms = [bn_layer() for _ in dims]
+ self.activation = tf_keras.activations.get(activation)
+
+ def call(
+ self, inputs: tf.Tensor, training: Any = None, mask: Any = None
+ ) -> tf.Tensor:
+ x = inputs
+ for i, layer in enumerate(self.dense):
+ x = (
+ self.activation(self.norms[i](layer(x)))
+ if i < self.num_layers - 1
+ else layer(x)
+ )
+ return x
diff --git a/official/projects/maskconver/modeling/layers/multiscale_maskconver_head.py b/official/projects/maskconver/modeling/layers/multiscale_maskconver_head.py
new file mode 100644
index 00000000000..c051c9495d5
--- /dev/null
+++ b/official/projects/maskconver/modeling/layers/multiscale_maskconver_head.py
@@ -0,0 +1,193 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definition for multi-scale MaskConver head."""
+
+from typing import Any, List, Optional, Union
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.vision.ops import spatial_transform_ops
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MultiScaleMaskConverHead(tf_keras.layers.Layer):
+ """Creates a MaskConver head."""
+
+ def __init__(
+ self,
+ num_classes: int,
+ min_level: Union[int, str],
+ max_level: Union[int, str],
+ num_convs: int = 2,
+ num_filters: int = 256,
+ use_depthwise_convolution: bool = False,
+ depthwise_kernel_size: int = 3,
+ prediction_kernel_size: int = 1,
+ upsample_factor: int = 1,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ use_layer_norm: bool = True,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_initializer: Optional[Any] = tf.constant_initializer(0.0),
+ **kwargs):
+ """Initializes a maskconver head.
+
+ Args:
+ num_classes: An `int` number of mask classification categories. The number
+ of classes does not include background class.
+ min_level: An `int` or `str`, min level to use to build maskconver head.
+ max_level: An `int` or `str`, max level to use to build maskconver head.
+ num_convs: An `int` number of stacked convolution before the last
+ prediction layer.
+ num_filters: An `int` number to specify the number of filters used.
+ Default is 256.
+ use_depthwise_convolution: A bool to specify if use depthwise separable
+ convolutions.
+ depthwise_kernel_size: An `int` for the depthwise kernel size.
+ prediction_kernel_size: An `int` number to specify the kernel size of the
+ prediction layer.
+ upsample_factor: An `int` number to specify the upsampling factor to
+ generate finer mask. Default 1 means no upsampling is applied.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ use_layer_norm: A `bool` whether to use layer norm.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ bias_initializer: Bias initializer for the classification layer.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+
+ self._config_dict = {
+ 'num_classes': num_classes,
+ 'min_level': min_level,
+ 'max_level': max_level,
+ 'num_convs': num_convs,
+ 'num_filters': num_filters,
+ 'use_depthwise_convolution': use_depthwise_convolution,
+ 'depthwise_kernel_size': depthwise_kernel_size,
+ 'prediction_kernel_size': prediction_kernel_size,
+ 'upsample_factor': upsample_factor,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'norm_momentum': norm_momentum,
+ 'norm_epsilon': norm_epsilon,
+ 'kernel_regularizer': kernel_regularizer,
+ 'bias_regularizer': bias_regularizer,
+ 'bias_initializer': bias_initializer,
+ 'use_layer_norm': use_layer_norm,
+ }
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ else:
+ self._bn_axis = 1
+ self._activation = tf_utils.get_activation(activation)
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the variables of the segmentation head."""
+ use_depthwise_convolution = self._config_dict['use_depthwise_convolution']
+ conv_op = tf_keras.layers.Conv2D
+ if self._config_dict['use_layer_norm']:
+ bn_layer = lambda: tf_keras.layers.LayerNormalization(epsilon=1e-6)
+ else:
+ bn_kwargs = {
+ 'axis': self._bn_axis,
+ 'momentum': self._config_dict['norm_momentum'],
+ 'epsilon': self._config_dict['norm_epsilon'],
+ }
+ if self._config_dict['use_sync_bn']:
+ bn_layer = lambda: tf_keras.layers.experimental.SyncBatchNormalization( # pylint: disable=g-long-lambda
+ **bn_kwargs)
+ else:
+ bn_layer = lambda: tf_keras.layers.BatchNormalization(**bn_kwargs)
+
+ # Segmentation head layers.
+ self._convs = []
+ self._norms = []
+ for level in range(
+ self._config_dict['min_level'], self._config_dict['max_level'] + 1 # pyrefly: ignore[bad-argument-type, unsupported-operation]
+ ):
+ level_norms = []
+ for i in range(self._config_dict['num_convs']):
+ # We use shared convolution layers across levels.
+ if use_depthwise_convolution:
+ if level == self._config_dict['min_level']:
+ self._convs.append(
+ tf_keras.layers.DepthwiseConv2D(
+ name='segmentation_head_depthwise_conv_{}'.format(i),
+ kernel_size=self._config_dict['depthwise_kernel_size'],
+ padding='same',
+ use_bias=False,
+ depth_multiplier=1))
+ level_norms.append(bn_layer())
+ if level == self._config_dict['min_level']:
+ conv_name = 'segmentation_head_conv_{}'.format(i)
+ self._convs.append(
+ conv_op(
+ name=conv_name,
+ filters=self._config_dict['num_filters'],
+ kernel_size=1 if use_depthwise_convolution else 3,
+ padding='same',
+ use_bias=False,
+ kernel_initializer=tf_keras.initializers.he_normal(),
+ kernel_regularizer=self._config_dict['kernel_regularizer']))
+ level_norms.append(bn_layer())
+ self._norms.append(level_norms)
+
+ self._classifier = conv_op(
+ name='segmentation_output',
+ filters=self._config_dict['num_classes'],
+ kernel_size=self._config_dict['prediction_kernel_size'],
+ padding='same',
+ bias_initializer=self._config_dict['bias_initializer'],
+ kernel_initializer=tf_keras.initializers.truncated_normal(stddev=0.01),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'])
+
+ super().build(input_shape)
+
+ def call(self, inputs):
+ """Forward pass of the multiscale maskconver head."""
+ outputs = {}
+ for i, level in enumerate(
+ range(self._config_dict['min_level'], # pyrefly: ignore[bad-argument-type]
+ self._config_dict['max_level'] + 1)): # pyrefly: ignore[unsupported-operation]
+ x = inputs[str(level)]
+ for conv, norm in zip(self._convs, self._norms[i]):
+ x = conv(x)
+ x = norm(x)
+ x = self._activation(x)
+ if self._config_dict['upsample_factor'] > 1:
+ x = spatial_transform_ops.nearest_upsampling(
+ x, scale=self._config_dict['upsample_factor'])
+ outputs[level] = self._classifier(x)
+
+ return outputs
+
+ def get_config(self):
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(self._config_dict.items()))
+
+ @classmethod
+ def from_config(cls, config):
+ return cls(**config)
diff --git a/official/projects/maskconver/modeling/layers/panoptic_segmentation_generator.py b/official/projects/maskconver/modeling/layers/panoptic_segmentation_generator.py
new file mode 100644
index 00000000000..14f91af0709
--- /dev/null
+++ b/official/projects/maskconver/modeling/layers/panoptic_segmentation_generator.py
@@ -0,0 +1,337 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Postprocessing layer to generate panoptic segmentations."""
+
+from typing import Any, Dict, List, Mapping, Optional
+import tensorflow as tf, tf_keras
+from official.modeling import activations
+from official.vision.ops import spatial_transform_ops
+
+
+class MaskConverProposalGenerator(tf_keras.layers.Layer):
+ """MaskConverProposalGenerator."""
+
+ def __init__(self,
+ max_proposals: int = 100,
+ peak_error: float = 1e-6,
+ peak_extract_kernel_size: int = 3,
+ **kwargs):
+ """Initialize MaskConverProposalGenerator.
+
+ Args:
+ max_proposals: An `int` specifying the maximum number of max_proposals.
+ peak_error: A `float` for determining non-valid heatmap locations to mask.
+ peak_extract_kernel_size: An `int` indicating the kernel size used when
+ performing max-pool over the heatmaps to detect valid center locations
+ from its neighbors. From the paper, set this to 3 to detect valid.
+ locations that have responses greater than its 8-connected neighbors
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super(MaskConverProposalGenerator, self).__init__(**kwargs)
+
+ # Object center selection parameters
+ self._max_proposals = max_proposals
+ self._peak_error = peak_error
+ self._peak_extract_kernel_size = peak_extract_kernel_size
+
+ def process_heatmap(self, feature_map: tf.Tensor,
+ kernel_size: int) -> tf.Tensor:
+ """Processes the heatmap into peaks for box selection.
+
+ Given a heatmap, this function first masks out nearby heatmap locations of
+ the same class using max-pooling such that, ideally, only one center for the
+ object remains. Then, center locations are masked according to their scores
+ in comparison to a threshold. NOTE: Repurposed from Google OD API.
+
+ Args:
+ feature_map: A Tensor with shape [batch_size, height, width, num_classes]
+ which is the center heatmap predictions.
+ kernel_size: An integer value for max-pool kernel size.
+
+ Returns:
+ A Tensor with the same shape as the input but with non-valid center
+ prediction locations masked out.
+ """
+
+ feature_map = tf.math.sigmoid(feature_map)
+ if not kernel_size or kernel_size == 1:
+ feature_map_peaks = feature_map
+ else:
+ feature_map_max_pool = tf.nn.max_pool(
+ feature_map,
+ ksize=kernel_size,
+ strides=1,
+ padding='SAME')
+
+ feature_map_peak_mask = tf.math.abs(
+ feature_map - feature_map_max_pool) < self._peak_error
+
+ # Zero out everything that is not a peak.
+ feature_map_peaks = (
+ feature_map * tf.cast(feature_map_peak_mask, feature_map.dtype))
+
+ return feature_map_peaks
+
+ def get_top_k_peaks(self,
+ feature_map_peaks: tf.Tensor,
+ batch_size: int,
+ width: int,
+ num_classes: int,
+ k: int = 100):
+ """Gets the scores and indices of the top-k peaks from the feature map.
+
+ This function flattens the feature map in order to retrieve the top-k
+ peaks, then computes the x, y, and class indices for those scores.
+ NOTE: Repurposed from Google OD API.
+
+ Args:
+ feature_map_peaks: A `Tensor` with shape [batch_size, height,
+ width, num_classes] which is the processed center heatmap peaks.
+ batch_size: An `int` that indicates the batch size of the input.
+ width: An `int` that indicates the width (and also height) of the input.
+ num_classes: An `int` for the number of possible classes. This is also
+ the channel depth of the input.
+ k: `int` that controls how many peaks to select.
+
+ Returns:
+ top_scores: A Tensor with shape [batch_size, k] containing the top-k
+ scores.
+ y_indices: A Tensor with shape [batch_size, k] containing the top-k
+ y-indices corresponding to top_scores.
+ x_indices: A Tensor with shape [batch_size, k] containing the top-k
+ x-indices corresponding to top_scores.
+ channel_indices: A Tensor with shape [batch_size, k] containing the top-k
+ channel indices corresponding to top_scores.
+ """
+ # Flatten the entire prediction per batch
+ feature_map_peaks_flat = tf.reshape(feature_map_peaks, [batch_size, -1])
+
+ # top_scores and top_indices have shape [batch_size, k]
+ top_scores, top_indices = tf.math.top_k(feature_map_peaks_flat, k=k)
+
+ # Get x, y and channel indices corresponding to the top indices in the flat
+ # array.
+ y_indices = (top_indices // num_classes) // width
+ x_indices = (top_indices // num_classes) - y_indices * width
+ channel_indices_temp = top_indices // num_classes
+ channel_indices = top_indices - channel_indices_temp * num_classes
+
+ embedding_indices = tf.stack([y_indices, x_indices], axis=2)
+
+ return top_scores, embedding_indices, channel_indices
+
+ def __call__(self, ct_heatmaps: tf.Tensor):
+ # Get heatmaps from decoded outputs via final hourglass stack output
+ shape = tf.shape(ct_heatmaps)
+
+ _, width = shape[1], shape[2]
+ batch_size, num_channels = shape[0], shape[3]
+
+ # Process heatmaps using 3x3 max pool and applying sigmoid
+ peaks = self.process_heatmap(
+ feature_map=ct_heatmaps,
+ kernel_size=self._peak_extract_kernel_size)
+
+ # Get top scores along with their x, y, and class
+ # Each has size [batch_size, k]
+ scores, embedding_indices, channel_indices = self.get_top_k_peaks(
+ feature_map_peaks=peaks,
+ batch_size=batch_size,
+ width=width,
+ num_classes=num_channels,
+ k=self._max_proposals)
+
+ num_proposals = tf.reduce_sum(tf.cast(scores > 0, dtype=tf.int32), axis=1)
+ return {
+ 'classes': channel_indices,
+ 'confidence': scores,
+ 'embedding_indices': embedding_indices,
+ 'num_proposals': num_proposals
+ }
+
+ def get_config(self) -> Mapping[str, Any]:
+ config = {
+ 'max_proposals': self._max_proposals,
+ 'peak_error': self._peak_error,
+ 'peak_extract_kernel_size': self._peak_extract_kernel_size,
+ }
+
+ base_config = super(MaskConverProposalGenerator, self).get_config()
+ return dict(list(base_config.items()) + list(config.items()))
+
+ @classmethod
+ def from_config(cls, config):
+ return cls(**config)
+
+
+class MaskConverPanopticGenerator(tf_keras.layers.Layer):
+ """MaskConver panoptic generator."""
+
+ def __init__(self,
+ output_size: List[int],
+ num_classes: int,
+ is_thing: List[bool],
+ num_instances: int = 100,
+ object_mask_threshold: float = 0.,
+ small_area_threshold: int = 0,
+ overlap_threshold: float = 0.,
+ rescale_predictions: bool = False,
+ use_hardware_optimization: bool = False,
+ **kwargs):
+ super(MaskConverPanopticGenerator, self).__init__(**kwargs)
+
+ self._output_size = output_size
+ assert num_classes == len(is_thing)
+ self._num_classes = num_classes
+ self._is_thing = tf.constant(is_thing, dtype=tf.bool)
+ self._num_instances = num_instances
+ self._object_mask_threshold = object_mask_threshold
+ self._small_area_threshold = small_area_threshold
+ self._overlap_threshold = overlap_threshold
+ self._rescale_predictions = rescale_predictions
+ self._use_hardware_optimization = use_hardware_optimization
+
+ self.config_dict = {
+ 'output_size': output_size,
+ 'num_classes': num_classes,
+ 'is_thing': is_thing,
+ 'num_instances': num_instances,
+ 'object_mask_threshold': object_mask_threshold,
+ 'small_area_threshold': small_area_threshold,
+ 'overlap_threshold': overlap_threshold,
+ 'rescale_predictions': rescale_predictions,
+ 'use_hardware_optimization': use_hardware_optimization,
+ }
+
+ def _resize_and_pad_masks(self, masks, images_info):
+ """Resizes masks to match the original image shape and pads to`output_size`.
+
+ Args:
+ masks: a padded binary mask tensor, shape [b, h, w, 1].
+ images_info: a tensor that holds information about original and
+ preprocessed images.
+
+ Returns:
+ resized and padded masks: tf.Tensor.
+ """
+ rescale_size = tf.cast(
+ tf.math.ceil(images_info[:, 1, :] / images_info[:, 2, :]), tf.int32)
+ image_shape = tf.cast(images_info[:, 0, :], tf.int32)
+ offsets = tf.cast(images_info[:, 3, :], tf.int32)
+
+ return spatial_transform_ops.bilinear_resize_with_crop_and_pad(
+ masks,
+ rescale_size,
+ crop_offset=offsets,
+ crop_size=image_shape,
+ output_size=self._output_size)
+
+ def _generate_panoptic_masks(self, scores: tf.Tensor, classes: tf.Tensor,
+ masks: tf.Tensor) -> Dict[str, tf.Tensor]:
+ """Compute category and instance masks.
+
+ Args:
+ scores: Class confidences, shape [num_proposals].
+ classes: Class IDs, shape [num_proposals]
+ masks: Predicted binary mask logits, shape [num_proposals, height, width].
+
+ Returns:
+ Category and instance masks, both have shape [num_proposals, height,
+ width].
+ """
+ # Collect valid proposals.
+ valid_proposals = scores > self._object_mask_threshold
+
+ cur_scores = tf.cast(valid_proposals, tf.float32) * scores # n
+ cur_classes = tf.cast(valid_proposals, tf.int32) * classes # n
+ cur_masks = tf.cast( # h x w x n
+ valid_proposals, tf.float32)[None, None, :] * masks
+
+ cur_scores, cur_indices = tf.math.top_k(cur_scores, k=self._num_instances)
+ cur_classes = tf.gather(cur_classes, cur_indices)
+ cur_masks = tf.gather(cur_masks, cur_indices, axis=2)
+
+ num_proposals = self._num_instances
+
+ # Find the proposal ID for each pixel.
+ cur_mask_ids = tf.argmax( # h x w
+ cur_masks * cur_scores[None, None, :],
+ axis=2,
+ output_type=tf.int32)
+ # Compute original areas from binary mask.
+ original_areas = tf.reduce_sum(
+ tf.cast(cur_masks > 0.5, tf.float32), axis=[0, 1]) # n
+ # Compute mask areas from the proposal ID mask.
+ proposal_masks = tf.range( # h x w x n
+ num_proposals, dtype=tf.int32)[None, None, :] == cur_mask_ids[..., None]
+ mask_areas = tf.reduce_sum( # n
+ tf.cast(proposal_masks, tf.float32), axis=[0, 1])
+ # Compute valid masks to filter results.
+ valid_masks = tf.logical_and(mask_areas > self._small_area_threshold, # n
+ original_areas > 0)
+ valid_masks = tf.logical_and( # n
+ valid_masks, mask_areas > self._overlap_threshold * original_areas)
+ # Compute category mask and instance mask.
+ category_mask = tf.gather(cur_classes * tf.cast(valid_masks, tf.int32),
+ cur_mask_ids)
+ is_thing_mask = tf.gather(self._is_thing,
+ cur_classes * tf.cast(valid_masks, tf.int32))
+ instance_mask = tf.gather(is_thing_mask, cur_mask_ids)
+ instance_mask = tf.where(instance_mask, cur_mask_ids + 1,
+ tf.zeros_like(cur_mask_ids))
+ return {'category_mask': category_mask, 'instance_mask': instance_mask}
+
+ def __call__(self,
+ inputs: Dict[str, tf.Tensor],
+ images_info: Optional[tf.Tensor] = None) -> Dict[str, tf.Tensor]:
+ batched_scores = tf.cast(inputs['confidence'], dtype=tf.float32)
+ batched_classes = tf.cast(inputs['classes'], dtype=tf.int32)
+
+ batched_masks = tf.cast(inputs['mask_proposal_logits'], dtype=tf.float32)
+ # For on-device run, we use the following codes to speed up.
+ if self._use_hardware_optimization:
+ batched_masks = activations.hard_sigmoid(batched_masks)
+ # Note that we assume batch size is always 1.
+ panoptic_masks = self._generate_panoptic_masks(batched_scores[0],
+ batched_classes[0],
+ batched_masks[0])
+ for k, v in panoptic_masks.items():
+ panoptic_masks[k] = v[None]
+ panoptic_masks[k].set_shape((1, *self._output_size))
+ return panoptic_masks
+
+ if self._rescale_predictions and images_info is not None:
+ batched_masks = self._resize_and_pad_masks(batched_masks, images_info)
+ else:
+ batched_masks = tf.image.resize(batched_masks, self._output_size,
+ 'bilinear')
+ batched_masks = tf.nn.sigmoid(batched_masks)
+ panoptic_masks = tf.map_fn(
+ fn=lambda x: self._generate_panoptic_masks(x[0], x[1], x[2]),
+ elems=(batched_scores, batched_classes, batched_masks),
+ fn_output_signature={
+ 'category_mask': tf.int32,
+ 'instance_mask': tf.int32
+ },
+ parallel_iterations=32)
+
+ return panoptic_masks
+
+ def get_config(self):
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config):
+ return cls(**config)
diff --git a/official/projects/maskconver/modeling/maskconver_model.py b/official/projects/maskconver/modeling/maskconver_model.py
new file mode 100644
index 00000000000..743acc8785f
--- /dev/null
+++ b/official/projects/maskconver/modeling/maskconver_model.py
@@ -0,0 +1,171 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic Segmentation model."""
+
+from typing import Mapping, Union, Any, Dict, Optional, List
+
+import tensorflow as tf, tf_keras
+
+layers = tf_keras.layers
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MaskConverModel(tf_keras.Model):
+ """A MaskConver class model."""
+
+ def __init__(
+ self,
+ backbone: tf_keras.Model,
+ decoder: tf_keras.Model,
+ # panoptic_fpn_fusion: tf_keras.layers.Layer,
+ class_head: tf_keras.layers.Layer,
+ embedding_head: tf_keras.layers.Layer,
+ per_pixel_embeddings_head: tf_keras.layers.Layer,
+ mlp_embedding_head: tf_keras.layers.Layer,
+ proposal_generator: tf_keras.layers.Layer,
+ panoptic_generator: Optional[tf_keras.layers.Layer] = None,
+ level: int = 3,
+ padded_output_size: Optional[List[int]] = None,
+ score_threshold: float = 0.1,
+ l2_regularizer: Optional[Any] = None,
+ embedding_size: int = 256,
+ num_classes: int = 201,
+ **kwargs):
+ """MaskConver initialization function.
+
+ Args:
+ backbone: a backbone network.
+ decoder: a decoder network. E.g. FPN.
+ # panoptic_fpn_fusion: a panoptic_fpn_fusion layer.
+ class_head: class head.
+ embedding_head: embedding head.
+ per_pixel_embeddings_head: per_pixel_embeddings_head.
+ mlp_embedding_head: mlp embedding head.
+ proposal_generator: proposal_generator.
+ panoptic_generator: panoptic generator.
+ level: int.
+ padded_output_size: padded output size. GPU or CPU only.
+ score_threshold: score threshold, used for filtering.
+ l2_regularizer: l2 regularizer.
+ embedding_size: `int`, embedding size.
+ num_classes: `int`, the total number of classes.
+ **kwargs: keyword arguments to be passed.
+ """
+ super(MaskConverModel, self).__init__(**kwargs)
+ self._config_dict = {
+ 'backbone': backbone,
+ 'decoder': decoder,
+ 'class_head': class_head,
+ 'embedding_head': embedding_head,
+ 'mlp_embedding_head': mlp_embedding_head,
+ 'proposal_generator': proposal_generator,
+ 'level': level,
+ 'padded_output_size': padded_output_size,
+ 'score_threshold': score_threshold,
+ 'per_pixel_embeddings_head': per_pixel_embeddings_head,
+ }
+ self.backbone = backbone
+ self.decoder = decoder
+ self.class_head = class_head
+ self.embedding_head = embedding_head
+ self.embedding_size = embedding_size
+ self.mlp = mlp_embedding_head
+ self.proposal_generator = proposal_generator
+ self.panoptic_generator = panoptic_generator
+ self.num_classes = num_classes
+
+ self.level = level
+ self.padded_output_size = padded_output_size
+ self.score_threshold = score_threshold
+ self.per_pixel_embeddings_head = per_pixel_embeddings_head
+ self.class_embeddings = tf_keras.layers.Embedding(
+ num_classes,
+ self.embedding_size,
+ embeddings_regularizer=l2_regularizer)
+
+ def call(self, inputs: tf.Tensor, # pytype: disable=annotation-type-mismatch
+ image_info: Optional[tf.Tensor] = None,
+ box_indices: Optional[tf.Tensor] = None,
+ classes: Optional[tf.Tensor] = None,
+ training: bool = None # pyrefly: ignore[bad-function-definition]
+ ) -> Dict[str, Optional[Any]]:
+ backbone_features = self.backbone(inputs)
+
+ if self.decoder:
+ decoder_features = self.decoder(backbone_features)
+ else:
+ decoder_features = backbone_features
+
+ class_heatmaps = self.class_head((backbone_features, decoder_features),
+ training=training)
+ dense_mask_embeddings = self.embedding_head(
+ (backbone_features, decoder_features), training=training)
+ per_pixel_embeddings = self.per_pixel_embeddings_head(
+ (backbone_features, decoder_features), training=training)
+
+ if not training:
+ proposals = self.proposal_generator(class_heatmaps)
+ classes = proposals['classes']
+ confidence = proposals['confidence']
+ box_indices = proposals['embedding_indices']
+ _ = proposals['num_proposals']
+
+ mask_embeddings = tf.gather_nd(
+ dense_mask_embeddings, box_indices, batch_dims=1)
+ class_embeddings = self.class_embeddings(tf.maximum(classes, 0))
+ mask_embeddings = mask_embeddings * tf.cast(
+ class_embeddings, mask_embeddings.dtype)
+ mask_embeddings = self.mlp(mask_embeddings)
+
+ mask_proposal_logits = tf.einsum('bqc,bhwc->bhwq',
+ mask_embeddings,
+ per_pixel_embeddings)
+ mask_proposal_logits = tf.cast(mask_proposal_logits, tf.float32)
+
+ if not training:
+ outputs = {'classes': classes,
+ 'confidence': confidence, # pyrefly: ignore[unbound-name]
+ 'mask_proposal_logits': mask_proposal_logits,
+ 'class_heatmaps': class_heatmaps}
+ if self.panoptic_generator is not None:
+ panoptic_outputs = self.panoptic_generator(
+ outputs, images_info=image_info)
+ outputs.update({'panoptic_outputs': panoptic_outputs})
+ else:
+ outputs['mask_proposal_logits'] = tf.image.resize(
+ mask_proposal_logits, self.padded_output_size, 'bilinear')
+ else:
+ outputs = {'class_heatmaps': class_heatmaps,
+ 'mask_proposal_logits': mask_proposal_logits}
+ return outputs
+
+ @property
+ def checkpoint_items(
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(backbone=self.backbone,
+ class_head=self.class_head,
+ embedding_head=self.embedding_head,
+ per_pixel_embeddings_head=self.per_pixel_embeddings_head)
+ if self.decoder is not None:
+ items.update(decoder=self.decoder)
+ return items
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
diff --git a/official/projects/maskconver/modeling/multiscale_maskconver_model.py b/official/projects/maskconver/modeling/multiscale_maskconver_model.py
new file mode 100644
index 00000000000..42f516b543f
--- /dev/null
+++ b/official/projects/maskconver/modeling/multiscale_maskconver_model.py
@@ -0,0 +1,211 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic Segmentation MaskConver model."""
+
+from typing import Mapping, Union, Any, Dict, Optional, List
+
+import tensorflow as tf, tf_keras
+
+layers = tf_keras.layers
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MultiScaleMaskConverModel(tf_keras.Model):
+ """Multiscale MaskConver class model."""
+
+ def __init__(self,
+ backbone: tf_keras.Model,
+ decoder: tf_keras.Model,
+ mask_decoder: tf_keras.Model,
+ class_head: tf_keras.layers.Layer,
+ embedding_head: tf_keras.layers.Layer,
+ per_pixel_embeddings_head: tf_keras.layers.Layer,
+ mlp_embedding_head: tf_keras.layers.Layer,
+ panoptic_generator: Optional[tf_keras.layers.Layer] = None,
+ min_level: int = 3,
+ max_level: int = 5,
+ max_proposals: int = 100,
+ padded_output_size: Optional[List[int]] = None,
+ score_threshold: float = 0.1,
+ l2_regularizer: Optional[Any] = None,
+ embedding_size: int = 256,
+ num_classes: int = 201,
+ **kwargs):
+ """MaskConver initialization function."""
+ super().__init__(**kwargs)
+ self._config_dict = {
+ 'backbone': backbone,
+ 'decoder': decoder,
+ 'mask_decoder': mask_decoder,
+ 'class_head': class_head,
+ 'embedding_head': embedding_head,
+ 'mlp_embedding_head': mlp_embedding_head,
+ 'min_level': min_level,
+ 'max_level': max_level,
+ 'max_proposals': max_proposals,
+ 'padded_output_size': padded_output_size,
+ 'score_threshold': score_threshold,
+ 'per_pixel_embeddings_head': per_pixel_embeddings_head,
+ }
+ self.backbone = backbone
+ self.decoder = decoder
+ self.mask_decoder = mask_decoder
+ self.class_head = class_head
+ self.embedding_head = embedding_head
+ self.embedding_size = embedding_size
+ self.mlp = mlp_embedding_head
+ self.panoptic_generator = panoptic_generator
+ self.num_classes = num_classes
+
+ self.min_level = min_level
+ self.max_level = max_level
+ self.max_proposals = max_proposals
+ self.padded_output_size = padded_output_size
+ self.score_threshold = score_threshold
+ self.per_pixel_embeddings_head = per_pixel_embeddings_head
+ self.class_embeddings = tf_keras.layers.Embedding(
+ num_classes,
+ self.embedding_size,
+ embeddings_regularizer=l2_regularizer)
+
+ def process_heatmap(self, heatmap: tf.Tensor) -> tf.Tensor:
+ scoremap = tf.sigmoid(heatmap)
+ scoremap_max_pool = tf.nn.max_pool(
+ scoremap, ksize=3, strides=1, padding='SAME')
+ valid_mask = tf.abs(scoremap - scoremap_max_pool) < 1e-6
+ return scoremap * tf.cast(valid_mask, scoremap.dtype)
+
+ def call(self, inputs: tf.Tensor, # pytype: disable=annotation-type-mismatch
+ image_info: Optional[tf.Tensor] = None,
+ box_indices: Optional[tf.Tensor] = None,
+ classes: Optional[tf.Tensor] = None,
+ training: bool = None # pyrefly: ignore[bad-function-definition]
+ ) -> Dict[str, Optional[Any]]:
+ batch_size = tf.shape(inputs)[0]
+ backbone_features = self.backbone(inputs, training=training)
+
+ if self.decoder:
+ decoder_features = self.decoder(backbone_features)
+ if self.mask_decoder:
+ decoder_features2 = self.mask_decoder(backbone_features)
+ else:
+ decoder_features2 = decoder_features
+ else:
+ decoder_features = backbone_features
+
+ level_class_heatmaps = self.class_head(decoder_features, training=training)
+ level_dense_mask_embeddings = self.embedding_head(
+ decoder_features, training=training)
+ embedding_size = self.embedding_size
+
+ class_heatmaps = []
+ dense_mask_embeddings = []
+ class_scoremaps = []
+ for level in range(self.min_level, self.max_level + 1):
+ if not training:
+ class_scoremaps.append(
+ tf.reshape(
+ self.process_heatmap(level_class_heatmaps[level]),
+ [batch_size, -1],
+ )
+ )
+ class_heatmaps.append(
+ tf.reshape(
+ level_class_heatmaps[level], [batch_size, -1, self.num_classes]
+ )
+ )
+ dense_mask_embeddings.append(
+ tf.reshape(
+ level_dense_mask_embeddings[level],
+ [batch_size, -1, embedding_size],
+ )
+ )
+
+ class_heatmaps = tf.concat(class_heatmaps, axis=1)
+ dense_mask_embeddings = tf.concat(dense_mask_embeddings, axis=1)
+
+ per_pixel_embeddings = self.per_pixel_embeddings_head(
+ (backbone_features, decoder_features2), training=training # pyrefly: ignore[unbound-name]
+ )
+
+ if not training:
+ class_scoremaps = tf.concat(class_scoremaps, axis=1)
+ confidence, top_indices = tf.nn.top_k(
+ class_scoremaps, k=self.max_proposals
+ )
+ box_indices = top_indices // self.num_classes
+ classes = top_indices % self.num_classes
+
+ mask_embeddings = tf.gather(
+ dense_mask_embeddings, box_indices, batch_dims=1
+ )
+ class_embeddings = tf.cast(
+ self.class_embeddings(tf.maximum(classes, 0)), mask_embeddings.dtype
+ )
+ mask_embeddings_inputs = mask_embeddings + class_embeddings
+ mask_embeddings = self.mlp(mask_embeddings_inputs)
+
+ mask_proposal_logits = tf.einsum(
+ 'bqc,bhwc->bhwq', mask_embeddings, per_pixel_embeddings
+ )
+ mask_proposal_logits = tf.cast(mask_proposal_logits, tf.float32)
+
+ if not training:
+ outputs = {
+ 'classes': classes,
+ 'confidence': confidence, # pyrefly: ignore[unbound-name]
+ 'mask_embeddings': mask_embeddings,
+ 'mask_proposal_logits': mask_proposal_logits,
+ 'class_heatmaps': class_heatmaps,
+ }
+ if self.panoptic_generator is not None:
+ panoptic_outputs = self.panoptic_generator(
+ outputs, images_info=image_info
+ )
+ outputs.update({'panoptic_outputs': panoptic_outputs})
+ else:
+ outputs['mask_proposal_logits'] = tf.image.resize(
+ mask_proposal_logits, self.padded_output_size, 'bilinear'
+ )
+ else:
+ outputs = {
+ 'class_heatmaps': class_heatmaps,
+ 'mask_proposal_logits': mask_proposal_logits,
+ 'mask_embeddings': mask_embeddings,
+ }
+ return outputs
+
+ @property
+ def checkpoint_items(
+ self,
+ ) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(
+ backbone=self.backbone,
+ heads=self.heads,
+ class_head=self.class_head,
+ embedding_head=self.embedding_head,
+ per_pixel_embeddings_head=self.per_pixel_embeddings_head,
+ )
+ if self.decoder is not None:
+ items.update(decoder=self.decoder)
+ return items
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
diff --git a/official/projects/maskconver/serving/export_saved_model.py b/official/projects/maskconver/serving/export_saved_model.py
new file mode 100644
index 00000000000..31438f576eb
--- /dev/null
+++ b/official/projects/maskconver/serving/export_saved_model.py
@@ -0,0 +1,125 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+r"""Maskconver model export binary for serving/inference.
+
+To export a trained checkpoint in saved_model format (shell script):
+
+CHECKPOINT_PATH = XX
+EXPORT_DIR_PATH = XX
+CONFIG_FILE_PATH = XX
+export_saved_model --export_dir=${EXPORT_DIR_PATH}/ \
+ --checkpoint_path=${CHECKPOINT_PATH} \
+ --config_file=${CONFIG_FILE_PATH} \
+ --batch_size=2 \
+ --input_image_size=224,224
+To serve (python):
+export_dir_path = XX
+input_type = XX
+input_images = XX
+imported = tf.saved_model.load(export_dir_path)
+model_fn = imported.signatures['serving_default']
+output = model_fn(input_images)
+"""
+
+from absl import app
+from absl import flags
+import tensorflow as tf, tf_keras
+
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.maskconver.configs import maskconver as maskconver_cfg # pylint: disable=unused-import
+from official.projects.maskconver.configs import multiscale_maskconver as multiscale_maskconver_cfg # pylint: disable=unused-import
+from official.projects.maskconver.modeling import factory
+from official.projects.maskconver.modeling import fpn # pylint: disable=unused-import
+from official.projects.maskconver.serving import maskconver
+from official.projects.maskconver.tasks import maskconver as maskconver_tasks # pylint: disable=unused-import
+from official.projects.maskconver.tasks import multiscale_maskconver as multiscale_maskconver_tasks # pylint: disable=unused-import
+from official.vision.serving import export_saved_model_lib
+
+FLAGS = flags.FLAGS
+
+flags.DEFINE_string('experiment', 'maskconver_coco',
+ 'experiment type, e.g. maskconver_coco')
+flags.DEFINE_string('export_dir', None, 'The export directory.')
+flags.DEFINE_string('checkpoint_path', None, 'Checkpoint path.')
+flags.DEFINE_multi_string(
+ 'config_file',
+ default=None,
+ help='YAML/JSON files which specifies overrides. The override order '
+ 'follows the order of args. Note that each file '
+ 'can be used as an override template to override the default parameters '
+ 'specified in Python. If the same parameter is specified in both '
+ '`--config_file` and `--params_override`, `config_file` will be used '
+ 'first, followed by params_override.')
+flags.DEFINE_string(
+ 'params_override', '',
+ 'The JSON/YAML file or string which specifies the parameter to be overriden'
+ ' on top of `config_file` template.')
+flags.DEFINE_integer('batch_size', None, 'The batch size.')
+flags.DEFINE_string('input_type', 'image_tensor',
+ 'One of `image_tensor`, `image_bytes`, `tf_example`.')
+flags.DEFINE_string(
+ 'input_image_size', '640,640',
+ 'The comma-separated string of two integers representing the height,width '
+ 'of the input to the model.')
+
+
+def main(_):
+
+ params = exp_factory.get_exp_config(FLAGS.experiment)
+ for config_file in FLAGS.config_file or []:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=False)
+ if FLAGS.params_override:
+ params = hyperparams.override_params_dict(
+ params, FLAGS.params_override, is_strict=False)
+
+ params.validate()
+ params.lock()
+
+ input_image_size = [int(x) for x in FLAGS.input_image_size.split(',')]
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[FLAGS.batch_size] + input_image_size + [3])
+ if FLAGS.experiment == 'maskconver_coco':
+ model = factory.build_maskconver_model(
+ input_specs=input_specs,
+ model_config=params.task.model,
+ l2_regularizer=None)
+ elif FLAGS.experiment == 'multiscale_maskconver_coco':
+ model = factory.build_multiscale_maskconver_model(
+ input_specs=input_specs,
+ model_config=params.task.model,
+ l2_regularizer=None)
+
+ export_module = maskconver.MaskConverModule(
+ params=params,
+ model=model, # pyrefly: ignore[unbound-name]
+ batch_size=FLAGS.batch_size,
+ input_image_size=[int(x) for x in FLAGS.input_image_size.split(',')],
+ input_type=FLAGS.input_type,
+ num_channels=3)
+ export_saved_model_lib.export_inference_graph(
+ input_type=FLAGS.input_type,
+ batch_size=FLAGS.batch_size,
+ input_image_size=input_image_size,
+ params=params,
+ checkpoint_path=FLAGS.checkpoint_path,
+ export_dir=FLAGS.export_dir,
+ export_module=export_module,
+ export_checkpoint_subdir='checkpoint',
+ export_saved_model_subdir='saved_model')
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/maskconver/serving/export_tflite.py b/official/projects/maskconver/serving/export_tflite.py
new file mode 100644
index 00000000000..837d1d49c4f
--- /dev/null
+++ b/official/projects/maskconver/serving/export_tflite.py
@@ -0,0 +1,122 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+r"""Binary to convert a saved model to tflite model.
+
+It requires a SavedModel exported using export_saved_model.py with batch size 1
+and input type `tflite`, and using the same config file used for exporting saved
+model. It includes optional post-training quantization. When using integer
+quantization, calibration steps need to be provided to calibrate model input.
+
+To convert a SavedModel to a TFLite model:
+
+EXPERIMENT_TYPE = XX
+TFLITE_PATH = XX
+SAVED_MOODEL_DIR = XX
+CONFIG_FILE = XX
+export_tflite --experiment=${EXPERIMENT_TYPE} \
+ --saved_model_dir=${SAVED_MOODEL_DIR} \
+ --tflite_path=${TFLITE_PATH} \
+ --config_file=${CONFIG_FILE} \
+ --quant_type=fp16 \
+ --calibration_steps=500
+"""
+from absl import app
+from absl import flags
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.maskconver.configs import maskconver as maskconver_cfg # pylint: disable=unused-import
+from official.projects.maskconver.tasks import maskconver as maskconver_task
+from official.vision.serving import export_tflite_lib
+
+FLAGS = flags.FLAGS
+
+_EXPERIMENT = flags.DEFINE_string(
+ 'experiment',
+ default=None,
+ help='Experiment type, e.g. maskconver_coco',
+ required=True)
+_CONFIG_FILE = flags.DEFINE_multi_string(
+ 'config_file',
+ default='',
+ help='YAML/JSON files which specifies overrides. The override order '
+ 'follows the order of args. Note that each file '
+ 'can be used as an override template to override the default parameters '
+ 'specified in Python. If the same parameter is specified in both '
+ '`--config_file` and `--params_override`, `config_file` will be used '
+ 'first, followed by params_override.')
+_PARAMS_OVERRIDE = flags.DEFINE_string(
+ 'params_override', '',
+ 'The JSON/YAML file or string which specifies the parameter to be overriden'
+ ' on top of `config_file` template.')
+_SAVED_MODEL_DIR = flags.DEFINE_string(
+ 'saved_model_dir', None, 'The directory to the saved model.', required=True)
+_TFLITE_PATH = flags.DEFINE_string(
+ 'tflite_path', None, 'The path to the output tflite model.', required=True)
+_QUANT_TYPE = flags.DEFINE_string(
+ 'quant_type',
+ default=None,
+ help='Post training quantization type. Support `int8_fallback`, '
+ '`int8_full_fp32_io`, `int8_full`, `fp16`, `qat`, `qat_fp32_io`, '
+ '`int8_full_int8_io` and `default`. See '
+ 'https://www.tensorflow.org/lite/performance/post_training_quantization '
+ 'for more details.')
+_CALIBRATION_STEPS = flags.DEFINE_integer(
+ 'calibration_steps', 500,
+ 'The number of calibration steps for integer model.')
+_DENYLISTED_OPS = flags.DEFINE_string(
+ 'denylisted_ops', '', 'The comma-separated string of ops '
+ 'that are excluded from integer quantization. The name of '
+ 'ops should be all capital letters, such as CAST or GREATER.'
+ 'This is useful to exclude certains ops that affects quality or latency.')
+
+
+def main(_) -> None:
+ params = exp_factory.get_exp_config(_EXPERIMENT.value)
+ if _CONFIG_FILE.value is not None:
+ for config_file in _CONFIG_FILE.value:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=True)
+ if _PARAMS_OVERRIDE.value:
+ params = hyperparams.override_params_dict(
+ params, _PARAMS_OVERRIDE.value, is_strict=True)
+
+ params.validate()
+ params.lock()
+
+ logging.info('Converting SavedModel from %s to TFLite model...',
+ _SAVED_MODEL_DIR.value)
+
+ denylisted_ops = None
+ if _DENYLISTED_OPS.value:
+ denylisted_ops = list(_DENYLISTED_OPS.value.split(','))
+ tflite_model = export_tflite_lib.convert_tflite_model(
+ saved_model_dir=_SAVED_MODEL_DIR.value,
+ quant_type=_QUANT_TYPE.value,
+ params=params,
+ task=maskconver_task.PanopticMaskRCNNTask(params.task),
+ calibration_steps=_CALIBRATION_STEPS.value,
+ denylisted_ops=denylisted_ops)
+
+ with tf.io.gfile.GFile(_TFLITE_PATH.value, 'wb') as fw:
+ fw.write(tflite_model)
+
+ logging.info('TFLite model converted and saved to %s.', _TFLITE_PATH.value)
+
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/maskconver/serving/maskconver.py b/official/projects/maskconver/serving/maskconver.py
new file mode 100644
index 00000000000..e94350de6e2
--- /dev/null
+++ b/official/projects/maskconver/serving/maskconver.py
@@ -0,0 +1,91 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Maskconver input and model functions for serving/inference."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.maskconver.modeling import factory
+from official.vision.ops import preprocess_ops
+from official.vision.serving import export_base
+
+
+class MaskConverModule(export_base.ExportModule):
+ """MaskConver Module."""
+
+ def _build_model(self):
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[self._batch_size] + self._input_image_size + [3])
+
+ return factory.build_maskconver_model(
+ input_specs=input_specs,
+ model_config=self.params.task.model,
+ l2_regularizer=None)
+
+ def _build_inputs(self, image):
+ """Builds MaskConver model inputs for serving."""
+
+ # Normalizes image with mean and std pixel values.
+ image = preprocess_ops.normalize_image(
+ image, offset=preprocess_ops.MEAN_RGB, scale=preprocess_ops.STDDEV_RGB)
+
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ self._input_image_size,
+ padded_size=self._input_image_size,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0)
+ return image, image_info
+
+ def serve(self, images):
+ """Cast image to float and run inference.
+
+ Args:
+ images: uint8 Tensor of shape [batch_size, None, None, 3]
+ Returns:
+ Tensor holding classification output logits.
+ """
+ # Skip image preprocessing when input_type is tflite so it is compatible
+ # with TFLite quantization.
+ image_info = None
+ if self._input_type != "tflite":
+ with tf.device("cpu:0"):
+ images = tf.cast(images, dtype=tf.float32)
+ images_spec = tf.TensorSpec(
+ shape=self._input_image_size + [3], dtype=tf.float32)
+ image_info_spec = tf.TensorSpec(shape=[4, 2], dtype=tf.float32)
+
+ images, image_info = tf.nest.map_structure(
+ tf.identity,
+ tf.map_fn(
+ self._build_inputs,
+ elems=images,
+ fn_output_signature=(images_spec, image_info_spec),
+ parallel_iterations=32))
+
+ outputs = self.inference_step(images, image_info)
+
+ if "panoptic_outputs" in outputs:
+ outputs.update({
+ "panoptic_category_mask":
+ outputs["panoptic_outputs"]["category_mask"],
+ "panoptic_instance_mask":
+ outputs["panoptic_outputs"]["instance_mask"],
+ })
+ del outputs["panoptic_outputs"]
+
+ if image_info is not None:
+ outputs.update({"image_info": image_info})
+
+ return outputs
diff --git a/official/projects/maskconver/tasks/__init__.py b/official/projects/maskconver/tasks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/maskconver/tasks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/maskconver/tasks/maskconver.py b/official/projects/maskconver/tasks/maskconver.py
new file mode 100644
index 00000000000..72fd1ebfa3c
--- /dev/null
+++ b/official/projects/maskconver/tasks/maskconver.py
@@ -0,0 +1,641 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic MaskRCNN task definition."""
+from typing import Any, Dict, List, Mapping, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import task_factory
+from official.projects.centernet.ops import loss_ops
+from official.projects.maskconver.configs import maskconver as exp_cfg
+from official.projects.maskconver.dataloaders import maskconver_segmentation_input
+from official.projects.maskconver.dataloaders import panoptic_maskrcnn_input
+from official.projects.maskconver.losses import maskconver_losses
+from official.projects.maskconver.modeling import factory
+from official.projects.volumetric_models.losses import segmentation_losses as volumeteric_segmentation_losses
+from official.vision.dataloaders import input_reader_factory
+from official.vision.dataloaders import segmentation_input
+from official.vision.evaluation import panoptic_quality_evaluator
+from official.vision.evaluation import segmentation_metrics
+from official.vision.tasks import maskrcnn
+from official.vision.tasks import semantic_segmentation
+
+
+@task_factory.register_task_cls(exp_cfg.MaskConverTask)
+class PanopticMaskRCNNTask(maskrcnn.MaskRCNNTask):
+
+ """A single-replica view of training procedure.
+
+ Panoptic Mask R-CNN task provides artifacts for training/evalution procedures,
+ including loading/iterating over Datasets, initializing the model, calculating
+ the loss, post-processing, and customized metrics with reduction.
+ """
+
+ def build_model(self) -> tf_keras.Model:
+ """Build Panoptic Mask R-CNN model."""
+
+ tf_keras.utils.set_random_seed(0)
+ tf.config.experimental.enable_op_determinism()
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + self.task_config.model.input_size)
+
+ l2_weight_decay = self.task_config.losses.l2_weight_decay
+ # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
+ # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
+ # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
+ l2_regularizer = (tf_keras.regularizers.l2(
+ l2_weight_decay / 2.0) if l2_weight_decay else None)
+
+ model = factory.build_maskconver_model(
+ input_specs=input_specs,
+ model_config=self.task_config.model,
+ l2_regularizer=l2_regularizer)
+ return model
+
+ def build_inputs(
+ self,
+ params: exp_cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None
+ ) -> tf.data.Dataset:
+ """Build input dataset."""
+ decoder_cfg = params.decoder.get()
+ if params.decoder.type == 'simple_decoder':
+ decoder = panoptic_maskrcnn_input.TfExampleDecoder(
+ regenerate_source_id=decoder_cfg.regenerate_source_id,
+ mask_binarize_threshold=decoder_cfg.mask_binarize_threshold,
+ include_panoptic_masks=decoder_cfg.include_panoptic_masks,
+ panoptic_category_mask_key=decoder_cfg.panoptic_category_mask_key,
+ panoptic_instance_mask_key=decoder_cfg.panoptic_instance_mask_key)
+ else:
+ raise ValueError('Unknown decoder type: {}!'.format(params.decoder.type))
+
+ parser = panoptic_maskrcnn_input.Parser(
+ output_size=self.task_config.model.input_size[:2],
+ min_level=self.task_config.model.min_level,
+ max_level=self.task_config.model.max_level,
+ num_scales=self.task_config.model.anchor.num_scales,
+ aspect_ratios=self.task_config.model.anchor.aspect_ratios,
+ anchor_size=self.task_config.model.anchor.anchor_size,
+ dtype=params.dtype,
+ rpn_match_threshold=params.parser.rpn_match_threshold,
+ rpn_unmatched_threshold=params.parser.rpn_unmatched_threshold,
+ rpn_batch_size_per_im=params.parser.rpn_batch_size_per_im,
+ rpn_fg_fraction=params.parser.rpn_fg_fraction,
+ aug_rand_hflip=params.parser.aug_rand_hflip,
+ aug_scale_min=params.parser.aug_scale_min,
+ aug_scale_max=params.parser.aug_scale_max,
+ skip_crowd_during_training=params.parser.skip_crowd_during_training,
+ max_num_instances=self.task_config.model.num_instances,
+ mask_crop_size=params.parser.mask_crop_size,
+ segmentation_resize_eval_groundtruth=params.parser
+ .segmentation_resize_eval_groundtruth,
+ segmentation_groundtruth_padded_size=params.parser
+ .segmentation_groundtruth_padded_size,
+ segmentation_ignore_label=params.parser.segmentation_ignore_label,
+ panoptic_ignore_label=params.parser.panoptic_ignore_label,
+ include_panoptic_masks=params.parser.include_panoptic_masks,
+ num_panoptic_categories=self.task_config.model.num_classes,
+ num_thing_categories=self.task_config.model.num_thing_classes,
+ level=self.task_config.model.level,
+ gaussian_iou=params.parser.gaussaian_iou,
+ aug_type=params.parser.aug_type,
+ max_num_stuff_centers=params.parser.max_num_stuff_centers)
+
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training))
+ dataset = reader.read(input_context=input_context)
+
+ return dataset
+
+ def build_losses(self,
+ outputs: Mapping[str, Any],
+ labels: Mapping[str, Any],
+ aux_losses: Optional[Any] = None,
+ step=None) -> Dict[str, tf.Tensor]:
+ """Build Panoptic Mask R-CNN losses."""
+ loss_params = self._task_config.losses
+
+ batch_size = tf.cast(tf.shape(labels['num_instances'])[0], tf.float32)
+ center_loss_fn = maskconver_losses.PenaltyReducedLogisticFocalLoss(
+ alpha=loss_params.alpha, beta=loss_params.beta)
+ mask_loss_fn = maskconver_losses.PenaltyReducedLogisticFocalLoss()
+
+ # Calculate center heatmap loss
+ # TODO(arashwan): add valid weights.
+ # output_unpad_image_shapes = labels['image_info'][:, 0, :]
+ # valid_anchor_weights = loss_ops.get_valid_anchor_weights_in_flattened_image( # pylint: disable=line-too-long
+ # output_unpad_image_shapes, h, w)
+ # valid_anchor_weights = tf.expand_dims(valid_anchor_weights, 2)
+
+ true_flattened_ct_heatmap = loss_ops.flatten_spatial_dimensions(
+ labels['panoptic_heatmaps'])
+ true_flattened_ct_heatmap = tf.cast(true_flattened_ct_heatmap, tf.float32)
+
+ pred_flattened_ct_heatmap = loss_ops.flatten_spatial_dimensions(
+ outputs['class_heatmaps'])
+ pred_flattened_ct_heatmap = tf.cast(pred_flattened_ct_heatmap, tf.float32)
+ center_padding_mask = 1 - labels['panoptic_padding_mask'][:, :, :, None]
+ center_padding_mask = tf.image.resize(
+ center_padding_mask, tf.shape(
+ labels['panoptic_heatmaps'])[1:3], method='nearest')
+ center_padding_mask = tf.maximum(center_padding_mask, 0.0)
+ center_padding_mask = center_padding_mask * tf.ones_like(labels['panoptic_heatmaps'])
+ weights_flattened_mask = loss_ops.flatten_spatial_dimensions(
+ center_padding_mask)
+ center_loss = center_loss_fn(
+ target_tensor=true_flattened_ct_heatmap,
+ prediction_tensor=pred_flattened_ct_heatmap,
+ weights=weights_flattened_mask)
+
+ center_loss = tf.reduce_sum(
+ center_loss / (labels['num_instances'][:, None, None] + 1.0)) / batch_size
+
+ gt_masks = labels['panoptic_masks']
+ gt_mask_weights = labels['panoptic_mask_weights'][:, None, None, :] * tf.ones_like(gt_masks)
+ panoptic_padding_mask = labels['panoptic_padding_mask'][:, :, :, None] * tf.ones_like(gt_masks)
+
+ true_flattened_masks = loss_ops.flatten_spatial_dimensions(
+ gt_masks)
+ true_flattened_ct_heatmap = tf.cast(true_flattened_ct_heatmap, tf.float32)
+ predicted_masks = tf.cast(outputs['mask_proposal_logits'], tf.float32)
+ predicted_masks = tf.image.resize(
+ predicted_masks, tf.shape(gt_masks)[1:3], method='bilinear')
+ pred_flattened_masks = loss_ops.flatten_spatial_dimensions(predicted_masks)
+ mask_loss = tf.cast(0.0, tf.float32)
+ mask_loss_fn = tf_keras.losses.BinaryCrossentropy(
+ from_logits=True,
+ label_smoothing=0.0,
+ axis=-1,
+ reduction=tf_keras.losses.Reduction.NONE,
+ name='binary_crossentropy')
+ mask_weights = tf.reshape(
+ tf.cast(true_flattened_masks >= 0, tf.float32),
+ [-1, 1]) * tf.reshape(gt_mask_weights, [-1, 1]) * tf.reshape(
+ (1 - panoptic_padding_mask), [-1, 1])
+ mask_loss = mask_loss_fn(
+ tf.reshape(gt_masks, [-1, 1]),
+ tf.reshape(pred_flattened_masks, [-1, 1]),
+ sample_weight=mask_weights)
+ mask_loss = tf.reduce_sum(mask_loss) / (tf.reduce_sum(mask_weights) + 1.0)
+
+ # Dice loss
+ _, h, w, _ = gt_masks.get_shape().as_list()
+ masked_predictions = tf.sigmoid(predicted_masks) * gt_mask_weights * (1 - panoptic_padding_mask)
+ masked_gt_masks = gt_masks * gt_mask_weights * (1 - panoptic_padding_mask)
+
+ masked_predictions = tf.transpose(masked_predictions, [0, 3, 1, 2])
+ masked_predictions = tf.reshape(masked_predictions, [-1, h, w, 1])
+ masked_gt_masks = tf.transpose(masked_gt_masks, [0, 3, 1, 2])
+ masked_gt_masks = tf.reshape(masked_gt_masks, [-1, h, w, 1])
+
+ dice_loss_fn = volumeteric_segmentation_losses.SegmentationLossDiceScore(
+ metric_type='adaptive', axis=(2, 3))
+ dice_loss = dice_loss_fn(logits=masked_predictions, labels=masked_gt_masks)
+
+ total_loss = (center_loss + loss_params.mask_weight * mask_loss + loss_params.mask_weight * dice_loss)
+
+ if aux_losses:
+ total_loss += tf.add_n(aux_losses)
+
+ total_loss = loss_params.loss_weight * total_loss
+
+ losses = {'total_loss': total_loss,
+ 'mask_loss': mask_loss,
+ 'center_loss': center_loss,
+ 'dice_loss': dice_loss}
+ return losses
+
+ def build_metrics(self, training: bool = True) -> List[
+ tf_keras.metrics.Metric]:
+ """Build detection metrics."""
+ metrics = []
+ if training:
+ metric_names = [
+ 'total_loss',
+ 'center_loss',
+ 'mask_loss',
+ 'dice_loss',
+ ]
+ for name in metric_names:
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
+
+ else:
+ pq_config = self.task_config.panoptic_quality_evaluator
+ self.panoptic_quality_metric = (
+ panoptic_quality_evaluator.PanopticQualityEvaluator(
+ num_categories=pq_config.num_categories,
+ ignored_label=pq_config.ignored_label,
+ max_instances_per_category=pq_config.max_instances_per_category,
+ offset=pq_config.offset,
+ is_thing=pq_config.is_thing,
+ rescale_predictions=pq_config.rescale_predictions))
+ return metrics
+
+ def train_step(self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
+ """Does forward and backward.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ images, labels = inputs
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+
+ with tf.GradientTape() as tape:
+ outputs = model(
+ images,
+ box_indices=labels['panoptic_box_indices'],
+ classes=labels['panoptic_classes'],
+ training=True)
+ outputs = tf.nest.map_structure(
+ lambda x: tf.cast(x, tf.float32), outputs)
+
+ # Computes per-replica loss.
+ losses = self.build_losses(
+ outputs=outputs,
+ labels=labels,
+ aux_losses=model.losses,
+ step=optimizer.iterations)
+ scaled_loss = losses['total_loss'] / num_replicas
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scales back gradient when LossScaleOptimizer is used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ logs = {self.loss: losses['total_loss']}
+
+ if metrics:
+ for m in metrics:
+ m.update_state(losses[m.name])
+
+ return logs
+
+ def validation_step(self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ images, labels = inputs
+
+ outputs = model(
+ images,
+ image_info=labels['image_info'],
+ training=False)
+
+ logs = {self.loss: 0}
+
+ pq_metric_labels = {
+ 'category_mask':
+ labels['groundtruths']['gt_panoptic_category_mask'],
+ 'instance_mask':
+ labels['groundtruths']['gt_panoptic_instance_mask'],
+ 'image_info': labels['image_info']
+ }
+ logs.update({
+ self.panoptic_quality_metric.name:
+ (pq_metric_labels, outputs['panoptic_outputs'])})
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ if state is None:
+ self.panoptic_quality_metric.reset_states()
+ state = [self.panoptic_quality_metric]
+
+ self.panoptic_quality_metric.update_state(
+ step_outputs[self.panoptic_quality_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ step_outputs[self.panoptic_quality_metric.name][1]) # pyrefly: ignore[unsupported-operation]
+
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ result = {}
+
+ report_per_class_metrics = (
+ self.task_config.panoptic_quality_evaluator.report_per_class_metrics)
+ panoptic_quality_results = self.panoptic_quality_metric.result()
+ for k, value in panoptic_quality_results.items():
+ if k.endswith('per_class'):
+ if report_per_class_metrics:
+ for i, per_class_value in enumerate(value):
+ metric_key = 'panoptic_quality/{}/class_{}'.format(k, i)
+ result[metric_key] = per_class_value
+ else:
+ continue
+ else:
+ result['panoptic_quality/{}'.format(k)] = value
+ return result
+
+
+@task_factory.register_task_cls(exp_cfg.MaskConverSegTask)
+class MaskConverSegmentation(semantic_segmentation.SemanticSegmentationTask):
+
+ """A single-replica view of training procedure.
+
+ MaskConver task provides artifacts for training/evalution procedures,
+ including loading/iterating over Datasets, initializing the model, calculating
+ the loss, post-processing, and customized metrics with reduction.
+ """
+
+ def build_model(self) -> tf_keras.Model:
+ """Build maskconver model."""
+
+ tf_keras.utils.set_random_seed(0)
+ tf.config.experimental.enable_op_determinism()
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + self.task_config.model.input_size)
+
+ l2_weight_decay = self.task_config.losses.l2_weight_decay
+ # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
+ # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
+ # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
+ l2_regularizer = (tf_keras.regularizers.l2(
+ l2_weight_decay / 2.0) if l2_weight_decay else None)
+
+ model = factory.build_maskconver_model(
+ input_specs=input_specs,
+ model_config=self.task_config.model,
+ l2_regularizer=l2_regularizer,
+ segmentation_inference=True)
+ return model
+
+ def build_inputs(self,
+ params: exp_cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None):
+ """Builds classification input."""
+
+ ignore_label = self.task_config.losses.ignore_label
+
+ decoder = segmentation_input.Decoder()
+
+ parser = maskconver_segmentation_input.Parser(
+ output_size=params.output_size,
+ num_classes=self.task_config.model.num_classes,
+ crop_size=params.crop_size,
+ ignore_label=ignore_label,
+ resize_eval_groundtruth=params.resize_eval_groundtruth,
+ groundtruth_padded_size=params.groundtruth_padded_size,
+ aug_scale_min=params.aug_scale_min,
+ aug_scale_max=params.aug_scale_max,
+ aug_rand_hflip=params.aug_rand_hflip,
+ preserve_aspect_ratio=params.preserve_aspect_ratio,
+ level=self.task_config.model.level,
+ aug_type=params.aug_type,
+ max_num_stuff_centers=params.max_num_stuff_centers,
+ dtype=params.dtype)
+
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training))
+
+ dataset = reader.read(input_context=input_context)
+
+ return dataset
+
+ def build_losses(self,
+ outputs: Mapping[str, Any],
+ labels: Mapping[str, Any],
+ aux_losses: Optional[Any] = None,
+ step=None) -> Dict[str, tf.Tensor]:
+ """Build Panoptic Mask R-CNN losses."""
+ loss_params = self._task_config.losses
+
+ # b, h, w, c = outputs['class_heatmaps'].get_shape().as_list()
+ batch_size = tf.cast(tf.shape(labels['num_instances'])[0], tf.float32)
+ center_loss_fn = maskconver_losses.PenaltyReducedLogisticFocalLoss(
+ alpha=loss_params.alpha, beta=loss_params.beta)
+ mask_loss_fn = maskconver_losses.PenaltyReducedLogisticFocalLoss()
+
+ true_flattened_ct_heatmap = loss_ops.flatten_spatial_dimensions(
+ labels['seg_ct_heatmaps'])
+ true_flattened_ct_heatmap = tf.cast(true_flattened_ct_heatmap, tf.float32)
+
+ pred_flattened_ct_heatmap = loss_ops.flatten_spatial_dimensions(
+ outputs['class_heatmaps'])
+ pred_flattened_ct_heatmap = tf.cast(pred_flattened_ct_heatmap, tf.float32)
+ center_valid_mask = labels['seg_valid_mask'][:, :, :, None]
+ center_valid_mask = tf.image.resize(
+ center_valid_mask, tf.shape(
+ labels['seg_ct_heatmaps'])[1:3], method='nearest')
+ center_valid_mask = tf.maximum(center_valid_mask, 0.0)
+ center_valid_mask = center_valid_mask * tf.ones_like(
+ labels['seg_ct_heatmaps'])
+ weights_flattened_mask = loss_ops.flatten_spatial_dimensions(
+ center_valid_mask)
+ center_loss = center_loss_fn(
+ target_tensor=true_flattened_ct_heatmap,
+ prediction_tensor=pred_flattened_ct_heatmap,
+ weights=weights_flattened_mask)
+
+ center_loss = tf.reduce_sum(
+ center_loss /
+ (labels['num_instances'][:, None, None] + 1.0)) / batch_size
+
+ gt_masks = labels['seg_masks']
+ gt_mask_weights = labels['seg_mask_weights'][:, None,
+ None, :] * tf.ones_like(
+ gt_masks)
+ valid_mask = labels['seg_valid_mask'][:, :, :,
+ None] * tf.ones_like(gt_masks)
+
+ true_flattened_masks = loss_ops.flatten_spatial_dimensions(gt_masks)
+ true_flattened_ct_heatmap = tf.cast(true_flattened_ct_heatmap, tf.float32)
+ predicted_masks = tf.cast(outputs['mask_proposal_logits'], tf.float32)
+ predicted_masks = tf.image.resize(
+ predicted_masks, tf.shape(gt_masks)[1:3], method='bilinear')
+ pred_flattened_masks = loss_ops.flatten_spatial_dimensions(predicted_masks)
+ mask_loss = tf.cast(0.0, tf.float32)
+
+ mask_loss_fn = tf_keras.losses.BinaryCrossentropy(
+ from_logits=True,
+ label_smoothing=0.0,
+ axis=-1,
+ reduction=tf_keras.losses.Reduction.NONE,
+ name='binary_crossentropy')
+ mask_weights = tf.reshape(
+ tf.cast(true_flattened_masks >= 0, tf.float32),
+ [-1, 1]) * tf.reshape(gt_mask_weights, [-1, 1]) * tf.reshape(
+ (valid_mask), [-1, 1])
+ mask_loss = mask_loss_fn(
+ tf.reshape(gt_masks, [-1, 1]),
+ tf.reshape(pred_flattened_masks, [-1, 1]),
+ sample_weight=mask_weights)
+
+ mask_loss = tf.reduce_sum(mask_loss) / (tf.reduce_sum(mask_weights) + 1.0)
+
+ total_loss = (center_loss + loss_params.mask_weight * mask_loss)
+
+ if aux_losses:
+ total_loss += tf.add_n(aux_losses)
+
+ total_loss = loss_params.loss_weight * total_loss
+
+ losses = {'total_loss': total_loss,
+ 'mask_loss': mask_loss,
+ 'center_loss': center_loss}
+ return losses
+
+ def build_metrics(self, training: bool = True) -> List[
+ tf_keras.metrics.Metric]:
+ """Build detection metrics."""
+ metrics = []
+ if training:
+ metric_names = [
+ 'total_loss',
+ 'center_loss',
+ 'mask_loss',
+ ]
+ for name in metric_names:
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
+ else:
+ self.iou_metric = segmentation_metrics.PerClassIoU(
+ name='per_class_iou',
+ num_classes=self.task_config.model.num_classes,
+ rescale_predictions=False,
+ dtype=tf.float32)
+
+ return metrics
+
+ def train_step(self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
+ """Does forward and backward.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ images, labels = inputs
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+
+ with tf.GradientTape() as tape:
+ outputs = model(
+ images,
+ box_indices=labels['seg_box_indices'],
+ classes=labels['seg_classes'],
+ training=True)
+ outputs = tf.nest.map_structure(
+ lambda x: tf.cast(x, tf.float32), outputs)
+
+ # Computes per-replica loss.
+ losses = self.build_losses(
+ outputs=outputs,
+ labels=labels,
+ aux_losses=model.losses,
+ step=optimizer.iterations)
+ scaled_loss = losses['total_loss'] / num_replicas
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scales back gradient when LossScaleOptimizer is used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ logs = {self.loss: losses['total_loss']}
+
+ if metrics:
+ for m in metrics:
+ m.update_state(losses[m.name])
+
+ return logs
+
+ def validation_step(self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ metrics: Optional[List[Any]] = None):
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ features, labels = inputs
+
+ outputs = model(
+ features,
+ image_info=labels['image_info'],
+ training=False)
+ outputs = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), outputs)
+
+ logs = {self.loss: 0}
+ outputs = tf.one_hot(
+ tf.cast(outputs['panoptic_outputs']['category_mask'], tf.int32),
+ self.task_config.model.num_classes)
+
+ self.iou_metric.update_state(labels, tf.cast(outputs, tf.float32))
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ if state is None:
+ self.iou_metric.reset_states()
+ state = self.iou_metric
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ result = {}
+ ious = self.iou_metric.result()
+ for i, value in enumerate(ious.numpy()):
+ result.update({'iou/{}'.format(i): value})
+ # Computes mean IoU
+ result.update({'mean_iou': tf.reduce_mean(ious).numpy()})
+ return result
diff --git a/official/projects/maskconver/tasks/multiscale_maskconver.py b/official/projects/maskconver/tasks/multiscale_maskconver.py
new file mode 100644
index 00000000000..ab0c8d9a13a
--- /dev/null
+++ b/official/projects/maskconver/tasks/multiscale_maskconver.py
@@ -0,0 +1,278 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic Multi-scale MaskConver task definition."""
+from typing import Any, Dict, List, Mapping, Optional, Tuple
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import task_factory
+from official.projects.maskconver.configs import multiscale_maskconver as exp_cfg
+from official.projects.maskconver.dataloaders import multiscale_maskconver_input
+from official.projects.maskconver.losses import maskconver_losses
+from official.projects.maskconver.modeling import factory
+from official.projects.maskconver.modeling.layers import copypaste
+from official.projects.maskconver.tasks import maskconver
+from official.projects.volumetric_models.losses import segmentation_losses as volumeteric_segmentation_losses
+from official.vision.dataloaders import input_reader_factory
+
+
+@task_factory.register_task_cls(exp_cfg.MultiScaleMaskConverTask)
+class PanopticMultiScaleMaskConverTask(maskconver.PanopticMaskRCNNTask):
+
+ """A single-replica view of training procedure.
+
+ Panoptic Mask R-CNN task provides artifacts for training/evalution procedures,
+ including loading/iterating over Datasets, initializing the model, calculating
+ the loss, post-processing, and customized metrics with reduction.
+ """
+
+ def build_model(self) -> tf_keras.Model:
+ """Build Panoptic Mask R-CNN model."""
+
+ tf_keras.utils.set_random_seed(0)
+ tf.config.experimental.enable_op_determinism()
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + self.task_config.model.input_size)
+
+ l2_weight_decay = self.task_config.losses.l2_weight_decay
+ # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
+ # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
+ # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
+ l2_regularizer = (tf_keras.regularizers.l2(
+ l2_weight_decay / 2.0) if l2_weight_decay else None)
+
+ model = factory.build_multiscale_maskconver_model(
+ input_specs=input_specs,
+ model_config=self.task_config.model,
+ l2_regularizer=l2_regularizer)
+
+ # Get images and labels with batch size of 1.
+ images, labels = next(
+ iter(self.build_inputs(self.task_config.validation_data)))
+ images = tf.nest.map_structure(lambda x: x[0:1, ...], images)
+ labels = tf.nest.map_structure(lambda x: x[0:1, ...], labels)
+ _ = model(
+ images,
+ image_info=labels['image_info'],
+ training=False)
+ return model
+
+ def build_inputs(
+ self,
+ params: exp_cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None
+ ) -> tf.data.Dataset:
+ """Build input dataset."""
+ decoder_cfg = params.decoder.get()
+
+ if params.decoder.type == 'simple_decoder':
+ decoder = multiscale_maskconver_input.TfExampleDecoder(
+ regenerate_source_id=decoder_cfg.regenerate_source_id,
+ mask_binarize_threshold=decoder_cfg.mask_binarize_threshold,
+ include_panoptic_masks=decoder_cfg.include_panoptic_masks,
+ panoptic_category_mask_key=decoder_cfg.panoptic_category_mask_key,
+ panoptic_instance_mask_key=decoder_cfg.panoptic_instance_mask_key)
+ else:
+ raise ValueError('Unknown decoder type: {}!'.format(params.decoder.type))
+
+ if params.parser.copypaste:
+ sample_fn = copypaste.CopyPaste(
+ self.task_config.model.input_size[:2],
+ copypaste_frequency=params.parser.copypaste.copypaste_frequency,
+ copypaste_aug_scale_max=params.parser.copypaste.copypaste_aug_scale_max,
+ copypaste_aug_scale_min=params.parser.copypaste.copypaste_aug_scale_min,
+ aug_scale_min=params.parser.copypaste.aug_scale_min,
+ aug_scale_max=params.parser.copypaste.aug_scale_max,
+ random_flip=params.parser.aug_rand_hflip,
+ num_thing_classes=self.task_config.model.num_thing_classes)
+ else:
+ sample_fn = None
+
+ parser = multiscale_maskconver_input.Parser(
+ output_size=self.task_config.model.input_size[:2],
+ min_level=self.task_config.model.min_level,
+ max_level=self.task_config.model.max_level,
+ fpn_low_range=params.parser.fpn_low_range,
+ fpn_high_range=params.parser.fpn_high_range,
+ dtype=params.dtype,
+ aug_rand_hflip=params.parser.aug_rand_hflip,
+ aug_scale_min=params.parser.aug_scale_min,
+ aug_scale_max=params.parser.aug_scale_max,
+ max_num_instances=params.parser.max_num_instances,
+ segmentation_resize_eval_groundtruth=params.parser
+ .segmentation_resize_eval_groundtruth,
+ segmentation_groundtruth_padded_size=params.parser
+ .segmentation_groundtruth_padded_size,
+ segmentation_ignore_label=params.parser.segmentation_ignore_label,
+ panoptic_ignore_label=params.parser.panoptic_ignore_label,
+ num_panoptic_categories=self.task_config.model.num_classes,
+ num_thing_categories=self.task_config.model.num_thing_classes,
+ mask_target_level=params.parser.mask_target_level,
+ level=self.task_config.model.level,
+ gaussian_iou=params.parser.gaussaian_iou,
+ aug_type=params.parser.aug_type,)
+
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ sample_fn=sample_fn.copypaste_fn(
+ params.is_training) if sample_fn else None,
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training))
+ dataset = reader.read(input_context=input_context)
+
+ return dataset
+
+ def build_losses(self,
+ outputs: Mapping[str, Any],
+ labels: Mapping[str, Any],
+ iteration: Any,
+ aux_losses: Optional[Any] = None,
+ step=None) -> Dict[str, tf.Tensor]:
+ """Build Panoptic Mask R-CNN losses."""
+ # pylint: disable=line-too-long
+ loss_params = self._task_config.losses
+ center_loss_fn = maskconver_losses.PenaltyReducedLogisticFocalLoss(
+ alpha=loss_params.alpha, beta=loss_params.beta)
+
+ true_flattened_ct_heatmap = labels['panoptic_heatmaps']
+ true_flattened_ct_heatmap = tf.cast(true_flattened_ct_heatmap, tf.float32)
+
+ pred_flattened_ct_heatmap = outputs['class_heatmaps']
+ pred_flattened_ct_heatmap = tf.cast(pred_flattened_ct_heatmap, tf.float32)
+
+ center_loss = center_loss_fn(
+ target_tensor=true_flattened_ct_heatmap,
+ prediction_tensor=pred_flattened_ct_heatmap,
+ weights=1.0)
+
+ replica_context = tf.distribute.get_replica_context()
+ global_num_instances = replica_context.all_reduce(
+ tf.distribute.ReduceOp.SUM, labels['num_instances'])
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+ num_instances = tf.cast(global_num_instances, tf.float32) / tf.cast(num_replicas, tf.float32) + 1.0
+
+ center_loss = tf.reduce_sum(center_loss) / num_instances
+
+ gt_masks = labels['panoptic_masks']
+ gt_mask_weights = labels['panoptic_mask_weights'][:, None, None, :] * tf.ones_like(gt_masks)
+ panoptic_padding_mask = labels['panoptic_padding_mask'][:, :, :, None] * tf.ones_like(gt_masks)
+
+ # gt_masks
+ _, h, w, q = gt_masks.get_shape().as_list()
+ predicted_masks = tf.cast(outputs['mask_proposal_logits'], tf.float32)
+ predicted_masks = tf.image.resize(
+ predicted_masks, tf.shape(gt_masks)[1:3], method='bilinear')
+
+ mask_loss_fn = tf_keras.losses.BinaryCrossentropy(
+ from_logits=True,
+ label_smoothing=0.0,
+ axis=-1,
+ reduction=tf_keras.losses.Reduction.NONE,
+ name='binary_crossentropy')
+
+ mask_weights = tf.cast(gt_masks >= 0, tf.float32) * gt_mask_weights * (
+ 1 - panoptic_padding_mask) # b, h, w, # max inst
+ mask_loss = mask_loss_fn(
+ tf.expand_dims(gt_masks, -1),
+ tf.expand_dims(predicted_masks, -1),
+ sample_weight=tf.expand_dims(mask_weights, -1))
+
+ mask_loss = tf.reshape(mask_loss, [-1, h * w, q])
+ mask_loss = tf.reduce_sum(tf.reduce_mean(mask_loss, axis=1)) / num_instances
+
+ # Dice loss
+ masked_predictions = tf.sigmoid(predicted_masks) * tf.cast(
+ gt_mask_weights > 0, tf.float32) * (1 - panoptic_padding_mask)
+ masked_gt_masks = gt_masks * tf.cast(gt_mask_weights > 0, tf.float32) * (
+ 1 - panoptic_padding_mask)
+
+ masked_predictions = tf.transpose(masked_predictions, [0, 3, 1, 2])
+ masked_predictions = tf.reshape(masked_predictions, [-1, h, w, 1])
+ masked_gt_masks = tf.transpose(masked_gt_masks, [0, 3, 1, 2])
+ masked_gt_masks = tf.reshape(masked_gt_masks, [-1, h, w, 1])
+
+ dice_loss_fn = volumeteric_segmentation_losses.SegmentationLossDiceScore(
+ metric_type='adaptive', axis=(2, 3))
+ dice_loss = dice_loss_fn(logits=masked_predictions, labels=masked_gt_masks)
+
+ total_loss = center_loss + loss_params.mask_weight * (mask_loss + dice_loss)
+ if aux_losses:
+ total_loss += tf.add_n(aux_losses)
+
+ total_loss = loss_params.loss_weight * total_loss
+
+ losses = {'total_loss': total_loss,
+ 'mask_loss': mask_loss,
+ 'center_loss': center_loss,
+ 'dice_loss': dice_loss,}
+ return losses
+
+ def train_step(self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
+ """Does forward and backward.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ images, labels = inputs
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+
+ with tf.GradientTape() as tape:
+ outputs = model(
+ images,
+ box_indices=labels['panoptic_box_indices'],
+ classes=labels['panoptic_classes'],
+ training=True)
+ outputs = tf.nest.map_structure(
+ lambda x: tf.cast(x, tf.float32), outputs)
+
+ # Computes per-replica loss.
+ losses = self.build_losses(
+ outputs=outputs,
+ labels=labels,
+ aux_losses=model.losses,
+ iteration=optimizer.iterations,
+ step=optimizer.iterations)
+ scaled_loss = losses['total_loss'] / num_replicas
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scales back gradient when LossScaleOptimizer is used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ logs = {self.loss: losses['total_loss']}
+
+ if metrics:
+ for m in metrics:
+ m.update_state(losses[m.name])
+
+ return logs
diff --git a/official/projects/maskconver/train.py b/official/projects/maskconver/train.py
new file mode 100644
index 00000000000..f703026be79
--- /dev/null
+++ b/official/projects/maskconver/train.py
@@ -0,0 +1,30 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic MaskRCNN trainer."""
+
+from absl import app
+
+from official.common import flags as tfm_flags
+from official.projects.maskconver.configs import maskconver as maskconver_cfg # pylint: disable=unused-import
+from official.projects.maskconver.configs import multiscale_maskconver as multiscale_maskconver_cfg # pylint: disable=unused-import
+from official.projects.maskconver.modeling import fpn # pylint: disable=unused-import
+from official.projects.maskconver.tasks import maskconver as maskconver_task # pylint: disable=unused-import
+from official.projects.maskconver.tasks import multiscale_maskconver as multiscale_maskconver_task # pylint: disable=unused-import
+from official.vision import train
+
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(train.main)
diff --git a/official/projects/maxvit/README.md b/official/projects/maxvit/README.md
new file mode 100644
index 00000000000..af83526717e
--- /dev/null
+++ b/official/projects/maxvit/README.md
@@ -0,0 +1,115 @@
+# MaxViT: Multi-Axis Vision Transformer (ECCV 2022)
+
+[](https://arxiv.org/abs/2204.01697)
+
+⚠️ **DISCLAIMER**: This implementation is still under development.
+
+[TOC]
+
+[MaxViT](https://arxiv.org/abs/2204.01697) is a family of hybrid (CNN + ViT)
+vision backbone models, that achieves better performances across the board
+for both parameter and FLOPs efficiency than both state-of-the-art ConvNets and
+Transformers ([Blog](https://ai.googleblog.com/2022/09/a-multi-axis-approach-for-vision.html)).
+They can also scale well on large dataset sizes like ImageNet-21K.
+Notably, due to the linear-complexity of the grid attention used, MaxViT scales
+well on tasks requiring large image sizes, such as object detection and
+segmentation.
+
+MaxViT meta-architecture: a homogeneously stacked backbone, wherein each MaxViT
+block contains [MBConv](https://arxiv.org/abs/2104.00298), block attention
+(window-based local attention), and grid attention (dilated global attention).
+
+
+
+
+
+Results on ImageNet-1k standard train and test:
+
+
+
+
+
+Results on ImageNet-21k and JFT pre-trained models:
+
+
+
+
+
+
+### Model Performance
+
+Note: Deit ImageNet pretrain experimental settings are different from the
+paper. These experiments follows the pre-training hyperparameters in
+[paper](https://arxiv.org/abs/2204.01697) and only run pre-training for similar
+number of steps. The paper suggested a short fine-tuning with different
+hyper-parameters and EMA.
+
+
+
+#### Deit ImageNet pretrain {.new-tab}
+
+Model | Eval Size | Top-1 Acc | Acc on Paper | #Param | #FLOPs | Config
+------------- | --------- | :---------: | :----------: | :----: | :----: | :----:
+MaxViT-Tiny | 224x224 | 83.1 (-0.5) | 83.6 | 31M | 5.6G | [config](configs/experiments/maxvit_tiny_imagenet.yaml)
+MaxViT-Small | 224x224 | 84.1 (-0.3) | 84.4 | 69M | 11.7G | [config](configs/experiments/maxvit_small_imagenet.yaml)
+MaxViT-Base | 224x224 | 84.2 (-0.7) | 84.9 | 120M | 23.4G | [config](configs/experiments/maxvit_base_imagenet.yaml)
+MaxViT-Large | 224x224 | 84.6 (-0.6) | 85.2 | 212M | 43.9G | [config](configs/experiments/maxvit_large_imagenet.yaml)
+MaxViT-XLarge | 224x224 | 84.8 | - | 475M | 97.9G | [config](configs/experiments/maxvit_xlarge_imagenet.yaml)
+
+#### Cascade RCNN models {.new-tab}
+
+Model | Image Size | Window Size | Epochs | box AP | box AP on paper | mask AP | Config
+------------ | ---------: | :---------: | :----: | :-----------: | :-------------: | :-----: | :----:
+MaxViT-Tiny | 640x640 | 20x20 | 200 | 49.97 | - | 42.69 | [config](configs/experiments/coco_maxvitt_i640_crcnn.yaml)
+MaxViT-Tiny | 896x896 | 28x28 | 200 | 52.35 (+0.25) | 52.1 | 44.69 | -
+MaxViT-Small | 640x640 | 20x20 | 200 | 50.79 | - | 43.36 | -
+MaxViT-Small | 896x896 | 28x28 | 200 | 53.54 (+0.44) | 53.1 | 45.79 | [config](configs/experiments/coco_maxvits_i896_crcnn.yaml)
+MaxViT-Base | 640x640 | 20x20 | 200 | 51.59 | - | 44.07 | [config](configs/experiments/coco_maxvitb_i640_crcnn.yaml)
+MaxViT-Base | 896x896 | 28x28 | 200 | 53.47 (+0.07) | 53.4 | 45.96 | [config](configs/experiments/coco_maxvitb_i896_crcnn.yaml)
+
+
+
+
+
+
+#### JFT-300M supervised pretrain {.new-tab}
+
+Model | Pretrain Size | #Param | #FLOPs | globalPR-AUC
+------------- | :------------ | :----: | :----: | :----------:
+MaxViT-Base | 224x224 | 120M | 23.4G | 52.75%
+MaxViT-Large | 224x224 | 212M | 43.9G | 53.77%
+MaxViT-XLarge | 224x224 | 475M | - | 54.71%
+
+
+#### ImageNet Finetuning {.new-tab}
+
+Model | Image Size | Top-1 Acc | Acc on Paper | #Param | #FLOPs | Config
+------------- | :--------- | :-------------: | :----------: | :----: | :----: | :----:
+MaxViT-Base | 384x384 | 88.37% (-0.32%) | 88.69% | 120M | 74.2G | [config](configs/experiments/finetune_maxvitb_imagenet_i384.yaml)
+MaxViT-Base | 512x512 | 88.63% (-0.19%) | 88.82% | 120M | 138.3G | [config](configs/experiments/finetune_maxvitb_imagenet_i512.yaml)
+MaxViT-Large | 384x384 | 88.86% (-0.26%) | 89.12% | 212M | 128.7G | [config](configs/experiments/finetune_maxvitl_imagenet_i384.yaml)
+MaxViT-Large | 512x512 | 89.02% (-0.39%) | 89.41% | 212M | 245.2G | [config](configs/experiments/finetune_maxvitl_imagenet_i512.yaml)
+MaxViT-XLarge | 384x384 | 89.21% (-0.15%) | 89.36% | 475M | 293.7G | [config](configs/experiments/finetune_maxvitxl_imagenet_i384.yaml)
+MaxViT-XLarge | 512x512 | 89.31% (-0.22%) | 89.53% | 475M | 535.2G | [config](configs/experiments/finetune_maxvitxl_imagenet_i512.yaml)
+
+#### Cascade RCNN models {.new-tab}
+
+Model | Image Size | Window Size | Epochs | box AP | box AP on paper | mask AP | Config
+------------ | ---------: | :---------: | :----: | :-----------: | :-------------: | :-----: | :----:
+MaxViT-Base | 896x896 | 28x28 | 200 | 54.31 (+0.91) | 53.4 | 46.31 | [config](configs/experiments/coco_maxvitb_i896_crcnn.yaml)
+MaxViT-Large | 896x896 | 28x28 | 200 | 54.69 | - | 46.59 | [config](configs/experiments/coco_maxvitl_i896_crcnn.yaml)
+
+
+
+### Citation
+
+Should you find this repository useful, please consider citing:
+
+```
+@article{tu2022maxvit,
+ title={MaxViT: Multi-Axis Vision Transformer},
+ author={Tu, Zhengzhong and Talebi, Hossein and Zhang, Han and Yang, Feng and Milanfar, Peyman and Bovik, Alan and Li, Yinxiao},
+ journal={ECCV},
+ year={2022},
+}
+```
diff --git a/official/projects/maxvit/__init__.py b/official/projects/maxvit/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/maxvit/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/maxvit/configs/__init__.py b/official/projects/maxvit/configs/__init__.py
new file mode 100644
index 00000000000..b3ec1ee80ce
--- /dev/null
+++ b/official/projects/maxvit/configs/__init__.py
@@ -0,0 +1,22 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Lint as: python3
+"""Configs package definition."""
+
+from official.projects.maxvit.configs import backbones # pylint:disable=unused-import
+from official.projects.maxvit.configs import rcnn # pylint:disable=unused-import
+from official.projects.maxvit.configs import retinanet # pylint:disable=unused-import
+from official.projects.maxvit.configs import semantic_segmentation # pylint:disable=unused-import
+from official.projects.maxvit.configs import image_classification # pylint:disable=unused-import
diff --git a/official/projects/maxvit/configs/backbones.py b/official/projects/maxvit/configs/backbones.py
new file mode 100644
index 00000000000..898263d3e01
--- /dev/null
+++ b/official/projects/maxvit/configs/backbones.py
@@ -0,0 +1,94 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""CoAtNet Image classification configuration definition."""
+import dataclasses
+from typing import Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import hyperparams
+from official.vision.configs import backbones
+
+
+@dataclasses.dataclass
+class MaxViT(hyperparams.Config):
+ """MaxViT config."""
+ model_name: str = 'maxvit-tiny'
+ # These configs are specified according to `model_name` in default.
+ # Set values will override the default configs.
+ stem_hsize: Optional[Tuple[int, ...]] = None
+ block_type: Optional[Tuple[str, ...]] = None
+ num_blocks: Optional[Tuple[int, ...]] = None
+ hidden_size: Optional[Tuple[int, ...]] = None
+
+ # specific to the multi-axis attention in MaxViT
+ # Note that the window_size and grid_size should be divisible by all the
+ # feature map sizes along the entire network. Say, if you train on ImageNet
+ # classification at 224x224, set both to 7 is almost the only choice.
+ # If you train on COCO object detection at 896x896, set it to 28 is suggested,
+ # as following Swin Transformer, window size should scales with feature size.
+ # You may as well set it as 14 or 7.
+ window_size: int = 7 # window size for conducting block attention module.
+ grid_size: int = 7 # grid size for conducting sparse global grid attention.
+
+ # tfm specific
+ head_size: int = 32
+ dropatt: Optional[float] = None
+ dropout: Optional[float] = None
+ rel_attn_type: str = '2d_multi_head'
+ num_heads: Optional[int] = None
+
+ # A string of `current_window_size/ckpt_window_size` for finetuning from a
+ # checkpoint trained with `ckpt_window_size`.
+ scale_ratio: Optional[str] = None
+ ln_epsilon: float = 1e-5
+ ln_dtype: Optional[tf.DType] = None
+
+ # conv specific
+ downsample_loc: str = 'depth_conv'
+ kernel_size: int = 3
+ se_ratio: float = 0.25
+ dropcnn: Optional[float] = None
+
+ # Only channels_last is supported for now.
+ data_format: str = 'channels_last'
+ norm_type: str = 'sync_batch_norm'
+
+ # shared
+ add_pos_enc: bool = False
+ pool_type: str = '2d:avg'
+ pool_stride: int = 2
+ expansion_rate: int = 4
+
+ # Stochastic depth keep probability for the residual connection in. Smaller
+ # value means stronger regularization. If using anneal, it decays linearly
+ # from 1.0 to this value with the depth of each layer."
+ survival_prob: Optional[float] = None # from [0, 1]
+ survival_prob_anneal: bool = True
+
+ kernel_initializer: str = 'glorot_uniform'
+ bias_initializer: str = 'zeros'
+
+ # For cls head, should be same as the last `hidden_size` of backbone.
+ representation_size: Optional[int] = None
+ # Only effective when representation_size > 0.
+ add_gap_layer_norm: bool = True
+
+
+@dataclasses.dataclass
+class Backbone(backbones.Backbone):
+ """Configuration for backbones."""
+ type: Optional[str] = 'maxvit'
+ maxvit: MaxViT = dataclasses.field(default_factory=MaxViT)
diff --git a/official/projects/maxvit/configs/experiments/coco_maxvitb_i640_crcnn.yaml b/official/projects/maxvit/configs/experiments/coco_maxvitb_i640_crcnn.yaml
new file mode 100644
index 00000000000..939f59d1cfb
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/coco_maxvitb_i640_crcnn.yaml
@@ -0,0 +1,40 @@
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ anchor:
+ anchor_size: 3.0
+ backbone:
+ maxvit:
+ model_name: 'maxvit-base'
+ window_size: 20
+ grid_size: 20
+ scale_ratio: '20/7'
+ survival_prob: 0.7
+ input_size: [640, 640, 3]
+ max_level: 7
+ min_level: 3
+ rpn_head:
+ num_convs: 2
+ train_data:
+ global_batch_size: 256
+ validation_data:
+ global_batch_size: 64
+trainer:
+ optimizer_config:
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ decay_steps: 115750
+ initial_learning_rate: 0.002
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.0001
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 6000
diff --git a/official/projects/maxvit/configs/experiments/coco_maxvitb_i896_crcnn.yaml b/official/projects/maxvit/configs/experiments/coco_maxvitb_i896_crcnn.yaml
new file mode 100644
index 00000000000..5f032105f7c
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/coco_maxvitb_i896_crcnn.yaml
@@ -0,0 +1,42 @@
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 2.0e-07
+ model:
+ anchor:
+ anchor_size: 3.0
+ backbone:
+ maxvit:
+ model_name: 'maxvit-base'
+ window_size: 28
+ grid_size: 28
+ scale_ratio: '28/7'
+ survival_prob: 0.2
+ input_size: [896, 896, 3]
+ max_level: 7
+ min_level: 3
+ rpn_head:
+ num_convs: 2
+ train_data:
+ global_batch_size: 256
+ validation_data:
+ global_batch_size: 128
+trainer:
+ train_steps: 90000
+ validation_steps: 39
+ optimizer_config:
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ decay_steps: 90000
+ initial_learning_rate: 0.003
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.05
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 6000
diff --git a/official/projects/maxvit/configs/experiments/coco_maxvitl_i896_crcnn.yaml b/official/projects/maxvit/configs/experiments/coco_maxvitl_i896_crcnn.yaml
new file mode 100644
index 00000000000..b34d90ef7cd
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/coco_maxvitl_i896_crcnn.yaml
@@ -0,0 +1,42 @@
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 0.0
+ model:
+ anchor:
+ anchor_size: 3.0
+ backbone:
+ maxvit:
+ model_name: 'maxvit-large'
+ window_size: 28
+ grid_size: 28
+ scale_ratio: '28/7'
+ survival_prob: 0.2
+ input_size: [896, 896, 3]
+ max_level: 7
+ min_level: 3
+ rpn_head:
+ num_convs: 2
+ train_data:
+ global_batch_size: 256
+ validation_data:
+ global_batch_size: 64
+trainer:
+ train_steps: 90000
+ validation_steps: 78
+ optimizer_config:
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ decay_steps: 90000
+ initial_learning_rate: 0.003
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.05
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 6000
diff --git a/official/projects/maxvit/configs/experiments/coco_maxvits_i896_crcnn.yaml b/official/projects/maxvit/configs/experiments/coco_maxvits_i896_crcnn.yaml
new file mode 100644
index 00000000000..fa52a5d978f
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/coco_maxvits_i896_crcnn.yaml
@@ -0,0 +1,42 @@
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ anchor:
+ anchor_size: 3.0
+ backbone:
+ maxvit:
+ model_name: 'maxvit-small'
+ window_size: 28
+ grid_size: 28
+ scale_ratio: '28/7'
+ survival_prob: 0.5
+ input_size: [896, 896, 3]
+ max_level: 7
+ min_level: 3
+ rpn_head:
+ num_convs: 2
+ train_data:
+ global_batch_size: 256
+ validation_data:
+ global_batch_size: 128
+trainer:
+ train_steps: 115750
+ validation_steps: 39
+ optimizer_config:
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ decay_steps: 90000
+ initial_learning_rate: 0.003
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.05
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 6000
diff --git a/official/projects/maxvit/configs/experiments/coco_maxvitt_i640_crcnn.yaml b/official/projects/maxvit/configs/experiments/coco_maxvitt_i640_crcnn.yaml
new file mode 100644
index 00000000000..50945f85111
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/coco_maxvitt_i640_crcnn.yaml
@@ -0,0 +1,40 @@
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ anchor:
+ anchor_size: 3.0
+ backbone:
+ maxvit:
+ model_name: 'maxvit-tiny'
+ window_size: 20
+ grid_size: 20
+ scale_ratio: '20/7'
+ survival_prob: 0.3
+ input_size: [640, 640, 3]
+ max_level: 7
+ min_level: 3
+ rpn_head:
+ num_convs: 2
+ train_data:
+ global_batch_size: 256
+ validation_data:
+ global_batch_size: 64
+trainer:
+ optimizer_config:
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ decay_steps: 115750
+ initial_learning_rate: 0.002
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.0001
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 6000
diff --git a/official/projects/maxvit/configs/experiments/coco_maxvitxl_i896_crcnn.yaml b/official/projects/maxvit/configs/experiments/coco_maxvitxl_i896_crcnn.yaml
new file mode 100644
index 00000000000..fca7da78f63
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/coco_maxvitxl_i896_crcnn.yaml
@@ -0,0 +1,40 @@
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ anchor:
+ anchor_size: 3.0
+ backbone:
+ maxvit:
+ model_name: 'maxvit-base'
+ window_size: 28
+ grid_size: 28
+ scale_ratio: '28/7'
+ survival_prob: 0.7
+ input_size: [896, 896, 3]
+ max_level: 7
+ min_level: 3
+ rpn_head:
+ num_convs: 2
+ train_data:
+ global_batch_size: 256
+ validation_data:
+ global_batch_size: 64
+trainer:
+ optimizer_config:
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ decay_steps: 115750
+ initial_learning_rate: 0.003
+ optimizer:
+ adamw:
+ weight_decay_rate: 0.0001
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 6000
diff --git a/official/projects/maxvit/configs/experiments/finetune_maxvitb_imagenet_i384.yaml b/official/projects/maxvit/configs/experiments/finetune_maxvitb_imagenet_i384.yaml
new file mode 100644
index 00000000000..85d578f1716
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/finetune_maxvitb_imagenet_i384.yaml
@@ -0,0 +1,52 @@
+runtime:
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: 'backbone'
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-base'
+ representation_size: 768
+ survival_prob: 0.8
+ window_size: 12
+ grid_size: 12
+ scale_ratio: '12/7'
+ input_size: [384, 384, 3]
+ train_data:
+ global_batch_size: 512
+ dtype: 'bfloat16'
+ aug_crop: false
+ mixup_and_cutmix: null
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 15
+ validation_data:
+ dtype: 'bfloat16'
+ aug_crop: false
+ losses:
+ label_smoothing: 0.1
+ use_binary_cross_entropy: false
+ one_hot: true
+trainer:
+ train_steps: 100080
+ steps_per_loop: 2000
+ summary_interval: 2000
+ validation_interval: 2000
+ checkpoint_interval: 2000
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 1.0e-4
+ gradient_clip_norm: 1.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ type: constant
+ constant:
+ learning_rate: 5.0e-5
+ warmup:
+ type: null
diff --git a/official/projects/maxvit/configs/experiments/finetune_maxvitb_imagenet_i512.yaml b/official/projects/maxvit/configs/experiments/finetune_maxvitb_imagenet_i512.yaml
new file mode 100644
index 00000000000..55839b1d476
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/finetune_maxvitb_imagenet_i512.yaml
@@ -0,0 +1,52 @@
+runtime:
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: 'backbone'
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-base'
+ representation_size: 768
+ survival_prob: 0.8
+ window_size: 16
+ grid_size: 16
+ scale_ratio: '16/7'
+ input_size: [512, 512, 3]
+ train_data:
+ global_batch_size: 512
+ dtype: 'bfloat16'
+ aug_crop: false
+ mixup_and_cutmix: null
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 15
+ validation_data:
+ dtype: 'bfloat16'
+ aug_crop: false
+ losses:
+ label_smoothing: 0.1
+ use_binary_cross_entropy: false
+ one_hot: true
+trainer:
+ train_steps: 100080
+ steps_per_loop: 2000
+ summary_interval: 2000
+ validation_interval: 2000
+ checkpoint_interval: 2000
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 1.0e-4
+ gradient_clip_norm: 1.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ type: constant
+ constant:
+ learning_rate: 5.0e-5
+ warmup:
+ type: null
diff --git a/official/projects/maxvit/configs/experiments/finetune_maxvitl_imagenet_i384.yaml b/official/projects/maxvit/configs/experiments/finetune_maxvitl_imagenet_i384.yaml
new file mode 100644
index 00000000000..1163b99ed38
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/finetune_maxvitl_imagenet_i384.yaml
@@ -0,0 +1,52 @@
+runtime:
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: 'backbone'
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-large'
+ representation_size: 1024
+ survival_prob: 0.7
+ window_size: 12
+ grid_size: 12
+ scale_ratio: '12/7'
+ input_size: [384, 384, 3]
+ train_data:
+ global_batch_size: 512
+ dtype: 'bfloat16'
+ aug_crop: false
+ mixup_and_cutmix: null
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 15
+ validation_data:
+ dtype: 'bfloat16'
+ aug_crop: false
+ losses:
+ label_smoothing: 0.1
+ use_binary_cross_entropy: false
+ one_hot: true
+trainer:
+ train_steps: 100080
+ steps_per_loop: 2000
+ summary_interval: 2000
+ validation_interval: 2000
+ checkpoint_interval: 2000
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 1.0e-4
+ gradient_clip_norm: 1.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ type: constant
+ constant:
+ learning_rate: 5.0e-5
+ warmup:
+ type: null
diff --git a/official/projects/maxvit/configs/experiments/finetune_maxvitl_imagenet_i512.yaml b/official/projects/maxvit/configs/experiments/finetune_maxvitl_imagenet_i512.yaml
new file mode 100644
index 00000000000..b07cf259477
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/finetune_maxvitl_imagenet_i512.yaml
@@ -0,0 +1,52 @@
+runtime:
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: 'backbone'
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-large'
+ representation_size: 1024
+ survival_prob: 0.8
+ window_size: 16
+ grid_size: 16
+ scale_ratio: '16/7'
+ input_size: [512, 512, 3]
+ train_data:
+ global_batch_size: 512
+ dtype: 'bfloat16'
+ aug_crop: false
+ mixup_and_cutmix: null
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 15
+ validation_data:
+ dtype: 'bfloat16'
+ aug_crop: false
+ losses:
+ label_smoothing: 0.1
+ use_binary_cross_entropy: false
+ one_hot: true
+trainer:
+ train_steps: 100080
+ steps_per_loop: 2000
+ summary_interval: 2000
+ validation_interval: 2000
+ checkpoint_interval: 2000
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 1.0e-4
+ gradient_clip_norm: 1.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ type: constant
+ constant:
+ learning_rate: 5.0e-5
+ warmup:
+ type: null
diff --git a/official/projects/maxvit/configs/experiments/finetune_maxvitxl_imagenet_i384.yaml b/official/projects/maxvit/configs/experiments/finetune_maxvitxl_imagenet_i384.yaml
new file mode 100644
index 00000000000..cf2e8b3d218
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/finetune_maxvitxl_imagenet_i384.yaml
@@ -0,0 +1,52 @@
+runtime:
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: 'backbone'
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-xlarge'
+ representation_size: 1536
+ survival_prob: 0.8
+ window_size: 12
+ grid_size: 12
+ scale_ratio: '12/7'
+ input_size: [384, 384, 3]
+ train_data:
+ global_batch_size: 512
+ dtype: 'bfloat16'
+ aug_crop: false
+ mixup_and_cutmix: null
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 15
+ validation_data:
+ dtype: 'bfloat16'
+ aug_crop: false
+ losses:
+ label_smoothing: 0.1
+ use_binary_cross_entropy: false
+ one_hot: true
+trainer:
+ train_steps: 100080
+ steps_per_loop: 2000
+ summary_interval: 2000
+ validation_interval: 2000
+ checkpoint_interval: 2000
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 1.0e-4
+ gradient_clip_norm: 1.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ type: constant
+ constant:
+ learning_rate: 5.0e-5
+ warmup:
+ type: null
diff --git a/official/projects/maxvit/configs/experiments/finetune_maxvitxl_imagenet_i512.yaml b/official/projects/maxvit/configs/experiments/finetune_maxvitxl_imagenet_i512.yaml
new file mode 100644
index 00000000000..5c421b55b5a
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/finetune_maxvitxl_imagenet_i512.yaml
@@ -0,0 +1,52 @@
+runtime:
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: 'backbone'
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-xlarge'
+ representation_size: 1536
+ survival_prob: 0.8
+ window_size: 16
+ grid_size: 16
+ scale_ratio: '16/7'
+ input_size: [512, 512, 3]
+ train_data:
+ global_batch_size: 512
+ dtype: 'bfloat16'
+ aug_crop: false
+ mixup_and_cutmix: null
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 15
+ validation_data:
+ dtype: 'bfloat16'
+ aug_crop: false
+ losses:
+ label_smoothing: 0.1
+ use_binary_cross_entropy: false
+ one_hot: true
+trainer:
+ train_steps: 100080
+ steps_per_loop: 2000
+ summary_interval: 2000
+ validation_interval: 2000
+ checkpoint_interval: 2000
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 1.0e-4
+ gradient_clip_norm: 1.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ type: constant
+ constant:
+ learning_rate: 5.0e-5
+ warmup:
+ type: null
diff --git a/official/projects/maxvit/configs/experiments/maxvit_base_imagenet.yaml b/official/projects/maxvit/configs/experiments/maxvit_base_imagenet.yaml
new file mode 100644
index 00000000000..105ca51b875
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_base_imagenet.yaml
@@ -0,0 +1,25 @@
+task:
+ init_checkpoint: ''
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-base'
+ representation_size: 768
+ input_size: [224, 224, 3]
+trainer:
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.003
+ alpha: 0.01
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 10000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_base_imagenet_gpu.yaml b/official/projects/maxvit/configs/experiments/maxvit_base_imagenet_gpu.yaml
new file mode 100644
index 00000000000..4127025b60b
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_base_imagenet_gpu.yaml
@@ -0,0 +1,43 @@
+runtime:
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float16'
+task:
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-base'
+ representation_size: 768
+ norm_type: 'batch_norm'
+ input_size: [224, 224, 3]
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ train_data:
+ global_batch_size: 192
+ shuffle_buffer_size: 256
+ dtype: 'float16'
+ validation_data:
+ global_batch_size: 256
+ dtype: 'float16'
+trainer:
+ train_steps: 1500000
+ steps_per_loop: 10000
+ summary_interval: 10000
+ validation_interval: 10000
+ validation_steps: 195
+ optimizer_config:
+ ema: null
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.0001
+ alpha: 0.01
+ decay_steps: 1500000
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 10000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_large_imagenet.yaml b/official/projects/maxvit/configs/experiments/maxvit_large_imagenet.yaml
new file mode 100644
index 00000000000..e6d32dfece0
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_large_imagenet.yaml
@@ -0,0 +1,29 @@
+task:
+ init_checkpoint: ''
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-large'
+ representation_size: 1024
+ input_size: [224, 224, 3]
+trainer:
+ max_to_keep: 5
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ gradient_clip_norm: 0.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.001
+ alpha: 0.00
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 10000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_large_imagenet_gpu.yaml b/official/projects/maxvit/configs/experiments/maxvit_large_imagenet_gpu.yaml
new file mode 100644
index 00000000000..b90c846a110
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_large_imagenet_gpu.yaml
@@ -0,0 +1,43 @@
+runtime:
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float16'
+task:
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-large'
+ representation_size: 1024
+ norm_type: 'batch_norm'
+ input_size: [224, 224, 3]
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ train_data:
+ global_batch_size: 128
+ shuffle_buffer_size: 256
+ dtype: 'float16'
+ validation_data:
+ global_batch_size: 256
+ dtype: 'float16'
+trainer:
+ train_steps: 1000000
+ steps_per_loop: 10000
+ summary_interval: 10000
+ validation_interval: 10000
+ validation_steps: 195
+ optimizer_config:
+ ema: null
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.0001
+ alpha: 0.01
+ decay_steps: 1000000
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 10000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_small_imagenet.yaml b/official/projects/maxvit/configs/experiments/maxvit_small_imagenet.yaml
new file mode 100644
index 00000000000..1afee041bed
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_small_imagenet.yaml
@@ -0,0 +1,27 @@
+task:
+ init_checkpoint: ''
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-small'
+ representation_size: 768
+ input_size: [224, 224, 3]
+trainer:
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ gradient_clip_norm: 0.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.002
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 10000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_small_imagenet_gpu.yaml b/official/projects/maxvit/configs/experiments/maxvit_small_imagenet_gpu.yaml
new file mode 100644
index 00000000000..bd971ac02b5
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_small_imagenet_gpu.yaml
@@ -0,0 +1,43 @@
+runtime:
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float16'
+task:
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-small'
+ representation_size: 768
+ norm_type: 'batch_norm'
+ input_size: [224, 224, 3]
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ train_data:
+ global_batch_size: 256
+ shuffle_buffer_size: 512
+ dtype: 'float16'
+ validation_data:
+ global_batch_size: 256
+ dtype: 'float16'
+trainer:
+ train_steps: 1500000
+ steps_per_loop: 8000
+ summary_interval: 8000
+ validation_interval: 8000
+ validation_steps: 195
+ optimizer_config:
+ ema: null
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.0001
+ alpha: 0.01
+ decay_steps: 1500000
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 8000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_tiny_imagenet.yaml b/official/projects/maxvit/configs/experiments/maxvit_tiny_imagenet.yaml
new file mode 100644
index 00000000000..fed93cb6eb3
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_tiny_imagenet.yaml
@@ -0,0 +1,31 @@
+task:
+ init_checkpoint: ''
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-tiny'
+ representation_size: 512
+ add_gap_layer_norm: true
+ kernel_initializer: 'glorot_uniform'
+ kernel_initializer: 'glorot_uniform'
+ input_size: [224, 224, 3]
+trainer:
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ gradient_clip_norm: 0.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.002
+ alpha: 0.0
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 10000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_xlarge_imagenet.yaml b/official/projects/maxvit/configs/experiments/maxvit_xlarge_imagenet.yaml
new file mode 100644
index 00000000000..63ea180909f
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_xlarge_imagenet.yaml
@@ -0,0 +1,29 @@
+task:
+ init_checkpoint: ''
+ losses:
+ l2_weight_decay: 1.0e-07
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-xlarge'
+ representation_size: 1536
+ input_size: [224, 224, 3]
+trainer:
+ max_to_keep: 5
+ optimizer_config:
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ gradient_clip_norm: 0.0
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.001
+ alpha: 0.01
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 10000
diff --git a/official/projects/maxvit/configs/experiments/maxvit_xlarge_imagenet_gpu.yaml b/official/projects/maxvit/configs/experiments/maxvit_xlarge_imagenet_gpu.yaml
new file mode 100644
index 00000000000..cfd0e519b3f
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/maxvit_xlarge_imagenet_gpu.yaml
@@ -0,0 +1,43 @@
+runtime:
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float16'
+task:
+ model:
+ backbone:
+ maxvit:
+ model_name: 'maxvit-xlarge'
+ representation_size: 1536
+ norm_type: 'batch_norm'
+ input_size: [224, 224, 3]
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ train_data:
+ global_batch_size: 32
+ shuffle_buffer_size: 64
+ dtype: 'float16'
+ validation_data:
+ global_batch_size: 64
+ dtype: 'float16'
+trainer:
+ train_steps: 3000000
+ steps_per_loop: 15000
+ summary_interval: 15000
+ validation_interval: 15000
+ validation_steps: 390
+ optimizer_config:
+ ema: null
+ optimizer:
+ type: 'adamw'
+ adamw:
+ weight_decay_rate: 0.05
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.00001
+ alpha: 0.01
+ decay_steps: 3000000
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 15000
diff --git a/official/projects/maxvit/configs/experiments/retinanet_maxvit_base_coco_i1280_tpu.yaml b/official/projects/maxvit/configs/experiments/retinanet_maxvit_base_coco_i1280_tpu.yaml
new file mode 100644
index 00000000000..90bb355cf51
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/retinanet_maxvit_base_coco_i1280_tpu.yaml
@@ -0,0 +1,25 @@
+# RetinaNet with MaxViT backbone COCO detection.
+# Required flags:
+# --experiment_type=retinanet_maxvit_coco
+#
+# Expected AP on DF TPU 8x8: 50.38%.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ type: 'maxvit'
+ maxvit:
+ model_name: 'maxvit-base'
+ window_size: 40
+ grid_size: 40
+ scale_ratio: '40/7'
+ survival_prob: 0.3
+ input_size: [1280, 1280, 3]
diff --git a/official/projects/maxvit/configs/experiments/retinanet_maxvit_base_coco_i640_tpu.yaml b/official/projects/maxvit/configs/experiments/retinanet_maxvit_base_coco_i640_tpu.yaml
new file mode 100644
index 00000000000..28e4a3549a2
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/retinanet_maxvit_base_coco_i640_tpu.yaml
@@ -0,0 +1,20 @@
+# RetinaNet with MaxViT backbone COCO detection.
+# Required flags:
+# --experiment_type=retinanet_maxvit_coco
+#
+# Expected AP on DF TPU 4x4: 46.63%.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ model:
+ backbone:
+ type: 'maxvit'
+ maxvit:
+ model_name: 'maxvit-base'
+ window_size: 20
+ grid_size: 20
+ scale_ratio: '20/7'
+ survival_prob: 0.3
diff --git a/official/projects/maxvit/configs/experiments/seg_coco_maxvits_i640.yaml b/official/projects/maxvit/configs/experiments/seg_coco_maxvits_i640.yaml
new file mode 100644
index 00000000000..889c256259e
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/seg_coco_maxvits_i640.yaml
@@ -0,0 +1,69 @@
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ model:
+ num_classes: 91
+ input_size: [640, 640, 3]
+ backbone:
+ type: 'maxvit'
+ maxvit:
+ model_name: 'maxvit-small'
+ window_size: 20
+ grid_size: 20
+ scale_ratio: '20/7'
+ survival_prob: 0.7
+ decoder:
+ fpn:
+ fusion_type: 'concat'
+ type: 'fpn'
+ head:
+ level: 3
+ losses:
+ l2_weight_decay: 0
+ top_k_percent_pixels: 1.0
+ train_data:
+ output_size: [640, 640]
+ global_batch_size: 32
+ dtype: 'bfloat16'
+ aug_rand_hflip: true
+ aug_scale_max: 1.5
+ aug_scale_min: 0.5
+ validation_data:
+ output_size: [640, 640]
+ global_batch_size: 32
+ dtype: 'bfloat16'
+ groundtruth_padded_size: [640, 640]
+trainer:
+ optimizer_config:
+ learning_rate:
+ type: cosine
+ cosine:
+ decay_steps: 64000
+ initial_learning_rate: 0.000001
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ weight_decay_rate: 0.0001
+ type: adamw
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 4000
+ type: linear
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ steps_per_loop: 200
+ summary_interval: 200
+ train_steps: 64000
+ checkpoint_interval: 200
+ validation_interval: 200
+ validation_steps: 39
diff --git a/official/projects/maxvit/configs/experiments/seg_pascal_maxvits_i512.yaml b/official/projects/maxvit/configs/experiments/seg_pascal_maxvits_i512.yaml
new file mode 100644
index 00000000000..a20a33da0dd
--- /dev/null
+++ b/official/projects/maxvit/configs/experiments/seg_pascal_maxvits_i512.yaml
@@ -0,0 +1,68 @@
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: 'Please provide'
+ init_checkpoint_modules: ['backbone']
+ model:
+ num_classes: 21
+ input_size: [512, 512, 3]
+ backbone:
+ type: 'maxvit'
+ maxvit:
+ model_name: 'maxvit-small'
+ window_size: 16
+ grid_size: 16
+ scale_ratio: '16/7'
+ survival_prob: 0.7
+ decoder:
+ fpn:
+ fusion_type: 'sum'
+ type: 'fpn'
+ head:
+ level: 3
+ losses:
+ l2_weight_decay: 0
+ top_k_percent_pixels: 1.0
+ train_data:
+ output_size: [512, 512]
+ global_batch_size: 32
+ dtype: 'bfloat16'
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ output_size: [512, 512]
+ global_batch_size: 32
+ dtype: 'bfloat16'
+ groundtruth_padded_size: [512, 512]
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.00001
+ alpha: 0.01
+ optimizer:
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ weight_decay_rate: 0.0001
+ type: adamw
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 500
+ type: linear
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ steps_per_loop: 330
+ summary_interval: 330
+ train_steps: 20000
+ validation_interval: 330
+ checkpoint_interval: 330
+ validation_steps: 45
diff --git a/official/projects/maxvit/configs/image_classification.py b/official/projects/maxvit/configs/image_classification.py
new file mode 100644
index 00000000000..b3cf92fd1d0
--- /dev/null
+++ b/official/projects/maxvit/configs/image_classification.py
@@ -0,0 +1,58 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""MaxViT Image classification configuration definition."""
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling.optimization.configs import optimization_config
+from official.projects.maxvit.configs import backbones
+from official.vision.configs import image_classification as img_cls_cfg
+
+
+@exp_factory.register_config_factory('maxvit_imagenet')
+def maxvit_imagenet() -> cfg.ExperimentConfig:
+ """Returns MaxViT-Tiny on imagenet-1k.
+
+ Expected to be trained on DF 4x4 or bigger. Can eval on DF 4x2.
+
+ Returns:
+ The full experiment config.
+ """
+ # Reuse ViT deit pretraining config.
+ exp = img_cls_cfg.image_classification_imagenet_deit_pretrain()
+ exp.task.model = img_cls_cfg.ImageClassificationModel(
+ num_classes=1001,
+ input_size=[224, 224, 3],
+ kernel_initializer='glorot_uniform',
+ backbone=backbones.Backbone(
+ type='maxvit',
+ maxvit=backbones.MaxViT(
+ model_name='maxvit-tiny', representation_size=768
+ ),
+ ),
+ norm_activation=img_cls_cfg.common.NormActivation(activation='relu'),
+ )
+
+ exp.task.train_data.aug_type.randaug.num_layers = 2
+ exp.task.train_data.aug_type.randaug.magnitude = 15
+ exp.runtime.mixed_precision_dtype = 'bfloat16'
+ exp.trainer.optimizer_config.optimizer.adamw.gradient_clip_norm = 0.0
+ exp.trainer.optimizer_config.warmup.linear.warmup_steps = 10000
+ exp.trainer.optimizer_config.ema = optimization_config.opt_cfg.EMAConfig(
+ average_decay=0.9999,
+ trainable_weights_only=False,
+ )
+
+ return exp
diff --git a/official/projects/maxvit/configs/image_classification_test.py b/official/projects/maxvit/configs/image_classification_test.py
new file mode 100644
index 00000000000..4a91443b2b8
--- /dev/null
+++ b/official/projects/maxvit/configs/image_classification_test.py
@@ -0,0 +1,45 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.maxvit.configs import image_classification # pylint:disable=unused-import
+from official.vision.configs import image_classification as img_cls_config
+
+
+class MaxViTImageClassificationConfigTest(tf.test.TestCase):
+
+ def test_maxvit_build_model(self):
+ config = exp_factory.get_exp_config('maxvit_imagenet')
+
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(
+ config.task, img_cls_config.ImageClassificationTask
+ )
+ self.assertIsInstance(
+ config.task.model, img_cls_config.ImageClassificationModel
+ )
+ self.assertIsInstance(
+ config.task.train_data, img_cls_config.DataConfig
+ )
+ config.validate()
+ config.task.train_data.is_training = None
+ with self.assertRaises(KeyError):
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/maxvit/configs/rcnn.py b/official/projects/maxvit/configs/rcnn.py
new file mode 100644
index 00000000000..8df30f272a1
--- /dev/null
+++ b/official/projects/maxvit/configs/rcnn.py
@@ -0,0 +1,132 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Mask R-CNN configuration definition."""
+import os
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling.optimization.configs import optimization_config
+from official.projects.maxvit.configs import backbones
+from official.vision.configs import common
+from official.vision.configs import decoders
+from official.vision.configs import maskrcnn
+
+
+Parser = maskrcnn.Parser
+Anchor = maskrcnn.Anchor
+Losses = maskrcnn.Losses
+ROISampler = maskrcnn.ROISampler
+DetectionHead = maskrcnn.DetectionHead
+DataConfig = maskrcnn.DataConfig
+MaskRCNN = maskrcnn.MaskRCNN
+MaskRCNNTask = maskrcnn.MaskRCNNTask
+
+
+COCO_INPUT_PATH_BASE = (
+ '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco'
+)
+
+
+@exp_factory.register_config_factory('rcnn_maxvit_coco')
+def rcnn_maxvit_coco() -> cfg.ExperimentConfig:
+ """COCO object detection with MaxViT and Cascade R-CNN."""
+ steps_per_epoch = 1848 # based on 463 steps @ bs=256
+ train_batch_size = 256
+ coco_val_samples = 5000
+ eval_batch_size = 64
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(mixed_precision_dtype='bfloat16'),
+ task=MaskRCNNTask(
+ annotation_file=os.path.join(COCO_INPUT_PATH_BASE,
+ 'instances_val2017.json'),
+ model=MaskRCNN(
+ anchor=Anchor(num_scales=3, anchor_size=3.0),
+ backbone=backbones.Backbone(
+ type='maxvit',
+ maxvit=backbones.MaxViT(model_name='maxvit-base')
+ ),
+ decoder=decoders.Decoder(type='fpn', fpn=decoders.FPN()),
+ num_classes=91,
+ input_size=[640, 640, 3],
+ include_mask=True,
+ roi_sampler=ROISampler(
+ cascade_iou_thresholds=[0.7], foreground_iou_threshold=0.6),
+ detection_head=DetectionHead(
+ cascade_class_ensemble=True, class_agnostic_bbox_pred=True),
+ norm_activation=common.NormActivation(
+ use_sync_bn=True,
+ activation='relu',
+ norm_epsilon=0.001,
+ norm_momentum=0.99),
+ min_level=3,
+ max_level=7,
+ ),
+ losses=Losses(l2_weight_decay=0.0),
+ train_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ parser=Parser(
+ aug_rand_hflip=True, aug_scale_min=0.1, aug_scale_max=2.5)),
+ validation_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=True)),
+ trainer=cfg.TrainerConfig(
+ train_steps=90000,
+ validation_steps=coco_val_samples // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ best_checkpoint_export_subdir='best_ckpt',
+ best_checkpoint_eval_metric='AP',
+ checkpoint_interval=steps_per_epoch * 4,
+ optimizer_config=optimization_config.OptimizationConfig({
+ 'ema': {
+ 'average_decay': 0.9998,
+ 'trainable_weights_only': False,
+ },
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.0001,
+ 'beta_1': 0.9,
+ 'beta_2': 0.999,
+ 'include_in_weight_decay': r'.*(kernel|weight):0$',
+ },
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'decay_steps': 90000,
+ 'initial_learning_rate': 0.0001,
+ 'alpha': 0.03,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 6000,
+ 'warmup_learning_rate': 0.,
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
diff --git a/official/projects/maxvit/configs/rcnn_test.py b/official/projects/maxvit/configs/rcnn_test.py
new file mode 100644
index 00000000000..c91120cafb7
--- /dev/null
+++ b/official/projects/maxvit/configs/rcnn_test.py
@@ -0,0 +1,37 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.maxvit.configs import rcnn as exp_cfg
+
+
+class MaskRCNNConfigTest(tf.test.TestCase):
+
+ def test_maskrcnn_configs(self):
+ config = exp_factory.get_exp_config('rcnn_maxvit_coco')
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task, exp_cfg.MaskRCNNTask)
+ self.assertIsInstance(config.task.model, exp_cfg.MaskRCNN)
+ self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
+ config.validate()
+ config.task.train_data.is_training = None
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/maxvit/configs/retinanet.py b/official/projects/maxvit/configs/retinanet.py
new file mode 100644
index 00000000000..fa79608afce
--- /dev/null
+++ b/official/projects/maxvit/configs/retinanet.py
@@ -0,0 +1,39 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""RetinaNet configuration definition."""
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.maxvit.configs import backbones
+from official.vision.configs import retinanet
+
+
+@exp_factory.register_config_factory('retinanet_maxvit_coco')
+def retinanet_maxvit_coco() -> cfg.ExperimentConfig:
+ """COCO object detection with RetinaNet using MaxViT backbone."""
+ config = retinanet.retinanet_resnetfpn_coco()
+ config.task.model.backbone = backbones.Backbone(
+ type='maxvit', maxvit=backbones.MaxViT(
+ model_name='maxvit-base',
+ window_size=20,
+ grid_size=20,
+ scale_ratio='20/7',
+ survival_prob=0.7,
+ )
+ )
+ config.task.validation_data.global_batch_size = 32
+ config.trainer.validation_steps = 156
+ config.trainer.validation_interval = 1560
+ return config
diff --git a/official/projects/maxvit/configs/retinanet_test.py b/official/projects/maxvit/configs/retinanet_test.py
new file mode 100644
index 00000000000..c9b274fd02d
--- /dev/null
+++ b/official/projects/maxvit/configs/retinanet_test.py
@@ -0,0 +1,44 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for retinanet."""
+# pylint: disable=unused-import
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.maxvit.configs import retinanet
+from official.vision.configs import retinanet as exp_cfg
+
+
+class RetinaNetConfigTest(tf.test.TestCase, parameterized.TestCase):
+
+ def test_retinanet_configs(self):
+ config = exp_factory.get_exp_config('retinanet_maxvit_coco')
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task, exp_cfg.RetinaNetTask)
+ self.assertIsInstance(config.task.model, exp_cfg.RetinaNet)
+ self.assertIsInstance(
+ config.task.model.backbone.maxvit, retinanet.backbones.MaxViT
+ )
+ self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
+ config.validate()
+ config.task.train_data.is_training = None
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
+ config.validate()
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/maxvit/configs/semantic_segmentation.py b/official/projects/maxvit/configs/semantic_segmentation.py
new file mode 100644
index 00000000000..0414b40d033
--- /dev/null
+++ b/official/projects/maxvit/configs/semantic_segmentation.py
@@ -0,0 +1,242 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Semantic segmentation configuration definition."""
+import os
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.projects.maxvit.configs import backbones
+from official.vision.configs import common
+from official.vision.configs import decoders
+from official.vision.configs import semantic_segmentation
+
+DataConfig = semantic_segmentation.DataConfig
+Losses = semantic_segmentation.Losses
+Evaluation = semantic_segmentation.Evaluation
+SegmentationHead = semantic_segmentation.SegmentationHead
+SemanticSegmentationModel = semantic_segmentation.SemanticSegmentationModel
+SemanticSegmentationTask = semantic_segmentation.SemanticSegmentationTask
+
+
+# PASCAL VOC 2012 Dataset
+PASCAL_TRAIN_EXAMPLES = 10582
+PASCAL_VAL_EXAMPLES = 1449
+PASCAL_INPUT_PATH_BASE = 'gs://**/pascal_voc_seg'
+
+
+@exp_factory.register_config_factory('maxvit_seg_pascal')
+def maxvit_seg_pascal() -> cfg.ExperimentConfig:
+ """Image segmentation on Pascal VOC with MaxViT."""
+ train_batch_size = 32
+ eval_batch_size = 32
+ steps_per_epoch = PASCAL_TRAIN_EXAMPLES // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=SemanticSegmentationTask(
+ model=SemanticSegmentationModel(
+ num_classes=21,
+ input_size=[512, 512, 3],
+ min_level=3,
+ max_level=7,
+ backbone=backbones.Backbone(
+ type='maxvit',
+ maxvit=backbones.MaxViT(
+ model_name='maxvit-tiny',
+ window_size=16,
+ grid_size=16,
+ scale_ratio='16/7',
+ ),
+ ),
+ decoder=decoders.Decoder(type='fpn', fpn=decoders.FPN()),
+ head=SegmentationHead(level=3, num_convs=3),
+ norm_activation=common.NormActivation(
+ use_sync_bn=True,
+ activation='relu',
+ norm_epsilon=0.001,
+ norm_momentum=0.99,
+ ),
+ ),
+ losses=Losses(l2_weight_decay=1e-5, top_k_percent_pixels=1.0),
+ train_data=DataConfig(
+ input_path=os.path.join(PASCAL_INPUT_PATH_BASE, 'train_aug*'),
+ output_size=[512, 512],
+ is_training=True,
+ global_batch_size=train_batch_size,
+ aug_rand_hflip=True,
+ aug_scale_min=0.2,
+ aug_scale_max=1.5,
+ ),
+ validation_data=DataConfig(
+ input_path=os.path.join(PASCAL_INPUT_PATH_BASE, 'val*'),
+ output_size=[512, 512],
+ is_training=True,
+ global_batch_size=eval_batch_size,
+ resize_eval_groundtruth=True,
+ groundtruth_padded_size=[512, 512],
+ drop_remainder=True,
+ ),
+ ),
+ trainer=cfg.TrainerConfig(
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ train_steps=20000,
+ validation_steps=PASCAL_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'ema': {
+ 'average_decay': 0.9998,
+ 'trainable_weights_only': False,
+ },
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'beta_1': 0.9,
+ 'beta_2': 0.999,
+ 'weight_decay_rate': 0.0001,
+ 'include_in_weight_decay': r'.*(kernel|weight):0$',
+ },
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.0005,
+ 'decay_steps': 20000,
+ 'alpha': 0.03,
+ },
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 500,
+ 'warmup_learning_rate': 0,
+ },
+ },
+ }),
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ],
+ )
+
+ return config
+
+
+# COCO segmentation.
+COCO_TRAIN_EXAMPLES = 25600
+COCO_VAL_EXAMPLES = 5000
+COCO_INPUT_PATH_BASE = 'mscoco'
+
+
+@exp_factory.register_config_factory('maxvit_seg_coco')
+def maxvit_seg_coco() -> cfg.ExperimentConfig:
+ """Image segmentation on COCO with MaxViT."""
+ train_batch_size = 32
+ eval_batch_size = 32
+ steps_per_epoch = COCO_TRAIN_EXAMPLES // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=SemanticSegmentationTask(
+ model=SemanticSegmentationModel(
+ num_classes=91,
+ input_size=[640, 640, 3],
+ backbone=backbones.Backbone(
+ type='maxvit',
+ maxvit=backbones.MaxViT(
+ model_name='maxvit-tiny',
+ window_size=20,
+ grid_size=20,
+ scale_ratio='20/7',
+ ),
+ ),
+ decoder=decoders.Decoder(type='fpn', fpn=decoders.FPN()),
+ head=SegmentationHead(level=3, num_convs=3),
+ norm_activation=common.NormActivation(
+ use_sync_bn=True,
+ activation='relu',
+ norm_epsilon=0.001,
+ norm_momentum=0.99,
+ ),
+ ),
+ losses=Losses(l2_weight_decay=1e-5, top_k_percent_pixels=1.0),
+ train_data=DataConfig(
+ input_path=os.path.join(
+ COCO_INPUT_PATH_BASE,
+ 'mscoco_alltasks_trainvalminusminival2014*',
+ ),
+ output_size=[640, 640],
+ is_training=True,
+ global_batch_size=train_batch_size,
+ aug_rand_hflip=True,
+ aug_scale_min=0.2,
+ aug_scale_max=2.0,
+ ),
+ validation_data=DataConfig(
+ input_path=os.path.join(
+ COCO_INPUT_PATH_BASE, 'mscoco_alltasks_minival2014*'
+ ),
+ output_size=[640, 640],
+ is_training=True,
+ global_batch_size=eval_batch_size,
+ resize_eval_groundtruth=True,
+ groundtruth_padded_size=[640, 640],
+ drop_remainder=True,
+ ),
+ ),
+ trainer=cfg.TrainerConfig(
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ train_steps=64000,
+ validation_steps=COCO_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'ema': {
+ 'average_decay': 0.9998,
+ 'trainable_weights_only': False,
+ },
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'beta_1': 0.9,
+ 'beta_2': 0.999,
+ 'weight_decay_rate': 0.00001,
+ 'include_in_weight_decay': r'.*(kernel|weight):0$',
+ },
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.00005,
+ 'decay_steps': 64000,
+ 'alpha': 0.03,
+ },
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 1600,
+ 'warmup_learning_rate': 0,
+ },
+ },
+ }),
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ],
+ )
+
+ return config
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/configs/panoptic_maskrcnn_test.py b/official/projects/maxvit/configs/semantic_segmentation_test.py
similarity index 61%
rename from official/vision/beta/projects/panoptic_maskrcnn/configs/panoptic_maskrcnn_test.py
rename to official/projects/maxvit/configs/semantic_segmentation_test.py
index 442d07bb309..908c63173bb 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/configs/panoptic_maskrcnn_test.py
+++ b/official/projects/maxvit/configs/semantic_segmentation_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -12,30 +12,29 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-"""Tests for panoptic maskrcnn config."""
-# pylint: disable=unused-import
+
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+# pylint: disable=unused-import
+from official import vision
from official.core import config_definitions as cfg
from official.core import exp_factory
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as exp_cfg
+from official.projects.maxvit.configs import semantic_segmentation as exp_cfg
-class PanopticMaskRCNNConfigTest(tf.test.TestCase, parameterized.TestCase):
+class ImageSegmentationConfigTest(tf.test.TestCase, parameterized.TestCase):
- @parameterized.parameters(
- ('panoptic_fpn_coco',),
- )
- def test_panoptic_maskrcnn_configs(self, config_name):
+ @parameterized.parameters(('maxvit_seg_pascal',),
+ ('maxvit_seg_coco',))
+ def test_semantic_segmentation_configs(self, config_name):
config = exp_factory.get_exp_config(config_name)
self.assertIsInstance(config, cfg.ExperimentConfig)
- self.assertIsInstance(config.task, exp_cfg.PanopticMaskRCNNTask)
- self.assertIsInstance(config.task.model, exp_cfg.PanopticMaskRCNN)
+ self.assertIsInstance(config.task, exp_cfg.SemanticSegmentationTask)
self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
config.validate()
config.task.train_data.is_training = None
- with self.assertRaisesRegex(KeyError, 'Found inconsistncy between key'):
+ with self.assertRaises(KeyError):
config.validate()
diff --git a/official/projects/maxvit/docs/i21k_jft_results.png b/official/projects/maxvit/docs/i21k_jft_results.png
new file mode 100644
index 00000000000..d2332e0e9cd
Binary files /dev/null and b/official/projects/maxvit/docs/i21k_jft_results.png differ
diff --git a/official/projects/maxvit/docs/imagenet_results.png b/official/projects/maxvit/docs/imagenet_results.png
new file mode 100644
index 00000000000..0a106e694a3
Binary files /dev/null and b/official/projects/maxvit/docs/imagenet_results.png differ
diff --git a/official/projects/maxvit/docs/maxvit_arch.png b/official/projects/maxvit/docs/maxvit_arch.png
new file mode 100644
index 00000000000..d759bb1675c
Binary files /dev/null and b/official/projects/maxvit/docs/maxvit_arch.png differ
diff --git a/official/projects/maxvit/modeling/__init__.py b/official/projects/maxvit/modeling/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/maxvit/modeling/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/maxvit/modeling/common_ops.py b/official/projects/maxvit/modeling/common_ops.py
new file mode 100644
index 00000000000..ec65bbe2931
--- /dev/null
+++ b/official/projects/maxvit/modeling/common_ops.py
@@ -0,0 +1,265 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Common operations."""
+
+import functools
+import math
+from typing import Optional
+
+from absl import logging
+import numpy as np
+import tensorflow as tf, tf_keras
+
+
+def activation_fn(features: tf.Tensor, act_fn: str):
+ """Customized non-linear activation type."""
+ if act_fn in ('silu', 'swish'):
+ return tf.nn.swish(features)
+ elif act_fn == 'silu_native':
+ return features * tf.sigmoid(features)
+ elif act_fn == 'hswish':
+ return features * tf.nn.relu6(features + 3) / 6
+ elif act_fn == 'relu':
+ return tf.nn.relu(features)
+ elif act_fn == 'relu6':
+ return tf.nn.relu6(features)
+ elif act_fn == 'elu':
+ return tf.nn.elu(features)
+ elif act_fn == 'leaky_relu':
+ return tf.nn.leaky_relu(features)
+ elif act_fn == 'selu':
+ return tf.nn.selu(features)
+ elif act_fn == 'mish':
+ return features * tf.math.tanh(tf.math.softplus(features))
+ elif act_fn == 'gelu':
+ return (
+ 0.5
+ * features
+ * (
+ 1
+ + tf.tanh(
+ np.sqrt(2 / np.pi) * (features + 0.044715 * tf.pow(features, 3))
+ )
+ )
+ )
+ else:
+ raise ValueError('Unsupported act_fn {}'.format(act_fn))
+
+
+def get_act_fn(act_fn):
+ if act_fn is None:
+ act_fn = 'gelu'
+ if isinstance(act_fn, str):
+ return functools.partial(activation_fn, act_fn=act_fn)
+ elif callable(act_fn):
+ return act_fn
+ else:
+ raise ValueError('Unsupported act_fn %s.' % act_fn)
+
+
+def pooling_2d(inputs, pool_type, stride, **kwargs):
+ """Perform 2D pooling."""
+ if stride > 1:
+ if pool_type == 'max':
+ pool_op = tf_keras.layers.MaxPool2D
+ elif pool_type == 'avg':
+ pool_op = tf_keras.layers.AveragePooling2D
+ else:
+ raise ValueError('Unsurpported pool_type %s' % pool_type)
+ output = pool_op(
+ pool_size=(stride, stride), strides=(stride, stride), **kwargs
+ )(inputs)
+ else:
+ output = inputs
+ return output
+
+
+def drop_connect(inputs, training, survival_prob):
+ """Drop the entire conv with given survival probability."""
+ # "Deep Networks with Stochastic Depth", https://arxiv.org/pdf/1603.09382.pdf
+ if not training:
+ return inputs
+
+ # Compute tensor.
+ batch_size = tf.shape(inputs)[0]
+ random_tensor = survival_prob
+ random_tensor += tf.random.uniform([batch_size], dtype=inputs.dtype)
+ for _ in range(inputs.shape.rank - 1):
+ random_tensor = tf.expand_dims(random_tensor, axis=-1)
+ binary_tensor = tf.floor(random_tensor)
+ # Unlike conventional way that multiply survival_prob at test time, here we
+ # divide survival_prob at training time, such that no addition compute is
+ # needed at test time.
+ output = inputs / survival_prob * binary_tensor
+ return output
+
+
+def residual_add(residual, shortcut, survival_prob, training):
+ """Combine residual and shortcut."""
+ if survival_prob is not None and 0 < survival_prob < 1:
+ residual = drop_connect(residual, training, survival_prob)
+ return shortcut + residual
+
+
+def maybe_reshape_to_2d(x, height=None):
+ """Reshape tensor to 2d if not already 2d."""
+ if x.shape.rank == 3:
+ _, length, num_channel = x.shape.as_list()
+ if height is None:
+ height = int(np.sqrt(length))
+ else:
+ assert length % height == 0
+ width = length // height
+ logging.debug(
+ 'Reshape %s -> %s', [length, num_channel], [height, width, num_channel]
+ )
+ return tf.reshape(x, [-1, height, width, num_channel])
+ elif x.shape.rank == 4:
+ return x
+ else:
+ raise ValueError('Unsupport shape {}'.format(x.shape))
+
+
+def maybe_reshape_to_1d(x):
+ """Reshape tensor to 1d if not already 1d."""
+ if x.shape.rank == 4:
+ _, h, w, num_channel = x.shape.as_list()
+ logging.debug('Reshape %s -> %s', [h, w, num_channel], [h * w, num_channel])
+ return tf.reshape(x, [-1, h * w, num_channel])
+ elif x.shape.rank == 3:
+ return x
+ else:
+ raise ValueError('Unsupport shape {}'.format(x.shape))
+
+
+def generate_lookup_tensor(
+ length: int,
+ max_relative_position: Optional[int] = None,
+ clamp_out_of_range: bool = False,
+ dtype: tf.DType = tf.float32) -> tf.Tensor:
+ """Generate a one_hot lookup tensor to reindex embeddings along one dimension.
+
+ Args:
+ length: the length to reindex to.
+ max_relative_position: the maximum relative position to consider.
+ Relative position embeddings for distances above this threshold
+ are zeroed out.
+ clamp_out_of_range: bool. Whether to clamp out of range locations to the
+ maximum relative distance. If False, the out of range locations will be
+ filled with all-zero vectors.
+ dtype: dtype for the returned lookup tensor.
+ Returns:
+ ret: [length, length, vocab_size] lookup tensor that satisfies
+ ret[n,m,v] = 1{m - n + max_relative_position = v}.
+ """
+ if max_relative_position is None:
+ max_relative_position = length - 1
+ vocab_size = 2 * max_relative_position + 1
+ ret = np.zeros((length, length, vocab_size))
+ for i in range(length):
+ for x in range(length):
+ v = x - i + max_relative_position
+ if abs(x - i) > max_relative_position:
+ if clamp_out_of_range:
+ v = np.clip(v, 0, vocab_size - 1)
+ else:
+ continue
+ ret[i, x, v] = 1
+ return tf.constant(ret, dtype)
+
+
+def reindex_2d_einsum_lookup(
+ relative_position_tensor: tf.Tensor,
+ height: int,
+ width: int,
+ max_relative_height: Optional[int] = None,
+ max_relative_width: Optional[int] = None,
+ h_axis=None) -> tf.Tensor:
+ """Reindex 2d relative position bias with 2 independent einsum lookups.
+
+ Args:
+ relative_position_tensor: tensor of shape
+ [..., vocab_height, vocab_width, ...].
+ height: height to reindex to.
+ width: width to reindex to.
+ max_relative_height: maximum relative height.
+ Position embeddings corresponding to vertical distances larger
+ than max_relative_height are zeroed out. None to disable.
+ max_relative_width: maximum relative width.
+ Position embeddings corresponding to horizontal distances larger
+ than max_relative_width are zeroed out. None to disable.
+ h_axis: Axis corresponding to vocab_height. Default to 0 if None.
+
+ Returns:
+ reindexed_bias: a Tensor of shape
+ [..., height * width, height * width, ...]
+ """
+ height_lookup = generate_lookup_tensor(
+ height, max_relative_position=max_relative_height,
+ dtype=relative_position_tensor.dtype)
+ width_lookup = generate_lookup_tensor(
+ width, max_relative_position=max_relative_width,
+ dtype=relative_position_tensor.dtype)
+
+ if h_axis is None:
+ h_axis = 0
+
+ non_spatial_rank = relative_position_tensor.shape.rank - 2
+ non_spatial_expr = ''.join(chr(ord('n') + i) for i in range(non_spatial_rank))
+ prefix = non_spatial_expr[:h_axis]
+ suffix = non_spatial_expr[h_axis:]
+
+ reindexed_tensor = tf.einsum(
+ '{0}hw{1},ixh->{0}ixw{1}'.format(prefix, suffix),
+ relative_position_tensor, height_lookup, name='height_lookup')
+ reindexed_tensor = tf.einsum(
+ '{0}ixw{1},jyw->{0}ijxy{1}'.format(prefix, suffix),
+ reindexed_tensor, width_lookup, name='width_lookup')
+
+ ret_shape = relative_position_tensor.shape.as_list()
+ ret_shape[h_axis] = height * width
+ ret_shape[h_axis + 1] = height * width
+ reindexed_tensor = tf.reshape(reindexed_tensor, ret_shape)
+
+ return reindexed_tensor
+
+
+def float32_softmax(x: tf.Tensor, *args, **kwargs) -> tf.Tensor:
+ y = tf.cast(tf.nn.softmax(tf.cast(x, tf.float32), *args, **kwargs), x.dtype)
+ return y
+
+
+def get_shape_from_length(length: int, height: int = 1, width: int = 1):
+ """Gets input 2D shape from 1D sequence length."""
+ input_height = int(math.sqrt(length * height // width))
+ input_width = input_height * width // height
+ if input_height * input_width != length:
+ raise ValueError(
+ f'Invalid sequence length: {length} or shape: ({height, width}).'
+ )
+ return (input_height, input_width)
+
+
+def absolute_position_encoding(
+ position: tf.Tensor, hidden_size: int, dtype=tf.float32) -> tf.Tensor:
+ """Create absoulte position encoding."""
+ position = tf.cast(position, dtype)
+ half_hid = hidden_size // 2
+ freq_seq = tf.cast(tf.range(half_hid), dtype=dtype)
+ inv_freq = 1 / (10000 ** (freq_seq / half_hid))
+ sinusoid = tf.einsum('S,D->SD', position, inv_freq)
+ sin = tf.sin(sinusoid)
+ cos = tf.cos(sinusoid)
+ return tf.concat([sin, cos], axis=-1)
diff --git a/official/projects/maxvit/modeling/layers.py b/official/projects/maxvit/modeling/layers.py
new file mode 100644
index 00000000000..5c4a0f6059b
--- /dev/null
+++ b/official/projects/maxvit/modeling/layers.py
@@ -0,0 +1,880 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Layers and Model class for MaxViT."""
+
+import functools
+import string
+from typing import Any, Callable, Optional, Tuple, Union
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.projects.maxvit.modeling import common_ops
+
+
+class TrailDense(tf_keras.layers.Layer):
+ """Dense module that projects multiple trailing dimensions."""
+
+ def __init__(
+ self,
+ output_trailing_dims: Union[int, Tuple[int, ...]],
+ begin_axis: int = -1,
+ use_bias: bool = True,
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'dense',
+ ):
+ super().__init__(name=name)
+
+ if isinstance(output_trailing_dims, int):
+ self._output_trailing_dims = [output_trailing_dims]
+ else:
+ assert isinstance(output_trailing_dims, (list, tuple)) and all(
+ isinstance(i, int) for i in output_trailing_dims
+ ), f'Invalid output shape: {output_trailing_dims}.'
+ self._output_trailing_dims = list(output_trailing_dims)
+ self.begin_axis = begin_axis
+ self.use_bias = use_bias
+
+ self.kernel_initializer = kernel_initializer
+ self.bias_initializer = bias_initializer
+
+ def build(self, input_shape: tf.TensorShape) -> None:
+ """Create variables and einsum expression based on input shape."""
+ # Create variables
+ weight_shape = input_shape[self.begin_axis :] + self._output_trailing_dims
+ self.weight = self.add_weight(
+ name='weight',
+ shape=weight_shape,
+ initializer=self.kernel_initializer,
+ trainable=True,
+ )
+ if self.use_bias:
+ self.bias = self.add_weight(
+ name='bias',
+ shape=self._output_trailing_dims,
+ initializer=self.bias_initializer,
+ trainable=True,
+ )
+
+ # Create einsum expression
+ input_rank = input_shape.rank
+ shared_size = self.begin_axis % input_rank
+ i_only_size = input_rank - shared_size
+ o_only_size = len(self._output_trailing_dims)
+
+ assert input_rank + o_only_size < len(
+ string.ascii_uppercase
+ ), 'Cannot use einsum as input rank + output rank > 26.'
+ einsum_str = string.ascii_uppercase[: input_rank + o_only_size]
+
+ offset = 0
+ shared_str = einsum_str[offset : offset + shared_size]
+ offset += shared_size
+ i_only_str = einsum_str[offset : offset + i_only_size]
+ offset += i_only_size
+ o_only_str = einsum_str[offset : offset + o_only_size]
+
+ input_str = f'{shared_str}{i_only_str}'
+ output_str = f'{shared_str}{o_only_str}'
+ weight_str = f'{i_only_str}{o_only_str}'
+ # Examples
+ # - For 4D tensors in conv, a common expr would be 'ABCD,DE->ABCE'.
+ # - For `q/k/v` head projection in multi-head attention with two output
+ # trailing dims, the expr is 'ABC,CDE->ABDE'
+ # - For `o` output projection in multi-head attention with begin_axis = -2,
+ # the expr is 'ABCD,CDE->ABE'
+ self.einsum_expr = f'{input_str},{weight_str}->{output_str}'
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ output = tf.einsum(self.einsum_expr, inputs, self.weight)
+ if self.use_bias:
+ output += self.bias
+ return output
+
+
+class Attention(tf_keras.layers.Layer):
+ """Multi-headed attention module."""
+
+ def __init__(
+ self,
+ hidden_size: int,
+ head_size: int,
+ input_origin_height: int = 1,
+ input_origin_width: int = 1,
+ num_heads: Optional[int] = None,
+ dropatt: float = 0.0,
+ attn_axis: int = 0,
+ rel_attn_type: Optional[str] = None,
+ scale_ratio: Optional[float] = None,
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'attention',
+ ):
+ super().__init__(name=name)
+
+ self.hidden_size = hidden_size
+ self.head_size = head_size
+ self.input_origin_height = input_origin_height
+ self.input_origin_width = input_origin_width
+ self.num_heads = num_heads or hidden_size // head_size
+ self.dropatt = dropatt
+ self.attn_axis = attn_axis
+ self.rel_attn_type = rel_attn_type
+ self.scale_ratio = scale_ratio
+
+ self.kernel_initializer = kernel_initializer
+ self.bias_initializer = bias_initializer
+
+ self._q_proj = TrailDense(
+ output_trailing_dims=(self.num_heads, self.head_size),
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='q',
+ )
+ self._k_proj = TrailDense(
+ output_trailing_dims=(self.num_heads, self.head_size),
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='k',
+ )
+ self._v_proj = TrailDense(
+ output_trailing_dims=(self.num_heads, self.head_size),
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='v',
+ )
+ self._o_proj = TrailDense(
+ output_trailing_dims=self.hidden_size,
+ begin_axis=-2,
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='o',
+ )
+
+ self.q_scale = self.head_size**-0.5
+ self.relative_bias = None
+
+ def build(self, query_shape: Any) -> None:
+ ##### Content attention
+ # Einsum expression:
+ # B = batch_size
+ # N = num_heads
+ # K = head_size
+ # S = query_len (of the given attn_axis)
+ # T = key/value_len (of the given attn_axis)
+ # [U-Z] = length of other attension axes
+ # Example for 5D query_heads, (e.g. images [B x H x W x N x K])
+ # - when attn_axis = 0 (H axis):
+ # symbols = 'U' => num_attn_dims = 2
+ # q_expr = 'BSUNK' => 'S' is inserted, prefix = 'B', suffix = 'NK'
+ # k_expr = 'BTUNK' => 'T' is inserted, prefix = 'B', suffix = 'NK'
+ # v_expr = 'BTUNK' => 'T' is inserted, prefix = 'B', suffix = 'NK'
+ # a_expr = 'BUNST' => 'N x S x T' attention map
+ num_attn_dims = query_shape.rank - 2 # -2 to account for bsz, hidden size
+ assert num_attn_dims < 6, 'Only support at most 6 attention dims.'
+ symbols = ''.join([chr(ord('U') + i) for i in range(num_attn_dims - 1)])
+ insert = lambda s, i, c: s[:i] + c + s[i:]
+ create_expr = lambda s, prefix='B', suffix='NK': prefix + s + suffix
+ self.q_expr = create_expr(insert(symbols, self.attn_axis, 'S'))
+ self.k_expr = create_expr(insert(symbols, self.attn_axis, 'T'))
+ self.v_expr = create_expr(insert(symbols, self.attn_axis, 'T'))
+ self.a_expr = create_expr(symbols, suffix='NST')
+
+ ##### Relative attention
+ if self.rel_attn_type in ['2d_multi_head', '2d_single_head']:
+ query_shape_list = query_shape.as_list()
+ if query_shape.rank == 4:
+ height, width = query_shape_list[1:3]
+ elif query_shape.rank == 3:
+ seq_len = query_shape_list[1]
+ height, width = common_ops.get_shape_from_length(
+ seq_len, self.input_origin_height, self.input_origin_width
+ )
+ if height * width != seq_len:
+ raise ValueError(
+ 'Sequence length: %s violates input size: (%s, %s).'
+ % (seq_len, height, width)
+ )
+ else:
+ raise ValueError(
+ 'Does not support relative attention for query shape: %s.'
+ % query_shape_list
+ )
+
+ if self.scale_ratio is not None:
+ scale_ratio = eval(self.scale_ratio) # pylint:disable=eval-used # pyrefly: ignore[bad-argument-type]
+ vocab_height = 2 * int(height / scale_ratio) - 1
+ vocab_width = 2 * int(width / scale_ratio) - 1
+ else:
+ vocab_height = 2 * height - 1
+ vocab_width = 2 * width - 1
+
+ if self.rel_attn_type == '2d_multi_head':
+ rel_bias_shape = [self.num_heads, vocab_height, vocab_width]
+ elif self.rel_attn_type == '2d_single_head':
+ rel_bias_shape = [vocab_height, vocab_width]
+ else:
+ raise NotImplementedError(
+ f'rel_attn_type {self.rel_attn_type} not implemented yet.'
+ )
+
+ self._feat_height = height
+ self._feat_width = width
+ self.relative_bias = self.add_weight(
+ 'relative_bias',
+ rel_bias_shape,
+ initializer=self.kernel_initializer,
+ trainable=True,
+ )
+
+ def call(
+ self,
+ query: tf.Tensor,
+ training: bool,
+ context: Optional[tf.Tensor] = None,
+ attn_mask: Optional[tf.Tensor] = None,
+ ) -> tf.Tensor:
+ if context is None:
+ context = query
+
+ q_heads = self._q_proj(query)
+ k_heads = self._k_proj(context)
+ v_heads = self._v_proj(context)
+ q_heads *= self.q_scale
+
+ # attention
+ attn_logits = tf.einsum(
+ f'{self.q_expr},{self.k_expr}->{self.a_expr}', q_heads, k_heads
+ )
+
+ if self.relative_bias is not None:
+ if self.rel_attn_type == '2d_multi_head':
+ h_axis = 1
+ else:
+ h_axis = 0
+
+ if self.scale_ratio is not None:
+ src_shape = self.relative_bias.shape.as_list()
+ relative_bias = tf.expand_dims(self.relative_bias, axis=-1)
+ relative_bias = tf.image.resize(
+ relative_bias, [2 * self._feat_height - 1, 2 * self._feat_width - 1]
+ )
+ relative_bias = tf.cast(
+ tf.squeeze(relative_bias, axis=-1), self.compute_dtype
+ )
+ tgt_shape = relative_bias.shape.as_list()
+ logging.info(
+ 'Bilinear resize relative position bias %s -> %s.',
+ src_shape,
+ tgt_shape,
+ )
+ else:
+ relative_bias = tf.cast(self.relative_bias, self.compute_dtype)
+
+ reindexed_bias = common_ops.reindex_2d_einsum_lookup(
+ relative_position_tensor=relative_bias,
+ height=self._feat_height,
+ width=self._feat_width,
+ max_relative_height=self._feat_height - 1,
+ max_relative_width=self._feat_width - 1,
+ h_axis=h_axis,
+ )
+ attn_logits += reindexed_bias
+
+ if attn_mask is not None:
+ # attn_mask: 1.0 means CAN attend, 0.0 means CANNOT attend
+ attn_logits += (1.0 - attn_mask) * attn_logits.dtype.min
+
+ attn_probs = common_ops.float32_softmax(attn_logits, axis=-1)
+ if self.dropatt:
+ attn_probs = tf_keras.layers.Dropout(self.dropatt, name='attn_prob_drop')(
+ attn_probs, training=training
+ )
+
+ attn_out = tf.einsum(
+ f'{self.a_expr},{self.v_expr}->{self.q_expr}', attn_probs, v_heads
+ )
+ output = self._o_proj(attn_out)
+
+ return output
+
+
+class FFN(tf_keras.layers.Layer):
+ """Positionwise feed-forward network."""
+
+ def __init__(
+ self,
+ hidden_size: int,
+ dropout: float = 0.0,
+ expansion_rate: int = 4,
+ activation: str = 'gelu',
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'ffn',
+ ):
+ super().__init__(name=name)
+
+ self.hidden_size = hidden_size
+ self.expansion_rate = expansion_rate
+ self.expanded_size = self.hidden_size * self.expansion_rate
+ self.dropout = dropout
+ self.activation = activation
+
+ self._expand_dense = TrailDense(
+ output_trailing_dims=self.expanded_size,
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='expand_dense',
+ )
+ self._shrink_dense = TrailDense(
+ output_trailing_dims=self.hidden_size,
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='shrink_dense',
+ )
+ self._activation_fn = common_ops.get_act_fn(self.activation)
+
+ def call(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ output = inputs
+ output = self._expand_dense(output)
+ output = self._activation_fn(output)
+ if self.dropout:
+ output = tf_keras.layers.Dropout(self.dropout, name='nonlinearity_drop')(
+ output, training=training
+ )
+ output = self._shrink_dense(output)
+
+ return output
+
+
+class TransformerBlock(tf_keras.layers.Layer):
+ """Transformer block = Attention + FFN."""
+
+ def __init__(
+ self,
+ hidden_size: int,
+ head_size: int,
+ input_origin_height: int = 1,
+ input_origin_width: int = 1,
+ num_heads: Optional[int] = None,
+ expansion_rate: int = 4,
+ activation: str = 'gelu',
+ pool_type: str = '2d:avg',
+ pool_stride: int = 1,
+ pool_query_only: bool = False,
+ dropatt: Optional[Union[float, tf.Tensor]] = None,
+ dropout: Optional[Union[float, tf.Tensor]] = None,
+ rel_attn_type: Optional[str] = None,
+ scale_ratio: Optional[str] = None,
+ survival_prob: Optional[Union[float, tf.Tensor]] = None,
+ ln_epsilon: float = 1e-5,
+ ln_dtype: Optional[tf.DType] = None,
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'transformer',
+ ) -> None:
+ super().__init__(name=name)
+
+ self._hidden_size = hidden_size
+ self._head_size = head_size
+ self._input_origin_height = input_origin_height
+ self._input_origin_width = input_origin_width
+ self._num_heads = num_heads
+ self._expansion_rate = expansion_rate
+ self._activation = activation
+ self._pool_type = pool_type
+ self._pool_stride = pool_stride
+ self._pool_query_only = pool_query_only
+ self._dropatt = dropatt
+ self._dropout = dropout
+ self._rel_attn_type = rel_attn_type
+ self._scale_ratio = scale_ratio
+ self._survival_prob = survival_prob
+ self._ln_epsilon = ln_epsilon
+ self._ln_dtype = ln_dtype
+ self._kernel_initializer = kernel_initializer
+ self._bias_initializer = bias_initializer
+
+ def build(self, input_shape: tf.TensorShape) -> None:
+ if len(input_shape.as_list()) == 4:
+ _, height, width, _ = input_shape.as_list()
+ elif len(input_shape.as_list()) == 3:
+ _, seq_len, _ = input_shape.as_list()
+ height, width = common_ops.get_shape_from_length(
+ seq_len, self._input_origin_height, self._input_origin_width
+ )
+ else:
+ raise ValueError(f'Unsupported input shape: {input_shape.as_list()}.')
+
+ self.height, self.width = height, width
+ input_size = input_shape.as_list()[-1]
+
+ if input_size != self._hidden_size:
+ self._shortcut_proj = TrailDense(
+ self._hidden_size,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='shortcut_proj',
+ )
+ else:
+ self._shortcut_proj = None
+
+ self._attn_layer_norm = tf_keras.layers.LayerNormalization(
+ axis=-1,
+ epsilon=self._ln_epsilon,
+ dtype=self._ln_dtype,
+ name='attn_layer_norm',
+ )
+
+ self._attention = Attention(
+ self._hidden_size,
+ self._head_size,
+ height // self._pool_stride,
+ width // self._pool_stride,
+ num_heads=self._num_heads,
+ dropatt=self._dropatt, # pyrefly: ignore[bad-argument-type]
+ rel_attn_type=self._rel_attn_type,
+ scale_ratio=self._scale_ratio, # pyrefly: ignore[bad-argument-type]
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ )
+
+ self._ffn_layer_norm = tf_keras.layers.LayerNormalization(
+ axis=-1,
+ epsilon=self._ln_epsilon,
+ dtype=self._ln_dtype,
+ name='ffn_layer_norm',
+ )
+
+ self._ffn = FFN(
+ self._hidden_size,
+ dropout=self._dropout, # pyrefly: ignore[bad-argument-type]
+ expansion_rate=self._expansion_rate,
+ activation=self._activation,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ )
+
+ def downsample(self, inputs: tf.Tensor, name: str) -> tf.Tensor:
+ output = inputs
+ if self._pool_stride > 1:
+ assert self._pool_type in [
+ '2d:avg',
+ '2d:max',
+ '1d:avg',
+ '1d:max',
+ ], f'Invalid pool_type {self._pool_type}'
+ if self._pool_type.startswith('2d'):
+ output = common_ops.maybe_reshape_to_2d(output, height=self.height)
+ output = common_ops.pooling_2d(
+ output,
+ self._pool_type.split(':')[-1],
+ self._pool_stride,
+ padding='same',
+ data_format='channels_last',
+ name=name,
+ )
+ else:
+ output = common_ops.pooling_1d(
+ output,
+ self._pool_type.split(':')[-1],
+ self._pool_stride,
+ padding='same',
+ data_format='channels_last',
+ name=name,
+ )
+ return output
+
+ def shortcut_branch(self, shortcut: tf.Tensor) -> tf.Tensor:
+ shortcut = self.downsample(shortcut, 'shortcut_pool')
+ shortcut = common_ops.maybe_reshape_to_1d(shortcut)
+ if self._shortcut_proj:
+ shortcut = self._shortcut_proj(shortcut)
+
+ return shortcut
+
+ def attn_branch(
+ self,
+ inputs: tf.Tensor,
+ training: bool,
+ attn_mask: Optional[tf.Tensor] = None,
+ ) -> tf.Tensor:
+ output = self._attn_layer_norm(inputs)
+ if self._pool_query_only:
+ query = self.downsample(output, 'query_pool')
+ query = common_ops.maybe_reshape_to_1d(query)
+ output = common_ops.maybe_reshape_to_1d(output)
+ output = self._attention(
+ query, training, context=output, attn_mask=attn_mask
+ )
+ else:
+ output = self.downsample(output, 'residual_pool')
+ output = common_ops.maybe_reshape_to_1d(output)
+ output = self._attention(output, training, attn_mask=attn_mask)
+ return output
+
+ def ffn_branch(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ output = self._ffn_layer_norm(inputs)
+ output = self._ffn(output, training)
+ return output
+
+ def call(
+ self,
+ inputs: tf.Tensor,
+ training: bool,
+ attn_mask: Optional[tf.Tensor] = None,
+ ) -> tf.Tensor:
+ logging.info(
+ 'Block %s input shape: %s, (%s).', self.name, inputs.shape, inputs.dtype
+ )
+
+ shortcut = self.shortcut_branch(inputs)
+ output = self.attn_branch(inputs, training, attn_mask)
+ if self._dropout:
+ output = tf_keras.layers.Dropout(self._dropout, name='after_attn_drop')(
+ output, training=training
+ )
+ output = common_ops.residual_add(
+ output, shortcut, self._survival_prob, training
+ )
+
+ shortcut = output
+ output = self.ffn_branch(output, training)
+ if self._dropout:
+ output = tf_keras.layers.Dropout(self._dropout, name='after_ffn_drop')(
+ output, training=training
+ )
+ output = common_ops.residual_add(
+ output, shortcut, self._survival_prob, training
+ )
+
+ return output
+
+
+class SqueezeAndExcitation(tf_keras.layers.Layer):
+ """Squeeze-and-excitation layer."""
+
+ def __init__(
+ self,
+ se_filters: int,
+ output_filters: int,
+ local_pooling: bool = False,
+ data_format: str = 'channels_last',
+ activation: str = 'swish',
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'se',
+ ):
+ super().__init__(name=name)
+
+ self._local_pooling = local_pooling
+ self._data_format = data_format
+ self._activation_fn = common_ops.get_act_fn(activation)
+
+ # Squeeze and Excitation layer.
+ self._se_reduce = tf_keras.layers.Conv2D(
+ se_filters,
+ kernel_size=[1, 1],
+ strides=[1, 1],
+ padding='same',
+ data_format=self._data_format,
+ use_bias=True,
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='reduce_conv2d',
+ )
+ self._se_expand = tf_keras.layers.Conv2D(
+ output_filters,
+ kernel_size=[1, 1],
+ strides=[1, 1],
+ padding='same',
+ data_format=self._data_format,
+ use_bias=True,
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ name='expand_conv2d',
+ )
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ h_axis, w_axis = [2, 3] if self._data_format == 'channels_first' else [1, 2]
+ if self._local_pooling:
+ se_tensor = tf.nn.avg_pool(
+ inputs,
+ ksize=[1, inputs.shape[h_axis], inputs.shape[w_axis], 1],
+ strides=[1, 1, 1, 1],
+ padding='VALID',
+ )
+ else:
+ se_tensor = tf.reduce_mean(inputs, [h_axis, w_axis], keepdims=True)
+ se_tensor = self._se_expand(self._activation_fn(self._se_reduce(se_tensor)))
+ return tf.sigmoid(se_tensor) * inputs
+
+
+def _config_batch_norm(
+ norm_type: str,
+ ln_epsilon: float = 1e-6,
+ bn_momentum: float = 0.99,
+ bn_epsilon: float = 1e-6,
+) -> Callable[..., Any]:
+ """Defines the normalization class for MbConv based on `norm_type`."""
+
+ if norm_type == 'layer_norm':
+ return functools.partial(
+ tf_keras.layers.LayerNormalization, epsilon=ln_epsilon
+ )
+ elif norm_type == 'batch_norm':
+ return functools.partial(
+ tf_keras.layers.BatchNormalization,
+ momentum=bn_momentum,
+ epsilon=bn_epsilon,
+ )
+ elif norm_type == 'sync_batch_norm':
+ return functools.partial(
+ tf_keras.layers.BatchNormalization,
+ momentum=bn_momentum,
+ epsilon=bn_epsilon,
+ synchronized=True,
+ )
+ else:
+ raise ValueError(f'Unsupported norm_type {norm_type}.')
+
+
+def _build_downsample_layer(
+ pool_type: str, pool_stride: int, data_format: str = 'channels_last'
+) -> tf_keras.layers.Layer:
+ """Builds a downsample layer for MbConv based on pool type."""
+ if pool_type == 'max':
+ return tf_keras.layers.MaxPooling2D(
+ pool_size=(pool_stride, pool_stride),
+ strides=(pool_stride, pool_stride),
+ padding='same',
+ data_format=data_format,
+ )
+ elif pool_type == 'avg':
+ return tf_keras.layers.AveragePooling2D(
+ pool_size=(pool_stride, pool_stride),
+ strides=(pool_stride, pool_stride),
+ padding='same',
+ data_format=data_format,
+ )
+ else:
+ raise ValueError(f'Unsurpported pool_type {pool_type}')
+
+
+class MBConvBlock(tf_keras.layers.Layer):
+ """Mobile Inverted Residual Bottleneck (https://arxiv.org/abs/1905.02244)."""
+
+ def __init__(
+ self,
+ hidden_size: int,
+ downsample_loc: str = 'depth_conv',
+ data_format: str = 'channels_last',
+ kernel_size: int = 3,
+ expansion_rate: int = 4,
+ se_ratio: float = 0.25,
+ activation: str = 'gelu',
+ pool_type: str = 'avg',
+ pool_stride: int = 1,
+ dropcnn: Optional[float] = None,
+ survival_prob: Optional[float] = None,
+ norm_type: str = 'sync_batch_norm',
+ bn_epsilon: float = 1e-3,
+ bn_momentum: float = 0.99,
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'mbconv',
+ ):
+ super().__init__(name=name)
+
+ self._hidden_size = hidden_size
+ self._downsample_loc = downsample_loc
+ self._data_format = data_format
+ self._kernel_size = kernel_size
+ self._expansion_rate = expansion_rate
+ self._se_ratio = se_ratio
+ self._activation = activation
+ self._pool_type = pool_type
+ self._pool_stride = pool_stride
+ self._dropcnn = dropcnn
+ self._survival_prob = survival_prob
+ self._norm_type = norm_type
+ self._bn_epsilon = bn_epsilon
+ self._bn_momentum = bn_momentum
+ self._kernel_initializer = kernel_initializer
+ self._bias_initializer = bias_initializer
+ self._pool_layer = _build_downsample_layer(
+ pool_type, pool_stride, data_format)
+ self._activation_fn = common_ops.get_act_fn(self._activation)
+
+ def build(self, input_shape: tf.TensorShape) -> None:
+ """Builds block according to the arguments."""
+
+ channel_axis = 3 if self._data_format == 'channels_last' else 1
+ input_size = input_shape[channel_axis]
+ inner_size = self._hidden_size * self._expansion_rate
+
+ norm_cls = _config_batch_norm(
+ self._norm_type,
+ bn_momentum=self._bn_momentum,
+ bn_epsilon=self._bn_epsilon,
+ )
+
+ # Shortcut projection.
+ if input_size != self._hidden_size:
+ self._shortcut_conv = tf_keras.layers.Conv2D(
+ filters=self._hidden_size,
+ kernel_size=1,
+ strides=1,
+ padding='same',
+ data_format=self._data_format,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ use_bias=True,
+ name='shortcut_conv',
+ )
+ else:
+ self._shortcut_conv = None
+
+ # Pre-Activation norm
+ self._pre_norm = norm_cls(name='pre_norm')
+
+ # Expansion phase. Called if not using fused convolutions and expansion
+ # phase is necessary.
+ if self._expansion_rate != 1:
+ self._expand_conv = tf_keras.layers.Conv2D(
+ filters=inner_size,
+ kernel_size=1,
+ strides=(
+ self._pool_stride if self._downsample_loc == 'expand_conv' else 1
+ ),
+ kernel_initializer=self._kernel_initializer,
+ padding='same',
+ data_format=self._data_format,
+ use_bias=False,
+ name='expand_conv',
+ )
+ self._expand_norm = norm_cls(name='expand_norm')
+
+ # Depth-wise convolution phase. Called if not using fused convolutions.
+ self._depthwise_conv = tf_keras.layers.DepthwiseConv2D(
+ kernel_size=self._kernel_size,
+ strides=(
+ self._pool_stride if self._downsample_loc == 'depth_conv' else 1
+ ),
+ depthwise_initializer=self._kernel_initializer,
+ padding='same',
+ data_format=self._data_format,
+ use_bias=False,
+ name='depthwise_conv',
+ )
+ self._depthwise_norm = norm_cls(name='depthwise_norm')
+
+ if self._se_ratio is not None and 0 < self._se_ratio <= 1:
+ se_filters = max(1, int(self._hidden_size * self._se_ratio))
+ self._se = SqueezeAndExcitation(
+ se_filters=se_filters,
+ output_filters=inner_size,
+ data_format=self._data_format,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='se',
+ )
+ else:
+ self._se = None
+
+ # Output phase.
+ self._shrink_conv = tf_keras.layers.Conv2D(
+ filters=self._hidden_size,
+ kernel_size=1,
+ strides=1,
+ padding='same',
+ data_format=self._data_format,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ use_bias=True,
+ name='shrink_conv',
+ )
+
+ def downsample(self, inputs: tf.Tensor, name: str) -> tf.Tensor:
+ output = inputs
+ if self._pool_stride > 1:
+ output = self._pool_layer(output)
+ return output
+
+ def shortcut_branch(self, shortcut: tf.Tensor) -> tf.Tensor:
+ shortcut = self.downsample(shortcut, name='shortcut_pool')
+ if self._shortcut_conv:
+ shortcut = self._shortcut_conv(shortcut)
+
+ return shortcut
+
+ def residual_branch(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ output = self._pre_norm(inputs, training=training)
+ if self._downsample_loc == 'inputs':
+ output = self.downsample(output, name='residual_pool')
+ if self._expansion_rate != 1:
+ output = self._expand_conv(output)
+ output = self._expand_norm(output, training=training)
+ output = self._activation_fn(output)
+ logging.debug('Expand shape: %s', output.shape)
+
+ output = self._depthwise_conv(output)
+ output = self._depthwise_norm(output, training=training)
+ output = self._activation_fn(output)
+ logging.debug('DConv shape: %s', output.shape)
+
+ if self._dropcnn:
+ output = tf_keras.layers.Dropout(self._dropcnn, 'after_dconv_drop')(
+ output, training=training
+ )
+
+ if self._se:
+ output = self._se(output)
+ self.endpoints = {'expansion_output': output}
+
+ output = self._shrink_conv(output)
+ logging.debug('Shrink shape: %s', output.shape)
+
+ return output
+
+ def call(
+ self,
+ inputs: tf.Tensor,
+ training: bool,
+ survival_prob: Optional[Union[float, tf.Tensor]] = None,
+ ) -> tf.Tensor:
+ """Implementation of call().
+
+ Args:
+ inputs: the inputs tensor.
+ training: boolean, whether the model is constructed for training.
+ survival_prob: float, between 0 to 1, drop connect rate.
+
+ Returns:
+ A output tensor.
+ """
+ logging.debug(
+ 'Block %s input shape: %s (%s)', self.name, inputs.shape, inputs.dtype
+ )
+
+ residual = self.residual_branch(inputs, training)
+ shortcut = self.shortcut_branch(inputs)
+ survival_prob = survival_prob or self._survival_prob
+ output = common_ops.residual_add(
+ residual, shortcut, survival_prob, training
+ )
+
+ return output
diff --git a/official/projects/maxvit/modeling/maxvit.py b/official/projects/maxvit/modeling/maxvit.py
new file mode 100644
index 00000000000..6fda1d867a8
--- /dev/null
+++ b/official/projects/maxvit/modeling/maxvit.py
@@ -0,0 +1,933 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# pylint: disable=logging-fstring-interpolation
+r"""MaxViT layers and model class."""
+
+import functools
+from typing import Any, Mapping, Optional, Tuple, Union
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.projects.maxvit.modeling import common_ops as ops
+from official.projects.maxvit.modeling import layers
+from official.vision.modeling.backbones import factory
+
+
+MAXVIT_SPECS = {
+ 'maxvit-tiny-for-test': dict(
+ survival_prob=None,
+ stem_hsize=(8, 8),
+ block_type=('maxvit', 'maxvit', 'maxvit', 'maxvit'),
+ num_blocks=(2, 3, 3, 2),
+ hidden_size=(32, 32, 32, 768),
+ ),
+ 'maxvit-tiny': dict(
+ survival_prob=0.8,
+ stem_hsize=(64, 64),
+ block_type=('maxvit', 'maxvit', 'maxvit', 'maxvit'),
+ num_blocks=(2, 2, 5, 2),
+ hidden_size=(64, 128, 256, 512),
+ ),
+ 'maxvit-small': dict(
+ survival_prob=0.7,
+ stem_hsize=(64, 64),
+ block_type=('maxvit', 'maxvit', 'maxvit', 'maxvit'),
+ num_blocks=(2, 2, 5, 2),
+ hidden_size=(96, 192, 384, 768),
+ ),
+ 'maxvit-base': dict(
+ survival_prob=0.6,
+ stem_hsize=(64, 64),
+ block_type=('maxvit', 'maxvit', 'maxvit', 'maxvit'),
+ num_blocks=(2, 6, 14, 2),
+ hidden_size=(96, 192, 384, 768),
+ ),
+ 'maxvit-large': dict(
+ survival_prob=0.4,
+ stem_hsize=(128, 128),
+ block_type=('maxvit', 'maxvit', 'maxvit', 'maxvit'),
+ num_blocks=(2, 6, 14, 2),
+ hidden_size=(128, 256, 512, 1024),
+ ),
+ 'maxvit-xlarge': dict(
+ survival_prob=0.3,
+ stem_hsize=(192, 192),
+ block_type=('maxvit', 'maxvit', 'maxvit', 'maxvit'),
+ num_blocks=(2, 6, 14, 2),
+ hidden_size=(192, 384, 768, 1536),
+ ),
+}
+
+
+class MaxViTBlock(tf_keras.layers.Layer):
+ """MaxViT block = MBConv + Block-Attention + FFN + Grid-Attention + FFN."""
+
+ def __init__(
+ self,
+ hidden_size: int,
+ head_size: int,
+ window_size: int,
+ grid_size: int,
+ num_heads: Optional[int] = None,
+ downsample_loc: str = 'depth_conv',
+ data_format: str = 'channels_last',
+ kernel_size: int = 3,
+ expansion_rate: int = 4,
+ se_ratio: float = 0.25,
+ activation: str = 'gelu',
+ pool_type: str = '2d:avg',
+ pool_stride: int = 1,
+ dropcnn: Optional[float] = None,
+ dropatt: Optional[Union[float, tf.Tensor]] = None,
+ dropout: Optional[Union[float, tf.Tensor]] = None,
+ rel_attn_type: Optional[str] = None,
+ scale_ratio: Optional[str] = None,
+ survival_prob: Optional[Union[float, tf.Tensor]] = None,
+ ln_epsilon: float = 1e-5,
+ ln_dtype: Optional[tf.DType] = None,
+ norm_type: str = 'sync_batch_norm',
+ bn_epsilon: float = 1e-3,
+ bn_momentum: float = 0.99,
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'maxvit_block',
+ ) -> None:
+ super().__init__(name=name)
+
+ self._hidden_size = hidden_size
+ self._head_size = head_size
+ self._window_size = window_size
+ self._grid_size = grid_size
+ self._num_heads = num_heads
+ self._downsample_loc = downsample_loc
+ self._data_format = data_format
+ self._kernel_size = kernel_size
+ self._expansion_rate = expansion_rate
+ self._se_ratio = se_ratio
+ self._dropcnn = dropcnn
+ self._activation = activation
+ self._norm_type = norm_type
+ self._bn_epsilon = bn_epsilon
+ self._bn_momentum = bn_momentum
+ self._pool_type = pool_type
+ self._pool_stride = pool_stride
+ self._dropatt = dropatt
+ self._dropout = dropout
+ self._rel_attn_type = rel_attn_type
+ self._scale_ratio = scale_ratio
+ self._survival_prob = survival_prob
+ self._ln_epsilon = ln_epsilon
+ self._ln_dtype = ln_dtype
+ self._kernel_initializer = kernel_initializer
+ self._bias_initializer = bias_initializer
+
+ def build(self, input_shape: tf.TensorShape) -> None:
+ input_size = input_shape.as_list()[-1]
+
+ if input_size != self._hidden_size:
+ self._shortcut_proj = layers.TrailDense(
+ self._hidden_size,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='shortcut_proj',
+ )
+ else:
+ self._shortcut_proj = None
+
+ self._block_attn_layer_norm = tf_keras.layers.LayerNormalization(
+ axis=-1,
+ epsilon=self._ln_epsilon,
+ dtype=self._ln_dtype,
+ name='attn_layer_norm',
+ )
+
+ self._grid_attn_layer_norm = tf_keras.layers.LayerNormalization(
+ axis=-1,
+ epsilon=self._ln_epsilon,
+ dtype=self._ln_dtype,
+ name='attn_layer_norm_1',
+ )
+
+ self._block_attention = layers.Attention(
+ self._hidden_size,
+ self._head_size,
+ num_heads=self._num_heads,
+ dropatt=self._dropatt,
+ rel_attn_type=self._rel_attn_type,
+ scale_ratio=self._scale_ratio,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='attention',
+ )
+
+ self._grid_attention = layers.Attention(
+ self._hidden_size,
+ self._head_size,
+ num_heads=self._num_heads,
+ dropatt=self._dropatt,
+ rel_attn_type=self._rel_attn_type,
+ scale_ratio=self._scale_ratio,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='attention_1',
+ )
+
+ self._block_ffn_layer_norm = tf_keras.layers.LayerNormalization(
+ axis=-1,
+ epsilon=self._ln_epsilon,
+ dtype=self._ln_dtype,
+ name='ffn_layer_norm',
+ )
+
+ self._grid_ffn_layer_norm = tf_keras.layers.LayerNormalization(
+ axis=-1,
+ epsilon=self._ln_epsilon,
+ dtype=self._ln_dtype,
+ name='ffn_layer_norm_1',
+ )
+
+ self._block_ffn = layers.FFN(
+ self._hidden_size,
+ dropout=self._dropout,
+ expansion_rate=self._expansion_rate,
+ activation=self._activation,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='ffn',
+ )
+
+ self._grid_ffn = layers.FFN(
+ self._hidden_size,
+ dropout=self._dropout,
+ expansion_rate=self._expansion_rate,
+ activation=self._activation,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='ffn_1',
+ )
+
+ self._mbconv = layers.MBConvBlock(
+ self._hidden_size,
+ downsample_loc=self._downsample_loc,
+ data_format=self._data_format,
+ kernel_size=self._kernel_size,
+ expansion_rate=self._expansion_rate,
+ se_ratio=self._se_ratio,
+ activation=self._activation,
+ pool_type='avg' if self._pool_type == '2d:avg' else 'max',
+ pool_stride=self._pool_stride,
+ dropcnn=self._dropcnn,
+ survival_prob=self._survival_prob,
+ norm_type=self._norm_type,
+ bn_epsilon=self._bn_epsilon,
+ bn_momentum=self._bn_momentum,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name='mbconv',
+ )
+
+ def downsample(self, inputs, name):
+ output = inputs
+ if self._pool_stride > 1:
+ output = ops.maybe_reshape_to_2d(output)
+ output = ops.pooling_2d(
+ output,
+ self._pool_type,
+ self._pool_stride,
+ padding='same',
+ data_format='channels_last',
+ name=name,
+ )
+ return output
+
+ def window_partition(self, features: tf.Tensor) -> tf.Tensor:
+ """Partition the input feature maps into non-overlapping windows.
+
+ Note that unsuitable feature or window sizes may be costly on TPU due to
+ padding sizes:
+ https://docs.google.com/document/d/1GojE1Q7hR2qyi0mIfnTHgERfl7Dmsj6xPQ31MQo3xUk/edit#
+
+ Args:
+ features: [B, H, W, C] feature maps.
+
+ Returns:
+ Partitioned features: [B, nH, nW, wSize, wSize, c].
+
+ Raises:
+ ValueError: If the feature map sizes are not divisible by window sizes.
+ """
+
+ _, h, w, c = features.shape
+ window_size = self._window_size
+
+ if h % window_size != 0 or w % window_size != 0:
+ raise ValueError(
+ f'Feature map sizes {(h, w)} '
+ f'not divisible by window size ({window_size}).'
+ )
+
+ features = tf.reshape(
+ features,
+ (-1, h // window_size, window_size, w // window_size, window_size, c),
+ )
+ features = tf.transpose(features, (0, 1, 3, 2, 4, 5))
+ features = tf.reshape(features, (-1, window_size, window_size, c))
+ return features
+
+ def window_stitch_back(
+ self, features: tf.Tensor, window_size: int, h: int, w: int
+ ) -> tf.Tensor:
+ """Reverse window_partition."""
+ features = tf.reshape(
+ features,
+ [
+ -1,
+ h // window_size,
+ w // window_size,
+ window_size,
+ window_size,
+ features.shape[-1],
+ ],
+ )
+ return tf.reshape(
+ tf.transpose(features, (0, 1, 3, 2, 4, 5)),
+ [-1, h, w, features.shape[-1]],
+ )
+
+ def grid_partition(self, features: tf.Tensor) -> tf.Tensor:
+ """Partition the input feature maps into non-overlapping windows.
+
+ Note that unsuitable feature or window sizes may be costly on TPU due to
+ padding sizes:
+ https://docs.google.com/document/d/1GojE1Q7hR2qyi0mIfnTHgERfl7Dmsj6xPQ31MQo3xUk/edit#
+
+ Args:
+ features: [B, H, W, C] feature maps.
+
+ Returns:
+ Partitioned features: [B, nH, nW, wSize, wSize, c].
+
+ Raises:
+ ValueError: If the feature map sizes are not divisible by window sizes.
+ """
+ _, h, w, c = features.shape
+ grid_size = self._grid_size
+ if h % grid_size != 0 or w % grid_size != 0:
+ raise ValueError(
+ f'Feature map sizes {(h, w)} '
+ f'not divisible by window size ({grid_size}).'
+ )
+ features = tf.reshape(
+ features, (-1, grid_size, h // grid_size, grid_size, w // grid_size, c)
+ )
+ features = tf.transpose(features, (0, 2, 4, 1, 3, 5))
+ features = tf.reshape(features, (-1, grid_size, grid_size, c))
+ return features
+
+ def grid_stitch_back(
+ self, features: tf.Tensor, grid_size: int, h: int, w: int
+ ) -> tf.Tensor:
+ """Reverse window_partition."""
+ features = tf.reshape(
+ features,
+ [
+ -1,
+ h // grid_size,
+ w // grid_size,
+ grid_size,
+ grid_size,
+ features.shape[-1],
+ ],
+ )
+ return tf.reshape(
+ tf.transpose(features, (0, 3, 1, 4, 2, 5)),
+ [-1, h, w, features.shape[-1]],
+ )
+
+ def block_attn_branch(
+ self, inputs: tf.Tensor, training: bool, attn_mask: tf.Tensor
+ ) -> tf.Tensor:
+ output = self._block_attn_layer_norm(inputs)
+ # If put grid-attention in front, we don't need to downsample.
+ # Apply local block-attention
+ _, h, w, _ = output.shape
+ output = self.window_partition(output)
+ output = ops.maybe_reshape_to_1d(output)
+ output = self._block_attention(output, training, attn_mask=attn_mask)
+ output = self.window_stitch_back(output, self._window_size, h, w)
+ return output
+
+ def grid_attn_branch(
+ self, inputs: tf.Tensor, training: bool, attn_mask: tf.Tensor
+ ) -> tf.Tensor:
+ output = self._grid_attn_layer_norm(inputs)
+ # output = self.downsample(output, 'residual_pool')
+ # Apply global grid
+ _, h, w, _ = output.shape
+ output = self.grid_partition(output)
+ output = ops.maybe_reshape_to_1d(output)
+ output = self._grid_attention(output, training, attn_mask=attn_mask)
+ output = self.grid_stitch_back(output, self._grid_size, h, w)
+ return output
+
+ def block_ffn_branch(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ output = self._block_ffn_layer_norm(inputs)
+ output = self._block_ffn(output, training)
+ return output
+
+ def grid_ffn_branch(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ output = self._grid_ffn_layer_norm(inputs)
+ output = self._grid_ffn(output, training)
+ return output
+
+ def mbconv_branch(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ output = self._mbconv(inputs, training=training)
+ return output
+
+ def call(
+ self,
+ inputs: tf.Tensor,
+ training: bool,
+ attn_mask: Optional[tf.Tensor] = None,
+ ) -> tf.Tensor:
+ logging.debug(
+ 'Block %s input shape: %s (%s)', self.name, inputs.shape, inputs.dtype
+ )
+
+ # MBConv
+ output = self.mbconv_branch(inputs, training)
+
+ # block self-attention
+ shortcut = output
+ output = self.block_attn_branch(output, training, attn_mask) # pyrefly: ignore[bad-argument-type]
+ if self._dropout:
+ output = tf_keras.layers.Dropout(
+ self._dropout, name='after_block_attn_drop'
+ )(output, training=training)
+ output = ops.residual_add(output, shortcut, self._survival_prob, training)
+
+ shortcut = output
+ output = self.block_ffn_branch(output, training)
+ if self._dropout:
+ output = tf_keras.layers.Dropout(
+ self._dropout, name='after_block_ffn_drop_1'
+ )(output, training=training)
+ output = ops.residual_add(output, shortcut, self._survival_prob, training)
+
+ # grid self-attention
+ shortcut = output
+ output = self.grid_attn_branch(output, training, attn_mask) # pyrefly: ignore[bad-argument-type]
+ if self._dropout:
+ output = tf_keras.layers.Dropout(
+ self._dropout, name='after_grid_attn_drop'
+ )(output, training=training)
+ output = ops.residual_add(output, shortcut, self._survival_prob, training)
+
+ shortcut = output
+ output = self.grid_ffn_branch(output, training)
+ if self._dropout:
+ output = tf_keras.layers.Dropout(
+ self._dropout, name='after_grid_ffn_drop'
+ )(output, training=training)
+ output = ops.residual_add(output, shortcut, self._survival_prob, training)
+
+ return output
+
+
+class MaxViT(tf_keras.Model):
+ """MaxViT's backbone that outputs the pre-global-pooled features."""
+
+ def __init__(
+ self,
+ block_type: Tuple[str, ...],
+ num_blocks: Tuple[int, ...],
+ hidden_size: Tuple[int, ...],
+ stem_hsize: Tuple[int, ...],
+ head_size: int = 32,
+ num_heads: Optional[int] = None,
+ dropatt: Optional[float] = None,
+ dropout: Optional[float] = None,
+ rel_attn_type: str = '2d_multi_head',
+ window_size: int = 7,
+ grid_size: int = 7,
+ scale_ratio: Optional[str] = None,
+ ln_epsilon: float = 1e-5,
+ ln_dtype: Optional[tf.DType] = None,
+ downsample_loc: str = 'depth_conv',
+ kernel_size: int = 3,
+ se_ratio: float = 0.25,
+ dropcnn: Optional[float] = None,
+ data_format: str = 'channels_last',
+ norm_type: str = 'sync_batch_norm',
+ bn_epsilon: float = 1e-3,
+ bn_momentum: float = 0.99,
+ add_pos_enc: bool = False,
+ pool_type: str = '2d:avg',
+ pool_stride: int = 2,
+ expansion_rate: int = 4,
+ activation: str = 'gelu',
+ survival_prob: Optional[float] = None,
+ survival_prob_anneal: bool = True,
+ representation_size: Optional[int] = None,
+ add_gap_layer_norm: bool = False,
+ kernel_initializer: Optional[str] = 'glorot_uniform',
+ bias_initializer: Optional[str] = 'zeros',
+ name: str = 'maxvit',
+ **kwargs,
+ ):
+ """Initializes MaxViT backbone.
+
+ Args:
+ block_type: a tuple of `str`, specify each block type.
+ num_blocks: a tuple of `int`, specify the number of blocks in each stage.
+ hidden_size: a tuple of `int`, specify hidden size of block in each stage.
+ stem_hsize: a tuple of `int`, specify the hidden size of stem network.
+ head_size: embedding size of each attention head.
+ num_heads: number of attention head.
+ dropatt: an optional float of attention dropout rate.
+ dropout: an optional float of dropping rate for dropout regularization.
+ rel_attn_type: =a `str` specify the type of relative attention head,
+ possible values are ['2d_multi_head', '2d_single_head'].
+ window_size: window size for conducting block attention module.
+ grid_size: grid size for conducting sparse global grid attention.
+ scale_ratio: a optional string for finetuning at different window size,
+ e.g. '14/7'.
+ ln_epsilon: layer normalization epsilon.
+ ln_dtype: layer normalization data type.
+ downsample_loc: location to conduct downsampleing to feature maps.
+ kernel_size: stem convoluation kernal size.
+ se_ratio: se ratio for `mbconv` block.
+ dropcnn: an optional float of CNN dropout rate.
+ data_format: image data format, usualy 'channels_last'.
+ norm_type: normalization type, one of ['batch_norm', 'sync_batch_norm',
+ 'layer_norm'].
+ bn_epsilon: batch normalization epsilon.
+ bn_momentum: batch normalization momentum.
+ add_pos_enc: if add position embedding.
+ pool_type: pooling operation type, one of ['2d:avg', '2d:max', '1d:avg',
+ '1d:max'].
+ pool_stride: pooling stride size.
+ expansion_rate: expansion rate value.
+ activation: activate function.
+ survival_prob: survival probability.
+ survival_prob_anneal: if anneal survival probability.
+ representation_size: an optional `int` of representation size.
+ add_gap_layer_norm: if add layer norm to GAP of backbone final output.
+ kernel_initializer: kernel initializer.
+ bias_initializer: bias initializer.
+ name: specify module name.
+ **kwargs: extra keyword arguments to be passed.
+ """
+
+ super().__init__(name=name)
+ self._block_type = block_type
+ self._num_blocks = num_blocks
+ self._hidden_size = hidden_size
+ self._stem_hsize = stem_hsize
+ self._head_size = head_size
+ self._num_heads = num_heads
+ self._dropatt = dropatt
+ self._dropout = dropout
+ self._rel_attn_type = rel_attn_type
+ self._window_size = window_size
+ self._grid_size = grid_size
+ self._scale_ratio = scale_ratio
+ self._ln_epsilon = ln_epsilon
+ self._ln_dtype = ln_dtype
+ self._downsample_loc = downsample_loc
+ self._kernel_size = kernel_size
+ self._se_ratio = se_ratio
+ self._dropcnn = dropcnn
+ self._data_format = data_format
+ self._norm_type = norm_type
+ self._bn_epsilon = bn_epsilon
+ self._bn_momentum = bn_momentum
+ self._add_pos_enc = add_pos_enc
+ self._pool_type = pool_type
+ self._pool_stride = pool_stride
+ self._expansion_rate = expansion_rate
+ self._activation = activation
+ self._survival_prob = survival_prob
+ self._survival_prob_anneal = survival_prob_anneal
+ self._representation_size = representation_size
+ self._add_gap_layer_norm = add_gap_layer_norm
+ self._kernel_initializer = kernel_initializer
+ self._bias_initializer = bias_initializer
+ self._output_specs = {}
+
+ def build(self, input_shape: tf.TensorShape) -> None:
+ if self._norm_type == 'layer_norm':
+ bn_class = functools.partial(
+ tf_keras.layers.LayerNormalization, epsilon=self._ln_epsilon
+ )
+ elif self._norm_type == 'batch_norm':
+ bn_class = functools.partial(
+ tf_keras.layers.BatchNormalization,
+ momentum=self._bn_momentum,
+ epsilon=self._bn_epsilon,
+ )
+ elif self._norm_type == 'sync_batch_norm':
+ bn_class = functools.partial(
+ tf_keras.layers.BatchNormalization,
+ momentum=self._bn_momentum,
+ epsilon=self._bn_epsilon,
+ synchronized=True,
+ )
+ else:
+ raise ValueError(f'Unsupported norm_type {self._norm_type}.')
+
+ _, self.height, self.width, _ = input_shape.as_list()
+ logging.info(
+ f'Build backbone with input size: ({self.height}, {self.width}).'
+ )
+
+ # Stem
+ stem_layers = []
+ for i, _ in enumerate(self._stem_hsize):
+ conv_layer = tf_keras.layers.Conv2D(
+ filters=self._stem_hsize[i],
+ kernel_size=self._kernel_size,
+ strides=2 if i == 0 else 1,
+ padding='same',
+ data_format=self._data_format,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ use_bias=True,
+ name='conv_{}'.format(i),
+ )
+ stem_layers.append(conv_layer)
+ if i < len(self._stem_hsize) - 1:
+ stem_layers.append(bn_class(name='norm_{}'.format(i)))
+ stem_layers.append(
+ tf_keras.layers.Activation(
+ ops.get_act_fn(self._activation), name=f'act_{i}'
+ )
+ )
+ self._stem = tf_keras.Sequential(layers=stem_layers, name='stem')
+
+ # Backbone
+ self._blocks = []
+ total_num_blocks = sum(self._num_blocks)
+ bid = 0
+ for i, _ in enumerate(self._block_type):
+ self._blocks.append([])
+ for j in range(self._num_blocks[i]):
+ # block name
+ block_name = f'block_{i:0>2d}_{j:0>2d}'
+
+ ##### Update per-block config
+ # No pooling if not the first block in the stage
+ if j == 0:
+ pool_stride = self._pool_stride
+ else:
+ pool_stride = 1
+
+ # anneal the survival prob
+ survival_prob = self._survival_prob
+ if survival_prob and self._survival_prob_anneal:
+ drop_rate = 1.0 - survival_prob
+ survival_prob = 1.0 - drop_rate * bid / total_num_blocks
+ logging.info(
+ '[%02d/%02d] %s survival_prob: %.4f',
+ bid,
+ total_num_blocks,
+ block_name,
+ survival_prob,
+ )
+
+ ##### Init block
+ if self._block_type[i] == 'tfm':
+ block = layers.TransformerBlock(
+ hidden_size=self._hidden_size[i],
+ head_size=self._head_size,
+ input_origin_height=self.height,
+ input_origin_width=self.width,
+ num_heads=self._num_heads,
+ expansion_rate=self._expansion_rate,
+ activation=self._activation,
+ pool_type=self._pool_type,
+ pool_stride=pool_stride,
+ dropatt=self._dropatt,
+ dropout=self._dropout,
+ rel_attn_type=self._rel_attn_type,
+ scale_ratio=self._scale_ratio,
+ survival_prob=survival_prob,
+ ln_epsilon=self._ln_epsilon,
+ ln_dtype=self._ln_dtype,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name=block_name,
+ )
+ elif self._block_type[i] == 'mbconv':
+ assert self._pool_type in ['2d:max', '2d:avg'], (
+ 'Invalid pool_type %s for MBConv block' % self._pool_type
+ )
+ pool_type = self._pool_type.split(':')[-1]
+ block = layers.MBConvBlock(
+ hidden_size=self._hidden_size[i],
+ downsample_loc=self._downsample_loc,
+ data_format=self._data_format,
+ kernel_size=self._kernel_size,
+ expansion_rate=self._expansion_rate,
+ se_ratio=self._se_ratio,
+ activation=self._activation,
+ pool_type=pool_type,
+ pool_stride=pool_stride,
+ dropcnn=self._dropcnn,
+ survival_prob=survival_prob,
+ norm_type=self._norm_type,
+ bn_epsilon=self._bn_epsilon,
+ bn_momentum=self._bn_momentum,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name=block_name,
+ )
+ elif self._block_type[i] == 'maxvit':
+ block = MaxViTBlock(
+ hidden_size=self._hidden_size[i],
+ head_size=self._head_size,
+ window_size=self._window_size,
+ grid_size=self._grid_size,
+ num_heads=self._num_heads,
+ downsample_loc=self._downsample_loc,
+ data_format=self._data_format,
+ kernel_size=self._kernel_size,
+ expansion_rate=self._expansion_rate,
+ se_ratio=self._se_ratio,
+ activation=self._activation,
+ pool_type=self._pool_type,
+ pool_stride=pool_stride,
+ dropcnn=self._dropcnn,
+ dropatt=self._dropatt,
+ dropout=self._dropout,
+ rel_attn_type=self._rel_attn_type,
+ scale_ratio=self._scale_ratio,
+ survival_prob=survival_prob,
+ ln_epsilon=self._ln_epsilon,
+ ln_dtype=self._ln_dtype,
+ norm_type=self._norm_type,
+ bn_epsilon=self._bn_epsilon,
+ bn_momentum=self._bn_momentum,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ name=block_name,
+ )
+ else:
+ raise ValueError(f'Unsupported block_type {self._block_type[i]}')
+ self._blocks[-1].append(block)
+ bid += 1
+
+ if self._representation_size and self._representation_size > 0:
+ self._dense = tf_keras.layers.Dense(
+ self._representation_size, name='pre_logits')
+ if self._add_gap_layer_norm:
+ self._final_layer_norm = tf_keras.layers.LayerNormalization(
+ epsilon=self._ln_epsilon, name='final_layer_norm')
+
+ def _add_absolute_position_encoding(self, inputs: tf.Tensor) -> tf.Tensor:
+ """Add absolute sinusoid position encoding, which is computed on the fly."""
+ output = ops.maybe_reshape_to_2d(inputs)
+ h, w = tf.shape(output)[1], tf.shape(output)[2]
+ enc_size = output.shape.as_list()[-1] // 2
+ # sinusoid positional encoding that can be generated online
+ h_seq = tf.range(-h / 2, h / 2)
+ w_seq = tf.range(-w / 2, w / 2)
+ pos_enc_h = ops.absolute_position_encoding(
+ h_seq, enc_size, dtype=output.dtype
+ )
+ pos_enc_w = ops.absolute_position_encoding(
+ w_seq, enc_size, dtype=output.dtype
+ )
+ abs_pos_enc = tf.concat(
+ [
+ tf.tile(pos_enc_h[:, None, :], [1, w, 1]),
+ tf.tile(pos_enc_w[None, :, :], [h, 1, 1]),
+ ],
+ axis=-1,
+ )
+ output += abs_pos_enc
+ if inputs.shape.rank == 3:
+ output = ops.maybe_reshape_to_1d(output)
+ return output
+
+ def call( # pytype: disable=annotation-type-mismatch
+ self, inputs: tf.Tensor, mask: Optional[Any] = None, training: bool = None # pyrefly: ignore[bad-function-definition]
+ ) -> Mapping[str, tf.Tensor]:
+ logging.info(
+ 'MaxViT inputs: shape %s, dtype %s.', inputs.shape, inputs.dtype
+ )
+ output = self._stem(inputs, training=training)
+ logging.info(
+ 'Stage 0 (stem) output: shape %s, dtype %s.', output.shape, output.dtype
+ )
+
+ endpoints = {}
+ add_pos_enc = self._add_pos_enc
+ for idx, stage_blocks in enumerate(self._blocks):
+ # Add position encoding
+ # Note: the position encoding is usually added to the input of the first
+ # transformer block. For MaxViT, it is the first block of stage 3.
+ if (isinstance(add_pos_enc, (tuple, list)) and add_pos_enc[idx]) or (
+ isinstance(add_pos_enc, bool) and add_pos_enc
+ ):
+ logging.info('Add position encoding at stage %d.', idx + 1)
+ output = self._add_absolute_position_encoding(output)
+
+ # Blocks forward
+ for block in stage_blocks:
+ output = block(output, training=training)
+
+ if self._block_type[idx] == 'tfm':
+ height, width = ops.get_shape_from_length(
+ output.shape[1], self.height, self.width
+ )
+ output = tf.reshape(output, [-1, height, width, output.shape[-1]])
+
+ endpoints[str(idx + 2)] = output
+ logging.info(
+ 'Stage %d output: feature level %s shape %s, dtype %s.',
+ idx + 1,
+ idx + 2,
+ output.shape,
+ output.dtype,
+ )
+
+ self._output_specs = {
+ idx: endpoint.get_shape() for idx, endpoint in endpoints.items()
+ }
+
+ if self._representation_size and self._representation_size > 0:
+ # Backbone's output is [batch_size, height, weight, channel_size].
+ output = tf_keras.layers.GlobalAveragePooling2D()(output)
+ # Maybe add a layer_norm after global average pooling.
+ if self._add_gap_layer_norm:
+ output = self._final_layer_norm(output)
+ endpoints['pre_logits'] = tf.nn.tanh(self._dense(output))
+
+ return endpoints
+
+ @property
+ def output_specs(self):
+ """A dict of {level: TensorShape} pairs for the model output."""
+ return self._output_specs
+
+
+def override_predefined_spec_and_build_maxvit(
+ predefined_maxvit_spec, backbone_cfg, norm_activation_config
+):
+ """Builds a MaxViT backbone.
+
+ Args:
+ predefined_maxvit_spec: a dict predefined maxvit specifications.
+ backbone_cfg: the MaxViT backbone config.
+ norm_activation_config: normalization and activation config.
+
+ Returns:
+ The built MaxViT backbone.
+ """
+ survival_prob = (
+ predefined_maxvit_spec['survival_prob']
+ if backbone_cfg.survival_prob is None
+ else backbone_cfg.survival_prob
+ )
+ stem_hsize = (
+ predefined_maxvit_spec['stem_hsize']
+ if backbone_cfg.stem_hsize is None
+ else backbone_cfg.stem_hsize
+ )
+ block_type = (
+ predefined_maxvit_spec['block_type']
+ if backbone_cfg.block_type is None
+ else backbone_cfg.block_type
+ )
+ num_blocks = (
+ predefined_maxvit_spec['num_blocks']
+ if backbone_cfg.num_blocks is None
+ else backbone_cfg.num_blocks
+ )
+ hidden_size = (
+ predefined_maxvit_spec['hidden_size']
+ if backbone_cfg.hidden_size is None
+ else backbone_cfg.hidden_size
+ )
+
+ logging.info(
+ (
+ 'Final MaxViT specs: survival_prob=%s, stem_hsize=%s, hidden_size=%s,'
+ 'block_type=%s, num_blocks=%s,.'
+ ),
+ survival_prob,
+ stem_hsize,
+ hidden_size,
+ block_type,
+ num_blocks,
+ )
+
+ return MaxViT(
+ block_type=block_type,
+ num_blocks=num_blocks,
+ hidden_size=hidden_size,
+ stem_hsize=stem_hsize,
+ head_size=backbone_cfg.head_size,
+ dropatt=backbone_cfg.dropatt,
+ dropout=backbone_cfg.dropout,
+ rel_attn_type=backbone_cfg.rel_attn_type,
+ window_size=backbone_cfg.window_size,
+ grid_size=backbone_cfg.grid_size,
+ scale_ratio=backbone_cfg.scale_ratio,
+ ln_epsilon=backbone_cfg.ln_epsilon,
+ ln_dtype=backbone_cfg.ln_dtype,
+ downsample_loc=backbone_cfg.downsample_loc,
+ kernel_size=backbone_cfg.kernel_size,
+ se_ratio=backbone_cfg.se_ratio,
+ dropcnn=backbone_cfg.dropcnn,
+ data_format=backbone_cfg.data_format,
+ norm_type=backbone_cfg.norm_type,
+ bn_epsilon=norm_activation_config.norm_epsilon,
+ bn_momentum=norm_activation_config.norm_momentum,
+ add_pos_enc=backbone_cfg.add_pos_enc,
+ pool_type=backbone_cfg.pool_type,
+ pool_stride=backbone_cfg.pool_stride,
+ expansion_rate=backbone_cfg.expansion_rate,
+ activation=norm_activation_config.activation,
+ survival_prob=survival_prob,
+ survival_prob_anneal=backbone_cfg.survival_prob_anneal,
+ representation_size=backbone_cfg.representation_size,
+ add_gap_layer_norm=backbone_cfg.add_gap_layer_norm,
+ kernel_initializer=backbone_cfg.kernel_initializer,
+ bias_initializer=backbone_cfg.bias_initializer,
+ )
+
+
+@factory.register_backbone_builder('maxvit')
+def build_maxvit(
+ input_specs,
+ backbone_config,
+ norm_activation_config,
+ l2_regularizer=None,
+):
+ """Builds a MaxViT backbone."""
+ del l2_regularizer
+ backbone_cfg = backbone_config.get()
+ maxvit = override_predefined_spec_and_build_maxvit(
+ predefined_maxvit_spec=MAXVIT_SPECS[backbone_cfg.model_name],
+ backbone_cfg=backbone_cfg,
+ norm_activation_config=norm_activation_config,
+ )
+ # Build the backbone to get a proper `output_specs`.
+ dummy_inputs = tf_keras.Input(input_specs.shape[1:])
+ _ = maxvit(dummy_inputs, training=False)
+ return maxvit
diff --git a/official/projects/maxvit/modeling/maxvit_test.py b/official/projects/maxvit/modeling/maxvit_test.py
new file mode 100644
index 00000000000..5ae0e4ef22f
--- /dev/null
+++ b/official/projects/maxvit/modeling/maxvit_test.py
@@ -0,0 +1,148 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for MaxViT."""
+import collections
+from typing import Optional, Sequence
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.maxvit.configs import backbones
+from official.projects.maxvit.modeling import maxvit
+from official.vision.configs import common
+
+
+class MaxViTBlockTest(tf.test.TestCase):
+ """Test the layers of MaxViT."""
+
+ def testMaxViTBlockCreation(self) -> None:
+ """Ensures that layers can be constructed and forward-props can run."""
+ inputs_shape = [2, 64, 64, 3]
+ inp = tf.random.uniform(
+ shape=inputs_shape, minval=-1.0, maxval=1.0, dtype=tf.float32
+ )
+
+ model = maxvit.MaxViTBlock(
+ hidden_size=8, head_size=4, window_size=4, grid_size=4
+ )
+ out = model(inp, training=False)
+
+ self.assertAllEqual([2, 64, 64, 8], out.get_shape().as_list())
+ self.assertDTypeEqual(tf.reduce_mean(out).numpy(), np.float32)
+
+
+class MaxViTTest(tf.test.TestCase, parameterized.TestCase):
+ """Test the layers of MaxViT."""
+
+ @parameterized.named_parameters(
+ collections.OrderedDict(
+ testcase_name='MaxViTTest',
+ input_shape=[2, 64, 64, 3],
+ input_dtype=tf.float32,
+ training=False,
+ stem_hsize=[12, 12],
+ num_blocks=[2, 2, 2, 2],
+ window_size=2,
+ grid_size=2,
+ block_type=['maxvit', 'maxvit', 'maxvit'],
+ hidden_size=[16, 32, 64],
+ expected_shape=[2, 4, 4, 64],
+ name='maxvit_test',
+ ),
+ collections.OrderedDict(
+ testcase_name='MaxViTTiny',
+ input_shape=[2, 64, 64, 3],
+ input_dtype=tf.float32,
+ training=False,
+ block_type=['maxvit', 'maxvit', 'maxvit', 'maxvit'],
+ stem_hsize=[64, 64],
+ num_blocks=[2, 3, 5, 2],
+ window_size=2,
+ grid_size=2,
+ hidden_size=[96, 192, 384, 768],
+ expected_shape=[2, 2, 2, 768],
+ name='maxvit_tiny',
+ ),
+ collections.OrderedDict(
+ testcase_name='MaxViTTinyWithPrelogits',
+ input_shape=[2, 64, 64, 3],
+ input_dtype=tf.float32,
+ training=False,
+ representation_size=16,
+ add_gap_layer_norm=True,
+ block_type=['maxvit', 'maxvit', 'maxvit', 'maxvit'],
+ stem_hsize=[64, 64],
+ num_blocks=[2, 3, 5, 2],
+ window_size=2,
+ grid_size=2,
+ hidden_size=[96, 192, 384, 768],
+ expected_shape=[2, 2, 2, 768],
+ name='maxvit_tiny',
+ ),
+ )
+ def testForward(
+ self,
+ input_shape: Sequence[int],
+ input_dtype: Optional[tf.DType] = tf.float32,
+ **kwargs
+ ) -> None:
+ """Ensures that layers can be constructed and forward-props can run."""
+
+ inp = tf.random.uniform(
+ input_shape,
+ minval=-1.0,
+ maxval=1.0,
+ dtype=input_dtype,
+ )
+
+ model = maxvit.MaxViT(**kwargs)
+ out = model(inp, training=kwargs.get('training', None))
+
+ add_gap_layer_norm = kwargs.get('add_gap_layer_norm', False)
+ if add_gap_layer_norm:
+ self.assertAllEqual([input_shape[0], kwargs['representation_size']],
+ out['pre_logits'].get_shape().as_list())
+
+ # Remove `pre_logits` if exists.
+ out.pop('pre_logits', None)
+ out = out[max(out.keys())]
+ self.assertAllEqual(kwargs['expected_shape'], out.get_shape().as_list())
+ self.assertDTypeEqual(tf.reduce_mean(out).numpy(), np.float32)
+
+ def testBuildMaxViTWithConfig(self):
+ backbone_config = backbones.Backbone(
+ type='maxvit',
+ maxvit=backbones.MaxViT(
+ stem_hsize=[32, 32],
+ num_blocks=[2, 3, 5, 2],
+ window_size=2,
+ grid_size=2,
+ hidden_size=[32, 32, 32, 32],
+ ),
+ )
+ backbone = maxvit.build_maxvit(
+ input_specs=tf_keras.layers.InputSpec(shape=[None] + [64, 64, 3]),
+ backbone_config=backbone_config,
+ norm_activation_config=common.NormActivation(),
+ )
+
+ self.assertSetEqual(
+ set(['2', '3', '4', '5']), set(backbone.output_specs.keys())
+ )
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/vision/beta/projects/centernet/common/registry_imports.py b/official/projects/maxvit/registry_imports.py
similarity index 64%
rename from official/vision/beta/projects/centernet/common/registry_imports.py
rename to official/projects/maxvit/registry_imports.py
index 0d3b946fd45..af13e58fd5f 100644
--- a/official/vision/beta/projects/centernet/common/registry_imports.py
+++ b/official/projects/maxvit/registry_imports.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,8 +15,7 @@
"""All necessary imports for registration."""
# pylint: disable=unused-import
+# pylint: disable=g-bad-import-order
from official.vision import registry_imports
-from official.vision.beta.projects.centernet.configs import centernet
-from official.vision.beta.projects.centernet.modeling import centernet_model
-from official.vision.beta.projects.centernet.modeling.backbones import hourglass
-from official.vision.beta.projects.centernet.tasks import centernet as centernet_task
+from official.projects.maxvit import configs # pylint: disable=unused-import
+from official.projects.maxvit.modeling import maxvit # pylint: disable=unused-import
diff --git a/official/projects/vit/train.py b/official/projects/maxvit/train.py
similarity index 71%
rename from official/projects/vit/train.py
rename to official/projects/maxvit/train.py
index 2ef9ebdfff1..3d9c82667bd 100644
--- a/official/projects/vit/train.py
+++ b/official/projects/maxvit/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -12,13 +12,11 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-"""TensorFlow Model Garden Vision training driver, including ViT configs.."""
-
+"""TensorFlow Model Garden Vision training driver, including MaxViT configs.."""
from absl import app
from official.common import flags as tfm_flags
-from official.projects.vit import configs # pylint: disable=unused-import
-from official.projects.vit.modeling import vit # pylint: disable=unused-import
+from official.projects.maxvit import registry_imports # pylint: disable=unused-import
from official.vision import train
diff --git a/official/projects/maxvit/train_test.py b/official/projects/maxvit/train_test.py
new file mode 100644
index 00000000000..bc73a3a62e8
--- /dev/null
+++ b/official/projects/maxvit/train_test.py
@@ -0,0 +1,103 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import json
+import os
+
+from absl import flags
+from absl.testing import flagsaver
+import gin
+import tensorflow as tf, tf_keras
+
+from official.projects.maxvit import train as train_lib
+from official.vision.dataloaders import tfexample_utils
+
+
+FLAGS = flags.FLAGS
+
+
+class TrainTest(tf.test.TestCase):
+
+ def setUp(self):
+ super().setUp()
+
+ self._model_dir = os.path.join(self.get_temp_dir(), 'model_dir')
+ tf.io.gfile.makedirs(self._model_dir)
+ self._test_tfrecord_file = os.path.join(
+ self.get_temp_dir(), 'test.tfrecord'
+ )
+ num_samples = 3
+ example = tf.train.Example.FromString(
+ tfexample_utils.create_classification_example(
+ image_height=224, image_width=224
+ )
+ )
+ examples = [example] * num_samples
+ tfexample_utils.dump_to_tfrecord(
+ record_file=self._test_tfrecord_file, tf_examples=examples
+ )
+
+ def test_run(self):
+ saved_flag_values = flagsaver.save_flag_values()
+ train_lib.tfm_flags.define_flags()
+ FLAGS.mode = 'train'
+ FLAGS.model_dir = self._model_dir
+ FLAGS.experiment = 'maxvit_imagenet'
+
+ params_override = json.dumps({
+ 'runtime': {
+ 'mixed_precision_dtype': 'float32',
+ },
+ 'trainer': {
+ 'train_steps': 1,
+ 'validation_steps': 1,
+ 'optimizer_config': {
+ 'ema': None,
+ },
+ },
+ 'task': {
+ 'init_checkpoint': '',
+ 'model': {
+ 'backbone': {
+ 'maxvit': {
+ 'model_name': 'maxvit-tiny-for-test',
+ 'representation_size': 64,
+ 'add_gap_layer_norm': True,
+ }
+ },
+ 'input_size': [224, 224, 3],
+ 'num_classes': 3,
+ },
+ 'train_data': {
+ 'global_batch_size': 2,
+ 'input_path': self._test_tfrecord_file,
+ },
+ 'validation_data': {
+ 'global_batch_size': 2,
+ 'input_path': self._test_tfrecord_file,
+ },
+ },
+ })
+ FLAGS.params_override = params_override
+ train_lib.train.main('unused_args')
+
+ FLAGS.mode = 'eval'
+ with gin.unlock_config():
+ train_lib.train.main('unused_args')
+
+ flagsaver.restore_flag_values(saved_flag_values)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mobilebert/__init__.py b/official/projects/mobilebert/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/mobilebert/__init__.py
+++ b/official/projects/mobilebert/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/mobilebert/distillation.py b/official/projects/mobilebert/distillation.py
index 87a621b0ec8..1f89e6965f5 100644
--- a/official/projects/mobilebert/distillation.py
+++ b/official/projects/mobilebert/distillation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl import logging
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
from official.modeling import optimization
@@ -73,22 +73,34 @@ class PretrainDistillConfig(base_config.Config):
class BertDistillationProgressiveConfig(policies.ProgressiveConfig):
"""Defines the specific distillation behavior."""
if_copy_embeddings: bool = True
- layer_wise_distill_config: LayerWiseDistillConfig = LayerWiseDistillConfig()
- pretrain_distill_config: PretrainDistillConfig = PretrainDistillConfig()
+ layer_wise_distill_config: LayerWiseDistillConfig = dataclasses.field(
+ default_factory=LayerWiseDistillConfig
+ )
+ pretrain_distill_config: PretrainDistillConfig = dataclasses.field(
+ default_factory=PretrainDistillConfig
+ )
@dataclasses.dataclass
class BertDistillationTaskConfig(cfg.TaskConfig):
"""Defines the teacher/student model architecture and training data."""
- teacher_model: bert.PretrainerConfig = bert.PretrainerConfig(
- encoder=encoders.EncoderConfig(type='mobilebert'))
-
- student_model: bert.PretrainerConfig = bert.PretrainerConfig(
- encoder=encoders.EncoderConfig(type='mobilebert'))
+ teacher_model: bert.PretrainerConfig = dataclasses.field(
+ default_factory=lambda: bert.PretrainerConfig( # pylint: disable=g-long-lambda
+ encoder=encoders.EncoderConfig(type='mobilebert')
+ )
+ )
+
+ student_model: bert.PretrainerConfig = dataclasses.field(
+ default_factory=lambda: bert.PretrainerConfig( # pylint: disable=g-long-lambda
+ encoder=encoders.EncoderConfig(type='mobilebert')
+ )
+ )
# The path to the teacher model checkpoint or its directory.
teacher_model_init_checkpoint: str = ''
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
def build_sub_encoder(encoder, target_layer_id):
@@ -106,7 +118,7 @@ def build_sub_encoder(encoder, target_layer_id):
layer_output, attention_score = encoder.transformer_layers[layer_idx](
layer_output, attention_mask, return_attention_scores=True)
- return tf.keras.Model(
+ return tf_keras.Model(
inputs=[input_ids, input_mask, type_ids],
outputs=[layer_output, attention_score])
@@ -152,7 +164,7 @@ def __init__(self,
raise ValueError('distill_ground_truth_ratio has to be within [0, 1].')
# A non-trainable layer for feature normalization for transfer loss
- self._layer_norm = tf.keras.layers.LayerNormalization(
+ self._layer_norm = tf_keras.layers.LayerNormalization(
axis=-1,
beta_initializer='zeros',
gamma_initializer='ones',
@@ -182,7 +194,7 @@ def _build_pretrainer(self, pretrainer_cfg: bert.PretrainerConfig, name: str):
masked_lm = layers.MobileBertMaskedLM(
embedding_table=encoder.get_embedding_table(),
activation=tf_utils.get_activation(pretrainer_cfg.mlm_activation),
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=pretrainer_cfg.mlm_initializer_range),
name='cls/predictions')
@@ -207,7 +219,7 @@ def num_steps(self, stage_id) -> int:
return self._progressive_config.pretrain_distill_config.num_steps
# override policies.ProgressivePolicy
- def get_model(self, stage_id, old_model=None) -> tf.keras.Model:
+ def get_model(self, stage_id, old_model=None) -> tf_keras.Model:
del old_model
return self.build_model(stage_id)
@@ -238,6 +250,9 @@ def get_optimizer(self, stage_id):
})
opt_factory = optimization.OptimizerFactory(params)
optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
+ if isinstance(optimizer, tf_keras.optimizers.experimental.Optimizer):
+ optimizer = tf_keras.__internal__.optimizers.convert_to_legacy_optimizer(
+ optimizer)
return optimizer
@@ -259,7 +274,7 @@ def get_eval_dataset(self, stage_id):
return self._the_only_eval_dataset
# override base_task.task
- def build_model(self, stage_id) -> tf.keras.Model:
+ def build_model(self, stage_id) -> tf_keras.Model:
"""Build teacher/student keras models with outputs for current stage."""
# Freeze the teacher model.
self._teacher_pretrainer.trainable = False
@@ -292,7 +307,7 @@ def build_model(self, stage_id) -> tf.keras.Model:
for i in range(stage_id):
student_encoder.transformer_layers[i].trainable = False
- return tf.keras.Model(
+ return tf_keras.Model(
inputs=inputs,
outputs=dict(
student_output_feature=student_output_feature,
@@ -311,7 +326,7 @@ def build_model(self, stage_id) -> tf.keras.Model:
for layer in student_encoder.transformer_layers:
layer.trainable = True
- model = tf.keras.Model(
+ model = tf_keras.Model(
inputs=inputs,
outputs=dict(
student_pretrainer_output=student_pretrainer_output,
@@ -363,9 +378,9 @@ def _get_distribution_losses(self, teacher, student):
def _get_attention_loss(self, teacher_score, student_score):
# Note that the definition of KLDivergence here is a little different from
- # the original one (tf.keras.losses.KLDivergence). We adopt this approach
+ # the original one (tf_keras.losses.KLDivergence). We adopt this approach
# to stay consistent with the TF1 implementation.
- teacher_weight = tf.keras.activations.softmax(teacher_score, axis=-1)
+ teacher_weight = tf_keras.activations.softmax(teacher_score, axis=-1)
student_log_weight = tf.nn.log_softmax(student_score, axis=-1)
kl_divergence = -(teacher_weight * student_log_weight)
kl_divergence = tf.math.reduce_sum(kl_divergence, axis=-1, keepdims=True)
@@ -383,7 +398,7 @@ def build_losses(self, labels, outputs, metrics) -> tf.Tensor:
teacher_feature = outputs['teacher_output_feature']
student_feature = outputs['student_output_feature']
- feature_transfer_loss = tf.keras.losses.mean_squared_error(
+ feature_transfer_loss = tf_keras.losses.mean_squared_error(
self._layer_norm(teacher_feature), self._layer_norm(student_feature))
feature_transfer_loss *= distill_config.hidden_distill_factor
beta_loss, gamma_loss = self._get_distribution_losses(teacher_feature,
@@ -438,7 +453,7 @@ def build_losses(self, labels, outputs, metrics) -> tf.Tensor:
sentence_outputs = tf.cast(
student_pretrainer_output['next_sentence'], dtype=tf.float32)
sentence_loss = tf.reduce_mean(
- tf.keras.losses.sparse_categorical_crossentropy(
+ tf_keras.losses.sparse_categorical_crossentropy(
sentence_labels, sentence_outputs, from_logits=True))
total_loss += sentence_loss
@@ -446,16 +461,16 @@ def build_losses(self, labels, outputs, metrics) -> tf.Tensor:
metrics = dict([(metric.name, metric) for metric in metrics])
if not last_stage:
- metrics['feature_transfer_mse'].update_state(feature_transfer_loss)
- metrics['beta_transfer_loss'].update_state(beta_loss)
- metrics['gamma_transfer_loss'].update_state(gamma_loss)
+ metrics['feature_transfer_mse'].update_state(feature_transfer_loss) # pyrefly: ignore[unbound-name]
+ metrics['beta_transfer_loss'].update_state(beta_loss) # pyrefly: ignore[unbound-name]
+ metrics['gamma_transfer_loss'].update_state(gamma_loss) # pyrefly: ignore[unbound-name]
layer_wise_config = self._progressive_config.layer_wise_distill_config
if layer_wise_config.if_transfer_attention:
- metrics['attention_transfer_loss'].update_state(attention_loss)
+ metrics['attention_transfer_loss'].update_state(attention_loss) # pyrefly: ignore[unbound-name]
else:
- metrics['lm_example_loss'].update_state(mlm_loss)
+ metrics['lm_example_loss'].update_state(mlm_loss) # pyrefly: ignore[unbound-name]
if 'next_sentence_labels' in labels:
- metrics['next_sentence_loss'].update_state(sentence_loss)
+ metrics['next_sentence_loss'].update_state(sentence_loss) # pyrefly: ignore[unbound-name]
metrics['total_loss'].update_state(total_loss)
return total_loss
@@ -464,18 +479,18 @@ def build_losses(self, labels, outputs, metrics) -> tf.Tensor:
def build_metrics(self, training=None):
del training
metrics = [
- tf.keras.metrics.Mean(name='feature_transfer_mse'),
- tf.keras.metrics.Mean(name='beta_transfer_loss'),
- tf.keras.metrics.Mean(name='gamma_transfer_loss'),
- tf.keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
- tf.keras.metrics.Mean(name='lm_example_loss'),
- tf.keras.metrics.Mean(name='total_loss')]
+ tf_keras.metrics.Mean(name='feature_transfer_mse'),
+ tf_keras.metrics.Mean(name='beta_transfer_loss'),
+ tf_keras.metrics.Mean(name='gamma_transfer_loss'),
+ tf_keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
+ tf_keras.metrics.Mean(name='lm_example_loss'),
+ tf_keras.metrics.Mean(name='total_loss')]
if self._progressive_config.layer_wise_distill_config.if_transfer_attention:
- metrics.append(tf.keras.metrics.Mean(name='attention_transfer_loss'))
+ metrics.append(tf_keras.metrics.Mean(name='attention_transfer_loss'))
if self._task_config.train_data.use_next_sentence_label:
- metrics.append(tf.keras.metrics.SparseCategoricalAccuracy(
+ metrics.append(tf_keras.metrics.SparseCategoricalAccuracy(
name='next_sentence_accuracy'))
- metrics.append(tf.keras.metrics.Mean(name='next_sentence_loss'))
+ metrics.append(tf_keras.metrics.Mean(name='next_sentence_loss'))
return metrics
@@ -495,8 +510,8 @@ def process_metrics(self, metrics, labels, student_pretrainer_output):
student_pretrainer_output['next_sentence'])
# overrides base_task.Task
- def train_step(self, inputs, model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer, metrics):
+ def train_step(self, inputs, model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer, metrics):
"""Does forward and backward.
Args:
@@ -533,7 +548,7 @@ def train_step(self, inputs, model: tf.keras.Model,
return {self.loss: loss}
# overrides base_task.Task
- def validation_step(self, inputs, model: tf.keras.Model, metrics):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics):
"""Validatation step.
Args:
diff --git a/official/projects/mobilebert/distillation_test.py b/official/projects/mobilebert/distillation_test.py
index 68ce71abb16..43fe846dd7a 100644
--- a/official/projects/mobilebert/distillation_test.py
+++ b/official/projects/mobilebert/distillation_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from absl import logging
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.modeling import optimization
@@ -119,7 +119,7 @@ def prepare_config(self, teacher_block_num, student_block_num,
masked_lm = layers.MobileBertMaskedLM(
embedding_table=teacher_encoder.get_embedding_table(),
activation=tf_utils.get_activation(pretrainer_config.mlm_activation),
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=pretrainer_config.mlm_initializer_range),
name='cls/predictions')
teacher_pretrainer = models.BertPretrainerV2(
@@ -153,7 +153,7 @@ def test_task(self, teacher_block_num, student_block_num,
eval_dataset = bert_distillation_task.get_eval_dataset(stage_id=0)
eval_iterator = iter(eval_dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.legacy.SGD(learning_rate=0.1)
# test train/val step for all stages, including the last pretraining stage
for stage in range(student_block_num + 1):
diff --git a/official/projects/mobilebert/export_tfhub.py b/official/projects/mobilebert/export_tfhub.py
index 184de577b57..78355902ff7 100644
--- a/official/projects/mobilebert/export_tfhub.py
+++ b/official/projects/mobilebert/export_tfhub.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.mobilebert import model_utils
@@ -43,12 +43,12 @@ def create_mobilebert_model(bert_config):
# For interchangeability with other text representations,
# add "default" as an alias for MobileBERT's whole-input reptesentations.
encoder_output_dict["default"] = encoder_output_dict["pooled_output"]
- core_model = tf.keras.Model(
+ core_model = tf_keras.Model(
inputs=encoder_inputs_dict, outputs=encoder_output_dict)
pretrainer_inputs_dict = {x.name: x for x in pretrainer.inputs}
pretrainer_output_dict = pretrainer(pretrainer_inputs_dict)
- mlm_model = tf.keras.Model(
+ mlm_model = tf_keras.Model(
inputs=pretrainer_inputs_dict, outputs=pretrainer_output_dict)
# Set `_auto_track_sub_layers` to False, so that the additional weights
# from `mlm` sub-object will not be included in the core model.
@@ -60,7 +60,7 @@ def create_mobilebert_model(bert_config):
def export_bert_tfhub(bert_config, model_checkpoint_path, hub_destination,
vocab_file, do_lower_case):
- """Restores a tf.keras.Model and saves for TF-Hub."""
+ """Restores a tf_keras.Model and saves for TF-Hub."""
core_model, pretrainer = create_mobilebert_model(bert_config)
checkpoint = tf.train.Checkpoint(**pretrainer.checkpoint_items)
diff --git a/official/projects/mobilebert/model_utils.py b/official/projects/mobilebert/model_utils.py
index 70be52a4f86..0ef6069ced1 100644
--- a/official/projects/mobilebert/model_utils.py
+++ b/official/projects/mobilebert/model_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -159,7 +159,7 @@ def create_mobilebert_pretrainer(bert_config):
masked_lm = layers.MobileBertMaskedLM(
embedding_table=mobilebert_encoder.get_embedding_table(),
activation=tf_utils.get_activation(bert_config.hidden_act),
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
name="cls/predictions")
diff --git a/official/projects/mobilebert/run_distillation.py b/official/projects/mobilebert/run_distillation.py
index 30aefcf8d7b..aec39cc0296 100644
--- a/official/projects/mobilebert/run_distillation.py
+++ b/official/projects/mobilebert/run_distillation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/mobilebert/tf2_model_checkpoint_converter.py b/official/projects/mobilebert/tf2_model_checkpoint_converter.py
index eae8e8b2d2a..7fd11ed5125 100644
--- a/official/projects/mobilebert/tf2_model_checkpoint_converter.py
+++ b/official/projects/mobilebert/tf2_model_checkpoint_converter.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/mobilebert/utils.py b/official/projects/mobilebert/utils.py
index f9e2a924c6f..34666ab3ca2 100644
--- a/official/projects/mobilebert/utils.py
+++ b/official/projects/mobilebert/utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/mosaic/README.md b/official/projects/mosaic/README.md
new file mode 100644
index 00000000000..daf4859d790
--- /dev/null
+++ b/official/projects/mosaic/README.md
@@ -0,0 +1,178 @@
+# MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded Context
+
+[](https://arxiv.org/abs/2112.11623)
+
+This repository is the official implementation of the following
+paper.
+
+* [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded Context](https://arxiv.org/abs/2112.11623)
+
+## Description
+
+MOSAIC is a neural network architecture for efficient and accurate semantic
+image segmentation on mobile devices. MOSAIC is designed using commonly
+supported neural operations by diverse mobile hardware platforms for flexible
+deployment across various mobile platforms. With a simple asymmetric
+encoder-decoder structure which consists of an efficient multi-scale context
+encoder and a light-weight hybrid decoder to recover spatial details from
+aggregated information, MOSAIC achieves better balanced performance while
+considering accuracy and computational cost. Deployed on top of a tailored
+feature extraction backbone based on a searched classification network, MOSAIC
+achieves a 5% absolute accuracy gain on ADE20K with similar or lower latency
+compared to the current industry standard MLPerf mobile v1.0 models and
+state-of-the-art architectures.
+
+[MLPerf Mobile v2.0]((https://mlcommons.org/en/inference-mobile-20/)) included
+MOSAIC as a new industry standard benchmark model for image segmentation.
+Please see details [here](https://mlcommons.org/en/news/mlperf-inference-1q2022/).
+
+You can also refer to the [MLCommons GitHub repository](https://github.com/mlcommons/mobile_open/tree/main/vision/mosaic).
+
+## History
+
+### Oct 13, 2022
+
+* First release of MOSAIC in TensorFlow 2 including checkpoints that have been
+ pretrained on Cityscapes.
+
+## Maintainers
+
+* Weijun Wang ([weijunw-g](https://github.com/weijunw-g))
+* Fang Yang ([fyangf](https://github.com/fyangf))
+* Shixin Luo ([luotigerlsx](https://github.com/luotigerlsx))
+
+## Requirements
+
+[](https://badge.fury.io/py/tensorflow)
+[](https://badge.fury.io/py/tf-models-official)
+
+## Results
+
+The following table shows the mIoU measured on the `cityscapes` dataset.
+
+| Config | Backbone | Resolution | branch_filter_depths | pyramid_pool_bin_nums | mIoU | Download |
+|-------------------------|:--------------------:|:----------:|:--------------------:|:---------------------:|:-----:|:--------:|
+| Paper reference config | MobileNetMultiAVGSeg | 1024x2048 | [32, 32] | [4, 8, 16] | 75.98 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/mosaic/MobileNetMultiAVGSeg-r1024-ebf32-nogp.tar.gz) [tensorboard](https://tensorboard.dev/experiment/okEog90bSwupajFgJwGEIw/#scalars) |
+| Current best config | MobileNetMultiAVGSeg | 1024x2048 | [64, 64] | [1, 4, 8, 16] | 77.24 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/mosaic/MobileNetMultiAVGSeg-r1024-ebf64-gp.tar.gz) [tensorboard](https://tensorboard.dev/experiment/l5hkV7JaQM23EXeOBT6oJg/#scalars) |
+
+* `branch_filter_depths`: the number of convolution channels in each branch at
+ a pyramid level after `Spatial Pyramid Pooling`
+* `pyramid_pool_bin_nums`: the number of bins at each level of the `Spatial
+ Pyramid Pooling`
+
+## Training
+
+It can run on Google Cloud Platform using Cloud TPU.
+[Here](https://cloud.google.com/tpu/docs/how-to) is the instruction of using
+Cloud TPU. Following the instructions to set up Cloud TPU and
+launch training by:
+
+```shell
+EXP_TYPE=mosaic_mnv35_cityscapes
+EXP_NAME="" # You can give any name to the experiment.
+TPU_NAME="" # The name assigned while creating a Cloud TPU
+MODEL_DIR="gs://"
+# Now launch the experiment.
+python3 -m official.projects.mosaic.train \
+ --experiment=$EXP_TYPE \
+ --mode=train \
+ --tpu=$TPU_NAME \
+ --model_dir=$MODEL_DIR \
+ --config_file=official/projects/mosaic/configs/experiments/mosaic_mnv35_cityscapes_tdfs_tpu.yaml
+```
+
+## Evaluation
+
+Please run this command line for evaluation.
+
+```shell
+EXP_TYPE=mosaic_mnv35_cityscapes
+EXP_NAME="" # You can give any name to the experiment.
+TPU_NAME="" # The name assigned while creating a Cloud TPU
+MODEL_DIR="gs://"
+# Now launch the experiment.
+python3 -m official.projects.mosaic.train \
+ --experiment=$EXP_TYPE \
+ --mode=eval \
+ --tpu=$TPU_NAME \
+ --model_dir=$MODEL_DIR \
+ --config_file=official/projects/mosaic/configs/experiments/mosaic_mnv35_cityscapes_tdfs_tpu.yaml
+```
+
+## Quantization Aware Training (QAT)
+
+We support quantization aware training (QAT) and convert trained model to a
+TFLite model for on-device inference.
+
+### QAT Training
+```shell
+EXP_TYPE=mosaic_mnv35_cityscapes_qat
+EXP_NAME="" # You can give any name to the experiment.
+TPU_NAME="" # The name assigned while creating a Cloud TPU
+MODEL_DIR="gs://"
+NON_QAT_CHECKPOINT="gs://" # The checkpoint of non-qat training
+python3 -m official.projects.mosaic.train \
+ --experiment=$EXP_TYPE \
+ --mode=eval \
+ --tpu=$TPU_NAME \
+ --model_dir=$MODEL_DIR \
+ --config_file=official/projects/mosaic/qat/configs/experiments/semantic_segmentation/mosaic_mnv35_cityscapes_tfds_qat_tpu.yaml \
+ --params_override="task.quantization.pretrained_original_checkpoint=${NON_QAT_CHECKPOINT}"
+```
+
+### Export TFLite
+```shell
+EXP_TYPE=mosaic_mnv35_cityscapes_qat
+QAT_CKPT_PATH="gs://" # The checkpoint of qat training
+EXPORT_PATH="" # The path of SavedModel to be exported
+INPUT_SIZE="" # The image size that the model is trained on e.g. 1024,2048 for Cityscapes
+python3 -m official.projects.mosaic.qat.serving.export_saved_model \
+--checkpoint_path=${QAT_CKPT_PATH} \
+--config_file=${QAT_CKPT_PATH}/params.yaml \
+--export_dir=${EXPORT_PATH} \
+--experiment=${EXP_TYPE} \
+--input_type=tflite \
+--input_image_size=${INPUT_SIZE} \
+--alsologtostderr
+```
+
+```shell
+EXP_TYPE=mosaic_mnv35_cityscapes_qat
+SAVED_MOODEL_DIR="" # The path to the SavedModel exported in the previous step
+TFLITE_PATH="" # The path of TFLite file to be exported
+python3 -m official.projects.mosaic.qat.serving.export_tflite \
+--experiment=${EXP_TYPE} \
+--saved_model_dir=${SAVED_MOODEL_DIR} \
+--tflite_path=${TFLITE_PATH} \
+--quant_type=qat \
+--alsologtostderr
+```
+
+### Results
+The benchmark results on Cityscapes are reported below:
+
+model | resolution | mIoU | mIoU (QAT INT8) | Latency (QAT INT8, ms per img on Pixel6) | download (ckpt) | download (tflite)
+:------------------------------ | :--------: | ----------: | --------------: | ----------------------------------------------------------------------: | --------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------: | ----------------:
+MobileNet Multi-HW AVG + MOSAIC | 1024x2048 | 77.24 | 77.13 | 524ms, 327ms (w/o, w/ XNNPACK) | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/mosaic/MobileNetMultiAVGSeg-r1024-ebf64-gp-qat.tar.gz) \| [tensorboard](https://tensorboard.dev/experiment/g0ZzmRDdRdGn5Xn07xXvwg/#scalars) | [QAT INT8](https://storage.googleapis.com/tf_model_garden/vision/mosaic/mobilenet_multiavgseg_r1024_ebf64_gp_qat/tflite_model_depthwise_qat_int8.tflite)
+
+
+## License
+
+[](https://opensource.org/licenses/Apache-2.0)
+
+This project is licensed under the terms of the **Apache License 2.0**.
+
+## Citation
+
+If you want to cite this repository in your work, please consider citing the
+paper.
+
+```
+@inproceedings{weijun2021mosaic,
+ title={MOSAIC: Mobile Segmentation via decoding Aggregated Information and
+ encoded Context},
+ author={Weijun Wang, Andrew Howard},
+ journal={arXiv preprint arXiv:2112.11623},
+ year={2021},
+}
+```
diff --git a/official/projects/mosaic/configs/experiments/mosaic_mnv35_cityscapes_tfds_tpu.yaml b/official/projects/mosaic/configs/experiments/mosaic_mnv35_cityscapes_tfds_tpu.yaml
new file mode 100644
index 00000000000..cbea6c2d1e4
--- /dev/null
+++ b/official/projects/mosaic/configs/experiments/mosaic_mnv35_cityscapes_tfds_tpu.yaml
@@ -0,0 +1,86 @@
+# Using Tensorflow datasets: 'cityscapes/semantic_segmentation'
+# Some expected flags to use:
+# --experiment=mosaic_mnv35_cityscapes
+# mIoU: 77.24%
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ num_classes: 19
+ input_size: [null, null, 3]
+ backbone:
+ type: 'mobilenet'
+ mobilenet:
+ model_id: 'MobileNetMultiAVGSeg'
+ output_intermediate_endpoints: true
+ output_stride: 16
+ neck:
+ branch_filter_depths: [64, 64]
+ conv_kernel_sizes: [3, 5]
+ pyramid_pool_bin_nums: [1, 4, 8, 16]
+ dropout_rate: 0.0
+ head:
+ num_classes: 19
+ decoder_input_levels: ['3/depthwise', '2/depthwise']
+ decoder_stage_merge_styles: ['concat_merge', 'sum_merge']
+ decoder_filters: [64, 64]
+ decoder_projected_filters: [19, 19]
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ init_checkpoint: 'gs://tf_model_garden/vision/mobilenet/v3.5multiavg_seg_float/'
+ init_checkpoint_modules: 'backbone'
+ losses:
+ l2_weight_decay: 1.0e-04
+ train_data:
+ output_size: [1024, 2048]
+ crop_size: [1024, 2048]
+ input_path: ''
+ tfds_name: 'cityscapes/semantic_segmentation'
+ tfds_split: 'train'
+ is_training: true
+ global_batch_size: 32
+ dtype: 'float32'
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ output_size: [1024, 2048]
+ input_path: ''
+ tfds_name: 'cityscapes/semantic_segmentation'
+ tfds_split: 'validation'
+ is_training: false
+ global_batch_size: 32
+ dtype: 'float32'
+ drop_remainder: false
+ resize_eval_groundtruth: true
+trainer:
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 100000
+ initial_learning_rate: 0.1
+ power: 0.9
+ type: polynomial
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 925
+ type: linear
+ steps_per_loop: 92 # 2975 / 32 = 92
+ summary_interval: 92
+ train_steps: 100000
+ validation_interval: 92
+ validation_steps: 16 # 500 / 32 = 16
+ checkpoint_interval: 92
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/mosaic/configs/mosaic_config.py b/official/projects/mosaic/configs/mosaic_config.py
new file mode 100644
index 00000000000..142bf061b72
--- /dev/null
+++ b/official/projects/mosaic/configs/mosaic_config.py
@@ -0,0 +1,367 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Configuration definition for Semantic Segmentation with MOSAIC."""
+import dataclasses
+import math
+import os
+from typing import List, Optional, Union
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.modeling import optimization
+from official.vision.configs import backbones
+from official.vision.configs import common
+from official.vision.configs import semantic_segmentation as seg_cfg
+
+
+@dataclasses.dataclass
+class MosaicDecoderHead(hyperparams.Config):
+ """MOSAIC decoder head config for Segmentation."""
+ num_classes: int = 19
+ decoder_input_levels: List[str] = dataclasses.field(default_factory=list)
+ decoder_stage_merge_styles: List[str] = dataclasses.field(
+ default_factory=list)
+ decoder_filters: List[int] = dataclasses.field(default_factory=list)
+ decoder_projected_filters: List[int] = dataclasses.field(default_factory=list)
+ encoder_end_level: int = 4
+ use_additional_classifier_layer: bool = False
+ classifier_kernel_size: int = 1
+ activation: str = 'relu'
+ kernel_initializer: str = 'glorot_uniform'
+ interpolation: str = 'bilinear'
+
+
+@dataclasses.dataclass
+class MosaicEncoderNeck(hyperparams.Config):
+ """MOSAIC encoder neck config for segmentation."""
+ encoder_input_level: Union[str, int] = '4'
+ branch_filter_depths: List[int] = dataclasses.field(default_factory=list)
+ conv_kernel_sizes: List[int] = dataclasses.field(default_factory=list)
+ pyramid_pool_bin_nums: List[int] = dataclasses.field(default_factory=list)
+ activation: str = 'relu'
+ dropout_rate: float = 0.1
+ kernel_initializer: str = 'glorot_uniform'
+ interpolation: str = 'bilinear'
+ use_depthwise_convolution: bool = True
+
+
+@dataclasses.dataclass
+class MosaicSemanticSegmentationModel(hyperparams.Config):
+ """MOSAIC semantic segmentation model config."""
+ num_classes: int = 19
+ input_size: List[int] = dataclasses.field(default_factory=list)
+ head: MosaicDecoderHead = dataclasses.field(default_factory=MosaicDecoderHead)
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='mobilenet', mobilenet=backbones.MobileNet()
+ )
+ )
+ neck: MosaicEncoderNeck = dataclasses.field(default_factory=MosaicEncoderNeck)
+ mask_scoring_head: Optional[seg_cfg.MaskScoringHead] = None
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ use_sync_bn=True, norm_momentum=0.99, norm_epsilon=0.001
+ )
+ )
+
+
+@dataclasses.dataclass
+class MosaicSemanticSegmentationTask(seg_cfg.SemanticSegmentationTask):
+ """The config for MOSAIC segmentation task."""
+ model: MosaicSemanticSegmentationModel = dataclasses.field(
+ default_factory=MosaicSemanticSegmentationModel
+ )
+ train_data: seg_cfg.DataConfig = dataclasses.field(
+ default_factory=lambda: seg_cfg.DataConfig(is_training=True)
+ )
+ validation_data: seg_cfg.DataConfig = dataclasses.field(
+ default_factory=lambda: seg_cfg.DataConfig(is_training=False)
+ )
+ losses: seg_cfg.Losses = dataclasses.field(default_factory=seg_cfg.Losses)
+ evaluation: seg_cfg.Evaluation = dataclasses.field(
+ default_factory=seg_cfg.Evaluation
+ )
+ train_input_partition_dims: List[int] = dataclasses.field(
+ default_factory=list)
+ eval_input_partition_dims: List[int] = dataclasses.field(
+ default_factory=list)
+ init_checkpoint: Optional[str] = None
+ init_checkpoint_modules: Union[
+ str, List[str]] = 'all' # all, backbone, and/or neck.
+ export_config: seg_cfg.ExportConfig = dataclasses.field(
+ default_factory=seg_cfg.ExportConfig
+ )
+
+
+# Cityscapes Dataset (Download and process the dataset yourself)
+CITYSCAPES_TRAIN_EXAMPLES = 2975
+CITYSCAPES_VAL_EXAMPLES = 500
+CITYSCAPES_INPUT_PATH_BASE = 'cityscapes/tfrecord'
+
+
+@exp_factory.register_config_factory('mosaic_mnv35_cityscapes')
+def mosaic_mnv35_cityscapes() -> cfg.ExperimentConfig:
+ """Instantiates an experiment configuration of image segmentation task.
+
+ This image segmentation experiment is conducted on Cityscapes dataset. The
+ model architecture is a MOSAIC encoder-decoer. The default backbone network is
+ a mobilenet variant called Mobilenet_v3.5-MultiAvg on top of which the MOSAIC
+ encoder-decoder can be deployed. All detailed configurations can be overridden
+ by a .yaml file provided by the user to launch the experiments. Please refer
+ to .yaml examples in the path of ../configs/experiments/.
+
+ Returns:
+ A particular instance of cfg.ExperimentConfig for MOSAIC model based
+ image semantic segmentation task.
+ """
+ train_batch_size = 16
+ eval_batch_size = 16
+ steps_per_epoch = CITYSCAPES_TRAIN_EXAMPLES // train_batch_size
+ output_stride = 16
+
+ backbone_output_level = int(math.log2(output_stride))
+ config = cfg.ExperimentConfig(
+ task=MosaicSemanticSegmentationTask(
+ model=MosaicSemanticSegmentationModel(
+ # Cityscapes uses only 19 semantic classes for train/evaluation.
+ # The void (background) class is ignored in train and evaluation.
+ num_classes=19,
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
+ backbone=backbones.Backbone(
+ type='mobilenet',
+ mobilenet=backbones.MobileNet(
+ model_id='MobileNetMultiAVGSeg',
+ output_intermediate_endpoints=True,
+ output_stride=output_stride)),
+ neck=MosaicEncoderNeck(
+ encoder_input_level=backbone_output_level,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=[1, 4, 8, 16], # paper default
+ activation='relu',
+ dropout_rate=0.1,
+ kernel_initializer='glorot_uniform',
+ interpolation='bilinear',
+ use_depthwise_convolution=True),
+ head=MosaicDecoderHead(
+ num_classes=19,
+ decoder_input_levels=['3/depthwise', '2/depthwise'],
+ decoder_stage_merge_styles=['concat_merge', 'sum_merge'],
+ decoder_filters=[64, 64],
+ decoder_projected_filters=[19, 19],
+ encoder_end_level=backbone_output_level,
+ use_additional_classifier_layer=False,
+ classifier_kernel_size=1,
+ activation='relu',
+ kernel_initializer='glorot_uniform',
+ interpolation='bilinear'),
+ norm_activation=common.NormActivation(
+ activation='relu',
+ norm_momentum=0.99,
+ norm_epsilon=1e-3,
+ use_sync_bn=True)),
+ losses=seg_cfg.Losses(l2_weight_decay=4e-5),
+ train_data=seg_cfg.DataConfig(
+ input_path=os.path.join(CITYSCAPES_INPUT_PATH_BASE,
+ 'train_fine**'),
+ crop_size=[1024, 2048],
+ output_size=[1024, 2048],
+ is_training=True,
+ global_batch_size=train_batch_size,
+ aug_scale_min=0.5,
+ aug_scale_max=2.0),
+ validation_data=seg_cfg.DataConfig(
+ input_path=os.path.join(CITYSCAPES_INPUT_PATH_BASE, 'val_fine*'),
+ output_size=[1024, 2048],
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ resize_eval_groundtruth=True,
+ drop_remainder=False),
+ # Imagenet pre-trained Mobilenet_v3.5-MultiAvg checkpoint.
+ init_checkpoint='gs://tf_model_garden/vision/mobilenet/v3.5multiavg_seg_float/',
+ init_checkpoint_modules='backbone'),
+ trainer=cfg.TrainerConfig(
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ train_steps=100000,
+ validation_steps=CITYSCAPES_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ best_checkpoint_eval_metric='mean_iou',
+ best_checkpoint_export_subdir='best_ckpt',
+ best_checkpoint_metric_comp='higher',
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 0.1,
+ 'decay_steps': 100000,
+ 'end_learning_rate': 0.0,
+ 'power': 0.9
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 5 * steps_per_epoch,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+
+ return config
+
+
+@exp_factory.register_config_factory('mosaic_mnv4_cityscapes')
+def mosaic_mnv4_cityscapes() -> cfg.ExperimentConfig:
+ """Instantiates an experiment configuration of image segmentation task.
+
+ This image segmentation experiment is conducted on Cityscapes dataset. The
+ model architecture is a MOSAIC encoder-decoer. The default backbone network is
+ an experimental mobilenet V4 variant on top of which the MOSAIC
+ encoder-decoder can be deployed. All detailed configurations can be overridden
+ by a .yaml file provided by the user to launch the experiments. Please refer
+ to .yaml examples in the path of ../configs/experiments/.
+
+ Returns:
+ A particular instance of cfg.ExperimentConfig for MOSAIC model based
+ image semantic segmentation task.
+ """
+ train_batch_size = 16
+ eval_batch_size = 16
+ steps_per_epoch = CITYSCAPES_TRAIN_EXAMPLES // train_batch_size
+ output_stride = 16
+
+ backbone_output_level = int(math.log2(output_stride))
+ config = cfg.ExperimentConfig(
+ task=MosaicSemanticSegmentationTask(
+ model=MosaicSemanticSegmentationModel(
+ # Cityscapes uses only 19 semantic classes for train/evaluation.
+ # The void (background) class is ignored in train and evaluation.
+ num_classes=19,
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
+ backbone=backbones.Backbone(
+ type='mobilenet',
+ mobilenet=backbones.MobileNet(
+ model_id='MobileNetV4ConvMediumSeg',
+ output_intermediate_endpoints=True,
+ output_stride=output_stride)),
+ neck=MosaicEncoderNeck(
+ encoder_input_level=backbone_output_level,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=[1, 4, 8, 16], # paper default
+ activation='relu',
+ dropout_rate=0.1,
+ kernel_initializer='glorot_uniform',
+ interpolation='bilinear',
+ use_depthwise_convolution=True),
+ head=MosaicDecoderHead(
+ num_classes=19,
+ decoder_input_levels=['3/depthwise', '2/depthwise'],
+ decoder_stage_merge_styles=['concat_merge', 'sum_merge'],
+ decoder_filters=[64, 64],
+ decoder_projected_filters=[19, 19],
+ encoder_end_level=backbone_output_level,
+ use_additional_classifier_layer=False,
+ classifier_kernel_size=1,
+ activation='relu',
+ kernel_initializer='glorot_uniform',
+ interpolation='bilinear',
+ ),
+ norm_activation=common.NormActivation(
+ activation='relu',
+ norm_momentum=0.99,
+ norm_epsilon=1e-3,
+ use_sync_bn=True,
+ ),
+ ),
+ losses=seg_cfg.Losses(l2_weight_decay=4e-5),
+ train_data=seg_cfg.DataConfig(
+ input_path=os.path.join(
+ CITYSCAPES_INPUT_PATH_BASE, 'train_fine**'
+ ),
+ crop_size=[1024, 2048],
+ output_size=[1024, 2048],
+ is_training=True,
+ global_batch_size=train_batch_size,
+ aug_scale_min=0.5,
+ aug_scale_max=2.0,
+ ),
+ validation_data=seg_cfg.DataConfig(
+ input_path=os.path.join(CITYSCAPES_INPUT_PATH_BASE, 'val_fine*'),
+ output_size=[1024, 2048],
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ resize_eval_groundtruth=True,
+ drop_remainder=False,
+ ),
+ # Imagenet pre-trained MobileNetV4ConvMediumSeg checkpoint.
+ init_checkpoint=(
+ 'gs://tf_model_garden/vision/mobilenet/v4_seg_float//'
+ ),
+ init_checkpoint_modules='backbone',
+ ),
+ trainer=cfg.TrainerConfig(
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ train_steps=100000,
+ validation_steps=CITYSCAPES_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ best_checkpoint_eval_metric='mean_iou',
+ best_checkpoint_export_subdir='best_ckpt',
+ best_checkpoint_metric_comp='higher',
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 0.1,
+ 'decay_steps': 100000,
+ 'end_learning_rate': 0.0,
+ 'power': 0.9
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 5 * steps_per_epoch,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+
+ return config
diff --git a/official/projects/mosaic/modeling/mosaic_blocks.py b/official/projects/mosaic/modeling/mosaic_blocks.py
new file mode 100644
index 00000000000..6540c2f4990
--- /dev/null
+++ b/official/projects/mosaic/modeling/mosaic_blocks.py
@@ -0,0 +1,885 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Definitions of building blocks for MOSAIC model.
+
+Reference:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ Context](https://arxiv.org/pdf/2112.11623.pdf)
+"""
+
+from typing import Any, Dict, List, Optional, Tuple, Union
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MultiKernelGroupConvBlock(tf_keras.layers.Layer):
+ """A multi-kernel grouped convolution block.
+
+ This block is used in the segmentation neck introduced in MOSAIC.
+ Reference:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ Context](https://arxiv.org/pdf/2112.11623.pdf)
+ """
+
+ def __init__(
+ self,
+ output_filter_depths: Optional[List[int]] = None,
+ kernel_sizes: Optional[List[int]] = None,
+ use_sync_bn: bool = False,
+ batchnorm_momentum: float = 0.99,
+ batchnorm_epsilon: float = 0.001,
+ activation: str = 'relu',
+ kernel_initializer: str = 'GlorotUniform',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ use_depthwise_convolution: bool = True,
+ **kwargs):
+ """Initializes a Multi-kernel Grouped Convolution Block.
+
+ Args:
+ output_filter_depths: A list of integers representing the numbers of
+ output channels or filter depths of convolution groups.
+ kernel_sizes: A list of integers denoting the convolution kernel sizes in
+ each convolution group.
+ use_sync_bn: A bool, whether or not to use sync batch normalization.
+ batchnorm_momentum: A float for the momentum in BatchNorm. Defaults to
+ 0.99.
+ batchnorm_epsilon: A float for the epsilon value in BatchNorm. Defaults to
+ 0.001.
+ activation: A `str` for the activation fuction type. Defaults to 'relu'.
+ kernel_initializer: Kernel initializer for conv layers. Defaults to
+ `glorot_uniform`.
+ kernel_regularizer: Kernel regularizer for conv layers. Defaults to None.
+ use_depthwise_convolution: Allows spatial pooling to be separable
+ depthwise convolusions.
+ **kwargs: Other keyword arguments for the layer.
+ """
+ super(MultiKernelGroupConvBlock, self).__init__(**kwargs)
+
+ if output_filter_depths is None:
+ output_filter_depths = [64, 64]
+ if kernel_sizes is None:
+ kernel_sizes = [3, 5]
+ if len(output_filter_depths) != len(kernel_sizes):
+ raise ValueError('The number of output groups must match #kernels.')
+ self._output_filter_depths = output_filter_depths
+ self._kernel_sizes = kernel_sizes
+ self._num_groups = len(self._kernel_sizes)
+ self._use_sync_bn = use_sync_bn
+ self._batchnorm_momentum = batchnorm_momentum
+ self._batchnorm_epsilon = batchnorm_epsilon
+ self._activation = activation
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._use_depthwise_convolution = use_depthwise_convolution
+ # To apply BN before activation. Putting BN between conv and activation also
+ # helps quantization where conv+bn+activation are fused into a single op.
+ self._activation_fn = tf_utils.get_activation(activation)
+ if self._use_sync_bn:
+ self._bn_op = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ self._bn_op = tf_keras.layers.BatchNormalization
+
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ self._group_split_axis = -1
+ else:
+ self._bn_axis = 1
+ self._group_split_axis = 1
+
+ def build(self, input_shape: tf.TensorShape) -> None:
+ """Builds the block with the given input shape."""
+ input_channels = input_shape[self._group_split_axis]
+ if input_channels % self._num_groups != 0:
+ raise ValueError('The number of input channels must be divisible by '
+ 'the number of groups for evenly group split.')
+ self._conv_branches = []
+ if self._use_depthwise_convolution:
+ for i, conv_kernel_size in enumerate(self._kernel_sizes):
+ depthwise_conv = tf_keras.layers.DepthwiseConv2D(
+ kernel_size=(conv_kernel_size, conv_kernel_size),
+ depth_multiplier=1,
+ padding='same',
+ depthwise_regularizer=self._kernel_regularizer,
+ depthwise_initializer=self._kernel_initializer,
+ use_bias=False)
+ # Add BN->RELU after depthwise convolution.
+ batchnorm_op_depthwise = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation_depthwise = self._activation_fn
+ feature_conv = tf_keras.layers.Conv2D(
+ filters=self._output_filter_depths[i],
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=None,
+ use_bias=False)
+ batchnorm_op = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ # Use list manually as current QAT API does not support sequential model
+ # within a tf_keras.Sequential block, e.g. conv_branch =
+ # tf_keras.Sequential([depthwise_conv, feature_conv, batchnorm_op,])
+ conv_branch = [
+ depthwise_conv,
+ batchnorm_op_depthwise,
+ activation_depthwise,
+ feature_conv,
+ batchnorm_op,
+ ]
+ self._conv_branches.append(conv_branch)
+ else:
+ for i, conv_kernel_size in enumerate(self._kernel_sizes):
+ norm_conv = tf_keras.layers.Conv2D(
+ filters=self._output_filter_depths[i],
+ kernel_size=(conv_kernel_size, conv_kernel_size),
+ padding='same',
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ activation=None,
+ use_bias=False)
+ batchnorm_op = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ conv_branch = [norm_conv, batchnorm_op]
+ self._conv_branches.append(conv_branch)
+ self._concat_groups = tf_keras.layers.Concatenate(
+ axis=self._group_split_axis)
+
+ def call(self,
+ inputs: tf.Tensor,
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Calls this group convolution block with the given inputs."""
+ inputs_splits = tf.split(inputs,
+ num_or_size_splits=self._num_groups,
+ axis=self._group_split_axis)
+ output_branches = []
+ for i, x in enumerate(inputs_splits):
+ conv_branch = self._conv_branches[i]
+ # Apply layers sequentially and manually.
+ for layer in conv_branch:
+ if isinstance(layer, tf_keras.layers.Layer):
+ x = layer(x, training=training)
+ else:
+ x = layer(x)
+ # Apply activation function after BN, which also helps quantization
+ # where conv+bn+activation are fused into a single op.
+ x = self._activation_fn(x)
+ output_branches.append(x)
+ x = self._concat_groups(output_branches)
+ return x
+
+ def get_config(self) -> Dict[str, Any]:
+ """Returns a config dictionary for initialization from serialization."""
+ config = {
+ 'output_filter_depths': self._output_filter_depths,
+ 'kernel_sizes': self._kernel_sizes,
+ 'num_groups': self._num_groups,
+ 'use_sync_bn': self._use_sync_bn,
+ 'batchnorm_momentum': self._batchnorm_momentum,
+ 'batchnorm_epsilon': self._batchnorm_epsilon,
+ 'activation': self._activation,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'use_depthwise_convolution': self._use_depthwise_convolution,
+ }
+ base_config = super(MultiKernelGroupConvBlock, self).get_config()
+ base_config.update(config)
+ return base_config
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MosaicEncoderBlock(tf_keras.layers.Layer):
+ """Implements the encoder module/block of MOSAIC model.
+
+ Spatial Pyramid Pooling and Multi-kernel Conv layer
+ SpatialPyramidPoolingMultiKernelConv
+ References:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ context](https://arxiv.org/pdf/2112.11623.pdf)
+ """
+
+ def __init__(
+ self,
+ encoder_input_level: Optional[Union[str, int]] = '4',
+ branch_filter_depths: Optional[List[int]] = None,
+ conv_kernel_sizes: Optional[List[int]] = None,
+ pyramid_pool_bin_nums: Optional[List[int]] = None,
+ use_sync_bn: bool = False,
+ batchnorm_momentum: float = 0.99,
+ batchnorm_epsilon: float = 0.001,
+ activation: str = 'relu',
+ dropout_rate: float = 0.1,
+ kernel_initializer: str = 'glorot_uniform',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ interpolation: str = 'bilinear',
+ use_depthwise_convolution: bool = True,
+ **kwargs):
+ """Initializes a MOSAIC encoder block which is deployed after a backbone.
+
+ Args:
+ encoder_input_level: An optional `str` or integer specifying the level of
+ backbone outputs as the input to the encoder.
+ branch_filter_depths: A list of integers for the number of convolution
+ channels in each branch at a pyramid level after SpatialPyramidPooling.
+ conv_kernel_sizes: A list of integers representing the convolution kernel
+ sizes in the Multi-kernel Convolution blocks in the encoder.
+ pyramid_pool_bin_nums: A list of integers for the number of bins at each
+ level of the Spatial Pyramid Pooling.
+ use_sync_bn: A bool, whether or not to use sync batch normalization.
+ batchnorm_momentum: A float for the momentum in BatchNorm. Defaults to
+ 0.99.
+ batchnorm_epsilon: A float for the epsilon value in BatchNorm. Defaults to
+ 0.001.
+ activation: A `str` for the activation function type. Defaults to 'relu'.
+ dropout_rate: A float between 0 and 1. Fraction of the input units to drop
+ out, which will be used directly as the `rate` of the Dropout layer at
+ the end of the encoder. Defaults to 0.1.
+ kernel_initializer: Kernel initializer for conv layers. Defaults to
+ `glorot_uniform`.
+ kernel_regularizer: Kernel regularizer for conv layers. Defaults to None.
+ interpolation: The interpolation method for upsampling. Defaults to
+ `bilinear`.
+ use_depthwise_convolution: Use depthwise separable convolusions in the
+ Multi-kernel Convolution blocks in the encoder.
+ **kwargs: Other keyword arguments for the layer.
+ """
+ super().__init__(**kwargs)
+
+ self._encoder_input_level = str(encoder_input_level)
+ if branch_filter_depths is None:
+ branch_filter_depths = [64, 64]
+ self._branch_filter_depths = branch_filter_depths
+ if conv_kernel_sizes is None:
+ conv_kernel_sizes = [3, 5]
+ self._conv_kernel_sizes = conv_kernel_sizes
+ if pyramid_pool_bin_nums is None:
+ pyramid_pool_bin_nums = [1, 4, 8, 16]
+ self._pyramid_pool_bin_nums = pyramid_pool_bin_nums
+ self._use_sync_bn = use_sync_bn
+ self._batchnorm_momentum = batchnorm_momentum
+ self._batchnorm_epsilon = batchnorm_epsilon
+ self._activation = activation
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._interpolation = interpolation
+ self._use_depthwise_convolution = use_depthwise_convolution
+ self._activation_fn = tf_utils.get_activation(activation)
+
+ if self._use_sync_bn:
+ self._bn_op = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ self._bn_op = tf_keras.layers.BatchNormalization
+
+ self._dropout_rate = dropout_rate
+ if dropout_rate:
+ self._encoder_end_dropout_layer = tf_keras.layers.Dropout(
+ rate=dropout_rate)
+ else:
+ self._encoder_end_dropout_layer = None
+
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ self._channel_axis = -1
+ else:
+ self._bn_axis = 1
+ self._channel_axis = 1
+
+ def _get_bin_pool_kernel_and_stride(
+ self,
+ input_size: int,
+ num_of_bin: int) -> Tuple[int, int]:
+ """Calculates the kernel size and stride for spatial bin pooling.
+
+ Args:
+ input_size: Input dimension (a scalar).
+ num_of_bin: The number of bins used for spatial bin pooling.
+
+ Returns:
+ The Kernel and Stride for spatial bin pooling (a scalar).
+ """
+ bin_overlap = int(input_size % num_of_bin)
+ pooling_stride = int(input_size // num_of_bin)
+ pooling_kernel = pooling_stride + bin_overlap
+ return pooling_kernel, pooling_stride
+
+ def build(
+ self, input_shape: Union[tf.TensorShape, Dict[str,
+ tf.TensorShape]]) -> None:
+ """Builds this MOSAIC encoder block with the given single input shape."""
+ input_shape = (
+ input_shape[self._encoder_input_level]
+ if isinstance(input_shape, dict) else input_shape)
+ self._data_format = tf_keras.backend.image_data_format()
+ if self._data_format == 'channels_last':
+ height = input_shape[1]
+ width = input_shape[2]
+ else:
+ height = input_shape[2]
+ width = input_shape[3]
+
+ self._global_pool_branch = None
+ self._spatial_pyramid = []
+
+ for pyramid_pool_bin_num in self._pyramid_pool_bin_nums:
+ if pyramid_pool_bin_num == 1:
+ global_pool = tf_keras.layers.GlobalAveragePooling2D(
+ data_format=self._data_format, keepdims=True)
+ global_projection = tf_keras.layers.Conv2D(
+ filters=max(self._branch_filter_depths),
+ kernel_size=(1, 1),
+ padding='same',
+ activation=None,
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ use_bias=False)
+ batch_norm_global_branch = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ # Use list manually instead of tf_keras.Sequential([])
+ self._global_pool_branch = [
+ global_pool,
+ global_projection,
+ batch_norm_global_branch,
+ ]
+ else:
+ if height < pyramid_pool_bin_num or width < pyramid_pool_bin_num:
+ raise ValueError('The number of pooling bins must be smaller than '
+ 'input sizes.')
+ assert pyramid_pool_bin_num >= 2, (
+ 'Except for the gloabl pooling, the number of bins in pyramid '
+ 'pooling must be at least two.')
+ pool_height, stride_height = self._get_bin_pool_kernel_and_stride(
+ height, pyramid_pool_bin_num)
+ pool_width, stride_width = self._get_bin_pool_kernel_and_stride(
+ width, pyramid_pool_bin_num)
+ bin_pool_level = tf_keras.layers.AveragePooling2D(
+ pool_size=(pool_height, pool_width),
+ strides=(stride_height, stride_width),
+ padding='valid',
+ data_format=self._data_format)
+ self._spatial_pyramid.append(bin_pool_level)
+
+ # Grouped multi-kernel Convolution.
+ self._multi_kernel_group_conv = MultiKernelGroupConvBlock(
+ output_filter_depths=self._branch_filter_depths,
+ kernel_sizes=self._conv_kernel_sizes,
+ use_sync_bn=self._use_sync_bn,
+ batchnorm_momentum=self._batchnorm_momentum,
+ batchnorm_epsilon=self._batchnorm_epsilon,
+ activation=self._activation,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ use_depthwise_convolution=self._use_depthwise_convolution)
+
+ # Encoder's final 1x1 feature projection.
+ # Considering the relatively large #channels merged before projection,
+ # enlarge the projection #channels to the sum of the filter depths of
+ # branches.
+ self._output_channels = sum(self._branch_filter_depths)
+ # Use list manually instead of tf_keras.Sequential([]).
+ self._encoder_projection = [
+ tf_keras.layers.Conv2D(
+ filters=self._output_channels,
+ kernel_size=(1, 1),
+ padding='same',
+ activation=None,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ use_bias=False),
+ self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon),
+ ]
+ # Use the TF2 default feature alignment rule for bilinear resizing.
+ self._upsample = tf_keras.layers.Resizing(
+ height,
+ width,
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ self._concat_layer = tf_keras.layers.Concatenate(axis=self._channel_axis)
+
+ def call(self,
+ inputs: Union[tf.Tensor, Dict[str, tf.Tensor]],
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Calls this MOSAIC encoder block with the given input."""
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+ input_from_backbone_output = (
+ inputs[self._encoder_input_level]
+ if isinstance(inputs, dict) else inputs)
+ branches = []
+ # Original features from the final output of the backbone.
+ branches.append(input_from_backbone_output)
+ if self._spatial_pyramid:
+ for bin_pool_level in self._spatial_pyramid:
+ x = input_from_backbone_output
+ x = bin_pool_level(x)
+ x = self._multi_kernel_group_conv(x, training=training)
+ x = self._upsample(x)
+ branches.append(x)
+ if self._global_pool_branch is not None:
+ x = input_from_backbone_output
+ for layer in self._global_pool_branch:
+ x = layer(x, training=training)
+ x = self._activation_fn(x)
+ x = self._upsample(x)
+ branches.append(x)
+ x = self._concat_layer(branches)
+ for layer in self._encoder_projection:
+ x = layer(x, training=training)
+ x = self._activation_fn(x)
+ if self._encoder_end_dropout_layer is not None:
+ x = self._encoder_end_dropout_layer(x, training=training)
+ return x
+
+ def get_config(self) -> Dict[str, Any]:
+ """Returns a config dictionary for initialization from serialization."""
+ config = {
+ 'encoder_input_level': self._encoder_input_level,
+ 'branch_filter_depths': self._branch_filter_depths,
+ 'conv_kernel_sizes': self._conv_kernel_sizes,
+ 'pyramid_pool_bin_nums': self._pyramid_pool_bin_nums,
+ 'use_sync_bn': self._use_sync_bn,
+ 'batchnorm_momentum': self._batchnorm_momentum,
+ 'batchnorm_epsilon': self._batchnorm_epsilon,
+ 'activation': self._activation,
+ 'dropout_rate': self._dropout_rate,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'interpolation': self._interpolation,
+ 'use_depthwise_convolution': self._use_depthwise_convolution,
+ }
+ base_config = super().get_config()
+ base_config.update(config)
+ return base_config
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class DecoderSumMergeBlock(tf_keras.layers.Layer):
+ """Implements the decoder feature sum merge block of MOSAIC model.
+
+ This block is used in the decoder of segmentation head introduced in MOSAIC.
+ It essentially merges a high-resolution feature map of a low semantic level
+ and a low-resolution feature map of a higher semantic level by 'Sum-Merge'.
+ """
+
+ def __init__(
+ self,
+ decoder_projected_depth: int,
+ output_size: Tuple[int, int] = (0, 0),
+ use_sync_bn: bool = False,
+ batchnorm_momentum: float = 0.99,
+ batchnorm_epsilon: float = 0.001,
+ activation: str = 'relu',
+ kernel_initializer: str = 'GlorotUniform',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ interpolation: str = 'bilinear',
+ **kwargs):
+ """Initialize a sum-merge block for one decoder stage.
+
+ Args:
+ decoder_projected_depth: An integer representing the number of output
+ channels of this sum-merge block in the decoder.
+ output_size: A Tuple of integers representing the output height and width
+ of the feature maps from this sum-merge block. Defaults to (0, 0),
+ where the output size is set the same as the high-resolution branch.
+ use_sync_bn: A bool, whether or not to use sync batch normalization.
+ batchnorm_momentum: A float for the momentum in BatchNorm. Defaults to
+ 0.99.
+ batchnorm_epsilon: A float for the epsilon value in BatchNorm. Defaults to
+ 0.001.
+ activation: A `str` for the activation function type. Defaults to 'relu'.
+ kernel_initializer: Kernel initializer for conv layers. Defaults to
+ `glorot_uniform`.
+ kernel_regularizer: Kernel regularizer for conv layers. Defaults to None.
+ interpolation: The interpolation method for upsampling. Defaults to
+ `bilinear`.
+ **kwargs: Other keyword arguments for the layer.
+ """
+ super(DecoderSumMergeBlock, self).__init__(**kwargs)
+
+ self._decoder_projected_depth = decoder_projected_depth
+ self._output_size = output_size
+ self._low_res_branch = []
+ self._upsample_low_res = None
+ self._high_res_branch = []
+ self._upsample_high_res = None
+
+ self._use_sync_bn = use_sync_bn
+ self._batchnorm_momentum = batchnorm_momentum
+ self._batchnorm_epsilon = batchnorm_epsilon
+ self._activation = activation
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._interpolation = interpolation
+ # Apply BN before activation. Putting BN between conv and activation also
+ # helps quantization where conv+bn+activation are fused into a single op.
+ self._activation_fn = tf_utils.get_activation(activation)
+ if self._use_sync_bn:
+ self._bn_op = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ self._bn_op = tf_keras.layers.BatchNormalization
+
+ self._bn_axis = (
+ -1
+ if tf_keras.backend.image_data_format() == 'channels_last' else 1)
+ self._channel_axis = (
+ -1
+ if tf_keras.backend.image_data_format() == 'channels_last' else 1)
+ self._add_layer = tf_keras.layers.Add()
+
+ def build(
+ self,
+ input_shape: Tuple[tf.TensorShape, tf.TensorShape]) -> None:
+ """Builds the block with the given input shape."""
+ # Assume backbone features of the same level are concated before input.
+ low_res_input_shape = input_shape[0]
+ high_res_input_shape = input_shape[1]
+ low_res_channels = low_res_input_shape[self._channel_axis]
+ high_res_channels = high_res_input_shape[self._channel_axis]
+
+ if low_res_channels != self._decoder_projected_depth:
+ low_res_feature_conv = tf_keras.layers.Conv2D(
+ filters=self._decoder_projected_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=None,
+ use_bias=False)
+ batchnorm_op = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ self._low_res_branch.extend([
+ low_res_feature_conv,
+ batchnorm_op,
+ ])
+ if high_res_channels != self._decoder_projected_depth:
+ high_res_feature_conv = tf_keras.layers.Conv2D(
+ filters=self._decoder_projected_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=None,
+ use_bias=False)
+ batchnorm_op_high = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ self._high_res_branch.extend([
+ high_res_feature_conv,
+ batchnorm_op_high,
+ ])
+ # Resize feature maps.
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ low_res_height = low_res_input_shape[1]
+ low_res_width = low_res_input_shape[2]
+ high_res_height = high_res_input_shape[1]
+ high_res_width = high_res_input_shape[2]
+ else:
+ low_res_height = low_res_input_shape[2]
+ low_res_width = low_res_input_shape[3]
+ high_res_height = high_res_input_shape[2]
+ high_res_width = high_res_input_shape[3]
+ if (self._output_size[0] == 0 or self._output_size[1] == 0):
+ self._output_size = (high_res_height, high_res_width)
+ if (low_res_height != self._output_size[0] or
+ low_res_width != self._output_size[1]):
+ self._upsample_low_res = tf_keras.layers.Resizing(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ if (high_res_height != self._output_size[0] or
+ high_res_width != self._output_size[1]):
+ self._upsample_high_res = tf_keras.layers.Resizing(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+
+ def call(self,
+ inputs: Tuple[tf.Tensor, tf.Tensor],
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Calls this decoder sum-merge block with the given input.
+
+ Args:
+ inputs: A Tuple of tensors consisting of a low-resolution higher-semantic
+ level feature map from the encoder as the first item and a higher
+ resolution lower-level feature map from the backbone as the second item.
+ training: a `bool` indicating whether it is in `training` mode.
+ Note: the first item of the input Tuple takes a lower-resolution feature map
+ and the second item of the input Tuple takes a higher-resolution branch.
+
+ Returns:
+ A tensor representing the sum-merged decoder feature map.
+ """
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+ x_low_res = inputs[0]
+ x_high_res = inputs[1]
+ if self._low_res_branch:
+ for layer in self._low_res_branch:
+ x_low_res = layer(x_low_res, training=training)
+ x_low_res = self._activation_fn(x_low_res)
+ if self._high_res_branch:
+ for layer in self._high_res_branch:
+ x_high_res = layer(x_high_res, training=training)
+ x_high_res = self._activation_fn(x_high_res)
+ if self._upsample_low_res is not None:
+ x_low_res = self._upsample_low_res(x_low_res)
+ if self._upsample_high_res is not None:
+ x_high_res = self._upsample_high_res(x_high_res)
+ output = self._add_layer([x_low_res, x_high_res])
+ return output
+
+ def get_config(self) -> Dict[str, Any]:
+ """Returns a config dictionary for initialization from serialization."""
+ config = {
+ 'decoder_projected_depth': self._decoder_projected_depth,
+ 'output_size': self._output_size,
+ 'use_sync_bn': self._use_sync_bn,
+ 'batchnorm_momentum': self._batchnorm_momentum,
+ 'batchnorm_epsilon': self._batchnorm_epsilon,
+ 'activation': self._activation,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'interpolation': self._interpolation,
+ }
+ base_config = super(DecoderSumMergeBlock, self).get_config()
+ base_config.update(config)
+ return base_config
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class DecoderConcatMergeBlock(tf_keras.layers.Layer):
+ """Implements the decoder feature concat merge block of MOSAIC model.
+
+ This block is used in the decoder of segmentation head introduced in MOSAIC.
+ It essentially merges a high-resolution feature map of a low semantic level
+ and a low-resolution feature of a higher semantic level by 'Concat-Merge'.
+ """
+
+ def __init__(
+ self,
+ decoder_internal_depth: int,
+ decoder_projected_depth: int,
+ output_size: Tuple[int, int] = (0, 0),
+ use_sync_bn: bool = False,
+ batchnorm_momentum: float = 0.99,
+ batchnorm_epsilon: float = 0.001,
+ activation: str = 'relu',
+ kernel_initializer: str = 'GlorotUniform',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ interpolation: str = 'bilinear',
+ **kwargs):
+ """Initializes a concat-merge block for one decoder stage.
+
+ Args:
+ decoder_internal_depth: An integer representing the number of internal
+ channels of this concat-merge block in the decoder.
+ decoder_projected_depth: An integer representing the number of output
+ channels of this concat-merge block in the decoder.
+ output_size: A Tuple of integers representing the output height and width
+ of the feature maps from this concat-merge block. Defaults to (0, 0),
+ where the output size is set the same as the high-resolution branch.
+ use_sync_bn: A bool, whether or not to use sync batch normalization.
+ batchnorm_momentum: A float for the momentum in BatchNorm. Defaults to
+ 0.99.
+ batchnorm_epsilon: A float for the epsilon value in BatchNorm. Defaults to
+ 0.001.
+ activation: A `str` for the activation function type. Defaults to 'relu'.
+ kernel_initializer: Kernel initializer for conv layers. Defaults to
+ `glorot_uniform`.
+ kernel_regularizer: Kernel regularizer for conv layers. Defaults to None.
+ interpolation: The interpolation method for upsampling. Defaults to
+ `bilinear`.
+ **kwargs: Other keyword arguments for the layer.
+ """
+ super(DecoderConcatMergeBlock, self).__init__(**kwargs)
+
+ self._decoder_internal_depth = decoder_internal_depth
+ self._decoder_projected_depth = decoder_projected_depth
+ self._output_size = output_size
+ self._upsample_low_res = None
+ self._upsample_high_res = None
+
+ self._use_sync_bn = use_sync_bn
+ self._batchnorm_momentum = batchnorm_momentum
+ self._batchnorm_epsilon = batchnorm_epsilon
+ self._activation = activation
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._interpolation = interpolation
+ # Apply BN before activation. Putting BN between conv and activation also
+ # helps quantization where conv+bn+activation are fused into a single op.
+ self._activation_fn = tf_utils.get_activation(activation)
+ if self._use_sync_bn:
+ self._bn_op = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ self._bn_op = tf_keras.layers.BatchNormalization
+
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ self._channel_axis = -1
+ else:
+ self._bn_axis = 1
+ self._channel_axis = 1
+
+ def build(
+ self,
+ input_shape: Tuple[tf.TensorShape, tf.TensorShape]) -> None:
+ """Builds this block with the given input shape."""
+ # Assume backbone features of the same level are concated before input.
+ low_res_input_shape = input_shape[0]
+ high_res_input_shape = input_shape[1]
+ # Set up resizing feature maps before concat.
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ low_res_height = low_res_input_shape[1]
+ low_res_width = low_res_input_shape[2]
+ high_res_height = high_res_input_shape[1]
+ high_res_width = high_res_input_shape[2]
+ else:
+ low_res_height = low_res_input_shape[2]
+ low_res_width = low_res_input_shape[3]
+ high_res_height = high_res_input_shape[2]
+ high_res_width = high_res_input_shape[3]
+ if (self._output_size[0] == 0 or self._output_size[1] == 0):
+ self._output_size = (high_res_height, high_res_width)
+ if (low_res_height != self._output_size[0] or
+ low_res_width != self._output_size[1]):
+ self._upsample_low_res = tf_keras.layers.Resizing(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ if (high_res_height != self._output_size[0] or
+ high_res_width != self._output_size[1]):
+ self._upsample_high_res = tf_keras.layers.Resizing(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ # Set up a 3-layer separable convolution blocks, i.e.
+ # 1x1->BN->RELU + Depthwise->BN->RELU + 1x1->BN->RELU.
+ initial_feature_conv = tf_keras.layers.Conv2D(
+ filters=self._decoder_internal_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=None,
+ use_bias=False)
+ batchnorm_op1 = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation1 = self._activation_fn
+ depthwise_conv = tf_keras.layers.DepthwiseConv2D(
+ kernel_size=(3, 3),
+ depth_multiplier=1,
+ padding='same',
+ depthwise_regularizer=self._kernel_regularizer,
+ depthwise_initializer=self._kernel_initializer,
+ use_bias=False)
+ batchnorm_op2 = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation2 = self._activation_fn
+ project_feature_conv = tf_keras.layers.Conv2D(
+ filters=self._decoder_projected_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=None,
+ use_bias=False)
+ batchnorm_op3 = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation3 = self._activation_fn
+ self._feature_fusion_block = [
+ initial_feature_conv,
+ batchnorm_op1,
+ activation1,
+ depthwise_conv,
+ batchnorm_op2,
+ activation2,
+ project_feature_conv,
+ batchnorm_op3,
+ activation3,
+ ]
+ self._concat_layer = tf_keras.layers.Concatenate(axis=self._channel_axis)
+
+ def call(self,
+ inputs: Tuple[tf.Tensor, tf.Tensor],
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Calls this concat-merge block with the given inputs.
+
+ Args:
+ inputs: A Tuple of tensors consisting of a lower-level higher-resolution
+ feature map from the backbone as the first item and a higher-level
+ lower-resolution feature map from the encoder as the second item.
+ training: a `Boolean` indicating whether it is in `training` mode.
+
+ Returns:
+ A tensor representing the concat-merged decoder feature map.
+ """
+ low_res_input = inputs[0]
+ high_res_input = inputs[1]
+ if self._upsample_low_res is not None:
+ low_res_input = self._upsample_low_res(low_res_input)
+ if self._upsample_high_res is not None:
+ high_res_input = self._upsample_high_res(high_res_input)
+ decoder_feature_list = [low_res_input, high_res_input]
+ x = self._concat_layer(decoder_feature_list)
+ for layer in self._feature_fusion_block:
+ if isinstance(layer, tf_keras.layers.Layer):
+ x = layer(x, training=training)
+ else:
+ x = layer(x)
+ return x
+
+ def get_config(self) -> Dict[str, Any]:
+ """Returns a config dictionary for initialization from serialization."""
+ config = {
+ 'decoder_internal_depth': self._decoder_internal_depth,
+ 'decoder_projected_depth': self._decoder_projected_depth,
+ 'output_size': self._output_size,
+ 'use_sync_bn': self._use_sync_bn,
+ 'batchnorm_momentum': self._batchnorm_momentum,
+ 'batchnorm_epsilon': self._batchnorm_epsilon,
+ 'activation': self._activation,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'interpolation': self._interpolation,
+ }
+ base_config = super(DecoderConcatMergeBlock, self).get_config()
+ base_config.update(config)
+ return base_config
diff --git a/official/projects/mosaic/modeling/mosaic_blocks_test.py b/official/projects/mosaic/modeling/mosaic_blocks_test.py
new file mode 100644
index 00000000000..7ac1c4f8636
--- /dev/null
+++ b/official/projects/mosaic/modeling/mosaic_blocks_test.py
@@ -0,0 +1,99 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for mosaic_blocks."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.mosaic.modeling import mosaic_blocks
+
+
+class MosaicBlocksTest(parameterized.TestCase, tf.test.TestCase):
+
+ def test_multi_kernel_group_conv_block(self):
+ block = mosaic_blocks.MultiKernelGroupConvBlock([64, 64], [3, 5])
+ inputs = tf.ones([1, 4, 4, 448])
+ outputs = block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 4, 4, 128])
+
+ def test_mosaic_encoder_block(self):
+ block = mosaic_blocks.MosaicEncoderBlock(
+ encoder_input_level=4,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=[1, 4, 8, 16])
+ inputs = tf.ones([1, 32, 32, 448])
+ outputs = block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 32, 32, 128])
+
+ def test_mosaic_encoder_block_odd_input_overlap_pool(self):
+ block = mosaic_blocks.MosaicEncoderBlock(
+ encoder_input_level=4,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=[1, 4, 8, 16])
+ inputs = tf.ones([1, 31, 31, 448])
+ outputs = block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 31, 31, 128])
+
+ def test_mosaic_encoder_non_separable_block(self):
+ block = mosaic_blocks.MosaicEncoderBlock(
+ encoder_input_level=4,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=[1, 4, 8, 16],
+ use_depthwise_convolution=False)
+ inputs = tf.ones([1, 32, 32, 448])
+ outputs = block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 32, 32, 128])
+
+ def test_mosaic_decoder_concat_merge_block(self):
+ concat_merge_block = mosaic_blocks.DecoderConcatMergeBlock(64, 32, [64, 64])
+ inputs = [tf.ones([1, 32, 32, 128]), tf.ones([1, 64, 64, 192])]
+ outputs = concat_merge_block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 64, 64, 32])
+
+ def test_mosaic_decoder_concat_merge_block_default_output_size(self):
+ concat_merge_block = mosaic_blocks.DecoderConcatMergeBlock(64, 32)
+ inputs = [tf.ones([1, 32, 32, 128]), tf.ones([1, 64, 64, 192])]
+ outputs = concat_merge_block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 64, 64, 32])
+
+ def test_mosaic_decoder_concat_merge_block_default_output_size_4x(self):
+ concat_merge_block = mosaic_blocks.DecoderConcatMergeBlock(64, 32)
+ inputs = [tf.ones([1, 32, 32, 128]), tf.ones([1, 128, 128, 192])]
+ outputs = concat_merge_block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 128, 128, 32])
+
+ def test_mosaic_decoder_concat_merge_block_default_output_size_4x_rec(self):
+ concat_merge_block = mosaic_blocks.DecoderConcatMergeBlock(64, 32)
+ inputs = [tf.ones([1, 32, 64, 128]), tf.ones([1, 128, 256, 64])]
+ outputs = concat_merge_block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 128, 256, 32])
+
+ def test_mosaic_decoder_sum_merge_block(self):
+ concat_merge_block = mosaic_blocks.DecoderSumMergeBlock(32, [128, 128])
+ inputs = [tf.ones([1, 64, 64, 32]), tf.ones([1, 128, 128, 64])]
+ outputs = concat_merge_block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 128, 128, 32])
+
+ def test_mosaic_decoder_sum_merge_block_default_output_size(self):
+ concat_merge_block = mosaic_blocks.DecoderSumMergeBlock(32)
+ inputs = [tf.ones([1, 64, 64, 32]), tf.ones([1, 128, 128, 64])]
+ outputs = concat_merge_block(inputs)
+ self.assertAllEqual(outputs.shape, [1, 128, 128, 32])
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/modeling/mosaic_head.py b/official/projects/mosaic/modeling/mosaic_head.py
new file mode 100644
index 00000000000..1d50f1cf9d1
--- /dev/null
+++ b/official/projects/mosaic/modeling/mosaic_head.py
@@ -0,0 +1,242 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definitions of segmentation head of the MOSAIC model."""
+from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.projects.mosaic.modeling import mosaic_blocks
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MosaicDecoderHead(tf_keras.layers.Layer):
+ """Creates a MOSAIC decoder in segmentation head.
+
+ Reference:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ Context](https://arxiv.org/pdf/2112.11623.pdf)
+ """
+
+ def __init__(
+ self,
+ num_classes: int,
+ decoder_input_levels: Optional[List[str]] = None,
+ decoder_stage_merge_styles: Optional[List[str]] = None,
+ decoder_filters: Optional[List[int]] = None,
+ decoder_projected_filters: Optional[List[int]] = None,
+ encoder_end_level: Optional[int] = 4,
+ use_additional_classifier_layer: bool = False,
+ classifier_kernel_size: int = 1,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ batchnorm_momentum: float = 0.99,
+ batchnorm_epsilon: float = 0.001,
+ kernel_initializer: str = 'GlorotUniform',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ interpolation: str = 'bilinear',
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initializes a MOSAIC segmentation head.
+
+ Args:
+ num_classes: An `int` number of mask classification categories. The number
+ of classes does not include background class.
+ decoder_input_levels: A list of `str` specifying additional
+ input levels from the backbone outputs for mask refinement in decoder.
+ decoder_stage_merge_styles: A list of `str` specifying the merge style at
+ each stage of the decoder, merge styles can be 'concat_merge' or
+ 'sum_merge'.
+ decoder_filters: A list of integers specifying the number of channels used
+ at each decoder stage. Note: this only has affects if the decoder merge
+ style is 'concat_merge'.
+ decoder_projected_filters: A list of integers specifying the number of
+ projected channels at the end of each decoder stage.
+ encoder_end_level: An optional integer specifying the output level of the
+ encoder stage, which is used if the input from the encoder to the
+ decoder head is a dictionary.
+ use_additional_classifier_layer: A `bool` specifying whether to use an
+ additional classifier layer or not. It must be True if the final decoder
+ projected filters does not match the `num_classes`.
+ classifier_kernel_size: An `int` number to specify the kernel size of the
+ classifier layer.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ batchnorm_momentum: A `float` of normalization momentum for the moving
+ average.
+ batchnorm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_initializer: Kernel initializer for conv layers. Defaults to
+ `glorot_uniform`.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ interpolation: The interpolation method for upsampling. Defaults to
+ `bilinear`.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super(MosaicDecoderHead, self).__init__(**kwargs)
+
+ # Assuming 'decoder_input_levels' are sorted in descending order and the
+ # other setting are listed in the order according to 'decoder_input_levels'.
+ if decoder_input_levels is None:
+ decoder_input_levels = ['3', '2']
+ if decoder_stage_merge_styles is None:
+ decoder_stage_merge_styles = ['concat_merge', 'sum_merge']
+ if decoder_filters is None:
+ decoder_filters = [64, 64]
+ if decoder_projected_filters is None:
+ decoder_projected_filters = [32, 32]
+ self._decoder_input_levels = decoder_input_levels
+ self._decoder_stage_merge_styles = decoder_stage_merge_styles
+ self._decoder_filters = decoder_filters
+ self._decoder_projected_filters = decoder_projected_filters
+ if (len(decoder_input_levels) != len(decoder_stage_merge_styles) or
+ len(decoder_input_levels) != len(decoder_filters) or
+ len(decoder_input_levels) != len(decoder_projected_filters)):
+ raise ValueError('The number of Decoder inputs and settings must match.')
+ self._merge_stages = []
+ for (stage_merge_style, decoder_filter,
+ decoder_projected_filter) in zip(decoder_stage_merge_styles,
+ decoder_filters,
+ decoder_projected_filters):
+ if stage_merge_style == 'concat_merge':
+ concat_merge_stage = mosaic_blocks.DecoderConcatMergeBlock(
+ decoder_internal_depth=decoder_filter,
+ decoder_projected_depth=decoder_projected_filter,
+ output_size=(0, 0),
+ use_sync_bn=use_sync_bn,
+ batchnorm_momentum=batchnorm_momentum,
+ batchnorm_epsilon=batchnorm_epsilon,
+ activation=activation,
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ interpolation=interpolation)
+ self._merge_stages.append(concat_merge_stage)
+ elif stage_merge_style == 'sum_merge':
+ sum_merge_stage = mosaic_blocks.DecoderSumMergeBlock(
+ decoder_projected_depth=decoder_projected_filter,
+ output_size=(0, 0),
+ use_sync_bn=use_sync_bn,
+ batchnorm_momentum=batchnorm_momentum,
+ batchnorm_epsilon=batchnorm_epsilon,
+ activation=activation,
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ interpolation=interpolation)
+ self._merge_stages.append(sum_merge_stage)
+ else:
+ raise ValueError(
+ 'A stage merge style in MOSAIC Decoder can only be concat_merge '
+ 'or sum_merge.')
+
+ # Concat merge or sum merge does not require an additional classifer layer
+ # unless the final decoder projected filter does not match num_classes.
+ final_decoder_projected_filter = decoder_projected_filters[-1]
+ if (final_decoder_projected_filter != num_classes and
+ not use_additional_classifier_layer):
+ raise ValueError('Additional classifier layer is needed if final decoder '
+ 'projected filters does not match num_classes!')
+ self._use_additional_classifier_layer = use_additional_classifier_layer
+ if use_additional_classifier_layer:
+ # This additional classification layer uses different kernel
+ # initializers and bias compared to earlier blocks.
+ self._pixelwise_classifier = tf_keras.layers.Conv2D(
+ name='pixelwise_classifier',
+ filters=num_classes,
+ kernel_size=classifier_kernel_size,
+ padding='same',
+ bias_initializer=tf.zeros_initializer(),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ use_bias=True)
+ self._activation_fn = tf_utils.get_activation(activation)
+
+ self._config_dict = {
+ 'num_classes': num_classes,
+ 'decoder_input_levels': decoder_input_levels,
+ 'decoder_stage_merge_styles': decoder_stage_merge_styles,
+ 'decoder_filters': decoder_filters,
+ 'decoder_projected_filters': decoder_projected_filters,
+ 'encoder_end_level': encoder_end_level,
+ 'use_additional_classifier_layer': use_additional_classifier_layer,
+ 'classifier_kernel_size': classifier_kernel_size,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'batchnorm_momentum': batchnorm_momentum,
+ 'batchnorm_epsilon': batchnorm_epsilon,
+ 'kernel_initializer': kernel_initializer,
+ 'kernel_regularizer': kernel_regularizer,
+ 'interpolation': interpolation,
+ 'bias_regularizer': bias_regularizer
+ }
+
+ def call(self,
+ inputs: Tuple[Union[tf.Tensor, Mapping[str, tf.Tensor]],
+ Union[tf.Tensor, Mapping[str, tf.Tensor]]],
+ training: Optional[bool] = None) -> tf.Tensor:
+ """Forward pass of the segmentation head.
+
+ It supports a tuple of 2 elements. Each element is a tensor or a tensor
+ dictionary. The first one is the final (low-resolution) encoder endpoints,
+ and the second one is higher-resolution backbone endpoints.
+ When inputs are tensors, they are from a single level of feature maps.
+ When inputs are dictionaries, they contain multiple levels of feature maps,
+ where the key is the level/index of feature map.
+ Note: 'level' denotes the number of 2x downsampling, defined in backbone.
+
+ Args:
+ inputs: A tuple of 2 elements, each element can either be a tensor
+ representing feature maps or 1 dictionary of tensors:
+ - key: A `str` of the level of the multilevel features.
+ - values: A `tf.Tensor` of the feature map tensors.
+ The first is encoder endpoints, and the second is backbone endpoints.
+ training: a `Boolean` indicating whether it is in `training` mode.
+ Returns:
+ segmentation mask prediction logits: A `tf.Tensor` representing the
+ output logits before the final segmentation mask.
+ """
+
+ encoder_outputs = inputs[0]
+ backbone_outputs = inputs[1]
+ y = encoder_outputs[str(
+ self._config_dict['encoder_end_level'])] if isinstance(
+ encoder_outputs, dict) else encoder_outputs
+ if isinstance(backbone_outputs, dict):
+ for level, merge_stage in zip(
+ self._decoder_input_levels, self._merge_stages):
+ x = backbone_outputs[str(level)]
+ y = merge_stage([y, x], training=training)
+ else:
+ x = backbone_outputs
+ y = self._merge_stages[0]([y, x], training=training)
+
+ if self._use_additional_classifier_layer:
+ y = self._pixelwise_classifier(y)
+ y = self._activation_fn(y)
+
+ return y # pyrefly: ignore[bad-return]
+
+ def get_config(self) -> Dict[str, Any]:
+ """Returns a config dictionary for initialization from serialization."""
+ base_config = super().get_config()
+ base_config.update(self._config_dict)
+ return base_config
+
+ @classmethod
+ def from_config(cls, config: Dict[str, Any]):
+ return cls(**config)
diff --git a/official/projects/mosaic/modeling/mosaic_head_test.py b/official/projects/mosaic/modeling/mosaic_head_test.py
new file mode 100644
index 00000000000..c6b88149de5
--- /dev/null
+++ b/official/projects/mosaic/modeling/mosaic_head_test.py
@@ -0,0 +1,62 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for mosaic_head."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.mosaic.modeling import mosaic_head
+
+
+class MosaicBlocksTest(parameterized.TestCase, tf.test.TestCase):
+
+ def test_mosaic_head(self):
+ decoder_head = mosaic_head.MosaicDecoderHead(
+ num_classes=32,
+ decoder_input_levels=['3', '2'],
+ decoder_stage_merge_styles=['concat_merge', 'sum_merge'],
+ decoder_filters=[64, 64],
+ decoder_projected_filters=[32, 32])
+ inputs = [
+ tf.ones([1, 32, 32, 128]), {
+ '2': tf.ones([1, 128, 128, 64]),
+ '3': tf.ones([1, 64, 64, 192])
+ }
+ ]
+ outputs = decoder_head(inputs)
+ self.assertAllEqual(outputs.shape, [1, 128, 128, 32])
+
+ def test_mosaic_head_3laterals(self):
+ decoder_head = mosaic_head.MosaicDecoderHead(
+ num_classes=32,
+ decoder_input_levels=[3, 2, 1],
+ decoder_stage_merge_styles=[
+ 'concat_merge', 'concat_merge', 'sum_merge'
+ ],
+ decoder_filters=[64, 64, 64],
+ decoder_projected_filters=[32, 32, 32])
+ inputs = [
+ tf.ones([1, 32, 32, 128]), {
+ '1': tf.ones([1, 256, 256, 64]),
+ '2': tf.ones([1, 128, 128, 64]),
+ '3': tf.ones([1, 64, 64, 192])
+ }
+ ]
+ outputs = decoder_head(inputs)
+ self.assertAllEqual(outputs.shape, [1, 256, 256, 32])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/modeling/mosaic_model.py b/official/projects/mosaic/modeling/mosaic_model.py
new file mode 100644
index 00000000000..cd46f777c86
--- /dev/null
+++ b/official/projects/mosaic/modeling/mosaic_model.py
@@ -0,0 +1,179 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builds the overall MOSAIC segmentation models."""
+from typing import Any, Dict, Optional, Union
+
+import tensorflow as tf, tf_keras
+
+from official.projects.mosaic.configs import mosaic_config
+from official.projects.mosaic.modeling import mosaic_blocks
+from official.projects.mosaic.modeling import mosaic_head
+from official.vision.modeling import backbones
+from official.vision.modeling.heads import segmentation_heads
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MosaicSegmentationModel(tf_keras.Model):
+ """A model class for segmentation using MOSAIC.
+
+ Input images are passed through a backbone first. A MOSAIC neck encoder
+ network is then applied, and finally a MOSAIC segmentation head is applied on
+ the outputs of the backbone and neck encoder network. Feature fusion and
+ decoding is done in the segmentation head.
+
+ Reference:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ Context](https://arxiv.org/pdf/2112.11623.pdf)
+ """
+
+ def __init__(self,
+ backbone: tf_keras.Model,
+ head: tf_keras.layers.Layer,
+ neck: Optional[tf_keras.layers.Layer] = None,
+ mask_scoring_head: Optional[tf_keras.layers.Layer] = None,
+ **kwargs):
+ """Segmentation initialization function.
+
+ Args:
+ backbone: A backbone network.
+ head: A segmentation head, e.g. MOSAIC decoder.
+ neck: An optional neck encoder network, e.g. MOSAIC encoder. If it is not
+ provided, the decoder head will be connected directly with the backbone.
+ mask_scoring_head: An optional mask scoring head.
+ **kwargs: keyword arguments to be passed.
+ """
+ super(MosaicSegmentationModel, self).__init__(**kwargs)
+ self._config_dict = {
+ 'backbone': backbone,
+ 'neck': neck,
+ 'head': head,
+ 'mask_scoring_head': mask_scoring_head,
+ }
+ self.backbone = backbone
+ self.neck = neck
+ self.head = head
+ self.mask_scoring_head = mask_scoring_head
+
+ def call(self, # pytype: disable=annotation-type-mismatch,signature-mismatch
+ inputs: tf.Tensor,
+ training: bool = None) -> Dict[str, tf.Tensor]: # pyrefly: ignore[bad-function-definition]
+ backbone_features = self.backbone(inputs)
+
+ if self.neck is not None:
+ neck_features = self.neck(backbone_features, training=training)
+ else:
+ neck_features = backbone_features
+
+ logits = self.head([neck_features, backbone_features], training=training)
+ outputs = {'logits': logits}
+
+ if self.mask_scoring_head:
+ mask_scores = self.mask_scoring_head(logits)
+ outputs.update({'mask_scores': mask_scores})
+
+ return outputs
+
+ @property
+ def checkpoint_items(
+ self) -> Dict[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(backbone=self.backbone, head=self.head)
+ if self.neck is not None:
+ items.update(neck=self.neck)
+ if self.mask_scoring_head is not None:
+ items.update(mask_scoring_head=self.mask_scoring_head)
+ return items
+
+ def get_config(self) -> Dict[str, Any]:
+ """Returns a config dictionary for initialization from serialization."""
+ base_config = super().get_config()
+ model_config = base_config
+ model_config.update(self._config_dict)
+ return model_config
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+
+def build_mosaic_segmentation_model(
+ input_specs: tf_keras.layers.InputSpec,
+ model_config: mosaic_config.MosaicSemanticSegmentationModel,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ backbone: Optional[tf_keras.Model] = None,
+ neck: Optional[tf_keras.layers.Layer] = None
+) -> tf_keras.Model:
+ """Builds MOSAIC Segmentation model."""
+ norm_activation_config = model_config.norm_activation
+ if backbone is None:
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=model_config.backbone,
+ norm_activation_config=norm_activation_config,
+ l2_regularizer=l2_regularizer)
+
+ if neck is None:
+ neck_config = model_config.neck
+ neck = mosaic_blocks.MosaicEncoderBlock(
+ encoder_input_level=neck_config.encoder_input_level,
+ branch_filter_depths=neck_config.branch_filter_depths,
+ conv_kernel_sizes=neck_config.conv_kernel_sizes,
+ pyramid_pool_bin_nums=neck_config.pyramid_pool_bin_nums,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ batchnorm_momentum=norm_activation_config.norm_momentum,
+ batchnorm_epsilon=norm_activation_config.norm_epsilon,
+ activation=neck_config.activation,
+ dropout_rate=neck_config.dropout_rate,
+ kernel_initializer=neck_config.kernel_initializer,
+ kernel_regularizer=l2_regularizer,
+ interpolation=neck_config.interpolation,
+ use_depthwise_convolution=neck_config.use_depthwise_convolution)
+
+ head_config = model_config.head
+ head = mosaic_head.MosaicDecoderHead(
+ num_classes=model_config.num_classes,
+ decoder_input_levels=head_config.decoder_input_levels,
+ decoder_stage_merge_styles=head_config.decoder_stage_merge_styles,
+ decoder_filters=head_config.decoder_filters,
+ decoder_projected_filters=head_config.decoder_projected_filters,
+ encoder_end_level=head_config.encoder_end_level,
+ use_additional_classifier_layer=head_config
+ .use_additional_classifier_layer,
+ classifier_kernel_size=head_config.classifier_kernel_size,
+ activation=head_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ batchnorm_momentum=norm_activation_config.norm_momentum,
+ batchnorm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_initializer=head_config.kernel_initializer,
+ kernel_regularizer=l2_regularizer,
+ interpolation=head_config.interpolation)
+
+ mask_scoring_head = None
+ if model_config.mask_scoring_head:
+ mask_scoring_head = segmentation_heads.MaskScoring(
+ num_classes=model_config.num_classes,
+ **model_config.mask_scoring_head.as_dict(),
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer)
+
+ model = MosaicSegmentationModel(
+ backbone=backbone,
+ neck=neck,
+ head=head,
+ mask_scoring_head=mask_scoring_head)
+ return model
diff --git a/official/projects/mosaic/modeling/mosaic_model_test.py b/official/projects/mosaic/modeling/mosaic_model_test.py
new file mode 100644
index 00000000000..3315c82006f
--- /dev/null
+++ b/official/projects/mosaic/modeling/mosaic_model_test.py
@@ -0,0 +1,129 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for the overall MOSAIC segmentation network modeling."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.mosaic.modeling import mosaic_blocks
+from official.projects.mosaic.modeling import mosaic_head
+from official.projects.mosaic.modeling import mosaic_model
+from official.vision.modeling import backbones
+from official.vision.modeling.heads import segmentation_heads
+
+
+class SegmentationNetworkTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (128, [4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (128, [1, 4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (128, [1, 4, 8], [3, 2], ['sum_merge', 'sum_merge']),
+ (128, [1, 4, 8], [3, 2], ['concat_merge', 'concat_merge']),
+ (512, [1, 4, 8, 16], [3, 2], ['concat_merge', 'sum_merge']),
+ (256, [4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (256, [1, 4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (256, [1, 4, 8, 16], [3, 2], ['concat_merge', 'sum_merge']),
+ )
+ def test_mosaic_segmentation_model(self,
+ input_size,
+ pyramid_pool_bin_nums,
+ decoder_input_levels,
+ decoder_stage_merge_styles):
+ """Test for building and calling of a MOSAIC segmentation network."""
+ num_classes = 32
+ inputs = np.random.rand(2, input_size, input_size, 3)
+ tf_keras.backend.set_image_data_format('channels_last')
+ backbone = backbones.MobileNet(model_id='MobileNetMultiAVGSeg')
+ encoder_input_level = 4
+
+ neck = mosaic_blocks.MosaicEncoderBlock(
+ encoder_input_level=encoder_input_level,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=pyramid_pool_bin_nums)
+ head = mosaic_head.MosaicDecoderHead(
+ num_classes=num_classes,
+ decoder_input_levels=decoder_input_levels,
+ decoder_stage_merge_styles=decoder_stage_merge_styles,
+ decoder_filters=[64, 64],
+ decoder_projected_filters=[32, 32])
+
+ mask_scoring_head = segmentation_heads.MaskScoring(
+ num_classes=num_classes,
+ fc_input_size=[4, 4],
+ num_convs=1,
+ num_filters=32,
+ fc_dims=32,
+ num_fcs=1)
+
+ model = mosaic_model.MosaicSegmentationModel(
+ backbone=backbone,
+ head=head,
+ neck=neck,
+ mask_scoring_head=mask_scoring_head,
+ )
+
+ # Calls the MOSAIC model.
+ outputs = model(inputs)
+ level = min(decoder_input_levels)
+ self.assertAllEqual(
+ [2, input_size // (2**level), input_size // (2**level), num_classes],
+ outputs['logits'].numpy().shape)
+ self.assertAllEqual(
+ [2, num_classes],
+ outputs['mask_scores'].numpy().shape)
+
+ def test_serialize_deserialize(self):
+ """Validate the mosaic network can be serialized and deserialized."""
+ num_classes = 8
+ backbone = backbones.ResNet(model_id=50)
+ neck = mosaic_blocks.MosaicEncoderBlock(
+ encoder_input_level=4,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=[1, 4, 8, 16])
+ head = mosaic_head.MosaicDecoderHead(
+ num_classes=num_classes,
+ decoder_input_levels=[3, 2],
+ decoder_stage_merge_styles=['concat_merge', 'sum_merge'],
+ decoder_filters=[64, 64],
+ decoder_projected_filters=[32, 8])
+ mask_scoring_head = segmentation_heads.MaskScoring(
+ num_classes=num_classes,
+ fc_input_size=[4, 4],
+ num_convs=1,
+ num_filters=32,
+ fc_dims=32,
+ num_fcs=1)
+ model = mosaic_model.MosaicSegmentationModel(
+ backbone=backbone,
+ head=head,
+ neck=neck,
+ mask_scoring_head=mask_scoring_head,
+ )
+
+ config = model.get_config()
+ new_model = mosaic_model.MosaicSegmentationModel.from_config(config)
+
+ # Validate that the config can be forced to JSON.
+ _ = new_model.to_json()
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(model.get_config(), new_model.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/mosaic_tasks.py b/official/projects/mosaic/mosaic_tasks.py
new file mode 100644
index 00000000000..ed0340826a3
--- /dev/null
+++ b/official/projects/mosaic/mosaic_tasks.py
@@ -0,0 +1,102 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Task definition for image semantic segmentation with MOSAIC models."""
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.core import task_factory
+from official.projects.mosaic.configs import mosaic_config
+from official.projects.mosaic.modeling import mosaic_model
+from official.vision.tasks import semantic_segmentation as seg_tasks
+
+
+@task_factory.register_task_cls(mosaic_config.MosaicSemanticSegmentationTask)
+class MosaicSemanticSegmentationTask(seg_tasks.SemanticSegmentationTask):
+ """A task for semantic segmentation using MOSAIC model."""
+
+ # Note: the `build_model` is overrided to add an additional `train` flag
+ # for the purpose of indicating the model is built for performing `training`
+ # or `eval`. This is to make sure the model is initialized with proper
+ # `input_shape` if the model will be trained and evaluated in different
+ # `input_shape`. For example, the model is trained with cropping but
+ # evaluated with original shape.
+ def build_model(self, training: bool = True) -> tf_keras.Model:
+ """Builds MOSAIC segmentation model."""
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + self.task_config.model.input_size)
+
+ l2_weight_decay = self.task_config.losses.l2_weight_decay
+ # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
+ # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
+ # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
+ l2_regularizer = (tf_keras.regularizers.l2(
+ l2_weight_decay / 2.0) if l2_weight_decay else None)
+
+ model = mosaic_model.build_mosaic_segmentation_model(
+ input_specs=input_specs,
+ model_config=self.task_config.model,
+ l2_regularizer=l2_regularizer)
+
+ # Note: Create a dummy input and call model instance to initialize.
+ # This ensures all the layers are built; otherwise some layers may be
+ # missing from the model and cannot be associated with variables from
+ # a loaded checkpoint. The input size is determined by whether the model
+ # is built for performing training or eval.
+ if training:
+ input_size = self.task_config.train_data.output_size
+ crop_size = self.task_config.train_data.crop_size
+ if crop_size:
+ input_size = crop_size
+ else:
+ input_size = self.task_config.validation_data.output_size
+
+ if len(self.task_config.model.input_size) == 3:
+ input_channel = self.task_config.model.input_size[-1]
+ else:
+ input_channel = 3
+
+ dummy_input = tf.ones(shape=[1] + input_size + [input_channel])
+ model(dummy_input)
+
+ return model
+
+ def initialize(self, model: tf_keras.Model):
+ """Loads pretrained checkpoint."""
+ if not self.task_config.init_checkpoint:
+ return
+
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+
+ # Restoring checkpoint.
+ if 'all' in self.task_config.init_checkpoint_modules:
+ ckpt = tf.train.Checkpoint(**model.checkpoint_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ else:
+ ckpt_items = {}
+ if 'backbone' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(backbone=model.backbone)
+ if 'neck' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(neck=model.neck)
+
+ ckpt = tf.train.Checkpoint(**ckpt_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
diff --git a/official/projects/mosaic/mosaic_tasks_test.py b/official/projects/mosaic/mosaic_tasks_test.py
new file mode 100644
index 00000000000..f4f831f2439
--- /dev/null
+++ b/official/projects/mosaic/mosaic_tasks_test.py
@@ -0,0 +1,91 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for mosaic task."""
+# pylint: disable=unused-import
+import os
+
+from absl.testing import parameterized
+import orbit
+import tensorflow as tf, tf_keras
+
+from official import vision
+from official.core import exp_factory
+from official.modeling import optimization
+from official.projects.mosaic import mosaic_tasks
+from official.projects.mosaic.configs import mosaic_config as exp_cfg
+from official.vision.dataloaders import tfexample_utils
+
+
+class MosaicTaskTest(parameterized.TestCase, tf.test.TestCase):
+
+ def _create_test_tfrecord(self, tfrecord_file, example, num_samples):
+ examples = [example] * num_samples
+ tfexample_utils.dump_to_tfrecord(
+ record_file=tfrecord_file, tf_examples=examples)
+
+ @parameterized.parameters(
+ ('mosaic_mnv35_cityscapes', True),
+ ('mosaic_mnv35_cityscapes', False),
+ )
+ def test_semantic_segmentation_task(self, test_config, is_training):
+ """Tests mosaic task for training and eval using toy configs."""
+ input_image_size = [1024, 2048]
+ test_tfrecord_file = os.path.join(self.get_temp_dir(), 'seg_test.tfrecord')
+ example = tfexample_utils.create_segmentation_test_example(
+ image_height=input_image_size[0],
+ image_width=input_image_size[1],
+ image_channel=3)
+ self._create_test_tfrecord(
+ tfrecord_file=test_tfrecord_file, example=example, num_samples=10)
+ config = exp_factory.get_exp_config(test_config)
+ # Modify config to suit local testing
+ config.task.model.input_size = [None, None, 3]
+ config.trainer.steps_per_loop = 1
+ config.task.train_data.global_batch_size = 1
+ config.task.validation_data.global_batch_size = 1
+ config.task.train_data.output_size = [1024, 2048]
+ config.task.validation_data.output_size = [1024, 2048]
+ config.task.train_data.crop_size = [512, 512]
+ config.task.train_data.shuffle_buffer_size = 2
+ config.task.validation_data.shuffle_buffer_size = 2
+ config.task.validation_data.input_path = test_tfrecord_file
+ config.task.train_data.input_path = test_tfrecord_file
+ config.train_steps = 1
+ config.task.model.num_classes = 256
+ config.task.model.head.num_classes = 256
+ config.task.model.head.decoder_projected_filters = [256, 256]
+
+ task = mosaic_tasks.MosaicSemanticSegmentationTask(config.task)
+ model = task.build_model(training=is_training)
+ metrics = task.build_metrics(training=is_training)
+
+ strategy = tf.distribute.get_strategy()
+
+ data_config = config.task.train_data if is_training else config.task.validation_data
+ dataset = orbit.utils.make_distributed_dataset(strategy, task.build_inputs,
+ data_config)
+ iterator = iter(dataset)
+ opt_factory = optimization.OptimizerFactory(config.trainer.optimizer_config)
+ optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
+
+ if is_training:
+ logs = task.train_step(next(iterator), model, optimizer, metrics=metrics)
+ else:
+ logs = task.validation_step(next(iterator), model, metrics=metrics)
+
+ self.assertIn('loss', logs)
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/mosaic_tutorial.ipynb b/official/projects/mosaic/mosaic_tutorial.ipynb
new file mode 100644
index 00000000000..08385e31f4f
--- /dev/null
+++ b/official/projects/mosaic/mosaic_tutorial.ipynb
@@ -0,0 +1,356 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Lhn6RdCkunJY"
+ },
+ "source": [
+ "##### Copyright 2022 The TensorFlow Authors."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nFvE5hKZuYy6"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n",
+ "# you may not use this file except in compliance with the License.\n",
+ "# You may obtain a copy of the License at\n",
+ "#\n",
+ "# https://www.apache.org/licenses/LICENSE-2.0\n",
+ "#\n",
+ "# Unless required by applicable law or agreed to in writing, software\n",
+ "# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
+ "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
+ "# See the License for the specific language governing permissions and\n",
+ "# limitations under the License."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cfjBLLxVb61A"
+ },
+ "source": [
+ "# MOSAIC with Model Garden\n",
+ "\n",
+ "This tutorial demonstrates how to load the\n",
+ "[MOSAIC](https://github.com/tensorflow/models/tree/master/official/projects/mosaic) model trained using [Tensorflow Model Garden](https://github.com/tensorflow/models/tree/master/official) library.\n",
+ "\n",
+ "Tensorflow Model Garden contains a collection of\n",
+ "state-of-the-art models, implemented with TensorFlow's high-level APIs. The\n",
+ "implementations demonstrate the best practices for modeling, letting users to\n",
+ "take full advantage of TensorFlow for their research and product development."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hmslpkDEcPk9"
+ },
+ "source": [
+ "## Install Necessary Dependencies"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "WjAMpVVAcTYh"
+ },
+ "source": [
+ "## Import libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "T0JMx3hmcXL5"
+ },
+ "outputs": [],
+ "source": [
+ "import matplotlib.pyplot as plt\n",
+ "import matplotlib.patches as mpatches\n",
+ "\n",
+ "from PIL import Image\n",
+ "from six import BytesIO\n",
+ "from IPython import display\n",
+ "from urllib.request import urlopen\n",
+ "\n",
+ "import numpy as np\n",
+ "import tensorflow as tf\n",
+ "import tensorflow_hub as hub\n",
+ "\n",
+ "import absl.logging\n",
+ "absl.logging.set_verbosity(absl.logging.ERROR)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "02j9_wFYd3H5"
+ },
+ "source": [
+ "## Load Pretrained Model\n",
+ "Note that you may also access the pretrained model provided on [TFHub](https://tfhub.dev/google/mosaic/mobilenetmultiavgseg):\n",
+ "\n",
+ "```\n",
+ "keras_layer = hub.KerasLayer(\n",
+ " 'https://tfhub.dev/google/mosaic/mobilenetmultiavgseg/2',\n",
+ " signature='serving_default',\n",
+ " output_key='logits')\n",
+ "model = tf.keras.Sequential([keras_layer])\n",
+ "model.build([None, IMAGE_HEIGHT, IMAGE_WIDTH, 3])\n",
+ "```"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "7MOWG8h5eEh6"
+ },
+ "source": [
+ "### Download SavedModel\n",
+ "The model uses the implementation from the TensorFlow Model Garden GitHub repository, and achieves 77.24% mIoU on Cityscapes dataset with 19 classes."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WvTrNdt7sjh4"
+ },
+ "outputs": [],
+ "source": [
+ "! curl https://storage.googleapis.com/tf_model_garden/vision/mosaic/mosaic_mobilenet_multiavgseg_r1024_ebf64_gp_model.tar.gz --output model.tar.gz"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "VEJf6ZTcsREl"
+ },
+ "outputs": [],
+ "source": [
+ "# Extract savedmodel\n",
+ "! tar -xvf model.tar.gz"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "xwD9Rd5EebNC"
+ },
+ "source": [
+ "### Load the SavedModel"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "SvOypbg0dDOa"
+ },
+ "outputs": [],
+ "source": [
+ "# Load saved model\n",
+ "export_dir = \"saved_model\"\n",
+ "imported = tf.saved_model.load(export_dir)\n",
+ "model_fn = imported.signatures['serving_default']"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Th_BzrD5eoPp"
+ },
+ "source": [
+ "## Run Inference"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "X8ACyzMDdDR7"
+ },
+ "outputs": [],
+ "source": [
+ "# Defines helper function to download sample image\n",
+ "def load_image_into_numpy_array(path):\n",
+ " \"\"\"Load an image from file into a numpy array.\n",
+ "\n",
+ " Puts image into numpy array to feed into tensorflow graph.\n",
+ " Note that by convention we put it into a numpy array with shape\n",
+ " (height, width, channels), where channels=3 for RGB.\n",
+ "\n",
+ " Args:\n",
+ " path: the file path to the image\n",
+ "\n",
+ " Returns:\n",
+ " uint8 numpy array with shape (img_height, img_width, 3)\n",
+ " \"\"\"\n",
+ " image = None\n",
+ " if(path.startswith('http')):\n",
+ " response = urlopen(path)\n",
+ " image_data = response.read()\n",
+ " image_data = BytesIO(image_data)\n",
+ " image = Image.open(image_data)\n",
+ " else:\n",
+ " image_data = tf.io.gfile.GFile(path, 'rb').read()\n",
+ " image = Image.open(BytesIO(image_data))\n",
+ "\n",
+ " (im_width, im_height) = image.size\n",
+ "\n",
+ " image = np.array(image.getdata()).reshape(\n",
+ " (1, im_height, im_width, 3)).astype(np.uint8)\n",
+ "\n",
+ " return image"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "n3BXyYune5p6"
+ },
+ "source": [
+ "### Download a Sample Image"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "D7KLJpIGdDVs"
+ },
+ "outputs": [],
+ "source": [
+ "image_path = \"https://storage.googleapis.com/tf_model_garden/vision/mosaic/cityscape_sample.png\"\n",
+ "image_array = load_image_into_numpy_array(image_path)\n",
+ "image_array.shape"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "u2eMxCEfe83m"
+ },
+ "outputs": [],
+ "source": [
+ "Image.fromarray(image_array[0])"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "gn99qmZAf75S"
+ },
+ "source": [
+ "### Run Inference on Sample Image"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "gDdehkLeNeqz"
+ },
+ "outputs": [],
+ "source": [
+ "outputs = model_fn(inputs=image_array)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "LCWNjuFxNvTZ"
+ },
+ "outputs": [],
+ "source": [
+ "outputs['logits'].shape"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "l9WHT3WsgePX"
+ },
+ "source": [
+ "Note that the output is in the shape of `[batch_size, height, width, num_classes]`, which is the raw `logits` prediction output for each pixel."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "wGusC66mvFw9"
+ },
+ "outputs": [],
+ "source": [
+ "detection_results = np.argmax(outputs['logits'][0], axis=-1)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "TpzoZOJ_u1dw"
+ },
+ "outputs": [],
+ "source": [
+ "fig, ax = plt.subplots(figsize=(12, 12))\n",
+ "ax.imshow(detection_results)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "05ZTuuAci3Bu"
+ },
+ "source": [
+ "## Model Training\n",
+ "\n",
+ "Please check out the [TF Model Garden MOSAIC implementation](https://github.com/tensorflow/models/tree/master/official/projects/mosaic) for model training.\n",
+ "\n",
+ "Please check out the [Image Classification Tutorial](https://github.com/tensorflow/models/blob/master/docs/vision/image_classification.ipynb) for fine-tuning models from the TensorFlow Model Garden package."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "kh7Q42fmi6NP"
+ },
+ "outputs": [],
+ "source": []
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "private_outputs": true,
+ "provenance": [
+ {
+ "file_id": "1G4xIYr_wrx-QXzboQLNT2DGAfA4jNYLQ",
+ "timestamp": 1671845366432
+ }
+ ],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/mosaic/qat/configs/experiments/semantic_segmentation/mosaic_mnv35_cityscapes_tfds_qat_tpu.yaml b/official/projects/mosaic/qat/configs/experiments/semantic_segmentation/mosaic_mnv35_cityscapes_tfds_qat_tpu.yaml
new file mode 100644
index 00000000000..2b924d28a03
--- /dev/null
+++ b/official/projects/mosaic/qat/configs/experiments/semantic_segmentation/mosaic_mnv35_cityscapes_tfds_qat_tpu.yaml
@@ -0,0 +1,86 @@
+# Using Tensorflow datasets: 'cityscapes/semantic_segmentation'
+# Some expected flags to use:
+# --experiment=mosaic_mnv35_cityscapes_qat
+# mIoU (unquantized fp32): 77.24
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ num_classes: 19
+ input_size: [null, null, 3]
+ backbone:
+ type: 'mobilenet'
+ mobilenet:
+ model_id: 'MobileNetMultiAVGSeg'
+ output_stride: 16
+ neck:
+ branch_filter_depths: [64, 64]
+ conv_kernel_sizes: [3, 5]
+ pyramid_pool_bin_nums: [1, 4, 8, 16]
+ dropout_rate: 0.0
+ head:
+ num_classes: 19
+ decoder_input_levels: ['3/depthwise', '2/depthwise']
+ decoder_stage_merge_styles: ['concat_merge', 'sum_merge']
+ decoder_filters: [64, 64]
+ decoder_projected_filters: [19, 19]
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ losses:
+ l2_weight_decay: 1.0e-06 # 1/100 of original value.
+ quantization:
+ pretrained_original_checkpoint: 'gs://tf_model_garden/vision/mosaic/mobilenet_multiavgseg_r1024_ebf64_gp/best_ckpt-857'
+ init_checkpoint: null
+ train_data:
+ output_size: [1024, 2048]
+ crop_size: [1024, 2048]
+ input_path: ''
+ tfds_name: 'cityscapes/semantic_segmentation'
+ tfds_split: 'train'
+ is_training: true
+ global_batch_size: 32
+ dtype: 'float32'
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ output_size: [1024, 2048]
+ input_path: ''
+ tfds_name: 'cityscapes/semantic_segmentation'
+ tfds_split: 'validation'
+ is_training: false
+ global_batch_size: 32
+ dtype: 'float32'
+ drop_remainder: false
+ resize_eval_groundtruth: true
+trainer:
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 20000
+ initial_learning_rate: 0.001 # 1/100 of original lr.
+ power: 0.9
+ type: polynomial
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 0 # No warmup
+ type: linear
+ steps_per_loop: 92 # 2975 / 32 = 92
+ summary_interval: 92
+ train_steps: 20000
+ validation_interval: 92
+ validation_steps: 16 # 500 / 32 = 16
+ checkpoint_interval: 92
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/mosaic/qat/configs/mosaic_config.py b/official/projects/mosaic/qat/configs/mosaic_config.py
new file mode 100644
index 00000000000..bcd4ad39800
--- /dev/null
+++ b/official/projects/mosaic/qat/configs/mosaic_config.py
@@ -0,0 +1,38 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Mosaic configuration definition."""
+import dataclasses
+from typing import Optional
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.mosaic.configs import mosaic_config
+from official.projects.qat.vision.configs import common
+
+
+@dataclasses.dataclass
+class MosaicSemanticSegmentationTask(
+ mosaic_config.MosaicSemanticSegmentationTask):
+ quantization: Optional[common.Quantization] = None
+
+
+@exp_factory.register_config_factory('mosaic_mnv35_cityscapes_qat')
+def mosaic_mnv35_cityscapes() -> cfg.ExperimentConfig:
+ """Experiment configuration of image segmentation task with QAT."""
+ config = mosaic_config.mosaic_mnv35_cityscapes()
+ task = MosaicSemanticSegmentationTask.from_args(
+ quantization=common.Quantization(), **config.task.as_dict())
+ config.task = task
+ return config
diff --git a/official/projects/mosaic/qat/configs/mosaic_config_test.py b/official/projects/mosaic/qat/configs/mosaic_config_test.py
new file mode 100644
index 00000000000..0b9c63ae457
--- /dev/null
+++ b/official/projects/mosaic/qat/configs/mosaic_config_test.py
@@ -0,0 +1,45 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for mosaic."""
+# pylint: disable=unused-import
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official import vision
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.mosaic.configs import mosaic_config as exp_cfg
+from official.projects.mosaic.qat.configs import mosaic_config as qat_exp_cfg
+from official.projects.qat.vision.configs import common
+
+
+class MosaicConfigTest(tf.test.TestCase, parameterized.TestCase):
+
+ def test_mosaic_configs(self):
+ config = exp_factory.get_exp_config('mosaic_mnv35_cityscapes_qat')
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task,
+ qat_exp_cfg.MosaicSemanticSegmentationTask)
+ self.assertIsInstance(config.task.model,
+ exp_cfg.MosaicSemanticSegmentationModel)
+ self.assertIsInstance(config.task.quantization, common.Quantization)
+ config.validate()
+ config.task.train_data.is_training = None
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/qat/modeling/factory.py b/official/projects/mosaic/qat/modeling/factory.py
new file mode 100644
index 00000000000..65915208ddf
--- /dev/null
+++ b/official/projects/mosaic/qat/modeling/factory.py
@@ -0,0 +1,97 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Factory methods to build models."""
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.projects.mosaic.modeling import mosaic_blocks
+from official.projects.mosaic.modeling import mosaic_head
+from official.projects.mosaic.modeling import mosaic_model
+from official.projects.mosaic.qat.modeling.heads import mosaic_head as qat_mosaic_head
+from official.projects.mosaic.qat.modeling.layers import nn_blocks as qat_nn_blocks
+from official.projects.qat.vision.configs import common
+from official.projects.qat.vision.modeling.layers import nn_layers as qat_nn_layers
+from official.projects.qat.vision.quantization import helper
+from official.projects.qat.vision.quantization import schemes
+
+
+def build_qat_mosaic_model(
+ model: tf_keras.Model,
+ quantization: common.Quantization,
+ input_specs: tf_keras.layers.InputSpec) -> tf_keras.Model:
+ """Applies quantization aware training for mosaic segmentation model.
+
+ Args:
+ model: The model applying quantization aware training.
+ quantization: The Quantization config.
+ input_specs: The shape specifications of input tensor.
+
+ Returns:
+ The model that applied optimization techniques.
+ """
+
+ original_checkpoint = quantization.pretrained_original_checkpoint
+ if original_checkpoint is not None:
+ ckpt = tf.train.Checkpoint(model=model, **model.checkpoint_items)
+ status = ckpt.read(original_checkpoint)
+ status.expect_partial().assert_existing_objects_matched()
+
+ scope_dict = {
+ 'L2': tf_keras.regularizers.l2,
+ }
+
+ model.use_legacy_config = True # Ensures old Keras serialization format
+ # Apply QAT to backbone (a tf_keras.Model) first, and then neck and head.
+ with tfmot.quantization.keras.quantize_scope(scope_dict):
+ annotated_backbone = tfmot.quantization.keras.quantize_annotate_model(
+ model.backbone)
+ optimized_backbone = tfmot.quantization.keras.quantize_apply(
+ annotated_backbone, scheme=schemes.Default8BitQuantizeScheme())
+
+ # Check for valid encoder and head.
+ if not isinstance(model.head, mosaic_head.MosaicDecoderHead):
+ raise ValueError('Only support MosaicDecoderHead for head.')
+ if not isinstance(model.neck, mosaic_blocks.MosaicEncoderBlock):
+ raise ValueError('Only support MosaicEncoderBlock for encoder.')
+
+ head = qat_mosaic_head.MosaicDecoderHeadQuantized.from_config(
+ model.head.get_config())
+ neck = qat_nn_blocks.MosaicEncoderBlockQuantized.from_config(
+ model.neck.get_config())
+
+ mask_scoring_head = None
+ if model.mask_scoring_head is not None:
+ mask_scoring_head = qat_nn_layers.MaskScoringQuantized.from_config(
+ model.mask_scoring_head.get_config()
+ )
+
+ optimized_model = mosaic_model.MosaicSegmentationModel(
+ backbone=optimized_backbone,
+ head=head,
+ neck=neck,
+ mask_scoring_head=mask_scoring_head,
+ )
+
+ dummpy_input = tf.zeros([1] + list(input_specs.shape[1:]))
+ optimized_model(dummpy_input, training=True)
+ helper.copy_original_weights(model.head, optimized_model.head)
+ helper.copy_original_weights(model.neck, optimized_model.neck)
+
+ if model.mask_scoring_head is not None:
+ helper.copy_original_weights(
+ model.mask_scoring_head, optimized_model.mask_scoring_head
+ )
+
+ return optimized_model
diff --git a/official/projects/mosaic/qat/modeling/factory_test.py b/official/projects/mosaic/qat/modeling/factory_test.py
new file mode 100644
index 00000000000..2c76330bb6b
--- /dev/null
+++ b/official/projects/mosaic/qat/modeling/factory_test.py
@@ -0,0 +1,96 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for factory.py."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.mosaic.modeling import mosaic_blocks
+from official.projects.mosaic.modeling import mosaic_head
+from official.projects.mosaic.modeling import mosaic_model
+from official.projects.mosaic.qat.modeling import factory as qat_factory
+from official.projects.qat.vision.configs import common
+from official.vision.modeling import backbones
+from official.vision.modeling.heads import segmentation_heads
+
+
+class SegmentationModelBuilderTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (128, [4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (128, [1, 4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (128, [1, 4, 8], [3, 2], ['sum_merge', 'sum_merge']),
+ (128, [1, 4, 8], [3, 2], ['concat_merge', 'concat_merge']),
+ (512, [1, 4, 8, 16], [3, 2], ['concat_merge', 'sum_merge']),
+ (256, [4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (256, [1, 4, 8], [3, 2], ['concat_merge', 'sum_merge']),
+ (256, [1, 4, 8, 16], [3, 2], ['concat_merge', 'sum_merge']),
+ )
+ def test_mosaic_segmentation_model(self, input_size, pyramid_pool_bin_nums,
+ decoder_input_levels,
+ decoder_stage_merge_styles):
+ """Test for building and calling of a MOSAIC segmentation network."""
+ num_classes = 32
+ tf_keras.backend.set_image_data_format('channels_last')
+ backbone = backbones.MobileNet(model_id='MobileNetMultiAVGSeg')
+ encoder_input_level = 4
+
+ # Create a regular FP32 MOSAIC model.
+ neck = mosaic_blocks.MosaicEncoderBlock(
+ encoder_input_level=encoder_input_level,
+ branch_filter_depths=[64, 64],
+ conv_kernel_sizes=[3, 5],
+ pyramid_pool_bin_nums=pyramid_pool_bin_nums)
+ head = mosaic_head.MosaicDecoderHead(
+ num_classes=num_classes,
+ decoder_input_levels=decoder_input_levels,
+ decoder_stage_merge_styles=decoder_stage_merge_styles,
+ decoder_filters=[64, 64],
+ decoder_projected_filters=[32, 32])
+ mask_scoring_head = segmentation_heads.MaskScoring(
+ num_classes=num_classes,
+ num_convs=1,
+ num_filters=32,
+ fc_dims=128,
+ num_fcs=2,
+ fc_input_size=[8, 8],
+ )
+
+ model = mosaic_model.MosaicSegmentationModel(
+ backbone=backbone,
+ head=head,
+ neck=neck,
+ mask_scoring_head=mask_scoring_head,
+ )
+
+ inputs = np.random.rand(2, input_size, input_size, 3)
+ input_specs = tf_keras.layers.InputSpec(shape=inputs.shape)
+ expected_outputs = model(inputs)
+
+ # Create a quantized MOSAIC model from the regular FP32 model instance.
+ quantization_config = common.Quantization()
+ quantized_model = qat_factory.build_qat_mosaic_model(
+ model=model,
+ quantization=quantization_config,
+ input_specs=input_specs)
+
+ actual_output = quantized_model(inputs)
+ self.assertAllEqual(actual_output['logits'].numpy().shape,
+ expected_outputs['logits'].numpy().shape)
+
+if __name__ == '__main__':
+ tf.test.main()
+
diff --git a/official/projects/mosaic/qat/modeling/heads/mosaic_head.py b/official/projects/mosaic/qat/modeling/heads/mosaic_head.py
new file mode 100644
index 00000000000..b3ff90ebd4b
--- /dev/null
+++ b/official/projects/mosaic/qat/modeling/heads/mosaic_head.py
@@ -0,0 +1,211 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definitions of segmentation head of the MOSAIC model."""
+from typing import List, Optional
+
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.modeling import tf_utils
+from official.projects.mosaic.modeling import mosaic_head
+from official.projects.mosaic.qat.modeling.layers import nn_blocks
+from official.projects.qat.vision.quantization import configs
+from official.projects.qat.vision.quantization import helper
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MosaicDecoderHeadQuantized(mosaic_head.MosaicDecoderHead):
+ """Creates a quantized MOSAIC decoder in segmentation head.
+
+ Reference:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ Context](https://arxiv.org/pdf/2112.11623.pdf)
+ """
+
+ def __init__(
+ self,
+ num_classes: int,
+ decoder_input_levels: Optional[List[str]] = None,
+ decoder_stage_merge_styles: Optional[List[str]] = None,
+ decoder_filters: Optional[List[int]] = None,
+ decoder_projected_filters: Optional[List[int]] = None,
+ encoder_end_level: Optional[int] = 4,
+ use_additional_classifier_layer: bool = False,
+ classifier_kernel_size: int = 1,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ batchnorm_momentum: float = 0.99,
+ batchnorm_epsilon: float = 0.001,
+ kernel_initializer: str = 'GlorotUniform',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ interpolation: str = 'bilinear',
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initializes a MOSAIC segmentation head.
+
+ Args:
+ num_classes: An `int` number of mask classification categories. The number
+ of classes does not include background class.
+ decoder_input_levels: A list of `str` specifying additional
+ input levels from the backbone outputs for mask refinement in decoder.
+ decoder_stage_merge_styles: A list of `str` specifying the merge style at
+ each stage of the decoder, merge styles can be 'concat_merge' or
+ 'sum_merge'.
+ decoder_filters: A list of integers specifying the number of channels used
+ at each decoder stage. Note: this only has affects if the decoder merge
+ style is 'concat_merge'.
+ decoder_projected_filters: A list of integers specifying the number of
+ projected channels at the end of each decoder stage.
+ encoder_end_level: An optional integer specifying the output level of the
+ encoder stage, which is used if the input from the encoder to the
+ decoder head is a dictionary.
+ use_additional_classifier_layer: A `bool` specifying whether to use an
+ additional classifier layer or not. It must be True if the final decoder
+ projected filters does not match the `num_classes`.
+ classifier_kernel_size: An `int` number to specify the kernel size of the
+ classifier layer.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ batchnorm_momentum: A `float` of normalization momentum for the moving
+ average.
+ batchnorm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_initializer: Kernel initializer for conv layers. Defaults to
+ `glorot_uniform`.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ interpolation: The interpolation method for upsampling. Defaults to
+ `bilinear`.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super().__init__(
+ num_classes=num_classes,
+ decoder_input_levels=decoder_input_levels,
+ decoder_stage_merge_styles=decoder_stage_merge_styles,
+ decoder_filters=decoder_filters,
+ decoder_projected_filters=decoder_projected_filters,
+ encoder_end_level=encoder_end_level,
+ use_additional_classifier_layer=use_additional_classifier_layer,
+ classifier_kernel_size=classifier_kernel_size,
+ activation=activation,
+ use_sync_bn=use_sync_bn,
+ batchnorm_momentum=batchnorm_momentum,
+ batchnorm_epsilon=batchnorm_epsilon,
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ interpolation=interpolation,
+ bias_regularizer=bias_regularizer,
+ **kwargs)
+
+ # Assuming decoder_input_levels and the following lists are sorted and
+ # follow the same order.
+ if decoder_input_levels is None:
+ decoder_input_levels = ['3', '2']
+ if decoder_stage_merge_styles is None:
+ decoder_stage_merge_styles = ['concat_merge', 'sum_merge']
+ if decoder_filters is None:
+ decoder_filters = [64, 64]
+ if decoder_projected_filters is None:
+ decoder_projected_filters = [32, 32]
+ self._decoder_input_levels = decoder_input_levels
+ self._decoder_stage_merge_styles = decoder_stage_merge_styles
+ self._decoder_filters = decoder_filters
+ self._decoder_projected_filters = decoder_projected_filters
+ if (len(decoder_input_levels) != len(decoder_stage_merge_styles) or
+ len(decoder_input_levels) != len(decoder_filters) or
+ len(decoder_input_levels) != len(decoder_projected_filters)):
+ raise ValueError('The number of Decoder inputs and settings must match.')
+ self._merge_stages = []
+ for (stage_merge_style, decoder_filter,
+ decoder_projected_filter) in zip(decoder_stage_merge_styles,
+ decoder_filters,
+ decoder_projected_filters):
+ if stage_merge_style == 'concat_merge':
+ concat_merge_stage = nn_blocks.DecoderConcatMergeBlockQuantized(
+ decoder_internal_depth=decoder_filter,
+ decoder_projected_depth=decoder_projected_filter,
+ output_size=(0, 0),
+ use_sync_bn=use_sync_bn,
+ batchnorm_momentum=batchnorm_momentum,
+ batchnorm_epsilon=batchnorm_epsilon,
+ activation=activation,
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ interpolation=interpolation)
+ self._merge_stages.append(concat_merge_stage)
+ elif stage_merge_style == 'sum_merge':
+ sum_merge_stage = nn_blocks.DecoderSumMergeBlockQuantized(
+ decoder_projected_depth=decoder_projected_filter,
+ output_size=(0, 0),
+ use_sync_bn=use_sync_bn,
+ batchnorm_momentum=batchnorm_momentum,
+ batchnorm_epsilon=batchnorm_epsilon,
+ activation=activation,
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ interpolation=interpolation)
+ self._merge_stages.append(sum_merge_stage)
+ else:
+ raise ValueError(
+ 'A stage merge style in MOSAIC Decoder can only be concat_merge '
+ 'or sum_merge.')
+
+ # Concat merge or sum merge does not require an additional classifer layer
+ # unless the final decoder projected filter does not match num_classes.
+ final_decoder_projected_filter = decoder_projected_filters[-1]
+ if (final_decoder_projected_filter != num_classes and
+ not use_additional_classifier_layer):
+ raise ValueError('Additional classifier layer is needed if final decoder '
+ 'projected filters does not match num_classes!')
+ self._use_additional_classifier_layer = use_additional_classifier_layer
+ if use_additional_classifier_layer:
+ # This additional classification layer uses different kernel
+ # initializers and bias compared to earlier blocks.
+ self._pixelwise_classifier = helper.Conv2DQuantized(
+ name='pixelwise_classifier',
+ filters=num_classes,
+ kernel_size=classifier_kernel_size,
+ padding='same',
+ bias_initializer=tf.zeros_initializer(),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ activation=helper.NoOpActivation(),
+ use_bias=True)
+
+ self._activation_fn = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(activation, use_keras_layer=True),
+ configs.Default8BitActivationQuantizeConfig())
+
+ self._config_dict = {
+ 'num_classes': num_classes,
+ 'decoder_input_levels': decoder_input_levels,
+ 'decoder_stage_merge_styles': decoder_stage_merge_styles,
+ 'decoder_filters': decoder_filters,
+ 'decoder_projected_filters': decoder_projected_filters,
+ 'encoder_end_level': encoder_end_level,
+ 'use_additional_classifier_layer': use_additional_classifier_layer,
+ 'classifier_kernel_size': classifier_kernel_size,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'batchnorm_momentum': batchnorm_momentum,
+ 'batchnorm_epsilon': batchnorm_epsilon,
+ 'kernel_initializer': kernel_initializer,
+ 'kernel_regularizer': kernel_regularizer,
+ 'interpolation': interpolation,
+ 'bias_regularizer': bias_regularizer
+ }
diff --git a/official/projects/mosaic/qat/modeling/heads/mosaic_head_test.py b/official/projects/mosaic/qat/modeling/heads/mosaic_head_test.py
new file mode 100644
index 00000000000..f3255a93796
--- /dev/null
+++ b/official/projects/mosaic/qat/modeling/heads/mosaic_head_test.py
@@ -0,0 +1,62 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for mosaic_head."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.mosaic.qat.modeling.heads import mosaic_head
+
+
+class MosaicBlocksTest(parameterized.TestCase, tf.test.TestCase):
+
+ def test_mosaic_head(self):
+ decoder_head = mosaic_head.MosaicDecoderHeadQuantized(
+ num_classes=32,
+ decoder_input_levels=['3', '2'],
+ decoder_stage_merge_styles=['concat_merge', 'sum_merge'],
+ decoder_filters=[64, 64],
+ decoder_projected_filters=[32, 32])
+ inputs = [
+ tf.ones([1, 32, 32, 128]), {
+ '2': tf.ones([1, 128, 128, 64]),
+ '3': tf.ones([1, 64, 64, 192])
+ }
+ ]
+ outputs = decoder_head(inputs)
+ self.assertAllEqual(outputs.shape, [1, 128, 128, 32])
+
+ def test_mosaic_head_3laterals(self):
+ decoder_head = mosaic_head.MosaicDecoderHeadQuantized(
+ num_classes=32,
+ decoder_input_levels=['3', '2', '1'],
+ decoder_stage_merge_styles=[
+ 'concat_merge', 'concat_merge', 'sum_merge'
+ ],
+ decoder_filters=[64, 64, 64],
+ decoder_projected_filters=[32, 32, 32])
+ inputs = [
+ tf.ones([1, 32, 32, 128]), {
+ '1': tf.ones([1, 256, 256, 64]),
+ '2': tf.ones([1, 128, 128, 64]),
+ '3': tf.ones([1, 64, 64, 192])
+ }
+ ]
+ outputs = decoder_head(inputs)
+ self.assertAllEqual(outputs.shape, [1, 256, 256, 32])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/qat/modeling/layers/nn_blocks.py b/official/projects/mosaic/qat/modeling/layers/nn_blocks.py
new file mode 100644
index 00000000000..17f3ad2df2c
--- /dev/null
+++ b/official/projects/mosaic/qat/modeling/layers/nn_blocks.py
@@ -0,0 +1,448 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains quantized neural blocks for the QAT."""
+
+from typing import Dict, Tuple, Union
+
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.modeling import tf_utils
+from official.projects.mosaic.modeling import mosaic_blocks
+from official.projects.qat.vision.quantization import configs
+from official.projects.qat.vision.quantization import helper
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MultiKernelGroupConvBlockQuantized(mosaic_blocks.MultiKernelGroupConvBlock
+ ):
+ """A quantized multi-kernel grouped convolution block.
+
+ This block is used in the segmentation neck introduced in MOSAIC.
+ Reference:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ Context](https://arxiv.org/pdf/2112.11623.pdf)
+ """
+
+ def build(self, input_shape: tf.TensorShape) -> None:
+ """Builds the block with the given input shape."""
+ input_channels = input_shape[self._group_split_axis]
+ if input_channels % self._num_groups != 0:
+ raise ValueError('The number of input channels must be divisible by '
+ 'the number of groups for evenly group split.')
+
+ # Override the activation and bn with their quantized version.
+ self._activation_fn = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(self._activation, use_keras_layer=True),
+ configs.Default8BitActivationQuantizeConfig())
+ norm_layer = (
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if self._use_sync_bn else tf_keras.layers.BatchNormalization)
+ norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
+ norm_no_quantize = helper.BatchNormalizationNoQuantized(norm_layer)
+ self._bn_op = helper.norm_by_activation(
+ self._activation, norm_with_quantize, norm_no_quantize)
+
+ self._conv_branches = []
+ if self._use_depthwise_convolution:
+ for i, conv_kernel_size in enumerate(self._kernel_sizes):
+ depthwise_conv = helper.DepthwiseConv2DQuantized(
+ kernel_size=(conv_kernel_size, conv_kernel_size),
+ depth_multiplier=1,
+ padding='same',
+ depthwise_regularizer=self._kernel_regularizer,
+ depthwise_initializer=self._kernel_initializer,
+ use_bias=False,
+ activation=helper.NoOpActivation())
+ # Add BN->RELU after depthwise convolution.
+ batchnorm_op_depthwise = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation_depthwise = self._activation_fn
+ feature_conv = helper.Conv2DQuantized(
+ filters=self._output_filter_depths[i],
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=helper.NoOpActivation(),
+ use_bias=False)
+ batchnorm_op = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ # Use list manually as current QAT API does not support sequential model
+ # within a tf_keras.Sequential block, e.g. conv_branch =
+ # tf_keras.Sequential([depthwise_conv, feature_conv, batchnorm_op,])
+ conv_branch = [
+ depthwise_conv,
+ batchnorm_op_depthwise,
+ activation_depthwise,
+ feature_conv,
+ batchnorm_op,
+ ]
+ self._conv_branches.append(conv_branch)
+ else:
+ for i, conv_kernel_size in enumerate(self._kernel_sizes):
+ norm_conv = helper.Conv2DQuantized(
+ filters=self._output_filter_depths[i],
+ kernel_size=(conv_kernel_size, conv_kernel_size),
+ padding='same',
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ activation=helper.NoOpActivation(),
+ use_bias=False)
+ batchnorm_op = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ conv_branch = [norm_conv, batchnorm_op]
+ self._conv_branches.append(conv_branch)
+
+ self._concat_groups = helper.ConcatenateQuantized(
+ axis=self._group_split_axis)
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MosaicEncoderBlockQuantized(mosaic_blocks.MosaicEncoderBlock):
+ """Implements the encoder module/block of MOSAIC model.
+
+ Spatial Pyramid Pooling and Multi-kernel Conv layer
+ SpatialPyramidPoolingMultiKernelConv
+ References:
+ [MOSAIC: Mobile Segmentation via decoding Aggregated Information and encoded
+ context](https://arxiv.org/pdf/2112.11623.pdf)
+ """
+
+ def build(
+ self, input_shape: Union[tf.TensorShape, Dict[str,
+ tf.TensorShape]]) -> None:
+ """Builds this MOSAIC encoder block with the given single input shape."""
+ input_shape = (
+ input_shape[self._encoder_input_level]
+ if isinstance(input_shape, dict) else input_shape)
+ self._data_format = tf_keras.backend.image_data_format()
+ if self._data_format == 'channels_last':
+ height = input_shape[1]
+ width = input_shape[2]
+ else:
+ height = input_shape[2]
+ width = input_shape[3]
+
+ self._global_pool_branch = None
+ self._spatial_pyramid = []
+
+ # Override the activation and bn with their quantized version.
+ self._activation_fn = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(self._activation, use_keras_layer=True),
+ configs.Default8BitActivationQuantizeConfig())
+ norm_layer = (
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if self._use_sync_bn else tf_keras.layers.BatchNormalization)
+ norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
+ norm_no_quantize = helper.BatchNormalizationNoQuantized(norm_layer)
+ self._bn_op = helper.norm_by_activation(
+ self._activation, norm_with_quantize, norm_no_quantize)
+
+ for pyramid_pool_bin_num in self._pyramid_pool_bin_nums:
+ if pyramid_pool_bin_num == 1:
+ global_pool = helper.GlobalAveragePooling2DQuantized(
+ data_format=self._data_format, keepdims=True)
+
+ global_projection = helper.Conv2DQuantized(
+ filters=max(self._branch_filter_depths),
+ kernel_size=(1, 1),
+ padding='same',
+ activation=helper.NoOpActivation(),
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ use_bias=False)
+ batch_norm_global_branch = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ # Use list manually instead of tf_keras.Sequential([])
+ self._global_pool_branch = [
+ global_pool,
+ global_projection,
+ batch_norm_global_branch,
+ ]
+ else:
+ if height < pyramid_pool_bin_num or width < pyramid_pool_bin_num:
+ raise ValueError('The number of pooling bins must be smaller than '
+ 'input sizes.')
+ assert pyramid_pool_bin_num >= 2, (
+ 'Except for the gloabl pooling, the number of bins in pyramid '
+ 'pooling must be at least two.')
+ pool_height, stride_height = self._get_bin_pool_kernel_and_stride(
+ height, pyramid_pool_bin_num)
+ pool_width, stride_width = self._get_bin_pool_kernel_and_stride(
+ width, pyramid_pool_bin_num)
+ bin_pool_level = helper.AveragePooling2DQuantized(
+ pool_size=(pool_height, pool_width),
+ strides=(stride_height, stride_width),
+ padding='valid',
+ data_format=self._data_format)
+ self._spatial_pyramid.append(bin_pool_level)
+
+ # Grouped multi-kernel Convolution.
+ self._multi_kernel_group_conv = MultiKernelGroupConvBlockQuantized(
+ output_filter_depths=self._branch_filter_depths,
+ kernel_sizes=self._conv_kernel_sizes,
+ use_sync_bn=self._use_sync_bn,
+ batchnorm_momentum=self._batchnorm_momentum,
+ batchnorm_epsilon=self._batchnorm_epsilon,
+ activation=self._activation,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ use_depthwise_convolution=self._use_depthwise_convolution)
+
+ # Encoder's final 1x1 feature projection.
+ # Considering the relatively large #channels merged before projection,
+ # enlarge the projection #channels to the sum of the filter depths of
+ # branches.
+ self._output_channels = sum(self._branch_filter_depths)
+ # Use list manually instead of tf_keras.Sequential([]).
+ self._encoder_projection = [
+ helper.Conv2DQuantized(
+ filters=self._output_channels,
+ kernel_size=(1, 1),
+ padding='same',
+ activation=helper.NoOpActivation(),
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ use_bias=False),
+ self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon),
+ ]
+ # Use the TF2 default feature alignment rule for bilinear resizing.
+ self._upsample = helper.ResizingQuantized(
+ height,
+ width,
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ self._concat_layer = helper.ConcatenateQuantized(axis=self._channel_axis)
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class DecoderSumMergeBlockQuantized(mosaic_blocks.DecoderSumMergeBlock):
+ """Implements the decoder feature sum merge block of MOSAIC model.
+
+ This block is used in the decoder of segmentation head introduced in MOSAIC.
+ It essentially merges a high-resolution feature map of a low semantic level
+ and a low-resolution feature map of a higher semantic level by 'Sum-Merge'.
+ """
+
+ def build(
+ self,
+ input_shape: Tuple[tf.TensorShape, tf.TensorShape]) -> None:
+ """Builds the block with the given input shape."""
+ # Assume backbone features of the same level are concated before input.
+ low_res_input_shape = input_shape[0]
+ high_res_input_shape = input_shape[1]
+ low_res_channels = low_res_input_shape[self._channel_axis]
+ high_res_channels = high_res_input_shape[self._channel_axis]
+
+ # Override the activation and bn with their quantized version.
+ self._activation_fn = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(self._activation, use_keras_layer=True),
+ configs.Default8BitActivationQuantizeConfig())
+ norm_layer = (
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if self._use_sync_bn else tf_keras.layers.BatchNormalization)
+ norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
+ norm_no_quantize = helper.BatchNormalizationNoQuantized(norm_layer)
+ self._bn_op = helper.norm_by_activation(
+ self._activation, norm_with_quantize, norm_no_quantize)
+
+ if low_res_channels != self._decoder_projected_depth:
+ low_res_feature_conv = helper.Conv2DQuantized(
+ filters=self._decoder_projected_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=helper.NoOpActivation(),
+ use_bias=False)
+ batchnorm_op = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ self._low_res_branch = [
+ low_res_feature_conv,
+ batchnorm_op,
+ ]
+ if high_res_channels != self._decoder_projected_depth:
+ high_res_feature_conv = helper.Conv2DQuantized(
+ filters=self._decoder_projected_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=helper.NoOpActivation(),
+ use_bias=False)
+ batchnorm_op_high = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ self._high_res_branch = [
+ high_res_feature_conv,
+ batchnorm_op_high,
+ ]
+ # Resize feature maps.
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ low_res_height = low_res_input_shape[1]
+ low_res_width = low_res_input_shape[2]
+ high_res_height = high_res_input_shape[1]
+ high_res_width = high_res_input_shape[2]
+ else:
+ low_res_height = low_res_input_shape[2]
+ low_res_width = low_res_input_shape[3]
+ high_res_height = high_res_input_shape[2]
+ high_res_width = high_res_input_shape[3]
+ if (self._output_size[0] == 0 or self._output_size[1] == 0):
+ self._output_size = (high_res_height, high_res_width)
+ if (low_res_height != self._output_size[0] or
+ low_res_width != self._output_size[1]):
+ self._upsample_low_res = helper.ResizingQuantized(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ if (high_res_height != self._output_size[0] or
+ high_res_width != self._output_size[1]):
+ self._upsample_high_res = helper.ResizingQuantized(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ self._add_layer = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_keras.layers.Add(), configs.Default8BitQuantizeConfig([], [], True))
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class DecoderConcatMergeBlockQuantized(mosaic_blocks.DecoderConcatMergeBlock):
+ """Implements the decoder feature concat merge block of MOSAIC model.
+
+ This block is used in the decoder of segmentation head introduced in MOSAIC.
+ It essentially merges a high-resolution feature map of a low semantic level
+ and a low-resolution feature of a higher semantic level by 'Concat-Merge'.
+ """
+
+ def build(
+ self,
+ input_shape: Tuple[tf.TensorShape, tf.TensorShape]) -> None:
+ """Builds this block with the given input shape."""
+ # Assume backbone features of the same level are concated before input.
+ low_res_input_shape = input_shape[0]
+ high_res_input_shape = input_shape[1]
+ # Set up resizing feature maps before concat.
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ low_res_height = low_res_input_shape[1]
+ low_res_width = low_res_input_shape[2]
+ high_res_height = high_res_input_shape[1]
+ high_res_width = high_res_input_shape[2]
+ else:
+ low_res_height = low_res_input_shape[2]
+ low_res_width = low_res_input_shape[3]
+ high_res_height = high_res_input_shape[2]
+ high_res_width = high_res_input_shape[3]
+
+ self._concat_layer = helper.ConcatenateQuantized(axis=self._channel_axis)
+
+ # Override the activation and bn with their quantized version.
+ self._activation_fn = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(self._activation, use_keras_layer=True),
+ configs.Default8BitActivationQuantizeConfig())
+ norm_layer = (
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if self._use_sync_bn else tf_keras.layers.BatchNormalization)
+ norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
+ norm_no_quantize = helper.BatchNormalizationNoQuantized(norm_layer)
+ self._bn_op = helper.norm_by_activation(
+ self._activation, norm_with_quantize, norm_no_quantize)
+
+ if (self._output_size[0] == 0 or self._output_size[1] == 0):
+ self._output_size = (high_res_height, high_res_width)
+ if (low_res_height != self._output_size[0] or
+ low_res_width != self._output_size[1]):
+ self._upsample_low_res = helper.ResizingQuantized(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ if (high_res_height != self._output_size[0] or
+ high_res_width != self._output_size[1]):
+ self._upsample_high_res = helper.ResizingQuantized(
+ self._output_size[0],
+ self._output_size[1],
+ interpolation=self._interpolation,
+ crop_to_aspect_ratio=False)
+ # Set up a 3-layer separable convolution blocks, i.e.
+ # 1x1->BN->RELU + Depthwise->BN->RELU + 1x1->BN->RELU.
+ initial_feature_conv = helper.Conv2DQuantized(
+ filters=self._decoder_internal_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=helper.NoOpActivation(),
+ use_bias=False)
+ batchnorm_op1 = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation1 = self._activation_fn
+ depthwise_conv = helper.DepthwiseConv2DQuantized(
+ kernel_size=(3, 3),
+ depth_multiplier=1,
+ padding='same',
+ depthwise_regularizer=self._kernel_regularizer,
+ depthwise_initializer=self._kernel_initializer,
+ use_bias=False,
+ activation=helper.NoOpActivation())
+ batchnorm_op2 = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation2 = self._activation_fn
+ project_feature_conv = helper.Conv2DQuantized(
+ filters=self._decoder_projected_depth,
+ kernel_size=(1, 1),
+ padding='same',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=self._kernel_initializer,
+ activation=helper.NoOpActivation(),
+ use_bias=False)
+ batchnorm_op3 = self._bn_op(
+ axis=self._bn_axis,
+ momentum=self._batchnorm_momentum,
+ epsilon=self._batchnorm_epsilon)
+ activation3 = self._activation_fn
+ self._feature_fusion_block = [
+ initial_feature_conv,
+ batchnorm_op1,
+ activation1,
+ depthwise_conv,
+ batchnorm_op2,
+ activation2,
+ project_feature_conv,
+ batchnorm_op3,
+ activation3,
+ ]
+ self._concat_layer = helper.ConcatenateQuantized(axis=self._channel_axis)
diff --git a/official/projects/mosaic/qat/modeling/layers/nn_blocks_test.py b/official/projects/mosaic/qat/modeling/layers/nn_blocks_test.py
new file mode 100644
index 00000000000..6f0b30feb01
--- /dev/null
+++ b/official/projects/mosaic/qat/modeling/layers/nn_blocks_test.py
@@ -0,0 +1,119 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for nn_blocks."""
+
+from typing import Any, Iterable, Tuple
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from tensorflow.python.distribute import combinations
+from tensorflow.python.distribute import strategy_combinations
+from official.projects.mosaic.qat.modeling.layers import nn_blocks
+
+
+def distribution_strategy_combinations() -> Iterable[Tuple[Any, ...]]:
+ """Returns the combinations of end-to-end tests to run."""
+ return combinations.combine(
+ distribution=[
+ strategy_combinations.default_strategy,
+ strategy_combinations.cloud_tpu_strategy,
+ strategy_combinations.one_device_strategy_gpu,
+ ],
+ )
+
+
+class NNBlocksTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (nn_blocks.MultiKernelGroupConvBlockQuantized, [32, 64]),
+ (nn_blocks.MultiKernelGroupConvBlockQuantized, [64, 128]),
+ )
+ def test_multi_kernel_grouped_convolution_block_creation(
+ self, block_fn, output_filter_depths):
+ input_size = 32
+ inputs = tf_keras.Input(shape=(input_size, input_size, 16), batch_size=1)
+ block = block_fn(
+ output_filter_depths=output_filter_depths, kernel_sizes=[3, 3])
+
+ features = block(inputs)
+
+ self.assertAllEqual([1, input_size, input_size,
+ sum(output_filter_depths)], features.shape.as_list())
+
+ @parameterized.parameters(
+ (nn_blocks.MosaicEncoderBlockQuantized, [32, 64], [3, 3], [2, 2]),
+ (nn_blocks.MosaicEncoderBlockQuantized, [64, 128], [3, 1], [2, 4]),
+ (nn_blocks.MosaicEncoderBlockQuantized, [128, 256], [1, 1], [1, 1]),
+ (nn_blocks.MosaicEncoderBlockQuantized, [128, 256], [3, 3], [4, 4]),
+ )
+ def test_mosaic_encoder_block_creation(self, block_fn, branch_filter_depths,
+ conv_kernel_sizes,
+ pyramid_pool_bin_nums):
+ input_size = 128
+ in_filters = 24
+ inputs = tf_keras.Input(
+ shape=(input_size, input_size, in_filters), batch_size=1)
+ block = block_fn(
+ branch_filter_depths=branch_filter_depths,
+ conv_kernel_sizes=conv_kernel_sizes,
+ pyramid_pool_bin_nums=pyramid_pool_bin_nums)
+
+ features = block(inputs)
+
+ self.assertAllEqual([1, input_size, input_size,
+ sum(branch_filter_depths)], features.shape.as_list())
+
+ @parameterized.parameters(
+ (nn_blocks.DecoderSumMergeBlockQuantized, 32, [128, 64]),
+ (nn_blocks.DecoderSumMergeBlockQuantized, 16, [32, 32]),
+ )
+ def test_decoder_sum_merge_block_creation(self, block_fn,
+ decoder_projected_depth,
+ output_size):
+ inputs = (tf_keras.Input(shape=(64, 64, 128), batch_size=1),
+ tf_keras.Input(shape=(16, 16, 256), batch_size=1))
+ block = block_fn(
+ decoder_projected_depth=decoder_projected_depth,
+ output_size=output_size)
+
+ features = block(inputs)
+
+ self.assertAllEqual(
+ [1, output_size[0], output_size[1], decoder_projected_depth],
+ features.shape.as_list())
+
+ @parameterized.parameters(
+ (nn_blocks.DecoderConcatMergeBlockQuantized, 64, 32, [128, 64]),
+ (nn_blocks.DecoderConcatMergeBlockQuantized, 256, 16, [32, 32]),
+ )
+ def test_decoder_concat_merge_block_creation(self, block_fn,
+ decoder_internal_depth,
+ decoder_projected_depth,
+ output_size):
+ inputs = (tf_keras.Input(shape=(64, 64, 128), batch_size=1),
+ tf_keras.Input(shape=(16, 16, 256), batch_size=1))
+ block = block_fn(
+ decoder_internal_depth=decoder_internal_depth,
+ decoder_projected_depth=decoder_projected_depth,
+ output_size=output_size)
+
+ features = block(inputs)
+
+ self.assertAllEqual(
+ [1, output_size[0], output_size[1], decoder_projected_depth],
+ features.shape.as_list())
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/qat/serving/export_module.py b/official/projects/mosaic/qat/serving/export_module.py
new file mode 100644
index 00000000000..0485a99f816
--- /dev/null
+++ b/official/projects/mosaic/qat/serving/export_module.py
@@ -0,0 +1,43 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Export modules for QAT model serving/inference."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.mosaic.modeling import mosaic_model
+from official.projects.mosaic.qat.modeling import factory as qat_factory
+from official.vision.serving import semantic_segmentation
+
+
+class MosaicModule(semantic_segmentation.SegmentationModule):
+ """MOSAIC Module."""
+
+ def _build_model(self) -> tf_keras.Model:
+ input_specs = tf_keras.layers.InputSpec(shape=[1] +
+ self._input_image_size + [3])
+
+ model = mosaic_model.build_mosaic_segmentation_model(
+ input_specs=input_specs,
+ model_config=self.params.task.model,
+ l2_regularizer=None)
+
+ dummy_input = tf.ones(shape=input_specs.shape)
+ model(dummy_input)
+ # Check whether "quantization" is in task config to support both
+ # `quantized` and `non-quantized` version of Mosaic.
+ if hasattr(self.params.task, "quantization"):
+ return qat_factory.build_qat_mosaic_model(
+ model, self.params.task.quantization, input_specs)
+ return model
diff --git a/official/projects/mosaic/qat/serving/export_saved_model.py b/official/projects/mosaic/qat/serving/export_saved_model.py
new file mode 100644
index 00000000000..e905f5b7d34
--- /dev/null
+++ b/official/projects/mosaic/qat/serving/export_saved_model.py
@@ -0,0 +1,133 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+r"""Vision models export binary for serving/inference.
+
+To export a trained checkpoint in saved_model format (shell script):
+
+EXPERIMENT_TYPE = XX
+CHECKPOINT_PATH = XX
+EXPORT_DIR_PATH = XX
+export_saved_model --experiment=${EXPERIMENT_TYPE} \
+ --export_dir=${EXPORT_DIR_PATH}/ \
+ --checkpoint_path=${CHECKPOINT_PATH} \
+ --batch_size=2 \
+ --input_image_size=224,224
+
+To serve (python):
+
+export_dir_path = XX
+input_type = XX
+input_images = XX
+imported = tf.saved_model.load(export_dir_path)
+model_fn = imported.signatures['serving_default']
+output = model_fn(input_images)
+"""
+from absl import app
+from absl import flags
+
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.mosaic import registry_imports # pylint: disable=unused-import
+from official.projects.mosaic.configs import mosaic_config
+from official.projects.mosaic.qat.serving import export_module
+from official.vision.serving import export_saved_model_lib
+
+
+FLAGS = flags.FLAGS
+
+_EXPERIMENT = flags.DEFINE_string(
+ 'experiment', None, 'experiment type, e.g. retinanet_resnetfpn_coco')
+_EXPORT_DIR = flags.DEFINE_string('export_dir', None, 'The export directory.')
+_CHECKPOINT_PATH = flags.DEFINE_string('checkpoint_path', None,
+ 'Checkpoint path.')
+_CONFIG_FILE = flags.DEFINE_multi_string(
+ 'config_file',
+ default=None,
+ help='YAML/JSON files which specifies overrides. The override order '
+ 'follows the order of args. Note that each file '
+ 'can be used as an override template to override the default parameters '
+ 'specified in Python. If the same parameter is specified in both '
+ '`--config_file` and `--params_override`, `config_file` will be used '
+ 'first, followed by params_override.')
+_PARAMS_OVERRIDE = flags.DEFINE_string(
+ 'params_override', '',
+ 'The JSON/YAML file or string which specifies the parameter to be overriden'
+ ' on top of `config_file` template.')
+_BATCH_SIZE = flags.DEFINE_integer('batch_size', None, 'The batch size.')
+_IMAGE_TYPE = flags.DEFINE_string(
+ 'input_type', 'image_tensor',
+ 'One of `image_tensor`, `image_bytes`, `tf_example` and `tflite`.')
+_INPUT_IMAGE_SIZE = flags.DEFINE_string(
+ 'input_image_size', '224,224',
+ 'The comma-separated string of two integers representing the height,width '
+ 'of the input to the model.')
+_EXPORT_CHECKPOINT_SUBDIR = flags.DEFINE_string(
+ 'export_checkpoint_subdir', 'checkpoint',
+ 'The subdirectory for checkpoints.')
+_EXPORT_SAVED_MODEL_SUBDIR = flags.DEFINE_string(
+ 'export_saved_model_subdir', 'saved_model',
+ 'The subdirectory for saved model.')
+_LOG_MODEL_FLOPS_AND_PARAMS = flags.DEFINE_bool(
+ 'log_model_flops_and_params', False,
+ 'If true, logs model flops and parameters.')
+_INPUT_NAME = flags.DEFINE_string(
+ 'input_name', None,
+ 'Input tensor name in signature def. Default at None which'
+ 'produces input tensor name `inputs`.')
+
+
+def main(_):
+
+ params = exp_factory.get_exp_config(_EXPERIMENT.value)
+ for config_file in _CONFIG_FILE.value or []:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=True)
+ if _PARAMS_OVERRIDE.value:
+ params = hyperparams.override_params_dict(
+ params, _PARAMS_OVERRIDE.value, is_strict=True)
+
+ params.validate()
+ params.lock()
+
+ input_image_size = [int(x) for x in _INPUT_IMAGE_SIZE.value.split(',')]
+
+ if isinstance(params.task, mosaic_config.MosaicSemanticSegmentationTask):
+ export_module_cls = export_module.MosaicModule
+ else:
+ raise TypeError(f'Export module for {type(params.task)} is not supported.')
+
+ module = export_module_cls(
+ params=params,
+ batch_size=_BATCH_SIZE.value,
+ input_image_size=input_image_size,
+ input_type=_IMAGE_TYPE.value,
+ num_channels=3)
+
+ export_saved_model_lib.export_inference_graph(
+ input_type=_IMAGE_TYPE.value,
+ batch_size=_BATCH_SIZE.value,
+ input_image_size=input_image_size,
+ params=params,
+ checkpoint_path=_CHECKPOINT_PATH.value,
+ export_dir=_EXPORT_DIR.value,
+ export_checkpoint_subdir=_EXPORT_CHECKPOINT_SUBDIR.value,
+ export_saved_model_subdir=_EXPORT_SAVED_MODEL_SUBDIR.value,
+ export_module=module,
+ log_model_flops_and_params=_LOG_MODEL_FLOPS_AND_PARAMS.value,
+ input_name=_INPUT_NAME.value)
+
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/mosaic/qat/serving/export_tflite.py b/official/projects/mosaic/qat/serving/export_tflite.py
new file mode 100644
index 00000000000..83e5ca437a2
--- /dev/null
+++ b/official/projects/mosaic/qat/serving/export_tflite.py
@@ -0,0 +1,24 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Binary to convert a saved model to TFLite model for the QAT model."""
+
+from absl import app
+
+from official.projects.mosaic import registry_imports # pylint: disable=unused-import
+from official.vision.serving import export_tflite
+
+
+if __name__ == '__main__':
+ app.run(export_tflite.main)
diff --git a/official/projects/mosaic/qat/tasks/mosaic_tasks.py b/official/projects/mosaic/qat/tasks/mosaic_tasks.py
new file mode 100644
index 00000000000..da08855bb1d
--- /dev/null
+++ b/official/projects/mosaic/qat/tasks/mosaic_tasks.py
@@ -0,0 +1,43 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Semantic segmentation task definition."""
+import tensorflow as tf, tf_keras
+
+from official.core import task_factory
+from official.projects.mosaic import mosaic_tasks
+from official.projects.mosaic.qat.configs import mosaic_config as exp_cfg
+from official.projects.mosaic.qat.modeling import factory
+
+
+@task_factory.register_task_cls(exp_cfg.MosaicSemanticSegmentationTask)
+class MosaicSemanticSegmentationTask(mosaic_tasks.MosaicSemanticSegmentationTask
+ ):
+ """A task for semantic segmentation with QAT."""
+
+ def build_model(self, training=True) -> tf_keras.Model:
+ """Builds semantic segmentation model with QAT."""
+ model = super().build_model(training)
+ if training:
+ input_size = self.task_config.train_data.output_size
+ crop_size = self.task_config.train_data.crop_size
+ if crop_size:
+ input_size = crop_size
+ else:
+ input_size = self.task_config.validation_data.output_size
+ input_specs = tf_keras.layers.InputSpec(shape=[None] + input_size + [3])
+ if self.task_config.quantization:
+ model = factory.build_qat_mosaic_model(
+ model, self.task_config.quantization, input_specs)
+ return model
diff --git a/official/projects/mosaic/qat/tasks/mosaic_tasks_test.py b/official/projects/mosaic/qat/tasks/mosaic_tasks_test.py
new file mode 100644
index 00000000000..7e1558f2f14
--- /dev/null
+++ b/official/projects/mosaic/qat/tasks/mosaic_tasks_test.py
@@ -0,0 +1,90 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for mosaic task."""
+# pylint: disable=unused-import
+import os
+
+from absl.testing import parameterized
+import orbit
+import tensorflow as tf, tf_keras
+
+from official import vision
+from official.core import exp_factory
+from official.modeling import optimization
+from official.projects.mosaic.configs import mosaic_config as exp_cfg
+from official.projects.mosaic.qat.tasks import mosaic_tasks
+from official.vision.dataloaders import tfexample_utils
+
+
+class MosaicSemanticSegmentationTask(parameterized.TestCase, tf.test.TestCase):
+
+ def _create_test_tfrecord(self, tfrecord_file, example, num_samples):
+ examples = [example] * num_samples
+ tfexample_utils.dump_to_tfrecord(
+ record_file=tfrecord_file, tf_examples=examples)
+
+ @parameterized.parameters(
+ ('mosaic_mnv35_cityscapes_qat', True),
+ ('mosaic_mnv35_cityscapes_qat', False),
+ )
+ def test_semantic_segmentation_task(self, test_config, is_training):
+ """Semantic segmentation task test for training and val using toy configs."""
+ input_image_size = [1024, 2048]
+ test_tfrecord_file = os.path.join(self.get_temp_dir(), 'seg_test.tfrecord')
+ example = tfexample_utils.create_segmentation_test_example(
+ image_height=input_image_size[0],
+ image_width=input_image_size[1],
+ image_channel=3)
+ self._create_test_tfrecord(
+ tfrecord_file=test_tfrecord_file, example=example, num_samples=2)
+ config = exp_factory.get_exp_config(test_config)
+ # modify config to suit local testing
+ config.task.model.input_size = [None, None, 3]
+ config.trainer.steps_per_loop = 1
+ config.task.train_data.global_batch_size = 1
+ config.task.validation_data.global_batch_size = 1
+ config.task.train_data.output_size = [1024, 2048]
+ config.task.validation_data.output_size = [1024, 2048]
+ config.task.train_data.crop_size = [512, 512]
+ config.task.train_data.shuffle_buffer_size = 2
+ config.task.validation_data.shuffle_buffer_size = 2
+ config.task.validation_data.input_path = test_tfrecord_file
+ config.task.train_data.input_path = test_tfrecord_file
+ config.train_steps = 1
+ config.task.model.num_classes = 256
+ config.task.model.head.num_classes = 256
+ config.task.model.head.decoder_projected_filters = [256, 256]
+
+ task = mosaic_tasks.MosaicSemanticSegmentationTask(config.task)
+ model = task.build_model(is_training)
+ metrics = task.build_metrics(training=is_training)
+
+ strategy = tf.distribute.get_strategy()
+
+ data_config = config.task.train_data if is_training else config.task.validation_data
+ dataset = orbit.utils.make_distributed_dataset(strategy, task.build_inputs,
+ data_config)
+ iterator = iter(dataset)
+ opt_factory = optimization.OptimizerFactory(config.trainer.optimizer_config)
+ optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
+
+ if is_training:
+ task.train_step(next(iterator), model, optimizer, metrics=metrics)
+ else:
+ task.validation_step(next(iterator), model, metrics=metrics)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/mosaic/registry_imports.py b/official/projects/mosaic/registry_imports.py
new file mode 100644
index 00000000000..9d5281587ef
--- /dev/null
+++ b/official/projects/mosaic/registry_imports.py
@@ -0,0 +1,21 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""All necessary imports for registration on MOSAIC project."""
+# pylint: disable=unused-import
+from official.projects.mosaic import mosaic_tasks
+from official.projects.mosaic.configs import mosaic_config
+from official.projects.mosaic.modeling import mosaic_model
+from official.projects.mosaic.qat.configs import mosaic_config as mosaic_qat_config
+from official.projects.mosaic.qat.tasks import mosaic_tasks as mosaic_qat_tasks
diff --git a/official/projects/mosaic/train.py b/official/projects/mosaic/train.py
new file mode 100644
index 00000000000..8f135892884
--- /dev/null
+++ b/official/projects/mosaic/train.py
@@ -0,0 +1,109 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Training driver for MOSAIC models."""
+
+from absl import app
+from absl import flags
+import gin
+
+from official.common import distribute_utils
+from official.common import flags as tfm_flags
+from official.core import base_trainer
+from official.core import config_definitions
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+
+# Import MOSAIC libraries to register the model into tf.vision
+# model garden factory.
+# pylint: disable=unused-import
+from official.projects.mosaic import mosaic_tasks
+from official.projects.mosaic import registry_imports as mosaic_registry_imports
+from official.vision import registry_imports
+from official.vision.utils import summary_manager
+# pylint: enable=unused-import
+
+FLAGS = flags.FLAGS
+
+
+# Note: we overrided the `build_trainer` due to the customized `build_model`
+# methods in `MosaicSemanticSegmentationTask.
+def _build_mosaic_trainer(params: config_definitions.ExperimentConfig,
+ task: mosaic_tasks.MosaicSemanticSegmentationTask,
+ model_dir: str, train: bool,
+ evaluate: bool) -> base_trainer.Trainer:
+ """Creates custom trainer."""
+ checkpoint_exporter = train_lib.maybe_create_best_ckpt_exporter(
+ params, model_dir)
+ model = task.build_model(train)
+ optimizer = train_utils.create_optimizer(task, params)
+ trainer = base_trainer.Trainer(
+ params,
+ task,
+ model=model,
+ optimizer=optimizer,
+ train=train,
+ evaluate=evaluate,
+ checkpoint_exporter=checkpoint_exporter)
+ return trainer
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ # Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
+ # dtype is float16
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu)
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+ mosaic_trainer = _build_mosaic_trainer(
+ task=task,
+ params=params,
+ model_dir=model_dir,
+ train='train' in FLAGS.mode,
+ evaluate='eval' in FLAGS.mode)
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir,
+ trainer=mosaic_trainer,
+ eval_summary_manager=summary_manager.maybe_build_eval_summary_manager(
+ params=params, model_dir=model_dir
+ ),
+ )
+
+ train_utils.save_gin_config(FLAGS.mode, model_dir)
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ flags.mark_flags_as_required(['experiment', 'mode', 'model_dir'])
+ app.run(main)
diff --git a/official/projects/movinet/__init__.py b/official/projects/movinet/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/movinet/__init__.py
+++ b/official/projects/movinet/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/movinet/configs/__init__.py b/official/projects/movinet/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/movinet/configs/__init__.py
+++ b/official/projects/movinet/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/movinet/configs/movinet.py b/official/projects/movinet/configs/movinet.py
index 49388583fed..38b721b63c1 100644
--- a/official/projects/movinet/configs/movinet.py
+++ b/official/projects/movinet/configs/movinet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -53,6 +53,8 @@ class Movinet(hyperparams.Config):
gating_activation: str = 'sigmoid'
stochastic_depth_drop_rate: float = 0.2
use_external_states: bool = False
+ average_pooling_type: str = '3d'
+ output_states: bool = True
@dataclasses.dataclass
@@ -117,19 +119,22 @@ class Backbone3D(backbones_3d.Backbone3D):
movinet: movinet backbone config.
"""
type: str = 'movinet'
- movinet: Movinet = Movinet()
+ movinet: Movinet = dataclasses.field(default_factory=Movinet)
@dataclasses.dataclass
class MovinetModel(video_classification.VideoClassificationModel):
"""The MoViNet model config."""
model_type: str = 'movinet'
- backbone: Backbone3D = Backbone3D()
- norm_activation: common.NormActivation = common.NormActivation(
- activation=None, # legacy flag, not used.
- norm_momentum=0.99,
- norm_epsilon=1e-3,
- use_sync_bn=True)
+ backbone: Backbone3D = dataclasses.field(default_factory=Backbone3D)
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ activation=None, # legacy flag, not used.
+ norm_momentum=0.99,
+ norm_epsilon=1e-3,
+ use_sync_bn=True,
+ )
+ )
activation: str = 'swish'
output_states: bool = False
diff --git a/official/projects/movinet/configs/movinet_test.py b/official/projects/movinet/configs/movinet_test.py
index 6efd069c97f..fb1b6a6ad64 100644
--- a/official/projects/movinet/configs/movinet_test.py
+++ b/official/projects/movinet/configs/movinet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for movinet video classification."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import exp_factory
diff --git a/official/projects/movinet/configs/yaml/movinet_a0_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a0_gpu.yaml
new file mode 100644
index 00000000000..e0e39df6a29
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a0_gpu.yaml
@@ -0,0 +1,67 @@
+# Video classification using MoViNet-A0 backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a0'
+ stochastic_depth_drop_rate: 0.2
+ causal: false
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.2
+ activation: 'swish'
+ train_data:
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 192
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 192
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a0_stream_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a0_stream_gpu.yaml
new file mode 100644
index 00000000000..e01dd9f0de2
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a0_stream_gpu.yaml
@@ -0,0 +1,76 @@
+# Video classification using MoViNet-A0-Stream backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a0'
+ causal: true
+ # Note: we train with '3d_2plus1d', but convert to '2plus1d' for inference
+ conv_type: '2plus1d'
+ se_type: '2plus3d'
+ activation: 'hard_swish'
+ gating_activation: 'hard_sigmoid'
+ use_positional_encoding: false
+ stochastic_depth_drop_rate: 0.2
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.2
+ activation: 'hard_swish'
+ train_data:
+ name: kinetics600
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ random_stride_range: 0
+ dtype: 'bfloat16'
+ min_image_size: 192
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ name: kinetics600
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 192
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ train_steps: 85785
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a1_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a1_gpu.yaml
new file mode 100644
index 00000000000..15ce740d2a8
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a1_gpu.yaml
@@ -0,0 +1,67 @@
+# Video classification using MoViNet-A1 backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a1'
+ stochastic_depth_drop_rate: 0.2
+ causal: false
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'swish'
+ train_data:
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 192
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 192
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a1_stream_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a1_stream_gpu.yaml
new file mode 100644
index 00000000000..57a28db71e2
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a1_stream_gpu.yaml
@@ -0,0 +1,76 @@
+# Video classification using MoViNet-A1-Stream backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a1'
+ causal: true
+ # Note: we train with '3d_2plus1d', but convert to '2plus1d' for inference
+ conv_type: '2plus1d'
+ se_type: '2plus3d'
+ activation: 'hard_swish'
+ gating_activation: 'hard_sigmoid'
+ use_positional_encoding: false
+ stochastic_depth_drop_rate: 0.2
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.2
+ activation: 'hard_swish'
+ train_data:
+ name: kinetics600
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ random_stride_range: 0
+ dtype: 'bfloat16'
+ min_image_size: 192
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ name: kinetics600
+ feature_shape: !!python/tuple
+ - 32
+ - 172
+ - 172
+ - 3
+ temporal_stride: 5
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 192
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ train_steps: 85785
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a2_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a2_gpu.yaml
new file mode 100644
index 00000000000..60c1ec50ff6
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a2_gpu.yaml
@@ -0,0 +1,67 @@
+# Video classification using MoViNet-A2 backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a2'
+ stochastic_depth_drop_rate: 0.2
+ causal: false
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'swish'
+ train_data:
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 224
+ - 224
+ - 3
+ temporal_stride: 5
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 256
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ feature_shape: !!python/tuple
+ - 32
+ - 224
+ - 224
+ - 3
+ temporal_stride: 5
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 256
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a2_stream_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a2_stream_gpu.yaml
new file mode 100644
index 00000000000..3849e7fbbb3
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a2_stream_gpu.yaml
@@ -0,0 +1,76 @@
+# Video classification using MoViNet-A2-Stream backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a2'
+ causal: true
+ # Note: we train with '3d_2plus1d', but convert to '2plus1d' for inference
+ conv_type: '2plus1d'
+ se_type: '2plus3d'
+ activation: 'hard_swish'
+ gating_activation: 'hard_sigmoid'
+ use_positional_encoding: false
+ stochastic_depth_drop_rate: 0.2
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'hard_swish'
+ train_data:
+ name: kinetics600
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 224
+ - 224
+ - 3
+ temporal_stride: 5
+ random_stride_range: 0
+ dtype: 'bfloat16'
+ min_image_size: 256
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ name: kinetics600
+ feature_shape: !!python/tuple
+ - 32
+ - 224
+ - 224
+ - 3
+ temporal_stride: 5
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 256
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ train_steps: 85785
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a3_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a3_gpu.yaml
new file mode 100644
index 00000000000..cc6fb324a17
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a3_gpu.yaml
@@ -0,0 +1,67 @@
+# Video classification using MoViNet-A3 backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a3'
+ stochastic_depth_drop_rate: 0.2
+ causal: false
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'swish'
+ train_data:
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 256
+ - 256
+ - 3
+ temporal_stride: 2
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 288
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ feature_shape: !!python/tuple
+ - 32
+ - 256
+ - 256
+ - 3
+ temporal_stride: 2
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 288
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a3_stream_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a3_stream_gpu.yaml
new file mode 100644
index 00000000000..0a160c805f1
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a3_stream_gpu.yaml
@@ -0,0 +1,76 @@
+# Video classification using MoViNet-A3-Stream backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a3'
+ causal: true
+ # Note: we train with '3d_2plus1d', but convert to '2plus1d' for inference
+ conv_type: '2plus1d'
+ se_type: '2plus3d'
+ activation: 'hard_swish'
+ gating_activation: 'hard_sigmoid'
+ use_positional_encoding: true
+ stochastic_depth_drop_rate: 0.2
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'hard_swish'
+ train_data:
+ name: kinetics600
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 256
+ - 256
+ - 3
+ temporal_stride: 2
+ random_stride_range: 0
+ dtype: 'bfloat16'
+ min_image_size: 288
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ name: kinetics600
+ feature_shape: !!python/tuple
+ - 32
+ - 256
+ - 256
+ - 3
+ temporal_stride: 2
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 288
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ train_steps: 85785
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a4_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a4_gpu.yaml
new file mode 100644
index 00000000000..ca88d5bd009
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a4_gpu.yaml
@@ -0,0 +1,67 @@
+# Video classification using MoViNet-A4 backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a4'
+ stochastic_depth_drop_rate: 0.2
+ causal: false
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'swish'
+ train_data:
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 290
+ - 290
+ - 3
+ temporal_stride: 3
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 320
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ feature_shape: !!python/tuple
+ - 32
+ - 290
+ - 290
+ - 3
+ temporal_stride: 3
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 320
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a4_stream_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a4_stream_gpu.yaml
new file mode 100644
index 00000000000..804ec3dc83a
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a4_stream_gpu.yaml
@@ -0,0 +1,76 @@
+# Video classification using MoViNet-A4-Stream backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a4'
+ causal: true
+ # Note: we train with '3d_2plus1d', but convert to '2plus1d' for inference
+ conv_type: '2plus1d'
+ se_type: '2plus3d'
+ activation: 'hard_swish'
+ gating_activation: 'hard_sigmoid'
+ use_positional_encoding: true
+ stochastic_depth_drop_rate: 0.2
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'hard_swish'
+ train_data:
+ name: kinetics600
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 290
+ - 290
+ - 3
+ temporal_stride: 3
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 320
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ name: kinetics600
+ feature_shape: !!python/tuple
+ - 32
+ - 290
+ - 290
+ - 3
+ temporal_stride: 3
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 320
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ train_steps: 85785
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a5_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a5_gpu.yaml
new file mode 100644
index 00000000000..757be487179
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a5_gpu.yaml
@@ -0,0 +1,67 @@
+# Video classification using MoViNet-A5 backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a5'
+ stochastic_depth_drop_rate: 0.2
+ causal: false
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'swish'
+ train_data:
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 320
+ - 320
+ - 3
+ temporal_stride: 2
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 368
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ feature_shape: !!python/tuple
+ - 32
+ - 320
+ - 320
+ - 3
+ temporal_stride: 2
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 368
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/configs/yaml/movinet_a5_stream_gpu.yaml b/official/projects/movinet/configs/yaml/movinet_a5_stream_gpu.yaml
new file mode 100644
index 00000000000..80d8d8073f1
--- /dev/null
+++ b/official/projects/movinet/configs/yaml/movinet_a5_stream_gpu.yaml
@@ -0,0 +1,76 @@
+# Video classification using MoViNet-A5-Stream backbone on multiple GPUs.
+# This configuration is incomplete - Some parameters will be set by model garden Docker.
+# --experiment_type=movinet_kinetics600
+
+runtime:
+ distribution_strategy: 'multi_worker_mirrored'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 0.00003
+ label_smoothing: 0.1
+ model:
+ backbone:
+ movinet:
+ model_id: 'a5'
+ causal: true
+ # Note: we train with '3d_2plus1d', but convert to '2plus1d' for inference
+ conv_type: '2plus1d'
+ se_type: '2plus3d'
+ activation: 'hard_swish'
+ gating_activation: 'hard_sigmoid'
+ use_positional_encoding: true
+ stochastic_depth_drop_rate: 0.2
+ norm_activation:
+ use_sync_bn: true
+ dropout_rate: 0.5
+ activation: 'hard_swish'
+ train_data:
+ name: kinetics600
+ variant_name: rgb
+ feature_shape: !!python/tuple
+ - 32
+ - 320
+ - 320
+ - 3
+ temporal_stride: 2
+ random_stride_range: 1
+ dtype: 'bfloat16'
+ min_image_size: 368
+ aug_max_area_ratio: 1.0
+ aug_max_aspect_ratio: 2.0
+ aug_min_area_ratio: 0.08
+ aug_min_aspect_ratio: 0.5
+ validation_data:
+ name: kinetics600
+ feature_shape: !!python/tuple
+ - 32
+ - 320
+ - 320
+ - 3
+ temporal_stride: 2
+ num_test_clips: 1
+ num_test_crops: 1
+ min_image_size: 368
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1.8
+ decay_steps: 85785
+ warmup:
+ linear:
+ warmup_steps: 300
+ optimizer:
+ type: 'rmsprop'
+ rmsprop:
+ rho: 0.9
+ momentum: 0.9
+ epsilon: 1.0
+ clipnorm: 1.0
+ train_steps: 85785
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/movinet/modeling/__init__.py b/official/projects/movinet/modeling/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/movinet/modeling/__init__.py
+++ b/official/projects/movinet/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/movinet/modeling/movinet.py b/official/projects/movinet/modeling/movinet.py
index 26487d11da9..97433f28d14 100644
--- a/official/projects/movinet/modeling/movinet.py
+++ b/official/projects/movinet/modeling/movinet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,8 @@
import math
from typing import Dict, Mapping, Optional, Sequence, Tuple, Union
-import tensorflow as tf
+from absl import logging
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.movinet.modeling import movinet_layers
@@ -49,7 +50,6 @@
@dataclasses.dataclass
class BlockSpec:
"""Configuration of a block."""
- pass
@dataclasses.dataclass
@@ -297,8 +297,8 @@ class HeadSpec(BlockSpec):
}
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class Movinet(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Movinet(tf_keras.Model):
"""Class to build Movinet family model.
Reference: https://arxiv.org/pdf/2103.11511.pdf
@@ -310,7 +310,7 @@ def __init__(self,
use_positional_encoding: bool = False,
conv_type: str = '3d',
se_type: str = '3d',
- input_specs: Optional[tf.keras.layers.InputSpec] = None,
+ input_specs: Optional[tf_keras.layers.InputSpec] = None,
activation: str = 'swish',
gating_activation: str = 'sigmoid',
use_sync_bn: bool = True,
@@ -322,6 +322,7 @@ def __init__(self,
stochastic_depth_drop_rate: float = 0.,
use_external_states: bool = False,
output_states: bool = True,
+ average_pooling_type: str = '3d',
**kwargs):
"""MoViNet initialization function.
@@ -349,9 +350,9 @@ def __init__(self,
norm_epsilon: small float added to variance to avoid dividing by
zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Defaults to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Defaults to None.
stochastic_depth_drop_rate: the base rate for stochastic depth.
use_external_states: if True, expects states to be passed as additional
@@ -360,11 +361,13 @@ def __init__(self,
the model in streaming mode. Inputting the output states of the
previous input clip with the current input clip will utilize a stream
buffer for streaming video.
+ average_pooling_type: The average pooling type. Currently supporting
+ ['3d', '2d', 'none'].
**kwargs: keyword arguments to be passed.
"""
block_specs = BLOCK_SPECS[model_id]
if input_specs is None:
- input_specs = tf.keras.layers.InputSpec(shape=[None, None, None, None, 3])
+ input_specs = tf_keras.layers.InputSpec(shape=[None, None, None, None, 3])
if conv_type not in ('3d', '2plus1d', '3d_2plus1d'):
raise ValueError('Unknown conv type: {}'.format(conv_type))
@@ -383,16 +386,14 @@ def __init__(self,
self._gating_activation = gating_activation
self._norm_momentum = norm_momentum
self._norm_epsilon = norm_epsilon
- if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
- else:
- self._norm = tf.keras.layers.BatchNormalization
+ self._norm = tf_keras.layers.BatchNormalization
self._kernel_initializer = kernel_initializer
self._kernel_regularizer = kernel_regularizer
self._bias_regularizer = bias_regularizer
self._stochastic_depth_drop_rate = stochastic_depth_drop_rate
self._use_external_states = use_external_states
self._output_states = output_states
+ self._average_pooling_type = average_pooling_type
if self._use_external_states and not self._causal:
raise ValueError('External states should be used with causal mode.')
@@ -417,8 +418,8 @@ def __init__(self,
def _build_network(
self,
- input_specs: tf.keras.layers.InputSpec,
- state_specs: Optional[Mapping[str, tf.keras.layers.InputSpec]] = None,
+ input_specs: tf_keras.layers.InputSpec,
+ state_specs: Optional[Mapping[str, tf_keras.layers.InputSpec]] = None,
) -> Tuple[TensorMap, Union[TensorMap, Tuple[TensorMap, TensorMap]]]:
"""Builds the model network.
@@ -434,10 +435,10 @@ def _build_network(
"""
state_specs = state_specs if state_specs is not None else {}
- image_input = tf.keras.Input(shape=input_specs.shape[1:], name='inputs')
+ image_input = tf_keras.Input(shape=input_specs.shape[1:], name='inputs')
states = {
- name: tf.keras.Input(shape=spec.shape[1:], dtype=spec.dtype, name=name)
+ name: tf_keras.Input(shape=spec.shape[1:], dtype=spec.dtype, name=name)
for name, spec in state_specs.items()
}
@@ -465,6 +466,7 @@ def _build_network(
batch_norm_layer=self._norm,
batch_norm_momentum=self._norm_momentum,
batch_norm_epsilon=self._norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
state_prefix='state_stem',
name='stem')
x, states = layer_obj(x, states=states)
@@ -504,6 +506,7 @@ def _build_network(
batch_norm_layer=self._norm,
batch_norm_momentum=self._norm_momentum,
batch_norm_epsilon=self._norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
state_prefix=f'state_{name}',
name=name)
x, states = layer_obj(x, states=states)
@@ -520,6 +523,8 @@ def _build_network(
batch_norm_layer=self._norm,
batch_norm_momentum=self._norm_momentum,
batch_norm_epsilon=self._norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
+ average_pooling_type=self._average_pooling_type,
state_prefix='state_head',
name='head')
x, states = layer_obj(x, states=states)
@@ -635,7 +640,7 @@ def _get_state_dtype(self, name: str) -> str:
return self.dtype
def initial_state_specs(
- self, input_shape: Sequence[int]) -> Dict[str, tf.keras.layers.InputSpec]:
+ self, input_shape: Sequence[int]) -> Dict[str, tf_keras.layers.InputSpec]:
"""Creates a mapping of state name to InputSpec from the input shape."""
state_shapes = self._get_initial_state_shapes(
self._block_specs,
@@ -643,7 +648,7 @@ def initial_state_specs(
use_positional_encoding=self._use_positional_encoding)
return {
- name: tf.keras.layers.InputSpec(
+ name: tf_keras.layers.InputSpec(
shape=shape, dtype=self._get_state_dtype(name))
for name, shape in state_shapes.items()
}
@@ -702,19 +707,19 @@ def from_config(cls, config, custom_objects=None):
@factory.register_backbone_builder('movinet')
def build_movinet(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
norm_activation_config: hyperparams.Config,
- l2_regularizer: tf.keras.regularizers.Regularizer = None) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds MoViNet backbone from a config."""
backbone_type = backbone_config.type
backbone_cfg = backbone_config.get()
if backbone_type != 'movinet':
raise ValueError(f'Inconsistent backbone type {backbone_type}')
if norm_activation_config.activation is not None:
- raise ValueError(
- 'norm_activation is not used in MoViNets, but specified: %s' %
- norm_activation_config.activation)
+ logging.warn('norm_activation is not used in MoViNets, but specified: '
+ '%s', norm_activation_config.activation)
+ logging.warn('norm_activation is ignored.')
return Movinet(
model_id=backbone_cfg.model_id,
@@ -725,9 +730,11 @@ def build_movinet(
input_specs=input_specs,
activation=backbone_cfg.activation,
gating_activation=backbone_cfg.gating_activation,
+ output_states=backbone_cfg.output_states,
use_sync_bn=norm_activation_config.use_sync_bn,
norm_momentum=norm_activation_config.norm_momentum,
norm_epsilon=norm_activation_config.norm_epsilon,
- kernel_regularizer=l2_regularizer,
+ kernel_regularizer=l2_regularizer, # pyrefly: ignore[bad-argument-type]
stochastic_depth_drop_rate=backbone_cfg.stochastic_depth_drop_rate,
- use_external_states=backbone_cfg.use_external_states)
+ use_external_states=backbone_cfg.use_external_states,
+ average_pooling_type=backbone_cfg.average_pooling_type)
diff --git a/official/projects/movinet/modeling/movinet_layers.py b/official/projects/movinet/modeling/movinet_layers.py
index 61264a34580..d44a6cfdcde 100644
--- a/official/projects/movinet/modeling/movinet_layers.py
+++ b/official/projects/movinet/modeling/movinet_layers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from typing import Any, Mapping, Optional, Sequence, Tuple, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.vision.modeling.layers import nn_layers
@@ -65,8 +65,8 @@ def normalize_tuple(value: Union[int, Tuple[int, ...]], size: int, name: str):
return value_tuple
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class Squeeze3D(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Squeeze3D(tf_keras.layers.Layer):
"""Squeeze3D layer to remove singular dimensions."""
def call(self, inputs):
@@ -74,8 +74,8 @@ def call(self, inputs):
return tf.squeeze(inputs, axis=(1, 2, 3))
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class MobileConv2D(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MobileConv2D(tf_keras.layers.Layer):
"""Conv2D layer with extra options to support mobile devices.
Reshapes 5D video tensor inputs to 4D, allowing Conv2D to run across
@@ -95,11 +95,11 @@ def __init__(
use_bias: bool = True,
kernel_initializer: str = 'glorot_uniform',
bias_initializer: str = 'zeros',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- activity_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- kernel_constraint: Optional[tf.keras.constraints.Constraint] = None,
- bias_constraint: Optional[tf.keras.constraints.Constraint] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ activity_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ kernel_constraint: Optional[tf_keras.constraints.Constraint] = None,
+ bias_constraint: Optional[tf_keras.constraints.Constraint] = None,
use_depthwise: bool = False,
use_temporal: bool = False,
use_buffered_input: bool = False, # pytype: disable=annotation-type-mismatch # typed-keras
@@ -108,7 +108,7 @@ def __init__(
**kwargs): # pylint: disable=g-doc-args
"""Initializes mobile conv2d.
- For the majority of arguments, see tf.keras.layers.Conv2D.
+ For the majority of arguments, see tf_keras.layers.Conv2D.
Args:
use_depthwise: if True, use DepthwiseConv2D instead of Conv2D
@@ -148,9 +148,9 @@ def __init__(
self._batch_norm_op = batch_norm_op
self._activation_op = activation_op
- kernel_size = normalize_tuple(kernel_size, 2, 'kernel_size')
+ kernel_size = normalize_tuple(kernel_size, 2, 'kernel_size') # pyrefly: ignore[bad-argument-type]
- if self._use_temporal and kernel_size[1] > 1:
+ if self._use_temporal and kernel_size[1] > 1: # pyrefly: ignore[bad-index]
raise ValueError('Temporal conv with spatial kernel is not supported.')
if use_depthwise:
@@ -255,8 +255,8 @@ def call(self, inputs):
return x
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class ConvBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class ConvBlock(tf_keras.layers.Layer):
"""A Conv followed by optional BatchNorm and Activation."""
def __init__(
@@ -267,14 +267,15 @@ def __init__(
depthwise: bool = False,
causal: bool = False,
use_bias: bool = False,
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] =
- tf.keras.regularizers.L2(KERNEL_WEIGHT_DECAY),
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] =
+ tf_keras.regularizers.L2(KERNEL_WEIGHT_DECAY),
use_batch_norm: bool = True,
- batch_norm_layer: tf.keras.layers.Layer =
- tf.keras.layers.BatchNormalization,
+ batch_norm_layer: tf_keras.layers.Layer =
+ tf_keras.layers.BatchNormalization, # pyrefly: ignore[bad-function-definition]
batch_norm_momentum: float = 0.99,
batch_norm_epsilon: float = 1e-3,
+ use_sync_bn: bool = False,
activation: Optional[Any] = None,
conv_type: str = '3d',
use_buffered_input: bool = False, # pytype: disable=annotation-type-mismatch # typed-keras
@@ -294,6 +295,7 @@ def __init__(
batch_norm_layer: class to use for batch norm, if applied.
batch_norm_momentum: momentum of the batch norm operation, if applied.
batch_norm_epsilon: epsilon of the batch norm operation, if applied.
+ use_sync_bn: if True, use synchronized batch normalization.
activation: activation after the conv and batch norm operations.
conv_type: '3d', '2plus1d', or '3d_2plus1d'. '3d' uses the default 3D
ops. '2plus1d' split any 3D ops into two sequential 2D ops with their
@@ -310,8 +312,8 @@ def __init__(
super(ConvBlock, self).__init__(**kwargs)
- kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size')
- strides = normalize_tuple(strides, 3, 'strides')
+ kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size') # pyrefly: ignore[bad-argument-type]
+ strides = normalize_tuple(strides, 3, 'strides') # pyrefly: ignore[bad-argument-type]
self._filters = filters
self._kernel_size = kernel_size
@@ -325,6 +327,7 @@ def __init__(
self._batch_norm_layer = batch_norm_layer
self._batch_norm_momentum = batch_norm_momentum
self._batch_norm_epsilon = batch_norm_epsilon
+ self._use_sync_bn = use_sync_bn
self._activation = activation
self._conv_type = conv_type
self._use_buffered_input = use_buffered_input
@@ -351,6 +354,7 @@ def get_config(self):
'use_batch_norm': self._use_batch_norm,
'batch_norm_momentum': self._batch_norm_momentum,
'batch_norm_epsilon': self._batch_norm_epsilon,
+ 'use_sync_bn': self._use_sync_bn,
'activation': self._activation,
'conv_type': self._conv_type,
'use_buffered_input': self._use_buffered_input,
@@ -369,19 +373,21 @@ def build(self, input_shape):
self._batch_norm = self._batch_norm_layer(
momentum=self._batch_norm_momentum,
epsilon=self._batch_norm_epsilon,
+ synchronized=self._use_sync_bn,
name='bn')
- if self._conv_type != '3d' and self._kernel_size[0] > 1:
+ if self._conv_type != '3d' and self._kernel_size[0] > 1: # pyrefly: ignore[bad-index]
self._batch_norm_temporal = self._batch_norm_layer(
momentum=self._batch_norm_momentum,
epsilon=self._batch_norm_epsilon,
+ synchronized=self._use_sync_bn,
name='bn_temporal')
self._conv_temporal = None
- if self._conv_type == '3d_2plus1d' and self._kernel_size[0] > 1:
+ if self._conv_type == '3d_2plus1d' and self._kernel_size[0] > 1: # pyrefly: ignore[bad-index]
self._conv = nn_layers.Conv3D(
self._filters,
- (1, self._kernel_size[1], self._kernel_size[2]),
- strides=(1, self._strides[1], self._strides[2]),
+ (1, self._kernel_size[1], self._kernel_size[2]), # pyrefly: ignore[bad-index]
+ strides=(1, self._strides[1], self._strides[2]), # pyrefly: ignore[bad-index]
padding='same',
groups=self._groups,
use_bias=self._use_bias,
@@ -391,8 +397,8 @@ def build(self, input_shape):
name='conv3d')
self._conv_temporal = nn_layers.Conv3D(
self._filters,
- (self._kernel_size[0], 1, 1),
- strides=(self._strides[0], 1, 1),
+ (self._kernel_size[0], 1, 1), # pyrefly: ignore[bad-index]
+ strides=(self._strides[0], 1, 1), # pyrefly: ignore[bad-index]
padding=padding,
groups=self._groups,
use_bias=self._use_bias,
@@ -403,29 +409,29 @@ def build(self, input_shape):
elif self._conv_type == '2plus1d':
self._conv = MobileConv2D(
self._filters,
- (self._kernel_size[1], self._kernel_size[2]),
- strides=(self._strides[1], self._strides[2]),
+ (self._kernel_size[1], self._kernel_size[2]), # pyrefly: ignore[bad-index]
+ strides=(self._strides[1], self._strides[2]), # pyrefly: ignore[bad-index]
padding='same',
use_depthwise=self._depthwise,
groups=self._groups,
use_bias=self._use_bias,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=self._kernel_initializer, # pyrefly: ignore[bad-argument-type]
kernel_regularizer=self._kernel_regularizer,
use_buffered_input=False,
batch_norm_op=self._batch_norm,
activation_op=self._activation_layer,
name='conv2d')
- if self._kernel_size[0] > 1:
+ if self._kernel_size[0] > 1: # pyrefly: ignore[bad-index]
self._conv_temporal = MobileConv2D(
self._filters,
- (self._kernel_size[0], 1),
- strides=(self._strides[0], 1),
+ (self._kernel_size[0], 1), # pyrefly: ignore[bad-index]
+ strides=(self._strides[0], 1), # pyrefly: ignore[bad-index]
padding=padding,
use_temporal=True,
use_depthwise=self._depthwise,
groups=self._groups,
use_bias=self._use_bias,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=self._kernel_initializer, # pyrefly: ignore[bad-argument-type]
kernel_regularizer=self._kernel_regularizer,
use_buffered_input=self._use_buffered_input,
batch_norm_op=self._batch_norm_temporal,
@@ -469,8 +475,8 @@ def call(self, inputs):
return x
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class StreamBuffer(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class StreamBuffer(tf_keras.layers.Layer):
"""Stream buffer wrapper which caches activations of previous frames."""
def __init__(self,
@@ -544,7 +550,7 @@ def call(
return full_inputs, states
-@tf.keras.utils.register_keras_serializable(package='Vision')
+@tf_keras.utils.register_keras_serializable(package='Vision')
class StreamConvBlock(ConvBlock):
"""ConvBlock with StreamBuffer."""
@@ -556,14 +562,15 @@ def __init__(
depthwise: bool = False,
causal: bool = False,
use_bias: bool = False,
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = tf.keras
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = tf.keras
.regularizers.L2(KERNEL_WEIGHT_DECAY),
use_batch_norm: bool = True,
- batch_norm_layer: tf.keras.layers.Layer =
- tf.keras.layers.BatchNormalization,
+ batch_norm_layer: tf_keras.layers.Layer =
+ tf_keras.layers.BatchNormalization, # pyrefly: ignore[bad-function-definition]
batch_norm_momentum: float = 0.99,
batch_norm_epsilon: float = 1e-3,
+ use_sync_bn: bool = False,
activation: Optional[Any] = None,
conv_type: str = '3d',
state_prefix: Optional[str] = None, # pytype: disable=annotation-type-mismatch # typed-keras
@@ -583,6 +590,7 @@ def __init__(
batch_norm_layer: class to use for batch norm, if applied.
batch_norm_momentum: momentum of the batch norm operation, if applied.
batch_norm_epsilon: epsilon of the batch norm operation, if applied.
+ use_sync_bn: if True, use synchronized batch normalization.
activation: activation after the conv and batch norm operations.
conv_type: '3d', '2plus1d', or '3d_2plus1d'. '3d' uses the default 3D
ops. '2plus1d' split any 3D ops into two sequential 2D ops with their
@@ -594,8 +602,8 @@ def __init__(
Returns:
A output tensor of the StreamConvBlock operation.
"""
- kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size')
- buffer_size = kernel_size[0] - 1
+ kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size') # pyrefly: ignore[bad-argument-type]
+ buffer_size = kernel_size[0] - 1 # pyrefly: ignore[bad-index]
use_buffer = buffer_size > 0 and causal
self._state_prefix = state_prefix
@@ -613,6 +621,7 @@ def __init__(
batch_norm_layer=batch_norm_layer,
batch_norm_momentum=batch_norm_momentum,
batch_norm_epsilon=batch_norm_epsilon,
+ use_sync_bn=use_sync_bn,
activation=activation,
conv_type=conv_type,
use_buffered_input=use_buffer,
@@ -675,8 +684,8 @@ def call(self,
return x, states
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class StreamSqueezeExcitation(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class StreamSqueezeExcitation(tf_keras.layers.Layer):
"""Squeeze and excitation layer with causal mode.
Reference: https://arxiv.org/pdf/1709.01507.pdf
@@ -690,8 +699,8 @@ def __init__(
gating_activation: nn_layers.Activation = 'sigmoid',
causal: bool = False,
conv_type: str = '3d',
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = tf.keras
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = tf.keras
.regularizers.L2(KERNEL_WEIGHT_DECAY),
use_positional_encoding: bool = False,
state_prefix: Optional[str] = None, # pytype: disable=annotation-type-mismatch # typed-keras
@@ -802,12 +811,14 @@ def call(self,
states = dict(states) if states is not None else {}
if self._se_type == '3d':
- x, states = self._spatiotemporal_pool(inputs, states=states)
+ x, states = self._spatiotemporal_pool(
+ inputs, states=states, output_states=True)
elif self._se_type == '2d':
x = self._spatial_pool(inputs)
elif self._se_type == '2plus3d':
x_space = self._spatial_pool(inputs)
- x, states = self._spatiotemporal_pool(x_space, states=states)
+ x, states = self._spatiotemporal_pool(
+ x_space, states=states, output_states=True)
if not self._causal:
x = tf.tile(x, [1, tf.shape(inputs)[1], 1, 1, 1])
@@ -826,8 +837,8 @@ def call(self,
return x * inputs, states
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class MobileBottleneck(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MobileBottleneck(tf_keras.layers.Layer):
"""A depthwise inverted bottleneck block.
Uses dependency injection to allow flexible definition of different layers
@@ -835,11 +846,11 @@ class MobileBottleneck(tf.keras.layers.Layer):
"""
def __init__(self,
- expansion_layer: tf.keras.layers.Layer,
- feature_layer: tf.keras.layers.Layer,
- projection_layer: tf.keras.layers.Layer,
- attention_layer: Optional[tf.keras.layers.Layer] = None,
- skip_layer: Optional[tf.keras.layers.Layer] = None,
+ expansion_layer: tf_keras.layers.Layer,
+ feature_layer: tf_keras.layers.Layer,
+ projection_layer: tf_keras.layers.Layer,
+ attention_layer: Optional[tf_keras.layers.Layer] = None,
+ skip_layer: Optional[tf_keras.layers.Layer] = None,
stochastic_depth_drop_rate: Optional[float] = None,
**kwargs):
"""Implementation for mobile bottleneck.
@@ -861,7 +872,7 @@ def __init__(self,
self._attention_layer = attention_layer
self._skip_layer = skip_layer
self._stochastic_depth_drop_rate = stochastic_depth_drop_rate
- self._identity = tf.keras.layers.Activation(tf.identity)
+ self._identity = tf_keras.layers.Activation(tf.identity)
self._rezero = nn_layers.Scale(initializer='zeros', name='rezero')
if stochastic_depth_drop_rate:
@@ -919,8 +930,8 @@ def call(self,
return x + skip, states
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class SkipBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class SkipBlock(tf_keras.layers.Layer):
"""Skip block for bottleneck blocks."""
def __init__(
@@ -928,13 +939,14 @@ def __init__(
out_filters: int,
downsample: bool = False,
conv_type: str = '3d',
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] =
- tf.keras.regularizers.L2(KERNEL_WEIGHT_DECAY),
- batch_norm_layer: tf.keras.layers.Layer =
- tf.keras.layers.BatchNormalization,
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] =
+ tf_keras.regularizers.L2(KERNEL_WEIGHT_DECAY),
+ batch_norm_layer: tf_keras.layers.Layer =
+ tf_keras.layers.BatchNormalization, # pyrefly: ignore[bad-function-definition]
batch_norm_momentum: float = 0.99,
batch_norm_epsilon: float = 1e-3, # pytype: disable=annotation-type-mismatch # typed-keras
+ use_sync_bn: bool = False,
**kwargs):
"""Implementation for skip block.
@@ -951,6 +963,7 @@ def __init__(
batch_norm_layer: class to use for batch norm.
batch_norm_momentum: momentum of the batch norm operation.
batch_norm_epsilon: epsilon of the batch norm operation.
+ use_sync_bn: if True, use synchronized batch normalization.
**kwargs: keyword arguments to be passed to this layer.
"""
super(SkipBlock, self).__init__(**kwargs)
@@ -963,6 +976,7 @@ def __init__(
self._batch_norm_layer = batch_norm_layer
self._batch_norm_momentum = batch_norm_momentum
self._batch_norm_epsilon = batch_norm_epsilon
+ self._use_sync_bn = use_sync_bn
self._projection = ConvBlock(
filters=self._out_filters,
@@ -974,17 +988,18 @@ def __init__(
batch_norm_layer=self._batch_norm_layer,
batch_norm_momentum=self._batch_norm_momentum,
batch_norm_epsilon=self._batch_norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
name='skip_project')
if downsample:
if self._conv_type == '2plus1d':
- self._pool = tf.keras.layers.AveragePooling2D(
+ self._pool = tf_keras.layers.AveragePooling2D(
pool_size=(3, 3),
strides=(2, 2),
padding='same',
name='skip_pool')
else:
- self._pool = tf.keras.layers.AveragePooling3D(
+ self._pool = tf_keras.layers.AveragePooling3D(
pool_size=(1, 3, 3),
strides=(1, 2, 2),
padding='same',
@@ -1002,6 +1017,7 @@ def get_config(self):
'kernel_regularizer': self._kernel_regularizer,
'batch_norm_momentum': self._batch_norm_momentum,
'batch_norm_epsilon': self._batch_norm_epsilon,
+ 'use_sync_bn': self._use_sync_bn
}
base_config = super(SkipBlock, self).get_config()
return dict(list(base_config.items()) + list(config.items()))
@@ -1023,8 +1039,8 @@ def call(self, inputs):
return self._projection(x)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class MovinetBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MovinetBlock(tf_keras.layers.Layer):
"""A basic block for MoViNets.
Applies a mobile inverted bottleneck with pointwise expansion, 3D depthwise
@@ -1045,13 +1061,14 @@ def __init__(
conv_type: str = '3d',
se_type: str = '3d',
use_positional_encoding: bool = False,
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = tf.keras
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = tf.keras
.regularizers.L2(KERNEL_WEIGHT_DECAY),
- batch_norm_layer: tf.keras.layers.Layer =
- tf.keras.layers.BatchNormalization,
+ batch_norm_layer: tf_keras.layers.Layer =
+ tf_keras.layers.BatchNormalization, # pyrefly: ignore[bad-function-definition]
batch_norm_momentum: float = 0.99,
batch_norm_epsilon: float = 1e-3,
+ use_sync_bn: bool = False,
state_prefix: Optional[str] = None, # pytype: disable=annotation-type-mismatch # typed-keras
**kwargs):
"""Implementation for MoViNet block.
@@ -1081,13 +1098,14 @@ def __init__(
batch_norm_layer: class to use for batch norm.
batch_norm_momentum: momentum of the batch norm operation.
batch_norm_epsilon: epsilon of the batch norm operation.
+ use_sync_bn: if True, use synchronized batch normalization.
state_prefix: a prefix string to identify states.
**kwargs: keyword arguments to be passed to this layer.
"""
super(MovinetBlock, self).__init__(**kwargs)
- self._kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size')
- self._strides = normalize_tuple(strides, 3, 'strides')
+ self._kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size') # pyrefly: ignore[bad-argument-type]
+ self._strides = normalize_tuple(strides, 3, 'strides') # pyrefly: ignore[bad-argument-type]
# Use a multiplier of 2 if concatenating multiple features
se_multiplier = 2 if se_type == '2plus3d' else 1
@@ -1109,6 +1127,7 @@ def __init__(
self._batch_norm_layer = batch_norm_layer
self._batch_norm_momentum = batch_norm_momentum
self._batch_norm_epsilon = batch_norm_epsilon
+ self._use_sync_bn = use_sync_bn
self._state_prefix = state_prefix
self._expansion = ConvBlock(
@@ -1122,6 +1141,7 @@ def __init__(
batch_norm_layer=self._batch_norm_layer,
batch_norm_momentum=self._batch_norm_momentum,
batch_norm_epsilon=self._batch_norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
name='expansion')
self._feature = StreamConvBlock(
expand_filters,
@@ -1137,6 +1157,7 @@ def __init__(
batch_norm_layer=self._batch_norm_layer,
batch_norm_momentum=self._batch_norm_momentum,
batch_norm_epsilon=self._batch_norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
state_prefix=state_prefix,
name='feature')
self._projection = ConvBlock(
@@ -1150,6 +1171,7 @@ def __init__(
batch_norm_layer=self._batch_norm_layer,
batch_norm_momentum=self._batch_norm_momentum,
batch_norm_epsilon=self._batch_norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
name='projection')
self._attention = None
if se_type != 'none':
@@ -1185,6 +1207,7 @@ def get_config(self):
'kernel_regularizer': self._kernel_regularizer,
'batch_norm_momentum': self._batch_norm_momentum,
'batch_norm_epsilon': self._batch_norm_epsilon,
+ 'use_sync_bn': self._use_sync_bn,
'state_prefix': self._state_prefix,
}
base_config = super(MovinetBlock, self).get_config()
@@ -1232,8 +1255,8 @@ def call(self,
return self._mobile_bottleneck(inputs, states=states)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class Stem(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Stem(tf_keras.layers.Layer):
"""Stem layer for video networks.
Applies an initial convolution block operation.
@@ -1247,13 +1270,14 @@ def __init__(
causal: bool = False,
conv_type: str = '3d',
activation: nn_layers.Activation = 'swish',
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = tf.keras
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = tf.keras
.regularizers.L2(KERNEL_WEIGHT_DECAY),
- batch_norm_layer: tf.keras.layers.Layer =
- tf.keras.layers.BatchNormalization,
+ batch_norm_layer: tf_keras.layers.Layer =
+ tf_keras.layers.BatchNormalization, # pyrefly: ignore[bad-function-definition]
batch_norm_momentum: float = 0.99,
batch_norm_epsilon: float = 1e-3,
+ use_sync_bn: bool = False,
state_prefix: Optional[str] = None, # pytype: disable=annotation-type-mismatch # typed-keras
**kwargs):
"""Implementation for video model stem.
@@ -1273,14 +1297,15 @@ def __init__(
batch_norm_layer: class to use for batch norm.
batch_norm_momentum: momentum of the batch norm operation.
batch_norm_epsilon: epsilon of the batch norm operation.
+ use_sync_bn: if True, use synchronized batch normalization.
state_prefix: a prefix string to identify states.
**kwargs: keyword arguments to be passed to this layer.
"""
super(Stem, self).__init__(**kwargs)
self._out_filters = out_filters
- self._kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size')
- self._strides = normalize_tuple(strides, 3, 'strides')
+ self._kernel_size = normalize_tuple(kernel_size, 3, 'kernel_size') # pyrefly: ignore[bad-argument-type]
+ self._strides = normalize_tuple(strides, 3, 'strides') # pyrefly: ignore[bad-argument-type]
self._causal = causal
self._conv_type = conv_type
self._activation = activation
@@ -1289,6 +1314,7 @@ def __init__(
self._batch_norm_layer = batch_norm_layer
self._batch_norm_momentum = batch_norm_momentum
self._batch_norm_epsilon = batch_norm_epsilon
+ self._use_sync_bn = use_sync_bn
self._state_prefix = state_prefix
self._stem = StreamConvBlock(
@@ -1304,6 +1330,7 @@ def __init__(
batch_norm_layer=self._batch_norm_layer,
batch_norm_momentum=self._batch_norm_momentum,
batch_norm_epsilon=self._batch_norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
state_prefix=self._state_prefix,
name='stem')
@@ -1320,6 +1347,7 @@ def get_config(self):
'kernel_regularizer': self._kernel_regularizer,
'batch_norm_momentum': self._batch_norm_momentum,
'batch_norm_epsilon': self._batch_norm_epsilon,
+ 'use_sync_bn': self._use_sync_bn,
'state_prefix': self._state_prefix,
}
base_config = super(Stem, self).get_config()
@@ -1343,8 +1371,8 @@ def call(self,
return self._stem(inputs, states=states)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class Head(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Head(tf_keras.layers.Layer):
"""Head layer for video networks.
Applies pointwise projection and global pooling.
@@ -1355,13 +1383,15 @@ def __init__(
project_filters: int,
conv_type: str = '3d',
activation: nn_layers.Activation = 'swish',
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = tf.keras
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = tf.keras
.regularizers.L2(KERNEL_WEIGHT_DECAY),
- batch_norm_layer: tf.keras.layers.Layer =
- tf.keras.layers.BatchNormalization,
+ batch_norm_layer: tf_keras.layers.Layer =
+ tf_keras.layers.BatchNormalization, # pyrefly: ignore[bad-function-definition]
batch_norm_momentum: float = 0.99,
batch_norm_epsilon: float = 1e-3,
+ use_sync_bn: bool = False,
+ average_pooling_type: str = '3d',
state_prefix: Optional[str] = None, # pytype: disable=annotation-type-mismatch # typed-keras
**kwargs):
"""Implementation for video model head.
@@ -1378,6 +1408,9 @@ def __init__(
batch_norm_layer: class to use for batch norm.
batch_norm_momentum: momentum of the batch norm operation.
batch_norm_epsilon: epsilon of the batch norm operation.
+ use_sync_bn: if True, use synchronized batch normalization.
+ average_pooling_type: The average pooling type. Currently supporting
+ ['3d', '2d', 'none'].
state_prefix: a prefix string to identify states.
**kwargs: keyword arguments to be passed to this layer.
"""
@@ -1391,6 +1424,7 @@ def __init__(
self._batch_norm_layer = batch_norm_layer
self._batch_norm_momentum = batch_norm_momentum
self._batch_norm_epsilon = batch_norm_epsilon
+ self._use_sync_bn = use_sync_bn
self._state_prefix = state_prefix
self._project = ConvBlock(
@@ -1403,9 +1437,18 @@ def __init__(
batch_norm_layer=self._batch_norm_layer,
batch_norm_momentum=self._batch_norm_momentum,
batch_norm_epsilon=self._batch_norm_epsilon,
+ use_sync_bn=self._use_sync_bn,
name='project')
- self._pool = nn_layers.GlobalAveragePool3D(
- keepdims=True, causal=False, state_prefix=state_prefix)
+ if average_pooling_type.lower() == '3d':
+ self._pool = nn_layers.GlobalAveragePool3D(
+ keepdims=True, causal=False, state_prefix=state_prefix)
+ elif average_pooling_type.lower() == '2d':
+ self._pool = nn_layers.SpatialAveragePool3D(keepdims=True)
+ elif average_pooling_type == 'none':
+ self._pool = None
+ else:
+ raise ValueError(
+ '%s average_pooling_type is not supported.' % average_pooling_type)
def get_config(self):
"""Returns a dictionary containing the config used for initialization."""
@@ -1417,6 +1460,7 @@ def get_config(self):
'kernel_regularizer': self._kernel_regularizer,
'batch_norm_momentum': self._batch_norm_momentum,
'batch_norm_epsilon': self._batch_norm_epsilon,
+ 'use_sync_bn': self._use_sync_bn,
'state_prefix': self._state_prefix,
}
base_config = super(Head, self).get_config()
@@ -1439,11 +1483,15 @@ def call(
"""
states = dict(states) if states is not None else {}
x = self._project(inputs)
- return self._pool(x, states=states)
+ if self._pool is not None:
+ outputs = self._pool(x, states=states, output_states=True)
+ else:
+ outputs = (x, states)
+ return outputs
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class ClassifierHead(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class ClassifierHead(tf_keras.layers.Layer):
"""Head layer for video networks.
Applies dense projection, dropout, and classifier projection. Expects input
@@ -1459,9 +1507,9 @@ def __init__(
activation: nn_layers.Activation = 'swish',
output_activation: Optional[nn_layers.Activation] = None,
max_pool_predictions: bool = False,
- kernel_initializer: tf.keras.initializers.Initializer = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] =
- tf.keras.regularizers.L2(KERNEL_WEIGHT_DECAY), # pytype: disable=annotation-type-mismatch # typed-keras
+ kernel_initializer: tf_keras.initializers.Initializer = 'HeNormal', # pyrefly: ignore[bad-function-definition]
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] =
+ tf_keras.regularizers.L2(KERNEL_WEIGHT_DECAY), # pytype: disable=annotation-type-mismatch # typed-keras
**kwargs):
"""Implementation for video model classifier head.
@@ -1494,7 +1542,7 @@ def __init__(
self._kernel_initializer = kernel_initializer
self._kernel_regularizer = kernel_regularizer
- self._dropout = tf.keras.layers.Dropout(dropout_rate)
+ self._dropout = tf_keras.layers.Dropout(dropout_rate)
self._head = ConvBlock(
filters=head_filters,
kernel_size=1,
@@ -1508,7 +1556,7 @@ def __init__(
self._classifier = ConvBlock(
filters=num_classes,
kernel_size=1,
- kernel_initializer=tf.keras.initializers.random_normal(stddev=0.01),
+ kernel_initializer=tf_keras.initializers.random_normal(stddev=0.01),
kernel_regularizer=None,
use_bias=True,
use_batch_norm=False,
@@ -1518,7 +1566,7 @@ def __init__(
self._squeeze = Squeeze3D()
output_activation = output_activation if output_activation else 'linear'
- self._cast = tf.keras.layers.Activation(
+ self._cast = tf_keras.layers.Activation(
output_activation, dtype='float32', name='cast')
def get_config(self):
diff --git a/official/projects/movinet/modeling/movinet_layers_test.py b/official/projects/movinet/modeling/movinet_layers_test.py
index b4027043c1a..621ca516f7f 100644
--- a/official/projects/movinet/modeling/movinet_layers_test.py
+++ b/official/projects/movinet/modeling/movinet_layers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for movinet_layers.py."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet.modeling import movinet_layers
from official.vision.modeling.layers import nn_layers
@@ -64,7 +64,7 @@ def test_mobile_conv2d(self):
self.assertAllClose(predicted, expected)
def test_mobile_conv2d_bn(self):
- batch_norm_op = tf.keras.layers.BatchNormalization(
+ batch_norm_op = tf_keras.layers.BatchNormalization(
momentum=0.9,
epsilon=1.,
name='bn')
diff --git a/official/projects/movinet/modeling/movinet_model.py b/official/projects/movinet/modeling/movinet_model.py
index 0b527f7c159..a1f6b8c2839 100644
--- a/official/projects/movinet/modeling/movinet_model.py
+++ b/official/projects/movinet/modeling/movinet_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from typing import Any, Dict, Mapping, Optional, Sequence, Tuple, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet.configs import movinet as cfg
from official.projects.movinet.modeling import movinet_layers
@@ -27,20 +27,20 @@
from official.vision.modeling import factory_3d as model_factory
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class MovinetClassifier(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class MovinetClassifier(tf_keras.Model):
"""A video classification class builder."""
def __init__(
self,
- backbone: tf.keras.Model,
+ backbone: tf_keras.Model,
num_classes: int,
- input_specs: Optional[Mapping[str, tf.keras.layers.InputSpec]] = None,
+ input_specs: Optional[Mapping[str, tf_keras.layers.InputSpec]] = None,
activation: str = 'swish',
dropout_rate: float = 0.0,
kernel_initializer: str = 'HeNormal',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
output_states: bool = False,
**kwargs):
"""Movinet initialization function.
@@ -62,7 +62,7 @@ def __init__(
"""
if not input_specs:
input_specs = {
- 'image': tf.keras.layers.InputSpec(shape=[None, None, None, None, 3])
+ 'image': tf_keras.layers.InputSpec(shape=[None, None, None, None, 3])
}
self._num_classes = num_classes
@@ -90,9 +90,9 @@ def __init__(
def _build_backbone(
self,
- backbone: tf.keras.Model,
- input_specs: Mapping[str, tf.keras.layers.InputSpec],
- state_specs: Optional[Mapping[str, tf.keras.layers.InputSpec]] = None,
+ backbone: tf_keras.Model,
+ input_specs: Mapping[str, tf_keras.layers.InputSpec],
+ state_specs: Optional[Mapping[str, tf_keras.layers.InputSpec]] = None,
) -> Tuple[Mapping[str, Any], Any, Any]:
"""Builds the backbone network and gets states and endpoints.
@@ -110,10 +110,10 @@ def _build_backbone(
state_specs = state_specs if state_specs is not None else {}
states = {
- name: tf.keras.Input(shape=spec.shape[1:], dtype=spec.dtype, name=name)
+ name: tf_keras.Input(shape=spec.shape[1:], dtype=spec.dtype, name=name)
for name, spec in state_specs.items()
}
- image = tf.keras.Input(shape=input_specs['image'].shape[1:], name='image')
+ image = tf_keras.Input(shape=input_specs['image'].shape[1:], name='image')
inputs = {**states, 'image': image}
if backbone.use_external_states:
@@ -148,10 +148,10 @@ def _build_backbone(
def _build_network(
self,
- backbone: tf.keras.Model,
- input_specs: Mapping[str, tf.keras.layers.InputSpec],
- state_specs: Optional[Mapping[str, tf.keras.layers.InputSpec]] = None,
- ) -> Tuple[Mapping[str, tf.keras.Input], Union[Tuple[Mapping[ # pytype: disable=invalid-annotation # typed-keras
+ backbone: tf_keras.Model,
+ input_specs: Mapping[str, tf_keras.layers.InputSpec],
+ state_specs: Optional[Mapping[str, tf_keras.layers.InputSpec]] = None,
+ ) -> Tuple[Mapping[str, tf_keras.Input], Union[Tuple[Mapping[ # pytype: disable=invalid-annotation # typed-keras
str, tf.Tensor], Mapping[str, tf.Tensor]], Mapping[str, tf.Tensor]]]:
"""Builds the model network.
@@ -185,7 +185,7 @@ def _build_network(
return inputs, outputs
def initial_state_specs(
- self, input_shape: Sequence[int]) -> Dict[str, tf.keras.layers.InputSpec]:
+ self, input_shape: Sequence[int]) -> Dict[str, tf_keras.layers.InputSpec]:
return self._backbone.initial_state_specs(input_shape=input_shape)
@tf.function
@@ -199,7 +199,7 @@ def checkpoint_items(self) -> Dict[str, Any]:
return dict(backbone=self.backbone)
@property
- def backbone(self) -> tf.keras.Model:
+ def backbone(self) -> tf_keras.Model:
"""Returns the backbone of the model."""
return self._backbone
@@ -221,21 +221,21 @@ def get_config(self):
def from_config(cls, config, custom_objects=None):
# Each InputSpec may need to be deserialized
# This handles the case where we want to load a saved_model loaded with
- # `tf.keras.models.load_model`
+ # `tf_keras.models.load_model`
if config['input_specs']:
for name in config['input_specs']:
if isinstance(config['input_specs'][name], dict):
- config['input_specs'][name] = tf.keras.layers.deserialize(
+ config['input_specs'][name] = tf_keras.layers.deserialize(
config['input_specs'][name])
return cls(**config)
@model_factory.register_model_builder('movinet')
def build_movinet_model(
- input_specs: Mapping[str, tf.keras.layers.InputSpec],
+ input_specs: Mapping[str, tf_keras.layers.InputSpec],
model_config: cfg.MovinetModel,
num_classes: int,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None):
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None):
"""Builds movinet model."""
logging.info('Building movinet model with num classes: %s', num_classes)
if l2_regularizer is not None:
@@ -252,7 +252,7 @@ def build_movinet_model(
backbone,
num_classes=num_classes,
kernel_regularizer=l2_regularizer,
- input_specs=input_specs_dict,
+ input_specs=input_specs_dict, # pyrefly: ignore[bad-argument-type]
activation=model_config.activation,
dropout_rate=model_config.dropout_rate,
output_states=model_config.output_states)
diff --git a/official/projects/movinet/modeling/movinet_model_test.py b/official/projects/movinet/modeling/movinet_model_test.py
index 3187e38a3f6..f8cbc8af9a8 100644
--- a/official/projects/movinet/modeling/movinet_model_test.py
+++ b/official/projects/movinet/modeling/movinet_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet.modeling import movinet
from official.projects.movinet.modeling import movinet_model
@@ -29,9 +29,9 @@ def test_movinet_classifier_creation(self, is_training):
"""Test for creation of a Movinet classifier."""
temporal_size = 16
spatial_size = 224
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None, temporal_size, spatial_size, spatial_size, 3])
backbone = movinet.Movinet(model_id='a0', input_specs=input_specs)
@@ -48,7 +48,7 @@ def test_movinet_classifier_creation(self, is_training):
def test_movinet_classifier_stream(self):
"""Test if the classifier can be run in streaming mode."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -75,7 +75,7 @@ def test_movinet_classifier_stream(self):
def test_movinet_classifier_stream_pos_enc(self):
"""Test if the classifier can be run in streaming mode with pos encoding."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -103,7 +103,7 @@ def test_movinet_classifier_stream_pos_enc(self):
def test_movinet_classifier_stream_pos_enc_2plus1d(self):
"""Test if the model can run in streaming mode with pos encoding, (2+1)D."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -132,7 +132,7 @@ def test_movinet_classifier_stream_pos_enc_2plus1d(self):
def test_movinet_classifier_mobile(self):
"""Test if the model can run with mobile parameters."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -184,8 +184,8 @@ def test_saved_model_save_load(self):
model.build([1, 5, 172, 172, 3])
model.compile(metrics=['acc'])
- tf.keras.models.save_model(model, '/tmp/movinet/')
- loaded_model = tf.keras.models.load_model('/tmp/movinet/')
+ tf_keras.models.save_model(model, '/tmp/movinet/')
+ loaded_model = tf_keras.models.load_model('/tmp/movinet/')
output = loaded_model(dict(image=tf.ones([1, 1, 1, 1, 3])))
@@ -202,7 +202,7 @@ def test_saved_model_save_load(self):
)
def test_movinet_models(self, model_id, expected_params_millions):
"""Test creation of MoViNet family models with states."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
model = movinet_model.MovinetClassifier(
backbone=movinet.Movinet(
@@ -216,7 +216,7 @@ def test_movinet_models(self, model_id, expected_params_millions):
def test_movinet_a0_2plus1d(self):
"""Test creation of MoViNet with 2plus1d configuration."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
model_2plus1d = movinet_model.MovinetClassifier(
backbone=movinet.Movinet(
diff --git a/official/projects/movinet/modeling/movinet_test.py b/official/projects/movinet/modeling/movinet_test.py
index 0b082c00a62..55c5edb37c8 100644
--- a/official/projects/movinet/modeling/movinet_test.py
+++ b/official/projects/movinet/modeling/movinet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for movinet.py."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet.modeling import movinet
@@ -24,13 +24,13 @@ class MoViNetTest(parameterized.TestCase, tf.test.TestCase):
def test_network_creation(self):
"""Test creation of MoViNet family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
network = movinet.Movinet(
model_id='a0',
causal=True,
)
- inputs = tf.keras.Input(shape=(8, 128, 128, 3), batch_size=1)
+ inputs = tf_keras.Input(shape=(8, 128, 128, 3), batch_size=1)
endpoints, states = network(inputs)
self.assertAllEqual(endpoints['stem'].shape, [1, 8, 64, 64, 8])
@@ -45,7 +45,7 @@ def test_network_creation(self):
def test_network_with_states(self):
"""Test creation of MoViNet family models with states."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -70,7 +70,7 @@ def test_network_with_states(self):
def test_movinet_stream(self):
"""Test if the backbone can be run in streaming mode."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -100,7 +100,7 @@ def test_movinet_stream(self):
def test_movinet_stream_nse(self):
"""Test if the backbone can be run in streaming mode w/o SE layer."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -142,7 +142,7 @@ def test_movinet_stream_nse(self):
msg=f'Expecting stream_buffer only, found {state_key}')
def test_movinet_2plus1d_stream(self):
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
@@ -172,7 +172,7 @@ def test_movinet_2plus1d_stream(self):
self.assertAllClose(predicted, expected, 1e-5, 1e-5)
def test_movinet_3d_2plus1d_stream(self):
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = movinet.Movinet(
model_id='a0',
diff --git a/official/projects/movinet/movinet_streaming_model_training_and_inference.ipynb b/official/projects/movinet/movinet_streaming_model_training_and_inference.ipynb
new file mode 100644
index 00000000000..5917d2d589e
--- /dev/null
+++ b/official/projects/movinet/movinet_streaming_model_training_and_inference.ipynb
@@ -0,0 +1,1017 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "9OEXtK6c6lyz"
+ },
+ "source": [
+ "# Movinet Streaming Model Tutorial\n",
+ "\n",
+ "This tutorial trains a [MoViNet](https://arxiv.org/abs/2103.11511) with `A0` configuration from the TensorFlow Model Garden package (tensorflow-models).\n",
+ "\n",
+ "[Model Garden](https://www.tensorflow.org/tfmodels) contains a collection of state-of-the-art models, implemented with TensorFlow's high-level APIs. The implementations demonstrate the best practices for modeling, letting users to take full advantage of TensorFlow for their research and product development.\n",
+ "\n",
+ "**Streaming Models**:\n",
+ "\n",
+ "Streaming models implement causal (2+1)D convolutions with stream buffers. Streaming models use (2+1)D convolution instead of 3D to utilize optimized tf.nn.conv2d operations, which offer fast inference on CPU. Streaming models can be run on individual frames or on larger video clips like base models.\n",
+ "\n",
+ "**Note**: A3, A4, and A5 models use a positional encoding in the squeeze-excitation blocks, while A0, A1, and A2 do not. For the smaller models, accuracy is unaffected without positional encoding, while for the larger models accuracy is significantly worse without positional encoding.\n",
+ "\n",
+ "**Dataset:** [UCF_101](https://www.tensorflow.org/datasets/catalog/ucf101)\n",
+ "* A 101-label video classification dataset\n",
+ "\n",
+ "**This tutorial demonstrates how to:**\n",
+ "\n",
+ "* Use models from the TensorFlow Models package.\n",
+ "* Train/Fine-tune a pre-built [MoViNet](https://arxiv.org/abs/2103.11511) for Video Classification.\n",
+ "* Export the trained/tuned MoViNet-A0-Stream model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "560IfSkP7fld"
+ },
+ "source": [
+ "## Install cecessary libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "MVX4IofucuZK"
+ },
+ "outputs": [],
+ "source": [
+ "!pip install -U -q \"tf-models-official\"\n",
+ "\n",
+ "\n",
+ "# Install the mediapy package for visualizing images/videos.\n",
+ "# See https://github.com/google/mediapy\n",
+ "!command -v ffmpeg \u003e/dev/null || (apt update \u0026\u0026 apt install -y ffmpeg)\n",
+ "!pip install -q mediapy remotezip\n",
+ "!pip install -U -q git+https://github.com/tensorflow/docs"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "K-clGdjV7nF5"
+ },
+ "source": [
+ "## Import necessary libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "AklLemGhc_Jm"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "import tqdm\n",
+ "import random\n",
+ "import pathlib\n",
+ "import imageio\n",
+ "import itertools\n",
+ "import collections\n",
+ "\n",
+ "import cv2\n",
+ "import numpy as np\n",
+ "import remotezip as rz\n",
+ "import seaborn as sns\n",
+ "import matplotlib\n",
+ "import matplotlib.pyplot as plt\n",
+ "from tensorflow_docs.vis import embed\n",
+ "\n",
+ "import keras\n",
+ "import tensorflow as tf\n",
+ "import tensorflow_hub as hub\n",
+ "from tensorflow.keras import layers\n",
+ "from tensorflow.keras.optimizers import Adam\n",
+ "from tensorflow.keras.losses import SparseCategoricalCrossentropy"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "epVTJ5BH7ps3"
+ },
+ "source": [
+ "## Import the MoViNet model from TensorFlow Models (tf-models-official)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "xli9cfvw7s3t"
+ },
+ "outputs": [],
+ "source": [
+ "from official.projects.movinet.modeling import movinet\n",
+ "from official.projects.movinet.modeling import movinet_model\n",
+ "from official.projects.movinet.tools import export_saved_model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "pbMYvcJX76KP"
+ },
+ "source": [
+ "## Download Subdataset of UCF_101"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "F5XwQuLzd5Id"
+ },
+ "outputs": [],
+ "source": [
+ "# @title Helper functions for loading data and visualizing\n",
+ "def list_files_per_class(zip_url):\n",
+ " \"\"\"\n",
+ " List the files in each class of the dataset given the zip URL.\n",
+ "\n",
+ " Args:\n",
+ " zip_url: URL from which the files can be unzipped.\n",
+ "\n",
+ " Return:\n",
+ " files: List of files in each of the classes.\n",
+ " \"\"\"\n",
+ " files = []\n",
+ " with rz.RemoteZip(URL) as zip:\n",
+ " for zip_info in zip.infolist():\n",
+ " files.append(zip_info.filename)\n",
+ " return files\n",
+ "\n",
+ "def get_class(fname):\n",
+ " \"\"\"\n",
+ " Retrieve the name of the class given a filename.\n",
+ "\n",
+ " Args:\n",
+ " fname: Name of the file in the UCF101 dataset.\n",
+ "\n",
+ " Return:\n",
+ " Class that the file belongs to.\n",
+ " \"\"\"\n",
+ " return fname.split('_')[-3]\n",
+ "\n",
+ "def get_files_per_class(files):\n",
+ " \"\"\"\n",
+ " Retrieve the files that belong to each class.\n",
+ "\n",
+ " Args:\n",
+ " files: List of files in the dataset.\n",
+ "\n",
+ " Return:\n",
+ " Dictionary of class names (key) and files (values).\n",
+ " \"\"\"\n",
+ " files_for_class = collections.defaultdict(list)\n",
+ " for fname in files:\n",
+ " class_name = get_class(fname)\n",
+ " files_for_class[class_name].append(fname)\n",
+ " return files_for_class\n",
+ "\n",
+ "def download_from_zip(zip_url, to_dir, file_names):\n",
+ " \"\"\"\n",
+ " Download the contents of the zip file from the zip URL.\n",
+ "\n",
+ " Args:\n",
+ " zip_url: Zip URL containing data.\n",
+ " to_dir: Directory to download data to.\n",
+ " file_names: Names of files to download.\n",
+ " \"\"\"\n",
+ " with rz.RemoteZip(zip_url) as zip:\n",
+ " for fn in tqdm.tqdm(file_names):\n",
+ " class_name = get_class(fn)\n",
+ " zip.extract(fn, str(to_dir / class_name))\n",
+ " unzipped_file = to_dir / class_name / fn\n",
+ "\n",
+ " fn = pathlib.Path(fn).parts[-1]\n",
+ " output_file = to_dir / class_name / fn\n",
+ " unzipped_file.rename(output_file,)\n",
+ "\n",
+ "def split_class_lists(files_for_class, count):\n",
+ " \"\"\"\n",
+ " Returns the list of files belonging to a subset of data as well as the remainder of\n",
+ " files that need to be downloaded.\n",
+ "\n",
+ " Args:\n",
+ " files_for_class: Files belonging to a particular class of data.\n",
+ " count: Number of files to download.\n",
+ "\n",
+ " Return:\n",
+ " split_files: Files belonging to the subset of data.\n",
+ " remainder: Dictionary of the remainder of files that need to be downloaded.\n",
+ " \"\"\"\n",
+ " split_files = []\n",
+ " remainder = {}\n",
+ " for cls in files_for_class:\n",
+ " split_files.extend(files_for_class[cls][:count])\n",
+ " remainder[cls] = files_for_class[cls][count:]\n",
+ " return split_files, remainder\n",
+ "\n",
+ "def download_ufc_101_subset(zip_url, num_classes, splits, download_dir):\n",
+ " \"\"\"\n",
+ " Download a subset of the UFC101 dataset and split them into various parts, such as\n",
+ " training, validation, and test.\n",
+ "\n",
+ " Args:\n",
+ " zip_url: Zip URL containing data.\n",
+ " num_classes: Number of labels.\n",
+ " splits: Dictionary specifying the training, validation, test, etc. (key) division of data\n",
+ " (value is number of files per split).\n",
+ " download_dir: Directory to download data to.\n",
+ "\n",
+ " Return:\n",
+ " dir: Posix path of the resulting directories containing the splits of data.\n",
+ " \"\"\"\n",
+ " files = list_files_per_class(zip_url)\n",
+ " for f in files:\n",
+ " tokens = f.split('/')\n",
+ " if len(tokens) \u003c= 2:\n",
+ " files.remove(f) # Remove that item from the list if it does not have a filename\n",
+ "\n",
+ " files_for_class = get_files_per_class(files)\n",
+ "\n",
+ " classes = list(files_for_class.keys())[:num_classes]\n",
+ "\n",
+ " for cls in classes:\n",
+ " new_files_for_class = files_for_class[cls]\n",
+ " random.shuffle(new_files_for_class)\n",
+ " files_for_class[cls] = new_files_for_class\n",
+ "\n",
+ " # Only use the number of classes you want in the dictionary\n",
+ " files_for_class = {x: files_for_class[x] for x in list(files_for_class)[:num_classes]}\n",
+ "\n",
+ " dirs = {}\n",
+ " for split_name, split_count in splits.items():\n",
+ " print(split_name, \":\")\n",
+ " split_dir = download_dir / split_name\n",
+ " split_files, files_for_class = split_class_lists(files_for_class, split_count)\n",
+ " download_from_zip(zip_url, split_dir, split_files)\n",
+ " dirs[split_name] = split_dir\n",
+ "\n",
+ " return dirs\n",
+ "\n",
+ "def format_frames(frame, output_size):\n",
+ " \"\"\"\n",
+ " Pad and resize an image from a video.\n",
+ "\n",
+ " Args:\n",
+ " frame: Image that needs to resized and padded.\n",
+ " output_size: Pixel size of the output frame image.\n",
+ "\n",
+ " Return:\n",
+ " Formatted frame with padding of specified output size.\n",
+ " \"\"\"\n",
+ " frame = tf.image.convert_image_dtype(frame, tf.float32)\n",
+ " frame = tf.image.resize_with_pad(frame, *output_size)\n",
+ " return frame\n",
+ "\n",
+ "def frames_from_video_file(video_path, n_frames, output_size = (172,172), frame_step = 15):\n",
+ " \"\"\"\n",
+ " Creates frames from each video file present for each category.\n",
+ "\n",
+ " Args:\n",
+ " video_path: File path to the video.\n",
+ " n_frames: Number of frames to be created per video file.\n",
+ " output_size: Pixel size of the output frame image.\n",
+ "\n",
+ " Return:\n",
+ " An NumPy array of frames in the shape of (n_frames, height, width, channels).\n",
+ " \"\"\"\n",
+ " # Read each video frame by frame\n",
+ " result = []\n",
+ " src = cv2.VideoCapture(str(video_path))\n",
+ "\n",
+ " video_length = src.get(cv2.CAP_PROP_FRAME_COUNT)\n",
+ "\n",
+ " need_length = 1 + (n_frames - 1) * frame_step\n",
+ "\n",
+ " if need_length \u003e video_length:\n",
+ " start = 0\n",
+ " else:\n",
+ " max_start = video_length - need_length\n",
+ " start = random.randint(0, max_start + 1)\n",
+ "\n",
+ " src.set(cv2.CAP_PROP_POS_FRAMES, start)\n",
+ " # ret is a boolean indicating whether read was successful, frame is the image itself\n",
+ " ret, frame = src.read()\n",
+ " result.append(format_frames(frame, output_size))\n",
+ "\n",
+ " for _ in range(n_frames - 1):\n",
+ " for _ in range(frame_step):\n",
+ " ret, frame = src.read()\n",
+ " if ret:\n",
+ " frame = format_frames(frame, output_size)\n",
+ " result.append(frame)\n",
+ " else:\n",
+ " result.append(np.zeros_like(result[0]))\n",
+ " src.release()\n",
+ " result = np.array(result)[..., [2, 1, 0]]\n",
+ "\n",
+ " return result\n",
+ "\n",
+ "def to_gif(images):\n",
+ " converted_images = np.clip(images * 255, 0, 255).astype(np.uint8)\n",
+ " imageio.mimsave('./animation.gif', converted_images, fps=10)\n",
+ " return embed.embed_file('./animation.gif')\n",
+ "\n",
+ "\n",
+ "class FrameGenerator:\n",
+ " def __init__(self, path, n_frames, training = False):\n",
+ " \"\"\" Returns a set of frames with their associated label.\n",
+ "\n",
+ " Args:\n",
+ " path: Video file paths.\n",
+ " n_frames: Number of frames.\n",
+ " training: Boolean to determine if training dataset is being created.\n",
+ " \"\"\"\n",
+ " self.path = path\n",
+ " self.n_frames = n_frames\n",
+ " self.training = training\n",
+ " self.class_names = sorted(set(p.name for p in self.path.iterdir() if p.is_dir()))\n",
+ " self.class_ids_for_name = dict((name, idx) for idx, name in enumerate(self.class_names))\n",
+ "\n",
+ " def get_files_and_class_names(self):\n",
+ " video_paths = list(self.path.glob('*/*.avi'))\n",
+ " classes = [p.parent.name for p in video_paths]\n",
+ " return video_paths, classes\n",
+ "\n",
+ " def __call__(self):\n",
+ " video_paths, classes = self.get_files_and_class_names()\n",
+ "\n",
+ " pairs = list(zip(video_paths, classes))\n",
+ "\n",
+ " if self.training:\n",
+ " random.shuffle(pairs)\n",
+ "\n",
+ " for path, name in pairs:\n",
+ " video_frames = frames_from_video_file(path, self.n_frames)\n",
+ " label = self.class_ids_for_name[name] # Encode labels\n",
+ " yield video_frames, label"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JbHsejeveDFW"
+ },
+ "outputs": [],
+ "source": [
+ "# Helper functions below used are taken from following tutorials\n",
+ "# https://www.tensorflow.org/tutorials/video/video_classification\n",
+ "# https://www.tensorflow.org/tutorials/video/transfer_learning_with_movinet\n",
+ "\n",
+ "URL = 'https://storage.googleapis.com/thumos14_files/UCF101_videos.zip'\n",
+ "download_dir = pathlib.Path('./UCF101_subset/')\n",
+ "subset_paths = download_ufc_101_subset(URL,\n",
+ " num_classes = 10,\n",
+ " splits = {\"train\": 40, \"val\": 10, \"test\": 10},\n",
+ " download_dir = download_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Riw5ZQDx8aA0"
+ },
+ "source": [
+ "## Prepare train, valid and test dataset"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "CJ2xmsdVeGMM"
+ },
+ "outputs": [],
+ "source": [
+ "batch_size = 10\n",
+ "num_frames = 8\n",
+ "\n",
+ "CLASSES = sorted(os.listdir('./UCF101_subset/train'))\n",
+ "\n",
+ "output_signature = (tf.TensorSpec(shape = (None, None, None, 3), dtype = tf.float32),\n",
+ " tf.TensorSpec(shape = (), dtype = tf.int16))\n",
+ "\n",
+ "train_ds = tf.data.Dataset.from_generator(FrameGenerator(subset_paths['train'], num_frames, training = True),\n",
+ " output_signature = output_signature)\n",
+ "train_ds = train_ds.batch(batch_size)\n",
+ "\n",
+ "val_ds = tf.data.Dataset.from_generator(FrameGenerator(subset_paths['val'], num_frames),\n",
+ " output_signature = output_signature)\n",
+ "val_ds = val_ds.batch(batch_size)\n",
+ "\n",
+ "test_ds = tf.data.Dataset.from_generator(FrameGenerator(subset_paths['test'], num_frames),\n",
+ " output_signature = output_signature)\n",
+ "test_ds = test_ds.batch(batch_size)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "HuyVEA5l8myh"
+ },
+ "source": [
+ "### Check the prepared training batch of the data"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "5tPCI35TeihA"
+ },
+ "outputs": [],
+ "source": [
+ "for frames, labels in train_ds.take(1):\n",
+ " print(f\"Shape: {frames.shape}\")\n",
+ " print(f\"Label: {labels.shape}\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "aNk_b5gk8s-g"
+ },
+ "source": [
+ "## Build the model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "fmdzt6Eu9Afp"
+ },
+ "source": [
+ "### Construct the backbone with proper parameters"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "msnxCHSa9Cb2"
+ },
+ "outputs": [],
+ "source": [
+ "model_id = 'a0'\n",
+ "use_positional_encoding = model_id in {'a3', 'a4', 'a5'}\n",
+ "resolution = 172\n",
+ "\n",
+ "backbone = movinet.Movinet(\n",
+ " model_id=model_id,\n",
+ " causal=True,\n",
+ " conv_type='2plus1d',\n",
+ " se_type='2plus3d',\n",
+ " activation='hard_swish',\n",
+ " gating_activation='hard_sigmoid',\n",
+ " use_positional_encoding=use_positional_encoding,\n",
+ " use_external_states=False,\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "iSCnGqY69IzG"
+ },
+ "source": [
+ "### Construct the model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "n5m9qZ__eo92"
+ },
+ "outputs": [],
+ "source": [
+ "# Note: this is a temporary model constructed for the\n",
+ "# purpose of loading the pre-trained checkpoint. Only\n",
+ "# the backbone will be used to build the custom classifier.\n",
+ "\n",
+ "model = movinet_model.MovinetClassifier(\n",
+ " backbone,\n",
+ " num_classes=600,\n",
+ " output_states=True)\n",
+ "\n",
+ "# Create your example input here.\n",
+ "# Refer to the paper for recommended input shapes.\n",
+ "inputs = tf.ones([1, 13, 172, 172, 3])\n",
+ "\n",
+ "# [Optional] Build the model and load a pretrained checkpoint.\n",
+ "model.build(inputs.shape)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "6jfwRvD59RyX"
+ },
+ "source": [
+ "### Load the pretrained weights"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "g7y8zwcz9SJA"
+ },
+ "outputs": [],
+ "source": [
+ "# Extract pretrained weights\n",
+ "!wget https://storage.googleapis.com/tf_model_garden/vision/movinet/movinet_a0_stream.tar.gz -O movinet_a0_stream.tar.gz -q\n",
+ "!tar -xvf movinet_a0_stream.tar.gz\n",
+ "\n",
+ "checkpoint_dir = 'movinet_a0_stream'\n",
+ "checkpoint_path = tf.train.latest_checkpoint(checkpoint_dir)\n",
+ "checkpoint = tf.train.Checkpoint(model=model)\n",
+ "status = checkpoint.restore(checkpoint_path)\n",
+ "status.assert_existing_objects_matched()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "X0MyBsZnTrnh"
+ },
+ "source": [
+ "### Set up the distribution strategy"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "IHyqT0csQvRl"
+ },
+ "outputs": [],
+ "source": [
+ "# Detect hardware\n",
+ "try:\n",
+ " tpu_resolver = tf.distribute.cluster_resolver.TPUClusterResolver() # TPU detection\n",
+ "except ValueError:\n",
+ " tpu_resolver = None\n",
+ " gpus = tf.config.experimental.list_logical_devices(\"GPU\")\n",
+ "\n",
+ "# Select appropriate distribution strategy\n",
+ "if tpu_resolver:\n",
+ " tf.config.experimental_connect_to_cluster(tpu_resolver)\n",
+ " tf.tpu.experimental.initialize_tpu_system(tpu_resolver)\n",
+ " distribution_strategy = tf.distribute.experimental.TPUStrategy(tpu_resolver)\n",
+ " print('Running on TPU ', tpu_resolver.cluster_spec().as_dict()['worker'])\n",
+ "elif len(gpus) \u003e 1:\n",
+ " distribution_strategy = tf.distribute.MirroredStrategy([gpu.name for gpu in gpus])\n",
+ " print('Running on multiple GPUs ', [gpu.name for gpu in gpus])\n",
+ "elif len(gpus) == 1:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on single GPU ', gpus[0].name)\n",
+ "else:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on CPU')\n",
+ "\n",
+ "print(\"Number of accelerators: \", distribution_strategy.num_replicas_in_sync)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "CmPbIku_9ejS"
+ },
+ "source": [
+ "### Construct custom classifier with required number of classes"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "lpXf9GFWlE2-"
+ },
+ "outputs": [],
+ "source": [
+ "def build_classifier(batch_size, num_frames, resolution, backbone, num_classes):\n",
+ " \"\"\"Builds a classifier on top of a backbone model.\"\"\"\n",
+ " model = movinet_model.MovinetClassifier(\n",
+ " backbone=backbone,\n",
+ " num_classes=num_classes)\n",
+ " model.build([batch_size, num_frames, resolution, resolution, 3])\n",
+ "\n",
+ " return model\n",
+ "\n",
+ "# Construct loss, optimizer and compile the model\n",
+ "with distribution_strategy.scope():\n",
+ " model = build_classifier(batch_size, num_frames, resolution, backbone, 10)\n",
+ " loss_obj = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)\n",
+ " optimizer = tf.keras.optimizers.Adam(learning_rate = 0.001)\n",
+ " model.compile(loss=loss_obj, optimizer=optimizer, metrics=['accuracy'])"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "j3S6pmEO9y9w"
+ },
+ "source": [
+ "## Create a callback for storing the checkpoints"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Jw-TR9LBPuP_"
+ },
+ "outputs": [],
+ "source": [
+ "checkpoint_path = \"trained_model/cp.ckpt\"\n",
+ "checkpoint_dir = os.path.dirname(checkpoint_path)\n",
+ "\n",
+ "# Create a callback that saves the model's weights\n",
+ "cp_callback = tf.keras.callbacks.ModelCheckpoint(filepath=checkpoint_path,\n",
+ " save_weights_only=True,\n",
+ " verbose=1)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "E99ucruN9-cu"
+ },
+ "source": [
+ "## Train the model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "9c7aikXAvpwz"
+ },
+ "outputs": [],
+ "source": [
+ "results = model.fit(train_ds,\n",
+ " validation_data=val_ds,\n",
+ " epochs=2,\n",
+ " validation_freq=1,\n",
+ " verbose=1,\n",
+ " callbacks=[cp_callback])"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "CSgJ8pq2-CTG"
+ },
+ "source": [
+ "## Evaluate the model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "xtPKfnsNl8SX"
+ },
+ "outputs": [],
+ "source": [
+ "model.evaluate(test_ds)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "fGckPjcx-GY7"
+ },
+ "source": [
+ "## Plot the test data confusion matrix"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "n4aQvflFmBho"
+ },
+ "outputs": [],
+ "source": [
+ "def get_actual_predicted_labels(dataset):\n",
+ " \"\"\"\n",
+ " Create a list of actual ground truth values and the predictions from the model.\n",
+ "\n",
+ " Args:\n",
+ " dataset: An iterable data structure, such as a TensorFlow Dataset, with features and labels.\n",
+ "\n",
+ " Return:\n",
+ " Ground truth and predicted values for a particular dataset.\n",
+ " \"\"\"\n",
+ " actual = [labels for _, labels in dataset.unbatch()]\n",
+ " predicted = model.predict(dataset)\n",
+ "\n",
+ " actual = tf.stack(actual, axis=0)\n",
+ " predicted = tf.concat(predicted, axis=0)\n",
+ " predicted = tf.argmax(predicted, axis=1)\n",
+ "\n",
+ " return actual, predicted"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "aXYqu7QOmEBZ"
+ },
+ "outputs": [],
+ "source": [
+ "def plot_confusion_matrix(actual, predicted, labels, ds_type):\n",
+ " cm = tf.math.confusion_matrix(actual, predicted)\n",
+ " ax = sns.heatmap(cm, annot=True, fmt='g')\n",
+ " sns.set(rc={'figure.figsize':(6, 16)})\n",
+ " sns.set(font_scale=1.4)\n",
+ " ax.set_title('Confusion matrix of action recognition for ' + ds_type)\n",
+ " ax.set_xlabel('Predicted Action')\n",
+ " ax.set_ylabel('Actual Action')\n",
+ " plt.xticks(rotation=90)\n",
+ " plt.yticks(rotation=0)\n",
+ " ax.xaxis.set_ticklabels(labels)\n",
+ " ax.yaxis.set_ticklabels(labels)\n",
+ " plt.show()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "SXk55UYDmGxl"
+ },
+ "outputs": [],
+ "source": [
+ "fg = FrameGenerator(subset_paths['train'], num_frames, training = True)\n",
+ "label_names = list(fg.class_ids_for_name.keys())"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "QyxeRFB2mKIz"
+ },
+ "outputs": [],
+ "source": [
+ "actual, predicted = get_actual_predicted_labels(test_ds)\n",
+ "plot_confusion_matrix(actual, predicted, label_names, 'test')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0WLYpxFz-P-m"
+ },
+ "source": [
+ "## Reconstruct the whole model with `use_external_states=True` to make the inference using states."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "S29Cemf_Py1h"
+ },
+ "outputs": [],
+ "source": [
+ "model_id = 'a0'\n",
+ "use_positional_encoding = model_id in {'a3', 'a4', 'a5'}\n",
+ "resolution = 172\n",
+ "\n",
+ "# Create backbone and model.\n",
+ "backbone = movinet.Movinet(\n",
+ " model_id=model_id,\n",
+ " causal=True,\n",
+ " conv_type='2plus1d',\n",
+ " se_type='2plus3d',\n",
+ " activation='hard_swish',\n",
+ " gating_activation='hard_sigmoid',\n",
+ " use_positional_encoding=use_positional_encoding,\n",
+ " use_external_states=True,\n",
+ ")\n",
+ "\n",
+ "model = movinet_model.MovinetClassifier(\n",
+ " backbone,\n",
+ " num_classes=10,\n",
+ " output_states=True)\n",
+ "\n",
+ "# Create your example input here.\n",
+ "# Refer to the paper for recommended input shapes.\n",
+ "inputs = tf.ones([1, 13, 172, 172, 3])\n",
+ "\n",
+ "# [Optional] Build the model and load a pretrained checkpoint.\n",
+ "model.build(inputs.shape)\n",
+ "\n",
+ "# Load weights from the checkpoint to the rebuilt model\n",
+ "checkpoint_dir = 'trained_model'\n",
+ "model.load_weights(tf.train.latest_checkpoint(checkpoint_dir))"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "KQDBRqnSIZ23"
+ },
+ "source": [
+ "## Inference using external states"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "aN7ExIA9JjK_"
+ },
+ "outputs": [],
+ "source": [
+ "def get_top_k(probs, k=5, label_map=CLASSES):\n",
+ " \"\"\"Outputs the top k model labels and probabilities on the given video.\"\"\"\n",
+ " top_predictions = tf.argsort(probs, axis=-1, direction='DESCENDING')[:k]\n",
+ " top_labels = tf.gather(label_map, top_predictions, axis=-1)\n",
+ " top_labels = [label.decode('utf8') for label in top_labels.numpy()]\n",
+ " top_probs = tf.gather(probs, top_predictions, axis=-1).numpy()\n",
+ " return tuple(zip(top_labels, top_probs))"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "gquiUVQGIZJ2"
+ },
+ "outputs": [],
+ "source": [
+ "# Create initial states for the stream model\n",
+ "init_states_fn = model.init_states\n",
+ "init_states = init_states_fn(tf.shape(tf.ones(shape=[1, 1, 172, 172, 3])))\n",
+ "\n",
+ "all_logits = []\n",
+ "\n",
+ "# To run on a video, pass in one frame at a time\n",
+ "states = init_states\n",
+ "for frames, label in test_ds.take(1):\n",
+ " for clip in frames[0]:\n",
+ " # Input shape: [1, 1, 172, 172, 3]\n",
+ " clip = tf.expand_dims(tf.expand_dims(clip, axis=0), axis=0)\n",
+ " logits, states = model.predict({**states, 'image': clip}, verbose=0)\n",
+ " all_logits.append(logits)\n",
+ "\n",
+ "logits = tf.concat(all_logits, 0)\n",
+ "probs = tf.nn.softmax(logits)\n",
+ "\n",
+ "final_probs = probs[-1]\n",
+ "top_k = get_top_k(final_probs)\n",
+ "print()\n",
+ "for label, prob in top_k:\n",
+ " print(label, prob)\n",
+ "\n",
+ "frames, label = list(test_ds.take(1))[0]\n",
+ "to_gif(frames[0].numpy())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "TYxRHeuUINi8"
+ },
+ "source": [
+ "## Export to saved model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "aHaiqMXbVxC6"
+ },
+ "outputs": [],
+ "source": [
+ "saved_model_dir = 'model'\n",
+ "tflite_filename = 'model.tflite'\n",
+ "input_shape = [1, 1, 172, 172, 3]\n",
+ "\n",
+ "# Convert to saved model\n",
+ "export_saved_model.export_saved_model(\n",
+ " model=model,\n",
+ " input_shape=input_shape,\n",
+ " export_path=saved_model_dir,\n",
+ " causal=True,\n",
+ " bundle_input_init_states_fn=False)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Ct2fyC00I63a"
+ },
+ "source": [
+ "## Convert to TF Lite"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "lHr__VEsVOea"
+ },
+ "outputs": [],
+ "source": [
+ "converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir)\n",
+ "tflite_model = converter.convert()\n",
+ "\n",
+ "with open(tflite_filename, 'wb') as f:\n",
+ " f.write(tflite_model)\n",
+ "\n",
+ "# Create the interpreter and signature runner\n",
+ "interpreter = tf.lite.Interpreter(model_path=tflite_filename)\n",
+ "runner = interpreter.get_signature_runner()\n",
+ "\n",
+ "init_states = {\n",
+ " name: tf.zeros(x['shape'], dtype=x['dtype'])\n",
+ " for name, x in runner.get_input_details().items()\n",
+ "}\n",
+ "del init_states['image']\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "bmp9vBZ_I-Bw"
+ },
+ "source": [
+ "## Inference using external states on tflite model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "-0Exy6QJgep-"
+ },
+ "outputs": [],
+ "source": [
+ "# To run on a video, pass in one frame at a time\n",
+ "states = init_states\n",
+ "for frames, label in test_ds.take(1):\n",
+ " for clip in frames[0]:\n",
+ " # Input shape: [1, 1, 172, 172, 3]\n",
+ " outputs = runner(**states, image=clip)\n",
+ " logits = outputs.pop('logits')[0]\n",
+ " states = outputs\n",
+ "\n",
+ "probs = tf.nn.softmax(logits)\n",
+ "top_k = get_top_k(probs)\n",
+ "print()\n",
+ "for label, prob in top_k:\n",
+ " print(label, prob)\n",
+ "\n",
+ "frames, label = list(test_ds.take(1))[0]\n",
+ "to_gif(frames[0].numpy())"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "private_outputs": true,
+ "provenance": [
+ {
+ "file_id": "1JB5Zzgt1DPwblAy9ScfajQtyIC-W5__f",
+ "timestamp": 1679078074504
+ }
+ ],
+ "toc_visible": true
+ },
+ "gpuClass": "standard",
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/movinet/movinet_tutorial.ipynb b/official/projects/movinet/movinet_tutorial.ipynb
index a29cc72fac9..4e0b71bc7ee 100644
--- a/official/projects/movinet/movinet_tutorial.ipynb
+++ b/official/projects/movinet/movinet_tutorial.ipynb
@@ -49,20 +49,13 @@
"\n",
"# tf-models-official is the stable Model Garden package\n",
"# tf-models-nightly includes latest changes\n",
- "!pip install -q tf-models-nightly\n",
+ "!pip install -U -q \"tf-models-official\"\n",
"\n",
- "# Install tfds nightly to download ucf101\n",
- "!pip install -q tfds-nightly\n",
"\n",
"# Install the mediapy package for visualizing images/videos.\n",
"# See https://github.com/google/mediapy\n",
"!command -v ffmpeg \u003e/dev/null || (apt update \u0026\u0026 apt install -y ffmpeg)\n",
- "!pip install -q mediapy\n",
- "\n",
- "# Due to a bug, we reinstall opencv\n",
- "# See https://stackoverflow.com/q/70537488\n",
- "!pip uninstall -q -y opencv-python-headless\n",
- "!pip install -q \"opencv-python-headless\u003c4.3\""
+ "!pip install -q mediapy"
]
},
{
@@ -75,7 +68,6 @@
"source": [
"# Run imports\n",
"import os\n",
- "\n",
"import matplotlib as mpl\n",
"import matplotlib.pyplot as plt\n",
"import mediapy as media\n",
@@ -86,7 +78,10 @@
"import tensorflow_datasets as tfds\n",
"import tensorflow_hub as hub\n",
"import tqdm\n",
+ "import absl.logging\n",
"\n",
+ "tf.get_logger().setLevel('ERROR')\n",
+ "absl.logging.set_verbosity(absl.logging.ERROR)\n",
"mpl.rcParams.update({\n",
" 'font.size': 10,\n",
"})"
@@ -398,19 +393,33 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 17681,
+ "status": "ok",
+ "timestamp": 1674679874816,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
"id": "P0bZfrAsqPv2",
- "outputId": "bd82571f-8dfd-4faf-ed10-e34708b0405d"
+ "outputId": "fe2074c1-684e-4973-af76-c6b6deee511b"
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
- "jumping jacks 0.9166437\n",
- "zumba 0.016020728\n",
- "doing aerobics 0.008053946\n",
- "dancing charleston 0.006083599\n",
- "lunge 0.0035062772\n"
+ "1/1 [==============================] - 18s 18s/step\n",
+ "jumping jacks 0.9166436\n",
+ "zumba 0.016020758\n",
+ "doing aerobics 0.008053949\n",
+ "dancing charleston 0.006083598\n",
+ "lunge 0.0035062768\n"
]
}
],
@@ -430,7 +439,7 @@
"source": [
"## Run Streaming Model Inference with TensorFlow Hub and Plot Predictions\n",
"\n",
- "We will load MoViNet-A0-Stream from TensorFlow Hub as part of the [MoViNet collection](https://tfhub.dev/google/collections/movinet/).\n",
+ "We will load MoViNet-A2-Stream from TensorFlow Hub as part of the [MoViNet collection](https://tfhub.dev/google/collections/movinet/).\n",
"\n",
"The following code will:\n",
"\n",
@@ -457,20 +466,47 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 10662,
+ "status": "ok",
+ "timestamp": 1674679945388,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
"id": "YqSkt7l8ltwt",
- "outputId": "6ccf1dd6-95d1-43b1-efdb-2e931dd3a19d"
+ "outputId": "cdee6358-cc78-48a8-dc0a-cb4c502a3368"
},
"outputs": [
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "100%|██████████| 13/13 [00:10\u003c00:00, 1.23it/s]"
+ ]
+ },
{
"name": "stdout",
"output_type": "stream",
"text": [
- "100%|██████████| 13/13 [00:08\u003c00:00, 1.58it/s]\n",
- "jumping jacks 0.9998123\n",
- "zumba 0.00011835508\n",
- "doing aerobics 3.3375818e-05\n",
- "dancing charleston 4.9819987e-06\n",
- "finger snapping 3.8673647e-06\n"
+ "\n",
+ "jumping jacks 0.9998122\n",
+ "zumba 0.00011835461\n",
+ "doing aerobics 3.3375778e-05\n",
+ "dancing charleston 4.9820073e-06\n",
+ "finger snapping 3.867353e-06\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "\n"
]
}
],
@@ -502,9 +538,31 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
- "id": "Xdox556CtMRb"
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 6945,
+ "status": "ok",
+ "timestamp": 1674679952309,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "Xdox556CtMRb",
+ "outputId": "41a242ce-87aa-430e-d75c-516935c1da3b"
},
- "outputs": [],
+ "outputs": [
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "100%|██████████| 13/13 [00:06\u003c00:00, 1.90it/s]\n"
+ ]
+ }
+ ],
"source": [
"# Generate a plot and output to a video tensor\n",
"plot_video = plot_streaming_top_preds(probs, video, video_fps=8.)"
@@ -546,11 +604,7 @@
},
"outputs": [],
"source": [
- "# Run imports\n",
- "from official.vision.configs import video_classification\n",
- "from official.projects.movinet.configs import movinet as movinet_configs\n",
"from official.projects.movinet.modeling import movinet\n",
- "from official.projects.movinet.modeling import movinet_layers\n",
"from official.projects.movinet.modeling import movinet_model\n",
"from official.projects.movinet.tools import export_saved_model"
]
@@ -559,37 +613,72 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
- "id": "RLkV0xtPvfkY"
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 15595,
+ "status": "ok",
+ "timestamp": 1674679969145,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "5DGwH9qi87Oe",
+ "outputId": "f28271e5-afc1-4fb2-9ad5-1e6fdd89a64f"
},
- "outputs": [],
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "movinet_a0_stream/\n",
+ "movinet_a0_stream/ckpt-1.data-00000-of-00001\n",
+ "movinet_a0_stream/ckpt-1.index\n",
+ "movinet_a0_stream/checkpoint\n"
+ ]
+ },
+ {
+ "data": {
+ "text/plain": [
+ "\u003ctensorflow.python.checkpoint.checkpoint.CheckpointLoadStatus at 0x7f0b0dc436a0\u003e"
+ ]
+ },
+ "execution_count": 12,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
"source": [
- "# Export to saved model\n",
- "saved_model_dir = 'model'\n",
- "tflite_filename = 'model.tflite'\n",
- "input_shape = [1, 1, 172, 172, 3]\n",
- "batch_size, num_frames, image_size, = input_shape[:3]\n",
- "\n",
- "tf.keras.backend.clear_session()\n",
+ "model_id = 'a0'\n",
+ "use_positional_encoding = model_id in {'a3', 'a4', 'a5'}\n",
"\n",
- "# Create the model\n",
- "input_specs = tf.keras.layers.InputSpec(shape=input_shape)\n",
+ "# Create backbone and model.\n",
"backbone = movinet.Movinet(\n",
- " model_id='a0',\n",
+ " model_id=model_id,\n",
" causal=True,\n",
" conv_type='2plus1d',\n",
" se_type='2plus3d',\n",
- " input_specs=input_specs,\n",
" activation='hard_swish',\n",
" gating_activation='hard_sigmoid',\n",
- " use_sync_bn=False,\n",
- " use_external_states=True)\n",
+ " use_positional_encoding=use_positional_encoding,\n",
+ " use_external_states=True,\n",
+ ")\n",
+ "\n",
"model = movinet_model.MovinetClassifier(\n",
- " backbone=backbone,\n",
- " activation='hard_swish',\n",
+ " backbone,\n",
" num_classes=600,\n",
- " output_states=True,\n",
- " input_specs=dict(image=input_specs))\n",
- "model.build([1, 1, 1, 1, 3])\n",
+ " output_states=True)\n",
+ "\n",
+ "# Create your example input here.\n",
+ "# Refer to the paper for recommended input shapes.\n",
+ "inputs = tf.ones([1, 13, 172, 172, 3])\n",
+ "\n",
+ "# [Optional] Build the model and load a pretrained checkpoint.\n",
+ "model.build(inputs.shape)\n",
+ "\n",
"\n",
"# Extract pretrained weights\n",
"!wget https://storage.googleapis.com/tf_model_garden/vision/movinet/movinet_a0_stream.tar.gz -O movinet_a0_stream.tar.gz -q\n",
@@ -597,6 +686,23 @@
"\n",
"checkpoint_dir = 'movinet_a0_stream'\n",
"checkpoint_path = tf.train.latest_checkpoint(checkpoint_dir)\n",
+ "checkpoint = tf.train.Checkpoint(model=model)\n",
+ "status = checkpoint.restore(checkpoint_path)\n",
+ "status.assert_existing_objects_matched()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "RLkV0xtPvfkY"
+ },
+ "outputs": [],
+ "source": [
+ "# Export to saved model\n",
+ "saved_model_dir = 'model'\n",
+ "tflite_filename = 'model.tflite'\n",
+ "input_shape = [1, 1, 172, 172, 3]\n",
"\n",
"# Convert to saved model\n",
"export_saved_model.export_saved_model(\n",
@@ -604,8 +710,7 @@
" input_shape=input_shape,\n",
" export_path=saved_model_dir,\n",
" causal=True,\n",
- " bundle_input_init_states_fn=False,\n",
- " checkpoint_path=checkpoint_path)"
+ " bundle_input_init_states_fn=False)"
]
},
{
@@ -638,19 +743,33 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 9,
+ "status": "ok",
+ "timestamp": 1674680160875,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
"id": "-TQ-7oSJIlTA",
- "outputId": "a15519ff-d08c-40bc-fbea-d3a58169450c"
+ "outputId": "2a7cf5f5-7648-44dd-a5d5-69da9ea82838"
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
- "jumping jacks 0.9791285\n",
- "jogging 0.0019550633\n",
- "riding unicycle 0.0017429002\n",
- "passing soccer ball 0.0016952101\n",
- "stretching arm 0.0014458151\n"
+ "\n",
+ "jumping jacks 0.9733523\n",
+ "jogging 0.0032490466\n",
+ "stretching arm 0.002780116\n",
+ "riding unicycle 0.0019377996\n",
+ "passing soccer ball 0.0016310472\n"
]
}
],
@@ -742,19 +861,21 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
"executionInfo": {
- "elapsed": 2957,
+ "elapsed": 7,
"status": "ok",
- "timestamp": 1619748263684,
+ "timestamp": 1674680161043,
"user": {
- "displayName": "",
- "photoUrl": "",
- "userId": ""
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
},
- "user_tz": 360
+ "user_tz": 480
},
"id": "boQHbcfDhXpJ",
- "outputId": "eabc3307-d6bf-4f29-cc5a-c8dc6360701b"
+ "outputId": "e3e8b59d-861a-4358-f19a-99b059d7eb47"
},
"outputs": [
{
@@ -780,14 +901,16 @@
" 256x256 UCF with the first action recognition split.\n",
" \"\"\",\n",
" homepage='https://www.crcv.ucf.edu/data-sets/ucf101/',\n",
- " data_path='/readahead/128M/placer/prod/home/tensorflow-datasets-cns-storage-owner/datasets/ucf101/ucf101_1_256/2.0.0',\n",
+ " data_path='~/tensorflow_datasets/ucf101/ucf101_1_256/2.0.0',\n",
+ " file_format=tfrecord,\n",
" download_size=6.48 GiB,\n",
- " dataset_size=Unknown size,\n",
+ " dataset_size=7.41 GiB,\n",
" features=FeaturesDict({\n",
" 'label': ClassLabel(shape=(), dtype=tf.int64, num_classes=101),\n",
" 'video': Video(Image(shape=(256, 256, 3), dtype=tf.uint8)),\n",
" }),\n",
" supervised_keys=None,\n",
+ " disable_shuffling=False,\n",
" splits={\n",
" 'test': \u003cSplitInfo num_examples=3783, num_shards=32\u003e,\n",
" 'train': \u003cSplitInfo num_examples=9537, num_shards=64\u003e,\n",
@@ -811,10 +934,8 @@
")"
]
},
- "execution_count": null,
- "metadata": {
- "tags": []
- },
+ "execution_count": 18,
+ "metadata": {},
"output_type": "execute_result"
}
],
@@ -917,9 +1038,34 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
- "id": "JpfxpeGSsbzJ"
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 6934,
+ "status": "ok",
+ "timestamp": 1674680181444,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "JpfxpeGSsbzJ",
+ "outputId": "83a49ab1-b28e-45c6-c0b3-2fc446944f65"
},
- "outputs": [],
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "movinet_a0_base/\n",
+ "movinet_a0_base/checkpoint\n",
+ "movinet_a0_base/ckpt-1.data-00000-of-00001\n",
+ "movinet_a0_base/ckpt-1.index\n"
+ ]
+ }
+ ],
"source": [
"model_id = 'a0'\n",
"\n",
@@ -1022,19 +1168,21 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
"executionInfo": {
- "elapsed": 982253,
+ "elapsed": 3426904,
"status": "ok",
- "timestamp": 1619750139919,
+ "timestamp": 1674683608342,
"user": {
- "displayName": "",
- "photoUrl": "",
- "userId": ""
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
},
- "user_tz": 360
+ "user_tz": 480
},
"id": "Zecc_K3lga8I",
- "outputId": "e4c5c61e-aa08-47db-c04c-42dea3efb545"
+ "outputId": "1946f687-aece-49f5-f5e7-09aad9f9882b"
},
"outputs": [
{
@@ -1042,11 +1190,11 @@
"output_type": "stream",
"text": [
"Epoch 1/3\n",
- "1192/1192 [==============================] - 551s 451ms/step - loss: 2.5050 - top_1: 0.6692 - top_5: 0.8753 - val_loss: 1.6310 - val_top_1: 0.8109 - val_top_5: 0.9701\n",
+ "1192/1192 [==============================] - 1151s 949ms/step - loss: 2.5097 - top_1: 0.6726 - top_5: 0.8745 - val_loss: 1.6358 - val_top_1: 0.8125 - val_top_5: 0.9666\n",
"Epoch 2/3\n",
- "1192/1192 [==============================] - 533s 447ms/step - loss: 1.3336 - top_1: 0.9024 - top_5: 0.9906 - val_loss: 1.4576 - val_top_1: 0.8451 - val_top_5: 0.9740\n",
+ "1192/1192 [==============================] - 1138s 951ms/step - loss: 1.3347 - top_1: 0.9062 - top_5: 0.9894 - val_loss: 1.4627 - val_top_1: 0.8400 - val_top_5: 0.9709\n",
"Epoch 3/3\n",
- "1192/1192 [==============================] - 531s 446ms/step - loss: 1.2298 - top_1: 0.9329 - top_5: 0.9943 - val_loss: 1.4351 - val_top_1: 0.8514 - val_top_5: 0.9762\n"
+ "1192/1192 [==============================] - 1138s 955ms/step - loss: 1.2301 - top_1: 0.9340 - top_5: 0.9943 - val_loss: 1.4386 - val_top_1: 0.8438 - val_top_5: 0.9751\n"
]
}
],
@@ -1075,9 +1223,57 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
- "id": "9fZhzhRJRd2J"
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 839
+ },
+ "executionInfo": {
+ "elapsed": 33,
+ "status": "ok",
+ "timestamp": 1674683608343,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "9fZhzhRJRd2J",
+ "outputId": "43a8b75e-28f2-456c-b6d7-5a02b65d2443"
},
- "outputs": [],
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "Reusing TensorBoard on port 43479 (pid 278134), started 19:51:44 ago. (Use '!kill 278134' to kill it.)"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/javascript": [
+ "\n",
+ " (async () =\u003e {\n",
+ " const url = new URL(await google.colab.kernel.proxyPort(43479, {'cache': true}));\n",
+ " url.searchParams.set('tensorboardColab', 'true');\n",
+ " const iframe = document.createElement('iframe');\n",
+ " iframe.src = url;\n",
+ " iframe.setAttribute('width', '100%');\n",
+ " iframe.setAttribute('height', '800');\n",
+ " iframe.setAttribute('frameborder', 0);\n",
+ " document.body.appendChild(iframe);\n",
+ " })();\n",
+ " "
+ ],
+ "text/plain": [
+ "\u003cIPython.core.display.Javascript object\u003e"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ }
+ ],
"source": [
"%reload_ext tensorboard\n",
"%tensorboard --logdir logs --port 0"
@@ -1085,20 +1281,24 @@
}
],
"metadata": {
+ "accelerator": "GPU",
"colab": {
- "collapsed_sections": [],
"last_runtime": {
"build_target": "//learning/deepmind/dm_python:dm_notebook3",
"kind": "private"
},
- "name": "movinet_tutorial.ipynb",
"provenance": [
+ {
+ "file_id": "1nV2uiAZgRk2Ble2kximcRZvCSv9c02Xd",
+ "timestamp": 1674684623688
+ },
{
"file_id": "11msGCxFjxwioBOBJavP9alfTclUQCJf-",
"timestamp": 1617043059980
}
]
},
+ "gpuClass": "standard",
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
diff --git a/official/projects/movinet/tools/__init__.py b/official/projects/movinet/tools/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/movinet/tools/__init__.py
+++ b/official/projects/movinet/tools/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/movinet/tools/convert_3d_2plus1d.py b/official/projects/movinet/tools/convert_3d_2plus1d.py
index 0349c607517..bf9aa604936 100644
--- a/official/projects/movinet/tools/convert_3d_2plus1d.py
+++ b/official/projects/movinet/tools/convert_3d_2plus1d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet.modeling import movinet
from official.projects.movinet.modeling import movinet_model
diff --git a/official/projects/movinet/tools/convert_3d_2plus1d_test.py b/official/projects/movinet/tools/convert_3d_2plus1d_test.py
index d2899c9b960..1d9ad6d6a1f 100644
--- a/official/projects/movinet/tools/convert_3d_2plus1d_test.py
+++ b/official/projects/movinet/tools/convert_3d_2plus1d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet.modeling import movinet
from official.projects.movinet.modeling import movinet_model
diff --git a/official/projects/movinet/tools/export_saved_model.py b/official/projects/movinet/tools/export_saved_model.py
index 86be6616477..d90eb0b74cf 100644
--- a/official/projects/movinet/tools/export_saved_model.py
+++ b/official/projects/movinet/tools/export_saved_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -54,7 +54,7 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet.modeling import movinet
from official.projects.movinet.modeling import movinet_model
@@ -110,31 +110,40 @@
flags.DEFINE_string(
'checkpoint_path', '',
'Checkpoint path to load. Leave blank for default initialization.')
+flags.DEFINE_bool(
+ 'assert_checkpoint_objects_matched',
+ True,
+ 'Whether to check the checkpoint objects exactly match those of the model.',
+)
FLAGS = flags.FLAGS
def export_saved_model(
- model: tf.keras.Model,
+ model: tf_keras.Model,
input_shape: Tuple[int, int, int, int, int],
export_path: str = '/tmp/movinet/',
causal: bool = False,
bundle_input_init_states_fn: bool = True,
- checkpoint_path: Optional[str] = None) -> None:
+ checkpoint_path: Optional[str] = None,
+ assert_checkpoint_objects_matched: bool = True,
+) -> None:
"""Exports a MoViNet model to a saved model.
Args:
- model: the tf.keras.Model to export.
- input_shape: The 5D spatiotemporal input shape of size
- [batch_size, num_frames, image_height, image_width, num_channels].
- Set the field or a shape position in the field to None for dynamic input.
+ model: the tf_keras.Model to export.
+ input_shape: The 5D spatiotemporal input shape of size [batch_size,
+ num_frames, image_height, image_width, num_channels]. Set the field or a
+ shape position in the field to None for dynamic input.
export_path: Export path to save the saved_model file.
causal: Run the model in causal mode.
bundle_input_init_states_fn: Add init_states as a function signature to the
- saved model. This is not necessary if the input shape is static (e.g.,
- for TF Lite).
+ saved model. This is not necessary if the input shape is static (e.g., for
+ TF Lite).
checkpoint_path: Checkpoint path to load. Leave blank to keep the model's
initialization.
+ assert_checkpoint_objects_matched: Whether to check the checkpoint objects
+ exactly match those of the model.
"""
# Use dimensions of 1 except the channels to export faster,
@@ -149,7 +158,8 @@ def export_saved_model(
if checkpoint_path:
checkpoint = tf.train.Checkpoint(model=model)
status = checkpoint.restore(checkpoint_path)
- status.assert_existing_objects_matched()
+ if assert_checkpoint_objects_matched:
+ status.assert_existing_objects_matched()
if causal:
# Call the model once to get the output states. Call again with `states`
@@ -185,11 +195,11 @@ def predict(inputs):
else:
signatures = predict_fn
- tf.keras.models.save_model(
+ tf_keras.models.save_model(
model, export_path, signatures=signatures)
else:
_ = model(tf.ones(input_shape_concrete))
- tf.keras.models.save_model(model, export_path)
+ tf_keras.models.save_model(model, export_path)
def build_and_export_saved_model(
@@ -205,23 +215,23 @@ def build_and_export_saved_model(
num_classes: int = 600,
input_shape: Optional[Tuple[int, int, int, int, int]] = None,
bundle_input_init_states_fn: bool = True,
- checkpoint_path: Optional[str] = None) -> None:
+ checkpoint_path: Optional[str] = None,
+ assert_checkpoint_objects_matched: bool = True,
+) -> None:
"""Builds and exports a MoViNet model to a saved model.
Args:
export_path: Export path to save the saved_model file.
model_id: MoViNet model name.
causal: Run the model in causal mode.
- conv_type: 3d, 2plus1d, or 3d_2plus1d. 3d configures the network
- to use the default 3D convolution. 2plus1d uses (2+1)D convolution
- with Conv2D operations and 2D reshaping (e.g., a 5x3x3 kernel becomes
- 3x3 followed by 5x1 conv). 3d_2plus1d uses (2+1)D convolution with
- Conv3D and no 2D reshaping (e.g., a 5x3x3 kernel becomes 1x3x3
- followed by 5x1x1 conv).
- se_type:
- 3d, 2d, or 2plus3d. 3d uses the default 3D spatiotemporal global average
- pooling for squeeze excitation. 2d uses 2D spatial global average pooling
- on each frame. 2plus3d concatenates both 3D and 2D global average
+ conv_type: 3d, 2plus1d, or 3d_2plus1d. 3d configures the network to use the
+ default 3D convolution. 2plus1d uses (2+1)D convolution with Conv2D
+ operations and 2D reshaping (e.g., a 5x3x3 kernel becomes 3x3 followed by
+ 5x1 conv). 3d_2plus1d uses (2+1)D convolution with Conv3D and no 2D
+ reshaping (e.g., a 5x3x3 kernel becomes 1x3x3 followed by 5x1x1 conv).
+ se_type: 3d, 2d, or 2plus3d. 3d uses the default 3D spatiotemporal global
+ average pooling for squeeze excitation. 2d uses 2D spatial global average
+ pooling on each frame. 2plus3d concatenates both 3D and 2D global average
pooling.
activation: The main activation to use across layers.
classifier_activation: The classifier activation to use.
@@ -230,17 +240,19 @@ def build_and_export_saved_model(
use_positional_encoding: Whether to use positional encoding (only applied
when causal=True).
num_classes: The number of classes for prediction.
- input_shape: The 5D spatiotemporal input shape of size
- [batch_size, num_frames, image_height, image_width, num_channels].
- Set the field or a shape position in the field to None for dynamic input.
+ input_shape: The 5D spatiotemporal input shape of size [batch_size,
+ num_frames, image_height, image_width, num_channels]. Set the field or a
+ shape position in the field to None for dynamic input.
bundle_input_init_states_fn: Add init_states as a function signature to the
- saved model. This is not necessary if the input shape is static (e.g.,
- for TF Lite).
+ saved model. This is not necessary if the input shape is static (e.g., for
+ TF Lite).
checkpoint_path: Checkpoint path to load. Leave blank for default
initialization.
+ assert_checkpoint_objects_matched: Whether to check the checkpoint objects
+ exactly match those of the model.
"""
- input_specs = tf.keras.layers.InputSpec(shape=input_shape)
+ input_specs = tf_keras.layers.InputSpec(shape=input_shape)
# Override swish activation implementation to remove custom gradients
if activation == 'swish':
@@ -268,11 +280,13 @@ def build_and_export_saved_model(
export_saved_model(
model=model,
- input_shape=input_shape,
+ input_shape=input_shape, # pyrefly: ignore[bad-argument-type]
export_path=export_path,
causal=causal,
bundle_input_init_states_fn=bundle_input_init_states_fn,
- checkpoint_path=checkpoint_path)
+ checkpoint_path=checkpoint_path,
+ assert_checkpoint_objects_matched=assert_checkpoint_objects_matched,
+ )
def main(_) -> None:
@@ -291,7 +305,9 @@ def main(_) -> None:
num_classes=FLAGS.num_classes,
input_shape=input_shape,
bundle_input_init_states_fn=FLAGS.bundle_input_init_states_fn,
- checkpoint_path=FLAGS.checkpoint_path)
+ checkpoint_path=FLAGS.checkpoint_path,
+ assert_checkpoint_objects_matched=FLAGS.assert_checkpoint_objects_matched,
+ )
print(' ----- Done. Saved Model is saved at {}'.format(FLAGS.export_path))
diff --git a/official/projects/movinet/tools/export_saved_model_test.py b/official/projects/movinet/tools/export_saved_model_test.py
index a06be1c9e5a..3835cc10306 100644
--- a/official/projects/movinet/tools/export_saved_model_test.py
+++ b/official/projects/movinet/tools/export_saved_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,13 +13,16 @@
# limitations under the License.
"""Tests for export_saved_model."""
-
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_hub as hub
+# pylint: disable=g-direct-tensorflow-import
+from tensorflow.lite.python import interpreter as tfl_interpreter
+# pylint: enable=g-direct-tensorflow-import
from official.projects.movinet.tools import export_saved_model
+
FLAGS = flags.FLAGS
@@ -37,13 +40,13 @@ def test_movinet_export_a0_base_with_tfhub(self):
encoder = hub.KerasLayer(saved_model_path, trainable=True)
- inputs = tf.keras.layers.Input(
+ inputs = tf_keras.layers.Input(
shape=[None, None, None, 3],
dtype=tf.float32)
outputs = encoder(dict(image=inputs))
- model = tf.keras.Model(inputs, outputs)
+ model = tf_keras.Model(inputs, outputs)
example_input = tf.ones([1, 8, 172, 172, 3])
outputs = model(example_input)
@@ -62,7 +65,7 @@ def test_movinet_export_a0_stream_with_tfhub(self):
encoder = hub.KerasLayer(saved_model_path, trainable=True)
- image_input = tf.keras.layers.Input(
+ image_input = tf_keras.layers.Input(
shape=[None, None, None, 3],
dtype=tf.float32,
name='image')
@@ -73,7 +76,7 @@ def test_movinet_export_a0_stream_with_tfhub(self):
for name, state in init_states_fn(tf.constant([0, 0, 0, 0, 3])).items()
}
states_input = {
- name: tf.keras.Input(shape[1:], dtype=dtype, name=name)
+ name: tf_keras.Input(shape[1:], dtype=dtype, name=name)
for name, (shape, dtype) in state_shapes.items()
}
@@ -81,7 +84,7 @@ def test_movinet_export_a0_stream_with_tfhub(self):
outputs = encoder(inputs)
- model = tf.keras.Model(inputs, outputs)
+ model = tf_keras.Model(inputs, outputs)
example_input = tf.ones([1, 8, 172, 172, 3])
frames = tf.split(example_input, example_input.shape[1], axis=1)
@@ -120,7 +123,7 @@ def test_movinet_export_a0_stream_with_tflite(self):
converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_path)
tflite_model = converter.convert()
- interpreter = tf.lite.Interpreter(model_content=tflite_model)
+ interpreter = tfl_interpreter.Interpreter(model_content=tflite_model)
runner = interpreter.get_signature_runner('serving_default')
def state_name(name: str) -> str:
diff --git a/official/projects/movinet/tools/quantize_movinet.py b/official/projects/movinet/tools/quantize_movinet.py
index 5e34c3c9e52..946012fdc2f 100644
--- a/official/projects/movinet/tools/quantize_movinet.py
+++ b/official/projects/movinet/tools/quantize_movinet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -172,7 +172,7 @@ def get_dataset() -> tf.data.Dataset:
def stateful_representative_dataset_generator(
- model: tf.keras.Model,
+ model: tf_keras.Model,
dataset_iter: Any,
init_states: Mapping[str, tf.Tensor],
save_dataset_to_tfrecords: bool = False,
@@ -224,7 +224,7 @@ def stateful_representative_dataset_generator(
predictions=predictions,
output_states=output_states,
groundtruth_label_id=groundtruth_label_id,
- output_dataset_dir=output_dataset_dir,
+ output_dataset_dir=output_dataset_dir, # pyrefly: ignore[bad-argument-type]
file_index=counter)
yield {'image': frame, **input_states}
counter += 1
@@ -270,7 +270,7 @@ def quantize_movinet(dataset_fn):
# Load model
encoder = hub.KerasLayer(FLAGS.saved_model_with_states_dir, trainable=False)
- inputs = tf.keras.layers.Input(
+ inputs = tf_keras.layers.Input(
shape=[1, FLAGS.image_size, FLAGS.image_size, 3],
dtype=tf.float32,
name='image')
@@ -283,14 +283,14 @@ def quantize_movinet(dataset_fn):
tf.constant([1, 1, FLAGS.image_size, FLAGS.image_size, 3])).items()
}
states_input = {
- name: tf.keras.Input(shape[1:], dtype=dtype, name=name)
+ name: tf_keras.Input(shape[1:], dtype=dtype, name=name)
for name, (shape, dtype) in state_shapes.items()
}
# The inputs to the model are the states and the video
inputs = {**states_input, 'image': inputs}
outputs = encoder(inputs)
- model = tf.keras.Model(inputs, outputs, name='movinet_stream')
+ model = tf_keras.Model(inputs, outputs, name='movinet_stream')
input_shape = tf.constant(
[1, FLAGS.num_frames, FLAGS.image_size, FLAGS.image_size, 3])
init_states = init_states_fn(input_shape)
diff --git a/official/projects/movinet/train.py b/official/projects/movinet/train.py
index ef42379ec7b..7bc06649165 100644
--- a/official/projects/movinet/train.py
+++ b/official/projects/movinet/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/movinet/train_test.py b/official/projects/movinet/train_test.py
index ad53802ac65..0d643543292 100644
--- a/official/projects/movinet/train_test.py
+++ b/official/projects/movinet/train_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl import flags
from absl import logging
from absl.testing import flagsaver
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.movinet import train as train_lib
from official.vision.dataloaders import tfexample_utils
diff --git a/official/projects/mtop/README.md b/official/projects/mtop/README.md
new file mode 100644
index 00000000000..8efe0961621
--- /dev/null
+++ b/official/projects/mtop/README.md
@@ -0,0 +1,11 @@
+# MTOP (All Birds with One Stone: Multi-task Text Classification for Efficient Inference with One Forward Pass)
+
+**Note:** This project is a work in progress; please stay tuned.
+
+MTOP is a text encoder multi-task method that can conduct one forward pass to
+make predictions for all tasks. We propose prompt-based modules tailored for
+the multi-task setting and a conditional pooler for flexible task
+representations, and initialization from single task models for effective
+knowledge transfer. Our proposed approach gets superior performance on news
+tasks and the GLUE benchmark. We also release a multi-task news dataset.
+
diff --git a/official/projects/nhnet/__init__.py b/official/projects/nhnet/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/nhnet/__init__.py
+++ b/official/projects/nhnet/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/nhnet/configs.py b/official/projects/nhnet/configs.py
index fa0a787f9a4..ef762e15206 100644
--- a/official/projects/nhnet/configs.py
+++ b/official/projects/nhnet/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/nhnet/configs_test.py b/official/projects/nhnet/configs_test.py
index 54678ddecf2..506ebb25041 100644
--- a/official/projects/nhnet/configs_test.py
+++ b/official/projects/nhnet/configs_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for configs."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.nhnet import configs
BERT2BERT_CONFIG = {
diff --git a/official/projects/nhnet/decoder.py b/official/projects/nhnet/decoder.py
index dc1d8e3fd86..f4cb33f4c58 100644
--- a/official/projects/nhnet/decoder.py
+++ b/official/projects/nhnet/decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,14 +14,14 @@
"""Transformer decoder that mimics a BERT encoder, to load BERT checkpoints."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import model_utils as transformer_utils
from official.modeling import tf_utils
from official.nlp.modeling import layers
-class TransformerDecoder(tf.keras.layers.Layer):
+class TransformerDecoder(tf_keras.layers.Layer):
"""Transformer decoder stack."""
def __init__(self,
@@ -60,7 +60,7 @@ def build(self, unused_input_shapes):
intermediate_activation=self.intermediate_activation,
dropout_rate=self.hidden_dropout_prob,
attention_dropout_rate=self.attention_probs_dropout_prob,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=self.initializer_range),
multi_channel_cross_attention=self.multi_channel_cross_attention,
name=("layer_%d" % i)))
@@ -163,7 +163,7 @@ def get_attention_bias(input_tensor,
return tf.where(bias < 0, tf.zeros_like(bias), tf.ones_like(bias))
-class AttentionBias(tf.keras.layers.Layer):
+class AttentionBias(tf_keras.layers.Layer):
def __init__(self, bias_type, **kwargs):
super(AttentionBias, self).__init__(**kwargs)
@@ -173,7 +173,7 @@ def call(self, inputs):
return get_attention_bias(inputs, self.bias_type)
-class EmbeddingPostprocessor(tf.keras.layers.Layer):
+class EmbeddingPostprocessor(tf_keras.layers.Layer):
"""Performs various post-processing on a word embedding tensor."""
def __init__(self,
@@ -194,7 +194,7 @@ def __init__(self,
self.initializer_range = initializer_range
if not initializer:
- self.initializer = tf.keras.initializers.TruncatedNormal(
+ self.initializer = tf_keras.initializers.TruncatedNormal(
stddev=initializer_range)
else:
self.initializer = initializer
@@ -212,7 +212,7 @@ def build(self, input_shapes):
self.type_embeddings = self.add_weight(
"type_embeddings",
shape=[self.token_type_vocab_size, width],
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=self.initializer_range),
dtype=self.dtype)
@@ -221,13 +221,13 @@ def build(self, input_shapes):
self.position_embeddings = self.add_weight(
"position_embeddings",
shape=[self.max_position_embeddings, width],
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=self.initializer_range),
dtype=self.dtype)
- self.output_layer_norm = tf.keras.layers.LayerNormalization(
+ self.output_layer_norm = tf_keras.layers.LayerNormalization(
name="layer_norm", axis=-1, epsilon=1e-12, dtype=tf.float32)
- self.output_dropout = tf.keras.layers.Dropout(
+ self.output_dropout = tf_keras.layers.Dropout(
rate=self.dropout_prob, dtype=tf.float32)
super(EmbeddingPostprocessor, self).build(input_shapes)
@@ -267,7 +267,7 @@ def call(self, inputs):
return output
-class Decoder(tf.keras.layers.Layer):
+class Decoder(tf_keras.layers.Layer):
"""The decoder network which can reuse encoder embeddings for target."""
def __init__(self, config, embedding_lookup=None, **kwargs):
@@ -284,7 +284,7 @@ def build(self, unused_input_shapes):
self.embedding_lookup = layers.OnDeviceEmbedding(
vocab_size=self.config.vocab_size,
embedding_width=self.config.hidden_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=self.config.initializer_range),
name="target_embeddings")
self.embedding_postprocessor = EmbeddingPostprocessor(
@@ -292,7 +292,7 @@ def build(self, unused_input_shapes):
use_position_embeddings=True,
max_position_embeddings=self.config.max_position_embeddings,
dropout_prob=self.config.hidden_dropout_prob,
- initializer=tf.keras.initializers.VarianceScaling(
+ initializer=tf_keras.initializers.VarianceScaling(
scale=self.config.initializer_gain,
mode="fan_avg",
distribution="uniform"),
@@ -352,7 +352,7 @@ def call(self,
if not isinstance(all_encoder_outputs, list):
all_encoder_outputs = [all_encoder_outputs]
- target_embeds = self.embedding_lookup(target_ids)
+ target_embeds = self.embedding_lookup(target_ids) # pyrefly: ignore[not-callable]
if decode_loop_step is None:
target_embeds = self.embedding_postprocessor(target_embeds)
else:
diff --git a/official/projects/nhnet/decoder_test.py b/official/projects/nhnet/decoder_test.py
index 1c0feb81abc..13ac22fe9f2 100644
--- a/official/projects/nhnet/decoder_test.py
+++ b/official/projects/nhnet/decoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for projects.nhnet.decoder."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.modeling import layers
from official.projects.nhnet import configs
from official.projects.nhnet import decoder
@@ -43,18 +43,18 @@ def test_transformer_decoder(self):
def test_bert_decoder(self):
seq_length = 10
- encoder_input_ids = tf.keras.layers.Input(
+ encoder_input_ids = tf_keras.layers.Input(
shape=(seq_length,), name="encoder_input_ids", dtype=tf.int32)
- target_ids = tf.keras.layers.Input(
+ target_ids = tf_keras.layers.Input(
shape=(seq_length,), name="target_ids", dtype=tf.int32)
- encoder_outputs = tf.keras.layers.Input(
+ encoder_outputs = tf_keras.layers.Input(
shape=(seq_length, self._config.hidden_size),
name="all_encoder_outputs",
dtype=tf.float32)
embedding_lookup = layers.OnDeviceEmbedding(
vocab_size=self._config.vocab_size,
embedding_width=self._config.hidden_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=self._config.initializer_range),
name="word_embeddings")
cross_attention_bias = decoder.AttentionBias(bias_type="single_cross")(
@@ -72,7 +72,7 @@ def test_bert_decoder(self):
encoder_input_ids=encoder_input_ids,
target_ids=target_ids,
all_encoder_outputs=encoder_outputs)
- model = tf.keras.Model(inputs=model_inputs, outputs=outputs, name="test")
+ model = tf_keras.Model(inputs=model_inputs, outputs=outputs, name="test")
self.assertLen(decoder_layer.trainable_weights, 30)
# Forward path.
fake_inputs = {
@@ -87,21 +87,21 @@ def test_multi_doc_decoder(self):
self._config = utils.get_test_params(cls=configs.NHNetConfig)
seq_length = 10
num_docs = 5
- encoder_input_ids = tf.keras.layers.Input(
+ encoder_input_ids = tf_keras.layers.Input(
shape=(num_docs, seq_length), name="encoder_input_ids", dtype=tf.int32)
- target_ids = tf.keras.layers.Input(
+ target_ids = tf_keras.layers.Input(
shape=(seq_length,), name="target_ids", dtype=tf.int32)
- encoder_outputs = tf.keras.layers.Input(
+ encoder_outputs = tf_keras.layers.Input(
shape=(num_docs, seq_length, self._config.hidden_size),
name="all_encoder_outputs",
dtype=tf.float32)
embedding_lookup = layers.OnDeviceEmbedding(
vocab_size=self._config.vocab_size,
embedding_width=self._config.hidden_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=self._config.initializer_range),
name="word_embeddings")
- doc_attention_probs = tf.keras.layers.Input(
+ doc_attention_probs = tf_keras.layers.Input(
shape=(self._config.num_decoder_attn_heads, seq_length, num_docs),
name="doc_attention_probs",
dtype=tf.float32)
@@ -124,7 +124,7 @@ def test_multi_doc_decoder(self):
target_ids=target_ids,
all_encoder_outputs=encoder_outputs,
doc_attention_probs=doc_attention_probs)
- model = tf.keras.Model(inputs=model_inputs, outputs=outputs, name="test")
+ model = tf_keras.Model(inputs=model_inputs, outputs=outputs, name="test")
self.assertLen(decoder_layer.trainable_weights, 30)
# Forward path.
fake_inputs = {
diff --git a/official/projects/nhnet/evaluation.py b/official/projects/nhnet/evaluation.py
index c762aeb5489..54d8ed76ac5 100644
--- a/official/projects/nhnet/evaluation.py
+++ b/official/projects/nhnet/evaluation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,11 +16,9 @@
import os
-# Import libraries
-
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.transformer import metrics as metrics_v2
from official.legacy.transformer.utils import metrics
@@ -103,7 +101,7 @@ def continuous_eval(strategy,
model = models.create_model(model_type, params)
metric_layer = metrics_v2.MetricLayer(params.vocab_size)
eval_summary_writer = tf.summary.create_file_writer(
- os.path.join(model_dir, "summaries/eval"))
+ os.path.join(model_dir, "summaries/eval")) # pyrefly: ignore[no-matching-overload]
global_step = tf.Variable(
0,
trainable=False,
@@ -135,10 +133,10 @@ def _test_step_fn(inputs):
return tf.nest.map_structure(strategy.experimental_local_results, outputs)
metrics_and_funcs = [
- (tf.keras.metrics.Mean("bleu", dtype=tf.float32), bleu_score),
- (tf.keras.metrics.Mean("rouge_2_fscore",
+ (tf_keras.metrics.Mean("bleu", dtype=tf.float32), bleu_score),
+ (tf_keras.metrics.Mean("rouge_2_fscore",
dtype=tf.float32), rouge_2_fscore),
- (tf.keras.metrics.Mean("rouge_l_fscore",
+ (tf_keras.metrics.Mean("rouge_l_fscore",
dtype=tf.float32), rouge_l_fscore),
]
eval_results = {}
diff --git a/official/projects/nhnet/input_pipeline.py b/official/projects/nhnet/input_pipeline.py
index 3bfe2bc5113..da4ea344033 100644
--- a/official/projects/nhnet/input_pipeline.py
+++ b/official/projects/nhnet/input_pipeline.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Input pipelines."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def decode_record(record, name_to_features):
@@ -247,4 +247,4 @@ def _dataset_fn(ctx=None):
if use_dataset_fn:
return strategy.distribute_datasets_from_function(_dataset_fn)
else:
- return strategy.experimental_distribute_dataset(_dataset_fn())
+ return strategy.experimental_distribute_dataset(_dataset_fn()) # pyrefly: ignore[missing-attribute]
diff --git a/official/projects/nhnet/models.py b/official/projects/nhnet/models.py
index 6832a96404e..2fb1f374859 100644
--- a/official/projects/nhnet/models.py
+++ b/official/projects/nhnet/models.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.modeling.hyperparams import params_dict
@@ -66,7 +66,7 @@ def remove_sos_from_seq(seq, pad_token_id):
return targets
-class Bert2Bert(tf.keras.Model):
+class Bert2Bert(tf_keras.Model):
"""Bert2Bert encoder decoder model for training."""
def __init__(self, params, bert_layer, decoder_layer, name=None):
@@ -207,7 +207,7 @@ def _init_cache(self, batch_size):
}
return cache
- def call(self, inputs, mode="train"):
+ def call(self, inputs, mode="train"): # pyrefly: ignore[bad-override]
"""Implements call().
Args:
@@ -311,7 +311,7 @@ def _symbols_to_logits_fn(ids, i, cache):
return _symbols_to_logits_fn
- def call(self, inputs, mode="training"):
+ def call(self, inputs, mode="training"): # pytype: disable=signature-mismatch # overriding-default-value-checks
input_shape = tf_utils.get_shape_list(inputs["input_ids"], expected_rank=3)
batch_size, num_docs, len_passage = (input_shape[0], input_shape[1],
input_shape[2])
@@ -390,13 +390,13 @@ def get_bert2bert_layers(params: configs.BERT2BERTConfig):
Returns:
two keras Layers, bert_model_layer and decoder_layer
"""
- input_ids = tf.keras.layers.Input(
+ input_ids = tf_keras.layers.Input(
shape=(None,), name="input_ids", dtype=tf.int32)
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(None,), name="input_mask", dtype=tf.int32)
- segment_ids = tf.keras.layers.Input(
+ segment_ids = tf_keras.layers.Input(
shape=(None,), name="segment_ids", dtype=tf.int32)
- target_ids = tf.keras.layers.Input(
+ target_ids = tf_keras.layers.Input(
shape=(None,), name="target_ids", dtype=tf.int32)
bert_config = utils.get_bert_config_from_params(params)
bert_model_layer = networks.BertEncoder(
@@ -410,7 +410,7 @@ def get_bert2bert_layers(params: configs.BERT2BERTConfig):
attention_dropout_rate=bert_config.attention_probs_dropout_prob,
max_sequence_length=bert_config.max_position_embeddings,
type_vocab_size=bert_config.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
return_all_encoder_outputs=True,
name="bert_encoder")
@@ -442,11 +442,11 @@ def get_nhnet_layers(params: configs.NHNetConfig):
Returns:
two keras Layers, bert_model_layer and decoder_layer
"""
- input_ids = tf.keras.layers.Input(
+ input_ids = tf_keras.layers.Input(
shape=(None,), name="input_ids", dtype=tf.int32)
- input_mask = tf.keras.layers.Input(
+ input_mask = tf_keras.layers.Input(
shape=(None,), name="input_mask", dtype=tf.int32)
- segment_ids = tf.keras.layers.Input(
+ segment_ids = tf_keras.layers.Input(
shape=(None,), name="segment_ids", dtype=tf.int32)
bert_config = utils.get_bert_config_from_params(params)
bert_model_layer = networks.BertEncoder(
@@ -460,19 +460,19 @@ def get_nhnet_layers(params: configs.NHNetConfig):
attention_dropout_rate=bert_config.attention_probs_dropout_prob,
max_sequence_length=bert_config.max_position_embeddings,
type_vocab_size=bert_config.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
return_all_encoder_outputs=True,
name="bert_encoder")
bert_model_layer([input_ids, input_mask, segment_ids])
- input_ids = tf.keras.layers.Input(
+ input_ids = tf_keras.layers.Input(
shape=(None, None), name="input_ids", dtype=tf.int32)
- all_encoder_outputs = tf.keras.layers.Input((None, None, params.hidden_size),
+ all_encoder_outputs = tf_keras.layers.Input((None, None, params.hidden_size),
dtype=tf.float32)
- target_ids = tf.keras.layers.Input(
+ target_ids = tf_keras.layers.Input(
shape=(None,), name="target_ids", dtype=tf.int32)
- doc_attention_probs = tf.keras.layers.Input(
+ doc_attention_probs = tf_keras.layers.Input(
(params.num_decoder_attn_heads, None, None), dtype=tf.float32)
# pylint: disable=protected-access
decoder_layer = decoder.Decoder(params, bert_model_layer._embedding_layer)
@@ -494,7 +494,7 @@ def get_nhnet_layers(params: configs.NHNetConfig):
def create_transformer_model(params,
init_checkpoint: Optional[Text] = None
- ) -> tf.keras.Model:
+ ) -> tf_keras.Model:
"""A helper to create Transformer model."""
bert_layer, decoder_layer = get_bert2bert_layers(params=params)
model = Bert2Bert(
@@ -516,7 +516,7 @@ def create_transformer_model(params,
def create_bert2bert_model(
params: configs.BERT2BERTConfig,
cls=Bert2Bert,
- init_checkpoint: Optional[Text] = None) -> tf.keras.Model:
+ init_checkpoint: Optional[Text] = None) -> tf_keras.Model:
"""A helper to create Bert2Bert model."""
bert_layer, decoder_layer = get_bert2bert_layers(params=params)
if init_checkpoint:
@@ -532,7 +532,7 @@ def create_bert2bert_model(
def create_nhnet_model(
params: configs.NHNetConfig,
cls=NHNet,
- init_checkpoint: Optional[Text] = None) -> tf.keras.Model:
+ init_checkpoint: Optional[Text] = None) -> tf_keras.Model:
"""A helper to create NHNet model."""
bert_layer, decoder_layer = get_nhnet_layers(params=params)
model = cls(
diff --git a/official/projects/nhnet/models_test.py b/official/projects/nhnet/models_test.py
index 3f487d08943..b5a40f653d4 100644
--- a/official/projects/nhnet/models_test.py
+++ b/official/projects/nhnet/models_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl import logging
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=g-direct-tensorflow-import
from tensorflow.python.distribute import combinations
@@ -224,7 +224,7 @@ def _count_params(self, layer, trainable_only=True):
else:
return int(
np.sum([
- tf.keras.backend.count_params(p) for p in layer.trainable_weights
+ tf_keras.backend.count_params(p) for p in layer.trainable_weights
]))
def test_create_nhnet_layers(self):
diff --git a/official/projects/nhnet/optimizer.py b/official/projects/nhnet/optimizer.py
index 85a9a79448d..dd2d0aad14f 100644
--- a/official/projects/nhnet/optimizer.py
+++ b/official/projects/nhnet/optimizer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,12 +14,12 @@
"""Optimizer and learning rate scheduler."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.hyperparams import params_dict
-class LearningRateSchedule(tf.keras.optimizers.schedules.LearningRateSchedule):
+class LearningRateSchedule(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Learning rate schedule."""
def __init__(self, initial_learning_rate, hidden_size, warmup_steps):
@@ -68,7 +68,7 @@ def create_optimizer(params: params_dict.ParamsDict):
"""Creates optimizer."""
lr_schedule = LearningRateSchedule(params.learning_rate, params.hidden_size,
params.learning_rate_warmup_steps)
- return tf.keras.optimizers.Adam(
+ return tf_keras.optimizers.Adam(
learning_rate=lr_schedule,
beta_1=params.adam_beta1,
beta_2=params.adam_beta2,
diff --git a/official/projects/nhnet/raw_data_process.py b/official/projects/nhnet/raw_data_process.py
index 3f5d15eab10..89f433a68cc 100644
--- a/official/projects/nhnet/raw_data_process.py
+++ b/official/projects/nhnet/raw_data_process.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/nhnet/raw_data_processor.py b/official/projects/nhnet/raw_data_processor.py
index 1e3316e8be7..237171c9361 100644
--- a/official/projects/nhnet/raw_data_processor.py
+++ b/official/projects/nhnet/raw_data_processor.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,11 +16,11 @@
import collections
import json
-import multiprocessing
+import multiprocessing.pool
import os
import urllib.parse
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.data import classifier_data_lib
from official.nlp.tools import tokenization
diff --git a/official/projects/nhnet/trainer.py b/official/projects/nhnet/trainer.py
index cbd56141ab9..037335b624b 100644
--- a/official/projects/nhnet/trainer.py
+++ b/official/projects/nhnet/trainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,13 +16,11 @@
import os
-# Import libraries
-
from absl import app
from absl import flags
from absl import logging
from six.moves import zip
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.legacy.transformer import metrics as transformer_metrics
@@ -91,7 +89,7 @@ def define_flags():
# pylint: disable=protected-access
-class Trainer(tf.keras.Model):
+class Trainer(tf_keras.Model):
"""A training only model."""
def __init__(self, model, params):
@@ -120,9 +118,13 @@ def train_step(self, inputs):
tvars = self.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
self.optimizer.apply_gradients(list(zip(grads, tvars)))
+ if isinstance(self.optimizer, tf_keras.optimizers.experimental.Optimizer):
+ learning_rate = self.optimizer.learning_rate
+ else:
+ learning_rate = self.optimizer._decayed_lr(var_dtype=tf.float32)
return {
"training_loss": loss,
- "learning_rate": self.optimizer._decayed_lr(var_dtype=tf.float32)
+ "learning_rate": learning_rate,
}
@@ -147,7 +149,7 @@ def train(params, strategy, dataset=None):
optimizer=opt,
steps_per_execution=FLAGS.steps_per_loop)
summary_dir = os.path.join(FLAGS.model_dir, "summaries")
- summary_callback = tf.keras.callbacks.TensorBoard(
+ summary_callback = tf_keras.callbacks.TensorBoard(
summary_dir, update_freq=max(100, FLAGS.steps_per_loop))
checkpoint = tf.train.Checkpoint(
model=model, optimizer=opt, global_step=opt.iterations)
@@ -179,9 +181,6 @@ def train(params, strategy, dataset=None):
def run():
"""Runs NHNet using Keras APIs."""
- if FLAGS.enable_mlir_bridge:
- tf.config.experimental.enable_mlir_bridge()
-
strategy = distribute_utils.get_distribution_strategy(
distribution_strategy=FLAGS.distribution_strategy, tpu_address=FLAGS.tpu)
if strategy:
diff --git a/official/projects/nhnet/trainer_test.py b/official/projects/nhnet/trainer_test.py
index 886c8b4cf2a..7137d3ca179 100644
--- a/official/projects/nhnet/trainer_test.py
+++ b/official/projects/nhnet/trainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl import flags
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=g-direct-tensorflow-import
from tensorflow.python.distribute import combinations
diff --git a/official/projects/nhnet/utils.py b/official/projects/nhnet/utils.py
index 23c3d571e70..19db82a0135 100644
--- a/official/projects/nhnet/utils.py
+++ b/official/projects/nhnet/utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Optional, Text
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.legacy.bert import configs
from official.modeling.hyperparams import params_dict
@@ -47,8 +47,8 @@ def encoder_common_layers(transformer_block):
def initialize_bert2bert_from_pretrained_bert(
- bert_encoder: tf.keras.layers.Layer,
- bert_decoder: tf.keras.layers.Layer,
+ bert_encoder: tf_keras.layers.Layer,
+ bert_decoder: tf_keras.layers.Layer,
init_checkpoint: Optional[Text] = None) -> None:
"""Helper function to initialze Bert2Bert from Bert pretrained checkpoint."""
ckpt = tf.train.Checkpoint(model=bert_encoder)
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/README.md b/official/projects/panoptic/README.md
similarity index 66%
rename from official/vision/beta/projects/panoptic_maskrcnn/README.md
rename to official/projects/panoptic/README.md
index f1cfefdf878..8da71ea4a57 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/README.md
+++ b/official/projects/panoptic/README.md
@@ -25,7 +25,7 @@ $ export PYTHONPATH=$(pwd)
## Preparing Dataset
```bash
-$ ./official/vision/beta/data/process_coco_panoptic.sh
+$ ./official/vision/data/process_coco_panoptic.sh
```
## Launch Training
@@ -73,7 +73,7 @@ $ python3 train.py \
--model_dir $MODEL_DIR \
--params_override=$OVERRIDES
```
-**Note**: The [PanopticSegmentationGenerator](https://github.com/tensorflow/models/blob/ac7f9e7f2d0508913947242bad3e23ef7cae5a43/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/panoptic_segmentation_generator.py#L22) layer uses dynamic shapes and hence generating panoptic masks is not supported on Cloud TPUs. Running evaluation on Cloud TPUs is not supported for the same reason. However, training is supported on both Cloud TPUs and GPUs.
+**Note**: The [PanopticSegmentationGenerator](https://github.com/tensorflow/models/blob/ac7f9e7f2d0508913947242bad3e23ef7cae5a43/official/projects/panoptic/modeling/layers/panoptic_segmentation_generator.py#L22) layer uses dynamic shapes and hence generating panoptic masks is not supported on Cloud TPUs. Running evaluation on Cloud TPUs is not supported for the same reason. However, training is supported on both Cloud TPUs and GPUs.
## Pretrained Models
### Panoptic FPN
Backbone | Schedule | Experiment name | Box mAP | Mask mAP | Overall PQ | Things PQ | Stuff PQ | Checkpoints
@@ -83,6 +83,15 @@ ResNet-50 | 3x | `panoptic_fpn_coco` | 40.64 | 36.29
**Note**: Here 1x schedule refers to ~12 epochs
+### Panoptic Deeplab
+Backbone | Experiment name | Overall PQ | Things PQ | Stuff PQ | Checkpoints
+:---------------------| :-------------------------------| ---------- | --------- | -------- | ------------:
+Dilated ResNet-50 | `panoptic_deeplab_resnet_coco` | 36.80 | 37.51 | 35.73 | [ckpt](gs://tf_model_garden/vision/panoptic/panoptic_deeplab/coco/resnet50)
+Dilated ResNet-101 | `panoptic_deeplab_resnet_coco` | 38.39 | 39.47 | 36.75 | [ckpt](gs://tf_model_garden/vision/panoptic/panoptic_deeplab/coco/resnet101)
+MobileNetV3 Large | `panoptic_deeplab_mobilenetv3_large_coco` | 30.50 | 30.10 | 31.10 | [ckpt](gs://tf_model_garden/vision/panoptic/panoptic_deeplab/coco/mobilenetv3_large)
+MobileNetV3 Small | `panoptic_deeplab_mobilenetv3_small_coco` | 25.06 | 23.46 | 27.48 | [ckpt](gs://tf_model_garden/vision/panoptic/panoptic_deeplab/coco/mobilenetv3_small)
+
+
___
## Citation
```
@@ -94,4 +103,12 @@ ___
archivePrefix={arXiv},
primaryClass={cs.CV}
}
+
+@article{Cheng2020PanopticDeepLabAS,
+ title={Panoptic-DeepLab: A Simple, Strong, and Fast Baseline for Bottom-Up Panoptic Segmentation},
+ author={Bowen Cheng and Maxwell D. Collins and Yukun Zhu and Ting Liu and Thomas S. Huang and Hartwig Adam and Liang-Chieh Chen},
+ journal={2020 IEEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR)},
+ year={2020},
+ pages={12472-12482}
+}
```
diff --git a/official/projects/panoptic/__init__.py b/official/projects/panoptic/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/panoptic/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/panoptic/configs/__init__.py b/official/projects/panoptic/configs/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/panoptic/configs/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/configs/experiments/r50fpn_1x_coco.yaml b/official/projects/panoptic/configs/experiments/r50fpn_1x_coco.yaml
similarity index 100%
rename from official/vision/beta/projects/panoptic_maskrcnn/configs/experiments/r50fpn_1x_coco.yaml
rename to official/projects/panoptic/configs/experiments/r50fpn_1x_coco.yaml
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/configs/experiments/r50fpn_3x_coco.yaml b/official/projects/panoptic/configs/experiments/r50fpn_3x_coco.yaml
similarity index 100%
rename from official/vision/beta/projects/panoptic_maskrcnn/configs/experiments/r50fpn_3x_coco.yaml
rename to official/projects/panoptic/configs/experiments/r50fpn_3x_coco.yaml
diff --git a/official/projects/panoptic/configs/panoptic_deeplab.py b/official/projects/panoptic/configs/panoptic_deeplab.py
new file mode 100644
index 00000000000..a9151ca0d80
--- /dev/null
+++ b/official/projects/panoptic/configs/panoptic_deeplab.py
@@ -0,0 +1,688 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic Deeplab configuration definition."""
+import dataclasses
+import math
+import os
+from typing import List, Optional, Union
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.modeling import optimization
+from official.vision.configs import common
+from official.vision.configs import decoders
+from official.vision.configs import backbones
+
+
+_COCO_INPUT_PATH_BASE = 'coco/tfrecords'
+_COCO_TRAIN_EXAMPLES = 118287
+_COCO_VAL_EXAMPLES = 5000
+
+
+@dataclasses.dataclass
+class Parser(hyperparams.Config):
+ """Panoptic deeplab parser."""
+ ignore_label: int = 0
+ # If resize_eval_groundtruth is set to False, original image sizes are used
+ # for eval. In that case, groundtruth_padded_size has to be specified too to
+ # allow for batching the variable input sizes of images.
+ resize_eval_groundtruth: bool = True
+ groundtruth_padded_size: List[int] = dataclasses.field(default_factory=list)
+ aug_scale_min: float = 1.0
+ aug_scale_max: float = 1.0
+ aug_rand_hflip: bool = True
+ aug_type: common.Augmentation = dataclasses.field(
+ default_factory=common.Augmentation
+ )
+ sigma: float = 8.0
+ small_instance_area_threshold: int = 4096
+ small_instance_weight: float = 3.0
+ dtype = 'float32'
+
+
+@dataclasses.dataclass
+class TfExampleDecoder(common.TfExampleDecoder):
+ """A simple TF Example decoder config."""
+ panoptic_category_mask_key: str = 'image/panoptic/category_mask'
+ panoptic_instance_mask_key: str = 'image/panoptic/instance_mask'
+
+
+@dataclasses.dataclass
+class DataDecoder(common.DataDecoder):
+ """Data decoder config."""
+ simple_decoder: TfExampleDecoder = dataclasses.field(
+ default_factory=TfExampleDecoder
+ )
+
+
+@dataclasses.dataclass
+class DataConfig(cfg.DataConfig):
+ """Input config for training."""
+ decoder: DataDecoder = dataclasses.field(default_factory=DataDecoder)
+ parser: Parser = dataclasses.field(default_factory=Parser)
+ input_path: str = ''
+ drop_remainder: bool = True
+ file_type: str = 'tfrecord'
+ is_training: bool = True
+ global_batch_size: int = 1
+
+
+@dataclasses.dataclass
+class PanopticDeeplabHead(hyperparams.Config):
+ """Panoptic Deeplab head config."""
+ level: int = 3
+ num_convs: int = 2
+ num_filters: int = 256
+ kernel_size: int = 5
+ use_depthwise_convolution: bool = False
+ upsample_factor: int = 1
+ low_level: List[int] = dataclasses.field(default_factory=lambda: [3, 2])
+ low_level_num_filters: List[int] = dataclasses.field(
+ default_factory=lambda: [64, 32])
+ fusion_num_output_filters: int = 256
+
+
+@dataclasses.dataclass
+class SemanticHead(PanopticDeeplabHead):
+ """Semantic head config."""
+ prediction_kernel_size: int = 1
+
+
+@dataclasses.dataclass
+class InstanceHead(PanopticDeeplabHead):
+ """Instance head config."""
+ prediction_kernel_size: int = 1
+
+
+@dataclasses.dataclass
+class PanopticDeeplabPostProcessor(hyperparams.Config):
+ """Panoptic Deeplab PostProcessing config."""
+ output_size: List[int] = dataclasses.field(
+ default_factory=list)
+ center_score_threshold: float = 0.1
+ thing_class_ids: List[int] = dataclasses.field(default_factory=list)
+ label_divisor: int = 256 * 256 * 256
+ stuff_area_limit: int = 4096
+ ignore_label: int = 0
+ nms_kernel: int = 7
+ keep_k_centers: int = 200
+ rescale_predictions: bool = True
+
+
+@dataclasses.dataclass
+class PanopticDeeplab(hyperparams.Config):
+ """Panoptic Deeplab model config."""
+ num_classes: int = 2
+ input_size: List[int] = dataclasses.field(default_factory=list)
+ min_level: int = 3
+ max_level: int = 6
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone(
+ type='resnet', resnet=backbones.ResNet()
+ )
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder(
+ type='aspp', aspp=decoders.ASPP(level=3)
+ )
+ )
+ semantic_head: SemanticHead = dataclasses.field(default_factory=SemanticHead)
+ instance_head: InstanceHead = dataclasses.field(default_factory=InstanceHead)
+ shared_decoder: bool = False
+ generate_panoptic_masks: bool = True
+ post_processor: PanopticDeeplabPostProcessor = dataclasses.field(
+ default_factory=PanopticDeeplabPostProcessor
+ )
+
+
+@dataclasses.dataclass
+class Losses(hyperparams.Config):
+ label_smoothing: float = 0.0
+ ignore_label: int = 0
+ class_weights: List[float] = dataclasses.field(default_factory=list)
+ l2_weight_decay: float = 1e-4
+ top_k_percent_pixels: float = 0.15
+ segmentation_loss_weight: float = 1.0
+ center_heatmap_loss_weight: float = 200
+ center_offset_loss_weight: float = 0.01
+
+
+@dataclasses.dataclass
+class Evaluation(hyperparams.Config):
+ """Evaluation config."""
+ ignored_label: int = 0
+ max_instances_per_category: int = 256
+ offset: int = 256 * 256 * 256
+ is_thing: List[float] = dataclasses.field(
+ default_factory=list)
+ rescale_predictions: bool = True
+ report_per_class_pq: bool = False
+
+ report_per_class_iou: bool = False
+ report_train_mean_iou: bool = True # Turning this off can speed up training.
+
+
+@dataclasses.dataclass
+class PanopticDeeplabTask(cfg.TaskConfig):
+ """Panoptic deeplab task config."""
+ model: PanopticDeeplab = dataclasses.field(default_factory=PanopticDeeplab)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ is_training=False, drop_remainder=False
+ )
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+ init_checkpoint_modules: Union[
+ str, List[str]] = 'all' # all, backbone, and/or decoder
+ evaluation: Evaluation = dataclasses.field(default_factory=Evaluation)
+
+
+@exp_factory.register_config_factory('panoptic_deeplab_resnet_coco')
+def panoptic_deeplab_resnet_coco() -> cfg.ExperimentConfig:
+ """COCO panoptic segmentation with Panoptic Deeplab."""
+ train_steps = 200000
+ train_batch_size = 64
+ eval_batch_size = 1
+ steps_per_epoch = _COCO_TRAIN_EXAMPLES // train_batch_size
+ validation_steps = _COCO_VAL_EXAMPLES // eval_batch_size
+
+ num_panoptic_categories = 201
+ num_thing_categories = 91
+ ignore_label = 0
+
+ is_thing = [False]
+ for idx in range(1, num_panoptic_categories):
+ is_thing.append(True if idx <= num_thing_categories else False)
+
+ input_size = [640, 640, 3]
+ output_stride = 16
+ aspp_dilation_rates = [6, 12, 18]
+ multigrid = [1, 2, 4]
+ stem_type = 'v1'
+ level = int(math.log2(output_stride))
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(
+ mixed_precision_dtype='bfloat16', enable_xla=True),
+ task=PanopticDeeplabTask(
+ init_checkpoint='gs://tf_model_garden/vision/panoptic/panoptic_deeplab/imagenet/resnet50_v1/ckpt-436800', # pylint: disable=line-too-long
+ init_checkpoint_modules=['backbone'],
+ model=PanopticDeeplab(
+ num_classes=num_panoptic_categories,
+ input_size=input_size,
+ backbone=backbones.Backbone(
+ type='dilated_resnet', dilated_resnet=backbones.DilatedResNet(
+ model_id=50,
+ stem_type=stem_type,
+ output_stride=output_stride,
+ multigrid=multigrid,
+ se_ratio=0.25,
+ last_stage_repeats=1,
+ stochastic_depth_drop_rate=0.2)),
+ decoder=decoders.Decoder(
+ type='aspp',
+ aspp=decoders.ASPP(
+ level=level,
+ num_filters=256,
+ pool_kernel_size=input_size[:2],
+ dilation_rates=aspp_dilation_rates,
+ use_depthwise_convolution=True,
+ dropout_rate=0.1)),
+ semantic_head=SemanticHead(
+ level=level,
+ num_convs=1,
+ num_filters=256,
+ kernel_size=5,
+ use_depthwise_convolution=True,
+ upsample_factor=1,
+ low_level=[3, 2],
+ low_level_num_filters=[64, 32],
+ fusion_num_output_filters=256,
+ prediction_kernel_size=1),
+ instance_head=InstanceHead(
+ level=level,
+ num_convs=1,
+ num_filters=32,
+ kernel_size=5,
+ use_depthwise_convolution=True,
+ upsample_factor=1,
+ low_level=[3, 2],
+ low_level_num_filters=[32, 16],
+ fusion_num_output_filters=128,
+ prediction_kernel_size=1),
+ shared_decoder=False,
+ generate_panoptic_masks=True,
+ post_processor=PanopticDeeplabPostProcessor(
+ output_size=input_size[:2],
+ center_score_threshold=0.1,
+ thing_class_ids=list(range(1, num_thing_categories)),
+ label_divisor=256,
+ stuff_area_limit=4096,
+ ignore_label=ignore_label,
+ nms_kernel=41,
+ keep_k_centers=200,
+ rescale_predictions=True)),
+ losses=Losses(
+ label_smoothing=0.0,
+ ignore_label=ignore_label,
+ l2_weight_decay=0.0,
+ top_k_percent_pixels=0.2,
+ segmentation_loss_weight=1.0,
+ center_heatmap_loss_weight=200,
+ center_offset_loss_weight=0.01),
+ train_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ parser=Parser(
+ aug_scale_min=0.5,
+ aug_scale_max=1.5,
+ aug_rand_hflip=True,
+ aug_type=common.Augmentation(
+ type='autoaug',
+ autoaug=common.AutoAugment(
+ augmentation_name='panoptic_deeplab_policy')),
+ sigma=8.0,
+ small_instance_area_threshold=4096,
+ small_instance_weight=3.0)),
+ validation_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ parser=Parser(
+ resize_eval_groundtruth=False,
+ groundtruth_padded_size=[640, 640],
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ aug_rand_hflip=False,
+ aug_type=None,
+ sigma=8.0,
+ small_instance_area_threshold=4096,
+ small_instance_weight=3.0),
+ drop_remainder=False),
+ evaluation=Evaluation(
+ ignored_label=ignore_label,
+ max_instances_per_category=256,
+ offset=256*256*256,
+ is_thing=is_thing, # pyrefly: ignore[bad-argument-type]
+ rescale_predictions=True,
+ report_per_class_pq=False,
+ report_per_class_iou=False,
+ report_train_mean_iou=False)),
+ trainer=cfg.TrainerConfig(
+ train_steps=train_steps,
+ validation_steps=validation_steps,
+ validation_interval=steps_per_epoch,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adam',
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 0.0005,
+ 'decay_steps': train_steps,
+ 'end_learning_rate': 0.0,
+ 'power': 0.9
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 2000,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
+
+
+@exp_factory.register_config_factory('panoptic_deeplab_mobilenetv3_large_coco')
+def panoptic_deeplab_mobilenetv3_large_coco() -> cfg.ExperimentConfig:
+ """COCO panoptic segmentation with Panoptic Deeplab."""
+ train_steps = 200000
+ train_batch_size = 64
+ eval_batch_size = 1
+ steps_per_epoch = _COCO_TRAIN_EXAMPLES // train_batch_size
+ validation_steps = _COCO_VAL_EXAMPLES // eval_batch_size
+
+ num_panoptic_categories = 201
+ num_thing_categories = 91
+ ignore_label = 0
+
+ is_thing = [False]
+ for idx in range(1, num_panoptic_categories):
+ is_thing.append(True if idx <= num_thing_categories else False)
+
+ input_size = [640, 640, 3]
+ output_stride = 16
+ aspp_dilation_rates = [6, 12, 18]
+ level = int(math.log2(output_stride))
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(
+ mixed_precision_dtype='float32', enable_xla=True),
+ task=PanopticDeeplabTask(
+ init_checkpoint='gs://tf_model_garden/vision/panoptic/panoptic_deeplab/imagenet/mobilenetv3_large/ckpt-156000',
+ init_checkpoint_modules=['backbone'],
+ model=PanopticDeeplab(
+ num_classes=num_panoptic_categories,
+ input_size=input_size,
+ backbone=backbones.Backbone(
+ type='mobilenet', mobilenet=backbones.MobileNet(
+ model_id='MobileNetV3Large',
+ filter_size_scale=1.0,
+ stochastic_depth_drop_rate=0.0,
+ output_stride=output_stride)),
+ decoder=decoders.Decoder(
+ type='aspp',
+ aspp=decoders.ASPP(
+ level=level,
+ num_filters=256,
+ pool_kernel_size=input_size[:2],
+ dilation_rates=aspp_dilation_rates,
+ use_depthwise_convolution=True,
+ dropout_rate=0.1)),
+ semantic_head=SemanticHead(
+ level=level,
+ num_convs=1,
+ num_filters=256,
+ kernel_size=5,
+ use_depthwise_convolution=True,
+ upsample_factor=1,
+ low_level=[3, 2],
+ low_level_num_filters=[64, 32],
+ fusion_num_output_filters=256,
+ prediction_kernel_size=1),
+ instance_head=InstanceHead(
+ level=level,
+ num_convs=1,
+ num_filters=32,
+ kernel_size=5,
+ use_depthwise_convolution=True,
+ upsample_factor=1,
+ low_level=[3, 2],
+ low_level_num_filters=[32, 16],
+ fusion_num_output_filters=128,
+ prediction_kernel_size=1),
+ shared_decoder=False,
+ generate_panoptic_masks=True,
+ post_processor=PanopticDeeplabPostProcessor(
+ output_size=input_size[:2],
+ center_score_threshold=0.1,
+ thing_class_ids=list(range(1, num_thing_categories)),
+ label_divisor=256,
+ stuff_area_limit=4096,
+ ignore_label=ignore_label,
+ nms_kernel=41,
+ keep_k_centers=200,
+ rescale_predictions=True)),
+ losses=Losses(
+ label_smoothing=0.0,
+ ignore_label=ignore_label,
+ l2_weight_decay=0.0,
+ top_k_percent_pixels=0.2,
+ segmentation_loss_weight=1.0,
+ center_heatmap_loss_weight=200,
+ center_offset_loss_weight=0.01),
+ train_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ parser=Parser(
+ aug_scale_min=0.5,
+ aug_scale_max=2.0,
+ aug_rand_hflip=True,
+ aug_type=common.Augmentation(
+ type='autoaug',
+ autoaug=common.AutoAugment(
+ augmentation_name='panoptic_deeplab_policy')),
+ sigma=8.0,
+ small_instance_area_threshold=4096,
+ small_instance_weight=3.0)),
+ validation_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ parser=Parser(
+ resize_eval_groundtruth=False,
+ groundtruth_padded_size=[640, 640],
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ aug_rand_hflip=False,
+ aug_type=None,
+ sigma=8.0,
+ small_instance_area_threshold=4096,
+ small_instance_weight=3.0),
+ drop_remainder=False),
+ evaluation=Evaluation(
+ ignored_label=ignore_label,
+ max_instances_per_category=256,
+ offset=256*256*256,
+ is_thing=is_thing, # pyrefly: ignore[bad-argument-type]
+ rescale_predictions=True,
+ report_per_class_pq=False,
+ report_per_class_iou=False,
+ report_train_mean_iou=False)),
+ trainer=cfg.TrainerConfig(
+ train_steps=train_steps,
+ validation_steps=validation_steps,
+ validation_interval=steps_per_epoch,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adam',
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 0.001,
+ 'decay_steps': train_steps,
+ 'end_learning_rate': 0.0,
+ 'power': 0.9
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 2000,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
+
+
+@exp_factory.register_config_factory('panoptic_deeplab_mobilenetv3_small_coco')
+def panoptic_deeplab_mobilenetv3_small_coco() -> cfg.ExperimentConfig:
+ """COCO panoptic segmentation with Panoptic Deeplab."""
+ train_steps = 200000
+ train_batch_size = 64
+ eval_batch_size = 1
+ steps_per_epoch = _COCO_TRAIN_EXAMPLES // train_batch_size
+ validation_steps = _COCO_VAL_EXAMPLES // eval_batch_size
+
+ num_panoptic_categories = 201
+ num_thing_categories = 91
+ ignore_label = 0
+
+ is_thing = [False]
+ for idx in range(1, num_panoptic_categories):
+ is_thing.append(True if idx <= num_thing_categories else False)
+
+ input_size = [640, 640, 3]
+ output_stride = 16
+ aspp_dilation_rates = [6, 12, 18]
+ level = int(math.log2(output_stride))
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(
+ mixed_precision_dtype='float32', enable_xla=True),
+ task=PanopticDeeplabTask(
+ init_checkpoint='gs://tf_model_garden/vision/panoptic/panoptic_deeplab/imagenet/mobilenetv3_small/ckpt-312000',
+ init_checkpoint_modules=['backbone'],
+ model=PanopticDeeplab(
+ num_classes=num_panoptic_categories,
+ input_size=input_size,
+ backbone=backbones.Backbone(
+ type='mobilenet', mobilenet=backbones.MobileNet(
+ model_id='MobileNetV3Small',
+ filter_size_scale=1.0,
+ stochastic_depth_drop_rate=0.0,
+ output_stride=output_stride)),
+ decoder=decoders.Decoder(
+ type='aspp',
+ aspp=decoders.ASPP(
+ level=level,
+ num_filters=256,
+ pool_kernel_size=input_size[:2],
+ dilation_rates=aspp_dilation_rates,
+ use_depthwise_convolution=True,
+ dropout_rate=0.1)),
+ semantic_head=SemanticHead(
+ level=level,
+ num_convs=1,
+ num_filters=256,
+ kernel_size=5,
+ use_depthwise_convolution=True,
+ upsample_factor=1,
+ low_level=[3, 2],
+ low_level_num_filters=[64, 32],
+ fusion_num_output_filters=256,
+ prediction_kernel_size=1),
+ instance_head=InstanceHead(
+ level=level,
+ num_convs=1,
+ num_filters=32,
+ kernel_size=5,
+ use_depthwise_convolution=True,
+ upsample_factor=1,
+ low_level=[3, 2],
+ low_level_num_filters=[32, 16],
+ fusion_num_output_filters=128,
+ prediction_kernel_size=1),
+ shared_decoder=False,
+ generate_panoptic_masks=True,
+ post_processor=PanopticDeeplabPostProcessor(
+ output_size=input_size[:2],
+ center_score_threshold=0.1,
+ thing_class_ids=list(range(1, num_thing_categories)),
+ label_divisor=256,
+ stuff_area_limit=4096,
+ ignore_label=ignore_label,
+ nms_kernel=41,
+ keep_k_centers=200,
+ rescale_predictions=True)),
+ losses=Losses(
+ label_smoothing=0.0,
+ ignore_label=ignore_label,
+ l2_weight_decay=0.0,
+ top_k_percent_pixels=0.2,
+ segmentation_loss_weight=1.0,
+ center_heatmap_loss_weight=200,
+ center_offset_loss_weight=0.01),
+ train_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ parser=Parser(
+ aug_scale_min=0.5,
+ aug_scale_max=2.0,
+ aug_rand_hflip=True,
+ aug_type=common.Augmentation(
+ type='autoaug',
+ autoaug=common.AutoAugment(
+ augmentation_name='panoptic_deeplab_policy')),
+ sigma=8.0,
+ small_instance_area_threshold=4096,
+ small_instance_weight=3.0)),
+ validation_data=DataConfig(
+ input_path=os.path.join(_COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ parser=Parser(
+ resize_eval_groundtruth=False,
+ groundtruth_padded_size=[640, 640],
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ aug_rand_hflip=False,
+ aug_type=None,
+ sigma=8.0,
+ small_instance_area_threshold=4096,
+ small_instance_weight=3.0),
+ drop_remainder=False),
+ evaluation=Evaluation(
+ ignored_label=ignore_label,
+ max_instances_per_category=256,
+ offset=256*256*256,
+ is_thing=is_thing, # pyrefly: ignore[bad-argument-type]
+ rescale_predictions=True,
+ report_per_class_pq=False,
+ report_per_class_iou=False,
+ report_train_mean_iou=False)),
+ trainer=cfg.TrainerConfig(
+ train_steps=train_steps,
+ validation_steps=validation_steps,
+ validation_interval=steps_per_epoch,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adam',
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 0.001,
+ 'decay_steps': train_steps,
+ 'end_learning_rate': 0.0,
+ 'power': 0.9
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 2000,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/configs/panoptic_maskrcnn.py b/official/projects/panoptic/configs/panoptic_maskrcnn.py
similarity index 81%
rename from official/vision/beta/projects/panoptic_maskrcnn/configs/panoptic_maskrcnn.py
rename to official/projects/panoptic/configs/panoptic_maskrcnn.py
index 0d98b9a15b1..7973933b25e 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/configs/panoptic_maskrcnn.py
+++ b/official/projects/panoptic/configs/panoptic_maskrcnn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,9 +23,11 @@
from official.modeling import hyperparams
from official.modeling import optimization
from official.projects.deepmac_maskrcnn.configs import deep_mask_head_rcnn as deepmac_maskrcnn
+from official.projects.uvit.configs import backbones as uvit_backbones
from official.vision.configs import common
from official.vision.configs import maskrcnn
from official.vision.configs import semantic_segmentation
+from official.vision.configs.google import backbones
SEGMENTATION_MODEL = semantic_segmentation.SemanticSegmentationModel
@@ -36,6 +38,7 @@
_COCO_VAL_EXAMPLES = 5000
# pytype: disable=wrong-keyword-args
+# pylint: disable=unexpected-keyword-arg
@dataclasses.dataclass
@@ -66,14 +69,16 @@ class TfExampleDecoder(common.TfExampleDecoder):
@dataclasses.dataclass
class DataDecoder(common.DataDecoder):
"""Data decoder config."""
- simple_decoder: TfExampleDecoder = TfExampleDecoder()
+ simple_decoder: TfExampleDecoder = dataclasses.field(
+ default_factory=TfExampleDecoder
+ )
@dataclasses.dataclass
class DataConfig(maskrcnn.DataConfig):
"""Input config for training."""
- decoder: DataDecoder = DataDecoder()
- parser: Parser = Parser()
+ decoder: DataDecoder = dataclasses.field(default_factory=DataDecoder)
+ parser: Parser = dataclasses.field(default_factory=Parser)
@dataclasses.dataclass
@@ -91,17 +96,37 @@ class PanopticSegmentationGenerator(hyperparams.Config):
rescale_predictions: bool = False
+@dataclasses.dataclass
+class Backbone(backbones.Backbone):
+ """Configuration for backbones.
+
+ Attributes:
+ type: "str", type of backbone be used, one the of fields below.
+ uvit: uvit backbone config.
+ """
+ type: Optional[str] = None
+ uvit: uvit_backbones.VisionTransformer = dataclasses.field(
+ default_factory=uvit_backbones.VisionTransformer
+ )
+
+
@dataclasses.dataclass
class PanopticMaskRCNN(deepmac_maskrcnn.DeepMaskHeadRCNN):
"""Panoptic Mask R-CNN model config."""
- segmentation_model: semantic_segmentation.SemanticSegmentationModel = (
- SEGMENTATION_MODEL(num_classes=2))
- include_mask = True
+ backbone: Backbone = dataclasses.field(
+ default_factory=lambda: Backbone(type='resnet', resnet=backbones.ResNet())
+ )
+ segmentation_model: SEGMENTATION_MODEL = dataclasses.field(
+ default_factory=lambda: SEGMENTATION_MODEL(num_classes=2)
+ )
+ include_mask: bool = True
shared_backbone: bool = True
shared_decoder: bool = True
stuff_classes_offset: int = 0
generate_panoptic_masks: bool = True
- panoptic_segmentation_generator: PanopticSegmentationGenerator = PanopticSegmentationGenerator() # pylint:disable=line-too-long
+ panoptic_segmentation_generator: PanopticSegmentationGenerator = (
+ dataclasses.field(default_factory=PanopticSegmentationGenerator)
+ )
@dataclasses.dataclass
@@ -109,9 +134,13 @@ class Losses(maskrcnn.Losses):
"""Panoptic Mask R-CNN loss config."""
semantic_segmentation_label_smoothing: float = 0.0
semantic_segmentation_ignore_label: int = 255
+ semantic_segmentation_gt_is_matting_map: bool = False
semantic_segmentation_class_weights: List[float] = dataclasses.field(
default_factory=list)
semantic_segmentation_use_groundtruth_dimension: bool = True
+ # If true, use binary cross entropy (sigmoid) in loss, otherwise, use
+ # categorical cross entropy (softmax).
+ semantic_segmentation_use_binary_cross_entropy: bool = False
semantic_segmentation_top_k_percent_pixels: float = 1.0
instance_segmentation_weight: float = 1.0
semantic_segmentation_weight: float = 0.5
@@ -133,12 +162,19 @@ class PanopticQualityEvaluator(hyperparams.Config):
@dataclasses.dataclass
class PanopticMaskRCNNTask(maskrcnn.MaskRCNNTask):
"""Panoptic Mask R-CNN task config."""
- model: PanopticMaskRCNN = PanopticMaskRCNN()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False,
- drop_remainder=False)
- segmentation_evaluation: semantic_segmentation.Evaluation = semantic_segmentation.Evaluation() # pylint: disable=line-too-long
- losses: Losses = Losses()
+ model: PanopticMaskRCNN = dataclasses.field(default_factory=PanopticMaskRCNN)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ is_training=False, drop_remainder=False
+ )
+ )
+ segmentation_evaluation: semantic_segmentation.Evaluation = dataclasses.field(
+ default_factory=semantic_segmentation.Evaluation
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
init_checkpoint: Optional[str] = None
segmentation_init_checkpoint: Optional[str] = None
@@ -151,7 +187,9 @@ class PanopticMaskRCNNTask(maskrcnn.MaskRCNNTask):
# 'all': Initialize all modules
init_checkpoint_modules: Optional[List[str]] = dataclasses.field(
default_factory=list)
- panoptic_quality_evaluator: PanopticQualityEvaluator = PanopticQualityEvaluator() # pylint: disable=line-too-long
+ panoptic_quality_evaluator: PanopticQualityEvaluator = dataclasses.field(
+ default_factory=PanopticQualityEvaluator
+ )
@exp_factory.register_config_factory('panoptic_fpn_coco')
diff --git a/official/projects/panoptic/dataloaders/panoptic_deeplab_input.py b/official/projects/panoptic/dataloaders/panoptic_deeplab_input.py
new file mode 100644
index 00000000000..d885946b6c6
--- /dev/null
+++ b/official/projects/panoptic/dataloaders/panoptic_deeplab_input.py
@@ -0,0 +1,359 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data parser and processing for Panoptic Deeplab."""
+
+from typing import List, Optional
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.vision.configs import common
+from official.vision.dataloaders import parser
+from official.vision.dataloaders import tf_example_decoder
+from official.vision.ops import augment
+from official.vision.ops import preprocess_ops
+
+
+def _compute_gaussian_from_std(sigma):
+ """Computes the Gaussian and its size from a given standard deviation."""
+ size = int(6 * sigma + 3)
+ x = np.arange(size, dtype=float)
+ y = x[:, np.newaxis]
+ x0, y0 = 3 * sigma + 1, 3 * sigma + 1
+ gaussian = tf.constant(
+ np.exp(-((x - x0)**2 + (y - y0)**2) / (2 * sigma**2)),
+ dtype=tf.float32)
+ return gaussian, size
+
+
+class TfExampleDecoder(tf_example_decoder.TfExampleDecoder):
+ """Tensorflow Example proto decoder."""
+
+ def __init__(
+ self,
+ regenerate_source_id: bool,
+ panoptic_category_mask_key: str = 'image/panoptic/category_mask',
+ panoptic_instance_mask_key: str = 'image/panoptic/instance_mask'):
+ super(TfExampleDecoder,
+ self).__init__(
+ include_mask=True,
+ regenerate_source_id=regenerate_source_id)
+ self._panoptic_category_mask_key = panoptic_category_mask_key
+ self._panoptic_instance_mask_key = panoptic_instance_mask_key
+
+ self._panoptic_keys_to_features = {
+ panoptic_category_mask_key:
+ tf.io.FixedLenFeature((), tf.string, default_value=''),
+ panoptic_instance_mask_key:
+ tf.io.FixedLenFeature((), tf.string, default_value='')
+ }
+
+ def decode(self, serialized_example):
+ decoded_tensors = super(TfExampleDecoder,
+ self).decode(serialized_example)
+ parsed_tensors = tf.io.parse_single_example(
+ serialized_example, self._panoptic_keys_to_features)
+
+ category_mask = tf.io.decode_image(
+ parsed_tensors[self._panoptic_category_mask_key], channels=1)
+ instance_mask = tf.io.decode_image(
+ parsed_tensors[self._panoptic_instance_mask_key], channels=1)
+ category_mask.set_shape([None, None, 1])
+ instance_mask.set_shape([None, None, 1])
+
+ decoded_tensors.update({
+ 'groundtruth_panoptic_category_mask': category_mask,
+ 'groundtruth_panoptic_instance_mask': instance_mask
+ })
+ return decoded_tensors
+
+
+class Parser(parser.Parser):
+ """Parser to parse an image and its annotations into a dictionary of tensors."""
+
+ def __init__(
+ self,
+ output_size: List[int],
+ resize_eval_groundtruth: bool = True,
+ groundtruth_padded_size: Optional[List[int]] = None,
+ ignore_label: int = 0,
+ aug_rand_hflip: bool = False,
+ aug_scale_min: float = 1.0,
+ aug_scale_max: float = 1.0,
+ aug_type: Optional[common.Augmentation] = None,
+ sigma: float = 8.0,
+ small_instance_area_threshold: int = 4096,
+ small_instance_weight: float = 3.0,
+ dtype: str = 'float32'):
+ """Initializes parameters for parsing annotations in the dataset.
+
+ Args:
+ output_size: `Tensor` or `list` for [height, width] of output image. The
+ output_size should be divided by the largest feature stride 2^max_level.
+ resize_eval_groundtruth: `bool`, if True, eval groundtruth masks are
+ resized to output_size.
+ groundtruth_padded_size: `Tensor` or `list` for [height, width]. When
+ resize_eval_groundtruth is set to False, the groundtruth masks are
+ padded to this size.
+ ignore_label: `int` the pixel with ignore label will not used for training
+ and evaluation.
+ aug_rand_hflip: `bool`, if True, augment training with random
+ horizontal flip.
+ aug_scale_min: `float`, the minimum scale applied to `output_size` for
+ data augmentation during training.
+ aug_scale_max: `float`, the maximum scale applied to `output_size` for
+ data augmentation during training.
+ aug_type: An optional Augmentation object with params for AutoAugment.
+ sigma: `float`, standard deviation for generating 2D Gaussian to encode
+ centers.
+ small_instance_area_threshold: `int`, small instance area threshold.
+ small_instance_weight: `float`, small instance weight.
+ dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
+ """
+ self._output_size = output_size
+ self._resize_eval_groundtruth = resize_eval_groundtruth
+ if (not resize_eval_groundtruth) and (groundtruth_padded_size is None):
+ raise ValueError(
+ 'groundtruth_padded_size ([height, width]) needs to be'
+ 'specified when resize_eval_groundtruth is False.')
+ self._groundtruth_padded_size = groundtruth_padded_size
+ self._ignore_label = ignore_label
+
+ # Data augmentation.
+ self._aug_rand_hflip = aug_rand_hflip
+ self._aug_scale_min = aug_scale_min
+ self._aug_scale_max = aug_scale_max
+
+ if aug_type and aug_type.type:
+ if aug_type.type == 'autoaug':
+ self._augmenter = augment.AutoAugment(
+ augmentation_name=aug_type.autoaug.augmentation_name,
+ cutout_const=aug_type.autoaug.cutout_const,
+ translate_const=aug_type.autoaug.translate_const)
+ else:
+ raise ValueError('Augmentation policy {} not supported.'.format(
+ aug_type.type))
+ else:
+ self._augmenter = None
+
+ self._dtype = dtype
+
+ self._sigma = sigma
+ self._gaussian, self._gaussian_size = _compute_gaussian_from_std(
+ self._sigma)
+ self._gaussian = tf.reshape(self._gaussian, shape=[-1])
+ self._small_instance_area_threshold = small_instance_area_threshold
+ self._small_instance_weight = small_instance_weight
+
+ def _resize_and_crop_mask(self, mask, image_info, is_training):
+ """Resizes and crops mask using `image_info` dict."""
+ height = image_info[0][0]
+ width = image_info[0][1]
+ mask = tf.reshape(mask, shape=[1, height, width, 1])
+ mask += 1
+
+ if is_training or self._resize_eval_groundtruth:
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ mask = preprocess_ops.resize_and_crop_masks(
+ mask,
+ image_scale,
+ self._output_size,
+ offset)
+ else:
+ mask = tf.image.pad_to_bounding_box(
+ mask, 0, 0,
+ self._groundtruth_padded_size[0], # pyrefly: ignore[unsupported-operation]
+ self._groundtruth_padded_size[1]) # pyrefly: ignore[unsupported-operation]
+ mask -= 1
+
+ # Assign ignore label to the padded region.
+ mask = tf.where(
+ tf.equal(mask, -1),
+ self._ignore_label * tf.ones_like(mask),
+ mask)
+ mask = tf.squeeze(mask, axis=0)
+ return mask
+
+ def _parse_data(self, data, is_training):
+ image = data['image']
+
+ if self._augmenter is not None and is_training:
+ image = self._augmenter.distort(image)
+
+ image = preprocess_ops.normalize_image(image)
+
+ category_mask = tf.cast(
+ data['groundtruth_panoptic_category_mask'][:, :, 0],
+ dtype=tf.float32)
+ instance_mask = tf.cast(
+ data['groundtruth_panoptic_instance_mask'][:, :, 0],
+ dtype=tf.float32)
+
+ # Flips image randomly during training.
+ if self._aug_rand_hflip and is_training:
+ masks = tf.stack([category_mask, instance_mask], axis=0)
+ image, _, masks = preprocess_ops.random_horizontal_flip(
+ image=image, masks=masks)
+ category_mask = masks[0]
+ instance_mask = masks[1]
+
+ # Resizes and crops image.
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ self._output_size,
+ self._output_size,
+ aug_scale_min=self._aug_scale_min if is_training else 1.0,
+ aug_scale_max=self._aug_scale_max if is_training else 1.0)
+
+ category_mask = self._resize_and_crop_mask(
+ category_mask,
+ image_info,
+ is_training=is_training)
+ instance_mask = self._resize_and_crop_mask(
+ instance_mask,
+ image_info,
+ is_training=is_training)
+
+ (instance_centers_heatmap,
+ instance_centers_offset,
+ semantic_weights) = self._encode_centers_and_offets(
+ instance_mask=instance_mask[:, :, 0])
+
+ # Cast image and labels as self._dtype
+ image = tf.cast(image, dtype=self._dtype)
+ category_mask = tf.cast(category_mask, dtype=self._dtype)
+ instance_mask = tf.cast(instance_mask, dtype=self._dtype)
+ instance_centers_heatmap = tf.cast(
+ instance_centers_heatmap, dtype=self._dtype)
+ instance_centers_offset = tf.cast(
+ instance_centers_offset, dtype=self._dtype)
+
+ valid_mask = tf.not_equal(
+ category_mask, self._ignore_label)
+ things_mask = tf.not_equal(
+ instance_mask, self._ignore_label)
+
+ labels = {
+ 'category_mask': category_mask,
+ 'instance_mask': instance_mask,
+ 'instance_centers_heatmap': instance_centers_heatmap,
+ 'instance_centers_offset': instance_centers_offset,
+ 'semantic_weights': semantic_weights,
+ 'valid_mask': valid_mask,
+ 'things_mask': things_mask,
+ 'image_info': image_info
+ }
+ return image, labels
+
+ def _parse_train_data(self, data):
+ """Parses data for training."""
+ return self._parse_data(data=data, is_training=True)
+
+ def _parse_eval_data(self, data):
+ """Parses data for evaluation."""
+ return self._parse_data(data=data, is_training=False)
+
+ def _encode_centers_and_offets(self, instance_mask):
+ """Generates center heatmaps and offets from instance id mask.
+
+ Args:
+ instance_mask: `tf.Tensor` of shape [height, width] representing
+ groundtruth instance id mask.
+ Returns:
+ instance_centers_heatmap: `tf.Tensor` of shape [height, width, 1]
+ instance_centers_offset: `tf.Tensor` of shape [height, width, 2]
+ """
+ shape = tf.shape(instance_mask)
+ height, width = shape[0], shape[1]
+
+ padding_start = int(3 * self._sigma + 1)
+ padding_end = int(3 * self._sigma + 2)
+
+ # padding should be equal to self._gaussian_size which is calculated
+ # as size = int(6 * sigma + 3)
+ padding = padding_start + padding_end
+
+ instance_centers_heatmap = tf.zeros(
+ shape=[height + padding, width + padding],
+ dtype=tf.float32)
+ centers_offset_y = tf.zeros(
+ shape=[height, width],
+ dtype=tf.float32)
+ centers_offset_x = tf.zeros(
+ shape=[height, width],
+ dtype=tf.float32)
+ semantic_weights = tf.ones(
+ shape=[height, width],
+ dtype=tf.float32)
+
+ unique_instance_ids, _ = tf.unique(tf.reshape(instance_mask, [-1]))
+
+ # The following method for encoding center heatmaps and offets is inspired
+ # by the reference implementation available at
+ # https://github.com/google-research/deeplab2/blob/main/data/sample_generator.py # pylint: disable=line-too-long
+ for instance_id in unique_instance_ids:
+ if instance_id == self._ignore_label:
+ continue
+
+ mask = tf.equal(instance_mask, instance_id)
+ mask_area = tf.reduce_sum(tf.cast(mask, dtype=tf.float32))
+ mask_indices = tf.cast(tf.where(mask), dtype=tf.float32)
+ mask_center = tf.reduce_mean(mask_indices, axis=0)
+ mask_center_y = tf.cast(tf.round(mask_center[0]), dtype=tf.int32)
+ mask_center_x = tf.cast(tf.round(mask_center[1]), dtype=tf.int32)
+
+ if mask_area < self._small_instance_area_threshold:
+ semantic_weights = tf.where(
+ mask,
+ self._small_instance_weight,
+ semantic_weights)
+
+ gaussian_size = self._gaussian_size
+ indices_y = tf.range(mask_center_y, mask_center_y + gaussian_size)
+ indices_x = tf.range(mask_center_x, mask_center_x + gaussian_size)
+
+ indices = tf.stack(tf.meshgrid(indices_y, indices_x))
+ indices = tf.reshape(
+ indices, shape=[2, gaussian_size * gaussian_size])
+ indices = tf.transpose(indices)
+
+ instance_centers_heatmap = tf.tensor_scatter_nd_max(
+ tensor=instance_centers_heatmap,
+ indices=indices,
+ updates=self._gaussian)
+
+ centers_offset_y = tf.tensor_scatter_nd_update(
+ tensor=centers_offset_y,
+ indices=tf.cast(mask_indices, dtype=tf.int32),
+ updates=tf.cast(mask_center_y, dtype=tf.float32) - mask_indices[:, 0])
+
+ centers_offset_x = tf.tensor_scatter_nd_update(
+ tensor=centers_offset_x,
+ indices=tf.cast(mask_indices, dtype=tf.int32),
+ updates=tf.cast(mask_center_x, dtype=tf.float32) - mask_indices[:, 1])
+
+ instance_centers_heatmap = instance_centers_heatmap[
+ padding_start:padding_start + height,
+ padding_start:padding_start + width]
+ instance_centers_heatmap = tf.expand_dims(instance_centers_heatmap, axis=-1)
+
+ instance_centers_offset = tf.stack(
+ [centers_offset_y, centers_offset_x],
+ axis=-1)
+
+ return (instance_centers_heatmap,
+ instance_centers_offset,
+ semantic_weights)
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/dataloaders/panoptic_maskrcnn_input.py b/official/projects/panoptic/dataloaders/panoptic_maskrcnn_input.py
similarity index 81%
rename from official/vision/beta/projects/panoptic_maskrcnn/dataloaders/panoptic_maskrcnn_input.py
rename to official/projects/panoptic/dataloaders/panoptic_maskrcnn_input.py
index 027cb1dfc2d..61312ca231b 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/dataloaders/panoptic_maskrcnn_input.py
+++ b/official/projects/panoptic/dataloaders/panoptic_maskrcnn_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Data parser and processing for Panoptic Mask R-CNN."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import maskrcnn_input
from official.vision.dataloaders import tf_example_decoder
@@ -31,7 +31,7 @@ def __init__(
include_panoptic_masks: bool,
panoptic_category_mask_key: str = 'image/panoptic/category_mask',
panoptic_instance_mask_key: str = 'image/panoptic/instance_mask'):
- super(TfExampleDecoder, self).__init__(
+ super().__init__(
include_mask=True,
regenerate_source_id=regenerate_source_id,
mask_binarize_threshold=None)
@@ -48,18 +48,24 @@ def __init__(
panoptic_category_mask_key:
tf.io.FixedLenFeature((), tf.string, default_value=''),
panoptic_instance_mask_key:
- tf.io.FixedLenFeature((), tf.string, default_value='')})
+ tf.io.FixedLenFeature((), tf.string, default_value='')
+ })
self._segmentation_keys_to_features = keys_to_features
+ def decode_segmentation_mask(self, parsed_tensors):
+ segmentation_mask = tf.io.decode_image(
+ parsed_tensors['image/segmentation/class/encoded'], channels=1)
+ segmentation_mask.set_shape([None, None, 1])
+ return segmentation_mask
+
def decode(self, serialized_example):
- decoded_tensors = super(TfExampleDecoder, self).decode(serialized_example)
+ decoded_tensors = super().decode(serialized_example)
parsed_tensors = tf.io.parse_single_example(
serialized_example, self._segmentation_keys_to_features)
- segmentation_mask = tf.io.decode_image(
- parsed_tensors['image/segmentation/class/encoded'],
- channels=1)
- segmentation_mask.set_shape([None, None, 1])
- decoded_tensors.update({'groundtruth_segmentation_mask': segmentation_mask})
+ decoded_tensors.update({
+ 'groundtruth_segmentation_mask':
+ self.decode_segmentation_mask(parsed_tensors)
+ })
if self._include_panoptic_masks:
category_mask = tf.io.decode_image(
@@ -94,10 +100,13 @@ def __init__(self,
rpn_batch_size_per_im=256,
rpn_fg_fraction=0.5,
aug_rand_hflip=False,
+ aug_rand_vflip=False,
aug_scale_min=1.0,
aug_scale_max=1.0,
+ aug_type=None,
skip_crowd_during_training=True,
max_num_instances=100,
+ outer_boxes_scale=1.0,
mask_crop_size=112,
segmentation_resize_eval_groundtruth=True,
segmentation_groundtruth_padded_size=None,
@@ -125,16 +134,21 @@ def __init__(self,
rpn_unmatched_threshold: `float`, unmatched threshold for anchors in RPN.
rpn_batch_size_per_im: `int` for batch size per image in RPN.
rpn_fg_fraction: `float` for forground fraction per batch in RPN.
- aug_rand_hflip: `bool`, if True, augment training with random
- horizontal flip.
+ aug_rand_hflip: `bool`, if True, augment training with random horizontal
+ flip.
+ aug_rand_vflip: `bool`, if True, augment training with random vertical
+ flip.
aug_scale_min: `float`, the minimum scale applied to `output_size` for
data augmentation during training.
aug_scale_max: `float`, the maximum scale applied to `output_size` for
data augmentation during training.
+ aug_type: An optional Augmentation object with params for AutoAugment.
skip_crowd_during_training: `bool`, if True, skip annotations labeled with
`is_crowd` equals to 1.
max_num_instances: `int` number of maximum number of instances in an
image. The groundtruth data will be padded to `max_num_instances`.
+ outer_boxes_scale: a float to scale up the bounding boxes to generate
+ more inclusive masks. The scale is expected to be >=1.0.
mask_crop_size: the size which groundtruth mask is cropped to.
segmentation_resize_eval_groundtruth: `bool`, if True, eval groundtruth
masks are resized to output_size.
@@ -149,7 +163,7 @@ def __init__(self,
will be parsed. Set this to true if PQ evaluator is enabled.
dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
"""
- super(Parser, self).__init__(
+ super().__init__(
output_size=output_size,
min_level=min_level,
max_level=max_level,
@@ -161,22 +175,33 @@ def __init__(self,
rpn_batch_size_per_im=rpn_batch_size_per_im,
rpn_fg_fraction=rpn_fg_fraction,
aug_rand_hflip=False,
+ aug_rand_vflip=False,
aug_scale_min=aug_scale_min,
aug_scale_max=aug_scale_max,
+ aug_type=aug_type,
skip_crowd_during_training=skip_crowd_during_training,
max_num_instances=max_num_instances,
include_mask=True,
+ outer_boxes_scale=outer_boxes_scale,
mask_crop_size=mask_crop_size,
- dtype=dtype)
+ dtype=dtype,
+ )
self.aug_rand_hflip = aug_rand_hflip
- self._segmentation_resize_eval_groundtruth = segmentation_resize_eval_groundtruth
+ self.aug_rand_vflip = aug_rand_vflip
+ self._segmentation_resize_eval_groundtruth = (
+ segmentation_resize_eval_groundtruth
+ )
if (not segmentation_resize_eval_groundtruth) and (
- segmentation_groundtruth_padded_size is None):
+ segmentation_groundtruth_padded_size is None
+ ):
raise ValueError(
'segmentation_groundtruth_padded_size ([height, width]) needs to be'
- 'specified when segmentation_resize_eval_groundtruth is False.')
- self._segmentation_groundtruth_padded_size = segmentation_groundtruth_padded_size
+ 'specified when segmentation_resize_eval_groundtruth is False.'
+ )
+ self._segmentation_groundtruth_padded_size = (
+ segmentation_groundtruth_padded_size
+ )
self._segmentation_ignore_label = segmentation_ignore_label
self._panoptic_ignore_label = panoptic_ignore_label
self._include_panoptic_masks = include_panoptic_masks
@@ -221,37 +246,49 @@ def _parse_train_data(self, data):
are supposed to be used in computing the segmentation loss while
training.
"""
+ # (height, width, num_channels = 1)
+ # All the operations below support num_channels >= 1.
segmentation_mask = data['groundtruth_segmentation_mask']
# Flips image randomly during training.
- if self.aug_rand_hflip:
- masks = data['groundtruth_instance_masks']
- image_mask = tf.concat([data['image'], segmentation_mask], axis=2)
-
- image_mask, boxes, masks = preprocess_ops.random_horizontal_flip(
- image_mask, data['groundtruth_boxes'], masks)
-
- segmentation_mask = image_mask[:, :, -1:]
- image = image_mask[:, :, :-1]
-
- data['image'] = image
- data['groundtruth_boxes'] = boxes
- data['groundtruth_instance_masks'] = masks
-
- image, labels = super(Parser, self)._parse_train_data(data)
+ image_mask = tf.concat([data['image'], segmentation_mask], axis=2)
+ boxes = data['groundtruth_boxes']
+ masks = data['groundtruth_instance_masks']
+ image_mask, boxes, masks = preprocess_ops.random_horizontal_flip(
+ image_mask,
+ boxes,
+ masks,
+ prob=tf.where(self.aug_rand_hflip, 0.5, 0.0),
+ )
+ image_mask, boxes, masks = preprocess_ops.random_vertical_flip(
+ image_mask,
+ boxes,
+ masks,
+ prob=tf.where(self.aug_rand_vflip, 0.5, 0.0),
+ )
+
+ num_image_channels = data['image'].shape.as_list()[-1]
+ image = image_mask[:, :, :num_image_channels]
+ segmentation_mask = image_mask[:, :, num_image_channels:]
+
+ data['image'] = image
+ data['groundtruth_boxes'] = boxes
+ data['groundtruth_instance_masks'] = masks
+
+ image, labels = super()._parse_train_data(data)
image_info = labels['image_info']
image_scale = image_info[2, :]
offset = image_info[3, :]
- segmentation_mask = tf.reshape(
- segmentation_mask, shape=[1, data['height'], data['width']])
- segmentation_mask = tf.cast(segmentation_mask, tf.float32)
+ # (height, width, num_channels = 1)
+ segmentation_mask = tf.cast(segmentation_mask, tf.int32)
# Pad label and make sure the padded region assigned to the ignore label.
# The label is first offset by +1 and then padded with 0.
segmentation_mask += 1
- segmentation_mask = tf.expand_dims(segmentation_mask, axis=3)
+ # (1, height, width, num_channels = 1)
+ segmentation_mask = tf.expand_dims(segmentation_mask, axis=0)
segmentation_mask = preprocess_ops.resize_and_crop_masks(
segmentation_mask, image_scale, self._output_size, offset)
segmentation_mask -= 1
@@ -259,6 +296,7 @@ def _parse_train_data(self, data):
tf.equal(segmentation_mask, -1),
self._segmentation_ignore_label * tf.ones_like(segmentation_mask),
segmentation_mask)
+ # (height, width, num_channels = 1)
segmentation_mask = tf.squeeze(segmentation_mask, axis=0)
segmentation_valid_mask = tf.not_equal(
segmentation_mask, self._segmentation_ignore_label)
@@ -291,9 +329,13 @@ def _parse_eval_data(self, data):
shape [height_l, width_l, 4] representing anchor boxes at each
level.
"""
+
def _process_mask(mask, ignore_label, image_info):
- mask = tf.cast(mask, dtype=tf.float32)
- mask = tf.reshape(mask, shape=[1, data['height'], data['width'], 1])
+ # (height, width, num_channels = 1)
+ # All the operations below support num_channels >= 1.
+ mask = tf.cast(mask, dtype=tf.int32)
+ # (1, height, width, num_channels = 1)
+ mask = tf.expand_dims(mask, axis=0)
mask += 1
if self._segmentation_resize_eval_groundtruth:
@@ -306,20 +348,22 @@ def _process_mask(mask, ignore_label, image_info):
else:
mask = tf.image.pad_to_bounding_box(
mask, 0, 0,
- self._segmentation_groundtruth_padded_size[0],
- self._segmentation_groundtruth_padded_size[1])
+ self._segmentation_groundtruth_padded_size[0], # pyrefly: ignore[unsupported-operation]
+ self._segmentation_groundtruth_padded_size[1]) # pyrefly: ignore[unsupported-operation]
mask -= 1
# Assign ignore label to the padded region.
mask = tf.where(
tf.equal(mask, -1),
ignore_label * tf.ones_like(mask),
mask)
+ # (height, width, num_channels = 1)
mask = tf.squeeze(mask, axis=0)
return mask
- image, labels = super(Parser, self)._parse_eval_data(data)
+ image, labels = super()._parse_eval_data(data)
image_info = labels['image_info']
+ # (height, width, num_channels = 1)
segmentation_mask = _process_mask(
data['groundtruth_segmentation_mask'],
self._segmentation_ignore_label, image_info)
diff --git a/official/projects/panoptic/losses/panoptic_deeplab_losses.py b/official/projects/panoptic/losses/panoptic_deeplab_losses.py
new file mode 100644
index 00000000000..b0cf03a8bd7
--- /dev/null
+++ b/official/projects/panoptic/losses/panoptic_deeplab_losses.py
@@ -0,0 +1,148 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Losses used for panoptic deeplab model."""
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.projects.panoptic.ops import mask_ops
+
+EPSILON = 1e-5
+
+
+class WeightedBootstrappedCrossEntropyLoss:
+ """Weighted semantic segmentation loss."""
+
+ def __init__(self, label_smoothing, class_weights, ignore_label,
+ top_k_percent_pixels=1.0):
+ self._top_k_percent_pixels = top_k_percent_pixels
+ self._class_weights = class_weights
+ self._ignore_label = ignore_label
+ self._label_smoothing = label_smoothing
+
+ def __call__(self, logits, labels, sample_weight=None):
+ _, _, _, num_classes = logits.get_shape().as_list()
+
+ logits = tf.image.resize(
+ logits, tf.shape(labels)[1:3],
+ method=tf.image.ResizeMethod.BILINEAR)
+
+ valid_mask = tf.not_equal(labels, self._ignore_label)
+ normalizer = tf.reduce_sum(tf.cast(valid_mask, tf.float32)) + EPSILON
+ # Assign pixel with ignore label to class 0 (background). The loss on the
+ # pixel will later be masked out.
+ labels = tf.where(valid_mask, labels, tf.zeros_like(labels))
+
+ labels = tf.squeeze(tf.cast(labels, tf.int32), axis=3)
+ valid_mask = tf.squeeze(tf.cast(valid_mask, tf.float32), axis=3)
+ onehot_labels = tf.one_hot(labels, num_classes)
+ onehot_labels = onehot_labels * (
+ 1 - self._label_smoothing) + self._label_smoothing / num_classes
+ cross_entropy_loss = tf.nn.softmax_cross_entropy_with_logits(
+ labels=onehot_labels, logits=logits)
+
+ if not self._class_weights:
+ class_weights = [1] * num_classes
+ else:
+ class_weights = self._class_weights
+
+ if num_classes != len(class_weights):
+ raise ValueError(
+ 'Length of class_weights should be {}'.format(num_classes))
+
+ weight_mask = tf.einsum('...y,y->...',
+ tf.one_hot(labels, num_classes, dtype=tf.float32),
+ tf.constant(class_weights, tf.float32))
+ valid_mask *= weight_mask
+
+ if sample_weight is not None:
+ valid_mask *= sample_weight
+
+ cross_entropy_loss *= tf.cast(valid_mask, tf.float32)
+
+ if self._top_k_percent_pixels >= 1.0:
+ loss = tf.reduce_sum(cross_entropy_loss) / normalizer
+ else:
+ loss = self._compute_top_k_loss(cross_entropy_loss)
+ return loss
+
+ def _compute_top_k_loss(self, loss):
+ """Computs top k loss."""
+ batch_size = tf.shape(loss)[0]
+ loss = tf.reshape(loss, shape=[batch_size, -1])
+
+ top_k_pixels = tf.cast(
+ self._top_k_percent_pixels *
+ tf.cast(tf.shape(loss)[-1], dtype=tf.float32),
+ dtype=tf.int32)
+
+ # shape: [batch_size, top_k_pixels]
+ per_sample_top_k_loss = tf.map_fn(
+ fn=lambda x: tf.nn.top_k(x, k=top_k_pixels, sorted=False)[0],
+ elems=loss,
+ parallel_iterations=32,
+ fn_output_signature=tf.float32)
+
+ # shape: [batch_size]
+ per_sample_normalizer = tf.reduce_sum(
+ tf.cast(
+ tf.not_equal(per_sample_top_k_loss, 0.0),
+ dtype=tf.float32),
+ axis=-1) + EPSILON
+ per_sample_normalized_loss = tf.reduce_sum(
+ per_sample_top_k_loss, axis=-1) / per_sample_normalizer
+
+ normalized_loss = tf_utils.safe_mean(per_sample_normalized_loss)
+ return normalized_loss
+
+
+class CenterHeatmapLoss:
+ """Center heatmap loss."""
+
+ def __init__(self):
+ self._loss_fn = tf.losses.mean_squared_error
+
+ def __call__(self, logits, labels, sample_weight=None):
+ _, height, width, _ = labels.get_shape().as_list()
+ logits = tf.image.resize(
+ logits,
+ size=[height, width],
+ method=tf.image.ResizeMethod.BILINEAR)
+
+ loss = self._loss_fn(y_true=labels, y_pred=logits)
+
+ if sample_weight is not None:
+ loss *= sample_weight
+
+ return tf_utils.safe_mean(loss)
+
+
+class CenterOffsetLoss:
+ """Center offset loss."""
+
+ def __init__(self):
+ self._loss_fn = tf.losses.mean_absolute_error
+
+ def __call__(self, logits, labels, sample_weight=None):
+ _, height, width, _ = labels.get_shape().as_list()
+ logits = mask_ops.resize_and_rescale_offsets(
+ logits, target_size=[height, width])
+
+ loss = self._loss_fn(y_true=labels, y_pred=logits)
+
+ if sample_weight is not None:
+ loss *= sample_weight
+
+ return tf_utils.safe_mean(loss)
diff --git a/official/projects/panoptic/modeling/factory.py b/official/projects/panoptic/modeling/factory.py
new file mode 100644
index 00000000000..052a08cd0b1
--- /dev/null
+++ b/official/projects/panoptic/modeling/factory.py
@@ -0,0 +1,253 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Factory method to build panoptic segmentation model."""
+from typing import Optional
+
+import tensorflow as tf, tf_keras
+
+from official.projects.deepmac_maskrcnn.tasks import deep_mask_head_rcnn
+from official.projects.panoptic.configs import panoptic_deeplab as panoptic_deeplab_cfg
+from official.projects.panoptic.configs import panoptic_maskrcnn as panoptic_maskrcnn_cfg
+from official.projects.panoptic.modeling import panoptic_deeplab_model
+from official.projects.panoptic.modeling import panoptic_maskrcnn_model
+from official.projects.panoptic.modeling.heads import panoptic_deeplab_heads
+from official.projects.panoptic.modeling.layers import panoptic_deeplab_merge
+from official.projects.panoptic.modeling.layers import panoptic_segmentation_generator
+from official.vision.modeling import backbones
+from official.vision.modeling.decoders import factory as decoder_factory
+from official.vision.modeling.heads import segmentation_heads
+
+
+def build_panoptic_maskrcnn(
+ input_specs: tf_keras.layers.InputSpec,
+ model_config: panoptic_maskrcnn_cfg.PanopticMaskRCNN,
+ l2_regularizer: tf_keras.regularizers.Regularizer = None) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ """Builds Panoptic Mask R-CNN model.
+
+ This factory function builds the mask rcnn first, builds the non-shared
+ semantic segmentation layers, and finally combines the two models to form
+ the panoptic segmentation model.
+
+ Args:
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
+ model_config: Config instance for the panoptic maskrcnn model.
+ l2_regularizer: Optional `tf_keras.regularizers.Regularizer`, if specified,
+ the model is built with the provided regularization layer.
+ Returns:
+ tf_keras.Model for the panoptic segmentation model.
+ """
+ norm_activation_config = model_config.norm_activation
+ segmentation_config = model_config.segmentation_model
+
+ # Builds the maskrcnn model.
+ maskrcnn_model = deep_mask_head_rcnn.build_maskrcnn(
+ input_specs=input_specs,
+ model_config=model_config,
+ l2_regularizer=l2_regularizer)
+
+ # Builds the semantic segmentation branch.
+ if not model_config.shared_backbone:
+ segmentation_backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=segmentation_config.backbone,
+ norm_activation_config=norm_activation_config,
+ l2_regularizer=l2_regularizer)
+ segmentation_decoder_input_specs = segmentation_backbone.output_specs
+ else:
+ segmentation_backbone = None
+ segmentation_decoder_input_specs = maskrcnn_model.backbone.output_specs
+
+ if not model_config.shared_decoder:
+ segmentation_decoder = decoder_factory.build_decoder(
+ input_specs=segmentation_decoder_input_specs,
+ model_config=segmentation_config,
+ l2_regularizer=l2_regularizer)
+ decoder_config = segmentation_decoder.get_config()
+ else:
+ segmentation_decoder = None
+ decoder_config = maskrcnn_model.decoder.get_config()
+
+ segmentation_head_config = segmentation_config.head
+ detection_head_config = model_config.detection_head
+ postprocessing_config = model_config.panoptic_segmentation_generator
+
+ segmentation_head = segmentation_heads.SegmentationHead(
+ num_classes=segmentation_config.num_classes,
+ level=segmentation_head_config.level,
+ num_convs=segmentation_head_config.num_convs,
+ prediction_kernel_size=segmentation_head_config.prediction_kernel_size,
+ num_filters=segmentation_head_config.num_filters,
+ upsample_factor=segmentation_head_config.upsample_factor,
+ feature_fusion=segmentation_head_config.feature_fusion,
+ decoder_min_level=segmentation_head_config.decoder_min_level,
+ decoder_max_level=segmentation_head_config.decoder_max_level,
+ low_level=segmentation_head_config.low_level,
+ low_level_num_filters=segmentation_head_config.low_level_num_filters,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ num_decoder_filters=decoder_config['num_filters'],
+ kernel_regularizer=l2_regularizer)
+
+ if model_config.generate_panoptic_masks:
+ max_num_detections = model_config.detection_generator.max_num_detections
+ mask_binarize_threshold = postprocessing_config.mask_binarize_threshold
+ panoptic_segmentation_generator_obj = (
+ panoptic_segmentation_generator.PanopticSegmentationGeneratorV2(
+ output_size=postprocessing_config.output_size,
+ max_num_detections=max_num_detections,
+ stuff_classes_offset=model_config.stuff_classes_offset,
+ mask_binarize_threshold=mask_binarize_threshold,
+ score_threshold=postprocessing_config.score_threshold,
+ things_overlap_threshold=postprocessing_config
+ .things_overlap_threshold,
+ things_class_label=postprocessing_config.things_class_label,
+ stuff_area_threshold=postprocessing_config.stuff_area_threshold,
+ void_class_label=postprocessing_config.void_class_label,
+ void_instance_id=postprocessing_config.void_instance_id,
+ rescale_predictions=postprocessing_config.rescale_predictions))
+ else:
+ panoptic_segmentation_generator_obj = None
+
+ # Combines maskrcnn, and segmentation models to build panoptic segmentation
+ # model.
+
+ model = panoptic_maskrcnn_model.PanopticMaskRCNNModel(
+ backbone=maskrcnn_model.backbone,
+ decoder=maskrcnn_model.decoder,
+ rpn_head=maskrcnn_model.rpn_head,
+ detection_head=maskrcnn_model.detection_head,
+ roi_generator=maskrcnn_model.roi_generator,
+ roi_sampler=maskrcnn_model.roi_sampler,
+ roi_aligner=maskrcnn_model.roi_aligner,
+ detection_generator=maskrcnn_model.detection_generator,
+ panoptic_segmentation_generator=panoptic_segmentation_generator_obj,
+ mask_head=maskrcnn_model.mask_head,
+ mask_sampler=maskrcnn_model.mask_sampler,
+ mask_roi_aligner=maskrcnn_model.mask_roi_aligner,
+ segmentation_backbone=segmentation_backbone,
+ segmentation_decoder=segmentation_decoder,
+ segmentation_head=segmentation_head,
+ class_agnostic_bbox_pred=detection_head_config.class_agnostic_bbox_pred,
+ cascade_class_ensemble=detection_head_config.cascade_class_ensemble,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ num_scales=model_config.anchor.num_scales,
+ aspect_ratios=model_config.anchor.aspect_ratios,
+ anchor_size=model_config.anchor.anchor_size,
+ outer_boxes_scale=maskrcnn_model.outer_boxes_scale)
+ return model
+
+
+def build_panoptic_deeplab(
+ input_specs: tf_keras.layers.InputSpec,
+ model_config: panoptic_deeplab_cfg.PanopticDeeplab,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> tf_keras.Model:
+ """Builds Panoptic Deeplab model.
+
+
+ Args:
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
+ model_config: Config instance for the panoptic deeplab model.
+ l2_regularizer: Optional `tf_keras.regularizers.Regularizer`, if specified,
+ the model is built with the provided regularization layer.
+ Returns:
+ tf_keras.Model for the panoptic segmentation model.
+ """
+ norm_activation_config = model_config.norm_activation
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=model_config.backbone,
+ norm_activation_config=norm_activation_config,
+ l2_regularizer=l2_regularizer)
+
+ semantic_decoder = decoder_factory.build_decoder(
+ input_specs=backbone.output_specs,
+ model_config=model_config,
+ l2_regularizer=l2_regularizer)
+
+ if model_config.shared_decoder:
+ instance_decoder = None
+ else:
+ # semantic and instance share the same decoder type
+ instance_decoder = decoder_factory.build_decoder(
+ input_specs=backbone.output_specs,
+ model_config=model_config,
+ l2_regularizer=l2_regularizer)
+
+ semantic_head_config = model_config.semantic_head
+ instance_head_config = model_config.instance_head
+
+ semantic_head = panoptic_deeplab_heads.SemanticHead(
+ num_classes=model_config.num_classes,
+ level=semantic_head_config.level,
+ num_convs=semantic_head_config.num_convs,
+ kernel_size=semantic_head_config.kernel_size,
+ prediction_kernel_size=semantic_head_config.prediction_kernel_size,
+ num_filters=semantic_head_config.num_filters,
+ use_depthwise_convolution=semantic_head_config.use_depthwise_convolution,
+ upsample_factor=semantic_head_config.upsample_factor,
+ low_level=semantic_head_config.low_level,
+ low_level_num_filters=semantic_head_config.low_level_num_filters,
+ fusion_num_output_filters=semantic_head_config.fusion_num_output_filters,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer)
+
+ instance_head = panoptic_deeplab_heads.InstanceHead(
+ level=instance_head_config.level,
+ num_convs=instance_head_config.num_convs,
+ kernel_size=instance_head_config.kernel_size,
+ prediction_kernel_size=instance_head_config.prediction_kernel_size,
+ num_filters=instance_head_config.num_filters,
+ use_depthwise_convolution=instance_head_config.use_depthwise_convolution,
+ upsample_factor=instance_head_config.upsample_factor,
+ low_level=instance_head_config.low_level,
+ low_level_num_filters=instance_head_config.low_level_num_filters,
+ fusion_num_output_filters=instance_head_config.fusion_num_output_filters,
+ activation=norm_activation_config.activation,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ kernel_regularizer=l2_regularizer)
+
+ if model_config.generate_panoptic_masks:
+ post_processing_config = model_config.post_processor
+ post_processor = panoptic_deeplab_merge.PostProcessor(
+ output_size=post_processing_config.output_size,
+ center_score_threshold=post_processing_config.center_score_threshold,
+ thing_class_ids=post_processing_config.thing_class_ids,
+ label_divisor=post_processing_config.label_divisor,
+ stuff_area_limit=post_processing_config.stuff_area_limit,
+ ignore_label=post_processing_config.ignore_label,
+ nms_kernel=post_processing_config.nms_kernel,
+ keep_k_centers=post_processing_config.keep_k_centers,
+ rescale_predictions=post_processing_config.rescale_predictions)
+ else:
+ post_processor = None
+
+ model = panoptic_deeplab_model.PanopticDeeplabModel(
+ backbone=backbone,
+ semantic_decoder=semantic_decoder,
+ instance_decoder=instance_decoder,
+ semantic_head=semantic_head,
+ instance_head=instance_head,
+ post_processor=post_processor)
+
+ return model
diff --git a/official/projects/panoptic/modeling/heads/panoptic_deeplab_heads.py b/official/projects/panoptic/modeling/heads/panoptic_deeplab_heads.py
new file mode 100644
index 00000000000..637e4ab3838
--- /dev/null
+++ b/official/projects/panoptic/modeling/heads/panoptic_deeplab_heads.py
@@ -0,0 +1,434 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definitions for Panoptic Deeplab heads."""
+
+from typing import List, Mapping, Optional, Tuple, Union
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.projects.panoptic.modeling.layers import fusion_layers
+from official.vision.ops import spatial_transform_ops
+
+
+class PanopticDeeplabHead(tf_keras.layers.Layer):
+ """Creates a panoptic deeplab head."""
+
+ def __init__(
+ self,
+ level: Union[int, str],
+ num_convs: int = 2,
+ num_filters: int = 256,
+ kernel_size: int = 3,
+ use_depthwise_convolution: bool = False,
+ upsample_factor: int = 1,
+ low_level: Optional[List[int]] = None,
+ low_level_num_filters: Optional[List[int]] = None,
+ fusion_num_output_filters: int = 256,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initializes a panoptic deeplab head.
+
+ Args:
+ level: An `int` or `str`, level to use to build head.
+ num_convs: An `int` number of stacked convolution before the last
+ prediction layer.
+ num_filters: An `int` number to specify the number of filters used.
+ Default is 256.
+ kernel_size: An `int` number to specify the kernel size of the
+ stacked convolutions before the last prediction layer.
+ use_depthwise_convolution: A bool to specify if use depthwise separable
+ convolutions.
+ upsample_factor: An `int` number to specify the upsampling factor to
+ generate finer mask. Default 1 means no upsampling is applied.
+ low_level: An `int` of backbone level to be used for feature fusion. It is
+ used when feature_fusion is set to `deeplabv3plus`.
+ low_level_num_filters: An `int` of reduced number of filters for the low
+ level features before fusing it with higher level features. It is only
+ used when feature_fusion is set to `deeplabv3plus`.
+ fusion_num_output_filters: An `int` number to specify the number of
+ filters used by output layer of fusion module. Default is 256.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super(PanopticDeeplabHead, self).__init__(**kwargs)
+
+ self._config_dict = {
+ 'level': level,
+ 'num_convs': num_convs,
+ 'num_filters': num_filters,
+ 'kernel_size': kernel_size,
+ 'use_depthwise_convolution': use_depthwise_convolution,
+ 'upsample_factor': upsample_factor,
+ 'low_level': low_level,
+ 'low_level_num_filters': low_level_num_filters,
+ 'fusion_num_output_filters': fusion_num_output_filters,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'norm_momentum': norm_momentum,
+ 'norm_epsilon': norm_epsilon,
+ 'kernel_regularizer': kernel_regularizer,
+ 'bias_regularizer': bias_regularizer
+ }
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ else:
+ self._bn_axis = 1
+ self._activation = tf_utils.get_activation(activation)
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the variables of the head."""
+ kernel_size = self._config_dict['kernel_size']
+ use_depthwise_convolution = self._config_dict['use_depthwise_convolution']
+ random_initializer = tf_keras.initializers.RandomNormal(stddev=0.01)
+ conv_op = tf_keras.layers.Conv2D
+ conv_kwargs = {
+ 'kernel_size': kernel_size if not use_depthwise_convolution else 1,
+ 'padding': 'same',
+ 'use_bias': True,
+ 'kernel_initializer': random_initializer,
+ 'kernel_regularizer': self._config_dict['kernel_regularizer'],
+ }
+ bn_op = (tf_keras.layers.experimental.SyncBatchNormalization
+ if self._config_dict['use_sync_bn']
+ else tf_keras.layers.BatchNormalization)
+ bn_kwargs = {
+ 'axis': self._bn_axis,
+ 'momentum': self._config_dict['norm_momentum'],
+ 'epsilon': self._config_dict['norm_epsilon'],
+ }
+
+ self._panoptic_deeplab_fusion = fusion_layers.PanopticDeepLabFusion(
+ level=self._config_dict['level'],
+ low_level=self._config_dict['low_level'],
+ num_projection_filters=self._config_dict['low_level_num_filters'],
+ num_output_filters=self._config_dict['fusion_num_output_filters'],
+ use_depthwise_convolution=self
+ ._config_dict['use_depthwise_convolution'],
+ activation=self._config_dict['activation'],
+ use_sync_bn=self._config_dict['use_sync_bn'],
+ norm_momentum=self._config_dict['norm_momentum'],
+ norm_epsilon=self._config_dict['norm_epsilon'],
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'])
+
+ # Stacked convolutions layers.
+ self._convs = []
+ self._norms = []
+ for i in range(self._config_dict['num_convs']):
+ if use_depthwise_convolution:
+ self._convs.append(
+ tf_keras.layers.DepthwiseConv2D(
+ name='panoptic_deeplab_head_depthwise_conv_{}'.format(i),
+ kernel_size=kernel_size,
+ padding='same',
+ use_bias=True,
+ depthwise_initializer=random_initializer,
+ depthwise_regularizer=self._config_dict['kernel_regularizer'],
+ depth_multiplier=1))
+ norm_name = 'panoptic_deeplab_head_depthwise_norm_{}'.format(i)
+ self._norms.append(bn_op(name=norm_name, **bn_kwargs))
+ conv_name = 'panoptic_deeplab_head_conv_{}'.format(i)
+ self._convs.append(
+ conv_op(
+ name=conv_name,
+ filters=self._config_dict['num_filters'],
+ **conv_kwargs))
+ norm_name = 'panoptic_deeplab_head_norm_{}'.format(i)
+ self._norms.append(bn_op(name=norm_name, **bn_kwargs))
+
+ super().build(input_shape)
+
+ def call(self, inputs: Tuple[Union[tf.Tensor, Mapping[str, tf.Tensor]],
+ Union[tf.Tensor, Mapping[str, tf.Tensor]]],
+ training=None):
+ """Forward pass of the head.
+
+ It supports both a tuple of 2 tensors or 2 dictionaries. The first is
+ backbone endpoints, and the second is decoder endpoints. When inputs are
+ tensors, they are from a single level of feature maps. When inputs are
+ dictionaries, they contain multiple levels of feature maps, where the key
+ is the index of feature map.
+
+ Args:
+ inputs: A tuple of 2 feature map tensors of shape
+ [batch, height_l, width_l, channels] or 2 dictionaries of tensors:
+ - key: A `str` of the level of the multilevel features.
+ - values: A `tf.Tensor` of the feature map tensors, whose shape is
+ [batch, height_l, width_l, channels].
+ training: A bool, runs the model in training/eval mode.
+
+ Returns:
+ A `tf.Tensor` of the fused backbone and decoder features.
+ """
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+
+ x = self._panoptic_deeplab_fusion(inputs, training=training)
+
+ for conv, norm in zip(self._convs, self._norms):
+ x = conv(x)
+ x = norm(x, training=training)
+ x = self._activation(x)
+
+ if self._config_dict['upsample_factor'] > 1:
+ x = spatial_transform_ops.nearest_upsampling(
+ x, scale=self._config_dict['upsample_factor'])
+
+ return x
+
+ def get_config(self):
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(self._config_dict.items()))
+
+ @classmethod
+ def from_config(cls, config):
+ return cls(**config)
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class SemanticHead(PanopticDeeplabHead):
+ """Creates a semantic head."""
+
+ def __init__(
+ self,
+ num_classes: int,
+ level: Union[int, str],
+ num_convs: int = 2,
+ num_filters: int = 256,
+ kernel_size: int = 3,
+ prediction_kernel_size: int = 3,
+ use_depthwise_convolution: bool = False,
+ upsample_factor: int = 1,
+ low_level: Optional[List[int]] = None,
+ low_level_num_filters: Optional[List[int]] = None,
+ fusion_num_output_filters: int = 256,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initializes a instance center head.
+
+ Args:
+ num_classes: An `int` number of mask classification categories. The number
+ of classes does not include background class.
+ level: An `int` or `str`, level to use to build head.
+ num_convs: An `int` number of stacked convolution before the last
+ prediction layer.
+ num_filters: An `int` number to specify the number of filters used.
+ Default is 256.
+ kernel_size: An `int` number to specify the kernel size of the
+ stacked convolutions before the last prediction layer.
+ prediction_kernel_size: An `int` number to specify the kernel size of the
+ prediction layer.
+ use_depthwise_convolution: A bool to specify if use depthwise separable
+ convolutions.
+ upsample_factor: An `int` number to specify the upsampling factor to
+ generate finer mask. Default 1 means no upsampling is applied.
+ low_level: An `int` of backbone level to be used for feature fusion. It is
+ used when feature_fusion is set to `deeplabv3plus`.
+ low_level_num_filters: An `int` of reduced number of filters for the low
+ level features before fusing it with higher level features. It is only
+ used when feature_fusion is set to `deeplabv3plus`.
+ fusion_num_output_filters: An `int` number to specify the number of
+ filters used by output layer of fusion module. Default is 256.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super(SemanticHead, self).__init__(
+ level=level,
+ num_convs=num_convs,
+ num_filters=num_filters,
+ use_depthwise_convolution=use_depthwise_convolution,
+ kernel_size=kernel_size,
+ upsample_factor=upsample_factor,
+ low_level=low_level,
+ low_level_num_filters=low_level_num_filters,
+ fusion_num_output_filters=fusion_num_output_filters,
+ activation=activation,
+ use_sync_bn=use_sync_bn,
+ norm_momentum=norm_momentum,
+ norm_epsilon=norm_epsilon,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ **kwargs)
+ self._config_dict.update({
+ 'num_classes': num_classes,
+ 'prediction_kernel_size': prediction_kernel_size})
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the variables of the semantic head."""
+ super(SemanticHead, self).build(input_shape)
+ self._classifier = tf_keras.layers.Conv2D(
+ name='semantic_output',
+ filters=self._config_dict['num_classes'],
+ kernel_size=self._config_dict['prediction_kernel_size'],
+ padding='same',
+ bias_initializer=tf.zeros_initializer(),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'])
+
+ def call(self, inputs: Tuple[Union[tf.Tensor, Mapping[str, tf.Tensor]],
+ Union[tf.Tensor, Mapping[str, tf.Tensor]]],
+ training=None):
+ """Forward pass of the head."""
+
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+ x = super(SemanticHead, self).call(inputs, training=training)
+ outputs = self._classifier(x)
+ return outputs
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class InstanceHead(PanopticDeeplabHead):
+ """Creates a instance head."""
+
+ def __init__(
+ self,
+ level: Union[int, str],
+ num_convs: int = 2,
+ num_filters: int = 256,
+ kernel_size: int = 3,
+ prediction_kernel_size: int = 3,
+ use_depthwise_convolution: bool = False,
+ upsample_factor: int = 1,
+ low_level: Optional[List[int]] = None,
+ low_level_num_filters: Optional[List[int]] = None,
+ fusion_num_output_filters: int = 256,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initializes a instance center head.
+
+ Args:
+ level: An `int` or `str`, level to use to build head.
+ num_convs: An `int` number of stacked convolution before the last
+ prediction layer.
+ num_filters: An `int` number to specify the number of filters used.
+ Default is 256.
+ kernel_size: An `int` number to specify the kernel size of the
+ stacked convolutions before the last prediction layer.
+ prediction_kernel_size: An `int` number to specify the kernel size of the
+ prediction layer.
+ use_depthwise_convolution: A bool to specify if use depthwise separable
+ convolutions.
+ upsample_factor: An `int` number to specify the upsampling factor to
+ generate finer mask. Default 1 means no upsampling is applied.
+ low_level: An `int` of backbone level to be used for feature fusion. It is
+ used when feature_fusion is set to `deeplabv3plus`.
+ low_level_num_filters: An `int` of reduced number of filters for the low
+ level features before fusing it with higher level features. It is only
+ used when feature_fusion is set to `deeplabv3plus`.
+ fusion_num_output_filters: An `int` number to specify the number of
+ filters used by output layer of fusion module. Default is 256.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super(InstanceHead, self).__init__(
+ level=level,
+ num_convs=num_convs,
+ num_filters=num_filters,
+ use_depthwise_convolution=use_depthwise_convolution,
+ kernel_size=kernel_size,
+ upsample_factor=upsample_factor,
+ low_level=low_level,
+ low_level_num_filters=low_level_num_filters,
+ fusion_num_output_filters=fusion_num_output_filters,
+ activation=activation,
+ use_sync_bn=use_sync_bn,
+ norm_momentum=norm_momentum,
+ norm_epsilon=norm_epsilon,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ **kwargs)
+ self._config_dict.update({
+ 'prediction_kernel_size': prediction_kernel_size})
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the variables of the instance head."""
+ super(InstanceHead, self).build(input_shape)
+ self._instance_center_prediction_conv = tf_keras.layers.Conv2D(
+ name='instance_centers_heatmap',
+ filters=1,
+ kernel_size=self._config_dict['prediction_kernel_size'],
+ padding='same',
+ bias_initializer=tf.zeros_initializer(),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'])
+
+ self._instance_center_regression_conv = tf_keras.layers.Conv2D(
+ name='instance_centers_offset',
+ filters=2,
+ kernel_size=self._config_dict['prediction_kernel_size'],
+ padding='same',
+ bias_initializer=tf.zeros_initializer(),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'])
+
+ def call(self, inputs: Tuple[Union[tf.Tensor, Mapping[str, tf.Tensor]],
+ Union[tf.Tensor, Mapping[str, tf.Tensor]]],
+ training=None):
+ """Forward pass of the head."""
+
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+
+ x = super(InstanceHead, self).call(inputs, training=training)
+ instance_centers_heatmap = self._instance_center_prediction_conv(x)
+ instance_centers_offset = self._instance_center_regression_conv(x)
+ outputs = {
+ 'instance_centers_heatmap': instance_centers_heatmap,
+ 'instance_centers_offset': instance_centers_offset
+ }
+ return outputs
diff --git a/official/projects/panoptic/modeling/layers/fusion_layers.py b/official/projects/panoptic/modeling/layers/fusion_layers.py
new file mode 100644
index 00000000000..1289a9c6ed3
--- /dev/null
+++ b/official/projects/panoptic/modeling/layers/fusion_layers.py
@@ -0,0 +1,180 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains feature fusion blocks for panoptic segmentation models."""
+from typing import Any, Callable, Dict, List, Mapping, Optional, Union
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+
+
+# Type annotations.
+States = Dict[str, tf.Tensor]
+Activation = Union[str, Callable]
+
+
+class PanopticDeepLabFusion(tf_keras.layers.Layer):
+ """Creates a Panoptic DeepLab feature Fusion layer.
+
+ This implements the feature fusion introduced in the paper:
+ Cheng et al. Panoptic-DeepLab
+ (https://arxiv.org/pdf/1911.10194.pdf)
+ """
+
+ def __init__(
+ self,
+ level: int,
+ low_level: List[int],
+ num_projection_filters: List[int],
+ num_output_filters: int = 256,
+ use_depthwise_convolution: bool = False,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ interpolation: str = 'bilinear',
+ **kwargs):
+ """Initializes panoptic FPN feature fusion layer.
+
+ Args:
+ level: An `int` level at which the decoder was appled at.
+ low_level: A list of `int` of minimum level to use in feature fusion.
+ num_projection_filters: A list of `int` with number of filters for
+ projection conv2d layers.
+ num_output_filters: An `int` number of filters in output conv2d layers.
+ use_depthwise_convolution: A bool to specify if use depthwise separable
+ convolutions.
+ activation: A `str` name of the activation function.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ interpolation: A `str` interpolation method for upsampling. Defaults to
+ `bilinear`.
+ **kwargs: Additional keyword arguments to be passed.
+ Returns:
+ A `float` `tf.Tensor` of shape [batch_size, feature_height, feature_width,
+ feature_channel].
+ """
+ super(PanopticDeepLabFusion, self).__init__(**kwargs)
+
+ self._config_dict = {
+ 'level': level,
+ 'low_level': low_level,
+ 'num_projection_filters': num_projection_filters,
+ 'num_output_filters': num_output_filters,
+ 'use_depthwise_convolution': use_depthwise_convolution,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'norm_momentum': norm_momentum,
+ 'norm_epsilon': norm_epsilon,
+ 'kernel_regularizer': kernel_regularizer,
+ 'bias_regularizer': bias_regularizer,
+ 'interpolation': interpolation
+ }
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._channel_axis = -1
+ else:
+ self._channel_axis = 1
+ self._activation = tf_utils.get_activation(activation)
+
+ def build(self, input_shape: List[tf.TensorShape]):
+ conv_op = tf_keras.layers.Conv2D
+ conv_kwargs = {
+ 'padding': 'same',
+ 'use_bias': True,
+ 'kernel_initializer': tf.initializers.VarianceScaling(),
+ 'kernel_regularizer': self._config_dict['kernel_regularizer'],
+ }
+ bn_op = (tf_keras.layers.experimental.SyncBatchNormalization
+ if self._config_dict['use_sync_bn']
+ else tf_keras.layers.BatchNormalization)
+ bn_kwargs = {
+ 'axis': self._channel_axis,
+ 'momentum': self._config_dict['norm_momentum'],
+ 'epsilon': self._config_dict['norm_epsilon'],
+ }
+
+ self._projection_convs = []
+ self._projection_norms = []
+ self._fusion_convs = []
+ self._fusion_norms = []
+ for i in range(len(self._config_dict['low_level'])):
+ self._projection_convs.append(
+ conv_op(
+ filters=self._config_dict['num_projection_filters'][i],
+ kernel_size=1,
+ **conv_kwargs))
+ if self._config_dict['use_depthwise_convolution']:
+ depthwise_initializer = tf_keras.initializers.RandomNormal(stddev=0.01)
+ fusion_conv = tf_keras.Sequential([
+ tf_keras.layers.DepthwiseConv2D(
+ kernel_size=5,
+ padding='same',
+ use_bias=True,
+ depthwise_initializer=depthwise_initializer,
+ depthwise_regularizer=self._config_dict['kernel_regularizer'],
+ depth_multiplier=1),
+ bn_op(**bn_kwargs),
+ conv_op(
+ filters=self._config_dict['num_output_filters'],
+ kernel_size=1,
+ **conv_kwargs)])
+ else:
+ fusion_conv = conv_op(
+ filters=self._config_dict['num_output_filters'],
+ kernel_size=5,
+ **conv_kwargs)
+ self._fusion_convs.append(fusion_conv)
+ self._projection_norms.append(bn_op(**bn_kwargs))
+ self._fusion_norms.append(bn_op(**bn_kwargs))
+
+ def call(self, inputs, training=None):
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+
+ backbone_output = inputs[0]
+ decoder_output = inputs[1][str(self._config_dict['level'])]
+
+ x = decoder_output
+ for i in range(len(self._config_dict['low_level'])):
+ feature = backbone_output[str(self._config_dict['low_level'][i])]
+ feature = self._projection_convs[i](feature)
+ feature = self._projection_norms[i](feature, training=training)
+ feature = self._activation(feature)
+
+ shape = tf.shape(feature)
+ x = tf.image.resize(
+ x, size=[shape[1], shape[2]],
+ method=self._config_dict['interpolation'])
+ x = tf.cast(x, dtype=feature.dtype)
+ x = tf.concat([x, feature], axis=self._channel_axis)
+
+ x = self._fusion_convs[i](x)
+ x = self._fusion_norms[i](x, training=training)
+ x = self._activation(x)
+ return x
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
diff --git a/official/projects/panoptic/modeling/layers/panoptic_deeplab_merge.py b/official/projects/panoptic/modeling/layers/panoptic_deeplab_merge.py
new file mode 100644
index 00000000000..9703b57ba66
--- /dev/null
+++ b/official/projects/panoptic/modeling/layers/panoptic_deeplab_merge.py
@@ -0,0 +1,568 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""This file contains functions to post-process Panoptic-DeepLab results.
+
+Note that the postprocessing class and the supporting functions are branched
+from:
+https://github.com/google-research/deeplab2/blob/main/model/post_processor/panoptic_deeplab.py
+with minor changes.
+"""
+
+import functools
+from typing import Dict, List, Text, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.projects.panoptic.ops import mask_ops
+
+
+def _add_zero_padding(input_tensor: tf.Tensor, kernel_size: int,
+ rank: int) -> tf.Tensor:
+ """Adds zero-padding to the input_tensor."""
+ pad_total = kernel_size - 1
+ pad_begin = pad_total // 2
+ pad_end = pad_total - pad_begin
+ if rank == 3:
+ return tf.pad(
+ input_tensor,
+ paddings=[[pad_begin, pad_end], [pad_begin, pad_end], [0, 0]])
+ else:
+ return tf.pad(
+ input_tensor,
+ paddings=[[0, 0], [pad_begin, pad_end], [pad_begin, pad_end], [0, 0]])
+
+
+def _get_semantic_predictions(semantic_logits: tf.Tensor) -> tf.Tensor:
+ """Computes the semantic classes from the predictions.
+
+ Args:
+ semantic_logits: A tf.tensor of shape [batch, height, width, classes].
+ Returns:
+ A tf.Tensor containing the semantic class prediction of shape
+ [batch, height, width].
+ """
+ return tf.argmax(semantic_logits, axis=-1, output_type=tf.int32)
+
+
+def _get_instance_centers_from_heatmap(
+ center_heatmap: tf.Tensor,
+ center_threshold: float,
+ nms_kernel_size: int,
+ keep_k_centers: int) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Computes a list of instance centers.
+
+ Args:
+ center_heatmap: A tf.Tensor of shape [height, width, 1].
+ center_threshold: A float setting the threshold for the center heatmap.
+ nms_kernel_size: An integer specifying the nms kernel size.
+ keep_k_centers: An integer specifying the number of centers to keep (K).
+ Non-positive values will keep all centers.
+ Returns:
+ A tuple of
+ - tf.Tensor of shape [N, 2] containing N center coordinates (after
+ non-maximum suppression) in (y, x) order.
+ - tf.Tensor of shape [height, width] containing the center heatmap after
+ non-maximum suppression.
+ """
+ # Threshold center map.
+ center_heatmap = tf.where(
+ tf.greater(center_heatmap, center_threshold), center_heatmap, 0.0)
+
+ # Non-maximum suppression.
+ padded_map = _add_zero_padding(center_heatmap, nms_kernel_size, rank=3)
+ pooled_center_heatmap = tf_keras.backend.pool2d(
+ tf.expand_dims(padded_map, 0),
+ pool_size=(nms_kernel_size, nms_kernel_size),
+ strides=(1, 1),
+ padding='valid',
+ pool_mode='max')
+ center_heatmap = tf.where(
+ tf.equal(pooled_center_heatmap, center_heatmap), center_heatmap, 0.0)
+ center_heatmap = tf.squeeze(center_heatmap, axis=[0, 3])
+
+ # `centers` is of shape (N, 2) with (y, x) order of the second dimension.
+ centers = tf.where(tf.greater(center_heatmap, 0.0))
+
+ if keep_k_centers > 0 and tf.shape(centers)[0] > keep_k_centers:
+ topk_scores, _ = tf.math.top_k(
+ tf.reshape(center_heatmap, [-1]), keep_k_centers, sorted=False)
+ centers = tf.where(tf.greater(center_heatmap, topk_scores[-1]))
+
+ return centers, center_heatmap
+
+
+def _find_closest_center_per_pixel(centers: tf.Tensor,
+ center_offsets: tf.Tensor) -> tf.Tensor:
+ """Assigns all pixels to their closest center.
+
+ Args:
+ centers: A tf.Tensor of shape [N, 2] containing N centers with coordinate
+ order (y, x).
+ center_offsets: A tf.Tensor of shape [height, width, 2].
+ Returns:
+ A tf.Tensor of shape [height, width] containing the index of the closest
+ center, per pixel.
+ """
+ height = tf.shape(center_offsets)[0]
+ width = tf.shape(center_offsets)[1]
+
+ x_coord, y_coord = tf.meshgrid(tf.range(width), tf.range(height))
+ coord = tf.stack([y_coord, x_coord], axis=-1)
+
+ center_per_pixel = tf.cast(coord, tf.float32) + center_offsets
+
+ # centers: [N, 2] -> [N, 1, 2].
+ # center_per_pixel: [H, W, 2] -> [1, H*W, 2].
+ centers = tf.cast(tf.expand_dims(centers, 1), tf.float32)
+ center_per_pixel = tf.reshape(center_per_pixel, [height*width, 2])
+ center_per_pixel = tf.expand_dims(center_per_pixel, 0)
+
+ # distances: [N, H*W].
+ distances = tf.norm(centers - center_per_pixel, axis=-1)
+
+ return tf.reshape(tf.argmin(distances, axis=0), [height, width])
+
+
+def _get_instances_from_heatmap_and_offset(
+ semantic_segmentation: tf.Tensor, center_heatmap: tf.Tensor,
+ center_offsets: tf.Tensor, center_threshold: float,
+ thing_class_ids: tf.Tensor, nms_kernel_size: int,
+ keep_k_centers: int) -> Tuple[tf.Tensor, tf.Tensor, tf.Tensor]:
+ """Computes the instance assignment per pixel.
+
+ Args:
+ semantic_segmentation: A tf.Tensor containing the semantic labels of shape
+ [height, width].
+ center_heatmap: A tf.Tensor of shape [height, width, 1].
+ center_offsets: A tf.Tensor of shape [height, width, 2].
+ center_threshold: A float setting the threshold for the center heatmap.
+ thing_class_ids: A tf.Tensor of shape [N] containing N thing indices.
+ nms_kernel_size: An integer specifying the nms kernel size.
+ keep_k_centers: An integer specifying the number of centers to keep.
+ Negative values will keep all centers.
+ Returns:
+ A tuple of:
+ - tf.Tensor containing the instance segmentation (filtered with the `thing`
+ segmentation from the semantic segmentation output) with shape
+ [height, width].
+ - tf.Tensor containing the processed centermap with shape [height, width].
+ - tf.Tensor containing instance scores (where higher "score" is a reasonable
+ signal of a higher confidence detection.) Will be of shape [height, width]
+ with the score for a pixel being the score of the instance it belongs to.
+ The scores will be zero for pixels in background/"stuff" regions.
+ """
+ thing_segmentation = tf.zeros_like(semantic_segmentation)
+ for thing_id in thing_class_ids:
+ thing_segmentation = tf.where(tf.equal(semantic_segmentation, thing_id),
+ 1,
+ thing_segmentation)
+
+ centers, processed_center_heatmap = _get_instance_centers_from_heatmap(
+ center_heatmap, center_threshold, nms_kernel_size, keep_k_centers)
+ if tf.shape(centers)[0] == 0:
+ return (tf.zeros_like(semantic_segmentation), processed_center_heatmap,
+ tf.zeros_like(processed_center_heatmap))
+
+ instance_center_index = _find_closest_center_per_pixel(
+ centers, center_offsets)
+ # Instance IDs should start with 1. So we use the index into the centers, but
+ # shifted by 1.
+ instance_segmentation = tf.cast(instance_center_index, tf.int32) + 1
+
+ # The value of the heatmap at an instance's center is used as the score
+ # for that instance.
+ instance_scores = tf.gather_nd(processed_center_heatmap, centers)
+ # This will map the instance scores back to the image space: where each pixel
+ # has a value equal to the score of its instance.
+ flat_center_index = tf.reshape(instance_center_index, [-1])
+ instance_score_map = tf.gather(instance_scores, flat_center_index)
+ instance_score_map = tf.reshape(instance_score_map,
+ tf.shape(instance_segmentation))
+ instance_score_map *= tf.cast(thing_segmentation, tf.float32)
+
+ return (thing_segmentation * instance_segmentation, processed_center_heatmap,
+ instance_score_map)
+
+
+@tf.function
+def _get_panoptic_predictions(
+ semantic_logits: tf.Tensor, center_heatmap: tf.Tensor,
+ center_offsets: tf.Tensor, center_threshold: float,
+ thing_class_ids: tf.Tensor, label_divisor: int, stuff_area_limit: int,
+ void_label: int, nms_kernel_size: int, keep_k_centers: int
+) -> Tuple[tf.Tensor, tf.Tensor, tf.Tensor, tf.Tensor]:
+ """Computes the semantic class and instance ID per pixel.
+
+ Args:
+ semantic_logits: A tf.Tensor of shape [batch, height, width, classes].
+ center_heatmap: A tf.Tensor of shape [batch, height, width, 1].
+ center_offsets: A tf.Tensor of shape [batch, height, width, 2].
+ center_threshold: A float setting the threshold for the center heatmap.
+ thing_class_ids: A tf.Tensor of shape [N] containing N thing indices.
+ label_divisor: An integer specifying the label divisor of the dataset.
+ stuff_area_limit: An integer specifying the number of pixels that stuff
+ regions need to have at least. The stuff region will be included in the
+ panoptic prediction, only if its area is larger than the limit; otherwise,
+ it will be re-assigned as void_label.
+ void_label: An integer specifying the void label.
+ nms_kernel_size: An integer specifying the nms kernel size.
+ keep_k_centers: An integer specifying the number of centers to keep.
+ Negative values will keep all centers.
+ Returns:
+ A tuple of:
+ - the panoptic prediction as tf.Tensor with shape [batch, height, width].
+ - the centermap prediction as tf.Tensor with shape [batch, height, width].
+ - the instance score maps as tf.Tensor with shape [batch, height, width].
+ - the instance prediction as tf.Tensor with shape [batch, height, width].
+ """
+ semantic_prediction = _get_semantic_predictions(semantic_logits)
+ batch_size = tf.shape(semantic_logits)[0]
+
+ instance_map_lists = tf.TensorArray(
+ tf.int32, size=batch_size, dynamic_size=False)
+ center_map_lists = tf.TensorArray(
+ tf.float32, size=batch_size, dynamic_size=False)
+ instance_score_map_lists = tf.TensorArray(
+ tf.float32, size=batch_size, dynamic_size=False)
+
+ for i in tf.range(batch_size):
+ (instance_map, center_map,
+ instance_score_map) = _get_instances_from_heatmap_and_offset(
+ semantic_prediction[i, ...], center_heatmap[i, ...],
+ center_offsets[i, ...], center_threshold, thing_class_ids,
+ nms_kernel_size, keep_k_centers)
+ instance_map_lists = instance_map_lists.write(i, instance_map)
+ center_map_lists = center_map_lists.write(i, center_map)
+ instance_score_map_lists = instance_score_map_lists.write(
+ i, instance_score_map)
+
+ # This does not work with unknown shapes.
+ instance_maps = instance_map_lists.stack()
+ center_maps = center_map_lists.stack()
+ instance_score_maps = instance_score_map_lists.stack()
+
+ panoptic_prediction = _merge_semantic_and_instance_maps(
+ semantic_prediction, instance_maps, thing_class_ids, label_divisor,
+ stuff_area_limit, void_label)
+ return (panoptic_prediction, center_maps, instance_score_maps, instance_maps)
+
+
+@tf.function
+def _merge_semantic_and_instance_maps(
+ semantic_prediction: tf.Tensor,
+ instance_maps: tf.Tensor,
+ thing_class_ids: tf.Tensor,
+ label_divisor: int,
+ stuff_area_limit: int,
+ void_label: int) -> tf.Tensor:
+ """Merges semantic and instance maps to obtain panoptic segmentation.
+
+ This function merges the semantic segmentation and class-agnostic
+ instance segmentation to form the panoptic segmentation. In particular,
+ the class label of each instance mask is inferred from the majority
+ votes from the corresponding pixels in the semantic segmentation. This
+ operation is first proposed in the DeeperLab paper and adopted by the
+ Panoptic-DeepLab.
+ - DeeperLab: Single-Shot Image Parser, T-J Yang, et al. arXiv:1902.05093.
+ - Panoptic-DeepLab, B. Cheng, et al. In CVPR, 2020.
+ Note that this function only supports batch = 1 for simplicity. Additionally,
+ this function has a slightly different implementation from the provided
+ TensorFlow implementation `merge_ops` but with a similar performance. This
+ function is mainly used as a backup solution when you could not successfully
+ compile the provided TensorFlow implementation. To reproduce our results,
+ please use the provided TensorFlow implementation (i.e., not use this
+ function, but the `merge_ops.merge_semantic_and_instance_maps`).
+
+ Args:
+ semantic_prediction: A tf.Tensor of shape [batch, height, width].
+ instance_maps: A tf.Tensor of shape [batch, height, width].
+ thing_class_ids: A tf.Tensor of shape [N] containing N thing indices.
+ label_divisor: An integer specifying the label divisor of the dataset.
+ stuff_area_limit: An integer specifying the number of pixels that stuff
+ regions need to have at least. The stuff region will be included in the
+ panoptic prediction, only if its area is larger than the limit; otherwise,
+ it will be re-assigned as void_label.
+ void_label: An integer specifying the void label.
+ Returns:
+ panoptic_prediction: A tf.Tensor with shape [batch, height, width].
+ """
+ prediction_shape = semantic_prediction.get_shape().as_list()
+ # This implementation only supports batch size of 1. Since model construction
+ # might lose batch size information (and leave it to None), override it here.
+ prediction_shape[0] = 1
+ semantic_prediction = tf.ensure_shape(semantic_prediction, prediction_shape)
+ instance_maps = tf.ensure_shape(instance_maps, prediction_shape)
+
+ # Default panoptic_prediction to have semantic label = void_label.
+ panoptic_prediction = tf.ones_like(
+ semantic_prediction) * void_label * label_divisor
+
+ # Start to paste predicted `thing` regions to panoptic_prediction.
+ # Infer `thing` segmentation regions from semantic prediction.
+ semantic_thing_segmentation = tf.zeros_like(semantic_prediction,
+ dtype=tf.bool)
+ for thing_class in thing_class_ids:
+ semantic_thing_segmentation = tf.math.logical_or(
+ semantic_thing_segmentation,
+ semantic_prediction == thing_class)
+ # Keep track of how many instances for each semantic label.
+ num_instance_per_semantic_label = tf.TensorArray(
+ tf.int32, size=0, dynamic_size=True, clear_after_read=False)
+ instance_ids, _ = tf.unique(tf.reshape(instance_maps, [-1]))
+ for instance_id in instance_ids:
+ # Instance ID 0 is reserved for crowd region.
+ if instance_id == 0:
+ continue
+ thing_mask = tf.math.logical_and(instance_maps == instance_id,
+ semantic_thing_segmentation)
+ if tf.reduce_sum(tf.cast(thing_mask, tf.int32)) == 0:
+ continue
+ semantic_bin_counts = tf.math.bincount(
+ tf.boolean_mask(semantic_prediction, thing_mask))
+ semantic_majority = tf.cast(
+ tf.math.argmax(semantic_bin_counts), tf.int32)
+
+ while num_instance_per_semantic_label.size() <= semantic_majority:
+ num_instance_per_semantic_label = num_instance_per_semantic_label.write(
+ num_instance_per_semantic_label.size(), 0)
+
+ new_instance_id = (
+ num_instance_per_semantic_label.read(semantic_majority) + 1)
+ num_instance_per_semantic_label = num_instance_per_semantic_label.write(
+ semantic_majority, new_instance_id)
+ panoptic_prediction = tf.where(
+ thing_mask,
+ tf.ones_like(panoptic_prediction) * semantic_majority * label_divisor
+ + new_instance_id,
+ panoptic_prediction)
+
+ # Done with `num_instance_per_semantic_label` tensor array.
+ num_instance_per_semantic_label.close()
+
+ # Start to paste predicted `stuff` regions to panoptic prediction.
+ instance_stuff_regions = instance_maps == 0
+ semantic_ids, _ = tf.unique(tf.reshape(semantic_prediction, [-1]))
+ for semantic_id in semantic_ids:
+ if tf.reduce_sum(tf.cast(thing_class_ids == semantic_id, tf.int32)) > 0:
+ continue
+ # Check stuff area.
+ stuff_mask = tf.math.logical_and(semantic_prediction == semantic_id,
+ instance_stuff_regions)
+ stuff_area = tf.reduce_sum(tf.cast(stuff_mask, tf.int32))
+ if stuff_area >= stuff_area_limit:
+ panoptic_prediction = tf.where(
+ stuff_mask,
+ tf.ones_like(panoptic_prediction) * semantic_id * label_divisor,
+ panoptic_prediction)
+
+ return panoptic_prediction
+
+
+class PostProcessor(tf_keras.layers.Layer):
+ """This class contains code of a Panoptic-Deeplab post-processor."""
+
+ def __init__(
+ self,
+ output_size: List[int],
+ center_score_threshold: float,
+ thing_class_ids: List[int],
+ label_divisor: int,
+ stuff_area_limit: int,
+ ignore_label: int,
+ nms_kernel: int,
+ keep_k_centers: int,
+ rescale_predictions: bool,
+ **kwargs):
+ """Initializes a Panoptic-Deeplab post-processor.
+
+ Args:
+ output_size: A `List` of integers that represent the height and width of
+ the output mask.
+ center_score_threshold: A float setting the threshold for the center
+ heatmap.
+ thing_class_ids: An integer list shape [N] containing N thing indices.
+ label_divisor: An integer specifying the label divisor of the dataset.
+ stuff_area_limit: An integer specifying the number of pixels that stuff
+ regions need to have at least. The stuff region will be included in the
+ panoptic prediction, only if its area is larger than the limit;
+ otherwise, it will be re-assigned as void_label.
+ ignore_label: An integer specifying the void label.
+ nms_kernel: An integer specifying the nms kernel size.
+ keep_k_centers: An integer specifying the number of centers to keep.
+ Negative values will keep all centers.
+ rescale_predictions: `bool`, whether to scale back prediction to original
+ image sizes. If True, image_info is used to rescale predictions.
+ **kwargs: additional kwargs arguments.
+ """
+ super(PostProcessor, self).__init__(**kwargs)
+
+ self._config_dict = {
+ 'output_size': output_size,
+ 'center_score_threshold': center_score_threshold,
+ 'thing_class_ids': thing_class_ids,
+ 'label_divisor': label_divisor,
+ 'stuff_area_limit': stuff_area_limit,
+ 'ignore_label': ignore_label,
+ 'nms_kernel': nms_kernel,
+ 'keep_k_centers': keep_k_centers,
+ 'rescale_predictions': rescale_predictions
+ }
+ self._post_processor = functools.partial(
+ _get_panoptic_predictions,
+ center_threshold=center_score_threshold,
+ thing_class_ids=tf.convert_to_tensor(thing_class_ids),
+ label_divisor=label_divisor,
+ stuff_area_limit=stuff_area_limit,
+ void_label=ignore_label,
+ nms_kernel_size=nms_kernel,
+ keep_k_centers=keep_k_centers)
+
+ def _resize_and_pad_masks(self, mask, image_info):
+ """Resizes masks to match the original image shape and pads to`output_size`.
+
+ Args:
+ mask: a padded mask tensor.
+ image_info: a tensor that holds information about original and
+ preprocessed images.
+ Returns:
+ resized and padded masks: tf.Tensor.
+ """
+ rescale_size = tf.cast(
+ tf.math.ceil(image_info[1, :] / image_info[2, :]), tf.int32)
+ image_shape = tf.cast(image_info[0, :], tf.int32)
+ offsets = tf.cast(image_info[3, :], tf.int32)
+
+ mask = tf.image.resize(
+ mask,
+ rescale_size,
+ method='bilinear')
+ mask = tf.image.crop_to_bounding_box(
+ mask,
+ offsets[0], offsets[1],
+ image_shape[0],
+ image_shape[1])
+ mask = tf.image.pad_to_bounding_box(
+ mask, 0, 0,
+ self._config_dict['output_size'][0],
+ self._config_dict['output_size'][1])
+ return mask
+
+ def _resize_and_pad_offset_mask(self, mask, image_info):
+ """Rescales and resizes offset masks and pads to`output_size`.
+
+ Args:
+ mask: a padded offset mask tensor.
+ image_info: a tensor that holds information about original and
+ preprocessed images.
+ Returns:
+ rescaled, resized and padded masks: tf.Tensor.
+ """
+ rescale_size = tf.cast(
+ tf.math.ceil(image_info[1, :] / image_info[2, :]), tf.int32)
+ image_shape = tf.cast(image_info[0, :], tf.int32)
+ offsets = tf.cast(image_info[3, :], tf.int32)
+
+ mask = mask_ops.resize_and_rescale_offsets(
+ tf.expand_dims(mask, axis=0),
+ rescale_size)[0]
+ mask = tf.image.crop_to_bounding_box(
+ mask,
+ offsets[0], offsets[1],
+ image_shape[0],
+ image_shape[1])
+ mask = tf.image.pad_to_bounding_box(
+ mask, 0, 0,
+ self._config_dict['output_size'][0],
+ self._config_dict['output_size'][1])
+ return mask
+
+ def call(
+ self,
+ result_dict: Dict[Text, tf.Tensor],
+ image_info: tf.Tensor) -> Dict[Text, tf.Tensor]:
+ """Performs the post-processing given model predicted results.
+
+ Args:
+ result_dict: A dictionary of tf.Tensor containing model results. The dict
+ has to contain
+ - segmentation_outputs
+ - instance_centers_heatmap
+ - instance_centers_offset
+ image_info: A tf.Tensor of image infos.
+
+ Returns:
+ The post-processed dict of tf.Tensor, containing the following keys:
+ - panoptic_outputs
+ - category_mask
+ - instance_mask
+ - instance_centers
+ - instance_score
+ """
+ if self._config_dict['rescale_predictions']:
+ segmentation_outputs = tf.map_fn(
+ fn=lambda x: self._resize_and_pad_masks(x[0], x[1]),
+ elems=(result_dict['segmentation_outputs'], image_info),
+ fn_output_signature=tf.float32,
+ parallel_iterations=32)
+ instance_centers_heatmap = tf.map_fn(
+ fn=lambda x: self._resize_and_pad_masks(x[0], x[1]),
+ elems=(result_dict['instance_centers_heatmap'], image_info),
+ fn_output_signature=tf.float32,
+ parallel_iterations=32)
+ instance_centers_offset = tf.map_fn(
+ fn=lambda x: self._resize_and_pad_offset_mask(x[0], x[1]),
+ elems=(result_dict['instance_centers_offset'], image_info),
+ fn_output_signature=tf.float32,
+ parallel_iterations=32)
+ else:
+ segmentation_outputs = tf.image.resize(
+ result_dict['segmentation_outputs'],
+ size=self._config_dict['output_size'],
+ method='bilinear')
+ instance_centers_heatmap = tf.image.resize(
+ result_dict['instance_centers_heatmap'],
+ size=self._config_dict['output_size'],
+ method='bilinear')
+ instance_centers_offset = mask_ops.resize_and_rescale_offsets(
+ result_dict['instance_centers_offset'],
+ target_size=self._config_dict['output_size'])
+
+ processed_dict = {}
+
+ (processed_dict['panoptic_outputs'],
+ processed_dict['instance_centers'],
+ processed_dict['instance_scores'],
+ _) = self._post_processor(
+ tf.nn.softmax(segmentation_outputs, axis=-1),
+ instance_centers_heatmap,
+ instance_centers_offset)
+
+ label_divisor = self._config_dict['label_divisor']
+ processed_dict['category_mask'] = (
+ processed_dict['panoptic_outputs'] // label_divisor)
+ processed_dict['instance_mask'] = (
+ processed_dict['panoptic_outputs'] % label_divisor)
+
+ processed_dict.update({
+ 'segmentation_outputs': result_dict['segmentation_outputs']})
+
+ return processed_dict
+
+ def get_config(self):
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config):
+ return cls(**config)
diff --git a/official/projects/panoptic/modeling/layers/panoptic_segmentation_generator.py b/official/projects/panoptic/modeling/layers/panoptic_segmentation_generator.py
new file mode 100644
index 00000000000..c4093da8a92
--- /dev/null
+++ b/official/projects/panoptic/modeling/layers/panoptic_segmentation_generator.py
@@ -0,0 +1,617 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definition for postprocessing layer to genrate panoptic segmentations."""
+
+from typing import Any, Dict, List, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.projects.panoptic.modeling.layers import paste_masks
+from official.vision.ops import spatial_transform_ops
+
+
+def _batch_count_ones(masks: tf.Tensor,
+ dtype: tf.dtypes.DType = tf.int32) -> tf.Tensor:
+ """Counts the ones/trues for each mask in the batch.
+
+ Args:
+ masks: A tensor in shape (..., height, width) with arbitrary numbers of
+ batch dimensions.
+ dtype: DType of the resulting tensor. Default is tf.int32.
+
+ Returns:
+ A tensor which contains the count of non-zero elements for each mask in the
+ batch. The rank of the resulting tensor is equal to rank(masks) - 2.
+ """
+ masks_shape = masks.get_shape().as_list()
+ if len(masks_shape) < 2:
+ raise ValueError(
+ 'Expected the input masks (..., height, width) has rank >= 2, was: %s' %
+ masks_shape)
+ return tf.reduce_sum(tf.cast(masks, dtype), axis=[-2, -1])
+
+
+class PanopticSegmentationGenerator(tf_keras.layers.Layer):
+ """Panoptic segmentation generator layer."""
+
+ def __init__(
+ self,
+ output_size: List[int],
+ max_num_detections: int,
+ stuff_classes_offset: int,
+ mask_binarize_threshold: float = 0.5,
+ score_threshold: float = 0.5,
+ things_overlap_threshold: float = 0.5,
+ stuff_area_threshold: float = 4096,
+ things_class_label: int = 1,
+ void_class_label: int = 0,
+ void_instance_id: int = -1,
+ rescale_predictions: bool = False,
+ **kwargs):
+ """Generates panoptic segmentation masks.
+
+ Args:
+ output_size: A `List` of integers that represent the height and width of
+ the output mask.
+ max_num_detections: `int` for maximum number of detections.
+ stuff_classes_offset: An `int` that is added to the output of the
+ semantic segmentation mask to make sure that the stuff class ids do not
+ ovelap with the thing class ids of the MaskRCNN outputs.
+ mask_binarize_threshold: A `float`
+ score_threshold: A `float` representing the threshold for deciding
+ when to remove objects based on score.
+ things_overlap_threshold: A `float` representing a threshold for deciding
+ to ignore a thing if overlap is above the threshold.
+ stuff_area_threshold: A `float` representing a threshold for deciding to
+ to ignore a stuff class if area is below certain threshold.
+ things_class_label: An `int` that represents a single merged category of
+ all thing classes in the semantic segmentation output.
+ void_class_label: An `int` that is used to represent empty or unlabelled
+ regions of the mask
+ void_instance_id: An `int` that is used to denote regions that are not
+ assigned to any thing class. That is, void_instance_id are assigned to
+ both stuff regions and empty regions.
+ rescale_predictions: `bool`, whether to scale back prediction to original
+ image sizes. If True, image_info is used to rescale predictions.
+ **kwargs: additional kewargs arguments.
+ """
+ self._output_size = output_size
+ self._max_num_detections = max_num_detections
+ self._stuff_classes_offset = stuff_classes_offset
+ self._mask_binarize_threshold = mask_binarize_threshold
+ self._score_threshold = score_threshold
+ self._things_overlap_threshold = things_overlap_threshold
+ self._stuff_area_threshold = stuff_area_threshold
+ self._things_class_label = things_class_label
+ self._void_class_label = void_class_label
+ self._void_instance_id = void_instance_id
+ self._rescale_predictions = rescale_predictions
+
+ self._config_dict = {
+ 'output_size': output_size,
+ 'max_num_detections': max_num_detections,
+ 'stuff_classes_offset': stuff_classes_offset,
+ 'mask_binarize_threshold': mask_binarize_threshold,
+ 'score_threshold': score_threshold,
+ 'things_class_label': things_class_label,
+ 'void_class_label': void_class_label,
+ 'void_instance_id': void_instance_id,
+ 'rescale_predictions': rescale_predictions
+ }
+ super().__init__(**kwargs)
+
+ def build(self, input_shape: tf.TensorShape):
+ grid_sampler = paste_masks.BilinearGridSampler(align_corners=False)
+ self._paste_masks_fn = paste_masks.PasteMasks(
+ output_size=self._output_size, grid_sampler=grid_sampler)
+ super().build(input_shape)
+
+ def _generate_panoptic_masks(
+ self, boxes: tf.Tensor, scores: tf.Tensor, classes: tf.Tensor,
+ detections_masks: tf.Tensor,
+ segmentation_mask: tf.Tensor) -> Dict[str, tf.Tensor]:
+ """Generates panoptic masks for a single image.
+
+ This function implements the following steps to merge instance and semantic
+ segmentation masks described in https://arxiv.org/pdf/1901.02446.pdf
+ Steps:
+ 1. resolving overlaps between different instances based on their
+ confidence scores
+ 2. resolving overlaps between instance and semantic segmentation
+ outputs in favor of instances
+ 3. removing any stuff regions labeled other or under a given area
+ threshold.
+ Args:
+ boxes: A `tf.Tensor` of shape [num_rois, 4], representing the bounding
+ boxes for detected objects.
+ scores: A `tf.Tensor` of shape [num_rois], representing the
+ confidence scores for each object.
+ classes: A `tf.Tensor` of shape [num_rois], representing the class
+ for each object.
+ detections_masks: A `tf.Tensor` of shape
+ [num_rois, mask_height, mask_width, 1], representing the cropped mask
+ for each object.
+ segmentation_mask: A `tf.Tensor` of shape [height, width], representing
+ the semantic segmentation output.
+ Returns:
+ Dict with the following keys:
+ - category_mask: A `tf.Tensor` for category masks.
+ - instance_mask: A `tf.Tensor for instance masks.
+ """
+
+ # Offset stuff class predictions
+ segmentation_mask = tf.where(
+ tf.logical_or(
+ tf.equal(segmentation_mask, self._things_class_label),
+ tf.equal(segmentation_mask, self._void_class_label)),
+ segmentation_mask,
+ segmentation_mask + self._stuff_classes_offset
+ )
+ # sort instances by their scores
+ sorted_indices = tf.argsort(scores, direction='DESCENDING')
+
+ mask_shape = self._output_size + [1]
+ category_mask = tf.ones(mask_shape,
+ dtype=tf.float32) * self._void_class_label
+ instance_mask = tf.ones(
+ mask_shape, dtype=tf.float32) * self._void_instance_id
+
+ # filter instances with low confidence
+ sorted_scores = tf.sort(scores, direction='DESCENDING')
+
+ valid_indices = tf.where(sorted_scores > self._score_threshold)
+
+ # if no instance has sufficient confidence score, skip merging
+ # instance segmentation masks
+ if tf.shape(valid_indices)[0] > 0:
+ loop_end_idx = valid_indices[-1, 0] + 1
+ loop_end_idx = tf.minimum(
+ tf.cast(loop_end_idx, dtype=tf.int32),
+ self._max_num_detections)
+ pasted_masks = self._paste_masks_fn((
+ detections_masks[:loop_end_idx],
+ boxes[:loop_end_idx]))
+
+ # add things segmentation to panoptic masks
+ for i in range(loop_end_idx):
+ # we process instances in decending order, which will make sure
+ # the overlaps are resolved based on confidence score
+ instance_idx = sorted_indices[i]
+
+ pasted_mask = pasted_masks[instance_idx]
+
+ class_id = tf.cast(classes[instance_idx], dtype=tf.float32)
+
+ # convert sigmoid scores to binary values
+ binary_mask = tf.greater(
+ pasted_mask, self._mask_binarize_threshold)
+
+ # filter empty instance masks
+ if not tf.reduce_sum(tf.cast(binary_mask, tf.float32)) > 0:
+ continue
+
+ overlap = tf.logical_and(
+ binary_mask,
+ tf.not_equal(category_mask, self._void_class_label))
+ binary_mask_area = tf.reduce_sum(
+ tf.cast(binary_mask, dtype=tf.float32))
+ overlap_area = tf.reduce_sum(
+ tf.cast(overlap, dtype=tf.float32))
+
+ # skip instance that have a big enough overlap with instances with
+ # higer scores
+ if overlap_area / binary_mask_area > self._things_overlap_threshold:
+ continue
+
+ # fill empty regions in category_mask represented by
+ # void_class_label with class_id of the instance.
+ category_mask = tf.where(
+ tf.logical_and(
+ binary_mask, tf.equal(category_mask, self._void_class_label)),
+ tf.ones_like(category_mask) * class_id, category_mask)
+
+ # fill empty regions in the instance_mask represented by
+ # void_instance_id with the id of the instance, starting from 1
+ instance_mask = tf.where(
+ tf.logical_and(
+ binary_mask,
+ tf.equal(instance_mask, self._void_instance_id)),
+ tf.ones_like(instance_mask) *
+ tf.cast(instance_idx + 1, tf.float32), instance_mask)
+
+ stuff_class_ids = tf.unique(tf.reshape(segmentation_mask, [-1])).y
+ for stuff_class_id in stuff_class_ids:
+ if stuff_class_id == self._things_class_label:
+ continue
+
+ stuff_mask = tf.logical_and(
+ tf.equal(segmentation_mask, stuff_class_id),
+ tf.equal(category_mask, self._void_class_label))
+
+ stuff_mask_area = tf.reduce_sum(
+ tf.cast(stuff_mask, dtype=tf.float32))
+
+ if stuff_mask_area < self._stuff_area_threshold:
+ continue
+
+ category_mask = tf.where(
+ stuff_mask,
+ tf.ones_like(category_mask) * stuff_class_id,
+ category_mask)
+
+ results = {
+ 'category_mask': category_mask[:, :, 0],
+ 'instance_mask': instance_mask[:, :, 0]
+ }
+ return results
+
+ def _resize_and_pad_masks(self, mask, image_info):
+ """Resizes masks to match the original image shape and pads to`output_size`.
+
+ Args:
+ mask: a padded mask tensor.
+ image_info: a tensor that holds information about original and
+ preprocessed images.
+ Returns:
+ resized and padded masks: tf.Tensor.
+ """
+ rescale_size = tf.cast(
+ tf.math.ceil(image_info[1, :] / image_info[2, :]), tf.int32)
+ image_shape = tf.cast(image_info[0, :], tf.int32)
+ offsets = tf.cast(image_info[3, :], tf.int32)
+
+ mask = tf.image.resize(
+ mask,
+ rescale_size,
+ method='bilinear')
+ mask = tf.image.crop_to_bounding_box(
+ mask,
+ offsets[0], offsets[1],
+ image_shape[0],
+ image_shape[1])
+ mask = tf.image.pad_to_bounding_box(
+ mask, 0, 0, self._output_size[0], self._output_size[1])
+ return mask
+
+ def call(self,
+ inputs: tf.Tensor,
+ image_info: Optional[tf.Tensor] = None) -> Dict[str, tf.Tensor]:
+ detections = inputs
+
+ batched_scores = detections['detection_scores']
+ batched_classes = detections['detection_classes']
+ batched_detections_masks = tf.expand_dims(
+ detections['detection_masks'], axis=-1)
+ batched_boxes = detections['detection_boxes']
+ batched_segmentation_masks = tf.cast(
+ detections['segmentation_outputs'], dtype=tf.float32)
+
+ if self._rescale_predictions:
+ scale = tf.tile(
+ tf.cast(image_info[:, 2:3, :], dtype=batched_boxes.dtype), # pyrefly: ignore[unsupported-operation]
+ multiples=[1, 1, 2])
+ batched_boxes /= scale
+
+ batched_segmentation_masks = tf.map_fn(
+ fn=lambda x: self._resize_and_pad_masks(x[0], x[1]),
+ elems=(
+ batched_segmentation_masks,
+ image_info),
+ fn_output_signature=tf.float32,
+ parallel_iterations=32)
+ else:
+ batched_segmentation_masks = tf.image.resize(
+ batched_segmentation_masks,
+ size=self._output_size,
+ method='bilinear')
+
+ batched_segmentation_masks = tf.expand_dims(tf.cast(
+ tf.argmax(batched_segmentation_masks, axis=-1),
+ dtype=tf.float32), axis=-1)
+
+ panoptic_masks = tf.map_fn(
+ fn=lambda x: self._generate_panoptic_masks( # pylint:disable=g-long-lambda
+ x[0], x[1], x[2], x[3], x[4]),
+ elems=(
+ batched_boxes,
+ batched_scores,
+ batched_classes,
+ batched_detections_masks,
+ batched_segmentation_masks),
+ fn_output_signature={
+ 'category_mask': tf.float32,
+ 'instance_mask': tf.float32
+ }, parallel_iterations=32)
+
+ for k, v in panoptic_masks.items():
+ panoptic_masks[k] = tf.cast(v, dtype=tf.int32)
+
+ return panoptic_masks
+
+ def get_config(self) -> Dict[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config: Dict[str,
+ Any]) -> 'PanopticSegmentationGenerator':
+ return cls(**config)
+
+
+class PanopticSegmentationGeneratorV2(tf_keras.layers.Layer):
+ """Panoptic segmentation generator layer V2."""
+
+ def __init__(self,
+ output_size: List[int],
+ max_num_detections: int,
+ stuff_classes_offset: int,
+ mask_binarize_threshold: float = 0.5,
+ score_threshold: float = 0.5,
+ things_overlap_threshold: float = 0.5,
+ stuff_area_threshold: float = 4096,
+ things_class_label: int = 1,
+ void_class_label: int = 0,
+ void_instance_id: int = -1,
+ rescale_predictions: bool = False,
+ **kwargs):
+ """Generates panoptic segmentation masks.
+
+ Args:
+ output_size: A `List` of integers that represent the height and width of
+ the output mask.
+ max_num_detections: `int` for maximum number of detections.
+ stuff_classes_offset: An `int` that is added to the output of the semantic
+ segmentation mask to make sure that the stuff class ids do not ovelap
+ with the thing class ids of the MaskRCNN outputs.
+ mask_binarize_threshold: A `float`
+ score_threshold: A `float` representing the threshold for deciding when to
+ remove objects based on score.
+ things_overlap_threshold: A `float` representing a threshold for deciding
+ to ignore a thing if overlap is above the threshold.
+ stuff_area_threshold: A `float` representing a threshold for deciding to
+ to ignore a stuff class if area is below certain threshold.
+ things_class_label: An `int` that represents a single merged category of
+ all thing classes in the semantic segmentation output.
+ void_class_label: An `int` that is used to represent empty or unlabelled
+ regions of the mask
+ void_instance_id: An `int` that is used to denote regions that are not
+ assigned to any thing class. That is, void_instance_id are assigned to
+ both stuff regions and empty regions.
+ rescale_predictions: `bool`, whether to scale back prediction to original
+ image sizes. If True, image_info is used to rescale predictions.
+ **kwargs: additional kewargs arguments.
+ """
+ self._output_size = output_size
+ self._max_num_detections = max_num_detections
+ self._stuff_classes_offset = stuff_classes_offset
+ self._mask_binarize_threshold = mask_binarize_threshold
+ self._score_threshold = score_threshold
+ self._things_overlap_threshold = things_overlap_threshold
+ self._stuff_area_threshold = stuff_area_threshold
+ self._things_class_label = things_class_label
+ self._void_class_label = void_class_label
+ self._void_instance_id = void_instance_id
+ self._rescale_predictions = rescale_predictions
+
+ self._config_dict = {
+ 'output_size': output_size,
+ 'max_num_detections': max_num_detections,
+ 'stuff_classes_offset': stuff_classes_offset,
+ 'mask_binarize_threshold': mask_binarize_threshold,
+ 'score_threshold': score_threshold,
+ 'things_class_label': things_class_label,
+ 'void_class_label': void_class_label,
+ 'void_instance_id': void_instance_id,
+ 'rescale_predictions': rescale_predictions
+ }
+ super().__init__(**kwargs)
+
+ def call(self,
+ inputs: tf.Tensor,
+ image_info: Optional[tf.Tensor] = None) -> Dict[str, tf.Tensor]:
+ """Generates panoptic segmentation masks."""
+ # (batch_size, num_rois, 4) in absolute coordinates.
+ detection_boxes = tf.cast(inputs['detection_boxes'], tf.float32)
+ # (batch_size, num_rois)
+ detection_classes = tf.cast(inputs['detection_classes'], tf.int32)
+ # (batch_size, num_rois)
+ detection_scores = inputs['detection_scores']
+ # (batch_size, num_rois, mask_height, mask_width)
+ detections_masks = inputs['detection_masks']
+ # (batch_size, height, width, num_semantic_classes)
+ segmentation_outputs = inputs['segmentation_outputs']
+
+ if self._rescale_predictions:
+ # (batch_size, 2)
+ original_size = tf.cast(image_info[:, 0, :], tf.float32) # pyrefly: ignore[unsupported-operation]
+ desired_size = tf.cast(image_info[:, 1, :], tf.float32) # pyrefly: ignore[unsupported-operation]
+ image_scale = tf.cast(image_info[:, 2, :], tf.float32) # pyrefly: ignore[unsupported-operation]
+ offset = tf.cast(image_info[:, 3, :], tf.float32) # pyrefly: ignore[unsupported-operation]
+ rescale_size = tf.math.ceil(desired_size / image_scale)
+ # (batch_size, output_height, output_width, num_semantic_classes)
+ segmentation_outputs = (
+ spatial_transform_ops.bilinear_resize_with_crop_and_pad(
+ segmentation_outputs,
+ rescale_size,
+ crop_offset=offset,
+ crop_size=original_size,
+ output_size=self._output_size))
+ # (batch_size, 1, 4)
+ image_scale = tf.tile(image_scale, multiples=[1, 2])[:, tf.newaxis]
+ detection_boxes /= image_scale
+ else:
+ # (batch_size, output_height, output_width, num_semantic_classes)
+ segmentation_outputs = tf.image.resize(
+ segmentation_outputs, size=self._output_size, method='bilinear')
+
+ # (batch_size, output_height, output_width)
+ instance_mask, instance_category_mask = self._generate_instances(
+ detection_boxes, detection_classes, detection_scores, detections_masks)
+
+ # (batch_size, output_height, output_width)
+ stuff_category_mask = self._generate_stuffs(segmentation_outputs)
+
+ # (batch_size, output_height, output_width)
+ category_mask = tf.where((stuff_category_mask != self._void_class_label) &
+ (instance_category_mask == self._void_class_label),
+ stuff_category_mask + self._stuff_classes_offset,
+ instance_category_mask)
+
+ return {'instance_mask': instance_mask, 'category_mask': category_mask}
+
+ def _generate_instances(
+ self, detection_boxes: tf.Tensor, detection_classes: tf.Tensor,
+ detection_scores: tf.Tensor,
+ detections_masks: tf.Tensor) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Generates instance & category masks from instance segmentation outputs."""
+ batch_size = tf.shape(detections_masks)[0]
+ num_rois = tf.shape(detections_masks)[1]
+ mask_height = tf.shape(detections_masks)[2]
+ mask_width = tf.shape(detections_masks)[3]
+ output_height = self._output_size[0]
+ output_width = self._output_size[1]
+
+ # (batch_size, num_rois, mask_height, mask_width)
+ detections_masks = detections_masks * (
+ tf.cast((detection_scores > self._score_threshold) &
+ (detection_classes != self._void_class_label),
+ detections_masks.dtype)[:, :, tf.newaxis, tf.newaxis])
+
+ # Resizes and copies the detections_masks to the bounding boxes in the
+ # output canvas.
+ # (batch_size, num_rois, output_height, output_width)
+ pasted_detection_masks = tf.reshape(
+ spatial_transform_ops.bilinear_resize_to_bbox(
+ tf.reshape(detections_masks, [-1, mask_height, mask_width]),
+ tf.reshape(detection_boxes, [-1, 4]), self._output_size),
+ shape=[-1, num_rois, output_height, output_width])
+
+ # (batch_size, num_rois, output_height, output_width)
+ instance_binary_masks = (
+ pasted_detection_masks > self._mask_binarize_threshold)
+
+ # Sorts detection related tensors by scores.
+ # (batch_size, num_rois)
+ sorted_detection_indices = tf.argsort(
+ detection_scores, axis=1, direction='DESCENDING')
+ # (batch_size, num_rois)
+ sorted_detection_classes = tf.gather(
+ detection_classes, sorted_detection_indices, batch_dims=1)
+ # (batch_size, num_rois, output_height, output_width)
+ sorted_instance_binary_masks = tf.gather(
+ instance_binary_masks, sorted_detection_indices, batch_dims=1)
+ # (batch_size, num_rois)
+ instance_areas = _batch_count_ones(
+ sorted_instance_binary_masks, dtype=tf.float32)
+
+ init_loop_vars = (
+ 0, # i: the loop counter
+ tf.ones([batch_size, output_height, output_width], dtype=tf.int32) *
+ self._void_instance_id, # combined_instance_mask
+ tf.ones([batch_size, output_height, output_width], dtype=tf.int32) *
+ self._void_class_label # combined_category_mask
+ )
+
+ def _copy_instances_loop_body(
+ i: int, combined_instance_mask: tf.Tensor,
+ combined_category_mask: tf.Tensor) -> Tuple[int, tf.Tensor, tf.Tensor]:
+ """Iterates the sorted detections and copies the instances."""
+ # (batch_size, output_height, output_width)
+ instance_binary_mask = sorted_instance_binary_masks[:, i]
+
+ # Masks out the instances that have a big enough overlap with the other
+ # instances with higher scores.
+ # (batch_size, )
+ overlap_areas = _batch_count_ones(
+ (combined_instance_mask != self._void_instance_id)
+ & instance_binary_mask,
+ dtype=tf.float32)
+ # (batch_size, )
+ instance_overlap_threshold_mask = tf.math.divide_no_nan(
+ overlap_areas, instance_areas[:, i]) < self._things_overlap_threshold
+ # (batch_size, output_height, output_width)
+ instance_binary_mask &= (
+ instance_overlap_threshold_mask[:, tf.newaxis, tf.newaxis]
+ & (combined_instance_mask == self._void_instance_id))
+
+ # Updates combined_instance_mask.
+ # (batch_size, )
+ instance_id = tf.cast(
+ sorted_detection_indices[:, i] + 1, # starting from 1
+ dtype=combined_instance_mask.dtype)
+ # (batch_size, output_height, output_width)
+ combined_instance_mask = tf.where(instance_binary_mask,
+ instance_id[:, tf.newaxis, tf.newaxis],
+ combined_instance_mask)
+
+ # Updates combined_category_mask.
+ # (batch_size, )
+ class_id = tf.cast(
+ sorted_detection_classes[:, i], dtype=combined_category_mask.dtype)
+ # (batch_size, output_height, output_width)
+ combined_category_mask = tf.where(instance_binary_mask,
+ class_id[:, tf.newaxis, tf.newaxis],
+ combined_category_mask)
+
+ # Returns the updated loop vars.
+ return (
+ i + 1, # Increment the loop counter i
+ combined_instance_mask,
+ combined_category_mask)
+
+ # (batch_size, output_height, output_width)
+ _, instance_mask, category_mask = tf.while_loop(
+ cond=lambda i, *_: i < num_rois,
+ body=_copy_instances_loop_body,
+ loop_vars=init_loop_vars,
+ parallel_iterations=32,
+ maximum_iterations=num_rois)
+ return instance_mask, category_mask
+
+ def _generate_stuffs(self, segmentation_outputs: tf.Tensor) -> tf.Tensor:
+ """Generates category mask from semantic segmentation outputs."""
+ num_semantic_classes = tf.shape(segmentation_outputs)[3]
+
+ # (batch_size, output_height, output_width)
+ segmentation_masks = tf.argmax(
+ segmentation_outputs, axis=-1, output_type=tf.int32)
+ stuff_binary_masks = (segmentation_masks != self._things_class_label) & (
+ segmentation_masks != self._void_class_label)
+ # (batch_size, num_semantic_classes, output_height, output_width)
+ stuff_class_binary_masks = ((tf.one_hot(
+ segmentation_masks, num_semantic_classes, axis=1, dtype=tf.int32) == 1)
+ & tf.expand_dims(stuff_binary_masks, axis=1))
+
+ # Masks out the stuff class whose area is below the given threshold.
+ # (batch_size, num_semantic_classes)
+ stuff_class_areas = _batch_count_ones(
+ stuff_class_binary_masks, dtype=tf.float32)
+ # (batch_size, num_semantic_classes, output_height, output_width)
+ stuff_class_binary_masks &= tf.greater(
+ stuff_class_areas, self._stuff_area_threshold)[:, :, tf.newaxis,
+ tf.newaxis]
+ # (batch_size, output_height, output_width)
+ stuff_binary_masks = tf.reduce_any(stuff_class_binary_masks, axis=1)
+
+ # (batch_size, output_height, output_width)
+ return tf.where(stuff_binary_masks, segmentation_masks,
+ tf.ones_like(segmentation_masks) * self._void_class_label)
+
+ def get_config(self) -> Dict[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config: Dict[str,
+ Any]) -> 'PanopticSegmentationGeneratorV2':
+ return cls(**config)
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/paste_masks.py b/official/projects/panoptic/modeling/layers/paste_masks.py
similarity index 96%
rename from official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/paste_masks.py
rename to official/projects/panoptic/modeling/layers/paste_masks.py
index e46ffae8ca3..9dc221fa39c 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/paste_masks.py
+++ b/official/projects/panoptic/modeling/layers/paste_masks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,10 +16,10 @@
from typing import List
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-class BilinearGridSampler(tf.keras.layers.Layer):
+class BilinearGridSampler(tf_keras.layers.Layer):
"""Bilinear Grid Sampling layer."""
def __init__(self, align_corners: bool = False, **kwargs):
@@ -127,7 +127,7 @@ def from_config(cls, config):
return cls(**config)
-class PasteMasks(tf.keras.layers.Layer):
+class PasteMasks(tf_keras.layers.Layer):
"""Layer to paste instance masks."""
def __init__(self, output_size: List[int],
diff --git a/official/projects/panoptic/modeling/panoptic_deeplab_model.py b/official/projects/panoptic/modeling/panoptic_deeplab_model.py
new file mode 100644
index 00000000000..a5eae863a9e
--- /dev/null
+++ b/official/projects/panoptic/modeling/panoptic_deeplab_model.py
@@ -0,0 +1,122 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Build Panoptic Deeplab model."""
+from typing import Any, Mapping, Optional, Union
+
+import tensorflow as tf, tf_keras
+from official.projects.panoptic.modeling.layers import panoptic_deeplab_merge
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class PanopticDeeplabModel(tf_keras.Model):
+ """Panoptic Deeplab model."""
+
+ def __init__(
+ self,
+ backbone: tf_keras.Model,
+ semantic_decoder: tf_keras.Model,
+ semantic_head: tf_keras.layers.Layer,
+ instance_head: tf_keras.layers.Layer,
+ instance_decoder: Optional[tf_keras.Model] = None,
+ post_processor: Optional[panoptic_deeplab_merge.PostProcessor] = None,
+ **kwargs):
+ """Panoptic deeplab model initializer.
+
+ Args:
+ backbone: a backbone network.
+ semantic_decoder: a decoder network. E.g. FPN.
+ semantic_head: segmentation head.
+ instance_head: instance center head.
+ instance_decoder: Optional decoder network for instance predictions.
+ post_processor: Optional post processor layer.
+ **kwargs: keyword arguments to be passed.
+ """
+ super(PanopticDeeplabModel, self).__init__(**kwargs)
+
+ self._config_dict = {
+ 'backbone': backbone,
+ 'semantic_decoder': semantic_decoder,
+ 'instance_decoder': instance_decoder,
+ 'semantic_head': semantic_head,
+ 'instance_head': instance_head,
+ 'post_processor': post_processor
+ }
+ self.backbone = backbone
+ self.semantic_decoder = semantic_decoder
+ self.instance_decoder = instance_decoder
+ self.semantic_head = semantic_head
+ self.instance_head = instance_head
+ self.post_processor = post_processor
+
+ def call( # pytype: disable=annotation-type-mismatch,signature-mismatch
+ self, inputs: tf.Tensor,
+ image_info: tf.Tensor,
+ training: bool = None): # pyrefly: ignore[bad-function-definition]
+ if training is None:
+ training = tf_keras.backend.learning_phase()
+
+ backbone_features = self.backbone(inputs, training=training)
+
+ semantic_features = self.semantic_decoder(
+ backbone_features, training=training)
+
+ if self.instance_decoder is None:
+ instance_features = semantic_features
+ else:
+ instance_features = self.instance_decoder(
+ backbone_features, training=training)
+
+ segmentation_outputs = self.semantic_head(
+ (backbone_features, semantic_features),
+ training=training)
+ instance_outputs = self.instance_head(
+ (backbone_features, instance_features),
+ training=training)
+
+ outputs = {
+ 'segmentation_outputs': segmentation_outputs,
+ 'instance_centers_heatmap':
+ instance_outputs['instance_centers_heatmap'],
+ 'instance_centers_offset':
+ instance_outputs['instance_centers_offset'],
+ }
+ if training:
+ return outputs
+
+ if self.post_processor is not None:
+ panoptic_masks = self.post_processor(outputs, image_info)
+ outputs.update(panoptic_masks)
+ return outputs
+
+ @property
+ def checkpoint_items(
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(
+ backbone=self.backbone,
+ semantic_decoder=self.semantic_decoder,
+ semantic_head=self.semantic_head,
+ instance_head=self.instance_head)
+ if self.instance_decoder is not None:
+ items.update(instance_decoder=self.instance_decoder)
+
+ return items
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/modeling/panoptic_maskrcnn_model.py b/official/projects/panoptic/modeling/panoptic_maskrcnn_model.py
similarity index 77%
rename from official/vision/beta/projects/panoptic_maskrcnn/modeling/panoptic_maskrcnn_model.py
rename to official/projects/panoptic/modeling/panoptic_maskrcnn_model.py
index 309f7fa7bf7..0162bbfab0e 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/modeling/panoptic_maskrcnn_model.py
+++ b/official/projects/panoptic/modeling/panoptic_maskrcnn_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import List, Mapping, Optional, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.deepmac_maskrcnn.modeling import maskrcnn_model
@@ -26,23 +26,23 @@ class PanopticMaskRCNNModel(maskrcnn_model.DeepMaskRCNNModel):
def __init__(
self,
- backbone: tf.keras.Model,
- decoder: tf.keras.Model,
- rpn_head: tf.keras.layers.Layer,
- detection_head: Union[tf.keras.layers.Layer,
- List[tf.keras.layers.Layer]],
- roi_generator: tf.keras.layers.Layer,
- roi_sampler: Union[tf.keras.layers.Layer,
- List[tf.keras.layers.Layer]],
- roi_aligner: tf.keras.layers.Layer,
- detection_generator: tf.keras.layers.Layer,
- panoptic_segmentation_generator: Optional[tf.keras.layers.Layer] = None,
- mask_head: Optional[tf.keras.layers.Layer] = None,
- mask_sampler: Optional[tf.keras.layers.Layer] = None,
- mask_roi_aligner: Optional[tf.keras.layers.Layer] = None,
- segmentation_backbone: Optional[tf.keras.Model] = None,
- segmentation_decoder: Optional[tf.keras.Model] = None,
- segmentation_head: tf.keras.layers.Layer = None,
+ backbone: tf_keras.Model,
+ decoder: tf_keras.Model,
+ rpn_head: tf_keras.layers.Layer,
+ detection_head: Union[tf_keras.layers.Layer,
+ List[tf_keras.layers.Layer]],
+ roi_generator: tf_keras.layers.Layer,
+ roi_sampler: Union[tf_keras.layers.Layer,
+ List[tf_keras.layers.Layer]],
+ roi_aligner: tf_keras.layers.Layer,
+ detection_generator: tf_keras.layers.Layer,
+ panoptic_segmentation_generator: Optional[tf_keras.layers.Layer] = None,
+ mask_head: Optional[tf_keras.layers.Layer] = None,
+ mask_sampler: Optional[tf_keras.layers.Layer] = None,
+ mask_roi_aligner: Optional[tf_keras.layers.Layer] = None,
+ segmentation_backbone: Optional[tf_keras.Model] = None,
+ segmentation_decoder: Optional[tf_keras.Model] = None,
+ segmentation_head: tf_keras.layers.Layer = None, # pyrefly: ignore[bad-function-definition]
class_agnostic_bbox_pred: bool = False,
cascade_class_ensemble: bool = False,
min_level: Optional[int] = None,
@@ -50,13 +50,14 @@ def __init__(
num_scales: Optional[int] = None,
aspect_ratios: Optional[List[float]] = None,
anchor_size: Optional[float] = None,
+ outer_boxes_scale: float = 1.0,
use_gt_boxes_for_masks: bool = False, # pytype: disable=annotation-type-mismatch # typed-keras
**kwargs):
"""Initializes the Panoptic Mask R-CNN model.
Args:
- backbone: `tf.keras.Model`, the backbone network.
- decoder: `tf.keras.Model`, the decoder network.
+ backbone: `tf_keras.Model`, the backbone network.
+ decoder: `tf_keras.Model`, the decoder network.
rpn_head: the RPN head.
detection_head: the detection head or a list of heads.
roi_generator: the ROI generator.
@@ -69,12 +70,12 @@ def __init__(
mask_head: the mask head.
mask_sampler: the mask sampler.
mask_roi_aligner: the ROI alginer for mask prediction.
- segmentation_backbone: `tf.keras.Model`, the backbone network for the
+ segmentation_backbone: `tf_keras.Model`, the backbone network for the
segmentation head for panoptic task. Providing `segmentation_backbone`
will allow the segmentation head to use a standlone backbone. Setting
`segmentation_backbone=None` would enable backbone sharing between the
MaskRCNN model and segmentation head.
- segmentation_decoder: `tf.keras.Model`, the decoder network for the
+ segmentation_decoder: `tf_keras.Model`, the decoder network for the
segmentation head for panoptic task. Providing `segmentation_decoder`
will allow the segmentation head to use a standlone decoder. Setting
`segmentation_decoder=None` would enable decoder sharing between the
@@ -95,10 +96,12 @@ def __init__(
aspect_ratios=[1.0, 2.0, 0.5] adds three anchors on each scale level.
anchor_size: A number representing the scale of size of the base anchor to
the feature stride 2^level.
+ outer_boxes_scale: a float to scale up the bounding boxes to generate
+ more inclusive masks. The scale is expected to be >=1.0.
use_gt_boxes_for_masks: `bool`, whether to use only gt boxes for masks.
**kwargs: keyword arguments to be passed.
"""
- super(PanopticMaskRCNNModel, self).__init__(
+ super().__init__(
backbone=backbone,
decoder=decoder,
rpn_head=rpn_head,
@@ -117,6 +120,7 @@ def __init__(
num_scales=num_scales,
aspect_ratios=aspect_ratios,
anchor_size=anchor_size,
+ outer_boxes_scale=outer_boxes_scale,
use_gt_boxes_for_masks=use_gt_boxes_for_masks,
**kwargs)
@@ -150,16 +154,21 @@ def call(self,
gt_boxes: Optional[tf.Tensor] = None,
gt_classes: Optional[tf.Tensor] = None,
gt_masks: Optional[tf.Tensor] = None,
+ gt_outer_boxes: Optional[tf.Tensor] = None,
training: Optional[bool] = None) -> Mapping[str, tf.Tensor]:
image_shape = image_info[:, 1, :]
- model_outputs = super(PanopticMaskRCNNModel, self).call(
- images=images,
- image_shape=image_shape,
- anchor_boxes=anchor_boxes,
- gt_boxes=gt_boxes,
- gt_classes=gt_classes,
- gt_masks=gt_masks,
- training=training)
+ model_kwargs = {
+ 'images': images,
+ 'image_shape': image_shape,
+ 'anchor_boxes': anchor_boxes,
+ 'gt_boxes': gt_boxes,
+ 'gt_classes': gt_classes,
+ 'gt_masks': gt_masks,
+ 'training': training,
+ }
+ if self.outer_boxes_scale > 1.0:
+ model_kwargs['gt_outer_boxes'] = gt_outer_boxes
+ model_outputs = super().call(**model_kwargs)
if self.segmentation_backbone is not None:
backbone_features = self.segmentation_backbone(images, training=training)
@@ -188,9 +197,9 @@ def call(self,
@property
def checkpoint_items(
- self) -> Mapping[str, Union[tf.keras.Model, tf.keras.layers.Layer]]:
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
"""Returns a dictionary of items to be additionally checkpointed."""
- items = super(PanopticMaskRCNNModel, self).checkpoint_items
+ items = super().checkpoint_items
if self.segmentation_backbone is not None:
items.update(segmentation_backbone=self.segmentation_backbone)
if self.segmentation_decoder is not None:
diff --git a/official/projects/panoptic/ops/mask_ops.py b/official/projects/panoptic/ops/mask_ops.py
new file mode 100644
index 00000000000..9d91f0f7c16
--- /dev/null
+++ b/official/projects/panoptic/ops/mask_ops.py
@@ -0,0 +1,55 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utility functions for masks."""
+
+import tensorflow as tf, tf_keras
+
+
+def resize_and_rescale_offsets(input_tensor: tf.Tensor, target_size):
+ """Bilinearly resizes and rescales the offsets.
+
+ Reference:
+ https://github.com/google-research/deeplab2/blob/main/model/utils.py#L157
+
+ Args:
+ input_tensor: A tf.Tensor of shape [batch, height, width, 2].
+ target_size: A list or tuple or 1D tf.Tensor that specifies the height and
+ width after resizing.
+
+ Returns:
+ The input_tensor resized to shape `[batch, target_height, target_width, 2]`.
+ Moreover, the offsets along the y-axis are rescaled by a factor equal to
+ (target_height - 1) / (reference_height - 1) and the offsets along the
+ x-axis are rescaled by a factor equal to
+ (target_width - 1) / (reference_width - 1).
+ """
+ input_size_y = tf.shape(input_tensor)[1]
+ input_size_x = tf.shape(input_tensor)[2]
+ dtype = input_tensor.dtype
+
+ scale_y = tf.cast(target_size[0] - 1, dtype=dtype) / tf.cast(
+ input_size_y - 1, dtype=dtype)
+ scale_x = tf.cast(target_size[1] - 1, dtype=dtype) / tf.cast(
+ input_size_x - 1, dtype=dtype)
+
+ target_y, target_x = tf.split(
+ value=input_tensor, num_or_size_splits=2, axis=3)
+ target_y *= scale_y
+ target_x *= scale_x
+ _ = tf.concat([target_y, target_x], 3)
+ return tf.image.resize(
+ input_tensor,
+ size=target_size,
+ method=tf.image.ResizeMethod.BILINEAR)
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/serving/export_saved_model.py b/official/projects/panoptic/serving/export_saved_model.py
similarity index 72%
rename from official/vision/beta/projects/panoptic_maskrcnn/serving/export_saved_model.py
rename to official/projects/panoptic/serving/export_saved_model.py
index 2a0579e9154..6df5677c6d7 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/serving/export_saved_model.py
+++ b/official/projects/panoptic/serving/export_saved_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -35,18 +35,27 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import exp_factory
from official.modeling import hyperparams
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as cfg # pylint: disable=unused-import
-from official.vision.beta.projects.panoptic_maskrcnn.modeling import factory
-from official.vision.beta.projects.panoptic_maskrcnn.serving import panoptic_segmentation
-from official.vision.beta.projects.panoptic_maskrcnn.tasks import panoptic_maskrcnn as task # pylint: disable=unused-import
+# pylint: disable=unused-import
+from official.projects.panoptic.configs import panoptic_deeplab as panoptic_deeplab_cfg
+from official.projects.panoptic.configs import panoptic_maskrcnn as panoptic_maskrcnn_cfg
+# pylint: enable=unused-import
+from official.projects.panoptic.modeling import factory
+from official.projects.panoptic.serving import panoptic_deeplab
+from official.projects.panoptic.serving import panoptic_maskrcnn
+# pylint: disable=unused-import
+from official.projects.panoptic.tasks import panoptic_deeplab as panoptic_deeplab_task
+from official.projects.panoptic.tasks import panoptic_maskrcnn as panoptic_maskrcnn_task
+# pylint: enable=unused-import
from official.vision.serving import export_saved_model_lib
FLAGS = flags.FLAGS
+flags.DEFINE_string('model', 'panoptic_maskrcnn',
+ 'model type, one of panoptic_maskrcnn and panoptic_deeplab')
flags.DEFINE_string('experiment', 'panoptic_fpn_coco',
'experiment type, e.g. panoptic_fpn_coco')
flags.DEFINE_string('export_dir', None, 'The export directory.')
@@ -87,18 +96,25 @@ def main(_):
params.lock()
input_image_size = [int(x) for x in FLAGS.input_image_size.split(',')]
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[FLAGS.batch_size, *input_image_size, 3])
- model = factory.build_panoptic_maskrcnn(
- input_specs=input_specs, model_config=params.task.model)
- export_module = panoptic_segmentation.PanopticSegmentationModule(
+ if FLAGS.model == 'panoptic_deeplab':
+ build_model = factory.build_panoptic_deeplab
+ panoptic_module = panoptic_deeplab.PanopticSegmentationModule
+ elif FLAGS.model == 'panoptic_maskrcnn':
+ build_model = factory.build_panoptic_maskrcnn
+ panoptic_module = panoptic_maskrcnn.PanopticSegmentationModule
+ else:
+ raise ValueError('Unsupported model type: %s' % FLAGS.model)
+
+ model = build_model(input_specs=input_specs, model_config=params.task.model)
+ export_module = panoptic_module(
params=params,
model=model,
batch_size=FLAGS.batch_size,
input_image_size=[int(x) for x in FLAGS.input_image_size.split(',')],
num_channels=3)
-
export_saved_model_lib.export_inference_graph(
input_type=FLAGS.input_type,
batch_size=FLAGS.batch_size,
@@ -110,6 +126,5 @@ def main(_):
export_checkpoint_subdir='checkpoint',
export_saved_model_subdir='saved_model')
-
if __name__ == '__main__':
app.run(main)
diff --git a/official/projects/panoptic/serving/panoptic_deeplab.py b/official/projects/panoptic/serving/panoptic_deeplab.py
new file mode 100644
index 00000000000..b95bac7bc0b
--- /dev/null
+++ b/official/projects/panoptic/serving/panoptic_deeplab.py
@@ -0,0 +1,103 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic Segmentation input and model functions for serving/inference."""
+
+from typing import List
+
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.projects.panoptic.modeling import factory
+from official.projects.panoptic.modeling import panoptic_deeplab_model
+from official.vision.serving import semantic_segmentation
+
+
+class PanopticSegmentationModule(
+ semantic_segmentation.SegmentationModule):
+ """Panoptic Deeplab Segmentation Module."""
+
+ def __init__(self,
+ params: cfg.ExperimentConfig,
+ *,
+ model: tf_keras.Model,
+ batch_size: int,
+ input_image_size: List[int],
+ num_channels: int = 3):
+ """Initializes panoptic segmentation module for export."""
+
+ if batch_size is None:
+ raise ValueError('batch_size cannot be None for panoptic segmentation '
+ 'model.')
+ if not isinstance(model, panoptic_deeplab_model.PanopticDeeplabModel):
+ raise ValueError('PanopticSegmentationModule module not '
+ 'implemented for {} model.'.format(type(model)))
+ params.task.train_data.preserve_aspect_ratio = True
+ super(PanopticSegmentationModule, self).__init__(
+ params=params,
+ model=model,
+ batch_size=batch_size,
+ input_image_size=input_image_size,
+ num_channels=num_channels)
+
+ def _build_model(self):
+ input_specs = tf_keras.layers.InputSpec(shape=[self._batch_size] +
+ self._input_image_size + [3])
+
+ return factory.build_panoptic_deeplab(
+ input_specs=input_specs,
+ model_config=self.params.task.model,
+ l2_regularizer=None)
+
+ def serve(self, images: tf.Tensor):
+ """Cast image to float and run inference.
+
+ Args:
+ images: uint8 Tensor of shape [batch_size, None, None, 3]
+
+ Returns:
+ Tensor holding detection output logits.
+ """
+ if self._input_type != 'tflite':
+ with tf.device('cpu:0'):
+ images = tf.cast(images, dtype=tf.float32)
+ images_spec = tf.TensorSpec(
+ shape=self._input_image_size + [3], dtype=tf.float32)
+ image_info_spec = tf.TensorSpec(shape=[4, 2], dtype=tf.float32)
+
+ images, image_info = tf.nest.map_structure(
+ tf.identity,
+ tf.map_fn(
+ self._build_inputs,
+ elems=images,
+ fn_output_signature=(images_spec, image_info_spec),
+ parallel_iterations=32))
+
+ outputs = self.model.call(
+ inputs=images, image_info=image_info, training=False) # pyrefly: ignore[unbound-name]
+
+ masks = outputs['segmentation_outputs']
+ masks = tf.image.resize(masks, self._input_image_size, method='bilinear')
+ classes = tf.math.argmax(masks, axis=-1)
+ scores = tf.nn.softmax(masks, axis=-1)
+ final_outputs = {
+ 'semantic_logits': masks,
+ 'semantic_scores': scores,
+ 'semantic_classes': classes,
+ 'image_info': image_info,
+ 'panoptic_category_mask': outputs['category_mask'],
+ 'panoptic_instance_mask': outputs['instance_mask'],
+ }
+
+ return final_outputs
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/serving/panoptic_segmentation.py b/official/projects/panoptic/serving/panoptic_maskrcnn.py
similarity index 79%
rename from official/vision/beta/projects/panoptic_maskrcnn/serving/panoptic_segmentation.py
rename to official/projects/panoptic/serving/panoptic_maskrcnn.py
index 0987dc417d9..0be2e1b5836 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/serving/panoptic_segmentation.py
+++ b/official/projects/panoptic/serving/panoptic_maskrcnn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,10 +16,10 @@
from typing import List
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
-from official.vision.beta.projects.panoptic_maskrcnn.modeling import panoptic_maskrcnn_model
+from official.projects.panoptic.modeling import panoptic_maskrcnn_model
from official.vision.serving import detection
@@ -29,7 +29,7 @@ class PanopticSegmentationModule(detection.DetectionModule):
def __init__(self,
params: cfg.ExperimentConfig,
*,
- model: tf.keras.Model,
+ model: tf_keras.Model,
batch_size: int,
input_image_size: List[int],
num_channels: int = 3):
@@ -42,7 +42,7 @@ def __init__(self,
raise ValueError('PanopticSegmentationModule module not implemented for '
'{} model.'.format(type(model)))
- super(PanopticSegmentationModule, self).__init__(
+ super().__init__(
params=params,
model=model,
batch_size=batch_size,
@@ -50,7 +50,7 @@ def __init__(self,
num_channels=num_channels)
def serve(self, images: tf.Tensor):
- """Cast image to float and run inference.
+ """Casts image to float and run inference.
Args:
images: uint8 Tensor of shape [batch_size, None, None, 3]
@@ -103,21 +103,28 @@ def serve(self, images: tf.Tensor):
detections.pop('box_outputs')
detections.pop('backbone_features')
detections.pop('decoder_features')
-
- # Normalize detection boxes to [0, 1]. Here we first map them to the
- # original image size, then normalize them to [0, 1].
- detections['detection_boxes'] = (
- detections['detection_boxes'] /
- tf.tile(image_info[:, 2:3, :], [1, 1, 2]) /
- tf.tile(image_info[:, 0:1, :], [1, 1, 2]))
-
if model_params.detection_generator.apply_nms:
+ # Normalize detection boxes to [0, 1]. Here we first map them to the
+ # original image size, then normalize them to [0, 1].
+ detections['detection_boxes'] = (
+ detections['detection_boxes'] /
+ tf.tile(image_info[:, 2:3, :], [1, 1, 2]) /
+ tf.tile(image_info[:, 0:1, :], [1, 1, 2]))
+
final_outputs = {
'detection_boxes': detections['detection_boxes'],
'detection_scores': detections['detection_scores'],
'detection_classes': detections['detection_classes'],
'num_detections': detections['num_detections']
}
+
+ if 'detection_outer_boxes' in detections:
+ detections['detection_outer_boxes'] = (
+ detections['detection_outer_boxes'] /
+ tf.tile(image_info[:, 2:3, :], [1, 1, 2]) /
+ tf.tile(image_info[:, 0:1, :], [1, 1, 2]))
+ final_outputs['detection_outer_boxes'] = (
+ detections['detection_outer_boxes'])
else:
final_outputs = {
'decoded_boxes': detections['decoded_boxes'],
@@ -126,12 +133,15 @@ def serve(self, images: tf.Tensor):
masks = detections['segmentation_outputs']
masks = tf.image.resize(masks, self._input_image_size, method='bilinear')
classes = tf.math.argmax(masks, axis=-1)
- scores = tf.nn.softmax(masks, axis=-1)
+ if self.params.task.losses.semantic_segmentation_use_binary_cross_entropy:
+ scores = tf.nn.sigmoid(masks)
+ else:
+ scores = tf.nn.softmax(masks, axis=-1)
final_outputs.update({
'detection_masks': detections['detection_masks'],
- 'masks': masks,
- 'scores': scores,
- 'classes': classes,
+ 'semantic_logits': masks,
+ 'semantic_scores': scores,
+ 'semantic_classes': classes,
'image_info': image_info
})
if model_params.generate_panoptic_masks:
diff --git a/official/projects/panoptic/tasks/__init__.py b/official/projects/panoptic/tasks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/panoptic/tasks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/panoptic/tasks/panoptic_deeplab.py b/official/projects/panoptic/tasks/panoptic_deeplab.py
new file mode 100644
index 00000000000..bb89d4ca6fe
--- /dev/null
+++ b/official/projects/panoptic/tasks/panoptic_deeplab.py
@@ -0,0 +1,393 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic Deeplab task definition."""
+from typing import Any, Dict, List, Mapping, Optional, Tuple
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import base_task
+from official.core import task_factory
+from official.projects.panoptic.configs import panoptic_deeplab as exp_cfg
+from official.projects.panoptic.dataloaders import panoptic_deeplab_input
+from official.projects.panoptic.losses import panoptic_deeplab_losses
+from official.projects.panoptic.modeling import factory
+from official.vision.dataloaders import input_reader_factory
+from official.vision.evaluation import panoptic_quality_evaluator
+from official.vision.evaluation import segmentation_metrics
+
+
+@task_factory.register_task_cls(exp_cfg.PanopticDeeplabTask)
+class PanopticDeeplabTask(base_task.Task):
+ """A task for Panoptic Deeplab."""
+
+ def build_model(self):
+ """Builds panoptic deeplab model."""
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + self.task_config.model.input_size)
+
+ l2_weight_decay = self.task_config.losses.l2_weight_decay
+ # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
+ # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
+ # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
+ l2_regularizer = (tf_keras.regularizers.l2(
+ l2_weight_decay / 2.0) if l2_weight_decay else None)
+
+ model = factory.build_panoptic_deeplab(
+ input_specs=input_specs,
+ model_config=self.task_config.model,
+ l2_regularizer=l2_regularizer)
+
+ # Builds the model through warm-up call.
+ dummy_images = tf_keras.Input(self.task_config.model.input_size)
+ # Note that image_info is always in the shape of [4, 2].
+ dummy_image_info = tf_keras.layers.Input([4, 2])
+ _ = model(dummy_images, dummy_image_info, training=False)
+ return model
+
+ def initialize(self, model: tf_keras.Model):
+ """Loads pretrained checkpoint."""
+ if not self.task_config.init_checkpoint:
+ return
+
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+
+ # Restoring checkpoint.
+ if 'all' in self.task_config.init_checkpoint_modules:
+ ckpt = tf.train.Checkpoint(**model.checkpoint_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ else:
+ ckpt_items = {}
+ if 'backbone' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(backbone=model.backbone)
+ if 'decoder' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(semantic_decoder=model.semantic_decoder)
+ if not self.task_config.model.shared_decoder:
+ ckpt_items.update(instance_decoder=model.instance_decoder)
+
+ ckpt = tf.train.Checkpoint(**ckpt_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
+
+ def build_inputs(self,
+ params: exp_cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None):
+ """Builds panoptic deeplab input."""
+ decoder_cfg = params.decoder.get()
+
+ if params.decoder.type == 'simple_decoder':
+ decoder = panoptic_deeplab_input.TfExampleDecoder(
+ regenerate_source_id=decoder_cfg.regenerate_source_id,
+ panoptic_category_mask_key=decoder_cfg.panoptic_category_mask_key,
+ panoptic_instance_mask_key=decoder_cfg.panoptic_instance_mask_key)
+ else:
+ raise ValueError('Unknown decoder type: {}!'.format(params.decoder.type))
+
+ parser = panoptic_deeplab_input.Parser(
+ output_size=self.task_config.model.input_size[:2],
+ ignore_label=params.parser.ignore_label,
+ resize_eval_groundtruth=params.parser.resize_eval_groundtruth,
+ groundtruth_padded_size=params.parser.groundtruth_padded_size,
+ aug_scale_min=params.parser.aug_scale_min,
+ aug_scale_max=params.parser.aug_scale_max,
+ aug_rand_hflip=params.parser.aug_rand_hflip,
+ aug_type=params.parser.aug_type,
+ sigma=params.parser.sigma,
+ dtype=params.parser.dtype)
+
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training))
+
+ dataset = reader.read(input_context=input_context)
+
+ return dataset
+
+ def build_losses(self,
+ labels: Mapping[str, tf.Tensor],
+ model_outputs: Mapping[str, tf.Tensor],
+ aux_losses: Optional[Any] = None):
+ """Panoptic deeplab losses.
+
+ Args:
+ labels: labels.
+ model_outputs: Output logits from panoptic deeplab.
+ aux_losses: auxiliarly loss tensors, i.e. `losses` in keras.Model.
+
+ Returns:
+ The total loss tensor.
+ """
+ loss_config = self._task_config.losses
+ segmentation_loss_fn = (
+ panoptic_deeplab_losses.WeightedBootstrappedCrossEntropyLoss(
+ loss_config.label_smoothing,
+ loss_config.class_weights,
+ loss_config.ignore_label,
+ top_k_percent_pixels=loss_config.top_k_percent_pixels))
+ instance_center_heatmap_loss_fn = panoptic_deeplab_losses.CenterHeatmapLoss(
+ )
+ instance_center_offset_loss_fn = panoptic_deeplab_losses.CenterOffsetLoss()
+
+ semantic_weights = tf.cast(
+ labels['semantic_weights'],
+ dtype=model_outputs['instance_centers_heatmap'].dtype)
+ things_mask = tf.cast(
+ tf.squeeze(labels['things_mask'], axis=3),
+ dtype=model_outputs['instance_centers_heatmap'].dtype)
+ valid_mask = tf.cast(
+ tf.squeeze(labels['valid_mask'], axis=3),
+ dtype=model_outputs['instance_centers_heatmap'].dtype)
+
+ segmentation_loss = segmentation_loss_fn(
+ model_outputs['segmentation_outputs'],
+ labels['category_mask'],
+ sample_weight=semantic_weights)
+ instance_center_heatmap_loss = instance_center_heatmap_loss_fn(
+ model_outputs['instance_centers_heatmap'],
+ labels['instance_centers_heatmap'],
+ sample_weight=valid_mask)
+ instance_center_offset_loss = instance_center_offset_loss_fn(
+ model_outputs['instance_centers_offset'],
+ labels['instance_centers_offset'],
+ sample_weight=things_mask)
+
+ model_loss = (
+ loss_config.segmentation_loss_weight * segmentation_loss +
+ loss_config.center_heatmap_loss_weight * instance_center_heatmap_loss +
+ loss_config.center_offset_loss_weight * instance_center_offset_loss)
+
+ total_loss = model_loss
+ if aux_losses:
+ total_loss += tf.add_n(aux_losses)
+
+ losses = {
+ 'total_loss': total_loss,
+ 'model_loss': model_loss,
+ 'segmentation_loss': segmentation_loss,
+ 'instance_center_heatmap_loss': instance_center_heatmap_loss,
+ 'instance_center_offset_loss': instance_center_offset_loss
+ }
+
+ return losses
+
+ def build_metrics(self, training: bool = True) -> List[
+ tf_keras.metrics.Metric]:
+ """Build metrics."""
+ eval_config = self.task_config.evaluation
+ metrics = []
+ if training:
+ metric_names = [
+ 'total_loss',
+ 'segmentation_loss',
+ 'instance_center_heatmap_loss',
+ 'instance_center_offset_loss',
+ 'model_loss']
+ for name in metric_names:
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
+
+ if eval_config.report_train_mean_iou:
+ self.train_mean_iou = segmentation_metrics.MeanIoU(
+ name='train_mean_iou',
+ num_classes=self.task_config.model.num_classes,
+ rescale_predictions=False,
+ dtype=tf.float32)
+ else:
+ rescale_predictions = (not self.task_config.validation_data.parser
+ .resize_eval_groundtruth)
+ self.perclass_iou_metric = segmentation_metrics.PerClassIoU(
+ name='per_class_iou',
+ num_classes=self.task_config.model.num_classes,
+ rescale_predictions=rescale_predictions,
+ dtype=tf.float32)
+
+ if self.task_config.model.generate_panoptic_masks:
+ self.panoptic_quality_metric = (
+ panoptic_quality_evaluator.PanopticQualityEvaluator(
+ num_categories=self.task_config.model.num_classes,
+ ignored_label=eval_config.ignored_label,
+ max_instances_per_category=eval_config
+ .max_instances_per_category,
+ offset=eval_config.offset,
+ is_thing=eval_config.is_thing,
+ rescale_predictions=eval_config.rescale_predictions))
+
+ return metrics
+
+ def train_step(
+ self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
+ """Does forward and backward.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ images, labels = inputs
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+
+ with tf.GradientTape() as tape:
+ outputs = model(
+ inputs=images,
+ image_info=labels['image_info'],
+ training=True)
+ outputs = tf.nest.map_structure(
+ lambda x: tf.cast(x, tf.float32), outputs)
+
+ # Computes per-replica loss.
+ losses = self.build_losses(
+ labels=labels,
+ model_outputs=outputs,
+ aux_losses=model.losses)
+ scaled_loss = losses['total_loss'] / num_replicas
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scales back gradient when LossScaleOptimizer is used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ logs = {self.loss: losses['total_loss']}
+
+ if metrics:
+ for m in metrics:
+ m.update_state(losses[m.name])
+
+ if self.task_config.evaluation.report_train_mean_iou:
+ segmentation_labels = {
+ 'masks': labels['category_mask'],
+ 'valid_masks': labels['valid_mask'],
+ 'image_info': labels['image_info']
+ }
+ self.process_metrics(
+ metrics=[self.train_mean_iou],
+ labels=segmentation_labels,
+ model_outputs=outputs['segmentation_outputs'])
+ logs.update({
+ self.train_mean_iou.name:
+ self.train_mean_iou.result()
+ })
+
+ return logs
+
+ def validation_step(
+ self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ images, labels = inputs
+
+ outputs = model(
+ inputs=images,
+ image_info=labels['image_info'],
+ training=False)
+
+ logs = {self.loss: 0}
+ segmentation_labels = {
+ 'masks': labels['category_mask'],
+ 'valid_masks': labels['valid_mask'],
+ 'image_info': labels['image_info']
+ }
+
+ self.perclass_iou_metric.update_state(segmentation_labels,
+ outputs['segmentation_outputs'])
+
+ if self.task_config.model.generate_panoptic_masks:
+ pq_metric_labels = {
+ 'category_mask': tf.squeeze(labels['category_mask'], axis=3),
+ 'instance_mask': tf.squeeze(labels['instance_mask'], axis=3),
+ 'image_info': labels['image_info']
+ }
+ panoptic_outputs = {
+ 'category_mask':
+ outputs['category_mask'],
+ 'instance_mask':
+ outputs['instance_mask'],
+ }
+ logs.update({
+ self.panoptic_quality_metric.name:
+ (pq_metric_labels, panoptic_outputs)})
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ if state is None:
+ self.perclass_iou_metric.reset_states()
+ state = [self.perclass_iou_metric]
+ if self.task_config.model.generate_panoptic_masks:
+ state += [self.panoptic_quality_metric]
+
+ if self.task_config.model.generate_panoptic_masks:
+ self.panoptic_quality_metric.update_state(
+ step_outputs[self.panoptic_quality_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ step_outputs[self.panoptic_quality_metric.name][1]) # pyrefly: ignore[unsupported-operation]
+
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ result = {}
+ ious = self.perclass_iou_metric.result()
+ if self.task_config.evaluation.report_per_class_iou:
+ for i, value in enumerate(ious.numpy()):
+ result.update({'segmentation_iou/class_{}'.format(i): value})
+
+ # Computes mean IoU
+ result.update({'segmentation_mean_iou': tf.reduce_mean(ious).numpy()})
+
+ if self.task_config.model.generate_panoptic_masks:
+ panoptic_quality_results = self.panoptic_quality_metric.result()
+ for k, value in panoptic_quality_results.items():
+ if k.endswith('per_class'):
+ if self.task_config.evaluation.report_per_class_pq:
+ for i, per_class_value in enumerate(value):
+ metric_key = 'panoptic_quality/{}/class_{}'.format(k, i)
+ result[metric_key] = per_class_value
+ else:
+ continue
+ else:
+ result['panoptic_quality/{}'.format(k)] = value
+
+ return result
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/tasks/panoptic_maskrcnn.py b/official/projects/panoptic/tasks/panoptic_maskrcnn.py
similarity index 55%
rename from official/vision/beta/projects/panoptic_maskrcnn/tasks/panoptic_maskrcnn.py
rename to official/projects/panoptic/tasks/panoptic_maskrcnn.py
index 14137733c38..6b5c2287ea9 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/tasks/panoptic_maskrcnn.py
+++ b/official/projects/panoptic/tasks/panoptic_maskrcnn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,15 +16,16 @@
from typing import Any, Dict, List, Mapping, Optional, Tuple
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import dataset_fn
from official.core import task_factory
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as exp_cfg
-from official.vision.beta.projects.panoptic_maskrcnn.dataloaders import panoptic_maskrcnn_input
-from official.vision.beta.projects.panoptic_maskrcnn.modeling import factory
+from official.projects.panoptic.configs import panoptic_maskrcnn as exp_cfg
+from official.projects.panoptic.dataloaders import panoptic_maskrcnn_input
+from official.projects.panoptic.modeling import factory
+from official.vision.dataloaders import input_reader
from official.vision.dataloaders import input_reader_factory
-from official.vision.evaluation import panoptic_quality_evaluator
+from official.vision.evaluation import panoptic_quality
from official.vision.evaluation import segmentation_metrics
from official.vision.losses import segmentation_losses
from official.vision.tasks import maskrcnn
@@ -40,27 +41,46 @@ class PanopticMaskRCNNTask(maskrcnn.MaskRCNNTask):
the loss, post-processing, and customized metrics with reduction.
"""
- def build_model(self) -> tf.keras.Model:
- """Build Panoptic Mask R-CNN model."""
+ def __init__(self,
+ params,
+ logging_dir: Optional[str] = None,
+ name: Optional[str] = None):
+ super().__init__(params, logging_dir=logging_dir, name=name)
+ self.segmentation_train_mean_iou = None
+ self.segmentation_perclass_iou_metric = None
+ self.panoptic_quality_metric = None
- input_specs = tf.keras.layers.InputSpec(
+ def build_model(self) -> tf_keras.Model:
+ """Builds Panoptic Mask R-CNN model."""
+
+ input_specs = tf_keras.layers.InputSpec(
shape=[None] + self.task_config.model.input_size)
l2_weight_decay = self.task_config.losses.l2_weight_decay
# Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
- l2_regularizer = (tf.keras.regularizers.l2(
+ l2_regularizer = (tf_keras.regularizers.l2(
l2_weight_decay / 2.0) if l2_weight_decay else None)
model = factory.build_panoptic_maskrcnn(
input_specs=input_specs,
model_config=self.task_config.model,
l2_regularizer=l2_regularizer)
+
+ if self.task_config.freeze_backbone:
+ model.backbone.trainable = False
+
+ # Builds the model through warm-up call.
+ dummy_images = tf_keras.Input(self.task_config.model.input_size)
+ # Note that image_info is always in the shape of [4, 2].
+ dummy_image_info = tf_keras.layers.Input([4, 2])
+ _ = model(dummy_images, image_info=dummy_image_info, training=False)
+
return model
- def initialize(self, model: tf.keras.Model) -> None:
- """Loading pretrained checkpoint."""
+ def initialize(self, model: tf_keras.Model) -> None:
+ """Loads pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -84,7 +104,18 @@ def _get_checkpoint_path(checkpoint_dir_or_file):
elif init_module == 'backbone':
checkpoint_path = _get_checkpoint_path(
self.task_config.init_checkpoint)
- ckpt = tf.train.Checkpoint(backbone=model.backbone)
+
+ if self.task_config.model.backbone.type == 'uvit':
+ model.backbone.load_checkpoint(ckpt_filepath=checkpoint_path)
+ else:
+ ckpt = tf.train.Checkpoint(backbone=model.backbone)
+ status = ckpt.read(checkpoint_path)
+ status.expect_partial().assert_existing_objects_matched()
+
+ elif init_module == 'decoder':
+ checkpoint_path = _get_checkpoint_path(
+ self.task_config.init_checkpoint)
+ ckpt = tf.train.Checkpoint(decoder=model.decoder)
status = ckpt.read(checkpoint_path)
status.expect_partial().assert_existing_objects_matched()
@@ -106,8 +137,8 @@ def _get_checkpoint_path(checkpoint_dir_or_file):
else:
raise ValueError(
- "Only 'all', 'backbone', 'segmentation_backbone' and/or "
- "segmentation_backbone' can be used to initialize the model, but "
+ "Only 'all', 'backbone', 'decoder', 'segmentation_backbone' and/or "
+ "'segmentation_decoder' can be used to initialize the model, but "
"got {}".format(init_module))
logging.info('Finished loading pretrained checkpoint from %s for %s',
checkpoint_path, init_module)
@@ -117,7 +148,7 @@ def build_inputs(
params: exp_cfg.DataConfig,
input_context: Optional[tf.distribute.InputContext] = None
) -> tf.data.Dataset:
- """Build input dataset."""
+ """Builds input dataset."""
decoder_cfg = params.decoder.get()
if params.decoder.type == 'simple_decoder':
decoder = panoptic_maskrcnn_input.TfExampleDecoder(
@@ -136,16 +167,18 @@ def build_inputs(
num_scales=self.task_config.model.anchor.num_scales,
aspect_ratios=self.task_config.model.anchor.aspect_ratios,
anchor_size=self.task_config.model.anchor.anchor_size,
- dtype=params.dtype,
rpn_match_threshold=params.parser.rpn_match_threshold,
rpn_unmatched_threshold=params.parser.rpn_unmatched_threshold,
rpn_batch_size_per_im=params.parser.rpn_batch_size_per_im,
rpn_fg_fraction=params.parser.rpn_fg_fraction,
aug_rand_hflip=params.parser.aug_rand_hflip,
+ aug_rand_vflip=params.parser.aug_rand_vflip,
aug_scale_min=params.parser.aug_scale_min,
aug_scale_max=params.parser.aug_scale_max,
+ aug_type=params.parser.aug_type,
skip_crowd_during_training=params.parser.skip_crowd_during_training,
max_num_instances=params.parser.max_num_instances,
+ outer_boxes_scale=self.task_config.model.outer_boxes_scale,
mask_crop_size=params.parser.mask_crop_size,
segmentation_resize_eval_groundtruth=params.parser
.segmentation_resize_eval_groundtruth,
@@ -153,13 +186,17 @@ def build_inputs(
.segmentation_groundtruth_padded_size,
segmentation_ignore_label=params.parser.segmentation_ignore_label,
panoptic_ignore_label=params.parser.panoptic_ignore_label,
- include_panoptic_masks=params.parser.include_panoptic_masks)
+ include_panoptic_masks=params.parser.include_panoptic_masks,
+ dtype=params.dtype,
+ )
reader = input_reader_factory.input_reader_generator(
params,
dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
decoder_fn=decoder.decode,
- parser_fn=parser.parse_fn(params.is_training))
+ combine_fn=input_reader.create_combine_fn(params),
+ parser_fn=parser.parse_fn(params.is_training),
+ )
dataset = reader.read(input_context=input_context)
return dataset
@@ -168,22 +205,26 @@ def build_losses(self,
outputs: Mapping[str, Any],
labels: Mapping[str, Any],
aux_losses: Optional[Any] = None) -> Dict[str, tf.Tensor]:
- """Build Panoptic Mask R-CNN losses."""
+ """Builds Panoptic Mask R-CNN losses."""
params = self.task_config.losses
- use_groundtruth_dimension = params.semantic_segmentation_use_groundtruth_dimension
+ use_groundtruth_dimension = (
+ params.semantic_segmentation_use_groundtruth_dimension)
segmentation_loss_fn = segmentation_losses.SegmentationLoss(
label_smoothing=params.semantic_segmentation_label_smoothing,
class_weights=params.semantic_segmentation_class_weights,
ignore_label=params.semantic_segmentation_ignore_label,
+ gt_is_matting_map=params.semantic_segmentation_gt_is_matting_map,
use_groundtruth_dimension=use_groundtruth_dimension,
+ use_binary_cross_entropy=params
+ .semantic_segmentation_use_binary_cross_entropy,
top_k_percent_pixels=params.semantic_segmentation_top_k_percent_pixels)
instance_segmentation_weight = params.instance_segmentation_weight
semantic_segmentation_weight = params.semantic_segmentation_weight
- losses = super(PanopticMaskRCNNTask, self).build_losses(
+ losses = super().build_losses(
outputs=outputs,
labels=labels,
aux_losses=None)
@@ -209,69 +250,58 @@ def build_losses(self,
})
return losses
- def build_metrics(self, training: bool = True) -> List[
- tf.keras.metrics.Metric]:
- """Build detection metrics."""
- metrics = []
- num_segmentation_classes = self.task_config.model.segmentation_model.num_classes
+ def build_metrics(
+ self, training: bool = True
+ ) -> List[tf_keras.metrics.Metric]:
+ """Builds detection metrics."""
+ metrics = super().build_metrics(training)
+
if training:
- metric_names = [
- 'total_loss',
- 'rpn_score_loss',
- 'rpn_box_loss',
- 'frcnn_cls_loss',
- 'frcnn_box_loss',
- 'mask_loss',
- 'maskrcnn_loss',
- 'segmentation_loss',
- 'model_loss'
- ]
+ metric_names = ['maskrcnn_loss', 'segmentation_loss']
for name in metric_names:
- metrics.append(tf.keras.metrics.Mean(name, dtype=tf.float32))
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
if self.task_config.segmentation_evaluation.report_train_mean_iou:
self.segmentation_train_mean_iou = segmentation_metrics.MeanIoU(
name='train_mean_iou',
- num_classes=num_segmentation_classes,
+ num_classes=self.task_config.model.segmentation_model.num_classes,
rescale_predictions=False,
- dtype=tf.float32)
-
+ dtype=tf.float32,
+ )
else:
- self._build_coco_metrics()
-
- rescale_predictions = (not self.task_config.validation_data.parser
- .segmentation_resize_eval_groundtruth)
-
+ rescale_predictions = (
+ not self.task_config.validation_data.parser.segmentation_resize_eval_groundtruth
+ )
self.segmentation_perclass_iou_metric = segmentation_metrics.PerClassIoU(
name='per_class_iou',
- num_classes=num_segmentation_classes,
+ num_classes=self.task_config.model.segmentation_model.num_classes,
rescale_predictions=rescale_predictions,
- dtype=tf.float32)
+ dtype=tf.float32,
+ )
- if isinstance(tf.distribute.get_strategy(), tf.distribute.TPUStrategy):
- self._process_iou_metric_on_cpu = True
- else:
- self._process_iou_metric_on_cpu = False
-
- if self.task_config.model.generate_panoptic_masks:
+ if (
+ self.task_config.model.generate_panoptic_masks
+ and self.task_config.panoptic_quality_evaluator is not None
+ ):
if not self.task_config.validation_data.parser.include_panoptic_masks:
- raise ValueError('`include_panoptic_masks` should be set to True when'
- ' computing panoptic quality.')
+ raise ValueError(
+ '`include_panoptic_masks` should be set to True when'
+ ' computing panoptic quality.'
+ )
pq_config = self.task_config.panoptic_quality_evaluator
- self.panoptic_quality_metric = panoptic_quality_evaluator.PanopticQualityEvaluator(
+ self.panoptic_quality_metric = panoptic_quality.PanopticQualityV2(
num_categories=pq_config.num_categories,
- ignored_label=pq_config.ignored_label,
- max_instances_per_category=pq_config.max_instances_per_category,
- offset=pq_config.offset,
is_thing=pq_config.is_thing,
- rescale_predictions=pq_config.rescale_predictions)
+ ignored_label=pq_config.ignored_label,
+ rescale_predictions=pq_config.rescale_predictions,
+ )
return metrics
def train_step(self,
inputs: Tuple[Any, Any],
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
"""Does forward and backward.
@@ -288,15 +318,18 @@ def train_step(self,
num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
with tf.GradientTape() as tape:
- outputs = model(
- images,
- image_info=labels['image_info'],
- anchor_boxes=labels['anchor_boxes'],
- gt_boxes=labels['gt_boxes'],
- gt_classes=labels['gt_classes'],
- gt_masks=(labels['gt_masks'] if self.task_config.model.include_mask
- else None),
- training=True)
+ model_kwargs = {
+ 'image_info': labels['image_info'],
+ 'anchor_boxes': labels['anchor_boxes'],
+ 'gt_boxes': labels['gt_boxes'],
+ 'gt_classes': labels['gt_classes'],
+ 'training': True,
+ }
+ if self.task_config.model.include_mask:
+ model_kwargs['gt_masks'] = labels['gt_masks']
+ if self.task_config.model.outer_boxes_scale > 1.0:
+ model_kwargs['gt_outer_boxes'] = labels['gt_outer_boxes']
+ outputs = model(images, **model_kwargs)
outputs = tf.nest.map_structure(
lambda x: tf.cast(x, tf.float32), outputs)
@@ -307,13 +340,13 @@ def train_step(self,
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
# Scales back gradient when LossScaleOptimizer is used.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
@@ -323,7 +356,8 @@ def train_step(self,
for m in metrics:
m.update_state(losses[m.name])
- if self.task_config.segmentation_evaluation.report_train_mean_iou:
+ if (self.task_config.segmentation_evaluation.report_train_mean_iou and
+ self.segmentation_train_mean_iou is not None):
segmentation_labels = {
'masks': labels['gt_segmentation_mask'],
'valid_masks': labels['gt_segmentation_valid_mask'],
@@ -340,10 +374,35 @@ def train_step(self,
return logs
- def validation_step(self,
- inputs: Tuple[Any, Any],
- model: tf.keras.Model,
- metrics: Optional[List[Any]] = None) -> Dict[str, Any]:
+ def _update_metrics(self, labels, outputs, logs):
+ super()._update_metrics(labels, outputs, logs)
+
+ if self.segmentation_perclass_iou_metric is not None:
+ segmentation_labels = {
+ 'masks': labels['groundtruths']['gt_segmentation_mask'],
+ 'valid_masks': labels['groundtruths']['gt_segmentation_valid_mask'],
+ 'image_info': labels['image_info'],
+ }
+ self.segmentation_perclass_iou_metric.update_state(
+ segmentation_labels, outputs['segmentation_outputs']
+ )
+
+ if self.panoptic_quality_metric is not None:
+ pq_metric_labels = {
+ 'category_mask': labels['groundtruths']['gt_panoptic_category_mask'],
+ 'instance_mask': labels['groundtruths']['gt_panoptic_instance_mask'],
+ 'image_info': labels['image_info'],
+ }
+ self.panoptic_quality_metric.update_state(
+ pq_metric_labels, outputs['panoptic_outputs']
+ )
+
+ def validation_step(
+ self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ metrics: Optional[List[Any]] = None,
+ ) -> Dict[str, Any]:
"""Validatation step.
Args:
@@ -360,99 +419,91 @@ def validation_step(self,
images,
anchor_boxes=labels['anchor_boxes'],
image_info=labels['image_info'],
- training=False)
+ training=False,
+ )
logs = {self.loss: 0}
- coco_model_outputs = {
- 'detection_masks': outputs['detection_masks'],
- 'detection_boxes': outputs['detection_boxes'],
- 'detection_scores': outputs['detection_scores'],
- 'detection_classes': outputs['detection_classes'],
- 'num_detections': outputs['num_detections'],
- 'source_id': labels['groundtruths']['source_id'],
- 'image_info': labels['image_info']
- }
- segmentation_labels = {
- 'masks': labels['groundtruths']['gt_segmentation_mask'],
- 'valid_masks': labels['groundtruths']['gt_segmentation_valid_mask'],
- 'image_info': labels['image_info']
- }
-
- logs.update(
- {self.coco_metric.name: (labels['groundtruths'], coco_model_outputs)})
- if self._process_iou_metric_on_cpu:
- logs.update({
- self.segmentation_perclass_iou_metric.name:
- (segmentation_labels, outputs['segmentation_outputs'])
- })
- else:
- self.segmentation_perclass_iou_metric.update_state(
- segmentation_labels,
- outputs['segmentation_outputs'])
-
- if self.task_config.model.generate_panoptic_masks:
- pq_metric_labels = {
- 'category_mask':
- labels['groundtruths']['gt_panoptic_category_mask'],
- 'instance_mask':
- labels['groundtruths']['gt_panoptic_instance_mask'],
- 'image_info': labels['image_info']
- }
- logs.update({
- self.panoptic_quality_metric.name:
- (pq_metric_labels, outputs['panoptic_outputs'])})
+ self._update_metrics(labels, outputs, logs)
return logs
def aggregate_logs(self, state=None, step_outputs=None):
- if state is None:
- self.coco_metric.reset_states()
- self.segmentation_perclass_iou_metric.reset_states()
- state = [self.coco_metric, self.segmentation_perclass_iou_metric]
- if self.task_config.model.generate_panoptic_masks:
- state += [self.panoptic_quality_metric]
-
- self.coco_metric.update_state(
- step_outputs[self.coco_metric.name][0],
- step_outputs[self.coco_metric.name][1])
-
- if self._process_iou_metric_on_cpu:
- self.segmentation_perclass_iou_metric.update_state(
- step_outputs[self.segmentation_perclass_iou_metric.name][0],
- step_outputs[self.segmentation_perclass_iou_metric.name][1])
-
- if self.task_config.model.generate_panoptic_masks:
- self.panoptic_quality_metric.update_state(
- step_outputs[self.panoptic_quality_metric.name][0],
- step_outputs[self.panoptic_quality_metric.name][1])
-
+ is_first_step = not state
+ super().aggregate_logs(state, step_outputs)
+
+ if is_first_step:
+ if not isinstance(state, list):
+ state = []
+ if self.segmentation_perclass_iou_metric is not None:
+ state.append(self.segmentation_perclass_iou_metric)
+ if self.panoptic_quality_metric is not None:
+ state.append(self.panoptic_quality_metric)
+
+ if not state:
+ # Create an arbitrary state to indicate it's not the first step in the
+ # following calls to this function.
+ state = True
return state
- def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
- result = {}
- result = super(
- PanopticMaskRCNNTask, self).reduce_aggregated_logs(
- aggregated_logs=aggregated_logs,
- global_step=global_step)
-
+ def _reduce_semantic_metrics(self, logs: Dict[str, Any]):
+ """Updates the per class and mean semantic metrics in the logs."""
+ assert self.segmentation_perclass_iou_metric is not None
ious = self.segmentation_perclass_iou_metric.result()
if self.task_config.segmentation_evaluation.report_per_class_iou:
for i, value in enumerate(ious.numpy()):
- result.update({'segmentation_iou/class_{}'.format(i): value})
- # Computes mean IoU
- result.update({'segmentation_mean_iou': tf.reduce_mean(ious).numpy()})
-
- if self.task_config.model.generate_panoptic_masks:
- report_per_class_metrics = self.task_config.panoptic_quality_evaluator.report_per_class_metrics
- panoptic_quality_results = self.panoptic_quality_metric.result()
- for k, value in panoptic_quality_results.items():
- if k.endswith('per_class'):
- if report_per_class_metrics:
- for i, per_class_value in enumerate(value):
- metric_key = 'panoptic_quality/{}/class_{}'.format(k, i)
- result[metric_key] = per_class_value
- else:
- continue
- else:
- result['panoptic_quality/{}'.format(k)] = value
+ logs.update({'segmentation_iou/class_{}'.format(i): value})
+ logs.update({'segmentation_mean_iou': tf.reduce_mean(ious)})
+
+ def _reduce_panoptic_metrics(self, logs: Dict[str, Any]):
+ """Updates the per class and mean panoptic metrics in the logs."""
+ assert self.panoptic_quality_metric is not None
+ result = self.panoptic_quality_metric.result()
+ valid_thing_classes = result['valid_thing_classes']
+ valid_stuff_classes = result['valid_stuff_classes']
+ valid_classes = valid_stuff_classes | valid_thing_classes
+ num_categories = tf.math.count_nonzero(valid_classes, dtype=tf.float32)
+ num_thing_categories = tf.math.count_nonzero(
+ valid_thing_classes, dtype=tf.float32
+ )
+ num_stuff_categories = tf.math.count_nonzero(
+ valid_stuff_classes, dtype=tf.float32
+ )
+ valid_thing_classes = tf.cast(valid_thing_classes, dtype=tf.float32)
+ valid_stuff_classes = tf.cast(valid_stuff_classes, dtype=tf.float32)
+
+ logs['panoptic_quality/All_num_categories'] = num_categories
+ logs['panoptic_quality/Things_num_categories'] = num_thing_categories
+ logs['panoptic_quality/Stuff_num_categories'] = num_stuff_categories
+ for metric in ['pq', 'sq', 'rq']:
+ metric_per_class = result[f'{metric}_per_class']
+ logs[f'panoptic_quality/All_{metric}'] = tf.math.divide_no_nan(
+ tf.reduce_sum(metric_per_class), num_categories
+ )
+ logs[f'panoptic_quality/Things_{metric}'] = tf.math.divide_no_nan(
+ tf.reduce_sum(metric_per_class * valid_thing_classes),
+ num_thing_categories,
+ )
+ logs[f'panoptic_quality/Stuff_{metric}'] = tf.math.divide_no_nan(
+ tf.reduce_sum(metric_per_class * valid_stuff_classes),
+ num_stuff_categories,
+ )
+ if self.task_config.panoptic_quality_evaluator.report_per_class_metrics:
+ for i, is_valid in enumerate(valid_classes.numpy()):
+ if is_valid:
+ logs[f'panoptic_quality/{metric}/class_{i}'] = metric_per_class[i]
+
+ def reduce_aggregated_logs(
+ self,
+ aggregated_logs: Dict[str, Any],
+ global_step: Optional[tf.Tensor] = None,
+ ) -> Dict[str, tf.Tensor]:
+ """Optional reduce of aggregated logs over validation steps."""
+ logs = super().reduce_aggregated_logs(aggregated_logs, global_step)
+
+ if self.segmentation_perclass_iou_metric is not None:
+ self._reduce_semantic_metrics(logs)
+ self.segmentation_perclass_iou_metric.reset_state()
+ if self.panoptic_quality_metric is not None:
+ self._reduce_panoptic_metrics(logs)
+ self.panoptic_quality_metric.reset_state()
- return result
+ return logs
diff --git a/official/projects/panoptic/train.py b/official/projects/panoptic/train.py
new file mode 100644
index 00000000000..192be437e7d
--- /dev/null
+++ b/official/projects/panoptic/train.py
@@ -0,0 +1,32 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Panoptic MaskRCNN trainer."""
+
+from absl import app
+
+from official.common import flags as tfm_flags
+# pylint: disable=unused-import
+from official.projects.panoptic.configs import panoptic_deeplab
+from official.projects.panoptic.configs import panoptic_maskrcnn
+from official.projects.panoptic.tasks import panoptic_deeplab as panoptic_deeplab_task
+from official.projects.panoptic.tasks import panoptic_maskrcnn as panoptic_maskrcnn_task
+from official.projects.uvit import configs
+from official.projects.uvit import tasks
+from official.vision import train
+# pylint: enable=unused-import
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(train.main)
diff --git a/official/projects/perceiver/README.md b/official/projects/perceiver/README.md
new file mode 100644
index 00000000000..891500a2cd1
--- /dev/null
+++ b/official/projects/perceiver/README.md
@@ -0,0 +1,55 @@
+# Perceiver IO: A General Architecture for Structured Inputs & Outputs
+
+TF2 implementation of [Perceiver](https://arxiv.org/abs/2107.14795).
+
+## Default setup command:
+Scripts to pretrain, finetune, train from scratch can be found under
+perceiver/experiments.
+
+## BERT Wiki Books Pretrain
+
+Configurations can be seen on Table 8 and Table 9 of the
+ [paper](https://arxiv.org/abs/2107.14795). Our model configuration can be
+ deduced in the configs and experiment folder, where we follow the
+ configuration in the paper except for the tokenization and data.
+
+Model | Tokenizer | Pretrain Data | Batch Size | Steps | Val MLM Accuracy
+----- | --------: | ------------: | ---------: | ----: | ---------------:
+Perceiver IO Base (paper) | SentencePiece | T5 + Wiki | 512 | 500 k | N/A
+Perceiver IO Base (ours) | WordPiece | Wiki + Books | 512 | 500 k | 68.69 %
+
+## GLUE Finetune
+
+Our perceiver model is fine-tuned on GLUE upon the pre-trained model shown
+ above. These are all single-task fine-tuning only.
+
+These are run with configurations shown on Table 10 in the [paper](https://arxiv.org/abs/2107.14795).
+
+Model | Tokenizer | Pretrain Data | CoLA | MNLI-m/mm | MRPC | QNLI | QQP | RTE | SST-2 | STS-B | Average
+----- | --------: | ------------: | ---: | --------: | ----:| ----:| --: | --: | ----: | ----: | -----:
+Perceiver IO Base (paper) | SentencePiece | T5 + Wiki | 47.11 % | 84.53/85.03 % | 87.25 % | 92.12 % | 90.22 % | 65.23 % | 94.38 % | 88.18 % | 81.16 %
+Perceiver IO Base (ours) | WordPiece | Wiki + Books | 63.23 % | 84.29/84.52 % | 87.74 % | 91.43 % | 91.22 % | 70.76 % | 94.15 % | 89.85 % | 84.09 %
+
+Note: The average is computed by first averaging the results of MNLI-matched and
+MNLI-mismatched, which is then counted as a single task in the overall average.
+
+`Average = (63.23 + (84.29 + 84.52) / 2 + 87.74 + 91.43 + 91.22 + 70.76 + 94.15 + 89.85) / 8`
+
+## Discrepancy with the paper:
+
+* ~+2.93 average GLUE accuracy compared to paper results.
+
+## Citing TensorFlow Model Garden
+
+If you find this codebase helpful in your research, please cite this repository.
+
+```
+@misc{tensorflowmodelgarden2022,
+ author = {Hongkun Yu and Chen Chen and Xianzhi Du and Yeqing Li and
+ Abdullah Rashwan and Le Hou and Pengchong Jin and Fan Yang and
+ Frederick Liu and Jaeyoun Kim and Jing Li},
+ title = {{TensorFlow Model Garden}},
+ howpublished = {\url{https://github.com/tensorflow/models}},
+ year = {2020}
+}
+```
\ No newline at end of file
diff --git a/official/projects/perceiver/configs/encoders.py b/official/projects/perceiver/configs/encoders.py
new file mode 100644
index 00000000000..31df522027f
--- /dev/null
+++ b/official/projects/perceiver/configs/encoders.py
@@ -0,0 +1,47 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Build perceiver sequence encoder."""
+
+from official.projects.perceiver.configs import perceiver as cfg
+from official.projects.perceiver.modeling.layers import encoder
+from official.projects.perceiver.modeling.networks import sequence_encoder
+
+
+def build_encoder(
+ encoder_config: cfg.SequenceEncoderConfig
+) -> sequence_encoder.SequenceEncoder:
+ """Instantiate a perceiver encoder network from SequenceEncoderConfig.
+
+ Args:
+ encoder_config:
+ The sequence encoder config, which provides encoder parameters.
+
+ Returns:
+ An sequence encoder instance.
+ """
+ encoder_ = encoder.Encoder(
+ **encoder_config.encoder.as_dict())
+ return sequence_encoder.SequenceEncoder(
+ encoder=encoder_,
+ d_model=encoder_config.d_model,
+ d_latents=encoder_config.d_latents,
+ z_index_dim=encoder_config.z_index_dim,
+ max_seq_len=encoder_config.max_seq_len,
+ vocab_size=encoder_config.vocab_size,
+ z_pos_enc_init_scale=encoder_config.z_pos_enc_init_scale,
+ embedding_width=encoder_config.embedding_width,
+ embedding_initializer_stddev=encoder_config.embedding_initializer_stddev,
+ input_position_encoding_intializer_stddev=encoder_config
+ .input_position_encoding_intializer_stddev)
diff --git a/official/projects/perceiver/configs/experiments/glue_cola.yaml b/official/projects/perceiver/configs/experiments/glue_cola.yaml
new file mode 100644
index 00000000000..8b86db97811
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_cola.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ metric_type: 'matthews_corrcoef'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 1000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 2670
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 8551 examples, 10 epochs.
+ train_steps: 2670
+ validation_interval: 133
+ # Eval data size = 1043 examples.
+ validation_steps: 33
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'matthews_corrcoef'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/glue_mnli_m.yaml b/official/projects/perceiver/configs/experiments/glue_mnli_m.yaml
new file mode 100644
index 00000000000..fa9fbb6bf0b
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_mnli_m.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 3
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 3000
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 392,702 examples, 10 epochs.
+ train_steps: 122710
+ validation_interval: 6135
+ # Eval data size = 9815 examples.
+ validation_steps: 307
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 122710
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
diff --git a/official/projects/perceiver/configs/experiments/glue_mnli_mm.yaml b/official/projects/perceiver/configs/experiments/glue_mnli_mm.yaml
new file mode 100644
index 00000000000..369e80628ec
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_mnli_mm.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 3
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 3000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 122710
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 392,702 examples, 10 epochs.
+ train_steps: 122710
+ validation_interval: 6135
+ # Eval data size = 9832 examples.
+ validation_steps: 308
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/glue_mrpc.yaml b/official/projects/perceiver/configs/experiments/glue_mrpc.yaml
new file mode 100644
index 00000000000..9eecf9e9a76
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_mrpc.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 1000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 1140
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 3668 examples, 10 epochs.
+ train_steps: 1140
+ validation_interval: 57
+ # Eval data size = 408 examples.
+ validation_steps: 13
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/glue_qnli.yaml b/official/projects/perceiver/configs/experiments/glue_qnli.yaml
new file mode 100644
index 00000000000..99070401254
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_qnli.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 1000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 32730
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 104,743 examples, 10 epochs.
+ train_steps: 32730
+ validation_interval: 1636
+ # Eval data size = 5463 examples.
+ validation_steps: 171
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/glue_qqp.yaml b/official/projects/perceiver/configs/experiments/glue_qqp.yaml
new file mode 100644
index 00000000000..c0f71a1e558
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_qqp.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 3000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 113700
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 363,849 examples, 10 epochs.
+ train_steps: 113700
+ validation_interval: 5685
+ # Eval data size = 40,430 examples.
+ validation_steps: 1264
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/glue_rte.yaml b/official/projects/perceiver/configs/experiments/glue_rte.yaml
new file mode 100644
index 00000000000..25ac49f3830
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_rte.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 1000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 7700
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 2490 examples, 10 epochs.
+ train_steps: 7700
+ validation_interval: 38
+ # Eval data size = 277 examples.
+ validation_steps: 9
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/glue_sst.yaml b/official/projects/perceiver/configs/experiments/glue_sst.yaml
new file mode 100644
index 00000000000..138a075762b
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_sst.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 2
+ metric_type: 'accuracy'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'int'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'int'
+trainer:
+ checkpoint_interval: 1000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 21040
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 67,349 examples, 10 epochs.
+ train_steps: 21040
+ validation_interval: 1052
+ # Eval data size = 872 examples.
+ validation_steps: 28
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'cls_accuracy'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/glue_stsb.yaml b/official/projects/perceiver/configs/experiments/glue_stsb.yaml
new file mode 100644
index 00000000000..160b82bd364
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/glue_stsb.yaml
@@ -0,0 +1,40 @@
+task:
+ hub_module_url: ''
+ model:
+ num_classes: 1
+ metric_type: 'pearson_spearman_corr'
+ train_data:
+ drop_remainder: true
+ global_batch_size: 32
+ input_path: ''
+ is_training: true
+ seq_length: 128
+ label_type: 'float'
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 32
+ input_path: ''
+ is_training: false
+ seq_length: 128
+ label_type: 'float'
+trainer:
+ checkpoint_interval: 1000
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ # 100% of train_steps.
+ decay_steps: 1790
+ end_learning_rate: 0.0
+ initial_learning_rate: 3.0e-05
+ power: 1.0
+ type: polynomial
+ steps_per_loop: 1000
+ summary_interval: 1000
+ # Training data size 5749 examples, 10 epochs.
+ train_steps: 1790
+ validation_interval: 89
+ # Eval data size = 1500 examples.
+ validation_steps: 47
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'pearson_spearman_corr'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/perceiver/configs/experiments/vizier_config_glue.pbtxt b/official/projects/perceiver/configs/experiments/vizier_config_glue.pbtxt
new file mode 100644
index 00000000000..88665776d8d
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/vizier_config_glue.pbtxt
@@ -0,0 +1,27 @@
+# proto-file: learning/vizier/service/vizier.proto
+# proto-message: StudyConfig
+name: "vizier_perceiver"
+description: "Perceiver wordpiece glue finetune"
+parameter_configs {
+ name: "config.task.train_data.global_batch_size"
+ type: DISCRETE
+ feasible_points: 16
+ feasible_points: 32
+ feasible_points: 64
+ external_type: AS_INTEGER
+}
+parameter_configs {
+ name: "config.trainer.optimizer_config.learning_rate.polynomial.initial_learning_rate"
+ type: DISCRETE
+ feasible_points: 1e-5
+ feasible_points: 2e-5
+ feasible_points: 5e-5
+ feasible_points: 1e-4
+ scale_type: UNIT_LINEAR_SCALE
+ external_type: AS_FLOAT
+}
+goal: MAXIMIZE
+max_num_trials: 1000
+write_user: "tensorflow-tpus"
+read_user: "all"
+observation_noise: AUTOMATIC
diff --git a/official/projects/perceiver/configs/experiments/wiki_books_pretrain.yaml b/official/projects/perceiver/configs/experiments/wiki_books_pretrain.yaml
new file mode 100644
index 00000000000..3e7ca698684
--- /dev/null
+++ b/official/projects/perceiver/configs/experiments/wiki_books_pretrain.yaml
@@ -0,0 +1,31 @@
+task:
+ init_checkpoint: ''
+ train_data:
+ drop_remainder: true
+ global_batch_size: 512
+ # Use glob pattern to match all shards except 00141-of-00500-eval which is reserved for eval.
+ input_path: ''
+ is_training: true
+ max_predictions_per_seq: 76
+ seq_length: 512
+ use_next_sentence_label: false
+ use_position_id: false
+ use_v2_feature_names: true
+ validation_data:
+ drop_remainder: false
+ global_batch_size: 512
+ input_path: ''
+ is_training: false
+ max_predictions_per_seq: 76
+ seq_length: 512
+ use_next_sentence_label: false
+ use_position_id: false
+ use_v2_feature_names: true
+trainer:
+ checkpoint_interval: 20000
+ max_to_keep: 5
+ steps_per_loop: 1000
+ summary_interval: 1000
+ train_steps: 500000
+ validation_interval: 1000
+ validation_steps: 64
diff --git a/official/projects/perceiver/configs/perceiver.py b/official/projects/perceiver/configs/perceiver.py
new file mode 100644
index 00000000000..b74373af7f2
--- /dev/null
+++ b/official/projects/perceiver/configs/perceiver.py
@@ -0,0 +1,295 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver configurations."""
+
+import dataclasses
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.modeling.hyperparams import base_config
+from official.nlp.data import pretrain_dataloader
+from official.nlp.data import sentence_prediction_dataloader
+
+
+_SENTENCE_PREDICTION_TRAINER = cfg.TrainerConfig(
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'lamb',
+ 'lamb': {
+ 'weight_decay_rate': 0.01,
+ 'exclude_from_weight_decay': [
+ 'LayerNorm', 'layer_norm', 'bias'
+ ],
+ }
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 3.0e-05,
+ 'end_learning_rate': 0.0,
+ 'decay_steps': 32730,
+ 'power': 1.0,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 200,
+ 'warmup_learning_rate': 0.,
+ }
+ }
+ }))
+
+_MLM_WORDPIECE_TRAINER = cfg.TrainerConfig(
+ train_steps=500_000,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'lamb',
+ 'lamb': {
+ 'weight_decay_rate': 0.01,
+ 'exclude_from_weight_decay': [
+ 'LayerNorm', 'layer_norm', 'bias'
+ ],
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 1.25e-3,
+ 'decay_steps': 500_000,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 1_000,
+ 'warmup_learning_rate': 0.,
+ }
+ }
+ }))
+
+
+@dataclasses.dataclass
+class EncoderConfig(base_config.Config):
+ """The perceiver encoder processor configuration."""
+
+ _attention_heads = 8
+ _per_attention_head_last_dim = 32
+ self_attention_widening_factor: int = 1
+ self_attention_num_heads: int = _attention_heads
+ cross_attention_widening_factor: int = 1
+ cross_attention_num_heads: int = _attention_heads
+ num_self_attends_per_block: int = 26
+ num_blocks: int = 1
+ qk_last_dim: int = _attention_heads * _per_attention_head_last_dim
+ v_last_dim: int = 1280
+ dropout_prob: float = 0.0
+ dropout_attn_prob: float = 0.0
+ att_init_scale: float = 1.0
+ dense_init_scale: float = 1.0
+ norm_epsilon: float = 1e-5
+
+
+@dataclasses.dataclass
+class DecoderConfig(base_config.Config):
+ """The perceiver decoder configuration."""
+ num_heads: int = 8
+ _per_attention_head_last_dim = 32
+ output_last_dim: int = 768
+ qk_last_dim: int = num_heads * _per_attention_head_last_dim
+ v_last_dim: int = 768
+ use_query_residual: bool = False
+
+
+@dataclasses.dataclass
+class PositionalDecoder(base_config.Config):
+ d_model: int = 768
+ decoder: DecoderConfig = dataclasses.field(default_factory=DecoderConfig)
+ position_encoding_intializer_stddev: float = 0.02
+ output_index_dim: int = 512
+ d_latents: int = 1280
+ z_index_dim: int = 256
+
+
+@dataclasses.dataclass
+class ClassificationDecoderConfig(PositionalDecoder):
+ output_index_dim: int = 1
+
+
+@dataclasses.dataclass
+class MaskedLMDecoderConfig(PositionalDecoder):
+ output_index_dim: int = 512
+
+
+@dataclasses.dataclass
+class SequenceEncoderConfig(base_config.Config):
+ """The perceiver sequence encoder configuration."""
+ d_model: int = 768
+ d_latents: int = 1280
+ z_index_dim: int = 256
+ max_seq_len: int = 512
+ vocab_size: int = 30_522
+ embedding_width: int = 768
+ embedding_initializer_stddev: float = 0.02
+ input_position_encoding_intializer_stddev: float = 0.02
+ z_pos_enc_init_scale: float = 0.02
+
+ encoder: EncoderConfig = dataclasses.field(default_factory=EncoderConfig)
+
+
+@dataclasses.dataclass
+class PretrainerConfig(base_config.Config):
+ """The pretrainer configuration."""
+ encoder: SequenceEncoderConfig = dataclasses.field(
+ default_factory=SequenceEncoderConfig
+ )
+ decoder: MaskedLMDecoderConfig = dataclasses.field(
+ default_factory=MaskedLMDecoderConfig
+ )
+
+ mlm_activation: str = 'gelu'
+ mlm_initializer_range: float = 0.02
+
+
+@dataclasses.dataclass
+class ClassificationConfig(base_config.Config):
+ """The classification configuration."""
+ num_classes: int = 0
+ use_encoder_pooler: bool = False
+ encoder: SequenceEncoderConfig = dataclasses.field(
+ default_factory=SequenceEncoderConfig
+ )
+ decoder: ClassificationDecoderConfig = dataclasses.field(
+ default_factory=ClassificationDecoderConfig
+ )
+
+
+@dataclasses.dataclass
+class SentencePredictionConfig(cfg.TaskConfig):
+ """The sentence prediction task config."""
+
+ model: ClassificationConfig = dataclasses.field(
+ default_factory=ClassificationConfig
+ )
+
+ hub_module_url: str = ''
+ init_checkpoint: str = ''
+ init_cls_pooler: bool = False
+
+ metric_type: str = 'accuracy'
+
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
+
+
+@dataclasses.dataclass
+class PretrainConfig(cfg.TaskConfig):
+ """The word piece pretrain task config."""
+
+ model: PretrainerConfig = dataclasses.field(default_factory=PretrainerConfig)
+ init_checkpoint: str = ''
+
+ scale_loss: bool = False
+
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
+
+
+@exp_factory.register_config_factory('perceiver/word_piece_sentence_prediction')
+def perceiver_word_piece_sentence_prediction() -> cfg.ExperimentConfig:
+ """Config for perceiver sentence prediction.
+
+ Returns:
+ cfg.ExperimentConfig
+ References:
+ Perceiver IO (https://arxiv.org/abs/2107.14795).
+ """
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(enable_xla=True),
+ task=SentencePredictionConfig(
+ train_data=sentence_prediction_dataloader
+ .SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader
+ .SentencePredictionDataConfig()),
+ trainer=_SENTENCE_PREDICTION_TRAINER,
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+ return config
+
+
+@exp_factory.register_config_factory(
+ 'perceiver/word_piece_raw_sentence_prediction'
+)
+def perceiver_word_piece_raw_sentence_prediction() -> cfg.ExperimentConfig:
+ """Config for perceiver sentence prediction.
+
+ Returns:
+ cfg.ExperimentConfig
+ References:
+ Perceiver IO (https://arxiv.org/abs/2107.14795).
+ """
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(enable_xla=True),
+ task=SentencePredictionConfig(
+ train_data=sentence_prediction_dataloader.SentencePredictionTextDataConfig(),
+ validation_data=sentence_prediction_dataloader.SentencePredictionTextDataConfig(),
+ ),
+ trainer=_SENTENCE_PREDICTION_TRAINER,
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ],
+ )
+ return config
+
+
+@exp_factory.register_config_factory('perceiver/wordpiece_pretrain')
+def perceiver_wordpiece_pretrain() -> cfg.ExperimentConfig:
+ """Config for perceiver wordpiece pretrain.
+
+ Returns:
+ cfg.ExperimentConfig
+ References:
+ Perceiver IO (https://arxiv.org/abs/2107.14795).
+ Bert pretraining data
+ (https://github.com/google-research/bert/blob/master/tokenization.py#L168)
+ """
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(enable_xla=True),
+ task=PretrainConfig(
+ train_data=pretrain_dataloader.BertPretrainDataConfig(
+ global_batch_size=512,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True),
+ validation_data=pretrain_dataloader.BertPretrainDataConfig(
+ global_batch_size=512,
+ is_training=False,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True)),
+ trainer=_MLM_WORDPIECE_TRAINER,
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ])
+ return config
diff --git a/official/projects/perceiver/configs/perceiver_test.py b/official/projects/perceiver/configs/perceiver_test.py
new file mode 100644
index 00000000000..69ad175ad16
--- /dev/null
+++ b/official/projects/perceiver/configs/perceiver_test.py
@@ -0,0 +1,77 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for official.nlp.tasks.masked_lm."""
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.data import pretrain_dataloader
+from official.nlp.data import sentence_prediction_dataloader
+from official.projects.perceiver.configs import perceiver
+
+
+class PerceiverWordPiecePretrainConfigTest(tf.test.TestCase):
+
+ def test_word_piece_pretrain_config(self):
+ config = perceiver.PretrainConfig(
+ train_data=pretrain_dataloader.BertPretrainDataConfig(
+ global_batch_size=512,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True),
+ validation_data=pretrain_dataloader.BertPretrainDataConfig(
+ global_batch_size=512,
+ is_training=False,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True))
+ self.assertIsNotNone(config)
+ self.assertIsNotNone(config.model)
+ self.assertFalse(config.scale_loss)
+
+
+class PerceiverWordPieceSentencePredictionConfigTest(tf.test.TestCase):
+
+ def test_word_piece_fine_tune_config(self):
+ config = perceiver.SentencePredictionConfig(
+ train_data=sentence_prediction_dataloader
+ .SentencePredictionDataConfig(),
+ validation_data=sentence_prediction_dataloader
+ .SentencePredictionDataConfig())
+ self.assertIsNotNone(config)
+ self.assertIsNotNone(config.model)
+ self.assertFalse(config.init_cls_pooler)
+
+ def test_perceiver_sentence_prediction_returns_valid_learning_rate(self):
+ experiment_cfg = perceiver.perceiver_word_piece_sentence_prediction()
+ self.assertIsNotNone(experiment_cfg.trainer.optimizer_config.learning_rate)
+
+
+class PerceiverWordPieceRawSentencePredictionConfigTest(tf.test.TestCase):
+
+ def test_word_piece_raw_sentence_fine_tune_config(self):
+ config = perceiver.SentencePredictionConfig(
+ train_data=sentence_prediction_dataloader
+ .SentencePredictionTextDataConfig(),
+ validation_data=sentence_prediction_dataloader
+ .SentencePredictionTextDataConfig())
+ self.assertIsNotNone(config)
+ self.assertIsNotNone(config.model)
+ self.assertFalse(config.init_cls_pooler)
+
+ def test_perceiver_raw_sentence_prediction_returns_valid_learning_rate(self):
+ experiment_cfg = perceiver.perceiver_word_piece_raw_sentence_prediction()
+ self.assertIsNotNone(experiment_cfg.trainer.optimizer_config.learning_rate)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/perceiver/modeling/layers/decoder.py b/official/projects/perceiver/modeling/layers/decoder.py
new file mode 100644
index 00000000000..209956cb3b9
--- /dev/null
+++ b/official/projects/perceiver/modeling/layers/decoder.py
@@ -0,0 +1,147 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver basic decoder."""
+
+import collections
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+from official.projects.perceiver.modeling.layers import utils
+
+
+class Decoder(tf_keras.layers.Layer):
+ """Perceiver Decoder layer.
+
+ Uses cross attention decoder layer.
+ This layer implements a Perceiver Decoder from
+ "Perceiver: General Perception with Iterative Attention".
+ (https://arxiv.org/abs/2103.03206)
+
+ References:
+ [Attention Is All You Need](https://arxiv.org/abs/1706.03762)
+ [Perceiver: General Perception with Iterative
+ Attention](https://arxiv.org/abs/2103.03206)
+ (https://github.com/deepmind/deepmind-research/blob/master/perceiver/perceiver.py)
+ (https://github.com/tensorflow/models/blob/871c4e0a393ef4385534bee55354a5df8aa1ccf4/official/nlp/modeling/layers/transformer_encoder_block.py)
+ """
+
+ def __init__(self,
+ output_last_dim,
+ qk_last_dim=None,
+ v_last_dim=None,
+ use_query_residual=False,
+ output_w_init=None,
+ num_heads=1,
+ name="decoder",
+ **kwargs):
+ """Init.
+
+ Args:
+ output_last_dim:
+ Last dim size for output.
+ qk_last_dim:
+ When set, determines the last dimension of the attention score output.
+ Check `qk_last_dim` doc in `utils.build_cross_attention_block_args`.
+ v_last_dim:
+ When set, determines the value's last dimension in the multi-head
+ attention.
+ Check `v_last_dim` doc in `utils._build_transformer_encoder_block_args`.
+ use_query_residual:
+ Toggle to execute residual connection after attention.
+ output_w_init:
+ Ouptut layer kernel initializer.
+ num_heads:
+ Number of attention heads for the `TransformerEncoderBlock`.
+ name:
+ Sets the `tf_keras.layers.Layer` name.
+ **kwargs:
+ Any keyword arguments to pass through to `tf_keras.layers.Layer`.
+ """
+ super().__init__(name=name, **kwargs)
+
+ self._output_last_dim = output_last_dim
+ self._output_w_init = output_w_init
+ self._use_query_residual = use_query_residual
+ self._qk_last_dim = qk_last_dim
+ self._v_last_dim = v_last_dim
+ self._final_project = False # Make variable if needed
+ self._num_heads = num_heads
+
+ # Omitted `concat_preprocessed_input` for MLM use-case.
+
+ def build(self, input_shape):
+ """Build layers using `input_shape`.
+
+ Args:
+ input_shape:
+ Input shape(s) of the layer call.
+ """
+ decoder_query_shape = input_shape[0]
+ z_shape = input_shape[1]
+ self._decoding_cross_attn = layers.TransformerEncoderBlock(
+ **utils.build_cross_attention_block_args(
+ (decoder_query_shape, z_shape),
+ widening_factor=1,
+ dropout_prob=0.0,
+ num_heads=self._num_heads,
+ shape_for_attn="kv",
+ qk_last_dim=self._qk_last_dim,
+ v_last_dim=self._v_last_dim,
+ use_query_residual=self._use_query_residual))
+
+ def call(self, inputs, training=None, query_mask=None):
+ """Return decoded output of latent vector via the query.
+
+ Args:
+ inputs:
+ Expect inputs to be a tuple of perceiver's decoder query tensor and
+ latent tensor (z). For the cross attention block, `z` is the key-value
+ tensor and decoder query is the query tensor.
+ Latent tensor comes from the self-attention processing blocks and
+ decoder query comes from users to query for the desired output.
+ training:
+ Flag to indicate training status.
+ query_mask:
+ mask used to create the attention mask for the query tensor in the
+ cross attention block.
+
+ Returns:
+ `tf.Tensor` decoded output of latent vector via the query.
+ """
+ if not isinstance(inputs, collections.abc.Sequence):
+ raise ValueError("`inputs` must be a sequence.")
+ if len(inputs) != 2:
+ raise ValueError("`inputs` must have two elements.")
+
+ query, z = inputs
+ # Cross-attention decoding.
+ # key, value: B x N x K; query: B x M x K
+ # Attention maps -> B x N x M
+ # Output -> B x M x K
+ # Construct cross attention and linear layer lazily, in case we don't need
+ # them.
+ if query_mask is None:
+ attention_mask = None
+ else:
+ attention_mask = utils.make_cross_attention_mask(
+ query_mask=query_mask,
+ kv_mask=tf.ones(tf.shape(z)[:2], dtype=tf.int32))
+
+ output = self._decoding_cross_attn(
+ (query, z, attention_mask),
+ training=training)
+
+ return output
diff --git a/official/projects/perceiver/modeling/layers/decoder_test.py b/official/projects/perceiver/modeling/layers/decoder_test.py
new file mode 100644
index 00000000000..f90326ff870
--- /dev/null
+++ b/official/projects/perceiver/modeling/layers/decoder_test.py
@@ -0,0 +1,103 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for decoder."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.perceiver.modeling.layers import decoder
+
+
+class PerceiverBasicDecoderTest(tf.test.TestCase):
+
+ def test_layer_creation(self):
+ sequence_length = 80
+ embedding_width = 800
+ test_layer = decoder.Decoder(
+ output_last_dim=embedding_width,
+ num_heads=8)
+ lantent_length = 8
+ latent_width = 80
+ query_input = tf_keras.Input(
+ shape=(sequence_length, embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(lantent_length, latent_width))
+
+ output_tensor = test_layer((query_input, latent_input))
+ self.assertEqual(
+ query_input.shape.as_list(),
+ output_tensor.shape.as_list())
+
+ def test_layer_creation_with_mask(self):
+ embedding_width = 800
+ sequence_length = 80
+ test_layer = decoder.Decoder(
+ output_last_dim=embedding_width,
+ num_heads=8)
+ lantent_length = 8
+ latent_width = 80
+ query_input = tf_keras.Input(
+ shape=(sequence_length, embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(lantent_length, latent_width))
+ mask_tensor = tf_keras.Input(
+ shape=(sequence_length),
+ dtype=tf.int32)
+ output_tensor = test_layer(
+ (query_input, latent_input),
+ query_mask=mask_tensor)
+ self.assertEqual(
+ query_input.shape.as_list(),
+ output_tensor.shape.as_list())
+
+ def test_layer_invocation(self):
+ embedding_width = 800
+ sequence_length = 80
+ test_layer = decoder.Decoder(
+ output_last_dim=embedding_width,
+ num_heads=8)
+ lantent_length = 8
+ latent_width = 80
+ query_input = tf_keras.Input(
+ shape=(sequence_length, embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(lantent_length, latent_width))
+ mask_tensor = tf_keras.Input(
+ shape=(sequence_length),
+ dtype=tf.int32)
+ output_tensor = test_layer(
+ (query_input, latent_input),
+ query_mask=mask_tensor)
+
+ # Create a model from the test layer.
+ model = tf_keras.Model(
+ ((query_input, latent_input), mask_tensor),
+ output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 6
+ latent_data = 10 * np.random.random_sample(
+ (batch_size, lantent_length, latent_width))
+ mask_data = tf.ones((batch_size, sequence_length), dtype=tf.int32)
+ query_data = tf.ones(
+ (batch_size, sequence_length, embedding_width),
+ dtype=tf.float32)
+ _ = model.predict(((query_data, latent_data), mask_data))
+
+# TODO(b/222634115) Add tests to validate logic and dims.
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/perceiver/modeling/layers/encoder.py b/official/projects/perceiver/modeling/layers/encoder.py
new file mode 100644
index 00000000000..acc9141d41f
--- /dev/null
+++ b/official/projects/perceiver/modeling/layers/encoder.py
@@ -0,0 +1,170 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver encode processor."""
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+from official.projects.perceiver.modeling.layers import utils
+
+
+class Encoder(tf_keras.layers.Layer):
+ """Perceiver Encoder and Processor(s) layer.
+
+ This layer implements the Perceiver Encoder and Processor stack from
+ "Perceiver: General Perception with Iterative Attention".
+ (https://arxiv.org/abs/2103.03206)
+ It uses SelfAttention and CrossAttention modules.
+ It allows the user to choose the initial latent positional encodings.
+
+ References:
+ [Perceiver: General Perception with Iterative
+ Attention](https://arxiv.org/abs/2103.03206)
+ (https://github.com/deepmind/deepmind-research/blob/master/perceiver/perceiver.py)
+ (https://github.com/tensorflow/models/blob/871c4e0a393ef4385534bee55354a5df8aa1ccf4/official/nlp/modeling/layers/transformer_encoder_block.py)
+ """
+
+ def __init__(self,
+ self_attention_num_heads=8,
+ self_attention_widening_factor=1,
+ cross_attention_num_heads=8,
+ cross_attention_widening_factor=1,
+ num_self_attends_per_block=6,
+ num_blocks=8,
+ qk_last_dim=None,
+ v_last_dim=None,
+ dropout_prob=0.0,
+ dropout_attn_prob=0.0,
+ att_init_scale=1.0,
+ dense_init_scale=1.0,
+ norm_epsilon=1e-5,
+ name="encode_processor",
+ **kwargs):
+ """Init.
+
+ Args:
+ self_attention_num_heads:
+ Number of attention heads in the self-attention transformer block.
+ self_attention_widening_factor:
+ Multiplier used to widen on the inner layer of the MLP step within the
+ self-attention transformer block.
+ cross_attention_num_heads:
+ Number of attention heads in the cross-attention transformer block.
+ cross_attention_widening_factor:
+ Multiplier used to widen on the inner layer of the MLP step within the
+ cross-attention transformer block.
+ num_self_attends_per_block:
+ Number of different self-attention encoders initialized per latent
+ perceiver block.
+ num_blocks:
+ Number of latent perceiver blocks.
+ qk_last_dim:
+ When set, determines the last dimension of the attention score output.
+ Check `qk_last_dim` doc in `utils.build_cross_attention_block_args` for
+ more details.
+ v_last_dim:
+ It can impact the last dimension size of value projection in mult-head
+ attention output and `TransformerEncoderBlock`'s output.
+ For more details, check `v_last_dim` doc in
+ `utils._build_transformer_encoder_block_args`.
+ dropout_prob:
+ Dropout probability for the post-attention and output dropout.
+ dropout_attn_prob:
+ Dropout probability for within the attention layer.
+ att_init_scale:
+ Scale for the `tf_keras.initializers.VarianceScaling` used in attention
+ kernel.
+ dense_init_scale:
+ Scale for the `tf_keras.initializers.VarianceScaling` used in MLP
+ kernel.
+ norm_epsilon:
+ Epsilon value to initialize normalization layers.
+ name:
+ Sets the `tf_keras.layers.Layer` name.
+ **kwargs:
+ Any keyword arguments to pass through to `tf_keras.layers.Layer`.
+ """
+ super().__init__(name=name, **kwargs)
+
+ self._input_is_1d = True
+
+ self._num_self_attends_per_block = num_self_attends_per_block
+ self._dropout_prob = dropout_prob
+ self._qk_last_dim = qk_last_dim
+ self._v_last_dim = v_last_dim
+ self._norm_epsilon = norm_epsilon
+ self._dropout_attn_prob = dropout_attn_prob
+ self._att_init_scale = att_init_scale
+ self._dense_init_scale = dense_init_scale
+ self._num_blocks = num_blocks
+
+ self._self_attention_widening_factor = self_attention_widening_factor
+ self._self_attention_num_heads = self_attention_num_heads
+
+ self._cross_attention_widening_factor = cross_attention_widening_factor
+ self._cross_attention_num_heads = cross_attention_num_heads
+ self._cross_attention_shape_for_attn = "kv"
+ self._cross_attention_use_query_residual = True
+
+ def build(self, input_shape):
+ embeddings_shape = input_shape[0]
+ z_shape = input_shape[1]
+ self._self_attention_encoder_blocks = []
+ for i in range(self._num_self_attends_per_block):
+ self._self_attention_encoder_blocks.append(layers.TransformerEncoderBlock(
+ name=f"self_attention_encoder_{i}",
+ **utils.build_self_attention_block_args(
+ (z_shape,),
+ widening_factor=self._self_attention_widening_factor,
+ dropout_prob=self._dropout_prob,
+ dropout_attn_prob=self._dropout_attn_prob,
+ num_heads=self._self_attention_num_heads,
+ att_init_scale=self._att_init_scale,
+ dense_init_scale=self._dense_init_scale,
+ qk_last_dim=self._qk_last_dim,
+ v_last_dim=self._v_last_dim,
+ norm_epsilon=self._norm_epsilon)))
+
+ self._cross_attention_encoder_block = layers.TransformerEncoderBlock(
+ name="cross_attention_encoder",
+ **utils.build_cross_attention_block_args(
+ (z_shape, embeddings_shape),
+ widening_factor=self._cross_attention_widening_factor,
+ dropout_prob=self._dropout_prob,
+ dropout_attn_prob=self._dropout_attn_prob,
+ num_heads=self._cross_attention_num_heads,
+ att_init_scale=self._att_init_scale,
+ dense_init_scale=self._dense_init_scale,
+ shape_for_attn=self._cross_attention_shape_for_attn,
+ use_query_residual=self._cross_attention_use_query_residual,
+ norm_epsilon=self._norm_epsilon,
+ qk_last_dim=self._qk_last_dim,
+ v_last_dim=self._v_last_dim))
+
+ def call(self, inputs, input_mask=None, training=None):
+ embeddings = inputs[0]
+ z = inputs[1]
+ if input_mask is None:
+ input_mask = tf.ones(tf.shape(embeddings)[:2], dtype=tf.int32)
+ attention_mask = utils.make_cross_attention_mask(
+ query_mask=tf.ones(tf.shape(z)[:2], dtype=tf.int32),
+ kv_mask=input_mask)
+ z = self._cross_attention_encoder_block(
+ (z, embeddings, attention_mask),
+ training=training)
+ for _ in range(self._num_blocks):
+ for self_attention_block in self._self_attention_encoder_blocks:
+ z = self_attention_block(z, training=training)
+ return z
diff --git a/official/projects/perceiver/modeling/layers/encoder_test.py b/official/projects/perceiver/modeling/layers/encoder_test.py
new file mode 100644
index 00000000000..85dc62303c2
--- /dev/null
+++ b/official/projects/perceiver/modeling/layers/encoder_test.py
@@ -0,0 +1,217 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for encoder."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.perceiver.modeling.layers import encoder
+
+
+class EncoderTest(tf.test.TestCase):
+
+ def test_layer_creation(self):
+ test_layer = encoder.Encoder(
+ self_attention_num_heads=8,
+ cross_attention_num_heads=8)
+ sequence_length = 80
+ embedding_width = 800
+ lantent_length = 8
+ latent_width = 80
+ data_input = tf_keras.Input(
+ shape=(sequence_length, embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(lantent_length, latent_width))
+
+ output_tensor = test_layer((data_input, latent_input))
+ self.assertEqual(
+ latent_input.shape.as_list(),
+ output_tensor.shape.as_list())
+
+ def test_layer_creation_with_mask(self):
+ test_layer = encoder.Encoder(
+ self_attention_num_heads=8,
+ cross_attention_num_heads=8)
+ sequence_length = 80
+ embedding_width = 800
+ lantent_length = 8
+ latent_width = 80
+ data_input = tf_keras.Input(
+ shape=(sequence_length, embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(lantent_length, latent_width))
+ mask_tensor = tf_keras.Input(
+ shape=(sequence_length),
+ dtype=tf.int32)
+ output_tensor = test_layer(
+ (data_input, latent_input),
+ input_mask=mask_tensor)
+ self.assertEqual(
+ latent_input.shape.as_list(),
+ output_tensor.shape.as_list())
+
+ def test_layer_invocation(self):
+ test_layer = encoder.Encoder(
+ self_attention_num_heads=8,
+ cross_attention_num_heads=8)
+ sequence_length = 80
+ embedding_width = 800
+ lantent_length = 8
+ latent_width = 80
+ data_input = tf_keras.Input(
+ shape=(sequence_length, embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(lantent_length, latent_width))
+ mask_tensor = tf_keras.Input(
+ shape=(sequence_length),
+ dtype=tf.int32)
+
+ output_tensor = test_layer(
+ (data_input, latent_input),
+ input_mask=mask_tensor)
+
+ # Create a model from the test layer.
+ model = tf_keras.Model(
+ ((data_input, latent_input), mask_tensor),
+ output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 6
+ input_data = 10 * np.random.random_sample(
+ (batch_size, sequence_length, embedding_width))
+ mask_data = tf.ones((batch_size, sequence_length), dtype=tf.int32)
+ latent_data = tf.ones((batch_size, lantent_length, latent_width),
+ dtype=tf.float32)
+ _ = model.predict(((input_data, latent_data), mask_data))
+
+ def test_self_attention_widening_factor(self):
+ last_dim = 160
+ self_attention_widening_factor = 2
+ test_layer = encoder.Encoder(
+ self_attention_widening_factor=self_attention_widening_factor,
+ v_last_dim=last_dim)
+
+ some_sequence_length = 80
+ some_embedding_width = 800
+ some_lantent_length = 8
+ some_latent_width = last_dim
+ data_input = tf_keras.Input(
+ shape=(some_sequence_length, some_embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(some_lantent_length, some_latent_width))
+ mask_tensor = tf_keras.Input(shape=(some_sequence_length), dtype=tf.int32)
+ test_layer((data_input, latent_input), input_mask=mask_tensor)
+ value = test_layer._self_attention_encoder_blocks[
+ 0]._intermediate_dense.get_config()['output_shape'].pop()
+ self.assertEqual(last_dim * self_attention_widening_factor, value)
+
+ def test_cross_attention_widening_factor(self):
+ last_dim = 160
+ cross_attention_widening_factor = 2
+ test_layer = encoder.Encoder(
+ cross_attention_widening_factor=cross_attention_widening_factor,
+ v_last_dim=last_dim)
+
+ some_sequence_length = 80
+ some_embedding_width = 800
+ some_lantent_length = 8
+ some_latent_width = last_dim
+ data_input = tf_keras.Input(
+ shape=(some_sequence_length, some_embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(some_lantent_length, some_latent_width))
+ mask_tensor = tf_keras.Input(shape=(some_sequence_length), dtype=tf.int32)
+ test_layer((data_input, latent_input), input_mask=mask_tensor)
+ value = test_layer._cross_attention_encoder_block._intermediate_dense.get_config(
+ )['output_shape'].pop()
+ self.assertEqual(last_dim * cross_attention_widening_factor, value)
+
+ def test_self_attention_num_heads(self):
+ # TODO(b/222634115) parameterize test.
+ self_attention_num_heads = 16
+ test_layer = encoder.Encoder(
+ self_attention_num_heads=self_attention_num_heads)
+
+ some_sequence_length = 80
+ some_embedding_width = 800
+ some_lantent_length = 8
+ some_latent_width = 64
+ data_input = tf_keras.Input(
+ shape=(some_sequence_length, some_embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(some_lantent_length, some_latent_width))
+ mask_tensor = tf_keras.Input(shape=(some_sequence_length), dtype=tf.int32)
+ test_layer((data_input, latent_input), input_mask=mask_tensor)
+ value = test_layer._self_attention_encoder_blocks[
+ 0]._attention_layer.get_config()['num_heads']
+ self.assertEqual(self_attention_num_heads, value)
+
+ def test_cross_attention_num_heads(self):
+ # TODO(b/222634115) parameterize test.
+ cross_attention_num_heads = 16
+ test_layer = encoder.Encoder(
+ cross_attention_num_heads=cross_attention_num_heads)
+
+ some_sequence_length = 80
+ some_embedding_width = 800
+ some_lantent_length = 8
+ some_latent_width = 64
+ data_input = tf_keras.Input(
+ shape=(some_sequence_length, some_embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(some_lantent_length, some_latent_width))
+ mask_tensor = tf_keras.Input(shape=(some_sequence_length), dtype=tf.int32)
+ test_layer((data_input, latent_input), input_mask=mask_tensor)
+ value = test_layer._cross_attention_encoder_block._attention_layer.get_config(
+ )['num_heads']
+ self.assertEqual(cross_attention_num_heads, value)
+
+ def test_num_self_attends_per_block(self):
+ # TODO(b/222634115) parameterize test.
+ num_self_attends_per_block = 3
+ test_layer = encoder.Encoder(
+ num_self_attends_per_block=num_self_attends_per_block)
+
+ some_sequence_length = 80
+ some_embedding_width = 800
+ some_lantent_length = 8
+ some_latent_width = 64
+ data_input = tf_keras.Input(
+ shape=(some_sequence_length, some_embedding_width))
+ latent_input = tf_keras.Input(
+ shape=(some_lantent_length, some_latent_width))
+ mask_tensor = tf_keras.Input(shape=(some_sequence_length), dtype=tf.int32)
+ test_layer((data_input, latent_input), input_mask=mask_tensor)
+ self.assertLen(
+ test_layer._self_attention_encoder_blocks,
+ num_self_attends_per_block)
+
+ # TODO(b/222634115) num_blocks
+ # TODO(b/222634115) qk_last_dim validations
+ # TODO(b/222634115) v_last_dim validations
+ # TODO(b/222634115) dropout_prob validation
+ # TODO(b/222634115) dropout_attn_prob validation
+ # TODO(b/222634115) att_init_scale validation
+ # TODO(b/222634115) dense_init_scale validation
+ # TODO(b/222634115) cross_attention_use_query_residual validation
+ # (value passed correctly)
+ # TODO(b/222634115) norm_epsilon
+ # TODO(b/222634115) check latent dims
+ # TODO(b/222634115) make cross att mask validation when input_mask is None
+ # TODO(b/222634115) make cross att mask validation when input_mask is not None
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/perceiver/modeling/layers/utils.py b/official/projects/perceiver/modeling/layers/utils.py
new file mode 100644
index 00000000000..019eb0a8f9d
--- /dev/null
+++ b/official/projects/perceiver/modeling/layers/utils.py
@@ -0,0 +1,351 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver modeling utils."""
+
+import functools
+import tensorflow as tf, tf_keras
+
+
+def make_cross_attention_mask(query_mask, kv_mask):
+ """Compute the outer product between `query_mask` and `kv_mask`."""
+ # Porting `mask = jax.vmap(jnp.outer)(query_mask, kv_mask)`
+ return tf.einsum("ab,ac->abc", query_mask, kv_mask)
+
+
+def build_cross_attention_block_args(
+ input_shape,
+ widening_factor=1,
+ dropout_prob=0.0,
+ dropout_attn_prob=0.0,
+ num_heads=8,
+ att_init_scale=1.0,
+ dense_init_scale=1.0,
+ shape_for_attn="kv",
+ use_query_residual=True,
+ norm_epsilon=1e-5,
+ qk_last_dim=None,
+ v_last_dim=None):
+ """Builds cross attention block arguments for `TransformerEncoderBlock`.
+
+ Build cross attention block arguments for `TransformerEncoderBlock` used
+ in Perceiver.
+
+ The last dimension of the output of the attention block or `output_last_dim`
+ of `TransformerEncoderBlocks` is set to the first `input_shape`'s last
+ dimension.
+
+ `diff_q_kv_att_layer_norm` is set to `True`.
+
+ `inner_dropout` is set to 0.
+
+ `norm_first` is set to `True`.
+
+ `inner_activation` is set to gelu.
+
+ `kernel_initializer` and `attention_initializer` are both
+ `tf_keras.initializers.VarianceScaling`.
+
+ Args:
+ input_shape:
+ Check `input_shape` doc in `_build_transformer_encoder_block_args`.
+ widening_factor:
+ Check `widening_factor` doc in `_build_transformer_encoder_block_args`.
+ dropout_prob:
+ Check `dropout_prob` doc in `_build_transformer_encoder_block_args`.
+ dropout_attn_prob:
+ Check `dropout_attn_prob` doc in `_build_transformer_encoder_block_args`.
+ num_heads:
+ Check `num_heads` doc in `_build_transformer_encoder_block_args`.
+ att_init_scale:
+ Check `att_init_scale` doc in `_build_transformer_encoder_block_args`.
+ dense_init_scale:
+ Check `dense_init_scale` doc in `_build_transformer_encoder_block_args`.
+ shape_for_attn:
+ Valid values are `q` or `kv`. This value is used to determine the last
+ dimension of the attention score output attention last dimension.
+ `qk_last_dim` has higher precedence over `shape_for_attn`.
+ use_query_residual:
+ Toggle to execute residual connection after attention.
+ norm_epsilon:
+ Check `norm_epsilon` doc in `_build_transformer_encoder_block_args`.
+ qk_last_dim:
+ When set, determines the last dimension of the attention score output.
+ When it's `None`, it uses the first `input_shape`'s last dimension as the
+ last dimension of the attention score output. `qk_last_dim` has higher
+ precedence over `shape_for_attn`.
+ v_last_dim:
+ Check `v_last_dim` doc in `_build_transformer_encoder_block_args`.
+
+ Returns:
+ A `dict` mapping `TransformerEncoderBlock` arguments.
+
+ References:
+ [Perceiver: General Perception with Iterative
+ Attention](https://arxiv.org/abs/2103.03206)
+ (https://github.com/deepmind/deepmind-research/blob/master/perceiver/perceiver.py)
+ (https://github.com/tensorflow/models/blob/871c4e0a393ef4385534bee55354a5df8aa1ccf4/official/nlp/modeling/layers/transformer_encoder_block.py)
+ """
+ inputs_q_shape = input_shape[0]
+ inputs_kv_shape = input_shape[1]
+
+ output_last_dim = inputs_q_shape[-1]
+
+ if shape_for_attn == "q":
+ f_qk_last_dim = inputs_q_shape[-1]
+ elif shape_for_attn == "kv":
+ f_qk_last_dim = inputs_kv_shape[-1]
+ else:
+ raise ValueError(f"Unknown value {shape_for_attn} for "
+ "shape_for_attention.")
+
+ f_v_last_dim = None
+ if qk_last_dim is not None:
+ f_qk_last_dim = qk_last_dim
+ if v_last_dim is not None:
+ f_v_last_dim = v_last_dim
+
+ return _build_transformer_encoder_block_args(
+ input_shape=input_shape,
+ widening_factor=widening_factor,
+ dropout_prob=dropout_prob,
+ dropout_attn_prob=dropout_attn_prob,
+ num_heads=num_heads,
+ att_init_scale=att_init_scale,
+ dense_init_scale=dense_init_scale,
+ use_query_residual=use_query_residual,
+ norm_epsilon=norm_epsilon,
+ qk_last_dim=f_qk_last_dim,
+ v_last_dim=f_v_last_dim,
+ diff_q_kv_att_layer_norm=True,
+ output_last_dim=output_last_dim)
+
+
+def build_self_attention_block_args(
+ input_shape,
+ widening_factor=4,
+ dropout_prob=0.0,
+ dropout_attn_prob=0.0,
+ num_heads=8,
+ att_init_scale=1.0,
+ dense_init_scale=1.0,
+ norm_epsilon=1e-5,
+ qk_last_dim=None,
+ v_last_dim=None):
+ """Builds self attention block arguments for `TransformerEncoderBlock`.
+
+ Light wrapper around `_build_transformer_encoder_block_args` with some
+ assumptions around self attention block. Builds the arguments for
+ `TransformerEncoderBlock` used in Perceiver.
+
+ The last dimension of the output of the attention block or `output_last_dim`
+ of `TransformerEncoderBlocks` is set using the logic described in the
+ doc associated with `output_last_dim` in
+ `_build_transformer_encoder_block_args`.
+
+ `diff_q_kv_att_layer_norm` is set to `False`.
+
+ `use_query_residual` is set to `True`.
+
+ `inner_dropout` is set to 0.
+
+ `norm_first` is set to `True`.
+
+ `inner_activation` is set to gelu.
+
+ `kernel_initializer` and `attention_initializer` are both
+ `tf_keras.initializers.VarianceScaling`.
+
+ Args:
+ input_shape:
+ Check `input_shape` doc in `_build_transformer_encoder_block_args`.
+ widening_factor:
+ Check `widening_factor` doc in `_build_transformer_encoder_block_args`.
+ dropout_prob:
+ Check `dropout_prob` doc in `_build_transformer_encoder_block_args`.
+ dropout_attn_prob:
+ Check `dropout_attn_prob` doc in `_build_transformer_encoder_block_args`.
+ num_heads:
+ Check `num_heads` doc in `_build_transformer_encoder_block_args`.
+ att_init_scale:
+ Check `att_init_scale` doc in `_build_transformer_encoder_block_args`.
+ dense_init_scale:
+ Check `dense_init_scale` doc in `_build_transformer_encoder_block_args`.
+ norm_epsilon:
+ Check `norm_epsilon` doc in `_build_transformer_encoder_block_args`.
+ qk_last_dim:
+ Check `qk_last_dim` doc in `_build_transformer_encoder_block_args`.
+ v_last_dim:
+ Check `v_last_dim` doc in `_build_transformer_encoder_block_args`.
+
+ Returns:
+ A `dict` mapping `TransformerEncoderBlock` arguments.
+
+ References:
+ [Perceiver: General Perception with Iterative
+ Attention](https://arxiv.org/abs/2103.03206)
+ (https://github.com/deepmind/deepmind-research/blob/master/perceiver/perceiver.py)
+ (https://github.com/tensorflow/models/blob/871c4e0a393ef4385534bee55354a5df8aa1ccf4/official/nlp/modeling/layers/transformer_encoder_block.py)
+ """
+
+ return _build_transformer_encoder_block_args(
+ input_shape=input_shape,
+ widening_factor=widening_factor,
+ dropout_prob=dropout_prob,
+ dropout_attn_prob=dropout_attn_prob,
+ num_heads=num_heads,
+ att_init_scale=att_init_scale,
+ dense_init_scale=dense_init_scale,
+ use_query_residual=True,
+ norm_epsilon=norm_epsilon,
+ qk_last_dim=qk_last_dim,
+ v_last_dim=v_last_dim,
+ diff_q_kv_att_layer_norm=False,
+ output_last_dim=None)
+
+
+def _build_transformer_encoder_block_args(
+ input_shape,
+ widening_factor,
+ dropout_prob,
+ dropout_attn_prob,
+ num_heads,
+ att_init_scale,
+ dense_init_scale,
+ use_query_residual,
+ norm_epsilon,
+ qk_last_dim,
+ v_last_dim,
+ diff_q_kv_att_layer_norm,
+ output_last_dim):
+ """Build arguments for `TransformerEncoderBlock`.
+
+ `inner_dropout` is set to 0.
+
+ `norm_first` is set to `True`.
+
+ `inner_activation` is set to gelu.
+
+ `kernel_initializer` and `attention_initializer` are both
+ `tf_keras.initializers.VarianceScaling`.
+
+ Args:
+ input_shape:
+ input shape(s). Usually passed through `build` method in
+ `tf_keras.layers.Layer`.
+ widening_factor:
+ Multiplier used to widen on the inner layer of the MLP step within a
+ transformer attention block.
+ dropout_prob:
+ Dropout probability for the post-attention and output dropout.
+ dropout_attn_prob:
+ Dropout probability for within the attention layer.
+ num_heads:
+ Number of attention heads.
+ att_init_scale:
+ Scale for the `tf_keras.initializers.VarianceScaling` used in attention
+ kernel.
+ dense_init_scale:
+ Scale for the `tf_keras.initializers.VarianceScaling` used in MLP kernel.
+ use_query_residual:
+ Toggle to execute residual connection after attention.
+ norm_epsilon:
+ Epsilon value to initialize normalization layers.
+ qk_last_dim:
+ When set, determines the last dimension of the attention score output.
+ When it's `None`, it uses the first `input_shape`'s last dimension as the
+ last dimension of the attention score output.
+ v_last_dim:
+ When set, determines the value's last dimension in the multi-head
+ attention.
+ When it's `None`, it uses the `qk_last_dim` for `inner_dim` and
+ `value_dim`.
+ If `qk_last_dim` is `None`, the first input_shape's last dimension is used
+ as the last dimension of the attention score output.
+ If `output_last_dim` is `None`, `v_last_dim` is used to set the
+ `TransformerEncoderBlock`'s output's last dimension.
+ diff_q_kv_att_layer_norm:
+ If `True`, create a separate attention layer norm layer for query and
+ key-value if `norm_first` is `True`. Invalid to set to `True` if
+ `norm_first` is `False`.
+ output_last_dim:
+ When set, the value determines the last dimension of the output of the
+ attention block or `output_last_dim`.
+ When it's `None`, it uses, in order of decreasing precedence,
+ `v_last_dim`, `qk_last_dim`, and finally first `input_shape`'s last
+ dimension. To clarify, if `v_last_dim` or `qk_last_dim` is `None`, the
+ next order of precedence is used. The value is used to determine the last
+ dimension of the output of the attention block or `output_last_dim`.
+
+ Returns:
+ A `dict` mapping `TransformerEncoderBlock` arguments.
+
+ References:
+ [Perceiver: General Perception with Iterative
+ Attention](https://arxiv.org/abs/2103.03206)
+ (https://github.com/deepmind/deepmind-research/blob/master/perceiver/perceiver.py)
+ (https://github.com/tensorflow/models/blob/871c4e0a393ef4385534bee55354a5df8aa1ccf4/official/nlp/modeling/layers/transformer_encoder_block.py)
+ """
+
+ inputs_q_shape = input_shape[0]
+ # Q and K must have the same number of last dim.
+ # Default to preserving Q's input's shape.
+ if qk_last_dim is None:
+ qk_last_dim = inputs_q_shape[-1]
+ # V's number of last dim determines the shape of the output of QKV-attention.
+ # Default to the same number of last dim used in the key-query operation.
+ if v_last_dim is None:
+ v_last_dim = qk_last_dim
+ # Project the output of QKV attention to a desired number of last dim.
+ # Default to the same number as the output of the QKV attention operation.
+ if output_last_dim is None:
+ output_last_dim = v_last_dim
+
+ assert qk_last_dim % num_heads == 0
+ assert v_last_dim % num_heads == 0
+ qk_last_dim_per_head = qk_last_dim // num_heads
+ v_last_dim_per_head = v_last_dim // num_heads
+
+ return {
+ "num_attention_heads":
+ num_heads,
+ "inner_dim":
+ output_last_dim * widening_factor,
+ "inner_activation":
+ functools.partial(tf_keras.activations.gelu, approximate=True),
+ "kernel_initializer":
+ tf_keras.initializers.VarianceScaling(scale=dense_init_scale),
+ "attention_initializer":
+ tf_keras.initializers.VarianceScaling(scale=att_init_scale),
+ "norm_first":
+ True,
+ "norm_epsilon":
+ norm_epsilon,
+ "output_dropout":
+ dropout_prob,
+ "attention_dropout":
+ dropout_attn_prob,
+ "inner_dropout":
+ 0.0,
+ "use_query_residual":
+ use_query_residual,
+ "value_dim":
+ v_last_dim_per_head,
+ "key_dim":
+ qk_last_dim_per_head,
+ "output_last_dim":
+ output_last_dim,
+ "diff_q_kv_att_layer_norm":
+ diff_q_kv_att_layer_norm,
+ }
diff --git a/official/projects/perceiver/modeling/layers/utils_test.py b/official/projects/perceiver/modeling/layers/utils_test.py
new file mode 100644
index 00000000000..9e5d715f1e8
--- /dev/null
+++ b/official/projects/perceiver/modeling/layers/utils_test.py
@@ -0,0 +1,73 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for utils."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.perceiver.modeling.layers import utils
+
+
+class PerceiverUtilsSelfAttentionBlockArgsTest(tf.test.TestCase):
+
+ def test_output_last_dim_is_same_as_input_last_dim(self):
+ q_seq_len = 10
+ input_last_dim = 30
+ some_num_heads = 2
+
+ some_input_shape = ((2, q_seq_len, input_last_dim),)
+ args = utils.build_self_attention_block_args(
+ some_input_shape,
+ num_heads=some_num_heads)
+ self.assertEqual(args['output_last_dim'], input_last_dim)
+
+ def test_value_dim_is_same_as_input_last_dim_div_num_heads(self):
+ q_seq_len = 10
+ input_last_dim = 30
+ some_num_heads = 2
+
+ some_input_shape = ((2, q_seq_len, input_last_dim),)
+ args = utils.build_self_attention_block_args(
+ some_input_shape,
+ num_heads=some_num_heads)
+ self.assertEqual(args['value_dim'], input_last_dim // some_num_heads)
+
+ # TODO(b/222634115) Add tests for `build_self_attention_block_args` for
+ # better coverage
+
+
+class PerceiverUtilsCrossAttentionBlockArgsTest(tf.test.TestCase):
+
+ def test_1(self):
+ some_batch_size = 2
+ q_seq_len = 10
+ q_input_last_dim = 30
+ kv_seq_len = 6
+ kv_input_last_dim = 60
+ some_num_heads = 2
+
+ some_input_shape = (
+ (some_batch_size, q_seq_len, q_input_last_dim),
+ (some_batch_size, kv_seq_len, kv_input_last_dim))
+ args = utils.build_cross_attention_block_args(
+ some_input_shape,
+ num_heads=some_num_heads)
+ self.assertEqual(args['output_last_dim'], q_input_last_dim)
+
+ # TODO(b/222634115) Add tests for `build_cross_attention_block_args` for
+ # better coverage
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/perceiver/modeling/models/classifier.py b/official/projects/perceiver/modeling/models/classifier.py
new file mode 100644
index 00000000000..3774833f38d
--- /dev/null
+++ b/official/projects/perceiver/modeling/models/classifier.py
@@ -0,0 +1,228 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver classifier."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+
+
+class Classifier(tf_keras.Model):
+ """Classifier model based on a shared encoder and optional decoder.
+
+ This is an implementation of the network structure surrounding a transformer
+ encoder as described in "Perceiver IO: A General Architecture for Structured
+ Inputs & Outputs" (https://arxiv.org/abs/2107.14795).
+
+ The Classifier allows a user to pass in an encoder stack and an optional
+ decoder stack (e.g. perceiver decoder), and instantiates a classification
+ network based on the passed `num_classes` argument. If `num_classes` is set
+ to 1, a regression network is instantiated.
+
+ This is forked from
+ (https://github.com/tensorflow/models/blob/master/official/nlp/modeling/models/bert_classifier.py)
+
+ Attributes:
+ network:
+ A perceiver encode and processor transformer network. This network
+ should output a classification output. Furthermore, it should expose its
+ embedding table via a "get_embedding_table" method.
+ num_classes:
+ Number of classes outputted by classification head.
+ inputs:
+ A `Dict[str, tf_keras.Input]` with `input_word_ids`, `input_mask`, and
+ `input_type_ids`. The shapes are all `(None)` with dtype `tf.int32`.
+ head_name:
+ Name of the classification head.
+ classifier:
+ Classification head layer.
+ initializer:
+ `tf_keras.initializers.Initializer` used for classification head layer.
+ """
+
+ def __init__(self,
+ network,
+ num_classes,
+ decoder=None,
+ initializer=None,
+ dropout_rate=0.0,
+ head_name='glue',
+ cls_head=None,
+ name='classifier',
+ **kwargs):
+ """Init.
+
+ Args:
+ network:
+ A perceiver encode and processor transformer network. This network
+ should output a classification output. Furthermore, it should expose its
+ embedding table via a "get_embedding_table" method.
+ num_classes:
+ Number of classes to predict from the classification network.
+ decoder:
+ A perceiver decoder network. This network should accept the
+ latent output of the encoder and emits logits.
+ initializer:
+ The initializer (if any) to use in the classification networks.
+ Defaults to a Glorot uniform initializer.
+ dropout_rate:
+ The dropout probability of the cls head.
+ head_name:
+ Name of the classification head.
+ cls_head:
+ (Optional) The layer instance to use for the classifier head.
+ It should take in the output from network and produce the final logits.
+ If set, the arguments ('num_classes', 'initializer', 'dropout_rate',
+ 'use_encoder_pooler', 'head_name') will be ignored.
+ name:
+ Sets the `tf_keras.Model` name.
+ **kwargs:
+ Any keyword arguments to pass through to `tf_keras.Model`.
+ """
+ super().__init__(name=name, **kwargs)
+
+ self._config = {
+ 'network': network,
+ 'decoder': decoder,
+ 'num_classes': num_classes,
+ 'initializer': initializer,
+ 'dropout_rate': dropout_rate,
+ 'head_name': head_name,
+ 'cls_head': cls_head,
+ 'name': name,
+ }
+
+ self.num_classes = num_classes
+ self.head_name = head_name
+ self.initializer = initializer
+ self._decoder = decoder
+ self._network = network
+
+ inputs = self._network.inputs
+ outputs = self._network(inputs)
+
+ if 'sequence_output' not in outputs:
+ if 'latent_output' in outputs and self._decoder is not None:
+ decoder_inputs = {
+ 'latent_output': outputs['latent_output'],
+ 'input_mask': inputs['input_mask'],
+ }
+ decoder_outputs = self._decoder(decoder_inputs)
+ sequence_output = decoder_outputs['sequence_output']
+ else:
+ raise ValueError('if `sequence_output` is not in encoder output, '
+ '`latent_output` must be in encoder output and'
+ 'decoder must exist.')
+ else:
+ sequence_output = outputs['sequence_output']
+
+ cls_inputs = sequence_output
+
+ if initializer is None:
+ stddev = 1. / np.sqrt(cls_inputs.shape[-1])
+ initializer = tf_keras.initializers.TruncatedNormal(stddev=stddev)
+
+ if cls_head:
+ classifier = cls_head
+ else:
+ classifier = layers.ClassificationHead(
+ inner_dim=cls_inputs.shape[-1],
+ num_classes=num_classes,
+ initializer=initializer,
+ dropout_rate=dropout_rate,
+ name=head_name)
+
+ _ = classifier(cls_inputs)
+ self.inputs = inputs
+ self._cls_head = cls_head
+ self._name = name
+ self.classifier = classifier
+
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ """Return perceiver classifier model output tensors in a dict.
+
+ Accepts inputs as dictionary of tensors.
+ Args:
+ inputs:
+ A `Dict[str, tf_keras.Input]` with `input_word_ids`, `input_mask`, and
+ `input_type_ids`. The shapes are all `(None)` with dtype `tf.int32`.
+
+ Returns:
+ `tf.Tensor` classification output.
+ """
+ if not isinstance(inputs, dict):
+ raise ValueError(f'Unexpected inputs type to {self.__class__}.')
+
+ word_ids = inputs['input_word_ids']
+ input_type_ids = inputs.get('input_type_ids')
+ input_mask = inputs.get('input_mask')
+
+ encoder_inputs = {
+ 'input_word_ids': word_ids,
+ 'input_mask': input_mask,
+ 'input_type_ids': input_type_ids,
+ }
+ encoder_outputs = self._network(encoder_inputs)
+
+ if 'sequence_output' not in encoder_outputs:
+ if 'latent_output' in encoder_outputs:
+ z = encoder_outputs['latent_output']
+ decoder_inputs = {'latent_output': z, 'input_mask': input_mask}
+ decoder_output = self._decoder(decoder_inputs) # pyrefly: ignore[not-callable]
+
+ outputs = dict()
+ if isinstance(decoder_output, dict):
+ outputs = decoder_output
+ else:
+ raise ValueError('decoder\'s output should be a dict,'
+ f'but got {decoder_output}')
+ else:
+ raise ValueError('If `sequence_output` is not in encoder output,'
+ '`latent_output` must be in encoder output.')
+ else:
+ outputs = encoder_outputs
+
+ return self.classifier(outputs['sequence_output'])
+
+ @property
+ def checkpoint_items(self):
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(encoder=self._network, decoder=self._decoder)
+ if hasattr(self.classifier, 'checkpoint_items'):
+ for key, item in self.classifier.checkpoint_items.items():
+ items['.'.join([self.classifier.name, key])] = item
+ return items
+
+ def get_config(self):
+ """Return the configuration to set up this object using `from_config`."""
+ return self._config
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ """Initialize object using config from `get_config`.
+
+ https://www.tensorflow.org/api_docs/python/tf/keras/models/model_from_config
+
+ Args:
+ config:
+ Return the configuration to set up this object.
+ custom_objects:
+ Optional dictionary mapping names (strings) to custom classes or
+ functions to be considered during deserialization.
+ Returns:
+ A Keras model instance (uncompiled).
+ """
+ return cls(**config)
diff --git a/official/projects/perceiver/modeling/models/classifier_test.py b/official/projects/perceiver/modeling/models/classifier_test.py
new file mode 100644
index 00000000000..3f8a8801037
--- /dev/null
+++ b/official/projects/perceiver/modeling/models/classifier_test.py
@@ -0,0 +1,210 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for classifier."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+from official.projects.perceiver.configs import encoders
+from official.projects.perceiver.configs import perceiver as cfg
+from official.projects.perceiver.modeling.layers import decoder
+from official.projects.perceiver.modeling.models import classifier
+from official.projects.perceiver.modeling.networks import positional_decoder
+
+
+class ClassifierTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(('single_cls', 1), ('3_cls', 3))
+ def test_perceiver_trainer(self, num_classes):
+ """Validate that the Keras object can be created."""
+ # Build a perceiver sequence encoder network to use within the perceiver
+ # trainer.
+
+ vocab_size = 100
+ sequence_length = 512
+ d_model = 64
+ d_latents = 48
+ num_layers = 2
+ encoder_cfg = cfg.EncoderConfig(
+ v_last_dim=d_latents,
+ num_self_attends_per_block=num_layers)
+ sequence_encoder_cfg = cfg.SequenceEncoderConfig(
+ d_model=d_model,
+ d_latents=d_latents,
+ vocab_size=vocab_size,
+ encoder=encoder_cfg)
+ test_network = encoders.build_encoder(sequence_encoder_cfg)
+
+ deocder_cfg = cfg.DecoderConfig(
+ output_last_dim=d_latents,
+ v_last_dim=d_latents)
+ perceiver_classification_decoder_cfg = cfg.ClassificationDecoderConfig(
+ d_model=d_model,
+ decoder=deocder_cfg,
+ d_latents=d_latents)
+ decoder_ = decoder.Decoder(
+ **perceiver_classification_decoder_cfg.decoder.as_dict())
+ positional_decoder_ = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=perceiver_classification_decoder_cfg.output_index_dim,
+ z_index_dim=perceiver_classification_decoder_cfg.z_index_dim,
+ d_latents=perceiver_classification_decoder_cfg.d_latents,
+ d_model=perceiver_classification_decoder_cfg.d_model,
+ position_encoding_intializer_stddev=perceiver_classification_decoder_cfg
+ .position_encoding_intializer_stddev)
+
+ # Create a classifier with the created network.
+ trainer_model = classifier.Classifier(
+ network=test_network,
+ decoder=positional_decoder_,
+ num_classes=num_classes)
+
+ # Create a set of 2-dimensional inputs (the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+
+ # Invoke the trainer model on the inputs. This causes the layer to be built.
+ cls_outs = trainer_model({
+ 'input_word_ids': word_ids,
+ 'input_mask': mask,
+ 'input_type_ids': type_ids})
+
+ # Validate that the outputs are of the expected shape.
+ expected_classification_shape = [None, num_classes]
+ self.assertAllEqual(expected_classification_shape, cls_outs.shape.as_list())
+
+ @parameterized.named_parameters(
+ ('single_cls', 1, False),
+ ('2_cls', 2, False),
+ ('single_cls_custom_head', 1, True),
+ ('2_cls_custom_head', 2, True))
+ def test_perceiver_trainer_tensor_call(self, num_classes, use_custom_head):
+ """Validate that the Keras object can be invoked."""
+ # Build a perceiver sequence encoder network to use within the perceiver
+ # trainer.
+ vocab_size = 100
+ d_model = 64
+ d_latents = 48
+ num_layers = 2
+ encoder_cfg = cfg.EncoderConfig(
+ v_last_dim=d_latents,
+ num_self_attends_per_block=num_layers)
+ sequence_encoder_cfg = cfg.SequenceEncoderConfig(
+ d_model=d_model,
+ d_latents=d_latents,
+ vocab_size=vocab_size,
+ encoder=encoder_cfg)
+ test_network = encoders.build_encoder(sequence_encoder_cfg)
+
+ deocder_cfg = cfg.DecoderConfig(
+ output_last_dim=d_latents,
+ v_last_dim=d_latents)
+ perceiver_classification_decoder_cfg = cfg.ClassificationDecoderConfig(
+ d_model=d_model,
+ decoder=deocder_cfg,
+ d_latents=d_latents)
+ decoder_ = decoder.Decoder(
+ **perceiver_classification_decoder_cfg.decoder.as_dict())
+ positional_decoder_ = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=perceiver_classification_decoder_cfg.output_index_dim,
+ z_index_dim=perceiver_classification_decoder_cfg.z_index_dim,
+ d_latents=perceiver_classification_decoder_cfg.d_latents,
+ d_model=perceiver_classification_decoder_cfg.d_model,
+ position_encoding_intializer_stddev=perceiver_classification_decoder_cfg
+ .position_encoding_intializer_stddev)
+
+ cls_head = layers.GaussianProcessClassificationHead(
+ inner_dim=0, num_classes=num_classes) if use_custom_head else None
+
+ # Create a classifier with the created network.
+ trainer_model = classifier.Classifier(
+ network=test_network,
+ decoder=positional_decoder_,
+ cls_head=cls_head,
+ num_classes=num_classes)
+
+ # Create a set of 2-dimensional data tensors to feed into the model.
+ word_ids = tf.constant([[1, 1], [2, 2]], dtype=tf.int32)
+ mask = tf.constant([[1, 1], [1, 0]], dtype=tf.int32)
+ type_ids = tf.constant([[1, 1], [2, 2]], dtype=tf.int32)
+
+ # Invoke the trainer model on the tensors. In Eager mode, this does the
+ # actual calculation. (We can't validate the outputs, since the network is
+ # too complex: this simply ensures we're not hitting runtime errors.)
+ _ = trainer_model({
+ 'input_word_ids': word_ids,
+ 'input_mask': mask,
+ 'input_type_ids': type_ids})
+
+ @parameterized.named_parameters(
+ ('default_cls_head', None),
+ ('sngp_cls_head', layers.GaussianProcessClassificationHead(
+ inner_dim=0, num_classes=4)))
+ def test_serialize_deserialize(self, cls_head):
+ """Validate that the trainer can be serialized and deserialized."""
+ del cls_head
+ vocab_size = 100
+ d_model = 64
+ d_latents = 48
+ num_layers = 2
+ encoder_cfg = cfg.EncoderConfig(
+ v_last_dim=d_latents,
+ num_self_attends_per_block=num_layers)
+ sequence_encoder_cfg = cfg.SequenceEncoderConfig(
+ d_model=d_model,
+ d_latents=d_latents,
+ vocab_size=vocab_size,
+ encoder=encoder_cfg)
+ test_network = encoders.build_encoder(sequence_encoder_cfg)
+
+ deocder_cfg = cfg.DecoderConfig(
+ output_last_dim=d_latents,
+ v_last_dim=d_latents)
+ perceiver_classification_decoder_cfg = cfg.ClassificationDecoderConfig(
+ d_model=d_model,
+ decoder=deocder_cfg,
+ d_latents=d_latents)
+ decoder_ = decoder.Decoder(
+ **perceiver_classification_decoder_cfg.decoder.as_dict())
+ positional_decoder_ = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=perceiver_classification_decoder_cfg.output_index_dim,
+ z_index_dim=perceiver_classification_decoder_cfg.z_index_dim,
+ d_latents=perceiver_classification_decoder_cfg.d_latents,
+ d_model=perceiver_classification_decoder_cfg.d_model,
+ position_encoding_intializer_stddev=perceiver_classification_decoder_cfg
+ .position_encoding_intializer_stddev)
+
+ # Create a classifier with the created network.
+ trainer_model = classifier.Classifier(
+ network=test_network,
+ decoder=positional_decoder_,
+ num_classes=4)
+
+ # Create another trainer via serialization and deserialization.
+ config = trainer_model.get_config()
+ new_trainer_model = classifier.Classifier.from_config(config)
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(trainer_model.get_config(),
+ new_trainer_model.get_config())
+
+# TODO(b/222634115) add test coverage.
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/perceiver/modeling/models/pretrainer.py b/official/projects/perceiver/modeling/models/pretrainer.py
new file mode 100644
index 00000000000..33c9225725b
--- /dev/null
+++ b/official/projects/perceiver/modeling/models/pretrainer.py
@@ -0,0 +1,218 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver networks."""
+
+import copy
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+
+
+class Pretrainer(tf_keras.Model):
+ """Perceiver Pretrainer.
+
+ Adds the masked language model head upon the encoder output. Optionally
+ incorporates decoder output.
+ Forked from
+ (https://github.com/tensorflow/models/blob/master/official/nlp/modeling/models/bert_pretrainer.py)
+
+ Attributes:
+ encoder:
+ A perceiver encode and processor transformer network. This network
+ should output a classification output. Furthermore, it should expose its
+ embedding table via a "get_embedding_table" method.
+ masked_lm:
+ Masked language model network head for language modeling with encoder
+ and optionally decoded output.
+ inputs:
+ A `Dict[str, tf_keras.Input]` with `input_word_ids`, `input_mask`, and
+ `input_type_ids`. The shapes are all `(None)` with dtype `tf.int32`.
+ If `masked_lm_positions` is included, it will run masked language
+ modeling layer to return sequence of logits.
+ """
+
+ def __init__(self,
+ encoder,
+ decoder=None,
+ mlm_activation=None,
+ mlm_initializer='glorot_uniform',
+ customized_masked_lm=None,
+ name='pretrainer',
+ **kwargs):
+ """Init.
+
+ Args:
+ encoder:
+ A perceiver encode and processor transformer network. It should expose
+ its embedding table via a "get_embedding_table" method. Decoder won't
+ be used if `sequence_output` is in the output of the encoder.
+ decoder:
+ A perceiver decoder network. This parameter is optional. This layer
+ accepts the latent output of the encoder and emits logits. Decoder must
+ accept a dictionary of `latent_output` and `input_mask` as inputs. This
+ will not be used if `sequence_output` is an output from `encoder`.
+ mlm_activation:
+ The activation (if any) to use in the masked LM network. If `None`, no
+ activation will be used.
+ mlm_initializer:
+ The initializer (if any) to use in the masked LM. Default
+ to a Glorot uniform initializer.
+ customized_masked_lm:
+ A customized masked_lm layer. If None, will create
+ a standard layer from `layers.MaskedLM`; if not None, will use the
+ specified masked_lm layer. Above arguments `mlm_activation` and
+ `mlm_initializer` will be ignored.
+ name:
+ Sets the `tf_keras.Model` name.
+ **kwargs:
+ Any keyword arguments to pass through to `tf_keras.Model`.
+ """
+ super().__init__(**kwargs, name=name)
+
+ self._config = {
+ 'encoder': encoder,
+ 'decoder': decoder,
+ 'mlm_initializer': mlm_initializer,
+ 'mlm_activation': mlm_activation,
+ 'customized_masked_lm': customized_masked_lm,
+ 'name': name,
+ }
+
+ self._decoder = decoder
+ self.encoder = encoder
+ encoder_inputs = self.encoder.inputs
+
+ # Makes sure the weights are built.
+ encoder_outputs = self.encoder(encoder_inputs)
+
+ if 'sequence_output' not in encoder_outputs:
+ if 'latent_output' in encoder_outputs and self._decoder is not None:
+ decoder_inputs = {
+ 'latent_output': encoder_outputs['latent_output'],
+ 'input_mask': encoder_inputs['input_mask'],
+ }
+ decoder_outputs = self._decoder(decoder_inputs)
+ if 'sequence_output' not in decoder_outputs:
+ raise ValueError('`sequence_output` must be in decoder output.')
+ else:
+ raise ValueError('if `sequence_output` is not in encoder output, '
+ '`latent_output` must be in encoder output and'
+ 'decoder must exist.')
+
+ encoder_inputs = copy.copy(self.encoder.inputs)
+ inputs = dict(encoder_inputs)
+
+ if self._decoder is not None:
+ inputs.update(copy.copy(self._decoder.inputs))
+
+ self.masked_lm = customized_masked_lm or layers.MaskedLM(
+ embedding_table=self.encoder.get_embedding_table(),
+ activation=mlm_activation,
+ initializer=mlm_initializer,
+ name='cls/predictions')
+ masked_lm_positions = tf_keras.layers.Input(
+ shape=(None,), name='masked_lm_positions', dtype=tf.int32)
+
+ if isinstance(inputs, dict):
+ inputs['masked_lm_positions'] = masked_lm_positions
+ else:
+ raise ValueError(f'Unexpected inputs type to {self.__class__}.')
+ self.inputs = inputs
+
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
+ """Return perceiver pretrainer model output tensors in a dict.
+
+ Accepts inputs as dictionary of tensors.
+ Args:
+ inputs:
+ A `Dict[str, tf_keras.Input]` with `input_word_ids`, `input_mask`, and
+ `input_type_ids`. The shapes are all `(None)` with dtype `tf.int32`.
+ If `masked_lm_positions` is included, it will run masked language
+ modeling layer to return sequence of logits.
+
+ Returns:
+ `Dict[str, tf.Tensor]` with `sequence_output` and optionally
+ `mlm_logits`.
+ """
+ if not isinstance(inputs, dict):
+ raise ValueError(f'Unexpected inputs type to {self.__class__}.')
+
+ word_ids = inputs['input_word_ids']
+ input_type_ids = inputs.get('input_type_ids')
+ input_mask = inputs.get('input_mask')
+
+ encoder_inputs = {
+ 'input_word_ids': word_ids,
+ 'input_mask': input_mask,
+ 'input_type_ids': input_type_ids,
+ }
+ encoder_outputs = self.encoder(encoder_inputs)
+
+ if 'sequence_output' not in encoder_outputs:
+ if 'latent_output' in encoder_outputs:
+ z = encoder_outputs['latent_output']
+ decoder_inputs = {'latent_output': z, 'input_mask': input_mask}
+ decoder_output = self._decoder(decoder_inputs) # pyrefly: ignore[not-callable]
+
+ outputs = dict()
+ if isinstance(decoder_output, dict):
+ outputs = decoder_output
+ else:
+ raise ValueError('decoder\'s output should be a dict,'
+ f'but got {decoder_output}')
+ else:
+ raise ValueError('If `sequence_output` is not in encoder output,'
+ '`latent_output` must be in encoder output.')
+ else:
+ outputs = encoder_outputs
+
+ sequence_output = outputs['sequence_output']
+ # Inference may not have masked_lm_positions and mlm_logits is not needed.
+ if 'masked_lm_positions' in inputs:
+ masked_lm_positions = inputs['masked_lm_positions']
+ outputs['mlm_logits'] = self.masked_lm(
+ sequence_output, masked_positions=masked_lm_positions)
+ return outputs
+
+ @property
+ def checkpoint_items(self):
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(
+ encoder=self.encoder,
+ masked_lm=self.masked_lm,
+ decoder=self._decoder)
+ return items
+
+ def get_config(self):
+ """Return the configuration to set up this object using `from_config`."""
+ return self._config
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ """Initialize object using config from `get_config`.
+
+ https://www.tensorflow.org/api_docs/python/tf/keras/models/model_from_config
+
+ Args:
+ config:
+ Return the configuration to set up this object.
+ custom_objects:
+ Optional dictionary mapping names (strings) to custom classes or
+ functions to be considered during deserialization.
+ Returns:
+ A Keras model instance (uncompiled).
+ """
+ return cls(**config)
diff --git a/official/projects/perceiver/modeling/models/pretrainer_test.py b/official/projects/perceiver/modeling/models/pretrainer_test.py
new file mode 100644
index 00000000000..c2755747201
--- /dev/null
+++ b/official/projects/perceiver/modeling/models/pretrainer_test.py
@@ -0,0 +1,164 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for Perceiver pretrainer model."""
+import itertools
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+from official.projects.perceiver.configs import encoders
+from official.projects.perceiver.configs import perceiver as cfg
+from official.projects.perceiver.modeling.layers import decoder
+from official.projects.perceiver.modeling.models import pretrainer
+from official.projects.perceiver.modeling.networks import positional_decoder
+
+
+class PretrainerTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.parameters(itertools.product(
+ (False, True),
+ (False, True),
+ ))
+ def test_perceiver_pretrainer(self, use_customized_masked_lm,
+ has_masked_lm_positions):
+ """Validate that the Keras object can be created."""
+ # Build a transformer network to use within the Perceiver trainer.
+ vocab_size = 100
+ sequence_length = 512
+ d_model = 64
+ d_latents = 48
+ num_layers = 2
+ encoder_cfg = cfg.EncoderConfig(
+ v_last_dim=d_latents,
+ num_self_attends_per_block=num_layers)
+ sequence_encoder_cfg = cfg.SequenceEncoderConfig(
+ d_model=d_model,
+ d_latents=d_latents,
+ vocab_size=vocab_size,
+ encoder=encoder_cfg)
+ test_network = encoders.build_encoder(sequence_encoder_cfg)
+
+ _ = test_network(test_network.inputs)
+
+ deocder_cfg = cfg.DecoderConfig(
+ output_last_dim=d_latents,
+ v_last_dim=d_latents)
+ perceiver_mlm_decoder_cfg = cfg.MaskedLMDecoderConfig(
+ d_model=d_model,
+ decoder=deocder_cfg,
+ d_latents=d_latents)
+ decoder_ = decoder.Decoder(
+ **perceiver_mlm_decoder_cfg.decoder.as_dict())
+ positional_decoder_ = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=perceiver_mlm_decoder_cfg.output_index_dim,
+ z_index_dim=perceiver_mlm_decoder_cfg.z_index_dim,
+ d_latents=perceiver_mlm_decoder_cfg.d_latents,
+ d_model=perceiver_mlm_decoder_cfg.d_model,
+ position_encoding_intializer_stddev=perceiver_mlm_decoder_cfg
+ .position_encoding_intializer_stddev)
+
+ if use_customized_masked_lm:
+ customized_masked_lm = layers.MaskedLM(
+ embedding_table=test_network.get_embedding_table())
+ else:
+ customized_masked_lm = None
+
+ # Create a Perceiver trainer with the created network.
+ perceiver_trainer_model = pretrainer.Pretrainer(
+ encoder=test_network,
+ decoder=positional_decoder_,
+ customized_masked_lm=customized_masked_lm)
+ num_token_predictions = 20
+ # Create a set of 2-dimensional inputs (the first dimension is implicit).
+ inputs = dict(
+ input_word_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(sequence_length,), dtype=tf.int32))
+ if has_masked_lm_positions:
+ inputs['masked_lm_positions'] = tf_keras.Input(
+ shape=(num_token_predictions,), dtype=tf.int32)
+
+ # Invoke the trainer model on the inputs. This causes the layer to be built.
+ outputs = perceiver_trainer_model(inputs)
+
+ expected_keys = ['sequence_output']
+ if has_masked_lm_positions:
+ expected_keys.append('mlm_logits')
+
+ self.assertSameElements(outputs.keys(), expected_keys)
+ # Validate that the outputs are of the expected shape.
+ expected_lm_shape = [None, num_token_predictions, vocab_size]
+ if has_masked_lm_positions:
+ self.assertAllEqual(expected_lm_shape,
+ outputs['mlm_logits'].shape.as_list())
+
+ expected_sequence_output_shape = [None, sequence_length, d_model]
+ self.assertAllEqual(expected_sequence_output_shape,
+ outputs['sequence_output'].shape.as_list())
+
+ def test_serialize_deserialize(self):
+ """Validate that the trainer can be serialized and deserialized."""
+ vocab_size = 100
+ d_model = 64
+ d_latents = 48
+ num_layers = 2
+ encoder_cfg = cfg.EncoderConfig(
+ v_last_dim=d_latents,
+ num_self_attends_per_block=num_layers)
+ sequence_encoder_cfg = cfg.SequenceEncoderConfig(
+ d_model=d_model,
+ d_latents=d_latents,
+ vocab_size=vocab_size,
+ encoder=encoder_cfg)
+ test_network = encoders.build_encoder(sequence_encoder_cfg)
+
+ _ = test_network(test_network.inputs)
+
+ deocder_cfg = cfg.DecoderConfig(
+ output_last_dim=d_latents,
+ v_last_dim=d_latents)
+ perceiver_mlm_decoder_cfg = cfg.MaskedLMDecoderConfig(
+ d_model=d_model,
+ decoder=deocder_cfg,
+ d_latents=d_latents)
+ decoder_ = decoder.Decoder(
+ **perceiver_mlm_decoder_cfg.decoder.as_dict())
+ positional_decoder_ = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=perceiver_mlm_decoder_cfg.output_index_dim,
+ z_index_dim=perceiver_mlm_decoder_cfg.z_index_dim,
+ d_latents=perceiver_mlm_decoder_cfg.d_latents,
+ d_model=perceiver_mlm_decoder_cfg.d_model,
+ position_encoding_intializer_stddev=perceiver_mlm_decoder_cfg
+ .position_encoding_intializer_stddev)
+
+ # Create a Perceiver trainer with the created network.
+ perceiver_trainer_model = pretrainer.Pretrainer(
+ encoder=test_network,
+ decoder=positional_decoder_)
+
+ config = perceiver_trainer_model.get_config()
+ new_perceiver_trainer_model = pretrainer.Pretrainer.from_config(config)
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(perceiver_trainer_model.get_config(),
+ new_perceiver_trainer_model.get_config())
+
+# TODO(b/222634115) add test coverage.
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/perceiver/modeling/networks/positional_decoder.py b/official/projects/perceiver/modeling/networks/positional_decoder.py
new file mode 100644
index 00000000000..e0e4fe6860f
--- /dev/null
+++ b/official/projects/perceiver/modeling/networks/positional_decoder.py
@@ -0,0 +1,127 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver networks."""
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+
+
+class PositionalDecoder(tf_keras.layers.Layer):
+ """Perceiver Positional Decoder Network.
+
+ Creates a position encoding for queries and composes basic decoder.
+ e.g. the positional decoder can be used to do MLM, classification, or
+ regression.
+
+ Currently only supports positional decoding.
+
+ Use `self.inputs` for inputs.
+
+ Attributes:
+ inputs: A `Dict[Text, tf_keras.Input]` with `latent_output` and
+ `input_mask`. The shape of `latent_output` is shape
+ `(z_index_dim, d_latents)` with dtype `tf.float32` and `input_mask` is
+ shape `(None)` with dtype `tf.int32`.
+ """
+
+ def __init__(self,
+ decoder,
+ output_index_dim,
+ z_index_dim,
+ d_latents,
+ d_model,
+ position_encoding_intializer_stddev=0.02,
+ name='positional_decoder',
+ **kwargs):
+ """Init.
+
+ Args:
+ decoder:
+ Instance of perceiver `Decoder`.
+ output_index_dim:
+ Sequence length for the query encoding.
+ z_index_dim:
+ Latent index dimension.
+ d_latents:
+ Latent last dimension.
+ d_model:
+ Model last dimension.
+ position_encoding_intializer_stddev:
+ `stddev` of `tf_keras.initializers.TruncatedNormal` used for the
+ learned position embedding table kernel initializer.
+ name:
+ Sets the `tf_keras.layers.Layer` name.
+ **kwargs:
+ Any keyword arguments to pass through to `tf_keras.layers.Layer`.
+ """
+ super().__init__(**kwargs, name=name)
+
+ self._decoder = decoder
+ self._output_index_dim = output_index_dim
+ self._z_index_dim = z_index_dim
+ self._d_latents = d_latents
+ self._d_model = d_model
+
+ self._output_pos_enc = self._create_decoder_query(
+ position_encoding_intializer_stddev)
+
+ self.inputs = dict(
+ latent_output=tf_keras.Input(
+ shape=(self._z_index_dim, self._d_latents),
+ dtype=tf.float32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32))
+
+ def _create_decoder_query(self, position_encoding_intializer_stddev):
+ """Create the position encoding for the output query."""
+ return layers.PositionEmbedding(
+ max_length=self._output_index_dim,
+ name='decoder_pos_enc',
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=position_encoding_intializer_stddev))
+
+ def call(self, inputs, training=None):
+ """Return decoded output of latent vector.
+
+ Uses the positional encoding as query for the decoder and uses the
+ `latent_output` as key-value for the decoder.
+ Args:
+ inputs:
+ A `Dict[Text, tf_keras.Input]` with `latent_output` and
+ `input_mask`. The shape of `latent_output` is shape
+ `(z_index_dim, d_latents)` with dtype `tf.float32` and `input_mask` is
+ shape `(None)` with dtype `tf.int32`.
+ training:
+ Flag to indicate training status. Default is `None`. It is passed to
+ the decoder as is.
+
+ Returns:
+ `Dict[Text, tf.Tensor]` decoded `sequence_output` of a latent vector.
+ """
+ if not isinstance(inputs, dict):
+ raise ValueError(f'Unexpected inputs type to {self.__class__}.')
+
+ latent_output = inputs['latent_output']
+ query_mask = inputs.get('input_mask')
+ decoder_query = self._output_pos_enc(tf.ones(
+ (tf.shape(latent_output)[0], self._output_index_dim, self._d_model),
+ dtype=latent_output.dtype))
+ z = latent_output
+
+ sequence_output = self._decoder(
+ [decoder_query, z],
+ query_mask=query_mask,
+ training=training)
+ return dict(sequence_output=sequence_output)
diff --git a/official/projects/perceiver/modeling/networks/positional_decoder_test.py b/official/projects/perceiver/modeling/networks/positional_decoder_test.py
new file mode 100644
index 00000000000..90ecc63e4ea
--- /dev/null
+++ b/official/projects/perceiver/modeling/networks/positional_decoder_test.py
@@ -0,0 +1,110 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for positional_decoder."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.perceiver.configs import perceiver as cfg
+from official.projects.perceiver.modeling.layers import decoder
+from official.projects.perceiver.modeling.networks import positional_decoder
+
+
+class PositionalDecoderTest(tf.test.TestCase):
+
+ def test_dict_outputs_network_creation(self):
+ sequence_length = 21
+ z_index_dim = 8
+ d_model = 64
+ d_latents = 48
+ decoder_cfg = cfg.DecoderConfig(
+ output_last_dim=d_latents,
+ v_last_dim=d_latents,
+ num_heads=2)
+ positional_decoder_cfg = cfg.PositionalDecoder(
+ decoder=decoder_cfg,
+ d_model=d_model,
+ d_latents=d_latents,
+ output_index_dim=sequence_length,
+ z_index_dim=z_index_dim)
+
+ decoder_ = decoder.Decoder(positional_decoder_cfg.decoder.as_dict())
+ mlm_decoder = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=positional_decoder_cfg.output_index_dim,
+ z_index_dim=positional_decoder_cfg.z_index_dim,
+ d_latents=positional_decoder_cfg.d_latents,
+ d_model=positional_decoder_cfg.d_model)
+
+ # Create the inputs (note that the first dimension is implicit).
+ latent_output = tf_keras.Input(
+ shape=(z_index_dim, d_latents), dtype=tf.float32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = mlm_decoder(
+ dict(latent_output=latent_output, input_mask=mask))
+ data = dict_outputs["sequence_output"]
+
+ expected_data_shape = [None, sequence_length, d_model]
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, data.dtype)
+
+ def test_serialize_deserialize(self):
+ # Create a network object that sets all of its config options.
+ sequence_length = 21
+ z_index_dim = 8
+ d_model = 64
+ d_latents = 48
+ decoder_cfg = cfg.DecoderConfig(
+ output_last_dim=d_latents,
+ v_last_dim=d_latents,
+ num_heads=2)
+ positional_decoder_cfg = cfg.PositionalDecoder(
+ decoder=decoder_cfg,
+ d_model=d_model,
+ d_latents=d_latents,
+ output_index_dim=sequence_length,
+ z_index_dim=z_index_dim)
+
+ decoder_ = decoder.Decoder(positional_decoder_cfg.decoder.as_dict())
+ mlm_decoder = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=positional_decoder_cfg.output_index_dim,
+ z_index_dim=positional_decoder_cfg.z_index_dim,
+ d_latents=positional_decoder_cfg.d_latents,
+ d_model=positional_decoder_cfg.d_model)
+
+ # Create the inputs (note that the first dimension is implicit).
+ latent_output = tf_keras.Input(
+ shape=(z_index_dim, d_latents), dtype=tf.float32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = mlm_decoder(
+ dict(latent_output=latent_output, input_mask=mask))
+ data = dict_outputs["sequence_output"]
+
+ # Create a model based off of this network:
+ # model =
+ _ = tf_keras.Model([latent_output, mask], [data])
+
+ # TODO(b/222634115) make save work.
+ # Tests model saving/loading.
+ # model_path = self.get_temp_dir() + "/model"
+ # model.save(model_path)
+ # _ = tf_keras.models.load_model(model_path)
+
+# TODO(b/222634115) add test coverage.
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/perceiver/modeling/networks/sequence_encoder.py b/official/projects/perceiver/modeling/networks/sequence_encoder.py
new file mode 100644
index 00000000000..72bc74bf53f
--- /dev/null
+++ b/official/projects/perceiver/modeling/networks/sequence_encoder.py
@@ -0,0 +1,156 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Perceiver sequence encoder."""
+
+from typing import Optional, Dict
+
+import tensorflow as tf, tf_keras
+
+from official.nlp.modeling import layers
+
+
+class SequenceEncoder(tf_keras.layers.Layer):
+ """Perceiver encoder for sequences.
+
+ Assumes positional learned encoding for latent inputs and embeddings. Creates
+ an embedding table with vocab size. It uses the perceiver encode processor
+ to encode the input and process the latent representation. It can be
+ pretrained on masked LM and reused for fine-tuning.
+
+ Use `self.inputs` for inputs.
+ """
+
+ def __init__(self,
+ encoder: tf_keras.layers.Layer,
+ d_model: int,
+ d_latents: int,
+ z_index_dim: int,
+ max_seq_len: int,
+ vocab_size: int,
+ z_pos_enc_init_scale: float = 0.02,
+ embedding_width: Optional[int] = None,
+ embedding_initializer_stddev: float = 0.02,
+ input_position_encoding_intializer_stddev: float = 0.02,
+ name: str = 'sequence_encoder',
+ **kwargs):
+ """Init.
+
+ Args:
+ encoder:
+ Instance of perceiver `Encoder`.
+ d_model:
+ Last dimension size of the input and output tensors. e.g.
+ `[batch_size, max_seq_len, d_model]`.
+ d_latents:
+ Last dimension size of the latent tensors. e.g.
+ `[batch_size, z_index_dim, d_latents]`.
+ z_index_dim:
+ Second dimension size of the latent tensors. e.g.
+ `[batch_size, z_index_dim, d_latents]`.
+ max_seq_len:
+ Second dimension size of the input and outputs tensors. e.g.
+ `[batch_size, max_seq_len, d_model]`.
+ vocab_size:
+ Vocabulary size of the embedding table.
+ z_pos_enc_init_scale:
+ Latent array's positional encoding's truncated_normal initializer's
+ `stddev`.
+ embedding_width:
+ Embedding dimension of the embedding table.
+ embedding_initializer_stddev:
+ `stddev` of `tf_keras.initializers.TruncatedNormal` used for the
+ embedding table kernel initializer.
+ input_position_encoding_intializer_stddev:
+ `stddev` of `tf_keras.initializers.TruncatedNormal` used for the
+ learned position embedding table kernel initializer.
+ name:
+ Sets the `tf_keras.layers.Layer` name.
+ **kwargs:
+ Any keyword arguments to pass through to `tf_keras.layers.Layer`.
+ """
+ super().__init__(**kwargs, name=name)
+
+ self._embedding_width = embedding_width
+
+ self._encoder = encoder
+
+ self._d_model = d_model
+ self._z_index_dim = z_index_dim
+ self._d_latents = d_latents
+ if self._embedding_width is None:
+ self._embedding_width = self._d_model
+
+ # Construct the embeddling layer for the sequence vocab.
+ self._embedding_layer = layers.OnDeviceEmbedding(
+ vocab_size=vocab_size,
+ embedding_width=self._embedding_width,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=embedding_initializer_stddev),
+ name='word_embeddings')
+
+ # Construct the input positional encoding layer.
+ self._input_pos_encoding = layers.PositionEmbedding(
+ max_length=max_seq_len,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=input_position_encoding_intializer_stddev),
+ name='input_pos_encoding')
+
+ # Construct the latent array initial state.
+ self._z_pos_enc = layers.PositionEmbedding(
+ max_length=z_index_dim,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=z_pos_enc_init_scale),
+ name='z_pos_enc')
+
+ self.inputs = dict(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32))
+
+ def get_embedding_table(self) -> tf.Variable:
+ """Get embedding table."""
+ return self._embedding_layer.embeddings
+
+ def call(self,
+ inputs: Dict[str, tf.Tensor],
+ training: Optional[bool] = None) -> Dict[str, tf.Tensor]:
+ """Return encoded and processed latent output of inputs.
+
+ Args:
+ inputs:
+ Expect inputs to be a dictionary of `input_word_ids` and `input_mask`.
+ training:
+ Flag to indicate training status.
+
+ Returns:
+ `Dict[str, tf.Tensor]` decoded output of latent vector via the query.
+ """
+ if not isinstance(inputs, dict):
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+ word_ids = inputs['input_word_ids']
+ input_mask = inputs.get('input_mask')
+
+ word_embeddings = self._embedding_layer(word_ids)
+ pos_encodings = self._input_pos_encoding(word_embeddings)
+ embeddings = word_embeddings + pos_encodings
+
+ tensor_for_shape = tf.ones(
+ [tf.shape(embeddings)[0], self._z_index_dim, self._d_latents],
+ dtype=embeddings.dtype)
+ encoder_query = self._z_pos_enc(tensor_for_shape)
+
+ z = self._encoder(
+ [embeddings, encoder_query], input_mask=input_mask, training=training)
+ return dict(latent_output=z)
diff --git a/official/projects/perceiver/modeling/networks/sequence_encoder_test.py b/official/projects/perceiver/modeling/networks/sequence_encoder_test.py
new file mode 100644
index 00000000000..b31e3ed3003
--- /dev/null
+++ b/official/projects/perceiver/modeling/networks/sequence_encoder_test.py
@@ -0,0 +1,158 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for sequence_encoder."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.perceiver.configs import encoders
+from official.projects.perceiver.configs import perceiver
+from official.projects.perceiver.modeling.layers import encoder
+from official.projects.perceiver.modeling.networks import sequence_encoder
+
+
+class SequenceEncoderTest(tf.test.TestCase):
+
+ def _create_small_network(
+ self,
+ sequence_length,
+ z_index_dim,
+ d_latents,
+ vocab_size=100):
+ d_model = 64
+ num_layers = 2
+ encoder_cfg = perceiver.EncoderConfig(
+ v_last_dim=d_latents,
+ num_self_attends_per_block=num_layers)
+ sequence_encoder_cfg = perceiver.SequenceEncoderConfig(
+ d_model=d_model,
+ d_latents=d_latents,
+ z_index_dim=z_index_dim,
+ max_seq_len=sequence_length,
+ vocab_size=vocab_size,
+ encoder=encoder_cfg)
+ return encoders.build_encoder(sequence_encoder_cfg)
+
+ def test_dict_outputs_network_creation(self):
+ sequence_length = 21
+ z_index_dim = 128
+ d_latents = 48
+ test_network = self._create_small_network(
+ sequence_length=sequence_length,
+ z_index_dim=z_index_dim,
+ d_latents=d_latents)
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = test_network(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ data = dict_outputs["latent_output"]
+
+ expected_data_shape = [None, z_index_dim, d_latents]
+ self.assertAllEqual(expected_data_shape, data.shape.as_list())
+
+ # The default output dtype is float32.
+ self.assertAllEqual(tf.float32, data.dtype)
+
+ def test_dict_outputs_network_invocation(self):
+ num_types = 7
+ vocab_size = 57
+ sequence_length = 21
+ z_index_dim = 128
+ d_latents = 48
+ test_network = self._create_small_network(
+ sequence_length=sequence_length,
+ z_index_dim=z_index_dim,
+ d_latents=d_latents,
+ vocab_size=vocab_size)
+ # Create the inputs (note that the first dimension is implicit).
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ dict_outputs = test_network(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ data = dict_outputs["latent_output"]
+
+ # Create a model based off of this network:
+ model = tf_keras.Model([word_ids, mask, type_ids], [data])
+
+ # Invoke the model. We can't validate the output data here (the model is too
+ # complex) but this will catch structural runtime errors.
+ batch_size = 3
+ word_id_data = np.random.randint(
+ vocab_size, size=(batch_size, sequence_length))
+ mask_data = np.random.randint(2, size=(batch_size, sequence_length))
+ type_id_data = np.random.randint(
+ num_types, size=(batch_size, sequence_length))
+ outputs = model.predict([word_id_data, mask_data, type_id_data])
+ self.assertEqual(outputs[0].shape[1], d_latents)
+
+ def test_serialize_deserialize(self):
+ # Create a network object that sets all of its config options.
+ sequence_length = 21
+ vocab_size = 57
+ d_model = 64
+ d_latents = 48
+ z_index_dim = 128
+ num_layers = 2
+ encoder_cfg = perceiver.EncoderConfig(
+ v_last_dim=d_latents,
+ num_self_attends_per_block=num_layers)
+ sequence_encoder_config = perceiver.SequenceEncoderConfig(
+ d_model=d_model,
+ d_latents=d_latents,
+ z_index_dim=z_index_dim,
+ max_seq_len=sequence_length,
+ vocab_size=vocab_size,
+ encoder=encoder_cfg)
+ encoder_ = encoder.Encoder(
+ **sequence_encoder_config.encoder.as_dict())
+ network = sequence_encoder.SequenceEncoder(
+ encoder=encoder_,
+ d_model=sequence_encoder_config.d_model,
+ d_latents=sequence_encoder_config.d_latents,
+ z_index_dim=sequence_encoder_config.z_index_dim,
+ max_seq_len=sequence_encoder_config.max_seq_len,
+ vocab_size=sequence_encoder_config.vocab_size,
+ z_pos_enc_init_scale=sequence_encoder_config.z_pos_enc_init_scale,
+ embedding_width=sequence_encoder_config.embedding_width,
+ embedding_initializer_stddev=sequence_encoder_config
+ .embedding_initializer_stddev,
+ input_position_encoding_intializer_stddev=sequence_encoder_config
+ .input_position_encoding_intializer_stddev)
+
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+
+ dict_outputs = network(
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ data = dict_outputs["latent_output"]
+
+ # Create a model based off of this network:
+ # model =
+ _ = tf_keras.Model([word_ids, mask, type_ids], [data])
+
+ # TODO(b/222634115) make save work.
+ # Tests model saving/loading.
+ # model_path = self.get_temp_dir() + "/model"
+ # model.save(model_path)
+ # _ = tf_keras.models.load_model(model_path)
+
+# TODO(b/222634115) add test coverage.
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/perceiver/perceiver.ipynb b/official/projects/perceiver/perceiver.ipynb
new file mode 100644
index 00000000000..6a8bb576992
--- /dev/null
+++ b/official/projects/perceiver/perceiver.ipynb
@@ -0,0 +1,4432 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "h2tmtetBZsTE"
+ },
+ "source": [
+ "##### Copyright 2023 The TensorFlow Authors."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "DsTOPNzMZsTT"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n",
+ "# you may not use this file except in compliance with the License.\n",
+ "# You may obtain a copy of the License at\n",
+ "#\n",
+ "# https://www.apache.org/licenses/LICENSE-2.0\n",
+ "#\n",
+ "# Unless required by applicable law or agreed to in writing, software\n",
+ "# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
+ "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
+ "# See the License for the specific language governing permissions and\n",
+ "# limitations under the License."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "YRlNzHiLN87u"
+ },
+ "source": [
+ "# Train Perceiver model for down streaming tasks\n",
+ "\n",
+ "This tutorial demonstrates how to use [Perciever](https://arxiv.org/abs/2107.14795) model for down streaming tasks using Tensorflow Model Garden.\n",
+ "\n",
+ "[Tensorflow Model Garden](https://www.tensorflow.org/tfmodels) contains a collection of state-of-the-art models, implemented with TensorFlow's high-level APIs. The implementations demonstrate the best practices for modeling, letting users to take full advantage of TensorFlow for their research and product development."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "pRFcanY4y5lV"
+ },
+ "source": [
+ "\n",
+ "## Clone models repository"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "IPIuUUMxy5Ei"
+ },
+ "outputs": [],
+ "source": [
+ "!git clone -q https://github.com/tensorflow/models.git"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "B7OUNfFxymIm"
+ },
+ "source": [
+ "## Install necessary dependencies"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 97312,
+ "status": "ok",
+ "timestamp": 1689009983377,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "bpSKXUuQpVaW",
+ "outputId": "7d0006a4-68bf-4aad-d49f-c93a115b32fe"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m524.1/524.1 MB\u001b[0m \u001b[31m3.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.7/1.7 MB\u001b[0m \u001b[31m86.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m5.6/5.6 MB\u001b[0m \u001b[31m95.2 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m440.8/440.8 kB\u001b[0m \u001b[31m44.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m241.2/241.2 kB\u001b[0m \u001b[31m5.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m175.1/175.1 kB\u001b[0m \u001b[31m19.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n",
+ " Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n",
+ " Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m43.6/43.6 kB\u001b[0m \u001b[31m4.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.3/1.3 MB\u001b[0m \u001b[31m35.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m118.9/118.9 kB\u001b[0m \u001b[31m14.3 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m17.6/17.6 MB\u001b[0m \u001b[31m84.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h Building wheel for pyyaml (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n",
+ " Building wheel for seqeval (setup.py) ... \u001b[?25l\u001b[?25hdone\n",
+ "\u001b[33m WARNING: The scripts f2py, f2py3 and f2py3.10 are installed in '/root/.local/bin' which is not on PATH.\n",
+ " Consider adding this directory to PATH or, if you prefer to suppress this warning, use --no-warn-script-location.\u001b[0m\u001b[33m\n",
+ "\u001b[0m\u001b[33m WARNING: The script sacrebleu is installed in '/root/.local/bin' which is not on PATH.\n",
+ " Consider adding this directory to PATH or, if you prefer to suppress this warning, use --no-warn-script-location.\u001b[0m\u001b[33m\n",
+ "\u001b[0m\u001b[31mERROR: pip's dependency resolver does not currently take into account all the packages that are installed. This behaviour is the source of the following dependency conflicts.\n",
+ "numba 0.56.4 requires numpy\u003c1.24,\u003e=1.18, but you have numpy 1.25.1 which is incompatible.\n",
+ "tensorflow 2.13.0 requires numpy\u003c=1.24.3,\u003e=1.22, but you have numpy 1.25.1 which is incompatible.\u001b[0m\u001b[31m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m6.5/6.5 MB\u001b[0m \u001b[31m48.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m17.3/17.3 MB\u001b[0m \u001b[31m14.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h\u001b[31mERROR: pip's dependency resolver does not currently take into account all the packages that are installed. This behaviour is the source of the following dependency conflicts.\n",
+ "numba 0.56.4 requires numpy\u003c1.24,\u003e=1.18, but you have numpy 1.24.3 which is incompatible.\u001b[0m\u001b[31m\n",
+ "\u001b[0m"
+ ]
+ }
+ ],
+ "source": [
+ "!pip install -q tensorflow==2.13.0\n",
+ "!pip install -q -U tensorflow_datasets\n",
+ "!pip install -q --user -r models/official/requirements.txt\n",
+ "!pip install -q tensorflow-text==2.13.0"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "SJ50D6KvW6cK"
+ },
+ "source": [
+ "**Note**: Please restart the runtime once libraries are installed"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "S2qSVWakysDC"
+ },
+ "source": [
+ "## Please set the Python path with `os.environ` for models directory"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "vnaE50DGAg5H"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "os.environ['PYTHONPATH'] += \":/content/models\"\n",
+ "\n",
+ "import sys\n",
+ "sys.path.append(\"/content/models\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "5BPkQZBCzMSm"
+ },
+ "source": [
+ "## Import necessary libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 16820,
+ "status": "ok",
+ "timestamp": 1689010089176,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "moOdkRC01sdZ",
+ "outputId": "ec50aab3-57c3-4b9e-dd9f-138338111e98"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "2.13.0\n"
+ ]
+ }
+ ],
+ "source": [
+ "import os\n",
+ "import pprint\n",
+ "import tensorflow as tf\n",
+ "import tensorflow_datasets as tfds\n",
+ "\n",
+ "from IPython import display\n",
+ "from official.core import task_factory\n",
+ "from official.core import train_lib\n",
+ "from official.core import train_utils\n",
+ "from official.projects.perceiver.tasks import sentence_prediction\n",
+ "from official.projects.perceiver.configs import perceiver as exp_cfg\n",
+ "from official.nlp.modeling.layers import FastWordpieceBertTokenizer\n",
+ "from official.nlp.modeling.layers import BertPackInputs\n",
+ "\n",
+ "\n",
+ "pp = pprint.PrettyPrinter(indent=4) # Set Pretty Print Indentation\n",
+ "print(tf.__version__) # Check the version of tensorflow used"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZZ4ZNn3ezW0c"
+ },
+ "source": [
+ "## Download `glue/mrpc` dataset."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 931,
+ "referenced_widgets": [
+ "6c4cee474109471886f7dff896ff6d33",
+ "a8e2948404bf4c54a7c3226e6d304dfa",
+ "e364594b9e2446bab3c4768fb390a74b",
+ "b23a6acacbbf4267b48a59db16e7c4e6",
+ "eb67f79b4670428a9a350bc65a509907",
+ "1a13e1780d5d4679abf7b6c818ef9479",
+ "0f0adf4fa05745eebf6e5a968119cdea",
+ "638768d270a24fb9850d7ce8a929ec59",
+ "37bc2f8195c44cd2b024c5e3135ff390",
+ "644a43e4d9e54fd3b93fdce77434c102",
+ "14b3d59bfd394eb88c898cd0ee61ff53",
+ "9b2125597d1c4f48b951ea64df65aaaf",
+ "649a46c0be87434fbc8be9f18e68d41a",
+ "12d14e22b9e641efba0adbf41ad6438d",
+ "523c5f3f84d5423c8f71235ddfc8d5e1",
+ "93fc8dc08be04e10aa1d888e819370d9",
+ "bf689daf2c8e45068470aebd95336a3a",
+ "a7e39eb98b3548ed874effd391f865cd",
+ "d95fbc1de8354300b15afd398c79e336",
+ "21648c51936d4bdf9ecd068db0b0f4c9",
+ "d41ee155cbb0450980c04d22ac8c929e",
+ "5876534149a64daaba4b897bff309963",
+ "1377ec55510c471f957c5be3042ef36c",
+ "b512af1a5c394753a730dd607333424a",
+ "2d9b2bcd183e4a84b592f10e4b5b5661",
+ "af45c55d868841118adfea916a722bee",
+ "41bade4b10744d25bd5b40b6792b05f5",
+ "d3e660d14c394f49868930d71a95384c",
+ "b2c10377d6e34c76a9b89e41fb897f63",
+ "5c6e53e251b04cb7b8a92f32108f9a02",
+ "5e7cf7f3dc6e4335b58322fa7264ae97",
+ "e7fd108139b3439786ac18511f2247de",
+ "0572513797a24275812f60488e034892",
+ "d8f3a76c74374d4ca022aba7b184714a",
+ "aeab4fa4eead43c897db9ecb90e1dfb3",
+ "466f51b0011d458baca0e9561c6204c6",
+ "0a934c9980ed49e583a3f1a52c77695a",
+ "e69d91b41ade449bb4a972b84997924a",
+ "773ada8d21b34faf9999e358cd2dc372",
+ "9d9a9bd520684f73ad60a927efaa40c0",
+ "bd4d39f87d4c432398cf8d5ac43058d5",
+ "40fd6b4a0c3b4e959ecef49478eabc32",
+ "6d38ee13f6ab4cf8893b718d324301f4",
+ "b56e6022dbc648629c65d596e69f12af",
+ "58dad5720b0f404fb3b898dba90055ad",
+ "4a1ebf7b14c44dae9345ca78b3a274b0",
+ "200d6fb5129544cebc8f396874082879",
+ "c84bc3497fdc4d36b4a15ad9e48c6875",
+ "1e096cb57ae64aeb8d5ffc4c7522db29",
+ "5a734cfa349d4a90848b19935f475c21",
+ "f10ea5e8e2a0403ea5eca98da3c6354e",
+ "d74ab8fa5be449e9a05068307f754ac8",
+ "68f3b53a76e041cea3198ba56585e4d1",
+ "5154599b265e42b2bfb1f54d9b665451",
+ "2392d845775a48488d264dd0201b9063",
+ "524dc0d66c4047c6a9d3304cc0a87f33",
+ "228e8ce9e801480aae3afa14d2c4ffda",
+ "692817128ce646dc83879a24395107f1",
+ "8adde4af9df84c4893367261e3e57112",
+ "b976b945c1694c1baeef4aaedcdb965e",
+ "fa0162c12521438698ae678f6e63380e",
+ "0cd494d8cbfd49fe825799a0c356a2d4",
+ "ca8375042d1f46dabf50718e75faeb44",
+ "04911865ef3a4e859308087a8c66e58a",
+ "9dafefc83e444c91957ed6a9b2417c40",
+ "5e40195e05924ed0a57ff7547720d654",
+ "10a08d86c6c049c7b4b2231db6207526",
+ "8c6a79d8c9ec49649204a4845043bf59",
+ "0c0213ec610342c483655a43caab35df",
+ "2d37688bd1cd4a729245dfb78a15b484",
+ "fdec6c12ccae4f5c8595577657271a97",
+ "2b0a72fa0ba2408c8f9cd22011513377",
+ "07aa444924644d5fa4258dff62d47c18",
+ "6330bb99a63942919f485e56ce93df21",
+ "befb207161d64e4ba48e7c62295eb73c",
+ "a1a243e0250b4d9a84b1a9f7c1dfb29b",
+ "6b6e187de14b4833a69ed92fc2e51feb",
+ "bc072d37e96a4e1d80817a3b9eefffa5",
+ "2044e66365cd4beb9f33ebe873114ccd",
+ "8caebb7b5a9a422288fb8f3e93d25fd8",
+ "482f773e56244fccab80ecdd83dfe177",
+ "053df80c139441fbbe46f34d93ef9961",
+ "49e0e2bac1674e27b7989b0b5640223f",
+ "e8bd8359b444412386d547807846e24b",
+ "64afa7c458604b3d9a3a768e49b5d5ef",
+ "b6b80c5a9618457dbc665bcb46f2cd74",
+ "fcfbf00603934e1493583ffd47483f30",
+ "cb1b27446b024b16b3bcb81109f20e18",
+ "e3caa87a46f948038953eb53ee9eb49c",
+ "b5070d07f199417eb5efa1c71d88f5a3",
+ "cf8dd28cdcbf4bedba5556d808b5416e",
+ "fc22d54cc8364265ae2a3b77fde7dfe1",
+ "66f22cd19edb4f02a9a6c7b28eb9da59",
+ "9f8d4c9b27b14a54aae171509c80b323",
+ "0389b5e7bd2f4ad9a040f9a4d48af95c",
+ "d476f68cac2c4c84a06ae7d9112f1ed5",
+ "ded3625ff4454025a41bb449d3b59629",
+ "271ab31a75424942a0f39e02885d80a5",
+ "5139fd5ddd5141aa8ecfc436813a1c60"
+ ]
+ },
+ "executionInfo": {
+ "elapsed": 5649,
+ "status": "ok",
+ "timestamp": 1689010130551,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "izlrDcj62Hlh",
+ "outputId": "c4009825-4f49-41a3-d774-29166d1eb352"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Downloading and preparing dataset 1.43 MiB (download: 1.43 MiB, generated: 1.74 MiB, total: 3.17 MiB) to /root/tensorflow_datasets/glue/mrpc/2.0.0...\n"
+ ]
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "6c4cee474109471886f7dff896ff6d33",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Dl Completed...: 0 url [00:00, ? url/s]"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "9b2125597d1c4f48b951ea64df65aaaf",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Dl Size...: 0 MiB [00:00, ? MiB/s]"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "1377ec55510c471f957c5be3042ef36c",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Generating splits...: 0%| | 0/3 [00:00\u003c?, ? splits/s]"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "d8f3a76c74374d4ca022aba7b184714a",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Generating train examples...: 0%| | 0/3668 [00:00\u003c?, ? examples/s]"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "58dad5720b0f404fb3b898dba90055ad",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Shuffling /root/tensorflow_datasets/glue/mrpc/2.0.0.incompleteGFRJBN/glue-train.tfrecord*...: 0%| |…"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "524dc0d66c4047c6a9d3304cc0a87f33",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Generating validation examples...: 0%| | 0/408 [00:00\u003c?, ? examples/s]"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "10a08d86c6c049c7b4b2231db6207526",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Shuffling /root/tensorflow_datasets/glue/mrpc/2.0.0.incompleteGFRJBN/glue-validation.tfrecord*...: 0%| …"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "bc072d37e96a4e1d80817a3b9eefffa5",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Generating test examples...: 0%| | 0/1725 [00:00\u003c?, ? examples/s]"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "data": {
+ "application/vnd.jupyter.widget-view+json": {
+ "model_id": "e3caa87a46f948038953eb53ee9eb49c",
+ "version_major": 2,
+ "version_minor": 0
+ },
+ "text/plain": [
+ "Shuffling /root/tensorflow_datasets/glue/mrpc/2.0.0.incompleteGFRJBN/glue-test.tfrecord*...: 0%| | …"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Dataset glue downloaded and prepared to /root/tensorflow_datasets/glue/mrpc/2.0.0. Subsequent calls will reuse this data.\n"
+ ]
+ },
+ {
+ "data": {
+ "text/plain": [
+ "tfds.core.DatasetInfo(\n",
+ " name='glue',\n",
+ " full_name='glue/mrpc/2.0.0',\n",
+ " description=\"\"\"\n",
+ " GLUE, the General Language Understanding Evaluation benchmark\n",
+ " (https://gluebenchmark.com/) is a collection of resources for training,\n",
+ " evaluating, and analyzing natural language understanding systems.\n",
+ " \"\"\",\n",
+ " config_description=\"\"\"\n",
+ " The Microsoft Research Paraphrase Corpus (Dolan \u0026 Brockett, 2005) is a corpus of\n",
+ " sentence pairs automatically extracted from online news sources, with human annotations\n",
+ " for whether the sentences in the pair are semantically equivalent.\n",
+ " \"\"\",\n",
+ " homepage='https://www.microsoft.com/en-us/download/details.aspx?id=52398',\n",
+ " data_path=PosixGPath('/tmp/tmpyoq7f3i8tfds'),\n",
+ " file_format=tfrecord,\n",
+ " download_size=1.43 MiB,\n",
+ " dataset_size=1.74 MiB,\n",
+ " features=FeaturesDict({\n",
+ " 'idx': int32,\n",
+ " 'label': ClassLabel(shape=(), dtype=int64, num_classes=2),\n",
+ " 'sentence1': Text(shape=(), dtype=string),\n",
+ " 'sentence2': Text(shape=(), dtype=string),\n",
+ " }),\n",
+ " supervised_keys=None,\n",
+ " disable_shuffling=False,\n",
+ " splits={\n",
+ " 'test': \u003cSplitInfo num_examples=1725, num_shards=1\u003e,\n",
+ " 'train': \u003cSplitInfo num_examples=3668, num_shards=1\u003e,\n",
+ " 'validation': \u003cSplitInfo num_examples=408, num_shards=1\u003e,\n",
+ " },\n",
+ " citation=\"\"\"@inproceedings{dolan2005automatically,\n",
+ " title={Automatically constructing a corpus of sentential paraphrases},\n",
+ " author={Dolan, William B and Brockett, Chris},\n",
+ " booktitle={Proceedings of the Third International Workshop on Paraphrasing (IWP2005)},\n",
+ " year={2005}\n",
+ " }\n",
+ " @inproceedings{wang2019glue,\n",
+ " title={{GLUE}: A Multi-Task Benchmark and Analysis Platform for Natural Language Understanding},\n",
+ " author={Wang, Alex and Singh, Amanpreet and Michael, Julian and Hill, Felix and Levy, Omer and Bowman, Samuel R.},\n",
+ " note={In the Proceedings of ICLR.},\n",
+ " year={2019}\n",
+ " }\n",
+ " \n",
+ " Note that each GLUE dataset has its own citation. Please see the source to see\n",
+ " the correct citation for each contained dataset.\"\"\",\n",
+ ")"
+ ]
+ },
+ "execution_count": 3,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "tfds_name = 'glue/mrpc'\n",
+ "ds,ds_info = tfds.load(tfds_name,\n",
+ " with_info=True)\n",
+ "ds_info"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZXjm5M2Tzey6"
+ },
+ "source": [
+ "## Download bert base checkpoint for vocab file"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 10341,
+ "status": "ok",
+ "timestamp": 1689010145419,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "xR3t8TY-voKQ",
+ "outputId": "945fac79-c397-4157-d33b-8adc6a9edb18"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "--2023-07-10 17:28:55-- https://storage.googleapis.com/tf_model_garden/nlp/bert/v3/uncased_L-12_H-768_A-12.tar.gz\n",
+ "Resolving storage.googleapis.com (storage.googleapis.com)... 74.125.196.128, 173.194.215.128, 173.194.216.128, ...\n",
+ "Connecting to storage.googleapis.com (storage.googleapis.com)|74.125.196.128|:443... connected.\n",
+ "HTTP request sent, awaiting response... 200 OK\n",
+ "Length: 405351325 (387M) [application/octet-stream]\n",
+ "Saving to: ‘./uncased_L-12_H-768_A-12.tar.gz’\n",
+ "\n",
+ "./uncased_L-12_H-76 100%[===================\u003e] 386.57M 120MB/s in 3.2s \n",
+ "\n",
+ "2023-07-10 17:28:58 (120 MB/s) - ‘./uncased_L-12_H-768_A-12.tar.gz’ saved [405351325/405351325]\n",
+ "\n",
+ "uncased_L-12_H-768_A-12/\n",
+ "uncased_L-12_H-768_A-12/vocab.txt\n",
+ "uncased_L-12_H-768_A-12/bert_model.ckpt.index\n",
+ "uncased_L-12_H-768_A-12/bert_model.ckpt.data-00000-of-00001\n",
+ "uncased_L-12_H-768_A-12/params.yaml\n",
+ "uncased_L-12_H-768_A-12/bert_config.json\n"
+ ]
+ }
+ ],
+ "source": [
+ "!wget https://storage.googleapis.com/tf_model_garden/nlp/bert/v3/uncased_L-12_H-768_A-12.tar.gz -O ./uncased_L-12_H-768_A-12.tar.gz\n",
+ "!tar -zxvf ./uncased_L-12_H-768_A-12.tar.gz -C ./\n",
+ "!rm ./uncased_L-12_H-768_A-12.tar.gz"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ILW4p_gVP4pg"
+ },
+ "source": [
+ "## Configure the perceiver model for custom dataset training"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "cAlksY6nzdQY"
+ },
+ "outputs": [],
+ "source": [
+ "gs_folder_bert = \"./uncased_L-12_H-768_A-12\""
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "AzCEMctzz5N4"
+ },
+ "source": [
+ "### Load the registered configuration"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qS4-ZWxFBH7K"
+ },
+ "outputs": [],
+ "source": [
+ "exp_config = exp_cfg.exp_factory.get_exp_config('perceiver/word_piece_raw_sentence_prediction')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "4c5gFEvwPd7_"
+ },
+ "source": [
+ "### Change the parameters required to train the model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "m6-OhvAUyfXT"
+ },
+ "outputs": [],
+ "source": [
+ "BATCH_SIZE = 8\n",
+ "epochs = 5\n",
+ "vocab_file = './uncased_L-12_H-768_A-12/vocab.txt'\n",
+ "\n",
+ "\n",
+ "train_data_size = ds_info.splits['train'].num_examples\n",
+ "validation_data_size = ds_info.splits['validation'].num_examples\n",
+ "steps_per_epoch = int(train_data_size / BATCH_SIZE)\n",
+ "num_train_steps = steps_per_epoch * epochs\n",
+ "validation_steps = int(validation_data_size / BATCH_SIZE)\n",
+ "warmup_steps = int(0.1 * num_train_steps)\n",
+ "initial_learning_rate = 2e-5\n",
+ "\n",
+ "\n",
+ "exp_config.runtime.num_gpus = 1\n",
+ "exp_config.runtime.enable_xla = False\n",
+ "exp_config.runtime.mixed_precision_dtype = 'mixed_bfloat16'\n",
+ "exp_config.task.model.num_classes = 2\n",
+ "\n",
+ "exp_config.task.train_data.tfds_name = 'glue/mrpc'\n",
+ "exp_config.task.train_data.tfds_split = 'train'\n",
+ "exp_config.task.train_data.text_fields = ['sentence1', 'sentence2']\n",
+ "exp_config.task.train_data.global_batch_size = BATCH_SIZE\n",
+ "exp_config.task.train_data.lower_case = True\n",
+ "exp_config.task.train_data.tokenization = 'WordPiece'\n",
+ "exp_config.task.train_data.vocab_file = vocab_file\n",
+ "\n",
+ "exp_config.task.validation_data.tfds_name = 'glue/mrpc'\n",
+ "exp_config.task.validation_data.tfds_split = 'validation'\n",
+ "exp_config.task.validation_data.text_fields = ['sentence1', 'sentence2']\n",
+ "exp_config.task.validation_data.global_batch_size = BATCH_SIZE\n",
+ "exp_config.task.validation_data.lower_case = True\n",
+ "exp_config.task.validation_data.tokenization = 'WordPiece'\n",
+ "exp_config.task.validation_data.vocab_file = vocab_file\n",
+ "\n",
+ "exp_config.trainer.checkpoint_interval = steps_per_epoch\n",
+ "exp_config.trainer.optimizer_config.learning_rate.polynomial.initial_learning_rate = initial_learning_rate\n",
+ "exp_config.trainer.optimizer_config.learning_rate.polynomial.decay_steps = num_train_steps\n",
+ "exp_config.trainer.optimizer_config.warmup.polynomial.warmup_steps = warmup_steps\n",
+ "exp_config.trainer.steps_per_loop = steps_per_epoch\n",
+ "exp_config.trainer.summary_interval = steps_per_epoch\n",
+ "exp_config.trainer.train_steps = num_train_steps\n",
+ "exp_config.trainer.validation_interval = steps_per_epoch\n",
+ "exp_config.trainer.validation_steps = validation_steps"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LYP-_WCHnDHd"
+ },
+ "source": [
+ "### Detect the hardware"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 166,
+ "status": "ok",
+ "timestamp": 1689010193979,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "tjp6Ekql21_U",
+ "outputId": "72c0c030-0c1a-4edc-e023-82c98b4b0de2"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Running on single GPU /device:GPU:0\n",
+ "Number of accelerators: 1\n"
+ ]
+ }
+ ],
+ "source": [
+ "try:\n",
+ " tpu_resolver = tf.distribute.cluster_resolver.TPUClusterResolver() # TPU detection\n",
+ "except ValueError:\n",
+ " tpu_resolver = None\n",
+ " gpus = tf.config.experimental.list_logical_devices(\"GPU\")\n",
+ "\n",
+ "# Select appropriate distribution strategy\n",
+ "if tpu_resolver:\n",
+ " tf.config.experimental_connect_to_cluster(tpu_resolver)\n",
+ " tf.tpu.experimental.initialize_tpu_system(tpu_resolver)\n",
+ " distribution_strategy = tf.distribute.experimental.TPUStrategy(tpu_resolver)\n",
+ " print('Running on TPU ', tpu_resolver.cluster_spec().as_dict()['worker'])\n",
+ "elif len(gpus) \u003e 1:\n",
+ " distribution_strategy = tf.distribute.MirroredStrategy([gpu.name for gpu in gpus])\n",
+ " print('Running on multiple GPUs ', [gpu.name for gpu in gpus])\n",
+ "elif len(gpus) == 1:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on single GPU ', gpus[0].name)\n",
+ "else:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on CPU')\n",
+ "\n",
+ "print(\"Number of accelerators: \", distribution_strategy.num_replicas_in_sync)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "GxMDFCQhPjSW"
+ },
+ "source": [
+ "### Print the modified configuration."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 500
+ },
+ "executionInfo": {
+ "elapsed": 376,
+ "status": "ok",
+ "timestamp": 1689010197252,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "pep0gabI24_R",
+ "outputId": "eaadce22-2a36-4fbe-e089-ba2fdc547d26"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "{ 'runtime': { 'all_reduce_alg': None,\n",
+ " 'batchnorm_spatial_persistent': False,\n",
+ " 'dataset_num_private_threads': None,\n",
+ " 'default_shard_dim': -1,\n",
+ " 'distribution_strategy': 'mirrored',\n",
+ " 'enable_xla': False,\n",
+ " 'gpu_thread_mode': None,\n",
+ " 'loss_scale': None,\n",
+ " 'mixed_precision_dtype': 'mixed_bfloat16',\n",
+ " 'num_cores_per_replica': 1,\n",
+ " 'num_gpus': 1,\n",
+ " 'num_packs': 1,\n",
+ " 'per_gpu_thread_count': 0,\n",
+ " 'run_eagerly': False,\n",
+ " 'task_index': -1,\n",
+ " 'tpu': None,\n",
+ " 'tpu_enable_xla_dynamic_padder': None,\n",
+ " 'use_tpu_mp_strategy': False,\n",
+ " 'worker_hosts': None},\n",
+ " 'task': { 'allow_image_summary': False,\n",
+ " 'differential_privacy_config': None,\n",
+ " 'hub_module_url': '',\n",
+ " 'init_checkpoint': '',\n",
+ " 'init_cls_pooler': False,\n",
+ " 'metric_type': 'accuracy',\n",
+ " 'model': { 'decoder': { 'd_latents': 1280,\n",
+ " 'd_model': 768,\n",
+ " 'decoder': { 'num_heads': 8,\n",
+ " 'output_last_dim': 768,\n",
+ " 'qk_last_dim': 256,\n",
+ " 'use_query_residual': False,\n",
+ " 'v_last_dim': 768},\n",
+ " 'output_index_dim': 1,\n",
+ " 'position_encoding_intializer_stddev': 0.02,\n",
+ " 'z_index_dim': 256},\n",
+ " 'encoder': { 'd_latents': 1280,\n",
+ " 'd_model': 768,\n",
+ " 'embedding_initializer_stddev': 0.02,\n",
+ " 'embedding_width': 768,\n",
+ " 'encoder': { 'att_init_scale': 1.0,\n",
+ " 'cross_attention_num_heads': 8,\n",
+ " 'cross_attention_widening_factor': 1,\n",
+ " 'dense_init_scale': 1.0,\n",
+ " 'dropout_attn_prob': 0.0,\n",
+ " 'dropout_prob': 0.0,\n",
+ " 'norm_epsilon': 1e-05,\n",
+ " 'num_blocks': 1,\n",
+ " 'num_self_attends_per_block': 26,\n",
+ " 'qk_last_dim': 256,\n",
+ " 'self_attention_num_heads': 8,\n",
+ " 'self_attention_widening_factor': 1,\n",
+ " 'v_last_dim': 1280},\n",
+ " 'input_position_encoding_intializer_stddev': 0.02,\n",
+ " 'max_seq_len': 512,\n",
+ " 'vocab_size': 30522,\n",
+ " 'z_index_dim': 256,\n",
+ " 'z_pos_enc_init_scale': 0.02},\n",
+ " 'num_classes': 2,\n",
+ " 'use_encoder_pooler': False},\n",
+ " 'name': None,\n",
+ " 'train_data': { 'apply_tf_data_service_before_batching': False,\n",
+ " 'autotune_algorithm': None,\n",
+ " 'block_length': 1,\n",
+ " 'cache': False,\n",
+ " 'cycle_length': None,\n",
+ " 'deterministic': None,\n",
+ " 'drop_remainder': True,\n",
+ " 'enable_shared_tf_data_service_between_parallel_trainers': False,\n",
+ " 'enable_tf_data_service': False,\n",
+ " 'file_type': 'tfrecord',\n",
+ " 'global_batch_size': 8,\n",
+ " 'include_example_id': False,\n",
+ " 'input_path': '',\n",
+ " 'is_training': True,\n",
+ " 'label_field': 'label',\n",
+ " 'label_type': 'int',\n",
+ " 'lower_case': True,\n",
+ " 'prefetch_buffer_size': None,\n",
+ " 'preprocessing_hub_module_url': '',\n",
+ " 'seed': None,\n",
+ " 'seq_length': 128,\n",
+ " 'sharding': True,\n",
+ " 'shuffle_buffer_size': 100,\n",
+ " 'text_fields': ['sentence1', 'sentence2'],\n",
+ " 'tf_data_service_address': None,\n",
+ " 'tf_data_service_job_name': None,\n",
+ " 'tfds_as_supervised': False,\n",
+ " 'tfds_data_dir': '',\n",
+ " 'tfds_name': 'glue/mrpc',\n",
+ " 'tfds_skip_decoding_feature': '',\n",
+ " 'tfds_split': 'train',\n",
+ " 'tokenization': 'WordPiece',\n",
+ " 'trainer_id': None,\n",
+ " 'vocab_file': './uncased_L-12_H-768_A-12/vocab.txt'},\n",
+ " 'validation_data': { 'apply_tf_data_service_before_batching': False,\n",
+ " 'autotune_algorithm': None,\n",
+ " 'block_length': 1,\n",
+ " 'cache': False,\n",
+ " 'cycle_length': None,\n",
+ " 'deterministic': None,\n",
+ " 'drop_remainder': True,\n",
+ " 'enable_shared_tf_data_service_between_parallel_trainers': False,\n",
+ " 'enable_tf_data_service': False,\n",
+ " 'file_type': 'tfrecord',\n",
+ " 'global_batch_size': 8,\n",
+ " 'include_example_id': False,\n",
+ " 'input_path': '',\n",
+ " 'is_training': True,\n",
+ " 'label_field': 'label',\n",
+ " 'label_type': 'int',\n",
+ " 'lower_case': True,\n",
+ " 'prefetch_buffer_size': None,\n",
+ " 'preprocessing_hub_module_url': '',\n",
+ " 'seed': None,\n",
+ " 'seq_length': 128,\n",
+ " 'sharding': True,\n",
+ " 'shuffle_buffer_size': 100,\n",
+ " 'text_fields': [ 'sentence1',\n",
+ " 'sentence2'],\n",
+ " 'tf_data_service_address': None,\n",
+ " 'tf_data_service_job_name': None,\n",
+ " 'tfds_as_supervised': False,\n",
+ " 'tfds_data_dir': '',\n",
+ " 'tfds_name': 'glue/mrpc',\n",
+ " 'tfds_skip_decoding_feature': '',\n",
+ " 'tfds_split': 'validation',\n",
+ " 'tokenization': 'WordPiece',\n",
+ " 'trainer_id': None,\n",
+ " 'vocab_file': './uncased_L-12_H-768_A-12/vocab.txt'}},\n",
+ " 'trainer': { 'allow_tpu_summary': False,\n",
+ " 'best_checkpoint_eval_metric': '',\n",
+ " 'best_checkpoint_export_subdir': '',\n",
+ " 'best_checkpoint_metric_comp': 'higher',\n",
+ " 'checkpoint_interval': 458,\n",
+ " 'continuous_eval_timeout': 3600,\n",
+ " 'eval_tf_function': True,\n",
+ " 'eval_tf_while_loop': False,\n",
+ " 'loss_upper_bound': 1000000.0,\n",
+ " 'max_to_keep': 5,\n",
+ " 'optimizer_config': { 'ema': None,\n",
+ " 'learning_rate': { 'polynomial': { 'cycle': False,\n",
+ " 'decay_steps': 2290,\n",
+ " 'end_learning_rate': 0.0,\n",
+ " 'initial_learning_rate': 2e-05,\n",
+ " 'name': 'PolynomialDecay',\n",
+ " 'offset': 0,\n",
+ " 'power': 1.0},\n",
+ " 'type': 'polynomial'},\n",
+ " 'optimizer': { 'lamb': { 'beta_1': 0.9,\n",
+ " 'beta_2': 0.999,\n",
+ " 'clipnorm': None,\n",
+ " 'clipvalue': None,\n",
+ " 'epsilon': 1e-06,\n",
+ " 'exclude_from_layer_adaptation': None,\n",
+ " 'exclude_from_weight_decay': [ 'LayerNorm',\n",
+ " 'layer_norm',\n",
+ " 'bias'],\n",
+ " 'global_clipnorm': None,\n",
+ " 'name': 'LAMB',\n",
+ " 'weight_decay_rate': 0.01},\n",
+ " 'type': 'lamb'},\n",
+ " 'warmup': { 'linear': { 'name': 'linear',\n",
+ " 'warmup_learning_rate': 0.0,\n",
+ " 'warmup_steps': 200},\n",
+ " 'type': 'linear'}},\n",
+ " 'preemption_on_demand_checkpoint': True,\n",
+ " 'recovery_begin_steps': 0,\n",
+ " 'recovery_max_trials': 0,\n",
+ " 'steps_per_loop': 458,\n",
+ " 'summary_interval': 458,\n",
+ " 'train_steps': 2290,\n",
+ " 'train_tf_function': True,\n",
+ " 'train_tf_while_loop': True,\n",
+ " 'validation_interval': 458,\n",
+ " 'validation_steps': 51,\n",
+ " 'validation_summary_subdir': 'validation'}}\n"
+ ]
+ },
+ {
+ "data": {
+ "application/javascript": [
+ "google.colab.output.setIframeHeight(\"500px\");"
+ ],
+ "text/plain": [
+ "\u003cIPython.core.display.Javascript object\u003e"
+ ]
+ },
+ "execution_count": 9,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "pp.pprint(exp_config.as_dict())\n",
+ "display.Javascript('google.colab.output.setIframeHeight(\"500px\");')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "aivSkEcMQI6c"
+ },
+ "source": [
+ "## Create the `Task` object (`tfm.core.base_task.Task`) from the `config_definitions.TaskConfig`.\n",
+ "\n",
+ "The `Task` object has all the methods necessary for building the dataset, building the model, and running training \u0026 evaluation. These methods are driven by `tfm.core.train_lib.run_experiment`."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ybBYjDU72754"
+ },
+ "outputs": [],
+ "source": [
+ "model_dir = './trained_model/'\n",
+ "\n",
+ "with distribution_strategy.scope():\n",
+ " task = task_factory.get_task(exp_config.task, logging_dir=model_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "uQvshzY-QoSi"
+ },
+ "source": [
+ "## Train and Evaluate the model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 2882433,
+ "status": "ok",
+ "timestamp": 1689013083414,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "6zS8rSCqgHBz",
+ "outputId": "05e314d0-fcca-4d3c-9d11-5822444f7b02"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "restoring or initializing model...\n",
+ "train | step: 0 | training until step 458...\n",
+ "train | step: 458 | steps/sec: 0.8 | output: \n",
+ " {'auc': 0.72914344,\n",
+ " 'cls_accuracy': 0.66348255,\n",
+ " 'learning_rate': 1.6e-05,\n",
+ " 'training_loss': 0.63586426}\n",
+ "saved checkpoint to ./trained_model/ckpt-458.\n",
+ " eval | step: 458 | running 51 steps of evaluation...\n",
+ " eval | step: 458 | steps/sec: 2.2 | eval time: 23.1 sec | output: \n",
+ " {'auc': 0.7645757,\n",
+ " 'cls_accuracy': 0.67401963,\n",
+ " 'steps_per_second': 2.2112783063481327,\n",
+ " 'validation_loss': 0.6187927}\n",
+ "train | step: 458 | training until step 916...\n",
+ "train | step: 916 | steps/sec: 0.8 | output: \n",
+ " {'auc': 0.8580948,\n",
+ " 'cls_accuracy': 0.7363537,\n",
+ " 'learning_rate': 1.2e-05,\n",
+ " 'training_loss': 0.5320095}\n",
+ "saved checkpoint to ./trained_model/ckpt-916.\n",
+ " eval | step: 916 | running 51 steps of evaluation...\n",
+ " eval | step: 916 | steps/sec: 2.9 | eval time: 17.8 sec | output: \n",
+ " {'auc': 0.800959,\n",
+ " 'cls_accuracy': 0.6960784,\n",
+ " 'steps_per_second': 2.861756377583562,\n",
+ " 'validation_loss': 0.5854072}\n",
+ "train | step: 916 | training until step 1374...\n",
+ "train | step: 1374 | steps/sec: 0.8 | output: \n",
+ " {'auc': 0.95597136,\n",
+ " 'cls_accuracy': 0.83406115,\n",
+ " 'learning_rate': 7.999999e-06,\n",
+ " 'training_loss': 0.38442203}\n",
+ "saved checkpoint to ./trained_model/ckpt-1374.\n",
+ " eval | step: 1374 | running 51 steps of evaluation...\n",
+ " eval | step: 1374 | steps/sec: 2.9 | eval time: 17.8 sec | output: \n",
+ " {'auc': 0.7954676,\n",
+ " 'cls_accuracy': 0.6495098,\n",
+ " 'steps_per_second': 2.861158673104126,\n",
+ " 'validation_loss': 0.6343756}\n",
+ "train | step: 1374 | training until step 1832...\n",
+ "train | step: 1832 | steps/sec: 0.8 | output: \n",
+ " {'auc': 0.9869786,\n",
+ " 'cls_accuracy': 0.91784936,\n",
+ " 'learning_rate': 4e-06,\n",
+ " 'training_loss': 0.26613}\n",
+ "saved checkpoint to ./trained_model/ckpt-1832.\n",
+ " eval | step: 1832 | running 51 steps of evaluation...\n",
+ " eval | step: 1832 | steps/sec: 2.8 | eval time: 18.1 sec | output: \n",
+ " {'auc': 0.7897566,\n",
+ " 'cls_accuracy': 0.65686274,\n",
+ " 'steps_per_second': 2.8241461717336493,\n",
+ " 'validation_loss': 0.66516936}\n",
+ "train | step: 1832 | training until step 2290...\n",
+ "train | step: 2290 | steps/sec: 0.8 | output: \n",
+ " {'auc': 0.9962825,\n",
+ " 'cls_accuracy': 0.9585153,\n",
+ " 'learning_rate': 0.0,\n",
+ " 'training_loss': 0.18938588}\n",
+ "saved checkpoint to ./trained_model/ckpt-2290.\n",
+ " eval | step: 2290 | running 51 steps of evaluation...\n",
+ " eval | step: 2290 | steps/sec: 2.4 | eval time: 21.0 sec | output: \n",
+ " {'auc': 0.7945166,\n",
+ " 'cls_accuracy': 0.6764706,\n",
+ " 'steps_per_second': 2.4306418804782743,\n",
+ " 'validation_loss': 0.6767317}\n"
+ ]
+ }
+ ],
+ "source": [
+ "model, eval_logs = train_lib.run_experiment(\n",
+ " distribution_strategy=distribution_strategy,\n",
+ " task=task,\n",
+ " mode='train_and_eval',\n",
+ " params=exp_config,\n",
+ " model_dir=model_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "weZwyXKYQsVA"
+ },
+ "source": [
+ "## Testing the trained model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZVZ33VldQ4vL"
+ },
+ "source": [
+ "### Helper functions for pre-processing test data"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "O-q-hYpQ6nQB"
+ },
+ "outputs": [],
+ "source": [
+ "tokenizer = FastWordpieceBertTokenizer(\n",
+ " vocab_file=vocab_file,\n",
+ " lower_case=exp_config.task.train_data.lower_case\n",
+ ")\n",
+ "\n",
+ "packer = BertPackInputs(\n",
+ " seq_length=exp_config.task.train_data.seq_length,\n",
+ " special_tokens_dict=tokenizer.get_special_tokens_dict()\n",
+ ")\n",
+ "\n",
+ "\n",
+ "class BertInputProcessor(tf.keras.layers.Layer):\n",
+ " def __init__(self, tokenizer, packer):\n",
+ " super().__init__()\n",
+ " self.tokenizer = tokenizer\n",
+ " self.packer = packer\n",
+ "\n",
+ " def call(self, inputs):\n",
+ " tok1 = self.tokenizer(inputs['sentence1'])\n",
+ " tok2 = self.tokenizer(inputs['sentence2'])\n",
+ "\n",
+ " packed = self.packer([tok1, tok2])\n",
+ "\n",
+ " if 'label' in inputs:\n",
+ " return packed, inputs['label']\n",
+ " else:\n",
+ " return packed"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Xu4OqanvROFz"
+ },
+ "source": [
+ "### Pre-process test data"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "TYsP7Soq7off"
+ },
+ "outputs": [],
+ "source": [
+ "bert_inputs_processor = BertInputProcessor(\n",
+ " tokenizer=tokenizer,\n",
+ " packer=packer\n",
+ ")\n",
+ "test_ds = ds['test'].batch(\n",
+ " 1).map(bert_inputs_processor)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "w_I68G51RSHO"
+ },
+ "source": [
+ "### Get the predictions"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 10328,
+ "status": "ok",
+ "timestamp": 1689013095845,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 420
+ },
+ "id": "zDABiqAn29oi",
+ "outputId": "2ef6b773-1302-4659-e6bf-45552e4fae7f"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Sentence 1:[b'Shares in BA were down 1.5 percent at 168 pence by 1420 GMT , off a low of 164p , in a slightly stronger overall London market .']\n",
+ "Sentence 2:[b'Shares in BA were down three percent at 165-1 / 4 pence by 0933 GMT , off a low of 164 pence , in a stronger market .']\n",
+ "1/1 [==============================] - 4s 4s/step\n",
+ "Prediction: 1\n",
+ "Sentence 1:[b'The South Korean Agriculture and Forestry Ministry also said it would throw out or send back all Canadian beef currently in store .']\n",
+ "Sentence 2:[b'The South Korean Agriculture and Forestry Ministry said it would scrap or return all Canadian beef in store .']\n",
+ "1/1 [==============================] - 0s 103ms/step\n",
+ "Prediction: 1\n",
+ "Sentence 1:[b'\" New Yorkers didn \\'t embrace these units like they could have , \" said Matthew Daus , chairman of the commission .']\n",
+ "Sentence 2:[b'\" New Yorkers didn \\'t embrace these units like they could have , \" Matthew W. Daus , the commission \\'s chairman , said yesterday .']\n",
+ "1/1 [==============================] - 0s 95ms/step\n",
+ "Prediction: 1\n",
+ "Sentence 1:[b'\" I really liked him and I still do , \" Cohen Alon told the Herald yesterday .']\n",
+ "Sentence 2:[b'And I really liked him , and I still do .']\n",
+ "1/1 [==============================] - 0s 96ms/step\n",
+ "Prediction: 0\n",
+ "Sentence 1:[b'Tight media controls and the remote location of the northern village where the fighting broke out made it impossible to confirm what happened .']\n",
+ "Sentence 2:[b'Tight media controls and the remote location of the clash made it impossible to confirm what happened .']\n",
+ "1/1 [==============================] - 0s 94ms/step\n",
+ "Prediction: 1\n",
+ "Sentence 1:[b'She had been critically ill after May 7 surgery to replace a heart valve .']\n",
+ "Sentence 2:[b'She had been critically ill since having surgery at Baptist Hospital on May 7 to replace a heart valve .']\n",
+ "1/1 [==============================] - 0s 97ms/step\n",
+ "Prediction: 1\n",
+ "Sentence 1:[b'\" In fact , I was physically sick several times at this stage , because he looked so desperate , \" she said .']\n",
+ "Sentence 2:[b'Speaking about the day before he was found dead , she said : \" I was physically sick several times at this stage because he looked so desperate .']\n",
+ "1/1 [==============================] - 0s 101ms/step\n",
+ "Prediction: 0\n",
+ "Sentence 1:[b'The technology-laced Nasdaq Composite Index added 11.98 points , or 0.72 percent , to 1,680.42 .']\n",
+ "Sentence 2:[b\"The broader Standard \u0026 Poor 's 500 Index .SPX was off 1.07 points , or 0.11 percent , at 1,010.59 .\"]\n",
+ "1/1 [==============================] - 0s 99ms/step\n",
+ "Prediction: 0\n"
+ ]
+ }
+ ],
+ "source": [
+ "for record in ds['test'].batch(1).take(8):\n",
+ " print(f\"Sentence 1:{record['sentence1'].numpy()}\")\n",
+ " print(f\"Sentence 2:{record['sentence2'].numpy()}\")\n",
+ " processed_rec = bert_inputs_processor(record)\n",
+ " prediction = tf.argmax(\n",
+ " model.predict(processed_rec[0]),\n",
+ " axis=1)\n",
+ " print(f\"Prediction: {prediction[0]}\")"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "last_runtime": {
+ "build_target": "//learning/grp/tools/ml_python:ml_notebook",
+ "kind": "private"
+ },
+ "provenance": [
+ {
+ "file_id": "1ODIzioYI5DjPOT7-4EbfalaDhD0WXCHE",
+ "timestamp": 1689019241476
+ }
+ ],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ },
+ "widgets": {
+ "application/vnd.jupyter.widget-state+json": {
+ "0389b5e7bd2f4ad9a040f9a4d48af95c": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "04911865ef3a4e859308087a8c66e58a": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "053df80c139441fbbe46f34d93ef9961": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": "hidden",
+ "width": null
+ }
+ },
+ "0572513797a24275812f60488e034892": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "07aa444924644d5fa4258dff62d47c18": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "0a934c9980ed49e583a3f1a52c77695a": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_6d38ee13f6ab4cf8893b718d324301f4",
+ "placeholder": "",
+ "style": "IPY_MODEL_b56e6022dbc648629c65d596e69f12af",
+ "value": " 3452/3668 [00:01\u0026lt;00:00, 3686.58 examples/s]"
+ }
+ },
+ "0c0213ec610342c483655a43caab35df": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_6330bb99a63942919f485e56ce93df21",
+ "max": 408,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_befb207161d64e4ba48e7c62295eb73c",
+ "value": 408
+ }
+ },
+ "0cd494d8cbfd49fe825799a0c356a2d4": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "0f0adf4fa05745eebf6e5a968119cdea": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "10a08d86c6c049c7b4b2231db6207526": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_8c6a79d8c9ec49649204a4845043bf59",
+ "IPY_MODEL_0c0213ec610342c483655a43caab35df",
+ "IPY_MODEL_2d37688bd1cd4a729245dfb78a15b484"
+ ],
+ "layout": "IPY_MODEL_fdec6c12ccae4f5c8595577657271a97"
+ }
+ },
+ "12d14e22b9e641efba0adbf41ad6438d": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "success",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_d95fbc1de8354300b15afd398c79e336",
+ "max": 1,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_21648c51936d4bdf9ecd068db0b0f4c9",
+ "value": 0
+ }
+ },
+ "1377ec55510c471f957c5be3042ef36c": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_b512af1a5c394753a730dd607333424a",
+ "IPY_MODEL_2d9b2bcd183e4a84b592f10e4b5b5661",
+ "IPY_MODEL_af45c55d868841118adfea916a722bee"
+ ],
+ "layout": "IPY_MODEL_41bade4b10744d25bd5b40b6792b05f5"
+ }
+ },
+ "14b3d59bfd394eb88c898cd0ee61ff53": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "1a13e1780d5d4679abf7b6c818ef9479": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "1e096cb57ae64aeb8d5ffc4c7522db29": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": "hidden",
+ "width": null
+ }
+ },
+ "200d6fb5129544cebc8f396874082879": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_d74ab8fa5be449e9a05068307f754ac8",
+ "max": 3668,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_68f3b53a76e041cea3198ba56585e4d1",
+ "value": 3668
+ }
+ },
+ "2044e66365cd4beb9f33ebe873114ccd": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_49e0e2bac1674e27b7989b0b5640223f",
+ "placeholder": "",
+ "style": "IPY_MODEL_e8bd8359b444412386d547807846e24b",
+ "value": "Generating test examples...: 86%"
+ }
+ },
+ "21648c51936d4bdf9ecd068db0b0f4c9": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "228e8ce9e801480aae3afa14d2c4ffda": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_fa0162c12521438698ae678f6e63380e",
+ "placeholder": "",
+ "style": "IPY_MODEL_0cd494d8cbfd49fe825799a0c356a2d4",
+ "value": "Generating validation examples...: 75%"
+ }
+ },
+ "2392d845775a48488d264dd0201b9063": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "271ab31a75424942a0f39e02885d80a5": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "2b0a72fa0ba2408c8f9cd22011513377": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "2d37688bd1cd4a729245dfb78a15b484": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_a1a243e0250b4d9a84b1a9f7c1dfb29b",
+ "placeholder": "",
+ "style": "IPY_MODEL_6b6e187de14b4833a69ed92fc2e51feb",
+ "value": " 0/408 [00:00\u0026lt;?, ? examples/s]"
+ }
+ },
+ "2d9b2bcd183e4a84b592f10e4b5b5661": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_5c6e53e251b04cb7b8a92f32108f9a02",
+ "max": 3,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_5e7cf7f3dc6e4335b58322fa7264ae97",
+ "value": 3
+ }
+ },
+ "37bc2f8195c44cd2b024c5e3135ff390": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "40fd6b4a0c3b4e959ecef49478eabc32": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "41bade4b10744d25bd5b40b6792b05f5": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": "hidden",
+ "width": null
+ }
+ },
+ "466f51b0011d458baca0e9561c6204c6": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_bd4d39f87d4c432398cf8d5ac43058d5",
+ "max": 3668,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_40fd6b4a0c3b4e959ecef49478eabc32",
+ "value": 3668
+ }
+ },
+ "482f773e56244fccab80ecdd83dfe177": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_fcfbf00603934e1493583ffd47483f30",
+ "placeholder": "",
+ "style": "IPY_MODEL_cb1b27446b024b16b3bcb81109f20e18",
+ "value": " 1489/1725 [00:00\u0026lt;00:00, 3979.27 examples/s]"
+ }
+ },
+ "49e0e2bac1674e27b7989b0b5640223f": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "4a1ebf7b14c44dae9345ca78b3a274b0": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_5a734cfa349d4a90848b19935f475c21",
+ "placeholder": "",
+ "style": "IPY_MODEL_f10ea5e8e2a0403ea5eca98da3c6354e",
+ "value": "Shuffling /root/tensorflow_datasets/glue/mrpc/2.0.0.incompleteGFRJBN/glue-train.tfrecord*...: 0%"
+ }
+ },
+ "5139fd5ddd5141aa8ecfc436813a1c60": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "5154599b265e42b2bfb1f54d9b665451": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "523c5f3f84d5423c8f71235ddfc8d5e1": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_d41ee155cbb0450980c04d22ac8c929e",
+ "placeholder": "",
+ "style": "IPY_MODEL_5876534149a64daaba4b897bff309963",
+ "value": " 0/0 [00:00\u0026lt;?, ? MiB/s]"
+ }
+ },
+ "524dc0d66c4047c6a9d3304cc0a87f33": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_228e8ce9e801480aae3afa14d2c4ffda",
+ "IPY_MODEL_692817128ce646dc83879a24395107f1",
+ "IPY_MODEL_8adde4af9df84c4893367261e3e57112"
+ ],
+ "layout": "IPY_MODEL_b976b945c1694c1baeef4aaedcdb965e"
+ }
+ },
+ "5876534149a64daaba4b897bff309963": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "58dad5720b0f404fb3b898dba90055ad": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_4a1ebf7b14c44dae9345ca78b3a274b0",
+ "IPY_MODEL_200d6fb5129544cebc8f396874082879",
+ "IPY_MODEL_c84bc3497fdc4d36b4a15ad9e48c6875"
+ ],
+ "layout": "IPY_MODEL_1e096cb57ae64aeb8d5ffc4c7522db29"
+ }
+ },
+ "5a734cfa349d4a90848b19935f475c21": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "5c6e53e251b04cb7b8a92f32108f9a02": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "5e40195e05924ed0a57ff7547720d654": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "5e7cf7f3dc6e4335b58322fa7264ae97": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "6330bb99a63942919f485e56ce93df21": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "638768d270a24fb9850d7ce8a929ec59": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": "20px"
+ }
+ },
+ "644a43e4d9e54fd3b93fdce77434c102": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "649a46c0be87434fbc8be9f18e68d41a": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_bf689daf2c8e45068470aebd95336a3a",
+ "placeholder": "",
+ "style": "IPY_MODEL_a7e39eb98b3548ed874effd391f865cd",
+ "value": "Dl Size...: "
+ }
+ },
+ "64afa7c458604b3d9a3a768e49b5d5ef": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "66f22cd19edb4f02a9a6c7b28eb9da59": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": "hidden",
+ "width": null
+ }
+ },
+ "68f3b53a76e041cea3198ba56585e4d1": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "692817128ce646dc83879a24395107f1": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_ca8375042d1f46dabf50718e75faeb44",
+ "max": 408,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_04911865ef3a4e859308087a8c66e58a",
+ "value": 408
+ }
+ },
+ "6b6e187de14b4833a69ed92fc2e51feb": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "6c4cee474109471886f7dff896ff6d33": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_a8e2948404bf4c54a7c3226e6d304dfa",
+ "IPY_MODEL_e364594b9e2446bab3c4768fb390a74b",
+ "IPY_MODEL_b23a6acacbbf4267b48a59db16e7c4e6"
+ ],
+ "layout": "IPY_MODEL_eb67f79b4670428a9a350bc65a509907"
+ }
+ },
+ "6d38ee13f6ab4cf8893b718d324301f4": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "773ada8d21b34faf9999e358cd2dc372": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "8adde4af9df84c4893367261e3e57112": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_9dafefc83e444c91957ed6a9b2417c40",
+ "placeholder": "",
+ "style": "IPY_MODEL_5e40195e05924ed0a57ff7547720d654",
+ "value": " 305/408 [00:00\u0026lt;00:00, 1558.00 examples/s]"
+ }
+ },
+ "8c6a79d8c9ec49649204a4845043bf59": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_2b0a72fa0ba2408c8f9cd22011513377",
+ "placeholder": "",
+ "style": "IPY_MODEL_07aa444924644d5fa4258dff62d47c18",
+ "value": "Shuffling /root/tensorflow_datasets/glue/mrpc/2.0.0.incompleteGFRJBN/glue-validation.tfrecord*...: 0%"
+ }
+ },
+ "8caebb7b5a9a422288fb8f3e93d25fd8": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_64afa7c458604b3d9a3a768e49b5d5ef",
+ "max": 1725,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_b6b80c5a9618457dbc665bcb46f2cd74",
+ "value": 1725
+ }
+ },
+ "93fc8dc08be04e10aa1d888e819370d9": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "9b2125597d1c4f48b951ea64df65aaaf": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_649a46c0be87434fbc8be9f18e68d41a",
+ "IPY_MODEL_12d14e22b9e641efba0adbf41ad6438d",
+ "IPY_MODEL_523c5f3f84d5423c8f71235ddfc8d5e1"
+ ],
+ "layout": "IPY_MODEL_93fc8dc08be04e10aa1d888e819370d9"
+ }
+ },
+ "9d9a9bd520684f73ad60a927efaa40c0": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "9dafefc83e444c91957ed6a9b2417c40": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "9f8d4c9b27b14a54aae171509c80b323": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "a1a243e0250b4d9a84b1a9f7c1dfb29b": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "a7e39eb98b3548ed874effd391f865cd": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "a8e2948404bf4c54a7c3226e6d304dfa": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_1a13e1780d5d4679abf7b6c818ef9479",
+ "placeholder": "",
+ "style": "IPY_MODEL_0f0adf4fa05745eebf6e5a968119cdea",
+ "value": "Dl Completed...: 100%"
+ }
+ },
+ "aeab4fa4eead43c897db9ecb90e1dfb3": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_773ada8d21b34faf9999e358cd2dc372",
+ "placeholder": "",
+ "style": "IPY_MODEL_9d9a9bd520684f73ad60a927efaa40c0",
+ "value": "Generating train examples...: 94%"
+ }
+ },
+ "af45c55d868841118adfea916a722bee": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_e7fd108139b3439786ac18511f2247de",
+ "placeholder": "",
+ "style": "IPY_MODEL_0572513797a24275812f60488e034892",
+ "value": " 3/3 [00:01\u0026lt;00:00, 1.67 splits/s]"
+ }
+ },
+ "b23a6acacbbf4267b48a59db16e7c4e6": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_644a43e4d9e54fd3b93fdce77434c102",
+ "placeholder": "",
+ "style": "IPY_MODEL_14b3d59bfd394eb88c898cd0ee61ff53",
+ "value": " 3/3 [00:00\u0026lt;00:00, 7.70 url/s]"
+ }
+ },
+ "b2c10377d6e34c76a9b89e41fb897f63": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "b5070d07f199417eb5efa1c71d88f5a3": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_9f8d4c9b27b14a54aae171509c80b323",
+ "placeholder": "",
+ "style": "IPY_MODEL_0389b5e7bd2f4ad9a040f9a4d48af95c",
+ "value": "Shuffling /root/tensorflow_datasets/glue/mrpc/2.0.0.incompleteGFRJBN/glue-test.tfrecord*...: 0%"
+ }
+ },
+ "b512af1a5c394753a730dd607333424a": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_d3e660d14c394f49868930d71a95384c",
+ "placeholder": "",
+ "style": "IPY_MODEL_b2c10377d6e34c76a9b89e41fb897f63",
+ "value": "Generating splits...: 100%"
+ }
+ },
+ "b56e6022dbc648629c65d596e69f12af": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "b6b80c5a9618457dbc665bcb46f2cd74": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "b976b945c1694c1baeef4aaedcdb965e": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": "hidden",
+ "width": null
+ }
+ },
+ "bc072d37e96a4e1d80817a3b9eefffa5": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_2044e66365cd4beb9f33ebe873114ccd",
+ "IPY_MODEL_8caebb7b5a9a422288fb8f3e93d25fd8",
+ "IPY_MODEL_482f773e56244fccab80ecdd83dfe177"
+ ],
+ "layout": "IPY_MODEL_053df80c139441fbbe46f34d93ef9961"
+ }
+ },
+ "bd4d39f87d4c432398cf8d5ac43058d5": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "befb207161d64e4ba48e7c62295eb73c": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "bf689daf2c8e45068470aebd95336a3a": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "c84bc3497fdc4d36b4a15ad9e48c6875": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_5154599b265e42b2bfb1f54d9b665451",
+ "placeholder": "",
+ "style": "IPY_MODEL_2392d845775a48488d264dd0201b9063",
+ "value": " 0/3668 [00:00\u0026lt;?, ? examples/s]"
+ }
+ },
+ "ca8375042d1f46dabf50718e75faeb44": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "cb1b27446b024b16b3bcb81109f20e18": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "cf8dd28cdcbf4bedba5556d808b5416e": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_d476f68cac2c4c84a06ae7d9112f1ed5",
+ "max": 1725,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_ded3625ff4454025a41bb449d3b59629",
+ "value": 1725
+ }
+ },
+ "d3e660d14c394f49868930d71a95384c": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "d41ee155cbb0450980c04d22ac8c929e": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "d476f68cac2c4c84a06ae7d9112f1ed5": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "d74ab8fa5be449e9a05068307f754ac8": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "d8f3a76c74374d4ca022aba7b184714a": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_aeab4fa4eead43c897db9ecb90e1dfb3",
+ "IPY_MODEL_466f51b0011d458baca0e9561c6204c6",
+ "IPY_MODEL_0a934c9980ed49e583a3f1a52c77695a"
+ ],
+ "layout": "IPY_MODEL_e69d91b41ade449bb4a972b84997924a"
+ }
+ },
+ "d95fbc1de8354300b15afd398c79e336": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": "20px"
+ }
+ },
+ "ded3625ff4454025a41bb449d3b59629": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "ProgressStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "ProgressStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "bar_color": null,
+ "description_width": ""
+ }
+ },
+ "e364594b9e2446bab3c4768fb390a74b": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "FloatProgressModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "FloatProgressModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "ProgressView",
+ "bar_style": "success",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_638768d270a24fb9850d7ce8a929ec59",
+ "max": 1,
+ "min": 0,
+ "orientation": "horizontal",
+ "style": "IPY_MODEL_37bc2f8195c44cd2b024c5e3135ff390",
+ "value": 1
+ }
+ },
+ "e3caa87a46f948038953eb53ee9eb49c": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HBoxModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HBoxModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HBoxView",
+ "box_style": "",
+ "children": [
+ "IPY_MODEL_b5070d07f199417eb5efa1c71d88f5a3",
+ "IPY_MODEL_cf8dd28cdcbf4bedba5556d808b5416e",
+ "IPY_MODEL_fc22d54cc8364265ae2a3b77fde7dfe1"
+ ],
+ "layout": "IPY_MODEL_66f22cd19edb4f02a9a6c7b28eb9da59"
+ }
+ },
+ "e69d91b41ade449bb4a972b84997924a": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": "hidden",
+ "width": null
+ }
+ },
+ "e7fd108139b3439786ac18511f2247de": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "e8bd8359b444412386d547807846e24b": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "eb67f79b4670428a9a350bc65a509907": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "f10ea5e8e2a0403ea5eca98da3c6354e": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "DescriptionStyleModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "DescriptionStyleModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "StyleView",
+ "description_width": ""
+ }
+ },
+ "fa0162c12521438698ae678f6e63380e": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "fc22d54cc8364265ae2a3b77fde7dfe1": {
+ "model_module": "@jupyter-widgets/controls",
+ "model_module_version": "1.5.0",
+ "model_name": "HTMLModel",
+ "state": {
+ "_dom_classes": [],
+ "_model_module": "@jupyter-widgets/controls",
+ "_model_module_version": "1.5.0",
+ "_model_name": "HTMLModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/controls",
+ "_view_module_version": "1.5.0",
+ "_view_name": "HTMLView",
+ "description": "",
+ "description_tooltip": null,
+ "layout": "IPY_MODEL_271ab31a75424942a0f39e02885d80a5",
+ "placeholder": "",
+ "style": "IPY_MODEL_5139fd5ddd5141aa8ecfc436813a1c60",
+ "value": " 0/1725 [00:00\u0026lt;?, ? examples/s]"
+ }
+ },
+ "fcfbf00603934e1493583ffd47483f30": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": null,
+ "width": null
+ }
+ },
+ "fdec6c12ccae4f5c8595577657271a97": {
+ "model_module": "@jupyter-widgets/base",
+ "model_module_version": "1.2.0",
+ "model_name": "LayoutModel",
+ "state": {
+ "_model_module": "@jupyter-widgets/base",
+ "_model_module_version": "1.2.0",
+ "_model_name": "LayoutModel",
+ "_view_count": null,
+ "_view_module": "@jupyter-widgets/base",
+ "_view_module_version": "1.2.0",
+ "_view_name": "LayoutView",
+ "align_content": null,
+ "align_items": null,
+ "align_self": null,
+ "border": null,
+ "bottom": null,
+ "display": null,
+ "flex": null,
+ "flex_flow": null,
+ "grid_area": null,
+ "grid_auto_columns": null,
+ "grid_auto_flow": null,
+ "grid_auto_rows": null,
+ "grid_column": null,
+ "grid_gap": null,
+ "grid_row": null,
+ "grid_template_areas": null,
+ "grid_template_columns": null,
+ "grid_template_rows": null,
+ "height": null,
+ "justify_content": null,
+ "justify_items": null,
+ "left": null,
+ "margin": null,
+ "max_height": null,
+ "max_width": null,
+ "min_height": null,
+ "min_width": null,
+ "object_fit": null,
+ "object_position": null,
+ "order": null,
+ "overflow": null,
+ "overflow_x": null,
+ "overflow_y": null,
+ "padding": null,
+ "right": null,
+ "top": null,
+ "visibility": "hidden",
+ "width": null
+ }
+ }
+ }
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/perceiver/tasks/pretrain.py b/official/projects/perceiver/tasks/pretrain.py
new file mode 100644
index 00000000000..0a773d0210c
--- /dev/null
+++ b/official/projects/perceiver/tasks/pretrain.py
@@ -0,0 +1,61 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Task for perceiver wordpiece tokenized masked language model (MLM)."""
+
+import tensorflow as tf, tf_keras
+
+from official.core import task_factory
+from official.modeling import tf_utils
+from official.nlp.tasks import masked_lm
+from official.projects.perceiver.configs import encoders
+from official.projects.perceiver.configs import perceiver
+from official.projects.perceiver.modeling.layers import decoder
+from official.projects.perceiver.modeling.models import pretrainer
+from official.projects.perceiver.modeling.networks import positional_decoder
+
+
+@task_factory.register_task_cls(perceiver.PretrainConfig)
+class PretrainTask(masked_lm.MaskedLMTask):
+ """Task for masked language modeling for wordpiece tokenized perceiver."""
+
+ def build_model(self, params=None):
+ """Creates perceiver pretrainer model architecture.
+
+ Args:
+ params:
+ The task configuration instance, which can be any of dataclass,
+ ConfigDict, namedtuple, etc.
+ Returns:
+ A model instance.
+ """
+ config = params or self.task_config.model
+ sequence_encoder_cfg = config.encoder
+ encoder_network = encoders.build_encoder(sequence_encoder_cfg)
+ decoder_cfg = config.decoder
+ decoder_ = decoder.Decoder(decoder_cfg.decoder.as_dict())
+ mlm_decoder = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ output_index_dim=decoder_cfg.output_index_dim,
+ z_index_dim=decoder_cfg.z_index_dim,
+ d_latents=decoder_cfg.d_latents,
+ d_model=decoder_cfg.d_model,
+ position_encoding_intializer_stddev=decoder_cfg
+ .position_encoding_intializer_stddev)
+ return pretrainer.Pretrainer(
+ mlm_activation=tf_utils.get_activation(config.mlm_activation),
+ mlm_initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=config.mlm_initializer_range),
+ encoder=encoder_network,
+ decoder=mlm_decoder)
diff --git a/official/projects/perceiver/tasks/pretrain_test.py b/official/projects/perceiver/tasks/pretrain_test.py
new file mode 100644
index 00000000000..aa0b488f336
--- /dev/null
+++ b/official/projects/perceiver/tasks/pretrain_test.py
@@ -0,0 +1,132 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for official.nlp.tasks.masked_lm."""
+
+import tensorflow as tf, tf_keras
+import tensorflow_datasets as tfds
+
+from official.nlp.data import pretrain_dataloader
+from official.projects.perceiver.configs import perceiver
+from official.projects.perceiver.tasks import pretrain as tasks
+
+
+_NUM_EXAMPLES = 10
+
+
+def _gen_fn():
+ word_ids = tf.constant([1, 1], dtype=tf.int32)
+ mask = tf.constant([1, 1], dtype=tf.int32)
+ lm_mask = tf.constant([1, 1], dtype=tf.int32)
+ return {
+ 'file_name': 'test',
+ 'masked_lm_positions': lm_mask,
+ 'input_word_ids': word_ids,
+ 'input_mask': mask,
+ }
+
+
+def _as_dataset(self, *args, **kwargs):
+ del args
+ del kwargs
+ return tf.data.Dataset.from_generator(
+ lambda: (_gen_fn() for i in range(_NUM_EXAMPLES)),
+ output_types=self.info.features.dtype,
+ output_shapes=self.info.features.shape,
+ )
+
+
+def _fake_build_inputs(self, params, input_context=None): # pylint: disable=unused-argument
+ def dummy_data(_):
+ dummy_ids = tf.zeros((1, params.seq_length), dtype=tf.int32)
+ dummy_lm = tf.zeros((1, params.max_predictions_per_seq), dtype=tf.int32)
+ return dict(
+ input_word_ids=dummy_ids,
+ input_mask=dummy_ids,
+ masked_lm_positions=dummy_lm,
+ masked_lm_ids=dummy_lm,
+ masked_lm_weights=tf.cast(dummy_lm, dtype=tf.float32))
+
+ dataset = tf.data.Dataset.range(1)
+ dataset = dataset.repeat()
+ dataset = dataset.map(
+ dummy_data, num_parallel_calls=tf.data.experimental.AUTOTUNE)
+ return dataset
+
+
+class PretrainTaskTest(tf.test.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ tasks.PretrainTask.build_inputs = _fake_build_inputs
+
+ def test_task(self):
+ config = perceiver.PretrainConfig(
+ train_data=pretrain_dataloader.BertPretrainDataConfig(
+ input_path='dummy',
+ global_batch_size=512,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True),
+ validation_data=pretrain_dataloader.BertPretrainDataConfig(
+ input_path='dummy',
+ global_batch_size=512,
+ is_training=False,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True))
+ task = tasks.PretrainTask(config)
+
+ model = task.build_model()
+ metrics = task.build_metrics()
+ dataset = task.build_inputs(config.train_data)
+
+ iterator = iter(dataset)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
+ task.train_step(next(iterator), model, optimizer, metrics=metrics)
+ task.validation_step(next(iterator), model, metrics=metrics)
+
+ # Saves a checkpoint.
+ _ = tf.train.Checkpoint(model=model, **model.checkpoint_items)
+ # ckpt.save(config.init_checkpoint)
+ # TODO(b/222634115) fix ckpt.save
+ task.initialize(model)
+
+ def test_train_step(self):
+ config = perceiver.PretrainConfig(
+ train_data=pretrain_dataloader.BertPretrainDataConfig(
+ input_path='dummy',
+ global_batch_size=512,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True),
+ validation_data=pretrain_dataloader.BertPretrainDataConfig(
+ input_path='dummy',
+ global_batch_size=512,
+ is_training=False,
+ use_next_sentence_label=False,
+ use_v2_feature_names=True))
+
+ with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
+ task = tasks.PretrainTask(config)
+ model = task.build_model()
+ dataset = task.build_inputs(config.train_data)
+ metrics = task.build_metrics()
+
+ iterator = iter(dataset)
+ opt_cfg = perceiver._MLM_WORDPIECE_TRAINER.optimizer_config
+ optimizer = tasks.PretrainTask.create_optimizer(opt_cfg)
+ task.train_step(next(iterator), model, optimizer, metrics=metrics)
+
+# TODO(b/222634115) add test coverage.
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/perceiver/tasks/sentence_prediction.py b/official/projects/perceiver/tasks/sentence_prediction.py
new file mode 100644
index 00000000000..74591d2e64c
--- /dev/null
+++ b/official/projects/perceiver/tasks/sentence_prediction.py
@@ -0,0 +1,54 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Sentence prediction (classification) task."""
+
+from official.core import task_factory
+from official.nlp.tasks import sentence_prediction
+from official.projects.perceiver.configs import encoders
+from official.projects.perceiver.configs import perceiver
+from official.projects.perceiver.modeling.layers import decoder
+from official.projects.perceiver.modeling.models import classifier
+from official.projects.perceiver.modeling.networks import positional_decoder
+
+
+@task_factory.register_task_cls(perceiver.SentencePredictionConfig)
+class SentencePredictionTask(sentence_prediction.SentencePredictionTask):
+ """Task object for sentence_prediction.
+
+ Note: Making this similar to nlp.tasks.sentence_prediction.py to potentially
+ merge.
+ """
+
+ def build_model(self):
+ """Creates perceiver classification model architecture.
+
+ Returns:
+ A model instance.
+ """
+ encoder_network = encoders.build_encoder(self.task_config.model.encoder)
+ decoder_config = self.task_config.model.decoder
+ decoder_ = decoder.Decoder(decoder_config.decoder.as_dict())
+ classification_decoder = positional_decoder.PositionalDecoder(
+ decoder=decoder_,
+ d_model=decoder_config.d_model,
+ output_index_dim=decoder_config.output_index_dim,
+ z_index_dim=decoder_config.z_index_dim,
+ d_latents=decoder_config.d_latents,
+ position_encoding_intializer_stddev=decoder_config
+ .position_encoding_intializer_stddev)
+ return classifier.Classifier(
+ network=encoder_network,
+ decoder=classification_decoder,
+ num_classes=self.task_config.model.num_classes)
diff --git a/official/projects/perceiver/tasks/sentence_prediction_test.py b/official/projects/perceiver/tasks/sentence_prediction_test.py
new file mode 100644
index 00000000000..3f0c0d1f917
--- /dev/null
+++ b/official/projects/perceiver/tasks/sentence_prediction_test.py
@@ -0,0 +1,280 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for official.nlp.tasks.sentence_prediction."""
+import functools
+import os
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.nlp.data import sentence_prediction_dataloader
+from official.nlp.tasks import sentence_prediction
+from official.projects.perceiver.configs import perceiver
+from official.projects.perceiver.tasks import pretrain
+from official.projects.perceiver.tasks import sentence_prediction as perceiver_pred
+
+
+def _create_fake_dataset(output_path, seq_length, num_classes, num_examples):
+ """Creates a fake dataset.
+
+ Args:
+ output_path:
+ output path for the writer to serialize the dataset.
+ seq_length:
+ sequence length of the data.
+ num_classes:
+ Number of classes in the sentence prediction output. This is used to
+ determine if the label id feature should be for regression or
+ classification.
+ num_examples:
+ number of fake examples to create.
+ """
+
+ with tf.io.TFRecordWriter(output_path) as writer:
+ def create_int_feature(values):
+ return tf.train.Feature(
+ int64_list=tf.train.Int64List(value=np.ravel(values)))
+
+ def create_float_feature(values):
+ return tf.train.Feature(
+ float_list=tf.train.FloatList(value=np.ravel(values)))
+
+ for i in range(num_examples):
+ features = {}
+ input_ids = np.random.randint(100, size=(seq_length))
+ features["input_ids"] = create_int_feature(input_ids)
+ features["input_mask"] = create_int_feature(np.ones_like(input_ids))
+ features["segment_ids"] = create_int_feature(np.ones_like(input_ids))
+ features["segment_ids"] = create_int_feature(np.ones_like(input_ids))
+ features["example_id"] = create_int_feature([i])
+
+ if num_classes == 1:
+ features["label_ids"] = create_float_feature([np.random.random()])
+ else:
+ features["label_ids"] = create_int_feature(
+ [np.random.random_integers(0, num_classes - 1, size=())])
+
+ tf_example = tf.train.Example(
+ features=tf.train.Features(feature=features))
+ writer.write(tf_example.SerializeToString())
+
+
+class SentencePredictionTaskTest(tf.test.TestCase, parameterized.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ self._train_data_config = (
+ sentence_prediction_dataloader.SentencePredictionDataConfig(
+ input_path="dummy", seq_length=128, global_batch_size=1))
+
+ def get_model_config(self, num_classes):
+ return perceiver.ClassificationConfig(
+ num_classes=num_classes,
+ encoder=perceiver.SequenceEncoderConfig(
+ vocab_size=30_522,
+ encoder=perceiver.EncoderConfig(
+ num_self_attends_per_block=2)))
+
+ def _run_task(self, config):
+ task = perceiver_pred.SentencePredictionTask(config)
+ model = task.build_model()
+ metrics = task.build_metrics()
+
+ strategy = tf.distribute.get_strategy()
+ dataset = strategy.distribute_datasets_from_function(
+ functools.partial(task.build_inputs, config.train_data))
+
+ iterator = iter(dataset)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
+ task.train_step(next(iterator), model, optimizer, metrics=metrics)
+ # model.save(os.path.join(self.get_temp_dir(), "saved_model"))
+ # TODO(b/222634115) fix save
+ return task.validation_step(next(iterator), model, metrics=metrics)
+
+ def test_task(self):
+ # Saves a checkpoint.
+ pretrain_cfg = perceiver.PretrainerConfig(
+ encoder=perceiver.SequenceEncoderConfig(
+ vocab_size=30_522,
+ encoder=perceiver.EncoderConfig(
+ num_self_attends_per_block=2)))
+ pretrain_model = pretrain.PretrainTask(
+ None).build_model(pretrain_cfg)
+
+ # The model variables will be created after the forward call.
+ _ = pretrain_model(pretrain_model.inputs)
+ ckpt = tf.train.Checkpoint(
+ model=pretrain_model, **pretrain_model.checkpoint_items)
+ init_path = ckpt.save(self.get_temp_dir())
+
+ # Creates the task.
+ config = perceiver.SentencePredictionConfig(
+ init_checkpoint=init_path,
+ model=self.get_model_config(num_classes=2),
+ train_data=self._train_data_config)
+ task = perceiver_pred.SentencePredictionTask(config)
+ model = task.build_model()
+ metrics = task.build_metrics()
+ dataset = task.build_inputs(config.train_data)
+
+ iterator = iter(dataset)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
+ task.initialize(model)
+ task.train_step(next(iterator), model, optimizer, metrics=metrics)
+ task.validation_step(next(iterator), model, metrics=metrics)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name":
+ "regression",
+ "num_classes":
+ 1,
+ "expected_loss_predicate":
+ lambda loss: loss > 1.0,
+ "metric": tf_keras.metrics.MeanSquaredError,
+ },
+ {
+ "testcase_name":
+ "classification",
+ "num_classes":
+ 2,
+ "expected_loss_predicate":
+ lambda loss: loss < 1.0,
+ "metric": tf_keras.metrics.SparseCategoricalAccuracy
+ },
+ )
+ def test_metrics_and_losses(self, num_classes, expected_loss_predicate,
+ metric):
+ config = perceiver.SentencePredictionConfig(
+ init_checkpoint=self.get_temp_dir(),
+ model=self.get_model_config(num_classes),
+ train_data=self._train_data_config)
+ task = perceiver_pred.SentencePredictionTask(config)
+ model = task.build_model()
+ metrics = task.build_metrics()
+ self.assertIsInstance(metrics[0], metric)
+
+ dataset = task.build_inputs(config.train_data)
+ iterator = iter(dataset)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
+ task.train_step(next(iterator), model, optimizer, metrics=metrics)
+
+ logs = task.validation_step(next(iterator), model, metrics=metrics)
+ loss = logs["loss"].numpy()
+ self.assertTrue(expected_loss_predicate(loss))
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "matthews_corrcoef",
+ "num_classes": 2,
+ "metric_type": "matthews_corrcoef"
+ }, {
+ "testcase_name": "pearson_spearman_corr",
+ "num_classes": 1,
+ "metric_type": "pearson_spearman_corr"
+ })
+ def test_np_metrics(self, metric_type, num_classes):
+ config = perceiver.SentencePredictionConfig(
+ metric_type=metric_type,
+ init_checkpoint=self.get_temp_dir(),
+ model=self.get_model_config(num_classes),
+ train_data=self._train_data_config)
+ task = perceiver_pred.SentencePredictionTask(config)
+ model = task.build_model()
+ dataset = task.build_inputs(config.train_data)
+
+ iterator = iter(dataset)
+ strategy = tf.distribute.get_strategy()
+ distributed_outputs = strategy.run(
+ functools.partial(task.validation_step, model=model),
+ args=(next(iterator),))
+ outputs = tf.nest.map_structure(strategy.experimental_local_results,
+ distributed_outputs)
+ aggregated = task.aggregate_logs(step_outputs=outputs)
+ aggregated = task.aggregate_logs(state=aggregated, step_outputs=outputs)
+ self.assertIn(metric_type, task.reduce_aggregated_logs(aggregated))
+
+ def test_np_metrics_cola_partial_batch(self):
+ train_data_path = os.path.join(self.get_temp_dir(), "train.tf_record")
+ num_examples = 5
+ global_batch_size = 8
+ seq_length = 16
+ _create_fake_dataset(
+ train_data_path,
+ seq_length=seq_length,
+ num_classes=2,
+ num_examples=num_examples)
+
+ train_data_config = (
+ sentence_prediction_dataloader.SentencePredictionDataConfig(
+ input_path=train_data_path,
+ seq_length=seq_length,
+ is_training=True,
+ label_type="int",
+ global_batch_size=global_batch_size,
+ drop_remainder=False,
+ include_example_id=True))
+
+ config = perceiver.SentencePredictionConfig(
+ metric_type="matthews_corrcoef",
+ model=self.get_model_config(2),
+ train_data=train_data_config)
+ outputs = self._run_task(config)
+ self.assertEqual(outputs["sentence_prediction"].shape.as_list(), [8, 1])
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "classification",
+ "num_classes": 5,
+ }, {
+ "testcase_name": "regression",
+ "num_classes": 1,
+ })
+ def test_prediction(self, num_classes):
+ config = perceiver.SentencePredictionConfig(
+ model=self.get_model_config(num_classes=num_classes),
+ train_data=self._train_data_config)
+ task = perceiver_pred.SentencePredictionTask(config)
+ model = task.build_model()
+
+ test_data_path = os.path.join(self.get_temp_dir(), "test.tf_record")
+ seq_length = 16
+ num_examples = 100
+ _create_fake_dataset(
+ test_data_path,
+ seq_length=seq_length,
+ num_classes=num_classes,
+ num_examples=num_examples)
+
+ test_data_config = (
+ sentence_prediction_dataloader.SentencePredictionDataConfig(
+ input_path=test_data_path,
+ seq_length=seq_length,
+ is_training=False,
+ label_type="int" if num_classes > 1 else "float",
+ global_batch_size=16,
+ drop_remainder=False,
+ include_example_id=True))
+
+ predictions = sentence_prediction.predict(task, test_data_config, model)
+ self.assertLen(predictions, num_examples)
+ for prediction in predictions:
+ self.assertEqual(prediction.dtype,
+ tf.int64 if num_classes > 1 else tf.float32)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/perceiver/train.py b/official/projects/perceiver/train.py
new file mode 100644
index 00000000000..d728f18a7ed
--- /dev/null
+++ b/official/projects/perceiver/train.py
@@ -0,0 +1,29 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""TensorFlow Model Garden training driver, register Perceiver configs."""
+
+from absl import app
+
+from official.common import flags as tfm_flags
+from official.nlp import train
+# pylint: disable=unused-import
+from official.projects.perceiver.configs import perceiver
+from official.projects.perceiver.tasks import pretrain
+from official.projects.perceiver.tasks import sentence_prediction
+# pylint: enable=unused-import
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(train.main)
diff --git a/official/projects/pix2seq/README.md b/official/projects/pix2seq/README.md
new file mode 100644
index 00000000000..0dbf04695e2
--- /dev/null
+++ b/official/projects/pix2seq/README.md
@@ -0,0 +1,60 @@
+# Pix2Seq: A Language Modeling Framework for Object Detection
+
+[](https://arxiv.org/abs/2109.10852).
+
+TensorFlow 2 implementation of A Language Modeling Framework for Object
+Detection.
+
+The official implementation of Pix2Seq in Tensorflow 2 is [Here]
+(https://github.com/google-research/pix2seq).
+
+⚠️ Disclaimer: All datasets hyperlinked from this page are not owned or
+distributed by Google. The dataset is made available by third parties. Please
+review the terms and conditions made available by the third parties before using
+the data.
+
+## Training
+To train the model on MS-COCO, try the following command:
+
+```
+python3 train.py \
+ --mode=train \
+ --experiment=pix2seq_r50_coco \
+ --model_dir=$MODEL_DIR \
+ --config_file=./configs/experiments/coco_pix2seq_r50_gpu.yaml
+```
+
+## Evaluation
+To evaluate the model on MS-COCO, try the following command:
+
+```
+python3 train.py \
+ --mode=eval \
+ --experiment=pix2seq_r50_coco \
+ --model_dir=$MODEL_DIR \
+ --config_file=./configs/experiments/coco_pix2seq_r50_gpu.yaml
+```
+
+## Cite
+
+[Pix2seq paper](https://arxiv.org/abs/2109.10852):
+
+
+```
+@article{chen2021pix2seq,
+ title={Pix2seq: A language modeling framework for object detection},
+ author={Chen, Ting and Saxena, Saurabh and Li, Lala and Fleet, David J and Hinton, Geoffrey},
+ journal={arXiv preprint arXiv:2109.10852},
+ year={2021}
+}
+```
+
+## Contributors
+
+
+* Gunho Park ([Github @gunho1123](https://github.com/gunho1123))
+* Jiageng Zhang ([Github @Zarjagen](https://github.com/Zarjagen))
+* Shicheng Xu ([Github @lightxu](https://github.com/lightxu))
+* Tyler Scott ([Github @tylersco](https://github.com/tylersco))
+* Yu Lou ([Github @LouYu2015](https://github.com/LouYu2015))
+
diff --git a/official/projects/pix2seq/configs/pix2seq.py b/official/projects/pix2seq/configs/pix2seq.py
new file mode 100644
index 00000000000..cde02f2fcff
--- /dev/null
+++ b/official/projects/pix2seq/configs/pix2seq.py
@@ -0,0 +1,280 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pix2Seq configurations."""
+
+import dataclasses
+import os
+from typing import List, Optional, Union
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.modeling import optimization
+from official.projects.uvit.configs import backbones as uvit_backbones
+from official.vision.configs import backbones
+from official.vision.configs import common
+
+
+# Vocab.
+
+# A shared vocab among tasks and its structure -
+# Special tokens: [0, 99).
+# Class tokens: [100, coord_vocab_shift).
+# Coordinate tokens: [coord_vocab_shift, text_vocab_shift).
+
+PADDING_TOKEN = 0
+
+# 10-29 reserved for task id.
+
+FAKE_CLASS_TOKEN = 30
+FAKE_TEXT_TOKEN = 30 # Same token to represent fake class and fake text.
+SEPARATOR_TOKEN = 40
+INVISIBLE_TOKEN = 41
+
+BASE_VOCAB_SHIFT = 100
+
+# Floats used to represent padding and separator in the flat list of polygon
+# coords, and invisibility in the key points.
+PADDING_FLOAT = -1.0
+SEPARATOR_FLOAT = -2.0
+INVISIBLE_FLOAT = -3.0
+FLOATS = [PADDING_FLOAT, SEPARATOR_FLOAT, INVISIBLE_FLOAT]
+TOKENS = [PADDING_TOKEN, SEPARATOR_TOKEN, INVISIBLE_TOKEN]
+FLOAT_TO_TOKEN = dict(zip(FLOATS, TOKENS))
+TOKEN_TO_FLOAT = dict(zip(TOKENS, FLOATS))
+
+OD_ID = 10
+
+
+@dataclasses.dataclass
+class DataConfig(cfg.DataConfig):
+ """Input config for training."""
+
+ input_path: str = ''
+ tfds_name: str = ''
+ tfds_split: str = 'train'
+ global_batch_size: int = 0
+ is_training: bool = False
+ dtype: str = 'float32'
+ decoder: common.DataDecoder = dataclasses.field(
+ default_factory=common.DataDecoder
+ )
+ shuffle_buffer_size: int = 10000
+ file_type: str = 'tfrecord'
+ drop_remainder: bool = True
+ aug_scale_min: float = 1.0
+ aug_scale_max: float = 1.0
+ aug_color_jitter_strength: float = 0.0
+ label_shift: int = 0
+
+
+@dataclasses.dataclass
+class Losses(hyperparams.Config):
+ noise_bbox_weight: float = 1.0
+ eos_token_weight: float = 0.1
+ l2_weight_decay: float = 1e-4
+
+
+@dataclasses.dataclass
+class Backbone(backbones.Backbone):
+ """Configuration for backbones.
+
+ Attributes:
+ type: "str", type of backbone be used, one the of fields below.
+ uvit: uvit backbone config.
+ """
+ type: Optional[str] = None
+ resnet: backbones.ResNet = dataclasses.field(default_factory=backbones.ResNet)
+ uvit: uvit_backbones.VisionTransformer = dataclasses.field(
+ default_factory=uvit_backbones.VisionTransformer)
+
+
+@dataclasses.dataclass
+class BackboneConfig(hyperparams.Config):
+ """Configuration for backbones."""
+
+ backbone: Backbone = dataclasses.field(default_factory=Backbone)
+ # Whether to freeze this backbone during training.
+ freeze: bool = False
+ # The endpoint name of the features to extract from the backbone.
+ endpoint_name: str = '5'
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
+ # Optional checkpoint to load for this backbone.
+ init_checkpoint: Optional[str] = None
+ # If loading an init_checkpoint, whether to assert that all objects in the
+ # Python program are matched by the checkpoint.
+ # If False, understand that only the weak assertion of a non-trivial match
+ # will be made.
+ assert_existing_objects_matched: bool = True
+
+
+@dataclasses.dataclass
+class Pix2Seq(hyperparams.Config):
+ """Pix2Seq model definitions."""
+
+ max_num_instances: int = 100
+ hidden_size: int = 256
+ num_heads: int = 8
+ num_encoder_layers: int = 6
+ num_decoder_layers: int = 6
+ vocab_size: int = 3000
+ use_cls_token: bool = False
+ shared_decoder_embedding: bool = True
+ decoder_output_bias: bool = True
+ # The shape of the input image. For example, [640, 640, 3] for RGB.
+ # If using multiple backbones, this input size is understood to be the same
+ # for all backbones. If you need a separate input size for each backbone,
+ # please implement this behavior.
+ input_size: List[int] = dataclasses.field(default_factory=list)
+ # Backbones for each image modality. The RGB backbone is always the first one.
+ # If just using RGB, you should only set one backbone.
+ backbones: List[BackboneConfig] = dataclasses.field(
+ default_factory=lambda: [ # pylint: disable=g-long-lambda
+ BackboneConfig(
+ backbone=Backbone(
+ type='resnet',
+ resnet=backbones.ResNet(model_id=50, bn_trainable=False),
+ )
+ )
+ ]
+ )
+ drop_path: float = 0.1
+ # The dropout rates applied to the features extracted from each backbone.
+ encoded_feature_dropout_rates: List[float] = dataclasses.field(
+ default_factory=lambda: [0.1]
+ )
+ # The dropout rate for the transformer.
+ drop_units: float = 0.1
+ drop_att: float = 0.0
+ norm_first: bool = True
+ temperature: float = 1.0
+ top_k: int = 0
+ top_p: float = 0.4
+ early_stopping_token: int | None = None
+
+
+@dataclasses.dataclass
+class Pix2SeqTask(cfg.TaskConfig):
+ model: Pix2Seq = dataclasses.field(default_factory=Pix2Seq)
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+ init_checkpoint_modules: Union[str, List[str]] = 'all' # all, backbone
+ annotation_file: Optional[str] = None
+ per_category_metrics: bool = False
+ coord_vocab_shift: int = 1000
+ quantization_bins: int = 1000
+
+
+COCO_INPUT_PATH_BASE = 'coco'
+COCO_TRAIN_EXAMPLES = 118287
+COCO_VAL_EXAMPLES = 5000
+
+
+@exp_factory.register_config_factory('pix2seq_r50_coco')
+def pix2seq_r50_coco() -> cfg.ExperimentConfig:
+ """Config to get results that matches the paper."""
+ train_batch_size = 128
+ eval_batch_size = 16
+ steps_per_epoch = COCO_TRAIN_EXAMPLES // train_batch_size
+ train_steps = 80 * steps_per_epoch
+ config = cfg.ExperimentConfig(
+ task=Pix2SeqTask(
+ init_checkpoint='',
+ init_checkpoint_modules='backbone',
+ annotation_file=os.path.join(
+ COCO_INPUT_PATH_BASE, 'instances_val2017.json'
+ ),
+ model=Pix2Seq(
+ input_size=[640, 640, 3],
+ backbones=[
+ BackboneConfig(
+ backbone=Backbone(
+ type='resnet',
+ resnet=backbones.ResNet(model_id=50),
+ ),
+ norm_activation=common.NormActivation(
+ norm_momentum=0.9, norm_epsilon=1e-5, use_sync_bn=True
+ ),
+ init_checkpoint='',
+ )
+ ],
+ ),
+ losses=Losses(l2_weight_decay=0.0),
+ train_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=train_batch_size * 10,
+ aug_scale_min=0.3,
+ aug_scale_max=2.0,
+ aug_color_jitter_strength=0.0,
+ ),
+ validation_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=True,
+ ),
+ ),
+ trainer=cfg.TrainerConfig(
+ train_steps=train_steps,
+ validation_steps=COCO_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=5 * steps_per_epoch,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ max_to_keep=10,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw_experimental',
+ 'adamw_experimental': {
+ 'epsilon': 1.0e-08,
+ 'weight_decay': 0.05,
+ 'beta_1': 0.9,
+ 'beta_2': 0.95,
+ 'global_clipnorm': -1.0,
+ },
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'initial_learning_rate': 0.0001,
+ 'end_learning_rate': 0.000001,
+ 'offset': 0,
+ 'power': 1.0,
+ 'decay_steps': 80 * steps_per_epoch,
+ },
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 2 * steps_per_epoch,
+ 'warmup_learning_rate': 0,
+ },
+ },
+ }),
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ],
+ )
+ return config
diff --git a/official/projects/pix2seq/configs/pix2seq_test.py b/official/projects/pix2seq/configs/pix2seq_test.py
new file mode 100644
index 00000000000..25af8276a71
--- /dev/null
+++ b/official/projects/pix2seq/configs/pix2seq_test.py
@@ -0,0 +1,40 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for Pix2Seq config."""
+
+# pylint: disable=unused-import
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.pix2seq.configs import pix2seq as exp_cfg
+
+
+class Pix2SeqTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.parameters(('pix2seq_r50_coco',))
+ def test_pix2seq_configs(self, config_name):
+ config = exp_factory.get_exp_config(config_name)
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task, exp_cfg.Pix2SeqTask)
+ self.assertIsInstance(config.task.train_data, cfg.DataConfig)
+ config.task.train_data.is_training = None
+ with self.assertRaises(KeyError):
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pix2seq/dataloaders/pix2seq_input.py b/official/projects/pix2seq/dataloaders/pix2seq_input.py
new file mode 100644
index 00000000000..8be3e990eef
--- /dev/null
+++ b/official/projects/pix2seq/dataloaders/pix2seq_input.py
@@ -0,0 +1,298 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""COCO data loader for Pix2Seq."""
+
+from typing import Tuple
+import tensorflow as tf, tf_keras
+
+from official.projects.pix2seq import utils
+from official.projects.pix2seq.configs import pix2seq as pix2seq_cfg
+from official.projects.simclr.dataloaders import preprocess_ops as simclr_preprocess_ops
+from official.vision.dataloaders import parser
+from official.vision.ops import box_ops
+from official.vision.ops import preprocess_ops
+
+RESIZE_SCALES = (480, 512, 544, 576, 608, 640)
+
+
+class Parser(parser.Parser):
+ """Parse an image and its annotations into a dictionary of tensors."""
+
+ def __init__(
+ self,
+ eos_token_weight: float = 0.1,
+ output_size: Tuple[int, int] = (1333, 1333),
+ max_num_boxes: int = 100,
+ aug_rand_hflip=True,
+ aug_scale_min=0.3,
+ aug_scale_max=2.0,
+ aug_color_jitter_strength: float = 0.5,
+ aug_color_jitter_impl='simclrv2',
+ coord_vocab_shift=1000,
+ quantization_bins=1000,
+ skip_crowd_during_training=True,
+ label_shift: int = 0,
+ ):
+ self._eos_token_weight = eos_token_weight
+ self._output_size = output_size
+ self._max_num_boxes = max_num_boxes
+ self._aug_rand_hflip = aug_rand_hflip
+ self._aug_scale_min = aug_scale_min
+ self._aug_scale_max = aug_scale_max
+ self._aug_color_jitter_strength = aug_color_jitter_strength
+ self._aug_color_jitter_impl = aug_color_jitter_impl
+ self._coord_vocab_shift = coord_vocab_shift
+ self._quantization_bins = quantization_bins
+ self._skip_crowd_during_training = skip_crowd_during_training
+ self._label_shift = label_shift
+
+ def _parse_train_data(self, data):
+ """Parses data for training and evaluation."""
+ classes = data['groundtruth_classes'] + self._label_shift
+ boxes = data['groundtruth_boxes']
+
+ is_crowds = data['groundtruth_is_crowd']
+ # Skips annotations with `is_crowd` = True.
+ if self._skip_crowd_during_training:
+ num_groundtruths = tf.shape(classes)[0]
+ with tf.control_dependencies([num_groundtruths, is_crowds]):
+ indices = tf.cond(
+ tf.greater(tf.size(is_crowds), 0),
+ lambda: tf.where(tf.logical_not(is_crowds))[:, 0],
+ lambda: tf.cast(tf.range(num_groundtruths), tf.int64),
+ )
+ classes = tf.gather(classes, indices)
+ boxes = tf.gather(boxes, indices)
+
+ # Gets original image.
+ image = data['image']
+
+ # Normalizes image with mean and std pixel values.
+ image = tf.image.convert_image_dtype(image, dtype=tf.float32)
+ # Color jitter.
+ image = simclr_preprocess_ops.random_color_jitter(
+ image=image,
+ color_jitter_strength=self._aug_color_jitter_strength,
+ impl=self._aug_color_jitter_impl,
+ )
+ image = tf.clip_by_value(image, 0.0, 1.0)
+ image, boxes, _ = preprocess_ops.random_horizontal_flip(image, boxes)
+
+ image_shape = tf.shape(image)[:2]
+ boxes = box_ops.denormalize_boxes(boxes, image_shape)
+
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ self._output_size,
+ padded_size=self._output_size,
+ aug_scale_min=self._aug_scale_min,
+ aug_scale_max=self._aug_scale_max)
+
+ boxes = preprocess_ops.resize_and_crop_boxes(
+ boxes, image_info[2, :], image_info[1, :], image_info[3, :]
+ )
+ boxes = box_ops.normalize_boxes(boxes, image_info[1, :])
+
+ # Filters out ground truth boxes that are all zeros.
+ indices = box_ops.get_non_empty_box_indices(boxes)
+ boxes = tf.gather(boxes, indices)
+ classes = tf.gather(classes, indices)
+
+ boxes, classes = utils.reorder_object_instances(boxes, classes, 'random')
+ boxes, classes = utils.inject_noise_bbox(
+ boxes, classes, self._max_num_boxes
+ )
+
+ boxes = utils.clip_or_pad_to_max_len(boxes, self._max_num_boxes, 0)
+ classes = utils.clip_or_pad_to_max_len(classes, self._max_num_boxes, 0)
+
+ outputs = self.build_response_seq_from_bbox(
+ boxes, classes, self._coord_vocab_shift, self._quantization_bins
+ )
+ response_seq, response_seq_class_m, token_weights = outputs
+ prompt_seq = utils.build_prompt_seq_from_task_id(
+ pix2seq_cfg.OD_ID, response_seq
+ ) # (1)
+ input_seq = tf.concat([prompt_seq, response_seq_class_m], -1)
+ target_seq = tf.concat([prompt_seq, response_seq], -1)
+
+ backgrnd_val = 0.3
+ image = backgrnd_val + tf.image.pad_to_bounding_box(
+ image - backgrnd_val, 0, 0, self._output_size[0], self._output_size[1]
+ )
+
+ input_seq = utils.clip_or_pad_to_max_len(
+ input_seq, self._max_num_boxes * 5 + 1, -1)
+ target_seq = utils.clip_or_pad_to_max_len(
+ target_seq, self._max_num_boxes * 5 + 1, -1
+ )
+
+ input_seq, target_seq = input_seq[..., :-1], target_seq[..., 1:]
+ token_weights = utils.clip_or_pad_to_max_len(
+ token_weights, self._max_num_boxes * 5, -1
+ )
+
+ # Assign lower weights for ending/padding tokens.
+ token_weights = tf.where(
+ target_seq == pix2seq_cfg.PADDING_TOKEN,
+ tf.zeros_like(token_weights) + self._eos_token_weight,
+ token_weights,
+ )
+
+ labels = {
+ 'targets': target_seq,
+ 'weights': token_weights,
+ 'inputs': input_seq,
+ }
+
+ return image, labels
+
+ def build_response_seq_from_bbox(
+ self,
+ bbox,
+ label,
+ coord_vocab_shift,
+ quantization_bins,
+ noise_bbox_weight=1.0,
+ class_label_corruption='rand_n_fake_cls',
+ ):
+ """Build target seq from bounding bboxes for object detection.
+
+ Objects are serialized using the format of yxyxc.
+
+ Args:
+ bbox: `float` bounding box of shape (n, 4).
+ label: `int` label of shape (n).
+ coord_vocab_shift: `int`, shifting coordinates by a specified integer.
+ quantization_bins: `int`.
+ noise_bbox_weight: `float` on the token weights for noise bboxes.
+ class_label_corruption: `string` specifying how labels are corrupted for
+ the input_seq.
+
+ Returns:
+ discrete sequences with shape (seqlen).
+ """
+ # Bbox and label quantization.
+ is_padding = tf.expand_dims(tf.equal(label, 0), -1)
+ quantized_bbox = utils.quantize(bbox, quantization_bins)
+ quantized_bbox = quantized_bbox + coord_vocab_shift
+ quantized_bbox = tf.where(
+ is_padding, tf.zeros_like(quantized_bbox), quantized_bbox
+ )
+ new_label = tf.expand_dims(label + pix2seq_cfg.BASE_VOCAB_SHIFT, -1)
+ new_label = tf.where(is_padding, tf.zeros_like(new_label), new_label)
+ lb_shape = tf.shape(new_label)
+
+ # Bbox and label serialization.
+ response_seq = tf.concat([quantized_bbox, new_label], axis=-1)
+
+ response_seq = tf.reshape(response_seq, [-1])
+ rand_cls = pix2seq_cfg.BASE_VOCAB_SHIFT + tf.random.uniform(
+ lb_shape,
+ 0,
+ coord_vocab_shift - pix2seq_cfg.BASE_VOCAB_SHIFT,
+ dtype=new_label.dtype,
+ )
+ fake_cls = pix2seq_cfg.FAKE_CLASS_TOKEN + tf.zeros_like(new_label)
+ rand_n_fake_cls = tf.where(
+ tf.random.uniform(lb_shape) > 0.5, rand_cls, fake_cls
+ )
+ real_n_fake_cls = tf.where(
+ tf.random.uniform(lb_shape) > 0.5, new_label, fake_cls
+ )
+ real_n_rand_n_fake_cls = tf.where(
+ tf.random.uniform(lb_shape) > 0.5, new_label, rand_n_fake_cls
+ )
+ label_mapping = {
+ 'none': new_label,
+ 'rand_cls': rand_cls,
+ 'real_n_fake_cls': real_n_fake_cls,
+ 'rand_n_fake_cls': rand_n_fake_cls,
+ 'real_n_rand_n_fake_cls': real_n_rand_n_fake_cls,
+ }
+ new_label_m = label_mapping[class_label_corruption]
+ new_label_m = tf.where(is_padding, tf.zeros_like(new_label_m), new_label_m)
+
+ response_seq_class_m = tf.concat([quantized_bbox, new_label_m], axis=-1)
+ response_seq_class_m = tf.reshape(response_seq_class_m, [-1])
+
+ # Get token weights.
+ is_real = tf.cast(
+ tf.not_equal(new_label, pix2seq_cfg.FAKE_CLASS_TOKEN), tf.float32
+ )
+ bbox_weight = tf.tile(is_real, [1, 4])
+ label_weight = is_real + (1.0 - is_real) * noise_bbox_weight
+ token_weights = tf.concat([bbox_weight, label_weight], -1)
+ token_weights = tf.reshape(token_weights, [-1])
+
+ return response_seq, response_seq_class_m, token_weights
+
+ def _parse_eval_data(self, data):
+ """Parses data for training and evaluation."""
+ classes = data['groundtruth_classes'] + self._label_shift
+ boxes = data['groundtruth_boxes']
+ is_crowd = data['groundtruth_is_crowd']
+
+ # Gets original image and its size.
+ image = data['image']
+ image = tf.image.convert_image_dtype(image, dtype=tf.float32)
+
+ image_shape = tf.shape(image)[:2]
+ boxes = box_ops.denormalize_boxes(boxes, image_shape)
+ gt_boxes = boxes
+ image, image_info = preprocess_ops.resize_image(
+ image, min(self._output_size), max(self._output_size)
+ )
+ boxes = preprocess_ops.resize_and_crop_boxes(
+ boxes, image_info[2, :], image_info[1, :], image_info[3, :]
+ )
+ scale = tf.cast(
+ tf.concat([self._output_size, self._output_size], -1), boxes.dtype
+ )
+ boxes = boxes / scale
+
+ # Filters out ground truth boxes that are all zeros.
+ indices = box_ops.get_non_empty_box_indices(boxes)
+ boxes = tf.gather(boxes, indices)
+ classes = tf.gather(classes, indices)
+ is_crowd = tf.gather(is_crowd, indices)
+
+ prompt_seq = tf.constant([pix2seq_cfg.OD_ID], dtype=tf.int64)
+ backgrnd_val = 0.3
+ image = backgrnd_val + tf.image.pad_to_bounding_box(
+ image - backgrnd_val, 0, 0, self._output_size[0], self._output_size[1]
+ )
+
+ labels = {
+ 'prompt': prompt_seq,
+ 'classes': preprocess_ops.clip_or_pad_to_fixed_size(
+ classes, self._max_num_boxes
+ ),
+ 'boxes': preprocess_ops.clip_or_pad_to_fixed_size(
+ boxes, self._max_num_boxes
+ ),
+ }
+ labels.update({
+ 'id': int(data['source_id']),
+ 'image_info': image_info,
+ 'is_crowd': preprocess_ops.clip_or_pad_to_fixed_size(
+ is_crowd, self._max_num_boxes
+ ),
+ 'gt_boxes': preprocess_ops.clip_or_pad_to_fixed_size(
+ gt_boxes, self._max_num_boxes
+ ),
+ })
+
+ return image, labels
diff --git a/official/projects/pix2seq/dataloaders/pix2seq_input_test.py b/official/projects/pix2seq/dataloaders/pix2seq_input_test.py
new file mode 100644
index 00000000000..2106af5b9a3
--- /dev/null
+++ b/official/projects/pix2seq/dataloaders/pix2seq_input_test.py
@@ -0,0 +1,136 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for Pix2Seq input."""
+import io
+
+import numpy as np
+from PIL import Image
+import tensorflow as tf, tf_keras
+
+from official.projects.pix2seq.dataloaders import pix2seq_input
+from official.vision.dataloaders import tf_example_decoder
+
+
+IMAGE_KEY = 'image/encoded'
+LABEL_KEY = 'image/object/class/label'
+
+
+def _bytes_feature(value):
+ """Returns a bytes_list from a string / byte."""
+ if isinstance(value, type(tf.constant(0))):
+ value = (
+ value.numpy()
+ ) # BytesList won't unpack a string from an EagerTensor.
+ return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))
+
+
+def _float_feature(value):
+ """Returns a float_list from a float / double."""
+ return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))
+
+
+def _int64_feature(value):
+ """Returns an int64_list from a bool / enum / int / uint."""
+ return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))
+
+
+def fake_seq_example():
+ # Create fake data.
+ random_image = np.random.randint(0, 256, size=(480, 640, 3), dtype=np.uint8)
+ random_image = Image.fromarray(random_image)
+ labels = [42, 5]
+ with io.BytesIO() as buffer:
+ random_image.save(buffer, format='JPEG')
+ raw_image_bytes = buffer.getvalue()
+
+ xmins = [0.23, 0.15]
+ xmaxs = [0.54, 0.60]
+ ymins = [0.11, 0.5]
+ ymaxs = [0.86, 0.72]
+
+ feature = {
+ 'image/encoded': _bytes_feature(raw_image_bytes),
+ 'image/height': _int64_feature(480),
+ 'image/width': _int64_feature(640),
+ 'image/object/bbox/xmin': tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmins)
+ ),
+ 'image/object/bbox/xmax': tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmaxs)
+ ),
+ 'image/object/bbox/ymin': tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymins)
+ ),
+ 'image/object/bbox/ymax': tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymaxs)
+ ),
+ 'image/object/class/label': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=labels)
+ ),
+ 'image/object/area': tf.train.Feature(
+ float_list=tf.train.FloatList(value=[1., 2.])
+ ),
+ 'image/object/is_crowd': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[0, 0])
+ ),
+ 'image/source_id': _bytes_feature(b'123'),
+ }
+
+ # Create a Features message using tf.train.Example.
+
+ example_proto = tf.train.Example(features=tf.train.Features(feature=feature))
+
+ return example_proto, labels
+
+
+class Pix2SeqParserTest(tf.test.TestCase):
+
+ def test_image_input_train(self):
+ decoder = tf_example_decoder.TfExampleDecoder()
+ parser = pix2seq_input.Parser(
+ eos_token_weight=0.1,
+ output_size=[640, 640],
+ max_num_boxes=10,
+ ).parse_fn(True)
+
+ seq_example, _ = fake_seq_example()
+
+ input_tensor = tf.constant(seq_example.SerializeToString())
+ decoded_tensors = decoder.decode(input_tensor)
+ output_tensor = parser(decoded_tensors)
+ image, _ = output_tensor
+
+ self.assertAllEqual(image.shape, (640, 640, 3))
+
+ def test_image_input_eval(self):
+ decoder = tf_example_decoder.TfExampleDecoder()
+ parser = pix2seq_input.Parser(
+ eos_token_weight=0.1,
+ output_size=[640, 640],
+ max_num_boxes=10,
+ ).parse_fn(False)
+
+ seq_example, _ = fake_seq_example()
+
+ input_tensor = tf.constant(seq_example.SerializeToString())
+ decoded_tensors = decoder.decode(input_tensor)
+ output_tensor = parser(decoded_tensors)
+ image, _ = output_tensor
+
+ self.assertAllEqual(image.shape, (640, 640, 3))
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pix2seq/modeling/pix2seq_model.py b/official/projects/pix2seq/modeling/pix2seq_model.py
new file mode 100644
index 00000000000..3fa7f3a3629
--- /dev/null
+++ b/official/projects/pix2seq/modeling/pix2seq_model.py
@@ -0,0 +1,759 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Implements A Language Modeling Framework for Object Detection.
+
+Model paper: https://arxiv.org/abs/2109.10852
+This module does not support Keras de/serialization. Please use
+tf.train.Checkpoint for object based saving and loading and tf.saved_model.save
+for graph serialization.
+"""
+
+import math
+from typing import Any, List, Mapping, Optional, Sequence, Union
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.projects.pix2seq.modeling import transformer
+
+
+def get_shape(x):
+ static = x.shape.as_list()
+ dynamic = tf.shape(x)
+ return [dynamic[i] if s is None else s for i, s in enumerate(static)]
+
+
+def get_variable_initializer(name=None):
+ if name is None:
+ return tf_keras.initializers.TruncatedNormal(mean=0.0, stddev=0.02)
+
+
+def add_seq_pos_emb(
+ self, pos_encoding, max_seq_len, dim, name_prefix=None, initializer=None
+):
+ """Add seq_pos_emb variable/tensor to model instance referenced by `self`."""
+ if name_prefix is None:
+ name_prefix = self.name
+ if initializer is None:
+ initializer = get_variable_initializer()
+ if pos_encoding == "learned":
+ self.seq_pos_emb = self.add_weight(
+ shape=(max_seq_len + 1, dim),
+ initializer=initializer,
+ name="%s/seq_pos_embedding" % name_prefix,
+ )
+ # (gunho) currently only 'learned' positional encoding is supported
+ elif pos_encoding == "sin_cos":
+ self.seq_pos_emb = None
+ else:
+ raise ValueError("Unknown pos encoding %s" % pos_encoding)
+
+
+def add_vocab_token_emb(
+ self,
+ vocab_size,
+ dim,
+ output_bias,
+ name_prefix=None,
+ initializer=None,
+):
+ """Add token_embedding variable to model instance referenced by `self`."""
+ if name_prefix is None:
+ name_prefix = self.name
+ if initializer is None:
+ initializer = get_variable_initializer()
+ self.token_embedding = self.add_weight(
+ shape=[vocab_size, dim],
+ initializer=initializer,
+ name="%s/token_embedding" % name_prefix,
+ )
+ if output_bias:
+ self.outp_bias = self.add_weight(
+ shape=[vocab_size],
+ initializer=initializer,
+ name="%s/outp_bias" % name_prefix,
+ )
+
+
+def get_ar_mask(seq_len, dtype=tf.float32):
+ """Get autoregressive causal mask so the model cannot attends to the future.
+
+ Args:
+ seq_len: a `int` or `int` tensor specifying the sequence length.
+ dtype: tf data type for the return tensor.
+
+ Returns:
+ tensor of shape [1, 1, seq_len, seq_len] with ones for locations to be
+ masked out.
+ """
+ valid_locs = tf.linalg.band_part(
+ tf.ones([seq_len, seq_len], dtype=dtype), -1, 0
+ )
+ valid_locs = tf.reshape(valid_locs, [1, 1, seq_len, seq_len])
+ return 1.0 - valid_locs
+
+
+def position_embedding_sine(
+ attention_mask,
+ num_pos_features=256,
+ temperature=10000.0,
+ normalize=True,
+ scale=2 * math.pi,
+):
+ """Sine-based positional embeddings for 2D images.
+
+ Args:
+ attention_mask: a `bool` Tensor specifying the size of the input image to
+ the Transformer and which elements are padded, of size [batch_size,
+ height, width]
+ num_pos_features: a `int` specifying the number of positional features,
+ should be equal to the hidden size of the Transformer network
+ temperature: a `float` specifying the temperature of the positional
+ embedding. Any type that is converted to a `float` can also be accepted.
+ normalize: a `bool` determining whether the positional embeddings should be
+ normalized between [0, scale] before application of the sine and cos
+ functions.
+ scale: a `float` if normalize is True specifying the scale embeddings before
+ application of the embedding function.
+
+ Returns:
+ embeddings: a `float` tensor of the same shape as input_tensor specifying
+ the positional embeddings based on sine features.
+ """
+ if num_pos_features % 2 != 0:
+ raise ValueError(
+ "Number of embedding features (num_pos_features) must be even when "
+ "column and row embeddings are concatenated."
+ )
+ num_pos_features = num_pos_features // 2
+
+ # Produce row and column embeddings based on total size of the image
+ # [batch_size, height, width]
+ attention_mask = tf.cast(attention_mask, tf.float32)
+ row_embedding = tf.cumsum(attention_mask, 1)
+ col_embedding = tf.cumsum(attention_mask, 2)
+
+ if normalize:
+ eps = 1e-6
+ row_embedding = row_embedding / (row_embedding[:, -1:, :] + eps) * scale
+ col_embedding = col_embedding / (col_embedding[:, :, -1:] + eps) * scale
+
+ dim_t = tf.range(num_pos_features, dtype=row_embedding.dtype)
+ dim_t = tf.pow(temperature, 2 * (dim_t // 2) / num_pos_features)
+
+ # Creates positional embeddings for each row and column position
+ # [batch_size, height, width, num_pos_features]
+ pos_row = tf.expand_dims(row_embedding, -1) / dim_t
+ pos_col = tf.expand_dims(col_embedding, -1) / dim_t
+ pos_row = tf.stack(
+ [tf.sin(pos_row[:, :, :, 0::2]), tf.cos(pos_row[:, :, :, 1::2])], axis=4
+ )
+ pos_col = tf.stack(
+ [tf.sin(pos_col[:, :, :, 0::2]), tf.cos(pos_col[:, :, :, 1::2])], axis=4
+ )
+
+ final_shape = tf_utils.get_shape_list(pos_row)[:3] + [-1]
+ pos_row = tf.reshape(pos_row, final_shape)
+ pos_col = tf.reshape(pos_col, final_shape)
+ output = tf.concat([pos_row, pos_col], -1)
+
+ embeddings = tf.cast(output, tf.float32)
+ return embeddings
+
+
+def top_logits(
+ logits: tf.Tensor, k: int = 0, p: float = 1.0, mask: float = -1e10
+) -> tf.Tensor:
+ """Remove low probability logits via masking.
+
+ Args:
+ logits: class logits in shape of (batch size, total_classes).
+ k: specifying top k largest logits to keep.
+ p: specifying a probability for finding a minimum set of largest logits to
+ keep, where their cumulative probability is no less than p (actually in
+ the following version, it is "...cumulative probability is the largest but
+ no more than p").
+ mask: a value that's used to replace logits that don't satisfy the keep
+ conditions.
+
+ Returns:
+ logits where low probability ones are replaced with mask.
+ """
+ mask = tf.ones_like(logits) * mask
+ if k > 0:
+ min_logits = tf.nn.top_k(logits, k=k)[0][:, -1:]
+ logits = tf.where(logits < min_logits, mask, logits)
+ if p < 1.0:
+ sorted_logits = tf.sort(logits, direction="DESCENDING", axis=-1)
+ cum_probs = tf.cumsum(tf.nn.softmax(sorted_logits, axis=-1), axis=-1)
+ min_logits = -tf.reduce_max(
+ tf.where(cum_probs <= p, -sorted_logits, mask), -1, keepdims=True
+ )
+ min_logits = tf.minimum(min_logits, sorted_logits[:, :1])
+ logits = tf.where(logits < min_logits, mask, logits)
+ return logits
+
+
+class Pix2Seq(tf_keras.Model):
+ """Pix2Seq model with Keras.
+
+ Pix2Seq consists of backbone, input token embedding, Pix2SeqTransformer.
+ """
+
+ def __init__(
+ self,
+ backbones: Sequence[tf_keras.Model],
+ backbone_endpoint_names: Sequence[str],
+ max_seq_len,
+ vocab_size,
+ hidden_size,
+ num_heads,
+ num_encoder_layers=6,
+ num_decoder_layers=6,
+ drop_path=0.1,
+ encoded_feature_dropout_rates: Sequence[float] = (0.1,),
+ drop_units=0.1,
+ drop_att=0.0,
+ temperature=1.0,
+ top_k=0,
+ top_p=0.4,
+ early_stopping_token: int | None = None,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self._backbones = backbones
+ self._backbone_endpoint_names = backbone_endpoint_names
+ self._max_seq_len = max_seq_len
+ self._vocab_size = vocab_size
+ self._hidden_size = hidden_size
+ self._num_heads = num_heads
+ self._num_encoder_layers = num_encoder_layers
+ self._num_decoder_layers = num_decoder_layers
+ self._drop_path = drop_path
+ self._drop_units = drop_units
+ self._drop_att = drop_att
+ if hidden_size % 2 != 0:
+ raise ValueError("hidden_size must be a multiple of 2.")
+ if len(encoded_feature_dropout_rates) != len(self._backbones):
+ raise ValueError(
+ "The length of encoded_feature_dropout_rates must be equal to the "
+ "number of backbones."
+ )
+
+ self._encoder_dropouts = [
+ tf_keras.layers.Dropout(r) for r in encoded_feature_dropout_rates
+ ]
+ # Separate projections and learned layer normalization for each image.
+ num_backbones = len(self._backbones)
+ self._stem_projections = [
+ tf_keras.layers.Dense(self._hidden_size, name="stem_projection")
+ for _ in range(num_backbones)
+ ]
+ self._stem_lns = [
+ tf_keras.layers.LayerNormalization(epsilon=1e-6, name="stem_ln")
+ for _ in range(num_backbones)
+ ]
+
+ self._transformer = Pix2SeqTransformer(
+ max_seq_len=self._max_seq_len,
+ vocab_size=self._vocab_size,
+ hidden_size=self._hidden_size,
+ num_sources=num_backbones,
+ pos_encoding="learned",
+ num_encoder_layers=self._num_encoder_layers,
+ num_decoder_layers=self._num_decoder_layers,
+ drop_path=self._drop_path,
+ drop_units=self._drop_units,
+ drop_att=self._drop_att,
+ num_heads=self._num_heads,
+ )
+ self._temperature = temperature
+ self._top_k = top_k
+ self._top_p = top_p
+ self._early_stopping_token = early_stopping_token
+
+ @property
+ def backbones(self) -> Sequence[tf_keras.Model]:
+ return self._backbones
+
+ @property
+ def transformer(self) -> tf_keras.Model:
+ return self._transformer # pyrefly: ignore[bad-return]
+
+ def get_config(self):
+ config = {
+ "max_seq_len": self._max_seq_len,
+ "vocab_size": self._vocab_size,
+ "hidden_size": self._hidden_size,
+ "num_encoder_layers": self._num_encoder_layers,
+ "num_decoder_layers": self._num_decoder_layers,
+ "drop_path": self._drop_path,
+ "drop_units": self._drop_units,
+ "drop_att": self._drop_att,
+ "temperature": self._temperature,
+ "top_k": self._top_k,
+ "top_p": self._top_p,
+ "early_stopping_token": self._early_stopping_token,
+ "num_heads": self._num_heads,
+ }
+ config["backbone"] = self._backbones[0]
+ config["backbone_endpoint_name"] = self._backbone_endpoint_names[0]
+ for i in range(1, len(self._backbones)):
+ config[f"backbone_{i+1}"] = self._backbones[i]
+ config[f"backbone_endpoint_name_{i+1}"] = self._backbone_endpoint_names[i]
+ return config
+
+ @classmethod
+ def from_config(cls, config): # pyrefly: ignore[bad-override]
+ return cls(**config)
+
+ @property
+ def checkpoint_items(
+ self,
+ ) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ # For backward-compatibility with prior checkpoints, the first backbone
+ # should be named "backbone" and the second one should be named
+ # "backbone_2", etc.
+ items = dict(
+ backbone=self.backbones[0],
+ transformer=self.transformer,
+ stem_projection=self._stem_projections[0],
+ stem_ln=self._stem_lns[0],
+ )
+ for i in range(1, len(self.backbones)):
+ items[f"backbone_{i+1}"] = self.backbones[i]
+ items[f"stem_projection_{i+1}"] = self._stem_projections[i]
+ items[f"stem_ln_{i+1}"] = self._stem_lns[i]
+
+ return items
+
+ def _generate_image_mask(
+ self, inputs: tf.Tensor, target_shape: tf.Tensor
+ ) -> tf.Tensor:
+ """Generates image mask from input image."""
+ mask = tf.expand_dims(
+ tf.cast(
+ tf.not_equal(tf.reduce_sum(inputs, axis=-1), 0.3), inputs.dtype
+ ),
+ axis=-1,
+ )
+ mask = tf.image.resize(
+ mask, target_shape, method=tf.image.ResizeMethod.NEAREST_NEIGHBOR
+ )
+ return mask
+
+ def call( # pytype: disable=annotation-type-mismatch
+ self,
+ inputs: tf.Tensor,
+ targets: Optional[tf.Tensor] = None,
+ training: bool = None, # pyrefly: ignore[bad-function-definition]
+ use_teacher_forcing_for_eval: bool = False,
+ use_input_as_backbone_features=False,
+ ) -> List[Any]:
+ transformer_inputs = {
+ "tokens": targets,
+ "inputs": [], # List of [B, H*W, C] tensors, one per image modality.
+ "pos_emb": [], # List of positional embeddings for each image modality.
+ }
+ # Inputs has shape [B, N, H, W, C] where N is the number of images.
+ for i in range(len(self.backbones)):
+ inputs_i = inputs[:, i, :, :, :]
+ if use_input_as_backbone_features:
+ features = inputs_i
+ else:
+ features = self._backbones[i](inputs_i)[
+ self._backbone_endpoint_names[i]
+ ]
+ mask = tf.ones_like(features)
+ batch_size, h, w, num_channels = get_shape(features)
+ features = tf.reshape(features, [batch_size, h * w, num_channels])
+ features = self._stem_lns[i](
+ self._stem_projections[i](
+ self._encoder_dropouts[i](features, training)
+ )
+ )
+
+ pos_emb = position_embedding_sine(
+ mask[:, :, :, 0], num_pos_features=self._hidden_size
+ )
+ pos_emb = tf.reshape(pos_emb, [batch_size, -1, self._hidden_size])
+ pos_emb = tf.cast(pos_emb, features.dtype)
+ transformer_inputs["inputs"].append(features)
+ transformer_inputs["pos_emb"].append(pos_emb)
+
+ tokens = None
+ if training:
+ logits = self._transformer(transformer_inputs, training=True)
+ elif use_teacher_forcing_for_eval:
+ logits = self._transformer(transformer_inputs, training=False)
+ else:
+ tokens, logits = self._transformer.infer(
+ transformer_inputs, # pyrefly: ignore[bad-argument-type]
+ temperature=self._temperature,
+ top_k=self._top_k,
+ top_p=self._top_p,
+ early_stopping_token=self._early_stopping_token,
+ )
+
+ return [tokens, logits]
+
+
+def _create_cond_fn(
+ seq_len: int, early_stopping_token: int | None, prompt_len: int
+):
+ """Returns a loop condition for decoder.
+
+ Args:
+ seq_len: the maximum sequence length.
+ early_stopping_token: if not None, enable early termination based on this
+ token.
+ prompt_len: the length of prompt sequence.
+ """
+
+ def cond(step, caches, tokens, logits):
+ del caches
+ del logits
+ within_seq_len = (seq_len > prompt_len) & (step < seq_len - 1)
+ if early_stopping_token is None:
+ return within_seq_len
+ else:
+ tokens = tokens[prompt_len:step]
+ reached_early_stopping = tf.reduce_all(
+ tf.reduce_any(tokens == early_stopping_token, axis=0)
+ )
+ return within_seq_len & tf.logical_not(reached_early_stopping)
+
+ return cond
+
+
+class Pix2SeqTransformer(tf_keras.layers.Layer):
+ """Encoder and Decoder of Pix2Seq."""
+
+ def __init__(
+ self,
+ max_seq_len,
+ vocab_size,
+ hidden_size,
+ num_sources,
+ pos_encoding="learned",
+ num_encoder_layers=6,
+ num_decoder_layers=6,
+ drop_path=0.1,
+ drop_units=0.1,
+ drop_att=0.0,
+ output_bias=True,
+ num_heads=8,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self._max_seq_len = max_seq_len
+ self._vocab_size = vocab_size
+ self._hidden_size = hidden_size
+ self._num_sources = num_sources
+ self._pos_encoding = pos_encoding
+ self._num_encoder_layers = num_encoder_layers
+ self._num_decoder_layers = num_decoder_layers
+ self._drop_path = drop_path
+ self._drop_units = drop_units
+ self._drop_att = drop_att
+ self._output_bias = output_bias
+ self._num_heads = num_heads
+
+ add_seq_pos_emb(
+ self, self._pos_encoding, self._max_seq_len, self._hidden_size
+ )
+ add_vocab_token_emb(
+ self,
+ self._vocab_size,
+ self._hidden_size,
+ self._output_bias,
+ )
+
+ if self._num_encoder_layers > 0:
+ self._encoders = [
+ transformer.TransformerEncoder(
+ num_layers=self._num_encoder_layers,
+ dim=self._hidden_size,
+ mlp_ratio=4,
+ num_heads=self._num_heads,
+ drop_path=self._drop_path,
+ drop_units=self._drop_units,
+ drop_att=self._drop_att,
+ )
+ for _ in range(self._num_sources)
+ ]
+ else:
+ self._encoders = None
+
+ self._output_ln_encs = [
+ tf_keras.layers.LayerNormalization(epsilon=1e-6, name="output_ln_enc")
+ for _ in range(self._num_sources)
+ ]
+
+ self._projs = [
+ tf_keras.layers.Dense(self._hidden_size, name="proj/linear")
+ for _ in range(self._num_sources)
+ ]
+ self._proj_lns = [
+ tf_keras.layers.LayerNormalization(epsilon=1e-6, name="proj/ln")
+ for _ in range(self._num_sources)
+ ]
+ self._proj_mlps = [
+ transformer.MLP(
+ num_layers=1,
+ dim=self._hidden_size,
+ mlp_ratio=4,
+ drop_path=self._drop_path,
+ drop_units=self._drop_units,
+ name="proj/mlp",
+ )
+ for _ in range(self._num_sources)
+ ]
+
+ self._decoder = transformer.TransformerDecoder(
+ num_layers=self._num_decoder_layers,
+ dim=self._hidden_size,
+ mlp_ratio=4,
+ num_heads=self._num_heads,
+ drop_path=self._drop_path,
+ drop_units=self._drop_units,
+ drop_att=self._drop_att,
+ )
+ self._output_ln_dec = tf_keras.layers.LayerNormalization(
+ epsilon=1e-6, name="output_ln_dec"
+ )
+
+ def get_config(self):
+ return {
+ "max_seq_len": self._max_seq_len,
+ "vocab_size": self._vocab_size,
+ "hidden_size": self._hidden_size,
+ "pos_encoding": self._pos_encoding,
+ "num_encoder_layers": self._num_encoder_layers,
+ "num_decoder_layers": self._num_decoder_layers,
+ "drop_path": self._drop_path,
+ "drop_units": self._drop_units,
+ "drop_att": self._drop_att,
+ "output_bias": self._output_bias,
+ "num_heads": self._num_heads,
+ }
+
+ def encode_sources(
+ self,
+ sources: Sequence[tf.Tensor],
+ mem_pos_embeds: Sequence[tf.Tensor],
+ training: bool,
+ ):
+ """Encodes and concatenates sources for the decoder."""
+ encoded_sources = []
+ for i in range(self._num_sources):
+ source = sources[i]
+ mem_pos_embed = mem_pos_embeds[i]
+ source = source + mem_pos_embed
+ if self._encoders is not None:
+ encoded = self._encoders[i](
+ source, None, training=training, ret_list=False
+ )
+ else:
+ encoded = source
+
+ encoded = self._output_ln_encs[i](encoded)
+ encoded = self._proj_lns[i](self._projs[i](encoded))
+ encoded = encoded + mem_pos_embed
+ encoded = self._proj_mlps[i](encoded, training=training)
+ encoded_sources.append(encoded)
+
+ # encoded_sources is of length N, each item having shape
+ # [B, H*W, self._hidden_size]. Reshape to [B, N*H*W, self._hidden_size]
+ # before passing to decoder.
+ return tf.concat(encoded_sources, axis=1)
+
+ def call(self, inputs: dict[str, tf.Tensor], training: bool = None): # pytype: disable=annotation-type-mismatch
+ encoded = self.encode_sources(inputs["inputs"], inputs["pos_emb"], training) # pyrefly: ignore[bad-argument-type]
+
+ targets = inputs["tokens"]
+ seq_len = tf.shape(targets)[1]
+ seq_pos_emb = tf.expand_dims(self.seq_pos_emb[:seq_len], 0)
+ inp_embedding = outp_embedding = self.token_embedding
+ target_emb = tf.gather(inp_embedding, targets) + seq_pos_emb
+
+ self_attention_mask = 1.0 - get_ar_mask(seq_len, target_emb.dtype)
+
+ decoded, _ = self._decoder(
+ target_emb, encoded, None, self_attention_mask, None, training
+ )
+ decoded = self._output_ln_dec(decoded)
+
+ decoded = tf.cast(decoded, seq_pos_emb.dtype)
+ outp_embedding = tf.cast(outp_embedding, seq_pos_emb.dtype)
+
+ logits = tf.matmul(decoded, outp_embedding, transpose_b=True)
+ if self._output_bias:
+ logits = tf.nn.bias_add(logits, self.outp_bias)
+
+ return logits
+
+ def infer(
+ self,
+ inputs: tf.Tensor,
+ max_seq_len=None,
+ temperature=1.0,
+ top_k=0,
+ top_p=0.4,
+ sampling_callback=None,
+ early_stopping_token: int | None = None,
+ ):
+ """Autoregressive (without teacher-forcing) prediction.
+
+ Note: the autoregressive sampling/inference time can be further optimized by
+ caching *transformed* key / value inside multi-head attention for the
+ `encoded` and previously generated tokens, but this may make the code less
+ readable.
+
+ Args:
+ inputs: prompt - `int` tokens with shape of (bsz, prompt_len). encoded -
+ `float` encoded representations for conditioning with shape of (bsz,
+ size, dim). This can be optional in case of pure decoder.
+ max_seq_len: `int` of max generated sequence length (including prompt).
+ temperature: `float` scalar for scaling the logits before sampling.
+ top_k: `int` scalar for truncating top-k tokens according to logits before
+ token sampling.
+ top_p: `float` scalar specifying the threshold of cumulative probability
+ for truncating tokens before token sampling.
+ sampling_callback: a callbak `function` that take `next_logits`, and
+ return `next_token`. This is used when users need a specific logic for
+ sampling. Default to `None` with standard free-form sampling.
+ early_stopping_token: if not None, stop inference early based on this
+ token. This won't change sequence length, however. For each sequence,
+ the tokens after the early stopping token will be filled with the early
+ stopping token and logit values will have undefined behavior based on
+ implementation detail.
+
+ Returns:
+ sampled tokens with shape of (bsz, max_seq_len-prompt_len).
+ logits (temperature-scaled) associated with sampled token, in shape of
+ (bsz, max_seq_len-prompt_len, vocab_size).
+ """
+ encoded = self.encode_sources(
+ inputs["inputs"], inputs["pos_emb"], training=False
+ )
+ prompt = inputs["tokens"]
+ bsz = tf.shape(prompt)[0]
+ prompt_len = tf.shape(prompt)[1]
+
+ seq_len = self._max_seq_len if max_seq_len is None else max_seq_len
+ # (gunho) 500 (self._max_seq_len) -> 501 for prompt seq
+ seq_len = seq_len + 1
+ seq_pos_emb = tf.expand_dims(self.seq_pos_emb, 0)
+ inp_embedding = self.token_embedding
+ outp_embedding = inp_embedding
+
+ # Each step reads caches[:step] and tokens[step:next_step] and updates
+ # tokens[next_step], logits[next_step] and caches[step:next_step].
+ # On the first step, step=0, next_step=prompt_len. On subsequent steps
+ # next_step = step + 1.
+ def loop_body(step, caches, tokens, logits, is_prompt=False):
+ if is_prompt:
+ assert step == 0
+ x = tf.gather(inp_embedding, tf.transpose(tokens[:prompt_len]))
+ input_pos_embed = seq_pos_emb[:, :prompt_len]
+ x += input_pos_embed
+ self_attention_mask = 1.0 - get_ar_mask(prompt_len, x.dtype)
+ caches_in = None
+ else:
+ x = tf.gather(inp_embedding, tf.transpose(tokens[step]))
+ input_pos_embed = seq_pos_emb[:, step]
+ x += input_pos_embed
+ x = tf.expand_dims(x, 1) # (bsz, 1, d)
+ self_attention_mask = tf.ones([1, 1, 1, 1])
+ caches_in = tf.transpose(caches[:step], [1, 2, 0, 3])
+ decoded, caches_out = self._decoder(
+ x, encoded, caches_in, self_attention_mask, None, training=False
+ )
+ decoded = self._output_ln_dec(decoded)
+
+ # (gunho) transformer.py uses tf.float32 for numeric stability.
+ decoded = tf.cast(decoded, seq_pos_emb.dtype)
+
+ next_logits = tf.matmul( # only take the last for sampling next token.
+ decoded, outp_embedding, transpose_b=True
+ )[:, -1]
+ if self._output_bias:
+ next_logits = tf.nn.bias_add(next_logits, self.outp_bias)
+
+ # Scale and trunctate logits and sample next token.
+ if sampling_callback:
+ next_token = sampling_callback(
+ next_logits, step, temperature, top_k, top_p
+ )
+ else:
+ sampling_logits = next_logits / tf.cast(temperature, tf.float32)
+ sampling_logits = top_logits(sampling_logits, k=top_k, p=top_p)
+ next_token = tf.random.categorical(
+ sampling_logits, num_samples=1, dtype=tf.int32
+ )[:, 0]
+
+ # Update internal states.
+ next_step = step + (prompt_len if is_prompt else 1)
+ caches_out = tf.transpose(caches_out, [2, 0, 1, 3])
+ if is_prompt:
+ caches = tf.tensor_scatter_nd_update(
+ caches,
+ tf.range(prompt_len)[:, tf.newaxis],
+ caches_out,
+ )
+ else:
+ caches = tf.tensor_scatter_nd_update(caches, [[step]], caches_out)
+ tokens = tf.tensor_scatter_nd_update(tokens, [[next_step]], [next_token])
+ logits = tf.tensor_scatter_nd_update(logits, [[next_step]], [next_logits])
+ return (next_step, caches, tokens, logits)
+
+ caches_var = tf.zeros(
+ [seq_len - 1, self._num_decoder_layers, bsz, self._hidden_size]
+ )
+ tokens_var = tf.zeros([seq_len, bsz], dtype=tf.int64)
+ logits_var = tf.zeros([seq_len, bsz, self._vocab_size], dtype=tf.float32)
+ indices = tf.expand_dims(tf.range(prompt_len), -1)
+ tokens_var = tf.tensor_scatter_nd_update(
+ tokens_var, indices, tf.transpose(prompt, [1, 0])
+ )
+
+ step = 0
+ step, caches_var, tokens_var, logits_var = loop_body(
+ step, caches_var, tokens_var, logits_var, is_prompt=True
+ )
+ step, _, tokens_var, logits_var = tf.while_loop(
+ cond=_create_cond_fn(
+ seq_len=seq_len,
+ early_stopping_token=early_stopping_token,
+ prompt_len=prompt_len,
+ ),
+ body=loop_body,
+ loop_vars=[step, caches_var, tokens_var, logits_var],
+ )
+
+ # If stopping early based on early_stopping_token, assign
+ # early_stopping_token to all tokens after stopping occurs.
+ if early_stopping_token is not None:
+ tokens_var = tf.where(
+ tf.range(seq_len)[:, tf.newaxis] >= step,
+ tf.cast(early_stopping_token, tokens_var.dtype),
+ tokens_var,
+ )
+
+ sampled_tokens = tf.transpose(tokens_var[prompt_len:], [1, 0])
+ sampled_tokens_logits = tf.transpose(logits_var[prompt_len:], [1, 0, 2])
+ return sampled_tokens, sampled_tokens_logits
diff --git a/official/projects/pix2seq/modeling/pix2seq_model_test.py b/official/projects/pix2seq/modeling/pix2seq_model_test.py
new file mode 100644
index 00000000000..a76dcae17bd
--- /dev/null
+++ b/official/projects/pix2seq/modeling/pix2seq_model_test.py
@@ -0,0 +1,264 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for Pix2Seq model."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.projects.pix2seq.modeling import pix2seq_model
+from official.vision.modeling.backbones import resnet
+
+
+class Pix2SeqTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ ('One backbone', 1),
+ ('Two backbones', 2),
+ )
+ def test_forward(self, num_backbones: int):
+ hidden_size = 256
+ num_heads = 8
+ max_seq_len = 50
+ vocab_size = 164
+ image_size = 224
+ batch_size = 2
+ backbones = [
+ resnet.ResNet(50, bn_trainable=False) for _ in range(num_backbones)
+ ]
+ backbone_endpoint_names = ['5' for _ in range(num_backbones)]
+ model = pix2seq_model.Pix2Seq(
+ backbones,
+ backbone_endpoint_names,
+ max_seq_len,
+ vocab_size,
+ hidden_size,
+ num_heads=num_heads,
+ encoded_feature_dropout_rates=[0.1] * num_backbones,
+ )
+ _, outs = model(
+ tf.ones((batch_size, num_backbones, image_size, image_size, 3)),
+ tf.ones((batch_size, max_seq_len), tf.int64),
+ True,
+ )
+
+ self.assertLen(outs, 2) # intermediate decoded outputs.
+
+ @parameterized.named_parameters(
+ ('One backbone', 1),
+ ('Two backbones', 2),
+ )
+ def test_forward_infer_teacher_forcing(self, num_backbones: int):
+ hidden_size = 256
+ num_heads = 8
+ max_seq_len = 50
+ vocab_size = 164
+ image_size = 224
+ batch_size = 2
+ backbones = [
+ resnet.ResNet(50, bn_trainable=False) for _ in range(num_backbones)
+ ]
+ backbone_endpoint_names = ['5' for _ in range(num_backbones)]
+ model = pix2seq_model.Pix2Seq(
+ backbones,
+ backbone_endpoint_names,
+ max_seq_len,
+ vocab_size,
+ hidden_size,
+ num_heads=num_heads,
+ encoded_feature_dropout_rates=[0.1] * num_backbones,
+ )
+ _, outs = model(
+ tf.ones((batch_size, num_backbones, image_size, image_size, 3)),
+ tf.ones((batch_size, max_seq_len), tf.int64),
+ training=False,
+ use_teacher_forcing_for_eval=True,
+ )
+
+ self.assertLen(outs, 2) # intermediate decoded outputs.
+
+ @parameterized.named_parameters(
+ ('One backbone', 1),
+ ('Two backbones', 2),
+ )
+ def test_forward_infer(self, num_backbones: int):
+ hidden_size = 256
+ num_heads = 8
+ max_seq_len = 50
+ vocab_size = 600
+ image_size = 640
+ batch_size = 2
+ backbones = [
+ resnet.ResNet(50, bn_trainable=False) for _ in range(num_backbones)
+ ]
+ backbone_endpoint_names = ['5' for _ in range(num_backbones)]
+ model = pix2seq_model.Pix2Seq(
+ backbones,
+ backbone_endpoint_names,
+ max_seq_len,
+ vocab_size,
+ hidden_size,
+ num_heads=num_heads,
+ encoded_feature_dropout_rates=[0.1] * num_backbones,
+ )
+ tokens, _ = model(
+ tf.ones((batch_size, num_backbones, image_size, image_size, 3)),
+ tf.ones((batch_size, 1), tf.int64) * 10,
+ False,
+ )
+
+ self.assertLen(tokens, 2) # intermediate decoded outputs.
+
+ def test_forward_infer_with_early_stopping(self):
+ hidden_size = 256
+ num_heads = 8
+ max_seq_len = 50
+ vocab_size = 600
+ image_size = 640
+ batch_size = 2
+ backbone = resnet.ResNet(50, bn_trainable=False)
+ backbone_endpoint_names = ['5']
+ model = pix2seq_model.Pix2Seq(
+ [backbone],
+ backbone_endpoint_names,
+ max_seq_len,
+ vocab_size,
+ hidden_size,
+ num_heads=num_heads,
+ early_stopping_token=0,
+ )
+ tokens, _ = model(
+ tf.ones((batch_size, 1, image_size, image_size, 3)),
+ tf.ones((batch_size, 1), tf.int64) * 10,
+ False,
+ )
+
+ self.assertLen(tokens, 2) # intermediate decoded outputs.
+
+ def test_forward_infer_with_long_prompt(self):
+ hidden_size = 256
+ num_heads = 8
+ max_seq_len = 50
+ vocab_size = 600
+ image_size = 640
+ batch_size = 2
+ backbone = resnet.ResNet(50, bn_trainable=False)
+ backbone_endpoint_names = ['5']
+ model = pix2seq_model.Pix2Seq(
+ [backbone],
+ backbone_endpoint_names,
+ max_seq_len,
+ vocab_size,
+ hidden_size,
+ num_heads=num_heads,
+ )
+ tokens, _ = model(
+ tf.ones((batch_size, 1, image_size, image_size, 3)),
+ tf.ones((batch_size, 2), tf.int64) * 10,
+ False,
+ )
+
+ self.assertLen(tokens, 2) # intermediate decoded outputs.
+ self.assertShapeEqual(tokens, np.ndarray([batch_size, max_seq_len - 2 + 1]))
+
+ def test_cond_fn_without_early_stopping(self):
+ tokens = tf.constant(
+ # pyformat: disable
+ [
+ [0, 0, 0],
+ [0, 0, 0],
+ [0, 1, 0],
+ [1, 0, 0],
+ [0, 0, 1], # Should not stop early.
+ [0, 0, 0],
+ [0, 0, 0], # Should stop inference here.
+ ],
+ # pyformat: enable
+ dtype=tf.int64
+ )
+ cond = pix2seq_model._create_cond_fn(
+ seq_len=tokens.shape[0],
+ early_stopping_token=None,
+ prompt_len=1,
+ )
+ expected_results = [True, True, True, True, True, True, False]
+
+ self.assertLen(expected_results, tokens.shape[0])
+ for step, expected_result in enumerate(expected_results):
+ self.assertEqual(
+ expected_result,
+ cond(step, None, tokens, None),
+ msg=f'step={step}',
+ )
+
+ def test_cond_fn_with_early_stopping(self):
+ tokens = tf.constant(
+ # pyformat: disable
+ [
+ [0, 0, 0],
+ [0, 0, 0],
+ [0, 1, 0],
+ [1, 0, 0],
+ [0, 0, 1], # Should stop inference here.
+ [0, 0, 0],
+ [0, 0, 0],
+ ],
+ # pyformat: enable
+ dtype=tf.int64
+ )
+ cond = pix2seq_model._create_cond_fn(
+ seq_len=tokens.shape[0],
+ early_stopping_token=1,
+ prompt_len=1,
+ )
+ expected_results = [True, True, True, True, True, False, False]
+
+ self.assertLen(expected_results, tokens.shape[0])
+ for step, expected_result in enumerate(expected_results):
+ self.assertEqual(
+ expected_result,
+ cond(step, None, tokens, None),
+ msg=f'step={step}',
+ )
+
+ def test_cond_fn_with_early_stopping_keep_inference_to_end(self):
+ tokens = tf.constant(
+ # pyformat: disable
+ [
+ [1, 1, 1], # Early stopping token within prompt should be ignored.
+ [0, 0, 0],
+ [0, 1, 0],
+ [1, 0, 0], # Should keep inferencing until the end.
+ ],
+ # pyformat: enable
+ dtype=tf.int64
+ )
+ cond = pix2seq_model._create_cond_fn(
+ seq_len=tokens.shape[0],
+ early_stopping_token=1,
+ prompt_len=1,
+ )
+ expected_results = [True, True, True, False]
+
+ self.assertLen(expected_results, tokens.shape[0])
+ for step, expected_result in enumerate(expected_results):
+ self.assertEqual(
+ expected_result,
+ cond(step, None, tokens, None),
+ msg=f'step={step}',
+ )
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pix2seq/modeling/transformer.py b/official/projects/pix2seq/modeling/transformer.py
new file mode 100644
index 00000000000..8d392e0e187
--- /dev/null
+++ b/official/projects/pix2seq/modeling/transformer.py
@@ -0,0 +1,536 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Specialized Transformers for Pix2Seq.
+
+the position embeddings are added to the query and key for every self- and
+cross-attention layer.
+"""
+
+import tensorflow as tf, tf_keras
+
+
+class TransformerEncoder(tf_keras.layers.Layer):
+ """Transformer encoder."""
+
+ def __init__(
+ self,
+ num_layers,
+ dim,
+ mlp_ratio,
+ num_heads,
+ drop_path=0.1,
+ drop_units=0.1,
+ drop_att=0.0,
+ self_attention=True,
+ use_ffn_ln=False,
+ ln_scale_shift=True,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self._num_layers = num_layers
+ self._dim = dim
+ self._mlp_ratio = mlp_ratio
+ self._num_heads = num_heads
+ self._drop_path = drop_path
+ self._drop_units = drop_units
+ self._drop_att = drop_att
+ self._self_attention = self_attention
+ self._use_ffn_ln = use_ffn_ln
+ self._ln_scale_shift = ln_scale_shift
+
+ self.enc_layers = [
+ TransformerEncoderLayer( # pylint: disable=g-complex-comprehension
+ dim,
+ mlp_ratio,
+ num_heads,
+ drop_path,
+ drop_units,
+ drop_att,
+ self_attention=self_attention,
+ use_ffn_ln=use_ffn_ln,
+ ln_scale_shift=ln_scale_shift,
+ name='transformer_encoder' + suffix_id(i),
+ )
+ for i in range(num_layers)
+ ]
+
+ def call(self, x, mask, training, ret_list=False):
+ x_list = [x]
+ for i in range(self._num_layers):
+ x = self.enc_layers[i](x, mask, training)
+ x_list.append(x)
+ return (x, x_list) if ret_list else x
+
+ def get_config(self):
+ config = super().get_config()
+ updates = {
+ 'num_layers': self._num_layers,
+ 'dim': self._dim,
+ 'mlp_ratio': self._mlp_ratio,
+ 'num_heads': self._num_heads,
+ 'drop_path': self._drop_path,
+ 'drop_units': self._drop_units,
+ 'drop_att': self._drop_att,
+ 'self_attention': self._self_attention,
+ 'use_ffn_ln': self._use_ffn_ln,
+ 'ln_scale_shift': self._ln_scale_shift,
+ }
+ config.update(updates)
+ return config
+
+
+class TransformerEncoderLayer(tf_keras.layers.Layer): # pylint: disable=missing-docstring
+
+ def __init__(
+ self,
+ dim,
+ mlp_ratio,
+ num_heads,
+ drop_path=0.1,
+ drop_units=0.1,
+ drop_att=0.0,
+ self_attention=True,
+ use_ffn_ln=False,
+ ln_scale_shift=True,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self._dim = dim
+ self._mlp_ratio = mlp_ratio
+ self._num_heads = num_heads
+ self._drop_path = drop_path
+ self._drop_units = drop_units
+ self._drop_att = drop_att
+ self.self_attention = self_attention
+ self._use_ffn_ln = use_ffn_ln
+ self._ln_scale_shift = ln_scale_shift
+
+ if self_attention:
+ self.mha_ln = tf_keras.layers.LayerNormalization(
+ epsilon=1e-6,
+ center=ln_scale_shift,
+ scale=ln_scale_shift,
+ name='mha/ln',
+ )
+ self.mha = tf_keras.layers.MultiHeadAttention(
+ num_heads, dim // num_heads, dropout=drop_att, name='mha'
+ )
+ self.mlp = MLP(
+ 1,
+ dim,
+ mlp_ratio,
+ drop_path,
+ drop_units,
+ use_ffn_ln=use_ffn_ln,
+ ln_scale_shift=ln_scale_shift,
+ name='mlp',
+ )
+ self.dropp = DropPath(drop_path)
+
+ def call(self, x, mask, training):
+ # x shape (bsz, seq_len, dim_att), mask shape (bsz, seq_len, seq_len).
+ if self.self_attention:
+ x_ln = self.mha_ln(x)
+ x_residual = self.mha(x_ln, x_ln, x_ln, mask, training=training)
+ x = x + self.dropp(x_residual, training)
+ x = self.mlp(x, training)
+ return x
+
+ def get_config(self):
+ config = super().get_config()
+ updates = {
+ 'dim': self._dim,
+ 'mlp_ratio': self._mlp_ratio,
+ 'num_heads': self._num_heads,
+ 'drop_path': self._drop_path,
+ 'drop_units': self._drop_units,
+ 'drop_att': self._drop_att,
+ 'self_attention': self._self_attention,
+ 'use_ffn_ln': self._use_ffn_ln,
+ 'ln_scale_shift': self._ln_scale_shift,
+ }
+ config.update(updates)
+ return config
+
+
+def suffix_id(i):
+ """Return suffix id for layer/variable name."""
+ return '' if i == 0 else '_%d' % i
+
+
+class DropPath(tf_keras.layers.Layer):
+ """For stochastic depth."""
+
+ def __init__(self, drop_rate=0.0, **kwargs):
+ """Initializes a drop path layer."""
+ super().__init__(**kwargs)
+ self._drop_rate = drop_rate
+ if self._drop_rate < 0 or self._drop_rate >= 1.0:
+ raise ValueError('drop_rate {} is outside [0, 1)'.format(self._drop_rate))
+
+ def call(self, x, training=False):
+ """Performs a forward pass.
+
+ Args:
+ x: An input tensor of type tf.Tensor with shape [batch, height, width,
+ channels].
+ training: A boolean flag indicating whether training behavior should be
+ used (default: False).
+
+ Returns:
+ The output tensor.
+ """
+ if self._drop_rate == 0.0 or not training:
+ return x
+
+ keep_rate = 1.0 - self._drop_rate
+ xshape = tf.shape(x)
+ drop_mask_shape = [xshape[0]] + [1] * (len(xshape) - 1)
+ drop_mask = keep_rate + tf.random.uniform(drop_mask_shape, dtype=x.dtype)
+ drop_mask = tf.math.divide(tf.floor(drop_mask), keep_rate)
+ return x * drop_mask
+
+ def get_config(self):
+ config = super().get_config()
+ updates = {
+ 'drop_rate': self._drop_rate,
+ }
+ config.update(updates)
+ return config
+
+
+class FeedForwardLayer(tf_keras.layers.Layer): # pylint: disable=missing-docstring
+
+ def __init__(
+ self,
+ dim_att=256,
+ dim_mlp=1024,
+ drop_units=0.1,
+ use_ln=False,
+ ln_scale_shift=False,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self._dim_att = dim_att
+ self._dim_mlp = dim_mlp
+ self._drop_units = drop_units
+ self._use_ln = use_ln
+ self._ln_scale_shift = ln_scale_shift
+
+ self.dense1 = tf_keras.layers.Dense(
+ dim_mlp, activation=tf.nn.gelu, name='dense1'
+ )
+ self.dropout = tf_keras.layers.Dropout(drop_units)
+ self.dense2 = tf_keras.layers.Dense(dim_att, name='dense2')
+ if use_ln:
+ self.ln = tf_keras.layers.LayerNormalization(
+ epsilon=1e-6,
+ center=ln_scale_shift,
+ scale=ln_scale_shift,
+ name='mlp_ln',
+ )
+ else:
+ self.ln = lambda x: x
+
+ def call(self, x, training):
+ return self.dense2(self.dropout(self.ln(self.dense1(x)), training=training))
+
+ def get_config(self):
+ config = super().get_config()
+ updates = {
+ 'dim_att': self._dim_att,
+ 'dim_mlp': self._dim_mlp,
+ 'drop_units': self._drop_units,
+ 'use_ln': self._use_ln,
+ 'ln_scale_shift': self._ln_scale_shift,
+ }
+ config.update(updates)
+ return config
+
+
+class MLP(tf_keras.layers.Layer): # pylint: disable=missing-docstring
+
+ def __init__(
+ self,
+ num_layers,
+ dim,
+ mlp_ratio,
+ drop_path=0.1,
+ drop_units=0.0,
+ use_ffn_ln=False,
+ ln_scale_shift=True,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self._num_layers = num_layers
+ self._dim = dim
+ self._mlp_ratio = mlp_ratio
+ self._drop_path = drop_path
+ self._drop_units = drop_units
+ self._use_ffn_ln = use_ffn_ln
+ self._ln_scale_shift = ln_scale_shift
+
+ self.mlp_layers = []
+ self.layernorms = []
+ for i in range(num_layers):
+ self.mlp_layers.append(
+ FeedForwardLayer(
+ dim,
+ dim * mlp_ratio,
+ drop_units,
+ use_ln=use_ffn_ln,
+ ln_scale_shift=ln_scale_shift,
+ name='ffn' + suffix_id(i),
+ )
+ )
+ self.layernorms.append(
+ tf_keras.layers.LayerNormalization(
+ epsilon=1e-6,
+ center=ln_scale_shift,
+ scale=ln_scale_shift,
+ name='ffn/ln' + suffix_id(i),
+ )
+ )
+ self.dropp = DropPath(drop_path)
+
+ def call(self, x, training, ret_list=False):
+ x_list = [x]
+ for i in range(self._num_layers):
+ x_residual = self.mlp_layers[i](self.layernorms[i](x), training)
+ x = x + self.dropp(x_residual, training)
+ x_list.append(x)
+ return (x, x_list) if ret_list else x
+
+ def get_config(self):
+ config = super().get_config()
+ updates = {
+ 'num_layers': self._num_layers,
+ 'dim': self._dim,
+ 'mlp_ratio': self._mlp_ratio,
+ 'drop_path': self._drop_path,
+ 'drop_units': self._drop_units,
+ 'use_ffn_ln': self._use_ffn_ln,
+ 'ln_scale_shift': self._ln_scale_shift,
+ }
+ config.update(updates)
+ return config
+
+
+class TransformerDecoderLayer(tf_keras.layers.Layer): # pylint: disable=missing-docstring
+
+ def __init__(
+ self,
+ dim,
+ mlp_ratio,
+ num_heads,
+ drop_path=0.1,
+ drop_units=0.1,
+ drop_att=0.0,
+ dim_x_att=None,
+ self_attention=True,
+ cross_attention=True,
+ use_mlp=True,
+ use_enc_ln=False,
+ use_ffn_ln=False,
+ ln_scale_shift=True,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self._dim = dim
+ self._mlp_ratio = mlp_ratio
+ self._num_heads = num_heads
+ self._drop_path = drop_path
+ self._drop_units = drop_units
+ self._drop_att = drop_att
+ self._dim_x_att = dim_x_att
+ self._self_attention = self_attention
+ self._cross_attention = cross_attention
+ self._use_mlp = use_mlp
+ self._use_enc_ln = use_enc_ln
+ self._use_ffn_ln = use_ffn_ln
+ self._ln_scale_shift = ln_scale_shift
+
+ if self_attention:
+ self.self_ln = tf_keras.layers.LayerNormalization(
+ epsilon=1e-6,
+ center=ln_scale_shift,
+ scale=ln_scale_shift,
+ name='self_mha/ln',
+ )
+ self.self_mha = tf_keras.layers.MultiHeadAttention(
+ num_heads, dim // num_heads, dropout=drop_att, name='self_mha'
+ )
+ if cross_attention:
+ self.cross_ln = tf_keras.layers.LayerNormalization(
+ epsilon=1e-6,
+ center=ln_scale_shift,
+ scale=ln_scale_shift,
+ name='cross_mha/ln',
+ )
+ if use_enc_ln:
+ self.enc_ln = tf_keras.layers.LayerNormalization(
+ epsilon=1e-6,
+ center=ln_scale_shift,
+ scale=ln_scale_shift,
+ name='cross_mha/enc_ln',
+ )
+ else:
+ self.enc_ln = lambda x: x
+ dim_x_att = dim if dim_x_att is None else dim_x_att
+ self.cross_mha = tf_keras.layers.MultiHeadAttention(
+ num_heads, dim_x_att // num_heads, dropout=drop_att, name='cross_mha'
+ )
+ if use_mlp:
+ self.mlp = MLP(
+ 1,
+ dim,
+ mlp_ratio,
+ drop_path,
+ drop_units,
+ use_ffn_ln=use_ffn_ln,
+ ln_scale_shift=ln_scale_shift,
+ name='mlp',
+ )
+ self.dropp = DropPath(drop_path)
+
+ def call(self, x, enc, cache, mask_self, mask_cross, training):
+ """x in (bsz, seq, d), enc in (bsz, seq', d)."""
+ x_for_cache = []
+ if self._self_attention:
+ x_for_cache = x_ln = kv_ln = self.self_ln(x)
+ if cache is not None: # Augment kv_ln with cache in (bsz, c_size, d).
+ q_size, k_size = tf.shape(x)[1], tf.shape(cache)[1]
+ mask_self = tf.concat([tf.ones([1, 1, q_size, k_size]), mask_self], -1)
+ kv_ln = tf.concat([cache, x_ln], axis=1)
+ x_res = self.self_mha(x_ln, kv_ln, kv_ln, mask_self, training=training)
+ x = x + self.dropp(x_res, training)
+ if self._cross_attention:
+ x_ln = self.cross_ln(x)
+ enc = self.enc_ln(enc)
+ x_res = self.cross_mha(x_ln, enc, enc, mask_cross, training=training)
+ x = x + self.dropp(x_res, training)
+ if self._use_mlp:
+ x = self.mlp(x, training)
+ return x, x_for_cache
+
+ def get_config(self):
+ config = super().get_config()
+ updates = {
+ 'dim': self._dim,
+ 'mlp_ratio': self._mlp_ratio,
+ 'num_heads': self._num_heads,
+ 'drop_path': self._drop_path,
+ 'drop_units': self._drop_units,
+ 'drop_att': self._drop_att,
+ 'dim_x_att': self._dim_x_att,
+ 'self_attention': self._self_attention,
+ 'cross_attention': self._cross_attention,
+ 'use_mlp': self._use_mlp,
+ 'use_enc_ln': self._use_enc_ln,
+ 'use_ffn_ln': self._use_ffn_ln,
+ 'ln_scale_shift': self._ln_scale_shift,
+ }
+ config.update(updates)
+ return config
+
+
+class TransformerDecoder(tf_keras.layers.Layer): # pylint: disable=missing-docstring
+
+ def __init__(
+ self,
+ num_layers,
+ dim,
+ mlp_ratio,
+ num_heads,
+ drop_path=0.1,
+ drop_units=0.1,
+ drop_att=0.0,
+ dim_x_att=None,
+ self_attention=True,
+ cross_attention=True,
+ use_mlp=True,
+ use_enc_ln=False,
+ use_ffn_ln=False,
+ ln_scale_shift=True,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self._num_layers = num_layers
+ self._dim = dim
+ self._mlp_ratio = mlp_ratio
+ self._num_heads = num_heads
+ self._drop_path = drop_path
+ self._drop_units = drop_units
+ self._drop_att = drop_att
+ self._dim_x_att = dim_x_att
+ self._self_attention = self_attention
+ self._cross_attention = cross_attention
+ self._use_mlp = use_mlp
+ self._use_enc_ln = use_enc_ln
+ self._use_ffn_ln = use_ffn_ln
+ self._ln_scale_shift = ln_scale_shift
+
+ self.dec_layers = [
+ TransformerDecoderLayer( # pylint: disable=g-complex-comprehension
+ dim,
+ mlp_ratio,
+ num_heads,
+ drop_path,
+ drop_units,
+ drop_att,
+ dim_x_att=dim_x_att,
+ self_attention=self_attention,
+ cross_attention=cross_attention,
+ use_mlp=use_mlp,
+ use_enc_ln=use_enc_ln,
+ use_ffn_ln=use_ffn_ln,
+ ln_scale_shift=ln_scale_shift,
+ name='transformer_decoder_layer' + suffix_id(i),
+ )
+ for i in range(num_layers)
+ ]
+
+ def call(self, x, enc, caches, mask_self, mask_cross, training):
+ """x in (bsz, seq, d), enc in (bsz, seq', d)."""
+ presents = []
+ for i in range(self._num_layers):
+ cache = None if caches is None else caches[i]
+ x, x_for_cache = self.dec_layers[i](
+ x, enc, cache, mask_self, mask_cross, training
+ )
+ presents.append(x_for_cache)
+
+ return x, tf.stack(presents)
+
+ def get_config(self):
+ config = super().get_config()
+ updates = {
+ 'num_layers': self._num_layers,
+ 'dim': self._dim,
+ 'mlp_ratio': self._mlp_ratio,
+ 'num_heads': self._num_heads,
+ 'drop_path': self._drop_path,
+ 'drop_units': self._drop_units,
+ 'drop_att': self._drop_att,
+ 'dim_x_att': self._dim_x_att,
+ 'self_attention': self._self_attention,
+ 'cross_attention': self._cross_attention,
+ 'use_mlp': self._use_mlp,
+ 'use_enc_ln': self._use_enc_ln,
+ 'use_ffn_ln': self._use_ffn_ln,
+ 'ln_scale_shift': self._ln_scale_shift,
+ }
+ config.update(updates)
+ return config
diff --git a/official/projects/pix2seq/modeling/transformer_test.py b/official/projects/pix2seq/modeling/transformer_test.py
new file mode 100644
index 00000000000..11a84cad967
--- /dev/null
+++ b/official/projects/pix2seq/modeling/transformer_test.py
@@ -0,0 +1,182 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for transformer."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.pix2seq.modeling import transformer
+
+
+class TransformerTest(tf.test.TestCase):
+
+ def test_transformer_encoder(self):
+ batch_size = 2
+ sequence_length = 100
+ feature_size = 256
+ model = transformer.TransformerEncoder(
+ num_layers=3,
+ dim=feature_size,
+ mlp_ratio=4.0,
+ num_heads=2,
+ )
+ input_tensor = tf.ones((batch_size, sequence_length, feature_size))
+ out = model(input_tensor, mask=None, training=False)
+ self.assertAllEqual(
+ tf.shape(out), (batch_size, sequence_length, feature_size)
+ )
+
+ def test_transformer_encoder_get_config(self):
+ model = transformer.TransformerEncoder(
+ num_layers=2,
+ dim=256,
+ mlp_ratio=4.0,
+ num_heads=2,
+ )
+ config = model.get_config()
+ expected_config = {
+ 'name': 'transformer_encoder',
+ 'trainable': True,
+ 'dtype': 'float32',
+ 'num_layers': 2,
+ 'dim': 256,
+ 'mlp_ratio': 4.0,
+ 'num_heads': 2,
+ 'drop_path': 0.1,
+ 'drop_units': 0.1,
+ 'drop_att': 0.0,
+ 'self_attention': True,
+ 'use_ffn_ln': False,
+ 'ln_scale_shift': True,
+ }
+ self.assertAllEqual(expected_config, config)
+
+ def test_transformer_decoder_layer(self):
+ batch_size = 2
+ sequence_length = 100
+ memory_length = 200
+ feature_size = 256
+ model = transformer.TransformerDecoderLayer(
+ dim=feature_size,
+ mlp_ratio=4.0,
+ num_heads=2
+ )
+ input_tensor = tf.ones((batch_size, sequence_length, feature_size))
+ memory = tf.ones((batch_size, memory_length, feature_size))
+ self_attention_mask = tf.ones(
+ (batch_size, sequence_length, sequence_length), dtype=tf.int64
+ )
+
+ out, _ = model(
+ input_tensor,
+ memory,
+ None,
+ self_attention_mask,
+ None,
+ training=False,
+ )
+ self.assertAllEqual(
+ tf.shape(out), (batch_size, sequence_length, feature_size)
+ )
+
+ def test_transformer_decoder_layer_get_config(self):
+ model = transformer.TransformerDecoderLayer(
+ dim=256,
+ mlp_ratio=4.0,
+ num_heads=2
+ )
+ config = model.get_config()
+ expected_config = {
+ 'name': 'transformer_decoder_layer',
+ 'trainable': True,
+ 'dtype': 'float32',
+ 'dim': 256,
+ 'mlp_ratio': 4.0,
+ 'num_heads': 2,
+ 'drop_path': 0.1,
+ 'drop_units': 0.1,
+ 'drop_att': 0.0,
+ 'dim_x_att': None,
+ 'self_attention': True,
+ 'cross_attention': True,
+ 'use_mlp': True,
+ 'use_enc_ln': False,
+ 'use_ffn_ln': False,
+ 'ln_scale_shift': True,
+ }
+ self.assertAllEqual(expected_config, config)
+
+ def test_transformer_decoder(self):
+ batch_size = 2
+ sequence_length = 100
+ memory_length = 200
+ feature_size = 256
+ num_layers = 3
+ model = transformer.TransformerDecoder(
+ num_layers=num_layers,
+ dim=feature_size,
+ mlp_ratio=4.0,
+ num_heads=2,
+ )
+ input_tensor = tf.ones((batch_size, sequence_length, feature_size))
+ memory = tf.ones((batch_size, memory_length, feature_size))
+ self_attention_mask = tf.ones(
+ (batch_size, sequence_length, sequence_length), dtype=tf.int64
+ )
+
+ out, cache = model(
+ input_tensor, memory, None, self_attention_mask, None, training=False
+ )
+ self.assertAllEqual(
+ tf.shape(out), (batch_size, sequence_length, feature_size)
+ )
+ self.assertAllEqual(
+ tf.shape(cache), (num_layers, batch_size, sequence_length, feature_size)
+ )
+
+ def test_transformer_decoder_get_config(self):
+ num_layers = 2
+ num_attention_heads = 2
+ intermediate_size = 256
+ model = transformer.TransformerDecoder(
+ num_layers=num_layers,
+ dim=intermediate_size,
+ mlp_ratio=4.0,
+ num_heads=num_attention_heads,
+ )
+ config = model.get_config()
+ expected_config = {
+ 'name': 'transformer_decoder',
+ 'trainable': True,
+ 'dtype': 'float32',
+ 'num_layers': 2,
+ 'dim': 256,
+ 'mlp_ratio': 4.0,
+ 'num_heads': 2,
+ 'drop_path': 0.1,
+ 'drop_units': 0.1,
+ 'drop_att': 0.0,
+ 'dim_x_att': None,
+ 'self_attention': True,
+ 'cross_attention': True,
+ 'use_mlp': True,
+ 'use_enc_ln': False,
+ 'use_ffn_ln': False,
+ 'ln_scale_shift': True,
+ }
+ self.assertAllEqual(expected_config, config)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pix2seq/tasks/pix2seq_task.py b/official/projects/pix2seq/tasks/pix2seq_task.py
new file mode 100644
index 00000000000..99f0a0facaa
--- /dev/null
+++ b/official/projects/pix2seq/tasks/pix2seq_task.py
@@ -0,0 +1,372 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pix2Seq detection task definition."""
+
+from typing import Optional
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import base_task
+from official.core import task_factory
+from official.projects.pix2seq import utils
+from official.projects.pix2seq.configs import pix2seq as pix2seq_cfg
+from official.projects.pix2seq.dataloaders import pix2seq_input
+from official.projects.pix2seq.modeling import pix2seq_model
+from official.projects.uvit.modeling import vit # pylint: disable=unused-import
+from official.vision.dataloaders import input_reader_factory
+from official.vision.dataloaders import tf_example_decoder
+from official.vision.dataloaders import tfds_factory
+from official.vision.dataloaders import tf_example_label_map_decoder
+from official.vision.evaluation import coco_evaluator
+from official.vision.modeling import backbones as backbones_lib
+
+
+@task_factory.register_task_cls(pix2seq_cfg.Pix2SeqTask)
+class Pix2SeqTask(base_task.Task):
+ """A single-replica view of training procedure.
+
+ Pix2Seq task provides artifacts for training/evalution procedures, including
+ loading/iterating over Datasets, initializing the model, calculating the loss,
+ post-processing, and customized metrics with reduction.
+ """
+
+ def _build_backbones_and_endpoint_names(
+ self,
+ ) -> tuple[list[tf_keras.Model], list[str]]:
+ """Build backbones and returns their corresponding endpoint names."""
+ config: pix2seq_cfg.Pix2Seq = self._task_config.model
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + config.input_size
+ )
+ backbones = []
+ endpoint_names = []
+ for backbone_config in config.backbones:
+ backbone = backbones_lib.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=backbone_config.backbone,
+ norm_activation_config=backbone_config.norm_activation,
+ )
+ backbone.trainable = not backbone_config.freeze
+ backbones.append(backbone)
+ endpoint_names.append(backbone_config.endpoint_name)
+ return backbones, endpoint_names
+
+ def build_model(self):
+ """Build Pix2Seq model."""
+ config: pix2seq_cfg.Pix2Seq = self._task_config.model
+ backbones, endpoint_names = self._build_backbones_and_endpoint_names()
+ model = pix2seq_model.Pix2Seq(
+ backbones=backbones,
+ backbone_endpoint_name=endpoint_names,
+ max_seq_len=config.max_num_instances * 5,
+ vocab_size=config.vocab_size,
+ hidden_size=config.hidden_size,
+ num_encoder_layers=config.num_encoder_layers,
+ num_decoder_layers=config.num_decoder_layers,
+ drop_path=config.drop_path,
+ encoded_feature_dropout_rates=config.encoded_feature_dropout_rates,
+ drop_units=config.drop_units,
+ drop_att=config.drop_att,
+ num_heads=config.num_heads,
+ temperature=config.temperature,
+ top_p=config.top_p,
+ top_k=config.top_k,
+ early_stopping_token=config.early_stopping_token,
+ )
+ return model
+
+ def _get_ckpt(self, ckpt_dir_or_file: str) -> str:
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ return tf.train.latest_checkpoint(ckpt_dir_or_file)
+ return ckpt_dir_or_file
+
+ def initialize(self, model: tf_keras.Model):
+ """Loading pretrained checkpoint."""
+ if self._task_config.init_checkpoint_modules == 'backbone':
+ raise ValueError(
+ 'init_checkpoint_modules=backbone is no longer supported. Specify'
+ ' backbone checkpoints in each backbone config.'
+ )
+
+ if self._task_config.init_checkpoint_modules not in ['all', 'partial', '']:
+ raise ValueError(
+ 'Unsupported init_checkpoint_modules: '
+ f'{self._task_config.init_checkpoint_modules}'
+ )
+
+ if self._task_config.init_checkpoint and any(
+ [b.init_checkpoint for b in self._task_config.model.backbones]
+ ):
+ raise ValueError(
+ 'A global init_checkpoint and a backbone init_checkpoint cannot be'
+ ' specified at the same time.'
+ )
+
+ if self._task_config.init_checkpoint:
+ global_ckpt_file = self._get_ckpt(self._task_config.init_checkpoint)
+ ckpt = tf.train.Checkpoint(**model.checkpoint_items)
+ status = ckpt.restore(global_ckpt_file).expect_partial()
+ if self._task_config.init_checkpoint_modules != 'partial':
+ status.assert_existing_objects_matched()
+ logging.info(
+ 'Finished loading pretrained checkpoint from %s', global_ckpt_file
+ )
+ else:
+ # This case means that no global checkpoint was provided. Possibly,
+ # backbone-specific checkpoints were.
+ for backbone_config, backbone in zip(
+ self._task_config.model.backbones, model.backbones
+ ):
+ if not backbone_config.init_checkpoint:
+ continue
+
+ backbone_init_ckpt = self._get_ckpt(backbone_config.init_checkpoint)
+ if backbone_config.backbone.type == 'uvit':
+ # The UVit object has a special function called load_checkpoint.
+ # The other backbones do not.
+ backbone.load_checkpoint(ckpt_filepath=backbone_init_ckpt)
+ else:
+ ckpt = tf.train.Checkpoint(backbone=backbone)
+ status = (
+ ckpt.restore(backbone_init_ckpt)
+ .expect_partial()
+ .assert_nontrivial_match()
+ )
+ if backbone_config.assert_existing_objects_matched:
+ status.assert_existing_objects_matched()
+
+ logging.info(
+ 'Finished loading pretrained backbone from %s', backbone_init_ckpt
+ )
+
+ def build_inputs(
+ self, params, input_context: Optional[tf.distribute.InputContext] = None
+ ):
+ """Build input dataset."""
+
+ if params.tfds_name:
+ decoder = tfds_factory.get_detection_decoder(params.tfds_name)
+ else:
+ decoder_cfg = params.decoder.get()
+ if params.decoder.type == 'simple_decoder':
+ decoder = tf_example_decoder.TfExampleDecoder(
+ regenerate_source_id=decoder_cfg.regenerate_source_id
+ )
+ elif params.decoder.type == 'label_map_decoder':
+ decoder = tf_example_label_map_decoder.TfExampleDecoderLabelMap(
+ label_map=decoder_cfg.label_map,
+ regenerate_source_id=decoder_cfg.regenerate_source_id,
+ )
+ else:
+ raise ValueError(
+ 'Unknown decoder type: {}!'.format(params.decoder.type)
+ )
+
+ parser = pix2seq_input.Parser(
+ eos_token_weight=self._task_config.losses.eos_token_weight,
+ output_size=self._task_config.model.input_size[:2],
+ max_num_boxes=self._task_config.model.max_num_instances,
+ coord_vocab_shift=self._task_config.coord_vocab_shift,
+ quantization_bins=self._task_config.quantization_bins,
+ aug_scale_min=params.aug_scale_min,
+ aug_scale_max=params.aug_scale_max,
+ aug_color_jitter_strength=params.aug_color_jitter_strength,
+ label_shift=params.label_shift,
+ )
+
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training),
+ )
+ dataset = reader.read(input_context=input_context)
+
+ return dataset
+
+ def build_losses(self, outputs, labels, aux_losses=None):
+ """Builds DETR losses."""
+ targets = labels['targets']
+ weights = labels['weights']
+
+ targets = tf.one_hot(targets, self._task_config.model.vocab_size)
+
+ loss = tf_keras.losses.CategoricalCrossentropy(
+ from_logits=True, reduction=tf_keras.losses.Reduction.NONE
+ )(targets, outputs)
+
+ weights = tf.cast(weights, loss.dtype)
+ loss = tf.reduce_sum(loss * weights) / tf.reduce_sum(weights)
+
+ aux_losses = tf.add_n(aux_losses) if aux_losses else 0.0
+
+ total_loss = loss + aux_losses
+ return total_loss
+
+ def build_metrics(self, training=True):
+ """Builds detection metrics."""
+ metrics = []
+ metric_names = ['loss']
+ for name in metric_names:
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
+
+ if not training:
+ self.coco_metric = coco_evaluator.COCOEvaluator(
+ annotation_file=self._task_config.annotation_file,
+ include_mask=False,
+ need_rescale_bboxes=False,
+ per_category_metrics=self._task_config.per_category_metrics,
+ )
+ return metrics
+
+ def train_step(self, inputs, model, optimizer, metrics=None):
+ """Does forward and backward.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ features, labels = inputs
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+
+ with tf.GradientTape() as tape:
+ _, outputs = model(features, labels['inputs'], training=True)
+ outputs = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), outputs)
+
+ loss = self.build_losses(
+ outputs=outputs, labels=labels, aux_losses=model.losses
+ )
+ scaled_loss = loss / num_replicas
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scales back gradient when LossScaleOptimizer is used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ # Trainer class handles loss metric for you.
+ logs = {self.loss: loss}
+
+ all_losses = {
+ 'loss': loss,
+ }
+
+ # Metric results will be added to logs for you.
+ if metrics:
+ for m in metrics:
+ m.update_state(all_losses[m.name])
+ return logs
+
+ def validation_step(self, inputs, model, metrics=None):
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ features, labels = inputs
+
+ tokens, logits = model(features, labels['prompt'], training=False)
+ # loss = self.build_losses(
+ # outputs=outputs, labels=labels, aux_losses=model.losses)
+ loss = 0.0
+
+ # Multiply for logging.
+ # Since we expect the gradient replica sum to happen in the optimizer,
+ # the loss is scaled with global num_boxes and weights.
+ # To have it more interpretable/comparable we scale it back when logging.
+ num_replicas_in_sync = tf.distribute.get_strategy().num_replicas_in_sync
+ loss *= num_replicas_in_sync
+
+ # Evaluator class handles loss metric for you.
+ logs = {self.loss: loss}
+
+ outputs = utils.decode_object_seq_to_bbox(
+ logits,
+ tokens,
+ self._task_config.quantization_bins,
+ self._task_config.coord_vocab_shift,
+ )
+ pred_classes, pred_bboxes, scores, pred_num = outputs
+
+ image_size = features.shape[1:3].as_list()
+ # scale points to original image size during eval.
+ scale = utils.tf_float32(image_size)[tf.newaxis, :] / utils.tf_float32(
+ labels['image_info'][:, 1:2, :]
+ )
+ scale = scale * utils.tf_float32(labels['image_info'][:, 0:1, :])
+ pred_bboxes = utils.scale_points(pred_bboxes, scale)
+
+ predictions = {
+ 'detection_boxes': pred_bboxes,
+ 'detection_scores': scores,
+ 'detection_classes': pred_classes,
+ 'num_detections': pred_num,
+ 'source_id': labels['id'],
+ 'image_info': labels['image_info'],
+ }
+
+ ground_truths = {
+ 'source_id': labels['id'],
+ 'height': labels['image_info'][:, 0:1, 0],
+ 'width': labels['image_info'][:, 0:1, 1],
+ 'num_detections': tf.reduce_sum(
+ tf.cast(tf.math.greater(labels['classes'], 0), tf.int32), axis=-1
+ ),
+ 'boxes': labels['gt_boxes'],
+ 'classes': labels['classes'],
+ 'is_crowds': labels['is_crowd'],
+ }
+ logs.update({'predictions': predictions, 'ground_truths': ground_truths})
+
+ all_losses = {
+ 'loss': loss,
+ }
+
+ # Metric results will be added to logs for you.
+ if metrics:
+ for m in metrics:
+ m.update_state(all_losses[m.name])
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ if state is None:
+ self.coco_metric.reset_states()
+ state = self.coco_metric
+
+ state.update_state(
+ step_outputs['ground_truths'], step_outputs['predictions'] # pyrefly: ignore[unsupported-operation]
+ )
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ return aggregated_logs.result()
diff --git a/official/projects/pix2seq/train.py b/official/projects/pix2seq/train.py
new file mode 100644
index 00000000000..edeb6760a07
--- /dev/null
+++ b/official/projects/pix2seq/train.py
@@ -0,0 +1,73 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""TensorFlow Model Garden Vision training driver."""
+
+from absl import app
+from absl import flags
+import gin
+
+from official.common import distribute_utils
+from official.common import flags as tfm_flags
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+# pylint: disable=unused-import
+from official.projects.pix2seq.configs import pix2seq
+from official.projects.pix2seq.tasks import pix2seq_task
+# pylint: enable=unused-import
+
+FLAGS = flags.FLAGS
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ # Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
+ # dtype is float16
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu,
+ )
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir,
+ )
+
+ train_utils.save_gin_config(FLAGS.mode, model_dir)
+
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ flags.mark_flags_as_required(['experiment', 'mode', 'model_dir'])
+ app.run(main)
diff --git a/official/projects/pix2seq/utils.py b/official/projects/pix2seq/utils.py
new file mode 100644
index 00000000000..b6eb0ee2eb1
--- /dev/null
+++ b/official/projects/pix2seq/utils.py
@@ -0,0 +1,396 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pix2Seq required utility library."""
+import copy
+
+import tensorflow as tf, tf_keras
+from official.projects.pix2seq.configs import pix2seq as pix2seq_cfg
+
+
+def decode_object_seq_to_bbox(
+ logits, pred_seq, quantization_bins, coord_vocab_shift
+):
+ """Decode objects (label & bbox) for seq from `build_response_seq_from_bbox`.
+
+ Assume yxyxc format with truncation at the end for any uneven extra tokens.
+
+ Replace class tokens with argmax instead of sampling.
+ Args:
+ logits: `float` output logits in shape of (bsz, max_seq_len, vocab_size).
+ pred_seq: `int` pred sequence in shape of (bsz, max_seq_len).
+ quantization_bins: `int` for bins.
+ coord_vocab_shift: `int`, shifting coordinates by a specified integer.
+
+ Returns:
+ pred_class: `int` of shape (bsz, max_instances_per_image).
+ pred_bbox: `float` of shape (bsz, max_instances_per_image, 4).
+ pred_score: `float` of shape (bsz, max_instances_per_image).
+ """
+ _, seqlen, vocab_size = logits.shape
+
+ if seqlen % 5 != 0: # truncate out the last few tokens.
+ pred_seq = pred_seq[..., : -(seqlen % 5)]
+ logits = logits[..., : -(seqlen % 5), :]
+ pred_class_p = tf.nn.softmax(logits)[:, 4::5] # (bsz, instances, vocab_size)
+ mask_s1 = [0.0] * (pix2seq_cfg.BASE_VOCAB_SHIFT) # reserved.
+ mask_s2 = [1.0] * (
+ coord_vocab_shift - pix2seq_cfg.BASE_VOCAB_SHIFT
+ ) # labels.
+ mask_s3 = [0.0] * (vocab_size - coord_vocab_shift) # coordinates and others.
+ mask = tf.constant(mask_s1 + mask_s2 + mask_s3)
+ pred_class = tf.argmax(pred_class_p * mask[tf.newaxis, tf.newaxis, :], -1)
+
+ pred_num = logits[:, 4::5] * mask[tf.newaxis, tf.newaxis, :]
+ pred_num = tf.reduce_sum(
+ tf.cast(
+ tf.math.greater(tf.math.reduce_max(pred_num, axis=-1), 0), tf.int32
+ ),
+ axis=-1,
+ )
+
+ pred_score = tf.reduce_sum(
+ pred_class_p * tf.one_hot(pred_class, vocab_size), -1
+ )
+ pred_class = tf.maximum(pred_class - pix2seq_cfg.BASE_VOCAB_SHIFT, 0)
+ pred_bbox = seq_to_bbox(pred_seq - coord_vocab_shift, quantization_bins)
+ return pred_class, pred_bbox, pred_score, pred_num
+
+
+def seq_to_bbox(seq, quantization_bins, seq_format='yxyx_name'):
+ """Returns [0, 1] normalized yxyx bbox from token sequence."""
+ # [batch, 5*num_instances]
+ assert seq.shape.rank == 2, seq.shape.as_list()
+ # [batch, num_instances, 1]
+ if seq_format.startswith('name'):
+ ymin = tf.expand_dims(seq[:, 1::5], -1)
+ xmin = tf.expand_dims(seq[:, 2::5], -1)
+ ymax = tf.expand_dims(seq[:, 3::5], -1)
+ xmax = tf.expand_dims(seq[:, 4::5], -1)
+ else:
+ ymin = tf.expand_dims(seq[:, 0::5], -1)
+ xmin = tf.expand_dims(seq[:, 1::5], -1)
+ ymax = tf.expand_dims(seq[:, 2::5], -1)
+ xmax = tf.expand_dims(seq[:, 3::5], -1)
+ if seq_format in ['name_cycxhw', 'cycxhw_name']:
+ ycnt, xcnt, ysize, xsize = ymin, xmin, ymax, xmax
+ ymin = ycnt - ysize // 2
+ xmin = xcnt - xsize // 2
+ ymax = ycnt + ysize // 2
+ xmax = xcnt + xsize // 2
+ quantized_box = tf.concat([ymin, xmin, ymax, xmax], axis=-1)
+ quantized_box = dequantize(quantized_box, quantization_bins)
+ return tf.minimum(tf.maximum(quantized_box, 0), 1)
+
+
+def quantize(coordinates, bins):
+ """Quantization of (normalized) coordinates in [0, 1]."""
+ coordinates = tf.cast(tf.round(coordinates * (bins - 1)), tf.int64)
+ coordinates = tf.clip_by_value(coordinates, 0, bins - 1)
+ return coordinates
+
+
+def dequantize(boxes, bins):
+ """Dequantization of discrete tokens of coordinates in [0, bins-1]."""
+ boxes = tf.cast(boxes, tf.float32)
+ boxes = boxes / (bins - 1)
+ return boxes
+
+
+def truncation_bbox(bbox):
+ return tf.minimum(tf.maximum(bbox, 0.0), 1.0)
+
+
+def jitter_bbox(bbox, min_range=0.0, max_range=0.05, truncation=True):
+ """Jitter the bbox.
+
+ Args:
+ bbox: `float` tensor of shape (n, 4), ranged between 0 and 1.
+ min_range: min jitter range in ratio to bbox size.
+ max_range: max jitter range in ratio to bbox size.
+ truncation: whether to truncate resulting bbox to remain [0, 1].
+ Note: To create noisy positives, set min_range=0, which enables truncated
+ normal distribution. max_range <=0.05: noisy duplicates, <=0.02: near
+ duplicate. To create negatives: set min_range >= 0.1 to avoid false
+ negatives; suggested max_range <=0.4 to avoid too much randomness.
+
+ Returns:
+ jittered bbox.
+ """
+ n = tf.shape(bbox)[0]
+ h = bbox[:, 2] - bbox[:, 0]
+ w = bbox[:, 3] - bbox[:, 1]
+ noise = tf.stack([h, w, h, w], -1)
+ if min_range == 0:
+ noise_rate = tf.random.truncated_normal(
+ [n, 4], mean=0, stddev=max_range / 2.0, dtype=bbox.dtype
+ )
+ else:
+ noise_rate1 = tf.random.uniform([n, 4], min_range, max_range)
+ noise_rate2 = tf.random.uniform([n, 4], -max_range, -min_range)
+ selector = tf.cast(tf.random.uniform([n, 4], 0, 1) < 0.5, tf.float32)
+ noise_rate = noise_rate1 * selector + noise_rate2 * (1.0 - selector)
+ bbox = bbox + noise * noise_rate
+ return truncation_bbox(bbox) if truncation else bbox
+
+
+def shift_bbox(bbox, truncation=True):
+ """Shifting bbox without changing the bbox height and width."""
+ n = tf.shape(bbox)[0]
+ # randomly sample new bbox centers.
+ cy = tf.random.uniform([n, 1], 0, 1)
+ cx = tf.random.uniform([n, 1], 0, 1)
+ h = bbox[:, 2:3] - bbox[:, 0:1]
+ w = bbox[:, 3:4] - bbox[:, 1:2]
+ bbox = tf.concat(
+ [
+ cy - tf.abs(h) / 2,
+ cx - tf.abs(w) / 2,
+ cy + tf.abs(h) / 2,
+ cx + tf.abs(w) / 2,
+ ],
+ -1,
+ )
+ return truncation_bbox(bbox) if truncation else bbox
+
+
+def random_bbox(n, max_size=1.0, truncation=True):
+ """Generating random n bbox with max size specified within [0, 1]."""
+ cy = tf.random.uniform([n, 1], 0, 1)
+ cx = tf.random.uniform([n, 1], 0, 1)
+ h = tf.random.truncated_normal([n, 1], 0, max_size / 2.0)
+ w = tf.random.truncated_normal([n, 1], 0, max_size / 2.0)
+ bbox = tf.concat(
+ [
+ cy - tf.abs(h) / 2,
+ cx - tf.abs(w) / 2,
+ cy + tf.abs(h) / 2,
+ cx + tf.abs(w) / 2,
+ ],
+ -1,
+ )
+ return truncation_bbox(bbox) if truncation else bbox
+
+
+def augment_bbox(bbox, bbox_label, max_jitter, n_noise_bbox, mix_rate=0.0):
+ """Augment bbox.
+
+ There are two types of noises to add:
+
+ 1. Bad bbox: jittered bbox, shifted bbox, or random bbox.
+ 2. Duplicated bbox.
+ Args:
+ bbox: `float` tensor of shape (n, 4), ranged between 0 and 1.
+ bbox_label: `int` tensor of shape (n,).
+ max_jitter: `float` scalar specifying max jitter range for positive bbox.
+ n_noise_bbox: `int` scalar tensor specifying size of the extra noise to add.
+ mix_rate: `float`. Probability of injecting the bad bbox in the middle of
+ original bbox, followed by dup bbox at the end; otherwise simply append
+ all noises at the end of original bbox.
+
+ Returns:
+ bbox_new: augmented bbox that's `n_noise_bbox` larger than original.
+ label_new: new label for bbox_new.
+ is_real: a `float` 0/1 indicator for whether a bbox is real.
+ is_noise: a `float` 0/1 indicator for whether a bbox is extra.
+ """
+ n = tf.shape(bbox)[0]
+ dup_bbox_size = tf.random.uniform([], 0, n_noise_bbox + 1, dtype=tf.int32)
+ dup_bbox_size = 0 if n == 0 else dup_bbox_size
+ bad_bbox_size = n_noise_bbox - dup_bbox_size
+ multiplier = 1 if n == 0 else tf.math.floordiv(n_noise_bbox, n) + 1
+ bbox_tiled = tf.tile(bbox, [multiplier, 1])
+
+ # Create bad bbox.
+ bbox_tiled = tf.random.shuffle(bbox_tiled)
+ bad_bbox_shift = shift_bbox(bbox_tiled[:bad_bbox_size], truncation=True)
+ bad_bbox_random = random_bbox(bad_bbox_size, max_size=1.0, truncation=True)
+ bad_bbox = tf.concat([bad_bbox_shift, bad_bbox_random], 0)
+ bad_bbox = tf.random.shuffle(bad_bbox)[:bad_bbox_size]
+ bad_bbox_label = tf.zeros([bad_bbox_size], dtype=bbox_label.dtype) + (
+ pix2seq_cfg.FAKE_CLASS_TOKEN - pix2seq_cfg.BASE_VOCAB_SHIFT
+ )
+
+ # Create dup bbox.
+ bbox_tiled = tf.random.shuffle(bbox_tiled)
+ dup_bbox = jitter_bbox(
+ bbox_tiled[:dup_bbox_size], min_range=0, max_range=0.1, truncation=True
+ )
+ dup_bbox_label = tf.zeros([dup_bbox_size], dtype=bbox_label.dtype) + (
+ pix2seq_cfg.FAKE_CLASS_TOKEN - pix2seq_cfg.BASE_VOCAB_SHIFT
+ )
+
+ # Jitter positive bbox.
+ if max_jitter > 0:
+ bbox = jitter_bbox(bbox, min_range=0, max_range=max_jitter, truncation=True)
+
+ if tf.random.uniform([]) < mix_rate:
+ # Mix the bbox with bad bbox, appneded by dup bbox.
+ bbox_new = tf.concat([bbox, bad_bbox], 0)
+ bbox_new_label = tf.concat([bbox_label, bad_bbox_label], 0)
+ idx = tf.random.shuffle(tf.range(tf.shape(bbox_new)[0]))
+ bbox_new = tf.gather(bbox_new, idx)
+ bbox_new_label = tf.gather(bbox_new_label, idx)
+ bbox_new = tf.concat([bbox_new, dup_bbox], 0)
+ bbox_new_label = tf.concat([bbox_new_label, dup_bbox_label], 0)
+ else:
+ # Merge bad bbox and dup bbox into noise bbox.
+ noise_bbox = tf.concat([bad_bbox, dup_bbox], 0)
+ noise_bbox_label = tf.concat([bad_bbox_label, dup_bbox_label], 0)
+
+ if n_noise_bbox > 0:
+ idx = tf.random.shuffle(tf.range(n_noise_bbox))
+ noise_bbox = tf.gather(noise_bbox, idx)
+ noise_bbox_label = tf.gather(noise_bbox_label, idx)
+
+ # Append noise bbox to bbox and create mask.
+ bbox_new = tf.concat([bbox, noise_bbox], 0)
+ bbox_new_label = tf.concat([bbox_label, noise_bbox_label], 0)
+
+ return bbox_new, bbox_new_label
+
+
+def inject_noise_bbox(boxes, classes, max_instances_per_image):
+ boxes = copy.copy(boxes)
+ classes = copy.copy(classes)
+ num_instances = tf.shape(boxes)[0]
+ if num_instances < max_instances_per_image:
+ n_noise_bbox = max_instances_per_image - num_instances
+ boxes, classes = augment_bbox(boxes, classes, 0.0, n_noise_bbox)
+ return boxes, classes
+
+
+def build_prompt_seq_from_task_id(
+ task_vocab_id: int, response_seq=None, prompt_shape=None
+):
+ """Build prompt seq just using task id.
+
+ Args:
+ task_vocab_id: Vocab id for the task.
+ response_seq: an (optional) discerte target sequen with shape (bsz, ..., k).
+ prompt_shape: an (optional) tuple for prompt shape. One and only one of
+ `response_seq` and `prompt_shape` should be specified.
+
+ Returns:
+ discrete input sequence of task id with shape (bsz, ..., 1).
+ """
+ task_id = tf.constant(task_vocab_id)
+ if response_seq is not None:
+ prompt_seq = tf.zeros_like(response_seq[..., :1]) + tf.cast(
+ task_id, response_seq.dtype
+ )
+ if prompt_shape is not None:
+ assert response_seq is None, 'double specification'
+ prompt_seq = tf.zeros(prompt_shape, dtype=tf.int64) + tf.cast(
+ task_id, dtype=tf.int64
+ )
+ return prompt_seq # pyrefly: ignore[unbound-name]
+
+
+def clip_or_pad_to_max_len(data, max_len, dim):
+ """Pad the data tensor to max length on dim."""
+ shape = shape_as_list(data)
+ padding_shape, clipped_shape = copy.copy(shape), copy.copy(shape)
+ padding_shape[dim] = tf.maximum(0, max_len - padding_shape[dim])
+ clipped_shape[dim] = tf.minimum(clipped_shape[dim], max_len)
+
+ paddings = tf.zeros(padding_shape, dtype=data.dtype)
+ clipped_data = tf.slice(data, tf.zeros_like(shape), clipped_shape)
+ return tf.concat([clipped_data, paddings], axis=dim)
+
+
+def shape_as_list(t):
+ # Assumes rank of `t` is statically known.
+ shape = t.shape.as_list()
+ dynamic_shape = tf.shape(t)
+ return [
+ shape[i] if shape[i] is not None else dynamic_shape[i]
+ for i in range(len(shape))
+ ]
+
+
+def reorder_object_instances(boxes, classes, order):
+ """Must be called _before_ padding to max instances."""
+ if order == 'none':
+ return classes, boxes
+
+ assert boxes.shape.rank == 2, 'Must be unbatched'
+ boxes = tf.reshape(boxes, [-1, 2, 2])
+
+ if order == 'random':
+ idx = tf.random.shuffle(tf.range(tf.shape(boxes)[0]))
+ elif order == 'area':
+ areas = tf.cast(
+ tf.reduce_prod(boxes[:, 1, :] - boxes[:, 0, :], axis=1), tf.int64
+ ) # approximated size.
+ idx = tf.argsort(areas, direction='DESCENDING')
+ elif order == 'dist2ori':
+ y, x = boxes[:, 0], boxes[:, 1] # using top-left corner.
+ dist2ori = tf.square(y) + tf.square(x)
+ idx = tf.argsort(dist2ori, direction='ASCENDING')
+ else:
+ raise ValueError('Unknown order {}'.format(order))
+
+ boxes = tf.reshape(boxes, [-1, 4])
+ boxes = tf.gather(boxes, idx)
+ classes = tf.gather(classes, idx)
+
+ return boxes, classes
+
+
+def scale_points(points, scale):
+ """Scales points.
+
+ Args:
+ points: Tensor with shape [num_points * 2], [batch, num_points * 2] or
+ [batch, instances, num_points * 2] where points are organized in (y, x)
+ format.
+ scale: Tensor with shape [2] or [batch, 2].
+
+ Returns:
+ Tensor with same shape as points.
+ """
+ points_orig = points
+ orig_shape = tf.shape(points)
+ coords_len = points.shape[-1]
+ if points.shape.rank == 1:
+ points = tf.reshape(points, [coords_len // 2, 2])
+ elif points.shape.rank == 2:
+ points = tf.reshape(points, [-1, coords_len // 2, 2])
+ else:
+ points = tf.reshape(points, [-1, orig_shape[1], coords_len // 2, 2])
+ scale = tf.expand_dims(scale, -2)
+ points = points * scale
+ points = tf.reshape(points, orig_shape)
+ points = preserve_reserved_tokens(points, points_orig)
+ return points
+
+
+def preserve_reserved_tokens(points, points_orig):
+ """Preserve reserved tokens in points according to points_orig."""
+ return replace_reserved_tokens(
+ points, points_orig, dict(zip(pix2seq_cfg.FLOATS, pix2seq_cfg.FLOATS))
+ )
+
+
+def replace_reserved_tokens(seq, ref_seq, replacements):
+ for key, replacement in replacements.items():
+ seq = tf.where(
+ tf.equal(ref_seq, key), tf.constant(replacement, seq.dtype), seq
+ )
+ return seq
+
+
+def tf_float32(t):
+ return tf.cast(t, tf.float32)
diff --git a/official/projects/pixel/README.md b/official/projects/pixel/README.md
new file mode 100644
index 00000000000..6fa64a30f5c
--- /dev/null
+++ b/official/projects/pixel/README.md
@@ -0,0 +1,38 @@
+# Language Modelling with Pixels (PIXEL)
+
+TF2 implementation of [PIXEL](https://arxiv.org/abs/2207.06991).
+
+### Setup
+
+The current setup requires a numpyfied pytorch pixel model and preprocessed
+data. For the pixel model, we directly convert its state_dict and saved as
+numpy. For the preprocessed data, we run their pytorch implementation and save
+the pixel transformed data.
+
+Let's put these data in the directory `PATH_TO_PIXEL_DATA_DIR`, then
+, to convert the numpyfied model to a tensorflow checkpoint, run
+
+```shell
+python3 utils/convert_numpy_weights_to_tf.py $PATH_TO_PIXEL_DATA_DIR
+```
+
+This will create a `pixel_encoder.ckpt`. Denote the path to this checkpoint as
+`PATH_TO_PIXEL_ENCODER_CKPT`.
+
+### Training
+
+```shell
+export PATH_TO_PIXEL_DATA_DIR=xxx
+export PATH_TO_PIXEL_ENCODER_CKPT=xxx
+PATH_TO_TRAINING_RECORD=$PATH_TO_PIXEL_DATA_DIR/train.tf_record-*-of-20 # path to the training record
+PATH_TO_TESTING_RECORD=$PATH_TO_PIXEL_DATA_DIR/eval.tf_record # path to the evaluation record
+TPU_NAME="" # The name assigned while creating a Cloud TPU
+MODEL_DIR=/tmp/pixel_sst2 # directory to store the experiment
+# Now launch the experiment.
+python3 -m official.projects.pixel.train \
+ --experiment=pixel_sst2_finetune \
+ --params_override="task.train_data.input_path=${PATH_TO_TRAINING_RECORD},task.validation_data.input_path=${PATH_TO_TESTING_RECORD},runtime.distribution_strategy=tpu,init_checkpoint=$PATH_TO_PIXEL_ENCODER_CKPT"
+ --mode=train_and_eval \
+ --tpu=$TPU_NAME \
+ --model_dir=$MODEL_DIR
+```
\ No newline at end of file
diff --git a/official/projects/pixel/configs/pixel.py b/official/projects/pixel/configs/pixel.py
new file mode 100644
index 00000000000..3ed725a080f
--- /dev/null
+++ b/official/projects/pixel/configs/pixel.py
@@ -0,0 +1,112 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pixel configurations."""
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+from official.projects.pixel.data_loader import PixelDataConfig
+from official.projects.pixel.tasks.classification import PixelConfig
+from official.projects.pixel.tasks.classification import PixelModelConfig
+
+
+@exp_factory.register_config_factory('pixel_sst2_finetune')
+def pixel_sst2_finetune() -> cfg.ExperimentConfig:
+ """Config to get results that matches https://github.com/xplip/pixel for sst2."""
+ train_batch_size = 256
+ eval_batch_size = 32
+ num_train_steps = 15000
+
+ input_size = (16, 4096)
+ patch_h, patch_w = 16, 16
+ num_channels = 3
+ num_classes = 2
+
+ config = cfg.ExperimentConfig(
+ task=PixelConfig(
+ train_data=PixelDataConfig(
+ input_path=None,
+ is_training=True,
+ global_batch_size=train_batch_size,
+ shuffle_buffer_size=10000,
+ drop_remainder=True,
+ input_size=input_size,
+ patch_h=patch_h,
+ patch_w=patch_w,
+ num_channels=num_channels,
+ ),
+ validation_data=PixelDataConfig(
+ input_path=None,
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ shuffle_buffer_size=10000,
+ drop_remainder=True,
+ input_size=input_size,
+ patch_h=patch_h,
+ patch_w=patch_w,
+ num_channels=num_channels,
+ ),
+ model=PixelModelConfig(
+ filters=768,
+ num_layers=12,
+ mlp_dim=3072,
+ num_heads=12,
+ dropout_rate=0.1,
+ attention_dropout_rate=0.1,
+ init_stochastic_depth_rate=0.0,
+ ),
+ init_checkpoint=None,
+ input_size=input_size,
+ patch_h=patch_h,
+ patch_w=patch_w,
+ num_channels=num_channels,
+ num_classes=num_classes,
+ ),
+ trainer=cfg.TrainerConfig(
+ train_steps=num_train_steps,
+ validation_steps=27,
+ steps_per_loop=100,
+ summary_interval=100,
+ checkpoint_interval=100,
+ validation_interval=100,
+ max_to_keep=1,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ },
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'cycle': False,
+ 'polynomial': {
+ 'decay_steps': num_train_steps,
+ 'end_learning_rate': 0.0,
+ 'initial_learning_rate': 3.0e-05,
+ 'power': 1.0,
+ },
+ },
+ 'warmup': {
+ 'type': 'polynomial',
+ 'polynomial': {
+ 'warmup_steps': 100,
+ 'power': 1.0,
+ },
+ },
+ }),
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ ],
+ )
+ return config
diff --git a/official/projects/pixel/data_loader.py b/official/projects/pixel/data_loader.py
new file mode 100644
index 00000000000..dac72da4b20
--- /dev/null
+++ b/official/projects/pixel/data_loader.py
@@ -0,0 +1,110 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Loads dataset for the Pixel Classification task."""
+import dataclasses
+from typing import Mapping, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import config_definitions as cfg
+from official.core import input_reader
+from official.nlp.data import data_loader
+from official.nlp.data import data_loader_factory
+
+LABEL_TYPES_MAP = {'int': tf.int64, 'float': tf.float32}
+
+
+@dataclasses.dataclass
+class PixelDataConfig(cfg.DataConfig):
+ """Data config for text classification task."""
+
+ input_path: str = ''
+ global_batch_size: int = 32
+ is_training: bool = True
+ label_type: str = 'int'
+ num_channels: int = 3
+ input_size: Tuple[int, int] = (16, 4096)
+ patch_h: int = 16
+ patch_w: int = 16
+ # Whether to include the example id number.
+ include_example_id: bool = False
+ # Maps the key in TfExample to feature name.
+ # Either tfrecord, sstable, or recordio.
+ file_type: str = 'tfrecord'
+
+
+@data_loader_factory.register_data_loader_cls(PixelDataConfig)
+class PixelDataLoader(data_loader.DataLoader):
+ """A class to load dataset for text classification task."""
+
+ def __init__(self, params):
+ self._params = params
+ self._include_example_id = params.include_example_id
+
+ def name_to_features_spec(self):
+ """Defines features to decode. Subclass may override to append features."""
+ h, w = self._params.input_size
+ positions = h // self._params.patch_h * w // self._params.patch_w
+ name_to_features = {
+ 'pixel_values': tf.io.FixedLenFeature(
+ [self._params.num_channels, h, w], tf.float32
+ ),
+ 'label': tf.io.FixedLenFeature([1], tf.int64),
+ 'attention_mask': tf.io.FixedLenFeature([positions], tf.float32),
+ }
+ if self._include_example_id:
+ name_to_features['example_id'] = tf.io.FixedLenFeature([], tf.int64)
+
+ return name_to_features
+
+ def _decode(self, record: tf.Tensor):
+ """Decodes a serialized tf.Example."""
+ example = tf.io.parse_single_example(record, self.name_to_features_spec())
+
+ # tf.Example only supports tf.int64, but the TPU only supports tf.int32.
+ # So cast all int64 to int32.
+ for name in example:
+ t = example[name]
+ if t.dtype == tf.int64:
+ t = tf.cast(t, tf.int32)
+ example[name] = t
+
+ return example
+
+ def _parse(self, record: Mapping[str, tf.Tensor]):
+ """Parses raw tensors into a dict of tensors to be consumed by the model."""
+ key_mapping = {
+ 'pixel_values': 'pixel_values',
+ 'label': 'label',
+ 'attention_mask': 'attention_mask',
+ }
+ ret = {}
+ for record_key in record:
+ if record_key in key_mapping:
+ ret[key_mapping[record_key]] = record[record_key]
+ else:
+ ret[record_key] = record[record_key]
+ return ret
+
+ def load(self, input_context: Optional[tf.distribute.InputContext] = None):
+ """Returns a tf.dataset.Dataset."""
+ reader = input_reader.InputReader(
+ dataset_fn=dataset_fn.pick_dataset_fn(self._params.file_type),
+ params=self._params,
+ decoder_fn=self._decode,
+ parser_fn=self._parse,
+ )
+ return reader.read(input_context)
diff --git a/official/projects/pixel/modeling/pixel.py b/official/projects/pixel/modeling/pixel.py
new file mode 100644
index 00000000000..3a4185df4ea
--- /dev/null
+++ b/official/projects/pixel/modeling/pixel.py
@@ -0,0 +1,187 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pixel models."""
+
+import tensorflow as tf, tf_keras
+
+from official.vision.modeling.backbones import vit
+
+layers = tf_keras.layers
+
+
+class ViTEncoder(vit.Encoder):
+ """ViT Encoder.
+
+ The original vit implementation in official/vision/modeling/backbones/vit.py
+ does not support attention masks. This version allows passing the attention
+ mask in call along with inputs as a (bs, seqlen) tensor.
+ """
+
+ def call(self, inputs, training=None):
+ x, mask = inputs
+ if self._add_pos_embed:
+ x = self._pos_embed(x, inputs_positions=self._inputs_positions)
+ x = self._dropout(x, training=training)
+
+ for encoder_layer in self._encoder_layers:
+ x = encoder_layer((x, mask), training=training)
+ x = self._norm(x)
+ return x
+
+
+class VisionTransformer(tf_keras.layers.Layer):
+ """ViT backbone."""
+
+ def __init__(
+ self,
+ patch_h,
+ patch_w,
+ filters,
+ num_layers,
+ mlp_dim,
+ num_heads,
+ dropout_rate,
+ attention_dropout_rate,
+ init_stochastic_depth_rate,
+ **kwargs
+ ):
+ super().__init__(**kwargs)
+ self.patch_h = patch_h
+ self.patch_w = patch_w
+
+ self.filters = filters
+ self.num_layers = num_layers
+ self.mlp_dim = mlp_dim
+ self.num_heads = num_heads
+ self.dropout_rate = dropout_rate
+ self.attention_dropout_rate = attention_dropout_rate
+ self.init_stochastic_depth_rate = init_stochastic_depth_rate
+
+ def build(self, input_shape):
+ self.patch_to_embed = tf_keras.layers.Conv2D(
+ filters=self.filters,
+ kernel_size=(self.patch_h, self.patch_w),
+ strides=(self.patch_h, self.patch_w),
+ padding='valid',
+ kernel_initializer='lecun_normal',
+ )
+
+ self.encoder = ViTEncoder(
+ num_layers=self.num_layers,
+ mlp_dim=self.mlp_dim,
+ num_heads=self.num_heads,
+ dropout_rate=self.dropout_rate,
+ attention_dropout_rate=self.attention_dropout_rate,
+ init_stochastic_depth_rate=self.init_stochastic_depth_rate,
+ add_pos_embed=True,
+ )
+ self.token_cls = vit.TokenLayer()
+ super().build(input_shape)
+
+ def to_embed(self, patches):
+ return self.patch_to_embed(patches)
+
+ def insert_cls(self, patch_embeds):
+ return self.token_cls(patch_embeds)
+
+ def call(self, inputs): # pylint:disable=signature-mismatch
+ if isinstance(inputs, dict):
+ images = inputs.get('pixel_values', None)
+ attention_mask = inputs.get('attention_mask', None)
+ attention_mask = tf.transpose(
+ tf.concat(
+ values=[
+ tf.ones((1, tf.shape(attention_mask)[0]), tf.float32),
+ tf.transpose(attention_mask),
+ ],
+ axis=0,
+ )
+ )
+ attention_mask = tf.einsum('ij,ik->ijk', attention_mask, attention_mask)
+ attention_mask = tf.cast(attention_mask, tf.int32)
+ else:
+ raise ValueError('Unexpected inputs type to %s.' % self.__class__)
+
+ images = tf.transpose(images, perm=[0, 2, 3, 1])
+ patch_embeds = self.to_embed(images)
+ patch_shape = tf.shape(patch_embeds)
+ patch_embeds = tf.reshape(
+ patch_embeds, (patch_shape[0], -1, patch_shape[-1])
+ )
+ patch_embeds = self.insert_cls(patch_embeds)
+
+ return self.encoder((patch_embeds, attention_mask))
+
+
+class PixelClassifier(tf_keras.layers.Layer):
+ """Pixel classifier for finetuning. Uses the cls token."""
+
+ def __init__(self, encoder, num_classes, **kwargs):
+ super().__init__(**kwargs)
+ self.encoder = encoder
+ self.linear = tf_keras.layers.Dense(
+ num_classes,
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01),
+ )
+
+ def call(self, inputs):
+ encoded = self.encoder(inputs)
+ return self.linear(encoded[:, 0])
+
+
+class PixelLinearClassifier(tf_keras.layers.Layer):
+ """Pixel classifier for finetuning.
+
+ This is a layer with additional layer norm and linear layer in the
+ classification head. Uses the average of all token representations
+ """
+
+ def __init__(self, encoder, num_classes, num_filters, **kwargs):
+ super().__init__(**kwargs)
+ self.encoder = encoder
+ self.num_filters = num_filters
+ self.linear_clas = tf_keras.layers.Dense(
+ num_classes,
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01),
+ )
+
+ self.norm = tf_keras.layers.LayerNormalization(
+ name='classification_layer_norm',
+ axis=-1,
+ epsilon=1e-6,
+ dtype=tf.float32,
+ )
+
+ self.linear_trans = tf_keras.layers.Dense(
+ num_filters,
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.01),
+ )
+ self.activation = tf_keras.layers.Activation('gelu')
+ self.dropout = tf_keras.layers.Dropout(0.1)
+
+ def call(self, inputs, training=False):
+ attention_mask = inputs.get('attention_mask')
+ mask_lengths = tf.expand_dims(tf.reduce_sum(attention_mask, axis=1), 1)
+ attention_mask = tf.tile(
+ tf.expand_dims(attention_mask, 2), [1, 1, self.num_filters]
+ )
+ encoded = self.encoder(inputs)
+ encoded = self.norm(self.activation(self.linear_trans(encoded)))
+ encoded = self.dropout(encoded, training=training)
+
+ mean_pooling = (
+ tf.reduce_sum(encoded[:, 1:, :] * attention_mask, axis=1) / mask_lengths
+ )
+ return self.linear_clas(mean_pooling)
diff --git a/official/projects/pixel/tasks/classification.py b/official/projects/pixel/tasks/classification.py
new file mode 100644
index 00000000000..aeda3facb0b
--- /dev/null
+++ b/official/projects/pixel/tasks/classification.py
@@ -0,0 +1,218 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Text classification task with ViT."""
+
+import dataclasses
+from typing import Tuple
+
+import numpy as np
+from scipy import stats
+from sklearn import metrics as sklearn_metrics
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import config_definitions as cfg
+from official.core import task_factory
+from official.modeling import tf_utils
+from official.modeling.hyperparams import base_config
+from official.nlp.data import data_loader_factory
+from official.projects.pixel.modeling import pixel
+
+
+@dataclasses.dataclass
+class PixelModelConfig(base_config.Config):
+ """The model configuration."""
+
+ filters: int = 768
+ num_layers: int = 12
+ mlp_dim: int = 3072
+ num_heads: int = 12
+ dropout_rate: float = 0.1
+ attention_dropout_rate: float = 0.1
+ init_stochastic_depth_rate: float = 0.0
+
+
+@dataclasses.dataclass
+class PixelConfig(cfg.TaskConfig):
+ """The task configuration."""
+
+ train_data: cfg.DataConfig = cfg.DataConfig()
+ validation_data: cfg.DataConfig = cfg.DataConfig()
+ patch_h: int = 16
+ patch_w: int = 16
+ num_classes: int = 2
+ num_channels: int = 3
+ input_size: Tuple[int, int] = (16, 4096)
+ model: PixelModelConfig = PixelModelConfig()
+
+
+@task_factory.register_task_cls(PixelConfig)
+class PixelClassificationTask(base_task.Task):
+ """Text classificaiton with Pixel and load checkpoint if exists."""
+
+ label_field: str = 'label'
+ metric_type: str = 'accuracy'
+
+ def build_model(self) -> tf_keras.Model:
+ encoder = pixel.VisionTransformer(
+ self.task_config.patch_h,
+ self.task_config.patch_w,
+ self.task_config.model.filters,
+ self.task_config.model.num_layers,
+ self.task_config.model.mlp_dim,
+ self.task_config.model.num_heads,
+ self.task_config.model.dropout_rate,
+ self.task_config.model.attention_dropout_rate,
+ self.task_config.model.init_stochastic_depth_rate,
+ )
+ model = pixel.PixelLinearClassifier(
+ encoder, self.task_config.num_classes, self.task_config.model.filters
+ )
+ h, w = self.task_config.input_size
+ positions = h // self.task_config.patch_h * w // self.task_config.patch_w
+ model({
+ 'label': tf.zeros((1,)),
+ 'pixel_values': tf.zeros((1, self.task_config.num_channels, h, w)),
+ 'attention_mask': tf.zeros((1, positions)),
+ })
+ return model
+
+ def build_inputs(self, params, input_context=None):
+ return data_loader_factory.get_data_loader(params).load(input_context)
+
+ def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
+ label_ids = labels[self.label_field]
+ if self.task_config.num_classes == 1:
+ loss = tf_keras.losses.mean_squared_error(label_ids, model_outputs)
+ else:
+ loss = tf_keras.losses.sparse_categorical_crossentropy(
+ label_ids, tf.cast(model_outputs, tf.float32), from_logits=True
+ )
+
+ if aux_losses:
+ loss += tf.add_n(aux_losses)
+ return tf_utils.safe_mean(loss)
+
+ def initialize(self, model: tf_keras.Model):
+ """Load encoder if checkpoint exists.
+
+ Args:
+ model: The keras.Model built or used by this task.
+ """
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+ if not ckpt_dir_or_file:
+ return
+
+ ckpt = tf.train.Checkpoint(encoder=model.encoder)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ def build_metrics(self, training=None):
+ del training
+ if self.task_config.num_classes == 1:
+ metrics = [tf_keras.metrics.MeanSquaredError()]
+ elif self.task_config.num_classes == 2:
+ metrics = [
+ tf_keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
+ tf_keras.metrics.AUC(name='auc', curve='PR'),
+ ]
+ else:
+ metrics = [
+ tf_keras.metrics.SparseCategoricalAccuracy(name='cls_accuracy'),
+ ]
+ return metrics
+
+ def process_metrics(self, metrics, labels, model_outputs):
+ for metric in metrics:
+ if metric.name == 'auc':
+ # Convert the logit to probability and extract the probability of True..
+ metric.update_state(
+ labels[self.label_field],
+ tf.expand_dims(tf.nn.softmax(model_outputs)[:, 1], axis=1),
+ )
+ if metric.name == 'cls_accuracy':
+ metric.update_state(labels[self.label_field], model_outputs)
+
+ def process_compiled_metrics(self, compiled_metrics, labels, model_outputs):
+ compiled_metrics.update_state(labels[self.label_field], model_outputs)
+
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
+ features, labels = inputs, inputs
+ outputs = self.inference_step(features, model)
+ loss = self.build_losses(
+ labels=labels, model_outputs=outputs, aux_losses=model.losses
+ )
+ logs = {self.loss: loss}
+ if metrics:
+ self.process_metrics(metrics, labels, outputs)
+ if model.compiled_metrics:
+ self.process_compiled_metrics(model.compiled_metrics, labels, outputs)
+ logs.update({m.name: m.result() for m in metrics or []})
+ logs.update({m.name: m.result() for m in model.metrics})
+ if self.metric_type == 'matthews_corrcoef':
+ logs.update({
+ 'sentence_prediction': (
+ tf.expand_dims( # Ensure one prediction along batch dimension.
+ tf.math.argmax(outputs, axis=1), axis=1
+ )
+ ),
+ 'labels': labels[self.label_field],
+ })
+ else:
+ logs.update({
+ 'sentence_prediction': outputs,
+ 'labels': labels[self.label_field],
+ })
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ if self.metric_type == 'accuracy':
+ return None
+ if state is None:
+ state = {'sentence_prediction': [], 'labels': []}
+ state['sentence_prediction'].append(
+ np.concatenate(
+ [v.numpy() for v in step_outputs['sentence_prediction']], axis=0 # pyrefly: ignore[unsupported-operation]
+ )
+ )
+ state['labels'].append(
+ np.concatenate([v.numpy() for v in step_outputs['labels']], axis=0) # pyrefly: ignore[unsupported-operation]
+ )
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ if self.metric_type == 'accuracy':
+ return None
+
+ preds = np.concatenate(aggregated_logs['sentence_prediction'], axis=0)
+ labels = np.concatenate(aggregated_logs['labels'], axis=0)
+ if self.metric_type == 'f1':
+ preds = np.argmax(preds, axis=1)
+ return {self.metric_type: sklearn_metrics.f1_score(labels, preds)}
+ elif self.metric_type == 'matthews_corrcoef':
+ preds = np.reshape(preds, -1)
+ labels = np.reshape(labels, -1)
+ return {
+ self.metric_type: sklearn_metrics.matthews_corrcoef(preds, labels)
+ }
+ elif self.metric_type == 'pearson_spearman_corr':
+ preds = np.reshape(preds, -1)
+ labels = np.reshape(labels, -1)
+ pearson_corr = stats.pearsonr(preds, labels)[0]
+ spearman_corr = stats.spearmanr(preds, labels)[0]
+ corr_metric = (pearson_corr + spearman_corr) / 2
+ return {self.metric_type: corr_metric}
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/train.py b/official/projects/pixel/train.py
similarity index 67%
rename from official/vision/beta/projects/panoptic_maskrcnn/train.py
rename to official/projects/pixel/train.py
index efbdc10fc6f..86629580143 100644
--- a/official/vision/beta/projects/panoptic_maskrcnn/train.py
+++ b/official/projects/pixel/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -12,15 +12,16 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-"""Panoptic MaskRCNN trainer."""
+"""TensorFlow Model Garden Vision training driver, register pixel configs."""
from absl import app
from official.common import flags as tfm_flags
+# pylint: disable=unused-import
+from official.projects.pixel.configs import pixel
+from official.projects.pixel.tasks import classification
+# pylint: enable=unused-import
from official.vision import train
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as cfg # pylint: disable=unused-import
-from official.vision.beta.projects.panoptic_maskrcnn.tasks import panoptic_maskrcnn as task # pylint: disable=unused-import
-
if __name__ == '__main__':
tfm_flags.define_flags()
diff --git a/official/projects/pixel/utils/convert_numpy_weights_to_tf.py b/official/projects/pixel/utils/convert_numpy_weights_to_tf.py
new file mode 100644
index 00000000000..05aef714e60
--- /dev/null
+++ b/official/projects/pixel/utils/convert_numpy_weights_to_tf.py
@@ -0,0 +1,146 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Convert pixel model from numpy weights to official.projects.pixel."""
+import json
+import sys
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pixel.tasks import classification
+
+
+def convert(vit_encoder, hf_model_param_dict):
+ """Convert pixel model from huggingface to official.projects.pixel."""
+ num_layers = 12
+ num_attention_heads = 12
+ hidden_size = 768
+ head_size = hidden_size // num_attention_heads
+ assert head_size * num_attention_heads == hidden_size
+ vit_encoder.encoder.patch_to_embed.set_weights([
+ hf_model_param_dict[
+ "vit.embeddings.patch_embeddings.projection.weight"
+ ].transpose(2, 3, 1, 0),
+ hf_model_param_dict["vit.embeddings.patch_embeddings.projection.bias"],
+ ])
+ # pylint: disable=protected-access
+ vit_encoder.encoder.encoder._pos_embed.pos_embedding.assign(
+ hf_model_param_dict["vit.embeddings.position_embeddings"][:, :257]
+ )
+ vit_encoder.encoder.encoder._norm.set_weights([
+ hf_model_param_dict["vit.layernorm.weight"],
+ hf_model_param_dict["vit.layernorm.bias"],
+ ])
+ vit_encoder.encoder.token_cls.cls.assign(
+ hf_model_param_dict["vit.embeddings.cls_token"]
+ )
+
+ for layer_num in range(num_layers):
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._attention_layer._query_dense.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.attention.query.weight"
+ ].T.reshape((hidden_size, num_attention_heads, head_size)),
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.attention.query.bias"
+ ].reshape((num_attention_heads, head_size)),
+ ])
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._attention_layer._key_dense.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.attention.key.weight"
+ ].T.reshape((hidden_size, num_attention_heads, head_size)),
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.attention.key.bias"
+ ].reshape((num_attention_heads, head_size)),
+ ])
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._attention_layer._value_dense.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.attention.value.weight"
+ ].T.reshape((hidden_size, num_attention_heads, head_size)),
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.attention.value.bias"
+ ].reshape((num_attention_heads, head_size)),
+ ])
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._attention_layer._output_dense.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.output.dense.weight"
+ ].T.reshape((num_attention_heads, head_size, hidden_size)),
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.attention.output.dense.bias"
+ ],
+ ])
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._attention_layer_norm.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.layernorm_before.weight"
+ ],
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.layernorm_before.bias"
+ ],
+ ])
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._intermediate_dense.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.intermediate.dense.weight"
+ ].T,
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.intermediate.dense.bias"
+ ],
+ ])
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._output_dense.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.output.dense.weight"
+ ].T,
+ hf_model_param_dict[f"vit.encoder.layer.{layer_num}.output.dense.bias"],
+ ])
+ vit_encoder.encoder.encoder._encoder_layers[
+ layer_num
+ ]._output_layer_norm.set_weights([
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.layernorm_after.weight"
+ ],
+ hf_model_param_dict[
+ f"vit.encoder.layer.{layer_num}.layernorm_after.bias"
+ ],
+ ])
+
+
+if __name__ == "__main__":
+
+ data_path = sys.argv[1]
+ output_model_name = sys.argv[2] if len(sys.argv) > 2 else "pixel_encoder.ckpt"
+
+ model_name = json.load(open(f"{data_path}/model_name.json"))
+ model_params = np.load(
+ open(f"{data_path}/model_param.npy", "rb"), allow_pickle=True
+ )
+
+ config = classification.PixelConfig()
+ task = classification.PixelClassificationTask(config)
+ model = task.build_model()
+
+ convert(model, {k: v for k, v in zip(model_name, model_params)})
+ tf.train.Checkpoint(encoder=model.encoder).write(output_model_name)
diff --git a/official/projects/pointpillars/README.md b/official/projects/pointpillars/README.md
new file mode 100644
index 00000000000..c199f89c802
--- /dev/null
+++ b/official/projects/pointpillars/README.md
@@ -0,0 +1,152 @@
+# PointPillars: Fast Encoders for Object Detection from Point Clouds
+
+[](https://arxiv.org/abs/1812.05784)
+
+This repository is the implementation of the following paper.
+
+* [PointPillars: Fast Encoders for Object Detection from Point Clouds](https://arxiv.org/abs/1812.05784)
+
+## Description
+
+PointPillars is point cloud detection model which creatively encodes raw 3D
+point cloud signals into a format (bird-eye-view image) appropriate for a
+downstream detection pipeline. The paper introduces an encoder
+which utilizes PointNets to learn a representation of point clouds organized in
+vertical columns (pillars), then the encoded features can be used with any
+standard 2D convolutional detection architecture.
+
+This model implementation is based on
+[TensorFlow Model Garden](https://github.com/tensorflow/models/tree/master/official/projects/mosaic).
+When trained on processed
+[Waymo Open Dataset 1.2.0](https://waymo.com/open/data/perception/),
+it achieves 45.96% mAP and 45.35% mAPH on the vehicle class. The inference time
+on 1 V100 GPU is 53ms with a batch size 1.
+
+## History
+
+### Nov, 2022
+
+* First release of PointPillars implementation in TensorFlow Model Garden.
+
+## Maintainers
+
+* Xiang Xu ([xiangxu-google](https://github.com/xiangxu-google))
+* Fang Yang ([fyangf](https://github.com/fyangf))
+
+## Requirements
+
+[](https://badge.fury.io/py/tensorflow)
+
+```shell
+pip install --upgrade pip
+pip install tensorflow==2.6.0
+pip install tf-models-official==2.7.2
+pip install apache-beam[gcp]==2.42.0 --user
+```
+
+## Prepare dataset
+
+Take Waymo-Open-Dataset as the example, you need to install the library first:
+
+```shell
+pip install waymo-open-dataset-tf-2-6-0
+```
+
+Then you can use the provided script `tools/process_wod.py` to convert the raw
+[lidar frame data](https://github.com/waymo-research/waymo-open-dataset/blob/master/waymo_open_dataset/dataset.proto#L370)
+into a format which can be fed into the model:
+
+```shell
+SRC_DIR="gs://waymo_open_dataset_v_1_2_0_individual_files"
+DST_DIR="gs://"
+# See https://beam.apache.org/documentation/#runners for distributed runners.
+RUNNER="DirectRunner"
+
+python3 process_wod.py \
+--src_dir=${SRC_DIR} \
+--dst_dir=${DST_DIR} \
+--pipeline_options="--runner=${RUNNER}"
+```
+
+NOTE: This script requires the `--src_dir` to have two sub-folders:
+`training` for training data, and `validation` for validation data.
+
+## Training
+
+You can run the model training on
+[Google Cloud Platform](https://cloud.google.com/) using
+[Cloud TPU](https://cloud.google.com/tpu). Follow this
+[instruction](https://cloud.google.com/tpu/docs/how-to) to set up Cloud TPU.
+
+```shell
+MODEL_DIR="gs://"
+TRAIN_DATA="gs://"
+EVAL_DATA="gs://"
+
+python3 train.py \
+--experiment="pointpillars_baseline" \
+--mode="train" \
+--model_dir=${MODEL_DIR} \
+--config_file="configs/vehicle/pointpillars_3d_baseline_tpu.yaml" \
+--params_override="task.train_data.input_path=${TRAIN_DATA},task.validation_data.input_path=${EVAL_DATA}" \
+--tpu=${TPU}
+```
+
+You can also run the model training using multiple GPUs.
+
+```shell
+MODEL_DIR="gs://"
+TRAIN_DATA="gs://"
+EVAL_DATA="gs://"
+
+python3 train.py \
+--experiment="pointpillars_baseline" \
+--mode="train_and_eval" \
+--model_dir=${MODEL_DIR} \
+--config_file="configs/vehicle/pointpillars_3d_baseline_gpu.yaml" \
+--params_override="task.train_data.input_path=${TRAIN_DATA},task.validation_data.input_path=${EVAL_DATA}"
+```
+
+NOTE: The provided config file `configs/vehicle/pointpillars_3d_baseline_gpu.yaml`
+uses 8 GPUs. If you prefer another number of GPUs, you may want to tune the
+batch size, learning rate and training steps accordingly.
+
+## Results
+
+We use the following experiment setup to get the benchmark result:
+* Lidar range
+ * X: [-76.8, 76.8]
+ * Y: [-76.8, 76.8]
+ * Z: [-3.0, 3.0]
+* Pillars:
+ * Number of pillars per frame: 24000
+ * Number of points per pillar: 100
+ * Number of features per point: 10
+* Bird-eye-view image resolution: [512, 512, 64]
+* Accelerator: Cloud TPU-v2 (16 cores)
+* Batch size: 64
+* Epochs: 75
+
+model | mAP | mAPH | tensorboard
+-------------------- | ------ | ------ | -----------
+PointPillars-vehicle | 45.96% | 45.35% | [link](https://tensorboard.dev/experiment/bDPO7cWxRYKMh5QcMWmVng)
+
+## License
+
+[](https://opensource.org/licenses/Apache-2.0)
+
+This project is licensed under the terms of the **Apache License 2.0**.
+
+## Citation
+
+If you want to cite this repository in your work, please consider citing the
+paper.
+
+```
+@inproceedings{alex2019pointpillars,
+ title={PointPillars: Fast Encoders for Object Detection from Point Clouds},
+ author={Alex H. Lang, Sourabh Vora, Holger Caesar, Lubing Zhou, Jiong Yang, Oscar Beijbom},
+ journal={arXiv preprint arXiv:1812.05784},
+ year={2019},
+}
+```
diff --git a/official/projects/pointpillars/configs/pointpillars.py b/official/projects/pointpillars/configs/pointpillars.py
new file mode 100644
index 00000000000..2df00563abe
--- /dev/null
+++ b/official/projects/pointpillars/configs/pointpillars.py
@@ -0,0 +1,250 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pointpillars experiment configuration definition."""
+import dataclasses
+from typing import List, Optional, Tuple, Union
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.modeling import optimization
+from official.vision.configs import common
+
+
+@dataclasses.dataclass
+class ImageConfig(hyperparams.Config):
+ """Bird-eye-view pseudo image config."""
+ # The range should be large enough to cover a 64-channels Lidar points.
+ # The default values are chosen empirically.
+ x_range: Tuple[float, float] = (-76.8, 76.8)
+ y_range: Tuple[float, float] = (-76.8, 76.8)
+ z_range: Tuple[float, float] = (-3.0, 3.0)
+ resolution: float = 0.3
+ height: int = dataclasses.field(init=False)
+ width: int = dataclasses.field(init=False)
+
+ # Image height and width should be auto computed.
+ def __post_init__(self, height: int, width: int): # pyrefly: ignore[bad-function-definition]
+ self.height = int((-self.x_range[0] + self.x_range[1]) / self.resolution)
+ self.width = int((-self.y_range[0] + self.y_range[1]) / self.resolution)
+
+
+@dataclasses.dataclass
+class PillarsConfig(hyperparams.Config):
+ """Pillars config."""
+ num_pillars: int = 24000
+ num_points_per_pillar: int = 100
+ num_features_per_point: int = 10
+
+
+@dataclasses.dataclass
+class DataDecoder(hyperparams.Config):
+ """Data decoder config."""
+
+
+@dataclasses.dataclass
+class DataParser(hyperparams.Config):
+ """Data parser config."""
+
+
+@dataclasses.dataclass
+class DataConfig(cfg.DataConfig):
+ """Input config for training."""
+ input_path: str = ''
+ global_batch_size: int = 0
+ is_training: bool = False
+ dtype: str = 'float32'
+ decoder: DataDecoder = dataclasses.field(default_factory=DataDecoder)
+ parser: DataParser = dataclasses.field(default_factory=DataParser)
+ shuffle_buffer_size: int = 256
+ prefetch_buffer_size: int = 256
+ file_type: str = 'tfrecord_compressed'
+
+
+@dataclasses.dataclass
+class Anchor(hyperparams.Config):
+ length: float = 1.0
+ width: float = 1.0
+
+
+@dataclasses.dataclass
+class AnchorLabeler(hyperparams.Config):
+ """Data parser config."""
+ match_threshold: float = 0.5
+ unmatched_threshold: float = 0.5
+
+
+@dataclasses.dataclass
+class Featurizer(hyperparams.Config):
+ num_blocks: int = 1
+ num_channels: int = 64
+
+
+@dataclasses.dataclass
+class Backbone(hyperparams.Config):
+ min_level: int = 1
+ max_level: int = 3
+ num_convs: int = 6
+
+
+@dataclasses.dataclass
+class Decoder(hyperparams.Config):
+ """Feature decoder."""
+ # No fields yet, just a placeholder.
+
+
+@dataclasses.dataclass
+class AttributeHead(hyperparams.Config):
+ name: str = ''
+ type: str = 'regression'
+ size: int = 1
+
+
+def _default_heads():
+ return [
+ AttributeHead(name='heading', type='regression', size=1),
+ AttributeHead(name='height', type='regression', size=1),
+ AttributeHead(name='z', type='regression', size=1)
+ ]
+
+
+@dataclasses.dataclass
+class SSDHead(hyperparams.Config):
+ attribute_heads: List[AttributeHead] = dataclasses.field(
+ default_factory=_default_heads)
+
+
+@dataclasses.dataclass
+class DetectionGenerator(hyperparams.Config):
+ """Generator."""
+ apply_nms: bool = True
+ pre_nms_top_k: int = 5000
+ pre_nms_score_threshold: float = 0.05
+ nms_iou_threshold: float = 0.5
+ max_num_detections: int = 100
+ nms_version: str = 'v1' # `v2`, `v1`, `batched`
+ use_cpu_nms: bool = False
+
+
+@dataclasses.dataclass
+class PointPillarsModel(hyperparams.Config):
+ """The model config. Used by build_example_model function."""
+ classes: str = 'all'
+ num_classes: int = 4
+ image: ImageConfig = dataclasses.field(default_factory=ImageConfig)
+ pillars: PillarsConfig = dataclasses.field(default_factory=PillarsConfig)
+ anchors: List[Anchor] = dataclasses.field(default_factory=list)
+ anchor_labeler: AnchorLabeler = dataclasses.field(
+ default_factory=AnchorLabeler
+ )
+
+ min_level: int = 1
+ max_level: int = 3
+ featurizer: Featurizer = dataclasses.field(default_factory=Featurizer)
+ backbone: Backbone = dataclasses.field(default_factory=Backbone)
+ decoder: Decoder = dataclasses.field(default_factory=Decoder)
+ head: SSDHead = dataclasses.field(default_factory=SSDHead)
+ detection_generator: DetectionGenerator = dataclasses.field(
+ default_factory=DetectionGenerator
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
+
+
+@dataclasses.dataclass
+class Losses(hyperparams.Config):
+ loss_weight: float = 1.0
+ box_loss_weight: int = 100
+ attribute_loss_weight: int = 10
+ focal_loss_alpha: float = 0.25
+ focal_loss_gamma: float = 1.5
+ huber_loss_delta: float = 0.1
+ l2_weight_decay: float = 0
+
+
+@dataclasses.dataclass
+class PointPillarsTask(cfg.TaskConfig):
+ """The task config."""
+ model: PointPillarsModel = dataclasses.field(
+ default_factory=PointPillarsModel
+ )
+ use_raw_data: bool = False
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+ init_checkpoint_modules: Union[str, List[str]] = 'all'
+ use_wod_metrics: bool = True
+
+
+@exp_factory.register_config_factory('pointpillars_baseline')
+def pointpillars_baseline() -> cfg.ExperimentConfig:
+ """PointPillars baseline config."""
+ return cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(mixed_precision_dtype='float32'),
+ task=PointPillarsTask(
+ model=PointPillarsModel(
+ classes='vehicle',
+ num_classes=2,
+ min_level=1,
+ max_level=1,
+ anchors=[Anchor(length=1.0, width=1.0)],
+ featurizer=Featurizer(),
+ backbone=Backbone(),
+ decoder=Decoder(),
+ head=SSDHead()
+ ),
+ train_data=DataConfig(is_training=True),
+ validation_data=DataConfig(is_training=False),
+ losses=Losses()
+ ),
+ trainer=cfg.TrainerConfig(
+ train_steps=100,
+ validation_steps=100,
+ validation_interval=10,
+ steps_per_loop=10,
+ summary_interval=10,
+ checkpoint_interval=10,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'decay_steps': 100,
+ 'initial_learning_rate': 0.16,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 10,
+ 'warmup_learning_rate': 0.016
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ])
diff --git a/official/projects/pointpillars/configs/pointpillars_test.py b/official/projects/pointpillars/configs/pointpillars_test.py
new file mode 100644
index 00000000000..40abfccf474
--- /dev/null
+++ b/official/projects/pointpillars/configs/pointpillars_test.py
@@ -0,0 +1,47 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for pointpillars."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.pointpillars.configs import pointpillars as exp_cfg
+
+
+class PointPillarsConfigTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.parameters(
+ ('pointpillars_baseline',),
+ )
+ def test_configs(self, config_name):
+ config = exp_factory.get_exp_config(config_name)
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task, exp_cfg.PointPillarsTask)
+ self.assertIsInstance(config.task.model, exp_cfg.PointPillarsModel)
+ self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
+ self.assertIsInstance(config.task.validation_data, exp_cfg.DataConfig)
+ self.assertIsInstance(config.task.losses, exp_cfg.Losses)
+ self.assertGreater(config.task.model.image.height, 0)
+ self.assertGreater(config.task.model.image.width, 0)
+ self.assertLen(config.task.model.head.attribute_heads, 3)
+ config.task.train_data.is_training = None
+ with self.assertRaises(KeyError):
+ config.validate()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_gpu.yaml b/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_gpu.yaml
new file mode 100644
index 00000000000..05305d7f752
--- /dev/null
+++ b/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_gpu.yaml
@@ -0,0 +1,80 @@
+# Use this config to train PointPillars baseline on vehicle class.
+# experiment type: pointpillars_baseline
+# strategy: train and evaluate on 8 V100 GPUs.
+# mAP: 0.46
+# mAPH: 0.45
+# duraing: 4 hrs/epoch, 8 hrs/50 epochs
+runtime:
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ classes: 'vehicle'
+ num_classes: 2
+ image:
+ x_range: [-76.8, 76.8]
+ y_range: [-76.8, 76.8]
+ z_range: [-3.0, 3.0]
+ resolution: 0.3
+ height: 512
+ width: 512
+ pillars:
+ num_pillars: 24000
+ num_points_per_pillar: 100
+ num_features_per_point: 10
+ min_level: 1
+ max_level: 1
+ anchors:
+ - length: 15.752693
+ width: 6.930973
+ anchor_labeler:
+ match_threshold: 0.6
+ unmatched_threshold: 0.45
+ featurizer:
+ num_blocks: 1
+ num_channels: 64 # C
+ backbone:
+ min_level: 1
+ max_level: 3
+ num_convs: 6
+ detection_generator:
+ pre_nms_score_threshold: 0.05
+ nms_iou_threshold: 0.5
+ max_num_detections: 200
+ train_data:
+ global_batch_size: 16 # 2 * 8
+ dtype: 'float32'
+ shuffle_buffer_size: 256
+ prefetch_buffer_size: 256
+ validation_data:
+ global_batch_size: 32 # 4 * 8
+ dtype: 'float32'
+ shuffle_buffer_size: 256
+ prefetch_buffer_size: 256
+ init_checkpoint: null
+ use_wod_metrics: true
+trainer:
+ # (158081 / 16) = 9880, * 50
+ train_steps: 494000
+ # 39987 / 32
+ validation_steps: 1250
+ validation_interval: 9880
+ steps_per_loop: 9880
+ summary_interval: 9880
+ checkpoint_interval: 9880
+ optimizer_config:
+ optimizer:
+ type: 'sgd'
+ sgd:
+ momentum: 0.9
+ global_clipnorm: 10.0
+ learning_rate:
+ type: 'cosine'
+ cosine:
+ decay_steps: 494000
+ initial_learning_rate: 0.0016
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_learning_rate: 0.00016
+ warmup_steps: 9880 # 1 epoch
diff --git a/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_local.yaml b/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_local.yaml
new file mode 100644
index 00000000000..b13197b54be
--- /dev/null
+++ b/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_local.yaml
@@ -0,0 +1,76 @@
+# This config is used to test PointPillars baseline model training pipeline locally.
+# It runs a very small amount of epochs to verify the functionality of the training pipeline.
+# experiment type: pointpillars_baseline
+# strategy: train and evaluate on CPU
+runtime:
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ classes: 'vehicle'
+ num_classes: 2
+ image:
+ x_range: [-76.8, 76.8]
+ y_range: [-76.8, 76.8]
+ z_range: [-3.0, 3.0]
+ resolution: 0.3
+ height: 512
+ width: 512
+ pillars:
+ num_pillars: 24000
+ num_points_per_pillar: 100
+ num_features_per_point: 10
+ min_level: 1
+ max_level: 3
+ anchors:
+ - length: 15.752693
+ width: 6.930973
+ anchor_labeler:
+ match_threshold: 0.6
+ unmatched_threshold: 0.45
+ featurizer:
+ num_blocks: 1
+ num_channels: 64 # C
+ backbone:
+ min_level: 1
+ max_level: 3
+ num_convs: 6
+ detection_generator:
+ pre_nms_score_threshold: 0.05
+ nms_iou_threshold: 0.5
+ max_num_detections: 200
+ train_data:
+ global_batch_size: 4
+ dtype: 'float32'
+ shuffle_buffer_size: 10
+ prefetch_buffer_size: 10
+ validation_data:
+ global_batch_size: 2
+ dtype: 'float32'
+ shuffle_buffer_size: 10
+ prefetch_buffer_size: 10
+ init_checkpoint: null
+ use_wod_metrics: true
+trainer:
+ train_steps: 2
+ validation_steps: 2
+ steps_per_loop: 2
+ validation_interval: 2
+ summary_interval: 2
+ checkpoint_interval: 2
+ optimizer_config:
+ optimizer:
+ type: 'sgd'
+ sgd:
+ momentum: 0.9
+ global_clipnorm: 10.0
+ learning_rate:
+ type: 'cosine'
+ cosine:
+ decay_steps: 2
+ initial_learning_rate: 0.0032
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_learning_rate: 0.00032
+ warmup_steps: 2
diff --git a/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_tpu.yaml b/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_tpu.yaml
new file mode 100644
index 00000000000..98673b8ec5c
--- /dev/null
+++ b/official/projects/pointpillars/configs/vehicle/pointpillars_3d_baseline_tpu.yaml
@@ -0,0 +1,80 @@
+# Use this config to train PointPillars baseline on vehicle class.
+# experiment type: pointpillars_baseline
+# strategy: train on TPU pod 32 cores, evaluate on 8 V100 GPUs
+# mAP: 0.45
+# mAPH: 0.44
+# duraing: 16 mins/epoch, 15 hrs/50 epochs
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ classes: 'vehicle'
+ num_classes: 2
+ image:
+ x_range: [-76.8, 76.8]
+ y_range: [-76.8, 76.8]
+ z_range: [-3.0, 3.0]
+ resolution: 0.3
+ height: 512
+ width: 512
+ pillars:
+ num_pillars: 24000
+ num_points_per_pillar: 100
+ num_features_per_point: 10
+ min_level: 1
+ max_level: 1
+ anchors:
+ - length: 15.752693
+ width: 6.930973
+ anchor_labeler:
+ match_threshold: 0.6
+ unmatched_threshold: 0.45
+ featurizer:
+ num_blocks: 1
+ num_channels: 64 # C
+ backbone:
+ min_level: 1
+ max_level: 3
+ num_convs: 6
+ detection_generator:
+ pre_nms_score_threshold: 0.05
+ nms_iou_threshold: 0.5
+ max_num_detections: 200
+ train_data:
+ global_batch_size: 64 # 2 * 32, 2 per core, 4x4 df (4 workers, 8 cores per worker)
+ dtype: 'float32'
+ shuffle_buffer_size: 256
+ prefetch_buffer_size: 256
+ validation_data:
+ global_batch_size: 32 # 4 * 8
+ dtype: 'float32'
+ shuffle_buffer_size: 256
+ prefetch_buffer_size: 256
+ init_checkpoint: null
+ use_wod_metrics: true
+trainer:
+ # (158081 / 64) * 50 = 2470 * 50
+ train_steps: 123500
+ # 39987 / 32
+ validation_steps: 1250
+ validation_interval: 2470
+ steps_per_loop: 2470
+ summary_interval: 2470
+ checkpoint_interval: 2470
+ optimizer_config:
+ optimizer:
+ type: 'sgd'
+ sgd:
+ momentum: 0.9
+ global_clipnorm: 10.0
+ learning_rate:
+ type: 'cosine'
+ cosine:
+ decay_steps: 123500
+ initial_learning_rate: 0.0016
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_learning_rate: 0.00016
+ warmup_steps: 2470 # 1 epoch
diff --git a/official/projects/pointpillars/dataloaders/decoders.py b/official/projects/pointpillars/dataloaders/decoders.py
new file mode 100644
index 00000000000..1f82bcb77a3
--- /dev/null
+++ b/official/projects/pointpillars/dataloaders/decoders.py
@@ -0,0 +1,141 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data decoder and parser for Pointpillars."""
+
+from typing import Any, Mapping, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.vision.dataloaders import decoder
+
+
+class ExampleDecoder(decoder.Decoder):
+ """The class to decode preprocessed tf.example to tensors.
+
+ Notations:
+ P: number of pillars in an example
+ N: number of points in a pillar
+ D: number of features in a point
+ M: number of labeled boxes in an example
+ """
+
+ def __init__(self,
+ image_config: cfg.ImageConfig,
+ pillars_config: cfg.PillarsConfig):
+ """Initialize the decoder."""
+ self._feature_description = {
+ 'frame_id': tf.io.FixedLenFeature([], tf.int64),
+ 'pillars': tf.io.FixedLenFeature([], tf.string),
+ 'indices': tf.io.FixedLenFeature([], tf.string),
+ 'bbox/ymin': tf.io.VarLenFeature(tf.float32),
+ 'bbox/xmin': tf.io.VarLenFeature(tf.float32),
+ 'bbox/ymax': tf.io.VarLenFeature(tf.float32),
+ 'bbox/xmax': tf.io.VarLenFeature(tf.float32),
+ 'bbox/class': tf.io.VarLenFeature(tf.int64),
+ 'bbox/heading': tf.io.VarLenFeature(tf.float32),
+ 'bbox/z': tf.io.VarLenFeature(tf.float32),
+ 'bbox/height': tf.io.VarLenFeature(tf.float32),
+ 'bbox/difficulty': tf.io.VarLenFeature(tf.int64),
+ }
+ self._pillars_config = pillars_config
+
+ def _decode_pillars(
+ self, parsed_tensors: Mapping[str, tf.Tensor]
+ ) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Decode pillars from parsed tensors.
+
+ Args:
+ parsed_tensors: A {name: tensor} dict of parsed tensors.
+
+ Returns:
+ pillars: A tensor with shape [P, N, D]
+ indices: A tensor with shape [P, 2]
+ """
+ pillars = tf.io.decode_raw(parsed_tensors['pillars'], tf.float32)
+ pillars = tf.reshape(pillars, [
+ self._pillars_config.num_pillars,
+ self._pillars_config.num_points_per_pillar,
+ self._pillars_config.num_features_per_point
+ ])
+ indices = tf.io.decode_raw(parsed_tensors['indices'], tf.int32)
+ indices = tf.reshape(indices, [self._pillars_config.num_pillars, 2])
+ return pillars, indices
+
+ def _decode_boxes(self, parsed_tensors: Mapping[str, tf.Tensor]) -> tf.Tensor:
+ """Decode boxes from parsed tensors.
+
+ Args:
+ parsed_tensors: A {name: tensor} dict of parsed tensors.
+
+ Returns:
+ boxes: A tensor with shape [M, 4], the last dim represents box yxyx
+ """
+ ymin = parsed_tensors['bbox/ymin']
+ xmin = parsed_tensors['bbox/xmin']
+ ymax = parsed_tensors['bbox/ymax']
+ xmax = parsed_tensors['bbox/xmax']
+ boxes = tf.stack([ymin, xmin, ymax, xmax], axis=-1)
+ return boxes
+
+ def decode(self, serialized_example: Any) -> Mapping[str, Any]:
+ """Decode the serialized example.
+
+ Args:
+ serialized_example: a single serialized tf.Example string.
+
+ Returns:
+ decoded_tensors: a dictionary of tensors with the following fields:
+ - frame_id: an int64 scalar tensor to identify an example.
+ - pillars: a float32 tensor of shape [P, N, D].
+ - indices: an int32 tensor of shape [P, 2].
+ - gt_classes: an int32 tensor of shape [M].
+ - gt_boxes: a float32 tensor of shape [M, 4].
+ - gt_attributes: a dict of (name, [M, 1]) float32 pairs.
+ - gt_difficulty: an int32 tensor of shape [M].
+ """
+ parsed_tensors = tf.io.parse_single_example(
+ serialized=serialized_example, features=self._feature_description)
+
+ # Convert sparse tensor to dense tensor.
+ for k in parsed_tensors:
+ if isinstance(parsed_tensors[k], tf.SparseTensor):
+ parsed_tensors[k] = tf.sparse.to_dense(
+ parsed_tensors[k], default_value=0)
+
+ # Decode features and labels.
+ frame_id = parsed_tensors['frame_id']
+ pillars, indices = self._decode_pillars(parsed_tensors)
+ classes = tf.cast(parsed_tensors['bbox/class'], tf.int32)
+ boxes = self._decode_boxes(parsed_tensors)
+ attr_heading = tf.expand_dims(parsed_tensors['bbox/heading'], axis=1)
+ attr_z = tf.expand_dims(parsed_tensors['bbox/z'], axis=1)
+ attr_height = tf.expand_dims(parsed_tensors['bbox/height'], axis=1)
+ difficulty = tf.cast(parsed_tensors['bbox/difficulty'], tf.int32)
+
+ decoded_tensors = {
+ 'frame_id': frame_id,
+ 'pillars': pillars,
+ 'indices': indices,
+ 'gt_classes': classes,
+ 'gt_boxes': boxes,
+ 'gt_attributes': {
+ 'heading': attr_heading,
+ 'z': attr_z,
+ 'height': attr_height,
+ },
+ 'gt_difficulty': difficulty,
+ }
+ return decoded_tensors
diff --git a/official/projects/pointpillars/dataloaders/decoders_test.py b/official/projects/pointpillars/dataloaders/decoders_test.py
new file mode 100644
index 00000000000..9bd447b1a57
--- /dev/null
+++ b/official/projects/pointpillars/dataloaders/decoders_test.py
@@ -0,0 +1,104 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for decoders."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.dataloaders import decoders
+from official.vision.data.tfrecord_lib import convert_to_feature
+
+
+def _mock_serialized_example(num_pillars, num_points_per_pillar,
+ num_features_per_point, num_boxes):
+ frame_id = np.random.randint(0, 10, dtype=np.int64)
+ pillars = np.random.rand(num_pillars, num_points_per_pillar,
+ num_features_per_point).astype(np.float32)
+ indices = np.random.randint(0, 10, size=[num_pillars, 2], dtype=np.int32)
+ classes = np.random.randint(0, 10, size=[num_boxes], dtype=np.int32)
+ ymin = np.random.rand(num_boxes).astype(np.float32)
+ xmin = np.random.rand(num_boxes).astype(np.float32)
+ ymax = np.random.rand(num_boxes).astype(np.float32)
+ xmax = np.random.rand(num_boxes).astype(np.float32)
+ heading = np.random.rand(num_boxes).astype(np.float32)
+ z = np.random.rand(num_boxes).astype(np.float32)
+ height = np.random.rand(num_boxes).astype(np.float32)
+ difficulty = np.random.randint(0, 10, size=[num_boxes], dtype=np.int32)
+
+ feature = {
+ 'frame_id': convert_to_feature(frame_id, 'int64'),
+ 'pillars': convert_to_feature(pillars.tobytes(), 'bytes'),
+ 'indices': convert_to_feature(indices.tobytes(), 'bytes'),
+ 'bbox/class': convert_to_feature(classes, 'int64_list'),
+ 'bbox/ymin': convert_to_feature(ymin, 'float_list'),
+ 'bbox/xmin': convert_to_feature(xmin, 'float_list'),
+ 'bbox/ymax': convert_to_feature(ymax, 'float_list'),
+ 'bbox/xmax': convert_to_feature(xmax, 'float_list'),
+ 'bbox/heading': convert_to_feature(heading, 'float_list'),
+ 'bbox/z': convert_to_feature(z, 'float_list'),
+ 'bbox/height': convert_to_feature(height, 'float_list'),
+ 'bbox/difficulty': convert_to_feature(difficulty, 'int64_list'),
+ }
+ example = tf.train.Example(features=tf.train.Features(feature=feature))
+ serialized_example = example.SerializeToString()
+ return serialized_example
+
+
+class ExampleDecoderTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.parameters(
+ (2, 10, 1, 1),
+ (3, 2, 10, 10),
+ )
+ def test_shape(self, num_pillars, num_points_per_pillar,
+ num_features_per_point, num_boxes):
+ image_config = cfg.ImageConfig()
+ pillar_config = cfg.PillarsConfig()
+ pillar_config.num_pillars = num_pillars
+ pillar_config.num_points_per_pillar = num_points_per_pillar
+ pillar_config.num_features_per_point = num_features_per_point
+
+ decoder = decoders.ExampleDecoder(image_config, pillar_config)
+ serialized_example = _mock_serialized_example(num_pillars,
+ num_points_per_pillar,
+ num_features_per_point,
+ num_boxes)
+ decoded_example = decoder.decode(
+ tf.convert_to_tensor(value=serialized_example))
+ results = tf.nest.map_structure(lambda x: x.numpy(), decoded_example)
+
+ self.assertAllEqual(
+ (num_pillars, num_points_per_pillar, num_features_per_point),
+ results['pillars'].shape)
+ self.assertAllEqual(
+ (num_pillars, 2), results['indices'].shape)
+ self.assertAllEqual(
+ (num_boxes,), results['gt_classes'].shape)
+ self.assertAllEqual(
+ (num_boxes, 4), results['gt_boxes'].shape)
+ self.assertAllEqual(
+ (num_boxes, 1), results['gt_attributes']['heading'].shape)
+ self.assertAllEqual(
+ (num_boxes, 1), results['gt_attributes']['z'].shape)
+ self.assertAllEqual(
+ (num_boxes, 1), results['gt_attributes']['height'].shape)
+ self.assertAllEqual(
+ (num_boxes,), results['gt_difficulty'].shape)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/dataloaders/parsers.py b/official/projects/pointpillars/dataloaders/parsers.py
new file mode 100644
index 00000000000..34a70ad96b5
--- /dev/null
+++ b/official/projects/pointpillars/dataloaders/parsers.py
@@ -0,0 +1,237 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data decoder and parser for Pointpillars."""
+
+from typing import Any, Dict, List, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.utils import utils
+from official.vision.dataloaders import parser
+from official.vision.ops import anchor
+from official.vision.ops import preprocess_ops
+
+
+class Parser(parser.Parser):
+ """The class to parse decoded tensors to features and labels.
+
+ Notations:
+ N: number of pillars in an example
+ P: number of points in a pillar
+ D: number of features in a point
+ M: number of labeled boxes in an example
+ L: number of anchor boxes per pixel/location
+ """
+
+ def __init__(self,
+ classes: str,
+ min_level: int,
+ max_level: int,
+ image_size: Tuple[int, int],
+ anchor_sizes: List[Tuple[float, float]],
+ match_threshold: float,
+ unmatched_threshold: float,
+ max_num_detections: int,
+ dtype: str):
+ """Initialize the parser.
+
+ Args:
+ classes: A str to indicate which classes should be predicted.
+ min_level: An `int` minimum level of multiscale outputs.
+ max_level: An `int` maximum level of multiscale outputs.
+ image_size: A tuple (height, width) of image size.
+ anchor_sizes: A list of tuple (length, width) of anchor boxes.
+ match_threshold: A float number for positive anchor boxes.
+ unmatched_threshold: A float number for negative anchor boxes.
+ max_num_detections: An `int` number of maximum number of instances in an
+ image. The groundtruth data will be clipped/padded to the number.
+ dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
+ """
+ self._classes = classes
+ self._image_size = image_size
+ self._match_threshold = match_threshold
+ self._unmatched_threshold = unmatched_threshold
+ self._max_num_detections = max_num_detections
+ self._dtype = dtype
+
+ # Generate anchors,
+ # multi-level anchor dict, {level: [h_1, w_l, anchors_per_location * 4]}.
+ self._anchor_boxes = utils.generate_anchors(min_level,
+ max_level,
+ image_size,
+ anchor_sizes)
+
+ def _fix_groundtruths_size(self, groundtruths: Dict[str, Any],
+ size: int) -> Dict[str, Any]:
+ """Clips or pads the first dimension of groundtruths to the fixed size.
+
+ Args:
+ groundtruths: A dictionary of {`str`: `tf.Tensor`} that contains
+ groundtruth annotations of `classes`, `boxes`, `attributes` and
+ `difficulty`.
+ size: An `int` that specifies the expected size of the first dimension of
+ padded tensors.
+
+ Returns:
+ A dictionary of the same keys as input and padded tensors as values.
+ """
+ groundtruths['classes'] = preprocess_ops.clip_or_pad_to_fixed_size(
+ groundtruths['classes'], size, -1)
+ groundtruths['boxes'] = preprocess_ops.clip_or_pad_to_fixed_size(
+ groundtruths['boxes'], size, -1)
+ if 'attributes' in groundtruths:
+ for k, v in groundtruths['attributes'].items():
+ groundtruths['attributes'][
+ k] = preprocess_ops.clip_or_pad_to_fixed_size(v, size, -1)
+ groundtruths['difficulty'] = preprocess_ops.clip_or_pad_to_fixed_size(
+ groundtruths['difficulty'], size, -1)
+ return groundtruths
+
+ def _filter_level_2_labels(
+ self, data: Dict[str, Any]) -> Dict[str, Any]:
+ """Filter labels whose level is 2 [only for training]."""
+ mask = tf.where(data['gt_difficulty'] < 2)
+ data['gt_classes'] = tf.gather_nd(data['gt_classes'], mask)
+ data['gt_boxes'] = tf.gather_nd(data['gt_boxes'], mask)
+ for k, v in data['gt_attributes'].items():
+ data['gt_attributes'][k] = tf.gather_nd(v, mask)
+ data['gt_difficulty'] = tf.gather_nd(data['gt_difficulty'], mask)
+ return data
+
+ def _filter_non_class_labels(
+ self, data: Dict[str, Any]) -> Dict[str, Any]:
+ """Filter labels whose class is not self._classes."""
+ if self._classes == 'all':
+ return data
+ mask = tf.where(data['gt_classes'] == utils.CLASSES[self._classes])
+ data['gt_classes'] = tf.gather_nd(data['gt_classes'], mask)
+ data['gt_boxes'] = tf.gather_nd(data['gt_boxes'], mask)
+ for k, v in data['gt_attributes'].items():
+ data['gt_attributes'][k] = tf.gather_nd(v, mask)
+ data['gt_difficulty'] = tf.gather_nd(data['gt_difficulty'], mask)
+ # Reset 'bbox/class' to 1 to be a binary classification.
+ data['gt_classes'] = tf.ones_like(data['gt_classes'], dtype=tf.int32)
+ return data
+
+ def _parse_feature_and_label(
+ self, data: Dict[str, Any]
+ ) -> Tuple[Dict[str, Any], Dict[str, Any]]:
+ """Parse decoded tensors to features and labels.
+
+ Args:
+ data: A {name: tensor} dict of decoded tensors.
+
+ Returns:
+ features:
+ - pillars: A tensor, shape: [N, P, D], type: self._dtype
+ - indices: A tensor with shape: [N, 2], type: int32
+ labels:
+ - cls_targets: A {level_i: [h_i, w_i, L]} dict, type: float32
+ - box_targets: A {level_i: [h_i, w_i, L * 4]} dict, type: float32
+ - attribute_targets: A {name: {level_i: [h_i, w_i, L * 1]}} dict,
+ type: float32
+ - cls_weights: A flattened tensor with shape [total_num_anchors],
+ total_num_anchors is anchors across all levels, type: float32
+ - box_weights: A flattened tensor with shape [total_num_anchors],
+ total_num_anchors is anchors across all levels, type: float32
+ """
+ data = self._filter_non_class_labels(data)
+
+ pillars = data['pillars']
+ indices = data['indices']
+ classes = data['gt_classes']
+ boxes = data['gt_boxes']
+ attributes = data['gt_attributes']
+
+ # Label anchors,
+ # multi-level labels, {level: [h_l, w_l, ...]}.
+ anchor_labeler = anchor.AnchorLabeler(self._match_threshold,
+ self._unmatched_threshold)
+ (cls_targets, box_targets, att_targets, cls_weights,
+ box_weights) = anchor_labeler.label_anchors(
+ self._anchor_boxes, boxes, tf.expand_dims(classes, axis=1), attributes)
+
+ # Casts input to desired data type.
+ pillars = tf.cast(pillars, dtype=self._dtype)
+
+ # Packs features and labels for model_fn outputs.
+ features = {
+ 'pillars': pillars,
+ 'indices': indices,
+ }
+ labels = {
+ 'cls_targets': cls_targets,
+ 'box_targets': box_targets,
+ 'attribute_targets': att_targets,
+ 'cls_weights': cls_weights,
+ 'box_weights': box_weights,
+ }
+ return features, labels
+
+ def _parse_train_data(
+ self, data: Dict[str, Any]
+ ) -> Tuple[Dict[str, Any], Dict[str, Any]]:
+ """Parse data for training."""
+ # Skip level 2 boxes for training.
+ data = self._filter_level_2_labels(data)
+ return self._parse_feature_and_label(data)
+
+ def _parse_eval_data(
+ self, data: Dict[str, Any]
+ ) -> Tuple[Dict[str, Any], Dict[str, Any]]:
+ """Parse data for evaluation.
+
+ Args:
+ data: A {name: tensor} dict of decoded tensors.
+
+ Returns:
+ Other than features and labels for training, evaluation needs groundtruths
+ to calculate metrics.
+ groundtruths:
+ - frame_id: An int64 tensor to identify an example.
+ - num_detections: An `int` tensor representing the real number of boxes
+ used for computing metrics.
+ - classes: A [max_num_detections] int32 tensor
+ - boxes: A [max_num_detections, 4] float32 tensor
+ - attributes: A {name: [max_num_detections, 1]} float32 dict
+ - difficulty: A [max_num_detections] int32 tensor
+ """
+ features, labels = self._parse_feature_and_label(data)
+
+ # Add for detection generator.
+ labels.update({
+ 'anchor_boxes': self._anchor_boxes,
+ 'image_shape': tf.convert_to_tensor(self._image_size),
+ })
+
+ # Add groundtruth for metric evaluator.
+ # The number of boxes to calculate evaluation metrics, will be used to
+ # remove padding in evaluator.
+ num_detections = tf.minimum(
+ tf.shape(data['gt_classes'])[0], self._max_num_detections)
+ groundtruths = {
+ 'frame_id': data['frame_id'],
+ 'num_detections': num_detections,
+ 'classes': data['gt_classes'],
+ 'boxes': data['gt_boxes'],
+ 'attributes': data['gt_attributes'],
+ 'difficulty': data['gt_difficulty'],
+ }
+ # Fix the size for batching
+ groundtruths = self._fix_groundtruths_size(groundtruths,
+ self._max_num_detections)
+ labels['groundtruths'] = groundtruths
+
+ return features, labels
diff --git a/official/projects/pointpillars/dataloaders/parsers_test.py b/official/projects/pointpillars/dataloaders/parsers_test.py
new file mode 100644
index 00000000000..4768d16f77b
--- /dev/null
+++ b/official/projects/pointpillars/dataloaders/parsers_test.py
@@ -0,0 +1,136 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for parsers."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.dataloaders import parsers
+
+
+def _mock_decoded_example(num_pillars, num_points_per_pillar,
+ num_features_per_point, num_boxes):
+ frame_id = np.random.randint(0, 10, dtype=np.int64)
+ pillars = np.random.rand(num_pillars, num_points_per_pillar,
+ num_features_per_point).astype(np.float32)
+ indices = np.random.randint(0, 10, size=[num_pillars, 2], dtype=np.int32)
+ classes = np.random.randint(0, 10, size=[num_boxes], dtype=np.int32)
+ boxes = np.random.rand(num_boxes, 4).astype(np.float32)
+ heading = np.random.rand(num_boxes, 1).astype(np.float32)
+ z = np.random.rand(num_boxes, 1).astype(np.float32)
+ height = np.random.rand(num_boxes, 1).astype(np.float32)
+ difficulty = np.random.randint(0, 10, size=[num_boxes], dtype=np.int32)
+
+ decoded_example = {
+ 'frame_id': tf.convert_to_tensor(frame_id, dtype=tf.int64),
+ 'pillars': tf.convert_to_tensor(pillars, dtype=tf.float32),
+ 'indices': tf.convert_to_tensor(indices, dtype=tf.int32),
+ 'gt_classes': tf.convert_to_tensor(classes, dtype=tf.int32),
+ 'gt_boxes': tf.convert_to_tensor(boxes, dtype=tf.float32),
+ 'gt_attributes': {
+ 'heading': tf.convert_to_tensor(heading, dtype=tf.float32),
+ 'z': tf.convert_to_tensor(z, dtype=tf.float32),
+ 'height': tf.convert_to_tensor(height, dtype=tf.float32),
+ },
+ 'gt_difficulty': tf.convert_to_tensor(difficulty, dtype=tf.int32),
+ }
+ return decoded_example
+
+
+class ParserTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.parameters(
+ ('all', 1, 10, True),
+ ('vehicle', 10, 2, True),
+ ('pedestrian', 1, 10, False),
+ ('cyclist', 10, 2, False),
+ )
+ def test_shape(self, classes, num_boxes, max_num_boxes, is_training):
+ min_level = 1
+ max_level = 3
+ image_size = (32, 32)
+ anchor_sizes = [(1.1, 2.2)]
+ num_anchors_per_location = len(anchor_sizes)
+ match_threshold = 0.5
+ unmatched_threshold = 0.5
+ parser = parsers.Parser(classes, min_level, max_level, image_size,
+ anchor_sizes, match_threshold, unmatched_threshold,
+ max_num_boxes, 'float32')
+
+ num_pillars = 2
+ num_points_per_pillar = 3
+ num_features_per_point = 4
+ decoded_example = _mock_decoded_example(num_pillars, num_points_per_pillar,
+ num_features_per_point, num_boxes)
+ features, labels = parser.parse_fn(is_training=is_training)(
+ decoded_tensors=decoded_example)
+ features = tf.nest.map_structure(lambda x: x.numpy(), features)
+ labels = tf.nest.map_structure(lambda x: x.numpy(), labels)
+
+ self.assertAllEqual(
+ (num_pillars, num_points_per_pillar, num_features_per_point),
+ features['pillars'].shape)
+ self.assertAllEqual(
+ (num_pillars, 2), features['indices'].shape)
+ total_num_anchors = 0
+ for level in range(min_level, max_level + 1):
+ stride = 2**level
+ h_i = image_size[0] / stride
+ w_i = image_size[1] / stride
+ total_num_anchors += h_i * w_i * num_anchors_per_location
+ self.assertAllEqual((h_i, w_i, num_anchors_per_location),
+ labels['cls_targets'][str(level)].shape)
+ self.assertAllEqual((h_i, w_i, num_anchors_per_location * 4),
+ labels['box_targets'][str(level)].shape)
+ self.assertAllEqual(
+ (h_i, w_i, num_anchors_per_location),
+ labels['attribute_targets']['heading'][str(level)].shape)
+ self.assertAllEqual(
+ (h_i, w_i, num_anchors_per_location),
+ labels['attribute_targets']['height'][str(level)].shape)
+ self.assertAllEqual(
+ (h_i, w_i, num_anchors_per_location),
+ labels['attribute_targets']['z'][str(level)].shape)
+ if not is_training:
+ self.assertAllEqual((h_i, w_i, num_anchors_per_location * 4),
+ labels['anchor_boxes'][str(level)].shape)
+
+ self.assertAllEqual((total_num_anchors,),
+ labels['cls_weights'].shape)
+ self.assertAllEqual((total_num_anchors,),
+ labels['box_weights'].shape)
+
+ if not is_training:
+ self.assertAllEqual((2,), labels['image_shape'].shape)
+ groundtruths = labels['groundtruths']
+ self.assertEmpty(groundtruths['frame_id'].shape)
+ self.assertEmpty(groundtruths['num_detections'].shape)
+ self.assertAllEqual(
+ (max_num_boxes,), groundtruths['classes'].shape)
+ self.assertAllEqual(
+ (max_num_boxes, 4), groundtruths['boxes'].shape)
+ self.assertAllEqual(
+ (max_num_boxes, 1), groundtruths['attributes']['heading'].shape)
+ self.assertAllEqual(
+ (max_num_boxes, 1), groundtruths['attributes']['height'].shape)
+ self.assertAllEqual(
+ (max_num_boxes, 1), groundtruths['attributes']['z'].shape)
+ self.assertAllEqual(
+ (max_num_boxes,), groundtruths['difficulty'].shape)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/modeling/backbones.py b/official/projects/pointpillars/modeling/backbones.py
new file mode 100644
index 00000000000..5ccefb83240
--- /dev/null
+++ b/official/projects/pointpillars/modeling/backbones.py
@@ -0,0 +1,131 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Backbone models for Pointpillars."""
+
+from typing import Any, Mapping, Optional
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import layers
+from official.projects.pointpillars.utils import utils
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Backbone(tf_keras.Model):
+ """The backbone to extract features from BEV pseudo image.
+
+ The implementation is from the network architecture of PointPillars
+ (https://arxiv.org/pdf/1812.05784.pdf). It downsamples the input image
+ through convolutions and output features with multiple levels.
+ """
+
+ def __init__(
+ self,
+ input_specs: tf.TensorShape,
+ min_level: int = 1,
+ max_level: int = 3,
+ num_convs: int = 4,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initialize the backbone.
+
+ The output of the backbone is a multi-level features.
+ 1 <= min_level <= max_level,
+ level_feature_size = input_image_size / 2 ^ level,
+ e.g. input size (32, 32), feature size should be:
+ (32, 32) at level 0, (16, 16) at level 1, (8, 8) at level 2, ...
+ Args:
+ input_specs: A `tf.TensorShape` of the input tensor.
+ min_level: An `int` of min level for output multiscale features.
+ max_level: An `int` of max level for output multiscale features.
+ num_convs: An `int` number of convolution layers in a downsample group.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default to None.
+ **kwargs: Additional keyword arguments to be passed.
+
+ Returns:
+ endpoints: A `dict` of {level: Tensor} pairs for the model output.
+ output_specs: A dict of {level: TensorShape} pairs for the model output.
+ """
+ utils.assert_channels_last()
+
+ self._config_dict = {
+ 'input_specs': input_specs,
+ 'min_level': min_level,
+ 'max_level': max_level,
+ 'num_convs': num_convs,
+ 'kernel_regularizer': kernel_regularizer,
+ }
+ # Onlly allow to output from level 1.
+ if min_level < 1:
+ raise ValueError(
+ 'The min_level must be >= 1, but {} found.'.format(min_level))
+
+ input_channels = input_specs[-1]
+ inputs = tf_keras.Input(shape=input_specs[1:])
+
+ # build the net
+ x = inputs
+ net = {}
+ scale = 1
+ for level in range(1, max_level + 1):
+ x = self._block_group(
+ inputs=x,
+ filters=input_channels * scale)
+ scale *= 2
+ net[level] = x
+
+ # build endpoints
+ endpoints = {}
+ for level in range(min_level, max_level + 1):
+ endpoints[str(level)] = net[level]
+
+ self._output_specs = {l: endpoints[l].get_shape() for l in endpoints}
+ super(Backbone, self).__init__(inputs=inputs, outputs=endpoints)
+
+ def _block_group(self,
+ inputs: tf.Tensor,
+ filters: int) -> tf.Tensor:
+ """A group of convolution layers to downsample inputs.
+
+ Args:
+ inputs: A tensor to be downsampled.
+ filters: An `int` number of filters of convolution.
+
+ Returns:
+ x: A tensor of downsampled feature.
+ """
+ x = layers.ConvBlock(
+ filters=filters,
+ kernel_size=3,
+ strides=2,
+ kernel_regularizer=self._config_dict['kernel_regularizer'])(inputs)
+ for _ in range(1, self._config_dict['num_convs']):
+ x = layers.ConvBlock(
+ filters=filters,
+ kernel_size=3,
+ strides=1,
+ kernel_regularizer=self._config_dict['kernel_regularizer'])(x)
+ return x
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]) -> tf_keras.Model: # pyrefly: ignore[bad-override]
+ return cls(**config)
+
+ @property
+ def output_specs(self) -> Mapping[str, tf.TensorShape]:
+ return self._output_specs
diff --git a/official/projects/pointpillars/modeling/backbones_test.py b/official/projects/pointpillars/modeling/backbones_test.py
new file mode 100644
index 00000000000..5aaf76b8236
--- /dev/null
+++ b/official/projects/pointpillars/modeling/backbones_test.py
@@ -0,0 +1,61 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for backbones."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import backbones
+
+
+class BackboneTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ([1, 32, 32, 3], 1, 1),
+ ([2, 32, 64, 4], 1, 3),
+ )
+ def test_network_creation(self, input_shape, min_level, max_level):
+ batch_size = input_shape[0]
+ inputs = tf_keras.Input(shape=input_shape[1:], batch_size=batch_size)
+ backbone = backbones.Backbone(input_shape, min_level, max_level)
+ endpoints = backbone(inputs)
+ _, h, w, c = input_shape
+ for level in range(min_level, max_level + 1):
+ self.assertAllEqual([
+ batch_size,
+ int(h / 2**level),
+ int(w / 2**level),
+ int(c * 2**(level - 1))
+ ], endpoints[str(level)].shape.as_list())
+
+ def test_serialization(self):
+ kwargs = dict(
+ input_specs=[1, 64, 64, 3],
+ min_level=2,
+ max_level=4,
+ num_convs=3,
+ kernel_regularizer=None,
+ )
+ net = backbones.Backbone(**kwargs)
+ expected_config = kwargs
+ self.assertEqual(net.get_config(), expected_config)
+
+ new_net = backbones.Backbone.from_config(net.get_config())
+ self.assertAllEqual(net.get_config(), new_net.get_config())
+ _ = new_net.to_json()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/modeling/decoders.py b/official/projects/pointpillars/modeling/decoders.py
new file mode 100644
index 00000000000..2534b51d929
--- /dev/null
+++ b/official/projects/pointpillars/modeling/decoders.py
@@ -0,0 +1,105 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Decoder models for Pointpillars."""
+
+from typing import Any, Mapping, Optional
+
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import layers
+from official.projects.pointpillars.utils import utils
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Decoder(tf_keras.Model):
+ """The decoder to process feature maps learned by a backbone.
+
+ The implementation is from the network architecture of PointPillars
+ (https://arxiv.org/pdf/1812.05784.pdf). It upsamples the feature image
+ to the same size and combine them to be the output.
+ """
+
+ def __init__(
+ self,
+ input_specs: Mapping[str, tf.TensorShape],
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initialize the Decoder.
+
+ Args:
+ input_specs: A dict of {level: tf.TensorShape} of the input tensor.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default to None.
+ **kwargs: Additional keyword arguments to be passed.
+
+ Returns:
+ endpoints: A `dict` of {level: Tensor} pairs for the model output.
+ output_specs: A dict of {level: TensorShape} pairs for the model output.
+ """
+ self._config_dict = {
+ 'input_specs': input_specs,
+ 'kernel_regularizer': kernel_regularizer,
+ }
+
+ utils.assert_channels_last()
+
+ # Only allow to process levels learned by a backbone.
+ min_level = int(min(input_specs.keys()))
+ max_level = int(max(input_specs.keys()))
+
+ # Build inputs
+ inputs = {}
+ # Set min_level as the output level.
+ output_level = min_level
+ for level, shape in input_specs.items():
+ # Set num_filters as 2c if the channels of backbone output level is c.
+ if int(level) == output_level:
+ num_filters = 2 * shape[-1]
+ inputs[level] = tf_keras.Input(shape=shape[1:])
+
+ # Build lateral features
+ lateral_feats = {}
+ for level in range(min_level, max_level + 1):
+ lateral_feats[level] = inputs[str(level)]
+
+ # Build scale-up path
+ feats = []
+ for level in range(min_level, max_level + 1):
+ x = layers.ConvBlock(
+ filters=num_filters, # pyrefly: ignore[unbound-name]
+ kernel_size=3,
+ strides=int(2 ** (level - output_level)),
+ use_transpose_conv=True,
+ kernel_regularizer=kernel_regularizer)(
+ lateral_feats[level])
+ feats.append(x)
+
+ # Fuse all levels feature into the output level.
+ endpoints = {}
+ endpoints[str(output_level)] = tf_keras.layers.Concatenate(axis=-1)(feats)
+
+ self._output_specs = {l: endpoints[l].get_shape() for l in endpoints}
+ super(Decoder, self).__init__(inputs=inputs, outputs=endpoints, **kwargs)
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]) -> tf_keras.Model: # pyrefly: ignore[bad-override]
+ return cls(**config)
+
+ @property
+ def output_specs(self) -> Mapping[str, tf.TensorShape]:
+ return self._output_specs
diff --git a/official/projects/pointpillars/modeling/decoders_test.py b/official/projects/pointpillars/modeling/decoders_test.py
new file mode 100644
index 00000000000..c97c5d50177
--- /dev/null
+++ b/official/projects/pointpillars/modeling/decoders_test.py
@@ -0,0 +1,66 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for decoders."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import decoders
+
+
+class DecoderTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ({'1': [1, 32, 32, 3]},
+ 1, 1),
+ ({'1': [1, 32, 32, 3],
+ '2': [1, 16, 16, 6]},
+ 1, 2)
+ )
+ def test_network_creation(self, input_shape, min_level, max_level):
+ """Test if network could be created and infer with expected shapes."""
+ inputs = {}
+ for k, v in input_shape.items():
+ if k == str(min_level):
+ batch_size, height, width, _ = v
+ inputs[k] = tf_keras.Input(shape=v[1:], batch_size=batch_size)
+ decoder = decoders.Decoder(input_shape)
+ endpoints = decoder(inputs)
+
+ self.assertLen(endpoints, 1)
+ self.assertEqual(list(endpoints.keys())[0], str(min_level))
+
+ self.assertIn(str(min_level), endpoints)
+ expected_channels = input_shape[str(min_level)][-1] * 2 * (
+ max_level - min_level + 1)
+ self.assertAllEqual(endpoints[str(min_level)].shape.as_list(),
+ [batch_size, height, width, expected_channels])
+
+ def test_serialization(self):
+ kwargs = dict(
+ input_specs={'1': [1, 64, 64, 3]},
+ kernel_regularizer=None,
+ )
+ net = decoders.Decoder(**kwargs)
+ expected_config = kwargs
+ self.assertEqual(net.get_config(), expected_config)
+
+ new_net = decoders.Decoder.from_config(net.get_config())
+ self.assertAllEqual(net.get_config(), new_net.get_config())
+ _ = new_net.to_json()
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/modeling/factory.py b/official/projects/pointpillars/modeling/factory.py
new file mode 100644
index 00000000000..1f63914756b
--- /dev/null
+++ b/official/projects/pointpillars/modeling/factory.py
@@ -0,0 +1,132 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Factory methods to build models."""
+
+from typing import Mapping, Optional
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.modeling import backbones
+from official.projects.pointpillars.modeling import decoders
+from official.projects.pointpillars.modeling import featurizers
+from official.projects.pointpillars.modeling import heads
+from official.projects.pointpillars.modeling import models
+from official.vision.modeling.layers import detection_generator
+
+
+def build_pointpillars(
+ input_specs: Mapping[str, tf_keras.layers.InputSpec],
+ model_config: cfg.PointPillarsModel,
+ train_batch_size: int,
+ eval_batch_size: int,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> tf_keras.Model:
+ """Build the PointPillars model.
+
+ Args:
+ input_specs: A {name: input_spec} dict used to construct inputs.
+ model_config: A PointPillarsModel config.
+ train_batch_size: An `int` of training batch size per replica.
+ eval_batch_size: An `int` of evaluation batch size per replica.
+ l2_regularizer: A L2 regularizer.
+
+ Returns:
+ model: A PointPillarsModel built from the config.
+ """
+ # Build inputs
+ inputs = {}
+ for k, v in input_specs.items():
+ inputs[k] = tf_keras.Input(shape=v.shape[1:], dtype=v.dtype)
+
+ # Build featurizer
+ image_size = (model_config.image.height, model_config.image.width)
+ pillars_size = input_specs['pillars'].shape[1:]
+ featurizer_config = model_config.featurizer
+ featurizer = featurizers.Featurizer(
+ image_size=image_size,
+ pillars_size=pillars_size,
+ num_blocks=featurizer_config.num_blocks,
+ num_channels=featurizer_config.num_channels,
+ train_batch_size=train_batch_size,
+ eval_batch_size=eval_batch_size,
+ kernel_regularizer=l2_regularizer)
+ image = featurizer(inputs['pillars'], inputs['indices'], training=True)
+
+ # Build backbone
+ backbone_config = model_config.backbone
+ backbone = backbones.Backbone(
+ input_specs=featurizer.output_specs,
+ min_level=backbone_config.min_level,
+ max_level=backbone_config.max_level,
+ num_convs=backbone_config.num_convs,
+ kernel_regularizer=l2_regularizer)
+ encoded_feats = backbone(image)
+
+ # Build decoder
+ decoder = decoders.Decoder(
+ input_specs=backbone.output_specs,
+ kernel_regularizer=l2_regularizer)
+ decoded_feats = decoder(encoded_feats)
+
+ # Build detection head
+ head_config = model_config.head
+ num_anchors_per_location = (len(model_config.anchors))
+ head = heads.SSDHead(
+ num_classes=model_config.num_classes,
+ num_anchors_per_location=num_anchors_per_location,
+ num_params_per_anchor=4,
+ attribute_heads=[
+ attr.as_dict() for attr in (head_config.attribute_heads or [])
+ ],
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ kernel_regularizer=l2_regularizer)
+ scores, boxes, attrs = head(decoded_feats)
+
+ generator_config = model_config.detection_generator
+ detection_generator_obj = detection_generator.MultilevelDetectionGenerator(
+ apply_nms=generator_config.apply_nms,
+ pre_nms_top_k=generator_config.pre_nms_top_k,
+ pre_nms_score_threshold=generator_config.pre_nms_score_threshold,
+ nms_iou_threshold=generator_config.nms_iou_threshold,
+ max_num_detections=generator_config.max_num_detections,
+ nms_version=generator_config.nms_version,
+ use_cpu_nms=generator_config.use_cpu_nms)
+
+ image_size = [model_config.image.height, model_config.image.width]
+ anchor_sizes = [(a.length, a.width) for a in model_config.anchors]
+ model = models.PointPillarsModel(
+ featurizer=featurizer,
+ backbone=backbone,
+ decoder=decoder,
+ head=head,
+ detection_generator=detection_generator_obj,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ image_size=image_size,
+ anchor_sizes=anchor_sizes)
+
+ logging.info('Train/Eval batch size per replica: %d/%d', train_batch_size,
+ eval_batch_size)
+ logging.info('Model inputs: %s', inputs)
+ logging.info('Outputs in training:')
+ logging.info('Featurizer output: %s', image)
+ logging.info('Backbone output: %s', encoded_feats)
+ logging.info('Decoder output: %s', decoded_feats)
+ logging.info('Detection head outputs: scores %s, boxes %s, atrributes %s',
+ scores, boxes, attrs)
+ return model
diff --git a/official/projects/pointpillars/modeling/factory_test.py b/official/projects/pointpillars/modeling/factory_test.py
new file mode 100644
index 00000000000..8528800cfe2
--- /dev/null
+++ b/official/projects/pointpillars/modeling/factory_test.py
@@ -0,0 +1,57 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for factory.py."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.modeling import factory
+from official.projects.pointpillars.modeling import models
+
+
+class PointPillarsBuilderTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (4, 4),
+ (1, 2),
+ (2, 1),
+ )
+ def test_builder(self, train_batch_size, eval_batch_size):
+ model_config = cfg.PointPillarsModel()
+ model_config.anchors = [cfg.Anchor(length=1.0, width=1.0)]
+ pillars_config = model_config.pillars
+ input_specs = {
+ 'pillars':
+ tf_keras.layers.InputSpec(
+ shape=(None, pillars_config.num_pillars,
+ pillars_config.num_points_per_pillar,
+ pillars_config.num_features_per_point)),
+ 'indices':
+ tf_keras.layers.InputSpec(
+ shape=(None, pillars_config.num_pillars, 2), dtype='int32'),
+ }
+ model = factory.build_pointpillars(
+ input_specs, model_config, train_batch_size, eval_batch_size
+ )
+ config = model.get_config()
+ new_model = models.PointPillarsModel.from_config(config)
+ _ = new_model.to_json()
+ self.assertAllEqual(model.get_config(), new_model.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/modeling/featurizers.py b/official/projects/pointpillars/modeling/featurizers.py
new file mode 100644
index 00000000000..2fb5eea82e5
--- /dev/null
+++ b/official/projects/pointpillars/modeling/featurizers.py
@@ -0,0 +1,166 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Featurizer layers for Pointpillars."""
+
+from typing import Any, List, Mapping, Optional, Tuple
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import layers
+from official.projects.pointpillars.utils import utils
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Featurizer(tf_keras.layers.Layer):
+ """The featurizer to convert pillars to a BEV pseudo image.
+
+ The implementation is from the network architecture of PointPillars
+ (https://arxiv.org/pdf/1812.05784.pdf). It extract features from pillar
+ tensors then scatter them back to bird-eye-view (BEV) image using indices.
+
+ Notations:
+ B: batch size
+ H: height of the BEV image
+ W: width of the BEV image
+ P: number of pillars in an example
+ N: number of points in a pillar
+ D: number of features in a point
+ C: channels of the BEV image
+ """
+
+ def __init__(
+ self,
+ image_size: Tuple[int, int],
+ pillars_size: Tuple[int, int, int],
+ train_batch_size: int,
+ eval_batch_size: int,
+ num_blocks: int,
+ num_channels: int,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initialize the featurizer.
+
+ Args:
+ image_size: A [int, int] tuple to define the [H, W] of BEV image.
+ pillars_size: A [int, int, int] tuple to define the [P, N, D] of pillars.
+ train_batch_size: An `int` training batch size per replica.
+ eval_batch_size: An `int` evaluation batch size per replica.
+ num_blocks: An `int` number of blocks for extracting features.
+ num_channels: An `int` number channels of the BEV image.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ block layers. Default to None.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super(Featurizer, self).__init__(**kwargs)
+
+ self._config_dict = {
+ 'image_size': image_size,
+ 'pillars_size': pillars_size,
+ 'train_batch_size': train_batch_size,
+ 'eval_batch_size': eval_batch_size,
+ 'num_blocks': num_blocks,
+ 'num_channels': num_channels,
+ 'kernel_regularizer': kernel_regularizer,
+ }
+ self._image_shape = [image_size[0], image_size[1], num_channels]
+
+ utils.assert_channels_last()
+
+ def build(self, input_specs: List[tf.TensorShape]):
+ """Creates variables for the featurizer."""
+ self._blocks = []
+ for _ in range(self._config_dict['num_blocks']):
+ self._blocks.append(
+ layers.ConvBlock(
+ filters=self._config_dict['num_channels'],
+ kernel_size=1,
+ strides=1,
+ kernel_regularizer=self._config_dict['kernel_regularizer']))
+
+ # These batch_dims are [B, P, 1] tensors that could be created before
+ # call(). They will be used for tf.scatter_nd to convert pillars to BEV
+ # images. Because tf.scatter_nd requires a concrete batch size, we need to
+ # prepare all possibilities of batch size for train, eval and test mode.
+ self._train_batch_dims = self._get_batch_dims(
+ self._config_dict['train_batch_size'])
+ self._eval_batch_dims = self._get_batch_dims(
+ self._config_dict['eval_batch_size'])
+ self._test_batch_dims = self._get_batch_dims(1)
+
+ super(Featurizer, self).build(input_specs)
+
+ def _get_batch_dims(self, batch_size: int) -> tf.Tensor:
+ p = self._config_dict['pillars_size'][0]
+ batch_dims = np.indices([batch_size, p])[0]
+ batch_dims = tf.convert_to_tensor(batch_dims, dtype=tf.int32)
+ batch_dims = tf.expand_dims(batch_dims, axis=-1)
+ return batch_dims
+
+ def _get_batch_size_and_dims(self, # pytype: disable=annotation-type-mismatch
+ training: bool = None) -> Tuple[int, tf.Tensor]: # pyrefly: ignore[bad-function-definition]
+ # We use training as a ternary indicator, None for test mode.
+ # Test mode will be used for saving model and model inference.
+ if training is None:
+ batch_size = 1
+ batch_dims = self._test_batch_dims
+ else:
+ if training:
+ batch_size = self._config_dict['train_batch_size']
+ batch_dims = self._train_batch_dims
+ else:
+ batch_size = self._config_dict['eval_batch_size']
+ batch_dims = self._eval_batch_dims
+ return batch_size, batch_dims
+
+ def call(self, # pytype: disable=annotation-type-mismatch
+ pillars: tf.Tensor,
+ indices: tf.Tensor,
+ training: bool = None) -> tf.Tensor: # pyrefly: ignore[bad-function-definition]
+ """Forward pass of the featurizer."""
+ # Add batch index to pillar indices.
+ # (B, P, 1)
+ batch_size, batch_dims = self._get_batch_size_and_dims(training)
+ # (B, P, 3)
+ batch_indices = tf.concat([batch_dims, indices], axis=-1)
+
+ # Extract features from pillars.
+ # (B, P, N, D)
+ x = pillars
+ # (B, P, N, C)
+ for block in self._blocks:
+ x = block(x)
+ # (B, P, C)
+ x = tf.reduce_max(x, axis=2, keepdims=False)
+
+ # Scatter pillars back to form a BEV image.
+ # (B, H, W, C)
+ image = tf.scatter_nd(
+ batch_indices,
+ x,
+ shape=[batch_size] + self._image_shape)
+ self._output_specs = image.get_shape()
+ return image
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]) -> tf_keras.Model:
+ return cls(**config) # pyrefly: ignore[bad-return]
+
+ @property
+ def output_specs(self) -> tf.TensorShape:
+ return self._output_specs
diff --git a/official/projects/pointpillars/modeling/featurizers_test.py b/official/projects/pointpillars/modeling/featurizers_test.py
new file mode 100644
index 00000000000..4caafb17473
--- /dev/null
+++ b/official/projects/pointpillars/modeling/featurizers_test.py
@@ -0,0 +1,81 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for backbones."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import featurizers
+
+
+class FeaturizerTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ([32, 32], [16, 4, 2], 4, 2, 1),
+ ([32, 16], [1, 3, 1], 2, 2, 3),
+ )
+ def test_network_creation(self, image_size, pillars_size, train_batch_size,
+ eval_batch_size, num_blocks):
+ num_channels = 3
+ h, w = image_size
+ n, _, _ = pillars_size
+ featurizer = featurizers.Featurizer(image_size, pillars_size,
+ train_batch_size, eval_batch_size,
+ num_blocks, num_channels)
+
+ # Train mode.
+ pillars = tf_keras.Input(shape=pillars_size, batch_size=train_batch_size)
+ indices = tf_keras.Input(
+ shape=[n, 2], batch_size=train_batch_size, dtype=tf.int32)
+ image = featurizer(pillars, indices, training=True)
+ self.assertAllEqual([train_batch_size, h, w, num_channels],
+ image.shape.as_list())
+
+ # Evaluation mode.
+ pillars = tf_keras.Input(shape=pillars_size, batch_size=eval_batch_size)
+ indices = tf_keras.Input(
+ shape=[n, 2], batch_size=eval_batch_size, dtype=tf.int32)
+ image = featurizer(pillars, indices, training=False)
+ self.assertAllEqual([eval_batch_size, h, w, num_channels],
+ image.shape.as_list())
+
+ # Test mode, batch size must be 1.
+ pillars = tf_keras.Input(shape=pillars_size, batch_size=1)
+ indices = tf_keras.Input(
+ shape=[n, 2], batch_size=1, dtype=tf.int32)
+ image = featurizer(pillars, indices, training=None)
+ self.assertAllEqual([1, h, w, num_channels],
+ image.shape.as_list())
+
+ def test_serialization(self):
+ kwargs = dict(
+ image_size=[4, 4],
+ pillars_size=[4, 5, 6],
+ train_batch_size=4,
+ eval_batch_size=2,
+ num_blocks=3,
+ num_channels=4,
+ kernel_regularizer=None,
+ )
+ net = featurizers.Featurizer(**kwargs)
+ expected_config = kwargs
+ self.assertEqual(net.get_config(), expected_config)
+
+ new_net = featurizers.Featurizer.from_config(net.get_config())
+ self.assertAllEqual(net.get_config(), new_net.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/modeling/heads.py b/official/projects/pointpillars/modeling/heads.py
new file mode 100644
index 00000000000..6e72cb0caad
--- /dev/null
+++ b/official/projects/pointpillars/modeling/heads.py
@@ -0,0 +1,175 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Head layers for Pointpillars."""
+
+from typing import Any, Dict, List, Mapping, Optional, Tuple
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import layers
+from official.projects.pointpillars.utils import utils
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class SSDHead(tf_keras.layers.Layer):
+ """A SSD head for PointPillars detection."""
+
+ def __init__(
+ self,
+ num_classes: int,
+ num_anchors_per_location: int,
+ num_params_per_anchor: int = 4,
+ attribute_heads: Optional[List[Dict[str, Any]]] = None,
+ min_level: int = 1,
+ max_level: int = 3,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initialize the SSD Head.
+
+ Args:
+ num_classes: An `int` number of classes to predict.
+ num_anchors_per_location: An `int` number of anchors per location.
+ num_params_per_anchor: An `int` number of parameters per anchor.
+ attribute_heads: If not None, a list that contains a dict for each
+ additional attribute head. Each dict consists of 3 key-value pairs:
+ `name`, `type` ('regression' or 'classification'), and `size` (number
+ of predicted values for each instance).
+ min_level: An `int` of min level for output mutiscale features.
+ max_level: An `int` of max level for output mutiscale features.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default to None.
+ **kwargs: Additional keyword arguments to be passed.
+
+ Returns:
+ endpoints: A `dict` of {level: Tensor} pairs for the model output.
+ output_specs: A dict of {level: TensorShape} pairs for the model output.
+ """
+ super(SSDHead, self).__init__(**kwargs)
+ self._config_dict = {
+ 'num_classes': num_classes,
+ 'num_anchors_per_location': num_anchors_per_location,
+ 'num_params_per_anchor': num_params_per_anchor,
+ 'attribute_heads': attribute_heads,
+ 'min_level': min_level,
+ 'max_level': max_level,
+ 'kernel_regularizer': kernel_regularizer,
+ }
+
+ utils.assert_channels_last()
+
+ def build(self, input_specs: Mapping[str, tf.TensorShape]):
+ self._decoder_output_level = int(min(input_specs.keys()))
+ if self._config_dict['min_level'] < self._decoder_output_level:
+ raise ValueError('The min_level should be >= decoder output '
+ 'level, but {} < {}'.format(
+ self._config_dict['min_level'],
+ self._decoder_output_level))
+
+ # Multi-level convs.
+ # Set num_filters as the one of decoder's output level.
+ num_filters = input_specs[str(self._decoder_output_level)].as_list()[-1]
+ self._convs = {}
+ for level in range(self._decoder_output_level + 1,
+ self._config_dict['max_level'] + 1):
+ self._convs[str(level)] = layers.ConvBlock(
+ filters=num_filters,
+ kernel_size=3,
+ strides=2,
+ kernel_regularizer=self._config_dict['kernel_regularizer'])
+
+ # Detection convs, share weights across multi levels.
+ self._classifier = tf_keras.layers.Conv2D(
+ filters=(self._config_dict['num_classes'] *
+ self._config_dict['num_anchors_per_location']),
+ kernel_size=3,
+ strides=1,
+ padding='same',
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=1e-5),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_initializer=tf.constant_initializer(-np.log((1 - 0.01) / 0.01)))
+ self._box_regressor = tf_keras.layers.Conv2D(
+ filters=(self._config_dict['num_params_per_anchor'] *
+ self._config_dict['num_anchors_per_location']),
+ kernel_size=3,
+ strides=1,
+ padding='same',
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=1e-5),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_initializer=tf.zeros_initializer())
+ if self._config_dict['attribute_heads']:
+ self._att_predictors = {}
+ for att_config in self._config_dict['attribute_heads']:
+ att_name = att_config['name']
+ att_type = att_config['type']
+ att_size = att_config['size']
+ if att_type != 'regression':
+ raise ValueError('Unsupported head type: {}'.format(att_type))
+ self._att_predictors[att_name] = tf_keras.layers.Conv2D(
+ filters=(att_size * self._config_dict['num_anchors_per_location']),
+ kernel_size=3,
+ strides=1,
+ padding='same',
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=1e-5),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_initializer=tf.zeros_initializer())
+
+ super(SSDHead, self).build(input_specs)
+
+ def call(
+ self, inputs: Mapping[str, tf.Tensor]
+ ) -> Tuple[Dict[str, Any], Dict[str, Any], Dict[Any, Dict[str, Any]]]:
+ # Build multi level features.
+ feats = {}
+ for level in range(self._decoder_output_level,
+ self._config_dict['max_level'] + 1):
+ if level == self._decoder_output_level:
+ x = inputs[str(level)]
+ else:
+ x = self._convs[str(level)](feats[level - 1])
+ feats[level] = x
+
+ # Get multi level detection.
+ scores = {}
+ boxes = {}
+ if self._config_dict['attribute_heads']:
+ attributes = {
+ att_config['name']: {}
+ for att_config in self._config_dict['attribute_heads']
+ }
+ else:
+ attributes = {}
+
+ for level in range(self._config_dict['min_level'],
+ self._config_dict['max_level'] + 1):
+ # The branch to predict box classes.
+ scores[str(level)] = self._classifier(feats[level])
+ # The branch to predict boxes.
+ boxes[str(level)] = self._box_regressor(feats[level])
+ # The branches to predict box attributes.
+ if self._config_dict['attribute_heads']:
+ for att_config in self._config_dict['attribute_heads']:
+ att_name = att_config['name']
+ attributes[att_name][str(level)] = self._att_predictors[att_name](
+ feats[level])
+
+ return scores, boxes, attributes
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]) -> tf_keras.layers.Layer:
+ return cls(**config)
diff --git a/official/projects/pointpillars/modeling/heads_test.py b/official/projects/pointpillars/modeling/heads_test.py
new file mode 100644
index 00000000000..570988363f8
--- /dev/null
+++ b/official/projects/pointpillars/modeling/heads_test.py
@@ -0,0 +1,90 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for decoders."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import heads
+
+
+class SSDHeadTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (2, [], 1, 1),
+ (3, [{'name': 'z', 'type': 'regression', 'size': 1}], 1, 3))
+ def test_network_creation(self, num_classes, attribute_heads, min_level,
+ max_level):
+ """Test if network could be created and infer with expected shapes."""
+ # Fix the input shape, anchor size and num of conv filters.
+ n, h, w, c = 1, 32, 32, 4
+ num_anchors_per_location = 3
+ num_params_per_anchor = 4
+ inputs = {'1': tf_keras.Input(shape=[h, w, c], batch_size=n)}
+
+ head = heads.SSDHead(num_classes, num_anchors_per_location,
+ num_params_per_anchor, attribute_heads, min_level,
+ max_level)
+ scores, boxes, attributes = head(inputs)
+ for level in range(min_level, max_level+1):
+ self.assertIn(str(level), scores)
+ self.assertIn(str(level), boxes)
+ scale = 2**(level - min_level)
+ self.assertAllEqual(scores[str(level)].shape.as_list(), [
+ n,
+ int(h / scale),
+ int(w / scale), num_classes * num_anchors_per_location
+ ])
+ self.assertAllEqual(boxes[str(level)].shape.as_list(), [
+ n,
+ int(h / scale),
+ int(w / scale), num_params_per_anchor * num_anchors_per_location
+ ])
+ for attr_head in attribute_heads:
+ name = attr_head['name']
+ size = attr_head['size']
+ self.assertIn(name, attributes)
+ attr = attributes[name]
+ for level in range(min_level, max_level+1):
+ self.assertIn(str(level), attr)
+ scale = 2**(level - min_level)
+ self.assertAllEqual(attr[str(level)].shape.as_list(), [
+ n,
+ int(h / scale),
+ int(w / scale), size * num_anchors_per_location
+ ])
+
+ def test_serialization(self):
+ kwargs = dict(
+ num_classes=2,
+ num_anchors_per_location=3,
+ num_params_per_anchor=4,
+ attribute_heads=[
+ {'name': 'z', 'type': 'regression', 'size': 1},
+ ],
+ min_level=1,
+ max_level=3,
+ kernel_regularizer=None
+ )
+ net = heads.SSDHead(**kwargs)
+ expected_config = kwargs
+ self.assertEqual(net.get_config(), expected_config)
+
+ new_net = heads.SSDHead.from_config(net.get_config())
+ self.assertAllEqual(net.get_config(), new_net.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/modeling/layers.py b/official/projects/pointpillars/modeling/layers.py
new file mode 100644
index 00000000000..8450979d65d
--- /dev/null
+++ b/official/projects/pointpillars/modeling/layers.py
@@ -0,0 +1,151 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Featurizer layers for Pointpillars."""
+
+from typing import Any, Mapping, Optional
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import tf_utils
+from official.projects.pointpillars.utils import utils
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class ConvBlock(tf_keras.layers.Layer):
+ """A conv2d followed by a norm then an activation."""
+
+ def __init__(
+ self,
+ filters: int,
+ kernel_size: int,
+ strides: int,
+ use_transpose_conv: bool = False,
+ kernel_initializer: Optional[tf_keras.initializers.Initializer] = tf.keras
+ .initializers.VarianceScaling(),
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ use_bias: bool = False,
+ bias_initializer: Optional[tf_keras.initializers.Initializer] = tf.keras
+ .initializers.Zeros(),
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ use_sync_bn: bool = True,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ bn_trainable: bool = True,
+ activation: str = 'relu',
+ **kwargs):
+ """Initialize a block with conv, bn and activation.
+
+ Args:
+ filters: An int number of filters of the conv layer.
+ kernel_size: An int number of kernel size of the conv layer.
+ strides: An int number of strides of the conv layer.
+ use_transpose_conv: A bool for wether to use transpose conv or not.
+ kernel_initializer: A tf Initializer object for the conv layer.
+ kernel_regularizer: A tf Regularizer object for the conv layer.
+ use_bias: A bool for whether to use bias for the conv layer.
+ bias_initializer: A tf Initializer object for the conv layer bias.
+ bias_regularizer: A tf Regularizer object for the conv layer bias.
+ use_sync_bn: A bool for wether to use synchronized batch normalization.
+ norm_momentum: A float of normalization momentum for the moving average.
+ norm_epsilon: A float added to variance to avoid dividing by zero.
+ bn_trainable: A bool that indicates whether batch norm layers should be
+ trainable. Default to True.
+ activation: A str name of the activation function.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+
+ super(ConvBlock, self).__init__(**kwargs)
+
+ self._filters = filters
+ self._kernel_size = kernel_size
+ self._strides = strides
+ self._use_transpose_conv = use_transpose_conv
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._use_bias = use_bias
+ self._bias_initializer = bias_initializer
+ self._bias_regularizer = bias_regularizer
+ self._use_sync_bn = use_sync_bn
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._bn_trainable = bn_trainable
+ self._activation = activation
+ self._activation_fn = tf_utils.get_activation(activation)
+
+ utils.assert_channels_last()
+
+ def build(self, input_shape: tf.TensorShape):
+ """Creates variables for the block."""
+ # Config conv
+ if self._use_transpose_conv:
+ conv_op = tf_keras.layers.Conv2DTranspose
+ else:
+ conv_op = tf_keras.layers.Conv2D
+ conv_kwargs = {
+ 'filters': self._filters,
+ 'kernel_size': self._kernel_size,
+ 'strides': self._strides,
+ 'padding': 'same',
+ 'use_bias': self._use_bias,
+ 'kernel_initializer': self._kernel_initializer,
+ 'bias_initializer': self._bias_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_regularizer': self._bias_regularizer,
+ }
+ self._conv = conv_op(**conv_kwargs)
+
+ # Config norm
+ if self._use_sync_bn:
+ bn_op = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ bn_op = tf_keras.layers.BatchNormalization
+ bn_kwargs = {
+ 'axis': -1,
+ 'momentum': self._norm_momentum,
+ 'epsilon': self._norm_epsilon,
+ 'trainable': self._bn_trainable,
+ }
+ self._norm = bn_op(**bn_kwargs)
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ """Forward pass of the block."""
+ x = inputs
+ x = self._conv(x)
+ x = self._norm(x)
+ outputs = self._activation_fn(x)
+ return outputs
+
+ def get_config(self) -> Mapping[str, Any]:
+ config = {
+ 'filters': self._filters,
+ 'kernel_size': self._kernel_size,
+ 'strides': self._strides,
+ 'use_transpose_conv': self._use_transpose_conv,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'use_bias': self._use_bias,
+ 'bias_initializer': self._bias_initializer,
+ 'bias_regularizer': self._bias_regularizer,
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'bn_trainable': self._bn_trainable,
+ 'activation': self._activation,
+ }
+ return config
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]) -> tf_keras.Model:
+ return cls(**config) # pyrefly: ignore[bad-return]
diff --git a/official/projects/pointpillars/modeling/layers_test.py b/official/projects/pointpillars/modeling/layers_test.py
new file mode 100644
index 00000000000..50fc9fb8568
--- /dev/null
+++ b/official/projects/pointpillars/modeling/layers_test.py
@@ -0,0 +1,75 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for backbones."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.modeling import layers
+
+
+class ConvBlockTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ([1, 8, 8, 3], 4, 1, False),
+ ([1, 8, 8, 3], 4, 2, False),
+ ([1, 8, 8, 3], 2, 1, True),
+ ([1, 8, 8, 3], 2, 2, True),
+ )
+ def test_creation(self, input_shape, filters, strides,
+ use_transpose_conv):
+ kernel_size = 3
+ n, h, w, _ = input_shape
+ inputs = tf_keras.Input(shape=input_shape[1:], batch_size=n)
+ block = layers.ConvBlock(filters, kernel_size, strides, use_transpose_conv)
+ outputs = block(inputs)
+
+ if not use_transpose_conv:
+ if strides == 1:
+ self.assertAllEqual([n, h, w, filters], outputs.shape.as_list())
+ elif strides == 2:
+ self.assertAllEqual([n, h/2, w/2, filters], outputs.shape.as_list())
+ else:
+ if strides == 1:
+ self.assertAllEqual([n, h, w, filters], outputs.shape.as_list())
+ elif strides == 2:
+ self.assertAllEqual([n, h*2, w*2, filters], outputs.shape.as_list())
+
+ def test_serialization(self):
+ kwargs = dict(
+ filters=3,
+ kernel_size=3,
+ strides=1,
+ use_transpose_conv=False,
+ kernel_initializer=None,
+ kernel_regularizer=None,
+ use_bias=False,
+ bias_initializer=None,
+ bias_regularizer=None,
+ use_sync_bn=True,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ bn_trainable=True,
+ activation='relu',
+ )
+ net = layers.ConvBlock(**kwargs)
+ expected_config = kwargs
+ self.assertEqual(net.get_config(), expected_config)
+
+ new_net = layers.ConvBlock.from_config(net.get_config())
+ self.assertAllEqual(net.get_config(), new_net.get_config())
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/modeling/models.py b/official/projects/pointpillars/modeling/models.py
new file mode 100644
index 00000000000..2f5deb83ada
--- /dev/null
+++ b/official/projects/pointpillars/modeling/models.py
@@ -0,0 +1,217 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PointPillars Model."""
+from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
+
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.utils import utils
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class PointPillarsModel(tf_keras.Model):
+ """The PointPillars model class."""
+
+ def __init__(self,
+ featurizer: tf_keras.layers.Layer,
+ backbone: tf_keras.Model,
+ decoder: tf_keras.Model,
+ head: tf_keras.layers.Layer,
+ detection_generator: tf_keras.layers.Layer,
+ min_level: int,
+ max_level: int,
+ image_size: Tuple[int, int],
+ anchor_sizes: List[Tuple[float, float]],
+ **kwargs):
+ """Initialize the model class.
+
+ Args:
+ featurizer: A `tf_keras.layers.Layer` to extract features from pillars.
+ backbone: A `tf_keras.Model` to downsample feature images.
+ decoder: A `tf_keras.Model` to upsample feature images.
+ head: A `tf_keras.layers.Layer` to predict targets.
+ detection_generator: A `tf_keras.layers.Layer` to generate detections.
+ min_level: An `int` minimum level of multiscale outputs.
+ max_level: An `int` maximum level of multiscale outputs.
+ image_size: A tuple (height, width) of image size.
+ anchor_sizes: A list of tuple (length, width) of anchor boxes.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ super(PointPillarsModel, self).__init__(**kwargs)
+ self._featurizer = featurizer
+ self._backbone = backbone
+ self._decoder = decoder
+ self._head = head
+ self._detection_generator = detection_generator
+ self._min_level = min_level
+ self._max_level = max_level
+ self._image_size = image_size
+ self._anchor_sizes = anchor_sizes
+
+ def generate_outputs(
+ self,
+ raw_scores: Dict[str, tf.Tensor],
+ raw_boxes: Dict[str, tf.Tensor],
+ raw_attributes: Dict[str, Dict[str, tf.Tensor]],
+ image_shape: Optional[tf.Tensor] = None,
+ anchor_boxes: Optional[Mapping[str, tf.Tensor]] = None,
+ generate_detections: bool = False) -> Mapping[str, Any]:
+ if not raw_attributes:
+ raise ValueError('PointPillars model needs attribute heads.')
+ # Clap heading to [-pi, pi]
+ if 'heading' in raw_attributes:
+ raw_attributes['heading'] = utils.clip_heading(raw_attributes['heading'])
+
+ outputs = {
+ 'cls_outputs': raw_scores,
+ 'box_outputs': raw_boxes,
+ 'attribute_outputs': raw_attributes,
+ }
+ # Cast raw prediction to float32 for loss calculation.
+ outputs = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), outputs)
+ if not generate_detections:
+ return outputs
+
+ if image_shape is None:
+ raise ValueError('Image_shape should not be None for evaluation.')
+ if anchor_boxes is None:
+ # Generate anchors if needed.
+ anchor_boxes = utils.generate_anchors(
+ self._min_level,
+ self._max_level,
+ self._image_size,
+ self._anchor_sizes,
+ )
+ for l in anchor_boxes:
+ anchor_boxes[l] = tf.tile(
+ tf.expand_dims(anchor_boxes[l], axis=0),
+ [tf.shape(image_shape)[0], 1, 1, 1])
+
+ # Generate detected boxes.
+ if not self._detection_generator.get_config()['apply_nms']:
+ raise ValueError('An NMS algorithm is required for detection generator')
+ detections = self._detection_generator(raw_boxes, raw_scores,
+ anchor_boxes, image_shape,
+ raw_attributes)
+ outputs.update({
+ 'boxes': detections['detection_boxes'],
+ 'scores': detections['detection_scores'],
+ 'classes': detections['detection_classes'],
+ 'num_detections': detections['num_detections'],
+ 'attributes': detections['detection_attributes'],
+ })
+ return outputs
+
+ def call(self, # pytype: disable=annotation-type-mismatch,signature-mismatch
+ pillars: tf.Tensor,
+ indices: tf.Tensor,
+ image_shape: Optional[tf.Tensor] = None,
+ anchor_boxes: Optional[Mapping[str, tf.Tensor]] = None,
+ training: bool = None) -> Mapping[str, Any]: # pyrefly: ignore[bad-function-definition]
+ """Forward pass of the model.
+
+ Notation:
+ B: batch size
+ H_i: image height at level i
+ W_i: image width at level i
+ D: number of anchors per location
+ C: number of classes to predict
+ M: number of detected boxes
+ T: attribute size
+ P: number of pillars in an example
+ N: number of points in a pillar
+ D: number of features in a point
+
+ Args:
+ pillars: A tensor with shape [B, P, N, D].
+ indices: A tensor with shape [B, P, 2].
+ image_shape: A tensor with shape [B, 2] representing size of images.
+ anchor_boxes: A {level: tensor} dict contains multi level anchor boxes.
+ - key: a `str` level.
+ - value: a tensor with shape [B, H_i, W_i, 4 * D].
+ training: A `bool` indicating whether it's in training mode.
+
+ Returns:
+ cls_outputs: A {level: tensor} dict, tensor shape is [B, H_i, W_i, C * D].
+ box_outputs: A {level: tensor} dict, tensor shape is [B, H_i, W_i, 4 * D].
+ attribute_outputs: A {name: {level: tensor}} dict, tensor shape is
+ [B, H_i, W_i, T * D].
+
+ (Below are only for evaluation mode)
+ num_detections: A `int` tensor represent number of detected boxes.
+ boxes: A tensor with shape [B, M, 4].
+ scores: A tensor with shape [B, M].
+ classes: A tensor with shape [B, M].
+ attributes: A {name: tensor} dict, tensor shape is [B, M, T].
+
+ """
+ images = self.featurizer(pillars, indices, training=training)
+ features = self.backbone(images)
+ features = self.decoder(features)
+ raw_scores, raw_boxes, raw_attributes = self.head(features)
+ return self.generate_outputs(raw_scores=raw_scores,
+ raw_boxes=raw_boxes,
+ raw_attributes=raw_attributes,
+ image_shape=image_shape,
+ anchor_boxes=anchor_boxes,
+ generate_detections=not training)
+
+ @property
+ def checkpoint_items(
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(featurizer=self.featurizer,
+ backbone=self.backbone,
+ decoder=self.decoder,
+ head=self.head)
+ return items
+
+ @property
+ def featurizer(self) -> tf_keras.layers.Layer:
+ return self._featurizer
+
+ @property
+ def backbone(self) -> tf_keras.Model:
+ return self._backbone
+
+ @property
+ def decoder(self) -> tf_keras.Model:
+ return self._decoder
+
+ @property
+ def head(self) -> tf_keras.layers.Layer:
+ return self._head
+
+ @property
+ def detection_generator(self) -> tf_keras.layers.Layer:
+ return self._detection_generator
+
+ def get_config(self) -> Mapping[str, Any]:
+ config_dict = {
+ 'featurizer': self._featurizer,
+ 'backbone': self._backbone,
+ 'decoder': self._decoder,
+ 'head': self._head,
+ 'detection_generator': self._detection_generator,
+ 'min_level': self._min_level,
+ 'max_level': self._max_level,
+ 'image_size': self._image_size,
+ 'anchor_sizes': self._anchor_sizes,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]) -> tf_keras.Model: # pyrefly: ignore[bad-override]
+ return cls(**config)
diff --git a/official/projects/pointpillars/modeling/models_test.py b/official/projects/pointpillars/modeling/models_test.py
new file mode 100644
index 00000000000..a469224b783
--- /dev/null
+++ b/official/projects/pointpillars/modeling/models_test.py
@@ -0,0 +1,184 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for PointPillars models."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from tensorflow.python.distribute import combinations
+from tensorflow.python.distribute import strategy_combinations
+from official.projects.pointpillars.modeling import backbones
+from official.projects.pointpillars.modeling import decoders
+from official.projects.pointpillars.modeling import featurizers
+from official.projects.pointpillars.modeling import heads
+from official.projects.pointpillars.modeling import models
+from official.projects.pointpillars.utils import utils
+from official.vision.modeling.layers import detection_generator
+
+
+class PointpillarsTest(parameterized.TestCase, tf.test.TestCase):
+
+ @combinations.generate(
+ combinations.combine(
+ strategy=[
+ strategy_combinations.cloud_tpu_strategy,
+ strategy_combinations.one_device_strategy,
+ strategy_combinations.mirrored_strategy_with_one_gpu,
+ strategy_combinations.mirrored_strategy_with_two_gpus,
+ ],
+ training=[True, False],
+ ))
+ def test_all(self, strategy, training):
+ tf_keras.backend.set_image_data_format('channels_last')
+ num_classes = 2
+
+ h, w, c = 8, 8, 2
+ n, p, d = 2, 3, 4
+ image_size = [h, w]
+ pillars_size = [n, p, d]
+ indices_size = [n, 2]
+ attribute_heads = [{'name': 'heading', 'type': 'regression', 'size': 1}]
+
+ min_level = 1
+ max_level = 2
+
+ anchor_sizes = [(1.1, 1.1)]
+ num_anchors_per_location = len(anchor_sizes)
+
+ global_batch_size = 4
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+ batch_size = int(global_batch_size / num_replicas)
+ pillars = tf_keras.Input(shape=pillars_size, batch_size=batch_size)
+ indices = tf_keras.Input(
+ shape=indices_size, batch_size=batch_size, dtype=tf.int32)
+ image_shape = tf.tile(tf.expand_dims([h, w], axis=0), [batch_size, 1])
+ max_num_detections = 4
+
+ # Test model creation.
+ with strategy.scope():
+ anchor_boxes = utils.generate_anchors(min_level,
+ max_level,
+ image_size,
+ anchor_sizes)
+ for l in anchor_boxes:
+ anchor_boxes[l] = tf.tile(
+ tf.expand_dims(anchor_boxes[l], axis=0), [batch_size, 1, 1, 1])
+
+ featurizer = featurizers.Featurizer(
+ image_size=image_size,
+ pillars_size=pillars_size,
+ train_batch_size=batch_size,
+ eval_batch_size=batch_size,
+ num_blocks=3,
+ num_channels=c
+ )
+ image = featurizer(pillars, indices, training)
+ backbone = backbones.Backbone(
+ input_specs=featurizer.output_specs,
+ min_level=min_level,
+ max_level=max_level,
+ num_convs=3
+ )
+ encoded_feats = backbone(image)
+ decoder = decoders.Decoder(
+ input_specs=backbone.output_specs)
+ decoded_feats = decoder(encoded_feats)
+ head = heads.SSDHead(
+ num_classes=num_classes,
+ num_anchors_per_location=num_anchors_per_location,
+ num_params_per_anchor=4,
+ attribute_heads=attribute_heads,
+ min_level=min_level,
+ max_level=max_level
+ )
+ _ = head(decoded_feats)
+ generator = detection_generator.MultilevelDetectionGenerator(
+ max_num_detections=max_num_detections,
+ nms_version='v1',
+ use_cpu_nms=True,
+ soft_nms_sigma=0.1)
+ model = models.PointPillarsModel(
+ featurizer=featurizer,
+ backbone=backbone,
+ decoder=decoder,
+ head=head,
+ detection_generator=generator,
+ min_level=min_level,
+ max_level=max_level,
+ image_size=image_size,
+ anchor_sizes=anchor_sizes)
+ outputs = model(
+ pillars,
+ indices,
+ image_shape,
+ anchor_boxes,
+ training)
+
+ # Test training and evaluation.
+ if training:
+ cls_outputs = outputs['cls_outputs']
+ box_outputs = outputs['box_outputs']
+ for level in range(min_level, max_level+1):
+ self.assertIn(str(level), cls_outputs)
+ self.assertIn(str(level), box_outputs)
+ self.assertAllEqual([
+ batch_size,
+ h // 2**level,
+ w // 2**level,
+ num_classes * num_anchors_per_location
+ ], cls_outputs[str(level)].shape)
+ self.assertAllEqual([
+ batch_size,
+ h // 2**level,
+ w // 2**level,
+ 4 * num_anchors_per_location
+ ], box_outputs[str(level)].shape)
+ att_outputs = outputs['attribute_outputs']
+ self.assertLen(att_outputs, 1)
+ self.assertIn('heading', att_outputs)
+ self.assertAllEqual([
+ batch_size,
+ h // 2**level,
+ w // 2**level,
+ 1 * num_anchors_per_location
+ ], att_outputs['heading'][str(level)].shape)
+ else:
+ self.assertIn('boxes', outputs)
+ self.assertIn('scores', outputs)
+ self.assertIn('classes', outputs)
+ self.assertIn('num_detections', outputs)
+ self.assertAllEqual([
+ batch_size,
+ ], outputs['num_detections'].shape)
+ self.assertAllEqual([batch_size, max_num_detections, 4],
+ outputs['boxes'].shape)
+ self.assertAllEqual([batch_size, max_num_detections],
+ outputs['scores'].shape)
+ self.assertAllEqual([batch_size, max_num_detections],
+ outputs['classes'].shape)
+ self.assertIn('attributes', outputs)
+ self.assertAllEqual(
+ [batch_size, max_num_detections, 1],
+ outputs['attributes']['heading'].shape)
+
+ # Test serialization.
+ config = model.get_config()
+ new_model = models.PointPillarsModel.from_config(config)
+ _ = new_model.to_json()
+ self.assertAllEqual(model.get_config(), new_model.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/registry_imports.py b/official/projects/pointpillars/registry_imports.py
new file mode 100644
index 00000000000..0084534a1e7
--- /dev/null
+++ b/official/projects/pointpillars/registry_imports.py
@@ -0,0 +1,24 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""All necessary imports for registration.
+
+Custom models, task, configs, etc need to be imported to registry so they can be
+picked up by the trainer. They can be included in this file so you do not need
+to handle each file separately.
+"""
+
+# pylint: disable=unused-import
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.tasks import pointpillars as task
diff --git a/official/projects/pointpillars/tasks/pointpillars.py b/official/projects/pointpillars/tasks/pointpillars.py
new file mode 100644
index 00000000000..5c7cac982ec
--- /dev/null
+++ b/official/projects/pointpillars/tasks/pointpillars.py
@@ -0,0 +1,387 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PointPillars task definition."""
+
+import functools
+from typing import Any, List, Mapping, Optional, Tuple
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import task_factory
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.dataloaders import decoders
+from official.projects.pointpillars.dataloaders import parsers
+from official.projects.pointpillars.modeling import factory
+from official.projects.pointpillars.utils import utils
+from official.vision.dataloaders import input_reader_factory
+from official.vision.losses import focal_loss
+from official.vision.losses import loss_utils
+
+
+def pick_dataset_fn(file_type: str) -> Any:
+ if file_type == 'tfrecord':
+ return tf.data.TFRecordDataset
+ if file_type == 'tfrecord_compressed':
+ return functools.partial(tf.data.TFRecordDataset, compression_type='GZIP')
+ raise ValueError('Unrecognized file_type: {}'.format(file_type))
+
+
+def get_batch_size_per_replica(global_batch_size: int) -> int:
+ """Get batch size per accelerator replica."""
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+ if global_batch_size < num_replicas:
+ logging.warning('Global batch size is smaller than num replicas. '
+ 'Set batch size per replica to 1.')
+ return 1
+ if global_batch_size % num_replicas != 0:
+ raise ValueError(
+ 'global_batch_size {} is not a multiple of num_replicas {}'
+ .format(global_batch_size, num_replicas))
+ batch_size = int(global_batch_size / num_replicas)
+ return batch_size
+
+
+@task_factory.register_task_cls(cfg.PointPillarsTask)
+class PointPillarsTask(base_task.Task):
+ """A single-replica view of training procedure."""
+
+ def __init__(self,
+ params: cfg.PointPillarsTask,
+ logging_dir: Optional[str] = None,
+ name: Optional[str] = None):
+ super().__init__(params, logging_dir, name)
+ self._model = None
+ self._attribute_heads = self.task_config.model.head.attribute_heads
+
+ def build_model(self) -> tf_keras.Model:
+ # Create only one model instance if this function is called multiple times.
+ if self._model is not None:
+ return self._model
+
+ pillars_config = self.task_config.model.pillars
+ input_specs = {
+ 'pillars':
+ tf_keras.layers.InputSpec(
+ shape=(None, pillars_config.num_pillars,
+ pillars_config.num_points_per_pillar,
+ pillars_config.num_features_per_point)),
+ 'indices':
+ tf_keras.layers.InputSpec(
+ shape=(None, pillars_config.num_pillars, 2), dtype='int32'),
+ }
+
+ train_batch_size = get_batch_size_per_replica(
+ self.task_config.train_data.global_batch_size)
+ eval_batch_size = get_batch_size_per_replica(
+ self.task_config.validation_data.global_batch_size)
+
+ l2_weight_decay = self.task_config.losses.l2_weight_decay
+ l2_regularizer = (tf_keras.regularizers.l2(
+ l2_weight_decay / 2.0) if l2_weight_decay else None)
+
+ self._model = factory.build_pointpillars(
+ input_specs=input_specs,
+ model_config=self.task_config.model,
+ train_batch_size=train_batch_size,
+ eval_batch_size=eval_batch_size,
+ l2_regularizer=l2_regularizer)
+ return self._model
+
+ def initialize(self, model: tf_keras.Model):
+ """Loading pretrained checkpoint."""
+ if not self.task_config.init_checkpoint:
+ return
+
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+
+ if self.task_config.init_checkpoint_modules == 'all':
+ ckpt = tf.train.Checkpoint(**model.checkpoint_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ else:
+ ckpt_items = {}
+ if 'backbone' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(backbone=model.backbone)
+ if 'decoder' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(decoder=model.decoder)
+
+ ckpt = tf.train.Checkpoint(**ckpt_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
+
+ def build_inputs(
+ self,
+ params: cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None
+ ) -> tf.data.Dataset:
+ """Build input dataset."""
+ model_config = self.task_config.model
+ if (model_config.classes != 'all' and
+ model_config.num_classes != 2):
+ raise ValueError('Model num_classes must be 2 when not for all classes.')
+
+ decoder = decoders.ExampleDecoder(model_config.image, model_config.pillars)
+
+ image_size = [model_config.image.height, model_config.image.width]
+ anchor_sizes = [(a.length, a.width) for a in model_config.anchors]
+ anchor_labeler_config = model_config.anchor_labeler
+ parser = parsers.Parser(
+ classes=model_config.classes,
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ image_size=image_size,
+ anchor_sizes=anchor_sizes,
+ match_threshold=anchor_labeler_config.match_threshold,
+ unmatched_threshold=anchor_labeler_config.unmatched_threshold,
+ max_num_detections=model_config.detection_generator
+ .max_num_detections,
+ dtype=params.dtype,
+ )
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training))
+ dataset = reader.read(input_context=input_context)
+
+ return dataset
+
+ def compute_attribute_losses(
+ self,
+ outputs: Mapping[str, Any],
+ labels: Mapping[str, Any],
+ box_sample_weight: tf.Tensor) -> Mapping[str, float]:
+ """Computes attribute loss."""
+ att_loss_fn = tf_keras.losses.Huber(
+ self.task_config.losses.huber_loss_delta,
+ reduction=tf_keras.losses.Reduction.SUM)
+
+ losses = {}
+ total_loss = 0.0
+ for head in self._attribute_heads:
+ if head.type != 'regression':
+ raise ValueError(f'Attribute type {head.type} not supported.')
+
+ y_true_att = loss_utils.multi_level_flatten(
+ labels['attribute_targets'][head.name], last_dim=head.size)
+ y_pred_att = loss_utils.multi_level_flatten(
+ outputs['attribute_outputs'][head.name], last_dim=head.size)
+ if head.name == 'heading':
+ # Direction aware loss, wrap the delta angle to [-pi, pi].
+ # Otherwise for a loss that is symmetric to direction (i.e., heading 0
+ # and pi are the same), we use a tf.sin transform.
+ delta = utils.wrap_angle_rad(y_pred_att - y_true_att)
+ loss = att_loss_fn(
+ y_true=tf.zeros_like(delta),
+ y_pred=delta,
+ sample_weight=box_sample_weight)
+ else:
+ loss = att_loss_fn(
+ y_true=y_true_att,
+ y_pred=y_pred_att,
+ sample_weight=box_sample_weight)
+ total_loss += loss
+ losses[head.name] = loss
+ losses['total'] = total_loss
+ return losses
+
+ def compute_losses(
+ self,
+ outputs: Mapping[str, Any],
+ labels: Mapping[str, Any],
+ aux_losses: Optional[Any] = None) -> Mapping[str, float]:
+ """Build losses."""
+ params = self.task_config
+
+ cls_loss_fn = focal_loss.FocalLoss(
+ alpha=params.losses.focal_loss_alpha,
+ gamma=params.losses.focal_loss_gamma,
+ reduction=tf_keras.losses.Reduction.SUM)
+ box_loss_fn = tf_keras.losses.Huber(
+ params.losses.huber_loss_delta,
+ reduction=tf_keras.losses.Reduction.SUM)
+
+ # Sums all positives in a batch for normalization and avoids zero
+ # num_positives_sum, which would lead to inf loss during training
+ cls_sample_weight = labels['cls_weights']
+ box_sample_weight = labels['box_weights']
+ num_positives = tf.reduce_sum(box_sample_weight) + 1.0
+ cls_sample_weight = cls_sample_weight / num_positives
+ box_sample_weight = box_sample_weight / num_positives
+
+ y_true_cls = loss_utils.multi_level_flatten(
+ labels['cls_targets'], last_dim=None)
+ y_true_cls = tf.one_hot(y_true_cls, params.model.num_classes)
+ y_pred_cls = loss_utils.multi_level_flatten(
+ outputs['cls_outputs'], last_dim=params.model.num_classes)
+ y_true_box = loss_utils.multi_level_flatten(
+ labels['box_targets'], last_dim=4)
+ y_pred_box = loss_utils.multi_level_flatten(
+ outputs['box_outputs'], last_dim=4)
+
+ cls_loss = cls_loss_fn(
+ y_true=y_true_cls, y_pred=y_pred_cls, sample_weight=cls_sample_weight)
+ box_loss = box_loss_fn(
+ y_true=y_true_box, y_pred=y_pred_box, sample_weight=box_sample_weight)
+ attribute_losses = self.compute_attribute_losses(outputs, labels,
+ box_sample_weight)
+ model_loss = (
+ cls_loss + box_loss * params.losses.box_loss_weight +
+ attribute_losses['total'] * params.losses.attribute_loss_weight)
+
+ total_loss = model_loss
+ if aux_losses:
+ reg_loss = tf.reduce_sum(aux_losses)
+ total_loss += reg_loss
+ total_loss = params.losses.loss_weight * total_loss
+
+ losses = {
+ 'class_loss': cls_loss,
+ 'box_loss': box_loss,
+ 'attribute_loss': attribute_losses['total'],
+ 'model_loss': model_loss,
+ 'total_loss': total_loss,
+ }
+ for head in self._attribute_heads:
+ losses[head.name + '_loss'] = attribute_losses[head.name]
+ return losses
+
+ def build_metrics(self, training: bool = True) -> List[tf.metrics.Metric]:
+ """Define metrics and how to calculate them."""
+ # train/validation loss metrics
+ loss_names = [
+ 'class_loss', 'box_loss', 'attribute_loss', 'model_loss', 'total_loss'
+ ]
+ for head in self._attribute_heads:
+ loss_names.append(head.name + '_loss')
+ metrics = []
+ for name in loss_names:
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
+
+ # Use a separate metric for WOD validation.
+ if not training:
+ if self.task_config.use_wod_metrics:
+ # To use Waymo open dataset metrics, please install one of the pip
+ # package `waymo-open-dataset-tf-*` from
+ # https://github.com/waymo-research/waymo-open-dataset/blob/master/docs/quick_start.md#use-pre-compiled-pippip3-packages-for-linux
+ # Note that the package is built with specific tensorflow version and
+ # will produce error if it does not match the tf version that is
+ # currently used.
+ try:
+ from official.projects.pointpillars.utils import wod_detection_evaluator # pylint: disable=g-import-not-at-top
+ except ModuleNotFoundError:
+ logging.error('waymo-open-dataset should be installed to enable Waymo'
+ ' evaluator.')
+ raise
+ self._wod_metric = wod_detection_evaluator.create_evaluator(
+ self.task_config.model)
+ return metrics
+
+ def train_step(
+ self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[tf.metrics.Metric]] = None) -> Mapping[str, Any]:
+ """Does forward and backward."""
+ features, labels = inputs
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+ with tf.GradientTape() as tape:
+ outputs = model(pillars=features['pillars'],
+ indices=features['indices'],
+ training=True)
+ losses = self.compute_losses(
+ outputs=outputs, labels=labels, aux_losses=model.losses)
+
+ # Computes per-replica loss.
+ scaled_loss = losses['total_loss'] / num_replicas
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scales back gradient when LossScaleOptimizer is used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ # For updating trainer.train_loss
+ logs = {self.loss: losses['total_loss']}
+ # For updating trainer.train_metrics
+ if metrics:
+ for m in metrics:
+ m.update_state(losses[m.name])
+ return logs
+
+ def validation_step(
+ self,
+ inputs: Tuple[Any, Any],
+ model: tf_keras.Model,
+ metrics: Optional[List[tf.metrics.Metric]] = None) -> Mapping[str, Any]:
+ """Validatation step."""
+ features, labels = inputs
+ outputs = model(pillars=features['pillars'],
+ indices=features['indices'],
+ image_shape=labels['image_shape'],
+ anchor_boxes=labels['anchor_boxes'],
+ training=False)
+ losses = self.compute_losses(
+ outputs=outputs, labels=labels, aux_losses=model.losses)
+
+ # For updating trainer.validation_loss
+ logs = {self.loss: losses['total_loss']}
+ # For updating trainer.validation_metrics
+ if metrics:
+ for m in metrics:
+ m.update_state(losses[m.name])
+ if self.task_config.use_wod_metrics:
+ logs.update(
+ {self._wod_metric.name: (labels['groundtruths'], outputs)})
+ return logs
+
+ def aggregate_logs(self,
+ state: Any = None,
+ step_outputs: Any = None) -> Any:
+ """Called after each validation_step to update metrics."""
+ logging.log_every_n(logging.INFO,
+ 'Aggregating metrics after one evaluation step.', 1000)
+ if self.task_config.use_wod_metrics:
+ if state is None:
+ self._wod_metric.reset_states()
+ self._wod_metric.update_state(step_outputs[self._wod_metric.name][0],
+ step_outputs[self._wod_metric.name][1])
+ if state is None:
+ state = True
+ return state
+
+ def reduce_aggregated_logs(self,
+ aggregated_logs: Any,
+ global_step: Optional[tf.Tensor] = None) -> Any:
+ """Called after eval_end to calculate metrics."""
+ logging.info('Reducing aggregated metrics after one evaluation cycle.')
+ logs = {}
+ if self.task_config.use_wod_metrics:
+ logs.update(self._wod_metric.result())
+ return logs
diff --git a/official/projects/pointpillars/tasks/pointpillars_test.py b/official/projects/pointpillars/tasks/pointpillars_test.py
new file mode 100644
index 00000000000..ee667a7c2d4
--- /dev/null
+++ b/official/projects/pointpillars/tasks/pointpillars_test.py
@@ -0,0 +1,117 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for pointpillars."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.core import exp_factory
+from official.modeling import optimization
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.tasks import pointpillars
+
+
+def _mock_inputs(model_config):
+ batch_size = 1
+ image_config = model_config.image
+ pillars_config = model_config.pillars
+ pillars = tf.ones([
+ batch_size, pillars_config.num_pillars,
+ pillars_config.num_points_per_pillar,
+ pillars_config.num_features_per_point
+ ], dtype=tf.float32)
+ indices = tf.ones([
+ batch_size, pillars_config.num_pillars, 2
+ ], dtype=tf.int32)
+ features = {
+ 'pillars': pillars,
+ 'indices': indices,
+ }
+
+ image_height = image_config.height
+ image_width = image_config.width
+ num_anchors_per_location = len(model_config.anchors)
+ cls_targets = {}
+ box_targets = {}
+ attribute_targets = {}
+ for attr in model_config.head.attribute_heads:
+ attribute_targets[attr.name] = {}
+ total_num_anchors = 0
+ for level in range(model_config.min_level, model_config.max_level + 1):
+ stride = 2**level
+ h_i = int(image_height / stride)
+ w_i = int(image_width / stride)
+ cls_targets[str(level)] = tf.ones(
+ [batch_size, h_i, w_i, num_anchors_per_location], dtype=tf.int32)
+ box_targets[str(level)] = tf.ones(
+ [batch_size, h_i, w_i, num_anchors_per_location * 4], dtype=tf.float32)
+ for attr in model_config.head.attribute_heads:
+ attribute_targets[attr.name][str(level)] = tf.ones(
+ [batch_size, h_i, w_i, num_anchors_per_location], dtype=tf.float32)
+ total_num_anchors += h_i * w_i * num_anchors_per_location
+ cls_weights = tf.ones([batch_size, total_num_anchors], dtype=tf.float32)
+ box_weights = tf.ones([batch_size, total_num_anchors], dtype=tf.float32)
+ image_shape = tf.ones([batch_size, 2], dtype=tf.int32)
+ labels = {
+ 'cls_targets': cls_targets,
+ 'box_targets': box_targets,
+ 'attribute_targets': attribute_targets,
+ 'cls_weights': cls_weights,
+ 'box_weights': box_weights,
+ 'anchor_boxes': None,
+ 'image_shape': image_shape,
+ }
+ return features, labels
+
+
+class PointPillarsTaskTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (True),
+ (False),
+ )
+ def test_train_and_eval(self, is_training):
+ exp_config = exp_factory.get_exp_config('pointpillars_baseline')
+ task_config = exp_config.task
+ # modify config to suit local testing
+ task_config.model.image.height = 32
+ task_config.model.image.width = 32
+ task_config.model.pillars.num_pillars = 2
+ task_config.model.pillars.num_points_per_pillar = 3
+ task_config.model.pillars.num_features_per_point = 4
+ task_config.model.anchors = [cfg.Anchor(length=2.1, width=1.2)]
+
+ task_config.train_data.global_batch_size = 1
+ task_config.train_data.shuffle_buffer_size = 2
+ task_config.validation_data.global_batch_size = 1
+ task_config.validation_data.shuffle_buffer_size = 2
+ task_config.use_wod_metrics = False
+
+ task = pointpillars.PointPillarsTask(task_config)
+ inputs = _mock_inputs(task_config.model)
+ model = task.build_model()
+ opt_factory = optimization.OptimizerFactory(
+ exp_config.trainer.optimizer_config)
+ optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
+ metrics = task.build_metrics(training=is_training)
+
+ if is_training:
+ logs = task.train_step(inputs, model, optimizer, metrics=metrics)
+ else:
+ logs = task.validation_step(inputs, model, metrics=metrics)
+ self.assertIn('loss', logs)
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/tools/export_model.py b/official/projects/pointpillars/tools/export_model.py
new file mode 100644
index 00000000000..cc6ff197e11
--- /dev/null
+++ b/official/projects/pointpillars/tools/export_model.py
@@ -0,0 +1,66 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""A script to export PointPillars model."""
+
+from absl import app
+from absl import flags
+from absl import logging
+
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.pointpillars import registry_imports # pylint: disable=unused-import
+from official.projects.pointpillars.utils import model_exporter
+
+_EXPERIMENT = flags.DEFINE_string(
+ 'experiment', None, 'experiment type, e.g. retinanet_resnetfpn_coco')
+_EXPORT_DIR = flags.DEFINE_string('export_dir', None, 'The export directory.')
+_CHECKPOINT_PATH = flags.DEFINE_string('checkpoint_path', None,
+ 'Checkpoint path.')
+_BATCH_SIZE = flags.DEFINE_integer('batch_size', None, 'Batch size.')
+_CONFIG_FILE = flags.DEFINE_string(
+ 'config_file',
+ default=None,
+ help='YAML/JSON files which specifies overrides.')
+_TEST_INFERENCE = flags.DEFINE_boolean(
+ 'test_inference',
+ default=False,
+ help='True if want to load saved model and run inference.')
+
+
+def main(_):
+ params = exp_factory.get_exp_config(_EXPERIMENT.value)
+ if _CONFIG_FILE.value:
+ params = hyperparams.override_params_dict(
+ params, _CONFIG_FILE.value, is_strict=True)
+ params.validate()
+ params.lock()
+
+ model_exporter.export_inference_graph(
+ batch_size=_BATCH_SIZE.value,
+ params=params,
+ checkpoint_path=_CHECKPOINT_PATH.value,
+ export_dir=_EXPORT_DIR.value)
+ logging.info('Successfully exported model to %s', _EXPORT_DIR.value)
+
+ if _TEST_INFERENCE.value:
+ predict_fn = model_exporter.load_model_predict_fn(_EXPORT_DIR.value)
+ pillars, indices = model_exporter.random_input_tensors(
+ batch_size=_BATCH_SIZE.value, params=params,
+ )
+ _ = predict_fn(pillars=pillars, indices=indices)
+ logging.info('Successfully test model inference')
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/pointpillars/tools/process_wod.py b/official/projects/pointpillars/tools/process_wod.py
new file mode 100644
index 00000000000..3808743f767
--- /dev/null
+++ b/official/projects/pointpillars/tools/process_wod.py
@@ -0,0 +1,119 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""A script to run waymo open dataset preprocessing."""
+
+import os
+
+from absl import app
+from absl import flags
+from absl import logging
+import apache_beam as beam
+from apache_beam.io import tfrecordio
+import tensorflow as tf, tf_keras
+
+from official.modeling import hyperparams
+from official.projects.pointpillars.configs import pointpillars
+from official.projects.pointpillars.utils.wod_processor import WodProcessor
+from waymo_open_dataset import dataset_pb2
+
+
+_SRC_DIR = flags.DEFINE_string(
+ 'src_dir', None,
+ 'The direcotry to read official wod tfrecords,')
+_DST_DIR = flags.DEFINE_string(
+ 'dst_dir', None,
+ 'The direcotry to write processed tfrecords.')
+_CONFIG_FILE = flags.DEFINE_string(
+ 'config_file', None,
+ 'YAML file to specify configurations.')
+_PIPELINE_OPTIONS = flags.DEFINE_string(
+ 'pipeline_options', None,
+ 'Command line flags to use in constructing the Beam pipeline options. '
+ 'See https://beam.apache.org/documentation/#runners for available runners.')
+
+# The --src_dir must contain these two sub-folders.
+_SRC_FOLDERS = ['training', 'validation']
+
+
+def read_dataset(pipeline: beam.Pipeline,
+ src_file_pattern: str) -> beam.PCollection:
+ reader = tfrecordio.ReadFromTFRecord(
+ src_file_pattern,
+ coder=beam.coders.ProtoCoder(dataset_pb2.Frame))
+ raw_frames = pipeline | f'Read frames: {src_file_pattern}' >> reader
+ return raw_frames
+
+
+def count_examples(examples: beam.PCollection, dst_path: str):
+ writer = beam.io.WriteToText(
+ dst_path,
+ file_name_suffix='.stats.txt',
+ num_shards=1)
+ _ = (examples
+ | 'Count examples' >> beam.combiners.Count.Globally()
+ | 'Write statistics' >> writer)
+
+
+def write_dataset(examples: beam.PCollection, dst_path: str):
+ writer = tfrecordio.WriteToTFRecord(
+ dst_path,
+ coder=beam.coders.ProtoCoder(tf.train.Example),
+ file_name_suffix='.tfrecord',
+ compression_type='gzip')
+ _ = examples | f'Write examples: {dst_path}' >> writer
+
+
+def process_wod(pipeline: beam.Pipeline,
+ src_file_pattern: str,
+ dst_path: str,
+ wod_processor: WodProcessor):
+ """Creates the process WOD dataset pipeline."""
+ raw_frames = read_dataset(pipeline, src_file_pattern)
+ examples = (
+ raw_frames
+ | 'Reshuffle post read' >> beam.Reshuffle()
+ | 'Process one frame' >> beam.Map(
+ wod_processor.process_and_convert_to_tf_example)
+ | 'Reshuffle post decode' >> beam.Reshuffle())
+ count_examples(examples, dst_path)
+ write_dataset(examples, dst_path)
+
+
+def main(_):
+ pipeline_options = beam.options.pipeline_options.PipelineOptions(
+ _PIPELINE_OPTIONS.value.split(',')) # pyrefly: ignore[missing-attribute]
+
+ if _CONFIG_FILE.value:
+ cfg = hyperparams.read_yaml_to_params_dict(_CONFIG_FILE.value)
+ image_config = cfg.task.model.image
+ pillars_config = cfg.task.model.pillars
+ else:
+ cfg = pointpillars
+ image_config = cfg.ImageConfig()
+ pillars_config = cfg.PillarsConfig()
+
+ wod_processor = WodProcessor(image_config, pillars_config)
+ for folder in _SRC_FOLDERS:
+ src_file_pattern = os.path.join(_SRC_DIR.value, folder, '*.tfrecord') # pyrefly: ignore[no-matching-overload]
+ dst_path = os.path.join(_DST_DIR.value, folder) # pyrefly: ignore[no-matching-overload]
+ logging.info('Processing %s, writing to %s', src_file_pattern, dst_path)
+
+ pipeline = beam.Pipeline(options=pipeline_options)
+ process_wod(pipeline, src_file_pattern, dst_path, wod_processor)
+ pipeline.run().wait_until_finish()
+
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/pointpillars/train.py b/official/projects/pointpillars/train.py
new file mode 100644
index 00000000000..34efe535650
--- /dev/null
+++ b/official/projects/pointpillars/train.py
@@ -0,0 +1,105 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PointPillars trainer."""
+
+import os
+
+from absl import app
+from absl import flags
+from absl import logging
+import gin
+
+import tensorflow as tf, tf_keras
+
+from official.common import distribute_utils
+from official.common import flags as tfm_flags
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+from official.projects.pointpillars import registry_imports # pylint: disable=unused-import
+from official.projects.pointpillars.utils import model_exporter
+
+FLAGS = flags.FLAGS
+
+
+def _check_if_resumed_job(model_dir: str, manual_checkpoint_path: str):
+ """Check if the job is a resumed job."""
+ logging.info('Check if the job is resumed from %s', model_dir)
+ if not tf.io.gfile.exists(model_dir):
+ logging.info('%s not found, this is a new job.', model_dir)
+ return
+ try:
+ tf.train.load_checkpoint(model_dir)
+ except ValueError:
+ logging.info('No checkpoints found in %s, this is a new job.', model_dir)
+ return
+ else:
+ logging.info('The job is resuming from %s', model_dir)
+ if manual_checkpoint_path:
+ logging.warning('Found manually indicated checkpoint path %s for a '
+ 'resuming job, the manual checkpoint path will be '
+ 'ignored because the model must restore from '
+ 'checkpoints in %s.', manual_checkpoint_path, model_dir)
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+
+ # A training job could be terminated and resumed at any time by machine
+ # scheduler. A resuming job will automatically restore states from the
+ # model_dir, like loading checkpoints. It will skip checkpointed training
+ # steps and start from there for subsequent training. This function simply
+ # checks if the job is a resumed job or not and logs info for that.
+ _check_if_resumed_job(model_dir, params.task.init_checkpoint)
+
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ # Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. 'loss_scale' takes effect only when
+ # dtype is float16.
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu)
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir)
+
+ model_exporter.export_inference_graph(
+ batch_size=1,
+ params=params,
+ checkpoint_path=model_dir,
+ export_dir=os.path.join(model_dir, 'saved_model'))
+
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(main)
diff --git a/official/projects/pointpillars/utils/model_exporter.py b/official/projects/pointpillars/utils/model_exporter.py
new file mode 100644
index 00000000000..8f0b662b77c
--- /dev/null
+++ b/official/projects/pointpillars/utils/model_exporter.py
@@ -0,0 +1,255 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PointPillars model export utility function for serving/inference."""
+
+import os
+from typing import Any, Dict, Mapping, Optional, Tuple
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import export_base
+from official.core import train_utils
+from official.projects.pointpillars.modeling import factory
+from official.projects.pointpillars.utils import utils
+
+
+def export_inference_graph(
+ batch_size: int,
+ params: cfg.ExperimentConfig,
+ checkpoint_path: str,
+ export_dir: str,
+ export_module: Optional[export_base.ExportModule] = None,
+):
+ """Exports inference graph for PointPillars model.
+
+ Saved model is stored at export_dir/saved_model, checkpoint is saved
+ at export_dir/checkpoint, and params is saved at export_dir/params.yaml.
+
+ Args:
+ batch_size: An int number specifying batch size for inference.
+ Saved PointPillars model doesn't support dynamic batch size.
+ Only three batch sizes are acceptable:
+ train batch size per replica, evaluation batch size per replica, and 1.
+ params: An instance of cfg.ExperimentConfig.
+ checkpoint_path: Trained checkpoint path or directory.
+ export_dir: Export directory path.
+ export_module: Optional export module to be used instead of using params
+ to create one.
+ """
+ logging.info('Exporting model.')
+ if not export_module:
+ export_module = PointPillarsModule(
+ params=params,
+ batch_size=batch_size)
+ # Disable custom_gradients to make trt-converter be able to work.
+ # Consider to use tf_keras.models.save_model/load_model APIs to fix
+ # the custom gradients saving problem.
+ # https://github.com/tensorflow/tensorflow/issues/40166
+ save_options = tf.saved_model.SaveOptions(experimental_custom_gradients=False)
+ export_base.export(
+ export_module,
+ function_keys=['tensors'],
+ export_savedmodel_dir=export_dir,
+ checkpoint_path=checkpoint_path,
+ timestamped=False,
+ save_options=save_options)
+
+ logging.info('Saving checkpoint.')
+ ckpt = tf.train.Checkpoint(model=export_module.model)
+ ckpt.save(os.path.join(export_dir, 'checkpoint', 'ckpt'))
+
+ logging.info('Saving experiment params.')
+ train_utils.serialize_config(params, export_dir)
+
+
+def load_model_predict_fn(export_dir: str) -> Any:
+ """Load PointPillars model from saved directory.
+
+ Args:
+ export_dir: Export directory path.
+ Returns:
+ predict_fn: A function can be run for model inference.
+ """
+ logging.info('Loading model from %s.', export_dir)
+ model = tf.saved_model.load(export_dir)
+ predict_fn = model.signatures[
+ tf.saved_model.DEFAULT_SERVING_SIGNATURE_DEF_KEY]
+ return predict_fn
+
+
+def random_input_tensors(
+ batch_size: int,
+ params: cfg.ExperimentConfig) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Create random input tensors for PointPillars model.
+
+ Args:
+ batch_size: An int number specifying batch size to inference.
+ params: An instance of cfg.ExperimentConfig.
+ Returns:
+ pillars: A tensor for input.
+ indices: A tensor for input.
+ """
+ model_config = params.task.model
+ pillars_config = model_config.pillars
+ pillars = tf.random.uniform(
+ shape=[batch_size,
+ pillars_config.num_pillars,
+ pillars_config.num_points_per_pillar,
+ pillars_config.num_features_per_point],
+ minval=0.0,
+ maxval=1.0,
+ dtype=tf.float32,
+ name='pillars')
+ indices = tf.random.uniform(
+ shape=[batch_size, pillars_config.num_pillars, 2],
+ minval=0,
+ maxval=model_config.image.height,
+ dtype=tf.int32,
+ name='indices')
+ return pillars, indices
+
+
+class PointPillarsModule(export_base.ExportModule):
+ """PointPillars model export module."""
+
+ def __init__(self, params: cfg.ExperimentConfig, batch_size: int):
+ """Initialize the module.
+
+ Args:
+ params: Experiment params.
+ batch_size: The batch size of the model input.
+ """
+ self._params = params
+ self._batch_size = batch_size
+ self._pillars_spec, self._indices_spec = self._build_input_specs()
+ model = self._build_model()
+ super().__init__(params=params, model=model)
+
+ def _build_input_specs(
+ self) -> Tuple[tf_keras.layers.InputSpec, tf_keras.layers.InputSpec]:
+ pillars_config = self._params.task.model.pillars
+ pillars_spec = tf_keras.layers.InputSpec(
+ shape=(self._batch_size,
+ pillars_config.num_pillars,
+ pillars_config.num_points_per_pillar,
+ pillars_config.num_features_per_point),
+ dtype='float32')
+ indices_spec = tf_keras.layers.InputSpec(
+ shape=(self._batch_size,
+ pillars_config.num_pillars,
+ 2),
+ dtype='int32')
+ return pillars_spec, indices_spec
+
+ def _build_model(self) -> tf_keras.Model:
+ logging.info('Building PointPillars model.')
+ input_specs = {
+ 'pillars': self._pillars_spec, 'indices': self._indices_spec
+ }
+ model = factory.build_pointpillars(
+ input_specs=input_specs,
+ model_config=self._params.task.model,
+ # Train and eval batch size will be ignored for inference.
+ train_batch_size=1,
+ eval_batch_size=1)
+ return model
+
+ def serve(self, pillars: tf.Tensor, indices: tf.Tensor) -> Mapping[str, Any]:
+ """Run model inference.
+
+ Args:
+ pillars: A float32 tensor.
+ indices: An int32 tensor.
+ Returns:
+ outputs: A dict of detected results.
+ """
+ # Build image_shape and anchor_boxes on CPU.
+ with tf.device('cpu'):
+ model_config = self._params.task.model
+ image_size = [model_config.image.height,
+ model_config.image.width]
+
+ image_shape = tf.tile(tf.expand_dims(
+ image_size, axis=0), [self._batch_size, 1])
+
+ anchor_sizes = [(a.length, a.width) for a in model_config.anchors]
+ anchor_boxes = utils.generate_anchors(
+ min_level=model_config.min_level,
+ max_level=model_config.max_level,
+ image_size=image_size,
+ anchor_sizes=anchor_sizes)
+ for l in anchor_boxes:
+ anchor_boxes[l] = tf.tile(
+ tf.expand_dims(anchor_boxes[l], axis=0),
+ [self._batch_size, 1, 1, 1])
+
+ # Run model.
+ detections = self.model.call(
+ pillars=pillars,
+ indices=indices,
+ image_shape=image_shape,
+ anchor_boxes=anchor_boxes,
+ training=None
+ )
+ outputs = {
+ 'detection_boxes': detections['boxes'],
+ 'detection_scores': detections['scores'],
+ 'detection_classes': detections['classes'],
+ 'num_detections': detections['num_detections']
+ }
+ # NOTE: Need to flatten attributes, because outputs for functions used as
+ # signatures must be a single Tensor, a sequence of Tensors, or a dictionary
+ # from string to Tensor.
+ outputs.update(detections['attributes'])
+ return outputs
+
+ @tf.function
+ def inference_from_tensors(
+ self, pillars: tf.Tensor, indices: tf.Tensor) -> Mapping[str, Any]:
+ return self.serve(pillars, indices)
+
+ def get_inference_signatures(
+ self, function_keys: Dict[str, str]) -> Mapping[str, Any]:
+ """Gets defined function signatures.
+
+ Args:
+ function_keys: A dictionary with keys as the function to create signature
+ for and values as the signature keys when returns.
+
+ Returns:
+ A dictionary with key as signature key and value as concrete functions
+ that can be used for tf.saved_model.save.
+ """
+ signatures = {}
+ for input_type, name in function_keys.items():
+ if input_type == 'tensors':
+ pillars = tf.TensorSpec(
+ shape=self._pillars_spec.shape,
+ dtype=self._pillars_spec.dtype,
+ name='pillars'
+ )
+ indices = tf.TensorSpec(
+ shape=self._indices_spec.shape,
+ dtype=self._indices_spec.dtype,
+ name='indices'
+ )
+ signatures[
+ name] = self.inference_from_tensors.get_concrete_function(
+ pillars, indices)
+ else:
+ raise ValueError('Unrecognized input_type: {}'.format(input_type))
+ return signatures
diff --git a/official/projects/pointpillars/utils/utils.py b/official/projects/pointpillars/utils/utils.py
new file mode 100644
index 00000000000..970c0e777d0
--- /dev/null
+++ b/official/projects/pointpillars/utils/utils.py
@@ -0,0 +1,272 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains utility functions for pointpillars."""
+
+import collections
+from typing import Any, List, Mapping, Tuple
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+
+CLASSES = {'vehicle': 1, 'pedestrian': 2, 'cyclist': 3}
+
+
+def assert_shape(x: np.ndarray, shape: List[int]):
+ if tuple(x.shape) != tuple(shape):
+ raise ValueError('Shape of array should be {}, but {} found'.format(
+ shape, x.shape))
+
+
+def assert_channels_last():
+ if tf_keras.backend.image_data_format() != 'channels_last':
+ raise ValueError('Only "channels_last" mode is supported')
+
+
+def pad_or_trim_to_shape(x: np.ndarray, shape: List[int]) -> np.ndarray:
+ """Pad and trim x to the specified shape, x should have same rank as shape.
+
+ Args:
+ x: An np array.
+ shape: A list of int indicating a array shape.
+
+ Returns:
+ y: An np array with padded/trimmed shape.
+ """
+ shape = np.array(shape) # pyrefly: ignore[bad-assignment]
+
+ # Try to pad from end
+ pad_end = shape - np.minimum(x.shape, shape)
+ pad_begin = np.zeros_like(pad_end)
+ padder = np.stack([pad_begin, pad_end], axis=1)
+ x = np.pad(x, padder)
+
+ # Try to trim from end.
+ slice_end = shape
+ slice_begin = np.zeros_like(slice_end)
+ slicer = tuple(map(slice, slice_begin, slice_end))
+ y = x[slicer].reshape(shape)
+ return y
+
+
+def clip_boxes(boxes: np.ndarray, image_height: int,
+ image_width: int) -> np.ndarray:
+ """Clip boxes to image boundaries.
+
+ Args:
+ boxes: An np array of boxes, [y0, x0, y1, y1].
+ image_height: An int of image height.
+ image_width: An int of image width.
+ Returns:
+ clipped_boxes: An np array of boxes, [y0, x0, y1, y1].
+ """
+ max_length = [image_height, image_width, image_height, image_width]
+ clipped_boxes = np.maximum(np.minimum(boxes, max_length), 0.0)
+ return clipped_boxes
+
+
+def get_vehicle_xy(image_height: int, image_width: int,
+ x_range: Tuple[float, float],
+ y_range: Tuple[float, float]) -> Tuple[int, int]:
+ """Get vehicle x/y in image coordinate.
+
+ Args:
+ image_height: A float of image height.
+ image_width: A float of image width.
+ x_range: A float tuple of (-x, +x).
+ y_range: A float tuple of (-y, +x).
+ Returns:
+ vehicle_xy: An int tuple of (col, row).
+ """
+ vehicle_col = (image_width * (-x_range[0] / (-x_range[0] + x_range[1])))
+ vehicle_row = (image_height * (-y_range[0] / (-y_range[0] + y_range[1])))
+ vehicle_xy = (int(vehicle_col), int(vehicle_row))
+ return vehicle_xy
+
+
+def frame_to_image_coord(frame_xy: np.ndarray, vehicle_xy: Tuple[int, int],
+ one_over_resolution: float) -> np.ndarray:
+ """Convert float frame (x, y) to int image (x, y).
+
+ Args:
+ frame_xy: An np array of frame xy coordinates.
+ vehicle_xy: An int tuple of (vehicle_x, vehicle_y) in image.
+ one_over_resolution: A float of one over image resolution.
+ Returns:
+ image_xy: An np array of image xy cooridnates.
+ """
+ image_xy = np.floor(frame_xy * one_over_resolution).astype(np.int32)
+ image_xy[..., 0] += vehicle_xy[0]
+ image_xy[..., 1] = vehicle_xy[1] - 1 - image_xy[..., 1]
+ return image_xy
+
+
+def image_to_frame_coord(image_xy: np.ndarray, vehicle_xy: Tuple[int, int],
+ resolution: float) -> np.ndarray:
+ """Convert int image (x, y) to float frame (x, y).
+
+ Args:
+ image_xy: An np array of image xy cooridnates.
+ vehicle_xy: An int tuple of (vehicle_x, vehicle_y) in image.
+ resolution: A float of image resolution.
+ Returns:
+ frame_xy: An np array of frame xy coordinates.
+ """
+ frame_xy = image_xy.astype(np.float32)
+ frame_xy[..., 0] = (frame_xy[..., 0] - vehicle_xy[0]) * resolution
+ frame_xy[..., 1] = (vehicle_xy[1] - 1 - frame_xy[..., 1]) * resolution
+ return frame_xy
+
+
+def frame_to_image_boxes(frame_boxes: Any, vehicle_xy: Tuple[int, int],
+ one_over_resolution: float) -> Any:
+ """Convert boxes from frame coordinate to image coordinate.
+
+ Args:
+ frame_boxes: A [N, 4] array or tensor, [center_x, center_y, length, width]
+ in frame coordinate.
+ vehicle_xy: An int tuple of (vehicle_x, vehicle_y) in image.
+ one_over_resolution: A float number, 1.0 / resolution.
+
+ Returns:
+ image_boxes: A [N, 4] array or tensor, [ymin, xmin, ymax, xmax] in image
+ coordinate.
+ """
+ center_x = frame_boxes[..., 0]
+ center_y = frame_boxes[..., 1]
+ box_length = frame_boxes[..., 2]
+ box_width = frame_boxes[..., 3]
+
+ image_box_length = box_length * one_over_resolution
+ image_box_width = box_width * one_over_resolution
+ image_box_center_x = (center_x * one_over_resolution + vehicle_xy[0])
+ image_box_center_y = (vehicle_xy[1] - 1 - center_y * one_over_resolution)
+
+ ymin = image_box_center_y - image_box_width * 0.5
+ xmin = image_box_center_x - image_box_length * 0.5
+ ymax = image_box_center_y + image_box_width * 0.5
+ xmax = image_box_center_x + image_box_length * 0.5
+
+ image_boxes = np.stack([ymin, xmin, ymax, xmax], axis=-1)
+ return image_boxes
+
+
+def image_to_frame_boxes(image_boxes: Any, vehicle_xy: Tuple[float],
+ resolution: float) -> Any:
+ """Convert boxes from image coordinate to frame coordinate.
+
+ Args:
+ image_boxes: A [N, 4] array or tensor, [ymin, xmin, ymax, xmax] in image
+ coordinate.
+ vehicle_xy: A float tuple of (vehicle_x, vehicle_y) in image.
+ resolution: A float number representing pillar grid resolution.
+
+ Returns:
+ frame_boxes: A [N, 4] array or tensor, [center_x, center_y, length, width]
+ in frame coordinate.
+ """
+ ymin = image_boxes[..., 0]
+ xmin = image_boxes[..., 1]
+ ymax = image_boxes[..., 2]
+ xmax = image_boxes[..., 3]
+
+ image_box_length = xmax - xmin
+ image_box_width = ymax - ymin
+ image_box_center_x = xmin + image_box_length * 0.5
+ image_box_center_y = ymin + image_box_width * 0.5
+
+ center_x = (image_box_center_x - vehicle_xy[0]) * resolution
+ center_y = (vehicle_xy[1] - 1 - image_box_center_y) * resolution # pyrefly: ignore[bad-index]
+ box_length = image_box_length * resolution
+ box_width = image_box_width * resolution
+
+ frame_boxes = np.stack([center_x, center_y, box_length, box_width], axis=-1)
+ return frame_boxes
+
+
+def clip_heading(heading: Any) -> Any:
+ """Clip heading to the range [-pi, pi]."""
+ heading = tf.nest.map_structure(lambda x: np.pi * tf.tanh(x), heading)
+ return heading
+
+
+def wrap_angle_rad(angles_rad: Any,
+ min_val: float = -np.pi,
+ max_val: float = np.pi) -> Any:
+ """Wrap the value of `angles_rad` to the range [min_val, max_val]."""
+ max_min_diff = max_val - min_val
+ return min_val + tf.math.floormod(angles_rad + max_val, max_min_diff)
+
+
+def generate_anchors(min_level: int, max_level: int, image_size: Tuple[int],
+ anchor_sizes: List[Tuple[float]]) -> Mapping[str, Any]:
+ """Generate anchor boxes without scale to level stride.
+
+ Args:
+ min_level: integer number of minimum level of the output.
+ max_level: integer number of maximum level of the output.
+ image_size: a tuple (image_height, image_width).
+ anchor_sizes: a list of tuples, each tuple is (anchor_length, anchor_width).
+
+ Returns:
+ boxes_all: a {level: boxes_i} dict, each boxes_i is a [h_i, w_i, 4] tensor
+ for boxes at level i, each box is (ymin, xmin, ymax, xmax).
+
+ Notations:
+ k: length of anchor_sizes, the number of indicated anchors.
+ w: the image width at a specific level.
+ h: the image height at a specifc level.
+ """
+ # Prepare k anchors' lengths and widths
+ k = len(anchor_sizes)
+ # (k,)
+ anchor_lengths = []
+ anchor_widths = []
+ for anchor_size in anchor_sizes:
+ anchor_lengths.append(anchor_size[0])
+ anchor_widths.append(anchor_size[1]) # pyrefly: ignore[bad-index]
+ anchor_lengths = tf.convert_to_tensor(anchor_lengths, dtype=tf.float32)
+ anchor_widths = tf.convert_to_tensor(anchor_widths, dtype=tf.float32)
+ # (1, 1, k)
+ half_anchor_lengths = tf.reshape(0.5 * anchor_lengths, [1, 1, k])
+ half_anchor_widths = tf.reshape(0.5 * anchor_widths, [1, 1, k])
+
+ boxes_all = collections.OrderedDict()
+ for level in range(min_level, max_level + 1):
+ # Generate anchor boxes for this level with stride.
+ boxes_i = []
+ stride = 2 ** level
+ # (w,)
+ x = tf.range(stride / 2, image_size[1], stride, dtype=tf.float32) # pyrefly: ignore[bad-index]
+ # (h,)
+ y = tf.range(stride / 2, image_size[0], stride, dtype=tf.float32)
+ # (h, w)
+ xv, yv = tf.meshgrid(x, y)
+ # (h, w, 1)
+ xv = tf.expand_dims(xv, axis=-1)
+ yv = tf.expand_dims(yv, axis=-1)
+ # (h, w, k, 1)
+ y_min = tf.expand_dims(yv - half_anchor_widths, axis=-1)
+ y_max = tf.expand_dims(yv + half_anchor_widths, axis=-1)
+ x_min = tf.expand_dims(xv - half_anchor_lengths, axis=-1)
+ x_max = tf.expand_dims(xv + half_anchor_lengths, axis=-1)
+ # (h, w, k, 4)
+ boxes_i = tf.concat([y_min, x_min, y_max, x_max], axis=-1)
+ # [h, w, k * 4]
+ shape = boxes_i.shape.as_list()
+ boxes_i = tf.reshape(boxes_i, [shape[0], shape[1], shape[2] * shape[3]])
+
+ boxes_all[str(level)] = boxes_i
+ return boxes_all
diff --git a/official/projects/pointpillars/utils/utils_test.py b/official/projects/pointpillars/utils/utils_test.py
new file mode 100644
index 00000000000..0deb61ee4c7
--- /dev/null
+++ b/official/projects/pointpillars/utils/utils_test.py
@@ -0,0 +1,96 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for pointpillars utils."""
+
+from absl.testing import parameterized
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.utils import utils
+
+
+class UtilsTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ([2, 1], [2, 1]),
+ ([1, 1], [4, 3]),
+ ([2, 2, 4], [2, 1, 5]),
+ )
+ def test_pad_or_trim_to_shape(self, original_shape, expected_shape):
+ x = np.ones(shape=original_shape)
+ x = utils.pad_or_trim_to_shape(x, expected_shape)
+ self.assertAllEqual(x.shape, expected_shape)
+
+ @parameterized.parameters(
+ ([[1.1, 1.1, 2.2, 2.2]], 10.0, 5.0),
+ ([[1.1, 10.1, 2.2, 10.2]], 10.0, 10.0),
+ ([[-1.1, 10.1, -2.2, 10.2]], 5.0, 2.0),
+ )
+ def test_clip_boxes(self, boxes, height, width):
+ boxes = np.array(boxes)
+ boxes = utils.clip_boxes(boxes, height, width)
+ self.assertGreaterEqual(boxes[:, 0], 0.0)
+ self.assertGreaterEqual(boxes[:, 1], 0.0)
+ self.assertLessEqual(boxes[:, 2], height)
+ self.assertLessEqual(boxes[:, 3], width)
+
+ def test_get_vehicle_xy(self):
+ vehicle_xy = utils.get_vehicle_xy(10, 10, (-50, 50), (-50, 50))
+ self.assertEqual(vehicle_xy, (5, 5))
+
+ @parameterized.parameters(
+ ([[1.0, 1.0]]),
+ ([[-2.2, 4.2]]),
+ ([[3.7, -10.3]]),
+ )
+ def test_frame_to_image_and_image_to_frame(self, frame_xy):
+ frame_xy = np.array(frame_xy)
+ vehicle_xy = (0, 0)
+ resolution = 1.0
+ image_xy = utils.frame_to_image_coord(frame_xy, vehicle_xy, 1 / resolution)
+ frame_xy_1 = utils.image_to_frame_coord(image_xy, vehicle_xy, resolution)
+ self.assertAllEqual(frame_xy_1, np.floor(frame_xy))
+
+ @parameterized.parameters(
+ ([[1.0, 1.0, 2.0, 2.0]]),
+ ([[-2.2, -4.2, 2.2, 4.2]]),
+ )
+ def test_frame_to_image_boxes_and_image_to_frame_boxes(self, frame_boxes):
+ frame_boxes = np.array(frame_boxes)
+ vehicle_xy = (0, 0)
+ resolution = 1.0
+ image_boxes = utils.frame_to_image_boxes(frame_boxes, vehicle_xy,
+ 1 / resolution)
+ frame_boxes_1 = utils.image_to_frame_boxes(image_boxes, vehicle_xy,
+ resolution)
+ self.assertAllClose(frame_boxes_1, frame_boxes)
+
+ def test_generate_anchors(self):
+ min_level = 1
+ max_level = 3
+ image_size = [16, 16]
+ anchor_sizes = [(2.0, 1.0)]
+ all_anchors = utils.generate_anchors(min_level, max_level, image_size,
+ anchor_sizes)
+ for level in range(min_level, max_level + 1):
+ anchors = all_anchors[str(level)]
+ stride = 2**level
+ self.assertAllEqual(anchors.shape.as_list(),
+ [image_size[0] / stride, image_size[1] / stride, 4])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/pointpillars/utils/wod_detection_evaluator.py b/official/projects/pointpillars/utils/wod_detection_evaluator.py
new file mode 100644
index 00000000000..079ff34a009
--- /dev/null
+++ b/official/projects/pointpillars/utils/wod_detection_evaluator.py
@@ -0,0 +1,309 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Detection evaluator for the Waymo Open Dataset."""
+
+import abc
+from typing import Any, Mapping
+
+from absl import logging
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.utils import utils
+from waymo_open_dataset import label_pb2
+from waymo_open_dataset.metrics.python import wod_detection_evaluator
+
+np.set_printoptions(precision=4, suppress=True)
+
+
+class _AbstractEvaluator(
+ wod_detection_evaluator.WODDetectionEvaluator, metaclass=abc.ABCMeta):
+ """WOD detection evaluation metric base class."""
+
+ def __init__(self, model_config, config=None):
+ super().__init__(config=config)
+
+ image_config = model_config.image
+ self._resolution = image_config.resolution
+ self._vehicle_xy = utils.get_vehicle_xy(image_config.height,
+ image_config.width,
+ image_config.x_range,
+ image_config.y_range)
+ self._classes = model_config.classes
+
+ def _remove_padding(self, tensor_dict: Mapping[str, Any],
+ num_valid: int) -> Mapping[str, Any]:
+ """Remove the paddings of the prediction/groundtruth data."""
+ result_tensor_dict = {}
+ gather_indices = tf.range(num_valid)
+ for k, v in tensor_dict.items():
+ if v.shape[0] < num_valid:
+ raise ValueError(
+ '{} does not have enough elements to gather, {} < {}'.format(
+ k, v.shape[0], num_valid))
+ result_tensor_dict[k] = tf.gather(v, gather_indices)
+ return result_tensor_dict
+
+ def _compact_tensors(self,
+ tensor_dict: Mapping[str, Any]) -> Mapping[str, Any]:
+ """Compact tensors by concatenating them in tuples."""
+ compact_tensor_dict = {}
+ for k, v in tensor_dict.items():
+ if isinstance(v, tuple):
+ compact_tensor_dict[k] = tf.concat(v, axis=0)
+ elif isinstance(v, dict):
+ compact_tensor_dict[k] = v
+ for dk, dv in v.items():
+ if isinstance(dv, tuple):
+ compact_tensor_dict[k][dk] = tf.concat(dv, axis=0)
+ else:
+ compact_tensor_dict[k] = v
+ return compact_tensor_dict
+
+ def _adjust_class(self, tensor_dict: Mapping[str, Any]) -> tf.Tensor:
+ """Change predicted class to what defiend by label.proto."""
+ original_type = tf.cast(tensor_dict['classes'], tf.uint8)
+ if self._classes == 'all':
+ adjusted_type = tf.where(
+ tf.equal(original_type, 3),
+ tf.ones_like(original_type) * 4,
+ original_type)
+ else:
+ adjusted_type = tf.where(
+ tf.equal(original_type, 1),
+ tf.ones_like(original_type) * utils.CLASSES[self._classes],
+ original_type)
+ return adjusted_type
+
+ @abc.abstractmethod
+ def _get_box(self, box2d: tf.Tensor, attributes: Mapping[str, tf.Tensor]):
+ """Get box from yxyx and attributes.
+
+ Args:
+ box2d: a [N, 4] tensor encoding as (ymin, xmin, ymax, xmax)
+ attributes: a {name: [N, 1]} dict
+ Returns:
+ box: a tensor representing a 2d or 3d box
+ """
+
+ def update_state(self,
+ groundtruths: Mapping[str, tf.Tensor],
+ predictions: Mapping[str, tf.Tensor]):
+ """Update the metrics state with prediction and groundtruth data.
+
+ Notations:
+ B: batch size.
+ N: number of ground truth boxes.
+ M: number of predicted boxes.
+ T: attribute size.
+
+ Args:
+ groundtruths: a dictionary of Tensors including the fields below.
+ Required fields:
+ - frame_id: a tensor of int64 of shape [B].
+ - num_detections: a tensor of int32 of shape [B].
+ - boxes: a tensor of float32 of shape [B, N, 4],
+ (ymin, xmin, ymax, xmax).
+ - classes: a tensor of int32 of shape [B, N].
+ - attributes: a dict of tensor of float32 of shape [B, N, T].
+ - difficulties: a tensor of int32 of shape [B, N].
+
+ predictions: a dictionary of tensors including the fields below.
+ Required fields:
+ - num_detections: a tensor of int32 of shape [B].
+ - boxes: a tensor of float32 of shape [B, M, 4],
+ (ymin, xmin, ymax, xmax).
+ - scores: a tensor of float32 of shape [B, M].
+ - classes: a tensor of int32 of shape [B, M].
+ - attributes: a dict of tensor of float32 of shape [B, M, T].
+ """
+ # Remove tuples from dataset.
+ groundtruths = self._compact_tensors(groundtruths)
+ predictions = self._compact_tensors(predictions)
+
+ # Adjust type.
+ gt_type = self._adjust_class(groundtruths)
+ pred_type = self._adjust_class(predictions)
+
+ batch_size = tf.shape(groundtruths['frame_id'])[0]
+ for i in tf.range(batch_size):
+ # Set ground truths
+ gt_num_detections = groundtruths['num_detections'][i]
+ gt_attributes = {}
+ for k, v in groundtruths['attributes'].items():
+ gt_attributes[k] = v[i]
+ frame_groundtruths = {
+ 'ground_truth_frame_id':
+ tf.tile([groundtruths['frame_id'][i]], [gt_num_detections]),
+ 'ground_truth_bbox':
+ self._get_box(groundtruths['boxes'][i], gt_attributes),
+ 'ground_truth_type':
+ gt_type[i],
+ 'ground_truth_difficulty':
+ tf.cast(groundtruths['difficulty'][i], tf.uint8),
+ }
+ frame_groundtruths = self._remove_padding(
+ frame_groundtruths, gt_num_detections)
+
+ # Set predictions
+ pred_num_detections = predictions['num_detections'][i]
+ pred_attributes = {}
+ for k, v in predictions['attributes'].items():
+ pred_attributes[k] = v[i]
+ frame_predictions = {
+ 'prediction_frame_id':
+ tf.tile([groundtruths['frame_id'][i]], [pred_num_detections]),
+ 'prediction_bbox':
+ self._get_box(predictions['boxes'][i], pred_attributes),
+ 'prediction_type':
+ pred_type[i],
+ 'prediction_score':
+ predictions['scores'][i],
+ 'prediction_overlap_nlz':
+ tf.zeros_like(predictions['scores'][i], dtype=tf.bool)
+ }
+ frame_predictions = self._remove_padding(
+ frame_predictions, pred_num_detections)
+
+ # Update state for this frame.
+ super().update_state(frame_groundtruths, frame_predictions)
+
+ def evaluate(self) -> Mapping[str, Any]:
+ """Compute the final metrics.
+
+ Returns:
+ metric_dict: A dict of metrics, contains following breakdown keys:
+ mAP/{class}_level_1
+ mAP/{class}_[0, 30)_level_1
+ mAP/{class}_[30, 50)_level_1
+ mAP/{class}_[50, +inf)_level_1
+ mAP/{class}_level_2
+ mAP/{class}_[0, 30)_level_2
+ mAP/{class}_[30, 50)_level_2
+ mAP/{class}_[50, +inf)_level_2
+ mAPH/{class}_level_1
+ mAPH/{class}_[0, 30)_level_1
+ mAPH/{class}_[30, 50)_level_1
+ mAPH/{class}_[50, +inf)_level_1
+ mAPH/{class}_level_2
+ mAPH/{class}_[0, 30)_level_2
+ mAPH/{class}_[30, 50)_level_2
+ mAPH/{class}_[50, +inf)_level_2
+ It also contains following keys used as public NAS rewards.
+ AP
+ APH
+ """
+ ap, aph, _, _, _, _, _ = super().evaluate()
+ metric_dict = {}
+ for i, name in enumerate(self._breakdown_names):
+ # Skip sign metrics since we don't use this type.
+ if 'SIGN' in name:
+ continue
+ # Make metric name more readable.
+ name = name.lower()
+ for c in utils.CLASSES:
+ pos = name.find(c)
+ if pos != -1:
+ name = name[pos:]
+ if self._classes == 'all' or self._classes in name:
+ metric_dict['mAP/{}'.format(name)] = ap[i]
+ metric_dict['mAPH/{}'.format(name)] = aph[i]
+
+ # Set public metrics as AP and APH.
+ if self._classes == 'all':
+ ap, aph = 0, 0
+ for c in utils.CLASSES:
+ ap += metric_dict['mAP/{}_level_1'.format(c)]
+ aph += metric_dict['mAPH/{}_level_1'.format(c)]
+ metric_dict['AP'] = ap / len(utils.CLASSES)
+ metric_dict['APH'] = aph / len(utils.CLASSES)
+ else:
+ metric_dict['AP'] = metric_dict['mAP/{}_level_1'.format(self._classes)]
+ metric_dict['APH'] = metric_dict['mAPH/{}_level_1'.format(self._classes)]
+ return metric_dict
+
+
+class Wod3dDetectionEvaluator(_AbstractEvaluator):
+ """WOD 3D detection evaluation metric class."""
+
+ def _get_box(self, box2d: tf.Tensor,
+ attributes: Mapping[str, tf.Tensor]) -> tf.Tensor:
+ """Get box from yxyx and attributes.
+
+ Args:
+ box2d: a float32 [N, 4] tensor encoding as (ymin, xmin, ymax, xmax)
+ attributes: a float32 {name: [N, 1]} dict
+ Returns:
+ box: a float32 [N, 7] tensor representing a 3d box
+ """
+ box2d = utils.image_to_frame_boxes(box2d, self._vehicle_xy,
+ self._resolution)
+ values = []
+ values.append(box2d[:, 0]) # center_x
+ values.append(box2d[:, 1]) # center_y
+ values.append(attributes['z'][:, 0]) # center_z
+ values.append(box2d[:, 2]) # length
+ values.append(box2d[:, 3]) # width
+ values.append(attributes['height'][:, 0]) # height
+ values.append(attributes['heading'][:, 0]) # heading
+ box3d = tf.stack(values, axis=-1)
+ return box3d
+
+
+class Wod2dDetectionEvaluator(_AbstractEvaluator):
+ """WOD 2D detection evaluation metric class."""
+
+ def __init__(self, image_config: Any, config: Any = None):
+ if config is None:
+ config = self._get_default_config()
+ config.box_type = label_pb2.Label.Box.TYPE_2D
+ super().__init__(image_config, config)
+
+ # use utils
+ def _get_box(self, box2d: tf.Tensor,
+ attributes: Mapping[str, tf.Tensor]) -> tf.Tensor:
+ """Get box from yxyx and attributes.
+
+ Args:
+ box2d: a float32 [N, 4] tensor encoding as (ymin, xmin, ymax, xmax)
+ attributes: a float32 {name: [N, 1]} dict
+ Returns:
+ box: a float32 [N, 5] tensor representing a 2d box with heading
+ """
+ box2d = utils.image_to_frame_boxes(box2d, self._vehicle_xy,
+ self._resolution)
+ values = []
+ values.append(box2d[:, 0]) # center_x
+ values.append(box2d[:, 1]) # center_y
+ values.append(box2d[:, 2]) # length
+ values.append(box2d[:, 3]) # width
+ values.append(attributes['heading'][:, 0]) # heading
+ box2d_h = tf.stack(values, axis=-1)
+ return box2d_h
+
+
+def create_evaluator(model_config: cfg.PointPillarsModel) -> _AbstractEvaluator:
+ """Create either 2d or 3d evaluator."""
+ attr_count = len(model_config.head.attribute_heads)
+ if attr_count == 1:
+ logging.info('Use 2D detection evaluator.')
+ return Wod2dDetectionEvaluator(model_config)
+ if attr_count == 3:
+ logging.info('Use 3D detection evaluator.')
+ return Wod3dDetectionEvaluator(model_config)
+ raise ValueError(
+ 'The length of attribute_heads should be 1 or 3, found {}'.format(
+ attr_count))
diff --git a/official/projects/pointpillars/utils/wod_processor.py b/official/projects/pointpillars/utils/wod_processor.py
new file mode 100644
index 00000000000..c81e291d27a
--- /dev/null
+++ b/official/projects/pointpillars/utils/wod_processor.py
@@ -0,0 +1,434 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""A class to process waymo open dataset."""
+
+from typing import Any, List, Mapping, Optional, Sequence, Tuple
+import zlib
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.pointpillars.configs import pointpillars as cfg
+from official.projects.pointpillars.utils import utils
+from official.vision.data.tfrecord_lib import convert_to_feature
+from waymo_open_dataset import dataset_pb2
+from waymo_open_dataset import label_pb2
+from waymo_open_dataset.utils import frame_utils
+
+# The minimum length of required labeling boxes.
+_MIN_BOX_LENGTH = 1e-2
+# The seed for random generator.
+_RANDOM_SEED = 42
+
+
+class WodProcessor:
+ """The class to process waymo-open-dataset-tf-2-6-0.
+
+ https://github.com/waymo-research/waymo-open-dataset
+ """
+
+ def __init__(self,
+ image_config: cfg.ImageConfig,
+ pillars_config: cfg.PillarsConfig):
+ self._x_range = image_config.x_range
+ self._y_range = image_config.y_range
+ self._z_range = image_config.z_range
+
+ self._resolution = image_config.resolution
+ self._one_over_resolution = 1.0 / self._resolution
+ self._image_height = image_config.height
+ self._image_width = image_config.width
+ self._vehicle_xy = utils.get_vehicle_xy(image_height=image_config.height,
+ image_width=image_config.width,
+ x_range=image_config.x_range,
+ y_range=image_config.y_range)
+
+ self._num_pillars = pillars_config.num_pillars
+ self._num_points_per_pillar = pillars_config.num_points_per_pillar
+ self._num_features_per_point = pillars_config.num_features_per_point
+
+ self._rng = np.random.default_rng(seed=_RANDOM_SEED)
+
+ def _parse_range_image_and_top_pose(
+ self, frame: dataset_pb2.Frame
+ ) -> Tuple[Mapping[int, List[dataset_pb2.MatrixFloat]],
+ Optional[dataset_pb2.MatrixFloat]]:
+ """Parse range images and top pose given a frame.
+
+ Args:
+ frame: A frame message in wod dataset.proto.
+
+ Returns:
+ range_images: A dict of {laser_name: [range_image_return]},
+ each range_image_return is a MatrixFloat with shape (H, W, 4).
+ range_image_top_pose: Range image pixel pose for top lidar,
+ a MatrixFloat with shape (H, W, 6).
+ """
+ range_images = {}
+ range_image_top_pose = None
+
+ # Parse lidar laser data from two returns, ri_return1 is the first return,
+ # ri_return2 is the second return. Also get the top lidar pose from the
+ # first return of the top lidar.
+ for laser in frame.lasers:
+ if laser.ri_return1.range_image_compressed:
+ ri_str = zlib.decompress(laser.ri_return1.range_image_compressed)
+ ri = dataset_pb2.MatrixFloat()
+ ri.ParseFromString(ri_str)
+ range_images[int(laser.name)] = [ri]
+
+ if laser.name == dataset_pb2.LaserName.TOP:
+ pos_str = zlib.decompress(
+ laser.ri_return1.range_image_pose_compressed)
+ range_image_top_pose = dataset_pb2.MatrixFloat()
+ range_image_top_pose.ParseFromString(pos_str)
+
+ if laser.ri_return2.range_image_compressed:
+ ri_str = zlib.decompress(laser.ri_return2.range_image_compressed)
+ ri = dataset_pb2.MatrixFloat()
+ ri.ParseFromString(ri_str)
+ range_images[int(laser.name)].append(ri)
+
+ return range_images, range_image_top_pose
+
+ def _convert_range_image_to_point_cloud(
+ self,
+ frame: dataset_pb2.Frame,
+ range_images: Mapping[int, List[dataset_pb2.MatrixFloat]],
+ range_image_top_pose: dataset_pb2.MatrixFloat,
+ ri_index: int) -> np.ndarray:
+ """Convert range images (polar) to point cloud (Cartesian).
+
+ Args:
+ frame: A frame message in wod dataset.proto.
+ range_images: A dict of {laser_name: [range_image_return]}.
+ range_image_top_pose: Range image pixel pose for top lidar.
+ ri_index: 0 for the first return, 1 for the second return.
+
+ Returns:
+ point_cloud: a np array with shape (M, F),
+ each point has F attributes [x, y, z, intensity, elongation].
+ """
+ calibrations = sorted(
+ frame.context.laser_calibrations, key=lambda c: c.name)
+ point_cloud = []
+
+ cartesian_tensor = frame_utils.convert_range_image_to_cartesian(
+ frame, range_images, range_image_top_pose, ri_index, False)
+ for calibration in calibrations:
+ # Get range_image for this lidar calibration.
+ range_image = range_images[calibration.name][ri_index]
+ range_image_tensor = tf.reshape(
+ tf.convert_to_tensor(value=range_image.data), range_image.shape.dims)
+
+ # Stack xyz, intensity, elongation together.
+ xyz_tensor = cartesian_tensor[calibration.name]
+ intensity_tensor = range_image_tensor[..., 1:2]
+ elongation_tensor = range_image_tensor[..., 2:3]
+ points_tensor = tf.concat(
+ [xyz_tensor, intensity_tensor, elongation_tensor], axis=-1)
+
+ # Only select points if:
+ # 1. its range is greater than 0m, and
+ # 2. it is not in any no-label-zone
+ distance_mask = range_image_tensor[..., 0] > 0
+ nlz_mask = range_image_tensor[..., 3] == -1.0
+ mask = tf.logical_and(distance_mask, nlz_mask)
+ points_tensor = tf.gather_nd(points_tensor, tf.where(mask))
+
+ point_cloud.append(points_tensor.numpy())
+ point_cloud = np.concatenate(point_cloud, axis=0)
+
+ # Shuffle points to make the order independent to the range image.
+ # Otherwise, the pillars close to the auto vehicle would be empty if distant
+ # pillars have exceeded the maximum number.
+ self._rng.shuffle(point_cloud)
+
+ return point_cloud
+
+ def extract_point_cloud(
+ self, frame: dataset_pb2.Frame) -> Tuple[np.ndarray, np.ndarray]:
+ """Extract point cloud from frame proto.
+
+ Args:
+ frame: A frame message in wod dataset.proto.
+
+ Returns:
+ points: The point cloud, a float array with shape (M, F).
+ points_location: The pseudo image col/row of points, an array (M, 2),
+ col/row, int32.
+ """
+ # Get point cloud from range images
+ range_images, range_image_top_pose = self._parse_range_image_and_top_pose(
+ frame)
+ points_r1 = self._convert_range_image_to_point_cloud(
+ frame, range_images, range_image_top_pose, 0) # pyrefly: ignore[bad-argument-type]
+ points_r2 = self._convert_range_image_to_point_cloud(
+ frame, range_images, range_image_top_pose, 1) # pyrefly: ignore[bad-argument-type]
+ points = np.concatenate([points_r1, points_r2], axis=0)
+
+ # Get image col/row of points
+ points_location = utils.frame_to_image_coord(
+ points[:, 0:2], self._vehicle_xy, self._one_over_resolution)
+
+ # Select points locating inside the range.
+ selection = np.where((points_location[:, 0] >= 0) &
+ (points_location[:, 0] < self._image_width) &
+ (points_location[:, 1] >= 0) &
+ (points_location[:, 1] < self._image_height) &
+ (points[:, 2] >= self._z_range[0]) &
+ (points[:, 2] <= self._z_range[1]))
+ points = points[selection]
+ points_location = points_location[selection]
+ return points, points_location
+
+ def compute_pillars(
+ self,
+ points: np.ndarray,
+ points_location: np.ndarray) -> Tuple[tf.Tensor, tf.Tensor, int]:
+ """Compute pillars from point cloud.
+
+ Args:
+ points: The point cloud, a np array with shape (M, F).
+ points_location: The pseudo image col/row of points, a np array (M, 2).
+
+ Returns:
+ pillar_features: A tensor with shape (P, N, D).
+ pillar_indices: A tensor with shape (P, 2), row/col, int32.
+ pillars_count: The number of computed pillars before pad/trim.
+
+ Notations:
+ h: image height
+ w: image widht
+ p: number of pillars per example after trimming or padding
+ n: number of points per pillar
+ d: number of features per point after processing
+ f: number of features per point before processing
+ k: number of pillars before trimming or padding
+ """
+ h, w = self._image_height, self._image_width
+ p, n, d = (self._num_pillars, self._num_points_per_pillar,
+ self._num_features_per_point)
+ f = points.shape[-1]
+
+ grid_num_points = np.zeros((h, w), dtype=np.int32)
+ grid_locations = np.zeros((h, w, 2), dtype=np.int32)
+ grid_points = np.zeros((h, w, n, f), dtype=np.float32)
+
+ # Fill points into 2D grid.
+ for i, (point, (c, r)) in enumerate(zip(points, points_location)):
+ point_count = grid_num_points[r][c]
+ if point_count == n:
+ continue
+ grid_num_points[r][c] += 1
+ grid_locations[r][c] = [c, r]
+ grid_points[r][c][point_count][:] = point[:]
+
+ # Select k non-empty pillars randomly.
+ selection = np.where(grid_num_points > 0)
+ selection = [(i, j) for i, j in zip(selection[0], selection[1])]
+ self._rng.shuffle(selection)
+ selection = ([i[0] for i in selection], [i[1] for i in selection])
+
+ k = len(selection[0])
+ # (k,)
+ pillar_num_points = grid_num_points[selection]
+ # (k, 2)
+ pillar_locations = grid_locations[selection]
+ # (k, n, f)
+ pillar_points = grid_points[selection]
+
+ # Pad or trim to p pillars.
+ # (p,)
+ pillar_num_points = utils.pad_or_trim_to_shape(pillar_num_points, [p])
+ # (p, 2)
+ pillar_locations = utils.pad_or_trim_to_shape(pillar_locations, [p, 2])
+ # (p, n, f)
+ pillar_points = utils.pad_or_trim_to_shape(pillar_points, [p, n, f])
+
+ # Compute pillar features.
+ # (p, n, 3)
+ pillar_xyz = pillar_points[..., 0:3]
+ # (p, n, f-3)
+ pillar_others = pillar_points[..., 3:]
+ # (p, 1, 3)
+ pillar_sum_xyz = np.sum(pillar_xyz, axis=1, keepdims=True)
+ num_points = np.maximum(
+ pillar_num_points, 1.0, dtype=np.float32).reshape(p, 1, 1)
+ pillar_mean_xyz = pillar_sum_xyz / num_points
+ # (p, n, 3)
+ pillar_dxyz = pillar_xyz - pillar_mean_xyz
+ # (p, 1, 2)
+ pillar_center_xy = utils.image_to_frame_coord(
+ pillar_locations, self._vehicle_xy, self._resolution).reshape(p, 1, 2)
+
+ # Concat all features together, (k, n, d).
+ pillar_features = np.concatenate([
+ pillar_dxyz,
+ pillar_others,
+ np.tile(pillar_mean_xyz, (1, n, 1)),
+ np.tile(pillar_center_xy, (1, n, 1))], axis=-1)
+ # Get pillar indices [row, col], (k, 2).
+ pillar_locations[:, [0, 1]] = pillar_locations[:, [1, 0]]
+
+ utils.assert_shape(pillar_features, [p, n, d])
+ utils.assert_shape(pillar_locations, [p, 2])
+ pillar_features = tf.convert_to_tensor(pillar_features, dtype=tf.float32)
+ pillar_locations = tf.convert_to_tensor(pillar_locations, dtype=tf.int32)
+ return pillar_features, pillar_locations, k
+
+ def _adjust_label_type(self, label: label_pb2.Label) -> int:
+ # Only care about (vehicle, pedestrian, cyclist) types, override sign type
+ # with cyclist. After this, the types mapping would be:
+ # 0: unknown, 1: vehicle, 2: pedestrian, 3: cyclist
+ if label.type == label_pb2.Label.TYPE_CYCLIST:
+ return 3
+ return int(label.type)
+
+ def _adjust_difficulty_level(self, label: label_pb2.Label) -> int:
+ # Extend level-2 difficulty labels with boxes which have very little lidar
+ # points, since the model is a single modality (lidar) model.
+ if (label.num_lidar_points_in_box <= 5 or
+ label.detection_difficulty_level == label_pb2.Label.LEVEL_2):
+ return 2
+ return 1
+
+ def extract_labels(self, frame: dataset_pb2.Frame) -> Sequence[tf.Tensor]:
+ """Extract bounding box labels from frame proto.
+
+ Args:
+ frame: A frame message in wod dataset.proto.
+
+ Returns:
+ labels: A sequence of processed tensors.
+ """
+ xmin = []
+ xmax = []
+ ymin = []
+ ymax = []
+ classes = []
+ heading = []
+ z = []
+ height = []
+ difficulty = []
+
+ for label in frame.laser_labels:
+ box = label.box
+
+ # Skip boxes if it doesn't contain any lidar points.
+ # WARNING: Do not enable this filter when using v.1.0.0 data.
+ if label.num_lidar_points_in_box == 0:
+ continue
+
+ # Skip boxes if it's type is SIGN.
+ if label.type == label_pb2.Label.TYPE_SIGN:
+ continue
+
+ # Skip boxes if its z is out of range.
+ half_height = box.height * 0.5
+ if (box.center_z - half_height < self._z_range[0] or
+ box.center_z + half_height > self._z_range[1]):
+ continue
+
+ # Get boxes in image coordinate.
+ frame_box = np.array([[box.center_x, box.center_y, box.length,
+ box.width]])
+ image_box = utils.frame_to_image_boxes(frame_box, self._vehicle_xy,
+ self._one_over_resolution)
+ # Skip empty boxes.
+ image_box = utils.clip_boxes(image_box, self._image_height,
+ self._image_width)[0]
+ y0, x0, y1, x1 = image_box
+ if np.abs(y0 - y1) < _MIN_BOX_LENGTH or np.abs(x0 - x1) < _MIN_BOX_LENGTH:
+ continue
+
+ label_cls = self._adjust_label_type(label)
+ level = self._adjust_difficulty_level(label)
+
+ classes.append(label_cls)
+ ymin.append(y0)
+ xmin.append(x0)
+ ymax.append(y1)
+ xmax.append(x1)
+ heading.append(box.heading)
+ z.append(box.center_z)
+ height.append(box.height)
+ difficulty.append(level)
+
+ classes = tf.convert_to_tensor(classes, dtype=tf.int32)
+ ymin = tf.convert_to_tensor(ymin, dtype=tf.float32)
+ xmin = tf.convert_to_tensor(xmin, dtype=tf.float32)
+ ymax = tf.convert_to_tensor(ymax, dtype=tf.float32)
+ xmax = tf.convert_to_tensor(xmax, dtype=tf.float32)
+ heading = tf.convert_to_tensor(heading, dtype=tf.float32)
+ z = tf.convert_to_tensor(z, dtype=tf.float32)
+ height = tf.convert_to_tensor(height, dtype=tf.float32)
+ difficulty = tf.convert_to_tensor(difficulty, dtype=tf.int32)
+
+ # NOTE: This function might be called by an online data loader in a
+ # tf.py_function wrapping fashion. But tf.py_function doesn't support
+ # dict return type, so we have to return a sequence of unpacked.
+ return classes, ymin, xmin, ymax, xmax, heading, z, height, difficulty
+
+ def process_one_frame(self, frame: dataset_pb2.Frame) -> Sequence[Any]:
+ """Compute features and labels.
+
+ Args:
+ frame: A frame message in wod dataset.proto.
+
+ Returns:
+ labels: A sequence of processed tensors.
+ """
+ timestamp = frame.timestamp_micros
+ timestamp = tf.convert_to_tensor(timestamp, dtype=tf.int64)
+ points, points_location = self.extract_point_cloud(frame)
+ pillars, indices, _ = self.compute_pillars(points, points_location)
+ (classes, ymin, xmin, ymax, xmax, heading, z, height,
+ difficulty) = self.extract_labels(frame)
+
+ # NOTE: This function might be called by an online data loader in a
+ # tf.py_function wrapping fashion. But tf.py_function doesn't support
+ # dict return type, so we have to return a sequence of unpacked.
+ return (timestamp, pillars, indices, classes, ymin, xmin, ymax, xmax,
+ heading, z, height, difficulty)
+
+ def process_and_convert_to_tf_example(
+ self, frame: dataset_pb2.Frame) -> tf.train.Example:
+ """Processes one wod source tfrecord.
+
+ Args:
+ frame: The parsed wod frame proto.
+
+ Returns:
+ example: The tf example converted from frame.
+ """
+ (timestamp, pillars, indices, classes, ymin, xmin, ymax, xmax,
+ heading, z, height, difficulty) = self.process_one_frame(frame)
+ feature = {
+ 'frame_id': convert_to_feature(timestamp.numpy(), 'int64'),
+ 'pillars': convert_to_feature(pillars.numpy().tobytes(), 'bytes'),
+ 'indices': convert_to_feature(indices.numpy().tobytes(), 'bytes'),
+ 'bbox/class': convert_to_feature(classes.numpy(), 'int64_list'),
+ 'bbox/ymin': convert_to_feature(ymin.numpy(), 'float_list'),
+ 'bbox/xmin': convert_to_feature(xmin.numpy(), 'float_list'),
+ 'bbox/ymax': convert_to_feature(ymax.numpy(), 'float_list'),
+ 'bbox/xmax': convert_to_feature(xmax.numpy(), 'float_list'),
+ 'bbox/heading': convert_to_feature(heading.numpy(), 'float_list'),
+ 'bbox/z': convert_to_feature(z.numpy(), 'float_list'),
+ 'bbox/height': convert_to_feature(height.numpy(), 'float_list'),
+ 'bbox/difficulty': convert_to_feature(difficulty.numpy(), 'int64_list'),
+ }
+ example = tf.train.Example(features=tf.train.Features(feature=feature))
+ return example
diff --git a/official/projects/pruning/configs/__init__.py b/official/projects/pruning/configs/__init__.py
index 4425d4bd55b..b9be2aee76a 100644
--- a/official/projects/pruning/configs/__init__.py
+++ b/official/projects/pruning/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/pruning/configs/experiments/image_classification/imagenet_mobilenetv2_pruning_gpu.yaml b/official/projects/pruning/configs/experiments/image_classification/imagenet_mobilenetv2_pruning_gpu.yaml
index 2849a5e7d91..af1dca89d30 100644
--- a/official/projects/pruning/configs/experiments/image_classification/imagenet_mobilenetv2_pruning_gpu.yaml
+++ b/official/projects/pruning/configs/experiments/image_classification/imagenet_mobilenetv2_pruning_gpu.yaml
@@ -18,12 +18,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 1024
dtype: 'float32'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 1024
dtype: 'float32'
diff --git a/official/projects/pruning/configs/experiments/image_classification/imagenet_resnet50_pruning_gpu.yaml b/official/projects/pruning/configs/experiments/image_classification/imagenet_resnet50_pruning_gpu.yaml
index 8d5d7808330..dfa298a8a44 100644
--- a/official/projects/pruning/configs/experiments/image_classification/imagenet_resnet50_pruning_gpu.yaml
+++ b/official/projects/pruning/configs/experiments/image_classification/imagenet_resnet50_pruning_gpu.yaml
@@ -15,12 +15,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 1024
dtype: 'float32'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 1024
dtype: 'float32'
diff --git a/official/projects/pruning/configs/image_classification.py b/official/projects/pruning/configs/image_classification.py
index 4ab5952a85f..a4b641fafd3 100644
--- a/official/projects/pruning/configs/image_classification.py
+++ b/official/projects/pruning/configs/image_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/pruning/configs/image_classification_test.py b/official/projects/pruning/configs/image_classification_test.py
index f52505d27cc..bb8b13408b6 100644
--- a/official/projects/pruning/configs/image_classification_test.py
+++ b/official/projects/pruning/configs/image_classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for image_classification."""
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
diff --git a/official/projects/pruning/registry_imports.py b/official/projects/pruning/registry_imports.py
index 847cca42e14..fb2798c497d 100644
--- a/official/projects/pruning/registry_imports.py
+++ b/official/projects/pruning/registry_imports.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/pruning/tasks/__init__.py b/official/projects/pruning/tasks/__init__.py
index 9320e3087f5..d6b1463a283 100644
--- a/official/projects/pruning/tasks/__init__.py
+++ b/official/projects/pruning/tasks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/pruning/tasks/image_classification.py b/official/projects/pruning/tasks/image_classification.py
index 6b81788289a..ed96e557f5a 100644
--- a/official/projects/pruning/tasks/image_classification.py
+++ b/official/projects/pruning/tasks/image_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Image classification task definition."""
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.core import task_factory
@@ -49,7 +49,7 @@ class ImageClassificationTask(image_classification.ImageClassificationTask):
),
}
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
"""Builds classification model with pruning."""
model = super(ImageClassificationTask, self).build_model()
if self.task_config.pruning is None:
@@ -57,7 +57,7 @@ def build_model(self) -> tf.keras.Model:
pruning_cfg = self.task_config.pruning
- prunable_model = tf.keras.models.clone_model(
+ prunable_model = tf_keras.models.clone_model(
model,
clone_function=self._make_block_prunable,
)
@@ -113,9 +113,9 @@ def build_model(self) -> tf.keras.Model:
return pruned_model
def _make_block_prunable(
- self, layer: tf.keras.layers.Layer) -> tf.keras.layers.Layer:
- if isinstance(layer, tf.keras.Model):
- return tf.keras.models.clone_model(
+ self, layer: tf_keras.layers.Layer) -> tf_keras.layers.Layer:
+ if isinstance(layer, tf_keras.Model):
+ return tf_keras.models.clone_model(
layer, input_tensors=None, clone_function=self._make_block_prunable)
if layer.__class__ not in self._BLOCK_LAYER_SUFFIX_MAP:
@@ -139,7 +139,7 @@ def collect_prunable_layers(model):
"""Recursively collect the prunable layers in the model."""
prunable_layers = []
for layer in model.layers:
- if isinstance(layer, tf.keras.Model):
+ if isinstance(layer, tf_keras.Model):
prunable_layers += collect_prunable_layers(layer)
if layer.__class__.__name__ == 'PruneLowMagnitude':
prunable_layers.append(layer)
diff --git a/official/projects/pruning/tasks/image_classification_test.py b/official/projects/pruning/tasks/image_classification_test.py
index 09c5a0aa05a..bb2a797626a 100644
--- a/official/projects/pruning/tasks/image_classification_test.py
+++ b/official/projects/pruning/tasks/image_classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl.testing import parameterized
import numpy as np
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official import vision
diff --git a/official/projects/pruning/train.py b/official/projects/pruning/train.py
index e1d6e3c9416..2561c11fbec 100644
--- a/official/projects/pruning/train.py
+++ b/official/projects/pruning/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/qat/__init__.py b/official/projects/qat/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/nlp/README.md b/official/projects/qat/nlp/README.md
new file mode 100644
index 00000000000..37d79231d0e
--- /dev/null
+++ b/official/projects/qat/nlp/README.md
@@ -0,0 +1,100 @@
+# Quantization Aware Training for NLP Models
+
+## Description
+
+This project includes quantization aware training (QAT) code for NLP models.
+These are examples to show how to apply the Model Optimization Toolkit's
+[quantization aware training API](https://www.tensorflow.org/model_optimization/guide/quantization/training).
+Compared to post-training quantization (PTQ), QAT can minimize the quality loss
+from quantization, while still achieving the speed-up from integer quantization.
+
+Currently, we support a limited number of NLP tasks & models. We will keep
+adding support for other tasks and models in the next releases.
+
+## Maintainers
+
+- Jaehong Kim ([Xhark](https://github.com/Xhark))
+- Rino Lee ([rino20](https://github.com/rino20))
+
+## Requirements
+
+[](https://badge.fury.io/py/tensorflow)
+[](https://badge.fury.io/py/tf-models-official)
+
+## Results
+### MobileBERT
+
+Model name | SQUAD F1 (float) | SQUAD F1 (PTQ) | SQUAD F1 (QAT) | download | links
+:-------------------- | ---------------: | -------------: | -------------: | ---------: | ----:
+MobileBERT-EdgeTPU-XS | 88.02% | 84.96% | 85.42% | [FP32](https://storage.googleapis.com/tf_model_garden/nlp/qat/mobilebert/model_fp32.tflite) \| [INT8](https://storage.googleapis.com/tf_model_garden/nlp/qat/mobilebert/model_int8_ptq.tflite) \| [QAT INT8](https://storage.googleapis.com/tf_model_garden/nlp/qat/mobilebert/model_qat.tflite) ([ckpt](https://storage.googleapis.com/tf_model_garden/nlp/qat/mobilebert/mobilebert_qat.tar.gz)) | [tensorboard](https://tensorboard.dev/experiment/ky0gSa6nQva2a5ppL4Mtzw/#scalars)
+
+Please follow
+[MobileBERT QAT Tutorial Colab notebook](https://colab.research.google.com/github/tensorflow/models/blob/master/official/projects/qat/nlp/docs/MobileBERT_QAT_tutorial.ipynb)
+to try exported models.
+
+## Training
+
+It can run on Google Cloud Platform using Cloud TPU.
+[Here](https://cloud.google.com/tpu/docs/how-to) is the instruction of using
+Cloud TPU. Follow the below instructions to set up Cloud TPU and launch
+training, using mobilebert as an exmaple:
+
+```shell
+
+# First, Download the pre-trained floating point model as QAT needs to finetune it.
+gsutil cp gs://tf_model_garden/nlp/qat/mobilebert/mobilebert_fp32_ckpt.tar.gz /tmp/qat/
+
+# Extract the checkpoint.
+tar -xvzf /tmp/qat/mobilebert_fp32_ckpt.tar.gz
+
+# Convert the float checkpoint to QAT checkpoint.
+$ python3 pretrained_checkpoint_converter.py \
+ --experiment=bert/squad \
+ --config_file=/edgetpu/nlp/experiments/downstream_tasks/mobilebert_edgetpu_xs.yaml \
+ --config_file=/edgetpu/nlp/experiments/downstream_tasks/squad_v1.yaml \
+ --experiment_qat=bert/squad_qat \
+ --config_file_qat=/edgetpu/nlp/experiments/downstream_tasks/mobilebert_edgetpu_xs.yaml \
+ --config_file_qat=/qat/nlp/configs/experiments/squad_v1_mobilebert_xs_qat_1gpu.yaml \
+ --pretrained_checkpoint= \ # Example: /tmp/qat/mobilebert_fp32_ckpt
+ --output_checkpoint= # Example: /tmp/qat/mobilebert_fp32_ckpt_qat
+
+# Launch training. Note that we override the checkpoint path in the config file by "params_override" to supply the correct checkpoint.
+PARAMS_OVERRIDE="task.quantization.pretrained_original_checkpoint=/tmp/qat/mobilebert_fp32_ckpt_qat"
+EXPERIMENT=bert/squad_qat # Experiment type according to the subtask. Example: 'bert/squad_qat'
+TPU_NAME="" # The name assigned while creating a Cloud TPU.
+MODEL_DIR="gs://" # Model artifacts directory for the training.
+$ python3 train.py \
+ --experiment=${EXPERIMENT} \
+ --config_file_qat=/edgetpu/nlp/experiments/downstream_tasks/mobilebert_edgetpu_xs.yaml \
+ --config_file_qat=/qat/nlp/configs/experiments/squad_v1_mobilebert_xs_qat_1gpu.yaml \
+ --model_dir=${MODEL_DIR} \
+ --tpu=$TPU_NAME \
+ --params_override=${PARAMS_OVERRIDE}
+ --mode=train
+```
+
+## Evaluation
+
+Please Run below command for evaluation.
+
+```shell
+EXPERIMENT=bert/squad_qat # Experiment type according to the subtask. Example: 'bert/squad_qat'
+TPU_NAME="" # The name assigned while creating a Cloud TPU.
+MODEL_DIR="gs://" # Model artifacts directory for the training.
+$ python3 train.py \
+ --experiment=${EXPERIMENT} \
+ --config_file_qat=/edgetpu/nlp/experiments/downstream_tasks/mobilebert_edgetpu_xs.yaml \
+ --config_file_qat=/qat/nlp/configs/experiments/squad_v1_mobilebert_xs_qat_1gpu.yaml \
+ --model_dir=${MODEL_DIR} \
+ --tpu=$TPU_NAME \
+ --mode=eval
+```
+
+## License
+
+[](https://opensource.org/licenses/Apache-2.0)
+
+This project is licensed under the terms of the **Apache License 2.0**.
+
+
+
diff --git a/official/projects/qat/nlp/__init__.py b/official/projects/qat/nlp/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/nlp/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/vit/configs/__init__.py b/official/projects/qat/nlp/configs/__init__.py
similarity index 81%
rename from official/projects/vit/configs/__init__.py
rename to official/projects/qat/nlp/configs/__init__.py
index fa295acb873..45a20dd0470 100644
--- a/official/projects/vit/configs/__init__.py
+++ b/official/projects/qat/nlp/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,4 +14,4 @@
"""Configs package definition."""
-from official.projects.vit.configs import image_classification
+from official.projects.qat.nlp.configs import finetuning_experiments
diff --git a/official/projects/qat/nlp/configs/experiments/squad_v1_mobilebert_xs_qat_1gpu.yaml b/official/projects/qat/nlp/configs/experiments/squad_v1_mobilebert_xs_qat_1gpu.yaml
new file mode 100644
index 00000000000..2aa87c6a268
--- /dev/null
+++ b/official/projects/qat/nlp/configs/experiments/squad_v1_mobilebert_xs_qat_1gpu.yaml
@@ -0,0 +1,52 @@
+task:
+ # hub_module_url: 'gs://**/panzf/mobilebert/tfhub/'
+ max_answer_length: 30
+ n_best_size: 20
+ null_score_diff_threshold: 0.0
+ init_checkpoint: gs://**/mobilebert_xs_120_29708148/ckpt-19420-1
+ train_data:
+ drop_remainder: true
+ global_batch_size: 4
+ input_path: gs://**/tp/bert/squad_v1.1/train.tf_record
+ is_training: true
+ seq_length: 384
+ validation_data:
+ do_lower_case: true
+ doc_stride: 128
+ drop_remainder: false
+ global_batch_size: 6
+ input_path: gs://**/squad/dev-v1.1.json
+ is_training: false
+ query_length: 64
+ seq_length: 384
+ tokenization: WordPiece
+ version_2_with_negative: false
+ vocab_file: gs://**/panzf/ttl-30d/mobilebert/tf2_checkpoint/vocab.txt
+trainer:
+ checkpoint_interval: 8000
+ max_to_keep: 5
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 56000
+ end_learning_rate: 0.0
+ initial_learning_rate: 1.5e-05
+ power: 1.0
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ # 10% of total training steps
+ warmup_steps: 5600
+ type: polynomial
+ steps_per_loop: 8000
+ summary_interval: 8000
+ # 7 epochs for training
+ train_steps: 56000
+ validation_interval: 8000
+ validation_steps: 1808
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'final_f1'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/qat/nlp/configs/experiments/squad_v1_qat_1gpu.yaml b/official/projects/qat/nlp/configs/experiments/squad_v1_qat_1gpu.yaml
new file mode 100644
index 00000000000..854585cd481
--- /dev/null
+++ b/official/projects/qat/nlp/configs/experiments/squad_v1_qat_1gpu.yaml
@@ -0,0 +1,50 @@
+task:
+ hub_module_url: ''
+ max_answer_length: 30
+ n_best_size: 20
+ null_score_diff_threshold: 0.0
+ init_checkpoint: gs://**/bert_base_squad_qat/ckpt-1
+ train_data:
+ drop_remainder: true
+ global_batch_size: 6
+ input_path: gs://**/tp/bert/squad_v1.1/train.tf_record
+ is_training: true
+ seq_length: 384
+ validation_data:
+ do_lower_case: true
+ doc_stride: 128
+ drop_remainder: false
+ global_batch_size: 6
+ input_path: gs://**/squad/dev-v1.1.json
+ is_training: false
+ query_length: 64
+ seq_length: 384
+ tokenization: WordPiece
+ version_2_with_negative: false
+ vocab_file: gs://cloud-tpu-checkpoints/bert/keras_bert/uncased_L-12_H-768_A-12/vocab.txt
+trainer:
+ checkpoint_interval: 8000
+ max_to_keep: 5
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 29599
+ end_learning_rate: 0.0
+ initial_learning_rate: 0.00003
+ power: 1.0
+ type: polynomial
+ optimizer:
+ type: adamw
+ warmup:
+ polynomial:
+ power: 1
+ warmup_steps: 2960
+ type: polynomial
+ steps_per_loop: 8000
+ summary_interval: 8000
+ train_steps: 29599
+ validation_interval: 8000
+ validation_steps: 1808
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_eval_metric: 'final_f1'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/qat/nlp/configs/finetuning_experiments.py b/official/projects/qat/nlp/configs/finetuning_experiments.py
new file mode 100644
index 00000000000..57c42389e6e
--- /dev/null
+++ b/official/projects/qat/nlp/configs/finetuning_experiments.py
@@ -0,0 +1,35 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Fine-tuning configuration definition."""
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.nlp.configs import finetuning_experiments
+from official.projects.qat.nlp.tasks import question_answering
+
+
+@exp_factory.register_config_factory('bert/squad_qat')
+def bert_squad() -> cfg.ExperimentConfig:
+ """BERT Squad V1/V2 with QAT."""
+ config = finetuning_experiments.bert_squad()
+ task = question_answering.QuantizedModelQAConfig.from_args(
+ **config.task.as_dict())
+
+ # Copy QADataConfig objects.
+ task.train_data = config.task.train_data
+ task.validation_data = config.task.validation_data
+ config.task = task
+
+ return config
diff --git a/official/projects/qat/nlp/docs/MobileBERT_QAT_tutorial.ipynb b/official/projects/qat/nlp/docs/MobileBERT_QAT_tutorial.ipynb
new file mode 100644
index 00000000000..8c40e6a45b6
--- /dev/null
+++ b/official/projects/qat/nlp/docs/MobileBERT_QAT_tutorial.ipynb
@@ -0,0 +1,249 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "3E96e1UKQ8uR"
+ },
+ "source": [
+ "# MobileBERT QAT Tutorial\n",
+ "\n",
+ "This notebook provides a basic example code to build, run, and fine-tune [MobileBERT](https://arxiv.org/pdf/2004.02984.pdf) with QAT toolkit.\n",
+ "\n",
+ "Pretrained models downloaded from the [TensorFlow Hub](https://tfhub.dev/google/qat/nlp/mobilebert_xs_qat) and the [TensorFlow Model Garden](https://github.com/tensorflow/models/tree/master/official/projects/qat/nlp), which are both trained on [SQuAD](https://deepmind.com/research/open-source/kinetics) dateset for Q\u0026A task. You will run inference the models with dummy inputs."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "8_oLnvJy7kz5"
+ },
+ "source": [
+ "## Setup"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "s3khsunT7kWa"
+ },
+ "outputs": [],
+ "source": [
+ "# Install packages\n",
+ "\n",
+ "# tf-models-official is the stable Model Garden package\n",
+ "# tf-models-nightly includes latest changes\n",
+ "!pip install -q tf-models-nightly"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "dI_1csl6Q-gH"
+ },
+ "outputs": [],
+ "source": [
+ "# Run imports\n",
+ "import os\n",
+ "\n",
+ "import numpy as np\n",
+ "import tensorflow as tf\n",
+ "import tensorflow_hub as hub"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hjFoFGXcv7cA"
+ },
+ "source": [
+ "## Launch QAT Training\n",
+ "\n",
+ "Follow the [training guideline](https://github.com/tensorflow/models/tree/master/official/projects/qat/nlp#training) to start QAT training using the pretrained checkpoint."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Tz68XJhdwzDk"
+ },
+ "source": [
+ "## Running model from TFHub\n",
+ "\n",
+ "Running QAT trained MobileBERT model from tfhub. Note that it contains Fake-quant op and all ops are float32. It becomes actual int8 op when you convert them to TFLite using TFLite converter."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "cJilACgPw10a",
+ "outputId": "a1d9ec5b-7a17-440f-e445-9fc868dcada4"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "(1, 384)\n",
+ "(1, 384)\n"
+ ]
+ }
+ ],
+ "source": [
+ "loaded_obj = hub.load(\"https://tfhub.dev/google/qat/nlp/mobilebert_xs_qat/1\")\n",
+ "serving_model = loaded_obj.signatures['serving_default']\n",
+ "\n",
+ "# Dummy inputs\n",
+ "input_type_ids = tf.zeros(shape=[1, 384], dtype=tf.int32)\n",
+ "input_word_ids = tf.zeros(shape=[1, 384], dtype=tf.int32)\n",
+ "input_mask = tf.zeros(shape=[1, 384], dtype=tf.int32)\n",
+ "\n",
+ "bert_inputs = dict(\n",
+ " input_type_ids=input_type_ids, input_word_ids=input_word_ids, input_mask=input_mask)\n",
+ "\n",
+ "bert_outputs = serving_model(**bert_inputs)\n",
+ "\n",
+ "start_logits = bert_outputs[\"start_logits\"]\n",
+ "end_logits = bert_outputs[\"end_logits\"]\n",
+ "\n",
+ "print(start_logits.shape)\n",
+ "print(end_logits.shape)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "6g0tuFvf71S9"
+ },
+ "source": [
+ "## Running TFLite Model Inference\n",
+ "Running inference with trained quantized TFLite model with dummy dataset. We assume that data is already converted to integer from an input string using vocabulary."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "QpsXxASSwfVZ",
+ "outputId": "452d210a-67ea-4dfc-a8b1-f6136992ab99"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "100 33.8M 100 33.8M 0 0 102M 0 --:--:-- --:--:-- --:--:-- 101M\n"
+ ]
+ }
+ ],
+ "source": [
+ "# First download the TFLite model.\n",
+ "! curl https://storage.googleapis.com/tf_model_garden/nlp/qat/mobilebert/model_qat.tflite --output model_qat.tflite"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Rw1OsQ_O1LJC"
+ },
+ "outputs": [],
+ "source": [
+ "def get_dequantized_tensor(interpreter, output_detail):\n",
+ " if ('quantization' not in output_detail or\n",
+ " np.dtype(output_detail['dtype']) == np.dtype(np.float32)):\n",
+ " return interpreter.get_tensor(output_detail['index'])\n",
+ " output_scale, output_zero_point = output_detail['quantization']\n",
+ " return (np.array(interpreter.get_tensor(output_detail['index']), dtype=np.float32) - output_zero_point) * output_scale\n",
+ "\n",
+ "def run_tflite(interpreter, input_word_ids, input_mask, input_type_ids):\n",
+ " input_word_ids_index, input_mask_index, input_type_ids_index = [\n",
+ " detail['index'] for detail in interpreter.get_input_details()]\n",
+ " interpreter.set_tensor(input_word_ids_index, input_word_ids)\n",
+ " interpreter.set_tensor(input_mask_index, input_mask)\n",
+ " interpreter.set_tensor(input_type_ids_index, input_type_ids)\n",
+ " interpreter.invoke()\n",
+ "\n",
+ " start_logits_detail, end_logits_detail = interpreter.get_output_details()\n",
+ "\n",
+ " return get_dequantized_tensor(interpreter, start_logits_detail), get_dequantized_tensor(interpreter, end_logits_detail)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "cnN9RLht1kOA"
+ },
+ "outputs": [],
+ "source": [
+ "tflite_file = 'model_qat.tflite'\n",
+ "with open(tflite_file, 'rb') as fp:\n",
+ " tflite_model = fp.read()\n",
+ "\n",
+ "interpreter = tf.lite.Interpreter(\n",
+ " model_content=tflite_model,\n",
+ " experimental_preserve_all_tensors=True)\n",
+ "interpreter.allocate_tensors()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "qzTgVDw51vGY",
+ "outputId": "9e75bd54-6d2f-4fb6-c462-fe7fac854a27"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "(1, 384)\n",
+ "(1, 384)\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Dummy inputs\n",
+ "input_type_ids = np.zeros(shape=[1, 384], dtype=np.int32)\n",
+ "input_word_ids = np.zeros(shape=[1, 384], dtype=np.int32)\n",
+ "input_mask = np.zeros(shape=[1, 384], dtype=np.int32)\n",
+ "\n",
+ "start_logits, end_logits = run_tflite(interpreter, input_type_ids, input_word_ids, input_mask)\n",
+ "\n",
+ "print(start_logits.shape)\n",
+ "print(end_logits.shape)"
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "provenance": [],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/qat/nlp/modeling/__init__.py b/official/projects/qat/nlp/modeling/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/nlp/modeling/layers/__init__.py b/official/projects/qat/nlp/modeling/layers/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/layers/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/nlp/modeling/layers/mobile_bert_layers.py b/official/projects/qat/nlp/modeling/layers/mobile_bert_layers.py
new file mode 100644
index 00000000000..608a4ccd8f9
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/layers/mobile_bert_layers.py
@@ -0,0 +1,485 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""MobileBERT embedding and transformer layers."""
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.nlp import modeling
+from official.projects.qat.nlp.modeling.layers.multi_head_attention import MultiHeadAttentionQuantized
+from official.projects.qat.nlp.quantization import configs
+from official.projects.qat.nlp.quantization import helper
+from official.projects.qat.nlp.quantization import wrappers
+
+
+def _quantized_multi_head_attention(*args, **kwargs):
+ layer = MultiHeadAttentionQuantized(*args, **kwargs)
+ return wrappers.MultiHeadAttentionQuantizeWrapper(
+ layer, configs.DefaultMultiHeadAttentionQuantizeConfig())
+
+
+def _quantized_einsum_dense(*args, **kwargs):
+ layer = tf_keras.layers.EinsumDense(*args, **kwargs)
+ return tfmot.quantization.keras.QuantizeWrapperV2(
+ layer, configs.DefaultEinsumDenseQuantizeConfig())
+
+
+def _output_quantize(layer):
+ return tfmot.quantization.keras.QuantizeWrapperV2(
+ layer, configs.Default8BitOutputQuantizeConfig())
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class NoNormQuantized(tf_keras.layers.Layer):
+ """Apply element-wise linear transformation to the last dimension."""
+
+ def __init__(self, name=None):
+ super().__init__(name=name)
+
+ def build(self, shape):
+ kernal_size = shape[-1]
+ self.bias = self.add_weight('beta',
+ shape=[kernal_size],
+ initializer='zeros')
+ self.scale = self.add_weight('gamma',
+ shape=[kernal_size],
+ initializer='ones')
+ self.multiply = _output_quantize(
+ tf_keras.layers.Multiply())
+
+ def call(self, feature):
+ broadcast_shape = tf.shape(feature)
+ scale = tf.broadcast_to(self.scale, broadcast_shape)
+ output = self.multiply([feature, scale])
+ return output + self.bias
+
+
+def _get_norm_layer(normalization_type='no_norm', name=None):
+ """Get normlization layer.
+
+ Args:
+ normalization_type: String. The type of normalization_type, only `no_norm`
+ and `layer_norm` are supported.
+ name: Name for the norm layer.
+
+ Returns:
+ layer norm class.
+ """
+ if normalization_type == 'no_norm':
+ layer = NoNormQuantized(name=name)
+ elif normalization_type == 'layer_norm':
+ layer = tf_keras.layers.LayerNormalization(
+ name=name,
+ axis=-1,
+ epsilon=1e-12,
+ dtype=tf.float32)
+ else:
+ raise NotImplementedError('Only "no_norm" and "layer_norm" are supported.')
+ return layer
+
+
+class MobileBertEmbeddingQuantized(helper.LayerQuantizerHelper,
+ tf_keras.layers.Layer):
+ """Performs an embedding lookup for MobileBERT.
+
+ This layer includes word embedding, token type embedding, position embedding.
+ """
+
+ def __init__(self,
+ word_vocab_size,
+ word_embed_size,
+ type_vocab_size,
+ output_embed_size,
+ max_sequence_length=512,
+ normalization_type='no_norm',
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
+ dropout_rate=0.1,
+ **kwargs):
+ """Class initialization.
+
+ Args:
+ word_vocab_size: Number of words in the vocabulary.
+ word_embed_size: Word embedding size.
+ type_vocab_size: Number of word types.
+ output_embed_size: Embedding size for the final embedding output.
+ max_sequence_length: Maximum length of input sequence.
+ normalization_type: String. The type of normalization_type, only `no_norm`
+ and `layer_norm` are supported.
+ initializer: The initializer to use for the embedding weights and linear
+ projection weights.
+ dropout_rate: Dropout rate.
+ **kwargs: keyword arguments.
+ """
+ super().__init__(**kwargs)
+ self.word_vocab_size = word_vocab_size
+ self.word_embed_size = word_embed_size
+ self.type_vocab_size = type_vocab_size
+ self.output_embed_size = output_embed_size
+ self.max_sequence_length = max_sequence_length
+ self.normalization_type = normalization_type
+ self.initializer = tf_keras.initializers.get(initializer)
+ self.dropout_rate = dropout_rate
+
+ self.word_embedding = modeling.layers.OnDeviceEmbedding(
+ self.word_vocab_size,
+ self.word_embed_size,
+ initializer=initializer,
+ name='word_embedding')
+ self.type_embedding = modeling.layers.OnDeviceEmbedding(
+ self.type_vocab_size,
+ self.output_embed_size,
+ initializer=initializer,
+ name='type_embedding')
+ self.pos_embedding = modeling.layers.PositionEmbedding(
+ max_length=max_sequence_length,
+ initializer=initializer,
+ name='position_embedding')
+ self.word_embedding_proj = _quantized_einsum_dense(
+ 'abc,cd->abd',
+ output_shape=[None, self.output_embed_size],
+ kernel_initializer=initializer,
+ bias_axes='d',
+ name='embedding_projection')
+ self.embedding_out_add_pos = _output_quantize(tf_keras.layers.Add())
+ self.layer_norm = _output_quantize(
+ _get_norm_layer(normalization_type, 'embedding_norm'))
+ self.dropout_layer = tf_keras.layers.Dropout(
+ self.dropout_rate,
+ name='embedding_dropout')
+ self.embedding_out_add_type = _output_quantize(tf_keras.layers.Add())
+
+ def build(self, input_shape):
+ self._add_quantizer('word_embedding_out')
+ self._add_quantizer('pos_embedding_out')
+ self._add_quantizer('type_embedding_out')
+
+ self._build_quantizer_vars()
+
+ def get_config(self):
+ config = {
+ 'word_vocab_size': self.word_vocab_size,
+ 'word_embed_size': self.word_embed_size,
+ 'type_vocab_size': self.type_vocab_size,
+ 'output_embed_size': self.output_embed_size,
+ 'max_sequence_length': self.max_sequence_length,
+ 'normalization_type': self.normalization_type,
+ 'initializer': tf_keras.initializers.serialize(self.initializer),
+ 'dropout_rate': self.dropout_rate
+ }
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(config.items()))
+
+ def call(self, input_ids, token_type_ids=None, training=None):
+ word_embedding_out = self.word_embedding(input_ids)
+ word_embedding_out = self._apply_quantizer(
+ 'word_embedding_out', word_embedding_out, training)
+ word_embedding_out = tf.concat(
+ [tf.pad(word_embedding_out[:, 1:], ((0, 0), (0, 1), (0, 0))),
+ word_embedding_out,
+ tf.pad(word_embedding_out[:, :-1], ((0, 0), (1, 0), (0, 0)))],
+ axis=2)
+ word_embedding_out = self.word_embedding_proj(word_embedding_out)
+
+ pos_embedding_out = self.pos_embedding(word_embedding_out)
+ pos_embedding_out = self._apply_quantizer(
+ 'pos_embedding_out', pos_embedding_out, training)
+ embedding_out = self.embedding_out_add_pos([
+ word_embedding_out, pos_embedding_out])
+ if token_type_ids is not None:
+ type_embedding_out = self.type_embedding(token_type_ids)
+ type_embedding_out = self._apply_quantizer(
+ 'type_embedding_out', type_embedding_out, training)
+ embedding_out = self.embedding_out_add_type([
+ embedding_out, type_embedding_out])
+ embedding_out = self.layer_norm(embedding_out)
+ embedding_out = self.dropout_layer(embedding_out)
+
+ return embedding_out
+
+
+class MobileBertTransformerQuantized(tf_keras.layers.Layer):
+ """Transformer block for MobileBERT.
+
+ An implementation of one layer (block) of Transformer with bottleneck and
+ inverted-bottleneck for MobilerBERT.
+
+ Original paper for MobileBERT:
+ https://arxiv.org/pdf/2004.02984.pdf
+ """
+
+ def __init__(self,
+ hidden_size=512,
+ num_attention_heads=4,
+ intermediate_size=512,
+ intermediate_act_fn='relu',
+ hidden_dropout_prob=0.1,
+ attention_probs_dropout_prob=0.1,
+ intra_bottleneck_size=128,
+ use_bottleneck_attention=False,
+ key_query_shared_bottleneck=True,
+ num_feedforward_networks=4,
+ normalization_type='no_norm',
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
+ **kwargs):
+ """Class initialization.
+
+ Args:
+ hidden_size: Hidden size for the Transformer input and output tensor.
+ num_attention_heads: Number of attention heads in the Transformer.
+ intermediate_size: The size of the "intermediate" (a.k.a., feed forward)
+ layer.
+ intermediate_act_fn: The non-linear activation function to apply to the
+ output of the intermediate/feed-forward layer.
+ hidden_dropout_prob: Dropout probability for the hidden layers.
+ attention_probs_dropout_prob: Dropout probability of the attention
+ probabilities.
+ intra_bottleneck_size: Size of bottleneck.
+ use_bottleneck_attention: Use attention inputs from the bottleneck
+ transformation. If true, the following `key_query_shared_bottleneck`
+ will be ignored.
+ key_query_shared_bottleneck: Whether to share linear transformation for
+ keys and queries.
+ num_feedforward_networks: Number of stacked feed-forward networks.
+ normalization_type: The type of normalization_type, only `no_norm` and
+ `layer_norm` are supported. `no_norm` represents the element-wise linear
+ transformation for the student model, as suggested by the original
+ MobileBERT paper. `layer_norm` is used for the teacher model.
+ initializer: The initializer to use for the embedding weights and linear
+ projection weights.
+ **kwargs: keyword arguments.
+
+ Raises:
+ ValueError: A Tensor shape or parameter is invalid.
+ """
+ super().__init__(**kwargs)
+ self.hidden_size = hidden_size
+ self.num_attention_heads = num_attention_heads
+ self.intermediate_size = intermediate_size
+ self.intermediate_act_fn = intermediate_act_fn
+ self.hidden_dropout_prob = hidden_dropout_prob
+ self.attention_probs_dropout_prob = attention_probs_dropout_prob
+ self.intra_bottleneck_size = intra_bottleneck_size
+ self.use_bottleneck_attention = use_bottleneck_attention
+ self.key_query_shared_bottleneck = key_query_shared_bottleneck
+ self.num_feedforward_networks = num_feedforward_networks
+ self.normalization_type = normalization_type
+ self.initializer = tf_keras.initializers.get(initializer)
+
+ if intra_bottleneck_size % num_attention_heads != 0:
+ raise ValueError(
+ (f'The bottleneck size {intra_bottleneck_size} is not a multiple '
+ f'of the number of attention heads {num_attention_heads}.'))
+ attention_head_size = int(intra_bottleneck_size / num_attention_heads)
+
+ self.block_layers = {}
+ # add input bottleneck
+ dense_layer_2d = _quantized_einsum_dense(
+ 'abc,cd->abd',
+ output_shape=[None, self.intra_bottleneck_size],
+ bias_axes='d',
+ kernel_initializer=initializer,
+ name='bottleneck_input/dense')
+ layer_norm = _output_quantize(
+ _get_norm_layer(self.normalization_type,
+ name='bottleneck_input/norm'))
+ self.block_layers['bottleneck_input'] = [dense_layer_2d,
+ layer_norm]
+
+ if self.key_query_shared_bottleneck:
+ dense_layer_2d = _quantized_einsum_dense(
+ 'abc,cd->abd',
+ output_shape=[None, self.intra_bottleneck_size],
+ bias_axes='d',
+ kernel_initializer=initializer,
+ name='kq_shared_bottleneck/dense')
+ layer_norm = _output_quantize(
+ _get_norm_layer(self.normalization_type,
+ name='kq_shared_bottleneck/norm'))
+ self.block_layers['kq_shared_bottleneck'] = [dense_layer_2d,
+ layer_norm]
+
+ # add attention layer
+ attention_layer = _quantized_multi_head_attention(
+ num_heads=self.num_attention_heads,
+ key_dim=attention_head_size,
+ value_dim=attention_head_size,
+ dropout=self.attention_probs_dropout_prob,
+ output_shape=self.intra_bottleneck_size,
+ kernel_initializer=initializer,
+ name='attention')
+ layer_norm = _output_quantize(
+ _get_norm_layer(self.normalization_type,
+ name='attention/norm'))
+ self.block_layers['attention'] = [attention_layer,
+ layer_norm]
+
+ # add stacked feed-forward networks (ffn)
+ self.block_layers['ffn'] = []
+ self.ffn_add_layers = []
+ for ffn_layer_idx in range(self.num_feedforward_networks):
+ layer_prefix = f'ffn_layer_{ffn_layer_idx}'
+ layer_name = layer_prefix + '/intermediate_dense'
+ intermediate_layer = _quantized_einsum_dense(
+ 'abc,cd->abd',
+ activation=self.intermediate_act_fn,
+ output_shape=[None, self.intermediate_size],
+ bias_axes='d',
+ kernel_initializer=initializer,
+ name=layer_name)
+ layer_name = layer_prefix + '/output_dense'
+ output_layer = _quantized_einsum_dense(
+ 'abc,cd->abd',
+ output_shape=[None, self.intra_bottleneck_size],
+ bias_axes='d',
+ kernel_initializer=initializer,
+ name=layer_name)
+ layer_name = layer_prefix + '/norm'
+ layer_norm = _output_quantize(
+ _get_norm_layer(self.normalization_type,
+ name=layer_name))
+ self.block_layers['ffn'].append([intermediate_layer,
+ output_layer,
+ layer_norm])
+ self.ffn_add_layers.append(_output_quantize(
+ tf_keras.layers.Add()))
+
+ # add output bottleneck
+ bottleneck = _quantized_einsum_dense(
+ 'abc,cd->abd',
+ output_shape=[None, self.hidden_size],
+ activation=None,
+ bias_axes='d',
+ kernel_initializer=initializer,
+ name='bottleneck_output/dense')
+ dropout_layer = tf_keras.layers.Dropout(
+ self.hidden_dropout_prob,
+ name='bottleneck_output/dropout')
+ layer_norm = _output_quantize(
+ _get_norm_layer(self.normalization_type,
+ name='bottleneck_output/norm'))
+ self.block_layers['bottleneck_output'] = [bottleneck,
+ dropout_layer,
+ layer_norm]
+ self.attention_output_add = _output_quantize(
+ tf_keras.layers.Add())
+ self.output_add = _output_quantize(
+ tf_keras.layers.Add())
+
+ def get_config(self):
+ config = {
+ 'hidden_size': self.hidden_size,
+ 'num_attention_heads': self.num_attention_heads,
+ 'intermediate_size': self.intermediate_size,
+ 'intermediate_act_fn': self.intermediate_act_fn,
+ 'hidden_dropout_prob': self.hidden_dropout_prob,
+ 'attention_probs_dropout_prob': self.attention_probs_dropout_prob,
+ 'intra_bottleneck_size': self.intra_bottleneck_size,
+ 'use_bottleneck_attention': self.use_bottleneck_attention,
+ 'key_query_shared_bottleneck': self.key_query_shared_bottleneck,
+ 'num_feedforward_networks': self.num_feedforward_networks,
+ 'normalization_type': self.normalization_type,
+ 'initializer': tf_keras.initializers.serialize(self.initializer),
+ }
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(config.items()))
+
+ def call(self,
+ input_tensor,
+ attention_mask=None,
+ return_attention_scores=False):
+ """Implementes the forward pass.
+
+ Args:
+ input_tensor: Float tensor of shape `(batch_size, seq_length,
+ hidden_size)`.
+ attention_mask: (optional) int32 tensor of shape `(batch_size, seq_length,
+ seq_length)`, with 1 for positions that can be attended to and 0 in
+ positions that should not be.
+ return_attention_scores: If return attention score.
+
+ Returns:
+ layer_output: Float tensor of shape
+ `(batch_size, seq_length, hidden_size)`.
+ attention_scores (Optional): Only when return_attention_scores is True.
+
+ Raises:
+ ValueError: A Tensor shape or parameter is invalid.
+ """
+ input_width = input_tensor.shape.as_list()[-1]
+ if input_width != self.hidden_size:
+ raise ValueError(
+ (f'The width of the input tensor {input_width} != '
+ f'hidden size {self.hidden_size}'))
+
+ prev_output = input_tensor
+ # input bottleneck
+ dense_layer = self.block_layers['bottleneck_input'][0]
+ layer_norm = self.block_layers['bottleneck_input'][1]
+ layer_input = dense_layer(prev_output)
+ layer_input = layer_norm(layer_input)
+
+ if self.use_bottleneck_attention:
+ key_tensor = layer_input
+ query_tensor = layer_input
+ value_tensor = layer_input
+ elif self.key_query_shared_bottleneck:
+ dense_layer = self.block_layers['kq_shared_bottleneck'][0]
+ layer_norm = self.block_layers['kq_shared_bottleneck'][1]
+ shared_attention_input = dense_layer(prev_output)
+ shared_attention_input = layer_norm(shared_attention_input)
+ key_tensor = shared_attention_input
+ query_tensor = shared_attention_input
+ value_tensor = prev_output
+ else:
+ key_tensor = prev_output
+ query_tensor = prev_output
+ value_tensor = prev_output
+
+ # attention layer
+ attention_layer = self.block_layers['attention'][0]
+ layer_norm = self.block_layers['attention'][1]
+ attention_output, attention_scores = attention_layer(
+ query_tensor,
+ value_tensor,
+ key_tensor,
+ attention_mask,
+ return_attention_scores=True,
+ )
+ attention_output = layer_norm(
+ self.attention_output_add([attention_output, layer_input]))
+
+ # stacked feed-forward networks
+ layer_input = attention_output
+ for ffn_idx in range(self.num_feedforward_networks):
+ intermediate_layer = self.block_layers['ffn'][ffn_idx][0]
+ output_layer = self.block_layers['ffn'][ffn_idx][1]
+ layer_norm = self.block_layers['ffn'][ffn_idx][2]
+ intermediate_output = intermediate_layer(layer_input)
+ layer_output = output_layer(intermediate_output)
+ layer_output = layer_norm(
+ self.ffn_add_layers[ffn_idx]([layer_output, layer_input]))
+ layer_input = layer_output
+
+ # output bottleneck
+ bottleneck = self.block_layers['bottleneck_output'][0]
+ dropout_layer = self.block_layers['bottleneck_output'][1]
+ layer_norm = self.block_layers['bottleneck_output'][2]
+ layer_output = bottleneck(layer_output) # pyrefly: ignore[unbound-name]
+ layer_output = dropout_layer(layer_output)
+ layer_output = layer_norm(self.output_add([layer_output, prev_output]))
+
+ if return_attention_scores:
+ return layer_output, attention_scores
+ else:
+ return layer_output
diff --git a/official/projects/qat/nlp/modeling/layers/multi_head_attention.py b/official/projects/qat/nlp/modeling/layers/multi_head_attention.py
new file mode 100644
index 00000000000..b3fa42020bf
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/layers/multi_head_attention.py
@@ -0,0 +1,169 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Quantized multi head attention layer."""
+import math
+
+import tensorflow as tf, tf_keras
+
+from tensorflow.python.ops import array_ops
+from tensorflow.python.ops import math_ops
+from tensorflow.python.ops import special_math_ops
+from official.projects.qat.nlp.quantization import helper
+
+
+# -6 for mask adder before softmax on int8 model. (e^-6 < 1/256)
+_MASK_CONSTANT_FOR_INT8_QUANTIZATION = 6
+
+
+class MultiHeadAttentionQuantized(helper.LayerQuantizerHelper,
+ tf_keras.layers.MultiHeadAttention):
+ """Quantized multi head attention layer.
+
+ This layer only quantized _compute_attention part. EinsumDense child layers
+ should be quantized from the QuantizeConfig.
+ """
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self._compute_attention_first_call = True
+
+ def _build_from_signature(self, *args, **kwargs):
+ super()._build_from_signature( # pytype: disable=attribute-error # typed-keras
+ *args, **kwargs)
+ self._add_quantizer('query')
+ self._add_quantizer('attention_scores')
+ self._add_quantizer('attention_output')
+
+ self._add_quantizer('masked_softmax_attention_mask',
+ all_value_quantizer=True)
+ self._add_quantizer('masked_softmax_sub1')
+ self._add_quantizer('masked_softmax_mask1')
+ self._add_quantizer('masked_softmax_sub2')
+ self._add_quantizer('masked_softmax_clamp', all_value_quantizer=True)
+ self._add_quantizer('masked_softmax_mask2', all_value_quantizer=True)
+ self._add_quantizer('masked_softmax_adder_sub', all_value_quantizer=True)
+ self._add_quantizer('masked_softmax_adder_mul', all_value_quantizer=True)
+ self._add_quantizer('masked_softmax_add', all_value_quantizer=True)
+
+ def _masked_softmax(
+ self, attention_scores, attention_mask=None, training=None):
+ """Normalize the attention scores to probabilities."""
+ # `attention_scores` = [B, N, T, S]
+ if attention_mask is None:
+ return self._softmax(attention_scores)
+
+ # The expand dim happens starting from the `num_heads` dimension,
+ # (, num_heads, )
+ mask_expansion_axes = [-len(self._attention_axes) * 2 - 1]
+ for _ in range(len(attention_scores.shape) - len(attention_mask.shape)):
+ attention_mask = array_ops.expand_dims(
+ attention_mask, axis=mask_expansion_axes)
+ if attention_scores.dtype != attention_mask.dtype:
+ attention_mask = tf.cast(attention_mask, attention_scores.dtype)
+ attention_mask = self._apply_quantizer(
+ 'masked_softmax_attention_mask', attention_mask, training)
+
+ # Makes attention_scores >= 0 to avoid masked maximum value be 0.
+ attention_scores -= math_ops.reduce_min(
+ attention_scores, axis=-1, keepdims=True)
+ attention_scores = self._apply_quantizer(
+ 'masked_softmax_sub1', attention_scores, training)
+ attention_scores *= attention_mask
+ attention_scores = self._apply_quantizer(
+ 'masked_softmax_mask1', attention_scores, training)
+
+ # Makes attention_scores <= 0, and become max value be 0.
+ attention_scores -= math_ops.reduce_max(
+ attention_scores, axis=-1, keepdims=True)
+ attention_scores = self._apply_quantizer(
+ 'masked_softmax_sub2', attention_scores, training)
+
+ # Clip the range of values [-6, 0].
+ attention_scores = tf.clip_by_value(
+ attention_scores, clip_value_min=-6, clip_value_max=0)
+ attention_scores = self._apply_quantizer(
+ 'masked_softmax_clamp', attention_scores, training)
+ # We basically hard-code the to-be-masked-out part have -6.
+ # Maximum number is 0. It"s reasonable for 8 bit quantization because
+ # e^(0) / e^(-6) < 1/256
+ attention_scores *= attention_mask
+ attention_scores = self._apply_quantizer(
+ 'masked_softmax_mask2', attention_scores, training)
+ adder = attention_mask - 1.0
+ adder = self._apply_quantizer('masked_softmax_adder_sub', adder, training)
+ adder *= _MASK_CONSTANT_FOR_INT8_QUANTIZATION
+ adder = self._apply_quantizer('masked_softmax_adder_mul', adder, training)
+ attention_scores += adder
+ attention_scores = self._apply_quantizer(
+ 'masked_softmax_add', attention_scores, training)
+
+ return self._softmax(attention_scores)
+
+ def _compute_attention(self,
+ query,
+ key,
+ value,
+ attention_mask=None,
+ training=None):
+ """Applies Dot-product attention with query, key, value tensors.
+
+ This function defines the computation inside `call` with projected
+ multi-head Q, K, V inputs. Users can override this function for customized
+ attention implementation.
+
+ Args:
+ query: Projected query `Tensor` of shape `[B, T, N, key_dim]`.
+ key: Projected key `Tensor` of shape `[B, T, N, key_dim]`.
+ value: Projected value `Tensor` of shape `[B, T, N, value_dim]`.
+ attention_mask: a boolean mask of shape `[B, T, S]`, that prevents
+ attention to certain positions.
+ training: Python boolean indicating whether the layer should behave in
+ training mode (adding dropout) or in inference mode (doing nothing).
+
+ Returns:
+ attention_output: Multi-headed outputs of attention computation.
+ attention_scores: Multi-headed attention weights.
+ """
+ if self._compute_attention_first_call:
+ self._build_quantizer_vars()
+ # Note: Applying scalar multiply at the smaller end of einsum improves
+ # XLA performance, but may introduce slight numeric differences in
+ # the Transformer attention head.
+ query = math_ops.multiply(query, 1.0 / math.sqrt(float(self._key_dim)))
+
+ query = self._apply_quantizer('query', query, training)
+ # Take the dot product between "query" and "key" to get the raw
+ # attention scores.
+ attention_scores = special_math_ops.einsum(self._dot_product_equation, key,
+ query)
+ attention_scores = self._apply_quantizer(
+ 'attention_scores', attention_scores, training)
+
+ attention_scores = self._masked_softmax(
+ attention_scores, attention_mask, training)
+
+ # This is actually dropping out entire tokens to attend to, which might
+ # seem a bit unusual, but is taken from the original Transformer paper.
+ attention_scores_dropout = self._dropout_layer(
+ attention_scores, training=training)
+
+ # `context_layer` = [B, T, N, H]
+ attention_output = special_math_ops.einsum(self._combine_equation,
+ attention_scores_dropout, value)
+ attention_output = self._apply_quantizer(
+ 'attention_output', attention_output, training)
+
+ self._compute_attention_first_call = False
+ return attention_output, attention_scores
diff --git a/official/projects/qat/nlp/modeling/layers/transformer_encoder_block.py b/official/projects/qat/nlp/modeling/layers/transformer_encoder_block.py
new file mode 100644
index 00000000000..69aa7520cd4
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/layers/transformer_encoder_block.py
@@ -0,0 +1,344 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras-based TransformerEncoder block layer."""
+from typing import Optional
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.projects.qat.nlp.modeling.layers.multi_head_attention import MultiHeadAttentionQuantized
+from official.projects.qat.nlp.quantization import configs
+from official.projects.qat.nlp.quantization import wrappers
+
+
+def _quantized_multi_head_attention(*args, **kwargs):
+ layer = MultiHeadAttentionQuantized(*args, **kwargs)
+ return wrappers.MultiHeadAttentionQuantizeWrapper(
+ layer, configs.DefaultMultiHeadAttentionQuantizeConfig())
+
+
+def _quantized_einsum_dense(*args, **kwargs):
+ layer = tf_keras.layers.EinsumDense(*args, **kwargs)
+ return tfmot.quantization.keras.QuantizeWrapperV2(
+ layer, configs.DefaultEinsumDenseQuantizeConfig())
+
+
+def _output_quantize(layer):
+ return tfmot.quantization.keras.QuantizeWrapperV2(
+ layer, configs.Default8BitOutputQuantizeConfig())
+
+
+class TransformerEncoderBlockQuantized(tf_keras.layers.Layer):
+ """TransformerEncoderBlock layer.
+
+ This layer implements the Transformer Encoder from
+ "Attention Is All You Need". (https://arxiv.org/abs/1706.03762),
+ which combines a `tf_keras.layers.MultiHeadAttention` layer with a
+ two-layer feedforward network.
+
+ References:
+ [Attention Is All You Need](https://arxiv.org/abs/1706.03762)
+ [BERT: Pre-training of Deep Bidirectional Transformers for Language
+ Understanding](https://arxiv.org/abs/1810.04805)
+ """
+
+ def __init__(self,
+ num_attention_heads,
+ inner_dim,
+ inner_activation,
+ output_range=None,
+ kernel_initializer="glorot_uniform",
+ bias_initializer="zeros",
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ activity_regularizer=None,
+ kernel_constraint=None,
+ bias_constraint=None,
+ use_bias=True,
+ norm_first=False,
+ norm_epsilon=1e-12,
+ output_dropout=0.0,
+ attention_dropout=0.0,
+ inner_dropout=0.0,
+ attention_initializer=None,
+ attention_axes=None,
+ **kwargs):
+ """Initializes `TransformerEncoderBlock`.
+
+ Args:
+ num_attention_heads: Number of attention heads.
+ inner_dim: The output dimension of the first Dense layer in a two-layer
+ feedforward network.
+ inner_activation: The activation for the first Dense layer in a two-layer
+ feedforward network.
+ output_range: the sequence output range, [0, output_range) for slicing the
+ target sequence. `None` means the target sequence is not sliced.
+ kernel_initializer: Initializer for dense layer kernels.
+ bias_initializer: Initializer for dense layer biases.
+ kernel_regularizer: Regularizer for dense layer kernels.
+ bias_regularizer: Regularizer for dense layer biases.
+ activity_regularizer: Regularizer for dense layer activity.
+ kernel_constraint: Constraint for dense layer kernels.
+ bias_constraint: Constraint for dense layer kernels.
+ use_bias: Whether to enable use_bias in attention layer. If set False,
+ use_bias in attention layer is disabled.
+ norm_first: Whether to normalize inputs to attention and intermediate
+ dense layers. If set False, output of attention and intermediate dense
+ layers is normalized.
+ norm_epsilon: Epsilon value to initialize normalization layers.
+ output_dropout: Dropout probability for the post-attention and output
+ dropout.
+ attention_dropout: Dropout probability for within the attention layer.
+ inner_dropout: Dropout probability for the first Dense layer in a
+ two-layer feedforward network.
+ attention_initializer: Initializer for kernels of attention layers. If set
+ `None`, attention layers use kernel_initializer as initializer for
+ kernel.
+ attention_axes: axes over which the attention is applied. `None` means
+ attention over all axes, but batch, heads, and features.
+ **kwargs: keyword arguments/
+ """
+ super().__init__(**kwargs)
+
+ if output_range is not None:
+ logging.warning("`output_range` is available as an argument for `call()`."
+ "The `output_range` as __init__ argument is deprecated.")
+
+ self._num_heads = num_attention_heads
+ self._inner_dim = inner_dim
+ self._inner_activation = inner_activation
+ self._attention_dropout = attention_dropout
+ self._attention_dropout_rate = attention_dropout
+ self._output_dropout = output_dropout
+ self._output_dropout_rate = output_dropout
+ self._output_range = output_range
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
+ self._use_bias = use_bias
+ self._norm_first = norm_first
+ self._norm_epsilon = norm_epsilon
+ self._inner_dropout = inner_dropout
+ if attention_initializer:
+ self._attention_initializer = tf_keras.initializers.get(
+ attention_initializer)
+ else:
+ self._attention_initializer = self._kernel_initializer
+ self._attention_axes = attention_axes
+
+ def build(self, input_shape):
+ if isinstance(input_shape, tf.TensorShape):
+ input_tensor_shape = input_shape
+ elif isinstance(input_shape, (list, tuple)):
+ input_tensor_shape = tf.TensorShape(input_shape[0])
+ else:
+ raise ValueError(
+ "The type of input shape argument is not supported, got: %s" %
+ type(input_shape))
+ if len(input_tensor_shape.as_list()) != 3:
+ raise ValueError("TransformerEncoderBlock expects a three-dimensional "
+ "input of shape [batch, sequence, width].")
+ hidden_size = input_tensor_shape[-1]
+ if hidden_size % self._num_heads != 0:
+ raise ValueError(
+ "The input size (%d) is not a multiple of the number of attention "
+ "heads (%d)" % (hidden_size, self._num_heads))
+ self._attention_head_size = int(hidden_size // self._num_heads)
+ common_kwargs = dict(
+ bias_initializer=self._bias_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ activity_regularizer=self._activity_regularizer,
+ kernel_constraint=self._kernel_constraint,
+ bias_constraint=self._bias_constraint)
+ self._attention_layer = _quantized_multi_head_attention(
+ num_heads=self._num_heads,
+ key_dim=self._attention_head_size,
+ dropout=self._attention_dropout,
+ use_bias=self._use_bias,
+ kernel_initializer=self._attention_initializer,
+ attention_axes=self._attention_axes,
+ name="self_attention",
+ **common_kwargs)
+ self._attention_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
+ # Use float32 in layernorm for numeric stability.
+ # It is probably safe in mixed_float16, but we haven't validated this yet.
+ self._attention_layer_norm = _output_quantize(
+ tf_keras.layers.LayerNormalization(
+ name="self_attention_layer_norm",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32))
+ self._intermediate_dense = _quantized_einsum_dense(
+ "abc,cd->abd",
+ output_shape=(None, self._inner_dim),
+ bias_axes="d",
+ kernel_initializer=self._kernel_initializer,
+ name="intermediate",
+ **common_kwargs)
+ policy = tf_keras.mixed_precision.global_policy()
+ if policy.name == "mixed_bfloat16":
+ # bfloat16 causes BERT with the LAMB optimizer to not converge
+ # as well, so we use float32.
+ # TODO(b/154538392): Investigate this.
+ policy = tf.float32
+ self._intermediate_activation_layer = _output_quantize(
+ tf_keras.layers.Activation(
+ self._inner_activation, dtype=policy))
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
+ rate=self._inner_dropout)
+ self._output_dense = _quantized_einsum_dense(
+ "abc,cd->abd",
+ output_shape=(None, hidden_size),
+ bias_axes="d",
+ name="output",
+ kernel_initializer=self._kernel_initializer,
+ **common_kwargs)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
+ # Use float32 in layernorm for numeric stability.
+ self._output_layer_norm = _output_quantize(
+ tf_keras.layers.LayerNormalization(
+ name="output_layer_norm",
+ axis=-1,
+ epsilon=self._norm_epsilon,
+ dtype=tf.float32))
+ self._add = _output_quantize(tf_keras.layers.Add())
+ self._output_add = tf_keras.layers.Add()
+
+ super().build(input_shape)
+
+ def get_config(self):
+ config = {
+ "num_attention_heads":
+ self._num_heads,
+ "inner_dim":
+ self._inner_dim,
+ "inner_activation":
+ self._inner_activation,
+ "output_dropout":
+ self._output_dropout_rate,
+ "attention_dropout":
+ self._attention_dropout_rate,
+ "output_range":
+ self._output_range,
+ "kernel_initializer":
+ tf_keras.initializers.serialize(self._kernel_initializer),
+ "bias_initializer":
+ tf_keras.initializers.serialize(self._bias_initializer),
+ "kernel_regularizer":
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
+ "bias_regularizer":
+ tf_keras.regularizers.serialize(self._bias_regularizer),
+ "activity_regularizer":
+ tf_keras.regularizers.serialize(self._activity_regularizer),
+ "kernel_constraint":
+ tf_keras.constraints.serialize(self._kernel_constraint),
+ "bias_constraint":
+ tf_keras.constraints.serialize(self._bias_constraint),
+ "use_bias":
+ self._use_bias,
+ "norm_first":
+ self._norm_first,
+ "norm_epsilon":
+ self._norm_epsilon,
+ "inner_dropout":
+ self._inner_dropout,
+ "attention_initializer":
+ tf_keras.initializers.serialize(self._attention_initializer),
+ "attention_axes": self._attention_axes,
+ }
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(config.items()))
+
+ def call(self, inputs, output_range: Optional[tf.Tensor] = None):
+ """Transformer self-attention encoder block call.
+
+ Args:
+ inputs: a single tensor or a list of tensors. `input tensor` as the single
+ sequence of embeddings. [`input tensor`, `attention mask`] to have the
+ additional attention mask. [`query tensor`, `key value tensor`,
+ `attention mask`] to have separate input streams for the query, and
+ key/value to the multi-head attention.
+ output_range: the sequence output range, [0, output_range) for slicing the
+ target sequence. `None` means the target sequence is not sliced. If you
+ would like to have no change to the model training, it is better to only
+ set the `output_range` for serving.
+
+ Returns:
+ An ouput tensor with the same dimensions as input/query tensor.
+ """
+ if isinstance(inputs, (list, tuple)):
+ if len(inputs) == 2:
+ input_tensor, attention_mask = inputs
+ key_value = None
+ elif len(inputs) == 3:
+ input_tensor, key_value, attention_mask = inputs
+ else:
+ raise ValueError("Unexpected inputs to %s with length at %d" %
+ (self.__class__, len(inputs)))
+ else:
+ input_tensor, key_value, attention_mask = (inputs, None, None)
+
+ if output_range is None:
+ output_range = self._output_range
+ if output_range:
+ if self._norm_first:
+ source_tensor = input_tensor[:, 0:output_range, :]
+ input_tensor = self._attention_layer_norm(input_tensor)
+ if key_value is not None:
+ key_value = self._attention_layer_norm(key_value)
+ target_tensor = input_tensor[:, 0:output_range, :]
+ if attention_mask is not None:
+ attention_mask = attention_mask[:, 0:output_range, :]
+ else:
+ if self._norm_first:
+ source_tensor = input_tensor
+ input_tensor = self._attention_layer_norm(input_tensor)
+ if key_value is not None:
+ key_value = self._attention_layer_norm(key_value)
+ target_tensor = input_tensor
+
+ if key_value is None:
+ key_value = input_tensor
+ attention_output = self._attention_layer(
+ query=target_tensor, value=key_value, attention_mask=attention_mask)
+ attention_output = self._attention_dropout(attention_output) # pyrefly: ignore[not-callable]
+ if self._norm_first:
+ attention_output = self._add([source_tensor, attention_output]) # pyrefly: ignore[unbound-name]
+ else:
+ attention_output = self._attention_layer_norm(
+ self._add([target_tensor, attention_output]))
+ if self._norm_first:
+ source_attention_output = attention_output
+ attention_output = self._output_layer_norm(attention_output)
+ inner_output = self._intermediate_dense(attention_output)
+ inner_output = self._intermediate_activation_layer(inner_output)
+ inner_output = self._inner_dropout_layer(inner_output)
+ layer_output = self._output_dense(inner_output)
+ layer_output = self._output_dropout(layer_output) # pyrefly: ignore[not-callable]
+
+ if self._norm_first:
+ return self._output_add([source_attention_output, layer_output]) # pyrefly: ignore[unbound-name]
+
+ # During mixed precision training, layer norm output is always fp32 for now.
+ # Casts fp32 for the subsequent add.
+ layer_output = tf.cast(layer_output, tf.float32)
+ return self._output_layer_norm(
+ self._output_add([layer_output, attention_output]))
diff --git a/official/projects/qat/nlp/modeling/layers/transformer_encoder_block_test.py b/official/projects/qat/nlp/modeling/layers/transformer_encoder_block_test.py
new file mode 100644
index 00000000000..f0aa8f7fdee
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/layers/transformer_encoder_block_test.py
@@ -0,0 +1,226 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for Keras-based quantized transformer block layer."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.qat.nlp.modeling.layers.transformer_encoder_block import TransformerEncoderBlockQuantized
+
+
+@parameterized.named_parameters(
+ ('base', TransformerEncoderBlockQuantized))
+class TransformerEncoderBlockQuantizedLayerTest(
+ tf.test.TestCase, parameterized.TestCase):
+
+ def tearDown(self):
+ super(TransformerEncoderBlockQuantizedLayerTest, self).tearDown()
+ tf_keras.mixed_precision.set_global_policy('float32')
+
+ def test_layer_creation(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048, inner_activation='relu')
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+ # The default output of a transformer layer should be the same as the input.
+ self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
+
+ def test_layer_creation_with_mask(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048, inner_activation='relu')
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ # Create a 2-dimensional input (the first dimension is implicit).
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
+ output_tensor = test_layer([data_tensor, mask_tensor])
+ # The default output of a transformer layer should be the same as the input.
+ self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
+
+ def test_layer_invocation(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048, inner_activation='relu')
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output_tensor = test_layer(data_tensor)
+
+ # Create a model from the test layer.
+ model = tf_keras.Model(data_tensor, output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 6
+ input_data = 10 * np.random.random_sample(
+ (batch_size, sequence_length, width))
+ _ = model.predict(input_data)
+
+ def test_layer_invocation_with_mask(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048, inner_activation='relu')
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ # Create a 2-dimensional input (the first dimension is implicit).
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
+ output_tensor = test_layer([data_tensor, mask_tensor])
+
+ # Create a model from the test layer.
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
+
+ # Invoke the model on test data. We can't validate the output data itself
+ # (the NN is too complex) but this will rule out structural runtime errors.
+ batch_size = 6
+ input_data = 10 * np.random.random_sample(
+ (batch_size, sequence_length, width))
+ # The attention mask should be of shape (batch, from_seq_len, to_seq_len),
+ # which here is (batch, sequence_length, sequence_length)
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length, sequence_length))
+ _ = model.predict([input_data, mask_data])
+
+ def test_layer_output_range(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048, inner_activation='relu')
+ sequence_length = 21
+ width = 80
+
+ batch_size = 6
+ input_data = 10 * np.random.random_sample(
+ (batch_size, sequence_length, width))
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length, sequence_length))
+ output_tensor = test_layer([input_data, mask_data])
+
+ # The layer only attends to the first token and outputs the first token
+ # embedding.
+ new_layer = transformer_cls(
+ num_attention_heads=10,
+ inner_dim=2048,
+ inner_activation='relu')
+ _ = new_layer([input_data, mask_data], output_range=1)
+ new_layer.set_weights(test_layer.get_weights())
+ new_output_tensor = new_layer([input_data, mask_data], output_range=1)
+ self.assertAllClose(
+ new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
+
+ def test_layer_output_range_without_mask(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048,
+ inner_activation='relu', norm_first=True)
+ sequence_length = 21
+ width = 80
+
+ batch_size = 6
+ input_data = 10 * np.random.random_sample(
+ (batch_size, sequence_length, width))
+ output_tensor = test_layer(input_data)
+
+ # The layer only attends to the first token and outputs the first token
+ # embedding.
+ new_layer = transformer_cls(
+ num_attention_heads=10,
+ inner_dim=2048,
+ inner_activation='relu',
+ norm_first=True)
+ _ = new_layer(input_data, output_range=1)
+ new_layer.set_weights(test_layer.get_weights())
+ new_output_tensor = new_layer(input_data, output_range=1)
+ self.assertAllClose(
+ new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
+
+ def test_layer_output_range_with_pre_norm(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10, inner_dim=2048,
+ inner_activation='relu', norm_first=True)
+ sequence_length = 21
+ width = 80
+
+ batch_size = 6
+ input_data = 10 * np.random.random_sample(
+ (batch_size, sequence_length, width))
+ mask_data = np.random.randint(
+ 2, size=(batch_size, sequence_length, sequence_length))
+ output_tensor = test_layer([input_data, mask_data])
+
+ # The layer only attends to the first token and outputs the first token
+ # embedding.
+ new_layer = transformer_cls(
+ num_attention_heads=10,
+ inner_dim=2048,
+ inner_activation='relu',
+ norm_first=True)
+ _ = new_layer([input_data, mask_data], output_range=1)
+ new_layer.set_weights(test_layer.get_weights())
+ new_output_tensor = new_layer([input_data, mask_data], output_range=1)
+ self.assertAllClose(
+ new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
+
+ def test_transform_with_initializer(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10,
+ inner_dim=2048,
+ inner_activation='relu',
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
+ sequence_length = 21
+ width = 80
+ # Create a 3-dimensional input (the first dimension is implicit).
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
+ output = test_layer(data_tensor)
+ # The default output of a transformer layer should be the same as the input.
+ self.assertEqual(data_tensor.shape.as_list(), output.shape.as_list())
+
+ def test_dynamic_layer_sequence(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=10,
+ inner_dim=2048,
+ inner_activation='relu',
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
+ # Create a 3-dimensional input (the first dimension is implicit).
+ width = 30
+ input_tensor = tf_keras.Input(shape=(None, width))
+ output_tensor = test_layer(input_tensor)
+ model = tf_keras.Model(input_tensor, output_tensor)
+
+ input_length = 17
+ input_data = np.ones((1, input_length, width))
+ output_data = model.predict(input_data)
+
+ self.assertAllEqual([1, input_length, width], output_data.shape)
+
+ def test_separate_qkv(self, transformer_cls):
+ test_layer = transformer_cls(
+ num_attention_heads=2,
+ inner_dim=128,
+ inner_activation='relu',
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
+ # Forward path.
+ q_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
+ kv_tensor = tf.zeros([2, 8, 16], dtype=tf.float32)
+ dummy_mask = tf.zeros([2, 4, 8], dtype=tf.float32)
+ inputs = [q_tensor, kv_tensor, dummy_mask]
+ output = test_layer(inputs)
+ self.assertEqual(output.shape, q_tensor.shape)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/qat/nlp/modeling/models/__init__.py b/official/projects/qat/nlp/modeling/models/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/models/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/nlp/modeling/models/bert_span_labeler.py b/official/projects/qat/nlp/modeling/models/bert_span_labeler.py
new file mode 100644
index 00000000000..39b34cf8265
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/models/bert_span_labeler.py
@@ -0,0 +1,125 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""BERT Question Answering model."""
+# pylint: disable=g-classes-have-attributes
+import collections
+import tensorflow as tf, tf_keras
+
+from official.projects.qat.nlp.modeling.networks import span_labeling
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class BertSpanLabelerQuantized(tf_keras.Model):
+ """Span labeler model based on a BERT-style transformer-based encoder.
+
+ This is an implementation of the network structure surrounding a transformer
+ encoder as described in "BERT: Pre-training of Deep Bidirectional Transformers
+ for Language Understanding" (https://arxiv.org/abs/1810.04805).
+
+ The BertSpanLabeler allows a user to pass in a transformer encoder, and
+ instantiates a span labeling network based on a single dense layer.
+
+ *Note* that the model is constructed by
+ [Keras Functional API](https://keras.io/guides/functional_api/).
+
+ Args:
+ network: A transformer network. This network should output a sequence output
+ and a classification output. Furthermore, it should expose its embedding
+ table via a `get_embedding_table` method.
+ initializer: The initializer (if any) to use in the span labeling network.
+ Defaults to a Glorot uniform initializer.
+ output: The output style for this network. Can be either `logit`' or
+ `predictions`.
+ """
+
+ def __init__(self,
+ network,
+ initializer='glorot_uniform',
+ output='logits',
+ **kwargs):
+
+ # We want to use the inputs of the passed network as the inputs to this
+ # Model. To do this, we need to keep a handle to the network inputs for use
+ # when we construct the Model object at the end of init.
+ inputs = network.inputs
+
+ # Because we have a copy of inputs to create this Model object, we can
+ # invoke the Network object with its own input tensors to start the Model.
+ outputs = network(inputs)
+ if isinstance(outputs, list):
+ sequence_output = outputs[0]
+ else:
+ sequence_output = outputs['sequence_output']
+
+ # The input network (typically a transformer model) may get outputs from all
+ # layers. When this case happens, we retrieve the last layer output.
+ if isinstance(sequence_output, list):
+ sequence_output = sequence_output[-1]
+
+ # This is an instance variable for ease of access to the underlying task
+ # network.
+ span_labeling_quantized = span_labeling.SpanLabelingQuantized(
+ input_width=sequence_output.shape[-1],
+ initializer=initializer,
+ output=output,
+ name='span_labeling')
+ start_logits, end_logits = span_labeling_quantized(sequence_output)
+
+ # Use identity layers wrapped in lambdas to explicitly name the output
+ # tensors. This allows us to use string-keyed dicts in Keras fit/predict/
+ # evaluate calls.
+ start_logits = tf_keras.layers.Lambda(
+ tf.identity, name='start_positions')(
+ start_logits)
+ end_logits = tf_keras.layers.Lambda(
+ tf.identity, name='end_positions')(
+ end_logits)
+
+ logits = [start_logits, end_logits]
+
+ # b/164516224
+ # Once we've created the network using the Functional API, we call
+ # super().__init__ as though we were invoking the Functional API Model
+ # constructor, resulting in this object having all the properties of a model
+ # created using the Functional API. Once super().__init__ is called, we
+ # can assign attributes to `self` - note that all `self` assignments are
+ # below this line.
+ super().__init__(
+ inputs=inputs, outputs=logits, **kwargs)
+ self._network = network
+ config_dict = {
+ 'network': network,
+ 'initializer': initializer,
+ 'output': output,
+ }
+ # We are storing the config dict as a namedtuple here to ensure checkpoint
+ # compatibility with an earlier version of this model which did not track
+ # the config dict attribute. TF does not track immutable attrs which
+ # do not contain Trackables, so by creating a config namedtuple instead of
+ # a dict we avoid tracking it.
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
+ self._config = config_cls(**config_dict)
+ self.span_labeling = span_labeling_quantized
+
+ @property
+ def checkpoint_items(self):
+ return dict(encoder=self._network)
+
+ def get_config(self):
+ return dict(self._config._asdict())
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
diff --git a/official/projects/qat/nlp/modeling/networks/__init__.py b/official/projects/qat/nlp/modeling/networks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/networks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/nlp/modeling/networks/span_labeling.py b/official/projects/qat/nlp/modeling/networks/span_labeling.py
new file mode 100644
index 00000000000..f662b0afd08
--- /dev/null
+++ b/official/projects/qat/nlp/modeling/networks/span_labeling.py
@@ -0,0 +1,116 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Span labeling network."""
+# pylint: disable=g-classes-have-attributes
+import collections
+import tensorflow as tf, tf_keras
+import tensorflow_model_optimization as tfmot
+
+from official.projects.qat.nlp.quantization import configs
+
+
+def _apply_paragraph_mask(logits, paragraph_mask):
+ """Applies a position mask to calculated logits."""
+ masked_logits = logits * (paragraph_mask) - 1e30 * (1 - paragraph_mask)
+ return tf.nn.log_softmax(masked_logits, -1), masked_logits
+
+
+@tf_keras.utils.register_keras_serializable(package='Text')
+class SpanLabelingQuantized(tf_keras.Model):
+ """Span labeling network head for BERT modeling.
+
+ This network implements a simple single-span labeler based on a dense layer.
+ *Note* that the network is constructed by
+ [Keras Functional API](https://keras.io/guides/functional_api/).
+
+ Args:
+ input_width: The innermost dimension of the input tensor to this network.
+ activation: The activation, if any, for the dense layer in this network.
+ initializer: The initializer for the dense layer in this network. Defaults
+ to a Glorot uniform initializer.
+ output: The output style for this network. Can be either `logits` or
+ `predictions`.
+ """
+
+ def __init__(self,
+ input_width,
+ activation=None,
+ initializer='glorot_uniform',
+ output='logits',
+ **kwargs):
+
+ sequence_data = tf_keras.layers.Input(
+ shape=(None, input_width), name='sequence_data', dtype=tf.float32)
+
+ logits_layer = tf_keras.layers.Dense(
+ 2, # This layer predicts start location and end location.
+ activation=activation,
+ kernel_initializer=initializer,
+ name='predictions/transform/logits')
+ logits_layer = tfmot.quantization.keras.QuantizeWrapperV2(
+ logits_layer,
+ configs.Default8BitQuantizeConfig(['kernel'], ['activation'], False))
+ intermediate_logits = logits_layer(sequence_data)
+ start_logits, end_logits = self._split_output_tensor(intermediate_logits)
+
+ start_predictions = tf_keras.layers.Activation(tf.nn.log_softmax)(
+ start_logits)
+ end_predictions = tf_keras.layers.Activation(tf.nn.log_softmax)(end_logits)
+
+ if output == 'logits':
+ output_tensors = [start_logits, end_logits]
+ elif output == 'predictions':
+ output_tensors = [start_predictions, end_predictions]
+ else:
+ raise ValueError(
+ ('Unknown `output` value "%s". `output` can be either "logits" or '
+ '"predictions"') % output)
+
+ # b/164516224
+ # Once we've created the network using the Functional API, we call
+ # super().__init__ as though we were invoking the Functional API Model
+ # constructor, resulting in this object having all the properties of a model
+ # created using the Functional API. Once super().__init__ is called, we
+ # can assign attributes to `self` - note that all `self` assignments are
+ # below this line.
+ super().__init__(
+ inputs=[sequence_data], outputs=output_tensors, **kwargs)
+ config_dict = {
+ 'input_width': input_width,
+ 'activation': activation,
+ 'initializer': initializer,
+ 'output': output,
+ }
+ # We are storing the config dict as a namedtuple here to ensure checkpoint
+ # compatibility with an earlier version of this model which did not track
+ # the config dict attribute. TF does not track immutable attrs which
+ # do not contain Trackables, so by creating a config namedtuple instead of
+ # a dict we avoid tracking it.
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
+ self._config = config_cls(**config_dict)
+ self.start_logits = start_logits
+ self.end_logits = end_logits
+
+ def _split_output_tensor(self, tensor):
+ transposed_tensor = tf.transpose(tensor, [2, 0, 1])
+ return tf.unstack(transposed_tensor)
+
+ def get_config(self):
+ return dict(self._config._asdict())
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
diff --git a/official/projects/qat/nlp/pretrained_checkpoint_converter.py b/official/projects/qat/nlp/pretrained_checkpoint_converter.py
new file mode 100644
index 00000000000..d954c33fd2e
--- /dev/null
+++ b/official/projects/qat/nlp/pretrained_checkpoint_converter.py
@@ -0,0 +1,138 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""A converter for BERT pretrained checkpoint to QAT BERT checkpoint."""
+import tempfile
+
+from absl import app
+from absl import flags
+import tensorflow as tf, tf_keras
+
+from official.common import registry_imports # pylint: disable=unused-import
+from official.core import exp_factory
+from official.core import task_factory
+from official.modeling import hyperparams
+from official.projects.qat.nlp import registry_imports as qat_registry_imports # pylint: disable=unused-import
+
+FLAGS = flags.FLAGS
+
+_EXPERIMENT = flags.DEFINE_string(
+ 'experiment', default=None,
+ help='The experiment type registered for the pretrained model.')
+
+_CONFIG_FILE = flags.DEFINE_multi_string(
+ 'config_file',
+ default=None,
+ help='YAML/JSON files which specifies overrides. The override order '
+ 'follows the order of args. Note that each file '
+ 'can be used as an override template to override the default parameters '
+ 'specified in Python. If the same parameter is specified in both '
+ '`--config_file` and `--params_override`, `config_file` will be used '
+ 'first, followed by params_override.')
+
+_PARAMS_OVERRIDE = flags.DEFINE_string(
+ 'params_override',
+ default=None,
+ help='a YAML/JSON string or a YAML file which specifies additional '
+ 'overrides over the default parameters and those specified in '
+ '`--config_file`. Note that this is supposed to be used only to override '
+ 'the model parameters, but not the parameters like TPU specific flags. '
+ 'One canonical use case of `--config_file` and `--params_override` is '
+ 'users first define a template config file using `--config_file`, then '
+ 'use `--params_override` to adjust the minimal set of tuning parameters, '
+ 'for example setting up different `train_batch_size`. The final override '
+ 'order of parameters: default_model_params --> params from config_file '
+ '--> params in params_override. See also the help message of '
+ '`--config_file`.')
+
+_PRETRAINED_CHECKPOINT = flags.DEFINE_string(
+ 'pretrained_checkpoint',
+ default=None,
+ help='The path of pretrained checkpoint for the original bert model.')
+
+_EXPERIEMNT_QAT = flags.DEFINE_string(
+ 'experiment_qat', default=None,
+ help='The experiment type registered for the pretrained model.')
+
+_CONFIG_FILE_QAT = flags.DEFINE_multi_string(
+ 'config_file_qat',
+ default=None,
+ help='config_file flag for the qat model.')
+
+_PARAMS_OVERRIDE_QAT = flags.DEFINE_string(
+ 'params_override_qat',
+ default=None,
+ help='params_override flag for the qat model.')
+
+_OUTPUT_CHECKPOINT = flags.DEFINE_string(
+ 'output_checkpoint',
+ default=None,
+ help='The output checkpoint path for QAT applied BERT model.')
+
+
+def _build_model(experiment, config_file, params_override):
+ """Build the model."""
+ params = exp_factory.get_exp_config(experiment)
+ for config_file in config_file or []:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=True)
+
+ if params_override:
+ params = hyperparams.override_params_dict(
+ params, params_override, is_strict=True)
+
+ task = task_factory.get_task(params.task, logging_dir=tempfile.mkdtemp())
+ return task.build_model()
+
+
+def _set_weights_to_qat(model_from, model_to):
+ """Set pretrained weight to QAT applied model."""
+ name_to_index = {}
+ for index, weight in enumerate(model_to.weights):
+ origin_name = weight.name.replace('quant_', '').replace(
+ 'mobile_bert_embedding_1', 'mobile_bert_embedding')
+ name_to_index[origin_name] = index
+
+ model_to_weights = model_to.get_weights()
+ for weight, value in zip(model_from.weights, model_from.get_weights()):
+ index = name_to_index[weight.name]
+ model_to_weights[index] = value
+ model_to.set_weights(model_to_weights)
+
+
+def main(_):
+ model = _build_model(
+ _EXPERIMENT.value, _CONFIG_FILE.value, _PARAMS_OVERRIDE.value)
+ if _PRETRAINED_CHECKPOINT.value is not None:
+ ckpt = tf.train.Checkpoint(model=model)
+ status = ckpt.restore(_PRETRAINED_CHECKPOINT.value)
+ status.expect_partial().assert_existing_objects_matched()
+
+ model_qat = _build_model(
+ _EXPERIEMNT_QAT.value, _CONFIG_FILE_QAT.value, _PARAMS_OVERRIDE_QAT.value)
+
+ _set_weights_to_qat(model, model_qat)
+
+ if hasattr(model_qat, 'checkpoint_items'):
+ checkpoint_items = model_qat.checkpoint_items
+ else:
+ checkpoint_items = {}
+ ckpt_qat = tf.train.Checkpoint(
+ model=model_qat,
+ **checkpoint_items)
+ ckpt_qat.save(_OUTPUT_CHECKPOINT.value)
+
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/qat/nlp/quantization/__init__.py b/official/projects/qat/nlp/quantization/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/nlp/quantization/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/nlp/quantization/configs.py b/official/projects/qat/nlp/quantization/configs.py
new file mode 100644
index 00000000000..03b52f8de3e
--- /dev/null
+++ b/official/projects/qat/nlp/quantization/configs.py
@@ -0,0 +1,421 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Custom quantize configs."""
+from typing import Sequence, Callable, Tuple, Any, Dict
+
+import tensorflow as tf, tf_keras
+import tensorflow_model_optimization as tfmot
+
+
+Quantizer = tfmot.quantization.keras.quantizers.Quantizer
+Layer = tf_keras.layers.Layer
+Activation = Callable[[tf.Tensor], tf.Tensor]
+WeightAndQuantizer = Tuple[tf.Variable, Quantizer]
+ActivationAndQuantizer = Tuple[Activation, Quantizer]
+
+
+class _QuantizeHelper(object):
+ """Mixin with helper functions for quantizers."""
+
+ def _add_range_weights(self, layer, name, per_axis=False, tensor_shape=None):
+ """Add min and max vars to layer."""
+ # Added naming index to avoid duplicated.
+ if hasattr(layer, 'quantize_helper_weight_idx'):
+ layer.quantize_helper_weight_idx += 1
+ name = '{}/{}'.format(layer.quantize_helper_weight_idx, name)
+ else:
+ layer.quantize_helper_weight_idx = 0
+
+ shape = None
+ if per_axis and tensor_shape is not None:
+ shape = (tensor_shape[-1])
+
+ min_weight = layer.add_weight(
+ name + '_min',
+ initializer=tf_keras.initializers.Constant(-6.0),
+ trainable=False,
+ shape=shape)
+ max_weight = layer.add_weight(
+ name + '_max',
+ initializer=tf_keras.initializers.Constant(6.0),
+ trainable=False,
+ shape=shape)
+
+ return {'min_var': min_weight, 'max_var': max_weight}
+
+
+class LastValueQuantizer(
+ _QuantizeHelper,
+ tfmot.quantization.keras.quantizers.LastValueQuantizer):
+ pass
+
+
+class MovingAverageQuantizer(
+ _QuantizeHelper,
+ tfmot.quantization.keras.quantizers.MovingAverageQuantizer):
+ pass
+
+
+class NoQuantizer(tfmot.quantization.keras.quantizers.Quantizer):
+ """Dummy quantizer do nothing."""
+
+ def __call__(self, inputs, training, weights, **kwargs):
+ return tf.identity(inputs)
+
+ def get_config(self):
+ return {}
+
+ def build(self, tensor_shape, name, layer):
+ return {}
+
+ def __eq__(self, other):
+ if not isinstance(other, NoQuantizer):
+ return False
+
+ return True
+
+ def __ne__(self, other):
+ return not self.__eq__(other)
+
+
+class DefaultEinsumDenseQuantizeConfig(tfmot.quantization.keras.QuantizeConfig):
+ """QuantizeConfig for EinsumDense layer."""
+
+ # Configure how to quantize weights.
+ def get_weights_and_quantizers(self, layer):
+ return [(layer.kernel, LastValueQuantizer(
+ num_bits=8, symmetric=True, narrow_range=False, per_axis=False))]
+
+ # Configure how to quantize activations.
+ def get_activations_and_quantizers(self, layer):
+ return [(layer.activation, MovingAverageQuantizer(
+ num_bits=8, symmetric=False, narrow_range=False, per_axis=False))]
+
+ def set_quantize_weights(self, layer, quantize_weights):
+ # Add this line for each item returned in `get_weights_and_quantizers`
+ # , in the same order
+ layer.kernel = quantize_weights[0]
+
+ def set_quantize_activations(self, layer, quantize_activations):
+ # Add this line for each item returned in `get_activations_and_quantizers`
+ # , in the same order.
+ layer.activation = quantize_activations[0]
+
+ # Configure how to quantize outputs (may be equivalent to activations).
+ def get_output_quantizers(self, layer):
+ return []
+
+ def get_config(self):
+ return {}
+
+
+# pylint: disable=protected-access
+class DefaultMultiHeadAttentionQuantizeConfig(
+ tfmot.quantization.keras.QuantizeConfig):
+ """Default quantize config for MultiHeadAttention layer.
+
+ It only quantize child EinsumDense layers. It should be applied to
+ MultiHeadAttentionQuantized layer.
+ """
+
+ def __init__(self):
+ self.einsum_dense_config = DefaultEinsumDenseQuantizeConfig()
+ self.num_weight_per_einsum_dense = 1
+ self.num_activation_per_einsum_dense = 1
+
+ def _get_einsum_dense_layers(self, layer):
+ return [
+ layer._query_dense,
+ layer._key_dense,
+ layer._value_dense,
+ layer._output_dense]
+
+ def get_weights_and_quantizers(self, layer):
+ ret = []
+ for einsum_dense_layer in self._get_einsum_dense_layers(layer):
+ ret += self.einsum_dense_config.get_weights_and_quantizers(
+ einsum_dense_layer)
+ return ret
+
+ def get_activations_and_quantizers(self, layer):
+ ret = []
+ for einsum_dense_layer in self._get_einsum_dense_layers(layer):
+ ret += self.einsum_dense_config.get_activations_and_quantizers(
+ einsum_dense_layer)
+ return ret
+
+ def set_quantize_weights(self, layer, quantize_weights):
+ idx = 0
+ for einsum_dense_layer in self._get_einsum_dense_layers(layer):
+ self.einsum_dense_config.set_quantize_weights(
+ einsum_dense_layer,
+ quantize_weights[idx:idx+self.num_weight_per_einsum_dense])
+ idx += self.num_weight_per_einsum_dense
+
+ def set_quantize_activations(self, layer, quantize_activations):
+ idx = 0
+ for einsum_dense_layer in self._get_einsum_dense_layers(layer):
+ self.einsum_dense_config.set_quantize_activations(
+ einsum_dense_layer,
+ quantize_activations[idx:idx+self.num_activation_per_einsum_dense])
+ idx += self.num_activation_per_einsum_dense
+
+ def get_output_quantizers(self, layer):
+ return []
+
+ def get_config(self):
+ return {}
+# pylint: enable=protected-access
+
+
+class Default8BitOutputQuantizeConfig(tfmot.quantization.keras.QuantizeConfig):
+ """QuantizeConfig which only quantizes the output from a layer."""
+
+ def get_weights_and_quantizers(
+ self, layer: Layer) -> Sequence[WeightAndQuantizer]:
+ return []
+
+ def get_activations_and_quantizers(
+ self, layer: Layer) -> Sequence[ActivationAndQuantizer]:
+ return []
+
+ def set_quantize_weights(self,
+ layer: Layer,
+ quantize_weights: Sequence[tf.Tensor]):
+ pass
+
+ def set_quantize_activations(self,
+ layer: Layer,
+ quantize_activations: Sequence[Activation]):
+ pass
+
+ def get_output_quantizers(self, layer: Layer) -> Sequence[Quantizer]:
+ return [
+ MovingAverageQuantizer(
+ num_bits=8, per_axis=False, symmetric=False, narrow_range=False)
+ ]
+
+ def get_config(self) -> Dict[str, Any]:
+ return {}
+
+
+class Default8BitActivationQuantizeConfig(
+ tfmot.quantization.keras.QuantizeConfig):
+ """QuantizeConfig for keras.layers.Activation.
+
+ `keras.layers.Activation` needs a separate `QuantizeConfig` since the
+ decision to quantize depends on the specific activation type.
+ """
+
+ def _assert_activation_layer(self, layer: Layer):
+ if not isinstance(layer, tf_keras.layers.Activation):
+ raise RuntimeError(
+ 'Default8BitActivationQuantizeConfig can only be used with '
+ '`keras.layers.Activation`.')
+
+ def get_weights_and_quantizers(
+ self, layer: Layer) -> Sequence[WeightAndQuantizer]:
+ """See base class."""
+ self._assert_activation_layer(layer)
+ return []
+
+ def get_activations_and_quantizers(
+ self, layer: Layer) -> Sequence[ActivationAndQuantizer]:
+ """See base class."""
+ self._assert_activation_layer(layer)
+ return []
+
+ def set_quantize_weights(
+ self,
+ layer: Layer,
+ quantize_weights: Sequence[tf.Tensor]):
+ """See base class."""
+ self._assert_activation_layer(layer)
+
+ def set_quantize_activations(
+ self,
+ layer: Layer,
+ quantize_activations: Sequence[Activation]):
+ """See base class."""
+ self._assert_activation_layer(layer)
+
+ def get_output_quantizers(self, layer: Layer) -> Sequence[Quantizer]:
+ """See base class."""
+ self._assert_activation_layer(layer)
+
+ if not hasattr(layer.activation, '__name__'):
+ raise ValueError('Activation {} not supported by '
+ 'Default8BitActivationQuantizeConfig.'.format(
+ layer.activation))
+
+ # This code is copied from TFMOT repo, but added relu6 to support mobilenet.
+ if layer.activation.__name__ in ['relu', 'relu6']:
+ # 'relu' should generally get fused into the previous layer.
+ return [MovingAverageQuantizer(
+ num_bits=8, per_axis=False, symmetric=False, narrow_range=False)]
+ elif layer.activation.__name__ in ['linear', 'softmax', 'sigmoid']:
+ return []
+
+ raise ValueError('Activation {} not supported by '
+ 'Default8BitActivationQuantizeConfig.'.format(
+ layer.activation))
+
+ def get_config(self) -> Dict[str, Any]:
+ """Get a config for this quantizer config."""
+ return {}
+
+
+class NoQuantizeConfig(tfmot.quantization.keras.QuantizeConfig):
+ """Empty quantize config."""
+
+ # Configure how to quantize weights.
+ def get_weights_and_quantizers(self, layer):
+ return []
+
+ # Configure how to quantize activations.
+ def get_activations_and_quantizers(self, layer):
+ return []
+
+ def set_quantize_weights(self, layer, quantize_weights):
+ # Add this line for each item returned in `get_weights_and_quantizers`
+ # , in the same order
+ pass
+
+ def set_quantize_activations(self, layer, quantize_activations):
+ # Add this line for each item returned in `get_activations_and_quantizers`
+ # , in the same order.
+ pass
+
+ # Configure how to quantize outputs (may be equivalent to activations).
+ def get_output_quantizers(self, layer):
+ return []
+
+ def get_config(self):
+ return {}
+
+
+class Default8BitQuantizeConfig(tfmot.quantization.keras.QuantizeConfig):
+ """QuantizeConfig for non recurrent Keras layers."""
+
+ def __init__(self, weight_attrs, activation_attrs, quantize_output):
+ self.weight_attrs = weight_attrs
+ self.activation_attrs = activation_attrs
+ self.quantize_output = quantize_output
+
+ # TODO(pulkitb): For some layers such as Conv2D, per_axis should be True.
+ # Add mapping for which layers support per_axis.
+ self.weight_quantizer = LastValueQuantizer(
+ num_bits=8, per_axis=False, symmetric=True, narrow_range=True)
+ self.activation_quantizer = MovingAverageQuantizer(
+ num_bits=8, per_axis=False, symmetric=False, narrow_range=False)
+
+ def get_weights_and_quantizers(self, layer):
+ return [(getattr(layer, weight_attr), self.weight_quantizer)
+ for weight_attr in self.weight_attrs]
+
+ def get_activations_and_quantizers(self, layer):
+ return [(getattr(layer, activation_attr), self.activation_quantizer)
+ for activation_attr in self.activation_attrs]
+
+ def set_quantize_weights(self, layer, quantize_weights):
+ if len(self.weight_attrs) != len(quantize_weights):
+ raise ValueError(
+ '`set_quantize_weights` called on layer {} with {} '
+ 'weight parameters, but layer expects {} values.'.format(
+ layer.name, len(quantize_weights), len(self.weight_attrs)))
+
+ for weight_attr, weight in zip(self.weight_attrs, quantize_weights):
+ current_weight = getattr(layer, weight_attr)
+ if current_weight.shape != weight.shape:
+ raise ValueError('Existing layer weight shape {} is incompatible with'
+ 'provided weight shape {}'.format(
+ current_weight.shape, weight.shape))
+
+ setattr(layer, weight_attr, weight)
+
+ def set_quantize_activations(self, layer, quantize_activations):
+ if len(self.activation_attrs) != len(quantize_activations):
+ raise ValueError(
+ '`set_quantize_activations` called on layer {} with {} '
+ 'activation parameters, but layer expects {} values.'.format(
+ layer.name, len(quantize_activations),
+ len(self.activation_attrs)))
+
+ for activation_attr, activation in zip(
+ self.activation_attrs, quantize_activations):
+ setattr(layer, activation_attr, activation)
+
+ def get_output_quantizers(self, layer):
+ if self.quantize_output:
+ return [self.activation_quantizer]
+ return []
+
+ @classmethod
+ def from_config(cls, config):
+ """Instantiates a `Default8BitQuantizeConfig` from its config.
+
+ Args:
+ config: Output of `get_config()`.
+
+ Returns:
+ A `Default8BitQuantizeConfig` instance.
+ """
+ return cls(**config)
+
+ def get_config(self):
+ # TODO(pulkitb): Add weight and activation quantizer to config.
+ # Currently it's created internally, but ideally the quantizers should be
+ # part of the constructor and passed in from the registry.
+ return {
+ 'weight_attrs': self.weight_attrs,
+ 'activation_attrs': self.activation_attrs,
+ 'quantize_output': self.quantize_output
+ }
+
+ def __eq__(self, other):
+ if not isinstance(other, Default8BitQuantizeConfig):
+ return False
+
+ return (self.weight_attrs == other.weight_attrs and
+ self.activation_attrs == self.activation_attrs and
+ self.weight_quantizer == other.weight_quantizer and
+ self.activation_quantizer == other.activation_quantizer and
+ self.quantize_output == other.quantize_output)
+
+ def __ne__(self, other):
+ return not self.__eq__(other)
+
+
+def _types_dict():
+ return {
+ 'NoQuantizer':
+ NoQuantizer,
+ 'LastValueQuantizer':
+ LastValueQuantizer,
+ 'MovingAverageQuantizer':
+ MovingAverageQuantizer,
+ 'DefaultEinsumDenseQuantizeConfig':
+ DefaultEinsumDenseQuantizeConfig,
+ 'DefaultMultiHeadAttentionQuantizeConfig':
+ DefaultMultiHeadAttentionQuantizeConfig,
+ 'Default8BitOutputQuantizeConfig':
+ Default8BitOutputQuantizeConfig,
+ 'Default8BitActivationQuantizeConfig':
+ Default8BitActivationQuantizeConfig,
+ 'NoQuantizeConfig':
+ NoQuantizeConfig,
+ 'Default8BitQuantizeConfig':
+ Default8BitQuantizeConfig,
+ }
diff --git a/official/projects/qat/nlp/quantization/configs_test.py b/official/projects/qat/nlp/quantization/configs_test.py
new file mode 100644
index 00000000000..24ca716f43d
--- /dev/null
+++ b/official/projects/qat/nlp/quantization/configs_test.py
@@ -0,0 +1,286 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for configs.py."""
+
+from absl.testing import parameterized
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+
+from official.modeling import tf_utils
+from official.projects.qat.nlp.quantization import configs
+
+
+class _TestHelper(object):
+
+ def _convert_list(self, list_of_tuples):
+ """Transforms a list of 2-tuples to a tuple of 2 lists.
+
+ `QuantizeConfig` methods return a list of 2-tuples in the form
+ [(weight1, quantizer1), (weight2, quantizer2)]. This function converts
+ it into a 2-tuple of lists. ([weight1, weight2]), (quantizer1, quantizer2).
+
+ Args:
+ list_of_tuples: List of 2-tuples.
+
+ Returns:
+ 2-tuple of lists.
+ """
+ list1 = []
+ list2 = []
+ for a, b in list_of_tuples:
+ list1.append(a)
+ list2.append(b)
+
+ return list1, list2
+
+ # TODO(pulkitb): Consider asserting on full equality for quantizers.
+
+ def _assert_weight_quantizers(self, quantizer_list):
+ for quantizer in quantizer_list:
+ self.assertIsInstance(
+ quantizer,
+ tfmot.quantization.keras.quantizers.LastValueQuantizer)
+
+ def _assert_activation_quantizers(self, quantizer_list):
+ for quantizer in quantizer_list:
+ self.assertIsInstance(
+ quantizer,
+ tfmot.quantization.keras.quantizers.MovingAverageQuantizer)
+
+ def _assert_kernel_equality(self, a, b):
+ self.assertAllEqual(a.numpy(), b.numpy())
+
+
+class Default8BitQuantizeConfigTest(tf.test.TestCase, _TestHelper):
+
+ def _simple_dense_layer(self):
+ layer = tf_keras.layers.Dense(2)
+ layer.build(input_shape=(3,))
+ return layer
+
+ def testGetsQuantizeWeightsAndQuantizers(self):
+ layer = self._simple_dense_layer()
+
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+ (weights, weight_quantizers) = self._convert_list(
+ quantize_config.get_weights_and_quantizers(layer))
+
+ self._assert_weight_quantizers(weight_quantizers)
+ self.assertEqual([layer.kernel], weights)
+
+ def testGetsQuantizeActivationsAndQuantizers(self):
+ layer = self._simple_dense_layer()
+
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+ (activations, activation_quantizers) = self._convert_list(
+ quantize_config.get_activations_and_quantizers(layer))
+
+ self._assert_activation_quantizers(activation_quantizers)
+ self.assertEqual([layer.activation], activations)
+
+ def testSetsQuantizeWeights(self):
+ layer = self._simple_dense_layer()
+ quantize_kernel = tf_keras.backend.variable(
+ np.ones(layer.kernel.shape.as_list()))
+
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+ quantize_config.set_quantize_weights(layer, [quantize_kernel])
+
+ self._assert_kernel_equality(layer.kernel, quantize_kernel)
+
+ def testSetsQuantizeActivations(self):
+ layer = self._simple_dense_layer()
+ quantize_activation = tf_keras.activations.relu
+
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+ quantize_config.set_quantize_activations(layer, [quantize_activation])
+
+ self.assertEqual(layer.activation, quantize_activation)
+
+ def testSetsQuantizeWeights_ErrorOnWrongNumberOfWeights(self):
+ layer = self._simple_dense_layer()
+ quantize_kernel = tf_keras.backend.variable(
+ np.ones(layer.kernel.shape.as_list()))
+
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+
+ with self.assertRaises(ValueError):
+ quantize_config.set_quantize_weights(layer, [])
+
+ with self.assertRaises(ValueError):
+ quantize_config.set_quantize_weights(layer,
+ [quantize_kernel, quantize_kernel])
+
+ def testSetsQuantizeWeights_ErrorOnWrongShapeOfWeight(self):
+ layer = self._simple_dense_layer()
+ quantize_kernel = tf_keras.backend.variable(np.ones([1, 2]))
+
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+
+ with self.assertRaises(ValueError):
+ quantize_config.set_quantize_weights(layer, [quantize_kernel])
+
+ def testSetsQuantizeActivations_ErrorOnWrongNumberOfActivations(self):
+ layer = self._simple_dense_layer()
+ quantize_activation = tf_keras.activations.relu
+
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+
+ with self.assertRaises(ValueError):
+ quantize_config.set_quantize_activations(layer, [])
+
+ with self.assertRaises(ValueError):
+ quantize_config.set_quantize_activations(
+ layer, [quantize_activation, quantize_activation])
+
+ def testGetsResultQuantizers_ReturnsQuantizer(self):
+ layer = self._simple_dense_layer()
+ quantize_config = configs.Default8BitQuantizeConfig(
+ [], [], True)
+
+ output_quantizers = quantize_config.get_output_quantizers(layer)
+
+ self.assertLen(output_quantizers, 1)
+ self._assert_activation_quantizers(output_quantizers)
+
+ def testGetsResultQuantizers_EmptyWhenFalse(self):
+ layer = self._simple_dense_layer()
+ quantize_config = configs.Default8BitQuantizeConfig(
+ [], [], False)
+
+ output_quantizers = quantize_config.get_output_quantizers(layer)
+
+ self.assertEqual([], output_quantizers)
+
+ def testSerialization(self):
+ quantize_config = configs.Default8BitQuantizeConfig(
+ ['kernel'], ['activation'], False)
+
+ expected_config = {
+ 'class_name': 'Default8BitQuantizeConfig',
+ 'config': {
+ 'weight_attrs': ['kernel'],
+ 'activation_attrs': ['activation'],
+ 'quantize_output': False
+ }
+ }
+ serialized_quantize_config = tf_utils.serialize_keras_object(
+ quantize_config
+ )
+
+ self.assertEqual(expected_config, serialized_quantize_config)
+
+ quantize_config_from_config = (
+ tf_utils.deserialize_keras_object(
+ serialized_quantize_config,
+ module_objects=globals(),
+ custom_objects=configs._types_dict(),
+ )
+ )
+
+ self.assertEqual(quantize_config, quantize_config_from_config)
+
+
+@parameterized.parameters(
+ configs.LastValueQuantizer,
+ configs.MovingAverageQuantizer,
+ configs.NoQuantizer)
+class QuantizersTest(tf.test.TestCase, parameterized.TestCase):
+
+ def _simple_dense_layer(self):
+ layer = tf_keras.layers.Dense(2)
+ layer.build(input_shape=(3,))
+ return layer
+
+ def _get_quant_params(self, quantizer_type):
+ if quantizer_type == configs.NoQuantizer:
+ return {}
+
+ return {
+ 'num_bits': 8,
+ 'per_axis': False,
+ 'symmetric': False,
+ 'narrow_range': False
+ }
+
+ def _test_quantizer(self, quantizer):
+ inputs = tf.Variable(
+ np.array([[-1.0, 0.5], [0.0, 1.0]]),
+ name='inputs',
+ dtype=tf.dtypes.float32)
+ min_var = tf.Variable(0.0)
+ max_var = tf.Variable(0.0)
+
+ weights = {'min_var': min_var, 'max_var': max_var}
+ quant_tensor = quantizer(inputs, training=True, weights=weights)
+
+ results = self.evaluate(quant_tensor)
+ min_max_values = self.evaluate([min_var, max_var])
+
+ # TODO(pulkitb): Assert on expected values for testing.
+ # Since the underlying code is already tested in quant_ops_test.py, this
+ # just ensures the Quantizers code is wired properly.
+ print('Result: ', results)
+ print('min_var: ', min_max_values[0])
+ print('max_var: ', min_max_values[1])
+
+ layer = self._simple_dense_layer()
+ weights = quantizer.build(tf.TensorShape([1, 1, 1]), 'test', layer)
+ if isinstance(quantizer, (
+ configs.LastValueQuantizer, configs.MovingAverageQuantizer)):
+ self.assertLen(weights, 2)
+ self.assertFalse(weights['min_var'].trainable)
+ self.assertFalse(weights['max_var'].trainable)
+ elif isinstance(quantizer, configs.NoQuantizer):
+ self.assertEmpty(weights)
+
+ def testQuantizer(self, quantizer_type):
+ quantizer = quantizer_type(**self._get_quant_params(quantizer_type))
+
+ self._test_quantizer(quantizer)
+
+ def testSerialization(self, quantizer_type):
+ quantizer = quantizer_type(**self._get_quant_params(quantizer_type))
+
+ expected_config = {
+ 'class_name': quantizer_type.__name__,
+ 'config': self._get_quant_params(quantizer_type),
+ }
+ serialized_quantizer = tf_utils.serialize_keras_object(
+ quantizer
+ )
+
+ self.assertEqual(expected_config, serialized_quantizer)
+
+ quantizer_from_config = tf_utils.deserialize_keras_object(
+ serialized_quantizer,
+ module_objects=globals(),
+ custom_objects=configs._types_dict(),
+ )
+
+ self.assertEqual(quantizer, quantizer_from_config)
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/qat/nlp/quantization/helper.py b/official/projects/qat/nlp/quantization/helper.py
new file mode 100644
index 00000000000..0cea885578f
--- /dev/null
+++ b/official/projects/qat/nlp/quantization/helper.py
@@ -0,0 +1,49 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Quantization helpers."""
+
+import tensorflow_model_optimization as tfmot
+
+
+class LayerQuantizerHelper(object):
+ """Helper class that handles quantizers."""
+
+ def __init__(self, *args, **kwargs):
+ self._quantizers = {}
+ self._quantizer_vars = {}
+ super().__init__(*args, **kwargs)
+
+ def _all_value_quantizer(self):
+ return tfmot.quantization.keras.quantizers.AllValuesQuantizer(
+ num_bits=8, per_axis=False, symmetric=False, narrow_range=False)
+
+ def _moving_average_quantizer(self):
+ return tfmot.quantization.keras.quantizers.MovingAverageQuantizer(
+ num_bits=8, per_axis=False, symmetric=False, narrow_range=False)
+
+ def _add_quantizer(self, name, all_value_quantizer=False):
+ if all_value_quantizer:
+ self._quantizers[name] = self._all_value_quantizer()
+ else:
+ self._quantizers[name] = self._moving_average_quantizer()
+
+ def _apply_quantizer(self, name, inputs, training, **kwargs):
+ return self._quantizers[name](
+ inputs, training, self._quantizer_vars[name], **kwargs)
+
+ def _build_quantizer_vars(self):
+ for name in self._quantizers:
+ self._quantizer_vars[name] = self._quantizers[name].build(
+ tensor_shape=None, name=name, layer=self)
diff --git a/official/projects/qat/nlp/quantization/schemes.py b/official/projects/qat/nlp/quantization/schemes.py
new file mode 100644
index 00000000000..25520f7c376
--- /dev/null
+++ b/official/projects/qat/nlp/quantization/schemes.py
@@ -0,0 +1,206 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Quantization schemes."""
+import numpy as np
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+
+from official.modeling import tf_utils
+from official.projects.qat.nlp.modeling.layers import mobile_bert_layers
+from official.projects.qat.nlp.modeling.layers import transformer_encoder_block
+from official.projects.qat.nlp.quantization import configs
+
+keras = tf.keras
+default_8bit_transforms = tfmot.quantization.keras.default_8bit.default_8bit_transforms
+LayerNode = tfmot.quantization.keras.graph_transformations.transforms.LayerNode
+LayerPattern = tfmot.quantization.keras.graph_transformations.transforms.LayerPattern
+
+
+class TransformerEncoderBlockQuantize(
+ tfmot.quantization.keras.graph_transformations.transforms.Transform):
+ """Add QAT support for Keras Custom layer."""
+
+ _QUANTIZATION_AWARE_TRAINING_WEIGHT_NAMES = frozenset({
+ 'optimizer_step',
+ 'output_max', 'output_min',
+ 'kernel_min', 'kernel_max',
+ 'depthwise_kernel_min', 'depthwise_kernel_max',
+ 'query_min', 'query_max',
+ 'attention_scores_min', 'attention_scores_max',
+ 'attention_output_min', 'attention_output_max',
+ 'masked_softmax_attention_mask_min',
+ 'masked_softmax_attention_mask_max',
+ 'masked_softmax_sub1_min', 'masked_softmax_sub1_max',
+ 'masked_softmax_mask1_min', 'masked_softmax_mask1_max',
+ 'masked_softmax_sub2_min', 'masked_softmax_sub2_max',
+ 'masked_softmax_clamp_min', 'masked_softmax_clamp_max',
+ 'masked_softmax_mask2_min', 'masked_softmax_mask2_max',
+ 'masked_softmax_adder_sub_min', 'masked_softmax_adder_sub_max',
+ 'masked_softmax_adder_mul_min', 'masked_softmax_adder_mul_max',
+ 'masked_softmax_add_min', 'masked_softmax_add_max',
+ 'post_activation_min', 'post_activation_max',
+ 'word_embedding_out_min', 'word_embedding_out_max',
+ 'pos_embedding_out_min', 'pos_embedding_out_max',
+ 'type_embedding_out_min', 'type_embedding_out_max',
+ 'bias_min', 'bias_max'
+ })
+
+ _SUPPOTRED_MODEL_WEIGHT_NAMES = frozenset({
+ 'kernel', 'depthwise_kernel', 'bias',
+ 'gamma', 'beta', 'moving_mean', 'moving_variance',
+ 'embeddings'
+ })
+
+ def __init__(self):
+ super().__init__()
+ self._original_layer_pattern = 'modeling>TransformerEncoderBlock'
+ self._quantized_layer_class = transformer_encoder_block.TransformerEncoderBlockQuantized
+
+ def pattern(self) -> LayerPattern:
+ """See base class."""
+ return LayerPattern(self._original_layer_pattern)
+
+ def _is_quantization_weight_name(self, name):
+ simple_name = name.split('/')[-1].split(':')[0]
+ if simple_name in self._QUANTIZATION_AWARE_TRAINING_WEIGHT_NAMES:
+ return True
+ if simple_name in self._SUPPOTRED_MODEL_WEIGHT_NAMES:
+ return False
+ raise ValueError('Variable name {} is not supported on '
+ 'CustomLayerQuantize({}) transform.'.format(
+ simple_name,
+ self._original_layer_pattern))
+
+ def replacement(self, match_layer: LayerNode) -> LayerNode:
+ """See base class."""
+ bottleneck_layer = match_layer.layer
+ bottleneck_config = bottleneck_layer['config']
+ bottleneck_names_and_weights = list(match_layer.names_and_weights)
+ quantized_layer = self._quantized_layer_class(
+ **bottleneck_config)
+
+ quantized_layer_config = quantized_layer.get_config()
+ if 'hidden_size' in quantized_layer_config:
+ dummy_input_shape = [
+ 1, 1, quantized_layer_config['hidden_size']]
+ quantized_layer.compute_output_shape(dummy_input_shape)
+ elif 'num_attention_heads' in quantized_layer_config:
+ dummy_input_shape = [
+ 1, 1, quantized_layer_config['num_attention_heads']]
+ quantized_layer.compute_output_shape(dummy_input_shape)
+ else:
+ dummy_input_shape = [1, 1]
+ quantized_layer(np.zeros(shape=dummy_input_shape, dtype=np.int32),
+ np.zeros(shape=dummy_input_shape, dtype=np.int32),
+ training=False)
+
+ quantized_names_and_weights = zip(
+ [weight.name for weight in quantized_layer.weights],
+ quantized_layer.get_weights())
+ match_idx = 0
+ names_and_weights = []
+ for name_and_weight in quantized_names_and_weights:
+ if not self._is_quantization_weight_name(name=name_and_weight[0]):
+ name_and_weight = bottleneck_names_and_weights[match_idx]
+ match_idx = match_idx + 1
+ names_and_weights.append(name_and_weight)
+
+ if match_idx != len(bottleneck_names_and_weights):
+ raise ValueError('{}/{} of Bottleneck weights is transformed.'.format(
+ match_idx, len(bottleneck_names_and_weights)))
+ quantized_layer_config = tf_utils.serialize_layer(
+ quantized_layer, use_legacy_format=True
+ )
+ quantized_layer_config['name'] = quantized_layer_config['config']['name']
+ layer_metadata = {
+ 'quantize_config':
+ configs.NoQuantizeConfig()}
+
+ return LayerNode(
+ quantized_layer_config,
+ metadata=layer_metadata,
+ names_and_weights=names_and_weights)
+
+
+class MobileBertTransformerQuantize(TransformerEncoderBlockQuantize):
+
+ def __init__(self):
+ super().__init__()
+ self._original_layer_pattern = 'Text>MobileBertTransformer'
+ self._quantized_layer_class = mobile_bert_layers.MobileBertTransformerQuantized
+
+
+class MobileBertEmbeddingQuantize(TransformerEncoderBlockQuantize):
+
+ def __init__(self):
+ super().__init__()
+ self._original_layer_pattern = 'Text>MobileBertEmbedding'
+ self._quantized_layer_class = mobile_bert_layers.MobileBertEmbeddingQuantized
+
+
+class QuantizeLayoutTransform(
+ tfmot.quantization.keras.QuantizeLayoutTransform):
+ """Default model transformations."""
+
+ def apply(self, model, layer_quantize_map):
+ """Implement default 8-bit transforms.
+
+ Currently this means the following.
+ 1. Pull activations into layers, and apply fuse activations. (TODO)
+ 2. Modify range in incoming layers for Concat. (TODO)
+ 3. Fuse Conv2D/DepthwiseConv2D + BN into single layer.
+
+ Args:
+ model: Keras model to be quantized.
+ layer_quantize_map: Map with keys as layer names, and values as dicts
+ containing custom `QuantizeConfig`s which may have been passed with
+ layers.
+
+ Returns:
+ (Transformed Keras model to better match TensorFlow Lite backend, updated
+ layer quantize map.)
+ """
+
+ transforms = [
+ default_8bit_transforms.SeparableConv1DQuantize(),
+ default_8bit_transforms.SeparableConvQuantize(),
+ default_8bit_transforms.Conv2DReshapeBatchNormReLUQuantize(),
+ default_8bit_transforms.Conv2DReshapeBatchNormActivationQuantize(),
+ default_8bit_transforms.Conv2DBatchNormReLUQuantize(),
+ default_8bit_transforms.Conv2DBatchNormActivationQuantize(),
+ default_8bit_transforms.Conv2DReshapeBatchNormQuantize(),
+ default_8bit_transforms.Conv2DBatchNormQuantize(),
+ default_8bit_transforms.ConcatTransform6Inputs(),
+ default_8bit_transforms.ConcatTransform5Inputs(),
+ default_8bit_transforms.ConcatTransform4Inputs(),
+ default_8bit_transforms.ConcatTransform3Inputs(),
+ default_8bit_transforms.ConcatTransform(),
+ default_8bit_transforms.LayerReLUQuantize(),
+ default_8bit_transforms.LayerReluActivationQuantize(),
+ TransformerEncoderBlockQuantize(),
+ MobileBertTransformerQuantize(),
+ MobileBertEmbeddingQuantize(),
+ ]
+ return tfmot.quantization.keras.graph_transformations.model_transformer.ModelTransformer(
+ model, transforms,
+ set(layer_quantize_map.keys()), layer_quantize_map).transform()
+
+
+class Default8BitQuantizeScheme(
+ tfmot.quantization.keras.default_8bit.Default8BitQuantizeScheme):
+
+ def get_layout_transformer(self):
+ return QuantizeLayoutTransform()
diff --git a/official/projects/qat/nlp/quantization/wrappers.py b/official/projects/qat/nlp/quantization/wrappers.py
new file mode 100644
index 00000000000..3f601aba739
--- /dev/null
+++ b/official/projects/qat/nlp/quantization/wrappers.py
@@ -0,0 +1,52 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Quantization Wrappers."""
+import tensorflow_model_optimization as tfmot
+
+
+class MultiHeadAttentionQuantizeWrapper(
+ tfmot.quantization.keras.QuantizeWrapperV2):
+ """Custom quantize wrapper for the MultiHeadAttention layer."""
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self._first_call_built = False
+
+ def build(self, input_shape):
+ self.layer.build(input_shape)
+
+ def call(self,
+ query,
+ value,
+ key=None,
+ attention_mask=None,
+ return_attention_scores=False,
+ training=None):
+ if not self._first_call_built:
+ # pylint: disable=protected-access
+ self.layer._build_from_signature(query=query, value=value, key=key)
+ # pylint: enable=protected-access
+ self.layer.call(
+ query, value, key=key, attention_mask=attention_mask,
+ return_attention_scores=return_attention_scores,
+ training=training)
+ super().build(input_shape=None)
+ self._first_call_built = True
+
+ return super().call(
+ query, value=value, key=key, attention_mask=attention_mask,
+ return_attention_scores=return_attention_scores,
+ training=training
+ )
diff --git a/official/projects/qat/nlp/registry_imports.py b/official/projects/qat/nlp/registry_imports.py
new file mode 100644
index 00000000000..54f046e42a8
--- /dev/null
+++ b/official/projects/qat/nlp/registry_imports.py
@@ -0,0 +1,21 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""All necessary imports for registration on qat project."""
+# pylint: disable=unused-import
+from official.projects.qat.nlp import configs
+from official.projects.qat.nlp.modeling.layers import mobile_bert_layers
+from official.projects.qat.nlp.modeling.layers import multi_head_attention
+from official.projects.qat.nlp.modeling.layers import transformer_encoder_block
+from official.projects.qat.nlp.tasks import question_answering
diff --git a/official/projects/qat/nlp/tasks/__init__.py b/official/projects/qat/nlp/tasks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/nlp/tasks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/nlp/tasks/question_answering.py b/official/projects/qat/nlp/tasks/question_answering.py
new file mode 100644
index 00000000000..7a6a26408b4
--- /dev/null
+++ b/official/projects/qat/nlp/tasks/question_answering.py
@@ -0,0 +1,84 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Question/Answering configuration definition."""
+import dataclasses
+
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.core import task_factory
+from official.nlp import modeling
+from official.nlp.tasks import question_answering
+from official.projects.qat.nlp.modeling.layers import mobile_bert_layers
+from official.projects.qat.nlp.modeling.layers import transformer_encoder_block
+from official.projects.qat.nlp.modeling.models import bert_span_labeler
+from official.projects.qat.nlp.quantization import configs
+from official.projects.qat.nlp.quantization import schemes
+
+
+@dataclasses.dataclass
+class QuantizedModelQAConfig(question_answering.QuestionAnsweringConfig):
+ pass
+
+
+@task_factory.register_task_cls(QuantizedModelQAConfig)
+class QuantizedModelQATask(question_answering.QuestionAnsweringTask):
+ """Task object for question answering with QAT."""
+
+ def build_model(self):
+ model = super(QuantizedModelQATask, self).build_model()
+ # pylint: disable=protected-access
+ encoder_network = model._network
+ # pylint: enable=protected-access
+
+ with tfmot.quantization.keras.quantize_scope({
+ 'TruncatedNormal':
+ tf_keras.initializers.TruncatedNormal,
+ 'MobileBertTransformerQuantized':
+ mobile_bert_layers.MobileBertTransformerQuantized,
+ 'MobileBertEmbeddingQuantized':
+ mobile_bert_layers.MobileBertEmbeddingQuantized,
+ 'TransformerEncoderBlockQuantized':
+ transformer_encoder_block.TransformerEncoderBlockQuantized,
+ 'NoQuantizeConfig':
+ configs.NoQuantizeConfig,
+ }):
+ def quantize_annotate_layer(layer):
+ if isinstance(layer, (tf_keras.layers.LayerNormalization)):
+ return tfmot.quantization.keras.quantize_annotate_layer(
+ layer, configs.Default8BitOutputQuantizeConfig())
+ if isinstance(layer, (tf_keras.layers.Dense,
+ tf_keras.layers.Dropout)):
+ return tfmot.quantization.keras.quantize_annotate_layer(layer)
+ if isinstance(layer, (modeling.layers.TransformerEncoderBlock,
+ modeling.layers.MobileBertTransformer,
+ modeling.layers.MobileBertEmbedding)):
+ return tfmot.quantization.keras.quantize_annotate_layer(
+ layer, configs.NoQuantizeConfig())
+ return layer
+
+ annotated_encoder_network = tf_keras.models.clone_model(
+ encoder_network,
+ clone_function=quantize_annotate_layer,
+ )
+ quantized_encoder_network = tfmot.quantization.keras.quantize_apply(
+ annotated_encoder_network, scheme=schemes.Default8BitQuantizeScheme())
+
+ encoder_cfg = self.task_config.model.encoder.get()
+ model = bert_span_labeler.BertSpanLabelerQuantized(
+ network=quantized_encoder_network,
+ initializer=tf_keras.initializers.TruncatedNormal(
+ stddev=encoder_cfg.initializer_range))
+ return model
diff --git a/official/projects/qat/nlp/tasks/question_answering_test.py b/official/projects/qat/nlp/tasks/question_answering_test.py
new file mode 100644
index 00000000000..d3436552faf
--- /dev/null
+++ b/official/projects/qat/nlp/tasks/question_answering_test.py
@@ -0,0 +1,106 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for official.nlp.tasks.question_answering."""
+import json
+import os
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.nlp.configs import encoders
+from official.nlp.data import question_answering_dataloader
+from official.nlp.tasks import question_answering as qa_cfg
+from official.projects.qat.nlp.tasks import question_answering
+
+
+class QuestionAnsweringTaskTest(tf.test.TestCase, parameterized.TestCase):
+
+ def setUp(self):
+ super(QuestionAnsweringTaskTest, self).setUp()
+ self._encoder_config = encoders.EncoderConfig(
+ bert=encoders.BertEncoderConfig(vocab_size=30522, num_layers=1))
+ self._train_data_config = question_answering_dataloader.QADataConfig(
+ input_path="dummy", seq_length=128, global_batch_size=1)
+
+ val_data = {
+ "version":
+ "1.1",
+ "data": [{
+ "paragraphs": [{
+ "context":
+ "Sky is blue.",
+ "qas": [{
+ "question":
+ "What is blue?",
+ "id":
+ "1234",
+ "answers": [{
+ "text": "Sky",
+ "answer_start": 0
+ }, {
+ "text": "Sky",
+ "answer_start": 0
+ }, {
+ "text": "Sky",
+ "answer_start": 0
+ }]
+ }]
+ }]
+ }]
+ }
+ self._val_input_path = os.path.join(self.get_temp_dir(), "val_data.json")
+ with tf.io.gfile.GFile(self._val_input_path, "w") as writer:
+ writer.write(json.dumps(val_data, indent=4) + "\n")
+
+ self._test_vocab = os.path.join(self.get_temp_dir(), "vocab.txt")
+ with tf.io.gfile.GFile(self._test_vocab, "w") as writer:
+ writer.write("[PAD]\n[UNK]\n[CLS]\n[SEP]\n[MASK]\nsky\nis\nblue\n")
+
+ def _get_validation_data_config(self, version_2_with_negative=False):
+ return question_answering_dataloader.QADataConfig(
+ is_training=False,
+ input_path=self._val_input_path,
+ input_preprocessed_data_path=self.get_temp_dir(),
+ seq_length=128,
+ global_batch_size=1,
+ version_2_with_negative=version_2_with_negative,
+ vocab_file=self._test_vocab,
+ tokenization="WordPiece",
+ do_lower_case=True)
+
+ @parameterized.named_parameters(("squad1", False), ("squad2", True))
+ def test_predict(self, version_2_with_negative):
+ validation_data = self._get_validation_data_config(
+ version_2_with_negative=version_2_with_negative)
+
+ config = question_answering.QuantizedModelQAConfig(
+ model=qa_cfg.ModelConfig(encoder=self._encoder_config),
+ train_data=self._train_data_config,
+ validation_data=validation_data)
+ task = question_answering.QuantizedModelQATask(config)
+ model = task.build_model()
+
+ all_predictions, all_nbest, scores_diff = qa_cfg.predict(
+ task, validation_data, model)
+ self.assertLen(all_predictions, 1)
+ self.assertLen(all_nbest, 1)
+ if version_2_with_negative:
+ self.assertLen(scores_diff, 1)
+ else:
+ self.assertEmpty(scores_diff)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/qat/nlp/train.py b/official/projects/qat/nlp/train.py
new file mode 100644
index 00000000000..af44edcaada
--- /dev/null
+++ b/official/projects/qat/nlp/train.py
@@ -0,0 +1,68 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""TFM common training driver."""
+
+from absl import app
+from absl import flags
+import gin
+
+from official.common import distribute_utils
+from official.common import registry_imports # pylint: disable=unused-import
+from official.common import flags as tfm_flags
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+from official.projects.qat.nlp import registry_imports as qat_registry_imports # pylint: disable=unused-import
+
+FLAGS = flags.FLAGS
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ # Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
+ # dtype is float16
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu,
+ **params.runtime.model_parallelism())
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir)
+
+ train_utils.save_gin_config(FLAGS.mode, model_dir)
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(main)
diff --git a/official/projects/qat/vision/README.md b/official/projects/qat/vision/README.md
index 4184fa6763f..3c12f95c5c6 100644
--- a/official/projects/qat/vision/README.md
+++ b/official/projects/qat/vision/README.md
@@ -1,63 +1,145 @@
-# Quantization Aware Training Project for Computer Vision Models
+# Quantization Aware Training for Computer Vision Models
⚠️ Disclaimer: All datasets hyperlinked from this page are not owned or
distributed by Google. The dataset is made available by third parties.
Please review the terms and conditions made available by the third parties
before using the data.
-## Overview
+## Description
-This project includes quantization aware training code for Computer Vision
+This project includes quantization aware training (QAT) code for computer vision
models. These are examples to show how to apply the Model Optimization Toolkit's
[quantization aware training API](https://www.tensorflow.org/model_optimization/guide/quantization/training).
+compared to post-training quantization (PTQ), QAT can minimize the quality loss
+from quantization, while still achieving the speed-up from integer quantization.
+Therefore, it is the preferable technique to use when there is strict
+requirement on model latency and quality. Please find our
+[blogpost](https://blog.tensorflow.org/2022/06/Adding-Quantization-aware-Training-and-Pruning-to-the-TensorFlow-Model-Garden.html)
+for more details.
-Note: Currently, we support a limited number of ML tasks & models (e.g., image
-classification and semantic segmentation)
-We will keep adding support for other ML tasks and models in the next releases.
+Currently, we support a limited number of vision tasks & models. We will keep
+adding support for other tasks and models in the next releases.
-## How to train a model
+You can follow this
+[Colab notebook](https://colab.research.google.com/github/tensorflow/models/blob/master/official/projects/qat/vision/docs/qat_tutorial.ipynb)
+to try QAT.
-```
-EXPERIMENT=xxx # Change this for your run, for example, 'mobilenet_imagenet_qat'
-CONFIG_FILE=xxx # Change this for your run, for example, path of imagenet_mobilenetv2_qat_gpu.yaml
-MODEL_DIR=xxx # Change this for your run, for example, /tmp/model_dir
-$ python3 train.py \
---experiment=${EXPERIMENT} \
---config_file=${CONFIG_FILE} \
---model_dir=${MODEL_DIR} \
---mode=train_and_eval
-```
+## History
+
+### Jun. 9, 2022
+
+- First release of vision models covering image classification and semantic
+ segmentation tasks. Support ResNet, MobileNetV2, MobileNetV3 large and
+ Multi-hardware MobileNet, and DeepLabV3/V3+.
+
+### Nov. 30, 2022
+
+- Release of support for object detection task (RetinaNet).
+
+## Maintainers
-## Image Classification
+- Jaehong Kim ([Xhark](https://github.com/Xhark))
+* Fang Yang ([fyangf](https://github.com/fyangf))
+* Shixin Luo ([luotigerlsx](https://github.com/luotigerlsx))
-
-
-Comparison of Imagenet top-1 accuracy for the classification models
-
+## Requirements
-Note: The Top-1 model accuracy is measured on the validation set of [ImageNet](https://www.image-net.org/).
+[](https://badge.fury.io/py/tensorflow)
+[](https://badge.fury.io/py/tf-models-official)
+## Results
+### Image Classification
-### Pre-trained Models
+Model is trained on ImageNet1K train set and evaluated on the validation set.
-|Model |Resolution|Top-1 Accuracy (FP32)|Top-1 Accuracy (Int8/PTQ)|Top-1 Accuracy (Int8/QAT)|Config |Download |
+
+|Model |Resolution|Top-1 Accuracy (FP32)|Top-1 Accuracy (INT8)|Top-1 Accuracy (QAT INT8)|Config |Download |
|----------------------|----------|---------------------|-------------------------|-------------------------|--------------------------------------------------------------------------------------------------------------------------------------------------------------------|------------------------------------------------------------------------------------------------------------------------------------------------|
-|MobileNetV2 |224x224 |72.782% |72.392% |72.792% |[config](https://github.com/tensorflow/models/blob/master/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv2_qat_gpu.yaml) |[TFLite(Int8/QAT)](https://storage.googleapis.com/tf_model_garden/vision/mobilenet/v2_1.0_int8/mobilenet_v2_1.00_224_int8.tflite) |
-|ResNet50 |224x224 |76.710% |76.420% |77.200% |[config](https://github.com/tensorflow/models/blob/master/official/projects/qat/vision/configs/experiments/image_classification/imagenet_resnet50_qat_gpu.yaml) |[TFLite(Int8/QAT)](https://storage.googleapis.com/tf_model_garden/vision/resnet50_imagenet/resnet_50_224_int8.tflite) |
-|MobileNetV3.5 MultiAVG|224x224 |75.212% |74.122% |75.130% |[config](https://github.com/tensorflow/models/blob/master/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_qat_gpu.yaml)|[TFLite(Int8/QAT)](https://storage.googleapis.com/tf_model_garden/vision/mobilenet/v3.5multiavg_1.0_int8/mobilenet_v3.5multiavg_1.00_224_int8.tflite)|
+|MobileNetV2 |224x224 |72.78 |72.39 |72.79 |[config](https://github.com/tensorflow/models/blob/master/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv2_qat_gpu.yaml) |[TFLite(Int8/QAT)](https://storage.googleapis.com/tf_model_garden/vision/mobilenet/v2_1.0_int8/mobilenet_v2_1.00_224_int8.tflite) |
+|ResNet50 |224x224 |76.71 |76.42 |77.20 |[config](https://github.com/tensorflow/models/blob/master/official/projects/qat/vision/configs/experiments/image_classification/imagenet_resnet50_qat_gpu.yaml) |[TFLite(Int8/QAT)](https://storage.googleapis.com/tf_model_garden/vision/resnet50_imagenet/resnet_50_224_int8.tflite) |
+|MobileNetV3.5 MultiAVG|224x224 |75.21 |74.12 |75.13 |[config](https://github.com/tensorflow/models/blob/master/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_qat_gpu.yaml)|[TFLite(Int8/QAT)](https://storage.googleapis.com/tf_model_garden/vision/mobilenet/v3.5multiavg_1.0_int8/mobilenet_v3.5multiavg_1.00_224_int8.tflite)|
+
+### Object Detection
+
+Model is trained on COCO train set from scratch and evaluated on COCO validation
+set.
+
+model | resolution | mAP | mAP (FP32) | mAP (INT8) | mAP (QAT INT8) | download
+:----------------------- | :--------: | ---: | ---------: | ---------: | -------------: | ----------------:
+MobileNet v2 + RetinaNet | 256x256 | 23.3 | 23.3 | 0.04 | 21.7 | [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/qat/mobilenetv2_ssd_coco/mobilenetv2_ssd_i256_qat_ckpt.tar.gz) \| [tensorboard](https://tensorboard.dev/experiment/fAat72iXSqW8ZoTY3clMsg) [FP32](https://storage.cloud.google.com/tf_model_garden/vision/qat/mobilenetv2_ssd_coco/model_fp32.tflite) \| [INT8](https://storage.cloud.google.com/tf_model_garden/vision/qat/mobilenetv2_ssd_coco/model_int8_ptq.tflite) \| [QAT INT8](https://storage.cloud.google.com/tf_model_garden/vision/qat/mobilenetv2_ssd_coco/model_int8_qat.tflite)
-## Semantic Segmentation
+### Semantic Segmentation
Model is pretrained using COCO train set. Two datasets, Pascal VOC segmentation
-dataset and Cityscapes dataset (only for DeepLab v3+), are used to train and
-evaluate models. Model accuracy is measured on full Pascal VOC segmentation
-validation set.
+dataset and Cityscapes dataset are used to train and
+evaluate models.
+
+#### Pascal VOC
+
+model | resolution | mIoU | mIoU (FP32) | mIoU (INT8) | mIoU (QAT INT8) | download (tflite)|
+:------------------------- | :--------: | ----: | ----------: | ----------: | --------------: | ----------------:
+MobileNet v2 + DeepLab v3 | 512x512 | 75.27 | 75.30 | 73.95 | 74.68 | [FP32](https://storage.googleapis.com/tf_model_garden/vision/qat/deeplabv3_mobilenetv2_pascal_coco_0.21/model_none.tflite) \| [INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/deeplabv3_mobilenetv2_pascal_coco_0.21model_int8_full.tflite) \| [QAT INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/deeplabv3_mobilenetv2_pascal_coco_0.21/Fmodel_default.tflite)
+MobileNet v2 + DeepLab v3+ | 1024x2048 | 73.82 | 73.84 | 72.33 | 73.49 | [FP32](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/model_none.tflite) \| [INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/model_int8_full.tflite) \| [QAT INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/Fmodel_default.tflite)
+
+#### Cityscapes
+
+model | resolution | mIoU | mIoU (FP32) | mIoU (INT8) | mIoU (QAT INT8) | download (tflite)
+:------------------------- | :--------: | ----: | ----------: | ----------: | --------------: | ----------------:
+MobileNet v2 + DeepLab v3+ | 1024x2048 | 73.82 | 73.84 | 72.33 | 73.49 | [FP32](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/model_none.tflite) \| [INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/model_int8_full.tflite) \| [QAT INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/Fmodel_default.tflite)
+
+## Training
+
+It can run on Google Cloud Platform using Cloud TPU.
+[Here](https://cloud.google.com/tpu/docs/how-to) is the instruction of using
+Cloud TPU. Following the instructions to set up Cloud TPU and launch training,
+using object detection as an example:
+
+```shell
+
+# First download the pre-trained floating point model as QAT needs to finetune it.
+gsutil cp gs://tf_model_garden/vision/qat/mobilenetv2_ssd_coco/mobilenetv2_ssd_i256_ckpt.tar.gz /tmp/qat/
+
+# Extract the checkpoint.
+tar -xvzf /tmp/qat/mobilenetv2_ssd_i256_ckpt.tar.gz
+
+# Launch training. Note that we override the checkpoint path in the config file by "params_override" to supply the correct checkpoint.
+PARAMS_OVERRIDE="task.quantization.pretrained_original_checkpoint=/tmp/qat/mobilenetv2_ssd_i256_ckpt"
+EXPERIMENT=retinanet_mobile_coco_qat # Change this for your run, for example, 'mobilenet_imagenet_qat'.
+CONFIG_FILE=xxx # Change this for your run, for example, path of coco_mobilenetv2_qat_tpu_e2e.yaml.
+TPU_NAME="" # The name assigned while creating a Cloud TPU.
+MODEL_DIR="gs://" # Change this for your run, for example, /tmp/model_dir.
+$ python3 train.py \
+ --experiment=${EXPERIMENT} \
+ --config_file=${CONFIG_FILE} \
+ --model_dir=${MODEL_DIR} \
+ --tpu=$TPU_NAME \
+ --params_override=${PARAMS_OVERRIDE}
+ --mode=train
+```
+
+## Evaluation
+
+Please run this command line for evaluation.
+
+```shell
+EXPERIMENT=retinanet_mobile_coco # Change this for your run, for example, 'mobilenet_imagenet_qat'.
+CONFIG_FILE=xxx # Change this for your run, for example, path of coco_mobilenetv2_qat_tpu_e2e.yaml.
+TPU_NAME="" # The name assigned while creating a Cloud TPU.
+MODEL_DIR="gs://" # Change this for your run, for example, /tmp/model_dir.
+$ python3 train.py \
+ --experiment=${EXPERIMENT} \
+ --config_file=${CONFIG_FILE} \
+ --model_dir=${MODEL_DIR} \
+ --tpu=$TPU_NAME \
+ --mode=eval
+```
+
+## License
+
+[](https://opensource.org/licenses/Apache-2.0)
+
+This project is licensed under the terms of the **Apache License 2.0**.
-### Pre-trained Models
-model | resolution | mIoU | mIoU (FP32) | mIoU (FP16) | mIoU (INT8) | mIoU (QAT INT8) | download (tflite)
-:------------------------- | :--------: | ----: | ----------: | ----------: | ----------: | --------------: | ------------------------------------------------------: | ------------------------------------------------------: | -------------------------------------------------------: | ------------------------------------------------------: | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------: | ----------------:
-MobileNet v2 + DeepLab v3 | 512x512 | 75.27 | 75.30 | 75.32 | 73.95 | 74.68 | [FP32](https://storage.googleapis.com/tf_model_garden/vision/qat/deeplabv3_mobilenetv2_pascal_coco_0.21/model_none.tflite) \| [FP16](https://storage.googleapis.com/tf_model_garden/vision/qat/deeplabv3_mobilenetv2_pascal_coco_0.21/model_fp16.tflite) \| [INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/deeplabv3_mobilenetv2_pascal_coco_0.21model_int8_full.tflite) \| [QAT INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/deeplabv3_mobilenetv2_pascal_coco_0.21/Fmodel_default.tflite)
-MobileNet v2 + DeepLab v3+ | 1024x2048 | 73.82 | 73.84 | 73.65 | 72.33 | 73.49 | [FP32](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/model_none.tflite) \| [FP16](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/Fmodel_fp16.tflite) \| [INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/model_int8_full.tflite) \| [QAT INT8](https://storage.googleapis.com/tf_model_garden/vision/qat/mnv2_deeplabv3plus_cityscapes/Fmodel_default.tflite)
diff --git a/official/projects/qat/vision/__init__.py b/official/projects/qat/vision/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/vision/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/vision/configs/__init__.py b/official/projects/qat/vision/configs/__init__.py
index c542ea9f528..b9b45215bc3 100644
--- a/official/projects/qat/vision/configs/__init__.py
+++ b/official/projects/qat/vision/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,6 +13,7 @@
# limitations under the License.
"""Configs package definition."""
-
+from official.projects.qat.vision.configs import common
from official.projects.qat.vision.configs import image_classification
+from official.projects.qat.vision.configs import retinanet
from official.projects.qat.vision.configs import semantic_segmentation
diff --git a/official/projects/qat/vision/configs/common.py b/official/projects/qat/vision/configs/common.py
index 96d2bccb1c2..1739490ae0d 100644
--- a/official/projects/qat/vision/configs/common.py
+++ b/official/projects/qat/vision/configs/common.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,13 +25,22 @@ class Quantization(hyperparams.Config):
"""Quantization parameters.
Attributes:
+ version: A string that indicates the version of QAT API. Support `v2` and
+ `v3`.
pretrained_original_checkpoint: A string indicate pretrained checkpoint
location.
change_num_bits: A `bool` indicates whether to manually allocate num_bits.
num_bits_weight: An `int` number of bits for weight. Default to 8.
num_bits_activation: An `int` number of bits for activation. Default to 8.
+ quantize_detection_decoder: A `bool` indicates whether to quantize detection
+ decoder. It only works for detection model.
+ quantize_detection_head: A `bool` indicates whether to quantize detection
+ head. It only works for detection model.
"""
+ version: str = 'v2'
pretrained_original_checkpoint: Optional[str] = None
change_num_bits: bool = False
num_bits_weight: int = 8
num_bits_activation: int = 8
+ quantize_detection_decoder: bool = False
+ quantize_detection_head: bool = False
diff --git a/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv2_qat_tpu.yaml b/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv2_qat_tpu.yaml
new file mode 100644
index 00000000000..cbd398c68c9
--- /dev/null
+++ b/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv2_qat_tpu.yaml
@@ -0,0 +1,52 @@
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ num_classes: 1001
+ input_size: [224, 224, 3]
+ backbone:
+ type: 'mobilenet'
+ mobilenet:
+ model_id: 'MobileNetV2'
+ filter_size_scale: 1.0
+ dropout_rate: 0.1
+ losses:
+ l2_weight_decay: 0.0000001
+ one_hot: true
+ label_smoothing: 0.1
+ train_data:
+ input_path: '/readahead/200M/placer/prod/home/distbelief/imagenet-tensorflow/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'float32'
+ validation_data:
+ input_path: '/readahead/200M/placer/prod/home/distbelief/imagenet-tensorflow/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'float32'
+ drop_remainder: false
+ quantization:
+ pretrained_original_checkpoint: 'gs://**/mobilenetv2_gpu/22984194/ckpt-625500'
+trainer:
+ # With below setting, the accuracy of QAT reaches to accuracy 0.7279 after 43 hours with 8 GPUS.
+ train_steps: 31200
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ type: 'exponential'
+ exponential:
+ decay_rate: 0.9
+ decay_steps: 156
+ initial_learning_rate: 0.0001
+ name: 'ExponentialDecay'
+ offset: 0
+ staircase: true
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_steps: 0
diff --git a/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_qat_tpu.yaml b/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_qat_tpu.yaml
new file mode 100644
index 00000000000..6be91ed228c
--- /dev/null
+++ b/official/projects/qat/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_qat_tpu.yaml
@@ -0,0 +1,52 @@
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ num_classes: 1001
+ input_size: [224, 224, 3]
+ backbone:
+ type: 'mobilenet'
+ mobilenet:
+ model_id: 'MobileNetMultiAVG'
+ filter_size_scale: 1.0
+ dropout_rate: 0.3
+ losses:
+ l2_weight_decay: 0.000001
+ one_hot: true
+ label_smoothing: 0.1
+ train_data:
+ input_path: '/readahead/200M/placer/prod/home/distbelief/imagenet-tensorflow/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'float32'
+ validation_data:
+ input_path: '/readahead/200M/placer/prod/home/distbelief/imagenet-tensorflow/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'float32'
+ drop_remainder: false
+ quantization:
+ pretrained_original_checkpoint: 'gs://**/tf2_mhave_nobias_bn_aug05/28334857/ckpt-156000'
+trainer:
+ # With below setting, the accuracy of QAT reaches to accuracy 0.7513 after 30 hours with 8 GPUS.
+ train_steps: 31200
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ type: 'exponential'
+ exponential:
+ decay_rate: 0.9
+ decay_steps: 156
+ initial_learning_rate: 0.0004
+ name: 'ExponentialDecay'
+ offset: 0
+ staircase: true
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_steps: 0
diff --git a/official/projects/qat/vision/configs/experiments/retinanet/coco_mobilenetv2_qat_tpu_e2e.yaml b/official/projects/qat/vision/configs/experiments/retinanet/coco_mobilenetv2_qat_tpu_e2e.yaml
new file mode 100644
index 00000000000..7238f9357f1
--- /dev/null
+++ b/official/projects/qat/vision/configs/experiments/retinanet/coco_mobilenetv2_qat_tpu_e2e.yaml
@@ -0,0 +1,72 @@
+# --experiment_type=retinanet_mobile_coco_qat
+# COCO mAP: 23.02 from QAT training and 21.62 from the TFLite after conversion.
+# QAT only supports float32 tpu due to fake-quant op.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ losses:
+ l2_weight_decay: 0.0
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetV2'
+ filter_size_scale: 1.0
+ type: 'mobilenet'
+ decoder:
+ type: 'fpn'
+ fpn:
+ num_filters: 128
+ use_separable_conv: true
+ use_keras_layer: true
+ head:
+ num_convs: 4
+ num_filters: 128
+ use_separable_conv: true
+ input_size: [256, 256, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'relu6'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'float32'
+ global_batch_size: 256
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ dtype: 'float32'
+ global_batch_size: 256
+ is_training: false
+ drop_remainder: false
+ quantization:
+ pretrained_original_checkpoint: 'gs://**/coco_mobilenetv2_mobile_tpu/ckpt-277200'
+ quantize_detection_decoder: true
+ quantize_detection_head: true
+trainer:
+ best_checkpoint_eval_metric: AP
+ best_checkpoint_export_subdir: best_ckpt
+ best_checkpoint_metric_comp: higher
+ optimizer_config:
+ learning_rate:
+ type: 'exponential'
+ exponential:
+ decay_rate: 0.96
+ decay_steps: 231
+ initial_learning_rate: 0.5
+ name: 'ExponentialDecay'
+ offset: 0
+ staircase: true
+ steps_per_loop: 462
+ train_steps: 46200
+ validation_interval: 462
+ validation_steps: 20
diff --git a/official/projects/qat/vision/configs/experiments/retinanet/coco_mobilenetv3.5_avg_qat_tpu_e2e.yaml b/official/projects/qat/vision/configs/experiments/retinanet/coco_mobilenetv3.5_avg_qat_tpu_e2e.yaml
new file mode 100644
index 00000000000..573358a1c2c
--- /dev/null
+++ b/official/projects/qat/vision/configs/experiments/retinanet/coco_mobilenetv3.5_avg_qat_tpu_e2e.yaml
@@ -0,0 +1,74 @@
+# --experiment_type=retinanet_mobile_coco_qat
+# --topology=4x4
+# --tpu_platform=df
+# COCO mAP: 24.43 from QAT training and 23.1 from the TFLite after conversion.
+# QAT only supports float32 tpu due to fake-quant op.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ losses:
+ l2_weight_decay: 0.0
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetMultiAVG'
+ filter_size_scale: 1.0
+ type: 'mobilenet'
+ decoder:
+ type: 'fpn'
+ fpn:
+ num_filters: 128
+ use_separable_conv: true
+ use_keras_layer: true
+ head:
+ num_convs: 4
+ num_filters: 128
+ use_separable_conv: true
+ input_size: [256, 256, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'relu6'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'float32'
+ global_batch_size: 256
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ dtype: 'float32'
+ global_batch_size: 256
+ is_training: false
+ drop_remainder: false
+ quantization:
+ pretrained_original_checkpoint: 'gs://**/coco_mobilenetv3.5_avg_mobile_tpu/ckpt-277200'
+ quantize_detection_decoder: true
+ quantize_detection_head: true
+trainer:
+ best_checkpoint_eval_metric: AP
+ best_checkpoint_export_subdir: best_ckpt
+ best_checkpoint_metric_comp: higher
+ optimizer_config:
+ learning_rate:
+ type: 'exponential'
+ exponential:
+ decay_rate: 0.96
+ decay_steps: 231
+ initial_learning_rate: 0.5
+ name: 'ExponentialDecay'
+ offset: 0
+ staircase: true
+ steps_per_loop: 462
+ train_steps: 46200
+ validation_interval: 462
+ validation_steps: 20
diff --git a/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_gpu.yaml b/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_gpu.yaml
index f3c360d2644..3bfcfb57d33 100644
--- a/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_gpu.yaml
+++ b/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_gpu.yaml
@@ -1,4 +1,4 @@
-# --experiment_type=retinanet_spinenet_mobile_coco_qat
+# --experiment_type=retinanet_mobile_coco_qat
runtime:
distribution_strategy: 'mirrored'
mixed_precision_dtype: 'float32'
diff --git a/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_tpu.yaml b/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_tpu.yaml
new file mode 100644
index 00000000000..0ce3d546210
--- /dev/null
+++ b/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_tpu.yaml
@@ -0,0 +1,66 @@
+# --experiment_type=retinanet_mobile_coco_qat
+# COCO mAP: 24.7
+# QAT only supports float32 tpu due to fake-quant op.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ losses:
+ l2_weight_decay: 3.0e-05
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ spinenet_mobile:
+ stochastic_depth_drop_rate: 0.2
+ model_id: '49'
+ se_ratio: 0.2
+ use_keras_upsampling_2d: true
+ type: 'spinenet_mobile'
+ decoder:
+ type: 'identity'
+ head:
+ num_convs: 4
+ num_filters: 48
+ use_separable_conv: true
+ input_size: [384, 384, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'float32'
+ global_batch_size: 128
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ dtype: 'float32'
+ global_batch_size: 16
+ is_training: false
+ quantization:
+ pretrained_original_checkpoint: 'gs://**/coco_spinenet49_mobile_tpu_33884721/ckpt-277200'
+trainer:
+ checkpoint_interval: 924
+ optimizer_config:
+ learning_rate:
+ stepwise:
+ boundaries: [531300, 545160]
+ values: [0.0016, 0.00016, 0.000016]
+ type: 'stepwise'
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0000335
+ warmup_steps: 4000
+ steps_per_loop: 924
+ train_steps: 554400
+ validation_interval: 924
+ validation_steps: 1250
+ summary_interval: 924
diff --git a/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_tpu_e2e.yaml b/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_tpu_e2e.yaml
new file mode 100644
index 00000000000..b4374f9cd61
--- /dev/null
+++ b/official/projects/qat/vision/configs/experiments/retinanet/coco_spinenet49_mobile_qat_tpu_e2e.yaml
@@ -0,0 +1,67 @@
+# --experiment_type=retinanet_mobile_coco_qat
+# COCO mAP: 23.2
+# QAT only supports float32 tpu due to fake-quant op.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ losses:
+ l2_weight_decay: 3.0e-05
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ spinenet_mobile:
+ stochastic_depth_drop_rate: 0.2
+ model_id: '49'
+ se_ratio: 0.2
+ use_keras_upsampling_2d: true
+ type: 'spinenet_mobile'
+ decoder:
+ type: 'identity'
+ head:
+ num_convs: 4
+ num_filters: 48
+ use_separable_conv: true
+ input_size: [384, 384, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'float32'
+ global_batch_size: 256
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ dtype: 'float32'
+ global_batch_size: 16
+ is_training: false
+ quantization:
+ pretrained_original_checkpoint: 'gs://**/coco_spinenet49_mobile_tpu_33884721/ckpt-277200'
+ quantize_detection_head: true
+trainer:
+ checkpoint_interval: 462
+ optimizer_config:
+ learning_rate:
+ stepwise:
+ boundaries: [263340, 272580]
+ values: [0.032, 0.0032, 0.00032]
+ type: 'stepwise'
+ warmup:
+ linear:
+ warmup_learning_rate: 0.00067
+ warmup_steps: 2000
+ steps_per_loop: 462
+ train_steps: 277200
+ validation_interval: 462
+ validation_steps: 625
+ summary_interval: 924
diff --git a/official/projects/qat/vision/configs/image_classification.py b/official/projects/qat/vision/configs/image_classification.py
index 08e01cb65fb..66ce216947a 100644
--- a/official/projects/qat/vision/configs/image_classification.py
+++ b/official/projects/qat/vision/configs/image_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -35,6 +35,8 @@ def image_classification_imagenet() -> cfg.ExperimentConfig:
task = ImageClassificationTask.from_args(
quantization=common.Quantization(), **config.task.as_dict())
config.task = task
+ runtime = cfg.RuntimeConfig(enable_xla=False)
+ config.runtime = runtime
return config
diff --git a/official/projects/qat/vision/configs/image_classification_test.py b/official/projects/qat/vision/configs/image_classification_test.py
index 31208890d70..de637a3a63a 100644
--- a/official/projects/qat/vision/configs/image_classification_test.py
+++ b/official/projects/qat/vision/configs/image_classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for image_classification."""
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -40,7 +40,7 @@ def test_image_classification_configs(self, config_name):
self.assertIsInstance(config.task.quantization, common.Quantization)
self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
config.task.train_data.is_training = None
- with self.assertRaisesRegex(KeyError, 'Found inconsistncy between key'):
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
config.validate()
diff --git a/official/projects/qat/vision/configs/retinanet.py b/official/projects/qat/vision/configs/retinanet.py
index 91f2b528234..011bf7a1813 100644
--- a/official/projects/qat/vision/configs/retinanet.py
+++ b/official/projects/qat/vision/configs/retinanet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -28,8 +28,8 @@ class RetinaNetTask(retinanet.RetinaNetTask):
quantization: Optional[common.Quantization] = None
-@exp_factory.register_config_factory('retinanet_spinenet_mobile_coco_qat')
-def retinanet_spinenet_mobile_coco() -> cfg.ExperimentConfig:
+@exp_factory.register_config_factory('retinanet_mobile_coco_qat')
+def retinanet_mobile_coco() -> cfg.ExperimentConfig:
"""Generates a config for COCO OD RetinaNet for mobile with QAT."""
config = retinanet.retinanet_spinenet_mobile_coco()
task = RetinaNetTask.from_args(
diff --git a/official/projects/qat/vision/configs/retinanet_test.py b/official/projects/qat/vision/configs/retinanet_test.py
index 1a5ac79e8d8..d6c2c74f85d 100644
--- a/official/projects/qat/vision/configs/retinanet_test.py
+++ b/official/projects/qat/vision/configs/retinanet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for retinanet."""
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -28,7 +28,7 @@
class RetinaNetConfigTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters(
- ('retinanet_spinenet_mobile_coco_qat',),
+ ('retinanet_mobile_coco_qat',),
)
def test_retinanet_configs(self, config_name):
config = exp_factory.get_exp_config(config_name)
@@ -39,7 +39,7 @@ def test_retinanet_configs(self, config_name):
self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
config.validate()
config.task.train_data.is_training = None
- with self.assertRaisesRegex(KeyError, 'Found inconsistncy between key'):
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
config.validate()
diff --git a/official/projects/qat/vision/configs/semantic_segmentation.py b/official/projects/qat/vision/configs/semantic_segmentation.py
index 0bfe94b4549..fee621e861e 100644
--- a/official/projects/qat/vision/configs/semantic_segmentation.py
+++ b/official/projects/qat/vision/configs/semantic_segmentation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/qat/vision/configs/semantic_segmentation_test.py b/official/projects/qat/vision/configs/semantic_segmentation_test.py
index 58c1b54af45..36813fcc721 100644
--- a/official/projects/qat/vision/configs/semantic_segmentation_test.py
+++ b/official/projects/qat/vision/configs/semantic_segmentation_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for retinanet."""
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -39,7 +39,7 @@ def test_semantic_segmentation_configs(self, config_name):
self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
config.validate()
config.task.train_data.is_training = None
- with self.assertRaisesRegex(KeyError, 'Found inconsistncy between key'):
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
config.validate()
diff --git a/official/projects/qat/vision/docs/qat_tutorial.ipynb b/official/projects/qat/vision/docs/qat_tutorial.ipynb
new file mode 100644
index 00000000000..b1a634580b1
--- /dev/null
+++ b/official/projects/qat/vision/docs/qat_tutorial.ipynb
@@ -0,0 +1,3176 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cfjBLLxVb61A"
+ },
+ "source": [
+ "# Quantization-aware Training (QAT) for Object Detection with Model Garden\n",
+ "\n",
+ "This tutorial demonstrates how to apply [quantization-aware training (QAT)](https://github.com/tensorflow/models/tree/master/official/projects/qat/vision) from a pre-trained checkpoint, export the checkpoint to a TFLite and run inference on an image, for object detection task, using [Tensorflow Model Garden](https://github.com/tensorflow/models/tree/master/official) library.\n",
+ "\n",
+ "Tensorflow Model Garden contains a collection of\n",
+ "state-of-the-art models, implemented with TensorFlow's high-level APIs. The\n",
+ "implementations demonstrate the best practices for modeling, letting users to\n",
+ "take full advantage of TensorFlow for their research and product development.\n",
+ "\n",
+ "In this tutorial, we will use MobileNetV2 backbone with RetinaNet framework as an example to walk you through the process of applying QAT. This assumes you have already trained a model using Tensorflow Model Garden."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hmslpkDEcPk9"
+ },
+ "source": [
+ "## Install Necessary Dependencies"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "86HlYwGVWJ1s",
+ "outputId": "6872225b-eefd-4ef0-c59e-02c818d73923"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Looking in indexes: https://pypi.org/simple, https://us-python.pkg.dev/colab-wheels/public/simple/\n",
+ "Collecting tf-models-nightly\n",
+ " Downloading tf_models_nightly-2.11.0.dev20221221-py2.py3-none-any.whl (2.5 MB)\n",
+ "\u001b[K |████████████████████████████████| 2.5 MB 6.6 MB/s \n",
+ "\u001b[?25hCollecting sacrebleu\n",
+ " Downloading sacrebleu-2.3.1-py3-none-any.whl (118 kB)\n",
+ "\u001b[K |████████████████████████████████| 118 kB 42.8 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: oauth2client in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (4.1.3)\n",
+ "Requirement already satisfied: matplotlib in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (3.2.2)\n",
+ "Requirement already satisfied: gin-config in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (0.5.0)\n",
+ "Collecting tensorflow-text-nightly\n",
+ " Downloading tensorflow_text_nightly-2.12.0.dev20221221-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (5.9 MB)\n",
+ "\u001b[K |████████████████████████████████| 5.9 MB 24.9 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: six in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (1.15.0)\n",
+ "Requirement already satisfied: tensorflow-datasets in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (4.6.0)\n",
+ "Collecting tensorflow-model-optimization\u003e=0.4.1\n",
+ " Downloading tensorflow_model_optimization-0.7.3-py2.py3-none-any.whl (238 kB)\n",
+ "\u001b[K |████████████████████████████████| 238 kB 47.7 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: pandas\u003e=0.22.0 in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (1.3.5)\n",
+ "Requirement already satisfied: opencv-python-headless in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (4.6.0.66)\n",
+ "Requirement already satisfied: Pillow in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (7.1.2)\n",
+ "Requirement already satisfied: kaggle\u003e=1.3.9 in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (1.5.12)\n",
+ "Collecting immutabledict\n",
+ " Downloading immutabledict-2.2.3-py3-none-any.whl (4.0 kB)\n",
+ "Collecting seqeval\n",
+ " Downloading seqeval-1.2.2.tar.gz (43 kB)\n",
+ "\u001b[K |████████████████████████████████| 43 kB 1.6 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: tensorflow-hub\u003e=0.6.0 in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (0.12.0)\n",
+ "Requirement already satisfied: pycocotools in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (2.0.6)\n",
+ "Requirement already satisfied: google-api-python-client\u003e=1.6.7 in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (2.70.0)\n",
+ "Collecting tensorflow-addons\n",
+ " Downloading tensorflow_addons-0.19.0-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (1.1 MB)\n",
+ "\u001b[K |████████████████████████████████| 1.1 MB 42.0 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: scipy\u003e=0.19.1 in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (1.7.3)\n",
+ "Collecting tf-nightly\n",
+ " Downloading tf_nightly-2.12.0.dev20221221-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (557.8 MB)\n",
+ "\u001b[K |████████████████████████████████| 557.8 MB 8.3 kB/s \n",
+ "\u001b[?25hRequirement already satisfied: Cython in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (0.29.32)\n",
+ "Collecting tf-slim\u003e=1.1.0\n",
+ " Downloading tf_slim-1.1.0-py2.py3-none-any.whl (352 kB)\n",
+ "\u001b[K |████████████████████████████████| 352 kB 53.7 MB/s \n",
+ "\u001b[?25hCollecting py-cpuinfo\u003e=3.3.0\n",
+ " Downloading py_cpuinfo-9.0.0-py3-none-any.whl (22 kB)\n",
+ "Collecting pyyaml\u003c6.0,\u003e=5.1\n",
+ " Downloading PyYAML-5.4.1-cp38-cp38-manylinux1_x86_64.whl (662 kB)\n",
+ "\u001b[K |████████████████████████████████| 662 kB 58.0 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: psutil\u003e=5.4.3 in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (5.4.8)\n",
+ "Requirement already satisfied: numpy\u003e=1.20 in /usr/local/lib/python3.8/dist-packages (from tf-models-nightly) (1.21.6)\n",
+ "Collecting sentencepiece\n",
+ " Downloading sentencepiece-0.1.97-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl (1.3 MB)\n",
+ "\u001b[K |████████████████████████████████| 1.3 MB 72.4 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: google-auth-httplib2\u003e=0.1.0 in /usr/local/lib/python3.8/dist-packages (from google-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (0.1.0)\n",
+ "Requirement already satisfied: httplib2\u003c1dev,\u003e=0.15.0 in /usr/local/lib/python3.8/dist-packages (from google-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (0.17.4)\n",
+ "Requirement already satisfied: google-api-core!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.0,\u003c3.0.0dev,\u003e=1.31.5 in /usr/local/lib/python3.8/dist-packages (from google-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (2.11.0)\n",
+ "Requirement already satisfied: google-auth\u003c3.0.0dev,\u003e=1.19.0 in /usr/local/lib/python3.8/dist-packages (from google-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (2.15.0)\n",
+ "Requirement already satisfied: uritemplate\u003c5,\u003e=3.0.1 in /usr/local/lib/python3.8/dist-packages (from google-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (4.1.1)\n",
+ "Requirement already satisfied: requests\u003c3.0.0dev,\u003e=2.18.0 in /usr/local/lib/python3.8/dist-packages (from google-api-core!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.0,\u003c3.0.0dev,\u003e=1.31.5-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (2.23.0)\n",
+ "Requirement already satisfied: protobuf!=3.20.0,!=3.20.1,!=4.21.0,!=4.21.1,!=4.21.2,!=4.21.3,!=4.21.4,!=4.21.5,\u003c5.0.0dev,\u003e=3.19.5 in /usr/local/lib/python3.8/dist-packages (from google-api-core!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.0,\u003c3.0.0dev,\u003e=1.31.5-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (3.19.6)\n",
+ "Requirement already satisfied: googleapis-common-protos\u003c2.0dev,\u003e=1.56.2 in /usr/local/lib/python3.8/dist-packages (from google-api-core!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.0,\u003c3.0.0dev,\u003e=1.31.5-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (1.57.0)\n",
+ "Requirement already satisfied: rsa\u003c5,\u003e=3.1.4 in /usr/local/lib/python3.8/dist-packages (from google-auth\u003c3.0.0dev,\u003e=1.19.0-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (4.9)\n",
+ "Requirement already satisfied: cachetools\u003c6.0,\u003e=2.0.0 in /usr/local/lib/python3.8/dist-packages (from google-auth\u003c3.0.0dev,\u003e=1.19.0-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (5.2.0)\n",
+ "Requirement already satisfied: pyasn1-modules\u003e=0.2.1 in /usr/local/lib/python3.8/dist-packages (from google-auth\u003c3.0.0dev,\u003e=1.19.0-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (0.2.8)\n",
+ "Requirement already satisfied: python-dateutil in /usr/local/lib/python3.8/dist-packages (from kaggle\u003e=1.3.9-\u003etf-models-nightly) (2.8.2)\n",
+ "Requirement already satisfied: python-slugify in /usr/local/lib/python3.8/dist-packages (from kaggle\u003e=1.3.9-\u003etf-models-nightly) (7.0.0)\n",
+ "Requirement already satisfied: tqdm in /usr/local/lib/python3.8/dist-packages (from kaggle\u003e=1.3.9-\u003etf-models-nightly) (4.64.1)\n",
+ "Requirement already satisfied: certifi in /usr/local/lib/python3.8/dist-packages (from kaggle\u003e=1.3.9-\u003etf-models-nightly) (2022.12.7)\n",
+ "Requirement already satisfied: urllib3 in /usr/local/lib/python3.8/dist-packages (from kaggle\u003e=1.3.9-\u003etf-models-nightly) (1.24.3)\n",
+ "Requirement already satisfied: pytz\u003e=2017.3 in /usr/local/lib/python3.8/dist-packages (from pandas\u003e=0.22.0-\u003etf-models-nightly) (2022.7)\n",
+ "Requirement already satisfied: pyasn1\u003c0.5.0,\u003e=0.4.6 in /usr/local/lib/python3.8/dist-packages (from pyasn1-modules\u003e=0.2.1-\u003egoogle-auth\u003c3.0.0dev,\u003e=1.19.0-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (0.4.8)\n",
+ "Requirement already satisfied: chardet\u003c4,\u003e=3.0.2 in /usr/local/lib/python3.8/dist-packages (from requests\u003c3.0.0dev,\u003e=2.18.0-\u003egoogle-api-core!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.0,\u003c3.0.0dev,\u003e=1.31.5-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (3.0.4)\n",
+ "Requirement already satisfied: idna\u003c3,\u003e=2.5 in /usr/local/lib/python3.8/dist-packages (from requests\u003c3.0.0dev,\u003e=2.18.0-\u003egoogle-api-core!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.0,\u003c3.0.0dev,\u003e=1.31.5-\u003egoogle-api-python-client\u003e=1.6.7-\u003etf-models-nightly) (2.10)\n",
+ "Requirement already satisfied: dm-tree~=0.1.1 in /usr/local/lib/python3.8/dist-packages (from tensorflow-model-optimization\u003e=0.4.1-\u003etf-models-nightly) (0.1.8)\n",
+ "Requirement already satisfied: absl-py\u003e=0.2.2 in /usr/local/lib/python3.8/dist-packages (from tf-slim\u003e=1.1.0-\u003etf-models-nightly) (1.3.0)\n",
+ "Requirement already satisfied: cycler\u003e=0.10 in /usr/local/lib/python3.8/dist-packages (from matplotlib-\u003etf-models-nightly) (0.11.0)\n",
+ "Requirement already satisfied: pyparsing!=2.0.4,!=2.1.2,!=2.1.6,\u003e=2.0.1 in /usr/local/lib/python3.8/dist-packages (from matplotlib-\u003etf-models-nightly) (3.0.9)\n",
+ "Requirement already satisfied: kiwisolver\u003e=1.0.1 in /usr/local/lib/python3.8/dist-packages (from matplotlib-\u003etf-models-nightly) (1.4.4)\n",
+ "Requirement already satisfied: text-unidecode\u003e=1.3 in /usr/local/lib/python3.8/dist-packages (from python-slugify-\u003ekaggle\u003e=1.3.9-\u003etf-models-nightly) (1.3)\n",
+ "Collecting colorama\n",
+ " Downloading colorama-0.4.6-py2.py3-none-any.whl (25 kB)\n",
+ "Requirement already satisfied: tabulate\u003e=0.8.9 in /usr/local/lib/python3.8/dist-packages (from sacrebleu-\u003etf-models-nightly) (0.8.10)\n",
+ "Collecting portalocker\n",
+ " Downloading portalocker-2.6.0-py2.py3-none-any.whl (15 kB)\n",
+ "Requirement already satisfied: regex in /usr/local/lib/python3.8/dist-packages (from sacrebleu-\u003etf-models-nightly) (2022.6.2)\n",
+ "Requirement already satisfied: lxml in /usr/local/lib/python3.8/dist-packages (from sacrebleu-\u003etf-models-nightly) (4.9.2)\n",
+ "Requirement already satisfied: scikit-learn\u003e=0.21.3 in /usr/local/lib/python3.8/dist-packages (from seqeval-\u003etf-models-nightly) (1.0.2)\n",
+ "Requirement already satisfied: threadpoolctl\u003e=2.0.0 in /usr/local/lib/python3.8/dist-packages (from scikit-learn\u003e=0.21.3-\u003eseqeval-\u003etf-models-nightly) (3.1.0)\n",
+ "Requirement already satisfied: joblib\u003e=0.11 in /usr/local/lib/python3.8/dist-packages (from scikit-learn\u003e=0.21.3-\u003eseqeval-\u003etf-models-nightly) (1.2.0)\n",
+ "Requirement already satisfied: typeguard\u003e=2.7 in /usr/local/lib/python3.8/dist-packages (from tensorflow-addons-\u003etf-models-nightly) (2.7.1)\n",
+ "Requirement already satisfied: packaging in /usr/local/lib/python3.8/dist-packages (from tensorflow-addons-\u003etf-models-nightly) (21.3)\n",
+ "Requirement already satisfied: importlib-resources in /usr/local/lib/python3.8/dist-packages (from tensorflow-datasets-\u003etf-models-nightly) (5.10.1)\n",
+ "Requirement already satisfied: etils[epath] in /usr/local/lib/python3.8/dist-packages (from tensorflow-datasets-\u003etf-models-nightly) (0.9.0)\n",
+ "Requirement already satisfied: toml in /usr/local/lib/python3.8/dist-packages (from tensorflow-datasets-\u003etf-models-nightly) (0.10.2)\n",
+ "Requirement already satisfied: tensorflow-metadata in /usr/local/lib/python3.8/dist-packages (from tensorflow-datasets-\u003etf-models-nightly) (1.12.0)\n",
+ "Requirement already satisfied: dill in /usr/local/lib/python3.8/dist-packages (from tensorflow-datasets-\u003etf-models-nightly) (0.3.6)\n",
+ "Requirement already satisfied: termcolor in /usr/local/lib/python3.8/dist-packages (from tensorflow-datasets-\u003etf-models-nightly) (2.1.1)\n",
+ "Requirement already satisfied: promise in /usr/local/lib/python3.8/dist-packages (from tensorflow-datasets-\u003etf-models-nightly) (2.3)\n",
+ "Requirement already satisfied: typing_extensions in /usr/local/lib/python3.8/dist-packages (from etils[epath]-\u003etensorflow-datasets-\u003etf-models-nightly) (4.4.0)\n",
+ "Requirement already satisfied: zipp in /usr/local/lib/python3.8/dist-packages (from etils[epath]-\u003etensorflow-datasets-\u003etf-models-nightly) (3.11.0)\n",
+ "Collecting keras-nightly~=2.12.0.dev\n",
+ " Downloading keras_nightly-2.12.0.dev2022122108-py2.py3-none-any.whl (1.7 MB)\n",
+ "\u001b[K |████████████████████████████████| 1.7 MB 50.2 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: astunparse\u003e=1.6.0 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (1.6.3)\n",
+ "Requirement already satisfied: tensorflow-io-gcs-filesystem\u003e=0.23.1 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (0.29.0)\n",
+ "Collecting tf-estimator-nightly~=2.12.0.dev\n",
+ " Downloading tf_estimator_nightly-2.12.0.dev2022122109-py2.py3-none-any.whl (439 kB)\n",
+ "\u001b[K |████████████████████████████████| 439 kB 59.8 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: wrapt\u003e=1.11.0 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (1.14.1)\n",
+ "Requirement already satisfied: setuptools in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (57.4.0)\n",
+ "Requirement already satisfied: opt-einsum\u003e=2.3.2 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (3.3.0)\n",
+ "Requirement already satisfied: h5py\u003e=2.9.0 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (3.1.0)\n",
+ "Collecting flatbuffers\u003e=2.0\n",
+ " Downloading flatbuffers-22.12.6-py2.py3-none-any.whl (26 kB)\n",
+ "Requirement already satisfied: grpcio\u003c2.0,\u003e=1.24.3 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (1.51.1)\n",
+ "Requirement already satisfied: jax\u003e=0.3.15 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (0.3.25)\n",
+ "Requirement already satisfied: google-pasta\u003e=0.1.1 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (0.2.0)\n",
+ "Requirement already satisfied: gast\u003c=0.4.0,\u003e=0.2.1 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (0.4.0)\n",
+ "Requirement already satisfied: libclang\u003e=13.0.0 in /usr/local/lib/python3.8/dist-packages (from tf-nightly-\u003etf-models-nightly) (14.0.6)\n",
+ "Collecting tb-nightly~=2.12.0.a\n",
+ " Downloading tb_nightly-2.12.0a20221220-py3-none-any.whl (5.7 MB)\n",
+ "\u001b[K |████████████████████████████████| 5.7 MB 26.0 MB/s \n",
+ "\u001b[?25hRequirement already satisfied: wheel\u003c1.0,\u003e=0.23.0 in /usr/local/lib/python3.8/dist-packages (from astunparse\u003e=1.6.0-\u003etf-nightly-\u003etf-models-nightly) (0.38.4)\n",
+ "Requirement already satisfied: markdown\u003e=2.6.8 in /usr/local/lib/python3.8/dist-packages (from tb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (3.4.1)\n",
+ "Requirement already satisfied: tensorboard-data-server\u003c0.7.0,\u003e=0.6.0 in /usr/local/lib/python3.8/dist-packages (from tb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (0.6.1)\n",
+ "Requirement already satisfied: google-auth-oauthlib\u003c0.5,\u003e=0.4.1 in /usr/local/lib/python3.8/dist-packages (from tb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (0.4.6)\n",
+ "Requirement already satisfied: tensorboard-plugin-wit\u003e=1.6.0 in /usr/local/lib/python3.8/dist-packages (from tb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (1.8.1)\n",
+ "Requirement already satisfied: werkzeug\u003e=1.0.1 in /usr/local/lib/python3.8/dist-packages (from tb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (1.0.1)\n",
+ "Requirement already satisfied: requests-oauthlib\u003e=0.7.0 in /usr/local/lib/python3.8/dist-packages (from google-auth-oauthlib\u003c0.5,\u003e=0.4.1-\u003etb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (1.3.1)\n",
+ "Requirement already satisfied: importlib-metadata\u003e=4.4 in /usr/local/lib/python3.8/dist-packages (from markdown\u003e=2.6.8-\u003etb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (5.2.0)\n",
+ "Requirement already satisfied: oauthlib\u003e=3.0.0 in /usr/local/lib/python3.8/dist-packages (from requests-oauthlib\u003e=0.7.0-\u003egoogle-auth-oauthlib\u003c0.5,\u003e=0.4.1-\u003etb-nightly~=2.12.0.a-\u003etf-nightly-\u003etf-models-nightly) (3.2.2)\n",
+ "Building wheels for collected packages: seqeval\n",
+ " Building wheel for seqeval (setup.py) ... \u001b[?25l\u001b[?25hdone\n",
+ " Created wheel for seqeval: filename=seqeval-1.2.2-py3-none-any.whl size=16179 sha256=68dce998c3a0c033d5ad8dee4c0773f4f0caf8c6aed075f786b23a2a8f39b3af\n",
+ " Stored in directory: /root/.cache/pip/wheels/ad/5c/ba/05fa33fa5855777b7d686e843ec07452f22a66a138e290e732\n",
+ "Successfully built seqeval\n",
+ "Installing collected packages: tf-estimator-nightly, tb-nightly, portalocker, keras-nightly, flatbuffers, colorama, tf-slim, tf-nightly, tensorflow-text-nightly, tensorflow-model-optimization, tensorflow-addons, seqeval, sentencepiece, sacrebleu, pyyaml, py-cpuinfo, immutabledict, tf-models-nightly\n",
+ " Attempting uninstall: flatbuffers\n",
+ " Found existing installation: flatbuffers 1.12\n",
+ " Uninstalling flatbuffers-1.12:\n",
+ " Successfully uninstalled flatbuffers-1.12\n",
+ " Attempting uninstall: pyyaml\n",
+ " Found existing installation: PyYAML 6.0\n",
+ " Uninstalling PyYAML-6.0:\n",
+ " Successfully uninstalled PyYAML-6.0\n",
+ "\u001b[31mERROR: pip's dependency resolver does not currently take into account all the packages that are installed. This behaviour is the source of the following dependency conflicts.\n",
+ "tensorflow 2.9.2 requires flatbuffers\u003c2,\u003e=1.12, but you have flatbuffers 22.12.6 which is incompatible.\u001b[0m\n",
+ "Successfully installed colorama-0.4.6 flatbuffers-22.12.6 immutabledict-2.2.3 keras-nightly-2.12.0.dev2022122108 portalocker-2.6.0 py-cpuinfo-9.0.0 pyyaml-5.4.1 sacrebleu-2.3.1 sentencepiece-0.1.97 seqeval-1.2.2 tb-nightly-2.12.0a20221220 tensorflow-addons-0.19.0 tensorflow-model-optimization-0.7.3 tensorflow-text-nightly-2.12.0.dev20221221 tf-estimator-nightly-2.12.0.dev2022122109 tf-models-nightly-2.11.0.dev20221221 tf-nightly-2.12.0.dev20221221 tf-slim-1.1.0\n"
+ ]
+ }
+ ],
+ "source": [
+ "!pip install -U tf-models-nightly"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "WjAMpVVAcTYh"
+ },
+ "source": [
+ "## Import libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "T0JMx3hmcXL5"
+ },
+ "outputs": [],
+ "source": [
+ "import matplotlib.pyplot as plt\n",
+ "import matplotlib.patches as patches\n",
+ "from PIL import Image\n",
+ "\n",
+ "from six import BytesIO\n",
+ "from IPython import display\n",
+ "from urllib.request import urlopen\n",
+ "\n",
+ "import numpy as np\n",
+ "import tensorflow as tf\n",
+ "\n",
+ "import absl.logging\n",
+ "absl.logging.set_verbosity(absl.logging.ERROR)\n",
+ "tf.get_logger().setLevel(absl.logging.ERROR)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "7MOWG8h5eEh6"
+ },
+ "source": [
+ "## Download Pretrained Model\n",
+ "The model uses the implementation from the TensorFlow Model Garden GitHub repository, and achieves 23.3 mAP on COCO validation set. It uses a MobileNetV2 backbone and RetinaNet decoder on a 256x256 input image."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "WvTrNdt7sjh4",
+ "outputId": "bdcfb3c6-163d-41e6-fa27-ef50955a1b3d"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "100 18.8M 100 18.8M 0 0 32.4M 0 --:--:-- --:--:-- --:--:-- 32.4M\n"
+ ]
+ }
+ ],
+ "source": [
+ "! curl https://storage.googleapis.com/tf_model_garden/vision/qat/mobilenetv2_ssd_coco/mobilenetv2_ssd_i256_ckpt.tar.gz --output model.tar.gz"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "VEJf6ZTcsREl",
+ "outputId": "ae0a4187-62e9-4edc-e8b0-c256ece73279"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "mobilenetv2_ssd_i256_ckpt/\n",
+ "mobilenetv2_ssd_i256_ckpt/ckpt-277200.index\n",
+ "mobilenetv2_ssd_i256_ckpt/ckpt-277200.data-00000-of-00001\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Extract pretrained checkpoint.\n",
+ "! tar -xvzf model.tar.gz"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "xwD9Rd5EebNC"
+ },
+ "source": [
+ "## Launch QAT Training\n",
+ "\n",
+ "You can follow the [training guideline](https://github.com/tensorflow/models/tree/master/official/projects/qat/vision#training) to start QAT training using the pretrained checkpoint."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "n9cro-ivfo92"
+ },
+ "source": [
+ "## Export Model\n",
+ "\n",
+ "After QAT training completes, we can export a SavedModel and convert it to a TFLite model. For demonstration purpose only, we download a QAT trained model checkpoint and work on it."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "h2a4fgXkfn1m",
+ "outputId": "31a295c0-40fa-4315-84f5-0bdbc983beba"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "100 9822k 100 9822k 0 0 27.0M 0 --:--:-- --:--:-- --:--:-- 27.0M\n",
+ "mobilenetv2_ssd_i256_qat_ckpt/\n",
+ "mobilenetv2_ssd_i256_qat_ckpt/ckpt-1.data-00000-of-00001\n",
+ "mobilenetv2_ssd_i256_qat_ckpt/checkpoint\n",
+ "mobilenetv2_ssd_i256_qat_ckpt/ckpt-1.index\n"
+ ]
+ }
+ ],
+ "source": [
+ "! curl https://storage.googleapis.com/tf_model_garden/vision/qat/mobilenetv2_ssd_coco/mobilenetv2_ssd_i256_qat_ckpt.tar.gz --output model_qat.tar.gz\n",
+ "! tar -xvzf model_qat.tar.gz"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "8IlqYeivly68",
+ "outputId": "ba101be5-db20-4b9b-f4da-08059b988054"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "100 1836 100 1836 0 0 8784 0 --:--:-- --:--:-- --:--:-- 8784\n"
+ ]
+ }
+ ],
+ "source": [
+ "! curl https://raw.githubusercontent.com/tensorflow/models/master/official/projects/qat/vision/configs/experiments/retinanet/coco_mobilenetv2_qat_tpu_e2e.yaml --output params.yaml"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "ez_0U9XxhCXp",
+ "outputId": "75b6f602-1654-4446-9edd-aa993bc057d7"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "2022-12-21 18:12:44.763114: E tensorflow/tsl/lib/monitoring/collection_registry.cc:81] Cannot register 2 metrics with the same name: /tensorflow/core/bfc_allocator_delay\n",
+ "2022-12-21 18:12:45.870254: W tensorflow/tsl/platform/default/dso_loader.cc:66] Could not load dynamic library 'libnvinfer.so.8'; dlerror: libnvinfer.so.8: cannot open shared object file: No such file or directory; LD_LIBRARY_PATH: /usr/local/nvidia/lib:/usr/local/nvidia/lib64\n",
+ "2022-12-21 18:12:45.870397: W tensorflow/tsl/platform/default/dso_loader.cc:66] Could not load dynamic library 'libnvinfer_plugin.so.8'; dlerror: libnvinfer_plugin.so.8: cannot open shared object file: No such file or directory; LD_LIBRARY_PATH: /usr/local/nvidia/lib:/usr/local/nvidia/lib64\n",
+ "2022-12-21 18:12:45.870421: W tensorflow/compiler/tf2tensorrt/utils/py_utils.cc:38] TF-TRT Warning: Cannot dlopen some TensorRT libraries. If you would like to use Nvidia GPU with TensorRT, please make sure the missing libraries mentioned above are installed properly.\n",
+ "/usr/local/lib/python3.8/dist-packages/tensorflow_addons/utils/ensure_tf_install.py:37: UserWarning: You are currently using a nightly version of TensorFlow (2.12.0-dev20221221). \n",
+ "TensorFlow Addons offers no support for the nightly versions of TensorFlow. Some things might work, some other might not. \n",
+ "If you encounter a bug, do not file an issue on GitHub.\n",
+ " warnings.warn(\n",
+ "2022-12-21 18:12:48.949468: E tensorflow/compiler/xla/stream_executor/cuda/cuda_driver.cc:267] failed call to cuInit: CUDA_ERROR_NO_DEVICE: no CUDA-capable device is detected\n",
+ "I1221 18:12:48.975446 140428941322112 export_module.py:55] Set `nms_version` to `tflite` because only TFLite NMS is supported for QAT detection models.\n",
+ "I1221 18:12:48.984446 140428941322112 nn_layers.py:73] round_filter input=32 output=32\n",
+ "I1221 18:12:48.984663 140428941322112 nn_layers.py:73] round_filter input=16 output=16\n",
+ "I1221 18:12:48.984754 140428941322112 nn_layers.py:73] round_filter input=24 output=24\n",
+ "I1221 18:12:48.984833 140428941322112 nn_layers.py:73] round_filter input=24 output=24\n",
+ "I1221 18:12:48.984907 140428941322112 nn_layers.py:73] round_filter input=32 output=32\n",
+ "I1221 18:12:48.984983 140428941322112 nn_layers.py:73] round_filter input=32 output=32\n",
+ "I1221 18:12:48.985080 140428941322112 nn_layers.py:73] round_filter input=32 output=32\n",
+ "I1221 18:12:48.985161 140428941322112 nn_layers.py:73] round_filter input=64 output=64\n",
+ "I1221 18:12:48.985234 140428941322112 nn_layers.py:73] round_filter input=64 output=64\n",
+ "I1221 18:12:48.985307 140428941322112 nn_layers.py:73] round_filter input=64 output=64\n",
+ "I1221 18:12:48.985380 140428941322112 nn_layers.py:73] round_filter input=64 output=64\n",
+ "I1221 18:12:48.985452 140428941322112 nn_layers.py:73] round_filter input=96 output=96\n",
+ "I1221 18:12:48.985523 140428941322112 nn_layers.py:73] round_filter input=96 output=96\n",
+ "I1221 18:12:48.985596 140428941322112 nn_layers.py:73] round_filter input=96 output=96\n",
+ "I1221 18:12:48.985671 140428941322112 nn_layers.py:73] round_filter input=160 output=160\n",
+ "I1221 18:12:48.985743 140428941322112 nn_layers.py:73] round_filter input=160 output=160\n",
+ "I1221 18:12:48.985818 140428941322112 nn_layers.py:73] round_filter input=160 output=160\n",
+ "I1221 18:12:48.985891 140428941322112 nn_layers.py:73] round_filter input=320 output=320\n",
+ "I1221 18:12:48.985962 140428941322112 nn_layers.py:73] round_filter input=1280 output=1280\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:48.995921 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.129718 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.133578 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.391274 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.394566 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.398295 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.493044 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.496581 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.499755 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.600895 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.603896 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.607452 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.907559 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.911341 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:49.914487 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.012121 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.015271 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.018665 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.102202 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.105965 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.109401 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.206521 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.209712 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.213917 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.314715 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.317814 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.320856 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.409135 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.412137 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.415633 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.498831 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.501913 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.504915 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.585606 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.589559 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.593839 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.699263 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.702293 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.705178 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.791925 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.798440 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.801911 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.901241 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.905492 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:50.909714 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:51.023372 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:51.026516 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:51.030344 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:51.121911 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:51.126329 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:51.130333 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:51.226366 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "I1221 18:12:52.052204 140428941322112 fpn.py:113] FPN input_specs: {'2': TensorShape([None, 64, 64, 24]), '3': TensorShape([None, 32, 32, 32]), '4': TensorShape([None, 16, 16, 96]), '5': TensorShape([None, 8, 8, 320])}\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.238768 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.255487 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.273425 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.294029 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.312827 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "/usr/local/lib/python3.8/dist-packages/keras/engine/functional.py:639: UserWarning: Input dict contained keys ['6'] which did not match any model input. They will be ignored by the model.\n",
+ " inputs = self._flatten_to_reference_inputs(inputs)\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.494805 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.497647 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.500221 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.502812 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.504198 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.505835 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.506938 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.508035 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.509414 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.510470 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.511570 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.512733 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.514034 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.515157 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.517510 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.518876 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.520385 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.521717 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.522964 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.524257 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.532784 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.535722 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.539136 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.542675 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.544339 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.545793 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.547256 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.548651 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.550717 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.552119 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.553351 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.554597 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.556050 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.557350 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.558460 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.559547 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.560868 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.562187 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.587617 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:52.589650 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:From /usr/local/lib/python3.8/dist-packages/tensorflow/python/autograph/pyct/static_analysis/liveness.py:83: Analyzer.lamba_check (from tensorflow.python.autograph.pyct.static_analysis.liveness) is deprecated and will be removed after 2023-09-23.\n",
+ "Instructions for updating:\n",
+ "Lambda fuctions will be no more assumed to be used in the statement where they are used, or at least in the same block. https://github.com/tensorflow/tensorflow/issues/56089\n",
+ "W1221 18:12:52.893411 140428941322112 deprecation.py:364] From /usr/local/lib/python3.8/dist-packages/tensorflow/python/autograph/pyct/static_analysis/liveness.py:83: Analyzer.lamba_check (from tensorflow.python.autograph.pyct.static_analysis.liveness) is deprecated and will be removed after 2023-09-23.\n",
+ "Instructions for updating:\n",
+ "Lambda fuctions will be no more assumed to be used in the statement where they are used, or at least in the same block. https://github.com/tensorflow/tensorflow/issues/56089\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.740771 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.774617 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.779552 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.836686 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.839742 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.843206 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.910278 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.913603 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.916628 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.986388 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.989350 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:55.992302 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.074421 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.077560 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.081003 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.162846 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.166671 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.169923 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.249732 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.253354 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.256552 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.337864 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.342231 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.345453 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.418314 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.421329 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.424687 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.497306 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.500809 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.503766 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.590878 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.594089 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.597266 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.680670 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.683774 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.686938 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.768160 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.771155 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.774649 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.867106 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.870980 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.874137 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.969863 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.973033 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:56.976126 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:57.070321 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:57.073601 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:57.078219 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:57.176870 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:57.184587 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:57.187711 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:57.273781 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:58.401833 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:12:58.407577 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.389935 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.394536 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.399035 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.570006 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.574292 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.578619 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.767030 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.771076 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.781522 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.958673 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.962962 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:00.967329 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.147796 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.152103 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.156511 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.325911 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.330133 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.334604 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.495894 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.500783 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.505789 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.686200 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.690671 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.695453 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.876759 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.881086 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:01.885523 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.079122 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.083544 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.087835 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.244486 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.249272 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.253500 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.432092 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.436747 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.441804 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.642753 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.652113 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.661883 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.842828 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.848103 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:02.852962 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.060823 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.065197 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.069736 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.252587 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.256874 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.261666 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.430074 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:03.567693 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.018400 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.078130 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.082708 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.187595 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.192500 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.196663 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.354555 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.359011 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.363258 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.539405 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.543802 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.548494 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.699594 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.703820 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.709585 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.884591 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.889714 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:04.894985 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.071151 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.075779 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.080574 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.238726 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.243139 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.247942 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.447090 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.453935 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.458965 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.665561 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.670214 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.677138 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.858140 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.864031 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:05.869343 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.036916 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.042925 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.049704 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.232244 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.238185 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.244055 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.451255 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.457202 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.462138 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.627806 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.632334 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.636740 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.825390 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.829787 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:06.834579 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:07.310726 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:07.315327 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:07.319634 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:07.503761 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.263352 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.267443 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.270847 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.273896 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.279514 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.885777 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.888006 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.916746 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.920492 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:10.924909 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:12.329041 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:12.331942 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:12.334774 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:12.337657 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:12.340547 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.266652 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.270228 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.273658 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.276334 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.278174 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.279712 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.281238 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.282753 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.284448 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.286583 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.288877 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.290829 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.292556 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.294855 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.296359 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.297827 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.299448 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.301113 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.302600 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.304022 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.308149 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.310338 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.312530 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.317194 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.318946 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.320610 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.322926 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.324409 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.325994 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.327439 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.328929 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.330470 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.332253 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.333715 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.335189 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.336761 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.338433 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.339956 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.341529 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "WARNING:tensorflow:`tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "W1221 18:13:14.343154 140428941322112 batch_normalization.py:1426] `tf.keras.layers.experimental.SyncBatchNormalization` endpoint is deprecated and will be removed in a future release. Please use `tf.keras.layers.BatchNormalization` with parameter `synchronized` set to True.\n",
+ "/usr/local/lib/python3.8/dist-packages/keras/engine/functional.py:639: UserWarning: Input dict contained keys ['6'] which did not match any model input. They will be ignored by the model.\n",
+ " inputs = self._flatten_to_reference_inputs(inputs)\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79dd318e0\u003e, because it is not built.\n",
+ "W1221 18:13:42.789539 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79dd318e0\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79dcbdaf0\u003e, because it is not built.\n",
+ "W1221 18:13:43.138483 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79dcbdaf0\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79dccf130\u003e, because it is not built.\n",
+ "W1221 18:13:44.155304 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79dccf130\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79db466d0\u003e, because it is not built.\n",
+ "W1221 18:13:45.123558 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79db466d0\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79d958190\u003e, because it is not built.\n",
+ "W1221 18:13:46.430141 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79d958190\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79d88b730\u003e, because it is not built.\n",
+ "W1221 18:13:47.421999 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79d88b730\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79ddba3a0\u003e, because it is not built.\n",
+ "W1221 18:13:48.411185 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ctensorflow_model_optimization.python.core.quantization.keras.quantize_wrapper.QuantizeWrapperV2 object at 0x7fb79ddba3a0\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003cofficial.vision.modeling.layers.detection_generator.MultilevelDetectionGenerator object at 0x7fb79e562e80\u003e, because it is not built.\n",
+ "W1221 18:14:06.033631 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003cofficial.vision.modeling.layers.detection_generator.MultilevelDetectionGenerator object at 0x7fb79e562e80\u003e, because it is not built.\n",
+ "/usr/local/lib/python3.8/dist-packages/keras/engine/functional.py:639: UserWarning: Input dict contained keys ['6'] which did not match any model input. They will be ignored by the model.\n",
+ " inputs = self._flatten_to_reference_inputs(inputs)\n",
+ "/usr/local/lib/python3.8/dist-packages/keras/engine/functional.py:639: UserWarning: Input dict contained keys ['6'] which did not match any model input. They will be ignored by the model.\n",
+ " inputs = self._flatten_to_reference_inputs(inputs)\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79dd31c10\u003e, because it is not built.\n",
+ "W1221 18:15:39.459379 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79dd31c10\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79dcbd250\u003e, because it is not built.\n",
+ "W1221 18:15:39.488603 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79dcbd250\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79dccdd30\u003e, because it is not built.\n",
+ "W1221 18:15:39.545520 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79dccdd30\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79db46310\u003e, because it is not built.\n",
+ "W1221 18:15:39.632600 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79db46310\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79d954d90\u003e, because it is not built.\n",
+ "W1221 18:15:39.760443 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79d954d90\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79d88b370\u003e, because it is not built.\n",
+ "W1221 18:15:39.858972 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79d88b370\u003e, because it is not built.\n",
+ "WARNING:tensorflow:Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79ddbd190\u003e, because it is not built.\n",
+ "W1221 18:15:39.961016 140428941322112 save_impl.py:66] Skipping full serialization of Keras layer \u003ckeras.layers.merging.add.Add object at 0x7fb79ddbd190\u003e, because it is not built.\n",
+ "W1221 18:15:51.893777 140428941322112 save.py:272] Found untraced functions such as inference_from_image_bytes, inference_from_image_tensors, inference_from_tf_example, quant_activation_36_layer_call_fn, quant_activation_36_layer_call_and_return_conditional_losses while saving (showing 5 of 891). These functions will not be directly callable after loading.\n",
+ "INFO:tensorflow:Assets written to: /content/mobilenetv2_ssd_i256_qat_savedmodel/saved_model/assets\n",
+ "I1221 18:16:40.766796 140428941322112 builder_impl.py:797] Assets written to: /content/mobilenetv2_ssd_i256_qat_savedmodel/saved_model/assets\n",
+ "I1221 18:16:47.671683 140428941322112 train_utils.py:371] Saving experiment configuration to /content/mobilenetv2_ssd_i256_qat_savedmodel/params.yaml\n",
+ "2022-12-21 18:16:55.190074: E tensorflow/tsl/lib/monitoring/collection_registry.cc:81] Cannot register 2 metrics with the same name: /tensorflow/core/bfc_allocator_delay\n",
+ "2022-12-21 18:16:56.394554: W tensorflow/tsl/platform/default/dso_loader.cc:66] Could not load dynamic library 'libnvinfer.so.8'; dlerror: libnvinfer.so.8: cannot open shared object file: No such file or directory; LD_LIBRARY_PATH: /usr/local/nvidia/lib:/usr/local/nvidia/lib64\n",
+ "2022-12-21 18:16:56.394805: W tensorflow/tsl/platform/default/dso_loader.cc:66] Could not load dynamic library 'libnvinfer_plugin.so.8'; dlerror: libnvinfer_plugin.so.8: cannot open shared object file: No such file or directory; LD_LIBRARY_PATH: /usr/local/nvidia/lib:/usr/local/nvidia/lib64\n",
+ "2022-12-21 18:16:56.394834: W tensorflow/compiler/tf2tensorrt/utils/py_utils.cc:38] TF-TRT Warning: Cannot dlopen some TensorRT libraries. If you would like to use Nvidia GPU with TensorRT, please make sure the missing libraries mentioned above are installed properly.\n",
+ "/usr/local/lib/python3.8/dist-packages/tensorflow_addons/utils/ensure_tf_install.py:37: UserWarning: You are currently using a nightly version of TensorFlow (2.12.0-dev20221221). \n",
+ "TensorFlow Addons offers no support for the nightly versions of TensorFlow. Some things might work, some other might not. \n",
+ "If you encounter a bug, do not file an issue on GitHub.\n",
+ " warnings.warn(\n",
+ "2022-12-21 18:16:59.562431: E tensorflow/compiler/xla/stream_executor/cuda/cuda_driver.cc:267] failed call to cuInit: CUDA_ERROR_NO_DEVICE: no CUDA-capable device is detected\n",
+ "I1221 18:16:59.582286 140086576433024 export_tflite.py:102] Converting SavedModel from /content/mobilenetv2_ssd_i256_qat_savedmodel/saved_model to TFLite model...\n",
+ "W1221 18:17:01.833969 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:01.861936 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:01.865788 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:01.869644 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:01.873024 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:01.876166 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:01.879133 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:01.882937 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.052915 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.058312 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.061555 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.064986 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.068034 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.071142 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.074179 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.077957 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.081155 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.084321 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.114849 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.140550 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.145730 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.150812 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.155328 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.171447 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.184572 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.198348 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.201503 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.204597 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.207746 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.210942 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.214095 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.230438 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.234662 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.238849 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.242617 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.246194 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.259609 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.263310 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.266796 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.270354 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.294215 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.313850 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.323557 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.326894 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.335715 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.339398 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.342511 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.345713 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.349169 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.352294 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.355502 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.358910 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.362264 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.365817 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.369478 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.373200 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.376570 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.381706 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.384860 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.398141 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.401365 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.404412 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.407401 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.410552 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.413597 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.416796 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.420989 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.426159 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.430129 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.433763 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.436969 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.440397 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.448548 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.502177 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.505624 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.508775 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.511936 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.515357 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.518662 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.521887 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.525193 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.528397 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.531902 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.535319 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.538745 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.548315 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.564213 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.568897 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.572901 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.576473 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.605628 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.608869 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.611902 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.614970 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.617921 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.621021 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.623998 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.627149 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.630239 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.643897 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.647306 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.657177 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.670082 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.673666 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.676617 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.680211 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.683294 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.719151 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.722737 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.726030 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.729394 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.732691 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.748246 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.782508 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.786039 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.789324 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.795159 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.798289 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.801659 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.805912 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.810325 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.845919 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.849429 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.852868 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.855949 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.858876 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.861700 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.864760 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.908745 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.912907 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.916320 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.919600 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.935488 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.938981 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.942121 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.945243 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.948584 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.953220 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.957818 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.961713 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.964847 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.968117 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.971222 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:03.974279 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.107447 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.110886 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.114997 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.119149 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.147309 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.151013 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.154567 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.159224 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.163697 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.168208 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.173471 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.177752 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.181086 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.202994 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.206389 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.209619 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.212644 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.233421 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.247704 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.251548 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.254912 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.258786 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.262451 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.266170 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.269505 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.273328 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.276695 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.279961 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.282993 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.288027 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.291179 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.294162 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.297415 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.322781 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.326290 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.329427 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.332554 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.335563 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.338658 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.341754 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.344726 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.347718 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.350617 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.353471 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.356332 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.359286 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.361908 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.364699 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.375616 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.379907 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.384355 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.404713 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.407927 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.410848 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.413902 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.422914 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.432307 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.435856 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.438938 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.442582 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.466192 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.477117 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.480378 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.483490 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.490653 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.506877 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.510191 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.513490 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.526803 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.536736 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.540112 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.543606 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.546913 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.550118 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.553213 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.556227 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.569894 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.573045 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.576130 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.584743 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.596255 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.599708 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.602864 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.605781 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.608765 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.611877 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.620948 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.624305 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.627453 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.630741 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.633754 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.637156 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.640164 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.643054 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.652145 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.660982 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.665268 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.668302 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.688696 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.692344 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.709217 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.712594 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.715839 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.741044 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.745321 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.748711 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.751701 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.755128 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.758254 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.761297 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.764314 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.767467 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.779908 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.783169 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.786270 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.789369 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.793402 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.809949 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.813189 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.816168 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.828976 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.832124 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.847917 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.862906 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.866246 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.869312 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.904531 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.908685 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.912815 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.917125 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.920877 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.925222 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.929171 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.933343 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.937471 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.940481 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.943409 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.946587 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.976125 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.979541 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.991085 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:04.995532 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.027043 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.035886 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.039024 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.043981 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.046830 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.050261 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.063824 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.067102 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.070440 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.073659 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.076789 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.080468 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.096276 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.099544 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.103052 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.106530 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.138786 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.142378 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.158682 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.162131 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.165479 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.168768 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.172398 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.176745 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.181904 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.186816 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.194457 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.203829 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.207687 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.213747 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.223487 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.227216 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.230774 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.235475 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.239763 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.243802 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.247493 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.253672 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.256967 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.261670 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.266303 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.276698 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.280329 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.283785 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.302742 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.317409 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.320932 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.330446 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.334429 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.338118 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.341509 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.367841 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.371371 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.374648 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.377922 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.381251 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.384517 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.389691 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.403438 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.406856 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.410313 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.416514 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.429830 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.433140 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.436570 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.439809 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.443114 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.446361 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.449738 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.474654 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.744011 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.747428 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.750843 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.754172 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.757675 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.761175 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.764672 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.768120 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.771256 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.797747 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.801266 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.804727 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.819864 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.834294 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.838075 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.842591 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.847512 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:05.851940 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.116385 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.126725 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.131955 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.147536 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.150867 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.154092 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.157359 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.160727 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.163857 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.167242 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.170463 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.191684 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.195420 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.199023 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.203723 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.208950 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.214232 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.219646 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.224884 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.228929 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.233477 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.237176 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.247684 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.251557 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.281649 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.286052 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.290530 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.305075 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.308497 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.311787 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.315222 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.319722 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.323191 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.326408 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.329697 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.333119 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.340851 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.343919 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.354887 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.358145 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.361649 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.370978 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.374388 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.377813 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.381133 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.384377 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.397868 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.409268 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.414274 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.419613 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.426602 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.430803 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.434582 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.438274 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.449470 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.452875 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.456176 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.459314 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.487223 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.514169 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.517635 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.521114 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.524713 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.527928 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.541170 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.544476 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.547725 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.551100 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.554398 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.557494 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.560993 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.564421 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.568038 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.571249 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.580783 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.584342 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.587789 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.590912 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.594242 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.604251 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.607754 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.611104 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.614663 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.617954 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.621334 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.624505 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.633766 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.637281 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.640666 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.697752 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.701835 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.719791 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.725742 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.746177 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.749778 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.753386 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.757148 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.760676 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.764390 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.791980 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.796098 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.800871 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.816631 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.827705 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.831426 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.848267 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.852104 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.855429 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.858717 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.861981 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.865283 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.868603 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.871844 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.875119 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.878321 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.881863 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.885080 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.888415 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.943252 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.946881 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.950295 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.953680 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.957105 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.960235 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:06.963521 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.000750 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.004267 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.008517 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.011874 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.015257 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.018533 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.021870 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.042692 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.046239 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.049512 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.052748 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.056222 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.059809 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.063688 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.100050 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.106213 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.112806 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.117253 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.125887 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.129535 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.133166 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.136979 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.170974 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.192716 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.196619 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.200316 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.203691 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.221095 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.226936 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.258171 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.261868 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.265200 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.278019 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.281431 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.284940 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.290034 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.293513 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.297159 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.301320 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.304610 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.307956 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.311109 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.314416 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.317967 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.340984 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.344346 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.347555 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.420205 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.425553 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.431327 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.435126 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.438396 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.454046 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.457379 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.460873 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.464310 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.467842 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.471678 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.475241 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.699171 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.702976 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.710402 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.713936 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.717756 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.724472 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.765430 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.791969 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.795641 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.799168 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.802536 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.805986 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.809114 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.812291 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.815780 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.819712 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.829185 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.832834 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.836369 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.839878 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.843221 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.865881 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.869240 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.874382 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.877667 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.881077 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.884175 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.887381 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.890818 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.894042 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.898533 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.928145 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.931679 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.935016 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.938378 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.952089 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.955767 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.959110 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.974690 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.980280 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.992444 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:07.997846 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.001766 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.016292 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.056184 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.060419 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.064239 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.067873 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.080268 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.084681 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.099530 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.109639 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.113098 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.116702 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.120901 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.124510 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.127958 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.131857 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.137104 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.147616 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.151750 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.155249 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.159149 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.165153 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.179173 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.184626 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.189784 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.237135 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.240974 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.244683 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.249735 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.278170 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.288978 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.313314 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.316711 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.356050 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.359699 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.363184 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.366510 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.369985 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.373274 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.376762 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.382484 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.385872 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.389365 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.392619 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.395927 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.399097 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.402358 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.438320 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.442040 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.469902 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.473126 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.576879 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.580506 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.583993 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.631122 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.636280 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.641656 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.645222 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.648649 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.652125 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.655687 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.658931 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.662459 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.665767 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.669183 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.672621 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.676161 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.679719 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.683626 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.688252 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.691667 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.695404 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.714014 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.717648 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.721372 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.731693 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.735987 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.740943 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.744492 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.748081 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.751753 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.761792 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.765272 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.768588 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.771529 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.792488 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.795797 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.799427 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.802870 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.806282 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.828375 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.842153 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.845905 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.850255 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.853610 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.857072 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.860489 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.863839 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.885736 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.889615 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.893055 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.933040 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.937955 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.944232 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.957445 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.961934 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.965838 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.969466 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.972985 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.986599 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.989947 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.993541 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:08.997099 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.001623 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.006665 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.011180 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.014971 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.018563 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.022289 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.026139 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.030132 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.033631 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.054745 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.058344 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.061694 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.065016 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.094580 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.108342 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.139621 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.145189 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.150313 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.155238 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.159265 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.162876 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.178361 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.182292 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.185814 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.195262 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.198665 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.202097 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.205284 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.208463 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.218022 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.221660 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.225235 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.230182 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.236034 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.243700 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.248043 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.251512 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.289870 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.295206 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.299448 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.304211 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.313689 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.317157 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.320791 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.336902 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.340455 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.343843 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.347224 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.350859 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.354276 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.357615 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.360878 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.364186 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.367591 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.370823 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.388605 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.398191 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.407805 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.416774 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.447537 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.451026 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.454380 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.457751 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.461126 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.464739 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.478188 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.481849 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.506472 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.510011 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.521458 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.562007 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.565765 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.569715 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.575752 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.579237 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.582800 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.586099 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.608114 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.617295 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.620908 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.633984 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.637388 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.640927 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.644096 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.647571 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.650965 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.654139 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.657228 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.660479 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.663810 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.666944 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.670170 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.673820 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.677191 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.690985 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.709991 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.714411 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.757736 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.761470 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.765824 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.769329 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.772712 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.776100 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.779682 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.782963 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.792534 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.796229 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.799595 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.802953 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.806600 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.810155 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.828525 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.834909 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.841656 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.866725 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.880411 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.883872 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:09.887470 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.380246 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.384031 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.387405 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.390568 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.393811 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.396743 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.399626 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.402737 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.406118 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.409373 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.412671 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.416128 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.447284 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.450841 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.454134 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.457281 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.460178 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.463365 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.466651 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.469816 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.490312 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.499675 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.503104 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.506135 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.509262 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.512225 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:11.515477 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.859961 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.864102 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.867522 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.870930 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.874287 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.877806 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.881302 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.884607 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.887863 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.892759 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.896772 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.932880 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.936589 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.940176 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.943874 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.947355 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.952220 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.958549 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.962620 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.966130 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.969499 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.973399 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.977267 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.980696 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.984004 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.987146 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.990390 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.993551 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:12.996796 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.000497 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.003879 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.028940 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.032582 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.063237 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.072974 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.104979 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.108300 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.111386 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.137033 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.406627 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.410323 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.413556 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.416704 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.419942 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.423282 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.426612 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.429895 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.445829 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.460580 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.463783 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.466830 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.475552 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.479035 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.488554 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.535811 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.539295 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.542571 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.559621 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.563085 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.566620 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.570357 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.583740 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.586701 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.589771 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.615017 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.618582 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.621965 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.625557 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.629014 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.656193 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.662812 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.693506 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.697512 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.702372 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.727816 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.731411 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.735275 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.740318 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.747999 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.751750 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.755823 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.759222 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.762349 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.765438 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.785982 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.804029 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.807533 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.810772 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.813855 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.817081 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.838980 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.863186 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.889398 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.906575 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.910286 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.981549 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.985255 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.988542 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.991670 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:13.994682 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.044481 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.048154 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.052955 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.058163 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.063627 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.068251 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.071775 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.088352 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.092599 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.117514 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.121092 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.124291 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.127571 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.155845 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.160341 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.164247 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.167918 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.171631 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.176335 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.189753 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.194213 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.216233 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.219878 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.247578 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.251151 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.261777 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.298707 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.308323 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.311790 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.315256 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.318626 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.325127 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.355377 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.372799 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.376256 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.379462 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.399880 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.403455 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.408412 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.413290 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.418487 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.429763 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.433157 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.436530 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.440209 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.443695 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.447009 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.450441 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.476763 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.480419 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.484207 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.487591 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.490741 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.493625 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.496446 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.499561 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:14.507075 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.650613 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.654544 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.657942 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.667723 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.717339 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.753307 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.756975 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.760371 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.764285 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.778933 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.782537 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.804883 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.808943 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.812377 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.815630 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.819078 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.845104 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.862405 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.865976 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.869420 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.914996 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.923331 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.926570 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.954534 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.957952 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.966401 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.989189 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:15.993260 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.010920 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.014562 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.017881 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.021469 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.024992 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.028378 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.031956 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.035501 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.039447 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.042986 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.046298 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.049777 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.053336 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.056737 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.060157 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.063496 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.066686 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.069981 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.073894 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.077928 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.083110 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.101417 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.106227 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.110207 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.113829 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.117427 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.127162 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.131602 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.137422 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.158755 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.162382 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.188891 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.192456 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.214292 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.218026 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.221681 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.225149 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.228744 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.232360 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.236602 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.240309 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.249967 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.253578 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.256854 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.278398 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.282218 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.286526 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.290588 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.300189 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.321446 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.325473 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.328968 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.332318 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.335642 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.340106 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.343753 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.347323 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.352087 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.356842 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.367917 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.832907 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.836878 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.875442 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.879554 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.894783 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.899507 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.915704 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.919184 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.922784 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.926071 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.929429 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.932909 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.954123 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.957621 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.966825 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.970100 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:16.973304 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.011782 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.015103 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.027839 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.031467 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.034857 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.038139 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.041603 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.044981 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.048582 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.075973 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.079706 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.084862 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.089626 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.095576 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.100528 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.117537 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.121253 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.124716 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.128036 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.131693 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.145826 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.173878 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.177869 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.182114 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.185548 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.189032 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_205932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.214442 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.218230 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.222048 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.225720 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.252039 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.281569 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.286227 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.317684 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.321407 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.324940 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.328202 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.331767 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.335497 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.745440 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.755185 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.776521 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.782672 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.786555 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.789676 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.793015 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.796393 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.799882 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.803161 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.812605 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.815901 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.840858 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.844508 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.848793 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.852081 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.877747 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.901267 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.916739 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.936423 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.965641 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.990335 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:17.995961 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.001501 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.005629 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.017251 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.038846 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.061704 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.096615 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.122194 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.149516 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.176313 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.180294 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.208218 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.212166 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.216342 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.220590 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.224512 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.228221 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.255647 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.283870 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.293432 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.298568 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.328976 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.347272 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.385216 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.390526 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.395360 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.401503 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219252) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.424666 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.428475 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.432905 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.437021 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.441464 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.468777 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.472658 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.498170 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.560678 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.564279 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.569800 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.573055 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.580288 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.583390 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.589074 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.618865 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.676215 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.698539 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.720730 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.724510 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.728030 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.731840 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.735733 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.741803 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.757672 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.778858 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.809658 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.812978 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.816548 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.820338 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.846379 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.868192 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.873206 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.902956 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.908740 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.913902 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.918790 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217002) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.925950 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.930484 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.935316 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.949949 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.954042 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.957440 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.960884 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.964178 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.991885 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.996299 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:18.999917 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.003411 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.030101 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.033831 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.037510 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.040903 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.044373 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.075994 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.118130 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.152720 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.156132 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.183839 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.195557 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.208549 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.211677 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.214920 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.219161 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.222868 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.226444 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.230100 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.233811 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.238306 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.271070 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.276725 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.289636 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.293103 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.297010 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.322201 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.325738 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.347548 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.387200 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.398613 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.432085 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.435412 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.438698 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.441800 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.444714 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.447696 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.450798 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.456239 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.482051 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.589859 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.594902 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.656869 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.660279 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.673463 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212652) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.677351 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.680887 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.708299 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.713903 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.718492 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.722109 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.725678 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.752430 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.756621 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.760170 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.763517 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.858235 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.861827 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.865223 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218302) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.868754 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.882656 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.886180 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210112) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.931889 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.945642 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.949256 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.952525 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:19.955875 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.058701 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.062221 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.161660 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.187965 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.213473 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.217771 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220832) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.221216 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.253943 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.257847 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.284482 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.376671 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219942) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.380404 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.383812 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.386985 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.390215 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.398900 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.402706 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.405975 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.419350 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218602) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.423054 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.445380 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.449004 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.452284 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.455751 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.458857 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.465581 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.468689 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.471742 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.474627 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.477546 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.480571 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.509932 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.514380 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.535500 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.538988 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.543703 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.576017 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.579215 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.582411 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.585538 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.588627 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.591545 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.625960 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.714346 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.736800 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.764368 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214742) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.767977 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.771270 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.796850 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.800310 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.899358 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.903876 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.952116 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.955687 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.959070 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:20.984297 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.016969 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.049997 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.133502 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.137189 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.167459 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.170924 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.186705 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.299217 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.330720 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.356075 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209122) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.382363 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215012) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.385750 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.416228 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.431223 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.436017 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.440994 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.445937 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.449652 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.453093 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.456395 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.459859 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.463176 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.484573 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209632) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.945703 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.949088 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.995489 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:21.998567 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.001747 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.004735 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.026091 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.047724 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.136078 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.139456 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.168763 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.171998 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.174903 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.199324 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.202543 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.205558 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.219938 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.320815 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211822) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.357006 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.360630 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.382521 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.386183 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.411787 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.447109 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211082) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.471663 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209522) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.475990 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.479561 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208322) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.505694 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.509307 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.512812 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.672244 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210692) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.698368 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.702395 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.815166 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.819513 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209332) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.842743 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.876266 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.969275 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:22.986627 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212212) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.088651 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.123626 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.131126 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.135137 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.161439 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220382) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.164948 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.188970 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.192306 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208092) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.195322 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.198230 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211432) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.246809 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.250227 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:23.341490 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.658573 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.663973 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.686860 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.690590 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.693896 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.719719 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.742932 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.746504 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:24.749776 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.084220 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.088300 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.110897 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219262) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.114678 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.118265 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.146471 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213402) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.150151 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.176784 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.206919 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.210271 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.299243 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.320614 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.348356 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206622) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.382872 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.386180 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.389336 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218922) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.392680 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.486292 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207532) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.523798 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214502) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.533243 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216272) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.538197 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211802) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.618787 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214592) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.622276 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211722) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.666480 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.686309 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.728825 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_221542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.734832 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.739776 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.780561 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.783869 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209222) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.807218 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.810728 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.836889 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210872) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.868515 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209512) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.889692 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207022) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.920125 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.923770 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219492) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.927197 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208062) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.955945 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:25.959595 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213052) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.037004 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209932) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.076018 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.102079 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.109168 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.299429 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.303213 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206202) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.324830 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.328349 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.400985 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214162) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.481688 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216812) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.485085 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.509008 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.569974 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220452) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.573508 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.576708 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206992) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.600915 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.632478 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.645693 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.650224 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.761604 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210882) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.789531 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215792) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.793088 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210042) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.814680 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214462) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.905212 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213442) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.909891 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_209662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.937836 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207892) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:26.964494 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211562) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.357815 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_219372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.361330 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220242) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.364870 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212732) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.479242 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.501867 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.523820 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.622606 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218132) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.626303 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.736762 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.742841 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216662) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.746523 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210192) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.773837 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208372) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.794400 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.797803 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213412) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.883451 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208962) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:27.912292 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_220902) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.534716 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210172) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.563454 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216232) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.566780 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.601331 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_217982) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.604281 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215712) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.607264 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215682) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.610337 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218102) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.613186 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207072) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.655866 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213482) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.772805 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.835780 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214782) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.927375 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.947834 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208362) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:31.969461 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.001708 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.025838 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218552) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.029020 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.049427 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.070763 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215752) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.074252 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.095528 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.120349 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206312) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.145288 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210642) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.166167 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207582) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.194713 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206702) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:32.215944 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_216342) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:33.844275 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_212772) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:33.957591 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215032) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:33.961204 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214672) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:34.044854 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_218152) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:35.419019 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214292) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:35.505406 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206352) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:35.545868 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:35.549408 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213862) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.036015 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208852) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.064692 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206182) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.091268 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_210542) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.112746 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206762) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.136169 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206972) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.162342 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211142) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.182882 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215912) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.186331 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_213612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.272706 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_211472) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:37.299903 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207422) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:40.093797 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_215612) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:41.648584 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_206282) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:41.870301 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214952) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:41.911100 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_207392) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:41.938111 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_214842) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "W1221 18:17:41.996572 140086576433024 function_deserialization.py:611] Importing a function (__inference_internal_grad_fn_208572) with ops with unsaved custom gradients. Will likely fail if a gradient is requested.\n",
+ "2022-12-21 18:17:53.603566: W tensorflow/compiler/mlir/lite/python/tf_tfl_flatbuffer_helpers.cc:363] Ignored output_format.\n",
+ "2022-12-21 18:17:53.603644: W tensorflow/compiler/mlir/lite/python/tf_tfl_flatbuffer_helpers.cc:366] Ignored drop_control_dependency.\n",
+ "W1221 18:18:06.386931 140086576433024 util.py:822] For model outputs containing unsupported operations which cannot be quantized, the `inference_output_type` attribute will default to the original type.\n",
+ "I1221 18:18:06.924709 140086576433024 export_tflite.py:119] TFLite model converted and saved to /content/mobilenetv2_ssd_i256_qat_tflite.\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Model export and convert.\n",
+ "# First export a SavedModel. Make sure batch_size=1 and input_type=tflite.\n",
+ "! python3 /usr/local/lib/python3.8/dist-packages/official/projects/qat/vision/serving/export_saved_model.py --experiment=retinanet_mobile_coco_qat --export_dir=${PWD}/mobilenetv2_ssd_i256_qat_savedmodel --checkpoint_path=${PWD}/mobilenetv2_ssd_i256_qat_ckpt --batch_size=1 --input_type=tflite --input_image_size=256,256 --config_file=${PWD}/params.yaml --params_override=\"task.quantization.pretrained_original_checkpoint='${PWD}/mobilenetv2_ssd_i256_ckpt/ckpt-277200'\"\n",
+ "\n",
+ "# Convert the SavedModel to TFLite\n",
+ "! python3 /usr/local/lib/python3.8/dist-packages/official/projects/qat/vision/serving/export_tflite.py --experiment=retinanet_mobile_coco_qat --saved_model_dir=${PWD}/mobilenetv2_ssd_i256_qat_savedmodel/saved_model --tflite_path=${PWD}/mobilenetv2_ssd_i256_qat_tflite --config_file=${PWD}/params.yaml --quant_type=qat"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Th_BzrD5eoPp"
+ },
+ "source": [
+ "## Run Inference\n",
+ "\n",
+ "Now we will show how to use the converted TFLite model to do inference and obtain detection results. We provide our converted TFLite model that can be directly used for this."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "ufltjgUHnae1",
+ "outputId": "b63ad6d3-9212-437d-fd60-a2271e7f7e20"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "100 3386k 100 3386k 0 0 7329k 0 --:--:-- --:--:-- --:--:-- 7329k\n"
+ ]
+ }
+ ],
+ "source": [
+ "# First download the TFLite model.\n",
+ "! curl https://storage.googleapis.com/tf_model_garden/vision/qat/mobilenetv2_ssd_coco/model_int8_qat.tflite --output model.tflite"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "X8ACyzMDdDR7"
+ },
+ "outputs": [],
+ "source": [
+ "# Defines helper function to download sample image.\n",
+ "def load_image_into_numpy_array(path, height, width):\n",
+ " \"\"\"Load an image from file into a numpy array.\n",
+ "\n",
+ " Puts image into numpy array to feed into tensorflow graph.\n",
+ " Note that by convention we put it into a numpy array with shape\n",
+ " (height, width, channels), where channels=3 for RGB.\n",
+ "\n",
+ " Args:\n",
+ " path: the file path to the image\n",
+ "\n",
+ " Returns:\n",
+ " uint8 numpy array with shape (height, width, 3)\n",
+ " \"\"\"\n",
+ " image = None\n",
+ " if(path.startswith('http')):\n",
+ " response = urlopen(path)\n",
+ " image_data = response.read()\n",
+ " image_data = BytesIO(image_data)\n",
+ " image = Image.open(image_data)\n",
+ " else:\n",
+ " image_data = tf.io.gfile.GFile(path, 'rb').read()\n",
+ " image = Image.open(BytesIO(image_data))\n",
+ "\n",
+ " (im_width, im_height) = image.size\n",
+ " image = image.resize((height, width))\n",
+ "\n",
+ " image = np.array(image.getdata()).reshape(\n",
+ " (1, height, width, 3)).astype(np.uint8)\n",
+ "\n",
+ " return image"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "n3BXyYune5p6"
+ },
+ "source": [
+ "### Download a Sample Image"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "D7KLJpIGdDVs",
+ "outputId": "163b84a5-6456-49c0-dd64-23903779432c"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "(1, 256, 256, 3)"
+ ]
+ },
+ "execution_count": 10,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "image_path = \"https://djl.ai/examples/src/test/resources/dog_bike_car.jpg\"\n",
+ "image_array = load_image_into_numpy_array(image_path, 256, 256)\n",
+ "image_array.shape"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "u2eMxCEfe83m"
+ },
+ "outputs": [],
+ "source": [
+ "Image.fromarray(image_array[0])"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "gn99qmZAf75S"
+ },
+ "source": [
+ "### Run Inference on Sample Image"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "gDdehkLeNeqz"
+ },
+ "outputs": [],
+ "source": [
+ "# Load TFLite model.\n",
+ "tflite_path = 'model.tflite'\n",
+ "\n",
+ "with tf.io.gfile.GFile(tflite_path, 'rb') as f:\n",
+ " tflite_model = f.read()\n",
+ "\n",
+ "interpreter = tf.lite.Interpreter(model_content=tflite_model)\n",
+ "interpreter.allocate_tensors()\n",
+ "\n",
+ "input_details = interpreter.get_input_details()\n",
+ "output_details = interpreter.get_output_details()\n",
+ "interpreter.set_tensor(input_details[0]['index'], image_array)\n",
+ "\n",
+ "interpreter.invoke()\n",
+ "\n",
+ "# The function `get_tensor()` returns a copy of the tensor data.\n",
+ "# Use `tensor()` in order to get a pointer to the tensor.\n",
+ "outputs = []\n",
+ "for i in range(len(output_details)):\n",
+ " outputs.append(interpreter.get_tensor(output_details[i]['index']))"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "lR-TaAFaoUV_",
+ "outputId": "1779640d-bd7e-4ede-bd44-3bd2d1b56fb0"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "[array([[[ 1.06062141e+02, 3.99226303e+01, 2.37835205e+02,\n",
+ " 1.03872078e+02],\n",
+ " [ 3.33423462e+01, 1.52555298e+02, 7.61054840e+01,\n",
+ " 2.30757751e+02],\n",
+ " [ 5.80201950e+01, 5.14244308e+01, 1.93264526e+02,\n",
+ " 1.90765381e+02],\n",
+ " [ 3.17780743e+01, 1.97641087e+01, 5.62219238e+01,\n",
+ " 3.72487183e+01],\n",
+ " [ 3.46833572e+01, 1.89357872e+01, 5.47490005e+01,\n",
+ " 2.96583672e+01],\n",
+ " [ 4.58791046e+01, 1.97192001e+01, 2.30809067e+02,\n",
+ " 2.28094421e+02],\n",
+ " [ 7.44852142e+01, 3.73519745e+01, 2.39239090e+02,\n",
+ " 1.48372314e+02],\n",
+ " [ 4.89471016e+01, 2.19980789e+02, 6.96205444e+01,\n",
+ " 2.38883926e+02],\n",
+ " [ 6.08561096e+01, 4.66774597e+01, 9.56913452e+01,\n",
+ " 8.47750854e+01],\n",
+ " [ 5.58034134e+00, 5.37203522e+01, 9.11066132e+01,\n",
+ " 1.00488342e+02],\n",
+ " [ 3.71521416e+01, 2.12559471e+01, 4.48743706e+01,\n",
+ " 2.73382072e+01],\n",
+ " [ 1.20761276e+02, 2.29318695e+02, 1.91238724e+02,\n",
+ " 2.55465759e+02],\n",
+ " [ 1.45676407e+02, 6.96358109e+00, 2.41676407e+02,\n",
+ " 1.11954163e+02],\n",
+ " [ 3.64898758e+01, 1.25909805e+02, 6.94331360e+01,\n",
+ " 2.02398117e+02],\n",
+ " [ 3.14057674e+01, 1.68498611e+01, 7.25942307e+01,\n",
+ " 4.07579079e+01],\n",
+ " [ 4.34987106e+01, 1.26912971e+02, 6.47983627e+01,\n",
+ " 1.57384125e+02],\n",
+ " [-2.49016190e+00, 3.00395966e-01, 1.09385818e+02,\n",
+ " 2.84604702e+01],\n",
+ " [ 8.83496094e+00, 2.37918518e+02, 1.43678268e+02,\n",
+ " 2.56055817e+02],\n",
+ " [ 1.94364510e+01, 1.87362534e+02, 2.68458252e+01,\n",
+ " 1.93701614e+02],\n",
+ " [ 2.89923038e+01, 2.71423073e+01, 5.16018524e+01,\n",
+ " 4.34253349e+01],\n",
+ " [ 4.56293755e+01, 2.00866791e+02, 5.65502586e+01,\n",
+ " 2.09158875e+02],\n",
+ " [ 5.19470749e+01, 2.27266525e+02, 6.62325592e+01,\n",
+ " 2.33797623e+02],\n",
+ " [ 5.07977371e+01, 2.30534622e+02, 6.73818970e+01,\n",
+ " 2.40452560e+02],\n",
+ " [ 5.01660538e+01, 2.24255310e+02, 7.18596115e+01,\n",
+ " 2.34783173e+02],\n",
+ " [ 2.79034042e+01, 6.99869690e+01, 8.81479187e+01,\n",
+ " 1.01961700e+02],\n",
+ " [ 2.61520538e+01, 1.58735535e+02, 2.61847961e+02,\n",
+ " 2.57264465e+02],\n",
+ " [ 2.64463615e+01, 1.77282410e+01, 5.16306267e+01,\n",
+ " 2.53102493e+01],\n",
+ " [ 4.11829567e+01, 1.87375717e+02, 4.90480003e+01,\n",
+ " 1.92675613e+02],\n",
+ " [ 5.87513504e+01, 2.12607803e+02, 7.52229843e+01,\n",
+ " 2.39340866e+02],\n",
+ " [ 4.74041977e+01, 2.23578598e+02, 1.04595802e+02,\n",
+ " 2.54852463e+02],\n",
+ " [-6.53881073e-01, 5.67382431e+01, 8.52013397e+01,\n",
+ " 1.37619385e+02],\n",
+ " [ 4.30309334e+01, 1.79960037e+02, 4.92256813e+01,\n",
+ " 1.85104111e+02],\n",
+ " [ 4.88146706e+01, 1.93546463e+02, 5.74162865e+01,\n",
+ " 2.00479202e+02],\n",
+ " [ 5.05325356e+01, 2.27127548e+02, 5.77240791e+01,\n",
+ " 2.32923782e+02],\n",
+ " [ 4.48982162e+01, 1.91046005e+02, 7.34940186e+01,\n",
+ " 2.36091843e+02],\n",
+ " [ 1.93680325e+01, 1.21709106e+02, 6.42412872e+01,\n",
+ " 1.78681580e+02],\n",
+ " [ 4.34505081e+00, 2.19482620e+02, 1.26698425e+02,\n",
+ " 2.52173904e+02],\n",
+ " [ 6.50462646e+01, 4.61425095e+01, 1.48058075e+02,\n",
+ " 9.02009659e+01],\n",
+ " [ 1.28264740e+02, -2.90789032e+00, 2.50982025e+02,\n",
+ " 1.94907898e+02],\n",
+ " [ 2.55440331e+01, 1.87228508e+01, 5.55508766e+01,\n",
+ " 4.87296944e+01],\n",
+ " [ 4.29371147e+01, 1.29693497e+02, 4.93195000e+01,\n",
+ " 1.35319321e+02],\n",
+ " [ 4.99539413e+01, 2.33106995e+02, 5.83026733e+01,\n",
+ " 2.40918671e+02],\n",
+ " [ 3.91108551e+01, 1.67909088e+01, 5.88121567e+01,\n",
+ " 4.12347527e+01],\n",
+ " [ 4.56910324e+01, 2.26641602e+02, 6.40384064e+01,\n",
+ " 2.43926025e+02],\n",
+ " [ 4.21584091e+01, 2.25309219e+02, 7.20502853e+01,\n",
+ " 2.46690781e+02],\n",
+ " [ 5.83352318e+01, 2.14019638e+02, 7.17930832e+01,\n",
+ " 2.26993179e+02],\n",
+ " [ 4.61224594e+01, 2.10781815e+02, 7.38775406e+01,\n",
+ " 2.45218185e+02],\n",
+ " [ 1.85833035e+01, 1.18999054e+02, 6.78477707e+01,\n",
+ " 2.23587372e+02],\n",
+ " [ 6.29499054e+00, 2.28007355e+02, 2.02015747e+02,\n",
+ " 2.55130493e+02],\n",
+ " [ 8.46007690e+01, 1.32339249e+01, 1.74692444e+02,\n",
+ " 7.96282196e+01],\n",
+ " [ 2.83830490e+01, 1.60392761e+01, 4.86811066e+01,\n",
+ " 4.19863853e+01],\n",
+ " [ 4.23666039e+01, 1.46869293e+02, 4.78643532e+01,\n",
+ " 1.52169189e+02],\n",
+ " [ 2.81077156e+01, 1.56691101e+02, 4.79436073e+01,\n",
+ " 1.69283234e+02],\n",
+ " [ 1.06909409e+02, 1.14706993e-01, 1.79521652e+02,\n",
+ " 2.54930611e+01],\n",
+ " [ 1.45406113e+01, 1.67513275e+02, 7.45542984e+01,\n",
+ " 2.50844360e+02],\n",
+ " [ 2.01761017e+01, 2.11647606e+01, 9.63713531e+01,\n",
+ " 9.08352356e+01],\n",
+ " [ 7.29417725e+01, 3.81889153e+01, 1.85248047e+02,\n",
+ " 1.05811081e+02],\n",
+ " [ 1.11883369e+02, 2.16834534e+02, 2.36978790e+02,\n",
+ " 2.56812073e+02],\n",
+ " [-1.84612274e-01, 2.67456284e+01, 9.61846161e+01,\n",
+ " 1.74089142e+02],\n",
+ " [ 4.36148262e+01, 1.69856461e+02, 4.86417885e+01,\n",
+ " 1.75156357e+02],\n",
+ " [ 4.24236183e+01, 1.93141953e+02, 4.98329964e+01,\n",
+ " 1.99870865e+02],\n",
+ " [ 4.32086182e+01, 2.01430054e+02, 5.10736618e+01,\n",
+ " 2.07582764e+02],\n",
+ " [ 5.40823860e+01, 2.22546677e+02, 7.47558289e+01,\n",
+ " 2.45156235e+02],\n",
+ " [ 4.76513481e+01, 2.37343201e+02, 1.14557343e+02,\n",
+ " 2.55761139e+02],\n",
+ " [ 1.07118378e+02, 3.94208450e+01, 1.69666077e+02,\n",
+ " 7.26179962e+01],\n",
+ " [ 3.05676422e+01, 2.16517334e+02, 7.85676422e+01,\n",
+ " 2.53212097e+02],\n",
+ " [ 5.05019989e+01, 5.75438004e+01, 9.06332855e+01,\n",
+ " 1.04132607e+02],\n",
+ " [ 1.55681335e+02, 1.77642731e+02, 2.54589233e+02,\n",
+ " 2.55545563e+02],\n",
+ " [ 3.34228935e+01, 1.91442375e+01, 4.36412621e+01,\n",
+ " 3.30097275e+01],\n",
+ " [ 4.09411125e+01, 1.22797295e+02, 4.92898445e+01,\n",
+ " 1.27228371e+02],\n",
+ " [ 4.43922653e+01, 1.56731995e+02, 4.98900146e+01,\n",
+ " 1.62357819e+02],\n",
+ " [ 4.24254799e+01, 1.60734451e+02, 5.36771660e+01,\n",
+ " 1.67265549e+02],\n",
+ " [ 4.63862305e+01, 2.08533630e+02, 5.57934036e+01,\n",
+ " 2.15466370e+02],\n",
+ " [ 4.98276215e+01, 2.17040039e+02, 5.64033356e+01,\n",
+ " 2.23972778e+02],\n",
+ " [ 3.86769180e+01, 2.25027588e+02, 5.64848785e+01,\n",
+ " 2.46972412e+02],\n",
+ " [ 4.89842186e+01, 2.40763367e+02, 5.92723961e+01,\n",
+ " 2.46223816e+02],\n",
+ " [ 2.92078056e+01, 2.53186157e+02, 1.05177109e+02,\n",
+ " 2.56890839e+02],\n",
+ " [ 3.38526840e+01, 1.96612396e+01, 8.93628540e+01,\n",
+ " 4.91232224e+01],\n",
+ " [ 4.46160049e+01, 7.77583923e+01, 9.13839951e+01,\n",
+ " 9.73090057e+01],\n",
+ " [ 9.73159485e+01, 6.98130798e+01, 1.41037445e+02,\n",
+ " 1.03010223e+02],\n",
+ " [ 2.07002075e+02, 4.82017288e+01, 2.39819122e+02,\n",
+ " 9.24296341e+01],\n",
+ " [ 8.73469925e+00, 2.08183563e+02, 7.48746185e+01,\n",
+ " 2.55816437e+02],\n",
+ " [ 5.24328232e+00, 7.11837959e+00, 1.06756714e+02,\n",
+ " 6.96660843e+01],\n",
+ " [ 7.88083954e+01, 2.17992599e+02, 1.93191605e+02,\n",
+ " 2.55654007e+02],\n",
+ " [ 1.05052139e+02, 2.34384880e+01, 1.98947861e+02,\n",
+ " 9.31089630e+01],\n",
+ " [ 2.27988815e+01, 9.83052063e+00, 2.14063263e+02,\n",
+ " 1.05462700e+02],\n",
+ " [ 9.54042816e+01, 1.05430603e+00, 1.97700058e+02,\n",
+ " 1.35780472e+02],\n",
+ " [ 4.27733917e+01, 2.20528381e+02, 5.38207626e+01,\n",
+ " 2.38336334e+02],\n",
+ " [ 3.36835365e+01, 2.14658142e+02, 5.97304420e+01,\n",
+ " 2.46755829e+02],\n",
+ " [ 3.59260406e+01, 1.87410698e+01, 6.60482941e+01,\n",
+ " 2.63230820e+01],\n",
+ " [ 4.21795197e+01, 2.01985092e+01, 5.89557648e+01,\n",
+ " 2.83956451e+01],\n",
+ " [ 4.84506264e+01, 2.39215210e+02, 6.77033463e+01,\n",
+ " 2.45746307e+02],\n",
+ " [ 7.02042618e+01, 6.32016373e+01, 9.34986649e+01,\n",
+ " 8.45012894e+01],\n",
+ " [ 1.38511826e+02, 1.07008301e+02, 1.61539505e+02,\n",
+ " 1.23978867e+02],\n",
+ " [ 1.15022446e+02, 1.99184158e+02, 1.88977554e+02,\n",
+ " 2.55720932e+02],\n",
+ " [ 5.58068085e+01, 1.32803238e+02, 1.76295837e+02,\n",
+ " 2.07042801e+02],\n",
+ " [ 2.85206337e+01, 2.44100628e+01, 4.36846581e+01,\n",
+ " 3.46284256e+01],\n",
+ " [ 3.15243740e+01, 2.91521416e+01, 3.87992096e+01,\n",
+ " 3.68743706e+01],\n",
+ " [ 3.84787064e+01, 1.87677402e+01, 5.38183670e+01,\n",
+ " 2.38529263e+01],\n",
+ " [ 4.91354561e+01, 1.36482269e+02, 5.91211586e+01,\n",
+ " 1.44530548e+02]]], dtype=float32),\n",
+ " array([100.], dtype=float32),\n",
+ " array([[0.7578125 , 0.62890625, 0.578125 , 0.2421875 , 0.203125 ,\n",
+ " 0.1875 , 0.14453125, 0.12109375, 0.109375 , 0.1015625 ,\n",
+ " 0.08984375, 0.08984375, 0.08984375, 0.08203125, 0.07421875,\n",
+ " 0.07421875, 0.07421875, 0.07421875, 0.06640625, 0.06640625,\n",
+ " 0.06640625, 0.06640625, 0.06640625, 0.06640625, 0.06640625,\n",
+ " 0.06640625, 0.0625 , 0.0625 , 0.0625 , 0.0625 ,\n",
+ " 0.0625 , 0.0546875 , 0.0546875 , 0.0546875 , 0.0546875 ,\n",
+ " 0.0546875 , 0.0546875 , 0.0546875 , 0.0546875 , 0.05078125,\n",
+ " 0.05078125, 0.05078125, 0.05078125, 0.05078125, 0.05078125,\n",
+ " 0.05078125, 0.05078125, 0.05078125, 0.05078125, 0.05078125,\n",
+ " 0.046875 , 0.046875 , 0.046875 , 0.046875 , 0.046875 ,\n",
+ " 0.046875 , 0.046875 , 0.046875 , 0.046875 , 0.04296875,\n",
+ " 0.04296875, 0.04296875, 0.04296875, 0.04296875, 0.04296875,\n",
+ " 0.04296875, 0.04296875, 0.04296875, 0.0390625 , 0.0390625 ,\n",
+ " 0.0390625 , 0.0390625 , 0.0390625 , 0.0390625 , 0.0390625 ,\n",
+ " 0.0390625 , 0.0390625 , 0.0390625 , 0.0390625 , 0.0390625 ,\n",
+ " 0.0390625 , 0.0390625 , 0.0390625 , 0.0390625 , 0.0390625 ,\n",
+ " 0.0390625 , 0.0390625 , 0.03515625, 0.03515625, 0.03515625,\n",
+ " 0.03515625, 0.03515625, 0.03515625, 0.03515625, 0.03515625,\n",
+ " 0.03515625, 0.03125 , 0.03125 , 0.03125 , 0.03125 ]],\n",
+ " dtype=float32),\n",
+ " array([[18., 8., 2., 1., 1., 2., 18., 3., 2., 64., 1., 62., 18.,\n",
+ " 8., 64., 3., 1., 1., 10., 1., 1., 64., 64., 2., 64., 62.,\n",
+ " 1., 1., 3., 3., 64., 1., 1., 3., 3., 8., 64., 62., 2.,\n",
+ " 64., 1., 3., 64., 3., 64., 3., 3., 8., 1., 62., 1., 1.,\n",
+ " 1., 70., 8., 64., 18., 15., 64., 1., 1., 1., 3., 3., 62.,\n",
+ " 3., 64., 62., 1., 1., 1., 1., 1., 3., 3., 3., 1., 64.,\n",
+ " 64., 18., 17., 64., 64., 15., 18., 18., 18., 3., 64., 1., 1.,\n",
+ " 1., 2., 44., 15., 2., 1., 1., 1., 64.]], dtype=float32),\n",
+ " array([[[256., 256.],\n",
+ " [256., 256.],\n",
+ " [ 1., 1.],\n",
+ " [ 0., 0.]]], dtype=float32)]"
+ ]
+ },
+ "execution_count": 37,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "# The final outputs is a list of [detection_boxes, num_detections, detection_scores, detection_classes, image_info].\n",
+ "outputs"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Pkxpgrjod9jB"
+ },
+ "source": [
+ "### Visualize Detection Outputs"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "TpzoZOJ_u1dw"
+ },
+ "outputs": [],
+ "source": [
+ "plt.imshow(Image.fromarray(image_array[0]))\n",
+ "\n",
+ "scores = outputs[2]\n",
+ "num_detection = outputs[1]\n",
+ "boxes = outputs[0]\n",
+ "classes = outputs[3]\n",
+ "\n",
+ "# We only show boxes that have detection score larger than 0.5.\n",
+ "threshold = 0.5\n",
+ "\n",
+ "for i in range(int(num_detection[0])):\n",
+ " if scores[0, i] \u003e threshold:\n",
+ " ax = plt.gca()\n",
+ " rect = patches.Rectangle((boxes[0, i, 1], boxes[0, i, 0]), boxes[0, i, 3] - boxes[0, i, 1], boxes[0, i, 2] - boxes[0, i, 0], linewidth=1,edgecolor='r',facecolor='none')\n",
+ " ax.add_patch(rect)"
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "provenance": [],
+ "toc_visible": true
+ },
+ "gpuClass": "standard",
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/qat/vision/modeling/__init__.py b/official/projects/qat/vision/modeling/__init__.py
index aa57bbd8683..fbf4818dd28 100644
--- a/official/projects/qat/vision/modeling/__init__.py
+++ b/official/projects/qat/vision/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,5 +13,5 @@
# limitations under the License.
"""Modeling package definition."""
-
+from official.projects.qat.vision.modeling import heads
from official.projects.qat.vision.modeling import layers
diff --git a/official/projects/qat/vision/modeling/factory.py b/official/projects/qat/vision/modeling/factory.py
index 59f321c2b07..c1f34fa459a 100644
--- a/official/projects/qat/vision/modeling/factory.py
+++ b/official/projects/qat/vision/modeling/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,38 +13,42 @@
# limitations under the License.
"""Factory methods to build models."""
-# Import libraries
-
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.projects.qat.vision.configs import common
from official.projects.qat.vision.modeling import segmentation_model as qat_segmentation_model
+from official.projects.qat.vision.modeling.heads import dense_prediction_heads as dense_prediction_heads_qat
+from official.projects.qat.vision.modeling.layers import nn_layers as qat_nn_layers
from official.projects.qat.vision.n_bit import schemes as n_bit_schemes
+from official.projects.qat.vision.quantization import configs as qat_configs
+from official.projects.qat.vision.quantization import helper
from official.projects.qat.vision.quantization import schemes
from official.vision import configs
from official.vision.modeling import classification_model
from official.vision.modeling import retinanet_model
from official.vision.modeling.decoders import aspp
+from official.vision.modeling.decoders import fpn
+from official.vision.modeling.heads import dense_prediction_heads
from official.vision.modeling.heads import segmentation_heads
from official.vision.modeling.layers import nn_layers
def build_qat_classification_model(
- model: tf.keras.Model,
+ model: tf_keras.Model,
quantization: common.Quantization,
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
model_config: configs.image_classification.ImageClassificationModel,
- l2_regularizer: tf.keras.regularizers.Regularizer = None
-) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None # pyrefly: ignore[bad-function-definition]
+) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Apply model optimization techniques.
Args:
model: The model applying model optimization techniques.
quantization: The Quantization config.
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
model_config: The model config.
- l2_regularizer: tf.keras.regularizers.Regularizer object. Default to None.
+ l2_regularizer: tf_keras.regularizers.Regularizer object. Default to None.
Returns:
model: The model that applied optimization techniques.
@@ -58,7 +62,7 @@ def build_qat_classification_model(
status.expect_partial().assert_existing_objects_matched()
scope_dict = {
- 'L2': tf.keras.regularizers.l2,
+ 'L2': tf_keras.regularizers.l2,
}
with tfmot.quantization.keras.quantize_scope(scope_dict):
annotated_backbone = tfmot.quantization.keras.quantize_annotate_model(
@@ -92,17 +96,19 @@ def build_qat_classification_model(
with tfmot.quantization.keras.quantize_scope(scope_dict):
def apply_quantization_to_dense(layer):
- if isinstance(layer, (tf.keras.layers.Dense,
- tf.keras.layers.Dropout,
- tf.keras.layers.GlobalAveragePooling2D)):
+ if isinstance(layer, (tf_keras.layers.Dense,
+ tf_keras.layers.Dropout,
+ tf_keras.layers.GlobalAveragePooling2D)):
return tfmot.quantization.keras.quantize_annotate_layer(layer)
return layer
- annotated_model = tf.keras.models.clone_model(
+ backbone_optimized_model.use_legacy_config = True
+ annotated_model = tf_keras.models.clone_model(
backbone_optimized_model,
clone_function=apply_quantization_to_dense,
)
+ annotated_model.use_legacy_config = True
if quantization.change_num_bits:
optimized_model = tfmot.quantization.keras.quantize_apply(
annotated_model,
@@ -117,9 +123,21 @@ def apply_quantization_to_dense(layer):
return optimized_model
+def _clone_function_for_fpn(layer):
+ if isinstance(layer, (
+ tf_keras.layers.BatchNormalization,
+ tf_keras.layers.experimental.SyncBatchNormalization)):
+ return tfmot.quantization.keras.quantize_annotate_layer(
+ qat_nn_layers.BatchNormalizationWrapper(layer),
+ qat_configs.Default8BitOutputQuantizeConfig())
+ if isinstance(layer, tf_keras.layers.UpSampling2D):
+ return layer
+ return tfmot.quantization.keras.quantize_annotate_layer(layer)
+
+
def build_qat_retinanet(
- model: tf.keras.Model, quantization: common.Quantization,
- model_config: configs.retinanet.RetinaNet) -> tf.keras.Model:
+ model: tf_keras.Model, quantization: common.Quantization,
+ model_config: configs.retinanet.RetinaNet) -> tf_keras.Model:
"""Applies quantization aware training for RetinaNet model.
Args:
@@ -140,7 +158,8 @@ def build_qat_retinanet(
status.expect_partial().assert_existing_objects_matched()
scope_dict = {
- 'L2': tf.keras.regularizers.l2,
+ 'L2': tf_keras.regularizers.l2,
+ 'BatchNormalizationWrapper': qat_nn_layers.BatchNormalizationWrapper,
}
with tfmot.quantization.keras.quantize_scope(scope_dict):
annotated_backbone = tfmot.quantization.keras.quantize_annotate_model(
@@ -148,22 +167,51 @@ def build_qat_retinanet(
optimized_backbone = tfmot.quantization.keras.quantize_apply(
annotated_backbone,
scheme=schemes.Default8BitQuantizeScheme())
+ decoder = model.decoder
+ if quantization.quantize_detection_decoder:
+ if not isinstance(decoder, fpn.FPN):
+ raise ValueError('Currently only supports FPN.')
+
+ decoder = tf_keras.models.clone_model(
+ decoder,
+ clone_function=_clone_function_for_fpn,
+ )
+ decoder = tfmot.quantization.keras.quantize_apply(decoder)
+ decoder = tfmot.quantization.keras.remove_input_range(decoder)
+
+ head = model.head
+ if quantization.quantize_detection_head:
+ if not isinstance(head, dense_prediction_heads.RetinaNetHead):
+ raise ValueError('Currently only supports RetinaNetHead.')
+ head = (
+ dense_prediction_heads_qat.RetinaNetHeadQuantized.from_config(
+ head.get_config()))
+
optimized_model = retinanet_model.RetinaNetModel(
- optimized_backbone,
- model.decoder,
- model.head,
- model.detection_generator,
+ backbone=optimized_backbone,
+ decoder=decoder,
+ head=head,
+ detection_generator=model.detection_generator,
+ anchor_boxes=model.anchor_boxes,
min_level=model_config.min_level,
max_level=model_config.max_level,
num_scales=model_config.anchor.num_scales,
aspect_ratios=model_config.anchor.aspect_ratios,
anchor_size=model_config.anchor.anchor_size)
+
+ if quantization.quantize_detection_head:
+ # Call the model with dummy input to build the head part.
+ dummpy_input = tf.zeros([1] + model_config.input_size)
+ height, width, _ = model_config.input_size
+ image_shape = [[height, width]]
+ optimized_model.call(dummpy_input, image_shape=image_shape, training=False)
+ helper.copy_original_weights(model.head, optimized_model.head)
return optimized_model
def build_qat_segmentation_model(
- model: tf.keras.Model, quantization: common.Quantization,
- input_specs: tf.keras.layers.InputSpec) -> tf.keras.Model:
+ model: tf_keras.Model, quantization: common.Quantization,
+ input_specs: tf_keras.layers.InputSpec) -> tf_keras.Model:
"""Applies quantization aware training for segmentation model.
Args:
@@ -186,10 +234,11 @@ def build_qat_segmentation_model(
model.backbone, model.decoder, model.head, input_specs)
scope_dict = {
- 'L2': tf.keras.regularizers.l2,
+ 'L2': tf_keras.regularizers.l2,
}
- # Apply QAT to backbone (a tf.keras.Model) first.
+ model.use_legacy_config = True # Ensures old Keras serialization format
+ # Apply QAT to backbone (a tf_keras.Model) first.
with tfmot.quantization.keras.quantize_scope(scope_dict):
annotated_backbone = tfmot.quantization.keras.quantize_annotate_model(
model.backbone)
@@ -212,10 +261,12 @@ def apply_quantization_to_layers(layer):
return tfmot.quantization.keras.quantize_annotate_layer(layer)
return layer
- annotated_model = tf.keras.models.clone_model(
+ backbone_optimized_model.use_legacy_config = True
+ annotated_model = tf_keras.models.clone_model(
backbone_optimized_model,
clone_function=apply_quantization_to_layers,
)
+ annotated_model.use_legacy_config = True
optimized_model = tfmot.quantization.keras.quantize_apply(
annotated_model, scheme=schemes.Default8BitQuantizeScheme())
diff --git a/official/projects/qat/vision/modeling/factory_test.py b/official/projects/qat/vision/modeling/factory_test.py
index e0875d77a01..2f359259e2a 100644
--- a/official/projects/qat/vision/modeling/factory_test.py
+++ b/official/projects/qat/vision/modeling/factory_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,42 +14,53 @@
"""Tests for factory.py."""
-# Import libraries
-
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.qat.vision.configs import common
from official.projects.qat.vision.modeling import factory as qat_factory
+from official.projects.qat.vision.modeling.heads import dense_prediction_heads as qat_dense_prediction_heads
from official.vision.configs import backbones
from official.vision.configs import decoders
from official.vision.configs import image_classification as classification_cfg
from official.vision.configs import retinanet as retinanet_cfg
from official.vision.configs import semantic_segmentation as semantic_segmentation_cfg
from official.vision.modeling import factory
+from official.vision.modeling.decoders import fpn
+from official.vision.modeling.heads import dense_prediction_heads
class ClassificationModelBuilderTest(parameterized.TestCase, tf.test.TestCase):
@parameterized.parameters(
- ('resnet', (224, 224), 5e-5),
- ('resnet', (224, 224), None),
- ('resnet', (None, None), 5e-5),
- ('resnet', (None, None), None),
- ('mobilenet', (224, 224), 5e-5),
- ('mobilenet', (224, 224), None),
- ('mobilenet', (None, None), 5e-5),
- ('mobilenet', (None, None), None),
+ ('resnet', 50, (224, 224), 5e-5),
+ ('resnet', 50, (224, 224), None),
+ ('resnet', 50, (None, None), 5e-5),
+ ('resnet', 50, (None, None), None),
+ ('mobilenet', 'MobileNetV2', (224, 224), 5e-5),
+ ('mobilenet', 'MobileNetV2', (224, 224), None),
+ ('mobilenet', 'MobileNetV2', (None, None), 5e-5),
+ ('mobilenet', 'MobileNetV2', (None, None), None),
+ ('mobilenet', 'MobileNetV4ConvLarge', (224, 224), 5e-5),
)
- def test_builder(self, backbone_type, input_size, weight_decay):
+ def test_builder(self, backbone_type, model_id, input_size, weight_decay):
num_classes = 2
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None, input_size[0], input_size[1], 3])
+
+ backbone = backbones.Backbone(type=backbone_type)
+ if backbone_type == 'resnet':
+ backbone.resnet.model_id = model_id
+ elif backbone_type == 'mobilenet':
+ backbone.mobilenet.model_id = model_id
+ else:
+ raise ValueError('Unexpected backbone_type', backbone_type)
+
model_config = classification_cfg.ImageClassificationModel(
- num_classes=num_classes,
- backbone=backbones.Backbone(type=backbone_type))
+ num_classes=num_classes, backbone=backbone
+ )
l2_regularizer = (
- tf.keras.regularizers.l2(weight_decay) if weight_decay else None)
+ tf_keras.regularizers.l2(weight_decay) if weight_decay else None)
model = factory.build_classification_model(
input_specs=input_specs,
model_config=model_config,
@@ -67,48 +78,99 @@ def test_builder(self, backbone_type, input_size, weight_decay):
class RetinaNetBuilderTest(parameterized.TestCase, tf.test.TestCase):
@parameterized.parameters(
- ('spinenet_mobile', (640, 640), False),
+ ('spinenet_mobile', 'identity', (640, 640), False, False),
+ ('spinenet_mobile', 'identity', (640, 640), True, False),
+ ('mobilenet', 'fpn', (640, 640), True, False),
+ ('mobilenet', 'fpn', (640, 640), True, True),
)
- def test_builder(self, backbone_type, input_size, has_attribute_heads):
+ def test_builder(self,
+ backbone_type,
+ decoder_type,
+ input_size,
+ quantize_detection_head,
+ quantize_detection_decoder):
num_classes = 2
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None, input_size[0], input_size[1], 3])
- if has_attribute_heads:
- attribute_heads_config = [
- retinanet_cfg.AttributeHead(name='att1'),
- retinanet_cfg.AttributeHead(
- name='att2', type='classification', size=2),
- ]
+
+ if backbone_type == 'spinenet_mobile':
+ backbone_config = backbones.Backbone(
+ type=backbone_type,
+ spinenet_mobile=backbones.SpineNetMobile(
+ model_id='49',
+ stochastic_depth_drop_rate=0.2,
+ min_level=3,
+ max_level=7,
+ use_keras_upsampling_2d=True))
+ elif backbone_type == 'mobilenet':
+ backbone_config = backbones.Backbone(
+ type=backbone_type,
+ mobilenet=backbones.MobileNet(
+ model_id='MobileNetV2',
+ filter_size_scale=1.0))
+ else:
+ raise ValueError(
+ 'backbone_type {} is not supported'.format(backbone_type))
+
+ if decoder_type == 'identity':
+ decoder_config = decoders.Decoder(type=decoder_type)
+ elif decoder_type == 'fpn':
+ decoder_config = decoders.Decoder(
+ type=decoder_type,
+ fpn=decoders.FPN(
+ num_filters=128,
+ use_separable_conv=True,
+ use_keras_layer=True))
else:
- attribute_heads_config = None
+ raise ValueError(
+ 'decoder_type {} is not supported'.format(decoder_type))
+
model_config = retinanet_cfg.RetinaNet(
num_classes=num_classes,
- backbone=backbones.Backbone(
- type=backbone_type,
- spinenet_mobile=backbones.SpineNetMobile(
- model_id='49',
- stochastic_depth_drop_rate=0.2,
- min_level=3,
- max_level=7,
- use_keras_upsampling_2d=True)),
+ input_size=[input_size[0], input_size[1], 3],
+ backbone=backbone_config,
+ decoder=decoder_config,
head=retinanet_cfg.RetinaNetHead(
- attribute_heads=attribute_heads_config))
- l2_regularizer = tf.keras.regularizers.l2(5e-5)
- quantization_config = common.Quantization()
+ attribute_heads=None,
+ use_separable_conv=True))
+
+ l2_regularizer = tf_keras.regularizers.l2(5e-5)
+ # Build the original float32 retinanet model.
model = factory.build_retinanet(
input_specs=input_specs,
model_config=model_config,
l2_regularizer=l2_regularizer)
- _ = qat_factory.build_qat_retinanet(
+ # Call the model with dummy input to build the head part.
+ dummpy_input = tf.zeros([1] + model_config.input_size)
+ model(dummpy_input, training=True)
+
+ # Build the QAT model from the original model with quantization config.
+ qat_model = qat_factory.build_qat_retinanet(
model=model,
- quantization=quantization_config,
+ quantization=common.Quantization(
+ quantize_detection_decoder=quantize_detection_decoder,
+ quantize_detection_head=quantize_detection_head),
model_config=model_config)
- if has_attribute_heads:
- self.assertEqual(model_config.head.attribute_heads[0].as_dict(),
- dict(name='att1', type='regression', size=1))
- self.assertEqual(model_config.head.attribute_heads[1].as_dict(),
- dict(name='att2', type='classification', size=2))
+
+ if quantize_detection_head:
+ # head become a RetinaNetHeadQuantized when we apply quantization.
+ self.assertIsInstance(qat_model.head,
+ qat_dense_prediction_heads.RetinaNetHeadQuantized)
+ else:
+ # head is a RetinaNetHead if we don't apply quantization on head part.
+ self.assertIsInstance(
+ qat_model.head, dense_prediction_heads.RetinaNetHead)
+ self.assertNotIsInstance(
+ qat_model.head, qat_dense_prediction_heads.RetinaNetHeadQuantized)
+
+ if decoder_type == 'FPN':
+ if quantize_detection_decoder:
+ # FPN decoder become a general keras functional model after applying
+ # quantization.
+ self.assertNotIsInstance(qat_model.decoder, fpn.FPN)
+ else:
+ self.assertIsInstance(qat_model.decoder, fpn.FPN)
class SegmentationModelBuilderTest(parameterized.TestCase, tf.test.TestCase):
@@ -117,7 +179,7 @@ class SegmentationModelBuilderTest(parameterized.TestCase, tf.test.TestCase):
('mobilenet', (512, 512), 5e-5),)
def test_deeplabv3_builder(self, backbone_type, input_size, weight_decay):
num_classes = 21
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None, input_size[0], input_size[1], 3])
model_config = semantic_segmentation_cfg.SemanticSegmentationModel(
num_classes=num_classes,
@@ -140,7 +202,7 @@ def test_deeplabv3_builder(self, backbone_type, input_size, weight_decay):
upsample_factor=2,
use_depthwise_convolution=True))
l2_regularizer = (
- tf.keras.regularizers.l2(weight_decay) if weight_decay else None)
+ tf_keras.regularizers.l2(weight_decay) if weight_decay else None)
model = factory.build_segmentation_model(
input_specs=input_specs,
model_config=model_config,
@@ -153,7 +215,7 @@ def test_deeplabv3_builder(self, backbone_type, input_size, weight_decay):
('mobilenet', (512, 1024), 5e-5),)
def test_deeplabv3plus_builder(self, backbone_type, input_size, weight_decay):
num_classes = 19
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None, input_size[0], input_size[1], 3])
model_config = semantic_segmentation_cfg.SemanticSegmentationModel(
num_classes=num_classes,
@@ -184,7 +246,7 @@ def test_deeplabv3plus_builder(self, backbone_type, input_size, weight_decay):
upsample_factor=1,
num_filters=256))
l2_regularizer = (
- tf.keras.regularizers.l2(weight_decay) if weight_decay else None)
+ tf_keras.regularizers.l2(weight_decay) if weight_decay else None)
model = factory.build_segmentation_model(
input_specs=input_specs,
model_config=model_config,
diff --git a/official/projects/qat/vision/modeling/heads/__init__.py b/official/projects/qat/vision/modeling/heads/__init__.py
new file mode 100644
index 00000000000..e593d407791
--- /dev/null
+++ b/official/projects/qat/vision/modeling/heads/__init__.py
@@ -0,0 +1,16 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Heads package definition."""
+from official.projects.qat.vision.modeling.heads.dense_prediction_heads import RetinaNetHeadQuantized
diff --git a/official/projects/qat/vision/modeling/heads/dense_prediction_heads.py b/official/projects/qat/vision/modeling/heads/dense_prediction_heads.py
new file mode 100644
index 00000000000..eb61b07737a
--- /dev/null
+++ b/official/projects/qat/vision/modeling/heads/dense_prediction_heads.py
@@ -0,0 +1,381 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains definitions of dense prediction heads."""
+from typing import List, Mapping, Union, Optional, Any, Dict
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.modeling import tf_utils
+from official.projects.qat.vision.quantization import configs
+from official.projects.qat.vision.quantization import helper
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class RetinaNetHeadQuantized(tf_keras.layers.Layer):
+ """Creates a RetinaNet quantized head."""
+
+ def __init__(
+ self,
+ min_level: int,
+ max_level: int,
+ num_classes: int,
+ num_anchors_per_location: int,
+ num_convs: int = 4,
+ num_filters: int = 256,
+ attribute_heads: Optional[List[Dict[str, Any]]] = None,
+ use_separable_conv: bool = False,
+ activation: str = 'relu',
+ use_sync_bn: bool = False,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ num_params_per_anchor: int = 4,
+ share_classification_heads: bool = False,
+ share_level_convs: bool = True,
+ **kwargs):
+ """Initializes a RetinaNet quantized head.
+
+ Args:
+ min_level: An `int` number of minimum feature level.
+ max_level: An `int` number of maximum feature level.
+ num_classes: An `int` number of classes to predict.
+ num_anchors_per_location: An `int` number of number of anchors per pixel
+ location.
+ num_convs: An `int` number that represents the number of the intermediate
+ conv layers before the prediction.
+ num_filters: An `int` number that represents the number of filters of the
+ intermediate conv layers.
+ attribute_heads: If not None, a list that contains a dict for each
+ additional attribute head. Each dict consists of 4 key-value pairs:
+ `name`, `type` ('regression' or 'classification'), `size` (number of
+ predicted values for each instance), and `prediction_tower_name`
+ (optional, specifies shared prediction towers.)
+ use_separable_conv: A `bool` that indicates whether the separable
+ convolution layers is used.
+ activation: A `str` that indicates which activation is used, e.g. 'relu',
+ 'swish', etc.
+ use_sync_bn: A `bool` that indicates whether to use synchronized batch
+ normalization across different replicas.
+ norm_momentum: A `float` of normalization momentum for the moving average.
+ norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default is None.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
+ num_params_per_anchor: Number of parameters required to specify an anchor
+ box. For example, `num_params_per_anchor` would be 4 for axis-aligned
+ anchor boxes specified by their y-centers, x-centers, heights, and
+ widths.
+ share_classification_heads: A `bool` that indicates whethere sharing
+ weights among the main and attribute classification heads. Not used in
+ the QAT model.
+ share_level_convs: An optional bool to enable sharing convs
+ across levels for classnet, boxnet, classifier and box regressor.
+ If True, convs will be shared across all levels. Not used in the QAT
+ model.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+ del share_classification_heads
+ del share_level_convs
+
+ super().__init__(**kwargs)
+ self._config_dict = {
+ 'min_level': min_level,
+ 'max_level': max_level,
+ 'num_classes': num_classes,
+ 'num_anchors_per_location': num_anchors_per_location,
+ 'num_convs': num_convs,
+ 'num_filters': num_filters,
+ 'attribute_heads': attribute_heads,
+ 'use_separable_conv': use_separable_conv,
+ 'activation': activation,
+ 'use_sync_bn': use_sync_bn,
+ 'norm_momentum': norm_momentum,
+ 'norm_epsilon': norm_epsilon,
+ 'kernel_regularizer': kernel_regularizer,
+ 'bias_regularizer': bias_regularizer,
+ 'num_params_per_anchor': num_params_per_anchor,
+ }
+
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ else:
+ self._bn_axis = 1
+ self._activation = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(activation, use_keras_layer=True),
+ configs.Default8BitActivationQuantizeConfig())
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the variables of the head."""
+ if self._config_dict['use_separable_conv']:
+ conv_op = helper.SeparableConv2DQuantized
+ else:
+ conv_op = helper.quantize_wrapped_layer(
+ tf_keras.layers.Conv2D,
+ configs.Default8BitConvQuantizeConfig(
+ ['kernel'], ['activation'], False))
+ conv_kwargs = {
+ 'filters': self._config_dict['num_filters'],
+ 'kernel_size': 3,
+ 'padding': 'same',
+ 'bias_initializer': tf.zeros_initializer(),
+ 'bias_regularizer': self._config_dict['bias_regularizer'],
+ }
+ if not self._config_dict['use_separable_conv']:
+ conv_kwargs.update({ # pyrefly: ignore[no-matching-overload]
+ 'kernel_initializer': tf_keras.initializers.RandomNormal(
+ stddev=0.01),
+ 'kernel_regularizer': self._config_dict['kernel_regularizer'],
+ })
+
+ base_bn_op = (tf_keras.layers.experimental.SyncBatchNormalization
+ if self._config_dict['use_sync_bn']
+ else tf_keras.layers.BatchNormalization)
+ bn_op = helper.norm_by_activation(
+ self._config_dict['activation'],
+ helper.quantize_wrapped_layer(
+ base_bn_op, configs.Default8BitOutputQuantizeConfig()),
+ helper.quantize_wrapped_layer(
+ base_bn_op, configs.NoOpQuantizeConfig()))
+
+ bn_kwargs = {
+ 'axis': self._bn_axis,
+ 'momentum': self._config_dict['norm_momentum'],
+ 'epsilon': self._config_dict['norm_epsilon'],
+ }
+
+ # Class net.
+ self._cls_convs = []
+ self._cls_norms = []
+ for level in range(
+ self._config_dict['min_level'], self._config_dict['max_level'] + 1):
+ this_level_cls_norms = []
+ for i in range(self._config_dict['num_convs']):
+ if level == self._config_dict['min_level']:
+ cls_conv_name = 'classnet-conv_{}'.format(i)
+ self._cls_convs.append(conv_op(name=cls_conv_name, **conv_kwargs))
+ cls_norm_name = 'classnet-conv-norm_{}_{}'.format(level, i)
+ this_level_cls_norms.append(bn_op(name=cls_norm_name, **bn_kwargs))
+ self._cls_norms.append(this_level_cls_norms)
+
+ classifier_kwargs = {
+ 'filters': (
+ self._config_dict['num_classes'] *
+ self._config_dict['num_anchors_per_location']),
+ 'kernel_size': 3,
+ 'padding': 'same',
+ 'bias_initializer': tf.constant_initializer(-np.log((1 - 0.01) / 0.01)),
+ 'bias_regularizer': self._config_dict['bias_regularizer'],
+ }
+ if not self._config_dict['use_separable_conv']:
+ classifier_kwargs.update({ # pyrefly: ignore[no-matching-overload]
+ 'kernel_initializer': tf_keras.initializers.RandomNormal(stddev=1e-5),
+ 'kernel_regularizer': self._config_dict['kernel_regularizer'],
+ })
+ self._classifier = conv_op(
+ name='scores', last_quantize=True, **classifier_kwargs)
+
+ # Box net.
+ self._box_convs = []
+ self._box_norms = []
+ for level in range(
+ self._config_dict['min_level'], self._config_dict['max_level'] + 1):
+ this_level_box_norms = []
+ for i in range(self._config_dict['num_convs']):
+ if level == self._config_dict['min_level']:
+ box_conv_name = 'boxnet-conv_{}'.format(i)
+ self._box_convs.append(conv_op(name=box_conv_name, **conv_kwargs))
+ box_norm_name = 'boxnet-conv-norm_{}_{}'.format(level, i)
+ this_level_box_norms.append(bn_op(name=box_norm_name, **bn_kwargs))
+ self._box_norms.append(this_level_box_norms)
+
+ box_regressor_kwargs = {
+ 'filters': (self._config_dict['num_params_per_anchor'] *
+ self._config_dict['num_anchors_per_location']),
+ 'kernel_size': 3,
+ 'padding': 'same',
+ 'bias_initializer': tf.zeros_initializer(),
+ 'bias_regularizer': self._config_dict['bias_regularizer'],
+ }
+ if not self._config_dict['use_separable_conv']:
+ box_regressor_kwargs.update({ # pyrefly: ignore[no-matching-overload]
+ 'kernel_initializer': tf_keras.initializers.RandomNormal(
+ stddev=1e-5),
+ 'kernel_regularizer': self._config_dict['kernel_regularizer'],
+ })
+ self._box_regressor = conv_op(
+ name='boxes', last_quantize=True, **box_regressor_kwargs)
+
+ # Attribute learning nets.
+ if self._config_dict['attribute_heads']:
+ self._att_predictors = {}
+ self._att_convs = {}
+ self._att_norms = {}
+
+ for att_config in self._config_dict['attribute_heads']:
+ att_name = att_config['name']
+ att_type = att_config['type']
+ att_size = att_config['size']
+ att_convs_i = []
+ att_norms_i = []
+
+ # Build conv and norm layers.
+ for level in range(self._config_dict['min_level'],
+ self._config_dict['max_level'] + 1):
+ this_level_att_norms = []
+ for i in range(self._config_dict['num_convs']):
+ if level == self._config_dict['min_level']:
+ att_conv_name = '{}-conv_{}'.format(att_name, i)
+ att_convs_i.append(conv_op(name=att_conv_name, **conv_kwargs))
+ att_norm_name = '{}-conv-norm_{}_{}'.format(att_name, level, i)
+ this_level_att_norms.append(bn_op(name=att_norm_name, **bn_kwargs))
+ att_norms_i.append(this_level_att_norms)
+ self._att_convs[att_name] = att_convs_i
+ self._att_norms[att_name] = att_norms_i
+
+ # Build the final prediction layer.
+ att_predictor_kwargs = {
+ 'filters':
+ (att_size * self._config_dict['num_anchors_per_location']),
+ 'kernel_size': 3,
+ 'padding': 'same',
+ 'bias_initializer': tf.zeros_initializer(),
+ 'bias_regularizer': self._config_dict['bias_regularizer'],
+ }
+ if att_type == 'regression':
+ att_predictor_kwargs.update(
+ {'bias_initializer': tf.zeros_initializer()})
+ elif att_type == 'classification':
+ att_predictor_kwargs.update({
+ 'bias_initializer':
+ tf.constant_initializer(-np.log((1 - 0.01) / 0.01))
+ })
+ else:
+ raise ValueError(
+ 'Attribute head type {} not supported.'.format(att_type))
+
+ if not self._config_dict['use_separable_conv']:
+ att_predictor_kwargs.update({
+ 'kernel_initializer':
+ tf_keras.initializers.RandomNormal(stddev=1e-5),
+ 'kernel_regularizer':
+ self._config_dict['kernel_regularizer'],
+ })
+
+ self._att_predictors[att_name] = conv_op(
+ name='{}_attributes'.format(att_name), **att_predictor_kwargs)
+
+ super().build(input_shape)
+
+ def call(self, features: Mapping[str, tf.Tensor]):
+ """Forward pass of the RetinaNet quantized head.
+
+ Args:
+ features: A `dict` of `tf.Tensor` where
+ - key: A `str` of the level of the multilevel features.
+ - values: A `tf.Tensor`, the feature map tensors, whose shape is
+ [batch, height_l, width_l, channels].
+
+ Returns:
+ scores: A `dict` of `tf.Tensor` which includes scores of the predictions.
+ - key: A `str` of the level of the multilevel predictions.
+ - values: A `tf.Tensor` of the box scores predicted from a particular
+ feature level, whose shape is
+ [batch, height_l, width_l, num_classes * num_anchors_per_location].
+ boxes: A `dict` of `tf.Tensor` which includes coordinates of the
+ predictions.
+ - key: A `str` of the level of the multilevel predictions.
+ - values: A `tf.Tensor` of the box scores predicted from a particular
+ feature level, whose shape is
+ [batch, height_l, width_l,
+ num_params_per_anchor * num_anchors_per_location].
+ attributes: a dict of (attribute_name, attribute_prediction). Each
+ `attribute_prediction` is a dict of:
+ - key: `str`, the level of the multilevel predictions.
+ - values: `Tensor`, the box scores predicted from a particular feature
+ level, whose shape is
+ [batch, height_l, width_l,
+ attribute_size * num_anchors_per_location].
+ Can be an empty dictionary if no attribute learning is required.
+ """
+ scores = {}
+ boxes = {}
+ if self._config_dict['attribute_heads']:
+ attributes = {
+ att_config['name']: {}
+ for att_config in self._config_dict['attribute_heads']
+ }
+ else:
+ attributes = {}
+
+ for i, level in enumerate(
+ range(self._config_dict['min_level'],
+ self._config_dict['max_level'] + 1)):
+ this_level_features = features[str(level)]
+
+ # class net.
+ x = this_level_features
+ for conv, norm in zip(self._cls_convs, self._cls_norms[i]):
+ x = conv(x)
+ x = norm(x)
+ x = self._activation(x)
+ scores[str(level)] = self._classifier(x)
+
+ # box net.
+ x = this_level_features
+ for conv, norm in zip(self._box_convs, self._box_norms[i]):
+ x = conv(x)
+ x = norm(x)
+ x = self._activation(x)
+ boxes[str(level)] = self._box_regressor(x)
+
+ # attribute nets.
+ if self._config_dict['attribute_heads']:
+ prediction_tower_output = {}
+ for att_config in self._config_dict['attribute_heads']:
+ att_name = att_config['name']
+
+ def build_prediction_tower(atttribute_name, features, feature_level):
+ x = features
+ for conv, norm in zip(
+ self._att_convs[atttribute_name],
+ self._att_norms[atttribute_name][feature_level]):
+ x = conv(x)
+ x = norm(x)
+ x = self._activation(x)
+ return x
+
+ prediction_tower_name = att_config['prediction_tower_name']
+ if not prediction_tower_name:
+ attributes[att_name][str(level)] = self._att_predictors[att_name](
+ build_prediction_tower(att_name, this_level_features, i))
+ else:
+ if prediction_tower_name not in prediction_tower_output:
+ prediction_tower_output[
+ prediction_tower_name] = build_prediction_tower(
+ att_name, this_level_features, i)
+ attributes[att_name][str(level)] = self._att_predictors[att_name](
+ prediction_tower_output[prediction_tower_name])
+
+ return scores, boxes, attributes
+
+ def get_config(self):
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config):
+ return cls(**config)
diff --git a/official/projects/qat/vision/modeling/heads/dense_prediction_heads_test.py b/official/projects/qat/vision/modeling/heads/dense_prediction_heads_test.py
new file mode 100644
index 00000000000..6a040477320
--- /dev/null
+++ b/official/projects/qat/vision/modeling/heads/dense_prediction_heads_test.py
@@ -0,0 +1,116 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Lint as: python3
+"""Tests for dense_prediction_heads.py."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.qat.vision.modeling.heads import dense_prediction_heads
+
+
+def get_attribute_heads(att_head_type):
+ if att_head_type == 'regression_head':
+ return [
+ dict(name='depth', type='regression', size=1, prediction_tower_name='')
+ ]
+ elif att_head_type == 'shared_prediction_tower_attribute_heads':
+ return [
+ dict(
+ name='attr_1', type='regression', size=1, prediction_tower_name=''),
+ dict(
+ name='attr_2',
+ type='classification',
+ size=1,
+ prediction_tower_name='tower_1'),
+ dict(
+ name='attr_3',
+ type='regression',
+ size=1,
+ prediction_tower_name='tower_1')
+ ]
+ else:
+ raise ValueError('Undefined attribute type.')
+
+
+class RetinaNetHeadQuantizedTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ (False, False, False, None),
+ (False, True, False, None),
+ (True, False, True, 'regression_head'),
+ (True, True, True, 'regression_head'),
+ (True, True, True, 'shared_prediction_tower_attribute_heads'),
+ )
+ def test_forward(self, use_separable_conv, use_sync_bn, has_att_heads,
+ att_head_type):
+ if has_att_heads:
+ attribute_heads = get_attribute_heads(att_head_type)
+ else:
+ attribute_heads = None
+
+ retinanet_head = dense_prediction_heads.RetinaNetHeadQuantized(
+ min_level=3,
+ max_level=4,
+ num_classes=3,
+ num_anchors_per_location=3,
+ num_convs=2,
+ num_filters=256,
+ attribute_heads=attribute_heads,
+ use_separable_conv=use_separable_conv,
+ activation='relu',
+ use_sync_bn=use_sync_bn,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ )
+ features = {
+ '3': np.random.rand(2, 128, 128, 16),
+ '4': np.random.rand(2, 64, 64, 16),
+ }
+ scores, boxes, attributes = retinanet_head(features)
+ self.assertAllEqual(scores['3'].numpy().shape, [2, 128, 128, 9])
+ self.assertAllEqual(scores['4'].numpy().shape, [2, 64, 64, 9])
+ self.assertAllEqual(boxes['3'].numpy().shape, [2, 128, 128, 12])
+ self.assertAllEqual(boxes['4'].numpy().shape, [2, 64, 64, 12])
+ if has_att_heads:
+ for att in attributes.values():
+ self.assertAllEqual(att['3'].numpy().shape, [2, 128, 128, 3])
+ self.assertAllEqual(att['4'].numpy().shape, [2, 64, 64, 3])
+
+ def test_serialize_deserialize(self):
+ retinanet_head = dense_prediction_heads.RetinaNetHeadQuantized(
+ min_level=3,
+ max_level=7,
+ num_classes=3,
+ num_anchors_per_location=9,
+ num_convs=2,
+ num_filters=16,
+ attribute_heads=None,
+ use_separable_conv=False,
+ activation='relu',
+ use_sync_bn=False,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ )
+ config = retinanet_head.get_config()
+ new_retinanet_head = (
+ dense_prediction_heads.RetinaNetHead.from_config(config))
+ self.assertAllEqual(
+ retinanet_head.get_config(), new_retinanet_head.get_config())
diff --git a/official/projects/qat/vision/modeling/layers/__init__.py b/official/projects/qat/vision/modeling/layers/__init__.py
index 534843dd658..c7f055a9113 100644
--- a/official/projects/qat/vision/modeling/layers/__init__.py
+++ b/official/projects/qat/vision/modeling/layers/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Layers package definition."""
-
+from official.projects.qat.vision.modeling.layers import nn_layers
from official.projects.qat.vision.modeling.layers.nn_blocks import BottleneckBlockQuantized
from official.projects.qat.vision.modeling.layers.nn_blocks import Conv2DBNBlockQuantized
from official.projects.qat.vision.modeling.layers.nn_blocks import InvertedBottleneckBlockQuantized
diff --git a/official/projects/qat/vision/modeling/layers/nn_blocks.py b/official/projects/qat/vision/modeling/layers/nn_blocks.py
index a5e2c145363..7732bf587fb 100644
--- a/official/projects/qat/vision/modeling/layers/nn_blocks.py
+++ b/official/projects/qat/vision/modeling/layers/nn_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,8 @@
"""Contains quantized neural blocks for the QAT."""
from typing import Any, Dict, Optional, Sequence, Tuple, Union
-# Import libraries
-
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.modeling import tf_utils
@@ -30,8 +28,8 @@
# This class is copied from modeling.layers.nn_blocks.BottleneckBlock and apply
# QAT.
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class BottleneckBlockQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class BottleneckBlockQuantized(tf_keras.layers.Layer):
"""A quantized standard bottleneck block."""
def __init__(self,
@@ -43,8 +41,8 @@ def __init__(self,
resnetd_shortcut: bool = False,
stochastic_depth_drop_rate: Optional[float] = None,
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: tf.keras.regularizers.Regularizer = None,
- bias_regularizer: tf.keras.regularizers.Regularizer = None,
+ kernel_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
+ bias_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
activation: str = 'relu',
use_sync_bn: bool = False,
norm_momentum: float = 0.99,
@@ -70,9 +68,9 @@ def __init__(self,
the stochastic depth layer.
kernel_initializer: A `str` of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2d.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2d.
Default to None.
activation: A `str` name of the activation function.
use_sync_bn: A `bool`. If True, use synchronized batch normalization.
@@ -100,12 +98,12 @@ def __init__(self,
self._bias_regularizer = bias_regularizer
norm_layer = (
- tf.keras.layers.experimental.SyncBatchNormalization
- if use_sync_bn else tf.keras.layers.BatchNormalization)
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if use_sync_bn else tf_keras.layers.BatchNormalization)
self._norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
self._norm = helper.BatchNormalizationNoQuantized(norm_layer)
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -115,7 +113,7 @@ def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
"""Build variables and child layers to prepare for calling."""
if self._use_projection:
if self._resnetd_shortcut:
- self._shortcut0 = tf.keras.layers.AveragePooling2D(
+ self._shortcut0 = tf_keras.layers.AveragePooling2D(
pool_size=2, strides=self._strides, padding='same')
self._shortcut1 = helper.Conv2DQuantized(
filters=self._filters * 4,
@@ -216,7 +214,7 @@ def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
else:
self._stochastic_depth = None
self._add = tfmot.quantization.keras.QuantizeWrapperV2(
- tf.keras.layers.Add(),
+ tf_keras.layers.Add(),
configs.Default8BitQuantizeConfig([], [], True))
super(BottleneckBlockQuantized, self).build(input_shape)
@@ -280,8 +278,8 @@ def call(
# This class is copied from modeling.backbones.mobilenet.Conv2DBNBlock and apply
# QAT.
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class Conv2DBNBlockQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Conv2DBNBlockQuantized(tf_keras.layers.Layer):
"""A quantized convolution block with batch normalization."""
def __init__(
@@ -293,8 +291,8 @@ def __init__(
use_explicit_padding: bool = False,
activation: str = 'relu6',
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
use_normalization: bool = True,
use_sync_bn: bool = False,
norm_momentum: float = 0.99,
@@ -316,9 +314,9 @@ def __init__(
activation: A `str` name of the activation function.
kernel_initializer: A `str` for kernel initializer of convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2D.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
Default to None.
use_normalization: If True, use batch normalization.
use_sync_bn: If True, use synchronized batch normalization.
@@ -347,12 +345,12 @@ def __init__(
self._padding = 'same'
norm_layer = (
- tf.keras.layers.experimental.SyncBatchNormalization
- if use_sync_bn else tf.keras.layers.BatchNormalization)
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if use_sync_bn else tf_keras.layers.BatchNormalization)
self._norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
self._norm = helper.BatchNormalizationNoQuantized(norm_layer)
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -381,7 +379,7 @@ def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
"""Build variables and child layers to prepare for calling."""
if self._use_explicit_padding and self._kernel_size > 1:
padding_size = nn_layers.get_padding_for_kernel_size(self._kernel_size)
- self._pad = tf.keras.layers.ZeroPadding2D(padding_size)
+ self._pad = tf_keras.layers.ZeroPadding2D(padding_size)
conv2d_quantized = (
helper.Conv2DQuantized
if self._use_normalization else helper.Conv2DOutputQuantized)
@@ -421,8 +419,8 @@ def call(
return self._activation_layer(x)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class InvertedBottleneckBlockQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class InvertedBottleneckBlockQuantized(tf_keras.layers.Layer):
"""A quantized inverted bottleneck block."""
def __init__(self,
@@ -467,9 +465,9 @@ def __init__(self,
the stochastic depth layer.
kernel_initializer: A `str` of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2d.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2d.
Default to None.
activation: A `str` name of the activation function.
se_inner_activation: A `str` name of squeeze-excitation inner activation.
@@ -528,21 +526,21 @@ def __init__(self,
self._output_intermediate_endpoints = output_intermediate_endpoints
norm_layer = (
- tf.keras.layers.experimental.SyncBatchNormalization
- if use_sync_bn else tf.keras.layers.BatchNormalization)
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if use_sync_bn else tf_keras.layers.BatchNormalization)
self._norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
self._norm = helper.BatchNormalizationNoQuantized(norm_layer)
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
if not depthwise_activation:
self._depthwise_activation = activation
if regularize_depthwise:
- self._depthsize_regularizer = kernel_regularizer
+ self._depthwise_regularizer = kernel_regularizer
else:
- self._depthsize_regularizer = None
+ self._depthwise_regularizer = None
def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
"""Build variables and child layers to prepare for calling."""
@@ -584,9 +582,10 @@ def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
dilation_rate=self._dilation_rate,
use_bias=False,
depthwise_initializer=self._kernel_initializer,
- depthwise_regularizer=self._depthsize_regularizer,
+ depthwise_regularizer=self._depthwise_regularizer,
bias_regularizer=self._bias_regularizer,
- activation=helper.NoOpActivation())
+ activation=helper.NoOpActivation(),
+ )
self._norm1 = helper.norm_by_activation(self._depthwise_activation,
self._norm_with_quantize,
self._norm)(
@@ -641,7 +640,7 @@ def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
else:
self._stochastic_depth = None
self._add = tfmot.quantization.keras.QuantizeWrapperV2(
- tf.keras.layers.Add(),
+ tf_keras.layers.Add(),
configs.Default8BitQuantizeConfig([], [], True))
super(InvertedBottleneckBlockQuantized, self).build(input_shape)
@@ -715,3 +714,353 @@ def call(
if self._output_intermediate_endpoints:
return x, endpoints
return x
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class UniversalInvertedBottleneckBlockQuantized(tf_keras.layers.Layer):
+ """A quantized inverted bottleneck block with optional depthwise convs."""
+
+ def __init__(
+ self,
+ in_filters: int,
+ out_filters: int,
+ expand_ratio: float,
+ strides: int,
+ middle_dw_downsample: bool = True,
+ start_dw_kernel_size: int = 0,
+ middle_dw_kernel_size: int = 3,
+ end_dw_kernel_size: int = 0,
+ stochastic_depth_drop_rate: float | None = None,
+ kernel_initializer: str = 'VarianceScaling',
+ kernel_regularizer: tf_keras.regularizers.Regularizer | None = None,
+ bias_regularizer: tf_keras.regularizers.Regularizer | None = None,
+ activation: str = 'relu',
+ depthwise_activation: str | None = None,
+ use_sync_bn: bool = False,
+ dilation_rate: int = 1,
+ divisible_by: int = 1,
+ regularize_depthwise: bool = False,
+ use_residual: bool = True,
+ use_layer_scale: bool = False,
+ layer_scale_init_value: float = 1e-5,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 0.001,
+ output_intermediate_endpoints: bool = False,
+ **kwargs,
+ ):
+ """Initializes a UniversalInvertedBottleneckBlockQuantized.
+
+ This is an extension of IB with optional depthwise convs before expansion (
+ "starting" conv) and after projection ("ending" conv). Both of these convs
+ are executed without activation. The standard depthwise conv of IB ("middle"
+ conv) is optional too. This last one is followed by an activation, as in
+ standard IBs. Squeeze-and-Excite or fused types of IBs are not supported.
+
+ Args:
+ in_filters: The number of filters of the input tensor.
+ out_filters: The number of filters of the output tensor.
+ expand_ratio: The filter multiplier for the first inverted bottleneck
+ stage.
+ strides: The block stride. If greater than 1, this block will ultimately
+ downsample the input.
+ middle_dw_downsample: If True, downsample in the middle depthwise
+ otherwise downsample in the starting one.
+ start_dw_kernel_size: The kernel size of the starting depthwise. A value
+ of zero means that no starting depthwise will be added.
+ middle_dw_kernel_size: The kernel size of the middle depthwise. A value of
+ zero means that no middle depthwise will be added.
+ end_dw_kernel_size: The kernel size of the ending depthwise. A value of
+ zero means that no ending depthwise will be added.
+ stochastic_depth_drop_rate: If not None, drop rate for the stochastic
+ depth layer.
+ kernel_initializer: The name of the convolutional layer
+ kernel_initializer.
+ kernel_regularizer: An optional kernel regularizer for the Conv2ds.
+ bias_regularizer: An optional bias regularizer for the Conv2ds.
+ activation: The name of the activation function.
+ depthwise_activation: The name of the depthwise-only activation function.
+ use_sync_bn: If True, use synchronized batch normalization.
+ dilation_rate: The dilation rate to use for convolutions.
+ divisible_by: Ensures all inner dimensions are divisible by this number.
+ regularize_depthwise: If True, apply regularization on depthwise.
+ use_residual: If True, include residual connection between input and
+ output.
+ use_layer_scale: If True, use layer scale.
+ layer_scale_init_value: The initial layer scale value.
+ norm_momentum: Momentum value for the moving average in normalization.
+ norm_epsilon: Value added to variance to avoid dividing by zero in
+ normalization.
+ output_intermediate_endpoints: This block does not output any intermediate
+ endpoint, but this argument is included for compatibility with other
+ blocks.
+ **kwargs: Additional keyword arguments to be passed to
+ tf_keras.layers.Layer.
+ """
+ super().__init__(**kwargs)
+ logging.info(
+ 'UniversalInvertedBottleneckBlockQuantized with depthwise kernel sizes '
+ '{%d, %d, %d}, strides=%d, and middle downsampling: %s',
+ start_dw_kernel_size,
+ middle_dw_kernel_size,
+ end_dw_kernel_size,
+ strides,
+ middle_dw_downsample,
+ )
+
+ self._in_filters = in_filters
+ self._out_filters = out_filters
+ self._expand_ratio = expand_ratio
+ self._strides = strides
+ self._middle_dw_downsample = middle_dw_downsample
+ self._start_dw_kernel_size = start_dw_kernel_size
+ self._middle_dw_kernel_size = middle_dw_kernel_size
+ self._end_dw_kernel_size = end_dw_kernel_size
+ self._divisible_by = divisible_by
+ self._stochastic_depth_drop_rate = stochastic_depth_drop_rate
+ self._dilation_rate = dilation_rate
+ self._use_sync_bn = use_sync_bn
+ self._regularize_depthwise = regularize_depthwise
+ self._use_residual = use_residual
+ self._activation = activation
+ self._depthwise_activation = depthwise_activation
+ self._kernel_initializer = kernel_initializer
+ self._use_layer_scale = use_layer_scale
+ self._layer_scale_init_value = layer_scale_init_value
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+ self._output_intermediate_endpoints = output_intermediate_endpoints
+
+ if strides > 1:
+ if middle_dw_downsample and not middle_dw_kernel_size:
+ raise ValueError(
+ 'Requested downsampling at a non-existing middle depthwise conv.'
+ )
+ if not middle_dw_downsample and not start_dw_kernel_size:
+ raise ValueError(
+ 'Requested downsampling at a non-existing starting depthwise conv.'
+ )
+
+ if use_sync_bn:
+ norm_layer = tf_keras.layers.experimental.SyncBatchNormalization
+ else:
+ norm_layer = tf_keras.layers.BatchNormalization
+ self._norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
+ self._norm = helper.BatchNormalizationNoQuantized(norm_layer)
+
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ self._bn_axis = -1
+ else:
+ self._bn_axis = 1
+ if not depthwise_activation:
+ self._depthwise_activation = activation
+ if regularize_depthwise:
+ self._depthwise_regularizer = kernel_regularizer
+ else:
+ self._depthwise_regularizer = None
+
+ def build(self, input_shape):
+ # Starting depthwise conv.
+ if self._start_dw_kernel_size:
+ self._start_dw_conv = helper.DepthwiseConv2DQuantized(
+ kernel_size=self._start_dw_kernel_size,
+ strides=self._strides if not self._middle_dw_downsample else 1,
+ padding='same',
+ depth_multiplier=1,
+ dilation_rate=self._dilation_rate,
+ use_bias=False,
+ depthwise_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer
+ ),
+ depthwise_regularizer=self._depthwise_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ )
+ # No activation -> quantized norm should be okay.
+ self._start_dw_norm = self._norm_with_quantize(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ )
+
+ # Expansion with 1x1 convs.
+ expand_filters = nn_layers.make_divisible(
+ self._in_filters * self._expand_ratio, self._divisible_by
+ )
+
+ self._expand_conv = helper.Conv2DQuantized(
+ filters=expand_filters,
+ kernel_size=1,
+ strides=1,
+ padding='same',
+ use_bias=False,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ )
+ self._expand_norm = helper.norm_by_activation(
+ self._activation, self._norm_with_quantize, self._norm
+ )(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ )
+ self._expand_act = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(self._activation, use_keras_layer=True),
+ configs.Default8BitActivationQuantizeConfig(),
+ )
+
+ # Middle depthwise conv.
+ if self._middle_dw_kernel_size:
+ self._middle_dw_conv = helper.DepthwiseConv2DQuantized(
+ kernel_size=self._middle_dw_kernel_size,
+ strides=self._strides if self._middle_dw_downsample else 1,
+ padding='same',
+ depth_multiplier=1,
+ dilation_rate=self._dilation_rate,
+ use_bias=False,
+ depthwise_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer
+ ),
+ depthwise_regularizer=self._depthwise_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ )
+ self._middle_dw_norm = helper.norm_by_activation(
+ self._activation, self._norm_with_quantize, self._norm
+ )(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ )
+ self._middle_dw_act = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(
+ self._depthwise_activation, use_keras_layer=True
+ ),
+ configs.Default8BitActivationQuantizeConfig(),
+ )
+
+ # Projection with 1x1 convs.
+ self._proj_conv = helper.Conv2DQuantized(
+ filters=self._out_filters,
+ kernel_size=1,
+ strides=1,
+ padding='same',
+ use_bias=False,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ )
+ # No activation -> quantized norm should be okay.
+ self._proj_norm = self._norm_with_quantize(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ )
+
+ # Ending depthwise conv.
+ if self._end_dw_kernel_size:
+ self._end_dw_conv = helper.DepthwiseConv2DQuantized(
+ kernel_size=self._end_dw_kernel_size,
+ strides=1,
+ padding='same',
+ depth_multiplier=1,
+ dilation_rate=self._dilation_rate,
+ use_bias=False,
+ depthwise_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer
+ ),
+ depthwise_regularizer=self._depthwise_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ )
+ self._end_dw_norm = self._norm_with_quantize(
+ axis=self._bn_axis,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ )
+
+ if self._use_layer_scale:
+ raise NotImplementedError
+
+ if self._stochastic_depth_drop_rate:
+ self._stochastic_depth = nn_layers.StochasticDepth(
+ self._stochastic_depth_drop_rate
+ )
+ else:
+ self._stochastic_depth = None
+
+ super().build(input_shape)
+
+ def get_config(self):
+ config = {
+ 'in_filters': self._in_filters,
+ 'out_filters': self._out_filters,
+ 'expand_ratio': self._expand_ratio,
+ 'strides': self._strides,
+ 'middle_dw_downsample': self._middle_dw_downsample,
+ 'start_dw_kernel_size': self._start_dw_kernel_size,
+ 'middle_dw_kernel_size': self._middle_dw_kernel_size,
+ 'end_dw_kernel_size': self._end_dw_kernel_size,
+ 'divisible_by': self._divisible_by,
+ 'stochastic_depth_drop_rate': self._stochastic_depth_drop_rate,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_regularizer': self._bias_regularizer,
+ 'activation': self._activation,
+ 'depthwise_activation': self._depthwise_activation,
+ 'dilation_rate': self._dilation_rate,
+ 'use_sync_bn': self._use_sync_bn,
+ 'regularize_depthwise': self._regularize_depthwise,
+ 'use_residual': self._use_residual,
+ 'use_layer_scale': self._use_layer_scale,
+ 'layer_scale_init_value': self._layer_scale_init_value,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'output_intermediate_endpoints': self._output_intermediate_endpoints,
+ }
+ base_config = super().get_config()
+ return {**base_config, **config}
+
+ def call(self, inputs, training=None):
+ shortcut = inputs
+ x = inputs
+ if self._start_dw_kernel_size:
+ x = self._start_dw_conv(x)
+ x = self._start_dw_norm(x)
+
+ x = self._expand_conv(x)
+ x = self._expand_norm(x)
+ x = self._expand_act(x)
+
+ if self._middle_dw_kernel_size:
+ x = self._middle_dw_conv(x)
+ x = self._middle_dw_norm(x)
+ x = self._middle_dw_act(x)
+
+ x = self._proj_conv(x)
+ x = self._proj_norm(x)
+
+ if self._end_dw_kernel_size:
+ x = self._end_dw_conv(x)
+ x = self._end_dw_norm(x)
+
+ if self._use_layer_scale:
+ x = self._layer_scale(x)
+
+ if (
+ self._use_residual
+ and self._in_filters == self._out_filters
+ and self._strides == 1
+ ):
+ if self._stochastic_depth:
+ x = self._stochastic_depth(x, training=training)
+ x = x + shortcut
+
+ # Return empty intermediate endpoints to be compatible with other blocks.
+ if self._output_intermediate_endpoints:
+ return x, {}
+ return x
+
+
+MaybeDwInvertedBottleneckBlockQuantized = (
+ UniversalInvertedBottleneckBlockQuantized
+)
diff --git a/official/projects/qat/vision/modeling/layers/nn_blocks_test.py b/official/projects/qat/vision/modeling/layers/nn_blocks_test.py
index be7389b7aed..18c96307608 100644
--- a/official/projects/qat/vision/modeling/layers/nn_blocks_test.py
+++ b/official/projects/qat/vision/modeling/layers/nn_blocks_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,8 @@
"""Tests for nn_blocks."""
from typing import Any, Iterable, Tuple
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -45,7 +44,7 @@ def test_bottleneck_block_creation(self, block_fn, strides, use_projection,
stochastic_depth_drop_rate, se_ratio):
input_size = 128
filter_size = 256
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(input_size, input_size, filter_size * 4), batch_size=1)
block = block_fn(
filter_size,
@@ -73,7 +72,7 @@ def test_invertedbottleneck_block_creation(
input_size = 128
in_filters = 24
out_filters = 40
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(input_size, input_size, in_filters), batch_size=1)
block = block_fn(
in_filters=in_filters,
@@ -90,6 +89,174 @@ def test_invertedbottleneck_block_creation(
[1, input_size // strides, input_size // strides, out_filters],
features.shape.as_list())
+ @parameterized.parameters(
+ (2, True, 0, 5, 0, 12, 12, 2),
+ (2, False, 5, 0, 0, 12, 18, 4),
+ (1, True, 0, 0, 0, 12, 12, 6),
+ (1, True, 3, 0, 0, 12, 18, 2),
+ (1, True, 3, 3, 0, 12, 12, 4),
+ (1, True, 3, 3, 3, 12, 18, 6),
+ (1, True, 0, 3, 3, 12, 12, 2),
+ (1, True, 0, 0, 3, 12, 18, 4),
+ (1, True, 3, 0, 3, 12, 12, 6),
+ )
+ def test_maybedwinvertedbottleneck_block_creation(
+ self,
+ strides,
+ middle_dw_downsample,
+ start_dw_kernel_size,
+ middle_dw_kernel_size,
+ end_dw_kernel_size,
+ in_filters,
+ out_filters,
+ expand_ratio,
+ ):
+ input_size = 128
+ inputs = tf_keras.Input(
+ shape=(input_size, input_size, in_filters), batch_size=1
+ )
+ block = nn_blocks.MaybeDwInvertedBottleneckBlockQuantized(
+ in_filters=in_filters,
+ out_filters=out_filters,
+ expand_ratio=expand_ratio,
+ strides=strides,
+ middle_dw_downsample=middle_dw_downsample,
+ start_dw_kernel_size=start_dw_kernel_size,
+ middle_dw_kernel_size=middle_dw_kernel_size,
+ end_dw_kernel_size=end_dw_kernel_size,
+ )
+
+ features = block(inputs)
+
+ self.assertAllEqual(
+ [1, input_size // strides, input_size // strides, out_filters],
+ features.shape.as_list(),
+ )
+
+ @parameterized.parameters(
+ (2, True, 0, 5, 0, 12, 12, 2),
+ (2, False, 5, 0, 0, 12, 18, 4),
+ (1, True, 0, 0, 0, 12, 12, 6),
+ (1, True, 3, 0, 0, 12, 18, 2),
+ (1, True, 3, 3, 0, 12, 12, 4),
+ (1, True, 3, 3, 3, 12, 18, 6),
+ (1, True, 0, 3, 3, 12, 12, 2),
+ (1, True, 0, 0, 3, 12, 18, 4),
+ (1, True, 3, 0, 3, 12, 12, 6),
+ )
+ def test_maybedwinvertedbottleneck_block_forward_pass_no_nans(
+ self,
+ strides,
+ middle_dw_downsample,
+ start_dw_kernel_size,
+ middle_dw_kernel_size,
+ end_dw_kernel_size,
+ in_filters,
+ out_filters,
+ expand_ratio,
+ ):
+ tf.random.set_seed(42)
+
+ input_size = 128
+ input_shape = (input_size, input_size, in_filters)
+ output_shape = [
+ 1,
+ input_size // strides,
+ input_size // strides,
+ out_filters,
+ ]
+ inputs = tf_keras.Input(shape=input_shape, batch_size=1)
+ block = nn_blocks.MaybeDwInvertedBottleneckBlockQuantized(
+ in_filters=in_filters,
+ out_filters=out_filters,
+ expand_ratio=expand_ratio,
+ strides=strides,
+ middle_dw_downsample=middle_dw_downsample,
+ start_dw_kernel_size=start_dw_kernel_size,
+ middle_dw_kernel_size=middle_dw_kernel_size,
+ end_dw_kernel_size=end_dw_kernel_size,
+ )
+ features = block(inputs)
+ self.assertAllEqual(features.shape.as_list(), output_shape)
+
+ model = tf_keras.Model(inputs=inputs, outputs=features)
+ input_data = tf.random.uniform(
+ (1, input_size, input_size, in_filters), minval=-1.0, maxval=1.0
+ )
+ predicted_outputs = model.predict(input_data)
+ self.assertAllEqual(
+ tf.math.is_nan(predicted_outputs),
+ tf.constant(False, shape=output_shape),
+ )
+
+ @parameterized.parameters(
+ (2, True, 0, 5, 0, 12, 12, 2),
+ (2, False, 5, 0, 0, 12, 18, 4),
+ (1, True, 0, 0, 0, 12, 12, 6),
+ (1, True, 3, 0, 0, 12, 18, 2),
+ (1, True, 3, 3, 0, 12, 12, 4),
+ (1, True, 3, 3, 3, 12, 18, 6),
+ (1, True, 0, 3, 3, 12, 12, 2),
+ (1, True, 0, 0, 3, 12, 18, 4),
+ (1, True, 3, 0, 3, 12, 12, 6),
+ )
+ def test_maybedwinvertedbottleneck_block_backward_pass_no_nans(
+ self,
+ strides,
+ middle_dw_downsample,
+ start_dw_kernel_size,
+ middle_dw_kernel_size,
+ end_dw_kernel_size,
+ in_filters,
+ out_filters,
+ expand_ratio,
+ ):
+ tf.random.set_seed(42)
+
+ input_size = 128
+ inputs = tf_keras.Input(
+ shape=(input_size, input_size, in_filters), batch_size=1
+ )
+ output_shape = [
+ 1,
+ input_size // strides,
+ input_size // strides,
+ out_filters,
+ ]
+ block = nn_blocks.MaybeDwInvertedBottleneckBlockQuantized(
+ in_filters=in_filters,
+ out_filters=out_filters,
+ expand_ratio=expand_ratio,
+ strides=strides,
+ middle_dw_downsample=middle_dw_downsample,
+ start_dw_kernel_size=start_dw_kernel_size,
+ middle_dw_kernel_size=middle_dw_kernel_size,
+ end_dw_kernel_size=end_dw_kernel_size,
+ )
+ features = block(inputs)
+ self.assertAllEqual(features.shape.as_list(), output_shape)
+ model = tf_keras.Model(inputs=inputs, outputs=features)
+ model.compile(
+ optimizer=tf_keras.optimizers.Adam(),
+ loss=tf_keras.losses.MeanSquaredError(),
+ metrics=[tf_keras.metrics.MeanSquaredError()],
+ )
+ input_train = tf.random.uniform(
+ (1, input_size, input_size, in_filters), minval=-1.0, maxval=1.0
+ )
+ output_train = tf.random.uniform(output_shape, minval=-1.0, maxval=1.0)
+ input_valid = tf.random.uniform(
+ (1, input_size, input_size, in_filters), minval=-1.0, maxval=1.0
+ )
+ output_valid = tf.random.uniform(output_shape, minval=-1.0, maxval=1.0)
+ model.fit(
+ input_train,
+ output_train,
+ batch_size=1,
+ epochs=1,
+ validation_data=(input_valid, output_valid),
+ )
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/projects/qat/vision/modeling/layers/nn_layers.py b/official/projects/qat/vision/modeling/layers/nn_layers.py
index 42d5f6ea9bc..9abfe72826c 100644
--- a/official/projects/qat/vision/modeling/layers/nn_layers.py
+++ b/official/projects/qat/vision/modeling/layers/nn_layers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,14 +14,16 @@
"""Contains common building blocks for neural networks."""
-from typing import Callable, Dict, List, Mapping, Optional, Sequence, Tuple, Union
+import enum
+from typing import Any, Callable, Dict, List, Mapping, Optional, Sequence, Tuple, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.modeling import tf_utils
from official.projects.qat.vision.quantization import configs
from official.projects.qat.vision.quantization import helper
+from official.vision.modeling import heads
from official.vision.modeling.decoders import aspp
from official.vision.modeling.layers import nn_layers
@@ -31,10 +33,18 @@
Activation = Union[str, Callable]
-@tf.keras.utils.register_keras_serializable(package='Vision')
+# String constants.
+class FeatureFusion(str, enum.Enum):
+ PYRAMID_FUSION = 'pyramid_fusion'
+ PANOPTIC_FPN_FUSION = 'panoptic_fpn_fusion'
+ DEEPLABV3PLUS = 'deeplabv3plus'
+ DEEPLABV3PLUS_SUM_TO_MERGE = 'deeplabv3plus_sum_to_merge'
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
class SqueezeExcitationQuantized(
helper.LayerQuantizerHelper,
- tf.keras.layers.Layer):
+ tf_keras.layers.Layer):
"""Creates a squeeze and excitation layer."""
def __init__(self,
@@ -62,9 +72,9 @@ def __init__(self,
use_3d_input: A `bool` of whether input is 2D or 3D image.
kernel_initializer: A `str` of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2d.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2d.
Default to None.
activation: A `str` name of the activation function.
gating_activation: A `str` name of the activation function for final
@@ -86,7 +96,7 @@ def __init__(self,
self._kernel_initializer = kernel_initializer
self._kernel_regularizer = kernel_regularizer
self._bias_regularizer = bias_regularizer
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
if not use_3d_input:
self._spatial_axis = [1, 2]
else:
@@ -108,10 +118,14 @@ def _create_gating_activation_layer(self):
tf_utils.get_activation('relu6', use_keras_layer=True),
configs.Default8BitActivationQuantizeConfig())
else:
- self._gating_activation_layer = tfmot.quantization.keras.QuantizeWrapperV2(
- tf_utils.get_activation(
- self._gating_activation, use_keras_layer=True),
- configs.Default8BitActivationQuantizeConfig())
+ self._gating_activation_layer = (
+ tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(
+ self._gating_activation, use_keras_layer=True
+ ),
+ configs.Default8BitActivationQuantizeConfig(),
+ )
+ )
def _apply_gating_activation_layer(
self, x: tf.Tensor, training: bool) -> tf.Tensor:
@@ -135,7 +149,7 @@ def build(self, input_shape):
strides=1,
padding='same',
use_bias=True,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activation=helper.NoOpActivation())
@@ -146,13 +160,13 @@ def build(self, input_shape):
strides=1,
padding='same',
use_bias=True,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer,
bias_regularizer=self._bias_regularizer,
activation=helper.NoOpActivation())
self._multiply = tfmot.quantization.keras.QuantizeWrapperV2(
- tf.keras.layers.Multiply(),
+ tf_keras.layers.Multiply(),
configs.Default8BitQuantizeConfig([], [], True))
self._reduce_mean_quantizer = (
tfmot.quantization.keras.quantizers.MovingAverageQuantizer(
@@ -190,13 +204,13 @@ def call(self, inputs, training=None):
x = self._reduce_mean_quantizer(
x, training, self._reduce_mean_quantizer_vars)
x = self._activation_layer(self._se_reduce(x))
- x = self._apply_gating_activation_layer(self._se_expand(x), training)
+ x = self._apply_gating_activation_layer(self._se_expand(x), training) # pyrefly: ignore[bad-argument-type]
x = self._multiply([x, inputs])
return x
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class SegmentationHeadQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class SegmentationHeadQuantized(tf_keras.layers.Layer):
"""Creates a segmentation head."""
def __init__(
@@ -215,11 +229,12 @@ def __init__(
low_level_num_filters: int = 48,
num_decoder_filters: int = 256,
activation: str = 'relu',
+ logit_activation: Optional[str] = None,
use_sync_bn: bool = False,
norm_momentum: float = 0.99,
norm_epsilon: float = 0.001,
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
**kwargs):
"""Initializes a segmentation head.
@@ -237,10 +252,11 @@ def __init__(
prediction layer.
upsample_factor: An `int` number to specify the upsampling factor to
generate finer mask. Default 1 means no upsampling is applied.
- feature_fusion: One of `deeplabv3plus`, `pyramid_fusion`, or None. If
- `deeplabv3plus`, features from decoder_features[level] will be fused
- with low level feature maps from backbone. If `pyramid_fusion`,
- multiscale features will be resized and fused at the target level.
+ feature_fusion: One of `deeplabv3plus`, `deeplabv3plus_sum_to_merge`,
+ `pyramid_fusion`, or None. If `deeplabv3plus`, features from
+ decoder_features[level] will be fused with low level feature maps from
+ backbone. If `pyramid_fusion`, multiscale features will be resized and
+ fused at the target level.
decoder_min_level: An `int` of minimum level from decoder to use in
feature fusion. It is only used when feature_fusion is set to
`panoptic_fpn_fusion`.
@@ -256,13 +272,14 @@ def __init__(
It is only used when feature_fusion is set to `panoptic_fpn_fusion`.
activation: A `str` that indicates which activation is used, e.g. 'relu',
'swish', etc.
+ logit_activation: Unused.
use_sync_bn: A `bool` that indicates whether to use synchronized batch
normalization across different replicas.
norm_momentum: A `float` of normalization momentum for the moving average.
norm_epsilon: A `float` added to variance to avoid dividing by zero.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default is None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2D.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
**kwargs: Additional keyword arguments to be passed.
"""
super().__init__(**kwargs)
@@ -288,13 +305,16 @@ def __init__(
'kernel_regularizer': kernel_regularizer,
'bias_regularizer': bias_regularizer,
}
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
self._activation_layer = tfmot.quantization.keras.QuantizeWrapperV2(
tf_utils.get_activation(activation, use_keras_layer=True),
configs.Default8BitActivationQuantizeConfig())
+ if logit_activation:
+ raise ValueError('Unused logit_activation option inherited from '
+ 'vision SegmentationHead modeling config.')
def build(self, input_shape: Sequence[tf.TensorShape]):
"""Creates the variables of the segmentation head."""
@@ -303,7 +323,7 @@ def build(self, input_shape: Sequence[tf.TensorShape]):
# fusion type is `deeplabv3plus`.
backbone_shape = input_shape[0]
use_depthwise_convolution = self._config_dict['use_depthwise_convolution']
- random_initializer = tf.keras.initializers.RandomNormal(stddev=0.01)
+ random_initializer = tf_keras.initializers.RandomNormal(stddev=0.01)
conv_kwargs = {
'kernel_size': 3 if not use_depthwise_convolution else 1,
'padding': 'same',
@@ -313,9 +333,9 @@ def build(self, input_shape: Sequence[tf.TensorShape]):
}
norm_layer = (
- tf.keras.layers.experimental.SyncBatchNormalization
+ tf_keras.layers.experimental.SyncBatchNormalization
if self._config_dict['use_sync_bn'] else
- tf.keras.layers.BatchNormalization)
+ tf_keras.layers.BatchNormalization)
norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
norm_no_quantize = helper.BatchNormalizationNoQuantized(norm_layer)
norm = helper.norm_by_activation(self._config_dict['activation'],
@@ -327,13 +347,15 @@ def build(self, input_shape: Sequence[tf.TensorShape]):
'epsilon': self._config_dict['norm_epsilon'],
}
- if self._config_dict['feature_fusion'] == 'deeplabv3plus':
+ if self._config_dict['feature_fusion'] in [
+ FeatureFusion.DEEPLABV3PLUS, FeatureFusion.DEEPLABV3PLUS_SUM_TO_MERGE
+ ]:
# Deeplabv3+ feature fusion layers.
self._dlv3p_conv = helper.Conv2DQuantized(
kernel_size=1,
padding='same',
use_bias=False,
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_utils.clone_initializer(random_initializer),
kernel_regularizer=self._config_dict['kernel_regularizer'],
name='segmentation_head_deeplabv3p_fusion_conv',
filters=self._config_dict['low_level_num_filters'],
@@ -353,7 +375,8 @@ def build(self, input_shape: Sequence[tf.TensorShape]):
kernel_size=3,
padding='same',
use_bias=False,
- depthwise_initializer=random_initializer,
+ depthwise_initializer=tf_utils.clone_initializer(
+ random_initializer),
depthwise_regularizer=self._config_dict['kernel_regularizer'],
depth_multiplier=1,
activation=helper.NoOpActivation()))
@@ -375,7 +398,7 @@ def build(self, input_shape: Sequence[tf.TensorShape]):
kernel_size=self._config_dict['prediction_kernel_size'],
padding='same',
bias_initializer=tf.zeros_initializer(),
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_utils.clone_initializer(random_initializer),
kernel_regularizer=self._config_dict['kernel_regularizer'],
bias_regularizer=self._config_dict['bias_regularizer'],
activation=helper.NoOpActivation())
@@ -384,10 +407,12 @@ def build(self, input_shape: Sequence[tf.TensorShape]):
size=(self._config_dict['upsample_factor'],
self._config_dict['upsample_factor']),
interpolation='nearest')
- self._resizing_layer = tf.keras.layers.Resizing(
+ self._resizing_layer = helper.ResizingQuantized(
backbone_shape[1], backbone_shape[2], interpolation='bilinear')
self._concat_layer = helper.ConcatenateQuantized(axis=self._bn_axis)
+ self._add_layer = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_keras.layers.Add(), configs.Default8BitQuantizeConfig([], [], True))
super().build(input_shape)
@@ -412,14 +437,16 @@ def call(self, inputs: Tuple[Union[tf.Tensor, Mapping[str, tf.Tensor]],
segmentation prediction mask: A `tf.Tensor` of the segmentation mask
scores predicted from input features.
"""
- if self._config_dict['feature_fusion'] in ('pyramid_fusion',
- 'panoptic_fpn_fusion'):
+ if self._config_dict['feature_fusion'] in (
+ FeatureFusion.PYRAMID_FUSION, FeatureFusion.PANOPTIC_FPN_FUSION):
raise ValueError(
'The feature fusion method `pyramid_fusion` is not supported in QAT.')
backbone_output = inputs[0]
decoder_output = inputs[1]
- if self._config_dict['feature_fusion'] == 'deeplabv3plus':
+ if self._config_dict['feature_fusion'] in {
+ FeatureFusion.DEEPLABV3PLUS, FeatureFusion.DEEPLABV3PLUS_SUM_TO_MERGE
+ }:
# deeplabv3+ feature fusion.
x = decoder_output[str(self._config_dict['level'])] if isinstance(
decoder_output, dict) else decoder_output
@@ -429,7 +456,10 @@ def call(self, inputs: Tuple[Union[tf.Tensor, Mapping[str, tf.Tensor]],
y = self._activation_layer(y)
x = self._resizing_layer(x)
x = tf.cast(x, dtype=y.dtype)
- x = self._concat_layer([x, y])
+ if self._config_dict['feature_fusion'] == FeatureFusion.DEEPLABV3PLUS:
+ x = self._concat_layer([x, y])
+ else:
+ x = self._add_layer([x, y])
else:
x = decoder_output[str(self._config_dict['level'])] if isinstance(
decoder_output, dict) else decoder_output
@@ -453,7 +483,7 @@ def from_config(cls, config):
return cls(**config)
-@tf.keras.utils.register_keras_serializable(package='Vision')
+@tf_keras.utils.register_keras_serializable(package='Vision')
class SpatialPyramidPoolingQuantized(nn_layers.SpatialPyramidPooling):
"""Implements the quantized Atrous Spatial Pyramid Pooling.
@@ -475,7 +505,7 @@ def __init__(
activation: str = 'relu',
dropout: float = 0.5,
kernel_initializer: str = 'GlorotUniform',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
interpolation: str = 'bilinear',
use_depthwise_convolution: bool = False,
**kwargs):
@@ -531,8 +561,8 @@ def build(self, input_shape):
channels = input_shape[3]
norm_layer = (
- tf.keras.layers.experimental.SyncBatchNormalization
- if self._use_sync_bn else tf.keras.layers.BatchNormalization)
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if self._use_sync_bn else tf_keras.layers.BatchNormalization)
norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
norm_no_quantize = helper.BatchNormalizationNoQuantized(norm_layer)
norm = helper.norm_by_activation(self._activation, norm_with_quantize,
@@ -543,7 +573,7 @@ def build(self, input_shape):
conv1 = helper.Conv2DQuantized(
filters=self._output_channels,
kernel_size=(1, 1),
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer,
use_bias=False,
activation=helper.NoOpActivation())
@@ -564,7 +594,8 @@ def build(self, input_shape):
kernel_size=kernel_size,
padding='same',
depthwise_regularizer=self._kernel_regularizer,
- depthwise_initializer=self._kernel_initializer,
+ depthwise_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
dilation_rate=dilation_rate,
use_bias=False,
activation=helper.NoOpActivation())
@@ -576,7 +607,8 @@ def build(self, input_shape):
kernel_size=kernel_size,
padding='same',
kernel_regularizer=self._kernel_regularizer,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
dilation_rate=dilation_rate,
use_bias=False,
activation=helper.NoOpActivation())
@@ -599,7 +631,8 @@ def build(self, input_shape):
conv2 = helper.Conv2DQuantized(
filters=self._output_channels,
kernel_size=(1, 1),
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer,
use_bias=False,
activation=helper.NoOpActivation())
@@ -616,23 +649,24 @@ def build(self, input_shape):
helper.Conv2DQuantized(
filters=self._output_channels,
kernel_size=(1, 1),
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer,
use_bias=False,
activation=helper.NoOpActivation()),
- norm_with_quantize(
+ norm(
axis=self._bn_axis,
momentum=self._batchnorm_momentum,
epsilon=self._batchnorm_epsilon)
]
- self._dropout_layer = tf.keras.layers.Dropout(rate=self._dropout)
+ self._dropout_layer = tf_keras.layers.Dropout(rate=self._dropout)
self._concat_layer = helper.ConcatenateQuantized(axis=-1)
def call(self,
inputs: tf.Tensor,
training: Optional[bool] = None) -> tf.Tensor:
if training is None:
- training = tf.keras.backend.learning_phase()
+ training = tf_keras.backend.learning_phase()
result = []
for i, layers in enumerate(self.aspp_layers):
x = inputs
@@ -649,11 +683,11 @@ def call(self,
x = self._concat_layer(result)
for layer in self._projection:
x = layer(x, training=training)
- x = self._activation_fn_no_quant(x)
+ x = self._activation_fn(x)
return self._dropout_layer(x)
-@tf.keras.utils.register_keras_serializable(package='Vision')
+@tf_keras.utils.register_keras_serializable(package='Vision')
class ASPPQuantized(aspp.ASPP):
"""Creates a quantized Atrous Spatial Pyramid Pooling (ASPP) layer."""
@@ -669,7 +703,7 @@ def __init__(
activation: str = 'relu',
dropout_rate: float = 0.0,
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
interpolation: str = 'bilinear',
use_depthwise_convolution: bool = False,
spp_layer_version: str = 'v1',
@@ -691,7 +725,7 @@ def __init__(
dropout_rate: A `float` rate for dropout regularization.
kernel_initializer: A `str` name of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default is None.
interpolation: A `str` of interpolation method. It should be one of
`bilinear`, `nearest`, `bicubic`, `area`, `lanczos3`, `lanczos5`,
@@ -748,3 +782,173 @@ def call(self, inputs: Union[tf.Tensor, Mapping[str,
level = str(self._config_dict['level'])
backbone_output = inputs[level] if isinstance(inputs, dict) else inputs
return self.aspp(backbone_output)
+
+
+class BatchNormalizationWrapper(tf_keras.layers.Wrapper):
+ """A BatchNormalizationWrapper that explicitly not folded.
+
+ It just added an identity depthwise conv right before the normalization.
+ As a result, given normalization op just folded into the identity depthwise
+ conv layer.
+
+ Note that it only used when the batch normalization folding is not working.
+ It makes quantize them as a 1x1 depthwise conv layer that just work as same
+ as inference mode for the normalization. (Basically mult and add for the BN.)
+ """
+
+ def call(self, inputs: tf.Tensor, *args: Any, **kwargs: Any) -> tf.Tensor:
+ channels = tf.shape(inputs)[-1]
+ x = tf.nn.depthwise_conv2d(
+ inputs, tf.ones([1, 1, channels, 1]), [1, 1, 1, 1], 'VALID')
+ outputs = self.layer.call(x, *args, **kwargs)
+ return outputs
+
+
+class MaskScoringQuantized(heads.MaskScoring):
+ """Creates a quantized mask scoring layer.
+
+ This implements mask scoring layer from the paper:
+
+ Zhaojin Huang, Lichao Huang, Yongchao Gong, Chang Huang, Xinggang Wang.
+ Mask Scoring R-CNN.
+ (https://arxiv.org/pdf/1903.00241.pdf)
+ """
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the variables of the mask scoring head."""
+ self._activation_layer = tfmot.quantization.keras.QuantizeWrapperV2(
+ tf_utils.get_activation(
+ self._config_dict['activation'], use_keras_layer=True
+ ),
+ configs.Default8BitActivationQuantizeConfig(),
+ )
+ conv_kwargs = {
+ 'filters': self._config_dict['num_filters'],
+ 'kernel_size': 3,
+ 'padding': 'same',
+ }
+ conv_kwargs.update({
+ 'kernel_initializer': tf_keras.initializers.VarianceScaling(
+ scale=2, mode='fan_out', distribution='untruncated_normal'
+ ),
+ 'bias_initializer': tf.zeros_initializer(),
+ 'kernel_regularizer': self._config_dict['kernel_regularizer'],
+ 'bias_regularizer': self._config_dict['bias_regularizer'],
+ })
+ norm_layer = (
+ tf_keras.layers.experimental.SyncBatchNormalization
+ if self._config_dict['use_sync_bn']
+ else tf_keras.layers.BatchNormalization
+ )
+ norm_with_quantize = helper.BatchNormalizationQuantized(norm_layer)
+ norm_no_quantize = helper.BatchNormalizationNoQuantized(norm_layer)
+ bn_op = helper.norm_by_activation(
+ self._config_dict['activation'], norm_with_quantize, norm_no_quantize
+ )
+ bn_kwargs = {
+ 'axis': self._bn_axis,
+ 'momentum': self._config_dict['norm_momentum'],
+ 'epsilon': self._config_dict['norm_epsilon'],
+ }
+
+ self._convs = []
+ self._conv_norms = []
+ for i in range(self._config_dict['num_convs']):
+ if self._config_dict['use_depthwise_convolution']:
+ self._convs.append(
+ helper.DepthwiseConv2DQuantized(
+ name='mask-scoring-depthwise-conv-{}'.format(i),
+ kernel_size=3,
+ padding='same',
+ use_bias=False,
+ depthwise_initializer=tf_keras.initializers.RandomNormal(
+ stddev=0.01),
+ depthwise_regularizer=self._config_dict['kernel_regularizer'],
+ depth_multiplier=1,
+ activation=helper.NoOpActivation()))
+ norm_name = 'mask-scoring-depthwise-bn-{}'.format(i)
+ self._conv_norms.append(bn_op(name=norm_name, **bn_kwargs))
+ conv_name = 'mask-scoring_{}'.format(i)
+ if 'kernel_initializer' in conv_kwargs:
+ conv_kwargs['kernel_initializer'] = tf_utils.clone_initializer(
+ conv_kwargs['kernel_initializer']
+ )
+ if self._config_dict['use_depthwise_convolution']:
+ conv_kwargs['kernel_size'] = 1
+ self._convs.append(
+ helper.Conv2DQuantized(
+ name=conv_name, activation=helper.NoOpActivation(), **conv_kwargs
+ )
+ )
+ bn_name = 'mask-scoring-bn_{}'.format(i)
+ self._conv_norms.append(bn_op(name=bn_name, **bn_kwargs))
+
+ self._fcs = []
+ self._fc_norms = []
+ for i in range(self._config_dict['num_fcs']):
+ fc_name = 'mask-scoring-fc_{}'.format(i)
+ self._fcs.append(
+ helper.DenseQuantized(
+ units=self._config_dict['fc_dims'],
+ kernel_initializer=tf_keras.initializers.VarianceScaling(
+ scale=1 / 3.0, mode='fan_out', distribution='uniform'
+ ),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'],
+ name=fc_name,
+ activation=helper.NoOpActivation(),
+ )
+ )
+ bn_name = 'mask-scoring-fc-bn_{}'.format(i)
+ self._fc_norms.append(bn_op(name=bn_name, **bn_kwargs))
+
+ self._classifier = helper.DenseOutputQuantized(
+ units=self._config_dict['num_classes'],
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
+ bias_initializer=tf.zeros_initializer(),
+ kernel_regularizer=self._config_dict['kernel_regularizer'],
+ bias_regularizer=self._config_dict['bias_regularizer'],
+ name='iou-scores',
+ )
+
+ self._resizing_layer = helper.ResizingQuantized(
+ self._config_dict['fc_input_size'][0],
+ self._config_dict['fc_input_size'][1],
+ interpolation='bilinear',
+ )
+
+ self._identity_layer = helper.IdentityQuantized(trainable=False)
+
+ super().build(input_shape)
+
+ def call(self, inputs: tf.Tensor, training: bool = None): # pytype: disable=annotation-type-mismatch
+ """Forward pass mask scoring head.
+
+ Args:
+ inputs: A `tf.Tensor` of the shape [batch_size, width, size, num_classes],
+ representing the segmentation logits.
+ training: a `bool` indicating whether it is in `training` mode.
+
+ Returns:
+ mask_scores: A `tf.Tensor` of predicted mask scores
+ [batch_size, num_classes].
+ """
+ x = tf.stop_gradient(inputs)
+ for conv, bn in zip(self._convs, self._conv_norms):
+ x = conv(x)
+ x = bn(x)
+ x = self._activation_layer(x)
+
+ x = self._resizing_layer(x)
+
+ _, h, w, filters = x.get_shape().as_list()
+ x = tf.reshape(x, [-1, h * w * filters])
+
+ for fc, bn in zip(self._fcs, self._fc_norms):
+ x = fc(x)
+ x = bn(x)
+ x = self._activation_layer(x)
+
+ ious = self._classifier(x)
+ ious = self._identity_layer(ious)
+ return ious
diff --git a/official/projects/qat/vision/modeling/layers/nn_layers_test.py b/official/projects/qat/vision/modeling/layers/nn_layers_test.py
index 56f96c19133..fcf834bd3f8 100644
--- a/official/projects/qat/vision/modeling/layers/nn_layers_test.py
+++ b/official/projects/qat/vision/modeling/layers/nn_layers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,8 @@
"""Tests for nn_layers."""
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.qat.vision.modeling.layers import nn_layers
@@ -24,12 +23,15 @@
class NNLayersTest(parameterized.TestCase, tf.test.TestCase):
@parameterized.parameters(
- ('deeplabv3plus', 1),
- ('deeplabv3plus', 2),
- ('deeplabv3', 1),
- ('deeplabv3', 2),
+ ('deeplabv3plus', 1, 128, 128),
+ ('deeplabv3plus', 2, 128, 128),
+ ('deeplabv3', 1, 128, 64),
+ ('deeplabv3', 2, 128, 64),
+ ('deeplabv3plus_sum_to_merge', 1, 64, 128),
+ ('deeplabv3plus_sum_to_merge', 2, 64, 128),
)
- def test_segmentation_head_creation(self, feature_fusion, upsample_factor):
+ def test_segmentation_head_creation(self, feature_fusion, upsample_factor,
+ low_level_num_filters, expected_shape):
input_size = 128
decoder_outupt_size = input_size // 2
@@ -42,14 +44,11 @@ def test_segmentation_head_creation(self, feature_fusion, upsample_factor):
level=4,
upsample_factor=upsample_factor,
low_level=2,
- low_level_num_filters=128,
+ low_level_num_filters=low_level_num_filters,
feature_fusion=feature_fusion)
features = segmentation_head((backbone_output, decoder_output))
- expected_shape = (
- input_size
- if feature_fusion == 'deeplabv3plus' else decoder_outupt_size)
self.assertAllEqual([
2, expected_shape * upsample_factor, expected_shape * upsample_factor, 5
], features.shape.as_list())
@@ -61,7 +60,7 @@ def test_segmentation_head_creation(self, feature_fusion, upsample_factor):
)
def test_spatial_pyramid_pooling_creation(self, pool_kernel_size,
dilation_rates):
- inputs = tf.keras.Input(shape=(64, 64, 128), dtype=tf.float32)
+ inputs = tf_keras.Input(shape=(64, 64, 128), dtype=tf.float32)
layer = nn_layers.SpatialPyramidPoolingQuantized(
output_channels=256,
dilation_rates=dilation_rates,
@@ -79,7 +78,7 @@ def test_spatial_pyramid_pooling_creation(self, pool_kernel_size,
)
def test_aspp_creation(self, level, dilation_rates, num_filters):
input_size = 128 // 2**level
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
endpoints = tf.random.uniform(
shape=(2, input_size, input_size, 64), dtype=tf.float32)
@@ -91,6 +90,43 @@ def test_aspp_creation(self, level, dilation_rates, num_filters):
self.assertAllEqual([2, input_size, input_size, num_filters],
feats.shape.as_list())
+ @parameterized.parameters(False, True)
+ def test_bnorm_wrapper_creation(self, use_sync_bn):
+ inputs = tf_keras.Input(shape=(64, 64, 128), dtype=tf.float32)
+ if use_sync_bn:
+ norm = tf_keras.layers.experimental.SyncBatchNormalization(axis=-1)
+ else:
+ norm = tf_keras.layers.BatchNormalization(axis=-1)
+ layer = nn_layers.BatchNormalizationWrapper(norm)
+ output = layer(inputs)
+ self.assertAllEqual([None, 64, 64, 128], output.shape)
+
+ @parameterized.parameters(
+ (1, 1, 64, [4, 4]),
+ (2, 1, 64, [4, 4]),
+ (3, 1, 64, [4, 4]),
+ (1, 2, 32, [8, 8]),
+ (2, 2, 32, [8, 8]),
+ (3, 2, 32, [8, 8]),
+ )
+ def test_mask_scoring_creation(
+ self, num_convs, num_fcs, num_filters, fc_input_size
+ ):
+ inputs = tf_keras.Input(shape=(64, 64, 16), dtype=tf.float32)
+
+ head = nn_layers.MaskScoringQuantized(
+ num_classes=2,
+ num_convs=num_convs,
+ num_filters=num_filters,
+ fc_dims=128,
+ num_fcs=num_fcs,
+ fc_input_size=fc_input_size,
+ use_depthwise_convolution=True,
+ )
+
+ scores = head(inputs)
+ self.assertAllEqual(scores.shape.as_list(), [None, 2])
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/projects/qat/vision/modeling/segmentation_model.py b/official/projects/qat/vision/modeling/segmentation_model.py
index 99511e2275d..5ed61344882 100644
--- a/official/projects/qat/vision/modeling/segmentation_model.py
+++ b/official/projects/qat/vision/modeling/segmentation_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,14 +15,13 @@
"""Build segmentation models."""
from typing import Any, Mapping, Union
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-layers = tf.keras.layers
+layers = tf_keras.layers
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class SegmentationModelQuantized(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class SegmentationModelQuantized(tf_keras.Model):
"""A Segmentation class model.
Input images are passed through backbone first. Decoder network is then
@@ -34,9 +33,9 @@ class SegmentationModelQuantized(tf.keras.Model):
different backbones, and decoders.
"""
- def __init__(self, backbone: tf.keras.Model, decoder: tf.keras.layers.Layer,
- head: tf.keras.layers.Layer,
- input_specs: tf.keras.layers.InputSpec, **kwargs):
+ def __init__(self, backbone: tf_keras.Model, decoder: tf_keras.layers.Layer,
+ head: tf_keras.layers.Layer,
+ input_specs: tf_keras.layers.InputSpec, **kwargs):
"""Segmentation initialization function.
Args:
@@ -46,7 +45,7 @@ def __init__(self, backbone: tf.keras.Model, decoder: tf.keras.layers.Layer,
input_specs: The shape specifications of input tensor.
**kwargs: keyword arguments to be passed.
"""
- inputs = tf.keras.Input(shape=input_specs.shape[1:], name=input_specs.name)
+ inputs = tf_keras.Input(shape=input_specs.shape[1:], name=input_specs.name)
backbone_features = backbone(inputs)
if decoder:
@@ -69,7 +68,7 @@ def __init__(self, backbone: tf.keras.Model, decoder: tf.keras.layers.Layer,
@property
def checkpoint_items(
- self) -> Mapping[str, Union[tf.keras.Model, tf.keras.layers.Layer]]:
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
"""Returns a dictionary of items to be additionally checkpointed."""
items = dict(backbone=self.backbone, head=self.head)
if self.decoder is not None:
diff --git a/official/projects/qat/vision/n_bit/__init__.py b/official/projects/qat/vision/n_bit/__init__.py
index 569b809ec7d..1e1d19256a1 100644
--- a/official/projects/qat/vision/n_bit/__init__.py
+++ b/official/projects/qat/vision/n_bit/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/qat/vision/n_bit/configs.py b/official/projects/qat/vision/n_bit/configs.py
index 941c3690f34..c0fa5bfdccc 100644
--- a/official/projects/qat/vision/n_bit/configs.py
+++ b/official/projects/qat/vision/n_bit/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,12 +15,12 @@
"""Default 8-bit QuantizeConfigs."""
from typing import Sequence, Callable, Tuple, Any, Dict
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
Quantizer = tfmot.quantization.keras.quantizers.Quantizer
-Layer = tf.keras.layers.Layer
+Layer = tf_keras.layers.Layer
Activation = Callable[[tf.Tensor], tf.Tensor]
WeightAndQuantizer = Tuple[tf.Variable, Quantizer]
ActivationAndQuantizer = Tuple[Activation, Quantizer]
@@ -242,12 +242,12 @@ def build(self,
min_weight = layer.add_weight(
name + '_min',
shape=(tensor_shape[-1],),
- initializer=tf.keras.initializers.Constant(-6.0),
+ initializer=tf_keras.initializers.Constant(-6.0),
trainable=False)
max_weight = layer.add_weight(
name + '_max',
shape=(tensor_shape[-1],),
- initializer=tf.keras.initializers.Constant(6.0),
+ initializer=tf_keras.initializers.Constant(6.0),
trainable=False)
return {'min_var': min_weight, 'max_var': max_weight}
@@ -302,7 +302,7 @@ def __init__(self, num_bits_weight: int = 8, num_bits_activation: int = 8):
self._num_bits_activation = num_bits_activation
def _assert_activation_layer(self, layer: Layer):
- if not isinstance(layer, tf.keras.layers.Activation):
+ if not isinstance(layer, tf_keras.layers.Activation):
raise RuntimeError(
'DefaultNBitActivationQuantizeConfig can only be used with '
'`keras.layers.Activation`.')
diff --git a/official/projects/qat/vision/n_bit/configs_test.py b/official/projects/qat/vision/n_bit/configs_test.py
index 5390f8d9c47..65d61eebd34 100644
--- a/official/projects/qat/vision/n_bit/configs_test.py
+++ b/official/projects/qat/vision/n_bit/configs_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,12 @@
"""Tests for configs.py."""
-# Import libraries
-
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
+from official.modeling import tf_utils
from official.projects.qat.vision.n_bit import configs
@@ -68,7 +67,7 @@ def _assert_kernel_equality(self, a, b):
class DefaultNBitQuantizeConfigTest(tf.test.TestCase, _TestHelper):
def _simple_dense_layer(self):
- layer = tf.keras.layers.Dense(2)
+ layer = tf_keras.layers.Dense(2)
layer.build(input_shape=(3,))
return layer
@@ -100,7 +99,7 @@ def testGetsQuantizeActivationsAndQuantizers(self):
def testSetsQuantizeWeights(self):
layer = self._simple_dense_layer()
- quantize_kernel = tf.keras.backend.variable(
+ quantize_kernel = tf_keras.backend.variable(
np.ones(layer.kernel.shape.as_list()))
num_bits_weight = 4
num_bits_activation = 4
@@ -113,7 +112,7 @@ def testSetsQuantizeWeights(self):
def testSetsQuantizeActivations(self):
layer = self._simple_dense_layer()
- quantize_activation = tf.keras.activations.relu
+ quantize_activation = tf_keras.activations.relu
num_bits_weight = 4
num_bits_activation = 4
@@ -125,7 +124,7 @@ def testSetsQuantizeActivations(self):
def testSetsQuantizeWeights_ErrorOnWrongNumberOfWeights(self):
layer = self._simple_dense_layer()
- quantize_kernel = tf.keras.backend.variable(
+ quantize_kernel = tf_keras.backend.variable(
np.ones(layer.kernel.shape.as_list()))
num_bits_weight = 4
num_bits_activation = 4
@@ -142,7 +141,7 @@ def testSetsQuantizeWeights_ErrorOnWrongNumberOfWeights(self):
def testSetsQuantizeWeights_ErrorOnWrongShapeOfWeight(self):
layer = self._simple_dense_layer()
- quantize_kernel = tf.keras.backend.variable(np.ones([1, 2]))
+ quantize_kernel = tf_keras.backend.variable(np.ones([1, 2]))
num_bits_weight = 4
num_bits_activation = 4
@@ -154,7 +153,7 @@ def testSetsQuantizeWeights_ErrorOnWrongShapeOfWeight(self):
def testSetsQuantizeActivations_ErrorOnWrongNumberOfActivations(self):
layer = self._simple_dense_layer()
- quantize_activation = tf.keras.activations.relu
+ quantize_activation = tf_keras.activations.relu
num_bits_weight = 4
num_bits_activation = 4
@@ -207,15 +206,19 @@ def testSerialization(self):
'num_bits_activation': 4
}
}
- serialized_quantize_config = tf.keras.utils.serialize_keras_object(
- quantize_config)
+ serialized_quantize_config = tf_utils.serialize_keras_object(
+ quantize_config
+ )
self.assertEqual(expected_config, serialized_quantize_config)
- quantize_config_from_config = tf.keras.utils.deserialize_keras_object(
- serialized_quantize_config,
- module_objects=globals(),
- custom_objects=configs._types_dict())
+ quantize_config_from_config = (
+ tf_utils.deserialize_keras_object(
+ serialized_quantize_config,
+ module_objects=globals(),
+ custom_objects=configs._types_dict(),
+ )
+ )
self.assertEqual(quantize_config, quantize_config_from_config)
diff --git a/official/projects/qat/vision/n_bit/nn_blocks.py b/official/projects/qat/vision/n_bit/nn_blocks.py
index 6f168fab7aa..d4592cb3ba9 100644
--- a/official/projects/qat/vision/n_bit/nn_blocks.py
+++ b/official/projects/qat/vision/n_bit/nn_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,8 @@
"""Contains quantized neural blocks for the QAT."""
from typing import Any, Dict, Optional, Sequence, Union
-# Import libraries
-
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.modeling import tf_utils
@@ -62,8 +60,8 @@ def constructor(*arg, **kwargs):
# This class is copied from modeling.layers.nn_blocks.BottleneckBlock and apply
# QAT.
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class BottleneckBlockNBitQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class BottleneckBlockNBitQuantized(tf_keras.layers.Layer):
"""A quantized standard bottleneck block."""
def __init__(self,
@@ -75,8 +73,8 @@ def __init__(self,
resnetd_shortcut: bool = False,
stochastic_depth_drop_rate: Optional[float] = None,
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: tf.keras.regularizers.Regularizer = None,
- bias_regularizer: tf.keras.regularizers.Regularizer = None,
+ kernel_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
+ bias_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
activation: str = 'relu',
use_sync_bn: bool = False,
norm_momentum: float = 0.99,
@@ -104,9 +102,9 @@ def __init__(self,
the stochastic depth layer.
kernel_initializer: A `str` of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2d.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2d.
Default to None.
activation: A `str` name of the activation function.
use_sync_bn: A `bool`. If True, use synchronized batch normalization.
@@ -138,23 +136,23 @@ def __init__(self,
self._num_bits_activation = num_bits_activation
if use_sync_bn:
self._norm = _quantize_wrapped_layer(
- tf.keras.layers.experimental.SyncBatchNormalization,
+ tf_keras.layers.experimental.SyncBatchNormalization,
configs.NoOpQuantizeConfig())
self._norm_with_quantize = _quantize_wrapped_layer(
- tf.keras.layers.experimental.SyncBatchNormalization,
+ tf_keras.layers.experimental.SyncBatchNormalization,
configs.DefaultNBitOutputQuantizeConfig(
num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
else:
self._norm = _quantize_wrapped_layer(
- tf.keras.layers.BatchNormalization,
+ tf_keras.layers.BatchNormalization,
configs.NoOpQuantizeConfig())
self._norm_with_quantize = _quantize_wrapped_layer(
- tf.keras.layers.BatchNormalization,
+ tf_keras.layers.BatchNormalization,
configs.DefaultNBitOutputQuantizeConfig(
num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -163,14 +161,14 @@ def __init__(self,
def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
"""Build variables and child layers to prepare for calling."""
conv2d_quantized = _quantize_wrapped_layer(
- tf.keras.layers.Conv2D,
+ tf_keras.layers.Conv2D,
configs.DefaultNBitConvQuantizeConfig(
['kernel'], ['activation'], False,
num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
if self._use_projection:
if self._resnetd_shortcut:
- self._shortcut0 = tf.keras.layers.AveragePooling2D(
+ self._shortcut0 = tf_keras.layers.AveragePooling2D(
pool_size=2, strides=self._strides, padding='same')
self._shortcut1 = conv2d_quantized(
filters=self._filters * 4,
@@ -279,7 +277,7 @@ def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
else:
self._stochastic_depth = None
self._add = tfmot.quantization.keras.QuantizeWrapperV2(
- tf.keras.layers.Add(),
+ tf_keras.layers.Add(),
configs.DefaultNBitQuantizeConfig(
[], [], True,
num_bits_weight=self._num_bits_weight,
@@ -348,8 +346,8 @@ def call(
# This class is copied from modeling.backbones.mobilenet.Conv2DBNBlock and apply
# QAT.
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class Conv2DBNBlockNBitQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class Conv2DBNBlockNBitQuantized(tf_keras.layers.Layer):
"""A quantized convolution block with batch normalization."""
def __init__(
@@ -360,8 +358,8 @@ def __init__(
use_bias: bool = False,
activation: str = 'relu6',
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
use_normalization: bool = True,
use_sync_bn: bool = False,
norm_momentum: float = 0.99,
@@ -382,9 +380,9 @@ def __init__(
activation: A `str` name of the activation function.
kernel_initializer: A `str` for kernel initializer of convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2D.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2D.
Default to None.
use_normalization: If True, use batch normalization.
use_sync_bn: If True, use synchronized batch normalization.
@@ -412,13 +410,13 @@ def __init__(
if use_sync_bn:
self._norm = _quantize_wrapped_layer(
- tf.keras.layers.experimental.SyncBatchNormalization,
+ tf_keras.layers.experimental.SyncBatchNormalization,
configs.NoOpQuantizeConfig())
else:
self._norm = _quantize_wrapped_layer(
- tf.keras.layers.BatchNormalization,
+ tf_keras.layers.BatchNormalization,
configs.NoOpQuantizeConfig())
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -447,7 +445,7 @@ def get_config(self) -> Dict[str, Any]:
def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
"""Build variables and child layers to prepare for calling."""
conv2d_quantized = _quantize_wrapped_layer(
- tf.keras.layers.Conv2D,
+ tf_keras.layers.Conv2D,
configs.DefaultNBitConvQuantizeConfig(
['kernel'], ['activation'], False,
num_bits_weight=self._num_bits_weight,
@@ -486,8 +484,8 @@ def call(
return self._activation_layer(x)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class InvertedBottleneckBlockNBitQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class InvertedBottleneckBlockNBitQuantized(tf_keras.layers.Layer):
"""A quantized inverted bottleneck block."""
def __init__(self,
@@ -532,9 +530,9 @@ def __init__(self,
the stochastic depth layer.
kernel_initializer: A `str` of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2d.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2d.
Default to None.
activation: A `str` name of the activation function.
se_inner_activation: A `str` name of squeeze-excitation inner activation.
@@ -592,23 +590,23 @@ def __init__(self,
if use_sync_bn:
self._norm = _quantize_wrapped_layer(
- tf.keras.layers.experimental.SyncBatchNormalization,
+ tf_keras.layers.experimental.SyncBatchNormalization,
configs.NoOpQuantizeConfig())
self._norm_with_quantize = _quantize_wrapped_layer(
- tf.keras.layers.experimental.SyncBatchNormalization,
+ tf_keras.layers.experimental.SyncBatchNormalization,
configs.DefaultNBitOutputQuantizeConfig(
num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
else:
self._norm = _quantize_wrapped_layer(
- tf.keras.layers.BatchNormalization,
+ tf_keras.layers.BatchNormalization,
configs.NoOpQuantizeConfig())
self._norm_with_quantize = _quantize_wrapped_layer(
- tf.keras.layers.BatchNormalization,
+ tf_keras.layers.BatchNormalization,
configs.DefaultNBitOutputQuantizeConfig(
num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -622,13 +620,13 @@ def __init__(self,
def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
"""Build variables and child layers to prepare for calling."""
conv2d_quantized = _quantize_wrapped_layer(
- tf.keras.layers.Conv2D,
+ tf_keras.layers.Conv2D,
configs.DefaultNBitConvQuantizeConfig(
['kernel'], ['activation'], False,
num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
depthwise_conv2d_quantized = _quantize_wrapped_layer(
- tf.keras.layers.DepthwiseConv2D,
+ tf_keras.layers.DepthwiseConv2D,
configs.DefaultNBitConvQuantizeConfig(
['depthwise_kernel'], ['activation'], False,
num_bits_weight=self._num_bits_weight,
@@ -729,7 +727,7 @@ def build(self, input_shape: Optional[Union[Sequence[int], tf.Tensor]]):
self._stochastic_depth_drop_rate)
else:
self._stochastic_depth = None
- self._add = tf.keras.layers.Add()
+ self._add = tf_keras.layers.Add()
super().build(input_shape)
diff --git a/official/projects/qat/vision/n_bit/nn_blocks_test.py b/official/projects/qat/vision/n_bit/nn_blocks_test.py
index e5778b4414f..b379456b335 100644
--- a/official/projects/qat/vision/n_bit/nn_blocks_test.py
+++ b/official/projects/qat/vision/n_bit/nn_blocks_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,8 @@
"""Tests for nn_blocks."""
from typing import Any, Iterable, Tuple
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
@@ -46,7 +45,7 @@ def test_bottleneck_block_creation(self, block_fn, strides, use_projection,
num_bits_weight, num_bits_activation):
input_size = 128
filter_size = 256
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(input_size, input_size, filter_size * 4), batch_size=1)
block = block_fn(
filter_size,
@@ -76,7 +75,7 @@ def test_invertedbottleneck_block_creation(
input_size = 128
in_filters = 24
out_filters = 40
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(input_size, input_size, in_filters), batch_size=1)
block = block_fn(
in_filters=in_filters,
diff --git a/official/projects/qat/vision/n_bit/nn_layers.py b/official/projects/qat/vision/n_bit/nn_layers.py
index feef66e7cd0..8d0e942d75f 100644
--- a/official/projects/qat/vision/n_bit/nn_layers.py
+++ b/official/projects/qat/vision/n_bit/nn_layers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Any, Callable, Dict, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.modeling import tf_utils
@@ -58,8 +58,8 @@ def constructor(*arg, **kwargs):
return constructor
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class SqueezeExcitationNBitQuantized(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class SqueezeExcitationNBitQuantized(tf_keras.layers.Layer):
"""Creates a squeeze and excitation layer."""
def __init__(self,
@@ -88,9 +88,9 @@ def __init__(self,
use_3d_input: A `bool` of whether input is 2D or 3D image.
kernel_initializer: A `str` of kernel_initializer for convolutional
layers.
- kernel_regularizer: A `tf.keras.regularizers.Regularizer` object for
+ kernel_regularizer: A `tf_keras.regularizers.Regularizer` object for
Conv2D. Default to None.
- bias_regularizer: A `tf.keras.regularizers.Regularizer` object for Conv2d.
+ bias_regularizer: A `tf_keras.regularizers.Regularizer` object for Conv2d.
Default to None.
activation: A `str` name of the activation function.
gating_activation: A `str` name of the activation function for final
@@ -113,7 +113,7 @@ def __init__(self,
self._bias_regularizer = bias_regularizer
self._num_bits_weight = num_bits_weight
self._num_bits_activation = num_bits_activation
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
if not use_3d_input:
self._spatial_axis = [1, 2]
else:
@@ -136,13 +136,13 @@ def __init__(self,
def build(self, input_shape):
conv2d_quantized = _quantize_wrapped_layer(
- tf.keras.layers.Conv2D,
+ tf_keras.layers.Conv2D,
configs.DefaultNBitConvQuantizeConfig(
['kernel'], ['activation'], False,
num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
conv2d_quantized_output_quantized = _quantize_wrapped_layer(
- tf.keras.layers.Conv2D,
+ tf_keras.layers.Conv2D,
configs.DefaultNBitConvQuantizeConfig(
['kernel'], ['activation'], True,
num_bits_weight=self._num_bits_weight,
@@ -174,7 +174,7 @@ def build(self, input_shape):
activation=NoOpActivation())
self._multiply = tfmot.quantization.keras.QuantizeWrapperV2(
- tf.keras.layers.Multiply(),
+ tf_keras.layers.Multiply(),
configs.DefaultNBitQuantizeConfig(
[], [], True, num_bits_weight=self._num_bits_weight,
num_bits_activation=self._num_bits_activation))
diff --git a/official/projects/qat/vision/n_bit/schemes.py b/official/projects/qat/vision/n_bit/schemes.py
index 31661f89e23..a10c7cc2c3b 100644
--- a/official/projects/qat/vision/n_bit/schemes.py
+++ b/official/projects/qat/vision/n_bit/schemes.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,7 @@
"""Quantization schemes."""
from typing import Type
-# Import libraries
-
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.projects.qat.vision.n_bit import configs
diff --git a/official/projects/qat/vision/quantization/__init__.py b/official/projects/qat/vision/quantization/__init__.py
index 67c06b5c832..98f6a0ccf5e 100644
--- a/official/projects/qat/vision/quantization/__init__.py
+++ b/official/projects/qat/vision/quantization/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/qat/vision/quantization/configs.py b/official/projects/qat/vision/quantization/configs.py
index 17eeb9c3fcc..b8f6c3dfa3d 100644
--- a/official/projects/qat/vision/quantization/configs.py
+++ b/official/projects/qat/vision/quantization/configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,12 +15,12 @@
"""Default 8-bit QuantizeConfigs."""
from typing import Sequence, Callable, Tuple, Any, Dict
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
Quantizer = tfmot.quantization.keras.quantizers.Quantizer
-Layer = tf.keras.layers.Layer
+Layer = tf_keras.layers.Layer
Activation = Callable[[tf.Tensor], tf.Tensor]
WeightAndQuantizer = Tuple[tf.Variable, Quantizer]
ActivationAndQuantizer = Tuple[Activation, Quantizer]
@@ -215,12 +215,12 @@ def build(self,
min_weight = layer.add_weight(
name + '_min',
shape=(tensor_shape[-1],),
- initializer=tf.keras.initializers.Constant(-6.0),
+ initializer=tf_keras.initializers.Constant(-6.0),
trainable=False)
max_weight = layer.add_weight(
name + '_max',
shape=(tensor_shape[-1],),
- initializer=tf.keras.initializers.Constant(6.0),
+ initializer=tf_keras.initializers.Constant(6.0),
trainable=False)
return {'min_var': min_weight, 'max_var': max_weight}
@@ -261,7 +261,7 @@ class Default8BitActivationQuantizeConfig(
"""
def _assert_activation_layer(self, layer: Layer):
- if not isinstance(layer, tf.keras.layers.Activation):
+ if not isinstance(layer, tf_keras.layers.Activation):
raise RuntimeError(
'Default8BitActivationQuantizeConfig can only be used with '
'`keras.layers.Activation`.')
diff --git a/official/projects/qat/vision/quantization/configs_test.py b/official/projects/qat/vision/quantization/configs_test.py
index d23a65e3a16..b031ebe2667 100644
--- a/official/projects/qat/vision/quantization/configs_test.py
+++ b/official/projects/qat/vision/quantization/configs_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,12 @@
"""Tests for configs.py."""
-# Import libraries
-
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
+from official.modeling import tf_utils
from official.projects.qat.vision.quantization import configs
@@ -68,7 +67,7 @@ def _assert_kernel_equality(self, a, b):
class Default8BitQuantizeConfigTest(tf.test.TestCase, _TestHelper):
def _simple_dense_layer(self):
- layer = tf.keras.layers.Dense(2)
+ layer = tf_keras.layers.Dense(2)
layer.build(input_shape=(3,))
return layer
@@ -96,7 +95,7 @@ def testGetsQuantizeActivationsAndQuantizers(self):
def testSetsQuantizeWeights(self):
layer = self._simple_dense_layer()
- quantize_kernel = tf.keras.backend.variable(
+ quantize_kernel = tf_keras.backend.variable(
np.ones(layer.kernel.shape.as_list()))
quantize_config = configs.Default8BitQuantizeConfig(
@@ -107,7 +106,7 @@ def testSetsQuantizeWeights(self):
def testSetsQuantizeActivations(self):
layer = self._simple_dense_layer()
- quantize_activation = tf.keras.activations.relu
+ quantize_activation = tf_keras.activations.relu
quantize_config = configs.Default8BitQuantizeConfig(
['kernel'], ['activation'], False)
@@ -117,7 +116,7 @@ def testSetsQuantizeActivations(self):
def testSetsQuantizeWeights_ErrorOnWrongNumberOfWeights(self):
layer = self._simple_dense_layer()
- quantize_kernel = tf.keras.backend.variable(
+ quantize_kernel = tf_keras.backend.variable(
np.ones(layer.kernel.shape.as_list()))
quantize_config = configs.Default8BitQuantizeConfig(
@@ -132,7 +131,7 @@ def testSetsQuantizeWeights_ErrorOnWrongNumberOfWeights(self):
def testSetsQuantizeWeights_ErrorOnWrongShapeOfWeight(self):
layer = self._simple_dense_layer()
- quantize_kernel = tf.keras.backend.variable(np.ones([1, 2]))
+ quantize_kernel = tf_keras.backend.variable(np.ones([1, 2]))
quantize_config = configs.Default8BitQuantizeConfig(
['kernel'], ['activation'], False)
@@ -142,7 +141,7 @@ def testSetsQuantizeWeights_ErrorOnWrongShapeOfWeight(self):
def testSetsQuantizeActivations_ErrorOnWrongNumberOfActivations(self):
layer = self._simple_dense_layer()
- quantize_activation = tf.keras.activations.relu
+ quantize_activation = tf_keras.activations.relu
quantize_config = configs.Default8BitQuantizeConfig(
['kernel'], ['activation'], False)
@@ -185,15 +184,19 @@ def testSerialization(self):
'quantize_output': False
}
}
- serialized_quantize_config = tf.keras.utils.serialize_keras_object(
- quantize_config)
+ serialized_quantize_config = tf_utils.serialize_keras_object(
+ quantize_config
+ )
self.assertEqual(expected_config, serialized_quantize_config)
- quantize_config_from_config = tf.keras.utils.deserialize_keras_object(
- serialized_quantize_config,
- module_objects=globals(),
- custom_objects=configs._types_dict())
+ quantize_config_from_config = (
+ tf_utils.deserialize_keras_object(
+ serialized_quantize_config,
+ module_objects=globals(),
+ custom_objects=configs._types_dict(),
+ )
+ )
self.assertEqual(quantize_config, quantize_config_from_config)
diff --git a/official/projects/qat/vision/quantization/helper.py b/official/projects/qat/vision/quantization/helper.py
index 456ea10018d..5aa20649ba0 100644
--- a/official/projects/qat/vision/quantization/helper.py
+++ b/official/projects/qat/vision/quantization/helper.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,14 +13,87 @@
# limitations under the License.
"""Quantization helpers."""
-from typing import Any, Dict
-import tensorflow as tf
+from __future__ import annotations
+
+import copy
+from typing import Any, Dict, List, Optional, Type, Union
+
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
from official.projects.qat.vision.quantization import configs
+_QUANTIZATION_WEIGHT_NAMES = [
+ 'output_max',
+ 'output_min',
+ 'optimizer_step',
+ 'kernel_min',
+ 'kernel_max',
+ 'add_three_min',
+ 'add_three_max',
+ 'divide_six_min',
+ 'divide_six_max',
+ 'depthwise_kernel_min',
+ 'depthwise_kernel_max',
+ 'pointwise_kernel_min',
+ 'pointwise_kernel_max',
+ 'reduce_mean_quantizer_vars_min',
+ 'reduce_mean_quantizer_vars_max',
+ 'quantize_layer_min',
+ 'quantize_layer_max',
+ 'quantize_layer_1_min',
+ 'quantize_layer_1_max',
+ 'quantize_layer_2_min',
+ 'quantize_layer_2_max',
+ 'quantize_layer_3_min',
+ 'quantize_layer_3_max',
+ 'post_activation_min',
+ 'post_activation_max',
+]
+
+_ORIGINAL_WEIGHT_NAME = [
+ 'kernel',
+ 'depthwise_kernel',
+ 'pointwise_kernel',
+ 'gamma',
+ 'beta',
+ 'moving_mean',
+ 'moving_variance',
+ 'bias',
+]
+
+
+def is_quantization_weight_name(name: str) -> bool:
+ simple_name = name.split('/')[-1].split(':')[0]
+ if simple_name in _QUANTIZATION_WEIGHT_NAMES:
+ return True
+ if simple_name in _ORIGINAL_WEIGHT_NAME:
+ return False
+ raise ValueError('Variable name {} is not supported.'.format(simple_name))
+
+
+def copy_original_weights(original_model: tf_keras.Model,
+ quantized_model: tf_keras.Model):
+ """Helper function that copy the original model weights to quantized model."""
+ original_weight_value = original_model.get_weights()
+ weight_values = quantized_model.get_weights()
+
+ original_idx = 0
+ for idx, weight in enumerate(quantized_model.weights):
+ if not is_quantization_weight_name(weight.name):
+ if original_idx >= len(original_weight_value):
+ raise ValueError('Not enought original model weights.')
+ weight_values[idx] = original_weight_value[original_idx]
+ original_idx = original_idx + 1
+
+ if original_idx < len(original_weight_value):
+ raise ValueError('Not enought quantized model weights.')
+
+ quantized_model.set_weights(weight_values)
+
+
class LayerQuantizerHelper(object):
"""Helper class that handles quantizers."""
@@ -94,36 +167,125 @@ def norm_by_activation(activation, norm_quantized, norm_no_quantized):
return norm_no_quantized
+class SeparableConv2DQuantized(tf_keras.layers.Layer):
+ """Quantized SeperableConv2D."""
+
+ def __init__(
+ self,
+ name: Optional[str] = None,
+ last_quantize: bool = False,
+ **conv_kwargs,
+ ):
+ """Initializes a SeparableConv2DQuantized.
+
+ Args:
+ name: The name of the layer.
+ last_quantize: A `bool` indicates whether add quantization for the output.
+ **conv_kwargs: A keyword arguments to be used for conv and dwconv.
+ """
+
+ super().__init__(name=name)
+ self._conv_kwargs = copy.deepcopy(conv_kwargs)
+ self._name = name
+ self._last_quantize = last_quantize
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Creates the child layers of the layer."""
+ depthwise_conv2d_quantized = quantize_wrapped_layer(
+ tf_keras.layers.DepthwiseConv2D,
+ configs.Default8BitConvQuantizeConfig(['depthwise_kernel'], [], True),
+ )
+ conv2d_quantized = quantize_wrapped_layer(
+ tf_keras.layers.Conv2D,
+ configs.Default8BitConvQuantizeConfig(
+ ['kernel'], [], self._last_quantize
+ ),
+ )
+
+ dwconv_kwargs = self._conv_kwargs.copy()
+ # Depthwise conv input filters is always equal to output filters.
+ # This filters argument only needed for the point-wise conv2d op.
+ del dwconv_kwargs['filters']
+ dwconv_kwargs.update({
+ 'activation': None,
+ 'use_bias': False,
+ })
+ self.dw_conv = depthwise_conv2d_quantized(name='dw', **dwconv_kwargs)
+
+ conv_kwargs = self._conv_kwargs.copy()
+ conv_kwargs.update({
+ 'kernel_size': (1, 1),
+ 'strides': (1, 1),
+ 'padding': 'valid',
+ 'groups': 1,
+ })
+
+ self.conv = conv2d_quantized(name='pw', **conv_kwargs)
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ """Call the separable conv layer."""
+ x = self.dw_conv(inputs)
+ outputs = self.conv(x)
+ return outputs
+
+ def get_config(self) -> Dict[str, Any]:
+ """Returns the config of the layer."""
+ config = self._conv_kwargs.copy()
+ config.update({
+ 'name': self._name,
+ 'last_quantize': self._last_quantize,
+ })
+ return config
+
+ @classmethod
+ def from_config(
+ cls: Type[SeparableConv2DQuantized], config: Dict[str, Any]
+ ) -> SeparableConv2DQuantized:
+ """Creates a layer from its config."""
+ return cls(**config)
+
+
Conv2DQuantized = quantize_wrapped_layer(
- tf.keras.layers.Conv2D,
+ tf_keras.layers.Conv2D,
configs.Default8BitConvQuantizeConfig(['kernel'], ['activation'], False))
Conv2DOutputQuantized = quantize_wrapped_layer(
- tf.keras.layers.Conv2D,
+ tf_keras.layers.Conv2D,
configs.Default8BitConvQuantizeConfig(['kernel'], ['activation'], True))
DepthwiseConv2DQuantized = quantize_wrapped_layer(
- tf.keras.layers.DepthwiseConv2D,
+ tf_keras.layers.DepthwiseConv2D,
configs.Default8BitConvQuantizeConfig(['depthwise_kernel'], ['activation'],
False))
DepthwiseConv2DOutputQuantized = quantize_wrapped_layer(
- tf.keras.layers.DepthwiseConv2D,
+ tf_keras.layers.DepthwiseConv2D,
configs.Default8BitConvQuantizeConfig(['depthwise_kernel'], ['activation'],
True))
GlobalAveragePooling2DQuantized = quantize_wrapped_layer(
- tf.keras.layers.GlobalAveragePooling2D,
+ tf_keras.layers.GlobalAveragePooling2D,
configs.Default8BitQuantizeConfig([], [], True))
AveragePooling2DQuantized = quantize_wrapped_layer(
- tf.keras.layers.AveragePooling2D,
+ tf_keras.layers.AveragePooling2D,
configs.Default8BitQuantizeConfig([], [], True))
ResizingQuantized = quantize_wrapped_layer(
- tf.keras.layers.Resizing, configs.Default8BitQuantizeConfig([], [], True))
+ tf_keras.layers.Resizing, configs.Default8BitQuantizeConfig([], [], True))
ConcatenateQuantized = quantize_wrapped_layer(
- tf.keras.layers.Concatenate, configs.Default8BitQuantizeConfig([], [],
+ tf_keras.layers.Concatenate, configs.Default8BitQuantizeConfig([], [],
True))
UpSampling2DQuantized = quantize_wrapped_layer(
- tf.keras.layers.UpSampling2D, configs.Default8BitQuantizeConfig([], [],
+ tf_keras.layers.UpSampling2D, configs.Default8BitQuantizeConfig([], [],
True))
ReshapeQuantized = quantize_wrapped_layer(
- tf.keras.layers.Reshape, configs.Default8BitQuantizeConfig([], [], True))
+ tf_keras.layers.Reshape, configs.Default8BitQuantizeConfig([], [], True))
+DenseQuantized = quantize_wrapped_layer(
+ tf_keras.layers.Dense,
+ configs.Default8BitQuantizeConfig(['kernel'], ['activation'], False),
+)
+DenseOutputQuantized = quantize_wrapped_layer(
+ tf_keras.layers.Dense,
+ configs.Default8BitQuantizeConfig(['kernel'], ['activation'], True),
+)
+IdentityQuantized = quantize_wrapped_layer(
+ tf_keras.layers.Identity, configs.Default8BitQuantizeConfig([], [], True)
+)
# pylint:disable=g-long-lambda
BatchNormalizationQuantized = lambda norm_layer: quantize_wrapped_layer(
diff --git a/official/projects/qat/vision/quantization/helper_test.py b/official/projects/qat/vision/quantization/helper_test.py
new file mode 100644
index 00000000000..bb032dca6c3
--- /dev/null
+++ b/official/projects/qat/vision/quantization/helper_test.py
@@ -0,0 +1,54 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for helper."""
+import numpy as np
+import tensorflow as tf, tf_keras
+
+import tensorflow_model_optimization as tfmot
+from official.projects.qat.vision.quantization import helper
+
+
+class HelperTest(tf.test.TestCase):
+
+ def create_simple_model(self):
+ return tf_keras.models.Sequential([
+ tf_keras.layers.Dense(8, input_shape=(16,)),
+ ])
+
+ def test_copy_original_weights_for_simple_model_with_custom_weights(self):
+ one_model = self.create_simple_model()
+ one_weights = [np.ones_like(weight) for weight in one_model.get_weights()]
+ one_model.set_weights(one_weights)
+
+ qat_model = tfmot.quantization.keras.quantize_model(
+ self.create_simple_model())
+ zero_weights = [np.zeros_like(weight) for weight in qat_model.get_weights()]
+ qat_model.set_weights(zero_weights)
+
+ helper.copy_original_weights(one_model, qat_model)
+
+ qat_model_weights = qat_model.get_weights()
+ count = 0
+ for idx, weight in enumerate(qat_model.weights):
+ if not helper.is_quantization_weight_name(weight.name):
+ self.assertAllEqual(
+ qat_model_weights[idx], np.ones_like(qat_model_weights[idx]))
+ count += 1
+ self.assertLen(one_model.weights, count)
+ self.assertGreater(len(qat_model.weights), len(one_model.weights))
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/qat/vision/quantization/layer_transforms.py b/official/projects/qat/vision/quantization/layer_transforms.py
index 8adecea6755..cb1c0f2ad3f 100644
--- a/official/projects/qat/vision/quantization/layer_transforms.py
+++ b/official/projects/qat/vision/quantization/layer_transforms.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,34 +13,28 @@
# limitations under the License.
"""Contains custom quantization layer transforms."""
-from typing import Type, Mapping
-
-import tensorflow as tf
+from typing import Any, Type, Mapping, List, Union, Tuple
+import tensorflow as tf, tf_keras
import tensorflow_model_optimization as tfmot
+from official.modeling import tf_utils
from official.projects.qat.vision.modeling.layers import nn_blocks as quantized_nn_blocks
from official.projects.qat.vision.modeling.layers import nn_layers as quantized_nn_layers
from official.projects.qat.vision.quantization import configs
+from official.projects.qat.vision.quantization import helper
keras = tf.keras
LayerNode = tfmot.quantization.keras.graph_transformations.transforms.LayerNode
LayerPattern = tfmot.quantization.keras.graph_transformations.transforms.LayerPattern
_LAYER_NAMES = [
- 'Vision>Conv2DBNBlock', 'Vision>InvertedBottleneckBlock',
- 'Vision>SegmentationHead', 'Vision>SpatialPyramidPooling', 'Vision>ASPP'
-]
-
-_QUANTIZATION_WEIGHT_NAMES = [
- 'output_max', 'output_min', 'optimizer_step', 'kernel_min', 'kernel_max',
- 'add_three_min', 'add_three_max', 'divide_six_min', 'divide_six_max',
- 'depthwise_kernel_min', 'depthwise_kernel_max',
- 'reduce_mean_quantizer_vars_min', 'reduce_mean_quantizer_vars_max'
-]
-
-_ORIGINAL_WEIGHT_NAME = [
- 'kernel', 'depthwise_kernel', 'gamma', 'beta', 'moving_mean',
- 'moving_variance', 'bias'
+ 'Vision>Conv2DBNBlock',
+ 'Vision>InvertedBottleneckBlock',
+ 'Vision>MaybeDwInvertedBottleneckBlock',
+ 'Vision>UniversalInvertedBottleneckBlock',
+ 'Vision>SegmentationHead',
+ 'Vision>SpatialPyramidPooling',
+ 'Vision>ASPP',
]
@@ -58,16 +52,6 @@ def pattern(self) -> LayerPattern:
"""See base class."""
return LayerPattern(self._original_layer_pattern)
- def _is_quantization_weight_name(self, name):
- simple_name = name.split('/')[-1].split(':')[0]
- if simple_name in _QUANTIZATION_WEIGHT_NAMES:
- return True
- if simple_name in _ORIGINAL_WEIGHT_NAME:
- return False
- raise ValueError('Variable name {} is not supported on '
- 'CustomLayerQuantize({}) transform.'.format(
- simple_name, self._original_layer_pattern))
-
def _create_layer_metadata(
self, layer_class_name: str
) -> Mapping[str, tfmot.quantization.keras.QuantizeConfig]:
@@ -79,17 +63,23 @@ def _create_layer_metadata(
}
return layer_metadata
+ def _create_dummy_input_shape(
+ self, quantized_layer: tf_keras.layers.Layer
+ ) -> Union[List[int], Tuple[Any, Any]]:
+ dummy_input_shape = [1, 128, 128, 1]
+ # SegmentationHead layer requires a tuple of 2 tensors.
+ if isinstance(quantized_layer,
+ quantized_nn_layers.SegmentationHeadQuantized):
+ dummy_input_shape = ([1, 1, 1, 1], [1, 1, 1, 1])
+ return dummy_input_shape
+
def replacement(self, match_layer: LayerNode) -> LayerNode:
"""See base class."""
bottleneck_layer = match_layer.layer
bottleneck_config = bottleneck_layer['config']
bottleneck_names_and_weights = list(match_layer.names_and_weights)
quantized_layer = self._quantized_layer_class(**bottleneck_config)
- dummy_input_shape = [1, 64, 128, 1]
- # SegmentationHead layer requires a tuple of 2 tensors.
- if isinstance(quantized_layer,
- quantized_nn_layers.SegmentationHeadQuantized):
- dummy_input_shape = ([1, 1, 1, 1], [1, 1, 1, 1])
+ dummy_input_shape = self._create_dummy_input_shape(quantized_layer)
quantized_layer.compute_output_shape(dummy_input_shape)
quantized_names_and_weights = zip(
[weight.name for weight in quantized_layer.weights],
@@ -97,7 +87,7 @@ def replacement(self, match_layer: LayerNode) -> LayerNode:
match_idx = 0
names_and_weights = []
for name_and_weight in quantized_names_and_weights:
- if not self._is_quantization_weight_name(name=name_and_weight[0]):
+ if not helper.is_quantization_weight_name(name=name_and_weight[0]):
name_and_weight = bottleneck_names_and_weights[match_idx]
match_idx = match_idx + 1
names_and_weights.append(name_and_weight)
@@ -105,7 +95,9 @@ def replacement(self, match_layer: LayerNode) -> LayerNode:
if match_idx != len(bottleneck_names_and_weights):
raise ValueError('{}/{} of Bottleneck weights is transformed.'.format(
match_idx, len(bottleneck_names_and_weights)))
- quantized_layer_config = keras.layers.serialize(quantized_layer)
+ quantized_layer_config = tf_utils.serialize_layer(
+ quantized_layer, use_legacy_format=True
+ )
quantized_layer_config['name'] = quantized_layer_config['config']['name']
layer_metadata = self._create_layer_metadata(bottleneck_layer['class_name'])
@@ -117,15 +109,30 @@ def replacement(self, match_layer: LayerNode) -> LayerNode:
CUSTOM_TRANSFORMS = [
- CustomLayerQuantize('Vision>BottleneckBlock',
- quantized_nn_blocks.BottleneckBlockQuantized),
- CustomLayerQuantize('Vision>InvertedBottleneckBlock',
- quantized_nn_blocks.InvertedBottleneckBlockQuantized),
- CustomLayerQuantize('Vision>Conv2DBNBlock',
- quantized_nn_blocks.Conv2DBNBlockQuantized),
- CustomLayerQuantize('Vision>SegmentationHead',
- quantized_nn_layers.SegmentationHeadQuantized),
- CustomLayerQuantize('Vision>SpatialPyramidPooling',
- quantized_nn_layers.SpatialPyramidPoolingQuantized),
- CustomLayerQuantize('Vision>ASPP', quantized_nn_layers.ASPPQuantized)
+ CustomLayerQuantize(
+ 'Vision>BottleneckBlock', quantized_nn_blocks.BottleneckBlockQuantized
+ ),
+ CustomLayerQuantize(
+ 'Vision>InvertedBottleneckBlock',
+ quantized_nn_blocks.InvertedBottleneckBlockQuantized,
+ ),
+ CustomLayerQuantize(
+ 'Vision>MaybeDwInvertedBottleneckBlock',
+ quantized_nn_blocks.MaybeDwInvertedBottleneckBlockQuantized,
+ ),
+ CustomLayerQuantize(
+ 'Vision>UniversalInvertedBottleneckBlock',
+ quantized_nn_blocks.UniversalInvertedBottleneckBlockQuantized,
+ ),
+ CustomLayerQuantize(
+ 'Vision>Conv2DBNBlock', quantized_nn_blocks.Conv2DBNBlockQuantized
+ ),
+ CustomLayerQuantize(
+ 'Vision>SegmentationHead', quantized_nn_layers.SegmentationHeadQuantized
+ ),
+ CustomLayerQuantize(
+ 'Vision>SpatialPyramidPooling',
+ quantized_nn_layers.SpatialPyramidPoolingQuantized,
+ ),
+ CustomLayerQuantize('Vision>ASPP', quantized_nn_layers.ASPPQuantized),
]
diff --git a/official/projects/qat/vision/quantization/schemes.py b/official/projects/qat/vision/quantization/schemes.py
index fca03e4cbf9..03e9763e016 100644
--- a/official/projects/qat/vision/quantization/schemes.py
+++ b/official/projects/qat/vision/quantization/schemes.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,8 +13,6 @@
# limitations under the License.
"""Quantization schemes."""
-# Import libraries
-
import tensorflow_model_optimization as tfmot
from official.projects.qat.vision.quantization import layer_transforms
diff --git a/official/projects/qat/vision/registry_imports.py b/official/projects/qat/vision/registry_imports.py
index 2c93ccd9afe..6cf8bccfc03 100644
--- a/official/projects/qat/vision/registry_imports.py
+++ b/official/projects/qat/vision/registry_imports.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/qat/vision/serving/__init__.py b/official/projects/qat/vision/serving/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/qat/vision/serving/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/qat/vision/serving/export_module.py b/official/projects/qat/vision/serving/export_module.py
new file mode 100644
index 00000000000..064473f55ff
--- /dev/null
+++ b/official/projects/qat/vision/serving/export_module.py
@@ -0,0 +1,62 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Export modules for QAT model serving/inference."""
+import tensorflow as tf, tf_keras
+
+from official.projects.qat.vision.modeling import factory as qat_factory
+from official.vision import configs
+from official.vision.serving import detection
+from official.vision.serving import image_classification
+from official.vision.serving import semantic_segmentation
+
+
+class ClassificationModule(image_classification.ClassificationModule):
+ """Classification Module."""
+
+ def _build_model(self):
+ model = super()._build_model()
+ input_specs = tf_keras.layers.InputSpec(shape=[self._batch_size] +
+ self._input_image_size + [3])
+ return qat_factory.build_qat_classification_model(
+ model, self.params.task.quantization, input_specs,
+ self.params.task.model)
+
+
+class SegmentationModule(semantic_segmentation.SegmentationModule):
+ """Segmentation Module."""
+
+ def _build_model(self):
+ model = super()._build_model()
+ input_specs = tf_keras.layers.InputSpec(shape=[self._batch_size] +
+ self._input_image_size + [3])
+ return qat_factory.build_qat_segmentation_model(
+ model, self.params.task.quantization, input_specs)
+
+
+class DetectionModule(detection.DetectionModule):
+ """Detection Module."""
+
+ def _build_model(self):
+ model = super()._build_model()
+
+ if isinstance(self.params.task.model, configs.retinanet.RetinaNet):
+ model = qat_factory.build_qat_retinanet(model,
+ self.params.task.quantization,
+ self.params.task.model)
+ else:
+ raise ValueError('Detection module not implemented for {} model.'.format(
+ type(self.params.task.model)))
+
+ return model
diff --git a/official/projects/qat/vision/serving/export_saved_model.py b/official/projects/qat/vision/serving/export_saved_model.py
new file mode 100644
index 00000000000..e73e86576d6
--- /dev/null
+++ b/official/projects/qat/vision/serving/export_saved_model.py
@@ -0,0 +1,138 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+r"""Vision models export binary for serving/inference.
+
+To export a trained checkpoint in saved_model format (shell script):
+
+EXPERIMENT_TYPE = XX
+CHECKPOINT_PATH = XX
+EXPORT_DIR_PATH = XX
+export_saved_model --experiment=${EXPERIMENT_TYPE} \
+ --export_dir=${EXPORT_DIR_PATH}/ \
+ --checkpoint_path=${CHECKPOINT_PATH} \
+ --batch_size=2 \
+ --input_image_size=224,224
+
+To serve (python):
+
+export_dir_path = XX
+input_type = XX
+input_images = XX
+imported = tf.saved_model.load(export_dir_path)
+model_fn = imported.signatures['serving_default']
+output = model_fn(input_images)
+"""
+from absl import app
+from absl import flags
+
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.qat.vision import registry_imports # pylint: disable=unused-import
+from official.projects.qat.vision.serving import export_module
+from official.vision import configs
+from official.vision.serving import export_saved_model_lib
+
+FLAGS = flags.FLAGS
+
+_EXPERIMENT = flags.DEFINE_string(
+ 'experiment', None, 'experiment type, e.g. retinanet_resnetfpn_coco')
+_EXPORT_DIR = flags.DEFINE_string('export_dir', None, 'The export directory.')
+_CHECKPOINT_PATH = flags.DEFINE_string('checkpoint_path', None,
+ 'Checkpoint path.')
+_CONFIG_FILE = flags.DEFINE_multi_string(
+ 'config_file',
+ default=None,
+ help='YAML/JSON files which specifies overrides. The override order '
+ 'follows the order of args. Note that each file '
+ 'can be used as an override template to override the default parameters '
+ 'specified in Python. If the same parameter is specified in both '
+ '`--config_file` and `--params_override`, `config_file` will be used '
+ 'first, followed by params_override.')
+_PARAMS_OVERRIDE = flags.DEFINE_string(
+ 'params_override', '',
+ 'The JSON/YAML file or string which specifies the parameter to be overriden'
+ ' on top of `config_file` template.')
+_BATCH_SIZE = flags.DEFINE_integer('batch_size', None, 'The batch size.')
+_IMAGE_TYPE = flags.DEFINE_string(
+ 'input_type', 'image_tensor',
+ 'One of `image_tensor`, `image_bytes`, `tf_example` and `tflite`.')
+_INPUT_IMAGE_SIZE = flags.DEFINE_string(
+ 'input_image_size', '224,224',
+ 'The comma-separated string of two integers representing the height,width '
+ 'of the input to the model.')
+_EXPORT_CHECKPOINT_SUBDIR = flags.DEFINE_string(
+ 'export_checkpoint_subdir', 'checkpoint',
+ 'The subdirectory for checkpoints.')
+_EXPORT_SAVED_MODEL_SUBDIR = flags.DEFINE_string(
+ 'export_saved_model_subdir', 'saved_model',
+ 'The subdirectory for saved model.')
+_LOG_MODEL_FLOPS_AND_PARAMS = flags.DEFINE_bool(
+ 'log_model_flops_and_params', False,
+ 'If true, logs model flops and parameters.')
+_INPUT_NAME = flags.DEFINE_string(
+ 'input_name', None,
+ 'Input tensor name in signature def. Default at None which'
+ 'produces input tensor name `inputs`.')
+
+
+def main(_):
+
+ params = exp_factory.get_exp_config(_EXPERIMENT.value)
+ for config_file in _CONFIG_FILE.value or []:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=True)
+ if _PARAMS_OVERRIDE.value:
+ params = hyperparams.override_params_dict(
+ params, _PARAMS_OVERRIDE.value, is_strict=True)
+
+ params.validate()
+ params.lock()
+
+ input_image_size = [int(x) for x in _INPUT_IMAGE_SIZE.value.split(',')]
+
+ if isinstance(params.task,
+ configs.image_classification.ImageClassificationTask):
+ export_module_cls = export_module.ClassificationModule
+ elif isinstance(params.task, configs.retinanet.RetinaNetTask):
+ export_module_cls = export_module.DetectionModule
+ elif isinstance(params.task,
+ configs.semantic_segmentation.SemanticSegmentationTask):
+ export_module_cls = export_module.SegmentationModule
+ else:
+ raise TypeError(f'Export module for {type(params.task)} is not supported.')
+
+ module = export_module_cls(
+ params=params,
+ batch_size=_BATCH_SIZE.value,
+ input_image_size=input_image_size,
+ input_type=_IMAGE_TYPE.value,
+ num_channels=3)
+
+ export_saved_model_lib.export_inference_graph(
+ input_type=_IMAGE_TYPE.value,
+ batch_size=_BATCH_SIZE.value,
+ input_image_size=input_image_size,
+ params=params,
+ checkpoint_path=_CHECKPOINT_PATH.value,
+ export_dir=_EXPORT_DIR.value,
+ export_checkpoint_subdir=_EXPORT_CHECKPOINT_SUBDIR.value,
+ export_saved_model_subdir=_EXPORT_SAVED_MODEL_SUBDIR.value,
+ export_module=module,
+ log_model_flops_and_params=_LOG_MODEL_FLOPS_AND_PARAMS.value,
+ input_name=_INPUT_NAME.value)
+
+
+if __name__ == '__main__':
+ app.run(main)
diff --git a/official/projects/qat/vision/serving/export_tflite.py b/official/projects/qat/vision/serving/export_tflite.py
new file mode 100644
index 00000000000..7ddb01a6b51
--- /dev/null
+++ b/official/projects/qat/vision/serving/export_tflite.py
@@ -0,0 +1,23 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Binary to convert a saved model to TFLite model for the QAT model."""
+
+from absl import app
+
+from official.projects.qat.vision import registry_imports # pylint: disable=unused-import
+from official.vision.serving import export_tflite
+
+if __name__ == '__main__':
+ app.run(export_tflite.main)
diff --git a/official/projects/qat/vision/tasks/__init__.py b/official/projects/qat/vision/tasks/__init__.py
index 42350ee0790..1c0fbccf92c 100644
--- a/official/projects/qat/vision/tasks/__init__.py
+++ b/official/projects/qat/vision/tasks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,5 +13,6 @@
# limitations under the License.
"""Tasks package definition."""
-
from official.projects.qat.vision.tasks import image_classification
+from official.projects.qat.vision.tasks import retinanet
+from official.projects.qat.vision.tasks import semantic_segmentation
diff --git a/official/projects/qat/vision/tasks/image_classification.py b/official/projects/qat/vision/tasks/image_classification.py
index 20824dcef55..058d70c1052 100644
--- a/official/projects/qat/vision/tasks/image_classification.py
+++ b/official/projects/qat/vision/tasks/image_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Image classification task definition."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import task_factory
from official.projects.qat.vision.configs import image_classification as exp_cfg
@@ -25,25 +25,37 @@
class ImageClassificationTask(image_classification.ImageClassificationTask):
"""A task for image classification with QAT."""
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
"""Builds classification model with QAT."""
- input_specs = tf.keras.layers.InputSpec(
- shape=[None] + self.task_config.model.input_size)
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + self.task_config.model.input_size
+ )
l2_weight_decay = self.task_config.losses.l2_weight_decay
# Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
- l2_regularizer = (tf.keras.regularizers.l2(
- l2_weight_decay / 2.0) if l2_weight_decay else None)
-
- model = super(ImageClassificationTask, self).build_model()
- if self.task_config.quantization:
+ l2_regularizer = (
+ tf_keras.regularizers.l2(l2_weight_decay / 2.0)
+ if l2_weight_decay
+ else None
+ )
+
+ model = super().build_model()
+
+ # Only build a QAT model when quantization version is v2; otherwise leave it
+ # for outer quantization scope.
+ if (
+ self.task_config.quantization
+ and hasattr(self.task_config.quantization, 'version')
+ and self.task_config.quantization.version == 'v2'
+ ):
model = factory.build_qat_classification_model(
model,
self.task_config.quantization,
input_specs=input_specs,
model_config=self.task_config.model,
- l2_regularizer=l2_regularizer)
+ l2_regularizer=l2_regularizer,
+ )
return model
diff --git a/official/projects/qat/vision/tasks/image_classification_test.py b/official/projects/qat/vision/tasks/image_classification_test.py
index eac971da323..bee74a45adc 100644
--- a/official/projects/qat/vision/tasks/image_classification_test.py
+++ b/official/projects/qat/vision/tasks/image_classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl.testing import parameterized
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import exp_factory
diff --git a/official/projects/qat/vision/tasks/retinanet.py b/official/projects/qat/vision/tasks/retinanet.py
index 4adc82025fc..5c34a2d4633 100644
--- a/official/projects/qat/vision/tasks/retinanet.py
+++ b/official/projects/qat/vision/tasks/retinanet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""RetinaNet task definition."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import task_factory
from official.projects.qat.vision.configs import retinanet as exp_cfg
@@ -25,12 +25,22 @@
class RetinaNetTask(retinanet.RetinaNetTask):
"""A task for RetinaNet object detection with QAT."""
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
"""Builds RetinaNet model with QAT."""
model = super(RetinaNetTask, self).build_model()
- if self.task_config.quantization:
+ # Call the model with dummy input to build the head part.
+ dummpy_input = tf.zeros([1] + self.task_config.model.input_size)
+ model(dummpy_input, training=True)
+
+ # Only build a QAT model when quantization version is v2; otherwise leave it
+ # for outer quantization scope.
+ if (
+ self.task_config.quantization
+ and self.task_config.quantization.version == 'v2'
+ ):
model = factory.build_qat_retinanet(
model,
self.task_config.quantization,
- model_config=self.task_config.model)
+ model_config=self.task_config.model,
+ )
return model
diff --git a/official/projects/qat/vision/tasks/retinanet_test.py b/official/projects/qat/vision/tasks/retinanet_test.py
index fd459fa09b3..e2701982dce 100644
--- a/official/projects/qat/vision/tasks/retinanet_test.py
+++ b/official/projects/qat/vision/tasks/retinanet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl.testing import parameterized
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import exp_factory
@@ -36,8 +36,8 @@ def _create_test_tfrecord(self, tfrecord_file, example, num_samples):
record_file=tfrecord_file, tf_examples=examples)
@parameterized.parameters(
- ('retinanet_spinenet_mobile_coco_qat', True),
- ('retinanet_spinenet_mobile_coco_qat', False),
+ ('retinanet_mobile_coco_qat', True),
+ ('retinanet_mobile_coco_qat', False),
)
def test_retinanet_task(self, test_config, is_training):
"""RetinaNet task test for training and val using toy configs."""
@@ -65,6 +65,7 @@ def test_retinanet_task(self, test_config, is_training):
task = retinanet.RetinaNetTask(config.task)
model = task.build_model()
+ self.assertLen(model.weights, 2393)
metrics = task.build_metrics(training=is_training)
strategy = tf.distribute.get_strategy()
diff --git a/official/projects/qat/vision/tasks/semantic_segmentation.py b/official/projects/qat/vision/tasks/semantic_segmentation.py
index a2c1bd0e972..7baeed81ed7 100644
--- a/official/projects/qat/vision/tasks/semantic_segmentation.py
+++ b/official/projects/qat/vision/tasks/semantic_segmentation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Semantic segmentation task definition."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import task_factory
from official.projects.qat.vision.configs import semantic_segmentation as exp_cfg
@@ -25,12 +25,21 @@
class SemanticSegmentationTask(semantic_segmentation.SemanticSegmentationTask):
"""A task for semantic segmentation with QAT."""
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
"""Builds semantic segmentation model with QAT."""
model = super().build_model()
- input_specs = tf.keras.layers.InputSpec(shape=[None] +
- self.task_config.model.input_size)
- if self.task_config.quantization:
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None] + self.task_config.model.input_size
+ )
+
+ # Only build a QAT model when quantization version is v2; otherwise leave it
+ # for outer quantization scope.
+ if (
+ self.task_config.quantization
+ and hasattr(self.task_config.quantization, 'version')
+ and self.task_config.quantization.version == 'v2'
+ ):
model = factory.build_qat_segmentation_model(
- model, self.task_config.quantization, input_specs)
+ model, self.task_config.quantization, input_specs
+ )
return model
diff --git a/official/projects/qat/vision/train.py b/official/projects/qat/vision/train.py
index 453cb9fe4e0..76a07b32221 100644
--- a/official/projects/qat/vision/train.py
+++ b/official/projects/qat/vision/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/roformer/__init__.py b/official/projects/roformer/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/projects/roformer/__init__.py
+++ b/official/projects/roformer/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/roformer/roformer.py b/official/projects/roformer/roformer.py
index 0474de3dac8..163e8c69aaf 100644
--- a/official/projects/roformer/roformer.py
+++ b/official/projects/roformer/roformer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Roformer model configurations and instantiation methods."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.modeling.hyperparams import base_config
@@ -46,7 +46,7 @@ def get_encoder(encoder_cfg: RoformerEncoderConfig):
attention_dropout_rate=encoder_cfg.attention_dropout_rate,
max_sequence_length=encoder_cfg.max_position_embeddings,
type_vocab_size=encoder_cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
output_range=encoder_cfg.output_range,
embedding_width=encoder_cfg.embedding_size,
diff --git a/official/projects/roformer/roformer_attention.py b/official/projects/roformer/roformer_attention.py
index 14dbb18c450..5a77eba7a89 100644
--- a/official/projects/roformer/roformer_attention.py
+++ b/official/projects/roformer/roformer_attention.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,29 +14,32 @@
"""Roformer attention layer."""
# pylint: disable=g-classes-have-attributes
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-EinsumDense = tf.keras.layers.experimental.EinsumDense
-MultiHeadAttention = tf.keras.layers.MultiHeadAttention
+EinsumDense = tf_keras.layers.EinsumDense
+MultiHeadAttention = tf_keras.layers.MultiHeadAttention
def _build_trig_vector(length, key_dim):
"""Builds the trig vector."""
- tf_dtype = tf.keras.mixed_precision.global_policy().compute_dtype
+ tf_dtype = tf_keras.mixed_precision.global_policy().compute_dtype
position_ids = tf.cast(tf.range(length), dtype=tf_dtype)
position_ids = tf.expand_dims(position_ids, axis=0)
steps = key_dim // 2
- indices = tf.cast(tf.range(steps), dtype=tf_dtype)
- indices = tf.pow(tf.constant(10000.0, dtype=tf_dtype), -2 * indices / steps)
- vec = tf.einsum('bl,d->bld', position_ids, indices)
+ # 2 (i - 1) / key_dim = (i - 1) / steps: (-1 achieved with zero-indexing)
+ wavenumber_exponent = -tf.cast(tf.range(steps), dtype=tf_dtype) / steps
+ wavenumbers = tf.pow(
+ tf.constant(10000.0, dtype=tf_dtype), wavenumber_exponent
+ )
+ vec = tf.einsum('bl,d->bld', position_ids, wavenumbers)
sin_vec = tf.repeat(tf.sin(vec), repeats=2, axis=-1)
cos_vec = tf.repeat(tf.cos(vec), repeats=2, axis=-1)
sin_vec, cos_vec = tf.expand_dims(sin_vec, 2), tf.expand_dims(cos_vec, 2)
return sin_vec, cos_vec
-@tf.keras.utils.register_keras_serializable(package='Text')
-class RoformerAttention(tf.keras.layers.MultiHeadAttention):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class RoformerAttention(tf_keras.layers.MultiHeadAttention):
"""Roformer Attention."""
def __init__(self,
@@ -87,7 +90,7 @@ def roformer_recompute_qkv(self, q, k, v):
...] + k2 * self.k_sin_vec[:, 0:k_len, ...]
return ret_q, ret_w, v
- def call(self,
+ def call(self, # pytype: disable=signature-mismatch # overriding-parameter-count-checks
query,
value,
key=None,
diff --git a/official/projects/roformer/roformer_attention_test.py b/official/projects/roformer/roformer_attention_test.py
index d131e876a7a..11a632e21b9 100644
--- a/official/projects/roformer/roformer_attention_test.py
+++ b/official/projects/roformer/roformer_attention_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,11 @@
"""Tests for the attention layer."""
+from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.projects.roformer import roformer_attention
@@ -62,8 +62,7 @@ def _create_mock_attention_data(num_heads,
return data
-@keras_parameterized.run_all_keras_modes
-class RoformerAttentionTest(keras_parameterized.TestCase):
+class RoformerAttentionTest(tf.test.TestCase, parameterized.TestCase):
def setUp(self):
super(RoformerAttentionTest, self).setUp()
@@ -79,7 +78,7 @@ def test_trig_vector(self, length, key_dim):
for m in range(0, length):
half_d = key_dim // 2
std_emb = tf.range(half_d, dtype=tf.float32)
- std_emb = tf.pow(10000.0, -2 * std_emb / float(half_d))
+ std_emb = tf.pow(10000.0, -std_emb / float(half_d))
std_emb = m * std_emb
std_sin_emb = tf.sin(std_emb)
std_cos_emb = tf.cos(std_emb)
diff --git a/official/projects/roformer/roformer_encoder.py b/official/projects/roformer/roformer_encoder.py
index 3a84f3c3d19..3c7b32fdd3f 100644
--- a/official/projects/roformer/roformer_encoder.py
+++ b/official/projects/roformer/roformer_encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,14 +17,15 @@
import collections
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
from official.projects.roformer import roformer_encoder_block
-@tf.keras.utils.register_keras_serializable(package='Text')
-class RoformerEncoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class RoformerEncoder(tf_keras.Model):
"""Bi-directional Transformer-based encoder network with Roformer.
Roformer paper: https://arxiv.org/abs/2104.09864
@@ -76,10 +77,10 @@ def __init__(
max_sequence_length=512,
type_vocab_size=16,
inner_dim=3072,
- inner_activation=lambda x: tf.keras.activations.gelu(x, approximate=True),
+ inner_activation=lambda x: tf_keras.activations.gelu(x, approximate=True),
output_dropout=0.1,
attention_dropout=0.1,
- initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02),
output_range=None,
embedding_width=None,
embedding_layer=None,
@@ -98,14 +99,14 @@ def __init__(
attention_dropout = kwargs['attention_dropout_rate']
del kwargs['attention_dropout_rate']
- activation = tf.keras.activations.get(inner_activation)
- initializer = tf.keras.initializers.get(initializer)
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
- word_ids = tf.keras.layers.Input(
+ word_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_word_ids')
- mask = tf.keras.layers.Input(
+ mask = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_mask')
- type_ids = tf.keras.layers.Input(
+ type_ids = tf_keras.layers.Input(
shape=(None,), dtype=tf.int32, name='input_type_ids')
if embedding_width is None:
@@ -115,7 +116,7 @@ def __init__(
embedding_layer_inst = layers.on_device_embedding.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
name='word_embeddings')
else:
embedding_layer_inst = embedding_layer
@@ -125,28 +126,28 @@ def __init__(
type_embedding_layer = layers.on_device_embedding.OnDeviceEmbedding(
vocab_size=type_vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_one_hot=True,
name='type_embeddings')
type_embeddings = type_embedding_layer(type_ids)
# Roformer does not have absolute position embedding
- embeddings = tf.keras.layers.Add()([word_embeddings, type_embeddings])
+ embeddings = tf_keras.layers.Add()([word_embeddings, type_embeddings])
- embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
embeddings = embedding_norm_layer(embeddings)
- embeddings = (tf.keras.layers.Dropout(rate=output_dropout)(embeddings))
+ embeddings = (tf_keras.layers.Dropout(rate=output_dropout)(embeddings))
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
if embedding_width != hidden_size:
- embedding_projection = tf.keras.layers.experimental.EinsumDense(
+ embedding_projection = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='embedding_projection')
embeddings = embedding_projection(embeddings)
else:
@@ -171,7 +172,7 @@ def __init__(
attention_dropout=attention_dropout,
norm_first=norm_first,
output_range=transformer_output_range,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='roformer/layer_%d' % i)
transformer_layers.append(layer)
data = layer([data, attention_mask])
@@ -182,10 +183,10 @@ def __init__(
# like this will create a SliceOpLambda layer. This is better than a Lambda
# layer with Python code, because that is fundamentally less portable.
first_token_tensor = last_encoder_output[:, 0, :]
- pooler_layer = tf.keras.layers.Dense(
+ pooler_layer = tf_keras.layers.Dense(
units=hidden_size,
activation='tanh',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='pooler_transform')
cls_output = pooler_layer(first_token_tensor)
@@ -212,10 +213,10 @@ def __init__(
'max_sequence_length': max_sequence_length,
'type_vocab_size': type_vocab_size,
'inner_dim': inner_dim,
- 'inner_activation': tf.keras.activations.serialize(activation),
+ 'inner_activation': tf_keras.activations.serialize(activation),
'output_dropout': output_dropout,
'attention_dropout': attention_dropout,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_keras.initializers.serialize(initializer),
'output_range': output_range,
'embedding_width': embedding_width,
'embedding_layer': embedding_layer,
@@ -227,7 +228,7 @@ def __init__(
# the config dict attribute. TF does not track immutable attrs which
# do not contain Trackables, so by creating a config namedtuple instead of
# a dict we avoid tracking it.
- config_cls = collections.namedtuple('Config', config_dict.keys())
+ config_cls = collections.namedtuple('Config', config_dict.keys()) # pyrefly: ignore[bad-class-definition]
self._config = config_cls(**config_dict)
self._pooler_layer = pooler_layer
self._transformer_layers = transformer_layers
diff --git a/official/projects/roformer/roformer_encoder_block.py b/official/projects/roformer/roformer_encoder_block.py
index e7d894ee26c..7375991caba 100644
--- a/official/projects/roformer/roformer_encoder_block.py
+++ b/official/projects/roformer/roformer_encoder_block.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,12 +14,13 @@
"""Roformer TransformerEncoder block layer."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.projects.roformer import roformer_attention
-@tf.keras.utils.register_keras_serializable(package="Text")
-class RoformerEncoderBlock(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package="Text")
+class RoformerEncoderBlock(tf_keras.layers.Layer):
"""RoformerEncoderBlock layer."""
def __init__(self,
@@ -94,13 +95,13 @@ def __init__(self,
self._output_dropout = output_dropout
self._output_dropout_rate = output_dropout
self._output_range = output_range
- self._kernel_initializer = tf.keras.initializers.get(kernel_initializer)
- self._bias_initializer = tf.keras.initializers.get(bias_initializer)
- self._kernel_regularizer = tf.keras.regularizers.get(kernel_regularizer)
- self._bias_regularizer = tf.keras.regularizers.get(bias_regularizer)
- self._activity_regularizer = tf.keras.regularizers.get(activity_regularizer)
- self._kernel_constraint = tf.keras.constraints.get(kernel_constraint)
- self._bias_constraint = tf.keras.constraints.get(bias_constraint)
+ self._kernel_initializer = tf_keras.initializers.get(kernel_initializer)
+ self._bias_initializer = tf_keras.initializers.get(bias_initializer)
+ self._kernel_regularizer = tf_keras.regularizers.get(kernel_regularizer)
+ self._bias_regularizer = tf_keras.regularizers.get(bias_regularizer)
+ self._activity_regularizer = tf_keras.regularizers.get(activity_regularizer)
+ self._kernel_constraint = tf_keras.constraints.get(kernel_constraint)
+ self._bias_constraint = tf_keras.constraints.get(bias_constraint)
self._use_bias = use_bias
self._norm_first = norm_first
self._norm_epsilon = norm_epsilon
@@ -108,10 +109,11 @@ def __init__(self,
self._q_max_sequence_length = q_max_sequence_length
self._kv_max_sequence_length = kv_max_sequence_length
if attention_initializer:
- self._attention_initializer = tf.keras.initializers.get(
+ self._attention_initializer = tf_keras.initializers.get(
attention_initializer)
else:
- self._attention_initializer = self._kernel_initializer
+ self._attention_initializer = tf_utils.clone_initializer(
+ self._kernel_initializer)
self._attention_axes = attention_axes
def build(self, input_shape):
@@ -151,42 +153,42 @@ def build(self, input_shape):
attention_axes=self._attention_axes,
name="self_attention",
**common_kwargs)
- self._attention_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ self._attention_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
# It is probably safe in mixed_float16, but we haven't validated this yet.
self._attention_layer_norm = (
- tf.keras.layers.LayerNormalization(
+ tf_keras.layers.LayerNormalization(
name="self_attention_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
dtype=tf.float32))
- self._intermediate_dense = tf.keras.layers.experimental.EinsumDense(
+ self._intermediate_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, self._inner_dim),
bias_axes="d",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
name="intermediate",
**common_kwargs)
- policy = tf.keras.mixed_precision.global_policy()
+ policy = tf_keras.mixed_precision.global_policy()
if policy.name == "mixed_bfloat16":
# bfloat16 causes BERT with the LAMB optimizer to not converge
# as well, so we use float32.
# TODO(b/154538392): Investigate this.
policy = tf.float32
- self._intermediate_activation_layer = tf.keras.layers.Activation(
+ self._intermediate_activation_layer = tf_keras.layers.Activation(
self._inner_activation, dtype=policy)
- self._inner_dropout_layer = tf.keras.layers.Dropout(
+ self._inner_dropout_layer = tf_keras.layers.Dropout(
rate=self._inner_dropout)
- self._output_dense = tf.keras.layers.experimental.EinsumDense(
+ self._output_dense = tf_keras.layers.EinsumDense(
einsum_equation,
output_shape=(None, hidden_size),
bias_axes="d",
name="output",
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(self._kernel_initializer),
**common_kwargs)
- self._output_dropout = tf.keras.layers.Dropout(rate=self._output_dropout)
+ self._output_dropout = tf_keras.layers.Dropout(rate=self._output_dropout)
# Use float32 in layernorm for numeric stability.
- self._output_layer_norm = tf.keras.layers.LayerNormalization(
+ self._output_layer_norm = tf_keras.layers.LayerNormalization(
name="output_layer_norm",
axis=-1,
epsilon=self._norm_epsilon,
@@ -209,19 +211,19 @@ def get_config(self):
"output_range":
self._output_range,
"kernel_initializer":
- tf.keras.initializers.serialize(self._kernel_initializer),
+ tf_keras.initializers.serialize(self._kernel_initializer),
"bias_initializer":
- tf.keras.initializers.serialize(self._bias_initializer),
+ tf_keras.initializers.serialize(self._bias_initializer),
"kernel_regularizer":
- tf.keras.regularizers.serialize(self._kernel_regularizer),
+ tf_keras.regularizers.serialize(self._kernel_regularizer),
"bias_regularizer":
- tf.keras.regularizers.serialize(self._bias_regularizer),
+ tf_keras.regularizers.serialize(self._bias_regularizer),
"activity_regularizer":
- tf.keras.regularizers.serialize(self._activity_regularizer),
+ tf_keras.regularizers.serialize(self._activity_regularizer),
"kernel_constraint":
- tf.keras.constraints.serialize(self._kernel_constraint),
+ tf_keras.constraints.serialize(self._kernel_constraint),
"bias_constraint":
- tf.keras.constraints.serialize(self._bias_constraint),
+ tf_keras.constraints.serialize(self._bias_constraint),
"use_bias":
self._use_bias,
"norm_first":
@@ -231,7 +233,7 @@ def get_config(self):
"inner_dropout":
self._inner_dropout,
"attention_initializer":
- tf.keras.initializers.serialize(self._attention_initializer),
+ tf_keras.initializers.serialize(self._attention_initializer),
"attention_axes":
self._attention_axes,
}
@@ -284,9 +286,9 @@ def call(self, inputs):
key_value = input_tensor
attention_output = self._attention_layer(
query=target_tensor, value=key_value, attention_mask=attention_mask)
- attention_output = self._attention_dropout(attention_output)
+ attention_output = self._attention_dropout(attention_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- attention_output = source_tensor + attention_output
+ attention_output = source_tensor + attention_output # pyrefly: ignore[unbound-name]
else:
attention_output = self._attention_layer_norm(target_tensor +
attention_output)
@@ -297,10 +299,10 @@ def call(self, inputs):
inner_output = self._intermediate_activation_layer(inner_output)
inner_output = self._inner_dropout_layer(inner_output)
layer_output = self._output_dense(inner_output)
- layer_output = self._output_dropout(layer_output)
+ layer_output = self._output_dropout(layer_output) # pyrefly: ignore[not-callable]
if self._norm_first:
- return source_attention_output + layer_output
+ return source_attention_output + layer_output # pyrefly: ignore[unbound-name]
# During mixed precision training, layer norm output is always fp32 for now.
# Casts fp32 for the subsequent add.
diff --git a/official/projects/roformer/roformer_encoder_block_test.py b/official/projects/roformer/roformer_encoder_block_test.py
index 99dd2b00c6c..b5ec9fefd12 100644
--- a/official/projects/roformer/roformer_encoder_block_test.py
+++ b/official/projects/roformer/roformer_encoder_block_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,20 +16,18 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.projects.roformer import roformer_encoder_block
-@keras_parameterized.run_all_keras_modes
@parameterized.named_parameters(
('base', roformer_encoder_block.RoformerEncoderBlock))
-class RoformerEncoderBlockTest(keras_parameterized.TestCase):
+class RoformerEncoderBlockTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(RoformerEncoderBlockTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy('float32')
+ tf_keras.mixed_precision.set_global_policy('float32')
def test_layer_creation(self, transformer_cls):
test_layer = transformer_cls(
@@ -37,7 +35,7 @@ def test_layer_creation(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -48,9 +46,9 @@ def test_layer_creation_with_mask(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
@@ -61,11 +59,11 @@ def test_layer_invocation(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output_tensor = test_layer(data_tensor)
# Create a model from the test layer.
- model = tf.keras.Model(data_tensor, output_tensor)
+ model = tf_keras.Model(data_tensor, output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -80,13 +78,13 @@ def test_layer_invocation_with_mask(self, transformer_cls):
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -184,19 +182,19 @@ def test_layer_output_range_with_pre_norm(self, transformer_cls):
new_output_tensor, output_tensor[:, 0:1, :], atol=5e-5, rtol=0.003)
def test_layer_invocation_with_float16_dtype(self, transformer_cls):
- tf.keras.mixed_precision.set_global_policy('mixed_float16')
+ tf_keras.mixed_precision.set_global_policy('mixed_float16')
test_layer = transformer_cls(
num_attention_heads=10, inner_dim=2048, inner_activation='relu')
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
# Create a 2-dimensional input (the first dimension is implicit).
- mask_tensor = tf.keras.Input(shape=(sequence_length, sequence_length))
+ mask_tensor = tf_keras.Input(shape=(sequence_length, sequence_length))
output_tensor = test_layer([data_tensor, mask_tensor])
# Create a model from the test layer.
- model = tf.keras.Model([data_tensor, mask_tensor], output_tensor)
+ model = tf_keras.Model([data_tensor, mask_tensor], output_tensor)
# Invoke the model on test data. We can't validate the output data itself
# (the NN is too complex) but this will rule out structural runtime errors.
@@ -214,11 +212,11 @@ def test_transform_with_initializer(self, transformer_cls):
num_attention_heads=10,
inner_dim=2048,
inner_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
sequence_length = 21
width = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(sequence_length, width))
+ data_tensor = tf_keras.Input(shape=(sequence_length, width))
output = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output.shape.as_list())
@@ -228,7 +226,7 @@ def test_separate_qkv(self, transformer_cls):
num_attention_heads=2,
inner_dim=128,
inner_activation='relu',
- kernel_initializer=tf.keras.initializers.TruncatedNormal(stddev=0.02))
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(stddev=0.02))
# Forward path.
q_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
kv_tensor = tf.zeros([2, 8, 16], dtype=tf.float32)
@@ -238,8 +236,7 @@ def test_separate_qkv(self, transformer_cls):
self.assertEqual(output.shape, q_tensor.shape)
-@keras_parameterized.run_all_keras_modes
-class RoformerArgumentTest(keras_parameterized.TestCase):
+class RoformerArgumentTest(tf.test.TestCase, parameterized.TestCase):
def test_raises(self):
num_attention_heads = 2
@@ -254,7 +251,7 @@ def test_raises(self):
norm_first=True,
norm_epsilon=1e-6,
inner_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
def test_use_bias_norm_first(self):
@@ -270,7 +267,7 @@ def test_use_bias_norm_first(self):
norm_first=True,
norm_epsilon=1e-6,
inner_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
# Forward path.
dummy_tensor = tf.zeros([2, 4, 16], dtype=tf.float32)
@@ -291,7 +288,7 @@ def test_get_config(self):
norm_first=True,
norm_epsilon=1e-6,
inner_dropout=0.1,
- attention_initializer=tf.keras.initializers.RandomUniform(
+ attention_initializer=tf_keras.initializers.RandomUniform(
minval=0., maxval=1.))
encoder_block_config = encoder_block.get_config()
new_encoder_block = roformer_encoder_block.RoformerEncoderBlock.from_config(
@@ -315,7 +312,7 @@ def test_several_attention_axes(self, attention_axes):
seq_len = 21
dimensions = 80
# Create a 3-dimensional input (the first dimension is implicit).
- data_tensor = tf.keras.Input(shape=(seq_len, dimensions))
+ data_tensor = tf_keras.Input(shape=(seq_len, dimensions))
output_tensor = test_layer(data_tensor)
# The default output of a transformer layer should be the same as the input.
self.assertEqual(data_tensor.shape.as_list(), output_tensor.shape.as_list())
diff --git a/official/projects/roformer/roformer_encoder_test.py b/official/projects/roformer/roformer_encoder_test.py
index 7fc77f3cf4e..9db27e26d9e 100644
--- a/official/projects/roformer/roformer_encoder_test.py
+++ b/official/projects/roformer/roformer_encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,20 +16,16 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.projects.roformer import roformer_encoder
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class RoformerEncoderTest(keras_parameterized.TestCase):
+class RoformerEncoderTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(RoformerEncoderTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
def test_network_creation(self):
hidden_size = 32
@@ -41,16 +37,16 @@ def test_network_creation(self):
num_attention_heads=2,
num_layers=3)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, 3)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
expected_data_shape = [None, sequence_length, hidden_size]
expected_pooled_shape = [None, hidden_size]
@@ -71,9 +67,9 @@ def test_all_encoder_outputs_network_creation(self):
num_attention_heads=2,
num_layers=3)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network([word_ids, mask, type_ids])
all_encoder_outputs = dict_outputs["encoder_outputs"]
pooled = dict_outputs["pooled_output"]
@@ -92,7 +88,7 @@ def test_all_encoder_outputs_network_creation(self):
def test_network_creation_with_float16_dtype(self):
hidden_size = 32
sequence_length = 21
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a small BertEncoder for testing.
test_network = roformer_encoder.RoformerEncoder(
vocab_size=100,
@@ -100,9 +96,9 @@ def test_network_creation_with_float16_dtype(self):
num_attention_heads=2,
num_layers=3)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
@@ -135,15 +131,15 @@ def test_network_invocation(self, output_range, out_seq_len):
type_vocab_size=num_types,
output_range=output_range)
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -168,7 +164,7 @@ def test_network_invocation(self, output_range, out_seq_len):
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[1], sequence_length)
@@ -184,7 +180,7 @@ def test_network_invocation(self, output_range, out_seq_len):
dict_outputs = test_network([word_ids, mask, type_ids])
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[-1], hidden_size)
self.assertTrue(hasattr(test_network, "_embedding_projection"))
@@ -209,10 +205,10 @@ def test_serialize_deserialize(self):
norm_first=False)
network = roformer_encoder.RoformerEncoder(**kwargs)
expected_config = dict(kwargs)
- expected_config["inner_activation"] = tf.keras.activations.serialize(
- tf.keras.activations.get(expected_config["inner_activation"]))
- expected_config["initializer"] = tf.keras.initializers.serialize(
- tf.keras.initializers.get(expected_config["initializer"]))
+ expected_config["inner_activation"] = tf_keras.activations.serialize(
+ tf_keras.activations.get(expected_config["inner_activation"]))
+ expected_config["initializer"] = tf_keras.initializers.serialize(
+ tf_keras.initializers.get(expected_config["initializer"]))
self.assertEqual(network.get_config(), expected_config)
# Create another network object from the first object's config.
new_network = roformer_encoder.RoformerEncoder.from_config(
@@ -227,7 +223,7 @@ def test_serialize_deserialize(self):
# Tests model saving/loading.
model_path = self.get_temp_dir() + "/model"
network.save(model_path)
- _ = tf.keras.models.load_model(model_path)
+ _ = tf_keras.models.load_model(model_path)
if __name__ == "__main__":
diff --git a/official/projects/roformer/roformer_experiments.py b/official/projects/roformer/roformer_experiments.py
index cb095847d3d..19a2b1098e6 100644
--- a/official/projects/roformer/roformer_experiments.py
+++ b/official/projects/roformer/roformer_experiments.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/roformer/train.py b/official/projects/roformer/train.py
index 6ea0aec4b35..eace52f9564 100644
--- a/official/projects/roformer/train.py
+++ b/official/projects/roformer/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/s3d/configs/s3d.py b/official/projects/s3d/configs/s3d.py
index 1dcd1424c2c..893f32f3955 100644
--- a/official/projects/s3d/configs/s3d.py
+++ b/official/projects/s3d/configs/s3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -83,7 +83,7 @@ class Backbone3D(backbones_3d.Backbone3D):
s3d: s3d backbone config.
"""
type: str = 's3d'
- s3d: S3D = S3D()
+ s3d: S3D = dataclasses.field(default_factory=S3D)
@dataclasses.dataclass
@@ -95,4 +95,4 @@ class S3DModel(video_classification.VideoClassificationModel):
backbone: backbone config.
"""
model_type: str = 's3d'
- backbone: Backbone3D = Backbone3D()
+ backbone: Backbone3D = dataclasses.field(default_factory=Backbone3D)
diff --git a/official/projects/s3d/modeling/inception_utils.py b/official/projects/s3d/modeling/inception_utils.py
index ed0e06d3221..5618d6ce72f 100644
--- a/official/projects/s3d/modeling/inception_utils.py
+++ b/official/projects/s3d/modeling/inception_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,8 +15,9 @@
"""Contains modules related to Inception networks."""
from typing import Callable, Dict, Optional, Sequence, Set, Text, Tuple, Type, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.projects.s3d.modeling import net_utils
from official.vision.modeling.layers import nn_blocks_3d
@@ -47,8 +48,8 @@
('Mixed_5c_local', [[384], [192, 384], [48, 128], [128]]), # 8x7x7x1024
]
-initializers = tf.keras.initializers
-regularizers = tf.keras.regularizers
+initializers = tf_keras.initializers
+regularizers = tf_keras.regularizers
def inception_v1_stem_cells(
@@ -70,7 +71,7 @@ def inception_v1_stem_cells(
kernel_regularizer: Union[Text, regularizers.Regularizer] = 'l2',
parameterized_conv_layer: Type[
net_utils.ParameterizedConvLayer] = net_utils.ParameterizedConvLayer,
- layer_naming_fn: Callable[[Text], Text] = lambda end_point: None,
+ layer_naming_fn: Callable[[Text], Text] = lambda end_point: None, # pyrefly: ignore[bad-function-definition]
) -> Tuple[tf.Tensor, Dict[Text, tf.Tensor]]:
"""Stem cells used in the original I3D/S3D model.
@@ -109,10 +110,10 @@ def inception_v1_stem_cells(
if self_gating_endpoints is None:
self_gating_endpoints = set()
if use_sync_bn:
- batch_norm = tf.keras.layers.experimental.SyncBatchNormalization
+ batch_norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- batch_norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ batch_norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
bn_axis = -1
else:
bn_axis = 1
@@ -120,13 +121,13 @@ def inception_v1_stem_cells(
end_points = {}
# batch_size x 32 x 112 x 112 x 64
end_point = 'Conv2d_1a_7x7'
- net = tf.keras.layers.Conv3D(
+ net = tf_keras.layers.Conv3D(
filters=net_utils.apply_depth_multiplier(64, depth_multiplier),
kernel_size=[first_temporal_kernel_size, 7, 7],
strides=[2, 2, 2],
padding='same',
use_bias=False,
- kernel_initializer=kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(kernel_initializer),
kernel_regularizer=kernel_regularizer,
name=layer_naming_fn(end_point))(
inputs)
@@ -144,7 +145,7 @@ def inception_v1_stem_cells(
return net, end_points
# batch_size x 32 x 56 x 56 x 64
end_point = 'MaxPool_2a_3x3'
- net = tf.keras.layers.MaxPool3D(
+ net = tf_keras.layers.MaxPool3D(
pool_size=[1, 3, 3],
strides=[1, 2, 2],
padding='same',
@@ -155,13 +156,13 @@ def inception_v1_stem_cells(
return net, end_points
# batch_size x 32 x 56 x 56 x 64
end_point = 'Conv2d_2b_1x1'
- net = tf.keras.layers.Conv3D(
+ net = tf_keras.layers.Conv3D(
filters=net_utils.apply_depth_multiplier(64, depth_multiplier),
strides=[1, 1, 1],
kernel_size=[1, 1, 1],
padding='same',
use_bias=False,
- kernel_initializer=kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(kernel_initializer),
kernel_regularizer=kernel_regularizer,
name=layer_naming_fn(end_point))(
net)
@@ -191,7 +192,7 @@ def inception_v1_stem_cells(
norm_momentum=norm_momentum,
norm_epsilon=norm_epsilon,
temporal_conv_initializer=temporal_conv_initializer,
- kernel_initializer=kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(kernel_initializer),
kernel_regularizer=kernel_regularizer,
name=layer_naming_fn(end_point))(
net)
@@ -205,7 +206,7 @@ def inception_v1_stem_cells(
return net, end_points
# batch_size x 32 x 28 x 28 x 192
end_point = 'MaxPool_3a_3x3'
- net = tf.keras.layers.MaxPool3D(
+ net = tf_keras.layers.MaxPool3D(
pool_size=[1, 3, 3],
strides=[1, 2, 2],
padding='same',
@@ -219,22 +220,22 @@ def _construct_branch_3_layers(
channels: int,
swap_pool_and_1x1x1: bool,
pool_type: Text,
- batch_norm_layer: tf.keras.layers.Layer,
+ batch_norm_layer: tf_keras.layers.Layer,
kernel_initializer: Union[Text, initializers.Initializer],
kernel_regularizer: Union[Text, regularizers.Regularizer],
):
"""Helper function for Branch 3 inside Inception module."""
kernel_size = [1, 3, 3] if pool_type == '2d' else [3] * 3
- conv = tf.keras.layers.Conv3D(
+ conv = tf_keras.layers.Conv3D(
filters=channels,
kernel_size=[1, 1, 1],
padding='same',
use_bias=False,
kernel_initializer=kernel_initializer,
kernel_regularizer=kernel_regularizer)
- activation = tf.keras.layers.Activation('relu')
- pool = tf.keras.layers.MaxPool3D(
+ activation = tf_keras.layers.Activation('relu')
+ pool = tf_keras.layers.MaxPool3D(
pool_size=kernel_size, strides=[1, 1, 1], padding='same')
if swap_pool_and_1x1x1:
branch_3_layers = [conv, batch_norm_layer, activation, pool]
@@ -243,7 +244,7 @@ def _construct_branch_3_layers(
return branch_3_layers
-class InceptionV1CellLayer(tf.keras.layers.Layer):
+class InceptionV1CellLayer(tf_keras.layers.Layer):
"""A single Tensorflow 2 cell used in the original I3D/S3D model."""
def __init__(
@@ -313,11 +314,11 @@ def __init__(
self._kernel_regularizer = kernel_regularizer
self._parameterized_conv_layer = parameterized_conv_layer
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
+ self._norm = tf_keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._channel_axis = -1
else:
self._channel_axis = 1
@@ -330,7 +331,8 @@ def _build_branch_params(self):
kernel_size=[1, 1, 1],
padding='same',
use_bias=False,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer),
# norm
dict(
@@ -349,7 +351,8 @@ def _build_branch_params(self):
kernel_size=[1, 1, 1],
padding='same',
use_bias=False,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer),
# norm
dict(
@@ -371,7 +374,8 @@ def _build_branch_params(self):
norm_momentum=self._norm_momentum,
norm_epsilon=self._norm_epsilon,
temporal_conv_initializer=self._temporal_conv_initializer,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer),
]
branch_2_params = [
@@ -381,7 +385,8 @@ def _build_branch_params(self):
kernel_size=[1, 1, 1],
padding='same',
use_bias=False,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer),
# norm
dict(
@@ -403,7 +408,8 @@ def _build_branch_params(self):
norm_momentum=self._norm_momentum,
norm_epsilon=self._norm_epsilon,
temporal_conv_initializer=self._temporal_conv_initializer,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer)
]
branch_3_params = [
@@ -413,7 +419,8 @@ def _build_branch_params(self):
kernel_size=[1, 1, 1],
padding='same',
use_bias=False,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
kernel_regularizer=self._kernel_regularizer),
# norm
dict(
@@ -453,38 +460,38 @@ def build(self, input_shape):
branch_params = self._build_branch_params()
self._branch_0_layers = [
- tf.keras.layers.Conv3D(**branch_params[0][0]),
+ tf_keras.layers.Conv3D(**branch_params[0][0]),
self._norm(**branch_params[0][1]),
- tf.keras.layers.Activation('relu', **branch_params[0][2]),
+ tf_keras.layers.Activation('relu', **branch_params[0][2]),
]
self._branch_1_layers = [
- tf.keras.layers.Conv3D(**branch_params[1][0]),
+ tf_keras.layers.Conv3D(**branch_params[1][0]),
self._norm(**branch_params[1][1]),
- tf.keras.layers.Activation('relu', **branch_params[1][2]),
+ tf_keras.layers.Activation('relu', **branch_params[1][2]),
self._parameterized_conv_layer(**branch_params[1][3]),
]
self._branch_2_layers = [
- tf.keras.layers.Conv3D(**branch_params[2][0]),
+ tf_keras.layers.Conv3D(**branch_params[2][0]),
self._norm(**branch_params[2][1]),
- tf.keras.layers.Activation('relu', **branch_params[2][2]),
+ tf_keras.layers.Activation('relu', **branch_params[2][2]),
self._parameterized_conv_layer(**branch_params[2][3])
]
if self._swap_pool_and_1x1x1:
self._branch_3_layers = [
- tf.keras.layers.Conv3D(**branch_params[3][0]),
+ tf_keras.layers.Conv3D(**branch_params[3][0]),
self._norm(**branch_params[3][1]),
- tf.keras.layers.Activation('relu', **branch_params[3][2]),
- tf.keras.layers.MaxPool3D(**branch_params[3][3]),
+ tf_keras.layers.Activation('relu', **branch_params[3][2]),
+ tf_keras.layers.MaxPool3D(**branch_params[3][3]),
]
else:
self._branch_3_layers = [
- tf.keras.layers.MaxPool3D(**branch_params[3][3]),
- tf.keras.layers.Conv3D(**branch_params[3][0]),
+ tf_keras.layers.MaxPool3D(**branch_params[3][3]),
+ tf_keras.layers.Conv3D(**branch_params[3][0]),
self._norm(**branch_params[3][1]),
- tf.keras.layers.Activation('relu', **branch_params[3][2]),
+ tf_keras.layers.Activation('relu', **branch_params[3][2]),
]
if self._use_self_gating_on_branch:
diff --git a/official/projects/s3d/modeling/inception_utils_test.py b/official/projects/s3d/modeling/inception_utils_test.py
index 3fa79658dba..1ecb7863a00 100644
--- a/official/projects/s3d/modeling/inception_utils_test.py
+++ b/official/projects/s3d/modeling/inception_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.s3d.modeling import inception_utils
@@ -30,7 +30,7 @@ def test_s3d_stem_cells(self, depth_multiplier, first_temporal_kernel_size,
num_frames = 64
height, width = 224, 224
- inputs = tf.keras.layers.Input(
+ inputs = tf_keras.layers.Input(
shape=(num_frames, height, width, 3), batch_size=batch_size)
outputs, output_endpoints = inception_utils.inception_v1_stem_cells(
@@ -64,7 +64,7 @@ def test_inception_v1_cell_endpoint_match(self, conv_type,
channels = 128
height, width = 28, 28
- inputs = tf.keras.layers.Input(
+ inputs = tf_keras.layers.Input(
shape=(num_frames, height, width, channels), batch_size=batch_size)
inception_v1_cell_layer = inception_utils.InceptionV1CellLayer(
diff --git a/official/projects/s3d/modeling/net_utils.py b/official/projects/s3d/modeling/net_utils.py
index e4f052d946e..d4ec924ec78 100644
--- a/official/projects/s3d/modeling/net_utils.py
+++ b/official/projects/s3d/modeling/net_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,15 +15,16 @@
"""Commonly used TensorFlow 2 network blocks."""
from typing import Any, Text, Sequence, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
WEIGHT_INITIALIZER = {
- 'Xavier': tf.keras.initializers.GlorotUniform,
- 'Gaussian': lambda: tf.keras.initializers.RandomNormal(stddev=0.01),
+ 'Xavier': tf_keras.initializers.GlorotUniform,
+ 'Gaussian': lambda: tf_keras.initializers.RandomNormal(stddev=0.01),
}
-initializers = tf.keras.initializers
-regularizers = tf.keras.regularizers
+initializers = tf_keras.initializers
+regularizers = tf_keras.regularizers
def make_set_from_start_endpoint(start_endpoint: Text,
@@ -44,7 +45,7 @@ def apply_depth_multiplier(d: Union[int, Sequence[Any]],
return [apply_depth_multiplier(x, depth_multiplier) for x in d]
-class ParameterizedConvLayer(tf.keras.layers.Layer):
+class ParameterizedConvLayer(tf_keras.layers.Layer):
"""Convolution layer based on the input conv_type."""
def __init__(
@@ -73,10 +74,10 @@ def __init__(
self._norm_momentum = norm_momentum
self._norm_epsilon = norm_epsilon
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._channel_axis = -1
else:
self._channel_axis = 1
@@ -94,7 +95,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[self._kernel_size] * 3,
strides=self._strides,
dilation_rate=self._rates,
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
))
elif self._conv_type == '2d':
conv_layer_params.append(
@@ -103,7 +105,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[1, self._kernel_size, self._kernel_size],
strides=[1, self._strides[1], self._strides[2]],
dilation_rate=[1, self._rates[1], self._rates[2]],
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
))
elif self._conv_type == '1+2d':
channels_in = input_shape[self._channel_axis]
@@ -113,7 +116,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[self._kernel_size, 1, 1],
strides=[self._strides[0], 1, 1],
dilation_rate=[self._rates[0], 1, 1],
- kernel_initializer=self._temporal_conv_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._temporal_conv_initializer),
))
conv_layer_params.append(
dict(
@@ -121,7 +125,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[1, self._kernel_size, self._kernel_size],
strides=[1, self._strides[1], self._strides[2]],
dilation_rate=[1, self._rates[1], self._rates[2]],
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
))
elif self._conv_type == '2+1d':
conv_layer_params.append(
@@ -130,7 +135,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[1, self._kernel_size, self._kernel_size],
strides=[1, self._strides[1], self._strides[2]],
dilation_rate=[1, self._rates[1], self._rates[2]],
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
))
conv_layer_params.append(
dict(
@@ -138,7 +144,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[self._kernel_size, 1, 1],
strides=[self._strides[0], 1, 1],
dilation_rate=[self._rates[0], 1, 1],
- kernel_initializer=self._temporal_conv_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._temporal_conv_initializer),
))
elif self._conv_type == '1+1+1d':
conv_layer_params.append(
@@ -147,7 +154,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[1, 1, self._kernel_size],
strides=[1, 1, self._strides[2]],
dilation_rate=[1, 1, self._rates[2]],
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
))
conv_layer_params.append(
dict(
@@ -155,7 +163,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[1, self._kernel_size, 1],
strides=[1, self._strides[1], 1],
dilation_rate=[1, self._rates[1], 1],
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
))
conv_layer_params.append(
dict(
@@ -163,7 +172,8 @@ def _build_conv_layer_params(self, input_shape):
kernel_size=[self._kernel_size, 1, 1],
strides=[self._strides[0], 1, 1],
dilation_rate=[self._rates[0], 1, 1],
- kernel_initializer=self._kernel_initializer,
+ kernel_initializer=tf_utils.clone_initializer(
+ self._kernel_initializer),
))
else:
raise ValueError('Unsupported conv_type: {}'.format(self._conv_type))
@@ -185,7 +195,7 @@ def _build_activation_layer_params(self, conv_param):
def _append_conv_layer(self, param):
"""Appends conv, normalization and activation layers."""
self._parameterized_conv_layers.append(
- tf.keras.layers.Conv3D(
+ tf_keras.layers.Conv3D(
padding='same',
use_bias=False,
kernel_regularizer=self._kernel_regularizer,
@@ -196,7 +206,7 @@ def _append_conv_layer(self, param):
relu_layer_params = self._build_activation_layer_params(param)
self._parameterized_conv_layers.append(
- tf.keras.layers.Activation('relu', **relu_layer_params))
+ tf_keras.layers.Activation('relu', **relu_layer_params))
def build(self, input_shape):
self._parameterized_conv_layers = []
diff --git a/official/projects/s3d/modeling/net_utils_test.py b/official/projects/s3d/modeling/net_utils_test.py
index d45c1142878..341d3c6bc47 100644
--- a/official/projects/s3d/modeling/net_utils_test.py
+++ b/official/projects/s3d/modeling/net_utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
from absl import logging
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.s3d.modeling import net_utils
@@ -52,7 +52,7 @@ def test_parameterized_conv_layer_creation(self, conv_type, strides,
name = 'ParameterizedConv'
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(temporal_size, spatial_size, spatial_size, channels),
batch_size=batch_size)
parameterized_conv_layer = net_utils.ParameterizedConvLayer(
diff --git a/official/projects/s3d/modeling/s3d.py b/official/projects/s3d/modeling/s3d.py
index 9b76ad177ed..9eb07f80fb7 100644
--- a/official/projects/s3d/modeling/s3d.py
+++ b/official/projects/s3d/modeling/s3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
"""
from typing import Any, Dict, Mapping, Optional, Sequence, Text, Tuple, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.s3d.configs import s3d as cfg
@@ -28,15 +28,15 @@
from official.vision.modeling import factory_3d as model_factory
from official.vision.modeling.backbones import factory as backbone_factory
-initializers = tf.keras.initializers
-regularizers = tf.keras.regularizers
+initializers = tf_keras.initializers
+regularizers = tf_keras.regularizers
-class S3D(tf.keras.Model):
+class S3D(tf_keras.Model):
"""Class to build S3D family model."""
def __init__(self,
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
final_endpoint: Text = 'Mixed_5c',
first_temporal_kernel_size: int = 3,
temporal_conv_start_at: Text = 'Conv2d_2c_3x3',
@@ -61,7 +61,7 @@ def __init__(self,
"""Constructor.
Args:
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
final_endpoint: Specifies the endpoint to construct the network up to.
first_temporal_kernel_size: Temporal kernel size of the first convolution
layer.
@@ -109,7 +109,7 @@ def __init__(self,
self._self_gating_endpoints = net_utils.make_set_from_start_endpoint(
gating_start_at, inception_utils.INCEPTION_V1_CONV_ENDPOINTS)
- inputs = tf.keras.Input(shape=input_specs.shape[1:])
+ inputs = tf_keras.Input(shape=input_specs.shape[1:])
net, end_points = inception_utils.inception_v1_stem_cells(
inputs,
depth_multiplier,
@@ -146,9 +146,9 @@ def _s3d_cell(
end_point: Text,
end_points: Dict[Text, tf.Tensor],
filters: Union[int, Sequence[Any]],
- non_local_block: Optional[tf.keras.layers.Layer] = None,
- attention_cell: Optional[tf.keras.layers.Layer] = None,
- attention_cell_super_graph: Optional[tf.keras.layers.Layer] = None
+ non_local_block: Optional[tf_keras.layers.Layer] = None,
+ attention_cell: Optional[tf_keras.layers.Layer] = None,
+ attention_cell_super_graph: Optional[tf_keras.layers.Layer] = None
) -> Tuple[tf.Tensor, Dict[Text, tf.Tensor]]:
if end_point.startswith('Mixed'):
conv_type = (
@@ -176,12 +176,13 @@ def _s3d_cell(
kernel_initializer=self._kernel_initializer,
temporal_conv_initializer=self._temporal_conv_initializer,
kernel_regularizer=self._kernel_regularizer,
+ parameterized_conv_layer=self._get_parameterized_conv_layer_impl(),
name=self._get_layer_naming_fn()(end_point))(
net)
else:
- net = tf.keras.layers.MaxPool3D(
- pool_size=filters[0],
- strides=filters[1],
+ net = tf_keras.layers.MaxPool3D(
+ pool_size=filters[0], # pyrefly: ignore[bad-index]
+ strides=filters[1], # pyrefly: ignore[bad-index]
padding='same',
name=self._get_layer_naming_fn()(end_point))(
net)
@@ -237,13 +238,13 @@ def _get_layer_naming_fn(self):
return lambda end_point: None
-class S3DModel(tf.keras.Model):
+class S3DModel(tf_keras.Model):
"""An S3D model builder."""
def __init__(self,
- backbone: tf.keras.Model,
+ backbone: tf_keras.Model,
num_classes: int,
- input_specs: Mapping[Text, tf.keras.layers.InputSpec],
+ input_specs: Mapping[Text, tf_keras.layers.InputSpec],
final_endpoint: Text = 'Mixed_5c',
dropout_rate: float = 0.0,
**kwargs):
@@ -252,7 +253,7 @@ def __init__(self,
Args:
backbone: S3D backbone Keras Model.
num_classes: `int` number of possible classes for video classification.
- input_specs: input_specs: `tf.keras.layers.InputSpec` specs of the input
+ input_specs: input_specs: `tf_keras.layers.InputSpec` specs of the input
tensor.
final_endpoint: Specifies the endpoint to construct the network up to.
dropout_rate: `float` between 0 and 1. Fraction of the input units to
@@ -274,13 +275,13 @@ def __init__(self,
}
inputs = {
- k: tf.keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
+ k: tf_keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
}
streams = self._backbone(inputs['image'])
pool = tf.math.reduce_mean(streams[self._final_endpoint], axis=[1, 2, 3])
- fc = tf.keras.layers.Dropout(dropout_rate)(pool)
- logits = tf.keras.layers.Dense(**self._build_dense_layer_params())(fc)
+ fc = tf_keras.layers.Dropout(dropout_rate)(pool)
+ logits = tf_keras.layers.Dense(**self._build_dense_layer_params())(fc)
super(S3DModel, self).__init__(inputs=inputs, outputs=logits, **kwargs)
@@ -306,11 +307,11 @@ def _build_dense_layer_params(self):
@backbone_factory.register_backbone_builder('s3d')
def build_s3d(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
norm_activation_config: hyperparams.Config,
- l2_regularizer: tf.keras.regularizers.Regularizer = None
-) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None # pyrefly: ignore[bad-function-definition]
+) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds S3D backbone."""
backbone_type = backbone_config.type
@@ -337,11 +338,11 @@ def build_s3d(
@model_factory.register_model_builder('s3d')
def build_s3d_model(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
model_config: cfg.S3DModel,
num_classes: int,
- l2_regularizer: tf.keras.regularizers.Regularizer = None
-) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None # pyrefly: ignore[bad-function-definition]
+) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds S3D model with classification layer."""
input_specs_dict = {'image': input_specs}
backbone = build_s3d(input_specs, model_config.backbone,
diff --git a/official/projects/s3d/modeling/s3d_test.py b/official/projects/s3d/modeling/s3d_test.py
index d9565aa4700..eb59a33f672 100644
--- a/official/projects/s3d/modeling/s3d_test.py
+++ b/official/projects/s3d/modeling/s3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for S3D model."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.s3d.modeling import s3d
@@ -36,11 +36,11 @@ def test_build(self, num_frames, height, width, first_temporal_kernel_size):
batch_size = 5
input_shape = [batch_size, num_frames, height, width, 3]
- input_specs = tf.keras.layers.InputSpec(shape=input_shape)
+ input_specs = tf_keras.layers.InputSpec(shape=input_shape)
network = s3d.S3D(
input_specs=input_specs
)
- inputs = tf.keras.Input(shape=input_shape[1:], batch_size=input_shape[0])
+ inputs = tf_keras.Input(shape=input_shape[1:], batch_size=input_shape[0])
endpoints = network(inputs)
temporal_1a = (num_frames - 1)//2 + 1
@@ -71,7 +71,7 @@ def test_build(self, num_frames, height, width, first_temporal_kernel_size):
def test_serialize_deserialize(self):
# Create a network object that sets all of its config options.
kwargs = dict(
- input_specs=tf.keras.layers.InputSpec(shape=(5, 64, 224, 224, 3)),
+ input_specs=tf_keras.layers.InputSpec(shape=(5, 64, 224, 224, 3)),
final_endpoint='Mixed_5c',
first_temporal_kernel_size=3,
temporal_conv_start_at='Conv2d_2c_3x3',
@@ -81,7 +81,7 @@ def test_serialize_deserialize(self):
use_sync_bn=False,
norm_momentum=0.999,
norm_epsilon=0.001,
- temporal_conv_initializer=tf.keras.initializers.TruncatedNormal(
+ temporal_conv_initializer=tf_keras.initializers.TruncatedNormal(
mean=0.0, stddev=0.01),
temporal_conv_type='2+1d',
kernel_initializer='truncated_normal',
diff --git a/official/projects/s3d/train.py b/official/projects/s3d/train.py
index 5f1819fc2ea..c0c7fae93e8 100644
--- a/official/projects/s3d/train.py
+++ b/official/projects/s3d/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/beta/projects/simclr/README.md b/official/projects/simclr/README.md
similarity index 95%
rename from official/vision/beta/projects/simclr/README.md
rename to official/projects/simclr/README.md
index 91b4375bd60..3644532032a 100644
--- a/official/vision/beta/projects/simclr/README.md
+++ b/official/projects/simclr/README.md
@@ -10,7 +10,7 @@
An illustration of SimCLR (from our blog here).
-## Enviroment setup
+## Environment setup
The code can be run on multiple GPUs or TPUs with different distribution
strategies. See the TensorFlow distributed training
@@ -25,7 +25,7 @@ install -r ./official/requirements.txt`
To pretrain the model on Imagenet, try the following command:
```
-python3 -m official.vision.beta.projects.simclr.train \
+python3 -m official.projects.simclr.train \
--mode=train_and_eval \
--experiment=simclr_pretraining \
--model_dir={MODEL_DIR} \
@@ -44,7 +44,7 @@ You can also find image IDs of these subsets in `imagenet_subsets/`.
To fine-tune the whole network, refer to the following command:
```
-python3 -m official.vision.beta.projects.simclr.train \
+python3 -m official.projects.simclr.train \
--mode=train_and_eval \
--experiment=simclr_finetuning \
--model_dir={MODEL_DIR} \
diff --git a/official/projects/simclr/common/registry_imports.py b/official/projects/simclr/common/registry_imports.py
new file mode 100644
index 00000000000..f83e092be56
--- /dev/null
+++ b/official/projects/simclr/common/registry_imports.py
@@ -0,0 +1,22 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""All necessary imports for registration."""
+
+# pylint: disable=unused-import
+from official.projects.simclr.configs import simclr
+from official.projects.simclr.losses import contrastive_losses
+from official.projects.simclr.modeling import simclr_model
+from official.projects.simclr.tasks import simclr as simclr_task
+from official.vision import registry_imports
diff --git a/official/vision/beta/projects/simclr/configs/experiments/cifar_simclr_pretrain.yaml b/official/projects/simclr/configs/experiments/cifar_simclr_pretrain.yaml
similarity index 100%
rename from official/vision/beta/projects/simclr/configs/experiments/cifar_simclr_pretrain.yaml
rename to official/projects/simclr/configs/experiments/cifar_simclr_pretrain.yaml
diff --git a/official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_finetune_gpu.yaml b/official/projects/simclr/configs/experiments/imagenet_simclr_finetune_gpu.yaml
similarity index 100%
rename from official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_finetune_gpu.yaml
rename to official/projects/simclr/configs/experiments/imagenet_simclr_finetune_gpu.yaml
diff --git a/official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_finetune_tpu.yaml b/official/projects/simclr/configs/experiments/imagenet_simclr_finetune_tpu.yaml
similarity index 100%
rename from official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_finetune_tpu.yaml
rename to official/projects/simclr/configs/experiments/imagenet_simclr_finetune_tpu.yaml
diff --git a/official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_multitask_tpu.yaml b/official/projects/simclr/configs/experiments/imagenet_simclr_multitask_tpu.yaml
similarity index 100%
rename from official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_multitask_tpu.yaml
rename to official/projects/simclr/configs/experiments/imagenet_simclr_multitask_tpu.yaml
diff --git a/official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_pretrain_gpu.yaml b/official/projects/simclr/configs/experiments/imagenet_simclr_pretrain_gpu.yaml
similarity index 100%
rename from official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_pretrain_gpu.yaml
rename to official/projects/simclr/configs/experiments/imagenet_simclr_pretrain_gpu.yaml
diff --git a/official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_pretrain_tpu.yaml b/official/projects/simclr/configs/experiments/imagenet_simclr_pretrain_tpu.yaml
similarity index 100%
rename from official/vision/beta/projects/simclr/configs/experiments/imagenet_simclr_pretrain_tpu.yaml
rename to official/projects/simclr/configs/experiments/imagenet_simclr_pretrain_tpu.yaml
diff --git a/official/vision/beta/projects/simclr/configs/multitask_config.py b/official/projects/simclr/configs/multitask_config.py
similarity index 72%
rename from official/vision/beta/projects/simclr/configs/multitask_config.py
rename to official/projects/simclr/configs/multitask_config.py
index 241a290521c..348d349d127 100644
--- a/official/vision/beta/projects/simclr/configs/multitask_config.py
+++ b/official/projects/simclr/configs/multitask_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,8 +20,8 @@
from official.core import exp_factory
from official.modeling import hyperparams
from official.modeling.multitask import configs as multitask_configs
-from official.vision.beta.projects.simclr.configs import simclr as simclr_configs
-from official.vision.beta.projects.simclr.modeling import simclr_model
+from official.projects.simclr.configs import simclr as simclr_configs
+from official.projects.simclr.modeling import simclr_model
from official.vision.configs import backbones
from official.vision.configs import common
@@ -31,8 +31,9 @@ class SimCLRMTHeadConfig(hyperparams.Config):
"""Per-task specific configs."""
task_name: str = 'task_name'
# Supervised head is required for finetune, but optional for pretrain.
- supervised_head: simclr_configs.SupervisedHead = simclr_configs.SupervisedHead(
- num_classes=1001)
+ supervised_head: simclr_configs.SupervisedHead = dataclasses.field(
+ default_factory=lambda: simclr_configs.SupervisedHead(num_classes=1001)
+ )
mode: str = simclr_model.PRETRAIN
@@ -40,13 +41,22 @@ class SimCLRMTHeadConfig(hyperparams.Config):
class SimCLRMTModelConfig(hyperparams.Config):
"""Model config for multi-task SimCLR model."""
input_size: List[int] = dataclasses.field(default_factory=list)
- backbone: backbones.Backbone = backbones.Backbone(
- type='resnet', resnet=backbones.ResNet())
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='resnet', resnet=backbones.ResNet()
+ )
+ )
backbone_trainable: bool = True
- projection_head: simclr_configs.ProjectionHead = simclr_configs.ProjectionHead(
- proj_output_dim=128, num_proj_layers=3, ft_proj_idx=1)
- norm_activation: common.NormActivation = common.NormActivation(
- norm_momentum=0.9, norm_epsilon=1e-5, use_sync_bn=False)
+ projection_head: simclr_configs.ProjectionHead = dataclasses.field(
+ default_factory=lambda: simclr_configs.ProjectionHead( # pylint: disable=g-long-lambda
+ proj_output_dim=128, num_proj_layers=3, ft_proj_idx=1
+ )
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ norm_momentum=0.9, norm_epsilon=1e-5, use_sync_bn=False
+ )
+ )
heads: Tuple[SimCLRMTHeadConfig, ...] = ()
# L2 weight decay is used in the model, not in task.
# Note that this can not be used together with lars optimizer.
diff --git a/official/vision/beta/projects/simclr/configs/multitask_config_test.py b/official/projects/simclr/configs/multitask_config_test.py
similarity index 83%
rename from official/vision/beta/projects/simclr/configs/multitask_config_test.py
rename to official/projects/simclr/configs/multitask_config_test.py
index fbb1aed8dee..3176932efbe 100644
--- a/official/vision/beta/projects/simclr/configs/multitask_config_test.py
+++ b/official/projects/simclr/configs/multitask_config_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,12 +14,12 @@
"""Tests for multitask_config."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import exp_factory
from official.modeling.multitask import configs as multitask_configs
-from official.vision.beta.projects.simclr.configs import multitask_config as simclr_multitask_config
-from official.vision.beta.projects.simclr.configs import simclr as exp_cfg
+from official.projects.simclr.configs import multitask_config as simclr_multitask_config
+from official.projects.simclr.configs import simclr as exp_cfg
class MultitaskConfigTest(tf.test.TestCase):
diff --git a/official/vision/beta/projects/simclr/configs/simclr.py b/official/projects/simclr/configs/simclr.py
similarity index 80%
rename from official/vision/beta/projects/simclr/configs/simclr.py
rename to official/projects/simclr/configs/simclr.py
index 051e4459a87..b750a4246e3 100644
--- a/official/vision/beta/projects/simclr/configs/simclr.py
+++ b/official/projects/simclr/configs/simclr.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from official.core import exp_factory
from official.modeling import hyperparams
from official.modeling import optimization
-from official.vision.beta.projects.simclr.modeling import simclr_model
+from official.projects.simclr.modeling import simclr_model
from official.vision.configs import backbones
from official.vision.configs import common
@@ -55,8 +55,8 @@ class DataConfig(cfg.DataConfig):
shuffle_buffer_size: int = 10000
cycle_length: int = 10
# simclr specific configs
- parser: Parser = Parser()
- decoder: Decoder = Decoder()
+ parser: Parser = dataclasses.field(default_factory=Parser)
+ decoder: Decoder = dataclasses.field(default_factory=Decoder)
# Useful when doing a sanity check that we absolutely use no labels while
# pretrain by setting labels to zeros (default = False, keep original labels)
input_set_label_to_zero: bool = False
@@ -99,13 +99,24 @@ class Evaluation(hyperparams.Config):
class SimCLRModel(hyperparams.Config):
"""SimCLR model config."""
input_size: List[int] = dataclasses.field(default_factory=list)
- backbone: backbones.Backbone = backbones.Backbone(
- type='resnet', resnet=backbones.ResNet())
- projection_head: ProjectionHead = ProjectionHead(
- proj_output_dim=128, num_proj_layers=3, ft_proj_idx=1)
- supervised_head: SupervisedHead = SupervisedHead(num_classes=1001)
- norm_activation: common.NormActivation = common.NormActivation(
- norm_momentum=0.9, norm_epsilon=1e-5, use_sync_bn=False)
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='resnet', resnet=backbones.ResNet()
+ )
+ )
+ projection_head: ProjectionHead = dataclasses.field(
+ default_factory=lambda: ProjectionHead( # pylint: disable=g-long-lambda
+ proj_output_dim=128, num_proj_layers=3, ft_proj_idx=1
+ )
+ )
+ supervised_head: SupervisedHead = dataclasses.field(
+ default_factory=lambda: SupervisedHead(num_classes=1001)
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ norm_momentum=0.9, norm_epsilon=1e-5, use_sync_bn=False
+ )
+ )
mode: str = simclr_model.PRETRAIN
backbone_trainable: bool = True
@@ -113,13 +124,21 @@ class SimCLRModel(hyperparams.Config):
@dataclasses.dataclass
class SimCLRPretrainTask(cfg.TaskConfig):
"""SimCLR pretraining task config."""
- model: SimCLRModel = SimCLRModel(mode=simclr_model.PRETRAIN)
- train_data: DataConfig = DataConfig(
- parser=Parser(mode=simclr_model.PRETRAIN), is_training=True)
- validation_data: DataConfig = DataConfig(
- parser=Parser(mode=simclr_model.PRETRAIN), is_training=False)
- loss: ContrastiveLoss = ContrastiveLoss()
- evaluation: Evaluation = Evaluation()
+ model: SimCLRModel = dataclasses.field(
+ default_factory=lambda: SimCLRModel(mode=simclr_model.PRETRAIN)
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ parser=Parser(mode=simclr_model.PRETRAIN), is_training=True
+ )
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ parser=Parser(mode=simclr_model.PRETRAIN), is_training=False
+ )
+ )
+ loss: ContrastiveLoss = dataclasses.field(default_factory=ContrastiveLoss)
+ evaluation: Evaluation = dataclasses.field(default_factory=Evaluation)
init_checkpoint: Optional[str] = None
# all or backbone
init_checkpoint_modules: str = 'all'
@@ -128,15 +147,26 @@ class SimCLRPretrainTask(cfg.TaskConfig):
@dataclasses.dataclass
class SimCLRFinetuneTask(cfg.TaskConfig):
"""SimCLR fine tune task config."""
- model: SimCLRModel = SimCLRModel(
- mode=simclr_model.FINETUNE,
- supervised_head=SupervisedHead(num_classes=1001, zero_init=True))
- train_data: DataConfig = DataConfig(
- parser=Parser(mode=simclr_model.FINETUNE), is_training=True)
- validation_data: DataConfig = DataConfig(
- parser=Parser(mode=simclr_model.FINETUNE), is_training=False)
- loss: ClassificationLosses = ClassificationLosses()
- evaluation: Evaluation = Evaluation()
+ model: SimCLRModel = dataclasses.field(
+ default_factory=lambda: SimCLRModel( # pylint: disable=g-long-lambda
+ mode=simclr_model.FINETUNE,
+ supervised_head=SupervisedHead(num_classes=1001, zero_init=True),
+ )
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ parser=Parser(mode=simclr_model.FINETUNE), is_training=True
+ )
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ parser=Parser(mode=simclr_model.FINETUNE), is_training=False
+ )
+ )
+ loss: ClassificationLosses = dataclasses.field(
+ default_factory=ClassificationLosses
+ )
+ evaluation: Evaluation = dataclasses.field(default_factory=Evaluation)
init_checkpoint: Optional[str] = None
# all, backbone_projection or backbone
init_checkpoint_modules: str = 'backbone_projection'
diff --git a/official/vision/beta/projects/simclr/configs/simclr_test.py b/official/projects/simclr/configs/simclr_test.py
similarity index 85%
rename from official/vision/beta/projects/simclr/configs/simclr_test.py
rename to official/projects/simclr/configs/simclr_test.py
index 92ac4d278ac..ef873caec07 100644
--- a/official/vision/beta/projects/simclr/configs/simclr_test.py
+++ b/official/projects/simclr/configs/simclr_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,12 +15,12 @@
"""Tests for SimCLR config."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import exp_factory
-from official.vision.beta.projects.simclr.common import registry_imports # pylint: disable=unused-import
-from official.vision.beta.projects.simclr.configs import simclr as exp_cfg
+from official.projects.simclr.common import registry_imports # pylint: disable=unused-import
+from official.projects.simclr.configs import simclr as exp_cfg
class SimCLRConfigTest(tf.test.TestCase, parameterized.TestCase):
diff --git a/official/vision/beta/projects/simclr/dataloaders/preprocess_ops.py b/official/projects/simclr/dataloaders/preprocess_ops.py
similarity index 99%
rename from official/vision/beta/projects/simclr/dataloaders/preprocess_ops.py
rename to official/projects/simclr/dataloaders/preprocess_ops.py
index 081621466e8..53141f02dbe 100644
--- a/official/vision/beta/projects/simclr/dataloaders/preprocess_ops.py
+++ b/official/projects/simclr/dataloaders/preprocess_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Preprocessing ops."""
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
CROP_PROPORTION = 0.875 # Standard for ImageNet.
diff --git a/official/vision/beta/projects/simclr/dataloaders/simclr_input.py b/official/projects/simclr/dataloaders/simclr_input.py
similarity index 96%
rename from official/vision/beta/projects/simclr/dataloaders/simclr_input.py
rename to official/projects/simclr/dataloaders/simclr_input.py
index d7ac57dc16e..3862e8eef05 100644
--- a/official/vision/beta/projects/simclr/dataloaders/simclr_input.py
+++ b/official/projects/simclr/dataloaders/simclr_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -38,10 +38,10 @@
from typing import List
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.simclr.dataloaders import preprocess_ops as simclr_preprocess_ops
-from official.vision.beta.projects.simclr.modeling import simclr_model
+from official.projects.simclr.dataloaders import preprocess_ops as simclr_preprocess_ops
+from official.projects.simclr.modeling import simclr_model
from official.vision.dataloaders import decoder
from official.vision.dataloaders import parser
from official.vision.ops import preprocess_ops
diff --git a/official/vision/beta/projects/simclr/heads/simclr_head.py b/official/projects/simclr/heads/simclr_head.py
similarity index 92%
rename from official/vision/beta/projects/simclr/heads/simclr_head.py
rename to official/projects/simclr/heads/simclr_head.py
index 9213a81f0bb..88fc8e62f2c 100644
--- a/official/vision/beta/projects/simclr/heads/simclr_head.py
+++ b/official/projects/simclr/heads/simclr_head.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,17 +14,17 @@
"""SimCLR prediction heads."""
-from typing import Text, Optional
+from typing import Optional, Text
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.simclr.modeling.layers import nn_blocks
+from official.projects.simclr.modeling.layers import nn_blocks
-regularizers = tf.keras.regularizers
-layers = tf.keras.layers
+regularizers = tf_keras.regularizers
+layers = tf_keras.layers
-class ProjectionHead(tf.keras.layers.Layer):
+class ProjectionHead(tf_keras.layers.Layer):
"""Projection head."""
def __init__(
@@ -48,9 +48,9 @@ def __init__(
ft_proj_idx: `int` index of layer to use during fine-tuning. 0 means no
projection head during fine tuning, -1 means the final layer.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
use_sync_bn: if True, use synchronized batch normalization.
norm_momentum: `float` normalization omentum for the moving average.
@@ -143,7 +143,7 @@ def call(self, inputs, training=None):
return proj_head_output, proj_finetune_output
-class ClassificationHead(tf.keras.layers.Layer):
+class ClassificationHead(tf_keras.layers.Layer):
"""Classification Head."""
def __init__(
@@ -160,9 +160,9 @@ def __init__(
num_classes: `int` size of the output dimension or number of classes
for classification task.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
name: `str`, name of the layer.
**kwargs: keyword arguments to be passed.
diff --git a/official/vision/beta/projects/simclr/heads/simclr_head_test.py b/official/projects/simclr/heads/simclr_head_test.py
similarity index 92%
rename from official/vision/beta/projects/simclr/heads/simclr_head_test.py
rename to official/projects/simclr/heads/simclr_head_test.py
index 20ed748ee11..e8f17b5c23d 100644
--- a/official/vision/beta/projects/simclr/heads/simclr_head_test.py
+++ b/official/projects/simclr/heads/simclr_head_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.simclr.heads import simclr_head
+from official.projects.simclr.heads import simclr_head
class ProjectionHeadTest(tf.test.TestCase, parameterized.TestCase):
@@ -33,7 +33,7 @@ def test_head_creation(self, num_proj_layers, proj_output_dim):
proj_output_dim=proj_output_dim)
input_dim = 64
- x = tf.keras.Input(shape=(input_dim,))
+ x = tf_keras.Input(shape=(input_dim,))
proj_head_output, proj_finetune_output = test_layer(x)
proj_head_output_dim = input_dim
@@ -90,7 +90,7 @@ def test_head_creation(self, num_classes):
test_layer = simclr_head.ClassificationHead(num_classes=num_classes)
input_dim = 64
- x = tf.keras.Input(shape=(input_dim,))
+ x = tf_keras.Input(shape=(input_dim,))
out_x = test_layer(x)
self.assertAllEqual(out_x.shape.as_list(),
diff --git a/official/vision/beta/projects/simclr/losses/contrastive_losses.py b/official/projects/simclr/losses/contrastive_losses.py
similarity index 98%
rename from official/vision/beta/projects/simclr/losses/contrastive_losses.py
rename to official/projects/simclr/losses/contrastive_losses.py
index f16a7b723f5..62e6b6512a1 100644
--- a/official/vision/beta/projects/simclr/losses/contrastive_losses.py
+++ b/official/projects/simclr/losses/contrastive_losses.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
LARGE_NUM = 1e9
diff --git a/official/vision/beta/projects/simclr/losses/contrastive_losses_test.py b/official/projects/simclr/losses/contrastive_losses_test.py
similarity index 93%
rename from official/vision/beta/projects/simclr/losses/contrastive_losses_test.py
rename to official/projects/simclr/losses/contrastive_losses_test.py
index 9da5078295c..8f4c02b213f 100644
--- a/official/vision/beta/projects/simclr/losses/contrastive_losses_test.py
+++ b/official/projects/simclr/losses/contrastive_losses_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.simclr.losses import contrastive_losses
+from official.projects.simclr.losses import contrastive_losses
class ContrastiveLossesTest(tf.test.TestCase, parameterized.TestCase):
diff --git a/official/vision/beta/projects/simclr/modeling/layers/nn_blocks.py b/official/projects/simclr/modeling/layers/nn_blocks.py
similarity index 87%
rename from official/vision/beta/projects/simclr/modeling/layers/nn_blocks.py
rename to official/projects/simclr/modeling/layers/nn_blocks.py
index 013a7be5201..62b6e0938d8 100644
--- a/official/vision/beta/projects/simclr/modeling/layers/nn_blocks.py
+++ b/official/projects/simclr/modeling/layers/nn_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,14 +15,14 @@
"""Contains common building blocks for simclr neural networks."""
from typing import Text, Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
-regularizers = tf.keras.regularizers
+regularizers = tf_keras.regularizers
-class DenseBN(tf.keras.layers.Layer):
+class DenseBN(tf_keras.layers.Layer):
"""Modified Dense layer to help build simclr system.
The layer is a standards combination of Dense, BatchNorm and Activation.
@@ -54,9 +54,9 @@ def __init__(
zero.
activation: `str` name of the activation function.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
name: `str`, name of the layer.
**kwargs: keyword arguments to be passed.
@@ -77,10 +77,10 @@ def __init__(
self._name = name
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -106,7 +106,7 @@ def get_config(self):
return dict(list(base_config.items()) + list(config.items()))
def build(self, input_shape):
- self._dense0 = tf.keras.layers.Dense(
+ self._dense0 = tf_keras.layers.Dense(
self._output_dim,
kernel_initializer=self._kernel_initializer,
kernel_regularizer=self._kernel_regularizer,
@@ -129,5 +129,5 @@ def call(self, inputs, training=None):
if self._use_normalization:
x = self._norm0(x)
if self._activation:
- x = self._activation_fn(x)
+ x = self._activation_fn(x) # pyrefly: ignore[not-callable]
return x
diff --git a/official/vision/beta/projects/simclr/modeling/layers/nn_blocks_test.py b/official/projects/simclr/modeling/layers/nn_blocks_test.py
similarity index 88%
rename from official/vision/beta/projects/simclr/modeling/layers/nn_blocks_test.py
rename to official/projects/simclr/modeling/layers/nn_blocks_test.py
index f1b6b4f4cf6..4d4abf74d35 100644
--- a/official/vision/beta/projects/simclr/modeling/layers/nn_blocks_test.py
+++ b/official/projects/simclr/modeling/layers/nn_blocks_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,9 @@
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.simclr.modeling.layers import nn_blocks
+from official.projects.simclr.modeling.layers import nn_blocks
class DenseBNTest(tf.test.TestCase, parameterized.TestCase):
@@ -33,7 +33,7 @@ def test_pass_through(self, output_dim, use_bias, use_normalization):
use_normalization=use_normalization
)
- x = tf.keras.Input(shape=(64,))
+ x = tf_keras.Input(shape=(64,))
out_x = test_layer(x)
self.assertAllEqual(out_x.shape.as_list(), [None, output_dim])
diff --git a/official/vision/beta/projects/simclr/modeling/multitask_model.py b/official/projects/simclr/modeling/multitask_model.py
similarity index 90%
rename from official/vision/beta/projects/simclr/modeling/multitask_model.py
rename to official/projects/simclr/modeling/multitask_model.py
index 0ca7d44f9ab..959585d74e7 100644
--- a/official/vision/beta/projects/simclr/modeling/multitask_model.py
+++ b/official/projects/simclr/modeling/multitask_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,12 +16,12 @@
from typing import Dict, Text
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling.multitask import base_model
-from official.vision.beta.projects.simclr.configs import multitask_config as simclr_multitask_config
-from official.vision.beta.projects.simclr.heads import simclr_head
-from official.vision.beta.projects.simclr.modeling import simclr_model
+from official.projects.simclr.configs import multitask_config as simclr_multitask_config
+from official.projects.simclr.heads import simclr_head
+from official.projects.simclr.modeling import simclr_model
from official.vision.modeling import backbones
PROJECTION_OUTPUT_KEY = 'projection_outputs'
@@ -36,7 +36,7 @@ def __init__(self, config: simclr_multitask_config.SimCLRMTModelConfig,
self._config = config
# Build shared backbone.
- self._input_specs = tf.keras.layers.InputSpec(shape=[None] +
+ self._input_specs = tf_keras.layers.InputSpec(shape=[None] +
config.input_size)
l2_weight_decay = config.l2_weight_decay
@@ -44,7 +44,7 @@ def __init__(self, config: simclr_multitask_config.SimCLRMTModelConfig,
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
self._l2_regularizer = (
- tf.keras.regularizers.l2(l2_weight_decay /
+ tf_keras.regularizers.l2(l2_weight_decay /
2.0) if l2_weight_decay else None)
self._backbone = backbones.factory.build_backbone(
@@ -67,7 +67,7 @@ def __init__(self, config: simclr_multitask_config.SimCLRMTModelConfig,
super().__init__(**kwargs)
- def _instantiate_sub_tasks(self) -> Dict[Text, tf.keras.Model]:
+ def _instantiate_sub_tasks(self) -> Dict[Text, tf_keras.Model]:
tasks = {}
for model_config in self._config.heads:
diff --git a/official/vision/beta/projects/simclr/modeling/multitask_model_test.py b/official/projects/simclr/modeling/multitask_model_test.py
similarity index 81%
rename from official/vision/beta/projects/simclr/modeling/multitask_model_test.py
rename to official/projects/simclr/modeling/multitask_model_test.py
index e0d6c2ccaf2..0cd0394cfd0 100644
--- a/official/vision/beta/projects/simclr/modeling/multitask_model_test.py
+++ b/official/projects/simclr/modeling/multitask_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,11 +16,11 @@
import os.path
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.simclr.configs import multitask_config
-from official.vision.beta.projects.simclr.modeling import multitask_model
-from official.vision.beta.projects.simclr.modeling import simclr_model
+from official.projects.simclr.configs import multitask_config
+from official.projects.simclr.modeling import multitask_model
+from official.projects.simclr.modeling import simclr_model
class MultitaskModelTest(tf.test.TestCase):
diff --git a/official/vision/beta/projects/simclr/modeling/simclr_model.py b/official/projects/simclr/modeling/simclr_model.py
similarity index 90%
rename from official/vision/beta/projects/simclr/modeling/simclr_model.py
rename to official/projects/simclr/modeling/simclr_model.py
index da8a6e3572c..a133bec9062 100644
--- a/official/vision/beta/projects/simclr/modeling/simclr_model.py
+++ b/official/projects/simclr/modeling/simclr_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,9 @@
from typing import Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-layers = tf.keras.layers
+layers = tf_keras.layers
PRETRAIN = 'pretrain'
FINETUNE = 'finetune'
@@ -27,13 +27,13 @@
SUPERVISED_OUTPUT_KEY = 'supervised_outputs'
-class SimCLRModel(tf.keras.Model):
+class SimCLRModel(tf_keras.Model):
"""A classification model based on SimCLR framework."""
def __init__(self,
- backbone: tf.keras.models.Model,
- projection_head: tf.keras.layers.Layer,
- supervised_head: Optional[tf.keras.layers.Layer] = None,
+ backbone: tf_keras.models.Model,
+ projection_head: tf_keras.layers.Layer,
+ supervised_head: Optional[tf_keras.layers.Layer] = None,
input_specs=layers.InputSpec(shape=[None, None, None, 3]),
mode: str = PRETRAIN,
backbone_trainable: bool = True,
@@ -45,7 +45,7 @@ def __init__(self,
projection_head: a projection head network.
supervised_head: a head network for supervised learning, e.g.
classification head.
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
mode: `str` indicates mode of training to be executed.
backbone_trainable: `bool` whether the backbone is trainable or not.
**kwargs: keyword arguments to be passed.
@@ -69,7 +69,7 @@ def __init__(self,
# Set whether the backbone is trainable
self._backbone.trainable = backbone_trainable
- def call(self, inputs, training=None, **kwargs):
+ def call(self, inputs, training=None, **kwargs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
model_outputs = {}
if training and self._mode == PRETRAIN:
diff --git a/official/vision/beta/projects/simclr/modeling/simclr_model_test.py b/official/projects/simclr/modeling/simclr_model_test.py
similarity index 86%
rename from official/vision/beta/projects/simclr/modeling/simclr_model_test.py
rename to official/projects/simclr/modeling/simclr_model_test.py
index 84e7b809ca9..ad3f3d657c3 100644
--- a/official/vision/beta/projects/simclr/modeling/simclr_model_test.py
+++ b/official/projects/simclr/modeling/simclr_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,10 @@
"""Test for SimCLR model."""
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.simclr.heads import simclr_head
-from official.vision.beta.projects.simclr.modeling import simclr_model
+from official.projects.simclr.heads import simclr_head
+from official.projects.simclr.modeling import simclr_model
from official.vision.modeling import backbones
@@ -33,10 +33,10 @@ class SimCLRModelTest(parameterized.TestCase, tf.test.TestCase):
def test_model_creation(self, project_dim, num_proj_layers, ft_proj_idx):
input_size = 224
inputs = np.random.rand(2, input_size, input_size, 3)
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None, input_size, input_size, 3])
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = backbones.ResNet(model_id=50, activation='relu',
input_specs=input_specs)
diff --git a/official/vision/beta/projects/simclr/multitask_train.py b/official/projects/simclr/multitask_train.py
similarity index 89%
rename from official/vision/beta/projects/simclr/multitask_train.py
rename to official/projects/simclr/multitask_train.py
index 65f23be960a..7887fab0150 100644
--- a/official/vision/beta/projects/simclr/multitask_train.py
+++ b/official/projects/simclr/multitask_train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,9 +25,9 @@
from official.modeling.multitask import train_lib
# pylint: disable=unused-import
-from official.vision.beta.projects.simclr.common import registry_imports
-from official.vision.beta.projects.simclr.configs import multitask_config
-from official.vision.beta.projects.simclr.modeling import multitask_model
+from official.projects.simclr.common import registry_imports
+from official.projects.simclr.configs import multitask_config
+from official.projects.simclr.modeling import multitask_model
# pylint: enable=unused-import
FLAGS = flags.FLAGS
diff --git a/official/vision/beta/projects/simclr/tasks/simclr.py b/official/projects/simclr/tasks/simclr.py
similarity index 91%
rename from official/vision/beta/projects/simclr/tasks/simclr.py
rename to official/projects/simclr/tasks/simclr.py
index 019fe99118e..6f914882e01 100644
--- a/official/vision/beta/projects/simclr/tasks/simclr.py
+++ b/official/projects/simclr/tasks/simclr.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -27,7 +27,7 @@
from typing import Dict, Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions
@@ -36,11 +36,11 @@
from official.modeling import optimization
from official.modeling import performance
from official.modeling import tf_utils
-from official.vision.beta.projects.simclr.configs import simclr as exp_cfg
-from official.vision.beta.projects.simclr.dataloaders import simclr_input
-from official.vision.beta.projects.simclr.heads import simclr_head
-from official.vision.beta.projects.simclr.losses import contrastive_losses
-from official.vision.beta.projects.simclr.modeling import simclr_model
+from official.projects.simclr.configs import simclr as exp_cfg
+from official.projects.simclr.dataloaders import simclr_input
+from official.projects.simclr.heads import simclr_head
+from official.projects.simclr.losses import contrastive_losses
+from official.projects.simclr.modeling import simclr_model
from official.vision.modeling import backbones
OptimizationConfig = optimization.OptimizationConfig
@@ -82,7 +82,7 @@ def create_optimizer(self,
def build_model(self):
model_config = self.task_config.model
- input_specs = tf.keras.layers.InputSpec(shape=[None] +
+ input_specs = tf_keras.layers.InputSpec(shape=[None] +
model_config.input_size)
l2_weight_decay = self.task_config.loss.l2_weight_decay
@@ -90,7 +90,7 @@ def build_model(self):
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
l2_regularizer = (
- tf.keras.regularizers.l2(l2_weight_decay /
+ tf_keras.regularizers.l2(l2_weight_decay /
2.0) if l2_weight_decay else None)
# Build backbone
@@ -138,7 +138,7 @@ def build_model(self):
return model
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loading pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -230,12 +230,12 @@ def build_losses(self,
labels = tf.concat([labels, labels], 0)
if self.task_config.evaluation.one_hot:
- sup_loss = tf.keras.losses.CategoricalCrossentropy(
- from_logits=True, reduction=tf.keras.losses.Reduction.NONE)(labels,
+ sup_loss = tf_keras.losses.CategoricalCrossentropy(
+ from_logits=True, reduction=tf_keras.losses.Reduction.NONE)(labels,
outputs)
else:
- sup_loss = tf.keras.losses.SparseCategoricalCrossentropy(
- from_logits=True, reduction=tf.keras.losses.Reduction.NONE)(labels,
+ sup_loss = tf_keras.losses.SparseCategoricalCrossentropy(
+ from_logits=True, reduction=tf_keras.losses.Reduction.NONE)(labels,
outputs)
sup_loss = tf.reduce_mean(sup_loss)
@@ -269,19 +269,19 @@ def build_metrics(self, training=True):
if self.task_config.model.supervised_head:
metric_names.extend(['supervised_loss', 'accuracy'])
for name in metric_names:
- metrics.append(tf.keras.metrics.Mean(name, dtype=tf.float32))
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
else:
k = self.task_config.evaluation.top_k
if self.task_config.evaluation.one_hot:
metrics = [
- tf.keras.metrics.CategoricalAccuracy(name='accuracy'),
- tf.keras.metrics.TopKCategoricalAccuracy(
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.TopKCategoricalAccuracy(
k=k, name='top_{}_accuracy'.format(k))
]
else:
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
- tf.keras.metrics.SparseTopKCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.SparseTopKCategoricalAccuracy(
k=k, name='top_{}_accuracy'.format(k))
]
return metrics
@@ -313,7 +313,7 @@ def train_step(self, inputs, model, optimizer, metrics=None):
scaled_loss = losses['total_loss'] / num_replicas
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
@@ -322,13 +322,13 @@ def train_step(self, inputs, model, optimizer, metrics=None):
logging.info(var.name)
grads = tape.gradient(scaled_loss, tvars)
# Scales back gradient when LossScaleOptimizer is used.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
logs = {self.loss: losses['total_loss']}
- for m in metrics:
+ for m in metrics: # pyrefly: ignore[not-iterable]
m.update_state(losses[m.name])
logs.update({m.name: m.result()})
@@ -395,7 +395,7 @@ def create_optimizer(self,
def build_model(self):
model_config = self.task_config.model
- input_specs = tf.keras.layers.InputSpec(shape=[None] +
+ input_specs = tf_keras.layers.InputSpec(shape=[None] +
model_config.input_size)
l2_weight_decay = self.task_config.loss.l2_weight_decay
@@ -403,7 +403,7 @@ def build_model(self):
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
l2_regularizer = (
- tf.keras.regularizers.l2(l2_weight_decay /
+ tf_keras.regularizers.l2(l2_weight_decay /
2.0) if l2_weight_decay else None)
backbone = backbones.factory.build_backbone(
@@ -445,7 +445,7 @@ def build_model(self):
return model
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loading pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -514,13 +514,13 @@ def build_losses(self, labels, model_outputs, aux_losses=None):
"""
losses_config = self.task_config.loss
if losses_config.one_hot:
- total_loss = tf.keras.losses.categorical_crossentropy(
+ total_loss = tf_keras.losses.categorical_crossentropy(
labels,
model_outputs,
from_logits=True,
label_smoothing=losses_config.label_smoothing)
else:
- total_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ total_loss = tf_keras.losses.sparse_categorical_crossentropy(
labels, model_outputs, from_logits=True)
total_loss = tf_utils.safe_mean(total_loss)
@@ -534,14 +534,14 @@ def build_metrics(self, training=True):
k = self.task_config.evaluation.top_k
if self.task_config.evaluation.one_hot:
metrics = [
- tf.keras.metrics.CategoricalAccuracy(name='accuracy'),
- tf.keras.metrics.TopKCategoricalAccuracy(
+ tf_keras.metrics.CategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.TopKCategoricalAccuracy(
k=k, name='top_{}_accuracy'.format(k))
]
else:
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
- tf.keras.metrics.SparseTopKCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(name='accuracy'),
+ tf_keras.metrics.SparseTopKCategoricalAccuracy(
k=k, name='top_{}_accuracy'.format(k))
]
return metrics
@@ -580,7 +580,7 @@ def train_step(self, inputs, model, optimizer, metrics=None):
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
@@ -590,7 +590,7 @@ def train_step(self, inputs, model, optimizer, metrics=None):
grads = tape.gradient(scaled_loss, tvars)
# Scales back gradient before apply_gradients when LossScaleOptimizer is
# used.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
diff --git a/official/vision/beta/projects/simclr/train.py b/official/projects/simclr/train.py
similarity index 93%
rename from official/vision/beta/projects/simclr/train.py
rename to official/projects/simclr/train.py
index 0eed5ee07c2..e4e3abc5ca9 100644
--- a/official/vision/beta/projects/simclr/train.py
+++ b/official/projects/simclr/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@
from official.core import train_lib
from official.core import train_utils
from official.modeling import performance
-from official.vision.beta.projects.simclr.common import registry_imports # pylint: disable=unused-import
+from official.projects.simclr.common import registry_imports # pylint: disable=unused-import
FLAGS = flags.FLAGS
diff --git a/official/projects/teams/__init__.py b/official/projects/teams/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/teams/__init__.py
+++ b/official/projects/teams/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/teams/teams.py b/official/projects/teams/teams.py
index d2833cfe5de..1ccce999607 100644
--- a/official/projects/teams/teams.py
+++ b/official/projects/teams/teams.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import dataclasses
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.modeling.hyperparams import base_config
@@ -43,8 +43,12 @@ class TeamsPretrainerConfig(base_config.Config):
num_shared_generator_hidden_layers: int = 3
# Number of bottom layers shared between different discriminator tasks.
num_discriminator_task_agnostic_layers: int = 11
- generator: encoders.BertEncoderConfig = encoders.BertEncoderConfig()
- discriminator: encoders.BertEncoderConfig = encoders.BertEncoderConfig()
+ generator: encoders.BertEncoderConfig = dataclasses.field(
+ default_factory=encoders.BertEncoderConfig
+ )
+ discriminator: encoders.BertEncoderConfig = dataclasses.field(
+ default_factory=encoders.BertEncoderConfig
+ )
class TeamsEncoderConfig(encoders.BertEncoderConfig):
@@ -72,7 +76,7 @@ def get_encoder(bert_config: TeamsEncoderConfig,
hidden_size=bert_config.hidden_size,
embedding_width=bert_config.embedding_size,
max_seq_length=bert_config.max_position_embeddings,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
dropout_rate=bert_config.dropout_rate,
)
@@ -83,7 +87,7 @@ def get_encoder(bert_config: TeamsEncoderConfig,
bert_config.hidden_activation),
dropout_rate=bert_config.dropout_rate,
attention_dropout_rate=bert_config.attention_dropout_rate,
- kernel_initializer=tf.keras.initializers.TruncatedNormal(
+ kernel_initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
)
if embedding_network is None:
@@ -97,7 +101,7 @@ def get_encoder(bert_config: TeamsEncoderConfig,
hidden_cfg=hidden_cfg,
num_hidden_instances=bert_config.num_layers,
pooled_output_dim=bert_config.hidden_size,
- pooler_layer_initializer=tf.keras.initializers.TruncatedNormal(
+ pooler_layer_initializer=tf_keras.initializers.TruncatedNormal(
stddev=bert_config.initializer_range),
dict_outputs=True)
diff --git a/official/projects/teams/teams_experiments.py b/official/projects/teams/teams_experiments.py
index 030e1393918..c15aed46778 100644
--- a/official/projects/teams/teams_experiments.py
+++ b/official/projects/teams/teams_experiments.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/teams/teams_pretrainer.py b/official/projects/teams/teams_pretrainer.py
index 8b61be0795e..fba7ebdf090 100644
--- a/official/projects/teams/teams_pretrainer.py
+++ b/official/projects/teams/teams_pretrainer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Trainer network for TEAMS models."""
# pylint: disable=g-classes-have-attributes
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.nlp.modeling import layers
@@ -24,7 +24,7 @@
_LOGIT_PENALTY_MULTIPLIER = 10000
-class ReplacedTokenDetectionHead(tf.keras.layers.Layer):
+class ReplacedTokenDetectionHead(tf_keras.layers.Layer):
"""Replaced token detection discriminator head.
Arguments:
@@ -58,15 +58,16 @@ def __init__(self,
intermediate_activation=self.activation,
dropout_rate=self.hidden_cfg['dropout_rate'],
attention_dropout_rate=self.hidden_cfg['attention_dropout_rate'],
- kernel_initializer=self.initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='transformer/layer_%d_rtd' % i))
- self.dense = tf.keras.layers.Dense(
+ self.dense = tf_keras.layers.Dense(
self.hidden_size,
activation=self.activation,
- kernel_initializer=self.initializer,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='transform/rtd_dense')
- self.rtd_head = tf.keras.layers.Dense(
- units=1, kernel_initializer=self.initializer,
+ self.rtd_head = tf_keras.layers.Dense(
+ units=1,
+ kernel_initializer=tf_utils.clone_initializer(self.initializer),
name='transform/rtd_head')
if output not in ('predictions', 'logits'):
@@ -94,7 +95,7 @@ def call(self, sequence_data, input_mask):
return tf.squeeze(rtd_logits, axis=-1)
-class MultiWordSelectionHead(tf.keras.layers.Layer):
+class MultiWordSelectionHead(tf_keras.layers.Layer):
"""Multi-word selection discriminator head.
Arguments:
@@ -116,15 +117,15 @@ def __init__(self,
super(MultiWordSelectionHead, self).__init__(name=name, **kwargs)
self.embedding_table = embedding_table
self.activation = activation
- self.initializer = tf.keras.initializers.get(initializer)
+ self.initializer = tf_keras.initializers.get(initializer)
self._vocab_size, self.embed_size = self.embedding_table.shape
- self.dense = tf.keras.layers.Dense(
+ self.dense = tf_keras.layers.Dense(
self.embed_size,
activation=self.activation,
kernel_initializer=self.initializer,
name='transform/mws_dense')
- self.layer_norm = tf.keras.layers.LayerNormalization(
+ self.layer_norm = tf_keras.layers.LayerNormalization(
axis=-1, epsilon=1e-12, name='transform/mws_layernorm')
if output not in ('predictions', 'logits'):
@@ -201,8 +202,8 @@ def _gather_indexes(self, sequence_tensor, positions):
return output_tensor
-@tf.keras.utils.register_keras_serializable(package='Text')
-class TeamsPretrainer(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Text')
+class TeamsPretrainer(tf_keras.Model):
"""TEAMS network training model.
This is an implementation of the network structure described in "Training
@@ -298,7 +299,7 @@ def __init__(self,
output=output_type,
name='discriminator_mws')
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
"""TEAMS forward pass.
Args:
diff --git a/official/projects/teams/teams_pretrainer_test.py b/official/projects/teams/teams_pretrainer_test.py
index 9a1fc2029d8..eee290630c7 100644
--- a/official/projects/teams/teams_pretrainer_test.py
+++ b/official/projects/teams/teams_pretrainer_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,19 +14,15 @@
"""Tests for TEAMS pre trainer network."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.modeling import activations
from official.nlp.modeling.networks import encoder_scaffold
from official.nlp.modeling.networks import packed_sequence_embedding
from official.projects.teams import teams_pretrainer
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class TeamsPretrainerTest(keras_parameterized.TestCase):
+class TeamsPretrainerTest(tf.test.TestCase):
# Build a transformer network to use within the TEAMS trainer.
def _get_network(self, vocab_size):
@@ -38,7 +34,7 @@ def _get_network(self, vocab_size):
'hidden_size': hidden_size,
'embedding_width': hidden_size,
'max_seq_length': sequence_length,
- 'initializer': tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ 'initializer': tf_keras.initializers.TruncatedNormal(stddev=0.02),
'dropout_rate': 0.1,
}
embedding_inst = packed_sequence_embedding.PackedSequenceEmbedding(
@@ -55,7 +51,7 @@ def _get_network(self, vocab_size):
'attention_dropout_rate':
0.1,
'kernel_initializer':
- tf.keras.initializers.TruncatedNormal(stddev=0.02),
+ tf_keras.initializers.TruncatedNormal(stddev=0.02),
}
return encoder_scaffold.EncoderScaffold(
num_hidden_instances=2,
@@ -83,12 +79,12 @@ def test_teams_pretrainer(self):
# Create a set of 2-dimensional inputs (the first dimension is implicit).
num_token_predictions = 2
sequence_length = 128
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- lm_positions = tf.keras.Input(
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ lm_positions = tf_keras.Input(
shape=(num_token_predictions,), dtype=tf.int32)
- lm_ids = tf.keras.Input(shape=(num_token_predictions,), dtype=tf.int32)
+ lm_ids = tf_keras.Input(shape=(num_token_predictions,), dtype=tf.int32)
inputs = {
'input_word_ids': word_ids,
'input_mask': mask,
diff --git a/official/projects/teams/teams_task.py b/official/projects/teams/teams_task.py
index c8da8c82743..df57b18c359 100644
--- a/official/projects/teams/teams_task.py
+++ b/official/projects/teams/teams_task.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""TEAMS pretraining task (Joint Masked LM, Replaced Token Detection and )."""
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -30,9 +30,13 @@
@dataclasses.dataclass
class TeamsPretrainTaskConfig(cfg.TaskConfig):
"""The model config."""
- model: teams.TeamsPretrainerConfig = teams.TeamsPretrainerConfig()
- train_data: cfg.DataConfig = cfg.DataConfig()
- validation_data: cfg.DataConfig = cfg.DataConfig()
+ model: teams.TeamsPretrainerConfig = dataclasses.field(
+ default_factory=teams.TeamsPretrainerConfig
+ )
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=cfg.DataConfig
+ )
def _get_generator_hidden_layers(discriminator_network, num_hidden_layers,
@@ -75,7 +79,7 @@ def _build_pretrainer(
candidate_size=config.candidate_size,
mlm_activation=tf_utils.get_activation(
generator_encoder_cfg.hidden_activation),
- mlm_initializer=tf.keras.initializers.TruncatedNormal(
+ mlm_initializer=tf_keras.initializers.TruncatedNormal(
stddev=generator_encoder_cfg.initializer_range))
@@ -95,7 +99,7 @@ def build_losses(self,
metrics = dict([(metric.name, metric) for metric in metrics])
# Generator MLM loss.
- lm_prediction_losses = tf.keras.losses.sparse_categorical_crossentropy(
+ lm_prediction_losses = tf_keras.losses.sparse_categorical_crossentropy(
labels['masked_lm_ids'],
tf.cast(model_outputs['lm_outputs'], tf.float32),
from_logits=True)
@@ -123,7 +127,7 @@ def build_losses(self,
# Discriminator MWS loss.
mws_logits = model_outputs['disc_mws_logits']
mws_labels = model_outputs['disc_mws_label']
- mws_loss = tf.keras.losses.sparse_categorical_crossentropy(
+ mws_loss = tf_keras.losses.sparse_categorical_crossentropy(
mws_labels, mws_logits, from_logits=True)
mws_numerator_loss = tf.reduce_sum(mws_loss * lm_label_weights)
mws_denominator_loss = tf.reduce_sum(lm_label_weights)
@@ -165,15 +169,15 @@ def dummy_data(_):
def build_metrics(self, training=None):
del training
metrics = [
- tf.keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
- tf.keras.metrics.Mean(name='masked_lm_loss'),
- tf.keras.metrics.SparseCategoricalAccuracy(
+ tf_keras.metrics.SparseCategoricalAccuracy(name='masked_lm_accuracy'),
+ tf_keras.metrics.Mean(name='masked_lm_loss'),
+ tf_keras.metrics.SparseCategoricalAccuracy(
name='replaced_token_detection_accuracy'),
- tf.keras.metrics.Mean(name='replaced_token_detection_loss'),
- tf.keras.metrics.SparseCategoricalAccuracy(
+ tf_keras.metrics.Mean(name='replaced_token_detection_loss'),
+ tf_keras.metrics.SparseCategoricalAccuracy(
name='multiword_selection_accuracy'),
- tf.keras.metrics.Mean(name='multiword_selection_loss'),
- tf.keras.metrics.Mean(name='total_loss'),
+ tf_keras.metrics.Mean(name='multiword_selection_loss'),
+ tf_keras.metrics.Mean(name='total_loss'),
]
return metrics
@@ -199,8 +203,8 @@ def process_metrics(self, metrics, labels, model_outputs):
model_outputs['disc_mws_label'], model_outputs['disc_mws_logits'],
labels['masked_lm_weights'])
- def train_step(self, inputs, model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer, metrics):
+ def train_step(self, inputs, model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer, metrics):
"""Does forward and backward.
Args:
@@ -229,7 +233,7 @@ def train_step(self, inputs, model: tf.keras.Model,
self.process_metrics(metrics, inputs, outputs)
return {self.loss: loss}
- def validation_step(self, inputs, model: tf.keras.Model, metrics):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics):
"""Validatation step.
Args:
diff --git a/official/projects/teams/teams_task_test.py b/official/projects/teams/teams_task_test.py
index df3c93a0f92..d2a307c6d3d 100644
--- a/official/projects/teams/teams_task_test.py
+++ b/official/projects/teams/teams_task_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for teams_task."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import encoders
from official.nlp.data import pretrain_dataloader
@@ -48,7 +48,7 @@ def test_task(self, num_shared_hidden_layers,
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
diff --git a/official/projects/teams/train.py b/official/projects/teams/train.py
index b13afe537e5..511c90f3339 100644
--- a/official/projects/teams/train.py
+++ b/official/projects/teams/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/text_classification_example/classification_data_loader.py b/official/projects/text_classification_example/classification_data_loader.py
index fa142bc03b5..64d7b2ef885 100644
--- a/official/projects/text_classification_example/classification_data_loader.py
+++ b/official/projects/text_classification_example/classification_data_loader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Dict, Mapping, Optional, Tuple
import dataclasses
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import input_reader
diff --git a/official/projects/text_classification_example/classification_example.py b/official/projects/text_classification_example/classification_example.py
index b8600a4e043..f25c2aa5661 100644
--- a/official/projects/text_classification_example/classification_example.py
+++ b/official/projects/text_classification_example/classification_example.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import dataclasses
from typing import List, Mapping, Text
from seqeval import metrics as seqeval_metrics
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -34,7 +34,9 @@
@dataclasses.dataclass
class ModelConfig(base_config.Config):
"""A base span labeler configuration."""
- encoder: encoders.EncoderConfig = encoders.EncoderConfig()
+ encoder: encoders.EncoderConfig = dataclasses.field(
+ default_factory=encoders.EncoderConfig
+ )
head_dropout: float = 0.1
head_initializer_range: float = 0.02
@@ -45,20 +47,22 @@ class ClassificationExampleConfig(cfg.TaskConfig):
# At most one of `init_checkpoint` and `hub_module_url` can be specified.
init_checkpoint: str = ''
hub_module_url: str = ''
- model: ModelConfig = ModelConfig()
+ model: ModelConfig = dataclasses.field(default_factory=ModelConfig)
num_classes = 2
class_names = ['A', 'B']
- train_data: cfg.DataConfig = classification_data_loader.ClassificationExampleDataConfig(
+ train_data: cfg.DataConfig = dataclasses.field(
+ default_factory=classification_data_loader.ClassificationExampleDataConfig
)
- validation_data: cfg.DataConfig = classification_data_loader.ClassificationExampleDataConfig(
+ validation_data: cfg.DataConfig = dataclasses.field(
+ default_factory=classification_data_loader.ClassificationExampleDataConfig
)
class ClassificationExampleTask(base_task.Task):
"""Task object for classification."""
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
if self.task_config.hub_module_url and self.task_config.init_checkpoint:
raise ValueError('At most one of `hub_module_url` and '
'`init_checkpoint` can be specified.')
@@ -71,12 +75,12 @@ def build_model(self) -> tf.keras.Model:
return models.BertClassifier(
network=encoder_network,
num_classes=len(self.task_config.class_names),
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=self.task_config.model.head_initializer_range),
dropout_rate=self.task_config.model.head_dropout)
def build_losses(self, labels, model_outputs, aux_losses=None) -> tf.Tensor:
- loss = tf.keras.losses.sparse_categorical_crossentropy(
+ loss = tf_keras.losses.sparse_categorical_crossentropy(
labels, tf.cast(model_outputs, tf.float32), from_logits=True)
return tf_utils.safe_mean(loss)
@@ -88,7 +92,7 @@ def build_inputs(self,
return loader.load(input_context)
def inference_step(self, inputs,
- model: tf.keras.Model) -> Mapping[str, tf.Tensor]:
+ model: tf_keras.Model) -> Mapping[str, tf.Tensor]:
"""Performs the forward step."""
logits = model(inputs, training=False)
return {
@@ -98,7 +102,7 @@ def inference_step(self, inputs,
def validation_step(self,
inputs,
- model: tf.keras.Model,
+ model: tf_keras.Model,
metrics=None) -> Mapping[str, tf.Tensor]:
"""Validatation step.
@@ -142,8 +146,8 @@ def id_to_class_name(batched_ids):
# Convert id to class names, because `seqeval_metrics` relies on the class
# name to decide IOB tags.
- state['predict_class'].extend(id_to_class_name(step_outputs['predict_ids']))
- state['label_class'].extend(id_to_class_name(step_outputs['label_ids']))
+ state['predict_class'].extend(id_to_class_name(step_outputs['predict_ids'])) # pyrefly: ignore[unsupported-operation]
+ state['label_class'].extend(id_to_class_name(step_outputs['label_ids'])) # pyrefly: ignore[unsupported-operation]
return state
def reduce_aggregated_logs(self,
diff --git a/official/projects/text_classification_example/classification_example_test.py b/official/projects/text_classification_example/classification_example_test.py
index 4de434f531d..bb68e63dc5a 100644
--- a/official/projects/text_classification_example/classification_example_test.py
+++ b/official/projects/text_classification_example/classification_example_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for nlp.projects.example.classification_example."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.nlp.configs import encoders
@@ -60,7 +60,7 @@ def test_task_with_dummy_data(self):
dataset = task.build_inputs(train_data_config)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.initialize(model)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
diff --git a/official/projects/text_classification_example/train.py b/official/projects/text_classification_example/train.py
index c2e8e16558c..9911f15d946 100644
--- a/official/projects/text_classification_example/train.py
+++ b/official/projects/text_classification_example/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/token_dropping/encoder.py b/official/projects/token_dropping/encoder.py
index bb8adcef458..890c3817d0b 100644
--- a/official/projects/token_dropping/encoder.py
+++ b/official/projects/token_dropping/encoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,18 +17,19 @@
from typing import Any, Callable, Optional, Union, Tuple
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
from official.nlp.modeling import layers
-_Initializer = Union[str, tf.keras.initializers.Initializer]
+_Initializer = Union[str, tf_keras.initializers.Initializer]
_Activation = Union[str, Callable[..., Any]]
-_approx_gelu = lambda x: tf.keras.activations.gelu(x, approximate=True)
+_approx_gelu = lambda x: tf_keras.activations.gelu(x, approximate=True)
-class TokenDropBertEncoder(tf.keras.layers.Layer):
+class TokenDropBertEncoder(tf_keras.layers.Layer):
"""Bi-directional Transformer-based encoder network with token dropping.
During pretraining, we drop unimportant tokens starting from an intermediate
@@ -64,9 +65,9 @@ class TokenDropBertEncoder(tf.keras.layers.Layer):
token_keep_k: The number of tokens you want to keep in the intermediate
layers. The rest will be dropped in those layers.
token_allow_list: The list of token-ids that should not be droped. In the
- BERT English vocab, token-id from 1 to 998 contains special tokens such
- as [CLS], [SEP]. By default, token_allow_list contains all of these
- special tokens.
+ BERT English vocab, token-id from 1 to 998 contains special tokens such as
+ [CLS], [SEP]. By default, token_allow_list contains all of these special
+ tokens.
token_deny_list: The list of token-ids that should always be droped. In the
BERT English vocab, token-id=0 means [PAD]. By default, token_deny_list
contains and only contains [PAD].
@@ -105,11 +106,11 @@ def __init__(
token_keep_k: int = 256,
token_allow_list: Tuple[int, ...] = (100, 101, 102, 103),
token_deny_list: Tuple[int, ...] = (0,),
- initializer: _Initializer = tf.keras.initializers.TruncatedNormal(
+ initializer: _Initializer = tf_keras.initializers.TruncatedNormal(
stddev=0.02),
output_range: Optional[int] = None,
embedding_width: Optional[int] = None,
- embedding_layer: Optional[tf.keras.layers.Layer] = None,
+ embedding_layer: Optional[tf_keras.layers.Layer] = None,
norm_first: bool = False,
with_dense_inputs: bool = False,
**kwargs):
@@ -128,8 +129,12 @@ def __init__(
attention_dropout = kwargs.pop('attention_dropout_rate')
super().__init__(**kwargs)
- activation = tf.keras.activations.get(inner_activation)
- initializer = tf.keras.initializers.get(initializer)
+ if output_range is not None:
+ logging.warning('`output_range` is available as an argument for `call()`.'
+ 'The `output_range` as __init__ argument is deprecated.')
+
+ activation = tf_keras.activations.get(inner_activation)
+ initializer = tf_keras.initializers.get(initializer)
if embedding_width is None:
embedding_width = hidden_size
@@ -138,38 +143,38 @@ def __init__(
self._embedding_layer = layers.OnDeviceEmbedding(
vocab_size=vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
name='word_embeddings')
else:
self._embedding_layer = embedding_layer
self._position_embedding_layer = layers.PositionEmbedding(
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
max_length=max_sequence_length,
name='position_embedding')
self._type_embedding_layer = layers.OnDeviceEmbedding(
vocab_size=type_vocab_size,
embedding_width=embedding_width,
- initializer=initializer,
+ initializer=tf_utils.clone_initializer(initializer),
use_one_hot=True,
name='type_embeddings')
- self._embedding_norm_layer = tf.keras.layers.LayerNormalization(
+ self._embedding_norm_layer = tf_keras.layers.LayerNormalization(
name='embeddings/layer_norm', axis=-1, epsilon=1e-12, dtype=tf.float32)
- self._embedding_dropout = tf.keras.layers.Dropout(
+ self._embedding_dropout = tf_keras.layers.Dropout(
rate=output_dropout, name='embedding_dropout')
# We project the 'embedding' output to 'hidden_size' if it is not already
# 'hidden_size'.
self._embedding_projection = None
if embedding_width != hidden_size:
- self._embedding_projection = tf.keras.layers.experimental.EinsumDense(
+ self._embedding_projection = tf_keras.layers.EinsumDense(
'...x,xy->...y',
output_shape=hidden_size,
bias_axes='y',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='embedding_projection')
# The first 999 tokens are special tokens such as [PAD], [CLS], [SEP].
@@ -203,15 +208,14 @@ def __init__(
output_dropout=output_dropout,
attention_dropout=attention_dropout,
norm_first=norm_first,
- output_range=output_range if i == num_layers - 1 else None,
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='transformer/layer_%d' % i)
self._transformer_layers.append(layer)
- self._pooler_layer = tf.keras.layers.Dense(
+ self._pooler_layer = tf_keras.layers.Dense(
units=hidden_size,
activation='tanh',
- kernel_initializer=initializer,
+ kernel_initializer=tf_utils.clone_initializer(initializer),
name='pooler_transform')
self._config = {
@@ -222,7 +226,7 @@ def __init__(
'max_sequence_length': max_sequence_length,
'type_vocab_size': type_vocab_size,
'inner_dim': inner_dim,
- 'inner_activation': tf.keras.activations.serialize(activation),
+ 'inner_activation': tf_keras.activations.serialize(activation),
'output_dropout': output_dropout,
'attention_dropout': attention_dropout,
'token_loss_init_value': token_loss_init_value,
@@ -230,7 +234,7 @@ def __init__(
'token_keep_k': token_keep_k,
'token_allow_list': token_allow_list,
'token_deny_list': token_deny_list,
- 'initializer': tf.keras.initializers.serialize(initializer),
+ 'initializer': tf_keras.initializers.serialize(initializer),
'output_range': output_range,
'embedding_width': embedding_width,
'embedding_layer': embedding_layer,
@@ -239,21 +243,21 @@ def __init__(
}
if with_dense_inputs:
self.inputs = dict(
- input_word_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- dense_inputs=tf.keras.Input(
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_inputs=tf_keras.Input(
shape=(None, embedding_width), dtype=tf.float32),
- dense_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- dense_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
+ dense_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ dense_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
)
else:
self.inputs = dict(
- input_word_ids=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_mask=tf.keras.Input(shape=(None,), dtype=tf.int32),
- input_type_ids=tf.keras.Input(shape=(None,), dtype=tf.int32))
+ input_word_ids=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_mask=tf_keras.Input(shape=(None,), dtype=tf.int32),
+ input_type_ids=tf_keras.Input(shape=(None,), dtype=tf.int32))
- def call(self, inputs):
+ def call(self, inputs, output_range: Optional[tf.Tensor] = None):
if isinstance(inputs, dict):
word_ids = inputs.get('input_word_ids')
mask = inputs.get('input_mask')
@@ -302,8 +306,11 @@ def call(self, inputs):
# 4. Finally, all tokens go through the last layer.
# Step 1.
- for layer in self._transformer_layers[:self._num_layers // 2 - 1]:
- x = layer([x, attention_mask])
+ for i, layer in enumerate(self._transformer_layers[:self._num_layers // 2 -
+ 1]):
+ x = layer([x, attention_mask],
+ output_range=output_range if i == self._num_layers -
+ 1 else None)
encoder_outputs.append(x)
# Step 2.
@@ -321,12 +328,17 @@ def call(self, inputs):
# Then, call transformer layer with cross attention.
x_selected = self._transformer_layers[self._num_layers // 2 - 1](
- [x_selected, x_all, attention_mask_token_pass])
+ [x_selected, x_all, attention_mask_token_pass],
+ output_range=output_range if self._num_layers // 2 -
+ 1 == self._num_layers - 1 else None)
encoder_outputs.append(x_selected)
# Step 3.
- for layer in self._transformer_layers[self._num_layers // 2:-1]:
- x_selected = layer([x_selected, attention_mask_token_drop])
+ for i, layer in enumerate(self._transformer_layers[self._num_layers //
+ 2:-1]):
+ x_selected = layer([x_selected, attention_mask_token_drop],
+ output_range=output_range if i == self._num_layers - 1
+ else None)
encoder_outputs.append(x_selected)
# Step 4.
@@ -338,7 +350,8 @@ def call(self, inputs):
x = tf.gather(x, reverse_indices, batch_dims=1, axis=1)
# Then, call transformer layer with all tokens.
- x = self._transformer_layers[-1]([x, attention_mask])
+ x = self._transformer_layers[-1]([x, attention_mask],
+ output_range=output_range)
encoder_outputs.append(x)
last_encoder_output = encoder_outputs[-1]
@@ -385,4 +398,3 @@ def from_config(cls, config, custom_objects=None):
logging.warn(warn_string)
return cls(**config)
-
diff --git a/official/projects/token_dropping/encoder_config.py b/official/projects/token_dropping/encoder_config.py
index b7809d46f81..2566c84c4b5 100644
--- a/official/projects/token_dropping/encoder_config.py
+++ b/official/projects/token_dropping/encoder_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Token dropping encoder configuration and instantiation."""
import dataclasses
from typing import Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.modeling.hyperparams import base_config
@@ -53,7 +53,7 @@ def get_encoder(encoder_cfg: TokenDropBertEncoderConfig):
attention_dropout_rate=encoder_cfg.attention_dropout_rate,
max_sequence_length=encoder_cfg.max_position_embeddings,
type_vocab_size=encoder_cfg.type_vocab_size,
- initializer=tf.keras.initializers.TruncatedNormal(
+ initializer=tf_keras.initializers.TruncatedNormal(
stddev=encoder_cfg.initializer_range),
output_range=encoder_cfg.output_range,
embedding_width=encoder_cfg.embedding_size,
diff --git a/official/projects/token_dropping/encoder_test.py b/official/projects/token_dropping/encoder_test.py
index d880c970f93..861e1ca2660 100644
--- a/official/projects/token_dropping/encoder_test.py
+++ b/official/projects/token_dropping/encoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,24 +14,19 @@
"""Tests for transformer-based bert encoder network."""
-# Import libraries
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from tensorflow.python.keras import keras_parameterized # pylint: disable=g-direct-tensorflow-import
from official.nlp.modeling.networks import bert_encoder
from official.projects.token_dropping import encoder
-# This decorator runs the test in V1, V2-Eager, and V2-Functional mode. It
-# guarantees forward compatibility of this code for the V2 switchover.
-@keras_parameterized.run_all_keras_modes
-class TokenDropBertEncoderTest(keras_parameterized.TestCase):
+class TokenDropBertEncoderTest(tf.test.TestCase, parameterized.TestCase):
def tearDown(self):
super(TokenDropBertEncoderTest, self).tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
def test_dict_outputs_network_creation(self):
hidden_size = 32
@@ -46,9 +41,9 @@ def test_dict_outputs_network_creation(self):
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
@@ -56,7 +51,7 @@ def test_dict_outputs_network_creation(self):
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, 3)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
expected_data_shape = [None, sequence_length, hidden_size]
expected_pooled_shape = [None, hidden_size]
@@ -81,9 +76,9 @@ def test_dict_outputs_all_encoder_outputs_network_creation(self):
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
all_encoder_outputs = dict_outputs["encoder_outputs"]
@@ -103,7 +98,7 @@ def test_dict_outputs_all_encoder_outputs_network_creation(self):
def test_dict_outputs_network_creation_with_float16_dtype(self):
hidden_size = 32
sequence_length = 21
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a small BertEncoder for testing.
test_network = encoder.TokenDropBertEncoder(
vocab_size=100,
@@ -115,9 +110,9 @@ def test_dict_outputs_network_creation_with_float16_dtype(self):
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
@@ -150,22 +145,22 @@ def test_dict_outputs_network_invocation(
num_attention_heads=2,
num_layers=3,
type_vocab_size=num_types,
- output_range=output_range,
dict_outputs=True,
token_keep_k=2,
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
- dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids),
+ output_range=output_range)
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -195,7 +190,7 @@ def test_dict_outputs_network_invocation(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[1], sequence_length)
@@ -216,7 +211,7 @@ def test_dict_outputs_network_invocation(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[-1], hidden_size)
self.assertTrue(hasattr(test_network, "_embedding_projection"))
@@ -234,9 +229,9 @@ def test_network_creation(self):
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
@@ -244,7 +239,7 @@ def test_network_creation(self):
self.assertIsInstance(test_network.transformer_layers, list)
self.assertLen(test_network.transformer_layers, 3)
- self.assertIsInstance(test_network.pooler_layer, tf.keras.layers.Dense)
+ self.assertIsInstance(test_network.pooler_layer, tf_keras.layers.Dense)
expected_data_shape = [None, sequence_length, hidden_size]
expected_pooled_shape = [None, hidden_size]
@@ -282,9 +277,9 @@ def test_all_encoder_outputs_network_creation(self):
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
all_encoder_outputs = dict_outputs["encoder_outputs"]
@@ -304,7 +299,7 @@ def test_all_encoder_outputs_network_creation(self):
def test_network_creation_with_float16_dtype(self):
hidden_size = 32
sequence_length = 21
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
# Create a small BertEncoder for testing.
test_network = encoder.TokenDropBertEncoder(
vocab_size=100,
@@ -315,9 +310,9 @@ def test_network_creation_with_float16_dtype(self):
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
@@ -349,21 +344,21 @@ def test_network_invocation(self, output_range, out_seq_len):
num_attention_heads=2,
num_layers=3,
type_vocab_size=num_types,
- output_range=output_range,
token_keep_k=2,
token_allow_list=(),
token_deny_list=())
# Create the inputs (note that the first dimension is implicit).
- word_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- mask = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
- type_ids = tf.keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ word_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ mask = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
+ type_ids = tf_keras.Input(shape=(sequence_length,), dtype=tf.int32)
dict_outputs = test_network(
- dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
+ dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids),
+ output_range=output_range)
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
# Create a model based off of this network:
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
# Invoke the model. We can't validate the output data here (the model is too
# complex) but this will catch structural runtime errors.
@@ -392,7 +387,7 @@ def test_network_invocation(self, output_range, out_seq_len):
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[1], sequence_length)
@@ -412,7 +407,7 @@ def test_network_invocation(self, output_range, out_seq_len):
dict(input_word_ids=word_ids, input_mask=mask, input_type_ids=type_ids))
data = dict_outputs["sequence_output"]
pooled = dict_outputs["pooled_output"]
- model = tf.keras.Model([word_ids, mask, type_ids], [data, pooled])
+ model = tf_keras.Model([word_ids, mask, type_ids], [data, pooled])
outputs = model.predict([word_id_data, mask_data, type_id_data])
self.assertEqual(outputs[0].shape[-1], hidden_size)
self.assertTrue(hasattr(test_network, "_embedding_projection"))
@@ -422,7 +417,7 @@ class TokenDropCompatibilityTest(tf.test.TestCase):
def tearDown(self):
super().tearDown()
- tf.keras.mixed_precision.set_global_policy("float32")
+ tf_keras.mixed_precision.set_global_policy("float32")
def test_checkpoint_forward_compatible(self):
batch_size = 3
@@ -498,7 +493,7 @@ def test_keras_model_checkpoint_forward_compatible(self):
old_net = bert_encoder.BertEncoderV2(**kwargs)
inputs = old_net.inputs
outputs = old_net(inputs)
- old_model = tf.keras.Model(inputs=inputs, outputs=outputs)
+ old_model = tf_keras.Model(inputs=inputs, outputs=outputs)
old_model_outputs = old_model(data)
ckpt = tf.train.Checkpoint(net=old_model)
path = ckpt.save(self.get_temp_dir())
@@ -509,7 +504,7 @@ def test_keras_model_checkpoint_forward_compatible(self):
**kwargs)
inputs = new_net.inputs
outputs = new_net(inputs)
- new_model = tf.keras.Model(inputs=inputs, outputs=outputs)
+ new_model = tf_keras.Model(inputs=inputs, outputs=outputs)
new_ckpt = tf.train.Checkpoint(net=new_model)
new_ckpt.restore(path)
diff --git a/official/projects/token_dropping/experiment_configs.py b/official/projects/token_dropping/experiment_configs.py
index 3f2fd6a85b7..bae3a0f88bd 100644
--- a/official/projects/token_dropping/experiment_configs.py
+++ b/official/projects/token_dropping/experiment_configs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/token_dropping/masked_lm.py b/official/projects/token_dropping/masked_lm.py
index f159a216d37..f159d345a9f 100644
--- a/official/projects/token_dropping/masked_lm.py
+++ b/official/projects/token_dropping/masked_lm.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import dataclasses
from typing import Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import task_factory
from official.nlp.tasks import masked_lm
@@ -40,7 +40,7 @@ def build_losses(self,
"""Return the final loss, and the masked-lm loss."""
with tf.name_scope('MaskedLMTask/losses'):
metrics = dict([(metric.name, metric) for metric in metrics])
- lm_prediction_losses = tf.keras.losses.sparse_categorical_crossentropy(
+ lm_prediction_losses = tf_keras.losses.sparse_categorical_crossentropy(
labels['masked_lm_ids'],
tf.cast(model_outputs['mlm_logits'], tf.float32),
from_logits=True)
@@ -55,7 +55,7 @@ def build_losses(self,
sentence_outputs = tf.cast(
model_outputs['next_sentence'], dtype=tf.float32)
sentence_loss = tf.reduce_mean(
- tf.keras.losses.sparse_categorical_crossentropy(
+ tf_keras.losses.sparse_categorical_crossentropy(
sentence_labels, sentence_outputs, from_logits=True))
metrics['next_sentence_loss'].update_state(sentence_loss)
total_loss = mlm_loss + sentence_loss
@@ -66,8 +66,8 @@ def build_losses(self,
total_loss += tf.add_n(aux_losses)
return total_loss, lm_prediction_losses
- def train_step(self, inputs, model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer, metrics):
+ def train_step(self, inputs, model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer, metrics):
"""Does forward and backward.
Args:
@@ -96,14 +96,14 @@ def train_step(self, inputs, model: tf.keras.Model,
scaled_loss = loss / tf.distribute.get_strategy().num_replicas_in_sync
tvars = model.trainable_variables
if self.task_config.scale_loss:
- grads = tape.gradient(scaled_loss, tvars)
+ grads = tape.gradient(scaled_loss, tvars) # pyrefly: ignore[unbound-name]
else:
grads = tape.gradient(loss, tvars)
optimizer.apply_gradients(list(zip(grads, tvars)))
self.process_metrics(metrics, inputs, outputs)
return {self.loss: loss}
- def validation_step(self, inputs, model: tf.keras.Model, metrics):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics):
"""Validatation step.
Args:
diff --git a/official/projects/token_dropping/masked_lm_test.py b/official/projects/token_dropping/masked_lm_test.py
index 2c0ea5af948..7066bed650b 100644
--- a/official/projects/token_dropping/masked_lm_test.py
+++ b/official/projects/token_dropping/masked_lm_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for official.nlp.tasks.masked_lm."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.nlp.configs import bert
from official.nlp.configs import encoders
@@ -49,7 +49,7 @@ def test_task(self):
dataset = task.build_inputs(config.train_data)
iterator = iter(dataset)
- optimizer = tf.keras.optimizers.SGD(lr=0.1)
+ optimizer = tf_keras.optimizers.SGD(lr=0.1)
task.train_step(next(iterator), model, optimizer, metrics=metrics)
task.validation_step(next(iterator), model, metrics=metrics)
diff --git a/official/projects/token_dropping/train.py b/official/projects/token_dropping/train.py
index e84d45f7724..c8571753c2d 100644
--- a/official/projects/token_dropping/train.py
+++ b/official/projects/token_dropping/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/token_dropping/wiki_books_pretrain.yaml b/official/projects/token_dropping/wiki_books_pretrain.yaml
index f585ca35357..f90720feed7 100644
--- a/official/projects/token_dropping/wiki_books_pretrain.yaml
+++ b/official/projects/token_dropping/wiki_books_pretrain.yaml
@@ -5,7 +5,8 @@ task:
train_data:
drop_remainder: true
global_batch_size: 512
- input_path: /path-to-data/wikipedia.tfrecord*,/path-to-data/books.tfrecord*
+ # Use glob pattern to match all shards except 00141-of-00500 which is reserved for validation.
+ input_path: /path-to-data/wikipedia.tfrecord-*-of-00500,/path-to-data/books.tfrecord-*-of-00500
is_training: true
max_predictions_per_seq: 76
seq_length: 512
@@ -15,7 +16,7 @@ task:
validation_data:
drop_remainder: false
global_batch_size: 512
- input_path: /path-to-data/wikipedia.tfrecord*,/path-to-data/books.tfrecord*
+ input_path: /path-to-data/wikipedia.tfrecord-00141-of-00500-eval,/path-to-data/books.tfrecord-00141-of-00500-eval
is_training: false
max_predictions_per_seq: 76
seq_length: 512
diff --git a/official/projects/triviaqa/__init__.py b/official/projects/triviaqa/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/triviaqa/__init__.py
+++ b/official/projects/triviaqa/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/triviaqa/dataset.py b/official/projects/triviaqa/dataset.py
index 706bbb6779b..f91988ddd65 100644
--- a/official/projects/triviaqa/dataset.py
+++ b/official/projects/triviaqa/dataset.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl import logging
import apache_beam as beam
import six
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets.public_api as tfds
from official.projects.triviaqa import preprocess
@@ -39,7 +39,7 @@
}
"""
_DOWNLOAD_URL_TMPL = (
- "http://nlp.cs.washington.edu/triviaqa/data/triviaqa-{}.tar.gz")
+ "https://nlp.cs.washington.edu/triviaqa/data/triviaqa-{}.tar.gz")
_TRAIN_FILE_FORMAT = "*-train.json"
_VALIDATION_FILE_FORMAT = "*-dev.json"
_TEST_FILE_FORMAT = "*test-without-answers.json"
@@ -198,7 +198,7 @@ def _info(self):
builder=self,
description=_DESCRIPTION,
supervised_keys=None,
- homepage="http://nlp.cs.washington.edu/triviaqa/",
+ homepage="https://nlp.cs.washington.edu/triviaqa/",
citation=_CITATION,
features=tfds.features.FeaturesDict({
"id": tfds.features.Text(),
@@ -274,7 +274,7 @@ def _info(self):
}),
supervised_keys=None,
- homepage="http://nlp.cs.washington.edu/triviaqa/",
+ homepage="https://nlp.cs.washington.edu/triviaqa/",
citation=_CITATION,
)
@@ -282,9 +282,9 @@ def _split_generators(self, dl_manager):
"""Returns SplitGenerators."""
cfg = self.builder_config
download_urls = dict()
- if not (cfg.unfiltered and cfg.exclude_context):
+ if not (cfg.unfiltered and cfg.exclude_context): # pyrefly: ignore[missing-attribute]
download_urls["rc"] = _DOWNLOAD_URL_TMPL.format("rc")
- if cfg.unfiltered:
+ if cfg.unfiltered: # pyrefly: ignore[missing-attribute]
download_urls["unfiltered"] = _DOWNLOAD_URL_TMPL.format("unfiltered")
file_paths = dl_manager.download_and_extract(download_urls)
@@ -297,7 +297,7 @@ def _split_generators(self, dl_manager):
os.path.join(qa_dir, _VALIDATION_FILE_FORMAT))
test_files = tf.io.gfile.glob(os.path.join(qa_dir, _TEST_FILE_FORMAT))
- if cfg.exclude_context:
+ if cfg.exclude_context: # pyrefly: ignore[missing-attribute]
web_evidence_dir = None
wiki_evidence_dir = None
else:
@@ -311,19 +311,19 @@ def _split_generators(self, dl_manager):
return [
tfds.core.SplitGenerator(
- name=tfds.Split.TRAIN,
+ name=tfds.Split.TRAIN, # pyrefly: ignore[missing-attribute]
gen_kwargs={"files": train_files,
"web_dir": web_evidence_dir,
"wiki_dir": wiki_evidence_dir,
"answer": True}),
tfds.core.SplitGenerator(
- name=tfds.Split.VALIDATION,
+ name=tfds.Split.VALIDATION, # pyrefly: ignore[missing-attribute]
gen_kwargs={"files": valid_files,
"web_dir": web_evidence_dir,
"wiki_dir": wiki_evidence_dir,
"answer": True}),
tfds.core.SplitGenerator(
- name=tfds.Split.TEST,
+ name=tfds.Split.TEST, # pyrefly: ignore[missing-attribute]
gen_kwargs={"files": test_files,
"web_dir": web_evidence_dir,
"wiki_dir": wiki_evidence_dir,
@@ -346,7 +346,7 @@ def _build_pcollection(self, pipeline, files, web_dir, wiki_dir, answer):
web_dir=web_dir)
parse_example_fn = functools.partial(parse_example,
- self.builder_config.exclude_context,
+ self.builder_config.exclude_context, # pyrefly: ignore[missing-attribute]
web_dir, wiki_dir)
return (pipeline
| beam.Create(files)
diff --git a/official/projects/triviaqa/download_and_prepare.py b/official/projects/triviaqa/download_and_prepare.py
index 1a3140c3dd8..6897283df85 100644
--- a/official/projects/triviaqa/download_and_prepare.py
+++ b/official/projects/triviaqa/download_and_prepare.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/triviaqa/evaluate.py b/official/projects/triviaqa/evaluate.py
index 6d19c58e606..0a777a3b2ae 100644
--- a/official/projects/triviaqa/evaluate.py
+++ b/official/projects/triviaqa/evaluate.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.triviaqa import evaluation
diff --git a/official/projects/triviaqa/evaluation.py b/official/projects/triviaqa/evaluation.py
index 80218cab90b..3d1f6fe482d 100644
--- a/official/projects/triviaqa/evaluation.py
+++ b/official/projects/triviaqa/evaluation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/triviaqa/inputs.py b/official/projects/triviaqa/inputs.py
index 426a1154110..887ad338db4 100644
--- a/official/projects/triviaqa/inputs.py
+++ b/official/projects/triviaqa/inputs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import os
from typing import Optional, Text, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
from official.modeling import tf_utils
diff --git a/official/projects/triviaqa/modeling.py b/official/projects/triviaqa/modeling.py
index 4df0f1b2b01..3d0d50b812c 100644
--- a/official/projects/triviaqa/modeling.py
+++ b/official/projects/triviaqa/modeling.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,13 +13,13 @@
# limitations under the License.
"""Modeling for TriviaQA."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.nlp.configs import encoders
-class TriviaQaHead(tf.keras.layers.Layer):
+class TriviaQaHead(tf_keras.layers.Layer):
"""Computes logits given token and global embeddings."""
def __init__(self,
@@ -29,17 +29,17 @@ def __init__(self,
attention_dropout_rate=0.0,
**kwargs):
super(TriviaQaHead, self).__init__(**kwargs)
- self._attention_dropout = tf.keras.layers.Dropout(attention_dropout_rate)
- self._intermediate_dense = tf.keras.layers.Dense(intermediate_size)
- self._intermediate_activation = tf.keras.layers.Activation(
+ self._attention_dropout = tf_keras.layers.Dropout(attention_dropout_rate)
+ self._intermediate_dense = tf_keras.layers.Dense(intermediate_size)
+ self._intermediate_activation = tf_keras.layers.Activation(
intermediate_activation)
- self._output_dropout = tf.keras.layers.Dropout(dropout_rate)
- self._output_layer_norm = tf.keras.layers.LayerNormalization()
- self._logits_dense = tf.keras.layers.Dense(2)
+ self._output_dropout = tf_keras.layers.Dropout(dropout_rate)
+ self._output_layer_norm = tf_keras.layers.LayerNormalization()
+ self._logits_dense = tf_keras.layers.Dense(2)
def build(self, input_shape):
output_shape = input_shape['token_embeddings'][-1]
- self._output_dense = tf.keras.layers.Dense(output_shape)
+ self._output_dense = tf_keras.layers.Dense(output_shape)
super(TriviaQaHead, self).build(input_shape)
def call(self, inputs, training=None):
@@ -59,14 +59,14 @@ def call(self, inputs, training=None):
return logits
-class TriviaQaModel(tf.keras.Model):
+class TriviaQaModel(tf_keras.Model):
"""Model for TriviaQA."""
def __init__(self, model_config: encoders.EncoderConfig, sequence_length: int,
**kwargs):
inputs = dict(
- token_ids=tf.keras.Input((sequence_length,), dtype=tf.int32),
- question_lengths=tf.keras.Input((), dtype=tf.int32))
+ token_ids=tf_keras.Input((sequence_length,), dtype=tf.int32),
+ question_lengths=tf_keras.Input((), dtype=tf.int32))
encoder = encoders.build_encoder(model_config)
x = encoder(
dict(
@@ -91,7 +91,7 @@ def encoder(self):
return self._encoder
-class SpanOrCrossEntropyLoss(tf.keras.losses.Loss):
+class SpanOrCrossEntropyLoss(tf_keras.losses.Loss):
"""Cross entropy loss for multiple correct answers.
See https://arxiv.org/abs/1710.10723.
diff --git a/official/projects/triviaqa/predict.py b/official/projects/triviaqa/predict.py
index 16ccdb83fae..bf5f59f5d2c 100644
--- a/official/projects/triviaqa/predict.py
+++ b/official/projects/triviaqa/predict.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,7 +22,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
import sentencepiece as spm
@@ -34,7 +34,7 @@
flags.DEFINE_string('data_dir', None, 'TensorFlow Datasets directory.')
flags.DEFINE_enum('split', None,
- [tfds.Split.TRAIN, tfds.Split.VALIDATION, tfds.Split.TEST],
+ [tfds.Split.TRAIN, tfds.Split.VALIDATION, tfds.Split.TEST], # pyrefly: ignore[missing-attribute]
'For which split to generate predictions.')
flags.DEFINE_string('predictions_path', None, 'Output for predictions.')
@@ -155,7 +155,7 @@ def main(argv):
FLAGS.data_dir, FLAGS.split, FLAGS.batch_size, include_answers=False)
# Initialize model and compile.
with strategy.scope():
- model = tf.keras.models.load_model(FLAGS.saved_model_dir, compile=False)
+ model = tf_keras.models.load_model(FLAGS.saved_model_dir, compile=False)
logging.info('Model initialized. Beginning prediction loop.')
logits_fn = tf.function(
functools.partial(prediction.distributed_logits_fn, model))
diff --git a/official/projects/triviaqa/prediction.py b/official/projects/triviaqa/prediction.py
index f2c96954fab..948a19a3207 100644
--- a/official/projects/triviaqa/prediction.py
+++ b/official/projects/triviaqa/prediction.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Functions for inference."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def split_and_pad(strategy, batch_size, x):
diff --git a/official/projects/triviaqa/preprocess.py b/official/projects/triviaqa/preprocess.py
index fb16ef8a058..187dd6336e9 100644
--- a/official/projects/triviaqa/preprocess.py
+++ b/official/projects/triviaqa/preprocess.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -383,7 +383,7 @@ def _handle_exceptional_examples(
spans = [span]
span = AnswerSpan(i, i + pattern.find(b'A') + 1, 'Vitamin A')
span = realign_answer_span(features, None, processor, span)
- return spans + [span]
+ return spans + [span] # pyrefly: ignore[bad-return]
if features.id == 'odql_292--Colombia.txt':
pattern = b'Colombia is the third-most populous country in Latin America'
i = features.context.find(pattern)
@@ -469,10 +469,10 @@ def make_example(
}
if labels:
answers = set((label.begin, label.end) for label in labels)
- feature['answers'] = np.array([list(answer) for answer in answers],
+ feature['answers'] = np.array([list(answer) for answer in answers], # pyrefly: ignore[bad-assignment]
np.int64)
else:
- feature['answers'] = np.zeros([0, 2], np.int64)
+ feature['answers'] = np.zeros([0, 2], np.int64) # pyrefly: ignore[bad-assignment]
metrics.Metrics.counter('_', 'examples').inc()
return f'{features.id}--{features.stride_index}', feature
diff --git a/official/projects/triviaqa/sentencepiece_pb2.py b/official/projects/triviaqa/sentencepiece_pb2.py
index 080682d35d8..c56192dc386 100755
--- a/official/projects/triviaqa/sentencepiece_pb2.py
+++ b/official/projects/triviaqa/sentencepiece_pb2.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,10 +17,12 @@
# pylint: disable=protected-access
# Generated by the protocol buffer compiler. DO NOT EDIT!
"""Generated protocol buffer code."""
-from google.protobuf import descriptor as _descriptor
-from google.protobuf import message as _message
-from google.protobuf import reflection as _reflection
-from google.protobuf import symbol_database as _symbol_database
+
+from google.protobuf import descriptor as _descriptor # pyrefly: ignore[missing-module-attribute]
+from google.protobuf import message as _message # pyrefly: ignore[missing-module-attribute]
+from google.protobuf import reflection as _reflection # pyrefly: ignore[missing-module-attribute]
+from google.protobuf import symbol_database as _symbol_database # pyrefly: ignore[missing-module-attribute]
+
# @@protoc_insertion_point(imports)
_sym_db = _symbol_database.Default()
diff --git a/official/projects/triviaqa/train.py b/official/projects/triviaqa/train.py
index ff84f8dc205..fe0eeb29ae2 100644
--- a/official/projects/triviaqa/train.py
+++ b/official/projects/triviaqa/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -24,7 +24,7 @@
from absl import flags
from absl import logging
import gin
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
import sentencepiece as spm
@@ -180,7 +180,7 @@ def fit(model,
logging.info(hparams)
learning_rate_schedule = nlp_optimization.WarmUp(
learning_rate,
- tf.keras.optimizers.schedules.PolynomialDecay(
+ tf_keras.optimizers.schedules.PolynomialDecay(
learning_rate,
num_decay_steps,
end_learning_rate=0.,
@@ -213,7 +213,7 @@ def init_fn(init_checkpoint_path):
model.fit(
train_dataset,
callbacks=[
- tf.keras.callbacks.TensorBoard(model_dir, write_graph=False),
+ tf_keras.callbacks.TensorBoard(model_dir, write_graph=False),
])
ckpt_path = ckpt_manager.save()
if evaluate_fn is None:
@@ -233,12 +233,12 @@ def evaluate(sp_processor, features_map_fn, labels_map_fn, logits_fn,
decode_logits_fn, split_and_pad_fn, distribute_strategy,
validation_dataset, ground_truth):
"""Run evaluation."""
- loss_metric = tf.keras.metrics.Mean()
+ loss_metric = tf_keras.metrics.Mean()
@tf.function
def update_loss(y, logits):
loss_fn = modeling.SpanOrCrossEntropyLoss(
- reduction=tf.keras.losses.Reduction.NONE)
+ reduction=tf_keras.losses.Reduction.NONE)
return loss_metric(loss_fn(y, logits))
predictions = collections.defaultdict(list)
@@ -318,12 +318,12 @@ def main(argv):
_ = tf.random.get_global_generator()
train_dataset = inputs.read_batches(
FLAGS.data_dir,
- tfds.Split.TRAIN,
+ tfds.Split.TRAIN, # pyrefly: ignore[missing-attribute]
FLAGS.batch_size,
shuffle=True,
drop_final_batch=True)
validation_dataset = inputs.read_batches(FLAGS.data_dir,
- tfds.Split.VALIDATION,
+ tfds.Split.VALIDATION, # pyrefly: ignore[missing-attribute]
FLAGS.batch_size)
def train_map_fn(x, y):
diff --git a/official/projects/unified_detector/README.md b/official/projects/unified_detector/README.md
new file mode 100644
index 00000000000..8c24aaf8739
--- /dev/null
+++ b/official/projects/unified_detector/README.md
@@ -0,0 +1,170 @@
+# Towards End-to-End Unified Scene Text Detection and Layout Analysis
+
+
+
+[](https://arxiv.org/abs/2203.15143)
+
+Official TensorFlow 2 implementation of the paper `Towards End-to-End Unified
+Scene Text Detection and Layout Analysis`. If you encounter any issues using the
+code, you are welcome to submit them to the Issues tab or send emails directly
+to us: `hiertext@google.com`.
+
+## Installation
+
+### Set up TensorFlow Models
+
+```bash
+# (Optional) Create and enter a virtual environment
+pip3 install --user virtualenv
+virtualenv -p python3 unified_detector
+source ./unified_detector/bin/activate
+
+# First clone the TensorFlow Models project:
+git clone https://github.com/tensorflow/models.git
+
+# Install the requirements of TensorFlow Models and this repo:
+cd models
+pip3 install -r official/requirements.txt
+pip3 install -r official/projects/unified_detector/requirements.txt
+
+# Compile the protos
+# If `protoc` is not installed, please follow: https://grpc.io/docs/protoc-installation/
+export PYTHONPATH=${PYTHONPATH}:${PWD}/research/
+cd research/object_detection/
+protoc protos/string_int_label_map.proto --python_out=.
+```
+
+### Set up Deeplab2
+
+```bash
+# Clone Deeplab2 anywhere you like
+cd
+git clone https://github.com/google-research/deeplab2.git
+
+# Compile the protos
+protoc deeplab2/*.proto --python_out=.
+
+# Add to PYTHONPATH the directory where deeplab2 sits.
+export PYTHONPATH=${PYTHONPATH}:${PWD}
+```
+
+## Running the model on some images using the provided checkpoint.
+
+### Download the checkpoint
+
+Model | Input Resolution | #object query | line PQ (val) | paragraph PQ (val) | line PQ (test) | paragraph PQ (test)
+---------------------------------------------------------------------------------------------------------------------------------- | ---------------- | ------------- | ------------- | ------------------ | -------------- | -------------------
+Unified-Detector-Line ([ckpt](https://storage.cloud.google.com/tf_model_garden/vision/unified_detector/unified_detector_ckpt.tgz)) | 1024 | 384 | 61.04 | 52.84 | 62.20 | 53.52
+
+### Demo on single images
+
+```bash
+# run from `models/`
+python3 -m official.projects.unified_detector.run_inference \
+--gin_file=official/projects/unified_detector/configs/gin_files/unified_detector_model.gin \
+--ckpt_path= \
+--img_file= \
+--output_path=/demo.jsonl \
+--vis_dir=
+
+```
+
+The output will be stored in jsonl in the same hierarchical format as required
+by the evaluation script of the HierText dataset. There will also be
+visualizations of the word/line/paragraph boundaries. Note that, the unified
+detector produces line-level masks and an affinity matrix for grouping lines
+into paragraphs. For visualization purpose, we split each line mask into pixel
+groups which are defined as connected components/pixels. We visualize these
+groups as `words`. They are not necessarily at the word granularity, though. We
+visualize lines and paragraphs as groupings of these `words` using axis-aligned
+bounding boxes.
+
+##### FAQ
+1. __Q: What is `ckpt_path`__? A: If you download the checkpoint as instructed
+above, you will obtain two files: `/ckpt.index` and
+`/ckpt.data-00000-of-00001`.
+You need to set `--ckpt_path=/ckpt`, i.e. removing the suffix.
+
+
+## Inference and Evaluation on the HierText dataset
+
+### Download the HierText dataset
+
+Clone the [HierText repo](https://github.com/google-research-datasets/hiertext)
+and download the dataset. The `requirements.txt` in this folder already covers
+those in the HierText repo, so there is no need to create a new virtual
+environment again.
+
+### Inference and eval
+
+The following command will run the model on the validation set and compute the
+score. Note that the test set annotation is not released yet, so only validation
+set is used here for demo purposes.
+
+#### Inference
+
+```bash
+# Run from `models/`
+python3 -m official.projects.unified_detector.run_inference \
+--gin_file=official/projects/unified_detector/configs/gin_files/unified_detector_model.gin \
+--ckpt_path= \
+--img_dir= \
+--output_path=/validation_output.jsonl
+
+```
+
+#### Evaluation
+
+```bash
+# Run from `hiertext/`
+python3 eval.py \
+--gt=gt/validation.jsonl \
+--result=/validation_output.jsonl \
+--output=./validation-score.txt \
+--mask_stride=1 \
+--eval_lines \
+--eval_paragraphs \
+--num_workers=0
+
+```
+
+## Train new models.
+
+First, you will need to convert the HierText dataset into TFrecords:
+
+```bash
+# Run from `models/official/projects/unified_detector/data_conversion`
+CUDA_VISIBLE_DEVICES='' python3 convert.py \
+--gt_file=/path/to/gt.jsonl \
+--img_dir=/path/to/image \
+--out_file=/path/to/tfrecords/file-prefix
+
+```
+
+To train the unified detector, run the following script:
+
+```bash
+# Run from `models/`
+python3 -m official.projects.unified_detector.train \
+--mode=train \
+--experiment=unified_detector \
+--model_dir='' \
+--gin_file='official/projects/unified_detector/configs/gin_files/unified_detector_train.gin' \
+--gin_file='official/projects/unified_detector/configs/gin_files/unified_detector_model.gin' \
+--gin_params='InputFn.input_paths = ["/path/to/tfrecords/file-prefix*"]'
+
+```
+
+## Citation
+
+Please cite our [paper](https://arxiv.org/pdf/2203.15143.pdf) if you find this
+work helpful:
+
+```
+@inproceedings{long2022towards,
+ title={Towards End-to-End Unified Scene Text Detection and Layout Analysis},
+ author={Long, Shangbang and Qin, Siyang and Panteleev, Dmitry and Bissacco, Alessandro and Fujii, Yasuhisa and Raptis, Michalis},
+ booktitle={Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition},
+ year={2022}
+}
+```
diff --git a/official/projects/unified_detector/configs/gin_files/unified_detector_model.gin b/official/projects/unified_detector/configs/gin_files/unified_detector_model.gin
new file mode 100644
index 00000000000..4dfcbc1d71c
--- /dev/null
+++ b/official/projects/unified_detector/configs/gin_files/unified_detector_model.gin
@@ -0,0 +1,43 @@
+# Defining the unified detector models.
+
+# Model
+## Backbone
+num_slots = 384
+SyncBatchNormalization.momentum = 0.95
+
+get_max_deep_lab_backbone.num_slots = %num_slots
+
+## Decoder
+intermediate_filters = 256
+num_entity_class = 3 # C + 1 (bkg) + 1 (void)
+
+_get_decoder_head.atrous_rates = (6, 12, 18)
+_get_decoder_head.pixel_space_dim = 128
+_get_decoder_head.pixel_space_intermediate = %intermediate_filters
+_get_decoder_head.num_classes = %num_entity_class
+_get_decoder_head.aux_sem_intermediate = %intermediate_filters
+_get_decoder_head.low_level = [
+ {'feature_key': 'res3', 'channels_project': 64,},
+ {'feature_key': 'res2', 'channels_project': 32,},]
+_get_decoder_head.norm_fn = @SyncBatchNormalization
+_get_embed_head.norm_fn = @LayerNorm
+
+# Loss
+# pq loss
+alpha = 0.75
+tau = 0.3
+_entity_mask_loss.alpha = %alpha
+_instance_discrimination_loss.tau = %tau
+_paragraph_grouping_loss.tau = %tau
+_paragraph_grouping_loss.loss_mode = 'balanced'
+
+
+# Other Model setting
+UniversalDetector.mask_threshold = 0.4
+UniversalDetector.class_threshold = 0.5
+UniversalDetector.filter_area = 32
+universal_detection_loss_weights.loss_segmentation_word = 1e0
+universal_detection_loss_weights.loss_inst_dist = 1e0
+universal_detection_loss_weights.loss_mask_id = 1e-4
+universal_detection_loss_weights.loss_pq = 3e0
+universal_detection_loss_weights.loss_para = 1e0
diff --git a/official/projects/unified_detector/configs/gin_files/unified_detector_train.gin b/official/projects/unified_detector/configs/gin_files/unified_detector_train.gin
new file mode 100644
index 00000000000..384fa4cbbaf
--- /dev/null
+++ b/official/projects/unified_detector/configs/gin_files/unified_detector_train.gin
@@ -0,0 +1,22 @@
+# Defining the input pipeline of unified detector.
+
+# ===== ===== Model ===== =====
+# Internal import 2.
+OcrTask.model_fn = @UniversalDetector
+
+# ===== ===== Data pipeline ===== =====
+InputFn.parser_fn = @UniDetectorParserFn
+InputFn.dataset_type = 'tfrecord'
+InputFn.batch_size = 256
+
+# Internal import 3.
+
+UniDetectorParserFn.output_dimension = 1024
+# Simple data augmentation for now.
+UniDetectorParserFn.rot90_probability = 0.0
+UniDetectorParserFn.use_color_distortion = True
+UniDetectorParserFn.crop_min_scale = 0.5
+UniDetectorParserFn.crop_max_scale = 1.5
+UniDetectorParserFn.crop_min_aspect = 0.8
+UniDetectorParserFn.crop_max_aspect = 1.25
+UniDetectorParserFn.max_num_instance = 384
diff --git a/official/projects/unified_detector/configs/ocr_config.py b/official/projects/unified_detector/configs/ocr_config.py
new file mode 100644
index 00000000000..b1f2e53579e
--- /dev/null
+++ b/official/projects/unified_detector/configs/ocr_config.py
@@ -0,0 +1,78 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""OCR tasks and models configurations."""
+
+import dataclasses
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import optimization
+
+
+@dataclasses.dataclass
+class OcrTaskConfig(cfg.TaskConfig):
+ train_data: cfg.DataConfig = dataclasses.field(default_factory=cfg.DataConfig)
+ model_call_needs_labels: bool = False
+
+
+@exp_factory.register_config_factory('unified_detector')
+def unified_detector() -> cfg.ExperimentConfig:
+ """Configurations for trainer of unified detector."""
+ total_train_steps = 100000
+ summary_interval = steps_per_loop = 200
+ checkpoint_interval = 2000
+ warmup_steps = 1000
+ config = cfg.ExperimentConfig(
+ # Input pipeline and model are configured through Gin.
+ task=OcrTaskConfig(train_data=cfg.DataConfig(is_training=True)),
+ trainer=cfg.TrainerConfig(
+ train_steps=total_train_steps,
+ steps_per_loop=steps_per_loop,
+ summary_interval=summary_interval,
+ checkpoint_interval=checkpoint_interval,
+ max_to_keep=1,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.05,
+ 'include_in_weight_decay': [
+ '^((?!depthwise).)*(kernel|weights):0$',
+ ],
+ 'exclude_from_weight_decay': [
+ '(^((?!kernel).)*:0)|(depthwise_kernel)',
+ ],
+ 'gradient_clip_norm': 10.,
+ },
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 1e-3,
+ 'decay_steps': total_train_steps - warmup_steps,
+ 'alpha': 1e-2,
+ 'offset': warmup_steps,
+ },
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_learning_rate': 1e-5,
+ 'warmup_steps': warmup_steps,
+ }
+ },
+ }),
+ ),
+ )
+ return config
diff --git a/official/projects/unified_detector/data_conversion/convert.py b/official/projects/unified_detector/data_conversion/convert.py
new file mode 100644
index 00000000000..574bebd8704
--- /dev/null
+++ b/official/projects/unified_detector/data_conversion/convert.py
@@ -0,0 +1,66 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+r"""Script to convert HierText to TFExamples.
+
+This script is only intended to run locally.
+
+python3 data_preprocess/convert.py \
+--gt_file=/path/to/gt.jsonl \
+--img_dir=/path/to/image \
+--out_file=/path/to/tfrecords/file-prefix
+
+"""
+
+import json
+import os
+import random
+
+from absl import app
+from absl import flags
+import tensorflow as tf, tf_keras
+import tqdm
+import utils
+
+
+_GT_FILE = flags.DEFINE_string('gt_file', None, 'Path to the GT file')
+_IMG_DIR = flags.DEFINE_string('img_dir', None, 'Path to the image folder.')
+_OUT_FILE = flags.DEFINE_string('out_file', None, 'Path for the tfrecords.')
+_NUM_SHARD = flags.DEFINE_integer(
+ 'num_shard', 100, 'The number of shards of tfrecords.')
+
+
+def main(unused_argv) -> None:
+ annotations = json.load(open(_GT_FILE.value))['annotations']
+ random.shuffle(annotations)
+ n_sample = len(annotations)
+ n_shards = _NUM_SHARD.value
+ n_sample_per_shard = (n_sample - 1) // n_shards + 1
+
+ for shard in tqdm.tqdm(range(n_shards)):
+ output_path = f'{_OUT_FILE.value}-{shard:05}-{n_shards:05}.tfrecords'
+ annotation_subset = annotations[
+ shard * n_sample_per_shard : (shard + 1) * n_sample_per_shard]
+
+ with tf.io.TFRecordWriter(output_path) as file_writer:
+ for annotation in annotation_subset:
+ img_file_path = os.path.join(_IMG_DIR.value,
+ f"{annotation['image_id']}.jpg")
+ tfexample = utils.convert_to_tfe(img_file_path, annotation)
+ file_writer.write(tfexample)
+
+
+if __name__ == '__main__':
+ flags.mark_flags_as_required(['gt_file', 'img_dir', 'out_file'])
+ app.run(main)
diff --git a/official/projects/unified_detector/data_conversion/utils.py b/official/projects/unified_detector/data_conversion/utils.py
new file mode 100644
index 00000000000..2cbaafd9878
--- /dev/null
+++ b/official/projects/unified_detector/data_conversion/utils.py
@@ -0,0 +1,182 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utilities to convert data to TFExamples and store in TFRecords."""
+
+from typing import Any, Dict, List, Tuple, Union
+
+
+import cv2
+import numpy as np
+import tensorflow as tf, tf_keras
+
+
+def encode_image(
+ image_tensor: np.ndarray,
+ encoding_type: str = 'png') -> Union[np.ndarray, tf.Tensor]:
+ """Encode image tensor into byte string."""
+ if encoding_type == 'jpg':
+ image_encoded = tf.image.encode_jpeg(tf.constant(image_tensor))
+ elif encoding_type == 'png':
+ image_encoded = tf.image.encode_png(tf.constant(image_tensor))
+ else:
+ raise ValueError('Invalid encoding type.')
+ if tf.executing_eagerly():
+ image_encoded = image_encoded.numpy()
+ else:
+ image_encoded = image_encoded.eval()
+ return image_encoded
+
+
+def int64_feature(value: Union[int, List[int]]) -> tf.train.Feature:
+ if not isinstance(value, list):
+ value = [value]
+ return tf.train.Feature(int64_list=tf.train.Int64List(value=value))
+
+
+def float_feature(value: Union[float, List[float]]) -> tf.train.Feature:
+ if not isinstance(value, list):
+ value = [value]
+ return tf.train.Feature(float_list=tf.train.FloatList(value=value))
+
+
+def bytes_feature(value: Union[Union[bytes, str], List[Union[bytes, str]]]
+ ) -> tf.train.Feature:
+ if not isinstance(value, list):
+ value = [value]
+ for i in range(len(value)):
+ if not isinstance(value[i], bytes):
+ value[i] = value[i].encode('utf-8')
+ return tf.train.Feature(bytes_list=tf.train.BytesList(value=value))
+
+
+def annotation_to_entities(annotation: Dict[str, Any]) -> List[Dict[str, Any]]:
+ """Flatten the annotation dict to a list of 'entities'."""
+ entities = []
+ for paragraph in annotation['paragraphs']:
+ paragraph_id = len(entities)
+ paragraph['type'] = 3 # 3 for paragraph
+ paragraph['parent_id'] = -1
+ entities.append(paragraph)
+
+ for line in paragraph['lines']:
+ line_id = len(entities)
+ line['type'] = 2 # 2 for line
+ line['parent_id'] = paragraph_id
+ entities.append(line)
+
+ for word in line['words']:
+ word['type'] = 1 # 1 for word
+ word['parent_id'] = line_id
+ entities.append(word)
+
+ return entities
+
+
+def draw_entity_mask(
+ entities: List[Dict[str, Any]],
+ image_shape: Tuple[int, int, int]) -> np.ndarray:
+ """Draw entity id mask.
+
+ Args:
+ entities: A list of entity objects. Should be output from
+ `annotation_to_entities`.
+ image_shape: The shape of the input image.
+ Returns:
+ A (H, W, 3) entity id mask of the same height/width as the image. Each pixel
+ (i, j, :) encodes the entity id of one pixel. Only word entities are
+ rendered. 0 for non-text pixels; word entity ids start from 1.
+ """
+ instance_mask = np.zeros(image_shape, dtype=np.uint8)
+ for i, entity in enumerate(entities):
+ # only draw word masks
+ if entity['type'] != 1:
+ continue
+ vertices = np.array(entity['vertices'])
+ # the pixel value is actually 1 + position in entities
+ entity_id = i + 1
+ if entity_id >= 65536:
+ # As entity_id is encoded in the last two channels, it should be less than
+ # 256**2=65536.
+ raise ValueError(
+ (f'Entity ID overflow: {entity_id}. Currently only entity_id<65536 '
+ 'are supported.'))
+
+ # use the last two channels to encode the entity id.
+ color = [0, entity_id // 256, entity_id % 256]
+ instance_mask = cv2.fillPoly(instance_mask,
+ [np.round(vertices).astype('int32')], color)
+ return instance_mask
+
+
+def convert_to_tfe(img_file_name: str,
+ annotation: Dict[str, Any]) -> tf.train.Example:
+ """Convert the annotation dict into a TFExample."""
+
+ img = cv2.imread(img_file_name)
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
+ h, w, c = img.shape
+ encoded_img = encode_image(img)
+
+ entities = annotation_to_entities(annotation)
+ masks = draw_entity_mask(entities, img.shape)
+ encoded_mask = encode_image(masks)
+
+ # encode attributes
+ parent = []
+ classes = []
+ content_type = []
+ text = []
+ vertices = []
+
+ for entity in entities:
+ parent.append(entity['parent_id'])
+ classes.append(entity['type'])
+ # 0 for annotated; 8 for not annotated
+ content_type.append((0 if entity['legible'] else 8))
+ text.append(entity.get('text', ''))
+ v = np.array(entity['vertices'])
+ vertices.append(','.join(str(float(n)) for n in v.reshape(-1)))
+
+ example = tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ # input images
+ 'image/encoded': bytes_feature(encoded_img),
+ # image format
+ 'image/format': bytes_feature('png'),
+ # image width
+ 'image/width': int64_feature([w]),
+ # image height
+ 'image/height': int64_feature([h]),
+ # image channels
+ 'image/channels': int64_feature([c]),
+ # image key
+ 'image/source_id': bytes_feature(annotation['image_id']),
+ # HxWx3 tensors: channel 2-3 encodes the id of the word entity.
+ 'image/additional_channels/encoded': bytes_feature(encoded_mask),
+ # format of the additional channels
+ 'image/additional_channels/format': bytes_feature('png'),
+ 'image/object/parent': int64_feature(parent),
+ # word / line / paragraph / symbol / ...
+ 'image/object/classes': int64_feature(classes),
+ # text / handwritten / not-annotated / ...
+ 'image/object/content_type': int64_feature(content_type),
+ # string text transcription
+ 'image/object/text': bytes_feature(text),
+ # comma separated coordinates, (x,y) * n
+ 'image/object/vertices': bytes_feature(vertices),
+ })).SerializeToString()
+
+ return example
diff --git a/official/projects/unified_detector/data_loaders/autoaugment.py b/official/projects/unified_detector/data_loaders/autoaugment.py
new file mode 100644
index 00000000000..e7ac8f13001
--- /dev/null
+++ b/official/projects/unified_detector/data_loaders/autoaugment.py
@@ -0,0 +1,754 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""AutoAugment and RandAugment policies for enhanced image preprocessing.
+
+AutoAugment Reference: https://arxiv.org/abs/1805.09501
+RandAugment Reference: https://arxiv.org/abs/1909.13719
+
+This library is adapted from:
+`models/official/efficientnet/autoaugment.py` of
+`https://github.com/tensorflow/tpu`.
+Several changes are made. They are inspired by the TIMM library:
+`https://github.com/rwightman/pytorch-image-models/`
+
+Changes include:
+(1) Random Erasing / Cutout is added, and separated from the random augmentation
+ pool (not sampled as an operation).
+(2) For `posterize` and `solarize`, the arguments are changed such that the
+ level of corruption increases as the `magnitude` argument increases.
+(3) `color`, `contrast`, `brightness`, `sharpness` are randomly enhanced or
+ diminished.
+(4) Magnitude is randomly sampled from a normal distribution.
+(5) Operations are applied with a probability.
+"""
+
+import inspect
+import math
+import tensorflow as tf, tf_keras
+import tensorflow_addons.image as tfa_image
+
+# This signifies the max integer that the controller RNN could predict for the
+# augmentation scheme.
+_MAX_LEVEL = 10.
+
+
+def policy_v0():
+ """Autoaugment policy that was used in AutoAugment Paper."""
+ # Each tuple is an augmentation operation of the form
+ # (operation, probability, magnitude). Each element in policy is a
+ # sub-policy that will be applied sequentially on the image.
+ policy = [
+ [('Equalize', 0.8, 1), ('ShearY', 0.8, 4)],
+ [('Color', 0.4, 9), ('Equalize', 0.6, 3)],
+ [('Color', 0.4, 1), ('Rotate', 0.6, 8)],
+ [('Solarize', 0.8, 3), ('Equalize', 0.4, 7)],
+ [('Solarize', 0.4, 2), ('Solarize', 0.6, 2)],
+ [('Color', 0.2, 0), ('Equalize', 0.8, 8)],
+ [('Equalize', 0.4, 8), ('SolarizeAdd', 0.8, 3)],
+ [('ShearX', 0.2, 9), ('Rotate', 0.6, 8)],
+ [('Color', 0.6, 1), ('Equalize', 1.0, 2)],
+ [('Invert', 0.4, 9), ('Rotate', 0.6, 0)],
+ [('Equalize', 1.0, 9), ('ShearY', 0.6, 3)],
+ [('Color', 0.4, 7), ('Equalize', 0.6, 0)],
+ [('Posterize', 0.4, 6), ('AutoContrast', 0.4, 7)],
+ [('Solarize', 0.6, 8), ('Color', 0.6, 9)],
+ [('Solarize', 0.2, 4), ('Rotate', 0.8, 9)],
+ [('Rotate', 1.0, 7), ('TranslateY', 0.8, 9)],
+ [('ShearX', 0.0, 0), ('Solarize', 0.8, 4)],
+ [('ShearY', 0.8, 0), ('Color', 0.6, 4)],
+ [('Color', 1.0, 0), ('Rotate', 0.6, 2)],
+ [('Equalize', 0.8, 4), ('Equalize', 0.0, 8)],
+ [('Equalize', 1.0, 4), ('AutoContrast', 0.6, 2)],
+ [('ShearY', 0.4, 7), ('SolarizeAdd', 0.6, 7)],
+ [('Posterize', 0.8, 2), ('Solarize', 0.6, 10)],
+ [('Solarize', 0.6, 8), ('Equalize', 0.6, 1)],
+ [('Color', 0.8, 6), ('Rotate', 0.4, 5)],
+ ]
+ return policy
+
+
+def policy_vtest():
+ """Autoaugment test policy for debugging."""
+ # Each tuple is an augmentation operation of the form
+ # (operation, probability, magnitude). Each element in policy is a
+ # sub-policy that will be applied sequentially on the image.
+ policy = [
+ [('TranslateX', 1.0, 4), ('Equalize', 1.0, 10)],
+ ]
+ return policy
+
+
+# pylint: disable=g-long-lambda
+blend = tf.function(lambda i1, i2, factor: tf.cast(
+ tfa_image.blend(tf.cast(i1, tf.float32), tf.cast(i2, tf.float32), factor),
+ tf.uint8))
+# pylint: enable=g-long-lambda
+
+
+def random_erase(image,
+ prob,
+ min_area=0.02,
+ max_area=1 / 3,
+ min_aspect=1 / 3,
+ max_aspect=10 / 3,
+ mode='pixel'):
+ """The random erasing augmentations: https://arxiv.org/pdf/1708.04896.pdf.
+
+ This augmentation is applied after image normalization.
+
+ Args:
+ image: Input image after all other augmentation and normalization. It has
+ type tf.float32.
+ prob: Probability of applying the random erasing operation.
+ min_area: As named.
+ max_area: As named.
+ min_aspect: As named.
+ max_aspect: As named.
+ mode: How the erased area is filled. 'pixel' means white noise (uniform
+ dist).
+
+ Returns:
+ Randomly erased image.
+ """
+
+ image_height = tf.shape(image)[0]
+ image_width = tf.shape(image)[1]
+ image_area = tf.cast(image_width * image_height, tf.float32)
+
+ # Sample width, height
+ erase_area = tf.random.uniform([], min_area, max_area) * image_area
+ log_max_target_ar = tf.math.log(
+ tf.minimum(
+ tf.math.divide(
+ tf.math.square(tf.cast(image_width, tf.float32)), erase_area),
+ max_aspect))
+ log_min_target_ar = tf.math.log(
+ tf.maximum(
+ tf.math.divide(erase_area,
+ tf.math.square(tf.cast(image_height, tf.float32))),
+ min_aspect))
+ erase_aspect_ratio = tf.math.exp(
+ tf.random.uniform([], log_min_target_ar, log_max_target_ar))
+ erase_h = tf.cast(tf.math.sqrt(erase_area / erase_aspect_ratio), tf.int32)
+ erase_w = tf.cast(tf.math.sqrt(erase_area * erase_aspect_ratio), tf.int32)
+
+ # Sample (left, top) of the rectangle to erase
+ erase_left = tf.random.uniform(
+ shape=[], minval=0, maxval=image_width - erase_w, dtype=tf.int32)
+ erase_top = tf.random.uniform(
+ shape=[], minval=0, maxval=image_height - erase_h, dtype=tf.int32)
+ pad_right = image_width - erase_w - erase_left
+ pad_bottom = image_height - erase_h - erase_top
+ mask = tf.pad(
+ tf.zeros([erase_h, erase_w], dtype=image.dtype),
+ [[erase_top, pad_bottom], [erase_left, pad_right]],
+ constant_values=1)
+ mask = tf.expand_dims(mask, -1) # [H, W, 1]
+ if mode == 'pixel':
+ fill = tf.random.truncated_normal(
+ tf.shape(image), 0.0, 1.0, dtype=image.dtype)
+ else:
+ fill = tf.zeros(tf.shape(image), dtype=image.dtype)
+
+ should_apply_op = tf.cast(
+ tf.floor(tf.random.uniform([], dtype=tf.float32) + prob), tf.bool)
+ augmented_image = tf.cond(should_apply_op,
+ lambda: mask * image + (1 - mask) * fill,
+ lambda: image)
+ return augmented_image
+
+
+def solarize(image, threshold=128):
+ # For each pixel in the image, select the pixel
+ # if the value is less than the threshold.
+ # Otherwise, subtract 255 from the pixel.
+ return tf.where(image < threshold, image, 255 - image)
+
+
+def solarize_add(image, addition=0, threshold=128):
+ # For each pixel in the image less than threshold
+ # we add 'addition' amount to it and then clip the
+ # pixel value to be between 0 and 255. The value
+ # of 'addition' is between -128 and 128.
+ added_image = tf.cast(image, tf.int64) + addition
+ added_image = tf.cast(tf.clip_by_value(added_image, 0, 255), tf.uint8)
+ return tf.where(image < threshold, added_image, image)
+
+
+def color(image, factor):
+ """Equivalent of PIL Color."""
+ degenerate = tf.image.grayscale_to_rgb(tf.image.rgb_to_grayscale(image))
+ return blend(degenerate, image, factor)
+
+
+def contrast(image, factor):
+ """Equivalent of PIL Contrast."""
+ degenerate = tf.image.rgb_to_grayscale(image)
+ # Cast before calling tf.histogram.
+ degenerate = tf.cast(degenerate, tf.int32)
+
+ # Compute the grayscale histogram, then compute the mean pixel value,
+ # and create a constant image size of that value. Use that as the
+ # blending degenerate target of the original image.
+ hist = tf.histogram_fixed_width(degenerate, [0, 255], nbins=256)
+ mean = tf.reduce_sum(tf.cast(hist, tf.float32)) / 256.0
+ degenerate = tf.ones_like(degenerate, dtype=tf.float32) * mean
+ degenerate = tf.clip_by_value(degenerate, 0.0, 255.0)
+ degenerate = tf.image.grayscale_to_rgb(tf.cast(degenerate, tf.uint8))
+ return blend(degenerate, image, factor)
+
+
+def brightness(image, factor):
+ """Equivalent of PIL Brightness."""
+ degenerate = tf.zeros_like(image)
+ return blend(degenerate, image, factor)
+
+
+def posterize(image, bits):
+ """Equivalent of PIL Posterize. Smaller `bits` means larger degradation."""
+ shift = 8 - bits
+ return tf.bitwise.left_shift(tf.bitwise.right_shift(image, shift), shift)
+
+
+def rotate(image, degrees, replace):
+ """Rotates the image by degrees either clockwise or counterclockwise.
+
+ Args:
+ image: An image Tensor of type uint8.
+ degrees: Float, a scalar angle in degrees to rotate all images by. If
+ degrees is positive the image will be rotated clockwise otherwise it will
+ be rotated counterclockwise.
+ replace: A one or three value 1D tensor to fill empty pixels caused by the
+ rotate operation.
+
+ Returns:
+ The rotated version of image.
+ """
+ # Convert from degrees to radians.
+ degrees_to_radians = math.pi / 180.0
+ radians = degrees * degrees_to_radians
+
+ # In practice, we should randomize the rotation degrees by flipping
+ # it negatively half the time, but that's done on 'degrees' outside
+ # of the function.
+ if isinstance(replace, list) or isinstance(replace, tuple):
+ replace = replace[0]
+ image = tfa_image.rotate(image, radians, fill_value=replace)
+ return image
+
+
+def translate_x(image, pixels, replace):
+ """Equivalent of PIL Translate in X dimension."""
+ return tfa_image.translate_xy(image, [-pixels, 0], replace)
+
+
+def translate_y(image, pixels, replace):
+ """Equivalent of PIL Translate in Y dimension."""
+ return tfa_image.translate_xy(image, [0, -pixels], replace)
+
+
+def autocontrast(image):
+ """Implements Autocontrast function from PIL using TF ops.
+
+ Args:
+ image: A 3D uint8 tensor.
+
+ Returns:
+ The image after it has had autocontrast applied to it and will be of type
+ uint8.
+ """
+
+ def scale_channel(image):
+ """Scale the 2D image using the autocontrast rule."""
+ # A possibly cheaper version can be done using cumsum/unique_with_counts
+ # over the histogram values, rather than iterating over the entire image.
+ # to compute mins and maxes.
+ lo = tf.cast(tf.reduce_min(image), tf.float32)
+ hi = tf.cast(tf.reduce_max(image), tf.float32)
+
+ # Scale the image, making the lowest value 0 and the highest value 255.
+ def scale_values(im):
+ scale = 255.0 / (hi - lo)
+ offset = -lo * scale
+ im = tf.cast(im, tf.float32) * scale + offset
+ im = tf.clip_by_value(im, 0.0, 255.0)
+ return tf.cast(im, tf.uint8)
+
+ result = tf.cond(hi > lo, lambda: scale_values(image), lambda: image)
+ return result
+
+ # Assumes RGB for now. Scales each channel independently
+ # and then stacks the result.
+ s1 = scale_channel(image[:, :, 0])
+ s2 = scale_channel(image[:, :, 1])
+ s3 = scale_channel(image[:, :, 2])
+ image = tf.stack([s1, s2, s3], 2)
+ return image
+
+
+def sharpness(image, factor):
+ """Implements Sharpness function from PIL using TF ops."""
+ orig_image = image
+ image = tf.cast(image, tf.float32)
+ # Make image 4D for conv operation.
+ image = tf.expand_dims(image, 0)
+ # SMOOTH PIL Kernel.
+ kernel = tf.constant([[1, 1, 1], [1, 5, 1], [1, 1, 1]],
+ dtype=tf.float32,
+ shape=[3, 3, 1, 1]) / 13.
+ # Tile across channel dimension.
+ kernel = tf.tile(kernel, [1, 1, 3, 1])
+ strides = [1, 1, 1, 1]
+ with tf.device('/cpu:0'):
+ # Some augmentation that uses depth-wise conv will cause crashing when
+ # training on GPU. See (b/156242594) for details.
+ degenerate = tf.nn.depthwise_conv2d(image, kernel, strides, padding='VALID')
+ degenerate = tf.clip_by_value(degenerate, 0.0, 255.0)
+ degenerate = tf.squeeze(tf.cast(degenerate, tf.uint8), [0])
+
+ # For the borders of the resulting image, fill in the values of the
+ # original image.
+ mask = tf.ones_like(degenerate)
+ padded_mask = tf.pad(mask, [[1, 1], [1, 1], [0, 0]])
+ padded_degenerate = tf.pad(degenerate, [[1, 1], [1, 1], [0, 0]])
+ result = tf.where(tf.equal(padded_mask, 1), padded_degenerate, orig_image)
+
+ # Blend the final result.
+ return blend(result, orig_image, factor)
+
+
+def equalize(image):
+ """Implements Equalize function from PIL using TF ops."""
+
+ def scale_channel(im, c):
+ """Scale the data in the channel to implement equalize."""
+ im = tf.cast(im[:, :, c], tf.int32)
+ # Compute the histogram of the image channel.
+ histo = tf.histogram_fixed_width(im, [0, 255], nbins=256)
+
+ # For the purposes of computing the step, filter out the nonzeros.
+ nonzero = tf.where(tf.not_equal(histo, 0))
+ nonzero_histo = tf.reshape(tf.gather(histo, nonzero), [-1])
+ step = (tf.reduce_sum(nonzero_histo) - nonzero_histo[-1]) // 255
+
+ def build_lut(histo, step):
+ # Compute the cumulative sum, shifting by step // 2
+ # and then normalization by step.
+ lut = (tf.cumsum(histo) + (step // 2)) // step
+ # Shift lut, prepending with 0.
+ lut = tf.concat([[0], lut[:-1]], 0)
+ # Clip the counts to be in range. This is done
+ # in the C code for image.point.
+ return tf.clip_by_value(lut, 0, 255)
+
+ # If step is zero, return the original image. Otherwise, build
+ # lut from the full histogram and step and then index from it.
+ result = tf.cond(
+ tf.equal(step, 0), lambda: im,
+ lambda: tf.gather(build_lut(histo, step), im))
+
+ return tf.cast(result, tf.uint8)
+
+ # Assumes RGB for now. Scales each channel independently
+ # and then stacks the result.
+ s1 = scale_channel(image, 0)
+ s2 = scale_channel(image, 1)
+ s3 = scale_channel(image, 2)
+ image = tf.stack([s1, s2, s3], 2)
+ return image
+
+
+def invert(image):
+ """Inverts the image pixels."""
+ image = tf.convert_to_tensor(image)
+ return 255 - image
+
+
+NAME_TO_FUNC = {
+ 'AutoContrast': autocontrast,
+ 'Equalize': equalize,
+ 'Invert': invert,
+ 'Rotate': rotate,
+ 'Posterize': posterize,
+ 'PosterizeIncreasing': posterize,
+ 'Solarize': solarize,
+ 'SolarizeIncreasing': solarize,
+ 'SolarizeAdd': solarize_add,
+ 'Color': color,
+ 'ColorIncreasing': color,
+ 'Contrast': contrast,
+ 'ContrastIncreasing': contrast,
+ 'Brightness': brightness,
+ 'BrightnessIncreasing': brightness,
+ 'Sharpness': sharpness,
+ 'SharpnessIncreasing': sharpness,
+ 'ShearX': tfa_image.shear_x,
+ 'ShearY': tfa_image.shear_y,
+ 'TranslateX': translate_x,
+ 'TranslateY': translate_y,
+ 'Cutout': tfa_image.random_cutout,
+ 'Hue': tf.image.adjust_hue,
+}
+
+
+def _randomly_negate_tensor(tensor):
+ """With 50% prob turn the tensor negative."""
+ should_flip = tf.cast(tf.floor(tf.random.uniform([]) + 0.5), tf.bool)
+ final_tensor = tf.cond(should_flip, lambda: -tensor, lambda: tensor)
+ return final_tensor
+
+
+def _rotate_level_to_arg(level):
+ level = (level / _MAX_LEVEL) * 30.
+ level = _randomly_negate_tensor(level)
+ return (level,)
+
+
+def _shrink_level_to_arg(level):
+ """Converts level to ratio by which we shrink the image content."""
+ if level == 0:
+ return (1.0,) # if level is zero, do not shrink the image
+ # Maximum shrinking ratio is 2.9.
+ level = 2. / (_MAX_LEVEL / level) + 0.9
+ return (level,)
+
+
+def _enhance_level_to_arg(level):
+ return ((level / _MAX_LEVEL) * 1.8 + 0.1,)
+
+
+def _enhance_increasing_level_to_arg(level):
+ level = (level / _MAX_LEVEL) * .9
+ level = 1.0 + _randomly_negate_tensor(level)
+ return (level,)
+
+
+def _shear_level_to_arg(level):
+ level = (level / _MAX_LEVEL) * 0.3
+ # Flip level to negative with 50% chance.
+ level = _randomly_negate_tensor(level)
+ return (level,)
+
+
+def _translate_level_to_arg(level, translate_const):
+ level = level / _MAX_LEVEL * translate_const
+ # Flip level to negative with 50% chance.
+ level = _randomly_negate_tensor(level)
+ return (level,)
+
+
+def _posterize_level_to_arg(level):
+ return (tf.cast(level / _MAX_LEVEL * 4, tf.uint8),)
+
+
+def _posterize_increase_level_to_arg(level):
+ return (4 - _posterize_level_to_arg(level)[0],)
+
+
+def _solarize_level_to_arg(level):
+ return (tf.cast(level / _MAX_LEVEL * 256, tf.uint8),)
+
+
+def _solarize_increase_level_to_arg(level):
+ return (256 - _solarize_level_to_arg(level)[0],)
+
+
+def _solarize_add_level_to_arg(level):
+ return (tf.cast(level / _MAX_LEVEL * 110, tf.int64),)
+
+
+def _cutout_arg(level, cutout_size):
+ pad_size = tf.cast(level / _MAX_LEVEL * cutout_size, tf.int32)
+ return (2 * pad_size, 2 * pad_size)
+
+
+def level_to_arg(hparams):
+ return {
+ 'AutoContrast':
+ lambda level: (),
+ 'Equalize':
+ lambda level: (),
+ 'Invert':
+ lambda level: (),
+ 'Rotate':
+ _rotate_level_to_arg,
+ 'Posterize':
+ _posterize_level_to_arg,
+ 'PosterizeIncreasing':
+ _posterize_increase_level_to_arg,
+ 'Solarize':
+ _solarize_level_to_arg,
+ 'SolarizeIncreasing':
+ _solarize_increase_level_to_arg,
+ 'SolarizeAdd':
+ _solarize_add_level_to_arg,
+ 'Color':
+ _enhance_level_to_arg,
+ 'ColorIncreasing':
+ _enhance_increasing_level_to_arg,
+ 'Contrast':
+ _enhance_level_to_arg,
+ 'ContrastIncreasing':
+ _enhance_increasing_level_to_arg,
+ 'Brightness':
+ _enhance_level_to_arg,
+ 'BrightnessIncreasing':
+ _enhance_increasing_level_to_arg,
+ 'Sharpness':
+ _enhance_level_to_arg,
+ 'SharpnessIncreasing':
+ _enhance_increasing_level_to_arg,
+ 'ShearX':
+ _shear_level_to_arg,
+ 'ShearY':
+ _shear_level_to_arg,
+ # pylint:disable=g-long-lambda
+ 'Cutout':
+ lambda level: _cutout_arg(level, hparams['cutout_const']),
+ # pylint:disable=g-long-lambda
+ 'TranslateX':
+ lambda level: _translate_level_to_arg(level, hparams['translate_const'
+ ]),
+ 'TranslateY':
+ lambda level: _translate_level_to_arg(level, hparams['translate_const'
+ ]),
+ 'Hue':
+ lambda level: ((level / _MAX_LEVEL) * 0.25,),
+ # pylint:enable=g-long-lambda
+ }
+
+
+def _parse_policy_info(name, prob, level, replace_value, augmentation_hparams):
+ """Return the function that corresponds to `name` and update `level` param."""
+ func = NAME_TO_FUNC[name]
+ args = level_to_arg(augmentation_hparams)[name](level)
+
+ # Add in replace arg if it is required for the function that is being called.
+ # pytype:disable=wrong-arg-types
+ if 'replace' in inspect.signature(func).parameters.keys(): # pylint: disable=deprecated-method
+ args = tuple(list(args) + [replace_value])
+ # pytype:enable=wrong-arg-types
+
+ return (func, prob, args)
+
+
+def _apply_func_with_prob(func, image, args, prob):
+ """Apply `func` to image w/ `args` as input with probability `prob`."""
+ assert isinstance(args, tuple)
+
+ # Apply the function with probability `prob`.
+ should_apply_op = tf.cast(
+ tf.floor(tf.random.uniform([], dtype=tf.float32) + prob), tf.bool)
+ augmented_image = tf.cond(should_apply_op, lambda: func(image, *args),
+ lambda: image)
+ return augmented_image
+
+
+def select_and_apply_random_policy(policies, image):
+ """Select a random policy from `policies` and apply it to `image`."""
+ policy_to_select = tf.random.uniform([], maxval=len(policies), dtype=tf.int32)
+ # Note that using tf.case instead of tf.conds would result in significantly
+ # larger graphs and would even break export for some larger policies.
+ for (i, policy) in enumerate(policies):
+ image = tf.cond(
+ tf.equal(i, policy_to_select),
+ lambda selected_policy=policy: selected_policy(image),
+ lambda: image)
+ return image
+
+
+def build_and_apply_nas_policy(policies, image, augmentation_hparams):
+ """Build a policy from the given policies passed in and apply to image.
+
+ Args:
+ policies: list of lists of tuples in the form `(func, prob, level)`, `func`
+ is a string name of the augmentation function, `prob` is the probability
+ of applying the `func` operation, `level` is the input argument for
+ `func`.
+ image: tf.Tensor that the resulting policy will be applied to.
+ augmentation_hparams: Hparams associated with the NAS learned policy.
+
+ Returns:
+ A version of image that now has data augmentation applied to it based on
+ the `policies` pass into the function.
+ """
+ replace_value = [128, 128, 128]
+
+ # func is the string name of the augmentation function, prob is the
+ # probability of applying the operation and level is the parameter associated
+ # with the tf op.
+
+ # tf_policies are functions that take in an image and return an augmented
+ # image.
+ tf_policies = []
+ for policy in policies:
+ tf_policy = []
+ # Link string name to the correct python function and make sure the correct
+ # argument is passed into that function.
+ for policy_info in policy:
+ policy_info = list(policy_info) + [replace_value, augmentation_hparams]
+
+ tf_policy.append(_parse_policy_info(*policy_info))
+ # Now build the tf policy that will apply the augmentation procedue
+ # on image.
+ def make_final_policy(tf_policy_):
+
+ def final_policy(image_):
+ for func, prob, args in tf_policy_:
+ image_ = _apply_func_with_prob(func, image_, args, prob)
+ return image_
+
+ return final_policy
+
+ tf_policies.append(make_final_policy(tf_policy))
+
+ augmented_image = select_and_apply_random_policy(tf_policies, image)
+ return augmented_image
+
+
+def distort_image_with_autoaugment(image, augmentation_name):
+ """Applies the AutoAugment policy to `image`.
+
+ AutoAugment is from the paper: https://arxiv.org/abs/1805.09501.
+
+ Args:
+ image: `Tensor` of shape [height, width, 3] representing an image.
+ augmentation_name: The name of the AutoAugment policy to use. The available
+ options are `v0` and `test`. `v0` is the policy used for all of the
+ results in the paper and was found to achieve the best results on the COCO
+ dataset. `v1`, `v2` and `v3` are additional good policies found on the
+ COCO dataset that have slight variation in what operations were used
+ during the search procedure along with how many operations are applied in
+ parallel to a single image (2 vs 3).
+
+ Returns:
+ A tuple containing the augmented versions of `image`.
+ """
+ available_policies = {'v0': policy_v0, 'test': policy_vtest}
+ if augmentation_name not in available_policies:
+ raise ValueError('Invalid augmentation_name: {}'.format(augmentation_name))
+
+ policy = available_policies[augmentation_name]()
+ # Hparams that will be used for AutoAugment.
+ augmentation_hparams = dict(cutout_const=100, translate_const=250)
+
+ return build_and_apply_nas_policy(policy, image, augmentation_hparams)
+
+
+# Cutout is implemented separately.
+_RAND_TRANSFORMS = [
+ 'AutoContrast',
+ 'Equalize',
+ 'Invert',
+ 'Rotate',
+ 'Posterize',
+ 'Solarize',
+ 'Color',
+ 'Contrast',
+ 'Brightness',
+ 'Sharpness',
+ 'ShearX',
+ 'ShearY',
+ 'TranslateX',
+ 'TranslateY',
+ 'SolarizeAdd',
+ 'Hue',
+]
+
+# Cutout is implemented separately.
+_RAND_INCREASING_TRANSFORMS = [
+ 'AutoContrast',
+ 'Equalize',
+ 'Invert',
+ 'Rotate',
+ 'PosterizeIncreasing',
+ 'SolarizeIncreasing',
+ 'SolarizeAdd',
+ 'ColorIncreasing',
+ 'ContrastIncreasing',
+ 'BrightnessIncreasing',
+ 'SharpnessIncreasing',
+ 'ShearX',
+ 'ShearY',
+ 'TranslateX',
+ 'TranslateY',
+ 'Hue',
+]
+
+# These augmentations are not suitable for detection task.
+_NON_COLOR_DISTORTION_OPS = [
+ 'Rotate',
+ 'ShearX',
+ 'ShearY',
+ 'TranslateX',
+ 'TranslateY',
+]
+
+
+def distort_image_with_randaugment(image,
+ num_layers,
+ magnitude,
+ mag_std,
+ inc,
+ prob,
+ color_only=False):
+ """Applies the RandAugment policy to `image`.
+
+ RandAugment is from the paper https://arxiv.org/abs/1909.13719,
+
+ Args:
+ image: `Tensor` of shape [height, width, 3] representing an image. The image
+ should have uint8 type in [0, 255].
+ num_layers: Integer, the number of augmentation transformations to apply
+ sequentially to an image. Represented as (N) in the paper. Usually best
+ values will be in the range [1, 3].
+ magnitude: Integer, shared magnitude across all augmentation operations.
+ Represented as (M) in the paper. Usually best values are in the range [5,
+ 30].
+ mag_std: Randomness of magnitude. The magnitude will be sampled from a
+ normal distribution on the fly.
+ inc: Whether to select aug that increases as magnitude increases.
+ prob: Probability of any aug being applied.
+ color_only: Whether only apply operations that distort color and do not
+ change spatial layouts.
+
+ Returns:
+ The augmented version of `image`.
+ """
+ replace_value = [128] * 3
+ augmentation_hparams = dict(cutout_const=40, translate_const=100)
+ available_ops = _RAND_INCREASING_TRANSFORMS if inc else _RAND_TRANSFORMS
+ if color_only:
+ available_ops = list(
+ filter(lambda op: op not in _NON_COLOR_DISTORTION_OPS, available_ops))
+
+ for layer_num in range(num_layers):
+ op_to_select = tf.random.uniform([],
+ maxval=len(available_ops),
+ dtype=tf.int32)
+ random_magnitude = tf.clip_by_value(
+ tf.random.normal([], magnitude, mag_std), 0., _MAX_LEVEL)
+ with tf.name_scope('randaug_layer_{}'.format(layer_num)):
+ for (i, op_name) in enumerate(available_ops):
+ func, _, args = _parse_policy_info(op_name, prob, random_magnitude,
+ replace_value, augmentation_hparams)
+ image = tf.cond(
+ tf.equal(i, op_to_select),
+ # pylint:disable=g-long-lambda
+ lambda s_func=func, s_args=args: _apply_func_with_prob(
+ s_func, image, s_args, prob),
+ # pylint:enable=g-long-lambda
+ lambda: image)
+ return image
diff --git a/official/projects/unified_detector/data_loaders/input_reader.py b/official/projects/unified_detector/data_loaders/input_reader.py
new file mode 100644
index 00000000000..0d81a97d9a7
--- /dev/null
+++ b/official/projects/unified_detector/data_loaders/input_reader.py
@@ -0,0 +1,270 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Input data reader.
+
+Creates a tf.data.Dataset object from multiple input sstables and use a
+provided data parser function to decode the serialized tf.Example and optionally
+run data augmentation.
+"""
+
+import os
+from typing import Any, Callable, List, Optional, Sequence, Union
+
+import gin
+from six.moves import map # pyrefly: ignore[missing-source-for-stubs]
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from research.object_detection.utils import label_map_util
+from official.core import config_definitions as cfg
+from official.projects.unified_detector.data_loaders import universal_detection_parser # pylint: disable=unused-import
+
+FuncType = Callable[..., Any]
+
+
+@gin.configurable(denylist=['is_training'])
+class InputFn(object):
+ """Input data reader class.
+
+ Creates a tf.data.Dataset object from multiple datasets (optionally performs
+ weighted sampling between different datasets), parses the tf.Example message
+ using `parser_fn`. The datasets can either be stored in SSTable or TfRecord.
+ """
+
+ def __init__(self,
+ is_training: bool,
+ batch_size: Optional[int] = None,
+ data_root: str = '',
+ input_paths: List[str] = gin.REQUIRED,
+ dataset_type: str = 'tfrecord',
+ use_sampling: bool = False,
+ sampling_weights: Optional[Sequence[Union[int, float]]] = None,
+ cycle_length: Optional[int] = 64,
+ shuffle_buffer_size: Optional[int] = 512,
+ parser_fn: Optional[FuncType] = None,
+ parser_num_parallel_calls: Optional[int] = 64,
+ max_intra_op_parallelism: Optional[int] = None,
+ label_map_proto_path: Optional[str] = None,
+ input_filter_fns: Optional[List[FuncType]] = None,
+ input_training_filter_fns: Optional[Sequence[FuncType]] = None,
+ dense_to_ragged_batch: bool = False,
+ data_validator_fn: Optional[Callable[[Sequence[str]],
+ None]] = None):
+ """Input reader constructor.
+
+ Args:
+ is_training: Boolean indicating TRAIN or EVAL.
+ batch_size: Input data batch size. Ignored if batch size is passed through
+ params. In that case, this can be None.
+ data_root: All the relative input paths are based on this location.
+ input_paths: Input file patterns.
+ dataset_type: Can be 'sstable' or 'tfrecord'.
+ use_sampling: Whether to perform weighted sampling between different
+ datasets.
+ sampling_weights: Unnormalized sampling weights. The length should be
+ equal to `input_paths`.
+ cycle_length: The number of input Datasets to interleave from in parallel.
+ If set to None tf.data experimental autotuning is used.
+ shuffle_buffer_size: The random shuffle buffer size.
+ parser_fn: The function to run decoding and data augmentation. The
+ function takes `is_training` as an input, which is passed from here.
+ parser_num_parallel_calls: The number of parallel calls for `parser_fn`.
+ The number of CPU cores is the suggested value. If set to None tf.data
+ experimental autotuning is used.
+ max_intra_op_parallelism: if set limits the max intra op parallelism of
+ functions run on slices of the input.
+ label_map_proto_path: Path to a StringIntLabelMap which will be used to
+ decode the input data.
+ input_filter_fns: A list of functions on the dataset points which returns
+ true for valid data.
+ input_training_filter_fns: A list of functions on the dataset points which
+ returns true for valid data used only for training.
+ dense_to_ragged_batch: Whether to use ragged batching for MPNN format.
+ data_validator_fn: If not None, used to validate the data specified by
+ input_paths.
+
+ Raises:
+ ValueError for invalid input_paths.
+ """
+ self._is_training = is_training
+
+ if data_root:
+ # If an input path is absolute this does not change it.
+ input_paths = [os.path.join(data_root, value) for value in input_paths]
+
+ self._input_paths = input_paths
+ # Disables datasets sampling during eval.
+ self._batch_size = batch_size
+ if is_training:
+ self._use_sampling = use_sampling
+ else:
+ self._use_sampling = False
+ self._sampling_weights = sampling_weights
+ self._cycle_length = (cycle_length if cycle_length else tf.data.AUTOTUNE)
+ self._shuffle_buffer_size = shuffle_buffer_size
+ self._parser_num_parallel_calls = (
+ parser_num_parallel_calls
+ if parser_num_parallel_calls else tf.data.AUTOTUNE)
+ self._max_intra_op_parallelism = max_intra_op_parallelism
+ self._label_map_proto_path = label_map_proto_path
+ if label_map_proto_path:
+ name_to_id = label_map_util.get_label_map_dict(label_map_proto_path)
+ self._lookup_str_keys = list(name_to_id.keys())
+ self._lookup_int_values = list(name_to_id.values())
+ self._parser_fn = parser_fn
+ self._input_filter_fns = input_filter_fns or []
+ if is_training and input_training_filter_fns:
+ self._input_filter_fns.extend(input_training_filter_fns)
+ self._dataset_type = dataset_type
+ self._dense_to_ragged_batch = dense_to_ragged_batch
+
+ if data_validator_fn is not None:
+ data_validator_fn(self._input_paths)
+
+ @property
+ def batch_size(self):
+ return self._batch_size
+
+ def __call__(
+ self,
+ params: cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None
+ ) -> tf.data.Dataset:
+ """Read and parse input datasets, return a tf.data.Dataset object."""
+ # TPUEstimator passes the batch size through params.
+ if params is not None and 'batch_size' in params:
+ batch_size = params['batch_size']
+ else:
+ batch_size = self._batch_size
+
+ per_replica_batch_size = input_context.get_per_replica_batch_size(
+ batch_size) if input_context else batch_size
+
+ with tf.name_scope('input_reader'):
+ dataset = self._build_dataset_from_records()
+ dataset_parser_fn = self._build_dataset_parser_fn()
+
+ dataset = dataset.map(
+ dataset_parser_fn, num_parallel_calls=self._parser_num_parallel_calls)
+ for filter_fn in self._input_filter_fns:
+ dataset = dataset.filter(filter_fn)
+
+ if self._dense_to_ragged_batch:
+ dataset = dataset.apply(
+ tf.data.experimental.dense_to_ragged_batch(
+ batch_size=per_replica_batch_size, drop_remainder=True))
+ else:
+ dataset = dataset.batch(per_replica_batch_size, drop_remainder=True)
+ dataset = dataset.prefetch(tf.data.AUTOTUNE)
+
+ return dataset
+
+ def _fetch_dataset(self, filename: str) -> tf.data.Dataset:
+ """Fetch dataset depending on type.
+
+ Args:
+ filename: Location of dataset.
+
+ Returns:
+ Tf Dataset.
+ """
+
+ data_cls = dataset_fn.pick_dataset_fn(self._dataset_type)
+
+ data = data_cls([filename])
+ return data
+
+ def _build_dataset_parser_fn(self) -> Callable[..., tf.Tensor]:
+ """Depending on label_map and storage type, build a parser_fn."""
+ # Parse the fetched records to input tensors for model function.
+ if self._label_map_proto_path:
+ lookup_initializer = tf.lookup.KeyValueTensorInitializer(
+ keys=tf.constant(self._lookup_str_keys, dtype=tf.string),
+ values=tf.constant(self._lookup_int_values, dtype=tf.int32))
+ name_to_id_table = tf.lookup.StaticHashTable(
+ initializer=lookup_initializer, default_value=0)
+ parser_fn = self._parser_fn( # pyrefly: ignore[not-callable]
+ is_training=self._is_training, label_lookup_table=name_to_id_table)
+ else:
+ parser_fn = self._parser_fn(is_training=self._is_training) # pyrefly: ignore[not-callable]
+
+ return parser_fn
+
+ def _build_dataset_from_records(self) -> tf.data.Dataset:
+ """Build a tf.data.Dataset object from input SSTables.
+
+ If the input data come from multiple SSTables, use the user defined sampling
+ weights to perform sampling. For example, if the sampling weights is
+ [1., 2.], the second dataset will be sampled twice more often than the first
+ one.
+
+ Returns:
+ Dataset built from SSTables.
+ Raises:
+ ValueError for inability to find SSTable files.
+ """
+ all_file_patterns = []
+ if self._use_sampling:
+ for file_pattern in self._input_paths:
+ all_file_patterns.append([file_pattern])
+ # Normalize sampling probabilities.
+ total_weight = sum(self._sampling_weights) # pyrefly: ignore[no-matching-overload]
+ sampling_probabilities = [
+ float(w) / total_weight for w in self._sampling_weights # pyrefly: ignore[not-iterable]
+ ]
+ else:
+ all_file_patterns.append(self._input_paths)
+
+ datasets = []
+ for file_pattern in all_file_patterns:
+ filenames = sum(list(map(tf.io.gfile.glob, file_pattern)), [])
+ if not filenames:
+ raise ValueError(
+ f'Error trying to read input files for file pattern {file_pattern}')
+ # Create a dataset of filenames and shuffle the files. In each epoch,
+ # the file order is shuffled again. This may help if
+ # per_host_input_for_training = false on TPU.
+ dataset = tf.data.Dataset.list_files(
+ file_pattern, shuffle=self._is_training)
+
+ if self._is_training:
+ dataset = dataset.repeat()
+
+ if self._max_intra_op_parallelism:
+ # Disable intra-op parallelism to optimize for throughput instead of
+ # latency.
+ options = tf.data.Options()
+ options.experimental_threading.max_intra_op_parallelism = 1
+ dataset = dataset.with_options(options)
+
+ dataset = dataset.interleave(
+ self._fetch_dataset,
+ cycle_length=self._cycle_length,
+ num_parallel_calls=self._cycle_length,
+ deterministic=(not self._is_training))
+
+ if self._is_training:
+ dataset = dataset.shuffle(self._shuffle_buffer_size)
+
+ datasets.append(dataset)
+
+ if self._use_sampling:
+ assert len(datasets) == len(sampling_probabilities) # pyrefly: ignore[unbound-name]
+ dataset = tf.data.experimental.sample_from_datasets(
+ datasets, sampling_probabilities)
+ else:
+ dataset = datasets[0]
+
+ return dataset
diff --git a/official/projects/unified_detector/data_loaders/tf_example_decoder.py b/official/projects/unified_detector/data_loaders/tf_example_decoder.py
new file mode 100644
index 00000000000..5881cfdd542
--- /dev/null
+++ b/official/projects/unified_detector/data_loaders/tf_example_decoder.py
@@ -0,0 +1,320 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tensorflow Example proto decoder for GOCR."""
+
+from typing import List, Optional, Sequence, Tuple, Union
+
+import tensorflow as tf, tf_keras
+from official.projects.unified_detector.utils.typing import TensorDict
+from official.vision.dataloaders import decoder
+
+
+class TfExampleDecoder(decoder.Decoder):
+ """Tensorflow Example proto decoder."""
+
+ def __init__(self,
+ use_instance_mask: bool = False,
+ additional_class_names: Optional[Sequence[str]] = None,
+ additional_regression_names: Optional[Sequence[str]] = None,
+ num_additional_channels: int = 0):
+ """Constructor.
+
+ keys_to_features is a dictionary mapping the names of the tf.Example
+ fields to tf features, possibly with defaults.
+
+ Uses fixed length for scalars and variable length for vectors.
+
+ Args:
+ use_instance_mask: if False, prevents decoding of the instance mask, which
+ can take a lot of resources.
+ additional_class_names: If not none, a list of additional class names. For
+ additional class name n, named image/object/${n} are expected to be an
+ int vector of length one, and are mapped to tensor dict key
+ groundtruth_${n}.
+ additional_regression_names: If not none, a list of additional regression
+ output names. For additional class name n, named image/object/${n} are
+ expected to be a float vector, and are mapped to tensor dict key
+ groundtruth_${n}.
+ num_additional_channels: The number of additional channels of information
+ present in the tf.Example proto.
+ """
+ self._num_additional_channels = num_additional_channels
+ self._use_instance_mask = use_instance_mask
+
+ self.keys_to_features = {}
+ # Map names in the final tensor dict (output of `self.decode()`) to names in
+ # tf examples, e.g. 'groundtruth_text' -> 'image/object/text'
+ self.name_to_key = {}
+
+ if use_instance_mask:
+ self.keys_to_features.update({
+ 'image/object/mask': tf.io.VarLenFeature(tf.string),
+ })
+
+ # Now we have lists of standard types.
+ # To add new features, just add entries here.
+ # The tuple elements are (example name, tensor name, default value).
+ # If the items_to_handlers part is already set up use None for
+ # the tensor name.
+ # There are other tensor names listed as None which we probably
+ # want to discuss and specify.
+ scalar_strings = [
+ ('image/encoded', None, ''),
+ ('image/format', None, 'jpg'),
+ ('image/additional_channels/encoded', None, ''),
+ ('image/additional_channels/format', None, 'png'),
+ ('image/label_type', 'label_type', ''),
+ ('image/key', 'key', ''),
+ ('image/source_id', 'source_id', ''),
+ ]
+ vector_strings = [
+ ('image/attributes', None, ''),
+ ('image/object/text', 'groundtruth_text', ''),
+ ('image/object/encoded_text', 'groundtruth_encoded_text', ''),
+ ('image/object/vertices', 'groundtruth_vertices', ''),
+ ('image/object/object_type', None, ''),
+ ('image/object/language', 'language', ''),
+ ('image/object/reorderer_type', None, ''),
+ ('image/label_map_path', 'label_map_path', '')
+ ]
+ scalar_ints = [
+ ('image/height', None, 1),
+ ('image/width', None, 1),
+ ('image/channels', None, 3),
+ ]
+ vector_ints = [
+ ('image/object/classes', 'groundtruth_classes', 0),
+ ('image/object/frame_id', 'frame_id', 0),
+ ('image/object/track_id', 'track_id', 0),
+ ('image/object/content_type', 'groundtruth_content_type', 0),
+ ]
+ if additional_class_names:
+ vector_ints += [('image/object/%s' % name, 'groundtruth_%s' % name, 0)
+ for name in additional_class_names]
+ # This one is not yet needed:
+ # scalar_floats = [
+ # ]
+ vector_floats = [
+ ('image/object/weight', 'groundtruth_weight', 0),
+ ('image/object/rbox_tl_x', None, 0),
+ ('image/object/rbox_tl_y', None, 0),
+ ('image/object/rbox_width', None, 0),
+ ('image/object/rbox_height', None, 0),
+ ('image/object/rbox_angle', None, 0),
+ ('image/object/bbox/xmin', None, 0),
+ ('image/object/bbox/xmax', None, 0),
+ ('image/object/bbox/ymin', None, 0),
+ ('image/object/bbox/ymax', None, 0),
+ ]
+ if additional_regression_names:
+ vector_floats += [('image/object/%s' % name, 'groundtruth_%s' % name, 0)
+ for name in additional_regression_names]
+
+ self._init_scalar_features(scalar_strings, tf.string) # pyrefly: ignore[bad-argument-type]
+ self._init_vector_features(vector_strings, tf.string) # pyrefly: ignore[bad-argument-type]
+ self._init_scalar_features(scalar_ints, tf.int64) # pyrefly: ignore[bad-argument-type]
+ self._init_vector_features(vector_ints, tf.int64) # pyrefly: ignore[bad-argument-type]
+ self._init_vector_features(vector_floats, tf.float32) # pyrefly: ignore[bad-argument-type]
+
+ def _init_scalar_features(
+ self,
+ feature_list: List[Tuple[str, Optional[str], Union[str, int, float]]],
+ ftype: tf.dtypes.DType) -> None:
+ for entry in feature_list:
+ self.keys_to_features[entry[0]] = tf.io.FixedLenFeature(
+ (), ftype, default_value=entry[2])
+ if entry[1] is not None:
+ self.name_to_key[entry[1]] = entry[0]
+
+ def _init_vector_features(
+ self,
+ feature_list: List[Tuple[str, Optional[str], Union[str, int, float]]],
+ ftype: tf.dtypes.DType) -> None:
+ for entry in feature_list:
+ self.keys_to_features[entry[0]] = tf.io.VarLenFeature(ftype)
+ if entry[1] is not None:
+ self.name_to_key[entry[1]] = entry[0]
+
+ def _decode_png_instance_masks(self, keys_to_tensors: TensorDict)-> tf.Tensor:
+ """Decode PNG instance segmentation masks and stack into dense tensor.
+
+ The instance segmentation masks are reshaped to [num_instances, height,
+ width].
+
+ Args:
+ keys_to_tensors: A dictionary from keys to tensors.
+
+ Returns:
+ A 3-D float tensor of shape [num_instances, height, width] with values
+ in {0, 1}.
+ """
+
+ def decode_png_mask(image_buffer):
+ image = tf.squeeze(
+ tf.image.decode_image(image_buffer, channels=1), axis=2)
+ image.set_shape([None, None])
+ image = tf.to_float(tf.greater(image, 0))
+ return image
+
+ png_masks = keys_to_tensors['image/object/mask']
+ height = keys_to_tensors['image/height']
+ width = keys_to_tensors['image/width']
+ if isinstance(png_masks, tf.SparseTensor):
+ png_masks = tf.sparse_tensor_to_dense(png_masks, default_value='')
+ return tf.cond(
+ tf.greater(tf.size(png_masks), 0),
+ lambda: tf.map_fn(decode_png_mask, png_masks, dtype=tf.float32),
+ lambda: tf.zeros(tf.to_int32(tf.stack([0, height, width]))))
+
+ def _decode_image(self,
+ parsed_tensors: TensorDict,
+ channel: int = 3) -> TensorDict:
+ """Decodes the image and set its shape (H, W are dynamic; C is fixed)."""
+ image = tf.io.decode_image(parsed_tensors['image/encoded'],
+ channels=channel)
+ image.set_shape([None, None, channel])
+ return {'image': image}
+
+ def _decode_additional_channels(self,
+ parsed_tensors: TensorDict,
+ channel: int = 3) -> TensorDict:
+ """Decodes the additional channels and set its static shape."""
+ channels = tf.io.decode_image(
+ parsed_tensors['image/additional_channels/encoded'], channels=channel)
+ channels.set_shape([None, None, channel])
+ return {'additional_channels': channels}
+
+ def _decode_boxes(self, parsed_tensors: TensorDict) -> TensorDict:
+ """Concat box coordinates in the format of [ymin, xmin, ymax, xmax]."""
+ xmin = parsed_tensors['image/object/bbox/xmin']
+ xmax = parsed_tensors['image/object/bbox/xmax']
+ ymin = parsed_tensors['image/object/bbox/ymin']
+ ymax = parsed_tensors['image/object/bbox/ymax']
+ return {
+ 'groundtruth_aligned_boxes': tf.stack([ymin, xmin, ymax, xmax], axis=-1)
+ }
+
+ def _decode_rboxes(self, parsed_tensors: TensorDict) -> TensorDict:
+ """Concat rbox coordinates: [left, top, box_width, box_height, angle]."""
+ top_left_x = parsed_tensors['image/object/rbox_tl_x']
+ top_left_y = parsed_tensors['image/object/rbox_tl_y']
+ width = parsed_tensors['image/object/rbox_width']
+ height = parsed_tensors['image/object/rbox_height']
+ angle = parsed_tensors['image/object/rbox_angle']
+ return {
+ 'groundtruth_boxes':
+ tf.stack([top_left_x, top_left_y, width, height, angle], axis=-1)
+ }
+
+ def _decode_masks(self, parsed_tensors: TensorDict) -> TensorDict:
+ """Decode a set of PNG masks to the tf.float32 tensors."""
+
+ def _decode_png_mask(png_bytes):
+ mask = tf.squeeze(
+ tf.io.decode_png(png_bytes, channels=1, dtype=tf.uint8), axis=-1)
+ mask = tf.cast(mask, dtype=tf.float32)
+ mask.set_shape([None, None])
+ return mask
+
+ height = parsed_tensors['image/height']
+ width = parsed_tensors['image/width']
+ masks = parsed_tensors['image/object/mask']
+ masks = tf.cond(
+ pred=tf.greater(tf.size(input=masks), 0),
+ true_fn=lambda: tf.map_fn(_decode_png_mask, masks, dtype=tf.float32),
+ false_fn=lambda: tf.zeros([0, height, width], dtype=tf.float32))
+ return {'groundtruth_instance_masks': masks}
+
+ def decode(self, tf_example_string_tensor: tf.string):
+ """Decodes serialized tensorflow example and returns a tensor dictionary.
+
+ Args:
+ tf_example_string_tensor: A string tensor holding a serialized tensorflow
+ example proto.
+
+ Returns:
+ A dictionary contains a subset of the following, depends on the inputs:
+ image: A uint8 tensor of shape [height, width, 3] containing the image.
+ source_id: A string tensor contains image fingerprint.
+ key: A string tensor contains the unique sha256 hash key.
+ label_type: Either `full` or `partial`. `full` means all the text are
+ fully labeled, `partial` otherwise. Currently, this is used by E2E
+ model. If an input image is fully labeled, we update the weights of
+ both the detection and the recognizer. Otherwise, only recognizer part
+ of the model is trained.
+ groundtruth_text: A string tensor list, the original transcriptions.
+ groundtruth_encoded_text: A string tensor list, the class ids for the
+ atoms in the text, after applying the reordering algorithm, in string
+ form. For example "90,71,85,69,86,85,93,90,71,91,1,71,85,93,90,71".
+ This depends on the class label map provided to the conversion
+ program. These are 0 based, with -1 for OOV symbols.
+ groundtruth_classes: A int32 tensor of shape [num_boxes] contains the
+ class id. Note this is 1 based, 0 is reserved for background class.
+ groundtruth_content_type: A int32 tensor of shape [num_boxes] contains
+ the content type. Values correspond to PageLayoutEntity::ContentType.
+ groundtruth_weight: A int32 tensor of shape [num_boxes], either 0 or 1.
+ If a region has weight 0, it will be ignored when computing the
+ losses.
+ groundtruth_boxes: A float tensor of shape [num_boxes, 5] contains the
+ groundtruth rotated rectangles. Each row is in [left, top, box_width,
+ box_height, angle] order, absolute coordinates are used.
+ groundtruth_aligned_boxes: A float tensor of shape [num_boxes, 4]
+ contains the groundtruth axis-aligned rectangles. Each row is in
+ [ymin, xmin, ymax, xmax] order. Currently, this is used to store
+ groundtruth symbol boxes.
+ groundtruth_vertices: A string tensor list contains encoded normalized
+ box or polygon coordinates. E.g. `x1,y1,x2,y2,x3,y3,x4,y4`.
+ groundtruth_instance_masks: A float tensor of shape [num_boxes, height,
+ width] contains binarized image sized instance segmentation masks.
+ `1.0` for positive region, `0.0` otherwise. None if not in tfe.
+ frame_id: A int32 tensor of shape [num_boxes], either `0` or `1`.
+ `0` means object comes from first image, `1` means second.
+ track_id: A int32 tensor of shape [num_boxes], where value indicates
+ identity across frame indices.
+ additional_channels: A uint8 tensor of shape [H, W, C] representing some
+ features.
+ """
+ parsed_tensors = tf.io.parse_single_example(
+ serialized=tf_example_string_tensor, features=self.keys_to_features)
+ for k in parsed_tensors:
+ if isinstance(parsed_tensors[k], tf.SparseTensor):
+ if parsed_tensors[k].dtype == tf.string:
+ parsed_tensors[k] = tf.sparse.to_dense(
+ parsed_tensors[k], default_value='')
+ else:
+ parsed_tensors[k] = tf.sparse.to_dense(
+ parsed_tensors[k], default_value=0)
+
+ decoded_tensors = {}
+ decoded_tensors.update(self._decode_image(parsed_tensors))
+ decoded_tensors.update(self._decode_rboxes(parsed_tensors))
+ decoded_tensors.update(self._decode_boxes(parsed_tensors))
+ if self._use_instance_mask:
+ decoded_tensors[
+ 'groundtruth_instance_masks'] = self._decode_png_instance_masks(
+ parsed_tensors)
+ if self._num_additional_channels:
+ decoded_tensors.update(self._decode_additional_channels(
+ parsed_tensors, self._num_additional_channels))
+
+ # other attributes:
+ for key in self.name_to_key:
+ if key not in decoded_tensors:
+ decoded_tensors[key] = parsed_tensors[self.name_to_key[key]]
+
+ if 'groundtruth_instance_masks' not in decoded_tensors:
+ decoded_tensors['groundtruth_instance_masks'] = None
+
+ return decoded_tensors
diff --git a/official/projects/unified_detector/data_loaders/universal_detection_parser.py b/official/projects/unified_detector/data_loaders/universal_detection_parser.py
new file mode 100644
index 00000000000..b36c418a6e2
--- /dev/null
+++ b/official/projects/unified_detector/data_loaders/universal_detection_parser.py
@@ -0,0 +1,606 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data parser for universal detector."""
+
+import enum
+import functools
+from typing import Any, Tuple
+
+import gin
+import tensorflow as tf, tf_keras
+
+from official.projects.unified_detector.data_loaders import autoaugment
+from official.projects.unified_detector.data_loaders import tf_example_decoder
+from official.projects.unified_detector.utils import utilities
+from official.projects.unified_detector.utils.typing import NestedTensorDict
+from official.projects.unified_detector.utils.typing import TensorDict
+
+
+@gin.constants_from_enum
+class DetectionClass(enum.IntEnum):
+ """As in `PageLayoutEntity.EntityType`."""
+ WORD = 0
+ LINE = 2
+ PARAGRAPH = 3
+ BLOCK = 4
+
+
+NOT_ANNOTATED_ID = 8
+
+
+def _erase(mask: tf.Tensor,
+ feature: tf.Tensor,
+ min_val: float = 0.,
+ max_val: float = 256.) -> tf.Tensor:
+ """Erase the feature maps with a mask.
+
+ Erase feature maps with a mask and replace the erased area with uniform random
+ noise. The mask can have different size from the feature maps.
+
+ Args:
+ mask: an (h, w) binay mask for pixels to erase with. Value 1 represents
+ pixels to erase.
+ feature: the (H, W, C) feature maps to erase from.
+ min_val: The minimum value of random noise.
+ max_val: The maximum value of random noise.
+
+ Returns:
+ The (H, W, C) feature maps, with pixels in mask replaced with noises. It's
+ equal to mask * noise + (1 - mask) * feature.
+ """
+ h, w, c = utilities.resolve_shape(feature)
+ resized_mask = tf.image.resize(
+ tf.tile(tf.expand_dims(tf.cast(mask, tf.float32), -1), (1, 1, c)), (h, w))
+ erased = tf.where(
+ condition=(resized_mask > 0.5),
+ x=tf.cast(tf.random.uniform((h, w, c), min_val, max_val), feature.dtype),
+ y=feature)
+ return erased
+
+
+@gin.configurable(denylist=['is_training'])
+class UniDetectorParserFn(object):
+ """Data parser for universal detector."""
+
+ def __init__(
+ self,
+ is_training: bool,
+ output_dimension: int = 1025,
+ mask_dimension: int = -1,
+ max_num_instance: int = 128,
+ rot90_probability: float = 0.5,
+ use_color_distortion: bool = True,
+ randaug_mag: float = 5.,
+ randaug_std: float = 0.5,
+ randaug_layer: int = 2,
+ randaug_prob: float = 0.5,
+ use_cropping: bool = True,
+ crop_min_scale: float = 0.5,
+ crop_max_scale: float = 1.5,
+ crop_min_aspect: float = 4 / 5,
+ crop_max_aspect: float = 5 / 4,
+ is_shape_defined: bool = True,
+ use_tpu: bool = True,
+ detection_unit: DetectionClass = DetectionClass.LINE,
+ ):
+ """Constructor.
+
+ Args:
+ is_training: bool indicating TRAIN or EVAL.
+ output_dimension: The size of input images.
+ mask_dimension: The size of the output mask. If negative or zero, it will
+ be set the same as output_dimension.
+ max_num_instance: The maximum number of instances to output. If it's
+ negative, padding or truncating will not be performed.
+ rot90_probability: The probability of rotating multiples of 90 degrees.
+ use_color_distortion: Whether to apply color distortions to images (via
+ autoaugment).
+ randaug_mag: (autoaugment parameter) Color distortion magnitude. Note
+ that, this value should be set conservatively, as some color distortions
+ can easily make text illegible e.g. posterize.
+ randaug_std: (autoaugment parameter) Randomness in color distortion
+ magnitude.
+ randaug_layer: (autoaugment parameter) Number of color distortion
+ operations.
+ randaug_prob: (autoaugment parameter) Probabilily of applying each
+ distortion operation.
+ use_cropping: Bool, whether to use random cropping and resizing in
+ training.
+ crop_min_scale: The minimum scale of a random crop.
+ crop_max_scale: The maximum scale of a random crop. If >1, it means the
+ images are downsampled.
+ crop_min_aspect: The minimum aspect ratio of a random crop.
+ crop_max_aspect: The maximum aspect ratio of a random crop.
+ is_shape_defined: Whether to define the static shapes for all features and
+ labels. This must be set to True in TPU training as it requires static
+ shapes for all tensors.
+ use_tpu: Whether the inputs are fed to a TPU device.
+ detection_unit: Whether word or line (or else) is regarded as an entity.
+ The instance masks will be at word or line level.
+ """
+ if is_training and max_num_instance < 0:
+ raise ValueError('In TRAIN mode, padding/truncation is required.')
+
+ self._is_training = is_training
+ self._output_dimension = output_dimension
+ self._mask_dimension = (
+ mask_dimension if mask_dimension > 0 else output_dimension)
+ self._max_num_instance = max_num_instance
+ self._decoder = tf_example_decoder.TfExampleDecoder(
+ num_additional_channels=3, additional_class_names=['parent'])
+ self._use_color_distortion = use_color_distortion
+ self._rot90_probability = rot90_probability
+ self._randaug_mag = randaug_mag
+ self._randaug_std = randaug_std
+ self._randaug_layer = randaug_layer
+ self._randaug_prob = randaug_prob
+ self._use_cropping = use_cropping
+ self._crop_min_scale = crop_min_scale
+ self._crop_max_scale = crop_max_scale
+ self._crop_min_aspect = crop_min_aspect
+ self._crop_max_aspect = crop_max_aspect
+ self._is_shape_defined = is_shape_defined
+ self._use_tpu = use_tpu
+ self._detection_unit = detection_unit
+
+ def __call__(self, value: str) -> Tuple[TensorDict, NestedTensorDict]:
+ """Parsing the data.
+
+ Args:
+ value: The serialized data sample.
+
+ Returns:
+ Two dicts for features and labels.
+ features:
+ 'source_id': id of the sample; only in EVAL mode
+ 'images': the normalized images, (output_dimension, output_dimension, 3)
+ labels:
+ See `_prepare_labels` for its content.
+ """
+ data = self._decoder.decode(value)
+ features = {}
+ labels = {}
+ self._preprocess(data, features, labels)
+ self._rot90k(data, features, labels)
+ self._crop_and_resize(data, features, labels)
+ self._color_distortion_and_normalize(data, features, labels)
+ self._prepare_labels(data, features, labels)
+ self._define_shapes(features, labels)
+ return features, labels
+
+ def _preprocess(self, data: TensorDict, features: TensorDict,
+ unused_labels: TensorDict):
+ """All kinds of preprocessing of the decoded data dict."""
+ # (1) Decode the entity_id_mask: a H*W*1 mask, each pixel equals to
+ # (1 + position) of the entity in the GT entity list. The IDs
+ # (which can be larger than 255) are stored in the last two channels.
+ data['additional_channels'] = tf.cast(data['additional_channels'], tf.int32)
+ entity_id_mask = (
+ data['additional_channels'][:, :, -2:-1] * 256 +
+ data['additional_channels'][:, :, -1:])
+ data['entity_id_mask'] = entity_id_mask
+
+ # (2) Write image id. Used in evaluation.
+ if not self._use_tpu:
+ features['source_id'] = data['source_id']
+
+ # (3) Block mask: area without annotation
+ data['image'] = _erase(
+ data['additional_channels'][:, :, 0],
+ data['image'],
+ min_val=0.,
+ max_val=256.)
+
+ def _rot90k(self, data: TensorDict, unused_features: TensorDict,
+ unused_labels: TensorDict):
+ """Rotate the image, gt_bboxes, masks by 90k degrees."""
+ if not self._is_training:
+ return
+
+ rotate_90_choice = tf.random.uniform([])
+
+ def _rotate():
+ """Rotation.
+
+ These will be rotated:
+ image,
+ rbox,
+ entity_id_mask,
+ TODO(longshangbang): rotate vertices.
+
+ Returns:
+ The rotated tensors of the above fields.
+ """
+ k = tf.random.uniform([], 1, 4, dtype=tf.int32)
+ h, w, _ = utilities.resolve_shape(data['image'])
+ # Image
+ rotated_img = tf.image.rot90(data['image'], k=k, name='image_rot90k')
+ # Box
+ rotate_box_op = functools.partial(
+ utilities.rotate_rboxes90,
+ rboxes=data['groundtruth_boxes'],
+ image_width=w,
+ image_height=h)
+ rotated_boxes = tf.switch_case(
+ k - 1, # Indices start with 1.
+ branch_fns=[
+ lambda: rotate_box_op(rotation_count=1),
+ lambda: rotate_box_op(rotation_count=2),
+ lambda: rotate_box_op(rotation_count=3)
+ ])
+ # Mask
+ rotated_mask = tf.image.rot90(
+ data['entity_id_mask'], k=k, name='mask_rot90k')
+ return rotated_img, rotated_boxes, rotated_mask
+
+ # pylint: disable=g-long-lambda
+ (data['image'], data['groundtruth_boxes'],
+ data['entity_id_mask']) = tf.cond(
+ rotate_90_choice < self._rot90_probability, _rotate, lambda:
+ (data['image'], data['groundtruth_boxes'], data['entity_id_mask']))
+ # pylint: enable=g-long-lambda
+
+ def _crop_and_resize(self, data: TensorDict, unused_features: TensorDict,
+ unused_labels: TensorDict):
+ """Perform random cropping and resizing."""
+ # TODO(longshangbang): resize & translate box as well
+ # TODO(longshangbang): resize & translate vertices as well
+
+ # Get cropping target.
+ h, w = utilities.resolve_shape(data['image'])[:2]
+ left, top, crop_w, crop_h, pad_w, pad_h = self._get_crop_box(
+ tf.cast(h, tf.float32), tf.cast(w, tf.float32))
+
+ # Crop the image. (Pad the images if the crop box is larger than image.)
+ if self._is_training:
+ # padding left, top, right, bottom
+ pad_left = tf.random.uniform([], 0, pad_w + 1, dtype=tf.int32)
+ pad_top = tf.random.uniform([], 0, pad_h + 1, dtype=tf.int32)
+ else:
+ pad_left = 0
+ pad_top = 0
+ cropped_img = tf.image.crop_to_bounding_box(data['image'], top, left,
+ crop_h, crop_w)
+ padded_img = tf.pad(
+ cropped_img,
+ [[pad_top, pad_h - pad_top], [pad_left, pad_w - pad_left], [0, 0]],
+ constant_values=127)
+
+ # Resize images
+ data['resized_image'] = tf.image.resize(
+ padded_img, (self._output_dimension, self._output_dimension))
+ data['resized_image'] = tf.cast(data['resized_image'], tf.uint8)
+
+ # Crop the masks
+ cropped_masks = tf.image.crop_to_bounding_box(data['entity_id_mask'], top,
+ left, crop_h, crop_w)
+ padded_masks = tf.pad(
+ cropped_masks,
+ [[pad_top, pad_h - pad_top], [pad_left, pad_w - pad_left], [0, 0]])
+
+ # Resize masks
+ data['resized_masks'] = tf.image.resize(
+ padded_masks, (self._mask_dimension, self._mask_dimension),
+ method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)
+ data['resized_masks'] = tf.squeeze(data['resized_masks'], -1)
+
+ def _get_crop_box(
+ self, h: tf.Tensor,
+ w: tf.Tensor) -> Tuple[Any, Any, tf.Tensor, tf.Tensor, Any, Any]:
+ """Get the cropping box.
+
+ Args:
+ h: The height of the image to crop. Should be float type.
+ w: The width of the image to crop. Should be float type.
+
+ Returns:
+ A tuple representing (left, top, crop_w, crop_h, pad_w, pad_h).
+ Then in `self._crop_and_resize`, a crop will be extracted with bounding
+ box from top-left corner (left, top) and with size (crop_w, crop_h). This
+ crop will then be padded with (pad_w, pad_h) to square sizes.
+ The outputs also are re-cast to int32 type.
+ """
+ if not self._is_training or not self._use_cropping:
+ # cast back to integers.
+ w = tf.cast(w, tf.int32)
+ h = tf.cast(h, tf.int32)
+ side = tf.maximum(w, h)
+ return 0, 0, w, h, side - w, side - h
+
+ # Get box size
+ scale = tf.random.uniform([], self._crop_min_scale, self._crop_max_scale)
+ max_edge = tf.maximum(w, h)
+ long_edge = max_edge * scale
+
+ sqrt_aspect_ratio = tf.math.sqrt(
+ tf.random.uniform([], self._crop_min_aspect, self._crop_max_aspect))
+ box_h = long_edge / sqrt_aspect_ratio
+ box_w = long_edge * sqrt_aspect_ratio
+
+ # Get box location
+ left = tf.random.uniform([], 0., tf.maximum(0., w - box_w))
+ top = tf.random.uniform([], 0., tf.maximum(0., h - box_h))
+ # Get crop & pad
+ crop_w = tf.minimum(box_w, w - left)
+ crop_h = tf.minimum(box_h, h - top)
+ pad_w = box_w - crop_w
+ pad_h = box_h - crop_h
+ return (tf.cast(left, tf.int32), tf.cast(top, tf.int32),
+ tf.cast(crop_w, tf.int32), tf.cast(crop_h, tf.int32),
+ tf.cast(pad_w, tf.int32), tf.cast(pad_h, tf.int32))
+
+ def _color_distortion_and_normalize(self, data: TensorDict,
+ features: TensorDict,
+ unused_labels: TensorDict):
+ """Distort colors."""
+ if self._is_training and self._use_color_distortion:
+ data['resized_image'] = autoaugment.distort_image_with_randaugment(
+ data['resized_image'], self._randaug_layer, self._randaug_mag,
+ self._randaug_std, True, self._randaug_prob, True)
+ # Normalize
+ features['images'] = utilities.normalize_image_to_range(
+ data['resized_image'])
+
+ def _prepare_labels(self, data: TensorDict, features: TensorDict,
+ labels: TensorDict):
+ """This function prepares the labels.
+
+ These following targets are added to labels['segmentation_output']:
+ 'gt_word_score': A (h, w) float32 mask for textness score. 1 for word,
+ 0 for bkg.
+
+ These following targets are added to labels['instance_labels']:
+ 'num_instance': A float scalar tensor for the total number of
+ instances. It is bounded by the maximum number of instances allowed.
+ It includes the special background instance, so it equals to
+ (1 + entity numbers).
+ 'masks': A (h, w) int32 mask for entity IDs. The value of each pixel is
+ the id of the entity it belongs to. A value of `0` means the bkg mask.
+ 'classes': A (max_num,) int tensor indicating the classes of each
+ instance:
+ 2 for background
+ 1 for text entity
+ 0 for non-object
+ 'masks_sizes': A (max_num,) float tensor for the size of all masks.
+ 'gt_weights': Whether it's difficult / does not have text annotation.
+
+ These following targets are added to labels['paragraph_labels']:
+ 'paragraph_ids': A (max_num,) integer tensor for paragprah id. if `-1`,
+ then no paragraph label for this text.
+ 'has_para_ids': A float scalar; 1.0 if the sample has paragraph labels.
+
+ Args:
+ data: The data dictionary.
+ features: The feature dict.
+ labels: The label dict.
+ """
+ # Segmentation labels:
+ self._get_segmentation_labels(data, features, labels)
+ # Instance labels:
+ self._get_instance_labels(data, features, labels)
+
+ def _get_segmentation_labels(self, data: TensorDict,
+ unused_features: TensorDict,
+ labels: NestedTensorDict):
+ labels['segmentation_output'] = {
+ 'gt_word_score': tf.cast((data['resized_masks'] > 0), tf.float32)
+ }
+
+ def _get_instance_labels(self, data: TensorDict, features: TensorDict,
+ labels: NestedTensorDict):
+ """Generate the labels for text entity detection."""
+
+ labels['instance_labels'] = {}
+ # (1) Depending on `detection_unit`:
+ # Convert the word-id map to line-id map or use the word-id map directly
+ # Word entity ids start from 1 in the map, so pad a -1 at the beginning of
+ # the parent list to counter this offset.
+ padded_parent = tf.concat(
+ [tf.constant([-1]),
+ tf.cast(data['groundtruth_parent'], tf.int32)], 0)
+ if self._detection_unit == DetectionClass.WORD:
+ entity_id_mask = data['resized_masks']
+ elif self._detection_unit == DetectionClass.LINE:
+ # The pixel value is entity_id + 1, shape = [H, W]; 0 for background.
+ # correctness:
+ # 0s in data['resized_masks'] --> padded_parent[0] == -1
+ # i-th entity in plp.entities --> i+1 in data['resized_masks']
+ # --> padded_parent[i+1]
+ # --> data['groundtruth_parent'][i]
+ # --> the parent of i-th entity
+ entity_id_mask = tf.gather(padded_parent, data['resized_masks']) + 1
+ elif self._detection_unit == DetectionClass.PARAGRAPH:
+ # directly segmenting paragraphs; two hops here.
+ entity_id_mask = tf.gather(padded_parent, data['resized_masks']) + 1
+ entity_id_mask = tf.gather(padded_parent, entity_id_mask) + 1
+ else:
+ raise ValueError(f'No such detection unit: {self._detection_unit}')
+ data['entity_id_mask'] = entity_id_mask
+
+ # (2) Get individual masks for entities.
+ entity_selection_mask = tf.equal(data['groundtruth_classes'],
+ self._detection_unit)
+ num_all_entity = utilities.resolve_shape(data['groundtruth_classes'])[0]
+ # entity_ids is a 1-D tensor for IDs of all entities of a certain type.
+ entity_ids = tf.boolean_mask(
+ tf.range(num_all_entity, dtype=tf.int32), entity_selection_mask) # (N,)
+ # +1 to match the entity ids in entity_id_mask
+ entity_ids = tf.reshape(entity_ids, (-1, 1, 1)) + 1
+ individual_masks = tf.expand_dims(entity_id_mask, 0)
+ individual_masks = tf.equal(entity_ids, individual_masks) # (N, H, W), bool
+ # TODO(longshangbang): replace with real mask sizes computing.
+ # Currently, we use full-resolution masks for individual_masks. In order to
+ # compute mask sizes, we need to convert individual_masks to int/float type.
+ # This will cause OOM because the mask is too large.
+ masks_sizes = tf.cast(
+ tf.reduce_any(individual_masks, axis=[1, 2]), tf.float32)
+ # remove empty masks (usually caused by cropping)
+ non_empty_masks_ids = tf.not_equal(masks_sizes, 0)
+ valid_masks = tf.boolean_mask(individual_masks, non_empty_masks_ids)
+ valid_entity_ids = tf.boolean_mask(entity_ids, non_empty_masks_ids)[:, 0, 0]
+
+ # (3) Write num of instance
+ num_instance = tf.reduce_sum(tf.cast(non_empty_masks_ids, tf.float32))
+ num_instance_and_bkg = num_instance + 1
+ if self._max_num_instance >= 0:
+ num_instance_and_bkg = tf.minimum(num_instance_and_bkg,
+ self._max_num_instance)
+ labels['instance_labels']['num_instance'] = num_instance_and_bkg
+
+ # (4) Write instance masks
+ num_entity_int = tf.cast(num_instance, tf.int32)
+ max_num_entities = self._max_num_instance - 1 # Spare 1 for bkg.
+ pad_num = tf.maximum(max_num_entities - num_entity_int, 0)
+ padded_valid_masks = tf.pad(valid_masks, [[0, pad_num], [0, 0], [0, 0]])
+
+ # If there are more instances than allowed, randomly sample some.
+ # `random_selection_mask` is a 0/1 array; the maximum number of 1 is
+ # `self._max_num_instance`; if not bound, it's an array with all 1s.
+ if self._max_num_instance >= 0:
+ padded_size = num_entity_int + pad_num
+ random_selection = tf.random.uniform((padded_size,), dtype=tf.float32)
+ selected_indices = tf.math.top_k(random_selection, k=max_num_entities)[1]
+ random_selection_mask = tf.scatter_nd(
+ indices=tf.expand_dims(selected_indices, axis=-1),
+ updates=tf.ones((max_num_entities,), dtype=tf.bool),
+ shape=(padded_size,))
+ else:
+ random_selection_mask = tf.ones((num_entity_int,), dtype=tf.bool)
+ random_discard_mask = tf.logical_not(random_selection_mask)
+
+ kept_masks = tf.boolean_mask(padded_valid_masks, random_selection_mask)
+ erased_masks = tf.boolean_mask(padded_valid_masks, random_discard_mask)
+ erased_masks = tf.cast(tf.reduce_any(erased_masks, axis=0), tf.float32)
+ # erase text instances that are obmitted.
+ features['images'] = _erase(erased_masks, features['images'], -1., 1.)
+ labels['segmentation_output']['gt_word_score'] *= 1. - erased_masks
+ kept_masks_and_bkg = tf.concat(
+ [
+ tf.math.logical_not(
+ tf.reduce_any(kept_masks, axis=0, keepdims=True)), # bkg
+ kept_masks,
+ ],
+ 0)
+ labels['instance_labels']['masks'] = tf.argmax(kept_masks_and_bkg, axis=0)
+
+ # (5) Write mask size
+ # TODO(longshangbang): replace with real masks sizes
+ masks_sizes = tf.cast(
+ tf.reduce_any(kept_masks_and_bkg, axis=[1, 2]), tf.float32)
+ labels['instance_labels']['masks_sizes'] = masks_sizes
+ # (6) Write classes.
+ classes = tf.ones((num_instance,), dtype=tf.int32)
+ classes = tf.concat([tf.constant(2, tf.int32, (1,)), classes], 0) # bkg
+ if self._max_num_instance >= 0:
+ classes = utilities.truncate_or_pad(classes, self._max_num_instance, 0)
+ labels['instance_labels']['classes'] = classes
+
+ # (7) gt-weights
+ selected_ids = tf.boolean_mask(valid_entity_ids,
+ random_selection_mask[:num_entity_int])
+
+ if self._detection_unit != DetectionClass.PARAGRAPH:
+ gt_text = tf.gather(data['groundtruth_text'], selected_ids - 1)
+ gt_weights = tf.cast(tf.strings.length(gt_text) > 0, tf.float32)
+ else:
+ text_types = tf.concat(
+ [
+ tf.constant([8]),
+ tf.cast(data['groundtruth_content_type'], tf.int32),
+ # TODO(longshangbang): temp solution for tfes with no para labels
+ tf.constant(8, shape=(1000,)),
+ ],
+ 0)
+ para_types = tf.gather(text_types, selected_ids)
+
+ gt_weights = tf.cast(
+ tf.not_equal(para_types, NOT_ANNOTATED_ID), tf.float32)
+
+ gt_weights = tf.concat([tf.constant(1., shape=(1,)), gt_weights], 0) # bkg
+ if self._max_num_instance >= 0:
+ gt_weights = utilities.truncate_or_pad(
+ gt_weights, self._max_num_instance, 0)
+ labels['instance_labels']['gt_weights'] = gt_weights
+
+ # (8) get paragraph label
+ # In this step, an array `{p_i}` is generated. `p_i` is an integer that
+ # indicates the group of paragraph which i-th text belongs to. `p_i` == -1
+ # if this instance is non-text or it has no paragraph labels.
+ # word -> line -> paragraph
+ if self._detection_unit == DetectionClass.WORD:
+ num_hop = 2
+ elif self._detection_unit == DetectionClass.LINE:
+ num_hop = 1
+ elif self._detection_unit == DetectionClass.PARAGRAPH:
+ num_hop = 0
+ else:
+ raise ValueError(f'No such detection unit: {self._detection_unit}. '
+ 'Note that this error should have been raised in '
+ 'previous lines, not here!')
+ para_ids = tf.identity(selected_ids) # == id in plp + 1
+ for _ in range(num_hop):
+ para_ids = tf.gather(padded_parent, para_ids) + 1
+
+ text_types = tf.concat(
+ [
+ tf.constant([8]),
+ tf.cast(data['groundtruth_content_type'], tf.int32),
+ # TODO(longshangbang): tricks for tfes that have not para labels
+ tf.constant(8, shape=(1000,)),
+ ],
+ 0)
+ para_types = tf.gather(text_types, para_ids)
+
+ para_ids = para_ids - 1 # revert to id in plp.entities; -1 for no labels
+ valid_para = tf.cast(tf.not_equal(para_types, NOT_ANNOTATED_ID), tf.int32)
+ para_ids = valid_para * para_ids + (1 - valid_para) * (-1)
+ para_ids = tf.concat([tf.constant([-1]), para_ids], 0) # add bkg
+
+ has_para_ids = tf.cast(tf.reduce_sum(valid_para) > 0, tf.float32)
+
+ if self._max_num_instance >= 0:
+ para_ids = utilities.truncate_or_pad(
+ para_ids, self._max_num_instance, 0, -1)
+ labels['paragraph_labels'] = {
+ 'paragraph_ids': para_ids,
+ 'has_para_ids': has_para_ids
+ }
+
+ def _define_shapes(self, features: TensorDict, labels: TensorDict):
+ """Define the tensor shapes for TPU compiling."""
+ if not self._is_shape_defined:
+ return
+ features['images'] = tf.ensure_shape(
+ features['images'], (self._output_dimension, self._output_dimension, 3))
+ labels['segmentation_output']['gt_word_score'] = tf.ensure_shape(
+ labels['segmentation_output']['gt_word_score'],
+ (self._mask_dimension, self._mask_dimension))
+ labels['instance_labels']['num_instance'] = tf.ensure_shape(
+ labels['instance_labels']['num_instance'], [])
+ if self._max_num_instance >= 0:
+ labels['instance_labels']['masks_sizes'] = tf.ensure_shape(
+ labels['instance_labels']['masks_sizes'], (self._max_num_instance,))
+ labels['instance_labels']['masks'] = tf.ensure_shape(
+ labels['instance_labels']['masks'],
+ (self._mask_dimension, self._mask_dimension))
+ labels['instance_labels']['classes'] = tf.ensure_shape(
+ labels['instance_labels']['classes'], (self._max_num_instance,))
+ labels['instance_labels']['gt_weights'] = tf.ensure_shape(
+ labels['instance_labels']['gt_weights'], (self._max_num_instance,))
+ labels['paragraph_labels']['paragraph_ids'] = tf.ensure_shape(
+ labels['paragraph_labels']['paragraph_ids'],
+ (self._max_num_instance,))
+ labels['paragraph_labels']['has_para_ids'] = tf.ensure_shape(
+ labels['paragraph_labels']['has_para_ids'], [])
diff --git a/official/projects/unified_detector/docs/images/task.png b/official/projects/unified_detector/docs/images/task.png
new file mode 100644
index 00000000000..342ecef630c
Binary files /dev/null and b/official/projects/unified_detector/docs/images/task.png differ
diff --git a/official/projects/unified_detector/external_configurables.py b/official/projects/unified_detector/external_configurables.py
new file mode 100644
index 00000000000..8e2ecfe478e
--- /dev/null
+++ b/official/projects/unified_detector/external_configurables.py
@@ -0,0 +1,22 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Wrap external code in gin."""
+
+import gin
+import gin.tf.external_configurables
+import tensorflow as tf, tf_keras
+
+# Tensorflow.
+gin.external_configurable(tf_keras.layers.experimental.SyncBatchNormalization)
diff --git a/official/projects/unified_detector/modeling/universal_detector.py b/official/projects/unified_detector/modeling/universal_detector.py
new file mode 100644
index 00000000000..ac19bfd4082
--- /dev/null
+++ b/official/projects/unified_detector/modeling/universal_detector.py
@@ -0,0 +1,888 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Universal detector implementation."""
+
+from typing import Any, Dict, Optional, Sequence, Tuple, Union
+
+import gin
+import tensorflow as tf, tf_keras
+
+from deeplab2 import config_pb2
+from deeplab2.model.decoder import max_deeplab as max_deeplab_head
+from deeplab2.model.encoder import axial_resnet_instances
+from deeplab2.model.loss import matchers_ops
+from official.legacy.transformer import transformer
+from official.projects.unified_detector.utils import typing
+from official.projects.unified_detector.utils import utilities
+
+
+EPSILON = 1e-6
+
+
+@gin.configurable
+def universal_detection_loss_weights(
+ loss_segmentation_word: float = 1e0,
+ loss_inst_dist: float = 1e0,
+ loss_mask_id: float = 1e-4,
+ loss_pq: float = 3e0,
+ loss_para: float = 1e0) -> Dict[str, float]:
+ """A function that returns a dict for the weights of loss terms."""
+ return {
+ "loss_segmentation_word": loss_segmentation_word,
+ "loss_inst_dist": loss_inst_dist,
+ "loss_mask_id": loss_mask_id,
+ "loss_pq": loss_pq,
+ "loss_para": loss_para,
+ }
+
+
+@gin.configurable
+class LayerNorm(tf_keras.layers.LayerNormalization):
+ """A wrapper to allow passing the `training` argument.
+
+ The normalization layers in the MaX-DeepLab implementation are passed with
+ the `training` argument. This wrapper enables the usage of LayerNorm.
+ """
+
+ def call(self,
+ inputs: tf.Tensor,
+ training: Optional[bool] = None) -> tf.Tensor:
+ del training
+ return super().call(inputs)
+
+
+@gin.configurable
+def get_max_deep_lab_backbone(num_slots: int = 128):
+ return axial_resnet_instances.get_model(
+ "max_deeplab_s",
+ bn_layer=LayerNorm,
+ block_group_config={
+ "drop_path_schedule": "linear",
+ "axial_use_recompute_grad": False
+ },
+ backbone_use_transformer_beyond_stride=16,
+ extra_decoder_use_transformer_beyond_stride=16,
+ num_mask_slots=num_slots,
+ max_num_mask_slots=num_slots)
+
+
+@gin.configurable
+class UniversalDetector(tf_keras.layers.Layer):
+ """Univeral Detector."""
+ loss_items = ("loss_pq", "loss_inst_dist", "loss_para", "loss_mask_id",
+ "loss_segmentation_word")
+
+ def __init__(self,
+ backbone_fn: tf_keras.layers.Layer = get_max_deep_lab_backbone,
+ mask_threshold: float = 0.4,
+ class_threshold: float = 0.5,
+ filter_area: float = 32,
+ **kwargs: Any):
+ """Constructor.
+
+ Args:
+ backbone_fn: The function to initialize a backbone.
+ mask_threshold: Masks are thresholded with this value.
+ class_threshold: Classification heads are thresholded with this value.
+ filter_area: In inference, detections with area smaller than this
+ threshold will be removed.
+ **kwargs: other keyword arguments passed to the base class.
+ """
+ super().__init__(**kwargs)
+
+ # Model
+ self._backbone_fn = backbone_fn()
+ self._decoder = _get_decoder_head()
+ self._class_embed_head, self._para_embed_head = _get_embed_head()
+ self._para_head, self._para_proj = _get_para_head()
+
+ # Losses
+ # self._max_deeplab_loss = _get_max_deeplab_loss()
+ self._loss_weights = universal_detection_loss_weights()
+
+ # Post-processing
+ self._mask_threshold = mask_threshold
+ self._class_threshold = class_threshold
+ self._filter_area = filter_area
+
+ def _preprocess_labels(self, labels: typing.TensorDict):
+ # Preprocessing
+ # Converted the integer mask to one-hot embedded masks.
+ num_instances = utilities.resolve_shape(
+ labels["instance_labels"]["masks_sizes"])[1]
+ labels["instance_labels"]["masks"] = tf.one_hot(
+ labels["instance_labels"]["masks"],
+ depth=num_instances,
+ axis=1,
+ dtype=tf.float32) # (B, N, H, W)
+
+ def compute_losses(
+ self, labels: typing.NestedTensorDict, outputs: typing.NestedTensorDict
+ ) -> Tuple[tf.Tensor, typing.NestedTensorDict]:
+ """Computes the loss.
+
+ Args:
+ labels: A dictionary of ground-truth labels.
+ outputs: Output from self.call().
+
+ Returns:
+ A scalar total loss tensor and a dictionary for individual losses.
+ """
+ loss_dict = {}
+
+ self._preprocess_labels(labels)
+
+ # Main loss: PQ loss.
+ _entity_mask_loss(loss_dict, labels["instance_labels"],
+ outputs["instance_output"])
+ # Auxiliary loss 1: semantic loss
+ _semantic_loss(loss_dict, labels["segmentation_output"],
+ outputs["segmentation_output"])
+ # Auxiliary loss 2: instance discrimination
+ _instance_discrimination_loss(loss_dict, labels["instance_labels"], outputs)
+ # Auxiliary loss 3: mask id
+ _mask_id_xent_loss(loss_dict, labels["instance_labels"], outputs)
+ # Auxiliary loss 4: paragraph grouping
+ _paragraph_grouping_loss(loss_dict, labels, outputs)
+
+ weighted_loss = [self._loss_weights[k] * v for k, v in loss_dict.items()]
+ total_loss = sum(weighted_loss)
+ return total_loss, loss_dict # pyrefly: ignore[bad-return]
+
+ def call(self,
+ features: typing.TensorDict,
+ training: bool = False) -> typing.NestedTensorDict:
+ """Forward pass of the model.
+
+ Args:
+ features: The input features: {"images": tf.Tensor}. Shape = [B, H, W, C]
+ training: Whether it's training mode.
+
+ Returns:
+ A dictionary of output with this structure:
+ {
+ "max_deep_lab": {
+ All the max deeplab outputs are here, including both backbone and
+ decoder.
+ }
+ "segmentation_output": {
+ "word_score": tf.Tensor, [B, h, w],
+ }
+ "instance_output": {
+ "cls_logits": tf.Tensor, [B, N, C],
+ "mask_id_logits": tf.Tensor, [B, H, W, N],
+ "cls_prob": tf.Tensor, [B, N, C],
+ "mask_id_prob": tf.Tensor, [B, H, W, N],
+ }
+ "postprocessed": {
+ "classes": A (B, N) tensor for the class ids. Zero for non-firing
+ slots.
+ "binary_masks": A (B, H, W, N) tensor for the N binary masks. Masks
+ for void cls are set to zero.
+ "confidence": A (B, N) float tensor for the confidence of "classes".
+ "mask_area": A (B, N) float tensor for the area of each mask.
+ }
+ "transformer_group_feature": (B, N, C) float tensor (normalized),
+ "para_affinity": (B, N, N) float tensor.
+ }
+
+ Class-0 is for void. Class-(C-1) is for background. Class-1~(C-2) is for
+ valid classes.
+ """
+ # backbone
+ backbone_output = self._backbone_fn(features["images"], training)
+ # split instance embedding and paragraph embedding;
+ # then perform paragraph grouping
+ para_fts = self._get_para_outputs(backbone_output, training)
+ affinity = tf.linalg.matmul(para_fts, para_fts, transpose_b=True)
+ # text detection head
+ decoder_output = self._decoder(backbone_output, training)
+ output_dict = {
+ "max_deep_lab": decoder_output,
+ "transformer_group_feature": para_fts,
+ "para_affinity": affinity,
+ }
+ input_shape = utilities.resolve_shape(features["images"])
+ self._get_semantic_outputs(output_dict, input_shape)
+ self._get_instance_outputs(output_dict, input_shape)
+ self._postprocess(output_dict)
+
+ return output_dict
+
+ def _get_para_outputs(self, outputs: typing.TensorDict,
+ training: bool) -> tf.Tensor:
+ """Apply the paragraph head.
+
+ This function first splits the features for instance classification and
+ instance grouping. Then, the additional grouping branch (transformer layers)
+ is applied to further encode the grouping features. Finally, a tensor of
+ normalized grouping features is returned.
+
+ Args:
+ outputs: output dictionary from the backbone.
+ training: training / eval mode mark.
+
+ Returns:
+ The normalized paragraph embedding vector of shape (B, N, C).
+ """
+ # Project the object embeddings into classification feature and grouping
+ # feature.
+ fts = outputs["transformer_class_feature"] # B,N,C
+ class_feature = self._class_embed_head(fts, training)
+ group_feature = self._para_embed_head(fts, training)
+ outputs["transformer_class_feature"] = class_feature
+ outputs["transformer_group_feature"] = group_feature
+
+ # Feed the grouping features into additional group encoding branch.
+ # First we need to build the attention_bias which is used the standard
+ # transformer encoder.
+ input_shape = utilities.resolve_shape(group_feature)
+ b = input_shape[0]
+ n = int(input_shape[1])
+ seq_len = tf.constant(n, shape=(b,))
+ padding_mask = utilities.get_padding_mask_from_valid_lengths(
+ seq_len, n, tf.float32)
+ attention_bias = utilities.get_transformer_attention_bias(padding_mask)
+ group_feature = self._para_proj(
+ self._para_head(group_feature, attention_bias, None, training))
+ return tf.math.l2_normalize(group_feature, axis=-1)
+
+ def _get_semantic_outputs(self, outputs: typing.NestedTensorDict,
+ input_shape: tf.TensorShape):
+ """Add `segmentation_output` to outputs.
+
+ Args:
+ outputs: A dictionary of outputs.
+ input_shape: The shape of the input images.
+ """
+ h, w = input_shape[1:3]
+ # B, H/4, W/4, C
+ semantic_logits = outputs["max_deep_lab"]["semantic_logits"]
+ textness, unused_logits = tf.split(semantic_logits, [2, -1], -1)
+ # Channel[0:2], textness. c0: non-textness, c1: textness.
+ word_score = tf.nn.softmax(textness, -1, "word_score")[:, :, :, 1:2]
+ word_score = tf.squeeze(tf.image.resize(word_score, (h, w)), -1)
+ # Channel[2:] not used yet
+ outputs["segmentation_output"] = {"word_score": word_score}
+
+ def _get_instance_outputs(self, outputs: typing.NestedTensorDict,
+ input_shape: tf.TensorShape):
+ """Add `instance_output` to outputs.
+
+ Args:
+ outputs: A dictionary of outputs.
+ input_shape: The shape of the input images.
+ These following fields are added to outputs["instance_output"]:
+ "cls_logits": tf.Tensor, [B, N, C].
+ "mask_id_logits": tf.Tensor, [B, H, W, N].
+ "cls_prob": tf.Tensor, [B, N, C], softmax probability.
+ "mask_id_prob": tf.Tensor, [B, H, W, N], softmax probability. They are
+ used in training. Masks are all resized to full resolution.
+ """
+ # Get instance_output
+ h, w = input_shape[1:3]
+ ## Classes
+ class_logits = outputs["max_deep_lab"]["transformer_class_logits"]
+ # The MaX-DeepLab repo uses the last logit for void; but we use 0.
+ # Therefore we shift the logits here.
+ class_logits = tf.roll(class_logits, shift=1, axis=-1)
+ class_prob = tf.nn.softmax(class_logits)
+
+ ## Masks
+ mask_id_logits = outputs["max_deep_lab"]["pixel_space_mask_logits"]
+ mask_id_prob = tf.nn.softmax(mask_id_logits)
+ mask_id_logits = tf.image.resize(mask_id_logits, (h, w))
+ mask_id_prob = tf.image.resize(mask_id_prob, (h, w))
+ outputs["instance_output"] = {
+ "cls_logits": class_logits,
+ "mask_id_logits": mask_id_logits,
+ "cls_prob": class_prob,
+ "mask_id_prob": mask_id_prob,
+ }
+
+ def _postprocess(self, outputs: typing.NestedTensorDict):
+ """Post-process (filtering) the outputs.
+
+ Args:
+ outputs: A dictionary of outputs.
+ These following fields are added to outputs["postprocessed"]:
+ "classes": A (B,N) integer tensor for the class ids.
+ "binary_masks": A (B, H, W, N) tensor for the N binarized 0/1 masks. Masks
+ for void cls are set to zero.
+ "confidence": A (B, N) float tensor for the confidence of "classes".
+ "mask_area": A (B, N) float tensor for the area of each mask. They are
+ used in inference / visualization.
+ """
+ # Get postprocessed outputs
+ outputs["postprocessed"] = {}
+
+ ## Masks:
+ mask_id_prob = outputs["instance_output"]["mask_id_prob"]
+ mask_max_prob = tf.reduce_max(mask_id_prob, axis=-1, keepdims=True)
+ thresholded_binary_masks = tf.cast(
+ tf.math.logical_and(
+ tf.equal(mask_max_prob, mask_id_prob),
+ tf.greater_equal(mask_max_prob, self._mask_threshold)), tf.float32)
+ area = tf.reduce_sum(thresholded_binary_masks, axis=(1, 2)) # (B, N)
+ ## Classification:
+ cls_prob = outputs["instance_output"]["cls_prob"]
+ cls_max_prob = tf.reduce_max(cls_prob, axis=-1) # B, N
+ cls_max_id = tf.cast(tf.argmax(cls_prob, axis=-1), tf.float32) # B, N
+
+ ## filtering
+ c = utilities.resolve_shape(cls_prob)[2]
+ non_void = tf.reduce_all(
+ tf.stack(
+ [
+ tf.greater_equal(area, self._filter_area), # mask large enough.
+ tf.not_equal(cls_max_id, 0), # class-0 is for non-object.
+ tf.not_equal(cls_max_id,
+ c - 1), # class-(c-1) is for background (last).
+ tf.greater_equal(cls_max_prob,
+ self._class_threshold) # prob >= thr
+ ],
+ axis=-1),
+ axis=-1)
+ non_void = tf.cast(non_void, tf.float32)
+
+ # Storing
+ outputs["postprocessed"]["classes"] = tf.cast(cls_max_id * non_void,
+ tf.int32)
+ b, n = utilities.resolve_shape(non_void)
+ outputs["postprocessed"]["binary_masks"] = (
+ thresholded_binary_masks * tf.reshape(non_void, (b, 1, 1, n)))
+ outputs["postprocessed"]["confidence"] = cls_max_prob
+ outputs["postprocessed"]["mask_area"] = area
+
+ def _coloring(self, masks: tf.Tensor) -> tf.Tensor:
+ """Coloring segmentation masks.
+
+ Used in visualization.
+
+ Args:
+ masks: A float binary tensor of shape (B, H, W, N), representing `B`
+ samples, with `N` masks of size `H*W` each. Each of the `N` masks will
+ be assigned a random color.
+
+ Returns:
+ A (b, h, w, 3) float tensor in [0., 1.] for the coloring result.
+ """
+ b, h, w, n = utilities.resolve_shape(masks)
+ palette = tf.random.uniform((1, n, 3), 0.5, 1.)
+ colored = tf.reshape(
+ tf.matmul(tf.reshape(masks, (b, -1, n)), palette), (b, h, w, 3))
+ return colored
+
+ def visualize(self,
+ outputs: typing.NestedTensorDict,
+ labels: Optional[typing.TensorDict] = None):
+ """Visualizes the outputs and labels.
+
+ Args:
+ outputs: A dictionary of outputs.
+ labels: A dictionary of labels.
+ The following dict is added to outputs["visualization"]: {
+ "instance": {
+ "pred": A (B, H, W, 3) tensor for the visualized map in [0,1].
+ "gt": A (B, H, W, 3) tensor for the visualized map in [0,1], if labels
+ is present.
+ "concat": Concatenation of "prediction" and "gt" along width axis, if
+ labels is present. }
+ "seg-text": {... Similar to above, but the shape is (B, H, W, 1).} } All
+ of these tensors have a rank of 4 (B, H, W, C).
+ """
+
+ outputs["visualization"] = {}
+ # 1. prediction
+ # 1.1 instance mask
+ binary_masks = outputs["postprocessed"]["binary_masks"]
+ outputs["visualization"]["instance"] = {
+ "pred": self._coloring(binary_masks),
+ }
+ # 1.2 text-seg
+ outputs["visualization"]["seg-text"] = {
+ "pred":
+ tf.expand_dims(outputs["segmentation_output"]["word_score"], -1),
+ }
+
+ # 2. labels
+ if labels is not None:
+ # 2.1 instance mask
+ # (B, N, H, W) -> (B, H, W, N); the first one is bkg so removed.
+ gt_masks = tf.transpose(labels["instance_labels"]["masks"][:, 1:],
+ (0, 2, 3, 1))
+ outputs["visualization"]["instance"]["gt"] = self._coloring(gt_masks)
+ # 2.2 text-seg
+ outputs["visualization"]["seg-text"]["gt"] = tf.expand_dims(
+ labels["segmentation_output"]["gt_word_score"], -1)
+
+ # 3. concat
+ for v in outputs["visualization"].values():
+ # Resize to make the size align. The prediction always has stride=1
+ # resolution, so we make gt align with pred instead of vice versa.
+ v["concat"] = tf.concat(
+ [v["pred"],
+ tf.image.resize(v["gt"],
+ tf.shape(v["pred"])[1:3])],
+ axis=2)
+
+ @tf.function
+ def serve(self, image_tensor: tf.Tensor) -> typing.NestedTensorDict:
+ """Method to be exported for SavedModel.
+
+ Args:
+ image_tensor: A float32 normalized tensor representing an image of shape
+ [1, height, width, channels].
+
+ Returns:
+ Dict of output:
+ classes: (B, N) int32 tensor == o["postprocessed"]["classes"]
+ masks: (B, H, W, N) float32 tensor == o["postprocessed"]["binary_masks"]
+ groups: (B, N, N) float32 tensor == o["para_affinity"]
+ confidence: A (B, N) float tensor == o["postprocessed"]["confidence"]
+ mask_area: A (B, N) float tensor == o["postprocessed"]["mask_area"]
+ """
+ features = {"images": image_tensor}
+ nn_outputs = self(features, False)
+ outputs = {
+ "classes": nn_outputs["postprocessed"]["classes"],
+ "masks": nn_outputs["postprocessed"]["binary_masks"],
+ "confidence": nn_outputs["postprocessed"]["confidence"],
+ "mask_area": nn_outputs["postprocessed"]["mask_area"],
+ "groups": nn_outputs["para_affinity"],
+ }
+ return outputs
+
+
+@gin.configurable()
+def _get_decoder_head(
+ atrous_rates: Sequence[int] = (6, 12, 18),
+ pixel_space_dim: int = 128,
+ pixel_space_intermediate: int = 256,
+ low_level: Sequence[Dict[str, Union[str, int]]] = ({
+ "feature_key": "res3",
+ "channels_project": 64,
+ }, {
+ "feature_key": "res2",
+ "channels_project": 32,
+ }),
+ num_classes=3,
+ aux_sem_intermediate=256,
+ norm_fn=tf_keras.layers.BatchNormalization,
+) -> max_deeplab_head.MaXDeepLab:
+ """Get the MaX-DeepLab prediction head.
+
+ Args:
+ atrous_rates: Dilation rate for astrou conv in the semantic head.
+ pixel_space_dim: The dimension for the final panoptic features.
+ pixel_space_intermediate: The dimension for the layer before
+ `pixel_space_dim` (i.e. the separable 5x5 layer).
+ low_level: A list of dicts for the feature pyramid in forming the semantic
+ output. Each dict represents one skip-path from the backbone.
+ num_classes: Number of classes (entities + bkg) including void. For example,
+ if we only want to detect word, then `num_classes` = 3 (1 for word, 1 for
+ bkg, and 1 for void).
+ aux_sem_intermediate: Similar to `pixel_space_intermediate`, but for the
+ auxiliary semantic output head.
+ norm_fn: The normalization function used in the head.
+
+ Returns:
+ A MaX-DeepLab decoder head (as a keras layer).
+ """
+
+ # Initialize the configs.
+ configs = config_pb2.ModelOptions()
+ configs.decoder.feature_key = "feature_semantic"
+ configs.decoder.atrous_rates.extend(atrous_rates)
+ configs.max_deeplab.pixel_space_head.output_channels = pixel_space_dim
+ configs.max_deeplab.pixel_space_head.head_channels = pixel_space_intermediate
+ for low_level_config in low_level:
+ low_level_ = configs.max_deeplab.auxiliary_low_level.add()
+ low_level_.feature_key = low_level_config["feature_key"] # pyrefly: ignore[bad-assignment]
+ low_level_.channels_project = low_level_config["channels_project"] # pyrefly: ignore[bad-assignment]
+ configs.max_deeplab.auxiliary_semantic_head.output_channels = num_classes
+ configs.max_deeplab.auxiliary_semantic_head.head_channels = aux_sem_intermediate
+
+ return max_deeplab_head.MaXDeepLab(configs.decoder,
+ configs.max_deeplab, 0, norm_fn)
+
+
+class PseudoLayer(tf_keras.layers.Layer):
+ """Pseudo layer for ablation study.
+
+ The `call()` function has the same argument signature as a transformer
+ encoder stack. `unused_ph1` and `unused_ph2` are place holders for this
+ purpose. When studying the effectiveness of using transformer as the
+ grouping branch, we can use this PseudoLayer to replace the transformer to
+ use as a no-transformer baseline.
+
+ To use a single projection layer instead of transformer, simply set `extra_fc`
+ to True.
+ """
+
+ def __init__(self, extra_fc: bool):
+ super().__init__(name="extra_fc")
+ self._extra_fc = extra_fc
+ if extra_fc:
+ self._layer = tf_keras.Sequential([
+ tf_keras.layers.Dense(256, activation="relu"),
+ tf_keras.layers.LayerNormalization(),
+ ])
+
+ def call(self,
+ fts: tf.Tensor,
+ unused_ph1: Optional[tf.Tensor],
+ unused_ph2: Optional[tf.Tensor],
+ training: Optional[bool] = None) -> tf.Tensor:
+ """See base class."""
+ if self._extra_fc:
+ return self._layer(fts, training)
+ return fts
+
+
+@gin.configurable()
+def _get_embed_head(
+ dimension=256,
+ norm_fn=tf_keras.layers.BatchNormalization
+) -> Tuple[tf_keras.Sequential, tf_keras.Sequential]:
+ """Projection layers to get instance & grouping features."""
+ instance_head = tf_keras.Sequential([
+ tf_keras.layers.Dense(dimension, use_bias=False),
+ norm_fn(),
+ tf_keras.layers.ReLU(),
+ ])
+ grouping_head = tf_keras.Sequential([
+ tf_keras.layers.Dense(dimension, use_bias=False),
+ norm_fn(),
+ tf_keras.layers.ReLU(),
+ ])
+ return instance_head, grouping_head
+
+
+@gin.configurable()
+def _get_para_head(
+ dimension=128,
+ num_layer=3,
+ extra_fc=False) -> Tuple[tf_keras.layers.Layer, tf_keras.layers.Layer]:
+ """Get the additional para head.
+
+ Args:
+ dimension: the dimension of the final output.
+ num_layer: the number of transformer layer.
+ extra_fc: Whether an extra single fully-connected layer is used, when
+ num_layer=0.
+
+ Returns:
+ an encoder and a projection layer for the grouping features.
+ """
+ if num_layer > 0:
+ encoder = transformer.EncoderStack(
+ params={
+ "hidden_size": 256,
+ "num_hidden_layers": num_layer,
+ "num_heads": 4,
+ "filter_size": 512,
+ "initializer_gain": 1.0,
+ "attention_dropout": 0.1,
+ "relu_dropout": 0.1,
+ "layer_postprocess_dropout": 0.1,
+ "allow_ffn_pad": True,
+ })
+ else:
+ encoder = PseudoLayer(extra_fc)
+ dense = tf_keras.layers.Dense(dimension)
+ return encoder, dense
+
+
+def _dice_sim(pred: tf.Tensor, ground_truth: tf.Tensor) -> tf.Tensor:
+ """Dice Coefficient for mask similarity.
+
+ Args:
+ pred: The predicted mask. [B, N, H, W], in [0, 1].
+ ground_truth: The ground-truth mask. [B, N, H, W], in [0, 1] or {0, 1}.
+
+ Returns:
+ A matrix for the losses: m[b, i, j] is the dice similarity between pred `i`
+ and gt `j` in batch `b`.
+ """
+ b, n = utilities.resolve_shape(pred)[:2]
+ ground_truth = tf.reshape(
+ tf.transpose(ground_truth, (0, 2, 3, 1)), (b, -1, n)) # B, HW, N
+ pred = tf.reshape(pred, (b, n, -1)) # B, N, HW
+ numerator = tf.matmul(pred, ground_truth) * 2.
+ # TODO(longshangbang): The official implementation does not square the scores.
+ # Need to do experiment to determine which one is better.
+ denominator = (
+ tf.math.reduce_sum(tf.math.square(ground_truth), 1, keepdims=True) +
+ tf.math.reduce_sum(tf.math.square(pred), 2, keepdims=True))
+ return (numerator + EPSILON) / (denominator + EPSILON)
+
+
+def _semantic_loss(
+ loss_dict: Dict[str, tf.Tensor],
+ labels: tf.Tensor,
+ outputs: tf.Tensor,
+):
+ """Auxiliary semantic loss.
+
+ Currently, these losses are added:
+ (1) text/non-text heatmap
+
+ Args:
+ loss_dict: A dictionary for the loss. The values are loss scalars.
+ labels: The label dictionary containing:
+ `gt_word_score`: (B, H, W) tensor for the text/non-text map.
+ outputs: The output dictionary containing:
+ `word_score`: (B, H, W) prediction tensor for `gt_word_score`
+ """
+ pred = tf.expand_dims(outputs["word_score"], 1)
+ gt = tf.expand_dims(labels["gt_word_score"], 1)
+ loss_dict["loss_segmentation_word"] = 1. - tf.reduce_mean(_dice_sim(pred, gt))
+
+
+@gin.configurable
+def _entity_mask_loss(loss_dict: Dict[str, tf.Tensor],
+ labels: tf.Tensor,
+ outputs: tf.Tensor,
+ alpha: float = gin.REQUIRED):
+ """PQ loss for entity-mask training.
+
+ This method adds the PQ loss term to loss_dict directly. The match result will
+ also be stored in outputs (As a [B, N_pred, N_gt] float tensor).
+
+ Args:
+ loss_dict: A dictionary for the loss. The values are loss scalars.
+ labels: A dict containing: `num_instance` - (B,) `masks` - (B, N, H, W)
+ `classes` - (B, N)
+ outputs: A dict containing:
+ `cls_prob`: (B, N, C)
+ `mask_id_prob`: (B, H, W, N)
+ `cls_logits`: (B, N, C)
+ `mask_id_logits`: (B, H, W, N)
+ alpha: Weight for pos/neg balance.
+ """
+ # Classification score: (B, N, N)
+ # in batch b, the probability of prediction i being class of gt j, i.e.:
+ # score[b, i, j] = pred_cls[b, i, gt_cls[b, j]]
+ gt_cls = labels["classes"] # (B, N)
+ pred_cls = outputs["cls_prob"] # (B, N, C)
+ b, n = utilities.resolve_shape(pred_cls)[:2]
+ # indices[b, i, j] = gt_cls[b, j]
+ indices = tf.tile(tf.expand_dims(gt_cls, 1), (1, n, 1))
+ cls_score = tf.gather(pred_cls, tf.cast(indices, tf.int32), batch_dims=2)
+
+ # Mask score (dice): (B, N, N)
+ # mask_score[b, i, j]: dice-similarity for pred i and gt j in batch b.
+ mask_score = _dice_sim(
+ tf.transpose(outputs["mask_id_prob"], (0, 3, 1, 2)), labels["masks"])
+
+ # Get similarity matrix and matching.
+ # padded mask[b, j, i] = -1 << other scores, if i >= num_instance[b]
+ similarity = cls_score * mask_score
+ padded_mask = tf.cast(tf.reshape(tf.range(n), (1, 1, n)), tf.float32)
+ padded_mask = tf.cast(
+ tf.math.greater_equal(padded_mask,
+ tf.reshape(labels["num_instance"], (b, 1, 1))),
+ tf.float32)
+ # The constant value for padding has no effect.
+ masked_similarity = similarity * (1. - padded_mask) + padded_mask * (-1.)
+ matched_mask = matchers_ops.hungarian_matching(-masked_similarity)
+ matched_mask = tf.cast(matched_mask, tf.float32) * (1 - padded_mask)
+ outputs["matched_mask"] = matched_mask
+ # Pos loss
+ loss_pos = (
+ tf.stop_gradient(cls_score) * (-mask_score) +
+ tf.stop_gradient(mask_score) * (-tf.math.log(cls_score)))
+ loss_pos = tf.reduce_sum(loss_pos * matched_mask, axis=[1, 2]) # (B,)
+ # Neg loss
+ matched_pred = tf.cast(tf.reduce_sum(matched_mask, axis=2) > 0,
+ tf.float32) # (B, N)
+ # 0 for void class
+ log_loss = -tf.nn.log_softmax(outputs["cls_logits"])[:, :, 0] # (B, N)
+ loss_neg = tf.reduce_sum(log_loss * (1. - matched_pred), axis=-1) # (B,)
+
+ loss_pq = (alpha * loss_pos + (1 - alpha) * loss_neg) / n
+ loss_pq = tf.reduce_mean(loss_pq)
+ loss_dict["loss_pq"] = loss_pq
+
+
+@gin.configurable
+def _instance_discrimination_loss(loss_dict: Dict[str, Any],
+ labels: Dict[str, Any],
+ outputs: Dict[str, Any],
+ tau: float = gin.REQUIRED):
+ """Instance discrimination loss.
+
+ This method adds the ID loss term to loss_dict directly.
+
+ Args:
+ loss_dict: A dictionary for the loss. The values are loss scalars.
+ labels: The label dictionary.
+ outputs: The output dictionary.
+ tau: The temperature term in the loss
+ """
+ # The normalized feature, shape=(B, H/4, W/4, D)
+ g = outputs["max_deep_lab"]["pixel_space_normalized_feature"]
+ b, h, w = utilities.resolve_shape(g)[:3]
+ # The ground-truth masks, shape=(B, N, H, W) --> (B, N, H/4, W/4)
+ m = labels["masks"]
+ m = tf.image.resize(
+ tf.transpose(m, (0, 2, 3, 1)), (h, w),
+ tf.image.ResizeMethod.NEAREST_NEIGHBOR)
+ m = tf.transpose(m, (0, 3, 1, 2))
+ # The number of ground-truth instance (K), shape=(B,)
+ num = labels["num_instance"]
+ n = utilities.resolve_shape(m)[1] # max number of predictions
+ # is_void[b, i] = 1 if instance i in batch b is a padded slot.
+ is_void = tf.cast(tf.expand_dims(tf.range(n), 0), tf.float32) # (1, n)
+ is_void = tf.cast(
+ tf.math.greater_equal(is_void, tf.expand_dims(num, 1)), tf.float32)
+
+ # (B, N, D)
+ t = tf.math.l2_normalize(tf.einsum("bhwd,bnhw->bnd", g, m), axis=-1)
+ inst_dist_logits = tf.einsum("bhwd,bid->bhwi", g, t) / tau # (B, H, W, N)
+ inst_dist_logits = inst_dist_logits - 100. * tf.reshape(is_void, (b, 1, 1, n))
+ mask_id = tf.cast(
+ tf.einsum("bnhw,n->bhw", m, tf.range(n, dtype=tf.float32)), tf.int32)
+ loss_map = tf.nn.sparse_softmax_cross_entropy_with_logits(
+ labels=mask_id, logits=inst_dist_logits) # B, H, W
+ valid_mask = tf.reduce_sum(m, axis=1)
+ loss_inst_dist = (
+ (tf.reduce_sum(loss_map * valid_mask, axis=[1, 2]) + EPSILON) /
+ (tf.reduce_sum(valid_mask, axis=[1, 2]) + EPSILON))
+ loss_dict["loss_inst_dist"] = tf.reduce_mean(loss_inst_dist)
+
+
+@gin.configurable
+def _paragraph_grouping_loss(
+ loss_dict: Dict[str, Any],
+ labels: Dict[str, Any],
+ outputs: Dict[str, Any],
+ tau: float = gin.REQUIRED,
+ loss_mode="vanilla",
+ fl_alpha: float = 0.25,
+ fl_gamma: float = 2.,
+):
+ """Instance discrimination loss.
+
+ This method adds the para discrimination loss term to loss_dict directly.
+
+ Args:
+ loss_dict: A dictionary for the loss. The values are loss scalars.
+ labels: The label dictionary.
+ outputs: The output dictionary.
+ tau: The temperature term in the loss
+ loss_mode: The type of loss.
+ fl_alpha: alpha value in focal loss
+ fl_gamma: gamma value in focal loss
+ """
+ if "paragraph_labels" not in labels:
+ loss_dict["loss_para"] = 0.
+ return
+ # step 1:
+ # obtain the paragraph labels for each prediction
+ # (batch, pred, gt)
+ matched_matrix = outputs["instance_output"]["matched_mask"] # B, N, N
+ para_label_gt = labels["paragraph_labels"]["paragraph_ids"] # B, N
+ has_para_label_gt = (
+ labels["paragraph_labels"]["has_para_ids"][:, tf.newaxis, tf.newaxis])
+ # '0' means no paragraph labels
+ pred_label_gt = tf.einsum("bij,bj->bi", matched_matrix,
+ tf.cast(para_label_gt + 1, tf.float32))
+ pred_label_gt_pad_col = tf.expand_dims(pred_label_gt, -1) # b,n,1
+ pred_label_gt_pad_row = tf.expand_dims(pred_label_gt, 1) # b,1,n
+ gt_affinity = tf.cast(
+ tf.equal(pred_label_gt_pad_col, pred_label_gt_pad_row), tf.float32)
+ gt_affinity_mask = (
+ has_para_label_gt * pred_label_gt_pad_col * pred_label_gt_pad_row)
+ gt_affinity_mask = tf.cast(tf.not_equal(gt_affinity_mask, 0.), tf.float32)
+
+ # step 2:
+ # get affinity matrix
+ affinity = outputs["para_affinity"]
+
+ # step 3:
+ # compute loss
+ loss_fn = tf_keras.losses.BinaryCrossentropy(
+ from_logits=True,
+ label_smoothing=0,
+ axis=-1,
+ reduction=tf_keras.losses.Reduction.NONE,
+ name="para_dist")
+ affinity = tf.reshape(affinity, (-1, 1)) # (b*n*n, 1)
+ gt_affinity = tf.reshape(gt_affinity, (-1, 1)) # (b*n*n, 1)
+ gt_affinity_mask = tf.reshape(gt_affinity_mask, (-1,)) # (b*n*n,)
+ pointwise_loss = loss_fn(gt_affinity, affinity / tau) # (b*n*n,)
+
+ if loss_mode == "vanilla":
+ loss = (
+ tf.reduce_sum(pointwise_loss * gt_affinity_mask) /
+ (tf.reduce_sum(gt_affinity_mask) + EPSILON))
+ elif loss_mode == "balanced":
+ # pos
+ pos_mask = gt_affinity_mask * gt_affinity[:, 0]
+ pos_loss = (
+ tf.reduce_sum(pointwise_loss * pos_mask) /
+ (tf.reduce_sum(pos_mask) + EPSILON))
+ # neg
+ neg_mask = gt_affinity_mask * (1. - gt_affinity[:, 0])
+ neg_loss = (
+ tf.reduce_sum(pointwise_loss * neg_mask) /
+ (tf.reduce_sum(neg_mask) + EPSILON))
+ loss = 0.25 * pos_loss + 0.75 * neg_loss
+ elif loss_mode == "focal":
+ alpha_wt = fl_alpha * gt_affinity + (1. - fl_alpha) * (1. - gt_affinity)
+ prob_pos = tf.math.sigmoid(affinity / tau)
+ pt = prob_pos * gt_affinity + (1. - prob_pos) * (1. - gt_affinity)
+ fl_loss_pw = tf.stop_gradient(
+ alpha_wt * tf.pow(1. - pt, fl_gamma))[:, 0] * pointwise_loss
+ loss = (
+ tf.reduce_sum(fl_loss_pw * gt_affinity_mask) /
+ (tf.reduce_sum(gt_affinity_mask) + EPSILON))
+ else:
+ raise ValueError(f"Not supported loss mode: {loss_mode}")
+
+ loss_dict["loss_para"] = loss
+
+
+def _mask_id_xent_loss(loss_dict: Dict[str, Any], labels: Dict[str, Any],
+ outputs: Dict[str, Any]):
+ """Mask ID loss.
+
+ This method adds the mask ID loss term to loss_dict directly.
+
+ Args:
+ loss_dict: A dictionary for the loss. The values are loss scalars.
+ labels: The label dictionary.
+ outputs: The output dictionary.
+ """
+ # (B, N, H, W)
+ mask_gt = labels["masks"]
+ # B, H, W, N
+ mask_id_logits = outputs["instance_output"]["mask_id_logits"]
+ # B, N, N
+ matched_matrix = outputs["instance_output"]["matched_mask"]
+ # B, N
+ gt_to_pred_id = tf.cast(tf.math.argmax(matched_matrix, axis=1), tf.float32)
+ # B, H, W
+ mask_id_labels = tf.cast(
+ tf.einsum("bnhw,bn->bhw", mask_gt, gt_to_pred_id), tf.int32)
+ loss_map = tf.nn.sparse_softmax_cross_entropy_with_logits(
+ labels=mask_id_labels, logits=mask_id_logits)
+ valid_mask = tf.reduce_sum(mask_gt, axis=1)
+ loss_mask_id = (
+ (tf.reduce_sum(loss_map * valid_mask, axis=[1, 2]) + EPSILON) /
+ (tf.reduce_sum(valid_mask, axis=[1, 2]) + EPSILON))
+ loss_dict["loss_mask_id"] = tf.reduce_mean(loss_mask_id)
diff --git a/official/projects/unified_detector/registry_imports.py b/official/projects/unified_detector/registry_imports.py
new file mode 100644
index 00000000000..33b8f17c590
--- /dev/null
+++ b/official/projects/unified_detector/registry_imports.py
@@ -0,0 +1,21 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""All necessary imports for registration."""
+
+# pylint: disable=unused-import
+from official.projects.unified_detector import external_configurables
+from official.projects.unified_detector.configs import ocr_config
+from official.projects.unified_detector.tasks import ocr_task
+from official.vision import registry_imports
diff --git a/official/projects/unified_detector/requirements.txt b/official/projects/unified_detector/requirements.txt
new file mode 100644
index 00000000000..081993d5d74
--- /dev/null
+++ b/official/projects/unified_detector/requirements.txt
@@ -0,0 +1,8 @@
+tf-nightly
+gin-config
+opencv-python==4.2.0.32
+absl-py>=1.0.0
+shapely>=1.8.1
+apache_beam>=2.37.0
+matplotlib>=3.5.1
+notebook>=6.4.10
diff --git a/official/projects/unified_detector/run_inference.py b/official/projects/unified_detector/run_inference.py
new file mode 100644
index 00000000000..78ce57ac4a4
--- /dev/null
+++ b/official/projects/unified_detector/run_inference.py
@@ -0,0 +1,222 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+r"""A binary to run unified detector."""
+
+import json
+import os
+from typing import Any, Dict, Sequence, Union
+
+from absl import app
+from absl import flags
+from absl import logging
+
+import cv2
+import gin
+import numpy as np
+import tensorflow as tf, tf_keras
+import tqdm
+
+from official.projects.unified_detector import external_configurables # pylint: disable=unused-import
+from official.projects.unified_detector.modeling import universal_detector
+from official.projects.unified_detector.utils import utilities
+
+
+# group two lines into a paragraph if affinity score higher than this
+_PARA_GROUP_THR = 0.5
+
+
+# MODEL spec
+_GIN_FILE = flags.DEFINE_string(
+ 'gin_file', None, 'Path to the Gin file that defines the model.')
+_CKPT_PATH = flags.DEFINE_string(
+ 'ckpt_path', None, 'Path to the checkpoint directory.')
+_IMG_SIZE = flags.DEFINE_integer(
+ 'img_size', 1024, 'Size of the image fed to the model.')
+
+# Input & Output
+# Note that, all images specified by `img_file` and `img_dir` will be processed.
+_IMG_FILE = flags.DEFINE_multi_string('img_file', [], 'Paths to the images.')
+_IMG_DIR = flags.DEFINE_multi_string(
+ 'img_dir', [], 'Paths to the image directories.')
+_OUTPUT_PATH = flags.DEFINE_string('output_path', None, 'Path for the output.')
+_VIS_DIR = flags.DEFINE_string(
+ 'vis_dir', None, 'Path for the visualization output.')
+
+
+def _preprocess(raw_image: np.ndarray) -> Union[np.ndarray, float]:
+ """Convert a raw image to properly resized, padded, and normalized ndarray."""
+ # (1) convert to tf.Tensor and float32.
+ img_tensor = tf.convert_to_tensor(raw_image, dtype=tf.float32)
+
+ # (2) pad to square.
+ height, width = img_tensor.shape[:2]
+ maximum_side = tf.maximum(height, width)
+ height_pad = maximum_side - height
+ width_pad = maximum_side - width
+ img_tensor = tf.pad(
+ img_tensor, [[0, height_pad], [0, width_pad], [0, 0]],
+ constant_values=127)
+ ratio = maximum_side / _IMG_SIZE.value
+ # (3) resize long side to the maximum length.
+ img_tensor = tf.image.resize(
+ img_tensor, (_IMG_SIZE.value, _IMG_SIZE.value))
+ img_tensor = tf.cast(img_tensor, tf.uint8)
+
+ # (4) normalize
+ img_tensor = utilities.normalize_image_to_range(img_tensor)
+
+ # (5) Add batch dimension and return as numpy array.
+ return tf.expand_dims(img_tensor, 0).numpy(), float(ratio)
+
+
+def load_model() -> tf_keras.layers.Layer:
+ gin.parse_config_file(_GIN_FILE.value)
+ model = universal_detector.UniversalDetector()
+ ckpt = tf.train.Checkpoint(model=model)
+ ckpt_path = _CKPT_PATH.value
+ logging.info('Load ckpt from: %s', ckpt_path)
+ ckpt.restore(ckpt_path).expect_partial()
+ return model
+
+
+def inference(img_file: str, model: tf_keras.layers.Layer) -> Dict[str, Any]:
+ """Inference step."""
+ img = cv2.cvtColor(cv2.imread(img_file), cv2.COLOR_BGR2RGB)
+ img_ndarray, ratio = _preprocess(img)
+
+ output_dict = model.serve(img_ndarray)
+ class_tensor = output_dict['classes'].numpy()
+ mask_tensor = output_dict['masks'].numpy()
+ group_tensor = output_dict['groups'].numpy()
+
+ indices = np.where(class_tensor[0])[0].tolist() # indices of positive slots.
+ mask_list = [
+ mask_tensor[0, :, :, index] for index in indices] # List of mask ndarray.
+
+ # Form lines and words
+ lines = []
+ line_indices = []
+ for index, mask in tqdm.tqdm(zip(indices, mask_list)):
+ line = {
+ 'words': [],
+ 'text': '',
+ }
+
+ contours, _ = cv2.findContours(
+ (mask > 0.).astype(np.uint8),
+ cv2.RETR_TREE,
+ cv2.CHAIN_APPROX_SIMPLE)[-2:]
+ for contour in contours:
+ if (isinstance(contour, np.ndarray) and
+ len(contour.shape) == 3 and
+ contour.shape[0] > 2 and
+ contour.shape[1] == 1 and
+ contour.shape[2] == 2):
+ cnt_list = (contour[:, 0] * ratio).astype(np.int32).tolist()
+ line['words'].append({'text': '', 'vertices': cnt_list})
+ else:
+ logging.error('Invalid contour: %s, discarded', str(contour))
+ if line['words']:
+ lines.append(line)
+ line_indices.append(index)
+
+ # Form paragraphs
+ line_grouping = utilities.DisjointSet(len(line_indices))
+ affinity = group_tensor[0][line_indices][:, line_indices]
+ for i1, i2 in zip(*np.where(affinity > _PARA_GROUP_THR)):
+ line_grouping.union(i1, i2)
+
+ line_groups = line_grouping.to_group()
+ paragraphs = []
+ for line_group in line_groups:
+ paragraph = {'lines': []}
+ for id_ in line_group:
+ paragraph['lines'].append(lines[id_])
+ if paragraph:
+ paragraphs.append(paragraph)
+
+ return paragraphs
+
+
+def main(argv: Sequence[str]) -> None:
+ if len(argv) > 1:
+ raise app.UsageError('Too many command-line arguments.')
+
+ # Get list of images
+ img_lists = []
+ img_lists.extend(_IMG_FILE.value)
+ for img_dir in _IMG_DIR.value:
+ img_lists.extend(tf.io.gfile.glob(os.path.join(img_dir, '*')))
+
+ logging.info('Total number of input images: %d', len(img_lists))
+
+ model = load_model()
+
+ vis_dis = _VIS_DIR.value
+
+ output = {'annotations': []}
+ for img_file in tqdm.tqdm(img_lists):
+ output['annotations'].append({
+ 'image_id': img_file.split('/')[-1].split('.')[0],
+ 'paragraphs': inference(img_file, model),
+ })
+
+ if vis_dis:
+ key = output['annotations'][-1]['image_id']
+ paragraphs = output['annotations'][-1]['paragraphs']
+ img = cv2.cvtColor(cv2.imread(img_file), cv2.COLOR_BGR2RGB)
+ word_bnds = []
+ line_bnds = []
+ para_bnds = []
+ for paragraph in paragraphs:
+ paragraph_points_list = []
+ for line in paragraph['lines']:
+ line_points_list = []
+ for word in line['words']:
+ word_bnds.append(
+ np.array(word['vertices'], np.int32).reshape((-1, 1, 2)))
+ line_points_list.extend(word['vertices'])
+ paragraph_points_list.extend(line_points_list)
+
+ line_points = np.array(line_points_list, np.int32) # (N,2)
+ left = int(np.min(line_points[:, 0]))
+ top = int(np.min(line_points[:, 1]))
+ right = int(np.max(line_points[:, 0]))
+ bottom = int(np.max(line_points[:, 1]))
+ line_bnds.append(
+ np.array([[[left, top]], [[right, top]], [[right, bottom]],
+ [[left, bottom]]], np.int32))
+ para_points = np.array(paragraph_points_list, np.int32) # (N,2)
+ left = int(np.min(para_points[:, 0]))
+ top = int(np.min(para_points[:, 1]))
+ right = int(np.max(para_points[:, 0]))
+ bottom = int(np.max(para_points[:, 1]))
+ para_bnds.append(
+ np.array([[[left, top]], [[right, top]], [[right, bottom]],
+ [[left, bottom]]], np.int32))
+
+ for name, bnds in zip(['paragraph', 'line', 'word'],
+ [para_bnds, line_bnds, word_bnds]):
+ vis = cv2.polylines(img, bnds, True, (0, 0, 255), 2)
+ cv2.imwrite(os.path.join(vis_dis, f'{key}-{name}.jpg'),
+ cv2.cvtColor(vis, cv2.COLOR_RGB2BGR))
+
+ with tf.io.gfile.GFile(_OUTPUT_PATH.value, mode='w') as f:
+ f.write(json.dumps(output, ensure_ascii=False, indent=2))
+
+
+if __name__ == '__main__':
+ flags.mark_flags_as_required(['gin_file', 'ckpt_path', 'output_path'])
+ app.run(main)
diff --git a/official/projects/unified_detector/tasks/all_models.py b/official/projects/unified_detector/tasks/all_models.py
new file mode 100644
index 00000000000..42bb899c30d
--- /dev/null
+++ b/official/projects/unified_detector/tasks/all_models.py
@@ -0,0 +1,23 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Import all models.
+
+All model files are imported here so that they can be referenced in Gin. Also,
+importing here avoids making ocr_task.py too messy.
+"""
+
+# pylint: disable=unused-import
+
+from official.projects.unified_detector.modeling import universal_detector
diff --git a/official/projects/unified_detector/tasks/ocr_task.py b/official/projects/unified_detector/tasks/ocr_task.py
new file mode 100644
index 00000000000..05c36e8c918
--- /dev/null
+++ b/official/projects/unified_detector/tasks/ocr_task.py
@@ -0,0 +1,108 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Task definition for ocr."""
+
+from typing import Callable, Dict, Optional, Sequence, Tuple, Union
+
+import gin
+import tensorflow as tf, tf_keras
+
+from official.core import base_task
+from official.core import config_definitions as cfg
+from official.core import task_factory
+from official.projects.unified_detector.configs import ocr_config
+from official.projects.unified_detector.data_loaders import input_reader
+from official.projects.unified_detector.tasks import all_models # pylint: disable=unused-import
+from official.projects.unified_detector.utils import typing
+
+NestedTensorDict = typing.NestedTensorDict
+ModelType = Union[tf_keras.layers.Layer, tf_keras.Model]
+
+
+@task_factory.register_task_cls(ocr_config.OcrTaskConfig)
+@gin.configurable
+class OcrTask(base_task.Task):
+ """Defining the OCR training task."""
+
+ _loss_items = []
+
+ def __init__(self,
+ params: cfg.TaskConfig,
+ logging_dir: Optional[str] = None,
+ name: Optional[str] = None,
+ model_fn: Callable[..., ModelType] = gin.REQUIRED):
+ super().__init__(params, logging_dir, name)
+ self._modef_fn = model_fn
+
+ def build_model(self) -> ModelType:
+ """Build and return the model, record the loss items as well."""
+ model = self._modef_fn()
+ self._loss_items.extend(model.loss_items)
+ return model
+
+ def build_inputs(
+ self,
+ params: cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None
+ ) -> tf.data.Dataset:
+ """Build the tf.data.Dataset instance."""
+ return input_reader.InputFn(is_training=params.is_training)({},
+ input_context)
+
+ def build_metrics(self,
+ training: bool = True) -> Sequence[tf_keras.metrics.Metric]:
+ """Build the metrics (currently, only for loss summaries in TensorBoard)."""
+ del training
+ metrics = []
+ # Add loss items
+ for name in self._loss_items:
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
+ # TODO(longshangbang): add evaluation metrics
+ return metrics
+
+ def train_step(
+ self,
+ inputs: Tuple[NestedTensorDict, NestedTensorDict],
+ model: ModelType,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[Sequence[tf_keras.metrics.Metric]] = None
+ ) -> Dict[str, tf.Tensor]:
+ features, labels = inputs
+ input_dict = {"features": features}
+ if self.task_config.model_call_needs_labels:
+ input_dict["labels"] = labels
+
+ is_mixed_precision = isinstance(optimizer,
+ tf_keras.mixed_precision.LossScaleOptimizer)
+
+ with tf.GradientTape() as tape:
+ outputs = model(**input_dict, training=True)
+ loss, loss_dict = model.compute_losses(labels=labels, outputs=outputs)
+ loss = loss / tf.distribute.get_strategy().num_replicas_in_sync
+ if is_mixed_precision:
+ loss = optimizer.get_scaled_loss(loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(loss, tvars)
+ if is_mixed_precision:
+ grads = optimizer.get_unscaled_gradients(grads)
+
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ logs = {"loss": loss}
+ if metrics:
+ for m in metrics:
+ m.update_state(loss_dict[m.name])
+ return logs
diff --git a/official/projects/unified_detector/train.py b/official/projects/unified_detector/train.py
new file mode 100644
index 00000000000..43a43283a63
--- /dev/null
+++ b/official/projects/unified_detector/train.py
@@ -0,0 +1,70 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""TensorFlow Model Garden Vision training driver."""
+
+from absl import app
+from absl import flags
+import gin
+
+from official.common import distribute_utils
+from official.common import flags as tfm_flags
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+# pylint: disable=unused-import
+from official.projects.unified_detector import registry_imports
+# pylint: enable=unused-import
+
+FLAGS = flags.FLAGS
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ # Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
+ # dtype is float16
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu)
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir)
+
+ train_utils.save_gin_config(FLAGS.mode, model_dir)
+
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ flags.mark_flags_as_required(['experiment', 'mode', 'model_dir'])
+ app.run(main)
diff --git a/official/projects/unified_detector/utils/typing.py b/official/projects/unified_detector/utils/typing.py
new file mode 100644
index 00000000000..4fbab159741
--- /dev/null
+++ b/official/projects/unified_detector/utils/typing.py
@@ -0,0 +1,28 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Typing extension."""
+
+from typing import Dict, Union
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+NpDict = Dict[str, np.ndarray]
+FeaturesAndLabelsType = Dict[str, Dict[str, tf.Tensor]]
+TensorDict = Dict[Union[str, int], tf.Tensor]
+NestedTensorDict = Dict[
+ Union[str, int],
+ Union[tf.Tensor,
+ TensorDict]]
diff --git a/official/projects/unified_detector/utils/utilities.py b/official/projects/unified_detector/utils/utilities.py
new file mode 100644
index 00000000000..a85a2cc7cb2
--- /dev/null
+++ b/official/projects/unified_detector/utils/utilities.py
@@ -0,0 +1,235 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utility functions."""
+
+import collections
+from typing import List, Optional, Union
+
+import tensorflow as tf, tf_keras
+
+
+def resolve_shape(
+ tensor: tf.Tensor,
+ resolve_batch_size: bool = True) -> List[Union[tf.Tensor, int]]:
+ """Fully resolves the shape of the tensor.
+
+ Args:
+ tensor: The tensor for which to resolve the shape.
+ resolve_batch_size: If True, fully resolve the batch size. If False,
+ return the batch size if it is statically known and -1 otherwise. This
+ can be more efficient when converting a model to TFLite.
+
+ Returns:
+ A list containing the static dimension where possible and the dynamic
+ dimension otherwise.
+ """
+ with tf.name_scope('resolve_shape'):
+ shape = tensor.get_shape().as_list()
+ if None in shape:
+ shape_dynamic = tf.shape(tensor)
+ if shape[0] is None:
+ shape[0] = shape_dynamic[0] if resolve_batch_size else -1
+ for i in range(1, len(shape)):
+ if shape[i] is None:
+ shape[i] = shape_dynamic[i]
+ return shape
+
+
+def set_shape_dim(tensor: tf.Tensor, index: int, size: int) -> None:
+ """Set value of index-th element of tensor shape to size."""
+ shape = tensor.get_shape().as_list()
+ if len(shape) <= index:
+ raise ValueError(
+ 'Tensor rank must be at least %d. Got %d' % (index + 1, len(shape)))
+ shape[index] = size
+ tensor.set_shape(shape)
+
+
+def truncate_or_pad(input_tensor: tf.Tensor,
+ new_size: int,
+ axis: int = 1,
+ constant_value: Union[int, float] = 0) -> tf.Tensor:
+ """Truncate or zeros pad the axis of input tensor to new size."""
+ rank = len(input_tensor.shape)
+
+ if rank <= axis:
+ raise ValueError(
+ 'Tensor rank must be at least %d. Got %d' % (axis + 1, rank))
+
+ orig_size = tf.shape(input_tensor)[axis]
+
+ def _new_size(dim):
+ if dim == axis:
+ return new_size
+ n = tf.shape(input_tensor)[dim]
+ return -1 if n is None else n
+
+ def _truncate():
+ begin = [0] * rank
+ size = [_new_size(dim) for dim in range(rank)]
+ return tf.slice(input_tensor, begin, size)
+
+ def _pad():
+ padding = [[0, 0] for _ in range(rank)]
+ padding[axis][1] = new_size - orig_size
+ return tf.pad(input_tensor, padding, constant_values=constant_value)
+
+ output = tf.cond(orig_size >= new_size, _truncate, _pad)
+ if isinstance(new_size, int):
+ set_shape_dim(output, axis, new_size)
+ return output
+
+
+def rotate_rboxes90(rboxes: tf.Tensor,
+ image_width: int,
+ image_height: int,
+ rotation_count: int = 1) -> tf.Tensor:
+ """Rotate oriented rectangles counter-clockwise by multiples of 90 degrees."""
+ image_width = tf.cast(image_width, dtype=tf.float32)
+ image_height = tf.cast(image_height, dtype=tf.float32)
+
+ rotation_count = rotation_count % 4
+ x, y, w, h, angle = tf.split(rboxes, 5, axis=1)
+
+ if rotation_count == 0:
+ return rboxes
+ elif rotation_count == 1:
+ angle = tf.where(angle < -90.0, angle + 270, angle - 90)
+ return tf.concat([y, image_width - x - 1, w, h, angle], axis=1)
+ elif rotation_count == 2:
+ angle = tf.where(angle < 0.0, angle + 180, angle - 180)
+ return tf.concat([image_width - x - 1, image_height - y - 1, w, h, angle],
+ axis=1)
+ else:
+ angle = tf.where(angle > 90.0, angle - 270, angle + 90)
+ return tf.concat([image_height - y - 1, x, w, h, angle], axis=1)
+
+
+def normalize_image_to_range(image: tf.Tensor,
+ original_minval: int = 0,
+ original_maxval: int = 255,
+ target_minval: float = -1.0,
+ target_maxval: float = 1.0) -> tf.Tensor:
+ """Normalizes pixel values in the image.
+
+ Moves the pixel values from the current [original_minval, original_maxval]
+ range to the [target_minval, target_maxval] range.
+
+ Args:
+ image: A tensor of shape [height, width, channels]. Input will be converted
+ to float32 type before normalization.
+ original_minval: current image minimum value.
+ original_maxval: current image maximum value.
+ target_minval: target image minimum value.
+ target_maxval: target image maximum value.
+
+ Returns:
+ A float tensor with the same shape as the input image.
+ """
+ if image.dtype is not tf.float32:
+ image = tf.cast(image, dtype=tf.float32)
+
+ original_minval = float(original_minval) # pyrefly: ignore[bad-assignment]
+ original_maxval = float(original_maxval) # pyrefly: ignore[bad-assignment]
+ target_minval = float(target_minval)
+ target_maxval = float(target_maxval)
+ image = tf.cast(image, dtype=tf.float32)
+ image = tf.subtract(image, original_minval)
+ image = tf.multiply(image, (target_maxval - target_minval) /
+ (original_maxval - original_minval))
+ image = tf.add(image, target_minval)
+
+ return image
+
+
+def get_padding_mask_from_valid_lengths(
+ valid_lengths: tf.Tensor,
+ max_length: Optional[int] = None,
+ dtype: tf.dtypes.DType = tf.bool) -> tf.Tensor:
+ """Gets a 2D mask of the padded region from valid lengths.
+
+ Args:
+ valid_lengths: A 1D int tensor containing the valid length of each row.
+ max_length: (optional, int) The maximum length of each row. If `None`, the
+ maximum value in `valid_lengths` will be used.
+ dtype: The output dtype.
+
+ Returns:
+ 2D padded region mask.
+ """
+ with tf.name_scope('get_padding_mask_from_valid_lengths'):
+ if max_length is None:
+ max_length = tf.reduce_max(valid_lengths)
+ padding_mask = tf.logical_not(tf.sequence_mask(valid_lengths, max_length))
+
+ return tf.cast(padding_mask, dtype=dtype)
+
+
+def get_transformer_attention_bias(padding_mask: tf.Tensor) -> tf.Tensor:
+ """Gets attention bias.
+
+ Bias tensor that is added to the pre-softmax multi-headed attention logits,
+ which has shape [batch_size, num_attention_heads, max_length, max_length].
+ The tensor is zero at non-padded locations, and -1e9 (negative infinity) at
+ padded locations.
+
+ Args:
+ padding_mask: A [batch_size, max_length] float tensor, the padding mask.
+
+ Returns:
+ Attention bias tensor of shape [batch_size, 1, 1, max_length].
+ """
+ with tf.name_scope('attention_bias'):
+ # Uses -1e9 to represent -infinity. We do not actually use -Inf, since we
+ # want to be able to multiply these values by zero to get zero.
+ # (-Inf * 0 = NaN)
+ attention_bias = padding_mask * -1e9
+ attention_bias = tf.expand_dims(
+ tf.expand_dims(attention_bias, axis=1), axis=1)
+
+ return attention_bias
+
+
+class DisjointSet:
+ """A disjoint set implementation."""
+
+ def __init__(self, num_elements: int):
+ self._num_elements = num_elements
+ self._parent = list(range(num_elements))
+
+ def find(self, item: int) -> int:
+ if self._parent[item] == item:
+ return item
+ else:
+ self._parent[item] = self.find(self._parent[item])
+ return self._parent[item]
+
+ def union(self, i1: int, i2: int) -> None:
+ r1 = self.find(i1)
+ r2 = self.find(i2)
+ self._parent[r1] = r2
+
+ def to_group(self) -> List[List[int]]:
+ """Return the grouping results.
+
+ Returns:
+ A list of integer lists. Each list represents the IDs belonging to the
+ same group.
+ """
+ groups = collections.defaultdict(list)
+ for i in range(self._num_elements):
+ r = self.find(i)
+ groups[r].append(i)
+ return list(groups.values())
diff --git a/official/projects/video_ssl/README.md b/official/projects/video_ssl/README.md
index 92626164955..86b00d819cc 100644
--- a/official/projects/video_ssl/README.md
+++ b/official/projects/video_ssl/README.md
@@ -17,10 +17,7 @@ from the same short video are pulled together in the embedding space, while
clips from different videos are pushed away. CVRL significantly closes the gap
between unsupervised and supervised video representation learning.
-We release the code and pre-trained models.
-
-More pre-trained model checkpoints and a detailed instruction about the code
-will be updated.
+Here we release the code and pre-trained models.
## Experimental Results
@@ -35,7 +32,7 @@ will be updated.
## Pre-trained Model Checkpoints
We provide model checkpoints pre-trained on unlabeled RGB videos from
-Kinetics-400 and Kinetics-600. All models are trained scratch with random
+Kinetics-400 and Kinetics-600. All models are trained from scratch with random
initialization.
We also provide a baseline model checkpoint of "ImageNet inflated" we used in
diff --git a/official/projects/video_ssl/__init__.py b/official/projects/video_ssl/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/video_ssl/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/video_ssl/configs/__init__.py b/official/projects/video_ssl/configs/__init__.py
index 976989d6c84..1557829e4e0 100644
--- a/official/projects/video_ssl/configs/__init__.py
+++ b/official/projects/video_ssl/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/video_ssl/configs/video_ssl.py b/official/projects/video_ssl/configs/video_ssl.py
index 80bf506d7f7..d91be1dcbc1 100644
--- a/official/projects/video_ssl/configs/video_ssl.py
+++ b/official/projects/video_ssl/configs/video_ssl.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -28,20 +28,12 @@
VideoClassificationTask = video_classification.VideoClassificationTask
-@dataclasses.dataclass
-class VideoSSLPretrainTask(VideoClassificationTask):
- pass
-
-
-@dataclasses.dataclass
-class VideoSSLEvalTask(VideoClassificationTask):
- pass
-
-
@dataclasses.dataclass
class DataConfig(video_classification.DataConfig):
"""The base configuration for building datasets."""
is_ssl: bool = False
+ is_training: bool = True
+ drop_remainder: bool = True
@dataclasses.dataclass
@@ -51,8 +43,11 @@ class VideoSSLModel(VideoClassificationModel):
hidden_dim: int = 2048
hidden_layer_num: int = 3
projection_dim: int = 128
- hidden_norm_activation: common.NormActivation = common.NormActivation(
- use_sync_bn=False, norm_momentum=0.997, norm_epsilon=1.0e-05)
+ hidden_norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation(
+ use_sync_bn=False, norm_momentum=0.997, norm_epsilon=1.0e-05
+ )
+ )
@dataclasses.dataclass
@@ -61,17 +56,46 @@ class SSLLosses(Losses):
temperature: float = 0.1
+@dataclasses.dataclass
+class VideoSSLPretrainTask(VideoClassificationTask):
+ model: VideoSSLModel = dataclasses.field(default_factory=VideoSSLModel)
+ losses: SSLLosses = dataclasses.field(default_factory=SSLLosses)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True, drop_remainder=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ is_training=False, drop_remainder=False
+ )
+ )
+ losses: SSLLosses = dataclasses.field(default_factory=SSLLosses)
+
+
+@dataclasses.dataclass
+class VideoSSLEvalTask(VideoClassificationTask):
+ model: VideoSSLModel = dataclasses.field(default_factory=VideoSSLModel)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True, drop_remainder=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ is_training=False, drop_remainder=False
+ )
+ )
+ losses: SSLLosses = dataclasses.field(default_factory=SSLLosses)
+
+
@exp_factory.register_config_factory('video_ssl_pretrain_kinetics400')
def video_ssl_pretrain_kinetics400() -> cfg.ExperimentConfig:
"""Pretrain SSL Video classification on Kinectics 400 with resnet."""
exp = video_classification.video_classification_kinetics400()
- exp.task = VideoSSLPretrainTask(**exp.task.as_dict())
- exp.task.train_data = DataConfig(is_ssl=True, **exp.task.train_data.as_dict())
- exp.task.train_data.feature_shape = (16, 224, 224, 3)
- exp.task.train_data.temporal_stride = 2
- exp.task.model = VideoSSLModel(exp.task.model)
- exp.task.model.model_type = 'video_ssl_model'
- exp.task.losses = SSLLosses(exp.task.losses)
+ task = VideoSSLPretrainTask()
+ task.override(exp.task)
+ task.train_data.is_ssl = True
+ task.train_data.feature_shape = (16, 224, 224, 3)
+ task.train_data.temporal_stride = 2
+ task.model.model_type = 'video_ssl_model'
+ exp.task = task
return exp
@@ -79,23 +103,22 @@ def video_ssl_pretrain_kinetics400() -> cfg.ExperimentConfig:
def video_ssl_linear_eval_kinetics400() -> cfg.ExperimentConfig:
"""Pretrain SSL Video classification on Kinectics 400 with resnet."""
exp = video_classification.video_classification_kinetics400()
- exp.task = VideoSSLEvalTask(**exp.task.as_dict())
- exp.task.train_data = DataConfig(is_ssl=False,
- **exp.task.train_data.as_dict())
- exp.task.train_data.feature_shape = (32, 224, 224, 3)
- exp.task.train_data.temporal_stride = 2
- exp.task.validation_data.feature_shape = (32, 256, 256, 3)
- exp.task.validation_data.temporal_stride = 2
- exp.task.validation_data = DataConfig(is_ssl=False,
- **exp.task.validation_data.as_dict())
- exp.task.validation_data.min_image_size = 256
- exp.task.validation_data.num_test_clips = 10
- exp.task.validation_data.num_test_crops = 3
- exp.task.model = VideoSSLModel(exp.task.model)
- exp.task.model.model_type = 'video_ssl_model'
- exp.task.model.normalize_feature = True
- exp.task.model.hidden_layer_num = 0
- exp.task.model.projection_dim = 400
+ task = VideoSSLEvalTask() # Replaces the task type.
+ task.override(exp.task)
+ task.train_data.is_ssl = False
+ task.train_data.feature_shape = (32, 224, 224, 3)
+ task.train_data.temporal_stride = 2
+ task.validation_data.is_ssl = False
+ task.validation_data.feature_shape = (32, 256, 256, 3)
+ task.validation_data.temporal_stride = 2
+ task.validation_data.min_image_size = 256
+ task.validation_data.num_test_clips = 10
+ task.validation_data.num_test_crops = 3
+ task.model.model_type = 'video_ssl_model'
+ task.model.normalize_feature = True
+ task.model.hidden_layer_num = 0
+ task.model.projection_dim = 600
+ exp.task = task
return exp
@@ -103,13 +126,13 @@ def video_ssl_linear_eval_kinetics400() -> cfg.ExperimentConfig:
def video_ssl_pretrain_kinetics600() -> cfg.ExperimentConfig:
"""Pretrain SSL Video classification on Kinectics 400 with resnet."""
exp = video_classification.video_classification_kinetics600()
- exp.task = VideoSSLPretrainTask(**exp.task.as_dict())
- exp.task.train_data = DataConfig(is_ssl=True, **exp.task.train_data.as_dict())
- exp.task.train_data.feature_shape = (16, 224, 224, 3)
- exp.task.train_data.temporal_stride = 2
- exp.task.model = VideoSSLModel(exp.task.model)
- exp.task.model.model_type = 'video_ssl_model'
- exp.task.losses = SSLLosses(exp.task.losses)
+ task = VideoSSLPretrainTask()
+ task.override(exp.task)
+ task.train_data.is_ssl = True
+ task.train_data.feature_shape = (16, 224, 224, 3)
+ task.train_data.temporal_stride = 2
+ task.model.model_type = 'video_ssl_model'
+ exp.task = task
return exp
@@ -117,21 +140,20 @@ def video_ssl_pretrain_kinetics600() -> cfg.ExperimentConfig:
def video_ssl_linear_eval_kinetics600() -> cfg.ExperimentConfig:
"""Pretrain SSL Video classification on Kinectics 400 with resnet."""
exp = video_classification.video_classification_kinetics600()
- exp.task = VideoSSLEvalTask(**exp.task.as_dict())
- exp.task.train_data = DataConfig(is_ssl=False,
- **exp.task.train_data.as_dict())
- exp.task.train_data.feature_shape = (32, 224, 224, 3)
- exp.task.train_data.temporal_stride = 2
- exp.task.validation_data = DataConfig(is_ssl=False,
- **exp.task.validation_data.as_dict())
- exp.task.validation_data.feature_shape = (32, 256, 256, 3)
- exp.task.validation_data.temporal_stride = 2
- exp.task.validation_data.min_image_size = 256
- exp.task.validation_data.num_test_clips = 10
- exp.task.validation_data.num_test_crops = 3
- exp.task.model = VideoSSLModel(exp.task.model)
- exp.task.model.model_type = 'video_ssl_model'
- exp.task.model.normalize_feature = True
- exp.task.model.hidden_layer_num = 0
- exp.task.model.projection_dim = 600
+ task = VideoSSLEvalTask() # Replaces the task type.
+ task.override(exp.task)
+ task.train_data.is_ssl = False
+ task.train_data.feature_shape = (32, 224, 224, 3)
+ task.train_data.temporal_stride = 2
+ task.validation_data.is_ssl = False
+ task.validation_data.feature_shape = (32, 256, 256, 3)
+ task.validation_data.temporal_stride = 2
+ task.validation_data.min_image_size = 256
+ task.validation_data.num_test_clips = 10
+ task.validation_data.num_test_crops = 3
+ task.model.model_type = 'video_ssl_model'
+ task.model.normalize_feature = True
+ task.model.hidden_layer_num = 0
+ task.model.projection_dim = 600
+ exp.task = task
return exp
diff --git a/official/projects/video_ssl/configs/video_ssl_test.py b/official/projects/video_ssl/configs/video_ssl_test.py
index 3b11ddec130..5f68646abea 100644
--- a/official/projects/video_ssl/configs/video_ssl_test.py
+++ b/official/projects/video_ssl/configs/video_ssl_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
diff --git a/official/projects/video_ssl/dataloaders/__init__.py b/official/projects/video_ssl/dataloaders/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/video_ssl/dataloaders/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/video_ssl/dataloaders/video_ssl_input.py b/official/projects/video_ssl/dataloaders/video_ssl_input.py
index 046fe048dce..7a846819524 100644
--- a/official/projects/video_ssl/dataloaders/video_ssl_input.py
+++ b/official/projects/video_ssl/dataloaders/video_ssl_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from typing import Dict, Optional, Tuple
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.video_ssl.configs import video_ssl as exp_cfg
from official.projects.video_ssl.ops import video_ssl_preprocess_ops
from official.vision.dataloaders import video_input
@@ -120,17 +120,20 @@ def _process_image(image: tf.Tensor,
# Cast the frames in float32, normalizing according to zero_centering_image.
if is_training and is_ssl:
- image_1 = preprocess_ops_3d.normalize_image(image_1, zero_centering_image)
- image_2 = preprocess_ops_3d.normalize_image(image_2, zero_centering_image)
+ image_1 = preprocess_ops_3d.normalize_image(image_1, zero_centering_image) # pyrefly: ignore[unbound-name]
+ image_2 = preprocess_ops_3d.normalize_image(image_2, zero_centering_image) # pyrefly: ignore[unbound-name]
else:
image = preprocess_ops_3d.normalize_image(image, zero_centering_image)
# Self-supervised pre-training augmentations.
if is_training and is_ssl:
+ if zero_centering_image:
+ image_1 = 0.5 * (image_1 + 1.0) # pyrefly: ignore[unbound-name]
+ image_2 = 0.5 * (image_2 + 1.0) # pyrefly: ignore[unbound-name]
# Temporally consistent color jittering.
- image_1 = video_ssl_preprocess_ops.random_color_jitter_3d(image_1)
- image_2 = video_ssl_preprocess_ops.random_color_jitter_3d(image_2)
+ image_1 = video_ssl_preprocess_ops.random_color_jitter_3d(image_1) # pyrefly: ignore[unbound-name]
+ image_2 = video_ssl_preprocess_ops.random_color_jitter_3d(image_2) # pyrefly: ignore[unbound-name]
# Temporally consistent gaussian blurring.
image_1 = video_ssl_preprocess_ops.random_blur(image_1, crop_size,
crop_size, 1.0)
@@ -139,6 +142,8 @@ def _process_image(image: tf.Tensor,
image_2 = video_ssl_preprocess_ops.random_solarization(image_2)
image = tf.concat([image_1, image_2], axis=0)
image = tf.clip_by_value(image, 0., 1.)
+ if zero_centering_image:
+ image = 2 * (image - 0.5)
return image
@@ -216,7 +221,7 @@ def __init__(self,
input_params: exp_cfg.DataConfig,
image_key: str = IMAGE_KEY,
label_key: str = LABEL_KEY):
- super(Parser, self).__init__(input_params, image_key, label_key)
+ super().__init__(input_params, image_key, label_key)
self._is_ssl = input_params.is_ssl
def _parse_train_data(
@@ -233,7 +238,8 @@ def _parse_train_data(
stride=self._stride,
num_test_clips=self._num_test_clips,
min_resize=self._min_resize,
- crop_size=self._crop_size)
+ crop_size=self._crop_size,
+ zero_centering_image=self._zero_centering_image)
image = tf.cast(image, dtype=self._dtype)
features = {'image': image}
@@ -255,7 +261,8 @@ def _parse_eval_data(
num_test_clips=self._num_test_clips,
min_resize=self._min_resize,
crop_size=self._crop_size,
- num_crops=self._num_crops)
+ num_crops=self._num_crops,
+ zero_centering_image=self._zero_centering_image)
image = tf.cast(image, dtype=self._dtype)
features = {'image': image}
diff --git a/official/projects/video_ssl/dataloaders/video_ssl_input_test.py b/official/projects/video_ssl/dataloaders/video_ssl_input_test.py
index 951f5bd0ec3..5393380ec7c 100644
--- a/official/projects/video_ssl/dataloaders/video_ssl_input_test.py
+++ b/official/projects/video_ssl/dataloaders/video_ssl_input_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,9 @@
import io
-# Import libraries
import numpy as np
from PIL import Image
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.video_ssl.configs import video_ssl as exp_cfg
from official.projects.video_ssl.dataloaders import video_ssl_input
diff --git a/official/projects/video_ssl/losses/__init__.py b/official/projects/video_ssl/losses/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/video_ssl/losses/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/video_ssl/losses/losses.py b/official/projects/video_ssl/losses/losses.py
index 2aa2085b80e..c1c881d3506 100644
--- a/official/projects/video_ssl/losses/losses.py
+++ b/official/projects/video_ssl/losses/losses.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,8 +14,7 @@
"""Define losses."""
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.compiler.tf2xla.python import xla
diff --git a/official/projects/video_ssl/modeling/__init__.py b/official/projects/video_ssl/modeling/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/video_ssl/modeling/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/video_ssl/modeling/video_ssl_model.py b/official/projects/video_ssl/modeling/video_ssl_model.py
index d782bf26361..4ae2e5509eb 100644
--- a/official/projects/video_ssl/modeling/video_ssl_model.py
+++ b/official/projects/video_ssl/modeling/video_ssl_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,19 +15,17 @@
"""Build video classification models."""
from typing import Mapping, Optional
-# Import libraries
-
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.projects.video_ssl.configs import video_ssl as video_ssl_cfg
from official.vision.modeling import backbones
from official.vision.modeling import factory_3d as model_factory
-layers = tf.keras.layers
+layers = tf_keras.layers
-class VideoSSLModel(tf.keras.Model):
+class VideoSSLModel(tf_keras.Model):
"""A video ssl model class builder."""
def __init__(self,
@@ -38,7 +36,7 @@ def __init__(self,
hidden_norm_args,
projection_dim,
input_specs: Optional[Mapping[str,
- tf.keras.layers.InputSpec]] = None,
+ tf_keras.layers.InputSpec]] = None,
dropout_rate: float = 0.0,
aggregate_endpoints: bool = False,
kernel_initializer='random_uniform',
@@ -53,15 +51,15 @@ def __init__(self,
hidden_dim: `int` number of hidden units in MLP.
hidden_layer_num: `int` number of hidden layers in MLP.
hidden_norm_args: `dict` for batchnorm arguments in MLP.
- projection_dim: `int` number of ouput dimension for MLP.
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ projection_dim: `int` number of output dimension for MLP.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
dropout_rate: `float` rate for dropout regularization.
aggregate_endpoints: `bool` aggregate all end ponits or only use the
final end point.
kernel_initializer: kernel initializer for the dense layer.
- kernel_regularizer: tf.keras.regularizers.Regularizer object. Default to
+ kernel_regularizer: tf_keras.regularizers.Regularizer object. Default to
None.
- bias_regularizer: tf.keras.regularizers.Regularizer object. Default to
+ bias_regularizer: tf_keras.regularizers.Regularizer object. Default to
None.
**kwargs: keyword arguments to be passed.
"""
@@ -93,19 +91,19 @@ def __init__(self,
self._backbone = backbone
inputs = {
- k: tf.keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
+ k: tf_keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
}
endpoints = backbone(inputs['image'])
if aggregate_endpoints:
pooled_feats = []
for endpoint in endpoints.values():
- x_pool = tf.keras.layers.GlobalAveragePooling3D()(endpoint)
+ x_pool = tf_keras.layers.GlobalAveragePooling3D()(endpoint)
pooled_feats.append(x_pool)
x = tf.concat(pooled_feats, axis=1)
else:
x = endpoints[max(endpoints.keys())]
- x = tf.keras.layers.GlobalAveragePooling3D()(x)
+ x = tf_keras.layers.GlobalAveragePooling3D()(x)
# L2 Normalize feature after backbone
if normalize_feature:
@@ -113,22 +111,21 @@ def __init__(self,
# MLP hidden layers
for _ in range(hidden_layer_num):
- x = tf.keras.layers.Dense(hidden_dim)(x)
+ x = tf_keras.layers.Dense(hidden_dim)(x)
if self._config_dict['use_sync_bn']:
- x = tf.keras.layers.experimental.SyncBatchNormalization(
+ x = tf_keras.layers.experimental.SyncBatchNormalization(
momentum=self._config_dict['norm_momentum'],
epsilon=self._config_dict['norm_epsilon'])(x)
else:
- x = tf.keras.layers.BatchNormalization(
+ x = tf_keras.layers.BatchNormalization(
momentum=self._config_dict['norm_momentum'],
epsilon=self._config_dict['norm_epsilon'])(x)
x = tf_utils.get_activation(self._config_dict['activation'])(x)
# Projection head
- x = tf.keras.layers.Dense(projection_dim)(x)
+ x = tf_keras.layers.Dense(projection_dim)(x)
- super(VideoSSLModel, self).__init__(
- inputs=inputs, outputs=x, **kwargs)
+ super().__init__(inputs=inputs, outputs=x, **kwargs)
@property
def checkpoint_items(self):
@@ -149,10 +146,10 @@ def from_config(cls, config, custom_objects=None):
@model_factory.register_model_builder('video_ssl_model')
def build_video_ssl_pretrain_model(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
model_config: video_ssl_cfg.VideoSSLModel,
num_classes: int,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None):
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None):
"""Builds the video classification model."""
del num_classes
input_specs_dict = {'image': input_specs}
diff --git a/official/projects/video_ssl/ops/__init__.py b/official/projects/video_ssl/ops/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/video_ssl/ops/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/video_ssl/ops/video_ssl_preprocess_ops.py b/official/projects/video_ssl/ops/video_ssl_preprocess_ops.py
index f6a2ef3aa31..b7c8d30d05a 100644
--- a/official/projects/video_ssl/ops/video_ssl_preprocess_ops.py
+++ b/official/projects/video_ssl/ops/video_ssl_preprocess_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import functools
from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def random_apply(func, p, x):
@@ -399,7 +399,7 @@ def cdf(k, power=1.0):
offset=offset_2)
indices = tf.concat([indices_1, indices_2], axis=0)
- indices.set_shape((num_windows * num_steps,))
+ indices.set_shape((num_windows * num_steps,)) # pyrefly: ignore[unsupported-operation]
output = tf.gather(sequence, indices)
return output
diff --git a/official/projects/video_ssl/ops/video_ssl_preprocess_ops_test.py b/official/projects/video_ssl/ops/video_ssl_preprocess_ops_test.py
index 7e1b61465a9..50040e860e3 100644
--- a/official/projects/video_ssl/ops/video_ssl_preprocess_ops_test.py
+++ b/official/projects/video_ssl/ops/video_ssl_preprocess_ops_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -12,7 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.video_ssl.ops import video_ssl_preprocess_ops
from official.vision.ops import preprocess_ops_3d
diff --git a/official/projects/video_ssl/tasks/__init__.py b/official/projects/video_ssl/tasks/__init__.py
index cb709251258..3e5bc4cc0b8 100644
--- a/official/projects/video_ssl/tasks/__init__.py
+++ b/official/projects/video_ssl/tasks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/video_ssl/tasks/linear_eval.py b/official/projects/video_ssl/tasks/linear_eval.py
index 5d7849422c7..2bdafa6215d 100644
--- a/official/projects/video_ssl/tasks/linear_eval.py
+++ b/official/projects/video_ssl/tasks/linear_eval.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Video ssl linear evaluation task definition."""
from typing import Any, Optional, List, Tuple
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.core import task_factory
@@ -28,7 +28,7 @@
class VideoSSLEvalTask(video_classification.VideoClassificationTask):
"""A task for video ssl linear evaluation."""
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loading pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -49,8 +49,8 @@ def initialize(self, model: tf.keras.Model):
def train_step(self,
inputs: Tuple[Any, Any],
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
metrics: Optional[List[Any]] = None):
"""Does forward and backward.
@@ -66,5 +66,4 @@ def train_step(self,
model.backbone.trainable = False
logging.info('Setting the backbone to non-trainable.')
- return super(video_classification.VideoClassificationTask,
- self).train_step(inputs, model, optimizer, metrics)
+ return super().train_step(inputs, model, optimizer, metrics)
diff --git a/official/projects/video_ssl/tasks/pretrain.py b/official/projects/video_ssl/tasks/pretrain.py
index f58db11ce58..97ed251a44b 100644
--- a/official/projects/video_ssl/tasks/pretrain.py
+++ b/official/projects/video_ssl/tasks/pretrain.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Video ssl pretrain task definition."""
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.core import input_reader
@@ -39,7 +39,7 @@ def build_model(self):
for d1, d2 in zip(self.task_config.train_data.feature_shape,
self.task_config.validation_data.feature_shape)
]
- input_specs = tf.keras.layers.InputSpec(shape=[None] + common_input_shape)
+ input_specs = tf_keras.layers.InputSpec(shape=[None] + common_input_shape)
logging.info('Build model input %r', common_input_shape)
model = factory_3d.build_model(
@@ -104,9 +104,9 @@ def build_losses(self, model_outputs, num_replicas, model):
def build_metrics(self, training=True):
"""Gets streaming metrics for training/validation."""
metrics = [
- tf.keras.metrics.Mean(name='contrast_acc'),
- tf.keras.metrics.Mean(name='contrast_entropy'),
- tf.keras.metrics.Mean(name='reg_loss')
+ tf_keras.metrics.Mean(name='contrast_acc'),
+ tf_keras.metrics.Mean(name='contrast_entropy'),
+ tf_keras.metrics.Mean(name='reg_loss')
]
return metrics
@@ -151,14 +151,14 @@ def train_step(self, inputs, model, optimizer, metrics=None):
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
if isinstance(
- optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
# Scales back gradient before apply_gradients when LossScaleOptimizer is
# used.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
diff --git a/official/projects/video_ssl/tasks/pretrain_test.py b/official/projects/video_ssl/tasks/pretrain_test.py
index 5f0bdbbb38f..ba690e6c707 100644
--- a/official/projects/video_ssl/tasks/pretrain_test.py
+++ b/official/projects/video_ssl/tasks/pretrain_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import random
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official import vision
@@ -33,7 +33,7 @@
class VideoClassificationTaskTest(tf.test.TestCase):
def setUp(self):
- super(VideoClassificationTaskTest, self).setUp()
+ super().setUp()
data_dir = os.path.join(self.get_temp_dir(), 'data')
tf.io.gfile.makedirs(data_dir)
self._data_path = os.path.join(data_dir, 'data.tfrecord')
diff --git a/official/projects/video_ssl/train.py b/official/projects/video_ssl/train.py
index 5d1f4e8f58f..1551e965ee2 100644
--- a/official/projects/video_ssl/train.py
+++ b/official/projects/video_ssl/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/video_ssl/video_ssl.ipynb b/official/projects/video_ssl/video_ssl.ipynb
new file mode 100644
index 00000000000..5e9574c9616
--- /dev/null
+++ b/official/projects/video_ssl/video_ssl.ipynb
@@ -0,0 +1,968 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "nS7qzvG6-TfI"
+ },
+ "source": [
+ "# Video SSL Tutorial\n",
+ "\n",
+ "This tutorial trains a video_ssl with 3D-ResNet-50 (R3D-50) as backbone model from the TensorFlow Model Garden package (tensorflow-models).\n",
+ "\n",
+ "[Model Garden](https://www.tensorflow.org/tfmodels) contains a collection of state-of-the-art models, implemented with TensorFlow's high-level APIs. The implementations demonstrate the best practices for modeling, letting users to take full advantage of TensorFlow for their research and product development.\n",
+ "\n",
+ "**Dataset:** [UCF_101](https://www.tensorflow.org/datasets/catalog/ucf101)\n",
+ "* A 101-label video classification dataset\n",
+ "\n",
+ "**This tutorial demonstrates how to:**\n",
+ "\n",
+ "* Use models from the TensorFlow Models package.\n",
+ "* Train/Fine-tune a pre-built [video_ssl](https://arxiv.org/abs/2008.03800) for Video Classification.\n",
+ "* Export the trained/tuned video_ssl model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "HYrARPmK-OZI"
+ },
+ "source": [
+ "## Install cecessary libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XRUierWnXulP"
+ },
+ "outputs": [],
+ "source": [
+ "!pip install -U -q \"tensorflow\" \"tensorflow_addons\" \"immutabledict\" \"tensorflow_datasets\"\n",
+ "!pip install -U -q remotezip tqdm opencv-python einops\n",
+ "!pip install -U -q git+https://github.com/tensorflow/docs"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "K5W_31qC-GgY"
+ },
+ "source": [
+ "## Clone models repository"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WJWNaEAXdgI4"
+ },
+ "outputs": [],
+ "source": [
+ "!git clone https://github.com/tensorflow/models.git"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "nPbayvv5Pyg-"
+ },
+ "source": [
+ "### Set models as current working dir"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "g8hcYIlicrEn"
+ },
+ "outputs": [],
+ "source": [
+ "%cd ./models/"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "DnF43pn0-ARp"
+ },
+ "source": [
+ "## Import cecessary libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "or0N8bbktR03"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "import tqdm\n",
+ "import random\n",
+ "import pathlib\n",
+ "import pprint\n",
+ "import itertools\n",
+ "import imageio\n",
+ "import collections\n",
+ "\n",
+ "import cv2\n",
+ "import einops\n",
+ "import numpy as np\n",
+ "import remotezip as rz\n",
+ "import seaborn as sns\n",
+ "import matplotlib.pyplot as plt\n",
+ "import tensorflow as tf\n",
+ "\n",
+ "from IPython import display\n",
+ "from tensorflow_docs.vis import embed\n",
+ "\n",
+ "pp = pprint.PrettyPrinter(indent=4) # Set Pretty Print Indentation\n",
+ "print(tf.__version__) # Check the version of tensorflow used"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "r1hmO0Zxkz_Q"
+ },
+ "source": [
+ "## Import required modules from vidoe_ssl for running the model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "EopVEolTkkpd"
+ },
+ "outputs": [],
+ "source": [
+ "from official.core import task_factory\n",
+ "from official.core import train_lib\n",
+ "from official.core import train_utils\n",
+ "from official.modeling import performance\n",
+ "from official.vision.data import tfrecord_lib\n",
+ "\n",
+ "from official.projects.video_ssl.configs import video_ssl as exp_cfg\n",
+ "from official.projects.video_ssl.modeling import video_ssl_model\n",
+ "from official.projects.video_ssl.tasks import linear_eval\n",
+ "from official.projects.video_ssl.tasks import pretrain\n",
+ "from official.vision import registry_imports\n",
+ "from official.vision.serving import export_saved_model_lib"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "sIS-qexL968P"
+ },
+ "source": [
+ "## Download pretrained weights"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "10t-SR6XAXgO"
+ },
+ "outputs": [],
+ "source": [
+ "!wget https://storage.googleapis.com/tf_model_garden/vision/cvrl/r3d_1x_k600_800ep.tar.gz -P /content/\n",
+ "!tar -xvf /content/r3d_1x_k600_800ep.tar.gz -C ../\n",
+ "!rm ../r3d_1x_k600_800ep.tar.gz"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "xXix2TCo-yAk"
+ },
+ "source": [
+ "## Download Subdataset of UCF_101"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "SAE7S0ZCcHOT"
+ },
+ "outputs": [],
+ "source": [
+ "# @title Load and preprocess video data\n",
+ "def list_files_per_class(zip_url):\n",
+ " \"\"\"\n",
+ " List the files in each class of the dataset given the zip URL.\n",
+ "\n",
+ " Args:\n",
+ " zip_url: URL from which the files can be unzipped.\n",
+ "\n",
+ " Return:\n",
+ " files: List of files in each of the classes.\n",
+ " \"\"\"\n",
+ " files = []\n",
+ " with rz.RemoteZip(URL) as zip:\n",
+ " for zip_info in zip.infolist():\n",
+ " files.append(zip_info.filename)\n",
+ " return files\n",
+ "\n",
+ "def get_class(fname):\n",
+ " \"\"\"\n",
+ " Retrieve the name of the class given a filename.\n",
+ "\n",
+ " Args:\n",
+ " fname: Name of the file in the UCF101 dataset.\n",
+ "\n",
+ " Return:\n",
+ " Class that the file belongs to.\n",
+ " \"\"\"\n",
+ " return fname.split('_')[-3]\n",
+ "\n",
+ "def get_files_per_class(files):\n",
+ " \"\"\"\n",
+ " Retrieve the files that belong to each class.\n",
+ "\n",
+ " Args:\n",
+ " files: List of files in the dataset.\n",
+ "\n",
+ " Return:\n",
+ " Dictionary of class names (key) and files (values).\n",
+ " \"\"\"\n",
+ " files_for_class = collections.defaultdict(list)\n",
+ " for fname in files:\n",
+ " class_name = get_class(fname)\n",
+ " files_for_class[class_name].append(fname)\n",
+ " return files_for_class\n",
+ "\n",
+ "def download_from_zip(zip_url, to_dir, file_names):\n",
+ " \"\"\"\n",
+ " Download the contents of the zip file from the zip URL.\n",
+ "\n",
+ " Args:\n",
+ " zip_url: Zip URL containing data.\n",
+ " to_dir: Directory to download data to.\n",
+ " file_names: Names of files to download.\n",
+ " \"\"\"\n",
+ " with rz.RemoteZip(zip_url) as zip:\n",
+ " for fn in tqdm.tqdm(file_names):\n",
+ " class_name = get_class(fn)\n",
+ " zip.extract(fn, str(to_dir / class_name))\n",
+ " unzipped_file = to_dir / class_name / fn\n",
+ "\n",
+ " fn = pathlib.Path(fn).parts[-1]\n",
+ " output_file = to_dir / class_name / fn\n",
+ " unzipped_file.rename(output_file,)\n",
+ "\n",
+ "def split_class_lists(files_for_class, count):\n",
+ " \"\"\"\n",
+ " Returns the list of files belonging to a subset of data as well as the remainder of\n",
+ " files that need to be downloaded.\n",
+ "\n",
+ " Args:\n",
+ " files_for_class: Files belonging to a particular class of data.\n",
+ " count: Number of files to download.\n",
+ "\n",
+ " Return:\n",
+ " split_files: Files belonging to the subset of data.\n",
+ " remainder: Dictionary of the remainder of files that need to be downloaded.\n",
+ " \"\"\"\n",
+ " split_files = []\n",
+ " remainder = {}\n",
+ " for cls in files_for_class:\n",
+ " split_files.extend(files_for_class[cls][:count])\n",
+ " remainder[cls] = files_for_class[cls][count:]\n",
+ " return split_files, remainder\n",
+ "\n",
+ "def download_ufc_101_subset(zip_url, num_classes, splits, download_dir):\n",
+ " \"\"\"\n",
+ " Download a subset of the UFC101 dataset and split them into various parts, such as\n",
+ " training, validation, and test.\n",
+ "\n",
+ " Args:\n",
+ " zip_url: Zip URL containing data.\n",
+ " num_classes: Number of labels.\n",
+ " splits: Dictionary specifying the training, validation, test, etc. (key) division of data\n",
+ " (value is number of files per split).\n",
+ " download_dir: Directory to download data to.\n",
+ "\n",
+ " Return:\n",
+ " dir: Posix path of the resulting directories containing the splits of data.\n",
+ " \"\"\"\n",
+ " files = list_files_per_class(zip_url)\n",
+ " for f in files:\n",
+ " tokens = f.split('/')\n",
+ " if len(tokens) \u003c= 2:\n",
+ " files.remove(f) # Remove that item from the list if it does not have a filename\n",
+ "\n",
+ " files_for_class = get_files_per_class(files)\n",
+ "\n",
+ " classes = list(files_for_class.keys())[:num_classes]\n",
+ "\n",
+ " for cls in classes:\n",
+ " new_files_for_class = files_for_class[cls]\n",
+ " random.shuffle(new_files_for_class)\n",
+ " files_for_class[cls] = new_files_for_class\n",
+ "\n",
+ " # Only use the number of classes you want in the dictionary\n",
+ " files_for_class = {x: files_for_class[x] for x in list(files_for_class)[:num_classes]}\n",
+ "\n",
+ " dirs = {}\n",
+ " for split_name, split_count in splits.items():\n",
+ " print(split_name, \":\")\n",
+ " split_dir = download_dir / split_name\n",
+ " split_files, files_for_class = split_class_lists(files_for_class, split_count)\n",
+ " download_from_zip(zip_url, split_dir, split_files)\n",
+ " dirs[split_name] = split_dir\n",
+ "\n",
+ " return dirs\n",
+ "\n",
+ "def format_frames(frame, output_size):\n",
+ " \"\"\"\n",
+ " Pad and resize an image from a video.\n",
+ "\n",
+ " Args:\n",
+ " frame: Image that needs to resized and padded.\n",
+ " output_size: Pixel size of the output frame image.\n",
+ "\n",
+ " Return:\n",
+ " Formatted frame with padding of specified output size.\n",
+ " \"\"\"\n",
+ " frame = tf.image.convert_image_dtype(frame, tf.uint8)\n",
+ " frame = tf.image.resize_with_pad(frame, *output_size)\n",
+ " return frame\n",
+ "\n",
+ "def frames_from_video_file(video_path, n_frames, output_size = (224,224), frame_step = 15):\n",
+ " \"\"\"\n",
+ " Creates frames from each video file present for each category.\n",
+ "\n",
+ " Args:\n",
+ " video_path: File path to the video.\n",
+ " n_frames: Number of frames to be created per video file.\n",
+ " output_size: Pixel size of the output frame image.\n",
+ "\n",
+ " Return:\n",
+ " An NumPy array of frames in the shape of (n_frames, height, width, channels).\n",
+ " \"\"\"\n",
+ " # Read each video frame by frame\n",
+ " result = []\n",
+ " src = cv2.VideoCapture(str(video_path))\n",
+ "\n",
+ " video_length = src.get(cv2.CAP_PROP_FRAME_COUNT)\n",
+ "\n",
+ " need_length = 1 + (n_frames - 1) * frame_step\n",
+ "\n",
+ " if need_length \u003e video_length:\n",
+ " start = 0\n",
+ " else:\n",
+ " max_start = video_length - need_length\n",
+ " start = random.randint(0, max_start + 1)\n",
+ "\n",
+ " src.set(cv2.CAP_PROP_POS_FRAMES, start)\n",
+ " # ret is a boolean indicating whether read was successful, frame is the image itself\n",
+ " ret, frame = src.read()\n",
+ " result.append(format_frames(frame, output_size))\n",
+ "\n",
+ " for _ in range(n_frames - 1):\n",
+ " for _ in range(frame_step):\n",
+ " ret, frame = src.read()\n",
+ " if ret:\n",
+ " frame = format_frames(frame, output_size)\n",
+ " result.append(frame)\n",
+ " else:\n",
+ " result.append(np.zeros_like(result[0]))\n",
+ " src.release()\n",
+ " result = np.array(result)[..., [2, 1, 0]]\n",
+ "\n",
+ " return result\n",
+ "\n",
+ "def to_gif(images):\n",
+ " imageio.mimsave('./animation.gif', images, fps=10)\n",
+ " return embed.embed_file('./animation.gif')\n",
+ "\n",
+ "class FrameGenerator:\n",
+ " def __init__(self, path, n_frames, training = False):\n",
+ " \"\"\" Returns a set of frames with their associated label.\n",
+ "\n",
+ " Args:\n",
+ " path: Video file paths.\n",
+ " n_frames: Number of frames.\n",
+ " training: Boolean to determine if training dataset is being created.\n",
+ " \"\"\"\n",
+ " self.path = path\n",
+ " self.n_frames = n_frames\n",
+ " self.training = training\n",
+ " self.class_names = sorted(set(p.name for p in self.path.iterdir() if p.is_dir()))\n",
+ " self.class_ids_for_name = dict((name, idx) for idx, name in enumerate(self.class_names))\n",
+ "\n",
+ " def get_files_and_class_names(self):\n",
+ " video_paths = list(self.path.glob('*/*.avi'))\n",
+ " classes = [p.parent.name for p in video_paths]\n",
+ " return video_paths, classes\n",
+ "\n",
+ " def __call__(self):\n",
+ " video_paths, classes = self.get_files_and_class_names()\n",
+ "\n",
+ " pairs = list(zip(video_paths, classes))\n",
+ "\n",
+ " if self.training:\n",
+ " random.shuffle(pairs)\n",
+ "\n",
+ " for path, name in pairs:\n",
+ " video_frames = frames_from_video_file(path, self.n_frames)\n",
+ " label = self.class_ids_for_name[name] # Encode labels\n",
+ " yield video_frames, label"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nnn5gt5rKyqu"
+ },
+ "outputs": [],
+ "source": [
+ "# Helper functions below used are taken from following tutorials\n",
+ "# https://www.tensorflow.org/tutorials/video/video_classification\n",
+ "# https://www.tensorflow.org/tutorials/video/transfer_learning_with_movinet\n",
+ "\n",
+ "URL = 'https://storage.googleapis.com/thumos14_files/UCF101_videos.zip'\n",
+ "download_dir = pathlib.Path('../UCF101_subset/')\n",
+ "subset_paths = download_ufc_101_subset(URL,\n",
+ " num_classes = 10,\n",
+ " splits = {\"train\": 40, \"val\": 10, \"test\": 10},\n",
+ " download_dir = download_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0ah-b7zZ-1zr"
+ },
+ "source": [
+ "## Prepare train, valid and test dataset"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Bvnz1T_CgxuT"
+ },
+ "outputs": [],
+ "source": [
+ "n_frames = 10\n",
+ "CLASSES = sorted(os.listdir('../UCF101_subset/train/'))\n",
+ "\n",
+ "output_signature = (tf.TensorSpec(shape = (None, None, None, 3), dtype = tf.uint8, name='image'),\n",
+ " tf.TensorSpec(shape = (), dtype = tf.int16, name='label'))\n",
+ "\n",
+ "train_ds = tf.data.Dataset.from_generator(FrameGenerator(subset_paths['train'], n_frames, training=True),\n",
+ " output_signature = output_signature)\n",
+ "\n",
+ "val_ds = tf.data.Dataset.from_generator(FrameGenerator(subset_paths['val'], n_frames),\n",
+ " output_signature = output_signature)\n",
+ "\n",
+ "test_ds = tf.data.Dataset.from_generator(FrameGenerator(subset_paths['test'], n_frames),\n",
+ " output_signature = output_signature)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "uRjtR3N__WDi"
+ },
+ "source": [
+ "## Write data as TFRecords"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "9FeJjc7J_erR"
+ },
+ "source": [
+ "### Helper function to convert data as TF Sequence Example"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "thYjfktWj5Bh"
+ },
+ "outputs": [],
+ "source": [
+ "def process_record(record):\n",
+ " \"\"\"\n",
+ " Convert training samples to SequenceExample format. For more detailed\n",
+ " explaination about SequenceExample, please check here\n",
+ " https://www.tensorflow.org/api_docs/python/tf/train/SequenceExample\n",
+ "\n",
+ " Args:\n",
+ " record: training example with image frames and corresponding label.\n",
+ "\n",
+ " Return:\n",
+ " Return a SequenceExample which represents a\n",
+ " sequence of features and some context.\n",
+ " \"\"\"\n",
+ " seq_example = tf.train.SequenceExample()\n",
+ " for example in record[0]:\n",
+ " seq_example.feature_lists.feature_list.get_or_create(\n",
+ " 'image/encoded').feature.add().bytes_list.value[:] = [\n",
+ " tf.io.encode_jpeg(example).numpy()\n",
+ " ]\n",
+ " seq_example.context.feature[\n",
+ " 'clip/label/index'].int64_list.value[:] = [record[1].numpy()]\n",
+ " return seq_example"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "i-9fsPOQp9Qy"
+ },
+ "outputs": [],
+ "source": [
+ "output_dir = '../ucf101_tfrecords/'\n",
+ "LOG_EVERY = 100\n",
+ "if not os.path.exists(output_dir):\n",
+ " os.mkdir(output_dir)\n",
+ "\n",
+ "\n",
+ "def write_tfrecords(dataset, output_path, num_shards=1):\n",
+ " \"\"\"\n",
+ " Convert training samples to tfrecords\n",
+ "\n",
+ " Args:\n",
+ " dataset: Dataset as a iterator in (tfds format).\n",
+ " output_path: Directory to store the tfrecords.\n",
+ " num_shards: Split the tfrecords to sepecific number of shards.\n",
+ " \"\"\"\n",
+ "\n",
+ " writers = [\n",
+ " tf.io.TFRecordWriter(\n",
+ " output_path + '-%05d-of-%05d.tfrecord' % (i, num_shards))\n",
+ " for i in range(num_shards)\n",
+ " ]\n",
+ " for idx, record in enumerate(dataset):\n",
+ " if idx % LOG_EVERY == 0:\n",
+ " print('On image %d', idx)\n",
+ " seq_example = process_record(record)\n",
+ " writers[idx % num_shards].write(seq_example.SerializeToString())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Z7gfIpOa_oIj"
+ },
+ "source": [
+ "### Write training data as TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "C2BygDlYnFi8"
+ },
+ "outputs": [],
+ "source": [
+ "output_train_tfrecs = output_dir + 'train'\n",
+ "write_tfrecords(train_ds, output_train_tfrecs, num_shards=10)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "daNmrC7T_t5i"
+ },
+ "source": [
+ "### Write validation data as TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "bxMiQmM8GfQG"
+ },
+ "outputs": [],
+ "source": [
+ "output_val_tfrecs = output_dir + 'valid'\n",
+ "write_tfrecords(val_ds, output_val_tfrecs, num_shards=5)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "5pTf5iv__voH"
+ },
+ "source": [
+ "### Write test data as TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "BRQad9B3GkHQ"
+ },
+ "outputs": [],
+ "source": [
+ "output_test_tfrecs = output_dir + 'test'\n",
+ "write_tfrecords(test_ds, output_test_tfrecs, num_shards=5)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "tHm2vpHy_3w_"
+ },
+ "source": [
+ "## Experiment Configuration"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ffEqt4bD_-bQ"
+ },
+ "source": [
+ "### Load the existing Configuration"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "tf4N6_uW4UgK"
+ },
+ "outputs": [],
+ "source": [
+ "import yaml\n",
+ "\n",
+ "with open('./official/projects/video_ssl/configs/experiments/cvrl_linear_eval_k600.yaml', 'r') as file:\n",
+ " override_params = yaml.full_load(file)\n",
+ "\n",
+ "\n",
+ "exp_config = exp_cfg.exp_factory.get_exp_config('video_ssl_linear_eval_kinetics600')\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "V9QDyhgMAHs2"
+ },
+ "source": [
+ "### Override the configuration parameters"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "HYDhSyC4AQ9c"
+ },
+ "outputs": [],
+ "source": [
+ "exp_config.override(override_params, is_strict=False)\n",
+ "\n",
+ "WIDTH, HEIGHT = 224, 224\n",
+ "\n",
+ "# Runtime configuration\n",
+ "exp_config.runtime.distribution_strategy = \"mirrored\"\n",
+ "\n",
+ "# Task configuration\n",
+ "exp_config.task.freeze_backbone = True\n",
+ "exp_config.task.init_checkpoint = \"../r3d_1x_k600_800ep/r3d_1x_k600_800ep_backbone-1\"\n",
+ "exp_config.task.init_checkpoint_modules = \"backbone\"\n",
+ "\n",
+ "# Model configuration\n",
+ "exp_config.task.model.projection_dim = 10\n",
+ "\n",
+ "# Training data configuration\n",
+ "exp_config.task.train_data.input_path = '../ucf101_tfrecords/train*'\n",
+ "exp_config.task.train_data.num_classes=10\n",
+ "exp_config.task.train_data.global_batch_size = 2\n",
+ "exp_config.task.train_data.min_image_size = WIDTH\n",
+ "exp_config.task.train_data.num_examples = 400\n",
+ "exp_config.task.train_data.feature_shape = (n_frames, HEIGHT, WIDTH, 3)\n",
+ "\n",
+ "# Validation data configuration\n",
+ "exp_config.task.validation_data.num_classes=10\n",
+ "exp_config.task.validation_data.input_path = '../ucf101_tfrecords/valid*'\n",
+ "exp_config.task.validation_data.global_batch_size = 2\n",
+ "exp_config.task.validation_data.min_image_size = WIDTH\n",
+ "exp_config.task.validation_data.num_examples = 100\n",
+ "exp_config.task.validation_data.feature_shape = (n_frames, HEIGHT, WIDTH, 3)\n",
+ "\n",
+ "# Trainer configuration\n",
+ "\n",
+ "exp_config.trainer.train_steps = 2000\n",
+ "exp_config.trainer.checkpoint_interval = 200\n",
+ "exp_config.trainer.steps_per_loop = 200\n",
+ "exp_config.trainer.summary_interval = 200\n",
+ "exp_config.trainer.validation_interval = 200\n",
+ "exp_config.trainer.validation_steps = 200\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.decay_steps = 2000\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.initial_learning_rate = 0.008\n",
+ "exp_config.trainer.optimizer_config.warmup.linear.warmup_learning_rate = 0.007\n",
+ "exp_config.trainer.optimizer_config.warmup.linear.warmup_steps = 200"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "rmxwvgimiSQu"
+ },
+ "source": [
+ "### Set up the distribution strategy"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "RpD__dO5xJUC"
+ },
+ "outputs": [],
+ "source": [
+ "# Detect hardware\n",
+ "try:\n",
+ " tpu_resolver = tf.distribute.cluster_resolver.TPUClusterResolver() # TPU detection\n",
+ "except ValueError:\n",
+ " tpu_resolver = None\n",
+ " gpus = tf.config.experimental.list_logical_devices(\"GPU\")\n",
+ "\n",
+ "# Select appropriate distribution strategy\n",
+ "if tpu_resolver:\n",
+ " tf.config.experimental_connect_to_cluster(tpu_resolver)\n",
+ " tf.tpu.experimental.initialize_tpu_system(tpu_resolver)\n",
+ " distribution_strategy = tf.distribute.experimental.TPUStrategy(tpu_resolver)\n",
+ " print('Running on TPU ', tpu_resolver.cluster_spec().as_dict()['worker'])\n",
+ "elif len(gpus) \u003e 1:\n",
+ " distribution_strategy = tf.distribute.MirroredStrategy([gpu.name for gpu in gpus])\n",
+ " print('Running on multiple GPUs ', [gpu.name for gpu in gpus])\n",
+ "elif len(gpus) == 1:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on single GPU ', gpus[0].name)\n",
+ "else:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on CPU')\n",
+ "\n",
+ "print(\"Number of accelerators: \", distribution_strategy.num_replicas_in_sync)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "iIT9AgxoAday"
+ },
+ "source": [
+ "### Display the final configuration"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "_GVF3uDw8QVt"
+ },
+ "outputs": [],
+ "source": [
+ "pp.pprint(exp_config.as_dict())\n",
+ "display.Javascript('google.colab.output.setIframeHeight(\"500px\");')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "VZF9G3M5i1Lm"
+ },
+ "source": [
+ "## Create the `Task` object (`tfm.core.base_task.Task`) from the `config_definitions.TaskConfig`.\n",
+ "\n",
+ "The `Task` object has all the methods necessary for building the dataset, building the model, and running training \u0026 evaluation. These methods are driven by `tfm.core.train_lib.run_experiment`."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "OPn2AMRZK-zZ"
+ },
+ "outputs": [],
+ "source": [
+ "model_dir = '../trained_model/'\n",
+ "\n",
+ "with distribution_strategy.scope():\n",
+ " task = task_factory.get_task(exp_config.task, logging_dir=model_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "6M0rlFn-igH7"
+ },
+ "source": [
+ "## Visualization of Train Data"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "CGxZEl6ijF2a"
+ },
+ "outputs": [],
+ "source": [
+ "frames, label = list(train_ds.take(1))[0]\n",
+ "print(CLASSES[label])\n",
+ "to_gif(frames)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "bnMqrCtgo4UO"
+ },
+ "source": [
+ "## Train and evaluate"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "watwZ5mIwtvk"
+ },
+ "outputs": [],
+ "source": [
+ "model, eval_logs = train_lib.run_experiment(\n",
+ " distribution_strategy=distribution_strategy,\n",
+ " task=task,\n",
+ " mode='train_and_eval',\n",
+ " params=exp_config,\n",
+ " model_dir=model_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZJBzAMJeo9JN"
+ },
+ "source": [
+ "## Load logs in tensorboard"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "c3smE05dIgAV"
+ },
+ "outputs": [],
+ "source": [
+ "%load_ext tensorboard\n",
+ "%tensorboard --logdir '../trained_model'"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "WSTLebJGo_8a"
+ },
+ "source": [
+ "## Saving and exporting the trained model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "o4lhmgy9I6Du"
+ },
+ "outputs": [],
+ "source": [
+ "export_dir = '../exported_model/'\n",
+ "\n",
+ "export_saved_model_lib.export_inference_graph(\n",
+ " input_type='image_tensor',\n",
+ " batch_size=1,\n",
+ " input_image_size=[n_frames, HEIGHT, WIDTH],\n",
+ " params=exp_config,\n",
+ " checkpoint_path=tf.train.latest_checkpoint(model_dir),\n",
+ " export_dir=export_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "10Yn7-pBpns-"
+ },
+ "source": [
+ "## Importing SavedModel"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "u3fghz-jVrcR"
+ },
+ "outputs": [],
+ "source": [
+ "imported = tf.saved_model.load(export_dir)\n",
+ "model_fn = imported.signatures['serving_default']"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "D6qzUv1lpsRL"
+ },
+ "source": [
+ "## Visualize predictions"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "1McexxfTh-Bl"
+ },
+ "outputs": [],
+ "source": [
+ "frames, label = list(test_ds.shuffle(buffer_size=90).take(1))[0]\n",
+ "frames = tf.expand_dims(frames, axis=0)\n",
+ "result = model_fn(frames)\n",
+ "predicted_label = tf.argmax(result['probs'][0])\n",
+ "print(f\"Actual: {CLASSES[label]}\")\n",
+ "print(f\"Predicted: {CLASSES[predicted_label]}\")\n",
+ "to_gif(frames[0])"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "name": "video_ssl.ipynb",
+ "provenance": [],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/videoglue/README.md b/official/projects/videoglue/README.md
new file mode 100644
index 00000000000..12d1fd2bd97
--- /dev/null
+++ b/official/projects/videoglue/README.md
@@ -0,0 +1,138 @@
+# VideoGLUE: Video General Understanding Evaluation of Foundation Models
+[](https://arxiv.org/abs/2307.03166)
+
+This repository provides the official TensorFlow 2 implementation of
+[VideoGLUE: Video General Understanding Evaluation of Foundation Models](https://arxiv.org/abs/2307.03166)
+
+
+
+
+
+
+ Figure 1: We study four adaptation methods to apply a foundation model (FM) to
+ video understanding downstream tasks: (a) end-to-end finetuning, (b) frozen
+ backbone, (c) frozen backbone with multi-layer attention pooler (MLAP), and
+ (d) a low-rank adapter.
+
+
+
+
+## Description
+
+We evaluate the video understanding capabilities of existing foundation models
+(FMs) using a carefully designed experiment protocol consisting of three
+hallmark tasks (action recognition, temporal localization, and spatiotemporal
+localization), eight datasets well received by the community, and four
+adaptation methods tailoring an FM for downstream tasks. Furthermore, we
+jointly profile FMs' efficacy and efficiency when adapting to general video
+understanding tasks using cost measurements during both training and inference.
+Our main findings are as follows. First, task-specialized models significantly
+outperform the seven FMs studied in this work, in sharp contrast to what FMs
+have achieved in natural language and image understanding. Second, video-native
+FMs, whose pretraining data mainly contains the video modality, are generally
+better than image-native FMs in classifying motion-rich videos, localizing
+actions in time, and understanding a video of more than one action. Third, the
+video-native FMs can perform well on video tasks under light adaptations to
+downstream tasks (e.g., freezing the FM backbones), while image-native FMs win
+in full end-to-end finetuning. The first two observations reveal the need and
+tremendous opportunities to conduct research on video-focused FMs, and the last
+confirms that both tasks and adaptation methods matter when it comes to the
+evaluation of FMs.
+
+## Requirements
+* [DMVR: DeepMind Video Readers](https://github.com/deepmind/dmvr)
+* [TensorFlow Object Detection API](https://github.com/tensorflow/models/blob/master/research/object_detection/g3doc/tf2.md)
+
+## Tasks
+### Video Classification
+Use the following script to run the `video classification` experiment. Update the config to fit your compute environment.
+
+```shell
+PARAMS_OVERRIDE="task.train_data.global_batch_size=2,\
+task.train_data.shuffle_buffer_size=2,\
+task.validation_data.global_batch_size=4,\
+trainer.steps_per_loop=2,\
+trainer.validation_steps=2,\
+runtime.distribution_strategy=mirrored,\
+runtime.mixed_precision_dtype=float32"
+
+CONFIG="${PWD}/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_sthv2.yaml"
+EXPERIMENT="mh_video_classification_strong_aug"
+MODE="train" # change to 'eval' for running evaluation loop.
+
+python3 official/projects/videoglue/train.py \
+--model_dir="/tmp/video_classification" \
+--mode="${MODE}" \
+--experiment="${EXPERIMENT}" \
+--config_file="${CONFIG}" \
+--params_override="${PARAMS_OVERRIDE}"
+```
+
+### Spatiotemporal Action Localization
+Use the following script to run the `spatiotemporal action localization`
+experiment. Update the config to fit your compute environment.
+
+```shell
+PARAMS_OVERRIDE="task.train_data.global_batch_size=2,\
+task.train_data.shuffle_buffer_size=2,\
+task.validation_data.global_batch_size=4,\
+trainer.steps_per_loop=2,\
+trainer.validation_steps=2,\
+runtime.distribution_strategy=mirrored,\
+runtime.mixed_precision_dtype=float32"
+
+CONFIG="${PWD}/third_party/tensorflow_models/official/projects/videoglue/configs/yaml/vmae/ft/stal_vmae_vit3d_ava.yaml"
+EXPERIMENT="spatiotemporal_action_localization_vit12"
+MODE="train" # change to 'eval' for running evaluation loop.
+
+python3 -m official/projects/videoglue/train.py \
+--model_dir="/tmp/spatiotemporal_action_localization" \
+--mode="${MODE}" \
+--experiment="${EXPERIMENT}" \
+--config_file="${CONFIG}" \
+--params_override="${PARAMS_OVERRIDE}"
+```
+
+### Temporal Action Localization
+
+Following prior works, we employ
+[G-TAD](https://arxiv.org/abs/1911.11462) as our task head for predicting
+action categories and start and end timestamps. Please follow this
+[implementation](https://github.com/frostinassiky/gtad) to run Temporal Action
+Localization benchmarks.
+
+To extract features on ActivityNet using FM of choice, we use clips of `16
+frames` at a frame rate of 15 fps and a stride of `16 frames` (i.e.,
+non-overlapping clips). This gives one feature vector per `16/15 ~= 1.067`
+seconds.
+
+## Code
+- [x] Tasks
+ - [x] Task: Video Classification.
+ - [x] Task: Spatiotemporal Action Localization.
+ - [x] Task: Temporal Action Localization.
+- [x] Adaptations
+ - [x] Adaptation: End-to-end fine-tuning.
+ - [x] Adaptation: Frozen backbone pooler head.
+ - [ ] Adaptation: Frozen backbone multi-layer attention pooling.
+ - [ ] Adaptation: Low-rank adapter fine-tuning.
+
+## License
+
+[](https://opensource.org/licenses/Apache-2.0)
+
+This project is licensed under the terms of the **Apache License 2.0**.
+
+## Citation
+```
+@inproceedings{yuan2024videoglue,
+ title={VideoGLUE: Video General Understanding Evaluation of Foundation Models}
+ author={Yuan, Liangzhe and Gundavarapu, Nitesh Bharadwaj and Zhao, Long and
+ Zhou, Hao and Cui, Yin and Jiang, Lu and Yang, Xuan and Jia, Menglin and
+ Weyand, Tobias and Friedman, Luke and Sirotenko, Mikhail and Wang, Huisheng
+ and Schroff, Florian and Adam, Hartwig and Yang, Ming-Hsuan and Liu, Ting and
+ Gong, Boqing}
+ booktitle={Transactions on Machine Learning Research},
+ year={2024}
+}
+```
\ No newline at end of file
diff --git a/official/projects/videoglue/configs/backbones_3d.py b/official/projects/videoglue/configs/backbones_3d.py
new file mode 100644
index 00000000000..93e6f546667
--- /dev/null
+++ b/official/projects/videoglue/configs/backbones_3d.py
@@ -0,0 +1,41 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""3D Backbones configurations."""
+import dataclasses
+from typing import Tuple, Union, Optional
+
+from official.vision.configs import backbones
+from official.vision.configs import backbones_3d
+
+
+@dataclasses.dataclass
+class VisionTransformer3D(backbones.VisionTransformer):
+ """VisionTransformer3D config."""
+ variant: str = 'native'
+ temporal_patch_size: int = 4
+ pos_embed_shape: Optional[Union[Tuple[int, int], Tuple[int, int, int]]] = None
+
+
+@dataclasses.dataclass
+class Backbone3D(backbones_3d.Backbone3D):
+ """Configuration for backbones.
+
+ Attributes:
+ type: type of backbone be used, one of the fields below.
+ vit_3d: vit_3d backbone config.
+ """
+ type: str = 'vit_3d'
+ vit_3d: VisionTransformer3D = dataclasses.field(
+ default_factory=VisionTransformer3D)
diff --git a/official/projects/videoglue/configs/dataset.py b/official/projects/videoglue/configs/dataset.py
new file mode 100644
index 00000000000..ea6aae911bc
--- /dev/null
+++ b/official/projects/videoglue/configs/dataset.py
@@ -0,0 +1,96 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Video coarse classification configuration definition."""
+
+import dataclasses
+from typing import Optional, Tuple, Union, List
+
+from official.modeling import hyperparams
+from official.vision.configs import common as common_cfg
+
+
+@dataclasses.dataclass
+class VGG(hyperparams.Config):
+ """Config for VGG style data augmentation."""
+ pass
+
+
+@dataclasses.dataclass
+class Inception(hyperparams.Config):
+ """Config for Inception style data augmentation."""
+ min_aspect_ratio: float = 0.5
+ max_aspect_ratio: float = 2.0
+ min_area_ratio: float = 0.3
+ max_area_ratio: float = 1.0
+
+
+@dataclasses.dataclass
+class AVA(hyperparams.Config):
+ """Config for AVA style data augmentation."""
+ scale_min: float = 0.5
+ scale_max: float = 2.0
+
+
+@dataclasses.dataclass
+class DataAugmentation(hyperparams.OneOfConfig):
+ """Configuration for data augmentation.
+
+ Attributes:
+ type: 'str', type of backbone be used, one of the fields below.
+ inception: resnet backbone config.
+ vgg: dilated resnet backbone for semantic segmentation config.
+ ava: revnet backbone config.
+ """
+ type: Optional[str] = None
+ vgg: VGG = dataclasses.field(default_factory=VGG)
+ inception: Inception = dataclasses.field(default_factory=Inception)
+ ava: AVA = dataclasses.field(default_factory=AVA)
+
+
+@dataclasses.dataclass
+class DataConfig(hyperparams.Config):
+ """The base configuration for building datasets."""
+ name: str = 'some_dataset'
+ is_training: bool = False
+ num_classes: Union[int, List[int]] = 1
+ label_names: Union[str, List[str]] = 'label'
+ num_examples: int = 10000
+ global_batch_size: int = 128
+ feature_shape: Tuple[int, ...] = (64, 224, 224, 3)
+ min_resize: int = 256
+ temporal_stride: int = 5
+ sample_from_segments: bool = False
+ zero_centering_image: bool = True
+ random_flip_image: bool = True
+ num_test_clips: int = 1
+ num_test_crops: int = 1
+ data_augmentation: DataAugmentation = dataclasses.field(
+ default_factory=lambda: DataAugmentation(type='vgg')
+ )
+ randaug: Optional[common_cfg.RandAugment] = None
+ autoaug: Optional[common_cfg.AutoAugment] = None
+ mixup_cutmix: Optional[common_cfg.MixupAndCutmix] = None
+ is_multilabel: bool = False
+ # Pipeline parameters.
+ drop_remainder: bool = True
+ dtype: str = 'float32'
+ prefetch_buffer_size: int = 512
+ shuffle_buffer_size: int = 256
+ num_process_threads: int = 32
+ num_parallel_calls_interleave: int = 32
+ cycle_length: int = 32
+ block_length: int = 1
+ cache: bool = False
+ use_tf_data_service: bool = False
diff --git a/official/projects/videoglue/configs/head.py b/official/projects/videoglue/configs/head.py
new file mode 100644
index 00000000000..149ce1d0e3b
--- /dev/null
+++ b/official/projects/videoglue/configs/head.py
@@ -0,0 +1,56 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Configs for different model heads."""
+
+import dataclasses
+from typing import Optional
+
+from official.modeling import hyperparams
+
+
+@dataclasses.dataclass
+class MLP(hyperparams.Config):
+ """Config for the MLP head."""
+ normalize_inputs: bool = True
+ num_hidden_channels: int = 2048
+ num_hidden_layers: int = 3
+ num_output_channels: int = 128
+ use_sync_bn: bool = False
+ norm_momentum: float = 0.997
+ norm_epsilon: float = 1e-5
+ activation: Optional[str] = 'relu'
+
+
+@dataclasses.dataclass
+class ActionTransformer(hyperparams.Config):
+ """Config for the action transformer head."""
+ # parameters for classifier
+ num_hidden_layers: int = 0
+ num_hidden_channels: int = 0
+ use_sync_bn: bool = True
+ activation: str = 'relu'
+ dropout_rate: float = 0.0
+ # parameters for RoiAligner
+ crop_size: int = 4
+ sample_offset: float = 0.5
+ # parameters for TxDecoder
+ num_tx_channels: int = 128
+ num_tx_layers: int = 3
+ num_tx_heads: int = 3
+ use_bias: bool = True
+ tx_activation: Optional[str] = 'gelu'
+ attention_dropout_rate: float = 0.0
+ layer_norm_epsilon: float = 1e-12
+ use_positional_embedding: bool = True
diff --git a/official/projects/videoglue/configs/spatiotemporal_action_localization.py b/official/projects/videoglue/configs/spatiotemporal_action_localization.py
new file mode 100644
index 00000000000..e361407e754
--- /dev/null
+++ b/official/projects/videoglue/configs/spatiotemporal_action_localization.py
@@ -0,0 +1,288 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Spatiotemporal action localization configuration definition."""
+
+import dataclasses
+from typing import Optional
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.mae import optimization
+from official.projects.videoglue.configs import backbones_3d
+from official.projects.videoglue.configs import dataset
+from official.projects.videoglue.configs import head as head_cfg
+from official.vision.configs import common
+from official.vision.configs import video_classification
+
+Losses = video_classification.Losses
+
+
+@dataclasses.dataclass
+class DataConfig(dataset.DataConfig):
+ """The dataset config."""
+ data_augmentation: dataset.DataAugmentation = dataclasses.field(
+ default_factory=lambda: dataset.DataAugmentation( # pylint: disable=g-long-lambda
+ type='ava', ava=dataset.AVA()
+ )
+ )
+ is_training: bool = True
+ drop_remainder: bool = True
+ num_instances: int = 32
+ num_classes: int = 80
+ one_hot_label: bool = True
+ merge_multi_labels: bool = True
+ import_detected_bboxes: bool = False
+ color_augmentation: bool = True
+
+
+@dataclasses.dataclass
+class VideoActionTransformerModel(hyperparams.Config):
+ """The model config."""
+ model_type: str = 'video_action_transformer_model'
+ backbone: backbones_3d.Backbone3D = dataclasses.field(
+ default_factory=lambda: backbones_3d.Backbone3D( # pylint: disable=g-long-lambda
+ type='vit_3d',
+ vit_3d=backbones_3d.VisionTransformer3D(pooler='none')))
+ endpoint_name: str = 'encoded_tokens'
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ use_sync_bn=True, norm_momentum=0.9, norm_epsilon=1e-5
+ )
+ )
+ head: head_cfg.ActionTransformer = dataclasses.field(
+ default_factory=lambda: head_cfg.ActionTransformer( # pylint: disable=g-long-lambda
+ use_sync_bn=True,
+ num_hidden_layers=1,
+ num_hidden_channels=1024,
+ crop_size=7,
+ )
+ )
+
+
+@dataclasses.dataclass
+class SpatiotemporalActionLocalizationTask(
+ video_classification.VideoClassificationTask):
+ """Task for video action localization."""
+ model: VideoActionTransformerModel = dataclasses.field(
+ default_factory=VideoActionTransformerModel
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ data_augmentation=dataset.DataAugmentation(type='vgg'),
+ is_training=True,
+ drop_remainder=True,
+ )
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ is_training=False, drop_remainder=False
+ )
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ init_checkpoint: Optional[str] = None
+ init_checkpoint_modules: str = 'all' # all or backbone
+
+
+# NOTE: This utility function includes different name conventions of the same
+# layer from different ViT implementation. They would be safely ignored as long
+# as no substring collision happens. However it is still error-prone. Use with
+# caution. See how it is in:
+# tensorflow_models/official/projects/mae/optimization.py;l=40
+def _get_vit_layers(num_tx_layers: int = 12):
+ """Gets ViT layers substring and index."""
+ layers_substr = [
+ # rgb projection layer supports VMAE/INTERNVIDEO/FLVID.
+ 'conv3d/kernel',
+ 'conv3d/bias',
+
+ # postional embedding for VMAE.
+ 'add_separable_position_embs/pos_embedding_time',
+ 'add_separable_position_embs/pos_embedding_space',
+
+ # rgb projection layer supports IMP.
+ 'rgb_to_embedding',
+ # positional embedding for IMP.
+ 'rgb_pos_encoding',
+ # pre-projection for IMP.
+ 'pre_projection/vision_dense',
+
+ # rgb projection layer supports COCA.
+ 'input_projection/kernel',
+ 'input_projection/bias',
+ # rgb projection layer supports FLAVA.
+ 'conv2d/kernel',
+ 'conv2d/bias',
+
+ # positional embedding for COCA/FLAVA.
+ 'encoder/posembed_input/pos_embedding',
+
+ # common encoder final layer norm.
+ 'encoder/layer_normalization',
+
+ # post-projection for IMP.
+ 'post_projection/vision_dense',
+ ]
+ layers_idx = [
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ 0,
+ num_tx_layers,
+ num_tx_layers,
+ ]
+ if len(layers_idx) != len(layers_substr):
+ raise ValueError('layers_idx and layers_substr should have same length.')
+
+ for idx in range(num_tx_layers):
+ if idx == 0:
+ var_substr = 'encoder/transformer_encoder_block/'
+ else:
+ var_substr = f'encoder/transformer_encoder_block_{idx}/'
+ layers_substr.append(var_substr)
+ layers_idx.append(idx + 1)
+ return layers_idx, layers_substr
+
+
+def _get_clip_layers(num_tx_layers: int = 12):
+ """Gets CLIP layers substring and index."""
+ layers_substr = [
+ # rgb projection layer
+ 'conv1/kernel',
+
+ # postional embedding.
+ 'clip/visual/positional_embedding',
+
+ # encoder pre layer norm.
+ 'ln_pre',
+
+ # class embedding.
+ 'clip/visual/class_embedding',
+
+ # post layer norm.
+ 'ln_post',
+
+ # post-projection.
+ 'proj/kernel',
+ ]
+ layers_idx = [
+ 0,
+ 0,
+ 0,
+ 0,
+ num_tx_layers,
+ num_tx_layers,
+ ]
+ if len(layers_idx) != len(layers_substr):
+ raise ValueError('layers_idx and layers_substr should have same length.')
+
+ for idx in range(num_tx_layers):
+ var_substr = f'transformer/resblocks.{idx}/'
+ layers_substr.append(var_substr)
+ layers_idx.append(idx + 1)
+ return layers_idx, layers_substr
+
+
+@exp_factory.register_config_factory('spatiotemporal_action_localization')
+def spatiotemporal_action_localization() -> cfg.ExperimentConfig:
+ """Spatio-temporal action localization."""
+ task = SpatiotemporalActionLocalizationTask()
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(mixed_precision_dtype='bfloat16'),
+ task=task,
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ 'task.train_data.num_classes == task.validation_data.num_classes',
+ ])
+ config = video_classification.add_trainer(
+ config, train_batch_size=1024, eval_batch_size=64)
+ return config
+
+
+@exp_factory.register_config_factory('spatiotemporal_action_localization_vit12')
+def spatiotemporal_action_localization_vit12() -> cfg.ExperimentConfig:
+ """Spatio-temporal action localization for ViT-B with layer decay."""
+ config = spatiotemporal_action_localization()
+ layers_idx, vars_substr = _get_vit_layers(num_tx_layers=12)
+ optimizer_config = optimization.OptimizerConfig({
+ 'type': 'vit_adamw',
+ 'vit_adamw': {
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0,
+ 'beta_2': 0.999,
+ 'layer_decay': 0.75,
+ 'vars_substr': vars_substr,
+ 'layers_idx': layers_idx,
+ 'exclude_from_weight_decay': ['cls'],
+ },
+ })
+ config.trainer.optimizer_config.optimizer = optimizer_config
+ return config
+
+
+@exp_factory.register_config_factory(
+ 'spatiotemporal_action_localization_clip12')
+def spatiotemporal_action_localization_clip12() -> cfg.ExperimentConfig:
+ """Spatio-temporal action localization for CLIP-B with layer decay."""
+ config = spatiotemporal_action_localization()
+ layers_idx, vars_substr = _get_clip_layers(num_tx_layers=12)
+ optimizer_config = optimization.OptimizerConfig({
+ 'type': 'vit_adamw',
+ 'vit_adamw': {
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0,
+ 'beta_2': 0.999,
+ 'layer_decay': 0.75,
+ 'vars_substr': vars_substr,
+ 'layers_idx': layers_idx,
+ 'exclude_from_weight_decay': ['cls'],
+ },
+ })
+ config.trainer.optimizer_config.optimizer = optimizer_config
+ return config
+
+
+@exp_factory.register_config_factory('spatiotemporal_action_localization_vit16')
+def spatiotemporal_action_localization_vit16() -> cfg.ExperimentConfig:
+ """Spatio-temporal action localization for ViT-L/H/G with layer decay."""
+ config = spatiotemporal_action_localization()
+ # ViT-L/H/G have 16 layers transformer encoder block.
+ layers_idx, vars_substr = _get_vit_layers(num_tx_layers=16)
+ optimizer_config = optimization.OptimizerConfig({
+ 'type': 'vit_adamw',
+ 'vit_adamw': {
+ 'weight_decay_rate': 0.05,
+ # Avoid AdamW legacy behavior.
+ 'gradient_clip_norm': 0.0,
+ 'beta_2': 0.999,
+ 'layer_decay': 0.75,
+ 'vars_substr': vars_substr,
+ 'layers_idx': layers_idx,
+ 'exclude_from_weight_decay': ['cls'],
+ },
+ })
+ config.trainer.optimizer_config.optimizer = optimizer_config
+ return config
diff --git a/official/projects/videoglue/configs/video_classification.py b/official/projects/videoglue/configs/video_classification.py
new file mode 100644
index 00000000000..dc83f54bab2
--- /dev/null
+++ b/official/projects/videoglue/configs/video_classification.py
@@ -0,0 +1,83 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Video classification configuration definition."""
+import dataclasses
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.projects.videoglue.configs import backbones_3d
+from official.projects.videoglue.configs import dataset
+from official.vision.configs import common
+from official.vision.configs import video_classification
+
+VideoClassificationTask = video_classification.VideoClassificationTask
+DataConfig = video_classification.DataConfig
+
+
+@dataclasses.dataclass
+class MultiHeadVideoClassificationModel(
+ video_classification.VideoClassificationModel):
+ """The model config."""
+ model_type: str = 'mh_video_classification'
+ backbone: backbones_3d.Backbone3D = dataclasses.field(
+ default_factory=lambda: backbones_3d.Backbone3D( # pylint: disable=g-long-lambda
+ type='vit_3d', vit_3d=backbones_3d.VisionTransformer3D()
+ )
+ )
+ classifier_type: str = 'pooler' # 'linear' or 'pooler'
+ # only useful when classifier_type == 'pooler'
+ attention_num_heads: int = 6
+ attention_hidden_size: int = 768
+ attention_dropout_rate: float = 0.0
+ add_temporal_pos_emb_pooler: bool = False
+
+
+@dataclasses.dataclass
+class MultiHeadVideoClassificationTask(VideoClassificationTask):
+ """The task config."""
+ _MHVCM = MultiHeadVideoClassificationModel
+ model: _MHVCM = dataclasses.field(default_factory=_MHVCM)
+ train_data: dataset.DataConfig = dataclasses.field(
+ default_factory=lambda: dataset.DataConfig( # pylint: disable=g-long-lambda
+ is_training=True,
+ data_augmentation=dataset.DataAugmentation(type='inception'),
+ )
+ )
+ validation_data: dataset.DataConfig = dataclasses.field(
+ default_factory=lambda: dataset.DataConfig(is_training=False)
+ )
+
+
+@exp_factory.register_config_factory('mh_video_classification')
+def mh_video_classification() -> cfg.ExperimentConfig:
+ """Multi-head video classification."""
+ exp = video_classification.video_classification_kinetics400()
+ task = MultiHeadVideoClassificationTask()
+ exp.task = task
+ return exp
+
+
+@exp_factory.register_config_factory('mh_video_classification_strong_aug')
+def mh_video_classification_strong_aug() -> cfg.ExperimentConfig:
+ """Multi-head video classification with strong augmentation."""
+ exp = video_classification.video_classification_kinetics400()
+ task = MultiHeadVideoClassificationTask()
+ task.train_data = dataset.DataConfig(
+ is_training=True,
+ data_augmentation=dataset.DataAugmentation(type='inception'),
+ randaug=common.RandAugment(magnitude=9, magnitude_std=0.5),
+ mixup_cutmix=common.MixupAndCutmix())
+ exp.task = task
+ return exp
diff --git a/official/projects/videoglue/configs/yaml/vmae/frozen/stal_vmae_vit3d_ava.yaml b/official/projects/videoglue/configs/yaml/vmae/frozen/stal_vmae_vit3d_ava.yaml
new file mode 100644
index 00000000000..b2a26448a32
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/frozen/stal_vmae_vit3d_ava.yaml
@@ -0,0 +1,88 @@
+# experiment_type=spatiotemporal_action_localization_vit12
+
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: true
+ model:
+ endpoint_name: 'encoded_tokens'
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ head:
+ dropout_rate: 0.5
+ num_hidden_layers: 0
+ num_hidden_channels: 1024
+ num_tx_channels: 768
+ num_tx_heads: 6
+ num_tx_layers: 1
+ use_positional_embedding: true
+ train_data:
+ name: ava
+ num_examples: 210634
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: ava
+ num_examples: 57371
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 8
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ import_detected_bboxes: true
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 41140 # 50eps
+ optimizer:
+ type: 'vit_adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ warmup:
+ linear:
+ warmup_steps: 1500
+ train_steps: 41140 # 50eps
+ validation_steps: 7171
+ steps_per_loop: 100
+ summary_interval: 100
+ validation_interval: 100
+ checkpoint_interval: 1000
diff --git a/official/projects/videoglue/configs/yaml/vmae/frozen/stal_vmae_vit3d_avakinetics.yaml b/official/projects/videoglue/configs/yaml/vmae/frozen/stal_vmae_vit3d_avakinetics.yaml
new file mode 100644
index 00000000000..71e73773eb0
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/frozen/stal_vmae_vit3d_avakinetics.yaml
@@ -0,0 +1,88 @@
+# experiment_type=spatiotemporal_action_localization_vit12
+
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: true
+ model:
+ endpoint_name: 'encoded_tokens'
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ head:
+ dropout_rate: 0.5
+ num_hidden_layers: 0
+ num_hidden_channels: 1024
+ num_tx_channels: 768
+ num_tx_heads: 6
+ num_tx_layers: 1
+ use_positional_embedding: true
+ train_data:
+ name: avakinetics
+ num_examples: 353201
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: avakinetics
+ num_examples: 91919
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 8
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ import_detected_bboxes: true
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 68984
+ optimizer:
+ type: 'vit_adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ warmup:
+ linear:
+ warmup_steps: 1500
+ train_steps: 68984 # 50 eps
+ validation_steps: 11489
+ steps_per_loop: 100
+ summary_interval: 100
+ validation_interval: 100
+ checkpoint_interval: 1000
diff --git a/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_charades.yaml b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_charades.yaml
new file mode 100644
index 00000000000..27ef6689ee5
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_charades.yaml
@@ -0,0 +1,92 @@
+# ViT video classification on Charades.
+#
+# --experiment_type=mh_video_classification
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: true
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ attention_num_heads: 12
+ attention_hidden_size: 768
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: charades
+ num_classes: 157
+ num_examples: 209392
+ zero_centering_image: false
+ is_multilabel: true
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 1 # setting temporal stride = 1 because raw video is 6fps.
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: charades
+ num_classes: 157
+ zero_centering_image: false
+ is_multilabel: true
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 1 # setting temporal stride = 1 because raw video is 6fps.
+ num_test_clips: 4 # video length 128.
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 40896 # 50 eps
+ warmup:
+ linear:
+ warmup_steps: 2044 # 5%
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 40896 # 50 eps
+ validation_steps: 436 # 6976 // 16
+ steps_per_loop: 1000
+ summary_interval: 1000
+ validation_interval: 1000
diff --git a/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_diving48.yaml b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_diving48.yaml
new file mode 100644
index 00000000000..d90a7fd9be2
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_diving48.yaml
@@ -0,0 +1,91 @@
+# ViT video classification on Diving48
+#
+# --experiment_type=mh_video_classification
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: true
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ attention_num_heads: 12
+ attention_hidden_size: 768
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: diving48
+ num_classes: 48
+ num_examples: 15027
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: diving48
+ num_classes: 48
+ zero_centering_image: false
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ num_test_clips: 4
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+ l2_weight_decay: 5e-5
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.0001
+ decay_steps: 5869 # 100 eps
+ warmup:
+ linear:
+ warmup_steps: 293
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 0.00001
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 5869 # 100 eps
+ validation_steps: 122 # 1970 // 16
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_k400.yaml b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_k400.yaml
new file mode 100644
index 00000000000..0623c6486a3
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_k400.yaml
@@ -0,0 +1,90 @@
+# ViT video classification on Kinetics-400.
+#
+# --experiment_type=mh_video_classification
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: true
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ attention_num_heads: 12
+ attention_hidden_size: 768
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: kinetics400
+ num_classes: 400
+ num_examples: 235693
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: kinetics400
+ num_classes: 400
+ zero_centering_image: false
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ num_test_clips: 4
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 138101 # 138k
+ warmup:
+ linear:
+ warmup_steps: 4600
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 138101 # 138k, 150 eps
+ validation_steps: 1196
+ steps_per_loop: 1000
+ summary_interval: 1000
+ validation_interval: 1000
diff --git a/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_mit.yaml b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_mit.yaml
new file mode 100644
index 00000000000..28696d0df23
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_mit.yaml
@@ -0,0 +1,91 @@
+# ViT video classification on Moments-in-Time.
+#
+# --experiment_type=mh_video_classification
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: true
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ attention_num_heads: 12
+ attention_hidden_size: 768
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: moments-in-time
+ num_classes: 339
+ num_examples: 791297
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: moments-in-time
+ num_classes: 339
+ zero_centering_image: false
+ num_examples: 33900
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ num_test_clips: 4
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 154550 # 50 eps
+ warmup:
+ linear:
+ warmup_steps: 7727
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 154550 # 50 eps
+ validation_steps: 2118 # 33900 // 16
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_sthv2.yaml b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_sthv2.yaml
new file mode 100644
index 00000000000..7cff730686d
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/frozen/vc_vmae_vit3d_sthv2.yaml
@@ -0,0 +1,95 @@
+# ViT video classification on Sth-sth-v2.
+#
+# --experiment_type=mh_video_classification
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: true
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ attention_num_heads: 12
+ attention_hidden_size: 768
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: sthv2
+ num_classes: 174
+ num_examples: 168913
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ sample_from_segments: true
+ random_flip_image: false
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: sthv2
+ num_classes: 174
+ num_examples: 24777
+ zero_centering_image: false
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ sample_from_segments: true
+ random_flip_image: false
+ num_test_clips: 1
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 60000 # 30k
+ warmup:
+ linear:
+ warmup_steps: 2000
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 60000 # 30k
+ validation_steps: 1548 # 24777 // 16
+ steps_per_loop: 2000
+ summary_interval: 2000
+ validation_interval: 2000
diff --git a/official/projects/videoglue/configs/yaml/vmae/ft/stal_vmae_vit3d_ava.yaml b/official/projects/videoglue/configs/yaml/vmae/ft/stal_vmae_vit3d_ava.yaml
new file mode 100644
index 00000000000..b32705cbfe0
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/ft/stal_vmae_vit3d_ava.yaml
@@ -0,0 +1,88 @@
+# experiment_type=spatiotemporal_action_localization_vit12
+
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: false
+ model:
+ endpoint_name: 'encoded_tokens'
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ head:
+ dropout_rate: 0.5
+ num_hidden_layers: 0
+ num_hidden_channels: 1024
+ num_tx_channels: 768
+ num_tx_heads: 6
+ num_tx_layers: 1
+ use_positional_embedding: true
+ train_data:
+ name: ava
+ num_examples: 210634
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: ava
+ num_examples: 57371
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 8
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ import_detected_bboxes: true
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.01
+ decay_steps: 41140 # 50eps
+ optimizer:
+ type: 'vit_adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 0.00001
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ warmup:
+ linear:
+ warmup_steps: 1500
+ train_steps: 41140 # 50eps
+ validation_steps: 7171
+ steps_per_loop: 100
+ summary_interval: 100
+ validation_interval: 100
+ checkpoint_interval: 1000
diff --git a/official/projects/videoglue/configs/yaml/vmae/ft/stal_vmae_vit3d_avakinetics.yaml b/official/projects/videoglue/configs/yaml/vmae/ft/stal_vmae_vit3d_avakinetics.yaml
new file mode 100644
index 00000000000..d0b06b57b82
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/ft/stal_vmae_vit3d_avakinetics.yaml
@@ -0,0 +1,89 @@
+# experiment_type=spatiotemporal_action_localization_vit12
+
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ freeze_backbone: false
+ model:
+ endpoint_name: 'encoded_tokens'
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ head:
+ dropout_rate: 0.5
+ num_hidden_layers: 0
+ num_hidden_channels: 1024
+ num_tx_channels: 768
+ num_tx_heads: 6
+ num_tx_layers: 1
+ use_positional_embedding: true
+ train_data:
+ name: avakinetics
+ num_examples: 353201
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: avakinetics
+ num_examples: 91919
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 256
+ - 256
+ - 3
+ temporal_stride: 4
+ global_batch_size: 8
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ drop_remainder: true
+ import_detected_bboxes: true
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.01
+ decay_steps: 68984
+ optimizer:
+ type: 'vit_adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 0.00001
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ warmup:
+ linear:
+ warmup_steps: 1500
+ train_steps: 68984 # 50 eps
+ validation_steps: 11489
+ steps_per_loop: 100
+ summary_interval: 100
+ validation_interval: 100
+ checkpoint_interval: 1000
+ max_to_keep: 2
diff --git a/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_charades.yaml b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_charades.yaml
new file mode 100644
index 00000000000..3ef6507e1f6
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_charades.yaml
@@ -0,0 +1,88 @@
+# ViT video classification on Charades.
+#
+# --experiment_type=mh_video_classification_strong_aug
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: charades
+ num_classes: 157
+ num_examples: 209392
+ zero_centering_image: false
+ is_multilabel: true
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 1 # setting temporal stride = 1 because raw video is 6fps.
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: charades
+ num_classes: 157
+ zero_centering_image: false
+ is_multilabel: true
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 1 # setting temporal stride = 1 because raw video is 6fps.
+ num_test_clips: 4 # video length 128.
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 40896 # 50 eps
+ warmup:
+ linear:
+ warmup_steps: 2044 # 5%
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 40896 # 50 eps
+ validation_steps: 436 # 6976 // 16
+ steps_per_loop: 1000
+ summary_interval: 1000
+ validation_interval: 1000
diff --git a/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_diving48.yaml b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_diving48.yaml
new file mode 100644
index 00000000000..fa28c28f017
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_diving48.yaml
@@ -0,0 +1,87 @@
+# ViT video classification on Diving48
+#
+# --experiment_type=mh_video_classification
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.2
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ dropout_rate: 0.5
+ classifier_type: 'pooler'
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: diving48
+ num_classes: 48
+ num_examples: 15027
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: diving48
+ num_classes: 48
+ zero_centering_image: false
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ num_test_clips: 4
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 5869 # 100 eps
+ warmup:
+ linear:
+ warmup_steps: 293
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 5869 # 100 eps
+ validation_steps: 122 # 1970 // 16
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_k400.yaml b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_k400.yaml
new file mode 100644
index 00000000000..ddfbfe9f7e3
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_k400.yaml
@@ -0,0 +1,87 @@
+# ViT video classification on Kinetics-400.
+#
+# --experiment_type=mh_video_classification_strong_aug
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.2
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: kinetics400
+ num_classes: 400
+ num_examples: 235693
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: kinetics400
+ num_classes: 400
+ zero_centering_image: false
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ num_test_clips: 4
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 138101 # 138k
+ warmup:
+ linear:
+ warmup_steps: 4600
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 138101 # 138k, 150 eps
+ validation_steps: 1196
+ steps_per_loop: 1000
+ summary_interval: 1000
+ validation_interval: 1000
diff --git a/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_mit.yaml b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_mit.yaml
new file mode 100644
index 00000000000..828e958f036
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_mit.yaml
@@ -0,0 +1,89 @@
+# ViT video classification on Moments-in-Time.
+#
+# --experiment_type=mh_video_classification
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.2
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ dropout_rate: 0.5
+ classifier_type: 'pooler'
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: moments-in-time
+ num_classes: 339
+ num_examples: 791297
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: moments-in-time
+ num_classes: 339
+ zero_centering_image: false
+ num_examples: 33900
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ num_test_clips: 4
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+ l2_weight_decay: 5e-5
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 154550 # 50 eps
+ warmup:
+ linear:
+ warmup_steps: 7727
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 154550 # 50 eps
+ validation_steps: 2118 # 33900 // 16
+ steps_per_loop: 500
+ summary_interval: 500
+ validation_interval: 500
diff --git a/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_sthv2.yaml b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_sthv2.yaml
new file mode 100644
index 00000000000..d6741e1d98a
--- /dev/null
+++ b/official/projects/videoglue/configs/yaml/vmae/ft/vc_vmae_vit3d_sthv2.yaml
@@ -0,0 +1,94 @@
+# ViT video classification on Sth-sth-v2.
+#
+# --experiment_type=mh_video_classification_strong_aug
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: '/tmp/video_mae/VIT_B_16x4_MAE_PT.npy'
+ init_checkpoint_modules: 'customized_vmae'
+ model:
+ backbone:
+ type: 'vit_3d'
+ vit_3d:
+ variant: 'mae'
+ pooler: 'none'
+ model_name: 'vit-b16'
+ temporal_patch_size: 2
+ representation_size: 0
+ init_stochastic_depth_rate: 0.2
+ pos_embed_shape:
+ - 8 # time
+ - 196 # space
+ transformer:
+ dropout_rate: 0.0
+ classifier_type: 'pooler'
+ dropout_rate: 0.5
+ norm_activation:
+ norm_momentum: 0.9
+ use_sync_bn: true
+ train_data:
+ name: sthv2
+ num_classes: 174
+ num_examples: 168913
+ zero_centering_image: false
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ sample_from_segments: true
+ random_flip_image: false
+ global_batch_size: 256
+ dtype: 'bfloat16'
+ shuffle_buffer_size: 512
+ prefetch_buffer_size: 512
+ drop_remainder: true
+ validation_data:
+ name: sthv2
+ num_classes: 174
+ num_examples: 24777
+ zero_centering_image: false
+ global_batch_size: 16
+ feature_shape: !!python/tuple
+ - 16
+ - 224
+ - 224
+ - 3
+ temporal_stride: 4
+ sample_from_segments: true
+ random_flip_image: false
+ num_test_clips: 1
+ num_test_crops: 3
+ dtype: 'bfloat16'
+ drop_remainder: true
+ losses:
+ label_smoothing: 0.1
+ l2_weight_decay: 5e-5
+trainer:
+ optimizer_config:
+ learning_rate:
+ cosine:
+ initial_learning_rate: 1e-4
+ decay_steps: 32990 # 50 eps
+ warmup:
+ linear:
+ warmup_steps: 6598 # 5%
+ optimizer:
+ type: 'adamw'
+ adamw:
+ beta_1: 0.9
+ beta_2: 0.999
+ amsgrad: false
+ weight_decay_rate: 1e-5
+ exclude_from_weight_decay: !!python/list
+ - 'LayerNorm'
+ - 'layer_norm'
+ - 'bias'
+ train_steps: 32990 # 50 eps
+ validation_steps: 1548 # 24777 // 16
+ steps_per_loop: 1000
+ summary_interval: 1000
+ validation_interval: 1000
+ recovery_max_trials: 10
diff --git a/official/projects/videoglue/datasets/action_localization.py b/official/projects/videoglue/datasets/action_localization.py
new file mode 100644
index 00000000000..dbe8f9e4462
--- /dev/null
+++ b/official/projects/videoglue/datasets/action_localization.py
@@ -0,0 +1,311 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""The dataset factory for the video action localization."""
+import functools
+import os
+from typing import Any, Dict, List, Mapping, Optional, Union
+
+from dmvr import video_dataset
+import tensorflow as tf, tf_keras
+
+from official.projects.videoglue.datasets.common import utils
+
+
+class ActionLocalizationBaseFactory(video_dataset.BaseVideoDatasetFactory):
+ """Action Localization dataset factory."""
+
+ _BASE_DIR = '/tmp'
+ _TABLES = {
+ 'train': 'example.tfrecord',
+ 'test': 'example.tfrecord',
+ }
+ _SUBSETS = ('train', 'test')
+
+ _KEYFRAME_INDEX_KEY = 'clip/key_frame/frame_index'
+ _GT_PREFIX = 'clip/key_frame'
+ _DETECTOR_PREFIX = 'centernet'
+ _ZERO_BASED_INDEX = False # whether the labels are 0-indexed in the table.
+ # threshold to be applied for filtering detected boxes.
+ _TRAIN_DETECTION_SCORE = 0.9
+ _EVAL_DETECTION_SCORE = 0.8
+
+ _NUM_CLASSES = 80
+
+ def __init__(
+ self,
+ subset: str = 'train'):
+ """Initializes the factory."""
+
+ if subset not in self._SUBSETS:
+ raise ValueError(f'Invalid subset "{subset}".'
+ f' The available subsets are: {self._SUBSETS}')
+
+ table_name = os.path.join(self._BASE_DIR, self._TABLES[subset])
+ shards = utils.get_shards(table_name)
+
+ super().__init__(shards)
+
+ def _build(self,
+ is_training: bool = True,
+ # Video related parameters.
+ num_frames: int = 32,
+ temporal_stride: int = 1,
+ num_instance_per_frame: int = 5,
+ # Image related parameters.
+ min_resize: int = 224,
+ crop_size: int = 200,
+ zero_centering_image: bool = False,
+ color_augmentation: bool = False,
+ augmentation_type: str = 'AVA',
+ augmentation_params: Optional[Mapping[str, Any]] = None,
+ # Test related parameters,
+ num_test_clips: int = 1,
+ # Label related parameters.
+ one_hot_label: bool = True,
+ merge_multi_labels: bool = False,
+ import_detected_bboxes: bool = False):
+ """Builds the data processing graph.
+
+ Args:
+ is_training: Whether or not in training mode. If `True`, random sample,
+ crop and left right flip are used.
+ num_frames: Number of frames per subclip.
+ temporal_stride: Temporal stride to sample frames.
+ num_instance_per_frame: The max number of instances per frame to keep.
+ min_resize: Frames are resized so that `min(height, width)` is
+ `min_resize`.
+ crop_size: Final size of the frame after cropping the resized frames. Both
+ height and width are the same.
+ zero_centering_image: If `True`, frames are normalized to values in
+ [-1, 1]. If `False`, values in [0, 1].
+ color_augmentation: Whether to apply color augmentation on video clips.
+ augmentation_type: The data augmentation style applied on images.
+ augmentation_params: A dictionary of params for data augmentation.
+ num_test_clips: Number of test clips (1 by default). If more than 1, this
+ will sample multiple linearly spaced clips within each video at test
+ time. If 1, then a single clip in the middle of the video is sampled.
+ The clips are aggreagated in the batch dimension.
+ one_hot_label: Whether to return one-hot label.
+ merge_multi_labels: Whether to merge multi_labels.
+ import_detected_bboxes: Whether to parse and return detected boxes.
+ """
+ if num_test_clips != 1:
+ raise ValueError('only support num_test_clips = 1 for action '
+ 'localization task. ')
+
+ # Parse keyframe index.
+ self.parser_builder.parse_feature(
+ feature_type=tf.io.FixedLenFeature(1, dtype=tf.int64),
+ feature_name=self._KEYFRAME_INDEX_KEY,
+ output_name='keyframe_index',
+ is_context=True)
+ # Add keyframe boxes.
+ for dim in ['ymin', 'xmin', 'ymax', 'xmax']:
+ utils.add_context_box_dim(
+ parser_builder=self.parser_builder,
+ sampler_builder=self.sampler_builder,
+ preprocessor_builder=self.preprocessor_builder,
+ input_box_dim_name=f'{self._GT_PREFIX}/bbox/{dim}',
+ output_box_dim_name=f'instances_{dim}',
+ num_instances_per_frame=num_instance_per_frame,
+ num_frames=num_frames) # > 1 to duplicate keyframe boxes to all.
+ utils.group_instance_box_dims(
+ preprocessor_builder=self.preprocessor_builder,
+ output_position_name='instances_position',
+ box_key_prefix='instances')
+ # Add boxes scores
+ utils.add_context_box_dim(
+ parser_builder=self.parser_builder,
+ sampler_builder=self.sampler_builder,
+ preprocessor_builder=self.preprocessor_builder,
+ input_box_dim_name=f'{self._GT_PREFIX}/bbox/score',
+ output_box_dim_name='instances_score',
+ num_instances_per_frame=num_instance_per_frame,
+ num_frames=num_frames, # duplicate keyframe boxes to all frames.
+ default_value=1.0)
+ if is_training:
+ filter_instances_position_by_score_fn = functools.partial(
+ utils.filter_instances_box_by_score,
+ box_key_prefix='instances',
+ score_threshold=self._TRAIN_DETECTION_SCORE)
+ self.preprocessor_builder.add_fn(
+ fn=filter_instances_position_by_score_fn,
+ fn_name='filter_training_boxes')
+ else:
+ self.preprocessor_builder.add_fn(
+ fn=lambda x: utils.infer_instances_mask_from_position(inputs=x),
+ fn_name='infer_instances_mask')
+
+ # Add detected boxes.
+ if (not is_training) and import_detected_bboxes:
+ for dim in ['ymin', 'xmin', 'ymax', 'xmax', 'score']:
+ utils.add_instance_box_dim(
+ parser_builder=self.parser_builder,
+ sampler_builder=self.sampler_builder,
+ preprocessor_builder=self.preprocessor_builder,
+ input_box_dim_name='{}/bbox/{}'.format(self._DETECTOR_PREFIX, dim),
+ output_box_dim_name='detected_instances_{}'.format(dim),
+ sample_around_keyframe=True,
+ sample_random=False,
+ num_instances_per_frame=num_instance_per_frame,
+ num_frames=num_frames,
+ temporal_stride=temporal_stride,
+ sync_random_state=True)
+ utils.group_instance_box_dims(
+ preprocessor_builder=self.preprocessor_builder,
+ box_key_prefix='detected_instances',
+ output_position_name='detected_instances_position')
+ filter_instances_position_by_score_fn = functools.partial(
+ utils.filter_instances_box_by_score,
+ box_key_prefix='detected_instances',
+ score_threshold=self._EVAL_DETECTION_SCORE)
+ self.preprocessor_builder.add_fn(
+ fn=filter_instances_position_by_score_fn,
+ fn_name='filter_detected_boxes')
+
+ # Add images.
+ utils.add_image(
+ parser_builder=self.parser_builder,
+ sampler_builder=self.sampler_builder,
+ decoder_builder=self.decoder_builder,
+ preprocessor_builder=self.preprocessor_builder,
+ postprocessor_builder=self.postprocessor_builder,
+ sample_around_keyframe=True,
+ is_training=is_training,
+ num_frames=num_frames,
+ temporal_stride=temporal_stride,
+ num_test_clips=num_test_clips,
+ crop_size=crop_size,
+ min_resize=min_resize,
+ multi_crop=False,
+ zero_centering_image=zero_centering_image,
+ augmentation_type=augmentation_type,
+ augmentation_params=augmentation_params,
+ sync_random_state=True)
+
+ # Adapt boxes to the image augmentations.
+ utils.adjust_positions(
+ preprocessor_builder=self.preprocessor_builder,
+ input_tensor_name='instances_position',
+ output_tensor_name='instances_position')
+ if import_detected_bboxes:
+ utils.adjust_positions(
+ preprocessor_builder=self.preprocessor_builder,
+ input_tensor_name='detected_instances_position',
+ output_tensor_name='detected_instances_position')
+
+ utils.add_context_label(
+ parser_builder=self.parser_builder,
+ sampler_builder=self.sampler_builder,
+ preprocessor_builder=self.preprocessor_builder,
+ input_label_index_feature_name=f'{self._GT_PREFIX}/bbox/label/index',
+ input_label_name_feature_name=f'{self._GT_PREFIX}/bbox/label/string',
+ num_instances_per_frame=num_instance_per_frame,
+ num_frames=num_frames,
+ zero_based_index=self._ZERO_BASED_INDEX,
+ # merge_multi_labels fn expects label in 0-index id.
+ one_hot_label=False if merge_multi_labels else one_hot_label,
+ num_classes=self._NUM_CLASSES,
+ add_label_name=False)
+ if merge_multi_labels:
+ self.preprocessor_builder.add_fn(
+ fn=functools.partial(
+ utils.merge_multi_labels,
+ num_classes=self._NUM_CLASSES),
+ fn_name='merge_multi_labels')
+
+ if is_training and color_augmentation:
+ utils.apply_default_color_augmentations(
+ preprocessor_builder=self.preprocessor_builder,
+ zero_centering_image=zero_centering_image)
+
+ self.postprocessor_builder.add_fn(
+ fn=utils.update_valid_instances_mask,
+ fn_name='update_valid_instances_mask')
+
+ select_keyframe_instances_fn = functools.partial(
+ self._select_keyframe_instances,
+ keyframe_index=(num_frames // 2),
+ import_detected_bboxes=import_detected_bboxes)
+ self.postprocessor_builder.add_fn(
+ fn=select_keyframe_instances_fn,
+ fn_name='slice_keyframe_instances')
+
+ def _select_keyframe_instances(
+ self,
+ inputs: Dict[str, tf.Tensor],
+ keyframe_index: int,
+ import_detected_bboxes: bool) -> Dict[str, tf.Tensor]:
+ """Slices only instance-related inputs on keyframes.
+
+ Args:
+ inputs: The inputs dictionary containing instance related tensors.
+ Tensors' rank should be >= 3, with the order of
+ [batch, time, instances, ...].
+ keyframe_index: The integar local index for the keyframe. Typically this
+ is the middle frame in the 5D tensor.
+ import_detected_bboxes: Whether the pipeline has imported the detected
+ instances boxes.
+
+ Returns:
+ The returned dictionary containing only keyframe instances inputs.
+ """
+ instances_name_list = [
+ 'label', 'instances_position', 'instances_score', 'instances_mask',
+ 'nonmerge_label', 'nonmerge_instances_position'
+ ]
+ if import_detected_bboxes:
+ instances_name_list += [
+ 'detected_instances_position', 'detected_instances_score',
+ 'detected_instances_mask'
+ ]
+ for name in instances_name_list:
+ tensor = inputs[name][:, keyframe_index, ...]
+ inputs[name] = tensor
+ return inputs
+
+ def tables(self) -> Mapping[str, Union[str, List[str]]]:
+ """Returns a dictionary from table name to relative path."""
+ return self._TABLES
+
+
+class AVAKineticsFactory(ActionLocalizationBaseFactory):
+ """AVA-Kinetics data reader."""
+
+ _BASE_DIR = '/abc'
+ _TABLES = {
+ 'train': 'tfse_avakinetics-train.tfrecord@1000',
+ 'test': 'tfse_avakinetics-val.tfrecord@1000',
+ }
+ _SPLITS = ('train', 'test')
+
+ # In AVA-K we use centernet detector. Choose lower score threshold for eval.
+ _TRAIN_DETECTION_SCORE = 0.9
+ _EVAL_DETECTION_SCORE = 0.2
+
+
+class AVAFactory(ActionLocalizationBaseFactory):
+ """AVA v2.2 data reader."""
+
+ _BASE_DIR = '/abc'
+ _TABLES = {
+ 'train': 'tfse_ava_v2.2_train.tfrecord@1000',
+ 'test': 'tfse_ava_v2.2_val.tfrecord@1000',
+ }
+ _SPLITS = ('train', 'test')
+
+ _KEYFRAME_INDEX_KEY = 'key_frame/frame_index'
+ _DETECTOR_PREFIX = 'detectron_frcnn_person/region'
+ _ZERO_BASED_INDEX = True
diff --git a/official/projects/videoglue/datasets/common/processors.py b/official/projects/videoglue/datasets/common/processors.py
new file mode 100644
index 00000000000..f7053c781b4
--- /dev/null
+++ b/official/projects/videoglue/datasets/common/processors.py
@@ -0,0 +1,789 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utils for processing datasets features."""
+
+from typing import Any, MutableMapping, Optional, Tuple
+
+from dmvr import processors
+import simclr.data_util as simclr_data
+import tensorflow as tf, tf_keras
+
+sample_sequence = processors.sample_sequence
+sample_linsapce_sequence = processors.sample_linspace_sequence
+decode_jpeg = processors.decode_jpeg
+random_flip_left_right = processors.random_flip_left_right
+normalize_image = processors.normalize_image
+
+_VGG_EXPANSION_RATIO = 1.25
+
+
+def update_image_info(state: MutableMapping[str, tf.Tensor],
+ current_image_info: tf.Tensor):
+ """Updates image info by merging current augmentation into the last one.
+
+ NOTE: this is not a generic-purposed function and should only be used in this
+ codebase. In this function, image_info tensor encodes the information of the
+ image and the applied preprocessing. It is in the format of
+ [[original_height, original_width],
+ [desired_height, desired_width],
+ [y_scale, x_scale],
+ [y_offset, x_offset]],
+ where [desired_height, desired_width] is the actual scaled image size;
+ [y_scale, x_scale] is the scaling factor, which is the ratio of scaled
+ dimension / original dimension; [y_offset, x_offset] is the upper-left
+ coordinates to perform image cropping.
+
+ Args:
+ state: dict containing 'image_info' of historical image augmentation.
+ current_image_info: the data augmentation info of current step.
+
+ Returns:
+ state: updated dict of 'image_info'
+ """
+ image_info = state.pop('image_info')
+ # image_info in shape [4, 2]
+ i1, t1, s1, o1 = tf.unstack(image_info, axis=0)
+ i2, t2, s2, o2 = tf.unstack(current_image_info, axis=0)
+
+ # last target image_size(t1) should equalt to current input image_size(i2).
+ tf.debugging.assert_equal(t1, i2,
+ message='last target size != current input size. '
+ 'wrong augmentation order?')
+ i3 = i1
+ t3 = t2
+ s3 = s1 * s2
+ o3 = o1 * s2 + o2
+ new_image_info = tf.stack([i3, t3, s3, o3], axis=0)
+ state['image_info'] = new_image_info
+
+
+def multi_crop_image(frames: tf.Tensor,
+ target_height: int,
+ target_width: int) -> tf.Tensor:
+ """Three uniform crops of the image sequence.
+
+ If requested size is bigger than image size, image is padded with 0.
+
+ Args:
+ frames: A Tensor of dimension [timesteps, in_height, in_width, channels].
+ target_height: Target cropped image height.
+ target_width: Target cropped image width.
+
+ Returns:
+ A Tensor of shape [timesteps, out_height, out_width, channels] of type uint8
+ with the cropped images.
+ """
+ shape = tf.shape(frames)
+ static_shape = frames.shape.as_list()
+ seq_len = shape[0] if static_shape[0] is None else static_shape[0]
+ height = shape[1] if static_shape[1] is None else static_shape[1]
+ width = shape[2] if static_shape[2] is None else static_shape[2]
+ channels = shape[3] if static_shape[3] is None else static_shape[3]
+
+ size = tf.convert_to_tensor(
+ (seq_len, target_height, target_width, channels))
+
+ offset_1 = tf.broadcast_to([0, 0, 0, 0], [4])
+
+ portrait_offset_2 = tf.cast(height, tf.float32) / 2 - target_height // 2
+ landscape_offset_2 = tf.cast(width, tf.float32) / 2 - target_width // 2
+ offset_2 = tf.cond(
+ tf.greater_equal(height, width),
+ true_fn=lambda: tf.broadcast_to([0, portrait_offset_2, 0, 0], [4]),
+ false_fn=lambda: tf.broadcast_to([0, 0, landscape_offset_2, 0], [4]))
+
+ portrait_offset_3 = tf.cast(height, tf.float32) - target_height
+ landscape_offset_3 = tf.cast(width, tf.float32) - target_width
+ offset_3 = tf.cond(
+ tf.greater_equal(height, width),
+ true_fn=lambda: tf.broadcast_to([0, portrait_offset_3, 0, 0], [4]),
+ false_fn=lambda: tf.broadcast_to([0, 0, landscape_offset_3, 0], [4]))
+
+ crops = []
+ for offset in [offset_1, offset_2, offset_3]:
+ offset = tf.cast(tf.math.round(offset), tf.int32)
+ crops.append(tf.slice(frames, offset, size))
+ frames = tf.concat(crops, axis=0)
+ return frames
+
+
+def resize_and_crop(
+ frames: tf.Tensor,
+ min_resize: int,
+ crop_size: int,
+ is_flow: bool = False,
+ is_random: bool = False,
+ seed: Optional[int] = None,
+ state: Optional[MutableMapping[str, Any]] = None) -> tf.Tensor:
+ """Resizes the smallest and crops frames.
+
+ Args:
+ frames: A Tensor of dimension [timesteps, input_h, input_w, channels].
+ min_resize: Minimum size of the final image dimensions.
+ crop_size: Crop size of the final image dimensions.
+ is_flow: If is flow, will modify the raw values to account for the resize.
+ For example, if the flow image is resized by a factor k, we need to
+ multiply the flow values by the same factor k since one pixel displacement
+ in the resized image corresponds to only 1/k pixel displacement in the
+ original image.
+ is_random: Whether perform random crop or central crop.
+ seed: Random seed.
+ state: the dictionary contains data processing states.
+ Returns:
+ A Tensor of shape [timesteps, output_h, output_w, channels] of type
+ frames.dtype where min(output_h, output_w) = min_resize.
+ """
+ if is_flow and frames.dtype != tf.float32:
+ raise ValueError('If is_flow, frames should be given in float32.')
+
+ if min_resize < crop_size:
+ raise ValueError('min_resize should be larger than crop_size. Got '
+ f'({min_resize}, {crop_size}).')
+
+ if is_random:
+ min_resize = tf.random.uniform((),
+ minval=min_resize,
+ maxval=_VGG_EXPANSION_RATIO * min_resize,
+ dtype=tf.float32)
+
+ shape = tf.shape(input=frames)
+ image_size = tf.cast(shape[1:3], tf.float32)
+ input_h = image_size[0]
+ input_w = image_size[1]
+
+ scale = tf.cast(min_resize / input_h, tf.float32)
+ scale = tf.maximum(scale, tf.cast(min_resize / input_w, tf.float32))
+
+ scale_h = input_h * scale
+ scale_w = input_w * scale
+
+ def resize_fn():
+ """Function wraper to perform bilinear image resizing."""
+ frames_resized = tf.image.resize(
+ frames, (scale_h, scale_w), method=tf.image.ResizeMethod.BILINEAR)
+ return tf.cast(frames_resized, frames.dtype)
+
+ should_resize = tf.math.logical_or(tf.not_equal(input_w, scale_w),
+ tf.not_equal(input_h, scale_h))
+ frames = tf.cond(
+ pred=should_resize, true_fn=resize_fn, false_fn=lambda: frames)
+
+ if is_flow:
+ # Apply a multiplier to keep the right magnitude in the flow.
+ frames = frames * tf.cast(scale_h / input_h, tf.float32)
+
+ shape = tf.shape(input=frames)
+ image_size = tf.cast(shape[1:3], tf.float32)
+ # If a static_shape is available (e.g. when using this method from add_image
+ # method), it will be used to have an output tensor with static shape.
+ static_shape = frames.shape.as_list()
+ seq_len = shape[0] if static_shape[0] is None else static_shape[0]
+ channels = shape[3] if static_shape[3] is None else static_shape[3]
+ size = tf.convert_to_tensor(value=(seq_len, crop_size, crop_size, channels))
+ if is_random:
+ # Limit of possible offset in order to fit the entire crop:
+ # [1, input_h - target_h + 1, input_w - target_w + 1, 1].
+ limit = shape - size + 1
+ offset = tf.random.uniform(
+ shape=(4,),
+ dtype=tf.int32,
+ maxval=tf.int32.max,
+ seed=seed) % limit # [0, offset_h, offset_w, 0]
+ else:
+ # Central spatial crop.
+ offset = tf.convert_to_tensor(
+ (0, tf.cast((image_size[0] - crop_size) / 2, dtype=tf.int32),
+ tf.cast((image_size[1] - crop_size) / 2, dtype=tf.int32), 0))
+
+ frames = tf.slice(frames, offset, size)
+
+ if state is not None:
+ # Note: image_info encodes the information of the image and the applied
+ # preprocessing. It is in the format of
+ # [[original_height, original_width], [desired_height, desired_width],
+ # [y_scale, x_scale], [y_offset, x_offset]],
+ # where [desired_height, desired_width] is the actual scaled image size,
+ # and [y_scale, x_scale] is the scaling factor, which is the ratio of scaled
+ # dimension / original dimension. [y_offset, x_offset] is the upper-left
+ # coordinates to perform image cropping.
+ image_info = tf.stack([
+ tf.convert_to_tensor((input_h, input_w), tf.float32),
+ tf.convert_to_tensor((crop_size, crop_size), tf.float32),
+ tf.convert_to_tensor((scale, scale), tf.float32),
+ tf.cast(offset[1:3], tf.float32)])
+
+ if 'image_info' not in state:
+ state['image_info'] = image_info
+ else:
+ update_image_info(state, image_info)
+
+ return frames
+
+
+def random_crop_resize(frames: tf.Tensor,
+ output_height: int,
+ output_width: int,
+ num_frames: int,
+ num_channels: int,
+ aspect_ratio: Tuple[float, float],
+ area_range: Tuple[float, float],
+ state: MutableMapping[str, tf.Tensor]) -> tf.Tensor:
+ """First crops clip with jittering and then resizes.
+
+ Args:
+ frames: A Tensor of dimension [timesteps, input_h, input_w, channels].
+ output_height: Resized image height.
+ output_width: Resized image width.
+ num_frames: Number of input frames per clip.
+ num_channels: Number of channels of the clip.
+ aspect_ratio: Float tuple with the aspect range for cropping.
+ area_range: Float tuple with the area range for cropping.
+ state: A mutable dictionary passed to the stateful functions and might be
+ modified in order to keep metadata.
+ Returns:
+ A Tensor of shape [timesteps, output_height, output_width, channels] of type
+ frames.dtype.
+ """
+ shape = tf.shape(frames)
+ image_size = tf.cast(shape[1:3], tf.float32)
+ seq_len, _, _, channels = shape[0], shape[1], shape[2], shape[3]
+ bbox = tf.constant([0.0, 0.0, 1.0, 1.0], dtype=tf.float32, shape=[1, 1, 4])
+ factor = output_width / output_height
+ aspect_ratio = (aspect_ratio[0] * factor, aspect_ratio[1] * factor)
+ sample_distorted_bbox = tf.image.sample_distorted_bounding_box(
+ shape[1:],
+ bounding_boxes=bbox,
+ min_object_covered=0.1,
+ aspect_ratio_range=aspect_ratio,
+ area_range=area_range,
+ max_attempts=100,
+ use_image_if_no_bounding_boxes=True)
+ bbox_begin, bbox_size, _ = sample_distorted_bbox
+ offset_y, offset_x, _ = tf.unstack(bbox_begin)
+ target_height, target_width, _ = tf.unstack(bbox_size)
+ size = tf.convert_to_tensor((
+ seq_len, target_height, target_width, channels))
+ offset = tf.convert_to_tensor((0, offset_y, offset_x, 0))
+ frames = tf.slice(frames, offset, size)
+ frames = tf.cast(
+ tf.image.resize(frames, (output_height, output_width)),
+ frames.dtype)
+ frames.set_shape((num_frames, output_height, output_width, num_channels))
+ image_scale = (
+ tf.convert_to_tensor((output_height, output_width), tf.float32) /
+ tf.cast(bbox_size[:2], tf.float32))
+
+ # Note: image_info encodes the information of the image and the applied
+ # preprocessing. It is in the format of
+ # [[original_height, original_width], [desired_height, desired_width],
+ # [y_scale, x_scale], [y_offset, x_offset]],
+ # where [desired_height, desired_width] is the actual scaled image size,
+ # and [y_scale, x_scale] is the scaling factor, which is the ratio of scaled
+ # dimension / original dimension. [y_offset, x_offset] is the upper-left
+ # coordinates to perform image cropping.
+ offset = tf.convert_to_tensor((offset_y, offset_x), tf.float32)
+ offset *= image_scale
+ image_info = tf.stack([
+ image_size,
+ tf.convert_to_tensor((output_height, output_width), tf.float32),
+ image_scale,
+ offset])
+ if 'image_info' not in state:
+ state['image_info'] = image_info
+ else:
+ update_image_info(state, image_info)
+ return frames
+
+
+def resize_smallest(
+ frames: tf.Tensor,
+ min_resize: int,
+ is_flow: bool = False,
+ state: Optional[MutableMapping[str, tf.Tensor]] = None) -> tf.Tensor:
+ """Resizes frames so that `min(height, width)` is equal to `min_resize`.
+
+ This function will do nothing if the `min(height, width)` is already equal to
+ `min_resize`. This allows to save compute time.
+
+ Args:
+ frames: A tensor of dimension [timesteps, input_h, input_w, channels].
+ min_resize: Minimum size of the final image dimensions.
+ is_flow: If is flow, will modify the raw values to account for the resize.
+ For example, if the flow image is resized by a factor k, we need to
+ multiply the flow values by the same factor k since one pixel displacement
+ in the resized image corresponds to only 1/k pixel displacement in the
+ original image.
+ state: A mutable dictionary passed to the stateful functions and might be
+ modified in order to keep metadata.
+
+ Returns:
+ A tensor of shape [timesteps, output_h, output_w, channels] of same type as
+ input, where `min(output_h, output_w)` is `min_resize`.
+ """
+ if is_flow and frames.dtype != tf.float32:
+ raise ValueError('If `is_flow`, frames should be given in `tf.float32`.')
+ shape = tf.shape(input=frames)
+ input_h = shape[1]
+ input_w = shape[2]
+
+ output_h = tf.maximum(min_resize, (input_h * min_resize) // input_w)
+ output_w = tf.maximum(min_resize, (input_w * min_resize) // input_h)
+
+ def resize_fn():
+ """Bilinear resize function wrapper."""
+ frames_resized = tf.image.resize(
+ frames, (output_h, output_w), method=tf.image.ResizeMethod.BILINEAR)
+ return tf.cast(frames_resized, frames.dtype)
+
+ should_resize = tf.math.logical_or(tf.not_equal(input_w, output_w),
+ tf.not_equal(input_h, output_h))
+ frames = tf.cond(
+ pred=should_resize, true_fn=resize_fn, false_fn=lambda: frames)
+
+ if is_flow:
+ # Apply a multiplier to keep the right magnitude in the flow.
+ frames = frames * tf.cast(output_h / input_h, tf.float32)
+
+ if state is not None:
+ # Note: image_info encodes the information of the image and the applied
+ # preprocessing. It is in the format of
+ # [[original_height, original_width], [desired_height, desired_width],
+ # [y_scale, x_scale], [y_offset, x_offset]],
+ # where [desired_height, desired_width] is the actual scaled image size,
+ # and [y_scale, x_scale] is the scaling factor, which is the ratio of scaled
+ # dimension / original dimension.
+ image_info = tf.stack([
+ tf.convert_to_tensor((input_h, input_w), tf.float32),
+ tf.convert_to_tensor((output_h, output_w), tf.float32),
+ tf.convert_to_tensor(
+ (output_h / input_h, output_w / input_w), tf.float32),
+ tf.zeros([2], tf.float32)])
+
+ if 'image_info' not in state:
+ state['image_info'] = image_info
+ else:
+ update_image_info(state, image_info)
+
+ return frames
+
+
+def random_square_crop_by_scale(
+ image: tf.Tensor,
+ max_border: int = 128,
+ scale_min: float = 0.6,
+ scale_max: float = 1.3,
+ num_scales: int = 8,
+ seed: Optional[int] = None,
+ state: Optional[MutableMapping[str, tf.Tensor]] = None) -> tf.Tensor:
+ """Randomly crops a square in proportion to scale and image size.
+
+ Extract a square sized crop from an image whose side length is sampled by
+ randomly scaling the maximum spatial dimension of the image. If part of
+ the crop falls outside the image, it is filled with zeros.
+ The augmentation is borrowed from [1]: https://arxiv.org/abs/1904.07850
+
+ Args:
+ image: rank 4 float32 tensor containing images in shape
+ [time, height, width, channels].
+ max_border: The maximum size of the border. The border defines distance in
+ pixels to the image boundaries that will not be considered as a center of
+ a crop. To make sure that the border does not go over the center of the
+ image, we chose the border value by computing the minimum k, such that
+ (max_border / (2**k)) < image_dimension/2.
+ scale_min: The minimum value for scale.
+ scale_max: The maximum value for scale.
+ num_scales: The number of discrete scale values to sample between
+ [scale_min, scale_max]
+ seed: Random seed.
+ state: The dictionary contains random state.
+
+ Returns:
+ output_image: image which is the same rank as input image.
+ """
+
+ def _random_integer(minval, maxval, seed):
+ return tf.random.uniform(
+ [], minval=minval, maxval=maxval, dtype=tf.int32, seed=seed)
+
+ def _get_crop_border(border, size):
+ border = tf.cast(border, tf.float32)
+ size = tf.cast(size, tf.float32)
+
+ i = tf.math.ceil(tf.math.log(2.0 * border / size) / tf.math.log(2.0))
+ divisor = tf.math.pow(2.0, i)
+ divisor = tf.clip_by_value(divisor, 1, border)
+ divisor = tf.cast(divisor, tf.int32)
+ return tf.cast(border, tf.int32) // divisor
+
+ img_shape = tf.shape(image)
+ height, width = img_shape[1], img_shape[2]
+ scales = tf.linspace(scale_min, scale_max, num_scales)
+ random_id = _random_integer(0, num_scales, seed=seed)
+ scale = scales[random_id]
+
+ image_size = scale * tf.cast(tf.maximum(height, width), tf.float32)
+ image_size = tf.cast(image_size, tf.int32)
+ h_border = _get_crop_border(max_border, height)
+ w_border = _get_crop_border(max_border, width)
+
+ y_center = _random_integer(h_border,
+ tf.cast(height, tf.int32) - h_border + 1, seed)
+
+ x_center = _random_integer(w_border,
+ tf.cast(width, tf.int32) - w_border + 1, seed)
+
+ half_size = tf.cast(image_size / 2, tf.int32)
+ crop_ymin, crop_ymax = y_center - half_size, y_center + half_size
+ crop_xmin, crop_xmax = x_center - half_size, x_center + half_size
+
+ ymin = tf.maximum(crop_ymin, 0)
+ xmin = tf.maximum(crop_xmin, 0)
+ ymax = tf.minimum(crop_ymax, height - 1)
+ xmax = tf.minimum(crop_xmax, width - 1)
+
+ cropped_image = image[:, ymin:ymax, xmin:xmax, :]
+ offset_y = tf.maximum(0, ymin - crop_ymin)
+ offset_x = tf.maximum(0, xmin - crop_xmin)
+
+ output_image = tf.image.pad_to_bounding_box(
+ cropped_image, offset_height=offset_y, offset_width=offset_x,
+ target_height=image_size, target_width=image_size)
+
+ if state is not None:
+ # if (padding) else (cropping)
+ padding_y = tf.cast(offset_y > 0, tf.int32)
+ box_offset_y = padding_y * (-offset_y) + (1 - padding_y) * crop_ymin
+ padding_x = tf.cast(offset_x > 0, tf.int32)
+ box_offset_x = padding_x * (-offset_x) + (1 - padding_x) * crop_xmin
+
+ # Note: image_info encodes the information of the image and the applied
+ # preprocessing. It is in the format of
+ # [[original_height, original_width], [desired_height, desired_width],
+ # [y_scale, x_scale], [y_offset, x_offset]],
+ # where [desired_height, desired_width] is the actual scaled image size,
+ # and [y_scale, x_scale] is the scaling factor, which is the ratio of scaled
+ # dimension / original dimension.
+ image_info = tf.stack([
+ tf.convert_to_tensor((height, width), tf.float32),
+ tf.convert_to_tensor((image_size, image_size), tf.float32),
+ tf.convert_to_tensor((1.0, 1.0), tf.float32),
+ tf.convert_to_tensor((box_offset_y, box_offset_x), tf.float32)])
+
+ if 'image_info' not in state:
+ state['image_info'] = image_info
+ else:
+ update_image_info(state, image_info)
+
+ return output_image
+
+
+def resize_and_pad(
+ frames: tf.Tensor,
+ max_resize: int,
+ pad_size: int,
+ random: bool = False,
+ seed: Optional[int] = None,
+ state: Optional[MutableMapping[str, tf.Tensor]] = None) -> tf.Tensor:
+ """Resizes the largest and pads frames.
+
+ Args:
+ frames: A Tensor of dimension [timesteps, input_h, input_w, channels].
+ max_resize: Maximum size of the final image dimensions.
+ pad_size: Pad size of the final image dimensions.
+ random: If true, perform random crop; otherwise, perform central crop.
+ seed: Random seed.
+ state: The dictionary contains random state.
+ Returns:
+ A Tensor of shape [timesteps, output_h, output_w, channels] of type
+ frames.dtype where min(output_h, output_w) = max_resize.
+ """
+ if max_resize > pad_size:
+ raise ValueError('max_resize should not be larger than pad_size. Got '
+ f'({max_resize}, {pad_size}).')
+
+ pad_color = tf.reduce_mean(frames, axis=[0, 1, 2])
+
+ shape = tf.shape(input=frames)
+ image_size = tf.cast(shape[1:3], tf.float32)
+ input_h = image_size[0]
+ input_w = image_size[1]
+
+ scale = tf.cast(max_resize / input_h, tf.float32)
+ scale = tf.minimum(scale, tf.cast(max_resize / input_w, tf.float32))
+
+ scale_h = input_h * scale
+ scale_w = input_w * scale
+
+ frames_resized = tf.image.resize(
+ frames, (scale_h, scale_w), method=tf.image.ResizeMethod.BILINEAR)
+ frames = tf.cast(frames_resized, frames.dtype)
+
+ shape = tf.shape(input=frames)
+ image_size = tf.cast(shape[1:3], tf.float32)
+ size = tf.convert_to_tensor(value=(pad_size, pad_size))
+ if random:
+ # Limit of possible offset in order to fit the entire crop:
+ # [target_h - input_h + 1, target_w - input_w + 1].
+ limit = size - tf.cast(image_size, tf.int32) + 1
+ offset = tf.random.uniform(
+ shape=(2,),
+ dtype=tf.int32,
+ maxval=tf.int32.max,
+ seed=seed) % limit # [offset_h, offset_w]
+ offset_height, offset_width = offset[0], offset[1]
+ else:
+ # Central spatial pad.
+ offset_height = tf.cast((pad_size - image_size[0]) / 2, tf.int32)
+ offset_width = tf.cast((pad_size - image_size[1]) / 2, tf.int32)
+ offset = tf.convert_to_tensor([offset_height, offset_width])
+
+ padded_frames = tf.image.pad_to_bounding_box(
+ frames,
+ offset_height=offset_height,
+ offset_width=offset_width,
+ target_height=pad_size,
+ target_width=pad_size)
+
+ # Setting color of the padded pixels
+ frames_ones = tf.ones_like(frames)
+ frames_ones_padded = tf.image.pad_to_bounding_box(
+ frames_ones,
+ offset_height=offset_height,
+ offset_width=offset_width,
+ target_height=pad_size,
+ target_width=pad_size)
+ frames_color_padded = (1 - frames_ones_padded) * pad_color
+ padded_frames += frames_color_padded
+
+ if state is not None:
+ # Note: image_info encodes the information of the image and the applied
+ # preprocessing. It is in the format of
+ # [[original_height, original_width], [desired_height, desired_width],
+ # [y_scale, x_scale], [y_offset, x_offset]],
+ # where [desired_height, desired_width] is the actual scaled image size,
+ # and [y_scale, x_scale] is the scaling factor, which is the ratio of scaled
+ # dimension / original dimension.
+ image_info = tf.stack([
+ tf.convert_to_tensor((input_h, input_w), tf.float32),
+ tf.convert_to_tensor((pad_size, pad_size), tf.float32),
+ tf.convert_to_tensor((scale, scale), tf.float32),
+ # set to -1. * offset because it's padding.
+ -1. * tf.cast(offset, tf.float32)])
+
+ if 'image_info' not in state:
+ state['image_info'] = image_info
+ else:
+ update_image_info(state, image_info)
+ return padded_frames
+
+
+def crop_or_pad_features(features: tf.Tensor,
+ max_num_features: int,
+ feature_dimension: int,
+ constant_values: int = 0) -> tf.Tensor:
+ """Crops or pads given sequence of features vectors.
+
+ Args:
+ features: Tensor features of shape [T, feature_length] or [feature_length],
+ features of shape [feature_length] is expanded as [1, feature_length].
+ max_num_features: Maximum number of words in final result.
+ feature_dimension: The dimensionality of feature vector.
+ constant_values: The constant value used to padd the input tensor.
+
+ Returns:
+ A Tensor of shape [T, max_num_features * feature_dimension].
+ """
+ if len(features.shape) == 1:
+ features = tf.expand_dims(features, 0)
+
+ num_features = tf.shape(input=features)[1]
+ max_length = max_num_features * feature_dimension
+ paddings = ((0, 0),
+ (0, tf.maximum(0, max_length - num_features)))
+ features = tf.pad(
+ tensor=features[:, :max_length],
+ paddings=paddings,
+ constant_values=constant_values)
+ features.set_shape((None, max_length))
+ return features
+
+
+def sample_sequence_by_segment(
+ inputs: MutableMapping[str, tf.Tensor],
+ num_steps: int,
+ sample_target_key: str = 'image',
+ is_training: bool = True,
+) -> MutableMapping[str, tf.Tensor]:
+ """Samples a single clip from a given sequence by segments.
+
+ Args:
+ inputs: dict with sample_target and keyframe_index features.
+ num_steps: Number of steps (e.g. frames) to take.
+ sample_target_key: the key for the sample target.
+ is_training: whether in training mode.
+
+ Returns:
+ Modified inputs with sampled target.
+ """
+ sequence = inputs[sample_target_key]
+ sequence_length = tf.shape(sequence)[0]
+ segment_size = (tf.cast(sequence_length, tf.float32) - 1) / num_steps
+ indices = []
+ for i in range(num_steps):
+ start = tf.cast(tf.math.round(segment_size * i), tf.int32)
+ end = tf.cast(tf.math.round(segment_size * (i + 1)), tf.int32)
+ # special hanle if end == start.
+ end = tf.maximum(end, start+1)
+ if is_training:
+ indices.append(
+ tf.random.uniform(shape=(), minval=start, maxval=end, dtype=tf.int32))
+ else:
+ indices.append((start + end) // 2)
+ indices = tf.stack(indices, axis=0)
+ output = tf.gather(sequence, indices)
+ inputs[sample_target_key] = output
+ return inputs
+
+
+def sample_sequence_around_keyframe(
+ inputs: MutableMapping[str, tf.Tensor],
+ num_steps: int,
+ stride: int,
+ sample_target_key: str = 'image',
+ keyframe_index_key: str = 'keyframe_index'
+) -> MutableMapping[str, tf.Tensor]:
+ """Samples a single segment around keyframe from a given sequence.
+
+ Args:
+ inputs: dict with sample_target and keyframe_index features.
+ num_steps: Number of steps (e.g. frames) to take.
+ stride: the stride of sampling.
+ sample_target_key: the key for the sample target.
+ keyframe_index_key: the key for the keyframe index.
+
+ Returns:
+ Modified inputs with sampled target.
+ """
+ if sample_target_key not in inputs:
+ raise ValueError(f'{sample_target_key} is not found in input dictionary.')
+ if keyframe_index_key not in inputs:
+ raise ValueError(f'{keyframe_index_key} is not found in input dictionary.')
+
+ sequence = inputs[sample_target_key]
+
+ keyframe_index = tf.cast(inputs[keyframe_index_key], dtype=tf.int32)
+ keyframe_index = tf.squeeze(keyframe_index) # assuming there's a single one
+
+ sequence_length = tf.shape(sequence)[0]
+ sequence_length = tf.cast(sequence_length, tf.int32)
+
+ early = keyframe_index - (num_steps * stride) // 2
+ late = keyframe_index + (num_steps * stride) // 2
+ offset = tf.maximum(0, early)
+
+ pad_before = tf.maximum(0, -early)
+ pad_after = tf.maximum(0, late - sequence_length)
+
+ # Repeat first and last frames appropriately.
+ repeat = sequence.shape.ndims - 1
+ pad_before = [pad_before] + [1] * repeat
+ pad_before_frame = tf.tile(sequence[:1], pad_before)
+ pad_after = [pad_after] + [1] * repeat
+ pad_after_frame = tf.tile(sequence[-1:], pad_after)
+ sequence = tf.concat([pad_before_frame, sequence, pad_after_frame], axis=0)
+
+ indices = tf.linspace(offset, offset + (num_steps - 1) * stride, num_steps)
+ indices = tf.cast(indices, dtype=tf.int32)[:num_steps]
+ indices.set_shape((num_steps))
+
+ output = tf.gather(sequence, indices)
+ inputs[sample_target_key] = output
+ return inputs
+
+
+def random_color_augmentation(frames: tf.Tensor,
+ zero_centering_image: bool = False,
+ color_jitter_prob: float = 0.8,
+ color_drop_prob: float = 0.0) -> tf.Tensor:
+ """Standard color augmentation for video.
+
+ Args:
+ frames: the input video frames.
+ zero_centering_image: Whether the image frames has been zero centered.
+ color_jitter_prob: The probability to apply color jittering.
+ color_drop_prob: The probability to apply color dropping.
+
+ Returns:
+ The frames with color augmentations.
+ """
+
+ def color_jitter_fn(video):
+ """Does the color augmentations."""
+ if zero_centering_image:
+ video = 0.5 * (video + 1.0)
+ video = tf.image.random_brightness(video, max_delta=32.0 / 255.0)
+ video = tf.image.random_saturation(video, lower=0.6, upper=1.4)
+ video = tf.image.random_contrast(video, lower=0.6, upper=1.4)
+ video = tf.image.random_hue(video, max_delta=0.2)
+ video = tf.clip_by_value(video, 0.0, 1.0)
+ if zero_centering_image:
+ video = 2 * (video - 0.5)
+ return video
+
+ def color_drop_fn(video):
+ """Does the color drop."""
+ video = tf.image.rgb_to_grayscale(video)
+ video = tf.tile(video, [1, 1, 1, 3])
+ return video
+
+ frames = simclr_data.random_apply(color_jitter_fn, color_jitter_prob, frames)
+ frames = simclr_data.random_apply(color_drop_fn, color_drop_prob, frames)
+ return frames
+
+
+def random_blur_and_solarize(
+ frames: tf.Tensor,
+ zero_centering_image: bool = False,
+ blur_prob: float = 1.0,
+ solarize_prob: float = 0.2) -> tf.Tensor:
+ """Randomly blur and solarize a video clip.
+
+ Args:
+ frames: The input image frames tensor.
+ zero_centering_image: Whether the images are zero centered.
+ blur_prob: The probability to apply random bluring.
+ solarize_prob: The probability to apply random solarization.
+
+ Returns:
+ The images with random bluriness and solarization.
+ """
+ if zero_centering_image:
+ frames = 0.5 * (frames + 1.0)
+
+ height = tf.shape(frames)[1]
+ width = tf.shape(frames)[2]
+ frames = simclr_data.random_blur(frames, height, width, blur_prob)
+
+ def solarize_fn(image):
+ """Randomly solarizes images."""
+ image = image * tf.cast(tf.less(image, 0.5), tf.float32) + (
+ 1.0 - image) * tf.cast(tf.greater_equal(image, 0.5), tf.float32)
+ return image
+ frames = simclr_data.random_apply(solarize_fn, solarize_prob, frames)
+ frames = tf.clip_by_value(frames, 0.0, 1.0)
+
+ if zero_centering_image:
+ frames = 2 * (frames - 0.5)
+ return frames
diff --git a/official/projects/videoglue/datasets/common/utils.py b/official/projects/videoglue/datasets/common/utils.py
new file mode 100644
index 00000000000..67eba6a95df
--- /dev/null
+++ b/official/projects/videoglue/datasets/common/utils.py
@@ -0,0 +1,1278 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utils for processing data."""
+import functools
+from typing import Any, Mapping, Optional, MutableMapping, Dict
+
+from dmvr import builders
+from dmvr import modalities
+import tensorflow as tf, tf_keras
+
+from official.projects.videoglue.datasets.common import processors
+from official.vision.ops import augment
+from official.vision.ops import box_ops
+
+add_label = modalities.add_label
+
+
+def pad_or_clip_nd(tensor: tf.Tensor, output_shape: list[int]):
+ """Pads or clips given tensor to the output shape.
+
+ Args:
+ tensor: Input tensor to pad or clip.
+ output_shape: A list of integers / scalar tensors (or None for dynamic dim)
+ representing the size to pad or clip each dimension of the input tensor.
+
+ Returns:
+ Input tensor padded and clipped to the output shape.
+ """
+ tensor_shape = tf.shape(tensor)
+ clip_size = [
+ tf.where(tensor_shape[i] - shape > 0, shape, -1)
+ if shape is not None else -1 for i, shape in enumerate(output_shape)
+ ]
+ clipped_tensor = tf.slice(
+ tensor,
+ begin=tf.zeros(len(clip_size), dtype=tf.int32),
+ size=clip_size)
+
+ # Pad tensor if the shape of clipped tensor is smaller than the expected
+ # shape.
+ clipped_tensor_shape = tf.shape(clipped_tensor)
+ trailing_paddings = [
+ shape - clipped_tensor_shape[i] if shape is not None else 0
+ for i, shape in enumerate(output_shape)
+ ]
+ paddings = tf.stack(
+ [
+ tf.zeros(len(trailing_paddings), dtype=tf.int32),
+ trailing_paddings
+ ],
+ axis=1)
+ padded_tensor = tf.pad(clipped_tensor, paddings=paddings)
+ output_static_shape = [
+ dim if not isinstance(dim, tf.Tensor) else None for dim in output_shape
+ ]
+ padded_tensor.set_shape(output_static_shape)
+ return padded_tensor
+
+
+def merge_boxes_with_multiple_labels(boxes: tf.Tensor,
+ classes: tf.Tensor,
+ confidences: tf.Tensor,
+ num_classes: int,
+ quantization_bins: int = 10000):
+ """Merges boxes with same coordinates and returns K-hot encoded classes.
+
+ Args:
+ boxes: A tf.float32 tensor with shape [N, 4] holding N boxes. Only
+ normalized coordinates are allowed.
+ classes: A tf.int32 tensor with shape [N] holding class indices.
+ The class index starts at 0.
+ confidences: A tf.float32 tensor with shape [N] holding class confidences.
+ num_classes: total number of classes to use for K-hot encoding.
+ quantization_bins: the number of bins used to quantize the box coordinate.
+
+ Returns:
+ merged_boxes: A tf.float32 tensor with shape [N', 4] holding boxes,
+ where N' <= N.
+ class_encodings: A tf.int32 tensor with shape [N', num_classes] holding
+ K-hot encodings for the merged boxes.
+ confidence_encodings: A tf.float32 tensor with shape [N', num_classes]
+ holding encodings of confidences for the merged boxes.
+ merged_box_indices: A tf.int32 tensor with shape [N'] holding original
+ indices of the boxes.
+ """
+
+ def _assert_shape_equal_along_first_dimension(shape_a, shape_b):
+ """Asserts that shape_a and shape_b are the same along the 0th-dimension.
+
+ If the shapes are static, raises a ValueError when the shapes
+ mismatch.
+
+ If the shapes are dynamic, raises a tf InvalidArgumentError when the shapes
+ mismatch.
+
+ Args:
+ shape_a: a list containing shape of the first tensor.
+ shape_b: a list containing shape of the second tensor.
+
+ Returns:
+ Either a tf.no_op() when shapes are all static and a tf.assert_equal() op
+ when the shapes are dynamic.
+
+ Raises:
+ ValueError: When shapes are both static and unequal.
+ """
+ if isinstance(shape_a[0], int) and isinstance(shape_b[0], int):
+ if shape_a[0] != shape_b[0]:
+ raise ValueError('Unequal first dimension {}, {}'.format(
+ shape_a[0], shape_b[0]))
+ else: return tf.no_op()
+ else:
+ return tf.assert_equal(shape_a[0], shape_b[0])
+
+ def _assert_box_normalized(boxes, maximum_normalized_coordinate=1.1):
+ """Asserts the input box tensor is normalized.
+
+ Args:
+ boxes: a tensor of shape [N, 4] where N is the number of boxes.
+ maximum_normalized_coordinate: Maximum coordinate value to be considered
+ as normalized, default to 1.1.
+
+ Returns:
+ a tf.Assert op which fails when the input box tensor is not normalized.
+
+ Raises:
+ ValueError: When the input box tensor is not normalized.
+ """
+ box_minimum = tf.reduce_min(boxes)
+ box_maximum = tf.reduce_max(boxes)
+ return tf.Assert(
+ tf.logical_and(
+ tf.less_equal(box_maximum, maximum_normalized_coordinate),
+ tf.greater_equal(box_minimum, 0)),
+ [boxes])
+
+ boxes_shape = tf.shape(boxes)
+ classes_shape = tf.shape(classes)
+ confidences_shape = tf.shape(confidences)
+ box_class_shape_assert = _assert_shape_equal_along_first_dimension(
+ boxes_shape, classes_shape)
+ box_confidence_shape_assert = _assert_shape_equal_along_first_dimension(
+ boxes_shape, confidences_shape)
+ box_dimension_assert = tf.assert_equal(boxes_shape[1], 4)
+ box_normalized_assert = _assert_box_normalized(boxes)
+
+ with tf.control_dependencies(
+ [box_class_shape_assert, box_confidence_shape_assert,
+ box_dimension_assert, box_normalized_assert]):
+ quantized_boxes = tf.cast(boxes * (quantization_bins - 1), tf.int64)
+ ymin, xmin, ymax, xmax = tf.unstack(quantized_boxes, axis=1)
+ hashcodes = (
+ ymin +
+ xmin * quantization_bins +
+ ymax * quantization_bins * quantization_bins +
+ xmax * quantization_bins * quantization_bins * quantization_bins)
+ unique_hashcodes, unique_indices = tf.unique(hashcodes)
+ num_boxes = tf.shape(boxes)[0]
+ num_unique_boxes = tf.shape(unique_hashcodes)[0]
+ merged_box_indices = tf.math.unsorted_segment_min(
+ tf.range(num_boxes), unique_indices, num_unique_boxes)
+ merged_boxes = tf.gather(boxes, merged_box_indices)
+ unique_indices = tf.cast(unique_indices, tf.int64)
+ classes = tf.cast(classes, tf.int64)
+
+ def map_box_encodings(i):
+ """Produces box K-hot and score encodings for each class index."""
+ box_mask = tf.equal(
+ unique_indices, i * tf.ones(num_boxes, dtype=tf.int64))
+ box_mask = tf.reshape(box_mask, [-1])
+ box_indices = tf.boolean_mask(classes, box_mask)
+ box_confidences = tf.boolean_mask(confidences, box_mask)
+ box_class_encodings = tf.compat.v1.sparse_to_dense(
+ box_indices, [num_classes], tf.constant(1, dtype=tf.int64),
+ validate_indices=False)
+ box_confidence_encodings = tf.compat.v1.sparse_to_dense(
+ box_indices, [num_classes], box_confidences, validate_indices=False)
+ return box_class_encodings, box_confidence_encodings
+
+ # Important to avoid int32 here since there is no GPU kernel for int32.
+ # int64 and float32 are fine.
+ class_encodings, confidence_encodings = tf.map_fn(
+ map_box_encodings,
+ tf.range(tf.cast(num_unique_boxes, tf.int64)),
+ back_prop=False,
+ dtype=(tf.int64, tf.float32))
+
+ merged_boxes = tf.reshape(merged_boxes, [-1, 4])
+ class_encodings = tf.cast(class_encodings, dtype=tf.int32)
+ class_encodings = tf.reshape(class_encodings, [-1, num_classes])
+ confidence_encodings = tf.reshape(confidence_encodings, [-1, num_classes])
+ merged_box_indices = tf.reshape(merged_box_indices, [-1])
+ return (merged_boxes, class_encodings, confidence_encodings,
+ merged_box_indices)
+
+
+def add_image(parser_builder: builders.BaseParserBuilder,
+ sampler_builder: builders.SamplerBuilder,
+ decoder_builder: builders.DecoderBuilder,
+ preprocessor_builder: builders.PreprocessorBuilder,
+ postprocessor_builder: builders.PostprocessorBuilder,
+ input_feature_name: str = 'image/encoded',
+ output_feature_name: str = builders.IMAGE_FEATURE_NAME,
+ is_training: bool = True,
+ sample_around_keyframe: bool = False,
+ sample_from_segments: bool = False,
+ # Video related parameters.
+ num_frames: int = 32,
+ temporal_stride: int = 1,
+ num_test_clips: int = 1,
+ crop_size: int = 200,
+ min_resize: int = 224,
+ multi_crop: bool = False,
+ zero_centering_image: bool = False,
+ random_flip_image: bool = True,
+ augmentation_type: str = 'Inception',
+ augmentation_params: Optional[Mapping[str, Any]] = None,
+ randaug_params: Optional[Mapping[str, Any]] = None,
+ autoaug_params: Optional[Mapping[str, Any]] = None,
+ sync_random_state: bool = True,
+ seed: Optional[int] = None):
+ """Adds functions to process image feature to builders.
+
+ Args:
+ parser_builder: An instance of a builders.BaseParserBuilder.
+ sampler_builder: An instance of a builders.SamplerBuilder.
+ decoder_builder: An instance of a builders.DecoderBuilder.
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ postprocessor_builder: An instance of a builders.PostprocessorBuilder.
+ input_feature_name: Name of the feature in the input SequenceExample.
+ Exposing this as an argument allows using this function for different
+ image features.
+ output_feature_name: Name of the feature in the output features dictionary.
+ Exposing this as an argument allows using this function for different
+ image features.
+ is_training: Whether or not perform random operations. If True, random
+ sample, crop and left right flip is used.
+ sample_around_keyframe: Whether to sample clip around the keyframe. If True,
+ the random temporal sampling will be overridden and disabled.
+ sample_from_segments: Whether to sample frames from segments of a video. If
+ True, the temporal_stride will be ignored.
+ num_frames: Number of frames per subclip.
+ temporal_stride: Temporal stride to sample frames.
+ num_test_clips: Number of test clips (1 by default). If more than 1, this
+ will sample multiple linearly spaced clips within each video at test time.
+ If 1, then a single clip in the middle of the video is sampled. The clips
+ are aggreagated in the batch dimension.
+ crop_size: Final size of the frame after cropping the resized frames. Both
+ height and width are the same.
+ min_resize: The minimal length resize before cropping.
+ multi_crop: Whether to perform 3-view crop or not. This is only enabled
+ in evaluation mode. If is_training=True, this is ignored.
+ zero_centering_image: If True, frames are normalized to values in [-1, 1].
+ If False, values in [0, 1].
+ random_flip_image: If True, frames are randomly horizontal flipped during
+ the training.
+ augmentation_type: The style of Crop+Resize procedure. Support options:
+ ['Inception', 'VGG'].
+ augmentation_params: A dictionary contains image augmentation parameters
+ associated with the augmentation style.
+ randaug_params: A dictionary of params for RandAug policy.
+ autoaug_params: A dictionary of params for AutoAug policy.
+ sync_random_state: Whether to use stateful option to keep random operations
+ in sync between different modalities. All modalities having this option
+ True will use the same outcome in random operations such as sampling and
+ cropping.
+ seed: the random seed.
+ """
+ # Validate parameters.
+ if sync_random_state and multi_crop:
+ raise ValueError('multi_crop is not supported with sync random states.')
+
+ if augmentation_type.lower() == 'ava' and multi_crop:
+ raise ValueError('multi_crop should not be combined with ava augmentation.')
+
+ if sample_from_segments and sample_around_keyframe:
+ raise ValueError('sample_from_segments and sample_around_keyframes cannot '
+ 'be True at the same time.')
+
+ if sample_from_segments and num_test_clips > 1:
+ raise ValueError(
+ 'sample_from_segments is set to True while got num_test_clips: %d'
+ % num_test_clips
+ )
+
+ # Parse frames.
+ if isinstance(parser_builder, builders.SequenceExampleParserBuilder):
+ parser_builder.parse_feature(
+ feature_name=input_feature_name,
+ feature_type=tf.io.FixedLenSequenceFeature((), dtype=tf.string),
+ output_name=output_feature_name)
+ elif isinstance(parser_builder, builders.ExampleParserBuilder):
+ parser_builder.parse_feature(
+ feature_name=input_feature_name,
+ feature_type=tf.io.FixedLenFeature((), dtype=tf.string),
+ output_name=output_feature_name)
+ sampler_builder.add_fn(
+ fn=lambda x: tf.expand_dims(x, axis=0),
+ feature_name=output_feature_name,
+ fn_name=f'{output_feature_name}_expand_dim')
+ else:
+ raise ValueError('`parser_builder` has an unexpected type.')
+
+ if sample_around_keyframe:
+ # Sample clip around keyframe.
+ sample_around_keyframe_fn = functools.partial(
+ processors.sample_sequence_around_keyframe,
+ num_steps=num_frames,
+ stride=temporal_stride,
+ sample_target_key=output_feature_name)
+ sampler_builder.add_fn(
+ fn=sample_around_keyframe_fn,
+ fn_name='{}_sample_around_keyframe'.format(output_feature_name))
+ elif sample_from_segments:
+ sample_segment_fn = functools.partial(
+ processors.sample_sequence_by_segment,
+ num_steps=num_frames,
+ sample_target_key=output_feature_name,
+ is_training=is_training)
+ sampler_builder.add_fn(
+ fn=sample_segment_fn,
+ fn_name='{}_segment_sample'.format(output_feature_name))
+ elif is_training:
+ # Sample random clip.
+ def sample_sequence_fn(x, state):
+ return processors.sample_sequence(
+ x,
+ num_steps=num_frames, random=True, stride=temporal_stride, seed=seed,
+ state=state)
+ sampler_builder.add_fn(
+ fn=sample_sequence_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_random_sample'.format(output_feature_name),
+ stateful=sync_random_state)
+ else:
+ if num_test_clips > 1:
+ sample_linespace_sequence_fn = functools.partial(
+ processors.sample_linsapce_sequence,
+ num_windows=num_test_clips,
+ num_steps=num_frames,
+ stride=temporal_stride)
+ # Sample linspace clips.
+ sampler_builder.add_fn(
+ fn=sample_linespace_sequence_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_linspace_sample'.format(output_feature_name))
+ else:
+ sample_sequence_fn = functools.partial(
+ processors.sample_sequence,
+ num_steps=num_frames, random=False, stride=temporal_stride, seed=None)
+ # Sample middle clip.
+ sampler_builder.add_fn(
+ fn=sample_sequence_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_middle_sample'.format(output_feature_name))
+
+ # Decode JPEG string to tf.uint8.
+ num_raw_channels = 3
+ decoder_builder.add_fn(
+ fn=lambda x: processors.decode_jpeg(x, channels=num_raw_channels),
+ feature_name=output_feature_name,
+ fn_name='{}_decode_jpeg'.format(output_feature_name))
+
+ # Image crop, resize or pad.
+ if is_training:
+ if augmentation_type.lower() == 'inception':
+ min_aspect_ratio = augmentation_params['min_aspect_ratio'] # pyrefly: ignore[unsupported-operation]
+ max_aspect_ratio = augmentation_params['max_aspect_ratio'] # pyrefly: ignore[unsupported-operation]
+ min_area_ratio = augmentation_params['min_area_ratio'] # pyrefly: ignore[unsupported-operation]
+ max_area_ratio = augmentation_params['max_area_ratio'] # pyrefly: ignore[unsupported-operation]
+ # Inception-style image crop: random crop -> resize.
+ def random_crop_resize_fn(x, state=None):
+ return processors.random_crop_resize(
+ x, output_height=crop_size, output_width=crop_size,
+ num_frames=num_frames, num_channels=num_raw_channels,
+ aspect_ratio=(min_aspect_ratio, max_aspect_ratio),
+ area_range=(min_area_ratio, max_area_ratio),
+ state=state)
+ preprocessor_builder.add_fn(
+ fn=random_crop_resize_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_random_crop_resize'.format(output_feature_name),
+ stateful=sync_random_state)
+ elif augmentation_type.lower() == 'vgg':
+ # VGG-style image crop: resize -> random crop.
+ def resize_and_crop_fn(x, state):
+ return processors.resize_and_crop(
+ x,
+ min_resize=min_resize,
+ crop_size=crop_size, is_flow=False, is_random=True,
+ state=state)
+ preprocessor_builder.add_fn(
+ fn=resize_and_crop_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_resize_random_crop'.format(output_feature_name),
+ stateful=sync_random_state)
+ elif augmentation_type.lower() == 'ava':
+ # AVA-style image aug: random_crop -> resize -> random pad.
+ def random_square_crop_by_scale_fn(x, state=None):
+ return processors.random_square_crop_by_scale(
+ image=x,
+ scale_min=augmentation_params['scale_min'], # pyrefly: ignore[unsupported-operation]
+ scale_max=augmentation_params['scale_max'], # pyrefly: ignore[unsupported-operation]
+ state=state)
+ preprocessor_builder.add_fn(
+ fn=random_square_crop_by_scale_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_random_square_crop_by_scale'.format(output_feature_name),
+ stateful=sync_random_state)
+ def resize_and_pad_fn(x, state=None):
+ return processors.resize_and_pad(
+ frames=x,
+ max_resize=crop_size,
+ pad_size=crop_size,
+ random=True,
+ state=state)
+ preprocessor_builder.add_fn(
+ fn=resize_and_pad_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_resize_random_pad'.format(output_feature_name),
+ # Use state to keep coherence between modalities if requested.
+ stateful=sync_random_state)
+ else:
+ raise ValueError('Unrecognized augmentation_type: %s' %
+ augmentation_type)
+
+ if random_flip_image:
+ def random_flip_left_right_fn(x, state=None):
+ return processors.random_flip_left_right(
+ x, seed=seed, is_flow=False, state=state)
+ preprocessor_builder.add_fn(
+ fn=random_flip_left_right_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_random_flip'.format(output_feature_name),
+ stateful=sync_random_state)
+ else:
+ # Crop images, either a 3-view crop or a central crop.
+ if multi_crop:
+ resize_smallest_fn = functools.partial(
+ processors.resize_smallest,
+ min_resize=min_resize,
+ is_flow=False)
+ # Resize images (resize happens only if necessary to save compute).
+ preprocessor_builder.add_fn(
+ fn=resize_smallest_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_resize_smallest'.format(output_feature_name))
+ # Multi crop of the frames.
+ preprocessor_builder.add_fn(
+ fn=lambda x: processors.multi_crop_image(x, crop_size, crop_size),
+ feature_name=output_feature_name,
+ fn_name='{}_multi_crop'.format(output_feature_name))
+ else:
+ if augmentation_type.lower() == 'ava':
+ def resize_and_pad_fn(x, state=None):
+ return processors.resize_and_pad(
+ frames=x,
+ max_resize=crop_size,
+ pad_size=crop_size,
+ random=False,
+ state=state)
+ preprocessor_builder.add_fn(
+ fn=resize_and_pad_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_resize_central_pad'.format(output_feature_name),
+ stateful=sync_random_state)
+ else:
+ def resize_and_crop_fn(x, state=None):
+ return processors.resize_and_crop(
+ x,
+ min_resize=min_resize,
+ crop_size=crop_size,
+ is_flow=False,
+ is_random=False,
+ state=state)
+ preprocessor_builder.add_fn(
+ fn=resize_and_crop_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_resize_central_crop'.format(output_feature_name),
+ stateful=sync_random_state)
+
+ # Apply extra augmentation policy.
+ if is_training:
+ if randaug_params is not None and autoaug_params is not None:
+ raise ValueError('Choose to apply one of data augmentation policies: '
+ 'randaug and autoaug.')
+
+ if autoaug_params is not None:
+ augmenter = augment.AutoAugment(
+ augmentation_name=autoaug_params['augmentation_name'],
+ cutout_const=autoaug_params['cutout_const'],
+ translate_const=autoaug_params['translate_const'])
+ preprocessor_builder.add_fn(
+ fn=augmenter.distort,
+ feature_name=output_feature_name,
+ fn_name='{}_autoaug'.format(output_feature_name))
+
+ if randaug_params is not None:
+ augmenter = augment.RandAugment(
+ num_layers=randaug_params['num_layers'],
+ magnitude=randaug_params['magnitude'],
+ cutout_const=randaug_params['cutout_const'],
+ translate_const=randaug_params['translate_const'],
+ prob_to_apply=randaug_params['prob_to_apply'],
+ exclude_ops=randaug_params['exclude_ops'])
+ preprocessor_builder.add_fn(
+ fn=augmenter.distort,
+ feature_name=output_feature_name,
+ fn_name='{}_randaug'.format(output_feature_name))
+
+ # Cast the frames in float32, normalizing according to zero_centering_image.
+ preprocessor_builder.add_fn(
+ fn=lambda x: processors.normalize_image(x, zero_centering_image),
+ feature_name=output_feature_name,
+ fn_name='{}_normalize'.format(output_feature_name))
+
+ if (num_test_clips > 1 or multi_crop) and not is_training:
+ # In this case, multiple clips are merged together in batch dimenstion which
+ # will be B * num_test_clips.
+ def reshape_fn(x):
+ target_shape = (-1, num_frames, x.shape[-3], x.shape[-2], x.shape[-1])
+ return tf.reshape(x, target_shape)
+ postprocessor_builder.add_fn(
+ fn=reshape_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_reshape'.format(output_feature_name))
+
+
+def add_context_label(
+ parser_builder: builders.SequenceExampleParserBuilder,
+ sampler_builder: builders.SamplerBuilder,
+ preprocessor_builder: builders.PreprocessorBuilder,
+ input_label_index_feature_name: str = 'clip/key_frame/bbox/label/index',
+ output_label_index_feature_name: str = builders.LABEL_INDEX_FEATURE_NAME,
+ input_label_name_feature_name: str = 'clip/key_frame/bbox/label/string',
+ output_label_name_feature_name: str = builders.LABEL_NAME_FEATURE_NAME,
+ # Label related parameters.
+ num_frames: int = 1,
+ num_instances_per_frame: int = 5,
+ zero_based_index: bool = False,
+ one_hot_label: bool = True,
+ num_classes: Optional[int] = None,
+ add_label_name: bool = False):
+ """Adds functions to process label feature to builders.
+
+ Args:
+ parser_builder: An instance of a builders.SequenceExampleParserBuilder.
+ sampler_builder: An instance of a builders.SamplerBuilder.
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ input_label_index_feature_name: Name of the label index feature in the input
+ SequenceExample. Exposing this as an argument allows using this function
+ for different label features.
+ output_label_index_feature_name: Name of the label index feature in the
+ output features dictionary. Exposing this as an argument allows using this
+ function for different label features.
+ input_label_name_feature_name: Name of the label name feature in the input
+ SequenceExample. Exposing this as an argument allows using this function
+ for different label features.
+ output_label_name_feature_name: Name of the label name feature in the
+ output features dictionary. Exposing this as an argument allows using this
+ function for different label features.
+ num_frames: The number of frames. If the num_frames > 1, the labels will be
+ duplicated.
+ num_instances_per_frame: The number of label instances per frames.
+ zero_based_index: Whether the raw index are zero based. If not, converted to
+ the zero based index as the output.
+ one_hot_label: Return labels as one hot tensors. If is_multi_label is True,
+ one hot tensor might have multiple ones.
+ num_classes: Total number of classes in the dataset. It has to be procided
+ if one_hot_label is True.
+ add_label_name: Also return the name of the label. Not yet supported for
+ multi label.
+ """
+ # Validate parameters.
+ if one_hot_label and not num_classes:
+ raise ValueError(
+ 'num_classes should be given when requesting one hot label.')
+
+ # Parse label.
+ parser_builder.parse_feature(
+ feature_name=input_label_index_feature_name,
+ feature_type=tf.io.VarLenFeature(dtype=tf.int64),
+ output_name=output_label_index_feature_name,
+ is_context=True)
+
+ # Densify labels tensor in order to support multi label case.
+ sampler_builder.add_fn(
+ fn=lambda x: tf.sparse.to_dense(x, default_value=-1),
+ feature_name=output_label_index_feature_name,
+ fn_name='{}_sparse_to_dense'.format(output_label_index_feature_name))
+
+ # Crop or pad labels to max_num_instance.
+ crop_or_pad_features_fn = functools.partial(
+ processors.crop_or_pad_features,
+ max_num_features=num_instances_per_frame,
+ feature_dimension=1,
+ constant_values=-1)
+ preprocessor_builder.add_fn(
+ fn=crop_or_pad_features_fn,
+ feature_name=output_label_index_feature_name,
+ fn_name='{}_crop_or_pad'.format(output_label_index_feature_name))
+
+ if num_frames > 1:
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.tile(x, [num_frames, 1]),
+ feature_name=output_label_index_feature_name,
+ fn_name='{}_duplicate'.format(output_label_index_feature_name))
+
+ # Reshape the feature vector in [T, N].
+ target_shape = [num_frames, num_instances_per_frame]
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.reshape(x, target_shape),
+ feature_name=output_label_index_feature_name,
+ fn_name='{}_reshape'.format(output_label_index_feature_name))
+
+ # Convert the label id to be zero-indexed.
+ if not zero_based_index:
+ preprocessor_builder.add_fn(
+ fn=lambda x: x - 1,
+ feature_name=output_label_index_feature_name,
+ fn_name='{}_zero_based'.format(output_label_index_feature_name))
+
+ # Replace label index by one hot representation.
+ if one_hot_label:
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.one_hot(x, num_classes),
+ feature_name=output_label_index_feature_name,
+ fn_name='{}_one_hot'.format(output_label_index_feature_name))
+
+ if add_label_name:
+ parser_builder.parse_feature(
+ feature_name=input_label_name_feature_name,
+ feature_type=tf.io.VarLenFeature(dtype=tf.string),
+ output_name=output_label_name_feature_name,
+ is_context=True)
+ sampler_builder.add_fn(
+ fn=tf.sparse.to_dense,
+ feature_name=output_label_name_feature_name,
+ fn_name='{}_sparse_to_dense'.format(output_label_name_feature_name))
+
+ # Crop or pad labels to max_num_instance.
+ crop_or_pad_features_fn = functools.partial(
+ processors.crop_or_pad_features,
+ max_num_features=num_instances_per_frame,
+ feature_dimension=1,
+ constant_values='')
+ preprocessor_builder.add_fn(
+ fn=crop_or_pad_features_fn,
+ feature_name=output_label_name_feature_name,
+ fn_name='{}_crop_or_pad'.format(output_label_name_feature_name))
+
+ if num_frames > 1:
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.tile(x, [num_frames, 1]),
+ feature_name=output_label_name_feature_name,
+ fn_name='{}_duplicate'.format(output_label_name_feature_name))
+
+ # Reshape the feature vector in [T, N].
+ target_shape = [num_frames, num_instances_per_frame]
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.reshape(x, target_shape),
+ feature_name=output_label_name_feature_name,
+ fn_name='{}_reshape'.format(output_label_name_feature_name))
+
+
+def add_frame_label(
+ parser_builder: builders.SequenceExampleParserBuilder,
+ sampler_builder: builders.SamplerBuilder,
+ preprocessor_builder: builders.PreprocessorBuilder,
+ input_label_name: str = 'region/bbox/xmin',
+ output_label_name: str = 'instance_xmin',
+ dtype: tf.dtypes.DType = tf.float32,
+ # Instance related parameters.
+ is_training: bool = True,
+ num_instances_per_frame: int = 5,
+ num_frames: int = 32,
+ temporal_stride: int = 1,
+ sync_random_state: bool = True,
+ seed: Optional[int] = None):
+ """Adds functions to process label feature to builders.
+
+ Args:
+ parser_builder: An instance of a builders.SequenceExampleParserBuilder.
+ sampler_builder: An instance of a builders.SamplerBuilder.
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ input_label_name: The per frame label key stored in the tfse for parsing.
+ output_label_name: The output feature name. Exposing this as an argument
+ allows using this function for
+ dtype: The data type to be parsed.
+ is_training: Whether in training mode.
+ num_instances_per_frame: The number of instances label per frame.
+ num_frames: The number of frames in a video clip.
+ temporal_stride: The temporal sample stride.
+ sync_random_state: Whether to sync random state between features.
+ seed: the random seed.
+ """
+ # Parse per-frame label.
+ parser_builder.parse_feature(
+ feature_name=input_label_name,
+ # Entire signal stored in one Feature.
+ feature_type=tf.io.VarLenFeature(dtype=dtype),
+ output_name=output_label_name)
+
+ sampler_builder.add_fn(
+ fn=tf.sparse.to_dense,
+ feature_name=output_label_name,
+ fn_name='{}_sparse_to_dense'.format(output_label_name))
+
+ # Temporal sampler.
+ if is_training:
+ def sample_sequence_fn(x, state=None):
+ return processors.sample_sequence(
+ x,
+ num_steps=num_frames,
+ random=True,
+ stride=temporal_stride,
+ seed=seed,
+ state=state)
+ # Sample random clip.
+ sampler_builder.add_fn(
+ fn=sample_sequence_fn,
+ feature_name=output_label_name,
+ fn_name='{}_random_sample'.format(output_label_name),
+ # Use state to keep coherence between modalities if requested.
+ stateful=sync_random_state)
+ else:
+ sample_sequence_fn = functools.partial(
+ processors.sample_sequence,
+ num_steps=num_frames,
+ random=False,
+ stride=temporal_stride,
+ seed=None)
+ # Sample middle clip.
+ sampler_builder.add_fn(
+ fn=sample_sequence_fn,
+ feature_name=output_label_name,
+ fn_name='{}_middle_sample'.format(output_label_name))
+
+ # Crop or pad labels to num_instances_per_frame.
+ crop_or_pad_features_fn = functools.partial(
+ processors.crop_or_pad_features,
+ max_num_features=num_instances_per_frame,
+ feature_dimension=1,
+ constant_values=-1)
+ preprocessor_builder.add_fn(
+ fn=crop_or_pad_features_fn,
+ feature_name=output_label_name,
+ fn_name='{}_crop_or_pad'.format(output_label_name))
+
+ # Reshape the feature vector in [T, N].
+ target_shape = [num_frames, num_instances_per_frame]
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.reshape(x, target_shape),
+ feature_name=output_label_name,
+ fn_name='{}_reshape'.format(output_label_name))
+
+
+def merge_multi_labels(
+ inputs: MutableMapping[str, tf.Tensor],
+ num_classes: int,
+ input_label_name: str = 'label',
+ input_boxes_name: str = 'instances_position',
+ input_masks_name: str = 'instances_mask',
+ output_nonmerge_boxes_name: str = 'nonmerge_instances_position',
+ output_nonmerge_label_name: str = 'nonmerge_label',
+ quantization_bins: int = 1000):
+ """Merges boxes with the same coordinates and returns k-hot labels.
+
+ Args:
+ inputs: An inputs dictionary. Containing at least the following fields:
+ * instances_position: A tf.float32 tensor with shape
+ [num_frames, num_boxes, 4] holding boxes. Only normalized coordinates in
+ form [ymin, xmin, ymax, xmax] are allowed.
+ * label: A tf.int32 tensor with shape [num_frames, num_boxes] holding
+ zero-indexed classes. -1 means that label is padded and thus invalid.
+ * instances_mask: A tf.bool tensor with shape [num_frames, num_boxes]
+ holding class confidences.
+ num_classes: The maximum number of classes for the dataset.
+ input_label_name: The input label tensor name. The label should be the
+ 0-indexed, unique label per box.
+ input_boxes_name: The input box tensor name.
+ input_masks_name: The input binary mask tensor name.
+ output_nonmerge_boxes_name: The output box name holding the none-merged
+ might be duplicated boxes.
+ output_nonmerge_label_name: The output label index per box.
+ quantization_bins: The number of bins used to quantize the box coordinate.
+
+ Returns:
+ The output dictionary contains the merged multilabel labels/boxes and
+ none-merged labels/boxes pairs.
+ """
+ unique_labels = inputs[input_label_name]
+ unique_boxes = inputs[input_boxes_name]
+ unique_weights = inputs[input_masks_name]
+
+ if unique_labels.shape.rank != 2:
+ raise ValueError('one_hot should be turned off if merge_multi_labels.')
+
+ num_instances_per_frame = unique_boxes.shape.as_list()[-2]
+
+ labels = tf.unstack(unique_labels, axis=0)
+ boxes = tf.unstack(unique_boxes, axis=0)
+ weights = tf.unstack(unique_weights, axis=0)
+
+ merged_boxes = []
+ merged_labels = []
+ merged_weights = []
+
+ def true_fn(box, label, weight,
+ num_classes=num_classes,
+ num_instances_per_frame=num_instances_per_frame,
+ quantization_bins=quantization_bins):
+ # The label is 0-index and the invalid labels are padded with -1. We create
+ # a binary mask here to filter out the invalid/padded label.
+ valid_mask = tf.greater(label, -1)
+
+ confidence = tf.cast(weight, dtype=tf.float32)
+ merged_box, merged_label, merged_confidence, _ = (
+ merge_boxes_with_multiple_labels(
+ tf.boolean_mask(box, valid_mask),
+ tf.boolean_mask(label, valid_mask),
+ tf.boolean_mask(confidence, valid_mask),
+ num_classes=num_classes,
+ quantization_bins=quantization_bins))
+
+ merged_box = pad_or_clip_nd(
+ merged_box, [num_instances_per_frame, 4])
+ merged_label = pad_or_clip_nd(
+ merged_label, [num_instances_per_frame, num_classes])
+ merged_confidence = pad_or_clip_nd(
+ merged_confidence, [num_instances_per_frame, num_classes])
+ merged_label = tf.cast(merged_label, dtype=tf.float32)
+ merged_weight = tf.cast(tf.reduce_max(merged_confidence, axis=-1), tf.bool)
+ return merged_box, merged_label, merged_weight
+
+ def false_fn(box, label, weight, num_classes=num_classes):
+ label = tf.one_hot(label, num_classes)
+ return box, label, weight
+
+ for i in range(len(labels)):
+ # determine whether this frame is the keyframe by examining the label id.
+ # any instance with label id > -1 means such frame is labeled.
+ contains_label = tf.greater(tf.reduce_max(labels[i]), -1)
+
+ # pylint: disable=cell-var-from-loop
+ merged_box, merged_label, merged_weight = tf.cond(
+ contains_label,
+ true_fn=lambda: true_fn(boxes[i], labels[i], weights[i]),
+ false_fn=lambda: false_fn(boxes[i], labels[i], weights[i]))
+ # pylint: enable=cell-var-from-loop
+
+ merged_boxes.append(merged_box)
+ merged_labels.append(merged_label)
+ merged_weights.append(merged_weight)
+
+ # All values to the input labels/boxes/masks will now store the merged value.
+ inputs[input_label_name] = tf.stack(merged_labels, axis=0)
+ inputs[input_boxes_name] = tf.stack(merged_boxes, axis=0)
+ inputs[input_masks_name] = tf.stack(merged_weights, axis=0)
+
+ # The difference between groundtruth classes/boxes and above is the
+ # groundtruth boxes maybe duplicated since one box may have multiple labels.
+ inputs[output_nonmerge_label_name] = unique_labels
+ inputs[output_nonmerge_boxes_name] = unique_boxes
+ return inputs
+
+
+def add_context_box_dim(
+ parser_builder: builders.SequenceExampleParserBuilder,
+ sampler_builder: builders.SamplerBuilder,
+ preprocessor_builder: builders.PreprocessorBuilder,
+ input_box_dim_name: str = 'clip/key_frame/bbox/xmin',
+ output_box_dim_name: str = 'keyframe_xmin',
+ num_frames: int = 1,
+ num_instances_per_frame: int = 5,
+ default_value: float = -1):
+ """Adds functions to process label feature to builders.
+
+ Args:
+ parser_builder: An instance of a builders.BaseParserBuilder.
+ sampler_builder: An instance of a builders.SamplerBuilder.
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ input_box_dim_name: The box dim key stored in the tfse for parsing.
+ output_box_dim_name: The output feature name. Exposing this as an argument
+ allows using this function for
+ num_frames: The number of frames the boxes need to be duplicated. This
+ parameter is added to adapt the format of boxes in this codebase.
+ num_instances_per_frame: The number of instance per keyframe.
+ default_value: The default value to pad to num_instances_per_frame.
+ """
+ # Parse box dim.
+ parser_builder.parse_feature(
+ feature_name=input_box_dim_name,
+ # Entire signal stored in one Feature.
+ feature_type=tf.io.VarLenFeature(dtype=tf.float32),
+ output_name=output_box_dim_name,
+ is_context=True)
+
+ sampler_builder.add_fn(
+ fn=tf.sparse.to_dense,
+ feature_name=output_box_dim_name,
+ fn_name='{}_sparse_to_dense'.format(output_box_dim_name))
+
+ # Crop or pad boxes to num_instances_per_frame.
+ crop_or_pad_features_fn = functools.partial(
+ processors.crop_or_pad_features,
+ max_num_features=num_instances_per_frame,
+ feature_dimension=1,
+ constant_values=default_value)
+ preprocessor_builder.add_fn(
+ fn=crop_or_pad_features_fn,
+ feature_name=output_box_dim_name,
+ fn_name='{}_crop_or_pad'.format(output_box_dim_name))
+
+ # Duplicate keyframe boxes to all frames.
+ if num_frames > 1:
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.tile(x, [num_frames, 1]),
+ feature_name=output_box_dim_name,
+ fn_name='{}_duplicate'.format(output_box_dim_name))
+
+ # Reshape the feature vector in [T, N].
+ target_shape = [num_frames, num_instances_per_frame]
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.reshape(x, target_shape),
+ feature_name=output_box_dim_name,
+ fn_name='{}_reshape'.format(output_box_dim_name))
+
+
+def add_context_features(
+ parser_builder: builders.SequenceExampleParserBuilder,
+ sampler_builder: builders.SamplerBuilder,
+ preprocessor_builder: builders.PreprocessorBuilder,
+ input_feature_name: str = 'clip/p_scores',
+ output_feature_name: str = 'p_scores',
+ max_num_features: int = 64,
+ feature_dimension: int = 35):
+ """Adds functions to process context feature to builders.
+
+ Args:
+ parser_builder: An instance of a builders.BaseParserBuilder.
+ sampler_builder: An instance of a builders.SamplerBuilder.
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ input_feature_name: The box dim key stored in the tfse for parsing.
+ output_feature_name: The output feature name. Exposing this as an argument
+ allows using this function for
+ max_num_features: The number of features to be processed.
+ feature_dimension: The feature dimension.
+ """
+ # Parse box dim.
+ parser_builder.parse_feature(
+ feature_name=input_feature_name,
+ # Entire signal stored in one Feature.
+ feature_type=tf.io.VarLenFeature(dtype=tf.float32),
+ output_name=output_feature_name,
+ is_context=True)
+
+ sampler_builder.add_fn(
+ fn=tf.sparse.to_dense,
+ feature_name=output_feature_name,
+ fn_name='{}_sparse_to_dense'.format(output_feature_name))
+
+ # Crop or pad boxes to num_instances_per_frame.
+ crop_or_pad_features_fn = functools.partial(
+ processors.crop_or_pad_features,
+ max_num_features=max_num_features,
+ feature_dimension=feature_dimension,
+ constant_values=-1)
+ preprocessor_builder.add_fn(
+ fn=crop_or_pad_features_fn,
+ feature_name=output_feature_name,
+ fn_name='{}_crop_or_pad'.format(output_feature_name))
+
+ # Reshape the feature vector in [T, N].
+ target_shape = [max_num_features, feature_dimension]
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.reshape(x, target_shape),
+ feature_name=output_feature_name,
+ fn_name='{}_reshape'.format(output_feature_name))
+
+
+def add_instance_box_dim(
+ parser_builder: builders.SequenceExampleParserBuilder,
+ sampler_builder: builders.SamplerBuilder,
+ preprocessor_builder: builders.PreprocessorBuilder,
+ input_box_dim_name: str = 'region/bbox/xmin',
+ output_box_dim_name: str = 'instance_xmin',
+ # Instance related parameters.
+ sample_around_keyframe: bool = False,
+ sample_random: bool = True,
+ num_instances_per_frame: int = 5,
+ num_frames: int = 32,
+ temporal_stride: int = 1,
+ sync_random_state: bool = True):
+ """Adds functions to process label feature to builders.
+
+ Args:
+ parser_builder: An instance of a builders.BaseParserBuilder.
+ sampler_builder: An instance of a builders.SamplerBuilder.
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ input_box_dim_name: The input feature name.
+ output_box_dim_name: The output feature name.
+ sample_around_keyframe: Whether to sample the sequence around the keyframe.
+ If True, it requires keyframe id is known.
+ sample_random: Whether to perform random sampling.
+ num_instances_per_frame: The number of instances per frame.
+ num_frames: The number of frames to sample.
+ temporal_stride: The temopral sampling stride.
+ sync_random_state: Whether to sync the random state.
+ """
+ if sample_random and sample_around_keyframe:
+ raise ValueError(
+ 'sample_random and sample_around_keyframe cannot be both True.')
+
+ # Parse box dim.
+ parser_builder.parse_feature(
+ feature_name=input_box_dim_name,
+ # Entire signal stored in one Feature.
+ feature_type=tf.io.VarLenFeature(dtype=tf.float32),
+ output_name=output_box_dim_name)
+
+ sampler_builder.add_fn(
+ fn=tf.sparse.to_dense,
+ feature_name=output_box_dim_name,
+ fn_name='{}_sparse_to_dense'.format(output_box_dim_name))
+
+ # Temporal sampler.
+ if sample_random:
+ # Sample random clip.
+ def sample_sequence_fn(x, state=None):
+ return processors.sample_sequence(
+ sequence=x,
+ num_steps=num_frames,
+ random=True,
+ stride=temporal_stride,
+ state=state)
+ sampler_builder.add_fn(
+ fn=sample_sequence_fn,
+ feature_name=output_box_dim_name,
+ fn_name='{}_random_sample'.format(output_box_dim_name),
+ # Use state to keep coherence between modalities if requested.
+ stateful=sync_random_state)
+ elif sample_around_keyframe:
+ sample_around_keyframe_fn = functools.partial(
+ processors.sample_sequence_around_keyframe,
+ num_steps=num_frames,
+ stride=temporal_stride,
+ sample_target_key=output_box_dim_name)
+ sampler_builder.add_fn(
+ fn=sample_around_keyframe_fn,
+ fn_name='{}_sample_around_keyframe'.format(output_box_dim_name))
+ else:
+ # Sample middle clip.
+ sample_sequence_fn = functools.partial(
+ processors.sample_sequence,
+ num_steps=num_frames,
+ random=False,
+ stride=temporal_stride)
+ sampler_builder.add_fn(
+ fn=sample_sequence_fn,
+ feature_name=output_box_dim_name,
+ fn_name='{}_middle_sample'.format(output_box_dim_name))
+
+ # Crop or pad boxes to num_instances_per_frame.
+ crop_or_pad_features_fn = functools.partial(
+ processors.crop_or_pad_features,
+ max_num_features=num_instances_per_frame,
+ feature_dimension=1,
+ constant_values=-1)
+ preprocessor_builder.add_fn(
+ fn=crop_or_pad_features_fn,
+ feature_name=output_box_dim_name,
+ fn_name='{}_crop_or_pad'.format(output_box_dim_name))
+
+ # Reshape the feature vector in [T, N, C].
+ target_shape = [num_frames, num_instances_per_frame]
+ preprocessor_builder.add_fn(
+ fn=lambda x: tf.reshape(x, target_shape),
+ feature_name=output_box_dim_name,
+ fn_name='{}_reshape'.format(output_box_dim_name))
+
+
+def group_instance_box_dims(
+ preprocessor_builder: builders.PreprocessorBuilder,
+ output_position_name: str = 'instances_position',
+ box_key_prefix: str = 'instance'):
+ """Groups instance box dims into a compact feature vector.
+
+ Args:
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ output_position_name: The output position tensor name.
+ box_key_prefix: The box key prefix.
+ """
+
+ def _group_box_dims_fn(inputs: Dict[str, tf.Tensor]) -> Dict[str, tf.Tensor]:
+ """Utility function to group box dimensions."""
+ box_dims = [
+ inputs.pop('{}_ymin'.format(box_key_prefix)),
+ inputs.pop('{}_xmin'.format(box_key_prefix)),
+ inputs.pop('{}_ymax'.format(box_key_prefix)),
+ inputs.pop('{}_xmax'.format(box_key_prefix)),
+ ]
+
+ inputs[output_position_name] = tf.stack(
+ box_dims, axis=-1, name='stack_{}'.format(output_position_name))
+ return inputs
+
+ preprocessor_builder.add_fn(
+ fn=_group_box_dims_fn,
+ fn_name='{}_stack'.format(output_position_name))
+
+
+def infer_instances_mask_from_position(
+ inputs: MutableMapping[str, tf.Tensor],
+ instances_position_name: str = 'instances_position'):
+ """Infers binary instances mask from instance positions."""
+ instances_position = inputs[instances_position_name]
+ instances_mask = tf.reduce_sum(instances_position, axis=-1) > 0
+ inputs['instances_mask'] = instances_mask
+ return inputs
+
+
+def filter_instances_box_by_score(
+ inputs: MutableMapping[str, tf.Tensor],
+ box_key_prefix: str = 'detected_instances',
+ score_threshold: float = 0.2):
+ """Filters the low-score detected instance.
+
+ Args:
+ inputs: The input dictionary contains boxes and masks.
+ box_key_prefix: The instances position key prefix.
+ score_threshold: The detection score to threshold.
+
+ Returns:
+ The output dictionary contains filtered boxes and mask.
+ """
+ instances_position = inputs.pop(f'{box_key_prefix}_position')
+ instances_score = inputs[f'{box_key_prefix}_score']
+ num_frames, num_instances = instances_score.shape.as_list()
+
+ masks_list = tf.unstack(instances_score > score_threshold, num_frames, axis=0)
+ position_list = tf.unstack(instances_position, num_frames, axis=0)
+ constant_padding = tf.ones_like(position_list[0], dtype=tf.float32) * -1.
+ instances_position_filtered = []
+ for mask, position in zip(masks_list, position_list):
+ position_filtered = tf.boolean_mask(position, mask)
+ position_filtered = tf.concat([position_filtered, constant_padding], axis=0)
+ position_filtered = tf.slice(position_filtered, [0, 0], [num_instances, -1])
+ instances_position_filtered.append(position_filtered)
+ instances_position_filtered = tf.stack(instances_position_filtered, axis=0)
+
+ inputs[f'{box_key_prefix}_position'] = instances_position_filtered
+ inputs[f'{box_key_prefix}_mask'] = tf.reduce_sum(
+ instances_position_filtered, axis=-1) > 0
+ return inputs
+
+
+def adjust_positions(preprocessor_builder: builders.PreprocessorBuilder,
+ input_tensor_name: str = 'instances_position',
+ output_tensor_name: str = 'instances_position'):
+ """Adjusts box positions based on image/video data augmentation.
+
+ Args:
+ preprocessor_builder: An instance of a builders.PreprocessorBuilder.
+ input_tensor_name: The name of input tensor.
+ output_tensor_name: The name of output tensor.
+ """
+
+ def _resize_and_crop(boxes, image_info, is_flip=None):
+ """Resizes and crop boxes."""
+ # input boxes in shape [T, N, 4]
+ # image_info in shape [4, 2]
+ original_image_size, target_image_size, scale, offset = tf.unstack(
+ image_info, axis=0)
+ # 1) Un-normalize boxes to absolute coordinate.
+ boxes = box_ops.denormalize_boxes(boxes, original_image_size)
+ # 2) Adjusts box coordinates based on image_scale and offset.
+ boxes *= tf.tile(scale[None, None, :], [1, 1, 2])
+ boxes -= tf.tile(offset[None, None, :], [1, 1, 2])
+ # 3) Clips the boxes.
+ boxes = box_ops.clip_boxes(boxes, target_image_size)
+ # 4) Normalize boxes.
+ boxes = box_ops.normalize_boxes(boxes, target_image_size)
+ # 5) Random flip boxes.
+ if is_flip is not None:
+ boxes = tf.cond(
+ tf.equal(is_flip, 1),
+ lambda: box_ops.horizontal_flip_boxes(boxes),
+ lambda: boxes)
+ return boxes
+
+ def _adjust_boxes_fn(boxes: tf.Tensor,
+ state: MutableMapping[str, Any]) -> tf.Tensor:
+ return _resize_and_crop(
+ boxes,
+ image_info=state['image_info'],
+ is_flip=state.get('flip_left_right_is_flipped', None))
+
+ preprocessor_builder.add_fn(
+ fn=_adjust_boxes_fn,
+ feature_name=input_tensor_name,
+ fn_name='{}_adjust'.format(output_tensor_name),
+ stateful=True)
+
+
+def update_valid_instances_mask(
+ inputs: Dict[str, tf.Tensor],
+ instances_position_name: str = 'instances_position',
+ instances_mask_name: str = 'instances_mask') -> Dict[str, tf.Tensor]:
+ """Filters empty boxes and mark them in instances_mask.
+
+ Args:
+ inputs: The dictionary contains the input tensors.
+ instances_position_name: The box name stored in the dictionary.
+ instances_mask_name: The binary instance mask key in the inputs.
+
+ Returns:
+ The input dictionary with the updated instances mask.
+ """
+ boxes = inputs[instances_position_name]
+ height = boxes[..., 2] - boxes[..., 0]
+ width = boxes[..., 3] - boxes[..., 1]
+ instances_mask = tf.logical_and(tf.greater(height, 0), tf.greater(width, 0))
+ if instances_mask_name in inputs:
+ instances_mask = tf.logical_and(inputs[instances_mask_name], instances_mask)
+ inputs[instances_mask_name] = instances_mask
+ return inputs
+
+
+def apply_default_color_augmentations(
+ preprocessor_builder: builders.PreprocessorBuilder,
+ image_feature_name: str = builders.IMAGE_FEATURE_NAME,
+ zero_centering_image: bool = False):
+ """Applies default color augmentations on images.
+
+ Args:
+ preprocessor_builder: The prerpocessor builder.
+ image_feature_name: The image feature name.
+ zero_centering_image: Whether the image has been zero centered.
+ """
+ preprocessor_builder.add_fn(
+ functools.partial(
+ processors.random_color_augmentation,
+ zero_centering_image=zero_centering_image),
+ feature_name=image_feature_name,
+ fn_name='random_jitter_color')
+ preprocessor_builder.add_fn(
+ functools.partial(
+ processors.random_blur_and_solarize,
+ zero_centering_image=zero_centering_image),
+ feature_name=image_feature_name,
+ fn_name='random_blur_and_solarize')
+
+
+def get_shards(table_name: str) -> list[str]:
+ """Expands table into a list of sharded filenames."""
+ filenames = []
+ if '@' in table_name:
+ base_filename, num_shards = table_name.split('@')
+ num_shards = int(num_shards)
+ for ind in range(num_shards):
+ filename = '{}.{:05d}-of-{:05d}'.format(base_filename, ind, num_shards)
+ filenames.append(filename)
+ else:
+ filenames.append(table_name)
+ return filenames
diff --git a/official/projects/videoglue/datasets/dataset_factory.py b/official/projects/videoglue/datasets/dataset_factory.py
new file mode 100644
index 00000000000..b32503f6d6d
--- /dev/null
+++ b/official/projects/videoglue/datasets/dataset_factory.py
@@ -0,0 +1,135 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Defines internal dataset factory."""
+from typing import Any, Mapping, Optional
+
+from absl import flags
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.projects.videoglue.datasets import action_localization
+from official.projects.videoglue.datasets import video_classification
+
+
+# Define tf data service flags.
+_TF_DATA_SERVICE_ADDRESS = flags.DEFINE_string(
+ 'tf_data_service_address', '',
+ 'tf.data.service main address, starts with grpc+loas://.')
+
+
+FACTORY = {
+ 'kinetics400': video_classification.Kinetics400Factory,
+ 'diving48': video_classification.Diving48Factory,
+ 'sthv2': video_classification.Sthv2Factory,
+ 'moments-in-time': video_classification.MomentsInTimeFactory,
+ 'ava': action_localization.AVAFactory,
+ 'avakinetics': action_localization.AVAKineticsFactory,
+}
+
+
+class DataLoader(object):
+ """Data loader that returns tf.data.Dataset."""
+
+ def __init__(self, params: cfg.DataConfig, dataset_config: Mapping[str, Any]):
+ """Constructor.
+
+ Args:
+ params: Instance of `Dataset` configuration,
+ dataset_config: The config dictionary used for data reader pipeline.
+ """
+ self._params = params
+ self._dataset_config = dataset_config
+ self._name = params.name
+ self._is_training = params.is_training
+ self._shuffle = params.is_training
+
+ dataset_cls = FACTORY[self._name]
+ self._dataset_cls = dataset_cls(
+ subset='train' if self._is_training else 'test')
+
+ if params.feature_shape[1] != params.feature_shape[2]:
+ raise ValueError('Only support square crop. Got feature shape: %s' %
+ params.feature_shape)
+
+ self._dataset_cls.configure(**dataset_config)
+ self._dataset_cls.tune(
+ prefetch_buffer_size=params.prefetch_buffer_size,
+ shuffle_buffer=params.shuffle_buffer_size,
+ num_process_threads=params.num_process_threads,
+ cycle_length=params.cycle_length,
+ num_parallel_calls_interleave=params.num_parallel_calls_interleave,
+ block_length=params.block_length)
+
+ def _dataset_fn(
+ self,
+ input_context: Optional[tf.distribute.InputContext] = None
+ ) -> tf.data.Dataset:
+ """Generates features and labels for training or evaluation.
+
+ This uses the input pipeline based approach using file name queue to read
+ data so that entire data are not loaded in memory.
+
+ Args:
+ input_context: Distributed context.
+
+ Returns:
+ tf.data.Dataset
+ """
+ logging.info('dataset params: %s', self._params)
+ if input_context:
+ global_batch_size = self._params.global_batch_size
+ batch_size = input_context.get_per_replica_batch_size(global_batch_size)
+ else:
+ batch_size = self._params.global_batch_size
+
+ dataset = self._dataset_cls.make_dataset(
+ shuffle=self._shuffle,
+ num_epochs=-1 if self._is_training else 1,
+ batch_size=batch_size,
+ padded_batch=False,
+ drop_remainder=self._params.drop_remainder,
+ keep_key=False,
+ override_preprocess_fn=None,
+ input_context=input_context,
+ multi_host_sharding=False,
+ name=self._name)
+
+ if self._params.cache:
+ # Be careful of the cache. For large dataset like video, it may lead to
+ # OOM issue.
+ dataset = dataset.cache()
+
+ # Add repetition required for tf data service
+ if self._params.use_tf_data_service:
+ dataset = dataset.repeat()
+
+ return dataset
+
+ def __call__(
+ self,
+ input_context: Optional[tf.distribute.InputContext] = None
+ ) -> tf.data.Dataset:
+
+ def dataset_fn():
+ return self._dataset_fn(input_context=input_context)
+
+ # Add tf data service
+ if self._params.use_tf_data_service:
+ raise ValueError('tf data service is not supported.')
+ dataset = dataset_fn()
+
+ dataset = dataset.prefetch(32) # this will be done on local hosts
+ return dataset
diff --git a/official/projects/videoglue/datasets/video_classification.py b/official/projects/videoglue/datasets/video_classification.py
new file mode 100644
index 00000000000..0c06020fd49
--- /dev/null
+++ b/official/projects/videoglue/datasets/video_classification.py
@@ -0,0 +1,210 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Thd video classification dataset generator."""
+import os
+from typing import Any, Mapping, Optional
+
+from dmvr import video_dataset
+import tensorflow as tf, tf_keras
+
+from official.projects.videoglue.datasets.common import utils
+from official.vision.ops import augment
+
+
+class VideoClassificationBaseFactory(video_dataset.BaseVideoDatasetFactory):
+ """VideoClassification dataset factory."""
+
+ _BASE_DIR = '/tmp'
+
+ _TABLES = {
+ 'train': 'sample.tfrecord',
+ 'test': 'sample.tfrecord',
+ }
+
+ _SUBSETS = ('train', 'test')
+
+ _NUM_CLASSES = 400
+ _INPUT_LABEL_INDEX_KEY = 'clip/label/index'
+ _OUTPUT_LABEL_INDEX_KEY = 'label'
+
+ def __init__(
+ self,
+ subset: str = 'train'):
+ """Initializes the factory."""
+
+ if subset not in self._SUBSETS:
+ raise ValueError(f'Invalid subset "{subset}".'
+ f' The available subsets are: {self._SUBSETS}')
+
+ table_name = os.path.join(self._BASE_DIR, self._TABLES[subset])
+ shards = utils.get_shards(table_name)
+
+ super().__init__(shards)
+
+ def _build(
+ self,
+ is_training: bool = True,
+ # Video related parameters.
+ num_frames: int = 32,
+ temporal_stride: int = 1,
+ sample_from_segments: bool = False,
+ # Image related parameters.
+ min_resize: int = 256,
+ crop_size: int = 224,
+ zero_centering_image: bool = False,
+ random_flip_image: bool = True,
+ augmentation_type: str = 'VGG',
+ augmentation_params: Optional[Mapping[str, Any]] = None,
+ randaug_params: Optional[Mapping[str, Any]] = None,
+ autoaug_params: Optional[Mapping[str, Any]] = None,
+ mixup_cutmix_params: Optional[Mapping[str, Any]] = None,
+ # Test related parameters,
+ num_test_clips: int = 1,
+ multi_crop: bool = False,
+ # Label related parameters.
+ one_hot_label: bool = True,
+ get_label_str: bool = False):
+ """Default builder for this dataset.
+
+ Args:
+ is_training: whether or not in training mode.
+ num_frames: number of frames per subclip.
+ temporal_stride: temporal stride to sample frames.
+ sample_from_segments: Whether to sample frames from segments of a video.
+ If True, the temporal_stride and num_test_clips will be ignored.
+ min_resize: frames are resized so that min width/height is min_resize.
+ crop_size: final size of the frame after cropping the resized frames.
+ zero_centering_image: whether to have images between [-1, 1] or [0, 1].
+ random_flip_image: If True, frames are randomly horizontal flipped during
+ the training.
+ augmentation_type: The data augmentation style applied on images.
+ augmentation_params: A dictionary of params for data augmentation.
+ randaug_params: A dictionary of params for RandAug policy.
+ autoaug_params: A dictionary of params for AutoAug policy.
+ mixup_cutmix_params: A dictionary of params for Mixup and Cutmix data
+ augmentation policy.
+ num_test_clips: number of test clip (1 by default). If more than one, this
+ will sample multiple linearly spaced clips within each video at test
+ time. If 1, then a single clip in the middle of the video is sampled.
+ multi_crop: if True, 3 crops will be sampled from one video clip.
+ one_hot_label: whether or not to return one hot version of labels.
+ get_label_str: whether or not to return label as text.
+ """
+ utils.add_image(
+ parser_builder=self.parser_builder,
+ sampler_builder=self.sampler_builder,
+ decoder_builder=self.decoder_builder,
+ preprocessor_builder=self.preprocessor_builder,
+ postprocessor_builder=self.postprocessor_builder,
+ input_feature_name='image/encoded',
+ is_training=is_training,
+ num_frames=num_frames,
+ temporal_stride=temporal_stride,
+ sample_from_segments=sample_from_segments,
+ num_test_clips=num_test_clips,
+ multi_crop=multi_crop,
+ min_resize=min_resize,
+ crop_size=crop_size,
+ zero_centering_image=zero_centering_image,
+ random_flip_image=random_flip_image,
+ augmentation_type=augmentation_type,
+ augmentation_params=augmentation_params,
+ randaug_params=randaug_params,
+ autoaug_params=autoaug_params,
+ sync_random_state=is_training)
+
+ utils.add_label(
+ parser_builder=self.parser_builder,
+ decoder_builder=self.decoder_builder,
+ preprocessor_builder=self.preprocessor_builder,
+ one_hot_label=one_hot_label,
+ input_label_index_feature_name=self._INPUT_LABEL_INDEX_KEY,
+ output_label_index_feature_name=self._OUTPUT_LABEL_INDEX_KEY,
+ num_classes=self._NUM_CLASSES,
+ add_label_name=get_label_str)
+
+ if is_training and mixup_cutmix_params is not None:
+ def mixup_and_cutmix_fn(inputs):
+ image = inputs['image']
+ label = inputs['label']
+ if one_hot_label:
+ label = tf.math.argmax(label, axis=-1)
+ augmenter = augment.MixupAndCutmix(
+ mixup_alpha=mixup_cutmix_params['mixup_alpha'],
+ cutmix_alpha=mixup_cutmix_params['cutmix_alpha'],
+ prob=mixup_cutmix_params['prob'],
+ label_smoothing=mixup_cutmix_params['label_smoothing'],
+ num_classes=self._NUM_CLASSES)
+ image, label = augmenter(image, label)
+ inputs['image'] = image
+ inputs['label'] = label
+ return inputs
+
+ self.postprocessor_builder.add_fn(
+ mixup_and_cutmix_fn, fn_name='mixup_and_cutmix')
+
+ def tables(self):
+ return self._TABLES
+
+
+class MomentsInTimeFactory(VideoClassificationBaseFactory):
+ """Moments-in-time dataset."""
+
+ _BASE_DIR = '/tmp'
+
+ _TABLES = {
+ 'train': 'tfse_moments_in_time-train.tfrecord@1024',
+ 'test': 'tfse_moments_in_time-validation.tfrecord@1024',
+ }
+ _NUM_CLASSES = 339
+
+
+class Sthv2Factory(VideoClassificationBaseFactory):
+ """Sth-sth v2 dataset."""
+
+ _BASE_DIR = '/tmp'
+
+ _TABLES = {
+ 'train': 'tfse_something-v2-train.tfrecord@128',
+ 'test': 'tfse_something-v2-validation.tfrecord@128',
+ }
+
+ _NUM_CLASSES = 174
+
+
+class Diving48Factory(VideoClassificationBaseFactory):
+ """Diving48 dataset."""
+
+ _BASE_DIR = '/tmp'
+
+ _TABLES = {
+ 'train': 'tfse_diving48-train.tfrecord@200',
+ 'test': 'tfse_diving48-test.tfrecord@200',
+ }
+
+ _NUM_CLASSES = 48
+
+
+class Kinetics400Factory(VideoClassificationBaseFactory):
+ """Kinetics400 dataset."""
+
+ _BASE_DIR = '/tmp'
+
+ _TABLES = {
+ 'train': 'tfse_kinetics400-train.tfrecord@496',
+ 'test': 'tfse_kinetics400-val.tfrecord@141',
+ }
+
+ _NUM_CLASSES = 400
diff --git a/official/projects/videoglue/docs/VideoGLUE-fig2-v2.png b/official/projects/videoglue/docs/VideoGLUE-fig2-v2.png
new file mode 100644
index 00000000000..a795fd02401
Binary files /dev/null and b/official/projects/videoglue/docs/VideoGLUE-fig2-v2.png differ
diff --git a/official/projects/videoglue/evaluation/spatiotemporal_action_localization_evaluator.py b/official/projects/videoglue/evaluation/spatiotemporal_action_localization_evaluator.py
new file mode 100644
index 00000000000..ade1a9af71a
--- /dev/null
+++ b/official/projects/videoglue/evaluation/spatiotemporal_action_localization_evaluator.py
@@ -0,0 +1,234 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""The evaluator for the spatiotemporal action localization task."""
+from typing import Mapping
+
+from absl import logging
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from object_detection.utils import object_detection_evaluation
+
+# 60 filtered classes used for reporting evaluation results.
+_AVA_LABELS_60 = frozenset([
+ 1, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 17, 20, 22, 24, 26, 27, 28,
+ 29, 30, 34, 36, 37, 38, 41, 43, 45, 46, 47, 48, 49, 51, 52, 54, 56, 57, 58,
+ 59, 60, 61, 62, 63, 64, 65, 66, 67, 68, 69, 70, 72, 73, 74, 76, 77, 78, 79,
+ 80
+])
+
+_AVA_LABELS_STR_80 = (
+ 'bend/bow (at the waist)', 'crawl', 'crouch/kneel', 'dance', 'fall down',
+ 'get up', 'jump/leap', 'lie/sleep', 'martial art', 'run/jog', 'sit',
+ 'stand', 'swim', 'walk', 'answer phone', 'brush teeth',
+ 'carry/hold (an object)', 'catch (an object)', 'chop',
+ 'climb (e.g., a mountain)', 'clink glass', 'close (e.g., a door, a box)',
+ 'cook', 'cut', 'dig', 'dress/put on clothing', 'drink',
+ 'drive (e.g., a car, a truck)', 'eat', 'enter', 'exit', 'extract',
+ 'fishing', 'hit (an object)', 'kick (an object)', 'lift/pick up',
+ 'listen (e.g., to music)', 'open (e.g., a window, a car door)', 'paint',
+ 'play board game', 'play musical instrument', 'play with pets',
+ 'point to (an object)', 'press', 'pull (an object)', 'push (an object)',
+ 'put down', 'read', 'ride (e.g., a bike, a car, a horse)', 'row boat',
+ 'sail boat', 'shoot', 'shovel', 'smoke', 'stir', 'take a photo',
+ 'text on/look at a cellphone', 'throw', 'touch (an object)',
+ 'turn (e.g., a screwdriver)', 'watch (e.g., TV)', 'work on a computer',
+ 'write', 'fight/hit (a person)', 'give/serve (an object) to (a person)',
+ 'grab (a person)', 'hand clap', 'hand shake', 'hand wave', 'hug (a person)',
+ 'kick (a person)', 'kiss (a person)', 'lift (a person)',
+ 'listen to (a person)', 'play with kids', 'push (another person)',
+ 'sing to (e.g., self, a person, a group)',
+ 'take (an object) from (a person)',
+ 'talk to (e.g., self, a person, a group)', 'watch (a person)'
+)
+
+_IMAGE_SIZE = 512
+
+
+class SpatiotemporalActionLocalizationEvaluator(object):
+ """Spatiotemporal action localization evaluation metric class."""
+
+ def __init__(self, iou_threshold: float = 0.5):
+ self._prediction_scores_list = []
+ self._prediction_boxes_list = []
+ self._groundtruth_classes_list = []
+ self._groundtruth_boxes_list = []
+
+ ava_categories = []
+ for idx, name in enumerate(_AVA_LABELS_STR_80):
+ ava_categories.append({'id': idx + 1, 'name': name})
+
+ self._evaluator = object_detection_evaluation.PascalDetectionEvaluator(
+ categories=ava_categories,
+ matching_iou_threshold=iou_threshold)
+
+ @property
+ def name(self):
+ return 'mAP@.5IOU'
+
+ def _maybe_convert_to_numpy(self, outputs):
+ """Converts tf.Tensor to numpy arrays."""
+ if outputs:
+ outputs = tf.nest.map_structure(
+ lambda x: x.numpy() if isinstance(x, tf.Tensor) else x, outputs)
+ numpy_outputs = {}
+ for key, val in outputs.items():
+ if isinstance(val, tuple):
+ val = np.concatenate(val)
+ numpy_outputs[key] = val
+ else:
+ numpy_outputs = outputs
+
+ return numpy_outputs
+
+ def _format_ava_eval_data(
+ self, scores, boxes, groundtruth_classes, groundtruth_boxes):
+ """Converts data in the correct evaluation format.
+
+ Args:
+ scores: float32 numpy array of shape [N, C] for prediction scores.
+ boxes: float32 numpy array of shape [N, 4] for prediction boxes.
+ groundtruth_classes: float32 numpy array of shape [N] indicates 0-index
+ groundtruth classes.
+ groundtruth_boxes: float32 numpy array of shape [N, 4] of corresponding
+ groundtruth boxes.
+
+ Returns:
+ output_dict: the dictionary contains formatted numpy arrays.
+ """
+ instances_mask = np.sum(boxes, axis=-1) > 0
+ scores = scores[instances_mask]
+ boxes = boxes[instances_mask]
+
+ groundtruth_mask = np.sum(groundtruth_boxes, axis=-1) > 0
+ valid_gt_classes = groundtruth_classes[groundtruth_mask]
+ valid_gt_boxes = groundtruth_boxes[groundtruth_mask]
+
+ # There are circumstances that no groundtruth is provided for current clip.
+ if valid_gt_classes.size == 0:
+ return None
+
+ formatted_groundtruth_boxes = []
+ formatted_groundtruth_classes = []
+ formatted_detection_boxes = []
+ formatted_detection_classes = []
+ formatted_detection_scores = []
+ for i in range(valid_gt_boxes.shape[0]):
+ # Only evaluate AVA 60-classes.
+ if (valid_gt_classes[i] + 1) not in _AVA_LABELS_60:
+ continue
+ formatted_groundtruth_boxes.append(valid_gt_boxes[i] * _IMAGE_SIZE)
+ formatted_groundtruth_classes.append(valid_gt_classes[i] + 1)
+
+ for i in range(scores.shape[0]):
+ one_scores = scores[i].tolist()
+ for cls_idx, score in enumerate(one_scores):
+ # Only evaluate AVA 60-classes.
+ if (cls_idx + 1) not in _AVA_LABELS_60:
+ continue
+ formatted_detection_boxes.append(boxes[i] * _IMAGE_SIZE)
+ formatted_detection_classes.append(cls_idx + 1)
+ formatted_detection_scores.append(score)
+
+ if not formatted_groundtruth_boxes or not formatted_detection_boxes:
+ return None
+ else:
+ output_dict = {
+ 'groundtruth_boxes': formatted_groundtruth_boxes,
+ 'groundtruth_classes': formatted_groundtruth_classes,
+ 'detection_boxes': formatted_detection_boxes,
+ 'detection_classes': formatted_detection_classes,
+ 'detection_scores': formatted_detection_scores,
+ }
+ return output_dict
+
+ def update_state(self, step_outputs: Mapping[str, tf.Tensor]):
+ """Updates per-step evaluation states by aggregating prediction results.
+
+ Args:
+ step_outputs: A dictionary contains tensors for the evaluation.
+ * predictions: the model prediction score in shape [B, N, C].
+ * instances_position: the corresponding boxes for each predictions.
+ * nonmerge_label: the 0-indexed groundtruth label.
+ * nonmerge_instances_position: the corresponding groundtruth boxes for
+ each label. Note that the boxes here could be duplicated due to the
+ multi-labels.
+ """
+ filtered_step_outputs = {
+ 'scores': step_outputs['predictions'],
+ 'boxes': step_outputs['instances_position'],
+ 'groundtruth_classes': step_outputs['nonmerge_label'],
+ 'groundtruth_boxes': step_outputs['nonmerge_instances_position'],
+ }
+ outputs_np = self._maybe_convert_to_numpy(filtered_step_outputs)
+
+ self._prediction_scores_list.append(outputs_np['scores'])
+ self._prediction_boxes_list.append(outputs_np['boxes'])
+ self._groundtruth_classes_list.append(outputs_np['groundtruth_classes'])
+ self._groundtruth_boxes_list.append(outputs_np['groundtruth_boxes'])
+
+ def reset_states(self):
+ """Resets evaluation states."""
+ self._evaluator.clear()
+
+ self._prediction_scores_list = []
+ self._prediction_boxes_list = []
+ self._groundtruth_classes_list = []
+ self._groundtruth_boxes_list = []
+
+ def result(self):
+ """Fetches the final evaluation results."""
+ groundtruth_classes = np.concatenate(self._groundtruth_classes_list)
+ groundtruth_boxes = np.concatenate(self._groundtruth_boxes_list)
+ prediction_scores = np.concatenate(self._prediction_scores_list)
+ prediction_boxes = np.concatenate(self._prediction_boxes_list)
+
+ num_samples = groundtruth_classes.shape[0]
+ skipped = 0
+ for batch_id in range(num_samples):
+ output_dict = self._format_ava_eval_data(
+ scores=prediction_scores[batch_id],
+ boxes=prediction_boxes[batch_id],
+ groundtruth_classes=groundtruth_classes[batch_id],
+ groundtruth_boxes=groundtruth_boxes[batch_id])
+ if output_dict is None:
+ skipped += 1
+ continue
+
+ groundtruth_dict = {
+ 'groundtruth_boxes': np.array(
+ output_dict['groundtruth_boxes'], dtype=float),
+ 'groundtruth_classes': np.array(
+ output_dict['groundtruth_classes'], dtype=int),
+ 'groundtruth_difficult': np.zeros(
+ len(output_dict['groundtruth_boxes']), dtype=bool),
+ }
+ detections_dict = {
+ 'detection_boxes': np.array(
+ output_dict['detection_boxes'], dtype=float),
+ 'detection_classes': np.array(
+ output_dict['detection_classes'], dtype=int),
+ 'detection_scores': np.array(
+ output_dict['detection_scores'], dtype=float),
+ }
+ self._evaluator.add_single_ground_truth_image_info(
+ image_id=batch_id, groundtruth_dict=groundtruth_dict)
+ self._evaluator.add_single_detected_image_info(
+ image_id=batch_id, detections_dict=detections_dict)
+
+ metrics = self._evaluator.evaluate()
+ logging.info('Evaluated on %d videos, skipped %d videos.',
+ num_samples - skipped, skipped)
+ return {'mAP@.5IOU': metrics['PascalBoxes_Precision/mAP@0.5IOU']}
diff --git a/official/projects/videoglue/modeling/backbones/vit_3d.py b/official/projects/videoglue/modeling/backbones/vit_3d.py
new file mode 100644
index 00000000000..1ace1ad8a85
--- /dev/null
+++ b/official/projects/videoglue/modeling/backbones/vit_3d.py
@@ -0,0 +1,357 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""The vision transformer using 3D projection for video inputs."""
+
+from typing import Any, Optional, Tuple, Union
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.projects.videoglue.configs import backbones_3d as cfg
+from official.vision.modeling.backbones import factory
+from official.vision.modeling.backbones import vit
+
+Encoder = vit.Encoder
+TokenLayer = vit.TokenLayer
+layers = tf_keras.layers
+
+
+class AddSeparablePositionEmbs(tf_keras.layers.Layer):
+ """Adds (optionally learned) positional embeddings to the inputs."""
+
+ def __init__(self,
+ posemb_init: Optional[tf_keras.initializers.Initializer] = None,
+ posemb_origin_shape: Optional[Tuple[int, int]] = None,
+ posemb_target_shape: Optional[Tuple[int, int]] = None,
+ **kwargs):
+ """Constructs Postional Embedding module.
+
+ The logic of this module is: the learnable positional embeddings length will
+ be determined by the inputs_shape or posemb_origin_shape (if provided)
+ during the construction. If the posemb_target_shape is provided and is
+ different from the positional embeddings length, the embeddings will be
+ interpolated during the forward call.
+
+ Args:
+ posemb_init: The positional embedding initializer.
+ posemb_origin_shape: The intended positional embedding shape.
+ posemb_target_shape: The potential target shape positional embedding may
+ be interpolated to.
+ **kwargs: other args.
+ """
+ super().__init__(**kwargs)
+ self.posemb_init = posemb_init
+ self.posemb_origin_shape = posemb_origin_shape
+ self.posemb_target_shape = posemb_target_shape
+
+ def build(self, inputs_shape):
+ """Builds the separable positional embedding layer."""
+ if self.posemb_origin_shape is not None:
+ nt = self.posemb_origin_shape[0]
+ nl = self.posemb_origin_shape[1]
+ nc = inputs_shape[-1]
+ else:
+ _, nt, nl, nc = inputs_shape
+
+ self._pos_embedding_time = self.add_weight(
+ 'pos_embedding_time',
+ (1, nt, nc),
+ dtype=tf.float32,
+ initializer=tf_keras.initializers.TruncatedNormal(0.02))
+ self._pos_embedding_space = self.add_weight(
+ 'pos_embedding_space',
+ (1, nl, nc),
+ dtype=tf.float32,
+ initializer=tf_keras.initializers.TruncatedNormal(0.02))
+
+ def _interpolate(self, pos_embedding: tf.Tensor,
+ from_shape: Tuple[int, int],
+ to_shape: Tuple[int, int]) -> tf.Tensor:
+ """Interpolates the positional embeddings."""
+ logging.info('Interpolating postional embedding from length: %s to %s',
+ from_shape, to_shape)
+ grid_emb = tf.reshape(pos_embedding, [1] + list(from_shape) + [-1])
+ # NOTE: Using BILINEAR interpolation by default.
+ grid_emb = tf.image.resize(grid_emb, to_shape)
+ return tf.reshape(grid_emb, [1, to_shape[0] * to_shape[1], -1])
+
+ def call(self, inputs: tf.Tensor, inputs_positions: Any = None) -> tf.Tensor:
+ # inputs.shape is (batch_size, time_len, seq_len, emb_dim).
+ del inputs_positions
+ pos_embedding_time = self._pos_embedding_time
+ if inputs.shape[1] != pos_embedding_time.shape[1]:
+ pos_embedding_time = self._interpolate(
+ pos_embedding_time,
+ from_shape=(1, self.posemb_origin_shape[0]), # pyrefly: ignore[unsupported-operation]
+ to_shape=(1, self.posemb_target_shape[0])) # pyrefly: ignore[unsupported-operation]
+
+ pos_embedding_space = self._pos_embedding_space
+ if inputs.shape[2] != pos_embedding_space.shape[1]:
+ pos_embedding_space = self._interpolate(
+ pos_embedding_space,
+ from_shape=(1, self.posemb_origin_shape[1]), # pyrefly: ignore[unsupported-operation]
+ to_shape=(1, self.posemb_target_shape[1])) # pyrefly: ignore[unsupported-operation]
+
+ pos_embedding_time = tf.cast(pos_embedding_time[:, :, None, :],
+ inputs.dtype)
+ pos_embedding_space = tf.cast(pos_embedding_space[:, None, :, :],
+ inputs.dtype)
+ return inputs + pos_embedding_time + pos_embedding_space
+
+
+class VisionTransformer3D(tf_keras.Model):
+ """Class to build VisionTransformer-3D family model.
+
+ The Vision Transformer architecture with the modification on the first
+ patch2token layer in order to process video inputs.
+ Reference: https://arxiv.org/abs/2010.11929
+ """
+
+ def __init__(
+ self,
+ variant: str = 'native',
+ mlp_dim: int = 3072,
+ num_heads: int = 12,
+ num_layers: int = 12,
+ attention_dropout_rate: float = 0.0,
+ dropout_rate: float = 0.1,
+ init_stochastic_depth_rate: float = 0.0,
+ input_specs: layers.InputSpec = layers.InputSpec(
+ shape=[None, None, None, None, 3]),
+ temporal_patch_size: int = 4,
+ spatial_patch_size: int = 16,
+ hidden_size: int = 768,
+ representation_size: int = 0,
+ pooler: str = 'token',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ original_init: bool = True,
+ pos_embed_shape: Optional[
+ Union[Tuple[int, int], Tuple[int, int, int]]] = None):
+ """VisionTransformer initialization function.
+
+ Args:
+ variant: the implementation variant to use. Currently supporting
+ ['native', 'mae'].
+ mlp_dim: the mlp dimension in the transformer encoder.
+ num_heads: number of heads in the transformer encoder.
+ num_layers: number of layers in the transformer encoder.
+ attention_dropout_rate: dropout probability within the attention layer.
+ dropout_rate: the output layer dropout rate.
+ init_stochastic_depth_rate: the initial stochastic depth rate.
+ input_specs: the input shape.
+ temporal_patch_size: the patch size for the temporal dimension.
+ spatial_patch_size: the patch size for the spatial dimension.
+ hidden_size: the projection hidden size for the first layer.
+ representation_size: the feature size of representation.
+ pooler: type of pooler to use. Accept 'none', 'token' or 'gap'.
+ kernel_regularizer: kernel regularizer.
+ original_init: whether to use the original init described in the paper.
+ pos_embed_shape: the original positional embedding shape to use. If None,
+ the positional embedding shape will be inferred from the inputs.
+ """
+ self._variant = variant
+ self._mlp_dim = mlp_dim
+ self._num_heads = num_heads
+ self._num_layers = num_layers
+ self._hidden_size = hidden_size
+ self._representation_size = representation_size
+ self._pooler = pooler
+ self._input_specs = input_specs
+ self._temporal_patch_size = temporal_patch_size
+ self._spatial_patch_size = spatial_patch_size
+ self._kernel_regularizer = kernel_regularizer
+ self._original_init = original_init
+ self._pos_embed_shape = pos_embed_shape
+
+ self._patch_size = (
+ self._temporal_patch_size,
+ self._spatial_patch_size,
+ self._spatial_patch_size,
+ )
+ nt = self._input_specs.shape[1] // self._temporal_patch_size
+ nh = self._input_specs.shape[2] // self._spatial_patch_size
+ nw = self._input_specs.shape[3] // self._spatial_patch_size
+
+ inputs = tf_keras.Input(shape=input_specs.shape[1:])
+ add_pos_embed = True
+ if self._variant == 'native':
+ x = self._tokenize(inputs)
+ elif self._variant == 'mae':
+ x = self._mae_tokenize(inputs)
+ # NOTE: MAE variant adds pos_embed in the tokenizer.
+ add_pos_embed = False
+ else:
+ raise ValueError(
+ 'Unrecognized ViT-3D implementation variant choice: %s' %
+ variant)
+
+ # If we want to add a class token, add it here.
+ if pooler == 'token':
+ x = TokenLayer(name='cls')(x)
+
+ x = vit.Encoder(
+ num_layers=num_layers,
+ mlp_dim=mlp_dim,
+ num_heads=num_heads,
+ dropout_rate=dropout_rate,
+ attention_dropout_rate=attention_dropout_rate,
+ kernel_regularizer=kernel_regularizer,
+ kernel_initializer='glorot_uniform' if original_init else dict(
+ class_name='TruncatedNormal', config=dict(stddev=.02)),
+ init_stochastic_depth_rate=init_stochastic_depth_rate,
+ pos_embed_origin_shape=pos_embed_shape,
+ pos_embed_target_shape=None,
+ add_pos_embed=add_pos_embed)(x)
+
+ if pooler == 'token':
+ x = x[:, 0]
+ elif pooler == 'gap':
+ x = tf.reduce_mean(x, axis=1)
+ elif pooler == 'none':
+ x = tf.reshape(x, [-1, nt, nh, nw, x.shape[-1]], name='encoded_tokens')
+ else:
+ raise ValueError(f'unrecognized pooler type: {pooler}')
+
+ if representation_size:
+ x = tf_keras.layers.Dense(
+ representation_size,
+ kernel_regularizer=kernel_regularizer,
+ name='pre_logits',
+ kernel_initializer='lecun_normal' if original_init else 'he_uniform')(
+ x)
+ x = tf.nn.tanh(x)
+ else:
+ x = tf.identity(x, name='pre_logits')
+
+ if pooler == 'none':
+ endpoints = {'encoded_tokens': x}
+ else:
+ endpoints = {
+ 'pre_logits':
+ tf.reshape(x, [-1, 1, 1, 1, representation_size or hidden_size])
+ }
+
+ super().__init__(inputs=inputs, outputs=endpoints)
+
+ def _tokenize(self, inputs: tf.Tensor):
+ """The first layer to tokenize and project the input tensor."""
+ x = tf_keras.layers.Conv3D(
+ filters=self._hidden_size,
+ kernel_size=self._patch_size,
+ strides=self._patch_size,
+ padding='valid',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=('lecun_normal'
+ if self._original_init else 'he_uniform'))(inputs)
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ time_axis, rows_axis, cols_axis = (1, 2, 3)
+ else:
+ time_axis, rows_axis, cols_axis = (2, 3, 4)
+ # The reshape below assumes the data_format is 'channels_last,' so
+ # transpose to that. Once the data is flattened by the reshape, the
+ # data_format is irrelevant, so no need to update
+ # tf_keras.backend.image_data_format.
+ x = tf.transpose(x, perm=[0, 2, 3, 4, 1])
+
+ nt = self._input_specs.shape[time_axis] // self._temporal_patch_size
+ nh = self._input_specs.shape[rows_axis] // self._spatial_patch_size
+ nw = self._input_specs.shape[cols_axis] // self._spatial_patch_size
+ seq_len = nt * nh * nw
+ x = tf.reshape(x, [-1, seq_len, self._hidden_size])
+ return x
+
+ def _mae_tokenize(self, inputs: tf.Tensor):
+ """The first layer to tokenize and project the input tensor."""
+ # Follow the same normalization setting as the original implementation:
+ # https://github.com/facebookresearch/mae_st/blob/d752324a4a59aab6454236f33b0cd5849f1e600a/util/kinetics.py#L48-L49
+ # The inputs are supposed to be normalized to [0, 1] before applying the
+ # following mean/std.
+ mean = tf.constant((0.45, 0.45, 0.45), dtype=inputs.dtype)
+ std = tf.constant((0.225, 0.225, 0.225), dtype=inputs.dtype)
+ inputs = (inputs - mean) / std
+ x = tf_keras.layers.Conv3D(
+ filters=self._hidden_size,
+ kernel_size=self._patch_size,
+ strides=self._patch_size,
+ padding='valid',
+ kernel_regularizer=self._kernel_regularizer,
+ kernel_initializer=('lecun_normal'
+ if self._original_init else 'he_uniform'))(inputs)
+ if tf_keras.backend.image_data_format() == 'channels_last':
+ time_axis, rows_axis, cols_axis = (1, 2, 3)
+ else:
+ time_axis, rows_axis, cols_axis = (2, 3, 4)
+ # The reshape below assumes the data_format is 'channels_last,' so
+ # transpose to that. Once the data is flattened by the reshape, the
+ # data_format is irrelevant, so no need to update
+ # tf_keras.backend.image_data_format.
+ x = tf.transpose(x, perm=[0, 2, 3, 4, 1])
+
+ nc = x.shape[-1]
+ nt = self._input_specs.shape[time_axis] // self._temporal_patch_size
+ nh = self._input_specs.shape[rows_axis] // self._spatial_patch_size
+ nw = self._input_specs.shape[cols_axis] // self._spatial_patch_size
+
+ x = tf.reshape(x, [-1, nt, nh * nw, nc])
+ pos_embed_target_shape = (nt, nh * nw)
+ x = AddSeparablePositionEmbs(
+ posemb_init=self._original_init, # pyrefly: ignore[bad-argument-type]
+ posemb_origin_shape=self._pos_embed_shape, # pyrefly: ignore[bad-argument-type]
+ posemb_target_shape=pos_embed_target_shape)(x)
+ x = tf.reshape(x, [-1, nt * nh * nw, nc])
+ return x
+
+
+@factory.register_backbone_builder('vit_3d')
+def build_vit_3d(
+ input_specs: tf_keras.layers.InputSpec,
+ backbone_config: cfg.Backbone3D,
+ norm_activation_config: Any,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None):
+ """Builds ViT-3D model.
+
+ Args:
+ input_specs: the input shape specs.
+ backbone_config: the config for the backbone.
+ norm_activation_config: deprecated. norm and activation config.
+ l2_regularizer: the l2 regularizer.
+
+ Returns:
+ A VisionTransformer3D backbone.
+ """
+ del norm_activation_config
+ backbone_type = backbone_config.type
+ backbone_cfg = backbone_config.get()
+ assert backbone_type == 'vit_3d', (f'Inconsistent backbone type '
+ f'{backbone_type}')
+ backbone_cfg.override(vit.VIT_SPECS[backbone_cfg.model_name])
+
+ return VisionTransformer3D(
+ variant=backbone_cfg.variant,
+ mlp_dim=backbone_cfg.transformer.mlp_dim,
+ num_heads=backbone_cfg.transformer.num_heads,
+ num_layers=backbone_cfg.transformer.num_layers,
+ attention_dropout_rate=backbone_cfg.transformer.attention_dropout_rate,
+ dropout_rate=backbone_cfg.transformer.dropout_rate,
+ init_stochastic_depth_rate=backbone_cfg.init_stochastic_depth_rate,
+ input_specs=input_specs,
+ temporal_patch_size=backbone_cfg.temporal_patch_size,
+ spatial_patch_size=backbone_cfg.patch_size,
+ hidden_size=backbone_cfg.hidden_size,
+ representation_size=backbone_cfg.representation_size,
+ pooler=backbone_cfg.pooler,
+ kernel_regularizer=l2_regularizer,
+ original_init=backbone_cfg.original_init,
+ pos_embed_shape=backbone_cfg.pos_embed_shape)
diff --git a/official/projects/videoglue/modeling/heads/action_transformer.py b/official/projects/videoglue/modeling/heads/action_transformer.py
new file mode 100644
index 00000000000..3690c66ae73
--- /dev/null
+++ b/official/projects/videoglue/modeling/heads/action_transformer.py
@@ -0,0 +1,244 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""The implementation of action transformer head."""
+from typing import Mapping, Optional
+
+import tensorflow as tf, tf_keras
+
+from official.projects.videoglue.modeling.heads import simple
+from official.projects.videoglue.modeling.heads import transformer_decoder
+from official.vision.modeling.layers import roi_aligner
+
+
+def _get_shape(x: tf.Tensor):
+ """Helper function to return shape of a given tensor."""
+ static = x.shape.as_list()
+ dynamic = tf.shape(x)
+ return [dynamic[i] if s is None else s for i, s in enumerate(static)]
+
+
+class ActionTransformerHead(tf_keras.layers.Layer):
+ """A Video Action Transformer Head.
+
+ Reference: Girdhar, Rohit et. al. "Video action transformer network." In CVPR
+ 2019. https://arxiv.org/abs/1812.02707
+ """
+
+ def __init__(
+ self,
+ # parameters for classifier
+ num_hidden_layers: int,
+ num_hidden_channels: int,
+ use_sync_bn: bool,
+ num_classes: int,
+ activation: str = 'relu',
+ dropout_rate: float = 0.0,
+ classifier_norm_epsilon: float = 1e-5,
+ # parameters for RoiAligner
+ crop_size: int = 7,
+ sample_offset: float = 0.5,
+ # parameters for TxDecoder
+ num_tx_channels: int = 768,
+ num_tx_layers: int = 12,
+ num_tx_heads: int = 12,
+ use_bias: bool = True,
+ tx_activation: str = 'gelu',
+ attention_dropout_rate: float = 0.0,
+ layer_norm_epsilon: float = 1e-6,
+ use_positional_embedding: bool = True,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ name: str = 'action_transformer_classifier',
+ **kwargs,
+ ):
+ """Initializer.
+
+ Args:
+ num_hidden_layers: The number of hidden layer in the final classifier.
+ num_hidden_channels: The number of hidden channels in the classifier.
+ use_sync_bn: Whether to use the sync batch norm in the classifier.
+ num_classes: The final number of classes prediction.
+ activation: The activation used in the classifier.
+ dropout_rate: The dropout rate for the classifier.
+ classifier_norm_epsilon: Batchnorm epsilon for the classifier.
+ crop_size: The RoI-Align output crop size.
+ sample_offset: The RoI-Align sample offset.
+ num_tx_channels: The number of channels in the transformer.
+ num_tx_layers: The number of transformer layer.
+ num_tx_heads: The number of transformer head.
+ use_bias: Whether to use bias in the transformer decoder.
+ tx_activation: The activation function to use in the transformer.
+ attention_dropout_rate: The attention dropout rate.
+ layer_norm_epsilon: The layer norm epsilon.
+ use_positional_embedding: Whether to use positional embedding.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object.
+ bias_regularizer: tf_keras.regularizers.Regularizer object.
+ name: The head name.
+ **kwargs: Keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+
+ self._num_hidden_layers = num_hidden_layers
+ self._num_hidden_channels = num_hidden_channels
+ self._use_sync_bn = use_sync_bn
+ self._num_classes = num_classes
+ self._dropout_rate = dropout_rate
+ self._use_positional_embedding = use_positional_embedding
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+
+ if self._use_positional_embedding:
+ self._spatial_mlp = [
+ tf_keras.layers.Dense(
+ 4,
+ use_bias=True,
+ activation='relu',
+ name='spatial_mlp_l1',
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer),
+ tf_keras.layers.Dense(
+ 8,
+ use_bias=True,
+ name='spatial_mlp_l2',
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer),
+ ]
+ self._temporal_mlp = [
+ tf_keras.layers.Dense(
+ 4,
+ use_bias=True,
+ activation='relu',
+ name='temporal_mlp_l1',
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer),
+ tf_keras.layers.Dense(
+ 8,
+ use_bias=True,
+ name='temporal_mlp_l2',
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer),
+ ]
+
+ self._roi_aligner = roi_aligner.MultilevelROIAligner(
+ crop_size=crop_size,
+ sample_offset=sample_offset)
+ self._max_pooler = tf_keras.layers.MaxPool2D(
+ pool_size=(crop_size, crop_size),
+ strides=1,
+ padding='valid')
+
+ if num_tx_layers > 0:
+ self._attention_decoder = transformer_decoder.TransformerDecoder(
+ num_channels=num_tx_channels,
+ num_layers=num_tx_layers,
+ num_heads=num_tx_heads,
+ use_bias=use_bias,
+ activation=tx_activation,
+ dropout_rate=attention_dropout_rate,
+ layer_norm_epsilon=layer_norm_epsilon,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer)
+ else:
+ self._attention_decoder = None
+
+ self._dropout_layer = tf_keras.layers.Dropout(dropout_rate)
+ self._classifier = simple.MLP(
+ num_hidden_layers=self._num_hidden_layers,
+ num_hidden_channels=self._num_hidden_channels,
+ num_output_channels=self._num_classes,
+ use_sync_bn=self._use_sync_bn,
+ norm_epsilon=classifier_norm_epsilon,
+ activation=activation,
+ normalize_inputs=False,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer)
+
+ def _add_positional_embedding(self, inputs):
+ """Adds positional embeddings to the inputs tensor."""
+ # Compute the locations using meshgrid.
+ b, t, h, w = _get_shape(inputs)[:4]
+
+ mesh = tf.meshgrid(tf.range(t), tf.range(h), tf.range(w), indexing='ij')
+ position = tf.cast(
+ tf.tile(
+ tf.expand_dims(tf.stack(mesh, axis=-1), axis=0), [b, 1, 1, 1, 1]),
+ tf.float32)
+
+ # Make the positions relative to center point
+ # The mean of all position coordinates would be the center point anyway.
+ center_position = tf.reduce_mean(position, axis=[1, 2, 3], keepdims=True)
+ position -= center_position
+
+ # Apply learneable layers.
+ temporal_position = position[..., :1]
+ for mlp in self._temporal_mlp:
+ temporal_position = mlp(temporal_position)
+ spatial_position = position[..., 1:]
+ for mlp in self._spatial_mlp:
+ spatial_position = mlp(spatial_position)
+
+ return tf.concat([inputs, temporal_position, spatial_position], axis=-1)
+
+ def call(self,
+ inputs: Mapping[str, tf.Tensor],
+ training: bool = False) -> Mapping[str, tf.Tensor]:
+ """Forward calls.
+
+ Args:
+ inputs: the inputs dictionary contains
+ 'features': the instance embeddings in shape [B, T', H, W, C].
+ 'instances_positions': the instance boxes in shape [B, T, N, 4].
+ 'instances_mask': the validity mask for each instance position, in
+ [B, T, N].
+ training: whether in training mode.
+
+ Returns:
+ the final action classification.
+ """
+ features = inputs['features']
+ instances_position = inputs['instances_position']
+ if features.shape.ndims != 5:
+ raise ValueError('Expected features is a rank-5 tensor. Got shape %s' %
+ features.shape)
+
+ if self._use_positional_embedding:
+ features = self._add_positional_embedding(features)
+
+ # Perform RoI-pooling.
+ h, w = _get_shape(features)[2:4]
+ roi_features = {'0': tf.reduce_mean(features, axis=1)}
+ unnormalized_boxes = instances_position * tf.convert_to_tensor(
+ [h, w, h, w], instances_position.dtype)
+ # roi_features in shape [B, N, h, w, C]
+ roi_features = self._roi_aligner(
+ roi_features, unnormalized_boxes, training=training)
+ # Perform average_pooling on ROI-pooled features.
+ b, n, ch, cw, cc = _get_shape(roi_features)
+ roi_features = tf.reshape(roi_features, [b * n, ch, cw, cc])
+ roi_features = self._max_pooler(roi_features)
+ roi_features = tf.reshape(roi_features, [b, n, cc])
+
+ if self._attention_decoder is None:
+ predictions = roi_features
+ else:
+ outputs = self._attention_decoder(inputs=roi_features,
+ memory=features,
+ training=training)
+ # Get last hidden states and perform final classification.
+ predictions = outputs['hidden_states'][-1]
+
+ predictions = self._dropout_layer(predictions, training=training)
+ outputs = self._classifier(predictions, training=training)
+ return outputs
diff --git a/official/projects/videoglue/modeling/heads/simple.py b/official/projects/videoglue/modeling/heads/simple.py
new file mode 100644
index 00000000000..4f7524f0b46
--- /dev/null
+++ b/official/projects/videoglue/modeling/heads/simple.py
@@ -0,0 +1,251 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Constructs simple task head layers."""
+
+from typing import Any, Mapping, Optional, Union
+
+import tensorflow as tf, tf_keras
+from official.modeling import tf_utils
+from official.vision.modeling.backbones import vit
+
+
+class AddTemporalPositionEmbs(tf_keras.layers.Layer):
+ """Adds learned temporal positional embeddings to the video features."""
+
+ def __init__(self,
+ posemb_init: Optional[tf_keras.initializers.Initializer] = None,
+ **kwargs):
+ """Constructs Postional Embedding module.
+
+ Args:
+ posemb_init: The positional embedding initializer.
+ **kwargs: other args.
+ """
+ super().__init__(**kwargs)
+ self.posemb_init = posemb_init
+
+ def build(self, inputs_shape: Union[tf.TensorShape, list[int]]) -> None:
+ pos_emb_length = inputs_shape[1]
+ pos_emb_shape = (1, pos_emb_length, inputs_shape[-1])
+ self.pos_embedding = self.add_weight(
+ 'pos_embedding', pos_emb_shape, initializer=self.posemb_init)
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ pos_embedding = self.pos_embedding
+ # inputs.shape is (batch_size, temporal_len, spatial_len, emb_dim).
+ pos_embedding = tf.cast(pos_embedding, inputs.dtype)
+ _, t, _, c = inputs.shape
+ inputs = tf.reshape(pos_embedding, [-1, t, 1, c]) + inputs
+ return inputs
+
+
+class MLP(tf_keras.layers.Layer):
+ """Constructs the Multi-Layer Perceptron head."""
+
+ def __init__(
+ self,
+ num_hidden_layers: int,
+ num_hidden_channels: int,
+ num_output_channels: int,
+ use_sync_bn: bool,
+ norm_momentum: float = 0.99,
+ norm_epsilon: float = 1e-5,
+ activation: Optional[str] = None,
+ normalize_inputs: bool = False,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Multi-Layer Perceptron initialization.
+
+ Args:
+ num_hidden_layers: the number of hidden layers in the MLP.
+ num_hidden_channels: the number of hidden nodes.
+ num_output_channels: the number of final output nodes.
+ use_sync_bn: whether to use sync batch norm.
+ norm_momentum: the batch norm momentum.
+ norm_epsilon: the batch norm epsilon.
+ activation: the activation function.
+ normalize_inputs: whether to normalize inputs.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object.
+ bias_regularizer: tf_keras.regularizers.Regularizer object.
+ **kwargs: keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+
+ self._num_hidden_layers = num_hidden_layers
+ self._num_hidden_channels = num_hidden_channels
+ self._num_output_channels = num_output_channels
+ self._use_sync_bn = use_sync_bn
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._activation = activation
+ self._normalize_inputs = normalize_inputs
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+
+ self._layers = []
+ # MLP hidden layers
+ for _ in range(num_hidden_layers):
+ self._layers.append(
+ tf_keras.layers.Dense(
+ num_hidden_channels,
+ use_bias=False,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer))
+ if use_sync_bn:
+ self._layers.append(
+ tf_keras.layers.experimental.SyncBatchNormalization(
+ momentum=norm_momentum,
+ epsilon=norm_epsilon))
+ else:
+ self._layers.append(
+ tf_keras.layers.BatchNormalization(
+ momentum=norm_momentum,
+ epsilon=norm_epsilon))
+ if activation is not None:
+ self._layers.append(tf_utils.get_activation(activation))
+
+ # Projection head
+ self._layers.append(tf_keras.layers.Dense(num_output_channels))
+
+ def call(self, inputs: tf.Tensor, training: bool) -> tf.Tensor:
+ """Forward calls with N-D inputs tensor."""
+ if self._normalize_inputs:
+ inputs = tf.nn.l2_normalize(inputs, axis=-1)
+
+ for layer in self._layers:
+ if isinstance(layer, tf_keras.layers.Layer):
+ inputs = layer(inputs, training=training)
+ else: # activation
+ inputs = layer(inputs)
+ return inputs
+
+ def get_config(self) -> Mapping[str, Any]:
+ """Gets class config parameters."""
+ config_dict = {
+ 'num_hidden_layer': self._num_hidden_layer,
+ 'num_hidden_channels': self._num_hidden_channels,
+ 'num_output_channels': self._num_output_channels,
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'activation': self._activation,
+ 'normalize_inputs': self._normalize_inputs,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_regularizer': self._bias_regularizer,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config: Mapping[str, Any]):
+ """Factory constructor from config."""
+ return cls(**config)
+
+
+class AttentionPoolerClassificationHead(tf_keras.layers.Layer):
+ """Head layer for attention pooling classification network.
+
+ Applies pooling attention, dropout, and classifier projection. Expects input
+ to be vector with shape [batch_size, n, num_channels].
+ """
+
+ def __init__(
+ self,
+ num_heads: int,
+ hidden_size: int,
+ num_classes: int,
+ attention_dropout_rate: float = 0.,
+ dropout_rate: float = 0.,
+ kernel_initializer: str = 'HeNormal',
+ kernel_regularizer: Optional[
+ tf_keras.regularizers.Regularizer] = tf_keras.regularizers.L2(1.5e-5),
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ add_temporal_pos_embed: bool = False,
+ **kwargs):
+ """Implementation for video model classifier head.
+
+ Args:
+ num_heads: number of heads in attention layer.
+ hidden_size: hidden size in attention layer.
+ num_classes: number of output classes for the final logits.
+ attention_dropout_rate: the dropout rate applied to the attention map.
+ dropout_rate: the dropout rate applied to the head projection.
+ kernel_initializer: kernel initializer for the conv operations.
+ kernel_regularizer: kernel regularizer for the conv operations.
+ bias_regularizer: bias regularizer for the conv operations.
+ add_temporal_pos_embed: whether to add temporal position embedding or not.
+ **kwargs: keyword arguments to be passed to this layer.
+ """
+ super().__init__(**kwargs)
+
+ self._num_heads = num_heads
+ self._num_classes = num_classes
+ self._dropout_rate = dropout_rate
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+
+ self._add_pooler_token = vit.TokenLayer(name='pooler_token')
+ self._add_temporal_pos_embed = add_temporal_pos_embed
+ if self._add_temporal_pos_embed:
+ self._pos_embed = AddTemporalPositionEmbs(
+ posemb_init=tf_keras.initializers.RandomNormal(stddev=0.02),
+ name='posembed_final_learnt',
+ )
+
+ self._pooler_attention_layer_norm = tf_keras.layers.LayerNormalization(
+ name='pooler_attention_layer_norm',
+ axis=-1,
+ epsilon=1e-6,
+ dtype=tf.float32)
+
+ self._pooler_attention_layer = tf_keras.layers.MultiHeadAttention(
+ num_heads=num_heads,
+ key_dim=(hidden_size // num_heads),
+ value_dim=None,
+ dropout=attention_dropout_rate,
+ use_bias=True,
+ kernel_initializer='glorot_uniform',
+ name='pooler_attention')
+
+ self._dropout = tf_keras.layers.Dropout(dropout_rate)
+ self._classifier = tf_keras.layers.Dense(
+ num_classes,
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer)
+
+ def call(self, inputs: tf.Tensor) -> tf.Tensor:
+ """Calls the layer with the given inputs."""
+ # Input Shape: [batch_size, n, input_channels]
+ x = inputs
+ tf.assert_rank(x, 4, message=
+ '(b, t, s, c) shaped inputs are required.')
+ if self._add_temporal_pos_embed:
+ x = self._pos_embed(x)
+ _, s, t, c = x.shape
+ x = tf.reshape(x, [-1, s * t, c])
+
+ x = self._pooler_attention_layer_norm(x)
+ x = self._add_pooler_token(x)
+ pooler_token = x[:, 0:1, :]
+ x = x[:, 1:, :]
+ x = self._pooler_attention_layer(query=pooler_token, value=x,
+ return_attention_scores=False)
+
+ if self._dropout_rate and self._dropout_rate > 0:
+ x = self._dropout(x)
+
+ return self._classifier(tf.squeeze(x, axis=1))
diff --git a/official/projects/videoglue/modeling/heads/transformer_decoder.py b/official/projects/videoglue/modeling/heads/transformer_decoder.py
new file mode 100644
index 00000000000..3ad5cb614be
--- /dev/null
+++ b/official/projects/videoglue/modeling/heads/transformer_decoder.py
@@ -0,0 +1,320 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Definition for Transformer decoder."""
+
+from typing import Mapping, Optional, Union, List, Sequence
+from absl import logging
+
+import tensorflow as tf, tf_keras
+
+
+def _get_shape(x: tf.Tensor):
+ """Helper function to return shape of a given tensor."""
+ static = x.shape.as_list()
+ dynamic = tf.shape(x)
+ return [dynamic[i] if s is None else s for i, s in enumerate(static)]
+
+
+class DecoderUnit(tf_keras.layers.Layer):
+ """Constructs the decoder MHA module used in Transformer layers."""
+
+ def __init__(
+ self,
+ num_channels: int,
+ use_bias: bool,
+ dropout_rate: float,
+ activation: str,
+ layer_norm_epsilon: float,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+
+ super().__init__(**kwargs)
+ self._num_channels = num_channels
+ self._use_bias = use_bias
+ self._dropout_rate = dropout_rate
+ self._activation = activation
+ self._layer_norm_epsilon = layer_norm_epsilon
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+
+ def build(self, input_shape: Union[tf.TensorShape, List[tf.TensorShape]]):
+ """Builds the layer.
+
+ Args:
+ input_shape: the input shape for the keras tensor.
+ """
+ # Query, key, and value mapping.
+ self.layer_q = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='query')
+ self.layer_k = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='key')
+ self.layer_v = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='value')
+
+ self.dropout = tf_keras.layers.Dropout(self._dropout_rate)
+ # Note here is a different behavior for contrib_layers.layer_norm and
+ # tf_keras.layers.LayerNormalization, where by default, the former
+ # calculates mean/variance across all axes except the first one
+ # (batch axis), while the latter one computes statistics only on the last
+ # axis.
+ self.layer_norm = tf_keras.layers.LayerNormalization(
+ epsilon=self._layer_norm_epsilon,
+ name='layer_norm')
+
+ self.ffn1 = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=self._activation,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='ffn1')
+ self.ffn2 = tf_keras.layers.Dense(
+ self._num_channels,
+ use_bias=self._use_bias,
+ activation=None,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer,
+ name='ffn2')
+
+ super().build(input_shape)
+
+ def call(self,
+ query: tf.Tensor,
+ memory: Optional[tf.Tensor],
+ training: bool = False) -> Mapping[str, tf.Tensor]:
+ """Forward pass of the Transformer decoder unit.
+
+ Args:
+ query: the input query tensor.
+ memory: the input memory tensor for key/value pairs. If None,
+ self-attention will be performed.
+ training: whether in training mode.
+
+ Returns:
+ outputs: the output dictionary contains 'hidden_states' and
+ 'attention weights' matrix.
+ """
+ if memory is None:
+ memory = query
+
+ tensor_q = self.layer_q(query) # (bs, qlen, inner_dim)
+ tensor_k = self.layer_k(memory) # (bs, klen, inner_dim)
+ tensor_v = self.layer_v(memory) # (bs, klen, inner_dim)
+
+ scores = tf.matmul(tensor_q, tensor_k, transpose_b=True)
+ # Scales attention_scores.
+ dk = tf.cast(_get_shape(tensor_k)[-1], dtype=scores.dtype)
+ scores = scores / tf.math.sqrt(dk)
+
+ # Shape: (bs, seq_len, seq_len)
+ attention_weights = tf.nn.softmax(scores, axis=-1)
+ # Shape: (bs, seq_len, dim_per_head)
+ attention_features = tf.matmul(attention_weights, tensor_v)
+ # Shape: (bs, seq_len, seq_len)
+ attention_features = self.dropout(attention_features, training=training)
+
+ hidden_states = attention_features + tensor_q
+ hidden_states = self.layer_norm(hidden_states)
+
+ # Shape: (bs, seq_len, out_dim)
+ hidden_states = self.ffn1(hidden_states)
+ hidden_states = self.ffn2(hidden_states)
+
+ outputs = {
+ 'hidden_states': hidden_states,
+ 'attention_weights': attention_weights,
+ }
+ return outputs
+
+
+class TransformerDecoderLayer(tf_keras.layers.Layer):
+ """Constructs the main Transformer decoder module which includes MHA + FFN."""
+
+ def __init__(
+ self,
+ num_channels: int,
+ num_heads: int,
+ use_bias: bool,
+ activation: str,
+ dropout_rate: float,
+ layer_norm_epsilon: float,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ name: str = 'decoder_layer',
+ **kwargs):
+ super().__init__(name=name)
+
+ self._num_channels = num_channels
+ self._num_heads = num_heads
+ self._use_bias = use_bias
+ self._activation = activation
+ self._dropout_rate = dropout_rate
+ self._layer_norm_epsilon = layer_norm_epsilon
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+ self._name = name
+
+ self._mha_units = []
+ for i in range(num_heads):
+ self._mha_units.append(
+ DecoderUnit(
+ num_channels=num_channels,
+ use_bias=use_bias,
+ dropout_rate=dropout_rate,
+ activation=activation,
+ layer_norm_epsilon=layer_norm_epsilon,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ name='mha_{}'.format(i)))
+
+ def call(
+ self,
+ inputs: tf.Tensor,
+ memory: Optional[tf.Tensor] = None,
+ training: bool = False
+ ) -> Mapping[str, Union[tf.Tensor, Sequence[tf.Tensor]]]:
+ """Forward pass of the Transformer decoder layer.
+
+ Args:
+ inputs: the input query tensor.
+ memory: the input memory tensor for key/value pairs. If None,
+ self-attention will be performed.
+ training: whether in training mode.
+
+ Returns:
+ outputs: the output dictionary contains 'hidden_states' and
+ 'attention weights' matrix.
+ """
+
+ if memory is None:
+ logging.info('No memory tokens are provided. Performing self-attention '
+ 'on input tokens in TransfomerDecoder.')
+
+ all_head_feats = []
+ all_head_attentions = []
+ for i in range(self._num_heads):
+ outputs = self._mha_units[i](
+ query=inputs, memory=memory, training=training)
+ all_head_feats.append(outputs['hidden_states'])
+ all_head_attentions.append(outputs['attention_weights'])
+
+ outputs = {
+ 'hidden_states': tf.concat(all_head_feats, axis=-1),
+ 'attention_weights': all_head_attentions,
+ }
+ return outputs
+
+
+class TransformerDecoder(tf_keras.layers.Layer):
+ """Constructs the final Transformer decoder stack."""
+
+ def __init__(
+ self,
+ num_channels: int,
+ num_layers: int,
+ num_heads: int,
+ use_bias: bool,
+ activation: str,
+ dropout_rate: float,
+ layer_norm_epsilon: float,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ name: str = 'transformer_decoder',
+ **kwargs):
+ super().__init__(name=name)
+
+ self._num_channels = num_channels
+ self._num_layers = num_layers
+ self._num_heads = num_heads
+ self._use_bias = use_bias
+ self._activation = activation
+ self._dropout_rate = dropout_rate
+ self._layer_norm_epsilon = layer_norm_epsilon
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+
+ self._layers = []
+ for n in range(self._num_layers):
+ self._layers.append(
+ TransformerDecoderLayer(
+ num_channels=num_channels,
+ num_heads=num_heads,
+ use_bias=use_bias,
+ activation=activation,
+ dropout_rate=dropout_rate,
+ layer_norm_epsilon=layer_norm_epsilon,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ name='layer_{}'.format(n)))
+
+ def call(self,
+ inputs: tf.Tensor,
+ memory: Optional[tf.Tensor] = None,
+ training: bool = False) -> Mapping[str, Sequence[tf.Tensor]]:
+ """Forward pass of the Transformer decoder.
+
+ Args:
+ inputs: the input query tensor.
+ memory: the input memory tensor for key/value pairs. If None,
+ self-attention will be performed.
+ training: whether in training mode.
+
+ Returns:
+ outputs: the output dictionary contains 'hidden_states' and
+ 'attention weights' matrix.
+ """
+
+ all_hidden_states = ()
+ all_attentions = ()
+
+ memory_shape = _get_shape(memory) # pyrefly: ignore[bad-argument-type]
+ memory = tf.reshape(memory, [memory_shape[0], -1, memory_shape[-1]])
+ hidden_states = inputs
+
+ for layer in self._layers:
+ layer_outputs = layer(inputs=hidden_states,
+ memory=memory,
+ training=training)
+
+ # layer_outputs is a dictionary with the following keys:
+ # hidden_states, self_attention_weights
+ hidden_states = layer_outputs['hidden_states']
+ all_attentions += (layer_outputs['attention_weights'],)
+
+ # Add last layer
+ all_hidden_states += (hidden_states,)
+
+ outputs = {
+ 'hidden_states': all_hidden_states,
+ 'attention_weights': all_attentions,
+ }
+
+ return outputs
diff --git a/official/projects/videoglue/modeling/video_action_transformer_model.py b/official/projects/videoglue/modeling/video_action_transformer_model.py
new file mode 100644
index 00000000000..b2bf500485e
--- /dev/null
+++ b/official/projects/videoglue/modeling/video_action_transformer_model.py
@@ -0,0 +1,216 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builds the Video Action Transformer Network."""
+from typing import Mapping, Optional, Tuple
+
+import tensorflow as tf, tf_keras
+
+from official.projects.videoglue.configs import spatiotemporal_action_localization as cfg
+from official.projects.videoglue.modeling.backbones import vit_3d # pylint: disable=unused-import
+from official.projects.videoglue.modeling.heads import action_transformer
+from official.vision.modeling import backbones
+from official.vision.modeling import factory_3d as model_factory
+
+
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class VideoActionTransformerModel(tf_keras.Model):
+ """A Video Action Transformer Network.
+
+ Reference: Girdhar, Rohit et. al. "Video action transformer network." In CVPR
+ 2019. https://arxiv.org/abs/1812.02707
+ """
+
+ def __init__(
+ self,
+ backbone: tf_keras.Model,
+ num_classes: int,
+ endpoint_name: str,
+ # parameters for classifier
+ num_hidden_layers: int,
+ num_hidden_channels: int,
+ use_sync_bn: bool,
+ activation: str = 'relu',
+ dropout_rate: float = 0.0,
+ # parameters for RoiAligner
+ crop_size: int = 7,
+ sample_offset: float = 0.5,
+ # parameters for TxDecoder
+ num_tx_channels: int = 128,
+ num_tx_layers: int = 3,
+ num_tx_heads: int = 3,
+ use_bias: bool = True,
+ tx_activation: str = 'gelu',
+ attention_dropout_rate: float = 0.0,
+ layer_norm_epsilon: float = 1e-6,
+ use_positional_embedding: bool = True,
+ input_specs: Optional[Mapping[str, tf_keras.layers.InputSpec]] = None,
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs):
+ """Initialization function.
+
+ Args:
+ backbone: A backbone network.
+ num_classes: The final number of classes prediction.
+ endpoint_name: The endpoint name from the backbone to extract features.
+ num_hidden_layers: The number of hidden layer in the final classifier.
+ num_hidden_channels: The number of hidden channels in the classifier.
+ use_sync_bn: Whether to use the sync batch norm in the classifier.
+ activation: The activation used in the classifier.
+ dropout_rate: The dropout rate for the classifier.
+ crop_size: The RoI-Align output crop size.
+ sample_offset: The RoI-Align sample offset.
+ num_tx_channels: The number of channels in the transformer.
+ num_tx_layers: The number of transformer layer.
+ num_tx_heads: The number of transformer head.
+ use_bias: Whether to use bias in the transformer decoder.
+ tx_activation: The activation function to use in the transformer.
+ attention_dropout_rate: The attention dropout rate.
+ layer_norm_epsilon: The layer norm epsilon.
+ use_positional_embedding: Whether to use positional embedding.
+ input_specs: Specs of the input tensor.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object.
+ bias_regularizer: tf_keras.regularizers.Regularizer object.
+ **kwargs: Keyword arguments to be passed.
+ """
+ if not input_specs:
+ input_specs = {
+ 'image':
+ tf_keras.layers.InputSpec(shape=[None, None, None, None, 3]),
+ 'instances_position':
+ tf_keras.layers.InputSpec(shape=[None, None, 4]),
+ }
+
+ self._num_classes = num_classes
+ self._endpoint_name = endpoint_name
+ self._num_hidden_layers = num_hidden_layers
+ self._num_hidden_channels = num_hidden_channels
+ self._use_sync_bn = use_sync_bn
+ self._activation = activation
+ self._dropout_rate = dropout_rate
+ self._crop_size = crop_size
+ self._sample_offset = sample_offset
+ self._num_tx_channels = num_tx_channels
+ self._num_tx_layers = num_tx_layers
+ self._num_tx_heads = num_tx_heads
+ self._use_bias = use_bias
+ self._tx_activation = tx_activation
+ self._attention_dropout_rate = attention_dropout_rate
+ self._layer_norm_epsilon = layer_norm_epsilon
+ self._use_positional_embedding = use_positional_embedding
+ self._input_specs = input_specs
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+
+ inputs, outputs = self._build_model(backbone, input_specs)
+ super().__init__(inputs=inputs, outputs=outputs, **kwargs)
+ # Move backbone after super() call so Keras is happy.
+ self._backbone = backbone
+
+ def _build_model(
+ self, backbone: tf_keras.Model,
+ input_specs: Mapping[str, tf_keras.layers.InputSpec]
+ ) -> Tuple[Mapping[str, tf.Tensor], tf.Tensor]:
+ """Builds the model network.
+
+ Args:
+ backbone: the model backbone.
+ input_specs: the model input spec to use.
+
+ Returns:
+ Inputs and outputs as a tuple. Inputs are expected to be a dict with
+ base input and positions. Outputs are predictions per instance.
+ """
+
+ inputs = {
+ k: tf_keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
+ }
+ endpoints = backbone(inputs['image'])
+ features = endpoints[self._endpoint_name]
+
+ tx_inputs = {
+ 'features': features,
+ 'instances_position': inputs['instances_position'],
+ }
+ outputs = action_transformer.ActionTransformerHead(
+ num_hidden_layers=self._num_hidden_layers,
+ num_hidden_channels=self._num_hidden_channels,
+ use_sync_bn=self._use_sync_bn,
+ num_classes=self._num_classes,
+ activation=self._activation,
+ dropout_rate=self._dropout_rate,
+ # parameters for RoiAligner
+ crop_size=self._crop_size,
+ sample_offset=self._sample_offset,
+ # parameters for TxDecoder
+ num_tx_channels=self._num_tx_channels,
+ num_tx_layers=self._num_tx_layers,
+ num_tx_heads=self._num_tx_heads,
+ use_bias=self._use_bias,
+ tx_activation=self._tx_activation,
+ attention_dropout_rate=self._attention_dropout_rate,
+ layer_norm_epsilon=self._layer_norm_epsilon,
+ use_positional_embedding=self._use_positional_embedding,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_regularizer=self._bias_regularizer)(tx_inputs)
+ return inputs, outputs
+
+ @property
+ def backbone(self) -> tf_keras.Model:
+ """Returns the backbone of the model."""
+ return self._backbone
+
+
+@model_factory.register_model_builder('video_action_transformer_model')
+def build_video_action_transformer_model(
+ input_specs_dict: Mapping[str, tf_keras.layers.InputSpec],
+ model_config: cfg.VideoActionTransformerModel,
+ num_classes: int,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> VideoActionTransformerModel:
+ """Builds the video action localziation model."""
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs_dict['image'],
+ backbone_config=model_config.backbone,
+ norm_activation_config=model_config.norm_activation,
+ l2_regularizer=l2_regularizer)
+
+ # Norm layer type in the MLP head should same with backbone.
+ if (model_config.norm_activation.use_sync_bn
+ != model_config.head.use_sync_bn):
+ raise ValueError('Should use the same batch normalization type.')
+
+ return VideoActionTransformerModel(
+ backbone=backbone,
+ input_specs=input_specs_dict,
+ num_classes=num_classes,
+ endpoint_name=model_config.endpoint_name,
+ # parameters for classifier
+ num_hidden_layers=model_config.head.num_hidden_layers,
+ num_hidden_channels=model_config.head.num_hidden_channels,
+ use_sync_bn=model_config.head.use_sync_bn,
+ activation=model_config.head.activation,
+ dropout_rate=model_config.head.dropout_rate,
+ crop_size=model_config.head.crop_size,
+ sample_offset=model_config.head.sample_offset,
+ num_tx_channels=model_config.head.num_tx_channels,
+ num_tx_layers=model_config.head.num_tx_layers,
+ num_tx_heads=model_config.head.num_tx_heads,
+ use_bias=model_config.head.use_bias,
+ tx_activation=model_config.head.tx_activation,
+ attention_dropout_rate=model_config.head.attention_dropout_rate,
+ layer_norm_epsilon=model_config.head.layer_norm_epsilon,
+ use_positional_embedding=model_config.head.use_positional_embedding,
+ kernel_regularizer=l2_regularizer)
diff --git a/official/projects/videoglue/modeling/video_classification_model.py b/official/projects/videoglue/modeling/video_classification_model.py
new file mode 100644
index 00000000000..5cf802d39fd
--- /dev/null
+++ b/official/projects/videoglue/modeling/video_classification_model.py
@@ -0,0 +1,212 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builds video classification models."""
+from typing import Any, Mapping, Optional, Union, List, Text
+
+import tensorflow as tf, tf_keras
+
+from official.projects.videoglue.configs import video_classification as cfg
+from official.projects.videoglue.modeling.backbones import vit_3d # pylint: disable=unused-import
+from official.projects.videoglue.modeling.heads import simple
+from official.vision.modeling import backbones
+from official.vision.modeling import factory_3d as model_factory
+
+layers = tf_keras.layers
+
+
+class MultiHeadVideoClassificationModel(tf_keras.Model):
+ """A multi-head video classification class builder."""
+
+ def __init__(
+ self,
+ backbone: tf_keras.Model,
+ num_classes: Union[List[int], int],
+ input_specs: Optional[Mapping[str, tf_keras.layers.InputSpec]] = None,
+ dropout_rate: float = 0.0,
+ attention_num_heads: int = 6,
+ attention_hidden_size: int = 768,
+ attention_dropout_rate: float = 0.0,
+ add_temporal_pos_emb_pooler: bool = False,
+ aggregate_endpoints: bool = False,
+ kernel_initializer: str = 'random_uniform',
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ require_endpoints: Optional[List[Text]] = None,
+ classifier_type: str = 'linear',
+ **kwargs):
+ """Video Classification initialization function.
+
+ Args:
+ backbone: a 3d backbone network.
+ num_classes: `int` number of classes in classification task.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
+ dropout_rate: `float` rate for dropout regularization.
+ attention_num_heads: attention pooler layer number of heads.
+ attention_hidden_size: attention pooler layer hidden size.
+ attention_dropout_rate: attention map dropout regularization.
+ add_temporal_pos_emb_pooler: `bool` adds a learnt temporal position
+ embedding to the attention pooler.
+ aggregate_endpoints: `bool` aggregate all end ponits or only use the
+ final end point.
+ kernel_initializer: kernel initializer for the dense layer.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object. Default to
+ None.
+ bias_regularizer: tf_keras.regularizers.Regularizer object. Default to
+ None.
+ require_endpoints: the required endpoints for prediction. If None or
+ empty, then only uses the final endpoint.
+ classifier_type: choose from 'linear' or 'pooler'.
+ **kwargs: keyword arguments to be passed.
+ """
+ if not input_specs:
+ input_specs = {
+ 'image': layers.InputSpec(shape=[None, None, None, None, 3])
+ }
+ self._self_setattr_tracking = False
+ self._config_dict = {
+ 'backbone': backbone,
+ 'num_classes': num_classes,
+ 'input_specs': input_specs,
+ 'dropout_rate': dropout_rate,
+ 'attention_dropout_rate': attention_dropout_rate,
+ 'attention_num_heads': attention_num_heads,
+ 'attention_hidden_size': attention_hidden_size,
+ 'aggregate_endpoints': aggregate_endpoints,
+ 'kernel_initializer': kernel_initializer,
+ 'kernel_regularizer': kernel_regularizer,
+ 'bias_regularizer': bias_regularizer,
+ 'require_endpoints': require_endpoints,
+ }
+ self._input_specs = input_specs
+ self._backbone = backbone
+
+ inputs = {
+ k: tf_keras.Input(shape=v.shape[1:]) for k, v in input_specs.items()
+ }
+ endpoints = backbone(inputs['image'])
+
+ if classifier_type == 'linear':
+ pool_or_flatten_op = tf_keras.layers.GlobalAveragePooling3D()
+ elif classifier_type == 'pooler':
+ pool_or_flatten_op = lambda x: tf.reshape( # pylint:disable=g-long-lambda
+ x,
+ [
+ tf.shape(x)[0],
+ tf.shape(x)[1],
+ tf.shape(x)[2] * tf.shape(x)[3],
+ tf.shape(x)[4],
+ ],
+ )
+ else:
+ raise ValueError('%s classifier type not supported.' % classifier_type)
+
+ if aggregate_endpoints:
+ pooled_feats = []
+ for endpoint in endpoints.values():
+ x_pool = pool_or_flatten_op(endpoint)
+ pooled_feats.append(x_pool)
+ x = tf.concat(pooled_feats, axis=1)
+ else:
+ if not require_endpoints:
+ # Use the last endpoint for prediction.
+ x = endpoints[max(endpoints.keys())]
+ x = pool_or_flatten_op(x)
+ else:
+ # Concat all the required endpoints for prediction.
+ outputs = []
+ for name in require_endpoints:
+ x = endpoints[name]
+ x = pool_or_flatten_op(x)
+ outputs.append(x)
+ x = tf.concat(outputs, axis=1)
+
+ input_embeddings = tf.identity(x, name='embeddings')
+ num_classes = [num_classes] if isinstance(num_classes, int) else num_classes
+ outputs = []
+ if classifier_type == 'linear':
+ for nc in num_classes:
+ x = tf_keras.layers.Dropout(dropout_rate)(input_embeddings)
+ x = tf_keras.layers.Dense(
+ nc, kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer)(x)
+ outputs.append(x)
+ elif classifier_type == 'pooler':
+ for nc in num_classes:
+ x = simple.AttentionPoolerClassificationHead(
+ num_heads=attention_num_heads,
+ hidden_size=attention_hidden_size,
+ attention_dropout_rate=attention_dropout_rate,
+ num_classes=nc,
+ dropout_rate=dropout_rate,
+ kernel_initializer=kernel_initializer,
+ kernel_regularizer=kernel_regularizer,
+ bias_regularizer=bias_regularizer,
+ add_temporal_pos_embed=add_temporal_pos_emb_pooler)(
+ input_embeddings)
+ outputs.append(x)
+ else:
+ raise ValueError('%s classifier type not supported.')
+
+ super().__init__(inputs=inputs, outputs=outputs, **kwargs)
+
+ @property
+ def checkpoint_items(
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ return dict(backbone=self.backbone)
+
+ @property
+ def backbone(self) -> tf_keras.Model:
+ return self._backbone
+
+ def get_config(self) -> Mapping[str, Any]:
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+
+@model_factory.register_model_builder('mh_video_classification')
+def build_mh_video_classification_model(
+ input_specs: tf_keras.layers.InputSpec,
+ model_config: cfg.MultiHeadVideoClassificationModel,
+ num_classes: Union[List[int], int],
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> MultiHeadVideoClassificationModel:
+ """Builds the video classification model."""
+ input_specs_dict = {'image': input_specs}
+ norm_activation_config = model_config.norm_activation
+ backbone = backbones.factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=model_config.backbone,
+ norm_activation_config=norm_activation_config,
+ l2_regularizer=l2_regularizer)
+
+ model = MultiHeadVideoClassificationModel(
+ backbone=backbone,
+ num_classes=num_classes,
+ input_specs=input_specs_dict,
+ dropout_rate=model_config.dropout_rate,
+ classifier_type=model_config.classifier_type,
+ attention_num_heads=model_config.attention_num_heads,
+ attention_hidden_size=model_config.attention_hidden_size,
+ attention_dropout_rate=model_config.attention_dropout_rate,
+ add_temporal_pos_emb_pooler=model_config.add_temporal_pos_emb_pooler,
+ aggregate_endpoints=model_config.aggregate_endpoints,
+ kernel_regularizer=l2_regularizer,
+ require_endpoints=model_config.require_endpoints)
+ return model
diff --git a/official/projects/videoglue/tasks/multihead_video_classification.py b/official/projects/videoglue/tasks/multihead_video_classification.py
new file mode 100644
index 00000000000..bf239653410
--- /dev/null
+++ b/official/projects/videoglue/tasks/multihead_video_classification.py
@@ -0,0 +1,276 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""HS Video Classification task."""
+from typing import Any, List, Optional, Mapping
+
+import tensorflow as tf, tf_keras
+
+from official.core import task_factory
+from official.projects.videoglue.configs import video_classification as exp_cfg
+from official.projects.videoglue.datasets import dataset_factory
+from official.projects.videoglue.tools import checkpoint_loader
+from official.vision.tasks import video_classification
+
+
+@task_factory.register_task_cls(exp_cfg.MultiHeadVideoClassificationTask)
+class MultiHeadVideoClassificationTask(
+ video_classification.VideoClassificationTask):
+ """Internal video classification task."""
+
+ def _is_multihead(self):
+ """Reports the joint accuracy or not."""
+ label_names = self.task_config.train_data.label_names
+ is_multihead = isinstance(label_names, list) and len(label_names) > 1
+ return is_multihead
+
+ def _get_label_names(self):
+ """Gets the label names."""
+ if self._is_multihead():
+ label_names = self.task_config.train_data.label_names
+ else:
+ label_names = [self.task_config.train_data.label_names]
+ return label_names
+
+ def build_inputs(self, params: exp_cfg.DataConfig, input_context=None):
+ """Builds classification input."""
+ augmentation_type = params.data_augmentation.type
+ augmentation_params = params.data_augmentation.get().as_dict()
+
+ randaug_params = None
+ if params.randaug is not None:
+ randaug_params = params.randaug.as_dict()
+
+ autoaug_params = None
+ if params.autoaug is not None:
+ autoaug_params = params.autoaug.as_dict()
+
+ mixup_cutmix_params = None
+ if params.mixup_cutmix is not None:
+ mixup_cutmix_params = params.mixup_cutmix.as_dict()
+
+ dataset_config = {
+ 'is_training': params.is_training,
+ 'num_frames': params.feature_shape[0],
+ 'temporal_stride': params.temporal_stride,
+ 'sample_from_segments': params.sample_from_segments,
+ 'min_resize': params.min_resize,
+ 'crop_size': params.feature_shape[1],
+ 'zero_centering_image': params.zero_centering_image,
+ 'random_flip_image': params.random_flip_image,
+ 'num_test_clips': params.num_test_clips,
+ 'augmentation_type': augmentation_type,
+ 'augmentation_params': augmentation_params,
+ 'randaug_params': randaug_params,
+ 'autoaug_params': autoaug_params,
+ 'mixup_cutmix_params': mixup_cutmix_params,
+ # TODO(lzyuan): Unify num_test_crops and multi_crop flags.
+ 'multi_crop': params.num_test_crops == 3,
+ }
+ data_loader = dataset_factory.DataLoader(
+ params=params, dataset_config=dataset_config)
+ return data_loader(input_context=input_context)
+
+ def initialize(self, model: tf_keras.Model):
+ """Loads pretrained checkpoint."""
+ if not self.task_config.init_checkpoint:
+ return
+
+ checkpoint_loader.get_checkpoint_loader(
+ model=model,
+ init_checkpoint=self.task_config.init_checkpoint,
+ init_checkpoint_type=self.task_config.init_checkpoint_modules)
+
+ def build_metrics(self, training: bool = True):
+ """Gets streaming metrics for training/validation."""
+ metrics = []
+ for label_name in self._get_label_names():
+ if self._is_multilabel():
+ metrics += [
+ tf_keras.metrics.AUC(
+ curve='ROC',
+ multi_label=self._is_multilabel(),
+ name=f'{label_name}/ROC-AUC'),
+ tf_keras.metrics.RecallAtPrecision(
+ precision=0.95, name=f'{label_name}/RecallAtPrecision95'),
+ tf_keras.metrics.AUC(
+ curve='PR',
+ multi_label=self._is_multilabel(),
+ name=f'{label_name}/PR-AUC'),
+ ]
+ else:
+ metrics += [
+ tf_keras.metrics.CategoricalAccuracy(
+ name=f'{label_name}/accuracy'),
+ tf_keras.metrics.TopKCategoricalAccuracy(
+ k=1, name=f'{label_name}/top_1_accuracy'),
+ tf_keras.metrics.TopKCategoricalAccuracy(
+ k=5, name=f'{label_name}/top_5_accuracy')
+ ]
+
+ if self._is_multihead():
+ metrics.append(
+ tf_keras.metrics.Mean(name='label_joint/accuracy'))
+ return metrics
+
+ def process_metrics(self, metrics: List[Any],
+ labels: List[tf.Tensor],
+ model_outputs: List[tf.Tensor]):
+ """Processes and updates metrics.
+
+ Called when using custom training loop API.
+
+ Args:
+ metrics: a nested structure of metrics objects. The return of function
+ self.build_metrics.
+ labels: a nested structure of tensors contains labels.
+ model_outputs: a list of output tensors. Assume the order is aligned with
+ the label_names list.
+ """
+ for i, label_name in enumerate(self._get_label_names()):
+ for metric in metrics:
+ if label_name in metric.name:
+ metric.update_state(labels[i], model_outputs[i])
+
+ if self._is_multihead():
+
+ def joint_accuracy_fn(y_true: List[tf.Tensor], y_pred: List[tf.Tensor]):
+ """Calculates the joint accuracy of predictions."""
+ hits = []
+ for label, pred in zip(y_true, y_pred):
+ label_id = tf.argmax(label, axis=-1)
+ pred_id = tf.argmax(pred, axis=-1)
+ hits.append(tf.equal(label_id, pred_id))
+ hits = tf.math.reduce_all(tf.stack(hits, axis=-1), axis=-1)
+ return tf.reduce_mean(tf.cast(hits, tf.float32))
+
+ for metric in metrics:
+ if 'joint' in metric.name:
+ values = joint_accuracy_fn(labels, model_outputs)
+ metric.update_state(values)
+
+ def build_losses(self,
+ labels: List[tf.Tensor],
+ model_outputs: List[tf.Tensor],
+ aux_losses: Optional[Any] = None):
+ """Builds losses."""
+
+ all_losses = 0
+ for model_output, label in zip(model_outputs, labels):
+ losses = super().build_losses(model_outputs=model_output,
+ labels=label,
+ aux_losses=aux_losses)
+ all_losses += losses[self.loss]
+ return all_losses
+
+ def train_step(self,
+ inputs: Mapping[str, Any],
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[Any]] = None):
+ """Does forward and backward pass.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ features = inputs['image']
+ labels = [inputs[k] for k in self._get_label_names()]
+
+ num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
+ with tf.GradientTape() as tape:
+ outputs = model(features, training=True)
+ # tf_keras.Model eliminates the list if the outputs list len is 1.
+ # Recover it here to be compatible with multihead settings.
+ outputs = [outputs] if isinstance(outputs, tf.Tensor) else outputs
+ # Casting output layer as float32 is necessary when mixed_precision is
+ # mixed_float16 or mixed_bfloat16 to ensure output is casted as float32.
+ outputs = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), outputs)
+
+ if self._is_multilabel():
+ outputs = tf.nest.map_structure(tf.math.sigmoid, outputs)
+ else:
+ outputs = tf.nest.map_structure(tf.math.softmax, outputs)
+
+ all_losses = self.build_losses(model_outputs=outputs,
+ labels=labels,
+ aux_losses=model.losses)
+
+ # Scale loss as the default gradients allreduce performs sum inside the
+ # optimizer.
+ scaled_loss = all_losses / num_replicas
+
+ # For mixed_precision policy, when LossScaleOptimizer is used, loss is
+ # scaled for numerical stability.
+ if isinstance(
+ optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ tvars = model.trainable_variables
+ grads = tape.gradient(scaled_loss, tvars)
+ # Scale back gradient before apply_gradients when LossScaleOptimizer is
+ # used.
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ grads = optimizer.get_unscaled_gradients(grads)
+ optimizer.apply_gradients(list(zip(grads, tvars)))
+
+ logs = {self.loss: all_losses}
+ if metrics:
+ self.process_metrics(metrics, labels, outputs)
+ logs.update({m.name: m.result() for m in metrics})
+ return logs
+
+ def validation_step(self,
+ inputs: Mapping[str, tf.Tensor],
+ model: tf_keras.Model,
+ metrics: Optional[List[Any]] = None):
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ features = inputs['image']
+ labels = [inputs[k] for k in self._get_label_names()]
+
+ input_partition_dims = self.task_config.eval_input_partition_dims
+ if input_partition_dims:
+ strategy = tf.distribute.get_strategy()
+ features['image'] = strategy.experimental_split_to_logical_devices(
+ features['image'], input_partition_dims)
+
+ outputs = self.inference_step(features, model)
+ # tf_keras.Model eliminates the list if the outputs list len is 1.
+ # Recover it here to be compatible with multihead settings.
+ outputs = [outputs] if isinstance(outputs, tf.Tensor) else outputs
+ # Casting output layer as float32 is necessary when mixed_precision is
+ # mixed_float16 or mixed_bfloat16 to ensure output is casted as float32.
+ outputs = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), outputs)
+ all_losses = self.build_losses(model_outputs=outputs, labels=labels,
+ aux_losses=model.losses)
+
+ logs = {self.loss: all_losses}
+ if metrics:
+ self.process_metrics(metrics, labels, outputs)
+ logs.update({m.name: m.result() for m in metrics})
+ return logs
diff --git a/official/projects/videoglue/tasks/spatiotemporal_action_localization.py b/official/projects/videoglue/tasks/spatiotemporal_action_localization.py
new file mode 100644
index 00000000000..7729e3fa919
--- /dev/null
+++ b/official/projects/videoglue/tasks/spatiotemporal_action_localization.py
@@ -0,0 +1,284 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Spatiotemporal action localization task."""
+from typing import Any, List, Optional, Mapping
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.core import task_factory
+from official.projects.videoglue.configs import spatiotemporal_action_localization as exp_cfg
+from official.projects.videoglue.datasets import dataset_factory
+from official.projects.videoglue.evaluation import spatiotemporal_action_localization_evaluator as eval_util
+from official.projects.videoglue.tools import checkpoint_loader
+from official.vision.modeling import factory_3d
+from official.vision.tasks import video_classification
+
+
+@task_factory.register_task_cls(exp_cfg.SpatiotemporalActionLocalizationTask)
+class SpatiotemporalActionLocalizationTask(
+ video_classification.VideoClassificationTask):
+ """Spatiotemporal action localization task."""
+
+ def _is_multilabel(self):
+ """Whether the dataset/task has multi-labels."""
+ return True
+
+ def build_model(self) -> tf_keras.Model:
+ """Builds video model."""
+ common_input_shape = [
+ d1 if d1 == d2 else None
+ for d1, d2 in zip(self.task_config.train_data.feature_shape,
+ self.task_config.validation_data.feature_shape)
+ ]
+
+ num_instances = self.task_config.train_data.num_instances
+ input_specs_dict = {
+ 'image':
+ tf_keras.layers.InputSpec(shape=[None] + common_input_shape),
+ 'instances_position':
+ tf_keras.layers.InputSpec(shape=[None, num_instances, 4]),
+ }
+
+ l2_weight_decay = self.task_config.losses.l2_weight_decay
+ # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
+ # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
+ # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
+ l2_regularizer = (tf_keras.regularizers.l2(
+ l2_weight_decay / 2.0) if l2_weight_decay else None)
+
+ model = factory_3d.build_model(
+ self.task_config.model.model_type,
+ input_specs=input_specs_dict,
+ model_config=self.task_config.model,
+ num_classes=self.task_config.train_data.num_classes,
+ l2_regularizer=l2_regularizer)
+
+ if self.task_config.freeze_backbone:
+ logging.info('Freezing model backbone.')
+ model.backbone.trainable = False
+ return model
+
+ def build_inputs(
+ self, params: exp_cfg.DataConfig, input_context: Any = None
+ ) -> tf.data.Dataset:
+ """Builds classification input."""
+ augmentation_type = params.data_augmentation.type
+ augmentation_params = params.data_augmentation.get().as_dict()
+ dataset_config = {
+ 'is_training': params.is_training,
+ 'num_frames': params.feature_shape[0],
+ 'temporal_stride': params.temporal_stride,
+ 'num_instance_per_frame': params.num_instances,
+ 'min_resize': params.min_resize,
+ 'crop_size': params.feature_shape[1],
+ 'zero_centering_image': params.zero_centering_image,
+ 'color_augmentation': params.color_augmentation,
+ 'num_test_clips': params.num_test_clips,
+ 'augmentation_type': augmentation_type,
+ 'augmentation_params': augmentation_params,
+ 'one_hot_label': params.one_hot_label,
+ 'merge_multi_labels': params.merge_multi_labels,
+ 'import_detected_bboxes': params.import_detected_bboxes,
+ }
+ data_loader = dataset_factory.DataLoader(
+ params=params, dataset_config=dataset_config)
+ return data_loader(input_context=input_context)
+
+ def initialize(self, model: tf_keras.Model):
+ """Loads pretrained checkpoint."""
+ if not self.task_config.init_checkpoint:
+ return
+
+ checkpoint_loader.get_checkpoint_loader(
+ model=model,
+ init_checkpoint=self.task_config.init_checkpoint,
+ init_checkpoint_type=self.task_config.init_checkpoint_modules)
+
+ def build_losses(self,
+ labels: Mapping[str, tf.Tensor],
+ model_outputs: tf.Tensor,
+ aux_losses: Optional[Any] = None):
+ """Sparse categorical cross entropy loss.
+
+ Args:
+ labels: labels dictionary contains multi-hot "class_target" and
+ "sample_weights".
+ model_outputs: Output logits of the classifier.
+ aux_losses: auxiliarly loss tensors, i.e. `losses` in keras.Model.
+
+ Returns:
+ The dictionary contains total loss tensors.
+ """
+ losses_config = self.task_config.losses
+
+ # in shape [B, N]
+ xent_loss_fn = tf_keras.losses.BinaryCrossentropy(
+ reduction=tf_keras.losses.Reduction.NONE,
+ from_logits=False,
+ label_smoothing=losses_config.label_smoothing)
+ class_targets = labels['class_target']
+ sample_weight = tf.cast(labels['instances_mask'], tf.float32)
+ model_loss = xent_loss_fn(y_true=class_targets,
+ y_pred=model_outputs,
+ sample_weight=sample_weight)
+ if self._is_multilabel():
+ # Re-scale the binary cross entropy loss with num_classes to make it
+ # comparable with categorical cross entropy in scale. This enables to use
+ # the same learning rate for different losses.
+ model_loss *= self._get_num_classes()
+
+ model_loss = tf.reduce_sum(model_loss) / (
+ tf.reduce_sum(sample_weight) + 1e-6)
+
+ if aux_losses:
+ regularization_loss = tf.add_n(aux_losses)
+ else:
+ regularization_loss = 0.0
+
+ total_loss = model_loss + regularization_loss
+ logs = {
+ 'model_loss': model_loss,
+ 'regularization_loss': regularization_loss,
+ # Include model predictions and corresponding boxes for the eval.
+ 'predictions': model_outputs,
+ # Register total loss.
+ self.loss: total_loss,
+ }
+ return logs
+
+ def build_metrics(self, training: bool = True):
+ """Gets streaming metrics for training/validation."""
+ metrics = [
+ tf_keras.metrics.AUC(curve='PR', multi_label=True, name='AUPR'),
+ tf_keras.metrics.AUC(curve='ROC', multi_label=True, name='AUROC'),
+ ]
+ self.evaluator = eval_util.SpatiotemporalActionLocalizationEvaluator()
+ return metrics
+
+ def process_metrics(self, metrics: List[Any],
+ labels: Mapping[str, tf.Tensor],
+ model_outputs: tf.Tensor):
+ """Processes and updates metrics.
+
+ Called when using custom training loop API.
+
+ Args:
+ metrics: a nested structure of metrics objects. The return of function
+ self.build_metrics.
+ labels: a tensor or a nested structure of tensors.
+ model_outputs: a tensor or a nested structure of tensors. For example,
+ output of the keras model built by self.build_model.
+ """
+ num_classes = self.task_config.train_data.num_classes
+ class_targets = tf.reshape(labels['class_target'], [-1, num_classes])
+ model_outputs = tf.reshape(model_outputs, [-1, num_classes])
+ sample_weight = tf.cast(tf.reshape(labels['instances_mask'], [-1]),
+ tf.float32)
+ for metric in metrics:
+ metric.update_state(
+ y_true=class_targets,
+ y_pred=model_outputs,
+ sample_weight=sample_weight)
+
+ def train_step(self,
+ inputs: Mapping[str, tf.Tensor],
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[Any]] = None):
+ """Does one forward and backward pass.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ features = {
+ 'image': inputs['image'],
+ 'instances_position': inputs['instances_position'],
+ }
+ labels = {
+ 'class_target': inputs['label'],
+ 'instances_mask': inputs['instances_mask'],
+ }
+ return super().train_step(
+ inputs=(features, labels),
+ model=model,
+ optimizer=optimizer,
+ metrics=metrics)
+
+ def validation_step(self,
+ inputs: Mapping[str, tf.Tensor],
+ model: tf_keras.Model,
+ metrics: Optional[List[Any]] = None):
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ if self.task_config.validation_data.import_detected_bboxes:
+ instances_position = inputs['detected_instances_position']
+ instances_mask = inputs['detected_instances_mask']
+ instances_score = inputs['detected_instances_score']
+ else:
+ instances_position = inputs['instances_position']
+ instances_mask = inputs['instances_mask']
+ instances_score = inputs['instances_score']
+
+ features = {
+ 'image': inputs['image'],
+ 'instances_position': instances_position,
+ }
+ labels = {
+ 'class_target': inputs['label'],
+ 'instances_mask': instances_mask,
+ }
+ logs = super().validation_step(
+ inputs=(features, labels), model=model, metrics=metrics)
+
+ # Final action_prob should be multiplied by the box score.
+ predictions = logs['predictions'] * instances_score[..., None]
+
+ # Add predictions and labels for the next step eval_end().
+ logs.update({
+ 'predictions': predictions,
+ 'instances_position': instances_position,
+ 'nonmerge_label': inputs['nonmerge_label'],
+ 'nonmerge_instances_position': inputs['nonmerge_instances_position'],
+ })
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ """Aggregates logs."""
+ if state is None:
+ self.evaluator.reset_states()
+ # Create an arbitrary state to indicate it's not the first step in the
+ # following calls to this function.
+ state = True
+ self.evaluator.update_state(step_outputs)
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ """Reduces aggregated logs."""
+ return self.evaluator.result()
diff --git a/official/projects/videoglue/tools/checkpoint_loader.py b/official/projects/videoglue/tools/checkpoint_loader.py
new file mode 100644
index 00000000000..4fda1e50af6
--- /dev/null
+++ b/official/projects/videoglue/tools/checkpoint_loader.py
@@ -0,0 +1,217 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Customized checkpoint loader."""
+import re
+from typing import List, Tuple
+
+from absl import logging
+import numpy as np
+import tensorflow as tf, tf_keras
+
+
+# pylint:disable=line-too-long
+_VMAE_CKPT_MAPPING = [
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/query/kernel:0',
+ r'blocks.\1.attn.q.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/query/bias:0',
+ r'blocks.\1.attn.q.bias'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/value/kernel:0',
+ r'blocks.\1.attn.v.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/value/bias:0',
+ r'blocks.\1.attn.v.bias'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/key/kernel:0',
+ r'blocks.\1.attn.k.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/key/bias:0',
+ r'blocks.\1.attn.k.bias'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/attention_output/kernel:0',
+ r'blocks.\1.attn.proj.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention/attention_output/bias:0',
+ r'blocks.\1.attn.proj.bias'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention_layer_norm/gamma:0',
+ r'blocks.\1.norm1.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/self_attention_layer_norm/beta:0',
+ r'blocks.\1.norm1.bias'),
+ (r'encoder/transformer_encoder_block_(.*?)/intermediate/kernel:0',
+ r'blocks.\1.mlp.fc1.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/intermediate/bias:0',
+ r'blocks.\1.mlp.fc1.bias'),
+ (r'encoder/transformer_encoder_block_(.*?)/output/kernel:0',
+ r'blocks.\1.mlp.fc2.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/output/bias:0',
+ r'blocks.\1.mlp.fc2.bias'),
+ (r'encoder/transformer_encoder_block_(.*?)/output_layer_norm/gamma:0',
+ r'blocks.\1.norm2.weight'),
+ (r'encoder/transformer_encoder_block_(.*?)/output_layer_norm/beta:0',
+ r'blocks.\1.norm2.bias'),
+
+ # ======= final layer norm
+ (r'encoder/layer_normalization/gamma:0', r'norm.weight'),
+ (r'encoder/layer_normalization/beta:0', r'norm.bias'),
+
+ # ======= input projection layer
+ (r'conv3d/kernel:0', r'patch_embed.proj.weight'),
+ (r'conv3d/bias:0', r'patch_embed.proj.bias'),
+
+ # ======= agg embedding.
+ (r'token_layer/cls:0', r'cls_token'),
+
+ # ======= positional embedding.
+ (r'add_separable_position_embs/pos_embedding_time:0',
+ r'pos_embed_temporal'),
+ (r'add_separable_position_embs/pos_embedding_space:0',
+ r'pos_embed_spatial'),
+]
+# pylint:enable=line-too-long
+
+
+class CheckpointLoaderBase(object):
+ """Checkpoint loader object."""
+
+ def __init__(self, model: tf_keras.Model,
+ init_checkpoint: str,
+ init_checkpoint_type: str):
+ self._init_checkpoint = init_checkpoint
+ self._init_checkpoint_type = init_checkpoint_type
+
+ ckpt_dir_or_file = self._init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+
+ self._load_checkpoint(model, ckpt_dir_or_file)
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
+
+ def _load_checkpoint(self, model: tf_keras.Model, ckpt_dir_or_file: str):
+ """Loads checkpoint."""
+ if self._init_checkpoint_type == 'all':
+ ckpt = tf.train.Checkpoint(model=model)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ elif self._init_checkpoint_type == 'backbone':
+ ckpt = tf.train.Checkpoint(backbone=model.backbone)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ else:
+ raise ValueError(
+ 'Unrecognized init_checkpoint_type: %s' % self._init_checkpoint_type)
+
+ def _remap_variable_name(self,
+ variable_name: str,
+ name_mapping: List[Tuple[str, str]]):
+ """Remaps variable name given the mapping."""
+ for source, dest in name_mapping:
+ variable_name = re.sub(source, dest, variable_name)
+ return variable_name
+
+
+class CheckpointLoaderVMAE(CheckpointLoaderBase):
+ """Checkpoint loader for Video MAE."""
+
+ def _maybe_transpose_pytorch_weight(self, ckpt_weight):
+ """Transposes pytorch weight to macth with the Tensorflow convention."""
+ if len(ckpt_weight.shape) == 2:
+ # fc kernel
+ ckpt_weight = np.transpose(ckpt_weight, [1, 0])
+ elif len(ckpt_weight.shape) == 4:
+ # conv2d kernel
+ ckpt_weight = np.transpose(ckpt_weight, [2, 3, 1, 0])
+ elif len(ckpt_weight.shape) == 5:
+ # conv3d kernel
+ ckpt_weight = np.transpose(ckpt_weight, [2, 3, 4, 1, 0])
+ return ckpt_weight
+
+ def _customized_vmae_initialize(self,
+ model: tf_keras.Model,
+ ckpt_dir_or_file: str):
+ """Loads pretrained Video MAE checkpoint."""
+ with tf.io.gfile.GFile(ckpt_dir_or_file, 'rb') as ckpt:
+ weights = np.load(ckpt, allow_pickle=True)
+
+ ckpt_names = list(weights[()].keys())
+ ckpt_names = [n for n in ckpt_names if 'pred_head' not in n]
+
+ skipped = []
+ loaded = []
+ for krs_w in model.weights:
+ krs_name = krs_w.name
+ # Handle the first block naming.
+ krs_name = krs_name.replace('encoder/transformer_encoder_block/',
+ 'encoder/transformer_encoder_block_0/')
+ ckpt_name = self._remap_variable_name(krs_name, _VMAE_CKPT_MAPPING)
+ if ckpt_name in ckpt_names:
+ ckpt_weight = weights[()][ckpt_name]
+ ckpt_weight = self._maybe_transpose_pytorch_weight(ckpt_weight)
+
+ if ckpt_weight.shape == krs_w.shape:
+ krs_w.assign(ckpt_weight)
+ loaded.append(ckpt_name)
+ elif 'kernel' in krs_name and any(
+ [keyword in krs_name for keyword in ['key', 'query', 'value']]):
+ cin, cout = ckpt_weight.shape
+ num_heads = krs_w.shape[1]
+ ckpt_weight = tf.reshape(
+ ckpt_weight, [cin, num_heads, cout // num_heads])
+ krs_w.assign(ckpt_weight)
+ loaded.append(ckpt_name)
+ elif 'bias' in krs_name and any(
+ [keyword in krs_name for keyword in ['key', 'query', 'value']]):
+ cout = ckpt_weight.shape[0]
+ num_heads = krs_w.shape[0]
+ ckpt_weight = tf.reshape(ckpt_weight, [num_heads, cout // num_heads])
+ krs_w.assign(ckpt_weight)
+ loaded.append(ckpt_name)
+ elif 'kernel' in krs_name and 'attention_output' in krs_name:
+ cin, cout = ckpt_weight.shape
+ num_heads = krs_w.shape[0]
+ ckpt_weight = tf.reshape(ckpt_weight,
+ [num_heads, cin // num_heads, cout])
+ krs_w.assign(ckpt_weight)
+ loaded.append(ckpt_name)
+ else:
+ skipped.append(krs_name)
+ else:
+ skipped.append(krs_name)
+
+ leftover = set(ckpt_names) - set(loaded)
+ logging.info('skipped: %s', skipped)
+ logging.info('leftover: %s', leftover)
+
+ if any([('encoder' in v or 'conv3d' in v or 'pos_embedding' in v)
+ for v in skipped]):
+ raise ValueError('ViT backbone is only partially loaded.')
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
+
+ def _load_checkpoint(self, model: tf_keras.Model, ckpt_dir_or_file: str):
+ """Loads checkpoint."""
+ self._customized_vmae_initialize(
+ model=model, ckpt_dir_or_file=ckpt_dir_or_file)
+
+
+def get_checkpoint_loader(
+ model: tf_keras.Model, init_checkpoint: str, init_checkpoint_type: str):
+ """Gets the corresponding checkpoint loader."""
+
+ if init_checkpoint_type == 'customized_vmae':
+ return CheckpointLoaderVMAE(
+ model=model,
+ init_checkpoint=init_checkpoint,
+ init_checkpoint_type=init_checkpoint_type)
+
+ else:
+ return CheckpointLoaderBase(
+ model=model,
+ init_checkpoint=init_checkpoint,
+ init_checkpoint_type=init_checkpoint_type)
diff --git a/official/projects/videoglue/train.py b/official/projects/videoglue/train.py
new file mode 100644
index 00000000000..6c5ca4c1658
--- /dev/null
+++ b/official/projects/videoglue/train.py
@@ -0,0 +1,80 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Training driver."""
+
+from absl import app
+from absl import flags
+import gin
+
+# pylint: disable=unused-import
+from official.common import distribute_utils
+from official.common import flags as tfm_flags
+from official.core import task_factory
+from official.core import train_lib
+from official.core import train_utils
+from official.modeling import performance
+from official.projects.videoglue.modeling import video_action_transformer_model
+from official.projects.videoglue.modeling import video_classification_model
+from official.projects.videoglue.modeling.backbones import vit_3d
+from official.projects.videoglue.tasks import multihead_video_classification
+from official.projects.videoglue.tasks import spatiotemporal_action_localization
+from official.vision import registry_imports
+# pylint: enable=unused-import
+
+
+FLAGS = flags.FLAGS
+
+
+def main(_):
+ gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
+ params = train_utils.parse_configuration(FLAGS)
+ model_dir = FLAGS.model_dir
+ if 'train' in FLAGS.mode:
+ # Pure eval modes do not output yaml files. Otherwise continuous eval job
+ # may race against the train job for writing the same file.
+ train_utils.serialize_config(params, model_dir)
+
+ if 'train_and_eval' in FLAGS.mode:
+ assert (params.task.train_data.feature_shape ==
+ params.task.validation_data.feature_shape), (
+ f'train {params.task.train_data.feature_shape} != validate '
+ f'{params.task.validation_data.feature_shape}')
+
+ # Set mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
+ # can have significant impact on model speeds by utilizing float16 in case of
+ # GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
+ # dtype is float16
+ if params.runtime.mixed_precision_dtype:
+ performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
+ distribution_strategy = distribute_utils.get_distribution_strategy(
+ distribution_strategy=params.runtime.distribution_strategy,
+ all_reduce_alg=params.runtime.all_reduce_alg,
+ num_gpus=params.runtime.num_gpus,
+ tpu_address=params.runtime.tpu)
+ with distribution_strategy.scope():
+ task = task_factory.get_task(params.task, logging_dir=model_dir)
+
+ train_lib.run_experiment(
+ distribution_strategy=distribution_strategy,
+ task=task,
+ mode=FLAGS.mode,
+ params=params,
+ model_dir=model_dir)
+
+ train_utils.save_gin_config(FLAGS.mode, model_dir)
+
+if __name__ == '__main__':
+ tfm_flags.define_flags()
+ app.run(main)
diff --git a/official/projects/vit/README.md b/official/projects/vit/README.md
deleted file mode 100644
index c42f381b0fd..00000000000
--- a/official/projects/vit/README.md
+++ /dev/null
@@ -1,14 +0,0 @@
-# Vision Transformer (ViT) and Data-Efficient Image Transformer (DEIT)
-
-**DISCLAIMER**: This implementation is still under development. No support will
-be provided during the development phase.
-
-- [](https://arxiv.org/abs/2010.11929)
-- [](https://arxiv.org/abs/2012.12877)
-
-This repository is the implementations of Vision Transformer (ViT) and
-Data-Efficient Image Transformer (DEIT) in TensorFlow 2.
-
-* Paper title:
-- [An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale](https://arxiv.org/pdf/2010.11929.pdf).
-- [Training data-efficient image transformers & distillation through attention](https://arxiv.org/pdf/2012.12877.pdf).
diff --git a/official/projects/vit/configs/backbones.py b/official/projects/vit/configs/backbones.py
deleted file mode 100644
index 655c05b972c..00000000000
--- a/official/projects/vit/configs/backbones.py
+++ /dev/null
@@ -1,57 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Backbones configurations."""
-from typing import Optional
-
-import dataclasses
-
-from official.modeling import hyperparams
-
-
-@dataclasses.dataclass
-class Transformer(hyperparams.Config):
- """Transformer config."""
- mlp_dim: int = 1
- num_heads: int = 1
- num_layers: int = 1
- attention_dropout_rate: float = 0.0
- dropout_rate: float = 0.1
-
-
-@dataclasses.dataclass
-class VisionTransformer(hyperparams.Config):
- """VisionTransformer config."""
- model_name: str = 'vit-b16'
- # pylint: disable=line-too-long
- classifier: str = 'token' # 'token' or 'gap'. If set to 'token', an extra classification token is added to sequence.
- # pylint: enable=line-too-long
- representation_size: int = 0
- hidden_size: int = 1
- patch_size: int = 16
- transformer: Transformer = Transformer()
- init_stochastic_depth_rate: float = 0.0
- original_init: bool = True
-
-
-@dataclasses.dataclass
-class Backbone(hyperparams.OneOfConfig):
- """Configuration for backbones.
-
- Attributes:
- type: 'str', type of backbone be used, one the of fields below.
- vit: vit backbone config.
- """
- type: Optional[str] = None
- vit: VisionTransformer = VisionTransformer()
diff --git a/official/projects/vit/configs/image_classification.py b/official/projects/vit/configs/image_classification.py
deleted file mode 100644
index d312719a301..00000000000
--- a/official/projects/vit/configs/image_classification.py
+++ /dev/null
@@ -1,273 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Image classification configuration definition."""
-import dataclasses
-import os
-from typing import Optional
-
-from official.core import config_definitions as cfg
-from official.core import exp_factory
-from official.core import task_factory
-from official.modeling import hyperparams
-from official.modeling import optimization
-from official.projects.vit.configs import backbones
-from official.vision.configs import common
-from official.vision.configs import image_classification as img_cls_cfg
-from official.vision.tasks import image_classification
-
-# pytype: disable=wrong-keyword-args
-
-DataConfig = img_cls_cfg.DataConfig
-
-
-@dataclasses.dataclass
-class ImageClassificationModel(img_cls_cfg.ImageClassificationModel):
- """The model config."""
- backbone: backbones.Backbone = backbones.Backbone(
- type='vit', vit=backbones.VisionTransformer())
-
-
-@dataclasses.dataclass
-class Losses(hyperparams.Config):
- loss_weight: float = 1.0
- one_hot: bool = True
- label_smoothing: float = 0.0
- l2_weight_decay: float = 0.0
- soft_labels: bool = False
-
-
-@dataclasses.dataclass
-class Evaluation(hyperparams.Config):
- top_k: int = 5
-
-
-@dataclasses.dataclass
-class ImageClassificationTask(cfg.TaskConfig):
- """The task config. Same as the classification task for convnets."""
- model: ImageClassificationModel = ImageClassificationModel()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- losses: Losses = Losses()
- evaluation: Evaluation = Evaluation()
- init_checkpoint: Optional[str] = None
- init_checkpoint_modules: str = 'all' # all or backbone
-
-
-IMAGENET_TRAIN_EXAMPLES = 1281167
-IMAGENET_VAL_EXAMPLES = 50000
-IMAGENET_INPUT_PATH_BASE = 'imagenet-2012-tfrecord'
-
-# TODO(b/177942984): integrate the experiments to TF-vision.
-task_factory.register_task_cls(ImageClassificationTask)(
- image_classification.ImageClassificationTask)
-
-
-@exp_factory.register_config_factory('deit_imagenet_pretrain')
-def image_classification_imagenet_deit_pretrain() -> cfg.ExperimentConfig:
- """Image classification on imagenet with vision transformer."""
- train_batch_size = 4096 # originally was 1024 but 4096 better for tpu v3-32
- eval_batch_size = 4096 # originally was 1024 but 4096 better for tpu v3-32
- num_classes = 1001
- label_smoothing = 0.1
- steps_per_epoch = IMAGENET_TRAIN_EXAMPLES // train_batch_size
- config = cfg.ExperimentConfig(
- task=ImageClassificationTask(
- model=ImageClassificationModel(
- num_classes=num_classes,
- input_size=[224, 224, 3],
- kernel_initializer='zeros',
- backbone=backbones.Backbone(
- type='vit',
- vit=backbones.VisionTransformer(
- model_name='vit-b16',
- representation_size=768,
- init_stochastic_depth_rate=0.1,
- original_init=False,
- transformer=backbones.Transformer(
- dropout_rate=0.0, attention_dropout_rate=0.0)))),
- losses=Losses(
- l2_weight_decay=0.0,
- label_smoothing=label_smoothing,
- one_hot=False,
- soft_labels=True),
- train_data=DataConfig(
- input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'train*'),
- is_training=True,
- global_batch_size=train_batch_size,
- aug_type=common.Augmentation(
- type='randaug',
- randaug=common.RandAugment(
- magnitude=9, exclude_ops=['Cutout'])),
- mixup_and_cutmix=common.MixupAndCutmix(
- label_smoothing=label_smoothing)),
- validation_data=DataConfig(
- input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'valid*'),
- is_training=False,
- global_batch_size=eval_batch_size)),
- trainer=cfg.TrainerConfig(
- steps_per_loop=steps_per_epoch,
- summary_interval=steps_per_epoch,
- checkpoint_interval=steps_per_epoch,
- train_steps=300 * steps_per_epoch,
- validation_steps=IMAGENET_VAL_EXAMPLES // eval_batch_size,
- validation_interval=steps_per_epoch,
- optimizer_config=optimization.OptimizationConfig({
- 'optimizer': {
- 'type': 'adamw',
- 'adamw': {
- 'weight_decay_rate': 0.05,
- 'include_in_weight_decay': r'.*(kernel|weight):0$',
- 'gradient_clip_norm': 0.0
- }
- },
- 'learning_rate': {
- 'type': 'cosine',
- 'cosine': {
- 'initial_learning_rate': 0.0005 * train_batch_size / 512,
- 'decay_steps': 300 * steps_per_epoch,
- }
- },
- 'warmup': {
- 'type': 'linear',
- 'linear': {
- 'warmup_steps': 5 * steps_per_epoch,
- 'warmup_learning_rate': 0
- }
- }
- })),
- restrictions=[
- 'task.train_data.is_training != None',
- 'task.validation_data.is_training != None'
- ])
-
- return config
-
-
-@exp_factory.register_config_factory('vit_imagenet_pretrain')
-def image_classification_imagenet_vit_pretrain() -> cfg.ExperimentConfig:
- """Image classification on imagenet with vision transformer."""
- train_batch_size = 4096
- eval_batch_size = 4096
- steps_per_epoch = IMAGENET_TRAIN_EXAMPLES // train_batch_size
- config = cfg.ExperimentConfig(
- task=ImageClassificationTask(
- model=ImageClassificationModel(
- num_classes=1001,
- input_size=[224, 224, 3],
- kernel_initializer='zeros',
- backbone=backbones.Backbone(
- type='vit',
- vit=backbones.VisionTransformer(
- model_name='vit-b16', representation_size=768))),
- losses=Losses(l2_weight_decay=0.0),
- train_data=DataConfig(
- input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'train*'),
- is_training=True,
- global_batch_size=train_batch_size),
- validation_data=DataConfig(
- input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'valid*'),
- is_training=False,
- global_batch_size=eval_batch_size)),
- trainer=cfg.TrainerConfig(
- steps_per_loop=steps_per_epoch,
- summary_interval=steps_per_epoch,
- checkpoint_interval=steps_per_epoch,
- train_steps=300 * steps_per_epoch,
- validation_steps=IMAGENET_VAL_EXAMPLES // eval_batch_size,
- validation_interval=steps_per_epoch,
- optimizer_config=optimization.OptimizationConfig({
- 'optimizer': {
- 'type': 'adamw',
- 'adamw': {
- 'weight_decay_rate': 0.3,
- 'include_in_weight_decay': r'.*(kernel|weight):0$',
- 'gradient_clip_norm': 0.0
- }
- },
- 'learning_rate': {
- 'type': 'cosine',
- 'cosine': {
- 'initial_learning_rate': 0.003 * train_batch_size / 4096,
- 'decay_steps': 300 * steps_per_epoch,
- }
- },
- 'warmup': {
- 'type': 'linear',
- 'linear': {
- 'warmup_steps': 10000,
- 'warmup_learning_rate': 0
- }
- }
- })),
- restrictions=[
- 'task.train_data.is_training != None',
- 'task.validation_data.is_training != None'
- ])
-
- return config
-
-
-@exp_factory.register_config_factory('vit_imagenet_finetune')
-def image_classification_imagenet_vit_finetune() -> cfg.ExperimentConfig:
- """Image classification on imagenet with vision transformer."""
- train_batch_size = 512
- eval_batch_size = 512
- steps_per_epoch = IMAGENET_TRAIN_EXAMPLES // train_batch_size
- config = cfg.ExperimentConfig(
- task=ImageClassificationTask(
- model=ImageClassificationModel(
- num_classes=1001,
- input_size=[384, 384, 3],
- backbone=backbones.Backbone(
- type='vit',
- vit=backbones.VisionTransformer(model_name='vit-b16'))),
- losses=Losses(l2_weight_decay=0.0),
- train_data=DataConfig(
- input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'train*'),
- is_training=True,
- global_batch_size=train_batch_size),
- validation_data=DataConfig(
- input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'valid*'),
- is_training=False,
- global_batch_size=eval_batch_size)),
- trainer=cfg.TrainerConfig(
- steps_per_loop=steps_per_epoch,
- summary_interval=steps_per_epoch,
- checkpoint_interval=steps_per_epoch,
- train_steps=20000,
- validation_steps=IMAGENET_VAL_EXAMPLES // eval_batch_size,
- validation_interval=steps_per_epoch,
- optimizer_config=optimization.OptimizationConfig({
- 'optimizer': {
- 'type': 'sgd',
- 'sgd': {
- 'momentum': 0.9,
- 'global_clipnorm': 1.0,
- }
- },
- 'learning_rate': {
- 'type': 'cosine',
- 'cosine': {
- 'initial_learning_rate': 0.003,
- 'decay_steps': 20000,
- }
- }
- })),
- restrictions=[
- 'task.train_data.is_training != None',
- 'task.validation_data.is_training != None'
- ])
-
- return config
diff --git a/official/projects/vit/modeling/nn_blocks.py b/official/projects/vit/modeling/nn_blocks.py
deleted file mode 100644
index 891c9ac2426..00000000000
--- a/official/projects/vit/modeling/nn_blocks.py
+++ /dev/null
@@ -1,119 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Keras-based TransformerEncoder block layer."""
-import tensorflow as tf
-
-from official.nlp import modeling
-from official.vision.modeling.layers.nn_layers import StochasticDepth
-
-
-class TransformerEncoderBlock(modeling.layers.TransformerEncoderBlock):
- """TransformerEncoderBlock layer with stochastic depth."""
-
- def __init__(self,
- *args,
- stochastic_depth_drop_rate=0.0,
- return_attention=False,
- **kwargs):
- """Initializes TransformerEncoderBlock."""
- super().__init__(*args, **kwargs)
- self._stochastic_depth_drop_rate = stochastic_depth_drop_rate
- self._return_attention = return_attention
-
- def build(self, input_shape):
- if self._stochastic_depth_drop_rate:
- self._stochastic_depth = StochasticDepth(self._stochastic_depth_drop_rate)
- else:
- self._stochastic_depth = lambda x, *args, **kwargs: tf.identity(x)
-
- super().build(input_shape)
-
- def get_config(self):
- config = {"stochastic_depth_drop_rate": self._stochastic_depth_drop_rate}
- base_config = super().get_config()
- return dict(list(base_config.items()) + list(config.items()))
-
- def call(self, inputs, training=None):
- """Transformer self-attention encoder block call."""
- if isinstance(inputs, (list, tuple)):
- if len(inputs) == 2:
- input_tensor, attention_mask = inputs
- key_value = None
- elif len(inputs) == 3:
- input_tensor, key_value, attention_mask = inputs
- else:
- raise ValueError("Unexpected inputs to %s with length at %d" %
- (self.__class__, len(inputs)))
- else:
- input_tensor, key_value, attention_mask = (inputs, None, None)
-
- if self._output_range:
- if self._norm_first:
- source_tensor = input_tensor[:, 0:self._output_range, :]
- input_tensor = self._attention_layer_norm(input_tensor)
- if key_value is not None:
- key_value = self._attention_layer_norm(key_value)
- target_tensor = input_tensor[:, 0:self._output_range, :]
- if attention_mask is not None:
- attention_mask = attention_mask[:, 0:self._output_range, :]
- else:
- if self._norm_first:
- source_tensor = input_tensor
- input_tensor = self._attention_layer_norm(input_tensor)
- if key_value is not None:
- key_value = self._attention_layer_norm(key_value)
- target_tensor = input_tensor
-
- if key_value is None:
- key_value = input_tensor
- attention_output, attention_scores = self._attention_layer(
- query=target_tensor, value=key_value, attention_mask=attention_mask,
- return_attention_scores=True)
- attention_output = self._attention_dropout(attention_output)
-
- if self._norm_first:
- attention_output = source_tensor + self._stochastic_depth(
- attention_output, training=training)
- else:
- attention_output = self._attention_layer_norm(
- target_tensor +
- self._stochastic_depth(attention_output, training=training))
-
- if self._norm_first:
- source_attention_output = attention_output
- attention_output = self._output_layer_norm(attention_output)
- inner_output = self._intermediate_dense(attention_output)
- inner_output = self._intermediate_activation_layer(inner_output)
- inner_output = self._inner_dropout_layer(inner_output)
- layer_output = self._output_dense(inner_output)
- layer_output = self._output_dropout(layer_output)
-
- if self._norm_first:
- if self._return_attention:
- return source_attention_output + self._stochastic_depth(
- layer_output, training=training), attention_scores
- else:
- return source_attention_output + self._stochastic_depth(
- layer_output, training=training)
-
- # During mixed precision training, layer norm output is always fp32 for now.
- # Casts fp32 for the subsequent add.
- layer_output = tf.cast(layer_output, tf.float32)
- if self._return_attention:
- return self._output_layer_norm(layer_output + self._stochastic_depth(
- attention_output, training=training)), attention_scores
- else:
- return self._output_layer_norm(layer_output + self._stochastic_depth(
- attention_output, training=training))
diff --git a/official/projects/vit/modeling/vit.py b/official/projects/vit/modeling/vit.py
deleted file mode 100644
index 084ba1ed053..00000000000
--- a/official/projects/vit/modeling/vit.py
+++ /dev/null
@@ -1,299 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""VisionTransformer models."""
-import tensorflow as tf
-
-from official.modeling import activations
-from official.projects.vit.modeling import nn_blocks
-from official.vision.modeling.backbones import factory
-from official.vision.modeling.layers import nn_layers
-
-
-layers = tf.keras.layers
-
-VIT_SPECS = {
- 'vit-ti16':
- dict(
- hidden_size=192,
- patch_size=16,
- transformer=dict(mlp_dim=768, num_heads=3, num_layers=12),
- ),
- 'vit-s16':
- dict(
- hidden_size=384,
- patch_size=16,
- transformer=dict(mlp_dim=1536, num_heads=6, num_layers=12),
- ),
- 'vit-b16':
- dict(
- hidden_size=768,
- patch_size=16,
- transformer=dict(mlp_dim=3072, num_heads=12, num_layers=12),
- ),
- 'vit-b32':
- dict(
- hidden_size=768,
- patch_size=32,
- transformer=dict(mlp_dim=3072, num_heads=12, num_layers=12),
- ),
- 'vit-l16':
- dict(
- hidden_size=1024,
- patch_size=16,
- transformer=dict(mlp_dim=4096, num_heads=16, num_layers=24),
- ),
- 'vit-l32':
- dict(
- hidden_size=1024,
- patch_size=32,
- transformer=dict(mlp_dim=4096, num_heads=16, num_layers=24),
- ),
- 'vit-h14':
- dict(
- hidden_size=1280,
- patch_size=14,
- transformer=dict(mlp_dim=5120, num_heads=16, num_layers=32),
- ),
- 'vit-g14':
- dict(
- hidden_size=1664,
- patch_size=14,
- transformer=dict(mlp_dim=8192, num_heads=16, num_layers=48),
- ),
-}
-
-
-class AddPositionEmbs(tf.keras.layers.Layer):
- """Adds (optionally learned) positional embeddings to the inputs."""
-
- def __init__(self, posemb_init=None, **kwargs):
- super().__init__(**kwargs)
- self.posemb_init = posemb_init
-
- def build(self, inputs_shape):
- pos_emb_shape = (1, inputs_shape[1], inputs_shape[2])
- self.pos_embedding = self.add_weight(
- 'pos_embedding', pos_emb_shape, initializer=self.posemb_init)
-
- def call(self, inputs, inputs_positions=None):
- # inputs.shape is (batch_size, seq_len, emb_dim).
- pos_embedding = tf.cast(self.pos_embedding, inputs.dtype)
-
- return inputs + pos_embedding
-
-
-class TokenLayer(tf.keras.layers.Layer):
- """A simple layer to wrap token parameters."""
-
- def build(self, inputs_shape):
- self.cls = self.add_weight(
- 'cls', (1, 1, inputs_shape[-1]), initializer='zeros')
-
- def call(self, inputs):
- cls = tf.cast(self.cls, inputs.dtype)
- cls = cls + tf.zeros_like(inputs[:, 0:1]) # A hacky way to tile.
- x = tf.concat([cls, inputs], axis=1)
- return x
-
-
-class Encoder(tf.keras.layers.Layer):
- """Transformer Encoder."""
-
- def __init__(self,
- num_layers,
- mlp_dim,
- num_heads,
- dropout_rate=0.1,
- attention_dropout_rate=0.1,
- kernel_regularizer=None,
- inputs_positions=None,
- init_stochastic_depth_rate=0.0,
- kernel_initializer='glorot_uniform',
- add_pos_embed=True,
- **kwargs):
- super().__init__(**kwargs)
- self._num_layers = num_layers
- self._mlp_dim = mlp_dim
- self._num_heads = num_heads
- self._dropout_rate = dropout_rate
- self._attention_dropout_rate = attention_dropout_rate
- self._kernel_regularizer = kernel_regularizer
- self._inputs_positions = inputs_positions
- self._init_stochastic_depth_rate = init_stochastic_depth_rate
- self._kernel_initializer = kernel_initializer
- self._add_pos_embed = add_pos_embed
-
- def build(self, input_shape):
- if self._add_pos_embed:
- self._pos_embed = AddPositionEmbs(
- posemb_init=tf.keras.initializers.RandomNormal(stddev=0.02),
- name='posembed_input')
- self._dropout = layers.Dropout(rate=self._dropout_rate)
-
- self._encoder_layers = []
- # Set layer norm epsilons to 1e-6 to be consistent with JAX implementation.
- # https://flax.readthedocs.io/en/latest/_autosummary/flax.deprecated.nn.LayerNorm.html
- for i in range(self._num_layers):
- encoder_layer = nn_blocks.TransformerEncoderBlock(
- inner_activation=activations.gelu,
- num_attention_heads=self._num_heads,
- inner_dim=self._mlp_dim,
- output_dropout=self._dropout_rate,
- attention_dropout=self._attention_dropout_rate,
- kernel_regularizer=self._kernel_regularizer,
- kernel_initializer=self._kernel_initializer,
- norm_first=True,
- stochastic_depth_drop_rate=nn_layers.get_stochastic_depth_rate(
- self._init_stochastic_depth_rate, i + 1, self._num_layers),
- norm_epsilon=1e-6)
- self._encoder_layers.append(encoder_layer)
- self._norm = layers.LayerNormalization(epsilon=1e-6)
- super().build(input_shape)
-
- def call(self, inputs, training=None):
- x = inputs
- if self._add_pos_embed:
- x = self._pos_embed(x, inputs_positions=self._inputs_positions)
- x = self._dropout(x, training=training)
-
- for encoder_layer in self._encoder_layers:
- x = encoder_layer(x, training=training)
- x = self._norm(x)
- return x
-
- def get_config(self):
- config = {
- 'num_layers': self._num_layers,
- 'mlp_dim': self._mlp_dim,
- 'num_heads': self._num_heads,
- 'dropout_rate': self._dropout_rate,
- 'attention_dropout_rate': self._attention_dropout_rate,
- 'kernel_regularizer': self._kernel_regularizer,
- 'inputs_positions': self._inputs_positions,
- 'init_stochastic_depth_rate': self._init_stochastic_depth_rate,
- 'kernel_initializer': self._kernel_initializer,
- 'add_pos_embed': self._add_pos_embed,
- }
- base_config = super().get_config()
- return base_config.update(config)
-
-
-class VisionTransformer(tf.keras.Model):
- """Class to build VisionTransformer family model."""
-
- def __init__(self,
- mlp_dim=3072,
- num_heads=12,
- num_layers=12,
- attention_dropout_rate=0.0,
- dropout_rate=0.1,
- init_stochastic_depth_rate=0.0,
- input_specs=layers.InputSpec(shape=[None, None, None, 3]),
- patch_size=16,
- hidden_size=768,
- representation_size=0,
- classifier='token',
- kernel_regularizer=None,
- original_init=True):
- """VisionTransformer initialization function."""
- inputs = tf.keras.Input(shape=input_specs.shape[1:])
-
- x = layers.Conv2D(
- filters=hidden_size,
- kernel_size=patch_size,
- strides=patch_size,
- padding='valid',
- kernel_regularizer=kernel_regularizer,
- kernel_initializer='lecun_normal' if original_init else 'he_uniform')(
- inputs)
- if tf.keras.backend.image_data_format() == 'channels_last':
- rows_axis, cols_axis = (1, 2)
- else:
- rows_axis, cols_axis = (2, 3)
- # The reshape below assumes the data_format is 'channels_last,' so
- # transpose to that. Once the data is flattened by the reshape, the
- # data_format is irrelevant, so no need to update
- # tf.keras.backend.image_data_format.
- x = tf.transpose(x, perm=[0, 2, 3, 1])
- seq_len = (input_specs.shape[rows_axis] // patch_size) * (
- input_specs.shape[cols_axis] // patch_size)
- x = tf.reshape(x, [-1, seq_len, hidden_size])
-
- # If we want to add a class token, add it here.
- if classifier == 'token':
- x = TokenLayer(name='cls')(x)
-
- x = Encoder(
- num_layers=num_layers,
- mlp_dim=mlp_dim,
- num_heads=num_heads,
- dropout_rate=dropout_rate,
- attention_dropout_rate=attention_dropout_rate,
- kernel_regularizer=kernel_regularizer,
- kernel_initializer='glorot_uniform' if original_init else dict(
- class_name='TruncatedNormal', config=dict(stddev=.02)),
- init_stochastic_depth_rate=init_stochastic_depth_rate)(
- x)
-
- if classifier == 'token':
- x = x[:, 0]
- elif classifier == 'gap':
- x = tf.reduce_mean(x, axis=1)
-
- if representation_size:
- x = tf.keras.layers.Dense(
- representation_size,
- kernel_regularizer=kernel_regularizer,
- name='pre_logits',
- kernel_initializer='lecun_normal' if original_init else 'he_uniform')(
- x)
- x = tf.nn.tanh(x)
- else:
- x = tf.identity(x, name='pre_logits')
- endpoints = {
- 'pre_logits':
- tf.reshape(x, [-1, 1, 1, representation_size or hidden_size])
- }
-
- super(VisionTransformer, self).__init__(inputs=inputs, outputs=endpoints)
-
-
-@factory.register_backbone_builder('vit')
-def build_vit(input_specs,
- backbone_config,
- norm_activation_config,
- l2_regularizer=None):
- """Build ViT model."""
- del norm_activation_config
- backbone_type = backbone_config.type
- backbone_cfg = backbone_config.get()
- assert backbone_type == 'vit', (f'Inconsistent backbone type '
- f'{backbone_type}')
- backbone_cfg.override(VIT_SPECS[backbone_cfg.model_name])
-
- return VisionTransformer(
- mlp_dim=backbone_cfg.transformer.mlp_dim,
- num_heads=backbone_cfg.transformer.num_heads,
- num_layers=backbone_cfg.transformer.num_layers,
- attention_dropout_rate=backbone_cfg.transformer.attention_dropout_rate,
- dropout_rate=backbone_cfg.transformer.dropout_rate,
- init_stochastic_depth_rate=backbone_cfg.init_stochastic_depth_rate,
- input_specs=input_specs,
- patch_size=backbone_cfg.patch_size,
- hidden_size=backbone_cfg.hidden_size,
- representation_size=backbone_cfg.representation_size,
- classifier=backbone_cfg.classifier,
- kernel_regularizer=l2_regularizer,
- original_init=backbone_cfg.original_init)
diff --git a/official/projects/vit/modeling/vit_test.py b/official/projects/vit/modeling/vit_test.py
deleted file mode 100644
index 7318846ed7b..00000000000
--- a/official/projects/vit/modeling/vit_test.py
+++ /dev/null
@@ -1,42 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Tests for VIT."""
-
-from absl.testing import parameterized
-import tensorflow as tf
-
-from official.projects.vit.modeling import vit
-
-
-class VisionTransformerTest(parameterized.TestCase, tf.test.TestCase):
-
- @parameterized.parameters(
- (224, 85798656),
- (256, 85844736),
- )
- def test_network_creation(self, input_size, params_count):
- """Test creation of VisionTransformer family models."""
- tf.keras.backend.set_image_data_format('channels_last')
- input_specs = tf.keras.layers.InputSpec(
- shape=[2, input_size, input_size, 3])
- network = vit.VisionTransformer(input_specs=input_specs)
-
- inputs = tf.keras.Input(shape=(input_size, input_size, 3), batch_size=1)
- _ = network(inputs)
- self.assertEqual(network.count_params(), params_count)
-
-
-if __name__ == '__main__':
- tf.test.main()
diff --git a/official/projects/volumetric_models/__init__.py b/official/projects/volumetric_models/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/configs/__init__.py b/official/projects/volumetric_models/configs/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/configs/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/configs/backbones.py b/official/projects/volumetric_models/configs/backbones.py
index faae1882e3a..10d0faf947c 100644
--- a/official/projects/volumetric_models/configs/backbones.py
+++ b/official/projects/volumetric_models/configs/backbones.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -38,4 +38,4 @@ class Backbone(hyperparams.OneOfConfig):
unet_3d: UNet3D backbone config.
"""
type: Optional[str] = None
- unet_3d: UNet3D = UNet3D()
+ unet_3d: UNet3D = dataclasses.field(default_factory=UNet3D)
diff --git a/official/projects/volumetric_models/configs/decoders.py b/official/projects/volumetric_models/configs/decoders.py
index 828eaa9898c..28637499271 100644
--- a/official/projects/volumetric_models/configs/decoders.py
+++ b/official/projects/volumetric_models/configs/decoders.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -45,5 +45,7 @@ class Decoder(hyperparams.OneOfConfig):
unet_3d_decoder: UNet3D decoder config.
"""
type: Optional[str] = None
- identity: Identity = Identity()
- unet_3d_decoder: UNet3DDecoder = UNet3DDecoder()
+ identity: Identity = dataclasses.field(default_factory=Identity)
+ unet_3d_decoder: UNet3DDecoder = dataclasses.field(
+ default_factory=UNet3DDecoder
+ )
diff --git a/official/projects/volumetric_models/configs/semantic_segmentation_3d.py b/official/projects/volumetric_models/configs/semantic_segmentation_3d.py
index 3f6987f43bc..85c30b9f99d 100644
--- a/official/projects/volumetric_models/configs/semantic_segmentation_3d.py
+++ b/official/projects/volumetric_models/configs/semantic_segmentation_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -65,12 +65,22 @@ class SemanticSegmentationModel3D(hyperparams.Config):
input_size: List[int] = dataclasses.field(default_factory=list)
min_level: int = 3
max_level: int = 6
- head: SegmentationHead3D = SegmentationHead3D()
- backbone: backbones.Backbone = backbones.Backbone(
- type='unet_3d', unet_3d=backbones.UNet3D())
- decoder: decoders.Decoder = decoders.Decoder(
- type='unet_3d_decoder', unet_3d_decoder=decoders.UNet3DDecoder())
- norm_activation: common.NormActivation = common.NormActivation()
+ head: SegmentationHead3D = dataclasses.field(
+ default_factory=SegmentationHead3D
+ )
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='unet_3d', unet_3d=backbones.UNet3D()
+ )
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder( # pylint: disable=g-long-lambda
+ type='unet_3d_decoder', unet_3d_decoder=decoders.UNet3DDecoder()
+ )
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
@dataclasses.dataclass
@@ -88,11 +98,17 @@ class Evaluation(hyperparams.Config):
@dataclasses.dataclass
class SemanticSegmentation3DTask(cfg.TaskConfig):
"""The model config."""
- model: SemanticSegmentationModel3D = SemanticSegmentationModel3D()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- losses: Losses = Losses()
- evaluation: Evaluation = Evaluation()
+ model: SemanticSegmentationModel3D = dataclasses.field(
+ default_factory=SemanticSegmentationModel3D
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ evaluation: Evaluation = dataclasses.field(default_factory=Evaluation)
train_input_partition_dims: List[int] = dataclasses.field(
default_factory=list)
eval_input_partition_dims: List[int] = dataclasses.field(default_factory=list)
diff --git a/official/projects/volumetric_models/configs/semantic_segmentation_3d_test.py b/official/projects/volumetric_models/configs/semantic_segmentation_3d_test.py
index e54b0f98f45..666800f7b7d 100644
--- a/official/projects/volumetric_models/configs/semantic_segmentation_3d_test.py
+++ b/official/projects/volumetric_models/configs/semantic_segmentation_3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import exp_factory
diff --git a/official/projects/volumetric_models/dataloaders/__init__.py b/official/projects/volumetric_models/dataloaders/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/dataloaders/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/dataloaders/segmentation_input_3d.py b/official/projects/volumetric_models/dataloaders/segmentation_input_3d.py
index 1d7ba4bfb86..b7d0c46c323 100644
--- a/official/projects/volumetric_models/dataloaders/segmentation_input_3d.py
+++ b/official/projects/volumetric_models/dataloaders/segmentation_input_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Data parser and processing for 3D segmentation datasets."""
from typing import Any, Dict, Sequence, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import decoder
from official.vision.dataloaders import parser
diff --git a/official/projects/volumetric_models/dataloaders/segmentation_input_3d_test.py b/official/projects/volumetric_models/dataloaders/segmentation_input_3d_test.py
index 7a71de141cd..d008226ee7c 100644
--- a/official/projects/volumetric_models/dataloaders/segmentation_input_3d_test.py
+++ b/official/projects/volumetric_models/dataloaders/segmentation_input_3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import os
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.dataloaders import segmentation_input_3d
from official.vision.dataloaders import tfexample_utils
diff --git a/official/projects/volumetric_models/evaluation/__init__.py b/official/projects/volumetric_models/evaluation/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/evaluation/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/evaluation/segmentation_metrics.py b/official/projects/volumetric_models/evaluation/segmentation_metrics.py
index fc265720f8b..89742e254fa 100644
--- a/official/projects/volumetric_models/evaluation/segmentation_metrics.py
+++ b/official/projects/volumetric_models/evaluation/segmentation_metrics.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,16 +15,16 @@
"""Metrics for segmentation."""
from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.losses import segmentation_losses
class DiceScore:
"""Dice score metric for semantic segmentation.
- This class follows the same function interface as tf.keras.metrics.Metric but
- does not derive from tf.keras.metrics.Metric or utilize its functions. The
- reason is a tf.keras.metrics.Metric object does not run well on CPU while
+ This class follows the same function interface as tf_keras.metrics.Metric but
+ does not derive from tf_keras.metrics.Metric or utilize its functions. The
+ reason is a tf_keras.metrics.Metric object does not run well on CPU while
created on GPU, when running with MirroredStrategy. The same interface allows
for minimal change to the upstream tasks.
diff --git a/official/projects/volumetric_models/evaluation/segmentation_metrics_test.py b/official/projects/volumetric_models/evaluation/segmentation_metrics_test.py
index 1eac720016f..4f11659bb67 100644
--- a/official/projects/volumetric_models/evaluation/segmentation_metrics_test.py
+++ b/official/projects/volumetric_models/evaluation/segmentation_metrics_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for segmentation_losses.py."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.evaluation import segmentation_metrics
diff --git a/official/projects/volumetric_models/losses/__init__.py b/official/projects/volumetric_models/losses/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/losses/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/losses/segmentation_losses.py b/official/projects/volumetric_models/losses/segmentation_losses.py
index 6d422a9243e..3d78ffdb3ee 100644
--- a/official/projects/volumetric_models/losses/segmentation_losses.py
+++ b/official/projects/volumetric_models/losses/segmentation_losses.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Losses used for segmentation models."""
from typing import Optional, Sequence
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class SegmentationLossDiceScore(object):
@@ -65,7 +65,7 @@ def __call__(self, logits: tf.Tensor, labels: tf.Tensor) -> tf.Tensor:
if labels.get_shape().ndims < 2 or logits.get_shape().ndims < 2:
raise ValueError('The labels and logits must be at least rank 2.')
- epsilon = tf.keras.backend.epsilon()
+ epsilon = tf_keras.backend.epsilon()
keep_label_axis = list(range(len(logits.shape) - 1))
keep_batch_axis = list(range(1, len(logits.shape)))
diff --git a/official/projects/volumetric_models/losses/segmentation_losses_test.py b/official/projects/volumetric_models/losses/segmentation_losses_test.py
index f2f444c2b23..f1f30f85c03 100644
--- a/official/projects/volumetric_models/losses/segmentation_losses_test.py
+++ b/official/projects/volumetric_models/losses/segmentation_losses_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for segmentation_losses.py."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.losses import segmentation_losses
diff --git a/official/projects/volumetric_models/modeling/__init__.py b/official/projects/volumetric_models/modeling/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/modeling/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/modeling/backbones/__init__.py b/official/projects/volumetric_models/modeling/backbones/__init__.py
index 8e220167f1d..b3d489ce29b 100644
--- a/official/projects/volumetric_models/modeling/backbones/__init__.py
+++ b/official/projects/volumetric_models/modeling/backbones/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/volumetric_models/modeling/backbones/unet_3d.py b/official/projects/volumetric_models/modeling/backbones/unet_3d.py
index c675315aa37..e6e096f844f 100644
--- a/official/projects/volumetric_models/modeling/backbones/unet_3d.py
+++ b/official/projects/volumetric_models/modeling/backbones/unet_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,27 +21,26 @@
from typing import Any, Mapping, Sequence
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.volumetric_models.modeling import nn_blocks_3d
from official.vision.modeling.backbones import factory
-layers = tf.keras.layers
+layers = tf_keras.layers
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class UNet3D(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class UNet3D(tf_keras.Model):
"""Class to build 3D UNet backbone."""
def __init__(
self,
model_id: int,
- input_specs: layers = layers.InputSpec(shape=[None, None, None, None, 3]),
+ input_specs: layers = layers.InputSpec(shape=[None, None, None, None, 3]), # pyrefly: ignore[not-a-type]
pool_size: Sequence[int] = (2, 2, 2),
kernel_size: Sequence[int] = (3, 3, 3),
base_filters: int = 32,
- kernel_regularizer: tf.keras.regularizers.Regularizer = None,
+ kernel_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
activation: str = 'relu',
norm_momentum: float = 0.99,
norm_epsilon: float = 0.001,
@@ -64,7 +63,7 @@ def __init__(
convolution network will have. Following layers will contain a multiple
of this number. Lowering this number will likely reduce the amount of
memory required to train the model.
- kernel_regularizer: A tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: A tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
activation: The name of the activation function.
norm_momentum: The normalization momentum for the moving average.
@@ -92,7 +91,7 @@ def __init__(
self._use_batch_normalization = use_batch_normalization
# Build 3D UNet.
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=input_specs.shape[1:], dtype=input_specs.dtype)
x = inputs
endpoints = {}
@@ -117,7 +116,7 @@ def __init__(
pool_size=pool_size,
strides=(2, 2, 2),
padding='valid',
- data_format=tf.keras.backend.image_data_format())(
+ data_format=tf_keras.backend.image_data_format())(
x2)
else:
x = x2
@@ -153,10 +152,10 @@ def output_specs(self) -> Mapping[str, tf.TensorShape]:
@factory.register_backbone_builder('unet_3d')
def build_unet3d(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
norm_activation_config: hyperparams.Config,
- l2_regularizer: tf.keras.regularizers.Regularizer = None) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds 3D UNet backbone from a config."""
backbone_type = backbone_config.type
backbone_cfg = backbone_config.get()
diff --git a/official/projects/volumetric_models/modeling/backbones/unet_3d_test.py b/official/projects/volumetric_models/modeling/backbones/unet_3d_test.py
index 01e86a9d1a3..cee56e42de3 100644
--- a/official/projects/volumetric_models/modeling/backbones/unet_3d_test.py
+++ b/official/projects/volumetric_models/modeling/backbones/unet_3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,8 @@
"""Tests for 3D UNet backbone."""
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.modeling.backbones import unet_3d
@@ -29,9 +28,9 @@ class UNet3DTest(parameterized.TestCase, tf.test.TestCase):
)
def test_network_creation(self, input_size, model_id):
"""Test creation of UNet3D family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
network = unet_3d.UNet3D(model_id=model_id)
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(input_size[0], input_size[0], input_size[1], 3), batch_size=1)
endpoints = network(inputs)
diff --git a/official/projects/volumetric_models/modeling/decoders/__init__.py b/official/projects/volumetric_models/modeling/decoders/__init__.py
index f699cffbcc9..f3fc3f19f88 100644
--- a/official/projects/volumetric_models/modeling/decoders/__init__.py
+++ b/official/projects/volumetric_models/modeling/decoders/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/volumetric_models/modeling/decoders/factory.py b/official/projects/volumetric_models/modeling/decoders/factory.py
index 0e1d17c4299..c8f71f19d27 100644
--- a/official/projects/volumetric_models/modeling/decoders/factory.py
+++ b/official/projects/volumetric_models/modeling/decoders/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -41,9 +41,7 @@ def build_my_decoder():
"""
from typing import Union, Mapping, Optional
-# Import libraries
-
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import registry
from official.modeling import hyperparams
@@ -58,7 +56,7 @@ def register_decoder_builder(key: str):
This decorator supports registration of decoder builder as follows:
```
- class MyDecoder(tf.keras.Model):
+ class MyDecoder(tf_keras.Model):
pass
@register_decoder_builder('mydecoder')
@@ -83,7 +81,7 @@ def builder(input_specs, config, l2_reg):
def build_identity(
input_specs: Optional[Mapping[str, tf.TensorShape]] = None,
model_config: Optional[hyperparams.Config] = None,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None) -> None:
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None) -> None:
del input_specs, model_config, l2_regularizer # Unused by identity decoder.
return None
@@ -91,15 +89,15 @@ def build_identity(
def build_decoder(
input_specs: Mapping[str, tf.TensorShape],
model_config: hyperparams.Config,
- l2_regularizer: tf.keras.regularizers.Regularizer = None,
- **kwargs) -> Union[None, tf.keras.Model, tf.keras.layers.Layer]: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
+ **kwargs) -> Union[None, tf_keras.Model, tf_keras.layers.Layer]: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds decoder from a config.
Args:
input_specs: A `dict` of input specifications. A dictionary consists of
{level: TensorShape} from a backbone.
model_config: A `OneOfConfig` of model config.
- l2_regularizer: A `tf.keras.regularizers.Regularizer` object. Default to
+ l2_regularizer: A `tf_keras.regularizers.Regularizer` object. Default to
None.
**kwargs: Additional keyword args to be passed to decoder builder.
diff --git a/official/projects/volumetric_models/modeling/decoders/factory_test.py b/official/projects/volumetric_models/modeling/decoders/factory_test.py
index bcd4df69450..c243b2d420b 100644
--- a/official/projects/volumetric_models/modeling/decoders/factory_test.py
+++ b/official/projects/volumetric_models/modeling/decoders/factory_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for factory functions."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from official.projects.volumetric_models.configs import decoders as decoders_cfg
diff --git a/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder.py b/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder.py
index 534c43f901f..50bc2211fe1 100644
--- a/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder.py
+++ b/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,17 +21,17 @@
from typing import Any, Dict, Mapping, Optional, Sequence
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.volumetric_models.modeling import nn_blocks_3d
from official.projects.volumetric_models.modeling.decoders import factory
-layers = tf.keras.layers
+layers = tf_keras.layers
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class UNet3DDecoder(tf.keras.Model):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class UNet3DDecoder(tf_keras.Model):
"""Class to build 3D UNet decoder."""
def __init__(self,
@@ -39,7 +39,7 @@ def __init__(self,
input_specs: Mapping[str, tf.TensorShape],
pool_size: Sequence[int] = (2, 2, 2),
kernel_size: Sequence[int] = (3, 3, 3),
- kernel_regularizer: tf.keras.regularizers.Regularizer = None,
+ kernel_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
activation: str = 'relu',
norm_momentum: float = 0.99,
norm_epsilon: float = 0.001,
@@ -57,7 +57,7 @@ def __init__(self,
{level: TensorShape} from a backbone.
pool_size: The pooling size for the max pooling operations.
kernel_size: The kernel size for 3D convolution.
- kernel_regularizer: A tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: A tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
activation: The name of the activation function.
norm_momentum: The normalization momentum for the moving average.
@@ -89,13 +89,13 @@ def __init__(self,
self._norm = layers.BatchNormalization
self._use_batch_normalization = use_batch_normalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
channel_dim = -1
else:
channel_dim = 1
# Build 3D UNet.
- inputs = self._build_input_pyramid(input_specs, model_id)
+ inputs = self._build_input_pyramid(input_specs, model_id) # pytype: disable=wrong-arg-types # dynamic-method-lookup
# Add levels with up-convolution or up-sampling.
x = inputs[str(model_id)]
@@ -141,7 +141,7 @@ def _build_input_pyramid(self, input_specs: Dict[str, tf.TensorShape],
inputs = {}
for level, spec in input_specs.items():
- inputs[level] = tf.keras.Input(shape=spec[1:])
+ inputs[level] = tf_keras.Input(shape=spec[1:])
return inputs
def get_config(self) -> Mapping[str, Any]:
@@ -161,19 +161,19 @@ def output_specs(self) -> Mapping[str, tf.TensorShape]:
def build_unet_3d_decoder(
input_specs: Mapping[str, tf.TensorShape],
model_config: hyperparams.Config,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None
-) -> tf.keras.Model:
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None
+) -> tf_keras.Model:
"""Builds UNet3D decoder from a config.
Args:
input_specs: A `dict` of input specifications. A dictionary consists of
{level: TensorShape} from a backbone.
model_config: A OneOfConfig. Model config.
- l2_regularizer: A `tf.keras.regularizers.Regularizer` instance. Default to
+ l2_regularizer: A `tf_keras.regularizers.Regularizer` instance. Default to
None.
Returns:
- A `tf.keras.Model` instance of the UNet3D decoder.
+ A `tf_keras.Model` instance of the UNet3D decoder.
"""
decoder_type = model_config.decoder.type
decoder_cfg = model_config.decoder.get()
@@ -184,7 +184,7 @@ def build_unet_3d_decoder(
model_id=decoder_cfg.model_id,
input_specs=input_specs,
pool_size=decoder_cfg.pool_size,
- kernel_regularizer=l2_regularizer,
+ kernel_regularizer=l2_regularizer, # pyrefly: ignore[bad-argument-type]
activation=norm_activation_config.activation,
norm_momentum=norm_activation_config.norm_momentum,
norm_epsilon=norm_activation_config.norm_epsilon,
diff --git a/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder_test.py b/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder_test.py
index d901a6e46eb..8c8604bc517 100644
--- a/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder_test.py
+++ b/official/projects/volumetric_models/modeling/decoders/unet_3d_decoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,8 @@
"""Tests for 3D UNet decoder."""
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.modeling.backbones import unet_3d
from official.projects.volumetric_models.modeling.decoders import unet_3d_decoder
@@ -30,10 +29,10 @@ class UNet3DDecoderTest(parameterized.TestCase, tf.test.TestCase):
)
def test_network_creation(self, input_size, model_id):
"""Test creation of UNet3D family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
# `input_size` consists of [spatial size, volume size].
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(input_size[0], input_size[0], input_size[1], 3), batch_size=1)
backbone = unet_3d.UNet3D(model_id=model_id)
network = unet_3d_decoder.UNet3DDecoder(
diff --git a/official/projects/volumetric_models/modeling/factory.py b/official/projects/volumetric_models/modeling/factory.py
index 85d3e7e44df..8330aead93b 100644
--- a/official/projects/volumetric_models/modeling/factory.py
+++ b/official/projects/volumetric_models/modeling/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,10 +13,8 @@
# limitations under the License.
"""Factory methods to build models."""
-
-# Import libraries
-
-import tensorflow as tf
+from typing import Sequence, Union
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
from official.projects.volumetric_models.modeling.decoders import factory as decoder_factory
@@ -26,9 +24,11 @@
def build_segmentation_model_3d(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: Union[tf_keras.layers.InputSpec,
+ Sequence[tf_keras.layers.InputSpec]],
model_config: hyperparams.Config,
- l2_regularizer: tf.keras.regularizers.Regularizer = None) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None # pyrefly: ignore[bad-function-definition]
+) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds Segmentation model."""
norm_activation_config = model_config.norm_activation
backbone = backbone_factory.build_backbone(
diff --git a/official/projects/volumetric_models/modeling/factory_test.py b/official/projects/volumetric_models/modeling/factory_test.py
index 2de27eeb833..8d0e3a34ff6 100644
--- a/official/projects/volumetric_models/modeling/factory_test.py
+++ b/official/projects/volumetric_models/modeling/factory_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for factory.py."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.projects.volumetric_models.configs import semantic_segmentation_3d as exp_cfg
@@ -30,19 +30,19 @@ class SegmentationModelBuilderTest(parameterized.TestCase, tf.test.TestCase):
((64, 64, 64), None, False))
def test_unet3d_builder(self, input_size, weight_decay, use_bn):
num_classes = 3
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None, input_size[0], input_size[1], input_size[2], 3])
model_config = exp_cfg.SemanticSegmentationModel3D(num_classes=num_classes)
model_config.head.use_batch_normalization = use_bn
l2_regularizer = (
- tf.keras.regularizers.l2(weight_decay) if weight_decay else None)
+ tf_keras.regularizers.l2(weight_decay) if weight_decay else None)
model = factory.build_segmentation_model_3d(
input_specs=input_specs,
model_config=model_config,
l2_regularizer=l2_regularizer)
self.assertIsInstance(
- model, tf.keras.Model,
- 'Output should be a tf.keras.Model instance but got %s' % type(model))
+ model, tf_keras.Model,
+ 'Output should be a tf_keras.Model instance but got %s' % type(model))
if __name__ == '__main__':
diff --git a/official/projects/volumetric_models/modeling/heads/__init__.py b/official/projects/volumetric_models/modeling/heads/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/modeling/heads/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d.py b/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d.py
index f1fc7ef0a67..36fb540f74e 100644
--- a/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d.py
+++ b/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,13 +15,13 @@
"""Segmentation heads."""
from typing import Any, Union, Sequence, Mapping, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class SegmentationHead3D(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class SegmentationHead3D(tf_keras.layers.Layer):
"""Segmentation head for 3D input."""
def __init__(self,
@@ -35,8 +35,8 @@ def __init__(self,
norm_momentum: float = 0.99,
norm_epsilon: float = 0.001,
use_batch_normalization: bool = False,
- kernel_regularizer: tf.keras.regularizers.Regularizer = None,
- bias_regularizer: tf.keras.regularizers.Regularizer = None,
+ kernel_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
+ bias_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
output_logits: bool = True, # pytype: disable=annotation-type-mismatch # typed-keras
**kwargs):
"""Initialize params to build segmentation head.
@@ -60,9 +60,9 @@ def __init__(self,
norm_epsilon: `float`, the epsilon parameter of the normalization layers.
use_batch_normalization: A bool of whether to use batch normalization or
not.
- kernel_regularizer: `tf.keras.regularizers.Regularizer` object for layer
+ kernel_regularizer: `tf_keras.regularizers.Regularizer` object for layer
kernel.
- bias_regularizer: `tf.keras.regularizers.Regularizer` object for bias.
+ bias_regularizer: `tf_keras.regularizers.Regularizer` object for bias.
output_logits: A `bool` of whether to output logits or not. Default
is True. If set to False, output softmax.
**kwargs: other keyword arguments passed to Layer.
@@ -84,28 +84,28 @@ def __init__(self,
'bias_regularizer': bias_regularizer,
'output_logits': output_logits
}
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
- self._activation = tf_utils.get_activation(activation)
+ self._activation = tf_utils.get_activation(activation, use_keras_layer=True)
def build(self, input_shape: Union[tf.TensorShape, Sequence[tf.TensorShape]]):
"""Creates the variables of the segmentation head."""
- conv_op = tf.keras.layers.Conv3D
+ conv_op = tf_keras.layers.Conv3D
conv_kwargs = {
'kernel_size': (3, 3, 3),
'padding': 'same',
'use_bias': False,
- 'kernel_initializer': tf.keras.initializers.RandomNormal(stddev=0.01),
+ 'kernel_initializer': tf_keras.initializers.RandomNormal(stddev=0.01),
'kernel_regularizer': self._config_dict['kernel_regularizer'],
}
final_kernel_size = (1, 1, 1)
bn_op = (
- tf.keras.layers.experimental.SyncBatchNormalization
+ tf_keras.layers.experimental.SyncBatchNormalization
if self._config_dict['use_sync_bn'] else
- tf.keras.layers.BatchNormalization)
+ tf_keras.layers.BatchNormalization)
bn_kwargs = {
'axis': self._bn_axis,
'momentum': self._config_dict['norm_momentum'],
@@ -133,7 +133,7 @@ def build(self, input_shape: Union[tf.TensorShape, Sequence[tf.TensorShape]]):
padding='valid',
activation=None,
bias_initializer=tf.zeros_initializer(),
- kernel_initializer=tf.keras.initializers.RandomNormal(stddev=0.01),
+ kernel_initializer=tf_keras.initializers.RandomNormal(stddev=0.01),
kernel_regularizer=self._config_dict['kernel_regularizer'],
bias_regularizer=self._config_dict['bias_regularizer'])
@@ -170,10 +170,10 @@ def call(self, inputs: Tuple[Union[tf.Tensor, Mapping[str, tf.Tensor]],
x = self._norms[i](x)
x = self._activation(x)
- x = tf.keras.layers.UpSampling3D(size=self._config_dict['upsample_factor'])(
+ x = tf_keras.layers.UpSampling3D(size=self._config_dict['upsample_factor'])(
x)
x = self._classifier(x)
- return x if self._config_dict['output_logits'] else tf.keras.layers.Softmax(
+ return x if self._config_dict['output_logits'] else tf_keras.layers.Softmax(
dtype='float32')(
x)
diff --git a/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d_test.py b/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d_test.py
index 6c3aee1ee92..839013594d3 100644
--- a/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d_test.py
+++ b/official/projects/volumetric_models/modeling/heads/segmentation_heads_3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.modeling.heads import segmentation_heads_3d
diff --git a/official/projects/volumetric_models/modeling/nn_blocks_3d.py b/official/projects/volumetric_models/modeling/nn_blocks_3d.py
index a7c459b0dc1..96769a86612 100644
--- a/official/projects/volumetric_models/modeling/nn_blocks_3d.py
+++ b/official/projects/volumetric_models/modeling/nn_blocks_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,15 +16,14 @@
from typing import Sequence, Union
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.vision.modeling.layers import nn_layers
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class BasicBlock3DVolume(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class BasicBlock3DVolume(tf_keras.layers.Layer):
"""A basic 3d convolution block."""
def __init__(self,
@@ -32,8 +31,8 @@ def __init__(self,
strides: Union[int, Sequence[int]],
kernel_size: Union[int, Sequence[int]],
kernel_initializer: str = 'VarianceScaling',
- kernel_regularizer: tf.keras.regularizers.Regularizer = None,
- bias_regularizer: tf.keras.regularizers.Regularizer = None,
+ kernel_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
+ bias_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
activation: str = 'relu',
use_sync_bn: bool = False,
norm_momentum: float = 0.99,
@@ -53,9 +52,9 @@ def __init__(self,
height and width of the 3D convolution window. Can be a single integer
to specify the same value for all spatial dimensions.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
activation: `str` name of the activation function.
use_sync_bn: if True, use synchronized batch normalization.
@@ -84,10 +83,10 @@ def __init__(self,
self._use_batch_normalization = use_batch_normalization
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -99,12 +98,12 @@ def build(self, input_shape: tf.TensorShape):
self._norms = []
for filters in self._filters:
self._convs.append(
- tf.keras.layers.Conv3D(
+ tf_keras.layers.Conv3D(
filters=filters,
kernel_size=self._kernel_size,
strides=self._strides,
padding='same',
- data_format=tf.keras.backend.image_data_format(),
+ data_format=tf_keras.backend.image_data_format(),
activation=None))
self._norms.append(
self._norm(
@@ -132,7 +131,7 @@ def get_config(self):
base_config = super(BasicBlock3DVolume, self).get_config()
return dict(list(base_config.items()) + list(config.items()))
- def call(self, inputs: tf.Tensor, training: bool = None) -> tf.Tensor:
+ def call(self, inputs: tf.Tensor, training: bool = None) -> tf.Tensor: # pytype: disable=annotation-type-mismatch
"""Runs forward pass on the input tensor."""
x = inputs
for conv, norm in zip(self._convs, self._norms):
@@ -143,8 +142,8 @@ def call(self, inputs: tf.Tensor, training: bool = None) -> tf.Tensor:
return x
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class ResidualBlock3DVolume(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class ResidualBlock3DVolume(tf_keras.layers.Layer):
"""A residual 3d block."""
def __init__(self,
@@ -176,9 +175,9 @@ def __init__(self,
stochastic_depth_drop_rate: `float` or None. if not None, drop rate for
the stochastic depth layer.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
activation: `str` name of the activation function.
use_sync_bn: if True, use synchronized batch normalization.
@@ -203,10 +202,10 @@ def __init__(self,
self._bias_regularizer = bias_regularizer
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -214,7 +213,7 @@ def __init__(self,
def build(self, input_shape):
if self._use_projection:
- self._shortcut = tf.keras.layers.Conv3D(
+ self._shortcut = tf_keras.layers.Conv3D(
filters=self._filters,
kernel_size=1,
strides=self._strides,
@@ -227,7 +226,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv1 = tf.keras.layers.Conv3D(
+ self._conv1 = tf_keras.layers.Conv3D(
filters=self._filters,
kernel_size=3,
strides=self._strides,
@@ -241,7 +240,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv2 = tf.keras.layers.Conv3D(
+ self._conv2 = tf_keras.layers.Conv3D(
filters=self._filters,
kernel_size=3,
strides=1,
@@ -315,8 +314,8 @@ def call(self, inputs, training=None):
return self._activation_fn(x + shortcut)
-@tf.keras.utils.register_keras_serializable(package='Vision')
-class BottleneckBlock3DVolume(tf.keras.layers.Layer):
+@tf_keras.utils.register_keras_serializable(package='Vision')
+class BottleneckBlock3DVolume(tf_keras.layers.Layer):
"""A standard bottleneck block."""
def __init__(self,
@@ -350,9 +349,9 @@ def __init__(self,
stochastic_depth_drop_rate: `float` or None. if not None, drop rate for
the stochastic depth layer.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
Default to None.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
Default to None.
activation: `str` name of the activation function.
use_sync_bn: if True, use synchronized batch normalization.
@@ -377,10 +376,10 @@ def __init__(self,
self._kernel_regularizer = kernel_regularizer
self._bias_regularizer = bias_regularizer
if use_sync_bn:
- self._norm = tf.keras.layers.experimental.SyncBatchNormalization
+ self._norm = tf_keras.layers.experimental.SyncBatchNormalization
else:
- self._norm = tf.keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ self._norm = tf_keras.layers.BatchNormalization
+ if tf_keras.backend.image_data_format() == 'channels_last':
self._bn_axis = -1
else:
self._bn_axis = 1
@@ -388,7 +387,7 @@ def __init__(self,
def build(self, input_shape):
if self._use_projection:
- self._shortcut = tf.keras.layers.Conv3D(
+ self._shortcut = tf_keras.layers.Conv3D(
filters=self._filters * 4,
kernel_size=1,
strides=self._strides,
@@ -401,7 +400,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv1 = tf.keras.layers.Conv3D(
+ self._conv1 = tf_keras.layers.Conv3D(
filters=self._filters,
kernel_size=1,
strides=1,
@@ -414,7 +413,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv2 = tf.keras.layers.Conv3D(
+ self._conv2 = tf_keras.layers.Conv3D(
filters=self._filters,
kernel_size=3,
strides=self._strides,
@@ -429,7 +428,7 @@ def build(self, input_shape):
momentum=self._norm_momentum,
epsilon=self._norm_epsilon)
- self._conv3 = tf.keras.layers.Conv3D(
+ self._conv3 = tf_keras.layers.Conv3D(
filters=self._filters * 4,
kernel_size=1,
strides=1,
diff --git a/official/projects/volumetric_models/modeling/nn_blocks_3d_test.py b/official/projects/volumetric_models/modeling/nn_blocks_3d_test.py
index cd18c6b27d0..d7271824f91 100644
--- a/official/projects/volumetric_models/modeling/nn_blocks_3d_test.py
+++ b/official/projects/volumetric_models/modeling/nn_blocks_3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,8 @@
"""Tests for 3D volumeric convoluion blocks."""
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.modeling import nn_blocks_3d
@@ -26,7 +25,7 @@ class NNBlocks3DTest(parameterized.TestCase, tf.test.TestCase):
@parameterized.parameters((128, 128, 32, 1), (256, 256, 16, 2))
def test_bottleneck_block_3d_volume_creation(self, spatial_size, volume_size,
filters, strides):
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(spatial_size, spatial_size, volume_size, filters * 4),
batch_size=1)
block = nn_blocks_3d.BottleneckBlock3DVolume(
@@ -46,7 +45,7 @@ def test_bottleneck_block_3d_volume_creation(self, spatial_size, volume_size,
@parameterized.parameters((128, 128, 32, 1), (256, 256, 64, 2))
def test_residual_block_3d_volume_creation(self, spatial_size, volume_size,
filters, strides):
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(spatial_size, spatial_size, volume_size, filters), batch_size=1)
block = nn_blocks_3d.ResidualBlock3DVolume(
filters=filters,
@@ -65,7 +64,7 @@ def test_residual_block_3d_volume_creation(self, spatial_size, volume_size,
@parameterized.parameters((128, 128, 64, 1, 3), (256, 256, 128, 2, 1))
def test_basic_block_3d_volume_creation(self, spatial_size, volume_size,
filters, strides, kernel_size):
- inputs = tf.keras.Input(
+ inputs = tf_keras.Input(
shape=(spatial_size, spatial_size, volume_size, filters), batch_size=1)
block = nn_blocks_3d.BasicBlock3DVolume(
filters=filters, strides=strides, kernel_size=kernel_size)
diff --git a/official/projects/volumetric_models/modeling/segmentation_model_test.py b/official/projects/volumetric_models/modeling/segmentation_model_test.py
index f5df0a4241d..fdccac3d071 100644
--- a/official/projects/volumetric_models/modeling/segmentation_model_test.py
+++ b/official/projects/volumetric_models/modeling/segmentation_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models.modeling import backbones
from official.projects.volumetric_models.modeling import decoders
from official.projects.volumetric_models.modeling.heads import segmentation_heads_3d
@@ -35,7 +35,7 @@ def test_segmentation_network_unet3d_creation(self, input_size, depth):
"""Test for creation of a segmentation network."""
num_classes = 2
inputs = np.random.rand(2, input_size[0], input_size[0], input_size[1], 3)
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
backbone = backbones.UNet3D(model_id=depth)
decoder = decoders.UNet3DDecoder(
diff --git a/official/projects/volumetric_models/registry_imports.py b/official/projects/volumetric_models/registry_imports.py
index 461a0028f14..e412d58c669 100644
--- a/official/projects/volumetric_models/registry_imports.py
+++ b/official/projects/volumetric_models/registry_imports.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/volumetric_models/serving/__init__.py b/official/projects/volumetric_models/serving/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/serving/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/serving/export_saved_model.py b/official/projects/volumetric_models/serving/export_saved_model.py
index 0fcf249ef9e..22d5f8d3a04 100644
--- a/official/projects/volumetric_models/serving/export_saved_model.py
+++ b/official/projects/volumetric_models/serving/export_saved_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/volumetric_models/serving/semantic_segmentation_3d.py b/official/projects/volumetric_models/serving/semantic_segmentation_3d.py
index a85399c43b6..8fb96360125 100644
--- a/official/projects/volumetric_models/serving/semantic_segmentation_3d.py
+++ b/official/projects/volumetric_models/serving/semantic_segmentation_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Mapping
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.projects.volumetric_models.modeling import backbones
@@ -28,10 +28,10 @@
class SegmentationModule(export_base.ExportModule):
"""Segmentation Module."""
- def _build_model(self) -> tf.keras.Model:
+ def _build_model(self) -> tf_keras.Model:
"""Builds and returns a segmentation model."""
num_channels = self.params.task.model.num_channels
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[self._batch_size] + self._input_image_size + [num_channels])
return factory.build_segmentation_model_3d(
diff --git a/official/projects/volumetric_models/serving/semantic_segmentation_3d_test.py b/official/projects/volumetric_models/serving/semantic_segmentation_3d_test.py
index 4001b829d68..9d0c15f1377 100644
--- a/official/projects/volumetric_models/serving/semantic_segmentation_3d_test.py
+++ b/official/projects/volumetric_models/serving/semantic_segmentation_3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: disable=unused-import
from official.core import exp_factory
@@ -60,9 +60,9 @@ def _get_dummy_input(self, input_type):
image_tensor = tf.convert_to_tensor(self._image_array, dtype=tf.uint8)
return tf.expand_dims(image_tensor, axis=0)
if input_type == 'image_bytes':
- return [self._image_array.tostring()]
+ return [self._image_array.tobytes()]
if input_type == 'tf_example':
- encoded_image = self._image_array.tostring()
+ encoded_image = self._image_array.tobytes()
example = tf.train.Example(
features=tf.train.Features(
feature={
diff --git a/official/projects/volumetric_models/tasks/__init__.py b/official/projects/volumetric_models/tasks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/volumetric_models/tasks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/volumetric_models/tasks/semantic_segmentation_3d.py b/official/projects/volumetric_models/tasks/semantic_segmentation_3d.py
index 928d5d26cbf..0cbc5cac2e0 100644
--- a/official/projects/volumetric_models/tasks/semantic_segmentation_3d.py
+++ b/official/projects/volumetric_models/tasks/semantic_segmentation_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from typing import Any, Dict, Mapping, Optional, Sequence, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import dataset_fn
from official.core import base_task
@@ -33,9 +33,9 @@
class SemanticSegmentation3DTask(base_task.Task):
"""A task for semantic segmentation."""
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
"""Builds segmentation model."""
- input_specs = tf.keras.layers.InputSpec(
+ input_specs = tf_keras.layers.InputSpec(
shape=[None] + self.task_config.model.input_size +
[self.task_config.model.num_channels],
dtype=self.task_config.train_data.dtype)
@@ -45,7 +45,7 @@ def build_model(self) -> tf.keras.Model:
# (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
# (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
l2_regularizer = (
- tf.keras.regularizers.l2(l2_weight_decay /
+ tf_keras.regularizers.l2(l2_weight_decay /
2.0) if l2_weight_decay else None)
model = factory.build_segmentation_model_3d(
@@ -66,7 +66,7 @@ def build_model(self) -> tf.keras.Model:
return model
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loads pretrained checkpoint."""
if not self.task_config.init_checkpoint:
return
@@ -143,13 +143,13 @@ def build_losses(self,
return total_loss
def build_metrics(self,
- training: bool = True) -> Sequence[tf.keras.metrics.Metric]:
+ training: bool = True) -> Sequence[tf_keras.metrics.Metric]:
"""Gets streaming metrics for training/validation."""
metrics = []
num_classes = self.task_config.model.num_classes
if training:
metrics.extend([
- tf.keras.metrics.CategoricalAccuracy(
+ tf_keras.metrics.CategoricalAccuracy(
name='train_categorical_accuracy', dtype=tf.float32)
])
else:
@@ -168,9 +168,9 @@ def build_metrics(self,
def train_step(
self,
inputs,
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
- metrics: Optional[Sequence[tf.keras.metrics.Metric]] = None
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[Sequence[tf_keras.metrics.Metric]] = None
) -> Dict[Any, Any]:
"""Does forward and backward.
@@ -211,14 +211,14 @@ def train_step(
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
# Scales back gradient before apply_gradients when LossScaleOptimizer is
# used.
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
optimizer.apply_gradients(list(zip(grads, tvars)))
@@ -236,8 +236,8 @@ def train_step(
def validation_step(
self,
inputs,
- model: tf.keras.Model,
- metrics: Optional[Sequence[tf.keras.metrics.Metric]] = None
+ model: tf_keras.Model,
+ metrics: Optional[Sequence[tf_keras.metrics.Metric]] = None
) -> Dict[Any, Any]:
"""Validatation step.
@@ -275,25 +275,25 @@ def validation_step(
return logs
- def inference_step(self, inputs, model: tf.keras.Model) -> tf.Tensor:
+ def inference_step(self, inputs, model: tf_keras.Model) -> tf.Tensor:
"""Performs the forward step."""
return model(inputs, training=False)
def aggregate_logs(
self,
state: Optional[Sequence[Union[segmentation_metrics.DiceScore,
- tf.keras.metrics.Metric]]] = None,
+ tf_keras.metrics.Metric]]] = None,
step_outputs: Optional[Mapping[str, Any]] = None
- ) -> Sequence[tf.keras.metrics.Metric]:
+ ) -> Sequence[tf_keras.metrics.Metric]:
"""Aggregates statistics to compute metrics over training.
Args:
- state: A sequence of tf.keras.metrics.Metric objects. Each element records
+ state: A sequence of tf_keras.metrics.Metric objects. Each element records
a metric.
step_outputs: A dictionary of [metric_name, (labels, output)] from a step.
Returns:
- An updated sequence of tf.keras.metrics.Metric objects.
+ An updated sequence of tf_keras.metrics.Metric objects.
"""
if state is None:
for metric in self.metrics:
@@ -301,8 +301,8 @@ def aggregate_logs(
state = self.metrics
for metric in self.metrics:
- labels = step_outputs[metric.name][0]
- predictions = step_outputs[metric.name][1]
+ labels = step_outputs[metric.name][0] # pyrefly: ignore[unsupported-operation]
+ predictions = step_outputs[metric.name][1] # pyrefly: ignore[unsupported-operation]
# If `step_output` is distributed, it contains a tuple of Tensors instead
# of a single Tensor, so we need to concatenate them along the batch
diff --git a/official/projects/volumetric_models/tasks/semantic_segmentation_3d_test.py b/official/projects/volumetric_models/tasks/semantic_segmentation_3d_test.py
index 08cf0e693d2..88c42a446e0 100644
--- a/official/projects/volumetric_models/tasks/semantic_segmentation_3d_test.py
+++ b/official/projects/volumetric_models/tasks/semantic_segmentation_3d_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl.testing import parameterized
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import registry_imports # pylint: disable=unused-import
from official.core import exp_factory
diff --git a/official/projects/volumetric_models/train.py b/official/projects/volumetric_models/train.py
index e04956fb722..a4faa5f34ea 100644
--- a/official/projects/volumetric_models/train.py
+++ b/official/projects/volumetric_models/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/volumetric_models/train_test.py b/official/projects/volumetric_models/train_test.py
index 50e8fa7e49a..13f5446fed7 100644
--- a/official/projects/volumetric_models/train_test.py
+++ b/official/projects/volumetric_models/train_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
from absl import flags
from absl import logging
from absl.testing import flagsaver
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.volumetric_models import train as train_lib
from official.vision.dataloaders import tfexample_utils
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/big_query_ops.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/big_query_ops.py
new file mode 100644
index 00000000000..4b8aeac3a8f
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/big_query_ops.py
@@ -0,0 +1,148 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Designed to interact with Google BigQuery.
+
+For the purpose of dataset and table management, as well as data ingestion
+from pandas DataFrames.
+"""
+
+import logging
+import os
+import subprocess
+from google.cloud import bigquery
+from google.cloud import exceptions
+import pandas as pd
+import pandas_gbq
+
+# Configure logging
+logging.basicConfig(
+ level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
+)
+
+# Centralized Schema Definition
+_BIGQUERY_SCHEMA = [
+ bigquery.SchemaField("particle", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("source_name", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("image_name", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("detection_scores", "FLOAT", mode="REQUIRED"),
+ bigquery.SchemaField("creation_time", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_0", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_1", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_2", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_3", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("detected_classes", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField(
+ "detected_classes_names", "STRING", mode="REQUIRED"
+ ),
+ bigquery.SchemaField("detected_colors", "STRING", mode="REQUIRED"),
+]
+
+
+class BigQueryManager:
+ """Manages interactions with Google BigQuery for dataset and table operations.
+
+ This class provides methods to create datasets and tables, ingest data from
+ pandas DataFrames, and manage related file operations in Google Cloud Storage.
+ """
+
+ def __init__(self, project_id: str, dataset_id: str, table_id: str):
+ """Initializes the BigQuery client and storage coordinates."""
+ self.client = bigquery.Client(project=project_id)
+ self.project_id = project_id
+ self.dataset_id = dataset_id
+ self.table_id = table_id
+ self.table_ref = f"{project_id}.{dataset_id}.{table_id}"
+
+ def _ensure_dataset(self):
+ """Checks if dataset exists, creates it if not."""
+ dataset_ref = self.client.dataset(self.dataset_id)
+ try:
+ self.client.get_dataset(dataset_ref)
+ except exceptions.NotFound:
+ logging.info("Dataset %s not found. Creating...", self.dataset_id)
+ dataset = bigquery.Dataset(dataset_ref)
+ self.client.create_dataset(dataset, timeout=30)
+
+ def create_table(self, overwrite: bool = False) -> None:
+ """Creates the table with the defined schema."""
+ self._ensure_dataset()
+
+ try:
+ self.client.get_table(self.table_ref)
+ if overwrite:
+ logging.info("Overwriting table %s...", self.table_id)
+ self.client.delete_table(self.table_ref)
+ else:
+ logging.info("Table %s already exists. Skipping.", self.table_id)
+ return
+ except exceptions.NotFound:
+ pass
+
+ table = bigquery.Table(self.table_ref, schema=_BIGQUERY_SCHEMA)
+ self.client.create_table(table)
+ logging.info("Table %s created successfully.", self.table_id)
+
+ def ingest_data(self, df: pd.DataFrame) -> None:
+ """Ingests data from a pandas DataFrame into BigQuery using pandas_gbq."""
+ pandas_gbq.to_gbq(
+ df,
+ destination_table=self.table_ref,
+ project_id=self.project_id,
+ if_exists="append",
+ )
+ logging.info("Data ingested successfully into %s", self.table_ref)
+
+ def upload_image_results_to_storage_bucket(
+ self, input_directory: str, prediction_folder: str, output_directory: str
+ ) -> None:
+ """Moves folders to the destination bucket and cleans up local directories.
+
+ Args:
+ input_directory: Path to the local input directory.
+ prediction_folder: Path to the local folder containing results.
+ output_directory: The GCS path (gs://...) for output.
+ """
+ try:
+ commands = [
+ f"rm -r {os.path.basename(input_directory)}",
+ f"gsutil -m cp -r {prediction_folder} {output_directory}",
+ f"rm -r {prediction_folder}",
+ ]
+ subprocess.run(" && ".join(commands), shell=True, check=True)
+ logging.info("Successfully moved to destination bucket")
+ except (
+ KeyError,
+ IndexError,
+ TypeError,
+ ValueError,
+ subprocess.CalledProcessError,
+ ) as e:
+ logging.info(
+ "Issue in moving folders to destination bucket, due to error : %s", e
+ )
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/big_query_ops_test.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/big_query_ops_test.py
new file mode 100644
index 00000000000..c0fc0f35170
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/big_query_ops_test.py
@@ -0,0 +1,223 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import subprocess
+import unittest
+from unittest import mock
+
+from google.cloud import exceptions
+import pandas as pd
+
+from official.projects.waste_identification_ml.Deploy.detr_cloud_deployment.client import big_query_ops
+
+MODULE_PATH = big_query_ops.__name__
+
+
+class BigQueryManagerTest(unittest.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ self.mock_bigquery_client_patch = mock.patch(
+ f"{MODULE_PATH}.bigquery.Client"
+ )
+ self.mock_bigquery_client = self.mock_bigquery_client_patch.start()
+ self.mock_pandas_gbq_patch = mock.patch(f"{MODULE_PATH}.pandas_gbq.to_gbq")
+ self.mock_pandas_gbq = self.mock_pandas_gbq_patch.start()
+ self.mock_subprocess_run_patch = mock.patch(f"{MODULE_PATH}.subprocess.run")
+ self.mock_subprocess_run = self.mock_subprocess_run_patch.start()
+
+ self.project_id = "test-project"
+ self.dataset_id = "test-dataset"
+ self.table_id = "test-table"
+ self.manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+
+ def tearDown(self):
+ super().tearDown()
+ mock.patch.stopall()
+
+ def test_init_sets_attributes(self):
+ self.assertEqual(self.manager.project_id, self.project_id)
+ self.assertEqual(self.manager.dataset_id, self.dataset_id)
+ self.assertEqual(self.manager.table_id, self.table_id)
+ self.assertEqual(
+ self.manager.table_ref,
+ f"{self.project_id}.{self.dataset_id}.{self.table_id}",
+ )
+
+ def test_init_creates_client(self):
+ self.mock_bigquery_client.assert_called_once_with(project=self.project_id)
+
+ def test_ensure_dataset_exists(self):
+ self.manager.client.get_dataset.return_value = True
+
+ self.manager._ensure_dataset()
+
+ self.manager.client.get_dataset.assert_called_once_with(
+ self.manager.client.dataset(self.dataset_id)
+ )
+ self.manager.client.create_dataset.assert_not_called()
+
+ @mock.patch(f"{MODULE_PATH}.bigquery.Dataset")
+ def test_ensure_dataset_not_found(self, mock_bq_dataset):
+ self.manager.client.get_dataset.side_effect = exceptions.NotFound(
+ "Dataset not found"
+ )
+ mock_dataset_ref = self.manager.client.dataset.return_value
+ mock_bq_dataset.return_value = "dataset_obj"
+
+ self.manager._ensure_dataset()
+
+ self.manager.client.get_dataset.assert_called_once_with(mock_dataset_ref)
+ self.manager.client.dataset.assert_called_once_with(self.dataset_id)
+ mock_bq_dataset.assert_called_once_with(mock_dataset_ref)
+ self.manager.client.create_dataset.assert_called_once_with(
+ "dataset_obj", timeout=30
+ )
+
+ @mock.patch(f"{MODULE_PATH}.bigquery.Table")
+ def test_create_table_new_table_created_if_not_exists(self, mock_bq_table):
+ self.manager.client.get_table.side_effect = exceptions.NotFound(
+ "Table not found"
+ )
+ table_obj = "table_obj"
+ mock_bq_table.return_value = table_obj
+ with mock.patch.object(self.manager, "_ensure_dataset"):
+
+ self.manager.create_table()
+
+ mock_bq_table.assert_called_once_with(
+ self.manager.table_ref, schema=self.manager._schema
+ )
+ self.manager.client.create_table.assert_called_once_with(table_obj)
+
+ def test_create_table_new_ensures_dataset_exists_if_table_not_exists(self):
+ self.manager.client.get_table.side_effect = exceptions.NotFound(
+ "Table not found"
+ )
+ with mock.patch.object(
+ self.manager, "_ensure_dataset"
+ ) as mock_ensure_dataset:
+
+ self.manager.create_table()
+
+ mock_ensure_dataset.assert_called_once()
+
+ def test_create_table_exists_no_overwrite_does_not_recreate(self):
+ self.manager.client.get_table.return_value = True
+ with mock.patch.object(self.manager, "_ensure_dataset"):
+
+ self.manager.create_table(overwrite=False)
+
+ self.manager.client.delete_table.assert_not_called()
+ self.manager.client.create_table.assert_not_called()
+
+ def test_create_table_exists_no_overwrite_ensures_dataset_exists(self):
+ self.manager.client.get_table.return_value = True
+ with mock.patch.object(
+ self.manager, "_ensure_dataset"
+ ) as mock_ensure_dataset:
+
+ self.manager.create_table(overwrite=False)
+
+ mock_ensure_dataset.assert_called_once()
+
+ @mock.patch(f"{MODULE_PATH}.bigquery.Table")
+ def test_create_table_exists_overwrite_recreates_table(self, mock_bq_table):
+ self.manager.client.get_table.return_value = True
+ table_obj = "table_obj"
+ mock_bq_table.return_value = table_obj
+ with mock.patch.object(self.manager, "_ensure_dataset"):
+
+ self.manager.create_table(overwrite=True)
+
+ self.manager.client.delete_table.assert_called_once_with(
+ self.manager.table_ref
+ )
+ mock_bq_table.assert_called_once_with(
+ self.manager.table_ref, schema=self.manager._schema
+ )
+ self.manager.client.create_table.assert_called_once_with(table_obj)
+
+ def test_create_table_exists_overwrite_ensures_dataset_exists(self):
+ self.manager.client.get_table.return_value = True
+ with mock.patch.object(
+ self.manager, "_ensure_dataset"
+ ) as mock_ensure_dataset:
+
+ self.manager.create_table(overwrite=True)
+
+ mock_ensure_dataset.assert_called_once()
+
+ def test_ingest_data(self):
+ df = pd.DataFrame({"col1": [1, 2], "col2": [3, 4]})
+
+ self.manager.ingest_data(df)
+
+ self.mock_pandas_gbq.assert_called_once_with(
+ df,
+ destination_table=self.manager.table_ref,
+ project_id=self.project_id,
+ if_exists="append",
+ )
+
+ def test_upload_image_results_to_storage_bucket_success(self):
+ input_dir = "/tmp/input"
+ pred_dir = "/tmp/pred"
+ output_dir = "gs://bucket/output"
+
+ self.manager.upload_image_results_to_storage_bucket(
+ input_dir, pred_dir, output_dir
+ )
+
+ self.mock_subprocess_run.assert_called_once()
+ args, _ = self.mock_subprocess_run.call_args
+ self.assertIn(f"gsutil -m cp -r {pred_dir} {output_dir}", args[0])
+
+ def test_upload_image_results_to_storage_bucket_failure(self):
+ input_dir = "/tmp/input"
+ pred_dir = "/tmp/pred"
+ output_dir = "gs://bucket/output"
+ self.mock_subprocess_run.side_effect = subprocess.CalledProcessError(
+ 1, "cmd"
+ )
+
+ with self.assertLogs(level="INFO") as cm:
+ self.manager.upload_image_results_to_storage_bucket(
+ input_dir, pred_dir, output_dir
+ )
+
+ self.mock_subprocess_run.assert_called_once()
+ self.assertIn(
+ "Issue in moving folders to destination bucket", cm.output[-1]
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/color_extraction.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/color_extraction.py
new file mode 100644
index 00000000000..93c5d373a99
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/color_extraction.py
@@ -0,0 +1,238 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Extract properties from each object mask and detect its color."""
+
+from typing import List, Tuple, TypeVar
+
+import numpy as np
+import numpy.typing as npt
+from skimage import color as skimage_color
+from sklearn import cluster as sklearn_cluster
+from sklearn import neighbors as sklearn_neighbors
+import webcolors
+
+DType = TypeVar('DType', bound=np.generic)
+# Color representation as numpy array of 3 elements of float64
+# Those values could be in different scales like
+# RGB ([0.0,255.0], [0.0,255.0], [0.0 to 255.0])
+# LAB ([0.0,100], [-128,127], [-128,127])
+# NColor = Annotated[npt.NDArray[DType], Literal[3]][np.float64]
+NColor = np.ndarray
+
+
+PROPERTIES = [
+ 'area',
+ 'bbox',
+ 'convex_area',
+ 'bbox_area',
+ 'major_axis_length',
+ 'minor_axis_length',
+ 'eccentricity',
+ 'centroid',
+]
+
+GENERIC_COLORS = [
+ ('black', '#000000'),
+ ('green', '#008000'),
+ ('green', '#00ff00'), # lime
+ ('green', '#3cb371'), # mediumseagreen
+ ('green', '#2E8B57'), # seagreen
+ ('green', '#8FBC8B'), # darkseagreen
+ ('green', '#adff2f'), # olive
+ ('green', '#008080'), # Teal
+ ('green', '#808000'),
+ ('blue', '#000080'), # navy
+ ('blue', '#00008b'), # darkblue
+ ('blue', '#4682b4'), # steelblue
+ ('blue', '#40E0D0'), # turquoise
+ ('blue', '#00FFFF'), # cyan
+ ('blue', '#00ffff'), # aqua
+ ('blue', '#6495ED'), # cornflowerBlue
+ ('blue', '#4169E1'), # royalBlue
+ ('blue', '#87CEFA'), # lightSkyBlue
+ ('blue', '#4682B4'), # steelBlue
+ ('blue', '#B0C4DE'), # lightSteelBlue
+ ('blue', '#87CEEB'), # skyblue
+ ('blue', '#0000CD'), # mediumBlue
+ ('blue', '#0000ff'),
+ ('purple', '#800080'),
+ ('purple', '#9370db'), # mediumpurple
+ ('purple', '#8B008B'), # darkMagenta
+ ('purple', '#4B0082'), # indigo
+ ('red', '#ff0000'),
+ ('red', '#B22222'), # fireBrick
+ ('red', '#DC143C'), # fireBrick
+ ('red', '#8B0000'), # crimson
+ ('red', '#CD5C5C'), # indianred
+ ('red', '#F08080'), # lightCoral
+ ('red', '#FA8072'), # salmon
+ ('red', '#E9967A'), # darkSalmon
+ ('red', '#FFA07A'), # lightSalmon
+ ('gray', '#c0c0c0'), # silver,
+ ('gray', '#a9a9a9'), # +darkgray
+ ('gray', '#708090'), # +slategray
+ ('blue', '#778899'), # +lightslategray
+ ('white', '#ffffff'),
+ ('white', '#F5F5DC'), # beige
+ ('white', '#FFFAFA'), # snow
+ ('white', '#F0F8FF'), # aliceBlue
+ ('white', '#FFE4E1'), # mistyRose
+ ('yellow', '#ffff00'),
+ ('yellow', '#ffffe0'), # lightyellow
+ ('yellow', '#8B8000'), # darkyellow,
+ ('orange', '#ffa500'),
+ ('orange', '#ff8c00'), # darkorange
+ ('pink', '#ffc0cb'),
+ ('pink', '#ff00ff'), # fuchsia
+ ('pink', '#C71585'), # mediumVioletRed
+ ('pink', '#DB7093'), # paleVioletRed
+ ('pink', '#FFB6C1'), # lightPink
+ ('pink', '#FF69B4'), # hotPink
+ ('pink', '#FF1493'), # deepPink
+ ('pink', '#BC8F8F'), # rosybrown
+ ('brown', '#a52a2a'),
+ ('brown', '#8b4513'), # saddlebrown
+ ('brown', '#f4a460'), # sandybrown
+ ('brown', '#800000'), # maroon
+]
+
+
+def find_dominant_color(
+ image: np.ndarray, black_threshold: int = 50
+) -> Tuple[int, int, int]:
+ """Determines the dominant color in a given image.
+
+ Args:
+ image: An array representation of the image.
+ black_threshold: The intensity threshold below which pixels are considered
+ 'black' or near-black.
+
+ Returns:
+ The dominant RGB color in the format (R, G, B).
+ """
+ pixels = image.reshape(-1, 3)
+
+ # Filter out black pixels based on the threshold
+ non_black_pixels = pixels[(pixels > black_threshold).any(axis=1)]
+
+ if non_black_pixels.size:
+ kmeans = sklearn_cluster.KMeans(
+ n_clusters=1, n_init=10, random_state=0
+ ).fit(non_black_pixels)
+ dominant_color = kmeans.cluster_centers_[0].astype(int)
+ else:
+ dominant_color = np.array([0, 0, 0], dtype=int)
+ return tuple(dominant_color)
+
+
+def rgb_int_to_lab(rgb_int_color: Tuple[int, int, int]) -> NColor:
+ """Convert RGB color to LAB color space.
+
+ Args:
+ rgb_int_color: RGB tuple color e.g. (128,128,128)
+
+ Returns:
+ Numpy array of 3 elements that contains LAB color space.
+ """
+ return skimage_color.rgb2lab(
+ (rgb_int_color[0] / 255, rgb_int_color[1] / 255, rgb_int_color[2] / 255)
+ )
+
+
+def color_distance(
+ a: Tuple[int, int, int], b: Tuple[int, int, int]
+) -> np.ndarray:
+ """The color distance following the ciede2000 formula.
+
+ See: https://en.wikipedia.org/wiki/Color_difference#CIEDE2000
+
+ Args:
+ a: Color a
+ b: Color b
+
+ Returns:
+ The distance between color a and b
+ """
+ return skimage_color.deltaE_ciede2000(a, b, kC=0.6)
+
+
+def build_color_lab_list(
+ generic_colors: List[Tuple[str, str]],
+) -> Tuple[npt.NDArray[np.str_], List[NColor]]:
+ """Get Simple colors names and lab values.
+
+ Args:
+ generic_colors: List of colors in this format (color_name, rgb_value in hex)
+ e.g. [ ('black', '#000000'), ('green', '#008000'), ]
+
+ Returns:
+ Numpy array of strings that contains color names
+ ['black', 'green']
+ List of color lab values in the format of Numpy array of 3 elements
+ e.g.
+ [
+ np.array([0., 0., 0.]),
+ np.array([ 46.2276577 , -51.69868348, 49.89707556])
+ ]
+ """
+ names: list[str] = []
+ lab_values = []
+ for color_name, color_hex in generic_colors:
+ names.append(color_name)
+ hex_color = webcolors.hex_to_rgb(color_hex)
+ lab_values.append(rgb_int_to_lab(hex_color))
+ color_names = np.array(names)
+ return color_names, lab_values
+
+
+def get_generic_color_name(
+ rgb_colors: List[Tuple[int, int, int]],
+ generic_colors: List[Tuple[str, str]] | None = None,
+) -> List[str]:
+ """Retrieves generic names of given RGB colors.
+
+ Estimates the closest matching color name.
+
+ Args:
+ rgb_colors: A list of RGB values for which to retrieve the name.
+ generic_colors: A list of color names and their RGB values in hex.
+
+ Returns:
+ The list of closest color names.
+
+ Example: get_generic_color_name([(255, 0, 0), (0,0,0)])
+ ['red','black']
+ """
+ names, rgb_simple_colors = build_color_lab_list(
+ generic_colors or GENERIC_COLORS
+ )
+ tree = sklearn_neighbors.BallTree(rgb_simple_colors, metric=color_distance)
+ rgb_query = [*map(rgb_int_to_lab, rgb_colors)]
+ _, index = tree.query(rgb_query)
+ return [x[0] for x in names[index]]
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/color_extraction_test.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/color_extraction_test.py
new file mode 100644
index 00000000000..891ca17665c
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/color_extraction_test.py
@@ -0,0 +1,88 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+import numpy as np
+from official.projects.waste_identification_ml.Deploy.detr_cloud_deployment.client import color_extraction
+
+
+class ColorExtractionTest(unittest.TestCase):
+
+ def test_find_dominant_color_with_non_black_pixels(self):
+ # Create an image with a clear dominant color (Red)
+ image = np.zeros((10, 10, 3), dtype=np.uint8)
+ image[0:5, 0:5] = [255, 0, 0] # Top-left quarter is Red
+ image[5:10, 5:10] = [100, 0, 0] # Bottom-right quarter is dark Red
+ dominant_color = color_extraction.find_dominant_color(image)
+ self.assertEqual(dominant_color, (177, 0, 0))
+
+ def test_find_dominant_color_with_only_black_pixels(self):
+ image = np.zeros((10, 10, 3), dtype=np.uint8)
+ dominant_color = color_extraction.find_dominant_color(
+ image, black_threshold=50
+ )
+ self.assertEqual(dominant_color, (0, 0, 0))
+
+ def test_rgb_int_to_lab(self):
+ rgb = (255, 255, 255)
+ lab = color_extraction.rgb_int_to_lab(rgb)
+ # White in LAB is approx (100, 0, 0)
+ self.assertIsInstance(lab, np.ndarray)
+ self.assertEqual(lab.shape, (3,))
+ np.testing.assert_allclose(lab, [100.0, 0.0, 0.0], atol=1e-2)
+
+ def test_color_distance(self):
+ color_a = (100, 0, 0) # LAB
+ color_b = (100, 0, 0) # LAB
+ distance = color_extraction.color_distance(color_a, color_b)
+ self.assertEqual(distance, 0.0)
+ color_c = (0, 0, 0)
+ distance_diff = color_extraction.color_distance(color_a, color_c)
+ self.assertGreater(distance_diff, 0.0)
+
+ def test_build_color_lab_list(self):
+ generic_colors = [('black', '#000000'), ('white', '#ffffff')]
+ names, lab_values = color_extraction.build_color_lab_list(generic_colors)
+ np.testing.assert_array_equal(names, ['black', 'white'])
+ self.assertEqual(len(lab_values), 2)
+ np.testing.assert_allclose(lab_values[0], [0.0, 0.0, 0.0], atol=1e-2)
+ np.testing.assert_allclose(lab_values[1], [100.0, 0.0, 0.0], atol=1e-2)
+
+ def test_get_generic_color_name(self):
+ rgb_colors = [(255, 0, 0), (0, 0, 255)] # Red, Blue
+ generic_colors = [
+ ('red', '#ff0000'),
+ ('blue', '#0000ff'),
+ ('green', '#00ff00'),
+ ]
+ names = color_extraction.get_generic_color_name(rgb_colors, generic_colors)
+ self.assertEqual(names, ['red', 'blue'])
+
+
+if __name__ == '__main__':
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/inference_pipeline.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/inference_pipeline.py
new file mode 100644
index 00000000000..fbd7911fcf6
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/inference_pipeline.py
@@ -0,0 +1,230 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Pipeline to run the prediction on the images folder with Triton server."""
+
+import cuml.accel # pylint: disable=g-bad-import-order, g-import-not-at-top
+
+cuml.accel.install() # pylint: disable=g-bad-import-order, g-import-not-at-top
+
+import os # pylint: disable=g-bad-import-order, g-import-not-at-top
+from absl import app
+from absl import flags
+from big_query_ops import BigQueryManager
+import cv2
+from object_tracking import ObjectTracker
+import pandas as pd
+from PIL import Image
+from triton_server_inference import TritonObjectDetector
+import utils
+
+
+INPUT_DIRECTORY = flags.DEFINE_string(
+ "input_directory", None, "The path to the directory containing images."
+)
+OUTPUT_DIRECTORY = flags.DEFINE_string(
+ "output_directory", None, "The path to the directory to save the results."
+)
+MODEL_NAME = flags.DEFINE_string("model_name", None, "Model name")
+PREDICTION_THRESHOLD = flags.DEFINE_float(
+ "threshold", None, "Threshold to filter the prediction results"
+)
+SEARCH_RANGE_X = flags.DEFINE_integer(
+ "search_range_x",
+ None,
+ "Pixels upto which every object needs to be tracked along X axis.",
+)
+SEARCH_RANGE_Y = flags.DEFINE_integer(
+ "search_range_y",
+ None,
+ "Pixels upto which every object needs to be tracked along Y axis.",
+)
+MEMORY = flags.DEFINE_integer(
+ "memory", None, "Frames upto which every object needs to be tracked."
+)
+OVERWRITE = flags.DEFINE_boolean(
+ "overwrite",
+ False,
+ "If True, delete the preexisting BigQuery table before creating a new one.",
+)
+PROJECT_ID = flags.DEFINE_string(
+ "project_id", None, "Project ID mentioned in Google Cloud Project"
+)
+BQ_DATASET_ID = flags.DEFINE_string(
+ "bq_dataset_id", "Circularnet_dataset", "Big query dataset ID"
+)
+BQ_TABLE_ID = flags.DEFINE_string(
+ "bq_table_id", "Circularnet_table", "BigQuery Table ID for features data"
+)
+
+_IMAGE_SAVING_WIDTH = 432
+_IMAGE_SAVING_HEIGHT = 432
+_TRACKING_IMAGE_WIDTH = 300
+_TRACKING_IMAGE_HEIGHT = 300
+# https://soft-matter.github.io/trackpy/dev/tutorial/adaptive-search.html
+_ADAPTIVE_STOP = 5
+_ADAPTIVE_STEP = 0.95
+
+
+def main(_) -> None:
+ # Check if the input and output directories are valid.
+ if (
+ not INPUT_DIRECTORY.value
+ or not OUTPUT_DIRECTORY.value
+ or not INPUT_DIRECTORY.value.startswith("gs://")
+ or not OUTPUT_DIRECTORY.value.startswith("gs://")
+ ):
+ raise ValueError("Bucket path must be non-empty starting with 'gs://'")
+
+ input_directory, prediction_folder, logger = (
+ utils.setup_logger_and_directories(input_dir=INPUT_DIRECTORY.value)
+ )
+ filepaths_to_capture_time_dict = {
+ filepath: utils.get_image_capture_time(filepath)
+ for filepath in utils.files_paths(os.path.basename(input_directory))
+ }
+
+ model_manager = TritonObjectDetector(model_name=MODEL_NAME.value)
+
+ tracking_manager = ObjectTracker(
+ search_range=(SEARCH_RANGE_Y.value, SEARCH_RANGE_X.value),
+ memory=MEMORY.value,
+ adaptive_stop=_ADAPTIVE_STOP,
+ adaptive_step=_ADAPTIVE_STEP,
+ )
+
+ storage_manager = BigQueryManager(
+ project_id=PROJECT_ID.value,
+ dataset_id=BQ_DATASET_ID.value,
+ table_id=BQ_TABLE_ID.value,
+ )
+
+ results = {}
+ tracking_df = pd.DataFrame()
+ for frame, (image_path, creation_time) in enumerate(
+ filepaths_to_capture_time_dict.items(), start=1
+ ):
+
+ logger.info(f"Processing {os.path.basename(image_path)}")
+
+ # Perform Inference
+ try:
+ results = model_manager.predict(
+ image_path=image_path,
+ confidence_threshold=PREDICTION_THRESHOLD.value,
+ max_boxes=100,
+ output_dims=(_IMAGE_SAVING_WIDTH, _IMAGE_SAVING_HEIGHT),
+ )
+ results["class_names"] = model_manager.get_class_names(results)
+ logger.info(
+ f"Successfully got prediction for {os.path.basename(image_path)},"
+ f" with total predictions: {len(results['labels'])}"
+ )
+ except (KeyError, TypeError, RuntimeError, ValueError) as e:
+ logger.info(
+ f"Failed to get prediction for {os.path.basename(image_path)}, due to"
+ f" error : {e}"
+ )
+
+ # Continue to next image if no objects detected
+ if results["class_names"].size == 0:
+ logger.info(f"No objects detected in {os.path.basename(image_path)}")
+ continue
+
+ # Image resizing for tracking and saving
+ original_image = cv2.imread(image_path)
+ image_for_tracking = cv2.resize(
+ original_image,
+ (_TRACKING_IMAGE_WIDTH, _TRACKING_IMAGE_HEIGHT),
+ interpolation=cv2.INTER_AREA,
+ )
+ image_for_saving = cv2.resize(
+ original_image,
+ (_IMAGE_SAVING_WIDTH, _IMAGE_SAVING_HEIGHT),
+ interpolation=cv2.INTER_AREA,
+ )
+
+ # Save the image with bounding boxes & masks
+ try:
+ pil_image = Image.fromarray(image_for_saving)
+ save_path = os.path.join(
+ prediction_folder, os.path.basename(image_path)
+ )
+ utils.draw_detections_and_save_image(pil_image, results, save_path)
+ logger.info("Image with bounding box saved")
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info(
+ f"Issue in saving visualization of results, due to error : {e}"
+ )
+
+ # Feature Extraction
+ try:
+ detected_colors = utils.extract_color_names(results, image_for_saving)
+ tracking_manager.extract_features_for_tracking(
+ image=image_for_tracking,
+ results=results,
+ tracking_image_size=(_TRACKING_IMAGE_WIDTH, _TRACKING_IMAGE_HEIGHT),
+ image_path=image_path,
+ creation_time=creation_time,
+ frame_idx=frame,
+ colors=detected_colors,
+ )
+
+ logger.info("Features extracted.\n")
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info(f"Failed to extract properties, due to error : {e}")
+
+ # Object Tracking
+ try:
+ particle_df = tracking_manager.run_tracking()
+ tracking_df = tracking_manager.process_tracking_results(particle_df)
+ counts = tracking_df.groupby("detected_classes_names").size()
+ counts.to_frame().to_csv(os.path.join(os.getcwd(), "count.csv"))
+ logger.info("Object tracking applied.")
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info(f"Failed to apply object tracking, due to error : {e}")
+
+ # Upload Results to BigQuery
+ if isinstance(tracking_df, pd.DataFrame) and not tracking_df.empty:
+ try:
+ storage_manager.create_table(overwrite=OVERWRITE.value)
+ storage_manager.ingest_data(tracking_df)
+ storage_manager.upload_image_results_to_storage_bucket(
+ input_directory=input_directory,
+ prediction_folder=prediction_folder,
+ output_directory=OUTPUT_DIRECTORY.value,
+ )
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info(f"Issue in creation of table, due to error : {e}")
+ return
+ else:
+ logger.info("No features to ingest.")
+ utils.shutdown_vm()
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/labels50.csv b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/labels50.csv
new file mode 100644
index 00000000000..cf6fc7955e7
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/labels50.csv
@@ -0,0 +1,51 @@
+id,names
+1,Aluminium_Can
+2,Aluminium_Foil
+3,Battery
+4,Brush
+5,Bulb
+6,Fiber_Cardboard
+7,Fiber_Cup-&-glass
+8,Fiber_Paper
+9,Footwear
+10,Glass_Bottle
+11,Lighter
+12,Metals
+13,Metals_Bottle
+14,Metals_Container
+15,Metals_Lid
+16,Plastics-ABS_Electronics
+17,Plastics-HDPE_Bottle
+18,Plastics-HDPE_Container
+19,Plastics-HDPE_Lid
+20,Plastics-HDPE_Toys
+21,Plastics-HDPE_Tube
+22,Plastics-LDPE_Flexibles
+23,Plastics-MLP_Flexibles
+24,Plastics-MLP_Tube
+25,Plastics-PC_CD
+26,Plastics-PC_Goggles
+27,Plastics-PET_Blister-pack
+28,Plastics-PET_Bottle
+29,Plastics-PET_Container
+30,Plastics-PET_Cup-&-glass
+31,Plastics-PP
+32,Plastics-PP_Comb
+33,Plastics-PP_Container
+34,Plastics-PP_Cup-&-glass
+35,Plastics-PP_Pen
+36,Plastics-PP_Spoon
+37,Plastics-PP_Straw
+38,Plastics-PP_Tray
+39,Plastics-PS
+40,Plastics-PS_Container
+41,Plastics-PS_Cup-&-glass
+42,Plastics-PS_Flexibles
+43,Plastics-PS_Hangers
+44,Plastics-PVC_Flexibles
+45,Plastics-PVC_Pipe
+46,Plastics-Tetrapak_Carton
+47,Textile_Clothes
+48,Textile_Flexibles
+49,Tire
+50,Wood
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/object_tracking.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/object_tracking.py
new file mode 100644
index 00000000000..a23a7b90ee0
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/object_tracking.py
@@ -0,0 +1,290 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Object tracking using trackpy."""
+
+import os
+from typing import Any, Dict, List
+import cv2
+import numpy as np
+import pandas as pd
+import skimage.measure
+import trackpy as tp
+
+
+class ObjectTracker:
+ """Tracks objects across multiple frames using trackpy.
+
+ This class collects object detections from multiple frames, extracts features,
+ links them using trackpy, and aggregates the tracking results.
+ """
+
+ def __init__(
+ self,
+ search_range: tuple[int, int] = (20, 20),
+ memory: int = 3,
+ adaptive_stop: int = 5,
+ adaptive_step: float = 0.95,
+ ):
+ """Initializes the tracker.
+
+ Args:
+ search_range: (y_range, x_range) pixels for tracking.
+ memory: Number of frames an object can vanish and still be linked.
+ adaptive_stop: Minimum search range (pixels) to try before giving up on
+ linking a subnet. Prevents SubnetOversizeException on dense frames.
+ adaptive_step: Multiplicative factor (0 < step < 1) by which the search
+ range is reduced each adaptive iteration.
+ """
+ self.search_range = search_range
+ self.memory = memory
+ self.adaptive_stop = adaptive_stop
+ self.adaptive_step = adaptive_step
+ self.all_detections: List[pd.DataFrame] = []
+
+ # Region properties to extract
+ self._properties = (
+ 'area',
+ 'bbox',
+ 'convex_area',
+ 'bbox_area',
+ 'major_axis_length',
+ 'minor_axis_length',
+ 'eccentricity',
+ 'centroid',
+ 'label',
+ 'mean_intensity',
+ 'max_intensity',
+ 'min_intensity',
+ 'perimeter',
+ )
+
+ def extract_features_for_tracking(
+ self,
+ image: np.ndarray,
+ results: Dict[str, Any],
+ tracking_image_size: tuple[int, int],
+ image_path: str,
+ creation_time: Any,
+ frame_idx: int,
+ colors: List[str],
+ ):
+ """Extracts features from detection results for tracking.
+
+ This method resizes masks, extracts region properties using skimage,
+ and compiles a DataFrame of features for each frame, which is then
+ stored internally for later use by the tracking algorithm.
+
+ Args:
+ image: The original image as a numpy array.
+ results: A dictionary containing detection results, including 'masks',
+ 'confidence', 'labels', and 'class_names'.
+ tracking_image_size: The target size (width, height) for resizing masks
+ before feature extraction.
+ image_path: The file path of the image.
+ creation_time: The timestamp of when the image was created.
+ frame_idx: The index of the current frame.
+ colors: A list of color strings corresponding to each detection.
+ """
+ results['resized_masks_for_tracking'] = np.array([
+ cv2.resize(
+ m,
+ tracking_image_size,
+ interpolation=cv2.INTER_NEAREST,
+ )
+ for m in results['masks'].astype('int')
+ ])
+
+ frame_features_list = []
+ for mask in results['resized_masks_for_tracking']:
+ mask = np.where(mask, 1, 0)
+ props = skimage.measure.regionprops_table(
+ mask.astype(np.uint8),
+ intensity_image=image,
+ properties=self._properties,
+ )
+ df = pd.DataFrame(props)
+ frame_features_list.append(df)
+
+ if frame_features_list:
+ frame_df = pd.concat(frame_features_list, ignore_index=True)
+ frame_df.rename(
+ columns={
+ 'centroid-0': 'y',
+ 'centroid-1': 'x',
+ 'bbox-0': 'bbox_0',
+ 'bbox-1': 'bbox_1',
+ 'bbox-2': 'bbox_2',
+ 'bbox-3': 'bbox_3',
+ },
+ inplace=True,
+ )
+
+ frame_df['source_name'] = os.path.basename(os.path.dirname(image_path))
+ frame_df['image_name'] = os.path.basename(image_path)
+ frame_df['creation_time'] = creation_time
+ frame_df['frame'] = frame_idx
+ frame_df['detection_scores'] = results['confidence']
+ frame_df['detection_classes'] = results['labels']
+ frame_df['detection_classes_names'] = results['class_names']
+ frame_df['color'] = colors
+ self.all_detections.append(frame_df)
+ else:
+ self.all_detections.append(pd.DataFrame(columns=self._properties))
+
+ def _select_class_with_model_scores(self, group: pd.DataFrame) -> pd.Series:
+ """Selects the most representative class for a tracked particle.
+
+ This method is used within a groupby operation on 'particle'. It determines
+ the best class for a given particle by first finding the class(es) with the
+ highest frequency. If there's a tie in frequency, it breaks the tie by
+ selecting the class with the highest maximum detection score among the tied
+ classes.
+
+ Args:
+ group: A pandas DataFrame containing all detections associated with a
+ single tracked particle.
+
+ Returns:
+ A pandas Series containing the 'class_id', 'class_name', and
+ 'color_name' of the selected class.
+ """
+ class_counts = group['detection_classes'].value_counts()
+ tied_classes = class_counts[class_counts == class_counts.iloc[0]].index
+
+ max_scores = {
+ cls: group[group['detection_classes'] == cls]['detection_scores'].max()
+ for cls in tied_classes
+ }
+ best_class = max(max_scores.items(), key=lambda x: x[1])[0]
+
+ class_name = group[group['detection_classes'] == best_class][
+ 'detection_classes_names'
+ ].iloc[0]
+ color_name = group[group['detection_classes'] == best_class]['color'].iloc[
+ 0
+ ]
+ return pd.Series({
+ 'class_id': best_class,
+ 'class_name': class_name,
+ 'color_name': color_name,
+ })
+
+ def run_tracking(self) -> pd.DataFrame:
+ """Runs the trackpy linking algorithm on all collected detections.
+
+ This method concatenates all extracted features from multiple frames,
+ applies trackpy's linking to connect detections across frames into tracks
+ (particles), and preserves additional metadata.
+
+ Returns:
+ A pandas DataFrame containing the linked particles, with each row
+ representing a detection instance and including a 'particle' ID.
+ Returns an empty DataFrame if no detections have been collected.
+ """
+ if not self.all_detections:
+ return pd.DataFrame()
+
+ full_df = pd.concat(self.all_detections, ignore_index=True)
+
+ tracking_cols = [
+ 'x',
+ 'y',
+ 'frame',
+ 'bbox_0',
+ 'bbox_1',
+ 'bbox_2',
+ 'bbox_3',
+ 'major_axis_length',
+ 'minor_axis_length',
+ 'perimeter',
+ ]
+
+ track_df = tp.link_df(
+ full_df[tracking_cols],
+ search_range=self.search_range,
+ memory=self.memory,
+ adaptive_stop=self.adaptive_stop,
+ adaptive_step=self.adaptive_step
+ )
+
+ additional_columns = [
+ 'source_name',
+ 'image_name',
+ 'detection_scores',
+ 'detection_classes_names',
+ 'detection_classes',
+ 'color',
+ 'creation_time',
+ ]
+ track_df[additional_columns] = full_df[additional_columns]
+
+ track_df.drop(columns=['frame'], inplace=True)
+ return track_df
+
+ def process_tracking_results(self, track_df):
+ """Aggregates tracking results by particle.
+
+ This method takes the DataFrame with linked particles and aggregates
+ information such as the best class, detection scores, and initial bounding
+ box for each unique particle.
+
+ Args:
+ track_df: A pandas DataFrame containing tracking results, including a
+ 'particle' column generated by trackpy.
+
+ Returns:
+ A pandas DataFrame where each row represents a unique tracked object
+ ('particle'), containing aggregated information.
+ """
+ # Select best class per particle
+ class_info = (
+ track_df.groupby('particle')
+ .apply(self._select_class_with_model_scores, include_groups=False)
+ .reset_index()
+ )
+
+ final_particles = (
+ track_df.groupby('particle')
+ .agg({
+ 'source_name': 'first',
+ 'image_name': 'first',
+ 'detection_scores': 'max',
+ 'creation_time': 'first',
+ 'bbox_0': 'first',
+ 'bbox_1': 'first',
+ 'bbox_2': 'first',
+ 'bbox_3': 'first',
+ })
+ .reset_index()
+ )
+
+ final_particles['detected_classes'] = class_info['class_id']
+ final_particles['detected_classes_names'] = class_info['class_name']
+ final_particles['detected_colors'] = class_info['color_name']
+
+ return final_particles
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/object_tracking_test.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/object_tracking_test.py
new file mode 100644
index 00000000000..55a23dae930
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/object_tracking_test.py
@@ -0,0 +1,277 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+from unittest import mock
+
+import numpy as np
+import pandas as pd
+
+from official.projects.waste_identification_ml.Deploy.detr_cloud_deployment.client import object_tracking
+
+MODULE_PATH = object_tracking.__name__
+
+
+class ObjectTrackerTest(unittest.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ self.tracker = object_tracking.ObjectTracker(
+ search_range=(10, 10),
+ memory=2,
+ adaptive_stop=3,
+ adaptive_step=0.9,
+ )
+
+ def test_init(self):
+ # Assert
+ self.assertEqual(self.tracker.search_range, (10, 10))
+ self.assertEqual(self.tracker.memory, 2)
+ self.assertEqual(self.tracker.adaptive_stop, 3)
+ self.assertEqual(self.tracker.adaptive_step, 0.9)
+ self.assertEqual(self.tracker.all_detections, [])
+
+ @mock.patch(f"{MODULE_PATH}.cv2.resize")
+ @mock.patch(f"{MODULE_PATH}.skimage.measure.regionprops_table")
+ def test_extract_features_for_tracking(self, mock_regionprops, mock_resize):
+ image = np.zeros((100, 100), dtype=np.uint8)
+ masks = np.random.randint(0, 2, size=(2, 50, 50), dtype=np.uint8)
+ results = {
+ "masks": masks,
+ "confidence": [0.9, 0.8],
+ "labels": [1, 2],
+ "class_names": ["class1", "class2"],
+ }
+ tracking_image_size = (20, 20)
+ image_path = "/path/to/image_dir/image1.png"
+ creation_time = "2025-01-01"
+ frame_idx = 0
+ colors = ["red", "blue"]
+
+ mock_resize.return_value = np.ones((20, 20), dtype=np.uint8)
+ mock_regionprops.return_value = {
+ "centroid-0": [10.0],
+ "centroid-1": [10.0],
+ "bbox-0": [5],
+ "bbox-1": [5],
+ "bbox-2": [15],
+ "bbox-3": [15],
+ "area": [100],
+ "convex_area": [100],
+ "bbox_area": [100],
+ "major_axis_length": [10],
+ "minor_axis_length": [10],
+ "eccentricity": [0],
+ "label": [1],
+ "mean_intensity": [0],
+ "max_intensity": [0],
+ "min_intensity": [0],
+ "perimeter": [40],
+ }
+
+ self.tracker.extract_features_for_tracking(
+ image,
+ results,
+ tracking_image_size,
+ image_path,
+ creation_time,
+ frame_idx,
+ colors,
+ )
+
+ self.assertEqual(mock_resize.call_count, 2)
+ self.assertEqual(mock_regionprops.call_count, 2)
+ self.assertEqual(len(self.tracker.all_detections), 1)
+ df = self.tracker.all_detections[0]
+ self.assertIsInstance(df, pd.DataFrame)
+ self.assertEqual(len(df), 2)
+ self.assertIn("y", df.columns)
+ self.assertEqual(df["frame"].iloc[0], 0)
+ self.assertEqual(df["color"].tolist(), ["red", "blue"])
+
+ @mock.patch(f"{MODULE_PATH}.cv2.resize")
+ @mock.patch(f"{MODULE_PATH}.skimage.measure.regionprops_table")
+ def test_extract_features_for_tracking_no_detections(
+ self, mock_regionprops, mock_resize
+ ):
+ image = np.zeros((100, 100), dtype=np.uint8)
+ results = {
+ "masks": np.empty((0, 50, 50)), # No masks
+ "confidence": [],
+ "labels": [],
+ "class_names": [],
+ }
+ tracking_image_size = (20, 20)
+ image_path = "/path/to/image_dir/image1.png"
+ creation_time = "2025-01-01"
+ frame_idx = 0
+ colors = []
+
+ self.tracker.extract_features_for_tracking(
+ image,
+ results,
+ tracking_image_size,
+ image_path,
+ creation_time,
+ frame_idx,
+ colors,
+ )
+
+ mock_resize.assert_not_called()
+ mock_regionprops.assert_not_called()
+ self.assertEqual(len(self.tracker.all_detections), 1)
+ df = self.tracker.all_detections[0]
+ self.assertTrue(df.empty)
+ self.assertListEqual(list(df.columns), list(self.tracker._properties))
+
+ def test_run_tracking_empty(self):
+ self.tracker.all_detections = []
+
+ result_df = self.tracker.run_tracking()
+
+ self.assertTrue(result_df.empty)
+
+ @mock.patch(f"{MODULE_PATH}.tp.link_df")
+ def test_run_tracking(self, mock_link_df):
+
+ data1 = {
+ "x": [10],
+ "y": [10],
+ "frame": [0],
+ "bbox_0": [5],
+ "bbox_1": [5],
+ "bbox_2": [15],
+ "bbox_3": [15],
+ "major_axis_length": [10],
+ "minor_axis_length": [10],
+ "perimeter": [40],
+ "source_name": ["d1"],
+ "image_name": ["i1"],
+ "detection_scores": [0.9],
+ "detection_classes_names": ["c1"],
+ "detection_classes": [1],
+ "color": ["red"],
+ "creation_time": ["t1"],
+ }
+ data2 = {
+ "x": [12],
+ "y": [12],
+ "frame": [1],
+ "bbox_0": [7],
+ "bbox_1": [7],
+ "bbox_2": [17],
+ "bbox_3": [17],
+ "major_axis_length": [10],
+ "minor_axis_length": [10],
+ "perimeter": [40],
+ "source_name": ["d1"],
+ "image_name": ["i2"],
+ "detection_scores": [0.95],
+ "detection_classes_names": ["c1"],
+ "detection_classes": [1],
+ "color": ["red"],
+ "creation_time": ["t2"],
+ }
+ df1 = pd.DataFrame(data1)
+ df2 = pd.DataFrame(data2)
+ self.tracker.all_detections = [df1, df2]
+
+ linked_df = pd.concat([df1, df2], ignore_index=True)
+ linked_df["particle"] = 0
+ mock_link_df.return_value = linked_df
+
+ result_df = self.tracker.run_tracking()
+
+ mock_link_df.assert_called_once()
+ self.assertIsInstance(result_df, pd.DataFrame)
+ self.assertIn("particle", result_df.columns)
+ self.assertNotIn("frame", result_df.columns)
+ self.assertEqual(len(result_df), 2)
+ self.assertEqual(result_df["particle"].iloc[0], 0)
+
+ def test_process_tracking_results(self):
+
+ track_data = {
+ "x": [10, 12],
+ "y": [10, 12],
+ "bbox_0": [5, 7],
+ "bbox_1": [5, 7],
+ "bbox_2": [15, 17],
+ "bbox_3": [15, 17],
+ "major_axis_length": [10, 10],
+ "minor_axis_length": [10, 10],
+ "perimeter": [40, 40],
+ "particle": [0, 0],
+ "source_name": ["d1", "d1"],
+ "image_name": ["i1", "i2"],
+ "detection_scores": [0.9, 0.8],
+ "detection_classes_names": ["apple", "banana"],
+ "detection_classes": [1, 2],
+ "color": ["red", "yellow"],
+ "creation_time": ["t1", "t2"],
+ }
+ track_df = pd.DataFrame(track_data)
+
+ final_df = self.tracker.process_tracking_results(track_df)
+
+ self.assertIsInstance(final_df, pd.DataFrame)
+ self.assertEqual(len(final_df), 1) # 1 particle
+ self.assertEqual(final_df["particle"].tolist(), [0])
+ particle0 = final_df[final_df["particle"] == 0].iloc[0]
+ self.assertEqual(particle0["detected_classes"], 1)
+ self.assertEqual(particle0["detected_classes_names"], "apple")
+ self.assertEqual(particle0["detected_colors"], "red")
+ self.assertEqual(particle0["detection_scores"], 0.9)
+
+ def test_select_class_with_scores_tie_break(self):
+ group_data = {
+ "detection_classes": [1, 2, 1, 2],
+ "detection_scores": [0.8, 0.9, 0.7, 0.85],
+ "detection_classes_names": ["c1", "c2", "c1", "c2"],
+ "color": ["red", "blue", "red", "blue"],
+ }
+ group_df = pd.DataFrame(group_data)
+ result = self.tracker._select_class_with_model_scores(group_df)
+
+ self.assertEqual(result["class_name"], "c2")
+
+ def test_select_class_with_higher_frequency(self):
+ group_data = {
+ "detection_classes": [1, 2, 1, 1],
+ "detection_scores": [0.8, 0.9, 0.7, 0.85],
+ "detection_classes_names": ["c1", "c2", "c1", "c1"],
+ "color": ["red", "blue", "red", "red"],
+ }
+ group_df = pd.DataFrame(group_data)
+ result = self.tracker._select_class_with_model_scores(group_df)
+
+ self.assertEqual(result["class_name"], "c1")
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/requirements.sh b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/requirements.sh
new file mode 100644
index 00000000000..8fc226bbc47
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/requirements.sh
@@ -0,0 +1,47 @@
+#!/bin/bash
+
+# Summary
+cat << EOF
+This script sets up the environment for running ML models by ensuring Bash
+execution, installing system dependencies, setting up a virtual environment,
+installing ML packages, and cloning TensorFlow Model Garden.
+EOF
+
+# Ensure the script is executed with /bin/bash
+if [ -z "$BASH_VERSION" ]; then
+ exec /bin/bash "$0" "$@"
+fi
+
+sudo apt-get update -y
+
+# Install Docker if not already installed
+if ! command -v docker &> /dev/null
+then
+ echo "Docker is not installed. Installing Docker..."
+ curl -fsSL https://get.docker.com -o get-docker.sh
+ sudo sh get-docker.sh
+ rm -f get-docker.sh
+ echo "Docker installation completed."
+else
+ echo "Docker is already installed. Skipping Docker installation."
+fi
+
+# Create a virtual environment and install packages
+sudo apt-get install -y python3-venv python3-pip
+
+python3.10 -m venv myenv
+source myenv/bin/activate
+
+echo "Activated python environment, installing dependencies."
+
+pip install -r requirements.txt
+
+# Clone TensorFlow Model Garden if the 'models' directory does not exist
+if [ ! -d "models" ]; then
+ git clone --depth 1 https://github.com/tensorflow/models.git
+else
+ echo "'models' directory already exists. Skipping cloning."
+fi
+
+deactivate
+echo "Environment setup is complete."
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/requirements.txt b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/requirements.txt
new file mode 100644
index 00000000000..cfc0cb282f9
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/requirements.txt
@@ -0,0 +1,21 @@
+--extra-index-url https://pypi.nvidia.com
+
+natsort==8.4.0
+absl-py==2.4.0
+opencv-python==4.13.0.92
+pandas==2.3.3
+pandas-gbq==0.33.0
+google-cloud-bigquery==3.40.1
+google-auth==2.48.0
+google-cloud-storage==3.9.0
+scikit-image==0.25.2
+scikit-learn==1.7.2
+webcolors==1.13
+ffmpeg-python==0.2.0
+tritonclient[all]==2.65.0
+supervision==0.26.1
+pillow==12.0.0
+trackpy==0.7
+cupy-cuda12x[cuda_dlls]
+cupy-cuda12x[ctk]
+cuml-cu12==25.12.*
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/run_images.sh b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/run_images.sh
new file mode 100644
index 00000000000..8aa22db7d13
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/run_images.sh
@@ -0,0 +1,50 @@
+#!/bin/bash
+
+cat << EOF
+This script automates the execution of an Circularnet pipeline for image
+processing.
+Steps Performed:
+ 1. Activates the Python virtual environment named 'myenv'.
+ 2. Validates successful activation of the virtual environment.
+ 3. Executes the 'pipeline_images.py' script with the following parameters:
+ Parameters:
+ --input_directory : GCS directory where the input images are stored for
+ inference.
+ --output_directory : GCS directory where the model inference outputs will be
+ saved.
+ --model_name : Name of the model to download and use for inference.
+ --threshold : Confidence threshold for detections during
+ inference.
+ --search_range_x : Max pixel movement allowed in the X direction for
+ object tracking between missed frames.
+ --search_range_y : Max pixel movement allowed in the Y direction for
+ object tracking between missed frames.
+ --memory : Number of frames an object can be missed and still
+ be tracked.
+ --project_id : Google Cloud Project ID for BigQuery operations.
+ --bq_dataset_id : BigQuery Dataset ID where results will be stored.
+ --bq_table_id : BigQuery Table ID where results will be stored.
+ --overwrite : If set to True, overwrites the pre-existing
+ BigQuery table.
+EOF
+#Activate the virtual environment
+source myenv/bin/activate
+# Check if the virtual environment is activated
+if [[ "$VIRTUAL_ENV" != "" ]]; then
+ echo "Virtual environment 'myenv' activated successfully."
+else
+ echo "Failed to activate virtual environment. Exiting."
+ exit 1
+fi
+python inference_pipeline.py \
+ --input_directory=gs://recykal/TestData/SmallTestData \
+ --output_directory=gs://recykal/TestData/SmallTestData \
+ --model_name=cn_segmentation_trt_model \
+ --threshold=0.50 \
+ --search_range_x=150 \
+ --search_range_y=20 \
+ --memory=3 \
+ --project_id=waste-identification-ml-330916 \
+ --bq_dataset_id=circularnet_dataset \
+ --bq_table_id=test_table1 \
+ --overwrite=True
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/triton_server_inference.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/triton_server_inference.py
new file mode 100644
index 00000000000..a010ab4a3b8
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/triton_server_inference.py
@@ -0,0 +1,271 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Prediction from the Triton server."""
+
+import os
+from typing import Any, Dict, List, Tuple
+import cv2
+import numpy as np
+import pandas as pd
+import tritonclient.http as httpclient
+
+
+def _sigmoid(x: np.ndarray) -> np.ndarray:
+ """Applies sigmoid function to an array of scores."""
+ return 1 / (1 + np.exp(-x))
+
+
+def _box_cxcywh_to_xyxyn(x: np.ndarray) -> np.ndarray:
+ """Converts bounding boxes from cxcywh format to xyxyn format."""
+ cx, cy, w, h = x[..., 0], x[..., 1], x[..., 2], x[..., 3]
+ xmin = cx - w / 2
+ ymin = cy - h / 2
+ xmax = cx + w / 2
+ ymax = cy + h / 2
+ return np.stack([xmin, ymin, xmax, ymax], axis=-1)
+
+
+class TritonObjectDetector:
+ """Client for performing object detection inference using a Triton server.
+
+ This class handles preprocessing, making inference requests to a Triton HTTP
+ server, and post-processing the results, including scaling bounding boxes
+ and masks.
+ """
+
+ def __init__(
+ self,
+ server_url: str = 'localhost:8000',
+ model_name: str = 'detection_model',
+ input_size: tuple[int, int] = (432, 432),
+ verbose: bool = False,
+ ):
+ """Initializes the Triton Client and Model configuration."""
+ self.client = httpclient.InferenceServerClient(
+ url=server_url, verbose=verbose
+ )
+ self.model_name = model_name
+ self.input_size = input_size
+
+ # Normalization constants
+ self.means = np.array([0.485, 0.456, 0.406], dtype=np.float32)
+ self.stds = np.array([0.229, 0.224, 0.225], dtype=np.float32)
+
+ def _resize_mask_batch(
+ self, masks: np.ndarray, target_dims: Tuple[int, int]
+ ) -> np.ndarray:
+ """Resizes a batch of masks to the target dimensions."""
+ target_w, target_h = target_dims
+ masks_transposed = np.transpose(masks, (1, 2, 0))
+
+ resized_batch = cv2.resize(
+ masks_transposed, (target_w, target_h), interpolation=cv2.INTER_NEAREST
+ )
+
+ # If N=1, cv2.resize might drop the last dim, so we ensure 3D
+ if resized_batch.ndim == 2:
+ return resized_batch[np.newaxis, ...]
+
+ return np.transpose(resized_batch, (2, 0, 1))
+
+ def _scale_bbox_and_masks(
+ self, results: Dict[str, Any], target_dims: Tuple[int, int]
+ ) -> Dict[str, Any]:
+ """Scales normalized boxes and small mask logits to target dimensions."""
+ target_w, target_h = target_dims
+
+ # Scale Bounding Boxes
+ results['xyxy'][..., [0, 2]] *= target_w
+ results['xyxy'][..., [1, 3]] *= target_h
+
+ # Scale Masks
+ if results['masks'] is not None:
+ rescaled_masks = self._resize_mask_batch(
+ results['masks'], target_dims
+ )
+ results['masks'] = (rescaled_masks > 0).astype(bool)
+
+ return results
+
+ def _get_input_batch_for_inference(self, image_path: str) -> np.ndarray:
+ """Preprocesses an image for Triton inference.
+
+ Loads an image, resizes it, converts it to RGB, normalizes pixel values,
+ and transposes it to the channel-first format expected by the model.
+
+ Args:
+ image_path: The path to the input image file.
+
+ Returns:
+ A numpy array representing the preprocessed image, ready for inference.
+
+ Raises:
+ FileNotFoundError: If the image file does not exist.
+ """
+ original_image = cv2.imread(image_path)
+ if original_image is None:
+ raise FileNotFoundError(f'Image not found at {image_path}')
+
+ rgb_image = cv2.cvtColor(original_image, cv2.COLOR_BGR2RGB)
+ resized_image = cv2.resize(
+ rgb_image, self.input_size, interpolation=cv2.INTER_AREA
+ )
+
+ # Normalize: (pixel / 255 - mean) / std
+ float_image = resized_image.astype(np.float32) / 255.0
+ normalized_image = (float_image - self.means) / self.stds
+
+ # Transpose to CHW and add batch dimension
+ transposed_image = np.transpose(normalized_image, (2, 0, 1))
+ batched_image = np.expand_dims(transposed_image, axis=0).astype(np.float32)
+ return batched_image
+
+ def _reformat_triton_output_to_dict(
+ self,
+ outputs: List[np.ndarray],
+ confidence_threshold: float,
+ max_boxes: int,
+ ) -> Dict[str, Any]:
+ """Reformats and filters the raw outputs from the Triton server.
+
+ Args:
+ outputs: A list of numpy arrays containing the raw outputs from the Triton
+ model. Expected to contain [boxes, probabilities, masks (optional)].
+ confidence_threshold: Boxes with a confidence score below this threshold
+ will be filtered out.
+ max_boxes: The maximum number of top-scoring boxes to consider before
+ applying the confidence threshold.
+
+ Returns:
+ A dict containing arrays of detection results for keys 'confidence',
+ 'labels', 'xyxy', and 'masks'. Bounding boxes are in [xmin, ymin,
+ xmax, ymax] format. Masks are `None` if not present in model output.
+ For example:
+
+ {
+ 'confidence': np.array([0.9, 0.8]),
+ 'labels': np.array([1, 2]),
+ 'xyxy': np.array([[0.1, 0.2, 0.3, 0.4], [0.5, 0.6, 0.7, 0.8]]),
+ 'masks': np.array([mask1_data, mask2_data]) or None,
+ }
+ """
+ raw_boxes = outputs[0].squeeze()
+ raw_probs = _sigmoid(outputs[1])
+ masks = outputs[2].squeeze() if len(outputs) == 3 else None
+
+ scores = np.max(raw_probs, axis=2).squeeze()
+ labels = np.argmax(raw_probs, axis=2).squeeze()
+
+ # Filter by top-k bounding boxes
+ sorted_idx = np.argsort(scores)[::-1][:max_boxes]
+
+ # Filter by confidence score
+ confidence_mask_filter = scores[sorted_idx] > confidence_threshold
+ final_idx = sorted_idx[confidence_mask_filter]
+
+ return {
+ 'confidence': scores[final_idx],
+ 'labels': labels[final_idx],
+ 'xyxy': _box_cxcywh_to_xyxyn(raw_boxes[final_idx]),
+ 'masks': masks[final_idx] if masks is not None else None,
+ }
+
+ def predict(
+ self,
+ image_path: str,
+ confidence_threshold: float = 0.5,
+ max_boxes: int = 100,
+ output_dims: tuple[int, int] = (1024, 1024),
+ ) -> Dict[str, Any]:
+ """Performs inference on a single image using the Triton server.
+
+ Args:
+ image_path: The path to the input image file.
+ confidence_threshold: Boxes with a confidence score below this threshold
+ will be filtered out.
+ max_boxes: The maximum number of top-scoring boxes to consider before
+ applying the confidence threshold.
+ output_dims: The dimensions (width, height) to which bounding boxes and
+ masks should be scaled in the output.
+
+ Returns:
+ A dictionary containing the inference results:
+ - 'confidence': A numpy array of confidence scores.
+ - 'labels': A numpy array of predicted class labels (integer IDs).
+ - 'xyxy': A numpy array of bounding boxes in [xmin, ymin, xmax, ymax]
+ format, scaled to `output_dims`.
+ - 'masks': A numpy array of boolean masks, rescaled to `output_dims`,
+ or None if masks are not part of the model output.
+ """
+
+ # Preprocessing
+ input_data = self._get_input_batch_for_inference(image_path)
+
+ # Prepare Triton Input
+ infer_input = httpclient.InferInput(
+ 'input', input_data.shape, datatype='FP32'
+ )
+ infer_input.set_data_from_numpy(input_data, binary_data=True)
+
+ # Execute Inference
+ response = self.client.infer(
+ model_name=self.model_name, inputs=[infer_input]
+ )
+
+ # Extract results based on known output names
+ raw_outputs = [
+ response.as_numpy('dets'),
+ response.as_numpy('labels'),
+ response.as_numpy('masks'),
+ ]
+
+ # Reformat Triton output
+ results = self._reformat_triton_output_to_dict(
+ raw_outputs, confidence_threshold, max_boxes
+ )
+
+ if results['labels'].size != 0:
+ # Scale to output dimensions
+ results = self._scale_bbox_and_masks(results, output_dims)
+
+ return results
+
+ def _get_class_id_to_class_name_mapping(self):
+ """Returns a mapping from class ID to class name."""
+ labels_path = os.path.join(os.getcwd(), 'labels50.csv')
+ labels_df = pd.read_csv(labels_path)
+ class_id_to_class_name_mapper = labels_df.set_index('id').to_dict()['names']
+ return class_id_to_class_name_mapper
+
+ def get_class_names(self, results):
+ """Returns the class names for the given results."""
+ class_name_mapper = self._get_class_id_to_class_name_mapping()
+ label_names = np.array(
+ [class_name_mapper.get(c + 1, 'None') for c in results['labels']]
+ )
+ return label_names
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/triton_server_inference_test.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/triton_server_inference_test.py
new file mode 100644
index 00000000000..d93c854151a
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/triton_server_inference_test.py
@@ -0,0 +1,271 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+from unittest import mock
+
+import numpy as np
+import pandas as pd
+
+from official.projects.waste_identification_ml.Deploy.detr_cloud_deployment.client import triton_server_inference
+
+MODULE_PATH = triton_server_inference.__name__
+
+
+class TritonObjectDetectorTest(unittest.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ # Mock triton client
+ self.mock_http_client_patch = mock.patch(
+ f"{MODULE_PATH}.httpclient.InferenceServerClient"
+ )
+ self.mock_http_client = self.mock_http_client_patch.start()
+ self.mock_client_instance = self.mock_http_client.return_value
+
+ # Mock pd.read_csv for labels
+ self.mock_pd_read_csv_patch = mock.patch(f"{MODULE_PATH}.pd.read_csv")
+ self.mock_pd_read_csv = self.mock_pd_read_csv_patch.start()
+ self.mock_pd_read_csv.return_value = pd.DataFrame(
+ {"id": [1, 2], "names": ["Class1", "Class2"]}
+ )
+
+ # Mock os.path.join to avoid filesystem access for labels
+ self.mock_os_path_join_patch = mock.patch(f"{MODULE_PATH}.os.path.join")
+ self.mock_os_path_join = self.mock_os_path_join_patch.start()
+ self.mock_os_path_join.return_value = "dummy_path/labels50.csv"
+
+ self.detector = triton_server_inference.TritonObjectDetector(
+ server_url="test:8000", model_name="test_model", input_size=(100, 100)
+ )
+
+ def tearDown(self):
+ super().tearDown()
+ mock.patch.stopall()
+
+ def test_init(self):
+ # Assert
+ self.mock_http_client.assert_called_once_with(
+ url="test:8000", verbose=False
+ )
+ self.assertEqual(self.detector.model_name, "test_model")
+ self.assertEqual(self.detector.input_size, (100, 100))
+
+ def test_sigmoid(self):
+ # Arrange
+ x = np.array([-1.0, 0.0, 1.0])
+
+ # Act
+ result = triton_server_inference._sigmoid(x)
+
+ # Assert
+ np.testing.assert_allclose(result, [0.26894142, 0.5, 0.73105858])
+
+ def test_box_cxcywh_to_xyxyn(self):
+ # Arrange
+ x = np.array([[0.5, 0.5, 0.2, 0.2], [0.5, 0.5, 1.0, 1.0]])
+
+ # Act
+ result = triton_server_inference._box_cxcywh_to_xyxyn(x)
+
+ # Assert
+ expected = np.array([[0.4, 0.4, 0.6, 0.6], [0.0, 0.0, 1.0, 1.0]])
+ np.testing.assert_allclose(result, expected)
+
+ @mock.patch(f"{MODULE_PATH}.cv2.resize")
+ def test_scale_bbox_and_masks(self, mock_cv2_resize):
+ # Arrange
+ mock_cv2_resize.return_value = np.ones(
+ (20, 10)
+ ) # h=20, w=10. cv2 resize returns h, w
+
+ results = {
+ "xyxy": np.array([[0.1, 0.1, 0.5, 0.5]]),
+ "masks": np.ones((1, 5, 5), dtype=np.uint8),
+ }
+ target_dims = (10, 20) # w=10, h=20
+
+ # Act
+ scaled_results = self.detector._scale_bbox_and_masks(results, target_dims)
+
+ # Assert
+ np.testing.assert_allclose(scaled_results["xyxy"], [[1.0, 2.0, 5.0, 10.0]])
+ self.assertEqual(scaled_results["masks"].shape, (1, 20, 10))
+ self.assertEqual(scaled_results["masks"].dtype, bool)
+ mock_cv2_resize.assert_called_once()
+ args, kwargs = mock_cv2_resize.call_args
+ np.testing.assert_array_equal(args[0], np.ones((5, 5, 1)))
+ self.assertEqual(args[1], (10, 20))
+ self.assertEqual(
+ kwargs["interpolation"], triton_server_inference.cv2.INTER_NEAREST
+ )
+
+ @mock.patch(f"{MODULE_PATH}.cv2.imread")
+ @mock.patch(f"{MODULE_PATH}.cv2.cvtColor")
+ @mock.patch(f"{MODULE_PATH}.cv2.resize")
+ def test_get_input_batch_for_inference_success(
+ self, mock_resize, mock_cvtcolor, mock_imread
+ ):
+ # Arrange
+ mock_imread.return_value = np.zeros((200, 200, 3), dtype=np.uint8)
+ mock_cvtcolor.return_value = np.zeros((200, 200, 3), dtype=np.uint8)
+ mock_resize.return_value = np.zeros((100, 100, 3), dtype=np.uint8)
+ image_path = "dummy.jpg"
+
+ # Act
+ processed_image = self.detector._get_input_batch_for_inference(image_path)
+
+ # Assert
+ mock_imread.assert_called_once_with(image_path)
+ mock_cvtcolor.assert_called_once()
+ mock_resize.assert_called_once_with(
+ mock_cvtcolor.return_value, (100, 100), interpolation=mock.ANY
+ )
+ self.assertEqual(processed_image.shape, (1, 3, 100, 100))
+ self.assertEqual(processed_image.dtype, np.float32)
+
+ @mock.patch(f"{MODULE_PATH}.cv2.imread")
+ def test_get_input_batch_for_inference_filenotfound(self, mock_imread):
+ # Arrange
+ mock_imread.return_value = None
+ image_path = "nonexistent.jpg"
+
+ # Act & Assert
+ with self.assertRaises(FileNotFoundError):
+ self.detector._get_input_batch_for_inference(image_path)
+
+ def test_reformat_triton_output_to_dict(self):
+
+ raw_boxes = np.array([[[0.5, 0.5, 0.2, 0.2], [0.1, 0.1, 0.1, 0.1]]])
+ raw_probs_logits = np.array([[[-1.0, 3.0, 0.0], [-2.0, -3.0, -4.0]]])
+ masks = np.zeros((1, 2, 10, 10))
+ outputs = [raw_boxes, raw_probs_logits, masks]
+
+ # Act
+ # max_boxes=2, threshold=0.5 -> box 0 should be kept, box 1 filtered
+ results = self.detector._reformat_triton_output_to_dict(
+ outputs, confidence_threshold=0.5, max_boxes=2
+ )
+
+ # Assert
+ expected_results = {
+ "confidence": np.array([0.95257413]),
+ "labels": np.array([1]),
+ "xyxy": np.array([[0.4, 0.4, 0.6, 0.6]]),
+ "masks": np.zeros((1, 10, 10)),
+ }
+ self.assertCountEqual(results.keys(), expected_results.keys())
+ np.testing.assert_allclose(
+ results["confidence"], expected_results["confidence"]
+ )
+ np.testing.assert_array_equal(results["labels"], expected_results["labels"])
+ np.testing.assert_allclose(results["xyxy"], expected_results["xyxy"])
+ np.testing.assert_array_equal(results["masks"], expected_results["masks"])
+
+ @mock.patch(
+ f"{MODULE_PATH}.TritonObjectDetector._get_input_batch_for_inference"
+ )
+ @mock.patch(
+ f"{MODULE_PATH}.TritonObjectDetector._reformat_triton_output_to_dict"
+ )
+ @mock.patch(f"{MODULE_PATH}.TritonObjectDetector._scale_bbox_and_masks")
+ @mock.patch(f"{MODULE_PATH}.httpclient.InferInput")
+ def test_predict(
+ self,
+ mock_infer_input,
+ mock_scale,
+ mock_reformat_triton_output_to_dict,
+ mock_get_input_batch_for_inference,
+ ):
+ # Arrange
+ image_path = "dummy.jpg"
+ mock_get_input_batch_for_inference.return_value = np.zeros(
+ (1, 3, 100, 100), dtype=np.float32
+ )
+ mock_input_instance = mock.Mock()
+ mock_infer_input.return_value = mock_input_instance
+
+ mock_response = mock.Mock()
+ mock_response.as_numpy.side_effect = [
+ np.array([1]),
+ np.array([2]),
+ np.array([3]),
+ ]
+ self.mock_client_instance.infer.return_value = mock_response
+
+ post_process_result = {"xyxy": np.array([[1, 1, 2, 2]]), "masks": None}
+ mock_reformat_triton_output_to_dict.return_value = post_process_result
+ scale_result = {
+ "xyxy": np.array([[10, 10, 20, 20]]),
+ "masks": None,
+ "some_key": 1,
+ }
+ mock_scale.return_value = scale_result
+
+ # Act
+ results = self.detector.predict(
+ image_path,
+ confidence_threshold=0.7,
+ max_boxes=50,
+ output_dims=(200, 200),
+ )
+
+ # Assert
+ mock_get_input_batch_for_inference.assert_called_once_with(image_path)
+ mock_infer_input.assert_called_once_with(
+ "input", (1, 3, 100, 100), datatype="FP32"
+ )
+ mock_input_instance.set_data_from_numpy.assert_called_once_with(
+ mock_get_input_batch_for_inference.return_value, binary_data=True
+ )
+ self.mock_client_instance.infer.assert_called_once_with(
+ model_name="test_model", inputs=[mock_input_instance]
+ )
+ mock_reformat_triton_output_to_dict.assert_called_once_with(
+ [np.array([1]), np.array([2]), np.array([3])], 0.7, 50
+ )
+ mock_scale.assert_called_once_with(post_process_result, (200, 200))
+ self.assertEqual(results, scale_result)
+
+ def test_get_class_id_to_class_name_mapping(self):
+ mapper = self.detector._get_class_id_to_class_name_mapping()
+
+ self.mock_pd_read_csv.assert_called_once_with("dummy_path/labels50.csv")
+
+ self.assertEqual(mapper, {1: "Class1", 2: "Class2"})
+
+ def test_get_class_names(self):
+ results = {"labels": np.array([0, 1, 99])}
+
+ class_names = self.detector.get_class_names(results)
+
+ np.testing.assert_array_equal(class_names, ["Class1", "Class2", "None"])
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/utils.py b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/utils.py
new file mode 100644
index 00000000000..97e0cdd9b05
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/utils.py
@@ -0,0 +1,312 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Utility functions for the pipeline."""
+
+import datetime
+import logging
+import os
+import pathlib
+import re
+import subprocess
+import time
+
+import color_extraction
+import natsort
+import numpy as np
+import pandas as pd
+import PIL
+from PIL import Image
+import supervision as sv
+
+_DETECTION_COLOR_PALETTE = sv.ColorPalette.from_hex([
+ '#ffff00',
+ '#ff9b00',
+ '#ff66ff',
+ '#3399ff',
+ '#ff66b2',
+ '#ff8080',
+ '#b266ff',
+ '#9999ff',
+ '#66ffff',
+ '#33ff99',
+ '#66ff66',
+ '#99ff00',
+])
+
+
+def _create_log_file(name: str, logs_folder_path: str) -> logging.Logger:
+ """Creates a logger and a log file given the name of the video.
+
+ Args:
+ name: The name of the video.
+ logs_folder_path: Path to the directory where logs should be saved.
+
+ Returns:
+ logging.Logger: Logger object configured to write logs to the file.
+ """
+ log_file_path = os.path.join(logs_folder_path, f'{name}.log')
+ logger = logging.getLogger(name)
+ logger.setLevel(logging.INFO)
+ file_handler = logging.FileHandler(log_file_path)
+ formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
+ file_handler.setFormatter(formatter)
+ logger.addHandler(file_handler)
+ return logger
+
+
+def setup_logger_and_directories(input_dir):
+ """Sets up directories and a logger for the pipeline.
+
+ This function copies the input directory from GCP, creates a prediction
+ folder, and initializes a logger for the current run.
+
+ Args:
+ input_dir: The path to the input directory on GCP.
+
+ Returns:
+ A tuple containing:
+ - input_directory: The local path of the copied input directory.
+ - prediction_folder: The path to the created prediction folder.
+ - logger: The configured logging.Logger object.
+ """
+
+ input_directory = (input_dir).rstrip('/\\')
+ local_dir = os.path.basename(input_directory)
+ # A failed previous run can leave a file where gcloud needs a directory.
+ if os.path.isfile(local_dir):
+ os.remove(local_dir)
+ command = f'gcloud storage cp --recursive {input_directory} .'
+ subprocess.run(command, shell=True, check=True)
+ prediction_folder = os.path.basename(input_directory) + '_prediction'
+ os.makedirs(prediction_folder, exist_ok=True)
+ log_name = os.path.basename(input_dir)
+ log_folder = os.path.join(os.getcwd(), 'logs')
+ os.makedirs(log_folder, exist_ok=True)
+ logger = _create_log_file(log_name, log_folder)
+ return input_directory, prediction_folder, logger
+
+
+def get_class_id_to_class_name_mapping(label_csv):
+ labels_path = os.path.join(os.getcwd(), label_csv)
+ labels_df = pd.read_csv(labels_path)
+ class_id_to_class_name_mapper = labels_df.set_index('id').to_dict()['names']
+ return class_id_to_class_name_mapper
+
+
+def get_image_capture_time(image_path):
+ """Retrieves the creation time of an image, trying multiple methods.
+
+ Args:
+ image_path: The path to the image file.
+
+ Returns:
+ A string representing the creation time in the format "%Y-%m-%d %H:%M:%S" if
+ found, otherwise returns "Creation time not found".
+ """
+ try:
+ filename_str = os.path.basename(image_path)
+ parts = filename_str.split('_')
+ date_part = parts[-2]
+ time_part = parts[-1].split('.')[0]
+
+ # Format Date: YYYY-MM-DD
+ formatted_date = f'{date_part[:4]}-{date_part[4:6]}-{date_part[6:]}'
+
+ # Format Time: HH:MM:SS.ms (Taking the last 2 digits as the decimal)
+ formatted_time = (
+ f'{time_part[:2]}:{time_part[2:4]}:{time_part[4:6]}.{time_part[6:]}'
+ )
+
+ return datetime.datetime.strptime(
+ f'{formatted_date} {formatted_time}', '%Y-%m-%d %H:%M:%S.%f'
+ ).strftime('%Y-%m-%d %H:%M:%S')
+
+ except (IndexError, ValueError):
+ try:
+
+ # 1. Try EXIF data (if available)
+ image = Image.open(image_path)
+ exif_data = image.getexif()
+ if exif_data:
+ datetime_tag_id = 36867 # Tag ID for "DateTimeOriginal"
+ datetime_str = exif_data.get(datetime_tag_id)
+ if datetime_str:
+ return datetime.datetime.strptime(
+ datetime_str, '%Y:%m:%d %H:%M:%S'
+ ).strftime('%Y-%m-%d %H:%M:%S')
+
+ # 2. Try file modification time (less accurate, but better than nothing)
+ file_modified_time = os.path.getmtime(image_path)
+ return datetime.datetime.fromtimestamp(file_modified_time).strftime(
+ '%Y-%m-%d %H:%M:%S'
+ )
+ except FileNotFoundError:
+ return 'Image not found'
+ except PIL.UnidentifiedImageError as e:
+ return f'Error: {e}'
+ except (OSError, PIL.ImageError) as e:
+ return f'Error processing image or file: {e}'
+
+
+def parse_imgage_names_to_datetime(filename):
+ stem = pathlib.Path(filename).stem
+ m = re.match(r'img_(\d{8})_(\d{6})(\d*)', stem)
+ dt_str = f'{m.group(1)}_{m.group(2)}{m.group(3)}'
+ return datetime.datetime.strptime(dt_str, '%Y%m%d_%H%M%S%f')
+
+
+def files_paths(folder_path):
+ """List the full paths of image files in a folder and sort them.
+
+ Args:
+ folder_path: The path of the folder to list the image files from.
+
+ Returns:
+ A list of full paths of the image files in the folder, sorted in ascending
+ order.
+ """
+ img_extensions = ('.jpg', '.jpeg', '.png', '.gif', '.bmp', '.tiff', '.webp')
+ image_files_full_path = []
+ for entry in os.scandir(folder_path):
+ if entry.is_file() and entry.name.lower().endswith(img_extensions):
+ image_files_full_path.append(entry.path)
+
+ # Sort the list of files by name
+ try:
+ image_files_full_path = sorted(
+ image_files_full_path, key=parse_imgage_names_to_datetime
+ )
+ except AttributeError:
+ print('Failed to parse image names to datetime, using natsort instead.')
+ image_files_full_path = natsort.natsorted(image_files_full_path)
+ return image_files_full_path
+
+
+def _crop_objects_from_masks(
+ image: np.ndarray, masks: np.ndarray
+) -> list[np.ndarray]:
+ """Crops objects from an image based on provided masks.
+
+ Args:
+ image: The image as a numpy array.
+ masks: A numpy array of binary masks, where each slice corresponds to an
+ object.
+
+ Returns:
+ A list of numpy arrays, where each array is a cropped object.
+ """
+ cropped_objects = [
+ np.where(np.expand_dims(m, -1), image, 0) for m in masks.astype(int)
+ ]
+
+ return cropped_objects
+
+
+def extract_color_names(results, image_for_saving):
+ """Extracts generic color names from detected objects in an image.
+
+ This function uses the provided masks in `results` to crop each detected
+ object from `image_for_saving`. It then finds the dominant color for each
+ cropped object and returns a list of generic color names.
+
+ Args:
+ results: A dictionary containing detection results, including 'masks' (numpy
+ array of masks).
+ image_for_saving: The image as a numpy array from which objects are cropped.
+
+ Returns:
+ A list of strings, where each string is a generic color name
+ corresponding to a detected object.
+ """
+ # Crop objects from an image using masks for color detection.
+ cropped_objects = _crop_objects_from_masks(
+ image_for_saving, results['masks']
+ )
+ # Perform color detection using clustering approach.
+ dominant_colors = [
+ *map(color_extraction.find_dominant_color, cropped_objects)
+ ]
+ generic_color_names = color_extraction.get_generic_color_name(dominant_colors)
+ return generic_color_names
+
+
+def draw_detections_and_save_image(img, results, save_path):
+ """Used for plotting the annotations on the image.
+
+ Args:
+ img: The PIL Image object to plot annotations on.
+ results: A dictionary containing detection results.
+ save_path: The file path to save the annotated image.
+
+ Returns:
+ A PIL Image object with the annotations plotted.
+ """
+
+ detection = sv.Detections(
+ xyxy=results['xyxy'],
+ mask=results['masks'],
+ confidence=results['confidence'],
+ class_id=results['labels'],
+ data={'class_names': results['class_names']},
+ )
+
+ text_scale = sv.calculate_optimal_text_scale(resolution_wh=img.size)
+ thickness = sv.calculate_optimal_line_thickness(resolution_wh=img.size)
+ color = _DETECTION_COLOR_PALETTE
+ detections_labels = [
+ f'{class_name} : {probability:.2f}'
+ for class_name, probability in zip(
+ detection.data['class_names'], detection.confidence
+ )
+ ]
+ detections_image = sv.MaskAnnotator(opacity=0.4).annotate(
+ scene=img.copy(), detections=detection
+ )
+ detections_image = sv.BoxAnnotator(color=color, thickness=thickness).annotate(
+ detections_image, detection
+ )
+ detections_image = sv.LabelAnnotator(
+ color=color, text_color=sv.Color.BLACK, text_scale=text_scale
+ ).annotate(detections_image, detection, detections_labels)
+ image_with_detections = Image.new(
+ 'RGB', (img.width + detections_image.width, img.height)
+ )
+ image_with_detections.paste(img, (0, 0))
+ image_with_detections.paste(detections_image, (img.width, 0))
+ image_with_detections.save(save_path)
+
+
+def shutdown_vm():
+ """Shuts down the system."""
+ time.sleep(60)
+ print('Attempting to shut down the VM instance...')
+ try:
+ command = ['sudo', 'poweroff']
+ subprocess.run(command, check=True)
+ except subprocess.CalledProcessError as e:
+ print(f'Failed to shut down: {e}')
diff --git a/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/server/triton_inference_server.sh b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/server/triton_inference_server.sh
new file mode 100644
index 00000000000..3f78aebbe91
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/server/triton_inference_server.sh
@@ -0,0 +1,42 @@
+#!/bin/bash
+
+# Summary
+cat << END
+Automates the setup and deployment of a Pytorch ONNX Model on the NVIDIA Triton Inference Server.
+It cleans up old setups, downloads and organizes models, creates the required configuration,
+ensures screen is installed, and launches Triton in a detached session with GPU support.
+END
+
+# Check and delete the model_repository directory if it exists
+if [ -d "model_repository" ]; then
+ echo "Removing existing model_repository directory..."
+ rm -rf model_repository
+fi
+
+# Define an associative array with model names and their URLs
+declare -A models=(
+ ["CircularNet_Segmentation_Model_v1"]="https://storage.googleapis.com/"\
+"tf_model_garden/vision/waste_identification_ml/"\
+"CN-ModelCheckpoints/CN-TritonInferenceServer/CircularNet_model_v2.zip"
+)
+
+
+# Download, unzip, and organize models
+for model_name in "${!models[@]}"; do
+ url=${models[$model_name]}
+ zip_file="${url##*/}"
+ wget "$url" && unzip "$zip_file"
+ rm "$zip_file"
+done
+
+# Install screen if not already installed
+command -v screen >/dev/null 2>&1 || { \
+ sudo apt update && sudo apt install -y screen; \
+}
+
+echo "Starting Triton server in a screen session."
+
+# Start Triton server
+screen -dmS server bash -c '
+sudo docker run --gpus all --rm -p 8000:8000 -p 8001:8001 -p 8002:8002 \
+-v ${PWD}/model_repository:/models nvcr.io/nvidia/tritonserver:25.05-py3 tritonserver --model-repository=/models'
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/big_query_ops.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/big_query_ops.py
new file mode 100644
index 00000000000..7e4a36577cd
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/big_query_ops.py
@@ -0,0 +1,139 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Designed to interact with Google BigQuery.
+
+For the purpose of dataset and table management, as well as data ingestion
+from pandas DataFrames.
+"""
+
+import logging
+import os
+import subprocess
+from google.cloud import bigquery
+from google.cloud import exceptions
+import pandas as pd
+import pandas_gbq
+
+# Configure logging
+logging.basicConfig(
+ level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
+)
+
+# Centralized Schema Definition
+_BIGQUERY_SCHEMA = (
+ bigquery.SchemaField("tracker_id", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("frame_name", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("best_class", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("best_probability", "FLOAT", mode="REQUIRED"),
+)
+
+
+class BigQueryManager:
+ """Manages interactions with Google BigQuery for dataset and table operations.
+
+ This class provides methods to create datasets and tables, ingest data from
+ pandas DataFrames, and manage related file operations in Google Cloud Storage.
+ """
+
+ def __init__(self, project_id: str, dataset_id: str, table_id: str):
+ """Initializes the BigQuery client and storage coordinates."""
+ self.client = bigquery.Client(project=project_id)
+ self.project_id = project_id
+ self.dataset_id = dataset_id
+ self.table_id = table_id
+ self.table_ref = f"{project_id}.{dataset_id}.{table_id}"
+
+ def _ensure_dataset(self):
+ """Checks if dataset exists, creates it if not."""
+ dataset_ref = self.client.dataset(self.dataset_id)
+ try:
+ self.client.get_dataset(dataset_ref)
+ except exceptions.NotFound:
+ logging.info("Dataset %s not found. Creating...", self.dataset_id)
+ dataset = bigquery.Dataset(dataset_ref)
+ self.client.create_dataset(dataset, timeout=30)
+
+ def create_table(self, overwrite: bool = False) -> None:
+ """Creates the table with the defined schema."""
+ self._ensure_dataset()
+
+ try:
+ self.client.get_table(self.table_ref)
+ if overwrite:
+ logging.info("Overwriting table %s...", self.table_id)
+ self.client.delete_table(self.table_ref)
+ else:
+ logging.info("Table %s already exists. Skipping.", self.table_id)
+ return
+ except exceptions.NotFound:
+ pass
+
+ table = bigquery.Table(self.table_ref, schema=_BIGQUERY_SCHEMA)
+ self.client.create_table(table)
+ logging.info("Table %s created successfully.", self.table_id)
+
+ def ingest_data(self, df: pd.DataFrame) -> None:
+ """Ingests data from a pandas DataFrame into BigQuery using pandas_gbq."""
+ pandas_gbq.to_gbq(
+ df,
+ destination_table=self.table_ref,
+ project_id=self.project_id,
+ if_exists="append",
+ )
+ logging.info("Data ingested successfully into %s", self.table_ref)
+
+ def upload_image_results_to_storage_bucket(
+ self, input_directory: str, prediction_folder: str, output_directory: str
+ ) -> None:
+ """Moves folders to the destination bucket and cleans up local directories.
+
+ Args:
+ input_directory: Path to the local input directory.
+ prediction_folder: Path to the local folder containing results.
+ output_directory: The GCS path (gs://...) for output.
+ """
+ try:
+ commands = [
+ f"rm -r {os.path.basename(input_directory)}",
+ f"gcloud storage cp -r {prediction_folder} {output_directory}",
+ f"rm -r {prediction_folder}",
+ ]
+ subprocess.run(" && ".join(commands), shell=True, check=True)
+ logging.info("Successfully moved to destination bucket")
+ except (
+ KeyError,
+ IndexError,
+ TypeError,
+ ValueError,
+ subprocess.CalledProcessError,
+ ) as e:
+ logging.info(
+ "Issue in moving folders to destination bucket, due to error : %s",
+ e,
+ )
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/big_query_ops_test.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/big_query_ops_test.py
new file mode 100644
index 00000000000..71f163e5cdb
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/big_query_ops_test.py
@@ -0,0 +1,224 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import subprocess
+import unittest
+from unittest import mock
+
+from google.cloud import exceptions
+import pandas as pd
+
+from official.projects.waste_identification_ml.Deploy.pet_grading_cloud_deployment import big_query_ops
+
+
+class BigQueryManagerTest(unittest.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ self.project_id = "test-project"
+ self.dataset_id = "test_dataset"
+ self.table_id = "test_table"
+ patcher = mock.patch.object(big_query_ops.bigquery, "Client")
+ self.mock_bigquery_client_cls = patcher.start()
+ self.addCleanup(patcher.stop)
+ self.mock_bigquery_client = self.mock_bigquery_client_cls.return_value
+
+ def test_init(self):
+ # Act
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+
+ # Assert
+ self.mock_bigquery_client_cls.assert_called_once_with(
+ project=self.project_id
+ )
+ self.assertEqual(manager.project_id, self.project_id)
+ self.assertEqual(manager.dataset_id, self.dataset_id)
+ self.assertEqual(manager.table_id, self.table_id)
+ self.assertEqual(
+ manager.table_ref,
+ f"{self.project_id}.{self.dataset_id}.{self.table_id}",
+ )
+
+ def test_ensure_dataset_exists(self):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+
+ # Act
+ manager._ensure_dataset()
+
+ # Assert
+ self.mock_bigquery_client.dataset.assert_called_once_with(self.dataset_id)
+ self.mock_bigquery_client.get_dataset.assert_called_once_with(
+ self.mock_bigquery_client.dataset.return_value
+ )
+ self.mock_bigquery_client.create_dataset.assert_not_called()
+
+ def test_ensure_dataset_not_found_creates(self):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+ self.mock_bigquery_client.get_dataset.side_effect = exceptions.NotFound(
+ "Dataset not found"
+ )
+
+ # Act
+ manager._ensure_dataset()
+
+ # Assert
+ self.mock_bigquery_client.dataset.assert_called_once_with(self.dataset_id)
+ self.mock_bigquery_client.get_dataset.assert_called_once()
+ self.mock_bigquery_client.create_dataset.assert_called_once()
+
+ def test_create_table_skips_when_exists_no_overwrite(self):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+
+ # Act
+ with mock.patch.object(manager, "_ensure_dataset") as mock_ensure_dataset:
+ manager.create_table(overwrite=False)
+ mock_ensure_dataset.assert_called_once()
+
+ # Assert
+ self.mock_bigquery_client.get_table.assert_called_once_with(
+ manager.table_ref
+ )
+ self.mock_bigquery_client.delete_table.assert_not_called()
+ self.mock_bigquery_client.create_table.assert_not_called()
+
+ def test_create_table_overwrites_when_exists_and_overwrite(self):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+
+ # Act
+ with mock.patch.object(manager, "_ensure_dataset") as mock_ensure_dataset:
+ manager.create_table(overwrite=True)
+ mock_ensure_dataset.assert_called_once()
+
+ # Assert
+ self.mock_bigquery_client.get_table.assert_called_once_with(
+ manager.table_ref
+ )
+ self.mock_bigquery_client.delete_table.assert_called_once_with(
+ manager.table_ref
+ )
+ self.mock_bigquery_client.create_table.assert_called_once()
+
+ def test_create_table_creates_when_not_exists(self):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+ self.mock_bigquery_client.get_table.side_effect = exceptions.NotFound(
+ "Table not found"
+ )
+
+ # Act
+ with mock.patch.object(manager, "_ensure_dataset") as mock_ensure_dataset:
+ manager.create_table(overwrite=False)
+ mock_ensure_dataset.assert_called_once()
+
+ # Assert
+ self.mock_bigquery_client.get_table.assert_called_once_with(
+ manager.table_ref
+ )
+ self.mock_bigquery_client.delete_table.assert_not_called()
+ self.mock_bigquery_client.create_table.assert_called_once()
+
+ @mock.patch.object(big_query_ops.pandas_gbq, "to_gbq")
+ def test_ingest_data(self, mock_to_gbq):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+ df = pd.DataFrame([{"tracker_id": 1, "frame_name": "frame_1"}])
+
+ # Act
+ manager.ingest_data(df)
+
+ # Assert
+ mock_to_gbq.assert_called_once_with(
+ df,
+ destination_table=manager.table_ref,
+ project_id=manager.project_id,
+ if_exists="append",
+ )
+
+ @mock.patch.object(big_query_ops.subprocess, "run")
+ def test_upload_image_results_success(self, mock_run):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+
+ # Act
+ manager.upload_image_results_to_storage_bucket(
+ input_directory="/path/to/my_input_dir",
+ prediction_folder="/path/to/my_pred_folder",
+ output_directory="gs://my_output_bucket",
+ )
+
+ # Assert
+ expected_commands = [
+ "rm -r my_input_dir",
+ "gcloud storage cp -r /path/to/my_pred_folder gs://my_output_bucket",
+ "rm -r /path/to/my_pred_folder",
+ ]
+ mock_run.assert_called_once_with(
+ " && ".join(expected_commands), shell=True, check=True
+ )
+
+ @mock.patch.object(big_query_ops.subprocess, "run")
+ def test_upload_image_results_handles_subprocess_error(self, mock_run):
+ # Arrange
+ manager = big_query_ops.BigQueryManager(
+ self.project_id, self.dataset_id, self.table_id
+ )
+ mock_run.side_effect = subprocess.CalledProcessError(
+ returncode=1, cmd="some-command"
+ )
+
+ # Act & Assert (Should handle the error gracefully without raising)
+ manager.upload_image_results_to_storage_bucket(
+ input_directory="/path/to/my_input_dir",
+ prediction_folder="/path/to/my_pred_folder",
+ output_directory="gs://my_output_bucket",
+ )
+ mock_run.assert_called_once()
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/constants.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/constants.py
new file mode 100644
index 00000000000..97580d7438e
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/constants.py
@@ -0,0 +1,61 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Constants for the pet grading cloud deployment."""
+
+# Specify all the input paths
+DINOV3_REPO_DIR = './dinov3'
+CLASSIFIER_CHECKPOINT_PATH = './model_weights/best_pet_grading_model.pth'
+SAM3_CHECKPOINT_PATH = './model_weights/sam3_original_weights_sam3.pt'
+
+# SAM3 configurations.
+BATCH_SIZE = 10
+DETECTION_PROMPT = 'bottles and containers'
+BOTTLE_EXTRACTION_CONFIDENCE_THRESHOLD = 0.5
+BOTTLE_EXTRACTION_SCORE_THRESHOLD = 0.20
+BOTTLE_EXTRACTION_CONTAINMENT_THRESHOLD = 0.98
+BOTTLE_EXTRACTION_MAX_SHORT_SIDE = 1024
+BOTTLE_EXTRACTION_CROP_SIZE = (256, 256)
+
+
+# Tracking configurations.
+BYTETRACK_MINIMUM_IOU_THRESHOLD = 0.1
+BYTETRACK_MINIMUM_CONSECUTIVE_FRAMES = 2
+
+# Classifier configuration. Must match what was used at training time.
+DINOV3_MODEL_NAME = 'dinov3_vitl16'
+CLASSIFICATION_THRESHOLD = 50
+CLASS_NAMES = (
+ 'brown_bottles_grade3',
+ 'clean_PET_cold_drink_bottles_with_label_cap_ring_grade1',
+ 'clean_PET_cold_drink_bottles_without_label_cap_ring_grade1',
+ 'clean_PET_mango_juice_bottles_without_label_cap_ring_grade1',
+ 'clean_PET_water_bottles_with_label_cap_ring_grade1',
+ 'clean_PET_water_bottles_without_label_cap_ring_grade1',
+ 'clean_jars_grade1',
+ 'clean_liquor_bottles_without_label_cap_ring_grade1',
+ 'coloured_PET_bottles_grade3',
+ 'dirt_PET_cold_drink_bottles_with_label_cap_ring_grade3',
+ 'dirt_PET_cold_drink_bottles_without_label_cap_ring_grade3',
+ 'dirt_PET_mango_juice_bottles_without_label_cap_ring_grade3',
+ 'dirt_PET_water_bottles_with_label_cap_ring_grade3',
+ 'dirt_PET_water_bottles_without_label_cap_ring_grade3',
+ 'dirt_jars_grade3',
+ 'dirt_liquor_bottles_without_label_cap_ring_grade3',
+ 'full_sleeved_bottles_grade3',
+ 'green_bottles_grade3',
+ 'liquor_bottles_with_label_cap_ring_grade3',
+ 'non_food_bottles_grade3',
+ 'partially_sleeved_mango_juice_bottles_with_label_cap_ring_grade3',
+)
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/inference_pipeline.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/inference_pipeline.py
new file mode 100644
index 00000000000..2107b024aaf
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/inference_pipeline.py
@@ -0,0 +1,241 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Pipeline to run the prediction on the images folder."""
+
+import os
+
+from absl import app
+from absl import flags
+from big_query_ops import BigQueryManager
+from constants import BATCH_SIZE
+from constants import BOTTLE_EXTRACTION_CONFIDENCE_THRESHOLD
+from constants import BOTTLE_EXTRACTION_CONTAINMENT_THRESHOLD
+from constants import BOTTLE_EXTRACTION_CROP_SIZE
+from constants import BOTTLE_EXTRACTION_MAX_SHORT_SIDE
+from constants import BOTTLE_EXTRACTION_SCORE_THRESHOLD
+from constants import BYTETRACK_MINIMUM_CONSECUTIVE_FRAMES
+from constants import BYTETRACK_MINIMUM_IOU_THRESHOLD
+from constants import CLASS_NAMES
+from constants import CLASSIFICATION_THRESHOLD
+from constants import CLASSIFIER_CHECKPOINT_PATH
+from constants import DETECTION_PROMPT
+from constants import DINOV3_MODEL_NAME
+from constants import DINOV3_REPO_DIR
+from constants import SAM3_CHECKPOINT_PATH
+from object_extraction import ObjectExtractor
+import pandas as pd
+from pet_grade_classifier import DinoClassifier
+from PIL import Image
+import tqdm
+from trackers import ByteTrackTracker
+import utils
+
+INPUT_DIRECTORY = flags.DEFINE_string(
+ "input_directory", None, "The path to the directory containing images."
+)
+OUTPUT_DIRECTORY = flags.DEFINE_string(
+ "output_directory", None, "The path to the directory to save the results."
+)
+
+OVERWRITE = flags.DEFINE_boolean(
+ "overwrite",
+ False,
+ "If True, delete the preexisting BigQuery table before creating a new one.",
+)
+PROJECT_ID = flags.DEFINE_string(
+ "project_id", None, "Project ID mentioned in Google Cloud Project"
+)
+BQ_DATASET_ID = flags.DEFINE_string(
+ "bq_dataset_id", "Circularnet_dataset", "Big query dataset ID"
+)
+BQ_TABLE_ID = flags.DEFINE_string(
+ "bq_table_id", "Circularnet_table", "BigQuery Table ID for features data"
+)
+
+
+object_extractor = ObjectExtractor(
+ checkpoint_path=SAM3_CHECKPOINT_PATH,
+ confidence_threshold=BOTTLE_EXTRACTION_CONFIDENCE_THRESHOLD,
+ score_threshold=BOTTLE_EXTRACTION_SCORE_THRESHOLD,
+ containment_threshold=BOTTLE_EXTRACTION_CONTAINMENT_THRESHOLD,
+ max_short_side=BOTTLE_EXTRACTION_MAX_SHORT_SIDE,
+ crop_size=BOTTLE_EXTRACTION_CROP_SIZE,
+)
+
+
+classifier = DinoClassifier(
+ classifier_checkpoint_path=CLASSIFIER_CHECKPOINT_PATH,
+ dinov3_repo_dir=DINOV3_REPO_DIR,
+ model_name=DINOV3_MODEL_NAME,
+ class_names=CLASS_NAMES,
+)
+
+
+tracker = ByteTrackTracker(
+ minimum_iou_threshold=BYTETRACK_MINIMUM_IOU_THRESHOLD,
+ minimum_consecutive_frames=BYTETRACK_MINIMUM_CONSECUTIVE_FRAMES,
+)
+
+
+def get_batch_of_crops(image_paths, tracker_instance, batch_size=10):
+ """Yields per-batch crop dicts; tracker_instance persists across batches.
+
+ Args:
+ image_paths: List of file paths to the images.
+ tracker_instance: The tracker instance that persists across batches.
+ batch_size: Number of images to process in each batch.
+
+ Yields:
+ A dict mapping tracker ID to a list of dicts with frame name and crop.
+ One dict per batch of batch_size frames. The dict only contains
+ crops collected in that batch — not accumulated across all batches.
+ """
+ for batch_start in range(0, len(image_paths), batch_size):
+ batch_paths = image_paths[batch_start : batch_start + batch_size]
+
+ batch_records = {}
+ for image_path in tqdm.tqdm(
+ batch_paths,
+ desc=f"Batch {batch_start//batch_size + 1} frames",
+ leave=False,
+ ):
+ rgb_image = Image.open(image_path).convert("RGB")
+ resized_image, state, detections = object_extractor.extract(
+ rgb_image, prompt=DETECTION_PROMPT
+ )
+ detections = tracker_instance.update(detections)
+ object_extractor.get_cropped_objects_for_each_tracking_id(
+ image=resized_image,
+ state=state,
+ detections=detections,
+ source_frame_name=os.path.basename(image_path),
+ crop_size=BOTTLE_EXTRACTION_CROP_SIZE,
+ track_crop_records=batch_records,
+ )
+
+ yield batch_records
+
+
+def main(_) -> None:
+ if (
+ not INPUT_DIRECTORY.value
+ or not OUTPUT_DIRECTORY.value
+ or not INPUT_DIRECTORY.value.startswith("gs://")
+ or not OUTPUT_DIRECTORY.value.startswith("gs://")
+ ):
+ raise ValueError("Bucket path must be non-empty starting with 'gs://'")
+
+ input_directory, prediction_folder, logger = (
+ utils.setup_logger_and_directories(input_dir=INPUT_DIRECTORY.value)
+ )
+
+ checkpoint_path = os.path.join(prediction_folder, "prediction.csv")
+
+ storage_manager = BigQueryManager(
+ project_id=PROJECT_ID.value,
+ dataset_id=BQ_DATASET_ID.value,
+ table_id=BQ_TABLE_ID.value,
+ )
+
+ filepaths = utils.files_paths(os.path.basename(input_directory))
+ num_batches = (len(filepaths) + BATCH_SIZE - 1) // BATCH_SIZE
+ logger.info(
+ f"Found {len(filepaths)} image files. Starting inference over"
+ f" {num_batches} batches."
+ )
+ for batch_records in tqdm.tqdm(
+ get_batch_of_crops(
+ image_paths=filepaths, tracker_instance=tracker, batch_size=BATCH_SIZE
+ ),
+ total=num_batches,
+ desc="Batches",
+ ):
+ batch_predictions = []
+ for tracker_id, crop_records in tqdm.tqdm(
+ batch_records.items(), desc="Classifying crops", leave=False
+ ):
+ crop_pil_images = [crop_record["crop"] for crop_record in crop_records]
+ crop_predictions = classifier.predict(pil_images=crop_pil_images)
+ for crop_record, prediction in zip(crop_records, crop_predictions):
+ batch_predictions.append({
+ "tracker_id": tracker_id,
+ "frame_name": crop_record["frame_name"],
+ "crop": crop_record["crop"],
+ "predicted_class": prediction["predicted_class"],
+ "predicted_probability": prediction["predicted_probability"],
+ })
+ batch_predictions = utils.get_class_with_majority_vote(
+ batch_predictions, min_probability=CLASSIFICATION_THRESHOLD
+ )
+ utils.save_output_image_grids(
+ all_predictions=batch_predictions, output_dr=prediction_folder
+ )
+
+ # Append this image's rows to checkpoint immediately
+ if batch_predictions:
+ rows_df = (
+ pd.DataFrame(batch_predictions)
+ .drop(columns=["crop", "predicted_class", "predicted_probability"])
+ .drop_duplicates(subset=["tracker_id"])
+ )
+ rows_df.to_csv(
+ checkpoint_path,
+ mode="a",
+ header=not os.path.exists(checkpoint_path)
+ or os.path.getsize(checkpoint_path) == 0,
+ index=False,
+ )
+
+ # Upload Results to BigQuery
+ results_df = (
+ pd.read_csv(checkpoint_path).dropna()
+ if os.path.exists(checkpoint_path)
+ else pd.DataFrame()
+ )
+ if isinstance(results_df, pd.DataFrame) and not results_df.empty:
+ try:
+ logger.info(
+ f"Starting BigQuery ingestion for {len(results_df)} records..."
+ )
+ storage_manager.create_table(overwrite=OVERWRITE.value)
+ storage_manager.ingest_data(results_df)
+ storage_manager.upload_image_results_to_storage_bucket(
+ input_directory=input_directory,
+ prediction_folder=prediction_folder,
+ output_directory=OUTPUT_DIRECTORY.value,
+ )
+ logger.info("Pipeline execution successfully completed.")
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info(f"Issue in creation of table, due to error : {e}")
+ else:
+ logger.info("No data to ingest.")
+ # utils.shutdown_vm()
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/object_extraction.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/object_extraction.py
new file mode 100644
index 00000000000..89745d7b307
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/object_extraction.py
@@ -0,0 +1,456 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""ObjectExtractor: SAM3-based object segmentation and crop extraction."""
+
+from typing import Any
+import cv2
+import numpy as np
+from PIL import Image
+from sam3 import build_sam3_image_model
+from sam3.model.sam3_image_processor import Sam3Processor
+import supervision
+import torch
+
+_INFERENCE_KEYS_TO_DROP = frozenset([
+ "backbone_out",
+ "geometric_prompt",
+ "image_embeddings",
+])
+_STATE_ARRAY_KEYS = ("masks", "masks_logits", "boxes", "scores")
+_IMAGENET_MEAN_RGB = (124, 116, 104)
+_CROP_BUFFER = 5
+
+
+class ObjectExtractor:
+ """Segments objects in an image using SAM3 and returns cropped detections.
+
+ Wraps model loading, inference, filtering, merging, and cropping into a
+ single reusable object. All stateless image-processing helpers are exposed
+ as static methods so they can also be called without an instance.
+
+ Usage:
+ extractor = ObjectExtractor(
+ checkpoint_path="sam3.pth",
+ confidence_threshold=0.5,
+ score_threshold=0.5,
+ containment_threshold=0.5,
+ max_short_side=1024,
+ crop_size=(224, 224),
+ )
+ resized_image, state, detections = extractor.extract(pil_image, "bottle")
+ """
+
+ def __init__(
+ self,
+ checkpoint_path: str,
+ confidence_threshold: float,
+ score_threshold: float,
+ containment_threshold: float,
+ max_short_side: int,
+ crop_size: tuple[int, int],
+ device: str | None = None,
+ ) -> None:
+ """Load and initialize the SAM3 model and processor.
+
+ Args:
+ checkpoint_path: Path to the SAM3 weights file.
+ confidence_threshold: Confidence threshold for the SAM3 processor.
+ score_threshold: Minimum score to keep a detection.
+ containment_threshold: Ratio above which a smaller mask is dropped if it
+ is contained within a larger one.
+ max_short_side: Maximum allowed length for the shorter dimension of the
+ input image during inference.
+ crop_size: Target (height, width) for cropped object images.
+ device: Torch device string. Defaults to 'cuda' if available, otherwise
+ 'cpu'.
+ """
+ self._confidence_threshold = confidence_threshold
+ self._score_threshold = score_threshold
+ self._containment_threshold = containment_threshold
+ self._max_short_side = max_short_side
+ self._crop_size = crop_size
+ self._device = device or ("cuda" if torch.cuda.is_available() else "cpu")
+ self._model = (
+ build_sam3_image_model(checkpoint_path=checkpoint_path)
+ .eval()
+ .to(device=self._device)
+ )
+ self._processor = Sam3Processor(
+ self._model, confidence_threshold=self._confidence_threshold
+ )
+
+ # ------------------------------------------------------------------
+ # Public pipeline
+ # ------------------------------------------------------------------
+
+ def extract(
+ self, image: Image.Image, prompt: str
+ ) -> tuple[Image.Image, dict[str, Any], supervision.Detections]:
+ """Run the full extraction pipeline on a single image.
+
+ Steps: resize → SAM3 inference → sub-mask filtering → detections.
+
+ Args:
+ image: Input RGB PIL image.
+ prompt: Text prompt for SAM3.
+
+ Returns:
+ A tuple of (resized_image, inference_state, detections).
+ """
+ resized_image = self.resize_image_for_inference(image, self._max_short_side)
+ state = self._run_inference(resized_image, prompt)
+ state = self.filter_contained_sub_masks(state, self._containment_threshold)
+ detections = self.convert_sam3_state_to_detections(
+ state, score_threshold=self._score_threshold
+ )
+ return resized_image, state, detections
+
+ # ------------------------------------------------------------------
+ # Private inference helpers
+ # ------------------------------------------------------------------
+
+ def _run_inference(self, image: Image.Image, prompt: str) -> dict[str, Any]:
+ """Runs SAM3 inference on the input image."""
+ with torch.no_grad(), torch.autocast(
+ device_type="cuda", dtype=torch.float16
+ ):
+ state = self._processor.set_image(image)
+ state = self._processor.set_text_prompt(state=state, prompt=prompt)
+
+ for key in _INFERENCE_KEYS_TO_DROP:
+ state.pop(key, None)
+
+ return self._move_state_to_cpu(state)
+
+ def _move_state_to_cpu(
+ self, inference_state: dict[str, Any]
+ ) -> dict[str, Any]:
+ """Recursively moves tensors in the state dict to CPU."""
+ for key, value in inference_state.items():
+ if isinstance(value, torch.Tensor):
+ inference_state[key] = value.cpu()
+ elif isinstance(value, dict):
+ self._move_state_to_cpu(value)
+ return inference_state
+
+ def filter_contained_sub_masks(
+ self, state: dict[str, Any], containment_threshold: float
+ ) -> dict[str, Any]:
+ """Removes smaller masks that are largely contained within larger masks.
+
+ Args:
+ state: Dict with 'masks', 'masks_logits', 'boxes', 'scores'. masks is a
+ bool tensor of shape [N, H, W].
+ containment_threshold: Ratio above which the smaller mask is dropped.
+
+ Returns:
+ Filtered state dict.
+ """
+ masks = state["masks"]
+ num_masks = masks.shape[0]
+ if num_masks == 0:
+ return state
+
+ flat_masks = masks.view(num_masks, -1).float()
+ areas = flat_masks.sum(dim=1)
+ pairwise_intersection = flat_masks @ flat_masks.T
+
+ indices_to_remove: set[int] = set()
+ for i in range(num_masks):
+ if i in indices_to_remove:
+ continue
+ for j in range(i + 1, num_masks):
+ if j in indices_to_remove:
+ continue
+
+ intersection = pairwise_intersection[i, j].item()
+ area_i, area_j = areas[i].item(), areas[j].item()
+ smaller_index = i if area_i <= area_j else j
+ smaller_area = min(area_i, area_j)
+
+ if smaller_area == 0:
+ indices_to_remove.add(smaller_index)
+ continue
+
+ if intersection / smaller_area > containment_threshold:
+ indices_to_remove.add(smaller_index)
+
+ keep = torch.tensor(
+ sorted(set(range(num_masks)) - indices_to_remove), dtype=torch.long
+ )
+ for key in _STATE_ARRAY_KEYS:
+ state[key] = state[key][keep]
+ return state
+
+ @staticmethod
+ def resize_image_for_inference(
+ image: Image.Image, max_short_side: int
+ ) -> Image.Image:
+ """Resize so the short side does not exceed max_short_side (aspect-ratio safe).
+
+ Args:
+ image: Input RGB PIL image.
+ max_short_side: Maximum allowed length for the shorter dimension.
+
+ Returns:
+ Resized PIL image, or the original if already within the limit.
+ """
+ w, h = image.size
+ short_side = min(w, h)
+ if short_side <= max_short_side:
+ return image
+ scale = max_short_side / short_side
+ return image.resize((int(w * scale), int(h * scale)), Image.LANCZOS)
+
+ @staticmethod
+ def fill_mask_holes(mask: np.ndarray) -> np.ndarray:
+ """Fills all interior holes in a binary mask via border flood-fill.
+
+ More robust than morphological closing: fills holes of any size.
+
+ Args:
+ mask: Binary mask of shape (H, W), dtype bool or uint8.
+
+ Returns:
+ Hole-filled binary mask, dtype bool.
+ """
+ mask_u8 = np.asarray(mask).astype(np.uint8) * 255
+ h, w = mask_u8.shape
+ padded = np.zeros((h + 2, w + 2), dtype=np.uint8)
+ padded[1 : h + 1, 1 : w + 1] = mask_u8
+
+ filled = padded.copy()
+ cv2.floodFill(filled, mask=None, seedPoint=(0, 0), newVal=255)
+ filled = filled[1 : h + 1, 1 : w + 1]
+
+ interior_holes = cv2.bitwise_not(filled)
+ return cv2.bitwise_or(mask_u8, interior_holes).astype(bool)
+
+ @staticmethod
+ def get_padded_box(
+ box: list[float], mask_shape: tuple[int, ...], buffer: int = _CROP_BUFFER
+ ) -> tuple[int, int, int, int]:
+ """Expands a bounding box by buffer pixels, clamped to mask bounds.
+
+ Args:
+ box: [x_min, y_min, x_max, y_max].
+ mask_shape: (H, W, ...) shape of the reference mask.
+ buffer: Pixels to add on each side.
+
+ Returns:
+ (x_min, y_min, x_max, y_max) clamped to valid image bounds.
+ """
+ mask_h, mask_w = mask_shape[:2]
+ x_min, y_min, x_max, y_max = map(round, box)
+ return (
+ max(0, x_min - buffer),
+ max(0, y_min - buffer),
+ min(mask_w, x_max + buffer),
+ min(mask_h, y_max + buffer),
+ )
+
+ @staticmethod
+ def letterbox_image(
+ image: np.ndarray,
+ size: tuple[int, int],
+ color: tuple[int, int, int] = (0, 0, 0),
+ ) -> np.ndarray:
+ """Resize onto a fixed canvas while preserving aspect ratio.
+
+ Args:
+ image: (H, W, 3) numpy array.
+ size: Target (height, width).
+ color: RGB fill color for padding.
+
+ Returns:
+ Letterboxed numpy array of shape (size[0], size[1], 3).
+ """
+ ih, iw = image.shape[:2]
+ th, tw = size
+ scale = min(tw / iw, th / ih)
+ nw, nh = int(iw * scale), int(ih * scale)
+ resized = cv2.resize(image, (nw, nh), interpolation=cv2.INTER_LINEAR)
+ canvas = np.full((th, tw, 3), color, dtype=np.uint8)
+ ox, oy = (tw - nw) // 2, (th - nh) // 2
+ canvas[oy : oy + nh, ox : ox + nw] = resized
+ return canvas
+
+ @staticmethod
+ def crop_with_mean_background_blend(
+ image_array: np.ndarray,
+ mask: np.ndarray,
+ box: list[float],
+ size: tuple[int, int],
+ background_color: tuple[int, int, int] = _IMAGENET_MEAN_RGB,
+ ) -> Image.Image:
+ """Soft-edged letterboxed crop blended against a background color.
+
+ Args:
+ image_array: (H, W, 3) RGB numpy array.
+ mask: Binary mask of shape (H, W).
+ box: [x_min, y_min, x_max, y_max].
+ size: Output (height, width) after letterboxing.
+ background_color: RGB tuple for the blended background.
+
+ Returns:
+ Letterboxed PIL image with soft mask blending.
+ """
+ x_min, y_min, x_max, y_max = ObjectExtractor.get_padded_box(box, mask.shape)
+ roi_image = image_array[y_min:y_max, x_min:x_max]
+ roi_mask = mask[y_min:y_max, x_min:x_max].astype(np.uint8) * 255
+
+ kernel = np.ones((5, 5), np.uint8)
+ dilated = cv2.dilate(roi_mask, kernel, iterations=1)
+ blurred = cv2.GaussianBlur(dilated, (5, 5), 0)
+
+ alpha = blurred.astype(np.float32) / 255.0
+ bg = np.array(background_color, dtype=np.float32)
+ blended = roi_image.astype(np.float32) * alpha[:, :, None] + bg * (
+ 1.0 - alpha[:, :, None]
+ )
+
+ lb = ObjectExtractor.letterbox_image(
+ blended.astype(np.uint8), size=size, color=background_color
+ )
+ return Image.fromarray(lb)
+
+ @staticmethod
+ def convert_sam3_state_to_detections(
+ state: dict[str, Any], score_threshold: float
+ ) -> supervision.Detections:
+ """Converts a SAM3 state dict to an ``sv.Detections`` object.
+
+ Only detections with score >= ``score_threshold`` are kept. All
+ detections are assigned ``class_id=0`` because the pipeline uses a
+ single text prompt.
+
+ Args:
+ state: SAM3 state dict with 'boxes' and 'scores' tensors.
+ score_threshold: Minimum score to keep a detection.
+
+ Returns:
+ ``sv.Detections`` with xyxy boxes, confidence scores, and
+ ``class_id=0``.
+ """
+ boxes = state["boxes"].numpy().astype(np.float32)
+ scores = state["scores"].numpy().astype(np.float32)
+
+ if boxes.ndim == 1:
+ boxes = boxes.reshape(-1, 4)
+
+ keep_mask = scores >= score_threshold
+ kept_boxes = boxes[keep_mask]
+ kept_scores = scores[keep_mask]
+
+ if kept_boxes.shape[0] == 0:
+ return supervision.Detections.empty()
+
+ class_ids = np.zeros(kept_boxes.shape[0], dtype=int)
+ return supervision.Detections(
+ xyxy=kept_boxes,
+ confidence=kept_scores,
+ class_id=class_ids,
+ )
+
+ @staticmethod
+ def get_cropped_objects_for_each_tracking_id(
+ image: Image.Image,
+ state: dict[str, Any],
+ detections: supervision.Detections,
+ source_frame_name: str,
+ crop_size: tuple[int, int],
+ track_crop_records: dict[int, list[dict[str, Any]]],
+ ) -> None:
+ """Adds blended crops for assigned detections in a single frame.
+
+ Detections without an assigned tracker_id (tracker_id == -1) are
+ skipped. Detections from sv.Detections are matched back to rows of
+ ``state`` by xyxy near-equality.
+
+ Args:
+ image: RGB PIL image used for SAM3 inference on this frame.
+ state: SAM3 state dict with 'boxes' and 'masks' tensors on CPU.
+ detections: sv.Detections returned by tracker.update() for this frame.
+ source_frame_name: Filename (no directory) of the source frame, used to
+ label thumbnails.
+ crop_size: Output (height, width) for each blended crop.
+ track_crop_records: Mutable dict keyed by tracker_id. Each value is a
+ list of dicts shaped ``{'frame_name': str, 'crop': PIL.Image}``. New
+ records are appended in place.
+ """
+ if len(detections) == 0:
+ return
+ if detections.tracker_id is None:
+ return
+
+ state_boxes = state["boxes"].numpy().astype(np.float32)
+ detection_boxes = detections.xyxy.astype(np.float32)
+
+ for detection_row in range(len(detections)):
+ tracker_id = int(detections.tracker_id[detection_row])
+ if tracker_id == -1:
+ continue
+
+ detection_box = detection_boxes[detection_row]
+ matching_state_rows = np.where(
+ np.all(
+ np.isclose(state_boxes, detection_box, atol=1e-3),
+ axis=1,
+ )
+ )[0]
+ if matching_state_rows.size == 0:
+ continue
+
+ state_index = int(matching_state_rows[0])
+ blended_crop = ObjectExtractor.extract_blended_crop(
+ image=image,
+ state=state,
+ detection_index=state_index,
+ crop_size=crop_size,
+ )
+ track_crop_records.setdefault(tracker_id, []).append({
+ "frame_name": source_frame_name,
+ "crop": blended_crop,
+ })
+
+ @staticmethod
+ def extract_blended_crop(
+ image: Image.Image,
+ state: dict[str, Any],
+ detection_index: int,
+ crop_size: tuple[int, int],
+ ) -> Image.Image:
+ """Builds the blended crop variant for a single SAM3 detection.
+
+ Args:
+ image: RGB PIL image already resized for SAM3 inference.
+ state: SAM3 state dict with 'masks' and 'boxes' tensors on CPU.
+ detection_index: Row index into state['masks'] / state['boxes'].
+ crop_size: Output (height, width) for the letterboxed crop.
+
+ Returns:
+ Blended PIL image of size crop_size.
+ """
+ image_array = np.array(image)
+ raw_mask = np.squeeze(state["masks"][detection_index])
+ filled_mask = ObjectExtractor.fill_mask_holes(raw_mask)
+ box = state["boxes"][detection_index].tolist()
+ blended_crop = ObjectExtractor.crop_with_mean_background_blend(
+ image_array=image_array,
+ mask=filled_mask,
+ box=box,
+ size=crop_size,
+ )
+ return blended_crop
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/object_extraction_test.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/object_extraction_test.py
new file mode 100644
index 00000000000..428a1902a56
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/object_extraction_test.py
@@ -0,0 +1,415 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import sys
+import unittest
+from unittest import mock
+
+import numpy as np
+from PIL import Image
+import torch
+
+# Mock supervision before it is imported anywhere
+mock_supervision = mock.MagicMock()
+
+
+def empty_mock_detections():
+ return MockDetections()
+
+
+class MockDetections:
+
+ def __init__(self, xyxy=None, confidence=None, class_id=None):
+ self.xyxy = xyxy
+ self.confidence = confidence
+ self.class_id = class_id
+ self.tracker_id = None
+
+ def __len__(self):
+ return len(self.xyxy) if self.xyxy is not None else 0
+
+
+mock_supervision.Detections = MockDetections
+sys.modules["supervision"] = mock_supervision
+
+# Mock sam3 before it is imported anywhere
+mock_sam3 = mock.MagicMock()
+mock_sam3_image_processor = mock.MagicMock()
+sys.modules["sam3"] = mock_sam3
+sys.modules["sam3.model"] = mock.MagicMock()
+sys.modules["sam3.model.sam3_image_processor"] = mock_sam3_image_processor
+
+mock_sam3.build_sam3_image_model = mock.MagicMock()
+
+
+class MockSam3Processor:
+
+ def __init__(self, model, confidence_threshold=0.5):
+ self.model = model
+ self.confidence_threshold = confidence_threshold
+
+
+mock_sam3_image_processor.Sam3Processor = MockSam3Processor
+
+from official.projects.waste_identification_ml.Deploy.pet_grading_cloud_deployment import object_extraction # pylint: disable=g-bad-import-order, g-import-not-at-top
+
+MODULE_PATH = object_extraction.__name__
+
+
+class ObjectExtractorTest(unittest.TestCase):
+
+ @mock.patch(f"{MODULE_PATH}.Sam3Processor", MockSam3Processor)
+ @mock.patch(f"{MODULE_PATH}.build_sam3_image_model")
+ def test_init(self, mock_build_model):
+ # Arrange
+ mock_model = mock.MagicMock()
+ mock_build_model.return_value = mock_model
+ mock_model.eval.return_value = mock_model
+ mock_model.to.return_value = mock_model
+
+ # Act
+ extractor = object_extraction.ObjectExtractor(
+ checkpoint_path="mock_checkpoint.pth",
+ confidence_threshold=0.6,
+ score_threshold=0.5,
+ containment_threshold=0.7,
+ max_short_side=800,
+ crop_size=(224, 224),
+ device="cpu",
+ )
+
+ # Assert
+ mock_build_model.assert_called_once_with(
+ checkpoint_path="mock_checkpoint.pth"
+ )
+ mock_model.eval.assert_called_once()
+ mock_model.to.assert_called_once_with(device="cpu")
+ self.assertEqual(extractor._confidence_threshold, 0.6)
+ self.assertEqual(extractor._score_threshold, 0.5)
+ self.assertEqual(extractor._containment_threshold, 0.7)
+ self.assertEqual(extractor._max_short_side, 800)
+ self.assertEqual(extractor._crop_size, (224, 224))
+ self.assertEqual(extractor._device, "cpu")
+
+ def test_resize_image_for_inference_no_resize(self):
+ # Image size 400x300, min side 300, max permitted is 400.
+ image = Image.new("RGB", (400, 300))
+ resized = object_extraction.ObjectExtractor.resize_image_for_inference(
+ image, max_short_side=400
+ )
+ self.assertEqual(resized.size, (400, 300))
+
+ def test_resize_image_for_inference_does_resize(self):
+ # Image size 800x600, min side 600, max permitted is 300
+ image = Image.new("RGB", (800, 600))
+ resized = object_extraction.ObjectExtractor.resize_image_for_inference(
+ image, max_short_side=300
+ )
+ self.assertEqual(resized.size, (400, 300))
+
+ def test_fill_mask_holes(self):
+ # Create a simple mask (10x10) with a hole (a zero surrounded by ones)
+ mask = np.ones((10, 10), dtype=bool)
+ mask[4, 4] = False # Hole at index (4, 4)
+
+ filled = object_extraction.ObjectExtractor.fill_mask_holes(mask)
+
+ # The hole should be filled and become True
+ self.assertTrue(filled[4, 4])
+ self.assertTrue(np.all(filled))
+
+ def test_get_padded_box(self):
+ # Box [10, 10, 50, 50], mask shape (100, 100), buffer 5
+ box = [10.0, 10.0, 50.0, 50.0]
+ padded = object_extraction.ObjectExtractor.get_padded_box(
+ box, (100, 100), buffer=5
+ )
+ self.assertEqual(padded, (5, 5, 55, 55))
+
+ def test_get_padded_box_clamps_to_boundaries(self):
+ # Box [2, 2, 98, 98], mask shape (100, 100), buffer 5
+ box = [2.0, 2.0, 98.0, 98.0]
+ padded = object_extraction.ObjectExtractor.get_padded_box(
+ box, (100, 100), buffer=5
+ )
+ self.assertEqual(padded, (0, 0, 100, 100))
+
+ def test_letterbox_image(self):
+ image = np.zeros((100, 200, 3), dtype=np.uint8)
+ # Target size: 200x200
+ letterboxed = object_extraction.ObjectExtractor.letterbox_image(
+ image, (200, 200)
+ )
+ self.assertEqual(letterboxed.shape, (200, 200, 3))
+ # Aspect ratio should preserve
+ # The padding should be at top/bottom (y range 50 to 150)
+ self.assertTrue(np.all(letterboxed[:50, :, :] == 0))
+ self.assertTrue(np.all(letterboxed[150:, :, :] == 0))
+
+ def test_crop_with_mean_background_blend(self):
+ image_array = np.ones((50, 50, 3), dtype=np.uint8) * 100
+ mask = np.ones((50, 50), dtype=bool)
+ box = [10.0, 10.0, 40.0, 40.0]
+ size = (30, 30)
+
+ cropped = object_extraction.ObjectExtractor.crop_with_mean_background_blend(
+ image_array=image_array, mask=mask, box=box, size=size
+ )
+
+ self.assertIsInstance(cropped, Image.Image)
+ self.assertEqual(cropped.size, size)
+
+ def test_convert_sam3_state_to_detections(self):
+ state = {
+ "boxes": torch.tensor([[10.0, 10.0, 50.0, 50.0], [2.0, 2.0, 9.0, 9.0]]),
+ "scores": torch.tensor([0.8, 0.4]),
+ "masks": torch.ones((2, 100, 100), dtype=torch.bool),
+ }
+
+ detections = (
+ object_extraction.ObjectExtractor.convert_sam3_state_to_detections(
+ state, score_threshold=0.5
+ )
+ )
+
+ self.assertEqual(len(detections), 1)
+ np.testing.assert_allclose(
+ detections.xyxy, np.array([[10.0, 10.0, 50.0, 50.0]], dtype=np.float32)
+ )
+ np.testing.assert_allclose(
+ detections.confidence, np.array([0.8], dtype=np.float32)
+ )
+ np.testing.assert_equal(detections.class_id, np.array([0]))
+
+ def test_filter_contained_sub_masks(self):
+ # Two masks: one large, and one small completely inside it.
+ mask1 = torch.zeros((5, 5), dtype=torch.bool)
+ mask1[1:4, 1:4] = True # Area 9
+
+ mask2 = torch.zeros((5, 5), dtype=torch.bool)
+ mask2[2, 2] = True # Area 1, fully inside mask1
+
+ state = {
+ "masks": torch.stack([mask1, mask2]),
+ "masks_logits": torch.stack([mask1.float(), mask2.float()]),
+ "boxes": torch.tensor([[1.0, 1.0, 3.0, 3.0], [2.0, 2.0, 2.0, 2.0]]),
+ "scores": torch.tensor([0.9, 0.85]),
+ }
+
+ extractor: object_extraction.ObjectExtractor = (
+ object_extraction.ObjectExtractor.__new__(
+ object_extraction.ObjectExtractor
+ )
+ )
+ filtered = extractor.filter_contained_sub_masks(
+ state, containment_threshold=0.5
+ )
+
+ # Only mask1 (the larger one) should be kept
+ self.assertEqual(filtered["masks"].shape[0], 1)
+ self.assertTrue(torch.equal(filtered["masks"][0], mask1))
+
+ @mock.patch.object(object_extraction.ObjectExtractor, "extract_blended_crop")
+ def test_get_cropped_objects_for_each_tracking_id(self, mock_extract_crop):
+ image = Image.new("RGB", (100, 100))
+ state = {
+ "boxes": torch.tensor([[10.0, 15.0, 50.0, 55.0]]),
+ "masks": torch.ones((1, 100, 100), dtype=torch.bool),
+ }
+ detections = MockDetections(
+ xyxy=np.array([[10.0, 15.0, 50.0, 55.0]], dtype=np.float32),
+ confidence=np.array([0.9], dtype=np.float32),
+ class_id=np.array([0]),
+ )
+ detections.tracker_id = np.array([42])
+
+ mock_crop_img = Image.new("RGB", (20, 20))
+ mock_extract_crop.return_value = mock_crop_img
+
+ track_crop_records = {}
+
+ object_extraction.ObjectExtractor.get_cropped_objects_for_each_tracking_id(
+ image=image,
+ state=state,
+ detections=detections,
+ source_frame_name="frame_0.png",
+ crop_size=(20, 20),
+ track_crop_records=track_crop_records,
+ )
+
+ mock_extract_crop.assert_called_once_with(
+ image=image, state=state, detection_index=0, crop_size=(20, 20)
+ )
+ self.assertIn(42, track_crop_records)
+ self.assertEqual(len(track_crop_records[42]), 1)
+ self.assertEqual(track_crop_records[42][0]["frame_name"], "frame_0.png")
+ self.assertEqual(track_crop_records[42][0]["crop"], mock_crop_img)
+
+ @mock.patch(f"{MODULE_PATH}.Sam3Processor", MockSam3Processor)
+ @mock.patch(f"{MODULE_PATH}.build_sam3_image_model")
+ def test_extract(self, mock_build_model):
+ # Set up mock model and processor
+ mock_model = mock.MagicMock()
+ mock_build_model.return_value = mock_model
+ mock_model.eval.return_value = mock_model
+ mock_model.to.return_value = mock_model
+
+ extractor = object_extraction.ObjectExtractor(
+ checkpoint_path="mock.pth",
+ confidence_threshold=0.5,
+ score_threshold=0.5,
+ containment_threshold=0.8,
+ max_short_side=800,
+ crop_size=(224, 224),
+ device="cpu",
+ )
+
+ # Mock the internal processor's set_image/set_text_prompt
+ mock_processor = mock.MagicMock()
+ mock_processor.set_image.return_value = {
+ "masks": torch.ones((1, 100, 100), dtype=torch.bool),
+ "masks_logits": torch.ones((1, 100, 100), dtype=torch.float32),
+ "boxes": torch.tensor([[10.0, 10.0, 40.0, 40.0]]),
+ "scores": torch.tensor([0.9]),
+ }
+ mock_processor.set_text_prompt.side_effect = lambda state, prompt: state
+ extractor._processor = mock_processor
+
+ image = Image.new("RGB", (100, 100))
+ resized, state, detections = extractor.extract(image, "bottle")
+
+ # Assertions
+ self.assertIsInstance(resized, Image.Image)
+ self.assertIn("masks", state)
+ self.assertIsInstance(detections, MockDetections)
+ mock_processor.set_image.assert_called_once_with(resized)
+ mock_processor.set_text_prompt.assert_called_once_with(
+ state=mock.ANY, prompt="bottle"
+ )
+
+ def test_move_state_to_cpu(self):
+ state = {
+ "tensor1": torch.tensor([1, 2], device="cpu"),
+ "nested": {
+ "tensor2": torch.tensor([3, 4], device="cpu"),
+ "non_tensor": "string",
+ },
+ "non_tensor2": 10,
+ }
+ extractor: object_extraction.ObjectExtractor = (
+ object_extraction.ObjectExtractor.__new__(
+ object_extraction.ObjectExtractor
+ )
+ )
+ res = extractor._move_state_to_cpu(state)
+ self.assertEqual(res["non_tensor2"], 10)
+ self.assertEqual(res["nested"]["non_tensor"], "string")
+ self.assertFalse(res["tensor1"].is_cuda)
+ self.assertFalse(res["nested"]["tensor2"].is_cuda)
+
+ def test_extract_blended_crop(self):
+ image = Image.new("RGB", (50, 50))
+ state = {
+ "masks": torch.ones((1, 50, 50), dtype=torch.bool),
+ "boxes": torch.tensor([[10.0, 10.0, 40.0, 40.0]]),
+ }
+ crop = object_extraction.ObjectExtractor.extract_blended_crop(
+ image=image, state=state, detection_index=0, crop_size=(30, 30)
+ )
+ self.assertIsInstance(crop, Image.Image)
+ self.assertEqual(crop.size, (30, 30))
+
+ def test_get_cropped_objects_empty_detections(self):
+ track_crop_records = {}
+ object_extraction.ObjectExtractor.get_cropped_objects_for_each_tracking_id(
+ image=Image.new("RGB", (100, 100)),
+ state={},
+ detections=empty_mock_detections(),
+ source_frame_name="frame.png",
+ crop_size=(20, 20),
+ track_crop_records=track_crop_records,
+ )
+ self.assertEqual(track_crop_records, {})
+
+ def test_get_cropped_objects_tracker_id_none(self):
+ track_crop_records = {}
+ detections = MockDetections(
+ xyxy=np.array([[10.0, 15.0, 50.0, 55.0]], dtype=np.float32),
+ confidence=np.array([0.9], dtype=np.float32),
+ class_id=np.array([0]),
+ )
+ detections.tracker_id = None
+ object_extraction.ObjectExtractor.get_cropped_objects_for_each_tracking_id(
+ image=Image.new("RGB", (100, 100)),
+ state={},
+ detections=detections,
+ source_frame_name="frame.png",
+ crop_size=(20, 20),
+ track_crop_records=track_crop_records,
+ )
+ self.assertEqual(track_crop_records, {})
+
+ def test_get_cropped_objects_tracker_id_minus_one(self):
+ track_crop_records = {}
+ detections = MockDetections(
+ xyxy=np.array([[10.0, 15.0, 50.0, 55.0]], dtype=np.float32),
+ confidence=np.array([0.9], dtype=np.float32),
+ class_id=np.array([0]),
+ )
+ detections.tracker_id = np.array([-1])
+ object_extraction.ObjectExtractor.get_cropped_objects_for_each_tracking_id(
+ image=Image.new("RGB", (100, 100)),
+ state={"boxes": torch.tensor([[10.0, 15.0, 50.0, 55.0]])},
+ detections=detections,
+ source_frame_name="frame.png",
+ crop_size=(20, 20),
+ track_crop_records=track_crop_records,
+ )
+ self.assertEqual(track_crop_records, {})
+
+ def test_get_cropped_objects_no_matching_state_box(self):
+ track_crop_records = {}
+ detections = MockDetections(
+ xyxy=np.array([[10.0, 15.0, 50.0, 55.0]], dtype=np.float32),
+ confidence=np.array([0.9], dtype=np.float32),
+ class_id=np.array([0]),
+ )
+ detections.tracker_id = np.array([42])
+ object_extraction.ObjectExtractor.get_cropped_objects_for_each_tracking_id(
+ image=Image.new("RGB", (100, 100)),
+ state={"boxes": torch.tensor([[20.0, 20.0, 60.0, 60.0]])},
+ detections=detections,
+ source_frame_name="frame.png",
+ crop_size=(20, 20),
+ track_crop_records=track_crop_records,
+ )
+ self.assertEqual(track_crop_records, {})
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/pet_grade_classifier.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/pet_grade_classifier.py
new file mode 100644
index 00000000000..a204905ab86
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/pet_grade_classifier.py
@@ -0,0 +1,343 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""DINOv3-based grade classifier: model definition and inference class.
+
+Checkpoints must be produced by our training scripts
+(``train_classifier_repurpose_finetune_2.py`` or
+``train_classifier_repurpose_linear_probe.py``). The pooling strategy
+('cls' or 'cls_mean_patch') is auto-detected from the saved head's
+input dimension; no manual override is required.
+
+Image preprocessing matches the training-time validation transform:
+``Resize((image_size, image_size))`` + ``ToTensor`` + ImageNet normalization.
+
+Class names are NOT stored in our checkpoints — the caller must pass them in.
+"""
+
+from typing import Any, Sequence
+
+import numpy as np
+from PIL import Image
+import torch
+from torch import nn
+import torch.nn.functional as torch_functional
+from torchvision import transforms
+
+IMG_MEAN = (0.485, 0.456, 0.406)
+IMG_STD = (0.229, 0.224, 0.225)
+DEFAULT_INFERENCE_IMAGE_SIZE = 256
+
+POOLING_CLS = "cls"
+POOLING_CLS_MEAN_PATCH = "cls_mean_patch"
+SUPPORTED_POOLING_STRATEGIES = (POOLING_CLS, POOLING_CLS_MEAN_PATCH)
+
+
+def _load_backbone(dinov3_repo_dir: str, model_name: str) -> nn.Module:
+ """Loads the DINOv3 backbone from a local repository.
+
+ Args:
+ dinov3_repo_dir: Path to the cloned DINOv3 repository.
+ model_name: Name of the DINOv3 backbone variant to load.
+
+ Returns:
+ The loaded DINOv3 backbone model.
+ """
+ return torch.hub.load(
+ dinov3_repo_dir,
+ model_name,
+ source="local",
+ pretrained=False,
+ )
+
+
+class Dinov3Classification(nn.Module):
+ """DINOv3 backbone with a linear classification head.
+
+ Mirrors the training-side class in ``models.py``. The pooling
+ strategy controls the feature vector fed to the head:
+
+ - 'cls': CLS token only; head input = backbone hidden size.
+ - 'cls_mean_patch': CLS + mean patch tokens; head input = 2x hidden size.
+ """
+
+ def __init__(
+ self,
+ dinov3_repo_dir: str,
+ model_name: str,
+ num_classes: int,
+ pooling: str = POOLING_CLS,
+ ):
+ """Initializes the Dinov3Classification model.
+
+ Args:
+ dinov3_repo_dir: Path to the cloned DINOv3 repository.
+ model_name: Name of the DINOv3 backbone variant to load.
+ num_classes: The number of output classes for the classification head.
+ pooling: The pooling strategy to use ('cls' or 'cls_mean_patch').
+ Defaults to 'cls'.
+
+ Raises:
+ ValueError: If an unsupported pooling strategy is provided.
+ """
+ super().__init__()
+
+ if pooling not in SUPPORTED_POOLING_STRATEGIES:
+ raise ValueError(
+ f"Unsupported pooling strategy: {pooling!r}. "
+ f"Expected one of {SUPPORTED_POOLING_STRATEGIES}."
+ )
+
+ self.pooling = pooling
+ self.backbone_model = _load_backbone(
+ dinov3_repo_dir=dinov3_repo_dir,
+ model_name=model_name,
+ )
+
+ backbone_hidden_size = self.backbone_model.norm.normalized_shape[0]
+ head_input_features = (
+ backbone_hidden_size
+ if pooling == POOLING_CLS
+ else 2 * backbone_hidden_size
+ )
+ self.head = nn.Linear(
+ in_features=head_input_features,
+ out_features=num_classes,
+ bias=True,
+ )
+
+ def extract_features(self, x: torch.Tensor) -> torch.Tensor:
+ if self.pooling == POOLING_CLS:
+ return self.backbone_model(x)
+
+ token_features = self.backbone_model.forward_features(x)
+ cls_token = token_features["x_norm_clstoken"]
+ patch_tokens = token_features["x_norm_patchtokens"]
+ return torch.cat([cls_token, patch_tokens.mean(dim=1)], dim=1)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return self.head(self.extract_features(x))
+
+
+class DinoClassifier:
+ """Loads a trained DINOv3 classifier and runs batched inference.
+
+ Usage::
+
+ classifier = GradeClassifier(
+ classifier_checkpoint_path="model.pth",
+ dinov3_repo_dir="/path/to/dinov3",
+ model_name="dinov3_vitl16",
+ class_names=["grade_a", "grade_b", "grade_c"],
+ device=torch.device("cuda"),
+ )
+ results = classifier.predict(pil_images)
+ """
+
+ def __init__(
+ self,
+ classifier_checkpoint_path: str,
+ dinov3_repo_dir: str,
+ model_name: str,
+ class_names: Sequence[str],
+ device: str | None = None,
+ image_size: int = DEFAULT_INFERENCE_IMAGE_SIZE,
+ ):
+ """Loads the checkpoint and prepares the model for inference.
+
+ Args:
+ classifier_checkpoint_path: Path to a checkpoint produced by our
+ training scripts. Must contain ``model_state_dict``.
+ dinov3_repo_dir: Path to the cloned Facebook DINOv3 repository.
+ model_name: DINOv3 backbone variant (e.g. ``'dinov3_vitl16'``). Must
+ match the backbone the checkpoint was trained with.
+ class_names: Class names in label-index order. Not stored in
+ checkpoints, so the caller must supply them.
+ device: Torch device on which the model will run.
+ image_size: Square input size for the eval transform. Must be a multiple
+ of the backbone's patch size (16 for vit*16).
+
+ Raises:
+ KeyError: If the checkpoint is missing ``model_state_dict`` or
+ ``head.weight``.
+ ValueError: If the saved head dimension cannot be matched to a
+ known pooling strategy for this backbone.
+ """
+ self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
+ self.class_names = list(class_names)
+
+ saved_state_dict = self._read_state_dict(
+ classifier_checkpoint_path,
+ self.device,
+ )
+ self.pooling = self._detect_pooling(
+ saved_state_dict,
+ dinov3_repo_dir,
+ model_name,
+ )
+ print(f"[INFO]: Detected pooling strategy from checkpoint: {self.pooling}")
+
+ self.model = Dinov3Classification(
+ dinov3_repo_dir=dinov3_repo_dir,
+ model_name=model_name,
+ num_classes=len(class_names),
+ pooling=self.pooling,
+ ).to(self.device)
+ self.model.load_state_dict(saved_state_dict)
+ self.model.eval()
+
+ self._eval_transform = self._build_eval_transform(image_size)
+
+ # ------------------------------------------------------------------
+ # Inference
+ # ------------------------------------------------------------------
+
+ @torch.no_grad()
+ def predict(self, pil_images: Sequence[Image.Image]) -> list[dict[str, Any]]:
+ """Classifies a list of PIL images in a single forward pass.
+
+ Args:
+ pil_images: Sequence of RGB PIL images (length >= 1).
+
+ Returns:
+ One prediction dict per input image, in the same order::
+
+ {
+ "predicted_class": str,
+ "predicted_probability": float, # 0-100
+ "all_probabilities": dict[str, float],
+ }
+
+ Raises:
+ ValueError: If ``pil_images`` is empty.
+ """
+ if not pil_images:
+ raise ValueError("predict requires at least one image.")
+
+ batch = torch.stack([self._eval_transform(img) for img in pil_images]).to(
+ self.device
+ )
+
+ logits = self.model(batch)
+ probabilities = torch_functional.softmax(logits, dim=1).cpu().numpy()
+
+ predictions = []
+ for probability_row in probabilities:
+ predicted_index = int(np.argmax(probability_row))
+ predictions.append({
+ "predicted_class": self.class_names[predicted_index],
+ "predicted_probability": float(
+ probability_row[predicted_index] * 100.0
+ ),
+ "all_probabilities": {
+ name: float(p * 100.0)
+ for name, p in zip(self.class_names, probability_row)
+ },
+ })
+ return predictions
+
+ # ------------------------------------------------------------------
+ # Private helpers
+ # ------------------------------------------------------------------
+
+ @staticmethod
+ def _read_state_dict(
+ checkpoint_path: str, device: torch.device
+ ) -> dict[str, Any]:
+ """Reads and validates the model state dict from a checkpoint.
+
+ Args:
+ checkpoint_path: Path to the checkpoint file.
+ device: Torch device to map the loaded parameters to.
+
+ Returns:
+ The state dict containing model weights.
+
+ Raises:
+ KeyError: If 'model_state_dict' or 'head.weight' is missing.
+ """
+ checkpoint = torch.load(
+ checkpoint_path,
+ map_location=device,
+ weights_only=False,
+ )
+ if "model_state_dict" not in checkpoint:
+ raise KeyError(
+ f"Checkpoint at '{checkpoint_path}' is missing "
+ "'model_state_dict'. Was it produced by our training scripts?"
+ )
+ state_dict = checkpoint["model_state_dict"]
+ if "head.weight" not in state_dict:
+ raise KeyError(
+ "Checkpoint state dict is missing 'head.weight'; cannot infer "
+ "pooling. Was the checkpoint produced by our Dinov3Classification?"
+ )
+ return state_dict
+
+ @staticmethod
+ def _detect_pooling(
+ state_dict: dict[str, Any],
+ dinov3_repo_dir: str,
+ model_name: str,
+ ) -> str:
+ """Infers the pooling strategy from the shape of the saved head's weights.
+
+ Args:
+ state_dict: The model state dict loaded from checkpoint.
+ dinov3_repo_dir: Path to the cloned DINOv3 repository.
+ model_name: Name of the DINOv3 backbone variant.
+
+ Returns:
+ The inferred pooling strategy ('cls' or 'cls_mean_patch').
+
+ Raises:
+ ValueError: If the pooling strategy cannot be inferred from the head
+ dimensions.
+ """
+ # Load a throwaway backbone only to read its hidden size, then discard it.
+ probe = _load_backbone(
+ dinov3_repo_dir=dinov3_repo_dir, model_name=model_name
+ )
+ backbone_hidden_size = probe.norm.normalized_shape[0]
+ del probe
+
+ head_input_features = state_dict["head.weight"].shape[1]
+ if head_input_features == backbone_hidden_size:
+ return POOLING_CLS
+ if head_input_features == 2 * backbone_hidden_size:
+ return POOLING_CLS_MEAN_PATCH
+
+ raise ValueError(
+ "Cannot infer pooling strategy: head input dim "
+ f"{head_input_features} matches neither {backbone_hidden_size} "
+ f"(cls) nor {2 * backbone_hidden_size} (cls_mean_patch). "
+ "This usually means the checkpoint was trained with a different "
+ "backbone than the one configured here."
+ )
+
+ @staticmethod
+ def _build_eval_transform(image_size: int) -> transforms.Compose:
+ """Builds the transformation pipeline for evaluation image preprocessing.
+
+ Args:
+ image_size: Target square image size.
+
+ Returns:
+ A torchvision transforms Compose object.
+ """
+ return transforms.Compose([
+ transforms.Resize((image_size, image_size)),
+ transforms.ToTensor(),
+ transforms.Normalize(mean=IMG_MEAN, std=IMG_STD),
+ ])
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/pet_grade_classifier_test.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/pet_grade_classifier_test.py
new file mode 100644
index 00000000000..ae05fc12da5
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/pet_grade_classifier_test.py
@@ -0,0 +1,317 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for pet_grade_classifier module."""
+
+import unittest
+from unittest import mock
+
+from absl.testing import parameterized
+from PIL import Image
+import torch
+from torch import nn
+from torchvision import transforms
+
+from official.projects.waste_identification_ml.Deploy.pet_grading_cloud_deployment import pet_grade_classifier
+
+
+MODULE_PATH = pet_grade_classifier.__name__
+
+
+class MockBackbone(nn.Module):
+
+ def __init__(self, hidden_size: int = 768):
+ super().__init__()
+ self.norm = nn.LayerNorm(hidden_size)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return torch.zeros((x.shape[0], self.norm.normalized_shape[0]))
+
+ def forward_features(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
+ hidden_size = self.norm.normalized_shape[0]
+ return {
+ "x_norm_clstoken": torch.zeros((x.shape[0], hidden_size)),
+ "x_norm_patchtokens": torch.zeros((x.shape[0], 10, hidden_size)),
+ }
+
+
+class PetGradeClassifierTest(parameterized.TestCase):
+
+ @mock.patch(f"{MODULE_PATH}.torch.hub.load")
+ def test_load_backbone(self, mock_hub_load):
+ # Act
+ pet_grade_classifier._load_backbone("mock_repo", "mock_model_name")
+
+ # Assert
+ mock_hub_load.assert_called_once_with(
+ "mock_repo",
+ "mock_model_name",
+ source="local",
+ pretrained=False,
+ )
+
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ def test_dinov3_classification_init_unsupported_pooling(self):
+ with self.assertRaises(ValueError) as context:
+ _ = pet_grade_classifier.Dinov3Classification(
+ dinov3_repo_dir="mock_repo",
+ model_name="mock_model",
+ num_classes=3,
+ pooling="unsupported_pooling",
+ )
+
+ self.assertIn("Unsupported pooling strategy", str(context.exception))
+
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ def test_dinov3_classification_forward_pooling_cls(self, mock_load):
+ # Arrange
+ mock_load.return_value = MockBackbone(hidden_size=256)
+ model = pet_grade_classifier.Dinov3Classification(
+ dinov3_repo_dir="mock_repo",
+ model_name="mock_model",
+ num_classes=3,
+ pooling="cls",
+ )
+ x = torch.zeros((2, 3, 224, 224))
+
+ # Act
+ out = model(x)
+
+ # Assert
+ self.assertEqual(out.shape, (2, 3))
+
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ def test_dinov3_classification_forward_pooling_cls_mean_patch(
+ self, mock_load
+ ):
+ # Arrange
+ mock_load.return_value = MockBackbone(hidden_size=256)
+ model = pet_grade_classifier.Dinov3Classification(
+ dinov3_repo_dir="mock_repo",
+ model_name="mock_model",
+ num_classes=3,
+ pooling="cls_mean_patch",
+ )
+ # Head input is 2x hidden = 512
+ self.assertEqual(model.head.in_features, 512)
+ x = torch.zeros((2, 3, 224, 224))
+
+ # Act
+ out = model(x)
+
+ # Assert
+ self.assertEqual(out.shape, (2, 3))
+
+ @mock.patch(f"{MODULE_PATH}.torch.load")
+ def test_dino_classifier_read_state_dict_missing_model_state_dict(
+ self, mock_torch_load
+ ):
+ mock_torch_load.return_value = {}
+
+ with self.assertRaises(KeyError) as context:
+ pet_grade_classifier.DinoClassifier._read_state_dict(
+ "mock_ckpt", torch.device("cpu")
+ )
+
+ self.assertIn("missing 'model_state_dict'", str(context.exception))
+
+ @mock.patch(f"{MODULE_PATH}.torch.load")
+ def test_dino_classifier_read_state_dict_missing_head_weight(
+ self, mock_torch_load
+ ):
+ mock_torch_load.return_value = {"model_state_dict": {}}
+
+ with self.assertRaises(KeyError) as context:
+ pet_grade_classifier.DinoClassifier._read_state_dict(
+ "mock_ckpt", torch.device("cpu")
+ )
+
+ self.assertIn("missing 'head.weight'", str(context.exception))
+
+ @parameterized.named_parameters(
+ ("cls", torch.zeros((3, 256)), "cls"),
+ ("cls_mean_patch", torch.zeros((3, 512)), "cls_mean_patch"),
+ )
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ def test_detect_pooling(self, head_weight, expected_pooling, mock_load):
+ mock_load.return_value = MockBackbone(hidden_size=256)
+ state_dict = {"head.weight": head_weight}
+
+ pooling = pet_grade_classifier.DinoClassifier._detect_pooling(
+ state_dict, "mock_repo", "mock_model"
+ )
+
+ self.assertEqual(pooling, expected_pooling)
+
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ def test_detect_pooling_value_error(self, mock_load):
+ mock_load.return_value = MockBackbone(hidden_size=256)
+ # Target head dim 400 doesn't match 256 or 512
+ state_dict = {"head.weight": torch.zeros((3, 400))}
+
+ with self.assertRaises(ValueError) as context:
+ pet_grade_classifier.DinoClassifier._detect_pooling(
+ state_dict, "mock_repo", "mock_model"
+ )
+
+ self.assertIn("Cannot infer pooling strategy", str(context.exception))
+
+ @mock.patch(f"{MODULE_PATH}.Dinov3Classification.load_state_dict")
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ @mock.patch(f"{MODULE_PATH}.DinoClassifier._read_state_dict")
+ def test_dino_classifier_init_success(
+ self, mock_read_state, mock_load, mock_load_state_dict
+ ):
+ # Arrange
+ state_dict = {"head.weight": torch.zeros((3, 256))}
+ mock_read_state.return_value = state_dict
+ mock_load.return_value = MockBackbone(hidden_size=256)
+
+ # Act
+ classifier = pet_grade_classifier.DinoClassifier(
+ classifier_checkpoint_path="mock_ckpt.pth",
+ dinov3_repo_dir="mock_repo",
+ model_name="mock_model",
+ class_names=["class_a", "class_b", "class_c"],
+ device="cpu",
+ )
+
+ # Assert
+ self.assertEqual(classifier.pooling, "cls")
+ self.assertEqual(classifier.class_names, ["class_a", "class_b", "class_c"])
+ mock_load_state_dict.assert_called_once_with(state_dict)
+
+ @mock.patch(f"{MODULE_PATH}.Dinov3Classification.load_state_dict")
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ @mock.patch(f"{MODULE_PATH}.DinoClassifier._read_state_dict")
+ def test_predict_empty_input(self, mock_read_state, mock_load):
+ # Arrange
+ state_dict = {"head.weight": torch.zeros((3, 256))}
+ mock_read_state.return_value = state_dict
+ mock_load.return_value = MockBackbone(hidden_size=256)
+
+ classifier = pet_grade_classifier.DinoClassifier(
+ classifier_checkpoint_path="mock_ckpt.pth",
+ dinov3_repo_dir="mock_repo",
+ model_name="mock_model",
+ class_names=["class_a", "class_b", "class_c"],
+ device="cpu",
+ )
+
+ # Act & Assert
+ with self.assertRaises(ValueError) as context:
+ classifier.predict([])
+
+ self.assertIn("predict requires at least one image", str(context.exception))
+
+ @mock.patch(f"{MODULE_PATH}.Dinov3Classification.load_state_dict")
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ @mock.patch(f"{MODULE_PATH}.DinoClassifier._read_state_dict")
+ def test_predict_success(self, mock_read_state, mock_load):
+ # Arrange
+ state_dict = {"head.weight": torch.zeros((2, 256))}
+ mock_read_state.return_value = state_dict
+ mock_load.return_value = MockBackbone(hidden_size=256)
+
+ classifier = pet_grade_classifier.DinoClassifier(
+ classifier_checkpoint_path="mock_ckpt.pth",
+ dinov3_repo_dir="mock_repo",
+ model_name="mock_model",
+ class_names=["class_a", "class_b"],
+ device="cpu",
+ )
+
+ # Mock model call: must return logits shape [num_images, num_classes]
+ # For image 1: logits [10.0, 0.0] -> class_a is highly likely.
+ # For image 2: logits [0.0, 10.0] -> class_b is highly likely.
+ mock_logits = torch.tensor([[10.0, 0.0], [0.0, 10.0]])
+ classifier.model = mock.MagicMock(return_value=mock_logits)
+
+ image1 = Image.new("RGB", (100, 100))
+ image2 = Image.new("RGB", (100, 100))
+
+ # Act
+ predictions = classifier.predict([image1, image2])
+
+ # Assert
+ self.assertLen(predictions, 2)
+
+ # Image 1 prediction
+ self.assertEqual(predictions[0]["predicted_class"], "class_a")
+ self.assertGreater(predictions[0]["predicted_probability"], 99.0)
+ self.assertIn("class_a", predictions[0]["all_probabilities"])
+ self.assertIn("class_b", predictions[0]["all_probabilities"])
+
+ # Image 2 prediction
+ self.assertEqual(predictions[1]["predicted_class"], "class_b")
+ self.assertGreater(predictions[1]["predicted_probability"], 99.0)
+
+ def test_build_eval_transform(self):
+ transform = pet_grade_classifier.DinoClassifier._build_eval_transform(224)
+ self.assertIsInstance(transform, transforms.Compose)
+
+ # Test on a dummy PIL Image
+ img = Image.new("RGB", (100, 100))
+ tensor = transform(img)
+ self.assertEqual(tensor.shape, (3, 224, 224))
+
+ @parameterized.named_parameters(
+ ("cpu", False, "cpu"),
+ ("cuda", True, "cuda"),
+ )
+ @mock.patch(f"{MODULE_PATH}.torch.cuda.is_available")
+ @mock.patch(f"{MODULE_PATH}.Dinov3Classification.load_state_dict")
+ @mock.patch(f"{MODULE_PATH}._load_backbone")
+ @mock.patch(f"{MODULE_PATH}.DinoClassifier._read_state_dict")
+ @mock.patch(f"{MODULE_PATH}.Dinov3Classification.to")
+ def test_dino_classifier_device_default(
+ self,
+ cuda_available,
+ expected_device,
+ mock_to,
+ mock_read_state,
+ mock_load,
+ mock_cuda_available,
+ ):
+ mock_cuda_available.return_value = cuda_available
+ state_dict = {"head.weight": torch.zeros((3, 256))}
+ mock_read_state.return_value = state_dict
+ mock_load.return_value = MockBackbone(hidden_size=256)
+ mock_to.side_effect = lambda dev: mock.MagicMock()
+
+ classifier = pet_grade_classifier.DinoClassifier(
+ classifier_checkpoint_path="mock_ckpt.pth",
+ dinov3_repo_dir="mock_repo",
+ model_name="mock_model",
+ class_names=["class_a", "class_b"],
+ )
+ self.assertEqual(classifier.device, expected_device)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/requirements.sh b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/requirements.sh
new file mode 100644
index 00000000000..c596c111317
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/requirements.sh
@@ -0,0 +1,38 @@
+#!/bin/bash
+
+# Summary
+cat << EOF
+This script sets up the environment for running ML models by ensuring Bash
+execution, installing system dependencies, setting up a virtual environment,
+installing ML packages, and cloning TensorFlow Model Garden.
+EOF
+
+# Ensure the script is executed with /bin/bash
+if [ -z "$BASH_VERSION" ]; then
+ exec /bin/bash "$0" "$@"
+fi
+
+# update linux packages
+sudo apt-get update -y
+
+# Create a virtual environment and install packages
+sudo apt-get install -y python3-venv python3-pip
+
+# Clone dinvov3 repo
+git clone https://github.com/facebookresearch/dinov3.git
+
+# Download Model weights
+mkdir -p ./model_weights
+[ -f ./model_weights/sam3_original_weights_sam3.pt ] || wget -P ./model_weights https://storage.googleapis.com/tf_model_garden/vision/waste_identification_ml/dairy_product_packet_detection/sam3_original_weights_sam3.pt
+[ -f ./model_weights/best_pet_grading_model.pth ] || wget -P ./model_weights https://storage.googleapis.com/tf_model_garden/vision/waste_identification_ml/dairy_product_packet_detection/best_pet_grading_model.pth
+
+python3.10 -m venv myenv
+source myenv/bin/activate
+
+echo "Activated python environment, installing dependencies."
+
+pip install -r requirements.txt
+pip install numpy==1.26.4
+
+deactivate
+echo "Environment setup is complete."
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/requirements.txt b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/requirements.txt
new file mode 100644
index 00000000000..58d27b35852
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/requirements.txt
@@ -0,0 +1,27 @@
+--extra-index-url https://download.pytorch.org/whl/cu124
+sam3 @ git+https://github.com/facebookresearch/sam3.git@967fdd651f71ca14949122fed4c918a778ca9334
+natsort
+absl-py==2.4.0
+einops==0.8.2
+google-api-core==2.30.3
+google-auth==2.49.2
+google-auth-oauthlib==1.3.1
+google-cloud-bigquery==3.41.0
+google-cloud-core==2.5.1
+google-cloud-storage==3.10.1
+huggingface_hub==1.11.0
+opencv-python>=4.8.0,<4.10.0
+pandas==2.3.3
+pandas-gbq==0.35.0
+pycocotools
+scikit-image
+scikit-learn
+scipy
+seaborn
+supervision
+matplotlib
+termcolor
+torch==2.6.0+cu124
+torchvision==0.21.0+cu124
+torchmetrics
+trackers
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/run_images.sh b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/run_images.sh
new file mode 100644
index 00000000000..d4ff99eeb43
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/run_images.sh
@@ -0,0 +1,40 @@
+#!/bin/bash
+
+cat << EOF
+This script automates the execution of an Circularnet pipeline for image
+processing.
+Steps Performed:
+ 1. Activates the Python virtual environment named 'myenv'.
+ 2. Validates successful activation of the virtual environment.
+ 3. Executes the 'pipeline_images.py' script with the following parameters:
+ Parameters:
+ --input_directory : GCS directory where the input images are stored for
+ inference.
+ --output_directory : GCS directory where the model inference outputs will be
+ saved.
+ --model_name : Name of the model to download and use for inference.
+ --threshold : Confidence threshold for detections during
+ inference.
+ --project_id : Google Cloud Project ID for BigQuery operations.
+ --bq_dataset_id : BigQuery Dataset ID where results will be stored.
+ --bq_table_id : BigQuery Table ID where results will be stored.
+ --overwrite : If set to True, overwrites the pre-existing
+ BigQuery table.
+EOF
+#Activate the virtual environment
+source myenv/bin/activate
+# Check if the virtual environment is activated
+if [[ "$VIRTUAL_ENV" != "" ]]; then
+ echo "Virtual environment 'myenv' activated successfully."
+else
+ echo "Failed to activate virtual environment. Exiting."
+ exit 1
+fi
+
+python inference_pipeline.py \
+ --input_directory=gs://circularnet_data/tmp/pet_pipeline_test \
+ --output_directory=gs://circularnet_data/tmp/pet_pipeline_test \
+ --project_id=waste-identification-ml-330916 \
+ --bq_dataset_id=pet_dataset \
+ --bq_table_id=pet_pipeline_test \
+ --overwrite=True
diff --git a/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/utils.py b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/utils.py
new file mode 100644
index 00000000000..6eb93befabd
--- /dev/null
+++ b/official/projects/waste_identification_ml/Deploy/pet_grading_cloud_deployment/utils.py
@@ -0,0 +1,230 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utility functions for the pet grading cloud deployment."""
+
+import collections
+import datetime
+import gc
+import itertools
+import logging
+import math
+import os
+import pathlib
+import re
+import subprocess
+import matplotlib.pyplot as plt
+import natsort
+import numpy as np
+
+
+def _create_log_file(name: str, logs_folder_path: str) -> logging.Logger:
+ """Creates a logger and a log file given the name of the video.
+
+ Args:
+ name: The name of the video.
+ logs_folder_path: Path to the directory where logs should be saved.
+
+ Returns:
+ logging.Logger: Logger object configured to write logs to the file.
+ """
+ log_file_path = os.path.join(logs_folder_path, f"{name}.log")
+ logger = logging.getLogger(name)
+ logger.setLevel(logging.INFO)
+ file_handler = logging.FileHandler(log_file_path)
+ formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s")
+ file_handler.setFormatter(formatter)
+ logger.addHandler(file_handler)
+ return logger
+
+
+def parse_imgage_names_to_datetime(filename):
+ """Parses an image filename to extract a datetime object.
+
+ The filename is expected to be in the format "img_YYYYMMDD_HHMMSSms.*".
+
+ Args:
+ filename: The path to the image file.
+
+ Returns:
+ A datetime.datetime object parsed from the filename.
+ """
+ stem = pathlib.Path(filename).stem
+ m = re.match(r"img_(\d{8})_(\d{6})(\d*)", stem)
+ dt_str = f"{m.group(1)}_{m.group(2)}{m.group(3)}"
+ return datetime.datetime.strptime(dt_str, "%Y%m%d_%H%M%S%f")
+
+
+def files_paths(folder_path):
+ """List the full paths of image files in a folder and sort them.
+
+ Args:
+ folder_path: The path of the folder to list the image files from.
+
+ Returns:
+ A list of full paths of the image files in the folder, sorted in ascending
+ order.
+ """
+ img_extensions = (".jpg", ".jpeg", ".png", ".gif", ".bmp", ".tiff", ".webp")
+ image_files_full_path = []
+ for entry in os.scandir(folder_path):
+ if entry.is_file() and entry.name.lower().endswith(img_extensions):
+ image_files_full_path.append(entry.path)
+
+ try:
+ image_files_full_path = sorted(
+ image_files_full_path, key=parse_imgage_names_to_datetime
+ )
+ except AttributeError:
+ print("Failed to parse image names to datetime, using natsort instead.")
+ image_files_full_path = natsort.natsorted(image_files_full_path)
+ return image_files_full_path
+
+
+def setup_logger_and_directories(input_dir):
+ """Sets up directories and a logger for the pipeline.
+
+ Args:
+ input_dir: The path to the input directory on GCP.
+
+ Returns:
+ A tuple containing:
+ - input_directory: The local path of the copied input directory.
+ - prediction_folder: The path to the created prediction folder.
+ - logger: The configured logging.Logger object.
+ """
+
+ input_directory = (input_dir).rstrip("/\\")
+ command = f"gcloud storage cp -r {input_directory} ."
+ subprocess.run(command, shell=True, check=True)
+ prediction_folder = os.path.basename(input_directory) + "_prediction"
+ os.makedirs(prediction_folder, exist_ok=True)
+ log_name = os.path.basename(input_dir)
+ log_folder = os.path.join(os.getcwd(), "logs")
+ os.makedirs(log_folder, exist_ok=True)
+ logger = _create_log_file(log_name, log_folder)
+ return input_directory, prediction_folder, logger
+
+
+def get_class_with_majority_vote(list_of_predictions, min_probability=50):
+ """Determines consensus predictions for each tracker using majority voting.
+
+ Args:
+ list_of_predictions: A list of prediction dictionaries, where each must
+ contain 'tracker_id', 'predicted_class', and 'predicted_probability'.
+ min_probability: The minimum required probability threshold for a tracker to
+ be considered valid. Default is 50.
+
+ Returns:
+ A list of prediction dictionaries for the trackers that met the minimum
+ probability threshold, with each dictionary updated with its tracker's
+ 'best_class' and 'best_probability'.
+ """
+ tracker_probs = collections.defaultdict(lambda: collections.defaultdict(list))
+ for p in list_of_predictions:
+ tracker_probs[p["tracker_id"]][p["predicted_class"]].append(
+ p["predicted_probability"]
+ )
+
+ tracker_id_to_best_class_dict = {}
+ for tid, class_probs in tracker_probs.items():
+ vote_counts = {c: len(probs) for c, probs in class_probs.items()}
+ max_votes = max(vote_counts.values())
+ max_voted_class = [c for c, v in vote_counts.items() if v == max_votes]
+
+ if len(max_voted_class) == 1:
+ best_class = max_voted_class[0]
+ else:
+ # In case of tiebreak, choose by average probability
+ best_class = max(
+ max_voted_class,
+ key=lambda c, class_probs=class_probs: max(class_probs[c]),
+ )
+
+ best_probability = max(class_probs[best_class])
+ tracker_id_to_best_class_dict[tid] = {
+ "best_class": best_class,
+ "best_probability": best_probability,
+ }
+
+ valid_trackers = {
+ tid: info
+ for tid, info in tracker_id_to_best_class_dict.items()
+ if info["best_probability"] >= min_probability
+ }
+
+ filtered = [
+ p for p in list_of_predictions if p["tracker_id"] in valid_trackers
+ ]
+ for p in filtered:
+ p.update(valid_trackers[p["tracker_id"]])
+
+ return filtered
+
+
+def save_output_image_grids(
+ all_predictions, output_dr, columns_per_row=5, thumbnail_size=3
+):
+ """Saves a grid of thumbnail images with prediction labels for each frame.
+
+ Args:
+ all_predictions: A list of prediction dictionaries, where each contains
+ 'frame_name', 'tracker_id', 'best_class', 'best_probability', and 'crop'
+ (the image array).
+ output_dr: The output directory path where the saved image files are stored.
+ columns_per_row: The maximum number of image columns in a grid row. Default
+ is 5.
+ thumbnail_size: The figure height/width factor (in inches) for each cell in
+ the grid. Default is 3.
+ """
+ try:
+
+ keyfunc = lambda p: p["frame_name"]
+ sorted_preds = sorted(all_predictions, key=keyfunc)
+
+ for frame_name, group in itertools.groupby(sorted_preds, key=keyfunc):
+ rows = list(group)
+ n = len(rows)
+
+ ncols = min(n, columns_per_row)
+ nrows = math.ceil(n / ncols)
+
+ fig, axes = plt.subplots(
+ nrows, ncols, figsize=(ncols * thumbnail_size, nrows * thumbnail_size)
+ )
+ fig.suptitle(frame_name, fontsize=9, fontweight="bold")
+
+ if n == 1:
+ axes = np.array([axes])
+ axes = axes.flatten()
+
+ for i, p in enumerate(rows):
+ title = (
+ f"#{p['tracker_id']} {p['best_class']}"
+ f"\n{p['best_probability']:.3f}%"
+ )
+ axes[i].imshow(p["crop"])
+ axes[i].set_title(title, fontsize=7)
+ axes[i].axis("off")
+
+ for j in range(n, len(axes)):
+ axes[j].axis("off")
+
+ plt.tight_layout()
+ fig.savefig(f"{output_dr}/{frame_name}")
+ plt.close(fig)
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ print(f"Issue in saving visualization of results, due to error : {e}")
+ finally:
+ gc.collect()
diff --git a/official/projects/waste_identification_ml/README.md b/official/projects/waste_identification_ml/README.md
new file mode 100644
index 00000000000..57375bb94b0
--- /dev/null
+++ b/official/projects/waste_identification_ml/README.md
@@ -0,0 +1,86 @@
+# CircularNet
+
+Instance segmentation models for identification of recyclables on conveyor
+belts.
+
+We provide retraining and fine-tuning utilities, but if you're interested in
+partnering more closely with us reach out to
+waste-innovation-external@google.com
+
+## Overview
+
+Circularnet is built using RF-DETR, a vision transformer model that includes
+both object detection and instance segmentation, which is a deep learning model
+for instance image segmentation, where the goal is to assign instance level
+labels (e.g. person1, person2, cat) to every pixel in an input image.
+
+## Model Categories
+
+- **Material Type:** Identifies the material type (metal, paper etc) of an
+ object. For plastic, resin types are also identified (HDPE, PET, LDPE, etc).
+- **Material Form:** Categorizes objects based on the form factor (cup,
+ bottle, bag etc)
+- **Example inference label:** Plastics-PET_Bottle
+
+### Latest model
+### Single unified model that performs material type and form detections
+
+Model categories | Model backbone | Model type | GCP bucket path |
+| ------ | ------ | ----- | ------ |
+Material Type & Form | Vision transformer | onnx model | [click here](https://storage.googleapis.com/tf_model_garden/vision/waste_identification_ml/CircularNet_Segmentation_Model_v1.zip)
+
+## Full Documentation
+
+The full documentation, covering everything from how to choose and install
+a camera to how to prepare and make use of the model is **[here](circularnet-docs/content/_index.md).**
+Below, we also provide a quicker guide for running inference using a GCP VM,
+assuming you already have a working camera taking pictures.
+
+## End to End Cloud Deployment Guide
+
+End to end deployment involves three key steps:
+
+1. **GCP GPU VM creation**
+
+2. **Code configuration**
+
+3. **Results analysis**
+
+We will go through each one of them in details below
+
+#### [A] Prerequisite - Create VM instance:
+Create a Google cloud account and a T4 GPU enabled VM:
+
+- [Create VM in GCP Cloud](circularnet-docs/content/deploy-cn/before-you-begin.md)
+
+#### [B] Code Setup - Clone and start the pipeline
+
+Run the following commands mentioned in each step on the **SSH-in-browser**
+window of your VM instance in Google Cloud
+
+Step 1:
+
+- [Clone the repository](circularnet-docs/content/deploy-cn/clone-repo.md)
+
+Step 2:
+
+- [Start the server](circularnet-docs/content/deploy-cn/start-server.md)
+
+Step 3:
+
+- [Run the prediction Pipeline](circularnet-docs/content/deploy-cn/start-client.md)
+
+For more details: [Click Here](circularnet-docs/content/analyze-data/prediction-pipeline-in-cloud.md)
+
+#### [C] Setup Dashboard - Visualize results
+
+For reporting purposes and to analyze image categories, we need to set up and
+connect looker dashboard with BigQuery table:
+
+- [Prepare and analyze images](circularnet-docs/content/view-data/configure-dashboard.md)
+
+## Authors and Maintainers
+Umair Sabir - Primary Developer
+Vinit Ganorkar - Primary developer
+Ethan Steele - Collaborator
+Sujit Sanjeev - Product Manager
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/big_query_ops.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/big_query_ops.py
new file mode 100644
index 00000000000..3223274b4d4
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/big_query_ops.py
@@ -0,0 +1,114 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Designed to interact with Google BigQuery.
+
+For the purpose of dataset and table management, as well as data ingestion
+from pandas DataFrames.
+"""
+
+import logging
+from google.cloud import bigquery
+from google.cloud import exceptions
+import pandas as pd
+import pandas_gbq
+
+logging.basicConfig(
+ level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
+)
+
+_SCHEMA = [
+ bigquery.SchemaField("particle", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("source_name", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("image_name", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("detection_scores", "FLOAT", mode="REQUIRED"),
+ bigquery.SchemaField("creation_time", "STRING", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_0", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_1", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_2", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("bbox_3", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("detection_classes", "INTEGER", mode="REQUIRED"),
+ bigquery.SchemaField("detection_classes_names", "STRING", mode="REQUIRED"),
+]
+
+
+def create_table(
+ project_id: str,
+ dataset_id: str,
+ table_id: str,
+ overwrite: bool = False, # New optional argument
+) -> None:
+ """Creates a table in a BigQuery dataset.
+
+ Args:
+ project_id: The Google Cloud project ID.
+ dataset_id: The ID of the dataset in which the table is to be created.
+ table_id: The ID of the table to be created.
+ overwrite: If True, deletes the preexisting table before creating a new
+ one.
+ """
+ client = bigquery.Client(project=project_id)
+ dataset_ref = client.dataset(dataset_id)
+
+ try:
+ # Check if the dataset already exists
+ dataset = client.get_dataset(dataset_ref)
+ except exceptions.NotFound:
+ # If the dataset does not exist, create it
+ dataset = bigquery.Dataset(dataset_ref)
+ dataset = client.create_dataset(dataset)
+
+ table_ref = dataset.table(table_id)
+
+ try:
+ # Check if the table already exists
+ table = client.get_table(table_ref)
+ if overwrite:
+ logging.info(
+ "Overwriting table '%s' in dataset '%s'...", table_id, dataset_id
+ )
+ client.delete_table(table_ref)
+ table = bigquery.Table(table_ref, schema=_SCHEMA)
+ client.create_table(table)
+ print(f"Table '{table_id}' has been overwritten.")
+ else:
+ print(f"Table '{table_id}' already exists. Skipping creation.")
+ except exceptions.NotFound:
+ # If the table does not exist, create it
+ table = bigquery.Table(table_ref, schema=_SCHEMA)
+ client.create_table(table)
+ print(f"Table '{table_id}' created successfully.")
+
+
+def ingest_data(
+ df: pd.DataFrame, project_id: str, dataset_id: str, table_id: str
+) -> None:
+ """Ingests data from a pandas DataFrame into a specified BigQuery table.
+
+ This function takes a pandas DataFrame and appends its contents to a BigQuery
+ table
+ identified by the provided dataset and table IDs within the specified project.
+ If the table does not exist, BigQuery automatically creates it with a schema
+ inferred from the DataFrame.
+
+ Args:
+ df: The pandas DataFrame containing the data to be ingested.
+ project_id: The Google Cloud project ID.
+ dataset_id: The ID of the dataset containing the target table.
+ table_id: The ID of the table where the data will be ingested.
+ """
+ table_ref = f"{project_id}.{dataset_id}.{table_id}"
+ pandas_gbq.to_gbq(
+ df, destination_table=table_ref, project_id=project_id, if_exists="append"
+ )
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/big_query_ops_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/big_query_ops_test.py
new file mode 100644
index 00000000000..ab4012c4acb
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/big_query_ops_test.py
@@ -0,0 +1,56 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import big_query_ops
+
+
+class TestSchemaDefinition(unittest.TestCase):
+
+ def test_schema_definition(self):
+ expected_schema = [
+ ("particle", "INTEGER", "REQUIRED"),
+ ("source_name", "STRING", "REQUIRED"),
+ ("image_name", "STRING", "REQUIRED"),
+ ("detection_scores", "FLOAT", "REQUIRED"),
+ ("creation_time", "STRING", "REQUIRED"),
+ ("bbox_0", "INTEGER", "REQUIRED"),
+ ("bbox_1", "INTEGER", "REQUIRED"),
+ ("bbox_2", "INTEGER", "REQUIRED"),
+ ("bbox_3", "INTEGER", "REQUIRED"),
+ ("detection_classes", "INTEGER", "REQUIRED"),
+ ("detection_classes_names", "STRING", "REQUIRED"),
+ ]
+
+ # Check schema length
+ self.assertEqual(
+ len(big_query_ops._SCHEMA),
+ len(expected_schema),
+ "Schema length mismatch.",
+ )
+
+ # Validate each field's name, type, and mode in order
+ for idx, (field, expected) in enumerate(
+ zip(big_query_ops._SCHEMA, expected_schema)
+ ):
+ expected_name, expected_type, expected_mode = expected
+ self.assertEqual(
+ (field.name, field.field_type, field.mode),
+ (expected_name, expected_type, expected_mode),
+ f"Mismatch at field index {idx}.",
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/feature_extraction.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/feature_extraction.py
new file mode 100644
index 00000000000..679422a582f
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/feature_extraction.py
@@ -0,0 +1,76 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Extract properties of the mask."""
+
+import numpy as np
+import pandas as pd
+import skimage.measure
+
+_PROPERTIES = (
+ 'area',
+ 'bbox',
+ 'convex_area',
+ 'bbox_area',
+ 'major_axis_length',
+ 'minor_axis_length',
+ 'eccentricity',
+ 'centroid',
+ 'label',
+ 'mean_intensity',
+ 'max_intensity',
+ 'min_intensity',
+ 'perimeter',
+)
+
+
+def _extract_dataframes(
+ image: np.ndarray, masks: np.ndarray
+) -> list[pd.DataFrame]:
+ """Helper function to extract DataFrames from mask properties."""
+ list_of_df = []
+ for mask in masks:
+ mask = np.where(mask, 1, 0)
+ df = pd.DataFrame(
+ skimage.measure.regionprops_table(
+ mask, intensity_image=image, properties=_PROPERTIES
+ )
+ )
+ list_of_df.append(df)
+ return list_of_df
+
+
+def extract_properties(
+ image: np.ndarray, results: dict[str, np.ndarray], masks: str
+) -> pd.DataFrame:
+ """Extract properties of the mask."""
+ list_of_df = _extract_dataframes(
+ image, results[masks]
+ ) # Use the helper function
+ if not list_of_df: # Handle case where there are no valid masks
+ return pd.DataFrame(columns=_PROPERTIES) # pyrefly: ignore[bad-argument-type]
+
+ features = pd.concat(list_of_df, ignore_index=True)
+ features.rename(
+ columns={
+ 'centroid-0': 'y',
+ 'centroid-1': 'x',
+ 'bbox-0': 'bbox_0',
+ 'bbox-1': 'bbox_1',
+ 'bbox-2': 'bbox_2',
+ 'bbox-3': 'bbox_3',
+ },
+ inplace=True,
+ )
+ return features
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/feature_extraction_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/feature_extraction_test.py
new file mode 100644
index 00000000000..e7b11e7640f
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/feature_extraction_test.py
@@ -0,0 +1,102 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+import numpy as np
+import pandas as pd
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import feature_extraction
+
+TEST_IMAGE = np.array(
+ [
+ [10, 20, 30, 40, 50],
+ [15, 25, 35, 45, 55],
+ [20, 30, 40, 50, 60],
+ [25, 35, 45, 55, 65],
+ [30, 40, 50, 60, 70],
+ ],
+ dtype=np.uint8,
+)
+
+# Create dummy masks (e.g., two masks)
+TEST_MASKS = np.array(
+ [
+ [
+ [0, 0, 0, 0, 0],
+ [0, 1, 1, 0, 0],
+ [0, 1, 1, 0, 0],
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ ],
+ [
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ [0, 0, 1, 1, 0],
+ [0, 0, 1, 1, 0],
+ ],
+ ],
+ dtype=np.int32,
+)
+
+# Create empty masks (all zeros)
+EMPTY_MASKS = np.zeros((2, 5, 5), dtype=np.int32)
+
+# Simulate the results dictionary, assuming masks are under the key 'masks'
+TEST_RESULTS = {'masks': TEST_MASKS}
+EMPTY_RESULTS = {'masks': EMPTY_MASKS}
+
+
+# Expected DataFrame for comparison
+COMPARISON_DATA = {
+ 'area': [4.0, 4.0],
+ 'bbox_0': [1, 3],
+ 'bbox_1': [1, 2],
+ 'bbox_2': [3, 5],
+ 'bbox_3': [3, 4],
+ 'convex_area': [4.0, 4.0],
+ 'bbox_area': [4.0, 4.0],
+ 'major_axis_length': [2.0, 2.0],
+ 'minor_axis_length': [2.0, 2.0],
+ 'eccentricity': [0.0, 0.0],
+ 'y': [1.5, 3.5],
+ 'x': [1.5, 2.5],
+ 'label': [1, 1],
+ 'mean_intensity': [32.5, 52.5],
+ 'max_intensity': [40.0, 60.0],
+ 'min_intensity': [25.0, 45.0],
+ 'perimeter': [4.0, 4.0],
+}
+
+
+class TestExtractProperties(unittest.TestCase):
+
+ def test_extract_properties(self):
+ # Call the function
+ features_df = feature_extraction.extract_properties(
+ TEST_IMAGE, TEST_RESULTS, 'masks'
+ )
+ # Check if the DataFrames are equal
+ self.assertTrue(features_df.equals(pd.DataFrame(COMPARISON_DATA)))
+
+ def test_extract_properties_empty_masks(self):
+ """Test feature extraction with empty masks."""
+ features_df = feature_extraction.extract_properties(
+ TEST_IMAGE, EMPTY_RESULTS, 'masks'
+ )
+ # Expecting an empty DataFrame if there are no valid masks
+ self.assertTrue(features_df.empty)
+
+
+if __name__ == '__main__':
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/ffmpeg_ops.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/ffmpeg_ops.py
new file mode 100644
index 00000000000..1302084de50
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/ffmpeg_ops.py
@@ -0,0 +1,106 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""This script is designed to handle video and image processing tasks.
+
+The script relies heavily on the ffmpeg library for video and image
+processing and ffprobe for metadata extraction.
+
+It focuses on three primary functionalities:
+1) Splitting a video into individual frames, and
+2) Extracting the creation time of the video from its metadata.
+3) Extracting the creation time of an image from its metadata.
+"""
+
+import datetime
+import os
+import ffmpeg
+import PIL
+from PIL import Image
+
+
+def split_video_to_frames(
+ video_name: str, folder_name: str, fps: int = 30
+) -> None:
+ """Split the video into frames using ffmpeg-python.
+
+ Args:
+ video_name: The name/path of the video file.
+ folder_name: The name/path of the folder to store frames.
+ fps: Frames per second to extract from the video.
+ """
+ # Ensure the folder exists
+ if not os.path.exists(folder_name):
+ os.makedirs(folder_name)
+
+ (
+ ffmpeg.input(video_name)
+ .filter('fps', fps=fps)
+ .output(os.path.join(folder_name, 'frame_%06d.png'))
+ .run(capture_stdout=True, capture_stderr=True)
+ )
+
+
+def find_creation_time(video: str) -> str:
+ """Find the creation time of a video file.
+
+ Args:
+ video: A string path to the video file.
+
+ Returns:
+ A string representing the formatted creation time of the video in
+ "YYYY-MM-DD HH:MM:SS" format.
+ """
+ metadata = ffmpeg.probe(video)['streams']
+ timestamp_str = metadata[0]['tags']['creation_time']
+ return datetime.datetime.strptime(
+ timestamp_str, '%Y-%m-%dT%H:%M:%S.%fZ'
+ ).strftime('%Y-%m-%d %H:%M:%S')
+
+
+def get_image_creation_time(image_path):
+ """Retrieves the creation time of an image, trying multiple methods.
+
+ Args:
+ image_path: The path to the image file.
+
+ Returns:
+ A string representing the creation time in the format "%Y-%m-%d %H:%M:%S" if
+ found, otherwise returns "Creation time not found".
+ """
+
+ try:
+ # 1. Try EXIF data (if available)
+ image = Image.open(image_path)
+ exif_data = image.getexif()
+ if exif_data:
+ datetime_tag_id = 36867 # Tag ID for "DateTimeOriginal"
+ datetime_str = exif_data.get(datetime_tag_id)
+ if datetime_str:
+ return datetime.datetime.strptime(
+ datetime_str, '%Y:%m:%d %H:%M:%S'
+ ).strftime('%Y-%m-%d %H:%M:%S')
+
+ # 2. Try file modification time (less accurate, but better than nothing)
+ file_modified_time = os.path.getmtime(image_path)
+ return datetime.datetime.fromtimestamp(
+ file_modified_time
+ ).strftime('%Y-%m-%d %H:%M:%S')
+
+ except FileNotFoundError:
+ return 'Image not found'
+ except PIL.UnidentifiedImageError as e:
+ return f'Error: {e}'
+ except Exception as e: # pylint: disable=broad-exception-caught
+ return f'An unexpected error occurred: {e}'
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/ffmpeg_ops_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/ffmpeg_ops_test.py
new file mode 100644
index 00000000000..121b2f8ed58
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/ffmpeg_ops_test.py
@@ -0,0 +1,79 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import datetime
+import time
+import unittest
+from unittest import mock
+import ffmpeg
+import PIL
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import ffmpeg_ops
+
+
+class TestVideoImageProcessing(unittest.TestCase):
+
+ @mock.patch.object(ffmpeg, "probe", autospec=True)
+ def test_find_creation_time(self, mock_ffprobe):
+ mock_ffprobe.return_value = {
+ "streams": [{"tags": {"creation_time": "2024-02-25T15:30:45.123Z"}}]
+ }
+ expected_time = "2024-02-25 15:30:45"
+ self.assertEqual(ffmpeg_ops.find_creation_time("test.mp4"), expected_time)
+
+ @mock.patch("PIL.Image.open")
+ def test_get_image_creation_time_exif(self, mock_open):
+ mock_image = mock.Mock()
+ mock_exif = mock.Mock()
+ mock_exif.get.return_value = "2024:02:25 15:30:45"
+ mock_image.getexif.return_value = mock_exif
+ mock_open.return_value = mock_image
+
+ result = ffmpeg_ops.get_image_creation_time("dummy.jpg")
+ self.assertEqual(result, "2024-02-25 15:30:45")
+
+ @mock.patch("os.path.getmtime")
+ @mock.patch("PIL.Image.open")
+ def test_get_image_creation_time_no_exif(self, mock_open, mock_getmtime):
+ # Mock image and EXIF returning nothing
+ mock_image = mock.Mock()
+ mock_image.getexif.return_value = {} # No EXIF data
+ mock_open.return_value = mock_image
+
+ # Return fixed timestamp (e.g., 2024-02-25 15:30:45)
+ dt = datetime.datetime(2024, 2, 25, 15, 30, 45)
+ mock_getmtime.return_value = time.mktime(dt.timetuple())
+
+ result = ffmpeg_ops.get_image_creation_time("dummy.jpg")
+ self.assertEqual(result, "2024-02-25 15:30:45")
+
+ @mock.patch.object(PIL.Image, "open", autospec=True)
+ def test_get_image_creation_time_file_not_found(self, mock_image_open):
+ mock_image_open.side_effect = FileNotFoundError
+ self.assertEqual(
+ ffmpeg_ops.get_image_creation_time("missing.jpg"), "Image not found"
+ )
+
+ @mock.patch.object(PIL.Image, "open", autospec=True)
+ def test_get_image_creation_time_unidentified_image(self, mock_image_open):
+ mock_image_open.side_effect = PIL.UnidentifiedImageError(
+ "Cannot identify image"
+ )
+ self.assertEqual(
+ ffmpeg_ops.get_image_creation_time("corrupt.jpg"),
+ "Error: Cannot identify image",
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/inference_pipeline.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/inference_pipeline.py
new file mode 100644
index 00000000000..b4950f13570
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/inference_pipeline.py
@@ -0,0 +1,444 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Pipeline to run the prediction on the images folder with Triton server."""
+
+import os
+import subprocess
+import sys
+
+from absl import app
+from absl import flags
+import big_query_ops # pyrefly: ignore[missing-import]
+import cv2
+import feature_extraction # pyrefly: ignore[missing-import]
+import ffmpeg_ops # pyrefly: ignore[missing-import]
+import mask_bbox_saver # pyrefly: ignore[missing-import]
+import numpy as np
+import object_tracking # pyrefly: ignore[missing-import]
+import object_tracking_postprocessing # pyrefly: ignore[missing-import]
+import pandas as pd
+import triton_server_inference # pyrefly: ignore[missing-import]
+import utils # pyrefly: ignore[missing-import]
+
+sys.path.append(
+ "models/official/projects/waste_identification_ml/model_inference/"
+)
+import color_and_property_extractor # pylint: disable=g-bad-import-order, g-import-not-at-top # pyrefly: ignore[missing-import]
+
+INPUT_DIRECTORY = flags.DEFINE_string(
+ "input_directory", None, "The path to the directory containing images."
+)
+
+OUTPUT_DIRECTORY = flags.DEFINE_string(
+ "output_directory", None, "The path to the directory to save the results."
+)
+
+HEIGHT = flags.DEFINE_integer(
+ "height", None, "Height of an image required by the model"
+)
+
+WIDTH = flags.DEFINE_integer(
+ "width", None, "Width of an image required by the model"
+)
+
+MODEL = flags.DEFINE_string("model", None, "Model name")
+
+PREDICTION_THRESHOLD = flags.DEFINE_float(
+ "score", None, "Threshold to filter the prediction results"
+)
+
+SEARCH_RANGE_X = flags.DEFINE_integer(
+ "search_range_x",
+ None,
+ "Pixels upto which every object needs to be tracked along X axis.",
+)
+
+SEARCH_RANGE_Y = flags.DEFINE_integer(
+ "search_range_y",
+ None,
+ "Pixels upto which every object needs to be tracked along Y axis.",
+)
+
+MEMORY = flags.DEFINE_integer(
+ "memory", None, "Frames upto which every object needs to be tracked."
+)
+
+OVERWRITE = flags.DEFINE_boolean(
+ "overwrite",
+ False,
+ "If True, delete the preexisting BigQuery table before creating a new one.",
+)
+
+PROJECT_ID = flags.DEFINE_string(
+ "project_id", None, "Project ID mentioned in Google Cloud Project"
+)
+
+BQ_DATASET_ID = flags.DEFINE_string(
+ "bq_dataset_id", "Circularnet_dataset", "Big query dataset ID"
+)
+
+BQ_TABLE_ID = flags.DEFINE_string(
+ "bq_table_id", "Circularnet_table", "BigQuery Table ID for features data"
+)
+
+TRACKING_VISUALIZATION = flags.DEFINE_boolean(
+ "tracking_visualization",
+ False,
+ "If True, visualize the tracking results.",
+)
+
+CROPPED_OBJECTS = flags.DEFINE_boolean(
+ "cropped_objects",
+ False,
+ "If True, save cropped objects per category from the output.",
+)
+
+AREA_THRESHOLD = None
+HEIGHT_TRACKING = 300
+WIDTH_TRACKING = 300
+CIRCLE_RADIUS = 7
+FONT = cv2.FONT_HERSHEY_SIMPLEX
+FONTSCALE = 1
+COLOR = (255, 0, 0)
+
+
+def main(_) -> None:
+ # Check if the input and output directories are valid.
+ if (
+ not INPUT_DIRECTORY.value
+ or not OUTPUT_DIRECTORY.value
+ or not INPUT_DIRECTORY.value.startswith("gs://")
+ or not OUTPUT_DIRECTORY.value.startswith("gs://")
+ ):
+ raise ValueError("Bucket path must be non-empty starting with 'gs://'")
+
+ # Copy the images folder from GCP to the present directory.
+ input_directory = (INPUT_DIRECTORY.value).rstrip("/\\")
+ command = f"gcloud storage cp --recursive {input_directory} ."
+ subprocess.run(command, shell=True, check=True)
+
+ # Create a folder to store the predictions.
+ prediction_folder = os.path.basename(input_directory) + "_prediction"
+ os.makedirs(prediction_folder, exist_ok=True)
+
+ # Create a log directory and a logger for logging.
+ log_name = os.path.basename(INPUT_DIRECTORY.value)
+ log_folder = os.path.join(os.getcwd(), "logs")
+ os.makedirs(log_folder, exist_ok=True)
+ logger = utils.create_log_file(log_name, log_folder)
+
+ # Read the labels which the model is trained on.
+ labels_path = os.path.join(os.getcwd(), "labels.csv")
+ labels, category_index = utils.load_labels(labels_path)
+
+ # Read files from a folder.
+ files = utils.files_paths(os.path.basename(input_directory))
+
+ tracking_images = {}
+ features_set = []
+ image_plot = None
+ agg_features = None
+ tracking_features = None
+
+ for frame, image_path in enumerate(files, start=1):
+ # Prepare an input for a Triton model server from an image.
+ logger.info(f"Processing {os.path.basename(image_path)}")
+ try:
+ inputs, original_image, _ = (
+ triton_server_inference.prepare_image(
+ image_path, HEIGHT.value, WIDTH.value
+ )
+ )
+
+ # Extract the creation time of an image.
+ creation_time = ffmpeg_ops.get_image_creation_time(image_path)
+ logger.info(
+ f"Successfully read an image for {os.path.basename(image_path)}"
+ )
+ except (cv2.error, ValueError, TypeError) as e:
+ logger.info("Failed to read an image.")
+ logger.exception("Exception occurred:", e)
+ continue
+
+ try:
+ model_name = MODEL.value
+ result = triton_server_inference.infer(model_name, inputs)
+ logger.info(
+ f"Successfully got prediction for {os.path.basename(image_path)}"
+ )
+ logger.info(f"Total predictions:{result['num_detections'][0]}")
+ except (KeyError, TypeError, RuntimeError, ValueError) as e:
+ logger.info(
+ f"Failed to get prediction for {os.path.basename(image_path)}"
+ )
+ logger.exception("Exception occurred:", e)
+ continue
+
+ try:
+ # Take predictions only above the threshold.
+ if result["num_detections"][0]:
+ scores = result["detection_scores"][0]
+ filtered_indices = scores > PREDICTION_THRESHOLD.value
+
+ if any(filtered_indices):
+ result = utils.filter_detections(result, filtered_indices)
+ logger.info(
+ "Total predictions after"
+ f" thresholding:{result['num_detections'][0]}"
+ )
+ else:
+ logger.info("Zero predictions after threshold.")
+ continue
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info("Failed to filter out predictions.")
+ logger.exception("Exception occured:", e)
+
+ try:
+ # Convert bbox coordinates into normalized coordinates.
+ if result["num_detections"][0]:
+ result["normalized_boxes"] = result["detection_boxes"].copy()
+ result["normalized_boxes"][:, :, [0, 2]] /= HEIGHT.value
+ result["normalized_boxes"][:, :, [1, 3]] /= WIDTH.value
+ result["detection_boxes"] = (
+ result["detection_boxes"].round().astype(int)
+ )
+
+ # Adjust the image size to ensure both dimensions are at least 1024
+ # for saving images with bbox and masks.
+ height_plot, width_plot = utils.adjust_image_size(
+ original_image.shape[0], original_image.shape[1], 1024
+ )
+
+ # Resize the original image to overlay bbox and masks on it.
+ image_plot = cv2.resize(
+ original_image,
+ (width_plot, height_plot),
+ interpolation=cv2.INTER_AREA,
+ )
+
+ # Reframe the masks according to the new image size.
+ result["detection_masks_reframed"] = utils.reframe_masks(
+ result, "normalized_boxes", height_plot, width_plot
+ )
+
+ # Filter the prediction results and remove the overlapping masks.
+ unique_indices = utils.filter_masks(
+ result["detection_masks_reframed"],
+ iou_threshold=0.08,
+ area_threshold=AREA_THRESHOLD,
+ )
+ result = utils.filter_detections(result, unique_indices)
+ logger.info(
+ f"Total predictions after processing: {result['num_detections'][0]}"
+ )
+ else:
+ logger.info("Zero predictions after processing.")
+ continue
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info("Issue in post processing predictions results.")
+ logger.exception("Exception occured:", e)
+
+ try:
+ if result["num_detections"][0]:
+ result["detection_classes_names"] = np.array(
+ [[str(labels[i - 1]) for i in result["detection_classes"][0]]]
+ )
+
+ # Save the prediction results as an image file with bbx and masks.
+ mask_bbox_saver.save_bbox_masks_labels(
+ result=result,
+ image=image_plot,
+ file_name=os.path.basename(image_path),
+ folder=prediction_folder,
+ category_index=category_index,
+ threshold=PREDICTION_THRESHOLD.value,
+ )
+ logger.info("Visualization saved.")
+ except (KeyError, IndexError, TypeError, ValueError) as e:
+ logger.info("Issue in saving visualization of results.")
+ logger.exception("Exception occured:", e)
+
+ try:
+ # Resize an image for object tracking..
+ tracking_image = cv2.resize(
+ original_image,
+ (WIDTH_TRACKING, HEIGHT_TRACKING),
+ interpolation=cv2.INTER_AREA,
+ )
+ tracking_images[os.path.basename(image_path)] = tracking_image
+
+ # Reducing mask sizes in order to keep the memory required for object
+ # tracking under a threshold.
+ result["detection_masks_tracking"] = np.array([
+ cv2.resize(
+ i,
+ (WIDTH_TRACKING, HEIGHT_TRACKING),
+ interpolation=cv2.INTER_NEAREST,
+ )
+ for i in result["detection_masks_reframed"]
+ ])
+
+ # Crop objects from an image using masks for color detection.
+ cropped_objects = [
+ np.where(np.expand_dims(i, -1), image_plot, 0) # pyrefly: ignore[no-matching-overload]
+ for i in result["detection_masks_reframed"]
+ ]
+
+ # Perform color detection using clustering approach.
+ dominant_colors = [
+ *map(
+ color_and_property_extractor.find_dominant_color, cropped_objects
+ )
+ ]
+ generic_color_names = color_and_property_extractor.get_generic_color_name(
+ dominant_colors
+ )
+
+ # Extract features.
+ features = feature_extraction.extract_properties(
+ tracking_image, result, "detection_masks_tracking"
+ )
+ features["source_name"] = os.path.basename(os.path.dirname(image_path))
+ features["image_name"] = os.path.basename(image_path)
+ features["creation_time"] = creation_time
+ features["frame"] = frame
+ features["detection_scores"] = result["detection_scores"][0]
+ features["detection_classes"] = result["detection_classes"][0]
+ features["detection_classes_names"] = result["detection_classes_names"][0]
+ features["color"] = generic_color_names
+ features_set.append(features)
+ logger.info("Features extracted.\n")
+ except (KeyError, IndexError, TypeError, ValueError):
+ logger.info("Failed to extract properties.")
+
+ try:
+ if features_set:
+ features_df = pd.concat(features_set, ignore_index=True)
+
+ # Apply object tracking to the features.
+ tracking_features = object_tracking.apply_tracking(
+ features_df,
+ search_range_x=SEARCH_RANGE_X.value,
+ search_range_y=SEARCH_RANGE_Y.value,
+ memory=MEMORY.value,
+ )
+
+ # Process the tracking results to remove errors.
+ agg_features = object_tracking_postprocessing.process_tracking_result(
+ tracking_features
+ )
+ counts = agg_features.groupby("detection_classes_names").size()
+ counts.to_frame().to_csv(os.path.join(os.getcwd(), "count.csv"))
+ logger.info("Object tracking applied.")
+ except (KeyError, IndexError, TypeError, ValueError):
+ logger.info("Failed to apply object tracking.")
+
+ try:
+ if TRACKING_VISUALIZATION.value:
+ # Create a folder to save the tracking visualization.
+ tracking_folder = os.path.basename(input_directory) + "_tracking"
+ os.makedirs(tracking_folder, exist_ok=True)
+
+ # Save the tracking results as an image files.
+ output_folder = mask_bbox_saver.visualize_tracking_results(
+ tracking_features=tracking_features,
+ tracking_images=tracking_images,
+ tracking_folder=tracking_folder,
+ )
+ logger.info(f"Tracking visualization saved to {output_folder}.")
+
+ # Move the tracking visualization to the output directory.
+ commands = [
+ f"gcloud storage cp --recursive {output_folder} {OUTPUT_DIRECTORY.value}",
+ f"rm -r {output_folder}",
+ ]
+ combined_command_1 = " && ".join(commands)
+ subprocess.run(combined_command_1, shell=True, check=True)
+ logger.info("Tracking visualization saved.")
+ except (KeyError, IndexError, TypeError, ValueError):
+ logger.info("Failed to visualize tracking results.")
+
+ try:
+ if CROPPED_OBJECTS.value:
+ cropped_obj_folder = mask_bbox_saver.save_cropped_objects(
+ agg_features=agg_features,
+ input_directory=input_directory,
+ height_tracking=HEIGHT_TRACKING,
+ width_tracking=WIDTH_TRACKING,
+ resize_bbox=utils.resize_bbox,
+ )
+ logger.info("Cropped objects saved in %s", cropped_obj_folder)
+
+ # Move the cropped objects to the output directory.
+ commands = [
+ f"gcloud storage cp --recursive {cropped_obj_folder} {OUTPUT_DIRECTORY.value}",
+ f"rm -r {cropped_obj_folder}",
+ ]
+
+ combined_command_2 = " && ".join(commands)
+ subprocess.run(combined_command_2, shell=True, check=True)
+ logger.info("Cropped objects saved.")
+ except (KeyError, IndexError, TypeError, ValueError):
+ logger.info("Issue in cropping objects")
+ logger.info("Failed to crop objects.")
+ logger.exception("Exception occured:", e) # pyrefly: ignore[unbound-name]
+
+ if isinstance(agg_features, pd.DataFrame) and not agg_features.empty:
+ try:
+ # Create a big query table to store the aggregated features data.
+ big_query_ops.create_table(
+ PROJECT_ID.value,
+ BQ_DATASET_ID.value,
+ BQ_TABLE_ID.value,
+ overwrite=OVERWRITE.value,
+ )
+ logger.info("Successfully created table.")
+ except (KeyError, IndexError, TypeError, ValueError):
+ logger.info("Issue in creation of table")
+ return
+
+ try:
+ # Ingest the aggregated features data into the big query table.
+ big_query_ops.ingest_data(
+ agg_features, PROJECT_ID.value, BQ_DATASET_ID.value, BQ_TABLE_ID.value
+ )
+ logger.info("Data ingested successfully.")
+ except (KeyError, IndexError, TypeError, ValueError):
+ logger.info("Issue in data ingestion.")
+ return
+
+ try:
+ # Move the folders to the destination bucket.
+ commands = [
+ (
+ "gcloud storage cp --recursive"
+ f" {os.path.basename(input_directory)} {OUTPUT_DIRECTORY.value}"
+ ),
+ f"rm -r {os.path.basename(input_directory)}",
+ f"gcloud storage cp --recursive {prediction_folder} {OUTPUT_DIRECTORY.value}",
+ f"rm -r {prediction_folder}",
+ ]
+
+ combined_command_3 = " && ".join(commands)
+ subprocess.run(combined_command_3, shell=True, check=True)
+ logger.info("Successfully moved to destination bucket")
+ except (KeyError, IndexError, TypeError, ValueError):
+ logger.info("Issue in moving folders to destination bucket")
+ else:
+ logger.info("No features to ingest.")
+ utils.shutdown_system()
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/labels.csv b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/labels.csv
new file mode 100644
index 00000000000..85f749e670d
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/labels.csv
@@ -0,0 +1,45 @@
+Aluminium_Can
+Aluminium_Foil
+Battery
+Brush
+Bulb
+Fiber_Cardboard
+Fiber_Cup-&-glass
+Fiber_Paper
+Footwear
+Glass_Bottle
+Lighter
+Metals
+Metals_Bottle
+Metals_Container
+Metals_Lid
+Plastics-ABS_Electronics
+Plastics-HDPE_Bottle
+Plastics-HDPE_Container
+Plastics-HDPE_Lid
+Plastics-HDPE_Toys
+Plastics-HDPE_Tube
+Plastics-LDPE_Flexibles
+Plastics-MLP_Flexibles
+Plastics-PC_CD
+Plastics-PC_Goggles
+Plastics-PET_Blister-pack
+Plastics-PET_Bottle
+Plastics-PET_Container
+Plastics-PET_Cup-&-glass
+Plastics-PP_Comb
+Plastics-PP_Container
+Plastics-PP_Pen
+Plastics-PP_Spoon
+Plastics-PP_Straw
+Plastics-PP_Tray
+Plastics-PS
+Plastics-PS_Cup-&-glass
+Plastics-PS_Flexibles
+Plastics-PS_Hangers
+Plastics-PVC_Pipe
+Plastics-Tetrapak_Carton
+Textile_Clothes
+Textile_Flexibles
+Tire
+Wood
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/mask_bbox_saver.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/mask_bbox_saver.py
new file mode 100644
index 00000000000..eb350f2c0c7
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/mask_bbox_saver.py
@@ -0,0 +1,224 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Processing and saving images with annotations from Mask R-CNN outputs.
+
+It handles the extraction of bounding boxes and segmentation masks from the
+output and saves these annotations in a visually interpretable format.
+Additionally, it saves binary masks for each detected object.
+
+The script performs two main functions:
+1) It overlays bounding boxes and segmentation masks on the original image and
+saves this annotated image.
+2) It extracts and saves binary masks for each detected object in the image.
+"""
+
+from collections.abc import Mapping
+import dataclasses
+import os
+from typing import Any, Callable, Dict
+
+import cv2
+import numpy as np
+import pandas as pd
+
+from official.vision.utils.object_detection import visualization_utils as viz_utils
+
+CIRCLE_RADIUS = 7
+CIRCLE_COLOR = (255, 133, 233)
+TEXT_FONT = cv2.FONT_HERSHEY_SIMPLEX
+TEXT_SCALE = 1.0
+TEXT_COLOR = (255, 0, 0)
+TEXT_THICKNESS = 2
+TEXT_LINE_TYPE = cv2.LINE_AA
+
+
+@dataclasses.dataclass
+class BoundingBox:
+ y1: int | float
+ x1: int | float
+ y2: int | float
+ x2: int | float
+
+
+@dataclasses.dataclass
+class ImageSize:
+ height: int
+ width: int
+
+
+def save_bbox_masks_labels(
+ *,
+ result: Mapping[Any, np.ndarray],
+ image: np.ndarray,
+ file_name: str,
+ folder: str,
+ category_index: Dict[int, Dict[str, str]],
+ threshold: float,
+) -> None:
+ """Saves an image with visualized bounding boxes, labels, and masks.
+
+ This function takes the output from Mask R-CNN, copies the original image,
+ and applies visualizations of detection boxes, classes, and scores.
+ If available, it also applies segmentation masks. The result is an image that
+ juxtaposes the original with the annotated version, saved to the specified
+ folder.
+
+ Args:
+ result: The output from theMask RCNN model, expected to contain detection
+ boxes, classes, scores, reframed detection masks, etc.
+ image: The original image as a numpy array.
+ file_name: The filename for saving the output image.
+ folder: The folder path where the output image will be saved.
+ category_index: A dictionary mapping class IDs to class labels.
+ threshold: Value between 0 and 1 to filter out the prediction results.
+ """
+ image_new = image.copy()
+ viz_utils.visualize_boxes_and_labels_on_image_array(
+ image_new,
+ result['normalized_boxes'][0],
+ (result['detection_classes'][0] + 0).astype(int),
+ result['detection_scores'][0],
+ category_index=category_index,
+ use_normalized_coordinates=True,
+ max_boxes_to_draw=100,
+ min_score_thresh=threshold,
+ agnostic_mode=False,
+ instance_masks=result.get('detection_masks_reframed', None),
+ line_thickness=4,
+ )
+
+ cv2.imwrite(
+ os.path.join(folder, file_name),
+ np.concatenate((image, image_new), axis=1),
+ )
+
+
+def save_binary_masks(
+ result: Dict[Any, np.ndarray], file_name: str, folder: str
+) -> None:
+ """Saves binary masks generated from object detection results.
+
+ This function processes the binary mask data extracted from the results of
+ Mask RCNN model and saves the combined binary masks as an image file. It
+ creates a mask image by layering each individual mask found in the result and
+ saves the final mask image to the specified location.
+
+ Args:
+ result: The output from the Mask RCNN model, expected to contain key
+ 'detection_masks_reframed'.
+ file_name: The filename for saving the output mask image.
+ folder: The folder path where the output mask image will be saved.
+ """
+ mask = np.zeros_like(result['detection_masks_reframed'][0], dtype=np.uint8)
+
+ # Process and accumulate binary masks
+ for single_mask in result['detection_masks_reframed']:
+ mask += single_mask.astype(np.uint8) * 255 # Convert to 0-255 range
+
+ cv2.imwrite(os.path.join(folder, file_name), mask)
+
+
+# TODO(umairsabir): Add helper function to remove nested loop.
+def visualize_tracking_results(
+ tracking_features: pd.DataFrame,
+ tracking_images: Mapping[str, np.ndarray],
+ tracking_folder: str,
+) -> str:
+ """Draws tracking results on images and saves them to an output folder.
+
+ Args:
+ tracking_features: DataFrame with columns ['image_name', 'x', 'y',
+ 'particle'].
+ tracking_images: Mapping from image_name to image (numpy array).
+ tracking_folder: Directory where tracking results are saved.
+
+ Returns:
+ str:Path to the output folder where annotated images are saved.
+ """
+ groups = tracking_features.groupby('image_name')
+ for name, group in groups:
+ img = tracking_images[name].copy() # pyrefly: ignore[bad-index]
+ for k in range(len(group)):
+ x, y = int(group.iloc[k]['x']), int(group.iloc[k]['y'])
+ cv2.circle(img, (x, y), CIRCLE_RADIUS, CIRCLE_COLOR, -1)
+ cv2.putText(
+ img,
+ str(int(group.iloc[k]['particle'])),
+ (x, y),
+ TEXT_FONT,
+ TEXT_SCALE,
+ TEXT_COLOR,
+ TEXT_THICKNESS,
+ TEXT_LINE_TYPE,
+ )
+ cv2.imwrite(os.path.join(tracking_folder, str(name)), img)
+ return tracking_folder
+
+
+def save_cropped_objects(
+ agg_features: pd.DataFrame,
+ input_directory: str,
+ height_tracking: int,
+ width_tracking: int,
+ resize_bbox: Callable[
+ [BoundingBox, ImageSize, ImageSize], tuple[int, int, int, int]
+ ],
+ output_suffix: str = '_cropped_objects',
+) -> str:
+ """Saves cropped object images by category from tracking results.
+
+ Args:
+ agg_features: DataFrame containing grouped tracking results.
+ input_directory: The directory with original images.
+ height_tracking: Height used during tracking.
+ width_tracking: Width used during tracking.
+ resize_bbox: Function to resize bounding box.
+ output_suffix: Suffix for cropped objects folder.
+
+ Returns:
+ str: Path to the output folder where cropped images are saved.
+ """
+ cropped_obj_folder = os.path.basename(input_directory) + output_suffix
+ os.makedirs(cropped_obj_folder, exist_ok=True)
+
+ if agg_features.empty:
+ return cropped_obj_folder
+
+ for group_name, df in agg_features.groupby('detection_classes_names'):
+ class_folder = os.path.join(cropped_obj_folder, str(group_name))
+ os.makedirs(class_folder, exist_ok=True)
+
+ for row in df.itertuples(index=False):
+ image = cv2.imread(
+ os.path.join(os.path.basename(input_directory), row.image_name)
+ )
+ image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
+ new_h, new_w = image_rgb.shape[0], image_rgb.shape[1]
+
+ y1, x1, y2, x2 = row.bbox_0, row.bbox_1, row.bbox_2, row.bbox_3
+ bbox = BoundingBox(y1, x1, y2, x2)
+ new_bbox = resize_bbox(
+ bbox,
+ ImageSize(height=height_tracking, width=width_tracking),
+ ImageSize(height=new_h, width=new_w),
+ )
+
+ score = getattr(row, 'detection_scores', 0.0)
+ name = f'{os.path.splitext(row.image_name)[0]}_{row.particle}_{score:.2f}.png'
+ crop = image_rgb[new_bbox[0] : new_bbox[2], new_bbox[1] : new_bbox[3]]
+
+ cv2.imwrite(os.path.join(class_folder, name), crop)
+
+ return cropped_obj_folder
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/mask_bbox_saver_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/mask_bbox_saver_test.py
new file mode 100644
index 00000000000..cb8d58f83e8
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/mask_bbox_saver_test.py
@@ -0,0 +1,163 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import os
+import unittest
+from unittest import mock
+
+import numpy as np
+import pandas as pd
+
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import mask_bbox_saver
+
+
+class VisualizeTrackingResultsTest(unittest.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ self.tracking_folder = "mock_tracking_output"
+ self.image_name = "img1.png"
+ self.mock_input_dir = "/mock/path/images"
+ self.output_suffix = "_cropped_objects"
+ self.mock_image = np.ones((200, 200, 3), dtype=np.uint8) * 255
+
+ self.agg_features = pd.DataFrame({
+ "detection_classes_names": ["Bottle"],
+ "image_name": ["img1.png"],
+ "bbox_0": [0.1],
+ "bbox_1": [0.1],
+ "bbox_2": [0.5],
+ "bbox_3": [0.5],
+ "particle": [1],
+ "detection_scores": [0.9],
+ })
+
+ self.tracking_df = pd.DataFrame({
+ "image_name": [self.image_name, self.image_name],
+ "x": [15, 45],
+ "y": [25, 65],
+ "particle": [1, 2],
+ })
+
+ self.tracking_images = {self.image_name: self.mock_image}
+
+ def test_dataframe_validity(self):
+ self.assertFalse(
+ self.tracking_df.isnull().values.any(),
+ msg="DataFrame contains NaN values",
+ )
+
+ expected_cols = {"image_name", "x", "y", "particle"}
+ actual_cols = set(self.tracking_df.columns)
+ self.assertEqual(
+ expected_cols, actual_cols, msg="DataFrame missing required columns"
+ )
+
+ @mock.patch("cv2.imwrite")
+ def test_visualize_tracking_function_runs(self, mock_imwrite):
+ result_folder = mask_bbox_saver.visualize_tracking_results(
+ tracking_features=self.tracking_df,
+ tracking_images=self.tracking_images,
+ tracking_folder=self.tracking_folder,
+ )
+
+ self.assertEqual(result_folder, self.tracking_folder)
+ mock_imwrite.assert_called() # At least one image was attempted to be saved
+
+ # Optional: Check that the correct file path was constructed
+ expected_output_path = os.path.join(self.tracking_folder, self.image_name)
+ mock_imwrite.assert_any_call(expected_output_path, mock.ANY)
+
+ def test_dataframe_is_valid(self):
+ self.assertFalse(self.agg_features.empty)
+ required_columns = {
+ "detection_classes_names",
+ "image_name",
+ "bbox_0",
+ "bbox_1",
+ "bbox_2",
+ "bbox_3",
+ "particle",
+ "detection_scores",
+ }
+ self.assertTrue(required_columns.issubset(set(self.agg_features.columns)))
+
+ @mock.patch("cv2.imwrite")
+ @mock.patch(
+ "cv2.cvtColor", return_value=np.ones((200, 200, 3), dtype=np.uint8)
+ )
+ @mock.patch("cv2.imread", return_value=np.ones((200, 200, 3), dtype=np.uint8))
+ @mock.patch("os.makedirs")
+ def test_save_cropped_objects_success(
+ self, mock_makedirs, mock_imread, mock_cvtcolor, mock_imwrite # pylint: disable=unused-argument
+ ):
+
+ agg_features = pd.DataFrame({
+ "detection_classes_names": ["Bottle"],
+ "image_name": ["img1.png"],
+ "bbox_0": [0.1],
+ "bbox_1": [0.1],
+ "bbox_2": [0.5],
+ "bbox_3": [0.5],
+ "particle": [1],
+ "detection_scores": [0.9],
+ })
+
+ input_directory = "/mock/path/images"
+ output_suffix = "_cropped_objects"
+ mock_bbox = (10, 10, 100, 100)
+
+ def mock_resize_bbox(bbox, old_size, new_size): # pylint: disable=unused-argument
+ return mock_bbox
+
+ output_path = mask_bbox_saver.save_cropped_objects(
+ agg_features=agg_features,
+ input_directory=input_directory,
+ height_tracking=100,
+ width_tracking=100,
+ resize_bbox=mock_resize_bbox,
+ output_suffix=output_suffix,
+ )
+
+ expected_output = os.path.basename(input_directory) + output_suffix
+ self.assertEqual(output_path, expected_output)
+ mock_makedirs.assert_any_call(expected_output, exist_ok=True)
+ mock_imread.assert_called_once()
+ mock_imwrite.assert_called_once()
+
+ @mock.patch("os.makedirs")
+ def test_empty_dataframe_returns_early(self, mock_makedirs):
+ agg_features = pd.DataFrame() # Empty DataFrame
+ input_directory = "/mock/path/images"
+ output_suffix = "_cropped_objects"
+
+ def mock_resize_bbox(bbox, old_size, new_size): # pylint: disable=unused-argument
+ return (10, 10, 100, 100)
+
+ output_path = mask_bbox_saver.save_cropped_objects(
+ agg_features=agg_features,
+ input_directory=input_directory,
+ height_tracking=100,
+ width_tracking=100,
+ resize_bbox=mock_resize_bbox,
+ output_suffix=output_suffix,
+ )
+
+ expected_output = os.path.basename(input_directory) + output_suffix
+ self.assertEqual(output_path, expected_output)
+ mock_makedirs.assert_called_once_with(expected_output, exist_ok=True)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking.py
new file mode 100644
index 00000000000..bd6af252926
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking.py
@@ -0,0 +1,72 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Object tracking using trackpy."""
+
+import pandas as pd
+import trackpy as tp
+
+
+def apply_tracking(
+ df: pd.DataFrame, search_range_x: int, search_range_y: int, memory: int
+) -> pd.DataFrame:
+ """Apply tracking to the dataframe.
+
+ Args:
+ df: The dataframe to apply tracking to.
+ search_range_x: The search range of pixels for tracking along x axis.
+ search_range_y: The search range of pixels for tracking along y axis.
+ memory: The number of frames that an object can skip detection in and still
+ be tracked.
+
+ Returns:
+ The tracking result dataframe.
+ """
+ # Define the columns to examine for tracking
+ tracking_columns = [
+ 'x',
+ 'y',
+ 'frame',
+ 'bbox_0',
+ 'bbox_1',
+ 'bbox_2',
+ 'bbox_3',
+ 'major_axis_length',
+ 'minor_axis_length',
+ 'perimeter',
+ ]
+
+ # Perform the tracking using the relevant columns
+ track_df = tp.link_df(
+ df[tracking_columns],
+ search_range=(search_range_y, search_range_x),
+ memory=memory,
+ )
+
+ # Preserve original columns not used directly in tracking.
+ additional_columns = [
+ 'source_name',
+ 'image_name',
+ 'detection_scores',
+ 'detection_classes_names',
+ 'detection_classes',
+ 'color',
+ 'creation_time',
+ ]
+ track_df[additional_columns] = df[additional_columns]
+
+ # Remove unnecessary columns from the tracking result and reset index.
+ track_df.drop(columns=['frame'], inplace=True)
+
+ return track_df
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_postprocessing.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_postprocessing.py
new file mode 100644
index 00000000000..245d2ba19a8
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_postprocessing.py
@@ -0,0 +1,85 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Object tracking functions."""
+
+import pandas as pd
+
+
+def process_tracking_result(df: pd.DataFrame) -> pd.DataFrame:
+ """Process the tracking result dataframe.
+
+ Args:
+ df: Dataframe to be aggregated.
+
+ Returns:
+ Processed dataframe.
+ """
+ # Apply special class selection logic to each particle
+ class_info = df.groupby('particle', as_index=False).apply(
+ _select_class_with_scores, include_groups=False
+ )
+
+ grouped_particles = (
+ df.groupby('particle')
+ .agg({
+ 'source_name': 'first',
+ 'image_name': 'first',
+ 'detection_scores': 'max',
+ 'creation_time': 'first',
+ 'bbox_0': 'first',
+ 'bbox_1': 'first',
+ 'bbox_2': 'first',
+ 'bbox_3': 'first',
+ })
+ .reset_index()
+ )
+
+ # Add class information
+ grouped_particles['detection_classes'] = class_info['class_id']
+ grouped_particles['detection_classes_names'] = class_info['class_name']
+
+ return grouped_particles
+
+
+def _select_class_with_scores(group: pd.DataFrame) -> pd.Series:
+ """Selects a class based on the most frequently occurring class (modal).
+
+ If there's a tie, selects the class with the highest detection score.
+
+ Args:
+ group: It contains 'detection_classes', 'detection_scores', and
+ 'detection_classes_names'.
+
+ Returns:
+ A Series with 'class_id' and 'class_name' of the selected class.
+ """
+ # Get the value counts of classes
+ class_counts = group['detection_classes'].value_counts()
+
+ # For ties, also look at the highest score amonst tied classes
+ tied_classes = class_counts[class_counts == class_counts.iloc[0]].index
+ max_scores_by_class = {
+ cls: group[group['detection_classes'] == cls]['detection_scores'].max()
+ for cls in tied_classes
+ }
+
+ class_id = max(max_scores_by_class.items(), key=lambda x: x[1])[0]
+
+ # Get corresponding class name
+ class_name = group[group['detection_classes'] == class_id][
+ 'detection_classes_names'
+ ].iloc[0]
+
+ return pd.Series({'class_id': class_id, 'class_name': class_name})
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_postprocessing_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_postprocessing_test.py
new file mode 100644
index 00000000000..06049216012
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_postprocessing_test.py
@@ -0,0 +1,131 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+import pandas as pd
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import object_tracking_postprocessing
+
+TEST_DATA = pd.DataFrame({
+ "particle": [1, 1, 2, 2, 3, 3, 3, 4],
+ "source_name": [
+ "src1",
+ "src1",
+ "src2",
+ "src2",
+ "src3",
+ "src3",
+ "src3",
+ "src4",
+ ],
+ "image_name": [
+ "img1",
+ "img1",
+ "img2",
+ "img2",
+ "img3",
+ "img3",
+ "img3",
+ "img4",
+ ],
+ "detection_scores": [0.8, 0.9, 0.5, 0.7, 0.6, 0.7, 0.8, 0.9],
+ "detection_classes": ["A", "A", "B", "C", "D", "D", "E", "A"],
+ "detection_classes_names": [
+ "ClassA",
+ "ClassA",
+ "ClassB",
+ "ClassC",
+ "ClassD",
+ "ClassD",
+ "ClassE",
+ "ClassA",
+ ],
+ "color": [
+ "red",
+ "red",
+ "blue",
+ "blue",
+ "green",
+ "green",
+ "green",
+ "orange",
+ ],
+ "creation_time": [
+ "2024-01-01",
+ "2024-01-01",
+ "2024-01-02",
+ "2024-01-02",
+ "2024-01-03",
+ "2024-01-03",
+ "2024-01-03",
+ "2024-01-04",
+ ],
+ "bbox_0": [10, 10, 20, 20, 30, 30, 30, 40],
+ "bbox_1": [15, 15, 25, 25, 35, 35, 35, 45],
+ "bbox_2": [50, 50, 60, 60, 70, 70, 70, 80],
+ "bbox_3": [55, 55, 65, 65, 75, 75, 75, 85],
+})
+
+
+class TestProcessTrackingResult(unittest.TestCase):
+
+ def test_particle_grouping(self):
+ """Test if the function correctly aggregates the tracking data."""
+ result = object_tracking_postprocessing.process_tracking_result(
+ TEST_DATA.copy()
+ )
+
+ self.assertEqual(len(result), 4)
+
+ def test_single_class_selection(self):
+ """Test if the function correctly selects a class when there's only one."""
+ result = object_tracking_postprocessing.process_tracking_result(
+ TEST_DATA.copy()
+ )
+
+ self.assertEqual(
+ result[result["particle"] == 4]["detection_classes"].values, "A"
+ )
+
+ def test_modal_class_selection(self):
+ """Test if the function correctly selects the most common class."""
+ result = object_tracking_postprocessing.process_tracking_result(
+ TEST_DATA.copy()
+ )
+
+ self.assertEqual(
+ result[result["particle"] == 3]["detection_classes"].values, "D"
+ )
+
+ def test_single_class_selection_with_tie(self):
+ """Test if the function correctly selects a class when there's a modal tie."""
+ result = object_tracking_postprocessing.process_tracking_result(
+ TEST_DATA.copy()
+ )
+
+ self.assertEqual(
+ result[result["particle"] == 1]["detection_classes"].values, "A"
+ )
+
+ def test_tie_multiple_class_selection(self):
+ """Test class selection using scores for modal tie with multiple classes."""
+ result = object_tracking_postprocessing.process_tracking_result(
+ TEST_DATA.copy()
+ )
+
+ self.assertEqual(
+ result[result["particle"] == 2]["detection_classes"].values, "C"
+ )
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_test.py
new file mode 100644
index 00000000000..0a9dd2e2d4d
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/object_tracking_test.py
@@ -0,0 +1,71 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+
+import pandas as pd
+
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import object_tracking
+
+
+TEST_IMAGES = pd.DataFrame({
+ 'x': [1, 2, 3],
+ 'y': [4, 5, 6],
+ 'frame': [0, 0, 1],
+ 'bbox_0': [1, 2, 3],
+ 'bbox_1': [4, 5, 6],
+ 'bbox_2': [7, 8, 9],
+ 'bbox_3': [10, 11, 12],
+ 'major_axis_length': [13, 14, 15],
+ 'minor_axis_length': [16, 17, 18],
+ 'perimeter': [19, 20, 21],
+ 'source_name': ['source_name_1', 'source_name_2', 'source_name_3'],
+ 'image_name': ['image_name_1', 'image_name_2', 'image_name_3'],
+ 'detection_scores': [0.1, 0.2, 0.3],
+ 'detection_classes_names': ['class_name_1', 'class_name_2', 'class_name_3'],
+ 'detection_classes': [1, 2, 3],
+ 'color': ['red', 'blue', 'green'],
+ 'creation_time': [100, 200, 300],
+})
+
+
+class ObjectTrackingTest(unittest.TestCase):
+
+ def test_object_tracking_retains_columns(self):
+ """Tests that object tracking correctly retains columns not used in tracking."""
+ df = TEST_IMAGES.copy()
+ expected_columns = [
+ 'source_name',
+ 'image_name',
+ 'detection_scores',
+ 'detection_classes_names',
+ 'detection_classes',
+ 'color',
+ 'creation_time',
+ ]
+
+ tracking_result = object_tracking.apply_tracking(df, 10, 10, 10)
+
+ self.assertTrue(all(key in tracking_result for key in expected_columns))
+
+ def test_object_tracking_drops_columns(self):
+ """Tests that object tracking correctly drops unneeded columns."""
+ df = TEST_IMAGES.copy()
+
+ tracking_result = object_tracking.apply_tracking(df, 10, 10, 10)
+
+ self.assertNotIn('frame', tracking_result.columns)
+
+if __name__ == '__main__':
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/requirement.sh b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/requirement.sh
new file mode 100644
index 00000000000..c010da631d1
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/requirement.sh
@@ -0,0 +1,64 @@
+#!/bin/bash
+
+# Summary
+cat << EOF
+This script sets up the environment for running ML models by ensuring Bash
+execution, installing system dependencies, setting up a virtual environment,
+installing ML packages, and cloning TensorFlow Model Garden.
+EOF
+
+# Ensure the script is executed with /bin/bash
+if [ -z "$BASH_VERSION" ]; then
+ exec /bin/bash "$0" "$@"
+fi
+
+# Update the package lists for upgrades and new package installations.
+sudo apt-get update -y
+
+# Check if Docker is installed
+if ! command -v docker &> /dev/null
+then
+ echo "Docker is not installed. Installing Docker..."
+ # Download the Docker installation script from Docker's official site.
+ curl -fsSL https://get.docker.com -o get-docker.sh
+
+ # Run the Docker installation script.
+ sudo sh get-docker.sh
+
+ # Clean up Docker installation script
+ rm -f get-docker.sh
+
+ echo "Docker installation completed."
+else
+ echo "Docker is already installed. Skipping Docker installation."
+fi
+
+# Install python3-venv for creating virtual environments.
+sudo apt-get install -y python3-venv python3-pip ffmpeg
+
+# Create a virtual environment and activate it.
+python3.10 -m venv myenv
+source myenv/bin/activate
+
+# Install Python packages inside the virtual environment:
+pip install --no-cache-dir natsort absl-py opencv-python pandas pandas-gbq \
+ google-cloud-bigquery google-auth trackpy google-cloud-storage tensorflow \
+ scikit-image scikit-learn webcolors==1.13 ffmpeg-python tf_keras \
+ tf_slim tritonclient[all]
+
+# Install the python package for the visualization.
+pip install tf-models-official
+
+# Check if the 'models' directory exists before cloning.
+if [ ! -d "models" ]; then
+ # Cloning project directory from TF Model Garden for postprocessing
+ # and preprocessing functions.
+ git clone --depth 1 https://github.com/tensorflow/models.git
+else
+ echo "'models' directory already exists. Skipping cloning."
+fi
+
+# Deactivate the virtual environment
+deactivate
+
+echo "Environment setup is complete."
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/run_images.sh b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/run_images.sh
new file mode 100644
index 00000000000..cb2e827faa9
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/run_images.sh
@@ -0,0 +1,67 @@
+#!/bin/bash
+
+cat << EOF
+This script automates the execution of an Circularnet pipeline for image
+processing.
+
+Steps Performed:
+ 1. Activates the Python virtual environment named 'myenv'.
+ 2. Validates successful activation of the virtual environment.
+ 3. Executes the 'pipeline_images.py' script with the following parameters:
+
+ Parameters:
+ --input_directory : GCS directory where the input images are stored for
+ inference.
+ --output_directory : GCS directory where the model inference outputs will be
+ saved.
+ --height : Height to which input images are resized for the Mask
+ R-CNN model.
+ --width : Width to which input images are resized for the Mask
+ R-CNN model.
+ --model : Name of the model to download and use for inference.
+ --score : Confidence threshold for detections during
+ inference.
+ --search_range_x : Max pixel movement allowed in the X direction for
+ object tracking between missed frames.
+ --search_range_y : Max pixel movement allowed in the Y direction for
+ object tracking between missed frames.
+ --memory : Number of frames an object can be missed and still
+ be tracked.
+ --project_id : Google Cloud Project ID for BigQuery operations.
+ --bq_dataset_id : BigQuery Dataset ID where results will be stored.
+ --bq_table_id : BigQuery Table ID where results will be stored.
+ --overwrite : If set to True, overwrites the pre-existing
+ BigQuery table.
+ --tracking_visualization : If set to True, visualizes the tracking results
+ from the tracking algorithm.
+ --cropped_objects : If set to True, crops the objects per category
+ according to the prediction and tracking results.
+EOF
+
+# Activate the virtual environment
+source myenv/bin/activate
+
+# Check if the virtual environment is activated
+if [[ "$VIRTUAL_ENV" != "" ]]; then
+ echo "Virtual environment 'myenv' activated successfully."
+else
+ echo "Failed to activate virtual environment. Exiting."
+ exit 1
+fi
+
+python inference_pipeline.py \
+ --input_directory=gs://recykal/TestData/Delterra \
+ --output_directory=gs://recykal/TestData/output \
+ --height=1024 \
+ --width=1024 \
+ --model=Jan2025_ver2_merged_1024_1024 \
+ --score=0.70 \
+ --search_range_x=150 \
+ --search_range_y=20 \
+ --memory=10 \
+ --project_id=waste-identification-ml-330916 \
+ --bq_dataset_id=circularnet_dataset \
+ --bq_table_id=circularnet_table \
+ --overwrite=True \
+ --tracking_visualization=False \
+ --cropped_objects=False
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/triton_server_inference.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/triton_server_inference.py
new file mode 100644
index 00000000000..a1c75ff3ee3
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/triton_server_inference.py
@@ -0,0 +1,79 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Prediction from the Triton server."""
+
+from typing import Any
+import cv2
+import numpy as np
+from tritonclient import grpc as triton_grpc
+
+
+_OUTPUT_KEYS = (
+ 'detection_classes',
+ 'detection_masks',
+ 'detection_boxes',
+ 'image_info',
+ 'num_detections',
+ 'detection_scores',
+)
+_OUTPUTS = tuple(triton_grpc.InferRequestedOutput(key) for key in _OUTPUT_KEYS)
+
+
+def prepare_image(
+ path: str, height: int, width: int
+) -> tuple[triton_grpc.InferInput, np.ndarray, np.ndarray]:
+ """Prepares an image and converts it to an input for a Triton model server.
+
+ Args:
+ path: The file path to the image that needs to be processed.
+ height: The height of the image to be resized.
+ width: The width of the image to be resized.
+
+ Returns:
+ A tuple with the triton InferInput and both the original and resized
+ image.
+ """
+ image_bgr = cv2.imread(path)
+ image = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
+ image_resized = cv2.resize(
+ image, (width, height), interpolation=cv2.INTER_AREA
+ )
+ expanded_image = np.expand_dims(image_resized, axis=0)
+ inputs = triton_grpc.InferInput(
+ 'inputs', expanded_image.shape, datatype='UINT8'
+ )
+ inputs.set_data_from_numpy(expanded_image)
+ return inputs, image, image_resized
+
+
+def infer(
+ model_name: str, inputs: triton_grpc.InferInput
+) -> dict[str, Any]: # pyrefly: ignore[bad-return]
+ """Wraps inference and converts the result to a dictionary of output keys.
+
+ Args:
+ model_name: Model name in Triton Server.
+ inputs: The input data for inference.
+
+ Returns:
+ A dictionary of output keys and their corresponding values from
+ InferResult.
+ """
+ result = triton_grpc.InferenceServerClient(url='localhost:8001').infer(
+ model_name=model_name, inputs=[inputs], outputs=_OUTPUTS
+ )
+ if result:
+ return {key: result.as_numpy(key) for key in _OUTPUT_KEYS}
+
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/triton_server_inference_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/triton_server_inference_test.py
new file mode 100644
index 00000000000..0f3a3da03ed
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/triton_server_inference_test.py
@@ -0,0 +1,101 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+from unittest import mock
+import cv2
+import numpy as np
+from tritonclient import grpc as triton_grpc
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import triton_server_inference
+
+# Create a small 1x4 BGR (open cv default) test image
+BGR_TEST_IMAGE = np.zeros((1, 4, 3), dtype=np.uint8)
+BGR_TEST_IMAGE[0, 0] = [0, 0, 255] # Red in BGR
+BGR_TEST_IMAGE[0, 1] = [0, 255, 0]
+BGR_TEST_IMAGE[0, 2] = [255, 0, 0] # Blue in BGR
+BGR_TEST_IMAGE[0, 3] = [0, 255, 255]
+
+
+class TestTritonPrediction(unittest.TestCase):
+
+ @mock.patch.object(cv2, 'imread')
+ def test_input_conversion_to_rgb(self, mock_imread):
+ mock_imread.return_value = BGR_TEST_IMAGE
+
+ _, test_image, _ = (
+ triton_server_inference.prepare_image('/path/test_img.jpg', 5, 5)
+ )
+
+ # Check that a single BRG pixel is converted to RGB
+ self.assertEqual(test_image[0, 0].tolist(), [255, 0, 0])
+
+ @mock.patch.object(cv2, 'imread')
+ def test_input_image_resized(self, mock_imread):
+ mock_imread.return_value = BGR_TEST_IMAGE
+
+ _, _, test_image_resized = (
+ triton_server_inference.prepare_image('/path/test_img.jpg', 5, 5)
+ )
+
+ self.assertEqual(test_image_resized.shape, (5, 5, 3))
+
+ @mock.patch.object(cv2, 'imread')
+ def test_batch_dimension_prepended_to_triton_input(self, mock_imread):
+ mock_imread.return_value = BGR_TEST_IMAGE
+
+ test_triton_input, _, _ = (
+ triton_server_inference.prepare_image('/path/test_img.jpg', 5, 5)
+ )
+
+ self.assertEqual(test_triton_input.shape(), [1, 5, 5, 3])
+
+ @mock.patch.object(cv2, 'imread')
+ def test_image_converted_to_infer_input(self, mock_imread):
+ mock_imread.return_value = BGR_TEST_IMAGE
+
+ test_triton_input, _, _ = (
+ triton_server_inference.prepare_image('/path/test_img.jpg', 5, 5)
+ )
+
+ self.assertIsInstance(test_triton_input, triton_grpc.InferInput)
+
+ @mock.patch.object(triton_grpc.InferInput, 'set_data_from_numpy')
+ @mock.patch.object(cv2, 'imread')
+ def test_infer_input_set(self, mock_imread, mock_set_data_from_numpy):
+ mock_imread.return_value = BGR_TEST_IMAGE
+
+ triton_server_inference.prepare_image('/path/test_img.jpg', 5, 5)
+
+ # Check that the set_data_from_numpy method is called once. Triton
+ # InferInput data is a black-box, so we just check that it was set.
+ mock_set_data_from_numpy.assert_called_once()
+
+ @mock.patch.object(triton_grpc.InferenceServerClient, 'infer')
+ def test_inference_output_converted_to_dict(self, mock_query_model):
+ test_output_data = np.array([[1, 0]])
+ mock_infer_result = mock.create_autospec(
+ triton_grpc.InferResult, instance=True
+ )
+ mock_infer_result.as_numpy = lambda key: test_output_data
+ mock_query_model.return_value = mock_infer_result
+
+ result = triton_server_inference.infer('test_model', mock.MagicMock())
+
+ for key in triton_server_inference._OUTPUT_KEYS:
+ self.assertIn(key, result)
+ self.assertIsInstance(result[key], np.ndarray)
+
+
+if __name__ == '__main__':
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/utils.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/utils.py
new file mode 100644
index 00000000000..20ba3d260d2
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/utils.py
@@ -0,0 +1,510 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utility functions for the pipeline."""
+
+from collections.abc import Mapping, Sequence
+import csv
+import dataclasses
+import logging
+import os
+import subprocess
+import time
+from typing import Any, TypedDict
+
+import cv2
+import natsort
+import numpy as np
+import tensorflow as tf, tf_keras
+
+
+class ItemDict(TypedDict):
+ id: int
+ name: str
+ supercategory: str
+
+
+@dataclasses.dataclass
+class BoundingBox:
+ y1: int | float
+ x1: int | float
+ y2: int | float
+ x2: int | float
+
+
+@dataclasses.dataclass
+class ImageSize:
+ height: int
+ width: int
+
+
+def _reframe_image_corners_relative_to_boxes(boxes: tf.Tensor) -> tf.Tensor:
+ """Reframe the image corners ([0, 0, 1, 1]) to be relative to boxes.
+
+ The local coordinate frame of each box is assumed to be relative to
+ its own for corners.
+
+ Args:
+ boxes: A float tensor of [num_boxes, 4] of (ymin, xmin, ymax, xmax)
+ coordinates in relative coordinate space of each bounding box.
+
+ Returns:
+ reframed_boxes: Reframes boxes with same shape as input.
+ """
+ ymin, xmin, ymax, xmax = (boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3])
+
+ height = tf.maximum(ymax - ymin, 1e-4)
+ width = tf.maximum(xmax - xmin, 1e-4)
+
+ ymin_out = (0 - ymin) / height
+ xmin_out = (0 - xmin) / width
+ ymax_out = (1 - ymin) / height
+ xmax_out = (1 - xmin) / width
+ return tf.stack([ymin_out, xmin_out, ymax_out, xmax_out], axis=1)
+
+
+def _reframe_box_masks_to_image_masks(
+ box_masks: tf.Tensor,
+ boxes: tf.Tensor,
+ image_height: int,
+ image_width: int,
+ resize_method='bilinear'
+) -> tf.Tensor:
+ """Transforms the box masks back to full image masks.
+
+ Embeds masks in bounding boxes of larger masks whose shapes correspond to
+ image shape.
+ Args:
+ box_masks: A tensor of size [num_masks, mask_height, mask_width].
+ boxes: A tf.float32 tensor of size [num_masks, 4] containing the box
+ corners. Row i contains [ymin, xmin, ymax, xmax] of the box
+ corresponding to mask i. Note that the box corners are in
+ normalized coordinates.
+ image_height: Image height. The output mask will have the same height as
+ the image height.
+ image_width: Image width. The output mask will have the same width as the
+ image width.
+ resize_method: The resize method, either 'bilinear' or 'nearest'. Note that
+ 'bilinear' is only respected if box_masks is a float.
+ Returns:
+ A tensor of size [num_masks, image_height, image_width] with the same dtype
+ as `box_masks`.
+ """
+ resize_method = 'nearest' if box_masks.dtype == tf.uint8 else resize_method
+ def reframe_box_masks_to_image_masks_default():
+ """The default function when there are more than 0 box masks."""
+
+ num_boxes = tf.shape(box_masks)[0]
+ box_masks_expanded = tf.expand_dims(box_masks, axis=3)
+
+ resized_crops = tf.image.crop_and_resize(
+ image=box_masks_expanded,
+ boxes=_reframe_image_corners_relative_to_boxes(boxes),
+ box_indices=tf.range(num_boxes),
+ crop_size=[image_height, image_width],
+ method=resize_method,
+ extrapolation_value=0)
+ return tf.cast(resized_crops, box_masks.dtype)
+
+ image_masks = tf.cond(
+ tf.shape(box_masks)[0] > 0,
+ reframe_box_masks_to_image_masks_default,
+ lambda: tf.zeros([0, image_height, image_width, 1], box_masks.dtype))
+ return tf.squeeze(image_masks, axis=3)
+
+
+def _read_csv_to_list(file_path: str) -> Sequence[str]:
+ """Reads a CSV file and returns its contents as a list.
+
+ This function reads the given CSV file, skips the header, and assumes
+ there is only one column in the CSV. It returns the contents as a list of
+ strings.
+
+ Args:
+ file_path: The path to the CSV file.
+
+ Returns:
+ The contents of the CSV file as a list of strings.
+ """
+ data_list = []
+ with open(file_path, 'r') as csvfile:
+ reader = csv.reader(csvfile)
+ for row in reader:
+ data_list.append(row[0]) # Assuming there is only one column in the CSV
+ return data_list
+
+
+def _categories_dictionary(objects: Sequence[str]) -> Mapping[int, ItemDict]:
+ """This function takes a list of objects and returns a dictionaries.
+
+ A dictionary of objects, where each object is represented by a dictionary
+ with the following keys:
+ - id: The ID of the object.
+ - name: The name of the object.
+ - supercategory: The supercategory of the object.
+
+ Args:
+ objects: A list of strings, where each string is the name of an object.
+
+ Returns:
+ A tuple of two dictionaries, as described above.
+ """
+ category_index = {}
+
+ for num, obj_name in enumerate(objects, start=1):
+ obj_dict = {'id': num, 'name': obj_name, 'supercategory': 'objects'}
+ category_index[num] = obj_dict
+ return category_index
+
+
+def load_labels(
+ labels_path: str,
+) -> tuple[Sequence[str], Mapping[int, ItemDict]]:
+ """Loads labels from a CSV file and generates category mappings.
+
+ Args:
+ labels_path: Path to the CSV file containing label definitions.
+
+ Returns:
+ category_indices: A list of category indices.
+ category_index: A dictionary mapping category indices to ItemDict objects.
+ """
+ category_indices = _read_csv_to_list(labels_path)
+ category_index = _categories_dictionary(category_indices)
+ return category_indices, category_index
+
+
+def files_paths(folder_path):
+ """List the full paths of image files in a folder and sort them.
+
+ Args:
+ folder_path: The path of the folder to list the image files from.
+
+ Returns:
+ A list of full paths of the image files in the folder, sorted in ascending
+ order.
+ """
+ img_extensions = ('.jpg', '.jpeg', '.png', '.gif', '.bmp', '.tiff', '.webp')
+ image_files_full_path = []
+
+ for entry in os.scandir(folder_path):
+ if entry.is_file() and entry.name.lower().endswith(img_extensions):
+ image_files_full_path.append(entry.path)
+
+ # Sort the list of files by name
+ image_files_full_path = natsort.natsorted(image_files_full_path)
+
+ return image_files_full_path
+
+
+def create_log_file(name: str, logs_folder_path: str) -> logging.Logger:
+ """Creates a logger and a log file given the name of the video.
+
+ Args:
+ name: The name of the video.
+ logs_folder_path: Path to the directory where logs should be saved.
+
+ Returns:
+ logging.Logger: Logger object configured to write logs to the file.
+ """
+ log_file_path = os.path.join(logs_folder_path, f'{name}.log')
+ logger = logging.getLogger(name)
+ logger.setLevel(logging.INFO)
+ file_handler = logging.FileHandler(log_file_path)
+ formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
+ file_handler.setFormatter(formatter)
+ logger.addHandler(file_handler)
+ return logger
+
+
+def reframe_masks(
+ results: Mapping[str, Any], boxes: str, height: int, width: int
+) -> np.ndarray:
+ """Reframe the masks to an image size.
+
+ Args:
+ results: The detection results from the model.
+ boxes: The detection boxes.
+ height: The height of the original image.
+ width: The width of the original image.
+
+ Returns:
+ The reframed masks.
+ """
+ detection_masks = results['detection_masks'][0]
+ detection_boxes = results[boxes][0]
+ detection_masks_reframed = _reframe_box_masks_to_image_masks(
+ detection_masks, detection_boxes, height, width
+ )
+ detection_masks_reframed = tf.cast(detection_masks_reframed > 0.5, np.uint8)
+ detection_masks_reframed = detection_masks_reframed.numpy()
+ return detection_masks_reframed
+
+
+def _calculate_area(mask: np.ndarray) -> int:
+ """Calculate the area of the mask.
+
+ Args:
+ mask: The mask to calculate the area of.
+
+ Returns:
+ The area of the mask.
+ """
+ return np.sum(mask)
+
+
+def _calculate_iou(mask1: np.ndarray, mask2: np.ndarray) -> float:
+ """Calculate the intersection over union (IoU) between two masks.
+
+ Args:
+ mask1: The first mask.
+ mask2: The second mask.
+
+ Returns:
+ The intersection over union (IoU) between the two masks.
+ """
+ intersection = np.logical_and(mask1, mask2).sum()
+ union = np.logical_or(mask1, mask2).sum()
+ return intersection / union if union != 0 else 0
+
+
+def _is_contained(mask1: np.ndarray, mask2: np.ndarray) -> bool:
+ """Check if mask1 is entirely contained within mask2.
+
+ Args:
+ mask1: The first mask.
+ mask2: The second mask.
+
+ Returns:
+ True if mask1 is entirely contained within mask2, False otherwise.
+ """
+ return np.array_equal(np.logical_and(mask1, mask2), mask1)
+
+
+# TODO: b/416838511 - Reduce the nesting statement in the for loop.
+def filter_masks(
+ masks: np.ndarray,
+ iou_threshold: float = 0.8,
+ area_threshold: int | None = None,
+) -> Sequence[int]:
+ """Filter the overlapping masks.
+
+ Filter the masks based on the area and intersection over union (IoU).
+
+ Args:
+ masks: The masks to filter.
+ iou_threshold: The threshold for the intersection over union (IoU) between
+ two masks.
+ area_threshold: The threshold for the area of the mask.
+
+ Returns:
+ The indices of the unique masks.
+ """
+ # Calculate the area for each mask
+ areas = np.array([_calculate_area(mask) for mask in masks])
+
+ # Sort the masks based on area in descending order
+ sorted_indices = np.argsort(areas)[::-1]
+ sorted_masks = masks[sorted_indices]
+ sorted_areas = areas[sorted_indices]
+
+ unique_indices = []
+
+ for i, mask in enumerate(sorted_masks):
+ if (
+ area_threshold is not None and sorted_areas[i] > area_threshold
+ ) or sorted_areas[i] < 4000:
+ continue
+
+ keep = True
+ for j in range(i):
+ if _calculate_iou(mask, sorted_masks[j]) > iou_threshold or _is_contained(
+ mask, sorted_masks[j]
+ ):
+ keep = False
+ break
+ if keep:
+ unique_indices.append(sorted_indices[i])
+
+ return unique_indices
+
+
+def resize_each_mask(
+ masks: np.ndarray, target_height: int, target_width: int
+) -> np.ndarray:
+ """Resize each mask to the target height and width.
+
+ Args:
+ masks: The masks to resize.
+ target_height: The target height of the resized masks.
+ target_width: The target width of the resized masks.
+
+ Returns:
+ The resized masks.
+ """
+ combined_masks = []
+ for i in masks:
+ mask = cv2.resize(
+ i, (target_width, target_height), interpolation=cv2.INTER_NEAREST
+ )
+ combined_masks.append(mask)
+ return np.array(combined_masks)
+
+
+def extract_and_resize_objects(
+ results: Mapping[str, Any],
+ masks: str,
+ boxes: str,
+ image: np.ndarray,
+ resize_factor: float = 0.5,
+) -> Sequence[np.ndarray]:
+ """Extract and resize objects from the detection results.
+
+ Args:
+ results: The detection results from the model.
+ masks: The masks to extract objects from.
+ boxes: The bounding boxes of the objects.
+ image: The image to extract objects from.
+ resize_factor: The factor by which to resize the objects.
+
+ Returns:
+ A list of cropped objects.
+ """
+ cropped_objects = []
+
+ for i, mask in enumerate(results[masks]):
+ ymin, xmin, ymax, xmax = results[boxes][0][i]
+ mask = np.expand_dims(mask, axis=-1)
+
+ # Crop the object using the mask and bounding box
+ cropped_object = np.where(
+ mask[ymin:ymax, xmin:xmax], image[ymin:ymax, xmin:xmax], 0
+ )
+
+ # Calculate new dimensions
+ new_width = int(cropped_object.shape[1] * resize_factor)
+ new_height = int(cropped_object.shape[0] * resize_factor)
+ cropped_object = cv2.resize(
+ cropped_object, (new_width, new_height), interpolation=cv2.INTER_AREA
+ )
+ cropped_objects.append(cropped_object)
+
+ return cropped_objects
+
+
+def adjust_image_size(
+ height: int, width: int, min_size: int
+) -> tuple[int, int]:
+ """Adjust the image size to ensure both dimensions are at least of min_size.
+
+ Args:
+ height: The height of the image.
+ width: The width of the image.
+ min_size: Minimum size of the image dimension needed.
+
+ Returns:
+ The adjusted height and width of the image.
+ """
+ if height < min_size or width < min_size:
+ return height, width
+
+ # Calculate the scale factor to ensure both dimensions remain at least 1024
+ scale_factor = min(height / min_size, width / min_size)
+ return int(height / scale_factor), int(width / scale_factor)
+
+
+def filter_detections(
+ results: Mapping[str, np.ndarray],
+ valid_indices: Sequence[int] | Sequence[bool],
+) -> Mapping[str, np.ndarray]:
+ """Filter the detection results based on the valid indices.
+
+ Args:
+ results: The detection results from the model.
+ valid_indices: The indices of the valid detections.
+
+ Returns:
+ The filtered detection results.
+ """
+ if np.array(valid_indices).dtype == bool:
+ new_num_detections = int(np.sum(valid_indices))
+ else:
+ new_num_detections = len(valid_indices)
+
+ # Define the keys to filter
+ keys_to_filter = [
+ 'detection_masks',
+ 'detection_masks_resized',
+ 'detection_masks_reframed',
+ 'detection_classes',
+ 'detection_boxes',
+ 'normalized_boxes',
+ 'detection_scores',
+ ]
+
+ filtered_output = {}
+
+ for key in keys_to_filter:
+ if key in results:
+ if key == 'detection_masks':
+ filtered_output[key] = results[key][:, valid_indices, :, :]
+ elif key in ['detection_masks_resized', 'detection_masks_reframed']:
+ filtered_output[key] = results[key][valid_indices, :, :]
+ elif key in ['detection_boxes', 'normalized_boxes']:
+ filtered_output[key] = results[key][:, valid_indices, :]
+ elif key in [
+ 'detection_classes',
+ 'detection_scores',
+ 'detection_classes_names',
+ ]:
+ filtered_output[key] = results[key][:, valid_indices]
+ filtered_output['image_info'] = results['image_info']
+ filtered_output['num_detections'] = np.array([new_num_detections])
+
+ return filtered_output
+
+
+def resize_bbox(
+ bbox: BoundingBox, old_size: ImageSize, new_size: ImageSize
+) -> tuple[int, int, int, int]:
+ """Resize bounding box coordinates based on new image size.
+
+ Args:
+ bbox: BoundingBox with original coordinates.
+ old_size: Original image size.
+ new_size: New image size.
+
+ Returns:
+ Rescaled bounding box coordinates.
+ """
+ scale_x = new_size.width / old_size.width
+ scale_y = new_size.height / old_size.height
+
+ new_y1 = int(bbox.y1 * scale_y)
+ new_x1 = int(bbox.x1 * scale_x)
+ new_y2 = int(bbox.y2 * scale_y)
+ new_x2 = int(bbox.x2 * scale_x)
+
+ return new_y1, new_x1, new_y2, new_x2
+
+
+def shutdown_system():
+ """Shuts down the system."""
+ time.sleep(60)
+ print('Attempting to shut down the VM instance...')
+ try:
+ command = ['sudo', 'poweroff']
+ subprocess.run(command, check=True)
+ except subprocess.CalledProcessError as e:
+ print(f'Failed to shut down: {e}')
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/utils_test.py b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/utils_test.py
new file mode 100644
index 00000000000..4b4412726dd
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/client/utils_test.py
@@ -0,0 +1,249 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import os
+import tempfile
+import unittest
+import numpy as np
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client import utils
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client.utils import BoundingBox
+from official.projects.waste_identification_ml.Triton_TF_Cloud_Deployment.client.utils import ImageSize
+
+
+class TestLoadLabels(unittest.TestCase):
+
+ def test_load_labels(self):
+ # Create a temporary CSV file within the test
+ with tempfile.NamedTemporaryFile(mode='w+', delete=False) as temp_csv:
+ temp_csv.write('Label\nBottle\nCan\nCup\n')
+ temp_csv_path = temp_csv.name
+
+ try:
+ # Call the function under test
+ category_indices, category_index = utils.load_labels(temp_csv_path)
+
+ # Expected results
+ expected_list = ['Label', 'Bottle', 'Can', 'Cup']
+ expected_dict = {
+ 1: {'id': 1, 'name': 'Label', 'supercategory': 'objects'},
+ 2: {'id': 2, 'name': 'Bottle', 'supercategory': 'objects'},
+ 3: {'id': 3, 'name': 'Can', 'supercategory': 'objects'},
+ 4: {'id': 4, 'name': 'Cup', 'supercategory': 'objects'},
+ }
+
+ self.assertEqual(category_indices, expected_list)
+ self.assertEqual(category_index, expected_dict)
+
+ finally:
+ # Ensure the temporary file is deleted even if assertions fail
+ os.remove(temp_csv_path)
+
+ def test_files_paths_with_images(self):
+ # Create a temporary directory
+ with tempfile.TemporaryDirectory() as temp_dir:
+ # Create some image and non-image files
+ filenames = ['img2.jpg', 'img1.png', 'doc1.txt', 'photo.gif']
+ for filename in filenames:
+ open(os.path.join(temp_dir, filename), 'a').close()
+
+ # Call the function under test
+ result = utils.files_paths(temp_dir)
+
+ # Expected image files sorted naturally
+ expected = [
+ os.path.join(temp_dir, 'img1.png'),
+ os.path.join(temp_dir, 'img2.jpg'),
+ os.path.join(temp_dir, 'photo.gif'),
+ ]
+
+ self.assertEqual(result, expected)
+
+ def test_files_paths_with_no_images(self):
+ with tempfile.TemporaryDirectory() as temp_dir:
+ # Create only non-image files
+ filenames = ['doc1.txt', 'readme.md']
+ for filename in filenames:
+ open(os.path.join(temp_dir, filename), 'a').close()
+
+ result = utils.files_paths(temp_dir)
+ self.assertEqual(result, []) # Should return an empty list
+
+ def test_files_paths_empty_folder(self):
+ with tempfile.TemporaryDirectory() as temp_dir:
+ result = utils.files_paths(temp_dir)
+ self.assertEqual(result, [])
+
+ def test_resize_multiple_masks(self):
+ # Create two 5x5 masks
+ mask1 = np.zeros((5, 5), dtype=np.uint8)
+ mask2 = np.ones((5, 5), dtype=np.uint8)
+ masks = np.array([mask1, mask2])
+
+ resized_masks = utils.resize_each_mask(masks, 3, 3)
+
+ self.assertEqual(resized_masks.shape, (2, 3, 3))
+ self.assertTrue((resized_masks[0] == 0).all())
+ self.assertTrue((resized_masks[1] == 1).all())
+
+ def test_keeps_biggest_mask(self):
+ # Create larger masks to satisfy the area check (area >= 4000)
+ mask_small = np.zeros((100, 100), dtype=int)
+ mask_small[0:40, 0:40] = 1 # Area = 1600 (will be skipped)
+
+ mask_medium = np.zeros((100, 100), dtype=int)
+ mask_medium[0:70, 0:70] = 1 # Area = 4900 (passes area condition)
+
+ mask_large = np.zeros((100, 100), dtype=int)
+ mask_large[:, :] = 1 # Area = 10000 (passes area condition)
+
+ masks = np.array([mask_small, mask_medium, mask_large])
+
+ # Run filter_masks without specifying area_threshold
+ result = utils.filter_masks(masks, iou_threshold=0.5)
+
+ # Expect only the largest mask (index 2) to remain
+ self.assertEqual(result, [2])
+
+ def test_filter_with_boolean_indices(self):
+ results = {
+ 'detection_masks': np.random.rand(1, 3, 5, 5),
+ 'detection_masks_resized': np.random.rand(3, 5, 5),
+ 'detection_boxes': np.random.rand(1, 3, 4),
+ 'detection_classes': np.array([[1, 2, 3]]),
+ 'detection_scores': np.array([[0.9, 0.8, 0.3]]),
+ 'image_info': np.array([[640, 480]]),
+ }
+
+ valid_indices = [True, False, True]
+
+ output = utils.filter_detections(results, valid_indices)
+
+ self.assertEqual(output['detection_masks'].shape[1], 2)
+ self.assertEqual(output['detection_masks_resized'].shape[0], 2)
+ self.assertEqual(output['detection_boxes'].shape[1], 2)
+ self.assertEqual(output['detection_classes'].shape[1], 2)
+ self.assertEqual(output['detection_scores'].shape[1], 2)
+ self.assertTrue(np.array_equal(output['image_info'], results['image_info']))
+ self.assertEqual(output['num_detections'][0], 2)
+
+ def test_filter_with_integer_indices(self):
+ results = {
+ 'detection_masks': np.random.rand(1, 4, 5, 5),
+ 'detection_masks_resized': np.random.rand(4, 5, 5),
+ 'detection_boxes': np.random.rand(1, 4, 4),
+ 'detection_classes': np.array([[1, 2, 3, 4]]),
+ 'detection_scores': np.array([[0.9, 0.8, 0.3, 0.6]]),
+ 'image_info': np.array([[640, 480]]),
+ }
+
+ valid_indices = [0, 2] # Keep detections at index 0 and 2
+
+ output = utils.filter_detections(results, valid_indices)
+
+ self.assertEqual(output['detection_masks'].shape[1], 2)
+ self.assertEqual(output['detection_masks_resized'].shape[0], 2)
+ self.assertEqual(output['detection_boxes'].shape[1], 2)
+ self.assertEqual(output['detection_classes'].shape[1], 2)
+ self.assertEqual(output['detection_scores'].shape[1], 2)
+ self.assertEqual(output['num_detections'][0], 2)
+
+ def test_both_dimensions_below_min_size(self):
+ height, width, min_size = 800, 900, 1024
+
+ result = utils.adjust_image_size(height, width, min_size)
+
+ self.assertEqual(result, (800, 900)) # No scaling should happen
+
+ def test_height_below_min_size(self):
+ height, width, min_size = 900, 1200, 1024
+
+ result = utils.adjust_image_size(height, width, min_size)
+
+ self.assertEqual(result, (900, 1200)) # No scaling
+
+ def test_width_below_min_size(self):
+ height, width, min_size = 1300, 800, 1024
+
+ result = utils.adjust_image_size(height, width, min_size)
+
+ self.assertEqual(result, (1300, 800)) # No scaling
+
+ def test_both_dimensions_above_min_size(self):
+ height, width, min_size = 2048, 1536, 1024
+ expected_scale = min(height / min_size, width / min_size)
+ expected_height = int(height / expected_scale)
+ expected_width = int(width / expected_scale)
+
+ result = utils.adjust_image_size(height, width, min_size)
+
+ self.assertEqual(result, (expected_height, expected_width))
+
+ def test_exact_min_size(self):
+ height, width, min_size = 1024, 1024, 1024
+
+ result = utils.adjust_image_size(height, width, min_size)
+
+ self.assertEqual(result, (1024, 1024)) # Already meets the requirement
+
+ def test_extract_and_resize_single_object(self):
+ image = np.ones((10, 10, 3), dtype=np.uint8) * 255 # white image
+
+ # Define a simple binary mask (1 in a 4x4 box)
+ mask = np.zeros((10, 10), dtype=np.uint8)
+ mask[2:6, 3:7] = 1
+
+ # Box coordinates match the mask
+ boxes = np.array([[[2, 3, 6, 7]]], dtype=np.int32) # shape (1, 1, 4)
+
+ results = {'masks': [mask], 'boxes': boxes}
+
+ cropped_objects = utils.extract_and_resize_objects(
+ results, 'masks', 'boxes', image, resize_factor=0.5
+ )
+
+ self.assertEqual(len(cropped_objects), 1)
+ obj = cropped_objects[0]
+
+ # Original crop size is (4, 4), so resized should be (2, 2)
+ self.assertEqual(obj.shape[:2], (2, 2))
+
+ # Should still be 3 channels
+ self.assertEqual(obj.shape[2], 3)
+
+ # The output pixels in mask area should be non-zero
+ self.assertTrue(np.any(obj > 0))
+
+ def test_resize_bbox_scaling(self):
+ bbox = BoundingBox(y1=50, x1=100, y2=150, x2=200)
+ old_size = ImageSize(height=200, width=400)
+ new_size = ImageSize(height=400, width=800)
+
+ expected = (100, 200, 300, 400)
+ result = utils.resize_bbox(bbox, old_size, new_size)
+
+ self.assertEqual(result, expected)
+
+ def test_resize_bbox_no_scaling(self):
+ bbox = BoundingBox(y1=10, x1=20, y2=30, x2=40)
+ old_size = ImageSize(height=100, width=100)
+ new_size = ImageSize(height=100, width=100)
+
+ expected = (10, 20, 30, 40)
+ result = utils.resize_bbox(bbox, old_size, new_size)
+
+ self.assertEqual(result, expected)
+
+
+if __name__ == '__main__':
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/server/triton_inference_server.sh b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/server/triton_inference_server.sh
new file mode 100644
index 00000000000..cfb0c4a159f
--- /dev/null
+++ b/official/projects/waste_identification_ml/Triton_TF_Cloud_Deployment/server/triton_inference_server.sh
@@ -0,0 +1,51 @@
+#!/bin/bash
+
+# Summary
+cat << END
+Automates the setup and deployment of a TensorFlow SavedModel on the NVIDIA Triton Inference Server.
+It cleans up old setups, downloads and organizes models, creates the required configuration,
+ensures screen is installed, and launches Triton in a detached session with GPU support.
+END
+
+# Check and delete the model_repository directory if it exists
+if [ -d "model_repository" ]; then
+ echo "Removing existing model_repository directory..."
+ rm -rf model_repository
+fi
+
+# Define an associative array with model names and their URLs
+declare -A models=(
+ ["Jan2025_ver2_merged_1024_1024"]="https://storage.googleapis.com/"\
+"tf_model_garden/vision/waste_identification_ml/"\
+"Jan2025_ver2_merged_1024_1024.zip"
+)
+
+# Download, unzip, and organize models
+for model_name in "${!models[@]}"; do
+ url=${models[$model_name]}
+ zip_file="${url##*/}"
+
+ wget $url && unzip $zip_file
+
+ mkdir -p model_repository/$model_name/1/model.savedmodel
+
+ echo -e "name: \"$model_name\"\nplatform: \"tensorflow_savedmodel\"\n\
+ max_batch_size : 0" > model_repository/$model_name/config.pbtxt
+
+ mv $model_name/* model_repository/$model_name/1/model.savedmodel/
+
+ rm -r $model_name
+ rm $zip_file
+done
+
+# Install screen if not already installed
+command -v screen >/dev/null 2>&1 || { \
+ sudo apt update && sudo apt install -y screen; \
+}
+
+# Start Triton server
+screen -dmS server bash -c '
+sudo docker run --gpus all --rm -p 8000:8000 -p 8001:8001 -p 8002:8002 \
+-v ${PWD}/model_repository:/models \
+nvcr.io/nvidia/tritonserver:24.03-py3 \
+tritonserver --model-repository=/models --backend-config=tensorflow,version=2'
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/__init__.py b/official/projects/waste_identification_ml/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/waste_identification_ml/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/waste_identification_ml/circularnet-docs/archetypes/default.md b/official/projects/waste_identification_ml/circularnet-docs/archetypes/default.md
new file mode 100644
index 00000000000..c6f3fcef6e3
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/archetypes/default.md
@@ -0,0 +1,5 @@
++++
+title = '{{ replace .File.ContentBaseName "-" " " | title }}'
+date = {{ .Date }}
+draft = true
++++
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/_index.md
new file mode 100644
index 00000000000..f6eef19a10b
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/_index.md
@@ -0,0 +1,81 @@
+## Table of contents
+
+**[CircularNet overview](#circularnet-overview)**
+
+ * [Get started with CircularNet](#get-started-with-circularnet)
+
+**[Discover CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/_index.md)**
+
+* [Benefits of CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/benefits-of-cn.md)
+* [When to use CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/when-to-use-cn.md)
+* [How CircularNet works](/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/how-cn-works.md)
+
+**[Choose a deployment solution](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md)**
+
+* [Where to host the models](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md#where-to-host-the-models)
+ * [Cloud deployment](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md#cloud-deployment)
+ * [Edge device deployment](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md#edge-device-deployment)
+ * [In-house server deployment](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md#in-house-server-deployment)
+
+**[Set up the system requirements](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/_index.md)**
+
+* [Choose a camera](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/_index.md)
+ * [Recommendations for selecting a machine vision camera](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/camera-recommendations.md)
+ * [Factors based on vision system placement](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md)
+ * [Sensor size](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md#sensor-size)
+ * [Focal length](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md#focal-length)
+ * [Aperture size (f-number)](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md#aperture-size-f-number)
+ * [Shutter speed](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md#shutter-speed)
+ * [Table of specifications](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/table-of-specs.md)
+* [Choose edge device hardware](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-edge-device/_index.md)
+
+**[Deploy CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/_index.md)**
+
+* [Before you begin](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/before-you-begin.md)
+* [Clone the repository and install packages](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/clone-repo.md)
+* [Start the server](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server.md)
+
+**[Prepare and analyze images](/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/_index.md)**
+
+* [Learn about the prediction pipeline](/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/learn-about-pipeline.md)
+* [Apply the prediction pipeline in Google Cloud](/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-cloud.md)
+* [Apply the prediction pipeline in an edge device](/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-edge.md)
+
+**[View data analysis and reporting](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/_index.md)**
+
+* [Before you begin](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/before-you-begin.md)
+* [Configure the dashboard](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/configure-dashboard.md)
+
+**[Retrain CircularNet models](/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/_index.md)**
+
+* [Before you begin](/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/before-you-begin.md)
+* [Prepare the training data](/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/prepare-data.md)
+* [Launch the training job](/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/launch-job.md)
+
+## CircularNet overview
+
+CircularNet is a free computer vision model developed by Google that utilizes
+artificial intelligence (AI) and machine learning (ML) to provide detailed and
+accurate identification of waste streams and recyclables. Trained on a diverse
+global dataset, CircularNet aims to make waste management analytics accessible
+and promote data-driven decision-making. It supports efforts to keep valuable
+resources out of landfills and in circulation. Open access and collaboration are
+fundamental to CircularNet's vision. Its open-source models, powered by
+[TensorFlow](https://www.tensorflow.org/) and available on
+[GitHub](https://github.com/tensorflow/models/tree/master/official/projects/waste_identification_ml),
+are free to use, customizable, and can help bring analytics to new markets while
+minimizing cost.
+
+This guide offers step-by-step instructions for setting up and integrating
+CircularNet, accommodating various deployment options. It describes different
+deployment preferences so you can install CircularNet according to your needs.
+
+### Get started with CircularNet
+
+Start exploring CircularNet by reviewing the following documentation:
+
+1. Discover [the benefits, features, components, and use cases](/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/) of CircularNet.
+1. Choose between [the different deployment options](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/) to install CircularNet models.
+1. Learn about [the recommendations for installing the camera](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/) you require to capture images.
+1. Follow a step-by-step solution example to [deploy CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/) and [prepare your captured images](/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/) for analysis and object tracking.
+1. Learn how to connect your data with [a dashboard for visualization and reporting](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/).
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/_index.md
new file mode 100644
index 00000000000..ab77ee117bd
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/_index.md
@@ -0,0 +1,31 @@
+CircularNet provides an image-analysis model. It detects _material types_, and
+_material forms_. The model utilizes a Mask R-CNN algorithm for image training
+and implements ResNet or MobileNet as the convolutional neural networks for
+image classification tasks.
+
+The model is loaded sequentially to achieve accurate predictions. When working
+with images, each image undergoes preprocessing before the model uses them for
+prediction. In the case of video files, the video is split into individual
+frames at a given frame rate. These frames are then processed in the same
+sequential manner as images.
+
+The predictions from the model result in two distinct outputs, which are
+then post-processed and combined into a single comprehensive output. This output
+includes critical information such as the number of detected objects, their
+bounding boxes, class names, class IDs, and masks for each object. Further
+computer vision techniques extract multiple properties of each object, including
+color detection. These properties facilitate object tracking and help eliminate
+duplicate object counts.
+
+If you [deploy the server](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server) on Google Cloud, you can
+automate the entire image analysis workflow within your VM instance. Integration
+with BigQuery tables, storage buckets, and dashboards allows for seamless data
+flow and real-time updates. A [prediction pipeline](./learn-about-pipeline) for
+Google Cloud pushes the data directly to storage buckets and BigQuery tables,
+which you can connect to the dashboard for [visualization and analysis](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/).
+
+On the other hand, direct data transfer to the cloud for edge device implementations needs a client-side configuration. A [prediction pipeline](/official/projects/waste_identification_ml/circularnet-docs/content/learn-about-pipeline) for devices lets you load the model and store image analysis results locally.
+
+This section describes how to apply the specialized CircularNet model using
+a prediction pipeline on the client side to prepare and analyze the images you
+capture.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/learn-about-pipeline.md b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/learn-about-pipeline.md
new file mode 100644
index 00000000000..e227306bc3c
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/learn-about-pipeline.md
@@ -0,0 +1,29 @@
+CircularNet offers prediction pipelines for processing, analyzing, and
+performing object recognition on video or image files. These pipelines
+facilitate systematic and automated video and image analysis using the Mask
+R-CNN algorithm and additional object detection and feature extraction
+processes.
+
+You can run a prediction pipeline from a script to automatically apply the two
+specialized models that analyze images or video frames. A prediction pipeline
+operates through the following series of actions in a specific order to ensure
+reliable and consistent results:
+
+1. Organize your videos or images chronologically according to their creation
+ time or another time-related metadata.
+1. Import files one at a time. If your files are videos, the pipeline decomposes
+ them into individual frames and runs two Mask R-CNN models for pixel-level
+ instance segmentation on each frame.
+1. Implement a color detection algorithm to identify and categorize the colors
+ of the detected objects within the frames or images.
+1. Extract and record features from the detected objects, facilitating analysis
+ and machine learning applications.
+1. Output prediction results of each frame or image with overlaid masks and
+ identification of detected objects.
+
+After processing all frames of a single video, the pipeline implements an
+object-tracking algorithm to identify and eliminate duplicate occurrences of
+objects across sequential frames, enhancing the accuracy of object detection and
+analysis.
+
+Moreover, [applying a prediction pipeline in Google Cloud](/official/projects/waste_identification_ml/circularnet-docs/content/prediction-pipeline-in-cloud) automatically uploads raw images and prediction results into [BigQuery](https://cloud.google.com/bigquery) tables. This seamless integration allows you to combine [visualization dashboards with analytical reports](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/) effortlessly.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-cloud.md b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-cloud.md
new file mode 100644
index 00000000000..3c7548436d0
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-cloud.md
@@ -0,0 +1,206 @@
+Apart from applying the prediction models to analyze images, the script that runs [the prediction pipeline](./learn-about-pipeline) on Google Cloud ingests data into [BigQuery](https://cloud.google.com/bigquery) to store all the image analysis details.
+
+After [setting up a server](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server) in a Google Cloud account, you can start recording videos of objects passing on your conveyor belt to gather data for analysis. The next step is transferring those video or image files to a [Cloud Storage](https://cloud.google.com/storage) bucket, where the prediction pipeline processes the images.
+
+The results of each video or image prediction are stored and appended to two
+BigQuery tables, ensuring efficient data management. After systematically
+processing each file in the Cloud Storage bucket, the pipeline exports and
+stores the results in another bucket for further use and analysis.
+
+This page explains how to create and manage Cloud Storage buckets on Google
+Cloud for the videos you record, run the prediction pipeline to apply the
+models, and process images for further analysis. You can then connect a
+visualization dashboard to the BigQuery tables to [display results as charts and reports](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/).
+
+{{< table_of_contents >}}
+
+---
+
+## Store videos in Cloud Storage buckets
+
+To effectively manage and process the videos or images recorded by the [machine vision camera](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/) capturing objects on the conveyor belt, you need the following storage devices:
+
+- **Local disk storage**: Temporarily cache or store files locally from the camera.
+- **Cloud Storage input bucket**: Store your recorded videos or images. You can automate uploads using a storage transfer agent.
+- **Cloud Storage output bucket**: Store the prediction results.
+
+By using these storage options, you ensure an organized, automated, and
+efficient workflow for managing and processing your recorded videos or images
+and prediction results.
+
+### Create the Cloud Storage input and output buckets
+
+[Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) in your Google
+Cloud account that serves as the _input_ repository for the recorded videos.
+This bucket stores all the videos or images captured by the machine vision
+camera and centralizes the storage of your recorded files, making them easily
+accessible for further processing.
+
+Before uploading to the input bucket, store the files temporarily on your local
+disk in a predetermined location. This temporary storage ensures you have a
+local backup and can manage uploads efficiently.
+
+To streamline the upload process to the input bucket, you can schedule a storage job to run at your expected frequency. This automated job takes care of transferring videos or images from the local disk to the Cloud Storage input bucket, relieving you of the manual task. Additionally, implementing a storage transfer agent further facilitates the transfer process, ensuring reliable and efficient uploads. For more information, see [Manage transfer agents](https://cloud.google.com/storage-transfer/docs/managing-on-prem-agents). Otherwise, manually upload your files stored or cached in your local storage to the input bucket before [running the prediction pipeline](#run-the-prediction-pipeline). For more information about manual uploads, see [Upload objects from a file system](https://cloud.google.com/storage/docs/uploading-objects).
+
+In addition to the input bucket, create another Cloud Storage _output_ bucket to
+store the results generated by your prediction pipeline. This output bucket
+serves as a well-organized repository for the processed data, such as
+predictions and analysis results, ensuring that the output from your models is
+easily accessible and you can manage your pipeline's output.
+
+## Grant required permissions
+
+The [service account](https://cloud.google.com/compute/docs/access/service-accounts#default_service_account)
+from your NVIDIA T4 GPU virtual machine (VM) instance must have access to Cloud
+Storage buckets and BigQuery tables to let the VM instance run the prediction
+pipeline and upload the results to BigQuery.
+
+The VM you created when [deploying CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/) has a service
+account that must act as a _principal_ with the following two roles:
+
+- Grant the Storage Admin role to access your Cloud Storage input and output buckets. [Add the VM service account as a principal to a bucket-level policy](https://cloud.google.com/storage/docs/access-control/using-iam-permissions#bucket-add) for both buckets.
+- Grant the BigQuery Admin role to access your BigQuery tables. [Manage roles of the VM service account using the Google Cloud console](https://cloud.google.com/iam/docs/manage-access-service-accounts).
+
+## Run the prediction pipeline
+
+Follow these steps to run the prediction pipeline and process your input files
+on Google Cloud:
+
+1. [Open the Google Cloud console](https://cloud.google.com/cloud-console).
+2. [Create the Cloud Storage input and output buckets](#create-the-cloud-storage-input-and-output-buckets).
+3. Upload the images or videos your machine vision camera captured to the Cloud
+ Storage input bucket.
+4. [Grant the required permissions to the VM service account](#grant-required-permissions).
+5. From the **Navigation menu** on the Google Cloud console, select **Compute Engine** > **VM instances**.
+6. On the **VM instances** page, find the VM instance you created with the
+ NVIDIA T4 GPU when [deploying CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/).
+7. If you stopped your VM instance, restart it by clicking **More actions** >
+ **Start / Resume** in the row of the instance that you want to restart.
+
+ **Note:** Stop the instance after you finish using it to avoid unnecessary
+ expenses.
+
+8. Click **SSH** in the row of the instance that you want to connect to. The
+ **SSH-in-Browser** tool opens. For more information, see [Connect to VMs](https://cloud.google.com/compute/docs/connect/standard-ssh#connect_to_vms).
+9. On the **SSH-in-browser** window, [start the server](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server).
+10. Display the names of the models you loaded to the Triton inference server:
+
+ ```
+ cat triton_inference_server.sh
+ ```
+
+ The first lines of the output show the names of the loaded models in square
+ brackets. You can call any of these models when running the prediction
+ pipeline.
+
+ **Important:** Run the previous command in the `server` folder, which
+ contains the `triton_inference_server.sh` script.
+
+11. Exit the `server` folder and open the `client` folder in the
+ `prediction_pipeline` directory:
+
+ ```
+ cd ..
+ cd client/
+ ```
+
+ This folder contains the `pipeline_images.py` and `pipeline_videos.py`
+ Python files that store the complete prediction pipelines for input images
+ or videos, respectively. The `run_gcp_images.sh` and `run_gcp_videos.sh`
+ scripts run these Python files automatically.
+
+12. If you have to modify the scripts to provide your specific paths and values
+ for the prediction pipeline, edit the corresponding parameter values on the
+ script. The following example modifies the image pipeline script:
+
+ ```
+ vim run_images.sh
+ ```
+
+ The Vim editor displays the following parameters:
+
+ ```
+ --input_directory=
+ --output_directory=
+ --height=
+ --width=
+ --model=
+ --score=
+ --search_range_x=
+ --search_range_y=
+ --memory=
+ --project_id=
+ --bq_dataset_id=
+ --bq_table_id=
+ --overwrite=
+ --tracking_visualization=
+ --cropped_objects=
+ ```
+
+ Replace the following:
+
+ - ``: The path to [the Cloud Storage input bucket you
+ created](#create-the-cloud-storage-input-and-output-buckets), for example
+ `gs://my-input-bucket/`.
+ - ``: The path to [the Cloud Storage output bucket
+ you created](#create-the-cloud-storage-input-and-output-buckets), for
+ example `gs://my-output-bucket/`.
+ - ``: The rate at which you want to capture images from
+ videos to split videos into frames, for example, `15`.
+ - ``: The height in pixels of the image or video frames that the
+ model expects for prediction, for example, `512`.
+ - ``: The width in pixels of the image or video frames that the
+ model expects for prediction, for example, `1024`.
+ - ``: The name of the CircularNet model in the Triton
+ inference server that you want to call, for example,
+ `Jan2025_ver2_merged_1024_1024`.
+ - ``: The threshold for model prediction, for example, `0.70`.
+ - ``: The pixels up to which you want to track an object for
+ object tracking in consecutive frames, for example, `100`.
+ - ``: The frames up to which you want to track an object, for
+ example, `20`.
+ - ``: The ID of your Google Cloud project, for example,
+ `my-project`.
+ - ``: The ID that you want to assign to a BigQuery \
+ dataset to store prediction results, for example, `circularnet_dataset`.
+ - ``: The ID that you want to assign to a BigQuery \
+ table to store prediction results, for example, `circularnet_table`. \
+ If the table already exists in your Google Cloud project, the pipeline \
+ appends results to that table.
+ - `` : If set to True, overwrites the pre-existing \
+ BigQuery table.
+ - ``: If set to True, visualizes the \
+ tracking results from the tracking algorithm.
+ - ``: If set to True, crops the objects per \
+ category according to the prediction and tracking results.
+
+ **Note:** If your input files are not videos but images, replace
+ `run_gcp_videos.sh` on the command with `run_gcp_images.sh` and remove the
+ `--fps=` parameter.
+
+ Save changes and exit the Vim editor. To do this, press the **Esc** key,
+ type `:wq`, and then press **Enter**.
+
+13. Enter the `screen` session for the client:
+
+ ```
+ screen -R inference
+ ```
+14. Run the prediction pipeline:
+
+ ```
+ bash run_images.sh
+ ```
+
+ **Note:** If you have a large amount of input files, you can run the
+ pipeline in a `screen` session in the background without worrying about the
+ terminal closing down. First, you launch the `screen` session with the
+ `screen -R client` command. A new session shell launches. Then, run the
+ `bash run_images.sh` script in the new shell.
+
+The script also creates a `logs` folder inside the `client` folder that saves
+the logs with the troubleshooting results and records from the models.
+
+You have finished running the prediction pipeline and applying the prediction models to your files for further analysis. You can find the image results with the applied masks in your output bucket. You can also open the generated [BigQuery](https://cloud.google.com/bigquery) table to see the model analytics results and [preview table data](https://cloud.google.com/bigquery/docs/quickstarts/load-data-console#preview_table_data). To create new results, repeat the steps in this section every time you modify the files in your input bucket.
+
+**Important:** If you rerun the prediction pipeline on the same video or image file, you must delete the results created the first time you ran the script from the output bucket to avoid conflicting issues. [Manage the lifecycle of objects in your Cloud Storage buckets](https://cloud.google.com/storage/docs/lifecycle) to help manage costs.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-edge.md b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-edge.md
new file mode 100644
index 00000000000..cc3b4f310a9
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-edge.md
@@ -0,0 +1,121 @@
+The script that runs [the prediction pipeline](/official/projects/waste_identification_ml/circularnet-docs/content/learn-about-pipeline) on an
+edge device applies the prediction models to analyze images.
+
+After [setting up a server](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server) in your edge device, you
+can start recording videos of objects passing on your conveyor belt to gather
+data for analysis and storing files locally from [the camera](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/). The next step is transferring those
+videos or image files to a folder in the edge device, where the prediction
+pipeline processes the images.
+
+The results of each video or image prediction are also stored locally in an
+output folder in the edge device, ensuring efficient data management. After
+systematically processing each file in your local directory, the pipeline
+creates an output directory with results for further use and analysis.
+
+This page explains how to run the prediction pipeline to apply the models to
+images stored locally in the edge device. You can then manage data according to
+your needs, such as exporting the results to BigQuery tables and connecting a
+visualization dashboard to [display results as charts and reports](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/).
+
+{{< table_of_contents >}}
+
+---
+
+## Run the prediction pipeline
+
+Follow these steps to run the prediction pipeline and process your input files
+in an edge device:
+
+1. Open the terminal of your edge device to interact with the operating system
+ through the command line.
+1. [Start the server](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server).
+1. Display the names of the models you loaded to the Triton inference server:
+
+ ```
+ cat triton_server.sh
+ ```
+
+ The first lines of the output show the names of the loaded models in square
+ brackets. You can call any of these models when running the prediction
+ pipeline.
+
+ **Important:** Run the previous command in the `server` folder, which
+ contains the `triton_server.sh` script.
+
+1. Exit the `server` folder and open the `client` folder in the
+ `prediction_pipeline` directory:
+
+ ```
+ cd ..
+ cd client/
+ ```
+
+ This folder contains the `pipeline_images.py` Python file that stores the
+ complete prediction pipeline for input images. The `run_edge_images.sh`
+ script runs this Python file automatically.
+
+1. If you have to modify the script to provide your specific paths and values
+ for the prediction pipeline, edit the corresponding parameter values on the
+ script. The following example modifies the image pipeline script:
+
+ ```
+ vim run_edge_images.sh
+ ```
+
+ The Vim editor displays the following parameters:
+
+ ```
+ --input_directory=
+ --output_directory=
+ --height=
+ --width=
+ --model=
+ --score=
+ --search_range=
+ --memory=
+ ```
+
+ Replace the following:
+
+ - ``: The path to the local folder for input images
+ in the edge device, for example `/home/images/input_files/`.
+ - ``: The path to the local folder for output image
+ results in the edge device, for example `/home/images/output_files/`.
+ - ``: The height in pixels of the image or video frames that the
+ model expects for prediction, for example, `512`.
+ - ``: The width in pixels of the image or video frames that the
+ model expects for prediction, for example, `1024`.
+ - ``: The name of the CircularNet model in the Triton
+ inference server that you want to call, for example,
+ `Jan2025_ver2_merged_1024_1024`.
+ - ``: The threshold for model prediction, for example, `0.70`.
+ - ``: The pixels up to which you want to track an object for
+ object tracking in consecutive frames, for example, `100`.
+ - ``: The frames up to which you want to track an object, for
+ example, `20`.
+
+ Save changes and exit the Vim editor. To do this, press the **Esc** key,
+ type `:wq`, and then press **Enter**.
+
+1. Run the prediction pipeline:
+
+ ```
+ bash run_edge_images.sh
+ ```
+
+The script also creates a `logs` folder inside the `client` folder that saves
+the logs with the troubleshooting results and records from the models.
+
+You have finished running the prediction pipeline and applying the prediction
+models to your files for further analysis. You can find the image results with
+the applied masks in your output folder in the edge device. You can also export
+your results manually to a [BigQuery](https://cloud.google.com/bigquery) table
+to connect it with a visualization dashboard for [data analysis and reporting](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/).
+
+**Important:** If you rerun the prediction pipeline on the same file, you must
+delete the results created the first time you ran the script from the output
+folder to avoid conflicting issues.
+
+## What's next
+
+- [View data analysis and reporting](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/)
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/_index.md
new file mode 100644
index 00000000000..bbd70cd15c7
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/_index.md
@@ -0,0 +1,12 @@
+You can deploy CircularNet using any cloud provider or recommended edge device.
+However, the following instructions describe the steps to create and run a
+Triton inference server deployed on Google Cloud or an NVIDIA device. The server
+is designed to efficiently process an image or video as an input, splitting it
+into frames and performing model predictions on every frame.
+
+The instructions for Google Cloud are for an [NVIDIA T4 GPU](https://www.nvidia.com/en-us/data-center/tesla-t4/)
+computing unit, and the instructions for the edge device are for an
+[NVIDIA Jetson](https://www.nvidia.com/en-us/autonomous-machines/embedded-systems/jetson-orin/)
+solution. Deploying CircularNet models requires technical expertise to manage
+infrastructure, run commands, and establish connection settings and permissions
+on the cloud or your edge device.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/before-you-begin.md b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/before-you-begin.md
new file mode 100644
index 00000000000..75a8b508e2b
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/before-you-begin.md
@@ -0,0 +1,79 @@
+Before deploying CircularNet, you must follow these steps. Depending on whether
+you want to deploy CircularNet on Google Cloud or an NVIDIA edge device, choose
+one of the following options:
OS and storage: Click
+ Change and select the following:
+
+
Operating system: Deep Learning on Linux
+
Version: Deep Learning VM with
+ CUDA 11.3 preinstalled. Debian 11, Python 3.10. You can choose
+ any M number with this configuration, for example, M126.
+
Boot disk type: Balanced persistent disk
+
Size (GB): 300 GB
+
+
+
Security: Navigate to the Identity and API access
+ section and select the following:
+
+
Service accounts:Compute Engine default service account
+
Access scopes: Allow full access
+ to all Cloud APIs
+
+
+
Networking: Navigate to the Firewall
+ section and select the following:
+
+
Allow HTTP traffic
+
Allow HTTPS traffic
+
+
+
+
Note: Give your VM a name that is easy to remember and deploy in a region and a zone close to your physical location that allows GPUs.
+
+
From the Navigation menu on the Google Cloud console, select Compute Engine > VM instances.
+
On the VM instances page, find the VM instance you created with the NVIDIA T4 GPU.
+
Click SSH in the row of the instance that you want to connect to. Let the SSH-in-browser tool open. For more information, see Connect to VMs.
+
Important: If the SSH connection fails, you must create a firewall rule for the VM to allow TCP ingress traffic. Use Identity-Aware Proxy (IAP) for TCP forwarding to allow ingress traffic from the IPv4 ranges 35.235.240.0/20, 0.0.0.0/0, and 0.0.0.0/22 on TCP ports 22,3389.
+
+
+
+
+
+
+ Edge device
+
+
+
Get an edge device, configure it, and connect it to your local machine.
+
Ensure you have an internet connection for the deployment and package installation.
+
Set up your device, install the software and dependencies you want, and prepare the corresponding developer kit. For information on essential configurations, refer to the software documentation of your edge device. For example, see the Developer Guide of NVIDIA Jetson devices.
+
Open the terminal of your edge device to interact with the operating system through the command line.
+
+
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/clone-repo.md b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/clone-repo.md
new file mode 100644
index 00000000000..25ac8be9f95
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/clone-repo.md
@@ -0,0 +1,39 @@
+After meeting the [prerequisites](before-you-begin.md),
+follow these steps to clone the project from the [GitHub repository](https://github.com/tensorflow/models/tree/master/official/projects/waste_identification_ml) and install all the required packages.
+
+Run the following commands on the **SSH-in-browser** window of your VM instance
+in Google Cloud or the terminal of your edge device:
+
+1. Install Git:
+
+ ```
+ sudo apt-get install git
+ ```
+
+1. Clone the tensorflow models, which contains Circularnet
+(waste_identification_ml) [Repo link](https://github.com/tensorflow/models)
+
+ ```
+ git clone --depth 1 https://github.com/tensorflow/models.git
+ ```
+
+1. Open the `client` folder within the `waste_identification_ml` project
+directory:
+
+ ```
+ cd models/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/
+ ```
+
+1. Run `requirements.sh` to install all the required packages and libraries:
+
+ ```
+ sh requirements.sh
+ ```
+
+1. Return to the root directory:
+
+ ```
+ cd\
+ ```
+
+Next, start the triton inference server that will serve inference requests. [Start server](start-server.md)
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-client.md b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-client.md
new file mode 100644
index 00000000000..10c2ba5ac68
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-client.md
@@ -0,0 +1,89 @@
+# Run the inference pipeline
+
+Follow these steps to leverage the triton inference server you
+[started earlier](start-server.md) to run the inference pipeline on your images.
+
+Go to the `client` folder within the the `waste_identification_ml` project:
+
+```bash
+cd models/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/
+```
+
+The inference pipeline uses a script to set various parameters for inference,
+post-processing, and subsequent data analytics. We'll need to adjust some of
+these:
+
+```bash
+vim run_images.sh
+```
+
+The Vim editor displays all pipeline parameters. The script contains
+documentation for each parameter, but at minimum you will need to change:
+
+```python
+--input_directory=
+
+# This should be the path to the input bucket you created containing your
+# images, e.g. gs://bucket/input-images
+```
+
+```python
+--output_directory=
+
+# Like input, this should be the gcs bucket path to where you want the images
+# with predictions to write to.
+```
+
+```python
+--project_id=
+
+# The ID of your Google Cloud project housing your gcs bucket, for example,
+# `my-gcp-project`.
+```
+
+```python
+--bq_table_id=
+
+'''
+The ID that you want to use for your BigQuery table storing inference
+results, e.g `circularnet_table`. If the table already exists BigQuery
+within your project, the pipeline will either overwrite or append results,
+depending on how you set the `overwrite` parameter
+'''
+```
+
+Save changes and exit the Vim editor. To do this, press the **Esc** key, then
+type `:wq`, and press **Enter**.
+
+**Note:** For creating cloud storage bucket and adding images, follow
+[this guide](https://github.com/tensorflow/models/blob/master/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-cloud.md#Create-the-Cloud-Storage-input-and-output-buckets)
+
+Next, enter a `screen` session for the inference client:
+
+```bash
+screen -R inference
+```
+
+Now run the inference pipeline:
+
+```bash
+bash run_images.sh
+```
+
+If you want to exit the screen session without stopping inference, press
+**Ctrl + a** then **d** to detach from the screen session.
+
+The script also creates a `logs` folder inside the `client` folder that saves
+the logs with the troubleshooting results and records from the models.
+
+Congratulations, you have finished running the inference pipeline!. You can find
+individual inference results as images with overlaid object predictions in your
+output bucket. You can also open the generated BigQuery table to see overall
+analytics across all your images. You'll want to navigate to BigQuery within
+your cloud project, find the appropriate table, and
+[preview table data](https://cloud.google.com/bigquery/docs/quickstarts/load-data-console#preview_table_data).
+
+**Important:** If you rerun the prediction pipeline on the same images, you
+should delete any existing image results (output-bucket) created previously.
+Also, see
+[Manage the lifecycle of objects in your Cloud Storage buckets](https://cloud.google.com/storage/docs/lifecycle) to help manage image storage costs.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server.md b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server.md
new file mode 100644
index 00000000000..72e29719f7c
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/start-server.md
@@ -0,0 +1,54 @@
+Follow these steps to start a Triton inference server and configure it to serve
+inference requests to the circularnet model.
+
+**Note:** If you are in Google Cloud but have closed your **SSH-in-browser**
+window, open the **VM instances** page and click **SSH** in the row of the
+NVIDIA T4 GPU instance you want to connect to. Then, let the **SSH-in-browser**
+tool open. For more information, see [Connect to VMs](https://cloud.google.com/compute/docs/connect/standard-ssh#connect_to_vms).
+
+1. On the **SSH-in-browser** window of your VM instance in Google Cloud or the
+terminal of your edge device, open the `server` folder in the
+`waste_identification_ml` project directory:
+
+ ```
+ cd models/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/server/
+ ```
+
+1. Run the `triton_server.sh` script to create the Triton inference server and
+load the most recent circularnet model on the server:
+
+ ```
+ bash triton_inference_server.sh
+ ```
+
+The server keeps running in the background (using a screen session).
+
+You can confirm the server is running by opening the screen session the server
+is running within:
+
+1. List the `screen` sessions:
+
+ ```
+ screen -ls
+ ```
+
+ The output from your server shows the `(Detached)` message because you are
+ outside of the session.
+
+1. Enter the `screen` session for the server:
+
+ ```
+ screen -r server
+ ```
+
+ The `screen` session opens and displays the ongoing operations on the
+ server. The model will show a `READY` status when it is successfully
+ deployed.
+
+1. If you want to exit the screen session without stopping the server, press
+ **Ctrl + a** then **d** keyboard shortcut. This will detach the session.
+
+## What's next
+
+Next, you are ready to send your images for inference. See the
+[start client](start-client.md) section.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/_index.md
new file mode 100644
index 00000000000..de8bec328f5
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/_index.md
@@ -0,0 +1,14 @@
+## [Benefits of CircularNet](./benefits-of-cn)
+
+Learn about the benefits of using CircularNet as your automated tool for waste
+identification.
+
+## [When to use CircularNet](./when-to-use-cn)
+
+Learn about the scenarios in which CircularNet implements useful
+functionalities.
+
+## [How CircularNet works](./how-cn-works)
+
+Learn about the arrangement of components and the operation of CircularNet
+models to identify objects.
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/benefits-of-cn.md b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/benefits-of-cn.md
new file mode 100644
index 00000000000..81907b12c95
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/benefits-of-cn.md
@@ -0,0 +1,85 @@
+The following table describes some of the benefits of using CircularNet as your
+automated tool for waste identification:
+
+
+
+
+
CircularNet benefits
+
+
+
+
+
Cost optimization
+
+
No upfront costs: Access free, open-source models and code, eliminating the need for expensive software licenses or proprietary solutions.
+
+
+
Reduced labor expenses: Minimize the need for manual material identification and sorting, freeing up personnel for other critical tasks and reducing labor costs.
+
+
+
Pay-as-you-go flexibility: Leverage cloud-based deployment options to pay only for the computing resources you use, scaling up or down as needed without investing in dedicated hardware.
+
+
+
Accelerated implementation: Jumpstart your AI-powered waste analysis initiatives with pre-trained models and avoid the time and expense of collecting and annotating large datasets.
+
+
+
Adaptable to your needs: Fine-tune the pre-trained models on your own data to further improve accuracy and tailor the solution to your specific waste stream characteristics.
+
+
+
+
+
Automation
+
+
Uninterrupted waste analysis: Capture and analyze images of waste streams continuously during operating hours, enabling consistent monitoring and identification of materials without manual intervention.
+
+
+
Visualized insights: Access a simple and user-friendly dashboard to visualize and analyze data in real-time or in aggregate, empowering decision-making and proactive process adjustments.
+
+
+
Streamlined operations: Reduce the need for manual labor in material identification and sorting, leading to increased operational efficiency and lower labor costs.
+
+
+
Minimized errors: Eliminate human errors associated with manual sorting, resulting in more accurate waste characterization and improved recycling outcomes.
+
+
+
Automated reporting: Generate detailed, customizable reports on material composition and contamination levels, saving time and effort compared to manual reporting methods.
+
+
+
+
+
Residual analysis
+
+
Realize lost value: Reduce and quantify recyclables sent for disposal through precise material identification.
+
+
+
Optimize bale quality: Accurately assess the composition and purity of materials on the baling line to ensure they meet market demands and command higher prices.
+
+
+
Pinpoint contamination sources: Identify and track contaminants within the waste stream to identify their sources and take corrective action to improve sorting processes and reduce contamination.
+
+
+
+
+
+
+## Use cases for CircularNet
+
+While you can deploy the fundamental technology underpinning CircularNet in
+various areas across the waste management ecosystem, CircularNet has been
+trained on and designed for deployment in material recovery facilities (MRFs)
+and recycling facilities. Typical users include MRF and recycling operators,
+technology integrators, system and machine learning engineers, and academic
+researchers. Examples of solutions you can implement through CircularNet
+include:
+
+- Bale quality analysis
+- Residual line analytics
+- Identification of recyclable materials and contaminants in mixed waste
+ streams
+- Recycling and contamination reporting in MRF facilities
+
+For case studies by industry and application, refer to the following external
+articles:
+
+- [CircularNet: How Recykal built Asia's largest circular economy marketplace using Google AI](https://sustainability.google/operating-sustainably/stories/circular-economy-marketplace/)
+- [CircularNet: Reducing waste with Machine Learning](https://blog.tensorflow.org/2022/10/circularnet-reducing-waste-with-machine.html)
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/how-cn-works.md b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/how-cn-works.md
new file mode 100644
index 00000000000..db6dfe16c57
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/how-cn-works.md
@@ -0,0 +1,90 @@
+CircularNet is a waste analysis system that combines computer vision and machine
+learning to optimize recycling processes. To implement CircularNet, you need
+these key components:
+
+- **Machine vision camera:** A specialized camera captures high-resolution
+ images of materials on the conveyor belt, ensuring consistent quality even
+ with fast-moving waste. While CircularNet can process images from various
+ sources, such as cell phones or GoPro cameras, a machine vision camera with a
+ global shutter is recommended for optimal results with moving objects.
+- **Computational unit:** The image analysis component contains the trained ML
+ models that perform the analysis. The computational unit can be one of the
+ following:
+
+ - An **edge device** for on-site, real-time processing.
+ - A **cloud-based solution** for scalable, remote processing.
+ - **Your own server** for complete control over the infrastructure.
+
+- **Machine Learning (ML) models:** These pre-trained models are the heart of
+ CircularNet. They use a mask-based image analysis algorithm to identify and
+ classify materials and their forms within the waste stream.
+- **Reporting dashboard:** The raw images and analysis results are compiled and
+ presented through a visual dashboard. This dashboard provides detailed
+ breakdowns of materials, classifications, and other insights, empowering you
+ to make data-driven decisions to improve sorting and recycling processes.
+
+You can stream results to a dashboard on various cloud platforms or your own
+server. This guide provides instructions specifically for deploying a dashboard
+on Google Cloud.
+
+## Training models
+
+CircularNet's machine learning models utilize a vision-based masking algorithm
+to identify and classify diverse materials within waste streams. These models
+have been pre-trained using images from various Material Recovery Facilities
+(MRFs), representing multiple waste compositions.
+
+While the pre-trained models are ready for use, you can refine their performance
+by continuing the training process. Retraining can help you improve performance
+on specific material types or adapt to unique waste stream characteristics in
+your geography or deployment area. Retraining involves gathering and annotating
+additional images from three primary sources:
+
+- **MRF images**: Images captured and labeled directly from conveyor belts at
+ MRFs.
+- **Google's internal annotation**: Images annotated by Google's dedicated
+ team.
+- **Automatic annotation**: A processing script automatically annotates images
+ of sorted materials, focusing on common types like PET, aluminum cans, and
+ cardboard boxes.
+
+The automatic annotation process involves the following tasks:
+
+- Generating masks for objects in the images, using either bounding boxes or
+ more precise mask detection depending on the image complexity.
+- Filtering and cleaning the generated masks to remove duplicates or errors.
+- Adding annotations to create the final training dataset.
+
+
+
+**Figure 1.** A model applies mask detection to an image containing aluminum
+cans passing on a conveyor belt.
+
+After preparing the data, the ML models are taught and trained with it to learn
+how to identify materials in new real-world scenarios.
+
+## Applying the models to your data
+
+Once your machine vision camera has captured images or videos of your waste
+stream, you can put CircularNet's models to work. Upload these captured files to
+a server where CircularNet analyzes each image or video frame. The models
+leverage their learned knowledge to identify and classify the various materials
+and their forms present in the waste stream. Models can analyze data in
+real-time or in batches, depending on your needs and chosen deployment
+constraints. The analysis results, including details about each detected object,
+such as material type and form, are then systematically organized and stored in
+a [BigQuery](https://cloud.google.com/bigquery) database table. This structured
+data fuels the [Looker](https://cloud.google.com/looker) dashboard, where you
+can access and visualize it through comprehensive reports and insightful
+visualizations. The reporting dashboard serves as your gateway to insights
+generated by CircularNet, allowing you to create customized reports and
+empowering you to make data-driven decisions that enhance the efficiency and
+effectiveness of your sorting and recycling operations.
+
+## What's next
+
+- [Learn more about the deployment solutions that CircularNet offers](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/).
+- [Set up a hardware solution for your needs](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/).
+- [Learn how to deploy CircularNet on Google Cloud or an edge device](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/).
+- [Learn how to prepare your data and use the models to analyze images](/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/).
+- [Learn more about the visual dashboard to get analytics and reports](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/).
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/when-to-use-cn.md b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/when-to-use-cn.md
new file mode 100644
index 00000000000..50013f95a4c
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/discover-cn/when-to-use-cn.md
@@ -0,0 +1,37 @@
+Consider CircularNet if you want to automate the analysis of waste composition
+and material identification within your Material Recovery Facility (MRF) or
+recycling center. It is particularly valuable for scenarios where you want to
+implement the following functionalities:
+
+- **Gain aggregate or real-time insights** into the material types and forms
+ moving through your facility on conveyor belts.
+- **Reduce reliance on manual inspection** and sorting, thereby improving
+ efficiency and minimizing human error.
+- **Identify and quantify contaminants** within waste streams to improve the
+ quality of recycled materials.
+- **Generate automated and historical reports** on material composition,
+ recycling rates, and contamination levels to support data-driven
+ decision-making and operational improvements.
+
+CircularNet utilizes RGB computer vision models and pixel-level instance
+segmentation to accurately identify and classify materials, making it a valuable
+tool for enhancing the efficiency and effectiveness of waste management
+operations.
+
+CircularNet identifies material forms and types. Furthermore, in the case of
+plastic, it identifies plastic types. The model employs pixel-level instance
+segmentation, a technique that precisely outlines the shape of each object
+within an image. This technique offers several advantages, such as the
+following, compared to traditional bounding box object detection methods:
+
+- **Accurate object delineation**: Pixel-level segmentation provides a precise
+ representation of object boundaries, which is critical for accurately
+ measuring object size, shape, and quantity, especially in cluttered waste
+ streams.
+- **Improved contamination detection**: By accurately segmenting objects,
+ CircularNet can identify and quantify contaminants mixed with recyclables,
+ leading to enhanced sorting and higher-quality recycled materials.
+- **Enhanced material characterization**: Pixel-level information enables
+ nuanced material analysis, letting CircularNet distinguish between
+ similar-looking materials or identify specific material attributes, such as
+ plastic types.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/_index.md
new file mode 100644
index 00000000000..8ff46863d33
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/_index.md
@@ -0,0 +1,12 @@
+# Retrain CircularNet models
+
+CircularNet's models are initially trained with images captured from Material
+Recovery Facilities (MRFs). As a result, these open-source models are
+specialized in a limited set of materials, which might not fully align with the
+specific materials you collect.
+
+To address your specific use case, you can
+retrain the models by utilizing a pipeline built with
+[Vertex AI](https://cloud.google.com/vertex-ai/docs). This pipeline is ideal for
+analyzing materials from diverse sources, adapting to different business
+scenarios, or retraining the model with your own images.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/before-you-begin.md b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/before-you-begin.md
new file mode 100644
index 00000000000..3ab2d041009
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/before-you-begin.md
@@ -0,0 +1,22 @@
+## Before you begin
+
+Before starting the retraining process, ensure you meet the following
+requirements:
+
+1. [Get access to Google Cloud](https://console.cloud.google.com/).
+1. [Open the Google Cloud console](https://cloud.google.com/cloud-console).
+1. [Create a project on your Google Cloud account](https://cloud.google.com/resource-manager/docs/creating-managing-projects).
+1. Enable the Vertex AI and Cloud Storage APIs to manage programmatic
+ access and authentication.
+
+ To enable APIs, see
+ [Enabling an API in your Google Cloud project](https://cloud.google.com/endpoints/docs/openapi/enable-api).
+
+1. [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets)
+ to store files.
+1. Allocate at least four GPUs for the training job on Vertex AI. For more
+ information, see
+ [Configure compute resources for custom training](https://cloud.google.com/vertex-ai/docs/training/configure-compute).
+1. Set up a service account for Vertex AI, with permissions to perform
+ training jobs. For more information, see
+ [Use a custom service account](https://cloud.google.com/vertex-ai/docs/general/custom-service-account).
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/launch-job.md b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/launch-job.md
new file mode 100644
index 00000000000..7580da444de
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/launch-job.md
@@ -0,0 +1,40 @@
+## Launch the training job
+
+Run
+[the `CircularNET_Vertex_AI_ReTraining_v1.ipynb` script](https://github.com/tensorflow/models/blob/master/official/projects/waste_identification_ml/model_retraining/CircularNET_Vertex_AI_ReTraining_v1.ipynb)
+to configure parameters, launch the training job on Vertex AI, and start the
+training.
+
+**Note:** Wait to export the checkpoints to a saved TF model until the script
+finishes running. The training job might take several days to complete. Vertex
+AI manages the process in the background, so you don't need to keep your
+computer running.
+
+Follow these steps to monitor the training progress:
+
+1. Log in to your Google Cloud account.
+1. Navigate to Vertex AI.
+1. From the Vertex AI menu, click **Training**.
+1. On the **Training** page, open the **Hyperparameter tuning jobs** tab
+ and select the name of your training job.
+1. To observe data about the training job, perform one of the following
+ actions:
+ - Click **Open TensorBoard** to view detailed training metrics.
+ - Click **View Logs** to open the log console and monitor the
+ detailed log messages.
+
+---
+
+## Export and test the model
+
+After training is complete, you can export the model in various formats, such
+as TensorFlow, for future use when you deploy the model.
+[The `CircularNET_Vertex_AI_ReTraining_v1.ipynb` script](https://github.com/tensorflow/models/blob/master/official/projects/waste_identification_ml/model_retraining/CircularNET_Vertex_AI_ReTraining_v1.ipynb)
+guides you through this process.
+
+Additionally, you can test the model locally and run inferences on images to
+validate its performance. The script also contains information about running
+local tests.
+
+Finally, deploy the exported model using an edge device or Google Cloud,
+depending on your deployment requirements.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/prepare-data.md b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/prepare-data.md
new file mode 100644
index 00000000000..29b7117467c
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/retrain-models/prepare-data.md
@@ -0,0 +1,68 @@
+## Prepare your training dataset
+
+Begin by capturing images with your camera and performing the necessary
+preprocessing steps. Then, annotate the captured images to identify the
+materials present in each one. These annotations allow the model to learn which
+materials it must recognize during training. You must save annotations in COCO
+JSON format.
+
+The GitHub repository provides the required
+[preprocessing scripts](https://github.com/tensorflow/models/tree/master/official/projects/waste_identification_ml/pre_processing)
+to prepare the model for retraining. These scripts convert your annotated images
+into [TFRecords](https://www.tensorflow.org/tutorials/load_data/tfrecord), the
+required input format for TensorFlow models.
+
+The preprocessing scripts of the repository let you perform the following
+actions:
+
+- Convert image annotations to the COCO-annotated JSON file format, which
+ is necessary for labels and metadata on a dataset.
+- Clean and prepare a COCO-annotated JSON file for training an ML model,
+ ensuring that your dataset is clean, consistent, and ready for effective
+ model training.
+- Filter out irrelevant or noisy annotations that could negatively impact
+ the training process.
+- Adjust annotations to ensure they are in the optimal format for model
+ training.
+- Verify that all annotated images exist and are not corrupted.
+- Merge COCO-annotated JSON files into a single file and convert it into
+ TFRecord files, the required format for training with the Mask R-CNN model.
+
+Before proceeding with the retraining pipeline, run the preprocessing scripts
+with your data on your local workstation, remote server, or database. Once your
+TFRecords are ready,
+[upload them to a Cloud Storage bucket](https://cloud.google.com/storage/docs/uploading-objects).
+The bucket can have two locations, one for the training dataset and another for
+the validation dataset.
+
+---
+
+## Configure the training job
+
+Configure your model training by customizing values in
+[the `CircularNET_Vertex_AI_ReTraining_v1.ipynb` script](https://github.com/tensorflow/models/blob/master/official/projects/waste_identification_ml/model_retraining/CircularNET_Vertex_AI_ReTraining_v1.ipynb),
+including your project ID, bucket URI, and region. Provide the following
+information in the corresponding script variables:
+
+- `input_train_data_path`: path to the TFRecords of your training dataset
+ in the Cloud Storage bucket.
+- `input_validation_data_path`: path to the TFRecords of your validation
+ dataset in the Cloud Storage bucket.
+- `init_checkpoint_path`: path to the initial checkpoints with the model's
+ weights. You can use the open-source initial checkpoints in this variable
+ from
+ [the configuration file](https://github.com/tensorflow/models/blob/master/official/projects/waste_identification_ml/model_retraining/config/config_v1.yaml).
+- `config_file_path`: path to the configuration file containing parameters
+ for fine-tuning the training. You can use the open-source
+ [configuration file](https://github.com/tensorflow/models/blob/master/official/projects/waste_identification_ml/model_retraining/config/config_v1.yaml)
+ in this variable.
+- `service_account`: name of the Vertex AI service account with
+ permissions to perform training jobs.
+
+The script's placeholders, such as `PROJECT_ID`, `REGION`, and
+`STAGING_BUCKET`, must be replaced with your Google Cloud project, region, and
+Cloud Storage bucket, respectively.
+
+**Note:** You can also customize other script variables, such as `num_classes`,
+which refers to the number of categories for materials or other classes in your
+annotated images.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md
new file mode 100644
index 00000000000..4048e2b954d
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/solutions/_index.md
@@ -0,0 +1,172 @@
+CircularNet's open-source nature allows for flexible deployment across various
+infrastructures, enabling you to tailor the system to your waste management
+needs and resources.
+
+Deployment solutions depend on several factors, including the following:
+
+- **Technical expertise**: The technical skills available within your team to
+ manage the deployment and ongoing maintenance.
+- **Existing infrastructure**: The hardware and software components you already
+ have in place.
+- **Data availability needs**: How quickly you need access to the analysis
+ results. Options include real-time and batch analysis.
+- **Budget and resources**: The financial and computational resources you can
+ allocate to CircularNet.
+
+This page guide you through different deployment options, helping you choose the
+one that best aligns with your requirements and capabilities.
+
+{{< table_of_contents >}}
+
+---
+
+## Where to host the models
+
+While CircularNet's models are freely downloadable from [GitHub](https://github.com/tensorflow/models/tree/master/official/projects/waste_identification_ml), they require a computational environment to perform their analysis. You have the following three main options for where to host and run these models:
+
+- **Cloud account**: A cloud provider manages the infrastructure for you so
+ that you can benefit from the advantages of the cloud, such as paying only
+ for how much you consume instead of maintaining your own data center. For
+ example, you create a self-managed server hosted on the cloud infrastructure
+ to deploy the ML models. The provider handles software updates, the
+ underlying network, and data encryption. This guide contains a sample
+ solution in [Google Cloud](https://cloud.google.com/gcp), but you can choose
+ any other cloud provider to host ML models.
+- **Edge device**: A self-contained physical device equipped with all the
+ necessary tools to run the models and perform predictions locally. It offers
+ a turnkey solution, handling the entire analysis process on-site. While the
+ device can connect to the cloud for dashboard visualization, it doesn't
+ require constant internet connectivity.
+- **Your server**: You have complete control over the hardware and software
+ environment, allowing full customization. This option is suitable if you have
+ the technical expertise and resources to manage your server infrastructure.
+
+The best choice depends on your priorities. An edge device is ideal for
+real-time analysis and situations where internet connectivity is limited. A
+cloud solution offers scalability and flexibility, while running the models on
+your server provides maximum control but requires more technical management.
+
+The following sections describe the advantages and considerations of each
+option:
+
+- [**Cloud deployment**](#cloud-deployment)
+- [**Edge device deployment**](#edge-device-deployment)
+- [**In-house server deployment**](#in-house-server-deployment)
+
+### Cloud deployment
+
+If you prefer a fully managed solution and value flexibility, deploying CircularNet on a cloud platform like [Google Cloud](https://cloud.google.com/) could be ideal. You benefit from the cloud's infrastructure while retaining control over resource configuration and model customization.
+
+This approach is well-suited for those who want the following functionalities:
+
+- **A hosted solution:** The entire system, including the [dashboard for reporting and visualization](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/), resides in the cloud, eliminating the need for on-site hardware management.
+- **Flexibility and customization:** You can modify models, install additional software or libraries, and experiment with different frameworks as your needs evolve.
+
+Using a cloud account requires technical expertise to manage infrastructure,
+commands, and connection settings.
+
+The following are the benefits of cloud deployment:
+
+- **Customization and experimentation:** Install additional software or
+ libraries and try different frameworks to tailor the models to your needs.
+- **Offline data access**: Store and access your captured images privately and
+ securely in [Google Cloud Storage](https://cloud.google.com/storage).
+- **Scalability**: A cloud infrastructure can scale compute resources with
+ graphics processing units (GPUs) up or down based on your workload. If you
+ need more power for complex tasks, you can scale up. On the other hand, if
+ you have lower traffic, you can scale down to save costs.
+- **Lower upfront costs**: Avoid the initial investment in hardware and
+ deployment setup required for edge devices or on-premise servers. For pricing
+ information on Google Cloud, see [Google Cloud pricing](https://cloud.google.com/pricing).
+- **Managed infrastructure**: Google Cloud handles patching, security updates,
+ and infrastructure management, freeing you to focus on your model.
+- **High availability**: To minimize downtime for CircularNet's image
+ recognition service, you can configure your infrastructure for redundancy and
+ automatic failover.
+
+Overall, Google Cloud offers a scalable, flexible, and cost-effective solution
+for running the ML model with offline data. It balances processing power,
+manageability, and cost efficiency.
+
+To learn how to install CircularNet on Google Cloud, see [Deploy CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/).
+
+### Edge device deployment
+
+An edge device brings computing power directly to your facility for on-site data
+processing. It requires local installation and setup, providing a comprehensive
+solution for hosting and running CircularNet's ML models with a self-managed
+experience.
+
+**In most cases, we recommend running the ML model on an [NVIDIA](https://www.nvidia.com/en-us/edge-computing/) edge device.** For information about the minimum hardware requirements for the device to run CircularNet, [choose edge device hardware](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-edge-device/).
+
+The configuration of an edge device has benefits such as the following:
+
+- **Real-time inference**: Edge devices can perform real-time inference,
+ allowing you to process data as it is generated without significant delay.
+ This attribute enables the image identification feature to respond quickly to
+ user inputs or changes in the environment.
+- **Limited connectivity:** Edge devices can operate in environments with
+ limited, unreliable, or non-existent internet connection.
+- **High-performance inference**: Edge devices deliver rapid data processing
+ and analysis for low-latency applications.
+- **Energy efficiency**: Edge devices are designed to be power-efficient,
+ allowing them to perform complex computations while consuming minimal power.
+ This fact makes them ideal for use in battery-powered or energy-constrained
+ devices, such as some [machine vision cameras](/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/)
+ that you can use to capture images.
+- **Compact form**: Edge devices are small and easily integrated into various
+ setups, eliminating the need for bulky external hardware.
+- **Compatibility**: Edge devices support multiple frameworks and models. This
+ compatibility ensures you can leverage existing models and tools without
+ extensive modifications or retraining.
+
+Because of its on-site processing capabilities, an edge device might be
+preferable for applications that demand split-second analysis. Due to reduced
+network latency, processing on an edge device is fast for real-time
+applications.
+
+While edge devices offer numerous advantages, they have less processing power
+than larger cloud-based systems and may not be suitable for high-volume or
+complex tasks. Additionally, a device failure could temporarily disrupt
+operations. Choosing and setting up an edge device requires technical expertise
+to manage the hardware and software.
+
+To learn how to install CircularNet on an edge device, see [Deploy CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/).
+
+### In-house server deployment
+
+If you have a server, you manage every configuration setting to host the models.
+You have complete control over the infrastructure, model customization, and data
+management. This alternative is ideal if you count on developers and staff with
+a high level of technical expertise and you already have a data center intended
+for analysis and waste identification.
+
+Using an in-house server to host the model provides benefits such as the
+following:
+
+- **Full control and customization**: Install libraries, frameworks, or custom
+ tools to modify the ML model, optimize the server, and gain complete control
+ over the hardware and software environment.
+- **Data privacy**: Keep the model and offline data entirely on-premise.
+- **Cost-effective in the long run (potentially)**: After the initial
+ investment in hardware, pay no recurring monthly fees. Depending on use case
+ demands, a server can be cost-effective, especially for long-term use with
+ consistent workloads.
+
+However, using your server also comes with some drawbacks:
+
+- **Scalability challenges**: Scaling up resources requires buying additional
+ hardware, which can be expensive and time-consuming. Scaling down is also not
+ as flexible.
+- **Management overhead**: The user is responsible for maintaining the server
+ hardware, updating software, applying security patches, and troubleshooting
+ hardware issues.
+- **Limited redundancy**: A single server failure can lead to downtime for the
+ image recognition service.
+
+This guide doesn't cover the detailed steps for deploying CircularNet on your
+server. The deployment process is highly dependent on your specific on-premise
+infrastructure, data center setup, customized settings, available resources, and
+the unique requirements of your use case.
+
+You must consult resources and documentation tailored to your specific server environment and adapt [the CircularNet models](https://github.com/tensorflow/models/tree/master/official/projects/waste_identification_ml) to deploy and integrate them successfully into your existing systems. However, if you choose this alternative, we recommend providing a computational unit to run AI models, perform inference, and connect to a [dashboard for data analysis and reporting](/official/projects/waste_identification_ml/circularnet-docs/content/view-data/).
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/system-req/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/_index.md
new file mode 100644
index 00000000000..17df5790d3e
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/_index.md
@@ -0,0 +1,12 @@
+This section contains guidelines for the system resources you need to get, set
+up, and install in your facility to capture images, identify materials, and
+generate analytics results with CircularNet.
+
+## [Choose a camera](./choose-camera)
+
+Learn about the recommended cameras and their installation to capture
+high-quality images from your materials.
+
+## [Choose edge device hardware](./choose-edge-device)
+
+If you [chose an edge device for your deployment solution](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/#edge-device-deployment), learn about the recommended hardware to deploy CircularNet models.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/_index.md
new file mode 100644
index 00000000000..0071bb62713
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/_index.md
@@ -0,0 +1,22 @@
+The quality of images captured by your camera directly impacts the accuracy of
+CircularNet's analysis. Therefore, selecting the right machine vision camera is
+crucial for successful waste identification and characterization.
+
+The following are the primary key points when choosing a camera:
+
+- **Machine vision camera:** A specialized camera with a global shutter is
+ recommended for optimal image quality, especially when capturing fast-moving
+ objects on the conveyor belt. This type of camera minimizes motion blur and
+ distortion.
+- **Image resolution:** CircularNet analyzes individual frames, so
+ high-resolution images are essential.
+- **Installation:** Position the camera directly above the conveyor belt for
+ consistent lighting and full coverage.
+
+As a general rule, avoid capturing images with motion blur, fisheye effect, and
+quality issues due to the vibrations from the conveyor belt movement. Global
+shutter cameras typically meet these recommendations.
+
+The CircularNet model supports inference on individual image frames. Therefore,
+if the camera captures videos, the system converts those videos into frames for
+pre- and post-image processing.
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/camera-recommendations.md b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/camera-recommendations.md
new file mode 100644
index 00000000000..eb17fc75c92
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/camera-recommendations.md
@@ -0,0 +1,265 @@
+The following list contains the essential recommendations when selecting a
+camera to capture high-quality images:
+
+- **Frame rate:** Ensure the camera can capture images at a high frame rate,
+ matching the speed of the conveyor belt to avoid motion blurriness. Look for
+ cameras with at least 30 FPS (frames per second) or higher.
+- **Shutter type:** A global shutter is preferable over a rolling shutter to
+ avoid distortion, especially for fast-moving objects.
+- **Lens compatibility:** Choose a camera with interchangeable lenses to adjust
+ the field of view and focus based on the conveyor belt width and object size.
+- **Interface:** Select a camera with an appropriate interface (for example,
+ USB 3.0, GigE, or Camera Link) that supports high-speed data transfer to the
+ processing unit.
+- **Integration with software libraries:** Verify compatibility with software
+ and libraries for seamless integration. For example, if you are using an edge
+ device, such as NVIDIA or Raspberry Pi, you need software integration
+ compatibility with it.
+- **Scan type:** An area scan camera is preferable over a line scan camera to
+ ensure flexibility and comprehensive imaging capabilities required to
+ identify varying waste materials effectively. Area scan cameras are suitable
+ when objects vary significantly in size and shape and for situations where
+ the conveyor belt speed varies or is not uniform.
+- **Resolution:** A suitable resolution is crucial for image analysis because
+ it corresponds to the level of detail, patterns, and textures a camera can
+ detect from an object. You must have an image resolution of at least
+ 1024x1024 pixels. For this reason, choose a camera with a resolution between
+ 0.5 and 1.0 MP (megapixels), depending on the size of the objects and the
+ level of detail required.
+
+The higher the pixel resolution, the more processing time and latency are
+introduced by having to process all those pixels through the model. Find a
+balance between image resolution and processing latency.
+
+The following list contains additional factors you might want to consider based
+on your facility's conditions:
+
+- **Enclosure:** Ensure the camera has an industrial-grade enclosure,
+ preferably NEMA, IP65, or higher, to withstand harsh environments and dust.
+ For example, [Basler has some offerings](https://www.baslerweb.com/en/products/accessories-and-bundles/basler-ip67-housing/).
+- **Lighting:** Ensure diffused or even lighting across all items on the
+ conveyor belt. If the conveyor belt's speed is high, you need a smaller
+ [aperture size](./factors/#aperture-size-f-number) and a higher [shutter
+ speed](./factors/#shutter-speed). The conveyor belt should be well-lit to
+ ensure bright images, reducing blurriness.
+- **Parameter control:** Opt for a camera with good low-light sensitivity and
+ high dynamic range to handle varying lighting conditions. Look for cameras
+ with adjustable exposure settings to accommodate different lighting
+ conditions and object speeds. Also, choose a camera that supports remote
+ configuration for tuning software parameters such as shutter speed, exposure,
+ and frame rate.
+- **Integrated lightning:** Accompany the camera with an integrated lightning
+ system to reach the required luminosity in your facility. You can find
+ accessories such as lightning bulbs or lamps to increase the brightness of
+ the conveyor belt.
+- **Synchronization:** If you need multiple cameras, ensure they can be
+ synchronized to capture images simultaneously.
+- **Power supply:** Consider power over ethernet (PoE).
+- **Temperature range:** Verify the camera can operate within the temperature
+ range of the recycling facility.
+- **Reliability and durability:** Choose cameras with a proven track record for
+ reliability and durability in industrial applications.
+
+### **Camera Installation and Placement**
+
+Proper installation involves positioning the camera **directly** **above** the
+conveyor belt for **consistent** **lighting** and **full** **coverage**. As a
+general rule, avoid capturing images with motion blur, fisheye effect,
+and quality issues due to vibrations from the conveyor belt movement;
+global shutter cameras typically meet these recommendations. Crucially, no
+hands or other foreign objects, **nor overflowing belt edges** should be
+visible in the camera's field of view during operation, as
+obstructions will interfere with object detection and classification.
+
+---
+
+#### **Camera Mounting Guidelines:**
+
+Camera mounting directly impacts image stability and clarity. Adhering to
+these guidelines is critical for reliable data acquisition and consistent
+model performance.
+
+* **Vertical and Centered Alignment**: The camera must be positioned directly
+ and vertically above (i.e. perpendicular) the exact center of the
+ conveyor belt. This ensures a consistent perspective, minimizes image
+ distortion from angled views, and provides uniform coverage of the
+ belt's width.
+* **Fixed and Optimized Height**: The camera must be mounted at a fixed,
+ immovable height above the conveyor belt. This height should be
+ precisely determined based on calculations for optimal focal length
+ and field of view (FoV), considering the conveyor belt's dimensions
+ (width and object height). Once established, this height must remain
+ constant to maintain consistent image scaling and object detection
+ accuracy. Typical camera heights in facilities are around 3-4 feet
+ (approximately 1-2 meters). A precise 1:1 ratio between the camera's
+ working distance (height) and the field of view (e.g., a 1-meter
+ height for a 1x1 meter FoV) is highly recommended for optimal image
+ proportionality and model performance. Smaller belt widths and lower
+ camera heights are generally preferred where feasible to achieve
+ higher pixel density per object.
+* **Robust and Vibration-Resistant Assembly**: Utilize industrial-grade
+ mounting hardware specifically designed for high-vibration
+ environments like MRFs. The mounting system must prevent any
+ perceptible movement, wobble, or vibration of the camera, which would
+ result in blurry images and severely impact image quality and model
+ performance. Regular inspections for mount integrity and tightness are
+ mandatory.
+* **Unobstructed Field of View**: The mounting hardware itself, all
+ associated cabling, and any other facility infrastructure must be
+ positioned entirely outside the camera's field of view. A clear and
+ unobstructed view of the materials on the conveyor belt is essential
+ for accurate object detection and classification.
+* **Accessibility for Maintenance (Secondary Consideration)**: While the
+ primary focus is on stability and an unobstructed view, consider
+ designing the mounting setup to allow for reasonable access for
+ routine maintenance, cleaning of the camera lens, and adjustments to
+ lighting components.
+
+---
+
+#### **Strict Lighting Guidelines:**
+
+Consistent and adequate lighting is paramount for optimal image quality and
+accurate material identification.
+
+* **Diffused and Even Illumination**: Ensure [diffused or even lighting](https://www.effilux.com/en/products/led-bar/effi-flex)
+ across all items on the conveyor belt. Diffused lighting refers to
+ light that has been spread out or softened, rather than being direct
+ and harsh. This can be achieved using [diffusers](https://www.amazon.com/Torjim-Photography-Professional-3000-7500K-Recording/dp/B0CF44WSPJ/ref=sr_1_6?dib=eyJ2IjoiMSJ9.t7O2HZqH9szXiU7jZ2GIdULPAq9kN2Wqwo5ESFO5NDPH47xTxiugxvEh1lnsvCbd38rzWAZnNci8eiJfJtdzL-FDVJT3uZAzdvVz8QcqUiZA96QcZ2YmoUxUFLrTuOYi9VL7GJM6nrc1gjbAyR6M__NuvtTgtJ8WJKVvDiMubuEfBM7OEkHWT_3tw00_bNHTvB95rotyGse14vDsH9O7KDnDJggn5fIgW09tyNDf4dc.f9-abHQLKPtm52XxLyl0oXmSzGHUqRKZMO53S76LdhQ&dib_tag=se&keywords=light%2Bdiffuser&qid=1750113487&sr=8-6&th=1)
+ (translucent materials placed between the light source and the subject)
+ or by using large, soft light sources. The goal is to minimize harsh
+ shadows and hot spots, which can obscure material features and
+ negatively impact model performance.
+* **High-Speed Belt Lighting for Motion Blur Prevention**: If the conveyor
+ belt's speed is high, a smaller aperture size and a higher shutter
+ speed are necessary to prevent motion blur. To compensate for the
+ reduced light intake at higher shutter speeds, the conveyor belt
+ must be exceptionally well-lit to ensure bright images.
+* **Integrated and Dedicated Lighting System**: Accompany the camera with a
+ dedicated integrated lighting system to achieve the required
+ luminosity in your facility. This system should provide ample
+ brightness for the entire field of view.
+* **Eliminate Lighting Fluctuations**: Minimize any lighting fluctuations
+ (e.g., flickering lights, inconsistent ambient light, direct and even
+ indirect sunlight exposure) that could introduce variability into the
+ images and hinder consistent object detection. The lighting
+ environment should be controlled and stable throughout operation.
+* **Color Temperature Consistency**: Maintain a consistent color temperature
+ for all light sources. Variations in color temperature can alter the
+ perceived color of materials, potentially impacting the model's
+ ability to accurately classify objects, especially when color is a
+ distinguishing feature.
+* **Adequate Lux/Luminosity**: The lighting system must provide sufficient
+ lux (lumens per square meter) to ensure that the camera sensor
+ receives enough light for clear image capture, even at higher shutter
+ speeds. This is crucial for optimal image quality and model
+ performance.
+
+*Diffused lighting helps capture the full spectrum of light for an object
+and minimizes shadows*
+
+---
+
+#### **Factors Based on Vision System Placement:**
+
+You must precisely measure and account for the following factors when
+installing the camera above the conveyor belt:
+
+* **Conveyor Belt Width**: Accurately determine the total operational width
+ of the conveyor belt that the camera needs to fully cover. This
+ measurement is crucial for selecting the appropriate lens focal length
+ and ensuring the camera's field of view encompasses the entire belt
+ width *without* extending beyond its edges. Typical belt widths range
+ between 1 and 1.5 meters.
+* **Camera Height (Critical for Focal Length and FoV Ratio)**: Confirm the
+ fixed vertical height above the conveyor belt at which the camera
+ will be mounted. This height directly influences the required focal
+ length. Standard camera heights in facilities are around 3-4 feet
+ (approximately 1-2 meters). A precise 1:1 ratio between the camera's
+ working distance (height) and the field of view (e.g., a 1-meter
+ height for a 1x1 meter FoV) is highly recommended for optimal image
+ proportionality and model performance. For example, if your FoV
+ captures an area of 1x1 square meters of the belt, you must mount the
+ camera one or 1.5 meters above the conveyor belt.
+* **Conveyor Belt Speed (Critical for Shutter Speed)**: Accurately measure
+ the average and maximum operational speed of the conveyor belt. This
+ speed is a primary determinant for calculating the adequate shutter
+ speed to minimize motion blur. Belt speeds typically range between 1
+ and 4 meters per second; slower speeds are always preferable for
+ achieving sharper images.
+* **Field of View (FoV) and Image Proportionality (Square Aspect Ratio)**:
+ While the conveyor belt length is continuous, the camera's FoV must
+ capture well-proportioned images. The ideal FoV should be
+ approximately square, meaning the length of the belt captured in the
+ frame should be similar to its width (e.g., a 1-meter belt width
+ should correspond to approximately 1 meter of belt length in the
+ image). This square aspect ratio assists in consistent object
+ detection and tracking.
+* **Minimize Item Overlap**: The camera's positioning, coupled with
+ optimized conveyor belt loading procedures, should actively aim to
+ minimize the overlap of items on the belt. Excessive overlap
+ significantly impedes accurate pixel-level instance segmentation and
+ reduces detection confidence.
+* **Consistent Object Size**: While not always controllable, the system is
+ optimized for objects similar in size to those typically found in
+ household recycling streams, as depicted in the provided visual
+ examples.
+
+---
+
+#### **Camera Setup Checklist for CircularNet Deployment:**
+
+**I. Camera Mounting & Placement:**
+
+- [ ] **Vertical Alignment:** Is the camera positioned directly and
+ vertically above the exact center of the conveyor belt?
+- [ ] **Fixed Height:** Is the camera mounted at a fixed, immovable
+ height determined by focal length and FoV calculations?
+- [ ] **Robust Mount:** Is industrial-grade, vibration-resistant
+ mounting hardware being used for a stable assembly?
+- [ ] **Unobstructed FoV:** Are all mounting hardware, cables, and
+ facility structures outside the camera's field of view?
+- [ ] **Belt Coverage:** Does the camera's FoV fully cover the entire
+ width of the conveyor belt without extending beyond its edges?
+- [ ] **Square FoV:** Is the FoV configured to be approximately square
+ (belt width:belt length ratio of ~1:1)?
+- [ ] **No Hands/Obstructions:** Are operational procedures in place to
+ ensure no hands, foreign objects, or overflowing belt edges are
+ visible in the camera's view?
+- [ ] **Minimize Overlap:** Are conveyor belt loading procedures
+ optimized to minimize item overlap for better detection?
+
+**II. Lighting System:**
+
+- [ ] **Diffused Illumination:** Is the lighting system designed to
+ provide diffused and even illumination across the entire conveyor
+ belt?
+- [ ] **Adequate Brightness:** Is the lighting system powerful enough to
+ ensure bright images, especially at high conveyor belt speeds and
+ faster shutter speeds?
+- [ ] **Dedicated Lighting:** Is an integrated and dedicated lighting
+ system being installed with the camera?
+- [ ] **Stable Lighting:** Are measures in place to eliminate lighting
+ fluctuations (flicker, ambient light changes)?
+- [ ] **Consistent Color Temperature:** Is the lighting system
+ maintaining a consistent color temperature?
+- [ ] **Adequate Lux:** Is the lighting system providing sufficient
+ lux/luminosity for clear image capture?
+
+**III. Environmental Measurements:**
+
+- [ ] **Conveyor Belt Width:** Has the precise operational width of the
+ conveyor belt been measured and are there no overflowing
+ portions visible in sample images?
+- [ ] **Camera Mounting Height:** Has the exact fixed height for camera
+ mounting above the belt been determined?
+- [ ] **Conveyor Belt Speed:** Has the average and maximum operational
+ speed of the conveyor belt been measured?
+
+---
+## Recommended models
+
+The following list contains some examples of recommended models for your camera:
+
+- [Arducam High Quality Camera](https://www.arducam.com/product/b0242-arducam-imx477-hq-camera/)
+- [GoPro HERO12 Black](https://gopro.com/en/us/shop/cameras/hero12-black/CHDHX-121-master.html)
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md
new file mode 100644
index 00000000000..ebdaaffa977
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/factors.md
@@ -0,0 +1,111 @@
+You must measure the following factors when installing the camera above the
+conveyor belt:
+
+- **Conveyor belt width:** Determine the total width of the conveyor belt that
+ the camera needs to cover. This value helps determine the [focal length](#focal-length).
+ The conveyor belt width is typically between one and
+ 1.5 meters, and the camera must capture the entire belt width on the image
+ frame.
+- **Camera height:** Confirm the fixed height above the conveyor belt at which
+ you will mount the camera. This value helps determine the [focal length](#focal-length).
+ The camera height is typically between one and two
+ meters above the conveyor belt.
+- **Conveyor belt speed:** Measure the conveyor belt average speed to determine
+ the camera's adequate [shutter speed](#shutter-speed). The conveyor belt
+ speed must be between one and four meters per second.
+
+The conveyor belt length is variable because the belt moves continuously.
+However, consider capturing well-proportioned images. The field of view (FoV) is
+the area the camera covers on the captured images. This area must be
+approximately a square, covering the entire belt width and a similar distance
+for the belt length. So, for example, if the belt width is one meter, the belt
+length captured in the frame should also be approximately one meter to cover a
+square area.
+
+Keep an approximate ratio of 1:1 between the FoV and the camera height. So, for
+example, if your FoV captures an area of 1x1 square meters of the belt, you must
+mount the camera one or 1.5 meters above the conveyor belt.
+
+Calculate the following camera specifications based on the factors you measured:
+
+{{< table_of_contents >}}
+
+---
+
+## Sensor size
+
+Decide on the camera sensor size according to the information you need to
+detect. Larger sensor sizes fit more information, while smaller sensors apply
+cropping to lenses.
+
+An appropriate sensor size is between 2/3" to 1" for balanced high image quality
+and versatility in an industrial setting like a recycling facility. This
+recommendation considers the following aspects:
+
+- **Image quality:** Large sensors offer better image quality, higher dynamic
+ range, and improved low-light performance.
+- **Field of view (FoV):** Large sensors support a variety of lens options to
+ cover the entire width of the conveyor belt.
+- **Depth of field (DoF):** Large sensors achieve a greater DoF, which helps in
+ keeping all parts of the objects on the conveyor belt in focus. Use DoF
+ calculators like the [DoF simulator](https://dofsimulator.net/en/) to see the
+ effect of the [aperture](#aperture-size-f-number) and the sensor size on the
+ DoF.
+
+To convert sensor sizes to imaging area dimensions, review external references such as the [Photo Review table of common sensor sizes](https://www.photoreview.com.au/tips/buying/unravelling-sensor-sizes/).
+
+## Focal length
+
+Choose a focal length for the lens depending on the height of the camera from
+the conveyor belt and the required FoV. Assuming that the camera height is
+fixed, you can calculate how much of the conveyor belt the camera can see based
+on the focal length. Up to a certain point, a higher focal length makes it
+easier for the model to recognize what the material on an image is. Use the
+[sensor size](#sensor-size), camera height, and conveyor belt width to determine
+the focal length using the following formula:
+
+_f = (s)(d) / 𝑤_
+
+Where:
+
+- _f_ is the focal length
+- _s_ is the sensor width (for example, 6 mm assuming a 1/2.3" sensor size)
+- _d_ is the distance from the camera to the conveyor belt (camera height)
+- _𝑤_ is the width of the area to be covered (conveyor belt width)
+
+To simplify the calculation, install the camera at a height that equals the conveyor belt width. For example, suppose you have a [sensor size](#sensor-size) of 1/2.3″ such as the one the [ArduCam IMX477 sensor](https://www.arducam.com/product/b0242-arducam-imx477-hq-camera/) has. This sensor size equals a sensor width of 6 mm. If the conveyor belt width is one meter and you install the camera at one meter above the conveyor belt, then you can use the formula as follows:
+
+_f = (6)(1) / 1 = 6 mm_
+
+Therefore, if the lens's focal length is approximately equal to the sensor
+width, you can accommodate different conveyor belt widths by varying the camera
+height to be the same as the belt width.
+
+To estimate the focal length of a camera lens, you can use online calculators such as the [ArduCam focal length calculator](https://www.arducam.com/focal-length-calculator/).
+
+## Aperture size (f-number)
+
+The aperture controls the amount of light that gets into the camera sensor and
+has an inverse relationship with DoF. A wide aperture gives a shallow depth of
+field. On the contrary, a narrow aperture gives you a deeper DoF. Select an
+f-number that captures all items in focus as they pass through the conveyor
+belt. Choose values between f/2.8 and f/11, depending on the lightning, camera
+height and belt width.
+
+## Shutter speed
+
+The shutter speed impacts motion blur and overall brightness. The appropriate
+shutter speed depends on the conveyor belt's speed and the objects' motion.
+Faster conveyor belts require faster shutter speeds to avoid motion blur. Longer
+times on shutter speed let more light in, but motion blur increases. Calculate
+how much conveyor belt movement is acceptable while the shutter is open to
+prevent motion blurriness.
+
+Use the following formula to estimate shutter speed:
+
+_1 / (2)(frame rate)_
+
+A shutter speed between 1/250 and 1/500 seconds is generally suitable. You can
+start with 1/350 seconds to provide a good balance between reducing motion blur
+and maintaining image quality. Adjust based on real-world testing and ensure
+adequate lighting conditions support these faster shutter speeds.
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/table-of-specs.md b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/table-of-specs.md
new file mode 100644
index 00000000000..60640e90a9f
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-camera/table-of-specs.md
@@ -0,0 +1,25 @@
+The adequate specifications, such as sensor size, lens, height, or shutter
+speed, depend entirely on your facility's requirements.
+
+Conducting a complete engineering analysis is necessary for a thorough
+understanding of your facility's conditions. This includes factors like the size
+of objects on the conveyor belt, the speed of the belt, the positioning of the
+camera, and lighting conditions. However, the following table provides a
+summarized overview that can aid in the initial camera selection process.
+
+| **Specification** | **Recommendation** |
+| --------------------- | ----------------------------------------------------- |
+| Resolution | Values between 0.5 and 1.0 MP |
+| Sensor size | Values between 2/3" and 1" |
+| Aperture size | Values between f/2.8 and f/11, depending on lightning |
+| Frame rate | 30 FPS or higher |
+| Shutter type | Global shutter |
+| Shutter speed | Values between 1/250 and 1/500 seconds |
+| Lens compatibility | Interchangeable lenses |
+| Interface | USB 3.0, GigE, or Camera Link |
+| Light sensitivity | High sensitivity, high dynamic range |
+| Enclosure | Industrial grade, NEMA, IP65 or higher |
+| Software integration | Compatible with NVIDIA or Raspberry Pi |
+| Synchronization | Support for multi-camera synchronization |
+| Power supply | Power over ethernet (PoE) |
+| Scan type | Area scan camera |
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-edge-device/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-edge-device/_index.md
new file mode 100644
index 00000000000..5d795f2f37e
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/system-req/choose-edge-device/_index.md
@@ -0,0 +1,53 @@
+This section assumes [you chose an edge device](/official/projects/waste_identification_ml/circularnet-docs/content/solutions/#edge-device-deployment)
+as the computing unit to run CircularNet models. Selecting an edge device
+requires technical expertise to manage hardware and software configurations.
+
+Selecting and configuring the right edge device significantly impacts the
+performance of CircularNet ML models. To ensure optimal performance, we
+recommend using NVIDIA Jetson Xavier or Jetson Orin devices, configured with a
+minimum of 30 GB of RAM and between 100 GB and 200 GB of storage. This
+configuration is necessary to effectively run a Triton inference server, with
+which CircularNet operates because it is scalable and works with every kind of
+machine learning framework.
+
+The following are the recommended device specifications for an edge device:
+
+
+
+To run the Triton inference server in one of the recommended Jetson devices, use
+the Triton container available in the [NVIDIA DeepStream NGC catalog](https://catalog.ngc.nvidia.com/orgs/nvidia/containers/deepstream-l4t).
+Each device has a different JetPack version, a library that depends on the
+hardware and model. The required JetPack version for this Triton container is
+JetPack 6.0. When using JetPack 6.0, the container image needed is
+`nvcr.io/nvidia/deepstream-l4t:7.0-triton-multiarch`.
+
+For information about installing the latest version of JetPack on the device, see [JetPack SDK](https://developer.nvidia.com/embedded/jetpack). For other JetPack versions, refer to the [JetPack archive](https://developer.nvidia.com/embedded/jetpack-archive).
+
+Connect [the machine vision camera](/official/projects/waste_identification_ml/circularnet-docs/content/choose-camera/) to the edge device, which leverages the graphics processing unit (GPU) to run model inference. The captured images or videos and the inference results can be streamed back to Google Cloud.
+
+By configuring your device and completing the installation, you can ensure that
+your edge device is properly configured, a crucial step in running CircularNet
+models efficiently and effectively. To learn how to install CircularNet on an
+edge device, see [Deploy CircularNet](/official/projects/waste_identification_ml/circularnet-docs/content/deploy-cn/).
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/view-data/_index.md b/official/projects/waste_identification_ml/circularnet-docs/content/view-data/_index.md
new file mode 100644
index 00000000000..1188d45b86e
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/view-data/_index.md
@@ -0,0 +1,11 @@
+After running the prediction pipeline for object tracking and uploading the
+model results to BigQuery, you can configure a [Looker Studio](https://cloud.google.com/looker-studio)
+dashboard to get customized visualizations and point-in-time analysis to gain
+insight and take action toward waste management.
+
+CircularNet offers a Looker Studio template you can copy to your Google Cloud
+account to make your own dashboard. The template provides a baseline to adapt to
+your analytics and reporting needs. Customize your dashboard for case-specific
+charts and reports based on the model results stored in
+[BigQuery](https://cloud.google.com/bigquery/docs/introduction) as the data
+source.
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/view-data/before-you-begin.md b/official/projects/waste_identification_ml/circularnet-docs/content/view-data/before-you-begin.md
new file mode 100644
index 00000000000..d2936dced24
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/view-data/before-you-begin.md
@@ -0,0 +1,14 @@
+Before visualizing your prediction results on a dashboard, you must follow these
+steps:
+
+1. [Create a Google Cloud account](https://console.cloud.google.com/).
+1. [Open the Google Cloud console](https://cloud.google.com/cloud-console).
+1. Store your model results in BigQuery.
+ - If you [run the prediction pipeline for object tracking in Google Cloud](/official/projects/waste_identification_ml/circularnet-docs/content/analyze-data/prediction-pipeline-in-cloud),
+ you automatically load the data to BigQuery, so no action is required.
+ - If you use an edge device or any other solution for object tracking
+ outside of Google Cloud, you must [manually load the data from your
+ database to
+ BigQuery](https://cloud.google.com/bigquery/docs/loading-data).
+
+**Tip:** You can push data from an edge device to the cloud by configuring cloud API access from your code running on the edge device. For information on Jetson NVIDIA devices, see their [Reference Cloud documentation](https://docs.nvidia.com/moj/cloud/cloud-overview.html).
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/content/view-data/configure-dashboard.md b/official/projects/waste_identification_ml/circularnet-docs/content/view-data/configure-dashboard.md
new file mode 100644
index 00000000000..9a076ae6d32
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/content/view-data/configure-dashboard.md
@@ -0,0 +1,27 @@
+Follow these steps to configure a Looker Studio dashboard and visualize the
+model results for data analysis and reporting:
+
+1. Open the public Looker Studio template.
+1. Click **More options** > **Make a copy**.
+1. On the **Copy this report** window, enter the following details:
+
+ 1. Leave the value on the **Original data source** menu as shown because it
+ is the value of the template's BigQuery data source.
+ 1. On the **New data source** menu, select the name of the BigQuery table
+ that contains the model results you want to use as the data source for
+ the dashboard.
+
+1. If the **New data source** menu doesn't show your BigQuery table as an
+ available data source, perform the following actions:
+
+ 1. Click **Create data source** on the **New data source** menu.
+ 1. On the **Google connector** page, select **BigQuery**.
+ 1. If prompted, click **Authorize**.
+ 1. Select the Google Cloud project that hosts your BigQuery dataset.
+ 1. Select the dataset and table you want to use.
+ 1. Click **Connect**.
+ 1. On the new page, click **Add to report**.
+
+1. Click **Copy report**.
+
+A copy of the template dashboard is created in your Google Cloud account. You can change the title and customize this copy to fit your needs. For information about customizing and using the Looker Studio dashboard, see the [Quick start guide of Looker Studio](https://support.google.com/looker-studio/answer/9171315) and the [Looker documentation](https://cloud.google.com/looker/docs/intro).
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/hugo.toml b/official/projects/waste_identification_ml/circularnet-docs/hugo.toml
new file mode 100644
index 00000000000..f6c94abcbf7
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/hugo.toml
@@ -0,0 +1,80 @@
+baseURL = 'https://example.org/'
+languageCode = 'en-us'
+title = 'CircularNet Docs'
+theme = 'hugo-theme-techdoc'
+
+hasCJKLanguage = true
+metaDataFormat = "yaml"
+
+defaultContentLanguage = "en"
+defaultContentLanguageInSubdir= false
+enableMissingTranslationPlaceholders = false
+
+# Markup configure section
+# See https://gohugo.io/getting-started/configuration-markup/
+[markup]
+ defaultMarkdownHandler = "goldmark"
+ [markup.goldmark]
+ [markup.goldmark.renderer]
+ unsafe= true
+ [markup.goldmark.extensions]
+ [markup.goldmark.extensions.passthrough]
+ enable = true
+ [markup.goldmark.extensions.passthrough.delimiters]
+ block = [['\[', '\]'], ['$$', '$$']]
+ inline = [['\(', '\)']]
+ [markup.tableOfContents]
+ endLevel = 3
+ ordered = false
+ startLevel = 2
+
+[params]
+
+ math = true
+
+ # Source Code repository section
+ description = "put your description"
+ github_repository = ""
+ version = "1.0.0"
+
+ # Documentation repository section
+ # documentation repository (set edit link to documentation repository)
+ github_doc_repository = ""
+ github_doc_repository_path = ""
+ github_doc_repository_branch = ""
+
+ # Analytic section
+ google_analytics_id = "" # Your Google Analytics tracking id
+ tag_manager_container_id = "" # Your Google Tag Manager container id
+ google_site_verification = "" # Your Google Site Verification for Search Console
+
+ # Theme settings section
+ # Theme color
+ # See color value reference https://developer.mozilla.org/en-US/docs/Web/CSS/color
+ custom_font_color = ""
+ custom_background_color = ""
+
+ # Documentation Menu section
+ # Menu style settings
+ menu_style = "slide-menu" # "open-menu" or "slide-menu" or "" blank is as no sidebar
+
+ # Date format
+ dateformat = "" # default "2 Jan 2006"
+ # See the format reference https://gohugo.io/functions/format/#hugo-date-and-time-templating-reference
+
+ # path name excluded from documentation menu
+ menu_exclusion = [
+ "archives",
+ "archive",
+ "blog",
+ "entry",
+ "post",
+ "posts",
+ ]
+
+ # Algolia site search section
+ # See https://www.algolia.com/doc/
+ algolia_search_enable = true
+ algolia_indexName = "hugo-demo-techdoc"
+ algolia_appId = "7W4SAN4PLK"
+ algolia_apiKey = "cbf12a63ff72d9c5dc0c10c195cf9128" # Search-Only API Key
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/circularnet-docs/layouts/shortcodes/table_of_contents.html b/official/projects/waste_identification_ml/circularnet-docs/layouts/shortcodes/table_of_contents.html
new file mode 100644
index 00000000000..bbad869204b
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/layouts/shortcodes/table_of_contents.html
@@ -0,0 +1,3 @@
+
diff --git a/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/prepend-body.html b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/prepend-body.html
new file mode 100644
index 00000000000..3cebec89470
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/prepend-body.html
@@ -0,0 +1,10 @@
+{{ if hugo.IsServer }}
+
+{{ else }}
+{{- with .Site.Params.tag_manager_container_id -}}
+
+
+
+{{- end -}}
+{{- end -}}
diff --git a/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/search.html b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/search.html
new file mode 100644
index 00000000000..4a7097ea7f2
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/search.html
@@ -0,0 +1,57 @@
+
+
+
+
+
+
+
+
+
+
diff --git a/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/sidebar-footer.html b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/sidebar-footer.html
new file mode 100644
index 00000000000..e69de29bb2d
diff --git a/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/sidebar.html b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/sidebar.html
new file mode 100644
index 00000000000..6db7cfcadcb
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/sidebar.html
@@ -0,0 +1,12 @@
+{{ if ne .Site.Params.menu_style "" }}
+
+{{ if eq .Site.Params.menu_style "open-menu" }}
+{{- partial "menu/open-menu.html" . -}}
+{{ else if eq .Site.Params.menu_style "slide-menu" }}
+{{- partial "menu/slide-menu.html" . -}}
+{{ end }}
+
+
+{{ end }}
diff --git a/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/site-header.html b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/site-header.html
new file mode 100644
index 00000000000..08871ad733a
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/site-header.html
@@ -0,0 +1,12 @@
+
+
{{ .Site.Title }}
+{{- with .Site.Params.version -}}
+ Version {{ . }}
+{{- end -}}
+{{- with .Site.Params.github_repository -}}
+
+{{- end -}}
+{{ with .Site.Params.description }}
+
{{ . }}
+{{end}}
+
diff --git a/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/table-of-contents.html b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/table-of-contents.html
new file mode 100644
index 00000000000..a183a6a8ad4
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/partials/table-of-contents.html
@@ -0,0 +1,6 @@
+{{ if .Params.TableOfContents }}
+
+{{ end }}
diff --git a/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/posts/list.html b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/posts/list.html
new file mode 100644
index 00000000000..19d9496f3c9
--- /dev/null
+++ b/official/projects/waste_identification_ml/circularnet-docs/themes/hugo-theme-techdoc/layouts/posts/list.html
@@ -0,0 +1,6 @@
+{{- define "main" -}}
+
OS and storage: Click
+ Change and select the following:
+
+
Operating system: Deep Learning on Linux
+
Version: Deep Learning VM with
+ CUDA 12.4 preinstalled. Debian 11, Python 3.10. You can choose
+ any M number with this configuration, for example, M129.
+
Boot disk type: Balanced persistent disk
+
Size (GB): 200 GB
+
+
+
+
+ Note: Give your VM a name that is easy to remember and deploy in a region and a zone close to your physical location that allows GPUs.
+
+
+
+### 2. Download the Setup Script
+
+SSH into your VM instance. If this is the first access, you will be prompted
+to install nVidia drivers. After this is complete, run:
+
+```bash
+curl -o setup.sh https://raw.githubusercontent.com/tensorflow/models/master/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/setup.sh
+```
+
+### 3. Run the Setup Script
+
+Execute the setup script to download all required files and dependencies:
+
+```bash
+bash setup.sh
+```
+
+This will automatically download all necessary files for running the
+detection pipeline.
+
+### 4. Process Your Images
+
+Given a gcs bucket path containing your test images, run:
+
+```bash
+bash milk_pouch_product/run_pipeline.sh --gcs_path=/path/to/test_images
+```
+
+Replace `/path/to/test_images` with the actual bucket path to your image
+folder, for example:
+
+```bash
+bash milk_pouch_product/run_pipeline.sh --gcs_path=gs://dairy_product_detection/test_images/
+
+# Results will be in:
+# $gcs_path/predictions/dairy/
+# $gcs_path/predictions/others/
+```
+
+### Troubleshooting
+
+- Ensure your VM has sufficient memory and disk space
+- Verify that all image files are in supported formats (JPG, PNG, etc.)
+- Check that you have proper read/write permissions for the input directory
+
+## Dataset Creation for Training ML Models
+
+This guide explains how to create datasets for training image classifier, object
+detection, or instance segmentation models from images of a particular
+category.
+
+### 1. Prepare Your Images
+
+Organize your images into a folder. These should be images containing objects of
+a particular category (e.g., dairy products, bottles, cans, etc.).
+
+### 2. Run the Extract Objects Script
+
+Execute the following command to extract objects and generate dataset files:
+
+```python
+python3 extract_objects.py --gcs_path=/test_path --category_name=${category}
+```
+
+Replace:
+
+- `/test_path` with the path to your image folder.
+- `category` with your category name (e.g., bottles, cans, plastic, etc.)
+
+### 3. Generated Outputs
+
+The script will generate two types of outputs:
+
+#### For Image Classification Models
+
+A folder named **objects_for_classification** will be created containing all
+cropped objects extracted from the images. These cropped images can be
+directly used to train an image classifier model
+
+#### For Object Detection/Segmentation Models
+
+A COCO JSON file will be generated containing:
+
+- Annotations for all detected objects
+- Bounding boxes and segmentation masks
+- This file can be used to train object detection or instance segmentation models
+
+### Example Usage
+
+```python
+# Extract dairy products from images
+python3 extract_objects.py --gcs_path=/home/user/dairy_images --category_name=dairy
+
+# Extract plastic bottles
+python3 extract_objects.py --gcs_path=/home/user/bottle_images --category_name=bottles
+
+# Extract metal cans
+python3 extract_objects.py --gcs_path=/home/user/can_images --category_name=cans
+```
+
+### Output Structure
+
+After running the script, your directory will look like:
+
+```
+/test_images/
+├── image1.jpg
+├── image2.jpg
+├── objects_for_classification
+│ ├── crop_001.jpg
+│ ├── crop_002.jpg
+│ └── ...
+└── annotations.json # COCO format file for detection/segmentation
+```
+
+### Use Cases
+
+- **Image Classification Training/Finetuning**: Use images from
+`objects_for_classification/` folder
+- **Object Detection Training/Finetuning**: Use the COCO JSON file with
+original images
+- **Instance Segmentation Training/Finetuning**: Use the COCO JSON file with
+segmentation masks
+
+### Tips
+
+- Ensure your images are clear and objects are visible
+- Use consistent naming for category names across your datasets
+- Verify the generated annotations before training your models
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/cloudbuild.yaml b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/cloudbuild.yaml
new file mode 100644
index 00000000000..764bdb992c3
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/cloudbuild.yaml
@@ -0,0 +1,20 @@
+steps:
+# Build the image.
+- name: 'gcr.io/cloud-builders/docker'
+ args:
+ - 'build'
+ - '--no-cache'
+ - '--progress=plain'
+ - '--build-arg'
+ - 'GCS_PATH=${_GCS_PATH}'
+ - '-t'
+ - '${_REGION}-docker.pkg.dev/${PROJECT_ID}/${_REPO_NAME}/${_IMAGE_NAME}:latest'
+ - '.'
+
+# Push the container image to Artifact Registry
+images:
+- '${_REGION}-docker.pkg.dev/${PROJECT_ID}/${_REPO_NAME}/${_IMAGE_NAME}:latest'
+
+# Specify a high-performance machine type for the build
+options:
+ machineType: 'E2_HIGHCPU_32'
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/deploy.sh b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/deploy.sh
new file mode 100755
index 00000000000..3800c5c824e
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/deploy.sh
@@ -0,0 +1,328 @@
+#!/bin/bash
+
+# ==============================================================================
+# Fully Automated Deployment Script for Milk Pouch Detection.
+# Deploys a Compute Engine service, along with the
+# necessary GCP components, including:
+# 1. GCP project configuration.
+# 2. Service API enablement.
+# 3. BigQuery dataset and table creation.
+# 4. GCS bucket creation.
+# 5. Artifact Registry repository creation.
+# 6. Container image build and push (via Cloud Build).
+# 7. Service account creation and permission granting.
+# 8. Compute Engine deployment.
+#
+# This script is designed to be run locally.
+# It requires the Google Cloud CLI (gcloud) to be installed and authenticated.
+#
+# Usage:
+# ./deploy.sh \
+# [--gcp_project_id=] \
+# [--region=] \
+# [--zone=] \
+# [--device=cpu|gpu] \
+# [--compute=gce] \
+# [--source_bucket_name=]
+#
+# Arguments:
+# --gcp_project_id: Specify the GCP project ID.
+# --region: Specify the region for the resources. Default: asia-south1.
+# --zone: Specify the zone for the resources. Default: asia-south1-a.
+# --device: Specify the device type (cpu or gpu). Default: cpu.
+# --compute: Specify the compute platform (gce). Default: gce.
+# --source_bucket_name: Specify the source GCS bucket name.
+#
+# Example:
+# ./deploy.sh --device=gpu --compute=gce
+#
+# ==============================================================================
+
+# If any command fails, the script will stop immediately
+set -euo pipefail
+
+# -----------------------------------------------------------------------------
+# Configuration Block
+# -----------------------------------------------------------------------------
+
+# --- GCE Configuration ---
+# Name of the GCE instance if using GCE for processing
+export INSTANCE_NAME="milk-pouch-processor-vm"
+
+# --- Script Configuration ---
+# Name of the Artifact Registry repository
+export REPO_NAME="milk-pouch-classification-repo"
+
+# Name of the container image and Cloud Run service
+export IMAGE_NAME="milk-pouch-classification-service"
+
+# [Output] Name of the BigQuery Dataset
+export BQ_DATASET="milk_pouch_classification"
+
+# [Output] Name of the BigQuery Table
+export BQ_TABLE="milk_pouch_classification_results"
+
+# -----------------------------------------------------------------------------
+# Script Logic - Do not modify the following
+# -----------------------------------------------------------------------------
+
+# --- Argument Parsing ---
+# Set default values for device and compute platform
+PROJECT_ID="project-id-placeholder"
+REGION="asia-south1"
+ZONE="asia-south1-a" # Zone for the GCE instance
+DEVICE="gpu" # Default to GPU
+COMPUTE="gce" # Default to gce
+SOURCE_BUCKET_NAME=""
+
+# Parse command-line arguments --device [cpu|gpu] and --compute [cloud-run|gce]
+while [[ "$#" -gt 0 ]]; do
+ case $1 in
+ --gcp_project_id) PROJECT_ID="$2"; shift ;;
+ --region) REGION="$2"; shift ;;
+ --zone) ZONE="$2"; shift ;;
+ --device) DEVICE="$2"; shift ;;
+ --compute) COMPUTE="$2"; shift ;;
+ --source_bucket_name) SOURCE_BUCKET_NAME="$2"; shift ;;
+ *) echo "Unknown parameter passed"; exit 1 ;;
+ esac
+ shift
+done
+
+# Validate Project ID
+if [[ "${PROJECT_ID}" == "project-id-placeholder" ]]; then
+ echo "❌ Project ID is not specified. Please provide a valid project ID."
+ exit 1
+fi
+
+# Validate compute platform
+if [[ "${COMPUTE}" != "gce" ]]; then
+ echo "❌ Invalid service type specified. Currently only GCE is supported."
+ exit 1
+fi
+
+# Validate device type
+if [[ "${DEVICE}" != "cpu" && "${DEVICE}" != "gpu" ]]; then
+ echo "❌ Invalid device specified. Choose 'cpu' or 'gpu'."
+ exit 1
+fi
+
+# [Input] GCS Bucket name for uploading original images
+if [[ -z "${SOURCE_BUCKET_NAME}" ]]; then
+ export SOURCE_BUCKET_NAME="milk-pouch-classification-uploads-${PROJECT_ID}"
+fi
+
+
+echo "🚀 Starting deployment for a '${DEVICE}' configuration..."
+echo ""
+
+# -----------------------------------------------------------------------------
+# Final Configuration Summary
+# -----------------------------------------------------------------------------
+echo "✅ Deployment script is going to run with the following configuration."
+echo ""
+echo "--------------------------------------------------"
+echo "Configuration Summary:"
+echo "--------------------------------------------------"
+echo "Project ID: ${PROJECT_ID}"
+echo "Region: ${REGION}"
+echo "Compute Platform: ${COMPUTE}"
+echo "Device: ${DEVICE}"
+echo ""
+echo "Artifact Registry Repo: ${REPO_NAME}"
+echo "Image Name: ${IMAGE_NAME}"
+echo ""
+echo "Source GCS Bucket: ${SOURCE_BUCKET_NAME}"
+echo ""
+echo "BigQuery Dataset: ${BQ_DATASET}"
+echo "BigQuery Table: ${BQ_TABLE}"
+
+if [[ "$COMPUTE" == "gce" ]]; then
+ echo ""
+ echo "--- GCE Configuration ---"
+ echo "Instance Name: ${INSTANCE_NAME}"
+ echo "Zone: ${ZONE}"
+fi
+echo "--------------------------------------------------"
+
+echo "✅ Step 1: Configure gcloud CLI..."
+gcloud config set project "${PROJECT_ID}"
+gcloud config set run/region "${REGION}"
+echo "Project has been set to ${PROJECT_ID}, and region has been set to ${REGION}."
+echo ""
+
+# ---
+
+echo "✅ Step 2: Enable required GCP services..."
+gcloud services enable \
+ run.googleapis.com \
+ compute.googleapis.com \
+ artifactregistry.googleapis.com \
+ cloudbuild.googleapis.com \
+ logging.googleapis.com \
+ storage.googleapis.com \
+ iam.googleapis.com \
+ bigquery.googleapis.com \
+ pubsub.googleapis.com \
+ cloudscheduler.googleapis.com \
+ cloudresourcemanager.googleapis.com
+echo "All APIs have been enabled."
+echo ""
+
+# ---
+
+echo "✅ Step 3: Create BigQuery Dataset and Table..."
+bq --location="${REGION}" mk --dataset "${PROJECT_ID}:${BQ_DATASET}" \
+ || echo "Dataset '${BQ_DATASET}' already exists."
+bq mk --table "${PROJECT_ID}:${BQ_DATASET}.${BQ_TABLE}" \
+ ./src/milk_pouch_results_schema.json \
+ || echo "Table '${BQ_TABLE}' already exists."
+echo "BigQuery resources are ready."
+echo ""
+
+# ---
+
+echo "✅ Step 4: Create GCS Buckets..."
+gcloud storage buckets create \
+ --project "${PROJECT_ID}" \
+ --location "${REGION}" \
+ --default-storage-class standard \
+ --uniform-bucket-level-access "gs://${SOURCE_BUCKET_NAME}" \
+ || echo "Source Bucket 'gs://${SOURCE_BUCKET_NAME}' already exists."
+echo "GCS Buckets are ready."
+echo ""
+
+# ---
+
+echo "✅ Step 5: Create Artifact Registry repository..."
+gcloud artifacts repositories create "${REPO_NAME}" \
+ --repository-format=docker \
+ --location="${REGION}" \
+ --description="Docker repository for ML models" \
+ || echo "Repository '${REPO_NAME}' already exists."
+echo "Artifact Registry repository is ready."
+echo ""
+
+# ---
+
+echo "✅ Step 6: Build container image using Cloud Build (with cloudbuild.yaml)..."
+gcloud builds submit --timeout=2h --config cloudbuild.yaml \
+ --substitutions=_REGION="${REGION}",_REPO_NAME="${REPO_NAME}",_IMAGE_NAME="${IMAGE_NAME}",_GCS_PATH="gs://${SOURCE_BUCKET_NAME}"
+
+echo "Skipping container build step. Assuming image already exists in Artifact Registry."
+echo ""
+
+# ---
+
+echo "✅ Step 7: Create a dedicated service account and grant permissions..."
+SERVICE_ACCOUNT_NAME="milk-pouch-classification-sa"
+SERVICE_ACCOUNT_EMAIL="${SERVICE_ACCOUNT_NAME}@${PROJECT_ID}.iam.gserviceaccount.com"
+
+# Create service account if it doesn't exist
+gcloud iam service-accounts create "${SERVICE_ACCOUNT_NAME}" \
+ --display-name="Service Account for ${IMAGE_NAME}" \
+ || echo "Service account '${SERVICE_ACCOUNT_NAME}' already exists."
+
+echo "Granting IAM permissions to service account..."
+PROJECT_NUMBER=$( \
+ gcloud projects describe "${PROJECT_ID}" --format="value(projectNumber)")
+
+# Allow the service to invoke itself (required for Pub/Sub push
+# subscriptions)
+gcloud projects add-iam-policy-binding "${PROJECT_ID}" \
+ --member="serviceAccount:${SERVICE_ACCOUNT_EMAIL}" \
+ --role="roles/run.invoker" \
+ --condition=None > /dev/null 2>&1
+
+# Allow the service to write to GCS
+gcloud projects add-iam-policy-binding "${PROJECT_ID}" \
+ --member="serviceAccount:${SERVICE_ACCOUNT_EMAIL}" \
+ --role="roles/storage.objectAdmin" \
+ --condition=None > /dev/null 2>&1
+
+# Allow the service to write to BigQuery
+gcloud projects add-iam-policy-binding "${PROJECT_ID}" \
+ --member="serviceAccount:${SERVICE_ACCOUNT_EMAIL}" \
+ --role="roles/bigquery.dataEditor" \
+ --condition=None > /dev/null 2>&1
+
+# Allow the service to use Datastore
+gcloud projects add-iam-policy-binding "${PROJECT_ID}" \
+ --member="serviceAccount:${SERVICE_ACCOUNT_EMAIL}" \
+ --role="roles/datastore.user" \
+ --condition=None > /dev/null 2>&1
+
+# Allow the service to write logs
+gcloud projects add-iam-policy-binding "${PROJECT_ID}" \
+ --member="serviceAccount:${SERVICE_ACCOUNT_EMAIL}" \
+ --role="roles/logging.logWriter" \
+ --condition=None > /dev/null 2>&1
+
+# Allow the service to act as a Pub/Sub subscriber
+gcloud projects add-iam-policy-binding "${PROJECT_ID}" \
+ --member="serviceAccount:${SERVICE_ACCOUNT_EMAIL}" \
+ --role="roles/pubsub.subscriber" \
+ --condition=None > /dev/null 2>&1
+
+# Allow GCS to publish messages to the Pub/Sub topic
+GCS_SERVICE_AGENT="service-${PROJECT_NUMBER}@gs-project-accounts.iam.gserviceaccount.com"
+gcloud projects add-iam-policy-binding "${PROJECT_ID}" \
+ --member="serviceAccount:${GCS_SERVICE_AGENT}" \
+ --role="roles/pubsub.publisher" \
+ --condition=None > /dev/null 2>&1
+
+echo "All necessary IAM permissions have been granted."
+echo ""
+
+# ---
+
+echo "✅ Step 8: Deploy Compute Service..."
+if [[ "${COMPUTE}" == "gce" ]]; then
+ # --- GCE Deployment ---
+ echo "Setting GCP project..."
+ gcloud config set project ${PROJECT_ID}
+
+ echo "Deploying on the latest Deep Learning VM Image..."
+ GCE_CREATE_CMD="gcloud compute instances create ${INSTANCE_NAME} \
+ --project=${PROJECT_ID} \
+ --zone=${ZONE} \
+ --machine-type="n1-standard-4" \
+ --image-family="common-cu128-ubuntu-2204-nvidia-570" \
+ --image-project="deeplearning-platform-release" \
+ --boot-disk-size=200GB \
+ --scopes="cloud-platform" \
+ --maintenance-policy=TERMINATE \
+ --no-shielded-secure-boot \
+ --metadata-from-file="startup-script=gce_startup.sh" \
+ --metadata="IMAGE_URI=${REGION}-docker.pkg.dev/${PROJECT_ID}/${REPO_NAME}/${IMAGE_NAME}:latest""
+
+ if [[ "${DEVICE}" == "gpu" ]]; then
+ GCE_CREATE_CMD="${GCE_CREATE_CMD} --accelerator=\"type=nvidia-tesla-t4,count=1\""
+ fi
+
+ if eval "${GCE_CREATE_CMD}"; then
+ echo ""
+ echo "✅ Success! Compute Engine VM created."
+ else
+ GCE_EXIT_CODE=$?
+ if gcloud compute instances describe "${INSTANCE_NAME}" --zone "${ZONE}" --project "${PROJECT_ID}" > /dev/null 2>&1; then
+ echo "Compute Engine VM '${INSTANCE_NAME}' already exists."
+ else
+ echo "❌ Compute Engine VM creation failed with exit code ${GCE_EXIT_CODE}."
+ echo "Use cmd: gcloud compute images list --project deeplearning-platform-release --no-standard-images --filter='family:common-cu128 AND ubuntu' to check the available images."
+ fi
+ exit 1
+ fi
+
+ echo ""
+ echo "The process is now driven by the Compute Engine with '${IMAGE_NAME}'."
+ echo ""
+fi
+echo ""
+
+# ---
+
+echo "🚀🚀🚀 Deployment complete! 🚀🚀🚀"
+echo "You can upload files to the source bucket to be processed in the next run:"
+echo "gcloud storage cp your-local-image.jpg gs://${SOURCE_BUCKET_NAME}/"
+echo "Check the results in the bucket's subfolders."
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/gce_startup.sh b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/gce_startup.sh
new file mode 100644
index 00000000000..ce3cbe234dc
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/gce_startup.sh
@@ -0,0 +1,98 @@
+#!/bin/bash
+# This script is executed on a Google Compute Engine (GCE) instance upon startup.
+# It sets up Docker, authenticates with Google Artifact Registry, pulls a
+# specified Docker image, and runs it as a service.
+#
+# Below are the commands to debug the script and the docker container.
+# 1. To debug this script, run
+# `sudo journalctl -u google-startup-scripts.service` on the VM.
+# 2. To debug the docker container, run `docker ps -a` and
+# `docker logs `.
+
+# Exit immediately if any command fails.
+set -e
+
+echo "VM is ready. Driver is pre-installed."
+
+echo "--- Installing Docker Engine ---"
+
+# Install Docker using the official APT repository to ensure reliability and up-to-date packages.
+# First, update package lists and install prerequisites.
+apt-get update
+apt-get install -y ca-certificates curl
+
+echo "--- Setting up Environment Variables from GCE Metadata ---"
+IMAGE_URI=$(curl http://metadata.google.internal/computeMetadata/v1/instance/attributes/IMAGE_URI -H "Metadata-Flavor: Google")
+echo "IMAGE_URI: ${IMAGE_URI}"
+
+# Create a directory for Docker's GPG key.
+install -m 0755 -d /etc/apt/keyrings
+
+# Download Docker's official GPG key.
+curl -fsSL https://download.docker.com/linux/ubuntu/gpg -o /etc/apt/keyrings/docker.asc
+
+# Grant read permissions for the Docker GPG key.
+chmod a+r /etc/apt/keyrings/docker.asc
+
+# Add the Docker repository to APT sources.
+echo \
+ "deb [arch=$(dpkg --print-architecture) signed-by=/etc/apt/keyrings/docker.asc] https://download.docker.com/linux/ubuntu \
+ $(. /etc/os-release && echo "${VERSION_CODENAME}") stable" | \
+ tee /etc/apt/sources.list.d/docker.list > /dev/null
+
+# Update APT package lists again to include the new Docker repository.
+apt-get update
+
+# Install Docker Engine, CLI, containerd, and buildx/compose plugins.
+apt-get install \
+ -y docker-ce docker-ce-cli containerd.io docker-buildx-plugin docker-compose-plugin
+
+echo "Docker installed."
+
+echo "--- Authenticating Docker with gcloud ---"
+
+# Authenticate Docker to pull images from Google Artifact Registry.
+# `gcloud` is pre-installed on Deep Learning VM images.
+# This command configures Docker to use gcloud credentials for the specified registry domain.
+REGISTRY_HOST=$(echo "${IMAGE_URI}" | cut -d'/' -f1)
+gcloud auth configure-docker "${REGISTRY_HOST}" --quiet
+echo "Docker authenticated."
+
+# Define the Docker image and container name.
+DOCKER_IMAGE="${IMAGE_URI}"
+CONTAINER_NAME="milk-pouch-processor-vm"
+
+echo "Pulling your Docker image (${DOCKER_IMAGE}) (this may take a while)..."
+# Pull the specified Docker image from Google Artifact Registry.
+docker pull "${DOCKER_IMAGE}"
+echo "Image pull complete."
+
+echo "--- Waiting for NVIDIA GPU driver to be ready ---"
+# Loop until nvidia-smi runs successfully, indicating the driver is loaded.
+# This is crucial because the startup script might run before the GPU driver
+# kernel modules are fully loaded (a common race condition).
+until nvidia-smi; do
+ echo "Waiting for nvidia-smi to be available... (driver loading?)"
+ sleep 5
+done
+echo "NVIDIA GPU driver is ready."
+
+echo "Starting container loop... (This will run on the host VM)"
+
+# This loop runs on the HOST VM, not inside the container.
+# It restarts the *entire container* in each iteration.
+while true; do
+ echo "--- Running new container instance ---"
+
+ # --rm ensures the container and its resources (incl. GPU)
+ # are fully released upon exit.
+ docker run --rm \
+ --name "${CONTAINER_NAME}" \
+ --gpus all \
+ -e CATEGORY_NAME="milk_pouch" \
+ -e PYTHONUNBUFFERED=1 \
+ "${DOCKER_IMAGE}"|| true
+
+ echo "Container run finished. Restarting in 10 seconds..."
+ sleep 10
+done
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/__init__.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/__init__.py
new file mode 100644
index 00000000000..42c05c38983
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/__init__.py
@@ -0,0 +1,15 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Top level package for milk pouch detection."""
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/batched_io.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/batched_io.py
new file mode 100644
index 00000000000..c4a6bd2daeb
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/batched_io.py
@@ -0,0 +1,128 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utilities for batched I/O operations to improve pipeline performance."""
+
+from concurrent import futures
+import os
+from typing import List
+import cv2
+import numpy as np
+
+ThreadPoolExecutor = futures.ThreadPoolExecutor
+
+
+def save_masked_objects(
+ image: np.ndarray,
+ masks: List[np.ndarray],
+ boxes: List[np.ndarray],
+ source_image_path: str,
+ output_dir: str,
+) -> List[str]:
+ """Saves a batch of masked objects.
+
+ This function is called by BatchedMaskWriter's thread pool, so parallelism
+ is handled at the image level by the outer executor. Each call processes
+ one image's worth of objects sequentially.
+
+ Args:
+ image: Source image.
+ masks: Binary masks to save.
+ boxes: Bounding boxes corresponding to masks.
+ source_image_path: Path to the original image file.
+ output_dir: Directory path to save cropped objects.
+
+ Returns:
+ Paths to saved images.
+ """
+ image_name = os.path.splitext(os.path.basename(source_image_path))[0]
+
+ files_to_save = []
+ for idx, (mask, box) in enumerate(zip(masks, boxes)):
+ try:
+ masked_object = extract_masked_object(image, mask, box)
+ output_path = os.path.join(output_dir, f"{image_name}_object_{idx}.png")
+ result = cv2.imwrite(output_path, masked_object)
+ if result:
+ files_to_save.append(output_path)
+ except (ValueError, SystemError, AttributeError) as e:
+ print(f"[ERROR] Skipped object {idx}: {e}")
+ continue
+
+ return files_to_save
+
+
+def extract_masked_object(
+ image: np.ndarray, mask: np.ndarray, box: np.ndarray
+) -> np.ndarray:
+ """Extracts object from image using mask and bounding box.
+
+ Args:
+ image: Source image (H, W, 3).
+ mask: Binary mask (H, W).
+ box: Bounding box [x1, y1, x2, y2].
+
+ Returns:
+ Cropped and masked object as numpy array.
+ """
+
+ x1, y1, x2, y2 = map(int, box)
+
+ # Crop the image and mask
+ cropped_image = image[y1:y2, x1:x2]
+ cropped_mask = mask[y1:y2, x1:x2]
+ cropped_mask = cropped_mask.astype(bool)
+ masked_object = cropped_image.copy()
+ masked_object[~cropped_mask] = 1 # Black for non-masked pixels (background)
+ masked_object_rgb = cv2.cvtColor(masked_object, cv2.COLOR_BGR2RGB)
+
+ return masked_object_rgb
+
+
+class BatchedMaskWriter:
+ """A thread pool manager to parallelize object saving from images.
+
+ Each worker takes an image and then sequentially writes all of its detected
+ mask objects.
+ """
+
+ def __init__(self, output_dir: str):
+ self.output_dir = output_dir
+ self.executor = ThreadPoolExecutor(max_workers=4)
+ self.futures = []
+ os.makedirs(output_dir, exist_ok=True)
+
+ def __exit__(self, exc_type, exc_val, exc_tb):
+ # Wait for all writes to complete
+ for future in self.futures:
+ future.result()
+ self.executor.shutdown(wait=True)
+
+ def add_batch(
+ self,
+ image: np.ndarray,
+ masks: List[np.ndarray],
+ boxes: List[np.ndarray],
+ source_path: str,
+ ):
+ """Add a batch of masks to be written asynchronously."""
+ future = self.executor.submit(
+ save_masked_objects,
+ image,
+ masks,
+ boxes,
+ source_path,
+ self.output_dir,
+ )
+ self.futures.append(future)
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/batched_io_test.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/batched_io_test.py
new file mode 100644
index 00000000000..af7414c4bec
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/batched_io_test.py
@@ -0,0 +1,81 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import tempfile
+import time
+import unittest
+
+import numpy as np
+
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src import batched_io
+
+TEST_IMAGE_PATH = "/path/to/test_image.jpg"
+TEST_IMAGE = np.zeros((10, 10, 3), dtype=np.uint8)
+TEST_MASKS = np.ones((10, 10), dtype=np.uint8)
+TEST_BOXES = [[0, 0, 5, 5]]
+
+
+class BatchedIoTest(unittest.TestCase):
+ """Tests for batched_io."""
+
+ def test_extract_masked_object(self):
+ image = TEST_IMAGE.copy()
+ image[0:6, 0:6] = [128, 128, 128] # Add gray square
+ mask = TEST_MASKS.copy()
+ mask[1:5, 1:5] = 1 # Mask within gray square
+ box = [0, 0, 6, 6] # Bounding box of gray square
+
+ masked_object = batched_io.extract_masked_object(image, mask, box)
+
+ with self.subTest(name="CropDimensionsCorrect"):
+ # Check crop dimension, should be full bounding box dimensions.
+ self.assertEqual(masked_object.shape, (6, 6, 3))
+ with self.subTest(name="BackgroundIsBlack"):
+ boolean_mask = masked_object.astype(bool)
+ self.assertTrue(np.all(masked_object[~boolean_mask] == 1))
+ with self.subTest(name="ObjectIsRGB"):
+ self.assertEqual(masked_object.shape[-1], 3) # RGB
+
+ def test_save_masked_objects_saves_files(self):
+ image = np.zeros((10, 10, 3), dtype=np.uint8)
+ masks = [
+ np.ones((10, 10), dtype=np.uint8),
+ np.zeros((10, 10), dtype=np.uint8),
+ ]
+ masks[1][5:, 5:] = 1
+ boxes = [[0, 0, 5, 5], [5, 5, 10, 10]]
+ source_image_path = TEST_IMAGE_PATH
+
+ with tempfile.TemporaryDirectory() as temp_dir:
+ saved_files = batched_io.save_masked_objects(
+ image, masks, boxes, source_image_path, temp_dir
+ )
+
+ self.assertEqual(len(saved_files), 2)
+
+ def test_batched_mask_writer_queues_saves(self):
+ with tempfile.TemporaryDirectory() as temp_dir:
+ writer = batched_io.BatchedMaskWriter(output_dir=temp_dir)
+
+ writer.add_batch(TEST_IMAGE, TEST_MASKS, TEST_BOXES, TEST_IMAGE_PATH)
+ writer.add_batch(
+ TEST_IMAGE, TEST_MASKS, TEST_BOXES, "/another/path/to/image.jpg"
+ )
+ time.sleep(1)
+
+ self.assertEqual(len(writer.futures), 2)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/classify_images.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/classify_images.py
new file mode 100644
index 00000000000..e66aecaa0ca
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/classify_images.py
@@ -0,0 +1,80 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2025 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Classifies images based on Image Classifier."""
+
+import glob
+import os
+import shutil
+
+from absl import app
+from absl import flags
+import torch
+import tqdm
+
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src.models import classification
+
+FLAGS = flags.FLAGS
+INPUT_DIR = "input_images"
+CLASSIFICATION_DIR = "objects_for_classification"
+
+# Path to the custom trained model for Image Classifier.
+IMAGE_CLASSIFIER_WEIGHTS = "models/vit/best_vit_model_epoch_131.pt"
+CLASS_NAMES = ["dairy", "other"]
+
+
+def main(_) -> None:
+ parent_dir = os.path.dirname(os.path.abspath(INPUT_DIR))
+ predictions_dir = os.path.join(parent_dir, "predictions")
+
+ dairy_predictions = os.path.join(predictions_dir, "dairy")
+ other_predictions = os.path.join(predictions_dir, "others")
+ os.makedirs(dairy_predictions, exist_ok=True)
+ os.makedirs(other_predictions, exist_ok=True)
+
+ classifier = classification.ImageClassifier(
+ model_path=IMAGE_CLASSIFIER_WEIGHTS,
+ class_names=CLASS_NAMES,
+ device="cuda" if torch.cuda.is_available() else "cpu",
+ )
+
+ files = glob.glob(os.path.join(parent_dir, CLASSIFICATION_DIR, "*"))
+ print(f"Found {len(files)} images to process...")
+
+ total_dairy_packets = 0
+ for path in tqdm.tqdm(files):
+ pred_class, confidence = classifier.classify(path)
+ output_filename = f"{confidence:.2f}_{os.path.basename(path)}"
+ if pred_class == "dairy":
+ total_dairy_packets += 1
+ shutil.move(path, os.path.join(dairy_predictions, output_filename))
+ else:
+ shutil.move(path, os.path.join(other_predictions, output_filename))
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/coco_annotation_writer.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/coco_annotation_writer.py
new file mode 100644
index 00000000000..9491634b34e
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/coco_annotation_writer.py
@@ -0,0 +1,150 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Handles COCO JSON file creation for training dataset preparation.
+
+This module provides utilities for creating COCO JSON annotations
+from detected objects, to be used for training or finetuning object detection
+and segmentation models.
+"""
+
+import json
+import os
+from typing import List
+
+import numpy as np
+
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src import models_utils
+
+
+class CocoAnnotationWriter:
+ """Manages creation and writing of COCO-format annotations.
+
+ This class handles the incremental building of COCO JSON annotations as
+ images are processed, including image metadata and object annotations with
+ bounding boxes and segmentation masks.
+
+ Attributes:
+ category_name: Name of the object category for annotations.
+ coco_output: Dictionary containing COCO format data structure.
+ annotation_id_counter: Counter for generating unique annotation IDs.
+ image_id_counter: Counter for tracking processed images.
+ """
+
+ def __init__(self, category_name: str):
+ """Initializes the COCO annotation writer.
+
+ Args:
+ category_name: Name of the category (e.g., 'packets', 'dairy').
+ """
+ self.category_name = category_name
+ self.coco_output = models_utils.initialize_coco_output(category_name)
+ self.annotation_id_counter = 0
+ self.image_id_counter = 0
+
+ def add_image(
+ self,
+ file_path: str,
+ width: int,
+ height: int,
+ ) -> int:
+ """Adds image metadata to COCO output.
+
+ Args:
+ file_path: Path to the image file.
+ width: Image width in pixels.
+ height: Image height in pixels.
+
+ Returns:
+ The image ID assigned to this image.
+ """
+ image_info = {
+ "id": self.image_id_counter,
+ "file_name": os.path.basename(file_path),
+ "width": width,
+ "height": height,
+ }
+ self.coco_output["images"].append(image_info)
+ current_image_id = self.image_id_counter
+ self.image_id_counter += 1
+ return current_image_id
+
+ def add_annotations(
+ self,
+ image_id: int,
+ boxes: List[np.ndarray],
+ masks: List[np.ndarray],
+ ) -> int:
+ """Adds object annotations for a single image.
+
+ Args:
+ image_id: ID of the image these annotations belong to.
+ boxes: Bounding boxes in [x, y, width, height] format.
+ masks: Binary segmentation masks.
+
+ Returns:
+ Number of annotations successfully added.
+ """
+ annotations_added = 0
+
+ for box, mask in zip(boxes, masks):
+ try:
+ # Get the polygon points of masks
+ segmentation = models_utils.extract_largest_contour_segmentation(mask)
+ bbox_width, bbox_height, area = models_utils.get_bbox_details(box)
+
+ # Create annotation in COCO format
+ annotation_info = {
+ "id": self.annotation_id_counter,
+ "image_id": image_id,
+ "category_id": 1,
+ "bbox": [
+ int(box[0]),
+ int(box[1]),
+ int(bbox_width),
+ int(bbox_height),
+ ],
+ "area": int(area),
+ "iscrowd": 0,
+ "segmentation": segmentation,
+ }
+ self.coco_output["annotations"].append(annotation_info)
+ self.annotation_id_counter += 1
+ annotations_added += 1
+ except (ValueError, SystemError) as e:
+ print(f"[ERROR] Failed to create annotation: {e}")
+ continue
+
+ return annotations_added
+
+ def save(self, output_path: str) -> None:
+ """Saves COCO annotations to JSON file.
+
+ Args:
+ output_path: Full path where the JSON file should be saved.
+ """
+ with open(output_path, "w") as f:
+ json.dump(self.coco_output, f, indent=4)
+
+ def get_statistics(self) -> dict[str, int | str]:
+ """Gets statistics about the annotations created.
+
+ Returns:
+ Stats on how many annotations were created for what category.
+ """
+ return {
+ "num_images": self.image_id_counter,
+ "num_annotations": self.annotation_id_counter,
+ "category": self.category_name,
+ }
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/coco_annotation_writer_test.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/coco_annotation_writer_test.py
new file mode 100644
index 00000000000..885fc6d1628
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/coco_annotation_writer_test.py
@@ -0,0 +1,179 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2025 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for CocoAnnotationWriter class."""
+
+import json
+import os
+import tempfile
+import unittest
+
+import numpy as np
+
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src import coco_annotation_writer
+
+# Simple binary mask (20x20 square in center of 100x100 image)
+TEST_MASK = np.zeros((100, 100), dtype=np.uint8)
+TEST_MASK[40:60, 40:60] = 1
+
+# Bounding box [x, y, width, height]
+TEST_BOX = np.array([40, 40, 20, 20])
+
+
+class TestCocoAnnotationWriterAddImage(unittest.TestCase):
+ """Tests for the add_image method."""
+
+ def test_add_image_returns_correct_id(self):
+ """Test that add_image returns the correct image ID."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ image_id = writer.add_image("/path/to/image1.jpg", 800, 600)
+
+ self.assertEqual(image_id, 0)
+
+ def test_add_image_increments_counter(self):
+ """Test that image ID counter increments with each image."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ writer.add_image("/path/to/image1.jpg", 800, 600)
+ writer.add_image("/path/to/image2.jpg", 1024, 768)
+ id3 = writer.add_image("/path/to/image3.jpg", 640, 480)
+
+ self.assertEqual(id3, 2)
+
+ def test_add_image_extracts_basename(self):
+ """Test that only the filename is stored, not the full path."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ writer.add_image("/long/path/to/directory/image.jpg", 800, 600)
+
+ image_info = writer.coco_output["images"][0]
+ self.assertEqual(image_info["file_name"], "image.jpg")
+
+ def test_add_image_appends_to_list(self):
+ """Test that multiple images are appended to the images list."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ writer.add_image("/path/to/image1.jpg", 800, 600)
+ writer.add_image("/path/to/image2.jpg", 1024, 768)
+
+ self.assertEqual(len(writer.coco_output["images"]), 2)
+
+
+class TestCocoAnnotationWriterAddAnnotations(unittest.TestCase):
+ """Tests for the add_annotations method."""
+
+ def test_add_annotations_returns_count(self):
+ """Test that add_annotations returns the number of annotations added."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ count = writer.add_annotations(
+ image_id=0, boxes=[TEST_BOX], masks=[TEST_MASK]
+ )
+
+ self.assertEqual(count, 1)
+
+ def test_add_annotations_multiple_objects(self):
+ """Test adding multiple annotations for one image."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ mask2 = np.zeros((100, 100), dtype=np.uint8)
+ mask2[10:30, 10:30] = 1
+ box2 = np.array([10, 10, 20, 20])
+
+ count = writer.add_annotations(
+ image_id=0, boxes=[TEST_BOX, box2], masks=[TEST_MASK, mask2]
+ )
+
+ self.assertEqual(count, 2)
+
+ def test_add_annotations_links_to_image_id(self):
+ """Test that annotations are linked to the correct image ID."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ writer.add_annotations(image_id=42, boxes=[TEST_BOX], masks=[TEST_MASK])
+
+ annotation = writer.coco_output["annotations"][0]
+ self.assertEqual(annotation["image_id"], 42)
+
+
+class TestCocoAnnotationWriterSave(unittest.TestCase):
+ """Tests for the save method."""
+
+ def test_save_preserves_data(self):
+ """Test that saved data matches the internal coco_output."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+ writer.add_image("/path/to/image.jpg", 800, 600)
+
+ with tempfile.TemporaryDirectory() as tmpdir:
+ output_path = os.path.join(tmpdir, "test_output.json")
+ writer.save(output_path)
+
+ with open(output_path, "r") as f:
+ saved_data = json.load(f)
+
+ self.assertEqual(saved_data["images"], writer.coco_output["images"])
+
+
+class TestCocoAnnotationWriterGetStatistics(unittest.TestCase):
+ """Tests for the get_statistics method."""
+
+ def test_get_statistics_returns_correct_image_count(self):
+ """Test that statistics reflect the correct number of images."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+ writer.add_image("/path/to/image1.jpg", 800, 600)
+ writer.add_image("/path/to/image2.jpg", 1024, 768)
+
+ stats = writer.get_statistics()
+
+ self.assertEqual(stats["num_images"], 2)
+
+ def test_get_statistics_returns_correct_annotation_count(self):
+ """Test that statistics reflect the correct number of annotations."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ writer.add_annotations(
+ image_id=0, boxes=[TEST_BOX, TEST_BOX], masks=[TEST_MASK, TEST_MASK]
+ )
+
+ stats = writer.get_statistics()
+
+ self.assertEqual(stats["num_annotations"], 2)
+
+ def test_get_statistics_initial_state(self):
+ """Test that statistics are correct for a new writer with no data."""
+ writer = coco_annotation_writer.CocoAnnotationWriter("test_category")
+
+ stats = writer.get_statistics()
+
+ self.assertEqual(stats["num_images"], 0)
+ self.assertEqual(stats["num_annotations"], 0)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/extract_objects.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/extract_objects.py
new file mode 100644
index 00000000000..98ddd50b613
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/extract_objects.py
@@ -0,0 +1,211 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2025 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Detects, segments, and saves objects from images in a directory.
+
+This script initializes a computer vision pipeline to process images, identify
+objects based on a text prompt, and save each detected object as a separate
+cropped image to a temporary directory.
+"""
+
+import glob
+import os
+from typing import List, Tuple
+import warnings
+
+from absl import app
+from absl import flags
+import natsort
+import numpy as np
+import torch
+import tqdm
+
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src import batched_io # pyrefly: ignore[missing-module-attribute]
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src import coco_annotation_writer # pyrefly: ignore[missing-module-attribute]
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src.models import detection_segmentation
+
+
+warnings.filterwarnings("ignore", category=FutureWarning)
+warnings.filterwarnings("ignore", category=UserWarning)
+
+
+GROUNDING_DINO_WEIGHTS = "models/grounding_dino/groundingdino_swint_ogc.pth"
+GROUNDING_DINO_CONFIG = "models/grounding_dino/GroundingDINO_SwinT_OGC.py"
+SAM2_WEIGHTS = "models/sam2/sam2.1_hiera_large.pt"
+SAM2_CONFIG = "configs/sam2.1/sam2.1_hiera_l.yaml"
+TEXT_PROMPT = "packets"
+INPUT_DIR = "input_images"
+CLASSIFICATION_DIR = "objects_for_classification"
+COCO_OUTPUT_PATH = "coco_output.json"
+
+# Filter generated masks that are less than or equal to this
+# percentage of the overall image area
+MASK_FILTER_THRESHOLD_PERCENT = 1.0
+
+FLAGS = flags.FLAGS
+flags.DEFINE_string(
+ "category_name",
+ None,
+ "Name of the category. If provided, creates a COCO JSON file.",
+)
+
+
+def filter_masks_by_area(
+ masks: List[np.ndarray],
+ boxes: List[np.ndarray],
+ image_area: int,
+ min_percent: float = 1.0,
+) -> Tuple[List[np.ndarray], List[np.ndarray]]:
+ """Filter masks and boxes by minimum area threshold.
+
+ Args:
+ masks: Binary masks.
+ boxes: Bounding boxes corresponding to masks.
+ image_area: Total area of the source image.
+ min_percent: Minimum mask area as percentage of image area.
+
+ Returns:
+ Filtered masks and boxes.
+ """
+ filtered = [
+ (mask, box)
+ for mask, box in zip(masks, boxes)
+ if (np.sum(mask.astype(np.uint8)) / image_area) * 100 > min_percent
+ ]
+
+ if not filtered:
+ return [], []
+
+ valid_masks, valid_boxes = zip(*filtered)
+ return list(valid_masks), list(valid_boxes)
+
+
+def main(_) -> None:
+ """Runs the main object detection and extraction pipeline."""
+ if not os.path.isdir(INPUT_DIR):
+ raise ValueError(f"Input directory not found at '{INPUT_DIR}'")
+
+ # Check if COCO output should be created
+ create_coco = FLAGS.category_name is not None
+
+ print("Initializing image extraction and classification...")
+ try:
+ pipeline = detection_segmentation.ObjectDetectionSegmentation(
+ dino_config_path=GROUNDING_DINO_CONFIG,
+ dino_weights_path=GROUNDING_DINO_WEIGHTS,
+ sam_config_path=SAM2_CONFIG,
+ sam_checkpoint_path=SAM2_WEIGHTS,
+ )
+ except FileNotFoundError as e:
+ print(
+ f"\n⚠️ Error: Could not find model files: {e}. "
+ "Please check the paths in the script."
+ )
+ return
+ print("✅ Pipeline ready.")
+ os.makedirs(CLASSIFICATION_DIR, exist_ok=True)
+
+ # Initialize COCO annotation writer only if category_name is provided
+ coco_writer = None
+ if create_coco:
+ coco_writer = coco_annotation_writer.CocoAnnotationWriter(
+ FLAGS.category_name
+ )
+ print(
+ f"COCO JSON output will be created for category: {FLAGS.category_name}"
+ )
+ else:
+ print("No category name provided. Skipping COCO JSON creation.")
+
+ # Get all image files.
+ all_files = glob.glob(os.path.join(INPUT_DIR, "*"))
+ image_extensions = (".jpg", ".jpeg", ".png", ".bmp")
+ files = [f for f in all_files if f.lower().endswith(image_extensions)]
+ files = natsort.natsorted(files)
+
+ writer = batched_io.BatchedMaskWriter(CLASSIFICATION_DIR)
+
+ try:
+ for file_path in tqdm.tqdm(files):
+ try:
+ with torch.no_grad():
+ results = pipeline.detect_and_segment(file_path, TEXT_PROMPT)
+ if not results:
+ print("No objects detected")
+ continue
+ except (RuntimeError, ValueError) as e:
+ print(
+ "An unexpected error occurred with"
+ f" {os.path.basename(file_path)}: {e}"
+ )
+ continue
+
+ image = results["image"]
+ h, w = image.shape[:2]
+ image_area = h * w
+
+ valid_masks, valid_boxes = filter_masks_by_area(
+ results["masks"],
+ results["boxes"],
+ image_area,
+ min_percent=MASK_FILTER_THRESHOLD_PERCENT,
+ )
+
+ writer.add_batch(
+ image,
+ valid_masks,
+ valid_boxes,
+ file_path,
+ )
+
+ # Add image info to COCO output only if create_coco is True
+ if create_coco and coco_writer:
+ current_image_id = coco_writer.add_image(file_path, w, h)
+ coco_writer.add_annotations(current_image_id, valid_boxes, valid_masks)
+
+ finally:
+ # Ensure all I/O operations complete
+ if writer:
+ writer.__exit__(None, None, None)
+
+ # Save COCO JSON file only if create_coco is True
+ if create_coco and coco_writer:
+ output_path = os.path.join(INPUT_DIR, COCO_OUTPUT_PATH)
+ coco_writer.save(output_path)
+ stats = coco_writer.get_statistics()
+ print(
+ f"\n✅ COCO JSON saved to '{COCO_OUTPUT_PATH}'"
+ f" ({stats['num_images']} images,"
+ f" {stats['num_annotations']} annotations)."
+ )
+
+ print(f"✅ Cropped images saved to '{CLASSIFICATION_DIR}'.")
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/milk_pouch_results_schema.json b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/milk_pouch_results_schema.json
new file mode 100644
index 00000000000..d652bd4a6bf
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/milk_pouch_results_schema.json
@@ -0,0 +1,32 @@
+[
+ {
+ "name": "event_timestamp",
+ "type": "TIMESTAMP",
+ "mode": "NULLABLE",
+ "description": "The timestamp when the event was processed."
+ },
+ {
+ "name": "source_bucket",
+ "type": "STRING",
+ "mode": "NULLABLE",
+ "description": "The GCS bucket of the source image."
+ },
+ {
+ "name": "source_image",
+ "type": "STRING",
+ "mode": "NULLABLE",
+ "description": "The path to the source image within the bucket."
+ },
+ {
+ "name": "classification",
+ "type": "STRING",
+ "mode": "NULLABLE",
+ "description": "The predicted class for the detected object."
+ },
+ {
+ "name": "classification_probability",
+ "type": "FLOAT",
+ "mode": "NULLABLE",
+ "description": "The probability score of the classification."
+ }
+]
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/__init__.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/__init__.py
new file mode 100644
index 00000000000..e4a532fc634
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/__init__.py
@@ -0,0 +1,30 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Model clients and weights for computer vision and LLM inference.
+
+This package contains:
+ - Python client classes for model inference
+ - Pre-trained model weights organized by model type
+
+Structure:
+ - classification.py: ImageClassifier (ViT-B/16)
+ - detection_segmentation.py: ObjectDetectionSegmentation (Grounding DINO +
+ SAM2)
+ - llm.py: LlmModels (using Ollama interface)
+
+ - grounding_dino_model/: Grounding DINO weights and config
+ - sam2_model/: SAM2 weights
+ - image_classifier_model/: Fine-tuned ViT classifier weights
+"""
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/classification.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/classification.py
new file mode 100644
index 00000000000..23c9d6c54a9
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/classification.py
@@ -0,0 +1,188 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""ViT-based image classifier client for categorizing images."""
+
+import pathlib
+from typing import Sequence
+import warnings
+
+from PIL import Image
+import torch
+import torchvision
+
+# Suppress common warnings for a cleaner console output.
+warnings.filterwarnings('ignore', category=UserWarning)
+warnings.filterwarnings('ignore', category=FutureWarning)
+
+FEATURE_DIM = 768 # ViT-B/16 embedding size
+
+
+# TODO: b/455871640 - Add unit tests for this class.
+
+
+class ImageClassifier:
+ """ViT-based image classifier for categorizing images.
+
+ This class loads a fine-tuned ViT-B/16 model and provides methods for
+ image classification with automatic preprocessing.
+
+ Attributes:
+ model: The loaded ViT classifier model.
+ device: The PyTorch device the model runs on.
+ transform: The image preprocessing pipeline.
+ class_names: List of class names for predictions.
+ """
+
+ def __init__(
+ self,
+ model_path: str,
+ class_names: Sequence[str],
+ device: str = 'cuda',
+ image_size: tuple[int, int] = (224, 224),
+ ) -> None:
+ """Initializes the image classifier.
+
+ Args:
+ model_path: Path to the saved model state_dict.
+ class_names: List of class names corresponding to model output indices.
+ device: The hardware device to run the model on (e.g., "cuda", "cpu").
+ image_size: Target size (height, width) for resizing images.
+ """
+ self.device = torch.device(device)
+ self.class_names = class_names
+ self.transform = self._get_default_transform(image_size)
+
+ print('Loading ViT image classifier...')
+ self.model = self._load_vit_classifier(
+ pathlib.Path(model_path), len(class_names)
+ )
+ print('✅ ViT classifier loaded.')
+
+ def _load_vit_classifier(
+ self, model_path: pathlib.Path, num_classes: int
+ ) -> torch.nn.Module:
+ """Loads a fine-tuned ViT-B-16 model for inference.
+
+ Args:
+ model_path: Path to the saved model state_dict.
+ num_classes: Number of output classes for the model head.
+
+ Returns:
+ A PyTorch model in evaluation mode.
+ """
+ print(f'Loading model to {self.device}')
+
+ # Load base architecture.
+ model = torchvision.models.vit_b_16(weights=None)
+
+ # Freeze params.
+ for parameter in model.parameters():
+ parameter.requires_grad = False
+
+ # Set custom head.
+ model.heads = torch.nn.Linear(
+ in_features=FEATURE_DIM, out_features=num_classes
+ )
+
+ # Load the state_dict.
+ model.load_state_dict(torch.load(model_path, map_location=self.device))
+
+ # Set to device and eval mode.
+ model.to(self.device)
+ model.eval()
+
+ return model
+
+ def _get_default_transform(
+ self, image_size: tuple[int, int]
+ ) -> torchvision.transforms.Compose:
+ """Returns the default ImageNet transformation pipeline.
+
+ Args:
+ image_size: The target size (height, width) for resizing images.
+
+ Returns:
+ A torchvision Compose object representing the transformation pipeline.
+ """
+ return torchvision.transforms.Compose([
+ torchvision.transforms.Resize(image_size),
+ torchvision.transforms.ToTensor(),
+ # Standard mean and std values for ImageNet pre-trained models.
+ torchvision.transforms.Normalize(
+ mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]
+ ),
+ ])
+
+ def _process_image(self, image_path: pathlib.Path) -> torch.Tensor:
+ """Loads an image, applies transforms, and adds a batch dimension.
+
+ Args:
+ image_path: Path to the input image file.
+
+ Returns:
+ A transformed image tensor with a batch dimension of 1.
+ """
+ img = Image.open(image_path)
+ # Transform and add an extra dimension (batch_size = 1).
+ return self.transform(img).unsqueeze(dim=0)
+
+ def _predict(self, image_tensor: torch.Tensor) -> torch.Tensor:
+ """Performs inference on a single image tensor.
+
+ Args:
+ image_tensor: The input image tensor (with batch dimension).
+
+ Returns:
+ The raw logits output from the model.
+ """
+ # Move tensor to the same device as the model.
+ image_tensor = image_tensor.to(self.device)
+
+ # Turn on inference mode.
+ with torch.inference_mode():
+ return self.model(image_tensor)
+
+ def _get_prediction_details(self, logits: torch.Tensor) -> tuple[str, float]:
+ """Converts raw logits to a predicted class and its probability.
+
+ Args:
+ logits: The raw logits output from the model.
+
+ Returns:
+ A tuple containing the predicted class name and its probability.
+ """
+ probs = torch.softmax(logits, dim=1)
+
+ # Get the top probability and index.
+ pred_prob, pred_idx = torch.max(probs, dim=1)
+
+ # Get the class name and probability value.
+ pred_class = self.class_names[pred_idx.item()]
+ pred_prob_value = pred_prob.item()
+
+ return (pred_class, pred_prob_value) # pyrefly: ignore[bad-return]
+
+ def classify(self, image_path: str) -> tuple[str, float]:
+ """Classifies an image and returns the predicted class and probability.
+
+ Args:
+ image_path: Path to the input image file.
+
+ Returns:
+ A tuple containing the predicted class name and its probability.
+ """
+ image_tensor = self._process_image(pathlib.Path(image_path))
+ logits = self._predict(image_tensor)
+ return self._get_prediction_details(logits)
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/detection_segmentation.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/detection_segmentation.py
new file mode 100644
index 00000000000..6271f294739
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/detection_segmentation.py
@@ -0,0 +1,332 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Main bounding box detection and image segmentation logic."""
+
+import math
+from typing import Any, Optional
+import warnings
+
+from groundingdino.util import inference
+import models_utils # pyrefly: ignore[missing-import]
+import numpy as np
+from sam2 import build_sam
+from sam2 import sam2_image_predictor
+import torch
+
+# Suppress common warnings for a cleaner console output.
+warnings.filterwarnings('ignore', category=UserWarning)
+warnings.filterwarnings('ignore', category=FutureWarning)
+
+
+class ObjectDetectionSegmentation:
+ """Encapsulates vision models for object detection and segmentation.
+
+ This class provides a high-level API for using Grounding DINO and SAM2.
+ Models are loaded into memory once during initialization to avoid redundant
+ loading and improve performance for sequential processing tasks.
+
+ Attributes:
+ dino_model: The loaded Grounding DINO model.
+ sam_predictor: The initialized SAM2 predictor instance.
+ device: The PyTorch device (e.g., 'cuda' or 'cpu') the models run on.
+ """
+
+ def __init__(
+ self,
+ dino_config_path: str,
+ dino_weights_path: str,
+ sam_config_path: str,
+ sam_checkpoint_path: str,
+ device: str = 'cuda',
+ ) -> None:
+ """Initializes the vision pipeline by loading and setting up models.
+
+ Args:
+ dino_config_path: Path to the Grounding DINO configuration file.
+ dino_weights_path: Path to the Grounding DINO model weights file.
+ sam_config_path: Path to the SAM2 model configuration file.
+ sam_checkpoint_path: Path to the SAM2 model checkpoint file.
+ device: The hardware device to run models on (e.g., "cuda", "cpu").
+ """
+ self.device = torch.device(device)
+
+ print('Loading Grounding DINO model...')
+ self.dino_model = inference.load_model(dino_config_path, dino_weights_path)
+ self.dino_model.to(self.device)
+ print('✅ Grounding DINO model loaded.')
+
+ print('Loading SAM2 model...')
+ sam2_model = build_sam.build_sam2(
+ sam_config_path, sam_checkpoint_path, device=self.device
+ )
+ self.sam_predictor = sam2_image_predictor.SAM2ImagePredictor(sam2_model)
+ print('✅ SAM2 predictor initialized.')
+
+ def _detect_objects(
+ self,
+ image_path: str,
+ text_prompt: str,
+ box_threshold: float = 0.25,
+ text_threshold: float = 0.25,
+ ) -> tuple[np.ndarray, np.ndarray, torch.Tensor, list[str]]:
+ """Detects objects in an image using Grounding DINO based on a prompt.
+
+ Args:
+ image_path: The file path to the input image.
+ text_prompt: The text description of objects to detect.
+ box_threshold: The confidence threshold for object bounding boxes.
+ text_threshold: The confidence threshold for text-based labels.
+
+ Returns:
+ A tuple containing:
+ - image: The original image loaded as a NumPy array.
+ - boxes: Detected bounding boxes in CXCYWH format.
+ - scores: Confidence scores for each detected box.
+ - labels: Text labels corresponding to each box.
+ """
+ image, transformed_image = inference.load_image(image_path)
+ transformed_image = transformed_image.to(self.device)
+
+ boxes, scores, labels = inference.predict(
+ model=self.dino_model,
+ image=transformed_image,
+ caption=text_prompt,
+ box_threshold=box_threshold,
+ text_threshold=text_threshold,
+ )
+
+ return image, boxes, scores, labels
+
+ def _segment_objects(
+ self,
+ image: np.ndarray,
+ boxes: np.ndarray,
+ ) -> tuple[list[np.ndarray], list[torch.Tensor]]:
+ """Generates segmentation masks for batches of bounding boxes.
+
+ Args:
+ image: The source image as a NumPy array.
+ boxes: A NumPy array of bounding boxes in [x1, y1, x2, y2] format.
+
+ Returns:
+ A tuple containing:
+ - A list of boolean segmentation masks.
+ - A list of confidence scores for each mask.
+ """
+ self.sam_predictor.set_image(image)
+
+ # Stack all boxes for batched prediction
+ # SAM2 expects shape (num_boxes, 4)
+ batched_boxes = np.stack(boxes, axis=0) # pyrefly: ignore[no-matching-overload]
+ masks, scores, _ = self.sam_predictor.predict(
+ point_coords=None,
+ point_labels=None,
+ box=batched_boxes,
+ multimask_output=False,
+ )
+
+ return [mask.squeeze() for mask in masks], list(scores)
+
+ def _filter_boxes_by_area(
+ self, boxes: np.ndarray, max_box_area: float
+ ) -> list[np.ndarray]:
+ """Filter bounding boxes by area to remove overly large detections.
+
+ Args:
+ boxes: Array of bounding boxes in [x1, y1, x2, y2] format.
+ max_box_area: Maximum box area to keep
+
+ Returns:
+ List of valid bounding boxes that passed the area filter.
+ """
+ return [
+ bbox
+ for bbox in boxes
+ if (bbox[2] - bbox[0]) * (bbox[3] - bbox[1]) < max_box_area
+ ]
+
+ def _filter_boxes_at_image_edge(
+ self, boxes: list[np.ndarray], image_height: int, edge_tolerance: int = 5
+ ) -> list[np.ndarray]:
+ """Filter out bounding boxes that touch the top or bottom edges of the image.
+
+ Objects at frame edges are likely partial views.
+ Since each object appears in multiple frames, we can skip these.
+
+ Args:
+ boxes: List of bounding boxes in [x1, y1, x2, y2] format.
+ image_height: Height of the image in pixels.
+ edge_tolerance: Pixel tolerance for edge detection (default: 5).
+
+ Returns:
+ List of bounding boxes that don't touch top or bottom edges.
+ """
+ return [
+ bbox
+ for bbox in boxes
+ if bbox[1] > edge_tolerance and bbox[3] < image_height - edge_tolerance
+ ]
+
+ def _calculate_texture_variance(
+ self, image: np.ndarray, bbox: np.ndarray
+ ) -> float:
+ """Calculate texture variance within a bounding box.
+
+ Args:
+ image: The input image array.
+ bbox: Bounding box in [x1, y1, x2, y2] format.
+
+ Returns:
+ Standard deviation of pixel intensities (grayscale) within the box.
+ """
+ x1, y1, x2, y2 = bbox.astype(int)
+ crop = image[y1:y2, x1:x2]
+
+ # Return and convert to grayscale if not already
+ return float(
+ np.std(np.mean(crop, axis=2) if len(crop.shape) == 3 else crop)
+ )
+
+ def _filter_boxes_by_texture(
+ self,
+ boxes: list[np.ndarray],
+ image: np.ndarray,
+ min_texture_variance: float = 20.0,
+ ) -> tuple[list[np.ndarray], dict[int, float]]:
+ """Filter out boxes with very low texture variance.
+
+ This filter exists to remove misdetections caused by lighting artifacts.
+
+ Args:
+ boxes: List of bounding boxes in [x1, y1, x2, y2] format.
+ image: The input image array.
+ min_texture_variance: Minimum std dev of pixel intensities to keep.
+
+ Returns:
+ Tuple of
+ (filtered boxes, dict of texture variances).
+ """
+ variance_dict = {}
+ filtered_boxes = []
+
+ for i, bbox in enumerate(boxes):
+ variance = self._calculate_texture_variance(image, bbox)
+ variance_dict[i] = variance
+
+ if variance <= min_texture_variance:
+ print(f' Box {i} filtered (low texture variance: {variance:.2f})')
+ else:
+ filtered_boxes.append(bbox)
+
+ return filtered_boxes, variance_dict
+
+ def _filter_valid_boxes(
+ self,
+ boxes: np.ndarray,
+ image: np.ndarray,
+ image_shape: tuple[int, ...],
+ max_box_to_area_ratio: float,
+ min_texture_variance: float = 10.0,
+ ) -> tuple[list[np.ndarray], dict[int, float]]:
+ """Apply area, edge, and texture-based filtering to boxes.
+
+ Args:
+ boxes: Array of bounding boxes in [x1, y1, x2, y2] format.
+ image: The input image array (for texture analysis).
+ image_shape: Shape of the image (height, width, channels).
+ max_box_to_area_ratio: Maximum box area as ratio of image area.
+ min_texture_variance: Minimum texture variance to keep box.
+
+ Returns:
+ Tuple of (list of valid bounding boxes, dict of texture variances).
+ """
+ # Filter boxes by overall image area
+ image_area = math.prod(image_shape[:2])
+ valid_boxes = self._filter_boxes_by_area(
+ boxes, image_area * max_box_to_area_ratio
+ )
+ if not valid_boxes:
+ print('No objects passed area filter.')
+ return [], {}
+
+ # Filter boxes by intersection with edge of image
+ image_height = image_shape[0]
+ valid_boxes = self._filter_boxes_at_image_edge(valid_boxes, image_height)
+ if not valid_boxes:
+ print('No objects passed edge filter.')
+ return [], {}
+
+ # Filter boxes by texture variance
+ valid_boxes, variance_dict = self._filter_boxes_by_texture(
+ valid_boxes, image, min_texture_variance
+ )
+ if not valid_boxes:
+ print('No objects passed texture filter.')
+ return [], {}
+
+ return valid_boxes, variance_dict
+
+ def detect_and_segment(
+ self,
+ image_path: str,
+ text_prompt: str,
+ max_box_to_area_ratio: float = 0.25,
+ min_texture_variance: float = 10.0,
+ ) -> Optional[dict[str, Any]]:
+ """Runs detection and batched segmentation pipeline on an image.
+
+ This first uses GroundingDINO to extract bboxes from an image
+ based on a prompt, then passes all those boxes to SAM2 for
+ mask extraction.
+
+ Args:
+ image_path: The file path to the input image.
+ text_prompt: The text description of objects to use for box detection
+ max_box_to_area_ratio: Maximum box area as ratio of image area
+ min_texture_variance: Minimum texture variance to keep box
+
+ Returns:
+ A dictionary containing the processed data ('image', 'boxes', 'masks',
+ 'texture_variances') or None if no objects were detected.
+ """
+ print(f"\nProcessing '{image_path}'")
+ image, cxchywh_boxes, _, _ = self._detect_objects(image_path, text_prompt)
+
+ if cxchywh_boxes.shape[0] == 0:
+ print('No objects detected.')
+ return None
+
+ xyxy_boxes = models_utils.convert_boxes_cxcywh_to_xyxy(
+ cxchywh_boxes, image.shape
+ )
+ valid_boxes, variance_dict = self._filter_valid_boxes(
+ xyxy_boxes,
+ image,
+ image.shape,
+ max_box_to_area_ratio,
+ min_texture_variance,
+ )
+ if not valid_boxes:
+ return None
+
+ masks, _ = self._segment_objects(image, valid_boxes) # pyrefly: ignore[bad-argument-type]
+ print(f'Segmentation complete. Generated {len(masks)} masks.')
+
+ return {
+ 'image': image,
+ 'boxes': valid_boxes,
+ 'masks': masks,
+ 'texture_variances': variance_dict,
+ }
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/llm.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/llm.py
new file mode 100644
index 00000000000..02b813cd6af
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models/llm.py
@@ -0,0 +1,75 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Client for interacting with local LLMs via Ollama."""
+
+import subprocess
+import ollama
+
+
+class LlmModels:
+ """Provides an interface to interact with a local LLM via Ollama."""
+
+ def query_image_with_llm(
+ self, image_path: str, prompt: str, model_name: str
+ ) -> str:
+ """Sends an image and a text prompt to a local Ollama LLM.
+
+ Args:
+ image_path: Path to the image file.
+ prompt: The question or prompt for the LLM.
+ model_name: The name of the Ollama model to use (e.g., 'llava').
+
+ Returns:
+ The text response from the LLM.
+ """
+ response: ollama.ChatResponse = ollama.chat(
+ model=model_name,
+ messages=[{'role': 'user', 'content': prompt, 'images': [image_path]}],
+ options={
+ 'temperature': 0.0,
+ },
+ )
+ return response['message']['content']
+
+ def stop_model(self, model_name: str) -> None:
+ """Stops a running Ollama model to free up system resources.
+
+ This function executes the 'ollama stop' command-line instruction.
+
+ Args:
+ model_name: The name of the Ollama model to stop.
+ """
+ print(f'Attempting to stop Ollama model: {model_name}...')
+ try:
+ result = subprocess.run(
+ ['ollama', 'stop', model_name],
+ capture_output=True,
+ text=True,
+ check=False,
+ )
+ if result.returncode == 0:
+ print(f'✅ Successfully sent stop command for model: {model_name}')
+ else:
+ # This may not be an error if the model wasn't running.
+ print(
+ 'Info: Could not stop model (may not be running):'
+ f' {result.stderr.strip()}'
+ )
+ except FileNotFoundError:
+ print(
+ "⚠️ 'ollama' command not found. Is Ollama installed and in your PATH?"
+ )
+ except subprocess.CalledProcessError as e:
+ print(f'⚠️ An unexpected error occurred: {e}')
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models_utils.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models_utils.py
new file mode 100644
index 00000000000..65396cda4fc
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models_utils.py
@@ -0,0 +1,292 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Utility functions for milk pouch detection."""
+
+from collections.abc import Mapping
+import dataclasses
+import pathlib
+from typing import Any
+import cv2
+import matplotlib.pyplot as plt
+import numpy as np
+from PIL import Image
+import torch
+import torchvision
+
+
+@dataclasses.dataclass(frozen=True)
+class _BoundingBox:
+ """A class representing a bounding box."""
+ x1: float
+ y1: float
+ x2: float
+ y2: float
+
+
+def _box_area(box: _BoundingBox) -> float:
+ """Calculates the area of a bounding box.
+
+ Args:
+ box: A _BoundingBox object.
+
+ Returns:
+ The area of the bounding box.
+ """
+ return max(0, box.x2 - box.x1) * max(0, box.y2 - box.y1)
+
+
+def _calculate_iou(
+ box1: _BoundingBox,
+ box2: _BoundingBox
+) -> float:
+ """Calculates the Intersection over Union (IoU) of two bounding boxes.
+
+ Args:
+ box1: The first bounding box in (x1, y1, x2, y2) format.
+ box2: The second bounding box in (x1, y1, x2, y2) format.
+
+ Returns:
+ The IoU score, a float between 0.0 and 1.0.
+ """
+ # Determine the coordinates of the intersection rectangle
+ x1 = max(box1.x1, box2.x1)
+ y1 = max(box1.y1, box2.y1)
+ x2 = min(box1.x2, box2.x2)
+ y2 = min(box1.y2, box2.y2)
+
+ # Calculate the area of intersection
+ inter_area = max(0, x2 - x1) * max(0, y2 - y1)
+
+ # Calculate the area of both bounding boxes
+ box1_area = _box_area(box1)
+ box2_area = _box_area(box2)
+
+ # Calculate the area of the union
+ union_area = box1_area + box2_area - inter_area
+
+ # Compute the IoU score
+ return inter_area / union_area if union_area != 0 else 0.0
+
+
+def _is_contained(
+ inner_box: _BoundingBox,
+ outer_box: _BoundingBox,
+ margin: int = 5,
+) -> bool:
+ """Checks if one bounding box is contained within another, with a margin.
+
+ Args:
+ inner_box: The bounding box that is potentially inside.
+ outer_box: The bounding box that is potentially surrounding.
+ margin: An optional pixel margin to allow for slight inaccuracies.
+
+ Returns:
+ True if the inner box is contained within the outer box, False
+ otherwise.
+ """
+ return (
+ inner_box.x1 >= outer_box.x1 - margin
+ and inner_box.y1 >= outer_box.y1 - margin
+ and inner_box.x2 <= outer_box.x2 + margin
+ and inner_box.y2 <= outer_box.y2 + margin
+ )
+
+
+def filter_boxes_keep_smaller(
+ data: Mapping[str, list[Any]],
+ iou_threshold: float = 0.8,
+ area_threshold: int | None = None,
+ min_area: int = 1000,
+ margin: int = 5,
+) -> dict[str, list[Any]]:
+ """Filters overlapping bounding boxes, preferentially keeping smaller ones.
+
+ This function sorts boxes by area and iterates through them, discarding any
+ box that has a high IoU with an already-kept box or is contained within one.
+ This is useful for eliminating duplicate or redundant detections.
+
+ Args:
+ data: A dictionary containing 'boxes' and 'masks' lists.
+ iou_threshold: The IoU value above which a box is considered an overlap.
+ area_threshold: An optional maximum area to consider for a box.
+ min_area: The minimum area required for a box to be kept.
+ margin: The pixel margin used for the containment check.
+
+ Returns:
+ A dictionary with the filtered 'boxes' and their corresponding 'masks'.
+ """
+ # Check if the input data is valid
+ bounding_boxes = [_BoundingBox(*b) for b in data['boxes']]
+
+ areas = ([_box_area(b) for b in bounding_boxes])
+
+ # Sort boxes from smallest to largest area
+ sorted_indices = np.argsort(areas)
+ sorted_bounding_boxes = [bounding_boxes[i] for i in sorted_indices]
+
+ masks = np.array(data['masks'])
+ sorted_masks = masks[sorted_indices]
+
+ kept_boxes = []
+ kept_masks = []
+ kept_bounding_boxes_for_check = []
+
+ for i, box in enumerate(sorted_bounding_boxes):
+ current_area = _box_area(box)
+ if (
+ area_threshold is not None and current_area > area_threshold
+ ) or current_area < min_area:
+ continue
+
+ keep = True
+ for kept_box in kept_bounding_boxes_for_check:
+ if _calculate_iou(box, kept_box) > iou_threshold or _is_contained(
+ kept_box, box, margin
+ ):
+ keep = False
+ break
+
+ if keep:
+ kept_boxes.append([box.x1, box.y1, box.x2, box.y2])
+ kept_masks.append(sorted_masks[i])
+ kept_bounding_boxes_for_check.append(box)
+
+ return {'boxes': kept_boxes, 'masks': kept_masks}
+
+
+def convert_boxes_cxcywh_to_xyxy(
+ boxes: torch.Tensor, image_shape: tuple[int, int, int]
+) -> np.ndarray:
+ """Converts bounding boxes from center-based to corner-based format.
+
+ Args:
+ boxes: A tensor of bounding boxes in (cx, cy, w, h) format.
+ image_shape: A tuple representing the image dimensions (h, w, c).
+
+ Returns:
+ A NumPy array of bounding boxes in (x1, y1, x2, y2) format.
+ """
+ h, w, _ = image_shape
+ scale_factors = torch.tensor([w, h, w, h], device=boxes.device)
+ scaled_boxes = boxes * scale_factors
+ xyxy_boxes = torchvision.ops.box_convert(
+ boxes=scaled_boxes, in_fmt='cxcywh', out_fmt='xyxy'
+ )
+ return xyxy_boxes.cpu().numpy().astype(int)
+
+
+def initialize_coco_output(category_name: str) -> dict[str, list[Any]]:
+ """Initializes the COCO format output structure.
+
+ Args:
+ category_name: Name of the object category.
+
+ Returns:
+ A dictionary with the COCO format structure.
+ """
+ return {
+ 'categories': [{
+ 'id': 1,
+ 'name': category_name,
+ 'supercategory': 'object',
+ }],
+ 'images': [],
+ 'annotations': [],
+ }
+
+
+def plot_prediction(
+ image_path: pathlib.Path, pred_class: str, pred_prob: float
+):
+ """Plots the original image with its prediction and probability.
+
+ Args:
+ image_path: Path to the input image file.
+ pred_class: The predicted class name for the image.
+ pred_prob: The predicted probability for the class.
+ """
+ img = Image.open(image_path)
+ plt.figure()
+ plt.imshow(img)
+ plt.title(f'Pred: {pred_class} | Prob: {pred_prob:.3f}%')
+ plt.axis(False)
+ plt.show()
+
+
+def extract_largest_contour_segmentation(mask: np.ndarray) -> list[float]:
+ """Extracts the largest external contour from a binary mask.
+
+ This function finds all external contours in a binary mask and returns
+ the flattened coordinate list of the largest valid contour.
+
+ Args:
+ mask: A binary mask image (2D array) where the object is marked with 1s or
+ 255s. Shape should be (height, width).
+
+ Returns:
+ A list containing a single flattened contour with coordinates in the
+ format [x1, y1, x2, y2, ...]. Returns an empty list if no valid
+ contour is found.
+
+ Examples:
+ >>> mask = np.zeros((100, 100), dtype=np.uint8)
+ >>> mask[20:80, 20:80] = 1
+ >>> segmentation = extract_largest_contour_segmentation(mask)
+ """
+ mask_uint8 = mask.astype(np.uint8)
+
+ contours, _ = cv2.findContours(
+ mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE
+ )
+
+ if not contours:
+ return []
+
+ valid_segmentations = []
+ for contour in contours:
+ flattened = contour.flatten().tolist()
+ # Need at least 3 points (6 coordinates) for a valid polygon
+ if len(flattened) >= 6:
+ valid_segmentations.append(flattened)
+
+ if not valid_segmentations:
+ return []
+
+ return [max(valid_segmentations, key=len)] # pyrefly: ignore[bad-return]
+
+
+def get_bbox_details(box: list[int]) -> tuple[int, int, int]:
+ """Calculates width, height, and area from bounding box coordinates.
+
+ Args:
+ box: Bounding box coordinates in the format [x1, y1, x2, y2], where (x1,
+ y1) is the top-left corner and (x2, y2) is the bottom-right corner.
+
+ Returns:
+ A tuple containing (width, height, area) of the bounding box.
+
+ Examples:
+ >>> box = [10, 20, 50, 80]
+ >>> width, height, area = get_bbox_details(box)
+ >>> print(f"Width: {width}, Height: {height}, Area: {area}")
+ Width: 40, Height: 60, Area: 2400
+ """
+ x1, y1, x2, y2 = box
+
+ bbox_width = x2 - x1
+ bbox_height = y2 - y1
+ bbox_area = bbox_width * bbox_height
+
+ return bbox_width, bbox_height, bbox_area
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models_utils_test.py b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models_utils_test.py
new file mode 100644
index 00000000000..a99826f883b
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/models_utils_test.py
@@ -0,0 +1,143 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+import numpy as np
+import torch
+from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src import models_utils
+
+
+class UtilsTest(unittest.TestCase):
+ """Tests for the utility functions."""
+
+ def test_filter_boxes_keep_smaller_on_given_data(self):
+ boxes = [
+ [402.24, 0.54, 1343.04, 350.46],
+ [402.24, 0.54, 955.20, 333.18],
+ [930.24, 0.54, 1343.04, 351.54],
+ [611.52, 0.54, 955.20, 334.26],
+ [402.24, 0.54, 751.68, 305.10],
+ [749.76, 592.38, 1055.04, 875.34],
+ [941.76, 1012.50, 1039.68, 1078.38],
+ ]
+ masks = [f"mask_{i}" for i in range(len(boxes))]
+ data = {"boxes": boxes, "masks": masks}
+
+ expected_boxes = [
+ [941.76, 1012.50, 1039.68, 1078.38],
+ [749.76, 592.38, 1055.04, 875.34],
+ [402.24, 0.54, 751.68, 305.10],
+ [611.52, 0.54, 955.20, 334.26],
+ [930.24, 0.54, 1343.04, 351.54],
+ ]
+
+ result = models_utils.filter_boxes_keep_smaller(data)
+ actual_boxes = result["boxes"]
+
+ # Sort both lists for comparison (optional depending on importance of order)
+ actual_sorted = sorted(actual_boxes)
+ expected_sorted = sorted(expected_boxes)
+
+ self.assertEqual(len(actual_sorted), len(expected_sorted))
+ for box1, box2 in zip(actual_sorted, expected_sorted):
+ np.testing.assert_almost_equal(box1, box2, decimal=2)
+
+ def test_convert_boxes_cxcywh_to_xyxy_withsinglebox_returnscorrectcoordinates(
+ self,
+ ):
+ """Tests that a single box is converted correctly."""
+ boxes = torch.tensor([[0.5, 0.5, 0.2, 0.4]]) # cx, cy, w, h
+ image_shape = (100, 200, 3) # h, w, c
+ expected_boxes = np.array([[80, 30, 120, 70]]) # x1, y1, x2, y2
+
+ converted_boxes = models_utils.convert_boxes_cxcywh_to_xyxy(
+ boxes, image_shape
+ )
+ np.testing.assert_array_equal(converted_boxes, expected_boxes)
+
+ def test_convert_boxes_cxcywh_to_xyxy_withmultipleboxes_returnscorrectcoordinates(
+ self,
+ ):
+ """Tests that multiple boxes are converted correctly."""
+ boxes = torch.tensor([
+ [0.5, 0.5, 0.2, 0.4],
+ [0.25, 0.25, 0.1, 0.1],
+ ])
+ image_shape = (100, 200, 3)
+ expected_boxes = np.array([
+ [80, 30, 120, 70],
+ [40, 20, 60, 30],
+ ])
+ converted_boxes = models_utils.convert_boxes_cxcywh_to_xyxy(
+ boxes, image_shape
+ )
+ np.testing.assert_array_equal(converted_boxes, expected_boxes)
+
+ boxes = torch.empty((0, 4))
+ image_shape = (100, 200, 3)
+ expected_boxes = np.empty((0, 4), dtype=int)
+ converted_boxes = models_utils.convert_boxes_cxcywh_to_xyxy(
+ boxes, image_shape
+ )
+ np.testing.assert_array_equal(converted_boxes, expected_boxes)
+
+ def test_initialize_coco_output_with_category_name(self):
+ """Tests that COCO output structure is initialized correctly."""
+ category_name = "milk_pouch"
+ result = models_utils.initialize_coco_output(category_name)
+
+ self.assertIn("categories", result)
+ self.assertIn("images", result)
+ self.assertIn("annotations", result)
+
+ self.assertEqual(len(result["categories"]), 1)
+ self.assertEqual(result["categories"][0]["id"], 1)
+ self.assertEqual(result["categories"][0]["name"], category_name)
+ self.assertEqual(result["categories"][0]["supercategory"], "object")
+
+ self.assertEqual(result["images"], [])
+ self.assertEqual(result["annotations"], [])
+
+ def test_simple_rectangle_mask(self):
+ """Test extraction of contour from a simple rectangular mask."""
+ mask = np.zeros((100, 100), dtype=np.uint8)
+ mask[20:80, 20:80] = 1
+
+ result = models_utils.extract_largest_contour_segmentation(mask)
+
+ self.assertIsInstance(result, list)
+ self.assertEqual(len(result), 1)
+ self.assertGreaterEqual(len(result[0]), 6)
+
+ def test_empty_mask(self):
+ """Test that an empty mask returns an empty list."""
+ mask = np.zeros((100, 100), dtype=np.uint8)
+
+ result = models_utils.extract_largest_contour_segmentation(mask)
+
+ self.assertEqual(result, [])
+
+ def test_basic_bbox_calculation(self):
+ """Test basic width, height, and area calculation."""
+ box = [10, 20, 50, 80]
+
+ width, height, area = models_utils.get_bbox_details(box)
+
+ self.assertEqual(width, 40)
+ self.assertEqual(height, 60)
+ self.assertEqual(area, 2400)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/run_pipeline.sh b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/run_pipeline.sh
new file mode 100755
index 00000000000..ce958c05157
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/run_pipeline.sh
@@ -0,0 +1,150 @@
+#!/bin/bash
+
+# Exit immediately if a command exits with a non-zero status.
+set -o errexit
+
+# --- Parse command-line arguments ---
+batch_size=0 # Default to 0, meaning process all at once
+while [[ "$#" -gt 0 ]]; do
+ case $1 in
+ --gcs_path=*) gcs_path="${1#*=}"; shift ;;
+ --batch_size=*) batch_size="${1#*=}"; shift ;;
+ *) echo "❌ Unknown parameter passed: $1"; exit 1 ;;
+ esac
+done
+
+# --- Check if required arguments were provided ---
+if [ -z "$gcs_path" ]; then
+ echo "❌ Error: --gcs_path must be specified"
+ echo "✅ Usage: ./run_pipeline.sh --gcs_path=/path/to/images [--batch_size=N]"
+ exit 1
+fi
+
+# --- Run pipeline ---
+echo "✅ Activating virtual environment..."
+source myenv/bin/activate
+cd milk_pouch_project
+
+# List the image files in the GCS path.
+# NOTE: Adjust the grep pattern if other image types are expected.
+echo "=== DEBUGGING START ==="
+echo "DEBUG: gcs_path variable is: '${gcs_path}'"
+echo "DEBUG: Running 'gsutil ls \"${gcs_path}\"' to check accessibility:"
+gsutil ls "${gcs_path}" || echo "❌ gsutil ls failed"
+echo "DEBUG: Running 'gsutil ls -r \"${gcs_path}\" | head -n 10' to check content:"
+gsutil ls -r "${gcs_path}" | head -n 10 || echo "❌ gsutil recursive ls failed"
+echo "=== DEBUGGING END ==="
+
+echo "🖨️ Listing image files from GCS bucket: $gcs_path"
+mapfile -t all_gcs_files < <(gsutil ls -r "${gcs_path}" | grep -iE '\.(png|jpg|jpeg)$' | grep -v "/predictions/" | grep -v "/processed/")
+num_files=${#all_gcs_files[@]}
+
+if (( num_files == 0 )); then
+ echo "No image files found in $gcs_path. Exiting."
+ deactivate
+ exit 0
+fi
+
+# Create directories if they don't exist
+mkdir -p input_images
+mkdir -p predictions
+
+# Determine batch size
+if (( batch_size <= 0 )); then
+ echo "Processing all $num_files files at once."
+ batch_size=$num_files
+else
+ echo "Processing files in batches of $batch_size."
+fi
+
+# Iterate through files in batches
+for (( i=0; i num_files )); then
+ batch_end=$num_files
+ fi
+
+ # Get the current batch of files and calculate the number of files in it.
+ current_batch=("${all_gcs_files[@]:batch_start:batch_size}")
+ num_in_batch=${#current_batch[@]}
+ echo "--- Processing batch $(( i / batch_size + 1 ))/$(( (num_files + batch_size - 1) / batch_size )) ($num_in_batch files) ---"
+
+ # Clear previous batch's inputs and predictions
+ echo "🧹 Clearing input_images/ and predictions/..."
+ rm -rf input_images/*
+ rm -rf predictions/*
+
+ # Copy current batch files from GCS
+ echo "🖨️ Copying $num_in_batch files from GCS to input_images/..."
+ gsutil -m cp "${current_batch[@]}" input_images/
+
+ # Extract objects
+ echo "🔎 Extracting objects from images..."
+ if ! python3 extract_objects.py; then
+ echo "⚠️ Batch $(( i / batch_size + 1 )) failed during object extraction. Skipping to next batch."
+ continue
+ fi
+
+ # Classify objects
+ echo "🧠 Classifying objects..."
+ if ! python3 classify_images.py; then
+ echo "⚠️ Batch $(( i / batch_size + 1 )) failed during image classification. Skipping to next batch."
+ continue
+ fi
+
+ # Move predictions back to GCS
+ if [ -d "predictions" ] && [ -n "$(find predictions -type f -print -quit)" ]; then
+ echo "🖨️ Moving predictions for this batch back to GCS bucket: $gcs_path"
+ gsutil -m cp -r predictions/ "$gcs_path"
+ else
+ echo "⚠️ No predictions generated for this batch."
+ fi
+
+ # --- Move processed input files to 'processed/' directory preserving structure ---
+
+ # Ensure clean_gcs_path ends with / for correct substitution
+ clean_gcs_path="$gcs_path"
+ [[ "$clean_gcs_path" != */ ]] && clean_gcs_path="$clean_gcs_path/"
+
+ target_root="${clean_gcs_path}processed/"
+
+ # Group files by their destination directory to optimize gsutil calls
+ declare -a current_move_batch
+ current_move_dir=""
+
+ echo "📦 Moving processed files to ${target_root}..."
+
+ for file_url in "${current_batch[@]}"; do
+ # Get the directory of the file (e.g., gs://bucket/dev/2025-12-24/)
+ dir_url="$(dirname "$file_url")/"
+
+ # Calculate destination directory by injecting 'processed/'
+ # 1. Remove the base gcs_path from the file's dir to get the relative subdir (e.g., 2025-12-24/)
+ relative_dir="${dir_url#$clean_gcs_path}"
+ # 2. Append this relative dir to the processed root
+ dest_dir="${target_root}${relative_dir}"
+
+ # If the destination directory changes, flush the current batch
+ if [[ "$dest_dir" != "$current_move_dir" ]]; then
+ if (( ${#current_move_batch[@]} > 0 )); then
+ gsutil -m mv "${current_move_batch[@]}" "$current_move_dir"
+ current_move_batch=()
+ fi
+ current_move_dir="$dest_dir"
+ fi
+ current_move_batch+=("$file_url")
+ done
+
+ # Flush any remaining files
+ if (( ${#current_move_batch[@]} > 0 )); then
+ gsutil -m mv "${current_move_batch[@]}" "$current_move_dir"
+ fi
+
+ unset current_move_batch
+
+done
+
+echo "🧹 Deactivating virtual environment..."
+deactivate
+echo "✅ Done."
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/setup.sh b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/setup.sh
new file mode 100755
index 00000000000..589be2d0984
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src/setup.sh
@@ -0,0 +1,129 @@
+#!/bin/bash
+
+# This script sets up the complete environment for a computer vision project
+# using a model ensemble approach of bounding box detection, segmentation,
+# and classification. The classification model is a pre-trained VIT model.
+#
+# Usage: ./setup.sh [--cuda-version cu124|cu128]
+# --cuda-version: CUDA version to use (default: cu124)
+
+# Exit immediately if a command exits with a non-zero status.
+set -o errexit
+# Treat unset variables as an error when substituting.
+set -o nounset
+# Pipes fail if any command in the pipe fails.
+set -o pipefail
+
+# Parse command line arguments
+CUDA_VERSION="cu124"
+while [[ "$#" -gt 0 ]]; do
+ case $1 in
+ --cuda-version)
+ CUDA_VERSION="$2"
+ shift 2
+ ;;
+ *)
+ echo "Unknown option: $1"
+ echo "Usage: $0 [--cuda-version cu124|cu128]"
+ exit 1
+ ;;
+ esac
+done
+
+echo "Using CUDA version: $CUDA_VERSION"
+echo "-----"
+
+echo "🔹 Starting: Install System Dependencies"
+
+# Remove attempts to update deprecated packages
+sudo sed -i 's/^deb.*bullseye-backports/#&/' /etc/apt/sources.list
+sudo apt-get update
+sudo apt-get install -y python3-venv python3-pip lsof curl
+echo "✅ Finished: Install System Dependencies"
+echo "-----"
+
+echo "🔹 Starting: Create Virtual Environment"
+python3.10 -m venv myenv
+source myenv/bin/activate
+echo "✅ Finished: Create Virtual Environment"
+echo "-----"
+
+echo "🔹 Starting: Install Torch, Torchvision, Torchaudio"
+pip uninstall -y torch torchvision torchaudio > /dev/null 2>&1 || true
+pip install torch torchvision torchaudio --index-url "https://download.pytorch.org/whl/${CUDA_VERSION}"
+echo "✅ Finished: Install Torch, Torchvision, Torchaudio"
+echo "-----"
+
+echo "🔹 Starting: Install Grounding DINO"
+git clone https://github.com/IDEA-Research/GroundingDINO.git
+
+# Fix out of date cuda references during compile, can remove once
+# https://github.com/IDEA-Research/GroundingDINO/pull/415 is merged.
+cd GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn
+sed -i 's/value.type()/value.scalar_type()/g' ms_deform_attn_cuda.cu
+sed -i 's/value.scalar_type().is_cuda()/value.is_cuda()/g' ms_deform_attn_cuda.cu
+
+cd /home/${USER}/GroundingDINO/
+pip install -e .
+pip install timm==0.6.12
+cd ..
+echo "✅ Finished: Install Grounding DINO"
+echo "-----"
+
+echo "🔹 Starting: Install SAM2 and Required Python Packages"
+pip install --no-cache-dir \
+ opencv-python \
+ numpy \
+ ollama \
+ Pillow \
+ absl-py \
+ natsort \
+ "git+https://github.com/facebookresearch/sam2.git"
+echo "✅ Finished: Install SAM2 and Required Python Packages"
+echo "-----"
+
+echo "🔹 Starting: Create Project Directory Structure"
+mkdir -p milk_pouch_project/models/sam2
+mkdir -p milk_pouch_project/models/grounding_dino
+mkdir -p milk_pouch_project/models/vit
+echo "✅ Finished: Create Project Directory Structure"
+echo "-----"
+
+echo "🔹 Starting: Download SAM2 Checkpoint"
+wget -P ./milk_pouch_project/models/sam2 https://dl.fbaipublicfiles.com/segment_anything_2/092824/sam2.1_hiera_large.pt
+echo "✅ Finished: Download SAM2 Checkpoint"
+echo "-----"
+
+echo "🔹 Starting: Download GroundingDINO Model and Config"
+wget -P ./milk_pouch_project/models/grounding_dino https://github.com/IDEA-Research/GroundingDINO/releases/download/v0.1.0-alpha/groundingdino_swint_ogc.pth
+wget -P ./milk_pouch_project/models/grounding_dino https://raw.githubusercontent.com/IDEA-Research/GroundingDINO/refs/heads/main/groundingdino/config/GroundingDINO_SwinT_OGC.py
+echo "✅ Finished: Download GroundingDINO Model and Config"
+echo "-----"
+
+echo "🔹 Starting: Download Image Classifier Model"
+wget -P ./milk_pouch_project/models/vit https://storage.googleapis.com/tf_model_garden/vision/waste_identification_ml/dairy_product_packet_detection/best_vit_model_epoch_131.pt
+echo "✅ Finished: Download Image Classifier Model"
+echo "-----"
+
+echo "🔹 Starting: Clone Required Files from TensorFlow Models Repo"
+git clone --depth 1 --filter=blob:none --sparse https://github.com/tensorflow/models.git temp_tf_models
+cd temp_tf_models
+git sparse-checkout set official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src
+cd ..
+cp -r "temp_tf_models/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection/src"/* milk_pouch_project/
+rm -rf temp_tf_models
+echo "✅ Finished: Clone Required Files from TensorFlow Models Repo"
+echo "-----"
+
+echo "🔹 Starting: Modify Imports for Local Project Structure"
+find milk_pouch_project -type f -name "*.py" -exec sed -i \
+ 's|from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src.models import |from models import |g' {} +
+find milk_pouch_project -type f -name "*.py" -exec sed -i \
+ 's|from official.projects.waste_identification_ml.llm_applications.milk_pouch_detection.src import |import |g' {} +
+
+echo "✅ Finished: Modify Imports for Loscal Project Structure"
+echo "Files downloaded and modified successfully!"
+echo "-----"
+
+echo "🎉🎉🎉 Environment setup complete! 🎉🎉🎉"
+echo "-----"
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_florence2-sam2-gemma3.ipynb b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_florence2-sam2-gemma3.ipynb
new file mode 100644
index 00000000000..a94e1bddb2b
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_florence2-sam2-gemma3.ipynb
@@ -0,0 +1,633 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "id": "73Uw4PfHv6sv",
+ "metadata": {
+ "id": "73Uw4PfHv6sv"
+ },
+ "source": [
+ "# **Automatic Mask Generation Using Unsupervised Approach with Florence-2, SAM2, and Gemma3**"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "kMY1UaQMwX3O",
+ "metadata": {
+ "id": "kMY1UaQMwX3O"
+ },
+ "source": [
+ "In this notebook, we build an end-to-end unsupervised pipeline for object detection, segmentation, classification, and tracking—focusing on identifying and following milk pouches without manual labels. This approach leverages cutting-edge vision and language models and concludes with lightweight object tracking based on extracted features from segmentation masks.\n",
+ "\n",
+ "Key Components:\n",
+ "\n",
+ "\n",
+ "\n",
+ "1. **Florence-2 Multimodal Model**\u003cbr\u003e\n",
+ "A powerful vision-language model that performs generic object detection by returning bounding boxes around visually significant regions—completely label-free and prompt-driven.\n",
+ "2. **SAM2 (Segment Anything Model v2)**\u003cbr\u003e\n",
+ "Using the bounding boxes from Florence-2, SAM2 generates precise segmentation masks, enabling instance-level understanding and clean extraction of objects.\n",
+ "3. **Gemma3 12B QAT Model**\u003cbr\u003e\n",
+ "Each cropped masked region is passed to an open source Gemma3 quantization-aware large language model to determine whether it contains a milk pouch or not, enabling robust classification without explicit supervised training.\n",
+ "4. **Object Tracking via Mask Features**\u003cbr\u003e\n",
+ "For the final step, we extract distinguishing features from the segmented masks of positively identified milk pouches and use them to track the same objects across frames.\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "8w6XSOwm7rto",
+ "metadata": {
+ "id": "8w6XSOwm7rto"
+ },
+ "source": [
+ "While this colab focuses on the specific requirement of distinguishing milk sachets from other types (such as oil), the general approach could easily be adapted for other objects or use cases."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "4h669bAtw-XX",
+ "metadata": {
+ "id": "4h669bAtw-XX"
+ },
+ "source": [
+ "## Install and upgrade the necessary packages."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "MrAxLioBwfWg",
+ "metadata": {
+ "id": "MrAxLioBwfWg"
+ },
+ "outputs": [],
+ "source": [
+ "# Install the SAM2 (Segment Anything Model v2) library directly from the official Facebook Research GitHub repository\n",
+ "!pip install 'git+https://github.com/facebookresearch/sam2.git'"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "pJgWgODlibbB",
+ "metadata": {
+ "id": "pJgWgODlibbB"
+ },
+ "outputs": [],
+ "source": [
+ "!sudo apt-get update\n",
+ "!sudo apt-get install -y pciutils lshw\n",
+ "!pip install ollama"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "p4FMrtmls2KM",
+ "metadata": {
+ "id": "p4FMrtmls2KM"
+ },
+ "outputs": [],
+ "source": [
+ "# download the sample image from the circularnet project\n",
+ "url = (\n",
+ " \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\n",
+ " \"projects/waste_identification_ml/pre_processing/config/sample_images/\"\n",
+ " \"IMG_6509.png\"\n",
+ ")\n",
+ "\n",
+ "!curl -O {url} \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "FzOqkxvkxUoy",
+ "metadata": {
+ "id": "FzOqkxvkxUoy"
+ },
+ "source": [
+ "## Import the libraries and configure resources."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "SA-Qv2UkxYWU",
+ "metadata": {
+ "id": "SA-Qv2UkxYWU"
+ },
+ "outputs": [],
+ "source": [
+ "import torch, torchvision\n",
+ "from transformers import AutoProcessor, AutoModelForCausalLM\n",
+ "import sys\n",
+ "from sam2.build_sam import build_sam2\n",
+ "from sam2.sam2_image_predictor import SAM2ImagePredictor\n",
+ "from PIL import Image\n",
+ "import matplotlib.pyplot as plt\n",
+ "import matplotlib.patches as patches\n",
+ "import numpy as np\n",
+ "import tqdm\n",
+ "import math\n",
+ "import cv2\n",
+ "import tempfile\n",
+ "from google.colab.patches import cv2_imshow\n",
+ "from ollama import chat\n",
+ "from ollama import ChatResponse\n",
+ "import gc\n",
+ "import os\n",
+ "import glob\n",
+ "\n",
+ "print(\"PyTorch version:\", torch.__version__)\n",
+ "print(\"Torchvision version:\", torchvision.__version__)\n",
+ "print(\"CUDA is available:\", torch.cuda.is_available())\n",
+ "\n",
+ "# select the device for computation\n",
+ "if torch.cuda.is_available():\n",
+ " device = torch.device(\"cuda\")\n",
+ "elif torch.backends.mps.is_available():\n",
+ " device = torch.device(\"mps\")\n",
+ "else:\n",
+ " device = torch.device(\"cpu\")\n",
+ "print(f\"using device: {device}\")\n",
+ "\n",
+ "if device.type == \"cuda\":\n",
+ " # use bfloat16 for the entire notebook\n",
+ " torch.autocast(\"cuda\", dtype=torch.bfloat16).__enter__()\n",
+ " # turn on tfloat32 for Ampere GPUs (https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices)\n",
+ " if torch.cuda.get_device_properties(0).major \u003e= 8:\n",
+ " torch.backends.cuda.matmul.allow_tf32 = True\n",
+ " torch.backends.cudnn.allow_tf32 = True\n",
+ "elif device.type == \"mps\":\n",
+ " print(\n",
+ " \"\\nSupport for MPS devices is preliminary. SAM 2 is trained with CUDA and might \"\n",
+ " \"give numerically different outputs and sometimes degraded performance on MPS. \"\n",
+ " \"See e.g. https://github.com/pytorch/pytorch/issues/84936 for a discussion.\"\n",
+ " )"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "YCcsbDHdlk8G",
+ "metadata": {
+ "id": "YCcsbDHdlk8G"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Utils\n",
+ "\n",
+ "\n",
+ "def free_gpu_vars(*var_names, scope=None):\n",
+ " \"\"\"\n",
+ " Deletes variables (by name) from the provided scope (globals or locals),\n",
+ " collects garbage, and empties the CUDA cache.\n",
+ " \"\"\"\n",
+ " if scope is None:\n",
+ " scope = globals()\n",
+ " for var in var_names:\n",
+ " if var in scope:\n",
+ " del scope[var]\n",
+ " gc.collect()\n",
+ " torch.cuda.empty_cache()\n",
+ "\n",
+ "\n",
+ "def plot_bbox(image, data):\n",
+ " # Create a figure and axes\n",
+ " fig, ax = plt.subplots()\n",
+ "\n",
+ " # Display the image\n",
+ " ax.imshow(image)\n",
+ "\n",
+ " # Plot each bounding box\n",
+ " for bbox, label in zip(data['bboxes'], data['labels']):\n",
+ " # Unpack the bounding box coordinates\n",
+ " x1, y1, x2, y2 = bbox\n",
+ " # Create a Rectangle patch\n",
+ " rect = patches.Rectangle((x1, y1), x2-x1, y2-y1, linewidth=1, edgecolor='r', facecolor='none')\n",
+ " # Add the rectangle to the Axes\n",
+ " ax.add_patch(rect)\n",
+ " # Annotate the label\n",
+ " plt.text(x1, y1, label, color='white', fontsize=8, bbox=dict(facecolor='red', alpha=0.5))\n",
+ "\n",
+ " # Remove the axis ticks and labels\n",
+ " ax.axis('off')\n",
+ "\n",
+ " # Show the plot\n",
+ " plt.show()\n",
+ "\n",
+ "\n",
+ "def run_example(task_prompt, text_input=None):\n",
+ " if text_input is None:\n",
+ " prompt = task_prompt\n",
+ " else:\n",
+ " prompt = task_prompt + text_input\n",
+ " inputs = processor(text=prompt, images=image, return_tensors=\"pt\").to('cuda', torch.float16)\n",
+ " generated_ids = model.generate(\n",
+ " input_ids=inputs[\"input_ids\"].cuda(),\n",
+ " pixel_values=inputs[\"pixel_values\"].cuda(),\n",
+ " max_new_tokens=1024,\n",
+ " early_stopping=False,\n",
+ " do_sample=False,\n",
+ " num_beams=3,\n",
+ " )\n",
+ " generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]\n",
+ " parsed_answer = processor.post_process_generation(\n",
+ " generated_text,\n",
+ " task=task_prompt,\n",
+ " image_size=(image.width, image.height)\n",
+ " )\n",
+ "\n",
+ " return parsed_answer\n",
+ "\n",
+ "\n",
+ "def show_mask(\n",
+ " mask,\n",
+ " ax,\n",
+ " random_color=False,\n",
+ " borders = True\n",
+ "):\n",
+ " if random_color:\n",
+ " color = np.concatenate([np.random.random(3), np.array([0.6])], axis=0)\n",
+ " else:\n",
+ " color = np.array([30/255, 144/255, 255/255, 0.6])\n",
+ " h, w = mask.shape[-2:]\n",
+ " mask = mask.astype(np.uint8)\n",
+ " mask_image = mask.reshape(h, w, 1) * color.reshape(1, 1, -1)\n",
+ " if borders:\n",
+ " import cv2\n",
+ " contours, _ = cv2.findContours(mask,cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n",
+ " # Try to smooth contours\n",
+ " contours = [cv2.approxPolyDP(contour, epsilon=0.01, closed=True) for contour in contours]\n",
+ " mask_image = cv2.drawContours(mask_image, contours, -1, (1, 1, 1, 0.5), thickness=2)\n",
+ " ax.imshow(mask_image)\n",
+ "\n",
+ "\n",
+ "def show_points(\n",
+ " coords,\n",
+ " labels,\n",
+ " ax,\n",
+ " marker_size=375\n",
+ "):\n",
+ " pos_points = coords[labels==1]\n",
+ " neg_points = coords[labels==0]\n",
+ " ax.scatter(pos_points[:, 0], pos_points[:, 1], color='green', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)\n",
+ " ax.scatter(neg_points[:, 0], neg_points[:, 1], color='red', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)\n",
+ "\n",
+ "\n",
+ "def show_box(box, ax):\n",
+ " x0, y0 = box[0], box[1]\n",
+ " w, h = box[2] - box[0], box[3] - box[1]\n",
+ " ax.add_patch(plt.Rectangle((x0, y0), w, h, edgecolor='green', facecolor=(0, 0, 0, 0), lw=2))\n",
+ "\n",
+ "\n",
+ "def show_masks(\n",
+ " image,\n",
+ " masks,\n",
+ " scores,\n",
+ " point_coords=None,\n",
+ " box_coords=None,\n",
+ " input_labels=None,\n",
+ " borders=True\n",
+ "):\n",
+ " for i, (mask, score) in enumerate(zip(masks, scores)):\n",
+ " plt.figure(figsize=(10, 10))\n",
+ " plt.imshow(image)\n",
+ " show_mask(mask, plt.gca(), borders=borders)\n",
+ " if point_coords is not None:\n",
+ " assert input_labels is not None\n",
+ " show_points(point_coords, input_labels, plt.gca())\n",
+ " if box_coords is not None:\n",
+ " # boxes\n",
+ " show_box(box_coords, plt.gca())\n",
+ " if len(scores) \u003e 1:\n",
+ " plt.title(f\"Mask {i+1}, Score: {score:.3f}\", fontsize=18)\n",
+ " plt.axis('off')\n",
+ " plt.show()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "f3w2HICPxpIu",
+ "metadata": {
+ "id": "f3w2HICPxpIu"
+ },
+ "source": [
+ "## Read an image."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "9A3H9atRxqbk",
+ "metadata": {
+ "id": "9A3H9atRxqbk"
+ },
+ "outputs": [],
+ "source": [
+ "path = 'IMG_6509.png'\n",
+ "image = Image.open(path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "bGPfN2m-x-7T",
+ "metadata": {
+ "id": "bGPfN2m-x-7T"
+ },
+ "source": [
+ "## Download Florence-2 model."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "xeoo3uRmyERg",
+ "metadata": {
+ "id": "xeoo3uRmyERg"
+ },
+ "outputs": [],
+ "source": [
+ "model_id = 'microsoft/Florence-2-large'\n",
+ "model = AutoModelForCausalLM.from_pretrained(model_id, trust_remote_code=True, torch_dtype='auto').eval().cuda()\n",
+ "processor = AutoProcessor.from_pretrained(model_id, trust_remote_code=True)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "8Qvac4GAxtDE",
+ "metadata": {
+ "id": "8Qvac4GAxtDE"
+ },
+ "outputs": [],
+ "source": [
+ "# Perform object detection using Florence-2 OD task to detect all bboxes.\n",
+ "task_prompt = '\u003cCAPTION_TO_PHRASE_GROUNDING\u003e'\n",
+ "results = run_example(task_prompt, text_input=\"packets.\")\n",
+ "plot_bbox(image, results['\u003cCAPTION_TO_PHRASE_GROUNDING\u003e'])"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "zxnfK3UExuzb",
+ "metadata": {
+ "id": "zxnfK3UExuzb"
+ },
+ "outputs": [],
+ "source": [
+ "free_gpu_vars('model', 'processor')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "KoxQ7DynyQcN",
+ "metadata": {
+ "id": "KoxQ7DynyQcN"
+ },
+ "source": [
+ "## Download SAM-2 model."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "O57RyfEayo-_",
+ "metadata": {
+ "id": "O57RyfEayo-_"
+ },
+ "outputs": [],
+ "source": [
+ "# Create the 'checkpoints' directory one level up if it doesn't already exist\n",
+ "!mkdir -p checkpoints/\n",
+ "\n",
+ "# Download the pre-trained SAM2.1 Hiera Large model checkpoint into the 'checkpoints' directory\n",
+ "!wget -P checkpoints/ https://dl.fbaipublicfiles.com/segment_anything_2/092824/sam2.1_hiera_large.pt\n",
+ "\n",
+ "# Path to the pre-trained SAM2 model checkpoint\n",
+ "sam2_checkpoint = \"checkpoints/sam2.1_hiera_large.pt\"\n",
+ "\n",
+ "# Path to the configuration file for the SAM2 model variant being used\n",
+ "model_cfg = \"configs/sam2.1/sam2.1_hiera_l.yaml\"\n",
+ "\n",
+ "# Build the SAM2 model using the config and checkpoint; `device` should be set to \"cuda\" or \"cpu\"\n",
+ "sam2_model = build_sam2(model_cfg, sam2_checkpoint, device=device)\n",
+ "\n",
+ "# Create a predictor object using the loaded SAM2 model for image-based mask prediction\n",
+ "sam2_predictor = SAM2ImagePredictor(sam2_model)\n",
+ "\n",
+ "# Perform segmentation on bbox cordinates using SAM2 model.\n",
+ "sam2_predictor.set_image(image)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "3QUFPffLyKaN",
+ "metadata": {
+ "id": "3QUFPffLyKaN"
+ },
+ "outputs": [],
+ "source": [
+ "# Create a directory to store the cropped object images.\n",
+ "os.makedirs('tempdir', exist_ok=True)\n",
+ "\n",
+ "# Use bounding boxes to extract mask for each object..\n",
+ "for idx, bbox in tqdm.tqdm(enumerate(results['\u003cCAPTION_TO_PHRASE_GROUNDING\u003e']['bboxes'])):\n",
+ " x1, y1, x2, y2 = list(map(round, bbox))\n",
+ " if (x2-x1)*(y2-y1) \u003c 0.25 * math.prod(image.size):\n",
+ " input_box = np.array([x1, y1, x2, y2])\n",
+ "\n",
+ " masks, scores, _ = sam2_predictor.predict(\n",
+ " point_coords=None,\n",
+ " point_labels=None,\n",
+ " box=input_box[None, :],\n",
+ " multimask_output=False,\n",
+ " )\n",
+ " # show_masks(image, masks, scores, box_coords=input_box)\n",
+ "\n",
+ " # Convert the first mask to 0-255 and expand its dimensions to match the image channels.\n",
+ " # Multiply the mask with the original image (preserves object, sets background to 0).\n",
+ " # Crop the masked image to the bounding box [y1:y2, x1:x2].\n",
+ " masked_object = Image.fromarray(\n",
+ " np.where(\n",
+ " np.expand_dims(masks[0]*255, -1),\n",
+ " np.array(image), 0\n",
+ " )[y1:y2, x1:x2]\n",
+ " )\n",
+ "\n",
+ " image_path = f'tempdir/{os.path.splitext(path)[0]}_{idx}.png'\n",
+ " masked_object.save(image_path)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "5vnrB5S6ygAJ",
+ "metadata": {
+ "id": "5vnrB5S6ygAJ"
+ },
+ "outputs": [],
+ "source": [
+ "free_gpu_vars('masks', 'scores', 'sam2_predictor', 'sam2_model')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "-GBtU5QtyrVc",
+ "metadata": {
+ "id": "-GBtU5QtyrVc"
+ },
+ "source": [
+ "## Download Gemma3 model using Ollama tool."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "Ju0RGkY3ivYV",
+ "metadata": {
+ "id": "Ju0RGkY3ivYV"
+ },
+ "source": [
+ "Run the following commands in the terminal within your colab notebook."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "UTJBO5HilTrY",
+ "metadata": {
+ "id": "UTJBO5HilTrY"
+ },
+ "source": [
+ "\n",
+ "\n",
+ "```\n",
+ "curl https://ollama.ai/install.sh | sh\n",
+ "ollama serve\n",
+ "```\n",
+ "\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "X_ICcwY6jazw",
+ "metadata": {
+ "id": "X_ICcwY6jazw"
+ },
+ "outputs": [],
+ "source": [
+ "# Pull the required open sourced LLM model.\n",
+ "!ollama pull gemma3:12b-it-qat"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "PEyam8fLjzyq",
+ "metadata": {
+ "id": "PEyam8fLjzyq"
+ },
+ "outputs": [],
+ "source": [
+ "# Check if the model is downloaded.\n",
+ "!ollama list"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "rnOHG6_5orqj",
+ "metadata": {
+ "id": "rnOHG6_5orqj"
+ },
+ "outputs": [],
+ "source": [
+ "# Prompt to analyze an image for milk packet vs others.\n",
+ "prompt = \"\"\"\n",
+ "Analyze the provided image of a packaging. Was this packaging used to contain milk or a milk-based product? Answer in yes or no only.\n",
+ "\"\"\""
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "id": "F5HuNwRczB45",
+ "metadata": {
+ "id": "F5HuNwRczB45"
+ },
+ "source": [
+ "Read an cropped images to perform inference using LLM."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "TZkeWYq_zEXs",
+ "metadata": {
+ "id": "TZkeWYq_zEXs"
+ },
+ "outputs": [],
+ "source": [
+ "images = glob.glob('tempdir/*.png')\n",
+ "\n",
+ "for path in images:\n",
+ " # Run the chat/inference API, sending the temporary masked object image as input.\n",
+ " response: ChatResponse = chat(model='gemma3:12b-it-qat', messages=[\n",
+ " {\n",
+ " 'role': 'user',\n",
+ " 'content': prompt,\n",
+ " 'images': [path]\n",
+ " },\n",
+ " ])\n",
+ " image = cv2.imread(path)\n",
+ " plt.imshow(image)\n",
+ " plt.axis('off')\n",
+ " plt.show()\n",
+ "\n",
+ " # Print the model's response content (the generated answer)\n",
+ " print(f\"\\n{response.message.content}\")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "id": "4a-YeFcFzvVV",
+ "metadata": {
+ "id": "4a-YeFcFzvVV"
+ },
+ "outputs": [],
+ "source": [
+ "!ollama stop gemma3:12b-it-qat"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "machine_shape": "hm",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "codemirror_mode": {
+ "name": "ipython",
+ "version": 3
+ },
+ "file_extension": ".py",
+ "mimetype": "text/x-python",
+ "name": "python",
+ "nbconvert_exporter": "python",
+ "pygments_lexer": "ipython3",
+ "version": "3.10.14"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 5
+}
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_groundingdino-sam2-gemma3.ipynb b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_groundingdino-sam2-gemma3.ipynb
new file mode 100644
index 00000000000..4ac3c4007db
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_groundingdino-sam2-gemma3.ipynb
@@ -0,0 +1,517 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZAppGMCpM9Qq"
+ },
+ "source": [
+ "# Automatic Mask Generation Using Unsupervised Approach with Grounding Dino, SAM2, and Gemma3"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cwNkr2MGNGQ3"
+ },
+ "source": [
+ "In this notebook, we build an end-to-end unsupervised pipeline for object detection, segmentation, classification, and tracking—focusing on identifying and following milk pouches without manual labels. This approach leverages cutting-edge vision and language models and concludes with lightweight object tracking based on extracted features from segmentation masks.\n",
+ "\n",
+ "Key Components:\n",
+ "\n",
+ "\n",
+ "\n",
+ "1. **Grounding Dino**\n",
+ "\n",
+ "A powerful vision-language model that performs generic object detection by returning bounding boxes around visually significant regions—completely label-free and prompt-driven.\n",
+ "\n",
+ "2. **SAM2 (Segment Anything Model v2)**\n",
+ "\n",
+ "Using the bounding boxes from Grounding Dino, SAM2 generates precise segmentation masks, enabling instance-level understanding and clean extraction of objects.\n",
+ "\n",
+ "3. **Gemma3 12B QAT Model**\n",
+ "\n",
+ "Each cropped masked region is passed to an open source Gemma3 quantization-aware large language model to determine whether it contains a milk pouch or not, enabling robust classification without explicit supervised training.\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "S8JMep1cNw9K"
+ },
+ "source": [
+ "## Install necessary packages.\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "dHhKN6bfylwZ"
+ },
+ "outputs": [],
+ "source": [
+ "!git clone 'https://github.com/IDEA-Research/Grounded-SAM-2'\n",
+ "!pip install 'git+https://github.com/IDEA-Research/Grounded-SAM-2'\n",
+ "\n",
+ "%cd 'Grounded-SAM-2'\n",
+ "\n",
+ "# Install SAM2\n",
+ "!pip install -e .\n",
+ "\n",
+ "# Install Grounding Dino\n",
+ "!pip install --no-build-isolation -e grounding_dino\n",
+ "\n",
+ "!pip install addict yapf supervision\u003e=0.22.0"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Oiv9cRLHZPMj"
+ },
+ "outputs": [],
+ "source": [
+ "# Required for Ollama to detect GPUs.\n",
+ "!sudo apt-get install -y pciutils lshw\n",
+ "!pip install ollama"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "yuJRa2S1N2CH"
+ },
+ "source": [
+ "## Import model weights and configuration files."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iZhkdMt68zPh"
+ },
+ "outputs": [],
+ "source": [
+ "# Download Grounding Dino weights.\n",
+ "!mkdir grounding_dino_weights\n",
+ "!wget -P ./grounding_dino_weights https://github.com/IDEA-Research/GroundingDINO/releases/download/v0.1.0-alpha/groundingdino_swint_ogc.pth\n",
+ "!wget -P ./grounding_dino_weights https://raw.githubusercontent.com/IDEA-Research/GroundingDINO/refs/heads/main/groundingdino/config/GroundingDINO_SwinT_OGC.py"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "pvYX8RPyAJ2y"
+ },
+ "outputs": [],
+ "source": [
+ "# Download SAM2 weights\n",
+ "!mkdir sam2_weights\n",
+ "!wget -P ./sam2_weights https://dl.fbaipublicfiles.com/segment_anything_2/092824/sam2.1_hiera_large.pt"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "NoNvKYb3AO9-"
+ },
+ "outputs": [],
+ "source": [
+ "# download the sample image from the circularnet project\n",
+ "url = (\n",
+ " \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\n",
+ " \"projects/waste_identification_ml/pre_processing/config/sample_images/\"\n",
+ " \"IMG_6509.png\"\n",
+ ")\n",
+ "\n",
+ "!curl -O {url} \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "GPO5L6UXN6LO"
+ },
+ "source": [
+ "## Import libraries."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "bYtQDcmzAjI-"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "import supervision as sv\n",
+ "import torch\n",
+ "import tqdm\n",
+ "import numpy as np\n",
+ "from torchvision.ops import box_convert\n",
+ "from PIL import Image\n",
+ "from ollama import chat, ChatResponse\n",
+ "import glob\n",
+ "import cv2\n",
+ "import matplotlib.pyplot as plt\n",
+ "import math"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "C0XDBDs0SszU"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Utils\n",
+ "\n",
+ "def show_mask(\n",
+ " mask,\n",
+ " ax,\n",
+ " random_color=False,\n",
+ " borders = True\n",
+ "):\n",
+ " if random_color:\n",
+ " color = np.concatenate([np.random.random(3), np.array([0.6])], axis=0)\n",
+ " else:\n",
+ " color = np.array([30/255, 144/255, 255/255, 0.6])\n",
+ " h, w = mask.shape[-2:]\n",
+ " binary_mask = mask.astype(np.uint8)\n",
+ " mask_image = binary_mask.reshape(h, w, 1) * color.reshape(1, 1, -1)\n",
+ " if borders:\n",
+ " contours, _ = cv2.findContours(binary_mask,cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n",
+ " # Try to smooth contours\n",
+ " contours = [cv2.approxPolyDP(contour, epsilon=0.01, closed=True) for contour in contours]\n",
+ " mask_image = cv2.drawContours(mask_image, contours, -1, (1, 1, 1, 0.5), thickness=2)\n",
+ " ax.imshow(mask_image)\n",
+ "\n",
+ "\n",
+ "def show_points(\n",
+ " coords,\n",
+ " labels,\n",
+ " ax,\n",
+ " marker_size=375\n",
+ "):\n",
+ " pos_points = coords[labels==1]\n",
+ " neg_points = coords[labels==0]\n",
+ " ax.scatter(pos_points[:, 0], pos_points[:, 1], color='green', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)\n",
+ " ax.scatter(neg_points[:, 0], neg_points[:, 1], color='red', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)\n",
+ "\n",
+ "\n",
+ "def show_box(box, ax):\n",
+ " x0, y0 = box[0], box[1]\n",
+ " w, h = box[2] - box[0], box[3] - box[1]\n",
+ " ax.add_patch(plt.Rectangle((x0, y0), w, h, edgecolor='green', facecolor=(0, 0, 0, 0), lw=2))\n",
+ "\n",
+ "\n",
+ "def show_masks(\n",
+ " image,\n",
+ " masks,\n",
+ " scores,\n",
+ " point_coords=None,\n",
+ " box_coords=None,\n",
+ " input_labels=None,\n",
+ " borders=True\n",
+ "):\n",
+ " for i, (mask, score) in enumerate(zip(masks, scores)):\n",
+ " plt.figure(figsize=(10, 10))\n",
+ " plt.imshow(image)\n",
+ " show_mask(mask, plt.gca(), borders=borders)\n",
+ " if point_coords is not None:\n",
+ " assert input_labels is not None\n",
+ " show_points(point_coords, input_labels, plt.gca())\n",
+ " if box_coords is not None:\n",
+ " # boxes\n",
+ " show_box(box_coords, plt.gca())\n",
+ " if len(scores) \u003e 1:\n",
+ " plt.title(f\"Mask {i+1}, Score: {score:.3f}\", fontsize=18)\n",
+ " plt.axis('off')\n",
+ " plt.show()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "h3Opdq9nOYkj"
+ },
+ "source": [
+ "## Load models."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JcGgDIar-u_S"
+ },
+ "outputs": [],
+ "source": [
+ "# Load Grounding Dino model.\n",
+ "from grounding_dino.groundingdino.util.inference import load_model, load_image, predict, annotate\n",
+ "\n",
+ "# Path to the pre-trained Grounding Dino model checkpoint\n",
+ "WEIGHTS_PATH = \"grounding_dino_weights/groundingdino_swint_ogc.pth\"\n",
+ "\n",
+ "# Path to the configuration file for the Grounding Dino model variant being used\n",
+ "CONFIG_PATH = \"grounding_dino_weights/GroundingDINO_SwinT_OGC.py\"\n",
+ "\n",
+ "model = load_model(CONFIG_PATH, WEIGHTS_PATH)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "zYIaWLD2OjKD"
+ },
+ "outputs": [],
+ "source": [
+ "# Load SAM2 model.\n",
+ "from sam2.build_sam import build_sam2\n",
+ "from sam2.sam2_image_predictor import SAM2ImagePredictor\n",
+ "\n",
+ "# Path to the pre-trained SAM2 model checkpoint\n",
+ "sam2_checkpoint = \"sam2_weights/sam2.1_hiera_large.pt\"\n",
+ "\n",
+ "# Path to the configuration file for the SAM2 model variant being used\n",
+ "model_cfg = \"configs/sam2.1/sam2.1_hiera_l.yaml\"\n",
+ "\n",
+ "# Build the SAM2 model using the config and checkpoint; `device` should be set to \"cuda\" or \"cpu\"\n",
+ "sam2_model = build_sam2(model_cfg, sam2_checkpoint, device=torch.device(\"cuda\"))\n",
+ "\n",
+ "# Create a predictor object using the loaded SAM2 model for image-based mask prediction\n",
+ "sam2_predictor = SAM2ImagePredictor(sam2_model)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LkTUJzhTRCy6"
+ },
+ "source": [
+ "## Inference"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "B6b5Z7-XAnkX"
+ },
+ "outputs": [],
+ "source": [
+ "# Inference via Grounding Dino\n",
+ "%%time\n",
+ "IMAGE_PATH = \"IMG_6509.png\"\n",
+ "TEXT_PROMPT = \"packet\"\n",
+ "BOX_TRESHOLD = 0.25\n",
+ "TEXT_TRESHOLD = 0.25\n",
+ "\n",
+ "image_source, image = load_image(IMAGE_PATH)\n",
+ "\n",
+ "boxes, logits, phrases = predict(\n",
+ " model=model,\n",
+ " image=image,\n",
+ " caption=TEXT_PROMPT,\n",
+ " box_threshold=BOX_TRESHOLD,\n",
+ " text_threshold=TEXT_TRESHOLD\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "fLeaXwJKLaKw"
+ },
+ "outputs": [],
+ "source": [
+ "# Visualize Grounding Dino results.\n",
+ "annotated_frame = annotate(image_source=image_source, boxes=boxes, logits=logits, phrases=phrases)\n",
+ "\n",
+ "%matplotlib inline\n",
+ "sv.plot_image(annotated_frame, (16, 16))"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "shhCTs2tMKqy"
+ },
+ "outputs": [],
+ "source": [
+ "# Perform segmentation on bbox cordinates using SAM2 model.\n",
+ "sam2_predictor.set_image(image_source)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "b5L_XkxCQ5eD"
+ },
+ "outputs": [],
+ "source": [
+ "# Create a directory to store the cropped object images.\n",
+ "os.makedirs('tempdir', exist_ok=True)\n",
+ "\n",
+ "# Convert bbox format\n",
+ "h, w, _ = image_source.shape\n",
+ "boxes = boxes * torch.Tensor([w, h, w, h])\n",
+ "xyxy = box_convert(boxes=boxes, in_fmt=\"cxcywh\", out_fmt=\"xyxy\").numpy().astype(int)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "4GwSB3R8fd5y"
+ },
+ "outputs": [],
+ "source": [
+ "for idx, bbox in tqdm.tqdm(enumerate(xyxy)):\n",
+ " x1, y1, x2, y2 = bbox\n",
+ "\n",
+ " if (x2-x1)*(y2-y1) \u003c 0.25 * math.prod(image.size):\n",
+ " masks, scores, _ = sam2_predictor.predict(\n",
+ " point_coords=None,\n",
+ " point_labels=None,\n",
+ " box=bbox[None, :],\n",
+ " multimask_output=False,\n",
+ " )\n",
+ "\n",
+ " # show_masks(image_source, masks, scores, box_coords=bbox)\n",
+ "\n",
+ " # Convert the first mask to 0-255 and expand its dimensions to match the image channels.\n",
+ " # Multiply the mask with the original image (preserves object, sets background to 0).\n",
+ " # Crop the masked image to the bounding box [y1:y2, x1:x2].\n",
+ " masked_object = Image.fromarray(\n",
+ " np.where(\n",
+ " np.expand_dims(masks[0]*255, -1),\n",
+ " image_source, 0\n",
+ " )[y1:y2, x1:x2]\n",
+ " )\n",
+ "\n",
+ " image_path = f'tempdir/{os.path.splitext(IMAGE_PATH)[0]}_{idx}.png'\n",
+ " masked_object.save(image_path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "5NHKvcPGXk7n"
+ },
+ "source": [
+ "## Download Gemma3 model using Ollama tool."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "xCiuCSZoXb-M"
+ },
+ "source": [
+ "Run the following commands in the terminal within your colab notebook.\n",
+ "\n",
+ "```\n",
+ "curl https://ollama.ai/install.sh | sh\n",
+ "ollama serve\n",
+ "```\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "MbX_0DHISjAY"
+ },
+ "outputs": [],
+ "source": [
+ "# Pull the required open sourced LLM model.\n",
+ "!ollama pull gemma3:12b-it-qat"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "H-qOl2eYZbvg"
+ },
+ "outputs": [],
+ "source": [
+ "# Check if the model is downloaded.\n",
+ "!ollama list"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "G1hKcs1yZdqZ"
+ },
+ "outputs": [],
+ "source": [
+ "# Prompt to analyze an image for milk packet vs others.\n",
+ "prompt = \"\"\"\n",
+ "Analyze the provided image of packaging. Was this packaging used to contain milk or a milk-based product? Answer in yes or no only.\n",
+ "\"\"\""
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iF5p-hECZfKA"
+ },
+ "outputs": [],
+ "source": [
+ "# Read an cropped images to perform inference using LLM.\n",
+ "images = glob.glob('tempdir/*.png')\n",
+ "\n",
+ "for path in images:\n",
+ " # Run the chat/inference API, sending the temporary masked object image as input.\n",
+ " response: ChatResponse = chat(model='gemma3:12b-it-qat', messages=[\n",
+ " {\n",
+ " 'role': 'user',\n",
+ " 'content': prompt,\n",
+ " 'images': [path]\n",
+ " },\n",
+ " ])\n",
+ " image = cv2.imread(path)\n",
+ " image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n",
+ " plt.imshow(image)\n",
+ " plt.axis('off')\n",
+ " plt.show()\n",
+ "\n",
+ " # Print the model's response content (the generated answer)\n",
+ " print(f\"\\n{response.message.content}\")"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_groundingdino_sam2_gemini.ipynb b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_groundingdino_sam2_gemini.ipynb
new file mode 100644
index 00000000000..a7022112206
--- /dev/null
+++ b/official/projects/waste_identification_ml/llm_applications/milk_pouch_detection_using_groundingdino_sam2_gemini.ipynb
@@ -0,0 +1,520 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZAppGMCpM9Qq"
+ },
+ "source": [
+ "# Automatic Mask Generation Using Unsupervised Approach with Grounding Dino, SAM2, and Gemini"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cwNkr2MGNGQ3"
+ },
+ "source": [
+ "In this notebook, we build an end-to-end unsupervised pipeline for object detection, segmentation and classification on identifying and following milk pouches without manual labels. This approach leverages cutting-edge vision and language models on extracted features from segmentation masks.\n",
+ "\n",
+ "Key Components:\n",
+ "\n",
+ "1. **Grounding Dino**\n",
+ "\n",
+ "A powerful vision-language model that performs generic object detection by returning bounding boxes around visually significant regions—completely label-free and prompt-driven.\n",
+ "\n",
+ "2. **SAM2 (Segment Anything Model v2)**\n",
+ "\n",
+ "Using the bounding boxes from Grounding Dino, SAM2 generates precise segmentation masks, enabling instance-level understanding and clean extraction of objects.\n",
+ "\n",
+ "3. **Gemini-pro**\n",
+ "\n",
+ "Each cropped masked region is passed to Geimi-pro model to determine whether it contains a milk pouch or not, enabling robust classification without explicit supervised training.\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "S8JMep1cNw9K"
+ },
+ "source": [
+ "## Install necessary packages.\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "dHhKN6bfylwZ"
+ },
+ "outputs": [],
+ "source": [
+ "!git clone 'https://github.com/IDEA-Research/Grounded-SAM-2'\n",
+ "!pip install 'git+https://github.com/IDEA-Research/Grounded-SAM-2'\n",
+ "\n",
+ "%cd 'Grounded-SAM-2'\n",
+ "\n",
+ "# Install Grounding Dino\n",
+ "!pip install --no-build-isolation -e grounding_dino\n",
+ "\n",
+ "!pip install addict yapf supervision\u003e=0.22.0"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "yuJRa2S1N2CH"
+ },
+ "source": [
+ "## Import model weights and configuration files."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iZhkdMt68zPh"
+ },
+ "outputs": [],
+ "source": [
+ "# Download Grounding Dino weights.\n",
+ "!mkdir grounding_dino_weights\n",
+ "!wget -P ./grounding_dino_weights https://github.com/IDEA-Research/GroundingDINO/releases/download/v0.1.0-alpha/groundingdino_swint_ogc.pth\n",
+ "!wget -P ./grounding_dino_weights https://raw.githubusercontent.com/IDEA-Research/GroundingDINO/refs/heads/main/groundingdino/config/GroundingDINO_SwinT_OGC.py"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "7C4ApEvE_2Mq"
+ },
+ "outputs": [],
+ "source": [
+ "# Download SAM2 weights\n",
+ "!mkdir sam2_weights\n",
+ "!wget -P ./sam2_weights https://dl.fbaipublicfiles.com/segment_anything_2/092824/sam2.1_hiera_large.pt"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "NoNvKYb3AO9-"
+ },
+ "outputs": [],
+ "source": [
+ "# download the sample image from the circularnet project\n",
+ "url = (\n",
+ " \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\n",
+ " \"projects/waste_identification_ml/pre_processing/config/sample_images/\"\n",
+ " \"IMG_6509.png\"\n",
+ ")\n",
+ "\n",
+ "!curl -O {url} \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "GPO5L6UXN6LO"
+ },
+ "source": [
+ "## Import libraries."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "bYtQDcmzAjI-"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "import supervision as sv\n",
+ "import torch\n",
+ "import tqdm\n",
+ "import numpy as np\n",
+ "from torchvision.ops import box_convert\n",
+ "from PIL import Image\n",
+ "import glob\n",
+ "import cv2\n",
+ "import matplotlib.pyplot as plt\n",
+ "import math\n",
+ "import sys\n",
+ "from google.colab import auth\n",
+ "import shutil\n",
+ "\n",
+ "from google import genai\n",
+ "from google.genai import types\n",
+ "import requests"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "C0XDBDs0SszU"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Utils\n",
+ "\n",
+ "def show_mask(\n",
+ " mask,\n",
+ " ax,\n",
+ " random_color=False,\n",
+ " borders = True\n",
+ "):\n",
+ " if random_color:\n",
+ " color = np.concatenate([np.random.random(3), np.array([0.6])], axis=0)\n",
+ " else:\n",
+ " color = np.array([30/255, 144/255, 255/255, 0.6])\n",
+ " h, w = mask.shape[-2:]\n",
+ " binary_mask = mask.astype(np.uint8)\n",
+ " mask_image = binary_mask.reshape(h, w, 1) * color.reshape(1, 1, -1)\n",
+ " if borders:\n",
+ " contours, _ = cv2.findContours(binary_mask,cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n",
+ " # Try to smooth contours\n",
+ " contours = [cv2.approxPolyDP(contour, epsilon=0.01, closed=True) for contour in contours]\n",
+ " mask_image = cv2.drawContours(mask_image, contours, -1, (1, 1, 1, 0.5), thickness=2)\n",
+ " ax.imshow(mask_image)\n",
+ "\n",
+ "\n",
+ "def show_points(\n",
+ " coords,\n",
+ " labels,\n",
+ " ax,\n",
+ " marker_size=375\n",
+ "):\n",
+ " pos_points = coords[labels==1]\n",
+ " neg_points = coords[labels==0]\n",
+ " ax.scatter(pos_points[:, 0], pos_points[:, 1], color='green', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)\n",
+ " ax.scatter(neg_points[:, 0], neg_points[:, 1], color='red', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)\n",
+ "\n",
+ "\n",
+ "def show_box(box, ax):\n",
+ " x0, y0 = box[0], box[1]\n",
+ " w, h = box[2] - box[0], box[3] - box[1]\n",
+ " ax.add_patch(plt.Rectangle((x0, y0), w, h, edgecolor='green', facecolor=(0, 0, 0, 0), lw=2))\n",
+ "\n",
+ "\n",
+ "def show_masks(\n",
+ " image,\n",
+ " masks,\n",
+ " scores,\n",
+ " point_coords=None,\n",
+ " box_coords=None,\n",
+ " input_labels=None,\n",
+ " borders=True\n",
+ "):\n",
+ " for i, (mask, score) in enumerate(zip(masks, scores)):\n",
+ " plt.figure(figsize=(10, 10))\n",
+ " plt.imshow(image)\n",
+ " show_mask(mask, plt.gca(), borders=borders)\n",
+ " if point_coords is not None:\n",
+ " assert input_labels is not None\n",
+ " show_points(point_coords, input_labels, plt.gca())\n",
+ " if box_coords is not None:\n",
+ " # boxes\n",
+ " show_box(box_coords, plt.gca())\n",
+ " if len(scores) \u003e 1:\n",
+ " plt.title(f\"Mask {i+1}, Score: {score:.3f}\", fontsize=18)\n",
+ " plt.axis('off')\n",
+ " plt.show()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "h3Opdq9nOYkj"
+ },
+ "source": [
+ "## Load models."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JcGgDIar-u_S"
+ },
+ "outputs": [],
+ "source": [
+ "# Load Grounding Dino model.\n",
+ "from grounding_dino.groundingdino.util.inference import load_model, load_image, predict, annotate\n",
+ "\n",
+ "# Path to the pre-trained Grounding Dino model checkpoint\n",
+ "WEIGHTS_PATH = \"grounding_dino_weights/groundingdino_swint_ogc.pth\"\n",
+ "\n",
+ "# Path to the configuration file for the Grounding Dino model variant being used\n",
+ "CONFIG_PATH = \"grounding_dino_weights/GroundingDINO_SwinT_OGC.py\"\n",
+ "\n",
+ "model = load_model(CONFIG_PATH, WEIGHTS_PATH)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "zYIaWLD2OjKD"
+ },
+ "outputs": [],
+ "source": [
+ "# Load SAM2 model.\n",
+ "from sam2.build_sam import build_sam2\n",
+ "from sam2.sam2_image_predictor import SAM2ImagePredictor\n",
+ "\n",
+ "# Path to the pre-trained SAM2 model checkpoint\n",
+ "sam2_checkpoint = \"sam2_weights/sam2.1_hiera_large.pt\"\n",
+ "\n",
+ "# Path to the configuration file for the SAM2 model variant being used\n",
+ "model_cfg = \"configs/sam2.1/sam2.1_hiera_l.yaml\"\n",
+ "\n",
+ "# Build the SAM2 model using the config and checkpoint; `device` should be set to \"cuda\" or \"cpu\"\n",
+ "sam2_model = build_sam2(model_cfg, sam2_checkpoint, device=torch.device(\"cuda\"))\n",
+ "\n",
+ "# Create a predictor object using the loaded SAM2 model for image-based mask prediction\n",
+ "sam2_predictor = SAM2ImagePredictor(sam2_model)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LkTUJzhTRCy6"
+ },
+ "source": [
+ "## Inference"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "B6b5Z7-XAnkX"
+ },
+ "outputs": [],
+ "source": [
+ "# Inference via Grounding Dino\n",
+ "%%time\n",
+ "IMAGE_PATH = \"IMG_6509.png\"\n",
+ "TEXT_PROMPT = \"packet\"\n",
+ "BOX_TRESHOLD = 0.25\n",
+ "TEXT_TRESHOLD = 0.25\n",
+ "\n",
+ "image_source, image = load_image(IMAGE_PATH)\n",
+ "\n",
+ "boxes, logits, phrases = predict(\n",
+ " model=model,\n",
+ " image=image,\n",
+ " caption=TEXT_PROMPT,\n",
+ " box_threshold=BOX_TRESHOLD,\n",
+ " text_threshold=TEXT_TRESHOLD\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "fLeaXwJKLaKw"
+ },
+ "outputs": [],
+ "source": [
+ "# Visualize Grounding Dino results.\n",
+ "annotated_frame = annotate(image_source=image_source, boxes=boxes, logits=logits, phrases=phrases)\n",
+ "\n",
+ "%matplotlib inline\n",
+ "sv.plot_image(annotated_frame, (16, 16))"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "shhCTs2tMKqy"
+ },
+ "outputs": [],
+ "source": [
+ "# Perform segmentation on bbox cordinates using SAM2 model.\n",
+ "sam2_predictor.set_image(image_source)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "dt8Fs_whSx0r"
+ },
+ "outputs": [],
+ "source": [
+ "if os.path.exists('tempdir'):\n",
+ " shutil.rmtree('tempdir')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "b5L_XkxCQ5eD"
+ },
+ "outputs": [],
+ "source": [
+ "# Create a directory to store the cropped object images.\n",
+ "os.makedirs('tempdir', exist_ok=True)\n",
+ "\n",
+ "# Convert bbox format\n",
+ "h, w, _ = image_source.shape\n",
+ "boxes = boxes * torch.Tensor([w, h, w, h])\n",
+ "xyxy = box_convert(boxes=boxes, in_fmt=\"cxcywh\", out_fmt=\"xyxy\").numpy().astype(int)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "4GwSB3R8fd5y"
+ },
+ "outputs": [],
+ "source": [
+ "for idx, bbox in tqdm.tqdm(enumerate(xyxy)):\n",
+ " x1, y1, x2, y2 = bbox\n",
+ "\n",
+ " if (x2-x1)*(y2-y1) \u003c 0.25 * math.prod(image_source.shape[:2]):\n",
+ " masks, scores, _ = sam2_predictor.predict(\n",
+ " point_coords=None,\n",
+ " point_labels=None,\n",
+ " box=bbox[None, :],\n",
+ " multimask_output=False,\n",
+ " )\n",
+ "\n",
+ " # show_masks(image_source, masks, scores, box_coords=bbox)\n",
+ "\n",
+ " # Convert the first mask to 0-255 and expand its dimensions to match the image channels.\n",
+ " # Multiply the mask with the original image (preserves object, sets background to 0).\n",
+ " # Crop the masked image to the bounding box [y1:y2, x1:x2].\n",
+ " masked_object = Image.fromarray(\n",
+ " np.where(\n",
+ " np.expand_dims(masks[0]*255, -1),\n",
+ " image_source, 0\n",
+ " )[y1:y2, x1:x2]\n",
+ " )\n",
+ "\n",
+ " image_path = f'tempdir/{os.path.splitext(os.path.basename(IMAGE_PATH))[0]}_{idx}.png'\n",
+ " masked_object.save(image_path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "5NHKvcPGXk7n"
+ },
+ "source": [
+ "## Gemini parameter setting."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "j5yhKZvW9xUT"
+ },
+ "outputs": [],
+ "source": [
+ "# Authenticate colab notebook.\n",
+ "if \"google.colab\" in sys.modules:\n",
+ " auth.authenticate_user()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "sLwjN_NT94Uz"
+ },
+ "outputs": [],
+ "source": [
+ "MODEL_ID = \"gemini-2.5-pro\" # @param {type: \"string\", placeholder: \"[your-model-id]\", isTemplate: true}\n",
+ "PROJECT_ID = \"waste-identification-ml-330916\" # @param {type: \"string\", placeholder: \"[your-project-id]\", isTemplate: true}\n",
+ "LOCATION = os.environ.get(\"GOOGLE_CLOUD_REGION\", \"us-central1\")\n",
+ "\n",
+ "client = genai.Client(vertexai=True, project=PROJECT_ID, location=LOCATION)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iVtGB8KL0CEB"
+ },
+ "outputs": [],
+ "source": [
+ "prompt = \"\"\"\n",
+ "Your task is to determine if the packaging in the image was for a milk or milk-based product. Please follow these steps in order:\n",
+ "\n",
+ "1. **Primary Analysis:** Carefully examine the image for any visible text, logos, or brand names. This is the most important evidence. Base your conclusion on this information if it is available.\n",
+ "2. **Secondary Analysis:** If, and only if, there is no readable text or clear branding, consider other visual cues. A package that is a solid, bright color is less likely to be a dairy product and could be something like cooking oil.\n",
+ "\n",
+ "Based on your full analysis, was the packaging used for a milk or milk-based product? Answer with 'yes' or 'no' only.\n",
+ "\"\"\""
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "iF5p-hECZfKA"
+ },
+ "outputs": [],
+ "source": [
+ "# Read cropped images to perform inference using LLM.\n",
+ "images = glob.glob('tempdir/*.png')\n",
+ "\n",
+ "for path in images:\n",
+ " with open(path, 'rb') as f:\n",
+ " image_bytes = f.read()\n",
+ "\n",
+ " response = client.models.generate_content(\n",
+ " model=MODEL_ID,\n",
+ " contents=[\n",
+ " types.Part.from_bytes(\n",
+ " data=image_bytes,\n",
+ " mime_type='image/png',\n",
+ " ),\n",
+ " prompt\n",
+ " ],\n",
+ " config = types.GenerateContentConfig(\n",
+ " temperature=0.1,\n",
+ " )\n",
+ " )\n",
+ "\n",
+ " image = cv2.imread(path)\n",
+ " image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n",
+ " plt.imshow(image)\n",
+ " plt.axis('off')\n",
+ " plt.show()\n",
+ "\n",
+ " # Print the model's response content (the generated answer)\n",
+ " print(f\"\\n{response.text}\")"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_conversion/checkpoints_to_savedModel_to_tflite.ipynb b/official/projects/waste_identification_ml/model_conversion/checkpoints_to_savedModel_to_tflite.ipynb
new file mode 100644
index 00000000000..9cf5417dd6b
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_conversion/checkpoints_to_savedModel_to_tflite.ipynb
@@ -0,0 +1,346 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "wm0ezXfhdp2P"
+ },
+ "source": [
+ "# Convert Tensorflow model checkpoints to saved model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "KAy1ciItdrDB"
+ },
+ "source": [
+ "Given the checkpoints exported from Tensorflow model training, our goal is to convert those checkpoints into saved model for inference purpose.\u003cbr\u003e\n",
+ "Checkpoints is a binary file which contains all the values of the weights, biases, gradients and all the other variables saved. This file has an extension .ckpt. Checkpoints do not contain any description of the computation defined by the model and thus are typically only useful when source code that will use the saved parameter values is available.\u003cbr\u003e\n",
+ "A saved model contains a complete tensorflow program, including trained parameters and computation. It does not require the original model building code to run, which makes it useful for sharing or deploying with TFLite, tensorflow.js, Tensorflow Serving or Tensorflow Hub."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Lg4uH43Qds0q"
+ },
+ "source": [
+ "**Note** - We also assume that the script will be used as a Google Colab notebook. But this can be changed according to the needs of users. They can modify this in case they are working on their local workstation, remote server or any other database. This colab notebook can be changed to a regular jupyter notebook running on a local machine according to the need of the users."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "BofYg406d4LV"
+ },
+ "source": [
+ "## Import libraries \u0026 clone the TF model directory"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "oegDgL7yaaAq"
+ },
+ "outputs": [],
+ "source": [
+ "# install model-garden official and RESTART RUNTIME of the colab\n",
+ "!pip install tf-models-official"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "fMROz-xXdx6c"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "from google.colab import drive\n",
+ "import yaml\n",
+ "import tensorflow as tf"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 1197,
+ "status": "ok",
+ "timestamp": 1659384634603,
+ "user": {
+ "displayName": "Umair Sabir",
+ "userId": "06940594206388957365"
+ },
+ "user_tz": 420
+ },
+ "id": "gz1ajpHgeAJT",
+ "outputId": "1187e44e-82eb-4be1-8adc-1b50f6d7d0ed"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Drive already mounted at /content/gdrive; to attempt to forcibly remount, call drive.mount(\"/content/gdrive\", force_remount=True).\n",
+ "ln: failed to create symbolic link '/mydrive/My Drive': File exists\n",
+ "Successful\n"
+ ]
+ }
+ ],
+ "source": [
+ "# use this if your model and data are stored in the google drive\n",
+ "drive.mount('/content/gdrive')\n",
+ "\n",
+ "try:\n",
+ " !ln -s /content/gdrive/My\\ Drive/ /mydrive\n",
+ " print('Successful')\n",
+ "except Exception as e:\n",
+ " print(e)\n",
+ " print('Not successful')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 290,
+ "status": "ok",
+ "timestamp": 1659384648663,
+ "user": {
+ "displayName": "Umair Sabir",
+ "userId": "06940594206388957365"
+ },
+ "user_tz": 420
+ },
+ "id": "rRGalo90e2my",
+ "outputId": "f0305c8b-cf06-4637-a83c-05f5471313e7"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "fatal: destination path 'models' already exists and is not an empty directory.\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Clone the tensorflow models repository\n",
+ "!git clone --depth 1 https://github.com/tensorflow/models"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 189,
+ "status": "ok",
+ "timestamp": 1659384681521,
+ "user": {
+ "displayName": "Umair Sabir",
+ "userId": "06940594206388957365"
+ },
+ "user_tz": 420
+ },
+ "id": "HalXsX7BqdyX",
+ "outputId": "cb5555e4-0b77-4036-9230-1c01fcf1afaf"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "/content/models\n"
+ ]
+ }
+ ],
+ "source": [
+ "%cd models"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "8o7QzpHOsFHS"
+ },
+ "source": [
+ "## **MUST CHANGE** - Define the parameters according to your need"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "xak5WwMXppDF"
+ },
+ "outputs": [],
+ "source": [
+ "# this parameter depends on the backbone you will be using. In our \n",
+ "# case we used resnet backbone\n",
+ "EXPERIMENT_TYPE = 'maskrcnn_resnetfpn_coco' #@param {type:\"string\"}\n",
+ "\n",
+ "# path to the folder where all the files and checkpoints after model training \n",
+ "# are exported to\n",
+ "CHECKPOINT_PATH = '/mydrive/plastics_model/version_1/' #@param {type:\"string\"}\n",
+ "\n",
+ "# path where the saved model will be exported to\n",
+ "EXPORT_DIR_PATH = '/mydrive/plastics_model/experiment/' #@param {type:\"string\"}\n",
+ "\n",
+ "# config files are always stored with the checkpoints\n",
+ "CONFIG_FILE= CHECKPOINT_PATH + 'params.yaml'"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ddZzH5KhqSAy"
+ },
+ "outputs": [],
+ "source": [
+ "# config files are always stored with the checkpoints\n",
+ "# read the params.yaml file in order to get the height and width of an image\n",
+ "with open(CONFIG_FILE) as f:\n",
+ " my_dict = yaml.safe_load(f)\n",
+ "\n",
+ "HEIGHT = my_dict['task']['model']['input_size'][0]\n",
+ "WIDTH = my_dict['task']['model']['input_size'][1]\n",
+ "print(HEIGHT, WIDTH)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZWrbqHqJt947"
+ },
+ "source": [
+ "## calling the function to convert checkpoints to saved model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "k3wuYiVte4t_"
+ },
+ "outputs": [],
+ "source": [
+ "# run the conversion command\n",
+ "!python -m official.vision.serving.export_saved_model --experiment=$EXPERIMENT_TYPE \\\n",
+ " --export_dir=$EXPORT_DIR_PATH \\\n",
+ " --checkpoint_path=$CHECKPOINT_PATH \\\n",
+ " --batch_size=1 \\\n",
+ " --input_image_size=$HEIGHT,$WIDTH \\\n",
+ " --input_type=tflite \\\n",
+ " --config_file=$CONFIG_FILE"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "jQb13BbW78bs"
+ },
+ "source": [
+ "# Convert saved model to TF Lite model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "QEKaddZ58Ab6"
+ },
+ "source": [
+ "Given the saved model after Tensorflow model training, our goal is to convert saved model to TFLite for inference purpose on edge devices. "
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZXpSX-w_8A1c"
+ },
+ "source": [
+ "Tensorflow Lite is a set of tools that enables on-device machine learning by helping developers run their models on mobile, embedded and edge devices."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "20sV-bx59GRD"
+ },
+ "source": [
+ "## **MUST CHANGE** - Define the parameters according to your need"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "icoDtIin9REv"
+ },
+ "outputs": [],
+ "source": [
+ "# path where the tflite model will be written with its name\n",
+ "TFLITE_PATH = '/mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/model_version_1/plastics_model/tflite_fan/model.tflite' #@param {type:\"string\"}\n",
+ "\n",
+ "# path where saved model parameters are saved\n",
+ "SAVED_MODEL_DIR = EXPORT_DIR_PATH + '/saved_model/'"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "YE1Z13xs9NtJ"
+ },
+ "source": [
+ "## conversion of saved model to tflite"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "1LY9yoUP6Gr4"
+ },
+ "outputs": [],
+ "source": [
+ "converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir=SAVED_MODEL_DIR) \n",
+ "tflite_model = converter.convert() \n",
+ "with open(TFLITE_PATH, 'wb') as f:\n",
+ " f.write(tflite_model)"
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "collapsed_sections": [],
+ "name": "checkpoints_to_saved_model_to_tflite.ipynb",
+ "provenance": []
+ },
+ "gpuClass": "standard",
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_conversion/detectron2_inference_and_trt_conversion.ipynb b/official/projects/waste_identification_ml/model_conversion/detectron2_inference_and_trt_conversion.ipynb
new file mode 100644
index 00000000000..fd7d06caae3
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_conversion/detectron2_inference_and_trt_conversion.ipynb
@@ -0,0 +1,629 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZtzVZpUyxS5B"
+ },
+ "source": [
+ "# Waste identification with instance segmentation in Detectron2"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "MTrqWFRdxYyJ"
+ },
+ "source": [
+ "Welcome to the Instance Segmentation Colab! This notebook will take you through the steps of running an \"out-of-the-box\" Mask RCNN Instance Segmentation model on image from Detectron2."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "F0Gxj-CPxnFe"
+ },
+ "source": [
+ "To finish this task, a proper path for the model and a single image needs to be provided. The path to the labels on which the models are trained is in the waste_identification_ml directory inside the Tensorflow Model Garden repository. The label files are inferred automatically for the model."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LgJKrYHEyAEv"
+ },
+ "source": [
+ "## Clone and install Detectron2"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XXxxDT_QW2mR"
+ },
+ "outputs": [],
+ "source": [
+ "# Clone the Detectron2 repository and install the required packages.\n",
+ "# Relax as installing packages might take a while.\n",
+ "!git clone -q 'https://github.com/facebookresearch/detectron2'\n",
+ "!pip install -q 'git+https://github.com/facebookresearch/detectron2.git'"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WCWF8h6TuNqZ"
+ },
+ "outputs": [],
+ "source": [
+ "# Install supervision package for the postprocessing of output results\n",
+ "# from Detectron2 Mask RCNN model.\n",
+ "!pip install -q supervision"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "I1oMgVZj4n9n"
+ },
+ "source": [
+ "## Clone the TF Model Garden repo where the waste identification project is located"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "SHm3oznC4uBe"
+ },
+ "outputs": [],
+ "source": [
+ "!git clone --depth 1 https://github.com/tensorflow/models 2\u003e/dev/null"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ckQUcFtq1P3w"
+ },
+ "source": [
+ "## Imports and Setup"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "LdgIDnYe1dgv"
+ },
+ "outputs": [],
+ "source": [
+ "# Third-Party Imports\n",
+ "import csv\n",
+ "import cv2\n",
+ "# Detectron2 Imports\n",
+ "import detectron2\n",
+ "from detectron2.config import get_cfg\n",
+ "from detectron2.data.catalog import Metadata\n",
+ "from detectron2.engine import DefaultPredictor\n",
+ "from detectron2.structures import Boxes, Instances\n",
+ "from detectron2.utils.logger import setup_logger\n",
+ "from detectron2.utils.visualizer import Visualizer\n",
+ "import matplotlib.pyplot as plt\n",
+ "from PIL import Image\n",
+ "import supervision as sv\n",
+ "import torch\n",
+ "\n",
+ "# Setup Detectron2 Logger\n",
+ "setup_logger()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "R48IuhsN__vG"
+ },
+ "outputs": [],
+ "source": [
+ "# @title Utilities\n",
+ "\n",
+ "\n",
+ "def convert_detections_to_instances(\n",
+ " outputs: dict,\n",
+ " image_size: tuple[int, int] = (1024, 1024),\n",
+ " nms_threshold: float = 0.8,\n",
+ " class_agnostic: bool = True,\n",
+ ") -\u003e dict[str, Instances]:\n",
+ " \"\"\"Convert Detectron2 model outputs to an Instances object with Non-Maximum Suppression (NMS) applied.\n",
+ "\n",
+ " Args:\n",
+ " outputs: Detectron2 model output containing instance predictions.\n",
+ " image_size: Image dimensions (height, width).\n",
+ " nms_threshold: Non-Maximum Suppression (NMS) threshold.\n",
+ " class_agnostic: Whether NMS should be applied in a class-agnostic manner.\n",
+ "\n",
+ " Returns:\n",
+ " Reformatted Detectron2 output as {\"instances\": Instances}.\n",
+ " \"\"\"\n",
+ " # Apply NMS and convert to supervision Detections format\n",
+ " detections = sv.Detections.from_detectron2(outputs).with_nms(\n",
+ " threshold=nms_threshold, class_agnostic=class_agnostic\n",
+ " )\n",
+ "\n",
+ " # Convert extracted values to PyTorch tensors\n",
+ " bboxes = torch.tensor(detections.xyxy, dtype=torch.float32)\n",
+ " scores = torch.tensor(detections.confidence, dtype=torch.float32)\n",
+ " classes = torch.tensor(detections.class_id, dtype=torch.int64)\n",
+ "\n",
+ " # Create an Instances object\n",
+ " output_instances = Instances(image_size)\n",
+ " output_instances.set(\"pred_boxes\", Boxes(bboxes))\n",
+ " output_instances.set(\"scores\", scores)\n",
+ " output_instances.set(\"pred_classes\", classes)\n",
+ "\n",
+ " # Add masks if available\n",
+ " if detections.mask is not None:\n",
+ " masks = torch.tensor(detections.mask, dtype=torch.uint8)\n",
+ " output_instances.set(\"pred_masks\", masks)\n",
+ "\n",
+ " return {\"instances\": output_instances}\n",
+ "\n",
+ "\n",
+ "def read_csv(file_path: str) -\u003e list[str]:\n",
+ " \"\"\"Reads a CSV file and returns its contents as a list.\n",
+ "\n",
+ " This function reads the given CSV file, skips the header, and assumes\n",
+ " there is only one column in the CSV. It returns the contents as a list of\n",
+ " strings.\n",
+ "\n",
+ " Args:\n",
+ " file_path: The path to the CSV file.\n",
+ "\n",
+ " Returns:\n",
+ " The contents of the CSV file as a list of strings.\n",
+ " \"\"\"\n",
+ " data_list = []\n",
+ " with open(file_path, \"r\") as csvfile:\n",
+ " reader = csv.reader(csvfile)\n",
+ " for row in reader:\n",
+ " data_list.append(row[0])\n",
+ " return data_list"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "B8Q3y2r04e57"
+ },
+ "source": [
+ "## Import and load the labels."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "p2426m2L41Ny"
+ },
+ "outputs": [],
+ "source": [
+ "LABELS_PATH = (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/data/45_labels.csv'\n",
+ ")\n",
+ "\n",
+ "labels = read_csv(LABELS_PATH)\n",
+ "\n",
+ "my_metadata = Metadata()\n",
+ "my_metadata.set(thing_classes=labels)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "oWQkSqN_3aTP"
+ },
+ "source": [
+ "## Import Detectron2 Mask RCNN model."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "-P64IPCT3fsr"
+ },
+ "outputs": [],
+ "source": [
+ "%%bash\n",
+ "wget https://storage.googleapis.com/tf_model_garden/vision/\\\n",
+ "waste_identification_ml/Detectron2_Jan2025_1024_1024.zip\n",
+ "\n",
+ "unzip Detectron2_Jan2025_1024_1024.zip \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "IPM9Sk9sKgS_"
+ },
+ "source": [
+ "## Load the model and perform inference (Non-TRT)\n",
+ "\n",
+ "You will need to supply the input and output folders below. Among other options, you can use local files or connect to a google drive where you have images. See examples [here](https://colab.sandbox.google.com/notebooks/io.ipynb)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "0SLE4h9kgQkE"
+ },
+ "outputs": [],
+ "source": [
+ "# Initialize the Detectron2 configuration object\n",
+ "cfg = get_cfg()\n",
+ "\n",
+ "# Load the model configuration from a YAML file.\n",
+ "cfg.merge_from_file(\"config.yaml\")\n",
+ "\n",
+ "# Set the confidence threshold.\n",
+ "cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = 0.5\n",
+ "\n",
+ "# Specify the path to the trained model weights.\n",
+ "cfg.MODEL.WEIGHTS = \"model_final.pth\"\n",
+ "\n",
+ "# Create a predictor object using the configured model.\n",
+ "predictor = DefaultPredictor(cfg)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "D2l20agB7IHG"
+ },
+ "outputs": [],
+ "source": [
+ "HEIGHT = 1024\n",
+ "WIDTH = 1024\n",
+ "\n",
+ "import collections\n",
+ "import os\n",
+ "\n",
+ "INPUT_FOLDER = \"images\"\n",
+ "OUTPUT_FOLDER = \"images_output\"\n",
+ "LABEL_COUNTS = collections.defaultdict(int)\n",
+ "\n",
+ "\n",
+ "for filename in os.listdir(INPUT_FOLDER):\n",
+ " if filename.lower().endswith(\n",
+ " (\".jpg\", \".jpeg\", \".png\")\n",
+ " ): # Adjust extensions if needed\n",
+ " input_path = os.path.join(INPUT_FOLDER, filename)\n",
+ " img = cv2.imread(input_path)\n",
+ "\n",
+ " original_height, original_width = img.shape[:2]\n",
+ "\n",
+ " resized_image = cv2.resize(\n",
+ " img, (WIDTH, HEIGHT), interpolation=cv2.INTER_AREA\n",
+ " )\n",
+ "\n",
+ " outputs = predictor(resized_image)\n",
+ " outputs = convert_detections_to_instances(outputs)\n",
+ "\n",
+ " # Extract the predicted instances\n",
+ " instances = outputs[\"instances\"].to(\"cpu\")\n",
+ "\n",
+ " # Re-scale bounding boxes back to the original image size\n",
+ " scale_x = original_width / WIDTH\n",
+ " scale_y = original_height / HEIGHT\n",
+ " instances.pred_boxes.scale(scale_x, scale_y)\n",
+ "\n",
+ " # Resize masks to match the original image size\n",
+ " if instances.has(\"pred_masks\"):\n",
+ " pred_masks = instances.pred_masks.numpy() # Convert to NumPy array\n",
+ " resized_masks = []\n",
+ "\n",
+ " for mask in pred_masks:\n",
+ " resized_mask = cv2.resize(\n",
+ " mask.astype(\"uint8\"),\n",
+ " (original_width, original_height),\n",
+ " interpolation=cv2.INTER_NEAREST,\n",
+ " )\n",
+ " resized_masks.append(resized_mask)\n",
+ "\n",
+ " instances.pred_masks = torch.tensor(resized_masks, dtype=torch.uint8)\n",
+ "\n",
+ " # Initialize the visualizer with the original image\n",
+ " visualizer = Visualizer(\n",
+ " img_rgb=img, # Use the original image\n",
+ " metadata=my_metadata, # Metadata containing class labels, colors, etc.\n",
+ " scale=1, # Scale factor for visualization\n",
+ " )\n",
+ "\n",
+ " # Draw predictions on the original image\n",
+ " visualized_image = visualizer.draw_instance_predictions(\n",
+ " instances\n",
+ " ).get_image()\n",
+ "\n",
+ " output_path = os.path.join(OUTPUT_FOLDER, filename)\n",
+ " cv2.imwrite(output_path, visualized_image)\n",
+ " print(f\"Processed {filename} and saved to {output_path}\")\n",
+ "\n",
+ " # Count the labels\n",
+ " for label in instances.pred_classes:\n",
+ " label_name = labels[label.item()] # Get label name from index\n",
+ " LABEL_COUNTS[label_name] += 1 # Increment the count for this label\n",
+ "\n",
+ "# Print the label counts\n",
+ "for label_name, count in LABEL_COUNTS.items():\n",
+ " print(f\"{label_name}: {count}\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "dyqLE59oavQq"
+ },
+ "source": [
+ "## TensortRT Conversion and Prediction"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "X53zNeCAbfCl"
+ },
+ "outputs": [],
+ "source": [
+ "!pip install -q onnx onnxruntime common\n",
+ "!pip install -q git+https://github.com/NVIDIA/TensorRT#subdirectory=tools/onnx-graphsurgeon\n",
+ "!git clone -q https://github.com/NVIDIA/TensorRT.git #-b v8.6.1\n",
+ "!pip install -q TensorRT"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "_oDk4Mw0ex3O"
+ },
+ "source": [
+ "# Create TRT Sample image\n",
+ "\n",
+ "We need a 1344x1344 image to provide for calibration during conversion, see docs [here](https://github.com/NVIDIA/TensorRT/blob/main/samples/python/detectron2/README.md)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "3N8650bBccJH"
+ },
+ "outputs": [],
+ "source": [
+ "original_image = cv2.imread('sample_image.jpg')\n",
+ "original_height, original_width = original_image.shape[:2]\n",
+ "\n",
+ "trt_sample_image = cv2.resize(\n",
+ " original_image, (1344, 1344), interpolation=cv2.INTER_AREA\n",
+ ")\n",
+ "\n",
+ "cv2.imwrite('trt_sample_image.jpg', trt_sample_image)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "le0gRKx8iCnN"
+ },
+ "source": [
+ "# Modify export_model.py\n",
+ "\n",
+ "According to the docs [here](https://github.com/NVIDIA/TensorRT/blob/main/samples/python/detectron2/README.md#detectron-2-deployment), we need to modify the export_model.py script in detectron2/tools/deploy as follows:\n",
+ "\n",
+ "```\n",
+ "aug = T.ResizeShortestEdge(\n",
+ " [cfg.INPUT.MIN_SIZE_TEST, cfg.INPUT.MIN_SIZE_TEST], cfg.INPUT.MAX_SIZE_TEST\n",
+ ")\n",
+ "```\n",
+ "\n",
+ "--\u003e\n",
+ "\n",
+ "\n",
+ "```\n",
+ "aug = T.ResizeShortestEdge(\n",
+ " [1344, 1344], 1344\n",
+ ")\n",
+ "```\n",
+ "\n",
+ "This is to match the ultimate trt pipeline required resolution, and ensure proper calibration. Note that ResizeShortestEdge creates a scaling factor that is later used during image inference pre-processing, where infer.py ultimately creates a model training resolution image that it then scales and adds padding to."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "zlD1ZSRObf2s"
+ },
+ "outputs": [],
+ "source": [
+ "!python detectron2/tools/deploy/export_model.py \\\n",
+ " --sample-image trt_sample_image.jpg \\\n",
+ " --config-file config.yaml \\\n",
+ " --export-method tracing \\\n",
+ " --format onnx \\\n",
+ " --output ./ \\\n",
+ " MODEL.WEIGHTS model_final.pth \\\n",
+ " MODEL.DEVICE cuda"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nDA06Freb7gg"
+ },
+ "outputs": [],
+ "source": [
+ "!python TensorRT/samples/python/detectron2/create_onnx.py \\\n",
+ " --exported_onnx model.onnx \\\n",
+ " --onnx converted.onnx \\\n",
+ " --det2_config config.yaml \\\n",
+ " --det2_weights model_final.pth \\\n",
+ " --sample_image trt_sample_image.jpg"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "BEnrWyhscHDi"
+ },
+ "outputs": [],
+ "source": [
+ "!python3 TensorRT/samples/python/detectron2/build_engine.py \\\n",
+ "--onnx converted.onnx --engine engine32.trt --precision fp32"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "J84J4iXvkEws"
+ },
+ "source": [
+ "# Change the labels in the infer.py script to match your labels.csv\n",
+ "\n",
+ "The imported package /TensorRT/samples/python.detectron2/infer.py has a preset list of labels. These need to modified manually to match your labels.csv / metadata\n",
+ "\n",
+ "The input for inference below can be a file or a folder, and the expected output is image(s) with visualizations and an accompanying .txt file:\n",
+ "\n",
+ "[Inference in Python reference](https://github.com/NVIDIA/TensorRT/blob/main/samples/python/detectron2/README.md#inference-in-python)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "UipdRZyTDU-K"
+ },
+ "source": [
+ "# Run TRT Inference\n",
+ "\n",
+ "You can also take the TRT converted model here (engine32.trt) and use the script below to run inference from it.\n",
+ "\n",
+ "The infer.py script comes from\n",
+ "https://github.com/NVIDIA/TensorRT/blob/main/samples/python/detectron2/infer.py"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nedMCmTttSWJ"
+ },
+ "outputs": [],
+ "source": [
+ "!rm -rf tensorrt_predictions\n",
+ "!mkdir tensorrt_predictions\n",
+ "!python TensorRT/samples/python/detectron2/infer.py \\\n",
+ " --engine engine32.trt \\\n",
+ " --input images \\\n",
+ " --det2_config config.yaml \\\n",
+ " --output tensorrt_predictions"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "yd1azz14pSRk"
+ },
+ "source": [
+ "# Summarize the results (optional)\n",
+ "\n",
+ "This colab features trt conversion and inference as a demo, and is not designed to be productionized. However, you can still take a look at your results below by parsing the outputted text files."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "vHoGG2_tpTVE"
+ },
+ "outputs": [],
+ "source": [
+ "from collections import Counter\n",
+ "import os\n",
+ "\n",
+ "\n",
+ "def count_prediction_classes(folder_path):\n",
+ " \"\"\"Count the predicted classes in TensorRT prediction files.\n",
+ "\n",
+ " Args:\n",
+ " folder_path (str): Path to the folder containing .txt prediction files\n",
+ "\n",
+ " Returns:\n",
+ " Counter: Counts of each predicted class\n",
+ " \"\"\"\n",
+ " # Counter to store class counts\n",
+ " class_counts = Counter()\n",
+ "\n",
+ " # Loop through all txt files in the folder\n",
+ " for filename in os.listdir(folder_path):\n",
+ " if filename.endswith(\".txt\"):\n",
+ " file_path = os.path.join(folder_path, filename)\n",
+ "\n",
+ " # Read the file\n",
+ " try:\n",
+ " with open(file_path, \"r\") as f:\n",
+ " lines = f.readlines()\n",
+ "\n",
+ " # Extract class values (last column) from each line\n",
+ " for line in lines:\n",
+ " if line.strip(): # Skip empty lines\n",
+ " columns = line.strip().split()\n",
+ " if len(columns) \u003e= 6: # Ensure we have enough columns\n",
+ " class_value = int(\n",
+ " columns[5]\n",
+ " ) # The class value is the 6th column (index 5)\n",
+ " class_counts[class_value] += 1\n",
+ "\n",
+ " except Exception as e:\n",
+ " print(f\"Error processing file {filename}: {e}\")\n",
+ "\n",
+ " return class_counts\n",
+ "\n",
+ "\n",
+ "folder_path = \"tensorrt_predictions\"\n",
+ "\n",
+ "# Count classes\n",
+ "class_counts = count_prediction_classes(folder_path)\n",
+ "\n",
+ "# Print results\n",
+ "print(\"Class Counts:\")\n",
+ "for class_id, count in sorted(class_counts.items()):\n",
+ " print(f\"Class {class_id}: {count}\")"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "private_outputs": true,
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_conversion/savedmodel_to_tensorrt.ipynb b/official/projects/waste_identification_ml/model_conversion/savedmodel_to_tensorrt.ipynb
new file mode 100644
index 00000000000..3a622cee4dd
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_conversion/savedmodel_to_tensorrt.ipynb
@@ -0,0 +1,878 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "jFZKt6cqrBQ_"
+ },
+ "source": [
+ "# TensorRT Optimization for Mask R-CNN"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "lYcFFUPsrWwr"
+ },
+ "source": [
+ "## Overview"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "wwxSwqwKrXvf"
+ },
+ "source": [
+ "This notebook outlines the process and results of converting a TensorFlow saved model developed with Mask R-CNN architecture into an optimized TensorRT model. The primary objective of this conversion is to enhance the inference speed on both edge devices and cloud infrastructure, thereby facilitating real-time application requirements and scalable deployment scenarios."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "n2EV0Abhrac7"
+ },
+ "source": [
+ "## Background"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "y9ZOWXbOrcZC"
+ },
+ "source": [
+ "The Mask R-CNN model, renowned for its efficiency in instance segmentation tasks, was initially trained using a high-quality dataset to identify and segment objects within images. Although the model achieved a high accuracy, its inference time on standard hardware was a considerable bottleneck, taking approximately 35 seconds per image."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "txhQMEv5riDY"
+ },
+ "source": [
+ "## Objective"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hK2IyrLwrlg8"
+ },
+ "source": [
+ "To significantly reduce the inference time of the 2 Mask R-CNN model without compromising its accuracy, ensuring it meets the latency requirements of real-time applications. The model should be capable of delivering prompt predictions on edge devices with limited computational resources as well as on cloud platforms.\n",
+ "\n",
+ "The TensorFlow saved model was converted into a TensorRT model using Tensorflow library."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "kyYWvlyHsF1h"
+ },
+ "source": [
+ "**Note:\n",
+ "To execute this Colab notebook effectively, please ensure that you switch the runtime to utilize a GPU. Additionally, for optimal performance, select the 'High-RAM' option which is available under the 'Runtime' tab at the top of the Colab notebook interface. This configuration is essential for handling compute-intensive operations and large datasets without running into memory constraints.**"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "GQ_GovZfQe4p"
+ },
+ "source": [
+ "## Download required files \u0026 scripts."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "uK_Bk5pD6vwz"
+ },
+ "outputs": [],
+ "source": [
+ "!apt-get install tensorrt uff-converter-tf"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "v_nKwP0wPoOL"
+ },
+ "outputs": [],
+ "source": [
+ "# Download preprocessing script.\n",
+ "url = (\n",
+ " \"https://raw.githubusercontent.com/\"\n",
+ " \"tensorflow/models/master/\"\n",
+ " \"official/projects/waste_identification_ml/\"\n",
+ " \"model_inference/preprocessing.py\"\n",
+ ")\n",
+ "\n",
+ "!wget -q {url}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "VFoknp3ZsgUL"
+ },
+ "outputs": [],
+ "source": [
+ "# Download the script to pull instance segmentation model weights from the\n",
+ "# TF Model Garden repo.\n",
+ "url = (\n",
+ " \"https://raw.githubusercontent.com/\"\n",
+ " \"tensorflow/models/master/\"\n",
+ " \"official/projects/waste_identification_ml/\"\n",
+ " \"model_inference/download_and_unzip_models.py\"\n",
+ ")\n",
+ "\n",
+ "!wget -q {url}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "ArBZijNKtL8v",
+ "outputId": "6793decd-1644-4f75-d2f5-b4d206ddd6fb"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "100 3303k 100 3303k 0 0 2120k 0 0:00:01 0:00:01 --:--:-- 2120k\n",
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "100 1913k 100 1913k 0 0 942k 0 0:00:02 0:00:02 --:--:-- 943k\n"
+ ]
+ }
+ ],
+ "source": [
+ "# download the sample image from the circularnet project\n",
+ "url1 = (\n",
+ " \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\n",
+ " \"projects/waste_identification_ml/pre_processing/config/sample_images/\"\n",
+ " \"image_2.png\"\n",
+ ")\n",
+ "\n",
+ "url2 = (\n",
+ " \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\n",
+ " \"projects/waste_identification_ml/pre_processing/config/sample_images/\"\n",
+ " \"image_4.png\"\n",
+ ")\n",
+ "\n",
+ "!curl -O {url1}\n",
+ "!curl -O {url2}"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "olZx5rVUQo5c"
+ },
+ "source": [
+ "## Import required packages."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "McVjuDBCvdGb",
+ "outputId": "1b097eea-5467-4f05-84d7-e2890e0f4ef6"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ " Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.7/1.7 MB\u001b[0m \u001b[31m7.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h Building wheel for tensorrt (setup.py) ... \u001b[?25l\u001b[?25hdone\n"
+ ]
+ }
+ ],
+ "source": [
+ "!python3 -m pip install -q -U tensorrt tf_keras"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "WSX-u14IvY-6",
+ "outputId": "16377c07-6190-4b30-a83a-9775df0863a7"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "8.6.1\n"
+ ]
+ }
+ ],
+ "source": [
+ "import tensorrt\n",
+ "print(tensorrt.__version__)\n",
+ "assert tensorrt.Builder(tensorrt.Logger())"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "4RuDEk9WtonR"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "import sys\n",
+ "\n",
+ "from PIL import Image\n",
+ "import matplotlib.pyplot as plt\n",
+ "import numpy as np\n",
+ "import tensorflow as tf\n",
+ "from six import BytesIO\n",
+ "from six.moves.urllib.request import urlopen\n",
+ "from typing import Any, Callable\n",
+ "import preprocessing\n",
+ "\n",
+ "import logging\n",
+ "logging.disable(logging.WARNING)\n",
+ "\n",
+ "%matplotlib inline"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LOpbotPFOaFv"
+ },
+ "source": [
+ "## Utils"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "YSbWKz1Atliw"
+ },
+ "outputs": [],
+ "source": [
+ "def load_image_into_numpy_array(path: str) -\u003e np.ndarray:\n",
+ " \"\"\"Load an image from file into a numpy array.\n",
+ "\n",
+ " Puts image into numpy array to feed into tensorflow graph.\n",
+ " Note that by convention we put it into a numpy array with shape\n",
+ " (height, width, channels), where channels=3 for RGB.\n",
+ "\n",
+ " Args:\n",
+ " path: the file path to the image\n",
+ "\n",
+ " Returns:\n",
+ " uint8 numpy array with shape (1, h, w, 3)\n",
+ " \"\"\"\n",
+ " image = None\n",
+ " if(path.startswith('http')):\n",
+ " response = urlopen(path)\n",
+ " image_data = response.read()\n",
+ " image_data = BytesIO(image_data)\n",
+ " image = Image.open(image_data)\n",
+ " else:\n",
+ " image_data = tf.io.gfile.GFile(path, 'rb').read()\n",
+ " image = Image.open(BytesIO(image_data))\n",
+ "\n",
+ " (im_width, im_height) = image.size\n",
+ " return np.array(image.getdata()).reshape(\n",
+ " (1, im_height, im_width, 3)).astype(np.uint8)\n",
+ "\n",
+ "\n",
+ "def load_model(model_handle: str) -\u003e Callable:\n",
+ " \"\"\"Loads a TensorFlow SavedModel and returns a function that can be used to\n",
+ " make predictions.\n",
+ "\n",
+ " Args:\n",
+ " model_handle: A path to a TensorFlow SavedModel.\n",
+ "\n",
+ " Returns:\n",
+ " A function that can be used to make predictions.\n",
+ " \"\"\"\n",
+ " print('loading model...')\n",
+ " print(model_handle)\n",
+ " model = tf.saved_model.load(model_handle)\n",
+ " print('model loaded!')\n",
+ " detection_fn = model.signatures['serving_default']\n",
+ " return detection_fn\n",
+ "\n",
+ "\n",
+ "def perform_detection(model: Callable, image: np.ndarray) -\u003e dict[str, Any]:\n",
+ " \"\"\"Performs Mask RCNN on an image using the specified model.\n",
+ "\n",
+ " Args:\n",
+ " model: A function that can be used to make predictions.\n",
+ " image_np: A NumPy array representing the image to be detected.\n",
+ "\n",
+ " Returns:\n",
+ " A list of detections.\n",
+ " \"\"\"\n",
+ " detection_fn = model(image)\n",
+ " detection_fn = {key: value.numpy() for key, value in detection_fn.items()}\n",
+ " return detection_fn\n",
+ "\n",
+ "\n",
+ "def create_directory(path: str):\n",
+ " \"\"\"Create a directory at the specified path if it does not exist.\n",
+ "\n",
+ " Args:\n",
+ " path (str): The path of the directory to create.\n",
+ " \"\"\"\n",
+ " try:\n",
+ " os.makedirs(path, exist_ok=True)\n",
+ " print(f'Directory {path} created successfully')\n",
+ " except Exception as e:\n",
+ " print(f'Failed to create directory {path}: {e}')\n",
+ "\n",
+ "\n",
+ "def convert_to_tensorrt(\n",
+ " saved_model_dir: str,\n",
+ " output_saved_model_dir: str\n",
+ " ) -\u003e Callable:\n",
+ " \"\"\"\n",
+ " Converts a TensorFlow SavedModel to TensorRT format.\n",
+ "\n",
+ " Args:\n",
+ " saved_model_dir: The directory where the original TensorFlow SavedModel is\n",
+ " stored.\n",
+ " output_saved_model_dir: The directory where the TensorRT-converted model\n",
+ " will be saved.\n",
+ "\n",
+ " Returns:\n",
+ " Callable: A generator function that yields input data for building TRT\n",
+ " engines.\n",
+ " \"\"\"\n",
+ " params = tf.experimental.tensorrt.ConversionParams(\n",
+ " precision_mode='FP16',\n",
+ " # Set this to a large enough number so it can cache all the engines.\n",
+ " maximum_cached_engines=16\n",
+ " )\n",
+ "\n",
+ " converter = tf.experimental.tensorrt.Converter(\n",
+ " input_saved_model_dir=saved_model_dir, conversion_params=params\n",
+ " )\n",
+ "\n",
+ " converter.convert()\n",
+ "\n",
+ " # Define a generator function that yields input data, and use it to execute\n",
+ " # the graph to build TRT engines.\n",
+ " def my_input_fn():\n",
+ " yield image1\n",
+ "\n",
+ " converter.build(input_fn=my_input_fn) # Generate corresponding TRT engines\n",
+ " converter.save(output_saved_model_dir) # Generated engines will be saved.\n",
+ "\n",
+ "\n",
+ "def process_image(image_path: str) -\u003e tf.Tensor:\n",
+ " \"\"\"\n",
+ " Processes an image from a given file path.\n",
+ "\n",
+ " This function reads an image from the specified path, resizes it, and applies\n",
+ " normalization preprocessing.\n",
+ "\n",
+ " Args:\n",
+ " image_path: The file path of the image to be processed.\n",
+ "\n",
+ " Returns:\n",
+ " A TensorFlow Tensor representing the processed image.\n",
+ " \"\"\"\n",
+ " image_np = load_image_into_numpy_array(image_path)\n",
+ " image_np_cp = tf.image.resize(image_np[0], (512, 1024), method=tf.image.ResizeMethod.AREA)\n",
+ " image_np_cp = tf.cast(image_np_cp, tf.uint8)\n",
+ " image_np = preprocessing.normalize_image(image_np_cp)\n",
+ " image_np = tf.expand_dims(image_np, axis=0)\n",
+ " return image_np"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hy1CO7FysqE1"
+ },
+ "source": [
+ "## Import both Mask RCNN saved model(material \u0026 material form) from the repo."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "EH_s2raFsYfq"
+ },
+ "outputs": [],
+ "source": [
+ "# 'material_model' output is both material and its sub type e.g. Plastics_PET.\n",
+ "# 'material_form_model' outputs the form of an object e.g. can, bottle, etc.\n",
+ "MODEL_WEIGHTS = {\n",
+ " 'material_url': (\n",
+ " 'https://storage.googleapis.com/tf_model_garden/vision/'\n",
+ " 'waste_identification_ml/two_model_strategy/material/'\n",
+ " 'material_version_2.zip'\n",
+ " ),\n",
+ " 'material_form_url': (\n",
+ " 'https://storage.googleapis.com/tf_model_garden/vision/'\n",
+ " 'waste_identification_ml/two_model_strategy/material_form/'\n",
+ " 'material_form_version_2.zip'\n",
+ " ),\n",
+ "}\n",
+ "\n",
+ "\n",
+ "SAVED_MODEL_PATH = {\n",
+ "'material_model' : 'material/saved_model/',\n",
+ "'material_form_model' : 'material_form/saved_model/',\n",
+ "}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "hLTB2Q0Tsmcx"
+ },
+ "outputs": [],
+ "source": [
+ "# Download the model weights from the Google's repo.\n",
+ "url1 = MODEL_WEIGHTS['material_url']\n",
+ "url2 = MODEL_WEIGHTS['material_form_url']\n",
+ "!python3 download_and_unzip_models.py $url1 $url2"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "euasOIi2tJHA"
+ },
+ "source": [
+ "## Preprocess an image."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ozr7NLHbtTvM"
+ },
+ "outputs": [],
+ "source": [
+ "image1 = process_image('image_2.png')\n",
+ "image2 = process_image('image_4.png')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "oi6abxdeugCh"
+ },
+ "source": [
+ "## Load original SavedModel."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "bFZnUjV7uj_r",
+ "outputId": "ac1818ab-bf34-42b3-c00e-50cd684779a9"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "loading model...\n",
+ "material/saved_model/\n",
+ "model loaded!\n",
+ "loading model...\n",
+ "material_form/saved_model/\n",
+ "model loaded!\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Loading both models.\n",
+ "detection_fns = [\n",
+ " load_model(model_path)\n",
+ " for model_path in SAVED_MODEL_PATH.values()\n",
+ "]"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "DRjqNjUUvZ7K"
+ },
+ "source": [
+ "# Convert to TensorRT model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ONvJgKJjv0dw"
+ },
+ "outputs": [],
+ "source": [
+ "TENSORRT_MODEL_PATH = {\n",
+ "'material_model' : 'tensorrt/material/saved_model/',\n",
+ "'material_form_model' : 'tensorrt/material_form/saved_model/',\n",
+ "}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "zaeRIsVXv3vG",
+ "outputId": "d5f0555b-29c2-4d87-b72c-848f67772994"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Directory tensorrt/material/saved_model/ created successfully\n",
+ "Directory tensorrt/material_form/saved_model/ created successfully\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Create directories to store TensorRT models.\n",
+ "for value in TENSORRT_MODEL_PATH.values():\n",
+ " create_directory(value)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "HNLEo7BRv-db",
+ "outputId": "189e9842-0bee-46ae-b0dc-8051d5a23774"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "material/saved_model/ tensorrt/material/saved_model/\n",
+ "material_form/saved_model/ tensorrt/material_form/saved_model/\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Convert Tensorflow saved models into TensorRT models.\n",
+ "for key in SAVED_MODEL_PATH.keys():\n",
+ " value1 = SAVED_MODEL_PATH.get(key)\n",
+ " value2 = TENSORRT_MODEL_PATH.get(key)\n",
+ " print(value1, value2)\n",
+ " convert_to_tensorrt(value1, value2)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Vsl1OIMnUpGm"
+ },
+ "source": [
+ "## Load TensorRT models."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "xq_8sCo5wX74",
+ "outputId": "dd4a9bd1-4cd5-4798-f7d2-71b0d402861b"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "loading model...\n",
+ "tensorrt/material/saved_model/\n",
+ "model loaded!\n",
+ "loading model...\n",
+ "tensorrt/material_form/saved_model/\n",
+ "model loaded!\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Loading both models.\n",
+ "detection_fns_tensorrt = [\n",
+ " load_model(model_path)\n",
+ " for model_path in TENSORRT_MODEL_PATH.values()\n",
+ "]"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "aOuzd9BZD2uP"
+ },
+ "source": [
+ "## Checking speed with SavedModel."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "NyCPXNTHArMN",
+ "outputId": "7e1f9354-7347-4d86-c831-92c35333d939"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "386 ms ± 1.43 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%timeit\n",
+ "# Inference speed with first image.\n",
+ "results = list(\n",
+ " map(\n",
+ " lambda model: perform_detection(model, image1),\n",
+ " detection_fns\n",
+ " )\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "UNU4CTIaA3oA",
+ "outputId": "2e82a72a-4d68-45ca-a269-a4a22c043451"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "169 ms ± 1.13 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%timeit\n",
+ "detection_fns[0](image2)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "ecJW0MU4BGF_",
+ "outputId": "009107b8-28c4-40d8-dd58-8a486dfcc7b0"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "200 ms ± 1.9 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%timeit\n",
+ "detection_fns[1](image2)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "sPK4-NdBwT3R"
+ },
+ "source": [
+ "## Create an inference engine for TensorRT by predicting over a single image."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "B9PCp7JRxx6o",
+ "outputId": "c142db41-9932-44cb-ba2d-2021b9057923"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "210 ms ± 4.24 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%timeit\n",
+ "# Inference speed with first image.\n",
+ "results = list(\n",
+ " map(\n",
+ " lambda model: perform_detection(model, image1),\n",
+ " detection_fns_tensorrt\n",
+ " )\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "WfO9Fp-fEGrT"
+ },
+ "source": [
+ "## Checking speed with TensorRT model."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "Rx1cd0RrDmQm",
+ "outputId": "bda8db46-c793-4f69-f1aa-8d5746660737"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "83.8 ms ± 985 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%timeit\n",
+ "detection_fns_tensorrt[0](image2)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "GB8YPex-Dno3",
+ "outputId": "37509a66-0087-44b9-b9c8-e249530c9c2b"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "122 ms ± 1.52 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%timeit\n",
+ "detection_fns_tensorrt[1](image2)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hmiE_SMMEga4"
+ },
+ "source": [
+ "## Conclusion"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "DhcCcNOxEiNZ"
+ },
+ "source": [
+ "Average inference speed of 1st saved model over image 2 = **169 ms**\\\n",
+ "Average inference speed of 1st TensorRT model over image 2 = **83.8 ms**\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ShAgOkmhEwKy"
+ },
+ "source": [
+ "Average inference speed of 2nd saved model over image 2 = **210 ms**\\\n",
+ "Average inference speed of 2nd TensorRT model over image 2 = **122 ms**"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "machine_shape": "hm",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_inference/Inference_PyTorch_experimental.ipynb b/official/projects/waste_identification_ml/model_inference/Inference_PyTorch_experimental.ipynb
new file mode 100644
index 00000000000..c8fe748dfa9
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/Inference_PyTorch_experimental.ipynb
@@ -0,0 +1,444 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ZtzVZpUyxS5B"
+ },
+ "source": [
+ "# Waste identification with instance segmentation in PyTorch"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "MTrqWFRdxYyJ"
+ },
+ "source": [
+ "Welcome to the Instance Segmentation Colab! This notebook will take you through the steps of running an \"out-of-the-box\" Mask RCNN Instance Segmentation model on image from Detectron2."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "F0Gxj-CPxnFe"
+ },
+ "source": [
+ "To finish this task, a proper path for the model and a single image needs to be provided. The path to the labels on which the models are trained is in the waste_identification_ml directory inside the Tensorflow Model Garden repository. The label files are inferred automatically for the model."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "LgJKrYHEyAEv"
+ },
+ "source": [
+ "## RESTART the colab notebook after installing packages of Detectron2."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XXxxDT_QW2mR"
+ },
+ "outputs": [],
+ "source": [
+ "# Clone the Detectron2 repository and install the required packages.\n",
+ "# Relax as installing packages might take a while.\n",
+ "!git clone 'https://github.com/facebookresearch/detectron2'\n",
+ "!pip install 'git+https://github.com/facebookresearch/detectron2.git'"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WCWF8h6TuNqZ"
+ },
+ "outputs": [],
+ "source": [
+ "# Install supervision package for the postprocessing of output results\n",
+ "# from Detectron2 Mask RCNN model.\n",
+ "!pip install -q supervision"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "I1oMgVZj4n9n"
+ },
+ "source": [
+ "## Clone the TF Model Garden repo where the waste identification project is located."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "SHm3oznC4uBe"
+ },
+ "outputs": [],
+ "source": [
+ "!git clone --depth 1 https://github.com/tensorflow/models 2\u003e/dev/null"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ckQUcFtq1P3w"
+ },
+ "source": [
+ "## Imports and Setup"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "LdgIDnYe1dgv"
+ },
+ "outputs": [],
+ "source": [
+ "# Third-Party Imports\n",
+ "import csv\n",
+ "import torch\n",
+ "import cv2\n",
+ "import matplotlib.pyplot as plt\n",
+ "import supervision as sv\n",
+ "from PIL import Image\n",
+ "\n",
+ "# Detectron2 Imports\n",
+ "import detectron2\n",
+ "from detectron2.utils.logger import setup_logger\n",
+ "from detectron2.engine import DefaultPredictor\n",
+ "from detectron2.config import get_cfg\n",
+ "from detectron2.structures import Instances, Boxes\n",
+ "from detectron2.data.catalog import Metadata\n",
+ "from detectron2.utils.visualizer import Visualizer\n",
+ "\n",
+ "# Setup Detectron2 Logger\n",
+ "setup_logger()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "R48IuhsN__vG"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Utilities\n",
+ "\n",
+ "\n",
+ "def convert_detections_to_instances(\n",
+ " outputs: dict,\n",
+ " image_size: tuple[int, int] = (1024, 1024),\n",
+ " nms_threshold: float = 0.8,\n",
+ " class_agnostic: bool = True\n",
+ ") -\u003e dict[str, Instances]:\n",
+ " \"\"\"Convert Detectron2 model outputs to an Instances object with Non-Maximum Suppression (NMS) applied.\n",
+ "\n",
+ " Args:\n",
+ " outputs: Detectron2 model output containing instance predictions.\n",
+ " image_size: Image dimensions (height, width).\n",
+ " nms_threshold: Non-Maximum Suppression (NMS) threshold.\n",
+ " class_agnostic: Whether NMS should be applied in a class-agnostic manner.\n",
+ "\n",
+ " Returns:\n",
+ " Reformatted Detectron2 output as {\"instances\": Instances}.\n",
+ " \"\"\"\n",
+ " # Apply NMS and convert to supervision Detections format\n",
+ " detections = (\n",
+ " sv.Detections.from_detectron2(outputs)\n",
+ " .with_nms(threshold=nms_threshold, class_agnostic=class_agnostic)\n",
+ " )\n",
+ "\n",
+ " # Convert extracted values to PyTorch tensors\n",
+ " bboxes = torch.tensor(detections.xyxy, dtype=torch.float32)\n",
+ " scores = torch.tensor(detections.confidence, dtype=torch.float32)\n",
+ " classes = torch.tensor(detections.class_id, dtype=torch.int64)\n",
+ "\n",
+ " # Create an Instances object\n",
+ " output_instances = Instances(image_size)\n",
+ " output_instances.set(\"pred_boxes\", Boxes(bboxes))\n",
+ " output_instances.set(\"scores\", scores)\n",
+ " output_instances.set(\"pred_classes\", classes)\n",
+ "\n",
+ " # Add masks if available\n",
+ " if detections.mask is not None:\n",
+ " masks = torch.tensor(detections.mask, dtype=torch.uint8)\n",
+ " output_instances.set(\"pred_masks\", masks)\n",
+ "\n",
+ " return {\"instances\": output_instances}\n",
+ "\n",
+ "\n",
+ "def read_csv(file_path: str) -\u003e list[str]:\n",
+ " \"\"\"Reads a CSV file and returns its contents as a list.\n",
+ "\n",
+ " This function reads the given CSV file, skips the header, and assumes\n",
+ " there is only one column in the CSV. It returns the contents as a list of\n",
+ " strings.\n",
+ "\n",
+ " Args:\n",
+ " file_path: The path to the CSV file.\n",
+ "\n",
+ " Returns:\n",
+ " The contents of the CSV file as a list of strings.\n",
+ " \"\"\"\n",
+ " data_list = []\n",
+ " with open(file_path, 'r') as csvfile:\n",
+ " reader = csv.reader(csvfile)\n",
+ " for row in reader:\n",
+ " data_list.append(row[0])\n",
+ " return data_list"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "B8Q3y2r04e57"
+ },
+ "source": [
+ "## Import and load the labels."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "p2426m2L41Ny"
+ },
+ "outputs": [],
+ "source": [
+ "LABELS_PATH = (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/data/45_labels.csv'\n",
+ ")\n",
+ "\n",
+ "labels = read_csv(LABELS_PATH)\n",
+ "\n",
+ "my_metadata = Metadata()\n",
+ "my_metadata.set(thing_classes=labels)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "oWQkSqN_3aTP"
+ },
+ "source": [
+ "## Import Detectron2 Mask RCNN model."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "-P64IPCT3fsr"
+ },
+ "outputs": [],
+ "source": [
+ "%%bash\n",
+ "wget https://storage.googleapis.com/tf_model_garden/vision/\\\n",
+ "waste_identification_ml/Detectron2_Jan2025_1024_1024.zip\n",
+ "\n",
+ "unzip Detectron2_Jan2025_1024_1024.zip \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "oJ8MSkWh5uLG"
+ },
+ "source": [
+ "## Load the model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "LVxcxg7GYQ1U"
+ },
+ "outputs": [],
+ "source": [
+ "# Initialize the Detectron2 configuration object\n",
+ "cfg = get_cfg()\n",
+ "\n",
+ "# Load the model configuration from a YAML file.\n",
+ "cfg.merge_from_file(\"config.yaml\")\n",
+ "\n",
+ "# Set the confidence threshold.\n",
+ "cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = 0.5\n",
+ "\n",
+ "# Specify the path to the trained model weights.\n",
+ "cfg.MODEL.WEIGHTS = \"model_final.pth\"\n",
+ "\n",
+ "# Create a predictor object using the configured model.\n",
+ "predictor = DefaultPredictor(cfg)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "nYQJfdL46GPD"
+ },
+ "source": [
+ "## Import and load an image"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "EKjCF3KY6IUz"
+ },
+ "outputs": [],
+ "source": [
+ "# Path to a sample image stored in the repo.\n",
+ "IMAGES_FOR_TEST = {\n",
+ " 'Image1': (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/sample_images/image_2.png'\n",
+ " )\n",
+ "}\n",
+ "\n",
+ "# The model is trained on 1024 x 1024 image dimensions\n",
+ "HEIGHT = 1024\n",
+ "WIDTH = 1024"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "lIsvxs4e884I"
+ },
+ "outputs": [],
+ "source": [
+ "original_image = cv2.imread(IMAGES_FOR_TEST['Image1'])\n",
+ "original_height, original_width = original_image.shape[:2]\n",
+ "\n",
+ "resized_image = cv2.resize(\n",
+ " original_image,\n",
+ " (WIDTH, HEIGHT),\n",
+ " interpolation=cv2.INTER_AREA\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ney4CQXu6gvN"
+ },
+ "source": [
+ "## Perform prediction"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "NQASihEu6icr"
+ },
+ "outputs": [],
+ "source": [
+ "outputs = predictor(resized_image)\n",
+ "outputs = convert_detections_to_instances(outputs)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "mwehoMHn7Ruw"
+ },
+ "source": [
+ "## Visualize the results"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "keAsVmTy9I_4"
+ },
+ "outputs": [],
+ "source": [
+ "# Extract the predicted instances\n",
+ "instances = outputs[\"instances\"].to(\"cpu\")\n",
+ "\n",
+ "# Rescale bounding boxes back to the original image size\n",
+ "scale_x = original_width / WIDTH\n",
+ "scale_y = original_height / HEIGHT\n",
+ "instances.pred_boxes.scale(scale_x, scale_y)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JAdEtkhQ98Z9"
+ },
+ "outputs": [],
+ "source": [
+ "# Resize masks to match the original image size\n",
+ "if instances.has(\"pred_masks\"):\n",
+ " pred_masks = instances.pred_masks.numpy() # Convert to NumPy array\n",
+ " resized_masks = []\n",
+ "\n",
+ " for mask in pred_masks:\n",
+ " resized_mask = cv2.resize(\n",
+ " mask.astype(\"uint8\"),\n",
+ " (original_width, original_height),\n",
+ " interpolation=cv2.INTER_NEAREST\n",
+ " )\n",
+ " resized_masks.append(resized_mask)\n",
+ "\n",
+ " instances.pred_masks = torch.tensor(resized_masks, dtype=torch.uint8)\n",
+ "\n",
+ "# Initialize the visualizer with the original image\n",
+ "visualizer = Visualizer(\n",
+ " img_rgb=original_image, # Use the original image\n",
+ " metadata=my_metadata, # Metadata containing class labels, colors, etc.\n",
+ " scale=1 # Scale factor for visualization\n",
+ ")\n",
+ "\n",
+ "# Draw predictions on the original image\n",
+ "visualized_image = visualizer.draw_instance_predictions(instances).get_image()\n",
+ "\n",
+ "# Convert BGR to RGB for correct visualization in Matplotlib\n",
+ "visualized_image = visualized_image[:, :, ::-1]\n",
+ "\n",
+ "# Display the final image with predictions overlaid on the original image\n",
+ "plt.figure(figsize=(20, 20))\n",
+ "plt.axis(\"off\")\n",
+ "plt.imshow(visualized_image)\n",
+ "plt.show()"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "provenance": [
+ {
+ "file_id": "1A48TxIzVaHghg_ZYAxIgjv3uL7XhQEtR",
+ "timestamp": 1740088542918
+ }
+ ]
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_inference/Inference_Tensorflow.ipynb b/official/projects/waste_identification_ml/model_inference/Inference_Tensorflow.ipynb
new file mode 100644
index 00000000000..a7e14eb9c14
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/Inference_Tensorflow.ipynb
@@ -0,0 +1,736 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "TtlIRiNXWlQ0"
+ },
+ "source": [
+ "# Waste identification with instance segmentation in TensorFlow"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ohoMgYgXWsIO"
+ },
+ "source": [
+ "Welcome to the Instance Segmentation Colab! This notebook will take you through the steps of running an \"out-of-the-box\" Mask RCNN Instance Segmentation model on image from TF Model Garden."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "8PKG9z4VYPEs"
+ },
+ "source": [
+ "To finish this task, a proper path for the saved models and a single image needs to be provided. The path to the labels on which the models are trained is in the waste_identification_ml directory inside the Tensorflow Model Garden repository. The label files are inferred automatically for both models."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "j7yl9CqgYWvS"
+ },
+ "source": [
+ "## Imports and Setup"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ELUFMVDDAopS"
+ },
+ "outputs": [],
+ "source": [
+ "import sys\n",
+ "import tensorflow as tf\n",
+ "import csv\n",
+ "from typing import Any, TypedDict, Callable\n",
+ "import cv2\n",
+ "import logging\n",
+ "import numpy as np\n",
+ "import matplotlib.pyplot as plt\n",
+ "logging.disable(logging.WARNING)\n",
+ "\n",
+ "%matplotlib inline"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "_77YK3a_BCg_"
+ },
+ "source": [
+ "To visualize the images with the proper detected boxes and segmentation masks, we will use the TensorFlow Object Detection API. To install it we will clone the repo.\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qhk_NujKO0mb"
+ },
+ "outputs": [],
+ "source": [
+ "# Clone the tensorflow models repository.\n",
+ "!git clone --depth 1 https://github.com/tensorflow/models 2\u003e/dev/null"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "o1dYyG55BtWb"
+ },
+ "outputs": [],
+ "source": [
+ "sys.path.append('models/research/')\n",
+ "from object_detection.utils import ops as utils_ops\n",
+ "from object_detection.utils import visualization_utils as viz_utils"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "GO488S78_2GJ"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Utilities\n",
+ "class ItemDict(TypedDict):\n",
+ " id: int\n",
+ " name: str\n",
+ " supercategory: str\n",
+ "\n",
+ "\n",
+ "def load_model(model_path: str) -\u003e Callable:\n",
+ " \"\"\"Loads a TensorFlow SavedModel and returns a function for making predictions.\n",
+ "\n",
+ " Args:\n",
+ " model_path: Path to the TensorFlow SavedModel.\n",
+ "\n",
+ " Returns:\n",
+ " A function that can be used to make predictions.\n",
+ " \"\"\"\n",
+ " try:\n",
+ " print('loading model...')\n",
+ " model = tf.saved_model.load(model_path)\n",
+ " print('model loaded!')\n",
+ " detection_fn = model.signatures['serving_default']\n",
+ " return detection_fn\n",
+ " except (OSError, ValueError, KeyError) as e:\n",
+ " print(f\"Error loading model: {e}\")\n",
+ " raise\n",
+ "\n",
+ "\n",
+ "def perform_detection(model: Callable, image: np.ndarray) -\u003e dict[str, np.ndarray]:\n",
+ " \"\"\"Perform Mask R-CNN object detection on an image using the specified model.\n",
+ "\n",
+ " Args:\n",
+ " model: A function that can be used to make predictions.\n",
+ " image: A NumPy array representing the image to be processed.\n",
+ "\n",
+ " Returns:\n",
+ " Detection results, where keys are output names and values are NumPy arrays.\n",
+ " \"\"\"\n",
+ " detection_results = model(image)\n",
+ " detection_results = {key: value.numpy() for key, value in detection_results.items()}\n",
+ " return detection_results\n",
+ "\n",
+ "\n",
+ "def _read_csv_to_list(file_path: str) -\u003e list[str]:\n",
+ " \"\"\"Reads a CSV file and returns its contents as a list.\n",
+ "\n",
+ " This function reads the given CSV file, skips the header, and assumes\n",
+ " there is only one column in the CSV. It returns the contents as a list of\n",
+ " strings.\n",
+ "\n",
+ " Args:\n",
+ " file_path: The path to the CSV file.\n",
+ "\n",
+ " Returns:\n",
+ " The contents of the CSV file as a list of strings.\n",
+ " \"\"\"\n",
+ " data_list = []\n",
+ " with open(file_path, 'r') as csvfile:\n",
+ " reader = csv.reader(csvfile)\n",
+ " for row in reader:\n",
+ " data_list.append(row[0]) # Assuming there is only one column in the CSV\n",
+ " return data_list\n",
+ "\n",
+ "\n",
+ "def _categories_dictionary(objects: list[str]) -\u003e dict[int, ItemDict]:\n",
+ " \"\"\"This function takes a list of objects and returns a dictionaries.\n",
+ "\n",
+ " A dictionary of objects, where each object is represented by a dictionary\n",
+ " with the following keys:\n",
+ " - id: The ID of the object.\n",
+ " - name: The name of the object.\n",
+ " - supercategory: The supercategory of the object.\n",
+ "\n",
+ " Args:\n",
+ " objects: A list of strings, where each string is the name of an\n",
+ " object.\n",
+ "\n",
+ " Returns:\n",
+ " A tuple of two dictionaries, as described above.\n",
+ " \"\"\"\n",
+ " category_index = {}\n",
+ " for num, obj_name in enumerate(objects, start=1):\n",
+ " obj_dict = {'id': num, 'name': obj_name, 'supercategory': 'objects'}\n",
+ " category_index[num] = obj_dict\n",
+ " return category_index\n",
+ "\n",
+ "\n",
+ "def load_labels(labels_path: str) -\u003e tuple[list[str], dict[int, ItemDict]]:\n",
+ " \"\"\"\n",
+ " Load label mappings from a CSV file and generate category indices.\n",
+ "\n",
+ " Args:\n",
+ " labels_path (str): Path to the CSV file containing label mappings.\n",
+ "\n",
+ " Returns:\n",
+ " Tuple[Dict[int, dict], Dict[int, dict]]:\n",
+ " - A dictionary mapping category IDs to label details.\n",
+ " - A processed category index dictionary.\n",
+ " \"\"\"\n",
+ " labels = _read_csv_to_list(labels_path)\n",
+ " category_index = _categories_dictionary(labels)\n",
+ " return labels, category_index\n",
+ "\n",
+ "\n",
+ "def preprocess_image(path: str, height: int, width: int) -\u003e tuple[np.ndarray, np.ndarray]:\n",
+ " \"\"\"\n",
+ " Load an image from a file into a NumPy array, resize it, and expand dimensions for batch processing.\n",
+ "\n",
+ " Args:\n",
+ " path: The file path to the image.\n",
+ " height: Desired height of the resized image.\n",
+ " width: Desired width of the resized image.\n",
+ "\n",
+ " Returns:\n",
+ " original_image: The original image with shape (original_height, original_width, 3).\n",
+ " resized_image: The resized image with shape (1, height, width, 3), suitable for model input.\n",
+ " \"\"\"\n",
+ " original_image = cv2.imread(path)\n",
+ " if original_image is None:\n",
+ " raise FileNotFoundError(f\"Image not found at path: {path}\")\n",
+ "\n",
+ " original_image = cv2.cvtColor(original_image, cv2.COLOR_BGR2RGB)\n",
+ " resized_image = cv2.resize(original_image, (width, height), interpolation=cv2.INTER_AREA)\n",
+ " resized_image = np.expand_dims(resized_image, axis=0)\n",
+ "\n",
+ " return original_image, resized_image\n",
+ "\n",
+ "\n",
+ "def filter_detection(results: dict[str, np.ndarray], valid_indices: np.ndarray) -\u003e dict[str, np.ndarray]:\n",
+ " \"\"\"Filter the detection results based on the valid indices.\n",
+ "\n",
+ " Args:\n",
+ " results: The detection results from the model.\n",
+ " valid_indices: The indices of the valid detections.\n",
+ "\n",
+ " Returns:\n",
+ " The filtered detection results.\n",
+ " \"\"\"\n",
+ " if np.array(valid_indices).dtype == bool:\n",
+ " new_num_detections = int(np.sum(valid_indices))\n",
+ " else:\n",
+ " new_num_detections = len(valid_indices)\n",
+ "\n",
+ " # Define the keys to filter\n",
+ " keys_to_filter = [\n",
+ " 'detection_masks',\n",
+ " 'detection_masks_resized',\n",
+ " 'detection_masks_reframed',\n",
+ " 'detection_classes',\n",
+ " 'detection_boxes',\n",
+ " 'normalized_boxes',\n",
+ " 'detection_scores',\n",
+ " 'detection_classes_names',\n",
+ " ]\n",
+ "\n",
+ " # Apply filtering to the specified keys\n",
+ " filtered_output = {}\n",
+ "\n",
+ " for key in keys_to_filter:\n",
+ " if key in results:\n",
+ " if key == 'detection_masks':\n",
+ " filtered_output[key] = results[key][:, valid_indices, :, :]\n",
+ " elif key in ['detection_masks_resized', 'detection_masks_reframed']:\n",
+ " filtered_output[key] = results[key][valid_indices, :, :]\n",
+ " elif key in ['detection_boxes', 'normalized_boxes']:\n",
+ " filtered_output[key] = results[key][:, valid_indices, :]\n",
+ " elif key in ['detection_classes', 'detection_scores', 'detection_classes_names']:\n",
+ " filtered_output[key] = results[key][:, valid_indices]\n",
+ " filtered_output['num_detections'] = np.array([new_num_detections])\n",
+ "\n",
+ " return filtered_output\n",
+ "\n",
+ "\n",
+ "\n",
+ "def reframe_masks(results: dict[str, np.ndarray], boxes: str, height: int, width: int) -\u003e np.ndarray:\n",
+ " \"\"\"Reframe the masks to an image size.\n",
+ "\n",
+ " Args:\n",
+ " results: The detection results from the model.\n",
+ " boxes: The detection boxes.\n",
+ " height: The height of the original image.\n",
+ " width: The width of the original image.\n",
+ "\n",
+ " Returns:\n",
+ " The reframed masks.\n",
+ " \"\"\"\n",
+ " detection_masks = results['detection_masks'][0]\n",
+ " detection_boxes = results[boxes][0]\n",
+ " detection_masks_reframed = utils_ops.reframe_box_masks_to_image_masks(\n",
+ " detection_masks, detection_boxes, height, width\n",
+ " )\n",
+ " detection_masks_reframed = tf.cast(detection_masks_reframed \u003e 0.5, np.uint8)\n",
+ " detection_masks_reframed = detection_masks_reframed.numpy()\n",
+ " return detection_masks_reframed\n",
+ "\n",
+ "\n",
+ "def _calculate_area(mask: np.ndarray) -\u003e int:\n",
+ " \"\"\"Calculate the area of the mask.\n",
+ "\n",
+ " Args:\n",
+ " mask: The mask to calculate the area of.\n",
+ "\n",
+ " Returns:\n",
+ " The area of the mask.\n",
+ " \"\"\"\n",
+ " return np.sum(mask)\n",
+ "\n",
+ "\n",
+ "def _calculate_iou(mask1: np.ndarray, mask2: np.ndarray) -\u003e float:\n",
+ " \"\"\"Calculate the intersection over union (IoU) between two masks.\n",
+ "\n",
+ " Args:\n",
+ " mask1: The first mask.\n",
+ " mask2: The second mask.\n",
+ "\n",
+ " Returns:\n",
+ " The intersection over union (IoU) between the two masks.\n",
+ " \"\"\"\n",
+ " intersection = np.logical_and(mask1, mask2).sum()\n",
+ " union = np.logical_or(mask1, mask2).sum()\n",
+ " return intersection / union if union != 0 else 0\n",
+ "\n",
+ "\n",
+ "def _is_contained(mask1: np.ndarray, mask2: np.ndarray) -\u003e bool:\n",
+ " \"\"\"Check if mask1 is entirely contained within mask2.\n",
+ "\n",
+ " Args:\n",
+ " mask1: The first mask.\n",
+ " mask2: The second mask.\n",
+ "\n",
+ " Returns:\n",
+ " True if mask1 is entirely contained within mask2, False otherwise.\n",
+ " \"\"\"\n",
+ " return np.array_equal(np.logical_and(mask1, mask2), mask1)\n",
+ "\n",
+ "\n",
+ "def filter_masks(masks: np.ndarray, iou_threshold=0.8, area_threshold=None) -\u003e np.ndarray:\n",
+ " \"\"\"Filter the overlapping masks.\n",
+ "\n",
+ " Filter the masks based on the area and intersection over union (IoU).\n",
+ "\n",
+ " Args:\n",
+ " masks: The masks to filter.\n",
+ " iou_threshold: The threshold for the intersection over union (IoU) between\n",
+ " two masks.\n",
+ " area_threshold: The threshold for the area of the mask.\n",
+ "\n",
+ " Returns:\n",
+ " The indices of the unique masks.\n",
+ " \"\"\"\n",
+ " # Calculate the area for each mask\n",
+ " areas = np.array([_calculate_area(mask) for mask in masks])\n",
+ "\n",
+ " # Sort the masks based on area in descending order\n",
+ " sorted_indices = np.argsort(areas)[::-1]\n",
+ " sorted_masks = masks[sorted_indices]\n",
+ " sorted_areas = areas[sorted_indices]\n",
+ "\n",
+ " unique_indices = []\n",
+ "\n",
+ " for i, mask in enumerate(sorted_masks):\n",
+ " if (area_threshold is not None and sorted_areas[i] \u003e area_threshold) or sorted_areas[i] \u003c 4000:\n",
+ " continue\n",
+ "\n",
+ " keep = True\n",
+ " for j in range(i):\n",
+ " if _calculate_iou(mask, sorted_masks[j]) \u003e iou_threshold or _is_contained(\n",
+ " mask, sorted_masks[j]\n",
+ " ):\n",
+ " keep = False\n",
+ " break\n",
+ " if keep:\n",
+ " unique_indices.append(sorted_indices[i])\n",
+ "\n",
+ " return unique_indices\n",
+ "\n",
+ "\n",
+ "def adjust_image_size(height: int, width: int, min_size: int) -\u003e tuple[int, int]:\n",
+ " \"\"\"Adjust the image size to ensure both dimensions are at least 1024.\n",
+ "\n",
+ " Args:\n",
+ " height: The height of the image.\n",
+ " width: The width of the image.\n",
+ " min_size: Minimum size of the image dimension needed.\n",
+ "\n",
+ " Returns:\n",
+ " The adjusted height and width of the image.\n",
+ " \"\"\"\n",
+ " if height \u003c min_size or width \u003c min_size:\n",
+ " return height, width\n",
+ "\n",
+ " # Calculate the scale factor to ensure both dimensions remain at least 1024\n",
+ " scale_factor = min(height / min_size, width / min_size)\n",
+ "\n",
+ " new_height = int(height / scale_factor)\n",
+ " new_width = int(width / scale_factor)\n",
+ "\n",
+ " return new_height, new_width\n",
+ "\n",
+ "\n",
+ "def display_bbox_masks_labels(\n",
+ " result: dict[Any, np.ndarray],\n",
+ " image: np.ndarray,\n",
+ " category_index: dict[int, dict[str, str]],\n",
+ " threshold: float,\n",
+ ") -\u003e None:\n",
+ " \"\"\"Saves an image with visualized bounding boxes, labels, and masks.\n",
+ "\n",
+ " This function takes the output from Mask R-CNN, copies the original image,\n",
+ " and applies visualizations of detection boxes, classes, and scores.\n",
+ " If available, it also applies segmentation masks. The result is an image that\n",
+ " juxtaposes the original with the annotated version, saved to the specified\n",
+ " folder.\n",
+ "\n",
+ " Args:\n",
+ " result: The output from theMask RCNN model, expected to contain detection\n",
+ " boxes, classes, scores, reframed detection masks, etc.\n",
+ " image: The original image as a numpy array.\n",
+ " file_name: The filename for saving the output image.\n",
+ " folder: The folder path where the output image will be saved.\n",
+ " category_index: A dictionary mapping class IDs to class labels.\n",
+ " threshold: Value between 0 and 1 to filter out the prediction results.\n",
+ " \"\"\"\n",
+ " image_new = image.copy()\n",
+ " image_new = cv2.cvtColor(image_new, cv2.COLOR_BGR2RGB)\n",
+ " viz_utils.visualize_boxes_and_labels_on_image_array(\n",
+ " image_new,\n",
+ " result['normalized_boxes'][0],\n",
+ " (result['detection_classes'][0] + 0).astype(int),\n",
+ " result['detection_scores'][0],\n",
+ " category_index=category_index,\n",
+ " use_normalized_coordinates=True,\n",
+ " max_boxes_to_draw=100,\n",
+ " min_score_thresh=threshold,\n",
+ " agnostic_mode=False,\n",
+ " instance_masks=result.get('detection_masks_reframed', None),\n",
+ " line_thickness=4,\n",
+ " )\n",
+ " return image_new"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "t7d00cJH-68Z"
+ },
+ "outputs": [],
+ "source": [
+ "# Path to a sample image stored in the repo.\n",
+ "IMAGES_FOR_TEST = {\n",
+ " 'Image1': (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/sample_images/image_2.png'\n",
+ " )\n",
+ "}"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "4XjfDEq--UlE"
+ },
+ "source": [
+ "## Import and load pre-trained models."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ZQ435YHN3Lr-"
+ },
+ "outputs": [],
+ "source": [
+ "%%bash\n",
+ "wget https://storage.googleapis.com/tf_model_garden/vision/\\\n",
+ "waste_identification_ml/Jan2025_ver2_merged_1024_1024.zip -q\n",
+ "\n",
+ "unzip Jan2025_ver2_merged_1024_1024.zip \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "vMRVdtYEN5bg"
+ },
+ "outputs": [],
+ "source": [
+ "detection_fn = load_model('Jan2025_ver2_merged_1024_1024/')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "W6mmyLsOJicF"
+ },
+ "source": [
+ "## Load label map data"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "PM2A29OrJqaU"
+ },
+ "source": [
+ "Label maps correspond index numbers to category names, so that when our convolution network predicts 5, we know that this corresponds to airplane. Here we use internal utility functions, but anything that returns a dictionary mapping integers to appropriate string labels would be fine.\n",
+ "\n",
+ "We will load our labels from the same repository that we loaded the TF Object Detection API from."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "5RUzrh0uegqt"
+ },
+ "outputs": [],
+ "source": [
+ "LABELS_PATH = (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/data/45_labels.csv'\n",
+ ")\n",
+ "\n",
+ "labels, category_index = load_labels(LABELS_PATH)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "VkdD8-QvGZ23"
+ },
+ "source": [
+ "## Loading and pre-process an image"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "BSXQF57FGba5"
+ },
+ "source": [
+ "Let's try the model on a simple image."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "UVdlDwchGim_"
+ },
+ "source": [
+ "Note: when using images with an alpha channel, the model expect 3 channels images and the alpha will count as a 4th."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "-7ZS7gHgGk9f"
+ },
+ "outputs": [],
+ "source": [
+ "# The model is trained on 1024 x 1024 image dimensions\n",
+ "HEIGHT = 1024\n",
+ "WIDTH = 1024\n",
+ "IMAGE_PATH = (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/sample_images/image_2.png'\n",
+ ")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "rPrNf-NnOnL3"
+ },
+ "outputs": [],
+ "source": [
+ "original_image, resized_image = preprocess_image(IMAGE_PATH, HEIGHT, WIDTH)\n",
+ "input_tensor = tf.convert_to_tensor(resized_image, dtype=tf.uint8)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "-IhKjOZnQDJh"
+ },
+ "outputs": [],
+ "source": [
+ "%matplotlib inline\n",
+ "plt.figure(figsize=(10,10))\n",
+ "plt.imshow(original_image)\n",
+ "plt.show()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "H3r73X-FGzz-"
+ },
+ "source": [
+ "## Perform inference"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "SqW1z96LGzmZ"
+ },
+ "outputs": [],
+ "source": [
+ "# Running inference with bith the models.\n",
+ "result = perform_detection(detection_fn, input_tensor)\n",
+ "print(f'Total number of detections: {result[\"num_detections\"][0]}')\n",
+ "print(result.keys())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "332GbmRuG5A9"
+ },
+ "source": [
+ "## Process the output to remove overlapping or duplicate predictions."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ly7Uzms9R9yP"
+ },
+ "outputs": [],
+ "source": [
+ "# Prediction threshold.\n",
+ "PREDICTION_THRESHOLD = 0.50\n",
+ "area_threshold = None\n",
+ "\n",
+ "if result[\"num_detections\"][0]:\n",
+ " scores = result[\"detection_scores\"][0]\n",
+ " filtered_indices = scores \u003e PREDICTION_THRESHOLD\n",
+ " result = filter_detection(result, filtered_indices)\n",
+ " print(\n",
+ " \"Total number of detections after threshold:\"\n",
+ " f\" {result['num_detections'][0]}\"\n",
+ " )\n",
+ "\n",
+ "if result[\"num_detections\"][0]:\n",
+ " # Normalize the bounding boxes according to the resized image size.\n",
+ " result[\"normalized_boxes\"] = result[\"detection_boxes\"].copy()\n",
+ " result[\"normalized_boxes\"][:, :, [0, 2]] /= HEIGHT\n",
+ " result[\"normalized_boxes\"][:, :, [1, 3]] /= WIDTH\n",
+ "\n",
+ " # Adjust the image size to ensure both dimensions are at least 1024\n",
+ " # for saving images with bbx and masks.\n",
+ " height_plot, width_plot = adjust_image_size(\n",
+ " original_image.shape[0], original_image.shape[1], 1024\n",
+ " )\n",
+ " image_plot = cv2.resize(\n",
+ " original_image,\n",
+ " (width_plot, height_plot),\n",
+ " interpolation=cv2.INTER_AREA,\n",
+ " )\n",
+ " # Reframe the masks to the new size.\n",
+ " result[\"detection_masks_reframed\"] = reframe_masks(\n",
+ " result, \"normalized_boxes\", height_plot, width_plot\n",
+ " )\n",
+ "\n",
+ " # Filter the prediction results based on the area threshold and\n",
+ " # remove the overlapping masks.\n",
+ " unique_indices = filter_masks(\n",
+ " result[\"detection_masks_reframed\"],\n",
+ " iou_threshold=0.08,\n",
+ " area_threshold=area_threshold,\n",
+ " )\n",
+ " result = filter_detection(result, unique_indices)\n",
+ " print(\n",
+ " \"Total number of detections after filtering:\"\n",
+ " f\" {result['num_detections'][0]}\"\n",
+ " )"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "bkN8UgZ0HcM-"
+ },
+ "source": [
+ "## Visualization of masks"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Y5l3TrsiUpS7"
+ },
+ "outputs": [],
+ "source": [
+ "labeled_image = display_bbox_masks_labels(\n",
+ " result,\n",
+ " image_plot,\n",
+ " category_index,\n",
+ " PREDICTION_THRESHOLD\n",
+ ")\n",
+ "%matplotlib inline\n",
+ "plt.figure(figsize=(20,20))\n",
+ "plt.axis('off')\n",
+ "plt.imshow(labeled_image)"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_inference/cn_model_run.ipynb b/official/projects/waste_identification_ml/model_inference/cn_model_run.ipynb
new file mode 100644
index 00000000000..f48a141d2af
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/cn_model_run.ipynb
@@ -0,0 +1,391 @@
+{
+ "cells": [
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "# CircularNet - Waste identification with instance segmentation\n",
+ "Welcome to the Instance Segmentation Notebook! This notebook will take you through the steps of running an Instance Segmentation model on Images. \\\n",
+ "There are two ways to run it :\n",
+ "1. Pytorch Model : How to quickly run and test the CircularNet model's Pytorch version.\n",
+ "2. ONNX Model : Running the converted ONNX version of the model."
+ ],
+ "id": "76xIEUuR1TEs"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "# 1. Pytorch Model"
+ ],
+ "id": "XQVvI_GHHMuo"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Install Required python packages"
+ ],
+ "id": "B7mjq4jAGSB-"
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "!pip install -q rfdetr==1.5.2 supervision"
+ ],
+ "id": "0P81MC2G0VoH"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Import Libraries"
+ ],
+ "id": "MSN0xoY6GU9_"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "import json\n",
+ "import os\n",
+ "import glob\n",
+ "import numpy as np\n",
+ "import pandas as pd\n",
+ "from tqdm.auto import tqdm\n",
+ "\n",
+ "import cv2\n",
+ "from PIL import Image\n",
+ "import supervision as sv\n",
+ "from rfdetr import RFDETRSegMedium"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "3EgtR7ci0nuN"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Load Image Paths"
+ ],
+ "id": "TsCVNrGeGPDn"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "LIMIT = 20\n",
+ "IMAGE_DIR_PATH = \"./images\"\n",
+ "\n",
+ "all_images = glob.glob(f\"{IMAGE_DIR_PATH}/*\")[:LIMIT]\n",
+ "print(f\"Number of Images : {len(all_images)})"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "lSWXOVi2GLf2"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Load Segmentation Model"
+ ],
+ "id": "QWXD3AP7GME6"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "!wget https://storage.googleapis.com/tf_model_garden/vision/waste_identification_ml/CN-ModelCheckpoints/July2026_checkpoint/checkpoint_best_total.pth"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "t5AfPJVRI-5s"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "seg_medium_model = RFDETRSegMedium(pretrain_weights=\"./checkpoint_best_total.pth\")\n",
+ "class_id_to_class_name_mapper = seg_medium_model.class_names"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "J3SdExtY0p90"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Run the model"
+ ],
+ "id": "5Ns1YZXuGbT9"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "detections = [seg_medium_model.predict(img_path, threshold=0.5) for img_path in tqdm(all_images)]"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "IXOR6ePsGfHY"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Visualize Results"
+ ],
+ "id": "29McgOkmG7nw"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "detections_images = []\n",
+ "for i in np.arange(len(all_images)-1):\n",
+ " path = all_images[i]\n",
+ " image = Image.open(path)\n",
+ "\n",
+ " detection = detections[i]\n",
+ "\n",
+ " text_scale = sv.calculate_optimal_text_scale(resolution_wh=image.size)\n",
+ " thickness = sv.calculate_optimal_line_thickness(resolution_wh=image.size)\n",
+ " color = sv.ColorPalette.from_hex([\n",
+ " \"#ffff00\", \"#ff9b00\", \"#ff66ff\", \"#3399ff\", \"#ff66b2\", \"#ff8080\",\n",
+ " \"#b266ff\", \"#9999ff\", \"#66ffff\", \"#33ff99\", \"#66ff66\", \"#99ff00\"\n",
+ " ])\n",
+ "\n",
+ " bbox_annotator = sv.BoxAnnotator(color=color, thickness=thickness)\n",
+ " mask_annotator = sv.MaskAnnotator()\n",
+ "\n",
+ " label_annotator = sv.LabelAnnotator(\n",
+ " color=color,\n",
+ " text_color=sv.Color.BLACK,\n",
+ " text_scale=text_scale)\n",
+ "\n",
+ " detections_labels = [\n",
+ " f\"{class_id_to_class_name_mapper[int(class_id+1)]} {confidence:.2f}\"\n",
+ " for class_id, confidence\n",
+ " in zip(detection.class_id, detection.confidence)\n",
+ " ]\n",
+ "\n",
+ " detections_image = image.copy()\n",
+ " detections_image = mask_annotator.annotate(detections_image, detection)\n",
+ " detections_image = bbox_annotator.annotate(detections_image, detection)\n",
+ " detections_image = label_annotator.annotate(detections_image, detection, detections_labels)\n",
+ "\n",
+ " detections_images.append(detections_image)\n",
+ "\n",
+ "# adjust the grid size based upon number of images, below it is set to 5(rows) x 4 (columns) for viewing 20 images\n",
+ "sv.plot_images_grid(images=detections_images, grid_size=(5, 4), size=(20, 20))"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "c-4x_Whs097J"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "# 2. How to run ONNX Version of the Model"
+ ],
+ "id": "lEK0yWmsHSKL"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Download model and essential codebase"
+ ],
+ "id": "9z9j-zhcHglk"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "!git clone --depth 1 https://github.com/tensorflow/models.git\n",
+ "!wget https://storage.googleapis.com/tf_model_garden/vision/waste_identification_ml/CN-ModelCheckpoints/July2026_checkpoint/inference_model.onnx\n",
+ "!wget https://storage.googleapis.com/tf_model_garden/vision/waste_identification_ml/CN-ModelCheckpoints/sample_image.jpg\n",
+ "!mv models/official/projects/waste_identification_ml/Deploy/detr_cloud_deployment/client/labels50.csv ./"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "QHXuzvaCHRpV"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Install all the required libraries\n"
+ ],
+ "id": "ul1Zy1s5HjFO"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "!pip install -q onnx==1.19.1 onnxruntime==1.23.2 supervision tritonclient[http]==2.58.0"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "DzY3Zg-lHeQt"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Import Python Libraries"
+ ],
+ "id": "A2LFwEfBHvc5"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "import onnx\n",
+ "import onnxruntime as ort\n",
+ "\n",
+ "from PIL import Image\n",
+ "import supervision as sv\n",
+ "\n",
+ "import sys\n",
+ "from unittest.mock import MagicMock\n",
+ "sys.modules[\"color_extraction\"] = MagicMock()\n",
+ "\n",
+ "from models.official.projects.waste_identification_ml.Deploy.detr_cloud_deployment.client.triton_server_inference import TritonObjectDetector\n",
+ "from models.official.projects.waste_identification_ml.Deploy.detr_cloud_deployment.client.utils import draw_detections_and_save_image"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "vVwZyXA7HsdJ"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Define essential variables"
+ ],
+ "id": "FV1DL5SDH9oP"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "input_image_path = \"./sample_image.jpg\"\n",
+ "onnx_model_path = \"./inference_model.onnx\"\n",
+ "output_image_dimension = (1920, 1080)"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "OvBh1HCnHsgX"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Initialize ONNX version of the model"
+ ],
+ "id": "8xeC6b4xIPbL"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "model_utils = TritonObjectDetector()\n",
+ "onnx_model = ort.InferenceSession(onnx_model_path)"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "sYcXoly3IL99"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Do model Inferencing"
+ ],
+ "id": "2L0oJEtaIUgd"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "image_array = model_utils._get_input_batch_for_inference(image_path=input_image_path)\n",
+ "outputs = onnx_model.run(None, {\"input\": image_array})"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "lGHtqGMyISTW"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Format the model output as per output dimensions"
+ ],
+ "id": "GdnE_BJtIZVt"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "results = model_utils._reformat_triton_output_to_dict(\n",
+ " outputs, confidence_threshold=0.5, max_boxes=100\n",
+ ")\n",
+ "\n",
+ "# output_image_dimension makes sure the output image dimession along with the bounding boxes and masks are resized.\n",
+ "results = model_utils._scale_bbox_and_masks(results, target_dims=output_image_dimension)\n",
+ "results[\"class_names\"] = model_utils.get_class_names(results)"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "6eucX1xOsyrr"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "## Save output and Visualize Results"
+ ],
+ "id": "tf1XzX-5Ihij"
+ },
+ {
+ "metadata": {},
+ "cell_type": "code",
+ "source": [
+ "draw_detections_and_save_image(img=Image.open(input_image_path), results=results, save_path= \"./output.jpg\")\n",
+ "Image.open(\"./output.jpg\")"
+ ],
+ "outputs": [],
+ "execution_count": null,
+ "id": "4tIuPGpls0h4"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [
+ "# END of Notebook"
+ ],
+ "id": "TA529zr_Imwg"
+ },
+ {
+ "metadata": {},
+ "cell_type": "markdown",
+ "source": [],
+ "id": "1GHOytpUInpI"
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "private_outputs": true
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 5
+}
diff --git a/official/projects/waste_identification_ml/model_inference/color_and_property_extractor.py b/official/projects/waste_identification_ml/model_inference/color_and_property_extractor.py
new file mode 100644
index 00000000000..21e4d031493
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/color_and_property_extractor.py
@@ -0,0 +1,346 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Extract properties from each object mask and detect its color."""
+
+from typing import Dict, List, Tuple, TypeVar, Union
+
+import numpy as np
+import numpy.typing as npt
+import pandas as pd
+from skimage import color as skimage_color
+import skimage.measure
+from sklearn import cluster as sklearn_cluster
+from sklearn import neighbors as sklearn_neighbors
+import webcolors
+
+DType = TypeVar('DType', bound=np.generic)
+# Color representation as numpy array of 3 elements of float64
+# Those values could be in different scales like
+# RGB ([0.0,255.0], [0.0,255.0], [0.0 to 255.0])
+# LAB ([0.0,100], [-128,127], [-128,127])
+# NColor = Annotated[npt.NDArray[DType], Literal[3]][np.float64]
+NColor = np.ndarray
+
+
+PROPERTIES = [
+ 'area',
+ 'bbox',
+ 'convex_area',
+ 'bbox_area',
+ 'major_axis_length',
+ 'minor_axis_length',
+ 'eccentricity',
+ 'centroid',
+]
+
+GENERIC_COLORS = [
+ ('black', '#000000'),
+ ('green', '#008000'),
+ ('green', '#00ff00'), # lime
+ ('green', '#3cb371'), # mediumseagreen
+ ('green', '#2E8B57'), # seagreen
+ ('green', '#8FBC8B'), # darkseagreen
+ ('green', '#adff2f'), # olive
+ ('green', '#008080'), # Teal
+ ('green', '#808000'),
+ ('blue', '#000080'), # navy
+ ('blue', '#00008b'), # darkblue
+ ('blue', '#4682b4'), # steelblue
+ ('blue', '#40E0D0'), # turquoise
+ ('blue', '#00FFFF'), # cyan
+ ('blue', '#00ffff'), # aqua
+ ('blue', '#6495ED'), # cornflowerBlue
+ ('blue', '#4169E1'), # royalBlue
+ ('blue', '#87CEFA'), # lightSkyBlue
+ ('blue', '#4682B4'), # steelBlue
+ ('blue', '#B0C4DE'), # lightSteelBlue
+ ('blue', '#87CEEB'), # skyblue
+ ('blue', '#0000CD'), # mediumBlue
+ ('blue', '#0000ff'),
+ ('purple', '#800080'),
+ ('purple', '#9370db'), # mediumpurple
+ ('purple', '#8B008B'), # darkMagenta
+ ('purple', '#4B0082'), # indigo
+ ('red', '#ff0000'),
+ ('red', '#B22222'), # fireBrick
+ ('red', '#DC143C'), # fireBrick
+ ('red', '#8B0000'), # crimson
+ ('red', '#CD5C5C'), # indianred
+ ('red', '#F08080'), # lightCoral
+ ('red', '#FA8072'), # salmon
+ ('red', '#E9967A'), # darkSalmon
+ ('red', '#FFA07A'), # lightSalmon
+ ('gray', '#c0c0c0'), # silver,
+ ('gray', '#a9a9a9'), # +darkgray
+ ('gray', '#708090'), # +slategray
+ ('blue', '#778899'), # +lightslategray
+ ('white', '#ffffff'),
+ ('white', '#F5F5DC'), # beige
+ ('white', '#FFFAFA'), # snow
+ ('white', '#F0F8FF'), # aliceBlue
+ ('white', '#FFE4E1'), # mistyRose
+ ('yellow', '#ffff00'),
+ ('yellow', '#ffffe0'), # lightyellow
+ ('yellow', '#8B8000'), # darkyellow,
+ ('orange', '#ffa500'),
+ ('orange', '#ff8c00'), # darkorange
+ ('pink', '#ffc0cb'),
+ ('pink', '#ff00ff'), # fuchsia
+ ('pink', '#C71585'), # mediumVioletRed
+ ('pink', '#DB7093'), # paleVioletRed
+ ('pink', '#FFB6C1'), # lightPink
+ ('pink', '#FF69B4'), # hotPink
+ ('pink', '#FF1493'), # deepPink
+ ('pink', '#BC8F8F'), # rosybrown
+ ('brown', '#a52a2a'),
+ ('brown', '#8b4513'), # saddlebrown
+ ('brown', '#f4a460'), # sandybrown
+ ('brown', '#800000'), # maroon
+]
+
+
+def extract_properties_and_object_masks(
+ final_result: Dict[str, np.ndarray],
+ height: int,
+ width: int,
+ original_image: np.ndarray,
+) -> Tuple[List[pd.DataFrame], List[np.ndarray]]:
+ """Extract specific properties from given detection masks.
+
+ Properties that will be computed includes the area of the masks, bbox
+ coordinates, area of that bbox, convex length, major_axis_length,
+ minor_axis_length, eccentricity and centroid.
+
+ Args:
+ final_result: A dictionary containing the num_detections, detection_classes,
+ detection_scores,detection_boxes,detection_classes_names,
+ detection_masks_reframed'
+ height: The height of the original image.
+ width: The width of the original image.
+ original_image: The actual image on which the objects were detected.
+
+ Returns:
+ A tuple containing two lists:
+ 1. List of dataframes where each dataframe contains properties for a
+ detected object.
+ 2. List of ndarrays where each ndarray is a cropped portion of the
+ original image
+ corresponding to a detected object.
+ """
+ list_of_df = []
+ cropped_masks = []
+
+ for i, mask in enumerate(final_result['detection_masks_reframed']):
+ mask = np.where(mask, 1, 0)
+ df = pd.DataFrame(
+ skimage.measure.regionprops_table(mask, properties=PROPERTIES)
+ )
+ list_of_df.append(df)
+
+ bb = final_result['detection_boxes'][0][i]
+ ymin, xmin, ymax, xmax = (
+ int(bb[0] * height),
+ int(bb[1] * width),
+ int(bb[2] * height),
+ int(bb[3] * width),
+ )
+ mask = np.expand_dims(mask, axis=2)
+ cropped_object = np.where(
+ mask[ymin:ymax, xmin:xmax], original_image[ymin:ymax, xmin:xmax], 0
+ )
+ cropped_masks.append(cropped_object)
+
+ return list_of_df, cropped_masks
+
+
+def find_dominant_color(
+ image: np.ndarray, black_threshold: int = 50
+) -> Tuple[int, int, int]:
+ """Determines the dominant color in a given image.
+
+ Args:
+ image: An array representation of the image.
+ black_threshold: The intensity threshold below which pixels
+ are considered 'black' or near-black.
+
+ Returns:
+ The dominant RGB color in the format (R, G, B).
+ """
+ pixels = image.reshape(-1, 3)
+
+ # Filter out black pixels based on the threshold
+ non_black_pixels = pixels[(pixels > black_threshold).any(axis=1)]
+
+ if non_black_pixels.size:
+ kmeans = sklearn_cluster.KMeans(
+ n_clusters=1, n_init=10, random_state=0
+ ).fit(non_black_pixels)
+ dominant_color = kmeans.cluster_centers_[0].astype(int)
+ else:
+ dominant_color = np.array([0, 0, 0], dtype=int)
+ return tuple(dominant_color) # pyrefly: ignore[bad-return]
+
+
+def color_difference(color1: int, color2: int) -> Union[float, int]:
+ """Computes the squared difference between two color components.
+
+ Args:
+ color1: First color component.
+ color2: Second color component.
+
+ Returns:
+ The squared difference between the two color components.
+ """
+ return (color1 - color2) ** 2
+
+
+def est_color(requested_color: Tuple[int, int, int]) -> str:
+ """Estimates the closest named color for a given RGB color.
+
+ The function uses the Euclidean distance in the RGB space to find the closest
+ match among the CSS3 colors.
+
+ Args:
+ requested_color: The RGB color value for which to find the closest named
+ color. Expected format is (R, G, B).
+
+ Returns:
+ The name of the closest matching color from the CSS3 predefined colors.
+
+ Example: est_color((255, 0, 0))
+ 'red'
+ """
+ min_colors = {}
+ for key, name in webcolors.CSS3_HEX_TO_NAMES.items():
+ r_c, g_c, b_c = webcolors.hex_to_rgb(key)
+ rd = color_difference(r_c, requested_color[0])
+ gd = color_difference(g_c, requested_color[1])
+ bd = color_difference(b_c, requested_color[2])
+ min_colors[(rd + gd + bd)] = name
+ return min_colors[min(min_colors.keys())]
+
+
+def get_color_name(rgb_color: Tuple[int, int, int]) -> str | None:
+ """Retrieves the name of a given RGB color.
+
+ If the RGB color exactly matches one of the CSS3 predefined colors, it returns
+ the exact color name.
+ Otherwise, it estimates the closest matching color name.
+
+ Args:
+ rgb_color: The RGB color value for which to retrieve the name.
+
+ Returns:
+ The name of the color if found, or None if the color is marked as 'Na' or
+ not found.
+
+ Example: get_color_name((255, 0, 0))
+ 'red'
+ """
+ if 'Na' not in rgb_color:
+ try:
+ closest_color_name = webcolors.rgb_to_name(rgb_color)
+ except ValueError:
+ closest_color_name = est_color(rgb_color)
+ return closest_color_name
+ else:
+ return None
+
+
+def rgb_int_to_lab(rgb_int_color: Tuple[int, int, int]) -> NColor:
+ """Convert RGB color to LAB color space.
+
+ Args:
+ rgb_int_color: RGB tuple color e.g. (128,128,128)
+
+ Returns:
+ Numpy array of 3 elements that contains LAB color space.
+ """
+ return skimage_color.rgb2lab(
+ (rgb_int_color[0] / 255, rgb_int_color[1] / 255, rgb_int_color[2] / 255)
+ )
+
+
+def color_distance(
+ a: Tuple[int, int, int], b: Tuple[int, int, int]
+) -> np.ndarray:
+ """The color distance following the ciede2000 formula.
+
+ See: https://en.wikipedia.org/wiki/Color_difference#CIEDE2000
+
+ Args:
+ a: Color a
+ b: Color b
+
+ Returns:
+ The distance between color a and b
+ """
+ return skimage_color.deltaE_ciede2000(a, b, kC=0.6)
+
+
+def build_color_lab_list(
+ generic_colors: List[Tuple[str, str]]
+) -> Tuple[npt.NDArray[np.str_], List[NColor]]:
+ """Get Simple colors names and lab values.
+
+ Args:
+ generic_colors: List of colors in this format (color_name, rgb_value in hex)
+ e.g. [ ('black', '#000000'), ('green', '#008000'), ]
+
+ Returns:
+ Numpy array of strings that contains color names
+ ['black', 'green']
+ List of color lab values in the format of Numpy array of 3 elements
+ e.g.
+ [
+ np.array([0., 0., 0.]),
+ np.array([ 46.2276577 , -51.69868348, 49.89707556])
+ ]
+ """
+ names: list[str] = []
+ lab_values = []
+ for color_name, color_hex in generic_colors:
+ names.append(color_name)
+ hex_color = webcolors.hex_to_rgb(color_hex)
+ lab_values.append(rgb_int_to_lab(hex_color))
+ color_names = np.array(names)
+ return color_names, lab_values
+
+
+def get_generic_color_name(
+ rgb_colors: List[Tuple[int, int, int]],
+ generic_colors: List[Tuple[str, str]] | None = None,
+) -> List[str]:
+ """Retrieves generic names of given RGB colors.
+
+ Estimates the closest matching color name.
+
+ Args:
+ rgb_colors: A list of RGB values for which to retrieve the name.
+ generic_colors: A list of color names and their RGB values in hex.
+
+ Returns:
+ The list of closest color names.
+
+ Example: get_generic_color_name([(255, 0, 0), (0,0,0)])
+ ['red','black']
+ """
+ names, rgb_simple_colors = build_color_lab_list(
+ generic_colors or GENERIC_COLORS
+ )
+ tree = sklearn_neighbors.BallTree(rgb_simple_colors, metric=color_distance)
+ rgb_query = [*map(rgb_int_to_lab, rgb_colors)]
+ _, index = tree.query(rgb_query)
+ return [x[0] for x in names[index]]
diff --git a/official/projects/waste_identification_ml/model_inference/color_and_property_extractor_test.py b/official/projects/waste_identification_ml/model_inference/color_and_property_extractor_test.py
new file mode 100644
index 00000000000..847bef12eab
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/color_and_property_extractor_test.py
@@ -0,0 +1,127 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import unittest
+
+import numpy as np
+import pandas as pd
+
+from official.projects.waste_identification_ml.model_inference import color_and_property_extractor
+
+
+class ColorAndPropertyExtractor(unittest.TestCase):
+
+ def test_extract_properties_and_object_masks(self):
+ final_array = {
+ 'detection_masks_reframed': np.array([[
+ [0, 1, 1, 0, 0],
+ ]]),
+ 'detection_boxes': np.array([[[0, 0, 5, 5]]]),
+ }
+ original_image = np.array([
+ [0, 0, 0, 0, 0],
+ [0, 1, 0, 0, 0],
+ [0, 1, 1, 0, 0],
+ [0, 1, 1, 1, 0],
+ [0, 0, 0, 0, 0],
+ ])
+ expected_list = [
+ pd.DataFrame({
+ 'area': [2.0],
+ 'bbox-0': [0],
+ 'bbox-1': [1],
+ 'bbox-2': [1],
+ 'bbox-3': [3],
+ 'convex_area': [2.0],
+ 'bbox_area': [2.0],
+ 'major_axis_length': [2.0],
+ 'minor_axis_length': [0.0],
+ 'eccentricity': [1.0],
+ 'centroid-0': [0.0],
+ 'centroid-1': [1.5],
+ }),
+ ]
+ expected_mask = [
+ np.array([[
+ [0, 0, 0, 0, 0],
+ [0, 1, 0, 0, 0],
+ [0, 1, 1, 0, 0],
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ ]])
+ ]
+
+ result_list, result_masks = (
+ color_and_property_extractor.extract_properties_and_object_masks(
+ final_array, 5, 5, original_image
+ )
+ )
+
+ self.assertTrue(expected_list[0].equals(result_list[0]))
+ self.assertTrue(np.array_equal(expected_mask[0], result_masks[0]))
+
+ def test_find_dominant_color(self):
+ original_image = np.array([
+ [
+ (0, 0, 0),
+ (255, 255, 255),
+ ],
+ [
+ (128, 128, 128),
+ (128, 128, 128),
+ ],
+ ])
+
+ result = color_and_property_extractor.find_dominant_color(
+ original_image, black_threshold=50
+ )
+
+ self.assertEqual(result, (170, 170, 170))
+
+ def test_find_dominant_color_black(self):
+ original_image = np.array([
+ [
+ (0, 0, 0),
+ (0, 0, 0),
+ ],
+ [
+ (0, 0, 0),
+ (0, 0, 0),
+ ],
+ ])
+
+ result = color_and_property_extractor.find_dominant_color(
+ original_image, black_threshold=50
+ )
+
+ self.assertEqual(result, (0, 0, 0))
+
+ def test_est_color(self):
+ result = color_and_property_extractor.est_color((255, 0, 0))
+
+ self.assertEqual(result, 'red')
+
+ def test_generic_color(self):
+ test_colors = np.array(
+ [(255, 0, 0), (55, 118, 171), (73, 128, 41), (231, 112, 13), (0, 0, 0)]
+ )
+ expected_colors = ['red', 'blue', 'green', 'orange', 'black']
+
+ result = color_and_property_extractor.get_generic_color_name(test_colors)
+
+ self.assertEqual(result, expected_colors)
+
+
+if __name__ == '__main__':
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/model_inference/download_and_unzip_models.py b/official/projects/waste_identification_ml/model_inference/download_and_unzip_models.py
new file mode 100644
index 00000000000..9cd435ec69e
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/download_and_unzip_models.py
@@ -0,0 +1,100 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""This module provides utilities for executing shell commands.
+
+It particularly downloads and extracts Mask RCNN models from the TensorFlow
+model garden. It includes a function to execute shell commands and
+a custom exception to handle errors that arise from command execution.
+
+Functions:
+ - execute_command(cmd: str) -> str: Executes a shell command and returns its
+ standard output. Raises
+ a CommandExecutionError if the command execution fails.
+
+Exceptions:
+ - CommandExecutionError: Custom exception that's raised when there's an
+ error executing a shell command.
+
+Usage:
+ The main purpose of this module is to download two specific Mask RCNN models
+ and unzip them. The module
+ performs these operations when imported.
+
+Note:
+ It's recommended to not perform actions like downloading files on module
+ import in production applications.
+ It's better to move such tasks inside a function or a main block to allow
+ for more controlled execution.
+"""
+import argparse
+import os
+import subprocess
+
+
+class CommandExecutionError(Exception):
+ """Raised when there's an error executing a shell command."""
+
+ def __init__(self, cmd, returncode, stderr):
+ super().__init__(f"Error executing command: {cmd}. Error: {stderr}")
+ self.cmd = cmd
+ self.returncode = returncode
+ self.stderr = stderr
+
+
+def execute_command(cmd: str) -> str:
+ """Executes a shell command and returns its output."""
+ result = subprocess.run(
+ cmd,
+ shell=True,
+ stdout=subprocess.PIPE,
+ stderr=subprocess.PIPE,
+ check=False,
+ )
+
+ if result.returncode != 0:
+ raise CommandExecutionError(
+ cmd, result.returncode, result.stderr.decode("utf-8")
+ )
+
+ return result.stdout.decode("utf-8")
+
+
+def main(_) -> None:
+ # Download the provided files
+ execute_command(f"wget {args.material_url}")
+ execute_command(f"wget {args.material_form_url}")
+
+ # Create directories
+ os.makedirs("material", exist_ok=True)
+ os.makedirs("material_form", exist_ok=True)
+
+ # Unzip the provided files
+ zip_file1 = os.path.basename(args.material_url)
+ zip_file2 = os.path.basename(args.material_form_url)
+ execute_command(f"unzip {zip_file1} -d material/")
+ execute_command(f"unzip {zip_file2} -d material_form/")
+
+
+if __name__ == "__main__":
+ parser = argparse.ArgumentParser(
+ description="Download and extract Mask RCNN models."
+ )
+ parser.add_argument("material_url", help="repo url for material model")
+ parser.add_argument(
+ "material_form_url", help="repo url for material form model"
+ )
+
+ args = parser.parse_args()
+ main(args)
diff --git a/official/projects/waste_identification_ml/model_inference/labels.py b/official/projects/waste_identification_ml/model_inference/labels.py
new file mode 100644
index 00000000000..50995fdab5f
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/labels.py
@@ -0,0 +1,118 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Load labels for model prediction.
+
+Given paths of CSV files, task is to import them and convert into a
+form required for mapping with the model output.
+"""
+import csv
+from typing import TypedDict
+
+
+class ItemDict(TypedDict):
+ id: int
+ name: str
+ supercategory: str
+
+
+def read_csv_to_list(file_path: str) -> list[str]:
+ """Reads a CSV file and returns its contents as a list.
+
+ This function reads the given CSV file, skips the header, and assumes
+ there is only one column in the CSV. It returns the contents as a list of
+ strings.
+
+ Args:
+ file_path: The path to the CSV file.
+
+ Returns:
+ The contents of the CSV file as a list of strings.
+ """
+ data_list = []
+ with open(file_path, 'r') as csvfile:
+ reader = csv.reader(csvfile)
+ next(reader) # Skip the header row if present
+ for row in reader:
+ data_list.append(row[0]) # Assuming there is only one column in the CSV
+ return data_list
+
+
+def categories_dictionary(objects: list[str]) -> dict[int, ItemDict]:
+ """This function takes a list of objects and returns a dictionaries.
+
+ A dictionary of objects, where each object is represented by a dictionary
+ with the following keys:
+ - id: The ID of the object.
+ - name: The name of the object.
+ - supercategory: The supercategory of the object.
+
+ Args:
+ objects: A list of strings, where each string is the name of an
+ object.
+
+ Returns:
+ A tuple of two dictionaries, as described above.
+ """
+ category_index = {}
+
+ for num, obj_name in enumerate(objects, start=1):
+ obj_dict = {'id': num, 'name': obj_name, 'supercategory': 'objects'}
+ category_index[num] = obj_dict
+
+ return category_index
+
+
+def load_labels(
+ label_paths: dict[str, str]
+) -> tuple[list[list[str]], dict[int, ItemDict]]:
+ """Loads labels, combines them, and formats them for prediction.
+
+ This function reads labels for multiple models, combines the labels in
+ order to predict a single label output, and formats them into the desired
+ structure required for prediction.
+
+ Args:
+ label_paths: Dictionary of label paths for different models.
+
+ Returns:
+ - A list of lists containing individual category indices for each
+ model.
+ - A dictionary of combined category indices in the desired format for
+ prediction.
+
+ Note:
+ - The function assumes there are exactly two models.
+ - Inserts a category 'Na' for both models in case there is no detection.
+ - The total number of predicted labels for a combined model is
+ predetermined.
+ """
+ # loading labels for both models
+ category_indices = [read_csv_to_list(label) for label in label_paths.values()]
+
+ # insert a cateory 'Na' for both models in case there is no detection
+ for i in [0, 1]:
+ category_indices[i].insert(0, 'Na')
+
+ # combine the labels for both models in order to predict a single label output
+ combined_category_indices = []
+ for i in category_indices[0]:
+ for j in category_indices[1]:
+ combined_category_indices.append(f'{i}_{j}')
+ combined_category_indices.sort()
+
+ # convert the list of labels into a desired format required for prediction
+ category_index = categories_dictionary(combined_category_indices)
+
+ return category_indices, category_index
diff --git a/official/projects/waste_identification_ml/model_inference/labels_test.py b/official/projects/waste_identification_ml/model_inference/labels_test.py
new file mode 100644
index 00000000000..4cdded51214
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/labels_test.py
@@ -0,0 +1,72 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import os
+import unittest
+from official.projects.waste_identification_ml.model_inference import labels
+
+TESTDATA = os.path.join(os.path.dirname(os.path.abspath(__file__)), "testdata")
+
+
+class LabelsTest(unittest.TestCase):
+
+ def test_read_csv_to_list(self):
+ result = labels.read_csv_to_list(f"{TESTDATA}/csv_to_list.csv")
+
+ self.assertEqual(first=result, second=["alpha", "beta", "gamma"])
+
+ def test_categories_dictionary(self):
+ expected = {
+ 1: {"id": 1, "name": "alpha", "supercategory": "objects"},
+ 2: {"id": 2, "name": "beta", "supercategory": "objects"},
+ 3: {"id": 3, "name": "gamma", "supercategory": "objects"},
+ }
+
+ result = labels.categories_dictionary(["alpha", "beta", "gamma"])
+
+ self.assertEqual(first=result, second=expected)
+
+ def test_load_labels(self):
+ label_paths = {
+ "1": f"{TESTDATA}/categories_1.csv",
+ "2": f"{TESTDATA}/categories_2.csv",
+ }
+ expected = [
+ [
+ "Na",
+ "alpha",
+ "beta",
+ ],
+ ["Na", "gamma", "delta"],
+ ]
+ expected_index = {
+ 1: {"id": 1, "name": "Na_Na", "supercategory": "objects"},
+ 2: {"id": 2, "name": "Na_delta", "supercategory": "objects"},
+ 3: {"id": 3, "name": "Na_gamma", "supercategory": "objects"},
+ 4: {"id": 4, "name": "alpha_Na", "supercategory": "objects"},
+ 5: {"id": 5, "name": "alpha_delta", "supercategory": "objects"},
+ 6: {"id": 6, "name": "alpha_gamma", "supercategory": "objects"},
+ 7: {"id": 7, "name": "beta_Na", "supercategory": "objects"},
+ 8: {"id": 8, "name": "beta_delta", "supercategory": "objects"},
+ 9: {"id": 9, "name": "beta_gamma", "supercategory": "objects"},
+ }
+
+ result, result_index = labels.load_labels(label_paths)
+
+ self.assertEqual(first=result, second=expected)
+ self.assertEqual(first=result_index, second=expected_index)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/official/projects/waste_identification_ml/model_inference/postprocessing.py b/official/projects/waste_identification_ml/model_inference/postprocessing.py
new file mode 100644
index 00000000000..ba1f98b3741
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/postprocessing.py
@@ -0,0 +1,529 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Post process the results output from the ML model.
+
+Given the output from the 2 mask RCNN models. The 3 main tasks are done by three
+functions mentioned below -
+1. reframing_masks : Reframe the masks according to the size of an image and
+their respective positions within an image.
+2. find_similar_masks : Given masks from the output of 2 models. Find masks
+which belong to the same object and combine all of their attributes like
+confidence score, bounding boxes, label names, etc. The masks are mapped to each
+other if their score is above a threshold limit. Two outputs are combined into
+a single output.
+3. filter_bounding_boxes : The combined output may have nested bounding boxes of
+the same object. The parent bounding boxes are removed in this step so that any
+object should not have more than a single bounding box.
+"""
+import copy
+from typing import Any, Optional, TypedDict, Dict, Tuple, List
+import numpy as np
+import tensorflow as tf, tf_keras
+
+
+class DetectionResult(TypedDict):
+ num_detections: np.ndarray
+ detection_classes: np.ndarray
+ detection_scores: np.ndarray
+ detection_boxes: np.ndarray
+ detection_classes_names: np.ndarray
+ detection_masks_reframed: np.ndarray
+
+
+class ItemDict(TypedDict):
+ id: int
+ name: str
+ supercategory: str
+
+
+def reframe_image_corners_relative_to_boxes(boxes: tf.Tensor) -> tf.Tensor:
+ """Reframe the image corners ([0, 0, 1, 1]) to be relative to boxes.
+
+ The local coordinate frame of each box is assumed to be relative to
+ its own for corners.
+
+ Args:
+ boxes: A float tensor of [num_boxes, 4] of (ymin, xmin, ymax, xmax)
+ coordinates in relative coordinate space of each bounding box.
+
+ Returns:
+ reframed_boxes: Reframes boxes with same shape as input.
+ """
+ ymin, xmin, ymax, xmax = tf.unstack(boxes, axis=1)
+
+ height = tf.maximum(ymax - ymin, 1e-4)
+ width = tf.maximum(xmax - xmin, 1e-4)
+
+ ymin_out = (0 - ymin) / height
+ xmin_out = (0 - xmin) / width
+ ymax_out = (1 - ymin) / height
+ xmax_out = (1 - xmin) / width
+ return tf.stack([ymin_out, xmin_out, ymax_out, xmax_out], axis=1)
+
+
+def reframe_box_masks_to_image_masks(
+ box_masks: tf.Tensor,
+ boxes: tf.Tensor,
+ image_height: int,
+ image_width: int,
+ resize_method: str = 'bilinear',
+) -> tf.Tensor:
+ """Transforms the box masks back to full image masks.
+
+ Embeds masks in bounding boxes of larger masks whose shapes correspond to
+ image shape.
+
+ Args:
+ box_masks: A tensor of size [num_masks, mask_height, mask_width].
+ boxes: A tf.float32 tensor of size [num_masks, 4] containing the box
+ corners. Row i contains [ymin, xmin, ymax, xmax] of the box corresponding
+ to mask i. Note that the box corners are in normalized coordinates.
+ image_height: Image height. The output mask will have the same height as the
+ image height.
+ image_width: Image width. The output mask will have the same width as the
+ image width.
+ resize_method: The resize method, either 'bilinear' or 'nearest'. Note that
+ 'bilinear' is only respected if box_masks is a float.
+
+ Returns:
+ A tensor of size [num_masks, image_height, image_width] with the same dtype
+ as `box_masks`.
+ """
+ resize_method = 'nearest' if box_masks.dtype == tf.uint8 else resize_method
+
+ def reframe_box_masks_to_image_masks_default():
+ """The default function when there are more than 0 box masks."""
+
+ num_boxes = tf.shape(box_masks)[0]
+ box_masks_expanded = tf.expand_dims(box_masks, axis=3)
+
+ resized_crops = tf.image.crop_and_resize(
+ image=box_masks_expanded,
+ boxes=reframe_image_corners_relative_to_boxes(boxes),
+ box_indices=tf.range(num_boxes),
+ crop_size=[image_height, image_width],
+ method=resize_method,
+ extrapolation_value=0,
+ )
+ return tf.cast(resized_crops, box_masks.dtype)
+
+ image_masks = tf.cond(
+ tf.shape(box_masks)[0] > 0,
+ reframe_box_masks_to_image_masks_default,
+ lambda: tf.zeros([0, image_height, image_width, 1], box_masks.dtype),
+ )
+ return tf.squeeze(image_masks, axis=3)
+
+
+def reframing_masks(
+ results: Dict[str, np.ndarray], height: int, width: int
+) -> Dict[str, np.ndarray]:
+ """Processes the output from Mask RCNN model to create a full size mask.
+
+ Args:
+ results: list of dictionaries containing the output of Mask RCNN.
+ height: The height of the image.
+ width: The width of the image.
+
+ Returns:
+ A processed list of dictionaries.
+ """
+ result = copy.deepcopy(results)
+ result['detection_boxes'][0][:, [0, 2]] /= height
+ result['detection_boxes'][0][:, [1, 3]] /= width
+
+ detection_masks = tf.convert_to_tensor(result['detection_masks'][0])
+ detection_boxes = tf.convert_to_tensor(result['detection_boxes'][0])
+ detection_masks_reframed = reframe_box_masks_to_image_masks(
+ detection_masks, detection_boxes, height, width
+ )
+ detection_masks_reframed = tf.cast(detection_masks_reframed > 0.8, np.uint8)
+ result['detection_masks_reframed'] = detection_masks_reframed.numpy()
+ return result
+
+
+def find_id_by_name(
+ dictionary: Dict[int, ItemDict], name: str
+) -> Optional[int]:
+ """Finds the id of a dictionary given its value.
+
+ Args:
+ dictionary: The dictionary containing the data.
+ name: The value to find.
+
+ Returns:
+ The id, or None if its not found.
+ """
+
+ # Iterate over the dictionary, and check if the name of the user
+ # matches the name that was passed in.
+ for value in dictionary.values():
+ if value['name'] == name:
+ # If the name matches, return the id of the user.
+ return value['id']
+
+ return None
+
+
+def combine_bounding_boxes(
+ box1: List[float],
+ box2: List[float],
+) -> List[float]:
+ """Combines two bounding boxes.
+
+ Args:
+ box1: A list of four numbers representing the coordinates of the first
+ bounding box.
+ box2: A list of four numbers representing the coordinates of the second
+ bounding box.
+
+ Returns:
+ A list of four numbers representing the coordinates of the combined
+ bounding box.
+ """
+
+ ymin = min(box1[0], box2[0])
+ xmin = min(box1[1], box2[1])
+ ymax = max(box1[2], box2[2])
+ xmax = max(box1[3], box2[3])
+
+ return [ymin, xmin, ymax, xmax]
+
+
+def calculate_combined_scores_boxes_classes(
+ i: int,
+ j: int,
+ results_1: DetectionResult,
+ results_2: DetectionResult,
+ category_indices: List[List[Any]],
+ category_index_combined: Dict[int, ItemDict],
+) -> Tuple[Any, List[float], Any, Optional[int]]:
+ """Calculate combined scores, boxes, and classes for matched masks.
+
+ Args:
+ i: Index of the mask from the results_1.
+ j: Index of the mask from the results_2.
+ results_1: A dictionary which contains the results from the first model.
+ results_2: A dictionary which contains the results from the second model.
+ category_indices: list of sub lists which contains the labels of 1st and
+ 2nd ML model.
+ category_index_combined: Combined category index.
+
+ Returns:
+ tuple: A tuple containing:
+ - avg_score: Average score of the matched masks.
+ - combined_box: Combined bounding box for the matched masks.
+ - combined_label: Combined label of the matched masks.
+ - result_id: ID associated with the combined label.
+ """
+ score_1 = results_1['detection_scores'][0][i]
+ score_2 = results_2['detection_scores'][0][j]
+ avg_score = (score_1 + score_2) / 2
+
+ box_1 = results_1['detection_boxes'][0][i]
+ box_2 = results_2['detection_boxes'][0][j]
+ combined_box = combine_bounding_boxes(box_1, box_2)
+
+ class_1 = results_1['detection_classes'][0][i]
+ class_2 = results_2['detection_classes'][0][j]
+ combined_label = (
+ category_indices[0][class_1] + '_' + category_indices[1][class_2]
+ )
+ result_id = find_id_by_name(category_index_combined, combined_label)
+
+ return avg_score, combined_box, combined_label, result_id
+
+
+def calculate_single_result(
+ index: int,
+ result: DetectionResult,
+ category_indices: List[List[Any]],
+ flag: Any | str,
+) -> Tuple[float, Tuple[float, float, float, float], str]:
+ """Calculate scores, boxes, and classes for non-matched masks.
+
+ Args:
+ index: Index of the mask in the results.
+ result: A dictionary containing detection results (either from results_1
+ or results_2).
+ category_indices: list of category indices.
+ flag: To identify whose model did not detected an object.
+
+ Returns:
+ score: Score of the mask.
+ box: Bounding box of the mask.
+ combined_label: Label of the mask with the added suffix.
+ """
+ combined_label = 'Default Value'
+ score = result['detection_scores'][0][index]
+ box = result['detection_boxes'][0][index]
+ class_idx = result['detection_classes'][0][index]
+ if flag == 'after':
+ combined_label = category_indices[class_idx] + '_Na' # pyrefly: ignore[unsupported-operation]
+ elif flag == 'before':
+ combined_label = 'Na_' + category_indices[class_idx] # pyrefly: ignore[unsupported-operation]
+ return score, box, combined_label
+
+
+def calculate_iou(
+ mask1: np.ndarray, mask2: np.ndarray
+) -> Tuple[float, np.ndarray]:
+ """Calculates the intersection over union (IoU) score for two masks.
+
+ Args:
+ mask1: The first mask.
+ mask2: The second mask.
+
+ Returns:
+ The IoU scorea and union of two masks.
+ """
+
+ # Check if the masks have the same dimensions.
+ if mask1.shape != mask2.shape:
+ raise ValueError('The masks must have the same dimensions.')
+
+ intersection = np.logical_and(mask1, mask2)
+ union = np.logical_or(mask1, mask2)
+ iou_score = np.sum(intersection) / np.sum(union)
+ return iou_score, union
+
+
+def find_similar_masks(
+ results_1: DetectionResult,
+ results_2: DetectionResult,
+ num_detections: int,
+ min_score_thresh: float,
+ category_indices: List[List[Any]],
+ category_index_combined: Dict[int, ItemDict],
+ area_threshold: float,
+ iou_threshold: float = 0.8,
+) -> Dict[str, np.ndarray]:
+ """Aligns the masks of the detections in `results_1` and `results_2`.
+
+ Args:
+ results_1: A dictionary which contains the results from the first model.
+ results_2: A dictionary which contains the results from the second model.
+ num_detections: The number of detections to consider.
+ min_score_thresh: The minimum score threshold for a detection
+ category_indices: list of sub lists which contains the labels of 1st and 2nd
+ ML model
+ category_index_combined: A dictionary with an object ID and nested
+ dictionary with name. e.g. {1: {'id': 1, 'name': 'Fiber_Na_Bag',
+ 'supercategory': 'objects'}}
+ area_threshold: Threshold for mask area consideration.
+ iou_threshold: IOU threshold to compare masks.
+
+ Returns:
+ A dictionary containing the following keys:
+ - num_detections: The number of aligned detections.
+ - detection_classes: A NumPy array of shape (num_detections,) containing
+ the classes for the aligned detections.
+ - detection_scores: A NumPy array of shape (num_detections,) containing
+ the scores for the aligned detections.
+ - detection_boxes: A NumPy array of shape (num_detections, 4) containing
+ the bounding boxes for the aligned detections.
+ - detection_classes_names: A list of strings containing the names of the
+ classes for the aligned detections.
+ - detection_masks_reframed: A NumPy array of shape (num_detections,
+ height, width) containing the full masks for the aligned detections.
+ """
+ detection_masks_reframed = []
+ detection_scores = []
+ detection_boxes = []
+ detection_classes = []
+ detection_classes_names = []
+
+ aligned_masks = 0
+ masks_list1 = results_1['detection_masks_reframed'][:num_detections]
+ masks_list2 = results_2['detection_masks_reframed'][:num_detections]
+ scores_list1 = results_1['detection_scores'][0]
+ scores_list2 = results_2['detection_scores'][0]
+ matched_masks_list2 = [False] * len(masks_list2)
+ matched_masks_list1 = [False] * len(masks_list1)
+
+ for i, mask1 in enumerate(masks_list1):
+ if (scores_list1[i] > min_score_thresh) and (
+ np.sum(mask1) < area_threshold
+ ):
+ is_similar = False
+
+ for j, mask2 in enumerate(masks_list2):
+ if scores_list2[j] > min_score_thresh and (
+ np.sum(mask2) < area_threshold
+ ):
+ iou, union = calculate_iou(mask1, mask2)
+
+ # masks which are present both in the 'detection_masks_reframed'
+ # key of 'results_1' & 'results_2' dictionary
+ if iou > iou_threshold:
+ aligned_masks += 1
+ is_similar = True
+ matched_masks_list2[j] = True
+ matched_masks_list1[i] = True
+
+ detection_masks_reframed.append(union)
+
+ avg_score, combined_box, combined_label, result_id = (
+ calculate_combined_scores_boxes_classes(
+ i,
+ j,
+ results_1,
+ results_2,
+ category_indices,
+ category_index_combined,
+ )
+ )
+ detection_scores.append(avg_score)
+ detection_boxes.append(combined_box)
+ detection_classes_names.append(combined_label)
+ detection_classes.append(result_id)
+ break
+
+ # masks which are only present in the 'detection_masks_reframed'
+ # of 'results_1' dictionary
+ if not is_similar:
+ aligned_masks += 1
+ detection_masks_reframed.append(mask1)
+ score, box, combined_label = calculate_single_result(
+ i, results_1, category_indices[0], 'after'
+ )
+ detection_scores.append(score)
+ detection_boxes.append(box)
+ detection_classes_names.append(combined_label)
+ result_id = find_id_by_name(category_index_combined, combined_label)
+ detection_classes.append(result_id)
+
+ # masks which are only present in the 'detection_masks_reframed'
+ # key of 'results_2' dictionary
+ for k, mask2 in enumerate(masks_list2):
+ if (
+ (not matched_masks_list2[k])
+ and (scores_list2[k] > min_score_thresh)
+ and (np.sum(mask2) < area_threshold)
+ ):
+ aligned_masks += 1
+ detection_masks_reframed.append(mask2)
+ score, box, combined_label = calculate_single_result(
+ k, results_2, category_indices[1], 'before'
+ )
+ detection_scores.append(score)
+ detection_boxes.append(box)
+ detection_classes_names.append(combined_label)
+ result_id = find_id_by_name(category_index_combined, combined_label)
+ detection_classes.append(result_id)
+
+ final_result = {
+ 'num_detections': np.array([aligned_masks]),
+ 'detection_classes': np.array(detection_classes),
+ 'detection_scores': np.array([detection_scores]),
+ 'detection_boxes': np.array([detection_boxes]),
+ 'detection_classes_names': np.array(detection_classes_names),
+ 'detection_masks_reframed': np.array(detection_masks_reframed),
+ }
+
+ return final_result
+
+
+def filter_bounding_boxes(
+ bounding_boxes: List[Tuple[int, int, int, int]],
+ iou_threshold: float = 0.5,
+ area_ratio_threshold: float = 0.8,
+) -> Tuple[List[Tuple[int, int, int, int]], List[int]]:
+ """Filters overlapping bounding boxes based on IoU and area ratio criteria.
+
+ This function filters out overlapping bounding boxes from a given list based
+ on Intersection over Union (IoU) and area ratio of the intersection to the
+ smaller bounding box's area.
+
+ Args:
+ bounding_boxes: A list of bounding boxes, where each bounding box is
+ represented as a tuple of (xmin, ymin, xmax, ymax).
+ iou_threshold: Threshold for Intersection over Union. Bounding boxes with
+ IoU above this threshold will be considered overlapping. Defaults to
+ 0.5.
+ area_ratio_threshold: Threshold for the area ratio of the intersection to
+ the smaller bounding box's area. Defaults to 0.8.
+
+ Returns:
+ tuple: A tuple containing:
+ - filtered_boxes: A list of bounding boxes that passed the filtering
+ criteria.
+ - eliminated_indices: Indices of the bounding boxes that didn't pass
+ the filtering criteria.
+
+ Example:
+ >>> bounding_boxes = [(10, 10, 50, 50), (20, 20, 60, 60)]
+ >>> filter_bounding_boxes(bounding_boxes)
+ ([(10, 10, 50, 50)], [1])
+ """
+ filtered_boxes = []
+ eliminated_indices = []
+
+ # Enumerate and sort the boxes based on their area in descending order
+ enumerated_boxes = list(enumerate(bounding_boxes))
+ sorted_boxes = sorted(
+ enumerated_boxes,
+ key=lambda item: (item[1][2] - item[1][0]) * (item[1][3] - item[1][1]),
+ reverse=True,
+ )
+
+ for idx, bbox in sorted_boxes:
+ skip_box = False
+
+ # Calculate areas of individual bounding boxes
+ area_bbox = (bbox[2] - bbox[0]) * (bbox[3] - bbox[1])
+
+ for jdx, other_bbox in sorted_boxes:
+ if idx == jdx:
+ continue
+
+ # Calculate intersection coordinates
+ xmin_inter = max(bbox[0], other_bbox[0])
+ ymin_inter = max(bbox[1], other_bbox[1])
+ xmax_inter = min(bbox[2], other_bbox[2])
+ ymax_inter = min(bbox[3], other_bbox[3])
+
+ # Calculate intersection area
+ width_inter = max(0, xmax_inter - xmin_inter)
+ height_inter = max(0, ymax_inter - ymin_inter)
+ area_inter = width_inter * height_inter
+
+ area_other_bbox = (other_bbox[2] - other_bbox[0]) * (
+ other_bbox[3] - other_bbox[1]
+ )
+
+ # Calculate area ratio
+ area_ratio = area_inter / min(area_bbox, area_other_bbox)
+
+ # Check for overlapping and area ratio thresholds
+ if area_ratio > area_ratio_threshold:
+ if area_bbox > area_other_bbox:
+ skip_box = True
+ eliminated_indices.append(idx)
+ break
+ elif (
+ area_inter > 0
+ and area_inter / (area_bbox + area_other_bbox - area_inter)
+ > iou_threshold
+ ):
+ if area_bbox > area_other_bbox:
+ skip_box = True
+ eliminated_indices.append(idx)
+ break
+
+ if not skip_box:
+ filtered_boxes.append(bbox)
+
+ return filtered_boxes, eliminated_indices
diff --git a/official/projects/waste_identification_ml/model_inference/postprocessing_test.py b/official/projects/waste_identification_ml/model_inference/postprocessing_test.py
new file mode 100644
index 00000000000..771b8feb0b3
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/postprocessing_test.py
@@ -0,0 +1,33 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import tensorflow as tf, tf_keras
+from official.projects.waste_identification_ml.model_inference import postprocessing
+
+
+class PostprocessingTest(tf.test.TestCase):
+
+ def test_reframe_image_corners_relative_to_boxes(self):
+ """Test if the function correctly reframes the image corners relative to the boxes."""
+ box1 = tf.constant([100.0, 200.0, 300.0, 400.0])
+ boxes = tf.stack([box1])
+ expected = tf.stack([tf.constant([-0.5, -1.0, -0.495, -0.995])])
+
+ result = postprocessing.reframe_image_corners_relative_to_boxes(boxes)
+
+ self.assertAllEqual(expected, result)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/waste_identification_ml/model_inference/preprocessing.py b/official/projects/waste_identification_ml/model_inference/preprocessing.py
new file mode 100644
index 00000000000..54a63fbd68d
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/preprocessing.py
@@ -0,0 +1,75 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""This module provides utilities to normalize image tensors.
+"""
+from typing import Sequence
+import tensorflow as tf, tf_keras
+
+MEAN_NORM = (0.485, 0.456, 0.406)
+STDDEV_NORM = (0.229, 0.224, 0.225)
+
+
+def normalize_image(
+ image: tf.Tensor,
+ offset: Sequence[float] = MEAN_NORM,
+ scale: Sequence[float] = STDDEV_NORM,
+) -> tf.Tensor:
+ """Normalizes the image to zero mean and unit variance.
+
+ If the input image dtype is float, it is expected to either have values in
+ [0, 1) and offset is MEAN_NORM, or have values in [0, 255] and offset is
+ MEAN_RGB.
+
+ Args:
+ image: A tf.Tensor in either (1) float dtype with values in range [0, 1) or
+ [0, 255], or (2) int type with values in range [0, 255].
+ offset: A tuple of mean values to be subtracted from the image.
+ scale: A tuple of normalization factors.
+
+ Returns:
+ A normalized image tensor.
+ """
+ image = tf.image.convert_image_dtype(image, dtype=tf.float32)
+ return normalize_scaled_float_image(image, offset, scale)
+
+
+def normalize_scaled_float_image(
+ image: tf.Tensor,
+ offset: Sequence[float] = MEAN_NORM,
+ scale: Sequence[float] = STDDEV_NORM,
+):
+ """Normalizes a scaled float image to zero mean and unit variance.
+
+ It assumes the input image is float dtype with values in [0, 1) if offset is
+ MEAN_NORM, values in [0, 255] if offset is MEAN_RGB.
+
+ Args:
+ image: A tf.Tensor in float32 dtype with values in range [0, 1) or [0, 255].
+ offset: A tuple of mean values to be subtracted from the image.
+ scale: A tuple of normalization factors.
+
+ Returns:
+ A normalized image tensor.
+ """
+ offset = tf.constant(offset)
+ offset = tf.expand_dims(offset, axis=0)
+ offset = tf.expand_dims(offset, axis=0)
+ image -= offset
+
+ scale = tf.constant(scale)
+ scale = tf.expand_dims(scale, axis=0)
+ scale = tf.expand_dims(scale, axis=0)
+ image /= scale
+ return image
diff --git a/official/projects/waste_identification_ml/model_inference/preprocessing_test.py b/official/projects/waste_identification_ml/model_inference/preprocessing_test.py
new file mode 100644
index 00000000000..025a538c32f
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/preprocessing_test.py
@@ -0,0 +1,71 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import tensorflow as tf, tf_keras
+from official.projects.waste_identification_ml.model_inference import preprocessing
+
+
+class PreprocessingTest(tf.test.TestCase):
+
+ def test_normalize_image(self):
+ image = tf.convert_to_tensor(
+ [[[1, 2, 3], [4, 5, 6]], [[7, 8, 9], [10, 11, 12]]], dtype=tf.int8
+ )
+ expected = tf.convert_to_tensor(
+ value=[
+ [
+ [-2.0835197, -1.9654106, -1.6994576],
+ [-1.9803666, -1.859955, -1.5944706],
+ ],
+ [
+ [-1.8772135, -1.7544994, -1.4894838],
+ [-1.7740605, -1.6490439, -1.384497],
+ ],
+ ],
+ dtype=tf.float32,
+ )
+
+ result = preprocessing.normalize_image(image=image)
+
+ self.assertAllEqual(expected, result)
+
+ def test_normalize_scaled_float_image(self):
+ image = tf.convert_to_tensor(
+ [
+ [[0.00787, 0.01575, 0.02362], [0.0315, 0.03937, 0.04724]],
+ [[0.0551, 0.063, 0.07086], [0.07874, 0.0866, 0.0945]],
+ ],
+ dtype=tf.float32,
+ )
+ expected = tf.convert_to_tensor(
+ value=[
+ [
+ [-2.0835197, -1.9654106, -1.6994576],
+ [-1.9803666, -1.859955, -1.5944706],
+ ],
+ [
+ [-1.8772135, -1.7544994, -1.4894838],
+ [-1.7740605, -1.6490439, -1.384497],
+ ],
+ ],
+ dtype=tf.float32,
+ )
+
+ result = preprocessing.normalize_scaled_float_image(image=image)
+
+ self.assertAllCloseAccordingToType(expected, result, rtol=1e-4)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/waste_identification_ml/model_inference/testdata/categories_1.csv b/official/projects/waste_identification_ml/model_inference/testdata/categories_1.csv
new file mode 100644
index 00000000000..2b80d2fa55f
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/testdata/categories_1.csv
@@ -0,0 +1,3 @@
+category
+alpha
+beta
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/model_inference/testdata/categories_2.csv b/official/projects/waste_identification_ml/model_inference/testdata/categories_2.csv
new file mode 100644
index 00000000000..d9ea95e492b
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/testdata/categories_2.csv
@@ -0,0 +1,3 @@
+category
+gamma
+delta
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/model_inference/testdata/csv_to_list.csv b/official/projects/waste_identification_ml/model_inference/testdata/csv_to_list.csv
new file mode 100644
index 00000000000..6e45e14f5e0
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference/testdata/csv_to_list.csv
@@ -0,0 +1,4 @@
+data
+alpha
+beta
+gamma
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/Inference_Detectron2_Pipeline_with_Tracking_experimental.ipynb b/official/projects/waste_identification_ml/model_inference_with_tracking/Inference_Detectron2_Pipeline_with_Tracking_experimental.ipynb
new file mode 100644
index 00000000000..e94a9a34607
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/Inference_Detectron2_Pipeline_with_Tracking_experimental.ipynb
@@ -0,0 +1,1002 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "TtlIRiNXWlQ0"
+ },
+ "source": [
+ "# Waste identification with instance segmentation in PyTorch"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ohoMgYgXWsIO"
+ },
+ "source": [
+ "This Colab notebook demonstrates an end-to-end pipeline for object detection, feature extraction, object tracking, and data aggregation using the Mask R-CNN model from Detectron2.\n",
+ "Key Steps in the Notebook:\n",
+ "\n",
+ "\n",
+ "\n",
+ "* Object Detection and segmentation – Detect objects in a set of images using Mask R-CNN.\n",
+ "* Feature Extraction \u0026 Tracking – Extract object features and track them across multiple frames to eliminate duplicate counts.\n",
+ "* Color Detection – Identify the color of each detected object.\n",
+ "* Postprocessing – Aggregate tracking results and apply filtering to reduce false positives and false negatives.\n",
+ "* Save detection and tracking results.\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "8PKG9z4VYPEs"
+ },
+ "source": [
+ "To finish this task, a proper path for the trained model and images need to be provided. The path to the labels on which the models are trained is in the waste_identification_ml directory inside the Tensorflow Model Garden repository."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "o2J5IWCgAOxC"
+ },
+ "source": [
+ "This notebook will output 3 folders and 1 csv file :\n",
+ "\n",
+ "\n",
+ "* **prediction_folder** : Will contain prediction results with bbox and masks.\n",
+ "* **tracking** : Will contain tracking visualization.\n",
+ "* **cropped_objects** : Will contain category level detected objects.\n",
+ "* **count.csv** : Will contain the individual counts of each category.\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "98YSgFKZaKVe"
+ },
+ "source": [
+ "## Install Detectron2 and RESTART the runtime"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "aHZ87DkPaN_A"
+ },
+ "outputs": [],
+ "source": [
+ "!git clone 'https://github.com/facebookresearch/detectron2'\n",
+ "!pip install 'git+https://github.com/facebookresearch/detectron2.git'"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "ELUFMVDDAopS"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Imports and Setup\n",
+ "\n",
+ "!pip install -q supervision trackpy openpyxl==3.1.2\n",
+ "\n",
+ "import sys\n",
+ "import tensorflow as tf\n",
+ "import csv\n",
+ "from typing import Any, TypedDict, Callable\n",
+ "import cv2\n",
+ "import logging\n",
+ "import numpy as np\n",
+ "import matplotlib.pyplot as plt\n",
+ "import glob\n",
+ "import natsort\n",
+ "import tqdm\n",
+ "import os\n",
+ "from PIL import Image\n",
+ "from scipy import ndimage\n",
+ "import pandas as pd\n",
+ "import skimage\n",
+ "import datetime\n",
+ "import trackpy as tp\n",
+ "import shutil\n",
+ "import supervision as sv\n",
+ "\n",
+ "# Detectron2 Utilities\n",
+ "import torch\n",
+ "import detectron2\n",
+ "from detectron2.utils.logger import setup_logger\n",
+ "from detectron2.engine import DefaultPredictor\n",
+ "from detectron2.config import get_cfg\n",
+ "from detectron2.structures import Instances, Boxes\n",
+ "from detectron2.data.catalog import Metadata\n",
+ "from detectron2.utils.visualizer import Visualizer\n",
+ "setup_logger()\n",
+ "\n",
+ "\n",
+ "logging.disable(logging.WARNING)\n",
+ "\n",
+ "%matplotlib inline"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "g5_XxVSFATYZ"
+ },
+ "outputs": [],
+ "source": [
+ "# Connect to Google drive if your data is stored there.\n",
+ "from google.colab import drive\n",
+ "drive.mount('/content/gdrive')\n",
+ "\n",
+ "try:\n",
+ " !ln -s /content/gdrive/My\\ Drive/ /mydrive\n",
+ " print('Successful')\n",
+ "except Exception as e:\n",
+ " print(e)\n",
+ " print('Not successful')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WAQ8ymV-Cgqp"
+ },
+ "outputs": [],
+ "source": [
+ "# # Connect to GCP bucket if your data is store there and copy them locally.\n",
+ "# !gcloud init\n",
+ "# gcloud storage cp --recursive gs://input ."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "_77YK3a_BCg_"
+ },
+ "source": [
+ "To visualize the images with the proper detected boxes and segmentation masks, we will use the TensorFlow Object Detection API. To install it we will clone the repo.\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qhk_NujKO0mb"
+ },
+ "outputs": [],
+ "source": [
+ "# Clone the tensorflow models repository.\n",
+ "!git clone --depth 1 https://github.com/tensorflow/models 2\u003e/dev/null"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "o1dYyG55BtWb"
+ },
+ "outputs": [],
+ "source": [
+ "sys.path.append('models/official/projects/waste_identification_ml/model_inference/')\n",
+ "import color_and_property_extractor"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "mvErM8kUirE2"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Utilities\n",
+ "\n",
+ "_PROPERTIES = (\n",
+ " 'area',\n",
+ " 'bbox',\n",
+ " 'convex_area',\n",
+ " 'bbox_area',\n",
+ " 'major_axis_length',\n",
+ " 'minor_axis_length',\n",
+ " 'eccentricity',\n",
+ " 'centroid',\n",
+ " 'label',\n",
+ " 'mean_intensity',\n",
+ " 'max_intensity',\n",
+ " 'min_intensity',\n",
+ " 'perimeter'\n",
+ ")\n",
+ "\n",
+ "def convert_detections_to_instances(\n",
+ " outputs: dict,\n",
+ " image_size: tuple[int, int] = (1024, 1024),\n",
+ " nms_threshold: float = 0.8,\n",
+ " class_agnostic: bool = True\n",
+ ") -\u003e dict[str, Instances]:\n",
+ " \"\"\"Convert Detectron2 model outputs to an Instances object with Non-Maximum Suppression (NMS) applied.\n",
+ "\n",
+ " Args:\n",
+ " outputs: Detectron2 model output containing instance predictions.\n",
+ " image_size: Image dimensions (height, width).\n",
+ " nms_threshold: Non-Maximum Suppression (NMS) threshold.\n",
+ " class_agnostic: Whether NMS should be applied in a class-agnostic manner.\n",
+ "\n",
+ " Returns:\n",
+ " Reformatted Detectron2 output as {\"instances\": Instances}.\n",
+ " \"\"\"\n",
+ " # Apply NMS and convert to supervision Detections format\n",
+ " detections = (\n",
+ " sv.Detections.from_detectron2(outputs)\n",
+ " .with_nms(threshold=nms_threshold, class_agnostic=class_agnostic)\n",
+ " )\n",
+ "\n",
+ " # Convert extracted values to PyTorch tensors\n",
+ " bboxes = torch.tensor(detections.xyxy, dtype=torch.float32)\n",
+ " scores = torch.tensor(detections.confidence, dtype=torch.float32)\n",
+ " classes = torch.tensor(detections.class_id, dtype=torch.int64)\n",
+ "\n",
+ " # Create an Instances object\n",
+ " output_instances = Instances(image_size)\n",
+ " output_instances.set(\"pred_boxes\", Boxes(bboxes))\n",
+ " output_instances.set(\"scores\", scores)\n",
+ " output_instances.set(\"pred_classes\", classes)\n",
+ "\n",
+ " # Add masks if available\n",
+ " if detections.mask is not None:\n",
+ " masks = torch.tensor(detections.mask, dtype=torch.uint8)\n",
+ " output_instances.set(\"pred_masks\", masks)\n",
+ "\n",
+ " return {\"instances\": output_instances}\n",
+ "\n",
+ "\n",
+ "def read_csv(file_path: str) -\u003e list[str]:\n",
+ " \"\"\"Reads a CSV file and returns its contents as a list.\n",
+ "\n",
+ " This function reads the given CSV file, skips the header, and assumes\n",
+ " there is only one column in the CSV. It returns the contents as a list of\n",
+ " strings.\n",
+ "\n",
+ " Args:\n",
+ " file_path: The path to the CSV file.\n",
+ "\n",
+ " Returns:\n",
+ " The contents of the CSV file as a list of strings.\n",
+ " \"\"\"\n",
+ " data_list = []\n",
+ " with open(file_path, 'r') as csvfile:\n",
+ " reader = csv.reader(csvfile)\n",
+ " for row in reader:\n",
+ " data_list.append(row[0])\n",
+ " return data_list\n",
+ "\n",
+ "\n",
+ "def adjust_image_size(height: int, width: int, min_size: int) -\u003e tuple[int, int]:\n",
+ " \"\"\"Adjust the image size to ensure both dimensions are at least 1024.\n",
+ "\n",
+ " Args:\n",
+ " height: The height of the image.\n",
+ " width: The width of the image.\n",
+ " min_size: Minimum size of the image dimension needed.\n",
+ "\n",
+ " Returns:\n",
+ " The adjusted height and width of the image.\n",
+ " \"\"\"\n",
+ " if height \u003c min_size or width \u003c min_size:\n",
+ " return height, width\n",
+ "\n",
+ " # Calculate the scale factor to ensure both dimensions remain at least 1024\n",
+ " scale_factor = min(height / min_size, width / min_size)\n",
+ "\n",
+ " new_height = int(height / scale_factor)\n",
+ " new_width = int(width / scale_factor)\n",
+ "\n",
+ " return new_height, new_width\n",
+ "\n",
+ "\n",
+ "def dilated_largest_component(mask: np.ndarray) -\u003e np.ndarray:\n",
+ " \"\"\"Extracts the largest connected component and fills holes.\n",
+ "\n",
+ " Args:\n",
+ " mask: Input binary mask (2D numpy array).\n",
+ "\n",
+ " Returns:\n",
+ " Binary mask of the largest connected component.\n",
+ " \"\"\"\n",
+ " mask = mask.astype(np.uint8)*255\n",
+ " num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask, connectivity=8)\n",
+ " largest_label = 1 + np.argmax(stats[1:, cv2.CC_STAT_AREA])\n",
+ " largest_component_mask = np.zeros(mask.shape, dtype=\"uint8\")\n",
+ " largest_component_mask[labels == largest_label] = 1\n",
+ " largest_component_mask = ndimage.binary_fill_holes(largest_component_mask).astype(int)\n",
+ " return largest_component_mask\n",
+ "\n",
+ "\n",
+ "def extract_properties(image, masks):\n",
+ " \"\"\"Extract properties of the mask.\n",
+ "\n",
+ " Args:\n",
+ " image: Corresponding image of the mask.\n",
+ " masks: The masks to extract properties from.\n",
+ "\n",
+ " Returns:\n",
+ " The extracted properties.\n",
+ " \"\"\"\n",
+ " list_of_df = []\n",
+ " for mask in masks:\n",
+ " mask = np.where(mask, 1, 0)\n",
+ " df = pd.DataFrame(\n",
+ " skimage.measure.regionprops_table(mask, intensity_image=image, properties=_PROPERTIES)\n",
+ " )\n",
+ " list_of_df.append(df)\n",
+ " features = pd.concat(list_of_df, ignore_index=True)\n",
+ " features.rename(\n",
+ " columns={\n",
+ " 'centroid-0': 'y',\n",
+ " 'centroid-1': 'x',\n",
+ " 'bbox-0': 'bbox_0',\n",
+ " 'bbox-1': 'bbox_1',\n",
+ " 'bbox-2': 'bbox_2',\n",
+ " 'bbox-3': 'bbox_3',\n",
+ " },\n",
+ " inplace=True,\n",
+ " )\n",
+ " return features\n",
+ "\n",
+ "\n",
+ "def get_image_creation_time(image_path):\n",
+ " \"\"\"\n",
+ " Retrieves the creation time of an image, trying multiple methods.\n",
+ "\n",
+ " Args:\n",
+ " image_path: The path to the image file.\n",
+ "\n",
+ " Returns:\n",
+ " A string representing the creation time in the format \"%Y-%m-%d %H:%M:%S\" if\n",
+ " found, otherwise returns \"Creation time not found\".\n",
+ " \"\"\"\n",
+ "\n",
+ " try:\n",
+ " # 1. Try EXIF data (if available)\n",
+ " image = Image.open(image_path)\n",
+ " exif_data = image._getexif()\n",
+ " if exif_data:\n",
+ " datetime_tag_id = 36867 # Tag ID for \"DateTimeOriginal\"\n",
+ " datetime_str = exif_data.get(datetime_tag_id)\n",
+ " if datetime_str:\n",
+ " datetime_obj = datetime.datetime.strptime(datetime_str, \"%Y:%m:%d %H:%M:%S\")\n",
+ " return datetime_obj.strftime(\"%Y-%m-%d %H:%M:%S\")\n",
+ "\n",
+ " # 2. Try file modification time (less accurate, but better than nothing)\n",
+ " file_modified_time = os.path.getmtime(image_path)\n",
+ " datetime_obj = datetime.datetime.fromtimestamp(file_modified_time)\n",
+ " return datetime_obj.strftime(\"%Y-%m-%d %H:%M:%S\")\n",
+ "\n",
+ " except FileNotFoundError:\n",
+ " return \"Image not found\"\n",
+ " except Exception as e:\n",
+ " return f\"Error: {e}\"\n",
+ "\n",
+ "\n",
+ "def process_tracking_result(df):\n",
+ " \"\"\"Process the tracking result dataframe.\n",
+ "\n",
+ " Args:\n",
+ " df: Dataframe to be aggregated.\n",
+ "\n",
+ " Returns:\n",
+ " Processed dataframe.\n",
+ " \"\"\"\n",
+ " # Get class information with the new include_groups parameter\n",
+ " class_info = df.groupby('particle', as_index=False).apply(\n",
+ " select_class_with_scores,\n",
+ " include_groups=False\n",
+ " )\n",
+ "\n",
+ " grouped = df.groupby('particle').agg({\n",
+ " 'source_name': 'first',\n",
+ " 'image_name': 'first',\n",
+ " 'detection_scores': 'max',\n",
+ " 'creation_time': 'first',\n",
+ " 'bbox_0': 'first',\n",
+ " 'bbox_1': 'first',\n",
+ " 'bbox_2': 'first',\n",
+ " 'bbox_3': 'first',\n",
+ " }).reset_index()\n",
+ "\n",
+ " # Add class information\n",
+ " grouped['detection_classes'] = class_info['class_id']\n",
+ " grouped['detection_classes_names'] = class_info['class_name']\n",
+ "\n",
+ " return grouped\n",
+ "\n",
+ "def select_class_with_scores(group):\n",
+ " \"\"\"\n",
+ " Select class based on modal class, falling back to highest score for ties.\n",
+ " Returns both class ID and class name.\n",
+ " \"\"\"\n",
+ " # Get the value counts of classes\n",
+ " class_counts = group['detection_classes'].value_counts()\n",
+ "\n",
+ " #print('class counts', class_counts)\n",
+ "\n",
+ " # If there's a clear winner (one mode), use it\n",
+ " if len(class_counts) == 1 or class_counts.iloc[0] \u003e class_counts.iloc[1]:\n",
+ " class_id = group['detection_classes'].mode().iloc[0]\n",
+ " else:\n",
+ " # If there's a tie, look at highest score for each tied class\n",
+ " tied_classes = class_counts[class_counts == class_counts.iloc[0]].index\n",
+ " #print('tied classes', tied_classes)\n",
+ " class_max_scores = {\n",
+ " cls: group[group['detection_classes'] == cls]['detection_scores'].max()\n",
+ " for cls in tied_classes\n",
+ " }\n",
+ " #print('class max scores', class_max_scores)\n",
+ " class_id = max(class_max_scores.items(), key=lambda x: x[1])[0]\n",
+ "\n",
+ " # Get corresponding class name\n",
+ " class_name = group[group['detection_classes'] == class_id]['detection_classes_names'].iloc[0]\n",
+ " #print('winner', pd.Series({'class_id': class_id, 'class_name': class_name}))\n",
+ " return pd.Series({'class_id': class_id, 'class_name': class_name})\n",
+ "\n",
+ "\n",
+ "def apply_tracking(df,\n",
+ " search_range_x,\n",
+ " search_range_y,\n",
+ " memory):\n",
+ " \"\"\"Apply tracking to the dataframe.\n",
+ "\n",
+ " Args:\n",
+ " df: The dataframe to apply tracking to.\n",
+ " search_range_x: The search range of pixels for tracking along x axis.\n",
+ " search_range_y: The search range of pixels for tracking along y axis.\n",
+ " memory: The frames memory for tracking.\n",
+ "\n",
+ " Returns:\n",
+ " The tracking result dataframe.\n",
+ " \"\"\"\n",
+ " # Define the columns to link for tracking.\n",
+ " # Additional features that can be used are 'area', 'label', 'color',\n",
+ " # 'eccentricity', 'convex_area', 'mean_intensity-0', 'mean_intensity-1',\n",
+ " # 'mean_intensity-2', 'max_intensity-0', 'max_intensity-1', 'max_intensity-2',\n",
+ " # 'min_intensity-0', 'min_intensity-1', 'min_intensity-2'.\n",
+ " tracking_columns = [\n",
+ " 'x',\n",
+ " 'y',\n",
+ " 'frame',\n",
+ " 'bbox_0',\n",
+ " 'bbox_1',\n",
+ " 'bbox_2',\n",
+ " 'bbox_3',\n",
+ " 'major_axis_length',\n",
+ " 'minor_axis_length',\n",
+ " 'perimeter',\n",
+ " ]\n",
+ "\n",
+ " # Perform the tracking operation on the specified columns\n",
+ " track_df = tp.link_df(df[tracking_columns], search_range=(search_range_y, search_range_x), memory=memory)\n",
+ "\n",
+ " # Copy the additional columns from the original dataframe\n",
+ " additional_columns = [\n",
+ " 'source_name',\n",
+ " 'image_name',\n",
+ " 'detection_scores',\n",
+ " 'detection_classes_names',\n",
+ " 'detection_classes',\n",
+ " 'color',\n",
+ " 'creation_time'\n",
+ " ]\n",
+ " track_df[additional_columns] = df[additional_columns]\n",
+ "\n",
+ " track_df.drop(columns=['frame'], inplace=True)\n",
+ " track_df.reset_index(drop=True, inplace=True)\n",
+ "\n",
+ " return track_df\n",
+ "\n",
+ "\n",
+ "def resize_bbox(y1, x1, y2, x2, old_height, old_width, new_height, new_width):\n",
+ " \"\"\"Resize bounding box coordinates based on new image size.\n",
+ "\n",
+ " Args:\n",
+ " y1, x1, y2, x2 (int/float): Original bounding box coordinates.\n",
+ " old_height, old_width (int): Original image dimensions.\n",
+ " new_height, new_width (int): New image dimensions.\n",
+ "\n",
+ " Returns:\n",
+ " (new_y1, new_x1, new_y2, new_x2): Rescaled bounding box coordinates.\n",
+ " \"\"\"\n",
+ " # Compute scale factors\n",
+ " scale_x = new_width / old_width\n",
+ " scale_y = new_height / old_height\n",
+ "\n",
+ " # Scale bounding box coordinates\n",
+ " new_y1 = int(y1 * scale_y)\n",
+ " new_x1 = int(x1 * scale_x)\n",
+ " new_y2 = int(y2 * scale_y)\n",
+ " new_x2 = int(x2 * scale_x)\n",
+ "\n",
+ " return new_y1, new_x1, new_y2, new_x2\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "4XjfDEq--UlE"
+ },
+ "source": [
+ "## Import and load pre-trained models."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ZQ435YHN3Lr-"
+ },
+ "outputs": [],
+ "source": [
+ "%%bash\n",
+ "wget https://storage.googleapis.com/tf_model_garden/vision/\\\n",
+ "waste_identification_ml/Detectron2_Jan2025_1024_1024.zip\n",
+ "\n",
+ "unzip Detectron2_Jan2025_1024_1024.zip \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "W6mmyLsOJicF"
+ },
+ "source": [
+ "## Import and load the labels"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "5RUzrh0uegqt"
+ },
+ "outputs": [],
+ "source": [
+ "LABELS_PATH = (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/data/45_labels.csv'\n",
+ ")\n",
+ "\n",
+ "labels = read_csv(LABELS_PATH)\n",
+ "\n",
+ "my_metadata = Metadata()\n",
+ "my_metadata.set(thing_classes=labels)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "VkdD8-QvGZ23"
+ },
+ "source": [
+ "## Load all images"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "dTMZWOCiECvL"
+ },
+ "outputs": [],
+ "source": [
+ "images_dir = \"/mydrive/circularnet/TestData/input-test-05022025\"\n",
+ "images = glob.glob(os.path.join(images_dir, \"*\"))\n",
+ "\n",
+ "# Make sure that the files are sorted.\n",
+ "images = natsort.natsorted(images)\n",
+ "len(images)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "-HOCxVfejxlT"
+ },
+ "outputs": [],
+ "source": [
+ "# Prediction confidence score.\n",
+ "PREDICTION_THRESHOLD = 0.70\n",
+ "\n",
+ "# The model is trained on 1024 x 1024 image dimensions.\n",
+ "HEIGHT = 1024\n",
+ "WIDTH = 1024\n",
+ "\n",
+ "# Object Tracking parameters.\n",
+ "SEARCH_RANGE_X=150\n",
+ "SEARCH_RANGE_Y=20\n",
+ "MEMORY=1\n",
+ "\n",
+ "# Create a folder for saving prediction results.\n",
+ "os.makedirs('prediction_folder', exist_ok=True)\n",
+ "prediction_folder = os.path.join(os.getcwd(), 'prediction_folder')\n",
+ "\n",
+ "# Dimensions for tracking images.\n",
+ "HEIGHT_TRACKING = 300\n",
+ "WIDTH_TRACKING = 300\n",
+ "\n",
+ "# Create a folder to troubleshoot tracking results.\n",
+ "os.makedirs('tracking', exist_ok=True)\n",
+ "\n",
+ "# Create a folder to save detected objects from Mask RCNN\n",
+ "# accpording to categories\n",
+ "output_dir = \"cropped_objects\"\n",
+ "os.makedirs(output_dir, exist_ok=True)\n",
+ "\n",
+ "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
+ "print(f\"Using device: {device}\")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "vMRVdtYEN5bg"
+ },
+ "outputs": [],
+ "source": [
+ "# Initialize the Detectron2 configuration object\n",
+ "cfg = get_cfg()\n",
+ "\n",
+ "# Load the model configuration from a YAML file.\n",
+ "cfg.merge_from_file(\"config.yaml\")\n",
+ "\n",
+ "# Set the confidence threshold.\n",
+ "cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = PREDICTION_THRESHOLD\n",
+ "\n",
+ "# Specify the path to the trained model weights.\n",
+ "cfg.MODEL.WEIGHTS = \"model_final.pth\"\n",
+ "\n",
+ "cfg.MODEL.DEVICE = \"cuda\"\n",
+ "\n",
+ "# Create a predictor object using the configured model.\n",
+ "predictor = DefaultPredictor(cfg)\n",
+ "predictor.model.to(device) # Ensure the model is on GPU"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "H3r73X-FGzz-"
+ },
+ "source": [
+ "## Perform inference"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "cm7IheKajzsW"
+ },
+ "outputs": [],
+ "source": [
+ "tracking_images = {}\n",
+ "features_set = []\n",
+ "\n",
+ "\n",
+ "for frame, image_path in tqdm.tqdm(enumerate(images, start=1)):\n",
+ " original_image = cv2.imread(image_path)\n",
+ " original_height, original_width = original_image.shape[:2]\n",
+ " resized_image = cv2.resize(\n",
+ " original_image,\n",
+ " (WIDTH, HEIGHT),\n",
+ " interpolation=cv2.INTER_AREA\n",
+ " )\n",
+ "\n",
+ " # Perform inference.\n",
+ " results = predictor(resized_image)\n",
+ "\n",
+ " # Implement class agnostic NMS.\n",
+ " results = convert_detections_to_instances(results)\n",
+ "\n",
+ "\n",
+ " # Extract the attributes from the prediction result.\n",
+ " fields = results[\"instances\"].to(\"cpu\").get_fields()\n",
+ " bboxes = fields[\"pred_boxes\"].tensor.numpy().astype(int)\n",
+ " if not len(bboxes):\n",
+ " continue\n",
+ "\n",
+ " scores = fields[\"scores\"].numpy()\n",
+ " classes = fields[\"pred_classes\"].numpy()\n",
+ " masks = fields[\"pred_masks\"].numpy()\n",
+ "\n",
+ " # Keep the predictions whose binary mask area \u003e 4000.\n",
+ " mask_areas = np.array([np.sum(i) for i in masks])\n",
+ " valid_indices = mask_areas \u003e 4000\n",
+ " bboxes = bboxes[valid_indices]\n",
+ " if not len(bboxes):\n",
+ " continue\n",
+ "\n",
+ " scores = scores[valid_indices]\n",
+ " classes = classes[valid_indices]\n",
+ " masks = masks[valid_indices]\n",
+ "\n",
+ " # Adjust the image size to ensure both dimensions are at least 1024\n",
+ " # for saving images with bbx and masks.\n",
+ " height_plot, width_plot = adjust_image_size(\n",
+ " original_image.shape[0], original_image.shape[1], 1024\n",
+ " )\n",
+ " image_plot = cv2.resize(\n",
+ " original_image,\n",
+ " (width_plot, height_plot),\n",
+ " interpolation=cv2.INTER_AREA,\n",
+ " )\n",
+ "\n",
+ " # Rescale bounding boxes\n",
+ " scale_x = width_plot / WIDTH\n",
+ " scale_y = height_plot / HEIGHT\n",
+ " bboxes = (bboxes * [scale_x, scale_y, scale_x, scale_y]).astype(int)\n",
+ "\n",
+ " # Rescale masks\n",
+ " if masks is not None:\n",
+ " resized_masks = np.array([\n",
+ " cv2.resize(mask.astype(\"uint8\"), (width_plot, height_plot), interpolation=cv2.INTER_NEAREST)\n",
+ " for mask in masks\n",
+ " ])\n",
+ " else:\n",
+ " resized_masks = None\n",
+ "\n",
+ " # Convert predictions to Detectron2 visualization format\n",
+ " pred_boxes = Boxes(torch.tensor(bboxes, dtype=torch.float32))\n",
+ " pred_classes = torch.tensor(classes, dtype=torch.int64)\n",
+ " pred_scores = torch.tensor(scores, dtype=torch.float32)\n",
+ "\n",
+ " predictions = {\n",
+ " \"pred_boxes\": pred_boxes,\n",
+ " \"scores\": pred_scores,\n",
+ " \"pred_classes\": pred_classes,\n",
+ " }\n",
+ "\n",
+ " if resized_masks is not None:\n",
+ " predictions[\"pred_masks\"] = torch.tensor(resized_masks, dtype=torch.uint8)\n",
+ "\n",
+ " # Save the prediction results as an image file with bbx and masks.\n",
+ " visualizer = Visualizer(\n",
+ " img_rgb=image_plot, metadata=my_metadata, scale=1\n",
+ " )\n",
+ " visualized_image = visualizer.draw_instance_predictions(\n",
+ " Instances((height_plot, width_plot),\n",
+ " **predictions\n",
+ " )\n",
+ " ).get_image()\n",
+ "\n",
+ " final_image = Image.fromarray(cv2.hconcat([image_plot[:,:,::-1], visualized_image[:, :, ::-1]]))\n",
+ " final_image.save(f'prediction_folder/{os.path.basename(image_path)}')\n",
+ "\n",
+ " # Create object tracking data.\n",
+ " tracking_image = cv2.resize(\n",
+ " original_image,\n",
+ " (WIDTH_TRACKING, HEIGHT_TRACKING),\n",
+ " interpolation=cv2.INTER_AREA,\n",
+ " )\n",
+ " tracking_images[os.path.basename(image_path)] = tracking_image\n",
+ "\n",
+ " tracking_masks = np.array([\n",
+ " cv2.resize(\n",
+ " mask.astype(\"uint8\"),\n",
+ " (WIDTH_TRACKING, HEIGHT_TRACKING),\n",
+ " interpolation=cv2.INTER_NEAREST\n",
+ " ) for mask in masks\n",
+ " ])\n",
+ " # In case of connected masks, keep the biggest mask and fill the holes\n",
+ " # in case of incomplete detections by Mask RCNN.\n",
+ " tracking_masks = np.array([\n",
+ " dilated_largest_component(i) for i in tracking_masks]\n",
+ " )\n",
+ "\n",
+ " # Crop objects from an image using masks for color detection.\n",
+ " cropped_objects = [\n",
+ " np.where(np.expand_dims(i, -1), image_plot[:,:,::-1], 0)\n",
+ " for i in resized_masks\n",
+ " ]\n",
+ "\n",
+ " # Perform color detection using clustering approach.\n",
+ " dominant_colors = [\n",
+ " *map(\n",
+ " color_and_property_extractor.find_dominant_color, cropped_objects\n",
+ " )\n",
+ " ]\n",
+ " generic_color_names = color_and_property_extractor.get_generic_color_name(dominant_colors)\n",
+ "\n",
+ " # Extract features.\n",
+ " features = extract_properties(\n",
+ " tracking_image, tracking_masks\n",
+ " )\n",
+ " features[\"source_name\"] = os.path.basename(os.path.dirname(image_path))\n",
+ " features[\"image_name\"] = os.path.basename(image_path)\n",
+ " features[\"creation_time\"] = get_image_creation_time(image_path)\n",
+ " features[\"frame\"] = frame\n",
+ " features[\"detection_scores\"] = list(scores)\n",
+ " features[\"detection_classes\"] = list(classes)\n",
+ " features[\"detection_classes_names\"] = [labels[i] for i in list(classes)]\n",
+ " features[\"color\"] = generic_color_names\n",
+ " features_set.append(features)\n",
+ "\n",
+ "\n",
+ "if features_set:\n",
+ " features_df = pd.concat(features_set, ignore_index=True)\n",
+ " tracking_features = apply_tracking(\n",
+ " features_df,\n",
+ " search_range_x=SEARCH_RANGE_X,\n",
+ " search_range_y=SEARCH_RANGE_Y,\n",
+ " memory=MEMORY\n",
+ " )\n",
+ " agg_features = process_tracking_result(tracking_features)\n",
+ " counts = agg_features.groupby(\"detection_classes_names\").size()\n",
+ " counts.to_frame().to_csv(os.path.join(os.getcwd(), \"count.csv\"))\n",
+ " print(counts)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "tcEjtdLL-CnZ"
+ },
+ "source": [
+ "## Visualize Object Tracking "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "8yZOHwsS5GLL"
+ },
+ "outputs": [],
+ "source": [
+ "CIRCLE_RADIUS =7\n",
+ "CIRCLE_THICKNESS = 3\n",
+ "font = cv2.FONT_HERSHEY_SIMPLEX\n",
+ "fontScale = 1\n",
+ "color = (255, 0, 0)\n",
+ "groups = tracking_features.groupby('image_name')\n",
+ "\n",
+ "for name, group in groups:\n",
+ " img = tracking_images[name].copy()\n",
+ " for k in range(len(group)):\n",
+ " cv2.circle(img,\n",
+ " (int(group.iloc[k]['x']),int(group.iloc[k]['y'])),\n",
+ " CIRCLE_RADIUS,\n",
+ " (255,133,233),\n",
+ " -1\n",
+ " )\n",
+ " cv2.putText(img,\n",
+ " str(int(group.iloc[k]['particle'])),\n",
+ " (int(group.iloc[k]['x']), int(group.iloc[k]['y'])),\n",
+ " font,\n",
+ " fontScale,\n",
+ " color,\n",
+ " 2,\n",
+ " cv2.LINE_AA\n",
+ " )\n",
+ "\n",
+ " cv2.imwrite(os.path.join('tracking',name), img)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "NW8GKKl7-Voy"
+ },
+ "source": [
+ "## Visualize Predictions by Categories"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "p5neH4_tsdZO"
+ },
+ "outputs": [],
+ "source": [
+ "if not agg_features.empty:\n",
+ " for group_name, df in tqdm.tqdm(agg_features.groupby(\"detection_classes_names\")):\n",
+ " os.makedirs(f'{output_dir}/{group_name}', exist_ok=True)\n",
+ "\n",
+ " for row in df.itertuples(index=False):\n",
+ " # Get the image\n",
+ " image = cv2.imread(os.path.join(images_dir, row.image_name))\n",
+ " image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n",
+ " new_h, new_w = image.shape[0], image.shape[1]\n",
+ "\n",
+ " # Get the bounding box and resize it\n",
+ " y1, x1, y2, x2 = row.bbox_0, row.bbox_1, row.bbox_2, row.bbox_3\n",
+ " new_bbox = resize_bbox(y1, x1, y2, x2, HEIGHT_TRACKING, WIDTH_TRACKING, new_h, new_w)\n",
+ "\n",
+ " # Include the score in the filename\n",
+ " score = row.detection_scores if hasattr(row, 'detection_scores') else 0.0\n",
+ " name = f'{os.path.splitext(row.image_name)[0]}_{row.particle}_{score:.2f}.png'\n",
+ "\n",
+ " # Save the cropped image\n",
+ " cv2.imwrite(f'{output_dir}/{row.detection_classes_names}/{name}',\n",
+ " image[new_bbox[0]:new_bbox[2], new_bbox[1]:new_bbox[3]])\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hgmZvRqCDrnh"
+ },
+ "source": [
+ "## Copying folders to my Google drive"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "jgmX0qaLDuj9"
+ },
+ "outputs": [],
+ "source": [
+ "destination_folder = '/mydrive/circularnet/TestModel'\n",
+ "os.makedirs(destination_folder, exist_ok=True)\n",
+ "\n",
+ "# Function to safely copy directory, removing destination first if it exists\n",
+ "def copytree_replace(src, dst):\n",
+ " if os.path.exists(dst):\n",
+ " shutil.rmtree(dst)\n",
+ " shutil.copytree(src, dst)\n",
+ "\n",
+ "# Function to safely copy file, overwriting if it exists\n",
+ "def copy_replace(src, dst):\n",
+ " if os.path.exists(dst):\n",
+ " os.remove(dst)\n",
+ " shutil.copy(src, dst)\n",
+ "\n",
+ "copytree_replace(os.path.join(os.getcwd(), \"prediction_folder\"), os.path.join(destination_folder, \"prediction_folder\"))\n",
+ "copytree_replace(os.path.join(os.getcwd(), \"cropped_objects\"),os.path.join(destination_folder, \"cropped_objects\"))\n",
+ "copytree_replace(os.path.join(os.getcwd(), \"tracking\"),os.path.join(destination_folder, \"tracking\"))\n",
+ "copy_replace(os.path.join(os.getcwd(), \"count.csv\"),os.path.join(destination_folder, \"count.csv\"))"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "private_outputs": true,
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/Inference_Tensorflow_Pipeline_with_Tracking.ipynb b/official/projects/waste_identification_ml/model_inference_with_tracking/Inference_Tensorflow_Pipeline_with_Tracking.ipynb
new file mode 100644
index 00000000000..0d0eac4702f
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/Inference_Tensorflow_Pipeline_with_Tracking.ipynb
@@ -0,0 +1,1234 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "TtlIRiNXWlQ0"
+ },
+ "source": [
+ "# Waste identification with instance segmentation in TensorFlow"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ohoMgYgXWsIO"
+ },
+ "source": [
+ "This Colab notebook demonstrates an end-to-end pipeline for object detection, feature extraction, object tracking, and data aggregation using the Mask R-CNN model from TensorFlow Model Garden.\n",
+ "Key Steps in the Notebook:\n",
+ "\n",
+ "\n",
+ "\n",
+ "* Object Detection and segmentation – Detect objects in a set of images using Mask R-CNN.\n",
+ "* Feature Extraction \u0026 Tracking – Extract object features and track them across multiple frames to eliminate duplicate counts.\n",
+ "* Color Detection – Identify the color of each detected object.\n",
+ "* Postprocessing – Aggregate tracking results and apply filtering to reduce false positives and false negatives.\n",
+ "* Save detection and tracking results.\n",
+ "* Push the final object count to a BigQuery table, which can be connected to a Looker dashboard for visualization in Google Cloud Platform (GCP).\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "8PKG9z4VYPEs"
+ },
+ "source": [
+ "To finish this task, a proper path for the saved models and images need to be provided. The path to the labels on which the models are trained is in the waste_identification_ml directory inside the Tensorflow Model Garden repository."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "o2J5IWCgAOxC"
+ },
+ "source": [
+ "This notebook will output 3 folders and 1 csv file :\n",
+ "\n",
+ "\n",
+ "* **prediction_folder** : Will contain prediction results with bbox and masks.\n",
+ "* **tracking** : Will contain tracking visualization.\n",
+ "* **cropped_objects** : Will contain category level detected objects.\n",
+ "* **count.csv** : Will contain the individual counts of each category.\n",
+ "\n",
+ "\n",
+ "\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "ELUFMVDDAopS"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Imports and Setup\n",
+ "\n",
+ "!pip install -q trackpy\n",
+ "\n",
+ "import sys\n",
+ "import tensorflow as tf\n",
+ "import csv\n",
+ "from typing import Any, TypedDict, Callable\n",
+ "import cv2\n",
+ "import logging\n",
+ "import numpy as np\n",
+ "import matplotlib.pyplot as plt\n",
+ "import glob\n",
+ "import natsort\n",
+ "import tqdm\n",
+ "import os\n",
+ "from PIL import Image\n",
+ "from scipy import ndimage\n",
+ "import pandas as pd\n",
+ "import skimage\n",
+ "import datetime\n",
+ "import trackpy as tp\n",
+ "import shutil\n",
+ "\n",
+ "logging.disable(logging.WARNING)\n",
+ "\n",
+ "%matplotlib inline"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "g5_XxVSFATYZ"
+ },
+ "outputs": [],
+ "source": [
+ "# Connect to Google drive if your data is stored there.\n",
+ "from google.colab import drive\n",
+ "drive.mount('/content/gdrive')\n",
+ "\n",
+ "try:\n",
+ " !ln -s /content/gdrive/My\\ Drive/ /mydrive\n",
+ " print('Successful')\n",
+ "except Exception as e:\n",
+ " print(e)\n",
+ " print('Not successful')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WAQ8ymV-Cgqp"
+ },
+ "outputs": [],
+ "source": [
+ "# # Connect to GCP bucket if your data is store there and copy them locally.\n",
+ "# !gcloud init"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "_77YK3a_BCg_"
+ },
+ "source": [
+ "To visualize the images with the proper detected boxes and segmentation masks, we will use the TensorFlow Object Detection API. To install it we will clone the repo.\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qhk_NujKO0mb"
+ },
+ "outputs": [],
+ "source": [
+ "# Clone the tensorflow models repository.\n",
+ "!git clone --depth 1 https://github.com/tensorflow/models 2\u003e/dev/null"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "o1dYyG55BtWb"
+ },
+ "outputs": [],
+ "source": [
+ "sys.path.append('models/research/')\n",
+ "from object_detection.utils import ops as utils_ops\n",
+ "from object_detection.utils import visualization_utils as viz_utils\n",
+ "\n",
+ "sys.path.append('models/official/projects/waste_identification_ml/model_inference/')\n",
+ "import color_and_property_extractor"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "cellView": "form",
+ "id": "GO488S78_2GJ"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Utilities\n",
+ "\n",
+ "_PROPERTIES = (\n",
+ " 'area',\n",
+ " 'bbox',\n",
+ " 'convex_area',\n",
+ " 'bbox_area',\n",
+ " 'major_axis_length',\n",
+ " 'minor_axis_length',\n",
+ " 'eccentricity',\n",
+ " 'centroid',\n",
+ " 'label',\n",
+ " 'mean_intensity',\n",
+ " 'max_intensity',\n",
+ " 'min_intensity',\n",
+ " 'perimeter'\n",
+ ")\n",
+ "\n",
+ "\n",
+ "class ItemDict(TypedDict):\n",
+ " id: int\n",
+ " name: str\n",
+ " supercategory: str\n",
+ "\n",
+ "\n",
+ "def load_model(model_path: str) -\u003e Callable:\n",
+ " \"\"\"Loads a TensorFlow SavedModel and returns a function for making predictions.\n",
+ "\n",
+ " Args:\n",
+ " model_path: Path to the TensorFlow SavedModel.\n",
+ "\n",
+ " Returns:\n",
+ " A function that can be used to make predictions.\n",
+ " \"\"\"\n",
+ " try:\n",
+ " print('loading model...')\n",
+ " model = tf.saved_model.load(model_path)\n",
+ " print('model loaded!')\n",
+ " detection_fn = model.signatures['serving_default']\n",
+ " return detection_fn\n",
+ " except (OSError, ValueError, KeyError) as e:\n",
+ " print(f\"Error loading model: {e}\")\n",
+ " raise\n",
+ "\n",
+ "\n",
+ "def perform_detection(model: Callable, image: np.ndarray) -\u003e dict[str, np.ndarray]:\n",
+ " \"\"\"Perform Mask R-CNN object detection on an image using the specified model.\n",
+ "\n",
+ " Args:\n",
+ " model: A function that can be used to make predictions.\n",
+ " image: A NumPy array representing the image to be processed.\n",
+ "\n",
+ " Returns:\n",
+ " Detection results, where keys are output names and values are NumPy arrays.\n",
+ " \"\"\"\n",
+ " detection_results = model(image)\n",
+ " detection_results = {key: value.numpy() for key, value in detection_results.items()}\n",
+ " return detection_results\n",
+ "\n",
+ "\n",
+ "def _read_csv_to_list(file_path: str) -\u003e list[str]:\n",
+ " \"\"\"Reads a CSV file and returns its contents as a list.\n",
+ "\n",
+ " This function reads the given CSV file, skips the header, and assumes\n",
+ " there is only one column in the CSV. It returns the contents as a list of\n",
+ " strings.\n",
+ "\n",
+ " Args:\n",
+ " file_path: The path to the CSV file.\n",
+ "\n",
+ " Returns:\n",
+ " The contents of the CSV file as a list of strings.\n",
+ " \"\"\"\n",
+ " data_list = []\n",
+ " with open(file_path, 'r') as csvfile:\n",
+ " reader = csv.reader(csvfile)\n",
+ " for row in reader:\n",
+ " data_list.append(row[0]) # Assuming there is only one column in the CSV\n",
+ " return data_list\n",
+ "\n",
+ "\n",
+ "def _categories_dictionary(objects: list[str]) -\u003e dict[int, ItemDict]:\n",
+ " \"\"\"This function takes a list of objects and returns a dictionaries.\n",
+ "\n",
+ " A dictionary of objects, where each object is represented by a dictionary\n",
+ " with the following keys:\n",
+ " - id: The ID of the object.\n",
+ " - name: The name of the object.\n",
+ " - supercategory: The supercategory of the object.\n",
+ "\n",
+ " Args:\n",
+ " objects: A list of strings, where each string is the name of an\n",
+ " object.\n",
+ "\n",
+ " Returns:\n",
+ " A tuple of two dictionaries, as described above.\n",
+ " \"\"\"\n",
+ " category_index = {}\n",
+ " for num, obj_name in enumerate(objects, start=1):\n",
+ " obj_dict = {'id': num, 'name': obj_name, 'supercategory': 'objects'}\n",
+ " category_index[num] = obj_dict\n",
+ " return category_index\n",
+ "\n",
+ "\n",
+ "def load_labels(labels_path: str) -\u003e tuple[list[str], dict[int, ItemDict]]:\n",
+ " \"\"\"\n",
+ " Load label mappings from a CSV file and generate category indices.\n",
+ "\n",
+ " Args:\n",
+ " labels_path (str): Path to the CSV file containing label mappings.\n",
+ "\n",
+ " Returns:\n",
+ " Tuple[Dict[int, dict], Dict[int, dict]]:\n",
+ " - A dictionary mapping category IDs to label details.\n",
+ " - A processed category index dictionary.\n",
+ " \"\"\"\n",
+ " labels = _read_csv_to_list(labels_path)\n",
+ " category_index = _categories_dictionary(labels)\n",
+ " return labels, category_index\n",
+ "\n",
+ "\n",
+ "def preprocess_image(path: str, height: int, width: int) -\u003e tuple[np.ndarray, np.ndarray]:\n",
+ " \"\"\"\n",
+ " Load an image from a file into a NumPy array, resize it, and expand dimensions for batch processing.\n",
+ "\n",
+ " Args:\n",
+ " path: The file path to the image.\n",
+ " height: Desired height of the resized image.\n",
+ " width: Desired width of the resized image.\n",
+ "\n",
+ " Returns:\n",
+ " original_image: The original image with shape (original_height, original_width, 3).\n",
+ " resized_image: The resized image with shape (1, height, width, 3), suitable for model input.\n",
+ " \"\"\"\n",
+ " original_image = cv2.imread(path)\n",
+ " if original_image is None:\n",
+ " raise FileNotFoundError(f\"Image not found at path: {path}\")\n",
+ "\n",
+ " original_image = cv2.cvtColor(original_image, cv2.COLOR_BGR2RGB)\n",
+ " resized_image = cv2.resize(original_image, (width, height), interpolation=cv2.INTER_AREA)\n",
+ " resized_image = np.expand_dims(resized_image, axis=0)\n",
+ "\n",
+ " return original_image, resized_image\n",
+ "\n",
+ "\n",
+ "def filter_detection(results: dict[str, np.ndarray], valid_indices: np.ndarray) -\u003e dict[str, np.ndarray]:\n",
+ " \"\"\"Filter the detection results based on the valid indices.\n",
+ "\n",
+ " Args:\n",
+ " results: The detection results from the model.\n",
+ " valid_indices: The indices of the valid detections.\n",
+ "\n",
+ " Returns:\n",
+ " The filtered detection results.\n",
+ " \"\"\"\n",
+ " if np.array(valid_indices).dtype == bool:\n",
+ " new_num_detections = int(np.sum(valid_indices))\n",
+ " else:\n",
+ " new_num_detections = len(valid_indices)\n",
+ "\n",
+ " # Define the keys to filter\n",
+ " keys_to_filter = [\n",
+ " 'detection_masks',\n",
+ " 'detection_masks_resized',\n",
+ " 'detection_masks_reframed',\n",
+ " 'detection_classes',\n",
+ " 'detection_boxes',\n",
+ " 'normalized_boxes',\n",
+ " 'detection_scores',\n",
+ " 'detection_classes_names',\n",
+ " ]\n",
+ "\n",
+ " # Apply filtering to the specified keys\n",
+ " filtered_output = {}\n",
+ "\n",
+ " for key in keys_to_filter:\n",
+ " if key in results:\n",
+ " if key == 'detection_masks':\n",
+ " filtered_output[key] = results[key][:, valid_indices, :, :]\n",
+ " elif key in ['detection_masks_resized', 'detection_masks_reframed']:\n",
+ " filtered_output[key] = results[key][valid_indices, :, :]\n",
+ " elif key in ['detection_boxes', 'normalized_boxes']:\n",
+ " filtered_output[key] = results[key][:, valid_indices, :]\n",
+ " elif key in ['detection_classes', 'detection_scores', 'detection_classes_names']:\n",
+ " filtered_output[key] = results[key][:, valid_indices]\n",
+ " filtered_output['num_detections'] = np.array([new_num_detections])\n",
+ "\n",
+ " return filtered_output\n",
+ "\n",
+ "\n",
+ "\n",
+ "def reframe_masks(results: dict[str, np.ndarray], boxes: str, height: int, width: int) -\u003e np.ndarray:\n",
+ " \"\"\"Reframe the masks to an image size.\n",
+ "\n",
+ " Args:\n",
+ " results: The detection results from the model.\n",
+ " boxes: The detection boxes.\n",
+ " height: The height of the original image.\n",
+ " width: The width of the original image.\n",
+ "\n",
+ " Returns:\n",
+ " The reframed masks.\n",
+ " \"\"\"\n",
+ " detection_masks = results['detection_masks'][0]\n",
+ " detection_boxes = results[boxes][0]\n",
+ " detection_masks_reframed = utils_ops.reframe_box_masks_to_image_masks(\n",
+ " detection_masks, detection_boxes, height, width\n",
+ " )\n",
+ " detection_masks_reframed = tf.cast(detection_masks_reframed \u003e 0.5, np.uint8)\n",
+ " detection_masks_reframed = detection_masks_reframed.numpy()\n",
+ " return detection_masks_reframed\n",
+ "\n",
+ "\n",
+ "def _calculate_area(mask: np.ndarray) -\u003e int:\n",
+ " \"\"\"Calculate the area of the mask.\n",
+ "\n",
+ " Args:\n",
+ " mask: The mask to calculate the area of.\n",
+ "\n",
+ " Returns:\n",
+ " The area of the mask.\n",
+ " \"\"\"\n",
+ " return np.sum(mask)\n",
+ "\n",
+ "\n",
+ "def _calculate_iou(mask1: np.ndarray, mask2: np.ndarray) -\u003e float:\n",
+ " \"\"\"Calculate the intersection over union (IoU) between two masks.\n",
+ "\n",
+ " Args:\n",
+ " mask1: The first mask.\n",
+ " mask2: The second mask.\n",
+ "\n",
+ " Returns:\n",
+ " The intersection over union (IoU) between the two masks.\n",
+ " \"\"\"\n",
+ " intersection = np.logical_and(mask1, mask2).sum()\n",
+ " union = np.logical_or(mask1, mask2).sum()\n",
+ " return intersection / union if union != 0 else 0\n",
+ "\n",
+ "\n",
+ "def _is_contained(mask1: np.ndarray, mask2: np.ndarray) -\u003e bool:\n",
+ " \"\"\"Check if mask1 is entirely contained within mask2.\n",
+ "\n",
+ " Args:\n",
+ " mask1: The first mask.\n",
+ " mask2: The second mask.\n",
+ "\n",
+ " Returns:\n",
+ " True if mask1 is entirely contained within mask2, False otherwise.\n",
+ " \"\"\"\n",
+ " return np.array_equal(np.logical_and(mask1, mask2), mask1)\n",
+ "\n",
+ "\n",
+ "def filter_masks(masks: np.ndarray, iou_threshold=0.8, area_threshold=None) -\u003e np.ndarray:\n",
+ " \"\"\"Filter the overlapping masks.\n",
+ "\n",
+ " Filter the masks based on the area and intersection over union (IoU).\n",
+ "\n",
+ " Args:\n",
+ " masks: The masks to filter.\n",
+ " iou_threshold: The threshold for the intersection over union (IoU) between\n",
+ " two masks.\n",
+ " area_threshold: The threshold for the area of the mask.\n",
+ "\n",
+ " Returns:\n",
+ " The indices of the unique masks.\n",
+ " \"\"\"\n",
+ " # Calculate the area for each mask\n",
+ " areas = np.array([_calculate_area(mask) for mask in masks])\n",
+ "\n",
+ " # Sort the masks based on area in descending order\n",
+ " sorted_indices = np.argsort(areas)[::-1]\n",
+ " sorted_masks = masks[sorted_indices]\n",
+ " sorted_areas = areas[sorted_indices]\n",
+ "\n",
+ " unique_indices = []\n",
+ "\n",
+ " for i, mask in enumerate(sorted_masks):\n",
+ " if (area_threshold is not None and sorted_areas[i] \u003e area_threshold) or sorted_areas[i] \u003c 4000:\n",
+ " continue\n",
+ "\n",
+ " keep = True\n",
+ " for j in range(i):\n",
+ " if _calculate_iou(mask, sorted_masks[j]) \u003e iou_threshold or _is_contained(\n",
+ " mask, sorted_masks[j]\n",
+ " ):\n",
+ " keep = False\n",
+ " break\n",
+ " if keep:\n",
+ " unique_indices.append(sorted_indices[i])\n",
+ "\n",
+ " return unique_indices\n",
+ "\n",
+ "\n",
+ "def adjust_image_size(height: int, width: int, min_size: int) -\u003e tuple[int, int]:\n",
+ " \"\"\"Adjust the image size to ensure both dimensions are at least 1024.\n",
+ "\n",
+ " Args:\n",
+ " height: The height of the image.\n",
+ " width: The width of the image.\n",
+ " min_size: Minimum size of the image dimension needed.\n",
+ "\n",
+ " Returns:\n",
+ " The adjusted height and width of the image.\n",
+ " \"\"\"\n",
+ " if height \u003c min_size or width \u003c min_size:\n",
+ " return height, width\n",
+ "\n",
+ " # Calculate the scale factor to ensure both dimensions remain at least 1024\n",
+ " scale_factor = min(height / min_size, width / min_size)\n",
+ "\n",
+ " new_height = int(height / scale_factor)\n",
+ " new_width = int(width / scale_factor)\n",
+ "\n",
+ " return new_height, new_width\n",
+ "\n",
+ "\n",
+ "def display_bbox_masks_labels(\n",
+ " result: dict[Any, np.ndarray],\n",
+ " image: np.ndarray,\n",
+ " category_index: dict[int, dict[str, str]],\n",
+ " threshold: float,\n",
+ ") -\u003e None:\n",
+ " \"\"\"Saves an image with visualized bounding boxes, labels, and masks.\n",
+ "\n",
+ " This function takes the output from Mask R-CNN, copies the original image,\n",
+ " and applies visualizations of detection boxes, classes, and scores.\n",
+ " If available, it also applies segmentation masks. The result is an image that\n",
+ " juxtaposes the original with the annotated version, saved to the specified\n",
+ " folder.\n",
+ "\n",
+ " Args:\n",
+ " result: The output from theMask RCNN model, expected to contain detection\n",
+ " boxes, classes, scores, reframed detection masks, etc.\n",
+ " image: The original image as a numpy array.\n",
+ " file_name: The filename for saving the output image.\n",
+ " folder: The folder path where the output image will be saved.\n",
+ " category_index: A dictionary mapping class IDs to class labels.\n",
+ " threshold: Value between 0 and 1 to filter out the prediction results.\n",
+ " \"\"\"\n",
+ " image_new = image.copy()\n",
+ " image_new = cv2.cvtColor(image_new, cv2.COLOR_BGR2RGB)\n",
+ " viz_utils.visualize_boxes_and_labels_on_image_array(\n",
+ " image_new,\n",
+ " result['normalized_boxes'][0],\n",
+ " (result['detection_classes'][0] + 0).astype(int),\n",
+ " result['detection_scores'][0],\n",
+ " category_index=category_index,\n",
+ " use_normalized_coordinates=True,\n",
+ " max_boxes_to_draw=100,\n",
+ " min_score_thresh=threshold,\n",
+ " agnostic_mode=False,\n",
+ " instance_masks=result.get('detection_masks_reframed', None),\n",
+ " line_thickness=4,\n",
+ " )\n",
+ " return image_new\n",
+ "\n",
+ "\n",
+ "def save_bbox_masks_labels(\n",
+ " result: dict[Any, np.ndarray],\n",
+ " image: np.ndarray,\n",
+ " file_name: str,\n",
+ " folder: str,\n",
+ " category_index: dict[int, dict[str, str]],\n",
+ " threshold: float,\n",
+ ") -\u003e None:\n",
+ " \"\"\"Saves an image with visualized bounding boxes, labels, and masks.\n",
+ "\n",
+ " This function takes the output from Mask R-CNN, copies the original image,\n",
+ " and applies visualizations of detection boxes, classes, and scores.\n",
+ " If available, it also applies segmentation masks. The result is an image that\n",
+ " juxtaposes the original with the annotated version, saved to the specified\n",
+ " folder.\n",
+ "\n",
+ " Args:\n",
+ " result: The output from theMask RCNN model, expected to contain detection\n",
+ " boxes, classes, scores, reframed detection masks, etc.\n",
+ " image: The original image as a numpy array.\n",
+ " file_name: The filename for saving the output image.\n",
+ " folder: The folder path where the output image will be saved.\n",
+ " category_index: A dictionary mapping class IDs to class labels.\n",
+ " threshold: Value between 0 and 1 to filter out the prediction results.\n",
+ " \"\"\"\n",
+ " image_new = image.copy()\n",
+ " viz_utils.visualize_boxes_and_labels_on_image_array(\n",
+ " image_new,\n",
+ " result['normalized_boxes'][0],\n",
+ " (result['detection_classes'][0] + 0).astype(int),\n",
+ " result['detection_scores'][0],\n",
+ " category_index=category_index,\n",
+ " use_normalized_coordinates=True,\n",
+ " max_boxes_to_draw=100,\n",
+ " min_score_thresh=threshold,\n",
+ " agnostic_mode=False,\n",
+ " instance_masks=result.get('detection_masks_reframed', None),\n",
+ " line_thickness=4,\n",
+ " )\n",
+ "\n",
+ " concatenated_image = np.concatenate((image, image_new), axis=1)\n",
+ " concatenated_image = Image.fromarray(concatenated_image)\n",
+ " concatenated_image.save(os.path.join(folder, file_name))\n",
+ "\n",
+ "\n",
+ "def dilated_largest_component(mask: np.ndarray) -\u003e np.ndarray:\n",
+ " \"\"\"Extracts the largest connected component and fills holes.\n",
+ "\n",
+ " Args:\n",
+ " mask: Input binary mask (2D numpy array).\n",
+ "\n",
+ " Returns:\n",
+ " Binary mask of the largest connected component.\n",
+ " \"\"\"\n",
+ " mask = mask.astype(np.uint8)*255\n",
+ " num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask, connectivity=8)\n",
+ " largest_label = 1 + np.argmax(stats[1:, cv2.CC_STAT_AREA])\n",
+ " largest_component_mask = np.zeros(mask.shape, dtype=\"uint8\")\n",
+ " largest_component_mask[labels == largest_label] = 1\n",
+ " largest_component_mask = ndimage.binary_fill_holes(largest_component_mask).astype(int)\n",
+ " return largest_component_mask\n",
+ "\n",
+ "\n",
+ "def extract_properties(image, results, masks):\n",
+ " \"\"\"Extract properties of the mask.\n",
+ "\n",
+ " Args:\n",
+ " image: Corresponding image of the mask.\n",
+ " results: The detection results from the model.\n",
+ " masks: The masks to extract properties from.\n",
+ "\n",
+ " Returns:\n",
+ " The extracted properties.\n",
+ " \"\"\"\n",
+ " list_of_df = []\n",
+ " for mask in results[masks]:\n",
+ " mask = np.where(mask, 1, 0)\n",
+ " df = pd.DataFrame(\n",
+ " skimage.measure.regionprops_table(mask, intensity_image=image, properties=_PROPERTIES)\n",
+ " )\n",
+ " list_of_df.append(df)\n",
+ " features = pd.concat(list_of_df, ignore_index=True)\n",
+ " features.rename(\n",
+ " columns={\n",
+ " 'centroid-0': 'y',\n",
+ " 'centroid-1': 'x',\n",
+ " 'bbox-0': 'bbox_0',\n",
+ " 'bbox-1': 'bbox_1',\n",
+ " 'bbox-2': 'bbox_2',\n",
+ " 'bbox-3': 'bbox_3',\n",
+ " },\n",
+ " inplace=True,\n",
+ " )\n",
+ " return features\n",
+ "\n",
+ "\n",
+ "def get_image_creation_time(image_path):\n",
+ " \"\"\"\n",
+ " Retrieves the creation time of an image, trying multiple methods.\n",
+ "\n",
+ " Args:\n",
+ " image_path: The path to the image file.\n",
+ "\n",
+ " Returns:\n",
+ " A string representing the creation time in the format \"%Y-%m-%d %H:%M:%S\" if\n",
+ " found, otherwise returns \"Creation time not found\".\n",
+ " \"\"\"\n",
+ "\n",
+ " try:\n",
+ " # 1. Try EXIF data (if available)\n",
+ " image = Image.open(image_path)\n",
+ " exif_data = image._getexif()\n",
+ " if exif_data:\n",
+ " datetime_tag_id = 36867 # Tag ID for \"DateTimeOriginal\"\n",
+ " datetime_str = exif_data.get(datetime_tag_id)\n",
+ " if datetime_str:\n",
+ " datetime_obj = datetime.datetime.strptime(datetime_str, \"%Y:%m:%d %H:%M:%S\")\n",
+ " return datetime_obj.strftime(\"%Y-%m-%d %H:%M:%S\")\n",
+ "\n",
+ " # 2. Try file modification time (less accurate, but better than nothing)\n",
+ " file_modified_time = os.path.getmtime(image_path)\n",
+ " datetime_obj = datetime.datetime.fromtimestamp(file_modified_time)\n",
+ " return datetime_obj.strftime(\"%Y-%m-%d %H:%M:%S\")\n",
+ "\n",
+ " except FileNotFoundError:\n",
+ " return \"Image not found\"\n",
+ " except Exception as e:\n",
+ " return f\"Error: {e}\"\n",
+ "\n",
+ "\n",
+ "def apply_tracking(df,\n",
+ " search_range_x,\n",
+ " search_range_y,\n",
+ " memory):\n",
+ " \"\"\"Apply tracking to the dataframe.\n",
+ "\n",
+ " Args:\n",
+ " df: The dataframe to apply tracking to.\n",
+ " search_range_x: The search range of pixels for tracking along x axis.\n",
+ " search_range_y: The search range of pixels for tracking along y axis.\n",
+ " memory: The frames memory for tracking.\n",
+ "\n",
+ " Returns:\n",
+ " The tracking result dataframe.\n",
+ " \"\"\"\n",
+ " # Define the columns to link for tracking.\n",
+ " # Additional features that can be used are 'area', 'label', 'color',\n",
+ " # 'eccentricity', 'convex_area', 'mean_intensity-0', 'mean_intensity-1',\n",
+ " # 'mean_intensity-2', 'max_intensity-0', 'max_intensity-1', 'max_intensity-2',\n",
+ " # 'min_intensity-0', 'min_intensity-1', 'min_intensity-2'.\n",
+ " tracking_columns = [\n",
+ " 'x',\n",
+ " 'y',\n",
+ " 'frame',\n",
+ " 'bbox_0',\n",
+ " 'bbox_1',\n",
+ " 'bbox_2',\n",
+ " 'bbox_3',\n",
+ " 'major_axis_length',\n",
+ " 'minor_axis_length',\n",
+ " 'perimeter',\n",
+ " ]\n",
+ "\n",
+ " # Perform the tracking operation on the specified columns\n",
+ " track_df = tp.link_df(df[tracking_columns], search_range=(search_range_y, search_range_x), memory=memory)\n",
+ "\n",
+ " # Copy the additional columns from the original dataframe\n",
+ " additional_columns = [\n",
+ " 'source_name',\n",
+ " 'image_name',\n",
+ " 'detection_scores',\n",
+ " 'detection_classes_names',\n",
+ " 'detection_classes',\n",
+ " 'color',\n",
+ " 'creation_time'\n",
+ " ]\n",
+ " track_df[additional_columns] = df[additional_columns]\n",
+ "\n",
+ " track_df.drop(columns=['frame'], inplace=True)\n",
+ " track_df.reset_index(drop=True, inplace=True)\n",
+ "\n",
+ " return track_df\n",
+ "\n",
+ "\n",
+ "def process_tracking_result(df):\n",
+ " \"\"\"Process the tracking result dataframe.\n",
+ "\n",
+ " Args:\n",
+ " df: Dataframe to be aggregated.\n",
+ "\n",
+ " Returns:\n",
+ " Processed dataframe.\n",
+ " \"\"\"\n",
+ " # Get class information with the new include_groups parameter\n",
+ " class_info = df.groupby('particle', as_index=False).apply(\n",
+ " select_class_with_scores,\n",
+ " include_groups=False\n",
+ " )\n",
+ "\n",
+ " grouped = df.groupby('particle').agg({\n",
+ " 'source_name': 'first',\n",
+ " 'image_name': 'first',\n",
+ " 'detection_scores': 'max',\n",
+ " 'creation_time': 'first',\n",
+ " 'bbox_0': 'first',\n",
+ " 'bbox_1': 'first',\n",
+ " 'bbox_2': 'first',\n",
+ " 'bbox_3': 'first',\n",
+ " }).reset_index()\n",
+ "\n",
+ " # Add class information\n",
+ " grouped['detection_classes'] = class_info['class_id']\n",
+ " grouped['detection_classes_names'] = class_info['class_name']\n",
+ "\n",
+ " return grouped\n",
+ "\n",
+ "def select_class_with_scores(group):\n",
+ " \"\"\"\n",
+ " Select class based on modal class, falling back to highest score for ties.\n",
+ " Returns both class ID and class name.\n",
+ " \"\"\"\n",
+ " # Get the value counts of classes\n",
+ " class_counts = group['detection_classes'].value_counts()\n",
+ "\n",
+ " # If there's a clear winner (one mode), use it\n",
+ " if len(class_counts) == 1 or class_counts.iloc[0] \u003e class_counts.iloc[1]:\n",
+ " class_id = group['detection_classes'].mode().iloc[0]\n",
+ " else:\n",
+ " # If there's a tie, look at highest score for each tied class\n",
+ " tied_classes = class_counts[class_counts == class_counts.iloc[0]].index\n",
+ " class_max_scores = {\n",
+ " cls: group[group['detection_classes'] == cls]['detection_scores'].max()\n",
+ " for cls in tied_classes\n",
+ " }\n",
+ " class_id = max(class_max_scores.items(), key=lambda x: x[1])[0]\n",
+ "\n",
+ " # Get corresponding class name\n",
+ " class_name = group[group['detection_classes'] == class_id]['detection_classes_names'].iloc[0]\n",
+ " return pd.Series({'class_id': class_id, 'class_name': class_name})\n",
+ "\n",
+ "\n",
+ "def resize_bbox(y1, x1, y2, x2, old_height, old_width, new_height, new_width):\n",
+ " \"\"\"Resize bounding box coordinates based on new image size.\n",
+ "\n",
+ " Args:\n",
+ " y1, x1, y2, x2 (int/float): Original bounding box coordinates.\n",
+ " old_height, old_width (int): Original image dimensions.\n",
+ " new_height, new_width (int): New image dimensions.\n",
+ "\n",
+ " Returns:\n",
+ " (new_y1, new_x1, new_y2, new_x2): Rescaled bounding box coordinates.\n",
+ " \"\"\"\n",
+ " # Compute scale factors\n",
+ " scale_x = new_width / old_width\n",
+ " scale_y = new_height / old_height\n",
+ "\n",
+ " # Scale bounding box coordinates\n",
+ " new_y1 = int(y1 * scale_y)\n",
+ " new_x1 = int(x1 * scale_x)\n",
+ " new_y2 = int(y2 * scale_y)\n",
+ " new_x2 = int(x2 * scale_x)\n",
+ "\n",
+ " return new_y1, new_x1, new_y2, new_x2"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "4XjfDEq--UlE"
+ },
+ "source": [
+ "## Import and load pre-trained models."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ZQ435YHN3Lr-"
+ },
+ "outputs": [],
+ "source": [
+ "%%bash\n",
+ "wget https://storage.googleapis.com/tf_model_garden/vision/\\\n",
+ "waste_identification_ml/Jan2025_ver2_merged_1024_1024.zip -q\n",
+ "\n",
+ "unzip Jan2025_ver2_merged_1024_1024.zip \u003e /dev/null 2\u003e\u00261"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "vMRVdtYEN5bg"
+ },
+ "outputs": [],
+ "source": [
+ "detection_fn = load_model('Jan2025_ver2_merged_1024_1024/')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "W6mmyLsOJicF"
+ },
+ "source": [
+ "## Load label map data"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "PM2A29OrJqaU"
+ },
+ "source": [
+ "Label maps correspond index numbers to category names, so that when our convolution network predicts 5, we know that this corresponds to airplane. Here we use internal utility functions, but anything that returns a dictionary mapping integers to appropriate string labels would be fine.\n",
+ "\n",
+ "We will load our labels from the same repository that we loaded the TF Object Detection API from."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "5RUzrh0uegqt"
+ },
+ "outputs": [],
+ "source": [
+ "LABELS_PATH = (\n",
+ " 'models/official/projects/waste_identification_ml/pre_processing/'\n",
+ " 'config/data/45_labels.csv'\n",
+ ")\n",
+ "\n",
+ "labels, category_index = load_labels(LABELS_PATH)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "VkdD8-QvGZ23"
+ },
+ "source": [
+ "## Load all images"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "dTMZWOCiECvL"
+ },
+ "outputs": [],
+ "source": [
+ "images_dir = \"/mydrive/circularnet/TestData/input-test-05022025\"\n",
+ "images = glob.glob(os.path.join(images_dir, \"*\"))\n",
+ "\n",
+ "# Make sure that the files are sorted.\n",
+ "images = natsort.natsorted(images)\n",
+ "len(images)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "-7ZS7gHgGk9f"
+ },
+ "outputs": [],
+ "source": [
+ "# The model is trained on 1024 x 1024 image dimensions.\n",
+ "HEIGHT = 1024\n",
+ "WIDTH = 1024\n",
+ "\n",
+ "# Object Tracking parameters.\n",
+ "# The distance an object can move along the x axis from one frame to another.\n",
+ "# The paramter was decided according to an image size of 300 x 300.\n",
+ "SEARCH_RANGE_X=150\n",
+ "# The distance an object can move along the y axis from one frame to another.\n",
+ "# The paramter was decided according to an image size of 300 x 300.\n",
+ "SEARCH_RANGE_Y=20\n",
+ "# No of frames you want to track.\n",
+ "MEMORY=1\n",
+ "\n",
+ "# Dimensions for tracking images.\n",
+ "HEIGHT_TRACKING = 300\n",
+ "WIDTH_TRACKING = 300\n",
+ "\n",
+ "# Prediction confidence score.\n",
+ "PREDICTION_THRESHOLD = 0.50\n",
+ "area_threshold = None\n",
+ "\n",
+ "# Create a folder for saving prediction results.\n",
+ "os.makedirs('prediction_folder', exist_ok=True)\n",
+ "prediction_folder = os.path.join(os.getcwd(), 'prediction_folder')\n",
+ "\n",
+ "# Create a folder to troubleshoot tracking results.\n",
+ "os.makedirs('tracking', exist_ok=True)\n",
+ "\n",
+ "# Create a folder to save detected objects from Mask RCNN\n",
+ "# accpording to categories\n",
+ "output_dir = \"cropped_objects\"\n",
+ "os.makedirs(output_dir, exist_ok=True)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "H3r73X-FGzz-"
+ },
+ "source": [
+ "## Perform inference"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "pB5NOjWKLGQo"
+ },
+ "outputs": [],
+ "source": [
+ "tracking_images = {}\n",
+ "features_set = []\n",
+ "\n",
+ "\n",
+ "for frame, image_path in tqdm.tqdm(enumerate(images, start=1)):\n",
+ " # Preprocess an image.\n",
+ " original_image, resized_image = preprocess_image(image_path, HEIGHT, WIDTH)\n",
+ " input_tensor = tf.convert_to_tensor(resized_image, dtype=tf.uint8)\n",
+ "\n",
+ " # Running inference with the model.\n",
+ " result = perform_detection(detection_fn, input_tensor)\n",
+ "\n",
+ " if result[\"num_detections\"][0]:\n",
+ " scores = result[\"detection_scores\"][0]\n",
+ " filtered_indices = scores \u003e PREDICTION_THRESHOLD\n",
+ " result = filter_detection(result, filtered_indices)\n",
+ "\n",
+ " if result[\"num_detections\"][0]:\n",
+ " # Normalize the bounding boxes according to the resized image size.\n",
+ " result[\"normalized_boxes\"] = result[\"detection_boxes\"].copy()\n",
+ " result[\"normalized_boxes\"][:, :, [0, 2]] /= HEIGHT\n",
+ " result[\"normalized_boxes\"][:, :, [1, 3]] /= WIDTH\n",
+ "\n",
+ " result['detection_boxes'] = result[\"detection_boxes\"].round().astype(int)\n",
+ "\n",
+ " # Adjust the image size to ensure both dimensions are at least 1024\n",
+ " # for saving images with bbx and masks.\n",
+ " height_plot, width_plot = adjust_image_size(\n",
+ " original_image.shape[0], original_image.shape[1], 1024\n",
+ " )\n",
+ " image_plot = cv2.resize(\n",
+ " original_image,\n",
+ " (width_plot, height_plot),\n",
+ " interpolation=cv2.INTER_AREA,\n",
+ " )\n",
+ " # Reframe the masks to the new size.\n",
+ " result[\"detection_masks_reframed\"] = reframe_masks(\n",
+ " result, \"normalized_boxes\", height_plot, width_plot\n",
+ " )\n",
+ "\n",
+ " # Filter the prediction results based on the area threshold and\n",
+ " # remove the overlapping masks.\n",
+ " unique_indices = filter_masks(\n",
+ " result[\"detection_masks_reframed\"],\n",
+ " iou_threshold=0.08,\n",
+ " area_threshold=area_threshold,\n",
+ " )\n",
+ " result = filter_detection(result, unique_indices)\n",
+ "\n",
+ " if result[\"num_detections\"][0]:\n",
+ " result['detection_classes_names'] = np.array(\n",
+ " [[str(labels[i-1]) for i in result['detection_classes'][0]]]\n",
+ " )\n",
+ "\n",
+ " # Save the prediction results as an image file with bbx and masks.\n",
+ " save_bbox_masks_labels(\n",
+ " result,\n",
+ " image_plot,\n",
+ " os.path.basename(image_path),\n",
+ " prediction_folder,\n",
+ " category_index,\n",
+ " PREDICTION_THRESHOLD,\n",
+ " )\n",
+ "\n",
+ " # Create object tracking data.\n",
+ " tracking_image = cv2.resize(\n",
+ " original_image,\n",
+ " (WIDTH_TRACKING, HEIGHT_TRACKING),\n",
+ " interpolation=cv2.INTER_AREA,\n",
+ " )\n",
+ " tracking_images[os.path.basename(image_path)] = tracking_image\n",
+ "\n",
+ " # Reducing mask sizes in order to keep the memory required for object\n",
+ " # tracking under a threshold.\n",
+ " result[\"detection_masks_tracking\"] = np.array([\n",
+ " cv2.resize(\n",
+ " i, (WIDTH_TRACKING, HEIGHT_TRACKING), interpolation=cv2.INTER_NEAREST\n",
+ " )\n",
+ " for i in result[\"detection_masks_reframed\"]\n",
+ " ])\n",
+ "\n",
+ " # In case of connected masks, keep the biggest mask and fill the holes\n",
+ " # in case of incomplete detections by Mask RCNN.\n",
+ " result[\"detection_masks_tracking\"] = np.array([\n",
+ " dilated_largest_component(i) for i in result[\"detection_masks_tracking\"]\n",
+ " ])\n",
+ "\n",
+ " # Crop objects from an image using masks for color detection.\n",
+ " cropped_objects = [\n",
+ " np.where(np.expand_dims(i, -1), image_plot, 0)\n",
+ " for i in result['detection_masks_reframed']\n",
+ " ]\n",
+ "\n",
+ " # Perform color detection using clustering approach.\n",
+ " dominant_colors = [\n",
+ " *map(\n",
+ " color_and_property_extractor.find_dominant_color, cropped_objects\n",
+ " )\n",
+ " ]\n",
+ " generic_color_names = color_and_property_extractor.get_generic_color_name(dominant_colors)\n",
+ "\n",
+ " # Extract features.\n",
+ " features = extract_properties(\n",
+ " tracking_image, result, \"detection_masks_tracking\"\n",
+ " )\n",
+ " features[\"source_name\"] = os.path.basename(os.path.dirname(image_path))\n",
+ " features[\"image_name\"] = os.path.basename(image_path)\n",
+ " features[\"creation_time\"] = get_image_creation_time(image_path)\n",
+ " features[\"frame\"] = frame\n",
+ " features[\"detection_scores\"] = result[\"detection_scores\"][0]\n",
+ " features[\"detection_classes\"] = result[\"detection_classes\"][0]\n",
+ " features[\"detection_classes_names\"] = result[\"detection_classes_names\"][0]\n",
+ " features[\"color\"] = generic_color_names\n",
+ " features_set.append(features)\n",
+ "\n",
+ "\n",
+ "if features_set:\n",
+ " features_df = pd.concat(features_set, ignore_index=True)\n",
+ " tracking_features = apply_tracking(\n",
+ " features_df,\n",
+ " search_range_x=SEARCH_RANGE_X,\n",
+ " search_range_y=SEARCH_RANGE_Y,\n",
+ " memory=MEMORY\n",
+ " )\n",
+ " agg_features = process_tracking_result(tracking_features)\n",
+ " counts = agg_features.groupby(\"detection_classes_names\").size()\n",
+ " counts.to_frame().to_csv(os.path.join(os.getcwd(), \"count.csv\"))\n",
+ " print(counts)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "tcEjtdLL-CnZ"
+ },
+ "source": [
+ "## Visualize Object Tracking "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "8yZOHwsS5GLL"
+ },
+ "outputs": [],
+ "source": [
+ "CIRCLE_RADIUS =7\n",
+ "font = cv2.FONT_HERSHEY_SIMPLEX\n",
+ "fontScale = 1\n",
+ "color = (255, 0, 0)\n",
+ "groups = tracking_features.groupby('image_name')\n",
+ "\n",
+ "for name, group in groups:\n",
+ " img = tracking_images[name].copy()\n",
+ " for k in range(len(group)):\n",
+ " cv2.circle(img,\n",
+ " (int(group.iloc[k]['x']),int(group.iloc[k]['y'])),\n",
+ " CIRCLE_RADIUS,\n",
+ " (255,133,233),\n",
+ " -1\n",
+ " )\n",
+ " cv2.putText(img,\n",
+ " str(int(group.iloc[k]['particle'])),\n",
+ " (int(group.iloc[k]['x']), int(group.iloc[k]['y'])),\n",
+ " font,\n",
+ " fontScale,\n",
+ " color,\n",
+ " 2,\n",
+ " cv2.LINE_AA\n",
+ " )\n",
+ "\n",
+ " cv2.imwrite(os.path.join('tracking',name), img)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "NW8GKKl7-Voy"
+ },
+ "source": [
+ "## Visualize Predictions by Categories"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JoJ9Bd9E-ezh"
+ },
+ "outputs": [],
+ "source": [
+ "if not agg_features.empty:\n",
+ " for group_name, df in tqdm.tqdm(agg_features.groupby(\"detection_classes_names\")):\n",
+ " os.makedirs(f'{output_dir}/{group_name}', exist_ok=True)\n",
+ "\n",
+ " for row in df.itertuples(index=False):\n",
+ " # Get the image\n",
+ " image = cv2.imread(os.path.join(images_dir, row.image_name))\n",
+ " image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n",
+ " new_h, new_w = image.shape[0], image.shape[1]\n",
+ "\n",
+ " # Get the bounding box and resize it\n",
+ " y1, x1, y2, x2 = row.bbox_0, row.bbox_1, row.bbox_2, row.bbox_3\n",
+ " new_bbox = resize_bbox(y1, x1, y2, x2, HEIGHT_TRACKING, WIDTH_TRACKING, new_h, new_w)\n",
+ "\n",
+ " # Include the score in the filename\n",
+ " score = row.detection_scores if hasattr(row, 'detection_scores') else 0.0\n",
+ " name = f'{os.path.splitext(row.image_name)[0]}_{row.particle}_{score:.2f}.png'\n",
+ "\n",
+ " # Save the cropped image\n",
+ " cv2.imwrite(f'{output_dir}/{row.detection_classes_names}/{name}',\n",
+ " image[new_bbox[0]:new_bbox[2], new_bbox[1]:new_bbox[3]])\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "hgmZvRqCDrnh"
+ },
+ "source": [
+ "## Copying folders to my Google drive"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "jgmX0qaLDuj9"
+ },
+ "outputs": [],
+ "source": [
+ "destination_folder = '/mydrive/circularnet/TestModel'\n",
+ "os.makedirs(destination_folder, exist_ok=True)\n",
+ "\n",
+ "# Function to safely copy directory, removing destination first if it exists\n",
+ "def copytree_replace(src, dst):\n",
+ " if os.path.exists(dst):\n",
+ " shutil.rmtree(dst)\n",
+ " shutil.copytree(src, dst)\n",
+ "\n",
+ "# Function to safely copy file, overwriting if it exists\n",
+ "def copy_replace(src, dst):\n",
+ " if os.path.exists(dst):\n",
+ " os.remove(dst)\n",
+ " shutil.copy(src, dst)\n",
+ "\n",
+ "copytree_replace(os.path.join(os.getcwd(), \"prediction_folder\"), os.path.join(destination_folder, \"prediction_folder\"))\n",
+ "copytree_replace(os.path.join(os.getcwd(), \"cropped_objects\"),os.path.join(destination_folder, \"cropped_objects\"))\n",
+ "copytree_replace(os.path.join(os.getcwd(), \"tracking\"),os.path.join(destination_folder, \"tracking\"))\n",
+ "copy_replace(os.path.join(os.getcwd(), \"count.csv\"),os.path.join(destination_folder, \"count.csv\"))"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "private_outputs": true,
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config.yaml b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config.yaml
new file mode 100644
index 00000000000..61a8cd6b612
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config.yaml
@@ -0,0 +1,242 @@
+# ---------------------------------------------------------
+# PET Bottle Grade Detection Pipeline Configuration
+# ---------------------------------------------------------
+paths:
+ # Local file-system source. Exactly one of `local.enable` / `gcs.enable`
+ # must be true (validated at config-load time).
+ #
+ # For BOTH sources, the input location must be the PARENT directory, not a
+ # folder of images. Its IMMEDIATE subfolders are each processed as an
+ # independent capture session; loose images placed directly in the parent
+ # are ignored.
+ #
+ # For accuracy reporting, name each subfolder EXACTLY after one of the class
+ # names listed under `classes` below (e.g. 'brown_bottles_grade3'). Matching
+ # is exact and case-insensitive, so a name that only contains a class as a
+ # substring (e.g. 'brown_bottles_grade3_batch1') is NOT treated as ground
+ # truth. When collapsed categories are enabled, the ground-truth CATEGORY
+ # (e.g. 'grade3') is derived from the matched class via the mapping below --
+ # the subfolder is NOT matched against category names. A subfolder whose
+ # name is not exactly a class name is reported with counts only (no class or
+ # category accuracy).
+ #
+ # Expected layout:
+ # /
+ # |-- brown_bottles_grade3/ -> images...
+ # |-- clean_jars_grade1/ -> images...
+ # `-- / -> images... (counts only, no accuracy)
+ #
+ # A run-wide 'summary.txt' collecting every subfolder's summary block is
+ # written to the active output root. In GCS mode it is written into
+ # temp_output_directory and uploaded to output_uri with the rest of the
+ # output; in local mode it is written into output_root_directory.
+ local:
+ enable: true
+ input_image_directory: "/home/umairsabir/new_data/test_data/ecosattva_pet_bottles_2026-07-28/2026-07-28_output" # TODO
+ output_root_directory: "/home/umairsabir/new_data/test_data/ecosattva_pet_bottles_2026-07-28/2026-07-28_output_model_output" # TODO
+ # Google Cloud Storage source. When enabled, input is downloaded from
+ # `input_uri` into `temp_input_directory`, the pipeline runs against the
+ # temp directories, and output is later uploaded to `output_uri`. The temp
+ # directories are pipeline-owned scratch space, kept separate from the
+ # local directories so real local data is never deleted or overwritten.
+ gcs:
+ enable: false
+ input_uri: "gs://recykal/latest_data/brown_bottles_grade3/images/" # TODO
+ output_uri: "gs://recykal/exp_output/" # TODO
+ temp_input_directory: "/home/umairsabir/bottle_grade/gcs_temp/input/" # TODO
+ temp_output_directory: "/home/umairsabir/bottle_grade/gcs_temp/output/" # TODO
+ # Output subfolders below are created INSIDE each per-subfolder result
+ # directory. They are shared across both sources.
+ output_frame_subfolder: "tracked_frames"
+ output_video_filename: "tracked_output.mp4"
+ track_grid_subfolder: "track_grids_by_category"
+
+
+# ---------------------------------------------------------
+# BigQuery result ingestion. Independent of the input/output source: results
+# can be written to BigQuery whether the pipeline runs in local or GCS mode.
+# When enable is true, project_id / dataset_id / table_id must be non-empty
+# (validated at config-load time).
+# ---------------------------------------------------------
+bigquery:
+ enable: false
+ project_id: "waste-identification-ml-330916"
+ dataset_id: "umair_packet_dataset"
+ table_id: "umair_bottle_table"
+ # When true, the table's existing rows are replaced by this run's rows.
+ # When false, this run's rows are appended.
+ overwrite: true
+
+
+models:
+ # RFDETR segmentation model. RFDETR replaces the previous SAM3 detector
+ # and is used purely as an object finder: its class predictions are
+ # discarded, and every detection contributes a mask + box + score to the
+ # tracking and DINOv3 grading stages downstream. Use the CircularNet
+ # fine-tuned checkpoint for waste-specific object recall.
+ #
+ # This block holds everything the model itself needs: checkpoint,
+ # device, its own confidence gate, and the file extensions that decide
+ # which images get fed into it.
+ rfdetr:
+ checkpoint_path: "/home/umairsabir/rfdetr_weights/checkpoint_best_total.pth" # TODO
+ device: "cuda" # Will fallback to 'cpu' automatically if unavailable
+ # File extensions used to discover the images that get fed to RFDETR.
+ image_file_extensions: ["*.png", "*.jpg", "*.jpeg"]
+ # RFDETR's own confidence gate: detections below this never reach
+ # the pipeline. Passed as `RFDETR.predict(threshold=...)`.
+ predict_threshold: 0.2
+ dinov3:
+ repo_dir: "/home/umairsabir/dinov3" # TODO
+ checkpoint_path: "/home/umairsabir/tracking/dinov3_21_classes_model/best_model.pth" # TODO
+ model_name: "dinov3_vitl16"
+ inference_image_size: 256
+ classification_batch_size: 32
+ # Standard ImageNet normalization means and standard deviations
+ image_mean: [0.485, 0.456, 0.406]
+ image_std: [0.229, 0.224, 0.225]
+
+
+# The complete list of target classes in exact label-index order.
+# To get per-subfolder class accuracy, name each input subfolder EXACTLY after
+# one of these names (see the note under `paths` above).
+classes:
+- 'brown_bottles_grade3'
+- 'clean_PET_cold_drink_bottles_with_label_cap_ring_grade1'
+- 'clean_PET_cold_drink_bottles_without_label_cap_ring_grade1'
+- 'clean_PET_mango_juice_bottles_without_label_cap_ring_grade1'
+- 'clean_PET_water_bottles_with_label_cap_ring_grade1'
+- 'clean_PET_water_bottles_without_label_cap_ring_grade1'
+- 'clean_jars_grade1'
+- 'clean_liquor_bottles_without_label_cap_ring_grade1'
+- 'coloured_PET_bottles_grade3'
+- 'dirt_PET_cold_drink_bottles_with_label_cap_ring_grade3'
+- 'dirt_PET_cold_drink_bottles_without_label_cap_ring_grade3'
+- 'dirt_PET_mango_juice_bottles_without_label_cap_ring_grade3'
+- 'dirt_PET_water_bottles_with_label_cap_ring_grade3'
+- 'dirt_PET_water_bottles_without_label_cap_ring_grade3'
+- 'dirt_jars_grade3'
+- 'dirt_liquor_bottles_without_label_cap_ring_grade3'
+- 'full_sleeved_bottles_grade3'
+- 'green_bottles_grade3'
+- 'liquor_bottles_with_label_cap_ring_grade3'
+- 'non_food_bottles_grade3'
+- 'partially_sleeved_mango_juice_bottles_with_label_cap_ring_grade3'
+
+
+# ---------------------------------------------------------
+# Optional grouping of fine-grained classes into broader categories.
+# When enable=true, every class above MUST appear in exactly one category
+# (validated at config-load time). When enable=false, the 'mapping' is
+# ignored and per-category reporting / nested folders are skipped.
+# Classes are grouped by their grade suffix.
+#
+# When enabled, per-subfolder CATEGORY accuracy is derived from the matched
+# ground-truth class via this mapping (e.g. a subfolder named exactly
+# 'brown_bottles_grade3' has category ground truth 'grade3'). Subfolders are
+# not matched against these category names directly.
+# ---------------------------------------------------------
+collapsed_categories:
+ enable: true
+ mapping:
+ grade1:
+ - 'clean_PET_cold_drink_bottles_with_label_cap_ring_grade1'
+ - 'clean_PET_cold_drink_bottles_without_label_cap_ring_grade1'
+ - 'clean_PET_mango_juice_bottles_without_label_cap_ring_grade1'
+ - 'clean_PET_water_bottles_with_label_cap_ring_grade1'
+ - 'clean_PET_water_bottles_without_label_cap_ring_grade1'
+ - 'clean_jars_grade1'
+ - 'clean_liquor_bottles_without_label_cap_ring_grade1'
+ grade3:
+ - 'brown_bottles_grade3'
+ - 'coloured_PET_bottles_grade3'
+ - 'dirt_PET_cold_drink_bottles_with_label_cap_ring_grade3'
+ - 'dirt_PET_cold_drink_bottles_without_label_cap_ring_grade3'
+ - 'dirt_PET_mango_juice_bottles_without_label_cap_ring_grade3'
+ - 'dirt_PET_water_bottles_with_label_cap_ring_grade3'
+ - 'dirt_PET_water_bottles_without_label_cap_ring_grade3'
+ - 'dirt_jars_grade3'
+ - 'dirt_liquor_bottles_without_label_cap_ring_grade3'
+ - 'full_sleeved_bottles_grade3'
+ - 'green_bottles_grade3'
+ - 'liquor_bottles_with_label_cap_ring_grade3'
+ - 'non_food_bottles_grade3'
+ - 'partially_sleeved_mango_juice_bottles_with_label_cap_ring_grade3'
+
+
+# ---------------------------------------------------------
+# Image preprocessing applied BEFORE any model sees the image. Detector
+# independent -- swapping RFDETR for another detector would still need
+# this stage.
+# ---------------------------------------------------------
+preprocessing:
+ # Maximum length of the shorter image side. Images larger than this are
+ # downscaled preserving aspect ratio; smaller images are left unchanged.
+ max_short_side: 1024
+
+
+# ---------------------------------------------------------
+# Filters applied to RFDETR's raw output before tracking. Order:
+# 1. containment_threshold -- mask-in-mask filter.
+# 2. merge_containment_threshold -- box-in-box merge (unconditional).
+# 3. score_threshold -- final cutoff before tracking.
+# ---------------------------------------------------------
+post_processing:
+ # Ratio above which a smaller mask is removed when contained inside a
+ # larger mask (intersection / smaller_mask_area).
+ containment_threshold: 0.98
+ # Ratio above which a smaller box is merged into a larger one
+ # (intersection_area / smaller_box_area).
+ merge_containment_threshold: 0.7
+ # Final score cutoff applied when handing detections to the tracker.
+ # Set to 0.0 to keep every detection that survived the earlier filters.
+ score_threshold: 0.20
+
+
+# ---------------------------------------------------------
+# Per-track crop geometry. These crops are what feed DINOv3, not RFDETR --
+# RFDETR sees the resized input image directly.
+# ---------------------------------------------------------
+cropping:
+ # Output letterbox size for track crops as [height, width].
+ crop_size: [256, 256]
+ # Pixel buffer added around tight bounding boxes before cropping (guards
+ # against edge clipping).
+ crop_buffer_pixels: 5
+
+
+tracking:
+ # Set to false when input images are independent (not consecutive video
+ # frames). When false, ByteTrack is bypassed entirely and every detection
+ # in every frame is assigned a fresh sequential ID. The two ByteTrack
+ # parameters below are ignored in that mode, and the per-subfolder summary
+ # block is skipped.
+ enable: true
+ bytetrack_minimum_iou_threshold: 0.1
+ bytetrack_minimum_consecutive_frames: 2
+
+
+visualization:
+ # ---------------------------------------------------------------------------
+ # Output Toggles
+ # Disable these to save disk space and drastically speed up processing time.
+ # ---------------------------------------------------------------------------
+ # If true, creates the folder specified in `output_frame_subfolder` (e.g., "tracked_frames/")
+ # inside each per-subfolder result directory and saves individual annotated PNGs.
+ save_frames: true
+ # If true, creates the MP4 file specified in `output_video_filename`
+ # (e.g., "tracked_output.mp4") inside each per-subfolder result directory.
+ save_video: false
+ # If true, creates the folder specified in `track_grid_subfolder` inside each
+ # per-subfolder result directory, with nested subfolders by category and class.
+ save_track_grids: true
+ # ---------------------------------------------------------------------------
+ # Render Settings
+ # ---------------------------------------------------------------------------
+ output_video_fps: 1
+ show_confidence_in_labels: true
+ # ImageNet mean used for crop backgrounds to match DINOv3 expectations
+ background_blend_color_rgb: [124, 116, 104]
+ track_grid_columns_per_row: 5
+ track_grid_thumbnail_size_inches: 3
+ track_grid_dpi: 150
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config_loader.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config_loader.py
new file mode 100644
index 00000000000..8b44384aa86
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config_loader.py
@@ -0,0 +1,609 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Loads and validates pipeline configurations using typed data containers."""
+
+import dataclasses
+from typing import Any, Self
+import yaml
+
+
+class ConfigurationError(Exception):
+ """Base exception for configuration errors."""
+
+ pass
+
+
+@dataclasses.dataclass(frozen=True)
+class LocalPathsConfig:
+ """Local file-system input and output locations.
+
+ Attributes:
+ enable: Whether local paths are the active source.
+ input_image_directory: Root directory containing immediate subfolders,
+ each holding the image frames of one capture session.
+ output_root_directory: Root directory where a per-subfolder result
+ directory will be created (mirroring the subfolder name).
+ """
+
+ enable: bool
+ input_image_directory: str
+ output_root_directory: str
+
+
+@dataclasses.dataclass(frozen=True)
+class GCSPathsConfig:
+ """Google Cloud Storage input and output locations.
+
+ Attributes:
+ enable: Whether GCS paths are the active source.
+ input_uri: GCS URI (e.g. 'gs://bucket/path/') to read input from.
+ output_uri: GCS URI (e.g. 'gs://bucket/path/') to write output to.
+ temp_input_directory: Local scratch directory that GCS input is downloaded
+ into. Pipeline-owned; cleared before each download.
+ temp_output_directory: Local scratch directory that pipeline output is
+ written to before being uploaded to GCS.
+ """
+
+ enable: bool
+ input_uri: str
+ output_uri: str
+ temp_input_directory: str
+ temp_output_directory: str
+
+
+@dataclasses.dataclass(frozen=True)
+class PathsConfig:
+ """File-system paths used by the pipeline.
+
+ Exactly one of `local` or `gcs` is enabled (validated at config-load
+ time). The output subfolder and file names are shared across both
+ sources.
+
+ Attributes:
+ local: Local input/output locations and their enable switch.
+ gcs: GCS input/output locations and their enable switch.
+ output_frame_subfolder: Name of the per-subfolder folder that stores
+ annotated frames.
+ output_video_filename: Filename of the annotated tracking video saved
+ inside each per-subfolder result directory.
+ track_grid_subfolder: Name of the per-subfolder folder that stores
+ track-grid PNGs grouped by collapsed category and class.
+ """
+
+ local: LocalPathsConfig
+ gcs: GCSPathsConfig
+ output_frame_subfolder: str
+ output_video_filename: str
+ track_grid_subfolder: str
+
+
+@dataclasses.dataclass(frozen=True)
+class BigQueryConfig:
+ """BigQuery ingestion settings.
+
+ Independent of the input/output source: results may be written to
+ BigQuery whether the pipeline runs in local or GCS mode.
+
+ Attributes:
+ enable: Whether pipeline results are written to BigQuery.
+ project_id: Google Cloud project ID that owns the dataset.
+ dataset_id: BigQuery dataset ID that holds the table.
+ table_id: BigQuery table ID that receives the rows.
+ overwrite: When True, existing rows are replaced by this run's rows. When
+ False, this run's rows are appended.
+ """
+
+ enable: bool
+ project_id: str
+ dataset_id: str
+ table_id: str
+ overwrite: bool
+
+
+@dataclasses.dataclass(frozen=True)
+class RFDETRConfig:
+ """RFDETR segmentation-model settings.
+
+ Everything the model itself needs lives here: checkpoint, device, its
+ own confidence gate, and the file extensions used to discover images
+ the model will be run on.
+
+ Attributes:
+ checkpoint_path: Absolute path to the RFDETR .pth checkpoint (typically
+ the CircularNet fine-tuned weights).
+ device: Target torch device string, e.g. 'cuda'. Falls back to 'cpu'
+ automatically when CUDA is unavailable.
+ image_file_extensions: Glob patterns used to discover input images that
+ will be fed to the model (e.g. ``["*.png", "*.jpg"]``).
+ predict_threshold: Confidence threshold passed to ``RFDETR.predict``.
+ Detections below this never reach the pipeline.
+ """
+
+ checkpoint_path: str
+ device: str
+ image_file_extensions: list[str]
+ predict_threshold: float
+
+
+@dataclasses.dataclass(frozen=True)
+class DINOv3Config:
+ """DINOv3 classification-model settings.
+
+ Attributes:
+ repo_dir: Absolute path to the DINOv3 repository directory.
+ checkpoint_path: Absolute path to the DINOv3 checkpoint.
+ model_name: Name of the DINOv3 model variant to load.
+ inference_image_size: Target image size for DINOv3 inference (assumed
+ square).
+ classification_batch_size: Batch size for DINOv3 classification.
+ image_mean: Mean RGB values for image normalization.
+ image_std: Standard deviation RGB values for image normalization.
+ """
+
+ repo_dir: str
+ checkpoint_path: str
+ model_name: str
+ inference_image_size: int
+ classification_batch_size: int
+ image_mean: tuple[float, float, float]
+ image_std: tuple[float, float, float]
+
+
+@dataclasses.dataclass(frozen=True)
+class ModelsConfig:
+ """Container for all model configurations.
+
+ Attributes:
+ rfdetr: Configuration for the RFDETR segmentation model.
+ dinov3: Configuration for the DINOv3 classification model.
+ """
+
+ rfdetr: RFDETRConfig
+ dinov3: DINOv3Config
+
+
+@dataclasses.dataclass(frozen=True)
+class PreprocessingConfig:
+ """Image preprocessing applied before any model sees the image.
+
+ Attributes:
+ max_short_side: Maximum length of the shorter image side. Images larger
+ than this are downscaled preserving aspect ratio.
+ """
+
+ max_short_side: int
+
+
+@dataclasses.dataclass(frozen=True)
+class PostProcessingConfig:
+ """Filters applied to RFDETR's raw output before tracking.
+
+ Applied in this order per image:
+
+ 1. containment_threshold -- mask-in-mask filter.
+ 2. merge_containment_threshold -- box-in-box merge (unconditional).
+ 3. score_threshold -- final cutoff before tracking.
+
+ Attributes:
+ containment_threshold: Mask-in-mask cutoff; the smaller mask is dropped
+ when its intersection with a larger mask exceeds this fraction of its
+ own area.
+ merge_containment_threshold: Box-in-box cutoff for the merge step; the
+ smaller detection is merged into the larger one when their box
+ intersection exceeds this fraction of the smaller box's area.
+ score_threshold: Final score cutoff applied when converting the state dict
+ to ``supervision.Detections``.
+ """
+
+ containment_threshold: float
+ merge_containment_threshold: float
+ score_threshold: float
+
+
+@dataclasses.dataclass(frozen=True)
+class CroppingConfig:
+ """Per-track crop geometry (crops that feed DINOv3).
+
+ Attributes:
+ crop_size: Output letterbox size as a ``(height, width)`` tuple.
+ crop_buffer_pixels: Pixel buffer added around each detection box before
+ cropping (guards against edge clipping).
+ """
+
+ crop_size: tuple[int, int]
+ crop_buffer_pixels: int
+
+
+@dataclasses.dataclass(frozen=True)
+class TrackingConfig:
+ """ByteTrack settings and a toggle to bypass tracking entirely.
+
+ Attributes:
+ bytetrack_minimum_iou_threshold: Minimum IoU for ByteTrack to link a
+ detection to an existing track. Ignored when `enable` is False.
+ bytetrack_minimum_consecutive_frames: Minimum frames a track must persist
+ before ByteTrack emits it. Ignored when `enable` is False.
+ enable: When True (default), ByteTrack runs normally and IDs are stable
+ across frames. When False, tracking is bypassed entirely and every
+ detection in every frame receives a fresh sequential ID. Use False when
+ input images are independent (not consecutive video frames).
+ """
+
+ bytetrack_minimum_iou_threshold: float
+ bytetrack_minimum_consecutive_frames: int
+ enable: bool = True
+
+
+@dataclasses.dataclass(frozen=True)
+class VisualizationConfig:
+ save_frames: bool
+ save_video: bool
+ save_track_grids: bool
+ output_video_fps: int
+ show_confidence_in_labels: bool
+ background_blend_color_rgb: tuple[int, int, int]
+ track_grid_columns_per_row: int
+ track_grid_thumbnail_size_inches: int
+ track_grid_dpi: int
+
+
+@dataclasses.dataclass(frozen=True)
+class CollapsedCategoriesConfig:
+ """Optional grouping of fine-grained classes into broader categories.
+
+ When enabled, every class in the pipeline's `classes` list must be
+ assigned to exactly one category. The category for a given class is
+ looked up via `get_category_for_class`. When disabled, the mapping is
+ empty and no per-category reporting is performed.
+
+ Attributes:
+ enable: Whether the collapsed-category feature is active.
+ mapping: Mapping from category name to the list of class names that fall
+ under it. Empty when disabled.
+ """
+
+ enable: bool
+ mapping: dict[str, list[str]] = dataclasses.field(default_factory=dict)
+
+ def get_category_for_class(self, class_name: str) -> str | None:
+ """Returns the category that contains the given class, or None if disabled.
+
+ Args:
+ class_name: The fine-grained class name to look up.
+
+ Returns:
+ The matching category name when the feature is enabled, or None
+ when the feature is disabled.
+
+ Raises:
+ ConfigurationError: If the feature is enabled but the class is
+ not present in any category (this indicates the validation
+ in `from_yaml` was bypassed).
+ """
+ if not self.enable:
+ return None
+ for category_name, class_list in self.mapping.items():
+ if class_name in class_list:
+ return category_name
+ raise ConfigurationError(
+ f"Class '{class_name}' is not assigned to any collapsed category."
+ )
+
+ def get_category_names(self) -> list[str]:
+ """Returns the configured category names in declaration order.
+
+ Returns:
+ List of category names. Empty list when the feature is disabled.
+ """
+ if not self.enable:
+ return []
+ return list(self.mapping.keys())
+
+
+@dataclasses.dataclass(frozen=True)
+class PipelineConfig:
+ """Root configuration object representing the entire pipeline state."""
+
+ paths: PathsConfig
+ bigquery: BigQueryConfig
+ models: ModelsConfig
+ classes: list[str]
+ preprocessing: PreprocessingConfig
+ post_processing: PostProcessingConfig
+ cropping: CroppingConfig
+ tracking: TrackingConfig
+ visualization: VisualizationConfig
+ collapsed_categories: CollapsedCategoriesConfig
+
+ @classmethod
+ def from_yaml(cls, yaml_path: str) -> Self:
+ """Parses the YAML file into a strictly typed configuration object.
+
+ Args:
+ yaml_path: Path to the YAML configuration file.
+
+ Returns:
+ A fully populated PipelineConfig.
+
+ Raises:
+ ConfigurationError: If the YAML file cannot be found, if exactly
+ one path source is not enabled, if BigQuery is enabled with a
+ missing ID, or if the collapsed_categories section is enabled
+ but invalid.
+ """
+ try:
+ with open(yaml_path, "r", encoding="utf-8") as file:
+ data = yaml.safe_load(file)
+ except FileNotFoundError as err:
+ raise ConfigurationError(f"Config file not found: {yaml_path}") from err
+
+ classes = data["classes"]
+ collapsed_categories = cls._build_collapsed_categories_config(
+ raw_section=data.get("collapsed_categories"),
+ classes=classes,
+ )
+ return cls(
+ paths=cls._build_paths_config(data["paths"]),
+ bigquery=cls._build_bigquery_config(data.get("bigquery")),
+ models=ModelsConfig(
+ rfdetr=cls._build_rfdetr_config(data["models"]["rfdetr"]),
+ dinov3=DINOv3Config(**{
+ **data["models"]["dinov3"],
+ "image_mean": tuple(data["models"]["dinov3"]["image_mean"]),
+ "image_std": tuple(data["models"]["dinov3"]["image_std"]),
+ }),
+ ),
+ classes=classes,
+ preprocessing=cls._build_preprocessing_config(data["preprocessing"]),
+ post_processing=cls._build_post_processing_config(
+ data["post_processing"]
+ ),
+ cropping=cls._build_cropping_config(data["cropping"]),
+ tracking=TrackingConfig(**data["tracking"]),
+ visualization=VisualizationConfig(**{
+ **data["visualization"],
+ "background_blend_color_rgb": tuple(
+ data["visualization"]["background_blend_color_rgb"]
+ ),
+ }),
+ collapsed_categories=collapsed_categories,
+ )
+
+ @staticmethod
+ def _build_paths_config(raw_section: dict[str, Any]) -> PathsConfig:
+ """Builds and validates the PathsConfig from raw YAML.
+
+ Exactly one of the local or GCS sources must be enabled. Enabling
+ both, or enabling neither, is a configuration error.
+
+ Args:
+ raw_section: The raw `paths` dict from YAML.
+
+ Returns:
+ A validated PathsConfig instance.
+
+ Raises:
+ ConfigurationError: If both sources are enabled or both are
+ disabled.
+ """
+ local_section = raw_section["local"]
+ gcs_section = raw_section["gcs"]
+ local_enabled = local_section["enable"]
+ gcs_enabled = gcs_section["enable"]
+
+ if local_enabled and gcs_enabled:
+ raise ConfigurationError(
+ "Both paths.local.enable and paths.gcs.enable are true. "
+ "Exactly one source must be enabled."
+ )
+ if not local_enabled and not gcs_enabled:
+ raise ConfigurationError(
+ "Both paths.local.enable and paths.gcs.enable are false. "
+ "Exactly one source must be enabled."
+ )
+
+ return PathsConfig(
+ local=LocalPathsConfig(**local_section),
+ gcs=GCSPathsConfig(**gcs_section),
+ output_frame_subfolder=raw_section["output_frame_subfolder"],
+ output_video_filename=raw_section["output_video_filename"],
+ track_grid_subfolder=raw_section["track_grid_subfolder"],
+ )
+
+ @staticmethod
+ def _build_bigquery_config(
+ raw_section: dict[str, Any] | None,
+ ) -> BigQueryConfig:
+ """Builds and validates the BigQueryConfig from raw YAML.
+
+ The section is optional: when it is missing or disabled, a disabled
+ config is returned. When enabled, the project, dataset, and table
+ IDs must all be non-empty.
+
+ Args:
+ raw_section: The raw `bigquery` dict from YAML, or None if the section
+ was omitted entirely.
+
+ Returns:
+ A validated BigQueryConfig instance.
+
+ Raises:
+ ConfigurationError: If enabled but any required ID is empty.
+ """
+ if raw_section is None or not raw_section.get("enable", False):
+ return BigQueryConfig(
+ enable=False,
+ project_id="",
+ dataset_id="",
+ table_id="",
+ overwrite=False,
+ )
+
+ required_fields = ("project_id", "dataset_id", "table_id")
+ missing_fields = [
+ field_name
+ for field_name in required_fields
+ if not raw_section.get(field_name)
+ ]
+ if missing_fields:
+ raise ConfigurationError(
+ "bigquery.enable is true but these required fields are empty: "
+ f"{missing_fields}"
+ )
+
+ return BigQueryConfig(
+ enable=True,
+ project_id=raw_section["project_id"],
+ dataset_id=raw_section["dataset_id"],
+ table_id=raw_section["table_id"],
+ overwrite=raw_section.get("overwrite", False),
+ )
+
+ @staticmethod
+ def _build_rfdetr_config(raw_section: dict[str, Any]) -> RFDETRConfig:
+ """Builds and validates the RFDETRConfig from raw YAML.
+
+ Args:
+ raw_section: The raw `models.rfdetr` dict from YAML.
+
+ Returns:
+ A validated RFDETRConfig instance.
+ """
+ return RFDETRConfig(
+ checkpoint_path=raw_section["checkpoint_path"],
+ device=raw_section["device"],
+ image_file_extensions=list(raw_section["image_file_extensions"]),
+ predict_threshold=float(raw_section["predict_threshold"]),
+ )
+
+ @staticmethod
+ def _build_preprocessing_config(
+ raw_section: dict[str, Any],
+ ) -> PreprocessingConfig:
+ """Builds and validates the PreprocessingConfig from raw YAML.
+
+ Args:
+ raw_section: The raw `preprocessing` dict from YAML.
+
+ Returns:
+ A validated PreprocessingConfig instance.
+ """
+ return PreprocessingConfig(
+ max_short_side=int(raw_section["max_short_side"]),
+ )
+
+ @staticmethod
+ def _build_post_processing_config(
+ raw_section: dict[str, Any],
+ ) -> PostProcessingConfig:
+ """Builds and validates the PostProcessingConfig from raw YAML.
+
+ Args:
+ raw_section: The raw `post_processing` dict from YAML.
+
+ Returns:
+ A validated PostProcessingConfig instance.
+ """
+ return PostProcessingConfig(
+ containment_threshold=float(raw_section["containment_threshold"]),
+ merge_containment_threshold=float(
+ raw_section["merge_containment_threshold"]
+ ),
+ score_threshold=float(raw_section["score_threshold"]),
+ )
+
+ @staticmethod
+ def _build_cropping_config(raw_section: dict[str, Any]) -> CroppingConfig:
+ """Builds and validates the CroppingConfig from raw YAML.
+
+ Args:
+ raw_section: The raw `cropping` dict from YAML.
+
+ Returns:
+ A validated CroppingConfig instance.
+ """
+ return CroppingConfig(
+ crop_size=tuple(raw_section["crop_size"]),
+ crop_buffer_pixels=int(raw_section["crop_buffer_pixels"]),
+ )
+
+ @staticmethod
+ def _build_collapsed_categories_config(
+ raw_section: dict[str, Any] | None, classes: list[str]
+ ) -> CollapsedCategoriesConfig:
+ """Builds and validates the CollapsedCategoriesConfig from raw YAML.
+
+ The YAML section is expected to look like:
+ collapsed_categories:
+ enable: true
+ mapping:
+ category_name: [class_a, class_b]
+
+ Validation rules when enabled:
+ - Every class in `classes` must appear in exactly one category.
+ - No class may appear in more than one category.
+ - Every class in the mapping must be present in `classes`.
+
+ Args:
+ raw_section: The raw `collapsed_categories` dict from YAML, or None if
+ the section was omitted entirely.
+ classes: The full list of class names from the config.
+
+ Returns:
+ A validated CollapsedCategoriesConfig instance.
+
+ Raises:
+ ConfigurationError: If validation fails.
+ """
+ if raw_section is None or not raw_section.get("enable", False):
+ return CollapsedCategoriesConfig(enable=False, mapping={})
+
+ raw_mapping = raw_section.get("mapping") or {}
+ if not isinstance(raw_mapping, dict) or not raw_mapping:
+ raise ConfigurationError(
+ "collapsed_categories.enable is true but 'mapping' is empty or"
+ " missing."
+ )
+
+ seen_classes: dict[str, str] = {}
+ for category_name, class_list in raw_mapping.items():
+ if not isinstance(class_list, list) or not class_list:
+ raise ConfigurationError(
+ f"Category '{category_name}' must map to a non-empty list of class"
+ " names."
+ )
+ for class_name in class_list:
+ if class_name not in classes:
+ raise ConfigurationError(
+ f"Class '{class_name}' in category '{category_name}' is not "
+ "present in the top-level 'classes' list."
+ )
+ if class_name in seen_classes:
+ raise ConfigurationError(
+ f"Class '{class_name}' is assigned to both "
+ f"'{seen_classes[class_name]}' and '{category_name}'."
+ )
+ seen_classes[class_name] = category_name
+
+ unmapped_classes = [
+ class_name for class_name in classes if class_name not in seen_classes
+ ]
+ if unmapped_classes:
+ raise ConfigurationError(
+ "The following classes are not assigned to any collapsed "
+ f"category: {unmapped_classes}"
+ )
+
+ return CollapsedCategoriesConfig(enable=True, mapping=dict(raw_mapping))
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config_loader_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config_loader_test.py
new file mode 100644
index 00000000000..93e13810370
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/config_loader_test.py
@@ -0,0 +1,379 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for config_loader.py."""
+
+import copy
+import dataclasses
+import pathlib
+from typing import Any
+
+from absl.testing import absltest
+import yaml
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import config_loader
+
+
+def _valid_config_mapping() -> dict[str, Any]:
+ """Returns a fully valid config mapping usable as a test baseline.
+
+ Individual tests deep-copy this and mutate a single section so each test
+ isolates exactly one rule.
+ """
+ return {
+ "paths": {
+ "local": {
+ "enable": True,
+ "input_image_directory": "/data/in",
+ "output_root_directory": "/data/out",
+ },
+ "gcs": {
+ "enable": False,
+ "input_uri": "gs://bucket/in/",
+ "output_uri": "gs://bucket/out/",
+ "temp_input_directory": "/tmp/in",
+ "temp_output_directory": "/tmp/out",
+ },
+ "output_frame_subfolder": "tracked_frames",
+ "output_video_filename": "tracked_output.mp4",
+ "track_grid_subfolder": "track_grids_by_category",
+ },
+ "bigquery": {
+ "enable": False,
+ "project_id": "proj",
+ "dataset_id": "ds",
+ "table_id": "tbl",
+ "overwrite": True,
+ },
+ "models": {
+ "rfdetr": {
+ "checkpoint_path": "/models/rfdetr.pth",
+ "device": "cuda",
+ "image_file_extensions": ["*.png", "*.jpg"],
+ "predict_threshold": 0.2,
+ },
+ "dinov3": {
+ "repo_dir": "/models/dinov3",
+ "checkpoint_path": "/models/dinov3.pth",
+ "model_name": "dinov3_vitl16",
+ "inference_image_size": 256,
+ "classification_batch_size": 32,
+ "image_mean": [0.485, 0.456, 0.406],
+ "image_std": [0.229, 0.224, 0.225],
+ },
+ },
+ "classes": ["clean_grade1", "dirty_grade3"],
+ "preprocessing": {"max_short_side": 1024},
+ "post_processing": {
+ "containment_threshold": 0.98,
+ "merge_containment_threshold": 0.7,
+ "score_threshold": 0.2,
+ },
+ "cropping": {"crop_size": [256, 256], "crop_buffer_pixels": 5},
+ "tracking": {
+ "enable": True,
+ "bytetrack_minimum_iou_threshold": 0.1,
+ "bytetrack_minimum_consecutive_frames": 2,
+ },
+ "visualization": {
+ "save_frames": True,
+ "save_video": False,
+ "save_track_grids": True,
+ "output_video_fps": 1,
+ "show_confidence_in_labels": True,
+ "background_blend_color_rgb": [124, 116, 104],
+ "track_grid_columns_per_row": 5,
+ "track_grid_thumbnail_size_inches": 3,
+ "track_grid_dpi": 150,
+ },
+ "collapsed_categories": {
+ "enable": True,
+ "mapping": {
+ "grade1": ["clean_grade1"],
+ "grade3": ["dirty_grade3"],
+ },
+ },
+ }
+
+
+class CollapsedCategoriesConfigTest(absltest.TestCase):
+ """Tests for the CollapsedCategoriesConfig behavior methods."""
+
+ def test_get_category_for_class_when_disabled_returns_none(self):
+ """Verifies a disabled config returns None for any class."""
+ config = config_loader.CollapsedCategoriesConfig(enable=False, mapping={})
+ self.assertIsNone(config.get_category_for_class("anything"))
+
+ def test_get_category_for_class_returns_matching_category(self):
+ """Verifies the containing category is returned for a mapped class."""
+ config = config_loader.CollapsedCategoriesConfig(
+ enable=True, mapping={"grade1": ["a", "b"], "grade3": ["c"]}
+ )
+ self.assertEqual(config.get_category_for_class("c"), "grade3")
+
+ def test_get_category_for_class_raises_for_unmapped_class(self):
+ """Verifies an enabled config raises when the class is unmapped."""
+ config = config_loader.CollapsedCategoriesConfig(
+ enable=True, mapping={"grade1": ["a"]}
+ )
+ with self.assertRaises(config_loader.ConfigurationError):
+ config.get_category_for_class("missing")
+
+ def test_get_category_names_when_disabled_is_empty(self):
+ """Verifies a disabled config reports no category names."""
+ config = config_loader.CollapsedCategoriesConfig(enable=False, mapping={})
+ self.assertEqual(config.get_category_names(), [])
+
+ def test_get_category_names_preserves_declaration_order(self):
+ """Verifies category names are returned in mapping insertion order."""
+ config = config_loader.CollapsedCategoriesConfig(
+ enable=True, mapping={"grade3": ["c"], "grade1": ["a"]}
+ )
+ self.assertEqual(config.get_category_names(), ["grade3", "grade1"])
+
+
+class BuildPathsConfigTest(absltest.TestCase):
+ """Tests for _build_paths_config source-selection validation."""
+
+ def test_local_only_is_valid(self):
+ """Verifies enabling only local produces a PathsConfig."""
+ section = _valid_config_mapping()["paths"]
+ result = config_loader.PipelineConfig._build_paths_config(section)
+ self.assertTrue(result.local.enable)
+ self.assertFalse(result.gcs.enable)
+
+ def test_raises_when_both_sources_enabled(self):
+ """Verifies enabling both local and GCS raises."""
+ section = _valid_config_mapping()["paths"]
+ section["local"]["enable"] = True
+ section["gcs"]["enable"] = True
+ with self.assertRaisesRegex(
+ config_loader.ConfigurationError, "Exactly one"
+ ):
+ config_loader.PipelineConfig._build_paths_config(section)
+
+ def test_raises_when_neither_source_enabled(self):
+ """Verifies enabling neither source raises."""
+ section = _valid_config_mapping()["paths"]
+ section["local"]["enable"] = False
+ section["gcs"]["enable"] = False
+ with self.assertRaisesRegex(
+ config_loader.ConfigurationError, "Exactly one"
+ ):
+ config_loader.PipelineConfig._build_paths_config(section)
+
+
+class BuildBigQueryConfigTest(absltest.TestCase):
+ """Tests for _build_bigquery_config."""
+
+ def test_missing_section_returns_disabled(self):
+ """Verifies a None section yields a disabled BigQuery config."""
+ result = config_loader.PipelineConfig._build_bigquery_config(None)
+ self.assertFalse(result.enable)
+ self.assertEqual(result.project_id, "")
+
+ def test_disabled_section_returns_disabled(self):
+ """Verifies an explicitly disabled section yields a disabled config."""
+ result = config_loader.PipelineConfig._build_bigquery_config(
+ {"enable": False}
+ )
+ self.assertFalse(result.enable)
+
+ def test_enabled_with_all_ids_is_valid(self):
+ """Verifies an enabled section with all IDs is accepted."""
+ result = config_loader.PipelineConfig._build_bigquery_config({
+ "enable": True,
+ "project_id": "p",
+ "dataset_id": "d",
+ "table_id": "t",
+ "overwrite": True,
+ })
+ self.assertTrue(result.enable)
+ self.assertTrue(result.overwrite)
+
+ def test_enabled_with_missing_id_raises(self):
+ """Verifies an enabled section missing an ID raises, naming the field."""
+ with self.assertRaisesRegex(config_loader.ConfigurationError, "table_id"):
+ config_loader.PipelineConfig._build_bigquery_config({
+ "enable": True,
+ "project_id": "p",
+ "dataset_id": "d",
+ "table_id": "",
+ })
+
+ def test_overwrite_defaults_to_false(self):
+ """Verifies overwrite defaults to False when omitted."""
+ result = config_loader.PipelineConfig._build_bigquery_config({
+ "enable": True,
+ "project_id": "p",
+ "dataset_id": "d",
+ "table_id": "t",
+ })
+ self.assertFalse(result.overwrite)
+
+
+class BuildCollapsedCategoriesConfigTest(absltest.TestCase):
+ """Tests for _build_collapsed_categories_config validation."""
+
+ def test_disabled_when_section_missing(self):
+ """Verifies a None section yields a disabled config."""
+ result = config_loader.PipelineConfig._build_collapsed_categories_config(
+ raw_section=None, classes=["a"]
+ )
+ self.assertFalse(result.enable)
+ self.assertEqual(result.mapping, {})
+
+ def test_valid_full_partition_is_accepted(self):
+ """Verifies a mapping covering every class exactly once is accepted."""
+ result = config_loader.PipelineConfig._build_collapsed_categories_config(
+ raw_section={
+ "enable": True,
+ "mapping": {"g1": ["a"], "g3": ["b", "c"]},
+ },
+ classes=["a", "b", "c"],
+ )
+ self.assertTrue(result.enable)
+ self.assertEqual(result.get_category_for_class("b"), "g3")
+
+ def test_raises_when_enabled_but_mapping_empty(self):
+ """Verifies enabling the feature with an empty mapping raises."""
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig._build_collapsed_categories_config(
+ raw_section={"enable": True, "mapping": {}}, classes=["a"]
+ )
+
+ def test_raises_when_category_maps_to_empty_list(self):
+ """Verifies a category with an empty class list raises."""
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig._build_collapsed_categories_config(
+ raw_section={"enable": True, "mapping": {"g1": []}},
+ classes=["a"],
+ )
+
+ def test_raises_when_class_not_in_top_level_list(self):
+ """Verifies a mapped class absent from `classes` raises."""
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig._build_collapsed_categories_config(
+ raw_section={"enable": True, "mapping": {"g1": ["ghost"]}},
+ classes=["a"],
+ )
+
+ def test_raises_when_class_in_two_categories(self):
+ """Verifies a class assigned to two categories raises."""
+ with self.assertRaisesRegex(config_loader.ConfigurationError, "both"):
+ config_loader.PipelineConfig._build_collapsed_categories_config(
+ raw_section={
+ "enable": True,
+ "mapping": {"g1": ["a"], "g3": ["a"]},
+ },
+ classes=["a"],
+ )
+
+ def test_raises_when_class_unmapped(self):
+ """Verifies a class present in `classes` but unmapped raises."""
+ with self.assertRaisesRegex(
+ config_loader.ConfigurationError, "not assigned"
+ ):
+ config_loader.PipelineConfig._build_collapsed_categories_config(
+ raw_section={"enable": True, "mapping": {"g1": ["a"]}},
+ classes=["a", "b"],
+ )
+
+
+class FromYamlTest(absltest.TestCase):
+ """Tests for the PipelineConfig.from_yaml entry point."""
+
+ def _write_config(self, mapping: dict[str, Any]) -> str:
+ """Writes a mapping to a temp YAML file and returns its path."""
+ config_path = pathlib.Path(self.create_tempdir().full_path) / "config.yaml"
+ config_path.write_text(yaml.safe_dump(mapping), encoding="utf-8")
+ return str(config_path)
+
+ def test_loads_valid_config(self):
+ """Verifies a valid config parses into a populated PipelineConfig."""
+ config_path = self._write_config(_valid_config_mapping())
+ config = config_loader.PipelineConfig.from_yaml(config_path)
+ self.assertIsInstance(config, config_loader.PipelineConfig)
+ self.assertEqual(config.classes, ["clean_grade1", "dirty_grade3"])
+ self.assertEqual(config.preprocessing.max_short_side, 1024)
+
+ def test_coerces_tuple_fields(self):
+ """Verifies list YAML values become tuples where the dataclass expects it."""
+ config_path = self._write_config(_valid_config_mapping())
+ config = config_loader.PipelineConfig.from_yaml(config_path)
+ self.assertEqual(config.models.dinov3.image_mean, (0.485, 0.456, 0.406))
+ self.assertEqual(config.cropping.crop_size, (256, 256))
+ self.assertEqual(
+ config.visualization.background_blend_color_rgb, (124, 116, 104)
+ )
+
+ def test_missing_file_raises_configuration_error(self):
+ """Verifies a missing YAML file raises ConfigurationError."""
+ with self.assertRaisesRegex(config_loader.ConfigurationError, "not found"):
+ config_loader.PipelineConfig.from_yaml("/nonexistent/config.yaml")
+
+ def test_both_sources_enabled_raises(self):
+ """Verifies the source-XOR rule is enforced end to end."""
+ mapping = copy.deepcopy(_valid_config_mapping())
+ mapping["paths"]["gcs"]["enable"] = True
+ config_path = self._write_config(mapping)
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig.from_yaml(config_path)
+
+ def test_bigquery_enabled_missing_id_raises(self):
+ """Verifies BigQuery validation runs through from_yaml."""
+ mapping = copy.deepcopy(_valid_config_mapping())
+ mapping["bigquery"] = {
+ "enable": True,
+ "project_id": "p",
+ "dataset_id": "d",
+ "table_id": "",
+ }
+ config_path = self._write_config(mapping)
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig.from_yaml(config_path)
+
+ def test_collapsed_categories_partition_validated(self):
+ """Verifies an incomplete category partition raises through from_yaml."""
+ mapping = copy.deepcopy(_valid_config_mapping())
+ # Empty the grade3 list so dirty_grade3 is left unmapped.
+ mapping["collapsed_categories"]["mapping"]["grade3"] = []
+ config_path = self._write_config(mapping)
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig.from_yaml(config_path)
+
+ def test_returned_config_is_frozen(self):
+ """Verifies the root PipelineConfig is immutable."""
+ config_path = self._write_config(_valid_config_mapping())
+ config = config_loader.PipelineConfig.from_yaml(config_path)
+ with self.assertRaises(dataclasses.FrozenInstanceError):
+ config.classes = []
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/dinov3_classifier.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/dinov3_classifier.py
new file mode 100644
index 00000000000..b1cccc308e5
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/dinov3_classifier.py
@@ -0,0 +1,413 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""DINOv3 classification module for batched image inference.
+
+This module exposes:
+ - ClassifierError: Exception raised for classifier construction or inference
+ issues.
+ - PoolingStrategy: Enum describing how backbone token features are pooled
+ before the classification head.
+ - Prediction: Typed result returned per input image by the classifier.
+ - DINOv3ClassificationModule: The underlying nn.Module wrapping a DINOv3
+ backbone and a linear head.
+ - DINOv3Classifier: High-level batched classifier that accepts an
+ already-loaded model; construct it via `DINOv3Classifier.from_config` for
+ the common case of loading from a checkpoint on disk.
+"""
+
+from collections.abc import Mapping, Sequence
+import enum
+import logging
+import pathlib
+from typing import Self, TypedDict
+
+import numpy
+from PIL import Image
+import torch
+from torch import nn
+from torch.nn import functional
+from torchvision.transforms import v2
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import config_loader
+
+_LOGGER = logging.getLogger(__name__)
+
+
+class PoolingStrategy(enum.Enum):
+ """Strategy for pooling backbone token features into a single vector.
+
+ Attributes:
+ CLS: Use only the CLS token embedding.
+ CLS_MEAN_PATCH: Concatenate the CLS token with the mean of patch tokens.
+ """
+
+ CLS = "cls"
+ CLS_MEAN_PATCH = "cls_mean_patch"
+
+
+class ClassifierError(Exception):
+ """Exception raised for classifier loading, construction, or shape issues."""
+
+
+class Prediction(TypedDict):
+ """Typed prediction result for a single image.
+
+ Callers index by string key exactly as with a plain dict; the TypedDict adds
+ static-analysis support without changing runtime behavior.
+
+ Attributes:
+ predicted_class: Name of the top-1 class.
+ predicted_probability_percent: Top-1 probability expressed as a percentage
+ in the range [0.0, 100.0].
+ all_probabilities_percent: Mapping from class name to that class's
+ probability as a percentage in the range [0.0, 100.0].
+ """
+
+ predicted_class: str
+ predicted_probability_percent: float
+ all_probabilities_percent: dict[str, float]
+
+
+def _resolve_device(requested_device: str) -> torch.device:
+ """Resolves the requested device string to an available torch.device.
+
+ Falls back to CPU with a warning if the requested device is not available.
+
+ Args:
+ requested_device: Device string requested by the caller.
+
+ Returns:
+ The resolved torch.device.
+ """
+ if requested_device.startswith("cuda") and not torch.cuda.is_available():
+ _LOGGER.warning(
+ "Requested device '%s' but CUDA is not available; falling back to CPU.",
+ requested_device,
+ )
+ return torch.device("cpu")
+
+ if requested_device == "mps" and not torch.backends.mps.is_available():
+ _LOGGER.warning(
+ "Requested device 'mps' but MPS is not available; falling back to CPU."
+ )
+ return torch.device("cpu")
+
+ return torch.device(requested_device)
+
+
+def _infer_pooling_from_state_dict(
+ saved_state_dict: Mapping[str, torch.Tensor], hidden_size: int
+) -> PoolingStrategy:
+ """Infers the pooling strategy from the shape of the saved head weights.
+
+ Args:
+ saved_state_dict: The checkpoint's model state dict.
+ hidden_size: Backbone hidden dimensionality.
+
+ Returns:
+ The inferred pooling strategy.
+
+ Raises:
+ ClassifierError: If the head input dimension matches neither the CLS nor
+ the CLS_MEAN_PATCH expectation.
+ """
+ head_input_features = saved_state_dict["head.weight"].shape[1]
+ if head_input_features == hidden_size:
+ return PoolingStrategy.CLS
+ if head_input_features == 2 * hidden_size:
+ return PoolingStrategy.CLS_MEAN_PATCH
+ raise ClassifierError(
+ "Cannot infer pooling strategy. Head input dimension "
+ f"{head_input_features} does not match hidden size {hidden_size} or "
+ f"{2 * hidden_size}."
+ )
+
+
+def _load_checkpoint_state_dict(
+ checkpoint_path: pathlib.Path | str, device: torch.device
+) -> dict[str, torch.Tensor]:
+ """Loads the model state dict from a checkpoint file.
+
+ Args:
+ checkpoint_path: Filesystem path to the checkpoint.
+ device: Target device for `map_location`.
+
+ Returns:
+ The `model_state_dict` mapping from parameter name to tensor.
+
+ Raises:
+ ClassifierError: If the checkpoint is missing required keys.
+ """
+ checkpoint = torch.load(
+ checkpoint_path, map_location=device, weights_only=True
+ )
+ if "model_state_dict" not in checkpoint:
+ raise ClassifierError(
+ "Checkpoint is missing required key 'model_state_dict'."
+ )
+ saved_state_dict = checkpoint["model_state_dict"]
+ if "head.weight" not in saved_state_dict:
+ raise ClassifierError(
+ "Checkpoint state dict is missing required key 'head.weight'."
+ )
+ return saved_state_dict
+
+
+def _build_image_transform(config: config_loader.DINOv3Config) -> v2.Compose:
+ """Builds the preprocessing pipeline used to feed PIL images to the model.
+
+ Args:
+ config: DINOv3 model configuration providing image size and normalization
+ statistics.
+
+ Returns:
+ A torchvision v2.Compose pipeline.
+ """
+ return v2.Compose([
+ v2.ToImage(),
+ v2.Resize(
+ (config.inference_image_size, config.inference_image_size),
+ antialias=True,
+ ),
+ v2.ToDtype(torch.float32, scale=True),
+ v2.Normalize(mean=config.image_mean, std=config.image_std),
+ ])
+
+
+class DINOv3ClassificationModule(nn.Module):
+ """DINOv3 backbone paired with a linear classification head.
+
+ Attributes:
+ pooling: The pooling strategy applied to backbone features.
+ backbone_model: The DINOv3 backbone.
+ head: Linear layer mapping pooled features to per-class logits.
+ """
+
+ def __init__(
+ self,
+ backbone_model: nn.Module,
+ hidden_size: int,
+ number_of_classes: int,
+ pooling: PoolingStrategy,
+ ):
+ """Initializes the module from an already-loaded backbone.
+
+ Args:
+ backbone_model: The DINOv3 backbone module. Injected so that this class
+ does not perform filesystem I/O in its constructor.
+ hidden_size: Dimensionality of the backbone's token embeddings.
+ number_of_classes: Number of output classes for the head.
+ pooling: Pooling strategy that determines the head's input dimension.
+ """
+ super().__init__()
+ self.pooling = pooling
+ self.backbone_model = backbone_model
+
+ if pooling is PoolingStrategy.CLS:
+ head_input_features = hidden_size
+ else:
+ head_input_features = 2 * hidden_size
+
+ self.head = nn.Linear(
+ in_features=head_input_features,
+ out_features=number_of_classes,
+ bias=True,
+ )
+
+ def extract_features(self, image_batch: torch.Tensor) -> torch.Tensor:
+ """Extracts pooled features from a batch of images.
+
+ Args:
+ image_batch: Tensor of shape (batch_size, 3, height, width).
+
+ Returns:
+ Tensor of pooled features. Shape is (batch_size, hidden_size) for the
+ CLS strategy and (batch_size, 2 * hidden_size) for CLS_MEAN_PATCH.
+ """
+ if self.pooling is PoolingStrategy.CLS:
+ return self.backbone_model(image_batch)
+
+ token_features = self.backbone_model.forward_features(image_batch)
+ cls_token = token_features["x_norm_clstoken"]
+ mean_patch_token = token_features["x_norm_patchtokens"].mean(dim=1)
+ return torch.cat([cls_token, mean_patch_token], dim=1)
+
+ def forward(self, image_batch: torch.Tensor) -> torch.Tensor:
+ """Runs the full backbone + head forward pass.
+
+ Args:
+ image_batch: Tensor of shape (batch_size, 3, height, width).
+
+ Returns:
+ Per-class logits of shape (batch_size, number_of_classes).
+ """
+ return self.head(self.extract_features(image_batch))
+
+
+class DINOv3Classifier:
+ """High-level batched classifier backed by a DINOv3 model.
+
+ The order of `class_names` is load-bearing: `class_names[i]` must correspond
+ to logit index `i` in the trained model. Reordering this list will silently
+ produce incorrect predictions.
+
+ Construct instances via `DINOv3Classifier.from_config` for the common case of
+ loading from a checkpoint. The `__init__` constructor accepts an
+ already-built model and is intended for callers that manage model loading
+ themselves (for example, unit tests).
+
+ Attributes:
+ class_names: Ordered list of class names indexed by logit position.
+ """
+
+ def __init__(
+ self,
+ model: DINOv3ClassificationModule,
+ class_names: Sequence[str],
+ image_transform: v2.Compose,
+ device: torch.device,
+ ):
+ """Initializes the classifier from injected dependencies.
+
+ Args:
+ model: The already-loaded classification module, expected to be on
+ `device` and in eval mode. This constructor does not move the model or
+ change its mode.
+ class_names: Ordered list of class names indexed by logit position.
+ image_transform: Preprocessing pipeline applied to each PIL image before
+ stacking into a batch.
+ device: Target device on which the model resides and to which input
+ batches will be moved.
+ """
+ self.class_names = list(class_names)
+ self._model = model
+ self._image_transform = image_transform
+ self._device = device
+
+ @classmethod
+ def from_config(
+ cls,
+ config: config_loader.DINOv3Config,
+ class_names: Sequence[str],
+ device: str,
+ ) -> Self:
+ """Builds a classifier by loading a checkpoint from disk.
+
+ Args:
+ config: DINOv3 model configuration.
+ class_names: Ordered list of class names indexed by logit position.
+ device: Requested device string (e.g., 'cuda', 'cpu', 'mps').
+
+ Returns:
+ A ready-to-use DINOv3Classifier with its model on the resolved device
+ and in eval mode.
+
+ Raises:
+ ClassifierError: If the checkpoint is missing required keys, if the
+ pooling strategy cannot be inferred, or if the state dict fails to
+ load into the constructed model.
+ """
+ resolved_device = _resolve_device(device)
+ saved_state_dict = _load_checkpoint_state_dict(
+ checkpoint_path=config.checkpoint_path, device=resolved_device
+ )
+
+ backbone_model = torch.hub.load(
+ config.repo_dir,
+ config.model_name,
+ source="local",
+ pretrained=False,
+ )
+ # DINOv3 Vision Transformers store token embedding dimension in
+ # norm.normalized_shape[0].
+ hidden_size = backbone_model.norm.normalized_shape[0]
+
+ pooling = _infer_pooling_from_state_dict(
+ saved_state_dict=saved_state_dict, hidden_size=hidden_size
+ )
+
+ model = DINOv3ClassificationModule(
+ backbone_model=backbone_model,
+ hidden_size=hidden_size,
+ number_of_classes=len(class_names),
+ pooling=pooling,
+ ).to(resolved_device)
+
+ try:
+ model.load_state_dict(saved_state_dict)
+ except RuntimeError as error:
+ raise ClassifierError(
+ f"Failed to load state dict into model: {error}"
+ ) from error
+
+ model.eval()
+ return cls(
+ model=model,
+ class_names=class_names,
+ image_transform=_build_image_transform(config),
+ device=resolved_device,
+ )
+
+ @torch.no_grad()
+ def predict_batch(self, images: Sequence[Image.Image]) -> list[Prediction]:
+ """Classifies a batch of PIL images in a single forward pass.
+
+ Args:
+ images: Sequence of PIL images to classify.
+
+ Returns:
+ List of Prediction dicts in the same order as `images`. Each Prediction
+ contains 'predicted_class', 'predicted_probability_percent', and
+ 'all_probabilities_percent'. Returns an empty list when `images` is
+ empty.
+ """
+ if not images:
+ return []
+
+ image_tensors = [self._image_transform(image) for image in images]
+ batch_tensor = torch.stack(image_tensors).to(self._device)
+
+ logits = self._model(batch_tensor)
+ probabilities = functional.softmax(logits, dim=1).cpu().numpy()
+
+ predictions: list[Prediction] = []
+ for probability_row in probabilities:
+ predicted_index = int(numpy.argmax(probability_row))
+ all_probabilities_percent: dict[str, float] = {}
+ for class_name, probability in zip(self.class_names, probability_row):
+ all_probabilities_percent[class_name] = float(probability * 100.0)
+ prediction: Prediction = {
+ "predicted_class": self.class_names[predicted_index],
+ "predicted_probability_percent": float(
+ probability_row[predicted_index] * 100.0
+ ),
+ "all_probabilities_percent": all_probabilities_percent,
+ }
+ predictions.append(prediction)
+ return predictions
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/dinov3_classifier_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/dinov3_classifier_test.py
new file mode 100644
index 00000000000..08e0ce16871
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/dinov3_classifier_test.py
@@ -0,0 +1,365 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for dinov3_classifier.py.
+
+The module exposes small pure helpers (device resolution, pooling inference,
+checkpoint reading) that are tested directly, plus two constructors on
+`DINOv3Classifier`: the DI-friendly `__init__` used by these tests, and the
+`from_config` factory whose I/O (`torch.hub.load`, `torch.load`) is patched
+so no real model or checkpoint is loaded.
+"""
+
+from unittest import mock
+
+from absl.testing import absltest
+import torch
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import dinov3_classifier
+
+
+def _make_probe_backbone(hidden_size: int) -> mock.Mock:
+ """Returns a fake backbone that exposes `norm.normalized_shape`."""
+ backbone = mock.Mock()
+ backbone.norm.normalized_shape = (hidden_size,)
+ return backbone
+
+
+def _make_dinov3_config() -> mock.Mock:
+ """Returns a stub DINOv3Config with only the fields the module reads."""
+ config = mock.Mock()
+ config.checkpoint_path = "/tmp/ckpt.pth"
+ config.repo_dir = "/tmp/dinov3"
+ config.model_name = "dinov3_vitl16"
+ config.inference_image_size = 32
+ config.image_mean = (0.485, 0.456, 0.406)
+ config.image_std = (0.229, 0.224, 0.225)
+ return config
+
+
+class ResolveDeviceTest(absltest.TestCase):
+ """Tests for _resolve_device."""
+
+ def test_returns_cpu_when_requested(self):
+ """Verifies an explicit CPU request always resolves to CPU."""
+ self.assertEqual(
+ dinov3_classifier._resolve_device("cpu"), torch.device("cpu")
+ )
+
+ def test_falls_back_to_cpu_when_cuda_unavailable(self):
+ """Verifies a CUDA request without CUDA falls back to CPU."""
+ with mock.patch.object(
+ dinov3_classifier.torch.cuda, "is_available", return_value=False
+ ):
+ self.assertEqual(
+ dinov3_classifier._resolve_device("cuda"), torch.device("cpu")
+ )
+
+ def test_returns_cuda_when_available(self):
+ """Verifies a CUDA request resolves to CUDA when available."""
+ with mock.patch.object(
+ dinov3_classifier.torch.cuda, "is_available", return_value=True
+ ):
+ self.assertEqual(
+ dinov3_classifier._resolve_device("cuda"), torch.device("cuda")
+ )
+
+
+class InferPoolingFromStateDictTest(absltest.TestCase):
+ """Tests for _infer_pooling_from_state_dict."""
+
+ def test_infers_cls_when_head_matches_hidden(self):
+ """Verifies a head width equal to hidden size yields CLS pooling."""
+ state_dict = {"head.weight": torch.zeros(3, 32)}
+ result = dinov3_classifier._infer_pooling_from_state_dict(
+ saved_state_dict=state_dict, hidden_size=32
+ )
+ self.assertIs(result, dinov3_classifier.PoolingStrategy.CLS)
+
+ def test_infers_cls_mean_patch_when_head_is_double_hidden(self):
+ """Verifies a head width of 2*hidden yields CLS_MEAN_PATCH pooling."""
+ state_dict = {"head.weight": torch.zeros(3, 64)}
+ result = dinov3_classifier._infer_pooling_from_state_dict(
+ saved_state_dict=state_dict, hidden_size=32
+ )
+ self.assertIs(result, dinov3_classifier.PoolingStrategy.CLS_MEAN_PATCH)
+
+ def test_raises_when_head_matches_neither(self):
+ """Verifies a mismatched head width raises ClassifierError."""
+ state_dict = {"head.weight": torch.zeros(3, 99)}
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError, "Cannot infer"
+ ):
+ dinov3_classifier._infer_pooling_from_state_dict(
+ saved_state_dict=state_dict, hidden_size=32
+ )
+
+
+class LoadCheckpointStateDictTest(absltest.TestCase):
+ """Tests for _load_checkpoint_state_dict."""
+
+ def test_returns_state_dict_when_valid(self):
+ """Verifies a well-formed checkpoint yields its model_state_dict."""
+ state_dict = {"head.weight": torch.zeros(2, 32)}
+ with mock.patch.object(
+ dinov3_classifier.torch,
+ "load",
+ return_value={"model_state_dict": state_dict},
+ ):
+ result = dinov3_classifier._load_checkpoint_state_dict(
+ checkpoint_path="/tmp/ckpt.pth", device=torch.device("cpu")
+ )
+ self.assertIs(result, state_dict)
+
+ def test_raises_when_model_state_dict_missing(self):
+ """Verifies a checkpoint without model_state_dict raises."""
+ with mock.patch.object(dinov3_classifier.torch, "load", return_value={}):
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError, "model_state_dict"
+ ):
+ dinov3_classifier._load_checkpoint_state_dict(
+ checkpoint_path="/tmp/ckpt.pth", device=torch.device("cpu")
+ )
+
+ def test_raises_when_head_weight_missing(self):
+ """Verifies a checkpoint without head.weight raises."""
+ with mock.patch.object(
+ dinov3_classifier.torch,
+ "load",
+ return_value={"model_state_dict": {"other.weight": torch.zeros(2, 2)}},
+ ):
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError, "head.weight"
+ ):
+ dinov3_classifier._load_checkpoint_state_dict(
+ checkpoint_path="/tmp/ckpt.pth", device=torch.device("cpu")
+ )
+
+
+class DINOv3ClassificationModuleTest(absltest.TestCase):
+ """Tests for DINOv3ClassificationModule head sizing and pooling dispatch."""
+
+ def test_head_input_dimension_for_cls_pooling(self):
+ """Verifies CLS pooling sizes the head to the hidden dimension."""
+ module = dinov3_classifier.DINOv3ClassificationModule(
+ backbone_model=_make_probe_backbone(32),
+ hidden_size=32,
+ number_of_classes=5,
+ pooling=dinov3_classifier.PoolingStrategy.CLS,
+ )
+ self.assertEqual(module.head.in_features, 32)
+ self.assertEqual(module.head.out_features, 5)
+
+ def test_head_input_dimension_for_cls_mean_patch_pooling(self):
+ """Verifies CLS_MEAN_PATCH sizing doubles the head input dimension."""
+ module = dinov3_classifier.DINOv3ClassificationModule(
+ backbone_model=_make_probe_backbone(32),
+ hidden_size=32,
+ number_of_classes=5,
+ pooling=dinov3_classifier.PoolingStrategy.CLS_MEAN_PATCH,
+ )
+ self.assertEqual(module.head.in_features, 64)
+
+ def test_extract_features_cls_calls_backbone_directly(self):
+ """Verifies CLS pooling passes the input straight through the backbone."""
+ backbone = _make_probe_backbone(32)
+ features = torch.ones((2, 32))
+ backbone.return_value = features
+ module = dinov3_classifier.DINOv3ClassificationModule(
+ backbone_model=backbone,
+ hidden_size=32,
+ number_of_classes=3,
+ pooling=dinov3_classifier.PoolingStrategy.CLS,
+ )
+ result = module.extract_features(torch.zeros((2, 3, 32, 32)))
+ self.assertEqual(result.shape, (2, 32))
+
+ def test_extract_features_cls_mean_patch_concatenates_tokens(self):
+ """Verifies CLS_MEAN_PATCH concatenates CLS with the mean patch token."""
+ backbone = _make_probe_backbone(32)
+ backbone.forward_features.return_value = {
+ "x_norm_clstoken": torch.ones((2, 32)),
+ "x_norm_patchtokens": torch.ones((2, 8, 32)),
+ }
+ module = dinov3_classifier.DINOv3ClassificationModule(
+ backbone_model=backbone,
+ hidden_size=32,
+ number_of_classes=3,
+ pooling=dinov3_classifier.PoolingStrategy.CLS_MEAN_PATCH,
+ )
+ result = module.extract_features(torch.zeros((2, 3, 32, 32)))
+ self.assertEqual(result.shape, (2, 64))
+
+
+class PredictBatchTest(absltest.TestCase):
+ """Tests for DINOv3Classifier.predict_batch post-processing."""
+
+ def _make_classifier(
+ self, class_names: list[str], logits: torch.Tensor
+ ) -> dinov3_classifier.DINOv3Classifier:
+ """Builds a classifier with a stubbed transform and model."""
+ model = mock.Mock(return_value=logits)
+ # Transform ignores the image and returns a fixed CHW tensor.
+ image_transform = lambda image: torch.zeros(3, 4, 4)
+ return dinov3_classifier.DINOv3Classifier(
+ model=model,
+ class_names=class_names,
+ image_transform=image_transform,
+ device=torch.device("cpu"),
+ )
+
+ def test_empty_images_returns_empty_list(self):
+ """Verifies an empty image list short-circuits to an empty result."""
+ classifier = self._make_classifier(["a"], torch.zeros(1, 1))
+ self.assertEqual(classifier.predict_batch([]), [])
+
+ def test_returns_argmax_class_per_row(self):
+ """Verifies each row's predicted class is the argmax of its logits."""
+ logits = torch.tensor([[10.0, 0.0, 0.0], [0.0, 0.0, 10.0]])
+ classifier = self._make_classifier(["a", "b", "c"], logits)
+ result = classifier.predict_batch([object(), object()])
+ self.assertEqual(result[0]["predicted_class"], "a")
+ self.assertEqual(result[1]["predicted_class"], "c")
+
+ def test_probabilities_are_percentages_summing_to_one_hundred(self):
+ """Verifies probabilities are percentages that cover every class."""
+ logits = torch.tensor([[10.0, 0.0, 0.0]])
+ classifier = self._make_classifier(["a", "b", "c"], logits)
+ prediction = classifier.predict_batch([object()])[0]
+ self.assertGreater(prediction["predicted_probability_percent"], 90.0)
+ self.assertCountEqual(
+ prediction["all_probabilities_percent"].keys(), ["a", "b", "c"]
+ )
+ total = sum(prediction["all_probabilities_percent"].values())
+ self.assertAlmostEqual(total, 100.0, places=3)
+
+
+class FromConfigTest(absltest.TestCase):
+ """Tests for DINOv3Classifier.from_config end-to-end wiring."""
+
+ def _patched_environment(
+ self, head_width: int, hidden_size: int = 32
+ ) -> mock.Mock:
+ """Patches torch.hub.load, torch.load, and load_state_dict as a group.
+
+ Args:
+ head_width: The width of the classification head.
+ hidden_size: The hidden size of the backbone.
+
+ Returns:
+ The mock for torch.hub.load so tests can assert its call count.
+ """
+ fake_backbone = _make_probe_backbone(hidden_size)
+ mock_hub_load = self.enter_context(
+ mock.patch.object(
+ dinov3_classifier.torch.hub,
+ "load",
+ return_value=fake_backbone,
+ )
+ )
+ self.enter_context(
+ mock.patch.object(
+ dinov3_classifier.torch,
+ "load",
+ return_value={
+ "model_state_dict": {
+ "head.weight": torch.zeros(2, head_width),
+ }
+ },
+ )
+ )
+ # Bypass the real state-dict load, which would complain about the
+ # fake backbone's missing parameters.
+ self.enter_context(
+ mock.patch.object(
+ dinov3_classifier.DINOv3ClassificationModule,
+ "load_state_dict",
+ autospec=True,
+ )
+ )
+ return mock_hub_load
+
+ def test_loads_backbone_exactly_once(self):
+ """Verifies from_config invokes torch.hub.load a single time."""
+ mock_hub_load = self._patched_environment(head_width=32)
+ dinov3_classifier.DINOv3Classifier.from_config(
+ config=_make_dinov3_config(),
+ class_names=["a", "b"],
+ device="cpu",
+ )
+ self.assertEqual(mock_hub_load.call_count, 1)
+
+ def test_returns_classifier_in_eval_mode(self):
+ """Verifies the returned classifier's model is in eval mode."""
+ self._patched_environment(head_width=32)
+ classifier = dinov3_classifier.DINOv3Classifier.from_config(
+ config=_make_dinov3_config(),
+ class_names=["a", "b"],
+ device="cpu",
+ )
+ self.assertFalse(classifier._model.training)
+
+ def test_selects_pooling_strategy_from_checkpoint(self):
+ """Verifies pooling is inferred from the head width in the checkpoint."""
+ self._patched_environment(head_width=64, hidden_size=32)
+ classifier = dinov3_classifier.DINOv3Classifier.from_config(
+ config=_make_dinov3_config(),
+ class_names=["a", "b"],
+ device="cpu",
+ )
+ self.assertIs(
+ classifier._model.pooling,
+ dinov3_classifier.PoolingStrategy.CLS_MEAN_PATCH,
+ )
+
+ def test_wraps_state_dict_load_error(self):
+ """Verifies a RuntimeError from load_state_dict is re-raised as ClassifierError."""
+ fake_backbone = _make_probe_backbone(32)
+ with mock.patch.object(
+ dinov3_classifier.torch.hub, "load", return_value=fake_backbone
+ ), mock.patch.object(
+ dinov3_classifier.torch,
+ "load",
+ return_value={"model_state_dict": {"head.weight": torch.zeros(2, 32)}},
+ ), mock.patch.object(
+ dinov3_classifier.DINOv3ClassificationModule,
+ "load_state_dict",
+ side_effect=RuntimeError("shape mismatch"),
+ ):
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError, "state dict"
+ ):
+ dinov3_classifier.DINOv3Classifier.from_config(
+ config=_make_dinov3_config(),
+ class_names=["a", "b"],
+ device="cpu",
+ )
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/gcs_ops.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/gcs_ops.py
new file mode 100644
index 00000000000..db43b2c7ad9
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/gcs_ops.py
@@ -0,0 +1,486 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Google Cloud Platform operations for the pipeline.
+
+Provides Cloud Storage transfers and BigQuery result ingestion.
+"""
+
+from collections.abc import Mapping
+import logging
+import os
+import shutil
+from typing import Any
+
+from google.cloud import bigquery
+from google.cloud import exceptions
+from google.cloud import storage
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import config_loader
+
+_LOGGER = logging.getLogger(__name__)
+
+_GCS_URI_PREFIX = "gs://"
+
+# Column order for the per-track table. This is also the column order for the
+# BigQuery table, so it is kept in one place.
+TRACK_TABLE_COLUMNS = (
+ "session_name",
+ "tracker_id",
+ "final_class",
+ "vote_count",
+ "collapsed_class",
+ "track_grid_filename",
+)
+
+_TRACK_GRID_NOT_AVAILABLE = "N/A"
+
+# BigQuery table schema as (column_name, field_type) pairs. The column names
+# and their order match TRACK_TABLE_COLUMNS. All fields are NULLABLE (the
+# default mode), so empty values such as a missing collapsed class ingest as
+# NULL. Kept as plain tuples so importing this module does not require the
+# BigQuery client; it is converted to SchemaField objects at ingestion time.
+_BIGQUERY_SCHEMA = (
+ ("session_name", "STRING"),
+ ("tracker_id", "INTEGER"),
+ ("final_class", "STRING"),
+ ("vote_count", "INTEGER"),
+ ("collapsed_class", "STRING"),
+ ("track_grid_filename", "STRING"),
+)
+
+
+def _parse_gcs_uri(gcs_uri: str) -> tuple[str, str]:
+ """Splits a gs:// URI into its bucket name and object prefix.
+
+ Args:
+ gcs_uri: A URI of the form 'gs://bucket-name/optional/prefix/'.
+
+ Returns:
+ A tuple of (bucket_name, object_prefix). The object prefix is an
+ empty string when the URI points at the bucket root.
+
+ Raises:
+ ValueError: If the URI does not start with 'gs://' or has no bucket
+ name.
+ """
+ if not gcs_uri.startswith(_GCS_URI_PREFIX):
+ raise ValueError(f"GCS URI must start with '{_GCS_URI_PREFIX}': {gcs_uri}")
+ uri_without_scheme = gcs_uri[len(_GCS_URI_PREFIX) :]
+ bucket_name, _, object_prefix = uri_without_scheme.partition("/")
+ if not bucket_name:
+ raise ValueError(f"GCS URI has no bucket name: {gcs_uri}")
+ return bucket_name, object_prefix
+
+
+def resolve_io_directories(
+ paths: config_loader.PathsConfig,
+) -> tuple[str, str]:
+ """Resolves the working input/output directories for the active source.
+
+ In local mode the configured local directories are returned unchanged.
+ In GCS mode the input is downloaded from the configured input URI into
+ the temp input directory, and the temp input/output directories are
+ returned as the working directories. The local branch does not construct
+ a Cloud Storage client, so no GCP credentials are needed when running
+ fully local.
+
+ Args:
+ paths: The validated paths configuration. Exactly one of its local or GCS
+ sources is enabled.
+
+ Returns:
+ A tuple of (input_root, output_root) that the pipeline reads from
+ and writes to.
+ """
+ if paths.gcs.enable:
+ _LOGGER.info("GCS mode: downloading input from %s ...", paths.gcs.input_uri)
+ storage_manager = CloudStorageManager()
+ storage_manager.download_directory(
+ gcs_uri=paths.gcs.input_uri,
+ local_directory=paths.gcs.temp_input_directory,
+ )
+ return (
+ paths.gcs.temp_input_directory,
+ paths.gcs.temp_output_directory,
+ )
+ return (
+ paths.local.input_image_directory,
+ paths.local.output_root_directory,
+ )
+
+
+def upload_output_directory(paths: config_loader.PathsConfig) -> None:
+ """Uploads pipeline output to GCS when GCS mode is active.
+
+ In local mode this is a no-op, so callers can invoke it unconditionally.
+ In GCS mode the entire temp output directory is uploaded to the
+ configured output URI, preserving folder structure.
+
+ Args:
+ paths: The validated paths configuration. Exactly one of its local or GCS
+ sources is enabled.
+ """
+ if not paths.gcs.enable:
+ return
+
+ _LOGGER.info("GCS mode: uploading output to %s ...", paths.gcs.output_uri)
+ storage_manager = CloudStorageManager()
+ storage_manager.upload_directory(
+ local_directory=paths.gcs.temp_output_directory,
+ gcs_uri=paths.gcs.output_uri,
+ )
+
+
+def build_track_rows(
+ track_summary: Mapping[Any, Any], session_name: str
+) -> list[dict[str, Any]]:
+ """Builds one table row per track from a subfolder's track summary.
+
+ The row shape is the BigQuery payload. `final_class` is the DINOv3 grade
+ resolved for the track, `collapsed_class` is its optional category (None
+ when collapsed categories are disabled), and `track_grid_filename` is the
+ saved grid PNG name ('N/A' when grids are disabled).
+
+ Args:
+ track_summary: Mapping from tracker_id to its resolved summary dict, as
+ returned by PipelineVisualizer.save_track_grids. Each value contains
+ 'final_class', 'category', 'vote_count', and 'output_path'.
+ session_name: Name of the capture session (the subfolder name).
+
+ Returns:
+ A list of row dicts keyed by TRACK_TABLE_COLUMNS, ordered by
+ tracker_id.
+ """
+ rows = []
+ for tracker_id in sorted(track_summary):
+ summary = track_summary[tracker_id]
+ rows.append({
+ "session_name": session_name,
+ "tracker_id": int(tracker_id),
+ "final_class": summary["final_class"],
+ "vote_count": int(summary["vote_count"]),
+ "collapsed_class": summary["category"],
+ "track_grid_filename": _track_grid_filename(summary["output_path"]),
+ })
+ return rows
+
+
+def print_track_rows(rows: list[dict[str, Any]]) -> None:
+ """Prints track rows as a simple table for reviewing the schema.
+
+ Temporary scaffolding used to review the BigQuery schema.
+
+ Args:
+ rows: Row dicts as produced by build_track_rows.
+ """
+ print(f"Schema: {' | '.join(TRACK_TABLE_COLUMNS)}")
+ for row in rows:
+ values = [str(row[column]) for column in TRACK_TABLE_COLUMNS]
+ print(" | ".join(values))
+
+
+def _track_grid_filename(output_path: str) -> str:
+ """Returns the track-grid PNG filename, or 'N/A' when grids are disabled.
+
+ When track grids are disabled, PipelineVisualizer.save_track_grids stores
+ a placeholder string rather than a real path. That case is reported as
+ 'N/A'; otherwise the basename of the saved PNG is returned.
+
+ Args:
+ output_path: The 'output_path' field from a track summary entry.
+
+ Returns:
+ The PNG filename (e.g. 'track_0011_dirt_jars_grade3.png') or 'N/A'.
+ """
+ if not output_path or output_path.startswith(_TRACK_GRID_NOT_AVAILABLE):
+ return _TRACK_GRID_NOT_AVAILABLE
+ return os.path.basename(output_path)
+
+
+def ingest_track_rows(
+ bigquery_config: config_loader.BigQueryConfig,
+ rows: list[dict[str, Any]],
+) -> None:
+ """Writes track rows to BigQuery when BigQuery ingestion is enabled.
+
+ This runs independently of the input/output source, so results can be
+ written whether the pipeline ran in local or GCS mode. It is a no-op when
+ BigQuery is disabled or when there are no rows, so callers can invoke it
+ unconditionally.
+
+ Args:
+ bigquery_config: Validated BigQuery configuration.
+ rows: Track rows as produced by build_track_rows, accumulated across all
+ processed subfolders.
+ """
+ if not bigquery_config.enable:
+ return
+ if not rows:
+ _LOGGER.info("BigQuery enabled but there are no rows to ingest.")
+ return
+
+ _LOGGER.info("Ingesting %d row(s) into BigQuery ...", len(rows))
+ manager = BigQueryManager(
+ project_id=bigquery_config.project_id,
+ dataset_id=bigquery_config.dataset_id,
+ table_id=bigquery_config.table_id,
+ )
+ manager.ingest_rows(rows, overwrite=bigquery_config.overwrite)
+
+
+class CloudStorageManager:
+ """Handles Cloud Storage transfers for pipeline input and output."""
+
+ def __init__(self, project_id: str | None = None):
+ """Initializes the Cloud Storage client.
+
+ Args:
+ project_id: Google Cloud project ID. When None, the client uses the
+ project inferred from the environment credentials.
+ """
+ self._client = storage.Client(project=project_id)
+
+ def download_directory(self, gcs_uri: str, local_directory: str) -> int:
+ """Downloads all objects under a GCS prefix into a local directory.
+
+ The local directory is cleared before downloading so that stale data
+ from a previous run does not remain. The last segment of the prefix
+ (the named folder) is preserved locally, so the folder itself is
+ recreated rather than having its contents flattened.
+
+ Args:
+ gcs_uri: Source URI of the form 'gs://bucket/prefix/'.
+ local_directory: Local destination directory. Created if missing,
+ emptied if it already exists.
+
+ Returns:
+ The number of files downloaded.
+
+ Raises:
+ ValueError: If the GCS URI is malformed.
+ """
+ bucket_name, object_prefix = _parse_gcs_uri(gcs_uri)
+ self._reset_local_directory(local_directory)
+
+ bucket = self._client.bucket(bucket_name)
+ blobs = list(self._client.list_blobs(bucket, prefix=object_prefix))
+
+ downloaded_count = 0
+ for blob in blobs:
+ # Skip directory placeholder objects whose name ends with a slash.
+ if blob.name.endswith("/"):
+ continue
+ relative_path = self._relative_object_path(
+ blob_name=blob.name, object_prefix=object_prefix
+ )
+ destination_path = os.path.join(local_directory, relative_path)
+ os.makedirs(os.path.dirname(destination_path), exist_ok=True)
+ blob.download_to_filename(destination_path)
+ downloaded_count += 1
+
+ _LOGGER.info(
+ "Downloaded %d file(s) from %s to %s",
+ downloaded_count,
+ gcs_uri,
+ local_directory,
+ )
+ return downloaded_count
+
+ def upload_directory(self, local_directory: str, gcs_uri: str) -> int:
+ """Uploads every file under a local directory to a GCS prefix.
+
+ The local directory tree is walked recursively and each file is
+ uploaded to the destination prefix, preserving the relative folder
+ structure. Existing objects at the destination are overwritten.
+
+ Args:
+ local_directory: Local source directory to upload.
+ gcs_uri: Destination URI of the form 'gs://bucket/prefix/'.
+
+ Returns:
+ The number of files uploaded.
+
+ Raises:
+ ValueError: If the GCS URI is malformed.
+ """
+ bucket_name, object_prefix = _parse_gcs_uri(gcs_uri)
+ bucket = self._client.bucket(bucket_name)
+
+ uploaded_count = 0
+ for current_directory, _, file_names in os.walk(local_directory):
+ for file_name in file_names:
+ local_path = os.path.join(current_directory, file_name)
+ relative_path = os.path.relpath(local_path, local_directory)
+ blob_name = self._join_gcs_path(object_prefix, relative_path)
+ bucket.blob(blob_name).upload_from_filename(local_path)
+ uploaded_count += 1
+
+ _LOGGER.info(
+ "Uploaded %d file(s) from %s to %s",
+ uploaded_count,
+ local_directory,
+ gcs_uri,
+ )
+ return uploaded_count
+
+ def _reset_local_directory(self, local_directory: str) -> None:
+ """Removes and recreates a local directory so it starts empty.
+
+ Args:
+ local_directory: Directory to reset.
+ """
+ if os.path.isdir(local_directory):
+ shutil.rmtree(local_directory)
+ os.makedirs(local_directory, exist_ok=True)
+
+ def _relative_object_path(self, blob_name: str, object_prefix: str) -> str:
+ """Returns a blob's path relative to the parent of the download prefix.
+
+ The last segment of the prefix (the named folder, e.g. 'images') is
+ preserved so that the folder itself is recreated locally rather than
+ having its contents flattened into the destination directory. For a
+ prefix 'a/b/images/' and blob 'a/b/images/img1.png', this returns
+ 'images/img1.png'.
+
+ Args:
+ blob_name: Full object name within the bucket.
+ object_prefix: The prefix the download was rooted at.
+
+ Returns:
+ The portion of the blob name below the prefix's parent, suitable
+ for joining onto the local destination directory.
+ """
+ parent_prefix = self._parent_prefix(object_prefix)
+ relative_path = blob_name[len(parent_prefix) :]
+ return relative_path.lstrip("/")
+
+ def _parent_prefix(self, object_prefix: str) -> str:
+ """Returns the parent portion of a prefix, keeping its last folder.
+
+ For 'a/b/images/' this returns 'a/b/'. For a single-segment prefix
+ like 'images/' this returns '' so the folder is kept at the top of
+ the destination. An empty prefix (bucket root) returns ''.
+
+ Args:
+ object_prefix: The prefix the download was rooted at.
+
+ Returns:
+ The prefix with its last folder segment removed.
+ """
+ trimmed_prefix = object_prefix.rstrip("/")
+ if "/" not in trimmed_prefix:
+ return ""
+ parent, _, _ = trimmed_prefix.rpartition("/")
+ return parent + "/"
+
+ def _join_gcs_path(self, object_prefix: str, relative_path: str) -> str:
+ """Joins a GCS prefix and a local relative path into a blob name.
+
+ GCS object names always use forward slashes, so any OS-specific
+ separators in the relative path are normalized.
+
+ Args:
+ object_prefix: Destination prefix (may be empty for bucket root).
+ relative_path: File path relative to the upload source directory.
+
+ Returns:
+ The full blob name under which the file should be stored.
+ """
+ normalized_relative = relative_path.replace(os.sep, "/")
+ clean_prefix = object_prefix.strip("/")
+ if not clean_prefix:
+ return normalized_relative
+ return clean_prefix + "/" + normalized_relative
+
+
+class BigQueryManager:
+ """Handles BigQuery dataset/table setup and result ingestion.
+
+ The BigQuery client is imported lazily so that importing this module does
+ not require google-cloud-bigquery to be installed when BigQuery is not
+ used.
+ """
+
+ def __init__(self, project_id: str, dataset_id: str, table_id: str):
+ """Initializes the BigQuery client and resolves the table reference.
+
+ Args:
+ project_id: Google Cloud project ID that owns the dataset.
+ dataset_id: BigQuery dataset ID that holds the table.
+ table_id: BigQuery table ID that receives the rows.
+ """
+ self._bigquery = bigquery
+ self._client = bigquery.Client(project=project_id)
+ self._project_id = project_id
+ self._dataset_id = dataset_id
+ self._table_id = table_id
+ self._table_reference = f"{project_id}.{dataset_id}.{table_id}"
+
+ def ingest_rows(self, rows: list[dict[str, Any]], overwrite: bool) -> None:
+ """Loads rows into the table, creating the dataset/table if needed.
+
+ The table is created from the module schema when it does not already
+ exist. When overwrite is True the existing rows are replaced; when
+ False the new rows are appended.
+
+ Args:
+ rows: Row dicts keyed by the schema column names.
+ overwrite: Whether to replace existing rows instead of appending.
+ """
+ self._ensure_dataset()
+
+ schema = [
+ self._bigquery.SchemaField(column_name, field_type)
+ for column_name, field_type in _BIGQUERY_SCHEMA
+ ]
+ if overwrite:
+ write_disposition = self._bigquery.WriteDisposition.WRITE_TRUNCATE
+ else:
+ write_disposition = self._bigquery.WriteDisposition.WRITE_APPEND
+
+ job_config = self._bigquery.LoadJobConfig(
+ schema=schema,
+ write_disposition=write_disposition,
+ )
+ load_job = self._client.load_table_from_json(
+ rows, self._table_reference, job_config=job_config
+ )
+ load_job.result()
+ _LOGGER.info("Ingested %d row(s) into %s", len(rows), self._table_reference)
+
+ def _ensure_dataset(self) -> None:
+ """Creates the dataset if it does not already exist."""
+ dataset_reference = self._bigquery.DatasetReference(
+ self._project_id, self._dataset_id
+ )
+ try:
+ self._client.get_dataset(dataset_reference)
+ except exceptions.NotFound:
+ _LOGGER.info("Dataset %s not found. Creating ...", self._dataset_id)
+ self._client.create_dataset(self._bigquery.Dataset(dataset_reference))
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/gcs_ops_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/gcs_ops_test.py
new file mode 100644
index 00000000000..4c935b38b61
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/gcs_ops_test.py
@@ -0,0 +1,295 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for gcs_ops.py.
+
+The Cloud Storage and BigQuery clients are patched so no network calls or
+credentials are needed. The pure URI/path/row helpers are exercised directly.
+"""
+
+from unittest import mock
+
+from absl.testing import absltest
+from absl.testing import parameterized
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import gcs_ops
+
+
+class ParseGcsUriTest(parameterized.TestCase):
+ """Tests for _parse_gcs_uri."""
+
+ @parameterized.named_parameters(
+ ("bucket_and_prefix", "gs://bucket/a/b/", ("bucket", "a/b/")),
+ ("bucket_root", "gs://bucket/", ("bucket", "")),
+ ("bucket_no_slash", "gs://bucket", ("bucket", "")),
+ ("nested_object", "gs://bkt/x/y/z.png", ("bkt", "x/y/z.png")),
+ )
+ def test_splits_bucket_and_prefix(self, uri, expected):
+ """Verifies the bucket and object prefix are split correctly."""
+ self.assertEqual(gcs_ops._parse_gcs_uri(uri), expected)
+
+ def test_raises_without_scheme(self):
+ """Verifies a URI missing the gs:// scheme raises ValueError."""
+ with self.assertRaises(ValueError):
+ gcs_ops._parse_gcs_uri("bucket/path/")
+
+ def test_raises_without_bucket(self):
+ """Verifies a URI with no bucket name raises ValueError."""
+ with self.assertRaises(ValueError):
+ gcs_ops._parse_gcs_uri("gs:///path/")
+
+
+class TrackGridFilenameTest(parameterized.TestCase):
+ """Tests for _track_grid_filename."""
+
+ def test_returns_basename_for_real_path(self):
+ """Verifies a real path is reduced to its basename."""
+ self.assertEqual(
+ gcs_ops._track_grid_filename("/out/grade3/track_0011_dirt.png"),
+ "track_0011_dirt.png",
+ )
+
+ @parameterized.named_parameters(
+ ("empty", ""),
+ ("na_marker", "N/A (grids disabled)"),
+ )
+ def test_returns_na_for_placeholder(self, output_path):
+ """Verifies an empty or placeholder path reports 'N/A'."""
+ self.assertEqual(gcs_ops._track_grid_filename(output_path), "N/A")
+
+
+class BuildTrackRowsTest(absltest.TestCase):
+ """Tests for build_track_rows."""
+
+ def test_builds_one_row_per_track_sorted_by_id(self):
+ """Verifies rows are produced per track, ordered by tracker_id."""
+ track_summary = {
+ 2: {
+ "final_class": "dirt_jars_grade3",
+ "category": "grade3",
+ "vote_count": 4,
+ "output_path": "/o/grade3/track_0002_dirt_jars_grade3.png",
+ },
+ 1: {
+ "final_class": "clean_jars_grade1",
+ "category": "grade1",
+ "vote_count": 7,
+ "output_path": "/o/grade1/track_0001_clean_jars_grade1.png",
+ },
+ }
+ rows = gcs_ops.build_track_rows(track_summary, session_name="session_a")
+ self.assertEqual([row["tracker_id"] for row in rows], [1, 2])
+ self.assertEqual(rows[0]["session_name"], "session_a")
+ self.assertEqual(rows[0]["final_class"], "clean_jars_grade1")
+ self.assertEqual(rows[0]["collapsed_class"], "grade1")
+ self.assertEqual(
+ rows[0]["track_grid_filename"], "track_0001_clean_jars_grade1.png"
+ )
+
+ def test_row_keys_match_table_columns(self):
+ """Verifies every row carries exactly the schema columns."""
+ track_summary = {
+ 0: {
+ "final_class": "c",
+ "category": None,
+ "vote_count": 1,
+ "output_path": "N/A (grids disabled)",
+ }
+ }
+ rows = gcs_ops.build_track_rows(track_summary, session_name="s")
+ self.assertCountEqual(rows[0].keys(), gcs_ops.TRACK_TABLE_COLUMNS)
+
+ def test_disabled_grid_reports_na_filename(self):
+ """Verifies a placeholder output path maps to an 'N/A' filename."""
+ track_summary = {
+ 0: {
+ "final_class": "c",
+ "category": None,
+ "vote_count": 1,
+ "output_path": "N/A (grids disabled)",
+ }
+ }
+ rows = gcs_ops.build_track_rows(track_summary, session_name="s")
+ self.assertEqual(rows[0]["track_grid_filename"], "N/A")
+ self.assertIsNone(rows[0]["collapsed_class"])
+
+
+class CloudStorageManagerPathHelpersTest(parameterized.TestCase):
+ """Tests for the pure path helpers on CloudStorageManager.
+
+ The storage client is patched so constructing the manager does not touch
+ GCP; only the string-manipulation helpers are exercised.
+ """
+
+ def setUp(self):
+ super().setUp()
+ self.enter_context(
+ mock.patch.object(gcs_ops.storage, "Client", autospec=True)
+ )
+ self.manager = gcs_ops.CloudStorageManager()
+
+ @parameterized.named_parameters(
+ ("nested_prefix", "a/b/images/", "a/b/"),
+ ("single_segment", "images/", ""),
+ ("empty_prefix", "", ""),
+ )
+ def test_parent_prefix(self, object_prefix, expected):
+ """Verifies the parent prefix keeps the prefix's final folder."""
+ self.assertEqual(self.manager._parent_prefix(object_prefix), expected)
+
+ def test_relative_object_path_preserves_last_folder(self):
+ """Verifies a blob path keeps the named download folder locally."""
+ result = self.manager._relative_object_path(
+ blob_name="a/b/images/img1.png", object_prefix="a/b/images/"
+ )
+ self.assertEqual(result, "images/img1.png")
+
+ @parameterized.named_parameters(
+ ("with_prefix", "out/run", "a/b.png", "out/run/a/b.png"),
+ ("empty_prefix", "", "a/b.png", "a/b.png"),
+ ("prefix_trailing_slash", "out/", "b.png", "out/b.png"),
+ )
+ def test_join_gcs_path(self, object_prefix, relative_path, expected):
+ """Verifies prefix and relative path join into a forward-slash blob name."""
+ self.assertEqual(
+ self.manager._join_gcs_path(object_prefix, relative_path), expected
+ )
+
+
+class ResolveIoDirectoriesTest(absltest.TestCase):
+ """Tests for resolve_io_directories (local branch needs no GCP client)."""
+
+ def _paths(self, local_enable: bool, gcs_enable: bool):
+ """Builds a minimal PathsConfig-like stub with local and gcs sub-objects."""
+ paths = mock.Mock()
+ paths.local = mock.Mock(
+ enable=local_enable,
+ input_image_directory="/data/in",
+ output_root_directory="/data/out",
+ )
+ paths.gcs = mock.Mock(
+ enable=gcs_enable,
+ input_uri="gs://b/in/",
+ output_uri="gs://b/out/",
+ temp_input_directory="/tmp/in",
+ temp_output_directory="/tmp/out",
+ )
+ return paths
+
+ def test_local_mode_returns_configured_dirs(self):
+ """Verifies local mode returns the configured local directories."""
+ paths = self._paths(local_enable=True, gcs_enable=False)
+ result = gcs_ops.resolve_io_directories(paths)
+ self.assertEqual(result, ("/data/in", "/data/out"))
+
+ def test_gcs_mode_downloads_and_returns_temp_dirs(self):
+ """Verifies GCS mode downloads input and returns the temp directories."""
+ paths = self._paths(local_enable=False, gcs_enable=True)
+ with mock.patch.object(
+ gcs_ops, "CloudStorageManager", autospec=True
+ ) as mock_manager_class:
+ result = gcs_ops.resolve_io_directories(paths)
+ mock_manager_class.return_value.download_directory.assert_called_once()
+ self.assertEqual(result, ("/tmp/in", "/tmp/out"))
+
+
+class UploadOutputDirectoryTest(absltest.TestCase):
+ """Tests for upload_output_directory."""
+
+ def test_local_mode_is_noop(self):
+ """Verifies local mode does not construct a storage manager."""
+ paths = mock.Mock()
+ paths.gcs = mock.Mock(enable=False)
+ with mock.patch.object(
+ gcs_ops, "CloudStorageManager", autospec=True
+ ) as mock_manager_class:
+ gcs_ops.upload_output_directory(paths)
+ mock_manager_class.assert_not_called()
+
+ def test_gcs_mode_uploads(self):
+ """Verifies GCS mode uploads the temp output directory."""
+ paths = mock.Mock()
+ paths.gcs = mock.Mock(
+ enable=True,
+ output_uri="gs://b/out/",
+ temp_output_directory="/tmp/out",
+ )
+ with mock.patch.object(
+ gcs_ops, "CloudStorageManager", autospec=True
+ ) as mock_manager_class:
+ gcs_ops.upload_output_directory(paths)
+ mock_manager_class.return_value.upload_directory.assert_called_once()
+
+
+class IngestTrackRowsTest(absltest.TestCase):
+ """Tests for ingest_track_rows."""
+
+ def test_disabled_is_noop(self):
+ """Verifies a disabled BigQuery config skips ingestion."""
+ config = mock.Mock(enable=False)
+ with mock.patch.object(
+ gcs_ops, "BigQueryManager", autospec=True
+ ) as mock_manager_class:
+ gcs_ops.ingest_track_rows(config, rows=[{"a": 1}])
+ mock_manager_class.assert_not_called()
+
+ def test_enabled_with_no_rows_is_noop(self):
+ """Verifies an empty row list skips ingestion even when enabled."""
+ config = mock.Mock(
+ enable=True, project_id="p", dataset_id="d", table_id="t"
+ )
+ with mock.patch.object(
+ gcs_ops, "BigQueryManager", autospec=True
+ ) as mock_manager_class:
+ gcs_ops.ingest_track_rows(config, rows=[])
+ mock_manager_class.assert_not_called()
+
+ def test_enabled_with_rows_ingests(self):
+ """Verifies enabled ingestion constructs the manager and writes rows."""
+ config = mock.Mock(
+ enable=True,
+ project_id="p",
+ dataset_id="d",
+ table_id="t",
+ overwrite=True,
+ )
+ rows = [{"session_name": "s"}]
+ with mock.patch.object(
+ gcs_ops, "BigQueryManager", autospec=True
+ ) as mock_manager_class:
+ gcs_ops.ingest_track_rows(config, rows=rows)
+ mock_manager_class.assert_called_once_with(
+ project_id="p", dataset_id="d", table_id="t"
+ )
+ mock_manager_class.return_value.ingest_rows.assert_called_once_with(
+ rows, overwrite=True
+ )
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/main.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/main.py
new file mode 100644
index 00000000000..290a1ba6fcc
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/main.py
@@ -0,0 +1,475 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PET Bottle Grade Detection Pipeline Orchestrator."""
+
+from collections.abc import Sequence
+import gc
+import glob
+import logging
+import os
+import pathlib
+import sys
+from typing import Any
+
+import cv2
+from PIL import Image
+import torch
+import tqdm
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import config_loader
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import dinov3_classifier
+
+try:
+ from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import gcs_ops # pylint: disable=g-import-not-at-top
+except ImportError:
+ gcs_ops = None
+
+try:
+ from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import rfdetr_detector # pylint: disable=g-import-not-at-top
+except ImportError:
+ rfdetr_detector = None
+
+try:
+ from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import track_manager # pylint: disable=g-import-not-at-top
+except ImportError:
+ track_manager = None
+
+try:
+ from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import visualization_utils # pylint: disable=g-import-not-at-top
+except ImportError:
+ visualization_utils = None
+
+_LOGGER_NAME = "waste_identification_pipeline"
+_LOG_FORMAT = "[%(asctime)s] %(levelname)s %(message)s"
+_LOG_DATE_FORMAT = "%Y-%m-%d %H:%M:%S"
+_LOG_FILE_NAME = "pipeline.log"
+_SUMMARY_FILE_NAME = "summary.txt"
+_MEMORY_FLUSH_FRAME_INTERVAL = 10
+
+
+def get_logger() -> logging.Logger:
+ """Returns the shared, configured pipeline logger."""
+ logger = logging.getLogger(_LOGGER_NAME)
+ if logger.handlers:
+ return logger
+ logger.setLevel(logging.INFO)
+ stream_handler = logging.StreamHandler(stream=sys.stdout)
+ stream_handler.setFormatter(
+ logging.Formatter(_LOG_FORMAT, datefmt=_LOG_DATE_FORMAT)
+ )
+ logger.addHandler(stream_handler)
+ logger.propagate = False
+ return logger
+
+
+def attach_file_handler(log_file_path: str) -> logging.FileHandler:
+ """Attaches a file handler to the pipeline logger."""
+ logger = get_logger()
+ file_handler = logging.FileHandler(
+ str(log_file_path), mode="w", encoding="utf-8"
+ )
+ file_handler.setLevel(logger.level)
+ file_handler.setFormatter(
+ logging.Formatter(_LOG_FORMAT, datefmt=_LOG_DATE_FORMAT)
+ )
+ logger.addHandler(file_handler)
+ return file_handler
+
+
+def detach_file_handler(file_handler: logging.FileHandler) -> None:
+ """Removes a file handler from the pipeline logger and closes it."""
+ logger = get_logger()
+ logger.removeHandler(file_handler)
+ file_handler.close()
+
+
+def create_summary_logger(summary_file_path: str) -> logging.Logger:
+ """Creates a dedicated summary logger that writes to a file."""
+ logger = logging.getLogger(f"summary_{summary_file_path}")
+ logger.setLevel(logging.INFO)
+ logger.propagate = False
+ handler = logging.FileHandler(summary_file_path, mode="w", encoding="utf-8")
+ handler.setFormatter(logging.Formatter("%(message)s"))
+ logger.addHandler(handler)
+ return logger
+
+
+def close_summary_logger(logger: logging.Logger | None) -> None:
+ """Closes all handlers on the summary logger."""
+ if logger is None:
+ return
+ for handler in list(logger.handlers):
+ logger.removeHandler(handler)
+ handler.close()
+
+
+def ensure_directories_exist(
+ directories: Sequence[str | pathlib.Path],
+) -> None:
+ """Creates each directory in the given sequence if it does not already exist."""
+ for directory in directories:
+ if directory:
+ os.makedirs(str(directory), exist_ok=True)
+
+
+def load_and_resize_image(
+ image_path: str | pathlib.Path, max_short_side: int
+) -> Image.Image:
+ """Loads an image via OpenCV and downscales it using Lanczos interpolation.
+
+ Args:
+ image_path: Path to the image file.
+ max_short_side: Maximum size for the shorter dimension.
+
+ Returns:
+ The resized PIL Image.
+
+ Raises:
+ OSError: If OpenCV cannot read the image file.
+ """
+ bgr_image = cv2.imread(str(image_path))
+ if bgr_image is None:
+ raise OSError(
+ f"OpenCV could not read the image at {image_path!s}. The file may be "
+ "missing, corrupt, or in an unsupported format."
+ )
+ height, width = bgr_image.shape[:2]
+ short_side = min(height, width)
+ if short_side > max_short_side:
+ scale = max_short_side / short_side
+ new_w, new_h = int(width * scale), int(height * scale)
+ bgr_image = cv2.resize(
+ bgr_image, (new_w, new_h), interpolation=cv2.INTER_LANCZOS4
+ )
+ rgb_image = cv2.cvtColor(bgr_image, cv2.COLOR_BGR2RGB)
+ return Image.fromarray(rgb_image)
+
+
+def flush_memory() -> None:
+ """Runs garbage collection and empties CUDA cache if available."""
+ gc.collect()
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+
+
+class PETBottlePipeline:
+ """Orchestrates image processing, object tracking, and classification."""
+
+ def __init__(self, config_path: str):
+ """Initializes pipeline components based on the configuration file.
+
+ Models are loaded once here and reused across all subfolders.
+
+ Args:
+ config_path: Path to the YAML configuration file.
+
+ Raises:
+ RuntimeError: If rfdetr_detector module is not available.
+ """
+ self._config = config_loader.PipelineConfig.from_yaml(config_path)
+ self._summary_logger = None
+
+ logger = get_logger()
+ logger.info("Using device: %s", self._config.models.rfdetr.device)
+ logger.info(
+ "Loading RFDETR model from %s ...",
+ self._config.models.rfdetr.checkpoint_path,
+ )
+ if rfdetr_detector is None:
+ raise RuntimeError("rfdetr_detector module is not available.")
+ self._detector = rfdetr_detector.RFDETRDetector(
+ rfdetr_config=self._config.models.rfdetr,
+ post_processing_config=self._config.post_processing,
+ )
+ logger.info("RFDETR model ready.")
+
+ logger.info(
+ "Loading DINOv3 classifier (%s) from %s ...",
+ self._config.models.dinov3.model_name,
+ self._config.models.dinov3.checkpoint_path,
+ )
+ self._classifier = dinov3_classifier.DINOv3Classifier.from_config(
+ config=self._config.models.dinov3,
+ class_names=self._config.classes,
+ device=self._config.models.rfdetr.device,
+ )
+ logger.info("DINOv3 classifier ready.")
+
+ def run(self) -> None:
+ """Executes the pipeline over every immediate subfolder of the input directory.
+
+ Raises:
+ RuntimeError: If gcs_ops module is not available.
+ """
+ logger = get_logger()
+ if gcs_ops is None:
+ raise RuntimeError("gcs_ops module is not available.")
+ input_root, output_root = gcs_ops.resolve_io_directories(self._config.paths)
+
+ subfolders = self._discover_subfolders(input_root)
+ if not subfolders:
+ logger.warning(
+ "No subfolders found in %s. Nothing to process.", input_root
+ )
+ return
+
+ ensure_directories_exist([output_root])
+ logger.info(
+ "Found %d subfolder(s) to process under %s",
+ len(subfolders),
+ input_root,
+ )
+
+ summary_file_path = os.path.join(output_root, _SUMMARY_FILE_NAME)
+ self._summary_logger = create_summary_logger(summary_file_path)
+ try:
+ if not self._config.tracking.enable:
+ self._summary_logger.info(
+ "Tracking disabled - no per-subfolder summaries were generated."
+ )
+
+ all_track_rows = []
+ for subfolder_path in subfolders:
+ subfolder_name = os.path.basename(subfolder_path)
+ subfolder_output_dir = os.path.join(output_root, subfolder_name)
+ try:
+ subfolder_rows = self._process_subfolder(
+ subfolder_path, subfolder_output_dir
+ )
+ all_track_rows.extend(subfolder_rows)
+ except Exception: # pylint: disable=broad-exception-caught
+ logger.exception(
+ "Subfolder '%s' failed. Skipping to next.",
+ subfolder_name,
+ )
+ continue
+
+ logger.info("Pipeline completed for all subfolders.")
+ logger.info("Summary written to %s", summary_file_path)
+ gcs_ops.upload_output_directory(self._config.paths)
+ gcs_ops.ingest_track_rows(self._config.bigquery, all_track_rows)
+ finally:
+ close_summary_logger(self._summary_logger)
+ self._summary_logger = None
+
+ def _process_subfolder(
+ self, input_subfolder: str, output_subfolder: str
+ ) -> list[dict[str, Any]]:
+ """Runs the per-frame loop and classification for a single subfolder.
+
+ Tracks are independent per subfolder: a fresh TrackManager and
+ PipelineVisualizer are instantiated for every call. A dedicated log
+ file is attached for the duration of this subfolder's processing.
+
+ Args:
+ input_subfolder: Absolute path to the input subfolder containing image
+ frames.
+ output_subfolder: Absolute path to the output subfolder where results,
+ logs, and intermediate artifacts will be written.
+
+ Returns:
+ The per-track rows for this subfolder, as produced by build_track_rows.
+ Empty when the subfolder has no images.
+
+ Raises:
+ RuntimeError: If track_manager, visualization_utils, or gcs_ops is not
+ available.
+ """
+ logger = get_logger()
+ subfolder_name = os.path.basename(input_subfolder)
+
+ frame_output_dir = os.path.join(
+ output_subfolder, self._config.paths.output_frame_subfolder
+ )
+ grid_output_dir = os.path.join(
+ output_subfolder, self._config.paths.track_grid_subfolder
+ )
+ video_output_path = os.path.join(
+ output_subfolder, self._config.paths.output_video_filename
+ )
+ log_file_path = os.path.join(output_subfolder, _LOG_FILE_NAME)
+
+ ensure_directories_exist([output_subfolder])
+ directories_to_create = []
+ if self._config.visualization.save_frames:
+ directories_to_create.append(frame_output_dir)
+ if self._config.visualization.save_track_grids:
+ directories_to_create.append(grid_output_dir)
+ ensure_directories_exist(directories_to_create)
+
+ file_handler = attach_file_handler(log_file_path)
+ try:
+ logger.info("Processing subfolder: %s", subfolder_name)
+
+ image_paths = self._collect_image_paths(input_subfolder)
+ if not image_paths:
+ logger.warning(
+ "No images found in subfolder '%s'. Skipping.", subfolder_name
+ )
+ return []
+
+ logger.info("Found %d image(s) in %s", len(image_paths), subfolder_name)
+
+ if (
+ track_manager is None
+ or visualization_utils is None
+ or gcs_ops is None
+ ):
+ raise RuntimeError(
+ "track_manager, visualization_utils, and gcs_ops modules are"
+ " required."
+ )
+
+ tracker = track_manager.TrackManager(
+ tracking_config=self._config.tracking,
+ cropping_config=self._config.cropping,
+ vis_config=self._config.visualization,
+ )
+ tracker.reset()
+ visualizer = visualization_utils.PipelineVisualizer(
+ config=self._config.visualization,
+ collapsed_categories=self._config.collapsed_categories,
+ out_video_path=video_output_path,
+ summary_logger=self._summary_logger,
+ )
+
+ self._run_frame_loop(image_paths, tracker, visualizer, frame_output_dir)
+
+ visualizer.close_video()
+ flush_memory()
+
+ track_predictions = tracker.classify_all_tracks(
+ self._classifier,
+ self._config.models.dinov3.classification_batch_size,
+ )
+ track_summary = visualizer.save_track_grids(
+ track_predictions, tracker.resolve_track_label, grid_output_dir
+ )
+ if self._config.tracking.enable:
+ visualizer.print_summary(
+ track_summary, input_subfolder, self._config.classes
+ )
+ else:
+ logger.info("Tracking disabled: skipping per-subfolder summary.")
+
+ track_rows = gcs_ops.build_track_rows(track_summary, subfolder_name)
+
+ logger.info("Finished subfolder: %s", subfolder_name)
+ return track_rows
+ finally:
+ detach_file_handler(file_handler)
+
+ def _run_frame_loop(
+ self,
+ image_paths: list[str],
+ tracker: Any,
+ visualizer: Any,
+ frame_output_dir: str,
+ ) -> None:
+ """Iterates over the image frames, running detection, tracking, and visualization.
+
+ Args:
+ image_paths: Sorted list of image file paths to process.
+ tracker: Per-subfolder TrackManager (already instantiated).
+ visualizer: Per-subfolder PipelineVisualizer (already instantiated).
+ frame_output_dir: Directory where annotated frames are saved.
+ """
+ progress_bar = tqdm.tqdm(image_paths, desc="Tracking")
+
+ for frame_index, image_path in enumerate(progress_bar):
+ resized_image = load_and_resize_image(
+ image_path, self._config.preprocessing.max_short_side
+ )
+
+ state = self._detector.detect(resized_image)
+ detections = self._detector.to_supervision_detections(state)
+
+ detections, unassigned_scores = tracker.update_and_extract_crops(
+ detections, state, resized_image, os.path.basename(image_path)
+ )
+
+ if unassigned_scores:
+ formatted_scores = ", ".join(
+ f"{score:.3f}" for score in unassigned_scores
+ )
+ progress_bar.write(
+ f" frame {frame_index:06d} unassigned scores: {formatted_scores}"
+ )
+
+ frame_output_path = os.path.join(
+ frame_output_dir, os.path.basename(image_path)
+ )
+ visualizer.annotate_and_write_frame(
+ resized_image, detections, frame_output_path
+ )
+
+ del resized_image, state, detections
+ if frame_index % _MEMORY_FLUSH_FRAME_INTERVAL == 0:
+ flush_memory()
+
+ def _collect_image_paths(self, directory: str) -> list[str]:
+ """Collects and sorts image file paths matching the configured extensions.
+
+ Args:
+ directory: Directory to scan for image files (non-recursive).
+
+ Returns:
+ Sorted list of image file paths.
+ """
+ image_paths = []
+ for extension_pattern in self._config.models.rfdetr.image_file_extensions:
+ image_paths.extend(glob.glob(os.path.join(directory, extension_pattern)))
+ return sorted(image_paths)
+
+ def _discover_subfolders(self, root_directory: str) -> list[str]:
+ """Returns sorted absolute paths of immediate (direct child) subfolders.
+
+ Args:
+ root_directory: Directory whose direct child folders should be returned.
+
+ Returns:
+ Sorted list of absolute subfolder paths.
+ """
+ logger = get_logger()
+ if not os.path.isdir(root_directory):
+ logger.error("Input directory does not exist: %s", root_directory)
+ return []
+
+ entries = os.listdir(root_directory)
+ subfolders = [
+ os.path.join(root_directory, entry)
+ for entry in entries
+ if os.path.isdir(os.path.join(root_directory, entry))
+ ]
+ return sorted(subfolders)
+
+
+if __name__ == "__main__":
+ pipeline = PETBottlePipeline(config_path="config.yaml")
+ pipeline.run()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/main_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/main_test.py
new file mode 100644
index 00000000000..dea7a3753bc
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/main_test.py
@@ -0,0 +1,107 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for main.py.
+
+The pipeline's model-loading __init__ is bypassed with __new__ so the
+filesystem discovery helpers can be tested without constructing RFDETR or
+DINOv3. The per-frame orchestration is out of scope for these focused tests.
+"""
+
+import pathlib
+from unittest import mock
+
+from absl.testing import absltest
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import main
+
+
+def _touch(path: pathlib.Path) -> None:
+ """Creates an empty file, making parent directories as needed."""
+ path.parent.mkdir(parents=True, exist_ok=True)
+ path.write_bytes(b"")
+
+
+def _make_pipeline(
+ image_file_extensions: list[str],
+) -> main.PETBottlePipeline:
+ """Builds a PETBottlePipeline via __new__ with only the config it needs."""
+ pipeline = object.__new__(main.PETBottlePipeline)
+ pipeline._config = mock.Mock()
+ pipeline._config.models.rfdetr.image_file_extensions = image_file_extensions
+ return pipeline
+
+
+class DiscoverSubfoldersTest(absltest.TestCase):
+ """Tests for _discover_subfolders."""
+
+ def test_returns_sorted_immediate_subfolders(self):
+ """Verifies only immediate child directories are returned, sorted."""
+ root = pathlib.Path(self.create_tempdir().full_path)
+ (root / "session_b").mkdir()
+ (root / "session_a").mkdir()
+ _touch(root / "loose.png")
+
+ pipeline = _make_pipeline(["*.png"])
+ result = pipeline._discover_subfolders(str(root))
+ names = [pathlib.Path(path).name for path in result]
+ self.assertEqual(names, ["session_a", "session_b"])
+
+ def test_missing_directory_returns_empty(self):
+ """Verifies a non-existent input directory returns an empty list."""
+ pipeline = _make_pipeline(["*.png"])
+ self.assertEqual(pipeline._discover_subfolders("/no/such/dir"), [])
+
+
+class CollectImagePathsTest(absltest.TestCase):
+ """Tests for _collect_image_paths."""
+
+ def test_collects_only_configured_extensions_sorted(self):
+ """Verifies only files matching configured globs are returned, sorted."""
+ directory = pathlib.Path(self.create_tempdir().full_path)
+ _touch(directory / "a.png")
+ _touch(directory / "b.jpg")
+ _touch(directory / "c.txt")
+ _touch(directory / "d.png")
+
+ pipeline = _make_pipeline(["*.png", "*.jpg"])
+ result = pipeline._collect_image_paths(str(directory))
+ names = [pathlib.Path(path).name for path in result]
+ self.assertEqual(names, ["a.png", "b.jpg", "d.png"])
+
+ def test_returns_empty_when_no_matches(self):
+ """Verifies a directory with no matching files yields an empty list."""
+ directory = pathlib.Path(self.create_tempdir().full_path)
+ _touch(directory / "notes.txt")
+
+ pipeline = _make_pipeline(["*.png"])
+ self.assertEqual(pipeline._collect_image_paths(str(directory)), [])
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/rfdetr_detector.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/rfdetr_detector.py
new file mode 100644
index 00000000000..fe1f8c85124
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/rfdetr_detector.py
@@ -0,0 +1,364 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""RFDETR detection and segmentation management.
+
+RFDETR is used purely as a "find and segment objects" engine: every
+detection contributes a mask, box, and score, and the RFDETR class
+prediction is intentionally discarded. Downstream tracking and DINOv3
+grading are unchanged.
+
+The returned state dict shape consumed by the rest of the pipeline
+(TrackManager crop extraction, visualization, BigQuery ingestion):
+
+ * ``masks`` torch.bool tensor of shape [N, 1, H, W]
+ * ``boxes`` torch.float32 tensor of shape [N, 4], (x_min, y_min, x_max,
+ y_max)
+ * ``scores`` torch.float32 tensor of shape [N]
+
+``merge_contained_boxes`` runs unconditionally because RFDETR has no
+prompt to gate it on. ``optimize_for_inference`` is called once at model
+build.
+"""
+
+import collections
+
+import numpy as np
+from PIL import Image
+import supervision
+import torch
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import config_loader
+
+try:
+ from rfdetr import RFDETRSegMedium # pylint: disable=g-import-not-at-top
+except ImportError:
+ RFDETRSegMedium = None
+
+_STATE_ARRAY_KEYS = ("masks", "boxes", "scores")
+
+
+class RFDETRDetector:
+ """Handles RFDETR inference and memory-safe post-processing.
+
+ Public surface:
+
+ * ``detect(image)`` returns the state dict.
+ * ``to_supervision_detections(state)`` returns
+ ``supervision.Detections`` with the score threshold applied.
+ """
+
+ def __init__(
+ self,
+ rfdetr_config: config_loader.RFDETRConfig,
+ post_processing_config: config_loader.PostProcessingConfig,
+ ):
+ """Initializes the RFDETR model and applies inference optimization.
+
+ Args:
+ rfdetr_config: Validated RFDETR-model config (checkpoint, device, and the
+ ``predict_threshold`` fed to ``RFDETR.predict``).
+ post_processing_config: Validated post-processing thresholds (mask-in-mask
+ filter, box-in-box merge, final score cutoff).
+
+ Raises:
+ RuntimeError: If the rfdetr package is not installed.
+ """
+ if RFDETRSegMedium is None:
+ raise RuntimeError(
+ "rfdetr is not installed. Please install it to use RFDETRDetector."
+ )
+ self._device = torch.device(
+ rfdetr_config.device if torch.cuda.is_available() else "cpu"
+ )
+ self._predict_threshold = rfdetr_config.predict_threshold
+ self._post_processing = post_processing_config
+ self._model = RFDETRSegMedium(
+ pretrain_weights=rfdetr_config.checkpoint_path
+ )
+ self._model.optimize_for_inference()
+
+ def detect(self, image: Image.Image) -> dict[str, torch.Tensor]:
+ """Runs RFDETR inference and returns a CPU-resident state dictionary.
+
+ Args:
+ image: PIL RGB image (already resized to the pipeline's target short
+ side).
+
+ Returns:
+ State dict with keys ``masks``, ``boxes``, ``scores``. All
+ tensors are on CPU.
+ """
+ image_width, image_height = image.size
+ raw_detections = self._model.predict(
+ image, threshold=self._predict_threshold
+ )
+
+ state = self._convert_detections_to_state(
+ detections=raw_detections,
+ image_height=image_height,
+ image_width=image_width,
+ )
+
+ state = self._filter_contained_sub_masks(state)
+ state = self._merge_contained_boxes(state)
+
+ return state
+
+ def to_supervision_detections(
+ self, state: dict[str, torch.Tensor]
+ ) -> supervision.Detections:
+ """Converts a state dict to ``supervision.Detections``.
+
+ Applies the configured ``score_threshold`` before returning. Every
+ surviving detection is emitted with ``class_id=0`` so downstream
+ tracking / voting does not depend on RFDETR's class predictions.
+
+ Args:
+ state: The state dict returned by ``detect``.
+
+ Returns:
+ A ``supervision.Detections`` instance.
+ """
+ boxes = state["boxes"].numpy().astype(np.float32)
+ scores = state["scores"].numpy().astype(np.float32)
+
+ if boxes.ndim == 1:
+ boxes = boxes.reshape(0, 4)
+
+ keep_mask = scores >= self._post_processing.score_threshold
+ kept_boxes = boxes[keep_mask]
+ kept_scores = scores[keep_mask]
+
+ if kept_boxes.shape[0] == 0:
+ return supervision.Detections.empty()
+
+ return supervision.Detections(
+ xyxy=kept_boxes,
+ confidence=kept_scores,
+ class_id=np.zeros(kept_boxes.shape[0], dtype=int),
+ )
+
+ def _convert_detections_to_state(
+ self,
+ detections: supervision.Detections,
+ image_height: int,
+ image_width: int,
+ ) -> dict[str, torch.Tensor]:
+ """Adapts a ``supervision.Detections`` into the pipeline state dict.
+
+ RFDETR's ``class_id`` is intentionally discarded. Masks are
+ reshaped to ``[N, 1, H, W]`` because the box-merge helper below
+ calls ``masks.squeeze(1)``.
+
+ Args:
+ detections: The ``supervision.Detections`` returned by
+ ``RFDETRSegMedium.predict``.
+ image_height: Height of the image that produced the detections.
+ image_width: Width of the image that produced the detections.
+
+ Returns:
+ A CPU-resident state dict.
+ """
+ if detections.mask is None or len(detections) == 0:
+ masks = torch.zeros((0, 1, image_height, image_width), dtype=torch.bool)
+ boxes = torch.zeros((0, 4), dtype=torch.float32)
+ scores = torch.zeros((0,), dtype=torch.float32)
+ else:
+ masks = torch.from_numpy(detections.mask.astype(bool)).unsqueeze(1)
+ boxes = torch.from_numpy(detections.xyxy.astype(np.float32))
+ scores = torch.from_numpy(detections.confidence.astype(np.float32))
+
+ return {
+ "masks": masks,
+ "boxes": boxes,
+ "scores": scores,
+ }
+
+ def _filter_contained_sub_masks(
+ self, state: dict[str, torch.Tensor]
+ ) -> dict[str, torch.Tensor]:
+ """Removes smaller masks that are contained inside larger masks.
+
+ Uses ``intersection / smaller_mask_area`` and drops the smaller
+ of the two when the ratio exceeds
+ ``post_processing.containment_threshold``.
+
+ Args:
+ state: State dict from ``_convert_detections_to_state``.
+
+ Returns:
+ The filtered state dict.
+ """
+ masks = state["masks"]
+ num_masks = masks.shape[0]
+ if num_masks == 0:
+ return state
+
+ flat_masks = masks.view(num_masks, -1).float()
+ pairwise_intersection = flat_masks @ flat_masks.T
+
+ indices_to_remove = set()
+ for i in range(num_masks):
+ if i in indices_to_remove:
+ continue
+ for j in range(i + 1, num_masks):
+ if j in indices_to_remove:
+ continue
+
+ area_i = pairwise_intersection[i, i].item()
+ area_j = pairwise_intersection[j, j].item()
+ smaller_index, smaller_area = (
+ (i, area_i) if area_i <= area_j else (j, area_j)
+ )
+
+ if smaller_area == 0:
+ indices_to_remove.add(smaller_index)
+ continue
+
+ intersection = pairwise_intersection[i, j].item()
+ if (
+ intersection / smaller_area
+ ) > self._post_processing.containment_threshold:
+ indices_to_remove.add(smaller_index)
+
+ keep_tensor = torch.tensor(
+ sorted(set(range(num_masks)) - indices_to_remove),
+ dtype=torch.long,
+ device=masks.device,
+ )
+
+ for key in _STATE_ARRAY_KEYS:
+ state[key] = state[key][keep_tensor]
+
+ return state
+
+ def _merge_contained_boxes(
+ self, state: dict[str, torch.Tensor]
+ ) -> dict[str, torch.Tensor]:
+ """Merges detections where a smaller box sits inside a larger one.
+
+ Uses ``intersection_area / smaller_box_area`` and merges the
+ smaller detection into the larger one when the ratio exceeds
+ ``post_processing.merge_containment_threshold``. Runs
+ unconditionally.
+
+ The merged detection's mask is the union of the group's masks,
+ its box is the enclosing box, and its score is the (clamped) sum
+ of the group's scores.
+
+ Args:
+ state: State dict from ``_filter_contained_sub_masks``.
+
+ Returns:
+ The merged state dict.
+ """
+ masks = state["masks"]
+ boxes = state["boxes"]
+ scores = state["scores"]
+ if len(scores) == 0:
+ return state
+
+ num_detections = len(masks)
+ box_areas = (boxes[:, 2] - boxes[:, 0]) * (boxes[:, 3] - boxes[:, 1])
+
+ is_absorbed = torch.zeros(
+ num_detections, dtype=torch.bool, device=boxes.device
+ )
+ absorb_target = list(range(num_detections))
+
+ for i in range(num_detections):
+ if is_absorbed[i]:
+ continue
+ for j in range(i + 1, num_detections):
+ if is_absorbed[j]:
+ continue
+
+ inter_x_min = torch.max(boxes[i, 0], boxes[j, 0])
+ inter_y_min = torch.max(boxes[i, 1], boxes[j, 1])
+ inter_x_max = torch.min(boxes[i, 2], boxes[j, 2])
+ inter_y_max = torch.min(boxes[i, 3], boxes[j, 3])
+
+ intersection_area = torch.clamp(
+ inter_x_max - inter_x_min, min=0
+ ) * torch.clamp(inter_y_max - inter_y_min, min=0)
+
+ smaller_index, larger_index, smaller_area = (
+ (i, j, box_areas[i])
+ if box_areas[i] <= box_areas[j]
+ else (j, i, box_areas[j])
+ )
+
+ if smaller_area == 0:
+ is_absorbed[smaller_index] = True
+ continue
+
+ containment_ratio = intersection_area / smaller_area
+ if (
+ containment_ratio
+ > self._post_processing.merge_containment_threshold
+ ):
+ is_absorbed[smaller_index] = True
+ absorb_target[smaller_index] = larger_index
+
+ groups: collections.defaultdict[int, list[int]] = collections.defaultdict(
+ list
+ )
+ for i in range(num_detections):
+ target = absorb_target[i] if is_absorbed[i] else i
+ groups[target].append(i)
+
+ merged_masks: list[torch.Tensor] = []
+ merged_boxes: list[torch.Tensor] = []
+ merged_scores: list[torch.Tensor] = []
+
+ for member_indices in groups.values():
+ member_tensor = torch.tensor(
+ member_indices, dtype=torch.long, device=boxes.device
+ )
+ merged_masks.append(masks[member_tensor].squeeze(1).any(dim=0))
+ group_boxes = boxes[member_tensor]
+ merged_boxes.append(
+ torch.stack([
+ group_boxes[:, 0].min(),
+ group_boxes[:, 1].min(),
+ group_boxes[:, 2].max(),
+ group_boxes[:, 3].max(),
+ ])
+ )
+ merged_scores.append(
+ torch.tensor(
+ min(scores[member_tensor].sum().item(), 1.0),
+ device=scores.device,
+ )
+ )
+
+ state["masks"] = torch.stack(merged_masks).unsqueeze(1)
+ state["boxes"] = torch.stack(merged_boxes)
+ state["scores"] = torch.stack(merged_scores)
+ return state
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/rfdetr_detector_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/rfdetr_detector_test.py
new file mode 100644
index 00000000000..c7beba34f74
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/rfdetr_detector_test.py
@@ -0,0 +1,373 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for rfdetr_detector.py.
+
+RFDETRSegMedium is patched at construction so no real model is loaded; the
+pure tensor post-processing (mask filter, box merge, state conversion,
+supervision export) is exercised on hand-built state dicts.
+"""
+
+import sys
+from unittest import mock
+
+from absl.testing import absltest
+import numpy as np
+import torch
+
+# Mock external dependencies before importing rfdetr_detector since they are
+# external pip packages not checked into //third_party/py.
+mock_rfdetr = mock.MagicMock()
+
+
+class MockRFDETRSegMedium:
+ """Mock class for RFDETRSegMedium."""
+
+ def __init__(self, pretrain_weights=None):
+ pass
+
+ def optimize_for_inference(self):
+ pass
+
+ def predict(self, image, threshold=0.0):
+ pass
+
+
+mock_rfdetr.RFDETRSegMedium = MockRFDETRSegMedium
+sys.modules.setdefault("rfdetr", mock_rfdetr)
+
+mock_supervision = mock.MagicMock()
+
+
+class MockDetections:
+ """Mock class for supervision.Detections."""
+
+ def __init__(self, xyxy=None, confidence=None, class_id=None, mask=None):
+ self.xyxy = xyxy
+ self.confidence = confidence
+ self.class_id = class_id
+ self.mask = mask
+
+ def __len__(self):
+ return len(self.xyxy) if self.xyxy is not None else 0
+
+ @classmethod
+ def empty(cls):
+ return cls(
+ xyxy=np.zeros((0, 4), dtype=np.float32),
+ confidence=np.zeros(0, dtype=np.float32),
+ class_id=np.zeros(0, dtype=int),
+ )
+
+
+mock_supervision.Detections = MockDetections
+sys.modules.setdefault("supervision", mock_supervision)
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import rfdetr_detector # pylint: disable=g-bad-import-order, g-import-not-at-top
+
+
+def _box_mask(height: int, width: int, y0: int, y1: int, x0: int, x1: int):
+ """Returns a single ``[1, H, W]`` bool mask with a filled rectangle."""
+ mask = torch.zeros((1, height, width), dtype=torch.bool)
+ mask[0, y0:y1, x0:x1] = True
+ return mask
+
+
+def _make_detector(
+ containment_threshold: float = 0.98,
+ merge_containment_threshold: float = 0.7,
+ score_threshold: float = 0.0,
+) -> rfdetr_detector.RFDETRDetector:
+ """Builds an RFDETRDetector with the model and CUDA checks mocked out."""
+ rfdetr_config = mock.Mock(
+ device="cpu", checkpoint_path="/tmp/ckpt.pth", predict_threshold=0.2
+ )
+ post_processing_config = mock.Mock(
+ containment_threshold=containment_threshold,
+ merge_containment_threshold=merge_containment_threshold,
+ score_threshold=score_threshold,
+ )
+ with mock.patch.object(
+ rfdetr_detector, "RFDETRSegMedium", autospec=True
+ ), mock.patch.object(
+ rfdetr_detector.torch.cuda, "is_available", return_value=False
+ ):
+ return rfdetr_detector.RFDETRDetector(
+ rfdetr_config=rfdetr_config,
+ post_processing_config=post_processing_config,
+ )
+
+
+class ConstructorTest(absltest.TestCase):
+ """Tests for RFDETRDetector construction."""
+
+ def test_optimizes_model_for_inference(self):
+ """Verifies the model is optimized for inference at build time."""
+ rfdetr_config = mock.Mock(
+ device="cpu", checkpoint_path="/tmp/ckpt.pth", predict_threshold=0.2
+ )
+ post_processing_config = mock.Mock(
+ containment_threshold=0.98,
+ merge_containment_threshold=0.7,
+ score_threshold=0.0,
+ )
+ with mock.patch.object(
+ rfdetr_detector, "RFDETRSegMedium", autospec=True
+ ) as mock_model_class, mock.patch.object(
+ rfdetr_detector.torch.cuda, "is_available", return_value=False
+ ):
+ rfdetr_detector.RFDETRDetector(
+ rfdetr_config=rfdetr_config,
+ post_processing_config=post_processing_config,
+ )
+ mock_model_class.assert_called_once_with(pretrain_weights="/tmp/ckpt.pth")
+ mock_model_class.return_value.optimize_for_inference.assert_called_once()
+
+
+class ConvertDetectionsToStateTest(absltest.TestCase):
+ """Tests for _convert_detections_to_state."""
+
+ def test_empty_detections_yield_zero_length_tensors(self):
+ """Verifies a None-mask detection object yields empty tensors."""
+ detector = _make_detector()
+ detections = mock.Mock()
+ detections.mask = None
+ detections.__len__ = mock.Mock(return_value=0)
+
+ state = detector._convert_detections_to_state(
+ detections, image_height=20, image_width=30
+ )
+ self.assertEqual(state["masks"].shape, (0, 1, 20, 30))
+ self.assertEqual(state["boxes"].shape, (0, 4))
+ self.assertEqual(state["scores"].shape, (0,))
+
+ def test_none_mask_with_non_zero_len_yields_empty_tensors(self):
+ """Verifies non-zero detections with None mask yields empty tensors."""
+ detector = _make_detector()
+ detections = mock.Mock()
+ detections.mask = None
+ detections.__len__ = mock.Mock(return_value=1)
+
+ state = detector._convert_detections_to_state(
+ detections, image_height=20, image_width=30
+ )
+ self.assertEqual(state["masks"].shape, (0, 1, 20, 30))
+ self.assertEqual(state["boxes"].shape, (0, 4))
+ self.assertEqual(state["scores"].shape, (0,))
+
+ def test_populated_detections_add_channel_dim(self):
+ """Verifies masks gain a channel dim and scores are converted."""
+ detector = _make_detector()
+ detections = mock.Mock()
+ detections.mask = np.ones((2, 8, 8), dtype=bool)
+ detections.xyxy = np.array([[0, 0, 4, 4], [1, 1, 6, 6]], dtype=np.float32)
+ detections.confidence = np.array([0.9, 0.6], dtype=np.float32)
+ detections.__len__ = mock.Mock(return_value=2)
+
+ state = detector._convert_detections_to_state(
+ detections, image_height=8, image_width=8
+ )
+ self.assertEqual(state["masks"].shape, (2, 1, 8, 8))
+ self.assertEqual(state["masks"].dtype, torch.bool)
+ self.assertEqual(state["scores"].shape, (2,))
+
+
+class FilterContainedSubMasksTest(absltest.TestCase):
+ """Tests for _filter_contained_sub_masks."""
+
+ def test_empty_state_unchanged(self):
+ """Verifies a zero-mask state passes through untouched."""
+ detector = _make_detector()
+ state = {
+ "masks": torch.zeros((0, 1, 10, 10), dtype=torch.bool),
+ "boxes": torch.zeros((0, 4)),
+ "scores": torch.zeros((0,)),
+ }
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 0)
+
+ def test_drops_contained_smaller_mask(self):
+ """Verifies a small mask fully inside a large one is removed."""
+ detector = _make_detector(containment_threshold=0.9)
+ big = _box_mask(20, 20, 0, 20, 0, 20)
+ small = _box_mask(20, 20, 5, 10, 5, 10)
+ state = {
+ "masks": torch.cat([big.unsqueeze(0), small.unsqueeze(0)], dim=0),
+ "boxes": torch.tensor(
+ [[0, 0, 20, 20], [5, 5, 10, 10]], dtype=torch.float32
+ ),
+ "scores": torch.tensor([0.9, 0.8]),
+ }
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 1)
+ self.assertAlmostEqual(result["scores"].item(), 0.9, places=5)
+
+ def test_keeps_disjoint_masks(self):
+ """Verifies two non-overlapping masks are both kept."""
+ detector = _make_detector(containment_threshold=0.9)
+ left = _box_mask(20, 20, 0, 10, 0, 5)
+ right = _box_mask(20, 20, 0, 10, 15, 20)
+ state = {
+ "masks": torch.cat([left.unsqueeze(0), right.unsqueeze(0)], dim=0),
+ "boxes": torch.tensor(
+ [[0, 0, 5, 10], [15, 0, 20, 10]], dtype=torch.float32
+ ),
+ "scores": torch.tensor([0.9, 0.8]),
+ }
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 2)
+
+
+class MergeContainedBoxesTest(absltest.TestCase):
+ """Tests for _merge_contained_boxes."""
+
+ def test_empty_state_unchanged(self):
+ """Verifies a zero-detection state passes through untouched."""
+ detector = _make_detector()
+ state = {
+ "masks": torch.zeros((0, 1, 10, 10), dtype=torch.bool),
+ "boxes": torch.zeros((0, 4)),
+ "scores": torch.zeros((0,)),
+ }
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["scores"].shape[0], 0)
+
+ def test_merges_contained_box_and_clamps_score(self):
+ """Verifies a contained box merges and the summed score is clamped."""
+ detector = _make_detector(merge_containment_threshold=0.7)
+ big = _box_mask(50, 50, 0, 40, 0, 40)
+ small = _box_mask(50, 50, 5, 15, 5, 15)
+ state = {
+ "masks": torch.cat([big.unsqueeze(0), small.unsqueeze(0)], dim=0),
+ "boxes": torch.tensor(
+ [[0, 0, 40, 40], [5, 5, 15, 15]], dtype=torch.float32
+ ),
+ "scores": torch.tensor([0.8, 0.9]),
+ }
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["masks"].shape[0], 1)
+ self.assertEqual(result["masks"].ndim, 4)
+ self.assertAlmostEqual(result["scores"][0].item(), 1.0, places=5)
+ torch.testing.assert_close(
+ result["boxes"][0], torch.tensor([0.0, 0.0, 40.0, 40.0])
+ )
+
+ def test_disjoint_boxes_not_merged(self):
+ """Verifies non-contained boxes stay separate."""
+ detector = _make_detector(merge_containment_threshold=0.7)
+ left = _box_mask(50, 50, 0, 10, 0, 10)
+ right = _box_mask(50, 50, 30, 40, 30, 40)
+ state = {
+ "masks": torch.cat([left.unsqueeze(0), right.unsqueeze(0)], dim=0),
+ "boxes": torch.tensor(
+ [[0, 0, 10, 10], [30, 30, 40, 40]], dtype=torch.float32
+ ),
+ "scores": torch.tensor([0.6, 0.7]),
+ }
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["masks"].shape[0], 2)
+
+
+class ToSupervisionDetectionsTest(absltest.TestCase):
+ """Tests for to_supervision_detections."""
+
+ def test_applies_score_threshold(self):
+ """Verifies detections below the score threshold are dropped."""
+ detector = _make_detector(score_threshold=0.5)
+ state = {
+ "masks": torch.zeros((2, 1, 10, 10), dtype=torch.bool),
+ "boxes": torch.tensor(
+ [[0, 0, 5, 5], [1, 1, 6, 6]], dtype=torch.float32
+ ),
+ "scores": torch.tensor([0.4, 0.9], dtype=torch.float32),
+ }
+ fake_detections = mock.Mock()
+ with mock.patch.object(
+ rfdetr_detector.supervision, "Detections"
+ ) as mock_detections:
+ mock_detections.return_value = fake_detections
+ detector.to_supervision_detections(state)
+ # One detection survives the 0.5 threshold.
+ _, kwargs = mock_detections.call_args
+ self.assertEqual(kwargs["xyxy"].shape[0], 1)
+ np.testing.assert_allclose(kwargs["confidence"], np.array([0.9]))
+ # class_id is always zeroed out.
+ self.assertTrue(np.all(kwargs["class_id"] == 0))
+
+ def test_returns_empty_when_all_below_threshold(self):
+ """Verifies an empty Detections is returned when nothing survives."""
+ detector = _make_detector(score_threshold=0.99)
+ state = {
+ "masks": torch.zeros((1, 1, 10, 10), dtype=torch.bool),
+ "boxes": torch.tensor([[0, 0, 5, 5]], dtype=torch.float32),
+ "scores": torch.tensor([0.4], dtype=torch.float32),
+ }
+ sentinel = mock.Mock()
+ with mock.patch.object(
+ rfdetr_detector.supervision, "Detections"
+ ) as mock_detections:
+ mock_detections.empty.return_value = sentinel
+ result = detector.to_supervision_detections(state)
+ mock_detections.empty.assert_called_once()
+ self.assertIs(result, sentinel)
+
+
+class DetectTest(absltest.TestCase):
+ """Tests for the detect orchestration method."""
+
+ def test_runs_predict_then_filters(self):
+ """Verifies detect calls predict and returns a filtered state dict."""
+ detector = _make_detector()
+ image = mock.Mock()
+ image.size = (30, 20) # (width, height)
+
+ raw_detections = mock.Mock()
+ detector._model.predict = mock.Mock(return_value=raw_detections)
+
+ built_state = {
+ "masks": torch.zeros((0, 1, 20, 30), dtype=torch.bool),
+ "boxes": torch.zeros((0, 4)),
+ "scores": torch.zeros((0,)),
+ }
+ with mock.patch.object(
+ detector, "_convert_detections_to_state", return_value=built_state
+ ), mock.patch.object(
+ detector, "_filter_contained_sub_masks", side_effect=lambda s: s
+ ) as mock_filter, mock.patch.object(
+ detector, "_merge_contained_boxes", side_effect=lambda s: s
+ ) as mock_merge:
+ result = detector.detect(image)
+
+ detector._model.predict.assert_called_once()
+ mock_filter.assert_called_once()
+ mock_merge.assert_called_once()
+ self.assertIs(result, built_state)
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/track_manager.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/track_manager.py
new file mode 100644
index 00000000000..3604a2ab628
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/track_manager.py
@@ -0,0 +1,338 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tracking, image blending, and track-level classification voting logic."""
+
+import collections
+import math
+from typing import Any
+
+import cv2
+import numpy as np
+from PIL import Image
+import torch
+import tqdm
+import trackers
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import config_loader
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import dinov3_classifier
+
+
+class TrackManager:
+ """Manages object trajectories, blended crop extraction, and final label voting."""
+
+ def __init__(
+ self,
+ tracking_config: config_loader.TrackingConfig,
+ cropping_config: config_loader.CroppingConfig,
+ vis_config: config_loader.VisualizationConfig,
+ tracker: Any | None = None,
+ ) -> None:
+ """Initializes TrackManager.
+
+ Args:
+ tracking_config: Configuration for tracking parameters.
+ cropping_config: Configuration for crop size and padding buffer.
+ vis_config: Configuration for visualization and background blend color.
+ tracker: Optional tracker instance to inject. If None, an instance of
+ ByteTrackTracker is constructed from tracking_config.
+ """
+ self._tracking_enabled = tracking_config.enable
+ if tracker is not None:
+ self._tracker = tracker
+ else:
+ self._tracker = trackers.ByteTrackTracker(
+ minimum_iou_threshold=tracking_config.bytetrack_minimum_iou_threshold,
+ minimum_consecutive_frames=tracking_config.bytetrack_minimum_consecutive_frames,
+ )
+ self._crop_size = cropping_config.crop_size
+ self._buffer = cropping_config.crop_buffer_pixels
+ self._bg_color = vis_config.background_blend_color_rgb
+ self._crop_records: dict[int, list[dict[str, Any]]] = {}
+ self._next_standalone_id = 1
+
+ def reset(self) -> None:
+ """Resets tracker ID counter and clears all collected crop records.
+
+ Call this when switching to a new video or scene so that track IDs in
+ the next session start from 1 and no crops carry over from the
+ previous session.
+ """
+ self._tracker.reset()
+ self._crop_records.clear()
+ self._next_standalone_id = 1
+
+ def update_and_extract_crops(
+ self,
+ detections: Any,
+ state: dict[str, torch.Tensor],
+ image: Image.Image,
+ frame_name: str,
+ ) -> tuple[Any, list[float]]:
+ """Updates tracks and extracts image crops.
+
+ When tracking is disabled in the config, ByteTrack is bypassed
+ entirely and every detection in the current frame is assigned a
+ fresh sequential ID. The returned `unassigned_scores` list is always
+ empty in that mode because no detection is ever dropped by the
+ tracker.
+
+ Args:
+ detections: Detection object (e.g. supervision.Detections) for the frame.
+ state: Detector state mapping containing bounding boxes and masks.
+ image: Full-frame PIL image.
+ frame_name: Identifier or filename for the current frame.
+
+ Returns:
+ A tuple of (detections, unassigned_scores).
+ """
+ if self._tracking_enabled:
+ return self._update_with_tracking(detections, state, image, frame_name)
+ return self._update_without_tracking(detections, state, image, frame_name)
+
+ def _update_with_tracking(
+ self,
+ detections: Any,
+ state: dict[str, torch.Tensor],
+ image: Image.Image,
+ frame_name: str,
+ ) -> tuple[Any, list[float]]:
+ """Runs ByteTrack and extracts crops for every successfully assigned detection."""
+ pre_update_scores = (
+ detections.confidence.copy() if len(detections) > 0 else None
+ )
+ detections = self._tracker.update(detections)
+
+ unassigned_scores = []
+ if pre_update_scores is not None and detections.tracker_id is not None:
+ unassigned_scores = [
+ float(score)
+ for score, tid in zip(pre_update_scores, detections.tracker_id)
+ if int(tid) == -1
+ ]
+
+ if len(detections) == 0 or detections.tracker_id is None:
+ return detections, unassigned_scores
+
+ state_boxes = state["boxes"].numpy().astype(np.float32)
+ det_boxes = detections.xyxy.astype(np.float32)
+
+ for row_idx in range(len(detections)):
+ tracker_id = int(detections.tracker_id[row_idx])
+ if tracker_id == -1:
+ continue
+
+ matching_rows = np.where(
+ np.all(np.isclose(state_boxes, det_boxes[row_idx], atol=1e-3), axis=1)
+ )[0]
+
+ if matching_rows.size == 0:
+ continue
+
+ crop = self._extract_blended_crop(image, state, int(matching_rows[0]))
+ self._crop_records.setdefault(tracker_id, []).append({
+ "frame_name": frame_name,
+ "crop": crop,
+ })
+
+ return detections, unassigned_scores
+
+ def _update_without_tracking(
+ self,
+ detections: Any,
+ state: dict[str, torch.Tensor],
+ image: Image.Image,
+ frame_name: str,
+ ) -> tuple[Any, list[float]]:
+ """Assigns a fresh sequential ID to every detection and extracts its crop.
+
+ ByteTrack is bypassed entirely. Detection box order is assumed to
+ match the detector state row order (they originate from the same
+ `to_supervision_detections` call upstream), so the mask for each
+ detection is looked up positionally.
+
+ Args:
+ detections: Detection object (e.g. supervision.Detections) for the frame.
+ state: Detector state mapping containing bounding boxes and masks.
+ image: Full-frame PIL image.
+ frame_name: Identifier or filename for the current frame.
+
+ Returns:
+ A tuple of (detections, unassigned_scores).
+ """
+ if len(detections) == 0:
+ return detections, []
+
+ num_detections = len(detections)
+ assigned_ids = np.arange(
+ self._next_standalone_id,
+ self._next_standalone_id + num_detections,
+ dtype=int,
+ )
+ detections.tracker_id = assigned_ids
+ self._next_standalone_id += num_detections
+
+ for row_idx in range(num_detections):
+ tracker_id = int(assigned_ids[row_idx])
+ crop = self._extract_blended_crop(image, state, row_idx)
+ self._crop_records.setdefault(tracker_id, []).append({
+ "frame_name": frame_name,
+ "crop": crop,
+ })
+
+ return detections, []
+
+ def classify_all_tracks(
+ self, classifier: dinov3_classifier.DINOv3Classifier, batch_size: int
+ ) -> dict[int, list[dict[str, Any]]]:
+ """Runs batch classification for every collected track crop with a progress bar."""
+ track_predictions = {}
+ total_crops = sum(len(records) for records in self._crop_records.values())
+
+ progress_bar = tqdm.tqdm(
+ total=total_crops, desc="Classifying crops", unit="crop"
+ )
+
+ for tracker_id, records in self._crop_records.items():
+ progress_bar.set_postfix_str(f"track {tracker_id:04d}")
+ per_crop_preds = []
+
+ chunks = [
+ records[i : i + batch_size]
+ for i in range(0, len(records), batch_size)
+ ]
+ for chunk in chunks:
+ pil_images = [r["crop"] for r in chunk]
+ chunk_preds = classifier.predict_batch(pil_images)
+
+ for record, pred in zip(chunk, chunk_preds):
+ per_crop_preds.append({
+ "frame_name": record["frame_name"],
+ "crop": record["crop"],
+ **pred,
+ })
+ progress_bar.update(len(chunk))
+
+ track_predictions[tracker_id] = per_crop_preds
+
+ progress_bar.close()
+ return track_predictions
+
+ def resolve_track_label(
+ self, per_crop_predictions: list[dict[str, Any]]
+ ) -> tuple[str, int]:
+ """Resolves a final class by majority vote, tie-breaking by confidence sum."""
+ vote_counter = collections.Counter(
+ pred["predicted_class"] for pred in per_crop_predictions
+ )
+ confidence_sums = collections.defaultdict(float)
+
+ for pred in per_crop_predictions:
+ confidence_sums[pred["predicted_class"]] += pred[
+ "predicted_probability_percent"
+ ]
+
+ highest_votes = max(vote_counter.values())
+ top_classes = [
+ cls for cls, votes in vote_counter.items() if votes == highest_votes
+ ]
+
+ if len(top_classes) == 1:
+ return top_classes[0], highest_votes
+
+ highest_conf = max(confidence_sums[cls] for cls in top_classes)
+ tied_conf_classes = sorted([
+ cls
+ for cls in top_classes
+ if math.isclose(confidence_sums[cls], highest_conf)
+ ])
+
+ return tied_conf_classes[0], highest_votes
+
+ def _extract_blended_crop(
+ self, image: Image.Image, state: dict[str, torch.Tensor], idx: int
+ ) -> Image.Image:
+ """Generates a soft-edged crop merged against the ImageNet mean background."""
+ img_arr = np.array(image)
+
+ # FIX: Explicitly convert the PyTorch tensor to a NumPy array before
+ # manipulating.
+ raw_mask_array = state["masks"][idx].numpy()
+ mask = self._fill_mask_holes(np.squeeze(raw_mask_array))
+
+ box = state["boxes"][idx].tolist()
+
+ h, w = mask.shape
+ x_min, y_min, x_max, y_max = [round(v) for v in box]
+ x_min, y_min = max(0, x_min - self._buffer), max(0, y_min - self._buffer)
+ x_max, y_max = min(w, x_max + self._buffer), min(h, y_max + self._buffer)
+
+ roi_img = img_arr[y_min:y_max, x_min:x_max]
+ roi_mask = mask[y_min:y_max, x_min:x_max].astype(np.uint8) * 255
+
+ dilated = cv2.dilate(roi_mask, np.ones((5, 5), np.uint8), iterations=1)
+ alpha = (
+ cv2.GaussianBlur(dilated, (5, 5), 0).astype(np.float32)[
+ :, :, np.newaxis
+ ]
+ / 255.0
+ )
+ bg = np.array(self._bg_color, dtype=np.float32)
+
+ blended = (roi_img.astype(np.float32) * alpha + bg * (1.0 - alpha)).astype(
+ np.uint8
+ )
+
+ canvas = np.full(
+ (self._crop_size[0], self._crop_size[1], 3),
+ self._bg_color,
+ dtype=np.uint8,
+ )
+ ch, cw = canvas.shape[:2]
+ bh, bw = blended.shape[:2]
+ scale = min(cw / bw, ch / bh)
+ rw, rh = int(bw * scale), int(bh * scale)
+
+ resized = cv2.resize(blended, (rw, rh), interpolation=cv2.INTER_LINEAR)
+ ox, oy = (cw - rw) // 2, (ch - rh) // 2
+ canvas[oy : oy + rh, ox : ox + rw] = resized
+
+ return Image.fromarray(canvas)
+
+ def _fill_mask_holes(self, mask: np.ndarray) -> np.ndarray:
+ """Fills interior mask holes via flood-fill."""
+ # Extra safety check to ensure it operates strictly as a numpy array
+ mask_u8 = np.asarray(mask).astype(np.uint8) * 255
+ h, w = mask_u8.shape
+ padded = np.zeros((h + 2, w + 2), dtype=np.uint8)
+ padded[1 : h + 1, 1 : w + 1] = mask_u8
+
+ flood_filled = padded.copy()
+ cv2.floodFill(flood_filled, mask=None, seedPoint=(0, 0), newVal=255)
+ holes = cv2.bitwise_not(flood_filled[1 : h + 1, 1 : w + 1])
+ return cv2.bitwise_or(mask_u8, holes).astype(bool)
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/track_manager_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/track_manager_test.py
new file mode 100644
index 00000000000..bd9c1b9b6ec
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/track_manager_test.py
@@ -0,0 +1,379 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for track_manager.py.
+
+ByteTrackTracker is mocked/injected at construction so no real tracker is built.
+The label-voting, hole-fill, blended-crop, and no-tracking assignment logic are
+exercised directly.
+"""
+
+import sys
+from typing import Any
+from unittest import mock
+
+from absl.testing import absltest
+import numpy as np
+from PIL import Image
+import torch
+
+# Mock trackers before it is imported anywhere since it is an external
+# pip package not checked into //third_party/py.
+mock_trackers = mock.MagicMock()
+
+
+class MockByteTrackTracker:
+
+ def __init__(
+ self,
+ minimum_iou_threshold: float = 0.1,
+ minimum_consecutive_frames: int = 2,
+ ) -> None:
+ pass
+
+ def update(self, detections: Any) -> Any:
+ return detections
+
+ def reset(self) -> None:
+ pass
+
+
+mock_trackers.ByteTrackTracker = MockByteTrackTracker
+sys.modules.setdefault("trackers", mock_trackers)
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import track_manager # pylint: disable=g-bad-import-order, g-import-not-at-top
+
+
+class _FakeDetections:
+ """Minimal stand-in for supervision.Detections used by the no-track path."""
+
+ def __init__(self, count: int) -> None:
+ self._count = count
+ self.tracker_id = None
+ self.confidence = np.full((count,), 0.9, dtype=np.float32)
+ self.xyxy = np.zeros((count, 4), dtype=np.float32)
+
+ def __len__(self) -> int:
+ return self._count
+
+
+def _make_manager(
+ tracking_enabled: bool = True,
+ crop_size: tuple[int, int] = (64, 64),
+ crop_buffer_pixels: int = 5,
+ tracker: Any | None = None,
+) -> track_manager.TrackManager:
+ """Builds a TrackManager with ByteTrackTracker injected or mocked."""
+ tracking_config = mock.Mock(
+ enable=tracking_enabled,
+ bytetrack_minimum_iou_threshold=0.1,
+ bytetrack_minimum_consecutive_frames=2,
+ )
+ cropping_config = mock.Mock(
+ crop_size=crop_size, crop_buffer_pixels=crop_buffer_pixels
+ )
+ vis_config = mock.Mock(background_blend_color_rgb=(124, 116, 104))
+ if tracker is None:
+ tracker = mock.create_autospec(MockByteTrackTracker, instance=True)
+ return track_manager.TrackManager(
+ tracking_config=tracking_config,
+ cropping_config=cropping_config,
+ vis_config=vis_config,
+ tracker=tracker,
+ )
+
+
+def _pred(class_name: str, probability: float) -> dict[str, Any]:
+ """Builds a single per-crop prediction dict for the voting tests."""
+ return {
+ "predicted_class": class_name,
+ "predicted_probability_percent": probability,
+ }
+
+
+class ResolveTrackLabelTest(absltest.TestCase):
+ """Tests for resolve_track_label (majority vote + tie-breaks)."""
+
+ def setUp(self):
+ super().setUp()
+ self.manager = _make_manager()
+
+ def test_clear_majority_wins(self):
+ """Verifies the class with the most votes wins."""
+ predictions = [
+ _pred("a", 90.0),
+ _pred("a", 80.0),
+ _pred("b", 99.0),
+ ]
+ label, votes = self.manager.resolve_track_label(predictions)
+ self.assertEqual(label, "a")
+ self.assertEqual(votes, 2)
+
+ def test_vote_tie_broken_by_confidence_sum(self):
+ """Verifies a vote tie is broken by the higher confidence sum."""
+ predictions = [
+ _pred("a", 60.0),
+ _pred("b", 95.0),
+ ]
+ label, votes = self.manager.resolve_track_label(predictions)
+ self.assertEqual(label, "b")
+ self.assertEqual(votes, 1)
+
+ def test_full_tie_broken_alphabetically(self):
+ """Verifies a vote and confidence tie is broken alphabetically."""
+ predictions = [
+ _pred("banana", 50.0),
+ _pred("apple", 50.0),
+ ]
+ label, _ = self.manager.resolve_track_label(predictions)
+ self.assertEqual(label, "apple")
+
+
+class FillMaskHolesTest(absltest.TestCase):
+ """Tests for _fill_mask_holes."""
+
+ def setUp(self):
+ super().setUp()
+ self.manager = _make_manager()
+
+ def test_fills_interior_hole(self):
+ """Verifies an interior hole is filled and dtype is bool."""
+ mask = np.zeros((30, 30), dtype=bool)
+ mask[5:25, 5:25] = True
+ mask[12:18, 12:18] = False
+ filled = self.manager._fill_mask_holes(mask)
+ self.assertEqual(filled.dtype, np.bool_)
+ self.assertTrue(filled[15, 15])
+
+ def test_leaves_solid_mask_unchanged(self):
+ """Verifies a hole-free mask is preserved."""
+ mask = np.zeros((10, 10), dtype=bool)
+ mask[2:8, 2:8] = True
+ filled = self.manager._fill_mask_holes(mask)
+ np.testing.assert_array_equal(filled, mask)
+
+
+class ExtractBlendedCropTest(absltest.TestCase):
+ """Tests for _extract_blended_crop."""
+
+ def test_returns_letterboxed_pil_image(self):
+ """Verifies the blended crop is a PIL image of the configured size."""
+ manager = _make_manager(crop_size=(64, 64))
+ image = Image.new("RGB", (50, 50), (10, 20, 30))
+ mask = torch.zeros((1, 1, 50, 50), dtype=torch.bool)
+ mask[0, 0, 10:40, 10:40] = True
+ state = {
+ "masks": mask,
+ "boxes": torch.tensor([[0, 0, 50, 50]], dtype=torch.float32),
+ }
+ crop = manager._extract_blended_crop(image, state, idx=0)
+ self.assertIsInstance(crop, Image.Image)
+ self.assertEqual(crop.size, (64, 64))
+
+
+class TrackerInjectionTest(absltest.TestCase):
+ """Tests for tracker dependency injection."""
+
+ def test_custom_tracker_injected(self):
+ """Verifies custom tracker instance is used when provided."""
+ custom_tracker = mock.Mock()
+ manager = _make_manager(tracker=custom_tracker)
+ self.assertIs(manager._tracker, custom_tracker)
+
+ def test_default_tracker_constructed_from_config(self):
+ """Verifies default tracker is constructed when tracker is None."""
+ tracking_config = mock.Mock(
+ enable=True,
+ bytetrack_minimum_iou_threshold=0.2,
+ bytetrack_minimum_consecutive_frames=3,
+ )
+ cropping_config = mock.Mock(crop_size=(64, 64), crop_buffer_pixels=5)
+ vis_config = mock.Mock(background_blend_color_rgb=(124, 116, 104))
+ manager = track_manager.TrackManager(
+ tracking_config=tracking_config,
+ cropping_config=cropping_config,
+ vis_config=vis_config,
+ )
+ self.assertIsInstance(manager._tracker, MockByteTrackTracker)
+
+
+class ResetTest(absltest.TestCase):
+ """Tests for reset."""
+
+ def test_clears_records_and_resets_counter(self):
+ """Verifies reset clears crop records and the standalone-ID counter."""
+ manager = _make_manager()
+ manager._crop_records[5] = [{"frame_name": "f", "crop": None}]
+ manager._next_standalone_id = 99
+ manager.reset()
+ self.assertEmpty(manager._crop_records)
+ self.assertEqual(manager._next_standalone_id, 1)
+ manager._tracker.reset.assert_called_once()
+
+
+class UpdateWithoutTrackingTest(absltest.TestCase):
+ """Tests for the no-tracking assignment path."""
+
+ def test_assigns_sequential_ids_and_records_crops(self):
+ """Verifies each detection gets a fresh sequential ID and a crop record."""
+ manager = _make_manager(tracking_enabled=False)
+ detections = _FakeDetections(count=2)
+ state = {
+ "masks": torch.zeros((2, 1, 20, 20), dtype=torch.bool),
+ "boxes": torch.zeros((2, 4), dtype=torch.float32),
+ }
+ with mock.patch.object(
+ manager, "_extract_blended_crop", return_value="CROP"
+ ):
+ returned, unassigned = manager.update_and_extract_crops(
+ detections, state, Image.new("RGB", (20, 20)), "frame_0.png"
+ )
+ np.testing.assert_array_equal(returned.tracker_id, np.array([1, 2]))
+ self.assertEqual(unassigned, [])
+ self.assertIn(1, manager._crop_records)
+ self.assertIn(2, manager._crop_records)
+
+ def test_ids_continue_across_calls(self):
+ """Verifies standalone IDs keep incrementing across frames."""
+ manager = _make_manager(tracking_enabled=False)
+ state = {
+ "masks": torch.zeros((1, 1, 20, 20), dtype=torch.bool),
+ "boxes": torch.zeros((1, 4), dtype=torch.float32),
+ }
+ with mock.patch.object(
+ manager, "_extract_blended_crop", return_value="CROP"
+ ):
+ manager.update_and_extract_crops(
+ _FakeDetections(1), state, Image.new("RGB", (20, 20)), "f0.png"
+ )
+ returned, _ = manager.update_and_extract_crops(
+ _FakeDetections(1), state, Image.new("RGB", (20, 20)), "f1.png"
+ )
+ np.testing.assert_array_equal(returned.tracker_id, np.array([2]))
+
+ def test_empty_detections_return_empty(self):
+ """Verifies zero detections produce no records and no unassigned scores."""
+ manager = _make_manager(tracking_enabled=False)
+ state = {
+ "masks": torch.zeros((0, 1, 20, 20), dtype=torch.bool),
+ "boxes": torch.zeros((0, 4), dtype=torch.float32),
+ }
+ returned, unassigned = manager.update_and_extract_crops(
+ _FakeDetections(0), state, Image.new("RGB", (20, 20)), "f.png"
+ )
+ self.assertEmpty(returned)
+ self.assertEqual(unassigned, [])
+ self.assertEmpty(manager._crop_records)
+
+
+class UpdateWithTrackingTest(absltest.TestCase):
+ """Tests for the ByteTrack-enabled path."""
+
+ def test_records_crops_for_assigned_tracks(self):
+ """Verifies assigned detections get crop records keyed by tracker_id."""
+ manager = _make_manager(tracking_enabled=True)
+
+ boxes = np.array([[0, 0, 5, 5], [6, 6, 9, 9]], dtype=np.float32)
+ tracked = _FakeDetections(count=2)
+ tracked.xyxy = boxes
+ tracked.tracker_id = np.array([10, 11])
+ manager._tracker.update = mock.Mock(return_value=tracked)
+
+ incoming = _FakeDetections(count=2)
+ incoming.xyxy = boxes
+ state = {
+ "masks": torch.zeros((2, 1, 20, 20), dtype=torch.bool),
+ "boxes": torch.from_numpy(boxes),
+ }
+ with mock.patch.object(
+ manager, "_extract_blended_crop", return_value="CROP"
+ ):
+ _, unassigned = manager.update_and_extract_crops(
+ incoming, state, Image.new("RGB", (20, 20)), "frame.png"
+ )
+ self.assertEqual(unassigned, [])
+ self.assertIn(10, manager._crop_records)
+ self.assertIn(11, manager._crop_records)
+
+ def test_unassigned_scores_reported_for_dropped_detections(self):
+ """Verifies detections ByteTrack leaves at id -1 surface as unassigned."""
+ manager = _make_manager(tracking_enabled=True)
+
+ boxes = np.array([[0, 0, 5, 5]], dtype=np.float32)
+ tracked = _FakeDetections(count=1)
+ tracked.xyxy = boxes
+ tracked.tracker_id = np.array([-1])
+ tracked.confidence = np.array([0.42], dtype=np.float32)
+ manager._tracker.update = mock.Mock(return_value=tracked)
+
+ incoming = _FakeDetections(count=1)
+ incoming.xyxy = boxes
+ incoming.confidence = np.array([0.42], dtype=np.float32)
+ state = {
+ "masks": torch.zeros((1, 1, 20, 20), dtype=torch.bool),
+ "boxes": torch.from_numpy(boxes),
+ }
+ _, unassigned = manager.update_and_extract_crops(
+ incoming, state, Image.new("RGB", (20, 20)), "frame.png"
+ )
+ self.assertLen(unassigned, 1)
+ self.assertAlmostEqual(unassigned[0], 0.42, places=5)
+ self.assertEmpty(manager._crop_records)
+
+
+class ClassifyAllTracksTest(absltest.TestCase):
+ """Tests for classify_all_tracks batching."""
+
+ def test_batches_and_merges_predictions(self):
+ """Verifies crops are classified in batches and merged with metadata."""
+ manager = _make_manager()
+ manager._crop_records = {
+ 7: [
+ {"frame_name": "f0.png", "crop": "C0"},
+ {"frame_name": "f1.png", "crop": "C1"},
+ {"frame_name": "f2.png", "crop": "C2"},
+ ]
+ }
+
+ classifier = mock.Mock()
+ # Return one prediction dict per crop passed in.
+ classifier.predict_batch.side_effect = lambda images: [
+ {"predicted_class": "a", "predicted_probability_percent": 90.0}
+ for _ in images
+ ]
+
+ result = manager.classify_all_tracks(classifier, batch_size=2)
+ # 3 crops with batch_size 2 -> two predict_batch calls (2 + 1).
+ self.assertEqual(classifier.predict_batch.call_count, 2)
+ self.assertLen(result[7], 3)
+ self.assertEqual(result[7][0]["frame_name"], "f0.png")
+ self.assertEqual(result[7][0]["predicted_class"], "a")
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/visualization_utils.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/visualization_utils.py
new file mode 100644
index 00000000000..d6862a188ec
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/visualization_utils.py
@@ -0,0 +1,586 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Rendering, plotting, video writing, and summary utilities."""
+
+import collections
+from collections.abc import Callable, Mapping, Sequence
+import dataclasses
+import logging
+import math
+import os
+from typing import Any
+
+import cv2
+import numpy as np
+from PIL import Image
+import tqdm
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import config_loader
+
+try:
+ import supervision # pylint: disable=g-import-not-at-top
+except ModuleNotFoundError:
+
+ class _SupervisionFallback:
+ """Fallback classes when supervision is not installed."""
+
+ class BoxAnnotator:
+
+ def annotate(self, scene, detections):
+ del detections
+ return scene
+
+ class LabelAnnotator:
+
+ def annotate(self, scene, detections, labels):
+ del detections, labels
+ return scene
+
+ class Detections:
+
+ def __len__(self) -> int:
+ return 0
+
+ supervision = _SupervisionFallback()
+
+_LOGGER = logging.getLogger(__name__)
+
+
+@dataclasses.dataclass(frozen=True)
+class TrackSummary:
+ """Resolved summary of a single tracker's predictions.
+
+ Attributes:
+ final_class: The class chosen after aggregating per-crop predictions.
+ category: The collapsed category for `final_class`, or None when the
+ collapsed-category feature is disabled.
+ vote_count: Number of per-crop predictions that voted for `final_class`.
+ output_path: Filesystem path where the track-grid PNG was saved, or a
+ placeholder string when grid saving is disabled.
+ """
+
+ final_class: str
+ category: str | None
+ vote_count: int
+ output_path: str
+
+ def __getitem__(self, item: str) -> Any:
+ return getattr(self, item)
+
+
+class PipelineVisualizer:
+ """Manages file I/O, frame annotation, grid generation, and summary printing."""
+
+ def __init__(
+ self,
+ config: config_loader.VisualizationConfig,
+ collapsed_categories: config_loader.CollapsedCategoriesConfig,
+ out_video_path: str,
+ summary_logger: logging.Logger | None = None,
+ ):
+ """Initializes the visualizer.
+
+ Args:
+ config: Visualization-specific configuration.
+ collapsed_categories: Optional grouping of fine-grained classes into
+ broader categories. When disabled, per-category reporting and folder
+ nesting are skipped.
+ out_video_path: Filesystem path where the annotated MP4 will be written
+ when `save_video` is true.
+ summary_logger: Optional dedicated logger whose handler writes the
+ per-subfolder summary blocks to a single summary file. When provided,
+ every summary line is emitted to both the main logger
+ (console/per-subfolder log) and this logger. When None, summaries are
+ logged only to the main logger.
+ """
+ self._config = config
+ self._collapsed_categories = collapsed_categories
+ self._box_annotator = supervision.BoxAnnotator()
+ self._label_annotator = supervision.LabelAnnotator()
+ self._video_path = out_video_path
+ self._video_writer = None
+ self._summary_logger = summary_logger
+
+ def annotate_and_write_frame(
+ self, image: Image.Image, detections: Any, frame_path: str
+ ) -> None:
+ """Draws bounding boxes onto the frame and saves based on config toggles."""
+ if not self._config.save_frames and not self._config.save_video:
+ return
+
+ frame_bgr = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
+ annotated = self._box_annotator.annotate(
+ scene=frame_bgr.copy(), detections=detections
+ )
+
+ labels = self._build_labels(detections)
+ annotated = self._label_annotator.annotate(
+ scene=annotated, detections=detections, labels=labels
+ )
+
+ if self._config.save_frames:
+ cv2.imwrite(frame_path, annotated)
+
+ if self._config.save_video:
+ if self._video_writer is None:
+ h, w = annotated.shape[:2]
+ fourcc = cv2.VideoWriter_fourcc(*"mp4v")
+ self._video_writer = cv2.VideoWriter(
+ self._video_path, fourcc, self._config.output_video_fps, (w, h)
+ )
+ self._video_writer.write(annotated)
+
+ def close_video(self) -> None:
+ """Releases the video writer buffer."""
+ if self._video_writer:
+ self._video_writer.release()
+
+ def save_track_grids(
+ self,
+ track_predictions: Mapping[int, Sequence[Mapping[str, Any]]],
+ resolve_fn: Callable[[Sequence[Mapping[str, Any]]], tuple[str, int]],
+ out_dir: str,
+ ) -> dict[int, TrackSummary]:
+ """Resolves labels and renders high-speed NumPy/OpenCV prediction grids.
+
+ Output folder structure depends on whether collapsed categories are
+ enabled:
+ - Enabled:
+ out_dir///track_NNNN_.png
+ - Disabled: out_dir//track_NNNN_.png
+
+ Args:
+ track_predictions: Mapping from tracker_id to its per-crop prediction
+ list.
+ resolve_fn: Callable returning (final_class, vote_count) for a list of
+ per-crop predictions.
+ out_dir: Root directory under which track-grid PNGs are saved.
+
+ Returns:
+ Mapping from tracker_id to a TrackSummary dataclass containing
+ `final_class`, `category`, `vote_count`, and `output_path`.
+ `category` is None when the feature is disabled.
+ """
+ summary: dict[int, TrackSummary] = {}
+ sorted_ids = sorted(track_predictions.keys())
+
+ desc = (
+ "Saving track grids"
+ if self._config.save_track_grids
+ else "Resolving track labels"
+ )
+ progress_bar = tqdm.tqdm(sorted_ids, desc=desc, unit="track")
+
+ # Pre-calculate grid dimensions based on config (convert inches to pixels)
+ tile_size = int(
+ self._config.track_grid_thumbnail_size_inches
+ * self._config.track_grid_dpi
+ )
+ img_size = int(tile_size * 0.75) # 75% of tile for image, 25% for text
+ header_height = 80
+ cols = self._config.track_grid_columns_per_row
+
+ for tracker_id in progress_bar:
+ per_crop_preds = track_predictions[tracker_id]
+ if not per_crop_preds:
+ continue
+
+ final_class, votes = resolve_fn(per_crop_preds)
+ progress_bar.set_postfix_str(f"track {tracker_id:04d} -> {final_class}")
+
+ category = self._collapsed_categories.get_category_for_class(final_class)
+ path = "N/A (grids disabled)"
+
+ if self._config.save_track_grids:
+ count = len(per_crop_preds)
+ num_cols = min(count, cols)
+ num_rows = math.ceil(count / num_cols)
+
+ # Create a blank white canvas for the entire grid
+ grid_w = num_cols * tile_size
+ grid_h = num_rows * tile_size + header_height
+ grid_canvas = np.full((grid_h, grid_w, 3), 255, dtype=np.uint8)
+
+ # Draw the title header
+ title = f"Track {tracker_id} | final: {final_class} ({count} crops)"
+ cv2.putText(
+ grid_canvas,
+ title,
+ (20, 50),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ 1.0,
+ (0, 0, 0),
+ 2,
+ cv2.LINE_AA,
+ )
+
+ # Stamp each crop onto the canvas
+ for idx, pred in enumerate(per_crop_preds):
+ r_idx = idx // num_cols
+ c_idx = idx % num_cols
+
+ x_offset = c_idx * tile_size
+ y_offset = header_height + (r_idx * tile_size)
+
+ # Create individual white tile
+ tile = np.full((tile_size, tile_size, 3), 255, dtype=np.uint8)
+
+ # Convert PIL RGB crop to OpenCV BGR and resize it
+ crop_bgr = cv2.cvtColor(np.array(pred["crop"]), cv2.COLOR_RGB2BGR)
+ crop_resized = cv2.resize(
+ crop_bgr, (img_size, img_size), interpolation=cv2.INTER_LINEAR
+ )
+
+ # Center the image horizontally in the tile
+ img_x_offset = (tile_size - img_size) // 2
+ tile[10 : img_size + 10, img_x_offset : img_x_offset + img_size] = (
+ crop_resized
+ )
+
+ # Draw text annotations below the image
+ text_y = img_size + 40
+ cv2.putText(
+ tile,
+ pred["frame_name"],
+ (15, text_y),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ 0.5,
+ (0, 0, 0),
+ 1,
+ cv2.LINE_AA,
+ )
+ cv2.putText(
+ tile,
+ pred["predicted_class"],
+ (15, text_y + 25),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ 0.45,
+ (0, 0, 0),
+ 1,
+ cv2.LINE_AA,
+ )
+ cv2.putText(
+ tile,
+ f"{pred['predicted_probability_percent']:.2f}%",
+ (15, text_y + 50),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ 0.5,
+ (0, 0, 0),
+ 1,
+ cv2.LINE_AA,
+ )
+
+ # Stamp the tile into the main grid canvas
+ grid_canvas[
+ y_offset : y_offset + tile_size, x_offset : x_offset + tile_size
+ ] = tile
+
+ class_dir = self._build_class_output_directory(
+ out_dir, category, final_class
+ )
+ os.makedirs(class_dir, exist_ok=True)
+ path = os.path.join(
+ class_dir, f"track_{tracker_id:04d}_{final_class}.png"
+ )
+
+ # Save the final numpy array directly to disk
+ cv2.imwrite(path, grid_canvas)
+
+ summary[tracker_id] = TrackSummary(
+ final_class=final_class,
+ category=category,
+ vote_count=votes,
+ output_path=path,
+ )
+
+ progress_bar.close()
+ return summary
+
+ def _log_summary_line(self, message: str, *args) -> None:
+ """Logs one summary line to both the main logger and the summary file.
+
+ The message is always emitted through the main logger (console and the
+ per-subfolder log file). When a summary logger was supplied, the same
+ fully-formatted line is also written to the dedicated summary file.
+
+ Args:
+ message: A printf-style logging message.
+ *args: Arguments interpolated into `message`.
+ """
+ _LOGGER.info(message, *args)
+ if self._summary_logger is not None:
+ self._summary_logger.info(message % args if args else message)
+
+ def print_summary(
+ self,
+ track_summary: Mapping[int, TrackSummary | Mapping[str, Any]],
+ input_directory: str,
+ class_names: list[str],
+ ) -> None:
+ """Logs per-class object counts and ground-truth class accuracy.
+
+ Every line is written both to the main logger (console and the
+ per-subfolder log) and, when a summary logger was supplied at
+ construction time, to the dedicated summary file. A blank separator
+ line is written to the summary file before each block so stacked
+ per-subfolder blocks stay readable.
+
+ When collapsed categories are enabled, also logs per-category counts
+ and ground-truth category accuracy. The ground-truth category is
+ derived from the ground-truth class via the collapsed-category
+ mapping, not matched from the subfolder name.
+
+ Args:
+ track_summary: Mapping from tracker_id to its resolved summary dict or
+ TrackSummary.
+ input_directory: Path to the subfolder that produced this summary.
+ class_names: Full list of class labels from the config.
+ """
+ if self._summary_logger is not None:
+ self._summary_logger.info("")
+
+ class_counts: collections.Counter[str] = collections.Counter(
+ s.final_class if isinstance(s, TrackSummary) else s["final_class"]
+ for s in track_summary.values()
+ )
+ total_objects = sum(class_counts.values())
+
+ ground_truth_class = self._infer_ground_truth_class(
+ input_directory, class_names
+ )
+
+ self._log_summary_line("Folder name: %s", os.path.basename(input_directory))
+ self._log_summary_line("Total tracked objects: %d", total_objects)
+
+ self._log_summary_line("By class:")
+ for class_name, count in class_counts.most_common():
+ self._log_summary_line(" %s: %d", class_name, count)
+
+ self._log_class_accuracy(
+ class_counts=class_counts,
+ total_objects=total_objects,
+ ground_truth_class=ground_truth_class,
+ )
+
+ if self._collapsed_categories.enable:
+ self._log_collapsed_category_section(
+ track_summary=track_summary,
+ total_objects=total_objects,
+ ground_truth_class=ground_truth_class,
+ )
+
+ def _log_class_accuracy(
+ self,
+ class_counts: collections.Counter[str],
+ total_objects: int,
+ ground_truth_class: str | None,
+ ) -> None:
+ """Logs the ground-truth class accuracy line.
+
+ Args:
+ class_counts: Counts of tracked objects per final class.
+ total_objects: Total number of tracked objects.
+ ground_truth_class: The class inferred from the subfolder name by exact
+ match, or None when it could not be inferred.
+ """
+ if ground_truth_class is None:
+ self._log_summary_line(
+ " Class accuracy: N/A (could not infer ground-truth class from"
+ " subfolder name)"
+ )
+ elif total_objects == 0:
+ self._log_summary_line(" Class accuracy: N/A (no tracked objects)")
+ else:
+ class_accuracy = class_counts.get(ground_truth_class, 0) / total_objects
+ self._log_summary_line(
+ " Class accuracy (vs %s): %.2f%%",
+ ground_truth_class,
+ class_accuracy * 100,
+ )
+
+ def _log_collapsed_category_section(
+ self,
+ track_summary: Mapping[int, TrackSummary | Mapping[str, Any]],
+ total_objects: int,
+ ground_truth_class: str | None,
+ ) -> None:
+ """Logs the 'By collapsed categories' counts and accuracy line.
+
+ The ground-truth category is derived from the ground-truth class via
+ the collapsed-category mapping (rather than matched from the subfolder
+ name), so class and category ground truth always agree.
+
+ Args:
+ track_summary: Mapping from tracker_id to its resolved summary dict or
+ TrackSummary.
+ total_objects: Total number of tracked objects.
+ ground_truth_class: The class inferred from the subfolder name, or None
+ when it could not be inferred.
+ """
+ category_counts: collections.Counter[str] = collections.Counter(
+ s.category if isinstance(s, TrackSummary) else s["category"]
+ for s in track_summary.values()
+ )
+
+ self._log_summary_line("By collapsed categories:")
+ for category_name, count in category_counts.most_common():
+ self._log_summary_line(" %s: %d", category_name, count)
+
+ ground_truth_category = self._derive_ground_truth_category(
+ ground_truth_class
+ )
+ if ground_truth_category is None:
+ self._log_summary_line(
+ " Category accuracy: N/A (could not infer ground-truth category "
+ "from subfolder name)"
+ )
+ elif total_objects == 0:
+ self._log_summary_line(" Category accuracy: N/A (no tracked objects)")
+ else:
+ category_accuracy = (
+ category_counts.get(ground_truth_category, 0) / total_objects
+ )
+ self._log_summary_line(
+ " Category accuracy (vs %s): %.2f%%",
+ ground_truth_category,
+ category_accuracy * 100,
+ )
+
+ def _derive_ground_truth_category(
+ self, ground_truth_class: str | None
+ ) -> str | None:
+ """Derives the ground-truth category from the ground-truth class.
+
+ The category is looked up through the collapsed-category mapping
+ rather than matched from the subfolder name, so a subfolder named
+ exactly after a class (e.g. 'brown_bottles_grade3') yields both its
+ class ground truth and, via the mapping, its category ground truth
+ (e.g. 'grade3'). Returns None when the class ground truth could not
+ be inferred.
+
+ Args:
+ ground_truth_class: The class inferred from the subfolder name, or None.
+
+ Returns:
+ The category that contains the ground-truth class, or None.
+ """
+ if ground_truth_class is None:
+ return None
+ return self._collapsed_categories.get_category_for_class(ground_truth_class)
+
+ def _infer_ground_truth_class(
+ self, input_directory: str, class_names: list[str]
+ ) -> str | None:
+ """Returns the configured class whose name exactly equals the subfolder name.
+
+ Matching is a case-insensitive exact comparison against the basename
+ of the input directory. If zero or multiple classes match, the result
+ is ambiguous and None is returned.
+
+ Args:
+ input_directory: Path to the per-subfolder input directory.
+ class_names: Full list of class labels from the config.
+
+ Returns:
+ The matched class name, or None if zero or multiple classes match.
+ """
+ return self._infer_ground_truth_from_names(input_directory, class_names)
+
+ def _infer_ground_truth_from_names(
+ self, input_directory: str, candidate_names: list[str]
+ ) -> str | None:
+ """Returns the candidate whose name exactly equals the subfolder basename.
+
+ Matching is a case-insensitive exact comparison against the basename
+ of the input directory: the subfolder must be named exactly after one
+ of the candidate names. A subfolder whose name merely contains a
+ candidate as a substring (e.g. 'brown_bottles_grade3_batch1') is not
+ treated as a match. If zero or multiple candidates match, returns None
+ to signal an ambiguous result.
+
+ Args:
+ input_directory: Path to the per-subfolder input directory.
+ candidate_names: List of names to match against.
+
+ Returns:
+ The single exactly-matched name, or None if zero or multiple
+ names match.
+ """
+ if not candidate_names:
+ return None
+ subfolder_name = os.path.basename(input_directory).lower()
+ matched_names = [
+ name for name in candidate_names if name.lower() == subfolder_name
+ ]
+ if len(matched_names) == 1:
+ return matched_names[0]
+ return None
+
+ def _build_class_output_directory(
+ self, out_dir: str, category: str | None, final_class: str
+ ) -> str:
+ """Returns the directory in which a track-grid PNG should be saved.
+
+ Args:
+ out_dir: Root output directory for track grids.
+ category: The collapsed category for the final class, or None if the
+ feature is disabled.
+ final_class: The resolved fine-grained class for the track.
+
+ Returns:
+ Filesystem path to the directory where the grid should be saved.
+ """
+ if category is None:
+ return os.path.join(out_dir, final_class)
+ return os.path.join(out_dir, category, final_class)
+
+ def _build_labels(self, detections: Any) -> list[str]:
+ """Formats the display string over bounded objects."""
+ if not detections:
+ return []
+
+ ids = (
+ list(detections.tracker_id)
+ if detections.tracker_id is not None
+ else [None] * len(detections)
+ )
+ confs = (
+ list(detections.confidence)
+ if detections.confidence is not None
+ else [None] * len(detections)
+ )
+
+ labels = []
+ for tid, conf in zip(ids, confs):
+ id_str = "?" if tid is None or int(tid) == -1 else f"ID {int(tid)}"
+ if self._config.show_confidence_in_labels and conf is not None:
+ labels.append(f"{id_str} {float(conf):.2f}")
+ else:
+ labels.append(id_str)
+ return labels
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/visualization_utils_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/visualization_utils_test.py
new file mode 100644
index 00000000000..3c7377be520
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/rfdetr_dinov3_tracking/visualization_utils_test.py
@@ -0,0 +1,287 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for visualization_utils.py.
+
+supervision's annotators are patched at construction so no real annotator is
+built. Ground-truth inference, category derivation, label building, output
+directory routing, and the grids-disabled resolve path are exercised
+directly.
+"""
+
+from unittest import mock
+
+from absl.testing import absltest
+from absl.testing import parameterized
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.rfdetr_dinov3_tracking import visualization_utils
+
+
+class _FakeDetections:
+ """Minimal stand-in for supervision.Detections used by _build_labels."""
+
+ def __init__(self, tracker_id, confidence):
+ self.tracker_id = tracker_id
+ self.confidence = confidence
+
+ def __len__(self) -> int:
+ if self.tracker_id is not None:
+ return len(self.tracker_id)
+ if self.confidence is not None:
+ return len(self.confidence)
+ return 0
+
+
+def _make_visualizer(
+ show_confidence: bool = True,
+ save_track_grids: bool = False,
+ collapsed_categories=None,
+) -> visualization_utils.PipelineVisualizer:
+ """Builds a PipelineVisualizer with supervision annotators patched out."""
+ config = mock.Mock(
+ show_confidence_in_labels=show_confidence,
+ save_frames=False,
+ save_video=False,
+ save_track_grids=save_track_grids,
+ track_grid_thumbnail_size_inches=3,
+ track_grid_dpi=150,
+ track_grid_columns_per_row=5,
+ )
+ if collapsed_categories is None:
+ collapsed_categories = mock.Mock(enable=False)
+ collapsed_categories.get_category_for_class.return_value = None
+ with mock.patch.object(
+ visualization_utils.supervision, "BoxAnnotator", autospec=True
+ ), mock.patch.object(
+ visualization_utils.supervision, "LabelAnnotator", autospec=True
+ ):
+ return visualization_utils.PipelineVisualizer(
+ config=config,
+ collapsed_categories=collapsed_categories,
+ out_video_path="/tmp/out.mp4",
+ summary_logger=None,
+ )
+
+
+class InferGroundTruthFromNamesTest(parameterized.TestCase):
+ """Tests for _infer_ground_truth_from_names."""
+
+ def setUp(self):
+ super().setUp()
+ self.visualizer = _make_visualizer()
+
+ def test_exact_case_insensitive_match(self):
+ """Verifies an exact case-insensitive subfolder match is returned."""
+ result = self.visualizer._infer_ground_truth_from_names(
+ "/data/Brown_Bottles_Grade3", ["brown_bottles_grade3", "clean_grade1"]
+ )
+ self.assertEqual(result, "brown_bottles_grade3")
+
+ def test_substring_is_not_a_match(self):
+ """Verifies a name that merely contains a class is not matched."""
+ result = self.visualizer._infer_ground_truth_from_names(
+ "/data/brown_bottles_grade3_batch1", ["brown_bottles_grade3"]
+ )
+ self.assertIsNone(result)
+
+ def test_no_candidates_returns_none(self):
+ """Verifies an empty candidate list returns None."""
+ result = self.visualizer._infer_ground_truth_from_names(
+ "/data/anything", []
+ )
+ self.assertIsNone(result)
+
+ def test_no_match_returns_none(self):
+ """Verifies a subfolder matching no class returns None."""
+ result = self.visualizer._infer_ground_truth_from_names(
+ "/data/unknown_folder", ["brown_bottles_grade3"]
+ )
+ self.assertIsNone(result)
+
+
+class DeriveGroundTruthCategoryTest(absltest.TestCase):
+ """Tests for _derive_ground_truth_category."""
+
+ def test_none_class_yields_none_category(self):
+ """Verifies a None ground-truth class yields a None category."""
+ visualizer = _make_visualizer()
+ self.assertIsNone(visualizer._derive_ground_truth_category(None))
+
+ def test_derives_category_via_mapping(self):
+ """Verifies the category is looked up from the class via the mapping."""
+ collapsed = mock.Mock(enable=True)
+ collapsed.get_category_for_class.return_value = "grade3"
+ visualizer = _make_visualizer(collapsed_categories=collapsed)
+ result = visualizer._derive_ground_truth_category("brown_bottles_grade3")
+ self.assertEqual(result, "grade3")
+ collapsed.get_category_for_class.assert_called_once_with(
+ "brown_bottles_grade3"
+ )
+
+
+class BuildClassOutputDirectoryTest(absltest.TestCase):
+ """Tests for _build_class_output_directory."""
+
+ def test_flat_layout_when_category_none(self):
+ """Verifies output nests only by class when categories are disabled."""
+ visualizer = _make_visualizer()
+ result = visualizer._build_class_output_directory(
+ "/out", category=None, final_class="clean_grade1"
+ )
+ self.assertEqual(result, "/out/clean_grade1")
+
+ def test_nested_layout_when_category_present(self):
+ """Verifies output nests by category then class when enabled."""
+ visualizer = _make_visualizer()
+ result = visualizer._build_class_output_directory(
+ "/out", category="grade1", final_class="clean_grade1"
+ )
+ self.assertEqual(result, "/out/grade1/clean_grade1")
+
+
+class BuildLabelsTest(absltest.TestCase):
+ """Tests for _build_labels."""
+
+ def test_empty_detections_yield_empty_labels(self):
+ """Verifies no detections produce no labels."""
+ visualizer = _make_visualizer()
+ detections = _FakeDetections(tracker_id=[], confidence=[])
+ self.assertEqual(visualizer._build_labels(detections), [])
+
+ def test_includes_confidence_when_configured(self):
+ """Verifies labels show the id and confidence when enabled."""
+ visualizer = _make_visualizer(show_confidence=True)
+ detections = _FakeDetections(tracker_id=[3], confidence=[0.9])
+ labels = visualizer._build_labels(detections)
+ self.assertEqual(labels, ["ID 3 0.90"])
+
+ def test_omits_confidence_when_disabled(self):
+ """Verifies labels show only the id when confidence is disabled."""
+ visualizer = _make_visualizer(show_confidence=False)
+ detections = _FakeDetections(tracker_id=[3], confidence=[0.9])
+ labels = visualizer._build_labels(detections)
+ self.assertEqual(labels, ["ID 3"])
+
+ def test_unassigned_id_renders_question_mark(self):
+ """Verifies a -1 tracker id renders as '?'."""
+ visualizer = _make_visualizer(show_confidence=False)
+ detections = _FakeDetections(tracker_id=[-1], confidence=[0.5])
+ labels = visualizer._build_labels(detections)
+ self.assertEqual(labels, ["?"])
+
+ def test_none_tracker_ids_render_question_marks(self):
+ """Verifies absent tracker ids render as '?' for every detection."""
+ visualizer = _make_visualizer(show_confidence=False)
+ detections = _FakeDetections(tracker_id=None, confidence=[0.5, 0.6])
+ labels = visualizer._build_labels(detections)
+ self.assertEqual(labels, ["?", "?"])
+
+
+class SaveTrackGridsResolveOnlyTest(absltest.TestCase):
+ """Tests for save_track_grids when grid rendering is disabled.
+
+ With save_track_grids=False the method resolves labels and builds the
+ summary without any OpenCV rendering, so it is exercised end to end here.
+ """
+
+ def test_builds_summary_without_rendering(self):
+ """Verifies the summary carries resolved labels and an 'N/A' path."""
+ collapsed = mock.Mock(enable=True)
+ collapsed.get_category_for_class.return_value = "grade3"
+ visualizer = _make_visualizer(
+ save_track_grids=False, collapsed_categories=collapsed
+ )
+
+ track_predictions = {
+ 1: [{
+ "predicted_class": "dirt_jars_grade3",
+ "predicted_probability_percent": 90.0,
+ }],
+ 2: [], # Empty prediction list is skipped.
+ }
+ resolve_fn = mock.Mock(return_value=("dirt_jars_grade3", 5))
+
+ summary = visualizer.save_track_grids(
+ track_predictions, resolve_fn, out_dir="/out"
+ )
+ self.assertIn(1, summary)
+ self.assertNotIn(2, summary) # empty list produced no entry
+ self.assertIsInstance(summary[1], visualization_utils.TrackSummary)
+ self.assertEqual(summary[1].final_class, "dirt_jars_grade3")
+ self.assertEqual(summary[1].category, "grade3")
+ self.assertEqual(summary[1].vote_count, 5)
+ self.assertEqual(summary[1].output_path, "N/A (grids disabled)")
+ # Also verify mapping subscription for backward compatibility.
+ self.assertEqual(summary[1]["final_class"], "dirt_jars_grade3")
+ self.assertEqual(summary[1]["category"], "grade3")
+ self.assertEqual(summary[1]["vote_count"], 5)
+ self.assertEqual(summary[1]["output_path"], "N/A (grids disabled)")
+
+
+class PrintSummaryTest(absltest.TestCase):
+ """Tests for PipelineVisualizer.print_summary."""
+
+ def test_print_summary_with_collapsed_categories(self):
+ """Verifies print_summary logs without errors when categories are enabled."""
+ collapsed = mock.Mock(enable=True)
+ collapsed.get_category_for_class.return_value = "grade3"
+ visualizer = _make_visualizer(collapsed_categories=collapsed)
+ track_summary = {
+ 1: visualization_utils.TrackSummary(
+ final_class="dirt_jars_grade3",
+ category="grade3",
+ vote_count=5,
+ output_path="N/A",
+ )
+ }
+ visualizer.print_summary(
+ track_summary,
+ input_directory="/data/dirt_jars_grade3",
+ class_names=["dirt_jars_grade3", "clean_grade1"],
+ )
+
+ def test_print_summary_categories_disabled(self):
+ """Verifies print_summary logs without errors when categories are disabled."""
+ visualizer = _make_visualizer()
+ track_summary = {
+ 1: visualization_utils.TrackSummary(
+ final_class="dirt_jars_grade3",
+ category=None,
+ vote_count=5,
+ output_path="N/A",
+ )
+ }
+ visualizer.print_summary(
+ track_summary,
+ input_directory="/data/dirt_jars_grade3",
+ class_names=["dirt_jars_grade3"],
+ )
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config.yaml b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config.yaml
new file mode 100644
index 00000000000..bdfa0d64670
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config.yaml
@@ -0,0 +1,218 @@
+# ==============================================================================
+# PET Bottle Grade Detection & Tracking Pipeline Configuration
+# ==============================================================================
+# This configuration file defines filesystem paths, model parameters, class
+# hierarchies, detection thresholds, tracking options, and visualization settings
+# used across the SAM3 + DINOv3 inference and tracking pipeline.
+# ==============================================================================
+
+# ------------------------------------------------------------------------------
+# 1. FILESYSTEM PATHS
+# ------------------------------------------------------------------------------
+# Defines input/output directories and naming conventions for generated artifacts.
+paths:
+ # Absolute path to the root directory containing immediate capture-session
+ # subfolders. Each subfolder holds individual image frames from one session.
+ input_image_directory: "/home/umairsabir/pfc_test_data/"
+
+ # Absolute path to the root directory where results will be written. A separate
+ # output directory is created for each subfolder in `input_image_directory`.
+ output_root_directory: "/home/umairsabir/pfc_tracking_output/version_1/"
+
+ # Subfolder created INSIDE each per-subfolder result directory to store
+ # individual annotated frames (active when `visualization.save_frames: true`).
+ output_frame_subfolder: "tracked_frames"
+
+ # Filename of the MP4 video compiled from annotated frames and saved INSIDE
+ # each per-subfolder result directory (active when `visualization.save_video: true`).
+ output_video_filename: "tracked_output.mp4"
+
+ # Subfolder created INSIDE each per-subfolder result directory to store
+ # thumbnail collages (track grids) grouped by category and class
+ # (active when `visualization.save_track_grids: true`).
+ track_grid_subfolder: "track_grids_by_category"
+
+
+# ------------------------------------------------------------------------------
+# 2. MODEL CONFIGURATIONS
+# ------------------------------------------------------------------------------
+# Paths and runtime hyperparameters for the segmentation (SAM3) and
+# classification (DINOv3) models.
+models:
+ sam3:
+ # Absolute path to the pretrained SAM3 open-vocabulary segmentation weights (.pt).
+ checkpoint_path: "/home/umairsabir/sam3_original_weight/sam3.pt"
+ # Hardware accelerator device for running SAM3 inference ("cuda" for GPU, "cpu" for CPU).
+ device: "cuda"
+
+ dinov3:
+ # Absolute path to the local DINOv3 repository directory containing code definition scripts.
+ repo_dir: "/home/umairsabir/dinov3"
+ # Absolute path to the trained DINOv3 fine-grained classifier checkpoint file (.pth).
+ checkpoint_path: "/home/umairsabir/training_output/version_1/model.pth"
+ # Vision Transformer architecture name identifying the backbone (e.g., "dinov3_vitl16").
+ model_name: "dinov3_vitl16"
+ # Target resolution (width & height in pixels) for resizing cropped detections before classification.
+ inference_image_size: 256
+ # Batch size used when passing cropped bounding box images through the DINOv3 classifier.
+ classification_batch_size: 32
+ # Standard RGB channel mean normalization values (ImageNet defaults required by DINOv3).
+ image_mean: [0.485, 0.456, 0.406]
+ # Standard RGB channel standard deviation normalization values (ImageNet defaults).
+ image_std: [0.229, 0.224, 0.225]
+
+
+# ------------------------------------------------------------------------------
+# 3. FINE-GRAINED TARGET CLASSES
+# ------------------------------------------------------------------------------
+# Complete list of target classes in exact label-index order.
+# IMPORTANT: The ordering MUST exactly match the index-to-class mapping output by
+# the DINOv3 classification model header (index 0 -> first class, index 1 -> second class, etc.).
+classes:
+- 'brown_bottles_grade3'
+- 'clean_PET_cold_drink_bottles_with_label_cap_ring_grade1'
+- 'clean_PET_cold_drink_bottles_without_label_cap_ring_grade1'
+- 'clean_PET_mango_juice_bottles_without_label_cap_ring_grade1'
+- 'clean_PET_water_bottles_with_label_cap_ring_grade1'
+- 'clean_PET_water_bottles_without_label_cap_ring_grade1'
+- 'clean_jars_grade1'
+- 'clean_liquor_bottles_without_label_cap_ring_grade1'
+- 'coloured_PET_bottles_grade3'
+- 'dirt_PET_cold_drink_bottles_with_label_cap_ring_grade3'
+- 'dirt_PET_cold_drink_bottles_without_label_cap_ring_grade3'
+- 'dirt_PET_mango_juice_bottles_without_label_cap_ring_grade3'
+- 'dirt_PET_water_bottles_with_label_cap_ring_grade3'
+- 'dirt_PET_water_bottles_without_label_cap_ring_grade3'
+- 'dirt_jars_grade3'
+- 'dirt_liquor_bottles_without_label_cap_ring_grade3'
+- 'full_sleeved_bottles_grade3'
+- 'green_bottles_grade3'
+- 'liquor_bottles_with_label_cap_ring_grade3'
+- 'non_food_bottles_grade3'
+- 'partially_sleeved_mango_juice_bottles_with_label_cap_ring_grade3'
+
+# ------------------------------------------------------------------------------
+# 4. COLLAPSED CATEGORIES (HIERARCHICAL GROUPING)
+# ------------------------------------------------------------------------------
+# Optional grouping of fine-grained classes into broader high-level categories
+# (e.g., reporting by acceptance status or PET grade instead of individual classes).
+#
+# Rules when `enable: true`:
+# - Every class in `classes` above MUST appear in exactly one category under `mapping`.
+# - If any class is missing or duplicated across categories, config loading raises an error.
+#
+# Rules when `enable: false`:
+# - The `mapping` dictionary is ignored.
+# - Per-category reporting tables and category-level folder structures are skipped.
+collapsed_categories:
+ enable: false
+ mapping:
+ accepted:
+ - 'milk_accepted'
+ rejected:
+ - 'milk_rejected'
+
+
+# ------------------------------------------------------------------------------
+# 5. DETECTION & PROMPT PARAMETERS
+# ------------------------------------------------------------------------------
+# Controls text-prompt selection and hyperparameter settings for SAM3 mask
+# generation and bounding box cropping.
+detection:
+ # Key identifying which prompt configuration from `configs` below is actively used during inference.
+ active_prompt: "bottles and containers"
+
+ # Glob patterns specifying which image file extensions to scan inside each session folder.
+ image_file_extensions: ["*.png", "*.jpg", "*.jpeg"]
+
+ # Hyperparameter sets indexed by prompt string.
+ configs:
+ "bottles and containers":
+ # Minimum SAM3 confidence score (0.0 to 1.0) required to generate a mask.
+ confidence_threshold: 0.5
+ # Minimum classification/detection score (0.0 to 1.0) required to retain a bounding box.
+ score_threshold: 0.20
+ # Containment/overlap ratio (0.0 to 1.0) above which a smaller bounding box is
+ # suppressed if it is enclosed inside a larger bounding box (prevents nested duplicates).
+ containment_threshold: 0.98
+ # Maximum pixel length of the image's shorter edge during resizing before SAM3 inference.
+ max_short_side: 1024
+ # Target resolution [width, height] in pixels for resizing bounding box crops for DINOv3.
+ crop_size: [256, 256]
+ # Number of padding pixels added around tight bounding boxes before cropping (captures context).
+ crop_buffer_pixels: 5
+
+ "packets":
+ confidence_threshold: 0.3
+ score_threshold: 0.0
+ containment_threshold: 0.98
+ # Intersection over Union (IoU) threshold above which overlapping bounding boxes
+ # of the same prompt class are merged together into a single bounding box.
+ merge_overlap_threshold: 0.7
+ max_short_side: 1024
+ crop_size: [256, 256]
+ crop_buffer_pixels: 5
+
+
+# ------------------------------------------------------------------------------
+# 6. MULTI-OBJECT TRACKING (BYTETRACK)
+# ------------------------------------------------------------------------------
+# Controls tracking behavior across consecutive video/capture frames using ByteTrack.
+tracking:
+ # Toggle tracking on/off across frames:
+ # - true: Runs ByteTrack to assign stable track IDs across consecutive frames.
+ # - false: Bypasses ByteTrack entirely; every detection in every frame receives a
+ # fresh sequential ID. Use `false` when images in subfolders are independent
+ # (not sequential video frames). When `false`, the thresholds below are ignored.
+ enable: true
+
+ # Minimum Intersection over Union (IoU, 0.0 to 1.0) between a detection box and a predicted
+ # track box required for ByteTrack to link them together across consecutive frames.
+ bytetrack_minimum_iou_threshold: 0.1
+
+ # Minimum number of consecutive frames a track must persist before ByteTrack confirms
+ # and emits the track ID in pipeline outputs.
+ bytetrack_minimum_consecutive_frames: 2
+
+
+# ------------------------------------------------------------------------------
+# 7. VISUALIZATION & ARTIFACT GENERATION
+# ------------------------------------------------------------------------------
+# Toggles and visual styling options for generating output charts, frames, and videos.
+visualization:
+ # ---------------------------------------------------------------------------
+ # Output Toggles (Disabling these saves disk space and speeds up processing)
+ # ---------------------------------------------------------------------------
+ # If true, creates `output_frame_subfolder` ("tracked_frames/") inside each per-subfolder
+ # result directory and saves individual annotated PNG images with bounding boxes and track IDs.
+ save_frames: true
+
+ # If true, compiles annotated frames into an MP4 video file named `output_video_filename`
+ # ("tracked_output.mp4") directly inside each per-subfolder result directory.
+ save_video: false
+
+ # If true, generates and saves track gallery collages (thumbnail grids of each tracked object)
+ # inside `track_grid_subfolder` ("track_grids_by_category/"), grouped by category/class.
+ save_track_grids: true
+
+ # ---------------------------------------------------------------------------
+ # Render & Styling Settings
+ # ---------------------------------------------------------------------------
+ # Playback frame rate (Frames Per Second) when compiling tracking output videos.
+ output_video_fps: 1
+
+ # If true, overlays the model's confidence percentage text on top of bounding box labels.
+ show_confidence_in_labels: true
+
+ # RGB background color [R, G, B] used to fill padding when letterboxing non-square crops
+ # to match `crop_size`. Defaults to DINOv3 / ImageNet mean background color to minimize bias.
+ background_blend_color_rgb: [124, 116, 104]
+
+ # Number of thumbnail columns per row in generated track summary grid figures.
+ track_grid_columns_per_row: 5
+
+ # Physical size (in inches) of each individual thumbnail cell in track summary grid figures.
+ track_grid_thumbnail_size_inches: 3
+
+ # Image resolution in Dots Per Inch (DPI) used when exporting track grid PNG figures.
+ track_grid_dpi: 150
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config_loader.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config_loader.py
new file mode 100644
index 00000000000..fe9910d4d6a
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config_loader.py
@@ -0,0 +1,411 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Loads and validates pipeline configurations using typed data containers."""
+
+from collections.abc import Mapping, Sequence
+import dataclasses
+from typing import Any, Self
+
+import yaml
+
+
+class ConfigurationError(Exception):
+ """Base exception for configuration errors."""
+
+
+@dataclasses.dataclass(frozen=True, kw_only=True)
+class PathsConfig:
+ """File-system paths used by the pipeline.
+
+ Attributes:
+ input_image_directory: Root directory containing immediate subfolders, each
+ holding the image frames of one capture session.
+ output_root_directory: Root directory where a per-subfolder result directory
+ will be created (mirroring the subfolder name).
+ output_frame_subfolder: Name of the per-subfolder folder that stores
+ annotated frames.
+ output_video_filename: Filename of the annotated tracking video saved inside
+ each per-subfolder result directory.
+ track_grid_subfolder: Name of the per-subfolder folder that stores
+ track-grid PNGs grouped by collapsed category and class.
+ """
+
+ input_image_directory: str
+ output_root_directory: str
+ output_frame_subfolder: str
+ output_video_filename: str
+ track_grid_subfolder: str
+
+
+@dataclasses.dataclass(frozen=True)
+class SAM3Config:
+ """Configuration settings for the SAM3 model.
+
+ Attributes:
+ checkpoint_path: Path to the SAM3 model checkpoint file.
+ device: Computational device for running inference (e.g., 'cuda' or 'cpu').
+ """
+
+ checkpoint_path: str
+ device: str
+
+
+@dataclasses.dataclass(frozen=True, kw_only=True)
+class DINOv3Config:
+ """Configuration settings for the DINOv3 model.
+
+ Attributes:
+ repo_dir: Directory containing the DINOv3 code repository.
+ checkpoint_path: Path to the DINOv3 model checkpoint file.
+ model_name: Architecture name of the DINOv3 model (e.g., 'dinov3_vitl16').
+ inference_image_size: Input image resolution for classification inference.
+ classification_batch_size: Batch size used when classifying image crops.
+ image_mean: Standard normalization RGB mean values.
+ image_std: Standard normalization RGB standard deviation values.
+ """
+
+ repo_dir: str
+ checkpoint_path: str
+ model_name: str
+ inference_image_size: int
+ classification_batch_size: int
+ image_mean: tuple[float, float, float]
+ image_std: tuple[float, float, float]
+
+
+@dataclasses.dataclass(frozen=True)
+class ModelsConfig:
+ """Container for all model-related configurations.
+
+ Attributes:
+ sam3: Configuration for the SAM3 segmentation model.
+ dinov3: Configuration for the DINOv3 classification model.
+ """
+
+ sam3: SAM3Config
+ dinov3: DINOv3Config
+
+
+@dataclasses.dataclass(frozen=True, kw_only=True)
+class PromptConfig:
+ """Detection and cropping hyperparameters for a specific text prompt.
+
+ Attributes:
+ confidence_threshold: Minimum confidence threshold for mask generation.
+ score_threshold: Minimum score threshold to retain detections.
+ containment_threshold: Ratio above which a smaller mask is removed if inside
+ a larger mask.
+ max_short_side: Maximum size of the image's short edge during resizing.
+ crop_size: Output resolution (width, height) for bounding box crops.
+ crop_buffer_pixels: Number of buffer pixels to expand around tight boxes.
+ merge_overlap_threshold: Intersection over Union threshold for merging
+ overlapping bounding boxes of the same prompt class.
+ """
+
+ confidence_threshold: float
+ score_threshold: float
+ containment_threshold: float
+ max_short_side: int
+ crop_size: tuple[int, int]
+ crop_buffer_pixels: int
+ merge_overlap_threshold: float = 0.7
+
+
+@dataclasses.dataclass(frozen=True)
+class DetectionConfig:
+ """Root configuration for text prompts and detection parameters.
+
+ Attributes:
+ active_prompt: Key identifying the currently active prompt configuration.
+ image_file_extensions: List of image file extensions to process.
+ configs: Mapping from prompt strings to their corresponding `PromptConfig`.
+ """
+
+ active_prompt: str
+ image_file_extensions: list[str]
+ configs: dict[str, PromptConfig]
+
+
+@dataclasses.dataclass(frozen=True)
+class TrackingConfig:
+ """ByteTrack settings and a toggle to bypass tracking entirely.
+
+ Attributes:
+ bytetrack_minimum_iou_threshold: Minimum IoU for ByteTrack to link a
+ detection to an existing track. Ignored when `enable` is False.
+ bytetrack_minimum_consecutive_frames: Minimum frames a track must persist
+ before ByteTrack emits it. Ignored when `enable` is False.
+ enable: When True (default), ByteTrack runs normally and IDs are stable
+ across frames. When False, tracking is bypassed entirely and every
+ detection in every frame receives a fresh sequential ID. Use False when
+ input images are independent (not consecutive video frames).
+ """
+
+ bytetrack_minimum_iou_threshold: float
+ bytetrack_minimum_consecutive_frames: int
+ enable: bool = True
+
+
+@dataclasses.dataclass(frozen=True, kw_only=True)
+class VisualizationConfig:
+ """Options and visual styling settings for pipeline output visualization.
+
+ Attributes:
+ save_frames: Whether to save individual annotated frames.
+ save_video: Whether to compile and save an output tracking video.
+ save_track_grids: Whether to generate and save track gallery grid images.
+ output_video_fps: Frames per second for generated tracking output videos.
+ show_confidence_in_labels: Whether to overlay confidence scores on labels.
+ background_blend_color_rgb: RGB background color used for crops.
+ track_grid_columns_per_row: Number of track thumbnails per row in grids.
+ track_grid_thumbnail_size_inches: Size in inches of each grid thumbnail.
+ track_grid_dpi: Dots per inch (DPI) for saving track grid figures.
+ """
+
+ save_frames: bool
+ save_video: bool
+ save_track_grids: bool
+ output_video_fps: int
+ show_confidence_in_labels: bool
+ background_blend_color_rgb: tuple[int, int, int]
+ track_grid_columns_per_row: int
+ track_grid_thumbnail_size_inches: int
+ track_grid_dpi: int
+
+
+@dataclasses.dataclass(frozen=True)
+class CollapsedCategoriesConfig:
+ """Optional grouping of fine-grained classes into broader categories.
+
+ When enabled, every class in the pipeline's `classes` list must be assigned
+ to exactly one category. The category for a given class is looked up via
+ `get_category_for_class`. When disabled, the mapping is empty and no
+ per-category reporting is performed.
+
+ Attributes:
+ enable: Whether the collapsed-category feature is active.
+ mapping: Mapping from category name to the list of class names that fall
+ under it. Empty when disabled.
+ """
+
+ enable: bool
+ mapping: dict[str, list[str]] = dataclasses.field(default_factory=dict)
+
+ def get_category_for_class(self, class_name: str) -> str | None:
+ """Returns the category that contains the given class, or None if disabled.
+
+ Args:
+ class_name: The fine-grained class name to look up.
+
+ Returns:
+ The matching category name when the feature is enabled, or None when the
+ feature is disabled.
+
+ Raises:
+ ConfigurationError: If the feature is enabled but the class is not present
+ in any category (this indicates the validation in `from_yaml` was
+ bypassed).
+ """
+ if not self.enable:
+ return None
+ for category_name, class_list in self.mapping.items():
+ if class_name in class_list:
+ return category_name
+ raise ConfigurationError(
+ f"Class '{class_name}' is not assigned to any collapsed category."
+ )
+
+ @property
+ def category_names(self) -> list[str]:
+ """Returns the configured category names in declaration order.
+
+ Returns:
+ List of category names. Empty list when the feature is disabled.
+ """
+ if not self.enable:
+ return []
+ return list(self.mapping.keys())
+
+
+def build_collapsed_categories_config(
+ raw_section: Mapping[str, Any] | None, classes: Sequence[str]
+) -> CollapsedCategoriesConfig:
+ """Builds and validates the CollapsedCategoriesConfig from raw YAML.
+
+ The YAML section is expected to look like:
+ collapsed_categories:
+ enable: true
+ mapping:
+ category_name: [class_a, class_b]
+
+ Validation rules when enabled:
+ - Every class in `classes` must appear in exactly one category.
+ - No class may appear in more than one category.
+ - Every class in the mapping must be present in `classes`.
+
+ Args:
+ raw_section: The raw `collapsed_categories` dict from YAML, or None if the
+ section was omitted entirely.
+ classes: The full list of class names from the config.
+
+ Returns:
+ A validated CollapsedCategoriesConfig instance.
+
+ Raises:
+ ConfigurationError: If validation fails.
+ """
+ if raw_section is None or not raw_section.get("enable", False):
+ return CollapsedCategoriesConfig(enable=False, mapping={})
+
+ raw_mapping = raw_section.get("mapping") or {}
+ if not isinstance(raw_mapping, Mapping) or not raw_mapping:
+ raise ConfigurationError(
+ "collapsed_categories.enable is true but 'mapping' is empty or "
+ "missing."
+ )
+
+ seen_classes: dict[str, str] = {}
+ mapping_dict: dict[str, list[str]] = {}
+ class_names = set(classes)
+ for category_name, class_list in raw_mapping.items():
+ if not isinstance(class_list, Sequence) or isinstance(class_list, str):
+ raise ConfigurationError(
+ f"Category '{category_name}' must map to a list of class names."
+ )
+ if not class_list:
+ raise ConfigurationError(
+ f"Category '{category_name}' must map to a non-empty list of class "
+ "names."
+ )
+ mapping_dict[str(category_name)] = list(class_list)
+ for class_name in class_list:
+ if class_name not in class_names:
+ raise ConfigurationError(
+ f"Class '{class_name}' in category '{category_name}' is not "
+ "present in the top-level 'classes' list."
+ )
+ if class_name in seen_classes:
+ raise ConfigurationError(
+ f"Class '{class_name}' is assigned to both "
+ f"'{seen_classes[class_name]}' and '{category_name}'."
+ )
+ seen_classes[class_name] = category_name
+
+ unmapped_classes = [c for c in classes if c not in seen_classes]
+ if unmapped_classes:
+ raise ConfigurationError(
+ "The following classes are not assigned to any collapsed category: "
+ f"{unmapped_classes}"
+ )
+
+ return CollapsedCategoriesConfig(enable=True, mapping=mapping_dict)
+
+
+@dataclasses.dataclass(frozen=True, kw_only=True)
+class PipelineConfig:
+ """Root configuration object representing the entire pipeline state."""
+
+ paths: PathsConfig
+ models: ModelsConfig
+ classes: list[str]
+ detection: DetectionConfig
+ tracking: TrackingConfig
+ visualization: VisualizationConfig
+ collapsed_categories: CollapsedCategoriesConfig
+
+ @classmethod
+ def from_yaml(cls, yaml_path: str) -> Self:
+ """Parses the YAML file into a strictly typed configuration object.
+
+ Args:
+ yaml_path: Path to the YAML configuration file.
+
+ Returns:
+ A fully populated PipelineConfig.
+
+ Raises:
+ ConfigurationError: If the YAML file cannot be found, if the file content
+ is invalid YAML, or if configuration validation fails.
+ """
+ try:
+ with open(yaml_path, "r", encoding="utf-8") as file:
+ data = yaml.safe_load(file)
+ except OSError as err:
+ raise ConfigurationError(
+ f"Config file not found or inaccessible: {yaml_path}"
+ ) from err
+ except yaml.YAMLError as err:
+ raise ConfigurationError(
+ f"Invalid YAML syntax in {yaml_path}: {err}"
+ ) from err
+
+ if not isinstance(data, dict):
+ raise ConfigurationError(
+ f"Configuration root in {yaml_path} must be a dictionary."
+ )
+
+ try:
+ prompt_configs = {}
+ for name, cfg in data["detection"]["configs"].items():
+ cfg_copy = dict(cfg)
+ if "crop_size" in cfg_copy:
+ cfg_copy["crop_size"] = tuple(cfg_copy["crop_size"])
+ prompt_configs[name] = PromptConfig(**cfg_copy)
+
+ classes = list(data["classes"])
+ collapsed_categories = build_collapsed_categories_config(
+ raw_section=data.get("collapsed_categories"),
+ classes=classes,
+ )
+
+ models_data = data["models"]
+ dinov3_data = dict(models_data["dinov3"])
+ if "image_mean" in dinov3_data:
+ dinov3_data["image_mean"] = tuple(dinov3_data["image_mean"])
+ if "image_std" in dinov3_data:
+ dinov3_data["image_std"] = tuple(dinov3_data["image_std"])
+
+ visualization_data = dict(data["visualization"])
+ if "background_blend_color_rgb" in visualization_data:
+ visualization_data["background_blend_color_rgb"] = tuple(
+ visualization_data["background_blend_color_rgb"]
+ )
+
+ return cls(
+ paths=PathsConfig(**data["paths"]),
+ models=ModelsConfig(
+ sam3=SAM3Config(**models_data["sam3"]),
+ dinov3=DINOv3Config(**dinov3_data),
+ ),
+ classes=classes,
+ detection=DetectionConfig(
+ active_prompt=data["detection"]["active_prompt"],
+ image_file_extensions=list(
+ data["detection"]["image_file_extensions"]
+ ),
+ configs=prompt_configs,
+ ),
+ tracking=TrackingConfig(**data["tracking"]),
+ visualization=VisualizationConfig(**visualization_data),
+ collapsed_categories=collapsed_categories,
+ )
+ except ConfigurationError:
+ raise
+ except Exception as err:
+ raise ConfigurationError(
+ f"Error validating configuration structure in {yaml_path}: {err}"
+ ) from err
+
+
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config_loader_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config_loader_test.py
new file mode 100644
index 00000000000..08535bfe85f
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/config_loader_test.py
@@ -0,0 +1,140 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests verifying pipeline configuration loading, validation, and category mapping."""
+
+import os
+from typing import Any
+from absl.testing import absltest
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import config_loader
+
+
+class ConfigLoaderTest(absltest.TestCase):
+ """Test suite covering YAML deserialization and category mapping validation rules."""
+
+ def setUp(self):
+ """Initializes common test fixtures and absolute paths to test configs."""
+ super().setUp()
+ self.valid_config_path = os.path.join(
+ os.path.dirname(os.path.abspath(__file__)), "config.yaml"
+ )
+
+ def assertRaisesConfigError(
+ self,
+ raw_section: dict[str, Any],
+ classes: list[str],
+ expected_message_regex: str,
+ ):
+ """Helper asserting that building collapsed categories config raises ConfigurationError."""
+ with self.assertRaisesRegex(
+ config_loader.ConfigurationError, expected_message_regex
+ ):
+ config_loader.build_collapsed_categories_config(raw_section, classes)
+
+ def test_from_yaml_valid_config(self):
+ """Verifies that a valid YAML configuration loads with expected typed data containers."""
+ config = config_loader.PipelineConfig.from_yaml(self.valid_config_path)
+ self.assertIsInstance(config, config_loader.PipelineConfig)
+ self.assertEqual(config.models.dinov3.inference_image_size, 256)
+ self.assertEqual(config.models.dinov3.classification_batch_size, 32)
+ self.assertIsInstance(config.models.dinov3.image_mean, tuple)
+ self.assertLen(config.models.dinov3.image_mean, 3)
+ self.assertIn("bottles and containers", config.detection.configs)
+ self.assertIsInstance(
+ config.detection.configs["bottles and containers"].crop_size, tuple
+ )
+ self.assertEqual(
+ config.detection.configs["bottles and containers"].crop_size, (256, 256)
+ )
+ self.assertFalse(config.collapsed_categories.enable)
+ self.assertEmpty(config.collapsed_categories.category_names)
+
+ def test_from_yaml_file_not_found(self):
+ """Ensures ConfigurationError is raised when the YAML file path does not exist."""
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig.from_yaml(
+ "/non_existent_path/bad_config.yaml"
+ )
+
+ def test_from_yaml_invalid_syntax(self):
+ """Ensures ConfigurationError is raised when the file contains malformed YAML syntax."""
+ temp_file = self.create_tempfile(
+ file_path="bad_config.yaml",
+ content="paths:\n input_image_directory: [unclosed_list\n",
+ )
+ with self.assertRaises(config_loader.ConfigurationError):
+ config_loader.PipelineConfig.from_yaml(temp_file.full_path)
+
+ def test_collapsed_categories_config_enabled_valid(self):
+ """Verifies valid category mappings when the collapsed categories feature is enabled."""
+ classes = ["class_a", "class_b", "class_c"]
+ raw_section = {
+ "enable": True,
+ "mapping": {
+ "cat_1": ["class_a", "class_b"],
+ "cat_2": ["class_c"],
+ },
+ }
+ cfg = config_loader.build_collapsed_categories_config(raw_section, classes)
+ self.assertTrue(cfg.enable)
+ self.assertEqual(cfg.get_category_for_class("class_a"), "cat_1")
+ self.assertEqual(cfg.get_category_for_class("class_c"), "cat_2")
+ self.assertEqual(cfg.category_names, ["cat_1", "cat_2"])
+
+ def test_collapsed_categories_config_missing_classes(self):
+ """Ensures an error is raised when top-level classes are omitted from category mappings."""
+ classes = ["class_a", "class_b", "class_c"]
+ raw_section = {
+ "enable": True,
+ "mapping": {
+ "cat_1": ["class_a"],
+ },
+ }
+ self.assertRaisesConfigError(
+ raw_section, classes, "not assigned to any collapsed category"
+ )
+
+ def test_collapsed_categories_config_duplicate_class(self):
+ """Ensures an error is raised if a class is assigned to multiple categories simultaneously."""
+ classes = ["class_a", "class_b"]
+ raw_section = {
+ "enable": True,
+ "mapping": {
+ "cat_1": ["class_a", "class_b"],
+ "cat_2": ["class_b"],
+ },
+ }
+ self.assertRaisesConfigError(raw_section, classes, "assigned to both")
+
+ def test_collapsed_categories_get_category_for_class_not_found(self):
+ """Verifies lookup behavior when querying an unassigned class on an active config."""
+ cfg = config_loader.CollapsedCategoriesConfig(
+ enable=True,
+ mapping={"cat_1": ["class_a"]},
+ )
+ with self.assertRaisesRegex(
+ config_loader.ConfigurationError,
+ "not assigned to any collapsed category",
+ ):
+ cfg.get_category_for_class("unknown_class")
+
+ def test_collapsed_categories_disabled_return_vals(self):
+ """Ensures lookups return None and empty lists when the collapsed categories feature is disabled."""
+ cfg = config_loader.CollapsedCategoriesConfig(enable=False, mapping={})
+ self.assertIsNone(cfg.get_category_for_class("any_class"))
+ self.assertEmpty(cfg.category_names)
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/dinov3_classifier.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/dinov3_classifier.py
new file mode 100644
index 00000000000..51b5f25184d
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/dinov3_classifier.py
@@ -0,0 +1,399 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""DINOv3 classification module for batched image inference.
+
+This module exposes:
+ - ClassifierError: Exception raised for classifier construction or inference
+ issues.
+ - PoolingStrategy: Enum describing how backbone token features are pooled
+ before the classification head.
+ - Prediction: Typed result returned per input image by the classifier.
+ - Dinov3ClassificationModule: The underlying nn.Module wrapping a DINOv3
+ backbone and a linear head.
+ - DINOv3Classifier: High-level batched classifier that accepts an
+ already-loaded model; construct it via `DINOv3Classifier.from_config` for
+ the common case of loading from a checkpoint on disk.
+"""
+
+from collections.abc import Mapping, Sequence
+import enum
+import logging
+import pathlib
+from typing import Self, TypedDict
+
+import numpy
+from PIL import Image
+import torch
+from torch import nn
+from torch.nn import functional
+from torchvision.transforms import v2
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import config_loader
+
+_LOGGER = logging.getLogger(__name__)
+
+
+class PoolingStrategy(enum.Enum):
+ """Strategy for pooling backbone token features into a single vector.
+
+ Attributes:
+ CLS: Use only the CLS token embedding.
+ CLS_MEAN_PATCH: Concatenate the CLS token with the mean of patch tokens.
+ """
+
+ CLS = "cls"
+ CLS_MEAN_PATCH = "cls_mean_patch"
+
+
+class ClassifierError(Exception):
+ """Exception raised for classifier loading, construction, or shape issues."""
+
+
+class Prediction(TypedDict):
+ """Typed prediction result for a single image.
+
+ Callers index by string key exactly as with a plain dict; the TypedDict adds
+ static-analysis support without changing runtime behavior.
+
+ Attributes:
+ predicted_class: Name of the top-1 class.
+ predicted_probability_percent: Top-1 probability expressed as a percentage
+ in the range [0.0, 100.0].
+ all_probabilities_percent: Mapping from class name to that class's
+ probability as a percentage in the range [0.0, 100.0].
+ """
+
+ predicted_class: str
+ predicted_probability_percent: float
+ all_probabilities_percent: dict[str, float]
+
+
+def _resolve_device(requested_device: str) -> torch.device:
+ """Resolves the requested device string to an available torch.device.
+
+ Falls back to CPU with a warning if the requested device is not available.
+
+ Args:
+ requested_device: Device string requested by the caller.
+
+ Returns:
+ The resolved torch.device.
+ """
+ if requested_device.startswith("cuda") and not torch.cuda.is_available():
+ _LOGGER.warning(
+ "Requested device '%s' but CUDA is not available; falling back to CPU.",
+ requested_device,
+ )
+ return torch.device("cpu")
+
+ if requested_device == "mps" and not torch.backends.mps.is_available():
+ _LOGGER.warning(
+ "Requested device 'mps' but MPS is not available; falling back to CPU."
+ )
+ return torch.device("cpu")
+
+ return torch.device(requested_device)
+
+
+def _infer_pooling_from_state_dict(
+ saved_state_dict: Mapping[str, torch.Tensor], hidden_size: int
+) -> PoolingStrategy:
+ """Infers the pooling strategy from the shape of the saved head weights.
+
+ Args:
+ saved_state_dict: The checkpoint's model state dict.
+ hidden_size: Backbone hidden dimensionality.
+
+ Returns:
+ The inferred pooling strategy.
+
+ Raises:
+ ClassifierError: If the head input dimension matches neither the CLS nor
+ the CLS_MEAN_PATCH expectation.
+ """
+ head_input_features = saved_state_dict["head.weight"].shape[1]
+ if head_input_features == hidden_size:
+ return PoolingStrategy.CLS
+ if head_input_features == 2 * hidden_size:
+ return PoolingStrategy.CLS_MEAN_PATCH
+ raise ClassifierError(
+ "Cannot infer pooling strategy. Head input dimension "
+ f"{head_input_features} does not match hidden size {hidden_size} or "
+ f"{2 * hidden_size}."
+ )
+
+
+def _load_checkpoint_state_dict(
+ checkpoint_path: pathlib.Path | str, device: torch.device
+) -> dict[str, torch.Tensor]:
+ """Loads the model state dict from a checkpoint file.
+
+ Args:
+ checkpoint_path: Filesystem path to the checkpoint.
+ device: Target device for `map_location`.
+
+ Returns:
+ The `model_state_dict` mapping from parameter name to tensor.
+
+ Raises:
+ ClassifierError: If the checkpoint is missing required keys.
+ """
+ checkpoint = torch.load(
+ checkpoint_path, map_location=device, weights_only=True
+ )
+ if "model_state_dict" not in checkpoint:
+ raise ClassifierError(
+ "Checkpoint is missing required key 'model_state_dict'."
+ )
+ saved_state_dict = checkpoint["model_state_dict"]
+ if "head.weight" not in saved_state_dict:
+ raise ClassifierError(
+ "Checkpoint state dict is missing required key 'head.weight'."
+ )
+ return saved_state_dict
+
+
+def _build_image_transform(config: config_loader.DINOv3Config) -> v2.Compose:
+ """Builds the preprocessing pipeline used to feed PIL images to the model.
+
+ Args:
+ config: DINOv3 model configuration providing image size and normalization
+ statistics.
+
+ Returns:
+ A torchvision v2.Compose pipeline.
+ """
+ return v2.Compose([
+ v2.ToImage(),
+ v2.Resize(
+ (config.inference_image_size, config.inference_image_size),
+ antialias=True,
+ ),
+ v2.ToDtype(torch.float32, scale=True),
+ v2.Normalize(mean=config.image_mean, std=config.image_std),
+ ])
+
+
+class Dinov3ClassificationModule(nn.Module):
+ """DINOv3 backbone paired with a linear classification head.
+
+ Attributes:
+ pooling: The pooling strategy applied to backbone features.
+ backbone_model: The DINOv3 backbone.
+ head: Linear layer mapping pooled features to per-class logits.
+ """
+
+ def __init__(
+ self,
+ backbone_model: nn.Module,
+ hidden_size: int,
+ number_of_classes: int,
+ pooling: PoolingStrategy,
+ ):
+ """Initializes the module from an already-loaded backbone.
+
+ Args:
+ backbone_model: The DINOv3 backbone module. Injected so that this class
+ does not perform filesystem I/O in its constructor.
+ hidden_size: Dimensionality of the backbone's token embeddings.
+ number_of_classes: Number of output classes for the head.
+ pooling: Pooling strategy that determines the head's input dimension.
+ """
+ super().__init__()
+ self.pooling = pooling
+ self.backbone_model = backbone_model
+
+ if pooling is PoolingStrategy.CLS:
+ head_input_features = hidden_size
+ else:
+ head_input_features = 2 * hidden_size
+
+ self.head = nn.Linear(
+ in_features=head_input_features,
+ out_features=number_of_classes,
+ bias=True,
+ )
+
+ def extract_features(self, image_batch: torch.Tensor) -> torch.Tensor:
+ """Extracts pooled features from a batch of images.
+
+ Args:
+ image_batch: Tensor of shape (batch_size, 3, height, width).
+
+ Returns:
+ Tensor of pooled features. Shape is (batch_size, hidden_size) for the
+ CLS strategy and (batch_size, 2 * hidden_size) for CLS_MEAN_PATCH.
+ """
+ if self.pooling is PoolingStrategy.CLS:
+ return self.backbone_model(image_batch)
+
+ token_features = self.backbone_model.forward_features(image_batch)
+ cls_token = token_features["x_norm_clstoken"]
+ mean_patch_token = token_features["x_norm_patchtokens"].mean(dim=1)
+ return torch.cat([cls_token, mean_patch_token], dim=1)
+
+ def forward(self, image_batch: torch.Tensor) -> torch.Tensor:
+ """Runs the full backbone + head forward pass.
+
+ Args:
+ image_batch: Tensor of shape (batch_size, 3, height, width).
+
+ Returns:
+ Per-class logits of shape (batch_size, number_of_classes).
+ """
+ return self.head(self.extract_features(image_batch))
+
+
+class DINOv3Classifier:
+ """High-level batched classifier backed by a DINOv3 model.
+
+ The order of `class_names` is load-bearing: `class_names[i]` must correspond
+ to logit index `i` in the trained model. Reordering this list will silently
+ produce incorrect predictions.
+
+ Construct instances via `DINOv3Classifier.from_config` for the common case of
+ loading from a checkpoint. The `__init__` constructor accepts an
+ already-built model and is intended for callers that manage model loading
+ themselves (for example, unit tests).
+
+ Attributes:
+ class_names: Ordered list of class names indexed by logit position.
+ """
+
+ def __init__(
+ self,
+ model: Dinov3ClassificationModule,
+ class_names: Sequence[str],
+ image_transform: v2.Compose,
+ device: torch.device,
+ ):
+ """Initializes the classifier from injected dependencies.
+
+ Args:
+ model: The already-loaded classification module, expected to be on
+ `device` and in eval mode. This constructor does not move the model or
+ change its mode.
+ class_names: Ordered list of class names indexed by logit position.
+ image_transform: Preprocessing pipeline applied to each PIL image before
+ stacking into a batch.
+ device: Target device on which the model resides and to which input
+ batches will be moved.
+ """
+ self.class_names = list(class_names)
+ self._model = model
+ self._image_transform = image_transform
+ self._device = device
+
+ @classmethod
+ def from_config(
+ cls,
+ config: config_loader.DINOv3Config,
+ class_names: Sequence[str],
+ device: str,
+ ) -> Self:
+ """Builds a classifier by loading a checkpoint from disk.
+
+ Args:
+ config: DINOv3 model configuration.
+ class_names: Ordered list of class names indexed by logit position.
+ device: Requested device string (e.g., 'cuda', 'cpu', 'mps').
+
+ Returns:
+ A ready-to-use DINOv3Classifier with its model on the resolved device
+ and in eval mode.
+
+ Raises:
+ ClassifierError: If the checkpoint is missing required keys, if the
+ pooling strategy cannot be inferred, or if the state dict fails to
+ load into the constructed model.
+ """
+ resolved_device = _resolve_device(device)
+ saved_state_dict = _load_checkpoint_state_dict(
+ checkpoint_path=config.checkpoint_path, device=resolved_device
+ )
+
+ backbone_model = torch.hub.load(
+ config.repo_dir,
+ config.model_name,
+ source="local",
+ pretrained=False,
+ )
+ # DINOv3 Vision Transformers store token embedding dimension in
+ # norm.normalized_shape[0].
+ hidden_size = backbone_model.norm.normalized_shape[0]
+
+ pooling = _infer_pooling_from_state_dict(
+ saved_state_dict=saved_state_dict, hidden_size=hidden_size
+ )
+
+ model = Dinov3ClassificationModule(
+ backbone_model=backbone_model,
+ hidden_size=hidden_size,
+ number_of_classes=len(class_names),
+ pooling=pooling,
+ ).to(resolved_device)
+
+ try:
+ model.load_state_dict(saved_state_dict)
+ except RuntimeError as error:
+ raise ClassifierError(
+ f"Failed to load state dict into model: {error}"
+ ) from error
+
+ model.eval()
+ return cls(
+ model=model,
+ class_names=class_names,
+ image_transform=_build_image_transform(config),
+ device=resolved_device,
+ )
+
+ @torch.no_grad()
+ def predict_batch(self, images: Sequence[Image.Image]) -> list[Prediction]:
+ """Classifies a batch of PIL images in a single forward pass.
+
+ Args:
+ images: Sequence of PIL images to classify.
+
+ Returns:
+ List of Prediction dicts in the same order as `images`. Each Prediction
+ contains 'predicted_class', 'predicted_probability_percent', and
+ 'all_probabilities_percent'. Returns an empty list when `images` is
+ empty.
+ """
+ if not images:
+ return []
+
+ image_tensors = [self._image_transform(image) for image in images]
+ batch_tensor = torch.stack(image_tensors).to(self._device)
+
+ logits = self._model(batch_tensor)
+ probabilities = functional.softmax(logits, dim=1).cpu().numpy()
+
+ predictions: list[Prediction] = []
+ for probability_row in probabilities:
+ predicted_index = int(numpy.argmax(probability_row))
+ all_probabilities_percent: dict[str, float] = {}
+ for class_name, probability in zip(self.class_names, probability_row):
+ all_probabilities_percent[class_name] = float(probability * 100.0)
+ prediction: Prediction = {
+ "predicted_class": self.class_names[predicted_index],
+ "predicted_probability_percent": float(
+ probability_row[predicted_index] * 100.0
+ ),
+ "all_probabilities_percent": all_probabilities_percent,
+ }
+ predictions.append(prediction)
+ return predictions
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/dinov3_classifier_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/dinov3_classifier_test.py
new file mode 100644
index 00000000000..590272034f4
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/dinov3_classifier_test.py
@@ -0,0 +1,375 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for the DINOv3 classification module and classifier."""
+
+from unittest import mock
+
+from absl.testing import absltest
+from PIL import Image
+import torch
+from torch import nn
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import config_loader
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import dinov3_classifier
+
+
+class DummyBackbone(nn.Module):
+ """Minimal backbone stand-in exposing the interface required by the module."""
+
+ def __init__(self, hidden_size: int = 64):
+ """Initializes the dummy backbone with a simulated norm layer.
+
+ Args:
+ hidden_size: Simulated token embedding dimensionality.
+ """
+ super().__init__()
+ self.norm = nn.LayerNorm(hidden_size)
+ self.hidden_size = hidden_size
+
+ def forward(self, image_batch: torch.Tensor) -> torch.Tensor:
+ """Returns dummy token embeddings of shape (batch_size, hidden_size).
+
+ Args:
+ image_batch: Batch of images.
+
+ Returns:
+ Tensor of ones with shape (batch_size, hidden_size).
+ """
+ batch_size = image_batch.shape[0]
+ return torch.ones((batch_size, self.hidden_size), dtype=torch.float32)
+
+ def forward_features(
+ self, image_batch: torch.Tensor
+ ) -> dict[str, torch.Tensor]:
+ """Returns dummy CLS and patch token features.
+
+ Args:
+ image_batch: Batch of images.
+
+ Returns:
+ Mapping matching the DINOv3 feature-extraction interface.
+ """
+ batch_size = image_batch.shape[0]
+ return {
+ "x_norm_clstoken": torch.ones(
+ (batch_size, self.hidden_size), dtype=torch.float32
+ ),
+ "x_norm_patchtokens": torch.ones(
+ (batch_size, 16, self.hidden_size), dtype=torch.float32
+ ),
+ }
+
+
+def _make_dummy_config() -> config_loader.DINOv3Config:
+ """Returns a DINOv3Config populated with values safe for tests."""
+ return config_loader.DINOv3Config(
+ repo_dir="/dummy/repo",
+ model_name="dinov3_vits14",
+ checkpoint_path="/dummy/checkpoint.pt",
+ inference_image_size=256,
+ image_mean=(0.485, 0.456, 0.406),
+ image_std=(0.229, 0.224, 0.225),
+ classification_batch_size=16,
+ )
+
+
+def _make_dummy_state_dict(
+ hidden_size: int = 32, number_of_classes: int = 3
+) -> dict[str, torch.Tensor]:
+ """Returns a state dict compatible with a DummyBackbone + CLS pooling head."""
+ return {
+ "backbone_model.norm.weight": torch.ones((hidden_size,)),
+ "backbone_model.norm.bias": torch.zeros((hidden_size,)),
+ "head.weight": torch.zeros((number_of_classes, hidden_size)),
+ "head.bias": torch.zeros((number_of_classes,)),
+ }
+
+
+class ResolveDeviceTest(absltest.TestCase):
+ """Tests for the module-level _resolve_device helper."""
+
+ def test_cpu_resolves_directly(self):
+ """Verifies that 'cpu' resolves to torch.device('cpu')."""
+ self.assertEqual(
+ dinov3_classifier._resolve_device("cpu"), torch.device("cpu")
+ )
+
+ @mock.patch.object(torch.cuda, "is_available", autospec=True)
+ def test_cuda_resolves_when_available(self, mock_is_available):
+ """Verifies that a CUDA device string resolves directly when available."""
+ mock_is_available.return_value = True
+ self.assertEqual(
+ dinov3_classifier._resolve_device("cuda:0"), torch.device("cuda:0")
+ )
+
+ @mock.patch.object(torch.cuda, "is_available", autospec=True)
+ def test_cuda_falls_back_to_cpu_when_unavailable(self, mock_is_available):
+ """Verifies fallback to CPU when CUDA is requested but unavailable."""
+ mock_is_available.return_value = False
+ self.assertEqual(
+ dinov3_classifier._resolve_device("cuda:0"), torch.device("cpu")
+ )
+
+ @mock.patch.object(torch.backends.mps, "is_available", autospec=True)
+ def test_mps_falls_back_to_cpu_when_unavailable(self, mock_is_available):
+ """Verifies fallback to CPU when MPS is requested but unavailable."""
+ mock_is_available.return_value = False
+ self.assertEqual(
+ dinov3_classifier._resolve_device("mps"), torch.device("cpu")
+ )
+
+
+class InferPoolingFromStateDictTest(absltest.TestCase):
+ """Tests for the module-level _infer_pooling_from_state_dict helper."""
+
+ def test_cls_pooling_from_matching_head_shape(self):
+ """Verifies CLS pooling inferred when head input equals hidden size."""
+ state_dict = {"head.weight": torch.zeros((5, 64))}
+ self.assertEqual(
+ dinov3_classifier._infer_pooling_from_state_dict(
+ state_dict, hidden_size=64
+ ),
+ dinov3_classifier.PoolingStrategy.CLS,
+ )
+
+ def test_cls_mean_patch_pooling_from_doubled_head_shape(self):
+ """Verifies CLS_MEAN_PATCH pooling inferred when head input is doubled."""
+ state_dict = {"head.weight": torch.zeros((5, 128))}
+ self.assertEqual(
+ dinov3_classifier._infer_pooling_from_state_dict(
+ state_dict, hidden_size=64
+ ),
+ dinov3_classifier.PoolingStrategy.CLS_MEAN_PATCH,
+ )
+
+ def test_invalid_head_shape_raises_classifier_error(self):
+ """Verifies an error is raised when head shape matches neither strategy."""
+ state_dict = {"head.weight": torch.zeros((5, 100))}
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError, "Cannot infer pooling strategy"
+ ):
+ dinov3_classifier._infer_pooling_from_state_dict(
+ state_dict, hidden_size=64
+ )
+
+
+class LoadCheckpointStateDictTest(absltest.TestCase):
+ """Tests for the module-level _load_checkpoint_state_dict helper."""
+
+ @mock.patch.object(torch, "load", autospec=True)
+ def test_missing_model_state_dict_raises_error(self, mock_load):
+ """Verifies error when checkpoint lacks 'model_state_dict' key."""
+ mock_load.return_value = {"other_key": {}}
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError,
+ "missing required key 'model_state_dict'",
+ ):
+ dinov3_classifier._load_checkpoint_state_dict(
+ checkpoint_path="/dummy/checkpoint.pt", device=torch.device("cpu")
+ )
+
+ @mock.patch.object(torch, "load", autospec=True)
+ def test_missing_head_weight_raises_error(self, mock_load):
+ """Verifies error when state dict lacks 'head.weight' parameter."""
+ mock_load.return_value = {
+ "model_state_dict": {"other.weight": torch.zeros(1)}
+ }
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError,
+ "missing required key 'head.weight'",
+ ):
+ dinov3_classifier._load_checkpoint_state_dict(
+ checkpoint_path="/dummy/checkpoint.pt", device=torch.device("cpu")
+ )
+
+ @mock.patch.object(torch, "load", autospec=True)
+ def test_valid_checkpoint_returns_state_dict(self, mock_load):
+ """Verifies a valid checkpoint returns the nested state dict."""
+ state_dict = _make_dummy_state_dict()
+ mock_load.return_value = {"model_state_dict": state_dict}
+ result = dinov3_classifier._load_checkpoint_state_dict(
+ checkpoint_path="/dummy/checkpoint.pt", device=torch.device("cpu")
+ )
+ self.assertIs(result, state_dict)
+
+
+class Dinov3ClassificationModuleTest(absltest.TestCase):
+ """Tests for the Dinov3ClassificationModule PyTorch module."""
+
+ def test_init_cls_strategy_sets_head_dimensions(self):
+ """Verifies module initialization with CLS pooling strategy."""
+ module = dinov3_classifier.Dinov3ClassificationModule(
+ backbone_model=DummyBackbone(hidden_size=32),
+ hidden_size=32,
+ number_of_classes=5,
+ pooling=dinov3_classifier.PoolingStrategy.CLS,
+ )
+ self.assertEqual(module.head.in_features, 32)
+ self.assertEqual(module.head.out_features, 5)
+
+ def test_init_cls_mean_patch_strategy_doubles_head_input(self):
+ """Verifies module initialization with CLS_MEAN_PATCH pooling strategy."""
+ module = dinov3_classifier.Dinov3ClassificationModule(
+ backbone_model=DummyBackbone(hidden_size=32),
+ hidden_size=32,
+ number_of_classes=5,
+ pooling=dinov3_classifier.PoolingStrategy.CLS_MEAN_PATCH,
+ )
+ self.assertEqual(module.head.in_features, 64)
+ self.assertEqual(module.head.out_features, 5)
+
+ def test_extract_features_cls_strategy_shape(self):
+ """Verifies extract_features returns backbone forward output under CLS."""
+ module = dinov3_classifier.Dinov3ClassificationModule(
+ backbone_model=DummyBackbone(hidden_size=32),
+ hidden_size=32,
+ number_of_classes=4,
+ pooling=dinov3_classifier.PoolingStrategy.CLS,
+ )
+ features = module.extract_features(torch.zeros((2, 3, 256, 256)))
+ self.assertEqual(features.shape, (2, 32))
+
+ def test_extract_features_cls_mean_patch_strategy_shape(self):
+ """Verifies extract_features concatenates CLS and mean patch tokens."""
+ module = dinov3_classifier.Dinov3ClassificationModule(
+ backbone_model=DummyBackbone(hidden_size=32),
+ hidden_size=32,
+ number_of_classes=4,
+ pooling=dinov3_classifier.PoolingStrategy.CLS_MEAN_PATCH,
+ )
+ features = module.extract_features(torch.zeros((2, 3, 256, 256)))
+ self.assertEqual(features.shape, (2, 64))
+
+ def test_forward_pass_produces_expected_logit_shape(self):
+ """Verifies the full forward pass produces per-class logits."""
+ module = dinov3_classifier.Dinov3ClassificationModule(
+ backbone_model=DummyBackbone(hidden_size=32),
+ hidden_size=32,
+ number_of_classes=3,
+ pooling=dinov3_classifier.PoolingStrategy.CLS,
+ )
+ logits = module(torch.zeros((2, 3, 256, 256)))
+ self.assertEqual(logits.shape, (2, 3))
+
+
+class DINOv3ClassifierTest(absltest.TestCase):
+ """Tests for the DINOv3Classifier public API."""
+
+ def setUp(self):
+ """Initializes common test fixtures."""
+ super().setUp()
+ self.class_names = ["class_a", "class_b", "class_c"]
+ self.image_transform = dinov3_classifier._build_image_transform(
+ _make_dummy_config()
+ )
+
+ def _make_classifier_with_mock_model(
+ self, mock_model: mock.MagicMock
+ ) -> dinov3_classifier.DINOv3Classifier:
+ """Constructs a DINOv3Classifier wrapping the given mock model."""
+ return dinov3_classifier.DINOv3Classifier(
+ model=mock_model,
+ class_names=self.class_names,
+ image_transform=self.image_transform,
+ device=torch.device("cpu"),
+ )
+
+ def test_predict_batch_empty_input_returns_empty_list(self):
+ """Verifies an empty input returns an empty list without inference."""
+ mock_model = mock.MagicMock()
+ classifier = self._make_classifier_with_mock_model(mock_model)
+ self.assertEmpty(classifier.predict_batch([]))
+ mock_model.assert_not_called()
+
+ def test_predict_batch_returns_top_class_and_percentages(self):
+ """Verifies predictions expose top class and percentage probabilities."""
+ mock_model = mock.MagicMock()
+ mock_model.return_value = torch.tensor([
+ [2.0, 0.0, 0.0], # class_a is the top class.
+ [0.0, 3.0, 0.0], # class_b is the top class.
+ ])
+ classifier = self._make_classifier_with_mock_model(mock_model)
+
+ predictions = classifier.predict_batch([
+ Image.new("RGB", (256, 256), color="red"),
+ Image.new("RGB", (256, 256), color="blue"),
+ ])
+
+ self.assertLen(predictions, 2)
+ self.assertEqual(predictions[0]["predicted_class"], "class_a")
+ self.assertEqual(predictions[1]["predicted_class"], "class_b")
+ self.assertIn("class_c", predictions[0]["all_probabilities_percent"])
+ self.assertGreater(
+ predictions[0]["predicted_probability_percent"],
+ predictions[0]["all_probabilities_percent"]["class_b"],
+ )
+
+ def test_from_config_builds_classifier_successfully(self):
+ """Verifies from_config wires the checkpoint into a ready-to-use model."""
+ mock_hub_load = self.enter_context(
+ mock.patch.object(torch.hub, "load", autospec=True)
+ )
+ mock_torch_load = self.enter_context(
+ mock.patch.object(torch, "load", autospec=True)
+ )
+ mock_hub_load.return_value = DummyBackbone(hidden_size=32)
+ mock_torch_load.return_value = {
+ "model_state_dict": _make_dummy_state_dict(
+ hidden_size=32, number_of_classes=3
+ )
+ }
+
+ classifier = dinov3_classifier.DINOv3Classifier.from_config(
+ config=_make_dummy_config(),
+ class_names=self.class_names,
+ device="cpu",
+ )
+
+ self.assertIsInstance(
+ classifier._model, dinov3_classifier.Dinov3ClassificationModule
+ )
+ self.assertFalse(classifier._model.training)
+
+ def test_from_config_wraps_state_dict_load_errors(self):
+ """Verifies load_state_dict RuntimeError is wrapped as ClassifierError."""
+ mock_hub_load = self.enter_context(
+ mock.patch.object(torch.hub, "load", autospec=True)
+ )
+ mock_torch_load = self.enter_context(
+ mock.patch.object(torch, "load", autospec=True)
+ )
+ mock_hub_load.return_value = DummyBackbone(hidden_size=32)
+ # head.bias shape mismatch triggers a RuntimeError in load_state_dict.
+ incompatible_state_dict = {
+ "backbone_model.norm.weight": torch.ones((32,)),
+ "backbone_model.norm.bias": torch.zeros((32,)),
+ "head.weight": torch.zeros((3, 32)),
+ "head.bias": torch.zeros((999,)),
+ }
+ mock_torch_load.return_value = {"model_state_dict": incompatible_state_dict}
+
+ with self.assertRaisesRegex(
+ dinov3_classifier.ClassifierError,
+ "Failed to load state dict into model",
+ ):
+ dinov3_classifier.DINOv3Classifier.from_config(
+ config=_make_dummy_config(),
+ class_names=self.class_names,
+ device="cpu",
+ )
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/image_ops.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/image_ops.py
new file mode 100644
index 00000000000..42e921e4d88
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/image_ops.py
@@ -0,0 +1,133 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Image loading, resizing, and general system utilities for the SAM3/DINOv3 tracking pipeline."""
+
+from collections.abc import Sequence
+import gc
+import pathlib
+
+import cv2
+from PIL import Image
+import torch
+
+
+class ImageReadError(OSError):
+ """Raised when an image cannot be read from disk.
+
+ Inherits from OSError so callers that broadly catch OSError (e.g. for I/O
+ errors) still catch this, while callers who want to be specific can catch
+ ImageReadError.
+ """
+
+
+def load_and_resize_image(
+ image_path: pathlib.Path, max_short_side: int
+) -> Image.Image:
+ """Loads an image via OpenCV and downscales it using Lanczos interpolation.
+
+ Lanczos is chosen to match the resampling used by the previous pipeline
+ (PIL.Image.LANCZOS), keeping SAM3 inputs close to the previous
+ implementation. The two implementations are not bit-identical (PIL and
+ OpenCV use slightly different Lanczos filter parameters) but are visually
+ equivalent and produce nearly identical downstream detections.
+
+ Args:
+ image_path: Path to the image file on disk.
+ max_short_side: Maximum allowed length of the shortest side. If the image is
+ larger, it is downscaled preserving aspect ratio.
+
+ Returns:
+ The loaded and (optionally) resized image as a PIL RGB Image.
+
+ Raises:
+ ImageReadError: If OpenCV is unable to decode the file (missing, corrupt,
+ or unsupported format).
+ """
+ bgr_image = cv2.imread(str(image_path))
+ if bgr_image is None:
+ raise ImageReadError(
+ f"OpenCV could not read the image at {image_path!s}. The file may be "
+ "missing, corrupt, or in an unsupported format."
+ )
+ height, width = bgr_image.shape[:2]
+ new_width, new_height = _compute_resized_dimensions(
+ width=width, height=height, max_short_side=max_short_side
+ )
+ if (new_width, new_height) != (width, height):
+ bgr_resized = cv2.resize(
+ bgr_image,
+ (new_width, new_height),
+ interpolation=cv2.INTER_LANCZOS4,
+ )
+ else:
+ bgr_resized = bgr_image
+ rgb_resized = cv2.cvtColor(bgr_resized, cv2.COLOR_BGR2RGB)
+ return Image.fromarray(rgb_resized)
+
+
+def ensure_directories_exist(directories: Sequence[pathlib.Path]) -> None:
+ """Creates each directory in the given sequence if it does not already exist.
+
+ Args:
+ directories: Paths to create. Each path must be non-empty and not point to
+ the current working directory (`.`).
+
+ Raises:
+ ValueError: If any path evaluates to `.` (e.g., from `pathlib.Path("")`) or
+ is empty (indicates a configuration bug).
+ """
+ for directory in directories:
+ if directory == pathlib.Path("."):
+ raise ValueError(
+ "ensure_directories_exist received an empty path or the current "
+ "working directory ('.'); this usually indicates a configuration "
+ "error."
+ )
+ directory.mkdir(parents=True, exist_ok=True)
+
+
+def release_caches() -> None:
+ """Runs Python garbage collection and empties the CUDA allocator cache.
+
+ Intended to be called between subfolders in a long pipeline run to release
+ RAM and VRAM held by unreferenced objects and cached allocations. This does
+ not free memory held by live references.
+ """
+ gc.collect()
+ if torch.cuda.is_available():
+ torch.cuda.empty_cache()
+
+
+def _compute_resized_dimensions(
+ width: int, height: int, max_short_side: int
+) -> tuple[int, int]:
+ """Returns the (width, height) that keeps the short side <= `max_short_side`.
+
+ If the image is already small enough, the original dimensions are returned
+ unchanged.
+
+ Args:
+ width: Original image width in pixels.
+ height: Original image height in pixels.
+ max_short_side: Maximum allowed length of the shorter side.
+
+ Returns:
+ (new_width, new_height), preserving the original aspect ratio.
+ """
+ short_side = min(width, height)
+ if short_side <= max_short_side:
+ return width, height
+ scale = max_short_side / short_side
+ return int(width * scale), int(height * scale)
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/image_ops_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/image_ops_test.py
new file mode 100644
index 00000000000..5ca79cbb55c
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/image_ops_test.py
@@ -0,0 +1,189 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for image_ops.py."""
+
+import pathlib
+from unittest import mock
+
+from absl.testing import absltest
+import numpy
+from PIL import Image
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import image_ops
+
+
+class LoadAndResizeImageTest(absltest.TestCase):
+ """Tests for load_and_resize_image."""
+
+ def setUp(self):
+ """Patches cv2 read/resize/cvtColor for deterministic behavior."""
+ super().setUp()
+ self.mock_imread = self.enter_context(
+ mock.patch.object(image_ops.cv2, "imread", autospec=True)
+ )
+ self.mock_resize = self.enter_context(
+ mock.patch.object(image_ops.cv2, "resize", autospec=True)
+ )
+ self.mock_cvt_color = self.enter_context(
+ mock.patch.object(image_ops.cv2, "cvtColor", autospec=True)
+ )
+ # cvtColor always returns a small deterministic RGB array.
+ self.mock_cvt_color.return_value = numpy.zeros((4, 4, 3), dtype=numpy.uint8)
+
+ def test_raises_image_read_error_when_cv2_returns_none(self):
+ """Verifies ImageReadError is raised when imread cannot decode the file."""
+ self.mock_imread.return_value = None
+ with self.assertRaisesRegex(
+ image_ops.ImageReadError, "OpenCV could not read the image"
+ ):
+ image_ops.load_and_resize_image(
+ image_path=pathlib.Path("/nowhere/x.png"), max_short_side=800
+ )
+
+ def test_skips_resize_when_image_is_small_enough(self):
+ """Verifies cv2.resize is not called when short side is within limit."""
+ self.mock_imread.return_value = numpy.zeros(
+ (100, 200, 3), dtype=numpy.uint8
+ )
+ result = image_ops.load_and_resize_image(
+ image_path=pathlib.Path("/x/y.png"), max_short_side=800
+ )
+ self.mock_resize.assert_not_called()
+ self.assertIsInstance(result, Image.Image)
+
+ def test_resizes_when_short_side_exceeds_max(self):
+ """Verifies cv2.resize is called with the scaled dimensions."""
+ # Input is 1000 x 2000, short side is 1000, max is 400 -> scale = 0.4.
+ self.mock_imread.return_value = numpy.zeros(
+ (1000, 2000, 3), dtype=numpy.uint8
+ )
+ self.mock_resize.return_value = numpy.zeros(
+ (400, 800, 3), dtype=numpy.uint8
+ )
+ image_ops.load_and_resize_image(
+ image_path=pathlib.Path("/x/y.png"), max_short_side=400
+ )
+ self.mock_resize.assert_called_once()
+ passed_size = self.mock_resize.call_args.args[1]
+ self.assertEqual(passed_size, (800, 400))
+
+
+class ComputeResizedDimensionsTest(absltest.TestCase):
+ """Tests for the private _compute_resized_dimensions helper."""
+
+ def test_returns_original_dimensions_when_within_limit(self):
+ """Verifies dimensions are returned unchanged when already small enough."""
+ self.assertEqual(
+ image_ops._compute_resized_dimensions(
+ width=200, height=100, max_short_side=800
+ ),
+ (200, 100),
+ )
+
+ def test_scales_down_when_short_side_exceeds_limit(self):
+ """Verifies scale factor is applied to both dimensions."""
+ # 2000 x 1000, short = 1000, limit = 400 -> scale = 0.4 -> (800, 400).
+ self.assertEqual(
+ image_ops._compute_resized_dimensions(
+ width=2000, height=1000, max_short_side=400
+ ),
+ (800, 400),
+ )
+
+ def test_preserves_aspect_ratio_when_scaling(self):
+ """Verifies the ratio width/height is preserved after scaling."""
+ new_width, new_height = image_ops._compute_resized_dimensions(
+ width=1600, height=900, max_short_side=450
+ )
+ # Short side becomes 450; long side scales by the same factor.
+ self.assertEqual((new_width, new_height), (800, 450))
+
+
+class EnsureDirectoriesExistTest(absltest.TestCase):
+ """Tests for ensure_directories_exist."""
+
+ def test_creates_each_directory(self):
+ """Verifies each real directory path is created on the filesystem."""
+ temp_root = pathlib.Path(self.create_tempdir().full_path)
+ dir_a = temp_root / "subdir_a"
+ dir_b = temp_root / "nested" / "subdir_b"
+
+ image_ops.ensure_directories_exist([dir_a, dir_b])
+
+ self.assertTrue(dir_a.exists())
+ self.assertTrue(dir_b.exists())
+ self.assertTrue(dir_a.is_dir())
+ self.assertTrue(dir_b.is_dir())
+
+ def test_raises_value_error_on_empty_or_current_dir_path(self):
+ """Verifies empty paths and '.' are caught as config errors."""
+ with self.assertRaisesRegex(
+ ValueError, "empty path or the current working directory"
+ ):
+ image_ops.ensure_directories_exist([pathlib.Path("")])
+ with self.assertRaisesRegex(
+ ValueError, "empty path or the current working directory"
+ ):
+ image_ops.ensure_directories_exist([pathlib.Path(".")])
+
+
+class ReleaseCachesTest(absltest.TestCase):
+ """Tests for release_caches."""
+
+ def test_runs_gc_and_skips_cuda_when_unavailable(self):
+ """Verifies gc.collect runs and cuda.empty_cache is skipped without GPU."""
+ mock_gc_collect = self.enter_context(
+ mock.patch.object(image_ops.gc, "collect", autospec=True)
+ )
+ self.enter_context(
+ mock.patch.object(
+ image_ops.torch.cuda,
+ "is_available",
+ autospec=True,
+ return_value=False,
+ )
+ )
+ mock_empty_cache = self.enter_context(
+ mock.patch.object(image_ops.torch.cuda, "empty_cache", autospec=True)
+ )
+
+ image_ops.release_caches()
+
+ mock_gc_collect.assert_called_once()
+ mock_empty_cache.assert_not_called()
+
+ def test_calls_cuda_empty_cache_when_available(self):
+ """Verifies cuda.empty_cache is called when CUDA is available."""
+ self.enter_context(
+ mock.patch.object(image_ops.gc, "collect", autospec=True)
+ )
+ self.enter_context(
+ mock.patch.object(
+ image_ops.torch.cuda,
+ "is_available",
+ autospec=True,
+ return_value=True,
+ )
+ )
+ mock_empty_cache = self.enter_context(
+ mock.patch.object(image_ops.torch.cuda, "empty_cache", autospec=True)
+ )
+
+ image_ops.release_caches()
+
+ mock_empty_cache.assert_called_once()
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/pipeline_logger.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/pipeline_logger.py
new file mode 100644
index 00000000000..59c9eab10a9
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/pipeline_logger.py
@@ -0,0 +1,101 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Logging configuration and handlers for the SAM3/DINOv3 tracking pipeline."""
+
+import logging
+import pathlib
+import sys
+
+_LOGGER_NAME = "waste_identification_pipeline"
+_LOG_FORMAT = "[%(asctime)s] %(levelname)s %(message)s"
+_LOG_DATE_FORMAT = "%Y-%m-%d %H:%M:%S"
+
+
+def get_logger() -> logging.Logger:
+ """Returns the shared, configured pipeline logger.
+
+ The logger writes to stdout, bypasses third-party root loggers
+ (`propagate=False`), and is idempotent: repeated calls return the same
+ instance without adding duplicate handlers.
+
+ Returns:
+ The shared pipeline logger instance.
+ """
+ logger = logging.getLogger(_LOGGER_NAME)
+ if logger.handlers:
+ return logger
+ logger.setLevel(logging.INFO)
+ stream_handler = logging.StreamHandler(stream=sys.stdout)
+ stream_handler.setFormatter(
+ logging.Formatter(_LOG_FORMAT, datefmt=_LOG_DATE_FORMAT)
+ )
+ logger.addHandler(stream_handler)
+ logger.propagate = False
+ return logger
+
+
+def attach_file_handler(log_file_path: pathlib.Path) -> logging.FileHandler:
+ """Attaches a file handler to the pipeline logger.
+
+ The new handler inherits the logger's current level and reuses the stream
+ handler's formatter so file logs match console logs. If no stream handler
+ is present, the default pipeline formatter is used instead.
+
+ Args:
+ log_file_path: Path where the log file will be written.
+
+ Returns:
+ The attached FileHandler so the caller can detach it later.
+ """
+ logger = get_logger()
+ file_handler = logging.FileHandler(
+ str(log_file_path), mode="w", encoding="utf-8"
+ )
+ file_handler.setLevel(logger.level)
+ file_handler.setFormatter(_get_stream_formatter(logger))
+ logger.addHandler(file_handler)
+ return file_handler
+
+
+def detach_file_handler(file_handler: logging.FileHandler) -> None:
+ """Removes a file handler from the pipeline logger and closes it.
+
+ Args:
+ file_handler: The handler previously returned by `attach_file_handler`.
+ """
+ logger = get_logger()
+ logger.removeHandler(file_handler)
+ file_handler.close()
+
+
+def _get_stream_formatter(logger: logging.Logger) -> logging.Formatter:
+ """Returns the formatter used by the first StreamHandler on the logger.
+
+ Falls back to the default pipeline formatter if no StreamHandler is
+ attached or the attached one has no formatter.
+
+ Args:
+ logger: Logger to inspect.
+
+ Returns:
+ A logging.Formatter suitable for consistent log output.
+ """
+ for handler in logger.handlers:
+ if (
+ isinstance(handler, logging.StreamHandler)
+ and handler.formatter is not None
+ ):
+ return handler.formatter
+ return logging.Formatter(_LOG_FORMAT, datefmt=_LOG_DATE_FORMAT)
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/pipeline_logger_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/pipeline_logger_test.py
new file mode 100644
index 00000000000..1dad6d80ab4
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/pipeline_logger_test.py
@@ -0,0 +1,108 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for pipeline_logger.py."""
+
+import logging
+import pathlib
+
+from absl.testing import absltest
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import pipeline_logger
+
+
+class GetLoggerTest(absltest.TestCase):
+ """Tests for get_logger."""
+
+ def setUp(self):
+ """Detaches all handlers before each test so the logger starts clean."""
+ super().setUp()
+ logger = logging.getLogger(pipeline_logger._LOGGER_NAME)
+ # Iterate over a shallow copy (`list(...)`) so removing handlers does not
+ # mutate the list during iteration.
+ for handler in list(logger.handlers):
+ logger.removeHandler(handler)
+ self.addCleanup(logger.handlers.clear)
+
+ def test_returns_configured_logger_with_stream_handler(self):
+ """Verifies a first call attaches exactly one StreamHandler."""
+ logger = pipeline_logger.get_logger()
+ stream_handlers = [
+ handler
+ for handler in logger.handlers
+ if isinstance(handler, logging.StreamHandler)
+ ]
+ self.assertLen(stream_handlers, 1)
+ self.assertEqual(logger.level, logging.INFO)
+ self.assertFalse(logger.propagate)
+
+ def test_repeated_calls_do_not_duplicate_handlers(self):
+ """Verifies the logger is idempotent across many get_logger calls."""
+ first = pipeline_logger.get_logger()
+ handler_count = len(first.handlers)
+ for _ in range(5):
+ pipeline_logger.get_logger()
+ self.assertLen(first.handlers, handler_count)
+
+
+class AttachAndDetachFileHandlerTest(absltest.TestCase):
+ """Tests for attach_file_handler and detach_file_handler."""
+
+ def setUp(self):
+ """Resets the pipeline logger before each test."""
+ super().setUp()
+ logger = logging.getLogger(pipeline_logger._LOGGER_NAME)
+ # Iterate over a shallow copy (`list(...)`) so removing handlers does not
+ # mutate the list during iteration.
+ for handler in list(logger.handlers):
+ logger.removeHandler(handler)
+ self.addCleanup(logger.handlers.clear)
+
+ def test_attach_adds_file_handler_with_matching_formatter(self):
+ """Verifies the attached FileHandler reuses the stream handler formatter."""
+ pipeline_logger.get_logger() # Ensure a stream handler is present.
+ log_path = pathlib.Path(self.create_tempdir().full_path) / "run.log"
+
+ file_handler = pipeline_logger.attach_file_handler(log_path)
+ try:
+ logger = pipeline_logger.get_logger()
+ self.assertIn(file_handler, logger.handlers)
+ stream_handler = next(
+ handler
+ for handler in logger.handlers
+ if isinstance(handler, logging.StreamHandler)
+ and not isinstance(handler, logging.FileHandler)
+ )
+ self.assertIs(file_handler.formatter, stream_handler.formatter)
+ self.assertEqual(file_handler.level, logger.level)
+ finally:
+ pipeline_logger.detach_file_handler(file_handler)
+
+ def test_detach_removes_and_closes_handler(self):
+ """Verifies detach_file_handler removes the handler and closes it."""
+ pipeline_logger.get_logger()
+ log_path = pathlib.Path(self.create_tempdir().full_path) / "run.log"
+ file_handler = pipeline_logger.attach_file_handler(log_path)
+
+ stream = file_handler.stream
+ self.assertIsNotNone(stream)
+ pipeline_logger.detach_file_handler(file_handler)
+
+ logger = pipeline_logger.get_logger()
+ self.assertNotIn(file_handler, logger.handlers)
+ self.assertTrue(stream.closed)
+ self.assertIsNone(file_handler.stream)
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/sam3_detector.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/sam3_detector.py
new file mode 100644
index 00000000000..4fe34193de4
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/sam3_detector.py
@@ -0,0 +1,408 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""SAM3 detection and segmentation management."""
+
+import collections
+from collections.abc import Mapping
+import logging
+import pathlib
+from typing import Any, Self
+
+import numpy
+from PIL import Image
+import supervision
+import torch
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import config_loader
+
+try:
+ from sam3 import model_builder as sam3_model_builder # pylint: disable=g-import-not-at-top
+ from sam3.model import sam3_image_processor # pylint: disable=g-import-not-at-top
+except ImportError:
+ sam3_model_builder = None
+ sam3_image_processor = None
+
+_LOGGER = logging.getLogger(__name__)
+
+# Prompt whose detections should have their contained boxes merged.
+_MERGE_PROMPT = "packets"
+
+# Intermediate state entries dropped before returning to callers to reduce
+# memory footprint. These are set by the SAM3 processor during inference and
+# are not needed downstream.
+_INFERENCE_KEYS_TO_DROP = frozenset(
+ ["backbone_out", "geometric_prompt", "image_embeddings"]
+)
+
+# State entries that are per-detection arrays, kept in lockstep after any
+# filtering step.
+_STATE_ARRAY_KEYS = frozenset({"masks", "masks_logits", "boxes", "scores"})
+
+
+class SAM3Detector:
+ """Runs SAM3 inference and post-processes detections.
+
+ The detector loads the SAM3 model onto the resolved device and exposes
+ `detect` for running inference on a single image with a text prompt.
+ """
+
+ def __init__(
+ self,
+ model: torch.nn.Module,
+ processor: Any,
+ device: torch.device,
+ prompt_config: config_loader.PromptConfig,
+ ):
+ """Initializes the SAM3 detector with pre-built dependencies.
+
+ Args:
+ model: The SAM3 PyTorch image model.
+ processor: The SAM3 image processor instance (`Sam3Processor`).
+ device: Resolved PyTorch device where the model resides (`torch.device`).
+ prompt_config: Detection and post-processing thresholds (`PromptConfig`).
+ """
+ self._model = model
+ self._processor = processor
+ self._device = device
+ self._config = prompt_config
+
+ @classmethod
+ def from_checkpoint(
+ cls,
+ checkpoint_path: pathlib.Path,
+ device: str,
+ prompt_config: config_loader.PromptConfig,
+ ) -> Self:
+ """Constructs a SAM3Detector by loading a model from a checkpoint.
+
+ Args:
+ checkpoint_path: Path to the SAM3 model checkpoint.
+ device: Requested device string (e.g., 'cuda', 'cpu', 'mps').
+ prompt_config: Detection and post-processing thresholds for the active
+ prompt.
+
+ Returns:
+ An initialized `SAM3Detector` instance on the resolved device.
+
+ Raises:
+ ImportError: If the 'sam3' package is not installed or available.
+ """
+ if sam3_model_builder is None or sam3_image_processor is None:
+ raise ImportError(
+ "The 'sam3' package is not installed or available in python path. "
+ "Cannot construct SAM3Detector from checkpoint."
+ )
+
+ resolved_device = _resolve_device(device)
+ model = sam3_model_builder.build_sam3_image_model(
+ checkpoint_path=str(checkpoint_path)
+ )
+ model.to(device=resolved_device)
+ processor = sam3_image_processor.Sam3Processor(
+ model,
+ confidence_threshold=prompt_config.confidence_threshold,
+ )
+ return cls(
+ model=model,
+ processor=processor,
+ device=resolved_device,
+ prompt_config=prompt_config,
+ )
+
+ def detect(self, image: Image.Image, prompt: str) -> dict[str, torch.Tensor]:
+ """Runs SAM3 inference and returns a CPU-bound state dictionary.
+
+ Post-processing (containment filtering, and for the merge prompt also box
+ merging) is performed on the active device to keep tensor operations on
+ the GPU where possible. The final state is moved to CPU before return.
+
+ Args:
+ image: PIL RGB image to detect on.
+ prompt: Text prompt to condition the detector.
+
+ Returns:
+ A dict containing per-detection tensors on CPU. Guaranteed keys include
+ those in _STATE_ARRAY_KEYS.
+ """
+ with torch.no_grad(), torch.autocast(
+ self._device.type, dtype=torch.float16
+ ):
+ state = self._processor.set_image(image)
+ state = self._processor.set_text_prompt(state=state, prompt=prompt)
+
+ # Post-process on the active device before the CPU transfer to avoid
+ # unnecessary device round-trips.
+ state = self._filter_contained_sub_masks(state)
+
+ if prompt == _MERGE_PROMPT:
+ state = self._merge_contained_boxes(state)
+
+ for key in _INFERENCE_KEYS_TO_DROP:
+ state.pop(key, None)
+
+ return self._move_state_to_cpu(state)
+
+ def to_supervision_detections(
+ self, state: Mapping[str, torch.Tensor]
+ ) -> supervision.Detections:
+ """Converts a SAM3 state dict into supervision.Detections.
+
+ Boxes whose score falls below `prompt_config.score_threshold` are
+ filtered out. Returns an empty Detections object when no box survives.
+
+ Args:
+ state: State dictionary as returned by `detect`.
+
+ Returns:
+ Detections containing the surviving boxes and their confidences.
+ """
+ boxes = state["boxes"].numpy().astype(numpy.float32)
+ scores = state["scores"].numpy().astype(numpy.float32)
+
+ # SAM3 uses a 1-D empty tensor to signal "no detections"; normalize to a
+ # (0, 4) shape so the slice below stays valid.
+ if boxes.ndim == 1:
+ boxes = boxes.reshape(0, 4) if boxes.size == 0 else boxes.reshape(-1, 4)
+
+ keep_mask = scores >= self._config.score_threshold
+ kept_boxes = boxes[keep_mask]
+ kept_scores = scores[keep_mask]
+
+ if kept_boxes.shape[0] == 0:
+ return supervision.Detections.empty()
+
+ return supervision.Detections(
+ xyxy=kept_boxes,
+ confidence=kept_scores,
+ class_id=numpy.zeros(kept_boxes.shape[0], dtype=int),
+ )
+
+ def _move_state_to_cpu(
+ self, inference_state: dict[str, Any]
+ ) -> dict[str, Any]:
+ """Recursively moves every tensor in the state dict to CPU in place.
+
+ Args:
+ inference_state: State dict returned by the SAM3 processor. Mutated in
+ place.
+
+ Returns:
+ The same dict reference, for call-site convenience.
+ """
+ for key, value in inference_state.items():
+ if isinstance(value, torch.Tensor):
+ inference_state[key] = value.cpu()
+ elif isinstance(value, dict):
+ self._move_state_to_cpu(value)
+ return inference_state
+
+ def _filter_contained_sub_masks(
+ self, state: dict[str, torch.Tensor]
+ ) -> dict[str, torch.Tensor]:
+ """Removes smaller masks that are largely contained in bigger masks.
+
+ Args:
+ state: State dict; mutated in place for the per-detection arrays.
+
+ Returns:
+ The same state dict reference.
+ """
+ masks = state["masks"]
+ number_of_masks = masks.shape[0]
+ if number_of_masks == 0:
+ return state
+
+ flat_masks = masks.view(number_of_masks, -1).float()
+ pairwise_intersection = flat_masks @ flat_masks.T
+
+ # TODO(umairsabir): Vectorize containment filtering using PyTorch's indexing
+ # and boolean operations over pairwise_intersection if nested loops become a
+ # bottleneck under large batches or many prompts.
+ indices_to_remove: set[int] = set()
+ for outer_index in range(number_of_masks):
+ if outer_index in indices_to_remove:
+ continue
+ for inner_index in range(outer_index + 1, number_of_masks):
+ if inner_index in indices_to_remove:
+ continue
+
+ outer_area = pairwise_intersection[outer_index, outer_index].item()
+ inner_area = pairwise_intersection[inner_index, inner_index].item()
+ if outer_area <= inner_area:
+ smaller_index = outer_index
+ smaller_area = outer_area
+ else:
+ smaller_index = inner_index
+ smaller_area = inner_area
+
+ if smaller_area == 0:
+ indices_to_remove.add(smaller_index)
+ if smaller_index == outer_index:
+ break
+ continue
+
+ intersection = pairwise_intersection[outer_index, inner_index].item()
+ if (intersection / smaller_area) > self._config.containment_threshold:
+ indices_to_remove.add(smaller_index)
+ if smaller_index == outer_index:
+ break
+
+ kept_indices = sorted(set(range(number_of_masks)) - indices_to_remove)
+ keep_tensor = torch.tensor(
+ kept_indices, dtype=torch.long, device=masks.device
+ )
+
+ for key in _STATE_ARRAY_KEYS:
+ state[key] = state[key][keep_tensor]
+
+ return state
+
+ def _merge_contained_boxes(
+ self, state: dict[str, torch.Tensor]
+ ) -> dict[str, torch.Tensor]:
+ """Merges detections where a smaller box is contained in a larger one.
+
+ Merging combines masks (element-wise OR), takes the axis-aligned union
+ of boxes, and sums the scores (clamped to 1.0).
+
+ Args:
+ state: State dict; mutated in place.
+
+ Returns:
+ The same state dict reference.
+ """
+ masks = state["masks"]
+ boxes = state["boxes"]
+ scores = state["scores"]
+ if len(scores) == 0:
+ return state
+
+ number_of_detections = len(masks)
+ box_areas = (boxes[:, 2] - boxes[:, 0]) * (boxes[:, 3] - boxes[:, 1])
+
+ is_absorbed = torch.zeros(
+ number_of_detections, dtype=torch.bool, device=boxes.device
+ )
+ absorb_target = list(range(number_of_detections))
+
+ for outer_index in range(number_of_detections):
+ if is_absorbed[outer_index]:
+ continue
+ for inner_index in range(outer_index + 1, number_of_detections):
+ if is_absorbed[inner_index]:
+ continue
+
+ intersection_x_min = torch.max(
+ boxes[outer_index, 0], boxes[inner_index, 0]
+ )
+ intersection_y_min = torch.max(
+ boxes[outer_index, 1], boxes[inner_index, 1]
+ )
+ intersection_x_max = torch.min(
+ boxes[outer_index, 2], boxes[inner_index, 2]
+ )
+ intersection_y_max = torch.min(
+ boxes[outer_index, 3], boxes[inner_index, 3]
+ )
+ intersection_area = torch.clamp(
+ intersection_x_max - intersection_x_min, min=0
+ ) * torch.clamp(intersection_y_max - intersection_y_min, min=0)
+
+ if box_areas[outer_index] <= box_areas[inner_index]:
+ smaller_index = outer_index
+ larger_index = inner_index
+ smaller_area = box_areas[outer_index]
+ else:
+ smaller_index = inner_index
+ larger_index = outer_index
+ smaller_area = box_areas[inner_index]
+
+ if smaller_area == 0:
+ is_absorbed[smaller_index] = True
+ absorb_target[smaller_index] = larger_index
+ if smaller_index == outer_index:
+ break
+ continue
+
+ containment_ratio = intersection_area / smaller_area
+ if containment_ratio > self._config.merge_overlap_threshold:
+ is_absorbed[smaller_index] = True
+ absorb_target[smaller_index] = larger_index
+ if smaller_index == outer_index:
+ break
+
+ groups: collections.defaultdict[int, list[int]] = collections.defaultdict(
+ list
+ )
+ for detection_index in range(number_of_detections):
+ if is_absorbed[detection_index]:
+ target = absorb_target[detection_index]
+ else:
+ target = detection_index
+ groups[target].append(detection_index)
+
+ merged_masks: list[torch.Tensor] = []
+ merged_boxes: list[torch.Tensor] = []
+ merged_scores: list[torch.Tensor] = []
+ for member_indices in groups.values():
+ member_tensor = torch.tensor(
+ member_indices, dtype=torch.long, device=boxes.device
+ )
+ merged_masks.append(masks[member_tensor].squeeze(1).any(dim=0))
+ group_boxes = boxes[member_tensor]
+ merged_boxes.append(
+ torch.stack([
+ group_boxes[:, 0].min(),
+ group_boxes[:, 1].min(),
+ group_boxes[:, 2].max(),
+ group_boxes[:, 3].max(),
+ ])
+ )
+ summed_score = scores[member_tensor].sum().item()
+ merged_scores.append(
+ torch.tensor(min(summed_score, 1.0), device=scores.device)
+ )
+
+ state["masks"] = torch.stack(merged_masks).unsqueeze(1)
+ state["boxes"] = torch.stack(merged_boxes)
+ state["scores"] = torch.stack(merged_scores)
+ return state
+
+
+def _resolve_device(requested_device: str) -> torch.device:
+ """Resolves the requested device string to an available torch.device.
+
+ Falls back to CPU with a warning if the requested device is not available.
+
+ Args:
+ requested_device: Device string requested by the caller.
+
+ Returns:
+ The resolved torch.device.
+ """
+ if requested_device.startswith("cuda") and not torch.cuda.is_available():
+ _LOGGER.warning(
+ "Requested device '%s' but CUDA is not available; falling back to CPU.",
+ requested_device,
+ )
+ return torch.device("cpu")
+
+ if requested_device == "mps" and not torch.backends.mps.is_available():
+ _LOGGER.warning(
+ "Requested device 'mps' but MPS is not available; falling back to CPU."
+ )
+ return torch.device("cpu")
+
+ return torch.device(requested_device)
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/sam3_detector_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/sam3_detector_test.py
new file mode 100644
index 00000000000..d39374172ab
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/sam3_detector_test.py
@@ -0,0 +1,570 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for sam3_detector.py."""
+
+import pathlib
+import sys
+from unittest import mock
+
+from absl.testing import absltest
+from absl.testing import parameterized
+import numpy
+from PIL import Image
+import torch
+
+# Mock supervision before it is imported anywhere since it is an external
+# pip package not checked into //third_party/py.
+mock_supervision = mock.MagicMock()
+
+
+class MockDetections:
+
+ def __init__(self, xyxy=None, confidence=None, class_id=None):
+ self.xyxy = xyxy
+ self.confidence = confidence
+ self.class_id = class_id
+ self.tracker_id = None
+
+ def __len__(self):
+ return len(self.xyxy) if self.xyxy is not None else 0
+
+ @classmethod
+ def empty(cls):
+ return cls(
+ xyxy=numpy.zeros((0, 4), dtype=numpy.float32),
+ confidence=numpy.zeros(0, dtype=numpy.float32),
+ class_id=numpy.zeros(0, dtype=int),
+ )
+
+
+mock_supervision.Detections = MockDetections
+sys.modules["supervision"] = mock_supervision
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import config_loader # pylint: disable=g-bad-import-order, g-import-not-at-top
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import sam3_detector # pylint: disable=g-bad-import-order, g-import-not-at-top
+
+
+def _make_prompt_config(
+ score_threshold: float = 0.3,
+ containment_threshold: float = 0.9,
+ merge_overlap_threshold: float = 0.7,
+) -> config_loader.PromptConfig:
+ """Returns a PromptConfig with values suitable for tests."""
+ return config_loader.PromptConfig(
+ confidence_threshold=0.5,
+ score_threshold=score_threshold,
+ containment_threshold=containment_threshold,
+ max_short_side=800,
+ crop_size=(224, 224),
+ crop_buffer_pixels=10,
+ merge_overlap_threshold=merge_overlap_threshold,
+ )
+
+
+def _make_detector_without_construction(
+ prompt_config: config_loader.PromptConfig | None = None,
+) -> sam3_detector.SAM3Detector:
+ """Returns a SAM3Detector instance initialized with mock dependencies.
+
+ Args:
+ prompt_config: Optional prompt configuration to inject into the detector.
+ """
+ return sam3_detector.SAM3Detector(
+ model=mock.MagicMock(),
+ processor=mock.MagicMock(),
+ device=torch.device("cpu"),
+ prompt_config=prompt_config or _make_prompt_config(),
+ )
+
+
+class ResolveDeviceTest(parameterized.TestCase):
+ """Tests for the module-level _resolve_device helper."""
+
+ @parameterized.named_parameters(
+ dict(
+ testcase_name="cpu_resolves_directly",
+ requested_device="cpu",
+ cuda_available=None,
+ mps_available=None,
+ expected_device=torch.device("cpu"),
+ ),
+ dict(
+ testcase_name="cuda_resolves_when_available",
+ requested_device="cuda:0",
+ cuda_available=True,
+ mps_available=None,
+ expected_device=torch.device("cuda:0"),
+ ),
+ dict(
+ testcase_name="cuda_falls_back_to_cpu_when_unavailable",
+ requested_device="cuda:0",
+ cuda_available=False,
+ mps_available=None,
+ expected_device=torch.device("cpu"),
+ ),
+ dict(
+ testcase_name="mps_falls_back_to_cpu_when_unavailable",
+ requested_device="mps",
+ cuda_available=None,
+ mps_available=False,
+ expected_device=torch.device("cpu"),
+ ),
+ )
+ def test_resolve_device(
+ self,
+ requested_device: str,
+ cuda_available: bool | None,
+ mps_available: bool | None,
+ expected_device: torch.device,
+ ):
+ """Verifies device resolution and fallback conditions."""
+ if cuda_available is not None:
+ self.enter_context(
+ mock.patch.object(
+ sam3_detector.torch.cuda,
+ "is_available",
+ autospec=True,
+ return_value=cuda_available,
+ )
+ )
+ if mps_available is not None:
+ self.enter_context(
+ mock.patch.object(
+ sam3_detector.torch.backends.mps,
+ "is_available",
+ autospec=True,
+ return_value=mps_available,
+ )
+ )
+ self.assertEqual(
+ sam3_detector._resolve_device(requested_device), expected_device
+ )
+
+
+class FromCheckpointTest(absltest.TestCase):
+ """Tests for SAM3Detector.from_checkpoint and __init__."""
+
+ def setUp(self):
+ """Patches SAM3 model construction and processor construction."""
+ super().setUp()
+ self.mock_sam3_model_builder = mock.MagicMock()
+ self.mock_sam3_image_processor = mock.MagicMock()
+ self.enter_context(
+ mock.patch.object(
+ sam3_detector,
+ "sam3_model_builder",
+ self.mock_sam3_model_builder,
+ )
+ )
+ self.enter_context(
+ mock.patch.object(
+ sam3_detector,
+ "sam3_image_processor",
+ self.mock_sam3_image_processor,
+ )
+ )
+ self.mock_build_model = self.mock_sam3_model_builder.build_sam3_image_model
+ self.mock_processor_class = self.mock_sam3_image_processor.Sam3Processor
+ self.fake_model = mock.MagicMock()
+ self.mock_build_model.return_value = self.fake_model
+
+ def test_raises_import_error_when_sam3_not_available(self):
+ """Verifies from_checkpoint raises ImportError if sam3 is missing."""
+ with mock.patch.object(sam3_detector, "sam3_model_builder", None):
+ with self.assertRaises(ImportError):
+ sam3_detector.SAM3Detector.from_checkpoint(
+ checkpoint_path=pathlib.Path("/tmp/ckpt.pt"),
+ device="cpu",
+ prompt_config=_make_prompt_config(),
+ )
+
+ def test_moves_model_to_resolved_device(self):
+ """Verifies the loaded model is moved to the resolved torch.device."""
+ sam3_detector.SAM3Detector.from_checkpoint(
+ checkpoint_path=pathlib.Path("/tmp/ckpt.pt"),
+ device="cpu",
+ prompt_config=_make_prompt_config(),
+ )
+ self.mock_build_model.assert_called_once_with(
+ checkpoint_path="/tmp/ckpt.pt"
+ )
+ self.fake_model.to.assert_called_once_with(device=torch.device("cpu"))
+
+ def test_processor_is_constructed_with_confidence_threshold(self):
+ """Verifies the processor receives the prompt's confidence threshold."""
+ prompt_config = _make_prompt_config()
+ sam3_detector.SAM3Detector.from_checkpoint(
+ checkpoint_path=pathlib.Path("/tmp/ckpt.pt"),
+ device="cpu",
+ prompt_config=prompt_config,
+ )
+ self.mock_processor_class.assert_called_once_with(
+ self.fake_model,
+ confidence_threshold=prompt_config.confidence_threshold,
+ )
+
+ def test_init_sets_attributes_directly(self):
+ """Verifies __init__ directly stores dependencies without I/O."""
+ model = mock.MagicMock()
+ processor = mock.MagicMock()
+ device = torch.device("cpu")
+ prompt_config = _make_prompt_config()
+ detector = sam3_detector.SAM3Detector(
+ model=model,
+ processor=processor,
+ device=device,
+ prompt_config=prompt_config,
+ )
+ self.assertIs(detector._model, model)
+ self.assertIs(detector._processor, processor)
+ self.assertIs(detector._device, device)
+ self.assertIs(detector._config, prompt_config)
+
+
+class ToSupervisionDetectionsTest(absltest.TestCase):
+ """Tests for SAM3Detector.to_supervision_detections."""
+
+ def test_returns_empty_when_boxes_tensor_is_one_dimensional(self):
+ """Verifies a 1-D boxes tensor is treated as an empty detection set."""
+ detector = _make_detector_without_construction()
+ state = {
+ "boxes": torch.zeros(0),
+ "scores": torch.zeros(0),
+ }
+ detections = detector.to_supervision_detections(state)
+ self.assertEmpty(detections)
+
+ def test_filters_out_boxes_below_score_threshold(self):
+ """Verifies boxes with score < threshold are dropped."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(score_threshold=0.5)
+ )
+ state = {
+ "boxes": torch.tensor([
+ [0.0, 0.0, 10.0, 10.0],
+ [5.0, 5.0, 15.0, 15.0],
+ ]),
+ "scores": torch.tensor([0.9, 0.2]),
+ }
+ detections = detector.to_supervision_detections(state)
+ self.assertLen(detections, 1)
+ numpy.testing.assert_array_equal(
+ detections.xyxy, numpy.array([[0.0, 0.0, 10.0, 10.0]])
+ )
+ numpy.testing.assert_array_equal(
+ detections.confidence, numpy.array([0.9], dtype=numpy.float32)
+ )
+
+ def test_returns_empty_when_all_boxes_are_below_threshold(self):
+ """Verifies an empty Detections is returned when nothing passes filter."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(score_threshold=0.99)
+ )
+ state = {
+ "boxes": torch.tensor([[0.0, 0.0, 10.0, 10.0]]),
+ "scores": torch.tensor([0.5]),
+ }
+ detections = detector.to_supervision_detections(state)
+ self.assertEmpty(detections)
+
+
+class MoveStateToCpuTest(absltest.TestCase):
+ """Tests for SAM3Detector._move_state_to_cpu."""
+
+ def test_moves_top_level_tensors_to_cpu(self):
+ """Verifies each tensor value at the top level is moved to CPU."""
+ detector = _make_detector_without_construction()
+ tensor_on_cpu = torch.zeros(2)
+ state = {"boxes": tensor_on_cpu, "scores": torch.ones(3)}
+ result = detector._move_state_to_cpu(state)
+ self.assertIs(result, state)
+ self.assertEqual(result["boxes"].device.type, "cpu")
+ self.assertEqual(result["scores"].device.type, "cpu")
+
+ def test_recurses_into_nested_dicts(self):
+ """Verifies nested-dict tensors are also moved."""
+ detector = _make_detector_without_construction()
+ state = {"outer": {"inner": torch.ones(2)}}
+ detector._move_state_to_cpu(state)
+ self.assertEqual(state["outer"]["inner"].device.type, "cpu")
+
+ def test_leaves_non_tensor_values_untouched(self):
+ """Verifies non-tensor values are passed through unchanged."""
+ detector = _make_detector_without_construction()
+ state = {"name": "detection", "boxes": torch.zeros(1)}
+ detector._move_state_to_cpu(state)
+ self.assertEqual(state["name"], "detection")
+
+
+class FilterContainedSubMasksTest(absltest.TestCase):
+ """Tests for SAM3Detector._filter_contained_sub_masks."""
+
+ def _make_state(self, masks: torch.Tensor) -> dict[str, torch.Tensor]:
+ """Returns a state dict with matching-length arrays for filtering."""
+ count = masks.shape[0]
+ return {
+ "masks": masks,
+ "masks_logits": torch.zeros_like(masks, dtype=torch.float32),
+ "boxes": torch.zeros((count, 4), dtype=torch.float32),
+ "scores": torch.ones(count, dtype=torch.float32),
+ }
+
+ def test_returns_unchanged_state_for_empty_input(self):
+ """Verifies an empty masks tensor is returned unchanged."""
+ detector = _make_detector_without_construction()
+ state = self._make_state(torch.zeros((0, 4, 4), dtype=torch.bool))
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 0)
+
+ def test_removes_smaller_mask_when_contained(self):
+ """Verifies a smaller mask fully contained in a larger one is removed."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(containment_threshold=0.9)
+ )
+ # A 4x4 large mask fully containing a 2x2 small mask (both non-zero
+ # everywhere they cover). The small mask sits at rows 1-2, cols 1-2.
+ large_mask = torch.ones((4, 4), dtype=torch.bool)
+ small_mask = torch.zeros((4, 4), dtype=torch.bool)
+ small_mask[1:3, 1:3] = True
+ masks = torch.stack([large_mask, small_mask])
+ state = self._make_state(masks)
+
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 1)
+ # The remaining mask should be the larger one (all True).
+ self.assertTrue(bool(result["masks"][0].all()))
+
+ def test_removes_smaller_mask_when_contained_first(self):
+ """Verifies a smaller mask appearing first is removed and breaks early."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(containment_threshold=0.9)
+ )
+ large_mask = torch.ones((4, 4), dtype=torch.bool)
+ small_mask = torch.zeros((4, 4), dtype=torch.bool)
+ small_mask[1:3, 1:3] = True
+ # Small mask first, so smaller_index == outer_index (0) when compared
+ # with large.
+ masks = torch.stack([small_mask, large_mask])
+ state = self._make_state(masks)
+
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 1)
+ # The remaining mask should be the larger one (all True).
+ self.assertTrue(bool(result["masks"][0].all()))
+
+ def test_keeps_both_when_containment_below_threshold(self):
+ """Verifies masks with negligible overlap are both kept."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(containment_threshold=0.9)
+ )
+ # Two 2x2 masks with only a 1-pixel overlap (25% containment of the
+ # smaller, well below the 90% threshold).
+ mask_a = torch.zeros((4, 4), dtype=torch.bool)
+ mask_a[0:2, 0:2] = True
+ mask_b = torch.zeros((4, 4), dtype=torch.bool)
+ mask_b[1:3, 1:3] = True
+ state = self._make_state(torch.stack([mask_a, mask_b]))
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 2)
+
+ def test_skips_inner_mask_if_already_marked_for_removal(self):
+ """Verifies that an already removed inner mask is skipped in nested loop."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(containment_threshold=0.9)
+ )
+ # Mask 0 (large, top half) does not overlap with mask 1 (bottom half).
+ # Mask 2 (small, top-left 1x1) is fully contained inside mask 0.
+ # When outer_index=0: mask 2 is added to indices_to_remove.
+ # When outer_index=1: inner_index=2 is checked and skipped via continue.
+ mask_0 = torch.zeros((4, 4), dtype=torch.bool)
+ mask_0[0:2, 0:4] = True
+ mask_1 = torch.zeros((4, 4), dtype=torch.bool)
+ mask_1[2:4, 0:4] = True
+ mask_2 = torch.zeros((4, 4), dtype=torch.bool)
+ mask_2[0:1, 0:1] = True
+ state = self._make_state(torch.stack([mask_0, mask_1, mask_2]))
+
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 2)
+
+ def test_removes_zero_area_mask(self):
+ """Verifies that a mask with zero area is removed during comparison."""
+ detector = _make_detector_without_construction()
+ mask_a = torch.ones((4, 4), dtype=torch.bool)
+ mask_b = torch.zeros((4, 4), dtype=torch.bool)
+ state = self._make_state(torch.stack([mask_a, mask_b]))
+
+ result = detector._filter_contained_sub_masks(state)
+ self.assertEqual(result["masks"].shape[0], 1)
+ self.assertTrue(bool(result["masks"][0].all()))
+
+
+class MergeContainedBoxesTest(absltest.TestCase):
+ """Tests for SAM3Detector._merge_contained_boxes."""
+
+ def _make_state(
+ self,
+ boxes: torch.Tensor,
+ scores: torch.Tensor,
+ masks: torch.Tensor | None = None,
+ ) -> dict[str, torch.Tensor]:
+ """Returns a state dict compatible with the merge pass."""
+ if masks is None:
+ # One-channel masks whose shape matches the number of boxes.
+ masks = torch.zeros((boxes.shape[0], 1, 4, 4), dtype=torch.bool)
+ return {"masks": masks, "boxes": boxes, "scores": scores}
+
+ def test_returns_unchanged_state_for_empty_input(self):
+ """Verifies an empty scores tensor is returned unchanged."""
+ detector = _make_detector_without_construction()
+ state = self._make_state(
+ boxes=torch.zeros((0, 4), dtype=torch.float32),
+ scores=torch.zeros(0, dtype=torch.float32),
+ )
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["scores"].shape[0], 0)
+
+ def test_merges_small_box_into_containing_larger_box(self):
+ """Verifies a smaller box fully inside a larger one is merged with it."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(merge_overlap_threshold=0.7)
+ )
+ boxes = torch.tensor([
+ [0.0, 0.0, 10.0, 10.0],
+ [2.0, 2.0, 6.0, 6.0],
+ ])
+ scores = torch.tensor([0.4, 0.5])
+ state = self._make_state(boxes=boxes, scores=scores)
+
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["boxes"].shape[0], 1)
+ numpy.testing.assert_array_equal(
+ result["boxes"].numpy(), numpy.array([[0.0, 0.0, 10.0, 10.0]])
+ )
+ # Summed score should be 0.9, well below the 1.0 clamp.
+ self.assertAlmostEqual(float(result["scores"][0]), 0.9, places=5)
+
+ def test_keeps_separate_boxes_when_no_containment(self):
+ """Verifies non-overlapping boxes are not merged."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(merge_overlap_threshold=0.7)
+ )
+ boxes = torch.tensor([
+ [0.0, 0.0, 5.0, 5.0],
+ [20.0, 20.0, 25.0, 25.0],
+ ])
+ scores = torch.tensor([0.6, 0.7])
+ state = self._make_state(boxes=boxes, scores=scores)
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["boxes"].shape[0], 2)
+
+ def test_clamps_summed_score_to_one(self):
+ """Verifies the merged score is clamped to at most 1.0."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(merge_overlap_threshold=0.7)
+ )
+ boxes = torch.tensor([
+ [0.0, 0.0, 10.0, 10.0],
+ [2.0, 2.0, 6.0, 6.0],
+ ])
+ scores = torch.tensor([0.8, 0.7])
+ state = self._make_state(boxes=boxes, scores=scores)
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["boxes"].shape[0], 1)
+ self.assertAlmostEqual(float(result["scores"][0]), 1.0, places=5)
+
+ def test_handles_zero_area_box(self):
+ """Verifies that a zero-area box is marked as absorbed during merge."""
+ detector = _make_detector_without_construction(
+ prompt_config=_make_prompt_config(merge_overlap_threshold=0.7)
+ )
+ boxes = torch.tensor([
+ [0.0, 0.0, 10.0, 10.0],
+ [5.0, 5.0, 5.0, 5.0],
+ ])
+ scores = torch.tensor([0.8, 0.2])
+ state = self._make_state(boxes=boxes, scores=scores)
+
+ result = detector._merge_contained_boxes(state)
+ self.assertEqual(result["boxes"].shape[0], 1)
+
+
+class DetectTest(absltest.TestCase):
+ """Tests for SAM3Detector.detect orchestration."""
+
+ def _make_detector(self, prompt_config=None):
+ """Returns a detector with mocked processor + short-circuited helpers."""
+ detector = _make_detector_without_construction(prompt_config=prompt_config)
+ # set_image / set_text_prompt return the same synthetic state dict.
+ self.state = {
+ "masks": torch.zeros((1, 1, 4, 4), dtype=torch.bool),
+ "masks_logits": torch.zeros((1, 1, 4, 4), dtype=torch.float32),
+ "boxes": torch.tensor([[0.0, 0.0, 10.0, 10.0]]),
+ "scores": torch.tensor([0.6]),
+ "backbone_out": torch.zeros(1),
+ "geometric_prompt": torch.zeros(1),
+ "image_embeddings": torch.zeros(1),
+ }
+ detector._processor.set_image.return_value = self.state
+ detector._processor.set_text_prompt.return_value = self.state
+ return detector
+
+ def test_drops_intermediate_keys_and_returns_cpu_state(self):
+ """Verifies inference-only keys are removed and result is on CPU."""
+ detector = self._make_detector()
+ # Bypass mask filtering by replacing it with a pass-through.
+ with mock.patch.object(
+ detector,
+ "_filter_contained_sub_masks",
+ autospec=True,
+ side_effect=lambda state: state,
+ ):
+ result = detector.detect(
+ image=Image.new("RGB", (16, 16), color="red"),
+ prompt="bottle",
+ )
+ for dropped_key in sam3_detector._INFERENCE_KEYS_TO_DROP:
+ self.assertNotIn(dropped_key, result)
+ self.assertEqual(result["boxes"].device.type, "cpu")
+
+ def test_merges_only_for_the_merge_prompt(self):
+ """Verifies merge runs only when the prompt matches _MERGE_PROMPT."""
+ detector = self._make_detector()
+ with mock.patch.object(
+ detector,
+ "_filter_contained_sub_masks",
+ autospec=True,
+ side_effect=lambda state: state,
+ ), mock.patch.object(
+ detector,
+ "_merge_contained_boxes",
+ autospec=True,
+ side_effect=lambda state: state,
+ ) as mock_merge:
+ detector.detect(
+ image=Image.new("RGB", (16, 16), color="red"),
+ prompt="bottle",
+ )
+ mock_merge.assert_not_called()
+
+ detector.detect(
+ image=Image.new("RGB", (16, 16), color="red"),
+ prompt=sam3_detector._MERGE_PROMPT,
+ )
+ mock_merge.assert_called_once()
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/visualization_utils.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/visualization_utils.py
new file mode 100644
index 00000000000..6d34f8c326f
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/visualization_utils.py
@@ -0,0 +1,594 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Rendering, plotting, video writing, and summary utilities."""
+
+import collections
+from collections.abc import Callable, Mapping, Sequence
+import logging
+import math
+import pathlib
+from typing import TypedDict
+
+import cv2
+import numpy
+from PIL import Image
+import tqdm
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import config_loader
+
+try:
+ import supervision # pylint: disable=g-import-not-at-top
+except ModuleNotFoundError:
+
+ class _SupervisionFallback:
+ """Fallback classes when supervision is not installed."""
+
+ class BoxAnnotator:
+
+ def annotate(self, scene, detections):
+ del detections
+ return scene
+
+ class LabelAnnotator:
+
+ def annotate(self, scene, detections, labels):
+ del detections, labels
+ return scene
+
+ class Detections:
+
+ def __len__(self) -> int:
+ return 0
+
+ supervision = _SupervisionFallback()
+
+
+_LOGGER = logging.getLogger(__name__)
+# Rendering constants for track-grid PNG output.
+# Image occupies this fraction of a tile; the remainder is used for text.
+_TILE_IMAGE_RATIO = 0.75
+_GRID_HEADER_HEIGHT_PIXELS = 80
+_IMAGE_TOP_PADDING_PIXELS = 10
+_TILE_TEXT_TOP_OFFSET_PIXELS = 40
+_TILE_TEXT_LINE_SPACING_PIXELS = 25
+_TILE_TEXT_LARGE_LINE_SPACING_PIXELS = 50
+_TILE_TEXT_LEFT_PADDING_PIXELS = 15
+_HEADER_TEXT_LEFT_PADDING_PIXELS = 20
+_HEADER_TEXT_BASELINE_PIXELS = 50
+_HEADER_FONT_SCALE = 1.0
+_HEADER_FONT_THICKNESS = 2
+_TILE_FONT_SCALE_STANDARD = 0.5
+_TILE_FONT_SCALE_SMALL = 0.45
+_TILE_FONT_THICKNESS = 1
+_TEXT_COLOR_BGR = (0, 0, 0)
+_CANVAS_COLOR_BGR = (255, 255, 255)
+_VIDEO_FOURCC = "mp4v"
+
+
+class PerCropPrediction(TypedDict):
+ """Per-crop prediction record consumed by the track-grid renderer.
+
+ Attributes:
+ crop: The PIL RGB crop that was classified.
+ frame_name: Name of the frame the crop came from.
+ predicted_class: Top-1 predicted class name.
+ predicted_probability_percent: Top-1 probability as a percentage in [0.0,
+ 100.0].
+ """
+
+ crop: Image.Image
+ frame_name: str
+ predicted_class: str
+ predicted_probability_percent: float
+
+
+class TrackSummary(TypedDict):
+ """Resolved summary of a single tracker's predictions.
+
+ Attributes:
+ final_class: The class chosen after aggregating per-crop predictions.
+ category: The collapsed category for `final_class`, or None when the
+ collapsed-category feature is disabled.
+ vote_count: Number of per-crop predictions that voted for `final_class`.
+ output_path: Filesystem path where the track-grid PNG was saved, or a
+ placeholder string when grid saving is disabled.
+ """
+
+ final_class: str
+ category: str | None
+ vote_count: int
+ output_path: str
+
+
+# Signature of the callable that resolves per-crop predictions into a single
+# (final_class, vote_count) result for a tracker.
+LabelResolver = Callable[[Sequence[PerCropPrediction]], tuple[str, int]]
+
+
+class PipelineVisualizer:
+ """Manages frame annotation, grid generation, video writing, and summaries."""
+
+ def __init__(
+ self,
+ config: config_loader.VisualizationConfig,
+ collapsed_categories: config_loader.CollapsedCategoriesConfig,
+ output_video_path: pathlib.Path,
+ ):
+ """Initializes the visualizer.
+
+ Args:
+ config: Visualization-specific configuration.
+ collapsed_categories: Optional grouping of fine-grained classes into
+ broader categories. When disabled, per-category reporting and folder
+ nesting are skipped.
+ output_video_path: Filesystem path where the annotated MP4 will be written
+ when `save_video` is true.
+ """
+ self._config = config
+ self._collapsed_categories = collapsed_categories
+ self._box_annotator = supervision.BoxAnnotator()
+ self._label_annotator = supervision.LabelAnnotator()
+ self._output_video_path = output_video_path
+ self._video_writer: cv2.VideoWriter | None = None
+
+ def annotate_and_write_frame(
+ self,
+ image: Image.Image,
+ detections: supervision.Detections,
+ frame_path: pathlib.Path,
+ ) -> None:
+ """Annotates a frame and writes it to disk and/or video per the config.
+
+ Args:
+ image: Source frame as a PIL RGB image.
+ detections: Detections to overlay on the frame.
+ frame_path: Filesystem path where the annotated frame PNG will be saved
+ when `save_frames` is true.
+ """
+ if not self._config.save_frames and not self._config.save_video:
+ return
+
+ annotated_frame = self._annotate_frame(image=image, detections=detections)
+
+ if self._config.save_frames:
+ cv2.imwrite(str(frame_path), annotated_frame)
+
+ if self._config.save_video:
+ self._write_video_frame(annotated_frame)
+
+ def close_video(self) -> None:
+ """Releases the video writer buffer if one was opened."""
+ if self._video_writer is not None:
+ self._video_writer.release()
+ self._video_writer = None
+
+ def save_track_grids(
+ self,
+ track_predictions: Mapping[int, Sequence[PerCropPrediction]],
+ resolve_labels: LabelResolver,
+ output_directory: pathlib.Path,
+ ) -> dict[int, TrackSummary]:
+ """Resolves labels and renders per-track prediction grids.
+
+ Output folder structure depends on whether collapsed categories are
+ enabled:
+ - Enabled: ///
+ track_NNNN_.png
+ - Disabled: //
+ track_NNNN_.png
+
+ Args:
+ track_predictions: Mapping from tracker_id to its per-crop prediction
+ list.
+ resolve_labels: Callable returning (final_class, vote_count) for a list of
+ per-crop predictions.
+ output_directory: Root directory under which track-grid PNGs are saved.
+
+ Returns:
+ Mapping from tracker_id to a TrackSummary. The `category` field is
+ None when the collapsed-category feature is disabled.
+ """
+ summary: dict[int, TrackSummary] = {}
+ sorted_tracker_ids = sorted(track_predictions.keys())
+
+ progress_description = (
+ "Saving track grids"
+ if self._config.save_track_grids
+ else "Resolving track labels"
+ )
+ progress_bar = tqdm.tqdm(
+ sorted_tracker_ids, desc=progress_description, unit="track"
+ )
+
+ for tracker_id in progress_bar:
+ per_crop_predictions = track_predictions[tracker_id]
+ if not per_crop_predictions:
+ continue
+
+ final_class, vote_count = resolve_labels(per_crop_predictions)
+ progress_bar.set_postfix_str(f"track {tracker_id:04d} -> {final_class}")
+
+ category = self._collapsed_categories.get_category_for_class(final_class)
+ output_path_string = "N/A (grids disabled)"
+
+ if self._config.save_track_grids:
+ output_path = self._render_and_save_track_grid(
+ tracker_id=tracker_id,
+ final_class=final_class,
+ category=category,
+ per_crop_predictions=per_crop_predictions,
+ output_directory=output_directory,
+ )
+ output_path_string = str(output_path)
+
+ summary[tracker_id] = TrackSummary(
+ final_class=final_class,
+ category=category,
+ vote_count=vote_count,
+ output_path=output_path_string,
+ )
+
+ progress_bar.close()
+ return summary
+
+ def print_summary(
+ self,
+ track_summary: Mapping[int, TrackSummary],
+ input_directory: pathlib.Path,
+ class_names: Sequence[str],
+ ) -> None:
+ """Logs per-class object counts and, if enabled, per-category counts.
+
+ Also logs class accuracy and, when collapsed categories are enabled,
+ category accuracy against the ground truth inferred from the input
+ subfolder name.
+
+ Args:
+ track_summary: Mapping from tracker_id to its resolved summary.
+ input_directory: Path to the subfolder that produced this summary.
+ class_names: Full list of class labels from the config.
+ """
+ class_counts = collections.Counter(
+ entry["final_class"] for entry in track_summary.values()
+ )
+ total_objects = sum(class_counts.values())
+
+ _LOGGER.info("Folder name: %s", input_directory.name)
+ _LOGGER.info("Total tracked objects: %d", total_objects)
+
+ _LOGGER.info("By class:")
+ for class_name, count in class_counts.most_common():
+ _LOGGER.info(" %s: %d", class_name, count)
+
+ self._log_class_accuracy(
+ class_counts=class_counts,
+ total_objects=total_objects,
+ input_directory=input_directory,
+ class_names=class_names,
+ )
+
+ if self._collapsed_categories.enable:
+ self._log_collapsed_category_section(
+ track_summary=track_summary,
+ total_objects=total_objects,
+ input_directory=input_directory,
+ )
+
+ def _annotate_frame(
+ self, image: Image.Image, detections: supervision.Detections
+ ) -> numpy.ndarray:
+ """Returns a BGR annotated frame with boxes and labels drawn on it."""
+ frame_bgr = cv2.cvtColor(numpy.array(image), cv2.COLOR_RGB2BGR)
+ annotated = self._box_annotator.annotate(
+ scene=frame_bgr.copy(), detections=detections
+ )
+ labels = self._build_labels(detections)
+ return self._label_annotator.annotate(
+ scene=annotated, detections=detections, labels=labels
+ )
+
+ def _write_video_frame(self, annotated_frame: numpy.ndarray) -> None:
+ """Writes a BGR frame to the video, opening the writer if not yet open."""
+ if self._video_writer is None:
+ height, width = annotated_frame.shape[:2]
+ fourcc = cv2.VideoWriter_fourcc(*_VIDEO_FOURCC)
+ self._video_writer = cv2.VideoWriter(
+ str(self._output_video_path),
+ fourcc,
+ self._config.output_video_fps,
+ (width, height),
+ )
+ self._video_writer.write(annotated_frame)
+
+ def _render_and_save_track_grid(
+ self,
+ tracker_id: int,
+ final_class: str,
+ category: str | None,
+ per_crop_predictions: Sequence[PerCropPrediction],
+ output_directory: pathlib.Path,
+ ) -> pathlib.Path:
+ """Renders one track grid PNG to disk and returns the output path."""
+ tile_size_pixels = int(
+ self._config.track_grid_thumbnail_size_inches
+ * self._config.track_grid_dpi
+ )
+ image_size_pixels = int(tile_size_pixels * _TILE_IMAGE_RATIO)
+
+ grid_canvas = self._build_grid_canvas(
+ tracker_id=tracker_id,
+ final_class=final_class,
+ per_crop_predictions=per_crop_predictions,
+ tile_size_pixels=tile_size_pixels,
+ image_size_pixels=image_size_pixels,
+ )
+
+ class_directory = _build_class_output_directory(
+ output_directory=output_directory,
+ category=category,
+ final_class=final_class,
+ )
+ class_directory.mkdir(parents=True, exist_ok=True)
+ output_path = class_directory / f"track_{tracker_id:04d}_{final_class}.png"
+ cv2.imwrite(str(output_path), grid_canvas)
+ return output_path
+
+ def _build_grid_canvas(
+ self,
+ tracker_id: int,
+ final_class: str,
+ per_crop_predictions: Sequence[PerCropPrediction],
+ tile_size_pixels: int,
+ image_size_pixels: int,
+ ) -> numpy.ndarray:
+ """Builds the grid canvas with a header and per-crop tiles drawn on it."""
+ crop_count = len(per_crop_predictions)
+ columns = min(crop_count, self._config.track_grid_columns_per_row)
+ rows = math.ceil(crop_count / columns)
+
+ grid_width_pixels = columns * tile_size_pixels
+ grid_height_pixels = rows * tile_size_pixels + _GRID_HEADER_HEIGHT_PIXELS
+ grid_canvas = numpy.full(
+ (grid_height_pixels, grid_width_pixels, 3),
+ _CANVAS_COLOR_BGR[0],
+ dtype=numpy.uint8,
+ )
+
+ header_text = (
+ f"Track {tracker_id} | final: {final_class} ({crop_count} crops)"
+ )
+ cv2.putText(
+ grid_canvas,
+ header_text,
+ (_HEADER_TEXT_LEFT_PADDING_PIXELS, _HEADER_TEXT_BASELINE_PIXELS),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ _HEADER_FONT_SCALE,
+ _TEXT_COLOR_BGR,
+ _HEADER_FONT_THICKNESS,
+ cv2.LINE_AA,
+ )
+
+ for tile_index, prediction in enumerate(per_crop_predictions):
+ tile = self._build_tile(
+ prediction=prediction,
+ tile_size_pixels=tile_size_pixels,
+ image_size_pixels=image_size_pixels,
+ )
+ row_index = tile_index // columns
+ column_index = tile_index % columns
+ x_offset = column_index * tile_size_pixels
+ y_offset = _GRID_HEADER_HEIGHT_PIXELS + row_index * tile_size_pixels
+ grid_canvas[
+ y_offset : y_offset + tile_size_pixels,
+ x_offset : x_offset + tile_size_pixels,
+ ] = tile
+
+ return grid_canvas
+
+ def _build_tile(
+ self,
+ prediction: PerCropPrediction,
+ tile_size_pixels: int,
+ image_size_pixels: int,
+ ) -> numpy.ndarray:
+ """Builds a single tile containing a resized crop and text annotations."""
+ tile = numpy.full(
+ (tile_size_pixels, tile_size_pixels, 3),
+ _CANVAS_COLOR_BGR[0],
+ dtype=numpy.uint8,
+ )
+
+ crop_bgr = cv2.cvtColor(numpy.array(prediction["crop"]), cv2.COLOR_RGB2BGR)
+ crop_resized = cv2.resize(
+ crop_bgr,
+ (image_size_pixels, image_size_pixels),
+ interpolation=cv2.INTER_LINEAR,
+ )
+
+ image_x_offset = (tile_size_pixels - image_size_pixels) // 2
+ image_bottom = _IMAGE_TOP_PADDING_PIXELS + image_size_pixels
+ tile[
+ _IMAGE_TOP_PADDING_PIXELS:image_bottom,
+ image_x_offset : image_x_offset + image_size_pixels,
+ ] = crop_resized
+
+ text_baseline_y = image_size_pixels + _TILE_TEXT_TOP_OFFSET_PIXELS
+ cv2.putText(
+ tile,
+ prediction["frame_name"],
+ (_TILE_TEXT_LEFT_PADDING_PIXELS, text_baseline_y),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ _TILE_FONT_SCALE_STANDARD,
+ _TEXT_COLOR_BGR,
+ _TILE_FONT_THICKNESS,
+ cv2.LINE_AA,
+ )
+ cv2.putText(
+ tile,
+ prediction["predicted_class"],
+ (
+ _TILE_TEXT_LEFT_PADDING_PIXELS,
+ text_baseline_y + _TILE_TEXT_LINE_SPACING_PIXELS,
+ ),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ _TILE_FONT_SCALE_SMALL,
+ _TEXT_COLOR_BGR,
+ _TILE_FONT_THICKNESS,
+ cv2.LINE_AA,
+ )
+ cv2.putText(
+ tile,
+ f"{prediction['predicted_probability_percent']:.2f}%",
+ (
+ _TILE_TEXT_LEFT_PADDING_PIXELS,
+ text_baseline_y + _TILE_TEXT_LARGE_LINE_SPACING_PIXELS,
+ ),
+ cv2.FONT_HERSHEY_SIMPLEX,
+ _TILE_FONT_SCALE_STANDARD,
+ _TEXT_COLOR_BGR,
+ _TILE_FONT_THICKNESS,
+ cv2.LINE_AA,
+ )
+ return tile
+
+ def _log_class_accuracy(
+ self,
+ class_counts: collections.Counter[str],
+ total_objects: int,
+ input_directory: pathlib.Path,
+ class_names: Sequence[str],
+ ) -> None:
+ """Logs the ground-truth class accuracy line."""
+ ground_truth_class = _infer_ground_truth_from_names(
+ input_directory=input_directory, candidate_names=class_names
+ )
+ if ground_truth_class is None:
+ _LOGGER.info(
+ " Class accuracy: N/A (could not infer ground-truth class from "
+ "subfolder name)"
+ )
+ return
+ if total_objects == 0:
+ _LOGGER.info(" Class accuracy: N/A (no tracked objects)")
+ return
+ class_accuracy = class_counts.get(ground_truth_class, 0) / total_objects
+ _LOGGER.info(
+ " Class accuracy (vs %s): %.2f%%",
+ ground_truth_class,
+ class_accuracy * 100,
+ )
+
+ def _log_collapsed_category_section(
+ self,
+ track_summary: Mapping[int, TrackSummary],
+ total_objects: int,
+ input_directory: pathlib.Path,
+ ) -> None:
+ """Logs the collapsed-category counts and accuracy line."""
+ category_counts = collections.Counter(
+ entry["category"] for entry in track_summary.values()
+ )
+
+ _LOGGER.info("By collapsed categories:")
+ for category_name, count in category_counts.most_common():
+ _LOGGER.info(" %s: %d", category_name, count)
+
+ ground_truth_category = _infer_ground_truth_from_names(
+ input_directory=input_directory,
+ candidate_names=self._collapsed_categories.category_names,
+ )
+ if ground_truth_category is None:
+ _LOGGER.info(
+ " Category accuracy: N/A (could not infer ground-truth category "
+ "from subfolder name)"
+ )
+ return
+ if total_objects == 0:
+ _LOGGER.info(" Category accuracy: N/A (no tracked objects)")
+ return
+ category_accuracy = (
+ category_counts.get(ground_truth_category, 0) / total_objects
+ )
+ _LOGGER.info(
+ " Category accuracy (vs %s): %.2f%%",
+ ground_truth_category,
+ category_accuracy * 100,
+ )
+
+ def _build_labels(self, detections: supervision.Detections) -> list[str]:
+ """Formats the label strings shown over each detected object."""
+ if not detections:
+ return []
+
+ if detections.tracker_id is not None:
+ tracker_ids = list(detections.tracker_id)
+ else:
+ tracker_ids = [None] * len(detections)
+
+ if detections.confidence is not None:
+ confidences = list(detections.confidence)
+ else:
+ confidences = [None] * len(detections)
+
+ labels: list[str] = []
+ for tracker_id, confidence in zip(tracker_ids, confidences):
+ if tracker_id is None or int(tracker_id) == -1:
+ identifier_string = "?"
+ else:
+ identifier_string = f"ID {int(tracker_id)}"
+ if self._config.show_confidence_in_labels and confidence is not None:
+ labels.append(f"{identifier_string} {float(confidence):.2f}")
+ else:
+ labels.append(identifier_string)
+ return labels
+
+
+def _infer_ground_truth_from_names(
+ input_directory: pathlib.Path, candidate_names: Sequence[str]
+) -> str | None:
+ """Returns the candidate whose name appears in the subfolder basename.
+
+ Matching is a case-insensitive substring match against the basename of
+ the input directory. If zero or multiple candidates match, returns None
+ to signal an ambiguous result.
+
+ Args:
+ input_directory: Path to the per-subfolder input directory.
+ candidate_names: Names to match against.
+
+ Returns:
+ The single matched name, or None if zero or multiple names match.
+ """
+ if not candidate_names:
+ return None
+ subfolder_name = input_directory.name.lower()
+ matched_names = [
+ name for name in candidate_names if name.lower() in subfolder_name
+ ]
+ if len(matched_names) == 1:
+ return matched_names[0]
+ return None
+
+
+def _build_class_output_directory(
+ output_directory: pathlib.Path,
+ category: str | None,
+ final_class: str,
+) -> pathlib.Path:
+ """Returns the directory in which a track-grid PNG should be saved."""
+ if category is None:
+ return output_directory / final_class
+ return output_directory / category / final_class
+
diff --git a/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/visualization_utils_test.py b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/visualization_utils_test.py
new file mode 100644
index 00000000000..3a2caf92d62
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_inference_with_tracking/sam3_dinov3_tracking_pipeline/visualization_utils_test.py
@@ -0,0 +1,553 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for the PipelineVisualizer public API."""
+
+import logging
+import pathlib
+from unittest import mock
+
+from absl.testing import absltest
+import numpy
+from PIL import Image
+
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import config_loader
+from official.projects.waste_identification_ml.model_inference_with_tracking.sam3_dinov3_tracking_pipeline import visualization_utils
+
+
+def _make_visualization_config(
+ save_frames: bool = True,
+ save_video: bool = True,
+ save_track_grids: bool = True,
+) -> config_loader.VisualizationConfig:
+ """Returns a VisualizationConfig with values suitable for tests."""
+ return config_loader.VisualizationConfig(
+ save_frames=save_frames,
+ save_video=save_video,
+ save_track_grids=save_track_grids,
+ output_video_fps=30,
+ show_confidence_in_labels=True,
+ background_blend_color_rgb=(0, 0, 0),
+ track_grid_columns_per_row=2,
+ track_grid_thumbnail_size_inches=2,
+ track_grid_dpi=100,
+ )
+
+
+def _make_collapsed_categories_disabled() -> (
+ config_loader.CollapsedCategoriesConfig
+):
+ """Returns a disabled CollapsedCategoriesConfig."""
+ return config_loader.CollapsedCategoriesConfig(enable=False, mapping={})
+
+
+def _make_collapsed_categories_enabled() -> (
+ config_loader.CollapsedCategoriesConfig
+):
+ """Returns an enabled CollapsedCategoriesConfig with two categories."""
+ return config_loader.CollapsedCategoriesConfig(
+ enable=True,
+ mapping={
+ "recyclable": ["bottle", "can"],
+ "trash": ["wrapper"],
+ },
+ )
+
+
+def _make_per_crop_prediction(
+ predicted_class: str = "bottle",
+ predicted_probability_percent: float = 85.5,
+ frame_name: str = "frame_0001.png",
+) -> visualization_utils.PerCropPrediction:
+ """Returns a PerCropPrediction backed by a small solid-color PIL image."""
+ return visualization_utils.PerCropPrediction(
+ crop=Image.new("RGB", (16, 16), color="red"),
+ frame_name=frame_name,
+ predicted_class=predicted_class,
+ predicted_probability_percent=predicted_probability_percent,
+ )
+
+
+class AnnotateAndWriteFrameTest(absltest.TestCase):
+ """Tests for PipelineVisualizer.annotate_and_write_frame."""
+
+ def setUp(self):
+ """Patches supervision annotators and cv2 calls used during annotation."""
+ super().setUp()
+ self.enter_context(
+ mock.patch.object(
+ visualization_utils.supervision,
+ "BoxAnnotator",
+ autospec=True,
+ create=True,
+ )
+ )
+ self.enter_context(
+ mock.patch.object(
+ visualization_utils.supervision,
+ "LabelAnnotator",
+ autospec=True,
+ create=True,
+ )
+ )
+ self.mock_cvt_color = self.enter_context(
+ mock.patch.object(
+ visualization_utils.cv2, "cvtColor", autospec=False, create=True
+ )
+ )
+ self.mock_imwrite = self.enter_context(
+ mock.patch.object(
+ visualization_utils.cv2, "imwrite", autospec=False, create=True
+ )
+ )
+ self.mock_video_writer_class = self.enter_context(
+ mock.patch.object(
+ visualization_utils.cv2,
+ "VideoWriter",
+ autospec=False,
+ create=True,
+ )
+ )
+ self.enter_context(
+ mock.patch.object(
+ visualization_utils.cv2,
+ "VideoWriter_fourcc",
+ autospec=False,
+ create=True,
+ return_value=0,
+ )
+ )
+
+ # cvtColor returns a synthetic BGR frame with known shape.
+
+ self.mock_cvt_color.return_value = numpy.zeros(
+ (10, 20, 3), dtype=numpy.uint8
+ )
+ # Both annotators return their scene unchanged.
+ box_annotator_instance = (
+ visualization_utils.supervision.BoxAnnotator.return_value
+ )
+ box_annotator_instance.annotate.side_effect = lambda scene, detections: (
+ scene
+ )
+ label_annotator_instance = (
+ visualization_utils.supervision.LabelAnnotator.return_value
+ )
+ label_annotator_instance.annotate.side_effect = (
+ lambda scene, detections, labels: scene
+ )
+
+ def _make_visualizer(
+ self, save_frames: bool = True, save_video: bool = True
+ ) -> visualization_utils.PipelineVisualizer:
+ """Returns a PipelineVisualizer with categories disabled."""
+ return visualization_utils.PipelineVisualizer(
+ config=_make_visualization_config(
+ save_frames=save_frames, save_video=save_video
+ ),
+ collapsed_categories=_make_collapsed_categories_disabled(),
+ output_video_path=pathlib.Path("/tmp/video.mp4"),
+ )
+
+ def _make_detections(self) -> mock.MagicMock:
+ """Returns a fake detections object satisfying visualizer expectations."""
+ detections = mock.MagicMock()
+ detections.__len__.return_value = 1
+ detections.tracker_id = numpy.array([1])
+ detections.confidence = numpy.array([0.9])
+ return detections
+
+ def test_skips_all_work_when_both_toggles_off(self):
+ """Verifies no annotation or I/O occurs when both save toggles are false."""
+ visualizer = self._make_visualizer(save_frames=False, save_video=False)
+ visualizer.annotate_and_write_frame(
+ image=Image.new("RGB", (20, 10), color="red"),
+ detections=self._make_detections(),
+ frame_path=pathlib.Path("/tmp/frame.png"),
+ )
+ self.mock_cvt_color.assert_not_called()
+ self.mock_imwrite.assert_not_called()
+ self.mock_video_writer_class.assert_not_called()
+
+ def test_writes_frame_when_save_frames_enabled(self):
+ """Verifies imwrite is called with the string form of the frame path."""
+ visualizer = self._make_visualizer(save_frames=True, save_video=False)
+ visualizer.annotate_and_write_frame(
+ image=Image.new("RGB", (20, 10), color="red"),
+ detections=self._make_detections(),
+ frame_path=pathlib.Path("/tmp/frame.png"),
+ )
+ self.mock_imwrite.assert_called_once()
+ called_path = self.mock_imwrite.call_args.args[0]
+ self.assertEqual(called_path, "/tmp/frame.png")
+ self.mock_video_writer_class.assert_not_called()
+
+ def test_opens_video_writer_lazily_on_first_frame(self):
+ """Verifies the video writer is constructed lazily on the first frame."""
+ visualizer = self._make_visualizer(save_frames=False, save_video=True)
+ video_writer_instance = self.mock_video_writer_class.return_value
+
+ visualizer.annotate_and_write_frame(
+ image=Image.new("RGB", (20, 10), color="red"),
+ detections=self._make_detections(),
+ frame_path=pathlib.Path("/tmp/frame.png"),
+ )
+ self.mock_video_writer_class.assert_called_once()
+ # Width and height come from the annotated BGR frame shape (10, 20, 3).
+ _, _, _, size = self.mock_video_writer_class.call_args.args
+ self.assertEqual(size, (20, 10))
+ video_writer_instance.write.assert_called_once()
+
+ # A second frame should reuse the same writer instance.
+ visualizer.annotate_and_write_frame(
+ image=Image.new("RGB", (20, 10), color="red"),
+ detections=self._make_detections(),
+ frame_path=pathlib.Path("/tmp/frame_2.png"),
+ )
+ self.mock_video_writer_class.assert_called_once()
+ self.assertEqual(video_writer_instance.write.call_count, 2)
+
+
+class CloseVideoTest(absltest.TestCase):
+ """Tests for PipelineVisualizer.close_video."""
+
+ def test_release_is_a_no_op_when_no_video_was_written(self):
+ """Verifies close_video does not fail when no writer was opened."""
+ visualizer = visualization_utils.PipelineVisualizer(
+ config=_make_visualization_config(),
+ collapsed_categories=_make_collapsed_categories_disabled(),
+ output_video_path=pathlib.Path("/tmp/video.mp4"),
+ )
+ # Should not raise.
+ visualizer.close_video()
+
+ def test_release_is_called_when_writer_exists(self):
+ """Verifies close_video releases the underlying writer."""
+ visualizer = visualization_utils.PipelineVisualizer(
+ config=_make_visualization_config(),
+ collapsed_categories=_make_collapsed_categories_disabled(),
+ output_video_path=pathlib.Path("/tmp/video.mp4"),
+ )
+ mock_writer = mock.MagicMock()
+ visualizer._video_writer = mock_writer
+ visualizer.close_video()
+ mock_writer.release.assert_called_once()
+ self.assertIsNone(visualizer._video_writer)
+
+
+class SaveTrackGridsTest(absltest.TestCase):
+ """Tests for PipelineVisualizer.save_track_grids."""
+
+ def setUp(self):
+ """Patches cv2 write operations and tqdm progress-bar iteration."""
+ super().setUp()
+ self.mock_imwrite = self.enter_context(
+ mock.patch.object(visualization_utils.cv2, "imwrite", autospec=True)
+ )
+ # tqdm wraps the iterable; we make it behave like a plain iterator with
+ # no-op set_postfix_str / close methods.
+ self.enter_context(
+ mock.patch.object(
+ visualization_utils.tqdm,
+ "tqdm",
+ autospec=True,
+ side_effect=lambda iterable, **kwargs: _FakeProgressBar(iterable),
+ )
+ )
+
+ def _make_visualizer(
+
+ self,
+ save_track_grids: bool,
+ collapsed_categories: config_loader.CollapsedCategoriesConfig,
+ ) -> visualization_utils.PipelineVisualizer:
+ """Returns a PipelineVisualizer with the given grid-saving toggle."""
+ return visualization_utils.PipelineVisualizer(
+ config=_make_visualization_config(save_track_grids=save_track_grids),
+ collapsed_categories=collapsed_categories,
+ output_video_path=pathlib.Path("/tmp/video.mp4"),
+ )
+
+ def test_returns_summary_and_skips_disk_when_grids_disabled(self):
+ """Verifies labels are resolved but no PNGs are written when disabled."""
+ visualizer = self._make_visualizer(
+ save_track_grids=False,
+ collapsed_categories=_make_collapsed_categories_disabled(),
+ )
+ track_predictions = {
+ 1: [_make_per_crop_prediction(predicted_class="bottle")],
+ }
+ resolve_labels = mock.MagicMock(return_value=("bottle", 1))
+
+ summary = visualizer.save_track_grids(
+ track_predictions=track_predictions,
+ resolve_labels=resolve_labels,
+ output_directory=pathlib.Path("/tmp/grids"),
+ )
+
+ self.assertEqual(summary[1]["final_class"], "bottle")
+ self.assertIsNone(summary[1]["category"])
+ self.assertEqual(summary[1]["vote_count"], 1)
+ self.assertEqual(summary[1]["output_path"], "N/A (grids disabled)")
+ self.mock_imwrite.assert_not_called()
+
+ def test_skips_tracks_with_no_predictions(self):
+ """Verifies tracker_ids with empty prediction lists are not included."""
+ visualizer = self._make_visualizer(
+ save_track_grids=False,
+ collapsed_categories=_make_collapsed_categories_disabled(),
+ )
+ track_predictions = {
+ 1: [],
+ 2: [_make_per_crop_prediction()],
+ }
+ resolve_labels = mock.MagicMock(return_value=("bottle", 1))
+
+ summary = visualizer.save_track_grids(
+ track_predictions=track_predictions,
+ resolve_labels=resolve_labels,
+ output_directory=pathlib.Path("/tmp/grids"),
+ )
+ self.assertNotIn(1, summary)
+ self.assertIn(2, summary)
+
+ def test_writes_png_under_flat_directory_when_categories_disabled(self):
+ """Verifies output path is //track_NNNN_.png."""
+ visualizer = self._make_visualizer(
+ save_track_grids=True,
+ collapsed_categories=_make_collapsed_categories_disabled(),
+ )
+ track_predictions = {
+ 7: [_make_per_crop_prediction(predicted_class="bottle")],
+ }
+ resolve_labels = mock.MagicMock(return_value=("bottle", 1))
+
+ with mock.patch.object(
+ visualization_utils.pathlib.Path, "mkdir", autospec=True
+ ):
+ summary = visualizer.save_track_grids(
+ track_predictions=track_predictions,
+ resolve_labels=resolve_labels,
+ output_directory=pathlib.Path("/tmp/grids"),
+ )
+
+ self.assertEqual(
+ summary[7]["output_path"], "/tmp/grids/bottle/track_0007_bottle.png"
+ )
+ self.mock_imwrite.assert_called_once()
+
+ def test_writes_png_under_category_directory_when_categories_enabled(self):
+ """Verifies output path is ///track_...png."""
+ visualizer = self._make_visualizer(
+ save_track_grids=True,
+ collapsed_categories=_make_collapsed_categories_enabled(),
+ )
+ track_predictions = {
+ 3: [_make_per_crop_prediction(predicted_class="bottle")],
+ }
+ resolve_labels = mock.MagicMock(return_value=("bottle", 1))
+
+ with mock.patch.object(
+ visualization_utils.pathlib.Path, "mkdir", autospec=True
+ ):
+ summary = visualizer.save_track_grids(
+ track_predictions=track_predictions,
+ resolve_labels=resolve_labels,
+ output_directory=pathlib.Path("/tmp/grids"),
+ )
+
+ self.assertEqual(summary[3]["category"], "recyclable")
+ self.assertEqual(
+ summary[3]["output_path"],
+ "/tmp/grids/recyclable/bottle/track_0003_bottle.png",
+ )
+
+
+class PrintSummaryTest(absltest.TestCase):
+ """Tests for PipelineVisualizer.print_summary."""
+
+ def _make_visualizer(
+ self,
+ collapsed_categories: config_loader.CollapsedCategoriesConfig,
+ ) -> visualization_utils.PipelineVisualizer:
+ """Returns a PipelineVisualizer with categories as specified."""
+ return visualization_utils.PipelineVisualizer(
+ config=_make_visualization_config(),
+ collapsed_categories=collapsed_categories,
+ output_video_path=pathlib.Path("/tmp/video.mp4"),
+ )
+
+ def _make_summary(
+ self,
+ ) -> dict[int, visualization_utils.TrackSummary]:
+ """Returns a small summary with two bottle tracks and one wrapper track."""
+ return {
+ 1: visualization_utils.TrackSummary(
+ final_class="bottle",
+ category="recyclable",
+ vote_count=5,
+ output_path="/tmp/a.png",
+ ),
+ 2: visualization_utils.TrackSummary(
+ final_class="bottle",
+ category="recyclable",
+ vote_count=4,
+ output_path="/tmp/b.png",
+ ),
+ 3: visualization_utils.TrackSummary(
+ final_class="wrapper",
+ category="trash",
+ vote_count=3,
+ output_path="/tmp/c.png",
+ ),
+ }
+
+ def test_logs_total_and_per_class_counts(self):
+ """Verifies total and per-class lines are logged."""
+ visualizer = self._make_visualizer(
+ collapsed_categories=_make_collapsed_categories_disabled()
+ )
+ with self.assertLogs(
+ visualization_utils._LOGGER.name, level=logging.INFO
+ ) as logs:
+ visualizer.print_summary(
+ track_summary=self._make_summary(),
+ input_directory=pathlib.Path("/data/bottle_session"),
+ class_names=["bottle", "wrapper", "can"],
+ )
+ joined = "\n".join(logs.output)
+ self.assertIn("Total tracked objects: 3", joined)
+ self.assertIn("bottle: 2", joined)
+ self.assertIn("wrapper: 1", joined)
+
+ def test_logs_class_accuracy_when_ground_truth_inferable(self):
+ """Verifies class accuracy is computed when the subfolder name matches."""
+ visualizer = self._make_visualizer(
+ collapsed_categories=_make_collapsed_categories_disabled()
+ )
+ with self.assertLogs(
+ visualization_utils._LOGGER.name, level=logging.INFO
+ ) as logs:
+ visualizer.print_summary(
+ track_summary=self._make_summary(),
+ input_directory=pathlib.Path("/data/bottle_session"),
+ class_names=["bottle", "wrapper", "can"],
+ )
+ joined = "\n".join(logs.output)
+ # 2 of 3 tracks are 'bottle' -> 66.67%.
+ self.assertIn("Class accuracy (vs bottle): 66.67%", joined)
+
+ def test_reports_class_accuracy_na_when_subfolder_unmatched(self):
+ """Verifies a clear message when no class name matches the subfolder."""
+ visualizer = self._make_visualizer(
+ collapsed_categories=_make_collapsed_categories_disabled()
+ )
+ with self.assertLogs(
+ visualization_utils._LOGGER.name, level=logging.INFO
+ ) as logs:
+ visualizer.print_summary(
+ track_summary=self._make_summary(),
+ input_directory=pathlib.Path("/data/unrelated_folder"),
+ class_names=["bottle", "wrapper", "can"],
+ )
+ self.assertIn("Class accuracy: N/A", "\n".join(logs.output))
+
+ def test_logs_category_section_when_categories_enabled(self):
+ """Verifies the category section is present when collapsed cats are on."""
+ visualizer = self._make_visualizer(
+ collapsed_categories=_make_collapsed_categories_enabled()
+ )
+ with self.assertLogs(
+ visualization_utils._LOGGER.name, level=logging.INFO
+ ) as logs:
+ visualizer.print_summary(
+ track_summary=self._make_summary(),
+ input_directory=pathlib.Path("/data/recyclable_session"),
+ class_names=["bottle", "wrapper", "can"],
+ )
+ joined = "\n".join(logs.output)
+ self.assertIn("By collapsed categories:", joined)
+ self.assertIn("recyclable: 2", joined)
+ self.assertIn("trash: 1", joined)
+ # 2 of 3 tracks are 'recyclable' -> 66.67%.
+ self.assertIn("Category accuracy (vs recyclable): 66.67%", joined)
+
+ def test_omits_category_section_when_categories_disabled(self):
+ """Verifies the category section is absent when collapsed cats are off."""
+ visualizer = self._make_visualizer(
+ collapsed_categories=_make_collapsed_categories_disabled()
+ )
+ with self.assertLogs(
+ visualization_utils._LOGGER.name, level=logging.INFO
+ ) as logs:
+ visualizer.print_summary(
+ track_summary=self._make_summary(),
+ input_directory=pathlib.Path("/data/bottle_session"),
+ class_names=["bottle", "wrapper", "can"],
+ )
+ self.assertNotIn("By collapsed categories:", "\n".join(logs.output))
+
+ def test_infer_ground_truth_from_names_returns_none_when_candidates_empty(
+ self,
+ ):
+ """Verifies _infer_ground_truth_from_names returns None if candidates empty."""
+ self.assertIsNone(
+ visualization_utils._infer_ground_truth_from_names(
+ input_directory=pathlib.Path("/data/bottle_session"),
+ candidate_names=[],
+ )
+ )
+
+
+class SupervisionFallbackTest(absltest.TestCase):
+ """Tests for _SupervisionFallback dummy classes."""
+
+ def test_fallback_annotators_return_scene_unchanged(self):
+ """Verifies fallback BoxAnnotator and LabelAnnotator return scene as-is."""
+ if not hasattr(visualization_utils, "_SupervisionFallback"):
+ self.skipTest("supervision package is installed; fallback not active.")
+ fallback = visualization_utils._SupervisionFallback()
+ box_annotator = fallback.BoxAnnotator()
+ label_annotator = fallback.LabelAnnotator()
+ dummy_scene = numpy.zeros((10, 10, 3), dtype=numpy.uint8)
+
+ self.assertIs(
+ box_annotator.annotate(dummy_scene, detections=None),
+ dummy_scene,
+ )
+ self.assertIs(
+ label_annotator.annotate(dummy_scene, detections=None, labels=[]),
+ dummy_scene,
+ )
+
+
+class _FakeProgressBar:
+ """Minimal stand-in for tqdm.tqdm used in save_track_grids tests."""
+
+ def __init__(self, iterable):
+ self._iterable = list(iterable)
+
+ def __iter__(self):
+ return iter(self._iterable)
+
+ def set_postfix_str(self, text: str) -> None:
+ del text
+ return None
+
+ def close(self) -> None:
+ return None
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/official/projects/waste_identification_ml/model_retraining/CircularNET_Vertex_AI_ReTraining_v1.ipynb b/official/projects/waste_identification_ml/model_retraining/CircularNET_Vertex_AI_ReTraining_v1.ipynb
new file mode 100644
index 00000000000..3289af06a73
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_retraining/CircularNET_Vertex_AI_ReTraining_v1.ipynb
@@ -0,0 +1,900 @@
+{
+ "nbformat": 4,
+ "nbformat_minor": 0,
+ "metadata": {
+ "colab": {
+ "provenance": []
+ },
+ "kernelspec": {
+ "name": "python3",
+ "display_name": "Python 3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## CircularNet Vertex AI Retraining Pipeline\n",
+ "\n",
+ "The goal is to train CircularNet model on Vertex AI using published checkpoints and configuration file.\n",
+ "\n",
+ "CircularNet team already open sourced the model on GitHub. But if users want to train or fine tune the model with their own training images, user can ran this notebook to launch a training job on Vertex AI. After training is completed, checkpoints will be output to the GCP storage bucket. User can then export the checkpoints to a saved TF model."
+ ],
+ "metadata": {
+ "id": "EeUlPs7V0ed9"
+ }
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "##Import Libraries and Setup Environment"
+ ],
+ "metadata": {
+ "id": "qrlHFLG6QT0s"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "if \"google.colab\" in str(get_ipython()):\n",
+ " # install google cloud API SDK\n",
+ " ! pip3 install -q --upgrade google-cloud-aiplatform[tensorboard]\n",
+ " # install model-garden official\n",
+ " ! pip3 install -q tf-models-official\n"
+ ],
+ "metadata": {
+ "id": "rHkSX9rpyUl_",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "5976efdb-c5c2-4840-a0e9-2e39fa002c14"
+ },
+ "execution_count": 1,
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stdout",
+ "text": [
+ "\u001b[?25l \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m0.0/5.7 MB\u001b[0m \u001b[31m?\u001b[0m eta \u001b[36m-:--:--\u001b[0m\r\u001b[2K \u001b[91m━━━━━━━━━━━\u001b[0m\u001b[91m╸\u001b[0m\u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.7/5.7 MB\u001b[0m \u001b[31m50.5 MB/s\u001b[0m eta \u001b[36m0:00:01\u001b[0m\r\u001b[2K \u001b[91m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[91m╸\u001b[0m \u001b[32m5.7/5.7 MB\u001b[0m \u001b[31m77.1 MB/s\u001b[0m eta \u001b[36m0:00:01\u001b[0m\r\u001b[2K \u001b[91m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[91m╸\u001b[0m \u001b[32m5.7/5.7 MB\u001b[0m \u001b[31m77.1 MB/s\u001b[0m eta \u001b[36m0:00:01\u001b[0m\r\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m5.7/5.7 MB\u001b[0m \u001b[31m46.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m289.2/289.2 kB\u001b[0m \u001b[31m14.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m6.9/6.9 MB\u001b[0m \u001b[31m51.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h\u001b[31mERROR: pip's dependency resolver does not currently take into account all the packages that are installed. This behaviour is the source of the following dependency conflicts.\n",
+ "flask 3.1.0 requires Werkzeug>=3.1, but you have werkzeug 2.0.3 which is incompatible.\u001b[0m\u001b[31m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m51.8/51.8 kB\u001b[0m \u001b[31m2.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m43.6/43.6 kB\u001b[0m \u001b[31m2.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m2.9/2.9 MB\u001b[0m \u001b[31m44.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m2.2/2.2 MB\u001b[0m \u001b[31m69.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m615.3/615.3 MB\u001b[0m \u001b[31m2.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m242.5/242.5 kB\u001b[0m \u001b[31m13.8 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m5.2/5.2 MB\u001b[0m \u001b[31m86.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.7/1.7 MB\u001b[0m \u001b[31m45.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m104.0/104.0 kB\u001b[0m \u001b[31m7.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m5.5/5.5 MB\u001b[0m \u001b[31m84.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n",
+ "\u001b[?25h Building wheel for seqeval (setup.py) ... \u001b[?25l\u001b[?25hdone\n"
+ ]
+ }
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "import os\n",
+ "from datetime import datetime\n",
+ "from google.cloud import aiplatform\n",
+ "from google.cloud.aiplatform import hyperparameter_tuning as hpt"
+ ],
+ "metadata": {
+ "id": "KXVVioDH0cGQ"
+ },
+ "execution_count": 2,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [],
+ "metadata": {
+ "id": "vVU-a93A1SgL"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# authenticate google colab\n",
+ "if \"google.colab\" in str(get_ipython()):\n",
+ "\n",
+ " from google.colab import auth as google_auth\n",
+ "\n",
+ " google_auth.authenticate_user()"
+ ],
+ "metadata": {
+ "id": "X2srDbpvyq5B"
+ },
+ "execution_count": 3,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Configure and Launch training job on Vertex AI"
+ ],
+ "metadata": {
+ "id": "9rfQE3CJQcDH"
+ }
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 4,
+ "metadata": {
+ "id": "ijP8BxDCyJPG"
+ },
+ "outputs": [],
+ "source": [
+ "# you can set the train job name prefix\n",
+ "TRAINING_JOB_PREFIX = \"training_material_model\" # @param {type:\"string\"}\n",
+ "OBJECTIVE = \"iod\"\n",
+ "\n",
+ "def get_job_name_with_datetime(prefix: str):\n",
+ " \"\"\"create a unique job name\n",
+ " Args:\n",
+ " prefix: prefix string for the training job name\n",
+ " Returns:\n",
+ " a unique training job name by appending a timestamp to prefix\n",
+ " \"\"\"\n",
+ " return prefix + datetime.now().strftime(\"_%Y%m%d_%H%M%S\")\n",
+ "\n",
+ "train_job_name = get_job_name_with_datetime(TRAINING_JOB_PREFIX + \"_\" + OBJECTIVE)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "train_job_name"
+ ],
+ "metadata": {
+ "id": "8v1_9YElzQkX",
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 36
+ },
+ "outputId": "097c737e-c49b-4025-82e3-2e54e5abd95a"
+ },
+ "execution_count": 5,
+ "outputs": [
+ {
+ "output_type": "execute_result",
+ "data": {
+ "text/plain": [
+ "'training_material_model_iod_20241218_190038'"
+ ],
+ "application/vnd.google.colaboratory.intrinsic+json": {
+ "type": "string"
+ }
+ },
+ "metadata": {},
+ "execution_count": 5
+ }
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "PROJECT_ID = \"waste-identification-ml-330916\" # @param {type:\"string\"}\n",
+ "BUCKET_URI = \"gs://circularnet\" # @param {type:\"string\"}\n",
+ "REGION = \"us-central1\" # @param {type:\"string\"}\n",
+ "\n",
+ "model_dir = os.path.join(BUCKET_URI, train_job_name)\n",
+ "\n",
+ "STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
+ "CHECKPOINT_BUCKET = os.path.join(BUCKET_URI, \"ckpt\")\n",
+ "\n",
+ "tensorboard_name = get_job_name_with_datetime(\"tensorboard\")\n",
+ "\n",
+ "tensorboard = aiplatform.Tensorboard.create(\n",
+ " display_name=tensorboard_name,\n",
+ " project=PROJECT_ID,\n",
+ " location=REGION,\n",
+ ")\n",
+ "\n",
+ "aiplatform.init(project=PROJECT_ID,\n",
+ " location=REGION,\n",
+ " staging_bucket=STAGING_BUCKET,\n",
+ " experiment_tensorboard=tensorboard)"
+ ],
+ "metadata": {
+ "id": "ifVgDBhqymtT",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "c57fa5fa-360b-499a-f639-4ec4a1c77eee"
+ },
+ "execution_count": 6,
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "INFO:google.cloud.aiplatform.tensorboard.tensorboard_resource:Creating Tensorboard\n",
+ "INFO:google.cloud.aiplatform.tensorboard.tensorboard_resource:Create Tensorboard backing LRO: projects/372354466754/locations/us-central1/tensorboards/3436752291032465408/operations/7568047832309956608\n",
+ "INFO:google.cloud.aiplatform.tensorboard.tensorboard_resource:Tensorboard created. Resource name: projects/372354466754/locations/us-central1/tensorboards/3436752291032465408\n",
+ "INFO:google.cloud.aiplatform.tensorboard.tensorboard_resource:To use this Tensorboard in another session:\n",
+ "INFO:google.cloud.aiplatform.tensorboard.tensorboard_resource:tb = aiplatform.Tensorboard('projects/372354466754/locations/us-central1/tensorboards/3436752291032465408')\n"
+ ]
+ }
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "model_dir"
+ ],
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 36
+ },
+ "id": "xYh1htheq-4J",
+ "outputId": "9d123e8f-0fdc-49a9-ae89-8530d2ff67c6"
+ },
+ "execution_count": 7,
+ "outputs": [
+ {
+ "output_type": "execute_result",
+ "data": {
+ "text/plain": [
+ "'gs://circularnet/training_material_model_iod_20241218_190038'"
+ ],
+ "application/vnd.google.colaboratory.intrinsic+json": {
+ "type": "string"
+ }
+ },
+ "metadata": {},
+ "execution_count": 7
+ }
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "tensorboard.resource_name"
+ ],
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 36
+ },
+ "id": "9rc3O9qVjb3F",
+ "outputId": "261238c4-e9de-46eb-b44f-204f7eca2cb4"
+ },
+ "execution_count": 8,
+ "outputs": [
+ {
+ "output_type": "execute_result",
+ "data": {
+ "text/plain": [
+ "'projects/372354466754/locations/us-central1/tensorboards/3436752291032465408'"
+ ],
+ "application/vnd.google.colaboratory.intrinsic+json": {
+ "type": "string"
+ }
+ },
+ "metadata": {},
+ "execution_count": 8
+ }
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "OBJECTIVE = 'iod'\n",
+ "REGION_PREFIX = REGION.split(\"-\")[0]\n",
+ "assert REGION_PREFIX in (\n",
+ " \"us\",\n",
+ " \"europe\",\n",
+ " \"asia\",\n",
+ "), f'{REGION} is not supported. It must be prefixed by \"us\", \"asia\", or \"europe\".'\n",
+ "\n",
+ "# set the Training constants.\n",
+ "TRAINING_JOB_PREFIX = \"train\"\n",
+ "TRAIN_CONTAINER_URI = f\"{REGION_PREFIX}-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/tfvision-oss-v2\"\n",
+ "TRAIN_MACHINE_TYPE = \"n1-highmem-16\"\n",
+ "TRAIN_ACCELERATOR_TYPE = \"NVIDIA_TESLA_V100\"\n",
+ "TRAIN_NUM_GPU = 4\n",
+ "\n",
+ "# set the Evaluation constants.\n",
+ "EVALUATION_METRIC = \"mean_iou\"\n"
+ ],
+ "metadata": {
+ "id": "POplG1swyzzJ"
+ },
+ "execution_count": 9,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# set the path to model training TF records, checkpoints and config\n",
+ "input_train_data_path = 'gs://circularnet/vertex_training/test0731/benjamin_30july2024/tfrecords_train/*.tfrecord' # @param {type:\"string\"}\n",
+ "input_validation_data_path = 'gs://circularnet/vertex_training/test0731/benjamin_30july2024/tfrecords_val/*.tfrecord' # @param {type:\"string\"}\n",
+ "init_checkpoint_path = 'gs://circularnet/ckpt/transfer-learning/material_form/ckpt-582000' # @param {type:\"string\"}\n",
+ "config_file_path = 'gs://circularnet/config/config_transferLearning_V5_7_31.yaml' # @param {type:\"string\"}\n",
+ "#total material form model labels\n",
+ "num_classes = 39 # @param {type:\"integer\"}\n",
+ "\n",
+ "# set the path of initial checkpoint and config yaml file\n",
+ "# all the args are configurable based on you specific use case\n",
+ "experiment_container_args_dict = {\n",
+ " # maskrcnn experiment args.\n",
+ " \"maskrcnn_resnetfpn_coco\": {\n",
+ " \"experiment\": \"maskrcnn_resnetfpn_coco\",\n",
+ " \"init_checkpoint\": init_checkpoint_path,\n",
+ " \"config_file\": config_file_path,\n",
+ " \"input_train_data_path\": input_train_data_path,\n",
+ " \"input_validation_data_path\": input_validation_data_path,\n",
+ " \"objective\": OBJECTIVE,\n",
+ " \"model_dir\": f\"{model_dir}/trained_model\",\n",
+ " \"num_classes\": num_classes,\n",
+ " \"global_batch_size\": 4,\n",
+ " \"prefetch_buffer_size\": 12,\n",
+ " \"train_steps\": 100,\n",
+ " }\n",
+ "}\n",
+ "\n",
+ "\n",
+ "experiment = \"maskrcnn_resnetfpn_coco\"\n",
+ "experiment_container_args = experiment_container_args_dict[experiment]\n",
+ "\n",
+ "#configure the training VM, GPU type, # of GPUs\n",
+ "worker_pool_specs = [\n",
+ " {\n",
+ " \"machine_spec\": {\n",
+ " \"machine_type\": TRAIN_MACHINE_TYPE,\n",
+ " \"accelerator_type\": TRAIN_ACCELERATOR_TYPE,\n",
+ " \"accelerator_count\": TRAIN_NUM_GPU,\n",
+ " },\n",
+ " \"replica_count\": 1,\n",
+ " \"container_spec\": {\n",
+ " \"image_uri\": TRAIN_CONTAINER_URI,\n",
+ " \"args\": [\n",
+ " \"--mode=train_and_eval\",\n",
+ " ]\n",
+ " + [\"--{}={}\".format(k, v) for k, v in experiment_container_args.items()],\n",
+ " },\n",
+ " },\n",
+ "]\n",
+ "\n",
+ "metric_spec = {\"model_performance\": \"maximize\"}\n",
+ "\n",
+ "\n",
+ "train_custom_job = aiplatform.CustomJob(\n",
+ " display_name=train_job_name,\n",
+ " project=PROJECT_ID,\n",
+ " worker_pool_specs=worker_pool_specs,\n",
+ " staging_bucket=STAGING_BUCKET,\n",
+ ")\n",
+ "\n",
+ "LEARNING_RATES = [0.001]\n",
+ "\n",
+ "MAX_TRIAL_COUNT = len(LEARNING_RATES)\n",
+ "\n",
+ "parameter_spec = {\n",
+ " \"learning_rate\": hpt.DiscreteParameterSpec(values=LEARNING_RATES, scale=\"linear\"),\n",
+ "}\n",
+ "\n",
+ "# create the Vertex AI hyperparameter training job\n",
+ "train_hpt_job = aiplatform.HyperparameterTuningJob(\n",
+ " display_name=train_job_name,\n",
+ " custom_job=train_custom_job,\n",
+ " metric_spec=metric_spec,\n",
+ " parameter_spec=parameter_spec,\n",
+ " max_trial_count=MAX_TRIAL_COUNT,\n",
+ " parallel_trial_count=1,\n",
+ " project=PROJECT_ID,\n",
+ " search_algorithm=None,\n",
+ ")\n",
+ "\n",
+ "# please change it to your own service_account\n",
+ "train_hpt_job.run()"
+ ],
+ "metadata": {
+ "id": "BheyZp-sy_zc",
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 1000
+ },
+ "outputId": "e82d9054-b059-49e9-959a-4e86c04042c6"
+ },
+ "execution_count": 10,
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "INFO:google.cloud.aiplatform.jobs:Creating HyperparameterTuningJob\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob created. Resource name: projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952\n",
+ "INFO:google.cloud.aiplatform.jobs:To use this HyperparameterTuningJob in another session:\n",
+ "INFO:google.cloud.aiplatform.jobs:hpt_job = aiplatform.HyperparameterTuningJob.get('projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952')\n",
+ "INFO:google.cloud.aiplatform.jobs:View HyperparameterTuningJob:\n",
+ "https://console.cloud.google.com/ai/platform/locations/us-central1/training/5866830459098365952?project=372354466754\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_PENDING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_RUNNING\n",
+ "INFO:google.cloud.aiplatform.jobs:HyperparameterTuningJob projects/372354466754/locations/us-central1/hyperparameterTuningJobs/5866830459098365952 current state:\n",
+ "JobState.JOB_STATE_FAILED\n"
+ ]
+ },
+ {
+ "output_type": "error",
+ "ename": "RuntimeError",
+ "evalue": "Job failed with:\ncode: 3\nmessage: \"Hyperparameter Tuning Trial #1 Failed before any other successful trials were completed. The failed trial had parameters: learning_rate=0.001, . The trial\\'s error message was: The replica workerpool0-0 exited with a non-zero status of 1. Termination reason: Error. To find out more about why your job exited please check the logs: https://console.cloud.google.com/logs/viewer?project=372354466754&resource=ml_job%2Fjob_id%2F5866830459098365952&advancedFilter=resource.type%3D%22ml_job%22%0Aresource.labels.job_id%3D%225866830459098365952%22\"\n",
+ "traceback": [
+ "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
+ "\u001b[0;31mRuntimeError\u001b[0m Traceback (most recent call last)",
+ "\u001b[0;32m\u001b[0m in \u001b[0;36m\u001b[0;34m()\u001b[0m\n\u001b[1;32m 81\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 82\u001b[0m \u001b[0;31m# please change it to your own service_account\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 83\u001b[0;31m \u001b[0mtrain_hpt_job\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mrun\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
+ "\u001b[0;32m/usr/local/lib/python3.10/dist-packages/google/cloud/aiplatform/jobs.py\u001b[0m in \u001b[0;36mrun\u001b[0;34m(self, service_account, network, timeout, restart_job_on_worker_restart, enable_web_access, tensorboard, sync, create_request_timeout, disable_retries, scheduling_strategy, max_wait_duration)\u001b[0m\n\u001b[1;32m 2990\u001b[0m \u001b[0mservice_account\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mservice_account\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0minitializer\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mglobal_config\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mservice_account\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2991\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 2992\u001b[0;31m self._run(\n\u001b[0m\u001b[1;32m 2993\u001b[0m \u001b[0mservice_account\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mservice_account\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 2994\u001b[0m \u001b[0mnetwork\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mnetwork\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
+ "\u001b[0;32m/usr/local/lib/python3.10/dist-packages/google/cloud/aiplatform/base.py\u001b[0m in \u001b[0;36mwrapper\u001b[0;34m(*args, **kwargs)\u001b[0m\n\u001b[1;32m 861\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 862\u001b[0m \u001b[0mVertexAiResourceNounWithFutureManager\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mwait\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 863\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mmethod\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 864\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 865\u001b[0m \u001b[0;31m# callbacks to call within the Future (in same Thread)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
+ "\u001b[0;32m/usr/local/lib/python3.10/dist-packages/google/cloud/aiplatform/jobs.py\u001b[0m in \u001b[0;36m_run\u001b[0;34m(self, service_account, network, timeout, restart_job_on_worker_restart, enable_web_access, tensorboard, sync, create_request_timeout, disable_retries, scheduling_strategy, max_wait_duration)\u001b[0m\n\u001b[1;32m 3128\u001b[0m )\n\u001b[1;32m 3129\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 3130\u001b[0;31m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_block_until_complete\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 3131\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 3132\u001b[0m \u001b[0;34m@\u001b[0m\u001b[0mproperty\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
+ "\u001b[0;32m/usr/local/lib/python3.10/dist-packages/google/cloud/aiplatform/jobs.py\u001b[0m in \u001b[0;36m_block_until_complete\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m 1672\u001b[0m \u001b[0;31m# JOB_STATE_FAILED or JOB_STATE_CANCELLED.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1673\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_gca_resource\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstate\u001b[0m \u001b[0;32min\u001b[0m \u001b[0m_JOB_ERROR_STATES\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1674\u001b[0;31m \u001b[0;32mraise\u001b[0m \u001b[0mRuntimeError\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"Job failed with:\\n%s\"\u001b[0m \u001b[0;34m%\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_gca_resource\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0merror\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 1675\u001b[0m \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 1676\u001b[0m \u001b[0m_LOGGER\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mlog_action_completed_against_resource\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"run\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m\"completed\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
+ "\u001b[0;31mRuntimeError\u001b[0m: Job failed with:\ncode: 3\nmessage: \"Hyperparameter Tuning Trial #1 Failed before any other successful trials were completed. The failed trial had parameters: learning_rate=0.001, . The trial\\'s error message was: The replica workerpool0-0 exited with a non-zero status of 1. Termination reason: Error. To find out more about why your job exited please check the logs: https://console.cloud.google.com/logs/viewer?project=372354466754&resource=ml_job%2Fjob_id%2F5866830459098365952&advancedFilter=resource.type%3D%22ml_job%22%0Aresource.labels.job_id%3D%225866830459098365952%22\"\n"
+ ]
+ }
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# load train logs into tensorboard\n",
+ "\n",
+ "tensorboard_log_dir = f\"{model_dir}/trained_model/trial_1/train\"\n",
+ "\n",
+ "aiplatform.start_upload_tb_log(\n",
+ " tensorboard_id=tensorboard.name,\n",
+ " tensorboard_experiment_name=train_hpt_job.name,\n",
+ " logdir=tensorboard_log_dir,\n",
+ " description=\"train\"\n",
+ " )\n",
+ "aiplatform.end_upload_tb_log()"
+ ],
+ "metadata": {
+ "id": "1_N4FSX3HW_T"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# load validation logs into tensorboard\n",
+ "\n",
+ "tensorboard_log_dir = f\"{model_dir}/trained_model/trial_1/validation\"\n",
+ "\n",
+ "aiplatform.start_upload_tb_log(\n",
+ " tensorboard_id=tensorboard.name,\n",
+ " tensorboard_experiment_name=train_hpt_job.name,\n",
+ " logdir=tensorboard_log_dir,\n",
+ " description=\"validation\"\n",
+ " )\n",
+ "aiplatform.end_upload_tb_log()"
+ ],
+ "metadata": {
+ "id": "xJeXd3RAL5XW"
+ },
+ "execution_count": null,
+ "outputs": []
+ }
+ ]
+}
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/model_retraining/config/config_v1.yaml b/official/projects/waste_identification_ml/model_retraining/config/config_v1.yaml
new file mode 100644
index 00000000000..a36b077e8b3
--- /dev/null
+++ b/official/projects/waste_identification_ml/model_retraining/config/config_v1.yaml
@@ -0,0 +1,278 @@
+runtime:
+ all_reduce_alg: null
+ batchnorm_spatial_persistent: false
+ dataset_num_private_threads: null
+ default_shard_dim: -1
+ distribution_strategy: multi_worker_mirrored
+ enable_xla: true
+ loss_scale: null
+ mixed_precision_dtype: float32
+ num_cores_per_replica: 1
+ num_gpus: 4
+ num_packs: 1
+ per_gpu_thread_count: 0
+ run_eagerly: false
+ task_index: -1
+ tpu: null
+ tpu_enable_xla_dynamic_padder: null
+ worker_hosts: null
+task:
+ allow_image_summary: false
+ allowed_mask_class_ids: null
+ annotation_file: "gs://circularnet/vertex_training/test0731/benjamin_30july2024/val.json"
+ differential_privacy_config: null
+ freeze_backbone: true
+ init_checkpoint_modules: all
+ losses:
+ frcnn_box_weight: 1.0
+ frcnn_class_loss_top_k_percent: 1.0
+ frcnn_class_use_binary_cross_entropy: false
+ frcnn_class_weight: 1.0
+ frcnn_huber_loss_delta: 1.0
+ l2_weight_decay: 4.0e-05
+ loss_weight: 1.0
+ mask_weight: 1.0
+ rpn_box_weight: 1.0
+ rpn_huber_loss_delta: 0.1111111111111111
+ rpn_score_weight: 1.0
+ model:
+ anchor:
+ anchor_size: 8.0
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 1
+ backbone:
+ resnet:
+ bn_trainable: true
+ depth_multiplier: 1.0
+ model_id: 50
+ replace_stem_max_pool: false
+ resnetd_shortcut: false
+ scale_stem: true
+ se_ratio: 0.0
+ stem_type: v0
+ stochastic_depth_drop_rate: 0.0
+ type: resnet
+ decoder:
+ fpn:
+ fusion_type: sum
+ num_filters: 256
+ use_keras_layer: false
+ use_separable_conv: false
+ type: fpn
+ detection_generator:
+ apply_nms: true
+ max_num_detections: 100
+ nms_iou_threshold: 0.5
+ nms_version: v2
+ pre_nms_score_threshold: 0.05
+ pre_nms_top_k: 5000
+ soft_nms_sigma: null
+ use_cpu_nms: false
+ use_sigmoid_probability: false
+ detection_head:
+ cascade_class_ensemble: false
+ class_agnostic_bbox_pred: false
+ fc_dims: 1024
+ num_convs: 4
+ num_fcs: 1
+ num_filters: 256
+ use_separable_conv: false
+ include_mask: true
+ input_size: [512, 1024, 3]
+ mask_head:
+ class_agnostic: false
+ num_convs: 4
+ num_filters: 256
+ upsample_factor: 2
+ use_separable_conv: false
+ mask_roi_aligner:
+ crop_size: 14
+ sample_offset: 0.5
+ mask_sampler:
+ num_sampled_masks: 128
+ max_level: 6
+ min_level: 2
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.0001
+ norm_momentum: 0.997
+ use_sync_bn: true
+ num_classes: 19
+ outer_boxes_scale: 1.0
+ roi_aligner:
+ crop_size: 7
+ sample_offset: 0.5
+ roi_generator:
+ nms_iou_threshold: 0.7
+ num_proposals: 1000
+ pre_nms_min_size_threshold: 0.0
+ pre_nms_score_threshold: 0.0
+ pre_nms_top_k: 2000
+ test_nms_iou_threshold: 0.7
+ test_num_proposals: 1000
+ test_pre_nms_min_size_threshold: 0.0
+ test_pre_nms_score_threshold: 0.0
+ test_pre_nms_top_k: 1000
+ use_batched_nms: false
+ roi_sampler:
+ background_iou_high_threshold: 0.5
+ background_iou_low_threshold: 0.0
+ cascade_iou_thresholds: null
+ foreground_fraction: 0.25
+ foreground_iou_threshold: 0.5
+ mix_gt_boxes: true
+ num_sampled_rois: 512
+ rpn_head:
+ num_convs: 1
+ num_filters: 256
+ use_separable_conv: false
+ name: null
+ per_category_metrics: true
+ train_data:
+ apply_tf_data_service_before_batching: false
+ block_length: 1
+ cache: false
+ cycle_length: null
+ decoder:
+ simple_decoder:
+ attribute_names: []
+ mask_binarize_threshold: null
+ regenerate_source_id: false
+ type: simple_decoder
+ deterministic: null
+ drop_remainder: true
+ dtype: bfloat16
+ enable_shared_tf_data_service_between_parallel_trainers: false
+ enable_tf_data_service: false
+ file_type: tfrecord
+ global_batch_size: 16
+ is_training: true
+ num_examples: -1
+ parser:
+ aug_rand_hflip: true
+ aug_rand_vflip: false
+ aug_scale_max: 1.25
+ aug_scale_min: 0.8
+ aug_type: null
+ mask_crop_size: 112
+ match_threshold: 0.5
+ max_num_instances: 100
+ num_channels: 3
+ rpn_batch_size_per_im: 256
+ rpn_fg_fraction: 0.5
+ rpn_match_threshold: 0.7
+ rpn_unmatched_threshold: 0.3
+ skip_crowd_during_training: true
+ unmatched_threshold: 0.5
+ prefetch_buffer_size: null
+ seed: null
+ sharding: true
+ shuffle_buffer_size: 10000
+ tf_data_service_address: null
+ tf_data_service_job_name: null
+ tfds_as_supervised: false
+ tfds_data_dir: ''
+ tfds_name: ''
+ tfds_skip_decoding_feature: ''
+ tfds_split: ''
+ trainer_id: null
+ weights: null
+ use_approx_instance_metrics: false
+ use_coco_metrics: true
+ use_wod_metrics: false
+ validation_data:
+ apply_tf_data_service_before_batching: false
+ block_length: 1
+ cache: false
+ cycle_length: null
+ decoder:
+ simple_decoder:
+ attribute_names: []
+ mask_binarize_threshold: null
+ regenerate_source_id: false
+ type: simple_decoder
+ deterministic: null
+ drop_remainder: false
+ dtype: bfloat16
+ enable_shared_tf_data_service_between_parallel_trainers: false
+ enable_tf_data_service: false
+ file_type: tfrecord
+ global_batch_size: 16
+ is_training: false
+ num_examples: -1
+ parser:
+ aug_rand_hflip: false
+ aug_rand_vflip: false
+ aug_scale_max: 1.0
+ aug_scale_min: 1.0
+ aug_type: null
+ mask_crop_size: 112
+ match_threshold: 0.5
+ max_num_instances: 100
+ num_channels: 3
+ rpn_batch_size_per_im: 256
+ rpn_fg_fraction: 0.5
+ rpn_match_threshold: 0.7
+ rpn_unmatched_threshold: 0.3
+ skip_crowd_during_training: true
+ unmatched_threshold: 0.5
+ prefetch_buffer_size: null
+ seed: null
+ sharding: true
+ shuffle_buffer_size: 10000
+ tf_data_service_address: null
+ tf_data_service_job_name: null
+ tfds_as_supervised: false
+ tfds_data_dir: ''
+ tfds_name: ''
+ tfds_skip_decoding_feature: ''
+ tfds_split: ''
+ trainer_id: null
+ weights: null
+trainer:
+ allow_tpu_summary: false
+ best_checkpoint_eval_metric: ''
+ best_checkpoint_export_subdir: ''
+ best_checkpoint_metric_comp: higher
+ checkpoint_interval: 500
+ continuous_eval_timeout: 3600
+ eval_tf_function: true
+ eval_tf_while_loop: false
+ loss_upper_bound: 1000000.0
+ max_to_keep: 5
+ optimizer_config:
+ ema: null
+ learning_rate:
+ stepwise:
+ boundaries: [15000, 20000]
+ name: PiecewiseConstantDecay
+ offset: 0
+ values: [0.2, 0.02, 0.002]
+ type: stepwise
+ optimizer:
+ sgd:
+ clipnorm: null
+ clipvalue: null
+ decay: 0.0
+ global_clipnorm: null
+ momentum: 0.9
+ name: SGD
+ nesterov: false
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0.0067
+ warmup_steps: 500
+ type: linear
+ preemption_on_demand_checkpoint: true
+ recovery_begin_steps: 0
+ recovery_max_trials: 0
+ steps_per_loop: 512
+ summary_interval: 512
+ train_steps: 160000
+ train_tf_function: true
+ train_tf_while_loop: true
+ validation_interval: 512
+ validation_steps: 512
+ validation_summary_subdir: validation
diff --git a/official/projects/waste_identification_ml/pre_processing/JSON_Generation_for_Training.ipynb b/official/projects/waste_identification_ml/pre_processing/JSON_Generation_for_Training.ipynb
new file mode 100644
index 00000000000..de571c5d5d5
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/JSON_Generation_for_Training.ipynb
@@ -0,0 +1,1086 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0JmF5ohLbPlF"
+ },
+ "source": [
+ "# Pre processing steps of a COCO JSON annotated file"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "uXwwz3PlbUX2"
+ },
+ "source": [
+ "Given a single COCO annotated JSON file, your goal is to pre-process in order to remove noise and manipulate it into a form which is suitable for training a ML model. This script will also check if the annotated images are broken or missing.The output of this notebook should be 2 JSON annotation files -\n",
+ "`Material and its sub type (e.g Plastics_HDPE)` and `'Material form and its sub type (e.g.Paper-Products-White-Paper)'`"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "E1SxGZD2bv8E"
+ },
+ "source": [
+ "The COCO annotation file includes the following -\n",
+ "\n",
+ "1. Name of the images.\n",
+ "\n",
+ "2. Dimensions of the images.\n",
+ "\n",
+ "3. Classes in the image category.\n",
+ "\n",
+ "4. Name of the super categories of the classes.\n",
+ "\n",
+ "5. Area acquired by the segmented pixels in an image.\n",
+ "\n",
+ "6. Bounding box co-ordinates.\n",
+ "\n",
+ "7. Annotated segmentation coordinates."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "j0v31gxTbweO"
+ },
+ "source": [
+ "There is a lot of noise in the real world annotation file. The images name could be wrong. The images mentioned in an annotation file may not be present in the image folder, which will disrupt the model training procedure. The contents within an annotation file may not match with each other. Even the files present in an image folder may be broken or truncated, which will cause errors while reading image files. Our goal is to eradicate all these problems."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "PyFn96EKb7A-"
+ },
+ "source": [
+ "Our goal is to make sure that all information in the key values corresponds to each other correctly. This notebook will help you achieve this task."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "W6aXxxox0DDa"
+ },
+ "source": [
+ "## Import labels and sample JSON file\n",
+ "To import total classes for the material, material_form and plastic_type we will import the label files from the waste_identification_ml project from Tensorflow Model Garden.\n",
+ "We will also import a noisy sample JSON file to illustrate an example."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "WluEHMZYm0zM",
+ "outputId": "7bdc802d-4b11-4959-8255-64099ff530a4"
+ },
+ "outputs": [
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 3862 100 3862 0 0 18499 0 --:--:-- --:--:-- --:--:-- 18567\n",
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 2427 100 2427 0 0 8695 0 --:--:-- --:--:-- --:--:-- 8667\n",
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 264 100 264 0 0 935 0 --:--:-- --:--:-- --:--:-- 936\n",
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 422 100 422 0 0 1525 0 --:--:-- --:--:-- --:--:-- 1528\n",
+ "mkdir: cannot create directory ‘image_folder’: File exists\n",
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 3303k 100 3303k 0 0 4518k 0 --:--:-- --:--:-- --:--:-- 4518k\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%bash\n",
+ "curl -O \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\\\n",
+ "\"projects/waste_identification_ml/two_model_inference/labels.py\"\n",
+ "\n",
+ "curl -O \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\\\n",
+ "\"projects/waste_identification_ml/pre_processing/config/sample_json/dataset.json\"\n",
+ "\n",
+ "\n",
+ "curl -O \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\\\n",
+ "\"projects/waste_identification_ml/pre_processing/config/data/\"\\\n",
+ "\"two_model_strategy_material.csv\"\n",
+ "\n",
+ "curl -O \"https://raw.githubusercontent.com/tensorflow/models/master/official/\"\\\n",
+ "\"projects/waste_identification_ml/pre_processing/config/data/\"\\\n",
+ "\"two_model_strategy_material_form.csv\"\n",
+ "\n",
+ "mkdir image_folder\n",
+ "\n",
+ "curl -o image_folder/image_2.png \"https://raw.githubusercontent.com/\"\\\n",
+ "\"tensorflow/models/master/official/projects/waste_identification_ml/\"\\\n",
+ "\"pre_processing/config/sample_images/image_2.png\""
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "MRhCAFlVcRm0"
+ },
+ "source": [
+ "## Import the required libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Mnxbo8GBcN2O"
+ },
+ "outputs": [],
+ "source": [
+ "import glob\n",
+ "import tqdm\n",
+ "import json\n",
+ "from PIL import Image\n",
+ "import subprocess\n",
+ "import copy\n",
+ "import os\n",
+ "from google.colab import files\n",
+ "from labels import load_labels"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "tGOCdeiucUgq"
+ },
+ "outputs": [],
+ "source": [
+ "# @title Utility Functions { display-mode: \"form\", run: \"auto\" }\n",
+ "def read_json(file):\n",
+ " \"\"\"Read any JSON file.\n",
+ "\n",
+ " Args:\n",
+ " file: path to the file\n",
+ " \"\"\"\n",
+ " with open(file) as json_file:\n",
+ " data = json.load(json_file)\n",
+ " return data\n",
+ "\n",
+ "\n",
+ "def search_dict_value(dic, id):\n",
+ " \"\"\"Returns the key of the dictionary from its value'\n",
+ "\n",
+ " Args:\n",
+ " dic = Mapping to search by value.\n",
+ " id = Value to search.\n",
+ " \"\"\"\n",
+ " key_list = list(dic.keys())\n",
+ " val_list = list(dic.values())\n",
+ " position = val_list.index(id)\n",
+ " return key_list[position]\n",
+ "\n",
+ "\n",
+ "def delete_truncated_images(folder_path: str) -\u003e None:\n",
+ " \"\"\"Find and delete truncated images.\n",
+ "\n",
+ " Args:\n",
+ " folder_path: path to the folder where images are saved.\n",
+ " \"\"\"\n",
+ " # path to the images folder to read its content\n",
+ " files = glob.glob(folder_path + '/*')\n",
+ " print('Total number of files in the folder:', len(files))\n",
+ "\n",
+ " num = 0\n",
+ "\n",
+ " # read all image files and remove them from the directory in case they are broken\n",
+ " for file in tqdm.tqdm(files):\n",
+ " if file.endswith(('.png','.jpg')):\n",
+ " try:\n",
+ " img = Image.open(file)\n",
+ " img.verify()\n",
+ " except:\n",
+ " num = num + 1\n",
+ " subprocess.run(['rm', file])\n",
+ " print('Broken file name: ' + file)\n",
+ " if num == 0:\n",
+ " print('\\nNo broken images found')\n",
+ " else:\n",
+ " print('Total number of broken images found:', num)\n",
+ "\n",
+ "\n",
+ "def spelling_correction(dic, typo_dict):\n",
+ " \"\"\"Correcting some common spelling mistakes.\"\"\"\n",
+ " for i in dic['categories']:\n",
+ " for old, new in typo_dict.items():\n",
+ " i['name'] = i['name'].replace(old, new)\n",
+ " return dic\n",
+ "\n",
+ "\n",
+ "def labeling_correction(dic, labels_dict, num):\n",
+ " \"\"\"Matching annotated labels with the correct labels.\n",
+ "\n",
+ " Mapping the modified labeling ID with the corresponding original ID for alignment\n",
+ " of categories.\n",
+ "\n",
+ " Args:\n",
+ " dic: JSON file read as a dictionary\n",
+ " num: keyword position inside the label\n",
+ " labels_dict: dictionary showing the labels ID of the original categories\n",
+ " \"\"\"\n",
+ " mapping_list = []\n",
+ " incorrect_labels = []\n",
+ "\n",
+ " for i in dic['categories']:\n",
+ " sp = i['name'].split('_')\n",
+ "\n",
+ " if num == 1:\n",
+ " target_value = sp[0].lower() + '_' + sp[1].lower()\n",
+ " elif num == 4:\n",
+ " target_value = sp[4].lower()\n",
+ " else:\n",
+ " raise ValueError(\"Invalid value for 'num'\")\n",
+ "\n",
+ " if target_value in labels_dict.values():\n",
+ " id_match = search_dict_value(labels_dict, target_value)\n",
+ " mapping_list.append((i['id'], target_value, id_match))\n",
+ " else:\n",
+ " incorrect_labels.append(i['id'])\n",
+ "\n",
+ " return mapping_list, incorrect_labels\n",
+ "\n",
+ "\n",
+ "def images_key(dic):\n",
+ " \"\"\"Align the data within the dictionary in the 'images' key.\n",
+ "\n",
+ " The 'image_id' parameter in the 'annotation' key is the same as 'id' in the 'images' key of the dictionary. This function\n",
+ " will also remove all image data from the 'images' key whose 'id' does not\n",
+ " match with 'image_id' in the 'annotation' key in the dictionary.\n",
+ "\n",
+ " Args:\n",
+ " dic: where the JSON file is read into\n",
+ " \"\"\"\n",
+ " image_ids = set(i['image_id'] for i in dic['annotations'])\n",
+ " new_images = [i for i in dic['images'] if i['id'] in image_ids]\n",
+ " return new_images\n",
+ "\n",
+ "\n",
+ "def annotations_key(dic, incorrect_labels, mapping_dict):\n",
+ " \"\"\"Align the data within the dictionary in the 'annotation' key.\n",
+ "\n",
+ " Notice that the 'category_id' in the 'annotation' key is same as 'id'\n",
+ " in the 'categories' key of the dictionary.\n",
+ "\n",
+ " Args:\n",
+ " dic: where the JSON file is read into\n",
+ " \"\"\"\n",
+ " new_annotation = []\n",
+ "\n",
+ " for i in dic['annotations']:\n",
+ " id = i['category_id']\n",
+ " if (id not in incorrect_labels) and (id in [tup[0] for tup in mapping_dict]):\n",
+ " new_id = [i[2] for i in mapping_dict if i[0] == id][0]\n",
+ " i['category_id'] = new_id\n",
+ " new_annotation.append(i)\n",
+ " return new_annotation\n",
+ "\n",
+ "\n",
+ "def annotated_images(folder_path, dic):\n",
+ " \"\"\"Get images infromation that are mentioned in an annotation file but are not present in an image folder.\n",
+ "\n",
+ " Args:\n",
+ " folder_path: path of an image folder.\n",
+ " \"\"\"\n",
+ " # read the file names from the directory\n",
+ " files = glob.glob(folder_path + '/*')\n",
+ " files = set(map(os.path.basename, files))\n",
+ "\n",
+ " # list of images in an annotation file\n",
+ " dic['images'] = [i for i in dic['images'] if i['file_name'] in files]\n",
+ " return dic\n",
+ "\n",
+ "\n",
+ "def image_annotation_key(dic):\n",
+ " \"\"\"Check if same images are present in both \"images\" key and \"annotations\" key.\n",
+ "\n",
+ " List of the image IDs which are in the \"images\" key but NOT in \"annotation\" key.\n",
+ " Remove information if they are not present in both keys.\n",
+ "\n",
+ " Args:\n",
+ " dic: annotation file read as a dictionary\n",
+ " \"\"\"\n",
+ " images_id = [i['id'] for i in dic['images']]\n",
+ " annotation_id = [i['image_id'] for i in dic['annotations']]\n",
+ " common_list = set(images_id).intersection(annotation_id)\n",
+ " dic['images'] = [i for i in dic['images'] if i['id'] in common_list]\n",
+ " dic['annotations'] = [i for i in dic['annotations'] if i['image_id'] in common_list]\n",
+ " return dic\n",
+ "\n",
+ "\n",
+ "def categories_dictionary(list_of_objects):\n",
+ " \"\"\"Generates a list of dictionaries representing categories of objects.\n",
+ "\n",
+ " Each dictionary has an 'id' corresponding to its order in the list, a 'name'\n",
+ " taken from the input list, and a fixed 'supercategory' set as 'objects'.\n",
+ "\n",
+ " Args:\n",
+ " list_of_objects: List of object names to be used as categories.\n",
+ "\n",
+ " Returns:\n",
+ " list: List of dictionaries, each representing a category.\n",
+ "\n",
+ " Example:\n",
+ " \u003e\u003e\u003e categories_dictionary(['car', 'bus'])\n",
+ " [{'id': 1, 'name': 'car', 'supercategory': 'objects'},\n",
+ " {'id': 2, 'name': 'bus', 'supercategory': 'objects'}]\n",
+ " \"\"\"\n",
+ " objects_dictionaries = []\n",
+ " for num, m in enumerate(list_of_objects, start=1):\n",
+ " objects_dictionaries.append({\n",
+ " 'id': num,\n",
+ " 'name': m,\n",
+ " 'supercategory': 'objects'\n",
+ " })\n",
+ "\n",
+ " return objects_dictionaries\n",
+ "\n",
+ "\n",
+ "def print_incorrect_labels(incorrect_labels, data_postprocessing, m):\n",
+ " \"\"\"Prints the incorrect labels and their count.\n",
+ "\n",
+ " Args:\n",
+ " incorrect_labels: List of incorrect label IDs.\n",
+ " data_postprocessing: The data containing postprocessing details.\n",
+ " m: A tuple where the element denotes a condition value.\n",
+ " \"\"\"\n",
+ " print('\\nTotal number of incorrect labels:', len(incorrect_labels))\n",
+ " print('Incorrect labels are below: ')\n",
+ "\n",
+ " for category in data_postprocessing['categories']:\n",
+ " if category['id'] in incorrect_labels:\n",
+ " name_parts = category['name'].split('_')\n",
+ "\n",
+ " if m == 1 and len(name_parts) \u003e= 2:\n",
+ " print(f'{name_parts[0]}_{name_parts[1]}')\n",
+ " elif m == 4 and len(name_parts) \u003e= 5:\n",
+ " print(name_parts[4])\n",
+ " print('')\n",
+ "\n",
+ "\n",
+ "def print_dict_characteristics(stage_num, dic):\n",
+ " \"\"\"Prints characteristics of the dictionary after post processing.\n",
+ "\n",
+ " Args:\n",
+ " stage_num: The stage number of post processing.\n",
+ " data_postprocessing: The data containing postprocessing details.\n",
+ " \"\"\"\n",
+ " print(f'Dictionary characteristics after post processing stage {stage_num}:')\n",
+ " print('images:', len(dic['images']),\n",
+ " 'categories:', len(dic['categories']),\n",
+ " 'annotations:', len(dic['annotations']))"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "LG7OgYECSL2g"
+ },
+ "outputs": [],
+ "source": [
+ "LABELS = {\n",
+ "'material_model' : 'two_model_strategy_material.csv',\n",
+ "'material_form_model' : 'two_model_strategy_material_form.csv',\n",
+ "}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XCyyokyRoydN"
+ },
+ "outputs": [],
+ "source": [
+ "# common labeling typo errors that have occured in the past data\n",
+ "_KNOWN_TYPOS_IN_MATERIAL = {\n",
+ " '_PET_':'_PETE_',\n",
+ " 'plastic_PP_':'_Plastics_PP_',\n",
+ " '_nothing_':'_Na_',\n",
+ " 'plastic_MLP_':'Plastics_Others-MLP_',\n",
+ " 'plastic_LDPE_':'Plastics_LDPE_',\n",
+ " 'plastic_HDPE_':'Plastics_HDPE_',\n",
+ " 'Metals_Aluminium_':'Metals_Na_',\n",
+ " 'Plastics_PETEE_':'Plastics_PET_',\n",
+ " 'Plastics_peTE_':'Plastics_PET_',\n",
+ " 'Plastics_PETE_':'Plastics_PET_',\n",
+ "}\n",
+ "\n",
+ "_KNOWN_TYPOS_IN_MATERIAL_FORM = {\n",
+ " 'and': '\u0026',\n",
+ " '_Cassete_': '_Cassette_',\n",
+ " '_Toy_':'_Toys_',\n",
+ " '_Toyss_':'toys',\n",
+ " '_Cup-and-Glass_':'_Cup-\u0026-glass_',\n",
+ " '_Tanglers_':'_Tangler_',\n",
+ " '_tub_':'_Container_',\n",
+ " '_Jar_':'_Jug-\u0026-Jar_',\n",
+ " '_Mug-\u0026-Tub_':'_Container_',\n",
+ " '_nothing_':'_Na_',\n",
+ " '_Jugs_':'_Jug-\u0026-Jar_',\n",
+ " '_Cans_':'_Can_',\n",
+ " '_Bottlee_':'_Bottle_',\n",
+ " '_Tub_':'_Container_',\n",
+ " '_Flexiblesiii_':'_Flexibles_',\n",
+ " '_Paper-products-Whitepaper':'_Paper-Products-White-Paper_',\n",
+ " '_Paper-products-Other_':'_Paper-Products_',\n",
+ "}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "f-05VwsL0mCi"
+ },
+ "outputs": [],
+ "source": [
+ "# reading labels\n",
+ "images_folder_path = 'image_folder/' #@param {type:\"string\"}\n",
+ "\n",
+ "category_indices, category_index = load_labels(LABELS)\n",
+ "\n",
+ "list_of_material = category_indices[0]\n",
+ "list_of_material.remove('Na')\n",
+ "\n",
+ "list_of_material_form = category_indices[1]\n",
+ "list_of_material_form.remove('Na')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "uZb1rvoWXKAr",
+ "outputId": "d63cbf50-c3e4-4828-ebe9-242f78bbffd2"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "['Fiber_Na',\n",
+ " 'Food_Na',\n",
+ " 'Glass_Na',\n",
+ " 'Inorganic-wastes_Na',\n",
+ " 'Metals_Na',\n",
+ " 'Plastics_HDPE',\n",
+ " 'Plastics_LDPE',\n",
+ " 'Plastics_Others-HIPC',\n",
+ " 'Plastics_Others-MLP',\n",
+ " 'Plastics_Others-Tetrapak',\n",
+ " 'Plastics_PET',\n",
+ " 'Plastics_PP',\n",
+ " 'Plastics_PS',\n",
+ " 'Plastics_PVC',\n",
+ " 'Rubber-\u0026-Leather_Na',\n",
+ " 'Textiles_Na',\n",
+ " 'Wood_Na',\n",
+ " 'Yard-trimming_Na']"
+ ]
+ },
+ "execution_count": 28,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "# display labels only for 'material' model\n",
+ "list_of_material"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "xK-ae7HqXcn-",
+ "outputId": "350620a2-ae3d-441c-9636-fea73ff6083e"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "['Bag',\n",
+ " 'Battery',\n",
+ " 'Blister-pack',\n",
+ " 'Book-\u0026-magazine',\n",
+ " 'Bottle',\n",
+ " 'Box',\n",
+ " 'Brush',\n",
+ " 'Bulb',\n",
+ " 'Can',\n",
+ " 'Cards',\n",
+ " 'Carton',\n",
+ " 'Cassette-\u0026-tape',\n",
+ " 'Clamshell',\n",
+ " 'Clothes',\n",
+ " 'Container',\n",
+ " 'Cosmetic',\n",
+ " 'Cup-\u0026-glass',\n",
+ " 'Cutlery',\n",
+ " 'Electronic-devices',\n",
+ " 'Flexibles',\n",
+ " 'Foil',\n",
+ " 'Foot-wear',\n",
+ " 'Hangers',\n",
+ " 'Jug-\u0026-Jar',\n",
+ " 'Lid',\n",
+ " 'Mirror',\n",
+ " 'Office-Stationary',\n",
+ " 'Paper-Products-Others',\n",
+ " 'Paper-Products-Others-Cardboard',\n",
+ " 'Paper-Products-Others-Newspaper',\n",
+ " 'Paper-Products-Others-Whitepaper',\n",
+ " 'Pipe',\n",
+ " 'Sachets-\u0026-Pouch',\n",
+ " 'Scissor',\n",
+ " 'Tangler',\n",
+ " 'Toys',\n",
+ " 'Tray',\n",
+ " 'Tube']"
+ ]
+ },
+ "execution_count": 29,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "# display labels only for 'material_form' model\n",
+ "list_of_material_form"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0OoDmNC22ycz"
+ },
+ "source": [
+ "## Find and delete truncated images from the image folder."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "bUUu3F6I20w3",
+ "outputId": "a1ae3164-3f68-4745-bdee-23210113fd64"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Total number of files in the folder: 1\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "100%|██████████| 1/1 [00:00\u003c00:00, 41.28it/s]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "\n",
+ "No broken images found\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "\n"
+ ]
+ }
+ ],
+ "source": [
+ "delete_truncated_images(images_folder_path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "65XuyPBSea7-"
+ },
+ "source": [
+ "## Perform operations on the file\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "l-uMtZK2edPY",
+ "outputId": "771215ed-8950-419e-c7e8-380529f6beb8"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "dict_keys(['images', 'annotations', 'categories'])\n"
+ ]
+ }
+ ],
+ "source": [
+ "# read json file and it should contain at least the three keys as shown below\n",
+ "path_to_json = 'dataset.json' #@param {type:\"string\"}\n",
+ "data = read_json(path_to_json)\n",
+ "print(data.keys())\n",
+ "\n",
+ "# create a copy to compare the results in the end\n",
+ "data_preprocessing = copy.deepcopy(data)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "G8w7MfDtvDIq",
+ "outputId": "6a48ea9d-084b-493a-e5ca-80c48fa5cbf7"
+ },
+ "outputs": [
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "100%|██████████| 6/6 [00:00\u003c00:00, 45590.26it/s]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "\n",
+ "Total number of wrong annotated labels are 5\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "\n"
+ ]
+ }
+ ],
+ "source": [
+ "# checking labeling mistakes as all annotated labels should have 6 keywords connected by '_'\n",
+ "num = 0\n",
+ "for i in tqdm.tqdm(data['categories']):\n",
+ " if len(i['name'].split('_')) != 6:\n",
+ " num += 1\n",
+ "print('\\nTotal number of wrong annotated labels are', num)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "q2jOWegZxPEp",
+ "outputId": "e988f4b0-e254-4fbe-80f3-99fe3a3504e7"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "\n",
+ "Total number of labels which has less than 6 keywords are 0\n"
+ ]
+ }
+ ],
+ "source": [
+ "# remove category labels which has less than 6 keywords\n",
+ "categories = []\n",
+ "num = 0\n",
+ "for i in data['categories']:\n",
+ " if len(i['name'].split('_')) \u003e= 6:\n",
+ " categories.append(i)\n",
+ " else:\n",
+ " num += 1\n",
+ "print('\\nTotal number of labels which has less than 6 keywords are', num)\n",
+ "data['categories'] = categories"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "qup_-ReIz-iv",
+ "outputId": "13fae63a-e34f-4e1a-f071-475dfacf028f"
+ },
+ "outputs": [
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "100%|██████████| 6/6 [00:00\u003c00:00, 51463.85it/s]\n"
+ ]
+ }
+ ],
+ "source": [
+ "# According to the collected data it was found that most issues occurs from the\n",
+ "# 6th keyword which are the sub category of the material form.\n",
+ "\n",
+ "for i in tqdm.tqdm(data['categories']):\n",
+ " l1 = i['name'].split('_')[:5]\n",
+ " l2 = i['name'].split('_')[5:]\n",
+ " l1.append('-'.join(l2))\n",
+ " i['name'] = '_'.join(l1)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "dYqGTRluopxb",
+ "outputId": "8fb71664-1bf1-41af-ae05-3b89492a0a44"
+ },
+ "outputs": [
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "100%|██████████| 6/6 [00:00\u003c00:00, 43464.29it/s]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "\n",
+ "Total number of wrong annotated labels are 0\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "\n"
+ ]
+ }
+ ],
+ "source": [
+ "# checking labeling mistakes as all annotated labels should have 6 keywords connected by '_'\n",
+ "num = 0\n",
+ "for i in tqdm.tqdm(data['categories']):\n",
+ " if len(i['name'].split('_')) != 6:\n",
+ " num += 1\n",
+ "print('\\nTotal number of wrong annotated labels are', num)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "MfXv6qrTpDA-",
+ "outputId": "111cf89b-6be4-4103-926a-9949fb235139"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Dictionary characteristics before processing :\n",
+ "images: 2 categories: 6 annotations: 6\n"
+ ]
+ }
+ ],
+ "source": [
+ "print('Dictionary characteristics before processing :')\n",
+ "print('images:',len(data_preprocessing['images']),'categories:', len(data_preprocessing['categories']),'annotations:',len(data_preprocessing['annotations']))"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "Ro8KNGaGFv7k",
+ "outputId": "05c388a4-942d-4620-d502-742e6cb989f7"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Dictionary characteristics after post processing stage 1:\n",
+ "images: 2 categories: 6 annotations: 6\n",
+ "\n",
+ "Total number of incorrect labels: 1\n",
+ "Incorrect labels are below: \n",
+ "Plastics_na\n",
+ "\n",
+ "Dictionary characteristics after post processing stage 2:\n",
+ "images: 2 categories: 18 annotations: 6\n",
+ "Dictionary characteristics after post processing stage 3:\n",
+ "images: 2 categories: 18 annotations: 5\n",
+ "Dictionary characteristics after post processing stage 4:\n",
+ "images: 1 categories: 18 annotations: 5\n",
+ "Dictionary characteristics after post processing stage 5:\n",
+ "images: 1 categories: 18 annotations: 4\n",
+ "\n",
+ "Dictionary characteristics after processing of material_type_annotation :\n",
+ "images: 1 categories: 18 annotations: 4\n",
+ "###################################################################\n",
+ "Dictionary characteristics after post processing stage 1:\n",
+ "images: 2 categories: 6 annotations: 6\n",
+ "\n",
+ "Total number of incorrect labels: 0\n",
+ "Incorrect labels are below: \n",
+ "\n",
+ "Dictionary characteristics after post processing stage 2:\n",
+ "images: 2 categories: 38 annotations: 6\n",
+ "Dictionary characteristics after post processing stage 3:\n",
+ "images: 2 categories: 38 annotations: 6\n",
+ "Dictionary characteristics after post processing stage 4:\n",
+ "images: 1 categories: 38 annotations: 6\n",
+ "Dictionary characteristics after post processing stage 5:\n",
+ "images: 1 categories: 38 annotations: 5\n",
+ "\n",
+ "Dictionary characteristics after processing of material_form_type_annotation :\n",
+ "images: 1 categories: 38 annotations: 5\n",
+ "###################################################################\n"
+ ]
+ }
+ ],
+ "source": [
+ "list_of_categories = [(list_of_material,1,'material_type_annotation.json',_KNOWN_TYPOS_IN_MATERIAL),\\\n",
+ " (list_of_material_form,4,'material_form_type_annotation.json',_KNOWN_TYPOS_IN_MATERIAL_FORM)]\n",
+ "\n",
+ "for m in list_of_categories:\n",
+ "\n",
+ " data_processing = copy.deepcopy(data)\n",
+ "\n",
+ " objects_dictionaries = categories_dictionary(m[0])\n",
+ "\n",
+ " # create a dict showing TDs corresponding to the labels \u0026 convert all words\n",
+ " # to lower case in order to eliminate case sensitive issues\n",
+ " labels_dict = dict([(i['id'], i['name'].lower()) for i in objects_dictionaries])\n",
+ "\n",
+ " # correcting grammatical errors\n",
+ " data_processing = spelling_correction(data_processing, m[3])\n",
+ " print_dict_characteristics(1, data_processing)\n",
+ "\n",
+ " # create a mapping table to map each label to the right label structure.\n",
+ " # find the incorrect labels.\n",
+ " mapping_dict, incorrect_labels = labeling_correction(data_processing, labels_dict, m[1])\n",
+ "\n",
+ " print_incorrect_labels(incorrect_labels, data_processing, m[1])\n",
+ "\n",
+ " # change the 'categories' key\n",
+ " data_processing['categories'] = objects_dictionaries\n",
+ " print_dict_characteristics(2, data_processing)\n",
+ "\n",
+ " # change the 'annotation' key\n",
+ " data_processing['annotations'] = annotations_key(data_processing, incorrect_labels, mapping_dict)\n",
+ " print_dict_characteristics(3, data_processing)\n",
+ "\n",
+ " # change the 'images' key\n",
+ " data_processing['images'] = images_key(data_processing)\n",
+ "\n",
+ " # remove data from the 'images' key not present in the image folder\n",
+ " data_processing = annotated_images(images_folder_path, data_processing)\n",
+ " print_dict_characteristics(4, data_processing)\n",
+ "\n",
+ " # align 'images' and 'annotations' key\n",
+ " data_processing = image_annotation_key(data_processing)\n",
+ " print_dict_characteristics(5, data_processing)\n",
+ "\n",
+ " # write to a new JSON file\n",
+ " with open(m[2], 'w') as opened_file:\n",
+ " opened_file.write(json.dumps(data_processing, indent=4))\n",
+ "\n",
+ " print('\\nDictionary characteristics after processing of', m[2].replace('.json','') ,':')\n",
+ " print('images:',len(data_processing['images']),'categories:', len(data_processing['categories']),'annotations:',len(data_processing['annotations']))\n",
+ " print('###################################################################')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 17
+ },
+ "id": "nC6XzQYL15Ki",
+ "outputId": "ea8e0a7d-79bc-4321-8a63-4855afc190e1"
+ },
+ "outputs": [
+ {
+ "data": {
+ "application/javascript": [
+ "\n",
+ " ((filepath) =\u003e {{\n",
+ " if (!google.colab.kernel.accessAllowed) {{\n",
+ " return;\n",
+ " }}\n",
+ " google.colab.files.view(filepath);\n",
+ " }})(\"/content/plastic_type_annotation.json\")"
+ ],
+ "text/plain": [
+ "\u003cIPython.core.display.Javascript object\u003e"
+ ]
+ },
+ "metadata": {},
+ "output_type": "display_data"
+ }
+ ],
+ "source": [
+ "# View the final JSON file\n",
+ "try:\n",
+ " files.view(m[2]) # use files.download to download the file\n",
+ "except ImportError:\n",
+ " pass"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "S7_PnincwKTE"
+ },
+ "source": [
+ "# Visualization of categories"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "Ty6WVqyywL61",
+ "outputId": "df519b20-2776-4a3c-d22a-14becb58992d"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "--2023-09-28 21:45:30-- https://raw.githubusercontent.com/tensorflow/models/master/official/projects/waste_identification_ml/pre_processing/config/visualization.py\n",
+ "Resolving raw.githubusercontent.com (raw.githubusercontent.com)... 185.199.108.133, 185.199.109.133, 185.199.110.133, ...\n",
+ "Connecting to raw.githubusercontent.com (raw.githubusercontent.com)|185.199.108.133|:443... connected.\n",
+ "HTTP request sent, awaiting response... 200 OK\n",
+ "Length: 3268 (3.2K) [text/plain]\n",
+ "Saving to: ‘visualization.py’\n",
+ "\n",
+ "\rvisualization.py 0%[ ] 0 --.-KB/s \rvisualization.py 100%[===================\u003e] 3.19K --.-KB/s in 0s \n",
+ "\n",
+ "2023-09-28 21:45:30 (24.5 MB/s) - ‘visualization.py’ saved [3268/3268]\n",
+ "\n"
+ ]
+ }
+ ],
+ "source": [
+ "# download visualization script\n",
+ "!wget https://raw.githubusercontent.com/tensorflow/models/master/official/projects/waste_identification_ml/pre_processing/config/visualization.py"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "CKXJsmvgwPpA",
+ "outputId": "884e658e-6e37-4ae1-9f28-18e03b55ba48"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "['material_type_annotation.json', 'material_form_type_annotation.json']\n"
+ ]
+ }
+ ],
+ "source": [
+ "from visualization import visualize_detailed_counts_horizontally\n",
+ "files = glob.glob('*annotation.json')\n",
+ "print(files)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "UjxwUBmOwdSe"
+ },
+ "outputs": [],
+ "source": [
+ "for file in files:\n",
+ " print(os.path.basename(file))\n",
+ " visualize_detailed_counts_horizontally(file)"
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "name": "JSON_Generation_for_Training.ipynb",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/coco_to_tfrecord.ipynb b/official/projects/waste_identification_ml/pre_processing/coco_to_tfrecord.ipynb
new file mode 100644
index 00000000000..d0a2ae0e401
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/coco_to_tfrecord.ipynb
@@ -0,0 +1,363 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "SsIv6LYT84gm"
+ },
+ "source": [
+ "# Conversion of COCO annotation JSON file to TFRecords"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "zl7o2xEW9IbX"
+ },
+ "source": [
+ "Given a COCO annotated JSON file, your goal is to convert it into a TFRecords file necessary to train with the Mask RCNN model.\n",
+ "\n",
+ "To accomplish this task, you will clone the TensorFlow Model Garden repo. The TensorFlow Model Garden is a repository with a number of different implementations of state-of-the-art (SOTA) models and modeling solutions for TensorFlow users.\n",
+ "\n",
+ "This notebook is an end to end example. When you run the notebook, it will take COCO annotated JSON train and test files as an input and will convert them into TFRecord files. You can also output sharded TFRecord files in case your training and validation data is huge. It makes it easier for the algorithm to read and access the data."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "g3OHfWQBpYVB"
+ },
+ "source": [
+ "**Note** - In this example, we assume that all our data is saved on Google drive and we will also write our outputs to Google drive. We also assume that the script will be used as a Google Colab notebook. But this can be changed according to the needs of users. They can modify this in case they are working on their local workstation, remote server or any other database. This colab notebook can be changed to a regular jupyter notebook running on a local machine according to the need of the users."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "CRwVTTPuED_1"
+ },
+ "source": [
+ "## Run the below command to connect to your google drive"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "pnsra7Zf0uGe"
+ },
+ "outputs": [],
+ "source": [
+ "!pip install -q tf-nightly\n",
+ "!pip install -q tensorflow-addons"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "bBN0CZWlD7zl"
+ },
+ "outputs": [],
+ "source": [
+ "# import libraries\n",
+ "from google.colab import drive\n",
+ "import sys"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "z5HNdeBp0x3G"
+ },
+ "outputs": [],
+ "source": [
+ "# \"opencv-python-headless\" version should be same of \"opencv-python\"\n",
+ "import pkg_resources\n",
+ "version_number = pkg_resources.get_distribution(\"opencv-python\").version\n",
+ "\n",
+ "!pip install -q opencv-python-headless==$version_number"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "i80tEP0pEJif",
+ "outputId": "cb0d8dde-8852-49eb-e6d7-33653722eee0"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Mounted at /content/gdrive\n",
+ "Successful\n"
+ ]
+ }
+ ],
+ "source": [
+ "# connect to google drive\n",
+ "drive.mount('/content/gdrive')\n",
+ "\n",
+ "# making an alias for the root path\n",
+ "try:\n",
+ " !ln -s /content/gdrive/My\\ Drive/ /mydrive\n",
+ " print('Successful')\n",
+ "except Exception as e:\n",
+ " print(e)\n",
+ " print('Not successful')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "w40-VpWXU-Hu"
+ },
+ "source": [
+ "## Clone TensorFlow Model Garden repository"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Vh42KtozpqeT"
+ },
+ "outputs": [],
+ "source": [
+ "# clone the Model Garden directory for Tensorflow where all the config files and scripts are located for this project.\n",
+ "# project folder name is - 'waste_identification_ml'\n",
+ "!git clone https://github.com/tensorflow/models.git"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "wm-k6-S4pr_B"
+ },
+ "outputs": [],
+ "source": [
+ "# Go to the model folder\n",
+ "%cd models"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "xNe2NuqjV4uW"
+ },
+ "source": [
+ "## Create TFRecord for training data"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "J9Nz75g0oJkI"
+ },
+ "outputs": [],
+ "source": [
+ "training_images_folder = '/mydrive/gtech/total_images/' #@param {type:\"string\"}\n",
+ "training_annotation_file = '/mydrive/gtech/_train.json' #@param {type:\"string\"}\n",
+ "output_folder = '/mydrive/gtech/train/' #@param {type:\"string\"}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "mjsai7PDAxgp",
+ "outputId": "c78c7eaa-36e0-48e0-ba2c-3e674bdc5402"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "I0422 00:06:23.072771 139705362556800 create_coco_tf_record.py:494] writing to output path: /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/train/\n",
+ "I0422 00:06:25.089654 139705362556800 create_coco_tf_record.py:366] Building bounding box index.\n",
+ "I0422 00:06:25.115955 139705362556800 create_coco_tf_record.py:377] 0 images are missing bboxes.\n",
+ "I0422 00:07:39.273266 139705362556800 tfrecord_lib.py:168] On image 0\n",
+ "I0422 00:09:03.214606 139705362556800 tfrecord_lib.py:168] On image 100\n",
+ "I0422 00:10:14.332473 139705362556800 tfrecord_lib.py:168] On image 200\n",
+ "I0422 00:11:11.556596 139705362556800 tfrecord_lib.py:168] On image 300\n",
+ "I0422 00:12:11.437826 139705362556800 tfrecord_lib.py:168] On image 400\n",
+ "I0422 00:13:13.166231 139705362556800 tfrecord_lib.py:168] On image 500\n",
+ "I0422 00:14:21.695016 139705362556800 tfrecord_lib.py:168] On image 600\n",
+ "I0422 00:15:24.191824 139705362556800 tfrecord_lib.py:168] On image 700\n",
+ "I0422 00:16:48.620902 139705362556800 tfrecord_lib.py:168] On image 800\n",
+ "I0422 00:17:48.565592 139705362556800 tfrecord_lib.py:168] On image 900\n",
+ "I0422 00:18:41.091029 139705362556800 tfrecord_lib.py:168] On image 1000\n",
+ "I0422 00:19:39.844225 139705362556800 tfrecord_lib.py:168] On image 1100\n",
+ "I0422 00:20:45.108587 139705362556800 tfrecord_lib.py:168] On image 1200\n",
+ "I0422 00:22:13.738559 139705362556800 tfrecord_lib.py:168] On image 1300\n",
+ "I0422 00:23:13.147292 139705362556800 tfrecord_lib.py:168] On image 1400\n",
+ "I0422 00:24:06.315325 139705362556800 tfrecord_lib.py:168] On image 1500\n",
+ "I0422 00:24:59.421572 139705362556800 tfrecord_lib.py:168] On image 1600\n",
+ "I0422 00:25:45.958540 139705362556800 tfrecord_lib.py:168] On image 1700\n",
+ "I0422 00:26:35.475085 139705362556800 tfrecord_lib.py:168] On image 1800\n",
+ "I0422 00:27:38.255803 139705362556800 tfrecord_lib.py:168] On image 1900\n",
+ "I0422 00:28:37.250636 139705362556800 tfrecord_lib.py:168] On image 2000\n",
+ "I0422 00:29:38.937792 139705362556800 tfrecord_lib.py:168] On image 2100\n",
+ "I0422 00:30:24.683607 139705362556800 tfrecord_lib.py:168] On image 2200\n",
+ "I0422 00:31:13.964802 139705362556800 tfrecord_lib.py:168] On image 2300\n",
+ "I0422 00:32:06.411041 139705362556800 tfrecord_lib.py:168] On image 2400\n",
+ "I0422 00:33:06.038232 139705362556800 tfrecord_lib.py:168] On image 2500\n",
+ "I0422 00:34:15.721037 139705362556800 tfrecord_lib.py:168] On image 2600\n",
+ "I0422 00:35:19.886712 139705362556800 tfrecord_lib.py:168] On image 2700\n",
+ "I0422 00:36:32.834578 139705362556800 tfrecord_lib.py:168] On image 2800\n",
+ "I0422 00:38:00.137243 139705362556800 tfrecord_lib.py:168] On image 2900\n",
+ "I0422 00:39:24.083769 139705362556800 tfrecord_lib.py:168] On image 3000\n",
+ "I0422 00:40:47.815561 139705362556800 tfrecord_lib.py:168] On image 3100\n",
+ "I0422 00:42:01.868806 139705362556800 tfrecord_lib.py:168] On image 3200\n",
+ "I0422 00:43:10.464518 139705362556800 tfrecord_lib.py:168] On image 3300\n",
+ "I0422 00:44:08.492330 139705362556800 tfrecord_lib.py:168] On image 3400\n",
+ "I0422 00:45:06.637591 139705362556800 tfrecord_lib.py:168] On image 3500\n",
+ "I0422 00:46:17.144057 139705362556800 tfrecord_lib.py:168] On image 3600\n",
+ "I0422 00:47:34.219212 139705362556800 tfrecord_lib.py:168] On image 3700\n",
+ "I0422 00:48:47.535176 139705362556800 tfrecord_lib.py:168] On image 3800\n",
+ "I0422 00:49:44.018001 139705362556800 tfrecord_lib.py:168] On image 3900\n",
+ "I0422 00:50:46.843277 139705362556800 tfrecord_lib.py:168] On image 4000\n",
+ "I0422 00:51:42.749161 139705362556800 tfrecord_lib.py:168] On image 4100\n",
+ "I0422 00:52:29.118489 139705362556800 tfrecord_lib.py:168] On image 4200\n",
+ "I0422 00:53:12.499863 139705362556800 tfrecord_lib.py:168] On image 4300\n",
+ "I0422 00:54:02.751904 139705362556800 tfrecord_lib.py:168] On image 4400\n",
+ "I0422 00:54:54.855237 139705362556800 tfrecord_lib.py:168] On image 4500\n",
+ "I0422 00:56:11.432259 139705362556800 tfrecord_lib.py:168] On image 4600\n",
+ "I0422 00:57:12.901312 139705362556800 tfrecord_lib.py:168] On image 4700\n",
+ "I0422 00:58:15.347571 139705362556800 tfrecord_lib.py:168] On image 4800\n",
+ "I0422 00:59:13.046698 139705362556800 tfrecord_lib.py:168] On image 4900\n",
+ "I0422 01:00:38.408758 139705362556800 tfrecord_lib.py:168] On image 5000\n",
+ "I0422 01:02:03.484946 139705362556800 tfrecord_lib.py:168] On image 5100\n",
+ "I0422 01:02:57.290261 139705362556800 tfrecord_lib.py:168] On image 5200\n",
+ "I0422 01:03:54.188467 139705362556800 tfrecord_lib.py:168] On image 5300\n",
+ "I0422 01:04:49.160263 139705362556800 tfrecord_lib.py:168] On image 5400\n",
+ "I0422 01:05:46.782065 139705362556800 tfrecord_lib.py:168] On image 5500\n",
+ "I0422 01:07:00.913060 139705362556800 tfrecord_lib.py:168] On image 5600\n",
+ "I0422 01:08:05.558512 139705362556800 tfrecord_lib.py:168] On image 5700\n",
+ "I0422 01:09:09.658477 139705362556800 tfrecord_lib.py:168] On image 5800\n",
+ "I0422 01:10:10.147291 139705362556800 tfrecord_lib.py:168] On image 5900\n",
+ "I0422 01:11:11.286698 139705362556800 tfrecord_lib.py:168] On image 6000\n",
+ "I0422 01:12:08.696386 139705362556800 tfrecord_lib.py:168] On image 6100\n",
+ "I0422 01:13:02.225769 139705362556800 tfrecord_lib.py:168] On image 6200\n",
+ "I0422 01:13:55.910152 139705362556800 tfrecord_lib.py:168] On image 6300\n",
+ "I0422 01:14:47.861520 139705362556800 tfrecord_lib.py:181] Finished writing, skipped 8 annotations.\n",
+ "I0422 01:14:47.862285 139705362556800 create_coco_tf_record.py:529] Finished writing, skipped 8 annotations.\n"
+ ]
+ }
+ ],
+ "source": [
+ "# run the script to convert your json file to TFRecord file\n",
+ "# --num_shards (how many TFRecord sharded files you want)\n",
+ "!python3 -m official.vision.data.create_coco_tf_record \\\n",
+ " --logtostderr \\\n",
+ " --image_dir=$training_images_folder \\\n",
+ " --object_annotations_file=$training_annotation_file \\\n",
+ " --output_file_prefix=$output_folder \\\n",
+ " --num_shards=100 \\\n",
+ " --include_masks=True \\\n",
+ " --num_processes=0"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "zwazp89SojMA"
+ },
+ "source": [
+ "## Create TFRecord for validation data"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "OVQn5DiFBUfv"
+ },
+ "outputs": [],
+ "source": [
+ "validation_annotation_file = '/mydrive/gtech/total_images/' #@param {type:\"string\"}\n",
+ "validation_data_folder = '/mydrive/gtech/_val.json' #@param {type:\"string\"}\n",
+ "output_folder = '/mydrive/gtech/val/' #@param {type:\"string\"}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "nWbKeLoVwXbi",
+ "outputId": "63f4fc03-43b1-424e-dfb2-200f9bbdf1e5"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "I0421 20:53:39.071351 140304098097024 create_coco_tf_record.py:494] writing to output path: /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/val/\n",
+ "I0421 20:53:40.622877 140304098097024 create_coco_tf_record.py:366] Building bounding box index.\n",
+ "I0421 20:53:40.627101 140304098097024 create_coco_tf_record.py:377] 0 images are missing bboxes.\n",
+ "I0421 20:54:41.275259 140304098097024 tfrecord_lib.py:168] On image 0\n",
+ "I0421 20:56:53.052898 140304098097024 tfrecord_lib.py:168] On image 100\n",
+ "I0421 20:59:01.886727 140304098097024 tfrecord_lib.py:168] On image 200\n",
+ "I0421 21:01:12.356394 140304098097024 tfrecord_lib.py:168] On image 300\n",
+ "I0421 21:03:03.635432 140304098097024 tfrecord_lib.py:168] On image 400\n",
+ "I0421 21:05:04.787051 140304098097024 tfrecord_lib.py:168] On image 500\n",
+ "I0421 21:06:52.991898 140304098097024 tfrecord_lib.py:168] On image 600\n",
+ "I0421 21:09:02.626780 140304098097024 tfrecord_lib.py:168] On image 700\n",
+ "I0421 21:11:39.070799 140304098097024 tfrecord_lib.py:168] On image 800\n",
+ "I0421 21:13:58.603258 140304098097024 tfrecord_lib.py:168] On image 900\n",
+ "I0421 21:16:23.214870 140304098097024 tfrecord_lib.py:168] On image 1000\n",
+ "I0421 21:18:25.072518 140304098097024 tfrecord_lib.py:168] On image 1100\n",
+ "I0421 21:20:29.223420 140304098097024 tfrecord_lib.py:168] On image 1200\n",
+ "I0421 21:22:34.431273 140304098097024 tfrecord_lib.py:168] On image 1300\n",
+ "I0421 21:24:29.066092 140304098097024 tfrecord_lib.py:168] On image 1400\n",
+ "I0421 21:26:33.851860 140304098097024 tfrecord_lib.py:168] On image 1500\n",
+ "I0421 21:28:25.426244 140304098097024 tfrecord_lib.py:168] On image 1600\n",
+ "I0421 21:28:59.923923 140304098097024 tfrecord_lib.py:181] Finished writing, skipped 2 annotations.\n",
+ "I0421 21:28:59.924295 140304098097024 create_coco_tf_record.py:529] Finished writing, skipped 2 annotations.\n"
+ ]
+ }
+ ],
+ "source": [
+ "# run the script to convert your json file to TFRecord file\n",
+ "# --num_shards (how many TFRecord sharded files you want)\n",
+ "!python3 -m official.vision.data.create_coco_tf_record --logtostderr \\\n",
+ " --image_dir=$validation_images_folder \\\n",
+ " --object_annotations_file=$validation_annotation_file \\\n",
+ " --output_file_prefix=$output_folder \\\n",
+ " --num_shards=100 \\\n",
+ " --include_masks=True \\\n",
+ " --num_processes=0"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "machine_shape": "hm",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/config/categories_list_of_dictionaries.py b/official/projects/waste_identification_ml/pre_processing/config/categories_list_of_dictionaries.py
new file mode 100644
index 00000000000..ee39fb59b0b
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/categories_list_of_dictionaries.py
@@ -0,0 +1,90 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Create a list of dictionaries for categories according to the taxonomy.
+
+Example usage-
+ build_material(MATERIAL_LIST,'material-types')
+ build_material(MATERIAL_FORM_LIST,'material-form-types')
+ build_material(MATERIAL_SUBCATEGORY_LIST,'material-subcategory-types')
+ build_material(MATERIAL_FORM_SUBCATEGORY_LIST,'material-form-subcategory-types')
+"""
+#! /usr/bin/env python
+
+from typing import List, Dict, Union
+
+MATERIAL_LIST = [
+ 'Inorganic-wastes', 'Textiles', 'Rubber-and-Leather', 'Wood', 'Food',
+ 'Plastics', 'Yard-trimming', 'Fiber', 'Glass', 'Metals'
+]
+
+MATERIAL_FORM_LIST = [
+ 'Flexibles', 'Bottle', 'Jar', 'Carton', 'Sachets-&-Pouch', 'Blister-pack',
+ 'Tray', 'Tube', 'Can', 'Tub', 'Cosmetic', 'Box', 'Clothes', 'Bulb',
+ 'Cup-&-glass', 'Book-&-magazine', 'Bag', 'Lid', 'Clamshell', 'Mirror',
+ 'Tangler', 'Cutlery', 'Cassette-&-tape', 'Electronic-devices', 'Battery',
+ 'Pen-&-pencil', 'Paper-products', 'Foot-wear', 'Scissor', 'Toys', 'Brush',
+ 'Pipe', 'Foil', 'Hangers'
+]
+
+MATERIAL_SUBCATEGORY_LIST = [
+ 'HDPE_Flexible_Color', 'HDPE_Rigid_Color', 'LDPE_Flexible_Color',
+ 'LDPE_Rigid_Color', 'PP_Flexible_Color', 'PP_Rigid_Color', 'PETE', 'PS',
+ 'PVC', 'Others-MLP', 'Others-Tetrapak', 'Others-HIPC', 'Aluminium',
+ 'Ferrous_Iron', 'Ferrous_Steel', 'Non-ferrous_Lead', 'Non-ferrous_Copper',
+ 'Non-ferrous_Zinc'
+]
+
+PLASTICS_SUBCATEGORY_LIST = [
+ 'HDPE', 'PETE', 'LDPE', 'PS', 'PP', 'PVC', 'Others-MLP', 'Others-Tetrapak',
+ 'Others-HIPC'
+]
+
+
+def build_material(category_list: List[str],
+ supercategory: str) -> List[Dict[str, Union[int, str]]]:
+ """Creates a list of dictionaries for the category classes.
+
+ Args:
+ category_list: list of categories from MATERIAL_LIST, MATERIAL_FORM_LIST,
+ MATERIAL_SUBCATEGORY_LIST, PLASTICS_SUBCATEGORY_LIST
+ supercategory: supercategory can be 'material-types', 'material-form-types',
+ 'material-subcategory-types', 'material-form-subcategory-types',
+ 'plastic-types'
+
+ Returns:
+ List of dictionaries returning categories with their IDs
+ """
+ list_of_dictionaries = []
+ for num, m in enumerate(category_list, start=1):
+ list_of_dictionaries.append({
+ 'id': num,
+ 'name': m,
+ 'supercategory': supercategory
+ })
+ return list_of_dictionaries
diff --git a/official/projects/waste_identification_ml/pre_processing/config/config.ini b/official/projects/waste_identification_ml/pre_processing/config/config.ini
new file mode 100644
index 00000000000..9c24c1c4de1
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/config.ini
@@ -0,0 +1,24 @@
+[config]
+config_folder_path = /mydrive/TFVision/pre-processing/config/
+
+[paths]
+annotation_path = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/18012022/annotations_18012022_coco.json
+images_folder_path = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/18012022/images
+new_annotation_path = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/18012022/material_annotations_18012022_coco.json
+
+[merge]
+input_files = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/20122021/material_annotations_20122021_coco.json,/mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/27122021/material_annotations_27122021_coco.json,/mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/03012022/material_annotations_03012022_coco.json,/mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/10012022/material_annotations_10012022_coco.json,/mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/18012022/material_annotations_18012022_coco.json
+output_file = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/output.json
+
+[split]
+input_file = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/output.json
+output_folder = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/
+
+[tfrecord]
+tensorflow_model_folder = /mydrive/TFVision/
+training_data_folder = /mydrive/TFVision/tfrecords/train/
+validation_data_folder = /mydrive/TFVision/tfrecords/val/
+training_images_folder = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/Total_images/train/
+training_annotation_file = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/_train.json
+validation_images_folder = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/Total_images/validation/
+validation_annotation_file = /mydrive/gtech/MRFs/Recykal/Latest_sharing_by_sanket/Google_Recykal/Taxonomy_version_2/_val.json
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/45_labels.csv b/official/projects/waste_identification_ml/pre_processing/config/data/45_labels.csv
new file mode 100644
index 00000000000..2c8efb25d9f
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/45_labels.csv
@@ -0,0 +1,45 @@
+Aluminium_Can
+Aluminium_Foil
+Battery
+Brush
+Bulb
+Fiber_Cardboard
+Fiber_Cup-&-glass
+Fiber_Paper
+Footwear
+Glass_Bottle
+Lighter
+Metals
+Metals_Bottle
+Metals_Container
+Metals_Lid
+Plastics-ABS_Electronics
+Plastics-HDPE_Bottle
+Plastics-HDPE_Container
+Plastics-HDPE_Lid
+Plastics-HDPE_Toys
+Plastics-HDPE_Tube
+Plastics-LDPE_Flexibles
+Plastics-MLP_Flexibles
+Plastics-PC_CD
+Plastics-PC_Goggles
+Plastics-PET_Blister-pack
+Plastics-PET_Bottle
+Plastics-PET_Container
+Plastics-PET_Cup-&-glass
+Plastics-PP_Comb
+Plastics-PP_Container
+Plastics-PP_Pen
+Plastics-PP_Spoon
+Plastics-PP_Straw
+Plastics-PP_Tray
+Plastics-PS
+Plastics-PS_Cup-&-glass
+Plastics-PS_Flexibles
+Plastics-PS_Hangers
+Plastics-PVC_Pipe
+Plastics-Tetrapak_Carton
+Textile_Clothes
+Textile_Flexibles
+Tire
+Wood
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/material_form_labels.csv b/official/projects/waste_identification_ml/pre_processing/config/data/material_form_labels.csv
new file mode 100644
index 00000000000..85c646e84e1
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/material_form_labels.csv
@@ -0,0 +1,35 @@
+String
+Flexibles
+Bottle
+Jar
+Carton
+Sachets-&-Pouch
+Blister-pack
+Tray
+Tube
+Can
+Tub
+Cosmetic
+Box
+Clothes
+Bulb
+Cup-&-glass
+Book-&-magazine
+Bag
+Lid
+Clamshell
+Mirror
+Tangler
+Cutlery
+Cassette-&-tape
+Electronic-devices
+Battery
+Pen-&-pencil
+Paper-products
+Foot-wear
+Scissor
+Toys
+Brush
+Pipe
+Foil
+Hangers
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/material_form_labels.pbtxt b/official/projects/waste_identification_ml/pre_processing/config/data/material_form_labels.pbtxt
new file mode 100644
index 00000000000..f0e398e74f1
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/material_form_labels.pbtxt
@@ -0,0 +1,139 @@
+# Mapping of material form label ids to string name and human readable string
+# label.
+# Source: CircularNet partners
+item {
+id: 1
+name: "Flexibles"
+}
+item {
+id: 2
+name: "Bottle"
+}
+item {
+id: 3
+name: "Jar"
+}
+item {
+id: 4
+name: "Carton"
+}
+item {
+id: 5
+name: "Sachets-&-Pouch"
+}
+item {
+id: 6
+name: "Blister-pack"
+}
+item {
+id: 7
+name: "Tray"
+}
+item {
+id: 8
+name: "Tube"
+}
+item {
+id: 9
+name: "Can"
+}
+item {
+id: 10
+name: "Tub"
+}
+item {
+id: 11
+name: "Cosmetic"
+}
+item {
+id: 12
+name: "Box"
+}
+item {
+id: 13
+name: "Clothes"
+}
+item {
+id: 14
+name: "Bulb"
+}
+item {
+id: 15
+name: "Cup-&-glass"
+}
+item {
+id: 16
+name: "Book-&-magazine"
+}
+item {
+id: 17
+name: "Bag"
+}
+item {
+id: 18
+name: "Lid"
+}
+item {
+id: 19
+name: "Clamshell"
+}
+item {
+id: 20
+name: "Mirror"
+}
+item {
+id: 21
+name: "Tangler"
+}
+item {
+id: 22
+name: "Cutlery"
+}
+item {
+id: 23
+name: "Cassette-&-tape"
+}
+item {
+id: 24
+name: "Electronic-devices"
+}
+item {
+id: 25
+name: "Battery"
+}
+item {
+id: 26
+name: "Pen-&-pencil"
+}
+item {
+id: 27
+name: "Paper-products"
+}
+item {
+id: 28
+name: "Footwear"
+}
+item {
+id: 29
+name: "Scissor"
+}
+item {
+id: 30
+name: "Toys"
+}
+item {
+id: 31
+name: "Brush"
+}
+item {
+id: 32
+name: "Pipe"
+}
+item {
+id: 33
+name: "Foil"
+}
+item {
+id: 34
+name: "Hangers"
+}
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/material_labels.csv b/official/projects/waste_identification_ml/pre_processing/config/data/material_labels.csv
new file mode 100644
index 00000000000..67843c87b84
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/material_labels.csv
@@ -0,0 +1,11 @@
+String
+Inorganic-waste
+Textiles
+Rubber-&-Leather
+Wood
+Food
+Plastics
+Yard-trimming
+Fiber
+Glass
+Metals
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/material_labels.pbtxt b/official/projects/waste_identification_ml/pre_processing/config/data/material_labels.pbtxt
new file mode 100644
index 00000000000..ccaeefaebcb
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/material_labels.pbtxt
@@ -0,0 +1,70 @@
+# Mapping of material label ids to string name and human readable string label.
+# Source: U.S. Environmental Protection Agency (http://shortn/_ktDlFMfXm3)
+item {
+id: 1
+name: "inorganic_waste"
+display_name: "Soil, bits of concrete and stones"
+}
+item {
+id: 2
+name: "textiles"
+display_name: "Discarded clothing, carpet, sheets and towels"
+}
+item {
+id: 3
+name: "rubber_and_leather"
+display_name:
+ "Rubber and leather products such as clothing footwear, tires, gaskets and "
+ "furniture"
+}
+item {
+id: 4
+name: "wood"
+display_name:
+ "Wood products such as cabinets, furniture, and packaging like crates and "
+ "pallets"
+}
+item {
+id: 5
+name: "food"
+display_name:
+ "Residential, commercial (supermarkets, food whoesale, restaurants, hotels, "
+ "sports venues) and insitutional (hospitals, offices, universities, schools, "
+ "food banks) food waste"
+}
+item {
+id: 6
+name: "plastics"
+display_name:
+ "Plastic products such as bags, sacks, wraps, cups, bottles, jugs, "
+ "containers, lids, utensils, medical devices and household items"
+}
+item {
+id: 7
+name: "yard_trimmings"
+display_name:
+ "Grass, leaves and tree and brush trimmings from residential, institutional "
+ "and commercial sources"
+}
+item {
+id: 8
+name: "fiber"
+display_name:
+ "Cardboard products such as office papers, newspapers, tissue paper, paper "
+ "plates, cups, corrugated boxes, milk cartons, and bags and sacks"
+}
+item {
+id: 9
+name: "glass"
+display_name:
+ "Glass products such as beer and soft drink bottles, wine and liquor "
+ "bottles, and bottles or jars for food, cosmetics and other products"
+}
+item {
+id: 10
+name: "metals"
+display_name:
+ "Ferrous (iron and steel), non-ferrous (lead, copper and zinc) and aluminum "
+ "products such as containers, packaging, appliances, batteries, electronics "
+ "and furniture"
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/one_model_labels.csv b/official/projects/waste_identification_ml/pre_processing/config/data/one_model_labels.csv
new file mode 100644
index 00000000000..50071f80379
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/one_model_labels.csv
@@ -0,0 +1,49 @@
+String
+Fiber_Na_Cup-&-Glass
+Fiber_Na_Paper-Products
+Fiber_Na_Paper-Products-Cardboard
+Fiber_Na_Paper-Products-Newspaper
+Fiber_Na_Paper-Products-White-Paper
+Glass_Na_Bottle
+Metals_na_Can
+Plastics_HDPE_Bottle
+Plastics_HDPE_Brush
+Plastics_HDPE_Can
+Plastics_HDPE_Cosmetic
+Plastics_HDPE_Cup-&-Glass
+Plastics_HDPE_Cutlery
+Plastics_HDPE_Jug-&-Jar
+Plastics_HDPE_Lid
+Plastics_HDPE_Mirror
+Plastics_HDPE_Na
+Plastics_HDPE_Office-Stationary
+Plastics_HDPE_Toys
+Plastics_HDPE_Tube
+Plastics_LDPE_Blister-Pack
+Plastics_LDPE_Flexibles
+Plastics_Na_Bag
+Plastics_Others-HIPC_Cup-&-Glass
+Plastics_Others-MLP_Flexibles
+Plastics_PET_Bottle
+Plastics_PET_Clamshell
+Plastics_PET_Cup-&-Glass
+Plastics_PET_Jug-&-Jar
+Plastics_PET_Lid
+Plastics_PP_Bottle
+Plastics_PP_Box
+Plastics_PP_Bulb
+Plastics_PP_Clamshell
+Plastics_PP_Jug-&-Jar
+Plastics_PP_Lid
+Plastics_PP_Office-Stationary
+Plastics_PP_Pipe
+Plastics_PS_Na
+Plastics_PVC_Na
+Rubber-&-Leather_Na_Foot-Wear
+Textiles_Na_Bag
+Textiles_Na_Clothes
+Plastics_PET_Can
+Plastics_PP_Cup-&-Glass
+Plastics_PP_Cutlery
+Plastics_PP_Hangers
+Plastics_PP_Tray
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/plastic_type_labels.csv b/official/projects/waste_identification_ml/pre_processing/config/data/plastic_type_labels.csv
new file mode 100644
index 00000000000..9d083999058
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/plastic_type_labels.csv
@@ -0,0 +1,10 @@
+String
+HDPE
+PET
+LDPE
+PS
+PP
+PVC
+Others-MLP
+Others-Tetrapak
+Others-HIPC
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/plastic_type_labels.pbtxt b/official/projects/waste_identification_ml/pre_processing/config/data/plastic_type_labels.pbtxt
new file mode 100644
index 00000000000..9d3b0488f4d
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/plastic_type_labels.pbtxt
@@ -0,0 +1,51 @@
+# Mapping of plastic type label ids to string name and human readable string
+# label.
+# Sources: U.S. Environmental Protection Agency (http://shortn/_VnW1VzdoKm) and
+# CirclarNet partners.
+item {
+id: 1
+name: "hdpe"
+display_name: "High-density polyethylene"
+}
+item {
+id: 2
+name: "pet"
+display_name: "Polyethylene terephthalate"
+}
+item {
+id: 3
+name: "ldpe"
+display_name: "Low-density polyethylene"
+}
+item {
+id: 4
+name: "ps"
+display_name: "Polystyrene"
+}
+item {
+id: 5
+name: "pp"
+display_name: "Polypropylene"
+}
+item {
+id: 6
+name: "pvc"
+display_name: "Polyvinyl chloride"
+}
+item {
+id: 7
+name: "Others-MLP"
+display_name: "Multi-Layer Plastic"
+}
+item {
+id: 8
+name: "Others-Tetrapak"
+display_name:
+ "A food-grade, aseptic carton made of layers of polyethylene, aluminum and "
+ "paperboard"
+}
+item {
+id: 9
+name: "Others-HIPC"
+display_name: "High Impact Polymer Composite"
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/two_model_strategy_material.csv b/official/projects/waste_identification_ml/pre_processing/config/data/two_model_strategy_material.csv
new file mode 100644
index 00000000000..24ead4e98bf
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/two_model_strategy_material.csv
@@ -0,0 +1,19 @@
+String
+Fiber_Na
+Food_Na
+Glass_Na
+Inorganic-wastes_Na
+Metals_Aluminium
+Plastics_HDPE
+Plastics_LDPE
+Plastics_Others-HIPC
+Plastics_Others-MLP
+Plastics_Others-Tetrapak
+Plastics_PET
+Plastics_PP
+Plastics_PS
+Plastics_PVC
+Rubber-&-Leather_Na
+Textiles_Na
+Wood_Na
+Yard-trimming_Na
diff --git a/official/projects/waste_identification_ml/pre_processing/config/data/two_model_strategy_material_form.csv b/official/projects/waste_identification_ml/pre_processing/config/data/two_model_strategy_material_form.csv
new file mode 100644
index 00000000000..ec54b76aecc
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/data/two_model_strategy_material_form.csv
@@ -0,0 +1,39 @@
+String
+Bag
+Battery
+Blister-pack
+Book-&-magazine
+Bottle
+Box
+Brush
+Bulb
+Can
+Cards
+Carton
+Cassette-&-tape
+Clamshell
+Clothes
+Container
+Cosmetic
+Cup-&-glass
+Cutlery
+Electronic-devices
+Flexibles
+Foil
+Foot-wear
+Hangers
+Jug-&-Jar
+Lid
+Mirror
+Office-Stationary
+Paper-Products-Others
+Paper-Products-Others-Cardboard
+Paper-Products-Others-Newspaper
+Paper-Products-Others-Whitepaper
+Pipe
+Sachets-&-Pouch
+Scissor
+Tangler
+Toys
+Tray
+Tube
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/IMG_6509.png b/official/projects/waste_identification_ml/pre_processing/config/sample_images/IMG_6509.png
new file mode 100644
index 00000000000..b41abf5a54c
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/IMG_6509.png differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/Sample_Rig_Illustration.png b/official/projects/waste_identification_ml/pre_processing/config/sample_images/Sample_Rig_Illustration.png
new file mode 100644
index 00000000000..ccbaa9bea98
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/Sample_Rig_Illustration.png differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg b/official/projects/waste_identification_ml/pre_processing/config/sample_images/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg
new file mode 100644
index 00000000000..10bcd073e8f
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_2.png b/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_2.png
new file mode 100644
index 00000000000..ee71c34aa44
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_2.png differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_3.jpg b/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_3.jpg
new file mode 100644
index 00000000000..34d5e2be84b
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_3.jpg differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_4.png b/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_4.png
new file mode 100644
index 00000000000..07c67c126f4
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_4.png differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/padded_image.png b/official/projects/waste_identification_ml/pre_processing/config/sample_images/padded_image.png
new file mode 100644
index 00000000000..9f4f6a39687
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/padded_image.png differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_images/sample_image_fastsam.jpeg b/official/projects/waste_identification_ml/pre_processing/config/sample_images/sample_image_fastsam.jpeg
new file mode 100644
index 00000000000..b063bab8a6b
Binary files /dev/null and b/official/projects/waste_identification_ml/pre_processing/config/sample_images/sample_image_fastsam.jpeg differ
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_json/dataset.json b/official/projects/waste_identification_ml/pre_processing/config/sample_json/dataset.json
new file mode 100644
index 00000000000..e33f12c33df
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/sample_json/dataset.json
@@ -0,0 +1 @@
+{"images":[{"height":2048,"width":2592,"id":1,"file_name":"ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg"},{"height":1080,"width":1920,"id":2,"file_name":"image_2.png"}],"annotations":[{"iscrowd":0,"image_id":1,"bbox":[832,255,729,697],"segmentation":[[994,255,990,398,989,536,996,565,971,624,832,772,852,801,937,817,1113,870,1247,903,1304,941,1363,941,1392,951,1413,834,1458,711,1511,572,1561,401,1561,386,1377,337,1059,258]],"category_id":0,"id":1,"area":311359},{"iscrowd":0,"image_id":2,"bbox":[84,305,105,116],"segmentation":[[84,346,87,377,100,399,137,420,174,417,189,378,186,340,169,319,147,304,115,311,94,325]],"category_id":1,"id":2,"area":9352},{"iscrowd":0,"image_id":2,"bbox":[671,80,232,105],"segmentation":[[689,107,697,174,742,177,798,185,887,172,903,162,897,98,887,88,858,80,829,80,810,104,766,103,671,104]],"category_id":2,"id":3,"area":17023},{"iscrowd":0,"image_id":2,"bbox":[646,235,234,376],"segmentation":[[645,243,655,282,652,311,655,346,661,383,679,419,710,483,745,554,760,590,768,603,774,607,790,611,831,591,871,556,879,530,836,465,760,311,697,261,668,235]],"category_id":3,"id":4,"area":41260},{"iscrowd":0,"image_id":2,"bbox":[54,640,342,284],"segmentation":[[355,640,187,761,60,875,53,888,87,914,105,924,160,891,227,845,281,798,303,774,348,745,394,704,395,678]],"category_id":4,"id":5,"area":30998},{"iscrowd":0,"image_id":2,"bbox":[513,622,248,202],"segmentation":[[513,733,631,665,660,664,687,659,724,641,731,632,745,622,761,646,761,667,744,678,716,693,684,714,676,743,550,824]],"category_id":5,"id":6,"area":19383}],"categories":[{"id":0,"name":"plastics_HDPE_flexible_color_SAchets-&-pouch_pouch","supercategory":"plastics_HDPE_flexible_color_SAchets-&-pouch_pouch"},{"id":1,"name":"Plastics_HDPE_Rigid_Blue_Lid_Bottle-Cap_Na_Na","supercategory":"Plastics_HDPE_Rigid_Blue_Lid_Bottle-Cap_Na_Na"},{"id":2,"name":"Plastics_peTE_Na_Clear_Bottle_Shampoo-Bottle_250Ml_Vlcc","supercategory":"Plastics_peTE_Na_Clear_Bottle_Shampoo-Bottle_250Ml_Vlcc"},{"id":3,"name":"Plastics_na_Rigid_Blue_Bottle_Hair-Oil-Bottle-500Ml_Parachute","supercategory":"Plastics_na_Rigid_Blue_Bottle_Hair-Oil-Bottle-500Ml_Parachute"},{"id":4,"name":"Plastics_HDPE_Rigid_Na_Cosmetic_Comb_Na_Na","supercategory":"Plastics_HDPE_Rigid_Na_Cosmetic_Comb_Na_Na"},{"id":5,"name":"Plastics_PETE_Na_Clear_Bottle_Energy-Drink-Bottle_250Ml_Sting-Energy","supercategory":"Plastics_PETE_Na_Clear_Bottle_Energy-Drink-Bottle_250Ml_Sting-Energy"}]}
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_json/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.json b/official/projects/waste_identification_ml/pre_processing/config/sample_json/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.json
new file mode 100644
index 00000000000..1b816e53340
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/sample_json/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.json
@@ -0,0 +1,98 @@
+{
+ "version": "4.5.13",
+ "flags": {},
+ "shapes": [
+ {
+ "label": "plastics_HDPE_flexible_color_SAchets-&-pouch_pouch",
+ "points": [
+ [
+ 994.3103448275863,
+ 255.0689655172414
+ ],
+ [
+ 990.8620689655174,
+ 398.1724137931035
+ ],
+ [
+ 989.1379310344828,
+ 536.1034482758621
+ ],
+ [
+ 996.0344827586207,
+ 565.4137931034484
+ ],
+ [
+ 971.8965517241381,
+ 624.0344827586207
+ ],
+ [
+ 832.2413793103449,
+ 772.3103448275863
+ ],
+ [
+ 852.9310344827588,
+ 801.6206896551724
+ ],
+ [
+ 937.4137931034484,
+ 817.1379310344828
+ ],
+ [
+ 1113.2758620689656,
+ 870.5862068965517
+ ],
+ [
+ 1247.7586206896553,
+ 903.344827586207
+ ],
+ [
+ 1304.6551724137933,
+ 941.2758620689656
+ ],
+ [
+ 1363.2758620689656,
+ 941.2758620689656
+ ],
+ [
+ 1392.586206896552,
+ 951.6206896551724
+ ],
+ [
+ 1413.2758620689656,
+ 834.3793103448277
+ ],
+ [
+ 1458.1034482758623,
+ 711.9655172413794
+ ],
+ [
+ 1511.5517241379312,
+ 572.3103448275863
+ ],
+ [
+ 1561.5517241379312,
+ 401.62068965517244
+ ],
+ [
+ 1561.5517241379312,
+ 386.1034482758621
+ ],
+ [
+ 1377.0689655172414,
+ 337.82758620689657
+ ],
+ [
+ 1059.8275862068967,
+ 258.51724137931035
+ ]
+ ],
+ "group_id": null,
+ "shape_type": "polygon",
+ "flags": {}
+ }
+ ],
+ "imagePath": "ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg",
+ "imageData": "/9j/4AAQSkZJRgABAQAAAQABAAD/2wBDAAgGBgcGBQgHBwcJCQgKDBQNDAsLDBkSEw8UHRofHh0aHBwgJC4nICIsIxwcKDcpLDAxNDQ0Hyc5PTgyPC4zNDL/2wBDAQkJCQwLDBgNDRgyIRwhMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjL/wAARCAgACiADASIAAhEBAxEB/8QAHwAAAQUBAQEBAQEAAAAAAAAAAAECAwQFBgcICQoL/8QAtRAAAgEDAwIEAwUFBAQAAAF9AQIDAAQRBRIhMUEGE1FhByJxFDKBkaEII0KxwRVS0fAkM2JyggkKFhcYGRolJicoKSo0NTY3ODk6Q0RFRkdISUpTVFVWV1hZWmNkZWZnaGlqc3R1dnd4eXqDhIWGh4iJipKTlJWWl5iZmqKjpKWmp6ipqrKztLW2t7i5usLDxMXGx8jJytLT1NXW19jZ2uHi4+Tl5ufo6erx8vP09fb3+Pn6/8QAHwEAAwEBAQEBAQEBAQAAAAAAAAECAwQFBgcICQoL/8QAtREAAgECBAQDBAcFBAQAAQJ3AAECAxEEBSExBhJBUQdhcRMiMoEIFEKRobHBCSMzUvAVYnLRChYkNOEl8RcYGRomJygpKjU2Nzg5OkNERUZHSElKU1RVVldYWVpjZGVmZ2hpanN0dXZ3eHl6goOEhYaHiImKkpOUlZaXmJmaoqOkpaanqKmqsrO0tba3uLm6wsPExcbHyMnK0tPU1dbX2Nna4uPk5ebn6Onq8vP09fb3+Pn6/9oADAMBAAIRAxEAPwDjv5+lHTpSHrmnAfjXzBIH170o6YpAD0NBJH0PakA4DPekwT9aXGMEUn60D6i985pCfbmgg9e1O4I/rQITjt1o4FA6+oo5I/rSAOeaAcLQegNABx1pjGnAOT3qNgG4HFPY8ciodzKxx09KtaiFJK1Mp45/Kojyc1Ko460mMcOeaQ8cGlXOABSbQT1qRCHHBHUUhBIzmhgDxSA7uDximMUYz1xikPJ/rQeD/KlGQOaAHHGMfrSnGODTAScd6cBzQwGOcDFQhsLkHNSTZAJHPtVWDL5Dqa0itLi6k+wSYweAalPCgDj2qOJTk4JqVhx71Mn0H0EZcr14ojxvxnIpzL8mBTF4wO1JPQCYbfem9TQAcHnik5PANSA0HdIARkCnkYb1pked7ZH60/PPoRTYXQqk4IHalXhcnmmByWbIxQrFhg0rNgPH3DSsQycCkHJ9qUccUCE6rTEPJJpzZJJpq5LY9e9CGPAzyOlIx4GeaU5UgZ6U103rgHk0dRDlKk9fwo25yc9KRFHQ9u9OHWk99AEbqOeaBg44o28kZGaTkNTAd/FTSeTjNKBgEnnNAwpOOhpDFJywHpSAkEenrQSWOQOlKDknngUwBmHQUNnb04pM8Eih92MUrBsJj5sZ9DQoGTg803IByfyp+ADwetMBpOW5PFMKqWzzT9oLc1GQBnOeaaAkGT3ppX5gM80IvPWnYH40bAMlBAzUZHOc1JtZpMk/WgxgAjtVJ2AWNcng04DbwTkU2PIPFSEDZnv6VLYEYX+LGCaccbR9KAMn0FIwyeegoAQsAuDnNEf3c9c0hG5hnPPWlQ4YKDx6UwBgF6jt0pYyETPb0oKlnz1xTcMc8/L6UXAG5YVEdqylycn0qYoWdTjgUMgbrjd2pxdg6D1IboMZp5xtwBzTEbFLkkj1rNgOx/CRg0wqMEflSk8lucAUiPvUcYp67gPC4RaD7CkOWOBSnJA9KT1HuICeVpc7QfWgFh+NIMgEY60CDaW6nAqLOCvGfWpG6gAnimkKASM5qtgHsc//AFqUAscCkB+X2pcgZYtgVO4ETJ833s4p4GScDANC49DgUhyzdc/SmA9yM4BppwRnvRhgTx06mmxD5cmkAxzgAnpUaFmLf3T0zUrqWUjpUEcO1wzc/wBK0ja2oIfuJBIPSmo7MfepVHGccVGATlj1NNWDYnU9TSLg7ifu0qLgAnnFKMHr0aodgDO4HApSG2kUAfxUnJ4qRlG4kHmhB3qaAbV5OWqGSJTMSV5X3qaFg2cdfSuh2cbISJT93rz6VVZSZDn8asgkjPpVRmbeSOfUVMBXIn5kpjAEjB+Y9z2p0pGfmbFMmIWMBercZNdEegbEse/zSp5FSpKFcq4BYjGaoQyHzlyeB1q5Gd02EXp1JonGwImkOFIxj+tLE5EZOOnQGoJMyBgD8wp8YkMAGfnA5zWTSsPcVSGcvtPFTeYpjO5fqKqRs0aBGbPvUmSH2miURDW+eJkA5HODSQSOcs4+UdBSxhBPJjOT1prCREA681T7B5kV26xoAF4brVeN2hQJnG6r86M0a/LkfxVnRyCZuBxmtaTXKJ2uX4c7OfmFWYyNpDfkKrxDIIB5qZEBPuOprCe7LJFwAcde1MKFZN2eD608AlgR1pr7jEdvb1qL6i0GIWBPOAKc25hUcfzwknr2p6HLZJ6DpVMQxQ2Sp4OKgRdrMAck1Kxw2Q2aVk+YNtwe9UnYEismd53dO4pJFJdcLgDrT3bMowOBQ/yuHY4A7etaX1ERSM24MnQdeKfGWRMtx6ZqKSTYiuiYJznJolUzxJJ1wK0a+4RKZizbjxjvUx+UqUPJ9qrxgJCOCSetTRFQmSeaiSS2GTI2H68VLkBqgiwqZ6ZpxkXYSOvpWMldlXI7ogqMiqkQKuQTlSKnu90kaqBzVSBSZME8D1ramvdF1LEbBQcDBJ4qQ7lbfk596rsfO6Hp0xVuIhkKk5xRLTUegqr824kYNTxg4JB4NVwcpgD6VLECIzk1jK7AlIDHbmo2cKox2qVFwpJ6mm4BRgeahMCiru8zKRk4qLyiHYNgZ7VoqFU5A+Y1UlULPluF7+9dFOeuiF0IvK2qMHkVKpBUbhjmnx5LMcYA6VWuiQgOe/Smm27MHoXTCAQwbj0p2SVIA5psDB4Rnlh1qY/eA9axno7DKiRTF/YVZ2nZkdRUh4PHSmFsfLnr2pOTYDSoZx6HoacuEO0mmK24jnpUq4dgTzUsBVI3EE0BiwAHQUcZI7mk3LtIFSMXpJ7Ckb7pb1p2Ts3U3GQCaBDSCw2nqKa5JAA5x1pwDEnNAG3pzVpgyFmyBj7vrTCo3E7c575qZowy4Bx7UEFFUKMmqUgKkgOwg1HGAqkEc96tSghGIIBPSq6BkBU8PW0XeIMevLZTpTo3G9hg7T3pi52HJxng0q52k8DFJvQC1EPkLEcjpTHjJO48E9KfG25M/lRKhfGOwrG9mFyJQVGT1HSplYYGOppjKCRk9KexwqAU3qAickkjFOiwWIPSk2/L1pVP7kkDBJqQHqwLfd4BpqkknAFDEgKB+dOx8yqPxpAImWIGPrRyDhqcBtXIzyaaRuYEdaXUBzfewaqOu9yPSrQclqiGC7eoHWqhowGMcRAelOBytNQZPzde9ObBl+Xp61QDSdoGOtOfLDjj3poUNJx271KQdh9qTdh2IOnI6CiIF3yOCKUghGBPBqFJjHMqEcVaV1oK9i5GcEk9vWhiShY/hR90kevSkz8hUj6VkARP5q5HX0FSdvTNNjQIgx+VPbHTtSe+g9iMA49hSsVzjrRg7WUfnR/CAetMBmcfWmAEHnvTuA2PSkY5eqQugq/KCSc5qIAx7UP8WamIGSAcUOoG0nqKEwHIuMA0EgDPQ0gJ3HPANKfmwKkYKMjrwKbxuyelEbZLLTmUADGBT6gN6H60AfKCTQ2OvejpxQIYSCrZHI6VHGS4K9PepGJUY71CM5AA4q4gT44x604gLwDxSL29R1owC3J6d6gBc7cLTUAEnTpTmXg5pMErkmgBOfMz1xSYyck8U4n5CRwO9M4EXFMNBrD5g2eKYFLuXzx6UpO1M9aMfu+ladAFRMbs8A1B1bKHjvU65bIaoTEcnacD0oQEjYHftS54Xjr1qIjL7Q3I74qUkBcDtQ1YCXGTjGBUWD5hXog7etSRjcDnt0qPe5mKlcKO9EUJj3CAHoAPSq+9WTap2k9KkeNvOB3cY6VFPKkW3C8+nrTiuga9SyB8nJ4FQSbjMpHAqWI74hnIqCYtkbRgA9aI7gxybd5apIgAearlgJMYPPepVOAGNDQyV2AXmiLhDu59BTQc807HzjHSo6WAcASOtRkhF96kOdlMx83rQgF4yDnrTSSJB79aQ5D57CjnzMntVWASTlto70sRdlORgUFvnOelLHjBUGjoFiK4O10565p0Tgrlu1FyMpkYyp4pIOVOQBT+yTZqQ4YJPapV5BqIjaGK9c1IpGM4qWUNIwQSKUHLZOcUcMh9aRMgbc8UAPxkEdqbjIxmkGQNo4ocFUAoEKTmQAdO/FIeZOOgoweWPfpTlI3Ak9KBsiVW83d1FS9WxTEOSxJ+Wl5346CmwHsNzE44qLbgEgcGpW+6DnrUbfKQo5xSQCtHhcDvVcqUlVR19atdZAPSomXbJnPJ704sBiHIJPU0ifKvHc0ABWIz1qIM20gtzmrtcRZByM5qBvkH16VN0TB600jO3PakhsryBQyt3NWGBwOaqyna4I6571Yyfvn8Kt7IXUZISGULywqcNhQp6k1WOQxl6k9BUwGEBPXvSYLyJvekxyMCkPORnFCD5dx5rOw9wIy5weBT3GRkUijGeOtGNwINADWBBpnVg3pUhGQQTxUbgbiAe1NAOBIbOOTScgZPegE5BIpexNADRxSFtuMjrTgMtz3pkiMQADTW4Owq428nmlGdmetM5ZeO1PPKr6DvQwGqcxlsc0nWM5HNK2do29aRsqOPxqhDWY8KRmkDDzNpNJMdq5A+aq7uQ+B+eatRugJJ5FcMg69qrrwTjoe/rTlxuJBwaEUkHJ5rSPuoHqPXBPTtxSpKsbEcn1pIxmmSRgMMdc0Oz0EXWkGM7qrzSY2nsaQAqSp5zTmj3OuRwBxUJJMY1jhA3XPQURkKB/eOeKe8RwqnjFMH+uGOi0JpoCEx+ZMHPah1UNhuhp7qc5Y8A0Tru28cVaewroEIBHrmpG4l3cADrUIT5uRn+lPZS74PQdKTSuNjCPkznAJqZgGUKp5FMRAThuTTosmVuwoYIdIoCKR97vTThJcevahQVfDnrTSAZcd/Q0kANHiUBTwBU0bAqCevpSIxxnHTtTkTG49j0pN6ARkb2wTwOalyGQknpTEPynPJFNgj2ux3daTAlAJXk4Jp4+7TMA9eoqQjAFSwEJwPWmscEUsh2gCoz97cTx2oSAeegqQfcqLqQc5qUDJIzxSYCjgfSjOQMdKDj0oGMnrUjFPtQeBScZJpRjIz0FAhTj9Kjmb9zn8qkPT61FIcjHpTjuBUCO4BJ4qYqAoGOaTHTnGKkwM49a0bAiHMnPQU1uuAKkx15prldp+lNC0IJCd3yihVO4nOcinAYBx0pibtw71fQNSRMc880rfN1NMCkg0rc4xS6jETljSEHJzxT1UKMd6TBUEnn1ouGxnzrsbOefX0qBWzksOau3QypzxVMHHSumDuiZb6AxBPTAFGRgkHmkLdAM0jDI6nrxV2DYToPU00bevSlIJ79DTRkemDVISHnnkdaYWwck8+mKOVO4HPtSOAccUINGa1vnaDjnvVkDJHFU7XLIuD9TVzJAz3rjmtSkWFODnNMugGtyM07jHeknXNuR14qae4nuYW7a3J49K1YSSFNY7gh/p61q2rboxmuiotAirEVwu2bpUO7IGO1TXQzKSTz9Kh2gH3qVsFmA5Az0pWO0etCgjmkIxz2pi2A8kHoc0gOCQaRiMgjk/yoB56nFMFYUFvz6VHK2IpOc/KaeAC3XimyA+Q5AxlSTTW410O66Cl4yMdaQD0P4Ud84rwiuouT60uf0pPfvSnnFIBc0h5pcgCk+vUUAGMigd89KdwRScYwR+NAWADjjpS9eOmaQcUvJApAIOeKGo42k96U/dzimPoNbGPSoecn07VO5BBB4qu4ZTxVREOGDn2pykAUwYPNPQnB44psZIOTnmjp1FIhOOeKcQQcmoAY3TjimM2Gw3BNSHkU0hT1PNNAHPQ0pHvTWX0p6gt1oAaBwcGlOT0oC88UpBFADSM8GoM4bb09Kkn+4dp5qrDJLvO9eP5VpCN0IuAEAZp+1WHNMHKg08gFc9DWb3GCsBxjimE5fgcVJkA01huz2oQBkEYHSkYfLxSoMCgjAODQDGlSRkHHtTicgAdcUmBtHcd6VRhCQec0+gAoKjJXn1o+6TzyaMlkz0psY8x85oSbAWNtx5GMVKOOT+FRRLt4NSjGQOopMQE84po4ORSvjcMDFBXJ4qRgevSjHYGmjOTk/Snd8Dk0wAfKcEc048Db+tJj5hyKUg9jSbEIPl5z1pMcE0pHHvTc9QRxQFx2RkYpOo9KUDA6U1SD1oH1HfdWkP3OaTA3dc0rc4FAB1XjrRyFJzyaCecUoXAxTYxjKAoycmnEfKOaCPlz3o/h5ouIbjABpjZIJp27HAxSchelNACnkDpSkH8B0pQeOBxQTwR3oBi44zjFMPQDvUhx0P400DGSe1JACsFGO59af27AVGuQN1PHK+3rQ2AY4BPamMSX4p7D5cimEkjg/hQgBgQhOeaZGgVSxPNPHPyk00dWUH6CruA9iTj3oAGwjOfahSOhpqck9qkB+CDwelMcZbI6mlGFHLZJpGIIB6e9C3AE+UgA8mpVxuIFQxquSSc+lSLgEZHWiQDmU7uRxioi22MnkEVMQMcmoHIVSD0pR1AmQgkHI9qNuR1xVaEEzFj930qxuy4OPlpyVnoA45yADnFABVie9BJ3cDHtQOpy1QAgY/ewOaQgFTnqad8vpmkbAIwKYDRnbSH7wz+FS4ycZ4pNofnPTvQgGA5DZ4H86WM7VyBj0pdu84PFNwdwx0FO4x4JHfk+tN4C5J5pxGRkkU1vvYFKwhCcAU04zmguGD/AC4pSSp+7iqsGpXB8wcHpSxpsJGc0qkK5QcUu0hunWtG7ASIeckZFP8AvHcRwKYCVBApxBC45zWYADnOafkAY60zO0le1OBwOec9qkCk+GZjjrTUQAkA81YkQbmFQg7FLA/N61vF3AcCMFfTtVeR1DY9amLKF3E8nqKjKK3OOe1ONk9QRDPh1GeB6UwgCFV9OntT3QEfN0piuxY8dOnvW8dtBDNojQBl6niplPOQMKe1RMeAGP8A9arJx5OAOn605PQB6ABsMwDEU75Q4yxLUwkFQevtT/vSE46isWUQNH+8zkkDoKkiJ+83BoeIvMGHHrSjDncvT0obTiShCSZxtGAetSOpkI28e3rTVIRgvc08HZx3qW7jILghjsJ6VjAlZjhsHtWzNgxneMH1rNCIHOOcdzXTQdtCWaFmyn5QD9c9anT5Ny54qpAdhBU8VdB3BhWVWPLIoeGyenI6U5vljAH5U1clRx+tCgkYJFc4PUZhQCP71NOFzxx6+lLnccGnliBtAzViK5jdGDbcf1p0hwu+luF3Kh3YIPSoZYf9HGTzVp3tcABVckncewFQuzMxYr1p8ayA+gHU0yRnO4DkDvWqtcWgy63ZIkIEZ6UMxARVb5ahOZFBboD1qTYX5Q7cdK15dLC6kjGQERjp24pimRjgrtUdTUYLxA/Mdx5zTN7NFhyck01HQNC2rF0GD8tTRncduOBVJSyrtQFSO5q7CrDnue9ZVIpalISaL5WfpVMKShJ9OK05OVZT1HNZrSBsYHTtSpSdgYRghDn5farMTBEAPfvURwqj3NTKoYhifwok77j6jSfLYEZHsKnSTLbSabG3mMcrgDvTkT958wx6Vm2uoExJ4B/KnYAABHWlHBFK3LZJrFgQOcNjvUUrbgFI571MwG/d3pj7Vbrya0iBHk7Rng9M1WniJ3ZboamkYqWJ7VWllLIxwRWsEJ7ElhKFyG/OrsblkJz9BWXbgtIFAwp71rIoUY/hHaisrO4Ifn93yOf50m3JHrSTHEZJ6elCksA+fwrn6XGIqCJ+BmpFJUnI69KTAxkD86RTvZiDkelD1C5Keme/rTCRjOKexIRRimOPlXP4VKAUkvEB3pjLuxk4K0/0wMUEAPnsaadgGZO0N29KUDHPc+9IxyDnOO1J95eOAKYDsDHNAG0HJ4pzcjgdqjyN/PpQOxXvIWkZSOg7VA6iPLHOTVyRSzA549KzTuW5Yk5U9a3pttWJsTqTJDjHGafjawIHXrUIJX5c5AqVWBAOeO4ptMZahyEJ7dqQEsrbuopIi2wqe3pTguQ2DxWLVmAwgleCaehEig+lNDqIyR3pyDbFkDAJ60dAEHzAegqYDaBtHNRPg4A71LkEcdMVLAQAY569jQuc4AzSMMhOePQUH+8vbtQC0HAn5eOlBOGpqgqMnr2pSBjoS38qVg3HHABNV2XCnPGetTuNxHp3qOYZyuKcQGnlB6UkxVQv940hU+Xgjp2qI5ccjGO9WkBID8+AfrUqH5SD2qqMhSyjkH86so3mLvApyQEOS0bkr83YVWJ3FAByD0q4AcEHr2zVMkROcDJ+tXCzEaH3l3e1G7KEEdOlMRgV+WpA+9iD0PWsWrMYiZyQTx2qVzuUBRTEOMg07O0dKT3HcAcHbScBSTS9gw60MBtwe/rSAjY8ZHfrTCAcNnkdfepGGFUAVH/CSBVIBwHy7sU45ZR3NA+77UZHlhgaBA+DnFBOFXA60ceX79qVeVx+dIBi4DMemaaSNvse9ObjJHQUn8IPaqHuLwQDjBoYcnmkOC2CcCkXDOcmiwgYnJPqKjHA54pzuZDgVHwzYqkgJ4sFST3pSMjApqnAwOnpT+FwT0qGAmTuHf2oc5alGA+cc0hA4OKAGEjBApG+RQPWhQS5PpTgQW5HSqAhLYYjGKa5JQY7dae45LE0xc7mxyDVruMcnHQfWkZdiFQM+tLGcE9zTZQeFPRvehbiGrhVJz0psYyuR29akyFQ45NRbWJHGA3aqCxahfdESOv8qQYIx+dEA2Ar0zSbgg9S3QVOiegDpWwnXpVRC03zSLx2FWWK4wxqvI77MKOf5VURMtp8sfA/CoJU3EAnGadBu2ZY80Ou5WPQCpWjGRgHJXINLGFC4J6UqIoYEEc08R9u3rTbAF4VvXtRxjg5xT1wFI7VGiBAdvepAkPQjNN6D+VOXncKYy8daSAa+e3SmTSBcZ6DrT/4A3X0qjcN+9CgnP8AOtIRuxbF0neVI6GnoACfemW5Jt1NPA+fpjHepfYaEmG5SOwpIMeWQDx61K3UioI4yqnvk0J6WH1JD92ljOSfemFsOB7U8NyuB170raC0YZxknpQp3MMdB3obAyO1JEcBvTtR0DQc5pHzgfypV5BJGQKQZOc8UgEBDKcHk0pICkU1CACKT5i3NOwCOpMW3PU9akXOzLfe9aiZsybccVIMkAY6U3sHUf1wBUXPmZ/Kpjycr0pmz5CalMBEJJ55pBycmmqdq5pxYFFIqrAVtu5/8Kaict2ParL/ACjC9TVbOWb2rRNtCJh69qWQZJx+dMRj5YHapTt249RUvRjKNyw3oeuO1PSXKA5xUc5UZAOCtSQ4II7dRWr+ES3EluEXCkfWnwyiZMKM4NVpELyHav4VegTZHkjBNEkkgJcYOW5BpEHBpzdME9KOikcVhcYn8IxThjBxTeqDHWlPAwOc0hCYx360zALGnE7jz2pBliT6CqQxucr9Kc5wPrSKuFJ6g0mflBPWmCHbSWHtTM5bp1pyktz+tIx/dEigBqqFYAdO9KuJARTAcHPf0pQCqYXrTEhsz+WAKZgtg8j1qSSIsV9utA5Gfzqk0kHUicknJGQarMQJAMc1cb7gC1UZfMlLHgCtIMCXAFRLjJIHHrUpBAGeaFVRuQcMaEwGr8q7iMg8VMiqXyR9KeECoEI7UkYDMfapcrhaw5Vy28jpTsdSaWMZxnoKBwxyKzuMa3K7ulUi6o+O5q58zrgGq01uN1aQtsybiSckLnp1pJCWTI5Ip0nXHpTSzBBkAVSGICSR39aU7ipx94UxDl8E8GntlQcnqab3GCgklj96mqzMhB6ipVVQhYH6UipuQvj5j1pXJEDAQgsec9aJcLmTJzSOAsLF+2KaxyWdsBCBTSBki5MYbpUpOMKOhFVkZnRuwPSpl/dovp60pIY/aBlfSmogBYimlsTHvT42YhgePp3qdbAPYEkEUpfIAFRhsDbnrUgOBgdKljB8lelRMCQuO9SFTjHrQRyoFCEBAHSnjIGPWmt1xTumDnOKTAXuT2FKOhApTnaPSow34mkA4/XinA85xwKaSAMZoB+XC0BqDMN4GaaVCqeeTQ3GDilOM5b8KaAhHHDd6efT86YwPmZJyKk/g9/T1qmCYxgCu0Go2UbCKlCjNMYAk800BARkYB4pCPLwe5qQ4GABTWU784yKtMTQpA2jnOaXBIHem5LqMHpTxkLxSH0EcYzg8n2prHkLjipMAHnp71EQQ2epoQnqyCfBJ6VnAHnHetWRPkJNZS48wjBHtXTSegpDiueD24+lNDcHilZfnpTzxnpWgm7iMefcCmZ44PXtTzyv6UzbtB5GaaBC4IwM0h6EenSnHkdefWmP32mmhI0LNh5eD3q/j8qzbHDIc9TitE8jjiuSqveKiTAjqKdICYCCOMU1eFxipCQYuDzis4aMGczMSs5NaVk2U9cVm3SgzHnBq9p+cEEDiuuprC4RJ7zd5gxxVbHIJ78/SrN4SFXPbvVVMFSWPPes4fCNvUd94Yzgd6YRheD9KdgeuKXysxHsfSnewrESsSQMd6eSMnIznpSMp3YHWkLZ+o/Wq3C/RkgCjJNMlP7hx/snmlB2gDH15pJATDIT02n+VJbgtztxkH2pyn/69IRtOQaUZ6968MoBw3tTsHGaaCMdaUkge3rSGBwckdaUHI+lJnABHPtTe/tQIeMjmg4oxkcUDmkMTOMYp+RigYA5FJgZoEBxt6GjOOO1BHbvR1XkUDGtjoRUTEk7SKm5I61XkGHzVxEKAEXIqWLleODUZOBg9KWNvmFN6oaHqTnmnMcjFCgbuvNK3NQAjKQvHSmALxk8ipM8cd6YcZHFCYhDkN60DJIINNkByCOKcoOeTVdBj+COvNJgg47+lHBbj86G21IEM7CM7jTIZVlfgU6ZkIwRTYAjZMfftWqso6gSqcvjtUmR93v60ALmlBG6smwEZsHGKRutOPUccUjDJGaEABentQRxmlxzjNAxnnigBjLnBBp+PlwBxTFA556VIHyopsCIhx0+6etOXCOVJ6jil3HGKTI3Ybr2p3YChtp5pQcsTjj2oJUn5aUHDDFSwFXr/jSYO76UuDhjigZA5/GkAhHPHWlA4NIAdxOKAOKBAhz1pSMDrQuDyxpRznJ6UMYh6H1pqnPv7U7v0pjSBSQOM9KEIcATnI/Omvk54p6k46/WgdwRzTGNUYBNA647Cjn5fSnsAOnNAA20EHHNBzwccUmB1xinEfMOaVwEIxgAUx87KecEjFJwTkZxQIhwF6DrSMCoyeakcsBwKbt4H61QwwAM0ucAkd6czL6cigEsMAcUasA/gzjn1pqg56inH7uKZnn0xSAewAAB5p6jjJ6elMALE5NPVgVAFHQBrZyB+dRjJYk9qlbIHA61DHkdRwaFsAucNwM5pGGG+Wk3ks3y0KSMHv2FVZoBwGF57UuApGPxobk4J4PeiQ7SAoyfWjcYhyQWPaoiWVcPzmp3AMe3PJpuDs4XJoTERpExbcDx6VMeARjmmxg7jT9pU+pNDYxQMjk81CRlhnoP1qyzELjpTMlFGBUpiGR7STx0p4BzwcCmjlicVIACen4ZobACcMMHJoIXpjI7mlwM4AprKQCSKSAECoMAUcMSQMUKflxjn3oLDO0cUAAA/CgHt70A9RmlHAx0NAwI+XmmnkcHNOwCCfyprZ2cHFACfwrtOaahy/zGiPIzuGT6UmN0/X5RVWEPAG7jikYbyKVgByPzpoXDZ6g0K24xrr+8znrSOTvzSlsMDTCCT1qguSKN2WzT1JOSTxUAJzyeKkUBsYY475pNAOA3Me9PwQAaaThQFIpRjHJzUsRG+Tk1CTjgD65qzJymKrsFJ6Zq4sCFwNytjrQWwcdfapSVb5e4qHaQzMTkDtWkbPcCvcBlKtglR1FMjBZFA4JqeQeZEGJx7etQoCi8nLGt4vSwuo1QROVHOKsruMQzzVWVNpJDdasRKEUMDwTRNXVxjkbLDjB9KdiR5cKcKDzToyWlzjj1oknEbfXtWTvfQNgKFXbDU2JlDFF6+1OfJjIU4ohUAk559aE7qzEkNVsNll+Y1OXXeR3qJnxFuxkjpT4h5vzMMHFJ9xkRbfG2RkA1nuQ0hKjjtWnIAsMgPSsu4JjIRBye9a0dWJk9oSPvdT0FXUlLTbXXGaoW4y/+7WhkuyAHAxRUV5WY+hJghgM4zTgefpTJs7wAcA09Bh/vfnXO0rAMxgZFMSR2JAXpUigKTmjI7Cm2AyYK8hLHGOgqDcZFZUOfepnT5855phVWjKq3ze1UnoSRyTeXHjOTiopSiw7cZLUkyMyCPIOOvFMQl3+78g963UNLivcSQqbcIvDHrSxIQN8jcdvekWTJY7eM9acJFCkE7V7E96vW1g03JRbqxDYK45JNQ3TBcFBk+tRPckLtZvlPX3qtJc7WGMY9KUYSb1C/YvLIPKyw+Y/pVtJQxUA9BzWMLklxgH8+lXIpVcArTqQ01GXHcvIQOgFUGZVYqOOefetF8BcDriqTxI31z+dZU3Yeo7crEDPFToiqwBbrUUiKArHoKdEUR9wGc0NXWgyfCqp+bAPSpIztwW5PYmkA+7xmlf5iCKyeoEoPzHFIxJfJOKcwwoHc96Qgck9RWYEbcZ71Cw53HoKnYkjIqCQgLjP1q4gyCT5iR696pyZRQNvHr61dkHykiq82MhlAxXTAmRLaYbAPWtDGBWTAcMGOQOma02fAGelZVk2xoGbzFZQelKgGMYqFXLTsAML71aUY781lLTQYjAqQfyFNT5WK9zT2UnDA4I70qEFt2M+tTfQe4Nk4XOKY+doGc4qRT82G/KmN9OBQhDi21BnvQMbG4zikUbskn8KHfCrjueaPQBqncmBRHjn2pzYXaBxmhPlY0D8hDnbx+VMGMAtUoB/CowOmeopgI3EZxz6VTlXLltn41cc5TcPyqtJ8w+Y/gK0pkuxCFLD3pFJAwDgeopSgXAzk9/apMEcY5HStWMlhfaAAOakyVBzUUTbXA6n0qcgsfrWMtwGEgLwKTl0xQzAEjPPrT0G1MNjnpRsMEcBCc9KfGcIe/pUez5CM1KjfKF9qmVrCGnrTxgcAdaao/eAHHFOD/eNJhcRCZJMdMUqnJPNRrhAafgI2PzoaAUcgEE0yXLOAOnen5w/HFRnHJ60LcCOQnzOPummlfMVh0pGyFyePenjc1ucNz64q9gAAbAnc0v3YPk596b0jx1Yd6WH/AFZU9qGBFI+QGx0qlIm99ynaasz4QMAOajVd20kYxW8LLURNbM5Xa4/GrCgB8Z4pIQNu48qKkRcAnrWMnqxgMmTI6VIR2J4qPaS+fTtUu3K+5rJgKOvFMbrzzTgcblHNNPDYPJpIYmcr9KYV9DxUi45FMHJz2q0ADOcUpB2YB6ULyxwKAOT+tIQzpx6ipEJGRjrTMdB0p6H5vp0pvYBknB579aRziPHenOuSM0h+YEd+xoQyu5wuScc9KkXhx6nrSbcg7eppeAwHcVd9BCcAn3qADYSQcgmpZGEilRnFN2rGFB5z0FVHYCdCTiiXKtk9D2pkRLDkYxUuSVwf1qHowEVtxzzSnkZNIvCdKXoce1SMbuwdvemnhWPepBgtu703G5iDTQiLrjceKY24wtt6joaGU+cGzwaVXC5B5OeK0AIiQDzzSuA6Fc81HIvzb+xp552gcetFuoMYOFxj6UzduK+gofckvXmg53gYzVpATp8w44NBI3AYyaYGZW3dj6U8H950qLANkQtIPanbflPYnpSHCvknPrUUjfKwVsGmk3oA+Bm37Xbd70Tndn07VHa7QrYbIBqd8jnHHWh6SBENupMahuCKsDLN/Sq0L7grMcN6VYXHmFv0NEr3EOAwORTFVtmCeAc9KfwXAIxTT8zEc4qUMWIswz7U1vlHpToxgHBpjcqCSfpR1AbKQIdoP41mMu51fPFaNwRsAqnDGjScdRW9N2VxO5fVfLixT84G7rTQcjJPB7UoOJCDyPT0rHcY8c8k4qFB+9YjoKkccDsRVZmKznnAHWiKuKTsTsMOpzxTsEtnt61HvLMMCno258AfWhj0HP8Ae4qJW2uUwcVJwRnvTc5YY696SGSDAGOwpp4WjlSM96OSB7daQhpIAwOp60BiTu7ChWy7DjFDDGAT1qhAMlv58UsfPGcCjHbPBpFVlJ560gHg4BWk6fLnFCZLYx0p2MyZxSKsRP6c0nChfWlf7pPekGFUAjNUibaDmUk5NV2UZY9qt5OwnFQyAeX9acWBCjfJleSO1PDYTcaaq7V9KVvuBaphsynd/dGOC3emxh0Jzx7VNOpYI2OR3pASSMnkdK2T92wrakrjOD3HpUwJaJSBxUJ+WTd2xUyYYAAcGspbFEn3icUmOT2Ip52q3sKQgF/as7gN6gU/ktxTBjOfSnZyCaBiNzz70wgjIpxG1AM9aaw4FNCGx57inEfLz9aR1x0PFP5JzmmwGcgdOKXOAVo+Y5A4oK85zQA30NN+bzTjpQDu2+goJJcDFMOgrltwI4zTcbV25xnvSndvPpTCcvtxj3poQL0C56VWcsSdvGKlLneT2qNhyMDBrSOgDwvzhakCqZdx6jvSRYMpYmpDnGKlsEG49TR9xs8c0HJ2k9BRnBDY4qQJPusPQ9aCAcntTA2TkHinEnO096VhjQSwHGMVCGJy2amJOeOwqI8OAKpARE8+uKTcCCPSnsmGxjimdIjkYz0rQSGIM4OfwqQ/MozzjtUWAp21JyF4Gc9qbHsKp2oTnK0AZiYimhxnaBgDqaUMM7OcUhAoEsRXdnPfFRspZlRW+UdRTyCrBV6jrTSMTM4XOelUgGpxuhDfN3qx9xlXOSewquFxmTHzDtU0coZQx4elIB0wyPlFNfcqjaeT1pZPl2sT/wDXodWyG9alAK2MLkYPpUwXcKjC5kBqUfMcgc96liAZU4/WkPLZA57U5jk59aaQakaQ1WDMfapV4AJqMJgZ7nqakHHy0MABODntQoAYmgYx7Gmgg7qQX0Eb7wAPWnEfpTVU5yaeuSOlNgJg7cE0j5x9aUZ4JpWC9+aXULEZ4x79aUfcpcZFKF4AzzTAZjCjFJs568VJ346Cmt1GOpouBXI2nHcUityVbrTyBuY800Lhtw5z1NagNOQOlSMSFFKRkHIx6VG2cYB60twYPINxGfwpjn5hjvTDkk5PFAPGBwKu1guNkcEkfnWYflkOevatTYu4k1lvkSnI4Fb0hNXHtkZ9ajJ3YGOOtOPIIx9M05SMYYEn3rTYm6GZOMdhTBhu+TUxXcMmmhdv3eeMU0xXQ3dk4x+VIPQHqORTjx9D1pCpHTBzQO66FnT8DI/Q1qDJx29qy7EAOfWtQH5RXNW+Ipaki46dhUyrlOtQoQKnTOKxjuDsc7egC4PrmrNhH8pY8mm3i/vz7dKmsfu8V1yfuWZMSS7w0SntVYDHNW7kZhUiqQ+6Mdu1Zw2G9GAGVwPzqRXALA9qZzyDz60h+7wv61VrhfQRhvYlT/8AXpccE0oO3oM+lGcjCjFMEtBCPxpspH2d+n3TT+jHAz61HOuIXI6FTQtx2SO5xzg9KXqf60nOaOnXvXhlCgfjQOOD0oGQM9qdwevakAhGF9jQOAcjmnA0DrzQAAEcUm0jvSnHQUgwep5pALjn0peh6cUg4PWjoRmgBD1GaX0/nQOtHPSgAI9aicDtUx/WmMvQ00wZH6A04cN0ppHPvS5IYCrBEo45xRjjI60hJHFJtNZgKpOKM46DFLjBGKTOMimA08nPpS/eAyaCCOe1JwTimA4ccCo2OT0zSM2Bj0pAABknihICGVhGvIyKmgKsAUoADEr1pyLtAwMAVbfuhrceRyMUhIwB3oByc0v8WRWYA5GMZpWPAHFAwTnHIpCMj6UAOI+br1pGIx70oOeaRuTQAgVefWlBx1HBpSRjAWjHy8np2oAjPL5xj6Um0hySM+hqQjtjFB6dM07gICQwAGacSW6daTdg8Uqj5uaTGKeuP1oYHqDRjLEAUAAHBpCEJwvvSjIHtSAckml+8fagBACFJzTupyOKQcnA4pTkMO1AAeDxUbKMin/xg9qGPpQgEx29aU8H+dJgdR1pecZPU8U0A05x049aZGX3ncAAOlSDPc044BzRewC8YwKG4PXNICcZoGMZI5qQ0AHAzjrSbsdKOSMU3O3ORjHSnYA6rkUzBOMnipB93rwaZnC4xVIAzngd6cGPQUgG3OB0oXLZ7CjYY7thulMG1csevpTxls4NMOCcZz60kIeCce1C8g0w5bODindFFAxRyMkdKQ8rgdaf0UAUwgAZPWkIYB8vHWg8hcUhO047mmngEHk1aAlB3L61GSwkx2NPiACnnmjGGyaNExgT34yKeh+TBPNMjXPUUpXkn34pMBwUfw04gKAM81GHIXCjqetKSSckUrAPIyuD3pmSflAJxSnDfdP50u1lwKEA3LMxB+WnqDnGaaQA24/SnADf/wDXoAdkDIGMUwgsOvIp20bvWm8GTnpUoBHG0DnrSDjkjrSuCwwOgqJt4bnkCqQE/AIx1pSNwyTgVCsue30oMgLYH5UcrETEAhQOPWmfyFODZbg/jTA2WIBBpahuR+WzSl1PWpsDhRTY3HY9aVvlwM8mm22May/MfSmZIAx0qUrxz6UxU3gt2FNANdRjvUL7lQEHmp2JYkDpUbHcRkVcWLYYQQp5/SpxhlCn8aR4wY8Y/OljTahLcn0pOzQ2LuVW4HIp4b2xTeAM+lLGwKkkVIBISEwO/WoGIC7jxmrHWPJ/Cq75Kgk4FOPmISMB0yPzpWQHqOD2qKB3kk25wB2qfkEg9farkmgM+VDvOOg7Coimx1G7IxyKtTFjIQQcdqpsdxw3HrXRTuxdCSRd7bQMDHekgcJ8mOlMmceSFDHA74pVRkQ46nHNaWVtQLSElm54HepGVJHANVm4i46/zp8D5cHJz6Vg03qhliXBBQfeqtHhZTGSSx6irG4o3PFRI37zex+b6UovRiJtp3FccVNGoigA7nvUcZKk5/KpQ2IznnHes5XGRyAEHA4I5rLfITdKOc1oTTMsZwvzGqUm+WIbh9TWtLuJphalZCccGrSMGKkceoqnb7VVmzk46VajcKuR+NaVNXoMtqqgD1pyEAmo2G4ginqQVH8q5HsBGT85oWRQSew7UZ5OT16UxRh8Hv2q1bqAnmFk4XJPvVb92qbN2GPXNSsXdiFGEXvVeRcsR/COhrWCRI9QoiODkVV8wQDO08+9WY8LBjGfeopVwcL8xIrePxCZXaUEdMfWq0sxC8dfWnXWShCtg1D5XyhSwHPWtopWEMZpJGwev0qWKLn5+3ap1RY/uj8ac3yrkY3Y6VT1HewsduGUtt57VPbKoJXuKbFuMe707U+2R2fI4zWFR73BFkuzE8cVXuMrOQB1q+iEAH9arzgySA45A5xXPBq5ZHE3mJlvujtUkW98jbgmoY3xnaOB2pRJIw4O0dq1aDQvK+xSOhpdw8sH9ahO7ywvr3qFRKqrt5weax9nd7hqaPZc9aTkYBHJ6mkALfM3ag8nJJ+tY9RjX44BwPaoMA8n8qnY5xg9ah9D3q4iIZiDGcHv0qqiggk849atSsGVtp4qrGCJGBFdNLYW7LhQAKNvbrUg+aNl7DuaijlG5RjOasZOw479aykA2JCGz2NSrlqi3EMqqflxUgTnA71nIe5Mw4A7VGNw7AZp+CNoJpOWcjNZoAwcFzzjoKTbujweN1PccHnimkgD1FFwGp+7GDSt93nrnig/MBxTu5GPpTGtRu45wRzSnlxSKS+CwwRTc4b+VAh7HaGpnYilYnfTeQ+3tQhkbPklcHFVJwdwweBVyRvmx1qmyEyFmOCegremIUKCwJPbmhHyCelJKAHHzfhTSoDAEcCrWoixHlSCtTIhAy3ftVeBt5CoRwas5O4kdPSspp3AY4C9BzTtpIFGB1P5U9MOpI7VNxg52oMfnSoQHHf1prYaMj0oiI3Nx24pdA3YsbYkfI6DvThIEjJPGaiRihOTy1SbQUG71oaQMGycdhS7d/zZ6UvBBJ/AUhBztz161IwHzDd6U1sEY/Wn42/L2pmMMQaAIpPmHA4pUPPB4xyKRiwY5HFRxgoCx5rVLQOo5goQsDn+lIm6VSy9TUigNEwJ+97U2CMQZQd+RRdWEVrgYbngCmiTanIP1qa5Tc2CMn0qFPuEHpWkXeIupbtiGgAHI7GpUBU9c1Xs0/dtg5FW/TcMVlU0k0MB/rc54pw+/wAdKYGAJ96dnAzzWbAX7h46U1sbt1OY8Z70zGaEAv8AEKYwJJNOOcZ4pvJB9u9NBcFHB7mnKcsQelM3YBAp6HmhjGr1IPQUA4YA80A9VFCH5yT/APqp+oK45+mO9RlNrEk44p5JPPvRs+Y56GhOwiGJh5hJHXp70AbASTyKiKsJiw5qX/e5NWwEOWxxUcrruwT16U94yXGB+tRSINucU42Amg2iLHX3p+Mqc9RUVuANwyD6VMc7frUy3AQMQgpxAA+vWk+7hsfSgk8mkCGsCFwOpoUFT16/pTm5xnoaQgksR+dCGQTEqmV5PrSx4Kbm/KiTmPAHPrTQx+Va06CuSFM7RSOFEgFAfDYJ/wDrU1PmDMf4aSuAFAMsw6dKGAK7847U1XMjEk09iCu3HFPVARby8m0Dg/pT3zu200KUnBqQq2S2eTVNh0GHDOUzketQBRNcMccL+tTogADHjNGPJ+X1pxaQDAVPOMYqRvnJzwB1qsz7JAMZJqywzFgnrRLQEVYiGG4dAeKtFi64xjIqsny5wOBVxdpRDjnvRN63AFOPlJqJ3ZAxU5qRz8x6VC4C4AGahbgyVMqM5zSEDd/WkWT93jHI7Uuc8dMjtRbUAYBu1Rx7Unxj8aljJ281EVwR71S6oCSIAMR1z3pVUBjnvS7cOT09qZgswcdKncCbPzjBxjrVVlzMeOtWVwTn1po6E+nSknYT1Io32y7CBx3qYDaxOKYiAyknrTy+CQeppvUfQM4YUxUxKWB69akIKn3qE/uwz5znsaEDJTkt7UP8rED8aXOYxgcmkYqPmPU0kDIogzOewp648zgZFNA44/ipwwAVHH41TAXk0A4brTEcliCelSKw59al6C0FQgtu9aUcsxNIvyg5/CkzhvY0hgRhD60w9NtSHA+lN/jJpoBwx09Ka6g9aVflO49KViR81LqBA2d4pjkc54pZhnBHWmSksR3rVIHcZLkIT19KhXJwe9T7BuG6q1zlWwDWsewiUP8AI27v3qW0c8DggVQUCVuWx7VoW0QRioomkkCLXRyfTrSNyc08csf5VHkliTXMMTj736U4cjPem459jTl9KbAQ5OMim4IOD1zTxyxpgxvyT160ICMbsNk85qXI256ioozuBJFORSEwfWqYlqO7FjTTwpJ5ob7xUnilIBG0Hp3oGRZ+U44Aprbsphuuc+9OK5DbTSbQZFHYVaEK4G5Qe1OwrAkdKCecnqKb0OPWkCYnqQKjPILEcipVOFf07UxxsiZtvXtTQ+gy3XjzDVjkntgdarRAkAetWU4ynf1onuJbCH5k+WmFsqoA4HWnglY84z7U1T8hAHekgHHkDHQU7vk03HTByKVjyPTtSGMAJkaogFBBHep+rAL3qLOXYjnHaqQMawO889qYCJFGOg61KMEknrjmo4sbHXGAKtbC6kS4Z2qRRg5zmmBMNjP4UpLBVwOD1qmGw84Y7KWNfvfzpCSgLkUxyXUnPHpSAdHhT8zZI6U1CI0yTkmlKhmQ44PWkPzMdwx6Ubhcay+WgAPf86mCKVZujGow+RtIwaTJVlzzmnqArMrqATyelSJlnIxwOlNKoQHxnFOhkJV26ntUvbQCQAj5qdEML6U1cshxxTgfl61DHsBOOM5zStycUh6gg9KXAL+1IQrAZNI3IAHWg8scdKMjPA4oAVzjgdqjBGM5pSct7dqjbKgCmkBNyFHHWl3BeFFRx5PB7VIQR2xSY+gvQe9DjJAxSgbjntTSDtJH4UgEJz0NO/hAFMjTCgHpUi47dKGAg4XFMJIGR1p/v29Ka3tTQEJTchDUgwSAKn2jZk1FtCuABVJisKcYPPFRSfKcdCO9TZJqN8FQe4604gUpSQcA496aWKlcHrSTOxOTgCo1ysy7zkV0paElsLuTJ4NZU+RMTWm0uJNoHWqV8pVgT36U6WjAiByB1pRwv8/amKDtyTxSn0/nWtiLEnBPtTCSBleaVQWA/WmycfgO1C3BDiNwyRimnOSaFJC4NIc9DmmVYlsmP2krgitgAbRWRanFx0/GtVhkg9vSuet8RUXoTKQ3SpU9DUKcj3FTL61gtwZlakAJsgcGksuExUupffB9egqKzX5yDyfWuq/ukdS1NjyfXBxVEgnGTWi4zG3oKznIyMHnPSs6b0sVKy1Anv1NKTwQOtK4O3I5FMzkHBwM1YNit0xnpSoGAw3QU0Dg9u+acpIGB3psSGjhicZHYdKbMuYnP+yf5UAknngCibHlSbuu04H4U1uO9zuv4RxTSc8d6eewqNlYtkdK8JFEmCAPSl7U3OeDTuc4pDAnHIHHrSAZ607tQB3oEJjNJ2penWgelIBe3FIOlKcDtSZ5zQOwDrjP4UfxYpD19/rQT0pgKcgU4jj8KaDhj3oOM9aAGFR36ik68kUucCgkAVb2AeOQMn8aVgMgd6aCduaXrz6VmAvfpTTyeacPXtSEe9NMBO/86aQB2p2M0xj0B/GmgFOCp4zTQu5cGnrwKN2Bmi4ERAxndUgYsvB4qvMu5dqtjNSW8bIuN2TV2VrgSHleKUEDrQo7ng00DIJ9KgCUEHikUfKe9RRbvWpQc/Wk1YNQH3eRSkAryOaOpApR0OT+VAxpHGccilHzNg0MeB0pCoG096AEbIODRnjJp+B3FJ6igVhknBU0/J3D3pGBIz0xQOvXjtR0AdyuTnNJg+tOGCeTxSdTikAmDnFKODSn72AKONvP5UgEHXJpwyOaOTik70AJzg+9DAYGKXqR0FDEAcUABHFIwGOKD1+lJuJ7CmgAD160uOfx4oP60HqMUXAdzyBxTTuBwOlKvU4PNGMdeKQAeSO2Ka+SDnmlBwcjqOtJyxp7MBG4XGOnQUZIA9aU4xjPPrSMSW46U7jGsTjHY0oyAATTRwCxpQMrmmIcBjnOKY4HOD0qQAljimnkEYpACYKUDnkmlXAXp+FIDg+x7UAPPIApj/K3HSnD7wJ/KhmH3j1pLRgMZctkDik2YDHNPIOOfyprEbcCqTHuNQBVHrmnZIBOPpSLyuPTvT25we1DAjGSMk809WJB4wKQrjGe9KeRtA5o0EIdyrxzinHJjOfypjoSyt0x1FSbfm3Fx9KHawAmAuPSnN0DZoXAHTk00nPrkdqnqPUWQEoB0peAg9qa5BxzSpzn2p6sB+Dx70irjJHel5YDFGQFzzUgNJ5AHApCnX3pSDgZxTVXdz2p3AaEAAOaCoAJ6E04gkjPNGNx5pgIEwM7qi2ghu2fSrABPQcUyVcY2jP0oT1Ar26gs2TwKsOcpu7CokAB45J61MwBAA5qpO7EOADx5BqFG3cAcfyqeFlXKnoBUQILEHgZrSajyKSFsxMDnB5/lUCqQwGc55qyRuJxUAAQMT1FRFlEudwpWOflGPamp843dBTud3HNSAsmFXAHSjJC4ppUscHpSjAbB6mjSwD8HHJ5HaoG+bqOKlMgAbiq7uTgA4ppMWhHDGEkyoOanIy2c02M84/WnFQq46k1UncNiu7ZZgnT1qkXxLlvujtVxoyQcjI+tV522kALW1NpMCCTf6ZUninRAO4DcA1I6KIuDn1qOES7o2b7vpW0ZJoXUnCEyHcdq9KJgpnTB4pXKM55ps6sSiL07ms1uPoTT8IpC5yaieRVYE5yKmjyISP4hVd9pbcRtI71MVrYC4rrjOc5FKJGznHyHpVYOVTKjKn9KfDIXiLbsHtUuG7AWdXYjB4Heqk7b4GGce9WJ+IRubp6VXIULtPSrp26iIYSIsK3OetaB4QKg69aoW5GfnbIHTirivsjBVck+taVFdoaJ2fDKo7inFtrKADuPamQEhAzH5qUyKimQ8t2rma1AWXiRCB+FMZ9nJIyadIyttbHJ6GoLlMocngVSV7Ji6ETXyPvjU/KeretVG3odhPLGkMZDeue1SRQ5LE5LfyrsUYw1RK13JI1Yrz61FLKcYHarTDblEOaY9sXQlmAJ6HFRzq+o2jNkQkjj60IndvXmr7iNAqj5jjk1WBwCDzW0JtisgXJBfHy/wA6j2EycdaegJfGDiraW/r1605TtqJXYsVqwhBHHNWrVRGrEj6Copn2KoU8E81YVsKRkYB/OuSUpSTbLHKSyZZar3CnepzyegFWieQBUF0mEB71jB+8DKqcM2BjJ5qRTnjbwOhpgUlzntSxum8hT+NbtNjJ0JIUEcjtTUYI5wCeadFGok3ZyakSPEj5GPSsrpMCViVG7v2pFk+bB4FGeTzQSADxWYwx1J4HaoCMbj7VYzlR3PpVd8BemeelOO4ELJsQYPWkH38nGalYEANj6Cq+4GTOMn0raOoixFsK57dqJEkcYQ4FReWzkYO3n0qwke0Hc2c9qUrLW4eoMoQL7CpsA7W6AVGwA2HHSl6MQelZvVATZxzxigHAPvSKcjHTFISQMdKzsApDEYFJuyp9qVMhQOp703aQMZ60xgWPCDqetBJ3D0pBh3yOop5+UBepo2EITu4PQUhHzY9KVnC496QkAD1NAxuOTk0dQcdqUghc0MMAGqEhgBII6VWYfMNx5q2G2vzxnpUEibWznOOvFXHQCuyqy+4PGaRgC5GcnHSnJh8/LTIwN5Yr0rYCa0Qo3zCrKdGJ5z0FQQk7wR0Pap0Vmfk8VjPfULCA4QF+MGlRgIs44PSiTkEdvbtUcQyoU5NT0AkzhCR60i8Yx+dKGIyCME0hbDBe9ACvjz8e1SdUGevaoZQA4AGSetOySFC9qTWiAkByuMUZ+UtikZhnAoLDp2pAOOSwNI5BUknmmqe+OfrTJN2zGeDQlqMTkgD0qGJiY2FKrf6OSpzikbPlEMdpNapCJUOEGepPNPABkPt3qFEzjHFTxklGyOalgRSqDhx+NRxrulP0596mZd0Z9DUcK7BjNUnoG5JbIIyVzwe1TE5zUEIByA1WAcHdjms5bgmNJ4yaf1FMIypx0pYj1PbpUvYBx5+tNGdoxwfWnDkk0mflxSGNJCgDrntQelBwAM009veqAQnC8U9cqg9TUanJ9hThlgCKbAXPBNJgkGl5yuBS4wD60gG/8ssGnDpk9aM8e5pqtuUetAA6jB4xUTEBSAeanJzx+VQOvzEetOICjJjGKjcZhC56VIgYRkCo3HyMCetWtxEcWVPDfj61azhc+tUYgQxLcD1qyzrgDnmqmtQJeqAHikLfLjsKQngj2oB+QEdqzsNoUHjI7UmcYOc7utL91vr2o3ArnHQ0CI5FwhOPwpqfMo46VJI3OfWoojjcD064q1sNoMFpNxxj0pVOY2z09KAwAznAxQFUoSB92mSNQAJjB56U8FtvoKarNJg9qHJK4Ud6Qxozktu4H61I25nBxwaiHO5T93ipuoHOMVTAawxhfTtTJt2zcvzGnqfmPFNdyDjvSV7gNZFyGPU96kUYTk8+lQASM+5vu9hVpl3KB3qpbAQRA7mDDg9KmLc7AKjTAlJXqe1TqCccfNUSeoEUi4QevrTZDsTd3NTOBs55zUEgO3kiiLuBDbnJYVMPvf0qor4Yhfu1ZSVdu7PJrSSESA8kjrRIuSMc0p/1bHODRDjZyc1n5jAtkDuKTcfMCgjHpQAEJBoK/vAwpgOU/OQOlIVGNo7Uqt+8Ixihzt5IpdQGoTvPHSlIUuNx+lIjEuec0rxCWRWz93pT6h0HMTuGPzqCTBwWPA/WpXI2deajkG/5B2oiBIhLJg9BSuN0eMcCmqcKEpdu2IjPWl1GQLL5iLjsccVKUBJJpsMYROfWpTjaAKqT10EQQfMr565p8YPU/hUkaAJ0pwGAaUpCtYZnnI7U5s8A03GQM9KGOSoFIYo55PShzg9KUEscCjrxSAOM0ONq89/0pMcn0FO5YZz2oArfMXORxTCcsCDTpZCqsOpqGOTcM461sk7XDyHFAW3d/So5kEiKpHWpiMHd3NMYEjkY96aeodCuyLGhbOMVctG+RiOtVAy4ZOT71aQiMbRyaqeqsxFnoM96D/EQOtMXcxHpTieTjoelc9hiDruPSlGfvf0pB0wOvc0AkqBTGKOmcc1GB8xJGQamGOnYd6jGfmxQhESAgEipduFGcVFEdrbfzqX7zYqpAhu3ex70oXacZoA2gnnmhTn5gOe1ICE4TIzyTQykYxTnAY7sUxH+fnpVgSbPnPpTSAWJHNOJAJPc0L8ozgAnpSARjiLOOaa67guT+BoTdsOe3rSsPlB7insBXUjzyuCCKsxnLn1FV3zhTjk9asJ8rHHWqmCBlIQ5NRMfn2E8GpgMZzycVWXGRn6ilENCyuA2CaavyknNCMGO7060rH5gRxmpAbjDFh27UxgFQlTzmnNkOB1HeonjkDDB+U9qpCDIUY6k0g+UEHqf0pSqhxz071E0m19p5zVpXGN5ViO3rUkZ3ptBqM5d854PanozKxUCqewbg+W2gHGKa6q8RTv2NDYx5YPJp+5UjPHNGwt9BJH8lFZRkGmzBywx932p0IDKAeQORTVJAPPJo2DcRZY3mDY59adCDzvGBUakiJsrhh0qTzCqIepPWm10QkKGIuMEcGnRhfNIz26UOc+zUu35jIOWPaoYyRSFXgcCnRjacVGP9XkVIp3HPXioYADiTngUq5IOaTjOe1LwDyKQB2JpF68dBSDk47UgPUetMYnAbIofLEKDyaVVzup8SY56YovYW44DkZFISRnNOBGMnr6U1snLHipGA4U+vajGfl7nrSBiWznilyTTFoLtG0gdqAcjAFHRetIPl696QDs/KM00jA96cTz7U12AGD3oQwH3MU1hyB29aCeOOlJknPpTExBgk1DNnacHGalQ5OO1R3AO0Y7Vcdw6GdMnOM8/zpkYDMNx6VNMM/N0NJENyngA/wA66k/dFy6kr5EecciqtyQY9x61d3fuunNVpkBiII5qYPUNTPUkY56U4NzyPxpnQ4xwOKUtyMY9q6WiScMFIwBxUTZ3Uq4yMUEAMe59KlaBoIThcA/NRyWwTzSemQc0q525PWmA+3LLcAnH1rZH3evNYkLfv0rYHKg5xWFZaopeZOnTBqVRk4qNRn6VIud49DXMtxso6n0Xn8ap2sqpNgn6VoanxFx61iLkSbsHPauyCvHUjc3V+ZW/WqIjLSH5uB0GKuw4ZOe4quuC7ZrKLs2XYgHy/L2FISG/Gny8N8vXvUY5JP8AEKtdybCnjORkCkzgAdT6+lAYYz3HSlBJGDTBajVBPB6UkikQPk5+U5qTqaZIc2z4PY9aaeoJHckHOSKCPSkBxmjqK8IsVcAg0vOcijHH0pQDjHekAvbrTT04/KjpxSgYbNILCHBBB607grjv9KQ/eFL60BuJnnGKByKMYHWlwCMUAIOKXBPeg8qKOelACcL0/OjAJ+alIGaRjkdKAGYw1LxnBNBJDc9O1NIwc1YD/wCEUu8YqNOnBpyg7uOlTYCQDK8U089OaVemKMcdaQdCPcccDpSSN8u7HNPJ68VG2c7atAPB4zim5BG3vQjMpA6ikCtvznigBBFkEZzT8BBgU7ZjnNKw+Wle4CBh9aFCgECow25iMU84BFDDoC4GeeacvHI70xfmPTFScCkwEHt+NOyMH2pCOetKpwOfypAGQ1NU9z2pM8tSnvgU9gFJDLgDrQeO3IpASFwB0pF+Yc8YosA8cmhhk+tC4xwaXI60gBeM+1Awp6UgI54zSjjJJpAAJyTQODgjgigEke1J93rTuAE9u9Lu+XHekOR+NKoB5NFtAFAHcc03DbuBg0DPmEjp2FOx82T3o2GR7uTinErkUoHByKTAI6UCDPPPSkJy3HSk/lSnAHvQADJyMcU48Ac0hPHJoOMc0MA6ZOcE0nReetKTk0rcEE0AMH3yCeQKQkkCjPJNNXJIJH1p2AVmJOMUuBjPf0pvmYz6+tHoBTsBIM8kU0nGQKcASNq0pAwMdqQDCCQMdKQ8MM/hSqc9e1N25PoaaAcANuPelb8sUq8KaAPkyTyO1LqAhbOcdKjcBouPzqQkhTTCPkx2NNAKu0LigYI/2RTI8ZKipCCOAeKGAHIYY6Ux3w45xSgkOWPelKggZ6igY4vgcd6EXdyaT5MZ60gbnC0rAPVsYOQRThkEk96Yy8cnGKU4xxyTQ9dxAMk9KWPqd1Cgn60kYGc9aBjxwM9qFAI9qUNzyOaEIOSTUgNzkHP50yN1IPPHpUmOPY1XRPLPr7VSsK5IMjnpntTj9KiVmDk4qUZbOfrQ9BjgRgnNNOQPc05emSKQ8jj8akBNoUdcmhfu+nvSMCWUg8UpO7AHamDGhOOvSkGOc08FdxUdqaWA7daeoC544OM9qgkXC7PXrUqhQMkGmbc7uxqloIE4iAzUg4AxyTUEWc47VP0zgUMYg6MQfpSAEcnOKcF2g980zduYjOAKQDxjGSOtQuvBFTZ2d+KjcbegxQhMjXP3etKzAYzSKQMn9aAORxVgNdhtz0qlOzg7VX5u5NW2Yk5NVptxcEHOe9aU9GACJVQM2STziojIZNuDhRT2yGCFvwqtjaAinOM5NbJXEWn2KDIBk1HIZGhDE4OeB61PCu2EBxwaVv3gYj+HpU81nYGriQHYjM5xu7VFIAkgwOSOTT45PNyD1HSmMVGcc018WoAsm5QmeR1p0EiuHHTHemxSB1I24NO3BEJK/Wh9rAOCukbFuR2qkFlZ9jNgEc+9XN0kkI6DPeofM3nCrwOhqotq4aCwGOKN1J3MtS243ghmz7CqCMI3YbOD71fgUKoZeM9aJq2oK5YyxYALgCkb5pNoX5fXNCEojFxgdvemgkneOB3rBIfUlcpsHHtTZEV4yrNxQASASacACcEfjQ2IoTRkTDyxx61EwmAZiwUGrVyVtixQ7mPX2qjMzTOARz3rppu9uxLLcLqVXec471KGSYbm4x2zVKKEEmNTk+pqwsHDc7qJRjfcepHcxBdo3c84FV2G0gPzzxVkFdwJDMB3zTXZGJkI5HQGnGTWgtNyRImKnafnPtVlUzFtbqKhSSQvvOAMcCpmZ3fKjjvWM22ytCG7QiMY6dzUtquUBPI60ksivHsUH5aWMlI8O2M0fYsHmWWbJGB1NV71ht6kYNSxYIBzxVadi0bA8gVnBe8HQr5ZrjcPlWp41DbsMCMelVDtmK7TjHarUKKqbd1bvRAOVFjXep59anik8yMkHNQtIix7cEgU6M/LuXH0rJ6rUaLC42liPpQVLc/pTEkyMd6kDZPTkVm0wHBdvbnFVwNxJJ4qwR83NQv/AKwY6UosZE7sQDVbYNxOOanOOVznFQlwPlI4/lW8V2JuTQyr1A49KeDliAeB1qCAYlL/AMB7UvmjziiDmiUVfQC3twAT60ucyH0FRu5IAA6deaepPGTz3rJpjJB8q4HWhiAo9aAcjPY00npisxigFBilJyFzSBtzYP40rOA3tQIQD5gR0FIQdxPtRnDkUYYggCmAJ80agdRQThQBzSIcKwyDQGUZHpTaBhjnmgHKY6Ypqj5jnNPUYGR1oYxGAY4B5qvLwxHIHrU/GCR1qKf/AFZbuKqO4ihuCEkHLe9IJixOe/WmRKSWY8DPQ1KQN2e3pXS0k7CJoZMSc9qs7uc96qRqU6/NVmOMdWrGdhj1TaSfypy5VSe/rQxLOB2ppy0pHas9WA8Z2EnqaYh+YGng5bGeKauVPP4UANfK5PUmnw8rymD6ZpBnIPXnvUgO3PHPpR5ABBAP97tUYAKkYyfrTucbv50jYQhc0kCHLkgjFRXTMiqUHSpORgA9KbOnmKfm5NOLs9Q3IGUqhIx81RyFUI3sOKQtsjCY3NURQyso6EHgVukJsucNgrUhG5eDg+1QxnnFTuuFTBrOWjsNjhzHUb4DDnrU2MIKhkIVVJ6moW4CIFRsDnNT5O0D0qCMbHy3Oam3E5AHWiQCsOdvrQvynBpGI25B5FIhJAJNK2gD+MmkbLDjmlBxu96Rcjj1pADcqOKTC9G7dKM7WyelITt69TTARBtBNAxtye5oPakZVJ56imAoYeYSe1IW5x60rDao44oU/NuAoAVc5C9qZG2OtPA6kjFNQAM2fwoABICpA602YcAjqaEOHOelPI3KST06U9mMZztG38TUcgL9O9SEHbk96jkz5BUDk96aAgwPlB4FWFGV5FVFQjaGOMVaDAocnirkIcSSeKUDnaPyoz8oHShc9c5JqBiMoPzDtSscDgYNAXaCSeaT7zg0CFbG3moI1wxBPNTM2ZCByKrbfLmb1PaqigHnBYBun8qlUYcqBxUJ+eVSBx2qWMFZSSc5pvYOpFuALJk4FCZZSOme9MaPDSMTg+tSFyUAX9Kpgxp24I6kVMpXyulRAfPnP5VIp3KOMVL2BDBICh44zUC72kJc8elTMuWGBgA/nUuxc4x1quawblcEBw6njuTVjHGT0qMxIzbQRgVIeQQOgqZagQh9kwG3g9DVjJ59T0qtIVDjPTtVkfKRjrjrRLa4ISQ/IeMkVCcs2GHFTn7pOeahdRyc8+lKIFURhrkj+EU8fMWVeg/ShVKxkjJI60pUCQPu47itWxLYkVskgYwKdkhcDvUAI3YH3T0qwg24yelRJWDoO5DcntTGIRjk8mhk3EN60SDdz6daSGOT5evJpzgMo9KaM+WPanSf6ril1GMUKsmB+dSAfMDUUODHnOSe9SqML160S3ERSJlsqOTUcqkOCDj196s9Fz+VQyMAAxPFOLAeOGwBRnepIP4U5R0b1phUpnApAxhzkLnin5+YYqsZGScKw+93qyACwPbvVSVgH9G20hJzwKZE6mY85pwJ31NrABPzjGKDwPemgddxpwxmgAjbjn0607AVAPWkABbPQUo9+lDAMYIUGhht47d6FBzk0jcDOevakBBMoydo5qJBh8AZNWRyMYyarKSSeOlaxegiQnccVXmZuh6DpUyk7uRVaUjecetXFajZAuRlsYp6TDzdxqyFUjrxVKSFxP8AJ/8AqrRNSEzViYOgOak/hHNU7XKOVY8mrr4U4zxXPNWdkCG4woI60oOM+9I2BHkUdQMfjUjFHApjZxgdaecHGelMY8ZAoQDFUeYSORUi52sTUceRIQKeCcg9u9UwVhwGcU3Oc4pScnrSAhQcde1IRGyjYQTjPpUf3QBjr0p8o+Ue/emtwi461aGSMOQW9KaSSRSrkoCfxpMnkjv0pBqBO77v40MenGQOtJnbgHgkc0mGWM5pgRuzKzEngdBUseMbvxxTJ1OxQDTkJKqM03qgRLz5h7+tQvjfn07VKxOQQagdgX2dc0ooB8YUA+/604EEDiqpdsKVGcVKp8xs9D3FU49QsOZwJBjHFIWLsQDz2pgUtLuHSnBduTmiyQiEIVby2Oc96Nm8lqerEsSRj1oZdhwOhq7gRS8KCOtKvDYHemktks3SnrgKOeab2AQsI5PfsKRRszkZ3dqJVG4EjJPanE5K8dKL6AhkYJX5D07Um0uVJGGFPClGYlgM+lJKMOqZx7076gLISHO4ja1KEWNV74pJFVnRD271IPmYEc461N9AWuo2Td5igDr1NKUPnbweD2pyuXcZ4x2pqK5mJJwKQEsrDbx2oRskc1G+Yz6iiJ+flHFK2gyf/lk2KbkFAacTxnHWmnAAHeoAQnAB9aD8oGBzSZ5wD+lOGDg9qYrgowBnqakH3sGgDPtQvKk+tSx7C+vpTW6bTwaMnO3vRJk4A60LcBFzik+U8E07OBt700qcZpiFz6Gmq2Xx2FA6UjfKMgcmnYB55OT0ppw79eRQTmM1GXCJuAyaEgJCVL0m7DdKQnjcO3SmykkLg4ppCHDIzmiXpz+FAPbNDnnpn0o6jKUnyjB7/pUMUe3p1NWbpdwOTjHeoEAZF52muiL90l7k7H5cr0quTlHye3SrAOPlPSoQoCsuevelEb1ZQZRsOBzUYAGARjmrCcORz9aZcxlXz0BrpT1sSiMNk0hPORz/AIUxxxkE49qeCMgdB/OqsABvmwRn15oHJPHFJ90nOPalHXPAoY29BVBEgb8q2UbKAY/GsRjhg3Vs1rofkB7+lY1VsCbLaMM9amQfMOeKrpx1qZTzmuR6MpkWoAm1JNY0RxIARW5fKDbHnkVgxlhIOc88+ldkNYmexsx54wO3SqcoZZTjjnmrcRzg5qtcgiQnt6Vkn7zuW3oQkYUnHNBx0znNK+SOfyqPAxkt+FWhPQViR0xu70uAT796eYwwOaRYz0/X1ouNpsQcHOOabKhEMhHTaT+lTeWGHoaScBbaTnOUP8qSeqHodl64FL6cUMSDxSg968QYoyB/WgDJNJk9KUYHT8qQB/FgilwM4A5ozk80hPbqaADdt60BjSAZGKTOOnIp2AcR6GlyDxTQQT70o6cikAo/Kge/akyMZHWl5wKAE4FKORRzjHagenegCN+tNyT0p8hC0xSDx71a2AcmDT8Eg5/OmIuG68VJg0mAhPYU1mwDnmndM5HPaoc8ncaFqMkDZXOPpUcnOD0BqQNkYHNBCggEdelGzEIqjlqVBtHHSiMgliB0pxGe1DYCYwacy/L1pCxzjtR2NSA3kHGOfWlIJHPBoyCKaCM89KYCjrjtT+hFIuBSkc0gDgHJH5UDknNLztPc0xSe3SgNBx68U0H160P8vTnPWnHJGetPoALuBPzUDqcUdgcUgwSc0AO4PQfWlAHpQSAvH40Z6YpBsIO/FAUchulLnA/pTdyjqOaYCkc59O1KQCQetJkcc0q4yKQAetBIAGKTO48GlJHIx9DSsAHO6gZOKCD1HWheOTxQAnfGaToadwOaRT8p45pgNx8x4pTjPPU0oHzZppHGe9MYoUls9qc4BbAPWmocinBOSc4pABwBgUHkZJwKU/3QKQggdqBETKN4I/hp3J4FHAc5600E7yCaYCYIY/LxSDc2SRilLEDg9aOTxjmmA9MIpHrSgZXcfwowvApSQO/HakwGnOaYgAyenpT8fORn60wMMHAxTQDxneaUqckt0NA44zQcgAGkxjCQSSelIWzg4oLBl2qetAVhjnAHWmIaV2ncO9S989qRiB7CgMAM4odxjFBYnjpSgNyaM7cDvTsAimxCKvy56+lKOQcYwOgpA4Iwp4pysGA9KTHuGcjk9KcNvU004LYI4NPboAOAKkQ1cbvSlTh8DNI2M47CnKcE/wAqBiybgfl5JpBkYBNOxzmk7jNF9LAKThScYFRFSXBHQcGnuT/9akIPc8+lJaARXEZIGCR9Kcm7bjv3NPcHtyBUOTnj8qtaqwupYU+mD70hITOT+VAYLGFpH+5x1qbajEDLwWPFNDZfGMj0puz5sE1LEoUlscmm7ICJPlkw/U0oYMxx0FPJHCk9e9VZUdH4YAHtVL3hFsnoevtUDjMhxxSw5GAxzT2+8QOppbMCNjjnnJp0bgR7mpJB8oxUYBKEenaq3QE6yeY2RQNrMSetCgIgA/SlwyDjAJqdLjHfdbJHJ6CoZRuOR2qYDoSevSo5jnoeKS3ERA4yMfjT1VicY69KhiO3Jzk+tE8zJt25LN0ArXlbYeYTBIW+bGfSq7KHTBGAaEtS8nmSnn0psxDPgE57AVoktk9Ra9RAg8wk/MQOKqbQspDHBPOKuQvtbbt57+9QSxh5Tx83pW0dHYCdZGKZSnqcR4I+dutQI7KFVehq0FLLvc4HvWclYfQqwFxJnGFHWibhs8Cp40cTljynao2jjkk3EYC9ad9bi6Dot7MQRTZFYTDJyvpUmcSccZ9KJCpbb3qU9RldmZpwByo9KiR0Rgq8n2qYL5XLcgDk1XdgmTjK9a2jZ6CehEoAcgn3Jq5HIcAKnA6HNVpFberj7rdvSr6qoh98dqqbSQLYsIwJwy5HbmklhJYEdPSiIExqcU+53syj7o9a49paFMbFzCSvHuacSBgkZ9adsxBhTx7UijbHucc0riK8q7CXCgk9BQtt+73FdrNUpZETMnXPGKa8zzMFQfKO5q03YWlysbYQt94ZNNGyI7pGwDVkxksGbqKgu4Y2j3M2AP1rSMr2TDYgZwZD83y9lpjP5kmCu1AKd5O5Dg7gKYoTBB5z6V0QiiXoOEo3ABcDsc1MLjaSx4UjvVYBdpz09PWlIMiHnIHQUuVALEWUeoHvRFKZJ2XGV/hOaWICPl+R2FOVkjfJG3j1qrIFZFpZGMRU8GoirLCR94+tSRAyOGU4X0qZogOfUVyyaiyvMz1XgYwTSqjeYHAzz60jqeV79qsRp5b/AFHStbgiRBuO4Yz6GozkRkY6HkUQLvdsHBp4LCMqeD2NZPcZNCoznsRUoGWz6dKiVWXae/epcbTx3rKfkHUUnJAA5pkg7t0qTjaeOfWmMu5KlaDKrqysSKgPXJPNWmY5wB0qGQAqDjpW8X3EwjJUgE9KVWVCXzx61EoXcVJyCOKRI8ckHAPrV2TEWhLvUdge9WFAP09agiVSpweT0qZDhc9xWErDFAO44HFNwMk4/DNOQlpT296Ch6g8VOwDFEitkEflUpAI5HJpgPzAD86lY5xgdKGwIyQASRyKJMhcg8mnMCWPpSbsuy9OKSGR5IUFV60vGDkdKUZRFAORml/5aYqgGpllXIIqQgBTTXOGAHFLznNS2A3gIOPwqORlMbbl49PWpQ2SVxz2qNgCrDtVIRl3kwYFUXGKhinZD83I9DUpiLBlHBFHkFVySCa7UlawvNk0UskjHamcj16VciV1ysh4zxWfDcNEpHb1q5FcCQ4zWM0+iC5Y3ZJAxgUxCQCT1PSn7hg479aj25ZUxWKKHgBVznmkfOFxxjrTzgJj070h/wBVk9+1LqISVsEAHP8ASpOhyRk+lRYCoXPXtTg24A55x0otoA7aDgNxmonx5mc5PanvgkKeopmdqDjn1pxGPY4kxtwD1NKwyOn401yNoyfm9acMkr3FIRnuD9o5HApZw4lQp0FWCFBZj+FNPKbh1FaqQdCRUCAHPJ7U9wAv0pqESIGx0pw+Y/N0Pas2McvKDmo5ANvJp6kklT0prZZueBSW4iJWEr4BzVkHCbR1NU0Ty8leT61aibcu4jBHaqn5AKwOR+oo79MUZG3PftSD7wJ4FQBKSNnI5ppGVUninNkjjpTW7VKAGxtxjpTSMAZ5NDklhig/eHNUgEz83FAySAep60EbRx3pcYXPemAPls+1J047UD7m4d6heZlkGFJHt2oUW9EBKrc98Gh/lk46Ug3Fhzk0si5JHrR1AQg7SR1PenA/lTAvPLcDrSkHBX8abGLu4J9KjU7yQ3SnBCcqelIF2yYH3fSmrAUmH78DaQPrVuMYGCOT0FMkUbwx6ipI+WBq5O6EkK/900A7eBxSgnI4BzQpGenIrMYPx+NNPEfv6048rk9KReRg9BQgEYc/KOtQPnz+FqctgA461ExIccdauIgxzkd6kQZ5brUKsemOPWpAcHLHr0oaDYhuOWyOvanqSUwBg0yYhWL54FORwQGz1q/sgCjb8uaeg2IWPWomB2sRx6VIhyoz1qWCHAZwe1K2Tj1pAxHNPOSATUgV2QeZj171MpBwQeBUU43jaDj3p0ChVwDwKt6oYhG6QE9BU4OX4HSoHAYHHWpkb5enJqXsIceOT3qJlwcgYJqUcqCaQdSam9gK2SGAzgfzpsqsPu8DvUojG/k/Smsfn9q0T1GQF8MigdOtWd3Udz3qmMi4Ynn3q0MGTJqpLQlDyQsYzTW+7txwetAYlsHoaFyZPl7VCRQowqhDStkkg9KAdykkcilPMZz1pANjOV+UdKkYnFRQgqtOIydoPNDWohxGDg9O1NfDYB4FDZLjBpJMEYHWhDHId454xQG+cg9BSKcYHrS4G0qOpoYiB490rMTnFSqcr0x/Wo5iysuOc5zUiNlFzwT2qneyAZGm2XGOTUgXk7j9KcBufmmnJfpxU3uA3O5cDOacFG0g9fWhsKVUHFIpH3e9MBy4MZ4peTgA01TjIFOwVB9TSAXhuKj5JK/pUgwG9aRRjJ6UloAzIxzUb/L82KkX7p54pkuc4A6VS3Agbhs1BIjecQvSrBByxHtTsDByOT1Nap2FYiD4jKj7wqHdg5J+apJMRFiPu+tQPnbvZvmq0gLdqD98jjtVsjJHHWo7Yjy+nWpVHzZPTtWE3djQxfmGOwp3UsMChOFbPekHKccVIB1HTgdaY2SuQMVIoyMdhUcnIFNbjEh4cn8qUcscDimbwJVGKcBywzgGqYhRgD3pG67c9OlOUHkk5xTcnLEDNIBjfN14xTGwxXb1HWpFHXcM+9NXkE96tOwIcckAU1yVwgHPrQhDpn07U0nc6kHigAlQs4OelOYF1GPxpzLlyBUasRMaFqATEKmfTrUUTnLKOvrTwA+4nPPalSNRLkDr1qlZKwkCK3lEnOaMBlB7j1qYYJI/IUxlLIam+oyJEYNnHHenRxAEmlTIGAe1LH8i9etU2AjDbggfrTW+UYB609sKhJyaQjKD60kBEzAAKKSTH389KWU4dQKa5IyT0FWhDJAMZz96moRjJzxTpArBTnHpTQBkjPFV0GLlZCjZ/OnIqqwQnLCm7FijAPPoakUkocjBoYh3HmFW6daiMgMXmMuSDxTw22YLnJ9PSmK3zlM8D2pJBqLH85LEfMKdAWKtkYpqyZjO3r2oDEFY2PJ70NASBGG4+h4p+Pk3Ac96ZIcZHYdaFY7VxyDUj0GyIfK9u9RD5WQIeD14q25G4JUDLiUccA9fWnF6CsWF+59KYcbuOtSSDCgCo8ZA45FQu4WQoBOeOaeewx9aQct9aeCBkmkwEbAIx0oPt07UuMqcULggccCkG4NgOCKM/Nk0bfm5pCPm/pQAnv60hPHSnuuMY/KkIKrjvTAZkZwKH4XJ70oHHHfpSkAjae9Metxh+ZeB8tRHjIx3qVs7doGKjeLcAx6rVInUcFLJhvzoYAIuecd6CDIAc4FLkNhD0FAWG9HLdc0jMVOSaUYXjPFQ3DKJVDcehppXYPzHyKChJ5zVYIpiUkVNO+1AB+lLsAVG7jtVJ2QxmwbqRlABwOtPJU8gc1G8mRgdqauwKQ/1+OlF7jywAaGb/SP60l3wpPb0roW6M0+hUXnoCO1P2EDC9TTEPzD/ADmpFJJ+9x6VqxsZxxxzT+QegpMYwB0zxxRuxnNINxrD5s9BWrbtugVsdOlZR9Ac1pWnzW4JNRU2Gi8owamAwQcfWo1I2+9Sg9K4pFMLrmEjuRXPFvm+ldHKuUNc84/eHd0zXTS21IZpWz7kUnrTLsZcgDj606DBQHvSXRUSj3HFRb3i1sV1UtnmnqoAxikjK789KlJwvH4VTYaCgce9AHGMZppJH0NOznnuKkBGxt6dKinP+jS+u0/ypzFg5zkCo58+RJx/Cf5VUVqgR2oPGMUZ9KX8KafavFQxwOT60cZ9qBjbjFCk5oCwvUmlHHbvSAtnPanMc+9IBM5470AYpTnbQDxii4De+fzpSBtpMYHFKSRweBQCGqvvTgOetCj5KMflQwDJzzSnjnvSDocUueKAGyAkZpgxjipGOVqHkNTQD0ILDmpM84qEdfSpgMHrRINQOQCM8VCy/NknipSM96jKMBknNOIDY9rHIPSpCMHPXNRRYwcHNTFcDg0PcBqnHIPXrT84NRqtPxnrz70gF6kZHFITgGgqccUvbHWkAxgGXBpBwADTj19qQpkcHvVXGOHBNKPvc9KaCOeeacOvX6VLAUHrg00DGcUqjjNJk5IPT19aBC9CKQEcjPNKM4zTSN2WxTAcMbOe/pTTw3Tj607gBQBSNxkUbMBzAE0relH8Oe9Gc84z6UgsGATQQDgCnDvxSY+YZNIOowIvbrT8ADOc0gGM0Y4z1pgCkDPFJn2o5xxQGByBzjrQA4gY649qUjJx1pByenHelHJpAMP3hQoIODnFI3ByOooUnv1p9AHY+Uc80jdBn8qX27ijkigBoBzjFOxgYz+FJk45oHHPc0MBxI4wKQ/L1/Clxhc5yaQngnPWgLjAAWxk5pG46DihjtOcc0NkYHrTAYSxUD+lO3OJCMcetI2W4DfhSkcnB/8ArVaAUA4yxpW+UZ7Uq427SMkU0ElRnpUN6ghSAqg4pqgFjxTs7+R+VC9/SgAAzinEAk5NNBx1HNO/hzSYEaqAwx0FLyQT2ppAB2jvS+irVMAxgcngU1mAQ8c05gBkDmkdQEOT+FCYxAMHkc07cd4B6Gohlk608HABptAJFGq9BUm0DGOnpUW4h8DBGeTUqngjOSaHcQHlwR0qQfdOaYyEAEdaegLVL2GMwpbn8qcpG4kCmkkSAgdKXOJMkfWgCQHc2TSEcjNHPJFIBlqkQufm9hSH7wApxOeKaQMgZzQMGwpwDUYH3hswD3qUZPOMAVC24ShgcL3FVERIq4XDHIFOJwoI60mM9+f50Ng8dKTd2MiLHeMnqelO3HGFHFNkVS4IBJNSAcgY+tN2AY+Ny88r6UhXLFsZPanOAWAAxTnJ4FNvsBEARIMD5cdaceBuPNGdpxx70nJ+nehiIAf3nPSnx9W4yPWmsv7xueKkUfLmrb0GSLkLj86UjpzSfMiHAwaUZ4zxnvWYASM4zn3pkybkPzU8LhuW6d6aBuBxQhFVMITx+NT7NxGRj3qMrtx6VICNvJ61q3dDGtjqOPrVTIVycdPWrTcjgfNUBUg7n7U4CKyTZuuRgdj605k/f7tvXpTJZMzLtXk9KR9/moWPTuK3SQiS2wpIYZYdqmb96rAnA96rxriV2dh7YqZk+UEnG7tSlo7hfQk2EIoRulMeMQoWHJJ/KhVVoMMcGnSRnyshufT1qNtAIwV84bVyQOaRoMyK2/DN2xTUKow2/e70tw5iiLnr9KrW9kHQZNIhUqRnHSo8hEIVOT0BodgkJmVefSmKyNKpYEMRWiWgX6EVzI7IpHyp6VJBdPGoTIP4VDcIwmKs/wAnYUkLZkVccHpWqimhK5r2zOZeu5aknBYhiMnPFVozJGTs5B6YqYOXfDDHpXJJa3RRPG2xCnfviocnYRnFSxqQpLHOKgP+rLY5qUr3AftQLnA9iam3bUGBUGFWIHHFOLFgARkU2rgRzkscKeD19qheNHBQnoKkdScqgxjnNROAFJP41tFaaElaSMklUOF7c9aa+xSkYyX/AIvapYoysu9uB2zURdYmJzuYnrW0X0ExPLImBI+UcVNHtTJxyTxUTb5AHlfA+lAKsQyHnuabu0JaMsOqKBu5PWqMil7gMR8ueKumPOCQfrTLmLZFvUdPWog7FMuxBTb5A4qZSAme9U7aYm3X1zVlTiM7xzXPUWrKRUchS248H9KVNxIDHK4pkwcT5A/+vT4+Ygxbn0rR/CA4Bo3KqvynvmnsWBBVeKaUkMmewqV1YxjBA9KhgO3M5XFStz05qujsE461ZA+XdjispKwMFOflPQUZHI/KgAEZprDIAxUgQNksTUMi7ozVlx8uD3qrICFx3zWsBDIwd4VeRTpnCxnHaolyr81OcODkZFaPe49bCWxkcZcYHarCEshz0pNwVANvHrTUbPAPHUmoeuoEwLHkHHrQ7sFOOcUqYIGaVvQ1l1AZA5JAIx61YY84FVlRUlLjoetWF5y3aiVm9ABulNYEy8L0pTgngcUgOGJPWpQWEUnLdyKTkHJHJpQOuODQ4bA557UxjQpGSxpeWXpgihwcrk4AozknFMQq/ePtTGPykHjNPIO48YHrTWPycDJFAGdIjI+MdaQn5gDwO9SXMTP85+9UOCevJNdkLONxbD9iM23PB70ptmVTtoYBAh6irJJIGOh9KiTe4CQlhHz1qYfdLnrQMB8GhSWfkYArFu+owI3D2ox8rbuKdgBTzTWG2HJ79KlANZ1cYpBjg4wccUseHwDx7U7bg+3aqemgdRoXD5HJ9aCMDkfjTzGSevAobDMO4pXAQgNtPelZsHAHHtSEjJFRIzGRg2ODxQlcZIwAOzHFRjHlkYzzUsgy3NNJCEcdaaEOjUY2kcUincp45pFBSUnP3u1PzgEAYpAIeFwetMPIIxknvTnX5hzR92QkChB1IWIQCM96nCgIuOcVAP8AbOT61MhxHjHP8qcgF+8pPf0oDAMAep6Uwn5gQOfWn4DFSTSsBN1OKQcnPpThjPtimrwG457VmAwH5s54FKflGepprEDg048ID1qgGycYJPFA+6M96STleuM0ZAXHWn0GO/gwaaR8v1peNme9JxwT0oEKjbVz6d6Zu+cYp/6CmIFBIFGgCIrKCWbJ9aFOGyelP45XOBQOFK07jHLgluaaD8+fwpOVHNEnQYHSiwiN1w2O1JGdxA9Kc+R8xpo+UgAYB71S2AeeCPalHABNJIT+FNzx/SlYY9sFR60wgYGOfWnHCihB8hoB6iEncMDApknXJGaepbPtTD1x6U1uIiX5mx0qQYYdDxQuCcdhSjlCMd6bYEF2BtGw4Pp606PKJhjlvWnXCAIh2/dqGKTeCxq94gwlYlAQe/WpY843GoPnPGMZ61YjBz14Wh6IEOPzsB0NKSQT6elGQrE9AaQ/KBUANnj3oOcCiFQEI3ZIokG9cA8GkgRVUKp+pql8IdSOcsqZQYzU8LFlHXPc0SDa3TPtSqTswetJvQES4zjjpSZ4OOKBwME0i5Yn0Hesxidf61C4Gzp9KnBwT3A60xgeB19KtCKjoQCQ2DSxPgDPPvUk0e4Fe9U2BhmBJwO/vW0feVg2LzYK7uhHQ0rEnGOCepqGNt68HrUwJVcdzWbVgHFsAjuKUn5PekUcHPWmkkOPQ+1SMSE/Mc/iKfuwCe4qOP7/AADk0rHLHpim1qIfET0bv0NK3I4qOMky9cg1J0U8UnuAi43Z7CncDn1qKIsPkK/jT5M7Bt7dqGtbADHeCQOlMUttA70g+SMYNL8288ZHpVCTbJhn73pScAFsUpyBgUxk3JjtUIY3GRkUqkk5/WlOfKAzjFNT1zVALwrE96cM4IpMZ57ClHUknrSAUdcd/WhhkH1pFYY4HNKMYY96QDP4sDmhxl8joKRQASTSghmPpVAQy/6z0ppZSnI5qRxl+aiUjYcdqtbAQyr+6xnv3qGRB5gyeBVh1DZDH/61V2j3MAG5HQ1tFiZfhbOFByKsHK8d6o24K47+9XAS7ZPIrCa1HqIOF69aOijmgAkkd6OG5PQVIDgOMnjNNO0KS1KCeBikY5XpSGVmcBlGMn2qUNwMr3qKQbHGO55qRvl5zxWrtYWvUkUAJhu9NHQgilzxQ3AFQAknCLt6+lMPLFcU9+x/KmvhXB9aaAYOG+lL/q3UetKQFfPXPelcDdj0p3Ew25lbJqKJwGKGnHPHHPehdpfg896Y3cUjtg5puMYGeacTmXrxSBgJcd+1Ah+ctkdqVjlgO1Nb5csOTTXYrEePm7Cla47jDkKQPWlChpAAenWkU7I2Yjk02IMg3Hk9verFsTn5+9NH+q6cCnISwYcimLkxFT2qUGpG4O9QO1NXknPA9afkj5yelRjkN7VaAa2GB44FNjGBtNPxuiLZ/CkQjbjp61XQLCsBgUhPzbQeDQ7GNCV5J7UqEKvTB7UdAuJtIuAexFQGQI8gIJ7GpnDGMAdaSVUCYYZL9aafcPQRZPLiBcYp64d/OJ4xRJscKvfHFKx8tAgXrQGrAsS2wjAPQ04blXapwR3px/1YYdaVW+UmobEOjyw56ikJBcZ49aSIAs0gOaXO5iRS6lIkRty81HngkDrSxOTnNAz09aWzEKo+XPTFKo3D+dICBxSjhtvrSAevTHrRjaaFO6TA4xR0yAOaQw6kikA3HJPTvS45PNIOFODQIDy2BSMQTzQOT1pNxMnA6UAKSS2MUhOW+9+lK3BPrTACG600ASED61G/zxkA81I2d2aTACnmmtAGxjanzHNKylfbNRk7VBAqQjIUsfwqmLcRhzmmuiyfhUj9R/SmkZGD0pJjvqVmfahO3pSQy+crNinykKQmM+tCgRttQYrXoK+pFkpJjHXpTVUJkZzntUrAqxOaYqkruLfQVSYFKUBZh9ajuyGRQvT1qe5UnDHrUE/EAA4PcVvHoStysisX9ccVIq9cmmxk7856d6k3YxjitHcJAevApM9QBS5IpoXHTrSBbCOBtxzV+w4g+YYFUsjkEVdsCfLOOamp8I7mnGuAD61IcbRjv3qOIkxipTjGO1cT3GP/AITnpXPXHy3BA6ZroR0H+NYF9tNywwRzXRS2J6lm0YsCD0FS3Qzg47VHZ8IAeamnA2KSamXxlR2KseQOetSqM/Sk28D3pVGOAfoaGMDnnJ4pchFOPxpD39KbJgrgjNIHsPDbuaim4tZBj+A06MEA80k2Ps0vrsP8qpaSQHZ55pDkcUh4OQKfjPX8q8UYg5FAxxmjbxxS59KLgLnBHFDYFB65obnmkIN3HJoHoKCAaMY6UDDoCaOD1oI55zQRgjBoAXmjoKTuKUj86ADjrRkg+tIQDTgccUAIcYqLHepN23qKZkDmmgGsoA47UoGQMGlyGzTlUY69KfqAY54NLtANGMc4zSgZ5pAR7VU4Ax7U7tkc57UjHDccmlGScE80ARbgHxjg1KQNvtTGBDgdqeAPXkU2wQZyOvNGRj3FKCAcUc49TSuAgwRk0DGarvMIm2+tTJ2aqcbK4CkDdTwBnpTf4gRzSqaljFwMn3qNyOgPIqTGD9arzgKQcdetEdWIfCzMGJqQArkikRcKCOh6UZI5I4oe4xUAzyaSQ4K80feOTSnrkil1AX27elKCM/WkTOOTQOu0/nikIcCM0hIJIoBpQMZJFAB9OlJjP0oH3R604egoATktUaxhXZwOTUjA7hRg00wDtmgEBeOtAHIz0owM5PSkHUYR8noaFB2kdae2NpPc03HGM0wDtkdfSl/hxmkC9eaAe+PpSGGf4cfjQSAPWjJHQU5fUimIQ+9I4yQAcUvUE0hxuyfu9qSGIdu4DqaViDjjGKjfgDjmlXJHNVYWojMqAYHJp+VCHC/MR1pnDHA6DqaavIz6dq0jLl2Fa5MGUDIHNNYH0pm8F+tPY87jWbVhiYC4OeaUDApAD1oKnbknkU2A4E8nHA70oJ7Cmqflx61IoIXrgVLAjYAMOcmkyAxIGQKdtwM0wqdnB5NNABC7gRTiMnJ6U042jnrSg44GcUDG4AJxx6UyJTsweual4IGOKCCqn1qriIZTtwuOT1NLCPnDYpJSCOe9LGwUDHAqtbDJ25yOtEBBU89KdhjyTxTQOTxhelZ20sINwXLA0u4Hqfxo2qVIYHGaTauwnBJoGPGdtA+Un1pAD349qU/6w+hqQA8nPakyN4wMmlJ+XnvRyAMfnQgDfg80n3OBzTSSFJUAn3pGkAwCRmqt2DQeD81BOW6UmARjPJpQc/IMGlYBrYGPWkLDdgdqXAJ2k8CkOOTjNMQ5jhMgUhYkgdzSHLKR6dqY0bOAx4I6U0kMceHZiKUEtximhlBIJ6/rTlOWPHHrSYERX5z+lPydgxTf4sLSRhhhTVdAJgN2STxSt8zZ7CmjO/incDJ71IDR6k4pRkodvWmfM2GPAFOALNwcAdaLAQsABz0FJuwTnp6050yT7VBLn0wB1NaRVxE5bOMH6VDLkggc1KuTENpzTNpyQDQtGMzxj7TgKc9M+tPKbGBznHU+lRyRBZ/mbBBzTsMZHJOF711ppokZDteY+h6Cr7IMA9SB0qkrKs2YeQavqCzBwcLioqAiPYHZCB81SuF27iMkd6bGHLNuPHvUcjBYuu7FZavQCmMpISrZOasNh1KyDrVKEyxuW7d60olWRee/Nb1NLMEitKGaMoRj39KhVguDsyRwD61daF5o8E8n9KpjdB8mOR/Ee1EGmrA97kd4F2n5fmHQ0QJuUcYIp14XYhWw2f0qKMGMfWtobWA2owAmM5I70jEBSM5xUNq+6MZPI64oYld7FfTFcsoNSsymTLuL5J496QcZB/E0yNmdVLce9WEiKBmJyT0B7VMly7h1K1wjuuF+6aT5cj5s461LMD5ezOCetU4DDFIVZ+e9ON2iXa5ZW6iYE549KjuNixgDPJzQLZZAxXqx5p4TbEwJ57E01yp3QdCk8YeQOx4qrsMkm8fcz3q95SojF2wveqjM08oBG1R0X0rri77E7BKAw44VeOtRxuikBRkd6dlvu9EB6UxYz83pVdBdTTRwVUAZJ7VDdneqrt4A5pIcsqlWxU0qsuTuyGrB2jIsitJNgwoBJ71fRjtIPas22QrIDzitBSNxQHp1qKyu7gV3O2fGe3OaYiB1yOCO1STgiTeBnjpUULlgQrAHPTFOOwE5BMgUMalwpUqD8w6VGpZlJz81SKufmPWsnoMdGhX5c5zU+MAjHTtUYOMsPyqQAkBqyb1GIV2mmthjkdutLjdkk0fwnihAQTZZhzxVaRuBnr6VakyFDYqrM2VJI4raBL2Ifly2484p4BMYU/pVcA7v1qxuZWH6Vs10QIlU4ULjpRGP3544poYnNKAwKleKzd9hlrdxgfhQVbbnPXrTedu6nBR5fORWICNgJg1KgUoOetRKoKHaeexNOgZwvzDkUNaASEfNxTRgDJ4FL0Y8UHBXB61ICFsHYD1pwBJwenrSAKzKcZxTg2XJoAhIdpOvSpOnApqr5bdetA+9jpTbuAM2wHnrS5HlnFGMtyOBTcjeQKe6AhuEJQY/SqTZAU45/lV5g2SWqHajZAPJram7aMGM+WUACpYFMbYJyB7VCOMqOKmRGUKGPPrVStYUbMmwVcnPHrQpwSTzijGeAcYpqks5BrHcexIuDlR3pJcNFgHpTkYByAKQqRnPep6gQLJ8wIqyTk+1VV2+dj3zVnOcZqp2AcnBNRMG3kjGKlbO3gUyXKqOOtQtwERSFyTyKFQb9zc5pyr8oU8Zp4XjA6DvTbsMjIyG9aY4byxjrUv070wsAh5oQhowApPWpeMmokwQd3SlZx5e7HApvVgNfJIOMYoYFWUqcr3pdwaMNjFIRtUEVSAhKsbgtk7e1TxMG65pmAMjuakg2gEd6begxmQjdPlPenMcSDA+lI6/KMetK+VTPepETg4wO9JkBvYUiHOD3xThyp5+tZ2sDGvy3OMVEG2sQTxnrUjgEAntTQMsGUdetUtgBsFyD09KYOuKRiQ+APxoUktj1qgHgYfpnNPJ+Yj06U0Ha4XGfehhz1qQHMTjHpTWACblpAfmwDT35Tb6UDI8gPn1oLZk46mlzwO+KQDHOOe1MBWO4/SmnkBj3oBxkHgmg9MY6UwGyjJGe9Io8wgdhSy53AjnimRn9420dOtUthdSaRcnjg96iB7D8akJ5yOc1XAAfOaSAsOcvSKeDxxSZ+bOeaRTxj9aLaAOOcZHamS8DPennOBikJU7h6UIBkTKQNtOXAY57VEhy/PSpQ2V6dabAjc+YxG6o4izM2RjFStgDgc0beCy9T1qk9A1Gn5U3YwaI8Kxbu3pTGYkkMakQjacc4o2QDn6qRyaR8Sps7inDkZJo3cEnj2qQGsVWPg4Aptso4A6d6V1BjP8qjhLs2R8oz0q7XTDqTS7uvAPY0kOQnz8mnSDNMt2zhT2qfsgT5y1Lj5cUm4ZNKoIJPf0rMBoBGV70vBHHahOGZjQBtBIpjIcY7nNZ05aSYgdK0ZAwwe9VWIUbsYPpW9N9REBaSMgY4zV5Tu2nFMPzx5A6inwuBHlu1EwRKxzkDqKa2BEM8mlGd2SaTO0HiskO4sY4OetMcYBQ/xVKnAJJ5qKZtqqcZPpTW4noMhVoRg8571MeMDNN6tT8gNkd6G7jG7gH+nWnEfJ6UwJmbcRUuQxznik/ICuEIxkZHapFXD57U8c4GaRhiQkdMU+a5PKODcfWgYA5pB8o+tDD5QScVJRE2dnzdO1Oj4OD0pzj5ABTBkMpzxVXuIeT8xAoHCn1pR/EaTPv9KQAvBApQB+JpHGBkdacRikFhh4IpqcuwH5U9gR82OaZnaSe/emtgGzZYYHWo8AAY/GpHYCPPpUe8CPNWtgKzsyNyOvc0j43Db371K0e9BlgSTwaR1wCPStroT2HQjMeOo9asw42YzzUER2Rkd+1SwZJGeKymMkA5xnk0mMMR696c33x7UmMg1ACjjPNJwA2TyKF4HP4UHJz70DIioKlienal+8oPWlbhc/pTQcMDVCZIQA2Mc0g5elxgk0z7o+tJAKeFPOaawyQaV+AOOtI2AAc8+lNIBrAK3PWnEfMxHIprYyDRgsH54xxTAT7/zZpAAwDdPWlTKIB69aU/KoGOtPYBrjKgjqTTGOx/c/pU0igAHt6VXk6Fu/anHUCZiA3P4Chj+dRqfMTJyCKkIw3J5pWsA0je209ulGPmAx0pwyC5602M4fPYCmBISQAAKbkqvtTozuOT6cVHtypA6Z/KkBDMWztHeg/Ku0c1IydGbr61HuAJ56+1WmIVxmIEGoQCRuJ61MQQgANQgFAozxVIdiUY4AGQaZlg7Zp4wFwvB7ZoCs0YBHJpBcYOUUE4P86VwWiBYcikn27VVjtP0okUsiKp+f+dMQjbQoOeMcUgdig3DOTSttZfLPLr1FBdlwmMCmG49wxKEfdFKWLPj+EUEB3x1p6fMDzwKjoHUk+4MY4pWAH1NA+cZP5U3krlhioKECAAgcUNyAwpQ3JBH1oIUrTJ6CZzyaXGTu7dqG4Tjt1NNBAAzzQBPGMjd+VITzn060qcJjsaTJORio6jFQkfQ01ecrjrQGIHtQBnJ7UxCINpZRTSeDjipFHBzURyR9KaBDxyvPemkd6VeFJ7cUADOc8GgYdsHrTOi+lPOQc0BAV4oAgAJj/lT8E8E03bhgAeBTygLE1ZIrcPxUcgAAyelS/wAZPb1qBvlQqeSe9KI7EcrkAMBz70EhVDnqaRgpIB/hqWMb4yDxWl7IXkV7gAlTjJqNt64wOKkmK5PPSoWzLHhWJ96uOwWTG3Cnygc4qjc8uvX8KvzIywlScmqM7ghRjtW9MV9dBtumc1KcYIBzjpxUEbADjNShhk/zq5biY0kk9KMAH19qOh69qU4JzzmgOwmPerOnnhhnvVfOBU+n53OPfpSl8LGtTYg5THcVIOGHNMgTcpFPxtbHpXDJdhiscKDise+4uG45NbGORn8KydR5uARg4FbUSR1o2QFxVmUZjHHTpVK0f94DWj/yyIGDTqaMtIrKfl5o4wMc0wcf1p24bR0x60mguIPUcClPzfjR3PGaFwKAALjtx2pk+fs0vH8B5qUdeajuObeXsQp/lTjug9Dsgueh60YJFB5Xg0ucYrxQE6UYyOnWilAGMj8aAHYwMCkHH0pqn0/Kn9TzSYwB4xik25p2MdfyoyMAd6XoA1s9aCDkGnYyKacg8mmGwpB6UZz9RQBxSgDHBoAT2PWl49cUgBGaMYBpAGMg+lQ7cMc8VN2NMKkjNUmA0HA5GP605MnjimnsMZxSp936VTAeTgUvaj5eBSkHoKgCCZM4p4IHXrQ7bSMmkyrDdVLYBcg801HD8DrTJOE3LSW5yCe9O2lwJgPmpxJJpFO7pTlHNSxkbRqxBYZx0pFDB+uRUmORSDIb29adxCng8Uq4Bz39KD1zS8enFSMaclun40jkBc45pwPAz0psi5XA60LcQK+ScUEnbgcmorYnDZPPc1LjOabVmG4R5GSRS5DLkd6Rd2Oo5pRjHAoDoGTtxnmlXqRnpSBcHGKVOAQR2pAKGyaUHOaQjPbmnDkcH8KQCDgc0Ln/AOvSkHFIQMcGhMAbls0vHOOaTbheDmkwcZ79qA0FU8884oOOmKaOR70/GAfWhgN5JxTWHzD2pycpkUm0k+lNbgGecE035yRt/Kn4wcUpGMmi4CbiCc8n+VNXKgg04HjgUi4yO9AxcjHrSEZpVIOSelKBuI9qBDTyuMUnOD60ue4pexHrQBEcAhfWkPUYX8adtJ6UnzbSo/OqQAFAkHenk4GMUxFKk05s7sDrS6gIdx4Izz1pW3HIxxSEkORS9CTk80wBV2qMfn6U8DC4znFJxge9AHP40rjBjxycUw5wCKe4znA/CkPQAUkA3aCoJpMjcABQ33cA00SKhAP3quwaDyuwgL070v8Avd6UYOSevakQFgcGpAYwGQAaiBCvtIqZ+O3ejC46VSdgBnYkYXI96VCWx81P+91pAvak2gBj1BHToc0xXJHoTxT8buAPxoxjP86EAq5A9+9L0OT1pVG0bj+NKDlOvWob1AaMZxRnI5pwAAyRSYyTgfWgNhrN8uBz61BKgJDj7wqdVUFsA8+tNYfPkDI71adg2BUbZuc9egp4BXBz1o3buaVwDyTj2qW9QIj8rEnqelCnnApZByT600EAAk4PpT3QD2O0YHpTV5XnoO1KoB69BQBuTimAwop2+x4p+RkgH8aRQM8nilxls8YpNgRt8inB5NGMHJP0oYgsRj8TTck4PeqWwakiDgknB9adjI65FBHIHtSqNvykVICBSfmJwD0pcfOAOaQ85PYU9TjPPNAiBztLEnpVPc4ccZBNXJRknI61Sn3gqVHydq2pK4FpQVOO2KjZfmzyB608tgUxycZHXFJXuPYpXZaPomB65pI12Bc/MzdaWdZJCDnA9MU2FgQFPWuqHwE9RyBd4KrgHmriNluO1UsFbjltoFXVxGMg9epqKnQYSB3k+UcLVedi67E5PepmctKc5xVaIl5XOOBUQX4C6jWAV4w/bvVtGXBOdxFVnRnnDE5FTMVQ5QfUVUtUPUe04XaO57VWOGnO5s4okQl/O7ikIVXkfPp2pxikKw+OSOR3GM471E+1iFHahwgtXeI8mqsU+0fMf0rWMbiv0NK2j2uWX7tPbJXGec5qK3mLjGeO9XAgZs9ARWFTSWpVkyOMb/mz0q0W+QkioRthQrnPPWlL8kelZSV9eg7kRIZ+aqPBGSQAQSetRm6DXBAbp2qxGPMyQMAVuoygrkpponVSEC+lLMFCcnpQjNsBC5x3zSrgnJ5rF7jKtynA7YqAxrFJkN8x7YqaZiFLMOc9KpmTILkdO1dVJNoh6akrxsVEYHzDkmo3iK27MO+OKk+0ogyw5xxUKXDMxy341olLoGg5Dj5AfrU5kb/VnAA7+lMhSIDG7Ldc1LcRqyfuzwOtZys5ajCFwZMOuQO9WkX5txHHb3qtDtwBtIPqasJuaTAPHesqm+hRDcHZksaihyGBwKs3KFlK9SargCNgOtELOID8lm3dPapyTGNuPxqBACSQelS7nKcjNKQdBUZmYgdPWrEfAK9hUK4R8etTopCkk81nJ6DDOGHFNY5XNP7ketMA28HpUICOUZXrx2qlICV4q5LkDHaqTZIxW9Mkgx830p4wSu3tTYwWJHf1py7kBGee9bMEOwcnOcVMiszKSah8w7iuOPWpgzEjb0FQ7gibP8O7gVIpHl8dahj+8cipBggDPFZSQw3cbh1qZOVAPU9agBAbA4J6U5ZMNjOWNJoESk4OOoHWmZJyT0NSZHIqPnOB0qUA9iVwVHNB+VCR1NNPXJ/Kn4yeD07UgI2HfNHO8kUEjqelCn+KqGO/i9qiDBZtoHPrTzkSZpuQWHH40IRCoLyfNnHelKhVLA/pQ7HzQf4aHKq39K0uwKyuqS8jJB6VaPQNnmqik/agTVk8gAetXIFsSOm1C2cZphfaqgAfN3qVk3xAt6dPWq8mVAYD5VqIq71DQnU4Iz1qbgOTiqkZEihzxjrVkncuV/ComtRopXDeVKHUZJq0H6etQ3UYYg469KkX5Cq4H1qnZxQtiYkkAjvSSg7Qc0pGeRn6UrqCQKyQDAPk607PyAioYW3TOvYVMpyD7U2rDWoEc57VGTg88g1Ifu+lQt1UdhTiIijO9mXp61KQMeWR+NNCqJMr2PNSBlPK9PWrbAaEKqUA4FOABXngUo+/imcjAzmpuxjHALhu9EbMkp5yDSsCX6cU0H9/gcmrWwiVhjr0psmSBgcU5huAz2PShTkAVC0AVBuGRTzkKFH402M/P/jTzw2al7gIw2pnrTIjkN2qRuW46VF9wZznNC2ACmAaarYXAHTqaeAQm88HtURznPY1SAk3YBc9qGOcN/Dilj5TpQ4wdv8AkUdQAAEFu9Ob1PUiolxwoqY9fWkxke3BPp3pGJDjH4UoPytk8mokALAHOVqkhEj8sBmkZjkE0wkbxnr608g5Ge1FrANYfLuqOJgwz3oZ/nKbuR2FHyqwAq0tAJ0C4OKg+bcwA4HvVrAAGKjIwDnqelQmA1Cducc1GSQuQMjPWpEBCnsaiyUypP41S3Gx+ePUUpAKZxz25pAcE+hpUO9SMACgBmRj0FOU56Dn0pARtI70ikq4HY0CFdty4B5qMFtgolYL060iglVJ+7VJaANaIswanIArkd+9S43N6ColH73PWnzASkkICaGI2g4600Hfx608BSnynINSBBNMq8IMmlSbJx0YUxLZUcsTk015c3GxVzitLJ6IC4x/dk96gtlZCzE5NTDJUE9Kjj4dj1NZrZoCwBkZ6UoOFPqe9MGcAk1JjAwPxrNgxh7e9GcqewFL04JqIuS20AkU0rjHt8xAH41RuSy7iB9KunhcDvUMyblwOtXB2YnewyEYQKeopsecsppTu3gj05psnEoIOK03AmV22NkY9KcCcDNN3EsCOnpTj90AnHrUMZIF/eZPTFMdgVzS8rFnrSFdwGOKlCDJJApwAwCRyKQjaMdaXkDGKBjUy0jE96XgJkdqR2IJCjjvTio2hV49aYhWPGR3owMZNLjkr2FBOWx2qQBenJ701s9QKM8GkdiGA7Yo6gOcKSKYuSBjvSv0yTTS2AG6KOppoOpIQNhpuMcDrS5+Xg9e9IOuO9AAOhJ604Y25zzTdvzBaeAM47UmA1+cDrUJYAliKlkPJxUE2ACauIDjh1PPWmHDKwPahSWA7CnZGSmee9VsBCpyrMO2MUhBKYY80jFY5NpztqXJ+8B16VQkMhYeYcrjFS227zSCcgdBUKjLsM1LbD5io9KJbMaLDdcjoKaxP3hTj93H503gjH5VkgHE8gAUMM/1pEPX1pWH50gI2AxzScEACntzkd6ZzgcdKpDHkjA9TTXBJxTiKGPAOeKQhrk7RjtTX9+tPbio2BY/SqQCEcBc0P8AIij1pWBJwB0/Wg/MNuOlMNwJygYDikyzDk9KTJMZ2njtSHhRzz6UWBjgdytk0x+SqgdO9OKgJjqTUeSyBf4jVIBRy+B90U4H7zHoKIhhST37UqgiNieaTENR87sDvQzERtgd6b/yy4NBHRCOO9MfQkTgA7uo6UhGzoSM9aaAC2R0p5JJHoKQaDD8sZyM5pjjCjI5qY4PGODUYdmBIHOKaYDGVsgdark5kYE8VYIOwknk1CVCbSeSe9aRFYVgXbH60oDljk5A6GhcsCwpTxECx5ov0Cw19rqCT909aHl+aM7OPrTgVMWxk5btSIxaI4TGOmaYWBsm5IA4pru0u5VPI9qBK6xbmBzmmkLG+4nhqEhlhHw20DpToxhTxnNRQIwBBP0p6RttIJ61DSAmYZICmlOCuD1qOPCDZnmnEYyvXNRYBUcbMHmmkYGMUuBs4pqsx4I57UwE2nHNOUD04FIeW680oIwQaBaEykFOO1Mx8vuaIjjp0pWzsbFTazAao2rilIK4B6GkDDgnoKfzv9KYDSSJAMcCmnJQk9DT1z1NMCk5GaBi4+UDOKGzt4NG7KkDrQeUxmgQ1iQB6045A+vak4PXqKOPTJpgN2AEk8e3rSEEgEcA9qcSdvvTGB+UDqetNA9RxxtCg9OtRSE/dHHvUrEIvTk1AWONx4NOIMacZCk1KoAHXimbd7gnpTyRjpmqYalaSMMGx19arQq+TtPAxVto9nC8A+tMwsanavB960UtLE22IZVMgyeVrOnztBrSkmBUj/JrJfJbOc47VvSQaXuLGcE4P4HvUqjGPWolC5PA5qZAeB3FaSFJjT95c/iKOSODwf0p7KByTjvimquM4OaVwvcQ8Hk496nsmzMahfpwOakszifHSiXwj3Nu3kPPPUVJkk1BCwyfSp161xTe6AcQeBisu/TBDGtdRnk1maiF2jOauiJ2K1uMSD2/StIZ8thWXan94MCtVQCGGecU6pUGUzkk4pTgrjH1pgG1iM8U5QcGgYig9BTwMnmk7YpTxxnrQAKWyd3Aplwc28v+4f5VISR0H41DMf3EvP8AAf5UR1aHc7P+Ling460nBPSkxXjAHXt1o6DApT6Ug9xQA4Yx0/CgdD600HnpxT+PSkADPel78im9ByKAcj2osApwelIOnFLn0poJ9KEA7j0oIwOOQaDSgDGc80gDj1pNwxzSYBB9aPSmAvy01j3pw9+aGAIwOKEBC3HNKmST6U0rkn68UIMd/wAa00sBOvfNBJBpgOXIp+MDkmswGlt6gjn0pFO9fmAFIGVTsx1oRUHGarZAIQPuk1FEpjmIHT0qdgOuc1GpJkI2496aegdSTNSAcdetQocL6mpF5GBx9aloBccE1EXIIAGRUx5pp5HTihMAPJHOPelzSYGcUighzjpQMdkdTQQAB1oIzxjFOPT1qRES7dxGOaeAOcimhMPnjkc05sHrVMCMHHUU4An2oyV7UqdDnOKLgLzuz3pTuI44pOv+NKeT7UgG5IYDPXrSgleCevekdSRwcUAHijQBwHXNLj9aQAlvalzgZHWkAYI6EZIpOwFOIpox60AKByMUMxK9+aM4OAKM5IHQUDE25A9KQ89+KcckVGeDyeDTEPGOSKbuLMQV/GnZG3gCnZ2g8UgG54xQF5/WlH3QcUDkAimABsjGKT+HJp2MDGRQRk4HTvSuA0dh270gJJ6ZpxwB059aQ9cUXAjPXAqNtyYAGTUpHfuaY27fg9a0QDwCcHpQjfMQxyaQfMxINO2hckdfWpYAxAOQKYGyfans2BzTFY+aSRxQloCHkjp3FO6DH60xeQSOtPAz1H40mAjY2+9McjbUhAI6dKj256857UIBjrhQx7VEoBIcjpU5UOOaijBLHn8K0ixixkvJkn8KlBG3CnioGUpyvT1p8XCFuuelEknqgHFhvK02QFxgVIVAGehNMYdCe3akmgH4KoTn8aGYpyeR6CnICw5HFOYApmpvqAhO2PINNB4z3ofOwY+6aANp+tC0AdnIAApzcAY60int0p3Q89KlgHTiggAEUfx5ob0xSAF4FR9jgYp7EDoKicEIT0qogOjJJGRStkjPemQtlAwPWpASfm702tQGsuBjrUOPMJ7be9TPjAHemnA6UJiHKOMHpQpyDnp6UkeDGcmkVssVAoAApK+lKBjHNHfGc007iAfegYP945/D3piDa3I59Key8jnn0pCPmx1qlsA8Esc9MUHABx3owfu+tKcb6lvUQLwPc0o4YU05zlRSB2KDCY9s0DEckkknFQ/LjJG7HanyLuPB5FVpnwAACWrSCvsLUnIyM4qJjhS6de1TZzGrE/hSSruh4HJo2dgKFx8y5zgdyO1LE6ttwBgd6bcn5AAuSPeiOTlVHIP6V1L4Q6jWOZ2Zjx2FWAC6BiOPrUM0ZklPO1V6mnrInljncB196l7aC2HDLyE56DpSWqs+8Z5PenysEi3KO1MR0jHnJ8xbrSV2tA6gqNGCCct9KflVhLk4x1ps5c+XtHrmmiElSDyM8ilvuGpGkjySli37kdMCmM5+0cnKt6VJISFZEHFIiL8u3r3rRW3BdhJF2KsYBPUkA1Q3ESZI4J/KtTq5Ab5hVQ4kmK7efWrpu+gblqFB8pGea0FACc9TVBSw2joBV1GIHAye1c9UY8rlskfUGoX2kEKPrUzEgfXvSbVAx+tYp9wMZ7Iiffk5qypIAXHXmrxIJIA61QlBV2AXO8810xqOejE9Cy2VjCLzTUcxRYJHFRSO7JtBwwFV+YYgcH5jxThC9wbY9y8r4Y529R6VESz4RMKvpVoEeXwPqaqO6u2ANoxya2UVsidhkgVZCOuOgBqLk424Bz0qZI8EsQQD0NKIl3/Ick9apSswDO08DJ9akDbCp24qUR5BOOV70KNygMKyckyrWFiwONwJ7VZEeFJz9aZEiq2QOBU23dyPyrCctR7ENwB5RKnBqqfnCBTx6VdlTcp4+tUJB0I6VdN9BdSYDY4DcE0+Phs+naoWkUojNnI4qaNMHI+6aclpqMkMe75ifwqdSM4H41WL7GyT17VNFkcj+L9KxlsFyYrjgUh7k9R3pThmyeAKjyAWyetZoCGQ7h6k1XdCT1wBVptoGfWqjFs1tERDgZIFIgxlieM0r/LxQqgpjOa26APB+bkYxSqcAknj09aYxw3SpQy7Bj8qlgSKe4p6HKnjmowTt4FLFnaazaGPGAR6ipQOOnNNJIXgAn1pFfGCDUvUETZ9etNY7R70A7n9hSjBBJ6nvU7AJkcDrS/Nv46UgAJPtSRFj9/qKYCyAFcHrQvQDvTQ4ZsZzTgOSP1o20BAwyuCcUnHUDr3p5IYfSmEbsDsKAIGYiY7gAKMAk8024iaU7lbgdqeoHfqO1a6WGQjBkPqalQAyEDJA6U1T82SPpUkQ4PbPSnLYlIkQl0IPODxVS5YrJgdDVqPKqTUMxyob1qY/EOwJhISoGT3qaLpwcio8ADk8ntT0xGAdw5ok7gPK5OCOlMI/eAdhTpDjGOaaSAflHWoVwHRyEybcfSpTgDA60xQB8+ORT85bNS9wsNjPXj8aFBB56GljJClcUEcjNIYpUE7R0qGRWKnaenSpuARUZA+bNOLEV0Uov1PNTgDoOhqEtsdVPQ96exbdgHAFaO7AeG+bPQCmxryec4oY52gevNOTAkIxxU7IYzccdM1Ep2ueMEVNnG7OMVWMpaIkLhs8VpHXoLRFlTuJOOvenjv61FESV5PWphwMVmw0BOpPenKSQRTV4XPrTshcD1qGAhB29elRcSL1qZuFJqIghfkHX9KpAKQTGMHgUzHz5PSnqeDSMM/UU7gxY+h9aRzhdxPfr6Uo+Uj0okGSV/ho6gNGMkg1IOYyTxTFHJxUgwFwe45pMZCRmMik/5ZjA59acxCggDrTMkLvzjHaqQhxAD7j1pOSVP50wyKzc/hTyAwAFOzW4yHy0Mhk6k1IIz5mSPpikQDOCcYpiMRJ97PPBqndk3LRIJAAqOZC3OcVIF6qaRvnOCazTsyhvAIUVAUO9stVjbgjnmo5gqnhcqaqLAiYkKe1CNhSe9OIynWolIZsZ6VaESfKGyeM0PhnAxTGYEjjv1qRSDJnHFADSADyM0PuAQIeO4qT+Ij9KQnGT1I7UrjY2TIIXGAKa3yMDT0yyc9aY3zHGOR2poSHAFVOeQaFYbQO1DHAx6CmMWXAXjPenYGK74IRRye/pUb7I3GM5NOecBguMmlMIRtwGc96a0WoEyndFUSv+9KAdKmzhMVGEXdzw1QuoEo4xuNPIycgYFRqcnJ6CpCc8/0qGA05Jz61G/ykbev8VSYw2MUxlG7JpoAY7QO1IRuzzSsNw9hQCOgpgUJN0TnOcdhUpXdCCRyKkuIg7E46UyErtw3HpW3NdXEJGwLqD26098GbJPymqxbZcEeveniMk8nNDXUNy9nBA7UhyslJG2UAznFKcZLnv0rHZjBc5LHsfyp2AWwOAaTJ2k049SR+VJh1InUuQFpVbdgjqKM7WJzxQo284HNMGKh/eHOc0ZO0896APmJz9KD6mkA78KGALbz27UikrnPJoJ7dR3oGNkTepGaYf8Ann2FTKepPFR5yPp1poBcYUA9RTicnIpMZwc8GndWz2pMQ0jMvNOUBSaaOXJoT7hBNACE8H0NVpWY4AFWX5xUEmWGAR71cQew3O9QcU5R1b+KkOB8tOUDAHc1TAp3oPBXvU8ajaC3GO9RzYZsEZFPgYscE9OlaP4BCPvWUsO46UtvM27kYzSnJdjnOKjVS0o9R0pdNQLqH5CT1oPHXvSjBGBSdM9xWIxRyQMUpGCDmmoeRmnHk0hjG+V+nWm8l+TTmP7ztmkXhwCKpALuBXIHNK33eeaai5bnkU4EMDnr6UCEfIP0PFRhcyEg4xUjN8wGOlMk+QknqaaARSGUsKVQdpyeaax2kLjjrSsvOQetMBqL5S8t3pSFVgx6Gkk27wD2pzHkLjpQAxy2/A6VERmTPcVJnMmQaY+T84HTrVINxyP5mRgDFPLjkCoY+fu9O5p3IOSeD0oa1ELI4VPTn1prOTEN38VN278A/nRjzV68qaqyAeseyMYOfSpQfkx60wqpVFHQdKfjbMoPPFQxjWOFUd6ZggAZ+tPPOSe1NjyYyepFNbANIBfbnr+lQPnzNmM471YZQZFYnBFV2IMxyDxVxAVT1XsO1O+9kHj2qND+9NKGYknGKbQtB+87kYDPrSPu3kDgU5Q+3Hc03yX8xWU/WloGpHNmUhRwRTSoZxGO3f0qRwDLu3cCkBRlLqefWrWwnpqSvlNqr0HWpCOB70kZUMA3XFNDZzx06VmUPQbZemc0KT5hyaFy2TQBt5x1pAOTBJxzSE4JyKWPAY8c0rjJFLqBX8zbLk08EYYnpTQh3tg8+lOQAgBs+9VoLQfFgnaM8+tSEZO2kiTaCx69qf8ANyazb1GRAAk4/KnA7QPem8ZGfvGpDhhg8YpsApoOFyDS5+THrTeDHjuaSEKBjPagfcJ7U1CTlfTrThwh/WmxjVGSfSl4De1KoAXr+lAwTnFAg+8fxqNzhxUmcKfeopMCIk9aa3Gxg+Z/akYZc+1PiHy7hzQM5JzVdSRoGcimqSVBPGO1OI2jjkml+6RmmP0IZsswHIHaq0hYMIz07VeI/eZ9O9Urhd0m/OMelaQfQTEeFcsepNZEq4kIPP0rcBzEdwxWPcjbMcduldFFu7QLR6DFx34IqYAHqc1ErbmI9qlUjbnmtGRIHUKcdaUDC8Ek9wKSRgBkeuKfGevqanWwxj4PtjoKbGMTZ64pz8MccVGv+tXniqWwG3D1BzVkHIz+tV4sccfWrC54wK4J7ldCZByKo6mn7rnkVejJDDjpUGogNAcjgVVJ2ZLMiBsSr27VqIQT06jisiAr5oYnvjFbEWcj+da1UVEqnO48ZPuaTvyfwqSQfvCSaYcZwBn2qEyhVOM9/pQQPzpqthumPejHOSf/AK1Ni2Fx19KiuMCCQHj5T/KpeSp4ps2Dayf7h/lTjug1OxHBwfwp3BoOMc0YABGa8UY3P50bsjvTgB3FHAwcdaBgMKTzmlB9qBjPpSj/APVSEJ9RRnIxS8kc9aOnA6UAJ05ozkc0dO9B47UAGM9aX26UDpS9RmkMaB70D68+lL7CjGDTAQd80hx60ufypDxQIbznpTORgn86cSCfekydvPStEIkQ+wpzc0meKCcAms3uMrFwLkqV/Gp0HOccU35Zc4HIpwGBj9K0b0AQqQSelNTGenNADBvvcehp5jIOR0NSCIlDISWPBqXOMDFKckY7U1SwbleKHqBLgUwnrnpTgw9OlNLYH9akAzgcUcjrTVbnjpTznNNgO6nOaaWGDzTvr2qFh0oSAchyuTTmA2kjvTNyj5R1pSeABQAAtj3pyn5eaGxxQfuY6Ug6B6AGnMcEU1j8w4wKcATzn9KGMYckgjilA+Uc0pGV45NNz2xQIeDzjoKU8kDPSkU5BowcZx1oACcdOtJkAe9O2jaOeaY+RjjOaEA7+VB5PoDSnigDjmku4wC8HnikKqRyOe1KeuBSEYGc9KELqKACMDgCgnikHUZGPWnHkfXrQ9wE6rgnikXoAKU8nB6GgkdADQAp4/CgA8kijHPPSgE84pIBDj8KRcE5J49KNxAIPegjPB6UwI365PamgAcZ6098FuenrTQcsAF+tWgEj3DPcU8HII7UzLFuBx6VIvMnP4UPuMJOMHP1qFZCz4xirDYJ5qHdhwNvJojqK5MMD5Bye9ICQCD19aUjAyTz7UwYx170rDHn7uetRxktnPX2qVskYxwaMbFOBST0EQMjqykH604jbkgVJjIyTig4L5xmq5tLDIiBsCkZzTsKgAUdKeACST3o/iHFLm0sIYc9TxTSoIPrUgQbSKaVXPXp0ouMarMG24+lPJBPtTSG8wHPApxBzx1pvUBAysg44pylWO7FCjbGT3x1qOIsI8sec9KGhEqkbm45pV6mmlsfj0pyn5eTUu4Bnr7UoIzkjioi53Y9akjVgx9KLAIRu5Jpkw3RFR3qRu9V4rhZHZVH404pvYCKzO5GQjGDVw+1QJ/rGGfwqX7yY7iqqO7uAONvGajYM0gHQCpcZb5u3amdXJJxUoBUbIYdMUin5uOtCKu445HrTs4kBHT2p9RhtCnnrScbgewpzZLZAzTMlc8c0kAMcyAAdaax5wCCafgkjceaY6Ddu/WmmgHoDjOefWlUc8n603Py8flTsYAyeT1pMBR0OO1HU+9IDsBPrSjHPqaVgI34XIPNU7jcUwAcnvVxsOMZxioWPUela03YXQIDhQrcn1pWLFGGeTUEVx8pLLyKas7SAsVwB0FW4Nu4wlzt2YznrUAMcTAjoDzToZHk3HHFM2qX2buB1rWKtoyRl07K4wchu1JE4QYPTualeEbAclj2zTYVU7t341aatYY9PMdsdv6VJt2SLGp4XqKkQHftUfLTygkYMy8r71m5oLDXQLyOPeowGG0E4z3qdgHGOg75qrdOyMu3pUwu3YB3dl6YqPhWGD81NSRhgvwaFQswYn6Vqo2eoDgxw7AfMahDNHw/3yODirMUYRSpOWaql2hZ0YnBXvTg1zWF0LMbhYzuOSO9W7di/OeDWeCpgzmp7aZMhVPWpqR0Y9i+MHOOaQgH5T1oVvn4HAoIYuO1coxEQDORzVSQblcM+xR6VbB+bGOnaq9ym4FV69zVxeoirGDsdlYY7U7IkjIAwFqJOIwFGRVqGMrAdwwTXTPTUS1KXzi2YYJ/HFNRVVc9D/OrrQ5iIzjPaofKB2oT1701UT0E12I9xCFR09alI+VQB0pzKI1IUYHqaiQs5VhwO/vQ2mtAJS6+ZsHSl3Yzg8elRudrKUAyfWnSoY/u8nvUWAsIMR7c5IqZABGBnB71UWURhj3HU1LEfMQ4ODWcovqMe/8AEBzWblvm3cgnpV9xz8p47mqjMoUqFByeuetVTQwDK8ZBGMVOrKcLnp+tV14cLjAI5qSKPY+evvVuwEzAlcdRmrI+UhRwKgUjd0wPSpSzM6ADjvWMtdBkpGCBUb4wQRxUuDnPaoiTg8Vkg2IJGygIqs+dp/lU8jdMniq0jDJ/WuiCJvoQyZBBNCMApyfxoMi9hzRgsSAMZrbpqA9Fzk9qlVFBA71EDsGMZxUism3cec1LvYCRV2kn9aXksOOaTePL3Ckjb5Sw6ntWdgHtIQAP1oDfJhTz601iY+TSKOQT37UWVgLCsvUnJ705WyWPYVCmN+M/jTwxJwOhqGhkyYLE0xm2Rtxk0qNjPrSE/uyW6+lStwEi2tyop4O7OOBUTfu0AXvUnOMjvTYx64+bFMxmIgdaVep4zSL/ABGkIaoIznpUIA5cHrUyrjPeolxtJA+X3q4gMzsYD8jTowQPmJzTX5lApQ204PNXqHUkjbLEHp6VWnkIkyF4Hc1ODtO8VWlb5yp7npRBaiZcTawDZ69aQgP+Bpsf3Nh79KUjy0AXoe1S1ZjHMd5+U0SADbk9KeAAhycEU0kj3I6ZqEDJRywPYdqc74z7HtUAJMgODUjAkcDilbuA4ZH40jEZGetIDuemSD5gSaVtRkuBgjrimOcgAdaMEtkd+lDk7hgdKFuLoRSfIq5GTS8uSccdqc/+rORlvSmOxGM9ParQBGuwsT0pwIDgdqZuYxlcfjUjj5F7kDim/MBG+WNs8VVmkEcA/vA1ckwFrOvW5Ax9BVU1dg9i5D86Bl6Y5qZeOpyTUFn/AMe4FWEXK5x0qKnxMBRjJXFKoHQjOOlIPv8A86Xq/HaoARuQR3puQTtzT2PIIqFW/ec/nQloA5Rljg5pOWbHSnA7Xz6005UEnr/KmA7gsM0mNjYNNXjJJ/Cnt90+poAjz+8J9KkBz170wcEZ7089+c02O4yQZJA4ppz5fPWidj5yAdKQEtMyn7oppOwh4RdhwM1Ex2JkcmpGJSMimIgEQ7mmg1IOBl8c+tOBzg9MVNhRwBzUJ4UknAzVJ3BltTnGSc0MCMelC48sMKVskjJ61m9wG4BDU1hhfTFSdWPHFMfqcUIZXIJUjP0qNHAHTn1qZlzt/pUMuMAjitlroJjN7edyv0qxnCk44qpM5WRSvQirQJdAKclswQ4uQpbHSmqmWLdM0/BKhBTwCqlc5rO9tgI2Pzr6elROx8wEdDUrjPPWo0+vIqkA7r2pSMgjPPekjzvyeaVR8xJ6UgGbFVmY/e45qOWUu4WPnHepnC7Sc8d6hV0ROOn0ql3AsxhtozzUZDFw2elPTlMnpUMqbtnPQ1K3AsA4jAqT7oGTUY52jsOtPA4OahgBPJOeDTCdoOOc05l5AAqKZyHVV6DrTSuFxzA+WfSkGP4aUk4OaaylR8vWmAFwmC3Q9qbIgYbh1PIoVcjB7+tPCgsB6U72ApznMg45FP3klV/OrMsYfIbrVKPzVkZXOQK0i1JCsWh8qkZ4FOUnAyeKAd0WMdajZvJUeufyqLXKJ0OWHHFKPmOO1MT71LnGKgQMMqVxmkPEYx1pwPzHmmK2c55HtTAeCfp70HDIcimgnzCB0p+PlPNIAjGF70uMnPamqSetLuwKGMRuMnNNxyxzg04dGyaRhyKBAv3TjtTk5Tnoaao4IzzS5wv0oADgLn1oTgeuaOABkUg+8PSgBTgL7+lQdHOe9Snk8jio2GQT3NVENiJBjcB1zUirhixpQMEsTzQoBBNU2BHIowxxxUcJX8KnlGAB+VVmcJOEBxVR1QdBjSIflUmp0bcuCfm7GoGAEh6Y9amA8yIlOPSqlawtUTx5BINOJyCMUgzs469M048gZ7Vi9xiYyhP5U8HKdOaa2dvWhOFGec0gGOPzpDkDOcetKxAyT0pMk9s5qgFQDmheuAaAMsBnpSoBubtQApyxODmmPzgnpThkA5pBkjH5UICJmywB6U8ff3flSDkE45oB/dcHmqAFG9txpqcg7uop4J4702MnaQSCaAYiD5OnzGos4Rt3HpUpG4hgcY7VG4DKWPQfrVIBkKt82aeIywAJwRSRMGVSBxUjZJGByKbbuHqRMAmcnGelKpMQVGGd3elnzkAjg0vzCQAAYFF7oQ7aEYf3R0pcjeD6d6awLSf7PanN99cDFSMjckNg9+tP6ZAHFJnPJHNJGxyVODT6AJKAWXnpUUoGX55qZgN4bqKgnXJOTx7VUQCPDHFPVhuK4qFQNvBzipBn7y8e9U0K4pBKD5sc4p5U7FXoT3pm3K5PXsfWjd+7Azz9KkegyVQPk9etAUJIgPHXileNWbOenakGHkwR06c1V9ABGKL82A2albh8gcetQohb7/Y9Kmk+ZcDnFJ7i6Cj/AFjjoKkTD/QVGzFWUY4qYABTjnPNQx2IxxMWHIpxGTk80KCVLd6d94Z/Ok2C2I1HLevrSKm47jQemB1p68AccGnewh45ApWPzZpAf/rUnQEnrUDuNP8ArDSt94UABl3ClBB6jmmIGPy4FJjAHFIBjAFOYjH9KB3EGA30o9DSDnjHNO4Az+VAluMDYXJpe2c/hSE7fnx1pQMg0wFP3D7VE/zfKegqXBNQsMyA04gRF8LtUcCnREsTzxUO4K2HPJ6VPGpCbQK0asgFY84HNKRjJGSaH9O9KBwDUAJt2k571TdNzlugqxKCx2ioWRiFycD0rSOmog3FYyAMmse5DeZhh71uL6YFZV+gDbhx61tRfvB5lZD7jJ7U7k8AAAVEnXJGOal3Y6da6GSxu44zj86lgUjJIyaaUyBTl+Qkg8dhSb0C+gr53YFQZAYYJzmpWbr2qI/KA3THanEEjZgA2oec/wA6uouW6/jWfatujVvWr6tkZFcdW9ylsSrw4pL9QbRzjNA4J/nT7gboWHqKVO3MTI5yLIkAHGDWvGeFYdBWQfllH171qxnCjjk1tV2HF3GTcSn3pmO5p85AlJ9aj657Cs1sUKTQOcAigkdcZ96cpzyKfQBD1yajuMfZpP8AdP8AKpTyT7dqjuADbS+yH+VEd0M7TqOKTjODx6UDOOaVc4NeMAH0zxQcdKOOtIeevSgA9PSnLwM03Bzg9PWlK5FIBR7nijFHbig5HOKQCE+tHajIJ5pc8DinsAnOelL3o6c0vJPNAwxg0gGWwKU9euTSHjpSEJmkIJHvS4GcjtQSc9Kdg6Ee2jt0pxGO3NBGVNVcBcgjgcihlLAgUKflp6+hH4VLAjRSg96U5xnv3p/fB5pCO9FwI1IIBY80YO4HPFGQw6YIpjqWAxxiqAlYtjilXcOc1CHIbZ37GpFYjg/eoadgHgk8AUx038EU8j8KPf1qUBGCF4zinjPGaTblge4pen0p7gL0600jPQU8EkcU3kdxSGJjPanADOR1pOgPpQFAB7GgQnO7pn3pTzgdDSE45zRg4HvTAc/PbJFKudoz1prbgAW/ClVgVyP0otoMUjaOucUh6YxmnA88j8aXsTUiGoCopcnHWgDjFGcdqBgBjgnpScFgDS88UDAYjFAAwycUoVVA596FBzmjHODzSAReFz60h+5mn54I7U3Axj+dAg6dTS9fpTfvAYpc5APSmMUnJx2poG1snle9KzYwCeT2pM5OOtC8xCnBOTzS53EnpQMHGOnc0uOcDFAdBmPnAzS5HfrTsBWx1phwuCRQA2QdqYpz0FSttPOKjODgdqpDEUlWJPWnrlmpm07xjp3px3dhihgSDkEY/Gq5XBDMcmpXB2gHjNRsASDnp3oiIch3Dd3pTuDDC00Y28d6kUHJo6jHLllyaODQpLE570rEA4FS9wEZcnApTyKQkYIxSdBQIUrlsZ5oPDjHWkQ4wfWlJPPFD3ATO4njGKiyc5/IVIBhMHk5pqg7CAfxqlYY1hgA5qViQv1pgH5fzokxtJ7Yo3DoAZmJHQUmRkLnmooJFcf3fQU2N0Mjj+Javld9RFgggDA6U8kkBaa+dgANKh4wO1Qxi7cEU8HC5P40nKL2/wAKM7hUMBhOcnFQpEqS5VeWNWAflPpTQPmOODVptAivgrfEDoRVrlTu4FNUDO49aD8xxRJ3sg0sHIzkfjTCMY4znvUjAnFMPynHT3oTDRDU5kIH3RSuuOhoUBe9K44JzxR1AcpwDTQAcmncmMDvTVycgkfSgBq8854pWGARnmlTHPPSh1AXOcnvR1AYOAAT1qViAoHeo2IAB709fu5xz702gAHcTjoKCSSfSmgcc0uDzzk0gGYGRUEuFcju3erJAB681DOPmGBVxeoFRIXUHnjPFTNGWiG0c1MjZG3GGIpo4UA9jWjm+wrdCuI9qCMHAHWoXjKyhUPWpZWMYYk4JqMlMZbODWkXITCdcoCh5FNgBRWcdTUs4VYQR0PU02GIo6qoyDyapNctg6ltXO7OOvSmuGEox071IgBIz1ppZyzKeKw66DGXG0LhRnNM+YSBWAxTyyuwBByaa5/iPJqo6aCGPtaXBXC+tI2PMwv3QOaVj8vy/lUTSsuRj7vaqSewxyFjMBjI9fSobq3Lys5I29vapYWZSC3U0kwKwuR8x7VabUhMiADxKqng0RIY3Ugd8CmwRkwBx1qzHIqoFP3x29KuTaYLzLDSFJlQDr1NTA4Y5JOelQp82CRj3qcfMBXJL0GRrgSnPQUyYF02r36mpDgc4yxqOYMVyOp7U0IjVFjjAxg9yaljwEXjHtTIV3A7jwD0omL57BexxVt3dmMHA37u1VZHMcgYH5e1Odi5wB+VRg4K7hkrWsIsnUaZGaMgj65NNjyAeMY65NROGYsxJp8JGxlPUnrWrVhE7KGKKDmnbwWNMJ+U4XgdOaeq/LuyPp61nfSwxyRAKW25FPCEAbeneiMMyY7DtUisSvWs222MTHyn371UcbXIHH9KvGPcgxVSZQjkY5NEN7ANUOxX6cU6N5cDIyKZExyCR9KkgYHkncO9aPUZLsfcflqxBuEW1ux60m7DY9akRwDsPSuaTvuMcxyCAOKiYYHFSHgEjgYqInI+tQhFWZfmx2FV5Cewq3OACcVVcgL710wYFcc8YqSLhSMZJpN3GTwKVWPVRWrJWg4JlSSePalUAuuefSgA4/pUayq0gHXFTqMtZJ4Hao0djJgjC9qch+fBHFJu5JI4FShjyEc4P86aFwwH6084YZpPkxwTU3AeME5xTlJ9KaoGCehoHD9PpSaGyRSOe9OLAqAKiXIUk9fSnRnIzipasIVwWcbTwKeAGbnqOlMUgEkdqVB82e1ICVTh8d6RTyR0FMX73AwaMHzAetKwwRSuSfXrULs+3IGAKmLsSSOc9qikXDbQcDuKtPXUncjkJyB60ZJQbTzmmSgmXGcAUpOFDDmtOiAmUMCnH1qG5iYyjaasRksoYDHrTZidoJ71KfvDYI+9sL2FE6kopHUURqIwXPU/rSyfMgUfnQ371w6AGPmCPim3HzLgtgCnqMTkhe1Lt38Y4qbpO4PYdGojTHrUi/KQKhw28ZPFTDlsj9KiQEWRlucD+dDZdaJAATt/Kl5K7afmAqA7V9aJByD2NPHGMio1I3nP3aV9RiPgNgdBUa7llOeVPSpernA4NRkkkYH4VSEIASh3cHtUuPlAP4UwncR29KHcFsA80O4IAW2HPSqFyQ+OenQVddWC4zxWdNtEhAH45rWkrsTL1kGERJ9eKtZ+UDvVOzLA4IyMVdJJIArOr8TGNIIYGpMbQMVG4OMA804HcACcAVmw2DA2Zpm3HJ69qk6E+lRuhYgg00xi4PUjmmZ3KQetOYlefSmg8HPWmhCc5CgZFPlH3cdqah5JH4U8DnmhsLjGUZJpwOFwKZk+YfSnKRuI/KmwGzDMi4HFNZh5u3qfWnXAwoNNYlOO57047DElw4KUKQilRzjpRIODng1Gm4SZHbrTS0AeCxJJGG7VAfmJQnnNTyZ3A96rOPLfryfWqiIuwsCowOKduIbOKhhOFA6VMxyAR2rNrUYoB3+gIppBINKMlj60o6HNIRBICUAXsaidcoDjJFSuSGwORTM87VHIrSLYylJKBIoPTtVmJy0ZH61UcFpsMORVmFkDbRyPWtppWJLC/KBzzUpzxxz3qFTuOcVOx+XpXPIZHtAJ5qHhXwOpqUnJx3Pao2QA7j26VUQCI/Mcinr8rnPSkA4HNLz0FDAbMhPygdaa6AsMfw9RUpGSCaaQFDE9KaYEUJEjdePSnyMA2WpI03MSq4H1pJBkNuGQO3rTerDUkQ5TI6GpuQAtVrdlMXAOfSrCnJ+lRLRghzELgVEwBJYGpOOp7VC7kcDn2pLUBFAOMHOaVsq+etKFKqMgCkcgLzzVdQGYPm5J6VIGwxPamZB6UvyhcDoaBjifMUkfhUHkjduBOR1FTZ2KFHUVWn37QVOCe3rVQEy1Gx2kkAehqPZ85yMk0kDED5jn2qXjJalsx2EOQevFOHB+XkCkOQc0LwNvc1ImOOOTimoqjO2gn5SDRDkAluKfQBw55oyVBI701gefanjlc45FSA1Sdoz3pZCQB2pPvc+lK2Mc9D0o6gN5OeOtOPTINQA/vzu7VKgATOeapqwIULkD370obgrQoJIx2pR9/NSMaTxjFD52jsc0vTqcCkBySe1AhrcgjuaaFwoyelOGVQ8c+tKBkDFUA3aNvNKg+Xkdac2M8dKRjyAOlK4Ecgz0PAqmSJZCq/exjNXXOFwOc9ahjRFcnFaQdkBVWAjKsevarKL+78tTjPQUMGV+KGcJKozircnIRYX7u2lx8pGaFXLFiaT/AJaYzWJQH7uKVDgY70h+ZxxRglhjt0oF6DZFyB60PwBg80smS2R+NHB5NMAAOCe9KBnLd6Dw2SaT7ufekDDdhfrSSHb0o5EnPT0okyQQadgGN97rxTkYlecACkAGCnekUlY9vc0wDkMWP4U0qAG96cwyoB4xTXAdOtNADfd2kc1GxJJToKmJ4GOajcZ5xz600GpFB8kmwjpVgDMh7VDHuK7ivIqWMkndTkC2Bss+R0pkQ4MjcZ4YU9w2xivXtQVDR7ScEnrST0FoCYVPX0pCW2bjSqvBBGMUo5AxQPoIMBTxUaN8p9RUoGWPtVcttdl65prUBzLmIjvULkOFUD8akDExHbUDnZGO1XFC0JVXkj9KcqfwkdahVi6ht3NTbscnnFDGLt3Hg5x0pN2+UAcUqsq4P8RpQu2QMO/TNIBvBfPQjr71GTvUsnBp5JM5JGKhyJGUhse1UkLYnUCRVfOCKI4ihYk9egpqMscwTsalBy5yfwpPQCASMmBJ2qzE4MeVHBqEgv8ALjmpIccryaUrWGSYxjH4U1yVAOKCW3nHSklY7QR1qUtRDMlHx2NSjpj0qIYO1+46VLGcg5okA4DPSj+BqXPXHakfAXGeTUgJjgY6U5zgZppbauBTnU+WCKB9RqjBLdqjVZDMzEfKelSgHYAcc07BAC55p3sA0jngc0gGRg9KU5ycUmMJn1pCGK5Iye1SZAPHWmSAlB6U8AFOKbCwnIyD/Ooj8oJP4VIeV9/SmyAFNuevemgKRjDMHJwB2q0rHbntUSIqnYTk1PkD5R0q5MEMbBbntSkkjrSnhh6UwHewOOM0g2Fx8nPFRnPQDNSMC3XiorjckX7sZNOO9g3EkZY0wW61lXcgYbBznvTnciPczZc9qpjcx4z+NddOnbUm4q4ztxk1JgEcDkUw4HTt6d6lQKVJPUVqxMQgggZ5pRnPSmMS7+/QVIv3uelILDSOPemEAg56ipnO5zxmoXPyHAzQhI1rIbrYcVdToB61RsjugAxn15q8uBXFU+JmiJOafNIBDyKbxkU24UtCD2HFTDQT0MJgBIQfWtGD/VCs+VQkxzyKvWx/djI610z+G5MVZi3Qy49D0pnfnrU10MBe4xUI+Ye/pWSehoBPPp7etKnBwDijGOAeKQcHGKoVtR+eOlRTg/ZZcf3D/KnAEnPbvTZyBby4/uH+VCXvIZ2YNLnimg5PNP4NeKwG4H50Hgc0uKQkmgAGQOv1pQSBzSDoM9KUHHTkUAHBzQCeaOnJoHTNAAByRS9BkU0MCcd6cwyeKAEzkDNKM5pGpRx9KADvzQTk0vbmkAxnNACHpRnjBpGOOaMA854poAPPGKTcNuOtBODSZGaAHg8YzSjGM0gJIwKU/pSYCg55pece2KTPGKPqaQDTjp60zo3JoYlT8oyKa+1yB0PpVpAMBw5yM4706IlgS3UU4EqwG3PvQOpA4qrgSZyCKbndmhEYAluD2pGJBGBn39KgGMRXRizNUwJIPNNdcrwM0Bjt6c0276gOB3ClxjvQp4FKajqMToeaQZGe9M35bA/GpABkZNVsIXBbA6Ad6CQWBJ4o7YpoAYkGktBiyHimqNg64BpSO4pFJJOaa0QiTsKOwA60Yxk5peAMkVI2JjLY9KVuGoAJ5FHcUgD600Dv60HnoaUA4HzdKYai55GKFHOT1o+n40inDZPIpAGQG6dKCCzE0pUFqXoc0XEQoGL9MAVJghTmgHDYHU0Hlvam3cBCuSGPOKQE7eetPK9BmmnAPJprXQBQSqnNO6DI70hGeB2oA7VLAMHI9KacFiTTs8im5w3NNDDgjntUQI7j6VJkbelR4Bb3qkIYhkeVgwwB0NObhz608nJprYZyM076hYeT0yahyoJXvUuM4ApjhQxOOaEA4Ek9OBTlyzZ28VAWKjk4Hc1YRhtwB9KGrILkmeRgdKGCD3P06UMcDjqaU7fTmoTtqgGjCgjGc9abnIx3pcnIGOtKRgc029R2GgbgO1DEgdeKccbcEU1sbMYpAGd3U8GgcIccj6UnGMU5gQpAyKAG4wMgUbfk6daQghBg08btv9KbAgiTYQxH0NK8QaTfjn1qTnG2kbIOKrmdxIVgPxpF+XHFMzuBzxiiBt+D2oUdB9Sc+gpDkdqQ8DrxmkzjJPXtUWAGO3HzZHpS5y5PrR1yfSk6DJpgK3J4HQYpucHpS8hcHvSnjFACL1yabxkk9akBKqB+VNPzEEnFAEYIbtxT327ePxpmUYgDseak4bim9wQpJKgDimYK9qUEnihwS47UX1AZGAFPPWhuhpoXB/SnZ+bHYU+ohjEAD1p6lmXFRynOSP8A9VOQ4XGeaq11oMVgWZeaeTjk0gJ280i5PbipAfx6ZHaq8gI7jcasEEgMTj2qtMBnLDpTjuIUbgoHG496XKyY9fWgSAKoC8mhsLnim9xkM0Akly3NVJ3WPK44HarjHaVOOvaqNzGxlXbnGcmt6T1syWOciWFcoVAp0DFJAM5U96cY8Qbi3XnpUES4bJPGeBV73SDqaQBV1Kn8KRyxJVfzqOQNIi7elPLMIQQPmrCwxWK4X1zyai3EbhjjtSjcIiXOWpANqnI600gAqpXcTnHOKrsjeYfmyCPSpXYhQRyKh3GVcIOe9aRT3FuPgztKt1HT3pscJhDM7bs9qshAFH94DrVTe7tIMcAURd27DYy2dZWP93tSTKUcsD1pIyqN5Y796luFbyl4zjrWzdmg6E9tKZRg9uoqyGBYY+7VK2TYoJJyatKeFzwfSueqrPQOg9sbgQP/AK1JIvDGlZizcVGyM3U1mgBBgcHk9qWTGAG64poYBsHmopZA8iYb61pytsVyNyqA+pqBwFxtHNP2sTuJ+XPFKyY3c9a2jZMncgkkCnYn4mmLcGGRs457Yp8qk4A69+KqXIxIoUk54zW0feDZlyKUy/Kv4nNWlTG3b2FULZW4CnBFasXC88Gsato7DHLGMbd340LgLgflTl5Vt1JtBIwOa5rlDlI2YPSq9wc/dHNT9MgflVeYEoD0Y9KqO4EK7VVlJ5p9oBsY5yM0xVwSzkj3qRVdkyhwPpWzDYtjAIY9aUBd5P501MeWB2HemjAYk1zvsBY6qR+tQnAOacsgkjB6HNMdhkc0krMGRyHuapuMk8VbbqT2qlIGYkr1raAiIEODz0pqn0PB61Ko2cE9aZsABHet09QJs8Ajr2pqrjoM4pMkY44pFlYpkKeOlTZhsSoMDJBp0o+UfLnPWk+cLjOWqQMNuG6ipbsxje2B1P6UKAr4PIphXc27OKTcEZQTkHrRa4l5lofN7YppzkADjvQPu8cilHIrMAGAM9aVcnOKTIwO+aXn8TQMdnCkDrQrELzTcAZAzSA4Bycn0pWAkDjOaTcwbg9ajVuOKN+Vz0xRYCRHwdoOfemSKrPlT1oDBEAI5NDkIvAp7PQCNypIAH1ojJMeMc56+lRsS0i7RT1yGKenetbOwixEdqnHUUyZd43HqvamhgDtz060s5VOOvrWdveDQIlxhSeakjYFsE7iPaoMeXsxzu6VKjLGRnv2pzT1GSA7mB7ClU8dcAUzdncQvSnIAVx0NZsNxYyOS1SIRtLdqYWHAHanKDtxjrUsERbixJAqRO5/Wo9w3FR2p+dsRJ4xTaHcechA3QmmqAvv6UBt69eKMYIFIBFOCR6VCSQ241M2dzEfhUEoUkAg5PariLoPK4Yn8hUeByeopzYUqvOe5px6YNPVANZv3XzGs2TDXDZyRWhJgAk8A9Kz/UsPp71rT6slmhaupXpyO1WuhAFVbThQw+6egqyp+bJrGpuUGOCetP4wMU0d80KdhxWbGPGTn3qINgnNSjhsk8VERu3Y70IQjN1zTCcMWPGaGXcuwHkUOmQOeatWAUYAx2NSD5ScnimDjAPanBc5pMCM4Zxg9aTJEoyOKcyhZBjjNKSCadwBgW+90B61C+TPtXp3NT5zGSOtNC4TPc9aadgGMg3gnNI3EhGTipiABzzVdzl8r27UJ3AUsTICThahlCySA5G4VKybqqKpU/WtIJAXIgFc/SpwPvDHSqlsTznqfWrQOQRnrUTVmHQQEkn17U4enrSYwcY6Uo6n1qBkLDDYNMGWbd0p8h6kj6U1AXPJ6da0WiAqTR4lZu3Y1HC6h+evrVucgLkDrVQkcEDgda2g7rUm2pfQ5hz04qX5miDVArYUc8EVYjz5OPWsJaFDAoJB71GR8p3dqmK8FR0qLA5BoTECYIIB4px4UHOPU01cfdxgCnEE59DTe4xWVWUDHNMlVWgKk496HOxsE8VFIBInB4oSJC3cZI6ilcHJxyP50RgdSBj0pZgSAV6dqp7j6CW3KkelWRjoKqW5KjZnkHnirajuR1qam4IXHy1CYwh3qDk1P94HnFNbBAPpUJ2Abjnio3xyoGTUhy3Ipoxv5POOtUgGINg559qYkgftjbUkYyCxPFQqzM7BRlatDFmkIkXPQ96lBD/e69qhuE8yFAODUyKQAT1WjTlF1FVkfAHDU8KCqk0kgVyMDDHpTg3YcmoYxjHqe1C/f+lKwzlfzqPBViKEApOWNETFx833hTZCd2FPP86Rgowc1VtBbkw+YN69qUAqMGkJyBilJyoqAEY/NwOKey/uwO/rTBzgHpTskrQxjWA4POe9AGCRT268VED94dD6+tNbC2HA8Fie/FIMhs5+9QRgD1NJwCfXtQHUUcnBP1poJz7UhXqenril3Z+UdKYCtljxSjgcd+tJyEI9acAScDt0pAGDjHSkHHSgkk7R270YwKAGEY96iPL5A6etS9s9SaikOCAauIEZLlu3NNl4XfjJz0qR0G7Pp1qAnCOc8elaREXo2LoG9R0pRyxPcdqr20g8sL+Aqwoxk1lJWYxe5waB8oyT16Ugydw70Z3IF7ikAMcpimkAgAnp3oY/MARxSMwEe/H4ZppAG4NgntTvvKfWoj8uCB9akAIx+tNjBs8USE9ulOGN3rSYyCKQrETErJ9elLgtJ7GjdmQemaaSfNP93tVCHSZLYxxTWAyFp0gOQBS/xAE0LYY08yY9O4okyQecUDHm570pXClSetAakYJIG09KfyMjv/OkQYfPan4AehsewN/qiT371EX2RqevPWpG+ZXB6VCoPmhMfKOhppCe48swn29iKQHDhfTrSBMSby3ApoJkl3DqP1p2Am4BLDp6VE2dpHcVKcY46VBIMfMO/WiIChQnGc8VCQ3kkOMmnh+FNNl5OQ2M1a3EQwcJz0qwrAtjHFRFDnJ5AqQMFHHfrTlqC0Ghju3AcetSmTEZLnAJGKimfZCCB1qNizQRlm5zkcUWuMsl8ZYjOOnvUawgluNuelPkJESjrmkkO11OTSXkIjKrLKG3crT3cPIjfpTZSkLk5xupyFFjPOQOlPzGxGm2uABye9TwrsLDuelVyGLjjHvVqPO75uc1MrWEKwxg96RgGT6UFtjMSfoKR1Bi9/SoGRuwDqB1NSYKkLUch+deKfEQx9feqewiX+IYprAFjntQvXFIcAH1qBsaDu6DipT0xmonJDD0p+ehpsB3OAPShjnk0MRuz+VNOcdeKQB2zQeRz2pDjAFDmmIap3bv0FOjBAOTyaYz4A29O9OALHOeKbAUgheaZnCcnntT2YAio5doQUIZFwBuY8g1KGGwAcZqvKQI2xTlJlKkDlRWjV1cCQg460jZXp+VP/iz2qKadI2xnJqVqw2JAOCWNUJrre7pwIx+tRXN6SCMhR6VnvIzHCjj1rpp0nuyb9EK7FyxJ5NNLbT1HvTvKcDJyamSykdfmH610XS3FYqFskkVNAcsM8Z61EVAYqDwKenyqCRzTew3sTSJggqT+PakQ4waRm3KDnJ9BSDhieMetT0IJCMHk9BVcnG4YwDUxPU+neoz83U9fSiI4o0dM5Q4/KtM8GsnS2AyMYrW7VyVviNFbYeCCKbNNtXbjn0pyDGBSXUWZAccYqIRu7kSMO4YtOTxj0q5aMGX6daq3aGOfaMHNT2qbQSe9dMtYi1uWbpsxpkcVAOepqzOMwA5xiqiMNmc9axjsXqSHkCm5ycAdKcWBOKaG45qg0FXrn0qObm3k/3G/lTl+U46e1NuCfs0g9VP8qa3Q7dTtcZpSc4o7Gjj06V4gBz2pAATTsU3BAoGA5PTilA9KQAHkdaUYweaBCMMfjQOgBpxIxQOO1FwEOC3HSlwM0duOKTHBoAB6UoOaTAPPel4X/CgEB5zQfSjkmg8UgGEBuKfj6UmBnjmlHSmAzsaOM89cU498U0nB6U0AKfmx3p/161GAd3FSA5PShgHUYpcDHIo6UgIIqQGsdlRMreZUpG4ZPakYAVSdgGKdrnJp2QACeCahKZkBqYfN1HSqa6jJM575xQMEnrSAnsKCDjrUADHnjpS4BAIpoHoMj1p5OKbuAgAOfWnH6UgbNB+7zyRSEIQByBQCARnvQTxik4JHOcUDHEcHNNCex560rD5sDpS4wMc470BoKDhfekHIpRwOtJ2pAOA4BIzSkf/AKqTHTJNIOck/nSAcOBzSDHUmkxkUdTk9KLAKc8cUdMijJ3E46UmMsM00Avp60HgUhOacf1pCDODxSjOfrTSeODijPT1oGNKfvM+lO/CnYG7NM/i9ae4DuBTcfNzS9cj9KYxIXO3PtQhC5IY0/ofemKRn0xTs8ZNDAXtmkAIPFB5zx1pACRz+FIYwgk/WoiBv9anAI4qM4UnJq4iFwpbgcikIGCcc+lGQnOKRWyxbH4UwHoeKiY8gEcH1NSRuOcDmmyReco3dqFvqNi7RnaTVhQq9KgOPlHU1IpAPtUyAU8Ek/hTgPlznJoyTmgEBRzzUi2GEbcAUpPzgdaa4OM0g4IpjDeGJ9BQrbjijpJ04pwC5yO9PQQinPanlm3Ypqk8nHNKRzn1pdRjSxD56YoD5OTzTXGTtU8jrTk6Zp9BDF++2ePanPgLnNKApb1pr5fChc880bsY1VBYgHGacqqnFOHUnHSkHEhNO7AeRwFHSk2j160oOc57UgPB45qUAzb8x9KHIOSBxSg+3NNXBbJ6CqAcxVVBpxUtgkYqu53dPu1MNzc44otoIXg9aTt9KfwBTFYjtxUj1IIwpclfxFTDnjoBTRwWzx705COfSrk7gNjIDN7d6CT1Y0JgSkYpflyaTEMBAkBoP3+KVhuHHHvSHqSeg9KaGMfOz3oTJ2ketNQlmORxU6HuOlU9NAF2Bc9ietIcqoxTvvNTVIkIwelQA9Rufio5ANjAjNSKQckVHKoyOePahaMVyGMFz6AVIQxGKhjOHKjIA7VLhlIPY1rJDGyZ4GenaoZB+7GRkVPKoUgnkelQz/NGVOB/SiIiB8mHDcAdqhx5rADsailYmL5ZMgdqUAhQq53dTXXGHmI01wsajNOCbEYKeT0qKNR5GD/+qpXDZTaPrXLLRsogYuYsFcN3qNyQNpXj1BqxLjay45Heq6kMnPJHSrjqhOwk3y2xCDJHaobNeSCcDtU6ssIY9zULR7yrxNtU1rF6WFbW5ac7iAO1MnZVi+bnPWlJ2OoWluCu0KoyT2rO2oMprCImKs456cVaQb3wRwvvVSWEzTkZwelWoX/eNGRkgda0lqr9QRY8vdEOwpyjAY4+lPDfuelJIOhXqK5rvYYzIVMAc96byBvHOamABC4qNuWMfWmtRDEwW+bk1RmTYpcnaT29augKikj5vxqvP8y/OM1rB2YOzIyQY1BJGehpWwsZD5yegp6JgB/7o4+tRkeap3cnqTWt1cW5DuySoGFqtKAGCY/GrOwcBj8nWmMiGfA/GtYaCe+o2PKHOa04JVIHGAazTGVfg8VatGG9jnkVNSKkgRojG4k9KOOT60yJtxOfriplHPtXC9GXuRdX46+tV7ltvyg4NWZwdhCcN61TfEcZYnLjvirpq+on2Iy5detP85gQoPBHXFRxsrR5fNELozkg4Hauiy2sGpY5MRAOWFK2dqKTzjk02MESqCSRU0mGOBWUtHYYqqAQBTSpMh46U/kEAdaQ8KQDzUdQI2G5sYqF8E9Kl6ZHQ96YRnnH41S0EVmI836VE2STjGTUzocZI4NQE7QSOtbw8gJFIKhWP1oimSQ+WhyfWqkzsYSASc8GnWsbKd/P0quRWuxasv8AzAqRyT1pzggYxn8aSM7xkg49KkUgHHftWD0GRDDdjkUeWCwNLKApLYyaSNicHFPfUZMAcY4xSEEDjnNPGG6mkIyOtQA3ndinEADrmm4ywzTzyelADTkA89elNVeOetPK4IJIxTM8nvTQCZ4z3NI3zfLmlB3EHFNAIYZpoB6sA3I+WlaQ+nFR5O4Y4A7U7O5ckUNAiN2IK+ntTVJ3EA89aRwcgk5FNj2mbcDWsdELqWFIM2MZ981JIo3Fic57VFvBBJOce1SxkMozWbvuAmCCpPTt70+RVLqc/lUTlRhX9eKkHzoSDz60n3GG05IHb9afECM5PNLtGMd6XbiQAVDfcCBlcTYJyh6n0qwAACQ2akIyCAKjI/dnHX1pc11YBqJtbPb3pWGYySetNAOQxyeOlOkJx8o4o6hpYWAc4PUUpOFJ70qYLZz+FD42kdKnqMaMBS3XiokCkbv4qlH+qwep6VHtK5z1NUg2Idm7BzzmpcsWU9qahY5HWnFiIgBxVNiG3IwrA8+wrNKl3HGK1XTcQTWdIMkn0rSk+gnuW7YhcLjNXDhUHrVKzZWw3f0q71UVnV+IY7gjJ60wLyKdjJIozufFZDEPXmkYEPT855PUVG3L7uxoQgAAcY5phO7J7d6cV+YkdKGYDjr61QCE8560/JzyaYcCQ+lKQSxHrQMHGXBpJOPm/WlfBwAeRSSLuG2hCF6cdR60EYHJ4FIv8Q7gUxg0kJVTh6aQEhO4jB61E4CngfjSxghVHQik/i65p7BbQHUg5/SoSAOTU+cnnpUUmMYXvTiAqH957VOcbg1VwcEAcVOT8g9qJLUAyQ/zD8qUYySRSKSd2Rye9KOPcVADJfmbFNUjeBin+uOTTEPz7sVa2GMnULBz+NUVy4x6d8VfuOVxjIPWqMZCSbRz71rT+ElvUniByFHOKuoQibfTvVSIneD1NWl7nvWc9RgeIz/OonB2g5wO9THpk9KhYZAGcAVMQA7RnHbk0pcNFkD8KYzgnBWnBSSQelV6gI+WB4zTRGwj4FSqAVNKBx157UXsGhWVSrDcDk1NIWC4UZppIUkHk1IwHl855ptgVrfLv0wKuAknA69qhRdikAcGpUGMt3pSdw6CqAF5NNxlTinAADpSHPOOlQA1ccr3puAvPUmnKPlO480xyMgd+lUtwGsGfgcAdqgJPm7Rke9WT8khYn8Ka/y9Vzn3q0waGKCJ+TkYqQMwkI/hPQ1GfmJJOO1SKuIwCelNrQRPG2QPWkztJIFIh3FcDn0pzH5MZyKy6lDdvzE+tRsOAKmThfrUbqSvB4NNPURE0mXK9u1LDncxbofakOETdSKDGjMe/Sr6aAShioHvT1BxzUJ+ZAx6jtU6fMgJOKhgIOW2gU7duOBUbvj7tORSrEnrSaGKeTgdKQj5ic0inDMemKMbhkGgBGIIyO9NHqOtPDZxim4JIHbvTFYBkKc9aRBgEnqKGGOB+dAXCk+tMY4ElQ350ozjdTT9wDpTj94AdBSAXgZPpSZB46UrdCBzSLjHX60gGdFwOgqJ13tzUp5/CoiCMtnpVoBu4M2CPwqDcgcoByetSsQPmHTFVo8Mxcjaf51tFaXJJ4GAbbj5TVlcYPtVaOLYx+bJqyuAmO4rOe4DycL9aQZNAHycjvREc5z2qAEdMgEDpTDjCgckVJu5PpUbYzkdaaKFHzYx2NLkFaapAIPWheAcimIeTwMCkJwQ1AzgdxTsfL/KlsGpFjBJFIoPfqKVBjJPc9KOdxJ6VQC5LNSKCcjHNAUrznj+dKmA2Qc0ANH3MDqOlOyQmT+dIAB1PJpXXkDP4UB0GN8xH61I2ScDgCo1JdwT261J6kChhqIQOeOtQghE3N+FSEhVBzVW5BcgRnk9SKqCuxA8wXIXvTo8kBs9faq0W0MEIyPrVlwFKlcVo420C5LF0bPXtTHBIHPPanlCWXaeR2pTHuOcdO1Z31uMrJuVnLDI9KZIQVyRz6VYxnJU9+ar3CGRiFPStItNitZCqN8Qx26+9SIhKAnt1pinbhc5p65EZOcihgRyneFZf4e1PQ+coOMYpwCqRkdaapAk2gYx2pdAHop3HzPwpYggyQOM96RR+9yDx6VJtTnHIqWwRWnjDRszDgnimcpGoY4FT5EiZfhRVSeMmMsTyOhrSOujDzJsu0yH+EVbCBBuBNUYHB2hjzV4HeCOgqJ3QCyKGXBApGGBwfwpCQwxnpQTn61CGNlJIHHHrSooXC0jj58E5Apw+8KfQQuct9O1IW6kDg0HgHb949Kbn5cY5osMNxyalAymT3qNuMA9al6gAGkwBhlevNMZCQBnAHWpG5OR0puTlqSYDOhPHFGCzc9BQecUg5HNUA1iDle9OjYFMelI33vamocKx9afQQrZK4H4VEzZG0jpSsScHOD6VVk8yWUrn9KuMQYrOqE4GT6UxLzjJUD2BqBw+4p3HWrlrbKEwR+NatRS1Em2Ma4lZThOtVdkk2ck/wCNa6xqAAaao+ckioVRLZBYzhYEj7ufrUoseRwBn9avoCDjrS5Uvyc0nVkFiotmEbB6+tLdgLFkDNW+Mk56VWvMGE+/alGTclcTOdPG7IqVeQfbpUcmPNOSevapQAeRwB1r0HsN6ikEAdAaRg27IwcDpS5+XknrTThiee9JE7DwueeoNM27c+op4ZcbT+VNI4yKQFrTM+aewFbIGByMmsbTjicjqOtbIwy89q5K/wARcRyHJHarLnc2Bz61TRcuoHrVx8CTaB25NTBaNky3MLUiyzhgOaLeQyEhvvCpdZGNjAH61Wszh8D866FrC4kncvzBntiB2xiqSL8orRxugNUVwQAKyi9Gi3uDfKuQM/jTYyxGT3pzZCnJ6GmxncTn/wDXVdAHAHJ557VHMcQSDqNp/lUhILkevSmzriCT/dP8qcd0Fjt88elKB+fvSAnFGOM5rwyhe3IoIpAT1pR1560AIOOnejJAwKTknNKM5oEHpmnHpmkJ9BRnvSGAIPSlB7CkUBc+ppcDv1oEJjnPegHmjjd9aTnNACjrmkI5xR16U49R6UDG87valPpQB1o7kimIae+O1IKdzn2po4yKaAQE5+tOyQaaflx2NP6gGm2AkhxggZNOQ57UbeeRS4AqbgGBjpTGweKeeDnNRF/nwRwaEMSVPMxjjFMOVcMR061LgAEZ60fLtwe1WpNAODZ57U7IxTFOenQ05flOOtQxETCQS46L3qQHNKTg9KaAQ3Jp3urD3FbGeaRnCocDOKaGJPIqJ5AHxnpTUbi6EyP5g3YpSvB9+9MRdigg/L6VKw3DFDVmMAe+aOcnnigKAvPNKSBz61ACK2TinfjmmjI7/jQefpQxD8469qUYC5pMDuMj60cMeBQAqgBTkZpMjgUuSRSAcZ70gFycmmofnyaUc5NC9KBidWpwz6dKQZwKUZ5x+dDEIOvNKMk5NHGaQdcUDY49etNPB9vSgZ3ClxzQHUaB3zRnge1HQjjikzkEDt7UwGA5/On4yMGmqMc5p64yMU2AHJPSk9/epemT+tRDpnPJpbh1DIU571DJhjj0qYr826oXyPm4NUgY4bnJBXgUm5lY8cU1ZGlb2pQrAnOTVWtuAROdhbFOZiwU0xFYSsG+72qUD5QBQ7XuAuO/eheaRjhelC9RmoAlA5znigYzzyaASc0Dgcjn1qRA/wAwxTduBwOaRmCqcfgKRGYA07MY8jDZ7U3PNDvgfdqEH5CzHimk2gJWZmO1eCe9ORcEknNQb8Ddjmpg2Vxng02gFJAJbGaY0h3AY4pxPJFV54i3KtinFK+ohGuVDFUxkVNGSIsseT1rPnhW2VXGST2q+i74st1P6VcopRugQ8HaOe9DcYxQpzxSbdzBs4FZdRkw+7z1ph56jpS+v6UoAB44qQuMbGKYCA20jnvUpHzDNRMMucdTVJiFUZcgDrRkldpoGVYCgDkk/lQMe3AGKaR8oNKM5we/ShxgA5/CgCGRyPlxxUy5K4qFpMuc9BT1J5JPB6VTWgIXADA9aF5zjoajdyCijuafnnjpRZ2uIRj89MkIx7U/cXbgUw8Lz2oSBEan5yoNSKQCeaABjI701UAywGc9qrcZNgn/ABpBHsXjp6UE5IxSkFiMHioAMbVHHWmSDAyT17VLkbvX2pj5cewoTERMNqkqOT3pSu5FA6jmjjdk8e1IxOeRxWl2xiO37sBhye1VLuMyxEo2D6VaILOc9DUeNr8Dj0qouzuhWuZkShoggPPfinIGEgxwD3qUuqyEYAJPao5eJRgZC9664u4tNzRUFl8sH8akIdEUL+OajgcMoK8ketOd96YY8muWafMMjdiNxHJPaqbt97OQRUskgJEasBioWQrO8jHIA6VrCIhscbzBVb7p7etXIkSNGBIwOntUVpuwXB+UdBUkeHDBsKPenUk3oNDdu0E/eNSRt+64X5hTQAoXHTNSN8vCLweprNu6AoCEi83hs56ire8Bxhee59arSIY5W29W5Bqxb4kjDMdzLWktYpsF2LcYDDFJk7/l5FFqd2SRxTwNucDnpXO3ZjGRHBAPNBbrt9aEVtxJ49KMj5j0AovqIglADYANLt3Lz2pwLKrMuOOhNQ8uQ7nnritYxbExwjOx1Y9e9VdgMLCM4HrV1W3EAD5aY6K0vlgcCnGVmDRlykso8v7qnk+tTNGirxknFEsZLFR90dvWo4SXYtn5fSulO6JSFZWyozjHaneWdxZTSctknr1qxFbXBj3hSUqrOwEkE2QoPU8VdQjPsKykJSUjoB1rQQjYSa5q0fIpMmZ1wc1lXTNggJwepzV2WQYAxk1SZm37W5FRTVmDZVBBjGCR61PEgLBc801wGwFOMURNsfJ+Y9q6r6AX4yRgjoOtMdij8HFLGGLkk5OM49KXy8ncRzWGiY9SZSABkdajIK5IGTTlOcgUjMSQBWXUGRFskD9KU8CkZWwD6UEZAy1VuIrzE7QCe9VnX5iM8EVZkjB5Oaif5cY49K3gARIuNo5qQKy4UHikgJJANPY7CQvPvQ3rYAicBWUcn1qbAI/2vWolA2ZIqUH5Md6ylvcfQiODztyfWmgOoOBx6U8qw4IyKQNsG0DjvVXEydTggkUNkio4pM9SMCpRtznHFZtNMYhwuPU0o27s4zTQvJbtTu+AKQCOCQDjAppwAP1p0gOAvekbAHqaaENIGRk/hQemaCuVA70HGCCKYxjEKuTyTQ2Aoz36CkznjsKVuVGenrVIBkhXyyuTVWPgFh0NTNhkJx+HrVduVIx07VrBaElkybWVSMin+YEXOeTVbP7sZ6ipY5OCCPwpNWGRmd5JN+MD0rSgwYd2OaqBlxtVetWYQyRYJ5NRPYCyOSW9Kb0Ge9J/AAac3EYGeawGPRTs5PXtSEFsgcAUKcADNKQVB5znrU9QIlwFGTk0HcFJNHy7yo/H2pzAsQDVAEJBbce/TFP+9knpUMS7CSKmxtJyc5pStcNxm7gZ61HIu5+vIFSH5CV/KmsuQGxyOppoEVkZyxAPHc1LtBTB61G/C57k1Ig2qVbrWj2uA4kiPHoOtZ7sofaT1q9ztK1l4YShcdetXTV2J3L0KAMpzir56E9cVQiG2TDdKvHHPPFZVNxhjIBHWk58wHPWgcZFGSAO5rMBxyDgDPrUcnIAFSdSGph5HPUUINxrH5sdjTSckinE5XBHNKFXcBVAIeTgDrQP1FO759KaB82aQaAxUIT60ikkbqVicEAUiH5SDwaa2GPA+Umo1yWNSDhMVBLOIGGRkURTeiEScFs1GUA3PT1YsMjoaawO3aBVLQBqnALY47UhOTnHSnA5G0jjtUcjH14701uGiGqCHAPrx71aPQDsO1Ved4B59KtHnNEgB8k46E0HIbFKeSWHagdvWoGRufmwOlIFw3P5U8JhiT36U18AHPWquLYbISVPYCs/LZJHr1q8zfuWLZxVFiWxxgDpWtNCJ0Ow4PBNWY8hST+VU/MCkHFXh8x3Y6ipmMJD8o9Kbjn605iCAajdj5mCOO1SkwGsA2QDyKcOVGDTWyH4PXrTecfL0FUtgJcEOBmkjLFzxgChXJAIpplCvzikGggQLKTznuamcjbnPFVHdnYbRx3NWxgIDinJNWuGrIwR54NTKCQear5zKuBx3NWM5bPrUyGOzkfSmlgBkc460N8vAOM00jC4PJJ9KlIQ0Y3FhSN1BFSDO08VFyBjv71S1AZgs5PanSklTxxQnEgHrTZsmMjpmq6gRICSy9c9aniA2YJ9xUdvHiDPUnrUqLscqRweRVyfQEPjYAZ6ntTiMDnvTU++RTnG4/SsXuFwBB4prfc2gcilHzEN0oP3j+lAyItlthH40hIDBMZHrShd8m7uO1GMtu646irEJuKudx4qUcpz0qFXLqfl796cGJfAPAoaAc3IYjk+lOXdySaRjkcdxSqNoCnrU9AE6Pjt6UrZzjtRjJ6cmlZSXxmkMRT8pwKRfmBY9qcAN2MdKaoyGxTEGOckU0AnrTiMkUmcyHH50AKo3DB6ignilGQM9PWkYAH6daOo9AOFz6mjvilPNGSfwoAjcEA4PNQBgflJyTVkgnOfwqBgMnPAq4sLCEYbAHAqmOLja3TnjFXT9zIPX9aryjDBh+NawepLJ42DSYxxT1OWYYzzUUQ5DA9etPDjecHioa1GTKOCOwoBwTg0KflOO/egHK47GswFP8VMI4Jp2OMdhQOXP0oAiC4p3UYHWkHGcmlTlh6VQxSD0z0pSe2OBTXIDZz+NOBDE+hpCGN8v0puDg84pSQz47UOA6jHY1QbjlAVOfwpirtJI5p5GeO9IDiQAdBSQAFBfNKMlicc96QggM2efSlTPl5xg0MBgHUng1IoJUjpxTCAIyR1qVMeWTQwKkrgQHB+aq0TAyDnoMGpJ3QFwO9V43Ajzkk/St4rQTJiY9+V/Gop5WXBUdKZuwCf1ozv69qtRtuJ7GjDISgYjBqXgAtVO0L4IbpmrZ+914rCasyiELtBUHk9Ka/yxnB5ApxwXPPSmSgcelNbgQRuSAc896nR8jBFVjIN7IBVhU+XAPNaSBMVSC4XGR606XKnKnn19aWCMqSWHNSnggEVm3qAkUfUmg4X8akBJyPyqOQ/ujz0qL3YytPIsWQ3SqBu1J29Bjn2p124kX3qiMDI6+tdlOCtqQ3bY0YFDSAgZx3zWmoAG7HNZunD5ck5NaZHy47GsKz96xS11RGCGXIpAx3Afd9KF4j5GPWlbPHrUC3QfxmlyQ/P4UDhsmk3fNx3oHsKeaagwct2oHf2poGWAznNCQEgO7dUi8DnvULcHaDj2p4J4XPSk0BIScUxTjLH8qCx4IyaCelKwgyAM00g4pxB2+1NYjA96aHcZIDkZqMP8wFSNzmoIY8MTWi21EO3g5cDmoyjZBH41YC4X+lNLYAoT7Ayqv70k9z096uRDamO/rUSAAbV4qUDYuO9OTvoCQoOc54NMRgQfWlxnJFMI644zU2ESK21KWJAMknkVFGCAdx4zUka/MT60PQexImew5qC4OUPHXvVgkLyBg1SvJlK7QcZogrsXkYEn+tYY96fGQVIqNz859KdFjOcYNem9htaWJSQeBUeRgg8E8VKyAkHPI7Uw8nGM/WpVibijIGGP5UDuMk0KPWlB5xjnpQBLYEC6z05rcB49KwrU4uhzW+p+XI7VyYj4homtApl5xxTmO6cmoI3wQRUikb/AK9ayv7tga1KOrAeUD096zrQ/N1GK09WUtERg4rHtiySgY4rpp/ATubaHMZHtVIrtY4P0q5DkJ17d6ptw5z+FZxvdou+gpxjk8ilGMZNBAOaaSVXFMNhynJJz+NRXP8Ax7yeyn+VPK748Y61DIrC3kzyAp/lVRWoandKMLj9aOnvS44FHbjpXhjEGc5pcHNGMCkOcUgA9cUuOlITgZFJyenFNDHUCkA45PNO6L70CDrQaAc4NHWkAdR0pOeaXkdqTnmgAUY5pW5NAORS8Y+tADccdKFGByaMZH9KUdB7UAJ3Ipp4zinEYphOMVSADyB2pwYdKTqpoXO7rkU7aAPoxjrQOvNB96gBjHnI55pHbcMgYpXUg5B/CmqrHqeapbDFyc+1OGOaQEgikGWBp2uDFDccCn4FMU4JFPB9ulJgJ7ZphCuSAeaczKrc0DBJIHNCEMZTt46imSICASBmpiSM4qGSYAbiKpXBkkbZGMfLT+AD6VFGzNGSDnPapVJ2ZY49aUlZjFPIprE9jxTzzwMU3OQRUIBvAb+lSDntUaL0qQHLE1TYgHPanfxHAxQCRQOxqAAZOaXvik/iIpQOc9aAE+Ygn9KQEgEGlI4IBpO+O1AwIOKXgAUmQx6YFBOSKBDgPmpAOetGefSkzg0AKT83ApCdvpSnge9Qv90Z9aErhqOcnIK07GOc/hRwBjtTj93pVJ2GRHPORTxwlIVyRzgenrR97gmkG49eQCenpTM9+9KOSRzxS47le1CAickHrxTXJAwB+NSlckelJjcM9qpMCEFo0680qlnX72CetNMbGYnPFIPlfg8+lXuA6EBmzu5BqVnCnrmq6MnmbAeh5q0oAGMcmlJdRDCxJ5p4GPrTSAD70pOTioGOUH14pxAwOOe9NBpxztJXgil1EJJz0GPWoQWJOCMVN1BP50zGVyDinF9xkbEshA4phUeXt/OpDECh+bAqBpFQ4JwTVx8hCK5w2eQvarEZ3RcdxVXzApB2dTirUX3CAaqS0HuSKSo+aonzu61IAcVDMDg+4rOK1AryyGQY7D1qzbv5inaeO1UPJxG0jHPoKvWS7Yl7A9B6VtNJQ0BE3OeaceMEdqAOWNKQOwrnAAR978qUcksfwo6rwKQZ28nikAfxEnp2qGQqjhu+KmUcHFU7pwoAz9auKu7CJVYj396enDe9VrZy/O7IHFWhjeDVSjZ2GOXO/NDjJxSA5cAmlyN2MVn1ApzQtI64ONtSWzBlw1S4yxqJDhyAOPWtOZtWEPKjdu44oJ5yT96lfG3GaUjr61NwEwVIOeaGUFgOuOtKeRk9qQfcP86B9BhX16UKoJwen8qVwNuaFOMkd6fQB5+YZoQbQec0p5OOwpgG5jzSAci+nekY4AxQMlvpTsYOaAICwByTx705gSck/QVG7bpACMg0pfMgCj5T3q7MQSIQRjoO1Rjg5PPoKnb5s4qqw+UoefanHXcDPuDliQ30p4ClfoKrTYjkCL0HY1NbkCJlPXFd6i1AlE9rKVITH1ppkZrojHyjvSWpZ2bDYx7VOqOzBSP+BVk7RlcdiDZ5l0B0/CleEfaGUOQO9SMwE4APPrSCLbO2Tz2ou7XAdbqnlNEO3eqyqFnaJzmrkMY8xlI2ke/Wm3ACTKQM5qVL3rdwaHyMFRdophzMoZT04xTtuRnoaSJVYMityOtKNlcfXUrTghlA+53HrU0BVx+7bKmoCrqkmW47GlthsBVTz3rRr3bIXUvROy/KBU53KpJqBSQqmpW++DniueW5Q3zC+QGxj2oDqVx1NI6gNz/F1NRK6RhiDRGNxMlDbfmI+Ufw+tVwdzFyuBngUnmoSCx2j+dPwZMnPyjpWiVhEi7gr45Pamy5WDIHzEcgUrgs/Bx61DKzLMuAQo79aSV5A2VplwFx171XRCGyatbWMhOM5qMAyZA4xXXBohmjoixGciZck9CegrcmRDb714x0FZmhRbzM7jjAArWlUBdo6V1wWmplUepnCCMOzlAWNAjjDFimSRj6VMw5qM96ThF7iUnYhW2iBJIzk0z7LEA2eSentUxph54zS9lDsUpMqPp0bIArYOeWNNGmKkmRJ8qjjIq+vCYpDwMil7OIczKYtpQVk3D6YqR4PmA3cnrUxcjGKaG5yal4eD6D52RrCxwAQB0pGt2ZRg9TjrUpbsKAxJxU/VYBzsgNtIJMFeBVds7sCt+3hxGWbqemarPZxsxyv5GolhrfCUqmtmYz7i30qCRPunGfWtt7KJg23IyKpz6e4XMbBsVHspxHzJmfCx8zgYHpU20c5OfSnLbywsfMjII74p74AOeKmaaZSIVAAwKkRfmyTUeAuGBOfep0HfA6VD2GthhZjg44qMgFTvOPap2H4D1qKQ9Oc1KBjNwEgXG5T3qyCqqD2NRgZA45FSZDdRgCm/IYOCycGnKA3TtSFcLxxQpOwepqOgCnPBApkgBOSakB4yeKZjJOR0pIBgPU5oweop2MjigYOBVCuRk4JH6UhGFHf2pT98jP0oQFnPPHYelVcZDIMjAX8c1XZRvxVqRjuZU/Oq7uQSAPm9T2rWF7EsawAxg8U4cOSDTHHyZJ5pgchyFTNWlcC8CWIweO9T+xNVICFXJqyvXBPHrWE1Yolz0DGnswyBnNVwQXxnNKzHnFZ8oE6PlzxwKm6k1ViJAyTUwYs+c8VEkA9QCTximsOjfrRkBcinEjYBjtSAauevWnsfUU1SQmCOacQSvvSe4DXBJPr2pnO0AHGetPJJIb9aY6jr37VSAru25ioHIoRw3zHtTZ8xjPUnqaZEWOQRx/OtlFWuBNBKHcnOc1UmU/aGK/gKtW6+WxGOO1VJWP2gnsegqofH7ouhZt3L53D5/SrqZKc9qz4pPmVW4YdKvrxgms6idxijls9qVh8xHQGh8ZA9KHHHHSsQBeeQaa3fsPSnINgGetI4AAo6gRg4XPfPWnDBO/1pgIbKnt3pyHOQR9KuzAcCCD6GheGIz+FGRuwKQnLCpAXgEVGoO7mpDkyYHamDJc/rTQD88E+lQvCkku88sv6VMpHPvTQoV/rQnYBGHyhhzSAkqW7UpY7sAYHpSAfIRjj3oATGGyaYxDOOODTh1z+VB4wq8mqQELkhs9qsqB5fXmojyxGcnsKkTlBmnLYFYXOAKTnYSB1pW6ClPYCoHuIrZ4z0pWG/JajAUfhRnAIPej0AifLg9h6VnS5LYORjvWizgKB3NZsoYz5ycHtXRS3E9iQHaVLd6uKucYOarAAldwzVmNgxwvUVMwWw6T7hAOKQfMuSOae3zHA6U0fTj1qFsBEWAX3PSkjUDI9ae/3TzRE3ylyKroHUgusxR4B4FQQxtIu4jipLhGeQE8A1OCEUAfdFWnaIhAdsYU9amzlcVA8a5z1zU6NvTgdKiXcZDIxThRx3NWF+VRzmqsrPuCoMrVheRgDpRJaBbUU4J+lI7HrtyKfjjA70jDKAenaougBMtnI6Uw4zup65JAzTCckjt6U+oDFOw7z2pkpO0kninHlSMZIpHyBs28iqW4MUBtg28DvTkDdc5xTIy3CmpBuCHufSnJgOU7sk9afkBN3eo4+CSeh7U/AyWxxWbGIBwaXHHvSDccU5sYbFAiIgr82eTTTwCfXtUrAFR60wgBuRVJ3C4wk4GO/akBLNlenpS7yJTGRx60EbXJ6ZHIqgJSQwGBxQpG801xtQbeM0gYByPyqLAObAcGnE4YsOaaDjAbrTgSSMigegg+9nOaBk5ApADuOaU54A60CsKcKQKZnLYxSsCDnGTTQ3z8jBoSAf1PvRIeRikxhqcSF4oGIPU00Ek+xpUw2T09KaCdp9+tAhA2Tx0qKUHYeOKmxtamvjBHc1S3GRIFMZPpUEg2OzdmqYMq5TuagkLfKoNaxvcklgIQEEcntUxAV8YqAZj2nrVk5Kg9MVMt7jHcBvrTcnoPwoUYY+goY5IxwBUAHOcA/WlK84HXvRg8N3/nSnqTQOxEMb29Keo+UgUzgSYp6gq+KbEIx+XPp2pI8+V6GhyMZ6AUL0HPNHQBCPlHqaRMcnNEuVlXb0ocfK3vT6DHRA7ueaRTtYk9O1LFkIW9KTJPvS6iFHVix4NOzhcU0jcfpSj5sn06UMYpA2AZx6UEjbims3Q0gZicY46ZosIpjG098mq0mSSAe/arOfLkKA1C0ZVycV1RRO5AxO7O3g+9OD8EfpRtxjPpTowB0XHY1pdByl615jzUy/cB7jtTYV2KB2p/cgVxyd2URMAWKmq96PlEanBPWrDf6zIqC6ciUELx3q4boTKcMWxuuQaujLfdGPQ1EigLnNSw7iM/w1pN31Yo7FiM569qfuOSRzUSNkFhUq4A9zWDKuPUg5NVroYtyAcH1FWOEO3saguQDCR3pR3QmY5YF8jmnxpvyPT2pAm1mwMkdqfFlWJ9a7XtoSm+pato/LXjp/Orh5I4qrbMG5A47VbU8ZrlqXuWNYbhxwaauCc1Jj5yajUEZGc1KARj69aMkNkUg5b2pQMg8cVQCYznj9aVSAy8c05eVOOM0mcHApABwAWPWk4VOOppMHZ835U1nxcADpina4Eyn5cdxQx5qsGKydeaskYwOuaTVhCZynH5UEcinBQoIzSA5P0pDItnBJ70IuO2KlOSuMUg5b2FO4hAOOenao2CupZutS7SwNR7d3y00AxQqZBPNKWHJxn0pmzcxz2pxI2gCqAAxOT696QEkUh5OB2pSSOBQITrwelToTx6dqhUcgU6S4SFctQ7vRDGzTFWY449aybuXnGSf6U+5vC+T1PpWeWZiScmuqlTtqxCEZHp/KiMnPtS5OemR29KaGwfUe1dI+hYG4ncemabnk8fTmnb+AOM46elMJAPPHvUIkcOvH5elJ8uPcU3jIAxjuacFHT05oBjo+LhDnPPaugTJXPcVzwA3qc4+natiOR5EABxzWFaLdgW5bX0Ip6/epi8VKDzjFcfUpkGp/PasOntWHCSGHHHoa37uINbPg4rnUfEmO1dlL4WiHc3Lc5FVpF/eE84qxbnK8VFICsjH86y2kaEfYEDOaBjAGORQcZFOHX6VQtxpBUYPQ0yb/j2k/3D/Kpe+cnnpUc4Atpcf3D/ACpxeqKsdoCDTgKaEx0pwOWrxGIKXkn6UnSg57UAB5NICM0GgA4zQAD5ue1ONNB5pelJhuKOv0pR1z1pODjj60vQfSgBDwe+KPrR9O9AySaAFGOfT0pB0OeaB0pfbtQAhFHQCj1J6ULg8GgBkpJXHrTdpyPepMZJqLG0hfWqQDugA7U5R1NBAHQ0Lk5xQwQ4cNyKM8mgAA80MeeKkNgOR1puCGHPFKVJOc8elJn1pjEP3uvShT83HQ0rKDjmm4yMCmgHEY60oIHNJnk0BlYZ7UMBSvzZoBw3XNLx600rnGO1CEKx5PFVbhS5CgYB71aOecDNQPLtcDbVQ30Abbtt/d/xd6snHAzUS7RJk8ZqY4HOAaJu7AAOp6VFEwkQgdKmzwRxUUcaxoVHIqVsHUeny08ADJpincDinr6VLAOlB54HbtSYpwP5YpDFzgYFIKU4z70dutIAxwab0HWnDgEU3GcjPFMBQT9KQcMTTHYrgKM1IvcU2uohQSaGA4Hp3oXqaRu4PSp6jBsEGkOAPeg84GOKNu7NUALgdOtHUAZoHTmlFIBOuaaqqAATStjb71H820euaqIE/wBKbgngULnvS8luKnqFhpJzig9gKXjp3pD0JpgQkYfkU0KM7sc1LwWx+dR8lDgc9qtMCJY0Mm8KQatKexFQK+MgdR1qSNspnuaqdwHElmyOnrSD5pMnpQwbaDSMoOBuwf51CQEobnHrTs5AwOaZtyRzTk5YY61IDdxJI7UgG7jOBT8Dmm8A5poBrlQh54FRSIjBWwM08lS20Ammq2e3Aq1oA0bQdpWnx52nd+FN2s7ZBp+TwPShsBxOSOuKa4+UkA4qOTzPNHoOtSkkjIHFK1tQuVLYuUIfAqRZ8TCPpnvULQSC6MgPy+lTiFWkVj1HStnyt3YlexayCOlHP0pSM55pFySO4rmGOAxxTfm5704YJHNJ/exyaXUBN21Tk1TuIxKhIHOat7VPHWkCjBBFWnZ3QFWFPLhK471ZjJAHr61BI2GCY696mKhcAdxVS11AcBjBJ5peAw4+tDY3AYpTgPnnFQBGcgs3Y0igAk9TjtQ6kg5HWmRknJYY9qpLS4h7fdOBzRGSVBPpTs4Bx3pI+Bz37UX0AR844pyjIOe9G0Z5NGfm+tSMGGGwKTB3ADpT2IHFMAIPuaaARmCnb1zS8dB/DQqjfkjJ7U7AGQOtABkjpS/wc0gwBzQAGQ7uhpMCA42lsUq5WMnGT2FNdlEeB90UisXVm6itLOwug9Gxz2qu7+WxY+vSmnKxmTOc1CQWwx4FaxhqBnyStJdMzqF9qnjOwb24B7VBNHKrgsw2ntU20mDrzXXpy3RK00JrXJdzjGauKx2ZQfMKpR53cdAKuodynHWueruVsVZoypVgPmJqcJiQsfu4oPMZwMMOhJqOBpNjl2yp9BRe6ESo7OpIXDUsyYIz1NEBWM57e9TOdzcCs27O4ytbMCXU8EHinFAoLqOT3pyR4uN5+93FE2SwGKd7vQCDaXgKSH6HFAi8n95GMkjrTbto/L8vf9RUcBUxBc/dNaq7VxX1sWcSMAGOOefapXbfIEUnAqBn8xwB92rRRQykMMkVm9FqOwyUMGx37VmXRkjbbyc88VrMpDAnk1SvCcsq9T1p0pWYpLQoQM1xLgdB1rUIIGF5xTbS2FvAT/E3OanhKkfKMj1q6s1ey2BEEoYxhOd1OmBaLPcdTUxjZkDZ+bvUbkBMZyfWoUr6AVjtjBJJNMYBU3HjdU77I8DrUOUZunH1raDbdxSOi0hAlmXX+I9asSN6U2zGyyiBGMjNJIccZr0YK0TmnrIgbjrUDMMladLKAQo6momUFc5zihsEGSeKToaTNHvjNIYuTnNNOSee1O7YprfpTQCGmnmkJ456etRvMiDlqNxNpbkoq1bRZYMelZsdwrOByavi+jiUDFUosh1YJbmiWCrUJ5bNUzfoxGe9TLdR8fN+lOzEpw7kpAprKKUOGHBpSBjPepsXchaMZ96RokZfmUHPtUzdfSm9c+vYVLimirsqtZwsc7SB9aYbLGAH+tXewphHI4qHSiylNmbPaysDhTj2qEQMowy9K2QKQrkEf0rL6vZWTK9p3MdFGMEY9qlQALjrWi0UZfeUGcYOKZ9nTcOSAetZSw8+haqIoN6ZpQvTnirJtnJ3DBJ7ColhkQHcpA6CsHCSWqKunsRsCwHPSlYnaPeo1cFip60/GRgHj0qWmtwG5IFAA65pG9KB9zJotoO5G4+YnPNRpKGYhRnPenynA9jUMK5OQMCtYK6EEhG7C/eqI4Mh5+vvUrhfMIHWmqoY9CKtOwbkcrAgkDBpkW9mIJ59aeQWySe/X0oQKzgqfrmtEwLEaNgDsasAEHB/CmKMLn9KeBxmuaTuMaMA47mlxwT29aRRls9qVjlcDOKV9Q6CKMKcmpQ5Ax7VApz9e9SjJ5NJoBwbsDz3p+c5z+FR4GQfzpvQ5qbXAtDvyMU7llNVUY5qXf8AL0qGgHD7u305qNm2gEdT3oRgST2oX5scDAqgIpcrG2ePQ1VKlFABOT1q1csrMCegqpK2cLuz7CtoarYTNCMBYQSOtUJYCs3mE9elXgu2BV9qiumIQYGamLtLQbFVRkAgcjip0O4bR2qFCOCDyBU8YwS/6UpgSH7pHekI+X3pRzkk03GUOT0rABwO4KcUh5U8UDGBinOM8E80AV2yQc9+lOUgJnPNI5OSR2pFTCGtOgDxwMmkJyM96D0we9JgYwelIRKRgZ7mo8c+3enMCVAFIR3BwO9JDBPlc9xQ2S3PNC/cPrQ3C5A5o6jE6g54NNXhcDmlzlBinDG4e1MRET82O1KpBbPb1pGxuJAzR/DgdTVAA4OTTl6Y6c0yTJAxTo28wAZ+tD2uMkI+cc8U0Dc30obJwfSlYYGFqRahjP1FAPHP4UZ+YAd+tJigNiJsAHIqhd5BHatIgndgc1SmQ+U2/gH1ram1cHsRwOQPm5q1EcEn7oPSqluV3kBs+1WAuUzkkk8VU1qCLTg7Bz1pr/cA75604fdGaiHB+Y8VihsUrhQfzo48sKvSj+LmhCAMDpVAMkjXHTJFQRK7SFs/LngVZC7ZAx60bPmJHSqUrIXUhuJCu1R1NWIcKuMc96gkXC7gMsaWFs4HOe9D+EENk4Y56A1Ojj1qKdWZCE/E1JEpMQJ5ND1iBIM7DzyKB64pQwwcdaQAFcGsxjM4fAFHSWkZRt59aGPyE9zTF1GAYJc9KVhl8+tJuIjIx+dKjMRyMelU+4xpiBbG45p64UbvWmhlExIPzUOckZ6GnZsSsOY/PjHPapS2Wx2qPKiRV64p54kwO9QxgM4yaP4M0pORgetIeBgcgUgG4zyKBgAHGTQM9M0u0AYzwaYCKo8wsOaYRuzuGMd6lQbSB+dR7ck01uIduOzaBTXyQOnFKVIUgdaRRnKfxH1oGPZhgAcmlLZHtTVGAaQdTxxSAcOjEDrSr2J6mmk5PFOHXmkAnUgfrRIAW560D7+BTW6ZPXNNCDocdzTnIBx+tNY8gn8adkBQaAEHQY60jcDHT1pyAAc9aQgb6AG8Bgc01gGGc8ilIA/pSsc/jTAhPbA6d6qkMX9wauZ+Xr0qFsibOK1iw3JM5Ycjb2qZsAgdKrQna+G4z0q0Qc8nj1qJaMEMzjFJgsPpTyMOB3pFBDEGkAqklPYUhBP5dKM4Kj1py5zgnNIe+5ERjGOvenHuO+aOkgHekXqfamIJR8wUUADAb0p3V+elNByzDtQA1hkikyGfHY9KG4Qk9V6UKSy5qhBGMoVPanD5cjFEf3cn86RcjcD+dDHYXggnrmhAFUnOc0AfLj8xSqNqgUgEdd6AgYHpSKxVsHpUn8PFQEnzB7daa1CwlxbhmDdKjNsdwG75fWrRO5aTK+WCD8opqbSsDIBaLtHXPepBAMdOlSRtgEHv3pQSMCk5MfQaBnHtScmTIHFOYYPt6VFu4LZyaFqITIVzznNRyrgfN0FSYXOAPxpsnMbZqluGpREbB8knAHFWI8shU8HtUbMhVMnjmnxsVXgc1s7tC2FRmyq4yBxVpHG/DdulV97GXGOKsbABnvWUhoUn5iarzygDnmrBAIPHOKzLrJkxngUU0mw6FcSZdiBz6U8Z2nacGmKACSVAx3FOUEjk8+3autkFuBtsSkDk1cBOzA61n2pLkpnge1XlIBA7iuaotSrjz8rYPekBwxHtSkZA55pv8RJHFZjuM5GVFAGMCkzngnB9qcw3Y5xiqAco2p9aMEAM/FIzgEDPNKw3oFzSDTYYpZ356elQyHEox1zU7AlhjjHeoDHulBzwauIhcDzt7cCrOcnOOvSoWKh8MKnA/KpkPqByc0D7pz1pVY0mBg+h7VAmN59aeq44HSj3xS4wMUNjFXAU+9QOwByP5VMzDO0dargkMQTyegpxQmRH7+c4pcHaTTHBL7j0pwOIyTwK2toAMeR7UnA5qvLeoh2gZNVpbkzEqDgemKuNOTFdFyO6UlsDkdKoXV1vZjnJPWmNJ5Y2g9epqqx3HnHAreFNJ3Ha4rNltx6mmkggmgFc9Mg009Mda3CwoGBzzR3wAMmkGPf6UqnnpzQBYUZyc0w48w+9KARg54pOScgcdsioJEYZbjoPSlAIPPNKOOcUL6dMcYpiEJ+ZcYIzWraZ3Vkv147etaliwZgfapkroG7WNFc7uTTwQT9KYnysad905FecaMllJMBGO1cq7BZjg4rqz80JHtXKXC4mYdSa68PrcjqbVky+UKLk/N6k1DaAiNeKsXCnKms3pIvoQgcDccfSlHUZpxC8CmcE570XuAueQVBxUVyP3Mv+4f5VKDnofrUM5P2aQnqVP8qqPxA2duO9OBzTSSM4pRyK8MAHGaDwOKM8cGjvzQAnelpDyaXnmmAoGTQOaB160oGOtIYnQYxQR6U7HekHJNAhCcHmjHy9eKOvejGRQADjOaPwozkgYpw9AaAG+1HApSc9uaTsfWmAZ70xhup4PFRsp3Dn60ILDs5oHDUhz2o54x+Oaq3cB5Pemr1J702RmBXAyKVCG5FKwDzwKTIJye1OwKYcgGpQCHAOM5oHJ60mcoDjFMjfOQ3bvV2AlK80KB0I5o5pw4AJqRiYw+TTSepH40pIB5FNcb12jgmmkA8kjJqPy9zfMOakXgYIpMgnPSjbYNBGQEjH5048AZpvAc5NOPI9MUgDbzk0rYxjHWlX1xSHJ570gCM+tKOme9NxtwB+NOXPTFN2EGfXrTs4wKb3pwOSBipYxWHekPTNDEY7807jb05NIQhPy0DAPSg/dxSk/KKBjcZJpuGI+WpMfLSDgnFO4CD06U1yQCQMmnc5x60Ec4ouIF+7Rgn8aDnj0ocHd/8AXoGIAT3ozyTSH5aUkdB1oAQnpxUL438d6mI5471EwJYHFVHcCbOQOaXO0Co4j8uT1qQ84PT1qWtQAdc5ozkig84I6CkzwaYEbIN5IPNIBjv07U4r1I600RjBI707gRJuDsSOKWPJJPRRUv3QOcUbVXgVVwETDE46UmxS2T1FSYIXjvSEbWznil6AKvLY9Kk74FRK2c8Yx1qQc55+tSwEPHyimNyCKcT84FIOW68UARkHbx+dCn5Md6czHHHWkHAxnmq6AKowCWpFABLHkU8rlRk800NzyMYpJgKQGNNkG5CBxinZ+UMOnrSkjGeoNF7CK+H2bRzjrUkS4Tjt3pcBl5GAaVeAFFU5XGPwSPp1pRwDk4oBxnNAHzc1AAuAaByTijOc0o4HPekA0DaT3zS8g47mkB5+lAAL57UwIZVAJbviiNg4DHp6Us3CkkcCmQybgMAYPatErxBE/G4cUZy9MJx83XFCOCT6VFhCyAlxUZHP1qXuajUZfAGTTQxR/q8E9abECAT2HSnJnkE0ozuIP50xDQ2H9zTo8liT+dRsD5mN3DdBThkEj8qLDJCFPPXFIDnqcCnAbevBpme/r0pAMIPnbqlAHT9agc/MMVLhuD0ptaAOUbm6/LSkgcAUigluDwKVSc1LXUCvIPmCHG2lTaFKKOgqVx83AwKhi+YE7vpxWl7oCJYz5ZVxj2pjhgOgx2zVh4yW24w1RvyoPcdKtPqIz5lDyEk5A7YqAyEkjpx0qa4faSoGSetQIpUhv4e9dcNhEttGUyrZyavxfJwD15qs8gG0Y69TVsNhM4ye1Z1HfcpMiEe9yzHHqKSLC5VTx2p+0jfg4B6UxVRU3KcsahO4hZFIdT0P86ldXJGOg61FJuUoSMgmp9h89jklSOnpSlawhGcDBxk96r327amw/X3qeBWX5ScnPWknHy5GCQetSmlJBujJ5djGQQT1Jq1bRIu+MNuPekZGZl2jB9alhAJZ1PzDrXU3oCJI1zH8w6VPHGcIScAfrVW3YnGW6Vei/eICK55pq4x7RjJNZ0h/0hhjPIrQXcoYH8KzpeLjd3qaO9gZcZdq46miLuoP1oR9y7h1pqgg4Y8Ut2wHDIibK9KqqP3fzjBParZO2Mjv6VXnCvtOM4px3sIaFV8EjJFMUK86gDj0FSsmIlHQntmn2UYN1DEeTnJroo25yWdBjZGqjsKpTSbQatTPgVkzybj7V6j0RzbsiLFpD61NnP0qsHVDz1Paplbcv1rPqVYdgZJozSdBmqs1ztOEGTVWJclHUtbgBkniq0t2q5C/Mapu8knLHk9qbk7jk81SS6mE6z+ySPLI5PPFMaPOAOlPQZbI/E05cMSdv05pmLlfdhGmApqXcCOtMLYb5RkCnLz15Iq7aamd+qFVWAII4NHBxg8+uKc3UYPSmcHODgelMhilmBxuI96nS7eM4J3Yqux24x/+qmngDPU9qNy4zcXoaSXitweD6VPuBHFY4z1p8c7Rng5qXE3jXa+I1Q34UdRmq8NwsuBn5qmBOTgVnY6VJSWg/vS4I+tN6fWnA8U9yuohHNAGenNHelH60rAN2cnmkAOfpTwc9aMZoaC7InhSRSrIDk54FNNtGVUYwR71PxmglSOnNRKnF7opSZTNopXhhuHrUQspckgZ56etXgpJ9qtBcp71n9Xg9ilUaOeuIHSM7gQQewqrECAa6VkI4wD9aryWsUmcpg5zkVHsHHYpVEzIOB2piyKxIxtPpWhLp+7JRx16Gs82N2D5mzCg84rJUmk7lcyZWdGBLbuPpTI2beMLx3q6ysF2laiKESDDYovYZYAG3Hc08gDFNUjdj1/SnMOpNYNWKIlO2Rs9KVpAvJFRbi2Rnn0pWXcMdfWr5epI1HjDnaM5681MZAAAKjt0VS64wakKgL60pWuNDg4bjsKVSGOT0pin5TkYpFTC59aloY8EHdz+FODEgBaRUG3BFOA2/U1OgDTkDAGKXe3Ao5xz0pjNgcLTAhmkKY2/nUCxjf5o9e9STAk4BqJW2zAZ4HSt4aLQWhqK2AKgugSoxzmlVs9e3elk+ZBtPSsFpK4bkNsxG7cKuhwoYk8VWjYBuRkYp8b7lyRxTmuo1oWFbcoYdDTiPl96arHkfw0ofH0rFgOHJz/D6UpIC5pOARml6DB6mkBWZvmOenanKSyNt601+CwPXFOjIEanHLVp0Adg4Bbt+tGf/rUuMIc9aaQcjmp3C5IGOwcfWoz8pIzmnMwAJ9Kp2szSbwR34OacY3u0IuKcbuc5pWHBHemqMOPXFOfuemKnqMYAelInD884p2c856dqQA59M0wIySSSBTlXuOlDZIIH40oHy4z0qugEZHBGafEMDk8mmnBGR2p3VgRQ9hjj1xS8bvalJyeByKaPepEKBk5HagcA4oAAYjpSdFI9aAAErniqlwNwYEfQVcYjcOaqzqX+ZTg1cHZ6g9TMViHwvBPatGNtsZPf3qikeJPl5PerkI8wkFutdFSzQix1Qn17U0AYCv37UoXBBB4Wkcgnn8KwQwfaueOtImFOD1NKxDEE/hTWPJbHBoAUuGbBHNKHyD13elNx+7JxSGQDap+8elUvILCyOqoGApsK8FqGh+fJbPtT8+W4UdKL2WgCufkYjvRG2YgFHSllOEHGfao0k8tOeppJXQmTrkA+vrRHk9RRkFAc4pScEY9KgoRgNp3dKiyShx68U9+Dg/xVFI3ygCqiIBkk8c0xVZwHJ5HQU5mwc5o2/OvHyDrVJgxhKmY5HI7VIwJVdpwaVlALN3poYYAPboad77B0FIJcMp5XrU6oCN3eo+RLwMCnAEOf7tQ72GxwA5pvIBNO3ZJ9O1ByBj1qQGAAlf50uMjFHOMAdKUcArTAap6FeT60SbVUYP1pIvkYgjGf1pSBg7hwelHUQirg+uRUUKsrEu3TpUoRgwOeKaRtkI65qkwJBgoTjFNU5XApx+VguKB8oII5qRiFh070Y+QYpcYO4jimqxKgdTmgQ77rdKTGSc0KcnPpSqOc9qNgGOGZBtHIpc5jO7qKUkqD2xSZwAQOtMBQQVFBPJ5pI+eT1o/iz6UWAQKCDmgDIxmlPr2obHH6UAR525yMVSuBuJYHk1ck+7x1HWq0y7gPWtYaMGTRYKBiMYqwx3IpFV1x5GM5I61NwiqFqJLUEO/izQp+Y8ZpOPM9qM4b61IDR8zClywXJNIvUnuO1Ky+vemMb3J70EbRkHk0MCR7ClzhB6CgQjZ4PpSb8HC/xd6cMlcdqNq//XpgMkxtzTc8DApzsCB2qM5Xp0NNbAPP3AV605tzoNvBpCpJBUcUr528cUDAg7Bjr3oT5iR+VIcjBB601RmQgcUCsTLkMR+tROD0HJNSKMHHtUbKOScjFJbgO3bQFIz60oG0n071ECXkx0AqTOZDg5FNoAT5myvQU4HLGo0woYg0m4oM+veiwEucnOOtRgfOQelPJxyKaSd3FJDuQlm8x88AdKjkO+E4ODU0uTGGxz3qtuO0FRwe9axFsRAYjXnFP3AxEA4YdDTTtz0xmnRqocjsa1YtSWJj5LEjoamjbfHuI5pGUBQqjIPWlHBwPyrF6jH8MvT9aoXPJA/Sr6r1HQVTuRiQH1zRT3DoUAu38KfECACPm4ok5Ye3akVsc45HeurdEOyRas9zKWPUVbVwFB7+tVbc7RgnOetPL/McVhNXZWxYJzgDrSD5vlY9aYHPLAdKapJBJ5YVFgHqQXI6YoZi2QBTGQtg9zUqqVGTz60aAMaMKwZuTTwMnPSkLgkMRx6UgJDFc0asegn3mYA4A7ClA6MBxUZYeYQvUdaljYlOKbVkAqkGTpSjLHBNMPBBPWhWfdwOD3qbAyUkkUp4x3pgBA570obAz3pAOYD0x9KcByKjXqCev8qlVjnpmkwFIOTjpUL4BIxn3qZ2CD5u9QMyjoecURAqStkUEExtkc0rxBn680pBCE56dK3ugMWX7xIwKVMKMseBTW+dzuHAP50x2D57L3rttfQzVhrtvbJ49qjzlsk/hSk5YcUFQfu/rWi0LStoJjJ60dB7+1ABwR0oH5UAG7DZFKAQR3+lM55zUqfMMgfWhg2PwG+npQQcnHrTlHzcLx0zQxIJJqOpLQ04yM0Y496U+9Gflz1NMBhHHJ4Hb0rRtQdgPaqLABcY61o2gL2+AOO5qXK2pL1NFHBAI6U4DINRwnK9OalBz+FefJWk0XuidOEx04rmLzH2xx/DnrXTRjK8Vz+ortvGwK6aL1Ey3Z8oOanu+Qp7d6p2D5GRV24bbCCevvUSVpFWuiuRuyMmmDBYZp24UDGelNAH3ff3qOfItpR/sGnbix9MUk5H2aX/AHD/ACprdAtTtieeKD2AowAeKMYya8QYdQfal5xRkZx3NH8qQCdvSjp3+tKR70h6UAKaATnvR70u7igYf5NLgCm9SBR25NAhR1PpSZIbnpTuKTvSAON1NHBIxTsZIo6GmAdB70g6nP4U49KTtQMaTzim5wORTj1H8qU4piG5wenWgk5AIpMHNBByM07BcJM8EHilGQRSkEDnmn44pX0GRY+Yc80FsHDU8qDzTCoJ55NCYhjDyx1+lPK5XPY02UZXB/OhB0I5AqugDgRtPP405Pf8qQAbfQ0o6+9SMD1welR8LL1qbIOeM1UuFbeG7DtVQ1dgLRx61G4Zl4oRg8fPWkLDICnnvRZp2AcBz6cU8YIwaZ1OR0HWlAzz39qkBVIJxnBp/QDt+FR4w+ehp55YDP40ARnG8Z6VKvcd6bhc0pbGMCh6iDuadkHk0h7U4A54FSAAD8qUdfakA7saD1HrSaGAOTSsT3FAxnijOe3SgQY70gyDxSjpQDxkCgBMHk0hz2p3amgZzQMXcCfWk75pw9aQYwT60AMHJ4FDdyo5p465xTQcdqYDYS2wF+GoA60/PYfhSA8bRTb1uHQaqEH2qTBLdaBxnH50Dqam4Ib2AH40cZyDRjGaRTyaYDX+bOKVV4GO1DDINJGeM96fQAZc49qTaN49acckZ9acqn73ei9gGsDnAphHzLg1P2IphTv6UJg9A25z/OlXO3A70KpDnPSlHPFAAQAc45pO/wBacRjC4zmmg/MPSkBHIcKfWmRgn5m/CpHUk4NNGcYrSLsgHcLlqbwwGe/WkHyoQx49acWCpnFIWgKUIx/D6Ubgw2qRgVG4XbjpmiMAAIOMU7ASqMj2FAOPqKTJyAtOZeBzUjHLyc9qB94k8ULjaQB+NNLgMCe1IABzkmnHkUmRjoKUAnnOKBDAwAGO9O6DimjHTtTwTQwI5ThPm6VCH8vAC/L61PJyMLyabgMAPzq4uyGOAJDEnFIFG0mnYpVXCVNwG4xHyOahZmDjB61OoJyDxUZUbh+lOL11DYWMcbqDjeAOppFbBC04580EijqAjAeYD3ozlx3pkjeXlsbvankkAn8qAQ5xuOQOTTHGTjsO9CnPH50rHjpnNACYy6le3WlOcn0pVBXj1oA5J7U7hYamVB3d6f0XimRg7iTzT+2BzSdwGnd1z1qpNdeTwi5bPWpbstFDxyTUMNu7srOAK1glbmewt9C1Eh2bifnPc1HKuFJFWVB4FQSqSNoNRzc0rgYtwd0wHGKhYkKMHI71Zu42EhYciqvAj4HJr0afwkO5JwADnmtKMF0GPxqhEolUE8Zq7C/A461NVX3KRNtDAc8A4pmxIhgc89aXICEbu9NhHyyE9M8e1cyuhisCQxzwDVjeWjHOKiJDBhjp1p8PKDFTLVDK6rKbvJPyY6VMyncyngY60ivliScntSkgkMT0/Wm3sJaEEg2Iu1frzVWRvs/yxjG4881oN/EScCs518mPeuXOeDWtJiZdhKvGXIw3pU9u2xGUjFU7Qt5TOw+Y9Kt2mJDJn8qmomrjJnT5Cc8msuf5XGT+NapAyKyL5f3x4yf5VNB+9YJMtWrb4846VJy0g42471HaNiPAH4mpGzg7m47VVRJS0AdIQ4Kgfj61XUFJNpb5fepicNkHqKg2puL/AMQ4HNTFCHS8kY5FWdIImvN44wKohmdMkYArR0SPy/OkPTAwa68NG0tSJPQuXsu3ispmLHA5NWJi1xMcdKesSxqOOfWu16sxitLlGWPbg0xJvLPzH6VLcvk+1ZskgeUEdqWiFOVi1NO0p64X0xUJ64Jz70qnI57dalAXHI5rVaI4Z6sZtz7Y7elKoIU+tKBhSDT0xkE9ewpWEmRn/a4Penp0LYyBUUn+sapF7p27+9XYkXrjIJz3FOA6c8fTpTCQEC8ihQeeM5p2Fckbpk8jsaaDzjqKD0HHFKB8v+0afQhbjSOQKQnPB5xUmNwHcCmEjHp6Uh7IQHb9KXrz1x3o55I4oAXBP5UaD1FDYIPQnvV23uA2FfqelUFPbGaeByOce2KLFRm42Nf+LFKPu9agtpTIm1j8471YGDmptY74yUldAAPSjqaUDmjHPvSGFFLjkj1oxkccVLAMdfSmAZbFSDoMHrRjHNFrDuCCpgcDNR45qQcgU0JsaRTCuDUp5zTD3NArkRTAA/WmbOtTH6Zoxmi1xkBjBIJA9BxUMlpHJtynTuKubc5prDio5E9yuZlBrNfMZkJUY4BqM2sgXd155xWjt4phWs5YeMilUfUyJIGVwxQgHocUnl7QMjFa7cDBHHoaFRJCqMo69cVk8NLZMtVO6MjaRxikc7EC55NbF3p3lRLIqkDPU9Kz5oHQKSuc+lYSpSi7NFKSepWjyzc1JtBpxwr5xxRkAfh0rFssQYLf0ozljzQnOTjvxQxAcAdaEtQA4HNNYZFOx82aaw4z2FAiowO84B96qvw+SMEng1NOx87gEY70sJ3/ACj5veumOiuK2pKG/dqM1MQQnp6UzpgEYNODbkJPbtWTKsRg4Kgjk+9PR/m2jt2qKZsRjC5JNPjbO3HFPoInUsDjP4U4NniowvOf1pcfe71m7ATb9xAz0qXfkYPWq2RjJpUfJ9qhxuMdKRxk80KMKhzUTtnknpSByQOOnSqS0AtPyoGeTTSCF6/MKicsyLt9aXfznrmlbQAuW2wsM9etFsixxgjqajlbe+zt3NWAgGMHim9I2BXGrlWapJOF2nvUa/ewe9SZDcHmoe4DFG1ORzTBIS4+X8acrZXp07UjLlsjpVLfUBXBUfWl5VevBobOeKVRu5zS6AiOQEIB60uTkAdutOPzZYio0bguKroBNnDEDvTeTkClByM96CM/N69qgBQMjn86acnjtTwflpp6UIBAApOarOxUsMcVOzfIcdaiZHdQe1XHzB7GftOSWOMdKswf3un9Kiu12sORSW7Zjbjdmuh6xuLbY0kYLHk9/Wmrtk5z9KQAPGBjpTe+0DpXOkO4bSFpIydrBuQD1qRVOME1FwMqGx6VSDYeM59qb5aiQkdadjJFOK5kzng0rgRknzQM5UdaRiPM+c/SpGAzn1qCeIFs9SaqNrgTSD5c54qq8TySqM8e1XGTEYB69qjgzk7jRF22CxIRhQOtKp9skUi8kDr709cHLd8VmwI5MswweKQjd26U4Z2t+lJuJkK9qYdSJhke9OdmIC46ClK44xzTRy33vpmrTAc2CoyvNRgoXCg/MOtPzgHPNVQ2yYHHFOMbgXAQZGYcjuKchLOQDUK8ZxUyKWJY/pUMB0h54oz90mkwSwB7UpyWwakAUYY5HNHY44NIGyxBoPGR3oGNlOCCOae3Yn8qiZ9gU4znrUi/McntTa0AcB8w9KRUHmgkUuc4J/KlJ+cHFSJ6iMfmJPamOuGJ7GpGwSPSopcjqeKaH0Bs+XjvQOEU4okG1RnoKAcjk9afQBBgAn1p44Q5PApg+6eOlP8A4RQwEI4NRg5U+3SpST0qPBDA9qEIVCSSW6UMMg4pAPnwetOkA6Cn1GNHQZNNOQMgcil6Ag0cggnqaYhmTkk84qBNvzEnOOtTtkHjj1qD5NrDdyauICwAqjjGe9TQuJFweopkQKqPpUiRKpyvSiXW4bEmfmph+8cj6VKOee1VpScg5wM1EdWD0JlHWggcg80nAIAFLwxzmkAwqdoA6d6M5G0UoO5qDg0wFGSMenemN90AU8jc3X60hwQT+lCH1EYAR471CxJY4qVsMAfSmEAsrAdRVIVhHb5l56ipG5QKe9QKdyc8bam3bsEH602gQHKKOKQDDbu9L/CTnOelN+9x6UgJQAF3Zpsp/d8/lT8AgAGmyegFStwIBuUAjrnvUrkbcjvTEClz/WpQTk56elXIBqDPGOlMYZI56U+LKuxJzmmsnDep6UluFtBSefXNI52y7R0ojB8vkZI70pOQM9TQFxhY5GeBVebAVm6DtVl8PjHFQyLvDpjiriwKRYBc461OoAKkVCxBIDdqnXO9cdK2lsJastMpCqAeRSlgqbsc0mRuHPWg4YMO4rAY7d8nIxmqt0uApPJFWc8cniq9wMx56kU4bi9TPlUlsAc+lJEWB9fSpGHybj+VEQ3DOOldV9Cbak8IYsdx4NTYy31qO2QsC2OlWsDA45FYSdmUN24HA5pQqh8EY9qceWHX6Un3ZMY5Pes7gBGSD07ik6kg/nTWnQS7DjNODDFOzARvm4x061DnbIxI+hqdwAuV/GoX5df1pxBiBuenWrMagKR3qkvN0SOgq6mWbnpTmF7sAOWNKBxjPBp3YgDimqwK/wC7WYwAyMdTSDABx2oLbTx3oGSCaAFQkpx0NSqvy1Gq7VJqbI2j2qZANkjV1O7tVfYpHvVn29ahP3qcWwIcKc+tV7mdIk29zVhiFVjWPcShpWkzxjjNdFOPM9RXK0rADaM46kVXJ3c44p8jFuemf0pgHA44rvSshJWQDA47+tOJ55FN6kc81OqADnJY027FERXjrx70nHXJqRuRgZ681Hz1x04xQiUJyakRSp98cYpigkEnnFTxDgBhwaUmDY9STgE/WkJ5wOaftKjjjHWo/oPrUIm40YBzzn3pTkICcUHgZowS2cfhTBdxM5UCtGyy1sVHUdaz8ZBA6ir2n5aJgDjHXFTLYPI0YfuYP5mpeADUUfCjFSjkYrhqP3mV0J4ieRWFrORdDP4VtxHaDWVrKYlQ+1bUXZiKun/LIQD15rUuFDW4OMAetZdgcvjPStWQZtmHenV0kUrlX7wz70AYwCeKQDvnOaVc7uR1FIYcAnHaopwPs8nptP8AKpG6cfnUdxzbyZ/un+VOO6E9zuO/WlzxQccUgA7V4ZQvSgUY656CjnnNAgpMUvNHSgYdenSimg8804YJyaAF4zRR3oXmkAtITwKCDmjpQICe2KMY60UvbrQADOPSm9qU9gM0jdMCmMQkZxingccdaacdTTgcigRG3yn2pBycU9+Rg9KRQFFWn7oD8ZwMUvTpQMkZ7UjHHNZgBP6VHI21d3SpDg496Y65Ug1Wz1AZlXj69aaQVQbDmnbSDwKTkfSqv2AkVtyc9aco+XPeo48t1oDMZCO1S0MceMY70joHUhuMUp5P8qUfdOfyp7BuRKFRCAaZEVLk5p7HqBxSQhQpIGD6VXmBKhBLKQfY0qZFIqkN14pQCCBUMBjKdzHNSKRtyKRgAMHuKRD8uAMUdAGnjgjg1LnAGKYzEMMjj1p/XBzwfWhgB4GaAe+eaVjgdeKaDzgVIh4HFA+8TSnBoXOKQAOufypGOBg0oFIRk+1AxSDt4pgkGdmeaeOQeaZsQPvx81NW6huOyelHXgcUmDt96UZ20gDoeKBj1x7UMOMUg+ReRzQhC9GznimkZzSNyowcYpgJZc96pIZICPWlGAaaBxmhevPrQA7OBRxjpzSHuKXoMDrUgO4C+9R5OcHpTyaYeDTQAwPGelMIO7g8VKR71CWIfaBj3poES4GRz9aXJ7VGGC/eNOX1J4osA/GRkUmDtwKUHPekyAp4/GpAb6dzSgjfxSDA7c0oXPNMBxznOelN6AilOduM/Wkbkn0FCQDMcdaTgc560/G4ZJxTHICjHWmgGP8AdAPQ0bsHZio7gEoGAztqEXBO3eMGtVG6AsSKeD1K08Yxk4z3xTSC0eR8pqSMAAipb0AZ5ygqCdu48e9TdgarzRCUYJ5Hf0qccoAMYFDStdbiF5A5NDqCcHpSLySCeBSuCV3Co6jFHC57mkBHQ8UdcZoCZDAdaQBxk4zxTl60g45pRk5PegCF8jJPSooX5IByT3qcgMeehqF1CkEdQe1aRfQROenHNO4UDnmo1znntT0JJOahoYuCRnNROcMMcD1qQnimMuQM9TTiA3o3HalcgH1zSdWIHQUrqAeKrcBW6DgUbQEBpCQF55Pahz8uO9IBQMKD0oO3f7UMfkHrTcc9aLAx4HXilC5Ix2oBwPenL3qdQK5BSQ5PXoKkUhAM/e/lTZVO8MegqshYzMO3vWqs0BO6GUgsOB0p4GDnsKQM3Tt2pSc4GOBUvsA5eSCaZMVWLnpS793TpSMAwOOvektHqJmVd5KFB1FUVGyUcZrQkTzHZiOB3zVDZukIB59K9Gloidi3KUW1UDANSQYyBu5HSq6ZfKv261PC6o54zzSnomhrclXc8ZJGMGpUKyg5GAKcNpYYHWkdlSZVxw3Wua9x2ECYiYA/jRCQYhzyKUhd20HjvSJkKaXQRGAPPznt0qQqSfm70mNr+5p77vLHHNNvYdhhCyKc/dHU1m3D5uEWLkHNaGR5WAOO4zVSUbIS4Hzjr7VrTdpakslTAkTnjtV2LCMduMHrWfp58yMbv4TxVu1IScgdqVVatFIuA5+UjBrMvI90uM/jWiBuYntmql1ES2RyaxpO0g6DbVQEwTUpwWGRkDtTIVYcE1LKdpAHTua2qau4Dcjf8owTVcjy2cg5x2xVrfu2kAYHU1GFR0bI79KiMmncQ0ooUE81fsELW84U8cYrPZwW6YHsa09HYSW0pA74rqw13O5FR6CBdnA4qC4k2irN0rAfLWVM5LdenWu5uxjuVbmQhMDqaqocmnStuctjP402M55xx2oRi3d3LCEcenrUybT7gfpUKHLelSISpI9K2Wxyz3JOckAc9qmt1JDMo5XpVZnznApwkOxl9cZp8tyFK0tS4I0ZiSAT603yogrZHfj3qAM6DaOM9fepY5N2FY4A74pcr6GjqR6rUb5B+0FM/d9vahoHVS3X1p7TcPhfmOMtTxMGEhJxuxgYp6kvl1+ZWeMqMnv7UYbHX6D1q3K3zBU+8TSbVLP8uT2pX0D2a6FTO0EY+opoJCjjj+dWTCGjQjvnJ9KiSMsenTpTuQ4voRjnkd6TkHA71KoYpnp608IQcGqsToiJUYEZPHrS8AY796kbjOOnpUHQUrWFe+xPG2xw+Mn0rUiIZQw71kJwMZ5q/YuChTutKWqOmhJxfL3LXoaB/nNAIxg0Gp6HWHTp607HFN5xg8U7PB9aW4CAbTtHSndaQCgnFKwx9KDUeT0pwPP1piY/8KYewxTutNJoATrSAd6OlLgCgAxTSOOKf+HFIxOOKAI6D0Pel70h64NIojbnNLAuJBS45zT4gQ4P507BfqX5pGGi3YKgj5cH0OappErA5H0qW6n/AOJdLbY4kYc/Q0q42Ciyf9ehC7/11M6az+cnOPcCqLoVYgjgV0SFQ4Mi7l7jNN1Czgx5tv8ANG3YjlD6GuWrhoy1RrCrbc5wHJ54qKU7fm25rQe3VW44PrURtGOGJHWuVUJxextzxZWJBGe9IQTz0FWmtDnAI9qhljZBgg1DhJdB8yZnTKC5fNFvgLkGknBV8BuvNOwFQY4HoK22WoD2bAJx1qRQApPPNRl8jCjvTkYeaygVn0GNmBBUjpTmOX+Xp6UXLFQqp1piNvIUnBHcU1ewFkqcCjgfjS4ycCkO3d1zWQAT8uKUDAANNDZHpSj5jigBsi5pozu9hTpFJGMVFtwQC31prYCQH5ST1oLg44xSAKPlA4pFUFgo7U9ACJ2d2zwvbFWTIAgxVZGYyEGpDxjNKS1H0JQ+Bk9e9P3ADPequTvPpTw+ee1S4h1FjYidgemalfBYZqHcAxPc07zMp05oYtiUgdKTtheKZnODT0IO7B61I9xozsII5NN2kAZPFPYYIHpTJPY/hTQtiYNlQRRnOTnv0pq/Km3pSjIGc9Kmwxw4BFV5WcEqhqcKeSTzSEcbs0J2YDdpCcDmmcmPHSpcHaD+lNH3eelO4yhPkfj2ojUhefyqacDIDdKi2FnPy8Y45rdPQV2WovlQgdM8Uv8Ay1wKjUHG3pT3J80ZwKza1AcVywBpmFG4Y4FPYrtK56UxT8hJ60kMUEhMdM09SuD6+1M5wM8UkvyHdnNG4gz1GeajWTfKVA4B609XDng596EGHPoPWq9QJJCOcHmo1AEpPf0qQKGOfSogCZ89/SkthknPJXrSrwOTwaXo+PzpBzlakBpOFPODnjigjHzDqe9DsARntSHAOO4qkIGYAkZ+aoBzIM9R2qR2P3sA0xiFZWUZYmqihjpFOdtRXKlcOM+9WCGMe8nmoGBePy/4jTi9RD4X3MOcg1YRsSkdj0qpG21toHK1ZUDeWxx2qZoCRhjI70mCq5PXtSk53Dv600cgCoQxFBKEg8g0u71+9Qo+YgUvA5I6d6BDcKCuacjcH3pGG4YzyKAoIBBoGLwoz60o6ikGCRSlgG5oDqJ94U3GQc9KVsliRxUcrMSFU9elNK4hcF1XPPWhRlee1OTCjBpF4Q46mgBV+Xp3pf4OO9JtBXk05jleKQxOnXqelMI+X5u9PJHA70jcqM00AwD591PY8bsU1t+7I+7SPk89MUxCEbmINBPHr60AZJOaGbCEjrQAh+/nqMc1XEK5JPLU+TOw4/8A1UwMAVHX19q0jsAsQ5LE8jtVlTuz7VAvyKSPXip1IAz3NKQC5+XpUUxVUGTnmpT92mSLkexqVow1BSMAikUEA56nvQh3Lx09qUqccGgYmTt+valAwgyOaac7qcjEsQRTYeQ4YB69aYynYcdaVAec9qAeuelINRhAZGUcZqJCQh3E8VMF4wfzqLOCVPU1a7C1EClwVbv3p0YABHb+dKCWJX8qIgUQg0NgO4ACkcHpSFgnJ4JpANwGe1RyZb5aEgJ0GCP0pGPGOppvIUc08fdLUh7EZXyxwKlJBwR3600/MnvSAnaTjFG4hMgzH1HSng9jUbE71K9+tSEEvkmhgIhJb2pDkZHY09eGOTx2pmCGyeaQDOduT+NRk7YsHrUjfMuccA1HwV6dK0QEDR/xr2oh+7zTw4MZ9fSoInO7k9a1V2hX1LkQyc9akxySOppu0IM54NSDPHvWLGRYyCD1qObG3jsOlTlcZx1NMdegxmhMFozI8xtxBHX1q3bKOp5OOlJLEpJ479KIpNqE4x6V0t3WhPUsRPhio698VI5AIxzUUWBEJMcnqakbG3p+NYvcY8sCciqv2kvcFAM+9TO4SHcD0quFJkEpOM0RS6juJO6EHHDZ496sxnGFAye9VmhY3JLj5KsNKsaFwOaqS0SROtx4DIct92opD83Tkd6py3UoAG/JPbFVDdyEYLZz3q40pPUGy35jxy44A71eW6hCZDYxWIbjecEfMaYZMjOeK0dHm3Frc3Wv7dejH8qFvIufSsBXbOcn6VNIXbGxSW9ql4eK0Hdm2LmMjlqesse3h6xlibydz/f9aIyZXUF8bal0V0YXaN0sGjHPHpTlbncDxWK9wyqCF2596sfaHG3a3DfpUOgwUjS3ZzzionI6iqEtxMjbQ4PrgUz+0SuUcDPekqMugcxYvDsiwD161iStubHUjvWjNcLOjEMBxWfFHEy7i3PpXTRXKtR3GLGXbAHalMShsMT+FWE+VCOg7UxlXG7OK15tRJ9hg28nGOMAmmMe+evWhm5461ETxjJqkgQE5I+Y9OKUD2/GkHbqalVCO+Se1NjegigEHk8VIrc57DpSKME8Uqr3B61LZDaY8OCCSPemknOADnPJpQMA/lxSFW3ZBpaFXsIGJzk80mSU6n6UhyDgGlI4689KYhc5BGRir2lnAf27VnjCgnHGKmspSjkE/e7UmroRsQPkAE9KmDZbGeKqQHcCasqMdeK46qtJlRelyeIgPgelZ2snIU4P5Vdib96Biq2sYMANOjpIl6GXaf63IGM1rghoWye1Y9v/AKwYNbEWTGVPpWlayLiVEUnrxntT8+v500ZBI7ZpeT1qWCQ0HvimznFtIMfwH+VSDABHpUc//HvJgfwn+VOO4zuDjHvRnHakPGaQ8cjrXhgO6jPejikBJHIwfSloADSZAyKUnjFBFMYnXoaXvxTT1pRSAd16UvageppDQIB1pcUcDr1pGbjGOp60gA+4pTxikBoyGoAXPem04c9OlGMnGKAGlQwGeopdxODQcUufancYNjFMwDjFLkHoaMkDIGadrCFXnIpxXKkHpTEBzzT+Se9L0AOFGKQjKkGlJB7Uozk0rjIc4KhjR646ilIBYE8nNKVG7NVoA1Rzmkxg4/Ohjuc89KQ5GDTEPGVNOApAemetKegqRjGHX1qGBSGIY1KQFFRwhmY5HPY1aegdSwPY0mMNnNKAdvFNZSMHr7VCAQDMmaEfIx0of7oxTAreepB+XvVKz3CxI2W4HSkhYc+opzbsHAqILjJIPHWkldATYz1pVHNKPm5HYUg61IDgd1KODn0pMY4oX1xSAUd8mjAoHTJ70vG0DvQ9BDcA0Ypcc8UuM55oGN5zQTjAxTlz9aTOTg0XAUnuR+FIwzwKU9elB6EUrgRugYYxxUOCg4HFWGJAyKQ8rjFUnYBqg7AfWlIIGCeKQBiu3sKCQOtU3qDHAnn2oHJyT0qESAOE7noKnxxSasIVuuB0prcNSk4FNbk0kMcCSeAaiKDO48kVIOmQeaaW2jHXNNANyrcHmnA/LjGKQqWBwMelLgnNGgCo3GKeetRk+Wueop3Xmk11AQfMMsOKUE44/KkINPxz16UAIcAdeaYfTvUgA6npTMDcPahMBO+KZL06U/IDcD61HKw2sapAMBDpkc1XWNkYlh06VPGQqgr07UNE0iAE89c1rezAdGG2Fm6U0Md5Yng8fSn7cKRmmrGAnJ/CpTQCsw2NtGTUicKS3TvTFRYuB0pwXKYHak7APyW5FKcAAZ+tRQhtuGPTv609sKuM8mk1ZgO3DI4pRxk9zSEZUAU4HLAVDAQnkGlJwOBxScFvWkIyelAETnDY71GmSSWHJ6VOyA/NUTDLIR1GcVpFgSKc5pQeMd6VQQM/rSEAsM9alhcUn5Av6U1cbqFZS5AJzTW+VS3NFugCZw5OeBTnHANNBDAcfhTjgimHQa2fL5oOAu79KaH3FccL3p74Iz2p2AcoDDnpTT3IpQQB0pDnAzxSAIxhQSakB28+tRLtJyvepByQMfShh0GyDcDk5qtnDKcYI71aUHqTxUEp5wB171UH0EKHzJz1IpGbCmkt12qcncT0NEh2kU2leyGLvC/Kenenn7vuaiC7n6VISS1DEUpx5aHPCjvWUhPmMw7HrWlely2BwvrWYWwhx949BXbQTauSydX8wKd314q2iKjFmJ3HoKqRphlB69TU/mb2Vzjb61U4paIeqLkQLgFjk9qjijuGkAkb5R0qWJyQO2akZSSpA6d65L2kMTG5zkdKYpI3KeCaeFIlbHTNRgYmORwaSSAQ5EoOODUrNuG5hnFQzuItrk/SnM25AxPWnbS7ENkIjK4XANU71ljAXYTu64NWp1OQxPy44qC6dVjEmzJ7c1rBaoBbYKNuG5xyAOlXIgqybwKo2Sfuw2cEHmrccm2cKOhp1FrYZc3EOQw4PQ1DPII8ZGT604yqHA+8e9JLgkF1zXKlZ6gIBnDHjNIzCkc+Zg/wjrShSSSzYA6GrYCE4jww59KqyuIY8nOancqFOCWAqAHc7E9PSrgiQRMIGPGe1bGjJttXx3NZDt5ezcfmPatfR5A8cuDnBGa68NdTZFT4Sa578Vz90xj3HNdBc9fasC/UriuyRz3aVzOK8n0pyNgc0xvmO3NPTBGPX9amOrM9lZEqnA9z+tTds5wy9qqqcPnH/wBap9wKhQetawMJaakhwVHqau+WAhQ9O3vVVVIYHsKkwQS2OSc1ruRzKN2y00KCPOM89KjkiUbirYI6jFNEziUHOfw6Uqv8uepY8rQk0EpRn0GGFw7DGdpFO8lvNRSuM571Is6bmyMKafkOU+bIXOfejUFCMtv61KxHlsrk7gfwpokdQ3YnrVpyCrMVyU7Z9ajkjUeYw6jFJClFrRP+v6TIvMyqqy8VYkJCP6cYqJrfDNt524/GmtCcqByTmnoxpzSd1/WxNaEk7QMA9amfariMnIXqSOtVYmMZyOKcZGY5PLd6dtbkuVopDXBBORnNQkhe9TDnkjH9aYVC4bFNmUbJkQwrcHIFWbJwJsZ5NVzxkZ4qS0OLhcfjUvY2h8SNYHFKAKZnP0pwOKg7xRyKcOlJnj3oUkUgFxxmkYUuabxg0wEHWlB5pvf6UYP4UrhYmDcUnXNNWgdTQA7IJ5oBpo4pT1+lMBaCPWkzSZ456UgEIPPt3oxS5o6H+tAABx7elOUcgmminqOefwoQXIrpsuq9utWkIMa/SqM7Zmq3AcwihB0JMjNOjlaMnHIYYINMPag59KNxFeaEKwwc5/SqxOFOOgq7IATVRoiCe4qWikRKS2aeACnI/OlI4+lNPIweo60mkVchktopT8yAmoRp0CZKAjPvVsYA68UFstxScEx8zKH9nfOSrZzTBbSRZLL+QrTwaeFBAB5rJ0U1YpVGjBu1YrnGD71HC7fKpA3d66IxRvwVGPpVefSombfEdrH8qz9jJKxSqK5QXGfSkI7/AKVYNlLGOBu+lQuhXORg1yyi0zRSTGJjBOKepx8wNJtGOfypetSA1zkZz+NRAZl6cevrU5QBQO1RZG4DHIqohsIMbiTwPXFBPznvSsQpIIzTE4YkdBRbqAmACOOakPZu3pTEYFjkc08cpTYCHoSDSr15pBjFOU5J45PepDoN4JHGKMtjilcjJA/GkxxwaY+g4Mdo9Kej7Rgd6hY7UGe1G48eppWuBYY5JPbtSBsgnv2qIkfxc+1N80bemaFEL2LUYyuSeaft3HFQo4G3B+lS5APXrWctwFByOPzpD83Hp1oHQ+lKMYxnr2ot2AQY3ECmsM5AHSj+I9qBkqTRsMqTbjNkfd7U5BgHnkUTrkEL1PWo4+Bj863Xwk9SRWABb8qeXAYHrTCPkHp6UpGQDipHqSyANjPp1oZVI4+9QCSFHahm5wO1RqAEkuAw4oIVztP3TTQSWJPSnFhs5oARECHHaopgwbB/MVKr5YqOTTLgsWAHGapXuG5IgPl9OtRgbX5+8elSKDjaaZKVQgn71JbgSoCTz1pAMSHFC42g5pVxgn16VIDJYg4+b/8AVTWVS/qaezlATg0blyB3NNNgREkgg9qZtLBeD3yKlZMKQ3PfNNfgBs8VSfYB6HMROarqfLds546VMCqpgHr2pWjEkYBHWmnZgVi4DbhyO9Th3zGVPy9/eoDDtYjP1NTRAbQufmzVSsGpZYYP1po4OT1px6jjIqLJeQjsKxSuA8EKxp2Pl2mmIQzb/SnqcqTihj0sNYnaMU5QQh/nSEYU8+9JG3yFiaOgCONuB3FKxG2k3Zbk80jcklhgUwG7vm24z7051w49qRD8+McU5uOQeTQAKNp65oAyGzSr8gye/ak5AJoAF6lSck96U8ZAFAHTnrRnnFIQHAIFICN2M9O1Byxx0pAOCe9MYgy2Se1KSC/B4xRjahPc00Db0FMQ0ZLHnGetPxwRTQMj3pRksc9KGAyUYIFVQSLkqenb3q03949RUK4O5+taRemoNBj971qZPmIXPSoAQxDMcEU6KYMSBzzTa0DQskgj6UjHMWQcknpS4AXHemAAL9azVgEC7FPqTUoAKD3phxjPpToyT1pMBpG760gGXK5xSquMk1GpySwpoCQHDY601gSpA6ilQjdv7Ui55z3oGKw+QCoJCQ6nrU4OVI9agkAyu48etVHcQq57jmliJYbs9KRiQu/qTQchRt79qqwbIA2ZcUck57dqVv3fzdSaYc7xg9aEA+NSyncOlPznHamqTkilUgg1LGOZQVxng0ijEZA7U5SSuD2pobIIxikIaciQCnYySKibdvB7U9BwGzyKbQdR+BuwBSMTknrTgQ2QKjYBkIHUUkAnIGD3/WmOAQABx6VIcbcntUSn5iCfrVICFkUs3bbUAwrgAdO5qYoVlPcnrUOTuYkdcVvERfByFB5xTix3Yx+NRRsSQalUlic9c1iyhrBs/KfqKjnkEQyanK5aq9yCwB6kdqI6vUWxXZg7MCcGlZQQF7jvUUxEcvPPvU8bq0vpWz0V0IkUbI1DdKJNoVF9akYq4JI6frTARgLjkCs0Mi2I2Ys4HXFOaNSyMOQOhFR+YF3Owxj9KW2YSRlAfpVu9riJWDScdx2qC6J+THvmpd6oS+elNMazDHWhaMRjz54BGCfeoSn5VduYwvA/KqXAGewrsg7rQTHIOTjqPWmsCBgU9GBbaD+NO2KZFx0qr6jHRRksMr17Z61fUrGDjjiqavtPT5uxz0qaRwY+cZ6Gs5aiQ/IK7jzUQZQcgcg9KiXkZJx9afEAc4wSaTSQBM5nkwB0pTE7LkNyOgqwkbKM9T/KmE7ZOTjPNF7bDZGIniJLck1VkIJGeM1bdgWbJ5HfNQeU2RuHXpTi+rE0Vhk5U8jPWpIEGSzdae0a7wOMd6lXaB7VTloUmDvxgCqkj5JJOcVLLJwQrZJ6mqpbsKcIgtWBOTwKTHrk8cUoPOelJz69a0GSBQOMH1p4wFyM+1Kg79xSY6knntipbJbDO7cB1FIpIA4HNHIyDjmkIxjnv1oF0JF4HNPPbjp+tNBJANB59etTuG4zuTikBPTv6+tOwQ3YfWkH3eB+JpgGMkZHSnW+BP8AWm5woOCaWM4nB7Uwd7G1ACT0xVoLxzVaBcd+e5xVoe1cVd++OOwqA7x7VDqag2x/wqVQRg0l4N1o4zz6VNNiZz8L7ZQDitqDAGOOlc+G/e9Ovet62OcGt660KRC4w7AnvRyOMU6QHzWJPGaZ15B/Gs1sGwEjcBUc4H2WQZ/gP8qfzuyajnP+jy8dVP8AKqjugO5oI96OKOnNeGMXoPrQOhpAckgGlzjpQAn8qTvml9sUA+tACGgcnNBx1FKP1pgLnjijt70YwKUE4qQA89uKMjHTmjiigYYIHSjgDFB7cUZAPpQIUe1J3IpRwKQ88+tACHik+p4pSMijGVNMBDjoetA5JxQp9aQLljz07VootvlsAo++MdKkBPrTMFTkjinA5GaJ0pQV2hXTDPpR1GKMc+9LwAc1kMY2FU47UxjuqQjnmmOOtUhjAQHPrQckilOTz60Ajqe1VsA7f82B1oYnAxzSABctnNOzkZNSBBPIAMAcmlRsgAdRSyFY0LHpUNs/mksv3a0SvHYXUudOtIScZApQMdTTh92sRkPUCgcMTn6UpBbIFIIwcVotwJDwc5qIkgMe1EjEMAKfjKcmlayuAsRBAweDSnJNCYHA6Cl6Gpe4DgOKXPPWgUH7w9BUhuGfm5oGASaX3FH8OO9ABkYNB9KB0weSKP4c0gAkj+lInbPWnYIAHSkHqaYB0OQKOoxnikBJNL/FQAE8YFJk/TFLkk5pGyTxyM0AN5zSMMnjpTz1HHHrTSCAxApoCMKvmZA56ZqRT7U0fKevNO7fWm3cA6jmhh37UuBgUY5qQGY44oY4470IMc0MoLDJ6VWlxCqSR707J28UxPvEA0gJ3kGk0MSQkAADJqRRj8KjY8lj0FOD4Qc5B71W6EPUil4LHmkGM9KAeM+tRYYpIxxmmkDnuTSM2CB2707rznoKdgI0Jwcjmop+EwR1qbnPWo5V3xnPNWnqJ7ESL8i7T061Z27R8oxUUSYUDHSpSMjdmnOV2MiCYNORcnHp3pACMmlXIYknOaQDiB09aRI/LTjpTyAxGOvelPIHrU3YESkqAOoqQgYyRzSFdz0442mi4CelAxnPrR1z9KUH5cUgEPA96aGLHGMU44znNNXnJpgK5wpFRjBG3OadnggZJpsfJwO3U00A/oMdhSgAkDv60AclqXsWzzSuA1QA5OOBS/wtilHTrTHJ2kLwaAGKOc5xihQcHNJhjhV/GlQKQ2CeDV2C4Ku0YIpCAelSZBXIzzTQDkjvSuIc3TtikOGwQKVuppg3bmz93sKEMaCd1PU1CzfNxUwxgD1psCBSys6g5/pT8FSxp/lqWJA5P60rAMuB1qua4EUZAOM9OtOdlxnORVXJWdhj5T39anJy4AHFNx6iEKhWD55NPZhgt60wkhDuOcnikXLHFFroCneYEBz1aqSIAA7dfStC/wBpATGT61TeJj32gdOK6qUtBJBFkb3Iy9Sxgu6sG+U1DBuj3ufmBHWnREtJv6Adq0dwRcRnMgGMY/WrS8ENu49KgVkKbieD3qZCojJ64rjndjEeRWUt/D61XRgZSN2anKAoDjj0FQbdkitt4pxtsIkkjV3UHnFNJSTdhuU7elPkbAJAwPSool2I2/71NaIYpAkiGfx9qqSgzSbAnyjoc1cg5jJ9eQarDzGdmKEIOnvWkL3dughlsVjcxbtzdT7Vadiz8Ljuaz7XEdwBj5s9avsyibYB9TV1E7p7gtiVTlconXqadK7BNuMYqUoERQDUd0Tt3DFczachsbGy7MZzUinOQRx6VQtZHdju71e3YO1eBVzp2BEZ2AEBM496ZkDPO0elSOApckfN6ZqqZW3KAOKcESyR4cqrYyxFaeixmNJgRtzjiqSZIyTj2rR00/O49q3w82p8oprQnnGc1g6kpKZHQd6356yLxdyMD6V6Mjl3Rhgc5x9DThnGD2HWk+8R7VIq5IHepijFvQjXOcg55q3CvOeoNKkBJ+tWBH5fNXFMylJLcFHXJ5pdpAyKdjg0hJ2njitUmc0tdxDjOKReGBxx9aeeSefmprHLYU07i21AZwM96ViB9RRnAyeKjLehzQN2SHFyhODwKckjKpx+NMztOT+VGMA54FPQabWxL9pIfJFPiZBIgUcDOTVYnJ560vIBNLlQ1Ud9SbKrbOUP3sfzpMgqMH8aiLHaSPu+lCkgYJwaS2Keu3YmBPbilPPfcaYjgqcnpTwBgc8CqMne5AcNlcc5qW2XMmey+1LJGxb5Rk/WrSReVGAOveoZ0UY3Y8Hmng96j+valB461B3Eo+tKOmcUwEZpd2D70yR2cUn9KiMgB54FPDAr1zSGHIp3Y/ypuMHNOzzTsAAil4zSZ60o6ZosAmaXPFGcikFCAAfyoznvRyab06UAOzz0oHSmjrinClsAuaeh5z2qMfp6Ub9q5oAgkIMpxVy2OYjjtVDqSelXbM4Uikh9CweaKCD0pM461SJIpSAf6VCeaWZ9zcHihBkc0hojchaiwXJx0p4UyyEdqS4nS3GxRlj2pWuDkoq7E2BeW59qMgGs97iYk/MQPT0qNpJSThjxT5dDF112NQNzQsobI7isz7TKp5OfUVPBMHwR19KVmjWNSMtDQUgCpAc96rqefepQ3FModnOfSmMFI5UH8Kdnimc5pNJhcY9pG6ZAwfaqklm6HI5HYVpK1GQfrWUqEZFKbRjMrKQSCMVF06da2WRScED8qryWiEk4xXNLDSWxqqqMpvvjdyTRlgDjp2q7JYEnPX2qrLBIq42kVm6clo0VzJjBkKWHJNPXO3p9ajVgq45zmpRyMZGKzZQdD8pyaVeuOKbkZwBSrw3X9KQwcYXgCmIwZTT3HHqaaAoAwOT3prYBpXIwKVeBxSqBu68GkI+bg0AEgIAC9e5qNio+XFSnOajK4bkcimgHoQO9S7uRz9agjBYjPbrTxnJ5pSQEgkYHnkVLv5FVhjpnmpUPGTUNATAAHIPNDAZ56CkyeMD5ac2WGKgZSmIW4Oe9Qx5MpGduO9WJ1DS5PI9KiCsJh2SuiNrCdyRRnPPNEjBSoY8ewoUfvDzS+WC2WzgVN9dQHrh8HdwKMfK/ejvtPek6HCn8KQMUZKDAzSuqiMc5oBIyMcUBMrhjSAYpVWCjk02dy0wX9aE2q3rilcM6ErgGqVrgLEHIIJzTbj5cMeT6U62ByQ1K6hnIalf3g0HjG1WJ605Sd2MUjAlhgU45wSOKhhsJ/wAszniokUkbj1FTHA6nrTcdRjrQnYY1sHFNIXoeRQrBnOPuikccHFUA2Nizt09qmTgkdj2qKJRGSM81KmQN3pRJgICCCMfhTBgMMjB7U52wAdv3qafnQZP40xEyjDAdqiORJjFTJjkY6VGx+cccdqlbjFA2jgc+lPzheOp7UD270h6kD8KQCNjGSelRu2AMflT5eRg9qrMxMYIGcGqiridiRcyS4IOB0OafJ8xIHX+dNX/UhhwetSkgqWHWmwRGj5jGBzTied35Uin5eO9M5YleeKQdBx3Bse/NSEjaQO1IWwT7d6aMg89xSYxVyVBPWlHDjHTvSMcU4ngE9aAEI28jvQDgf0pCScU4qM4oC4YyAW/Cm87sU9iMCo0J5J/CkhCD5X4pQflPqaaDhs9KBwMA81QxJQDGFzVYuEGxTzVqTAH9ahkGDjjJ6VcWIjkx5YBGc96jgIhfA/OkdnlKxg9OtDxMvIPHrWttLMPQvhgTkfjSkZ2kCoraRWTH8Q61Lk4x3rBpp2ACvPSnRkFeO1Jn5Wx1pkZ+TIpWuhiyElwB071GrnBwKkI24INMb74AGPSqQhCMHI6089B6U2XceF60pzsx3oAU/N0qGZWkfpxUwwF9Ce9RkNtbnB7U4uzAa4zH14HSlUsyDHT+VJyYeaImGwr6U+gxJ2KquB+PpTCpdVGelTOQsagDrTCmGLZwMcU09BCjKjjr3p69iPxqNE+VjnipFGV9hSYDgSUPakPbHWnIMkr2NRuTu47VK3ARsggA4x2pMgHrwaGJ3BjxTguVPY9qoaHKNnFIehAHNKx5BY/hQSxJPSpFuMf7pAqF2+QgdRT+EG7Oc9qFGSxPSrWgXIHOAD3PvUZwegzUjA55P/16Yy7mGOGFaoRMjsQPlx9KnJxhh36iokIC57GnK23is2McTjn1qORx9DUjBQoOaglwzZ6Ed/WiK1ArOSsu1uR2psRJmOV5qVk3ujg446UnmAMxXqtbdBE8TDB56UBgwMmOlRAbo/Ru1PTEkYRjyahrqAhQZPc46VDaI0YY9/SpSyfaShb5wKrCSQzGNueeDVJNqwtkWUtyGK7utLDmFsE5FOCGNuTxioQjt8jtn0pbjFu0UsGYcGsmZVBwuAT1FbN0h8tV64rFlI87BH4+lbUHoFtbkaAliB0FTJG4bKr9DTYf9bxjFW5ZfKjQqBuPetpN7IV77kEihNhJ5PtQMlSOpPemplj1478Vft0Qoxxz71LdheZEtodnXNPRFT147VNJchFEcfI/iOKjChfmJyaht9R6XHyTeWmFHzEVCoWQ7mHT9KVmH3iu4npRI2IsDvQ7sOpEYNzj0NWI12g7xTBlF3bgeKRbgtEzMMY6UpKTHcrOU+0HFQSSlmKqfyoY7fm3EFqgPK8dRW8YiQshx8o/OowfYdKDnr+lA5rRFWDafpxS4HU5p4TAAPc0N0IHHNFxXHpJk9h/WnZwfbtUW3vgipFAXAJGfWpaJ0WqEAJ64+tKqtxjpU3llDyBjNSAggDjOKhyJbGFGwT2PtTMkDB6VOAfM61HIvzYHfvSTuUrsh5z070hPXuKcVOc/qaYR19PSrQJDhtbA5+hpi/LMvJANPwdobAzTRlnGOvvTQjbhb5RnrVxTyM1RtyRGrMeKuoc5FcVdalQ2Hc0kik278dR3pw5NKTuVgM9Kzg9Qe5yhX96RzweK3LXmND2x1rHYYmYdPwrVtz+6AB6d8V11noCJLn7xyahUYPWp7jO5TUK49OBWS2Ktd3AHHbiop+bd/8Acb+VSBxgjPPrimTkC3kH+yaqKakgR3OMHpRQfXtSA814QxVxzS44wOtIOTzQRmkIOOuab3HFLigdPemApOM8UgwetKfekPvQMVemaUZpCDgelKeOlIAxk+9HAAoxRxQAHHFGMkc0gBz9aUcc0AKTyKD6CkHJp34UgEbp70z3zTjnk0iDH3jxWkNxDQMtk9KXnd9KX+LHYUAbsj0qk5c9+oaWH+YNoGMikXPbrTTgLwOR2p+4HHGMd61qTc1aTJSs7h35pKXtmgHHPFc8rdChpI5FMbd68U84pNwxikhkYGeQc5oc4XB/OnKwAwKJFLLgHAp31AihOOBzmp1Hy81BHCEYkc1K77B04qpavQBsrBUYk5WoYDvQlRgfzp24NGXxkUW5DKTjaPSqWkWImzxjNOzxzSYwM0D5lxms3YGNwWUkcZpAOgzyOlAJ5C8j0peBz1NNaDsOcDjpk0MxWMEj8KXIxu60h+cfMfwoWoCxBcAjpQx+c5pEYEDAxTzgNipe4DhwuaDyBR2ApD+VSAvXoKdggUgBxuoXnJpAC4xyaacHinAcdKAMmgAbJ4HIoIx70h+9SnpigAAx3oBO6gnj/CkPFAACSSCaGODjtQOB70hXvTEB4o4A9u9KelNOT36UDGsCwFKQCCf0pXB/CgkHgDApgKTkZ/nR7UmcgjFKAdvApANU4OMUpHNB5NBwAeetMBOD+NNCAgjPNKOD1/ClGC5INMBNgC7TwKUqCm0805goHJ+lMxuOVNCbAAMUrdPXHpSN2wKMYXOeaAHE4A460vb2NRhsLlulO5YAYoaACBjjoaYRkcdBUmAB14pgwoOKEwA5AGM5p38OPzprA7s5604j5SKAGD/WEdqUjninKMGmtnOaNwHe/Wl4HNNBHfpTgMLSEIeoNDcD3oIOKAQeTQMawz0704nCgCg0oBPPpQAhAx05poOFwBTs89KaOOKYDuhHqe1MwRkA/hTwDnrQMEjHWjYBMDp696G+VMZ/Gjqwz260HDYBoAA2D+FIwwh460fx+1DjkZORQBGVIAK9O9EY65602Mkj5sgelPjGXLZ/Cr6CHY+bFIy7D1NOU7jkdqCc8Ede9SBG+ScflSgkqeOlIxbeD2FPz7cU+gyvt3EknAqWNgSBt5FIqHn+dLGQuR1NV0AeAN3J96bwGGM4pduT9acMAdOtTcChNG3nggcCpcn7xPSrLLlTVd4uM/5NaqpdWEMQckk06PhyOo9aiXDfd9efapQCku2m0BHNbl5c4wB1qneFwuVPyCtRt2fb1NZl0FYHDjI7EVVJ3aQnsNt9iQllGT3pGkUxBgOT1FEEgJ8sDLH07U2din7sDr2rosuYOhcj+dAAMCpUIL7exqvASI8ZyfepoVaNenPasKisV1HyOQrADA7VXPCKe+atKcr8+CwqtKSwJ6AdD6Uo7kse4JYsRkDpSLtUqxHzsKQtttwQNxPeoirtKrD/APVVJdwJcgExg4LelMIIt2ER5PenhB529uCe1AQJGyDjPSmn2DcpIyLOARlx1arUpB+fByaoqu2ZlZiT7Vbhc4wVxjgVtNXsxosF2ZY37DrirDASqMDIIqrFI7fu2XC+tW92FG01zVAKcEflOSR1qwmPMLAde9NYFmHNSIM5TPHrSc9NREDDEmFGSe9RFXPHQZ496sttXkH5j1NVEQo5Lcj61dN3AnSMlipbt0rR0vPmOcdutZ1uqbWkY9Oc1p6cwMhxwMVtSTVZEy+EsS/xVmzr1yK1ZQN1Z9yvU16bOWJzpGJmHcGrSR5xnrUVzHtuiQPQ1bjwFGDyelKnq7HPU0WpLFCNpJPHfilk+YZB+laUFjLLZtMo/dggE+mahvYY47gxROHCjlsd63Ry1GloUOpIpDgAAfepxXvjp70wqOOMCixFxCRnKmndMfpSZGfUd6bg7uRxQO3YUZxx09KaOADn/wCtTgCBTTnkdc0h9QPrnkUYyOOtKoOSD1p3kyPkqlMOVsZk5IHFAwRVhbN+jHA+lPFmMDc36UuZItUZvoV9p24J4qPucVoJbooHen+VGOi80uZGqoSuZsZO8Vet4XlACjA9aeFC9uamicqwI7UuZlrDrqydbZIuerHqajdT6VbY71DAcGoG5od2axSWiKxHWk9xTnGCaaMdKkseKRun86PU0p547UbhYoXoLx98Z7VAzuEjChhwT1rV2An2pzRowBKjP0pa3HfSxRW7dQ7PzsUEipftuOfLJC43HPTNPe1ikcuQcn0NNazDM2H2q2Ny+uKaYWQ9ryFXKlvrx0qbcNu7PB6Gs26geMMowUkYcY5q1drLJF5ccecjrnpQDWmhOXU455PSojdwiTYXG4daoyW0/mEKp4Awf5010LykBcbiMjHX3ouh8t0a+aYScVnR3khMQzksxyPQCpIbp2aMuBiQnGD0p2J3/r+u5dHtTs881U+2IrOCD8pANTecm7buG49BSGS7vXkVG7dfSgnimHODSYbiKKtW3ElV1HNToCq7u9OIMudqgnk24XuasKQygjvVS5zuFNsmxEBvOc04cHihV+UU7HNLUphGAi5rIlYyuzsevStlhmMisQjaOfyqo7nNiNkiNs96Z8obAPFPbGTxnHWo84wSab1RgtxMk9RikBKtnHSl4DE569qaMAUinoWo7tlxu596tpdqeuRWZkinLnBH8J6U7ItVZI1VuY2xh8/hS+emeGrIZscDpTs+lFkw9tJLY1hPGM/NSG5jBHzj8qy85OM8GlwMN+FLlQe3k+hom7jzjd+lMN7H0zn8KpYGfX2pudvHUUcqF7efUvfa0LdOKPtMTnB/lVDHXIwPX0p3QHn8KHFMPbzLRigkJxj6gU1rQKvyniqinaM557Vctpy/yP1rGVCLOiniG9GQm2YHgikS3Yt2zmrzD0PNCRlmAxkn0rJ4aLOn2jM6SNlchgRTBjIHrWrcxhZtp545zVV7WMsH28+xrF4Z9ClUXUq5HI6UwjPep2gKtkdPSo5Cc4I6VjKDi9TRO6E5xwKiYHbzxUp+6Oaidfl3E8mkh7ixnjBNOIx/hUaYAHrUvGKJbgKF+bBqaMbiR2qAfM2asxcMOwqJDFPOAO1K7AjHtzQ4K9OhoICrzyazArSYCcEgiofMAUADJPrU0uRk9QetQohJzjBHeuiLVidRY3A3EqR9KldtsajPXrTZQUPynANKV+Qbh070nZ6jsSEZXd2ppIVyQuBUYOT04+tPdsj3qbWGPQ5jz3NOJyMn8qiBAGAODT0OAVxSaENVBycdetODKF54zQRlCSPamOgKgg4I7U1qNiwnDEEYzSuNrYPSkQjfkcilkwrbyelD3AdGWJBNPx82M9e1RxtuAPTFSDhgah7iQrdOO1RPuC5AzU2MZz3pmMA5PFCYERwo4HJHWldcJwcHvSId0gzx7U446tVDIYiS2TzinqWww/nSKMHAFPb5ehqm9RCdVznmo5XBgx09cU6RVIIHX0o2b4Fz1oVtxj7Vw6/KcinMcjHemQJ5Uf1/SntkEYqZfFoF9B3II7inN6+lIxOcUMeMVADJDlT6mq0bseGGPappRuG3GSKZE3z7ZBgjpWkdhdR3JYDGQf0p5B2gflQuMk9hTnAbnt6UmwIo5P3zoR0xinxjlmPeowNzk9x1FSrj8utEgA7VYe9KVxg/pSAhiPUd6cMt+FSG40HL80mc5ZulHOSe9DckAUxinsacSAfc0hJyAvpSMQvBpCHNgYzUOctmpJT8vpUIYMxz0FVFALgl/QCl6tkDgUhPBI/CjOFGBTAHPysTyB2qnIWaTOfyq3JymAe9V3O1/YVcAGiIJnscU5cE9aVs5HvTE+/tHarF0JFQrLuBwM8j1qyWwd3vUGz5QQcEVJg7Oe1ZysyhxwM80kakJyaByATSJgqc1PQQrnIH86bnkcZFSOuVx3oUZX6dKL6AN9M0zd1c9O4p7Dnnk1CW6+lNagTc7FyfrTDnkgU4nCg02RsRggUIGIcBQAODUAfy5goGRUudsYGMCoJCpkBB+atIq4MsuQD6+1I3IViab+73Hnk0OBkKTyKm3QAyqrnPHpUq4K9OKps5TZxx3qeNjtOfwpyjoJE+SWpkoJcY6U4HgYHWkdSzrULcY1zgEYo+9CpPBpXGc55zSEDaq9KfQbBfnQe1OxkgZ4qOPmQqp4qbgAjvSYlqQ/L1xwKQAhRnkGhgdhUHvzS5DYFWBHIg3DjPpUcnBbHSpXbaTkdBUTAKhOfvVSuHoKmdhwKUENTF+9149KfGvzE9qbAeeWK9jTXVXIAPIokYrggdaU7QA46ntUoCEqpnyP1qCMhpmAGM981MvNycimR4aVkUY9q1WiF1BM+cVbjHao5CwxJ90j3qzAoxhjlhUMkJk3bnwB2ppq4CQsGuS5XORwacYwJ1YN0/WnWo2pjGQvGadHGXnchcelJvUBzSbyD19aAR53I4pjZZvmGCKmVVfDVLsgsxt58sXA5J4rGmi+83Gc5xWzdKXjz0x0rLuB8mP61pQdkDsV45Ajbuw7Uxn3vknPpTMEsOM+wqxGir1Gfaup2WpI3B249e9TLISmzP0xUZQsSVB206CJic7sYNS9hpE8ILFQTirEkBUjPXtToItqgnmpDGS2WPNZt32ERiMx8lcE1XmAc7RV2XPA6ioEQ5LEgKeAKmLHcrhGZC38NU2lIQqBjJq7dBUh2bufSs6YjcOMY7da2grhuRtuLHPSmnI7+9SAlW5IJ/lTAQSQRz2zWqGN5YelKFAOTTlzjkZFDqfWncLi/N65pMEnGOTRg556+nSghuucUhEjKRj5eacsLlge7dBTR5hUA5wKmgLb9x6VDbSF1LAjB4J6VHhUfnp2qYRsTkHmomBCncOayTG1ZDXkGcDkdqaRuIyKbyWGB070buenT1q7W2JQxjh8cmlzgZUA560Y2sev40mcZNUGghyRTGyCGHBp7HjPrTJGJABFUgRrQcxghsn09a0V6cVl2pDQH0HStNMhRgVyV+hUd2SdDTlIcn1pnen445rGna42c7cYE7egNXbVt0QPSql4MXT/mD6VYsmBBx/wDrrqnrG5K7Fi5Y5XPGahzxgHjtU1x91W7VCucjjgd6zTuihQpJyKZOMW0mT/CefwqYAJ3/AAqG5INvJ6BT/KnF6jSO49qXqaCaQivDDqKemc0tIemKPT0pDDHGaQdM0p96TpxTFYMg8009OKcOaQYLUAOUGl5zSd8UNnAANINheOaQcAilOMYozgUDD2pRkCk78Uo9e9AB0NAPFHekNIQdqQrkdeaft/8A100jHSmMTGMZPNSRrupmec46UB2XoBmtqco895Cd+g9kHOTUajrnoO9TB1JBYfWosZyc5Fb1YRlDmiyE9dRQcKR1obp+FNO0HANG/I4rlbuixxAPIqMsATTgQeRTWYDqM1KGNVtw6VIOaQ4IIoX7oApsBcAHjpTXGRilbg4zQCc89PWkrgVYw4VyflNS25DRetOZVCsuM5plsixgqp5Na3TTuJEw5HXAoUja2OgoyAOOaZH948/hWYxAMAg96eOwxTSS0nPajJVTimA9j8nAqMk7cA809DuTcetJ04Ao2AIiyrgjI+tPP3s9KYWxgdM05h8vzdqXW4D+tLkk+wpoYdacDlalgO5Ix+lIp4o6ClAz1qQDqDQCccijAPAoJ/hHagQgHzM1KMlc01yVXIFOzhRTYwHK5FJjnFL0A5pAOev0oEDntSYp3Xj+dIxGaEMTPIoxSdOAeTS8H60xCHlhQRg89aUYHU0zcAc0AKnzAkDJp3Udfwpqk7iOgp3TIpsYYwRmmjB+nelJAIOaCMikBGSOvrThhfmzyaTaDwRxSgZGDTAU525A4oXgEKcU1GOTkUAYY5oAUk9PXrQyruGT+FLuXcB60pA3HigQ0JvHzdRT+2BSg7u1Jkc+tJsYEfKPWo2GMCpCTyBUeQ30FCAaTucegqTq2B2pmzHIP4ULk8HimwHDJbA70P1IpVGD7etMOWUkULcB/AAA70E84HUUAZGSeQKXOOg696OoEKtukKdAKl+9igbeooU8+1DYugvXnHFA5BOefShuOPWkUgnbS6DBuoHb0pvJGelPH3vWmsCehppgKR8oOaDkA460DGRTG56UgFzjA796DnfziqxndVLSevygd6f5uV46ntWjgxE27Jx39KHOR3zUTSYkAHc81PkEVLVhkLjp/KlXIOSOKbuzkCnMPlFPyAGUnBBwuaeT8lNfhfpQchRih6iI1Y5K1IOBUZBUbgOCetP3DcP84qmrjEJ2cAfWoupyDzmpjwDUZXA3Ac0kBKGGeTSodwNRKqSKXPapI8bsg/LSskA4DPXpUMwO4sDwBU5OAQKr3S5wR+NEdxMigAKZHAzz71OcblNRIAI9oFTsAY9p4x0q5PUfQjLsWPpWW8Rd2OMgd601JCHjgVUELeS5LYLHOa0pPldxENujIGwMmmTKQzOcZpHuGBCxcgfeNEjgRls5J/St4xd7sV0WInUQgkZbvUhdhEQDzVWAEKMc81ak+4AvFRO3MMmXCxnI+Y9KoOzPkYxjuatMzeYik5x2qJ0ALKRwamNkweoiMWCkHOOtOEjq7Z4SnRKkaHHUmm3Wfs5A+9nrVNrm0DoS4bzM9VNPlK+YNwxgcUkG4RAucnFMIzzJ1PSoe9uwMoFP9Kd2PHaliDxoN3OT0pLoq10qrxjnOadvLH3xxXV9lMSHh5fNBz8p/StMfMmAazYwzbkAAJ6n0qzbsdjRlssB1rCrFyWgyYYBI701MtKRnpQAAdx5ApQVKd8nvWKsAkxwOBkn9KhOSxNStuDg44HU0yYr9xPujrWkBCq3G0cn2q7pIYS4Y8gVkW7N5jEdAa1tMOLpc9e9dFP3aiIlrE0psCqU4yOauy1Ul5r0mcyMK9TEsZp43Iq7uakv1/ck9xVaKXPBNODZy17o17DUfJjkt5GxFKME+nvU6TS2lnudEkjlyM+uKzPIGA6/lTCzrkNkrWxx819epO+05z1Hb0pFjDHI6UxGDHrg9/er9vbmOBZ8hlz0pMqKTM8DJx1xSNwpXqKmlAMpYLhT09qjSN2bgZoHyvoMSNiduMk1ajs+d0h5qxGixgjv3NPOMgVDdzqhRS1ZTuYxGBs4qzE2+JW7kdKc8ayKVbpTYYfKTbuJ54qNeY6VZRFYUmeTSyKV5FRnODTEOJpDgj3NNIYjJzRnnrQOw7GBSgc0gORThz0oQF2BtyFTTHBViKZCxUg5471YlTg+tV0Je5VYYPSojkGpiOc5qNuvNTYa1GDmnAZFN5GacvrQMkUE0jfrUiimMKAbGr7ml9aMUgPpQhMDgimSzpBGWkOB3NPPpVTUFxbiT+4wP1oYMU3bHHlwu1L582ebZgPqKWa7WHyyRw/cdqnDBlBB47UlZlPQrxSwSPtACyDsRzSG0jXJTKt2OelTMiFg2BuHQ0ZNMZnXUTQ2xQHd5hAJ70qRZvgPLIzjcc9MDitAqrEbgDg5FBxwKSbAqXFzsmQKwK87qfBM0hZXTaw5xntSNaRMHKjBbuKiZTBvPnL5pHOR2FJuw0kX4WWUZQ5HTIqd+gANVtOj22aHpuyasPniq0IZLDJxjPFRXA3SCmKxRs1KzK3zCnuApXHXrTRx1NLuyaQ8fWlYQjnris66RlfevQ9avtwKjIDAg9KE7ClHmVjHYEckcUxwc8mrVzB5ZLDp/Kqpwc4OBVaM4+VxdmIeBkjJ+tMHOOOnvTjgZHakGR3xQ3YdrsU53HHIpQcA4/GkHH3eaTOT1oTG1Yew9OtLt4wB9aaBuX0NKD701qyHsPHGB60uRwB0pMfn6UA85zVWM3dMAdp4oIGPejPJ9R3poPT19aNxod6mmnnaOopAcnA6npSZBH070ilqLnPJ6d6cuVcfzpoj+XJ4pQo7jJos2Unbc0iMgE1ZtfmUsBVF34CZ6da0LZcW6j1rNHffQhuP9f68U0qO1PmGZfwpoGaVh9CMx1E8QZSCM/hVk9DSHpUuKYJ2M6W3KDKZx3HpVeQg/LWozAA8cVnSfKx4riq01F3RvCV1YjUKW9xThnB9fWoVU7jk4FSjA71izQcgOOT3q0qnbxVeMDNTqCV9s1Ehj+pINHcgUDO760dWweg71mLcqzg7h6jtUaAlh2qzLg5PU1BH3z1zxWqegaCSgswOeKOCQx/KpAAck/hUbOQjED5qa7AKOJBxgChTz7HtTEcyKcn5vpTy/wAofP1ptAKvykk9KlXqxJ4FRghkJFPGCvXg1DGHOMdjQQNuD19KRhlc9MU7B4IPGOlHmBGuFbbzzUpjVoxxk00jLk/pT/4N2cmk2A2IhiUHan8FtuOO9RKPlOOD608MPLDE9e9JrsIeDliT2pGX5OKXHyA012+UbevekhjFJxnvSFgW47VIRyABUDAgZBq1qIUffJHNGAxMgOaSJcA5796cgwNp5B71TsDBhukBHINIrFl/3e1KTtHH4UR5Lcjr3pLYb3HPkgBaex5GKbGCWOT0pc/K3tUAKTkA0j8KCPypAMJjNK5+UY70BuNONwbPFRfMZNpH41K6/LTI2Bzg/WqWwMdkOMA9KUSZPFQodshJFSIVJOKbQl3DlXbOOe1OiXCHJ5NRgESE9akXLsCePahgKCViye1KOue1AHBWlOWwF7VADcdeOfWmu20cdTTjuIx6Uh2qp/nTQCrx+PWlXB4I5piHrnvSliCTQMVgSTjrUBGwA1MpPp161FIM42iqj2EHKnBPBpB9zg0kvKdfu1GhYIcfePSqSugJgCo5OeKrsWCEEAc1YUFo8mqzkI2TTjuDHMPl+Y803ZhkOetQNcblPPPelhIcYY8g8VpytLUVzT+XYM9BSHGDnoajGXXbnPtQx568CsLDFjHUkdTUijCEdaaeAMGpB16UmwI2JPPagtngYBpzgkk+lNjHOSevajoAj5CcH9KQhTtJpZCMcVXSXewBHToapK6BlhgDJgjimcAM+ePSkaYb8UBl2EHvTswIiA6A785pXh+bcn3u9KSFOPTpTmkXJIqrvoC8xmOM9zTXYBNzctUZly+M04KZG5PFVa24tCNtz4zyP5VJAwZ8A5FOlG1TRCAjA9j2p3vEa0ZYU8n2pR1BJ6UIRuPegg4zWIDH6ZB5NB+6M9TRIMKF6UgcDAJz6VS2AX7km2nbiTkVG5xhj96nE4APrSaAiwUYjNCna4GKe2S696a0Z3kg8elVcXoI+MNUUjFYxjkVKygq1McHy8Dr2qkAwYPzDrT1yFApm4BhTskZY9KbGP8AMBUnGQOlMdT5GBz9KkQAJgHg+tNGVVmzxSAaAEQt1ZarxI7P5h4NSrlYWwdxNRANHCSTz6VouoiSMbZeOpqTYA7PnPtSx7WUSng9KSNSJDk/hUNjQRDMZHQtTYy+GGelSRMHU4AGKicMko+baD196Fq2hdCFizSZUZB61cUEQ8DmmqixOqjknvUik78E/hRJ32DYhmz5bKTg4rHnc7lGDgjrW6Y8yE+tZN5EUkcgDA7VrRkr2B7FMHy5BxmnoT5m7OKhchhkfe9aesnYDBrqaJJJJMHBY4PU4oSXC4H3jULtvOaaDuALcGjlVgsayTlYznjFOS4z6gistJHztY5FSOSw+U4qHHoGxfnuVyAOgqubgCPbjJFVZARg5+tRMRj7x+tNQBofKWaTcx3HpmopBjJJ5pSSTnJ/Cmk7vTFaJDXcTnnHegAsCD160qqpOc/nUqYCk9/pTbHcYh4xnmkBzySTjvTmyT93ApoT6cUCEYbiR3HelTIfn8vSnpxknFIxwNw60r9A8h+8NxnihX2S5Uex5pi4DepqUQs6sQM0tCdi7HJkZzim3G3G8HJaooUdTtbOPSnzbQQAeKwtaQ90V+RwcmkGQxHPvSt8xp2cRk461qSROuMkdaYQdmdxFPOAMH88U0MAKpFChcDcOajbPIHT1NPHBJ6CkKgAnHOM00F9TQsSfKweRitNT8uD0rKsWPlEZxxWnGBt55rmrrqEdydRj+lPXuaYADjFPB46VzR1Y2YOpcXLdOewpdPyFIB59KdqyBZs+vrUOnlvMxXb/wAuwRpy8wKMc1WZvmAzVqQEwnvzVQrhgRzWMRseFByc1HMc20mBj5D/ACqQAdxyBUU5xbyA91P8qqO42d3z60px2o4C5H5UAd68EYduKTovHWlIIOc01ec96aEO/GkI4FGM8ijkjFIAHB6UgGDmkbIOcU6mMXuTigfypOhIoAJFIAzmigHLYI6Uo4PNAAfbrS9OKTqcilI4oATIHNLzgZpCBTvTmkAnfNIfrTs0h+7QAd6ae/H0p2Bz700dPYU0AhGQDjmnIwCkGljCknd0p2Yx05INdMU+W6ZL3I1UAdKUYz/Sgn5s5pR8vNc7KGMCOegqIMsi/KcjvUszHY2Bk+lVrZ9qkOu0mritGxEvmAjjgVIi/wB2omX51P6UqsdxUfjSa00GSkc80oAH9ab3p/U47ioAjdQimq1tKsrNxg1bkUFcetVVKJPtC7QB1rWnqIkclchetQxBkkK5y55NWSgY5HUUiwjfv/ioUkhjFL/eI5pzEKyknr2p7YxxTWQEUr6gOHAp2OdwNIRgDJ60pbAHFS/ICME8k8HsKWTLqvOKXhgSTQoPGOlACL8mB1zU2QMGmDGBmnjBxmkwA5JUjpS8/hQDx14pVzjFSAdDignH1pufn+lL1b2oAMcYpe3tSHlsA0o6kUhB2xRjAxR3xSH73qaYw5zS4yetNznmlIxQAjDBpuc09ueagAZZGz92mtREhxnnmoZW2gk9u1RC8Ugt0qtflhGjBvvda2hTbaTFfS5fRsybycg1NwWrMspWaJ0A+6ODVu0LNFz606tLlHuTnmnEnv2pnQH+dKDxWAxjcHNOXJO7tUTgy5XpUigrFjPSrsrAAPJAoHykbuc0xT5ZJZqd8uzk5NKwDguW5FO/izQGycY60J3OakBSRnpR1OTR3OfypevSkAwDuKTGB9aeBxgdaa4zkdPemAYwuKNoK8nrScc8/Skxk57UwHFgqn0pVXK5/Sm4yPr0pw44pAKCFQggZPSkbpn9KVlXA9RQVGAT+VACcgZwPYUi53bqVjxnHFNAO/Han0BARyCeaSIDJFOI5NKv3Tn8qEAdHxSEdqActkdKAw3ZNKwChRwT0prA7TxTh8yjFJ2wecdaEAwDPDDFRKF+8eADUpO5SO9RsoPy5q0AvlK5zjBNSA5TNIgHPPWlA5wO1F76AQsSDxwKcTggnpSORnOPpSK2WIPSnYQ92G1j09KVOYwc5pGA2tmhF3RDacA0ugwnIwB2qudz7cdqsuu5c4pkaFN3p2FVF6CeoXG4RkR9aRf0NOkDMCAPpQqfIFbkii6sMXZwV7elG3bgA07HB5xTYzlemSKnUCXAODTWILbaQHYpzikBABf1pLuBGMKzYWnn5o8Hio1YbiSeBTwd/wBKtgiNQzxkr271E4CxMpJqwgwhAH61WmByV3YP0qlvYTKDx+XKAoyGpJoj5fIyewp7shkAUliKcVBj3Ht2zXWnohbobbHBIz+FXYlAU7qpQPtUNjnP51bXcyMccVlUWo1sCH5iyjkU2Z2ZFZhhh2pYpRtYgfdolO9ckUrWYESSIVZW5I6UkrMqjB96ZAgEzMMU9zmVSvJ6YrTlVxdCyxKIrA8mmvzFvY4IoMgMyqy4A6GmyoZWbOdvvWVtQZnXTo2zGSW/iq3EA3ydwOKrXI2YjQ/d6n1qaFX3KT1I611O3LoCvcswoQ5z92lt0SMM2Op+WnrHkktx61MCqqRwB61yudh2Kpc7sA/Wpg+UyKr4YSkAZJ71YTdtXIzx0pu1gsxJ9r7VLYPpUMjmNwBgg/rUpjLvkrg/WopwIgG38n2pRa2JfciiVt4x0HUVqWJUXgIP1rNRZNxccj1rSs1CSh88Z61vf3kFtDUmHQdqqPycVbmH5VUfqT6V6m5xopXCblIIrHjXCZ71vSrmsaRdly6460RVmY11pctwSHywGNSnDLgj/wCtVS34bbVlGYCuhWaPOkmpXRG0BVuOlCTspA5461PkZ+XjPao3j3ZGMGlYFIsSPFJaoI0IkXOTn71X7ewkSxlmUfKqhn9hmqEVq0MgZznFbw1MNpxh2quBtcY+8KzlLsd0ISV5GQxHY8+tN3Yz60AHFGzNSdHqIXOaehNIFA5NLlR0NMZPwRjtUM6pG/yHcCKaXAzzUZcHrRcLMA3ODRnjikLD1oz0qRi9McUBsH+tNzSgU0DLUfParq5e3Z+6dR7VRQ4HWrtnIolUHG1hginbQllZ1I7VAcVenhaN3jbqh7d6qOMHmjoNMixzT4+KaRzUkY4oGP5ppOc8U5jjimE4yKBbjWyQMHFHTrSBiSRT16ZpDsNzjrTJV8yNl7GpHXcMd6aBgA0b6AZWWe2VWUtJC3zD2qeNEaUrGzxMecGnzxPFL58K5P8AEvrURmNzcQmBDuU/MT2FJdylsOKpazIXd3Y9D6VZ78Cq6QStOZHPIODxwwqxxmhMBQaXr+FMHcU7r9KGIVcHFDW0U3LoCfWnquakxhfenuGwiYjUKowFHApd4OPamHimnrTEPbFCnmmdenrTgeaQEhPHHSjqcYoH3eaO1ADHPNMzzinNUROKAGOMgjsazp4TGSyDIPb0rRJzUTDK4I4pX1JnFSMwHLcfhTM9KsSJ5b4xnPSomHykY61WhzuDQnAPTJqaIAAkDiq+3HTrT8YUelO2glJp3sTyRr8zZ9BUbKFPHJNIXI6Hg9qXeSV2jjsM0o6ClyyHMhTCtx6UKecLwasckA4BYdPamjEce9xuY/hVJilSs7pkA5NIAT9KneNFmPzbQRwKZ5ZjTeePajmF7O25EFLckcetOUbcAUmSpFKBnp19KZNwHBDClX5m2gY9aCehB5p8Yyu4d6luyLpw5pWJVBLe9awBVFXPOKzIVHmKMda1T2xUK1juexVm/wBZxTeQKdMP3ppBwKAGvg5ppOBSls9Kikb5SOtD2GtSGV8tjtVOZiHqwx9KzLyYrIAuCa56qurGsNCXkFQeD61KvbFVEkEnSrkX3MD865JKxqiaMcnvUobAAPekjBC59aeVwAawk9ShTkUmCcA0ZJwe1JzuJ7VItxsv3zgdagGBnB/SrLE7jjoRVY4UHPWriA4E7OBgnpTHXKbe5qRcqi+ops2WbevT2qluBURWSYq2APWrGMJwevWq6hixyPxqYYJwegrSW4IkAx8o6U4t5fFRoe+OBTh8/PesxiFjkDt3qYBuhPB/So1G5sZxUke4E7uaTBDZMZUipGHQY4pp657UsZJPUVIhgy8rKBinkfe/SmqR5o/nUn8eaGAhI24zimSHYBx+NLt+c549RTGIbnGR9aFuMfnvng9KjbAXOeKk3A4BHNRP8tUhDVXcOvSpYxuBUnpVUv5YIPOelTofLOAMk96pp2BCgDGP0pVPG7PA7UEHZgDn60kWFTpwanoA5PvFx3p5BCk54pqENkA9OgpxJVT79RUvcdxMggHv2phIJ3dBSghVz09KYzBU9jVJCHl8jPQVC0qgtgfNTiw8v61TeUhCWWrjEGWlkAXnhqehOSSMr2qrGUmAbtjirYIEYHeiStoA5hwOOvWhBhAc8GlGARn8KTPyEDtWYx5IDZ60A+hqtvbOPWnRlwCD0HenygTE4fAqGbcsJPcU3JRwTUjAPEwJoWgiOGUFFb1qUOfmI6VSVc4C54qdMgsp6etXKKAcH4I7mlPCjJ6UwKVH8qeq/KS3akwIpSQPlXOe9OUjKjpmmzfdwDTVB4z1HeqtdB1Ji2B8vT0rNYPOeh68mtBGJYkdKYIsklePeqg1HcCq0Y27cc01EeNgc8VNIpLgHtT/ACgVyOtXzaCJY3UAleTikjJLHcKZHGEIKdakyRz+dZNWGP24apAcgn0pF45pE/nWbCw4thT71GAfM3H/APVUhwSCe3amLhmJ/Omg6A6naSTwKrBdgDgZzVsnIPp6VUcDcBnp1FVB30ARlZW3A8U4ZD8ngCnMo2gn+GogxI+Y8nuKvdAMkYuCAKpNIxG3NXXOGK557cVVeI8EVrCwnqNRSDy5Bq/C/wAhHTFVxESODmrKRcUptNagiNyzHGeO9SBf3YIzx0oaMJ07/rSoeNo61O60AfC43N61KRlR7darrHtk+Xn1qwecgGspbjQjLu+Y9ulQxEtuLc4qc5LAdhTFbkjFC2DqMfL4Y/iKcrE9O1Mb5WzQrb4gRVW0BD0ByW64pUBAO45poYqoXuf0oBxlSeaQB/Dn1qCTgZJx71ICWyMdO9NkA2rmqWjAgRQ3OeTUxxgAmoY8YK09mGcHnFW1qC2Hj51Cg0FF8sxnv3pxxtBUfhRLJjkDPvU630AgjCgtEp570sgEUQU8+tPbaHD45NMmcNJtK1Sd2HQUKWiHzYAqRVUAMOSBUZzsfngHg0SkrEqp1PeluA+3JOSRigKWk3Mv3e/rUgbEQJGCetKTlcjr6VN9QsNlALqM4Ap6DahPbtVGORmJ3HBFWFZnjUjv1pyi0rC3JxgjdnioJYI5k3Y5qUD5Sq00NtQY/GpV09BmFd2ZhfCjGTnIqspwpyCO1dNPEsi5YViXFt5T7cY7/Wu2lW5lZiejKhXjI/Sl6j+Qp/lPkYOKeECtg9uorW6EIqNnOSMU7KnIz1pXZuMdKiAAYn/JqdWGxG2c47dM0ioXAUg59ae7DIGOD/OhZCr9+mKvWwXGn7u33poQckjn0p6ocZPOTTiCOvancXkN4CkbetIcEHAOR3pcc9sUE4wuPxoDYXDcE/WoyrZPtUjna23JbioxnmhDQnJGec04fOOByPSnDp2xinoNue1JsV7Eactz+tPdsN8pwKVFO8kj2ocj5vSjqK5YSRnbkZ296mdVboOKqwMxJA60sk53EdCO1ZOLvoCtswPDfL+BpnDE4NG4MQB+GKQjHuatIFcQj5SD9BTFBA9+lPCjac9femkbl5zTHdWsNPOfQdqVuQO9GD1/MmnOODjrTET2IyG4zjpzWvFjYMcVjaewV8mtiE7k5PWsK97Dj8RKCQxzwKmQgrUYxjJqReDntXKnqNmRrYyYzVOxyJK0tZTMKsDisyzOJe/Su2LvTEaxOYGx19KqxhguKtxj5Gz0xVZSd3TpWEeo99x23B4PIqKcn7PKD12n+VS7gCcDrUU/MEuOyn+VVHdFHd9RQpHSk7UCvCAB60i4BNOHvSbRnNMBR60g6/WgEHil70hjcfNS9etNOQwCjI9aVj0zTEKfYdaU4A4o6dKQUDFxzmlFIScilzSAByaOvNGM0E8dKBBjJzQDRg4xQvQk0DDqRSkZ70cUZ4pANxSCnY796aD045pgNDHdilQ5fGOKQn5utCsVI4xmtY2vqIk70mctg0Acc0p+9WQxrgEEEVXjk3ZLDGDVvC8881AFQMVFVF6ANIDgNnj1pwIJ46etGAyso7UKNhA7VQDiQuMdT0FOLcbhSHHBxTX3FfkqQHglhzUZwX6fjTsgDPfvSOR3OKqO4WDO1TjmnrkioUXGRmpULY5FKSAjZdj/AFqQ/NjuKV/lGcZpF9cdaXmAh4ZRinFgASRQoAYnH1oGCOTxQ3oLoMCZINPBGDx0pGOGG0daM8UDAjjmnj5V6cU0DmnrzyealgItOBAPJwaaOKOv1pMELtGQ3cUoPGaCc8Cmg5GO1ADyNvpSAY70mfyp1AAOTR70opMUgEHNIMdPSlXkZozhs0wDPTpUcgJB96kXnnHNNYHGcc01oIxvs5SQK5OD0xVmS2ZoQpOcdvWrLg+YBtyKlBBU8Vu6krJjtpYqWyOI9vTtirMcYiXaOSe9LDkL82M08kHGOgqJzbYEZIzinbsrwKaeTTgflOOtQBFucP8AN0FOVgyZFRZfHPU/pT1Ji4PfpV20AGALYxRCVY7SOnapNoZwQvHrTVXa5YDk0r6WAkB5B/MUBtwo4xuFOxyDnioAQc4p3CtSNk4ycD2oA+akAY28008H607rz2pjEZzTQC8c01m6e9KGBAJ60NwRxTAXGGGOlCkAmg/f56CkRwTjPWiwC7vmpw5XJpn3eTTxlhmkwDAIGelJ/FmjjgA0MTuAB4FAB0ye9NByfen8EYqPbhyc/lQAqnDtinEDGaANqnj86YeUxRuA/oKQ8HHrSbsDA6U4kEe9ADcdR69ahdRnH51LjqM9Kry5KY3YP0q47gyVVCpgfpT16HPWmA4UetOznGKJXAa2Onc9Ki+dCBj61JKo3Bs9KiwxlPcDpzVR2EiUjMZ4zT4VxEO3tTQCBz09KSAu5b+72qXqgJR90AmmkHHXANOA5x+tIxxnAzUjEOSM0e/ejrkUg/8A10wBVyGyeTSj7o5pTgDigj5ff0ouAHJY8cU1topxORgmo5Dt+XHT1ppXF0IpQoVgOT6U6DO0lu9IcmQnsKeOY8jrV9BjyO3aq8mQ2SKsLnbmoJSWx/KpjuJlGQCN92MM3pUaYbeXyRjiluA5BCnaw6GljjJhAZsgdeOtdi+G4iCJuM9cHirgLPHkk5qqnG3yxxnkGrWWH0oqO9hofF+7Qpjml2Ehlc5z7UNuVgN2M06R3QA9T2FZNtu4FWGKQXABPBpZFWOUFR3/ADp6SIHZu1V4YpDvZjuz0rVXe4rGjtDgOOCOlI+9kXJ5FQWR2yMrnO7pmp8OFfBGe1ZSjZjKd9GXiLdSccVFZy5DK35mpJIZfI2s2HY9ahtxsmEbde9dEbcrQupowtyQxz705kBZQowD29aZHGDI3PGOlTLhRkdelcstNhleQnzd3AA6VMnKZXrjBqtdhioPXmrFqT5W41b0hzAGzplse5qFlDRY647+lWiu4HjJqpNCPJZS35VMZXepLutiusjyShc/Kfu+1adsu1AjNuOeRVWOAJGjKPmI6ntTkkHnKiD5m6mt21JpRFY6Cb/Vgj0qkpycmmSX8zB0gs3kEeBuDAZqG3ecsftEYj9ADmvVXwo5uV3uR3l8YC48h3CjlugrO85pXEjjaG+6OtSXskTXszXG544VBCA4GTSRBJI5gsJjdcHBOeDTIklyPuOjOJl45qyVGSeuKpjrnGMe9aGAxOPbitos8uoiMgqMHg0oOSATmnA4Ax2p8Sl5UQHv6dKrYhas1mUSJg9aquhRiKuDrRIgcfyrBo9YoAdhTyo24BocFDg0wSbetJj6EDnb9ahZmxgH8amlYNJ7EUwKDxSKIwrHnvTwpFSqoHNB96Y7keCOO1L94CnY9qcq8etLqA3GenapFX1FKFGBxinYpoWwuR0pVfZiozxmlPoO9Mk0XukniXjEq8E/3hVcoCDiq2cH+tOEhXvT0Cw0gg9KeOuM0hbJzinFlA/rSGI7bcnFRse+eKRnH0qPcAetTqw0JY6l4AqBDgiph1wPxpgBHFJ+OKXBB+tB6UdR3GqBjmgKFJAA5p+D1ox7UBciYY+tRkVK/WoT9KQAKcvWkWpFXNAEkY/SpD06Ui8D2pc1VgbIyB1puOv86kY03HakAzAH1oB9RTivpSdMg96AFU8U7OBTBwf60pJwRQAwnINRMTjmpCeeOtNIy1Fw6EZ6/wBKawqRkI+lQtkH2pFEUkYkXb+tUSCoIK4x71ojBPp7VUuIyCX7HrQmjKrG6uiA56/rUkChjhqZ3yeooJPXrVNaGF7O71JDGE6N17etM2FXYdx2pUcAgtnI71LGzPK2F69aepNotEaSOpO1sD6Zp6yjbtdMge9PlXguq8H+LNKyxxr8y/OegzTElJbPQiLhpDIy8HtmppJVZGw+eRgY6VGsavsVTknOaJI9igg5z3osJNpNjMZOB2pcEICKTHQ9qQ4AyeKb0IWvQdje4Uj6+1TDOfb0ptumE56n9KkK4XpwKybO+lC0bsmtVzMoNaXJOe1ULFf34PpWhknimti3oVZR+8OPxph4U/yqSYAMTVdznNAIRm4FQscg1ITlfeo3+UZqWV1K07hUJJ4xWA7s8jOe9X9SlyBGOh6kVRGCw71g2tzWIqMVPXA71r2zllAxzWMnzMa2LNW2g9RWNZaGiL6gBQtLztPqKXuTSNwMY571wFEMsmxQAPrTzjGBzSMFLLgVCzEZFWlcRLICFUioZUyV96mPMfTpUMjHYN3PpVRuDJCoKg0PlF6U5ANqkc8UEhY8nvU31Apy+YWDKOnaliZn5xzUsqEncDzio1baOOprW+gPckGR1XGKcrAsMfpVeRyr5ySO9SZKY20uUOopkGcnimrOQQF5z2qvIjSybR0qVYmjGQM8VXKragWnPGexpIiNnX8KhkmBXB+8KcjAthR9ajl0DqSOOBt6inli6hW4IqnOzb/lOCKsLkjJ/MUmrINxeQ4Lcio+snTIPaldv3mCaTGHBzj3poADF8E5BHaiTkjB/Cq3nbmK557VaRCQoxz3qmrBuQT8HFSRsztwfk+lF0B5eQOe5ogUGEY6U/s3Ddku4BR3Pr6UZDbkBxSHG0uOlHO0FRz61AxsSbSeeamI+XPehVwD/epwHyVLd2LYi6xg1AzZcqTkVaKgoPaq/ks1yHzx6VUWgAZYgEYqCcgjy2Q/WrxwhBHNNICjJXJPSmpWYadSnDbMgJHA7Yq0qsMFjwaeWycDoKUYIAPSiU29xkZyxBp4X5cmnD7uR27UbSV4qLgRsVxnbxTv4Oc0oxtIpQfkFDYhCq7RSBcD60MM4py9QKAIWUYJPFNUfNzyR3qaQZPIppUKuO9UmBGFbBJ6npUm0qoHUjrSbfn5HPtSplckjNDYDDgIeOe1RNkR8dRVjb8hzzUT4LAbeTTiwYkRwgYjmpOB83T0pqld23OfakmTIx3HSjdgRSH95u9etLnjHemEh5AoP3adjAxnOa0t0ESJ0x3pxX5uO9JEh3YznFTHJPHWs29RrYaFKcDmlC4HPWnEkPmk6YJ71NwEAOckcVHyM8Y9qkDb2I96CuXHPyjrT9QGqdxORwKiZdzFhUyuNxXrikxjOFpp2GVn3NIR1WkX51bDc+mKkY7IiSOKjyQmEX73NaLYWxFMnyEjJI96SMqUAHOKLklU2jqOtNhIEY9q0XwkkwXkCp1AznNRxcnIx7VJEDtJ6c1nJ6DHYyv8qaFAXJPOaeTggLiog4DFSalXKH/NnPUGn/KPqah8zJ9hUq4b5j/Kk0IHyE4700t84AHNSdVyOlQuT5gA4oQEN6G2Ag8mnQZaFTnle9LMhkxzimqNmABkd60T92wW1JQAZATzQCS5+Xj1pSArA+tKucnPB7VAEYTqW702ZQYxkVLjPB71GxLFl7CmnqHQpFiW4/Op9h31HEoGRjmp17E9PStZMF5kmT5nXim52uEHPrSZ2qTmmK25SduGArNIBTEC5B6Z602RwWMZ496EfzFcE4A74pFUPww+Zf1q0u4NjfLWOMRk8mp94WNQepqMkSLuA5HFNlzvVQPu80PXcXoWASzbe4ppbHGOTSI2MyMvP1qYYK5xUPQZWntiQGj4I9KfBbhQrknNSsAo44PelSRJEwpzjrQ5PlsIdjatRlTt6d/yqYnI64FJknB6VCYxhJz0wRUbwrcKA/UVKcB6ANowOvemnbVCMie2eNy3YdDVfYSfmHPWtiVw0wjHfrVfULaNIg4O3HUetdMKuyYW6mYcqBgYBoIXbyOT+lKAV+bduBqCV/nBU5FdKV2TpcYXy3Sl53EZ69/Wm4/PtU0CFn5GR3q3oU+xJCAFOeTTZGLKMgCpvLBYcYA6VM2wR7QM57VlfW5LM7GT9f1qRclOevpSGPkjpjpViMrHxIM5HWrbF0IC464BHtSblVVVV5PWp2SIfdGD71DIBnrzzzikmmK4xQrNt6CnQOA+HPApYiuCSenapraOOaQhhnA4psba6kZciRmGQp5qMrvBPbPrU0ltIuRg9eaiVSF/pSVugriqdrD1HSkb5m54pQm31z3NDArzyT60+oAvynH5U4oSCcDjrzTRyQamZQwBAx6+9S3YEtSEjPAyB600H5e9Dr0AakRSAcH8qoaegpbCDik55I7nkUpXHXmkI5BzjnpTQ0SWhCyjcMc8VtxMNvHFYMD/AL70rdh+ZeOMDmsa690I7lleeo4NKDjkdRTVFPXg9MfjXEtGU9SnqgBtaxIG2zAZ47k1v6h/x6txkVz0QPnAEdDXbSd4smxuR5+Ye1U43wTVyIZPHUiqIyJNv9KiNncp6E+OfSoblgbd8j+E/wAqeWBGc4FRzYNvJn+4f5U4rVBdbHe9KBkcmkI5FLnK814JQDgn0oHSjtQB19aAGDhvan5z24pGGRUMc480o1UlcCbpnFAPegjK8UoHGakBKVRSDrThnNABR9PzoOM4pO5oAUUDrQCM0uDuoAQ0vHHvSHrSng0gEPJ6UHIPpS5xRjP40BYQg1HuIOMVJzTDgHpzTQETZEo2jqacchuufSgth8Y5oxg8itb6CJiMj3FAB696ByM0Dn6VkMUjNRkAHOKkA5pG6nn/AOvQgGEjGRTOc/ypVOGAxxSuN3fFXsAgcliuOnepBgCoAWDk5+WpEbIHrScQHMAVqpKwTDOePSrh68VVktw8qsWJA7VVNq9mJ3Jo8PGHI69qdnA56ULjbig84Hb1pPcYjoHUYoU8EZ6U4HAxjA9abHnBz2pAOORwOtLkYpCSWxigj5uDxSYDWPz5YfSnfxYpGGQcdqTB4bNO4DiBxk809QPypnAIBNOjOV9TUvYBeCaC3OaXABoHWkAdF96TGAMjFKeGGKCTyaAG9DinAd6bj5ulKOCRigBRkk4FKpyOaay/KR3pR93B7UAKDwaaRyKUdd1GepPegBAT+NIw55NKOpNJIgKc00BEGEg+U9PSlX5eM/hUVpB9mBVuhqTK72bsK0aV7LYEPBwOep6UuOMZ5PamEhjnFPwSc1DVkAzBHGaMYbPSl2ZYt2pM5PFMBjglhjpS8lj09qR2KqTQhBG7tVWdgJlODjjNBPOPWkAJ+bjNKRgZHWoe4DmGOAabzS56U0seR19qQDiCTSjj5RTV5TPelVecKKAFJzmoSMN15qc8DmoSAHyTinEBFYFgKccg802IgHI4pWB696b3DUcDuHTNRuQrqSCSakUkDB61GWBkIx070LcCTd839KcDxjtTD1yo/GnAZUGlsAvAxigkdaRxlV54pCcgACkA77qnvTclcD1708kDg96jfJIweKEA4DL9aQjOaUPuHFRnzPOC4wmKaV2A88LjFH8AOaXdkgdqOMHNADMfPkmmtgE4wMdacBlsnv2prBXU5PFNbgVpd6SjAyD3q1j0FI8YYgU8BccDj1q3K6QDXB2+1CDHOOPWlYDbmkTIT5j+dStgA7scdKSBQm5c/gKQschc81KBzn0ovZBYM/MfalIJUHHFJkEkd+9ABxgfhUgDfeAHamnglcU9vu+561BJIIvmPJNON3sBICOg6ilwM+3eo0IUZPNSPyKclYBT1PFMJJfpwelG4sSAKa5I6H8aEIaRt3dzSwtleetMII388GnRBsDPFU9hkgztbHSoZcfhUwACFT60yRcKxJpJ6i2M6YARsC+B60r7WiWP7tNuGjC/MuaQBWZAeoFdcbtCK5G5+Acg1ciUTAj0qtK/lylAOD1PrViCTYwBXrVS1iNErNtHIyRUbBmRCT0zmpSCX2HnNRujeYV6KOlZRYepVlkGAB/EelTRSskBY4x6VHNGVZQFB/GoVQqCuCCa6FaURXZdjKRqH6t2pw3NPnPHXNV4VypB4T0qYquAgbHrWUopOwBcxrNIrluFPFU0KtOZlbI+nSreAkbKT14qsBGD5Y+UnqKunorA1qXYGKqMnlqlUjaWNQRBFYL3z1qw+Bx09hWU1rcfQinBaLIGamgAVMduwolX9xjviobSQMfYdqnWUWHUtL94tVUwqskmT17VaIwD71DJGuckHnrzUp2YmhhcGHngdqW2HzoWXDUqophYAfWljy0g2A7hW8bXsLUtrbxvPMiXTRsQCyis+M2kF/5SOzy4+8WyKvXP2KGSOe5UmRwFVQDkmq8k1vbO2LNo0B++R6168b2OW13/AMAc9lbSzCSVMt6Z60t0YIk28K0nA96hvbZrmWCRJGAU/wANQSJunQXEwbYflGMVOwNXViNRnIx+FXEbKDnFVkOZFz1qWL+Ijp2reJ5VRaE5OOR1P6VasEy7N3Wqn3e/14rVtEKW656nOaqWwYeN56lgdc0Ug6/Sl6DrWR6OhHcrvi3gcjrWcRjvVqe6HzRqe3JrEkYqxJ59aOVmLrxi7IuvzSqORWcxAGV4zViG5UKA/Ud6XKXGspbl/Hy001Et1Hkjd+lIbqLGS3HrijlZftI9yUDr6U4cVAJ0Gfm5FOW4jPO79KLD513JwMYpf85qD7RHnG6hruIcZ/SnZkucbbk9LjNVXvMcKPoaha7kYcce1HK2S60Ey8cCoHuolAxzVIuz/ebOaTn05H6VSijKVd/ZJ5LqTnA21CZ5HPLE+9MP3cf5NKFGz1p6Iy5pN6saeeSc+1NPLE5xSj5jzzS4B6cCle4721Qscj9mKj6danF3KOFbNQnoARz2oVVUEdfagOZ9HYtresANy9alS8jY/NxmqYwAaB6Y4osV7WS6mqrq/Q80/jGKyTkZI4HrTkuJVP3silyFrEK9mi/J04NQHv60wXit94YNPJB5H4VNrHRGcZaocBU8a81Cgqyi4zQVcdSd6UgetJnFAhr03Hp0qQgEAmmd80AhoHOKNueRS+1ISBQAh4HPWo2fPA602WQBgPWogTu570hokU5xUgGM0irUirzRoFwC5X3qpKNrYq63H5VSnwTR0BO7IcnIx1pJBvQrRn5uPzpQGbpUlNX0M8t8o9+1ADbc/rUpi2OQfvA9fSjBzurRHDLTQbgcbvvUu4hSoYgemKCeAMZzTcZJUc1WxGrJBK4XaOlSPNuDfJhj3zUAHAwCaAOMk0aMOZrQsWzLvyT83amyY8w5PXrUXIGTyRSjjJzgUtxt6coH7pJ/ColBlcDGRTZH3fKvWp7VQqs3ftRJlUoXauWU54p4PzCmxjIx+lXIocHLdazudw+0jALH2qyeO1UxIY589u4q4envVLYkqTHMmKgx9KlmOWOahIPNK40I3A5qvPIFUnNTuwC8Vj6lcYwmeTWc3ZFx3M6WUyys3Smg5BxigY6dhSAZzWRuloSRj5+OlbNqNuAOlZMCkuOOB6VsQDADGuevsUWsfN160rsB+NRq3zhj0NOOAx9K47DBQBwvNRzLwDTkPXNV7iQquc9KqKbYh7ZKgHk1DIO2TT4G3qCeae6gsCe9XswHKMIBSOGC4B4p4weP1p3t2qL6gyCUt5YyfrUKSADp1q24yzAjIArPklCHGOlaQ1VgHzEEEnjBqAszlsH6USt56AqMgU6GHIyelaxVlqII2CnaDzV1MEqDVAoN4GPxq7CwDcHkUpWsCuOmgWRvQ96ZbxtEzDqPU9qeZR5pxUo9OuaybdrDImXDHPNKgDDI6GnMCcgHrTI2CYUdqnoAwjLYPLDrU23LdKjVSSx71IjZPJ6U2MjSIGUHAFSImJdxPPpTlBAJ/WkXiTrUtsBjRJJlT0pPlKsgH4CnKxIJoUH5mPGKq4hqIqptH4ipcBW9umKQYP8AjSjBG49e9JtsYnfJ/hpW+4cU2TIUYOM09QdmDyKkBBgpUbDjI6ntQqkFuMjNKSAxbtTENGDk46UgPmDJ+UCnjJ4x171Hgg46LVaASE/LnHNCEBeRyelJxuAPelIG7A4HakAqkc9qcDxmmcpnJ4NAOFHrSaAXaAcHqetNQlmIxwKewO0GjgtjpmhAI33sKaQHjp1pf4iR1oU8jtmgeor8ggU09cAU7HPFISBnBoQCY/eUmcjHanNncPemqQDyOlMQhGTtxxUch5JPapfXFRP82fQdacQGR4znPOKl4Zd3eqh3LMCo+U9fargAUAdquStqHkVzCiTZAx+NRySbQQB196kcjzQf8imvGJG54x0q10bES2zjpjk1Z7g1VgiEZBFWW+4DWU7X0KQj/dLE8U2Nt0eSabJgqQe9LEAUotoAqJtYvnO6oZnKFSDgHqKnYnyyR27VmyyPJ8/QDsaqC5ndivqWN/7wRjqe9WP+WfBxVG0LGZnYcYq0imOM85qpqzsC8hCQ2VY5HpUQbYuwHpTiRtJLYz7VnPOVkZQMj1qoRctBXSJbiddhz1qOHfjJqAlmJ3c/Spo+gHp+ldCjZWFqy5G3z46E0jThSEBxzzUKMdxPTFOkXeOn0rNxVwWhKLjk7aRCznJP0qiVZCMjIP6VfgUsi96UoqKuMmhQ5JxU+MjFNU4GOgp5wDgda527jGMwI244FQSkb1I6mpyAT0xVWfdkEDPNVDcHsSSKQwIFO6LuI4PamOzD5T0Pf0pUyqHJ6elPoGor8gHpigZLZ7GnYzknvSfdIGeaQMBjdxUJyspOeKm/jIqKVcvnFOO4FR1YOCOnerAOYwffmmSZDrxkU8jGK0bukLYVstFwKaN4G0cEjvTy3lgAjOe9MBYTZIPFJDQkAYKS45pIyVLs/alRjKpZ+nakKhlWMvz1p+or3BYxGCSep496bGzGYkrxT2RZDtGfkNOkGYCV7d6L9wbGo5aYqwG3tU6OHBC9uKqAskIAGf6VaiORnFKa0GJIrPx9096YyFXUR4AFTMrSDKn5qPLIAH8XepUrCDJJC4yTSktvAxxS4DEdjThkOam4/IaDlSx7VBPMVUKp5apWYkMM4zUSqGY+gqo92IIIto3gZPpWdqUryNtbK47VrjhdgFULq23BmkfcB2rSlJc92DvYzAAy8D5qrOpUnPUdKuvtXleKryblfkE+ortixJ2ZFtHbp3NWbWPMnH8+tRxLVyJ1j4PalOWlg5iUR7XIpwgLS7cgYFJ5vm3GFIp7s/DqMMvasU2ndg30KmzLtu4NRldq5PUVaYllDdN2aau3ILn61aZDZWBcrzUbLk5x+NaNyYEYbE4x61TZtzltv4VSbAgXI7cVZs2Ky5CkmonbgD1PNWILwRQlFGW9abvYLk5DKzTNjB7ZqB91wSVAwOwqB5ndiWzj1pbc4468+tTytahsPZiMIccdaikXPQ8VLKyiUgDjvg9ai3cE/lTQ2gVeOfyqQvjAI4qIk4z29cU5c7sHk+tNoNdxjnrn8KMYUMDg+wpWBZ8+n60eoDE80xtAOVOO3eoz6GjJH9cUoGcenqaewrBCQLgdAK3LdiRyP1rBTd5w5OK3Lc9OKiqrxZXUuJz1pwOeRzTAcE4p4PvxXn9RsbcjNowPJxXNrlXUgc555rpp/mhYD0rmAcNnk12UVoyTZg/hxgiq8hxO9WIcbVI6ioJ1P2hh2qFuPZaDFBL+tNmCi1cA8bT1qUAjjuKbcKDayf7pqk9UM7ntSgcUgHtTsV4IxvU8UE4pccZo6mgYY4FV5BhmNWDnOMU10DgVUXZhYZA5dSD1FSg+tIqBG46ml78USd3oFhO2cfSnc0gzS98EVIgPHXmk5zS9TQAT1oGIFwc45p3vR3pSaQDTg0o6mjsaBQAnXNKeV96MYNIc80AwOccUhzjmnL0P9KQ88mgCNsZyaXgjjvQcGhQAtaJiHKMilHANIB6UpGSMdKhjFWhjjJ70Hg/Wmt7UgI2wWz3pGxkjuaG+YHHamqpYBj1rRICInEwXbmrEXA+bioo1zKT1x0qZepz0q5fDYB+cKcVFHJuyOlP9j0qpLdbJDGq8+tRGN9guXB7UdeBUULSMfnXANSAkOR2olFpgRbG8xiThe1KrkttHapRjILcVHs/eZB+X0ovfRgSZ5xjOKUDGT0zTc9aUE9Mc9qkAbgECmIxlTA6inkEKfemrlMjGKaAe3JxnpSqAuTTG+7Tl+YA0ugDxxzQvQk0hyWxSnnjFSAdOc0dVz0FKR2o9BSATJye+KQH5+e4oJpD196pAOwM57UHp9aO+KOpFIBWAyMde9NYd6U80mc8UAKMg+gpGIx1pTzTCm85OeKAGsucH0pNoA6dakYGkXkc+lUmAEEKOefSlGQBQAfwpfek2AhXjGaZ9xj6VJ2OaaSB75oQEZOQDtzSIwbKgcCnEYJGOKZGFDEAc960TugHowPHel+Yv7GkJ3PgDigNt4UcDtUgSfebpSd+RilyduRxQRkYNK9wDcMU5RjNM2jtTzwMDrUgBGRUBHzcnPtVg9hUDDk046CI8/Nt6AU9csCTwDUbSAOoAyTQCQfmP0FaWGSYypOaQggdMmgliKY7eWmT1pJATfw4HNAcjC/nUYzw45GOlSoV2Fsc96GrALSg5P4UDpj86CMnANQAAAjPp2pCMilBxxRkhMjrQBEP3OBnipiMDd2qu8ZdwQ2B6Y61YPPaqdrXAYQMYoXlTk0jgKpBPWmL80RKEZ9KaV0A5228dz2puQoFD54z+NQF2Z2XHSnFXQFkYHPc05eV5qFG+cqGzjrUisBgHvScbAOyCpwaaxGMHvS4ULjNHBAHalYRDIP3i8cntU+QCAT17UxtvmKMdelPdvnHFD1GKw+aj+LpSk9800+p5pAHRaglYBR3IqwRk81WnUDnHFOO4nsO3kgYx9KdJkoR04qIF1kGMHNTOoK4NXKy1GR7wFAFEjBU5/wD1VDGcSuMcetJKzM+3rnrV8moD0lV/lHU1ISysAD9aZGhUe/apMkyYFS7X0AeeFGPxpsgJXG79KceG9qR+nHSo6iM+4AKsoGQKrQOGVM/e9KvSxjcwPeqCqscwbGCK7KTVmgCRczh2PK9qsQH5AwxUTMmRx3/M063y2SOAP1qmroS0J4iS+89KSWMglxzmptu5g3AHenBW83I+4axctdB2M65gMilgcMOtKsYjg2B/m9almQrNuB+Wq6l/MB4CmtoXktxeZJFkBQclu9SzMBgDAJ6VAjkOSOD2pzRFpt7NwOlOSu7h0JM7iAW5Aqu8arMGZv8A69KR5btLtJb+dDuzW8cjLhvSmr9BX7k8CCUsM4OBzVuPGzbjoKqQuYYc55PXirKZYblbArGpcpD5HJiOKqWHLEjvTriQiB8c/wBaltI1WEHGD6Uorlg2F9SyRkY9OlMlA3DPann5WzjBqORmDDPWsUu4mMt1IaYnpxilj3CNjnJHpSqSWbHU9aIh5aSO5OAOa6aSvUSJloiyy+ZFBdM4SSI8FulOvLSe/jCSlUQ/eAHWp7eWK4jOV4H3gw6U+S/t/KJEgYD0r14ps55NXM+6ti9u0MZ28cVU+wSTNI0xCblAGB3HetLz45w5gIfb1AquZpd3zQMB2OaNUyVrqZn/AC2GB0PSrUShXwOPepvseZDt/KnxWbeYMtj3xWkWjidOXYZFEZJNozWx0AApI7VIcleWPenkDjBocr7F0KXs467jf61DeSbICVOGNS47dqoagTmMduaFuVWk4wbRUTIPBqO6Tdhwc5p4568injDxlDwD0rR6nnrQz87W69aYDiTgcU5sq5FMP3h696zN07lleAueDTX64/Kmo5K8flS44xjnuc01oD1egdwRwKE6Z7fyozyOfwppOOMUDv1JAQc4OCe1LuGABz/Sog2O2Me9GSRwetMTWhKWyOvT260vIPJ+lMBHO48UBs8CmJd2PwcD5frzSA7mzmjg9RwKOnPSpGvwGYyRjgd6eQNuAeaQqCoxSKOpA5Hai+g7dLDckv0pVXJGOKe2Dj+9SJjrQkD7Md3z6dKUHBz0zS4yCR+VNIC8jOaSBPQdyOc59KCSpz3pBwMk49Pagkb+OCKrUmXkPxyp6Um0EAjpSY52gcd6U85osxP0GPkK3HFNjlaMnB6U9yTntnpUTZyT0NKxalZ3RpW1wkmMnB9KuKyuOCMVjIPk4708fMfYUcqLWIezNc+1NPSs1Z5EyA2P6VMl638S5HrS5WaKvB7l7Pyg008GoVuomHDdO1PL7hwaVjaMlLZgx61DI+O9DvwfWqzvk1LNLBuLyDJ4qxGhzk1HDHxk1aQYoG9BQvIqULjvTFGT71J+FMkik6e9Z05ycVoS8c1nzfepSCI1FGf5VMAAKZGM1KOaUblNlS6jxIHHeq7fLnA+taU8fmRlc8+tZrH5a1jscNWPLO5GrlicnpSseMY4+tKRnOeD/KnYwCM0EN6DVPzcnmkB4x6dqVR19utNGM4zk9hRoCTa1HdTx3qKSQ52jrQ8g27QMHtzUQXJ6c1Lk+g4xuKilm6/Wr8A3IFUVVjyoOenpWlaxFYQ3c0PY3pfEWIYwvXrWtOGSztjgDzC3b0rNTqPetK9bdaWCDJ27yfxpK3U2m3bTy/NGZcId2R3qWGfegRuGXv6ilmGDjrVNshsjrSK3JCNzE/lUTnAIpfNAOahklVuM80m7IZDcSBUznmuekl86Un8BWxev8uDWHxnOKxk7s2ghS2TnHApAQB9KUYJNLHwcVJpuXbOMnIArSVMoB3FVbFcFiRxir0SnjNcdWWox3HA7CpPl3DikQAZH40AZbmudu4xpAAOeDmobmNSmMde9Whz2zTHUOWB6CnGVmIpW2F+THSrTKcZ7VQG9Zzjpmr7HIIzzWk1rcSYRkhDmnJnO7saYpySBTic8YrN6lASCSDWVNCZZGGeB0FarDrj0/Os0RO0zbiRmtqWjE9yO1Cxhst+GKmRjgjvRs8vJOCDUCEhs4JxWr97UTLDkLCS3A9ajj5JPbNNnJkXjpSWZ8wlSaErRuBpsRgEHp1qSM/MB04qsjKG2g8jqKsIwD+9c0kUIy9SByKrxkB2yOnarQYDOapxcSMM9TTjqmIlRuOKeFAQY6mollBkKZ6VIJFJ5PA60NMEPbPyr3pMZY0m4M2ecdqUMM5zUgM3AKRnFRFuSpPXvSyHDE9qqeeWcbhgGtYxuF0XlOOKVmAA55PaolmAGO9PUcEnrUW7jJSQwVT260/dxj0qCN/lyTknvUqnCknqaloEG7bn1NRnG7p1p6Huw+lJwV57UIAVSMknrSMPmBpQ3FDcKD1oERk+nB9KBuOB0pwO5t+OKVeofsKq4xHUs+CeKcRgkUHlh6HvS/x7ccetSAoyRjNMyxBx19KkJyDUcZG7nvQhWFyTk5pq9ATTm4Yg96apGcelPoA843D0pnJY+9Oz8wFAU/jSGO25GfSoc/nUvzAHmqs0ojbr9eKcVcGyUsCQCeKiVwXx2odyvAHPeoo3zyw4rRR0FoEpQjB5KmrBIWMfSqbKAxx1zVhgXjJU4NU1oIg3AztipsZbNRxqTIS3Wkd9rkVW+iC6FW6YShAM5q4x+QiqJQLIJKuE5UEdDUTiNWRWklIccZBqWNztIXtTCm1elJCdnzfnQ7WAtAYUj1qvKFLYI4FWR82OailUc461EXqA12zFtSmHITGMmgEkgDp3FKSA/H41ewyrIrBQ2eaqqvm5weQamupN3Cn60kUaoo2nj3rojorsm+thPL2x9qaq847+tSFgRTWZYlz2NVFMTHxjlgOtTKMjrz61AOmR37ipkcDJPSpkAOikYPJ9PSkgR0bAOAe1OLAn8aejLk1Dvsh6E8fCnJ5pwwST2NQJJywqbcSPb0rBoYpHX0qs2SDjg1aJ5HeqV0DuAB4HaqhqwFiZ2jJJ5p1vv2MGHIqFixjzGOSeeas25YoS3UVctEw0uOI3AKvGKCOc4pVIKnmkQsQSazGhO54qJgcgipFUqTknB7U0qC/P51SEVXLLLz0NTMAyhs/hUchUy8jmnKcOCenetHsA/KkbTnjtSR72brlRStgEkimMCn3OAaSAdGzbmVhkVGgXzy5HIp25kYszcU2YK6bVb71NbiHxsnmEAfjUxAzt7VHHGFjwT8wqVmYKCBkiolvoMpyKRI0SnrzVm3bCbGOWA5qvdLK8mUOPwp9v5gDFuSeKtq8RdS2pwhIpwwRupgORnFOXlCawYxFOCeKaD8hxRnCkU1QdjAdadgIp38qMEj8aSM4BbHXpUksJfAB6dqdjPBq7qwWBCA5J6iorhQ0BGfvU89Cx61BcuEVSOfUU4q8hGYIzuwxJx1FQuAzBQfrVuV1Y7yMEds1WKMXyua7YvqTYYUIY7TT42Kv7U5owIzxzUBO0Z71S94Hoy/bRr5u/0NWpY1kkDRqR6mq1o25CRxgd+1XYMDlzwehqH5kuxVkYJ8uPzqvwVPpVi6Ch2P8ACagjUgYzg+9Fw6kEm7arDnFA+4GcZzU06eWQB3qAkjHv0qk7oLNOzIztyM96la2ZEXaQRjrmmlTtyR07CkLMeRkADiq1BCbdxGPyp5j8oZZs5HGKaqbQD37U0nLjvRcL6gQSTzS9hxxQF+ZsHg0p6YA/GgB2eBTVJyR07U/kLwMGowckHsaQW7DueRUa4yc1KTkk56VGQF7Z79KaKEzjOPz6Ugwy4LE96c2D74oVTz9afQTehFnEoOa3LU5wRx71itw+QOhrTgUuisDj3pTs1qNvZmqPc0q8CmrjAI/WlHNeaWxzZKn6VzUx2yMCSDXUgDac+lcvegC5fjHtXZRVnYhLuadmcQcenei74cHHWmWWBEOc8d6kuiSE+XtUP49B9CAc9/xpszjyJcf3T/KpSoAHoagl/wBRJgfwn+VXHco708Uo7UAjGcUteAxhnjFIBxSk0AigAPXNHXn9KCAKPegBB96l6UA80mB+VACZYcfrS57d6OhoPTrQAp4X39KU565pO1LztxSAMd6UDvR2x3pDwaAE7YpRjgZoyfxox3NACkc0h5GD0paQ8nFABzSDoRmlIPODSdBxQAjcDI5pD2460uBgjFIDx09qpAKvIpaEGePTrR0FJgNZyoGe9IT8mc0koDLt9aaseFKE1SSsIijdPPIByTUq7y+COKaI/LI4B96eefunntVtroAgAV8Y5pQuHyT9KFyCQTSOw5B5NStxjicsTu4FNAXcDjPvSMFEeGPXvSAlVUKN1NICfoc01ieop3J60hPOKi7AUDI5pduBTQfmx2p/RTSYERQheD1oBxjNPPTpUTnA5PPpVLUCTrg0hHP86iJPG3rSsCoBzz6UWAk5II7etOTiMUY2rSRtn6ilfQY8ZHNL/WkPv0pe+PapuIU9qQg/jS98YoYZODSAToM0Yy5b0pCc8UuecCmALzzRx260gYAex70K+c4HFFgBgMYzzSjA6daYxPX1pFXHIPFVbQB+B1JoznvxSFevoKUDj0qQAAY5pmD17VIQAcdBSHuBTQCDB46UoOEwPyppzj2py4x70NgDfdHvTSBxjk0HuM0bhux6d6AGMTxg5zSFSB6UpJJ68UwAGQdzVoQ8MCcDjHekABYnNKxAycU4crnHWlcY4AnHpQBnOKTHHajomKkBBncc9M1JwH4Oc0xepJOKcF5FDAcRlqgcDNTDHNRt9/px3oiBWOFfA4qMjJzVpgHPNQkHzxgcVrFgSqWxjGTSOmSNx49KFVs80jKTIDnj0pLcBwUbNqmnRBhGd/BzR1XI7Uqlig7UnsA4ABjik3HcePoaMkngcetL1PoBSsAoIPFIWyMCkBxjBp33RmpAjkViRg4FP5xihiSelJvGfrTvcBH54NRwR+WWqXgYNGec002lYBj/AHevJpCPmO0daftAAJpoPOfXvQmBEkW2QyYxn3pzLuOc08AYyelO2jB5qnK4DYsOODkjvQW5AFOjURoQtNJ2rnHNJ2bBkbsPOQEcip3G41FuzKMj8alPQ+npQ3oHQCTtNJztB/ShT8vIxS9SakGB4am4yCadnLf1pjMCCBT1YCAEEnP40OAEJx9KYpLHA6elTbRsA6mqasBnsWLjHGe1KjMZthGCKnZB5y0qgGfeOla86AfgBcjPFMVRjeKlzmMgfhUS8jB7d6yQbk3Qgd6jOTyacrbsN6UrgkA/oKXUCk5Jy7HpVbzYpHO3qOpq7L8q56is+QRq2/HzHoK6aVmTsPkjAUtTbMNECGHOaSdT5Ab35NPjxuVepxmtk3y2HrcsABARnOfWpIiGtgrHBqKRWaIFV5zyak37LcfLlqwewEd1JsiBKkkdqjjjDxAkc1NJ8wKtz61XUbQ+0nGOKqHwtdRWF2CNuSAD29KjfPnrz8w7UkRfjI+YetK8Rebf3HWtErPVhqPkZ4z8p/GozLMLYd5M8cU+VjHKCo+tRXCSNdK6NxgcelXHXcTJLNpDC5fv0zV6HlTlaoyxstyoHCDpzWkpDNtHXvWNaz1j1KVytM4iUgr1p1jJvXnin3SlYuAC3Y1HYtISd/Oe2OlSmnTYdS7t5JJ5qB97PgjAFWFYk/pUNypaPCdawT1AjwVwFHPYetXIkUq2RxjmqMbgAE9elaNryWP8NddFfvERLYqWLRTtPCZWYMoHIxxTVeW2t0T7ZAEYERny8k4qxFahb95FBQAgkZ4aolsCdTU206fucsI2XO3P869aLVznT1V/6/rUZphtwZVjm8yd+XOMfpVo5Mq46Z5FJDprQXUl1NKJZm4GF2gCmzziGaNNufMJ5z0xTdugnuWEABpynD1FaOZYQzDk9amIwc0yC6OVppp0fMYzSMKZOpGeMVnakMGM/WtI8nNZupdY+fWqjuZYj+GympBPHWnK4B44PpUfbmkB29DzV7Hn85HeHzG8wfeP3hVU5JzjrVydcg8YFVTxyBUPU2i+7GoSCSDUvGAfzqI/dJq5bRxyxEEcj3pXaLjDmdkVixQkEZ/pTm68DJp8sO1yp6DpSCN2JwMkUFWe1iIdSfWlOFWgqdwDde9IO9KwOXSw4EHOBkelSxD5WkbkL+tQ5HfkelSK+CRjKnqKsm6JXwY0YjBOacYRlQWwf7uKa8gfGBhRT1kCQZB+c9aSuW+W+pCAwVm7CkbcePWp4RlthXIPbNMmYBggHA6n1ovdkWtHcjGcZz06U5Rjr37VK8YzIB1GMCleLe4APXpQmhyi1/XyGjBX5etMwCCCf0qRUY7do4bOOfSkZTzkYpmbXUjDc/WkAxkZxTmwFX1o4POM4pNMadgyuzNL1I7+1NAJ46YoJYY/nT9BO3UccEnHfvUW0lxnvUgPHIxSZIxximib6Cjjp+dGRnpx3o25yc036Ggm1txxkIPvSEg8E5pAMrmgkbiMUaFLUU4UDngdKSK6aI7Typ7VHI3Un0qCRsnHYVLZpC6ldGkZN44ORSxx72yRwOpqnasGl2Ho3Q1qqoX5QKyt2PRjK6HKB09KkAxTF61KBmqSBjkHXmn4LHAp8EDzOFQZJqVkWBmVGDberCnsSyhP8oOazX+971embIOOapsPm6VLKiCttXpzU8ZyM/pUQXNSohFBQ/vjt2rNni8uYjOPQVqYxyao3yAoH9KpGFaN437FENgnPanZ5G7vTX6Ag800Nkgkc+lCObdCsdpqJpMdqJJMd6iUbvf0pN9Bwj1F6nn8amjjzinRQndk08gqOAKaXUbkkxFwCMdq1o+YlI9Kxlfb2rYiYGMDvTexVBavUmj6jNXriTe0ajkImKpwgbsdqmLYBNT0Ol3K91L5Q6ZJqgZHPUEn2q5PhxzVUgq3FSy1YjIlY7cbfxqtcKUYLnPerjuAR69qpSEMST1qJbFx1Kkp4JzWYx74zite4RFtnbPIrJGFXnisloaR7IAQVyO1SRruYHvUXT5e9XLNPm68UpOyuXY04FyqjHOKn4BHPIpIwVTd2pV2hhnqe1cDd3cZMuBIaBjBPX0oQZJPpQejAVl1AUHkHsab1LCnZwBQDkk0gM9o2SbPY1OEO8EnrSyqMg45p+AoUdsda1croBqLmYntQxwScYpSNkn1qOclQMdaS1Yx+SGGenrUbklSAaWLzPLzJyfWsy4vZIJwn3h/KrhBydkJuxNPNHkRNx71GrrvZAeV7+tUrm5+0y56cU1JGjcZ6V1KnaJN9dC6W+YDPJ9qkRSiAKtVFlDTKxGPxq6kq7OeD3pSTS0BEiSBMDOT3NKJ+pHHvUWwHmM4HpTC5Q8jis2kx7Fv7RuUgj5vrVZZirncOSaaJsnAGM1BLI3UkZFOMA0Jwz+YeafExfIZeazopz5gJbPbpVvLGYFelXKFg3L6sSuewp0fGRnntUIZlT696kVzyfXpiud7DGOqsTHzkdazZ5Sp2k5btWjekpDv6+orHO+XLHrW9FXVyJPoaFrKqEBvxNTvMSw2CqES4QnNTiTaenPelOPvFrUkhmEsjKeParakjCdfeqkMaNudeCfSrMQIQ7jkr0rOdgV7ErDaAPSkbjHp3FAUsRk04kM2D0rIfQRVwNtKw3KV6U0k/ePbtTv4M+tDAYhxgEcHoaWPgEEU05LgAZFPK5xTYDj8oAo/iJpdvPzdqaThiMdagQ4YIJqJj94jqKlHU56d6jJJY4HHrTQA/JGaReJDTscZPSkIGN2etMBwJyKXv14ppAKilAwpPTNIBd2X5HGarywGTk+vFWOiHvTRkke9NNrVDKrkZ57daijZy2c8Z6VPIgB2N3qGEYdkHNbJ6EhOv70kH6CpoM7B6+9RzBv3e3tUykDBPWk/hGNlX5N3cVXKZh3AdegNW3AYbD3qrNIoKoDjZ1pwbBjl/eRjnGOtTJztB596qxM3OMEZ6VeUbiGA/GiegIgnGMgVXhLZOTVwoOT14qPy9oBFKL0sMlTO3dTX6E1KE+Qc0x0yoGcVF9RWKY3BmcHLDqKkUtsLYzmkeMpIMDjvQ24YCfdrXcFoV5rZVI9+ajGQPlHNXpE+XGMGq2wxg7jnPerjK61E73I1XuTxUM6Ox4PFXUAIJ7mmsBkccd60UxabkcKEQ8mpVQY9M09F4PanHgcGoc7sERlTjDCnR8Nn1p7rujzTdnAI7dam47EgTDlvXrUo52mmAZAA6mpF67aybuUhGxvz7VXuVJTPepz1xjmmOu8HPGKIuzEVlG2PPQ1NG4xjPJqmqPKxw2CKuQY2gkYYdRWs0K5MOFPamZLR475qRcMDgUwncRiskMQk5AprDJHNB4fJ701iFyM8dqpICvLxJu6Yp4AwuRSOApPcmmq2F+70rToBO+GOT07VDvO5eMLSupZEwTjNSEAgJjg0tgI5cP8ALjkdKY3yhVC5YVM42jeME9qY3IV/4sU0xWJAoDFievWgMGO4E7ajkO1FBGSeuKep/cEKMexpWHYY+55CQcqadaOzFgV4PQmkjPmnAPQU+CPy49ueR1odrWFcnZhu29+9RtvIAH40mN0m48etPUjf9ajYYbRjJp6A4x60hGT7Uqd+OKl7AwTk5700jgmnEYUe9I3ypk8UICJjhBnrVK4zsJ7CrrrvUDHNVLgEoQp/Ktqe4O5m27FiwwfrUzOEY5PPpTIx5RbjnPc1DIT5oYDFddk2R2HgiQn5vwpVhj/iOTTAyjIwOfSnI4GRt6dKb8gTuTxAISVGRWpbxpIvzGswDdjYCCOvvV6FpIyOOCKnzZL7BfxBQGXIHSs0khav3V0ynYUyOxzVCbIHPSklqF2Ru5ON3amhvcc0rpuGf4aWOFXQgN8w6D1q9LCZHI+3g5z/ADqMFgoBPIpzJk9+O1OEO1NxPQVd0h3VhvzYyetCryc9+lA475o24z3oBDgNo6/WlwfWm9Op60o5wAakQrHHBNM74PU0uAPWm4OcnmmgJBjkd6jPAHr61ISAOnNRsScZGc5oRaTAtgcnkc0gyST29aQ7cDOKVMnA5pgBbnr17VqWJ/dnJwPSsuQDbweavWP+pOQSRSlsS+5rKAenapB2FRoeB6U8dc/rXms1ZKucj3rntQA+1Px1NdAPXtWHqoCz+mffrXRQepDH2WAgGc1Ncg4Q55qtZOT9B0qzccxBs9KqWkx20IGbJOSMg9qimJ8h/m/hOeKl4JGDxUcseIJG6YU8fhVx0Ytb3O+HFOAx3pD6fyoA6V8+aCj19aQCl7c0YoAD09aRvQUo4NNZuaBC+/agilx3pOe9Awz8/tSkd80g+lKv0oC4dqXt1o7UUgDqOlGCSKXGBSYyaAA8H6UdaXHzdc0ZG3pQAn8XA/Cgct0pcYxQDigBByaQDLHjFPHfH5U3AoAaOaZ061K3HFROeSuKqKuwFU8kA9aeBzTF64HSnBcYoYCPnGQOaYu5vvDn1qfHUCmniknoAzqx3HjtUZcg5XmnSgEj09aiOF3Ecn0q0hDs557+lObLAYwDUCKRIMGrLgDBHUVT0GRuMqM/nT4yS2MfKOh9aQkscDkVKFCrjpSb0sApwOaZnB5GKc/3c56VEu4qCxzzUpaAOA3kipTw2BxTIznkcU4knpSYAxyaqumZSRVsng+tVXG1jg96cHqHUd0zxz2NPXDHJ/Km9uT0pAG3e1U9QJh8wxmm7DggHnNLkD8KDyS3Y1CAeRkYpehpOQmaFxjJqQF/i6UuM0HvSHpSAFGKABzxmjtzQD8xFMLAAATn8KTseM+1L3NBoBEZBZTup0aADHSnEAcdzSZz0p3Cw7gMKDgH3ox8wz3o28+9IQnAI4prcg9s0/ryaQ8845NAxoIzt6GgcDAHFLgdSOaQdKYAcBT60zp160/IbgU0jkjrntTAQ8g5psajOSKkIJz6VG2cjHT2prsAv8R9KVV+Y0cYIPWl5OKQCBc9e1PLA4xwaQcf/WpB8vJHPakA1sZHFSjgioyCDknPpTgcnkVTAeMDg96a30pfc9ewoYg89/SoAi52jsabtHyk9RT2GcYpp4AzWiYCKCzdc0bTvyKTeFfAHXqakC4JPQUPuA2MNwCKQK3mcfdFPjyX4FMjQLIxDZOaE9WBLtVF29qa8iBwh6npSkgjPek8pfN3n7xFLTqIUqBzSk5IpT39aG6D0qUMQ8ZzTdo2jjk0pYHgmlB+bGMUwEYcD2pMcDNOPBxSDByTQgsIeetNcBjjOBTweCBx6VEcY4GSKaAfwcKOacwx8o6UxCSfp2qRQSpY9RQ9AYjfKvpUTEhSBUrZY89KY4I5HNEQGRAkEkc+lT/8s80wKQMHqaeemKbDoIhPJP50A4B44pqnt704H5cY4pNAIuc1GPmbJGPrUgOTUExABBJ9qaWoAj4Y9xU0KhWxVSNQkoAzk9qtx/fJA6+9VJW2BCBfmORwKQ4VCVFSkHOOgqNwBnvUBshM8cDgVGELqVJzkVKCG+UdKQjbIT1qk7AKqhUCjgClfKgCmqp28dacMg4I5pPe4FWdcKDnGO1Z8ilRuI3HPFak67shjxWfcEkFSuB65roosWxG4Mluc9venw7WKMThwOAKSNSBt7AU2LKuuOQO/pW/RoXUvQPuDqwzUiZO7ptpFEcZOSNxHWmAiFDznPSubqMjkiLNIw6mm2u5lcPHhvrVhFIi5GSe9IPklJPTuKpStdBYqgqZTsXDDrUwDNGTjJ9KNudwHHvTN7pGxUZ44xVaNaCGthQAKryIEcOWOfSljdiuZOGNQzuSrMG6e1axTTF01JpAzTl+o4rThHAYcEisJJGRNxbLfyrUtJTJEP7w7UVYO2nQaZauFLxL2qOLcspJ6VYyDj5fwqq5YghOD/KuWD6D2LIbMykHgUMAwYYzxSIpUAnrTiNueo460tAZTKsMHG49ga1rUEQ8jn0rPiD4Geea0bUHyyD+FdeHledjOfwle7WeOUy28ZcyDDHPTFKLu4hkTZamUyL0GAQR71NdQPLsML7XVutRu1zaQ7+JXDdhjI9K9KL1MFomNd9TkvFkMG23YYKEjg+uaZfxxtES6bsHgZxUl9qEzWRNvGy3GRhGFPlBFuDKvO35hWjJaGWW8IUkK5XsvYVZI61n6c5xjaFV+QoOT+JrQb0oFJ9S5DzHTmHGajtj8uPSpiODTIZX2kZzWfqIBaPnJ5q5NOEyT+FZcszSPnP4043uYV5LlsVzz7Ypyrgbs8jvTlwoI/M0oHYj9a1scOwxk59hVSRSv0q6QDn0/nTHGeo/+vSepSdjP5HFTQP5cueuaR0CvntTCDjgYNZ3RtF2aZoTqJId44x3p1uoMO71qC0YlCrfdPSrQXagApX0sdcVzS5ypMqvcqh9Kka3QjhcfjVaVv37MDwe9Ikzpxn8KLGanG7uic2zgnacj0xUJUp8p4q7C/mRBjVe4G6YKPxoi7OzCrBcvMiNVOOKfgZPHJ61oGJCmCvFU4IzNJjotUnfUzlT5Go9xgYqc9MdKaRjt+tTT25iGQ3y/SoQMjryKF5EyunysdG7BsEnGadHMVwNudvTn1qLb+Hv604eoHFOyEpSVizHMp2nGAvFPjbdLgH5fT1qpgDjr6U4M6k9iaLXBVWrXJTHu8pR781BKoDfIPlYcc1IZSVdT3xzTZP9WmOCuaNUxNpr+vT8hTF8oGc5600xMGX5eWzipDMDE2R8xxj3oMi4Ujqo4H1pK5bUWRrGzYwOvSkRSy4AzTlmZR8o49M9KI3VYmB6noKZFkKISWYHjGKaYmTbng84pwkb5iepxT1ZTs3t93NDBRT23/4JG8LKCM5IqB12E56ipppPMZQvckmoXkTLOG5Pt0qbmnJEglY9+BUWPx9ealcb244x29KhGXPHWluVy21J7ZOd/TFalvN5g2ty3r61RVFRAoGaVR8vWqSViVVcZGuoOalUZaqlvchwFY/N61qRQZiMjHCiptY61NSV0xEyoxnApkjcYHSnM4IIFQuabGVZTVcj5+anlqJVqCkPUcfzp+OKRR+dLnAAoAG46/lUEyh0KmnM2SaYTxmgLXMzBByRzULucEjk1Nc8SHBwGqocsRjmndHKqbvYOW5x0PrU8UeY8nrREnyjip8YUDtQo3Cc0tBVAB6UrgbcDrSoOc55pcDnmtLGLV1crsvPsa0F+RQKrxx73FWiuTxzUSZ00E9WyxbuN+D1qeU4U+lUAfm+lW1kWVQrH6Ur3N2hh5HWoWXI4qSRClQtJtWloUilcPh9uarscZapJSTKWPeq1wXaEhRye1YyZa2IbuRTEQDwao8YLVJITv2njHao2x2HSoRqtBAvzggf/WrTtVCr71QQBjjPX1rUtB6DpWVV6FIvqQI8DpUbEDk8Yp6kY2+lI+1cjPNca3HYmDcZz1oIAPNRjJwSKe2eT69Kh7gKo3cDt3pc/KaRc454zSqACR1FJgMf/V89aReUVT1pMqxAJpVwWHOfeq6BqNJzMM//AK6eQS49qidT5gz0qXOQcdcdab6AMncRqWPSsC6PnSFgML2NX7m4LEIecVnzDDH1rqoQ5dRNjGUBAV5NNHzAipVTcuAM+9SmABM1vcCGJSQcCp0DYOR9KfCm5eByalVB5gDDAHSldNkiIGBxU7jfDgdaXIGPX1p5TaOOnqKylYq5QeF1A/nVVlOTn8DW1sDdenaoJIEXnv3ojV11BmdDEDnAyR2rRiQjHFRW0ZLMw+7nGauxqFUsRzSqy1BDljV1IOcCmvEUiVgMj0qfOB05pWxtArm5mMz7l2MSjPXvVSKPrgd+taLxhxhjwelQxx7ZXQdBiuiElawDEg9OnerCxjByvSk8ghwQeBUhB6HvUSlcLDUXacgYJpFkKybccNThncATzTmwQCBUXGSKME/pSDAPH4049Bim/wARI71IMQggewp5+77U04JIPWlx8oyfwpCAZEgPakzmQjrQTwu3vSYG5vWmMeKTOBzQDuUetOPXHXNIBRznnA9KbkAdOO1BPPHTvSnk4FIXoIRng9D2ph5XAFSd8npTFHzEdBTQxP4QSOlKpDIT6GgjLdeKagwDgUxDs5XPpT+gA9KYB8vHWnAjk0hkcq5G8jkVAiBXZhVpiCpqIBiSB261cZaWFoN4/GkjA+9mhzggAcU5cYAHan0GLJkHd+VVJFBLMwyTVtgGbmqUoOScYXNVTEx6AIM96tI3yjB+tU8kfw5/GpbcNvy3enNXVwLJ68dDSbfmHen4wxFNPc1jcbFBy2O3rSkZwuKUDaT6+tITn60CEIDH2qBV+YjqKnXAUnFQ4ALDHXpVRGJtyxz0NROmcqw4NSBSIxzyKUqCMVadmIgMZRQOw7VHIhWTIqyctwexpGVi59KpTAYGyFUDr61KsYV+e9RYGalEgcg4pN6Ah5HJXHB7VHtCZNSlsnio8jcBUIeg7kg4oXhs96XBB68UKPmLDpR0EBwGOTTH44PellbAPqKazAxg00hkEKhGc1LCd6EAc1Citl88A1LDtWM45PStJCLC8AjrUY6+9PHTINMPQ5FZIBrqCwJphw2T6U8keWDmkCjH1q0BXlTABzRGQ+aSTLKCeTQibQcHpWnQW5Lg7Tjih/vAY4xzSIpIbNIyswO049KnqO+gm0zDI4A7UAhkIPAWnDhD2b0FNjYMhBNMNxHGQsmelLK4ygA69ajOJM4JwtPQ+YQ56jpTsJjyfLxsXrUrZPI4ppkxLg9acq/McnioYyu6kuFBwD3qwsewDnmkyGlAxwRxUvGAO4pSbEgOdvvS/dT0Jo/h570FsBTioGA7A0xxuYL2FPOBlqickIzD71NANMn7wgjj1qvIOGz1NPUENj24pAuCWYcmtVoIzHYqd2z8arysQMYqzKGyRk4zVduijqBXZEXkxq5Bxj/61WIQADkDj1pkahR82cnpVhQARsPBokyW7EygMAq/eHWpYZd0TKRynf1pRLGyrhfnHBxVafcrFj0PYUr3Fa+o2bM8o2N8nvT5LfICqdw9h0qusq7gUGFParcMriLcgzg0NMLFUR7XVSKkaEQqzE4PrUsxwUI6k5PtUNzIG+6MZ96LhddRgdVff1PoRTeZnDEHk9KkhjVgWIyR096uxZZAPL4HSmK5jSRlWwMkeppADurUvXG0YA9aolWYb8YHcVSkO7YzaQCD2po5H1pwyvuDQABgY4oF6DduDmg5zjGfWnlcHp/9eo1Oc49aaCzFK8e5PrTSORzTwCGHX2NMbpwRk8UIaEOeTyfalXLcY+lIARznP1pVOOPWmV0AqNpz271c01gUYHoKpupzgnGKsaaPnb3NGjWotkbXA47VKMcVEh+WpVrzZbmjHqM+1Y+sLiZCB1zzWyDznriszVhuEbDvWlC6mRLYo2LDeeMVoTD9wwrOsztlOeD6Vp7d8b9icVtU0lca1KqAZyRj2pJx/o8vH8J/lT1UgnvTbji2k5/gYfpST95FHdDpnvS9Of0pgzml69K8IY7vjFJ3zmmsxzS5yOtFgDv1pSOBQMemaQZ5pAHGOtN4zyaRzngflSMMD3q0gJCPShetMB4qQcEVLQC/WkGM+tL1JzzRSAB8wzS54pM8UA0AA+8c0uBj2puRvp2cj2NABj8qTHIpTn1pc5wKQCUo4NHbmk/lQBFM20ZqNfmJPcVO6bhVZyY8Y6VrDsIepYk5qVCSuT1qHBz1qcLwDRNJDsLnqaa33SaV+CKUc/SswIo/ni+Yc1DINsoIPFWcjJB6UgVXXitE7agRiPDbuuelOJ4zTj8vGeRUbpkZB6UbgQwOfNbPFXeCO2KoRI4m3E8ehq8D0zTqJAthOg9qYpO3kYNEhxgAfWhJEkJC/wAPeps7AOBII9KfnHFRAEvkmpM5bFJgAzmq9wAV5OCe9WCcfSoJgpXcRyKcNxCqoCgGlOSTz0FCklQf0p4HBB70NjeofKOtLkEUgIz60Z6A9aQD/urjFKPWm5xk07HyDBqbAxTggikzwD2pwH403tSAGwRSdOopTxSdulMBRxuOcUHtgUirzyOaDn8aatfUOg4j5h7UDjNGMH3NHC8YpAIc5o55pV60EZpCGgnpTsfLSE4OB0PegZC4pjE5zijocUucHj8aQ88/lQgDhenek285pxGMHFNIO4DtQADGN1RvkJn1qToDUZI3HI5qogKvC8jmnAM2TjApFyx5pRgnApsBMDNIpO456CndWzxxSDAzkdaSEI4LA0YIQ7e9Dsdo96aeinmmkMlUcDJobtTc7kwODT+rcmlYBrZAqLkEccVKOcseBSHJahAVyMTAEcetWNpZjzxULH98B+tTDJbmqYDclCPemSYWQA5571Ix3SDFMchnw3WhBuObGQPWpOOMnmomBJGD8opsiN5gbdj2osmBOAOWzQcEDHakGdvTFKDlTUMBq4BwO1Ln0FIFGemKcDyeKGA1idpoUfjTJjsj5NOiyVJaqtpcBxwF9KjKnPTAp5wME9aYz/MVPU0RVwFB2g45qQfdx2qFyU2qD161Mp+UCiSsAFcJxUT7tlTNwoqBsrF1waURDhnGTSqcse1Ql0Q7STx696sDG3OKpqwwUbcnuelAJ20o4bH5Uh6ACpAQcKT61XmQbdx5NTtx1OfaoXw2BjjFVDcGV4943N1OOnpVmAt9nJfg57VX2MJ0xxjr71aVQQADmtZtWAcAWkzu4HJpSB96nKNvy0jZCsccVjcBkZyWbAFBJJyO9CngEcg0nmA5I+lV1AWFSifMeTTnPIwetNj3bMNz7045B5/Kk9wIJueT27Vmzxlpevy9xWk4whzVV4kHzZ5NbUnYRGihV6YAGOtQDeH3Z+UGpGILMT1JFRurCbGePaumKuxF1o/M2sfwqQkbgG6LSJIC6xAZIHWnPtYMx7VzSvfUoQEqDx16U3I3ZJ+tKzq8GR0HSo9xTDMM07CJW+aMkN3qvGR0B+X1p8nzRsyg+1Rb1SNU67s8gVUVoIQorBm79M1XBGwxnlRVl8W8O3aW7moIcu7nGB1reFnr0HsUztxjv1FXbF28zkfLVICTJ3cmtCzUqpPetJ25SY6mvuCqDmqzfLIxNOikyoJHA6God7NckHoa4VD3mUy4nOPTvQxG045IoT7jHr6UuMRk4rIHuQLuK4z83fitG0z5XPXNZuf3hA6dzV+xkD+YAMAHg124XWoZz+EtY702RdynB59cU+kJ7Yr1LaHMLDMj23mZDYHJxzmqiTw3sHmRHcmSPSmGECV1W4MJ+8w7EUtilulpstn3KCcmqVrA11I7eO3tpfLjAV27ZzVw9OvFZ+DHqAOwImeW7tWmVouxPYltzjimX9yLeDj7zdBSwna/PSsm8mN5cnB+Rfu1cdTCtPlWm7I3leQ4JpEyy46H1oAwQB27UpUYAH51otjgcne4mD0796Ugryef0oCncRnpUkcTSPs6g07iSbGqglzgZbsKlWxkIyeDV6OIQ4UDk9TUhx3qHJ9DrjQS1luZM2nvyQRVCSEocMuB610eQT2zVS/jUW7Njkdqi7Rr7OK2MpblCnoR2qeK4EkBbPTtWdINm5TzihZP3fTpwfep06BGpLVS7DlyG68fTrSL0wOM091JPTGf1p0ShpQPU8Ve5gtHyl1FKoBjtVeIiS6J7E/0q23Ck9gKgslyxJHSkup01NXFFqU7YmaorFdoYkc067OIcZ60tmu2DPrT6Cd3VXoR3rYK45qmDkf7PpV2aIzXPJ+VR6UTWqbCVGCOlO6SMZ05Sbl0GWsKvuyOmKLqFIyu04zmpbJf3RYjBaor07pQuORQr3HKKVK9ivknFIcngDjsKOn1q5aQAr5jD6Cq2VzCEXN8qK628jnIU4+tI1rKuCOKtz3G1yiDJHX2qKGSSVnyemO1RdnQqdO/L1KzKVABXHtTkidxlVyO5q3dOPLAJ5PtUcd0Y0CFenvTTbRnKMYztJkRt5D1GfSkMZRckVeiuhI6oFIJp9y2IGJ6CjVFOnGSckzMJyfc9qjkcbc96uWyht7MOmMVTvFVWDDv2octSFTfLzEBcjBB5pjZyAep96cI2kJI7d6jIOeeCKm99C1F7sMgLgDn1zViBADu7moY0LH271awcDHSmlcUp2QcjkHgVIBtUe1AG3g8jvQozx3PSrME7D0UgnJ/Sr0V4QNkx47GqvI4PT+dMPAG4fjT9RqTTvE2XdPLAHUnrUR71mpO0UnqvpV5JVmXINZuJ206in6jJAT9ahxgcGrJ4GMVFgA5qLG4ox61GzfmKR3xwKYT/wDWoGLnNRs3YUMeMDmm4ORkUhlO6UsQfSo0jUrk961BAGBLdD2qtLbtGT3HrVxsznrJp3RCg29OPWngcZPFHfA69qXBz0571djk3EC4b09KMDOOnrTthYkAZqzFbZwXpNmkIczsNt4SgLE8npxUjcnFPZwMKtMHH0qNzsjFRVkJt7ZoAIOe1BOKaz45pdSx8k+yJtx+lZ7yFuKdNJk8dKhI9KiTZSiO3bRzTHZcfd59amWIbct2qnduERsdah6IvqZszfvSR0zULHn2PWpCBxUT4Jz1FSh3JYjkHHFatm+5TismPcpAC8Dqc1pWUgMgTpms6q0NEzU4UZ7mmFcbifwp5BySaaccZHPvXAhj4yzKPapCQUHtUUTHLDNSAAkgDrUvcfQUjsT0pBnHqaPcjp1FGSpOCKQtiCfHG3hj3pYiQQD270St8uQKjQsy528545rVfDYB0jhWBPFSFuBg4Hemyqro2TTIgPLBGaNGgKs4VrkqDzUBtwXzke9WZ0y4ZR06kVCi7yw610w0QrtCrGijAOaeEUoAR+NC4RsBf1qUck5XihvsBAmA2BwMYpQpDEHkEdaFVlfHUfzqTD9x+FLmC1xwClAf5U6PLHb7VX3MqcHpUscm3nGM/rSeiYE6MnAJyaiuV2KzjpUayFZNp6nvUswZiqdVPWo+0gFgixE2O4qVAFUg0qjHB/Ol4IPFZuTuyrAoBxntSqN2c9qaCRkHrT1HyHvmoAidexqIDZkdzU7YOO/FRgbkJ7+tWmAZOd275cUm4OgP60xCfKKk/nQMbCvY07CH8eaOM5/SiNQGxjIpsJKnkdKmGAvHFJu2gxCcg460nsDTsYGc00/KcCpQDsgketGCCR1FMY7W4pzHnNOwDQ21M+lEbblLMMe1B+/7elPIAWncBQAFB9KM4Hr70DpjNA4OD2qQFz8pP50hwB160qkscnpTTzyfXigQ7jG0Cox3FKpJYN2o4yfagBVO0Hjmm5Krg9aA2Tg9qCoLA84NMBRwp4pV4waQffxil9BnpSHYU8VH0z2Jp4wWznIppXOaEJEXG0+9NXIc/rTnxjAHIpoyWyOtadBki8VHIoYfdyKe2S4PSmyHgDNC3ERIVTGOhqQsFnABx7VCoAY7jyPbpUjYUB8c9KtrULlj+LJPFKeRQThaM4AwKxGKfmP1oJ7DtQq8Fs9KXPynHSkIb0UgCoJAN2e5qde+elROAM+oq47jtcjIYuR2pzEDB7UIWzg9KUDazAiqEMcHAIOKGIA3dj+lLjIOelR9AwI49KaBCEhVBB475qRcIeBxVV3OMYwKsxZ8v1AqpKyGSAYBPajaOvWgnKfSkIyAAazFuOPLDn8KSQgEbelC/eBJqN8Bl5xmmkO4SjJwT1pgG0FSeKlIwckU1tuCaafQREF3hkB5p8ShSVpAmJS2etPjHz7jwapsSJeTntTQOaVeeo59KQlRz3rMZG3oeaVxnHtSngjjrRx1FUBWkXex2dqjUl2APSpyMkt2qCM7nII4Nap6ATjG8DPSmyMMjqMU7AUnPWmyH5VAGR3PpUrcOghGJi3Y05FAJZRw1MkO5PlHIpfLZowOlPoAwAJIVA47+9LGVwUB6dTTpSy27EDLDvUcYAh68mq3Qth7/NKDnjsamUMinJ5J4qugKIcjI9KsAlwrGpl2AeeCD2PeiP5huHNEnEYz3pwwFGOM1n0Hsx2MkAmlP3fpTD90E9acGzFUgIUJQDtSEbFIIznmnEkoOOD0pr/eApgN2jfux2quykvnPFWZAcZU1G2VXdiqixGTckxO4A+oqvChlfmrF7mSdeSBSADYQnau6LfKKyuRoqtKRuOKtBeVAxUcartOeBT0+9hfuHpSbuIliQFvcHip5YM5BGD61JEsckeQeR1pSw5J79KW6F1Mi4hMJyvI/kaliYkLk1JPsywPOemKgTK4Zs47e9X0F1sWbu1Xyg6nioUgLuRyQKnuZX8hAcYan2u5YmJYZ96lN2bE0UVmEU2OoFWDc7nGDkelVyivKXc8A9BUkTxkENxz0NPQbeg6W4SSQgNt444phjBJBf8ACpPIU4ZTkd6rTAZyM0JIEwbZjaOlRgNu4/AUsTEHJ+7SqjSAnPFMSuxg6kjHvTAPQYqUKUbk0w5OT2qkBHnDdTjHWgHOSTzUhTJORnNRgBFwvWndMvoCjkA9+lKANwwcmmvwfmOKUdc9jQKwn8ZORjvmrFgSLhgelQnaxwvI65p1l8t2OSKN0M3IyNuBU6/dHvVZWwOeKe02xQcde9cE4u5aehYRgSB3qhrCkWwxU0cm6UU3UxutGGelVTVpokx7RwZgM4FbEeMkDpisO34nUY6VuRKePet6ysxogY5dgKhmz9ml/wBw/wAqeSN5U0y45t5PZD/Kpjuh6M7njPtS9c0Ac9KUDrXhAHekb1xQOTQ3PWgYnIxS0uc8ZoAA+lADCOen40bQScmpDzSEYpqQWGAYGTSgZz6UtGR2pAKDtoOAucc0YyaUnNACAnaO1CjJz3oPU0KQPwoAQAgmnn7vSmswFG4kcUgH9qQYpRjH86Qc0gA+uaX26Ck6GlBwOKAAjjNVpQM4PFWexqCZSw9vSrhuJkfRfapk+UYqNP3iAY6VYQfJirnYZBcP5abvSlgm8xMjvTpVDrjHFMhh2Slh0xwKSceWzF1Jiufek2gLkU7iggY4696zuMgB2sc05BletRyRkyK2cVMi9RWjtYEQ+W3zZOc9BSQMWBDdQalkHB9qhgZjuJPA74prVMB04dlITr6U22QKhUn5u9SnldwNKjcZPJo5rRsHW40qWI2nGOtTDGelIpGeOlJyB0/CobuAuD0HJqBtxl28Y7806WULlyaigK4MnO5u1aQXVh1JsDGKcuAuOc0LktSsMDIPNRfUBuOnenDAByKbznBpcgNkikA8LuHJpdvAwaarZyM0vf3FSA7HbPFKOpJOBTSecUpGaQCEcnnikXuaGztpVOF4pgIBwWpeOtHVaU9sfnTfYABoyNxNID0/nRu5PakAoJzz17UZ55pAOc96U84pAJzk+lBGeAaUc0dD7mmAxgSQKUYfhTkCgn5sUqrgelPoAdx60w8mn4wM0m0bTQgEP3aiGS7E9Km/hGKrureZx0qo6ASF8HAFG4A8fjTGKgFh1FOVgRn0ptASAEnApvJA4pcnJ96B9045NQAjkY29RikC7VGDxSucMPWozn1qkAu7a3HOak6Lx3qJVLYYjBFSI6sTtOcU35ALznHbvQxG2gc59PWkIzwKgCntP2hmJ+Q1OrgsAD16VHIjOrqeF7UW5VY85yRW+koiRMx5GBz3qN2w/wB3PoaexPA7etIZF3suORUIYihjIp/OpXXkEnkVHbR4VmB6nrUr8sPalJ6iEMnz7ACfU05TjgUBhk+1IuQOO3ep0GLtGQM0q8EnPPpTVOSSaUY5xQAyZPMG3+GlQghVHAFOONhz61GFwCRwapbWAkPzCsy4klE+0dK0lPGc1VaAOTvPJqqTUXdiaJnkCquepqQc5AqHyQSueVWpV2qpCjAok09Rkncmo3XOOKeSQMGkI+SsguVGHz7n/CrZPKr2qrtJmBI+UCp33B1x0rV62QEo4cnvSdN2TSN270OQuM9KzQDTgcnmo9jBsBs59qVpACvoe9OHHGc+9Uk0BXZwGYk9akh/1pPYCo7kCOAkdu9Ot5d6qAfmIzWiV43QLsWfcd+lKxBOwUxCQMdTTzxmsdgIiS0m0Y46CmlPLXap5qDesc5JPzHpVjPO4jtWtrWYeY4Dr6UrcjJ4H86QfM2e3elZsnFZgRSDenBxiqkgG3PX0q3Iu6JsVSJYRFh83t0ramIZuAlDkc46UxkxKZc9qkVw4JCkkH8qhmUs23HWuhOzAuoVKxtnkd6acmFxnOe9RwsQm0jIx1qcoBCVY4FZyVmBBEQijafl7mm+dgNnoMY96swIkkLRiq8kCiIx9QOtCcb6i9CQSl1AI4NIhR5CMcDtSIymFSKZCTtYkZY8VXKnewdR0TO8xUj5e2aZBkTyxleCKdCXVwM8YpkDLHJKWbcw61asLaxVZdjsCfmq/ZqdpB+bPeqcYSWV3B4PUGp7EsbjOfkI4NaSfuPuCLhVmQqrYA70lumZeeoFQ3GRHwOnf0qa22xxhs4J65rnivduV1LQA3AdadnnB59qavynqOelJE2Xk46Vz8ru0DIZTtchepq5p3AYHrjrVSQBZCw6mrVkwExAHUV04aXvoiS0L5oAo9aSvXOUhltIppFaRc46c1ZWzjhi3xJtB4OKZ9a0LMG5gkgH8QyPqOad9AbdtDn7iB2ug0KgSYx5h5C/hV6MMqKrNuPc+tK6c59ab9KEIr3sxUBFJBPQiqSpgYHPtUszFp2H92mpIw78H2reC0PNrTvOwKv938qFjBXI4qdJo1gdNmXPR/T8Kj3bc4H05qjLTQRVz347VetIwiF+7VQMmeQM+9acSkRIo6YqZ7G+HV5egrEVnXNwxcoDhR+tXpnEUTMe1YvAHPWlHUvEzcfdRMCWYLFwT1rQ1FPLsF3dZThT9OtVrCINJvHQVb1ZxJshA+SFMZ9SeTSmyqEWo37/AHHNSgeac9KiQAdf/wBdSSAeeSPX+lRoQWUdvSoSCW5YyOAe3epbchZVJ6c80w4zt7+tIxHGKsy5rO5enYeS+DzxilsxmNj3NUMEAEHntUkcjKpEZIosbKreXM0WbxvmQdBzVmFcQKvSs4yNI2XPNTi+GzlKTT2CFWPM5NlmWVYlLH8Kjgm85WyMEdqou7SPljzToH8qQMTx3p8ugKs3PyNKJdqgdqpTENcNjrWh0BxWbncd5Gc+9KJVdqyQzGBjseta6oFQDtissAYAxk1pwSeZEPWnLUyoaXMyflmbHJqzZpthLetOmtCXJDDae1PjTbGF9KJNWLpU5KbbILpQxQ56ZzUotIzGuev1qFyzOwx0PFQ+c6cg9PamlpoZyqRU3dF5IFifIPTtUd637gjuTS20jurFz06VWv5OUUe9Rrc3bXs7x2YQtiLjoaq3Zyg56d6kdxFFgnp3prqXG3P1qW7u5oo2jyjUPlwjNMgjUkufwpz/AL2QRjoOtIwJOMfL3prXYU2kr9jSt4kdc9RUU1v5DAjkH9KLPK3AxwD1H4VbuwPszc+lVaxlNqpBy6ozwPlB6YpUO1z2pVB3qexq7bxKqBsferRuxzU4OTsim2Oc9O1CqHkPHA96kulUT7V4A61EowTjg+tHQm1pWGNu3Zz+NAkdDlTg0Mwz1wBUTkZ/nSGnpuX4roSjDcN6U4nHIrLzzgdKsW9wGGxjz6+tZtdjspVW9Jbk7nvTAcgetEmelCjgd6g6hcVKqdKQA+lSoOBVAxcdsU3OafRjA56GkSM8pG5K80eQn939alxtPFID8vvTE4xGbVX7tRtJtB/SnuwUVSeTcxH5UnoUl2FD/Mf1qZW46cVTLEHPWp42yMUIqxMiiViTxVadwCasg4U4OBWfM+5jg1LGhnU9PpU0cODk0kEZbBI6nirQFEUU3oQy4VTisS9clggGc9q17pvl96wpmzKx9azqPUcFcjIycHpSbdx9qXIHH6U6L5nrPZGpcghGASRirqwJgNwGHeq9tz8hGRVwcn5ce9cs5O+41qSkZwSeDTW/1gwe1PUhgKawBbPoKwW4xE68D61ODtb2FQoBwATzUp6EY5pSELxu9qR3GPelOdvTkVGW+XIH4UtxkVxkoNv41ErBCFzx3FSSEmLJ7VXRhJlwelbwV0STzk7Co6UiIBF1x603+EZ546URHdCWPal0GP2AKSepquqGPOD1qVSAMnp2FKm8tkr8h7+lUm0LcRow5CnjHekMgLFV5PrUsqcqfSmqqghvyFLnurDsLGmASOlOcYVmFOI2rwOtAPykHqazvqG5AI/kzjqacy4fA7VIvzIOKbKPkJHU1fNcfQbHGnnF26GlkXEnmE8CmQNuHzDp0qWfGzaelJ72ErDl+dN2etPBwgxnFNQfuxzxSg/KQO361EtxiNn86evbntUZ+aMAnpT1549KT2ARiOaiPCZ/KpDyGNMb7o96aAiBO/2FIvBJJ5pxIX7x69Khd/k46D1raKuBLC21csc+gqdcEH0qhFNuzn8jV5GBXIHFRNWFYUqSfY03aN1SKdymoxkx89agOo0NhsN0p/GBkcmopCcA5p4fDqpHUdapoAfLfj3p3QJmmvwMYyKUkgA0hjgCCT6UpxtHPWm8gsppxHyj2pALnAxTZCAoJ6UZBY880rDI60dRMQ52njjtTTjZ9aeCNvB4pvGwkUIAHBwPxoAw3B6imqM9TyRS4AB702MVSPMApRjdjFICN44pckHJpCFPDccU3OOTSqcOM00ksPYUwGdGJpkY+cn1qVgC5zUaD94ewFVfQYYKtyaYenNPKktk1G+V78VSEyMgAlv0q0hDwjsfpUWAVAJ4qXlUJFEndAPx8opSfu00Hgd80o/2qzGhcnGBS85IpiscCnE9TQA0scYFRNnzDjoKlz8uRUZJPOelUgQZIYY6UoGST0pqHOc9O1OPPSmxDCPkPb1qEyBiRirDcrzVduAcDBx+dVEBpTcD7dKlj+SMknrTYxhQjHrT9oaLaelVIBwA25B4p+QAD60wDbxngClJG0Gs2MCuDyPyqByGTLcFasnJwT0qvOhKnbxVRYnsNSTeAcmnHP8AFxUFucKeOtTsPk5/CrasxJ3EH+tJJ4I4qRF2xkk5PamJkKMmlVSeCeKlj2JVBYZJoxznjFICcEUrDBH61HUBjgFs+lNBA5zStwTzxTcAKKpbB1IRkOwB4qOM/MRjpUr9+eT2qJcg4zz6VqgZKCduHFLu2xnbzSZEigj8qGZkUBRmkADcqk46+lRvvddoH196kZiqHGc+tNRAJRuPNNdxDiG8jCjntTQoCqo645pyLIJcnhRTWOyUnP0o8gbFQHewPHtVgMowvc1AMhyT+FPjYBwW61MlcZKQDwccUqjB5ph5Y8YpyuGxxUPYBX+/jHFOYDAA6UOeeKjZmJIA4pIB7HkelDAdaTqcHt1pC+4EdqLBohMZQjOc1BNIiJljgCpmYbDjrWZcI0znn5e4rWnG71BlaaTzpDtOR2qSOIjb+tWbfThFGzkkgdqgkvD0QDiunmv7sSX5gBznPBpfMCDIGfan20fm5GOfWrf2LHoaiU0nZit2Iba5UFgq/Iex7U6NkLYccdjU0dmAuM/NTmtMrjNSqsEPlbKe4I5IG9Cfyqwyh1BPAx0pwstmcHr1qKaNh7qO1WqsXoLlYyWMeVtB5+lOtImMWCcr704Sq4ZeOnSoLOfbIUbj0q1pci2xG0ardsqnAqCXYGIUjFX5m35YJjI61mPbvjcWPPaiLT3HYsR3qxR4UZ9KVruORMBMH61T8soeRimlstgfjVpLoGvQlGMkZ7U4sQuP4T2qL7o6YpeT15zQK9ieSNcAryPSoXX5SAelOWJmIB6e9K0Xz4J60k7dSlZjEY9SKacFiQOKmEOABjFRS4V8r0BpppvQFfYYw5BPPtTT6MeT6UmckcmlwSeDz29qsfqO4xxTo/lcPnnuPSlROBkZ/GkJGGA61N+wrIvxzLKx54AFSqxZAP0rPtYyWLMfpWlBGxUbutYVLJlLQkgUlwR1p98oe2epY1CLmm3JzGwx1FYqScrjOdhAWTORnNbcTYKmsEArKp963I87V9a6awokEobeeRntmmTLi2kPYof5U6RT55ps64tpOf4D/Kpjuikup3fuTSZIpRx160nvivBAB1prZxnNOHJxSMuafUYvUAY60oJAxSAY6Gl70gYox+NIw44NL0PSjryKQaCdOppAKXqc0dvegAXk0ZO7mgDp6UvXjtTAQ/e61CyMHJ3cGpjjmjGTimnYNyLa2do6etSL+tL1OO9C8Gk2A4DHenKcZpoAwaPWpAUEE80h4zQxAxSfWmA7PGMUyQEpx2p+c8Y6U1ydpoW4ES8LwKlT7vXmowcjkc0sfcd6p7ASY4+tJjnO6l9KVQMc1IDSdvvThg/jScHOaQZDZ/SgAZQW47U5QB9aDwDSBuaNw6CPjHJ+tVIcb2IPFWpACp7U2OPatXF2TAZJGJEwDgUqIUG0dKey4UmqqXZ8wI67RVRu1oHUtIMHJobIB5zShgATnioZn+QAHGe9Qk7gyo7iSQxgc1bCKVGG6VAY2RiwAPqasKgUBvXrWs2ktAQ5SBnmnZ4AqOPB5zkVJnOQRWbQCLyc00FTkDrTsYGRTVjCEsOpoAVRg570/IL4x1pF6n0pSfz9anqFx4PPSkLfKBSZxx3NP4BGaQDWwOM0N2ApT6mgYyaAEP3eO1Ieoz2ozgmgfdyadgFyevrRgZxjmlzkbRQx2/WkAfQc0DpnP1oz78mjoMkUgAUg4Io7UvTJxQAg5bNJnGSaXjBpq8gk1XqId/BxSE8Ypw4HIphOBmkMXOABUcmRz69qf1pGXrg801uDICpEfI5pIiGTFOxhm3dKZCOfatVsBYPC8DmlGQOeAaGYYAxQfvAZ6VmASEKATURYA5HU1IwGcNz61DID5mVHFVEBVcsOakRFRMAcmm4AGcZPYVKCSQaJO+wCKOQT27UDnJozyc80ueMDpU7gRyY2kdqiUZiO0VPJnaBjimKmItoHFUnZCGbigJJ6UiBGUSfxHrilIzG2RmqscjIxH8K1pFXvYC4GKkjt7VLnAJ65rMiZ3JbdyT6U8XmbkxDqKHRbdkFy+WAHpSLkD2qtLlsKzc+1SRviIc5Oan2dkO+o5pAmA3ANKjlmYYwvrVeYI2QWxT4yQcA5HQU+Rct+oLct5wtRgAClyTwRwKZvJznpWaQC4pj5EnTOaVmZlytIUYjr0NNaAOZxtK5+tEbZUe1MZM/lT4xsGMU3awEvuTSPkjHtStxwOlB5PWswKqsypyMk9qmjYn5j2707b7daFXC4q3JMBe3WmyjMWM809uVCgc0hxsHcVIFYxkqnPSlJOdxPFSlRzUQU4YA1opBqQXDGQbFHIFFqGWTawxgU9gVdHUfU0LnzyScj1rWL92wluXIzz704jCk96jiKk5HSnk7s8cVzNWYzHJMeoHcMg9q0c7hjPFUr2I7g6/iasxnMXycmumbUoISJEYEHC0/cPTmo0P73HpTz97kVi10GD/cPf2qjMr+XsUAfWr+3CZPWqjLlDuNVTdmJlVAwR1J6EYpkrnecjIP6UkcmCwfr296UvsIfqT0BrqSd7sVyS1Pm25xyfSre1UXY3PHFU7XeIWZRtJ7VbyWiDMvI7VnVTTGthYiIowxXDk4pZFyxKnHFNk+eNW29KkaPLBj0rNvqxlVSqxgnkj2qEoXdQuc85qwQUZl/h/lTJZmj2Kida0i3fQXTUV2Krkc+lNjPmRSMyhSR0p6pwMvhf7uKbNJhRtO3J61UWkxMqLG7gFhtUdB60Wk+ZX9OgFSXskiSQ905yBUcQSKVdgz1z7VstVqLqWpAXQYXGB0qG1H7za5yaugB1JA4IqrDCRdf3vesqclqimaKKdxJ79KlJDErjmmJheAacOGzmuV6sZDOQrgZ6UtuSrBwep6Ukw5Dbc4PWoSwV9o/A1tTv0JbN0HKg5pSPwqpZylotpPIq2pyM17EGpRTOWUbMT6VPaymGdWBxUOOaPequSSyoMkDkZqIL2qVmBiL9+lNhjLI78bV71W5OxlTjFzIpHpUZ3ADsat3qEOHUdetVTuPPX3rdao82pFqTFRQTz1pJV8uQY6dqac7sAfNSSSkj5hTRi7LYQ5CnPFbQOVFYSZZwMcZrbBNRM68Kt2VdQfbEEH8VZg5PNX9QUlox25qrEhklVQOacdjKveVS3ojW0yDd5SqOXb+tVrwqDLhsgMRmt7Q4g2px5+6oLfkK5a8ctaH1ck/rWbd9P66nZ8Fu1n+Fv8AMykiMuWHWo8fMBjn61Zt2wxB79qJoyrbwPqKV7MjkvBNDgML6D1pOr460/BAH86Il3TKDzWmm5hZuViwtmu3BbJqGa3MXzZyPWtDGOc1XvGHlhe56VKbOqpTgot2KTcnaO1KefcCjlRz071atYVcFmXn602c8IczsupVwTkgcdqGAzgn8a0mto2HTn61RkjAkKjt3ppphUpOC1FEzhQA2R6YoA5yDxTVQsdqjmp0tZT2wPrRcSjKXmMzjGOCO9AZoslDg96m+xSntSPbyLy6fr0ppolxmtbCSXsgG0jOOtMW/Y8CP681DKfvevpUX8XIwDSdi41Z9y/aSL525zg1bdYpATwayozlPWlbOOD070WuEarjo0aR2RLheBWRcP5t1t6ge9RvOyfKGqKN2jbJXr71m9zdTUklayJLp+Qo6d6nZ9qkkVSLb5CWHB61LLKJCFH3fWi2lilPdlm2i3Ddjr1pFXuTgVoWSAwg0jWJDHa3BqouzIrRlKK5RtjHmQv121PekLAR3NSwxCCMKOT3PrVK5l82Xj7q/rVbsl/u6VnuyJY/mAzxWoy7VGOBVGzXdMv+zVy7bZbsfyofYKKUYuZmu3mFnPfqaYxO2hCMcdRTT3/nVs5GxDnJ7mozhjwacRn/AAqLkjnoKgtIQjOSD0qMkj7p6U85xx3phAZqV2aJdi7HMJFx/EKsIvvWYhIbPQitW3kEqbuh7ipa6nZTqc2jHquDUw+7zTCO4p/QEUjS4mOc07FA6/jS9Rk0wEPOc0xuAPSnZ4qCRsnHSjYCKZ89O1Uw+3nv6mrEjYBxVWRd3PQVnIuIm8MMipIm+QZP41TtUfccj8KukYU46elFyhZJWAIBqGNfMfjv+lLGNx9quQxKgyKa1Buw4IFwAOgpHIAp596glbAznimTqUbyTapOelYoyxJ/Or1/ISMevSs/7rc4xXM3d3NY6IXryR+VWLZcnI5qvkD8Ks2vBHHWplsXqaUChTnHNWMbVJxyaijBC/1qdcFAT2rhk3fUrYFTCk5xTScLxTyM5phxnpx3qb3GKpCyL6mpumM1XjIaQ/zqZzlcClIQ84+btmmbQI8UuMEE0Dhj1xUjIZk6H+EdqgREi3Fjhe1W8Et0+Wq0kAaYEdPStYS6CZFldzAGnRPxjHA96hmhaGQbjkVImAPY1q0rXC4+PLMSelSD5BwcgdqiTO/j7vp6UiNvkKqcj37VDVwuSK5YYYcU7kttUYoQ9iOaTdnLenekMkf5wFHUUx2A69qco79+9NIG/wBalaCHBgU+Xoe9JKu+HFLjBHYY5pGGeh4FNb3DoMQ/MM8badOcxlgM/WoY8vMW5xnpVmXBGAKb0YwhH7sBetKOFYgc02EBV2g896cB8vvUPcRGpAUN3NTdMEc1XQDyzk5IqXPy8dPWnJAOYYqFmwOR16VKxx9DURBZytEfMBDg4P5VWm4QA9T3qy6gkAHpVUjz5gFztB5961gDHeSXVWB5q5GCsY5qNgFAPYCnxuWiz3FTO7HckX5lOKjcghSBT87QMDg01sAAdqzQhkg3H2p0Qwoz+FKR940gOBg/hT6DsK3Kkjt0ppyY+elLIeOOlBGEA6imA4c5bsacv3SBSKvAHal/jGBUsBOACetB+70zmgn5hStxj0pCGJhu9KwzkDpQECMQOlJzsJI5NMY1AB+FOyACcZNJH9w96VuBgU+oCAHBPrSq2QSetBBIGBil6DFAAuWBP5VGeF46ZqUfdz6VGQTnd0NCEDcfN3pAMhhQxyBzS57CmA1xhBioWxs+Y49TVlx8uB1qrIm6PaDgZqogH3mQjpipR88YK8gVGq4hAI5FTryhHTNOTAcx4xSSZ2+lJkZIHX1pWJYD3rMBCeOe1O4wPekP3sE8GlPv26UDExg7QOKYxBPHapM+3NQuRg5NUtRCEkMDjg1KvyjJpmcpxTwflUUMBrctkGoXUM3B4qw3BwOhqHoxX86aGNBDd+lO4IK96aAqk+tOJ6EDr1qmJDuMZ7GgEADikbgYAoPzAYHbmpHYDkggdulOYZjINJkgChSDxnn0oFoZ5YouRnrVhTlQx9KiYkShcdT1p8bbuCOa3auhEnBAyMYNPIHyimZOTgdKlBGBismMUA5JpH5OT0pSDjIPFNf/AFRFSANhs56VH976UuRsIPWmKSBgnrVpC6jGjUuJD1qJOJm469KmkB2cfnUIzvbNaR1GOUjBXFSoMjB69qYnPOPwpwU5BY4xSYCRlg5U0mwl8k8k04nchK9KMZTr8w6UBvoIwYzgbvlH601g/mbh92ljyyHfwabGSdwz8vamIeF3vuP3e1PClnAx0qJGcRZwfap0Uht/c9qmWg1qPI+emxhg7ZHFKd2STzikR8AMfxqegEhPJFRA5YkfjUrkH5sVWiZjkgd6EtALG75tvfvSH5RjuaMcg4oc5f2FSBESFABFVVTdJgjvVp8F8DoKjzh1f17VrFh5Focxbc8VlvZASHb0FaLt5YP6UyPkkk9falCTjdoTSYW0IUYNWD97NMHygDvTvXNRJtvUY7Pf8qQN82BzigYznFKMDPrUgBJye9V7j/VEirHYio5VBQinHcLGCs2LrpxVyGISSs2MAfrVa4TbNngA1PAWZh2AHSu96x0Mt3YvSAGLFRNHuC8cip1+6BSYxxiuROxoyvJaBxx09KqSWRJyvB6fWtMAheaTg49auNSSCyZkNbsCc9KkRPLU561qyRKV6daz5YMAYGcVpGrzaC5bbDVYMgGeaaQCQveogWXIPBP6UGQHnFXyivclYZyBVR2H3gPwNTq3mN0x+NR+SCCd3J6cVUdNwbV7kKbWYf4VIUVX3DkfSpFtBwx6VMIwFwabmugW1K0suwhUHPrUe4KCWHPp61JK6IcBcnvVWQE4OOvf0qoq6G1d6mnbfM+cY46Vpw/6vpzVCxICDg5q9GeTmuWurMIbEw70jEGMgDPFNZsDpTlz5XPFYxXUZzM3E3rg1swcxqRWRcoBO44Az2rUtTiBfpXZV+FCW4koCzEY61DMW+zygjjYf5VPcKDKp9B+dRT8W8v+4cflURd2iup3PWk5B4pT2pozXhpaXAf2zS4yOtITgUVI9wAHelHBpMUhINMBcYNL60Drn9aOvTrSAD0zTScc0rZ7CkPvQA4etJgjJFKffmg9DnrQAYy1HfjpQp4pFBBJPSgBRxQPftQTnpTcbmHtQA/vxQOv0pAfakzzgUWAceoo/lQAR1oyQM0AHcmmE5XmnEdxUW7JK9BTSuAisSxGKeuB35oJ9OtIMgj9at6gSsSMGhm+TpyKQ/d60pwUrMYiEkZ6U4e9MyFGM0xmbPBp2uImboaYBkjHGKXJI60KOnNCAV14xSAcDmnHBpi8dTSBD2GfamFF9OacMgZpr9MUIBjZ6DpUEgMrqAcgUO5UEE5qO2DLMSfu10Rg7XAsynacgfWnbS3PpSsCVx39aOnHWsmwEyR0GBS+lITk7c0vcZpXAMsXxjAp5Ucc1GXzIB6U9j8wBFJoBMkZI6UA7gSDgUrsFGM8Go2JIAU4rSFNz2FexN0OP1pSMnnpSY9SPandiP6UnTs9Q5hG5oA2j3pcZIxSNkfWp5bBcZtzznin8HtxSEdqTpgdqm9xjhwT60jHFKT7UNyDxQA0tgCk+YOfSlTJX6UoySPagBQBg85pMZGD+dKeBR7UANPzLmkH3Vz09KRXGCopzHgVXQBxPPPam9TTs46dKTp+NSAcY4pjHnrTwPl54qI4cnmmgIuCWJ6Uq7eijrSGRVJGPrTlfitWBIFbaTijaCo3HJoZwE5amI6njuKizAkcjAPeoSWLrxgVNJyRxUMm7AcGnATJFb1p6khMnqaiAw24/jUg5P1pSQxTnIHSlPTgdKRsl8D86GOcACpAjlHy4HJpiI5HB+X0qfaT+NIcg8dKtSewaFbcAGPTAqsHDkjGMcnmrshKr8q9az2O0s68g1rT94VyfzljjaRjgdqWFPMIlYc9agnjDQhegFSrMojGOcVq4WWm4J6knlvJOGJ4Ht1pyKBKcUxHcg4HLY/ClhkAJX06mplGSVuwaDJlDZTscVZChefTtURXdMJN3ygdMVNjcO3rWbegx5YZ56mmlQDim7gz8c4pHlCHPUnpWdncB5O0AYper4qJctLz0qZCBnj86TVgGkrupyn5xTQcKTikjztJ7UMCUfe5oYZP1oHQt0FIvGcmpAQNl/QUqYBJNIVwD705RhunGKegB7+tBBxilUZP0pFHJzSAjZcZI6GonbB9vSpzgqwqGQHIG3p3q467gOTAG3FVpGCScjjPrUq7w+SOCeTUFwRIzRgcY6+taxjqBdhZSQAeOxqRhlSO1Vrb93Eqnqasg/JWVTfQZWmjxbsBzS2+RB83WpioC89DUY69PwpqWlifMQbgSxp5wMGkPQgn8aaxxDx8x70rXACSyn9KryDI9R6VaGPLAPWq0gIcKD1q4PWwyjJm2m3dc9Kfcu3lbxjHenu6o67uTnApZkAJB/KuhPa5NyO3KmHAPXvVsMVUDGc9qowyhZAmOnUVorgNtBwMUq3coaAxg4GGHamjzFtgWPzDrmnN5hgOF+bsCaeobyP9rvWT2EU5X3Zx1oZi3B+8BRLIUcsQMHr7U2RVKb+2eea3SXKhEgb92wJAX6VDcvhUULk0pUb+ASg6VG/mICdvHcUJWdwbYk8qpAsmzLdhmoLaTPBG4noafc828YI+pNNtiN54yFHGBW8bJMWtzRhV9i4bIpRI5mIYD5e9NtwYGAP8XSmOrJc+Yw4btWGnNqN3NBUwpPU05VypPegHKA4pc5Bwea5r+8MZMoeM5NZ0mxGAwd3XJrSlU+WeM4H51nyq0h3MdoXvW1GWu+hLJLK5EVwu7q55NbqPXLQBg+EPfrXSKzqF3DjHWvVp/DYxmWep4pD0NNVuPanHOOK0MthoJK7ew5qVRiP6/rTFAIp7HGAOKqJLIp4/NiKevSsotuGM4PfitodKzr2IROXPRq1i7HJiI6cyKpYL1+YelRD52yep6e1BfzGz09KmjGDuPIFWzjjq7ssW1g8qs3ZRuJPYVcjkWWMMvQ1Uju5I7SdFbCyDaeO2aqwXLQOQBlT2pSV9jenVVN26M1ZYVlXDCmR2yRSBh17U1L6EjJbGe2KeLmI87/0rOzOrmpt3urmtYv8AZrC9u+8UWB+PFclenfJsJ4QYzXSz3KR+Frkj5vMkVcdM81yk2Tlm5NCtd9zGtLWy2/4f/gFQt8xJ6fyq3bSb2AY5YVTwC3I60KzLJuHBo0ZUJtWZrXMaBd2Oc1Dapm4zjpTGlaQKSORViyB+dj1NC0iPSVVNFsDjnpVG8I3oPTPFX6ilgjkOWXn1zSWhtUi5RsjNBBOB0HWtG1AFuGHeozZKSSrYNWY02IFHbvTbRnShKL1FOcHFZsS+bPz0atCVgsZPaspSy4IOCOlCVya7SaNdQFHAwBRkDFUPtrleFyT3zVcyNICXOcd6XKU68ehtK4JHNTAcVz3mFWyOK1rK581CCclabjYdOqpuxVvogspcD6iqcg5zitOddztuOQay9y4IfqKaszKtDlenUsW0AkUsxxmm3kXkRFg26pYXAXb3UVUupjI+wnPrUqTvoaSpQ5LtalUIzKWxz2pqrnHOPStBjgVSjXftX1pRfccqajaxGxBqeCLJywq1Hao+Bt4+tOliMUpUHOKqOpnVi4K7FjlaJ8qcD6VZF+wGTH+tVApYk9RQQcY9avQwU5x2ZPJdySpx8q/zqa1t45IiSM5qjjpx+FPUsvKEg07CjU968lc1IYFiB296hvEklChBkc5qql5IpGTmrMd6GHzLgfnUtNanQqlOceW9jPZTGxBGDQ+ATk8/SpbiQSzM38I6VXblevNVqcrST0GM34VEWxz3pznjGaj3c8nNI0XmxST17CkCn04pQB1JpyLuGcfWp2GLjBz2/nUsLmNtw6d6YqhuvSn4A2g9O1PQLu90akTh0479qU5rOjlMEn+yevvV8NnBFTJWOylU51ruHIpwft0qMtt601G3c+lTc2Q9269s1Wd6sbcj3qCRMEe1JjSIGB60xvQ07Izg9aDg4FSythiIOAB9aeeWAHSgDFOAxzjmjoGwRR4zjvVkYC7c81GvT2oLYbgdaaSQtxW9M81Vnbg1ZdgFzWfO+QaUnZDjuZV4wMuM521WGTx0qSSQM5Y1GvQn8vasOhsttBQAcY6elXbVQpBx9KqRjcwJ61ehTdgbjWc3YpF+PnkniplbceetRR8rjuKkxxkfjXFK1xkinAqJFYKwPJNSL8wGKSTcT8vXFShkQbbgAc+tTZG7PemJ6cEinRg5OR0pyEO5wfQUoYFgBSYJYjPFKBlgDUjEQ5bg5ppHz5p6jGT3pjcd6FuIR8Sls1WOAMY/GrZ5+YdO9U5VIcnv2rSIMSPAY4HJFSxonXGM1XiJZtzD8c1bJG4cYqpvUOgM2FyPXrUZZQM9zTnYBODzUTEFiq9aUVoBKjAITj/61AGMGkjyx24o5ZuaHuMe+SQKa5AQ7accjLHj0powY84wT3pLRCI42Hn7Qccc1NMCwIHb0qJCnnEL1qYgtCSTjmqlowEt02qMnJ708kc4FRxk7cN+FOboQOtQ9WNCDAVm9e1EQ3AE96QZCYGKIFKoM0PZgSDkcioZDgFulTAEsc9B2qtMTg8de1EVqJg5CgN602DCxMwABJpEGVCmpFC5cKOOtap2DdCyj9zgUtoQY8EcGlBEibVHI7UL8ucDp3qZO61AlAIXb1NREHzDzwalJPDN19KifHmFs9KiIDieQD+dHfB4zRgvSZO4Z7Gge6HP1ORUYO77tSPw3pmmD5BgfjQtgJP4hx0pQeN2KYSfrzTiMg+lIOoEchvyobPFK2NoHX3pvagBwIyc9aQMdlJGwJ4FN3FuMdKLAKv3cY60gPyUgOJCT92lGcZPNOwWBfuFjS9/ak3DgHqaVR2J4oYgdip9qMDeO+KDSKDnOelAw5GcjNNPb1p2SeeoBpoOScimhCs3FQMdpyeFFTsMKTVdwGfB71UR3DaCcnpUi8YyM1A6EbQeBU6uOD1x3qnsSiXjJJH0pFGTzxSDlQQKVun0rMoTvnHSngZfOKiJJJHapVOB14oYDcDJI7VA6k7sfnVjgE+9RsPl4604sBqkkEflT1+6faojngipU4XJ9abQAe1RuDxjt1qRuQvvTGbbJ70RAhJ5zT8kHPbrSLwT6GnHgBcVTEDNlsUHP8PANBOeR+NGdqZpAL0XJ60m4Abu9KeSD29KQj5sdqAIJc7t9EQbzN3appFDcEcelMOVcbelWnoHUOTIQelSxghqYxOSAM+9LGxZ8ZwBUvYB5OW6cUNyeOvpRnCn3pFIA5ODSsNDCuCTmoeWYbunrUwORnNRupwSOxq0AFjk5FQMp8wc8VOThc96hkJ4FVEQ9R+8wDxTn3E9fwxUaY35z+FPAzz1oe4DcGPCg5BpZQp29jTh9/27Ukadd46Ur9Q3Gs2JFXqppPLxvRDz60sRPJI4ojAQszNnd0qthWG/NkIOo71KSS2PzqLckTgEc5qww+XI5z1pMe4krgAEHmjcrbRjg0bdoAx8tKgxISegFSrASEcYHSo5FwVIOMdR609nAxjoajmOCh60o3uPQkyPrSOOMDqahBKNk9PSnq258j0osICSqgjkmmhckYpoLHI79qlQbTyetU9AHum9M9qjjHyZzz2qSQHyz79agXoPm6dalaoCxj1OKTdk015ARkdKhDZbg0JNhcsk4PFOyAOTUR5IwaPmJ68UrASk5H1pkmdvvQCenakZiCcDNJDMu+jb7x57/SrFr+8QNjmnzn93yKbYMPKBxXQ5PkI6lo9OnNNbrk0rMR247Ur/AHARWJQ3t70gOCM/lQwI6GkGOPUdKYhzEbuvBppUNnH4U0njp0pEfI4Oadh3RE0CSv8AMPxqvPGIuNuPerjMNowM+9VLjrjpitYN3EVrfrkc5qYDBz3PakjUAYUYpxIwT+Vat3Ykug4ttXOeBVSWfcAAc+vamzTlxsHQ96r56HHWrjDqwYvfkfWmyHKZo/iJGTnvRJ93HpWq3DqjVs2BjTIGe1XomOSMVlWjYjFaUTfMeOa5Ky1CBM+GbFTIAVI9KgQZep04BrDYfU5u+Upctn16Vo2fzQDIqpqS4vCQefarVmcwCuqesEJPoTTr8qnFVpzm3l7/ACH+VWLg4hBPaq0mTayn/YP8qiG6KR3R7c0vX6UH7wHalzz0rwx7jScrTgeMGkxzR296QATz7UgHzCnDgc0ZyaAD2oFHvSKSQc8GgBMEnOaTPtTs8kUg60wA9c0Y79jSj3pQPyoATPYUv3uD0oB60dBkUgBQB24oPJ96DkCl980XAQDnrRx+VKMUg9PWgBw60nWj3oGAfehAIeRiowmOTUnJPNHBFC0AiXA5FLjnJpCCDxTmGB71bDoK+PLyOo7UqsdozSEfLmiNtyip6AKRnNKeOMUEjOO1KOTmkAzBLDjGKefWg5xxQ2MUAG4D61ExwR147Ux/M5206NixI9KpR0AkDdCelEoLRnYfmokUkUqghB396S7huZsME2S0jZqzuCsoYcmrG3gjHFV5UIbIHStfaOT1DyLAYbcmm5y1MRhIuMU7IXIxgfzqbW3AXABJoyMFvSmnLDI60pcLhccmpAQLvbI4IqTkuBTcsCABmn/xe1DYCNwTnp2pjMMptH+9UhPTjg0hAAxt4qqdTkE1cQZJPpninnn8KTcB0prybcD1pc0r3GPGST6U4YB5JOe9MB+T60p+9/SpbYNA33smlPWmnt605iSMVIDGO3mhiTTiMn2FKOcUwExt46Ug5JJp3Q0p69OPpQA0Zx70AHd+FO4zkUg5BJ60AIqgvuxSkc89KUdvWlPJNFwGd+elA9f507gmg9hRcCNs/jUZABqXBPOKYVPNNAU3U568GnxsFByCcUoCg4JyR3pdqo2c5zW+mwDmYCPcaYjHcDjOak4ZtuOlOMZUAKKi6QC5yp9aaQSuO1KUJGKXjp3qdgEZSFXHepFxtFIRuXbSqu0AHk0r6AJwGzQMluKUglTRjGM0gGkkPzzSkkGl2gHOefSkI+frx6UARSDI21SaJo9qL3Oa0BhiWPakdQ4znpWtOo4CaRmfOJSW/Clt7eTe4PINXTEsjDByRU+3aMenWt5YlpWQ0hiLtAz9Kgk227Fz/FVzjHNRTQrKmCK54T968gdyKGNghYtndUwA27aj3rDGF9O9PUjZuJ5PSnK+62AgjYRuQOc0XRITeB07U5gu4Ed6LpDJGVHGad9ULWxDFcbnBA4FXsjbx+dUoQUAXoB3qcPuBIpzSvoPoSAEqeKADtyeKarlhkfhTjkL6CsgH5wAMmlx3x9KTqBTVYs5FKzAcffvSnkcfjQ/JBpQw69qXS4C5wMjoaaThgKcPvVG/HShAOGFP1qNixU461IxIw1Jj1600MrRnafLJOe5olUckDI9KfMhIyvBHWkZjvAx161pe+oiKKTMeSvzA9KthjsziqO4/aNw4HerkbB1O05qqi05gJFGVBPT0qIqxY7T16U5CTlaXgc5rK9ncRGANh3Go8iKE4FSL8uTjimzOFUvjIq1vYAjO8Nmo3I3g9xUkTMcEnjHSq07M52R+vJqlH3rBcikycgLkjpSXAxGD/ER+VPlAdPLzx60HAiDKCy9K1QioFRJAwHzH9K0I2IQbuKzHV/tBOOPWtBW+YIPxq6msR3Jg5aElRg9qeAy2+GOT60xSRIVHA9aAhzgtXMwKV44VCGOdxqZvmhCjqe9Mu1VlCdcnrjpQ8Z8r5eg9+tbq3KhCIVQhN+FUZzUe4ywSZbHOAaabdrhVz8oHXHNStCwHlg4A/WtPd67gVzFLEmxxkdRk0RF0m2KuAasSwhyrF+UHNV7Vy5kcjjstaKd7hfUvQt5kYGckVYYBioPOe5qskfkodvO6rog+VWP3hXJNxvcaHBNqjvRgMwBqRuQtREFWz3xWKd3uAsuSSo6Vn3CE/Ic7R+taKkkciqdyC67c8+laU3aViWR6bB5l2ncKeTXQ4GMYqhpVr5MRY/eatMKK9umvdOeo9SAxkcqad5gVCW4x1qccDPasm/uN0hRTgDrx1rRRuznqVOSNzTC/dFDn5unHam2b+baox69D+dLJ97NOwJppMcvPWql8GkhJA6dBVs/Iqk96jPzfSmDjdWMVPmApxLbsAZPY0+6gMTZUfKe1RLw2Ca0Wx58otPlkSSNtQADr1qAEcD+EfrUrHf2qLPA45PeqIk9QOD2yaVeSCe9BOeCaei5YAjHv6UInU2NVHleHrGHvI7Nj6Vzb5GQeDWrqV0908SsRsiXaoA6VltnBqbJaml05lbnJ7UbehI+tBI5zzmgAkhc9azaOlLQsJu2da0bRdsAPcms0BtmM5rWhXbEoNN6IqivebJOgqmL8E4K496tSPtQmscc4BpxSY61SUGuU00uo2IAPJ7VOazLRS0659a0+p96TRpSm5xuyvdki3bn0rPPJ461cv2wqr2OarpE8jEqnH1prY563vTsRYzwvSggdu3epGRkYKRg1GCM+zdRTUrmXLZ6iMuCdx59auWXy76pu+D1/GrEA8uHnq3WlJqxrRj+8LMkmDnNZYPnTseo61PcS/LgHk9BSJH5aYUcnr71m3bU65R5pJdhrzCJGP8AEaqo2JQN3SpDF5wkJPC4AFNhtycFhikpIU4Sk7ssyHqfyptpFlA3r0qfbvFXIIQoAA6ULU05dbj4I9q1nSSbpWbGQa052McLMOtZsKF5FX1raK0uceJk21FGrDHshUDn1pssKyRnjntUqnBx2NOYcZFT1OhxXLYwxk9ev86ezDBwevanSIEmZKYo3sF7dhWlzzVFrQB8xz0zQThcdD61rKqhQuOPSqt3AirvXip5lsbzoNRvcpFcEAjFROx55qysbSOccEdahmhkTJ28Cm2r6mXI7XS0Kzcf1pig/jTiMtjtS4PBPH61JokKoP3acAADxz3oUYJBpwAxnHAp2uTKVgQE5x09KNwxgik34zwCaQ5zjpilYL3tYRn5GD06VbtJtyFf7tUjyeetS2h/esDRLY2o3U0XGbI+tOjbb9KTOeOlIMLWezO5FheMgfhTJVyBSRtk/SpetULUzpAA2fXrUJcjj86vTxjPtVKWPvWb3LTHh1wAPzqQHIzVRwx2E5JBqymSOmD6UJ3dgew9WIYg9OtSSD5QRTQO/enn7hpi3Ktwx257Vm3ku2MjPzGr8x3AjpisW7fdKFzytZzfQuJA2NuTTQcA4696djPbOKbzjms0bK6JoeXB7CtCFe+Oao2yljnr71pqPlGBwa56rBFmIccCnLkr7GnKvTFKQe54rkuUCMQDgc0hU7CSaRnGSR+FOU7h7+lJ9wIUAUE9aljGVzmopMZc9qkzmNSOlUwHMOBt7UvGCT370g4wM/WgDjB49KkBR93HemyDkD0pQxLc8UnTqeaFuAcZIxx2qJxuJX0qTJA4FMON2apCK3yqx29u1SIS/JphD+bweMU9D8h71owE4KEGmiPkHPy1KSFTJFIDmNfWlfQY1S3mDavA96eAQ/XmmKoRueppwYbuvIoEKWVnAzz3pm478dj2pwAVwe4qORHO7BxnvTSC45GVZSalDbkIquAIgVqZMMgGORRJX1ASJ/nIPbtUhAYnNQq4+0Nxg1PxtBPU1MgGY2qacOFUA0j5C5A+Y0gJ2Fu9LcZI7FSD2qrcHkAHr2qyOV3EdarXIYyALwfWnDcT1QDO3jginxhSpZTyetNVixY4pEO1gB+NUMmjHlZ9SOaFwX45Bpp/1mQcCl3hcnt60nqIcVy30pHxu2jvRu2496GP8TDmpQbCg7WGOtGRwc0iYwT3pSMZagYZ3LmmchT6mpCcjJ49Kiwe1CAchx8x6VLk7WNRjJOz8qkIAJT9aHuAi/cGaAOMf1pc8H26Ui989akQigANilwNhyOaTBCE+9DHC5z160wuNPzA+lDZJGKO1K2ccDkUxjGCl89SKeGH3aZj594POO9KuDzmmxDzxwO9IQOMHHrQAd3NNP3tvU0kMVjtNJ0GT3oc54HJpDjaBTQANx71G5wm4nmpMkBh3qFgGyCeKcQGKCVLHkr0qSMYXbnOe9MjYjd7VIpAjLDnNW7isOjbGB2p/U5qFHBJA61N1U1D3GRnPmAeoqVQCNpqNm+YYPNSZA+tJiBhlx6Uw8qeOlPY+nU0zB3HHT3oQxv3VGelA+7170duTQOeOtUA7GVX+dRyY3Akc1J3IxxTXUN8ppIRAxwwzznoal9+pqPgvinLjeSDgVbAACS27ikLYXj8qcuGyTQOPxpDHHkZpjnOCOo60qcKcimIu2QnP4U0hbjiN7UwqfOPPShAQ2eoqTjv1ansG41hk8Hk0RoMlh1PWnMoCDHBpUYBPelfTQWgMelMcA8dKUkjjpzxSA9N350IYDDH0xSdXI7UAck0mcDd2pgMwHOBxio5mHmZx+NS8rkgcGoJBtXbng1cdxdBykKTk9alHMfA4P6VASDtFWBxwRSkMbIpYjaenWlGNxz+FBwnJ5JprhVG0c5oQAjHzCH+76Uhj3MSDlKVlDQEE9aAwjwvXPSj0D1GKquwOMnvTonczMhGF9aCpC4XijcRID1B6U9xbE6OGOOtRM7dCPrUoADcdKRVIY89e1QrDBsLGMU1U3R5PUdKk+8uKQjaSKSYdSIjdu7e9ICF24HHrRLuHKjJpVwUBI5Wr6CDZ8+PT9afwwDYozlwR0oHyPjPFIY6VvlxUaLznHBqV+FH86anG6knoMhZdqkDvTFG7APWrL4O3iomG3AHeqTJJFTA3E9KcTjp3pgDDHNKMZ55FSx9B+cnHekbgYFIT196QkgZzSQFef09epqKzdVRl64NWW5GaqMhikZ15B61vHWNhPcstKCoGeacZvkFZP2gMxUrj3zUiy8/Menem6IrmiXyM+lMLcgj9ar+YdpI60ivx8xPFLkGTM45x1qGNyY+evrRjIPy004QHn8KpLSw0TAlfxqCbllJP1oeUEYU/jUDHPGOB0qox1uTpYnzzwPxqGaT93he9P5YE9frVaViQFYfjVxWo3oRsRjnGaZjCe/X6UEH04B60hXjr1963RKQKPm5NK7DAAHWkPGO/rS9VIPB6igfmWrJiU6nitFM7vr3rNsclSFHPcVpxjL4xwK5624Jq7LSjCj1qUZqP29KkjIycVxjMXVEIuM9j61JZEGPgdqXVuXU45plkTtPHFdd700JaMtSDMH41UmP+jzf7h/lV0qTCwB5zVKYYtnHfacn14qae5R3uOc0jZx70juFBNCOWTJHJrw7PcYozty1Lj5aDg0hwQcUgFBGB700A560qAgY9+tKOBgU3psAg9KM0A5zQAMUAMbJb5eKcucmggntQv3sd6AF4zil74qNmAPWpB16UWAOnPrSNnbxTgd31pO/vSAQ8gnvSRk7PenDikB4p30AWjoOO9A5PTNKfve1IBfbpRjJ4o680AcYpANI75o9ajl5AwaeCCPWq6ANPUUhHPHanNx+FRbiSapaoCbqpAGKiVwh2jk09WBUDPzd6ruRHk5+Y0Ja2YFjeD1/KlXluOKzvP8AnQF8nNaCrzmqnTcVcOg4NzjvQTx0qJ22yDn60m4lyD0qeUCdcMDUBIikJPVqcp2nikZhvycZFC0Yibceh5pe2BVRHLHGetTJKrHAOaTixkgBGaRhkc08cjJpMc89KV7iK4UoSc5pRkxjd1p7CqxuMBsdqtLmGTxkdBQVAO4nntTYSpTI61K67hg0nowsIHwMDmngg+4phIXgUJyME9KVtAJemOKTGQaGOFOOKTOeKlK4ETKSD7U0bivPap1OVIoIDjHNVzWAbGxYA4pzcNkdaRBjgDFO4Umpb1ARctyVxTu9Jnnk/hSZC+5otcB3t2ppJB4pA3YDnvQSR707WAf1YGlBIJ9aa2QeOlKpyRgUmrAL93J700AjjNL/ABe1Iwwc96AHKOc+lGeo6GlGcDtSAYJPelcBcYpKByfpSnpQAh4H1pm3A9KfzkZPFNZgz8dKaAjaMbhxzUZGCSxyKsHDNxUTAKMY4q4yAF+blQBT0zzmotwzj0qRPlFEkFw7c01sAg55pXYDtmkIBXNJAP5J46UpweD171VS5VpjGvBqyegyOabjbcAJIAP5UHtmgjDY7Ug5qRjx94k1GzYPPepD6Z/Gopck8EURtcTHEnIAFAGcgDAFRo+/p071LnAqmrAR/Kkh96UfdIzzTSNxGePao2m8p8FTj1pqN9AJ2OOD2p23Ke1VjKWIbPFS72Z14+U9aOUBk8IkhbqCe9McmO3BJ6deKsscDaOlRPCvlAZqoyVuVi1IUkQypjgnmrDDMYquIVVhg1M7hcDNVUUbrlH0KU8+xuO1TxkeVjPLVj3DkXRBbr04qazn3T4LVu6Xu6CTubMC7UNPwdoBHFNjO6nk8VxybvqMUcyYNNAwCRTl6ZoPpSAQds9aQHGaEUliTS7QTn8qBhu55phyxzmnADPWkx+8AoQkNcnGSacTwCB0pNwbKnpSfxFex6VVgFDBlyetVnkAk68CpsBgR0qKS3G7HU96qKXUCrKspYkn93VqxULEWXhfSlETeT5fUnjNTImxFQHmtJz05UA5GyScUvUYxTGAC49KeBhR6msHYCCXco2huCeTiop2XyiCcZqS4xGMn86x2Z7qcY+6DgCuilDmJfY0bWWSZiM8Y4qRY9iuSOvSlt08h9vTtmp5MKcetTUdnoUZ8oCYXcctwKUvhVQDk9qkmjWRhzhh3qF929Ao+6elXF3QiG4i2zjJwFqaBkkQsOPU0kjKtyWfnjgGpIQjFVxjBzWl/cDqWOShXG044NJAphiJc5NJu3OVU9Opohbzmdzwo4rDVJjvqR3LMEDIu4k9fSoLiUkop4DVZcEAfL8tM2iTBHPoauLS3JsNjDAKENG9Fmw3zNin7tqkAgNVdAW3bkwT05q/idwAMzJKdu0djUEQaCLIO4sQasrhYDE5+ZxxVdojbW8Qc5fJzW0FoxN2ZfQswVv4ccir8RLIMHvWbGojUyBvvdFq3ayEqQThq5auzKJnIXqeKiZyRxUz/nVdchvqORWMUh3Ho5aMMtNMfmSqvY9abGWBIIyO1XLJN8vmEdK6KVK80kRJ2RdRNihB0AqXGDSDjNMllWJGkY8DrXtRjZWOKUurIb258mMhT+8bpWJnqD3qWeZ53LuMZ6CmxxtI21eT3rX4UedUk6k9OhpaTIdzoehxWgykvyKgtIlhi2r17mrTcMH7Um76nXTi4xSkMuyDcMFPygACoD6CnE7iXPegDNSkbaWIZVDqVbmsx0eGXaenY+tazqc+1RSRrIhVhwapOxjVpqW25mE+h+pppyQCKmliaLhvmB71HgYH8q1umefJNOwRx+a4Ve/f0oD88DHpTlwpDH7o/Wn3iLFOwjB8pgCo/nRcFsRu2VKHj0qnKWBY9xVjeCwxyPSoboDaCTmlLYIv3inxuxTxngg59/Sm/wAR9R2pyZOB3NZnVsrEwwFwRj0rUS6iK9cVmrjORyadk59j0qrXIVRw2Ll1OhiKq3J9qoHHrT1AzzSdeW6U0rCnNz1ZaslzMWPYVeIrIRipLL2qQ3Uq8F+R3xU8tzanWUY2Y+9bMygdVFXItoiUoPl7VluXkJZuW4p8N35eEJ+U/pQ1oFOr+8d+pcuIRKvuOhrNlVw2CMCr5l3DIPFRluKi9jedFTdyokO4guOB+tSyyBVyxwBTZZ1UfMcegqkzmR8kZ9B6UavcPdpq0dyeI+a4kbt0AqeSTy4y2M4pFj2KAD0qvcPulVF5x1FLcv4I+Y+3V9hBPDdaV5QJBGg/+tT1IVDjqKrRA7hzmiMdbk1JOMVFGjaHzXA24z3zWgny8Gq2npgMxHpgVeNXpewqcpON2Ur18bF9etNs1/fEnt3qG8bNw3OQO1XbIAQbgPvU3pE54vmrPyJufyqSNs/KaYcYNRwyebGHHGaR1X1sVdRXbIp/vdaZapukXI+71q5eKJLUkdRUdomYSw70/snNyWrfiWArEnHIFVLuUb1QdRnmr6XL26OIyFJHJxXOMweZ2PQ1MTWs/dsaNs42OQck9TUp+YfWskErg5IJ6ipY7iQDn5gPfFNq5FOrGKUWWnt0b7w+lRNZgncDU8E/nAjbjFSMwUZNKzWhvaE1czGjMR2t+VMZ8Zxz706eQzXJGeF6VHjIz+lXucU42loHVicZPpQMEcHmmlSD1oGcdOlFhIUg8AHPpUttGfPUY571EGIyfSpLM/vc/wB2k72NqfxIvshGcimnHAoMjNxmgZ5zWZ2iKxRxVvsMVUwTwOtWUOUA9KaEBGQQe9U5o8HHaruOwpsqqy4P50mBnkZOfSlTr/KnshB9hTABmkXclpTnZ/OkHPJpXPGQKegtSlctgEVhO++RmA61q6hLsQkdelY+OenesJO7NYIU89BmmgnGcf8A1qUZ7GhSd23HFSWy3aLkgnrWqFBj56e9ULNBgH9K1cAIOO1cdZ6lII8mMHvTz0BNJwF46UHlMDkj1rneoxrcgDNKud27FNz2HanITgCmAxwQMDjNOA/dc9O1K3HBpwyRg8UdA6CHg/WlJ5B6gUAnoetKe+DUgBG5ge1JIACCKXPy7e9NkGQFHWjqDEC/IQe9QlQo+pqZuPqKikwIwxq0IjkJxkUqHGR6ClOGAz2pq52sD1FX0GK67uQeD0pFIKL6460770fB6U1CBGSOfejVAKQdykDJoHL7vzp+eFbvSEAMQeppdAI97FjmnMPlDA/UetJuBfbjj1p4bkr1AHFVsIilY7RtXLHvUkXER/vVC8gTaP4qnQnyx2zRLYBioDJuJqcHcM+lQ7wkwQj71WB8pwBUSHoMPzKOeRTEJKYPrUuOM55qu5KtzRHXQCcEsPpUTttbPXNOjYP7EDmmyNlgccYoS1sDEAw+R3pVUFQQORTQd+0jgd6enRmpi3GnBb8aYj+aCvcdalOW59O1RRN+8Jx+NPoMkYFgNo6GnTLkAUhT5DjgmlcExY/iqQEUEAkinEDgd+9RxqRgMcgU8MC30oe4gPPHaog37056CpExhvUVXLAS4A+tVFAWUO5sg8UoOXwKRQAoxT1A5wKhjDIDe1IMKDnvQTgdOaCMADqaEAfwYNI+SRgc4obluKE5A5oAjB3ufSn/APLPmmHdnAHHrThjyxTYEZ+8PT0qRQAMY69qY2QAMUpIXH6n0psRKB8hzUeMnd6VICMU1skgipQDE6mmgYUkmnquD1ppUHjtVJgCndyDwahmHyEAdDUxPyFcU1shgo79TVLcBgG1OnJ60vAQEjNNAYyHJ5pHJjwoGQe/pVASQL85Yd6kP3qgt3fzSCOKsYBJ71MtGMgl+WUHt61NkZyaibJf2pSeFHc0dBEvRuaQ8gg96UkBQP1phPyf1qUMRvucdfWlUjaOOaTovvR2HGKoQ7pwDzTWOAT3p49TTSeDxxSQyuV2SDJ4NJHyWH8IqSRRgA1EDtbbjjua1WqEPwQQVNPOMgYoOMewpQMgipbGN3HPP4GkUnJzSFyXGKAeSD1NMQBmXII5NPbAxnpTW5kGKc6k456UmxjifkDYzTExhgBzTn+VcU1SEZnpLYW7BjhDxk0yTIUAZ5FPY/JuA5pGPSmg6CLnywO9DEBQp701WO5iRxQvPzHk07AMkO2PA6d6a2PJBA5xUrYDEVE4wcg5q4sCLIG39Kspkr71AMHaanT5WIxwaJbAK3MYzTFOVzjpUjrkgYyBTH5AAFSgGq4KfMOKVyAi4HIqAMACPTrTopA7EmrceoiQHMjMDSRNhdobkUkeDJgDAIqJQEnLd/5U7X0GXI2JBYjmnbA2CTyKqRzfvcA5WrW7awTv61nKLQkOwck9qbn56XJ3lew70hA35PWpRTGnl+e1CtkFaXjGeppAcEYGCaoQhU54PTrTgAzg01Fw+D3pVOGxjgUMB7njntSLlmBpX7tTQCExzmpWwMkPPbpTGwScVJ/D9ajJIBHehAD8gU0HGMUDnJzSMcAAU/IBz9ajc9vSpAQQCajmIGP5U1vYTtYXP50xlIOCO1C/6zJ6Gl5L4PSq2GUpLVSxbbVZoSBuT8jWyUG0k96zblTFuK8VtTqN6CsVhOyfKV6dfeni4AB3J+tRzspO7rTFiJ+ZWzW1k1dkotLOr5w1RsUHBOT61WdNgJ+76+9RM5GMHI7ZpqCewWuWvOG75R+Oal87bHkfjjtVPI25JBJ7CkSXa5PX2puFw1JvNZm3AceppHJDfMM+uKZ9o745PIFDyAgHHI6mnbyGJIQcfXpUbcL1x7U5myCR+VJknoODVIBTkdMGgMdh5yaAMIegPel6Z4Oe2OtANE9lIRuHQ+la8IJUZ6VjWmPO284xWxEwwOeprCutBJe82W/SnLnJ7UwDrzUgODmuEpmfqoJVcHHrVSyyAw7D9av6qu63rOsGw7D0rrp60hdTSjBIYe1Urj/VTAcjaf5VdQ8kHvVCdv3M2c5Knr9KUPiKZ2Lgy3G3sKuAdhShRyT1oGR16V4rldWGg70dqOnBpuSTwPxqAHAnPtSY5pBkYzSg85xTDqL24o6UqnqMUmDzSAM5HXikA54pBQSOtMBGGWGaeMk+9RMu8gg9KlxzQ9gFz81Bo7HPWm5wTSsAjHn2oDHdSbh3704AZx6d6bAXPpR0waVQDR1br0pAAHPsaUkZopOhzigCCYgcU23P7vBqw6jOcVERgcGrTurCBiR7VD82d2MmpJAd4/nTlGRx0qo2SuMWIZ571HMEzuJxU6deetNkwqkkZFQn7wMzfIjd/wDaB/KtKPlOazZ2O7MY4P3q0YTmMcYrepdx1FEbcQ7ypzkikjB3crirHbjmjgGsOZ2sMj2AE00xLu3VMRk004xjFJSYMj8lduRUNvGiSbyME9KtI2BtqOQIWAbr2q4y6MNifPGRTWOOT0oU4GCOlZ97dqkgjz8xPSlGHNKyE3YvlsqSOKoSRnzBgcHrVu2/1AJ6ntSSEbgF5qleMrDIrVRubn8DVskYx6VAF8mEydAzcVIBg5FVVjqmxXGlwq8nnNM83D5B4HaiZN688DtUYAVsdcVMUmhlhHLjOOKcPmIxwBUcbFQe+e1PRwc+nepa7APzycHk0hJAoUrkHtSowZjjmpSbEABUcilPyrjvQDzTQ258UrDHNjqTUWScCnvwTTRjIJpoQ4Dj6U4YAoHXnpSj7uaTGC5PagHBpFJIzTmwOB2FKwCMDik53Ajp6044qMyYG7tTV2BISBx1pR1+gpucnPagGkAoPBzS9Fz3ozgGmseOvNG4CA880uQCAajzhcseaVGB6n6U7AKflNRSMC23FPYjIxVaZyp3Y5q4xuwJNwPFKHXacH2qssoxyaRZIy2Qcj0rVw7gWwxckHpSkfKV9ahimjGRnmnM+05rOzTAhCqsioOMHJq3vyw28iqUkyxncPmJqzFIHQcYq5J8twJZvu5B/GmxksuacSGjzVSCbDFOhzzUQjeLQFxQMZNQXDfIRipPMzjC8GopgzJgcZpRWuorCRMoxg8UrykPhR1qtHEV+UHJzU20yN06dDW0krjFlcJhmODSeaCnI4pJYSwCenehoS0YXuOppe5Yl3CNm35bj2qVDnJAwO1Aj2qSRyKcn3QMYxUyfYocH3rx27+tRzhscDNTKoBwvQUj8jAGahNXBlGPBcA9e9SSLkFhk47U54u+MVLGmIya1lPqgMOWHfOWwRilhg2SqSOta8kajkjk0xlVCOK6FWVhbEsaFEwOnY1KfmGBQgLRjNIGOQoHJrjk23cY8Zzt7ClOeeOtGMkUHkjHWswDpwOtMdhtxTgTnnrQRg5IpoBoBHy5pGXLEjrUhGST6dqTcNpPr0p31Bke3C80Z4/lmpR70xlAOCM0+YCIffwKUrlj7Uu0jLHr6U9SCm4d6dwGbj/jT9v7sHPNCkY9KjeRg/A4pJXdgFJAiBbnmlDfLnoaZMucUucxkdBV2vsBDeqJI9hPzHpS2tmsHzNjdj8qmRN3zEdO9SZ4OaFUcVyoEV2VvtGSeP505zkZ9KV1y2TTgFEZwevWlfYRSDFmLlearorm4yD0qzNGFJI6d6rx7hcrgbc10QtuhdSKcM8wLcbetTLJiXr0HQVL5aCTI5qvaqpDNk7geeKtWcQ2LsBV1JHPvUi7hGEx83qKZEgVshsii3Em9i7cnpWL6jHsMQnP4mqyMQhK8ValJaBgoySKoxfKCCcY6jFEFoxX1Ji21dyjrVeYOk6sv3T1NW4nLIRtx6Cmzs25V2bs1d7SGyplt7vH1j7UzcksSSOD5h6ZPWrbHZINi8N941WuFhdomUEqSR9K6KbuiHZaEsMhe3D7cAHGDVu3JCg7eWqnHMu8wFMJ2OetX4Fw20HgDpWFVLexasywT8oxjNNSMcs3U96Uc4HQZp/BzmuTXoMqjK5A7nitW2j8uFQfvHrVC2jV5tvZeSK0N4AzXrYKDtzM56sug8vgGsm+uPOkMYPyr+tT3dwwG1T8x71QRfMYIvXt7V6UF1PNxE38CFiieUbQOPWr8UaxLtUU5IxEu0DpTlBHFJsulS5F5lq3Kk7SeO1XL6W3+yLHGDvHVs8VnIOc1KMuwA5pdDaw0KTxUwT5CSvatmDS7eO38yaZUH8UrdF+nrU76holvY/uojNPjGCSAfenZ9FcUnY5p/pURAz9Kss6MORk1HIYyoK53dwadhpkLIrqVYZHpVKe0ZAGXlRV/vkUmSBSTsRUpxmtTIA568ipprlpLSOPAxDnafrVuW3WUkjgnvVK5gaGEhmBPsKtNM4pUJwv2K0fLjjn60ky74yB1xQhO8d8d6lcZHtTM5JaNGdt4znpTlyTj8aGGAQetKoPPNZ6HRfTQnQDPB/ClcnvzmkTLAnt39qdgf3cZ960VjKXkMIySAKUA5IB6UZxg9fekzjBFAkuog570gJIz3NNJ4470wuFUkjJ+tK9hodLIe3AqE4PJ6UhILDjj0oJJYDPSk3oaKOoJNInCHinvcydO/rUeccdT6U0elTp1NU2la4EktnrmpbbHmgEcVGFyOB1oPoe9GguZpl+Rwi7qp27B5wx6nOajJLDlsiliBLqBwRUpGrqXaZrRRbxkj8KG07+JG/DFPspUb5XOGrQwKaTRcuWoiOCMxwhe4609jtUk07PFVrxwLdl7n/Gnuwk1GJnElzk/wAVbCJsQL6Vm2yb7hQR06/lWpnIq5M58LHRshupCkBI6/yqGwbMZTuKbfyfdT1qG0k2zBexoS90JTtWRqD5gRTUTyoVXOcdaM46/lSSSALjvUXOnlu7le4l2RMT0FZEaFztHX+dXL2bcFj6ZqC3ZVfLHBpq9jCpZzSJTa8Da3Ttio/s8gA44q6mGOS2B61MqjBC/nSTaNJUIMjhj8qLnqetV7iQqhOOlXXTjHas+8OCijvnNG7Kl7sNCpHjbk9f50jMMdqkwQBn7tRbgCQOP1rSxwbvUaMnHY+tKBwR0FPJxwKTGRwMZ6Uir9CNsggY6VoWkW5WkYct2qrjJztrUhTZEq4qZJWN6HxMSQCMKABzTRg4ptw/zrz0p0RGBmpOskK4Hamq2GpSfam49qAJcg0hUmmKTux+dSE7cGjoBG6h6qshRmH5Vc6sc0jruFJj2KinHHanP0ORTHjKPkdKa78Yz2qb6FNaGJqb7pgoPSqfpk1Jcyhp2YZ9KjzkdKwNo7CAcYOOKcnUZHJpFzjH6VLCpD5PI9KTYzStB04rTUhR04qlbKVIJNXiDs46mvPqu7L2I4gQpLdzSqckmmlsHHang44qGKxEANzYBGfWnFV2ACnEOAB2pFPykYp3GMkPy8/nT4yGQnvUMp429qkjxkKORTtoBKuAM004Lbh2604j5sY60EAtjHHpUAGRvpDyw9qDg8Ac0jgkjFCAa4AP1pkoyoHSpGIb8aZktkY6VaQiIfcB7imgspJHOTTkUhTz+lKp6qelXcYjAgYA696WJBGCrHOadIDx6AUwMpfgc96S1QiUndjA6U0DPJ6+lOz8nufSkHB5qQ6keCU3Y/Cmq4AOTyeBUjZCnuPSoYwhmyB04rWOtwGpESxdu3SrC4G3A4qGVz/q1HHerCnCKvrSnsC8hoTEzOTz2FTDJJz2qtMmZV56H86nHzSEdKmSukwFxnJqJh+9B6g9KlHH9KhBzJtPWlEZKAFcjHXvUbdCO1OXOeeSKHAJORS6gRhNsfqaMkJ+NOBB+7TWUkhR26iqQCkZTaDio1wgODlv5VOcbCfSodyM3IOTTiIeuQyntTs5BAHJpAcjjt0p/QfWkxkSAqSv61IvHIqM4WTJFOx3/ShiFXnOaimG0ZHWpVGW5FMcYP0oW4wGVQZP1qYdFx1qHILZx171KB8x9BSkID900cFdzD6UoHXnik6g1IAfUVHGrCRj/D2p4bDGlwDmnsBG/BAz1pxUZApHQY2nkjpQfvAjimMafvgUE5koJG/FDYVgB+dOwD29KUg7cZoPXNI3UipENXIGQOaRsgA9zTsZwtI33cZzVANB5pkjEAse1PX7nTmkkH7vHrTW4MhUhWIAJNE0e+MjPNOPD5JxSgc59Ku/ULdBsHyRkZ59anUZHv61VJ8o5x161LDMJAdvSiUW9QGvwzEDJ9KRSzKM8USHYxIP1FIvLse1C2Asjpg0je/50rADn2ob7uazAamOaQjj6mhT3oPTA6jvVALznFNfoFp6DJ54ppANFxiHlee1QrgsSeo/WpDwPXNNOVI2/jVIQ5eAc96aWxk9PalK5YnNNGSSaNAsKcbunSmghmBPBpy8Kc9D0pmQflA4BphYcnyuV7mn4wuc0wHkHuO9PHIBNJgDfMCT+VRhtzkY4qXqcdc1GAAx9McUIQ5G3oTjH1oQ/utx+Y05f9VkcZpNu2PAHNAyNSxyW5FGSq7icgelOH3yPSmjOdvQUwGtjaHFQPhSSDgmrLLkEelV5wrKD6VcXqAyFsoc9u9WV5UNjpVdMADC5FWcYU4FEhD8ncKjPIYd6C3yq36VJgEEDnHrUbDRmgDcQR9aSNmEbc8cdKeyguWI6Gm5Ee3HU10X0FqidmCqv96lkVQ2cZLVGWwBnle1TptMbZPNQ9B21IYgsaDI+Y1YKuXV93AHOariISylycgdBViMiToeKUu4l5ksbBwcc0jj5jSRAruUUrk5xWfUYiD5OnFIBuIBPSlj4J70KpDfWgBm4+YTj8KeSOc0wAb8k8elPUZ4P4U2JEq/NFyKiEgYqKnTBQj0qqOGx3NTEbJs4A9qYGySMUrHB56VErZfPpTSESoBz2pnfNOOd3HekGS1AxVHBzzUcoIbcB8tTqBuzUcvK+woT1Ar5xJnNTouWJUVAMl+BViOPBJHQ9atuwEhwcVRvgTgAZNXycAetVZskkkUqbs7iaMGSJlkIIJojkaNgPWtWcbE+tZSq8jjg5J4rujLmRLRPPIrxggc1VMWFzuzV82QA5NU2GH2nnB6U4tdBaohZNuOOcUgUjnPerDn5QDz6UkSFm4PNXzaDUtCNY+AT6cUi8cEVcNq4YE9O9NZcNjHFTzpjvZEKpx2HNIOuByKncbc4HHc1COGJ68flQncQDocUpzy2KBkq3P6UmOSe/TimG460/4+OnX0rYjACKM1jW5xcgDjFbMGCvynis6/w3HHctqe+OafkA02MEgHNKPeuC2pTIdQANsx96x7FgZzjsK2r5Qbdu4ArFswFnHAyOldVJWg0S7XNVRk+9Z06sI5Qeyn+VacfJBNUboACZfRT/KlTb5irand8cg9aa7FV+UZPpTl5PNOFeEMgRGc7pPyqYAAYxR3PpQSTzTbuAdTTG+XkU/14pGU0kwQ3cSOaXJxml28CkYc027gRnK5I5zUbudhxUzp3qEqSR2q46gNjdic/pVmNtxzVfDCTnp2qaLA4FOequBIeTQRxR3pW4FZAM24xTqU5BxSMTup3AUcUDuaM0e1IBc0DpzSY496XFIBG6VAg/eHk4qYnnpxSLgDIFWnZC6iOuWFKoAzSjJpCcYxQmMcAM0jLkYPSjP500HJ5pagRpEpYqFqZBgkUgIJBHSnDqD0obAdwfwpFAI+lIeMmhTlemKkAAy2B0oxS55xQPfrQAiqAKQqDyeTQx+al4AOafmAmeDmsK+gYXRc87q2ZCD90ZNQyw7wMjJrejLkdxPUZayFbfJPSpVYs/X5azbi5EJ8tfvHtUmnTsW2vWs6bleSQXuy5M4aNF/utmpkYsPaq8qbmIU0sLbU2g/N3rOfvR1BInmxgKByaqhvnJNWlHzbiajKKVKjvURaWgCRvklgc01nIXj73epEXYm3HNOMYAOBii6uMrSTGMBcfepATEnX5T2FO8giQl+QBwamEXK4H1rS6iroLEKTNs2hfmNTCVio2jLUrhQoHvVRpJI5wq9D1pK03sBf6KcmmRkMST2pchlAzUe7tisrAWQc9elNZ9vQdaYkm4EChmxtwKVmmBIp2gL+dKcseKgViCe9P3En6UNMNyRjnoaTaGBU9BUJkZn4HFPUkDmi1gJQRnApMkEY4pBnHpSNuyCOnelbUB7HHamuR0pJOBxTDkZ4zTSAic7W6/WmRyHzAexpZIXHIH1oVOQScY7VqrWESyZ28dazJg7cgc/WtCRsfKOaj8jK49aqGmoWuYy7/O61ciQqhz6dBVlLBRJkDipja4z6mtpVo7BYoQBnkzjA9KvmLzNuTgelRyRNEo2j5jT4XYpgjB9DUSlzRuhoY1sodjnipNm1SR1+tSsA3Xp3pqLlzjvWXO2tQJlXdFiqYjEcoI//AF1dQbVPWqdzA8kiMhwRSpv3rLQNCYtjB/IUq/dOTxTTGWAB7dKaxO8KvTvSDqCKEBJ5JpYmcy8/dFRxks7FhhBU0ZDox6LVNWAfKw2FgKYM475pwBZee3akKMqHJyfWpAYjALtB3GpQoKEHg1DFGqYwck04E5ODmhpPYCeMBRnrQ3rik3gKPWlycAVGoMYUyBTgoUYzwKXPG2lAGRijmewDHjzznFMeEYUDrU5+8BSN1HrSUmG41eFxQOBzTcYfjvTgPXpTeoC/dHvTuA1N7cZpW6kk8UgEHTJNDcn3pepwegpV6A/lQA2TBHvTSgKgEdKeOeTSHhf500wAdCuaUfKMnrSYCqcUHp05pDE2kkkngUAYSnJ0zTRkKQadxDR6flQUJGR+VKR0HSnKcYB6U72Aryq3mghuvUelSIBnGKRsF8dDSxBhy3WqesQJDwnHXvTN2FORTgfmpuOCTx7VCARzwDQw4ApTyB6UoG5uarZAVZVLJj+VUtkqzZ3ckVpyLyWFZuG+0eo7Ct6MhdSN96szZyPUVNbMS7HZyx9aZIGSI7evpQjtFLGz8BvStk7qwi6hCow/SmxKyygueT0pQd2+jgCPOCe1YLcCSTKKRWdH5imRiN2Dgc1oSLliSetUo2LMy5wCeKuk9wHxSlVOQQc9KlRy7bs59hUMgCZAH0pbWMRxNjGPSqai02wC3c+aysMg9DVd0SNJEjH45qxCWSV8gbfWqk6DdKIjlmxkelXHdCexPbSYhK4y6jirtqQwyB1rMtWEbDJ+Q8ZrVhjCLheB1pVtL6FE4JI+YYNJ5mRgDpSMegB5pnJcAd+tccVzOwFq2Ajj3fxMaWWbZk5pobtVWV/MOQflr6CjDlioo4K1S2pG7HJc9TU9goLOxHIxVQleeM5/SrWnthnHriuh7HnUpc1VXL+fyoGPyo6ilVcr71megP8A4cUpyBkUL0A70/BJAFAyN3lkYFmLYGBk9BTCWqYRtnipUspnwBGxPoBT1YFPJowcVuQ+Gr6WLzWj8uPH3nNZgMKZKnf7ii9xNpOxCqluec1I3lRRE/ekPbsPekaQlcAYHtUJo0DUFbnk80jgE80o4PSkxlqBlY2WctHx7VUkjZMhlI/GttBgCo5Qr/Kw4qlKxz1KEZbaHOypknP8qYuWJx+Va1xp6uMocGstoWikIZelK6MuSUPiRMM8DPNISRknqaCcDg/hUDPnB/WqMndsecHjP14pCe35UwnPOc561GWJJA7dTSuUo9BzORwo5qEk5OKM4zjr60pIAIHNK5oooRvvYJ+tKDjbnpTV7bV5+tKRnj8qXQe2ofeIPrSqvBx/+qlIxweTTuQu0dqaFJiHj8KiPHX8aczHPXj1ppOePWjcEhB0OOlWoE5B7ioo03OB1q0owcDrQgk7Eka/MRnmnrIUb5GwDTe/PelHXFarY5nKzuiRbuXOC2ffFMdy7Zc03bkZNG3HbHrS0TK55NalqxdBI248t0q/kYwelYhxtzn6cU5pZQuA5xScbm1OtyRs0SXMm65Zs5A6flUBfDAg4NGdoIJz70xjknjp15o2MZXbNhpPlyKhJ5JqGCXfEPmyw60sz+XGT39Kza6HpKSceYoyMJZmb+HtSKgXk8tVi1g3sdw4FWmtIyOBgelVdbHMqcpLmRnAMDhTzV+xMjBi54HSmmwZeVOcdB6VatoTFDgj5u9DaaHThKMtR7Hisq7O+dh2Fakh3A1kZDSZ9epoiVXfupdxzHamQOtVcE81NO6lgo/GogAB0zVHNZ31AYzilB4zS4PU9qYFPalsN+RZhjLuB2NaEjYFV7RdiFm79DT2YEk1Ld2dlCHKiEgvLnsKkLbR70g+9ijHeo6m5KCCKUfSmLS568U9gFpwbK4703IzT1XmjqAcilAyKcB3oyKfoIhmwENY13ceTG35CtK8lCoR3rm7qUTTkDkLWE2rmkFcrkdux60dDxSseoxTd3y9PoKzVzfSw5TnIzVy2U76pqAWPU1o24IIrOpsMvxqf7vFWgflJ6VXViSAOlT4GMHqK4JBqRBQsnXJbr7U4EZOetRxgMx9j1qbAJ5PFEtxjXPzAg4pfupjHWkfAHp6UZLKCT0o6ARsuY24/CliUoMCnEDqeBTU++SvSnfQCcdMimZ3EUo3cjvQR8g9RUAAwTigH14pCPlGOtKRyM02A08fjTW5XPcU4gYGajYYYgdM00AiMCpzQgHmHP3fWlCggHtTVHP16VW4hXyVODzSbRwwXmnNnO3HSg9RjpRcB7dB2qI4Dc/lUjcke9QyOPOC0R1HcXIZsE4x3qB2CsQgOWqRztbPb0pFZRJsPJP6VqhA0oACD79CyYUE8sTTmQK5I5pVTCkt1PSldbhYSYbdpJxnpU7fKoIxk9aimOFGeo7VJ1TPUCoewdRSp3Z6imMuH3DinnkA9u9RFm347VKGPzhvY0rEq3I5ozzgj6Up5OOM4oAhAGfSlQ5JPPPWkQjPI6U9RtQnrVCQ1O4xgU0HLMvAP0oBcEBuT6CoSQ0+cZFVYZZXcOOuaevApCASD+VKeQT6VmxETcnNKW3Y7YodhznpTDxkdqpDZKGww9x1qOUfK36Gn8ZRu1I4+QmjqAsWfK9xT1OQTnio4wFQnHJ7U8DCY/vUnuKwowV46U0gCM+tOPC46Ckb7gpDAAYz+VLxyaCuVFN+81AgY5cgH8aTqDk8mkfIfjigHC57U+gxuzY2fSlBDHeOlNPGXY8elCglCB3qhXJh830pjsME5pynHHWmH7wTsalLUY4EA5J5PSjAyxxTCQJCT0FPz8vPrTsBGzHdSvjPXik2jOT0oJDZA7UwI3wzc0EYTA5NLhckdx1pM7jjFWIZL/qwPbmn267UGPzqIBsEt09KIX5Ct/FV/ZAjvHZXHpUkTgqcc1FfBXQEHkcZogUiME81SV4C6mh1Gc8Gkzke1NQ5j57U8ZCCuYrcaDhio4pgc7iB0qR2AAOeTTUxyaoBTgY55PWkGC2Pzo5PSmKNjHdx7UIQqqDk56dM01wfLzT1wwGKQDPHYU09QEjXbGATilQYBJ70rD5hg0dytAIYxGAO1NIKH3pzkgY7ioJHdcEd+9UlcCQfIOnNSIMDA6GmbsxhjgmpI2GwdhSewABhSAcmonyAuOpqT+HI4NMbDYweRQgJAcKOKXgLj1qPDA4P3afnoM0hkIXDHPUUr8fjTyqmTk80jY5PpVXFuNBAJHc9aiZRtKkYFLk5DikOWYYPB7YqloGhBE3BHWrIZgeO9QiPDMR+FWEB8sZFOQJAygHA4oPVsDn1pc5Gf0pgJx6VKAbIFLA4wtVjyG46HirMo/cjA/CqxkCnaR1rSIMYCrkkjlaIiFJPr2oAYxsnQ9qnWMGFC3ymrbSFsw/1WMc56VbiVVyw71XXIn2sflHSrHQk4rGY9xw6sR0pCuQfU96Tf0wO9PIGCO4qNgGxgAY70jLz9OlOAxn9KHAK80X1Ah+6OaepHc9elMfgA5+tCDcFB6iq6BqWQdoIB61VlP71QPxqzjkVC4zOM9qmO4mBOPf1ph4x7VIQMk9aaOhFUhjjgCl6DIoA9etOA4J9KQxcDbUMhDKAODTpGO3C9aYsRJye1CVtReQkZ5J7VYVvlzTUUgHnmk6DtQ3cQOSTzUL/ADY44qX73/66gmm8hC7YwKuPZDsVrqVQNoPzelZ3zIQ+elLcS/aZd2OOwFKiByB0A6A12QjyrUhslnuPNQAN07VWKMe9PZRzj5SO1Ot4wzfOxxTvZEve40ptUE8mpEjCkMCB3NaAVQmNvAqJrZHYED61n7S+hfL1FUNJEDnBqq6sspPY1oZVExis+WQO3qB2qYN3B7ETEYyKib7wJGakZiWB6moyfpk1uifMTOG9c9qM5Y8/lQ3Bzjp3pFHGcdTVFWshYuLlCea3IcYIFYI4lXHIzW/AFxwPrWdZ+6C3J4zwBT8c5HamJyODxT+a4HuU1ZBOA0BB444rAhyJxznrW/JhkIrnsKLkAnoe1dNHZom5sxH8qp3nKTeyn+VWosVBfqAkp7bD/KlD4i9TuBxg+tKO9NzTuhxivDAQDFIOmaUjPWgn1HFAAKQgluaUdKOtIBaSjGc0AUwE75NMIHT0pzdfejOM01sBGw+XIoiYA4pW+6eKhUjrVrYC0TSjrUYbjn8Kd71GwMUkZzSA5Y56Uo+71puOlAD+CaQHg0mdoNA5HWgBw4peP8c0AjPHSkPpUgB/QU3gkjt3px4FBAxnvTAB0pjEKc05Tlfas+9nVSR7VUI8zsDdiyJ0OQDzUo2qOTWCLgbiQOB3qZbrc3J4HWumWHdrhc1z2C1KO3NU4rpJMAfeqyOnPeueUWgHE80h44qMyYPzE+1CyKTU2AlH3hmnZzUTsFIqN5Tkc8UcrYEjsBimCdW4FQqzScYOfWlEW1lx071r7O24hk8zxkHkA0qyt5BzyxqSVSxO/wC72qLMbNgHn0qlbawGZNCPtPXk9adbyNHNk1YkiInZgOSKqYYPyPmFdlN80bCWhou5QlxyMU22BEhLHg9KlAPlZx2ojjbeN68Cua6V0VbUmVt3TpSFt3QfjShCnQd6cykpgcVjoA2KT94fanncwznHpUNvE4bL/e71ZOGcKOlErJ6AtRgOBgcmnDO454FAAySKUpkDmobAiZd3TtVGRHSfeGyD2x0rV2e9M8sMoq4VOVgMjU7QDTCnzuAewqyEIU0gQYJNTzagyJYyFGypBhsHr60qqQOaXgcipbDUQKBk0oAPOOvanHgAdaDxg+lK9wGKozkCnAZIOPrQucUoHGfWlcBrEDBIpjMoG5jwO1SMNy+9QyR7lKnoaqNuoEqkSjcOlOxnj1psUQiiCA07sKUrX0ACueDTGjBA4qTkjkUhPAPpSuDI9inp2p+0ZPpSAAknHFOPIHFNsBDjAHfvSMMmgnJGKXuM9KAGtyMEUx4w3AOKlwCxNIACelNOwETJgAA9acseDnOM9BQ2WcBqcOhzRfQAC9fSkcd6dSnr7UuoEZAVOtV1YK4x/FVpgGQg8YqsSFYse1aQAnbA7Uw7Pu8+9AfKAj86GU4460lpoA/IAwvAqORW8s460sjYAA+9T+SMd6HpqDIVBEZOMk0gBI6VLjAODUi/KuOmafNYCDDA4xzUxBCYHWnE8cCkx69qltgNJOcgU5chST1pCeDS9PpSAVehY0h5YUHgD0pTywPal1GRuWzkUNlU+XqaWQ8Uh5P0FUIVD+75pTyg7D1pOCMYp38PHbihgNZhx2Jpc5xikPJGOwoGcfTtR0AVs547Uqj5cd6CRkClPQ5pa2Ab/Bz1pTkgcfhSE5xTjzxnmlcBFIyVP4GmnpilPDe1Obg57HpVbgRP15pynd/9amMcgE04DaBg0dAGyAbt+M4oU/NntilbrgdTTQcsEPUdauOwD1OWzSHkknpSRZDPx9DSvxg1LWoXD+D+VKhHOepprAgccU7r9aN0BFJjHJ696oygqGKNz61dnBOVA/EVTlC+YUzlu/tWtMTGNnkIfvVDEGIORwOhNWN2Y2Venriq4kIi2j15PpXTTe4mWomKKXY5qVCpKbvwqhFvSNyBx296kjctECTgjuRSlC+o7lxHO5t3Iqo6/ZwXHBPf0q2SVwOmepqMgs64GQOtZJ2Ymhiu7R5x16Ggv8oP50+Y+XExA+X19KijBNsrHlu1XurgRyysJFwpPrzS3SqN7pgSnFJIxgtt+3c3bn3qK4UvIrKe3INaxWxPkFvEURUk4DGtCCUGQoSfl6VlxHZcKjZIWtKK3LFioIyeTVzTl8xrRFmQ8k45/nU1tAz5dhgds1JDBgZPP1qSWTykzjJ7Cqw+G5XdmVSokiC5YRgxJjJ6n0qk+SM9hSnJYsfmbOSaDngfxfzr1Iqx4lWp7SWo1iwOO/pToH8u5pCNwH86hP3s9R61TJi2nzG6owc09cYNQ2z+ZCGqwozWTPUi01ccowATW3oOgyaw0hWQRpGBuOM1lW8L3E6RxrlmOAK6c6hF4bt209ELXRIMuD0HUUWYSkkjS07wfBHMslxdBlU8pt603WPFllZrJbabGksv3PNA4T/GuX1PX73UVaNn8q3YYMaf49ayc4HHbtVci3lqZ2lLfRf1/X+RdvNXvby3NvNcFom+8MYJ/Gs88DFITzSE8UNmiilqBOBSHpSkZpQvtUjE205VGacBjNAHzUwHE1CT+VStwCe9QnPFAACOneoJdoJJxipWOB7VSnck4FJ6oZHIscn8OKrNajBAapge9KTzzS1Bwi90U2tXyPmzSfZmHNXD0560gGPxouxeyiiobV+mRQLTtnmreO1Kq0XY/ZQKyWgDZ3cn2qRbRBjmrAH/AOulI4obHyRtsVTaKMbWwKjktmUHBq2W460zPOMUXYvZRM54XUklcfjUW3PGOR71qleakhshKSzLgDvQn3IdJLYpRKFGR1NPHBGB+FWmtCFynzj06VWZAHx0rWNjinGSepIAAc/nSgYPTn60AbsY6U4ZU4PSqRk1qhjDA54IppAAIzz3pznHIPFRlumB/wDXouJJCtgEHP8A9am7iD1puTu60vVfU0h2uI3T6U1lG04PSndTgDOe1CDC+3rSNEkMAK4NO9+/rTivtTgu3ORnFO5D9Sa3uBCTkZDVcSeJu+D9KzcAnA4FKeVyPyo5UaxrSirI1wwI4NGcc1jHkHJo+0SIPlf9KXKaxxHdGjcy7YyxrILbVxnBFLJM0uN5z6VF978e1K1kTOfO9B0Y4yD+FO7nPagAA89PSnbR25qkQ7sTgkcYzT4YvNPP3e9LDCZMgHj1q4ECKFUYUVLkka0qXNZvYa79h2pUT5Nzd+lRnk5qyAGTOenaoOzoR47UYAqSMHdgdqRh1PFFx9BuaQkdM06IfPgtj3pDHy2DwO9ACYJq0owgpiJz0qbGBTQmMPTNROwWpJDxVC5lwhqW7AtTP1C5IzjmsXOGPf1NWLl/Ok4PC1VIw/17VzN8zudEVZDyc8EcGmAkt/KlJwQT0/lSL1PvQVdXJrdccVqwgYBrPtl+fbitaKMDHauatLUpDyMNwMgHmp1Ib5h6U0pt4JyKdEPkJH4GuWWw2QcJKoA61OB+VQy5aROM4qQcKMdfSh6q4C4AHPXtTSAVBz0pdvQetBQluvFK4AQTxTAMFVHrTxzkkciolP7zkdOlNATt8p6801+CPShATmnH52BpdQQHl/Wg0A4fOKQZz160gEUAnFMf5QxHJpxb58ZqJuWYjkdqpIQi9Cd3NIxwq46+tHRwO3YUkxPmAY4rRbgTf8s+etRscIPXtmnFQCaGQFQDUdQHICU61CyqGznmpgOdoPFRyJ869/amnqMHBLYxxUCAeaeMc1OWJyB+dNC7pcZ+Yd8VcXYW5JgMG3cCmQ5kcsfunpmgglSAeKIHDYQdhx7UdNBCTZYMPSnIcqPcUTKTGcHFLCu2Neck96Sa5R7DhnaV60wjoT1FS7dp+tRSH5tnc96hajFIOdxP0oD5TOOnamnJB5o3EKO9VYNgGWQk0udygDpTSSGIHSlGdox1oFYF6qeKacKzADOad5YyG6MO9I7KHIB+btT6h0HBsgHGMUrE7wFHBpgJZOvPcVJ3GOuKlqzGMbbkr6VE2PvZ5PapJANpbGaiXhASKpbAPBIVealcZIHaoMnYM8c1ZYYAI9KUhERb59oHGaecAimKf3+O9Sg5JbsO1DGO6vjHFMyDJil6gkZpq43Ek5qQF38/ShTjPHNIR1alAGC3rQAxjtByfxpATg0kxKrnHBpVH7vOcmqtpcRHMWKY7UW7Eox6Glc71KjqKSOMI2OlVpYNCVfli9TQcbORSIpDENz6U4ruJB6VPUaRH8u0Gn53KOOKaQB8vSnsCEAA6U2wGPuGBSABScDmlYcZP5U1T3oWwCAEMzdKAM9OBQp+dgTSIApPv3qgGOyiNsmoVwI1z+FSPny2AHJ7UgUhMbeRWi2ERzkYZR09aSE5hyRzSS5A2twR1p0CgLsPFV0Ey+n+rBHWlblR7UicDHalBwOa5nuUMIDHpxSAkZOKXOOaCfkPvTEKOPm9ahY7X5NStnZ16VFKAMMe9OIx6kKMdaRMbCfWkU5GfSlwWAPSmIFBK89qByD6incgcU0cbjmgBHyYz61GfujPenEkK1RtzHnvVJCY4ZCtnpTk+5QD8uR1pyngYFJjFb7o9TUT4wQB81S4PUd6Y3JzmkgFAyoU0n3uo6UKRuz6U4HA68CmAsnO1h27Ux8N8pHXtQxPlnBpOgBNCQbkeVPyjp3FRyEx7euKfwCTjmmuWO0YzWi3DQYWy27PWp1fEdVekntUytlWHanJASHBjIpgGGVCetLnaFPrSM4Eg4/GpAjnB3Eqx61E8oxgjJFWJx8pAbBzxUI+QncMse1aR2E/IbIu5gR19KnGHCZOah8s72Y8A9qltwRFn+L0olsA99pbDD6NUmcHb2FIMHaCOe1EXLsT0rN7DQ8Akk/lSnpmmqCG4PymnN8wxnioYAxAIB/KnOMgLikJ5FKp5OaBEE4/d4A5psQOQelSuCwYflUcQ2gnuKtPQbSLWefSoT9/dTlOST27UuByajYe4wDINJg4zjkdqeM5NBX5+/Ip3AFHOaGPJWpNvHFCKAxJ5FTcGNjTn5hxS4HTHJqUgA5FQuewpJ3DoNY84z0prEk8UEEd6jmlWFCWrRdkIqyXYhu8H7nqKp3t2JjtX7uagmzJIeODUv2PcgG/H4V2RhGNmyI3aIDtJG36fSrUUBUB1x9KprHiUjdzV4QhYwDN+FaS0FqRuU8zcx3DrUIlJLbCVA6CpTGiS5B3e1VnfY5GOfSiKHuX1mKx4JzntSrdsozWesjAAc4JpSzenXpUummK7RYmuWfgcetQnPU5JNAXcRu7UpGM4PFNJLRCuKCQASM1Gc5AzUjD92AMgng4qNyeAMYzVIaa0AkHoT0703ccH2FLgfj0pAuOKYxr8FBwMVuW8oQDI4IrDccqf0rXt1LJH9OlTUtyDXRl6MjO4VMDyagTCk+lSqQc968+RXQkXkHniufnJS6bjBJroF+7iufvlAvGNb0PiZN7GpAc4qHUP9Qxx/AafanMYI5yKbqJ/wBGfH90/wAqS/iIrU7XPrTC+HwPzpzEAc9ajCnJJrxUhkwJIppNImVXBoYHGPWlYAU55p2M0yMEcVIRxih7gImQKQH5jinfyoA9aAGkfnTcHnjvTiepqMNub600mA185+tQN8pxjipiSX+lJgkHcK0joAxZMgDNTAk9KqmN4wWAzjpUsLlW2kc02uqAsDimlvlIzTuvemN1GazSAieUghRyD3qZGynH4moHUqc9AaWBjnHaraVgsWhyaUZz70ict9KdnPPpWTAD1pCM8d6XB60E9c0gGudqHFZN+DgAnk9BWsy7gazbiBpGGeq9K3oNKWomUHg2oeeR1FVlBjyM8GtkWjZGRwfWmvYAtjqTXa60FpcVjOikMc6lTzW+r7gB0wOazk05vODYxjqa1QAoyRXPXqQktBrYpN5pOScY6VNGhA5H41NgNmlXlOKwc9NA2IPLIzUohGASKkIOMinA8YqHNjIvLG09hTDDngHirAOB0pB0PHFHPITK1xE8kW0GmW1osWXPLHrVw8LTc8DirjUklZDtqVzACf61XNq3mgjmtDaAuD3oCYzTjVaCxAikEZ4X0qQkZxml6E+p7UuM8mobu7sBcAY70AkjFGQScDk05cKp9akCMg+ZuB4oQjdnNGV6E80jcR1XQBRxnHWng7kB71FATyCePSpsDaeaJaAAG3GO/Wlbjp19KTuAOlKB15qQEb270EZAHpRjqD0puGC+9ADhgLwKQLxnuacF4I7jpQOetAAvAJpO5z0pep9qBwKQCck4pf4cetIuRzQ5OwkU7XEKAc4Jpr9M4yewpyEhQaVuRmkMTPFA5GfSmjPmEU4ZAxQAZzmmn5iQOlL0INKMYIHBosAzgdKcPugU04JxQTjAFMA64p2SeB0qPcATzUo4X60MAz81RgndTx60g7e9IBrDceOtKB70dCaOnO6q1BC9WwelIDkHHelUkNk96AcHjrSACdpwe9QFOT2zUx5YZpCOc007AMCLGoUc0pUbqXbxn9aUYJznj1ptgRyDJXHFPztCgU1gSflNOVTuB9O9PdAIFfeR27U/g4ApoGXLY6U5eCalgOfpgUhPA4o6nk0CkAwgd6cwO0ZoI5JNDZIzjimAHr0oZuTijHyk55puDnGaAFK5Q0mMKKd/DjpSMQOBQmADnmlBIyKYudvoKcDnp+VACkYIoA5z60Y3UmSG6UgAYz6YpR096QjHFIzbfenYAPBGafgFiTSEDdjFLk8EnoaNAEPPal654pqjB69O9KxAPHBoArzg7QOnvToZN0fzDmllbLqAOBSJ8rFSKtfDYeg+QnjA57moyq+aH796fIDjrwaZtAbPfvTT6i6j1Yb9ppXPOM4FQEncD3qUsTGeKloBCTu56CnE5Y7TgUEErkdabuOAB1p3AHyCMDPvVGZR525Tz3FaDN8oFUpYxKGAbDVpTeuoEAZgrkDK9qrOrZTAzjrVqJB9wMfl9qrnzDLKG4XFdMH72hLJHnTyldRle5q00IaBQByKihjVLaNSOvWrc4yEAOD6VNRpOyH01IpsiAEDcVqM+YIt5wCe1SpI/mAEZBFNlLFAWOGHWs9tBPUi+8ShqZF8uAozVXaVTsYVOuHUknrVyVo6hoyl5j5bJ6/dGKluZWEUZVfnPBpZAGKhR83UCrlpC0qiWdMf3VrWCU7WRLditbWLTSLLMCFHQetbMcWecfhTooSW3N9KsKoAHFd8YKxjKYwrsGTwO9ZcsjTSFscfw+1WL243sI1Pyj73vVTI24zx9K6YQseZiK13yoTuc9aTv7ilJyee1KAcYzWhyIa+Mbh931qLgDnmpypPyk4FRNGVUjp7VLLTNDTJAcxmtPAAwDWDBIYplJ4FdFAiyyAuwWMcsx7Coe9zvoyvG3Y2tJeLSNPl1iZPMZOIYs4LnpWDNcTXNzLczNullOWOPyFOu7x7yYMQUjUbY48/dH/16r54p7F2u+b+vX/h/wDMUnuaaT79aD06008gZpFgaBR70uBikAD2qQcUwCnDimgFJ5xSE00nHWkB6+lACs+R1qEmnt161A7gA0DSEkfk1Ufk81HdX8MQxu3H0FZ8uoyEkKoUeprOUkkbxozlsjQzj6Um9B1YfnWM9xK+Qzn3pjMduO31qPaJuxusK+rNsyx7fvr+dAlU9GX86xWOR14puMdz+dHtEV9VXc3lwe9S8ZrngzqQQx9ualS7nTo549aFUTJeElbRm5kYx3prNzWampuow6Z9SO9TR3kTkZO0n1q+ZPYylRmt0WevNKM0qYI4OQfSrUMAxuPSmYsjjgLHcenpV5QEj20gXpxQTj8aroS3crEmOUjPB7U7asvDAGklXJzio1cpIKXUGrofJYYQGNse2KqtG8Qw4wPWtjh48jpUT8jBq1JnNLDxbutDELZOM9aZgYOeDWlJZxscjiq0lrKoGBketUpXMZUJR6XKwHH3v0o6nJPNSeWQcHinhTngY9KaMdRsUO9gp71JeQrbztEjBguOQKWIgNn06e1NlBYnA5PvTsJ3toQknPI+tGD1FOOGY4/EUmcANQFrik/l3prHGe1LuGRtP6VCxGc+vpRcqwpYYNRMxPPWntyTUf8AGDjipZSjrYb1yc9KeACMmhQTjA69atR2rtyRgUmzSMG9isAc4x+FW4rYvgvwPSp0hSPkDn1qXOBz+NJyOiFFLcbgL0FMPNPbJGBTTkKR+tTc2sRHg4pytjjNBx0FIc96Ciyhw3pUR6n9KI2yMU+TBT3oAYhwQTUqL601Vzg1MgosIcPl7c0rEAYzQTgZNRuSAeaYiKVwKxb+fAY56dK0bqXap9fWueu5C8uBztrGq9LGkFdlcYPXim4zjB6Ucbjk0jNz14rKxuCt+lPiTLkj8KjXGOatRjaRgVMnYa11LEK/MP51qQ+hHbiqMCgkEmr8a45rjqMoeQSDk8jpSREkHOMEdKVyEUkmkiGVaslogsMmYLtx1p5YEgio7jgD370+LBiPPXkU7LluAkrfMNvalyW4oUfLkil+XYOeaV9LBoAYrkY4qE4Rvc1MOTz6UyTBAJFCeoEgJVenWgnYAMcetJnKqT0pGO4j2pdQHE4YjOPShQB9fem7S3z54peW+YmgAJ5B6k9KaFJk5FOkOHG2msD5maaFsRM2XwKWQ5JI60xhtkBB5p7c5NWAKSfl6kU4nkgY6Ukfyrk8nFDkq2QOtLqA6P73SmvkduTTlPzA9qZKeQewpLcBseQaGbbLu9fSmqpD56ZoQk5P61otwYTgY2Z4brTodoO1T071HcMETnk1HC2xVI6+lO3uhsy243HApsClOM09uBnrmoooyhJzyahbWGTquetRyYHzVIM5H86bIpIwKnqBEQXQbeKCGApwPykZ5ppyfl7HrVgOVgRmmx520vA47mkAAXaDikANgSgk8U5xk5XqKQqPXk9KHXamcdaABTmQYHFLyzBs0yJRHyO9PXKjnvQ/IW4mQQcVF1jwD9RT0YGQ8YxTFJ3nK4FUgHDlNp6ipmJ+UCooxjcMYzTlOEPPSpYIRTiUn1608HJwPyqDcqrz1qWPaQSOlOSDqSDhRjtUYAWQn+9UiknPFRSSASBf0qVcb0HHPIoYYXkUrc8Dn0pzEcAjmkIZ/Btz1pq8tinyHawIGfakwQc4qhkaR5nYnvUpGWGaF4JOeRThjlqTYhuAhNJyGyRwKd/CfejoozSGNYDfnFOc4BFIBjk846UrE5FAeo0gEA96jyFjIHWnucZxUJPHAq0gBmPHHWlb7vHFNkyWGOlOI53HpVCI5AQgcGpozlA36VA54UdjU68AH+E9aJbAVJVzIcnFNTIcY7damuE+ViBUX8KAVpF6CZejOY8/hQMbMd6ZGdqj3p+fm5GDWLWpQ1uSAKANxANLk7gSMCkJPJoAD97HbvUUg3fKPu1ISBx1zTHBUjb+NNCI0fGcD5akIygBPNRthWOB1608g5BFW+4hxIAA7Goy207M809gAwFMlIMmMUkMCQJMetNYnbjuKCuXznkUcmTPaqAkQYiBPNJGpBbd0PSlT5hn1pM4dVPPuKnuArEUMODtoPdumKG+6PU0ARq2GCnjNPJVX2VGMedyadt3sSe1UwFY/KABignCqT1oJLJkdqXOeCOD1pAMZScEfyqMdMnoKmBxIRj5ajIJDDpimmBXGM5UcinofbApHwMAfnQDhST0FaPUB2DuB6j1pHYFsD/9dDghAAOKSRwo3Y60biJXQEbeSe1RSAFwG49KlV8rvIye1I6ApvYVKdhjcYwD1HaowQEDHio3P78sM9Rk1MrK0hQjIq7WQbk/zEgk5NNhkLErj60jI3nDB+UCiNwXOOB3rO2gE4AjYKOlC9xmo9xwpyc9qkBxyT1qWhdNBe9HRsUdQCaM5570hic4I71EhIyMcU8/e9DikUDZk5poB/QcdDRnjjpTTwgpMkLjuaLASg8nNKTz7UxR3zTiARUjJF+bJ9acRhRjj1qNSFFIz7jgCptcBzP6VExGDjnFIxxjPFRmZA3XjvVqPYRKoOMk5rNnmV5pI3HA+6fepJ73J2IfqfSqoj8zkHrXRTi0+aQr30GKgTJZsNT96BsE804RLgZ5IpsQC5Rug6e9a3uCuVrgANlSOahXEvV8Yq7IqhDuHAHSqAbco+XBrWDuhW1uWXUrCGDAiqr84JAzS43cbuvalwuAOtUlYSD04+hpygseaUE4+7j+tOiVvN5HtQ3oImWMqAxyPamNGAcbuPerZIKYB5qBkxmsYyuNxSIvM+XHYdqjcjOevvUhXAOPu4qIkgjA7VogtYM/KQKaBx3ODmnZAzxz603IPOO1UhiOQOeccYrWsnxGh4xWU+AMEcHoa0rPBiUEYx0pTScdSTQOA3X9KdH0ZsUwAb+n41MMjivPZoyRe5zWFqZIuic4963Vxk1jaqNsoYDrWlB++Is2f+rXjgjrRfkfY3z12mmWT7ol9adet/ok3+6f5VWvtA6HYOgY08cCmbDvznpUnVa8R7FC9uaXPNHX6UZx9akNA79KB170H1oI5oAU8jjpSLz3pR15puOcdqEAh6HFRKoLYqYrTQAMDGKpAM25YZpyjANP2gNx2oAxRcCNwSpAqMQBWJBqfHNL0OarmaFYj25o2nFSED/Cg8dam4yMx5xninCMDOBjNPxnHP4UEdqOZgGMY9aOg+tKeTgnmkOelSA4gHGKa+SODTjyMimseM4o6gIBg9cUhTLEijduA4p3XmnqgGgfkKGTgbTg07vzS9cDFNtgCjjNNKA5zTicN0pu7nrxS66ALwMjFIB60dD9elLjk8/SkAjHIwOnrSg9c0jthcClB4xTAONtGT1pGAAApjS7cDHWmlcCKdZG2+W2PWnW4ZQfMOaftBODTsHPXitFKysC3D0pQ3JzRjGTQR8uegrIA2/NuxzQ2eRQ2dvFB+Zc460AIV+U4oXBzQcdBSFdq8UwIZtuMgHOaHYLFkj6e9LKD5RwM00YaMZFaK1gFiZixboPQ1ZyNuRVJ0bZheKngkVkUA5I60SjdXQEv8eKcOMjH0qNn2tUi881i0Ajruwe9LgDnvS+9NPUCi4C9TSjqee1NJ+ajknOelCAXOM0D5ePWjqR60Hk57CgAHI57UpOT7d6D3HSmY4oAdk5zS9RyOaBwOaXqRSuAjc4Pek6NSjqMikA5OaADGTzxRxnA60meo6mjG3AHWmAjcfWkJB69ac4II/nSHnJIxmmnoAwKM7qeO3pR90il4x9aGwDj86RuD7UqjgmjNADT92mHhRx1p4Hf0qKTpknpTiBIj5we1OyBnJqnHLujBUd+tWdvy5NXKDi9QEVuc/pTuSOaZnpt59aeCM4z0qWgFOcYHXvTVGQQOaPU5+tIvAxQAvUnjilU4GAaQkbT6+lCA7ulAyOWQx4UdSanGM4xzULRK8oY/w1NjJP6UO1hAR83tRzuzQvOT60Zy3vUgI3UCjOSfahuTj9aTpgdqaAd1HHSgck0E4A96MYGaQCDjJag/MR/OlPFIR0PegBmGLYz8opw4GTQ2TnIo5JxVMBQPlJ6Uig4z3p3qKCCFBNIBgPK0oIyffpTNwVgKdjPApgO/jyRTjjAyajZiWUAdepp+A2KTAYTgnufSnAlsH0pG++Md6UHj8aOgDG+6cnpQuMNmnEDdTcgOQDwaa2AVlynJ5pgUCQtUn8WDUbYxnpmqiwE25BOMUq58sA9aCMIOacT8nHAoYDWViQo4FKrKpxnkdaYQBJuPamucS7qaV9AJSMk+mKrMqBWY8kVZbP59qrMudwY8HrRHQGUxOAGKL8tRuzBvmHyntnrSoQZjGozGOtLLsa42n7vb2rtSSaJWpMil4EJOAamyQ4bHy+uarI7LGIsZ4qwMyBY9uFxWc431WwxgjZLppd3B7VKy8Nk/QUxAjMGB+5xintjfnuelS/ICrCyzImVzjNTbQIyyj2ptsqCRxwGz0q3HAZnCjoOuKtpylZE30GW9qJCHIP09a1o4ABk06KEIB8vSpgK9KlSUUc86l9hAPaq17P5ShQ2GbpxViR1jUuxwO9Y8khlkZyMM3b0rpirs469XkjpuyMAZ6/jilOM5NKSfTGKbjPbJ+tbHmvTYATyBSjbj607aVzzgGkHBPy/hSC2uoo54UcetHIA4z+NGCvQfhSr8ynPOaNEV+YwooyQMAdK10ffAprJZcck/WtGzbNqvtn+dTLY6cM2pNEp96TsacwwaaajQ7hDzTadjrSY74oGHX6GngU0D1pw9qBC/jQeKO3X8KTP5UAN7Cm5xk0pPtVG+vBbIAvLnoPSkVCLm7IddXkduuWOW7CsOe7lnbJbaPQUyRmd2LHJ9TTMZJxWEptvQ9Olh4w1erGBc8AVKLOZxnbge/etC1tRFHvYAseme1V59QxJtiAYDqxo5V1FKtJy5YFKSJ4mAcY96ZjODVmadZoFyMOG6VWJqJK0jWDbWoY/wAmmkelOOeR+tFJFjQCSKUrgdaUc4zS9c89aEwGYyenSnKmTT1jJIqYKAOKVuxSQQySwtuVz9DWnDqx4WWP/gQrKPfnFKPzq1UaepM6MJ7o6SK6hl+44z3FPfoR61zS5XoSB9asx388OMncuOhrVVEzkngn9hmu+CvvVY5LZPbpTYb6OYYzg/3TUhUirORxcHZl21k3RkUN1NVIXMcnfHpV7b+fpTM2tSPHPHWnKuRT1Xn3p2MU7CGGFXHIzUT2kbAgDBNWCePTNNJo2IcYy3RSawBzh/0qFrWXdx+daOabT5mR7CBlm1lAzs59M00xSd14FaueaYTnJ9aOZkrDxWzMjypBk7D+dJ5EpBGz9a1G96hLE/SjmH9Xj3KQtXbBY05bJQPmOT6Va7+9H0NTdmqowXQZHGijgVKDzzTOKO/t60GiJMg9KTqc0uAqrjvSE+lABkZ56U4hNvJzTRz0o5zQwuRsoyeKYeBUjZJ4qNjikPcTcVORU4G5Qw6VAg3ZAHNXYoSFwaauAIvoKkUACl+6MD86aTT6EgT1qvK2AR3qVjVK4fCmk2NalC9mChiTwKxByctyx71bv5MkR9cnJqpyR6Vyyd3c3hETpyaaw2/N2pxG7jJGO9NZdw5pI09BVG5s9M9auxqAQeKqxAkAc5rRjjJGOmKzqOwWJIiB26VeiUeX/Kq8SDO31q5Go2nPauOpK5Qjr+7IPpUcW4kAfdFT9VOajUnjaevUYqE9BDZkyuT6VHbLsTkcnrViVtqkgZxUYbJ5FNN2sMU4JyOB3pvHYZx3pXBLjHQ0zOGCj160JASElxuPGe1MYZcZHAp+cEHio5Mg8jg0LcCQkZx2oU5UZ5pqHK5/KnBSFyetJgGcY9DQ2M/L2pNi8qefSkRNrNjpQA+TkDFMk5fGaeOuT0NNBw/0oWgELqBNk0ZyDjoaa4/eZPPpTkycYrToAqAFVXuKVwMhD0pI/vZp2c84+tS9xDQfnGB7UsuOCelCk+ZihyT1o6jGE4Kj1FRsXKgg8Z5FSsNwyKY/Dn0HariIaAZpgH4Vf1qNIczOSMLnirBKiQDv2pJTl938I7VTb2ESNnp0UVFGpM7NnipA5eIYHNRM4SRVJ5PSkk1oO5Yxls0rgluKAMKBQ59KyGREfPnOMU1+HxT+ApOeaYAME+9WgEJHD54peoyByTzSS7drAjPoKVSRHuA+b0qr6CCQOHAUU9zkYbkUEFxu/iozuUj0qb6BuMXYWwOBSp945PSl3BlGTz6+tIMYI70AhGUbCVGGphDMR2qbrHk8GogCuWzwKaYx4UEmmgZHB4p65WMnrTYwqpuHekBXaQb/AFBqyOSCPu+lVXKxv83RqsowaP8AlVy2EPVtwY9CaGUEHIyaVfuH1pygbcd6yegxnVcg0KARkdqFwCRT+4GKLgNJ+Use1IrFl+bpTsYQ5NRq6qSopoQ8c9Pwoxn5R2pkZGeO9OGc/TrSAVjkgfpSEHccnNA6gjvRgFgRTGBxnpxSnqCfWg8OaQjpmkAxjnPvUZAJIFSHnjNIg+cirWgiLaWx7dqUcjngLS9Cc80meCp71VwImHmAAdu9WYx+7296rB/n44xVsMNmaJ3sAjKpXaRVeVBsHFWR3JH40yUYAIFTF2YDIjlMdx29KmYcZ71XHEo461YB/MUSBDWUnkU11J5BpzElD2pQDtx3ouAx+O/NNc/KD69Kcw9OtRNgDB7d6aGhSp2DHbrR2GOlCnEQPUGlDcY9ulUIVgNoPpTCd7AjtT1AMfNNx8obGKSAZ94EmgYyUFO2jafQmmK3zMwXmrEySIgAqKEUeZk9aSMAucDk9aUEruA61I7AeeP0pW6jGAKaemaUkBcUAyEKN/mfw1IpHOO9ARdpHajaM5pt3AWMEKc0gXABPXNEZO72pWBwfSkA0ja+OxpjglvQ/wA6lOSvXBpjDJB61SYNMryjawwOlPABJ4/+tTXwR7CnRud/PQ1fQQsgJQhRz6VC6lkUDtVhQN+O5pu0gMoNJOwCI4VgCeBUrsrEE8ioGXanA601shUBYe9PluMa+eVx35p8BG4nFI4ZnI3Z4pbYHO0dqb+ERK/DYAxQhSQ8ClyepGQKQfIAQOtR0GO2qYgAehpTzg9MUpUFOMDNOQAxdc1IC546ZFIeUoOSME0LknjpSAQ4ZwaAc8Y4p4xndTRy3PFAxCPm9hSKDgnqaeQSQBTgMcDtRcCPBJGR9acWVD24qK5uFgHqxqhLNIwGRwauMHIV0Tz3oY/L2qP7Yyp935j05qswOeox6VG7EsBnOPwroVONrCvrqTGZ5hjPNDRlY+v0OaIU3wEgcmq8SyxFlbOKtLoidicxBxkHB7mlUmPtwO1Ee4AjoPSlOScVPkXvqKr7sjH45pFiAzu609Bs6UOdx4NTfsPcZLGrJ6EdSKoMT5hUDOO9W3jYo3zdaqHOTgcj1ramZv0JGMeB+757nNNdo1IKcUpimmCkjgelSNa5QZqrpbiGBw6FsHdUkeGIOMYqEgB9oORUsYZRkH5aTWhS1JimDuFRM25iDT3POQ36VEQzHg8DpUR8wGHK5yPpUbc5JPJp7DBz26UxlB46HFaqxLeo0DIyBzSbeNvqetPU5B5pB8hPHy+tO4CNt2DJOK0LA4i96zzhgpq5YEFSucH160StbUUrmqp/Gpk781CAB361KGG7Ga82RsyRDk4/OszV0yEIHrk1pp1HrVDVhmIGrpP30Q9UV9PPy/Sp70f6HLx/Car6eeWx/wDrqzeDdaS98Ia2f8QFsdoc5pTkdKBjIOaRh83FeAWL2o7/AM6TvS96A2DvSdaCeaXtTAAaBQABTQBgigB3FIR3xRmjdgYNIBcZGaUY5GOaQHFBKhfc0AN5zg0vVvwoGAcmjjk9aYCjpg0o44PNJ3GRSn2NLoAmKQ8AUucUh+8PShAISPMA707OaTA8zOOe9KRgcUwF5A6UnbBpxPHWmsM8ikAgx930oPFGcfjQ4+XrQHQbuAbg808Mcc1Ao6E1IDyBVSiA58gE0yJPl+Y5qZxkYNMPT2pJ6ABO5gAKXoeelRu+BkUIxcDinbS4D8A9qT+LA6UbvzpzHHSkHQjZju5H40Bg+MCiUEpnqfSiFSVJPU1fS4D9vHHWmqBnnrTiQTjPSoZGf/lmN2etStQ2JW56dKUjPA7U1OE560pJ6gcGi2tgAj5TjqBxSKCIPenv90mmqu0Be1CYCrgjmjuaT+IfzoPzE57UWAZISsRIpF4jIApHy0i4Hy96c3Tiq1sBGflUDNOTAkA9aY2d4FLFIrEHABqrOwFkLuJFL1PFIDhR796cOFNYgISMj1NNIJbJ7UvOKceaBDRkmlHcYpv8ORTsELQMBw1BHOKXqQBQRlcjrSAQnAzS4wBSHB6UZyopgGQFOaVeRmoid0ZzwRT4+FGT2pgKDxnrSsMLmkzuPsKUnJ296QCDCnPWlOcZ9aa5O7FO/hzQAjDjNNxyKeOF5PJppGVoARsZxSfe6dKQqduKRCenYd6YEgORxTSdq4NIhG0nNEnY0W1DoKCGGKi253Bu/SpOoABqMpmQmqVkBEqEA56VMDlM5pjAk4xx607btQAHpVSdwE6HgU5Vwc5yaaq4XPf1p8Y43Gk9ABj2zQOM+tGRuNC/MTS6AB+5x1NCkg0pGSD2FIPmP407gIPv7c4OKf8ANxjvTSBvDVIT0Pek7AG7oaRThs+tKRzimNy3XFJAOJ44pvIAJokYbQO/p601nwvPWiwXHn5jx0p2ckCo423DFSg4GPyoasAlJ/KlPSg9AAKQCenNIGwxNIWAIGM0vVj7U0AoGCSe/SkXOSD0oXPPHNBOX9hQG5EAMkkd6lB2kDHJpDnfjFJgl/51W6DYefujPWheB+HNA+b8KAcqT3FSw2E+6ORSqBg5pshJC0qjCfyp9LgDD34pjrnlakHTB6VGTgHPQUIB6fd96Y+AOelCAeYcNTmwUKsOafUBhOF3evaiP5huJ/ClOBwewpF570xiE+1EgOVYDJpfvMVx0FDgk8cD1pq4g/i571CQiuQRk09FdW+Ztx9aYCGlYHn0qtmBRKubjaSAo6gCidSu4qMsf0pZGUXL5baB+tLJKEYFuRXSr2RPQWPKgEnJHapGl2wqU6noDUMcpNySVwuOaSeRWkTbyD0FJxbauMnZY44sAkBzzxUhIVMgZYCkbh4sjk9qka382bap/Cpd27BcitbdpH2qDyck+lbcMKxqAB0/WktoPIjA6nuasbTXp0aKirs5qk76C9KDgHNJgj6VWvbgxIEU/O3T2roSMJSUVcqXdx5srIDhE6n1qvtIJJ7UoUKuBwO9IAUGMfMa2ijy6knJ3YDgcHk9DShBzkceuaX3Azn9aB93rx64qifUTdk/Nx/WkDDG3t60ozkk9u1BbP1oJvfqGc5I+9SE5H86XaCBjg9qQ8rnOKExtOwhXqMdK0rIYtgO4rOH8QBwOK07NcWq8Y6/zqZPQ6cKvfuiQ9Sc803680403FZneJycGjkUHngUvYUhi0oFID6/lS5HegQZx9KYT+lKxzTCePajUNRkjhELnoBXNTyme4aRiSM4HsK2NUlK2bAE5bjisTOB7VlOV3ZHoYOFk5iZ/wAmprKHzrhcj5ByagPNatjF5duWxy1RDc6K0+SNyW45ibDAZ457VWSzhSMDG73NRXDtPeLCpO0dauTuIoS390Vo7M4UpRirPcx/J3XDxxqSd1WhpwC8v82O1Q28dy4Yx/LuPLGrEcVzG43vvU9alRXU6pyktE0U54DCeuVzwaYkLytlEJ/CtC5jEgVD3YVeCCKE7V6DoKOUhV2o+ZhtDIn3kIB9aEj4x2rasJheKyyRAFeoqPUbNIIvNj4GeRUyhZXNoVve5JqzM4cDFNwccUo7c8U9IZZBlEJA7gVOp0t2GfjRg5p20qSGXB9xQP8AJpWS3AXHp0pkpCj2qR2Cr7+tVHbJyTQrMTYgcpIGU4NbVrdrMm1jhh+tYmMt71KrNGQQeR6VcZOPoYVKSqHT2iAzozjcuenrWlqSw/aBLbgrG/VD/D+NZmiTpcj96wUr1JravpbQ2cYhRhcEndk8AVvpa6PLqRcHZmeOnXijIAOKaTmj+VUSISSfagilHHWkI4oFYQ0hp2aYeKBiE1GTgZoJx1qMksPagEI7bqYaceBTWNSUNJwfekzgUH1o7UBYUc8UoHIpOlO6DNAATz7UdKUDHWlHPeqEIPWkPtSgjOAKazYBNK4DXYdBUbDPNGckfWnqpbFKw7k1hGc7iOnSpycEinQLsT2P6Uw88/rVE9Qz6U0nOaCMcUwk9qQWGu3BNZl1Lkn0q7PJtU1iX837raPvHpWc3oaRM+RzJKzHofWk684zim5PTtmlBG3rWBuKck+vrTTz3NKGPbinxjnt+NLYepKiZwTk4rQhB2896rxj061eRcIAe1c9RlIVeGAFWwpUeuetQAE89DUqE556VzSuA5zkYFMDYYAflTiw8wjHNRgZkzUpaCHP80JHbPNQk/MmOnOasuBsIHWoAoIUenenFjFdtoAJwfWmykL82OakC7+GHSopMHcT0qluIfGVZAR3okxt5/KmwMFBApGO4MBRbUdtCRSAAB0NPI2n2pinpjsKkkJI/mamW4EWSBuA6d/WlOMZB60EhnK9qcuAOSOKbCwp+Y+1MfpntThkrgcUjAFcHoKSYELAby3Y0ifMAM4pzYMRNRrgJnNaIQ6M4cgU9FBBXPPamR43HHU1IpAUlfvUpAM3YJPU0EFuc0BuSDSrja1A7iYBAIPHpTZPm6cUqAeUcdaUjEfH16U+oiFUZWLP36Uk7K7KuenWpk+barVWl/dyFQvzHrWq3EWkZYo9o/Cq8nNwDjnuRTbc7mJbqOxqXkSBuoNLaQFgZCk56UHlMZzSj7xAoIxisL6lXImYYwe9KiDYKZKcOmOpqYEFT7VT2AhnBO3A5pzchQODRISQMfjTX+6G6E01sFx+coo70Bc8E9OtCjG0k00vsJBHWhaiGgblbvinMcuMfnQq7Tx0PelHCZ4z3psZJ1jU96rkneAOnepgWzjHFV2+Q5HJpRAmB525pQBgj0NIhz84p2Mkkd6QFS4jMigD1p6syIoAp4Yhie2OKYilhk+vNaX0sInXOPrThyabkl+KeeCGFZMZHlhnA5pdzDFCt85J705+Vz0FAETltxUnrTFHzErTwCZCT+FICFOAOasQ3BDE5wBT1cYHcGmumRgnrSAhFAFPdDJk+91pcfMBUSt8x9ewqXOGGRmodwAH5+RTe5zzTicbjTcHHXrQIaxwCfypFP7zgcGnOoICnigAK2BT6BsyPrk5603IKMe9SkfeBpn8BziqTGQEqQWxzVi2JMPPWq+DtODViAho+O3WqlsJD26Ckbpn8qVeckc0x/un2rNBsRAjOfSpl+6fU1UBbeADgVaHAXAq5KwD89zQB1HelwNgIpqnnGOKgBp6lhUBbc5U+lWeuRjiq4PLdiOlXEBY/m4zxS4w24DrSQlTkDpS/wC0Ogoe4tx3JNDEbPajdlQaCTszikUR4CoARx2oQ/KwIobL7SOnelAJLZqxBGAHYilXGSDSYyp296UjC570mABcCmMuXHbFSEFlHrTHJJBHPrSQxrHJP900gz5mB/DSkDG3PWn/AMBI/MVQtWKp+c47U1k5LU5Pu5P3vWlPJ61OzAhALoQ1L0AFIxw+0daUgHqOR3qwsV5OGI9aRAVXOMH60+YZPHakwTgAdatbCHIcDIGcU5cKxJ700rgDHUU5gTnoKkERPhcs5+U9qiPI3qc+1WZANqs3QdRUJCpEWXPJq09BDY2aRSx4Ip8Z8tT/AHqTLYU9PWnu4RzkcChj1HgblwetG37vPIpEJYNk00ghgQfxqbajsiY8gqeBTl+6AO1Mbk4PU05CAMAmoew+o8kHmgfKOBxSHkDFKMYGDxUiF28EenSm45P1px+7yM+tKDgmgYdxinjk+9NxlhkcU8DPI6VLERGBN5bGc+tZV87bio7VsSHnNZd5BkMy9x1rei/e1GZhYn5v0qLJcjvn0p5bEQx17mmRZBB5x79671sRaxp2vyKB2qRpSTgjApiYKZz+FLn5eenpXM9yxjOFcE9DT8rnGeaacMQMZxTRFsORxRoLW4khJJC9aVV2qOxqaPay+9RyyqMqD83pTv0QN21I59ygbOTnmoY4W5Zl6ds1Os2wncMk9KiS4Z5iCpGTWkU7aEN+8P8Am2hc8e1RyqzRFRnIqdVKPwcg9aqvdyozKqgD1pxWoDY4njYOwwCPWrGQyHpj1NQh2lwC3B7VOkSsu48kUpPuNbDQ2wHvioTgk469qtqipnAqs4BdhmiL1E72GMCQccn2pjEkZzyeue1Sn5jweO1RlBw2MkVafcS0GIME8jFJgkU4YwRmm5ODjPvVDWrF+6vHIqbTzh2GDios8DFSWBPnHFHTUH8Jtggopp6kcGoFPybehBqdfmxmvPmrMvoSrwcZqrqa7rYjJFWg2D71Df8ANuwxkUqfxIPIy9NOGYZzV26/49Jv+ubfyqhpxJkOfyrQusC0m56xt/Kumf8AEQlc7Icrml6mkXgD1o5NfPlCnFHXg0rc96QZ/wDr0AMYZPHWn44FNJw3qKcDke9MALetIOego6DNKOv1pAIMn8KGpccmkYc0ADc4pQvGDyaQjpmn9utHQBp7g0mM04n86Tgn6UAA569aXsaQYJyaAcE5oDoJ14pSckYox0ozg0AHf3oHUfyp3H4UdSD6UgDuKTvzS96QnkUAJnB5FK316Uhx3oLAjjvTAYBgn07UJweaUkbulN6GqAlBG2kO0rt9aFA24PWkP3QPSkgM6/mkt5EEY68VdtA4tx5h+ajy/MHzjp0qQMB8prVzXJy2BC474pp4ann71Mfr05rJANJO7I5p6Aht3Y1DCDkgnpUp3djxVNW0AYTh1p64Ge2aY43sMHgGlOC238qAHgd+1NYkKcCng4OKY33aQBGSUG7rTyQRjNIuSAKAN2TnFLqA0nC0q5I5qNPmLK2TTgGBx1qmrANLANwelIxJyR3oACEkjAoLcAdzTAaRkqfSmRgRy5657U9yVHHNJMQiBsdateQFs4NN39R6d6ZC2YxnipVGFJFZNWdmHUXgAZHFD5oPAFHO2puIaeFA/WlJ5FBACA0dgaBjjw1NPHtmlYng0HpntSAQ8CgcDmhTngik70wGAEKx7U6IZSmlxyp4FEDZXBqrOwIlH3SKQc0q9zQMjmoAbjJINL0BoJx0oXhcHkmmAcnB70zqDjqafnIyenpTD13AU0ArL79Kbld5GKU9eaaCMnnmmgFHynHpSkDYSetIvJz3pQMt7UWQCAcZNMIIYY/GguTPsH3akYgHbVbBuR7QQM8mg8cdKlABPHQVC4JJ5wKSd3YL2E3Z4A6VJkFcDqKrgiNiM5JqZV+TOetVJIEHPAHelxhQO9L90U3+HHepAfgAbaRT82MVHv8ALHPU1Ip+bP5UWDqDoCVI6in4GRUYLbvmp4PPtSYCkjJPpUbHc6kdaeRxgdKaeF9qEBE5JkBxxQ/zAr61L/Bn0pmAWwfwq7iGwrgj9asscN7VGFO3gUu4ADNS3djF45NGRRwPoaQjjrUgIckkimyFsDb1708EBDQR90dqrYBQMKOaGO3PpQo2gg0E4UDrS6gMlI3D3pevPaiUBipPIx+VIOme1PoA7oCBSq2cgdqF4ySetIvAZvypALg7SO9LjkDHakBywpEbdJ7iizsA0NvY+gpNgG7HQ09AAWJpjKCxx1NVfsBFbkb2x+NWBgLnvVW1jeOeTJ4Y1ZB+cqeKqdrh0EIwC3WoufMqfIJwB+NQbiZCuOM0kA9sDr3604n5BjrTZcA8nj0pxwQFz+FF9AI8/dJOD/Omsyqy+ppzYZgtQyIPMTuKpJdQKN8IxNxywNRzjcVXPzdamu9iSlsDNIX3kbhyBXXB+6ibakKg+Z5Y4bHOanhEQlKfeZfWoYWLzHHIA61ZSIq28jOeBVyv1AsRB5xg+vFa1tbiIbm5Y96g021MUYd+p6A9q0Vroo0LPmZjUqdEATP0p+3gZpwxijpz2rsOfcgnmWCNnboKx2Z3kMh+8etTXU/nyE/wL0PrUARSfYfrWkUcFepzOyBFIbHQmhQNp55pzHC5pCOASelWc1lcQYOOMUq/dwVz+NO5A55NIeTx1HSmOwDJJPakYbl4oVsEt69KUZz9OlAXTVhrAZIHBo25yM/pS5zwck96T5Qc0g0Yp+9joTWtGmyJV9Ko2yebMuenetJqiT6Hdh42uyM/ypuKefpTT1yKg6RveijtQT0OaYAOvFBzmkphagYpP6UxjxSMc1GW4pDsUNWOY4x23VmflWnqSlrfd/dOazeornnfmZ6OFf7sSJN8qoOcnvWvI6xQn+6orGJKtkHBB7UkkjsvLHHcZpxloaVaTqW10LungGSSXqafeYcxRk8M/NZ0UrxElGxnqKN7O5kd8sOlCldEOi3O5tTN5FuSi52jgCqts08gZ5eB2XFNS/YRjemT7VFLeSSLhBtHr3quZXMlRnqrfMWUSXFyqQqW2HJI9a1YJ9uIrhNjn16GqukzRQb1l4Lchj61cvLYXnllJB8pyCKcdrhUiotQktO457TDGSA7H/Q1kXUs8svlzvkL/CBjNb6kRxKC2cDk1Tjiiu76SX7yoAOOhNOS7Co1LN82tjFI454FbWmJttM4xuNGoW7SQKsMYznsOgq3bx+VbonoOamMbMqtXU6ehnaoVAjQKMseuKZFp2UDSvtyOFFGphpL+NE5I9Kry2U7Nu80l/8AeNTaLk3YuDcaS1sMurUjJjyQtUAASPWte3LlSsnVeM461RngC3gRRw3NJwSZVOq7uMmMitZZeVHHrT3tpY+XUgfStPIjUBeiin6fP9rEiOgKjofWq5OhH1iV20tCgsxjA8vK/jWxYXguo9jH94OvvWbfW4tpRt+43T2qCKUwyq6nGDShJx0ZvUhGvC6Oi6NjNGODTUkEkayrzkc07p3ro2PJaa0Y7gDpTD1px6VExOaOghxNNOBTicDPeoyQetAEbHPSmdqe3TnpUZODSGBPFRtzTiaZnNJMdg680Z7Gk70gzk80xDxz9aeB61Gp5IzUnFAwPSl/GkJOaTPFAhDyeKbLnAwKkHPPemSbcDFGoyJOD/OrESlmAqEDPFWrYHG7vQgZZOQp+lQjoKmb/V+9RdBTIuNc9qhJ61I5xmq0jYWk9BpFS7kIAxWHcv5krHOQOlX9QuNkbNnnoKyff171zzdzohET+LJH4UoJHIxxSZwPxoA461BfUUDDcjpVmFM4HeoU5Of0q7CuTis5sZPCgYcjFWo/u89aZGuMgnkjipYARGSfvdq5ZtMoXJZsKMipFPmcE4HeoN2CRmpI2VYyT3NQ1oA1gRk5pwYAr60ofvjrTCcHe1LcCzg7MZxmqQdYjs71bHKZ7dqpbB5xOOtEEtbgWVIyfeoZXUHaOSalCMWUZpCqkg45oVrgMQ7MDuOtOk5XpiiTHmcDmlfAwccU3vcAjPA9qkz8wJ6VXhP7wr2BqwxB4HSlJWYEbjL5FO2jkihj82B1pA+GPekA8fdGOopkgPlY7mlBxJgDGaG5bmlsBCMmFuMY70zgocVJlizLioSfnK9q1QD93C8VIgAGAeopnBUZ/Knqw8tieD2pPyEMXO3FLuIAXvSDcpBxSgZbkZNAxExk47U5+OM/hSKQZWGfwpRyDmn1EMVVEh2nk1BINs+9uppxOwExkDPGakWNSvzN+Nbq2gtyBWDSnPANSuyI2O2OM08xo2GA5pJh8hYgE9qiTXNoG6JUJ6k808ANnNRQfNGGbtU/8AIrCSsyiDGZRkZOetSsBjgUi/ez3penfk0NgMbO3Hc1E6+Zx6VM4ww4Oe5qsGIZi4x7Zq4oCRPmA9KWUjg46U2HhDg0SHYuSu72p9bCFc8jFK/8Jz+FN3b0UqccU7BMgNDVnYe47q22opU6cVIMli2OB1odj1qVowIVLYKAY9KsDJjx6VXBO5mHUdqnjYFDg5pyDQa2CvtUewRLkE/NU38AHXvTdwJb0FCYgHyjp1p4J289RTBnac084JqWMaGCmhgcYpxPb9aDycGgCNyQQRTcHdnFKxJAPSmMWJ9qpAxzHL4PTtUeSze/pUz/AHVIAqLDAs56GqT0EOXrzUjHBB9ajT73TipB82M8AVLAUnn60vf6U0n5jnpSiTg0hjXI3dfpSDlzRjJye1CnLmmIM5Ykmo93XHOKewGeKjAyrU0MiZwXUA1PbkD5Qee9VUiO48/KOhq1D1rSVkrCJE+XcPSmtwn1qRMb2NMcZGemKyW4yswwRipwSoHqarsfm9s1Y+/gj8at7CHj7v0ppznHXNOOc+1NUHHvUDEfgcVVuGYOhUcmrT8NjtUbkbzgdKuDsIaCu8inD7pzTcbWJPSmkllBJp2uA4bhxnApzZ27c8n0qNmBXr0p6v8ALk80ArEZY7Sg/OnRH90R70jIM59ajhYh9hqrXQPuWFwo6UnUdc05gCcA0hGABmoGKuRlqYByx6U8d1IzTWbggflQhEKHJ+cYqRWXYaa/ysH601cOwXsKvcCYcDjvSsMEGmr8wOOq0pHzYqWFhMHrjmkUZAz1pWB3D065pCecjgUwIZAfm55pi9OvTvT5l3Nn86iXljWi2AmiDZ5pMkyEGlUkDA6CnFjy2M9qnqMYGLsSOVpqsjFlK9O1Ku5VLA/ShiqqWK0yb6EYJaUbPuinqfO3buxpoIhOQOD0pNxjYYHBNVa4MkwXjI7Zp8Y2Agmm9QQvSnqTnHpUMaGIf3x3Dr3pskojfpkHqaUZ+0Ajp6VHLGWYe55NUkr6gXF5UHrSgYAI7UirtGO1KuCp5zWQC4I/Gnd8Uxfmb2qRfmPuKTAcBwSDR92PjvQMcjp6Ucg4J4qQI2G5CD1NMZMRcmrGMqfSoZuI8E444NUncZgXluUGVHyn9Kqjk4zyOlbHlloCDWabcxOAw4r0Kc9LGbVh0c7BQP5VKGZgSegpscWeM9P1qd0dlAVTx3pSauVfQE3INxH1pruDyzcjsKQ2tw0fBqJbKRiSRznrUrl3bDqBuAhwBxT1mj7jB96jazkX73OO1QeW4YZBx/KtEovYH3LjPkE1Wabc6vwCDTfO2MMg0jI0h3BcU4xsJ+ZaikR5CSwGOlHkGUlwcD6VRMTbTjqKfHcPCvOT9e9Pl6oXKW1hCKSOaeEfGA+BVM3jB94H3u3pVxJlYdTUSUlqVFdBMMH+ZsilMCMM9zSMwcZ46Um/AGT/APXqdQRXYbWK4pj/AOyOM9qsysjqCD83tVaTBAxnGa0i7idrjCRytKeAeT+NBA6gDnoaTB4yeTVhYVRweKfZjbcg9TTOWbqM063O26U447+9AmbOBwe4qfsOKgJ5J9qeu5owcc1xVFqPoWR/KmXS74G47dKcp9OtOl+aEjHUVir3H1MOxIFxjpntV+6A+xTcfwH+VZlm3+mAEcc84rTus/ZJ8f3G/lXXPSauCWh2YB3UvOcCjHPWjjNfPlaB0FHtRmk45PegAGc8GjGOtHU8daX+dACE4wMUAkn37UYzg0mMg+tMBx5/CkzkH1pAMcZp1ACD5l54IpRzjHWkxzinYHbqKQAwxzTMYOfWnHmjg8GmtBijpTQDg56UoGCR2pRxnn8KQhO5oChutDcAClHbnFAIMEmlBy1HQ0fSi4B396TqaX3zRjGSaQCMeeRkUmBtz3oU84xxSjimA0+tN5wM09xgD+VNYdKuOwDh06U3gvxzSoecUi7fMJH41IEmOen0pDgkcUvf6U0feqQEJ+fFMfrink/MMioJ3KMuT+FXFXYgQ5jPqe9PIbys5+amRkkkEjFSswA69aqQ2MUsUGPvUCLFwHJyfQdqUDZgZ61IFxytCdtgB+tGARzR357UcmoABxnnmkIPShD8woJ4P6UwIo5CZNvQ1NjL89BTVUbsgc0/PzZ5pya6D3InBdiucCgKM80rEBj3Jpcc9aLiGEZyPypCn7kjripGwASR0pg+aMjoapPQQ+EZUA9al9R2qNN20CpONp9azlqxsByMn8KXHAxSdFFBO1c96V2AnsaH7jPFKMYyaOM4ouDA+mKQ9h605uOaaefmoQAQdwxRxn6UpOSPX1pMDBzQgIN+WIIxSKdslS+Vkk0vlLng4NWpIOg8YZeKM4TFG3A46CkPPFS9GAYJx2oPJ6YNGehI5peAR70gGtzz6U0jPB6VIQQetR9sU0MCueaYPlJNPIyvFR87gvaqQh3A5HWnIcjkU3Hy8c4o5XAA4PWjcB6jJ3d6CD1700ZBFCMSDu6etKwh/Az61E68ZpxYck02RyMYHHpQlqMjCAMD61MBkBe1M9APxp4wBx1pvYBXxmmOMNz1qRuF/rUTkb6UQGt8zgDjmpCAWHPApvDc9vSnbc554qr9AEBO7mn5OaQYwTnpSJwDmk9QHBdvHrQewx0oB43UvHXtUh6CBecHpSYx25pxzkUNwck00wAg5460bcjmhs8UoyB7mlcBoxjnrRnPfilK7Rj8c00H5fp3p2ANoPPOBR1Ix07UvYZoHYUagKvzE89KG68UmNpNLwOaQCY4I703aSNp70p4ORzTiRnnrTARsYx6UYwMUv8AF7U0j5+tACOxUAgfWlQc+5pCMPg9DTv4wR0pgGeCTzSYzz3oc4bb2NNB5ANCQCQx+WxYnOacQGfPenEDIFBG3kcZo5rvUBq4DAD8ajfCy7j+VPRiDz1qOQBjkjNUtwEJ3/MTzT0ILtxn1qNMIucY9KfFtUH5cCq8gFQKWzjmo3JDHj6VKqlZCc5zUbkde4pLcNjMvmJYKV5PP1p8UJLMfQU+63GVfl69DnpTbdgkjAtu9q64v3LE9Ssoa2YnGM9BWxpUDTossgyo6Z71UtbZ7y8IYfulxk100UQVQqL8o6YrtpU+ZJtGU5WQigYqTaMe9af9j3MNkt1NAVhbo1Sazb6fZeVFas73DANIDwEBGR9a6ktNDknO2/X+v6+8yMDvVLULgonlKfmb9KuSyKiFm7ViSOZXMrryTVxVzKvPljYaCqjFK2QeKcO2etIMYBHWtErHnttsBgY44780w8gkjFS7STz+dRkY6nJ7UxO47PHHUdfam/eJ5yPSlP3f60gJAGRRsO93YAcHnrTh1x2pikgYI5HalY5znvQPpqKFwefwFDZBx37ClBy571ZtoBJJvIyoobLjC/uosWUPlxbiPmbrVgj1p1IRWV7u56MYqC5URkYzTDx3pzkDvVdmOcUrlEhOe9JnvTQeaCcUDAtgcVCTxSuwHNRFvSgLDmI4xzTPrSgZ56U4L7UtyiN4hIhQ9xWJLC1vIUbOB0PrXQgY60ya3S4Qq469D6VMo32N6Fbkfkc/gYqJgQM1oT6fNEcoNy+1UnRwcMp+lY2aPRU4yV4sjAyTgc0BCef0qRFb0PIp6q/9zn6UncEmxpHGBSonHNPEbZ+4fyqQRSAY2N+VFi0IF4xSEsrfKzD6GnlJMcofyppBHUH8qdmth3Q15JGGC7Y9CamtLprQnAyp6ioeOho68mi7WonGLVraF6fVSyFYoyCe57VZtb6FoVVnAcDkHvWR2qF27Z71SqPqYSw1Nx5Uadxcxf2ijI2WIwfSm3Nw0EYYDcxOKymJByOoq7BfRugWcAfWiMr3sZTo8tuqRcti8sO912k9qryANqCf7K5NSy6hbxR8MG9AKgsyziSaTqeB7VclczSavO1kPun2QSEelLpd4lspDjAPeor9vlROxOTUMZwgFZttSub0KalT16l3ULqO5KKnKrzmqR6nPSjPGKM8fzpSd3c6KcFCNka2lS7laI9uRV8n5fSsPT5gt4AD9a1TIe1dEL8p5+KhafqTFhg881GZBTBlupoVfm6UzlStqLuIFIqlzinlCW4HJqbaEX3pgRSr8oAqoeAR6VbkI2EniqIOTSBCk4600HtTscdabjnGOKQwI5PqaQEHrSscKTmqL3TMcJwPWmkTOajuWZJ0iPJqJ7wgDbxVQHJz39aOijH41SSOWVaTWhZN1Iepwe4pUuZDjuKqg+o5p46Y/OnZGftJXLiXv9/gmpjIrqAjZrNPBz19KZuIckHBpWNIVpLc1V5P86vQrhRWTa3QLASce9bCY28UrWOlTUloKxytRnhaex5xUbnBpsERSGqFxJtUircr4B/lWReT7EZuwrOUrFxRkahMZLgRgcL1NM/hGDmoHZjli3OakRiUGTWDWh0RuhSQDwKVevPWmlRn39acnIAz9aT2KtYmiwOcVfgXkDFVokOMda0YlGCcdO1c1WQ+hMqFnBHapIlzuB6etNyygY5p8mdhCjk1yt9BjWCnKr1pNuFGRzTIoXiPXr61K2QQB1709NkwFIJAYcAUxsupGOfWn53bkHUUmCqnnmpQEqcJ6Y7VAihpD7GpYhleTye1RB1ErK/GenvQt3YCQACTdnPHFNOSCAOaAMdBnFOPzcigNyB3wVycZp7cpj8qZKFOw45B6U9yEAFXbRARRNtcg9/arIBLYqrhhNnOVP6VaJO0nvUyEICSSMc0zHGKcpyd54FLkbunFIe4pHzHNNbJb6UFgXBB4HWmshMmQxxihIQ3IMm4VEc+duC8U+NfKJyeDSEYf1zWishiduakz+7BHSmHjjHSngHhe1JhsJKDsU55oGNuSeaVgd5B7U0klsY60AIuA/uacTh8djSEjcPQU5iAw45oYiBosux547VLnbHk9BQGKk4HWorjeQqqep5q97INhY5MnJOAe1LOe+OBTG2+eikdOtPuAWTAOPehrVALby+YoIHFWG46Cqlq2F2ZyB3q0CCevFRUVmOw1vvqM/WnsM4Xt61Fged0qXq2RUsBHIGf51Wkx5mMc1ZJ+UsaqyHv3qoBcSIkg44ouctFkA5pItwJDd6kfHk+wrRuzE1oRQbuATjFWD9/0+lV43y644JGan3kqWI+YUpDFU4RgBQSQeOc0oJ342/SmkfOV3cVAEYGJiR3FSxjCnAqNcAbjyakjOAWpyAXPy5z0qMkFGx3p5GB1qPpKCB2oQXJFBCj2FPXlT60mSPTFKvAwTUsBASF5oIx060ucAjv2phb5hQhCOBtxnmmMOQKlxuOO1MPD00xg+BxTRySpqRgBjNNUfMTTTAYnDEU8H5cU3B38CpBknAPWhgNYtt44pR8qc07AYFc8009dvpSAY2NuD3ojOELd6JMA0o+4P1qughJP4SBTI8bjUr/AHA1RKeW4+tC2GJwHAHQUsR5J7etK7AHI9KbEwxtBzVdBE6cg00j93SsSMAdaH4XPSoGQbPlOKkQEJxUOSGPHFTodwyauV7BYcPu4ppb5selOHI6UxiS/HeoQhJCcZB5FIvIBxzT2+9zTQdq7T3qlsHUYeFyKTZhB608jBxnBNMOCdw5FUA0qF20o+VT3FKx3rn0pEOWIxjFNbajFkXeqjuO9VcFLnJb9Ksu5WAkVCRuYNjntVR0EWFwMUjH5sEdKFOYwDQ4BYfrUdQ2QpOWOKhOeVA+b1qZRyaVo8kHihOwFSUOse1eSDT0zsT1PWlkU+YQTUHnnJBGMd60SugLaLtHHUnmnHhc55piklgw9KeBlTnrWbATPycjmm9OOx7U4fdIJJJ6UDGw569qB2K7g+ZgdKiUkPuxx2qdyA31qEYJwa1jsSOI/I1JjCbd3PrUSZY/SpY8s2TSY0RMu2JcHnNPkj3BdppflLH07U0H97k9MUXBWGOm4KndelKcbtuOQKEU+cWPSkZt3z45FV5B1JEYrFnHzU5XDDI60xnyqsPyoyQdwHB61NgJFGM+tNbDLg0bvmXJppwGOaSB7E/OAKei9QetRFxntUqnnOeDUNOwAo5K0/BGPWmgjNPXjAPWpBDguCBnilwFznv+lCnnGaflckVAwKgYGar3CBwcHkVNJJtYDFRsMNkelVHQRReQKmcY9azgGuLgMD+7q84+8OuaZCnBxwPSuuDUVcROsKKoUD6VMEG7HaokDYB/Q1MSO3WsZNjFIFGzPbIo+nSkL471AepXuiojyDzVGKF2Xcep7VemwV+b8qAvygrwB2reMrRsJrqU5YQwHf60LHj5fSrJALinlAxziq59AM6W2+Xgc1Ue2kHOK2WTDDjINDJ8ucYzVRqtBbqYPksT8wx9KciPyMHAFbDR8ZxVOYbR0/KtVV5hq5WTIfg44qUqxjxjBqJGBYH8K0MEoePeicrMSXQqKoRCD2NMILDjinu5Y56NUYBOQMZql3JfQYAu4k+lIflPPOKVRyctml4yeBnHU1XUXUQjBJI4psbfv0xwKOBzng0m798pI6elVYpq5txHcCOnHNTIx59Kgt+SD1HrVhRwea4qytJoaXukyfXFPYEx4qvHvV8EZHrVlcFOBWTWodTnUO2+/GtS6x9jm4/gP8qzJm2X+COAelaNz/x5zD/pm38q6p2bixI7U5wKTFLnNHVh7V88WKaTHel69KQngg0AGBkUYBOKQc5zSgYNAAMUnalJ70HtQAgFL0OaQ8tQT2xQAozmj+dGTilA4zigAwCMUmOetH1FLgZJoATByaUDJPtSDgZ70oPU5psLCAUEZxQDkmgjAxSAcevPSl7E0gx6dKB8w68UgDHFB9ulHfij1oAaQaEyV+brSt0pobcDiq9BoHJwaQ8rxTsEjJpnQHjiqg+ghY8F89KcBgk4qONhjI/CpsDmlJWdgBR3oxj6mgelHapBjW7VUYHzBu6VafJbGelV3XDbiauAupFCVaY7WyR1qzIuMEtwKrJF5YZgfmNWkG+PBPPetZ+QxGVXZTmpVyODULYEgKjNTkZGe9ZMBucscUg5wKXjqBx3pPu/NSAUHB60HoTjimk/P9akPK4FSAz5tvFO6DihgNvFIB6GnuA0ng5o4AwOppTx8vrSlQvSgBGAC4NJj5fl6mnMQcU0ggDmmrjEhLEndUpPP9aiXIkJPTtUq4pPe4hTjj1pDwOnNKDQPfpSAXgKBmmgcZ9aU4IGO9LwARSAb0bnpS9ge3pTSOMd+9KOgHp3pgHTkUgGRj1pS2OKReWGe1AC9zSkHAI70HqB+tB4HP4UgAcn+dICN59KB8v1oxwR3poBSBkk9qQAEjnmgHIIpB9888UIEOLDOCaYCB2o/izjik3dqaQAwywx096acg4Ap3Uc005xkUwFRNtBHH9KRSQcHmnk85ptgB4AFIR8vXrSglskijrnPQUl2Ahc/KQvLCl3Zxk8+lBbD9OTQcbwRyCKoY4DHJ60sY5pFPIyaEBJJ7VIhWPOCeKjc4Gep7VI2MVEw4PPI5qoq4D4R8n+NOAPPFInKY6085247mpb1Aj3Bfc0xCMkg9TzTzg/N6VDnEwAPB61aV9ALJAUClPL4oHKhcfjSj7xrNgB5FKcHjFGMMAOtI5wc0kADoaQElRk96XntSgjb7UAIfu+tNHJwaccYqFTmTPaqirgSkjI9KZks/I4FEmdpwKUce5oWgD+rdOnWmDOwgDv1py45Ao/2e9C3AQ/c4600qS49O9Lxg880nfnrTWgCnlTSFfmHrTtwJxTGJ6j1o1AkbhuRTAdzHHQU9hknNIAFb+tGgCN9wetMYguq9DTzyB7VEybjnJpxSAkLDf7jpT2OepqJQCSSORUhIIGB160mhkaL+9PtUVySkZYcZqwB+8zj8KrXhPllcYzVR1aE9hskn7oADJ7inRSMTsI6d6p2rF5QXFX1QE8ZHNdEkoqzDcRHDyEA9OtJnJbPSkDBbsLt5PenEkE/X1rOyVmgKsuQQGGfeks7XzJGEY5bGfappt8hKgZzwAK1tPshaQhTy5+8a68NBzRnOSiTQQrDGFUf/XrY0y8jgcxOqKHIzIy7to+lZ4XaR61Kibz6e9epFJKxyyk27s7rW9btbLR4xFIs0kq4iTH3sdSfSuCmmeaVpZW3SMcsaJZdxx2HvVWXeykR/epuy2MYx1cv6/r+u96t7KJG8rt/FVbGQcdR2qRoGj++PrTN3HHJrSK0OKrK8tRCAepyRTMAYYd/wBKc2DJz1PT2oPI61WxlZNiKWye1Jj5SwpyLlgx/KnNk4+Xp2qkyeUYBg8dKaV/ADpTiPmJ6A96XJxnuKXUtJNCemDTe3zUo4c/5zU8Ns03XhPWloi4py2EhiaVsDhR3rVSMKgVTgCmRoI12qOlSj0rNu53U6fKhCMfWkxxmnt+tNA/WgvUrTHnFQA881JOcPiogOKm5aJBx9aRumSeOtNDcj3ocFuB0oAhc7jgdKcIz6VIiYp2B0o1AYI+aeFAFOHFI3TihBcryPg09XGKikwTSISaB2LJwQTVUxb5D8o+tTjJx6VIqBR9aLJhsRJAu77g49qkEan+EflUmMD39KTgDJ6UWQczGkIozgAD2qidTi+0eUibhnGRVe5uJb6U29v/AKv+Jh3q3aWEdsA2A0nr6UKw9bXZcKj0H5UhRc8ouPpS9SaOooshXZE9tC/3o1/Kq8mnwN0BU+oNXCaiZ+1LlRcatRbMx7qxeIMyncvp3qiBztbg+9bsr8cnj3qqIYLs4VhuHcVlKn2OulintMyHGO1MAFXLu1ktwQ449fWq0aEjOazt3OvmTWgoA6DnmteKAx2yDueTWUMLye1asN9HMFQHa2Mc1pTSsY4pScbIoXz5uipH3QBTFkBA9cVLdLvuWb0OKqkYPNZbu5tSTjBEwYMaa567elCbTwRQc56U0yzQ0qENKznsK1ti4yaqabGUtixGCx/Srgbgg11R2PLxMrzdhMqD0qB5j5nHQdKSeXaMDr3qGH5mA75oMDVjAKbu5qNjzTx8sY4qJyOneqJ6kE7ZUgVXVdoHrUsuS2PzpAKSKGY5ppZUXLHAqQ8Gs25lEjlAcKKErmU6nIrjbidpOn3aiIz1pc85FKFGAPStNEcjbk9xp65PNGQrZB4NHpRyDjv60hJ9UCnnp+tOAzyelM75NO5HHWnoK90J7kcUwfNxjn60ORu46U9FyhJ7dOKAWw9FAYgfhV+1uTFw33O1VccAEcipM5//AFVW6I53GV0a5YMRt5FRzHAz2qna3Hlvsb7h/Sprt+B3rKSsd9Kopq5VuJMDk1z2oS72CDnHUVrXUmEJNYEhMpZj1b9Kwm7s6oIhdRkADrT4zxikwGXI6jgUqZUGoexprcUn5s459KkhAJ9TUJYg1PbZLc8Y7UpLQpNMvwoCASOvpWjEvAbHFVbdeM9qvopwFxzXBVepQ2QDB21IuWA9aGGFII4pkDgADH0rLoA8nkU3cGOfQ0uQuT3zRgEnHakG4AgNkdG70zoOacAduQKbIWJximtwFt3JyMc5qK4QCcE9PWn24IfGf/r02ddzbe/WqTtILEsZG04OfenEcndVaI4AGeanbA60mrMCCRV2gn161IzYjyRmkkYL82MACiMmSINnj+dV0EM3YIzgZNTg7gRnOKgKjKqeB2qYYOB60pAOKYyueKZt2g8U/PNN7kHrUXAYo+Tg8GpFTanFNZBvAHbtTsljz+NNsZHg55pkm3HB6VK4H3qimUBOOhq467gRodw4BzUqnOdvJFNxsQEH8aVBtbecYpsQ5shTn71N3EjAHNICVck9D096EOHwe9KwAw2gDHPenOQWA9KMcDjn0o/iz2oAYxxICTx6Um0DBpzxhyDjkdKApOCe1UnoFhjJ5sxboBwKJH2RbgMn0pvmjeFX8RUzDehpt2tcd7kFsuAXJ4NWhghsdKgQkAqOlT7tq/zqZ6sQx2xg45NSL9zg801jhQepoCsQDnBqWMf8xQiqs2Syr271ZGQSDUE3UUQ3AiGDLj0qfG5Sg61ABl8mrMYyuT24rSWwFFBtdWbnHFXBySwNQrGBKCe2eKmTGSB0NE3cUVYduG/PY0yQENuHPtSlRt56g0rAbB7VAyuzfPgdKerH7ucYqPaRnJyaSP75APIrRrQC0RzTdxEu3blfWnnOQRQSF471ncAUZJB59aXjjNCgbOKUkkDFIAf749KQgZAokOeRSc7Q3t0oSAAw3e1JtypzzS5+UYozlsdqAsIAH69qbkluKfnk46Uh4cY60w2E6N1696VPvCmn74GM04DnPehgOHG496aQSpIpy8AnFMyxVqQEM2DjFOUgrTZRtjJHUU2JvkPfNaW0F1J2A2Dmojgkg9CKeeE+tRuxXOelKIEUzAKEzz3piyMuGz8vpSuQ5BNP2AqQfxra1kF9SVG8wkg0pO9celMgQpk54p8fDEnpWT3Aa4IbPrUicrimEnnPWnR57npSewD/ALvA/KozkEmpBkNmo92WK0IYrD+L1o7DJ4ozuGKU/dCj8KAGOMHJqMnb8p71My9z0HWqzdd+flHAqo6iBDjIHrzUpAB+XtUeQVLelCn5N1UwuOc7iRjg0xlwynd0FPORHnFISNoY8YoQDV+U7TyKlb73rTOdxYdKcuTF15oYDyMMKHJPTGB0pischj1PajJ5z27VNh3sMmUMu3pzVG8BWVSDxWifuDd1qvPB5iN39K1pys9RDbWUtDt61bByuAapWwKoRVuLJGTU1Er6CTHYIxjpTScv8tOQ/Ln16VGhO9vQVKKGSjPI6iq+QZdpq3IMLgdarEHIJ61pDYl3HA4Py9BTyTvBXpTFBEZPelibcMEcChoY8qqsOnNRybfKPGacBh3Oc0ZzHwOKEAwlhGjYwO9KSPmUcFh1pSysh3Doc0oKmQEdPSmAnKopA56Gj5geO9NB81TzyDT9/B9u1AbjuCw4poOWyR0pRyMjrQ5Cj+dSD2BTmQg1IhHmY7VEeQWpwAAz60NCLH8XFSEfOPaoo8kg56jmpecgVk30GOVMvnNLnbkmndDSPjGSai9wGEggE1EWO8mnOcNTG6n3FWhkBTJLAc1Go2svvVgcEg9fSoggLBhzitUxD0wV57U4UzGORTs4WpYDwcA0jcAUEgADvSNwBikBUc+Y5INODbohg96R1WMk9jSRDgEHOa26CJBy9SIN3ao1b5yD1NTRjHWolsC0GOgBPpQcFAT2p04OQM496QZzt65pLYOpE6jGKzp2+YgnFaLcAiqU6gjkdK2pvUGUkGCTgEVdV2ZcdKqqTgtjj0q1DKXjJAraeoluVnIBIzluKT5dxBPX2p7KC3uP1qMjngVSFdsiYheQOhpD05GRSncSSelBXIOOa0ARFABHpzTW++CMCncZ4xTH+Ugk5Gaa3C6Nq3JVUyeexqygzkE1TgG5Eara53EiuSv8VxxtYkTrj9amj6H0qID56mQDaeeO9YNjMC+BW+yO5FX7gk2cp7bG/lVPUl/0kMoyAatTHNjKf+mbfyrpe0QS1O3xjIoxxS9ce1FfPlAee1IetLjn2pPxoQB9OKF4HNL3OaOoxQAg4HPekPb2pwIJ9qToOetAAM59qQ96U/eyKaTg4pgKoz1pcnvSA4FOx60MA7ZpuT0xTgQT6UhHuKEmAE4XpSZBUUrHoO1A6g0dAFH/AOulI5FIc0e9SAp46UoGPxpMDB5/CgkDBzQAZOaXPNJyeac2Rz2oAaQOlMVcEkdDTzzimpk5zTQDhnac1CxJyp7VMRwKiOd5qoasCGNtpwas55wfSs9H/fkE89qvKxZQTW1WKspALnPNOHBz3oAwOTR75rABD940zYMH3pxYg/WjsCKFdARsuRhetCHBA70uAG96ZEp3liOKrdAiRfmPTmpMg/LTOM8UAfOT3pNAiKWOQnCthe9SLymPSlJweelIoG4471XN7tg6iAlufSpFBNRPkNjtUgzjipkAcFcetNkyEAFSkDb16VGPmHWkgEzuIpWO7gUBTkjtSBcEDtVbgKy7ufemsWBHpTycDPamYycZpIBcDdk0sTBOPWmquGyDmlBG7B6imA8/Lx+tKrcnPSkAY8U4AknPFS7IBVAIFGQTSD5aGGMelTa4DSOaeOpIpu0jluuaD8zECqsAjDOTQQDj1Ap47gdBUYGcnPehBoLuwvB5peFAGaQA7h6U8glsikAnfOOlGD1xxRyRig5II7UAMIwxNOX+feggdKBgDGOadwEcdfamRrgEnk+tSN9zA71DGSGZfSqjsBITgEd6YRuAxT+NoPemOMlcdqLALn58Yp2N3NNBwTnk0oIABHWhrQB2du0DvSZ/h70EYGe9IQSwIP1FSGozeOuKjQ5XJ6Z4p0xCkdhSoNyYq+lwQ8c4alB+faKRRyKVR8xzUtgDfcOOtR7+ACOTUpHHPSo2GOccnpTiAgcqcDtUmfUc1ArMse4jnNSqd2CeDTlGwCsuQw7moQmG4qbOAfWkIHbvQnYBwJ445p4I3gU3bxzS8Bxz2qHqMX+KkPzf0pM4YelHQepoEDH34FIMZz605gOuOtIB6HpRoGorZKmqcb7pSmc1bYlh161VtozHK7N1PStKaVncOpZcZGBR9089KAeRxTTyc+lRcB2doFKODnuetN69acenPWkAuBvqMj5icdKc2Q3FJICehpoBEBC89TTpMBQBSKMrnNAwQSKbVmA4jKDNL6e3WkbnHFITgnikMXr9KaVyPrStkAACl7Ak8DrQIYAFPH40JhW4OaDgrx1NP24ANUwEUbnJP5VBdp5kLKKnRcsSetKwyDjpSTs7g9VYo2MBSJWk61YGQrEfhQfliGT0qC2n3yhc5xWzbneQJpCR75Z1c/KFqxIAGx69qacox44NXLW28xxKw4q4RdSSS2Jbsh1jZ+Vl3GSelaC4H1po4pzY3kqMA9BmvYpU1CNjklLmY+NGkdVAyxPSnTSRmNEjGGUne2fvf4YqKT5VAB5qLJB29a0uQ1fQVueKkRNvJpI4znJ61IR2oSGIQDwelV3tI2yQNpNWD16c0nrVdCHFS3RlyWciZK/MvbtVd1bGG79q289KjdVcYIzTUn1MJYaP2THBweTwKlKk7yOVGM1YewR2BHy46CrS20ix+Xb/ADJKuJQR0xVqSMZUZrdGOTgj0FPRWY/Jzn9a0/7JaJEeTJRunHWpFjWMYUYHehy7FU6En8WhTisQGDSckdquAccdKKXBNRdvc6oQjHRABz70duKX3zzSigsaBxzSqpINPxx14ozgEg0bgUJk+cnvUY4TGPwq1IM5OKgZT2qRrYYBg9KOlPxSHoaLDAn8qBwaXGfpTT1phoxxpjtjjNBOBULvmkA2TnnNEY7ilCE1YjTCihIrYEUkA1LxjPpSDgUZ70yRGYAZJwB1NY15eNeKYrf/AFf8T9Kv31vLcqqxy7Uz8wx1FK9pC9qbcjCEDJFCH6mDINsapE7Me2Bge9W4dTENmYxkyDoa0PsMcULCD5HIwGPOKzJNNe3ZTEPMUckHualuxadyxHqEkMe6c7nb7iAcmprTUBdFgV2kVkC5ld2aMfvZTgMf4QOwqKMvbySxO+AT85HWnzdwsdAbmNn2Kw3emar3V3HajLcsfur61iB3hl81FKrn5VNOkWZw11OhKj14ovcOXQeZmvJN0hJA6Iverum25hYyuME9FHanacySRsUiCsvf1q0RtIxSKbtoW5Y0liIZQwI71iXFg0RLR5I9K2kYNFjP4VC3LGlKKkFOrKm9DmpCc4IOc1PCmFz0rX/s5buXlcAdSKZPpkinETAqKycWd9PEQe+hnZJNMdM1a+w3G4/IT74py6fcE4K7fc0uVmrqQ7mf9w5q9ZWxuZAeQg6k1bi01FO6U7j6CrSBUG1BgVcYPqc9XEpK0CUYVQo4A4AqN2wpJp2QQcjmmkA5A9K2R57uUnfc+fWprTBm56U149nbg96YhKNlTyKnZjuacjjn26VETn61EJN/B608dc5qrk20GsAx9KQjB9acemKjkcIm4ngUWBtEF3NtQBeprOOB0qR3MkhfHXt6UhCluvFaJWOGpPndyM5HHT3oweT0J7U5gQfvZ/Cmjrkmgi3QDxgCkwePWlIJx29qQntjgdqSDYOQM9CKaTkfzpx68d6jPJp3BIDyevXvVhB8nPWoIly2TVpRuAoQpPoh/LfeHFP6gDvTQCRjNB4UE1ZnJjTlvc1YuTsjApluqtMD6VDeyckVjUdjtwsNOYy9QmHCHPzVnuCwPOPSnzv5kzOvfgVCFY5Y9BXLbqejFCscDjFNjGTycmkIGCOMfzoi4OM9f0p9CluOdTuGDV20hIJOeD+lQBe5GDV2yAOPr0rOcny6DsX4FOPpV5BkZ/Oq8I+bFWsfMOwrzZu7KEZcCoWTb8q8YqdhkHIqBwSQB0NTFgOGFXNDAkYo2nAAP4Uh37Tg89qYWGKzB8frSsrby2e/Sm8gHI5Ao3Eq7Z4Hara7AKoKEEnvTL0MgBQn3xToDuYk8j1qSdd0ecdKW0gKsGfMGT81WuGYntVZSGl75FWTzjHT1pz3uMrXufsxPenWWPsgAbPrTrkLtO44X0pLdVSIqOnpVJ+5Ym2txXwGVc9akTG4AcZqCXjbj1/KnwkFgSKTXujJQAT7ChiN27vTjgnGOaQY6VmmHQbJ8zfKaUMFJFIoGSKVRtPPpxT6WAD8ygDqKhYhuB2qVSMnmmSYVckcmmtwGgg9RzQp+XBHFOyoUY70KMSD3pgMljYjKngGlByAW6091+bHY1CjBFYEdDTWqEOzx64pQcnmkZsIeOewpUP7oH9KLaDDd8wpJZAis56UMSVJHBFMdRNEVJxjqaaSArWqBmLBcZq7gCPbn6moFVgFJPXjirAULgkdquo9RIYmwTFVpw4J9KhUgXIBH1qcn5xgcetS0NO+wgYuAacrfNg011+UAdM80DAOTUW0AcGy2fSoJjycdfWnxkqSD0NRyYckEYGKqKswGehPep14Ut+QquV3R/SrCf6geoqpbAI/3c5pDJhgcUrgFSCuFPeo15C47etJAS42tjPFOblABUTEu3HAqcNtQVL0AhkGWOORiqofZOcj6VeQAocCoXjTGSPpVwkr2YEyEOvXpTSPkLZpyY2nHU0hbAK45qAANjAI4NPY7eBQoAQD8qQYZsn8qQDjgJ6GmLkxj2pxI2mk48sEdKOgCDJOTSLwSO9A5GP1pQvzEkUAKoA3A8033/WkTncD0pSNqgE9aYDWXH1pw7jvQ3BGKULls0AKcLgdhTR8vPY0/gqaYRmIetJAV2b5GqJOAfWppABj360yML5h9K2T0EywOR9BUDgFcHvU6dDio36CoT1GQMMALjnrSEl1wDUjgIu8ck8VCp2ux9RxWsdhdS1GSYxxQDkYHamrMEgBxmlVlxuFQ0wFbnjuKWMZJOaO2eooXgDJ4NLoMexIOe1RlcSFs089qayjdnPSkgI0YkHPHNSHO4+1GR2603li3tTEEhJQ89aiOD8hFSMSMY6U3GGJqloAg5GM8DrTgoIxjj0psQ5yDk1IMkj0oYDAS7Fe1MI3ZGadyJB3FNxy3HJ6UwHscAY6UAfugMUjAgBehFP6L70mOzIWZlIGOvekO7HWpANy9MGnSJhaq4rETZUAA/hR5u3IP40jPs296ikcFd36U0ri2JgNke7HU04sPxqsLhWwBxVlSGQN60pJrcY4fd96aBgH1NIh+Uk01SWkYEdKSAc+CPcVSO4H5jnmrrHIx29KptiTGKuAEwJUCnEjmmJnOM9OtKvzNnFNoBy8RnjNICVtyf0pPmM3HQUp+aPjjcaQMaQEUAdGpigxt81SuAR83BHQ0x5FbDKM4OKaEwTPmMQPlNC5KnPWnAhNx6gmnIMGhsLCRqUUg8k96D3zSoCBubj2oYAoXB61PUYm0c81KuNuDVcPtBGOanz+7AB5odxE8IHTsKk781DDwAAevWpwAD1rKS1GOHzDrikfg57UDnjPSlf0NQBXlbkHtTJDwD6U6TjFNOCo9K1QXA4Lhs/Wkx8/FROwQ7wc0vmchh0NVZiuOzjCmlAAYfypWGRSA4A96RQ7HPSkIBbP5U7vSA54pCdiq45I6+lJGfnxjgVMUzJ0/GlVfnPpWnNoBEhw2SOP51aTk5qIDc+0dqnThsDvUSdwI5cli3p2pjbSQR1qdlOcGonwCOKIsCKQ8k96oznAIA59KtuSGz61QkfdIynt3remgIYnHmHnn3qx5mQMD5arIoaX61IwZTsH3c8e1btK5N7bhKvfHSo2BAyGz9akkTZj361EV+bpj1NOOxNtbEfVjzRg4OOfpQ2QeD+dGce2PSrDcao259KRwD17U8knIHemMMDJyc9qaK3NSzcmBS3SrvG/gVn2D/6LjsKvIcnj/wDXXNWWtwWiJkzjrzUyYOQR1qupO/ParCYzjpXNsMx79cSqccVLKd2nyEdPLb+VGqEiRWxwPSkY502T/rm38q6Y6xiGzO6HHNHQc0dVH6UYzzXz5QEdDSHpSjpTTzTAd3owCMUcf/Wo5z1/KgA7ADikznr1oUYzQfzoAN2R1poGV3UgwDilC8cU9gF7Yp3PGKbnJB7UoPzUALkZNA6nNJ60o60gDGeMUEdqQHjNLztBzQAZBpR060hFH8qAFHIOaY0QcqT2p/fpQM/hSu0G4bsnFO6jNNGM59aXt70AIR0pqH5icU489BTVzimA48DBpjDLg5p5HI9ajPtinFgQeQPOZu5qZANuPSkGA2aWIcZrR6x1BdiT+dBGeKQHB4pcZ7VkBDJuABU/WjfiL0PvSzEgDioHcs5THHr6VpFXAkBJU88mp0UeXmqao3nE5+XFW0bCc9Kc1ZABGRx1pVAC470jeucU5QMDmswI3UEjPamREnLdD6VJIARUfG3K9M84q47BsPDBjinggLg9aaAMEilGGA9c0rJgLjIPpQCAMEfSlPXHSmt6jqKlAJ0pAcdOSac/MYxUZBU8fnTSDQeT83Wo05OT0PenOhPTt1o3AoDjvTQCY8sZA5pN2ZRx1pzsdmMdaQAK2O9OOoExPINLnBGRzTDktxTz1FZvzAB94ntRjccULznIpP8Aao0ADn6il9qQn5uKUdd36UbgDHb+PamEfKc08gs+SMGkYdsdaemwDVDKgOc+9Sd854xTB1IxxTm68mkwAetIT6UpGHxQeAevvQAHoM96D0z39KM/Lk01icjj8aOoDmU4xUePnJp5PoeaQnBp7AISM8nj0pj4Vc+tP25AzTGBYEU0AoxnNKoIY56UwAOMr261IDwcdapgL97AoYcHFIGIOCKdnOSetQFitOoYbT1psGQSSelOn5LNnFVbJHU4YcZ9a2j8LC+poAfLmlGMelOI4x2pvYDvmsdwGE7lBFI3A96lz83tUZGe9UmAwk52+tPQBE2t1qMk78U8qWPXB71XQB4HA3etOI3EHtSDBOB1xSnjGOtTJNOzDcRsfhSDPJpJD8goLcAAdP1qbAOXsT3owRzQeTign5eOc0WYCk5NNXhxS9Bk9aUABjx+NHQBeC2PzqLI3ZPan5C5akCbifcU0MXoM0YAUGkUArj3p2MjHf1pCAgADP8A+umuOfalJzkdu1HbbigA9SBTecHFLuOdpHWhvlbH6U0AgHBz3pF+4c9zStjHHak7YFPYBwzjFJ/HntSg5FIflHSkgANg8UvO05PFM53cU/jPNAxDkKBSqcDk9elMZjj3qEvu2r/EDwa0jDmEWl+VD70uPlNM3YYoacflwKhoOgyRPMjIBxxyar2dqYZQxNWh6CnwwNLKo7CtKcpL3Y9Q21HJb+eduOM8mtBVCJtHGOgpY/3CELj61HuHpzXr4ah7KOu5y1J82xJjBpwGULZxt5qMyHHNOC+YuK6TIaSWO0fnT0QLT0jCdB9aXt04o3GKBSdOM0h4FITg571ZIHpTc+valJyDik9BQAh9aaeDS5z1o7cUbBuABzVqwuWtbpJRg4PIPcVVWlHWmh7HUXcK3CzWunjzoZUDlAP9WfasWOwmmdo4oyzL1FFjqE9jMXgcrkYb3Fbdh4gSE4eIfPnc460WdvdISszl2TB5pOn+NaWpQ26SA28vmBiSfas5xz/KkNPQbjIxT15OKYD0pwODTGPB20wsB0qRfmAxz61DIfmOKA2I3ySAelMKg9acx646Uwkke1TsMMc+1JgUckjjFLimhjQO1NYdKf60jqdhI9KGgKsjcY701IiTkikiQzSE9QDVtUAwKkeyGpH3PSpRgClIwMU0+1UidwPJxSHoPWimk4xTGKaYzADrSM/HHWoTk89qQyTO9c/lTTnNOHHFKRkUgMC6H2aKW3ZDnduicVYjFoChm2/aCuWzWg2HxuAOPWqzaZBJIXIbcTnrS2KjZFXNlA7SSSedIeg9KikW51OUALtiB4Facen2sTAiPJ75NWVGBwMYoHdFaC3W2h8tOvc+tOI4xUrDk0zvmmLcauUJx0NKDh8UHhaWBS7g44FLUC7AoWM+rcmkdeOalHAqKRh+FX0Ie4zcACDUZf8AOmuxPQUAAnrj60FDSxJ69aFGDTiOgpSOaQDWz2/OjBHSnYpD096AQrAOuPSqrRlD04qznBoOGBFFgKnTGKesm0cnikkQryOlR5xSHuWQ4I61Su5dzGMdutSbtvtWf5mWJ9TVRSMK7tG3cd1bPam5+XHTNJkAc9aUcH1rQ4XdB6EUnJXHb0px5BpCuM4H1o5RJhtzyDjFNYYxmlPGecfhSNkMfTvSNLoYzcZzgdqZ60E9qciljycUrDJY1AUHH1qZeDn0oC4AI/KlI9uatGU7jgM8dM9KRuT1/GlAIWmsDnnvQZk8PyxO56nvWNqU5UOueW4Fa9wwitggOCK5u6l8y4I7L0rlqy1sevQp2iiALgYHI9KeOh/lSE9x1pN2cjNYnWQHITb2pYwSMKdwpThnxjrSxZDj07ir6BsXFTjOOasWkRV8etMTt71ehUbFcjFcs5NJlF6IDIIqcc81FGoznFSscAV58tygPzCoGLADHQVYYZUVD1UnHSnF2FYa+4kY70ucAEU5BgZPNNPyjJHFNdgGkqqnd1qMjIKjvSkAyZJ6+1M2MZcKdoxVxQD4PlZU6+9TSZ8rFV0TZLtPOeasv9z09amW9w3KrKFnGTwetTA5T2qF2LkNjr3oj3dCc1bWgDbuMyHaODTLVfKmZC24nvUkr7QWPLdqjjmAB+XJNXG/LYNCa4UYDY471LCQUJH4UyY/ujkD3oiGEwetZ/ZBkw4yaY2QwGKc5CqB3pjckE1CAQnCjHWnNxjNB4baB0ofBHIquwBt4BxUcn3sEcHoTTweATwKbMMsMjgU1uAwH5MZpw+8uaQYDZpynkuBQIVxvfNQ7eeD9TU+ACWNRL9/HY0LYYmGCdeR3ojz5fNPdeuaQLhAAOad9BDXJ8naep60iIqRlQP1pzbVUihVVUwf51V9BkKMZGZMY2nrVhc4xnNV1cQjgZYnNToxJDYxgc05LTQRESPMB2/U1Kpyh4xUcj7RnHXtUiDdFx681L2uCAnI2kUjgI4GM8c+1LxgMOvemTSfvAB1PWkhjiML9aikOQBjBqbcGAIqOVQdxHamt9QGEEHHanoccDmokOAA3NSQjap4yO1UwJWPJAHHeosYfg9qkBOxs96gfOC3T2pJAKqk5Hc1Og4UegqFRkqRUyDbnP50pAInG4HuaV1+Wmg7Tz+dOdiybhU9QG/dHI+lJLg8gYNOYjGRQcEDBphuCZZMkfMO1LzhscY70g4dz27U6PBy3r2oYCqD5ePzpsXKYPWlySDxTfugc0twAClU8H1pOAvBpRwvX8aGAIMggjrTWztJB6U4HA5HWgj5T+lHUBMnH8qXqwFIRyPanDg7qBi4BLc8U085A7Hind8VGxG1h70IQx1zkdqhTAJqfgJnqagYYIwK0iGhMny8CmyHZjJz/ShCD2xRLgjHejqLoQPIFIUHIqszO5JB2mpJThwvQEVEiknOa3igLMKmSLG7pUgKgbM4qHcEXOfxpAfMYFeveoauFy2v3MetBX+GjaUT3pTyQazGKcjGR0pGIB3HoaCc8GkcAsE7UgE24y3agcvQWw/l5pp3eYcdKYh5OF96jY4Axz609hySRxTTjZz2poYm0bQBxmnMQuKQfcB7Ck4kTHQimCEYbG3Z4pTuIBAokYKqKeaC7BgMfLRqIR88d6BhhjPzCkYYlqMcTM2eO+KpJMCwcAAjqKGGRg0gO404gZ96gZGyA4BNVXQOWVe1XJFBbFNWJUcn+lXGVkLUppa4G4np7VLEsm5lPSrIA2sfWq8JfJzT5nJMB0QwzKfwoIIBb3pVP74g0jnORR1Aa5IQnv61WwAu6rMp+XA6GqsmQqkmrgJ6j1YqykHr1qd/lCntVbcQwxyKs5JXZn8KUkO+g2QOpLA9elMXiEr3FOkLAqAaVQokbdwf50J6AwVcJsPINRRgR/L39KVNzPlu1SOMbSB1609tBdBIgGUgml3FZQMcUjKC4x0HUU4r+85/CkUKVzuG7pSghcACmxcn5hTkBRjnpUsRGQVP3anVR1qMk55p4b5AKHdgSJxjHWp1ALZquTtKH86nBIGe9ZyGPU4bmkGMc0BuCRzR1XpUBYhcH8e1RD0PSppMECoWHBOa0QiBlLKQOualCjygpHIpqDv3p2SWz1yKtgKmTkUqZAINAPJA705epzUsfqBwQCOD60zIDEg4p5BIx2phADY9aEIVuSSDyaMYPND8YAFMEnPNNK4DhnPrUsZ+eoiQWwevrT4zs46mk0ATMQeDUUv3KklwTn8qrPLuU1UUIbIAU61UlUbj64qdhuTPYVAxz3reGg7JlUHDYHNTN82c9PWoUI8z5gOTU4KDgCtnYiwkmDg9+9RMACMDnuKeynJx0z0pqt8xB/GhbA99SN89v1pijGcmnygiQgcjFR55P+c1a2BisQzf4UmCQT+Rp33WO4f/AF6Tttz0pgXNOI8l8jqc1oKoUjHFZmnnlkrUPQH0rCu9Rx8yTkZFSxj5ulQknAx2qVGCkHNcrQ3Yoank/KOoNN6adKM5/dt29jVjUlHlE+tVhj+z5e/yN/Wt6fwoOp3gyOKMYH86B1pcc5614JQmMCkxxmnDoRSfw0gEIPWgN7UbRjOaD7UwAep70dqU/wAqTjGO9ACYyc0p6UgHFGfWmADjrS+v86OpBpduVpBYOg4pM9xS8dO1JwMDtTAMZGKQ9OtOo9DikgDHNAJ4NByD1oxxRsA7PemsSO9KMYpcZbBpARROZATjgHj3qUDOf0pAAOBxTsDORxxVNpsBvagcCgEHIx0pm6kBIMfjUEoG3rU3b3pkgGz1oi7AQshO01LHjbSZAAA6mhOCAK1lqtQHgcUoOeKU4HTrQQB0OayQEUnTDVUkCtNtzj+tXZBlM96yoJGnuyT0HGK3pRbuLrY0lAQClGScDtQMleRzTgMsMVnrfUYjrnApyHqO1BOeaiDEyY7ULVWAkIBVsn8qhjUQR7VHWp9uM4PWmN8x+7wPenFtaB5i4I6/jQPvZ9aVWBzmgDLcVLDQVjgimv1PPWlJyDmmsA2PQUkA48DNNY8A0v3+9RgELtJ5FNAOY8ZPemthsqaeCCPcUz75ytNMCQDIBxxSFcnd6dKEyUKk07G0gd6WzAcRgA0oPAzTTuAA6ih1DpweQaVgBjwSadgheKQgEAdxSAnknv0pACjI64NO7YHemBsMVIoLFWHqTTsA/Ocikznj0pRjJx0pGHpStqAzJLdcVIRhR61GEXjHUU8HcOKbAcxA5pCfyo7jHUUrcjjikBGCQcZpzDK8d6DxijOE/lQAoA2g/rSYzil6gCkyMdaLgKeGz2qM4YZFOzyBTWcRsF9aaTAiUEq2O/al3BMetOZcgjdgU1lGA2MntWmjAmIzinDnIqOJiw54p/Y4rOS1ArXK7oyvTNQQMRIo7/yq5PwMnqKqNv8AOUjpW1OWlg6l3sPrStw1A9TQfUisBiYpjYyAO1PY4471G/3gO5qoiEfKHA5zSpkvntQxyOuaF+UAZ5q4rUQI+fmIwc81KT39aiOVPzDmpAcqMdqus25XYLYZNwAS3FIrK4zSzrvTFCqBGAOvaslblHsP7Y7+tAGARSHt6imxPyD60gJOqYHPtQfujBoUjJowNuT2pWAR+B0o5AFHVeaG4wRQMXOBgdaMdB60YH3qUbQu4/lQIQjAJ70h4PHNKWGz2piALkjr3poBScHHpSZzg0BdxzSHO3B4pgNY5OTxTgf4qYzAMARxTgSR0xQFxVUhx6CnPyDzzUZbEqjOalbG7+dGoDCoJBz060hO7nPfilXHSm4OcDgU1uArA78ZqBoj5gKnCg/nU+R5nHeo84l2HnNXBvoBLgFwx79qUkMw9aYcSA4NPBUMOPmqWrgOjQvLsQZJ6VqwxLEgQfe71nWshguQ7YwePpWsjxt8ykHPevRwUIWcuphWb2GsmcioWiIq3ikKgjFeiYFIggn2qWFscZqR4gelQ7Sjc0rBe5Z6jNNI70Alh70pOFqxaAenSonyWpSxLjHT1oPINGoCA8UCkBGcdacB+dMBM85oxS8UeuaBMQCnCjGKUHFACjApN+MntTgOKaRTAf5hIqNuccUY4ApSOOaTYrDPWj+dKRj6UAjFMYK21ufxqGVwZTtPFPfgY9aYq+tIYwr70u0DJp+OlIeTgUANxxyaPSl7UdeDQAEelJIP3TD1FP8ASmtyvtQIqWQUKyY71Y2hckVXjBSY1bIyKCiM5JpMU4+1N70wGngYqJmP41I2cYqJhzUgtRMZwccUqrk8ClHvTgMUWKI34PtSrypGaV+Tz+VIB27UCGbQBSK46DvT25ye1QnjNA0Sg+tOBFQ9OKdk+tFwYsoK1Go5qRzuGDSAAfSi4CMvHFWLeHYgJ71ArBWBPTNXHYHp07U0gewjNgGoGOc+lSPgkYphwKCUQEYNAzTmHynFR9zzQUPU/nTxyajB/WpAM896NwYh703q3tUhGB0qMD1piAnn2oA9KU8DrTc9BQMSUZFVSOetWnGc1UkOMHvUNAhGBxmstGHIPetKWQLCazQhGO3vVI5671SH4A6n9KTIBBzx9KRiAoGc++Kj5zk9qpSOfls9iyCDznI7cU5sBsZ6U2JMjHrTmyOnQ1Zm0kN2nPB4pjAYJBqRSCfSkdetAlpuV8Et61YiHyYA5NQouX4NW0wOccH9KSLbshcY9xTck5pzMAvB+lCEcDpnvVGO7HjA+6ee1Ea7pkH1poHzZx+tSwfKjyH8DSk7I0pR5popapcbQwHWsAL3A59T2q5qEwluNuehqm3XAbj6VxN3Z7UNEH16Ui4PPrQwOetAypAHfpSK6kZADAAYqSIZbk80zb82CcntVmNPurn8abY7GhHHuAUj5atLE6J8p4FJApKhWq1jcmK8+cnexRLEQ2MHipGOBwKhhTYcr0qbluRXNLceoN96mOAeM8Cn4JI5/CmYw3WkgIkHzHJpQOME5ApxUItBI3BVHFXuC2ImY7wueaJd29dv3ae5wwxyaHVmYA8Y6U7gGAyBjyKkJygOKiy2MAdKlU/u/elICJkw+TxSrEAMn8KRyTP/ALNPdi0lPWwFSeNll+7gUxQVkJXAz0p9zMImGTkmoYWMjA7ePrW0L21E7FkHMbNmlibbjrnvTQm1Su760qbd4Ge1S7WYywV3cU1uFyKUc5pqxnaTms15gDY8wY70p+8aRuSDjgClPIznrQA1gNm3pilkGR3ye1JyVyeooz+7HrTAjU4JFKOOOxpp+Vh64pSpY8VTEKBw2aYpIkXsKkYZYdtv61GcFxQgHMSQT1xTYySzc05wdvB6dqiUkPt9f0prVAPKghifxpkT7wzEfhT1BAIxkmlChEJ9e9NPQCNnU/KoyfWpkBKg9u9VzKmNi/epyS7wAGyPpVOOgXEnd4wVVcj1qZDsiz370yXA4p6EMQCOtQ9hir93djp2qvLhGDE8n9anZuq9qgfa5Ue9EdwQ2FSHOT9BU2Pl68GmKrBznsetK6jZ9DTe4yNwQTg9KfG259pPFNcZG4cZoVgAhzmrWwiZlyx56VFK21vr1qQgsxZTg9xTDyQDUIBI1Kp83IqcMXUDpiqyORJknj0qyCCgxRMBrE5xT/uqR1BprcyfUUu7J2jtUhsCjrzUfHIp4Pzse1ISGVjQgAfNnHSnEBTwaSPKoPWg5JB9KOoDidrBe1IMEkGlIwTzSL03YpANwNxBpCQFx3oHMhPUU4KQuaYC4yuM0HoATxTCxHbrTxxjPbtQIDxgelKc54/ClAyeehpDnNSMUcPjtUJGG6cVKxw38qjf5ozmqiA0tk81EwJkFTAZUjvULgg5yciriARgtNuHCmpJQv3j0qOOTLlBUknMfXFD3EZ7jM2c8elPUEfU1HISbhRjiplXBNdGtkLUbLGzRfKeafaHC4x0p4AJJPSkQ7RkGpburAW2y46cVFuwRjtT0bzAV9OlNzhiMc1gtNChQd20AUO4X5qZvAfinY3qwwDmnYQxArP5g6HvTlwXPcUIu2ML3pwwOnUdabGNJwp9KbgbDngmn8Y55pvVvakheRGpIQg0DHlnB/CgE72pQNicHrVgNkYcKRnFNSQtJsPAp0iqBv60ijCeYF5NNWsHUdyXyfXFJINjAKKXIJBPWlCky7s0tgFyUbd2pyjg5P0qMjL896kBwCKTGJ97mg9R70mfl20N91eOfWkKwjr8uKgc7oxg45qw5C4z1IqJsFSmOauICIeSetNfOeOT6UD5UOOcUp5Ut1xT6j3B8YKjrVZ1yRk9KkXcY/MzlqY7fIQvUVotCNxVzuxjFTFgzY/WoEIByetSDkkDr2pMYrYabaDz2NDpufk5NKFyN3fFIpKozNS9BjQw2MMcdqcziOJcnOaQOuwuFwaUhNi7+uKdgTEVV3iTPHenOv7wPn5abtUqAOjdKc27aoHWk9xWHBSJDjpQoJOSaBy5zx/WlY/JuzU3Huhu75tuKfn1qtNL5bkimM8ruW6DtVKN9QvqXHmXjJ6VWuL5tqhO3emfZXl+fOQasiyXaobqKFyR3Jd2LBJLLleh9augkjntTY0VQMdqk6nIH41jKSb0RUU0MYZWq0jEAgDNWj2NRFPnxiiLsD2ItuEpDnIA71Jg+lMYEMoHSqTFa44/KMntSryC1NyG+WnKMHH60hjFZgp4zTdzqpJXmpl4ySeaVsFelO4rFYy/vAAuc1VeSbecJj3zWi6jOR2poUZbirjJLoFmzN8+SKQlgTu6VKLqTHzLirJiQ9RQYlkIXHSqc4voTZrqSAiVOB071UmO3Cgcmrm1YyFTr71RnBDkt61MNyhGH7vA7VWLZHHB96sSHByDwagB3MRW8RlfGJSMcCrFuoJHGTjrVZgBNjHWrUL4J449M1pJuwna4yRFWRjnB7e1Q9CRjr1qSVFJ3Bsn6VEMqtNbE9QkOAf7x71ErdcjGKlfng9fWoVwoq0hMeXyOBTGIDZIJpcf3ex7UmCDj35oQySxYi4OCcnvWzGdwxWNZ8TEZwa1wVQZPWsq2ug09SfGADjkUpOQPSmMdwBp0YDJg81y2uNjLs5hY/zqhE4fT5ypzlD+HBrSlVfLYEZGKy4lC2dwB/cP8jW1Pb5g90ehYPXvRnGM0Zpe2a8AoX0pp5HrS5zR0FACEZIoI5HajnP0ozxTAQsKQdc011OOKcBhQM0xh0H40Be9L2Hak5PA6UCEU54qTHA9KaFA6cUvQUtAAHIJ9KUDOTTR8oz+Zp38NDAQ8cUdOKBxS4GeTxQF7CHk0h5FOP3gc0hGTQADHAp5yDTMcAjn2pzHAAod0AAYNKOuT3oOcCgngDtUgRvlWJGaaCCcd6fMe65461WhZjKcjFarVAWAeKRhkdetLk5wR9KRh6VOlwI1Xng8U6M/Nn9KXgJSR9j3zVvVATYGcZppxupQCTTcgMAeKzACDg+lVUtljkZ1FWyDk0EY6VSk47BoNyWUkHioY2O454xTwWCkEZOahdgX5PJqoq4E+7jJpCCV3KetAxtIFKikEEdPSkA/H7sZPzVWa5ZZCgTPvVv+HPeonyr528etKDV9RMd1GPXrQAVajjHvS7fl5pDAnCk0zJC5xg07dhSKbkk5PHFNAOUgEcfhUbjBJ706Nhvye9LwTijZgInTAoAxzilPXAND/cUg/WgCNsl+PyqYEO/PaoyfnGPzqUcYo9RiITnkcUqoNxweKQfLmnHjpUiEUYZvalB5pRkLkUAAn0AoYDWGF3d+1BIGDjmlY+nNIvzE+9AAFJBx2pYsqpzzSjnNHc5o2AReC2O9CYVTnilAxk96buAoAUDBJHShQcZNGOAKB0AHagBGBJwKByQPxpexpqc/WhOwC8jk0hPygmlHIxSHOf8AZoAT0pAFJwecUZDAEdBQrdzwKr0AYeWwTSueATSFQrhifpRjIG5uBT0AlAwoAA57UhOBjPNIMlM5yaFOTtPakASYZfes9zIJcrwtaBAwfaqfIJ/StaTsBcQ4QeuKc/PTtUEJYoPWpyAV2g81lJWkAgPH171G5xHnHNPAOPpSMA3FNLUGRhjt6c04DLAntUigbCCKQjAGea6KsIwiu7JUriuCV560dFGOPWl6nHpSkfL9K5XJ9SthjjK/L3pqcfnTmbao45pSu0DuaYCBsscihRytOPHGeTSIACcjnFAbCk8kA4pW7KetJjBxSMMnJPSkADuCOaDzjH5UqkHOaQcE0dQAZYYHamuwC+1Lt9DgmkKApsY07oB3Hljb+dMAwmM8nvUihfLA9KMAjPpSuAwbgOetDcoP1p45wCKQ8AigOhXZmLA7eKmHIGDVd1cYA/HNTrny1ycAVpJKwCbQX3GnkgISep6mmvwpwOKcRxg1DAjyQucc0EHO6nsDk0xhgrzxVXuAvIGe9NcDO89e1SEHb1qNwzLhRk+9C3AbAd6tjpTyM4PekhUKWHenDLE54xVSetwFdiGFP854kBDkU0jpkdKVvm+lEZuOqBlq31Fjgt92ryXcb9GGaxA21gijC05TgkZrsp42a+JXMpUkzfDBhxTSVIGRnFY8dzNESA2QOuatxapEx2yjYa7qWKp1PUxlSaLjOAMKmKjKluT0qVBHIu9JFNIV/KulbGdiPHoKNpx1qQLjNGMUAMCjrnpS4znFKR0pMc0AHr6UDgYo5o+tMQmfSjNITRQMkUZ605hxTF4HWnZ+T3oEMdtiFqbv+Xn0HNOcjbzjNQk+g+tTqMc8mfpTN+cGkC5GTSYyTVIBTJnt0pu7I4pSO9GOaBgDkc9aXI6UmOelBoEL2oHFNz3FJmmBIWGDTSeDTRSZ9aQDCvzg1JIRwM/Wo3cjgUInHPNIY/Pp1pp60/AzSEVQETL1qPvU7DODUTDBzUh5CD/9VOHv0pueaUntR0GIeSD6Uhpf0pMdKAEPI61G6+34VJ0pw5BHrR0ArqM8U7tweKcUCnI703GCaEhtiE9u1APQZpaOMHIpANf0p0EvYnio35xTBwcii9mFkaAwRmo3HFNhk3fWpGGTVE7EPVfeoieeOgqV+Afeou49KRSHVMmCBUA9KnXCr7UxMbIccCk7cc0uCxJoPGf0piGkDrTcd6djkUh60MYx+AeeapOSWzVpzxzVR+DWbZSK87khU/vdah3bhgc5pJ2zMc9Fpvt6Vcdjjqu8rjT2PWkX5m5HApXyeO9S28RHzVS3M72RMo9uaGBI9aTIHHSnY4BP/wCutDBtsiZeGpPM28HmnsRzxmmFMgdqkpPsSRru+6eDUy5UH1qrGxU4bp6+lWWA9ePWmtRTYwg5yOaFOOaCCCMik7ZB/SmRYkzlgPfp6U+9cW9jhjg45pkA3SoOmap69ch2WEH8KyquysdmFjdtmIxZyxJ5PWk3H15pcd8f/XpMgnI71ynqrQFAwPm6Uo+UkY4pGyD0qRV5JApMFYjAAc+9XIVG0nHFQtHlgRwasQqB8gPJqXIaWpoxtuVcVdXLYGMcVn2vCn0WtGL51B7VwVVZlbkicD2pc7Rn1pisCxHpTxggisGAcdc9KYeWzmnALhqYSQRTQCnL02TgA0o4B9aD93A/KmA0nBz3pz5KgjtSFjj5RyO1KRlAKYyEsTH8o5qaM5wfaoSwACY5bpUqFdu0HkU3sIY43NnHTrTj94KD9TSycKAOaRXwMKeT3o6AVp7YTXAXIwOtPhhETMv8PrUnl7X3DkkUgXcABVcztYBowshJ702NfnyeMd6exBBHU1GCNuCeTTWwFgHLZA+Wmk4GB1pV+77UKMN9KzAZu2rkUwvmQKp6dak7cjOTUZ+WTGMe9WmBKpDnBOB2oOdhz+VKVx0objHoanqHUrgcsx59KerbQD3NRv8AIdx6elSKASSat7AKyDfz0qNc96kIH1zTHGAF64pIQ7uR371CExKOeTUoXAyTyRUQAzknoapAyZhjaO560hBxg9D1pxYMwPWo95J56etLcZVlRYMbRnPapYFATc3btUjIC8ZPXmkZ0TPPXtWl20IV8MRzSqw3Hb19aQp8wOeMUiEknI+lTbQY6VQeSeKgBABOO/FWH+ZiPSq8339q9qIvoMlDjYMDr1ofkkHvUfQdM+1OzuX09qLARnrtpI8jAxxQwy+4jAohHzMc5Har6CsSocytRIoDg+1LbkByCOtIwOSCvPpU9QIAT5gJHQ4Jq0pAGB0qvHnhG+b3qeIhl+nenMQ8gbgKF45pMfMT60jEgbcVmNjkA3MP1oxwR0pucHg9Kee9DAi6Hd0qVz8oH51C2VYYPHenkb364A7U2uoIeO470E8AA0Ej/wCvQVzkipAYxw+PXpT25GDTJBu24605jznGabENZcODninpjH0prc4ApxIBAFD1AUMd34U0EZPrSqRvNIeD7mkMc4zimkfMTmpGPz461C52qTQgE3ZY46e1QNwakHTAqNh+9yx4rVBsMjUiRsHnNTzAmLHeoVJDk44NWG+570S3QbGfNlXRutSjgfWkmA+vpT+ij9BW3Qlig80n06UvQnNNHGTUtAOj3h89AandeBzzUUUgKg4qYHqcVnK9ykVpgUQ1NCMQjmkkXdGQw60sYCQndQ3dB1FUBEJpVG3DE9ajRtyEkcGl3FjjHApNAObGCaRsAAj8qXOeO1IfvZ6ChB6EbHa4B705AFiIpGw0g9RRuDNj0qughHODsoJAOw8j0pWyZBUZAaTK9aEA4M2WyOnShTliwpByretAO1OO1MbHxnIBIp5Ix9KbkADjrSr0qWhDX+8Mc0jfe3E4ApxHGB3700fNFzTGByzA54oGDIfWkTmP2oQ/MAecUxXIjHtBUHk0YOwr1qWReeOtRIMA7jzVc10PRDFPzhRULHy1bPXNSqv704NQzKSp3dPatY7kvYcjBqs7dqkjriqcSbFFWsngnoegqZLXQa2G7S6DB5FKGPKEZppby3AycHilA2HJPSkCVgRlMbAcUrAyIpbqOoqPJZCgxmldjsVQMUWAH3AqF6dxT2bOMDn0pGUv35HWms4BUqc4600LqJG0jOd3WkhSWRmEnQ1Ju8uQY5FPTJm+U8d6TYNEa2v747+1WPKQY46cGn55z3HXNIR6n6VDk2OyHqMIAOlObnHYU0EkH1pwzngVmxgQQ2O1Sfd47U3qRTiNxNSwY1hyO1IfvGnEc4zSHAwe9AEWSGIqN8lTgcipGXElMP0GK0QhqHd8x4pTnc2TSkHjHSlYEKMU2Go5R6+tKBxSKMfjS52r1/GpYwIB5znFR544/KlZhnHaq7SbXxngVSQgZ8sFHUd6mHygMOTUGAE3dz0p2c45q2gJQw5JHNVJ8MxAqzgEc9arcbjg0R3F0GMoUYP4VCSc5AqaQY2+vSoZAQMjr61tEeiKkhHmnjmrUZBGB+FVJuJetTxnAHHXvWstiU9RxQ7snv61Fgg8/iKnkwwGec1WaQl8Y+h/pRG7FuI45+9UROFwckHvTmAJOT1ppBDY6+1aIEAODjr7HtSnjnOcdaYRk55/GnAZwQeKA6D7STN106itaEDkY/Ose3wk3OMVrxnLNjpUVrcugfaLBB2YHSljAC+vvSDG3mnD0Fcl2imDAmPA9KyVOYLnJ52H+RrVR2+YHp61klsfaB1zGx/Q1vT1bJe6PQx1I7U4Y70gGc+lOGMV88zQD06UhweKXqKTjNIBMZ/Gk7EClH60YJ5qgEUgZ9KUZOfSkHU4oyOnegOgdc0i4BwKUjAFJg7uOvrQA7ODxQRnmjHGe4oI4FABjIxQRmkJ4z60L0NADsDGaM4GelABwaXrSAQkAAYoBzx2pCcn6UqsSaYB0OKUDDHdRyaUZLUrgID1o5I9KXIPtQ3oDSARgMc9ajZcDIFSNzgdvWkcEkCqQDc5ySOaQn5eeBT8ZxkU11AFFwI8rk96FPp2pBGQ7HPFGdpNaWAmYYXKnmkKgkE9RTk+6M9KGYK3Wo6gGSKQ9D/Ondeh5pMcHNSMjbCr6Cqa4YsRzirTnduA6VWSMorEDmt6duoia2LNEQ45qUZVjSxHKBSOQKc2Q31qJP3gF4xj+dMkLBfWnvggY601ueCPpUJgIhyB2JFKT8tNypOB1p/BWqktdAI+hH0pkjYX5qdIxyAKbt3BlbpTXdh0Hbctu7YpyjcfcUoUAAHrScqKGxjVXaxyadgbeTSHA+b1p2ARzSERoxMhBHA71KeWwKhLbD+PNS5wc9qGgHgZ5HWhehJPApFbjjqajbJcKTx3qUtbATKOPajHc0YwSBSsQSBjFIBO2aauSc96c3HPag8GgAHU/ShgTwKXk/jQOOn50IBCeRzxSHHWjmlbgYNAB6HpSCnEEjr700H5KADAJOOvekB2se5py45NIB81MBF9O9BHyn0oHrQelF9QIlACFemaUEYwegpSMEnPSo3IJAq0Ay6ycBf51MmNoDdQKgldxKoxle5p+4iQk8AjirtokGg9chuuaXGKauc1Lgg896zegB/AapT8dOM96txggvn8Kq3KGVCB1qoaMOgtmAq7c596tZ+fGKpQbQqgHO3rVxWGN1OonzXYCNxkUv3V9zQRk7vSngZ5NQmA3ovPWmP/AA+gNPGMnNNP3/amnrqIc2AcDvS85A65pMAEUuepqWMG5JzQuc4z1pT1pFIOfahABxkLS4y3HSkyVPI5NA4NIY4ngk0wjK47U5jngDikOAadhCKP4f1oXIBxS9KQd/SgBenU0hwx57UNyQKTBUZz160JAO7Yz9aTOOaDww560HsB+FAIXOelI3cU7GDx0pFHUk0ICu6fMCT0606PBQ80rfxcUKAEPHFaXutQFyeP7ppxPzY7U1TuVTTv4hmpegXEzjg0xgSTmn4OTmmsuc88ZoQdBx4I/rTRy+O9JJlgAD0pyncAetPpoAuzO45pMZWng44FRBgZCKSbYEmTtHtS5wDxTWzil6DApN3Aapwp3c0Ah2z3pcAHFNjJ3Yb8KrfUBZDjIrOmys7HqSPyrRC89c1A6iXcpXkVdOVncTRn21/PHcN5bMB9a0D4gkiZQV3g9+lY0ETxTkO+cH0p0jc4UDk5zXqxqtWsRyprVHTQa0knVSD71aXUIm43jPvXL2z5zt5z61OU34YnJFJ4pxlsS6KZ0wnDZxil80celcbcTPEoKOysD601NavIM5beCOM9q6IVlNXIdJna71NG4etctD4kJdVeE89SDWjDrFvK+3ftb3rXmIdNo1/TNL3zVRLgFSVII+tSCbIwelCaJcWi4qkrntUZb3pst68kMUIUBUzggdc1CHOOe9Vo0KzJWIz1pDg1GW9aXOaAsOyCaSm55pM+/PahASH/ACaTNMDZozg0AOPoaBzjNM7daN3GKAFJ5ppNIWz9KQtmkMU/yozSU3sDTuAoAzk07cBUbHimlqWgExYdaYZOaj6mlwKLjsO8wk07buBOKI03HpVjZgYp6idiieDzSZyRipZV5qNRQFxT0pD9Kdn1pMfjnvSGNPtSrxml6c+tJwKYCk5pnUmhmyfT0pvrg8Uhi7c80jLT8kDGeKUDimFyuwwOlMI69hVllyKidOOOtTYCNWKnNWo5A4qqQRQMocrQNonkGKhByDnpT3lDj3pm32piRInbNSZ4GDTUPy5z0p6jgk1SBhgYFIehxSkc80dR0oJGsKjbipT1qN/ekMgfr7VSun8tCT+dXGIAyayLqbzHwp+UUnqKUuVFfJLHn8akBGNvf1qPpj09KmRcr057043OWSG4ycn86tQfKuM8mq5+UnHUdTUynPB4NXEyaurEjISPpTSMLg9akXkcH/61JIpPGOlWZX6ELEMaUEgU1/UdaVeRn9KmxVxHXJyKtWrCZfIcgFfun1qsTyf50KSj7h1FNOzBWe5O6FWx0IqLPH0q1Id6hiMnvVSQheD+dO5FtbE9u21XkPRRXPXU5luHc846Vu3Enl2uzGCetc65xI2K5Kruz1sPDljYaDn696DxmkJIPoBSDGOnH1rOx1J3Q4De3r6VbiTkjFQRg4X1q9GpBHv1rObGkQyj5we3rRGfn/rUsq4UHHAqJRvXJOPekndAy9BxEfU1ftjletZ8Iyi8Z960YfvAegrlq7lIeAvmdKlXjPFMAwpxTlDDgVysBACCTj6VGy/xE9OlSfNj+dN6nFNMBFIJyaXqeelNByMY6UN8+McY60+oDXOzgc5p44jwetRlgZevNSAg9TTYyDG08nr+lSx/dPr3pmCegpyblQ5GTTeqEOkPQCmHaPlFPYgrnFMB79AO1JAISxQjvSLmMhTT2GCSPyqIgB92OapAOO7Zx3NRFVVwDzUrqWKqpx61FIo3g9TVICVDlDT1PDGmgA4xQWxnPSoeoC9qSQcdMmmM5KgAc56VIx9egptWAEPHPalbIUEUDBCjtQD8x9KQEEvK4PShfmBAGCKQsHcCnDLSZA4rToCHA4QjvTW5XFOVcNio8ZkB7YqQHEqFAPWo3HzD9af93nvTSSWqhEhACYxULcMpJx7VOeufTtUUkQaUMfyoi1cdhZWwnPWqgi53McsfWrrIWP1qvOB5AAPNXF9hDpGAUMOcd6SF2b5yMUkQLQlDT4iQChGAKNk0MmIyOD+NVJflcHPPerStklcfSq8q54NRDRjQIx5BNSEBshTUacHpT0OGOfzpsCGYkJ6YNNQ/IzZ47UrYKtt61FESBg1qloJ6lyMZkXFG5iXBGPSo0cBwxNTEgZYDOazYEKgKvB+Y9DUkJIb8OajcEJuA6VIoBQHuetD2EOkHy+npQc7+eh6U2cEpn0p6/NEMHOKnoManMnPSpFXk5NRqcsTUwHOT3pMCD+LaRTzhJSQeKRh8xIpduUXPJp3AevGc0hyAcUq8g+1ISdpHvU9QEK5XI6g03gEe/an8EE005JA/WmgFJycDqKX+IkjpSKPmJ7U9eQc9aTAYM7s+lKATQMDOaVclSaAFBG0evemsoIHrSj7pB70v8LetAkVCfmIpshwVX1obnoeaVsNCD0IrZDIW3I4XtVlQGPWqRmMhCISTWggwg7mnO6QkVJflB70KCSOecU65TkioUk3HOcj0q46rQHuSjBySeaUKTmhPuHvToQDESeoonohIjtEYEhjxVhjtYjPFJENq5x1pexNZSd2Uh7gFKjXgY9BzUh6Uw9Dk4NSgRAXBG4dKdFJ5iZyRUc3yNjPXtRDwuFPIrVrQnUtD7p9TTDnbxT8Ejig4C4zWZVyM84IHXvTCwEhqXuFqMrgOapC2AZYvtoTiMnHPpQMqoxn5utOXKAA/jTYdRir196YilefWpRy5zxUUfGc9qLgyRWywJpwO4Go0BwfQVJHwxH5UmA5ulNIwmKc3JJzTGByDSQCgbCAKRgARj8qeGJPI+lMA5IJoQdBzgZPtVIMxxk1dyG5xgVWMQAJFVBpbhuQxsWkbHGe9J8uTz9RSqDFKec5pyoUf6962uhEUZPIqcdhjpVdEIYkj8BUyZ39aJBuLIFaTBFDhPvk/d60shG4ADrTWAVuvHcetSgGl1VC6jmhHIQZOCaSXCrgDrSeUrxKCcVWlg6i5KMXLcmlACgsPmBp0oCxJ6etJHgoSOMdqL6XAIDvYntViE5YjFRRR4O5TkCrCKFIrObQ0tBVA645NDYPB60i5XPH0oPDfWo6gSDgYpFJA9qccDlabtJB7ipGh6gKc9qdn5jUaD5cU8HAx3PekxDj0570xjhc96dztpp6UkBG+Rn1NRRZPBqZhlc+9RJGQ5YmtE9A6jlHzkU4csDTe3Xk0owBzQMdt+bNMkYAHnmpk5GCKgu4i0eTkGlF62YiMNmLJGTUEzLtUEcipI23Dj06VTeTfJjHTrmt4x1Eyyr7nK+nQVOMcnvVZFJO6rL/dznmpluPUY7E9OlVgQrZqduQOaY4wADVRExkgyhY5qDeASMVbYfKcVA2ARkcVcWMo3JUYODT4SAAMY9TTbrqN3IpLfAXJPFb/AGSXuTOc1XIO8AdPX0qy3zLjOWFRgbTjA5pRdhMYQAOMcVByGJ4B7VO+AeOneonHJwfYGriKxHkk85/xpwGOQelDcEdcjtSZ4AwRVXuNiAgT/e71sQnd24xWOfvBiM81rW2Nq8fhUVdhtlvKsp5zUkfQYqJQuPT0p8RwTXFJDQ8ICc5xWTc4WSYZz8jfyrYBOMise7JWWYY+8h/lW1F+8LoehdKcD0qMAk80/PH0rwGixw6U09+KUZx05NBJpANzmgCjHPWkB5qgDkc0pG5qTPJHTFGDnINABzigfzoJzSnAXPegA2gDjmkJ6U7oBTGFAxPSnLx06UuMnijnIwKAHDOMigdDSdKdkYHvSEN6DnrQBigjkUAkkmgAGdvHNKpyCTwaYcmQYp5zzxT6DAsCMmgk4yetNY7gBS4wpGaBCkZ5NI2S4HrTx93qKaeB70gExluvSkZcsMHmk/iHNPHTIFN6AQDdvO7ml43EYp5AAyaicsG4HB71oncVydW3DpTWUORkUyF859KerqWKg81LTuMd0bFIxG0k07HUdzTWAIx2qQ6ERTCDbSFTlTnpTgWDBR92myHDj0rRAPjkPmMCOKkYYwQahVv3mOgNSvwoHSlLdADEFhzRgk5B6Um7LYx260nTGW5NTYAQckE80jcg04kKMHrUbYIBzTWoDX+8BilXoQx6GlZS2Mdacq9zTvoA8qdu4HrTWBxgmkIPy4NBYmpsMQnouPpThzye1MJ2pnqRT04j96roIjIXIz3p5wQUx0qKUN8uOcGnOadgHqcEe1KgBG71pqj5M96dHnHHSpYEvGKQ9RSgcZoHzAe1QMRiOnUUGlB3Z9qTO4YxQIUYwDSKcDnvQcDHFBHIPpSC4ZPT1pCCRz0oJ9Oppc5GOtMB3JYU3sQO9O4A4600/d/rRcAwAuBR0B9aFPB4zQccmgBMAgUrD5aQcn6UdWA7UANAwTkdarTBirKnX1q2x+f6Um3kDtVqVtQ3KEaPtxI30FSxkuMt26U9kBBI/CmDEca5OQK15r6hcnUAZx1zTk569ajyGCsp6085zxxWTT6gO/iJxVZ87mI6dqsAE7gDVa5IVVJPFOK10Ar24InlVj+lXAwYDFU/NQ7gGwfXFWoBtQE1tNaXYLsTevPSnDoDTOpAp55GK5mAwEFs005wTUjADAprL/D3poBoYMikHNO5LYFRyZRlHYVN6Y796uVt0A48fjSKOp7UjH8zQAQuOtZ2DcM7jn0pQSTTd2047UvbcKdgFIyMnrSY3UvBGaaDnil6ALjNIB2pQcODTzweB1ouG4zGcmm/wknvTjwnWm5yuBTAcO5I5xSAkjPpS+9MJwvFCGOHtQRn6UKckelK3fFAEcgCnn9KQNkfL27U5vve5pqNjO0YHerS0ELGCFJPWnJknd6UDkMKFPXHpS1YbCEFj1picggnjtT/AOEmkAA4ppgIBs5Y+1KDtPApp55JOPpRkcGjcB+4g03aVJ9T39KfwQRjHtUchHynvmktwHnJj96FwU96XkikzgYHNCAC2QMdaaQMjnoKHHzjFJuPmlewFNbDJEwcgDrUBJV39McVYB21Tmfa7gnA7UQV2J7mPsl+0tk5BPU09fkmZCPxpiZ890VuRzmpmZzFk9upr029EiY7DI5ZVlPGcHpV6PKggVQTCqAp3E9T61aiLFnVu3SsqkdBoS4iEkZB4asa5DKxU9V6Y71vsARjvWVdRjz8+varw07OwpLsVYnBxkcg81e8szFdpxtqNLdd+X6MOBSQSFXKj6c1tN3fugWmkaDbtYrn3qaPUbgSYLce9QxxneXds57VHNks20/hURqSTsDiaSa1shMkicD0NW7fW7S4YKJAGP8ACeK5mdM7I85zTDb7HUDqOpFdEKjtqZuCbO3EqtzkU/dnHNc/byNHGArYPpmrsd4VUbxk9yKSxEU7Ml0rbGpuJozkc1QF/HkDdjNWUlVhwwPHat41ItaGbgybNB5FR7/mwaXfxjNXe4tdyTvTPWgMPWjNCCwmcUHnvSZzSd6GAE9aQnrQWxTD04oAXPFIM5pR1yKcOT7UncYgXOKkRNxoVc1bRAo+tNEgibBz1oJwppaZIcd6rYlleTrUPT2qWTpUBNSy0KOme9GelMBJp4Ukc0BYXPNNbr0qQLjtTWXigCE+lAHQ0hOCeKcoxQiiQCnDpSAAU4dKZI3GQPWmkA9qfkU36daLARMlRMv5Va7VGVznik0NMr471KOVHPIpjKVb2oyFQnpS0GAfEmO3cVaU/lWfGctk9T0q9GeOuaaeoPYeRSYz9Kdn0pD0qiBh9ahkIUZNTOQAT271k3Vx5rEA4T+dBM6ihuQ3NwZiVThB+tU24Y+tTbsk+lQyY3dc0nojmu5SuxgAGKsQjknNQ4zg4xmpI2GQDwCelKLKkSumBnHHrmmMcd+T7U55AznPT0pYIxJID271aMmWkXCjP3h2qNyADmnySbFyOKqhnkY+i9Ksx+J3HHacZNJg9ugp6wmQhcYzQflcqcHHX3palX0IwQBz2pQeeOPSl5Y4ApVjJwSKYJk4yoz+XNUifNuAvbvVmV1WIL1Aqpa58xnPU9Kic7Oxph4c8yS/blV9uaw3ILsfetPUJAXYk8KBWXyw9jXI3eTPWggHqOp/Sg8HGMe9GcdKRcnOetI0ZYgB3Zx0q3FnqeKhhTIyOnerUaMeAM4rCbVwEmHynGcd6gRccE5z0rQMYMZqoF/eegqac9GPcli5bANXkU7+DxVSMEyjHBq+uFYAd6yqSGh6nc+BxTlJJB71GnyyHHXvThgNjua5mMlYbeKhOMZ96fISCOaaQFHJzmkguIvQ89aQnIPtRx07+lIVxGfWqsIjCAncTgVMgG3PrVeIOFO7nnipYhtbB6GrltYBDlVI70qswHT5j1pkm8Sb+wpYVcHdIeT0HpRbS4IkIbnJz7U3apJPrTgDnJPXr70mdrHFSArHA96jZFwGb1p5U8GmOpZtuOtND6CvtTJx1qsxAO49T0qZ12vswfrUUg+fHQDpVxsJliIYBY0MMr0zzRFjaaePuYx1qJbgV3bHI71IDuUN3pCgwTjpSBsnOOKe4Dl+Yj2oIAHHekBCxjvTmb5SR0pah0ITgSYU/U0MfmAHTNLgAM4H4U1fu55zWgEhI3nHU0N1A796aMZBJ5px/wBZ/WpAYQc460O2ZAB0pyj96cnNNPXPpTAUdjnrTbglXDD8qRMsBnORTpSTxt4BprRh0HDkk47VVSFnLMT0q0CVx6VXmfMixp071Ub3shbEicDA70xMLOdzdR0qRSF3L36VXaPZOCDn1oSu9Q2LOCeB1pjLlqUsUIx3pZPlAIqBjApQFjTD90NzUzncmO1Qk/IR3q4sCKf5WHHymhFUAYoZhJlSOR0NIAUUitOlgHjhs+tSg4iHoahXnCnoOlSswJOeAOKloCPzXzhRkUqEpIob0PNKq/Kc9R0oCnhm6jvRdbAWHwQOOtRxjy0x61IxygPYUxlyM/nWaAaj4Zk7VYUdOeKhOFYMO9SAjG6iWoxJMhQRSjqhx1pxUNgEcjmmo27FT0ECDEh54NKMBzRj5eOuaQ9N1ADVO0EUFvLI96PvAUEDd7Gq0vqA4jqKdnAA/WjdycimqeT71ICMQq896QH5R70SYKnPaoVfIX2qkroCwOCBSsc59aYGGQx/KgnBNTYCjKSCAp4zzUzxh4gvaq84KSL9an6rmt3shDYolRiQMcdamt24OetRtIyrkDk06BuDnrSldq7GRTyHO3FV9m1xg81ZuFycioEfkFgM1pB6aEtE46+9PjOxmXHDVGSMjHSpVGWUr0HWpltqBIACcGmvwf60rEg5FLyxGRWWxQoGFqCWRB8jfeNToTgg1Tul3spHB9aqCu9RPYZO5dkIHFSQp5ZOO9R+WSEU9u9WIxuz7da0k7KyAmAOcA0vQYI6UxWJbJNODcZ71ixkJyGx3FPIAQgGnNhXzURcKTmq3AONvHUUqjdgk5HeoGfJBFEDnew7elXysTJiTuyKiIGDubk9KkblSagkXODnpRFAx4G1Rg8nvT15cmgxhkGOnekwVUenejcLD1+Y9elKTvA9RSKNik+ppASpyRUsYEkkY4pXGDuB5pI3DKcCl4b5R2o2ESMMx8Go3AKgYqQYJI9KjY5U9+eKSGUpmJl2gfdpzLkqxPSlmJSUY9OaYTtB3VuttCbKwKRlhnr1p8a7RnqaiVs5AH0qWMko3cihjHErs3Ht0pNgZ8scg0q8qBjilLHzNgHyip1ATajrgDOKaYw4BB6VJFkliy8U1h5atgdeuKLu4+g2Vl2qjDIaq/mBZQgPBq3KiMi5H4VSCjzCRwRVw1Jd7lxCIjtJ61OCS+DVKOUSj5ecVZUnyw2eaia7hoTgflRtHBFIoOwgnrSBgU45xWQxwY5OaFIqNM5yacBnoeBTsBJgYwTSMeaUcnNA4NSA8nI4pCNpxSLwCc0HrzSGNbjGaDjByefamtz05pR905pisIVG6hME800Nn8KevGabDckRskVHfEmE461Ihx+FRSEknjOKS3uGxlW5kQlQMjuaqzI288kZ6CttFUEnHvTWgWUEkV0qqlK9hcpVtifKAJyasZzxniohEU6cU8HOeOnWplq7jQrcA80wgZGBk0vB70hzsBHUUWE9RrHjioJPm7/jUjEAgHvTGAwQeoq1oV5FO74HIz7+lEHIwBS3IDDk/jSWqnG7d0HNdC+AhoeVHIz1703AZSCee9SOADkLz3pm5AzcfSpTFsxj9cdcdDVdgAec8+narPH5VC/Bz6VpFg27kZJLE4z7Zox8y/3u9KckEcAUmTgc8dxVgNZiuB0+talpgxoTyQKzZAGK9M1ft8Ki4OamUeaISZoRj86kTGfeoo8Dv1qRcZ9DXBIpE6L2rI1JihYEdVP8q11Pzf1rL1LBDkjkKefwq6NudB0O8xgU444ppPHIpc+leGUO/Ggk96YCT1pR93pzRYA5zk0hyfqKO3Wjpz+VMBpI655pRkDmmsCRx1pycAA9aYwB5JxQo9+aXv7UL1OaW4twIyaQnjPenA9cnmkOTzSABxS5OD70mc8Gl9u1MBR060nQUvBIFAHzYzRoAoJxk00dKXvSH0pAKvc96QMSM96QfKpFCDCdeaaAF+9z1pcAnJ7UKo3HNIeM0MBxxnNBHI96BhUzQeSDSGNOMZI5FLnNNHBNKDlAehpiCTDDFMYbgOafnkcU3I6np61UUAxDiYY6GlCgOSKRR+99u1Ko/eNz24qmguS5BIPSg9OaOg5FDcdqyAiOASvPPf0pvDECpmByD2phXgkHpVpgKBg9MUsjfLyfxpFG5eaWQjbyOKXXUGMDYOM59KkwMjiqysgO/wBOtWVYOc44qpqwIgmkGdtQxy+bJ8vQVDeh2lBH3Kls4QqZAwK1UUoczEncthcOuDzTmPy45oJG4DvSnLA47Vg2MQcKKQrkE05mIAIpSeCBQu4yPaOc9qUYJx2oK8E5pVwBnHNMQ1+GHpTOTJjHApXIHzGnrjbzxijoAgUCXJP0FSAcgdqii4OW59Kl96UgHA8Y7UmRt4oA+UnNL1UYqQBeBTR0p2ccGk7gCgBcdMnmgdaaTyATSnO7rxRYBD604dee9N9qU8gZoDUAM5BFGeQOooH3/rS9SMdKAuHGSPWmtx8lLxnn86G64P50wEAw2M9R1pehNDYBU96DweKAGsMjFIGpcAnPSgDk57UIYxl/h7E81GyhQTUzcgHNQjOPm6Va0EMk+XaB2p6uQBnv3pvO85GRTwuWHHA5qm9NQJF+UEnv2qvIgLKWFWEId8dxUcmDu9qmLswKixKWZj+VTxnIGOgqM4iBK85o8wI23pW0ryAsq2cEDmpRycmoYjiMMe9Srx1rCStoMXbk89KReTnuKVe5PQ0i8c9zUiGtyT6Y60obAxQ/AJHPtSFgBuH4UwHkALuNJ91c4pCf4jSMS3y96LAxVUGnMeKFGMgDpS4HSkwEIx1HaoyD1HFSZycH8BTB95himgHBep68Uo6YJoUYJpeoz3pANYELgd6ihVgMMeT2qbJYZI6U0rhsmqTaQARhSPypjZAwBUuMqKjcFlyKSAchJA+lB4x601SQMnrSscH5qb30GNbhvU00Lyy55NO3BpNvcUgcLJz3HSqV9hDgCBgHmgHAwOTSIwyeKcAFwalgKR8hGaYMYwe/WntzTSeDgcU0BFITxz+FOxjGaZIdsYIGTmpcZA96q9kA4A5J7YppUlQfSnA5b5TwOtBHBwakYKcpjuKRRh8Uo+YcdaFHBpXAa3yP6mmnKnnqe1Oc8hvSkYZG4dTTWgh7HCg4ye9VZ+QPTvU6nKj2qGZsggVUV7wMzJgIpN4HLHrUc83AQj5mpbxSjYB+btUCR+dIHzyK9GMU43ZOosQIkIHbqa0EZGbg8gfnWbvZnI7dxVm1Ty3JxkHvSnFPVjTRYjbc7evoartD50xkPBHapVTazvng09fmjLLWN3F3AhMarcFjyagmXY25O9TuQMHGSar3XykAj5W960g22hdB2WWEFm+YdvSo2YbCEOW7inRvvibHftSxw+XzV6XfMN7FTy3Jy3BXpToZTIdrdD+lSznJXDCqhG3Iz81bQ95E7GtAAZCQ3TqKtKCF56Gs/T2JBwMVeBYcY4rjqRtKxfQiJIYgA8dKk3shDA8inkA/N6dariX5iCPwohLsJl4XkkZ3HDexp8GrQsxEg2+9U3GACenvVBlCyFR26iuijWlfUiUUdOl1DIMq4I9qlBB5ByK5aMsHDbqvpeSR87uBW/t0nZkez7G3jimkc+tUI9VTpIMe9W47mKRQVYHI7Gt4yTIcWtxxHWmjrUhweRUTDFUTsPWplXngVHEpbOKuRxYAzTEwjjx2qTHOKdgUh46GqIEFROc5qUgkHH5VDINoIPX0oGVpM881DjOakbkmmKvNSUPijzzipQgAp6LtX3pH+WnYBjGoXfrTnbtUJ+Y0AMwSeKmVcUqR/nUu00hiYwfak+tOxgk0nX6UxDSPagL09adijqcUAMxycUBQRTuAKUcCkIhlUBTmqLkkEA1cmbJIHU1TkGDxUstIamd1XY84qog+arSelOISJskZoJxzRjK1QvZHJ2DIUdfeqtcynPkVxk85mbav3PX1qmeR149Klwctx6YqGRiD15HWrfkcV5SldkRO05A571G3A5HPepGxu3GmEEk+veoexor9CE9PpUiAY5FNI9OFpy4470kUrjlOQCx57VahIS3yB8xqoMlgvXNWwu0Bf0rSLMpqy1IiGlbJPyjvVq2jUROCcHHy471BIfmAxT4HwvXr2qkZ7qw85HvmopUP8PUVdZAY1ZevNVmj49D6U7E3sV84HJx7U5nyMUkm1ST3qAyqSccnsanmsVy82w6f5Y255pIx5S568U5IjJ8zH6UxsE4HQdDXNVkmehh6XKrsp3hwm0dzzVIgeuMVZu23ShSelVMYLe9Yxu1qdm2wvGTg8elPiUsQoOBTMDJ96sW4Gfah7FJFmFSkZJPSrsIyC3aq65L7cZFWkUYGOnpXLN33GSlW2EY4qqyFjjt3q+E3R/0qhehvMAVsZrOm7uyGydDiMEDmrG7aoOPxqpEyrtVz261cKlselKorPUBekgIGakBz8wHWo87WGPxpy4U9etYsB74JGPSmNyQR1pX4K8UcBM96S0AYwUH3poHysM0Hkmg5PQ9Opqh6bkYycgnAFOUbn4PSk8thghvqKTOx19TV77CHyHfnJwaAHyDu4x0pzFen50xcB0IOOtSthErjgnNNXCxjjk96kkAA9c1GANgx+AqVsMd1jB7j9KbvG3jrTlOenQ1HhskAcdhTS7gSOAWHqKrOMgselTBicEjmoZMiQjOQaqIDrcggDpUw+Y4HSqtqTkqRyKuHggCiasw6ETkb8EYFNHTpzUhVWyT0oVQWJPalfQBjfLGARz3pHQkgDpSsCwb1FNRyUBYfd4NNXFYY7YBT25pqHCE0pG5i+c0RcglqvoAhG0AH8ae5KEY6UmMpk9RTzkru74pXsMQYVvrSHCqR78UhB4B60kqfuwQeaNLgHRQ2etOcbk+Y8Got4+VX+9zUwy0YBptNagRvG5jUBsCmhUhAY55qYg4A7dqilG5gOwpp3FYM7iH9etKh3ScjAPel27TgfhTQRkA8nPWgexLtyCSM01iSnTFO6OBnGKc3IAx9KgCFQBGaREwnXr0qTaGX8Kjj4BwPrVdGBUnyhGO3WnBi5pJ+WIzgUkQIBOc4rZbAOxubjqO9SlQ8RHcU1M5JHSpFwFII5NS2FiCQllDJ2ODTnIVgFbGeTRCoSR0J69jSMRIxPTHT3p9bAWyDsCk8UhHyEfrSD7ozUbEkkKfeskhjiOMU8jIXHamZBye9Tg5Qe9D0FoL0Jz+FNXkFhwPSlOACKSPPGeRUdAFUAqaBjYcdKM5YntSjAVsUARBB5fXvxQ2PLB9KVTx14NNJG0gdqrqBITmk4496Y2TjHSlQjH0otYAcgRketVU4AqZgc5qszspJI6nirgtBE7N8w9KkBy3tUePlHrTwMoB6UPYZC6+YcMaahxlfSppCAfrUKDGf7xqlsIaBuHNSwLtbIORSY6enepUHIwOKG9AsI+CxBGRis18FnI61qH5mrPmj2ucdD0qqbWwMRGbcgIq/v2rx0qkcqfw4qWBw6BSM05pPUS0LJwFGKcSTg1B94sOmKlG4qM+lZMocB8/HT0pjgHgflT/4M96jyCx5pILkRJJHYjtUqjCqfWqzsQ+BVpMmNQ3XFW9EIhLYcjJxUpPygVFJHgZHBJqc8jFD7ghJPmOM1GwDtt71IAN2fypq4I3UkxlKVikgSpl4y4PWiRQwLYp0ar5Z960voTqPc8ADoaj2Hdsx8v8AOlck7Rj8Kexyy4qdUNigbUwPxpdvycnkUhH70YNKTz1qRiDlM0jrkD1FP/2u1NJGOD1poOhCTtlHHWpDiMZzTP8AlrjNOJ38GmwZIDwMd6ay4I96dnJApr5Lj3qUDK9wCfnBquo3nLfnmrUwOODkVUlG1QBwe9bw2JHkBWOPxqVMYNQ7iSOKkGVPH50MN0OY7ZcAU0MTPnP4elOY/vDjpUZBE5IPekhssZBPXmiTAlFIdgl/2jTto5BPNQCGyjKAg4FVn2lcjrVoMGByOB1qJseYd3T1qo6AQwqyZ461bjYsSvpVUO212Ayo6Vctm3RZ6GnUva4kSuNo68VHENuc9ulSSjgA9KBtBIFYp6FdSIrhsjOTUgXCnnrSFscU8HGM9KbbsLQXGQMUgbLA9qFOCRQMA+9SA5eOO1IR19aRcqSO1LwV680uoIb2NJwB9acen1603b0NUCGnsM9aeCAcVGx2gZ6Uu75s5p2AlJxUTHccDp3pdxAB9aAAAaS0AG+XAowFGPWkOSPegckUwsRSqc5U00BS59TU0gJ/pUWAZOKtPQCI8DBppJCnmnlOCvWkA2Z9KsXQjZR5maY3X2pxYGQ+tRseAcZ9KtB1Ks/3TnvTrZSqc45ps5OQByDUqA7cZyK2+yIYwJOOgFR9TgHipmxuIxyKjZSeM00ybdgXgEY5qu+TnAxU3JBGcVEwI5qloHYTHy/h1pGHPyj8aMk+lJ0/AZqhjX4A5q/akGIHv/OqD4xwT71cs3+Qgt2zxVCktDUAyRipFGDnpUCNk/jU4zv/AKV589GWSpyTVDVVzEzZ6Kavg4IqlqY/0dzjjaaKb99AztjxnNKo9aUjPBpBXiFCgAZpemKaud/PSl68UAITzzRnFKQMGm+lPQBW9BTMENkHIp5PU0n1oACcjFHRfekAG4frTjnccUDDqaQHHBoIORzQfvciiwheg65px7U0D5utKTn60gAHn6UEgEnNJ7UFQSDQAoIA4PNByMk9elJtxzSHk4poAyCCTS5zz2pCCBjqaVQcYPahh0HDls9qT+8KUqSM0H7nI5pAJncmKU8Y4oHC4oIJ470wG9Dj1pO+D09aVsnHrSKOAKfQBXG0EdaZJjyxUjLx1pgwwK96I6gRLk4PUdqkGTP0wQKjkUhRjselToCSp7960k0CHNnFMYbu/TtT88kHpTcc1n6AMUZbOaX5SxxQAQ5z0pg4kI7UwJOQRzT2AwQR+dR44FPPzL70gIPK2g7eM1KPlTkc0pJDAdBRnqKfM3uBEgV+tKFw4A6UFCXBU8d6cMliKd+oCYJl3Z4qUe9MChSADT+tTIBmMEdcetC8A96cTgZ7UmAF4PWkBGR8tPB2j60pFNPLc9qdwEZflA9etM5dGWpHyEHrTVyoJAprYBygEjHbrUhUgE1HDnYSeuak9+1JrWwai9VoUHb1pABtoHTGanYBTywPrSYGcjqKBwRRyCaAA4OCaF4wTTT90L39aUEcA0wA4LZxxS8Fee1BYAhccdqQqckdhSYxVPyg96XqeKQntSk7T60hAcbOKHIwMDFB6ZxQ3AzQgFwD+FNGKM/LyeaO2KYCMBt6U3PPTrTxkDpkCmAHGfSmgAcZHSoj1296mBGM00gBs+tNAROAxBB60/oMj86ZgE5/KnxnBPpVgPHHbrUcyqV2+tPXJyT0pCATkDOahaMCFIwoC46VXih/0lmbn0NWcHzCvtTZARKMH5fat4ydwJ+o5HSl/ipR9yowATk1juA8jjA70OeMA80vce1MOXPoKSQClugA4ppIBwegoQ8EelABLZFUAucjd2oAyetIo+8AeppVjCyM1GgDojjvT+re3rTO1P4wAe9S97gIw5zTRw/FOAwTnkUnfPakAZx2p/QgYpmDuFOcigBrA7+DS/eJGaVsA9Kb1ORTvoAueCBTSBgYpxyehpOpzSQEZGWwBxTmXLYPpTQNshPUU9+CB3NU/IBhUhge9RYJmGOnepmb58Z+lRg7SR6dapNjAN84A7VI2BgmmxD58n04p+AT9KT0AUAZx603GB9KcMsfxpshJHH5Ur9AG53ADbTh0pNwI2jk0meg7U9RDl+UbsUrcggd6cOo9KaASx54FLqOw1c7QBRjBx6UoGMUgJDEnvTAHwqj0zzTZOMAcA9KeQDGAeSaY6h3UEcjoaa8xaiRsI1C/nTJwNh/nTNv7wk9c9KknXMQHTNU97h0Mi7BdCVHK96rWhI3sxyCO1XJuI3TOcd6oRny/mHHt6V6FNvlaJ0JTnaQF5HU1I7h4xFnBxzVdJG3ZHCk8mlhUm4zjgdBTaAvqyhDFu5A5pVxbxkFuD3psi5JZfvd6STlVZuAKwT1KG3GVjV1PQ9KZcMNikjg1LITzt7U0ESIAehqkIiswHQqwwx5p8jsc4HSpYoVjPXOaaygqRjHrQ5pyuC7FR0STD55UVVZt0nI59QelWUdUHzDk9qp75DM3y53da6oN3JNKzlQ5XpV+JsjH6VjQnyyXA5JrWhfepx1rmrQ6lIdEclwDUe0JKWI4NOQBcgcUrY6VknroMUnIAxketU2A80+lWgx8vp0qtIoBz61UdxMh6P1yB0qypLgrVZwNwVRz61OCOGxz71tLVWEhsy7V+Xk9Kj37QTuO4dBmn3Em4jYOlVpMAHPWnTbQMuxatPCQr4Ye9acWq28qgE7WPQGudZgzKuOlPMilG+UjngV0qpJIzcUzs7aSMRgkjnoc1aE8fTPNcLDNLjMbMMd81eS/uRIE80kgdcVftorch07nXeapIGaXcDXLDUrhG6hsH0rTg1WI7UkOxiOKqFaMtEyZUna5pse9QSksM5pyuHwQwI9qae4zxWlyLEO2nqvNKRTh1pjJRwOarzNg1NuxUUgznFDF1INpz7VIqDFOC+1O4BpFABjilOKOOKQnvTJEPH9KQ9qUj9KTjPpQMQ9OTQO1JjjNJ0OcUgA8cd6RmwMCgkcnvTT06UBYhkHOfWq8v3zVmTpn3qvKvzn3qGUhqcVZiPP9KrjGSc1NGOR2phYtDpmmNg/4UueKQ9KokrvBG3VapS2m3O0/hV9jg+1RMAe1Fxezi3sZTxSLgkZ9ah6Nlhj2rUIAzUbKrZBHFK5Do32M1iFHTj+dMzjJ7npV2a2TbxxSQ2LyMMcgUc2pDpyQyCMIhc9aldiOAPpTpUZDgjAFV2O4HA5Pv0rRNWOWSbY9iApJOSe1PjIUbS3Prio4/nY9qmjXJyOapahK9tCwjdeOD71PCUcjcMr/OqoPbsaUS7frVWuZ3tqZl5vFy64+UHjmo0yGUg8elT33/HwSKihHmycng1zvS510dkkSrO7gk8ZqPODgU8qFzjsaY5wh9RXNJtnopGbOwaZjUEjD5j3p5Y7uPypCATgmmtCtWiNTwFHWr9umF69KoKD5h75rRtcMCScipqbFR2LcZZQSwqaBshiBgVDBkuTnKkVPArgc9BXLPzGW1cYwKp3x8vDYyatRggZJqrcRmUEk4A7etZQaUhu4yEhzufHHQVo5G1SOlYpkCSJtHHpWtbsWQMR1qq0XuC1DID7ux6VNtz3qDcDIVA4qbkHBrFgI56eopA/zkevWlILOTikBzIBj5aSAbIM5A7UEnFEjHOR070nRvYU0AgAKc/lTQ4J2nj0p+Qec4yelQyA5+XtVICV8YGeTmhkXeM9qCXJBzwPWgn5uaWwEhPy465pqHD7T1NPUZAPamOMuSo5FJAPTGCR9KYflYt+VPXBQHuajfATaTmhbgG7EfPemyDcN+MU7aBgE02QZQAiqQthlpwSSc1ZUjv+FUoSCQAMEVbOeBRNajEbIX60p4AAPPahifM5qIMyt81JK4EmR2pAAUIxxTQdnHftTgQFXjrRsBCQCh21GpyFxUs0btG3POajj+TvkmtFsIlySoxSgg8DpUSklTk9KcG2DIPU0rAOcDI9RThjaWPbpUGNrqGYkmpTkqwxx60WAgkQ4Deh5qRCSFweKjknKMRjAp6MRHz3qmnYfUkLAKfUVGW+UEjpUbyBAwzySKqSTNzzjmnGncRaLszZFLGdqbupNR27MQPSpn5Yf5xTemgEygHk/rSk7QSRz2pF5jC96VuUArJjEBBwBTHXD7ecUiMy9unWpGzt96ezApFc5L9qZ5oUkY61IUJbrx6VWc7W29QOoreOo9i3EMgNn608E7ic/KKigOUAI/CpmGcgDjuaiWjEROyiTk9f0pVVQx/nTJ1DbcA4HU0MAmATwadtBeZZXhDkcVCTmTpxU4wYwCccVE3BA21CKHydKfESY1B60wkYORUv90jpjrUvYT3A/dIP4U2MEAnt6U898mmpwmB3pJ6B1FGCoxS8MMCmrgZpUP3qGG4g5yOwqN2CnHrT0H7wjPWopMkMwHSqW4bIcGOFUc08YCUg5XiiVuMCjfQOpExLybQfxpsq5AwMrUkYPmHI4NOx82Oop3sw3GBcEn0qQL8pz+dIBljTywyBSbERyINnTpUUa4NWWIINRn7pxQmMhUgy7ccVKnB5pFwCcdTUijCg96psBp4BIHAqjcnZsH6VoEfIwHSqN2o2Ddyc1VN6iYxT37VPCCHJx8ppiDHAFSQn5eeAKub0EhCzKxU1MrfKBTCRJJ7r1p6AA9OlZsaJAOMVXK5Y1Y6g+3SmlOM/nUJ2GVnQswxxipskgH0oAy54pVHysOtW2IYRl89qepGOTSAYBNKgAQ+9JgIBgZ7UzAX8akB+UjvUbjgYHShAMcEODSYIcDsKfIBgEckU1wWAweR1q0AOP3gXtREOoPanN0B7imoAzHntR0DqOBwQO9KwyxHSkwC4PYdqMYkznrSAdgFNtNOBxjoaUgAmjG5jzxSCxDu/fEZxToyeWqMqA5J61LGwxjHBq3sA5eoIPBpJjiRcU9R+8waSTBYZ7VC3BldjmUMTximsu+Q8dKDkyEHtSk7hkd612AgPJXHepVUMpJ7cVCTsAYjpUkbZB96tiHMN0PJwaaQoIOeR29Kcv+r59elIcCUPjg0ttB7jxIrygEfjRL/rsg/WjpKxoeMDDP3qeoaD5cBODgHrUII8k0rFSqIw69DR0kMZHHamthdSNCPKJ96vW23y93Y1WiCum0Dg1Yj+7sUcipm76D8yRhkBjUKSoZmXvU8uRCfas/afN3L1PepgroLlkPg4YfSniTJ9+tRt8yDjpTiMAEcUNICVuGye4prH5waVTlcHtSYzipQDgfkzik7U7ov9KaMZpAAbPB6U5jwABTCSOetLgkCgdyNxnAamnJK4HFOf5gD6UFsOAOhqwHAg49qCuXz60oxjFJ0UDvUgID+dHGaaThuetKxCrjHWnYVxJGOQRUIBU5PU96lJ+Tkc1E4JwBVx7A2ISM5zimOTkipgmee4HNQty3oapAQP97GaR2Az7UBcOeckdqRs4zjJ9K1Qr3K9wMLuK1IjgjbnrUc4JGelJGcgHPFapaC1JnXkktknpUOT1HSnS/KOBmm7hsYd6IoHawgHzHHUdailO5STwKl3ckjioZPlPTiqiJeQ0rkelDJ3JBxxxSHBGOcH3oU/OAOFHf1qwWwjrkHHcVNYnkr2phA2nHFPsfvcU7uwdDYwoHtUmeeKiUjbkgVIeRx1rgkUTDiquptm0cf7JqwvKg9+9V9QXdZyEcYUkilTXvoDts4OaQcHPb0oA5xQeuM14pY8AnmmZ5oUHHWg5NAg+9QajQFSc08U9hoD96gAZ9qQcNS9jQIAc/hQp4Pr2oAxzQBjnFAMUdcmlxk5po5Oe1OxxxSAReRnFA5XNKvpjg0oXg0wI8DJ5p5GBmmjGcnrT/agBMZFJ0IPelLY6UpHFIBvIOT1FKvQnNJ15NP+opgIGyMDmgA0iJs+7+tO55odgEHLHil/ioxtJI9KQdRzzS6gNJwTxyKFOfqaUj5iO9MBOfSqAf1PXp2pmBu9CadjByOppGGGzRHcCNuDUiNkZ6ComyDx0FSLwBxmrltqA/IzQOc0mOKUdRjtWYDcd+9MZRgDual6nNNcAMMc01dsBKfxt5phLBsY4oEin2ptdgBi3mDutKM8Zx16U1BluTxUhOGwOlICPHzEdqcuBz1qpdJO0oKHC1aiUhQpOTVOPu3BDguGzTgPSkyQcflSqxzk9fSoeoAwyCKiQnqR9KlOQoNN25X3o2ACcnikHLZbvSnG0+tNB60AKcMcUgUE4xQcKtNRiDnmqQEuB9BSscDFNPB9jQ/bFSA7vTl6k0hHGaFYhakAAy3FGCWOB0o4OcdKOi+9MLCKOTmlXqTQw5A7UEggADFAxOMlqQE4OB1pSOcUAA4oEGOppTyP60exo6EgUgF4B+tNUlsg4/GnAYHBpBypNMNRp5/ClOcjHegDqaQff/nQMXnJFBFOUjJNIBnPpQIjxRL2NO6tj0ob72KaAiYAgACmnOPl7damJyeKhbKnGOtUmA9DlfrQ2MbfQ1Hg7xjpTsHbu702hiOQELZoHzRkjg0AZBzyKUYVc44pqwhysQdvqKacb1Hen9cNTSP3hNK+gxzHnFMJxgVIeT7VGR85yKSEDHYG9KUHKg469aaCHGaej7l3CnbQBFGAW6E0oOR6Zox8gz60JyuOlDAXqQT0p/JA9aZjOBTs4HvUDBh1zSEZHHGKUe9B4UZpiDn8aU880m4EgnvSknr2pAIeg70c4AzxSAggjvS5wADQAfw4X8KQegpRwvf6U3ODx1NMA6cChuFFKewpGPHtSQDOCwwcU04Pyj8aQfePPanr8owRz3rTYBI+4HbpTiQFJNJswSSacR8vIpN3GCk5JPQ0jkdBQuSDnimP1IHFJIBoJVeDzTlGcc8U0bViOaQyeWowuR61o12BFgYA47UoP3simxHIJpSwrNpgIMbue1Mblx6U7O3NNO4oCfWmtwFX7uM5NL0U+uKjJxk55oU7lYk8in1uIaTn7vWmOSeCeaZ5jKQoHU0rA+Xlvvd6tKwIrMVlVlz+NZRVVl2sfkq/LlQwXvVWRAAXxlsV2UnYWpEvyHDnAPSprZiOc9OhqpMS7g9h1qxanJKtytbSXuhrcvKwO5wefSo0kMseCOlMYvFIMngelNQ7ZWGMDrWCjcLkqHzMqx6ikiwXWMrnHBNLG22Zi4xnpirG1cAqKG2g3BUYIBnJzUMp5C4xVhDuyM81H1X5qzjuMzJo/wB4rk8A/nUErYlBBAB6VouqmE54xWfJCMb92R9K66bT3It2HNu3KAOvNaNmDuOTzVS2ZXABP3elX4ABlh0NRVl0KSJXGVOOtRyvtxzz3qUjBOelRTICmRwa5o2GKGA56Z61FImBz1qVgNg9R1qtN+8Tk1ot9AZC52sWWnwMzhtxyajY5BXv6VLbAKCT1HStre6StyJkOWUcE9Kgdc4UdR3q6QHbPf0qCVCHYAciqg9bMOliHGJN+M+lPd2mJ5AxUfOFXPJPT1pAeSB1xV21J9SxEA0rE/ePUVOcrJx95hVeIMg55JqfJBDOPmUcCspWuVsPDKqBnPTtUUUpaQFh8uevpUUtwZQqDjPWiUlEULxzTjG2rEXv7QngkzG3A7Z4q/DrwkAEkWPcGsQPui3dPWoQjcAHGBWtKo4qwSgmdbHqFvMcI4z6GpxIOxrkFAVsg5NXILuWME7yB6Voq6IdI6TzCaC+axV1RlKgqGz3FW0vUkJBBFWq8e5Hs2i/v46UeZxmqwnQkAMPpTwykcEVammieWxKHBp28Yqu0gUetMM/HTpT5g5S15gxmlDDPWqokLHApQSMgn6UXCxOSMdaQ9KiCE9+KkCAcE0XFbQCc008igjBpOecd6YhrqShH61XPrVrJqN0B/rSaGiuOpzUqjFMZMEn1p8bDoaSKJ1ag8imgEDpmlOSM1RJC/DUxjgU58Goj0NJjGtzULnbzUtQSsW4HepexSQAGVgorRRRCoQdR1qK2jEUYJ69qduDZpx7iYMc1XeJHOSvNTMMjIFMxmmS1dDBZqythsenFIbSUdfm/pVmPuKtoAPxq4toxnRjIyGBXjpjvTHYHJ6ela0+GXGzdmqs1pG+SPlNXz3OZ4aXTUx7nLp7iobfqx9O9X57MrGSDkD2qnEhRMdwa56ktbHRh4ST94kB+bPrVa6bbCzYqfJAFU705CgVg3rY7UUGcB8ZPtTiAeKRgN+cZ9qaCxJyefSmXfoOUdSfpV+zIKkY5qqi7l4HJrStogqbgeazqSsivIdGu1ynqKmjRyCm7p3qIptGc85496tKjbvfvXNJ6XGiRMKi7qgvV8yIKp2n1qwVDYOc024B8vA9OayjL3kwZlmAGL3Hf1q5Z+YoIzlajV484HpxU9rH5Z3A4HpW05NpiRMVCPuHX0qQf3j1pg4JB5HrTwNoU+tczGOJw2MdaQkbelKOEOevrTRyCB0qQGtgoecU3eMfKeKV8AjPOe1NICgACrQMbjDDP4CmMwQ5QU7aWPTHv60m0AD0NUtAHgHYfU03lSctmnrxAf0piLuXJGaAJoznoPlHrTpMBiB070xCTx2qSQgnANQ9xjeBx6VGvL5zS4O4segofAUbR1poQoOWIpx+brTGLAAAcmlUZGc8ikwK+4ifAHAqyPue5qDgSEhc81YXnPFVJoBOhyarZbfkmrEhOfrUL4OMcYoiBLjcc0Abhz2pA3IUH8qUrjv0pDGNuYgdv51BsPngZ+UVZfOR6VXKnziRWkXoIePmdl7dqQDenpjrSKwDDaeMd6dGxLkY4oATCyHf/d7088KMc5601wFXaBx3pAMp160txEM5DygAVIpyvSmNGFkBzyaeq7MFu1aO3KPqUPKZ5iCT1p4hCNjGfap1/dXDEnk05mAkDkdavmZKKzpJHhhn6VL53mxlc4JqSSReBVbyyrBl6UbjRctlKL855NTZDLj0qtFgAFmzVhjnGOpFYy3GRBiSTUw4TfVRTl5F78VOhGzbmiSAi4ycGoJMFsL261K/yMQeKq7gGrWKGWImwST0p8L7pXA/EVCGwSB2pQ20sw4NFhFjAUgEYqs5V/lIwRU5cPDv6mocpsZzwTSiDRZjIJwe1MkbnHcU6NgUDd+9MkzvDCpXxBqK7Dbjv3qZGDRgZwRUDD5t2KlUYIYDr1olYCX+HHb1qJidygHjFTMMKBUL/wCuUYzWcQJBjBB60oxjp1pGHznHag/KKQAcZGOtDDqAKZkq3PrT24OaYxAuzkdKY3GO59Kdg8EnFMkA88CmtyR6j5hxS4CmlxhevJoP3sH1pbjEAwMmkYcY7077zUw/eoQDFPzeoHegkb8CjIX60jHbFu6GrAUOA3SpByBxVaNvWp1Hykmhqwajj8x46VTufvIO3c1ZDfuzioJxlc9+1OGjF0GqO5GB35qWAggrjgGodwIHanrkwsRxg1pLUWxIV2yEjvT04znrUKMJQSD8w7VOjDB9aykUKMhdtDAcCgY2k/kaXOfmPapDcry7hLkelEMgK7vWlmYKAx/ColBBx2NaLVAS55A7HrTuA2AeKa2RjH40cjr360g8wOQMimMdrbzUvfPao2HynPI70IAXABJNNQANnPShW3RdMY6UgH7pvX1qxXHk4DGonfysMBwam6xjPbrUUuc4WiO4O48jjjrSHggYqOTJwU/KpRzyeeKLBuDHDDjPtSn7opuNwJ9KB8yDPBpWAZKoON3HenxMuwkGmykBkBB5pwUBWC/Wn0AVeBkmh8YyaRceVjNK/MeM/hS6gNkABz2qBeEOOnap5VBjANQun7tccVUdgd2QEfKwzzkUjcKCO1Iys2cdf50DIjyT0rYRYiYNGQetITtwhGc02PK4J609kJkLDoKjqPWw1wN+AeRinu4BVSOnWkjwy7yOacrIYw7dRSYDSNsoGOAODSjBZSKJTmPg8mmqCkJAHPajoIkYhZGx2qSI4IPrVdstDkfeqW3UrEmalrQPIsu2Aaobikygd+lXZuYwcVSOHnx3opgyz96PgZJob7gp6rhuOlIRjjHNTcYoGCDQ3H40cdcdKcemaQBggc801uCCTT16VHJ1ApLcY7+dCnJzQpoBz8opgMPA9qbjgDPNPHIINMYAAU0Jj15Oe9OKlqZ1AIpwfBx3pNDIZAUkBHekZ1kbCn7tSTuFXpnFVFKr8wGM9a0irq5PUsOcp9KjU5b6ds0g5+hoQAsSKdtBjzkjI61FKuBnvUynacdaikywOTRHcTIjjO6o2I3Yp+OeOlR5Bycc1qgbIpc4J7Cq8bBVXB/DFWZBlTzmqkZK7gOfrW0dUJtlmRjgEdKgxlenX1qbOV44x2qFhg/e5+lOIpXYBiTjqaZLlgKVW2sfXqDQWAGD2qtmCZEMcALyPWhcjjp9aCexPFGMnOec5x6VYxDycMwA6VLZkCQ/MOlR5BBOe9LbnEvy00r6Cb0NmNsx+lSrjGcVBBhUI71MvA47dq4Zq0mikyUMcAetRXmPsMpH9xv5VKnNQ3v/AB5SjH8B/lUw+JB1O0HqKMY+bHNRs3XFSKzMh9a8bYoQEjrSnpjNLjjFIwODik9WNjd3FKOufSmKQBk09ORTaAAeSaPU0q8A8Ug5GO9IQHNITzk8U7HT2pjNzg96EAoxwafxgU2Mceg7UueMHvQAHpn9aXowwaTG0gUvVuOtACkE8Ug5PtS9sGlwKm4xuOlKRkc0AUgySc0xCLkgCnE4Gc01mCf405e5NPzAcOcUmMZ5oGenagggZqQFAwMnvTWHFKev8qUkbQMUwG85zmkJ5x605hgVGT3HamtQH9QAeKZLkNjuaeDwCTTJDwSKaAZn5aehO0ZqPdwO1PRsqBitZaq4C8k5/SnDIIz3peuPWgNzj0rEOgh6n3ppbHOKdwOvSm4PB7UIAOduPWmZUjaRz604kEY7ikAB59KpO2wDsYA96eOW296Zzn2pUU7sk85pDFdSWA9KTAL5zTiM5qNFCtnP/wBakthDwcjnr6UcHFG4ACgqRjigBSeRSAjnmkl4GccgVCyuXXH400vMZIO+aF+XORTyoJGKR8BBQnZiImI5PpSRESHmnFdwIAoiUDOePStE1bULEhIzz+FOA5+nam7doGaeAAuD1NZMAHTrQpJUAetJyMUo+4fWkAdAT3pcZ9zSD7vPekB4x39aA6iqckk0i5H40d+vJpSQMj2oAGBU0ikk0rYxmmqNvU0LYB2eRjmhfvZNJ34pfQZoAAMDk9aQdD9aUD5TzQvQ5HXpQrAAPJpFGMmhRxk96X6fjSAQcGkxg49TTjjOR0pBnbn9aYCfdyB1oYc9OTTiOQaOrMaAGngYFNOM9OacASue1MYnbkU0BEcgZ7CnDI+lJ0RRmgZCn1rQBEcOSQOKk28bT0psYXFADEGk9wHDJcqOgpX4HFKhyCe9JnnBqeoxRjCjNMcnnPen4yfYU08t7UANBAHTmnHHXPy00c9ead7CmxXAk5xinYwMAfjSZAZR3o6H2pMAIPXv60Dpz1pRyCBTRkPQApBPOeKUn5dtDfKMCjblQfWgBMDaPUU7llxTVG6Pr1p3Q4x0obAQECTJGaUemaaMls9xThgtnHWgYA5pG6Ejr2FKc4PrSMDj60kICeKCPkNJjpk5p3Y4NN6AUgxWXBHWrOC3XtUUq8e4NSKCY855NW3dXGAGdoPepMjr7Uw4Cg/rQCCWP5VO4hpkywXNMkGDntSQx7JPmOTnrT1TcWB6VWiGQ+X5qnB609tqxKD1p4i2r8vFKsa4XdzjpT5hahG+U579qkxnBqL/AJbAEcVKSAMipYyCaQrJgH8KWE74lx+dVbp381doyO9Tw/LEVB6VsopQEhwHJ4zmnuDtwBSRAhMnr3p38JP+RWbYyiXCThTyT3p7guGwcD1qvck7t2cAVYBzCo7GtWtExFOaMldu78apOGXdhuB0FasygBT6VmleXzW9J3BmerdM9KngbdwtV3QpJjPGantsbiR+ddjtykF+FD5vzc57U64jYkBRwe9Is0byIF++etPckgjPSuN3Url9BXiztOOnU09QRuyaaOcDPHpTlIYEZ5FTdsBB8pDDvSOd/Q8imsMtt7UyBTkuxzmhJWuA05CnIqhKFSQpjiteVQxIHJFYspZ3ZiOa2o6ksljbYpVR171ftmO3BHFZqncuc9PSr9o5dMDtV1Yqw0XSflA7GmH7xQ/hUgI7+lM431xoYgQAnPWq7A4ORyatOoI3dx3qJxk89DVxYFBmaOTOOnapLfDHO7J7illR2woGB3NNt+XdcYxXRdOBNrMezbBwe9VZZC8vXGKtSqMZB5z0qjIzxuRnrRSSbBjlXZImaRk8uTOeKI38wAt94Usp3cAYHpW17MnYWKXypAWbcD7U6W5MjnsR2qqNzEd/Sp9uV6cjqKcox6he6HLH5ZU4/WrlwrSIy46dDUCkICd3zLU6yNK/yjj1rGVykuhSj4yuOe9TwKHVlJ5xUbZhuif4fepUG19w/OhvS4xkUZids/nVhh8ny9T2qMY8x3PTuKnVuh7Gom20gXYiLbQu78farMXDAjvVdlbb8wyB3qUTKoGe9RJaaDLQ+9k96CW27Qx9etMPQMDwacQMZFZp22AYZJAN288dqWO8kU/MAw70xjjvUO/jdj8K2p1JoXKiydXWFwDGcHqc1cg1O0nY/vNp9CK5+d9yZHT6UtogIyeldcartdmbgmdcJFYAqRilDCsGKdoRncQB0pBrbRvtePcD3FVTrRm7Eum+hvEZFNIqpbX0UwBRhz1HpVsPvGa3uRawg680mPfrTzTSPamSMK5qAoQcirP86ayg9aTQxscvY9fWpH+7x09arsmGzTo3IHNK4xrjFREH8anmPmHcO9QyfKtLYOpAzcnAp9vDvbcRwKYib29qvKAkQUcULcb0GP6VGOGNPbmm4+bmqbJEOcUAc5pcZpM+vFAWHpwT6VZVuKrKc1KDTEOLknFM604cmlxiqEVbg4ibFZGcD8a1r5gsGB1rJAOK55mkLidqzbt90pAH3e9aTkBSaxy5ZiSOp6Vl1NF5iHkkCm9TzS06NMngcU7louQphA2M1eVBsXb3qKGPMWDwasKAiqntXNN3bH5iSRB5UGamQESYzjHWkjwxwDz2NNJ/f5xyR+dZWbQ9ixGoCnnPpUUysInycA01ZMoABjnmppwTE2fu45NRy2eoaWMo24coFfPWr1sG2BN+5lFUSwiX72SOnFW7N1aTI+8etdE03HQmJZdmyFH408EsduMAUxiFfHTPWpB9zPrXI9itBUPzHPQUoABNNjPB96TBYnFIBkhyAewqMsSy7Oc96lcfL0qvCzGQkcAVpFaB1JJFbcCDz6UFcxjHalYDGATTM4GM0LYNx6dAvY0JkDGKVfmUY4pVJ6elJsYq53cinkAc03PYU4D5huqGJDWBz160hHzAD86A26XHanDOwk8U9gI3BElIwx3p2cDJamEkkMTmq1ADxjHWpUyE571CD827sKnXPJFEgEJIGcVE4ycAVO4wcdRUbn5RxxUpgGMEilAxk56Ug4bce9GOvvTBiN8yHB/SqrjcpXvVvChQBwKry7YlwO/WriIYFCFTUsY+fcDwRUZ3NtJ7dqkjO5j6DpTewDX3FXbPFORt0YwODSsPm2fw02LKHH4AUPYZE/LZP8PepInLjJ6d6SVPlOehotRtRiORVbxFfUrOh+159elE0oA2FelPlXEyv0xUDzCSfAGR3Nax1SFsOjKHA/yKsnlTgZzUCxL5zE8HtipXYxkDdmpbuxkaS4bb3B6VcIyPf3qiMfad2OT2rRAyQT1qamgyiRiQMf4qst/q896bLjzFOOlSceVkj8alu9gsVpWyx9R1qk+8S46jtVqQ8sRUBjIbd2raGgxwbGGI61Lj5WLcDtTFACjPXsKfLlogmPxpMRGHGxlzgdqR03x4XpSKu1S22pVwqtxxT21QD4MMAvcdaWTLDB9ajgTq4bIqR37YqH8QXAZMYANTq37setQc4zUqt8qVEkBNnjOeahYHzlIP1qYAYBB570xmCz89WqY7gPBy5YGkHzMSaacqeKXAC0rAMZtpBPSpjywz3qJ+Ex1p5HAz2oAQ9ge1MOfOPH1qViPvGmKwLsQKaYDxyc0nJyaCML9aTB2rikCFzgf1pkgww9aeRjFMYfNmhAMCjPPX0ptwPlPFSIFHuaJCuatPUCqhLgKByO9TbyVCntUCHbLj1qaNgc54q5IB64CBKjlHykDv0pS2MnH0pScgZ6etTswKSkh2U9Fq3BjBHY1WZAsh5+Y9qsW3+rOeCK1n8JI0xulzlRxVpR1NBII2mlPCgfrWLbYxjL8m0HrQAWXGeaeRwM8UmNqnBxU3GV5VAQHrt7U3eC4b8qmkHygdzUTKAq4H1rRO6AewO8mhvmGc0rjBFDZGBSAARvANIwJyo4J700EF6eWAbr0o2AaTldvpTCNylQORTxjJ479aFAElNaCY3ouzOSKQ5Ke9Kp+YsaYCTLtzTSGRAMAW7irCcrnNRk7Ztp6GpM7cjsKqWothwChDgcGmngZpWJXGKP4QKgY2RcgZpG+Z+Dz6UjMBtGaUqA24d6pCEdv3ZAPIpxyYxmmsh84MOlP6uVA4FADW+aIHHSmMDhfSnDmM+tIgLIQaa0ArthJSfWoRIAvTrU8qblweo6VCi78g9K1jawEsZGBvPHpUg++wJyDUWw5UHtUiMvnlSaT8hBGu+NgvFBCyDYOopy5RsE9KjG2Jy3c9KlasOhKME7O4qGEvvO/j+lTgBZGJOc1GrK7hcdKEx9RsQIOSeGqWFmDFG5ApjZDjH3RUqLtmJ/hbt6USETuPQ/SqUjBJQSOfWrgPTnp2qtcIMjvmphvYbLCHJ4Hag8EE02MbYwBTyAVANR1BDM7RTs5QD9aY4JQe1BGFPPSmDJscEelRv2bHIp6cqT60FeMDpS2YaiYwQc9aAPnz2pzdeeRTAMOTnjtSATG4HFNJzx2pyKVB54pO3NUAgJJz6Um/cxA6r1oHcUwDGSDjPWqSQEjAYJJ59KrDaUIHSrHbGetIseSQBihOwbsrqpXPPephwTk4PfiiRCOAKcB+NNu4bgByDUUg5APepsDecfjUE/XfjpRHcPMiwAPY9DTDyTinbcoR696YSAoU9TWqAY6jZg1QABkxnj1rQYZFUCpSbA61tTE7llAcYPSo9mTkjgdBUnBHIppBXH60ybEZTnI60yUDtxmptoJIDVE4I7c1SEyI/KfbpTWJJJIwfanMQT15ob7zA5NWhoQkA+vHGKWA/vhzjP6UDgDI4PSmj5ZlYDGe1NOwdDZtwCD0JqYZ21DAR2GRUq5B5FclTWTY4p8upNGwOaivObObPXY38qmTC5x1qK6B+xz/APXNv5Gso/EPqdfgMSemKeqbelGzJPFKFJ7/AIV47ehQPmmOwAxnn0okViODg+tRiIKdzHLetEUuobiheQCOakHC1EoLE81KcY96cgEPGMUv3SPemEbgPanAgkD0pNDHH25zTCuWp4wR7ijGfmpbCEUnPHanYBGAOaaB1IpwGAPWgAORjIzS4+bIoLdM0HgZxSAXA4oLc00fcp2fl96AEB/OjseaRcLkmlHTHaiwDcAn+tOAOPekAAJFA+tADuhFCjjrRjKnjpSgYUeppAJjOeeKB9/noKXGDj1ox8/0p37ANcjFRjuRT5EYkBeKRUI4qo7ASYAAOPwqJwNpqQDrTHwFY9qSt0AgIyvsKfB0/Wmgfuz1xQvBBFb30sBYB7img5BJ60KxIB9acOhPasNgG8c0A0h4zxQrdeKYDcgc+vWmv1A7GnMSRx0pdoyGPaqTSAcVHy80HIwD1pxwQMetIRzzU3AVhgZph2jOOvpTyOT6VBFEEmc5PvRG3UBHIznPA6ipI334BNQmPKttOTT0TJ9DVvlsCJpAFcZORQwHOetJjK4PUU12DYweBUAPQ44po6HPSlRw2MGnEdRn6UgIxwOtNJ3HA7U5guz3B6VChbzuenrVpdQ6lj5iq46elPGNw5ppBGMGnHGdw9Khu4XEXO4nsKUDJPpSDhOvHpSqQUx3pNAHGTQMlselKOFxTexNCAXkc9/WlOCM0hOO1H8PPWgYvU9KYAGB/vCnjhcn8aiUnexApoQ/PygEc0cgZHU0ingClJwKXUAODwO9JyeO9K2QR6CndfmoAQYCigcCk6rT+pA70bjEUcHBpp6bRTjxn0pMYIxSEJn5DinD7pNRkfNgUoYnPpTGPAyoqF84OKlJ4zULN2Hemr3ERIhCsPelGSSfUc04hgxwMj1oK4XGOTWjeoDMCOPBPOanHI461G0YcJz0qXgfUUnK4DUBVeeaGGSD2px4FA+4TUeYxq9CAaaeuAacn3MnvTTwOKfUQ37m3JpVONx7VESXC5HINSjO4jtV20AZ/wAtBjqelT9AQeaYQAQwp5YGk3oAqYB/Ck2/N9aaDh/rT+nPeo2AG6mjOF56UOORQ2SxzQAHhR6Uq/czmgj5QfWkU4BzQHUQDJLGnLwMetNBBXFL95R607MA6nNI2QB3J7UvU89qG+YA/pSQDVyXGaevXBprYx0oGMZPX1psBrrnJpAM8U88jk9e1IOhPpRcBknbHegKVUeuKHbAApVJLZ7CqWiARF/eZ9qeOdw9Kap+fHrTj6Ck2ADtzxTQOfXHSlckMBn8Kb0bNAACC2ccjrTmUPGVPfpimYKqT605AAmTTt1QyuIzkAjjsPSnqPLyD1qwcY3YqBjvfk4pqVxEgHzZxTX5O31pw+U+9IWwhpeYyleJsGFGaI04RTyAOasTLmLOMk9KrwOxdVI571vBtxF1FnQ5HHSqFyVUnBwT1xWo6/MT37CsyVVEmWHTrWlFq4MrNal4wQAPTmoIYjFN1yDWlIVwc8g9qpqGZx7V0xm2mS1YuW8IRuR8x70pGAx71KoPBB4qOVSoPPB5zWDd2UNQBot2ORR5qqyseM9qRGHLLy3emqoLqW6im/MQTy+UwbbkNRuVE5OO/NNnPXmozl4wSd39apRugCa4EeGA5I4qvhstvbcx6cVZnRSV4yapllB3knjvitI2ewuoyJArEE4HpVyz4cqDkYqMbGRmxgn9KmsRhTnrV1JKzGlY0BjyAf0pG7MBx2pqH92xHUdKdHkjB71xeY2JJg4XoaRhyRjNOI3H3FI3ElNgVzhEIJ5FNjG1Sx71JMnX+dRqnmIcHp2rWLVrkj+CPpVO7jVUZx0PU1PGCg6daimG4HHenHSWgPYoowwwJqRV4AfqelRMDvHHNPSXJ5bn1rrd2tCdyVBH5rYHGOeacBIJAV47kGmqRuYgUPI2CQTz0qNWO45tpjJc8k061kZW2gjFIFDrk9qWJCG46/zqW1ZodncdejO1sDr1oQ5j9fepZfmjK9cdqqKfKOTUw1jYNmTLIoYIc5PWpJQ/CAYPrTFCLIJSfl+lSNIJs4pPVjBC2xgfzpsg3qAOg60qeWpK9Se1SRDaOR97rUvRhuTKq+QAOmKcqtgLmkXBjwOPQUA/ICOo7Vi27jB1znHUVAVwnvU+4HJAxg800gFenU80RbQMz7j7uBUtmN8ZxxTpIgwPalUrEvYetb8/u2RNtbi3YdYsdF71niM+Zj071bu5g6dfyqrEP7x5PBq6SajqDeppWyhYcDqP1q2LySJcfeHvVO1OMjPSkuWdVBU96i8lPRjdrGtFqCYy3y8d6uLIrrkHINc8oyBu9KsiRo1yCRitfrDWkiHTT2Nkj0pMe9Uor8cBx26irccqSjKmuqFSMlozFxaFIBBqMJgH1zUx68UmO3ertfURB93INVJCXfA6VcdeMVAkeJOaljQ+FMU9zjikXIBprUweohPPvSgYpmfmp/cUMAHTOKQinL0P6UoHNMBF4NSimYpwzTRLHrSO22lU49qgnlCxsc8027ICjeSGQ47VV6DPrUznI3E9ahccdOKwle5rHREFwdsLGsrbgZ6VfvnCxqhP3ulUf4Tmsl3NIsQAk/SrdtEWznn2qsqHP8q2LKHEdZ1Z8qKJFiyMDintsVxkdqlUBfwqvcAlS6CuRO7sN6DYrhPNORj3p7OUkU4GKiWJHj6c5p7xbol3HpWrSuK+hZSPDc9D3qXBYFT0qJTuIwelWNwBOOK5pPUow5IRHKfQmrEECxyo0bcGrk0Cu2cde1JDEqrgCt/a3iIkdeVx+dMjJZsHp2pYg25geg6e9C/IwB61jtoGg9OM4py8N1/CmngjinN9zAqGAnG75ariPa5JPB7VOBgg0x2xwBk1UXbYBSBkMOBULnBJ29alA2xHuRTCQQT+VNaB1GwsX+bGPapQQXJbpTYCGAIGKcD+8IPSnLcAVsyk9u9K7YNKBxupVGeo5qboBiIA/B5FB3MSO1PUfNj9aMYPBouBGePlXkVGQGYgHjvT2YKPSm5GcY5NUkAzOCMngdasJyP9mqzDZkt0NTQsCAac9UBI2OPWmMcnrTmPynv6VFuYtu21KQEg+8FzSNyxxTY9xyx7VIW2oWAyaHuAwjB44x1qKZQzkg1LId0ec8Z5pjhQg5696qLERAfNkHp1FLGcMfQ0gHJPU0oBIxjjvVMBWDL8xNJErMwJNPlAwB+YpAP3hHpS6AK6/vD2psJ+9xxilk5Lc1FbZU/N0ppaDZDesxMajvT4IdrYxjNPuFCyB26LTZpT5JkXp/KtE7xSQtL6kEsgExC1MF8yEf3hUVrCW2ueT71dUBFIHWibtogVymI2V8960F+Y4/Sqw5JBGamgQrluuaibuhiFQxIPWkDcYPSnvkMMd6Zt7GpQFd13Mx9KjPCipgAN2D1qAfK/Oea2jqMev3Bmlf5oiTTVY8dqWQ7lwKOoWI45VaJlODUYYiM88mkNufIL96bGB5JJboa1shbkys6bcdT1qV3+UE9c1CX4UnvT2kDpnrioa1BkpcHjtUiHgHOapoxJ25zV2MZUe1RNWFclC85zSSoGZTTxk4HcVFMDvUj9KyjuMUtls96cwyf50wg7ht6Z4p7Aq+Rz7UAD/dBFOfjHpSEndjHFNBJ2g9aXQAfIB9KapAH1pxOWYU1epAFNbAOznGe1KCSM44pP4cAUDhcZoAd1BpGHpSjgEHigDA+tICLIBB7USD5S56dqR/lU/Xigf6k7ufaq8wfZFTIDjiplYBgMZzTCuSMipQMyD2Fat6BsNcbeMcGhvugHtT2wzcngUjLnINSmBWmyZEx9c1LC24/WmnJJHYVJEPkz3q7+6IkQN5hJ6dqkyScClycLQVIBx1rFvXUY1yeBTs7gB6VG5OC3XFOQ/L7mi2gCPg44qME+bg96e2SCMfSogm1y2fwqogOzukI7U5xls0xSFOT1PelQlj/jQ0AjKFDEdaUcoDnmkYcHd3pei0+gWEDbgR6VGWPm4xxT3wpyO9QucsAD1qkhNk5GwADp6VHIuJQFODT1yWx2oxySeopJ2AiYHzM9qepOSDTTzuPShW528mq6ASMRnGKQgl8E0HqGzQ2TJnsKlaDGN82R6Uo+9z+NNY8EigA7Af4qYhxyJjnpThy5zSEEsvcd6VBgE96XQY1OmR09KYCTKeOKVMjcO9HITjrVC6DJc5yT17YqJcg8DHrUs2cDPeo/m3k9hVx2C7uOLfNnvT+NzHvUR+cg+tDjLjApWGErZjGM7h196Y7nCufuinMCBlaQoZI2QniqVib30JG+Yh+mB0oBQIJPWmZ3qhHToaRhtVU9OlFgRP8AxZx1pxLK5PUVHuYwDdwTxUyjbDis2UBwwBU0xjtbGM804qTEcU5RtGTzS2ESHpSY457UE5FGMkioAYwxwehpeqcdO9IzYxn8KP4cf0qgFQnaRjpTycAmmJyCDQSeQeaVtRj2+57mmHOwZ7U7O5Mj9abjIwTwaEAoyTjHSow25sYqUAg1GoGf8aaDQQ4BphTdTm4IP50q5I9PSmIOhyB0qRF4zmmjrUoA2dfrUtjG4DDgcDrTOvSpSAq5qLOAaEAhGDVeUHI9DVkDIzVa4z69auG4mRZ2jBOaYw5PHFKGCggc03OB7elbICMcECqUuVkAyB/SrZ+Yniqk65ZWAyfc1vDcTLKDuKjZiGwf/wBVCEnn8vakIJJ4z05zR1IYrnBGKa+XIPrS4ymCcCmYwo9+2KpWC5CeDgj8KUcHIFD53bSeT3pDx0zzVj8hQSBk8kdqYThh7c5pctke9I2VIPHuaYJGzb42DuKsAfPxwKo2bZC81e6OMflXHVVpFIkU85plz/x53H/XNv5U8YwT27VHc5Npccn/AFbfyrOPxIDsweDSjOM96D9aD0FeKUKcHFNYA4oHINA55NCAZjDUrDjinDvntSDkU7gMCkDGeRSqhwT0pVzyaXBC+maLgRcgEmnpynNOwCOBSlTs4/Kne4wQZBXvSg4ByKQDaM9D60buOnWluICQR70vYA00E7s07+VIA5BA7UHP5UmS1O3cUWDoNX5sehpw647Cjtx2qGF3Zn3DAHT3qrX2Amxk5puAMY605uMUwnJqUBIAMjnFKDzgGoy3QU/ocGnJW3AUjBGOtIRgc/nS9KMAnApAHRhnt3phPNPye9Nxw3rT6gGeetMbGKf1IFI45x0pICAYwabH8zY9aYzcsv40QZAG7Oa6be7cC3t2ucUq9xSDk5z9aVcg5Ga53oAmPlppBFPyCMYpp69KEwE6Lg9KQgbR9acV5yaOCBmncLD8kDJFNGNwJ60vPIpOCPpSAdkdD371GDlyMU5ucHORUTFg67R9aaWobD8BQfSl2ADilxxSNynSnYBgbnnvUEhIlA7GrATA460NFlTnqKcWrgQROqyiPGKtMMkCqxiPnBuwqUSZl2Z57U5WeqBEhwDz3qu7bWwOtTjvmq8igOD3pQ3AsjG4Eeneh+KZkH5gOlODbkwRzSadwF7Z/SlA54HXrTFPBA6CnMcHrSGOPJx60dgBSM4DBT1NOUAKSam1hCZ+YUvUkdqbjpThkNgdaOgAO+aYBgH9aevB560f3s0LQZHGmEFOLbmDGl5VfrzSEYGabvcQucrz60rcqAODTSDwB3pXIIAHWkAoXC+hFJ3yKViaANq0gG53P04pxwXoHB5pOrE9qBgANwxzmm7SPlFPXjOe9IBy1O4DWbIx6VC3QYNPcFM9zUYG7vwKtCHDO0kdqUEsgPehDnA9afgDA/Om+wyNsgYHBFPXLLk9aXYPM60L0x1zUXEhOMZ9KU54A70FcEjHNHcHvQAh7AU1uOKUDJJPShhxkimBEwwRTyTkj2pJFyFJp2PmDGqvoMGX93gcUY5oY5+lKvJx60mABl3bQec0vbBqLy1jmDY5NSSfLy34UCHngUi9M/lSjn6U1Tyc/hU2YxRyD6DpTX+5xTwcZH500jjimCI40MaKM/jUiYHPWkH3MUowFAFNu4hSMg0i9KH+6QDzQowoB7d6WyHqL700fd/lS9uelOACjJoYDSQWyKaTwR2PelTIZs9O1AGWyelOwiNiAcDqKcemaR8ZLAZHakTJHzDFPoMV2w1PPT3qiBKbhn3fL6Gryc8kdaqULJMW4jgZBPaq6uXkJJ+Wp5DgHjIqsf3ZI9aIoZMcl/Y05BjI7EU3GYwO9PX5WAparQQ9sDiox8xx6VI/BNQq4LnHWpim0A+XnGMUxSC+D0omJUcCkQAxBu5qktAHHG/B6UiKASfalwQAxFGNw4GKE9A6kTjcAR1rLnLCbceh7VpOrMw2npWfcFUBJ/P1roo7ifcYx+RiDk1CpCgcZNT4JTgdTR5YVzgZz+lbp2HYtQqfLz7UyZd0JGMntUsSlYwOlGw7cMMkCsG2pA9SijBCcc46insWkQMo6dTTZgN4PTPapY4zsIBwp7Vs7NXERtHkEdz0oESqQAKncBNpxkihQcFjyTU82gyjNuVuvy1SkUBWOMr1rRulJC7TgmqpwkTFhn2ropvS5LId+1Gcj5RV20fzVHpVRseQxXoe1W4V2KmOTRP4bDW5cyQwx0p4yvb6VE+VUhal6KM/jXI9hi55bimbjjJ59KcchCTUDyZUAggg0KN9gFlXK5zUMaFgMHgVNKoZB6imREnjPQVpF6C0uNmzgbTgVVOWUg/hV6ZcopHHpUDptk3cbQKuDsgZm4JJ3HkdKaqAEqDjHQ1O/BbPT1qM/dZh2rrjK6It1JUJCYH8VK/C7R0PaoIwfM68CpFYfdz06moktR3uOiYMNp5NPUYlAB+YVHGu2cqDwO9Tl9khxjJpN2d0MkQ5lJI69aieMGTJFTKAHyRg0wMzTMm3g1kpNPQYKEUeo71CRscYJCGpWUmTaPunrTuQx4GzFNMYzZkAng9qtJ86qc9O9UmuIZHCN1FXYWUxntg8VM07ASk9cDGKUYppYSADnPenBwDkjI7Vi13AXHy9KhOB06085xz+VR7g5IpxQyMg4B/OqU5IkYE8elXncAACs+4YmUkcYraluQxY8cnr6UiDD8nvTIj7CpWIDAmt3oBaiYBwe3rVljuztqijBmBH5Vc5IYY/DNc1RajGx4Vm+bmpc71qBSF5x8x61IsnyAD73aplqNErMqIMngUyOXEuQ5Ge1JJ91c0Kqg9MEURdtRblmPVMTeU4yc9RWkkqyDgiufljUXG+M85q6jlcODXQsRypEOmmapAxURXJyagjvApCuME1YSSOUfIwNdUKkZq6MXFoYRgVGeanIwKayj61eoEe3ABp4BIpuMHipBxQgY3H5GngcCkPU0FgB707CFOMYpgmGcDtUbsSPSo0/WgLEsk5GcdaicEoxenKRnJHNNkbj1FJ6jKZAJxngd6bIMMRTsbWxTZAdx9xWUnZGiMm++adefu54qt354Iqe5Obhv50sEJkYZHH86yvZGi0H21uWOTmtqIBQoqK3RQmCKnYAR8dq4qs3JlbCOwXI7npUZDGEgHGKLrhEYnB71XFw0hCp3ojF2ugLCoCQfUU4L8xyeKai/Ivzcr1qRTtXDdT0pNsByBVPsKfncQQOtRFiMg8GpVBEYz1rNrqwBzlQDTYwF6nk06ThB71GeQCByOnNEQZOAFJOOneqbsfMJziry/6pSeSaz72dY5VWnBPmsDLEedwyaev3yM1DE/mRkjpnip0wRnvSkmmAZwwz3qMggkd6lGd2DUWSsntSQFYO6ZU8mpI33cEYNNX5nbPJFKEYjk9etauwDoGBZh6dKkQHqwqNUHmEAYan5JA3cVMh9BVIzz+VPBO4kDoMUzA3D3p4ztzUMAB6YoB6n1oxgAUHjHHWkJETx56/nUYQk7vSp2O5uKjLjJFWmwGGPK8t+NPiJHy0nA7UsOTJyeRTeqAlIxyKgzs681OcHjtUDfPgGpiA5AAue9LnkLilC/jR0fmgBrDhiRxTJhkJgdOtTMcKc1DKAwx27VUWIiOcjHSnkEqMDFCgEkA8GlXknPWrbHYkbAC9jio3IB9SakxlMHr2qu+Q5NStWAScdByetMQ5bOeOw9KkdginuahQ7V5H0rRbCJ5xmI5/GkSPdbBW5z3pzDIHoaUMVOMcY6VKelhjIRsO0/hSscBmxnnpVaScmXPYCmIZDLz901pyt7i0J2clgPWrSggBR071TkQhlPXbVsMRgnqazlsrDGTcsBnoaYSQ545qSU7RkdaYuOtJbARso35xnNQTfK4Bq0/LYBqtOQSAx5rSD1GHAHHWgtgHBzS8E9Oe1NGAjAiqVgVhg3upC9O9QRR8lScGpo5ijdPwqJyBNkdK1V9hW1I5nHygDr2q2q7VUY4FRmBd3mk9etSsTgBRSbukkIZGwaUleQKuxtg+3aqqgK5JHLVciUbRmsqlhkiH5uPxprfIp5yacPkyKjl5RuaxW4dAU7wpHSnk7ZOelRW6bYz61KMYGetU1qA8gk5NRx7SxNPZzu5HFMIEfTr3qUHUQcbu/pSLlcljwaeTxwOKbkkkH8qYhvKkenapduTknioi2GC55NS9GHpQxg33hRnc1L/ABZprYzkUhjGGX5qNTnjuKmOCwxUXKnnvVIRG38Ip/TAPegD07CgfNjNXcEKRxz1FI/3cCn4y20dqaRweOlSmBC52y4A5NOjP7wL2qCf5pAQcE9akhIV+eT3rW2gi5nIOO1ITk9c0sfUgUgxk46DtWPUY1vule/rSDJAp2OSKSMYJLUX0EObOAKrP8si9xVo/dBqtIp876U4A9RZMcY/CnAYX3picj6U9TtGTTYxGHY9e9J0A5pTy5PpSgdSenpQA0jgjOTmoX+XDZ5qbop9zUE2dpx1q47iJFZsZ9aWRuOKhycDHGKkzlAx4IptdQ0sNk6AjuacWwQBUb43KueDTufN/lRYLkpwFwO1AwGH0pTgIPakJwBioGhrjC7e9MySQtOcZbOKavzPxzVrYQ7J872NPXCk46VGv3m4pYsknd0qXsAgU7m9+tNC8cnnsKk6sSOlR5+cGqQCSjKehFQbiQc96st056VXyBuGeneqjsDFDEBQBUrErzjJNQKSzkZqSP5+o6U2Am4JKVPSkVMM5BoOJm3EY2nil8wLMcjrR6BfqMxvj2qfmFEmfLVv4qVPlnYt17UgYtNuA+U1QiRpCqgnn2qRHVIhuPWoo8qH3j6VIv70YI5FRJIZPGdvTvRKMqcdaEIbgDgUE5b0BrLqNakMDEj5uo61Oo5Y9qjTBB4p+f3QxTlqJDZF6Z4pWboPyokU4z2psnIB9KEPqO9AtNfhsinLgYOc4pH/AJ0LcQ8EAYFIy9D+dMi+6c84NSt8ynHalsxjc7RkdxQB8uaTJ246U4H5cGgCJhmmoSGK96kfuR1HeqysS5Iq46oViwX2nHc01ZkZ9meajlY7gSOtVHjKzlyeneqjBPcZpFwcjuKADnA61XjkLknOamHTNQ1YW47JIGarXH3cVZ4xkdqgl9e4px3ArqhBIxz71Gc8+oqTfvCkdTTOfqK2QXIAMuR2qvcZ357GrKgMW45qCc4U8ZraL1FZhFyc9c04kbulRwgsBz35xUjDkbTg55qnuTbuDHa3bHamPJnHGFHP1pWQYznpTHwV6de9ONgvbcjbhvT/ABo6sMEUr4FNPG0989aoFYXjPAqN8E4DcU45HXHvSMAq4HJHtTA0LQ4UH25rQOOwzWZatmPPY/pWkDnAzXNWWo4koOBz+VMuf+PScZ/5Zt/KlXOeeaS5x9jnx/zzb+RrGPxIq9zsyQTjH40o64NJzzQMdq8QY0daUjFBIzmkJJ61QCsMrSDIT60PnFByMEUAAGB9aXtz+VJwOaB94mgB2SBj1oxjvxQPmJPpSk/LjtSs0AjjvTEO5tvQipSOaaUw+RxVJrqGzGgEZpx5T8KXHPtSDADD8qVwFAxSZxmjHB9aafvfWkA7kNg9xSjrn07UvbnrSDpyeKe4Ds7qYVJ4HenAClBAB/SkBEibW5NS8ZHejtnvQAM5ocm1qAvBbgU0DNO7Eim9PqaSAUZOfypvOcD8ad0br1o6Mefxp2Aa4II5pJAeCetOwQeaWTO3mhB0M1uXcDrSbtrbfQ1LIu2fcBwetMfgE4zxXWpLlEXFO0Djg0/PHTAqBGBRe5qbouO9cslqMMfL70Z+ajPyg4pMZXrikAnfmgYB5pSAcE0DGee1NWAEOQTQAAPrSc4OPwpxHHvQwAgjkdqaVz061IRgjnimj72D+dCuA1SSSD+FO9RR0PI59aGPcdaHqAnORjrSleTzSglgPWlJ5P0pAQBcJg1Ep+8+OnepCWXpTtmPl9a0Wm40IGLID61C4LsQTggVYRfLBFMx85b14pxaUritcSDg8/dPGKkDAsyAdO9Rx/u5COoqcKAxI70SaQbiMnU9DTTyckVIG9T9KMZA4rO4C4BYMRyBQDuLAUmcZOKUDAJ70twAdh3zTed+acRke9IT8vSgBWYEjHWkbp1yaD644pGpIBzA8evpQ3IwB7UjY3DmjOW96abQxB1A7UYJ/CjsTjpRg4xQIUEluRwKXrnnjvSKfl60o4GaTACCWAob9aByc+lH3m4600Ar8Ck6Lz1oY84pCSGI7UrDI5V3oQeM1WgfLMlXH5UEVTA8t2yeTW1PXQXUnQbcetSEZYDp61GnJHtUjdRjrUy3GNbhhjmhOuaBgZz1pMHd1pCW4p53DvQpJ4o7k0J1NIYo64pGAweaXJUGmsMITTQhrHcM+lAy2MmlIzGB+dRuwXK4yKpdgHtzgUgfJG08CggBB60qKMcdfehLQB5w59xSSZOB2pIgRy/Wnk85Hap2YAOTxSgAnimZx8tLnaePzoaAUZL+2KQ9QKcR8uR0ppAxk0kwBh1INOAAGfzphPvSqctin0GLj5qB9/NOHU4/Cm9PpSACM8UHkAelB44FNOcgg0IAQjJHXNNJKo2epoGMDHelk6FQKoBkefL5704HC4qJdwyvepPY9qpgNeNXz79KkHypjvQchaXAOGqb6agJIBsAPWmFRnJp23djPakIBO000AgGB7U5QN2e1J1BAoh4yDT31AfnPFVlT5zj7pNWCeTxzTG4U84pRdhBMQQMCo4m2rnOeaecGKolJ3DjpVx0Vh6kyncT3p3bjvTGwE60pXHWpYMjk4GfzrPkRZmzn5ferjhgG+br0GKqSjy4yOrHrXRTVtSSAhgcZ6/pUozuGOaiRTtwzcVJtMUOc5FbSb2YIm37QCOSKlxmTIPBqpCTJz+dXgqsMg8is6iSGUhGochhuKnipQu0c+tWNo3jjr60gwxYdu1S56ARSKC2SOKIiPLIzzUhHIBpqptTGMClzaWAp3CPyUGPeqTLtX5q1JwrJtP5etULhSEyo6Gt6UmS+5DtG9V7DtV2LEjtxgiqygNEcnn3FWosKqEfjWk2UWUBIw3Wl6nFPQHg55p20E7hXG3qMicfJiomj4XIzjmrLrkGmKo8sHsKcWIinyR6VBGhLHP41bkXIzUEeA4yM1cJaWEx4ClQCOlVpeYyCcfSrRU7lPbmoJhtQelOL1BmXMuYs9s0xRsiwe54q3MFUDcOBVaUbSOO1dsJXRDSTIlx0x+FWIyF4YcntUabA+W61NEplUZ+XHeqnJNDTu9BoIQK/epGG7kDr3qNl3YOOBT0UyAEnA9Kzb6gSElWJHNTOVwjDhqrS4QJxzmpnZQAScH0qGiriXBJfA49TVSWTI2Z7cmrM7ExhhxmqojLfN1qqa01E97ENvA7zYH3fU1rR4SLBHNZ43pkHp6VfiysS5O7NFVtgkWUb5B701htUZPSmw5xnFOYbzk1yvRlDjhhuFRFcZx+NSAjIXtTJPlPWmtwIXTGO+ao3EYD8nOa0CyudoNU7j5iBnkGuinJ3JZBGN2RQTuYY5qxPDtiVlH1NUwSrHAzW0WnqIuW8hB9farwcMD6VnwYY5/Or6r82O1YVbX1LVyJUxx6VJjEfH3u9NY7W4pTu2DbWbDYkjxJEGPUU/aXbNMh2iPing/KQKzk9QRDKjMRs4bvVpQTGA3UVSZm3Kw+lXVOYye9Od0kBXmyHB7UiyFFDBiDTJXxNtxn2pxUbPpVrSKDcupfsFw43VZS+tyPmJU+9ZXXB7UxwS2Oo7VtCvJOzJdNM3A0bjKMCO2KaTisjLKF2McD0p7XUiL97mtliImbpvoanTJJpCaorduV7E1GL2UqTtX6VXt4h7Nl89etIcf/AFqoR6hI4JKgGp0uS2OAKbrwW4vZtkx4prNTDMwzwMU+MrcISnBHUUo1Yz2BwaIGwT1zUMjhAWY9BVlk2kk1myt9qn2D7i9feplKw4opxRNNLuPc1pxRKhGKSOPypduPlNPJKEcVyTnd6Gy00J1HIIHFTOOML1/lTIsbh/KplYksR2rmm3cZWu4DJb7e/wDKqcMWOD07GtV1IUg55qpOFgQEjNaUqjtyishiSIjcdCeRUnmLJNsFZzFXYfLtP1q9bx4fLccVco2QIlILNx+NSLk8DrTJDtBYdqWNsoBjmsHsFxWICEY4FQxsXd1IIWpmwqkNyDVMyyC52j7npVxTYjRRt0WCcZrDvG33DD8jWyw3x4U84rKkgI3ue/StKFoydwlqi5ayMIwOtWQTsHPNVLcKsAHU96tKdsIrOr8Q0O/h3Go5BjaSafjKD1psgyg/nWa3DoIqqu4jvTFbgkU7qQB+NGdpwAOKoBiAtID+ftTnywIHam84Ldz6U9wQ6lRkHrTGNzs/hzUoLE5x8vemMeevWpEOQTnipYg9DR0PSjotIGNSA0j5+tJxvwRyaSQYBOaFYMA2OavoAhGCc+lIo6N0BpCzbMnBo5ZFqguTgjaeeaZjnIFLgYxQT8p96zAXPz9aYxHOKUnGMdaCcLnvTAMfLg1DKA38XGKmP3Q5PHaq5HTPPPFXEBAfQ1IudnSmMMPgU7oMDvTeoEvBwcdKi43kEU85+ULyBSONrHj60kIhmAJLHtTDtyRyakPLEt+BqFeSSa0joHQsHJTimKBtBJ5zTkyY+KMBgFpdwIREfOcHpSLxcbAancZmB7VC2I7jIOPUVad0A2ZykgxzmrEbZjDHn2qvNNsbgZY1PEMquep7UpL3R2HTY2swPPamRElDupZWwyrimKcSEZ5qEvdACBkk9arXHCLgZNWTkDFOaMMgJFXGVmPqQDnbjrjrSFflIB57UsPQk9u1Kp+8QKpvURTJ2k7uD2pFGWz2NSzqMBhjBqIEsoIHArVaoCXA8sKTz2p6KQdpqrK5fDA421oIAQr55I5NTJWWoJ6iTLlR61JGf3YA602VfmHNEeV49+Ky3QFlsFAR1pkuPLGO9LL8igEVEAWGM8CoS6jFhOB5ZPPepAOc54qvFIAWyuDVmMkg9qc1YQvU5IprYZxxyacpJB9SaRcByc/NjioQA2BtApFA8wmgr8tNVsAepp9A6iYBkD46VKmOTTGySacn3D6UPYBQfl96Cdoye9B+Y5NNPI3HpSAP4f5Ux8ZGetOZvmGPxprAM+R+VUgGcqcGjIyF/OnE/OPSmL97NUBKow2egppGVOOp70vO4E9qVuc8YHY1IFKQgMq04He64FMlTOM1IeAuP4a36CLcWGOR1oxleOopI8ZyKUdSDzWDGL1Ix19aaAGBGaUdAOwpwHBOOtINxB90YqKX5Tip8/IOKhkPIOOtEdwIxxHg/epQdwwaaMlj70iA72960sA4HJ570EdR3FGSccU04B4oACdwz6UpBKlh3qIsGUqOoqQ5VcU7CGjiMHHNAGEzT85Xa3WmnDDk8jtTuMY54UYohJYgHrTn5kHHSkC4uAT0p7oW5Lw2cDimHpxTxg7h6mmHhTUIewN905PFRx4Ryewp/WP3JprMFQg1S7CBARubPBpIPk4bvQmfK/Gn7cspXpTYxyqF/wAKjOACae33t2KQjHXpUoQzqtVipUDJzVsY8s46VWkzkVpFhuKBs59aWIsgyeeaTAYdfu0DJxjt2pgOYFFLDnPSkdAwX0XrS/NwT0HWkUnLMBQA5SpBfqelMAIiYDk9qFwIyV/EUjkiRXBwDQDERy0JWSpQdrxheh6UxMCQq3rx7VIF3ShT2PFDsJ3JoRtLEHr2p7rSIP3pI6GlbIU1i9xjFOGK0qnIx6U1eCfanBTg8/SmwH54PFMIDAe1KDke9ICCCAKS0AQDcfcU7qR3piA4JNAJIPrTaAkGMgYp4xnjpUOc4xTwduDUtDEKjt2pTnIIpQDkjH40uflwaAImB5AOQaqSqyy7l6HrVuTcpFMLALnGc1pF2DcinYsq7RnHao5MsuD0NTgjPA5qAxEyHd0NXFkiquwAA1OuScfzqJU59h0qZSdvvUyGP6KfUVBM21OVzU5Py9MnvUMmGyp5pR3AqkjPyjBHemgkZHc08jYcjp6U1shjx1rYEQjljkdPeoJwQmT+VSrnzGzzTZgCpOCR71qtGLZMgjAx/wDWqVvvZ5NRRkBTx8o7VKx3dOlW9ydOo3Ayc8570xsk7ccDnr0pzEkjOfek4weKYDG+8STkZ5pmc4z+NPkxuxjn61HuAwcYB61aHqK2OqgetNYgqTnp+tKMDI6ihiAuMcUxdSzZcJnrWsp6cVjWRw5Xt7d614W4HNc9cpbkwqO7P+iTjvsb+VSjmoLnm1n56I38q54fEhs7YD1pBxxTiOM00Z6HpXi7lAeMc0gHJJpkr4cDFSHjiqs7XENJ4zRncRjvSkYXjvTNwAx3oAeehoBGORyaP4OlJjK5pIB4+XJz9aX0AppO1eR1pw+6CKB3HEYGO9BBIoAyM9aCTggjpUCGhsCgHrmgDkjNBwADVARmQMxReT61J0yO9IqquTjrTgCRmm32DoH8PIpOCAO9OOQuKFU59qkAA5I9qQHC0vqTTeq/jTAUDNKD8p9aOwFGMYNJgKQMgDp35pMZbntSuMHg0AEHOaH2AOpppA7U48k/rTCeSBQAuN3WlOT3pP6Un8JpgRSp3qtGh8s7u9WpBvHHaowuF255HetoP3QK+5ouF5wavDlAarhMSsRVkD5P6Uqj0QW0F6qPak7DNKD8uM80HHFZANYbhx1poYbiMcUMcMfemDJeqSQEqcnBFOPBzTAQzYA6elSP0461L3AQZxStxx2peozTT92jYBpYg4PSlKg45pGcIAGwc9OKVjiqYDl5IBpSCOO1MQkHIp7fd4qXuBGT8446dacQN/HU037xzTs7irelMBjjORzzQCCuMdqefvZpm0hutNNARfclHORU5IU8HrUTIM+lPVVyD1py1VwAkgr9akycU0YUkUu75wMVICnO3A9aTBHGOtOBBzzikXk+9IBTkc+1NIwpFOPOVpp6A5pABOMA0pHzHvig9ckdBQCMnNMY0dSxpw5yQc0gwaP4TjrQLoIRxn1pW46Hig9MZpX44z0oYDSMAc0v8IPrSEAk89qMgDb3oAcOFNCrg9cE00HC9KeBuG4flSAD96mHqakbkBgPrUEj4YelOK1GLnIqBiOrHmpmOACKayAjkc1adgHLgJkd6U8nNIoAUKPxp2flwKUtHoDEPdh1pjHrjqac3THSm5AfHehB1H9APWlQc8d6Yw4J9KeCQoIPNAdRG5bFD/dxStwc0AUCGN8oxUTEb+etTnBFQsBknuKqIDypyPSkQfMSenangkqKjclOgpeQDx82Qe1O4K80fw5pCCcCkFhCMHI70oBxjvTsAkD9aQjJI9KQBnI9jSH7pzS4wABSPyOaEAncU9ccmm4wuO4pVU7R6etAwXjJ7Ui8kkjgUpwfxpRwGz36UANY5Bp3GMd6bzsyaUjjNPYQ3B8wDtR1Yk80DpnNOxxu9aEwZAQSc56Uqglsk9aVvu59KEIHJ/CqvoCFB/hxS5O7j8qRSM57mlQEvntSGIxwoFNBw4yOtPOScGony0eU600hDn4A96ZG4APqetPbKx5rKVpFnYfwCtYR5tAvbQ11YuMioZchWyKjW49TT1GZg27g9qXI0wI47hGTC9R2pGKEZz05IqSa1V5Pl+U57VF9mZZCucrjrWsXTs77hqywhV03U9hkj1qNECRqgqU8cVg7X0GQPgZIHaqmdwLyLjHSrkgOCcY9KrzIQDt71pBiKkZBBUjCZqWYhIht5U1GrF4yCPmzVllVbYfzraV1qxdNCvF8hAA4Iq3GxA5PWq67jOC3AI49qmiVt+AcgUTsxkof5xkdf0pCAspbt6U5/lXOOTSKwOA1Y3ESEAoWA60zAMefengl1YHgDpUSDBKg/KamIxsq4UEVnSDaWBOQe1apXAx6VQkQq4z3ralLUTKaAlWJarVsNyc9R0qtlQ5GcZNWYFBZiD90cV0z1QIvxNldv609sqCB1xREoMZ9zUgwF964ZPUoj/gBPXvTCCUK+tTLg59KYVyTjrST1EVSzMCDwBTYyDJtHIFWHQeWeKrDIkX3rVNCejJivyAfrUMyDYw/u1aUkqaikX9yzDknpSjJ3GZ0mNgz96q8zK2FA5HQ1NKD5oPXHaoZCD8zHAHeu2naxm33K/dc9Ksl0BAJ7cmoQpOMHBp8e1SAy5rVpBqSvGzJhVyCaVIdw3E1NEuVyrYLe1SIAW2AjcOprHmtdFbkDxNJChP3c8054gNpOMVKQUG1uQKQqGzk8Co5nYERMq7eemOaqxk7iByBVgYMTHO4ioEZt2CME9M1pAbJvJ3qezVNbx+XHhjmmRu7A8Yx3pzOpUZNRK+w1oSxAhiD1p5XKUJj5sdqFYlfasZN3AQKAgFMOSMkZxU0oxyB9KgkZhEABz60R1GQMGU7hxntTMbXwfzqflU4PPcVDMc9Ovetot7EsWRd6kZzkVUMChAQfwqyjAAjH4USBEQYP1Bq03FidggjVPxNXCAyA4x2qsZkiUAjipI7hJVwtTK8tRrQWSMBVbvTzjOT0pzjcgHekUny+nNZXbGRKygZzxnpUg4yRyKbGhIOex5pEcoOmeelP0AnwAG56c04fMSw4qJ8qNyHr1FSg/u+O9Q1oGpWdRnJpAfkIHQGnyDzFGB061GflOwjINaJ6CHR8p14FKR17A0mdhAA605wDgjik9xkIYxqQOcVBub5mJ5Pera5bIxxUEn3cdM1pFq4rBbbixYkn0qfbzjpUEBZS2eeetTZJlxSlqxoTythIx71Mu1QPalJwxB70H5VAPWs+ZsEIzZU4/MUW8rQktnp1peiDaOabsK5yetClbVD30FuJzKhVGGSOneqdnCqZ5zip2Hyc+tRxqQ7MPxrZzckyUrMsMMnOeaCdxBI5pScrgjrT41Cr8w4zxWHQaJkbChyOtTD5lGOvemRgBfmHA6UqN3HQ1jLUB+dwyegrIv/APXDHTua1WzVDUE5yBkd6ui7SE9ipEqnBYc/WtFGXev6VRWNgmSMD1qxbuH6jkVvUVwJ5S2084AqO1bzFJzkA1NIokgOepqvApQlV5BNYq1mPZliT7oJ6VUmm8ubaF4Perb/ADRbetRhVZcHtTg0twJrcho9wPJ4FRbkLbCOaIpRtOB8oqkbjNwUPHNJRcpNhoi6qhcjGAeMU9chNoqKZzvjUdTT0dWLD0pNXVwHKcPtpWxg7ulRrkOQvHvS7S0hJbgdqVgEzskHqac2BUe3DDceKeOSPUdqbsAEbH659qe4JOO1DYAyetMVyzHJ+g9KnfUBVzxnoacuV49ajQ8DOaec5znIoYDiMYHekOScY6UsnODQCAnvSBkcqk9OlRgjgZ4qRuB/vVWyAOOoNaR2sDZOvzdOnvTV5c88UKcsM9KXaqg0XAlwMcmhjnOKbEfl3ehpcnPtUABBwAOvejqB6jrQM7/c00NhPemDHHggdRVU5adsn5O1WWOBg9aikUlRjjnmqg7MTIt2XA/iPWpFX5if0qHcFlANTKOw6mrktBk2MAgdRTG4Yg9aeOZD2HrTGABLHOayW4ELbt7HHNQnJLNn8KsMcZbvVfeHYk962iLUnh/1R96IjjcO4ogI2k04Dbnb1pN7jGyhvLB7iq06sSGzkmrch4yT0qjNIxkGPudqqne4tmR/OWGBWjCAsSk5zUCxAKG71NGSYgzclac5JoNRkjbJsVEz7n989acZN82e9MbhsDmkiuhKpy+D1qXevQn5ahPOAPTrTJPu0Wuw3JUwGIAwKQHCnNMiY4yaUjLHii2oELYJ246UxRtyAeKdKNr5HT+dMjztOB+tarYBgUYPHTvV2HLRjjAqsRuGFq3CCI8d6mbuhBICXpI1MjemKecqxJ7UyFyCQ3BrNXsMlkO8D2poJyaQt+FKWAx3zSEMZS5449anTd5fHWqyH98cfjVtTgAjpRMBRk4wfrTWXMgI4x3py/KcnoaYSeAOM8VC3AGyymmjnBpzjGVHSoiGBCk8mqWwEj8An1pyfLH9aawBGc9KXOYwQaXQQ7tikIyMZpSQCDSEhVJ/KkMaSq5NRF9q7v0pzn5DgfNUco/dZq0gFJ+UsPypD9z60u44APT0oPJA9KYeRL91MdTSsPkA9aYQQueDntUhPQ1IFSRCScdAaaQGTp0qYr1HaoXyxK5IC1qmFidR8oIP1qUD95kGq8TExKOvrUxJBU/nUzWokOxxk0gJKkCnZG3rnNNRs5xyKzAUnjAFI3IBzxT1zvx2xTXXKn6ULcZACVHB6VEsu12Bp6gsjqKjSHkjuOhrZW6iJkYNkZpD90mkVsHBGCOtIDvX2pWGR7SsmeoNObCkYNBUlgd1DOAxOOmKrcQ/GXpsnO3bTlwTnpnpTOiFfWkg6Em4ZB7Gk25aq0blWC469easkYcYPShxsCY0EbjzzT8HOTTf4wR+NSYy2Klj3GY+c+gqFyQ+asDA3ZqB13lcdKqO4DnXMYpRwijtQ2XX6UE5AA7UALjrk0jMDGWoY8gnvSMF2ECgBFOYsKelV2xvqePkk+vaq8nCnA5+tXHcnpqLF99hT05QgdaYh2n+tSINrE02V1EJ3KOOKdkBiPUU0fPGw7miL5IgTyKQkxsRCxtg5xSD54xxjFOACEqBwaiLB5lxwR2qtxWHr+8+Y1IjAhCevrTAoDkhvlqUHfFuxwOlJjZLGdr4NSn5aihfzEVutOdsA+lZS1YxigZp7E5HFQwn96STxU4Byc05KwhpOMEGkAw5PY0oC07PJUj6VNw3IlzuYE8Uq9KF64NKuAeOlU2MUAjjvSqRkikB+ahRluKQiTqT70h9D0pBz07Up5qRjH6DmonAG0nkU9j+8xnimn58irWgg6ckfhSJznNL1wD2p38QxTAawyc05TjkilByT+lMPBx3NIY7kgmomwckd+pqQncpUVDIoAx+dOIiDgjB7GmspyeeDT3TBBH41DkgbetbIV76DRwfpTHB8rjJ5pR80n1/ShuIjxxWi0AoHbvxk88H2qwhwoA6GoGA81qlXGwYP4VtLYlroGQzYHSkUqVIpQ68gDjNIVCp0waQXI5MlzyB9aM9mGaWQhs5600MMkkcHir6BuxCMHOcD0xS4IXd39KMkEnIPpQp792oGx1uTu2ZySa2YuVxWLbn99n8a17c/Jjp7VnXXu3BbloHjmobkD7JP67G/lUoIwKju/8Aj0l4/wCWbc/hXJDSSKbOzY5+lO5wKQDIx3oycc14zGIcE5NIPm5zQG9qFPoaB3Bm52nv0pq/6w8U8/fHGTTeVY+9NCHEcYpcgYFJ0/Gl71ICj5jz2pM4ancHNCYyFJp3AVASvXGaXBHWkAw1OcBTzzSAawGcigfdx37Uv3QKDyeBxSTAQ8DHcUoPy8UEhsUA54xinuADJyfSkA4zTjwSKaOnue1IBRkD1pCMjg0pBo75Hai/QBjEnGOnepR0wR9KafbpS5x26ii7AACenQUuSRyRSdOlBXGKQagOVIxzScdMU4nvigtjOcGmFyDzAH2Y5NSFflFN2qeepqXAIp6W0AiZcVCifOec1YbDAjuKgDbecdulXC4Cpw5PapAMrxUXJO7FTjhaKm4dBMZXPelHJOaF+6KRhznNZgJjBxUY+XJxxUpOGBpp4bB5FWgGLkksKlBBXnvUagq3HSnqc5GKcgHqMjihsYHH1poOMY7U8HI9qzYEZwzAHtQVBUg02aUxgOFzSrllBxj1q+gDQQFHPFPBwuaQoG4PFKncUnYBjqcAinoMD3oJI2nHBp5AxkU76DGAEknrikc5x609jt5FJ6ccGpENPqaRflA4oOCmehpQN0Qb3qgEAz160uSz5I4FKecGlxhc9RUgKQNwI/Gjq2R0FBHBNOAwMGlcBgAyecUgBCc+tOI+U8UnYU76AJu+brkUp+9TSoDEA8VJx37inoMaeCT6UEYHXmkHAOacCGHNIQjAhhmnfxZpM5I9qOvOaG9AEUDzMZpDyfbNKxzgilb7wFOwdBGPGO9Ge3ekcDA9TQc7gR0pLYBwyFbniq8uNme2anPI471Unf5CKqKd7AOLZXb6VIG3nkdKrxSAEE88VMTwCOhq5RsMkjPBz1NKRtWkXAXGMU7qAKz1ENbtUbYGCTzUm4b8YzioWYK+O1OKu7AP6j61KPuA4qLdnIFPX0/KhoYrHc2aSPuGOKX+LikThsd6EwA8A+tROduQvWpScnaetQyODgjr3px3ESD7i561GS24nOQacCGUY5pS3HTp2p7MB6t8oBFKOPxpqAcZ4pevyjtUseods0p7Yo4C+wpFPHtmkAvTp1pnUin/AMRpuOeKEAvalUnb9ab6YOadnIGKQBj5aQcjHU0pPGMUnpzTQCj5hgnHpTW6YpWwOO1DHkdKYCY+anNwcdqQDrzyaCwVx60gZG3cUijJxjipGAGT69KYTgrjvVgJnPSndwBSOAAfWhfvbehFHQB3YmkYjGe3elYgDHUVHg5wO1JK4CKPMXHUVTmfIKqtW0Rkjz/E3WlKdPetE0mLoRRwjYDtxipo4vl+lSAYXJ7UYwpJ6miVSTuGgfxHBpjdG9akHFMcY4HU1mgGMvyfKaFyV+alJCjb39aTGGx2q1roA1juA9qqz5OMVZcrtJ7DtVTczAcEZ7VpBAQnKsoUct1qZR/o7AjHvUEpKAKOWPWpwfNhxnBFay2EQK5IAbhxU8JWIZY8N1qonzszLnHTNWVASPDHrVSsC2LCnIxyfSoxMcZUbsdBT1BjTeeppQqodo4U1joDHRMWjznBPao9u3G09KkbCrheBUMbgHymY5Pei19hj33Zwo4I5NVJm/dBgMke9XSMDAFULhdqP9elVTethSKk5RZsY5YU+JykqheuOaUETcsvzj3otxkj1rqvpYEbEQDdOAKlVRkiq1uSBj86tIAck1589GUIoAyvSkPBBIp+MqajkOdvoOtJO4iJwQCAcVVmBIBA6HrVzIILH+HpVZyc8D8K2hdCewolAj3Z4pC58nP5U0xK4MfY8mnICFKd6bSAoSDEjErz2NRsnlopJ3Z61embIwwHtVTKouWHOO9dEG2rCKrhgvXgU5G8tTjqKe4Dxq5PPekAYxFmXOa32WpL8iSD515OCO9SRMzDAHJ6tTAjLFuwDkc81IuY4hgZP96spdRonjALkMc+9Dxn+DnNMiOFIP3z1qeHPlkHrWElZ3K3ZRChV+UdajZxI6/LgCrPlYZx27VWjDKSMfnW0LPUROVAJHaoyqoVAOSanVQwDY5qER5OSMAVKae5RZA2fl0oJy2AOBScblBb2xT3UsMDisnuFxHLE4xwajbIYAnj6VKTtbnnFMOAu4VUdhkbrzkVTYh2IBx7U+eZklwnFVw+2Usv3W9q3hBkt6luMFh9KdKiuM46fpUUM37wnHBqbLlTnv2qWmnqGhUdfMcAngdKtRRLEMA9etLHGAPm61IsakZ7dqUpXVgsPdsjINIgDrwaULuJHbtTVGOorIoXDbyOoqObd5RKnGOtSs4XnPJ6CoJUbPXrVREyK1m+YDfkHtitDHyYzVSGII2T19KtIdy7c8iirZvQFsRvnbjuahfO4jH0NTkjJ45HWo2BLZFEbgNTrtPNSH5c1CWORt+93qVuRmm1cLjWJDAfnVaZvMY54C9KnyQ2c9egqCWNmIJ4X0q4qzAjt2BY7fzq0pbPvSRxFD8oyKc6nqopyabDoSDJAYnmld9wUdc1GWxtBGfSlT7uR26VkwJSeADQzB/l7im5BRN1IoCFmJwDSt1DcY+QxXqKRcGQ+mOlKHBPB/OmIMZYdT1q0BPnKhSO/WrCL2bkCool8yJT6VbOBGMisp6aAD44xQRwMflTXbCrgU85I5OD6CsgEwchc8io5UDqc9KlboG/SklG5AfUZoj3Az3jZiFB4p8oZcKh+vvTn3Ky4qtcGRJsqP8A61dEbuwi5na20dahhdw+CODTowFZWY5z1NSbGFwSF4xxzU6XsMcDhW9RURO1JCejVYJUZB4PaqAEhDhjgD1ogrgWLcHzNuPlpDaqbrf69qS0YkOT0P3atEgMDRK8WG4kgBOe470xkOSAPxpzlmx6ZpwO4YzjFZ7A9So0m2QrtJX1qYL94gU1kLSYzxUgAVWX2qnsIhDZ2jjNKrLvK4x71XkypAxUykDp3q3HS40MuSQMjqKkhIYE5+YdqhlGYmHftTrN8Agjj1p/YBE8b+YmT970qRflxgdaiRTGCTzntUik/gKydguPboeeaaFwOetObCtTMkscdqlADY2gmoGGBweDUz4OPQVGMMMsOgwKuOwBtzijcCWHbFMjLEnPFKRmTOPl71QE0f3do6etA6YPemRZJp7ZLe1Q9wBTTO54pWyAAKYzE8jtTQCsNy7s0ucnFLkAc496RiAh9T+lCYEDq27OKfHyRnmo5mITleT+tSRkHArSzsG5IvzZzxSHIU7qVMEjHT0pZACo9u1ZgR4woJ6mq8ihQ/r2qw24Kvt2qsysxdscVpEWhLAu6MHp609SfM9qZAGEPXOac5cYAOPXih6sY489qqyFfMK7cAVZ3fKeMYqkiFmd39eBVQW4FvbtwT07CmxuWV9w79Kdncny88UyNvmYHk96ndaiDaN5wKjBAJJq3gFTVEkh2H5U46lEm75wP1psmMdeRQn3M9+1LKu4Z6GqWjAbHhsjNSZwcDvUUci9DxUnQggcUO6YMhuMPwOoqusZMRINWpCMkdh3qBRkkE9K0jdIQsZGAVxz1q9HwobH1FZq7h8w/Gr8ZLRrU1F1Gth8nIz+VQNnzdvftU7jCgVCYwzr6iojYRKASx3daZEGeM7hg9qkTPmkbutSH5QFxU36ARf8s8Yw1TIOMZqFTtkz61MBhTz1pSAUZJ9qYwLSAinoCOD3pOnIPTpSQA5waZIBuBNPOXPFMl5YEihAPK5ABpoGDgdqCWLYA4NID+8z2FCEPxnOKbyVxml3Zbg0gA34PagY1+9V2zIi4FWnA2/WoFG1SpHNXFhoNA6D1qUfeGPSokGAuTUob5Se9NgOBznPalzkEelRq2Bx3p4HyE1LDcRhxgdaqyt8oC9TVpyQTge1VpBjCVcBDIpCYyV7dqthsqDiqkcZ8xh29KtxAqu3FOpboCFX2pwx60wNwKXkAGsxkozgmhun1o4wBnpQ3IHNRYZXXI3jtSRnB3DvSkfORSKMcVoIQrukY44pEHBU1LkbOnNRgEN15NNMCNlLKwHWnEYUA88UBQGNDnagxzVXEEZLDHTFIDg7j2psbA8Dv2oZwJGRugp2dxgMMQ2PenbskY+7TUz5RA60IN0e096BDsqrAgdakQEAmoiQXVSOanY4GKlh1Gdc+p61EM4J9KlOFGe9MPyqfehB0AD93TcjzNtLkk7SMik4bPtxTHuPIypPaousbL3zUw4TGKaFCuQfShMViOMYjFQynI+tSRNuDAjvUcg/d8Va3DoNHAPHToKkB3EY6iomG1AadEcPkcg1b2CxIcgbl4FJnzISc4xSMxVtmODTtoQMp+71FSAjL+5B702JN5LYwe9K+5o/l5ANSL/qj6+tF3YCOOPMZDN1p0YRB5XY0jEiRQo4NKV2lpKGFiaJtrGNRSuxO5DSK4VRIw5PWllHzZqOoDUGJMelTHPaoowVYNjkip8jbnuKmTDoMICnpTXJ4p+ckE00ckZFINALBW96ToD6UrDLYob0pjA8YPrSAAOcdaUkZC0mD5hOeDQK45Rgn3704jikXAbFPYEDHepYyu4II5/Gkxg7hUjEHgVGGPSrQCknbgY570cgDmnbMKAKRQQvIoC4KNvbg9KJRxmgsDxjGKWXheTxS6gQKwwT2oYh+arTS7F5NBkIAA71tydRD5OXwT19qhciMEsPxqwzZUH2qsXDAjHSqiLYjYfPuH3TTiAFbk470Dbs+tKefoa0uC0M5iN5bA5qZADGeOlMlQeY23qKlX7h9utavYmyIxkSHn8aRjnGM89KfjqRjFNKlcHPFNBr1GleQcD3pinL+tPbgnjIpgIBzg5JqkCEYA9M5FAB2Zz+NKqdSDj2pVHBLDBxRcOgltlZc5IrVgyB71mW7YmODmtKFjgg9M9aiqm02CfvFtTnjqajuf8Aj0mz/wA82/lT17c025P+iT/9c2/lXJH4kUztQSWwOKOvWlI4zSE9zXiljSPm9qRB82RTlOWpBwTjpRqIAcsc0hHze1PHrikIyeKAFONopACPmp3FIRn6UhijjnPFNzhs044IGBzSMeFwKdgJFO4dKAN2cdqPunApckHipAMlgOOlIhHPFKQRSMORgU7iAgLg5oA+U+tOwMAdxSH5eQaOoCE8Z70oy4PsKQ8MB0pQDk9hQAg9yaTOBTsbqTAK/SluAo4GCOvc0YIAyc+lBwwGOvvStx70AHtij/IoJyfc0uc9fSkxjQePpQemD0NBHpmlHJx27U9BCdOAKUYJJ7UAk/KaTG0+1IBCN2cVXVsAg9qs4I57VAcLIQe9aU97ANzhvrUqknPpUS8E+napE+6fWqnYB6sOQaQnB680gPBpN3GcdazsAuf4fypCMNSDIOfSlJzginYBFBHfmpFbLYIqMZzg+tPXrkGgB6qM4zj0oHcCkJO7PpTsbW4qQIJ2KIxC5Paktyxg+Y4apWHUd8U1cbc9/SrTtGwCnrS5G6jgnjrR71AAxGetITgEUE80isHXOKYDiAQFpAR36U7jaDQcbaLq4yF8Dd6U5FzFjtiklGDg9xToiVUCq6CFRvkxingZB9qaCCucY9qdkZOB2qXqA3OVpxBzzSEdOlHJPXikAMdo9jSEAcdqcynfgik68nt1oAGA3D3pDz2pwbkZ5ApOjAjtTGQzo0ilVOD605cqAO9SEA5I6im5zzT5nawgXgmgAcnrRj5T60v8IpMAT5uv5ULznPUU3+EAGnKcZFAAME8njtSfdBz+VKQOnal4cD9aNQGdFx3qhOS0jAdq0ODiqrDOT3q6bSeoFeJHMYB4JNWScqBnkU2Mg8+lPVAGznk1pKWox6ZZKcCOlIMD8KR1zg5rN6gHABOeKYY9wJJ+lOc5xgcUjHHHpTWghyIB3py8P1qCSbylUep5NTAjaeOtEr7jHA/Mc/hSZ+fA9aCcBcd6Mjd15qUArjEmaryjb+dS7syYAqO4UswYnpVR31EJETkk9KlGGYFRn1qJXyOB8tSIw496clqA5SS3PanDGc5601MljkflTlxuPHFQ0ADjg96TGFx3NGdz0hGWAFIY4EZHqaOxNAAyAaNvY96AEA6gdaVSAMU0ffPpRkKcUAOJximsvzCnEHcM9aG+8F9KA3EXhqbzuJFPB+YgU3OGAFC3C448etIwwc9z2px+ZifSm/eIo2YAwzVXedwParZ69eDVdwASuelVHewmKTk0qthh796UrtQU1Rt6UxkhHIA6UAgg8UceWT0oBDEUgDbk4pNp3Y96UA7s9KaW6k0IBw5IHYUpXDc9aCCVFGDmh2AQkEDFI5xjHeiPhuaU/eyRmkBH1fPp2pCTgnHNJgLMcck0df8ACrsIhlwkZzUQJlZG7d81I/J3VBlgpLDaB0NbwdkBFOQr47+tPDBVBHU1HOyu4BGB60sQaRge1a2bjdi6itIq4yOp6VY2CSNQBiqUhUyA54z+dXhkqpU0prS4LUk6qoJ4FNdlViD1pJAWYFeg60hI3HJyQM1kkMfgOVAHHemlVEm7vRG25Nw6mjbnLE8UvINAl3BMg8iqs67yQeFPWrceVLFzle1V7o+YRsGc1cL3B7FJABIR0FNhJEwXHAPOKGlAkzjHPGaFHkMz4yO1daSsSadq3zlW456VdAHzY6dqo2pYv5rHr2x0q7G2QfTtXFW3KDJJOBTSflAxTyOR+tI4yhxwM1kgIWwrgY4qGUNvJQdOtTuMKf0qDed7buhFaxTExqjaQ3Vj+lTAANx171WZx5yqozjrVrIWTAHFVJbMEVZsO2SPqKqyFASRzxxmrkwLS4Tj1NU8qH29VzyK1p32QmVQuFBb9asK+UBx8o71HPseUBeRTow7RhK6ZLmjdk7MlbeIlwOG5NIw3FQvPqKnVQIcntTIW3ZKDr3rJy0KGSTGNQdvI4xnrU8Lb8spxnrUYDLlQM46+1OtYydx6j1qfdabBMHws0gPQVVkmCuQ3TtV0hSzZ61lz/f2jqD09adJX0G3Yf53zKRwF5NXGAlQkDtVEKAwY9DxV9VAA2PwRzVVErXQlcgQl2yBypq8Rkbsdaase1NuMc9akbgYHIrCck9ihpUFj6VA+VYjHBqyUIjGOaiMTFt3cdqmLSGVXgV8fzqHy0A4GPerZU7WFQmI4ZieK2jIkjijIl6cCrTDkE8iqbSbSoU9etSiYswQDNXKLeoXsL1fkYJ7VYjIzgdKhfHmYzUiKwiz3HQVm1dBsKCwk9u1KcjGe1P2ADeex6UydZHxsH1HrU31HsEkILZPUdKGUmJexpzB2iBPDDrShcocdqTdtwWxCwYMhH5VIRgYHU0IN5X0qTI3k+n6UNgMK4U4696gfK8nqelXCvB/OoZlBA9qSlqNlVAxGT941Ow+TJ603aQwbt6VIx3HIq2xLYrsflGKR2BUY5PpTZGZGO3n2pwJKjI5q7aXAVQNwYZyRyKe4G3rgGoQ7ZHAqZseXg/WkwEI+VQadHnBGMCoz8ye/apxxgAdetKQARlSeoHSmSKWh6VJjKkDrS7AVC/pUJ2DcoKCp5NT7gIxnvUhgDPyOlPdQq424xV8yYdR0CnaozmrWMJ8w5NRxYRF55xT5HwhA61jN3lYBQCQufxoALZ547molLFFO7j0qUA+XwevWoasAp6nvimMcopz9aeD8pHamt/qwSOtJARDDSEHqKiZC/Uc+lTE7QWAppOFL4rVMRAhJGCMEHpVhjh1H61Ep+YtjjtUrIGU5P0py1Y+gMozx+dU745BA4IxV1CNxjqldsBcnnAop35gexKCfLi2D61MSwcdqrKzh02j5T0q4ozISac3bUBrZVdx6UxG3JjoamAyhB6U0ABj+tZ3DQZj5sZ6U88AdaAqM5yOnSl3fKc9BQ2BTlT94ewpsbtkblGQetPuBjc6nr2qqjHv1HrW8VdagWol3huxohXa3l96liACggde9RyAhg2fmFTfVoCfbtJHf0pxOVUAc0gbLD/aozhyD0rIB7YZOKjZQGGKeMlaa2dmaSGImCGyOlMB3IQD0qUjioSMHFUhAQTgelB+YEdaXopPemEBGJ3darcZKhwuKd05qNMZznipDndgVL3EMB5yagdyX/nT3BLbs8CmMVHzqOauNgJByeRSvIARgcUxHDjOOaXGEyR06UW1BkU5y64PFLEyO+0ClfCxE0xCR8y+tWtVYC0gAYjvSfdPPOaVcdR1NNbjAx3rLqAxixOMUmQi89DTmJJ7+9IccA+lUgFBO7jt2okB2k/xU2I5dsdKcw+XPWm1ZhchjDE5c5PeoJHIbBH3qsl8grnr3qvs23J7jt7VrHVgTxkCEDPNKu0T4796YYgsu4c0/wCUS5J+Y1D3YEhb5Sc9KqcurMeuauAAgqeh61R2t5hCnAHWlDqA4vtGAOKJZCIfc9DSgAA+tE6jycetXpdFXsVUO446VaGQnv2qqi7QGz1q0q5UDmrqCYxgO/WoovLZzg81ZlGeR2qjbtmQkjGe1EdVcT3JmJU89c1YjBwCKrMdwD+nWpoZMKB61MloMsSFg5OKiXmTJ6CppG5FREZHHSs47CJUUCQmpM7m3U1DkYzSkcYqWFyIrnp1HepASyA55qPOOMcGpF4UDrTYdBx46GkfIX5aVgA2KOgI7ipAROF470xwcGpUwF56mmMCQ3FNPUBE+8cGkI4b1ojGFPPNGCQeafUBUOF460YIbNAARVpQP3nPakAr8vUDYGX/AAqYHjJpn8Iz0JpxEQuBkUpYcA8UMv7zJ6dqQruOKsNQT5SRn6VNGTtOeM1EFG7ipYzlj6UpDBx8ozxnrUEsZMm49B2qdhufOeBTHwVJJ4oi7AV4m+bnrmrAwqE+tQx7VfrwanQh0K1UwQ1Rs4zwakX5hnsKYpXdjuOlSEfIKhgOxk8cinNwT6ikX7uO4oJ6kVAEBGCWqNCVyx5B7VKQWyc9aYwwAa1QDyAQDTW+9jtmlZsdOlNY8j0NJAMkHXFQ78ADPFT872yOlR+UCWJq4vuJkcfyyEg8U9lWQlgPrSIuBT4mwhJFU31QK41F/drjjaOlBYkZ7UqZdGB4phJ4U9fSjdhqS7lJyKlzgmoAoZgemKlBBB9ahoY3GV+lJNyFwKevJamk5UhvyprcBqnkYH1pAoWMnPU0kakbsnr2pdvG09BTe4EhPyjFNIz82OTSngj0oPDAdqlBYgY4lwB9aSRSX46VIv3ue9JIMqcDmtEwaIHA24PHrSLhHC9qUnIORz60wEZIPX0q1sK1ycc7t38PNEcgaFiwwD2qNHY25xye3vSod7lcfKRSsDehJu8tCccU1tywll6ZoGWJjYYzQpIiK46daQWY4OHwo43DilGVjKsODUQbaoYD5RUgcyJyOfSk0Aq4Vtp+6elSScgAdaZEBtJI78U7OTkjFJ7gh8RIAU1KRwarj261Pkbvc1EkMacZGKAcv14oI59KacjOOlJCFYbmzSsMsTnpSBuRmlbrkdaB3I85PXpT2PJAzTR97IH1pcbst6UxbD4/ennJPpUa9alPXJ6VL3GQ8BiR3qPHzewqZlw24ZqJh37VSYiVeQKafQcCnKML7dqOD9RSGQsCWBxSTtxT2IB55pk5GNvrVrdCW5m71ZirdqcvzIDio3RUmzjJqWFv3Y4GDXS9roXkxxcZ6/hUbuFII6USgbsGkwGhwBkCkkg1YoAyw9ehpN2G6fWlDcKMfNQzdSKYbFO4wJc1ZjUSJxkZqvcgLgjr2NT25+Xnr2rR/CmhW11IGBVsYGPanOGYbQKsMMnPaoZAqk7eBQpJgo6kHscgU0Ekg8e/rUpwxyahPB4P/wBarQkw6hgDjFIFAI5P1pygc46E9aapYAjrnvVDashqtibPbNakBOS2eD0rKYZYZGOe1alsPlz1wOKJ/DYWzLqnP1ptz/x6T88+W38qVPu8j8KZcH/Q5v8AcP8AKuGK95FrXU7gZHrSZxxSLnnNKeQfX0rxRiHsBSYApxwBnvTTnA6UDFHTNKMBQB1pD2FA4BxQIUY2kd6cDximjkGgjgYoGKCByaXrwD1qONvMQ5GMU5/lQEU3poxEn3Wz1pexNNU71DA9aep45FSwEByMnqKaDnqKXBB+tC85GKAF6N6ikYelKpyT6UAcZ71IDGzjOeaepxjPNNPKk+9KoyvPSqbGOxgdOtIBg4NL9aQnJ3DipEKFJJ44o5xyeKOh9qBkZIHHQZoAVgMgqOB1obGM45oyQMjvR/CM9ab1DQQYAoAJXPbNKcA49aBwM9jUgNP3hQCcc0DhqFGSRVJXACMjioXGWI71PnBqN/8AWDtTiBHt/hFPUDaeKbxuOOualxgYqpbARkZzxTUVgOalI2mmHPI9qSYDcZQinLwKQj5T605fudMmmg2AjD8HinLjPTk1G3AIxT0A4HehrQNR+MKfWlHP1pp604ZK/SoGISue9QKhjdiTkGpTjuKaqHexBzkVcWIVVAxilHI6UoIU0nT6GoAQktz6UgPyZHU05VIzjpTTnA46GqWoD1PG3FKclMelJ6EdaB1b9KkBjqXw1OUFVH160DlfelbsOg7U7gC8sR3pEI380hyGI70MDuGKFuA7ufanD7maaD8xz0pV6MPyosAAnGTSLj5j1FKDtXB70nIwBSuALyrGkUbgQfwoK7gcdKTPTFMBR3A4oXqQ3ShuwBoOMcfjQgF6ZyKQfdGfwpeqUhyVC45oAAMAH0oGMn3pTnkdMCkz0A4NF+wApC5FKMKee9HBIz1prttJosMFHWodoBPpQjuS2RwOlOPJxVWsFiFlBUkflT+C4PYUbfvEU2NSF+aruIkJ+bJFIxIYDtmjAJB701mDe1F9Rj3JDYHfpSOAHznBNICGOPTvTnUMx9qHZMRBInmgIx+tT4wCMcdqYoUNnFSMu4iiUk35B0FJwAaawy+RUnRcflTWJBAHPrU3s9AYgG0knvTZfuA0srfOAOc0kpzHgU9eoEChjtxxzVlxgAjt1qPBMi47DmpDyDzzTk7gIhySeaex6BeDTFHzgAVJ0bB9Kh7jEIAUHvULOBIpqYY21ATkZ25prcCUEEDFOYYfimJkgN+lOJyMnqaT0YCBhvIFNj+YnNNGFmLetPRvmOOtN7AhzZ3delKx5yevrTXOFz6UM/zA4pbqwDuhppPzE0pOTzTMgEZ70JAPUjaST16U0sVwR3puSWwDkCnAcYNPYB5wV461WkUB+epqwxx2pr7CwJ5yKI7hYZx5XPSqskuYzg/Lmrm0NGRWPOTFKU6e1a043eomayEBF560oJVx2zVe2O5FzmpC2T16VMo6sZOOAc9aaDmPOOtIfmjyDzSg4QZqGrAO68jihW3bh/OmoxYnPpQgGTjrQtHcQvABweaWRhszg1GPv7SOacxGDimHQhdCzKwPSlVsqWxyOtOzx7mmRnCtV3urAMxknPT1qCU71K5w1SykE59KpbnMhyCFrWnG4MJMZUMOPap7dVUNg9elMLgRLu79DTD+7wc5z2rR3asAjKCxUDOKmtpP3ZbsPaoGm3bgOp71JbSfJtpyT5dQLLKFQsvPHSoomBQSEYp4ciFj3oVd8IKdPSs9rpi6ixffY5yKkbke1V0DecRngfrVh8+XjPOOtRKLTGNG0rn0qvO52tg89qljztIbr3NRzqVjPP1prRg9jOkCggMcsfbpSTN8yAU+Ur5gbHIHFNB3orkdK7LbMk0LRi0QI6mrqAY681nWLj1wK0EHft2rmr7lLYlBwB70xjwwpxIC8daRj8oPtXMguQkkxjPX0qvIASMDrVl8DHPOKqyLjq1bR8xEUSlSQvNXeDEM8+tUUPzqE9Pmq0WEaZbpWlSOqBbCMBg7hVF8ZZ8YC8VaeRuoHFV5FRQF+8WPNXTVmTIgIRiCBgnr7Uqy7vkjOR60PtEoCnnvTc7ZgxHPYV0JoRoY+XYw4HamAoqZxg+gpA7uQzDjFSPGWK4XjrXM9NyiIDbGRgjP61YhVVi+XqetRBt9wF/OpgctlR04pSukPQqTARMcdR61SdmyZB0J6elXNR3KDiqCuQQrcit6KbjcXUltznknr2q/AimLkfSqduyOxQD5q0F+U88YFTX7DRHDMGXLA8VZT51GOhqu0SrJkH71Xol/d5xj2rCpJboaItp3EjoP1qNzglh+NTjgt6moXTgndxUJ9xkGMn61DIMK/HFTuNig/lUUgz361tF9iTPYHcH7etSROeMjmnyKoxu6CpABuDA8Yro5rrYLDFVjJ8xyatw5Ksv61VVZBMCO9WVyswKjr1rKfw7gSugMJDUkTMQm6rCruOO2OKgkXHso6VjF9GMJH8wtsOc9qcqkRYHBqNRiHfnkVKr5UA8Fu1Np9AREDhRgHrU0Y3KW9eopGQogHXHenbgCAB9ahu+wDiBn3xUDjKfjU/Vhx0pJAO460LQCDHBHakZQACOlSgBFwelM+8Djt2p31GVpkbd8g/GmJH8m0nnvUwzvweCT0qIq5lLYx7VtF6WFoQ/Z/MbeDkDtU0yHaAp5FSAgEKvA7mpWVcAmk5MCoAQgJ5NWgGMgIHaoRncRnIJ9KskEAUTYIaqlXZienanQ5c5IwaeV4HWlI8sbv0rHmuBCQRMOP1p8qljgmjAJPrSoBvOabloA+NAr47etLtUnnoKYmfMwW5xUrYLYqZ7gNC568Uozg56UpBVQc/nSDIUjNK4C8CP3qPOUI9KGc4AxTI5PvE44pxiwBvmB9R1pH+ZBt70iuXlbimRMXhYEYbtV2sguL32DpTwFLjB5FRoQAUIyaljXCg9Tmh6DY8j96CetU7oCQlOh9avOMPk1nXgdbjcp47iilrITZGkmyAHdkr7VeilWVN6/lVBm2AKQNp61fgREjDIeDWlW1rgiYEbWOODSBQfxpcgAgDpUe/K4BrBAP2qGHFNkGVI3Y9KVsECmSDdhe1NAUZXbIPYdqFjVRvbketSeWRMQeV7CkSMBZEY/LW91awdSxG6vGMVFcKFQH0p8WEiBI60yTB+8MjNTpzaBrYnX5Y155A60DnLAcGhsOeOlKg+UjHFZjJVGWGBgUwjOR+VKnQD9KG4ap6iYh5bJPAqN1JcYHFSEDP1pjc4wcU0Aj4ABzgYpi4YZIyBUkgBXHUim7CF/pVJpgJAW5PXnpUvIyxPWolIU4H41LjBApS3AjlQkD0qF1wu3sasMc5FQPuLg9RVRYDVPyAU44Che3YU0sSflHA6ipEw/JGD6VTuAOgMeDyfSqu/ao4/CrDlmBIqsWG8qw6HrVQ13DcuROSmSMHFPYsFz3qK2kDxlsfQVKSQnuelZyVmHQYoOSxOAetMB+Y5pS+2MDrk0ZGzgUwIkby8j1OKmRdqk5z6VHHzx39alUbkOaqTERTkDBHXtUTrl19assu5BgcioJfkTj5jVQeg2IJxzu+9T8pkFvXpUXkZQMx+bvT8DJ9BTdgJXz5bEVXRuSB1qyD8hA4qrGf3pwOaiOzAXb8hbNO3Ax8jk0qyKysuOaQBdhJ7VTYyH5RUqk9uhqJz8px37ULxjnp0qnqHUlLbVwOtUE3RSk4yKvY4P6VQcnzcE1VPqJlpPusR3p0IxzimwnKFB+Bp8OfMZScjtRLQCyE3LknrTsBV6U1DtBp5H69653uANgj6UckdaRtpjx3pAvy4PWgQjYJGetOixg5PApj8jj/8AXT1UKuRR0GKfvDnNK3JNJ3z6UEgk89e1IAJ7Dpig8R89TQv97FAO/igBqcDHWngDP1pkagyHNP8AvNxQwEAyeaTHU5o42UmcDBpgL/AKiYgYp5OQopkg3AEdqaAax3MDilGAxPYUNjaB60YzgDmqENGd5btUkZJXrUYbkqOtSxYBoYwfhSB19ajI+XbnNSEDJ9qjXkmkgIuNwI7U5SRLweDTRtUYP5UqqsS+orV2sJihSsxUdO9TqAxx6UxuW3A4z0qVeKzkx2HA/NStx0pq45pT6+lZ9RjOinHeo3XIqU9AKZKu7j161S3F0Gnr60h9e1PGPL7cVGwyvWqQAWwxUjr0qIg429cVI2WYYHSmsB5nFUgGPgqMnpSkM0ftSNz8ueKcxMcYA7dTVCGyZjX3NRBCR5ncnpTpgW2nsOtOjKn5R0NUtIgyTO3inoM9TyaYzBSPUVJySMVmxkeDg+tIw+THennO4+lJjq3agLDQB8tICWJx1oUDeeeBSfxbx+VNBsSnkDBpj5Lr7U7OBux+FIM4JxyaSAiBxMeODT2IGPemsCHUDpQcbiOw71QELfKT2FR7fm3djT5QDk5xURfmtVsIcHIVcfdpSQqPIDwelRbyoKkcHpimfdRUzwx4NVyibLJmKxK2OelPBZYyRyGqDcylEI49KmXLQkdSOntUtDGIu5SG4zUxGCpTsKacZRT1FEZPmM5HA60mJolRtr7Cc5qWReQBVfaWnyegqywPX1rOWhRXc7WHPep0YEk1Xk+8RViMAxZ9KJbCHeuaaeRilIBQg9aReVOe1QPqKFATntSLx1NEZyeaGznANPyAAMEkdaaCcEZp4GAKaSQR60APBx9ak6rmocZJFSdfpUtBsIx+Tj8ajJwlSsMVD/s5zTQh6ncgNKelC9MHvSk4HtQxleTmUZ6UyYd81Kx+eopVyoYj8K0i9gKcqgvuHJFSIcEDHBFNVNrHnOaWE5x7Vs9iSOVAWOTyelOAGMdu9E8eTmnDlRgc076BYh2GI4x+tOzxnGKe529smmg8HjrRe4LTYqXJ2nkUkUgXHfHWluvmXlunao7YkcfhkVsvhE2XCdw5FQsQvzbTmniRl+UL+NEu5gT+lQtGPoVw3GehqHjf6j0qdgQx5x/WoeCc9Cf1rWJOoYHK+1MRsHODinuSOOw9aTJA46DtVDtoNbG/IzxzmtG1bhTWa2QcgcVo2hzj9DSn8JL6F5CO4+lR3AxbTf8AXNv5U+Pr1pLoEWs3psb+VckdJl9DtweaDgtxRwM0mRn3rwygJ+bFGc8HvQRzQV7+tAwGRkmk5IP607POD0ozjrTbARTk96cBzx3pF4PtThnOe1SxDUTa554NNm5jI9af3okTAwe9VruHQban9ypFTjGc+1Q24Bi2qelSLwSCOtEtwHdDik6n+dB4bNAGM5qLaAOzgDim5wOKDyc0Y+XGaLAIPu0qLxyaQggUD1PQUwHEEAD0pGGQB+VKCSxPakAw2aT30Ad1GKQghfagctzSEH1pAByFAp3AA70YBHvR3IFAAcc5/CkAyvenZHTrQMhcdqAGjOMgdKaAxy3anr93rTRxmmmAKMknFMbqPWn8gZApr560LyAYOCSaep3AE9BUR6CpE+70q5LQBxG4k0189hyKf2puMsSOoqEBGfmX3pQCo60E4GRSK/XIq0rjFYBuvelT27CmvwCe+KSJsLnvQloImxuYY60ZHHFIjc0MMKagBxwDzmq4ZzKu04Wp84TJPNVmnSOUA9TWkN9ALHU9aF+aQegppIHPanjnJqGmhiEkA+tJ/B9TTgMnmkUcsPShCBhzxSsMKD3PWkHJ5NKRk460aDuB4UADmhs7ee1DcrkHpSFiQBjvQroQqghxk0vR8+lI3XmhmOwACgYAAk46UqfLn0NI3O0DilPWjVABzgE80nXLClbIIB6UMCBx0qRCHrj1oAG7HX0oBG35hzQOue1MA6Hmm8ZxTu5zTc/MM9qEgHdG9qM5l5FBHzZHOKM8ggU9mAhzu57UzkSc1Kw688Go+DzQAjDPAPSlKZGaQ8L1p4GIxg0xjHZQgH50jA5DDpRIoYZamthEReoFUrfMQoxuI9aXbgk01iFwfbmmpOrjPc0cr3GKMpk4yDTNoOWxzSgkAg/hRC+/O7p6VSTYhikrLwODU/XJzxTVQM2aOAgU9aTdx6DevOealU4UH1NRdJCMcCpc4HHpQ7gOYbSATQTuY9qCOQaVsbvrUX0Ai7+4puT5TEdR0p+3Jye/amqw5UdKtWsIijY4LEdanAwuT1qEkCQKBipBy5PYUMaEDBHJ71KSOc9KibpvNK2TFk9u1KyYhzHbyOlRAgKPennHlj0PaoJHKKAORVRXQbsTqfkAFOPEXWmoPkzmnnCrt71Mk0w2IDn8RT4jgjPBNQSErLn+GpV5k3fw9q0toBMfnbHpSHkbsUvGDjvQvEePU1kIa52kH1NNmypyRTm+Zue1KwJX5qegFeMHcfm6jpVgcjOelUmQpcI2eKuDA3Y71c0rXGI+CuQTimg7lA/I0DA+lIx2ksePSpQCx5BIz0rMubZ5LoyAH3q/A27d6VLtG0+9aJuDuKyZUtfmTJ7U8Eh2GOKmWPaOBSOvB45PFDldgxYjvUkd6IzkEH1psY2jBGKX5gdxqdBj1bJwOlNUgtgdRSxDcWz26UKuFJNJqwC4+Y5NCqCMnpTQdx9zTk7n1pMBhPzfSo8kFuODUzDnpmqxGCW9Ogq42Yrg5LpxxUTJ8m3PPrUjktGCvGOoqPbuI9R3rRBqNlRvLCoMkd6ZECjfMM+9Tvuw2TgHvVNz+9A5Y1qrtiJZIgEJBwD7UyDiQpjoOtWJsyQAY+b0qvGd7rhvmHDe9NbO4dSdUbYwxkmpokK5UGo/LIm3A5FSKAZWIPNYyfYaI4wwlYsRt7VLksM9KiZGF3uzhPTFSMTkjsKT1sCKi7zICW47mnyYdCc59Krxu8aHd17VOWwp4yfStmmSilNwCB1700DMIXOamZQyBwc1AiFWbf2HFbx2BF22QY3A1pI2UHqKzbRixI6D0q9CSgG7FctbfUol65WmM53baQMPMz+dI2AC1YJdwGkbvm7dqguRuYAcjHIq2pG3dj8KqyKTMDnitYfFqLoVceSwQLz65qxHHiLaT071WkJa4bP3RVqJt0ZPYV0Tb5LiQ2RgFC9u1VVcq7nH0qzKAwyDz2qvJGEXk/SlTaW4WZXcAT4b72OtO3FeTyD0NNJGTjO6njdJGo24C9a6L2ikL0L8JXysAcntUg3DK+oqGJiIwenHJqZXBTcT9K5KnxNopEHy+YO1SRy5G3nOahfJl3Af/XoPODQo33Aivs7s9OetZrSYbkEZrUvnSGPB5J7VjhWkkJxy1dVBPlF1LNqQjls4Na6OHiHrWTFBhgG71qwRgEYHAqcQhrUaYXaRecCr4/1YHSog285xyKcTxjNcc23oMRgAQKjnAYDFSOPmFIycnNJaDKcoIcE9D2qMsWcgcYqa5zhQO1V8hi3qa3hqiXoQy8oQRyalSMiAbfvVE8hCKD17Cr8OPL3evStpScY6iWpEgIJ38471IhIdGJqLAkY4PNOjOUyD0rF6oZdCkZNQN055qZWGMHmo5CdxGKxW5RBEmGIYkoegqZwPOQDqOlQROxyrcAmrYZXbA6rWsk07iWxDNuaQ9gOopYgd554p8vL56ZpMbCR1zUX0DYlC4A75pGyME07rED39KHGUUms0w9COQbmGaYFPm+1OkOVGPpSYZQDnJ9atWAidVDZJqMqSTg9e9TOoZ8HriiJVOMDgVpFguxXhjIc849ql52kk9aJGAbc3QUsbLKpK9B0p6tXAhZfKUKtWSCyhqY+GOMcjrUiDcAQ31qZPS4DwflHrTLhiowKkyCRgcVUu2bzB6VEVdgTRODLz1Ip+3GT61DAMnzNvapX+Z8U5aMYkI2y9asDAJI61XBHnnHX0qaNick9KiSEO5OFzwOlRdJOBUhzs68mmOQF4HX0pIBrEcg5qMIdrDvU4UbQ3em4wu7H3qpOwFZW2OCwPHeiIfO3Zc1LLHnC5xUKuu4RA8jqMVotVoHUcG3SZ29OKkiU5wTxSRNucjbzipFJEhUUpXWg7EhGTk9RVHUCQvTg1bOd4x361BeRiSM56ipp6STEYqDcSgHPbNbFpGY4FG6smZgkmAc+tatnKJIwC3Irrq3cboFbYs56E9DTGVVYgdD0pYiHDH0p2Ny8iuPYBhB4/zigY2kE9KlUb+o5FQKNsh709w0EfCqTxmq6MrZU81ORvYDFRmLZJ04q422ARxmMBeo6U2VnS3L9/SnIVM3HQ1IFAkC9j0Bqr2BbDo/mQc4zSocuadtwoUnmkON4rO4EiHv3pCO/alB5IFIxBX6VHUBoximNwSR0p55FMkyVAH41SAFG1iSc0SOAuc8UvABH403qmKrzBjFxnP97vU55OKrElXIA6VZUkoDj5jTn3AGxjaOtNblQOuKf0Bz3pCDk46VAFZjtkAA5NLGSu5mp0igtk02M+ZuBFaboNhBnYwPQHNVJMOG7H1q0CDI8ftVJscjoBWsNxFiwcFcflVx8lTnqKo2JwWI61obTg4HJ61nV+MFsVBnYAaex4AFNwVkOe9HRuTT6j6CKcSAg81NJlVXFQphW/kakJDKRjNJ7gI8myEkH61TR2CZqNpChZM0JlgB0A/Wto09AuPhMsrYap5B5anbzinxAKmOhpHA2H9alu7EPjcMintioInUSkr0NSW8bKvz9D2qKMYmb5eM8UWV3YaFUcMOnNOAymO9MLZfrgU9WwW4pMZCVOC2enakVdwbPWpQvDE02NQvTpVX0E9RSSCAO9VXU+YSO3WrbMAC3pVaRycnoKqDdwCFjvx29Klh+SYhsmooxh+TnNSBw8m7ofSnILGgqgj2pWzz3x0oTKgLSMD2NcvUBuP3Z7UL0NOJ5OOlAAAJFAEJ6e/pTgwKADt1oBzJzTdpUcGrAkbgAZpxAGF6mk+8AKcfvZqWA0cOVJpR0NMb7+e3enbsrxSAbGCpY5605fu5zTc7QPU07rHgU2CGkfMAD060SAbwBytCAEE9xQTxQArHEgpo7nFKeWBoI+Y579qaAjIGz6GlXAYChsfdpCNhHtVWAQDa/PQ0IOc4xmkbiQHt0ojGGJPQnp6UPYCZxxgVFnHI6d6kYe1QKdoJbpRFaAVp3IYsO3arMRDRDf1NV5EyxyOPSpo1Jix3BrV25RdSXAA6cU8YUHnOabIvyjNCDA+lYsZIv3dxpQOME0gJ2gCnFeAPzqWCGHhvYUxmIc89Kkc46fjUTp8w5yDVIAiwQRnik4b8KFGwkdqTnnHen1AarZJI/Khc8k/nS42cjvTW5WqAGG5T60zJIA7intyBxSsnfp9adw0Ix8+W7CkUKMAdR0pSoT5eqmkjQBx7VWlhDiVcn1zU27Cqaq/fLdsdKs8bQOuKiQxp+XPqaXqhB61GxyxJPShSTg0W0AAQCW9KXgj6UwfLJsI605FxkGm0IcvzJmgkMuRQM4Ge9KV29OM0g1IiNzjPal25QtSknYSabjC4B5qg2RC6BzikyBHwOae33s+tRHAPPINaINBr5yAo/GouWZY2XmrQi+Yn8qZMgMqsBkj9KqMlewmEhKlQR+NORHTeSc+lOcFnAPSnorbyT0qW9AIpV2OJM9eKcCXdkIwDSyZJU449KaHOXOOBilrYHYe4bamD+NWxho19agRMxjJqePDJ16VnNj6lWVD5gx0qZMBQMdO1Nk6kU+LlfehvQBxOGAppwNwFOIyMUzjp1zUoBUGF680YIyaVhhRijovA/CmA0/MRQTlgR2oXg88+lJ1JPSgBVzu61IBlcelRKSOtSqfrSYxHJzxz7VHI23nH41I3OKYRntnFERXJF4GcdaRz70fewMUkn3Dn8qOoyM7g2O1NYnB/rThyi0jZJ56VQFNl2rwetEa7QSfSpJ1/WmKcg4NbJ3RIsgPl5BxTFO4YP4U+Rhs20yMbFBPX3prYb3Gy5UDNRwsZEyBn61YkxnnkDtVYAJl+x7VUdguRuuxc4GDUUHU8Z9xViUbos9AKqo4jJHODWsdUR1RcPIP60wklMdz3oSRWTHSoWlG7nGKlR1Kv2CZdq8nI74qBjhwT9OlSOCWAHSoiW3YOPStookUct/9emsSCAeP60p6ggc03o36GqQ90LIuenBq5btx9OtUG5bDA9at24OfYevajpZia6GmmM5wKbcOPskw7lD/KljIIxz7U25UfZZfUI38q49Oco7V5MP14o3jOaz7p5EdvfvU1uTKmTXk+ytHmKTuXFbcM0q5Y0xAFUD0pwwOlYtdhj/AGpW+bGBTAD1J4p2No4NINxOQcU8klRximjBXPcUoJPUcUgD0zTmAZRik4JApQeo60CI7dCgI9TU/Ac5/CkGOOaQAg8daG77jHAZOM03GCfSnN8pyeppcEJntS1EIeRjFITk9KXcdpz2poA38dxRYBetIuTkUpHz8UhIAzjmhAPI+XjvSZ+XGBQMj6UEY5pbAIO5NKOBzSD2pScrjHNFh3ADI9s0oHyk00H5qcOKQByTS84waUck4FNGR1oATA2mjOAcUZwDxSdBgHg0wA8jpxTX4XANOIIOOwpp9R0oERsNuBRG2MgHNAO4n2oQYzWoyZeORTegJ70ueCKTPUGswIwOOe9KAVOMZBpTyBxzRn5wMVSYhHHHNIqgHApXXHBpq4BIpp6ASKcdqMZFIThCRyRSI+W296XKxiswXOetVzCGnEjHirDDc39arzwu23B781dN2ESg7wQRjHIqQZBGKAOn0pyDJx3pTncYdQSO1Jxwe9OGSGFIctwOMVmIAfXvSrgNzS/eAxwRSAZOTwRT22ATheMcGlA5AxSZ+b1xSLnf70h9Rem7P4Uv/LP60H0PSkP3sDpQLoL0U8fSkwQvNOkzwD3obJAQjntQgEYcZ9e9IeIwc05uBtI+lJIMHHancBGOVHGKHwF4PalYA4x0NIBk47UWACflHWkP3DxzSjnqKCMbvpQD1BSdhPr2pEHqfpSDnilBO4jFHkA3dk4pRz8p6U3aQ/FOA+YntTasBEYzu9OamJ2iml8kA0FwzEA80atDEwTgGmSJlN3entzICPSo7hiqbR1prcHsDIGAz2oMShVKjBFMYERg+tP3jYPar1ENl+8vPNNGFfk9aHcEhcndSkYbGOQOtOzW4x5GXGKMgjd6UAZUjvSlcAZqQEbrx1NKOynqaGGWU9jTmI7daN0AuTuAPanPkNTSOv8AOpOj81FxELHGCeBTGOFJ7npT3XKge9MkUsvy8kVSAgdwHXPWnSS/OoHeo5YwwO84xToSWZs446Vsop6gTgngHtQPlDMRwabGN4IaldCse31rPTYeonXHcGqpdvNwOgoDFJE5yBVgKCxOOtaWUdSSSHOz5qeTuYY7Ui560/o4rGTuyiKVC64pdpGPQUp6cnIpWGQc0J6WEO25ApAflAA6U7oox6U0dMCle2iAcx5wBSH7v4UL97GKO5PagZWuB8qnHINTIpKc+lB+bp0pycniq5vdsBCqt5nNSSDAyOlPXhjlaZL/AKule9gK8bbWOB9anH3fWqysC/HepxlRkGtJ7AIjfOwzTn6ACkC7ZCe57VJ/EOnNQ9xEbjgfrSMPlBzTpBjp1NMckELjPpT3GLGfn/Cn4y+B0pqgrgdSetOztJUdaTATHzn86EG7jGKV/lXd3poY7R2JpbgKTlWquQFyD/FVhiAlVpxwOfxq4biA/KmB0/nUTAmMEHFSx42kA02Rf4QeBWi3DUa53RhT19aprGUm9SRzV2RfkBGOO1V2JBBHWrg2thdSRssNx7Cq0OA2/mrOHVc9STUKuUk2EZFaRVhvckVirbSc5qRHUOQg49aZChNxu6ikeVI2xnbUy10Qh0jFmUKelTAEkZqtJIABzwKsxr+5BPU1ElaIyEW4BBLZOfSnSIu3kdaUjbyDknoKScgR7zz6ikm2w2RS2E9fuioUJYEEd6t4O7noelVsgsynhV710wuSWYMK+1fvGro4JyazbcqJlIPP0q+AGjJ6+1ZVYbFIewxgp3prMflGM0yNwNoPQ0SyFE3DrnpWVgJxwpx07VXlAUMSevSpeTGpB4zTZUSWPDj6U7NPUTRWkgQxBl4J7jvS2sbeVyeOwpMgRBc4z0qaBuMnpjitW5crEtWRsqlsdhVQgM+5ugNXiDhm7tVF3KKFxmnSfQGRzNib92eKnWH5VIP1FRMFdRt6L3qzBxGHHIraTtEEWUCqu3uKQ4VOR8vamAlXLBs/hUrkNEBjNcreoyup3AkdadHhwc9RTypRAMU0MEQlR1qr6gVNRjZ49w/h7VUiTpj8T6Vau5WC4HQ98VXRtqgDrXVQuo6ie4pkw+3Iq5bF2Ugt92qKgKzFhmtO1ZCCq/e7ilW0VrDRKiny845PWpY+ck1AWYHPb0p0D8t71yONxkzEBs+lRhwVLjvT+CSCeT2pjLtXHpUWWwDHGY/u5NVgrA5H5elWtpKdahbgHnIrSL1sBn3L/vQCOKfBI5XA4ANJcgE56n0qaBQBu7CuptKGxPUsrsB3dyOTUSnZMFI4pWdGQ7fvDvTFLM+T0/lWNh7l2Ej5u+KHyxJA4HWmQqcbs8GlDsGYAdaxt7xRBIpJbavB96lhPlgKfvnrQGckqg5HUmhQwcSGqb0sxKxJKRtDDkimMD5YenuwWM5qKNi8RUjjsalJ2AnRsgU52OQF6VDA38Pep16ZNTLRjIX+706UjjKDmnKOTk8USLlF9qd7CGSrwCvWkRcIOep61IcYGajUksf0NWtgGyqshMeOM0+CMKpUcAU3K7yc5NOHCDBoV7WDYjjBDZbnmnr8qkE45pjPhgOhp46YJ6d6HewEoIGSaTasjgN2pEbIzj5TQN3mA1CWoDio6LxTXHVgelPQ4JNMZ8gg9KFe9guRKdp342lqsAkKMDg1VYMGQE8A/nU7kjj0qpK4EvU+3em4HzEn6VGJSWIHWnkYxx1qLWGOGNpAoAzH6elIuDwOgpx44HSkIryNgBieRwKrsgWTzsZPfmrEqKAoIyKhnBSPCflW0GJiuWDAx9+tSjJIbuarpncGPfrU4O6VR2FDXQq5PkBhiobhiFJFOLDsKinYSfuz19KiK1FcyZkQuWUZNXNNBEZ3rUXl84JqaENnPPoa65O8LCWhobcAYoPy/KOho6MAaVmxnPauIfUT5up4FRsQTjHTpUpO5QR3qNgFYNTQDCcYxUhwMjrmoiflBzT8/KdvWqtoBGx+cgDBojx5vPU9aGzkkHmmphZfcCqWwyVRubr0pepBNRwupBUn5iOlKucjI4qWmImT71NblvQetLjBpH3A89KnqNjeGfA6UjD5ie1KwwwCjpTTliFPSqQhT69zTSDn27U/OSABTe3vQMiwAST171Oo+Wo1ALNnpUqnKk05CFJ+RaBw+T3FGAOvSkJ5J9KkZGcKST3oXaM4pHOQKBkZ5q9bCKzuUnz69qrTqxYnt3qeQKZR6io7k7Si+oNbx6CY21fY4Hc1rZxJWPAD54UnitaM5J9RU12gRXuERWEhODTQA+GPSluCA+OlMDHdtxipS90Y6TluPyp/f5B1pqZ35PWpFGOaT0Ay5EbcX7jqKswKDjdx6e9LcJngcFu9QQsU3ZP0roi7onYuyk7wNvFDgEFegpVJdPcUydCYgFOCetRNWY0whc8LnIqJ2IkOOxpbdgoKnrRKyidv50re8UmMRM5OcU5CwYgdcUi7mHBpVyG4oBoVerKTyajXO4jqDT3XD4zSRnBPtTXkAP1NV5UPlkGrMnKkYxmo3j4XJxVRdhFeMFo8j7w61PAjscntUIk8tiAvBNW7SRvNYMOo6VU77iLanAOT0oIyMUBMKaco+TOOtcjfUYg+7x1FJwU/Ghe+aVvucHmgBpX5fcUi4yTQoypGOaUY4zTACepFO6pilOFyAKjGdpz07UgDOB60sRPJPSmPyQtOHyrtqugDThnz2FKDkH+tIcKMdzS+nPNAAowuKcRwCKbyD1p7YIGKQAoxim/x80pJ7Ck4Iz0NAhgX5z6UAZ/CnEgc+tHUU7hcikJHI6UqdciiUHbx2psPIIqugybHyk+tRMm5MHg5qU9AKjkAG0UohuQyj7w9KkTPl/7RqCRsqWHWpYGLIu771ate6LqSEFl561IikDNJjBIp+QMCsWwsNyc0/OAfem5y+3rTjwMdqQxkuQhIHWolYlBkcmrJAKjNMK/KKaelgGKoxjsKaw5K07HbNMb71NB0EGMhe9Q5xnPGKVc7yc0zOZGX1HWtUgJMhz8tDE5GemKagzH6UjOAQppW1F0FYl22gcChciXGO1PBAJIHNRb2JLqPanuHUUfe296mDYOSOO1V4CGJJ+8KnUZQZpSHdEbjKnAzTo2BjpGX5gR0FMUFTx0p6NC2HuBvDY6dKegJOSPxpvBIz0p4PA9Kl7DGhixAxTnwUIoUgue2KOdpNLqFiMAkEUh5XjHFOUnnJ7UiD5T71QELDL46YqOQBsc8VK/EhNQEZXJNaoVrFnaODnoPzqJk3sGNSIPl59KP4eKlPUOhFhmcN2PWniYFcenSkmby1wO9QSKfL469apK4rWJZpCyBV79aeg3x7B1HWqkMgJx3q8gDHjgnrRJcqsUOTIBA7VPHjy+KhiY7iuKmiyFIIrGQiGQYZsdadCcJuJpZM7Tx9KRDlaOgEnPJHemFSGGKkBHIFRljkUkAp69OKaDknB4pcEqf50gUhcUwAHL470xj8wxTx1JNMwTID2poB561Ip+SolyW+lSLzzSaGBAwcUnGKGPpQB8tIBQSAaYelL2I9aQfc5/OmIFyVGR0pv8AEB61IvHA/GkYDNFwKs3zE5JFVkfDFSOT2q44G7J6VWEIZy4PFbwaSEI3KcClQ5UA9aawIbI6mnhecDoab2AbISOajkj3oMcEVZZMgg9KiZcE80RYFeQExbehPaqmMPg9u9abR7uOmKozLtmzn6e9awlfQVrEsYymT6VCdqk5WpoXAATGSagmTOcnv0NVHcHawhbJA7iowc8Dn6mnNjaCoBPvTNvGT3rRJAJ0zjnPQ0zPHTdUg2heOmaQgDAB4p3ERlstk9at2zEjB4Iqs2GYA9TU8OY2zjtTvYJamjGRkjFLcn/RZcdNjfyqpHJ83BzSzTkwuD1KnP5Vz8nvaD2PQJYUkGSOaEiVF4qQ980fw4r5/mZY0DijHJ/WlxzR2pXGIPukUZ2kDPJpw96a0YMgfOaqKvqIXnoaeOwphIB5p5NSxhxilB+Ukc1HyBjFNBZVwDyeKEkBNnJzTuxpqjaoGKcvNSwHckAtSDOOvFAPbJoC/LmgQvQ00jByKXOSDS8NxilcBp6jmjufSg4pSRjAp310AX1JpM4OMc0EcUo5BzxS3H0EGMdOaFzye9JGTznpmlBwCRQ0wAEbs/zoHB6UYwMmhj6ChiHHK89qMYPNJnd24FJkng0dAF/hPvTGzwf0p+Sq4NNbIX2pq3UBQMrnvTQMAUvVRjj1pAfl6UMCNRsJ9ajDHkmpW4GRTcdcmtFqA5WyNw5xS7t+CBTQuE9qfjuBUNAHAc8U1j8wxUhIIGTUb5yD2oS1AJiMGoFcbwO56VOxyAccVAoywK9quNuodSRchW3CkCYl3DoRUqEMpUHmmngfSlezAewxllqMPsQhjnJ4p4XKGomiVvvdulOLsDJk+6TSjOc5waav3Rt60pPy8dah7gSKSQfWmAZJB4pyLkdeRzTSTgsKXmPqC/KCaXjNNYHAx3pWG4gjp3oEMwQT604DkjHahuWz0NIXxlj1qrgGSfl7U/B5zSHkA0pwFIJ7cUgBicAnoBgUN98dzRj5duc0iZCncKHqApO5hntSEhjj1707GF560gwfmFK4CHAfAPSlAznFC4zmheGbjIoARm3YGOaOvGaXpyKTgcmi4DR8rY9O4pwbndSDAOPWgjBxTAB8xYr0qByxxtP1qZEEbZ3Zz1pHUDBFVpcCEgsxNPSMoSc5BqVhhwTxkUh+Xj1pt6WGKww44pkyhucU9+o5pCB5ZB61KAjQbkwego8tTFyKcgxHzTuiU3IRB5aghu+Kbgs3zde1TkfKah2/PmrUrjH4w+SaV/mx2qKRwJAc1JkPGGHep1EK7DIxyKFXHJ70jr8gwfxpcnYMUugxxGU44571I2NoGOaiYbVA7VLgCMEnmh7CIplZTUAnAwp6npVuQBl3D05qgY8OpPatIWejFqMvGAQn0602xlWSNip//XSzDk7uQahsgsLMq9BzzXRBLkaGty5GwVmJPSnxzLJ34PQ1UnTbA7DvVfLRIig9OlQqal1C9i6YmQ5681aAAU46iocktHkcnrUo2o7KWyaiV0rBsPxkY70J9/nrSryM/nTQTng1hYYp+6frSn5lHahsKABSk5WhsQE5OBTGyCBintjAx6UN1DEUbMAx824UrYGcDrQB8v1oUZU5HIob7jGH7gwOKcBtAI70nLfL2oU44I4pgKDnc2OajcHHrmpAe4701xkil1Ao+S0cint3q7gKCKi4JxT8ktWkndahsKACcntTiPlPrQeHPGAaZu5IFRqwEkBZcjimIcgZoKkEjPJ6Un+rUY5FaRStYB4B3jPalZvmHqaQHK5J5NIzBSHb7tTa4dCSTg4x1qFfmkYHtUpw4BPaoQxDMT0oVrC6j85PT8arz8R5I/Gpz8y49KrTqJYMZ5FaQSuDuCdRjkHvRLzIADg0luSAFx9KVlY4L9c9au9mMSTAYLnINVXBTgtgDvVyQBgMVRkDsd459quD6CZMwZhHgmkj/dvtY5bvT42ZkG3qn61E/wAkpJPLVot7CJ4zmXcB06iq88Pm3PzfWrKuEjJpVRXkL96mLabkGhQlOEZgMhaLW7ecrGM4HepJx5cLBz16CkgCQrgHk1d04hcsNKQpQcuKAx2BW+9TMKCznqMc0rAMisO9Z2SGK6gfMT06VRndR8rD5m61ZcjYBnBHQVUfDEFxnmtYLUTFQ4IdTyK04CSCSOO/vWSMRuAvIzWjE7YYbue3FOsrxBFjIaMfLkjpUTkOwAPOeRRFcf8ALNhg1IyLIy47da5vh3Q0OU7EAxxTLgtuUqflpszbRhemetPGTCB1OOKF3BoqTbm+meakR/nUKflx0qFjsRscsP0qO0YmQMe1bqF0Qty5NIQSO3eq8koUYxwRUzx73Yk9elV3YBeRg+tRHUrUiVyGHbI4FWbdXKlehP6VVEYyXbipbJmbdk961n8N0TG5aUY75IqdSSoI6VWCM778cipxk7cH8K5ZFiyxmUHnj0pHUpGAODT5CduRxVa4LYPcelKN2BBcx5jHPAqlGVV+RkCr4bzARjiq8kCgnJ612UnZWYhf3bn/ABqazARietVTGxTC9atRMcxgjAHpTnK6eo0TSnagz3p8KAqCKZchWKAcKO1TRgowHYVyv4dB9R5G5+OtNkHy8cZp+0iTOeKRhyVbtWa3AhfptHSqcsiopXNXHyCTiqflcEtzmtoabksifDYI/CpITtcjqppo5TaRzS2qkH5vwrbe4IlO7YeAPSm8AlR3qV0YQHHPNVvXOeO1RHUHoW4d3lBVPPenRBiM9xTQRGwOeGqUnAxnrWb0KCN2DkE8HvSO5Bwecmm78soqTIfnsOtJ6boSFlUbD6moYQFIBp9xJ8ikd+1Vo7hfMAIxTgnbQC7Gu0n+9TxkA5pqrly3YinA5PIrJ9xkbHORnkUsXI5NOA6k0wfcOBzT02ARuh54pi42kdj0p8gLKPWo25Ueg71cRDI1ABY0qPu5xjNNZv3TADOKIydhyOnatOVvcBSTuJA69DT8YXGOtNGW60DDbc8YqWtAJVA2gdqVickDp2NMAY8U99uR6VnbUBzHbwPWmyJnpyfWhAdxJ6UMeTn86QEIUlicggDpUkikqCTyajVdqEqe/NTAZPBGBVy7hqRxsCxGPwqQ5ZcVFko4bPGamH3vapl3AXbtjIA5FAJKgHrSKcEhuc0uQCCT+FICM4MhU/hUdwAoBboafwZuOtDoJsqfyq1o0w3RXUh3YA8Zp3zomRyc9aWIBQyKMU5CSVQ/jV37ASggKc9agdNrGUd+oqQEecy5pGyykYqFdMNzPLiO5OR8pp7yOsoQdTURQyTHJ69qeSFmVmGdvBNdOmjZPQ0o2Pl8jnvTiq7SCaarfP8AKPlPendTyOK5XvcsWNcJsHamMu4Y/u0/gPkd6QkDfnqelSm7iIcDDgdBT0ACZHOe9KuAp460n3UIB+lVdsCJ0Ccjk0gBwzeoqVhlcevemMCU2jsapagQwptYyY4FSqfMbI6imupTAUZDdfapAPLUM/BFOTvqA9O5NObld2OlIuMgdqc3C+lZPcBh56CmYKpnvmnZAXaTzSMQCozk1SDYYjFhmkk3BAB+NPjGAwHGKD8wPNVfUCE5UYzz3NTR8L71XONxB6VYAHlgU5bAO5ZQMUmBt96UEgfSkxhRg/WoAh+ZskjAFCMEGOu6pCuN3vUIGCGH3R2rRagQuFVtzfhTLgc7gKWdGeUuvY0rqwiIPJrVbLURXCnzA/r3q8kmJAg796obyqYb8a0YUUxhyOQODRU2uwuNnbGPlBye5oxg9Bz0p04yAaYF37SB0rNbAOY9/SlBLNimudsRxToskA9KHtcBjgGUZ/CqcpIlGBxV5xtJZRkmqk6nZlmwa0pvUHoWE3JCWHWoVldlKntVmEFIl3DNRybQxYDFNvUERRlCTtHIqJ3MkgYL7UyNiJiR0/lT2J8wgD8aq1mPcf8AcJ5pIywbPrStwRkZFNLAP1pIOo45YZ9Kag+fJ6d6VsoeDn1FMDYOOx5pq9gJD1B7e9JIxLBe1PZcxDj60knyw7sc9qnqHkU34nx2FXbdhI4w33etZz7yVY5Oat6eGKsexrWa93UVzUUlkJApGwI+tERGzI701gMdK4uoxFI25PbpTuhz2pinDEHoKcPmWqYDFG1x6GkdwrjJ4HSnOACuOahl5cd8VUdRE7Nl1I6Gm8lwAflNISGQAdKcAVUY4pWsMRwSfekfcwHqac5JIIpoLbgD0FCAJRgilHC80si8bjyabJjINC2EKW3EZpeg4pifezj8aeB1AND0GA+6TSL8yH2pcEng00DarYPBpAJu3Nt7UcgEk0zO0k4p0RLDJPWrsFxrk7D70QkYNOlUYxnjvUEeQwA6U1qgLPBAJH0qNvmU+tS5BxUTZDgY61K3AhdAdvFGdsu0dKdIdsgpu/LFhzWi1EXAOQTyKVhnnFNUZA7UoIbtWLHcZ5gVvenscH2NMKqzH26UoO4kGmBKeQR2prAjHpRyEHFI3LAY4qUIYfvn0pjg8Y6GnBck5NBHAA7VaGRNwwGODTWjUEYNTSAEjFRTkRgn1qosCJXU4Udqey5YED6VRhY+ecjGa0HwMc4xWko8rEmIOByOTSBcBs/hTgULZP4UHhTU3AYq+UCe5qRTlcHvRjKgUA4cik9RjQ2SV9qUMQpJphGGznFCuWVefxp2FewE4IGOD+lSKQMp6VG/J207gxt/eFAwLDe3PSnrytU3DMmR1qxHuCYPem46BfqC8YNLk+WRTSMDrk0v8BFJgQygsoqAZIKmp5SccVBxgfqK1jsIsQsW4I5HFPKhRwOlNiJD57GpBlmJPSoe4KxSkcu5zxtp8eNnJ5p0gUOcHrUBlGcAdO1ar3lZEvQsJCquG7mpeVJwOR3qsku5dverfYA9azmnfUpAhJ+YdTU68D3qGMFenepuc8dDWcgGTZDe1RwgEkVNKA3X8qjj+STHrQvhH1HqeDSFDgEHmlxjIFKR2pXAifKrmkR8ipvQEVC8QByOKaa6gP8A4TSLkcZ5qESFTtanhhjOcmm1YQ5RyQRT1JxxTRwc04H0pMYpHH1pDnkGnDPemknn2/WkIZu+bFOBOMU3oScU7PfFMaYing0vUn9KQZ65oY5GelAdSJlODVaNidwA/Cp5ZNsYPc1AQWfcK1jtqSV45JPP2v0z3q2xUDIqB1LSg9BT8A4ParlZjsS9ec/SmBCz89KcB8oHf3pyYzz+dTewhxUbeKzriMq2T09q0/w/GqFyeCGPAp027gyFSAd3X+lNm6hhyppIyu055I6ZoJ+TacZHeui2pNrEJc4FBBDZI6Dr6UNnjA/+tSHdgkkf41YakbdPl9aTJA6cinZy3HTrSEhumR171QdBEJ35x0PNWwoIyOT6VVVvmHfnNWjwoIpSQmOEIfleo5xUcsTCN+eApqWLKMCenSi5YiJx/snNJXvqEddD0ZiWOMUHGMUpamqea+YNROrU44zkdKZ2pR0oAZKX24WpAuEH601h8opwOCKq+lhgyjilYE/d70o5PtR938OlShCjDcDrQPu5pqn0p2cHHpSYxynCjNH3eaDyPelyAAKeqAUqCeKTnHBobk80AgmkxWEGQcdadkZIFJwDmgHGRihgJ/DQBwPSmsG4A6VKOMCk9BjUJyRilI55pVxk5pCfSjcQ0YGRzS9jR1H40oznFHmMRfm6mkxkkH8KXpTjjaMDmmu4hASvTvQeWB6UmctTmK4GOtJXAQnBprDBApX6+woOC49qBiY4oB+Q0MSppM4UjFNaiIpMkYHekPC47+tPbIBphGcAdO9WnoBJGcJgjrTjwuBSR4f5fSnLy3PT3qXdjEONuBSD5vlp2MMfSlOCRjrQIZgn5fwquoKBj1q0D83PaoJBtDFelVF66gEPQnoe9OzuIBpsQwCQfpSuDncKbtcY7dk4BpHJdCM4pi5yAe9OddgwDmiyuIWMbFwKHdYo2Y9KbEXOQ3TtSvgqAact9w6DomDID609eje1RqdrKvapv4zjvUStcY0Hil/hagnLg/pS55OOpqQGDkZNIyB1yfyp2MqcdRS44ANFwAcJj0pNo2gHtSryCp64oThc4zTewhGbGGWlck496Rlwo44NObJKkcClYBCSflP500ZU4p4wzE+lN+vFACtwABjmg8dDQRgqfSg/fz2ptAIeMfyoYYOKMjfg0x8gn0pJADE78U5shvamEninuMkZ4GKq1hi479qCQF2kc9qCcKMGjOWBPXFIQreh/CmkZXJpc7lx3zRnJIA4pDEx8gz1oIyBS8HAFAyXo1QDG6YpOQozSuCXHr6UPycYqt0A1+WAHSoJcgqB0NTycMR6U1gHC8VSfUCKRVQk46in27b4gR0FE3KdM0sKCOPAqtHFi2FyfLbNEXzQgEYINIvCcnrT0OYio6ip6ACZMmDyDUh5bB6HtUcYA9jUi/Mdx7VLGI+QWUdKgkQ5z/COtWc/Lz+dQMTsIPenG4mVblgI+vzCs/zGR3x1OBV26ixDknn0rPiUvNtk6dQa7qVrCZpOFaEIe9Q2sJe4zIPu9BVqKPcu70OKVUaOZj29ay5lG6Q+pMow5pCmST/F/Ongbm+tL0OfTpXNezGNUbc89aNowfUU4n5lP6Uo++eOtDkxMQjgUv3ie2KQZJPtSjAU/WkApwVHrTTljgfjTlAwST700ZzntTsAvYHsKXOMCkHOQOlLn5fepsMUgK49KjPJ2+tKeOe9BbGDjnvVIBQAV6dKaST16U7oOOtNwSpJFICJvvcCnlsqAO1JjOfakTO3NX5ASE7gBjmmrxnPWjoMjmg8AD1peoERJLdORSsoC4PUU4kA4PWk2sx5/GmLQQYIx2FOdcRgdfamjOcrTyxxRce4EbVC1HIMZFSMQ546jrUZ5f1NCVgGuSoHpULJ5sRFTOxOeOKgkcqwAbFaRuAsJb5TjjpTnB3A9B6URsGUHHBokA2EHOTVO/NYQx8pHhOtUpldoepBz1q8q7FO4/Q1VuB845/CtIPUTRLbgMh2cjHNQlQJMOenSmwv5bYGQCOtPlDCUL1XFXrzXAdbsAWZuh9anGPM3BsA+1QRKjMVB+tTAMAw7DpUzeoyC/TeAnrUKW+1Sd2D6GrF4WZE2c+tMijaa15b5hVRbUNyRI4svt3fL3FShQGIz8v06UxGVZNmfmxzTiAFODj1qZD0GyAlMqOB3qncR+cY36Lk1bmYsgANVLiQxRqqnnsKune+j1BjAQsx3AjHSr1uXznHJ6Vn7yGQHkjrWhCu+beG49K0qbAiSCEtIxIz71OI9vQc1InOT6UPktkdK45Sb0HsRLFmNge9IwEaD3qYY2bSOagZlIKkfd6UK7YbFFo9yMxPQ/nTIgQUVR1/SpJ5FjhIB5qK3kMbKxxk12QvymdrGiY8E5NVbhAq7R9avjn5vX1qpOAASB+Fc0W7lsps25hkYHarKBONnpyapkM0m5qltW3Pkr8xrqcLxuSnqXwcRA5xmnLg55wMdajUNkp0NPixlkHFcskigUMehyKryEs53dqsMRGdqmq7Asm9uOaIb3BkE0vkxZAqh5zOST3rXnUPEq9/SsyeMqBjoOtdVGS6oTJAASCT06VZMpeUDt24qkgOd2fxq9ACxDdRRO240WJMM4TuKsnAwKr537mxUyfMQT0rkkMlZQRmmdM5xzTmb90DUYbcckVmhjSMZJHOarFhuwfwqyMnPTAqGQHKnGK1T1EQ+Xkg/pTYsl8HrU7MoXBPPYU2IqZN36Vd3Zh1FffuOTnPb0qDYQCS3XpVh9oJGee9QoMqKIvQBfMChQRkg/lUxG5gd2KpgMZSO1XVBO3HSnJWQ1YJdypgdqBKqWu/qKdIwLbV70wrlTH2NZLbUXoPK74we9ReQjjd3FPXcAfQDg0yPJU46VS06gWhxgDpTj1x2pqLhTT+vPasWMawO72FMLASfWpHPOewqEspG7t6046iDHzEVCyiTK+lTAknOOlR52tkde1axTQiNtyKFC59aZErhWJ6+lWCpZOep65phY5KqenUetWpMYbiMD86bxjk8DpUYfDnsPSlDZ68UNNCLCP8o9T1p7kfLUcSg7jQys7+wrNpXGPGQQCfoKcOQQeopuQZA1KxwwNSGhFhgvHSpPmROOfUUyX5QoHGaU8LnGardArjUJZMP371LvA+tCxqFzTMKp9T3pOzAlYcgZ5pQo3njioIvmkAOQwqdehqWrAQyIwmBX8aeSM5WnMoGMd6YDtGCe9O90GxF8ysGA6dqFHzsw60/GFJbrUcQZWbuKtAOUkgkDmngExk+lCZYE9KcpLBhipbAz/s5Fw0hOFx09KlUrKpIGTTbpmDgL34NMhnRcJ/F3ro1cbiuXkwuCRgGnjg5IqIHMKk9BUrckHtXM9xjvuyAU1u/r2p2AJgaQkA89KkBrnCe/eo2IaPjqKfMCSCelRKArls8VcdrgPhcGLJpqjG4k8UhBCEGlYgKD0FOwIAdz4PQU2R/wB4y+tSBdyDBx3z61AX/fLuX15qlZsCZRjapPSphwxqpkknH51Zjzt56jqaiSAjYHGR1pshCKCB8xp6ltzbvu54pG+4SR06U+oDATxz9aUnOAvSmMCQM07OEHrTGQy8tjsKsIQwBxUbJ8nP1zT4WyhFN7aCFDnO0CjPYdaCflFG78SakBG5HHXvUZ6lew7VICwye9IFXzCcYpp2AhmIGSO9M35hLYzmifcQcHj1psa4jK5ya2ilyiIVUN97vV+McBe1UYzlg5+7V6FgQT7UqlwQsgEkeMdKZGp2AHinM21Nxpd25ARxWabsNjXUbdpNIhxF7CnYyCcU2MliVPSmtgELDAqtIrSyKMYAqeYFcMB06Coss8eehParj3QiUyBmCDqKY8e5SGP0p6qAgJPTqaRcyncOR3o2egaFNovnCp1FWAuMDGSOtKiATMc8npTo2Dl8DPvVSlcYjMq4GOagc4kwKmYDcDmojgk8/jREZPHGOCBn1pk0AkYMOMVLDhVGT1p8g4PYVHM1LQRU3nOCOKafmTZ6U4ccUdCWA/GtBFeJVLbT0HSpoMozA/dNIyYBZevamq5JBJ4Hf1q37yEX4Tg47VI3K+lRDDMADUr1yvcrYgLAqaWM4znqOlQSbkl2r0708tgrxWltAJBvYsGbBpjDa3B/Gnsx3/7NMfAz60INhYz8n9akQkgAiogD5AA696lXJC/SlIA6yHmlz83TpS/wk0meN2OtSAPk0mAR/SlJHl+5oH3Ae9ADScEYpWO0gZoYZUUmD1PXtTAXocDrQwwOOM9aRc9acB8pz3pCIGJDjA4pRuxwKe4zHimJlflzV30GK4JXNMC7cEU8g4xTOelCESJ79qY+eT3qUAADFMbAB7UluMrzAtESvJpYE+UButOIAQmlUchyefStL6WQmWP4cenel3A8Y5pFAxnikUEEgCsRhg7+O1GCDx1pyDPNKBkk0XAUj5sDoKRh1IpeQMntTTwMdaQDGBIX1pQMMBmnOOBmo1ycntVLUAAO/HaoLlgFOevarJILYFVLsgrwMmqhqwM8NvYk/hThI4PJNRwqckd6ey5SuxpEolWXp61Oj569Kz2UqPepEkZcZ6A96hwvsHU0AxKg/lUgHGO9RxOrRnAxUvTJrCQyN1G5h1qOIYWpyedxPX0qGPPmHjg009AFXGQx6im4Ayex61KoHJFI0eUx2ouOxCg+UY7VMuCMdMUiIFJyfwpUOMkjmhu4rDHOHC96RCSWJ4oYbju70i8tnsKfQBsmeSDkelQA88jA71LKQvHWo1UliK0WwEyHIBHQVIWO446VEh+Uips/KAO1QxoqyxOz/KarTxmNyc8nvWm2OMHms+63B+RxWlOTuS0rajIf9bjPNaEQLKd3aqURXdnP/wBatBR3FFVghY+AfrUo5GDVcuE/rUyk1hJFEh7ZqBhhsg1YXlc1Xfhs9qUQJs5A9aa3DCo423dKlPI96LWDcQt82D1pm47xmlIy2e9NxliQeO9MBWUM3Tj1qF0+cgHpUoYbiKazfPmqVwYm7AAqUZxn86jH3jUgGeKTAUcd+tQyM6yDHIqXgHHX1pCATSTswE5JJpc5GKbz0oOQCRTEKGPGRSsNwpp659KecYoYdCtcDjjnHamJhuQOtPlHcUyIHLE9K0XwgNf5uBTAvyFRz71KRhvU0dMEU7gLHggCnd6Tjd/SngZ74qWwEOQfSqV2Bv6fT3q83I471Ru8DqaunuJlJTkgdT7UDODnrSIeDnoBS7uwHWusjqNyQMA0HITjrTWGTgD8aVepyc0xDRkNgc56UzdkYOAfX1p207iO3XPpSYHP8qZa2Gq2zBA5z+VWwcxdeB3NVTjAOBnPrVhB+7AJ7dKGr6ilpqPkbaQvUio5GJibf6HFOlOGHAyT+dMc5ib/AHTStqHU9LZs4K9qY7FAMD8akYHtTSuTivmb9zRgD1NOGQp4600DBp/H4UhgCCp9qOMc0rfKDgUL93mhgKBhaVRlskcUikkgUdDUiHHBIxSEfNnHNOHIPrTfmHFMB3BoPPQdKACM03HPBpIBc5OaUYUkHrTcAAUvO8+lN2AXkHPanDIOaYecYNL2bPWkkAHrmk9cUuTg0AYWgB3Xk0Kcg+1JtyOTSLk8UrAOGO44NIMZJ604YI5pF4A4p9AF3bs5HSmPllYKOemaco+Y+hoA2Ej1ouA2NNqAHk04jJJA4peeRmmqeSM0fEwFzxTcgtxSgjkGmkEgii1hibug/KnNyaYPenE5I9KAA/cJ71GOB7mpJOKjfgU0IIgdxHepgMtz2qJTg57mpeibqbB6DRkA04jjcOMUnTBzyaUr8uQanyAj5Kk4pki7lx3pdxQ9ePSkcDO71FXbqA1FwoUmpFHBBpvofapCuctn8KcgG9wCKSThuKccsPeiRQB61OwMMADIPNNbAxQxG6mkbx8pOKduoEmcYPU1IzZPHBx0qHr0604nJBBpWGCnI56+tPHQ0xsA5HWnHIxz1FAbC5zk9zSsSVB7UDqQeppcHaR0xUiEB+fmkU4DAikbhNw60pbKBsc0xgTlMH8DStkYHam7gPlpGY4C460XAeCI8DGBQVwBk8UhySM9PelYfN14ouIRiPmINH8Io9eODSk4BUj6UrAMwDjnGOtDJkgZyKAucikJ4PrT6gDbsDpRIVKjnp29KbE5YDNRI264OV46VcVd2AmOCgamGQflUsgC4A5HtTAi5OelKyW4CByCCKej7i3uKgbHFKoYS57GmrDLC5AJpUGSaQ5B+lHK4I6ms1a4EYJ3k0cEimtnfgU0MUVi3artcBWPJ7k01cBNvrTgc7T196ZgA57VT00AlxgKPQUBgFzUcshSHdVdAVAIOST0oUebUCc8YWnAHjnmoxww3DmpUwfmoasgHjDNtz0p6njbjBNRk/vAR+NSIdpyeal6CFODkfnUTHJAxxUpG0Z9ab0BB6ilsBXnjC/Mec1nSxFJCwXj0rXmGY/XIrLuZD5gTbwOpropXvZBsXLPIjOep7VN7GmwgGDK9QKdn5ApHJqarTlcYKMA4p3fPTNMxiQkfdNP6nIrJiDOPmoc85zQTu6dKaDmkA/IXIJ60Yx16UmNzg0vU/WgYEfPjsaZ0crnipF/1mD1puMSDNAgHyr9acAA2expuAX9AKXlWBPNPoAD72Ka4yQR2px44HU0jDbUjHEDaG/Om8HcRSryCKa/EZPpTQiMNuyKamdxUnpTwo27h1pDgZJ6mrTsMcTlcelAIZ+ORSoRyO9NBCZapAcQNpPekzkkikU71bHQUJksfSnYBMFRQSMAnrQeGpX9PbrQlcBkbfeOeelGMEtn8KjVcsPUVI3PNW9BDHPydKrT5GDjnFXGGBjtVaaN3QY61VNq4SvYSIiSEDuDzUz4UAd6ZCmyPBPWpGzkD2qpSTemwIhkbMftVa4LZBxwoqyjNySKglQlyc4U1UHZgRLIuQNuSB1qaY8iTOB6VViOyfGflqe7jMsQdTxW0lqmLoChUbcOpp8kxU7QDuqH5goIHI/WrgKBFlfgjr70ppL3lqMgmz5QR+rdTTLZtjFM4AHWpLoBoWY9RVKDeZA3RaUVzRYupIke67BYnaO/rUsh3SMnUelL5nJyOB3pFGPmHIPf0obvYLDZGCR4x81U5VLpuJ4FXJ49y7welVHKtbnjjvV0+4O5E0vmOuFwPWtNNqoPmxurKD7ztAAHpV+H94wy33fatai90ImlAuIyM596XeN22oUlG7ywOKc6gKOeRXC4tsoUttJzzUPmbldT17VYUKyZH3+9U1JV371UY6iGyQhsKST6imxwq0ir0I6VI+4ODgc9qEwpJPWtIyYi2WGVBNQTthzgfSnH5sHb8xpZhnntWfW4zNcASkseTxUlvGxk3YzinSbXfpwKfF/riccHpW/M+UksBS24sMmnxquM4xSDhS2cA0DKn73FczbehXmI6K+Co5qGYMDyflqwSE5NRFAzYY5FOLAYUL/h+tVJVI3A8+1XoQQ7bhx2pXiUqxxkmtVPlFYzo1DRAd6kicphc0ip8+F4pWQiQk4+laaO4y1GxKnNTpxF65qtE37v69BViPAXBFYTTuO44jKelJvBPNPB/LNMkZcj0rJLUNUM3DcOKa44xUjkBhn86idgBnqM1aTArTJghgaQHa49D0p93/qS/pVWKTJwx4rognJE7MtSjjdimQnc2CcZ6cU5SxA54Bpr5DlumOlLyAdENshfqasAYG4dPWq1vIZEIK4NW41LKcjj0qaiaY1sRysq4OOKbHKWVmI5pZFV/lbp7UxAUUjuKlJW1DqSRMWyG705QArAdO9RxfdIHWpMHIHY9aHuBKjgqMHpTzwAahVCrjHAFTHk4x2rJpJgMYk5XHBpmwKChp78rx2pqMGY9OKYCMOh7VG4YuMDipSxIwvao3JzgH61cWwHZJye4pjZBz+VOOQB7CkJC4GKasBTeE7yxPFPTaCVI7VJMMrwOfWmJt3Z74rRNtCJo8lOeKeylUJB60yJiYjnrUueid6yejH0GY2xAd6eGDjGelID83Tg0igE7QealgNbBQgjp0pofzIgQMHpTnH8NPhVQh74qugCOoVMZ5pC4xv/AIT7U5gGfBHFOQg/L2qb6B1GxklGJHXpT06c96jAAfYe9SKMLtB4FDADnk0zGQWp7k4+tNyQDQmwInO6LI7UyJjtYg4U9OKBIGBQHmmQkBzGOmK1SsmBPGSUYA09W4IPU1EhKPjHB70MGzjPPeosgMy+uSs2zt3FRRuVOcfQiptQj/fAkdelCJmPpy3SuyDSiiS7bSSNbjfVpMbRk8iq9qS0AVv4asj7lctTdldBRjGfzpHyy5xQ5wgIFKD+7+tZeYDQBs9abgMvI+lSPwo9aiIO4enemgI5Cxkwo4FOJ3EinMmMkVFsI+Yd6tWYE4Ofl9OlQzD94uBlqcrDaG9KOrn170LcBqDa+TVg/LwO9QAMGwRmrHalINBp4BJFNB3Zz0FKTztPUUhB6VIFaRnwxHJ65qK3di5VzxVjaOR/DUaJ/pGOordNWYEgORz+VJExGeKB3JHIohxgv2qbaAPJzkdhS7gE+lAHPXimAjYfapGPPGTjrTXOGGOhpUJYZNGPlJxkijqIinXd8o6YyajiOQ+DyKmckjIHNRIoVj71pF6CtYh2FcDsatWuMevtSRkbeeDUkYCscCiUrppjGyjdTAQY156VIwIU8cGowu0YFSnpYRJkBWAOaYDgcdRShSE9zSqMSfWjQY4jcDkVUfLNjbx2NWySCDTHwR6CnB2DchCgDY5zSwFQCRwKSeMSOCDwOtOgwUPqTVvWNw6iJky5I4NOAESsAcelJID8u31plwTkCluA0fePPWmYGSOlJnY4+br7VK6YX39ap6ANkTCgA0okMh2nt0pQpZc4wB1pkI6nvQthjyuZMCmvlkxUi4OWHWiNRgjuDSvYQhUBce1VTGcgCtBiAvI6VGu3G8DinGQrEkQBweMgVLkMpPcVGAByB1qTGVOBgmsZbjIiuZST0I603ZjI7VIwBwuKOrH0p3AhBOd3ao5ZPmD/AKVIoIY4OVHtTZI8sAOlaKyYdAUEAgHg1MhOzrzTFAx9KfjB4qJO4C/wHNOzhABTT94qaUdDmpAa2ByaVQRECPxpAuRk9KVW3L7UwBie1I2cAU7HYCmEdMdqEAYwcGng5/Cmk9DSj7uB260AxjN82O1InX+VIw9+tBOGXFVbQAbOPemA+owakPLFuwqML1BqkIlUkpimuCQB69acvQikc4UDGTU9RkLDHyrTWORkHkVI3yn3NRrgOB6mrQmWEB8sVICQTzTRkcdTSnPmAVmxiryKQcH2pT9/A6UvBPoKQCj5mNIQaBwR6UdcZNICK5fAApqONgFR3rZABqCGTDlDW0Y3iK+pbLAc96pzFjk5+lWDhCB3NUppcz7AOBVQWoMjRcZxxUyAbySOB0qNgQx29+tPibKDPatne1ybEc3A9MVGmXx6elTS89Pz9ajUlAc01sPqORmBIq7HIduOlUUY7gOOatxHL8Hp2qKmw9idiARmmjiXjoRTnwcrTdpAGKxQx42tkCk3H8KVcUdMjHApAIgIJJ6UDgGhQdpoBCj1zQBEQEUn1qNFznIwc1LIw4HX2qMckkVa2EEmNjECoVzuGDyalYcYPeoiAATmriBJHxnmrChcEgdagQAEe9ToMKDnjvUSAaU9O1VLqEuwOOKtsxAPYU2YhlxjOacW0w3KNvHsYZOcVoA/Iaox/wCs29KvFccVdR3eotkRkrnGKsLwFqopYztkVYRi3FZyQyxwMYqCZQRUytlelRSfdNZx3AZENrnng05iR06VEhwpx1p+7MYJFW1qF0ODAdaG4QnvTGHI/nSCQFip60WGKGHI6n0puMgDPIoz8270pHO35jnmqEPzyacD8uai3YORyDUoOF+tS0ABsHOaM4NLjaPWkyBSAMkHNKODTScLS8baA6gcdaVTk9M00jKj1pqtgY7U7DGyDlh60yMBQVqZxyO1RgAMcDn0qk9BdRAuCaaTlSRSysVUkdTUSsSD+lUlfUQ5WLH5ulShhjiockHpzUi8rntTkgHYwapXZIYirxPy8cVRulBzzkinT3AoqR6/hSdQDg4pFGDkrmnHK7gRz9a6zNDTkvnPFA5XI6+9J1frkfSlGMY5OegNMryGFsOeR9aQn5fWlO0E9/6Umfm7HP60x2uNIA7cntVuH7o9fSqpzuBPSp4iNmM5PrQyW+o+VipHcU0tshf0KkdfalY5X3FMb/VuM5+X0pJDiz01vuCkAJH0pT3xTCcDivmDUQg54/GnDoB3pwxjmk4C+9ADiD+FH8OD1pAflGaTqeKAHqBtOaQnAp2AMUj98Cp6iEDqzhc4PpUmOcHmo0jVZA+OcVIM88VTt0GIPlODSgcHijIDYJoB64qRWGNncMD61IxByBSDAHIFC8cnNNtWGA+9xSgZPtSKec0fxjFSIVhjFJwSPpStgmlyMYqrgAPoKawYHmjnOOlPPuaV7IY3HGf0o5xS5OMdaTnjilcQfw49KUjgGg8NSgDoaLj0D7y5FN7dOaU5AobIFFwGZ4LUhPGRSuCQMU1gVIx+VUhB1I4pzHcQOlMyRzTtwLDH40IYrDORUajJb2qRsdaj57dTTiIRBlvm6VOwIzVf+LjrU/VMjrRLYYuQcE9TSrwpPXPamkYAANI3IzzgUPe4mRyDceOvWm7twC9x3p+Q2CBimEbZPrVeQIWM561LjnHao14QmpM7kJJ5qbNgNVTuwaU84PYdaAPm560ZBBHrSuMiaQKC2M4pITuj3EHPenEY+lOA446VfMrCEACk4PFKOAcUwH5f5U7sB60twE6ufpUgOVz6cUhU5NID29aTGSkYwRwaTOc54NGPl60ucEcVPUQ3GRSxjK4Pahu7L+VNXJTrzmqSGKgEkmO460Y/ec0u1VbINC/6zJ6Una2gkIAWcDPAPFSA8nPNRAfvB2zUjfKMGh3toA3PBB4x0qNpB1qSQHiqk6kN7U4q71DYmQhmzmngZzkfjUScKMfWpAdp47mqlGzGNjwrcj6VG0TFnbPHapT6+hpWwfyoUrO4hvmKIxuPQVAk4ZiAM56UvlmQBj07ipFhjRgyiqTilruHURY26N0zT8YPsKeTzxzR1IzWbk2MU4yKOrD07UfKX9qHOG9s1IDSB5n9ahlXIKk9amcFW6daicEyk9hVoW+gKuwBfQVGyjYeeDUynJIqNlAjI9aq+txjADJAQ/4VGu6PDYznpTzvU7QeAPzqTbwBV81mBGEYuXbqBwKmjz34pHbOCBkU3zBuGTUv3gJVBJbPanjBGfzpucPj1p64UgGoewgbqOaQkO2B+NK/yPtByKQDa3FKwBnHB71XlgQkk8kVO+S+aR+S2etNSAZGuyIY+mKRjiMH3pQSVIpRjgD9ap6sBC+I8gcmpB9zk8ioUBEhY9MdKnHLY7YpOwDHXMXB5pD8pAHelPAJpp4BY9+lJASg7WBxx3oxzTU+Zc9qVjuIPak1Z2GN+6eetOOQA3ekPTHftStwcDp6U7gIc/nSjlc88UEbQBmgE7AKHoAuQ3zN+FMYZ6dBTmI2jHamnpmgAjOFpw5PSmjGCfegk7SfWgBAcZWm4w/PX0p/YetNyCaLgB5cnHIpucnBHBqToD71E33xzzTAeMKAB0pM7BgCoo0ZJCXY47DFSOTuIFNrzEKfmwelNkO3n1oZTtGDkg0k+BGM9RQtxlZCxmGelWcnOfTrUBO0rjvU+4rgE9at6u4CjBY5/CowS3BFPUbWpB/rM9BU7MQi5G4DtTTzwKEbDnBzQWWOTHc1aWoEWcZQDrUTA7QrHJFWHUIxPcd6gBy5bORVxAhWMbX2Dn3oVh5LAnGOlJvXc2ynAB1Leg5rV92JEUWXVOvvVtj+7KjkiqkO8ucE4FXTtDDIPNOrugIpAZIWGcZptvGuzr0qR4xHCxY8DvUVqPlLbsipjZp6gO3p55DfdPGKeu0Ss46Y4FRMPMnYMcgVIpCqR0Hb3o2QyuCzlsiqzRqythvu9qvSDODjDdxVJ8bztOKuD7CaKrDDAY5q5DuKHB7VTRT5jN2rQhCjDZ4NbzaURRJoNzqGHbrU20c7Tk1XBI6HA9auRKFBGe3WuOXkUQgeWcdT3prMfM29jUoQhyW/KnLEGlyR0pXW4IqTHJz6UyFyPvA896fcIqMWJ+UVWD73GDjnpW0bNEs0I3DDdjiomk8wYHAPSnwY2sSc4qJsBwM8CoSSlYZBOxRting1I0qoqp3qKXmQnqO+ajMgY/Mc46VuldCNLJZR6GnEE45/AVGpDKpxQkm2Q4OR3Nc7jq7FEsjkDG3HvUZZA20HBqTmQZI6dqiYBnBxipj2AsLgCopdwPGDSwNuDbugPWmEeY5x0NO1mBQVnFxkninxqXnLA49amkjRSV6GmJhXwBzW3NcXqTqpY4BqeNSq8nkmoV+TgHJqZmGfoKxd2NDt+2mvGrg5pVG5SRzmnDBGB+NZvfQYyRAUFQEBlwO1WHHynmq6j7x7VcW2IiuXAgK+lZSsQ3fArVdRIhUHJNZzxlDnHArrotJNPcmSLMJLlhk9qdI+xCoGSOSadbEMpb0qGcjJK9O9Ta8rBsizC2QuOCf1q5H8u4HpVKEfukI/OrSH5R796xqK2xSIWIWXJOQelOYBjx2pj/Nn2p0Qwewo6AhloxMjBuuasINrEHrUKDbKWHSrCAbuvWlPcELt5IJyDUhJ6jpTF5y3pTxyvPWsWwGHhSxquoCFj0z61aYAgqelVXU8qpzWkGA8NnIB/GmnPUURng8dqQkMVx1prRgOJ4AY0hKh8DgCjhXOe/So5f7x79aaSbANzOGGMCq6S/vMEYxVlTl8AfjVeZCWPHTvWsbJ2B6bFiFz90DFTNwvPWqFrI7zEHj+tXxkqWPQVFSPLIQiqQTk04fKVLCkDAkU7OTtx0rNjGt8oOaibeCVi796mYDLA8mkhwG54NNMPIjbPmEZ+tTJhEHH1qBl3yMew61JGysjcdKJbAO/i4OSehpwyAWpkfHJHNOGeSTwam+oBuzz6CkzjPPalGNhGKibIYHPApoCnHjJJ69qWNg0pTHTvSTE70dPypVbZIWC7s9K6OlxdC1GdzMT0p2CSMmmD7hOME9qlxuUHPNYS3GULkKi7n554qtkdQfpV65jD5U/d7VQUeXKBj5DXTSd0K9jTtJPMj5GMdfepwcLUFvkKR6dKmA3J1rmmlzMY7+DHc0gB2+4owVIoBA61ICuMtx9ajxjLHvT8cHJphPydKEArdBmmgnHNPYjAqMttbA5JpoCKRTjg8VMoGBnqaZICE57U/fiLeRTbuhaDk65z0pVJLZ7UwsdoI6CnRnoKTQwP+syelLnJ6fWlbGeaaehA70g2IwQuRjiog4j+dhjNSyN5YAxk1HIhkYZ4ArRB0BH8zJIxRFhnIIwKf5SpTCSJDxTTvsBJnDE0hGwginfePtSNy2MVHUYjew6U4kYHpQxOSe9HWM5oAhIOT6VErHJz1p7AgkZ4PSo4wBuBPNarYQA5YgcVNCSW5qBMnd7dKdbly3zGm1oBYfoSajPzqO1SycnIPFNC7s44wKzQDC3yr7UKcOM96aq9FJ4BomLqwMa5NV1sIkduMnoKqSTYUEnr0qrLLMdys2OeRVdnJOCefStoUgZoRuXYc8VY2jZnOOazImKqAPzq9bucgNzRONhk0u5VAQfNUNyCZFx3qdmwc5qsxaWVT6VELgKECMS3XtUshOA3X0FRTsOlSEZiVm6Ub6sALYiOahU4I9+tSLzgGnC3KuSeRTukAg44FSqu1AR3phYD5s4HfinRuZEBH5VD2C9mKxXp69ai+6cDpTwpLZ6561FIeSoGMdaaXYCwrZIwOD+lTH72RyDVWFwVA6nvVlTw2O1RJWAY/DA0D7v16UA5AJoU8ZPbpSFuNUAdupoAOTTuSM5pW7U7jTGRLheep6VJjCUxG3vxTydzdeKTARhyCeKRT83Jp2QZMDoKYTnFAC8hcHqaagwCtLIdx44AqJcCfnvVLYWxMDgcmkH8S9/ahuvNGMMeaSGJjgDuKceOlA6cjmmvnaP1oAjdSDmlIy2RTZDlttOXoF796roHQAPkOTzQoHBzSkDmm/dx70APU5U+9JgE4NOGMZpGPOaQFW5J529qZDkDL/yqW4QscAcdzSAfu+RWqa5RdSVXHBJ4NSk5cZqrtY4Cmpx9zFQ0P1JARu6dKP4sCkQ/Lk9fWnr1z+VQ9wBgN/HSk52le5p2ecjpTd1JAUrxvkBx061UjIL7gOtaM0YcEE8VnRpsUkevOa6abXLYT3J9xcjuarzEq545qdM4GOPQ1DK6gkmqjuGghG7Ehp8Y2ryKjB3ADt1qZRz+HFUyXZjWwSMflUMi/KQMZ/nVvaCBnriq9yvoOPWiL1G0QpycGp4GAfNVjnZ706JuSO1aSV0BqcbwxpjEkkZ5qo1ycAk/UUsUvmnjpWHs2VdFkSlSR19qfuLLVdSSwqcHIAFS0G48/dwKaxwnvTeVOc07GVBqQGHAbJpCcocdaViN3vSFhtyKoOgjjOM5ziq/wCvtVmQcZqv9361cdhMeG3MB3FTK2OOoqAY3g5xmpwQF9DSkAvOWyeO1I2MBu+KdgFfrTXwVAqA1K2FM4z15qfzBnpVdVAnBPI9amY7kGOtaSD0F3MHOePSnjKjPao3YAjIyPWpV+7gng96l7AyZD8ufX0pkn3M0K2Rt9KcSSNvpWfUZVgRo3JPerBGVHpSlcDBoI52027u4tiNhhSR2pqAEk+tSEg5FQs4LYHWqV2MCcNjtSFWkTA6UinLcnnNSphSQDTeghqRFVGTyKcCAADUh6e1V3OTnOMUlqPYn3Zzmmn7ppucqKU5DetIAJyuKcfu9KT+PB6UYGQM8dqABsADNIo6UkoJ4pwGFHrR0DqBAB56U3kMeMZpx6E4zQB69qBaEU65XPFRwr6ipZRxTEGGq0/dBoGUmQ805R8vFDcNz1oHHGeBQAY+TB7VTu8lMDmrrdBtPB71WucKvFVDcEZaseQR0pc7V246+9JGw3sKdgMcHiuwzvdgcZzjrRnewPWkyASMflTM4IGOe9FhvYawwxbPWkYcAgHnkGnFgPwoOMH09aq4yPBA4JB96ntyRgYxnv6VE2OMt7VJGc845HFDdweqHsefXNI5/ctjutOPO3A47j0pknCvkZyvrUrclbnphOOo5qJmOQKkxk9elQuTkDHHrXzcVc1J85UUds1GGwcfw1IoIT60NDAcgk9BTiDjcOlIRlcdDRvAAHrU2Akzux2GKODx2po9M0MDj3pLcNx2CDzR5gb5QOaOoxmmxph2J6Gny23EKw3MDTiMZI6UEc4xQxx3pbgJnJzSnk8Himhs5A6U9SAhGM0JABGCfpSKD16ULndg805uoz3pPsMBwMnpTeuOaeQTwT+FNHAx70lYQoGCPSl6YpAQc/SlGQMnNNrTUBp5IPeg5zz1FKeTRjk5NIBwXcTz+NIPSgcjIFAxk5NDADknnpSEZJIPFL97Jz0pABtOe1LWwwPIyPxprEEc0oPy0wnLYxVIQ1vuY70sQwME0NySaB83zYp30GKu3oaaBhz0zStzj2psnyhTmnHcQ0feNTocEfyqH1NPXnkdRTmtAHEZYEU48gr2NNjJcA9DTmySM9ajYNxpwBtxUTrlt3pUud1JgK2D0NVrcBicxbadkgAdxSdDgdKj8wh8DqaesgJlbIz6U4cA1FFnkkU85yamWjAhkBKHBojlEowBgintxz2qGCYSyMAMAVpDVMOpYdAMDtSg5PIxikyeaVSM9OKhjQhOcHNNbjPtTsZfFMON+DQgJk5X6GkOMYHUUDg+1IOpOeDS3YAHUEZpQPvEfhUKqM+wqX0YHrVNJCFIwBTipEfp6UMR+lKW4UHkVOwCEDrnmggnBJ4NBQYGD36UEZ47ilqAjAsdvWo3UMCDU2DzzyKaMEHjmmgISuDxT8YQHPNIp4bIp0gHH0qugwLFG5pvQYx1p7DJWmEkdBUiEgOVZTSomM/SorbO9+CT1zU+DtGDzVy0YDANoPvTgSFAIxSv0460dcA9u9Te6AQAY+tKnzPgnpSbv3m3HFEY/fZ7Hii3ce4PliM9RSNz0pzA8ioCHG4HtRFX3EKqkEt601xyBTwxLAHpRI21wMZqtRiAruJ9KUD5GqHPzbl61MSAMk07dgIyDsJHWofLdzjOGzxVjJLD07UrHa68U02hbj0Hy89qd3+lGPl9qDgof71ZttgKMt16mmg4BzSkHAB4IpH5wBSQCtwOmaRxkEgfWhm7DpTW479arzARSCVA6ihuATim8A5709jjPoRzVaXAYowQc/hUv8AHcUwLke3rTwdw6c96l6jAj5ck00DfHjvTjgkrTUIGR6UJaXAcvCkYoP3R60Dkn3ppPzDHSluApxng8ilY46UjgD8aYzE4PSiwD3OVXnmnE7cg+lJ/D7indRyabsA3HAzQDkEmlyeF70h+83FTYBo5HtTgBsFN65FOz/DTAac7c5pq/KB3Bp5AxtPQ1HgBgpqkArMAxGKYxA+bHzVI+M0yRQSox0oiIQkswPpQDnL/AKU8rg49aZuGcDpTGLHkH60OMqTjpQzBZFHrQSce5o13ApMyphT061YJyARzmoruDzIxt6ipIhtCqelaOzimgFVwenUUrKcEg8UIAH4HWnAbQc96l9xbkS/KOlLtBbJPXpSt933pJAAy4HPc0xiEg7srzUPVHx19Knk4jI6Gq6EBiDwW6VUdRFVA/mEOAKlOI43TOTimDc6Nnt2oto/MileT8K6PUkInbyCQMkHrVklB5Ybr2qrFKISIwMmp0InjYAfMpoqRvqMfOSwIxwapwKQDngDtVyUkgcbsdqqRylZGBXjPWohonYfUcMiUkDAapU27fXHSoVk3zEDAIqYnByvU9qqSaQDHCsxIHJqlPkjEfT+dXXxvx27moN0Y3KvTpVQuhdSrFtkhI6Pmp7XJhIxjHSqoyhLAdOtXYGUqpNbT2uCFUncEJxzzxWjEVKZqF41dg461Jjn+lck2pKyGI3zk4PIp8XI3E8jrSEjzl/wpF6Mx6VMQK92g8ksetZ6KONvWtG4BaMqx4NVJVVYt3YfrXRRdkS9NSaEjlMUSY6Y4HvTYCfM69OtJIV25YZyaU174FeYtJGoHU9ahARenzEVYDbYSMdarJtj5xye1dMOpLuaNvIFgBY801iytuxlT71WE7KpPcngVbHzLz1rKUbPmZXkWQzEDHGaQfOpwce9JbsfJ296NpwecHvXPJWditwVc5UdqEOxflqWBFU4NNZSCwX8KltXsBVIO8sT1pRuDfNg0kp+YDPSoHXftJbBrZK6EXpB5SbwM5pqszAZB5oRjs2k5A7U5JCHHpWa0uMnRSMYPFRSkJyOlPWUE7McetJIvOOtZ2s7sADBkGOtMcApt7mmwoVPLbs9qkbDOD/Diqe4+hVRDHIu3vTpIQUIp7dRgc1I5G0jvVNvcWljOmIt1A9eopttJvDIxxmn30ZdAfSqEMhjc8V1U4qcPMm9mbSRhVwealGRGBVezm3x/P1HepmbgjsDXLOLUrMpEbo7HA4BqFmEEyozZB7VPM2NpH5VWntw0quT+lXTfcT8i20JABU/Wnrzzmkjd2gBxlhTkwFPHNZybvZjHR/c4608/dGaro/zY7VKpJJDdKhphccPve1Vww87A6VOceSagCEoSOGpxAdgIxx3qNsIAB1qQruIJ7UhQEE5+lUmBE3LgU5uTilICkMRSbgBxVWb2ASMHzDmkcFeSeO9MQt5hJ5oZ8kg9Kqz0BjUH73I6/wA6sxtuQ89Kqknzl44A4qzCpwzHvRNLcEPwBtwKeuBIeeajLZAPYUjLlkYGsQHN90sBUOSRkmp2OB9ajGFwpHNUtEJkittABHWlcbQO2aFOV5HI96UqH6ioHqIeoHt1oJA4FGcNz0FA5yfyoGLjHy561CxyCamAI4/Go5BjBPQ047iM6Uckk7QKWN1wBv8AocVLImZCMcVS8hyzY7dq64WasxGk3MWB1qQN+7Hc1BCcRcjkU+M4B9DWEkPUfLho8Ac1mOw84Kf4e9aKlmUt27VnSqTIR0OeK1o7kvYtWzMZWHb0q4oCqMVTtUZGLNyDVxMAetZ1LXZVh7kD2NGAW680w4YAGnLzzWNrAI3JJ7U3OCPQ9qXGTgdKGADAdxTQDmALAdqjA/eYp5OeRSM2MEfnTQDJhl/U05z8mO1DfK3HWmgkgk9PSgBpbMeAcU2ENnPXFOIBOKIwVbaOlXfQCYkkZIwaaOHHHFPfgFRTVJ24PWswI5Bubd6Ug5O7vQ6Yj2+tIFbIBGMVa2AmAywz0NR/x5NPflvSoyeCw5pK+4Cq2W+lA4f+tKo4JA+tNViSOMUwFB/enPTFLn5TzwaZuJk4H408AFKTAhkwPrUDYVlwOtWXKnr1FQK7hWGOe1aRExCwSX+lPiIByeM1AyAyb26mpIxkgdPera0GXFBBIPemDC555NOI596VgvXvWOwMhPBCnvTnUFcd8UhUbx7d6Qtg0wKssMMfRMn61Suk2upUc/zrTkQlMmqMg3XEYbjGa6aTbZMtCSJeMY+tSQu3mMAOKcnJ6UOpHzqOlEhhcA7Dt6moI+MAg5q4rboiTjJqONRsLHhvSoTsrDIZ844HNTRMZYwoFKyl0wvaktmKAgjk9KL3iF9SeOBVw3cU9+VNO4WPBPPekZSVx2rC7b1AqSsEUkDIzU0LBk3Y4NIQAMdRQDnAHStHqgsxBkMS3bpVcMGdlI4HfNLdkhl2nr0qLbh+auK0uw6k8ZIdcDj1q506dKpQny/lPfpVxQWjI9KznuAj44ApDyuO9K+PlpFOHwamwCdT16U/jOT0piHEmKfjg/rQxkC53gAcetTNkAZqMj5QelPA+XnFN6k9BTxTemfSnc55prHJwKQxrtjgUbAGz3prMRkD2qRQSAD1qtkAAZOSeaXHf0ox8+c8GgcZFSArEED3pkgPNO4HAqN+RTW4bDSoJB70g+UscUo9KGGRt7VQCH5hn1oVdzc9BQQQAoqTjIA60XsAi8Z4px+6fSkH3setDAKD3zUgRsc8jrTJAc8dO5qXHIwOKY442d+9WmA0MOgp0ZLcngUxQRJtxwKsYCrmmxD+nTpSZ546UA5UYoPOO2KyGO5yAeaXAJPvSAHcMmhuDxQBAD/k1WZcMfWpixYse4qOTCxg/rW0dAGZCrVeRBnB6ntU4GXpGzndjgVpHQTINpWNecmplbI4/KoiTjOeRU3uKt7CVyRcZHpUVwQOD0qQY3dKq3IZhkdB2qYq7B7Fadhn5eo7VHG3zdOvvT1Ubj1471FLiNcjiulaaEt22JJhkjb1FFs5jba3HHSkjfeqk5+hpyBWckjpS6WHe+peRs42irADKg45qvbsFbHFWwS2ea5Z6FIiZgowe9S8bQaZje2e1SZGcdqhj6EbjgHrUbDap9alYEjAwKY3Qk9qaDcST5lA7+tVycE1PySMiomGSQSauImR/wDLTJ61YVjuGahU5PT8akwf4abDzJerE9qjY7lOT0qQckcU1FAz3zUBcro4G4mpEAQN3FROmdyYqWNv4TyBWj20E2KPmXGMiljfI57UD/VkrTx8qA45qWFhyEmU46ZqQj5qI1wMgc04rk59Kyb1GDmo1J3nj61KFz1pr+3ahdgI1BVsmoNhWVnJ+gqzn5c9faopPu8dO9XF6gMxgbs0qNlsU0kHGOKkA4yetNisI7MSOMVC2N4HerEnf1qr95uDg+lOJVtCTOceoqTndjpUY+ZhipAM9uaGIcvLUvO49jSBcN04px4P86gYnXvRk8UhAXkdKTknOaBDgfmxSDqaBgPUbSbcknpTSuK46UZ+lM6Ng1KTuXOeKicH8qaH1BvmOSacMAD1qqrFuD271bQfKKclYN2KQNvNVLhd0ec1bxxmq8p+U+1ENwuZIGGzjn0pSeM4/CnZ/fHjNHr04rtIuRserAe1Jt+ct/kVJwpBpmcueOBTJTGOdvzEc03joRwRmlAIYg989aUL8uSMiqKehGRu6kjPapYW/eMcHHrTCCrAj1p0WMnv7U3sHQmkOGGDnPtTHH7thjGQak/gyeDUUp+RvxpWV7Ind6Hpx4xUWQxwcfSn96QIN+QOa+ZRsOC+nSlIxQCefalI4qeowOMdc0hQFgAacuCOwNGOQB0o2F1AnOacQDg0HniggAD0FK+oxw2ls9OKFABoO0420Y9TRre4CMce9Nb5iBThzxUb5HPpRuxDgNoIHSnLwenFNHbPFPH3sUmMAMg0KpPOaBx8tIOD7UIQ85OTTc4HsaDkUADuaNBgFx0PWnMG70AUdTzQ2wE/i60mOTS8bvrSkZODSYhqnPT16Uqj5jmlXA4xzSHJNDGIuF5xQcAZ6g04/c/rTT83yjtQIRsKme1MVt2SPpT2xtxTFAU7fxqojALmnRndkdOaQ9OO1Bz17e1HqAowXwfWmSKGOKkGCfeo3GHyKaauIawwpWpIx+79801vujGOtPjI5qpLsCHKB1BxSkZPXrSqRhgR9DSKOTms3qBEMZIpCpYemKftG8nPHtSHjg9DVLQERMCDuzzTWUFwM4NSnBbp2qHHJbuKqIEi4BYdqkU7vqKgbcilu9SHO0Nnr1ocbq4ClflJqJcBvlFTggpyOveqrSMj8DqaIJvQCbfuFOXIbNNA+TgcmpAccGk9Bi98+3aomUlt1SE7V5pDwAe3pSQABk4pXH7rA65pV+8TQOeM8GjdgQpktk9KfgLEcUAYDZNKOVU980bgOVQfm9BzT/4h6UiYWQf3e9GcOfahrqIQD5hz1pQMsTmjAL8elOGD07UrvqAjH5t46mmgnHTmlx8pIpBnJNO4WEzndxzSKc8N26ULnJNDH5d2e9IBCfl4oUZBpW6DFJzjg0BYZEdkzD+9UoHJz2FQjIdT71LyQxBqpBsC/KMk9aXjZ7mkIyopwyVBOMVICYxgdqD7DvTZAWA2nHrUnAIz0p9AGvgMBUTt2qQ48wnqtRNjIUjg04oBeGPFMIJI5zn9KkCjBx1NI6gkN0209nYZBHwxVjyaQ4IxmpFTPz9+9NCbctiruhCFsqNpxSLuKAH71IVOwgcZPFKNy7SadgLI+XAPQ0jMdw4peSBzxQ3BA7Vk9HYZK5JI57UxiGLHFOONynORUTtsk2nvSS1EHTG7rQMlDTZTkgDtSMT1HSqsMid2CgnqKc0gaME9fSopSCxzwKJDlwo6YrayYFlGAYqehFSIQF4qjA/msM/eFXOVHH5VE4WB2HHC00E7844pswYuooY5IA7VmloANKC4APfpQXJbgcUzZmXPQ08nGR3FOy6AR+Yxfp8oqWFSSGYVXhYyT4xwKuA7ce1OStoCGkZNDHc2M4p3CnikAxzUbCHMQCCPSoxksfWnOMnjpTMkMrDoO9NIYqn94c96U/6wUoHGfWhR1J60X6CQjAjIPrxTWGGyOtSYOCCajxljQncY7rzSNhgppQcpj0NMkbaN3YUAKcFiaiwenenx5KE44pgbMnHQdTVWdwIypbHqO9Sl8gH0pickk0gBDkdj0qnroCHP3z+FRjKgbjUudzAHpUEshD4FEewtCVsfIe9Pc5+8eaZwEHrSjLfe7U2Am5dwJFI/DHNGQT9KYoc5dzxQlcYkwLBi3AqLIEYZuSKsSKCuc4qvw0dXF6WEQEEyM56d6fGVFqQo/Co5AzMOw/nSgFU3KOPSt91a4iMEoCdvI6VZhfjcMe9QAEyEtxxUsGPL2IeDTldoY4h1Of4WqBYyJmB+YetTyuRDtxkr2qIu4O3GCeamKeogmjDy5TgDv70sUnyldvNRnzHygPTvTzIqRkAfMOM1TjpbcNhH+/uI4qrMwlkG0bVHf1q2isSpYcHqKZcgZAC8etKLs7DZSnOxlCjOak+ZVwg6daVTlcZ/E0gkyCeua6OlkhF5JCyICOo5qwgJc+4qrbnfj9KsRP8AvSSeK5JpJsZJj5iSKIcHII4NOZlDbe5prDy+FHJrJbAV7lSVZcYA71SZvLUE9R2rTmyeT6c1nzRgK2fwrelLUlha7pJDk9aluogMAd6ZZbt+ccduasTLuHB5FVUdpjWxQaMoWb24FRGPzAG7gdKsySKI8EfMahwu4bfujqa0hJiYiRncA3ze1XYwuGB4JqmzDeWUc1bjYuQTwe9Od3G4LYsKgEan3prg4JAJqTHOAfvdqUgmTB6VyJu5RDbuwUlj9OKkU8H19ahupfLb5eecZpkF2pkwTVuLkuZIVxLggHGOfSq6EuxA61bnXnf1Bqoi/vtw6Z6VpB6ahsWYVZASRxmpiw5YdfSmsWUYxyadEoKk4+tZyd9RjhFtG8mnZUxh/Wlxv+X9KSUZiIHUCs3fqMrRSqX25q3tIQcc1kwI4lwevvWueVBz0rSrFKwk9CNjtbeRUbHc+cVLglOtQt8nNZoBi4YhSKp3FqqAyL0yM1ajDeZntSagpMHHX0ranJxkrEy2uNgT5Rg1YmGJFHrVa12kZXOfT1qckvtJ4Ip1G+a4xsz4kAHOKilk3x7U6ilk3Fzt6jvVZoyZCGbCjuKcYoGaUQKxjnr1qQL0Hp3qG3x5WAeR0qdWJQHPXrWE9xjFADkj8qmAG1vWolOGJ7VIgLNnNSwFGPLweuKgQElmx17VYznk1FIjKchsL9KIvUGGcDmo2baAR0NPY/uzULMAAM1UVcBZGAVcc+1QSKzMD2NSZAXnkCkZt6qRwBWsdNUAijD7CO3Bo2bfk6mlV0Ugk89qkIz8xpttiIZI/ukHGOtWFJOE9utQqOCBzSjIIY/jUt30GTNuEgA+7UgyF96aHIPrmlJypx61m1oAwgY96YWz0608EDt0qLGec0ICdCCvTmnoT81MT7gIPWnAnk55NSwG8nrRyD7ClA+Qknk0Y7+tAC9hTJBuXaB0NSA/MFHSmt8uT3oTswKk2BKuPzokkWNicZzT5kyoNRyN0X2reLQh0QXJ/wBrmnoQeemO1RoxDqC3Jp7BQXB70pajBDiPHr3qrcxE4cdquEYTGOMVBI4CBMZ96cG73QrD4sGNT61MQB0/Oq8TK67B26VNJnJUHilUXvDHIRIvHNPUbAR0zTI0Ea/KfenMCWyO9ZPyAMgN70i8vRglwQKAeS3eiwCDGCBSDBBHpS9FJ65pBwoJpgHJyw7dqVRkknoOtIWODgc0K21h79aNbBYMjBwOaRRk5z9KRid+Og9aeuN2T2FPYBxzu5pjEk5FPA4JPSmjvnpUgMYkPyKXcCM560EZBFQquZMHt0q1ZjJzgICTTFOIhT+B8opO20UkxBGMBu9N7E4p4yJCKTopFFwGOOAwP4Ug3beOlOcDgGmo/wA5X1qkBG4ZlBXjHWmlWAHqetTEfIcnHNRF8E4/CqTAinGApzSQbmO4nIFNkXzowTnimwb1kz29K1XwkmkrZAPamsMuD60nzcHPSpHAIHqa53oyiEt82B1poJdT7U7G3vzTTkIzd6tCF3ZIxzVG5X/Sd3UdqnEhXt9KbNnn5cgVrBNSBjVbgYHNThiEORmoB978elWBnHHSnK2wEBmxHt70kMu1Bzzmlmt8fMDio1hP4ihctguXCScY6GpOBj+dMSQF/L64p4G5eeawYyRuTSHK4pAwJ96R89e1SkFyOQj5h60xGH3BSyA8sDmiI5+Y9fStLaANdPmXPY9aZMmVJHapiu8MDUShwNpGQe9XFi3IozuYH0q8pwlV0QrwfzqwMY659KibuCEf/VikVvmHH40+XtUAB654qUroY4gBwx7mpwRgn1qMDcwbtT1HJ9DSbCxH05I4p/8AH/IVGOZsdqkB3DdTegAf9YRTJeGB9Kd1BNI53DmkgGkDdinoQX5FJwXpExuJphoKeHB60M2fu0qn5hnrSdiKAFbP40xzxxTicvUcnzc46UIGIOG255NOHGQKjGd+6peoqmBGeMc8mlLZBNDDB69KRT8p9aaAcW5DGlJPpzTQMxj607JPJqWIMY4HWmOfvGnn1z1qOQblAFNANQ4Yt1HrUpJ78jFRqdvBFSodx29hTkOwD7oIpw6k5oVcDHakHYetSwHrk5J7UuSSaTJU+1LnvUgVOrv2ppHykY6US/K7ehpFYlK2QDVwAT0phJCNmnsmI+mSajkPyjJq1qBEXUxcdKkBGwcUwqBH8owe1PODGPar3JJ1GUBxUM4wuD0qZD8nPaqd03O3dilFajKxfauAM++agkkEjKD0z3olY52jj3pm0Zw2M11RXUknRg2R3FMVijHsfShIzncadGqyS9aLi1Zci5jzmrMDEgA9cVWZcIqD86tW6lY1rmm7q5auTgbelB4AzSsOAKaRzjtWAxHHHpnpTCuUOetPYFjijqaaYEXJABqBj8/SrDjnIFV3OCSetaREKnA5qUctmokHHrTzkAHvRJAhynaCe1A4BOeKiKkFQDkU5sLGQOcUWAarjGD+dIjBF6UrKuxQBzU4jBjGRyetDaQbsYM+WMetWETCgNzUK/JgVYBzUSYDlwFxijJxzSFgO/NGcrWYw3HaBSA5znvSDAHNKpBP1pgNAwDz+FQkcEHn6VKTjJxUeQeAOtWg0IgBvyOhqQNz71Dkq2MU4FsiraFcc4yT3qFfvcZ4qU9cnrTFXIyTTWw9hY+DkVIjZJ9ajUYPtTk4z/KkxXJlBI5p2M0BTgD0pWwDWbYDcZBpuB6dKeRyMdaaV5NACZ3H3qtMrMCO/rU469aCuTzVJ2YbjY1KrtNB54zShgR3poO0kdh3p6h5AsQAOOpp6L2pAeQTTlPzZzSbYC8AEVXkGVINWCeeDUcgw2O1EQMWVSsoJPNDDIwB06A1NdDMmMbQP1qBSxbjoPWu5O6IvbQOQoweaFYZ6c9qMKX2+p6U3BOD0ye9UOw1uT04ppPGAPpT3PGeuKZhtv600LcQkj8+5pVbL44x9KQqD1xwegpQRuBPr1qh2RNuJGOQfemSNiJl9ulPAJXnBFJIF8pscHBNT1Ej0wDndSdyac3TApGzjivmFc0EU4zinnrTCMDHepOOjUPcYE5xgDNMVwJdueacVyBg80KqiTcRye9NWvqA7ktTiuRimjBbinAepqBDVXHXr7VIACpzTQcEkc05Tub0+tVqAwrjnpQw3rmnHP3T+dJnB4pbAA5Iz0pwHzkgZFIx7igZA4PWlfuA0H5iT1pw9DQAQwJ6UuNxz2poBvBGc8CnKowSTzTSBuGCPenHk8UPToFxFJBOKTnJNOBIpNn61NwEBx1oY45FGORR/EfTpQMUjIBBpd2eOlGOAQaTJ3NjvRboIB904oHXNLwT6cVGGy5UdutNK4w3fMaaRlwy8etKP1pfukHtTWmwCdWoRsdelDdcDvSHGKEA9QAd3akY7mz2FOQYTb260g/iB/ChqwiAdcetSRkCSofmEv8As1MikyZxWregIkK/LnNCZxhqVwcfL1oXlR61jqMRQN3tTTySM/Snn71MfJ6UIRGxwvzdaaOI2GKWVsx803G6MYOK2SdrgSKODu7jikIOw8cCnOhMeKFPy81nfQfUQYePHemuNnykU6NcJk0OpZBk9arqJkUUmWPoKmYHPWmhFAz7U7qM+nNJ67AL1603nHPNOwdp55FMAzF75pW7j6jkYsc9KU428dQajU4OKkVckkDr6U3oA0L+8LN0pYvlOG9acTkADtSE/McflRzaACg+YRnipCRtqJTluakBBb+dQxCjhs0cHkHHNAHBpoO1GJ79Ka7DHD5vl6DtTXBBwfWhDmNWPHPNI3JyelPTYQAkP7Ck/wCWgBHFKeCfcUoHy7u9D2AacAkCkAyuOhFKDlt3oOlNx8pPektAEYBQCTUvBjJB59KjkBePHrT1zj1p30AQHaPanA/u9uKYVzuzx3p4Pyr7UgF6JuoJyBSKSQy+nNOJDHpgAUP1AY3C7O570wqDjPJxxTnHzbgelNP3t3rTQCD5ADjNKwOPXPWhuAOOtK3Kj1p3GRDgFB1NIuAm3vSsfmHbNM2gEN6VSWgiKQES49KmVg3z+lRhhISR+NAfKEA8iradrD6FlTuQd6djoT+NRxZ2g96kJ6k96yemwCKAuR27UjcsoNGDywoZtxBx0FC7gMfJbHvxSIMbgaXPIJpwx5g4xVp6BYreT5mGJ6VAVdpWHQdjWgcY2gVA6hj8vUVUZaiILZVjDKDhhVtMsMk5qnAmJXJ4JPSri/6vFaVW7XDoSNt2j1qNiQwHc0//AJZjjkVG4y4PoOK50tbsfUZkgnnFSFcpnNU3Ek0mU7VctlJj2v1zWrikr3DccBhgQKQEux9BQHPn+w6URkF39BWS8wuPJOAKVTnI71Gud59KkI5460SSQMUcKQe9Nz+746CntjFNPQAdKhMTF7A9hSAkMWPelYjA9BTW+YcdqYwUfNzQcIwXvSseQ3amsMNmnZrcAGVX602Q8kdqcuD3oYZXOeTQ9xEavgBB0qNhtYgGq7yeRMAOc9Kex8xg46CtlDqMkyAhBPNL0Ge5HFRycE4p7A7Fx1xS8wAkLEWPJFQDMiM4HzVKgMgKnv1qaNFBIA4p3UUG4xVwqhjk+tPQH5uKYMnI7Zp6nCkDqKloLjdoUH3puRgDt6U8DueaY65Oc9KEwGTKdxK/dNQohSMjpzVtyCOKryMGGBxVRbtYNiDaGTJpiPiVgTxUgI8zaPug1XEZMrYOFroik2IllbeN4H1p0KARs2cDtUT8Q4jP4mltCXiIbv0qmvd0EWX2sofrxQqhm3MMGmiMG225p2dsdZSlfYZWmTBwhIo8ktgDrT1iJyzN9B6UgBJUKeQeTV30shCkjdjsKayqIyTSomZSTyO9JdYHz54NK2thlGVw3OcMfamwqE+V2xnpSyoJZd44FJHAHfeTlT+ldUfh1JV7l2GQLHkcHpmr0XPzEYNZiOpJVBwKvw5Yhh+VctSLKRLIcAMaUDeqUSkF8UxcpJnPyVlHYAk+ccGqsqgx4z+NW32hWY9DVEPtlKkY9KuF0J2HW5EZAByaszRgr15HeqEZzIHY4IPFXJAN2See1VU3uC2K2cKVxlm71WkDKmwdQeTV2WPbk7sHFZrAscHoe9bU9dSWKmN3qe5qyjtkDoPWoEA2453etTxKcbiSAOlaVJKwIuQbix3damXluTwarwfOCSeasdGA9K452uWhsyBht7Vkuub1lHWtiVjtJz81ZMqsl0GH3m/WtKDtcl7osLN5fySDPpRHHvLEfL7VK1usqZP3uxqoWMUu0tVpqT8xlkybAFZuR3pIZPl5PWoUQGUZbOO9TsVAXaOR2qXoMnbK4NTDawIPeqsbs747D1qwByeaxkhiGMb+BQhyG9KcD8pPem52qe1RcQyRgq8VCxyoJ602d23ZXpS8+V8xrVRsrgNaQrj0qtfynKYPynqKlLDaOOe9V75AGRhj3NbU0uZXJlsTWbsOgqxnIxnpVa2fEYBPNSb+hHQ0VIvmGSSL5aE+vSoCdiZPU9KsuTJGFP8A+uqUkgjkwy5I6c0qYbFi3kJfBHarkTZQCqVpmTDnrVxWwM9azq72GIePxqSIE8k9KbJxIGxwadGwCH0NZboEPPOajOXba3QdKeDgE+lQbCHZy3FJAKVySMcVWeXbjK8VZVuM/wANMP3iTWsXZ6g/IjICqc9qZIuyL5O9KTjIbmmud2VPArRIB0YAXkVKxEZBJyKiDKCoPJqSRAyAdqmWjBiKQvPcmpRGSDmmAYxzmnKzYKd8VEmAvXAqTAxwOKjwFUkdaeuSuAKlgByOKgbgcd6mIOc9qh+8wGORVRBko+4PagdPTHWkGCdppGYBtgqQJTjcCOlJnccDqKaThyKeM5JxxS2ACAMHNI4CsM9KCBtznpRKOQc0AyMk4Yn8KqBxkE85OM1Zkb5Sw7VBIdsWAv3q2gkHUTYTKMcBe/rUofdMdtRKdoYZ4FNgf/SSijgVbTe4FzOd2Bk1QG9XbeMrWgCN5x0qAEMHVhjniopyaArwFxKSRlfWruO5qiqtDIPQ9qvg8ACnV7ghY8EEE/Slzz05pF5U0vbjn1rEBvBOAaXHbPBpdoABAxSEDGT1oAaBmPntSk8cjgUHlcAdKQ/c2ntTAaoOMmnpyD60HA2jNIh+fvTvcEAHGDQPuHinMfn5pgfDEelICRTvjIxxTSOAtP42U0Yxu70gG7u2aYuAfpTgRnOOKi6SNnpVoCVWyRSBhyaah/eYxz607O09KAHrhsHP4UjEKQf1piEg57U/G9GB61OzBDc5OajUYYnH1qSMFQM9qYfv57mqAbNlk4POahOQetTORg56mq6RMGLMc1pHYOou0+VtzyagjUl2571MFyGYnpUK4cHBzk1onuLQ0VwYxTmIyDjgVHDwpB6gc0pbjaKwe4xDtZgfWmM4DkfpTuQ+2mFQSc84poBkm1W46UpPXP3aQnf1GKe2Scdq06ARNFtye5709HCJg9adMCWyT8v86qyM0cvJ3E9qcfeQrFkEuSfWo0JMhjxgCnQksMHvTlYLJgDrSegDo4AHODzTkx8wBzg088DPrUUceCcdPWs733AeFw3HSnEcAdKYAcAmpM5XNJjIGwxx2pFXGdtLICACo5NIuDnFWthDiSAAKC4GBnk0ijoT+VI3pjkU9L2YxzEdD3p0GGyPSoHyx69KniG2PniplogB8lGB/OmIMxmpWwVIPeosEEAdqSegDg21ip/OpMjA9ajONxJPFSYwAKTAjJCksacOFPpTWwx25wacCMYpgKT8oxSSYOeelKcbaY6/uyPWgBASScj8qcAAcEc0wZ24604nnpzQxBty5PcUKec/rTm44/OkC8YxRcBFIBIzUbsNjCnZwSfemMflIHWmkNiJ9wGpcfKD0qKMHHPSpVByR2FNgNkIZgKAoAJoUZbJHIpP4j7UeQC9Vx607aANvao1YsSfSpFyTk0PQBowMim4yMnrQCcn1oPy/U0wEdecjrToiwYjFRuSX3dgOlTI+8bh2oewdRd2c49aUgfeFNA4NL90EVIXHnlaU42jHXFIDxRkhc+tSBUmG1wx6U1VOCafcglUxSMACMHit0/dAQN+7OetQTKCvI4qd1wpFRtyvTpRECsxAYY5NTKMcdxUJG1gfWpEYnn1rZ7E9ScHCmsy4fJJ71c3k8GoJ4txGBTg7MGVduABnjNBQk5zn+lSFCGCAUOjAHP3fStuYSE28e9JEhEp9fWgljjHWrKxgnHSk3ZXDyLaRgqD39alI4FV4WKkKenY1Z4IGK5JXTLH5696b/Fjt60hbBA70jHn+tRYQ4d6YRkE+lOB3HP60z+MimgEbJIH51CRmQ81N6nvURX+LvVxYPVgBkYz+FOYYwKaoJ5BqZUOMnrQ2OxHtAYNQAuc461K+FjJx+FV2YtGOcEULUB6tmXGOKmZ+CaqhdqYxUjfKvPWhpCAOd/PSnKWJYZ5qDaY0wD1qVc7d3Q0NAKobfgmpd20/MarhsSqD1NT/efn8KUg6Cqc5pwIAx3NJj56THz1I9QkIKkflVcHa+KlkbaM9qhLfPnFVHYNLiltwyKUDnpUceTUy8nJFU9BDdvNOCAsPalx83P/AOqgfezmpuDG7QcmlUAgcc04Dg5oUbFAouMlHSlYYGTQp4GaRiehrMQjfeBpnO72p2CeKTtzVBuREfOaTd5mRQzAEntUayDy9ynrWiV0BKq4BpWHy4NMDEnJpxbIpNMY3J79qVWpMkgrmkTAoJ8iUnJI7U1jk0DOP601+DSQ9ijdDLkdTVJHIcc1busq2c1S6s3G0e1dtNXiTsyRyS27PTk0wDcwHXvn0p2MH1HpTckHj6VfoJtgcHPc9qF4yMigsGJA4/pSFCpyOgp2BpJEXrzzmlHPXj2pc5IUnHNM53c+vbtVD3LYxtWmvzFIMDoaejcYB4qN/lWTudpH6VNtSUenD73NB4cUvfPahQNpJ618ubajOS2acRkUA45oB9KQg5HApeMjPWnEdD60hHHvTv3AVfunPelHcZpM889aew4yO1CATlV9zQPlwTSnAQc81G8hVMY5o5XewEh4YUuRgjGD6UgwUHPPrQRu+ak3YAwBwaaAMHnmnH7vFNXgZNADgTyAOMUgGQeaTOG9qU539aPUBVwMk8mkUkD+lA6n0FCkbvrRqkMXPzE9qQ5xn8qMenegg7etGghDnFL/ABA9qTp1FOBzj0pALlSuO4pF4Ge/emg8cU44IBHei47iYy+KaEAJP504/K2QaTnJJ79qBDfukE96RuBmlbnGOlMmIGNp4xVWuA9m5BI7U04L8UFtyAnt0qMNggnv0p21KJ2ODx0oxgrjvQPmIJ4PpThgMKl6CIJOHPHanA8g0FC789qQZU498Vo1eKYidCMEkUKvJ56dKUfKgB/OkHrWbGAHPI4pruFYgdM1JyFyelQTuFXeRyKI9hXsQtzIRjrTsbTgU6H95lsc0MQOtaNPYAZyHHNNMg38ClKncCelRznZyOp6CkkBKvINRyRswBDYHpRESq/Wnh+CD3FOzTCwKcxj60q989KTO0Ee9HPk9cmkwH5G7jpUfZsd/wBKmAHlj3qNTtPTg0k7DBRk5NSIODk8VC7HcQDT4myGzT1e4biqc/WhQQ/PWlxjBFBYEknriklcBjnGSR+FSr8wzmoZMtgA8ipFG1QBRpbUQ/khvb9ajb7qg08naRg5HpTM4kII4PSlbUCRsFRjtTZAAKUH+E9DTXz2pAAIAIPPpRk+Xt6YpONqkdad97J9qd2AmMDrTXAwADRk5C54oZR0zSHZDjyQR0HFIOCSaCNoA7UpGF98U/NCE6cetKOvzUhztH6UueSD+FG4BnGTTs7Upj4CEZ5pc5IX86VgFIAjPqabjgDHIp+ctxQBg4bg0IZGBu4zwKQk7SaeRt69fakQcEHr1qrrqCIW4ZQeTTGGVb0NSk4bJ6UwKGIVjxVLYCNUADD1p8cahd/rTguXwacgBytNsBEO0nPQVHvLkc9DRMTsA7Zpig8BeacY6XAsId2Vx1pyDccDjbSINp5oyccdTUPcCKRyJAAPrRGx3c1KRxu6k03ABJHT0p30EK+FQ+9Q42ru71IwJQbu1NbDHae1NaAQSJtR2657VPEw2AdxSD5mOR8tOQKJDjvV82lmA9MnJprjDnninxnBYUyToTWa3GEaBRwOlSLycjoKQEhd2OtKo2rnsaXqAZBUkdTSABc+tNDAZoDeveiwDiAE4604DAyaCAZBzkUrY5GaTTAQDge/SkHG4H8KcSANp69qTBCe9HmJic8DrTQTu5FOX/V5PU0dVJ70WGIVwAKSQ9O4pw5XjvRIFCj0oAiDYGcdKdgFcnrUbHavtTwSU56etUIqna0vI57GlRAqkfnTZXAlUdjUwHyGtHdJDIWx5nzZxTpG24x1pSSzgAdOtINrPzx709OoD4lxyTUgOMqe9AGaO59qzbbAbjBIoVtjEY60pyOT0pFAYlvTtTTsAdBjPFMJ+Qg96ccAHPQdKa3XceBVLuAmAIyMcmq86HIyKsPliCOgplwpdc57VcW76iZXQZygOO9RLlpiAMinmMopI5NV/MYEknn1rda7CJiNsZjI57UyJXZsAnaKmYnyQy8kVFCXR+eF9aE27h1LrZNufWoJQWjG3qKlVtpAz16UyUgMFHUVkr3GyM5U+461FuYAkHGakIMk20/lTDiPcCck9OK0iIa7GMNzmmM3mQgvz7VMIyyDdzUUyKF9KpNbIGQod3B796Yjtaq4bnJ4oY5wN2CKbdZlaMDtW6T26C8whkKt8w61sRN5YzjI9ayI2wDEetaSyELGuKzra7IcdieZwoBIyTTXIUKueKWRN4Vj2pm3cPm6iuWNkhkxxxnpVKVWMzDGV9at5DkZ6iq7vlyOgq4J3E0V0wZPlHTrUzSNkbecVVYlXIzyelWIiY2X0xxWslpcSG3MmWBbrVXOJiP+WdWpxvG7GCR0qmxIYjr61pS+EHuPIyxI4A6VYRy4weFFMyqQrnrSlVCDDcn0ptXWoti3b4JJHftVjORweaigXbyRzUrfdI7muKXxWLIHAaQsPWo5wGZWPLLT2BXJppBKkBc571oklsxDlYSJtHXFZ7qQWGOR61eWIgFWOM1W8rG4k7gKum1dgwtBhgSOTU+QJSuMg1CGU/N0xVlMMQ2KJdxkm3Z8oHepFTHJ79KYdxk2mnA/Ng1hJaAPxxjFNdeMZ6d6eTwRSSgleKhDZWO0DoM0wjJ/2ajnJVKW3YsnzcV0KOlxXG7QJckUGJZlIfv0omYsy7fzqaJfk3n7wqulxFfygkZ3dqjUrLjAxirMyZXIPXrUDLtXaOCacZX3Ak3hJNoGfQ1FM2CPl+b1q0iqqKx5zVS6Uv06UR+KwCWzncOa0wcKADzWZDHJDgkZBq8EARWJ4FKsle4R2JZcFDzj3pEy0SbaVvmjDD0py42KBWPQY4jCbe5qscrG+4elWpBnHaoJMM22lHcGRdAAOBTpByOeDSBhIxx29afIcYJrQCvOcHbjNI5Hl9ecU+Z8DGKrSSDCqOpq4+gEkTb6sljwVGTVKLPGD061cj+7x1pVI2YxBgdelPTjB700ninKPkHtWb1EOA45p4O0HPemjqM0qnJJP3agBG4Xk/hUR5BapHy341HwMgciqWgDlbI460mCpUjvSE7flWjc24L3FMBWyXVqlVmIOaZ820jFOAwoGKlgCgc9eaGzsJ7ClzhTxzSt93b60gK+MIx/KqxDyL8pxVog4bNVSXGFFbQDS42OJmk3D7vcUiEi4WTOB3zUikgHnOODUDjcoI55rXdgaSHrxioZWMeW61LGSIlB60ycjaBjNY2tKwXIVYSruI5Wpw3yhh1qEbgoA5PrT4yfLAb8DVSS6AWRgN9aBgMaZESVBb1qQfM/NYtWARyV4POaQ8ZA6UpzuC02T73FCAF5Y1ETjnHHentlXBHeh8bSPzqkBErlpMHmlVsSE44B5NLyqE9+1NLEKBjrVA9ySVgMv2prA+ZkflTmGIgDzTcY6nOaSAmAIB9DSEYXFLzuGTxQ+ckjvUARn7maibLjBGM09jtXHUikb7oJ49atXAcuQBQxIU+tRqd+CTgVJkeWfWnawW10GopPU1InB4pgG7BB6dqkjHNKQbinG0iqkuQwwfpVpSPmFQT8Rk0Q3EBIxgjkioypXA/OnMVODUbowk3KauIyNhtkZM5zTdu36UsincHb9KYjFhnHQ1sldC2ZcjOEHPXvTXOwYHU0sWDGcUp/1fPWsdmMa7kj39acBlvfFMBAB44qQE9cdqGLUjVRhjnkdaZNKY13EfjUiKMO3c1HdDfEMdqtWbSB3Io5mkO0mnxR5JdupqrEywqTn5jUqO8vJ4WtJRtsCZbjKgYXvTWIV17k9KajqSFWnSYDg+lZW1AsA5bBpqtiQ+nalGR360zdiUelZpDJDnGMfSlAPIpCx5PcUoy3zUgI5OCOc1FnD4zxUsgydw4FQheQauIDjnPA4FBIILAHNPYkqSKj3fIQKa1BsjXO7dn61Zj55/SoQv7s8Z9KdEMAEUS1QE7dvWoujHmpWGRUXJZsjipQdQY5IwKlUjYM1ERnbjt1qXA2fShgQvkSB8/WiNgXJI4pJwyoOcZ706EEJkn8afQB5OOQeaH5GPSk7cinHueoqQIieuOopeQwGKT1x+NPHzbT3qhCsOKYWO7GakPOeOKhkYJjIA9aUQB8eYABxTABv+tG8B8/pQARJVjHJlhgdPWn45FNACAgUoOfrU7i8hxGcYqOQd/WpO2O9RP0+lCGAPynHepOcY9OtQbgW2jqKlXOdx7VTQCNgNjPSgnnp2pW55ppHy+9ICNSC5UnmkiDLIR/DTVAWXJPJ7Vai7jFXLRCQ7gNijHWmnk5NOBwRWQx2SRjFHbHWk/izS5weKQFadd478VHgcAHirB5cg1WlAQnnitY9hDpWw2fWoXJAyDxSvh9vGBSZG/b/DWkUMrEgAe1SqwYcflTJQCxAxgUqLtXPc1rfQnYf904zxRuBbkcCoXfDZzzQrYbr+FCQXJHKgggZY9KjeM+WT3PanlQzA+lTRofLySM9qXNyginbwYXc34Cp+p4FSpE6qRniosEEhhTc1JiSsNYksB3q0jDb16VSdsHjtUsAPXtUyjoV1LgGWDUL1JPFKp3HHYUDjI61i2MFHy4pD3xRjjFHO80gFxk4BwKjI42kVLgk1G3HNNMLDUADYFSlgBVYvl+KcWyOKpxbEEjlmwKiRCVbnkU6QHeMelIg2q1WtEBLxtAxzUfmK+R3FOJ2puHWoGxuLdjSihO/QkDYADU45832pCOSx6UPu2AjrTGNOXlDelTK+RvNMyRIrDoaUH5ipHy0mBYzzn1pD9/NGRgYpD9/rWYxJORjvUHAJzVhge1Vic9e1VEByD5+KmXHJqBDuFTDpgUSFsHQ5Hel9CKRuuO/rTsknFIBcfLSkcA0jckCnY5ANSMdt+UetIx4GO9Ln5TmmY3AZNJAxTwKYRxTjzimtnNNARSqDgYzmoGjUAoBirZ5ORVeVSH3dBWsWDRGr4ypFS4KgZqBQ28sRipWO4HIqpIQZGd1OYkjI5qID5SBUq4xSeghwOFxSjAHJ6UnBODSj0IqR9Sjdgc1m5G7jjFal2QFz0zWbn5yR+lddLYh7jsMflHTsKQcEAmnZO3OaTknj8avckQHnI+9TSdxDd80rcNjFIMAn19aofqMzt9evem54wAaeePx4IpmSGIPSqSuUtUWYWyMZpJfuOcc4OKWHgHI4FJKCEc54IqOoj04/KM9qMU0fNTj157CvlzQQ8DFOUcH2FMPUU5fmJOaewxwHejtnsaQ8qPWlCkLg9ucUumohRksKdkAn0oXbs/2qAp6flTemgCNxx61FJncmOlTggZyKjwGzz06CiOgDlcGPjpQOVIqOMjlSOKfjIJHFOSV7oLgc9B600klSaXGBzTB7/lRDsA9SHI56dafnOfamInzEnjjgU5xzilJALjCk5oX0pMZbril4B4qUhjehxTtuFBoxg/WgHOc0AJgYo4A4oOQKTv05piDHy9fwpV44PcUh9KQZ34NCYwzgml5PNNfHBxigDaRzxQIXH3geuKg2krg9TUwyCSajLZm6fLVIAcYIAoxng+tOchnPt0obaFU96NRjxkAH0pQwLAmhTiPpQOevSpAa2ex5pgPy80StgD1o42j3rRfCKxKrFkI/KnFspjHShB+76c9aRepxWbAcHwpHWoJQGUkjipgVAOaaSM+1CfUGRr9zIPamMw3jd1qRxtAI6VTuAzsCp6VpHWWoFkOd+O1Nmj3sDSwEvHk9RUo5OO9ElaWgEOxtwbtSpncVI47VN6rimscGjmAbkbtvakPD7c9aRiM9MZqTqoIHOKTAB9zrwKMcgU1Qd5B6VIwGetSMhkTIJHBpVAjUKTnNPI+Q561GeisB0/WqTurATMNi5yDmmnJjHHIpW4z70zJDgEdKXXQBjKfNHrU8fvTPlZ9wqTGSDmh7WAaTjJ4oUh3PPNH8XNNjBWVjxQlcRIecA9aQ5Bbvx3oJ+UtjpSA7wKSuMIyBnI/GkXkHFGdoK9+1KvygjvSuAKBhvUUzO7JzyKeB94jqabgEcdT1p6CHE5Qcc0L1AoPJI9KM4ak9R7DedxHpSnluaRuJRS/eJIGcVVtBCyDt3IoUYJ47UpJzk0h55PSi4DlGNvFKfvAnpmkY5CkUjckCkAdH+bmggbz79KcV+Y80wZz83boaW4EE7fIRUCyFQpPNWJUxn0qsQwBzXRBpqwXJomJlPoRUqgAsQelV0JiAP51OpABI6GlOwyNyHAPYHkVJkb1IwM0LgqwIqMrkjaeBUeQErMN2O+aTdwSetM2+YpPQipEUMMntSaQhWzxjtUZXJ571KpJPrimno2fwNJAI4yQe1Vt21ye/pU0hKqu3tVSQlf3mM+tbU1cB8UheQjtUgYAk+h60kIDEdsipkQKCD0Jpydg2Fi5Ut3NNlOQBjk04MoGBSN2xWXW4+g9T+7x3pVGRtzSttwGFA++SOlJ+YEDZRWYjkUQt5i7jwe1S4B3Z6VEEwhA4HrVJq2oEqNkgHrQZAGBY8Ux1K4b0pMCQqWHFKyYEp+Yn0FKTlBjvTVOGI7GnKAQV79qVr6IBCPlGDSZOcelOAGCD1FJk4xilqAY2cdjSSc8AdKVRuOD1HSkzkMaOoFaQFlOOop4ZmG0CgqTk/mKSLIkyRxWt7oCC5iLEbOuanA+XJFR3C7VMnrUqkMnXiht8thEMSt5rMT0qR0Hf8ACnogBYikb7mTScrsNhOVxS8bSwoPJFIGG5k79qW4xDkx8dCacnCg46mhuDt7UqkcgUX0EIRluelRsMNtPSpjjy+etQBiQe5qojYYO32zQuCCSchaASYwM8g0HAyKt76i6kEw4JUZzVLaDcKDwKtch8Z4qtGgWUSs2Wz0xW1PQTsWVYJuB/AVH1IAHyjoajaQMxPT0p6qDHszgjqa0S6sCyMYbaefWq28M+O/Q1OE2glRgU1yojbP4GsYtJgxFDKRnr3NIWBcts5PfNS5AA3Dr1pxjBwMYFDkMjROfm+7VaVF6HOT0q+VVuh4FVgy+ZxShJ3ugsZslu0Wd34VLFG/2ViRhhVjZvck9c025XMDEdeOK6lUk4pE2sU4Y8SgueT0rTtyhJUnc30qhCu0Ayn5vWtGBUjHmnvU1mUkTSqxx2xUDklTt6GrTZKFgeMVnxs7Kydeea54K4/Itp8qBmPOKhZVRSzdxQN23k9KZIQ0JJOAKtb3FLYpmII3zHrVnnA29qhdfMXcO3NXkw0K4PB7VpKXu6kohuIWaNW79zVGXMTjjJxWw65ReKzrlQZc5GTRSnrqEhu/918vzMe1PhUbhznHtVeIHczAfKKnhkxJtzx6VrONloKOxfTl+O1S9j61GpGRgVIeQTXDLcshLA9TTVOGIX8adKhKe/rUUcg2lQctnmtIq6dhD3B2k1XXJTHSrh+eDIHNUmRsbs0R00YbCMrRZ4zmmJcyB1X1qyE3xE+lRbDt4APPWtYySequPUtZYNvqVenNQRnaFDHJFSx5LHPSsZJDZJ0WgnjI70h+9k9KXJ3Z7VkxFN1OTmmKhUHJ7VPIuZODVV9xyK3i29AIy5DZNXYjhMk5FUioPLfhVmJgV255NaSj7okLIARk8Co3RZMEVOVyu2oXMcSkE/Wsovog8yVQNmPSqsojSY7+vYVZjw8RZTmq9zb/AGgbjxjvWkPdlqwYyAvI5AHyt3q+FPkYY/U1StSbd9hfg9K0G4TkVNXfTYa2EX5YwKFDKFGcigcIB60uGHPT0rJ9hkrfMc1mySus+4r8oNaBBC89TVe7Py4HX0p0nroJioTJGX6Z7UhxtXIzTEJWFQzdaewIQBRVvRgRXI3cCmCNV25HJqWTIBPc9KZGd8ZY9acdgGR/6w4HB6mpVYbs9jUOSFJ7mnIMqw6GqltcNhS+HJxwKlQgoT61EANuTT4QwfDdKmS0AsAZznpSg4yTSMS3QfhQDye+KxGIOVJqLdt4PenA7k49aik4kyau2pI9WG6myBtxcHHpRgKR2J71KVDRjJ/Sm+4EQYlox39KsKB5gOaZ5Sjn+Knj5BkDmplK4x3J600jGT3oJwc5pzgE1IELYA+tVZARIrj8atSk5wO1VrhguNwxWsL3AaI0XLYznnio5E2oCpwPpQJVDBFzjvUjSDYOOla6pisOhyY9zGpG/wBWcUp2iML2PWlACAL1461nJ3ZRUVWGWHXNWgfkweKqgkTbc4qcgMAKuWu5JNuBUYPA71MTlFIqsmCoUdBUqtj71YMZJwR70xsLGGPOaeQM0xhuXb3qUBG33M0JyuKCBggUiHaCpq7AMbLsMnGOlLgNIOeKamDux29acqZPHFVcB8pwDjp2qFZFaTgcgVYGO/Iqq4HmYUYyaIB1LIYZXJzTmPPpzUUWVAzy3epW+Zjzj2qWrMLjSBu6VC/Mu01M3D8VEygTAjrTiAi8nDDkVIfvewpFIJG0c05yOABQwGAEEj1pyrsPXJphyScU9GznJ5pvyAeGUE1WmYBOanFQyAMSG6Uo7jEwDCrCo5N2AN3TqKnfaoRR0qBiEJI6mtIsQ0smwLUY2gHd+NJjaw3HrTm+Vie1aC0J7dlYYXoKkyCx4qG3+Q49anXiVsjj1rKW41sRAHABpyk7sE8etNDMpYt0pclgD2xzTYgVfLYjrmkk4j9RTScMMGlkIEBB4NFtRmaMTysqnitCOLanP0qtZQ4jbJyxNXXz5XHNbTdtESmNjj8onNI+cg9iaeMEfSmzk7RjpWN7sZKANwyeMU1lOSewpV7Y5ofcVHbFT1GxwDDmnrgA/wAhSK+RjHNC43GpYajGwKj74xnNSOMH2qLPPFUgHK2EwahOccDgVMOU4HWmqvzHPTtVIBN2Gx/kU4DLALURU+ZkninqR5uO5pgWc7gF9B1qNzhQcU8ZIxTZAMe1ZrcCMfdI6E1KoK9T1qDpzjpUyEumT1FUwCUZGCOaYoPBp7ksMVGHJJ9qS2DqSY7UpGRweKQ9etOOAfYUgIgMs1KmcEdqXHzk460LgggUxDioFVZm3OeOlWmAK9OaqzErzj9acNx7EaNgK2KnBLOCfyqBWAAUHipjwBx+VXIQ9PvmlBO8mkPTFITggdqjcEPGSxPao8ckGn5IJ9BTByvvQhjFA3n1qXgAimH5QSODSp6t0qnqAP8Ad56U3G0D1pzAMM+lITlQfShAQyDEwwMCrCsQarzHco5xipoiHjBHbpTltqIdzg96Izzn0pTzgUAAJ9ai+gaDlPJFKCM4zSKMrml2jrSGthD3NQzAFOnNSv8Ac461C7YUt1AqogV2GWGO1IcLye9SRg7s01mCvgjrWyfQWhXZRksenrTgCsZJ6VO8SthO9DR7kKdhVc4WM5yCpJ601fmwO/ekfKtt96FOWyO1dCWhJZQ7SQalikyQOnpVaP5iWI/CrSqC2D1rKaGi0o444qvdAgAqKnRugxTZcnpzisIu0iuhkq24HGcirVvkpTVjOSzCpgCACo4rok+hCJYiyp8x4qRDgnNQsMOoJqYjIGK55FIP4eRyaVj8wpcc801+gNIfQXdyBUcjZBFK5+bmonPPFOKAjLZPFPRuopiY28gZpx4we1aPsIRjhznoKkAG3f2NMkIyMjrQzlABijdA2Okwie1Q7k2j0FTScnaehqMRhXIIyKIvQGPUMHOTlD0pXyVKihR8pBNL0YMe9JgtCN3ZY19utPHfHemAHcwPQ0+Jh5eR2pvYCaL7nzdRSqysSPypo4XPrTVYK9Z2uHQlYfLVJxsXJ71cz8vrUE6jZTg7MbuxkTZz6VZTpyOapxZIxmrif6v2pzEGSfwpVPX1oP3cCkAwefSoGP7jNHJNAxnngUfxUgHE0mcYFBPNHXikAnNNPSn9qYcYxTQXEJ49+9RghkxinPnbkcU04RcVaDqRbSWyOlB6Y9KfECD83WmOpYZHFX1JGgHJ5x6VIi5U+o71GFwcU7oMmhgKvB9qeM4zTB1yOlPB4pMZUu2B4bpWa2A5AOCe1aF2u4cCs8hWP3ea6qWxDeo7tzRuAY8nn0oU5bFLgYK46d6sNRrE46DioywHU5qRmTIJHNRHGec/gKpAEhyQR16UhIBy3P0FDdV4P4UucfL1NUNksDAqQTj+lLPjacf3TUKY3525qeVt0LZ64pW1JPTDjj6UueKTPfFO+8RmvlnqbCY7UIRnBpCcn2p2cdRQAoA39acp4w1NU4+alU5YntSuIcvytk9KXBJ9BQTu4WlUEtgnGKe+gCPnKjApgUBsjrUu0eYc0zADE9jQAwZ7jmn4+Xg8UoClsk8U5Rhvam092BEV3AcUEDeOKkGSCRSDAUZBzmpTAafYc0SH5aey4UnP5U0hSo9aHuA2Mgj3p+3uKFUdO1AzgjJ605asAx3oIGBzSbCDljxS4C9eaVrAID1HegkgjIxikU/Nx60uTiiyBjehz3pDkEN60hcbxmgnsc4oQxX+9waDksB+FHTj170dZPSmkAj4AIqBeoqaQZyOuKg28jHaqilqhErc84pzKCFPp/Kk6Lg+lHYH1FPRIZInJx2puTyaVBjn9KcRjJIrPQCqzeYAcd6dyCKJMrIABwaUg1sthInBIXg0KN2cHpTF9KeE2k81k9HqAu0MDzimHlR61IeAOaa3ABAqQIHOIyGqFIssST8pFWceYp3DmmKBxgY9a2i7aCCHGCAeKlYKuMdaZFEqNx61I4XcMA496l7jGnHBphwTgjpULsyXGP4akdiDuUU+WyDQXcC5XFSEYApsY5BPpTsHr70TsgGY596cwyBigYywHGaU8MBmo2GJjtTVz0IJxSk4I7+9KvGT2prTUB75YA46VFL8sZYHmpOQu7tUMy5dR2PWmtXcQltkxEMOanUrs5qFThgM9eKkRTuI9KqervYAYkDaelJHlQR1olAJyKQnC+uai+gxWUlBzjNSKBleMU1emDQGDfUUXurAIciQU4H5+nWgDCn1oXk7vSpQhOhJpQmEBNH3grY9qGHGQelUwEbpnviozh+tPkYBgcdqZ+hosMI1LPyelPRwrEEdqB8rDmkIBO0UX1AXB2k0vWLHvRjLY7UmSFNCDqLnhfUU4DdJzTEJZQcU/wDhBHak10AMlsjuKX70fuKTlQWB60DIHsaG2hEb4KAGoXHIGOO9WD8w+lREZJz0qoyYDJEG/wBiBSjj5RSsBuxnikU85NW25bjSJNv7o4qAKyHGeTViNvmK44NRTDBDds0k3cBIiwbaw61Oo4YCqj7jOCDxVrO1gc9e9OS0uG4KcHNNXktnqaCP3mD0NHIc+1Z7MBg+4Q3XtUJXBK/jVllDcVDtO7dWkXrdiGbSZV28VYj5Bz1qLpKeeoqSIc5qptNDFEXBI7UD7pHcVIGJb+dRSY8wbfxrLcWomSMr2xUq9CBUTH5vapVG1QevrTkrq4xOiEEd+tIx+UY5pzZC/Wkxg7am4C/8slzUa9M9ealUERnNQYIbJ6VSswHryGNPx8w200MNhApQcq3rS2YBz5lBypyeh7UDoD3pTyRQxCYIGD3pWyB6UDlsGjOSQenpSeoyMDuelC4ZTTdxC4PGabCxVTVW0uAsiCSHYRQiARYz0pz8pn25oXmMHoabbsIQHarZpoOY/rTuqnHemDJOOlCYwEgJ4OQKMAMTTCgDHA4NEmVJwabS6ASblyCe9IB3pqlSQp9KcpHQ/rRawDmOT7+lVEJFwRjg1aJ2nINRE5cmnB2YuoMmDkcZprL8ufTvUxz92mAqCyk8U7tsLFeRC6Ar1qtcHaCW+9jpVpz12jiqNwd0xVzj0ropq+4gtW5CsuW/lUxKq7HselNt3O1vbqaZKQu3A49c1s03MCxHI7DPYjmpIoxErA8k0qYeJcU8AnAXrXNOWtig2fu+efSmCRi+3HyjvUuOAM04hRgL071F7iImzj5e9RPFjp171PKSpwOfSo5yUXcB1oi3cZX8vHzevaoTC0gcM1WVGz3zUbEq7AdD3roi2ToVLeNp2DZ+UdBV5R5qdPu9qpxK6Sqg+4O9XlcqzoBk1dV3s0EdiyR+5GO1UdwCs449a0I8Abcc1QnYQyFcZBrmjdysMmiXdFleMVWeNvLYDrVhQRHgHg0MpEBPc1fNZh0KO8iIgfdWr1oF+zBhVB/3cXJz7VPYyhoWUde1ayXNFk9S6fnQc9Kz7jKkFRVwthcDrVSYhFK96yhuORESIwU6lu1CECQAJTQwQ/MPmI4NTopRlAPXvXQ3ZaiLka4XcakTlfbrUJJKhe1SKuF46VySepWwroHX2NRxqseVXuetT54FVWLeYeMYpR6oCYNiMgCs66zFIAOg61eY5TjiqMimYN6jvWtNWd2IfCxdRgYFEYZVZVohPlDYT1qQ/K20D5j6Vc2BVieR5Nx4x2q+mSPfvVWWBlkBU1Yhb5QDSm01dDSLR5ApgIJPoKcW28Y5qINknFc6AU881DIBk+hqXPyE4phYKvTmqTAqyqQcjpUkJ/jNKwJAB701DuO3HFaqTsIlDZ3VTZRu2HmrioBHVVx94qvIpwsmBJGCnIHPelldkXFIHAXB+8etK545HWm99RlOSKR5BKo4rUWXfEuRz3NRhS2FB6VJGMk9gKmcubQFoxcFXz2qXpgnrUeS5wKUEk4z0rEB/JOKr3pIIC9TVluuar3Ss+3acAdacLcyEytGP3mxjyTVp8qQB+NVQR5+AOR1NWpH+f1Naz3H0IpVJGSeewqFBiIgjAqxIQzccnFQk5Y+1EW9gsRKfWpIj83t3pqtl+V4HShiV+7/ABVdrh0J3yQm3pRg78E0i5RB6mnqMjcetRsBKw2gDv8AWmlcH2pwG4gnt1of5vwrLqAwtlsLUMxLOBjkUpz5gP5UjNtkPrVpdgHJ8zKOtTn2HSqwG18A/jU2SOR0pT3AecOcilYE9PxpqtggAc96exwT71GoxvJzTxnG4nmmbcMD69acRgYNIRFK2MnFVHKMvPOatzELhR19KpSROOAM1tTsJixogBZRjNKEBUkjpTt2AO2KacggD8qq7Y9CZPmQED8akOWIGOlMjGwYqUjCHHU1EtwKrqPOyKft3OORxSStsXfUQmVjuHBq7NrQCWIghwDkipEYSrkdqbGuzJ6A1Gs6+YVAqbXvYRcX5unWm9zg8Ukbj+HqKcRtODWfUZGMnJHQU0qFU45zSsSGwOlOXK8dR61SYEYUbgB+NOYBRildSkgI796a3Jb0o3DYemFIBOaozShMuB0OM1biYNnjpUFzHuJA7mrhpLURJExIB7etTMeM+tNVQsQGORTmB49utS2mxjhyB61C3Q4FSFtzcUwfK+PXrSQCISCR0pw9c1GrYds09WLc9KbQEfQl8Zp6/dJPWkxlivY06McnPSm9gFDggD3qOU4Yt2p+35+OlQyF95AHy0LVhcQfOhLdulRudy/dyO1SlQY2UGmR7hDj0q0+oFYncgPSnEc89cUnLq2045p4A2jHJrRiJEwTznFTOpBGKixkg96sOOhrKW4yCVsYBHWmM5CAYpbhtqhR1pin5Fz96qS0AeAMj2pZR5hwOppnIAxTi209aOoIrIXVsIOh5qwXYnYOh6+1RkOFO3jJoaXaenTrWj1EtCccDGeaJlDRYquZyZMkcVZkPACis2mmhsehwoocZyx5HamRZPBPTrTmYMQuOPSoe4EikcUDlsYoycY9KXpjFQA1uFz61CQRyancYj9zUJ5Qe3eqiADkClJBIFIDuIXNITh8DpVdQEfnikLYw2OlEnHTn3pcjJBAqugEsbCQAj8aJDuOMUkYCJwOtObgZ71n1BELDJK4qZRtXAqIn5j71JHwpHWnLYAPIAHUVGR8xAqUDnJFMAIckjrQgHZDMM9u1KDjNN27SD3pwGDmk7AIAMc/hTUBDHHalJz0o/iIH40xDj9w8c1TmYnGenpV0Z2kZ6VVYAnOM04bjIIQCT7dKsD5sVFGm1ifWphgkYH1q5O7EPII4xz3phOBjHNPbKn3phyRnvUIY5M7femZw3P4Cnk4+b8qjzyCaaACAuSTkmlHI5ppx19e1PX5iQegpsBWPygY4ppPGMU49+aYCTn9aSAhki3jFTQDEePSoyxDgdqli4Bx0q3sKw/qcmjbk8dKPQ0AjOAazHoPXnindcjHSmqcE05SNpOalhuRtgLg/hVEyMW2+lW5mO0HHNZ7OWOelbU0It7v3fPaoSVZAwHNCybh7UwsfMx2NWogOLhWDHpThINuc9artIpJBPA9qZEp9eO1VyXWoXG3qgsCDgGq0WN4wcVfuYmaOqCKU69q3ptctiZblhfl4zUyyc471UA5yOtSIcYPrSaV9Q6GlG2RinkcE1WhO5gatnnAHWuWSsyyLZuBOPxp4jAQe1LngYpxHA/pU3YDDHgjjpSE4IpXfDZNMQ7vmzTVwQ8d6Q5IHfNKOATmkfp06UuoEb/cwetV2LLJk9CKtNz1qpOCV3DotawExv8AEc9qeudtRgng54NSNliR2q2GqFYkkjvSnJA4+btQcl8jpSJkhnNSDEYsZQO1SOwzjFRx/MofsOoqRjwDihivdA2BtPXNOkAyCRwKYCC5TtTmbcpXuKQyDa/m5J+U1YVAGxUbFgR6etWMYAJokw0bK80oQ+3pSxYcBm61XmUAkVZgG6PmqaSjcCZfmyD2qORd2QwqZeB702dsJkVknqMqhQpOTirCkKgwcms3zcvtOeavW/I5Oa0nGy1JvroTA/KMHFI3XJPHrS+wFMcc7Tzis0PyJQcr7mnc8DvTegB9KcMnnNSMQDk/0pxyo5NIM556U4ikA08gCmkce1OPPNNY5HWmgGuflwahkBbGO1Snnk9ulIScVa0FuR4G7PakfkD0pTjAHrTAedpqkMU5xkDpTQfl9qcxNNzheelNC2FHC4qRAAMVEOnPBzT1OTmkxEUwwSTWVISj8dh0rVnIIPrWTMx3n5etdFETHJnaABwelKeCcdKF+vBpQdxOFPNak7iMo79KjOSM4HNSA4z6VCwPIximh7oH9jg+tNK8Z4GR1pzcAYHuKQgfp0qkAsfDcn2FSyDEL45+Wod2D6HOOaeSTC59jQDvzHqOcCkbjnvSkfMM9qOvNfLGo1eMcUqepHSkGc5p33cCnuLYerZBBFKMKopoxtzSg8VIDx6gYFKBkEk8012Ix3FOwPlweSKrcBANzDt70MozilH3D6561E7srAqCaElsIkx8wFGG69OaIiXOSOaVz831qWrIdxwboCKaCQMY5pTkgHPTtSMScc5odugCHjrTZRxgU8EYwaTOe3NJIBp+VaVTnmlx603GRjtTWu4Duc9e1GB0PWgAEDFGCDik9wQzBDEd6OSuO9PXcW5ppHJNAED/AHuuMU8n5RSHGcnvS9AP0q3ohin5mAoYZPvTTweKUEH60gF4P1HU1GcBR6mpD3JHFMbBf2AoiAuMdaUdM0EZQ+uKVFzGR6VXmARkZqQ5YVGgC5461Jjt2NRJANkAIz6VDnmpXVguT0qEdD7VcL2aAeh+YNUx61W3EDOKsZyoY8UpRsCFBwMHrSHhc9aMgn+dIy/KcVO4ADg5psg4JHU07qmB1o3AxEY5Bp2sxCKBtz0Ip20spJ69qafvcHink5cY6UNajICvqKhDhW21dlXBOKyplkaTKjr3q6S5nZiuaO7O3b9KRjtXFJEpRACOgp+QfpmlNWkMYVyAelOx1Gacwx9KAOCaT7ARsvB9qRTztHSnMPlznrTUGzOeaaXcB4JIwajdWd1C9M1KpPX04owd2RU7agiHZgj1FS8ryKRid/PWhWPQ07t6gNmkCKFNIh3Y9DRKoyCaegUccCmtEAoGGJ5xSKu0MR3oUgE88Uikl8ZyKWgug8cpikUAEkU4kbuKQAAGp2GC8Zz+NLjaueooBwhPY0gBKZPY0LUCJyCAKTG8DFJIGB+lOBBxjtVbIRKq/MAaaBgnmhSTyacQGYYpDEA+UjnIoJwCuOTQn3z6CgH580AKCACKX+CkA+bnvSjBUiloAoB2nn6U052j1HanY4+lIflcFTQthMaAST60xfm3VIONzUxTt3HtTWoxmPlbPU1Cp7dhVjGUJFNVQACepq07IB0Qzkk84pJAWiI9aF+RsCnNwxxRfW4uhUlwUAzzVlRlUJqEpkZx05qZBhAapuy0AV8sc+lNJIbPc1J0BqPBOM9qzuMcCN2R6VHj73tSpnaT37UuDtY4561fKBDnOSeoqSJgKZn5R7nkU5BsUDvVva4Eq8E0meN3bsKcD8h9aaSfLwelYgRnuT0qRfmjWmsuAMdaFP7wA9KppNASS52Dj8aiDknkc1Ln5iOtIoBelsAFs49O9Ml+97U89enHrRIqlPcjFNAQM+MMDwKmUgqXH3ajVM/KRwOlPC7I2X+VNtNCDPGaeThQe1QRksOTTixIAHenysZIuMk5pc4OTUYYK208YqRQP4u9Q4tAQsDkZ6ZqJSTIQKlLBd3r2qFBgN69aqOwidsHrSKwK8flTVJMZOOtLF90nFDVkMcwKrimnk8dKevzde1Rqp2GkkAj/McjpTWUYFKFK4549Kbk4wRj8apIBCNhLDrimRvvbnj1p7kZx2qRFXJ+lVewASD8uO9R7QHA7VJxt+lNHUHbSiDHH71NKKT8y81ICC+T0oH+s5piKkoKKdvGelZ9wp37pccd61JE8w7D3NY90zLdFT9wV04e7ZLJbdsEAnhu1PZQgcH7tVA6tIFT7vrWhKu1Qp6VtUVvUNxYJlZAFHXpUxdUyCeaqKAnT8Kju9wlBUcY5NYOmpS0Kb0L3nq0gUHipd2wk9qxbaTNwfQVcluwEb0HenKha1hX0uy6GDZJH0qOXlMdfSqVvcu8gAGBV6Qbo2OMe1RKm4PUe6K5LKTkcdKVFUoT1I6UgkLtgjpToeWI7HtVtvUClA7x7mkXn61ahYAGXPFOKKzFWqRNoHlgcU5Sv0ElYDO4fAHXvUdwjvMOOB1NSh/LbAGfepSQU3fxGsr8rukMht/MYlHHI6U9wD8vQCnoCW3enalYAnP50nK7Cxi3Jy+0c0/TnAc7vwzTbhcb2J5/nTIEV3A6Emu6KThYi+pp7euDz61VmwwwOoq22VOKqSgo5IGRXLHSRTK0zNkqT8wxgVfiONu8Y4qi4aVxIBgjpVlGLYJbgfpW9RNxEty6h254+lOByKiz0AOakjyeT2rkaRQrHOB0xVeWQF92eB1FSMxUnI4qpIQJgG+7ThHUHsWQ4ChUHJqKUeW2Txn0p6sNxPrSXCl4xzyDVpWdwKzsowD1J/Kp/MVSB1NVhGA4Jp6ASS7vzrSUU+oExZnc46dxQhBY47UI/wAzACkIyMA4JrP1AnjOOpye9IrfPgdO9NhQKN27JNKRtJAHJqNLgSjGcCoWwXI9KMbQR3pjA7femogMZuD3NPRRxVWfzF/GmRySJjLYNbezbWguppO6KuO1ZsswBKxnii5nfGzvUCWxGeauFJR1kDbJIC3mcnmrc4J2nPSok/d4J6UssZyMHg0pau6H0HW0+ZcE1c3BuFHNUFQxjcODV1d25Xzx6VnVS3QD14HPBoRw4JFKy8A+9Io2DNY6APU7sqeabIG3j0qWIgc1HKwUYJ61KeoPYryfJLggc96c/A+7xRKvKkjJ7UM/ygVrukGhGqhCMH8KWQhQMDk05hmNTjmoZM5X3qlqAbiFyOtIFKxlm6jvSlS0m0HAp0i7gIyeKvQAVg8OeoqTI2A+lRwgB9npUjDBx6VL1Ak/hBBzT2wU9qh3YyM1IvCjvmspRsCK8pCmnkAgN2qGTmU5HFS7s/IOcdatqyQDGBDqD1NWBjOBUZG4gg5xUqEgk5x7VMgGB8tk9qkxk59Kh25QnqamXlB79ql7DGMx4ftTwcvkdDUc52KuPWmrNuIKiny3WgWJWxuyapyTbCxA4q3IN0ZNVBEGTJ79quCS3ENgOQS2CT29KciCT5weBTRGMMBx7UkMuxAuOtW9dYgWsDaOefWnMdpHemIwYKRQG3Env2rNphcjmOCSV+Wqx2zKm1asMHZcNzzVZAwfC/w962hsBbOViKDmo0h2tuI+Y96kkkwA4HWlQk9e/as9UgHRKQxOepqZ/v0yPgHtTjwPc1m3cCMn5j7UpIA96U4ByaYeWz1FMBsmWcKKQ/IDnmnnAAPekPJ/pTTDqKg74oZcHFOhBxz0FDv8xOOKXUB2CSTjpSPkqT370m8jtwacuefSkBEhwMnnFIn7z5yMU8D5iRTN23r0NXcAU73wBxQOFHOcUxNqt1z704deuAKGANnORTwdpxTM5xgc1IMFyKTAGX5TjimDO33NOJ3HHakJH3R0oQWI0AAJzyaa2VVj6jilyN20UkgLA+gq+oFeLq46GhcAc1HCwJY+p4qeVdqj3rZ2F0HRj5zViMfISTUK/cx3FTIPl255rGY2U7n539MdaSLOBgcVJdYWTg1FHkHA61pHWIupMeVxShcjJ5qDdtbrxVhAWyualphcYRuwxP0FJOAxMYHXrT8g5z0FRly/KDrTVwHeUiooBqRiViz3FRoF3dee9PwHB5pMYyFtoJJ61KeDnFVUcLIdw/CrbEbFPaiSsxKyHquU60vTBPNNVh2604AHjtWTGIeT7VGBkNUudwOOoqP+EknmmgYiAD5qYB87UBuflGAKCu0luuKodxjELHwKVBv2nNI5HXoKevyk46VXQRMDkU1j14py4VaOo5rMLEJJbjPSpIcoCvrTR/rOacOHFU9rAPPFNA+YH0pxIJx+NMP3mOeBUoBQMkA0ZwKABgHNKTnjHSgBrcLkUn8I560r42dOlIcnacU0LQcwPlVAgJGKs9TjtUDYUEDrTi3sNDdrK2fWnYOeDxRj5SSaAAAPemApGW6nFDkA4FHU80HuetINxq84BppwT0+lKp4pDgD3qgGN1Abt6VIOAMHk1GSuT1zUitxx1psBzcDGaacKMDrQclv50ZyxGKSQEMhVWHr2p8WRn0NNbDOcjinD5lG08VT2ETdgO1KijmkXgZ70qdcZ5rMY8ck85pOdp9KBw2KUtjIxSAikI8tiay1Xcp9+9acw+Uj1rOPBY9AK3pbEsUAeWfmyRTQG2k9acQoXK96bE/7skjFa9B7EZGG+tTRgbsDtUKLvl3E8VfjjXG4fnSm7IQMoIx2qhdx7OR0PWtAHFV7hBKRUU5WY3axnp8yMR+VSRjcSPSnJH5aso/KpbaPbHubnmurmViUnZIsQx4U4/OpgNvXrTI3AQD86m6sea45u7KSGqcZB6mhCdvNRN8pJNOU5GaVgEmPBz0pAMIADxT5AGGCKiDHft6VS2HuTDpn1ofHfoaFxtxilIyCPSo6gRMcHOKjcZGB0p5GVPtTGOWFaIXmQbeeODTuVXHvUnG6mkc/N+FXcByL8hyaVI8Rk0Jksc9DQ2VyM8dqnULkYGI9qml3kptxz6UxmCv8AypokPnE4q7CuTqQBnHPpUgwctg1CCdwbH4VKpAVuwNQxiLyAae7bRUbkY+U8U4qGjz3pW6gU2G+QsOavwriPp15qkGPmcrwa0VAVcVVR6WEhBwOaUqCCCaMd+9AyRWJRnyWq+ZxwafApQnnpUkzgGooWy2T3re7cdSbalrcCRikc+nNIBjOeaVyduay6lbD1O5c4FOHSgdM0Y4xUsQHPFOJxyKT+KkJoGKMD8aYw9KceRTSNwBzQhEa5BweTTW+U4J49akzk+9RTLvB5rRbgRbsSY9KVCGZs0xGByoOaemAp9TVtWDYUjjrxTM5U+tPON3Wm8bsjvQgEHI6c08Ec005G7tmhWGOKLCGzdM9x3rLl+/u9RWpJ8wIrLmyGPOK2oiY4cLj8uaM4Ddx6UkeCvPbqKOxJOR6YrXS4vUD8zECmSA9MHPtTxnOOn0pknqTzjtTW42Ax19DSEZJwBgc4oBAP6UEEn5fzpk37jCc8kfjUi48mQHB+U80xhjOB+VOQfupCRztPP4U2NnqOfmpCeoFCA7SxFOUZycfWvlrGgwHg5pQcKCaUL1OKUKCKaBighunelHOAOlR7wnJ6elSoQVzSa6gPOCg/WkGMdeaM4Xtn0pRgYPX6UPuAoJ24FIDgUZJJC0Z28YotZAC8n5fvGhgMZzzSAlcGn4Bzzk/zpW0AD2zxTRnPFK2COvNICRxUsBFxkk807oxOKAcKVwKXaVwab02AQDJOetB9unejOM+poHOM0ABHygDpmkJYGlxz1pCcjntRsAYO7FNOdvtT+M5A5pvLHFCAYyjikP3N3oKd1yD0FNUDJz0p6gIASu6lj6NxQvyEihflHsaobFA3cHpUZx5nWpAcE80xuSOKSEL1enRkgmmryGqSNcEkUxjfUAU8n5fQ0gO0545pSccVDuAHlcZzUDjmphwvB5qNjhh7U4t3AgGQcHtVleY8GoQB8x71PGd0fPUCrntcQA4BH60o4XFIB8ppOpx6VG4xGyoyKYZFU4PGetSkBlIPeoZlxHkCnG2zEJ1fg8VOMCMHv3qGMZXI4anBsZz0q2ruwE2RtI74pu0fL6dxRGdxzTlwAwNZrcBCQGyD+FR/w+1O8vd9aEyQU9Ke4CsCevFN53ZHQ088qM9qQ8IPSlsMYw3U2MnJz+tSuPl44NM6EZ6+1UthD1wEI5zTc45zwKc3B4pGAwAPxpeYxpznOOtJu5+tSSYG0DvVcZV/UVSXMIbM/wAh70IxJ60rqWJPZutL5Wz8arS1hgCd3tRDu3lz37U0Ah8npUqcrz1z1pPRASEZXg0dcKKFGD7UDljWYAOAy+1KOEPr2pqHLe1OOA2KdwGsAF5FRbSuSDwT0qY4KkEZzSYGwqfSnEBEP7sjnNPXqAelMHCgdqccBc5pMBqjDZPSnEcmmZ4APSpAcD1Jo2AQkHGOtKRyMde9NHOcdac2OPXpUiFLDjI60xuAxp5OQM9BUbHIqlYBR901GT8oPWnMT07GmZJBXFPQBWbGBjikfOQBSM23aO1Pxlhjp3qmMYWJXPepjnAb1qED5zUzH92KTdwGFcttpwBG5T6UoxvGOuKVRuJJ7Um2A3qu7v6Uw/dOTUmMsPSmH7rUXuAmACP1pJTgnFPP+rHHNI65I9ad9QID97AowTKMnpUcoKSM5P4U6Ft67iOa1t7ugFl88EUNyMYpeMqDQvBJI4rHcAZcAHt6UxhwcU4DDdOPSjqwp6ACDKLu+93oP3uKTkNu7CnHHJ/Kh6iE5Y7elKQW59KAQT+FKCQp7g0DGjlgc4pz/NjHfrTQPlI7U6Mhs57dKQiNkwAQOtRqwZmA6g1ZHzZqvKohc471pHXQBkj7QeOalQloqhlzIeBg9qVWI+X0qnHQNhXGeagdygUgdTzU7oSAKgmXohPJpQ3AsJzGc0i5GTnikiVlU7qcn3QKTWoxy8ZPrTQAA2DSnDr6YpuQCFpagMLlmAA+ppxBbB7elI3HTvTiuIwQaYiNRukyelOHBYYpSRigDKY7k073GOUgLjvSN8y5HWkJwSKV/lUDFTbUBP4ABTuCh9aSNdvfvTsgFsU9gIZAdu7oaybxN27aOTya1yC27I6Cs0l3mK5xtrpoXuTLUpvCIijJ161eaQ7BvHXvVUykzEAAgcAVZfLFD14roqa7ij5DkQZIqvdNlgDU0bHftHeo7zCxH1H61nT+MqWiK9vGDNgHj1qSdP3oXHA70tunK4GfWpZ4ysgduncVo5LmQtiCP93J8g4B61rOR5We57VmKwN1tA49K0QQy+3asa2yGiByqjdjk01WIIx3oI+fqeO1MlYhwB0HahagR3M5jdSTRbzl+2D2qOZizkdTRAn7wHtmteT3b2FfUvBiHVc9afK2yVABxSRqZHOfujoamCfNk8+lczsmrlCglXx60MOGJNNBA5xmnSZ28dfSosIyrlSze1MtU/fo36VYuUKktjp1qtbg+fnP5120neNiXuabsCSpHWonXPyjv3pHJz04p4GTkdBXMlbUopCNlGAacirgAdcc1IwUng80eXtHynJ9K2crrUmw5SEPXp0qeNj3PWqbMqPhx+FWIHBzzxWco6FIkkJY5XtVSZA7DJwKub+eKhK881MHZieo1VEeDnjtTpHKgkDinhQVUHrTygzz09qObXUeqRkzS7WwTUkDLIzc4GOlNvECHAHXvUKg+VlT9a64RTjcV9TRjC5HPzUr5U5qrbAswPUirsg8zC9KxqR5ZWH0HRDIBFOdDgeg60+MBfbFIjhyyg81zttu6DYQKAnNV92ZGAHFWGGTgVEwKg88d6cWBXuNxbPSqfU5/hFXpFwlVNh3Hjmuqk7CZEWLvzzipyMR4zSmMbM4zn0oddpUYyKvn6CCPbINhJBpzqxfaTwO9Mdv3qsg61Ll2Yris3fcq6I3ydoB69avQPmMAjnvVdVKY4zVlEyoPeoqSTQvMlAymD2qN1OB7VIxxgA018KPWsOox6HHTv1qKdQyHNKpwd3anr/Ee1NaO4GZHLIG74FSJJvyo9cmp32oGLLx296hChCxHet4yTBE3mBl44AqMkScgdKRMJuzznoKcBlD7+lT5gNDBTvHenllZyOlVn3KwCingFnBqrdQJl+Tc3cnpT1+bk1GAWfOacODz26VmwBRng8EdKWEtkqx+lIoILHNPXgA5603sCGIu9yc9KFAQs2OakHCkjikwMGpuA1DjLAU4EseO9M5WPGKVSAo9SKbWgAmFZgDUw6ZPWoADuBA5NTH7uAeRUyGMkUEcim7QsgNSZ4560oTaBkcmi9gFzuUjoTUQXYv41K5OBjrTBjYwNCbEQNhWLVVbOCy8AGrpTIOOCaqybg4AHy9DWsGALKcKccDrTizCUFRxTotpcKBUhJ3ADpTk9Q1sDFjyO3WoioycdT1pzkKSQc55qOFt+W7miMWBNtLRqM8CnKuV9qI1PPpmn/dUjvWbb2ARSTIAakbnGOveo1Hzcdakj4Y4PNSwEbnOe1MCkPn1pxHOcU0HDHJ49aaAYeTSx5PJ5pcfK2KRMk89qYEgGBjv6UyQfKOehqRSQ3PSo2bJOelJbgMkkYr8ozUsRJi5HNV1OHJI69qtZGBxTlorARscKTTHIMX0p0r4j9qiEnyEAZFCQEaPu75Hapc/MABwe9RRqPNB24FSgjfgirkK48H26UKfUUgP60pPQCpaGOX5cmkIyQ2eKcThTx0pgOVGOtIBoH7wmkZcKdvU1IVz0qNsL1PTvTTAqLbssofPAqWU/u1OaNw+7/D2pWw0YArbVvUWwRMd4zVmM/Nk1TTII4G6rUR4561E0Vcr6jkY2iqUe98kHkfrWvMFKYbvWV9nEb/ACGrpSXLYh7khYFMH8TU8b8ADj0qLYGQ8U+LsPSnKwyYpkMPWmLtjjY4wKkY4cnHGKrkl2JP3ahJselxsEfmKXY9elWVAAPPNU0L/wAI4FTxZOf1NVJdRdCIR+YzlevpUwLGIKeo7U2IETt6dqnCDefU0pSAcmNowfxp4PzYpqkDk9KOvzGsmMd0UepqKQEdKmHUt2qFxwcULcBvGBjvSNhzsB6UuQoGOtNYY+YdSetWAnHK44pFPIboKkBAJzwTUTA4wOmaYMnVwWJ7Up6EmogCH2/rU5Hy4xUPQCBix5XipFOefWq7yBWKk81IpwMYqmtAJ0HB/Sm/wkHvTz2A61FKBszULcBVOR9KccOeBxTE5jwOvrTlIxTYA47UgPp1pWPFJwMZ60IB6855qq0q7jjpVn+A1Qk+aRgOxqoK7BlpSWAzxS/xkdqRBlVNDEbsZ6UdRCtzg0McEEUjHAxS4+UUh3EUdqQjjJpw5FNfAT1piGe/60JjB9aASDSKDvI7VQx/fPQUD7vpQ2CcDvSDqRSAikJJwBTxtXC9fWlyQDkU1OpFV0AsY5FGSCRjmkGMDNOBw/NZ2AXow46Ugxu5pCcsfam5/M0WAinbJNUsZJFXJACcZ61CsJEhJraDSQrNkLIVAOelRM4ETKehq1Ku1VFU3GJB6VtF3EWIMjk88VcjA2cjmqcLdB+lXzwKyqtXGiMgHJ9KiQBl4696lchQeKjjwqgd6lbB1IGj3SA4471YUZXaBxRjnHrSq23mqcm1YLDNyxuwPepoyCu7rVGRg7sSKsWxwoFEo6CWrFlX5OaZExDY7VYP3ScVXfIGR+NKLurFMlzu98U3cA54+lMRiASPwpZSTHnvRbWwdCUfLzinDG0kd6jUlgPang9cdKloCMkbcDk0wgKCKkI+Y/pTOp96pAMz8+RQemabyZPankZ+UVYhik9RTXcnGT0oflQB1qJzzjuKpIAdvMIAFORSW4HSoouwqzGrctjAqpaIQ5SV3HsKN4ZSenpT8ZI9KjChXKjpWegahDgoQ3envkHA6U1OdwxgdqfHyMGk97jWwqqGAyM+lTA4XHam42qAOaUVDYCk5p4446U0YzS96kZRufllIB60yPPm+xqa9C8MQOO/pTLQZ+Yjmt0/cuT1LYA2+9RuQuBnipB93mmuuayW5Q8ZyB2p2QQaavC9adngnrUsAPOOaM8D+dC8LSHgdM0BuMLfN06UmcigndSA84qrCuGPm56UxhyakwM570xuT0600BTC5kPp61NGPl9/rRgEf54pQAmMdq1buCIyfmP60AfLjv2p5wwPFNOBwOlAahghTSYGRinbs8d6TbhvY0INBOe9Z94Bk/5zWkwx1qlege/4VdN+8JlVcjbwf8aeoKggjA7mmIcZ9egp5XPJ4wa6BMUj5RgdRUbnrxwOxqQYIx6VGcrkZJ9aSFsiPB2hu30pQCepx9O1OHDHHftSZDHA/OrbuJXGkjHekT7r8/wH+VOIGcY7dPSmbgquOny4FNK5S3PVyMDHrSqflIpF569BTR96vlTQcAD19KN2V46imk46daRec+gosBIFG0560qdCP502Jl3FWIqQfT6UPQQhUbc5pcAID39qFUHOaFOB/Wl5ACnApNpK7u1AHBJ4px7DPHej1ADg4NP+8OnFMAwM5p2fSh36AIBjBx1NIQSwz0p6jdzTc5fB6UrAGOcA8UMx3AZo6E/SlTuTR1ACCRn0oY55xg+lKMkZ7CgHc3A5xTAbg4z3FLzt56mkP3iKHGMUt1qAnQY70zlTwaWQ7Sp9aDzlvSjYBCcAjuTTWBC4PrR97mhyWIx261a2GBGVx3owcgHoKF+9n8qMBiR3xQlqAYzKMU5jlt2OlIDz0ob5GA7GpTAFwFOe9KuR9KY5xjA6VKpGeehq1roIjkYKRt5pA3QtUMj5lCjoDVhADhaHHQN0K/UelQHcH5PBq2wXZ79qgdMuKUXyy0AYMZbFSx/cOO1QSOI8dqkjf06VUrtDJQvORzntSZ+bJNK2VOAajl4/Hms0mA5fmpuDsYelEeNoIPXtT9gV+vWnaz1FuRDIXPpUoAPzd+4oRRuINAwGOelNt7AOQKN304pv8OKUYB60pwQR6VLdtBjGmKsvH0pA4Ynjk06QYZQelQCQAkNwR71S1QtSyMFfek6ttHTFHACsO4ozyGB61NmMCMknPSmgZJYUr9GXvUe8KAvrQkCHk8j170rbTjafrUe4Hg9cUpGNpzVeTAazDcTUaA7gOxqbYATx1pi5V+evamnbYQpXHB7VIwBAFR7vmLVISWJIpO9hiSAHGKG4VcDnvTnwpA9RTZCAq0baCJB/qsd/SmA4DUsmCwI7im8ZOOvpSaGKV2pnoTTunJ7ikc7hkCkbG0H2pWEDcAimnld1ObBXdmmZ4BpoY48quaVl+UEUhxwe1OGdntTWwEYxj3qQEHB70zblc56Uqj91+NIAGOTTgRnnvSDheRxmjALYFK3UBQcKQabwxz0AFDEEAAUdMLRsxARuDEDtUbHC8daeTt4zUbj5TzzVIYjDIFPiPzYPSkUZXHelC7RkdaH2AQnY2afnD47Ux14PaiMleWp28wQ8fMwK9DSpgsQePSkUY4HQ0qgK/PSkxCn7hXuKYBkFT1p2ME/pSLkSmlcBAMqAO3FBOePShMFyD0oA689KBlcjzH2kcGnYC5XFOXBfJNBU7ua02BaDkOUBPXtS53KMdBTRzgClAwhHeoYdR5IKk55FIv3Ce/amjhR60oOBjFKwCc7ce9B4QDPGaRTlsZ607HHzVTAUEbMdxSqcKf0pu3DfWnZy+AMVNwGk4IU8Zp0ZIbI7etIP9Zg+lCZ349etVpZASxnLnnANMdNxIPalUZY4pUJO7HXvQhFeRQSGH4ioAD55OPwq2VBJHeoiMybgO1XGQCB/MGahlXc2T1FS7QvQ0xzhgAOvWne7ugHAtkDNBYKzJSNzICKjnOAjA8k4NCV2NkinCZzQc7d3c0i/cx0ozhRmlqFxz/MigmmqcnpxSnoO9Nz82aPIZJtDOD2pP4/pRyBkGk5UA4pWEBA3bicUp5Gaaw3KKA2eg4FPcBUJLZ7U/gCmqdwP6Uq8j6daTAYWwpx1NZ8sJaQkcFu9XpXLAhevpUDLzhgQTW9JtO5LSZlMmHz3q6MiLntSPGAxXv704qwjHHSumcuaILQQELjHBqtfOOMdKWa42JuxwKptIXiJbqKUKbT5gb0LcKlXGDVq9IEQBOCarRSdCvUVYvTut+R+NXNLmTDoV7UL527rWgzALx+VY9vKISQelasIEiFs8VnVjfUaGMQCWxUROW3Hr2FSSnnA/KqySMScfrUw2uK4mMknuKWGY+d2xUxjBwB1PaqwgdZGAHA6Gt4u902LVGnG+MjHGKmjHy5qm2dqnPNWiQIwO/rXHURY37vFO/hJxzRjigdcHoOlSpWApzA5YkZzWcWb7R8p71qXHHPrVGQCMEg5JrqpSsrmbLqsJAuOo608/KTjpVOM+XDu6k1YibfH81ROLTKWowr+85OCe1MYlDvPOKZcsRKCBmmkk5BX6VpGLsmIbOwlcM4wO1LEzAnDcVIiCUDceB0ojgG9ufpTclZoOpJHITFweakxvOT1pDAiqGxSxqp+bHSsW1uhkqKqgHNOYdcd6BjGQKG+WMn8qyvqOyKF0oLhT93+dUsIVx3HNTvI00pz2qsFYHAHNd1NNKxJbtZgNqgVoOuBx171QtFEcnz9zwK0H+UM3asKz10Gth8eBnP40FRkleMd6ihbINTdOB0rFpp3GMJOCSMetQOxYketSShjxUMjeVjnJqooGOUnYQRUBLAEKM496spJiInjNQhWDEjv0q0rN3EMt1KjaTgHtTJJMbvSpUHzAMPxqOdSzDHStLq6F0Io2I53celWEkLMABx61WVRu5H1q3CmDjue9OeiuUrj4wDnjmlR9gLMetNWQI+COaVAAh3dT0rACbOTkConceYRn8KUFgQB0oMWZdwGaSsgY8kcLmnqCqEHvUTMep60OTIvy9aGhjyUKlTyBVNZA7uM9OlWSp8kjP1qqYwrcDr3q6dkLqSRrwWb8qWM/KeOlCgcCmSs2QFH1o3YDPNOOmKdGxDHHX1phYLjNPAAJwcVdulgJ4yfofWkYbuDSKu1c+tK+RjI5rProAiHBKN0p4Cj5aZF9/mpUwSc9aUmA7AwBSNw1K/WhuBxzWYED58z1FOI7mkBJc4oJ3ACtAHAkPtHNSc5PHNMUFZMg8e9SA8A1DGCuGOMYpf4aQAZ3etKBhevNLQQh+4Cfxpoxj604nK81GpJ4poAZQG3dsVnyuI3HGQavsQwKc+9U2CyKSR0rWnpuJjbU/6Rx0qyf9acVFEFwJMfN6VYDBXOO1VN6jWiK7NsJBGc9qjQFTt7setF2sm8kcg1GGZ5B2xVxV4hsaMZwvuKcxAbkcmmJj5jTXY4U1g1djBG/eHJqQnbj37UwoS27NDJk/e4odhbASTkHgmkZN6YHbrT+OTSKcc0X7AxehAzk01GBOCMH0p3/LQtTGRc7zwaPUB+4BiBUSksjBvwpZTtUjHNNX93DtPWmloA5ANvNTL8/HaqUJI4JwD2q4nAODilJWAaF3Mc9PSo5CUBwMrU2OOvWmleAPWhMCILt7dacSNwx+JoPXPpQnqaYCZBJB7UmDuFC/eb3oj+Y7aYakrAbsYpW+U4pAeST+FJks270qBgOmO9RT5KCpBkc+9RznEeR1qluIgGEcA9O9OjJYNxx2prvkDufWpPup6ZrV7CIEcct6dKtxbSue1U8DdxyKswHEeBRNKwIfI/c9KgAzlhzirEmM4H41Wf/WDBwKmAxVIGSafDGE567qjB5x271MDuYAVTAcRgYbp6VEVHQU9yWbPamyybRntikr2Aa+FiYrzSWyFVBJ5PWmpKHGw9e1Pzg4HaqeisLrcUHbccc+tS4yc1DnbKPepuQBjiokPcU4BwKaGLMAPxFO4BqKTd1FSkBZHI+lQsSHORxT1OVHpUb9OTSW4XG9JCe1NJwSSeO1PK7qZ1jYVaAjLEyKfWh2xMB1NNUYIU09sbw1XomGtifO07jUvG3J71DHyuSelS5BH1rKQFC5ADhsVPbMSM/wA6ey7iVxwKau1cgCrbvGwE455prHOc05fu49ajkwYyPWs1uDAYwcGkUEkY6UJ8i7WPPanAYOB0qthB2oI4wRSgdsfjRk7sE0h2FxlME1UdNkuRU4cknNRspeQEcVUbpiJBkqBnimnAyfWpAcDHpTZBtQnvST1GMBBSpBkAVXglEkZPYdqnViU6cinJWEu4i5UnPTtTWOFJ7inhsnkYNNfpR1GRE5OB09KXcA+M4pSeRiozjdnFWhE49aZn5sU4HnmjocmpGNPzCoo8LKxJ/CpOVBpgAVixPWriwJ93vRuG4A01SCM5xUJkHnCpUQLbHGaq7+eamJ3CqyjJOaIoTZLj5gcVJ3wetNHB96UsFbnuKGPoMKZJqnOgwSB0q8eDg96gkX5NvfvVwdmDK1uhxwcVo8BRmqkS54HSrRPygfrTqu7ErDSAxJNQxH96e4qViQDmo12glgOvWpWwDgcyHPT1pWwUPvUIfDexNPdgoAA607DRUQYY98VYjBLjHSmxgeYfSpFBMgx92tJSuSybOTxTZQPLoLbSQfypF+eP6Vl5jGKDgDHNKy70wOo7084AGO9IDtGO/rTuDQA4yKbG+OetDghQR+NJGM59PWi2g72HD3600kZ707nk4pF6Egc0C1IpOMH1pcjvSN0x60EfIc9RVdAdxJVxyBz6VXkHII6VaGNgaoXQFyauLExkf3lOeParSH5ScVFFDtf2FTnBUjpSm0MAcR5PBpjuFJLUpHyjPSq9ydxVcfhSiruwhYmbL7unarAULjHWqwyxwBg/zq0pyBRPcESHIGPWge9GM5GaEPJ5rMoeep/lQPSk69aMnt2qQRWvV+Xj8ajs8E8fpUtznyGqKxDFNx6itl8Ana6LoyeOtNf7uBTgKY/TtWaGKDxz1oyR+NNXpyKcpy1AD14HvSNgAZoXr9KM88VIhh+8eaQc0rjPbrQq7V9aoEGABTD93g08dOfwprD5aaDYrqME+9BGOnHvUcgKvgGpuSOTWj7jEIGfX2pDnrQVZTwTTWcgDOKaELgswBP40rc4FL23elDYCZ9eKQhGx1NVbnlQTnmrTAkDFV7jcY+gx9auG4WM9ePrUrNuHPNQ4YvgEYzUgXPbiulkryBenNEgHUDAPvQN3A6elIxJXDYPrS3YeowgDOMnHSkI5PcntS5xx29aAcEgdSOtUG4HAPA4NQkqQexweBUmQx756cU05CtnHTtTRXmerHOcCkPpSsQoJ96aD39a+XsyhC3zEk80oPFRMSScDPNWlClORzVNWVwK0QBnYdcdat5yNuO/FMQbclV5p27tik3ewdBcELt96RgfwpQcDB5oXpg9KnrYOgoUNHnNG3JwOmKT7oxmgE9BT0fQBePunrSjg4HJpoUFix60/KgZ71EkAAY6n8KZnnNSdScjio9vzdaOgdRRyOetP4AxTSMYp3KoaFqwEwSAf4aMjHHFOBO3BpvfFLoA0HL804NjORQRtcGg84PY9aS1AjdSQcULwM09lwevFRt1wOatsLCAfMaaeWx0pcZ5pOME96Bi+gH4UgGCaVcnAPFL3IP509QYHKqCRTnwwyeuKQHfgZ4FPYZbaewoAbgYHrigglBTuGAxwRR2x3FPRgVjw24ip1XLgj61GRlyCOKmYDCstDvYB+Pn5FRsMk1JgkgihvlyDUqwjLvguFz61PCB5YwelQ3wBKsegPNSWyfKWHeupJSpNgty0Tnn1qjqLSBB5fNWnfGF70pwTjtXPHR3Bq5QsfMZssT09K1AMpnvUaIEk4AHtT1BVSBVVJqb1GA4FRTOFUntUxICEY5qGVVYjPT0rNbiYsDhoiSMkjFSAgAe/Wo0Ty+O1POGI9KqT1ArXjMg+Uc9RWUWmkYZPU1uSrvqtFaheT1Fa06kYrULalhM+UAeoFPGQg96anzP04A5p6nPA6VlJ3dxjJcFSwqnKSW8z0q+ygj/AGe9V2iAfGOKqEktWIrxuXlBJ4PrVw8tt7VA9sPNUirBHyDHXNVNqS0GK3PP4VCQW+btT14X15pWUADB61nsLca6BkH601VIZjngipc7lx0NOwNgPf0qm2vQCtIHLK272xUjjIUY6U8xjAOaimlCOg7mi/Np1AlJHYdqRP8AWZPTvSSZVc96cFwufWpfcBWH3vSmZ4x+dOJ/d571GoYjpSSAkI4PpUbgYFSD7v8AOo3IBGBwKSGNRuMCpl3KNvY1FCAdzY47VOvOQe3SrasA0jDfShMFT6U3OWIpqE7tvSo6CQ/qGyOKQcEEU5hxx+NNAz06UDFHJz6U1QCxY9ah3ujMTypqTzV+XbzV8rAkyD/WmMe+M4pV4BPc9KCP3ZpWsxdBFwNpHWnHlsCoVbaox2qc4wHHWql3YDZOST6Uq/MgpJOBu7Gl4ES1NtBiryD2xSYyp7mnZ+Qg96apwAe4pAPPIA7imY3SZHBFO6EkU0cEnPUULfUQgzv6fWkzl+O/WlXJYnPIpuM8jrTuAwn959DUuPmNMPP1pVJBFDYx0YCkilUYfB600DaAQacRkAjr3pgNOfMyBwDzUZbMv1qbjkGmMF3cdRQpWARR+86cinyHAowQNwFLJ0BxSYbDSMKCepoHC7hTZHJUAdRTYzwR3p8ugEq/MSx60qE5JHakVRsz0NKjYDYHal1AUdc+valViHYgdqamDzmlXcSxHp0oV0IaZUSYA8+tKByTjgiiMI7ZK0pyGI7Y6VSaQEaoGDD2qu+dvHBBq0nAZhUZUZLEfhTTsBABxg/pSjaSFPTtSrGdzMDxUciFowynkdatWbsFxyttJ9R60rnag7g0hOW55zTn+XAI4pMY13xtxSjDHpyKjYkqCB0NVxdBWbJ59KahfYC6pycjoKVct3qOJg444FSD72BwKhqzsMbk4IojGRjuac2AwB/GkA2vuFPoIIv9YeKcvAbJ701gQc5705hTWoEMnAIXrUDvmRcnkj8qnkwqHNVVTJ5PJ6VtSstRMY/NwCevanFm2lcVFdROrB0OTU0QIiDnknqK6LXjcRRvkAh245rP+fysAYrYvIx1PJI4qBLZWCqR071pCXLC4mrsowzOnJHHerU97I8e3APFWY7JGgJHUnimCzCNyMCplVg3qgV7WMdGOMk59BW5ZyFoQSeB2rN8nMuAOpq3DlBs71pUd47AixK6liR1quqMWJ7etDkghgM9jT0JZCAfrWMVZaBclXsQeakxmMhzVXa4XKndirUSMyZY8mpkla9yiREG5W6gdKn2hicVED8qgDAqfG2ueTGMPTC1IcYyeTTAPnzjinEEtjtUgVrlAY93eslonLMCeO1bUw3AqDxVR/l+U4xW9GTiQ0MhjJjCnpUoAwSOlS7QEBHSkBGw8cUTm5NsrYqsMOWPT0pFJLFsdelSSrlduee9VI2IfbjpWkdVcm+pcVW8vBx60csR2ApAuBhTSMoXqeam4eZM6h0BzxSAMQB1FIPnOegpQ2T1wQKTjuhlhRgY9KOWU+tIjHbyaXPy4xWFtRmTdIYJcgdaayNI6tHxmrt08P8AGMmoI7iHzFUV2U5Sa0QtCdLZV2sxyRU7cgAjioZM7lC9KWWTAyO3esHdjHxsqEgLxTxncxPQVVLk9Dx1qyGLJSkmlcBGcnp2pGiEnbmngDpj8TTg2AWqb66Ayq0YQbAe/SoywDrz8tSyEbiR1qs4AY/oK2jqA5sbzinsS0fI59ariVs7QMn+VTPkxEHrVSi1YRBuIclgPwqyxZUVqoniQn17Vc3bkRTnkVU46XASZjuDKOD15p+7euf0qGVR5uF4AqQkIVA6Gotohk4yG6YFOUtkn1pnmbpSh69qmK5IHcVnLTRgMK5bB6GhCI3K1NjPP50xlDYOKi+lmARquTtPWqzu3mFSOnSrBURr8o4qGRCrFyeKqNrgGMNTXbBPvTDJ85Hakzvc9lrRRYBIQFA7npTWbCjHJp7knaoHQdfSoXVllAxxVqwXLO7Cgk/hUjkbATjFQKv73DNx2qQNkeoqJL3gIwSHBGSKtHHljHXvUEeN3NTIPlJJ4pTa6AG/I57UBi6k9h0pGXjNIT+RrPQCI85IOCTSAGNAcU7AYAe9Oydwz0q76APydnSnjikBG3mnADH1rNsYEgc0dDk9KGAOMijsTngUhAeWAx9KjOBkHg0/JIzjpTONwPemgI2YJ061WYqEKk5zUrgbzmqznEwXbx25rohG4MsMf3S4oYbAuT96lVdwX2p+3zD0wBUtoCG5VsIQfrVaNNspJPTpWhJHv46Yqt5W52yenSqhJctmLqSwF2yD3qQgkewqO3JTg9alI+Umoa94fQQMBhc09sAjuahkAyPQ1MMABzxxxUtdQE2gsB6Uxx8hx+VSAA89qaMLGx7UIAQlohkYqN2Bk2HtT9wA3GokQlyxGDTS1bDVjpBlgTzRIu9Dn6UFgDuIqMOzKwHFUkwAIvy5NTxjA3eoqJTtjwRkjvT4W/dHmiYdSYZOO2aQkDPpSqcjFM9c1kMa2CKaRjpTm4Ax0poGRu7VorCEc5kO3vSJneSOKcRwR3pgGMiqWwdSYklqUdMDoKahzjnpT+/tWbQDOFOD2qCXJ+b9KnPcnoaicEJjHFVHcCDI2hgOvSn/AHlyeophIEagdQaCSY/atbC1GghSQfyqeIHOO1Q9+e9WEOQp60pMY6TlSB1NVXBCgdzVth8uT1qBsOcgZqYuwmNiTK5FPQFjxSeWRkZ5709F2jNNse4pBAK02RMxDPWnqpbNKw3fL6dqSdhEGAWBA/GmkgyEZ5qXaFHvVYD9+QPxPrVxsw1JpnCSKetTgDAZjUMu1XBYZ9BT9+8dOtQ1oh7MeWy456Uh+/xzR0BHenEYAIHJqdgFY7eB+NNkUZAz0ocEEZ/KnScKKQyInDcUN9zjvQ5x2605eQwNV5iKyHk5605V8wZzxSEBJsHnPakU4OB0rQVydBhSM5FSrycVAgJGM/Spx8q4BrOQ7jT98+lQKRvNTvx0796jKgNmhASISxB/OmyjIyOlOX7mPWkbp060uoEZYMMn8qeOvIqNMu2zHSphjBwOlVLQAY4+WkLZx6UNhRj9aYxIwP8AIpJXAawzIT6+9TIpCjjiq6ZZyCeKtg5IFEtBDMfMTUV25EXA6ipT98gUydPNUBe1C0abGULP+LPer6A7M+vSo1hWMjAqRMgGrnLmd0StBOAfeg/dIxRk7uKOxJqSiFs4pvA25p7kYwO9RcHjqatATjk9MUH/ACaFPbr6UvTpUgNYnjpUDYMnPb9alcHgniogRuPHJq4gTsR8tVpCvnHHrUs5wuc9qgjA4PeqigLseCABTigwRjJpkZGciphwKxlowKcrMpJH40Rt5kasetSzcD2qtD8zj0rWOsRdSZsnJpDnbz3oTJc04jJqdh77DVUA8dKcSSKReBikY/LnPFMLA/PTtUTj5MdM08YQAk5zTSNx3dqpaAM2bRn86N4aMkHp0pdpMbEd6bHjyunNVuIRDxz1qdSd2ccH1qEAbgTyamU5GCcmlISBgC2KEbOQOo7UxNwJ3VKgC5PrUvQr0EYngdqUkZ479adggE96YMEA460gG4LKaI87eRSHIYEflT85X61XQXmB47U3dtT1pzD5MCoiMD5s0kMieQA0B8sB2prrhsHoaYoJJ9q2SVhEqNg+3pTvLwMjrTeOp7VMOg5qGAKMKBSswyKUc9KRRuycVACcgqD0qrMhabg8CrRB+XNQspEpOauLsxW0GxhgQSPrVhAR9KgG4RkZqePhOaJAiQkiiP759KilJ4209c7/AOtRbQZKOTSDoSacSRnim4wuakZHJhkI7U1Cq9BUhB2UxegGODVLYCRTx701jk0oXFMYktSW4CoTxmnD0qNex7d6lyARn8KbEhyjDmgDkn0oz83FFSA0nJFHUU1iMdeaaG4OOlOwD8bgaa+duKXjaKGAI56npTQFOcgEHHFPVwRx39KbMMKfeo4VyfTHatrXiHUn3bSQelNbBxxTsbhg8e9OPYelSL0BR1B5pG+9z09aeOR7CmOeSfzpLcA/lUEw+XDDg1MnB61HPkoRiqjuD2uZuMScj60Kf1okBUlSOnTFCcgZHPvXUQLt+QDsKac8j065pTxtwcGkxkkE0DGDk47HvSH7+Bn6UEncOmc+tByZAT1qwuL9ABz0qPLDcOmQak5z1yP5VG4IUnJPGOlC3Gmeqn5m68UMo4A6UpUbelOxjA9q+Y0sWNVOwFPH3MY5pQMNzTj8z5FLRjGKCDSjB5o3EEg0oUgZX8aLJ9RCDnjtQX25BHSl2lRnFGAcA9akBByMmnkcZ9ajJXOKkAyMZ4HSmwAZwBinJjbSAkjAFKCcYApJdwYrtkgdKY+Cwx+NOByCT1PFMbg5pMBQPm5PFOI9DxTTyRx1pwG047GlYBeCvvSDgHFKcDoeKT+LNLqAbcnk0Dj5BSA7myRxSA9SOtUuwxrN2PrTDnfxTmAC5pMhTkd6dwEJJwBTeScU4jaNx70gI307a6gLjj6UpGAD2oxzilbr14xSuG4nGRt4HapCAW46004wMdqXJJB70mA1SAzA07o1IMbyW60HOM02BG/D+pzUyn5MfpSBVKnPXrSgkc+lOS0AVcjgnkUAcnPXFHRuT1oYYcHNTvqIrOgfIYcZoXKkYGBUz8OB61Hn5wK1T0sAjRAyBh1oZcNmpv4gaR8En1rNyewDNx3A08N8nFQu5Q89qehDRZHc1TWlxjySwJ9Kgm4+bNWRyBj0qFovMUr+lJaMTK4m2vuLcGraspxg1VltA5U56VIkRVutaT5JK63AsDO4gVHLmMFhzUnRQR360hUFcd6yS6A7jIXYDLDBPFOBAb3NKxwAfSoPMxISR1q3rsBP3weQKRzgMccUKOPrVe6fYABURTbAmHMeR2ppOEHPeiN8RD3pxUfKfzqmnEbBT1U/nSHiQAmkkG1t2eKVh+730XYhEbLkmpZBj7vU1GhVwOelPXnpnilNagDHvUboJCGI5HSnsc8HrTZjgEjtxQkwY5+eKVh0+lRoDsBY809ugIptNaDuDEbeetIcgLSyNkrxRI3tU2EA4TIPPeo3UkEipG4wR3pOmOaEAxCQCuOKcrYYYNMZwjqp708Lgk9h0qnewDHBH0pEkBmYHripHXceDx7VFHESxbFONuoEjEY96QfKvWkdQTnvTc7hgdaVuwCPtYgZz61QeURXXlg8+lIgmW/bunpVm6hQOs23kdTXQlGNtdxX6ky5DDnj0qRySMdqj25cPnHHSpHP3cjisZIor+WR0bPNSlugH40Yxg+valeLZtbsfSqs3qA9jkADoKE+eLjtzUQkA+UjrUsJ2gqPSk4iHkblHtTW+6PWlAyhI4NJ3GelZ3GKTnA701iAQBwaYzkHcBTuGB9qdnuAbgKTI2+56U7GVBH0NAG5umMUCGOSiDihDlc0snX29KhRz09KpK6C5MMhRnvUnKEEelQhs/hUpOVyalqwwORz1pML97uacPu5JzimgblPTjmhINheRgdqJfQHI7UoOVHqKb1BzR11EQscMCBT414Y01UyTk9KkX+L+VVLYaFVsgjHNCH7w9RSqQVPqKavc1KWoDh8o+tBbaVx1pBlvwpGO5hjjFAh7dcDtSfeYknBFCH58mgfOxHrQ1YYIBsPHeo3Py4XtU6jDbe9RSqU3jHNOOrERLxE3uaT7sYp0f3Sp70mcA+lX1GRnHykjikdt689qc4JhJA5Hao2BCE/pTXdiEbcqADkN1rMuLdjJuU81ovuMRA69uahgDKh8ytoPl1QWuPtgUA96t5G2oEwyH2qY42g1lUd3djFY7SM0n8WKAC2Se1C5due1RoA6Q8gHtSnkAjrTJMsv0p2cgU3awEMu3G1jzUKKvY8np7VNLtWQOTn2qBH3FytawvbQQjA7+TRb4dWRjzTGzyxqThQsifia2jpEQydQWUNxikIAYsRSXD7my34VVuc/KRwDWig7IVy0JCrcfd9PSoJ53ZznOD3pGbbGAODTozuKxsODT5EnewyBEZjkCplQlBt61biiCowx16VCo2MQehqJVHJ2BIjICjaOhqMgxjH50MdkjZ6jpSNIzFcgimrhdEkTrtZT1arEMpJCgcVURFMvPfvT0PlN170pq6sNGizBQB0pvnAtgHnFVHn+Q56023crJuPNZqk2m2O5o8henNOA7k80i8+4NOb2rne4iCQZQ1S2kMMnPrWg+MHFUZv3SCt4PoJ6EwIMdKoG3AqGOQBKlQgx5HHtQ00MZIQG5FVpBtIYAA5q1lfKLMMmq7bGHzdauF0JjwNy5zg03OXAI6c06MjZnHAPSkdiG44p7MHoS/LnAqB9wnCgfLUyEZG7rSgbZCcZ9KSdmDJVNKpLA8jFMVuDmms2M46HrWVtRmdezKHZR1qvbLlgWPfin3YWSfcv40kUZfJUV307KBBqYKgY5HelVFkQ8daZbMxUg9BTkOzg1yPcsRYiuSO1WUXaAx/KgMDgU7kjArKUm9GAxvvcUx1I47GpcHP0pknShMNCowZSTmoXbcCN2DVidjnGPxqmV3N9K6Ya6sTCJ8S425OKnmj8xc54xVUqzvxxj9atxnChm/Crnb4kHQoq4WTpnB6mrqSEOAFyD09qqXGBk7efSpLeR8qSKcveRKepclC7txXk+9MBG36UrndGT0z0pEUtH8351jsiyRVLtu7jvUsZYIWIyajXKqMHipQwCfWoeoWJRw31pE6MCKRRwD3p5bAK4rMCCRiFxnGaquHIC7uO/FWZAGIGelQTEKAo61rBgNSP5z6UshKjAAxUoXC7u1RSHIHFVfUYREsN2aR082QEGo1YoT3WhDhiQeCfSnbW4h7IA2SaXOA2DTjhx8wqFVAlIxwetCV9wLGcbTjr1qYcKR1zUWzgCpmOOlZyYAVDEZ4OKixlue1K77XU0hky3HApWaAiVW3jng9qlzuwD61FnaxNSD5hmqYEnYgdBTl6n0pucp7+9ABB296iwEoJ69qaACfbuacOExRjC4qABiOmKiZQuAetPOcZA5pHwQPXvVICrMQY8/lUShC+CeTU0o+YjHFUw37wg10QV1YC4hBIA4xxTwx8zHpUMfIJHWrCgbs1EtAGyZ2nHeqQnDMFI57mrxPJPpUDosnAHOaqnbqJiqckAdBUvRTmoFyshU1YGGQ1MxojJGQuamH3BmoY0ycn71SR5bJPQUmBIAAnPWmlvlwR1qQ/dDVFIeVIqYq4EMw3Rhe4NDMVkVMfjSMQWAxT329T1q9gGuSxxio4PusO/c0snJIWliyGx3NNbAG1mBI59qfECAQKYC0Yxmnpnpnmk2BKAQM/rSNkcAU/GEpjdc1HUBCdwpNuFGKa2MgKacemO/rVANLANuzTcgzfzpZFIIHWkzubPamrASIQwJ6inEjINMRuDinE5GR1qXuAw5K+hpWGVC07jB4pGPGR2p3DTqV3H3j0x1qNVA6dDUjnjFNbmIAdetaLYBp5Y4qSI5HHUdqhYkYwOtSwn5eOlOWwE7cY96iRfmqZjxjHNMZenNZpgNcHfuHSkXg7fWnu4UYPaohwwJ6Y4qlsBMMggClOM5zzTEyWJJ6d6egHJPNLYCIpvPB5qGOFg+TyR3qdwQrYPNNifJx1PTNWm7ANmGSOKlPOFFRS5UE55qSJs8flSe1w6gowTnmng5WoyxLgAVL0XHepYAecc9KUnOc0LwpzSN9361IEeMjmgNiPOaRehpY16g9KoBsiqYyT94c5qtFJz845PSp5lGOD0qGRPmTGcCto2sLqWUJIz3FTKOAT1quuVwBVhRheT1rCQxrYznNR53EmnsPmOD1phBOKasA9W9BxQ4JIGOKPQCnE8HIpdQI1AVxj8aCxBOKeMYJ71A7nHHemtQ2JC2VzionJJyTRnC4PSlJG0DpmqSsAkQ3SA5qyRtXg1XiO0helWAQTzUy3EIF/ecUFdqkCnDO7NN3HOO9TdjGlflyaRevJp38WKQHLGqu7AN4DU48/wCFHG40gOKAIZPlxgc03o1SPyaic4ya0WoiReG59KU9M1GpAA7mpD93nvSYEbEHOf0pFXgHtT9oyTSEEAjvTuFyCaUFttRupVOvPpTjHmQnGTUhgYkHPFaq0dxDrI7g2ecVazxwKht1MeQRT2YIRWM9ZaDWxDdZdOBUcMZVRzzU75INQfNuAqovSwEqKeTSAlcYpwGRQOnpijcYwkhyKa5I2g85pGGZCeh9aax3Nt6VSQrg+flXHFPJ2gKfxpHyAOeaVuRjqTQFhp+QAL3p2AqA+veowcMQR0pxJYbRQ0AwkGQY6VJGc85zioip8z6U9DiQjtVPYWtyRlpyfc9c0vXmmJz3rPdFbEnO2mKRt5px4BPWm5G3A/KkAgGQTScLj2qQdMYqORSVNNbgKfmbIpGUbcd6ZGx24p5OTt9O9PYCtN94DvUS53c1alTkkAZ96jxhsgfWtU9BNCRgMNrdamBAPHJFRjO9j2pgYq2elFrgTo3c05AN2R3qDdg47GpIjk9yKhoNRzHjgdKjZdyg1K3CnbimKwB56UINxjHIxjFPB3RgDoKTv7URkMrDPSmFkOZQVxjvTlOCfSmjk85p44bjrUhYkzknNL/D/Sm9Tk0vUZHNQAxj8opv3WxinP1poPz/AEqlsA/OcAHgUyQccdacO9BztyaBkIzuGRUvHFNC55qTbjjNNsWgZAPpTug5pnUdadn5RUgRsefSlVfTvTuAelLkbs9KLgNYYFR7/rU7cmomGR04poCtNjIHY9qbDEQ2M8ZqyybsHFSKoUg1fPpYVtSPABwSaTAzjvSv94ADP9KarY5/nSGKCCCAaYSMH0pwUA01x2PFNWARR0Oaa4/dljT1x/8AWpMcYPamD20M24A35HJ71Hz7HHrU9yuGJFQrnB4rqi9CPQcxJxjmm5z0GMUD5W5ppOe3bsaaB6kbrkjtilPOQc/WlfqM5ApA2Op4H61QhW25I9eaa5IQgk9KXAIJzzUcn3CCOccGmiluesAHk0Hh80jZ3H3pzLtYc54r5fU0HkAZJ60A5GAOaDyopM7Su2klZiFJzxjn1oBIalf+8Oc0gb+HvQvxAVjuGQefSmrlSM9KUBcYPFHl5HXP9Kd23cBvlgygmnnhqCTwDzRgs+RSdgHfw5/GgEquRQcjINGeB1ovrqADkhQOaD79KcT0K8Uhye3HrUsQnBPXijG48GlJG3AoUDHJxS12GKPm47jpSLwM9qUA9V/OjGFByeetFrgJj5dwHemc4JqQcZBJwe1RPycDpVuNlcEIW3Jj070gwefajOF4FNZgFx3pK4wwSPpScY3U4NlOmKTjGDT12AUdMmnqBnHtTQcrzT1GEPrSAaxwV4+tPOQ2T070iLuBbuKVOUI79qLaANYDcSOlL1wR0oI+Qn3oUkLg9KI+YAADICTgGlY/vOOlIeefemMCSSKd76CJsCTJ9KAQV55xQFO0kcDHNCgeXgnmk1qwIZCQQT+dMOSwI6VLNxgHFRcL3q47APU+nbtSucYI70xOeRUgHGaWiGVrmIzL9KWP5CqdsVYIGeahdMnI6inGba5RaXJ1OG56UhAEgINRxFiSHxT2B8ylKPLsA5cEkGkIHmDPTHFA5BbNRuTxUgPDcE+lL1TOeaZtJB9TVeeVo1AAyaainsDZbI3qBjmm+WC4z0FMt5mZQXGOKkOc59etOScWA4DDZzVe4RnjbbyxqdSVyeoo2jbuPeknZhYrwKyoA46VOwOwHNGAV60i9OTxmqlLmd2gEYBsKT2pGGPlY8U5sDkUjgFRU31AigBUtnv0qdOGx2NRqpx157U5D8w3VUrghwx5nrULgmQAdKmxz9ahCZm3E8CpWgE3BQ+tNYYAzUir8575phOeCelKw9BX9cY4496RvvYPWhnOBkc0rDK59e9AhpB3EelB6Z7UZ+UHv0oJzGR2pu1wEwHIbHSnIQy7c01AFU8446UqAFT60WGL0QjjFIuNmRRwq/XtSA7elCYCMcgkdTSIAu71p5HG7FN6nPShbARSLh9wGPWnygEbTT2TB2npTWUFSR0FVdCIHyjqB0pJXKKBTpyRFkdRVYlpXXcMZ61pH3kHkWlIdQRyRUrDcBioo1CAr2p6keXxSlZ7AMkjDYYU5CVkHHFIW2g5NMEh3DI4o1aAsuMEYoYDAxSkZYEUmdx9qzemgxrAY9xSkAEe9B5A9aUjJ57dKSegDFyFYYpw5BPfvTQfmPvS52rx1qtAI5FLRsR1pkS7eG71YPBCiq7/AH81UNg6geGz2zU45Bz07VBkupFSxsWXaO1GyAeBgHn6UidDRyRj0pRguozjHBqEhCqMg9qacYxSjlitJ1bpwKA0I8E59qcnKn1pqbsuOop4+Ugr09KpruPUFGG5PHegf6wgdKeVwgPrTMbOtLVCFVsMc/SlKlVB96TOF6c0rcKAenWkMYCSSQKemdpI60gyI8+tLESsZ568U/UB6ctuzzTGBl3Z64pRlevelAClgT1FNWVhFMljMAp6U5gcgelORQJCT1peknTAqr9B9BkTbic9+tNlUYBHSnRbQz0xQAjc8CqtqIrEsSOeM9KbcxksjA9KkXDSZB4pCQWbceD0rS70DcVPlbA71KNzEjtUMWAzGp8kpkdaiS1GOB2r1pRwOM1E7MQoA69TT2OMAdKloBWOI+KXsPem43EL29aceGA9KQETsO9Q7Skjehp8pAkwvU04Z2kMK0XuoCqRvRoiuDShWFsVx0qRmCy/7Xc00NtLgn3FbJtokayq8ZZuoFU5Pm+X0qfeAGTs3eqxQ78AcntXRRVtWKW4sX76Zhj61egiUjJHIqraxqGJz83arkIIZgeM1NbXRBEkz2HQ1SmBMgUd6tLuzsFNkTnpyK5o6SuPdEK24kOD0Heo3TOVzx2NTyTJCuCeaijCGFnJrSLer6BYqghHwDzTZEZVLHqehpHk5IQZNVvNdzsJ4rpjBvUTNOGMSQ89qmRQ/C1TtpnRAg6HrV2IOSMDjvWU7q9xourlYxgUhfC/WnOWIA9Kpo5aQr2rkUebUrYnxkgVTuAS5QtyO2KtIvJ56VWkZRIXPWtIaSJI0PQdMVKCQ+PWokmAJK96Y9wMk55FauLb2BWLDrz0zj9agl5Geh7VJA5ki5HTvSTHkY70tU7AxsbLGoz3p5UvnAxnoaijIIw1TPkKCKb3uhKzQRRksc806V/LxzTbdwxGOpqWVBuBNZttOzH00GgqVyTTZHyh+nGKfNwRxxUJJIwRxTS6jMtmJY9Oau2y7VyKhkjXceevepoMogyetdUrchKWpY2GL51PHemrMs8ny1YaJhCSeeKzbZtsp471jCKlfuir2LspMeJAeO9Lb3Xmr0wafcYMB46dKgs043AdqjRwbYa3LmdnXvTZR6U1yeoHAo8zCj1rJLqBHMm0D3qs64cY79asySDAzUDMN+R3rWF7BYhaH94MdKtwRjA54FU3JL+gqb7RiIKDWkoyaSFsSXiIOAapxNztBzio72STZvzk+nrSWTBYsv1NVGFoC3ZeclowyjB700vhAOnrTCzvEQvTNRxl5HGVwRSUNdR9S8nzKpPT0qQDGeKiAG4g8DsKlQZfr2rGSGSA/NSlsHPYUq9/pTZkMkLBaz0GU5mZ5QF4HemzJsKsWyKRIgpJZulPUfuue/eujRbCH8lB6UrBdhA9KSPkEk0Ac4BqAKpBBwaSNWQnJ6+tKzM0mQMDvSuQSWJrXZBuMUkOM1Nz5oIHFRNlsEdBU8XKZX8aHawEgPzHr9KUMVk56U0YOT0NLj5TmsmgHTkEAioEcng8VKVYoCRR5OQPUUJpAEijHHfrQhyh7Yp5XA680KoII6Ur6AKoy+PWlUfPnPNNUDbndSs208VO4DycnGeaXJIzTFOW2jqacOpWpaCwoG4+1MdcsB6U/ODwahnco/FOKu9AIrgheO9Zpkw24jBq/dksuQeR1NZZXrjkV10loKW5oW8ynOepq2G3AhTWVFGT93mtGOMIM55qKsUnce5IFO0mkXAIxStk4Hao+VOAayWoDFX9+SanQbVbJqtIpFwMdD2q067VzVSeiBaDU+U9etTxjCsCelQgjcAKmx1xwazkwFH3eaimBIAqU9hTGGSfapW9wIiuACfTmmBSGNSH5gTTZDtwPatLtgRvnOakQAkH1qOI74ixprlgV29O9XboFx74LdfxqSMLzzUUoCqMHNLE4BXila6AsLzz6dqRuT/KhDjJo2881kBCVC5Pf2pyngk9RSMQM460i5IHNaW0AeRkbjUZz1FOJyNtI42qMcUIBw+QZ60vQZ70xeSfapf4OaTGNPH0oYErkUAEnBPFGcnFIRXdRkkAkUwKd2e1SSkiTbnrTemV71qgGtkNjjFPjXaRTHGMg8n3p4OFH0ptXWgbEx5P0phOPnJ4pynj6UhwyY71AEUjKVBPOaQsDxS8YxjmgkMQO5qhChvlOepqWM7V571EwHTrTl5Ue1J6ofUH+XdnoajhTauc802aXbjP0p0RJTLGqSdhdRZFJyaRW+7t6U45OfSmxKAgH60dA1FfII4696mx0XvUbqSRS7vnC5qdx2JTxkdqRvmTmkDBjxTuDUbAQBWAx29acPlIHrQzDHFKB8vvV3AimGGOOnrRE29AT0p7AFNppin5CMcDpVLYOpID8201LyFx3qupIQE9asZDAH1rOSAY3AGKbIdqe9PYZJFMcZPPQU0AoI696eOc4qBGy23vUqkkH0oaATcclajc/IR3pTgSe9C8luapICLJZeB0p4G1AT2FNGBuweCakHyqAe9OQDImUvnvVlOVJzVaIfMRip14BFTMB6nvjrTVHzUitg4I5pWOSCamwDOvXtSg9xSH26UE/Ngd6dhMVs7h7U1jwD2NPfsDUB5THpTQwflqiGScCpJD8tR4xjHWrQCquBknk9qkHbmoxnA9af8AwgUMBerYpOrZp2QDTDwKSDqRfMrAAVZHIGetRblHJ/Cp1OVzTkArYVcmql1NsA9D6VNd+YYfl61WmjxEATjFOmluxeRPE+6PPTFRSndwv4022b5CKHIR85H0qre8HQmBIGDTM8Y70qupzjvSsO56VIxjjKcH6mmNyy0ZGCD0NKFJcHtVgwlGXBp6ggDPNEg4OBgU9CCmcVLegtCvKNrZXqaVTkjJ5pJQRNTAf3vNWtUMcxIbFEfLUS5L5HSnj7wHrR0J0JmAAIqOMYbNSEgDFNxls4/Os0VYUHH+FNGd9PYA80zcduaEIf1YelNJyvFOVsAZprkKcdjQtx2I8YYY4FLwec80jnBIqNCSGA/OrtoGiEkJI65qMP0GeKfK2046im7QoJq0tCbjhkIc96b5fmQ8cU8NhMYo6IcdT+tF2PqMCkqAeop8ZwCfWmJkDB60MfTjHah6i0JwSQBiogMt7jvUsZz8wqMgFyOlSlYGB4QHrTxhV9jTgAU4HFMwxT2pXuND+OD1zSj5W6daVRhMDmkzkYHakCHjgY707tTc846048j6VAxkn3enNNGc/hTpD0pAMYNUtg6gB83Jp4G4YI4ppPHNPHApMQwDgAU7GB0pB+FOXJWi4LcaBz0ob0pGPYcUOcD0NMY0dTmn4OSajjA209TkUMSHduDUeeeaeaaenWkgFUe2DSc8g0uSOnaow2XJNNAJJnHHWopM9+vtUpBJNRbepFaR0AVVO4HOcUrA46c08Umc59KVxDYzzjvUbnCkA4Ge1SKpDH0pJANuQOapbgZ9yQe9QocBeME+lT3Awq4PNQDBGeOK6Y7Etajjz060wYA+venCQlfm4yeCaRQAMBh9aY3ohj53Dpig5ZwcUMfm9fwpF55zk5q7WEvICGHPvTHXCZ9qkHU5pkn3DjFC3BHqvU7j2p27ctNX7rA9qOcAivmGaEuflC44pPugUu7O3inHaEx3xSd7h0G5+XBpgyGzintkkce1JnjA7UOXYLD3GQDjio8kcUoYlaAN34Cpu2wHHAHXmlwUwTTPvY9c09m3Ag9aLdRgQcbs5FKv3envSg/Limrz16U7i6j8lh0pN2BgetNZiDxSr6Ur3CwBhyCKX6DNGNuaEOQR2NJXbAcCeccKajck4HQVIVzhQcimPkgAD5hTW+oCfw9ec8VGX2tsPNTdRioQi7ye/aqvdhqBwI/eqxRy+as9X5HSjGee1K9hjQpCgnPSlGOSacxyo9qFA3Y7UPVgAHzAEfhTiQDweDSIefm+lCLk80SXRCQ/7qjHQ0DgbcdaTBKkehzS8ZUHrSkAjjpg0meSBTiMHApoxt570kwHcAU0HjGMUoHy80Hke4pjHRk4KnpS4GMZ6UxD8/PTFOYHBYUSQiGX7uSOPWoTngmrEwzGoqH+EccVatawDoz6dDUwwVwOoqJTt46g1JjoPWpkxikDy/xqFxzvB6Cpv4So9KMAxkYoi+omREgru4zQGGzrg1DckJgjr6VXebcACMEVrGDkrgX06EN+FOHIIPaoIJ/MwMcipiOKykrOwxxOGGBxTWVSelPGMHOeKQ5MePel5IRHx0FPfnGO3WmAYfpzT/VsVTAGyuRSscoBjj1pM7+aMnyyD+BqUAxsqw9KU4IwKaOnNK/yKD+lNAIwzHmnZBQeoqCObfMRjAxU2ME4qmkhh0xxzQwyd1KwKqM1GJONpFJK4Eh+6CKXho8Ed6T7oHSlyVTHr3pCHR8Nj16VE43PgcVIAd4GaiQN57BunY0RH1FBz17U7dlSO1CjaxzSD75yOPSh2BETMc4x0NSjlCPTmkYAP0zSrgk9uKadwG/xgjkUHG7ikVwvB70rjkEfpRYBxG44x0pjcuPapGzu4PJqFgQ+O9G4E+c8VEzZYLipeAcgdKiYZYkcULQGPLZ470MMcDkU2Ibmye1OJ+fNJ7i3GE5OAOKZsDMD3p4Ayx707GcN271eyAhYgEqR060yJg6ZX7uetSOobc3rUSRiCHaOlOLVgsEgzkjt0FNQkvz+tTKMYPao1X5yT2NNMC10cflQDgsD+FMXLZY0N1VvWo3GPyWXihjkjPWmqcScHg0pGSc/nU2sBGjZcg0oJB5prHBzikzwSTg1e+oEgORntUMgIJ70+NgVK96H5IDcUlowIV++MVLFhWwOTSEbVz3p8YIG6rctAHrwuO9NzjB96dkbs/lSdV5rIAXIJPrTuh6U052jFOHQqevamGwn3QT2oDfLj8qU9OaDwAMdKOgB82zb2zTW+Y5FPGNn86aPuMfyp6iEJyARSv8ANtFCHIOe9KPl6jrxQ9xiFsqR3FL0QcUjDZkEUMfX04pbgP8AvBSe1ISGU54I7UZ2hR2IoCgsQfSmt9QIyMkelHDHntSfeGD2oTIAzQAIuxskcGoZCqDp1NSXO7zEC0lypYZGMr29atJCK5hUAkDgis91kaYK3AB4NasiEwAKcGqq8sQ/Ve9b05NMLDJUZRj+E1J5pjgUnmlbngjk9qR8tGcDn0qb9wJcjywRQCZGGKrgO0Iwakjm/e+WBxUuN9h6E5PzADpT1UZ5pgznPen+1ZvQCLaFff6UkkioCx6UsmVWopP3ox6VcdQKTTmTJA5odiWXHXvUkUIVyx79qQJsuOTwa7IuOyJDyMJkdcc05UBlX1p8QPmMCcikXKzAr0p82rVwHNDiZWTv2p4yScHkHmnS5ByKXbhgAMZrFyaQIY52uD0pp+ZwR2p0oJPT61E7hHG3qRSWugGfeqxlbdwe1RiVzb7APrVy6VpZVIHFVSfIc1102nBJ7k7DAVjAIHJpjIdwY8Z7VI20qHzwOlLuDQkkcirU+6HYsxxfLjP0NXLYER5PXNZcbuyA9hWlC7MQgrnqJ20GWB6ntUYT5mcDrUw/1Z9aagCk+9cquMUqdo2/jVC6wBkVojIUmqF0vDA9KuD94T1MySYgkKenWktiHm+bBqth+jE59asWSFZCSfpXoWSiQtdzYXbECOh7VE3zD3z1qTeCuD3qOcfLhevtXEr3L6DBiMAnGKe+XAC8iqzpJt+9+FTwM7xbR1FaNaXAWB9xAC4x71O4ZnqCMMzkZxipN7ZIxnHfNZy1egdB0g7E5qNjj5R0qRSCpJPNOAD8Yqb2GVXiAAPenttBUcYpzJ8wBqOQBeSa0i79RFl5Qw2dqzShhnzjg1Ya4RgoQ05XWT7y9OnNOKcdRvUly0kZJ4qSJREoAOc1CGIJ44PSpYxlAxrOT08g1JGA3AGkdAW74py88nrUZfB2ngGs1cCF0ycHgVWk+ViCeMVdkYEY9KzroEsMHit6d27MT2K00mG3ZpizsScjHHrULlg5zTSx2nA59a7Yx0IvqOaYs/zGrlopeIsazDknFadmSYNpOBRNWhoUt7E6MP4e9NDDccDJpfK8sE9u1JCysAzcYrnVtxl0KwwT3qTvzSId6DB4pSnbNc0rp2YxU475FS9cg9KgGABt6nrUoOeF61LQbEPkxvlj0qvcRldu08elSSiRZCoGBUEk+QoY/MK2ppt3uBJC2UPrUoHyk1WgkZicDr3q23y0S0ewFRox5uT07iolUNLxyvpViRQZMfnUQysuxRiri9AFEPJLdKsqqJhR1piYaTae1O27mxnkVMr7MBvAJAFPU7gcCmMu1ztpYD8p9al6q4DxnBU08DC4oiPJNOyBk/hUNh1IzhT6+tIpy3HSmt1Jp6jByopvYBQBnA6ihT14pFyDkmpAPlzSvYBsfJyakAB6U0AgU/IHAqW7gNCjPHNMkAK5b8qlXIU+9RS8CmtwKsy5X2qnAi+YS1XJlLxEDjFQxw5YbuvrXTB2ixbjwqghU6nmp0G0FSajVQG3d6eDwSetRJ32HsSkAJjgmoW4kFTD1NVZ2+ZdtTBNvQCRyplH86lPIHPBqvESUJI57VICTHkdRTkraAPQbfzqbIyD3qFT8uG61KOAG71nLcOgMfm4pvPNBzu9qG68/hStoAcA47UzAJ56Cn579qOd3rTQEQAXOBwajI/ekc4qQkbgO9RSuFYccmtEHQY4CrntToU+XrmhsMgXPNLEu0c9Kq+gFnkgeppz4UD1pq8jJ60EcZNY9QGOQc5FRp2AqRvmPT8qYoKsMGrWwAPvYNI7biMDjvS9/rQ4Cr6GjqA1OW21PznFQxcHIqbGVOetKW4AwGAB+NNxg4FOPQGkXnmkBWlGJOe9KfvU6UBSM1GhLNzWibsAPjv1pI2JGaGGWye9Ih29Dn2qugbMnXgEmlDYU570i9fXHekK5qA2Ij8zmkPD5FOjyrEAZxQ3f0NWIWI/Kc9TQhO4g9KEx5ZGeadgkc4zmkMiljDkE9qCdqAGpuMHFV3IEZGOKafQWhID8pxyaSFVEfv6ULlh6CkRssee9HRh1JTwu40MoKBx3pzn5Md6bngLUIYsQ55PSn8c8UkfCYNO3fN7UncCPGASxoxuGB0NAHzsKFyrgdqYABlhz0pBgk8cUuQXzimIMvkdKdgFQ5BBFPHTFMQ/NmnIwJx3pMBzbjn3pnIXJ6mnvhl44xTDyoz0oQABhi1PyBkY4xSA8A9qF+Y/0pARnGTnr60yMseD2pspwwPpSx7jkseK0toA/AI4pzEYBxzTY1PpRg4/SpC46Ljd609Tz9aiXIJ/nTipyDngUNAPA+c0nUnHWl43ZzSg8+1SA3ikXpmnFSc5PWmngDFMBx+4SetREDbzU2DtIpjLxQmBC4HWmNkjPWpivy5NR545q0wGAEZpyEYAFMBOxjQnUk9PSqYEynJ96R+cnqKTOABjn0oc/LjNK2oivLnGB61chXEY5qizfMKuxjaASaua90OpYO3YKrTJvT61NjLU1uenQVinZ3GUygt0bnBNUXZieTmrl0SzYxz6VVZckcc9ua6qe12S2SWrEcHrVxh8vNRRquc8ZHepCSBnOcVnJ3Y9iNgMfj1pNwA4PSpOikkVE6tsz3oQMmzujGaWMARkd6ihbKj+VT4AUmk9NA31IgpJ+lDIpwcc07qnHU012AG70ouxkTYCDPU06IngntUMrFm44qZeMHtV9CbWZID8uR+VKvA/rTVUjPp2pRzzmoYx5zmojwcetStyv+NRnBOMURYxx9cUOhYBvSg/dpwAKGk2DK7tk801W2pg02WM9ffinhCOSeK10sIWRcjIHNQseOv1qfOATUG4b+aIsTtYGBAzSDPl8HkdanYZHHeoNnB56U0wYqkEHnmmuxaM8UKQvODn0NNZsqVxj0qrajJbc546+9SbMSZxUNqu3GWzVkDrng1EtGKwozzzxTfuj2pxwc03INQV1FZgiVEGzKQDijJZc45pAMvuFWkK5YA6H0qRRxz3qBWBOD0qYdcVk0MY5wdtGKJBkk0Ic8dqfQB/UjNKehpoILAU481IDO5J79qeoPNMPUU5c7etNi2AjnNNfBXOcinP9056elNcALx+dCDyEQAA0vaoxnjin4702CHDAJppOe1OGCSMU0jIxnrSDYCcr6Gojw3tUgHUU0rk4PeqQCKSV5pnG5vapQuOQagP+t74PWqQmSKcj2FIQAopUXAwTxSsSRj0ovqMOi/WmSHPOOacRyP5UgXPOelCAguEBjOOlZysBxjrWrIP3ZHesmU7WJA+bNdFF3VhPcVlyvGaEAHXmpD/AKr1OeKhXIFap3VhDuje1MGOmKVwQMDnH6Uik9jx70xAhO4hcH60kn+rOenal/j4+lNf5kPrg80+oz1Qcg5zSqD0PTrSHIAHv1qUgeXnPPpXy7TsWKpAUkigAnJxmkRxtwacGwBgUm9kADJOMc0D5Rk+tBODSsMpjPOaFduwDcfMSOmaBwD60EYAHenkLt4PPvRZsBoXlcetLIuGz70vK8Hn0oJyQDxTaiAEZBFNB+Q0sgG4bfxxSnBI4NJANUAjnvSgelBG3tQPun0pabAL14pwwoIwaaM9AKUHg56mjbqAoO1c9zRjK5zz1pDjGADTHYKm4k+lNMALbQTTDlTu9aXhgMdxQ+Qdv60mrPUYgGMk80gBOQKcABwPzpOQKABSF4NNHQtTyBt3cUjcKKaXcQZyBn60/gICeuaYR8uO9P6rg+lG24xx4ANJnccj9aT7qYzT8qUx0NTuIQkr70mAQPrSE4X8aAR05waaWgwPQ+1J0XJpemfehiSAKLWFcRFO/g8GnFypKnvTCSuKcxIHPUim7AEnzIAOwqDoKnyChA6iq5G2QemeaqL1AevoKkB45/CoxnzDipAAeg+tTIYmfnPpTgeM+vam5HP1pwwrKc8UvQRA8G4lgeR2pjxq4HY1OxwTjoaZkbcGqUmgsRqm3kdcdamUkjOeTTEyKkUAginJ32GOH3CO9NBwCKcPuEj8qADgnvUN9hEPmYccdalxxnrmmsoKnjmlQHJB6VT1QCg4X8aUjMRI6ZpoyQR2py8gjOKm9gI3PAI655pcBhz0FRyEIp3cYp8JDocHtVRTaHuQOhU716ipFk3kYoY7R9arPIYDkdDVr3lYRakbcPcVGqB8e1Kjqy5zzTUk2Sr6HihJxGycHJwRT2YYAx0pmQH6d6Vs5HvUSuxDlHPPehflLE9e9K3RcdPamupUE+tTbUGAO8H1pSeMfxetMUcbsU9jllanoAw8AjvSL1pJGCjcT1pYyGG8HjFNJ2GQlA9wGz07VPjcwxVcjE24HFWC21lI6U3H3Uw2F4Y4Pao29+tTAZfP41C+dxGOKS1AkQ4JBPXpScDcDTevHelbkAY6UtgEQ4B/Sn5+XJHJqPaVHvUm7KCloAzaQpJ60bwRtXvTyCeAKhVGWQjvTjruA49lPeopV3My9jUj53L6jpTUJJ9xVrRXEMORHgdjTA37znvUzYHI6CqqkSSZ6H0q46gW1PyGnHkgjtUYIPyj6UsZwCGqGrDHt0LdCe1OP3eDTc7SM9KVjhj6HpUu4DHA2g1E6ny/cVIw6D1oPCYPaqTsAxOV4696Ujfhs9KTPy5HGKF5p2YEm3ccntSryG9KReQQeM0qHlgBnilZtgLwV4HSkbJwB0p0eMsDTDwM5pb6iF6kY7dae3J3U04A+tB6L6UulhsWTlhjpikHQmlcbDtPekGVH1piFXoc0LgIc0KOSD3pobHykd6LMYqfdI6Uuf3f40Dkbew6ULynP1FCTAUZcbj2pgHX2p3/ACzJoIG0UhAScD0o5IK02Zgjrj8acjAhicc00mMiZhuGOaVSSMVHjBJ9altyA/PenbQBWwWU96HAZsDpQRiTpSgZJycUriIGO2QpzjtTH2qpHc1LIuQWPWs+aRvJY9xWsVzbBsRySgS9elSh9q5zms5WLgg9adMzAZGcjtXU6Q0yytzulwOnegnMm5PxqvFEzx7jxipbdeSD0qnGKegi9GxZcnrU6Nj73GarMwTAXqalZjgHriuSaTegxzqFXJOartlE56Gl88SEjOcUrNwcDPFCi1uBWkyCv50kwB2uO/FPUbyxf8qidwAVHT0roitrEsmRgAr0ws28qeB60kJDREt2ofcJ1I+7VaJ2GW8ZTk/SlU/PwOafIf3agD8ajkYJGGB+tY3bdgY5iCp9qquirH5ncVaAJhJzScGD271K91g0QyfdDY5IrMKq2S3LVrOVMWD1FZLqVYr0B6V1UHq1bcmRAEDPwePSrKp+6PHHpVYq24EDHNW1YmEgdu9a1ItWYIhjlHKAZq5b5Uk9+9V4IBvJPHvVuKLDdeO9ZzkktRouJ6npTCcS7fyoLdOaCu2QHFchRKTkHFVbgAxjPepkJZz6ZqKZcqQOlOOhL2Mme2JAYDr0qzbWuFDGlYkcVZQhrcNXTOpJRsJJXEQc8j6VDIxD8fnU7kMgC9qhcnbk1lFjehH85fJFPiJ5XOKjJO8EHNTQrlsetavRCGhT5uCalIYNntTnTbgL1p65KbsVk5aDGrhSRjrQWMTHjNCndkk4xTJUOQ5PSp3DYqXU7kfu+PWqYZ2Zdx/StVIlmjfB7VnrGY5gG/A11UpJaWE7lpYlVenNSBMrkHj0p6qOMmgfM21elYtu5VhqfeyRViNsj2qMJ0BNOT5flqG7oaJiAE46VXmO7juKm3ALgmoDyxGKmK1EV2DMrFvxNRBCyDH41alAXgnikEaqm5ehrZS0F1Mu4i/ecHrUDrhMDrWjcj0NU2AAwT7dK6ac7ol6MrCFt2Mc1Yx5e0jPFPiT5smpSozWjk7WAuW7CRenQUySMqTtXrTbeQJnNSea3llhj6VyNPmuix6HygAe9TjBFVg+9M9TipUkJxkfWs5xfUOo1TtkOelTxggZzyBULIplG4cdqnxhTUtqwIN3dhwaoXESh9wHBq5uVmA6mkni8yPApxaiwepQQmGVR/Ca0AQeewHWq8qAIq9/Wnx/d2GrbUkABcvupCN0iyAfWlJABBOOODSRj5Sn8PrUoAXG8svepFYEn2piIEHB49aYpGSBzRbmAkyGyRTV++R09Kc4AxTNp3AmgCVeGxg1IcDFNTG4MTSg7jkjpUMCJmIbJXg96EbPTpUpA2g9zQiYTIp3VtQIhy5B4qfOU4NQvnsM1IhJXmkxjvlBApcgZxTB0z3pcnbyOamwhwOVxTZB8g9KXquBTXyU69KEMrSvtHoD1qBHAY+9STfKeR9apAZYdua6qcboRcbaR8tSoRtqFE2pgY61MuMAEc1MlbQB/Lg+lQSrtOcVOp+RqJBlBnnNRF2YEAbDgY69DT2ztAX1qN2w/T6U+IEqVNW7biRLj5l9KkzkEflTEX5CT0FKCGGc81kxj2HAXvTWGDxTlODk0jcKQKnYBOdp9ab/AAk0v3TTW+5kU0BGTgkjpUEwzg9hVkkbeOopkiZWtIvUBIdsmG70vzY9BTY42jYipQMjnrQ3qA9MlcE8U4k4NNjOFwad16ms3uFxp7UzHNPzk4GKickkimtwAkAe1D8oCacnJOaa3JIFUtwEjPz4U8VKeuB0qvGoSTirBOPqaJbiF7AUnvSgYxmkH3j3qRkE6/JxTDkYGKsyLubnpVV8q2c1pF3Vg2A8tz0qMEBiO1SOTt6cmhI1PBPSqWiAkjPPsalPzNnHFQxgB8+lSg461EtwI+xx1NMc7QOOtOk+8NvSmkc9apeYAhHQU5h82aahCnpTzkkk/lR10AYXIXNV3YBsk59qlcsYzgc1CsTbckdauKQnctRcrkH60yPDMw9KWEYXk8d6apzOeOBU73C5MuTwaaMmXJp4GQaQDJyOc1IxwG6QDt3p+AT9KZ0k4PXrTxhWqGAzJDGmI3dqWUkE4qPIOcVaWgdR6sDKQSM0ityVxVa6byhuB+am285ckHk561ooXVw6ljCiT73JqRAEGPWoGO6XGKkMgJAFS07ATk/Jx1puMDFKCAv0oA4zWYCLgHFO7+gqMZDGnAlhTsFyKVSx4GMU1GJbbjinMevamKwzgZq1sBMow1K4wODzQo24zQR175qOoajSBgGlb1xyKAflxjmlIBXGelMBqN+lPUk8/pUa4z7VICOlJgOPem9sEUp6ZPU0N97ikgFwWxTTy/NOBySKYeDjNNBsNxkH0qFxyQO1SvnFRNgtgnr3q4iGDutNQA8A5p2D5ufakiGJCe1X0GS4OdvpSSd/WheHPpSOeuO9IWlyuFGSTySasx/MODVR5NmABVm3AwBnmrltcXUtAgCkYY4FJ1IHejqCawKRXni3NkfjUXkcirnUHNM2knmtIzaFYaqAE8U2UgU845NVrltwA700rsewqThjtx9KnPzDB71nKMMpFXUbgc8mqnG2wr3HbdnSnpyp9qQjIx1pQMLj1rMbECjPNQS4IYDpip+QMfnUEvyqc1cdWIgiOQTjJqcZKjtUMA5PHHpUw5BGKuW4IlU4G3tSnAUDvQvIHNL1FZMYEZxUZXnPcVLnOKYwJzQmFtAHTPanAADrTFOF57UqjjPehhuDDIPvTXACjsBT26+lRuuQc00FxnfGeKiwCxz3qUDA+tNbjJNWiQX5SajICg9cetPIGDSLwuD0qkMi35HeogpkBHSrXlDbj/IpsaYVs45qlJIXSwyBGDhgc561aBO7NMUfJx+NKDlcn9KiTuC0HE7c0xcgE9aRiAuP4u1MGdpVjx7UJDuSxnI6daUJgHtTfuvjPWpT8zYqWAwLkjNTDnpTQuOvBp4HGB0qWwSYhGCc01COlPz196iQ4Y0LYB44yc08ZxTOrc9KcD7fhSYxh5PTHbipBwMU0+v6UoHGaGIcoyOaZJjFPx+VMcnZg9DSW4yIgtjAqQdc03bgUuOM1bEtxR7U4YC01eCR2pWbavtUgNHU801h8wOaXjbgUyQ/LVJajsOB+b1FRtgt9KQkYAzz/Ko5QdgC5z61aQiZTil/hIBqJSSceneps49qTVgQnQkN3pMFQf5U4rucD0pWXPP8qVwIz8yn2rLnQCUntWpjjI/KqF0mGGBW1J6ksrluMZ4H60wg7c9OfyoCDfwKdhgR0GeldAhshwcknFIvrgCk25fnPHrSngnpyelPQHrsL1PBzjrimyAbHyeccU4rtwR+P401jhXzz8vFNIOp6nklefwpQflP1pBgZ5pyfd6cV8voaC5A/GnZAI54phUDBzkU4dAMc0NW3AcuWI9adgYOTyOaaPlpFJJNLzAdzjJFMOcnmlDZYjPFKQM5zSAVWOQMc0gHz89c075Vfr0oblsjtVPTVh1FI2PjIxQOGyPShQCNxpQvXJ4qVvcQEE5J6UidfYU3OBzS5G4HoKVrMe48Zzj9aXBBHGaB6gUgLAntgU2gAthDkcmoZUDbR19RTwRtJOaanOSe1O70FYV9oCqgprctzxSqQCe9Iehz1NJu5Qh+Y8HinZ3AjtSHGzihOMZHWpEKCCpFNOcr9OaVgOMdKcSA3vVIBmctkmndSADQEwCKEGFJxT8hjmIOcdabnjHOaCSrZpD1BqUrCFHpijb8360pOBSgHAB70IBMknpSEbgPWnYwD9KaOmO9VrfUBrnPB6+tOHKDJprcrg9acOQFqpbAKw+QetQORnnrU/RhUEg/eHFTFAOBww4qUZQ+9QZxj1qUEk805bDHEE0qAHr0FN5BJzTl+6W6A1K8gIJPvHbVWScJJ83WpbiXb06isyRWlXeeCK6aMLrUV7GnBKGUGrHToazImPyhccCtFCWPPTFTXpqL0HuSAAqSKM9c+lCnaR79aD8pyelYLuIjZivJp2dw4HSqVzdiKTbUttMZxkZAFaKnJxuGjLAO0daA24mk8v5s0kKkOxPTHAqIq4EVyN8ZFVbKRo5ijDjHWr0iZB96j8oBlOOnetaU0k4vqGtxX+/jsKr3HOABwauMvzn07VG0YKipTtICnFA8TDJOGpblH8xQCcA1cIJj56jpSFc/Meoq3Vd7gRbmaVVPbvVpPmbr0qEjEgbFS9Bkd6ipK4xytkkHoKN3UHvQcbRilPzc9MVmFhFIU7DzxQ7cY9KTkNn0pz9j607oRDIA8OO9NiBRFBqV4wGGOlNkGI++QKpO2gXI5Yy7gjgCnjBTB4xSod0WfXpTtvyH3oUnsAqfcY0EfKGP0pqjamO1KDjANK6GO24Ix3owGUk8Y/WkU8EH8KXdg4x1FKwDR92kblQR69KdxyPWgDABx0oSAcCCVx1FN480sf0o+8xxSgHaR3NO+oiGfruWo42yVz1NOkVhmmDAUEjla0jsA+TAVh69KoM224Wrso3whu4NZ90rAKyjLA8VdONmD0LxkCHFSZCsD2qrb7nh+blhU8gJAx1FTJJSsMkJzxTshyvoKqyOE25PFTQYzwc1PLbUBzECQZ6HpRnexXHFHVsnoKEwMkUk+4Ee3DYp38JA60rkKN3rSRtuBfoPSmr9AHH7g9aSMbT6E0fxcUrAls1ADsfNTThiQc04H5jmmEndz0NCAcpyCPypSPkBNIegxSv/AKsA9qGxCyZYjI5xRkOw9hSZIxznijG3nHNFxgvJ69KActkjkUDsT19KVud3HWn0Aa2cZoZuF2+lNZ8gChTuXgU9kBIcqo9+tLj5/bHSkyWUUo+bccdBUsBkiB9xNNVNqA08H5TnrSZwBx1oTaVgImPIOOAaVCN5OKf1yPWmFeNoOKpO+gD85bcO1PHzOOetRLuBIpyEZG7pQ9QElAKnB4FZ9xCdvycVddvnz2qJyTz1rSDlHYRizDZKDjrTirTAccVovbLIwJFSeQscZUDmuj26UbdQsQQgNAV6URhXBA6CkVXj3Y6GnW6HJJ6mk7P3rj3LIjBKk9u9NlyoOBmphym3vTGPO2ufmdwKkaInT73epjwhI69qQx7WPpSSMOo6CtG0wIQ+QeOe9QSFScZwe1WNxyDjC1BJFhixrWFk9SWOtysZ8vPXrRdEo67fSo48sw54FF0DkAD8TW1nz7h0L6FiqjHGO9NmQFVHYmmW8y+UAaSeUEBPXpWUoycwZOx+XAHahMLEc9KcihIhnr2qOQkYXsetY3bdkMikYFWYfhWaAZUyeq1qyoAvXjFULOPgg966aMrXExAh8olh1pbYEK2allIDBR1xzUYyj+1PmbWobMcqnO/P4VY3YYharAhZCD93tU8cZMue1TNpoEWduEUn86Vhu+tOwMAdqH4cAdq5b6jAHCnHWmAdiOTTmGCTSMcsT0FEfICg6Y3E9R0qSFt1v9O1LcDv602D5gUxxXRvHUkVWBOMY/Gmlsttzx/KpHUKPkxmo9x27gMGp0eoxilcFccU9MBRnr60wEHLVKgVwARzWj8xDw/PFPY/ujg8Go3bDBRwPWkZeozWVluMYmc8HipygkTHpUO0YABwf51PFwCDyfanJ9gFjRE3AdKzbudVZlI57GtTGFasK7Qm5c1VD3pag9BomkwPm4q7ZzFnxntWbtIGSMA96ntmCE8nPrXTOKsCZsnbuyOcUwHJqvDKd+CPl7c1P/Fj0rkceXcY84+7SgAkjpinfxDFMIOOOlQmDGkKw4phIC8dqcUJ6cAVGVwpBNWrAVLhSQOKoBTu3fma0p/uY7DrVQplAK6qcrIhrUIcluOhqyLdiCelNhACjPBq6pLIQozjrU1ZNPQaWhSMZwWpA277vpUrkpGyk1WjYdBzVRs2Mtw8Lg8n3qxEpJIqpCGlyAelWkJX5M1jN6jJWUEA80nIix1PrSbmztNBYJnd0rICv/qT7nvU8TtIuccfWozKksuBVhcIuB3qpPTVagiNwDH05pqgbTnk06TIHA4qLIwB60RV0BXuclFA+9T4H3oQfSo5G2kZHSlSQbSV61tb3RdSwjgptxTFOMADmm7lBAHDnvSxK2CWqdh3JHxkH1puQQpHUUxGDMRjpTg4DFaVgJQQFwvWnq2EORzUaAb+OlTH7x9Kzkw1Qq8gZOPag5HGeBTV+8WxS5+Q5NT1GJJjcO9OJwOKjcgc0qElckfhRYQucA5p/XGaYOvNPBAznrSYwPXGaa2COnAp2BtGOtNkQAcChbiKNwSTwTmoRFgj9amckSY/SpVUNz3ro5rIOpGobfjsP1qTeAxJ702RiDtA980jLuAGelF77h0JsjApxOQB6VB82dq9qeDlTk1Dj1Arzod+4H6VODtVT61BKxULngdxTynmopB4FaP4dREy/IMHkU8DPHX0pg5jBz0qQH5R61kxjsZYDtQSASMdKAMDNDgdAOtQAzPBpmcpj0NOYEd6Q4C545q1YLjCwXFKTlge1MdegJ+tDHgCqSAn6mo8dQRzTlG4DNJja1SA+PbtA7inscY4yajQgGpNxz04qZbhsIFG4+lQkbWY54qU+mKiYfNTiAiH5OTketKTtyRRxkDPFBCg4JqnuA3A4Oeaf1+opmBjPcU/HIzxQwHde1BIyKTOQcCj+H3FSgA4zjJxVdxg5PHtVk4xUEhzg/nVR3AYW4yKbu4JPBpzrlcUxsYPrWisBImd3HT1qVsFCKiQ8r6VKT14qHuBHn5SO9NCkJnvSDIbg4pxyeO9UxEY28knkVMpDRgn8KgwA3rU4OMKBjinJaAMZSBx0pkuQvFTvgL05qq7s74UcCiN2PQkTIXbimklD06+9PQ4TJ+9TXBZlp9QJ04T69qbGBk0q8KAaM7TjHJqHuApGTmpMDYPX1poGOvenAgk4qWG5HKOR7VFlVBNSSZAPFRnhRniri9AIbiMPIjdaesQjJCjmhlxID27VJ0Oavm2CxXaQow460yIqJCe5p0+QwKjnNNhi2uWx1q1blEXx0xQrEg+1GcCmNkDgVgtRg7bTyM0qk4zTW5UHp60RNwfanbQBXUBTmoYWDMQetTS8ZNQDjBXinHYCyuSSCelKR1NRk7Tn1pxO3HvU7gIen1px4Az0NNbk8Dihvm4J/CgOoKc844pQcE+9CrhQMYFIx+b60bgPHIBzTiflpoz0FKRhRmpARF2nce9I3y8nrTwcgZpGGW6cU766gRnLZqGXg9e9WCMDgVFIueaqL1AhJ5zQp4HFNdiAeMgdaNx4ArSwE2c+3FNYgAZoHYd8c0Fc9RnApdRMrS8EDHWnI2zpTnQyKD2HSmxoXkOa0urA2TxufMyTUwPOe1VHOyUL0xVlDuUY7VnJaXDoP6jIph9QKf0FMfK9O9QhjTgqQTVeVN2ak3Zc8dKaCXINaLQCsQQ2T0/nUo+Q5qZ1Gc4qJ/mbbirvcWw+BixOamHJz2qKKMqORT6zdrhcVqjfBbJ5NSGUCLjr2qm7sDVRTYbIl4VsjvQMZ6VApbJzUkRDD1I71TVgRYAwMD0px6ADj6UxT0p2eD6VmMU/d6cUw4I4NPzkE1E7fKSOaEIceRjPIo3AoQDzUSudw96VRgnmqsBKDlOaa3bBp2cjpTWJA96SHYTgfSo3ILbaGOFHrUZ+9z1PSrSFuPGCTg8UzJJxQuVzzz2oJxJnH1qgH79y59KarB1OB+NKoAckdKZEqncBS0DUezYXHvzT8AJUY547ipMfuhk80mGtyCRT5gIpT1wKQsTNg8YoVsscjkVeomTKPn9+9SucHjrUYGOc96lxnGazfcrYO+R0p/f3qMd8n6U/oOuahggyM4qM/K3SpB0Oaay5WgAAHrSg8mmc9DwPWnHrmmCAn5v6U8fNxULf6zj8aeDh8UNC6knGOOlMkwSBThkDPpTBk59alAJjjFGR0xQ+Dk96MZqgAZLc8e1EgBXFKOpNNc4WjqPSwwZ38dKc6gAnFNBORmnt2yaYiJUAb3oKbhjFSqAvUUoXinzARrHgA05ABnNPIyMY5pSnNS5AlYM96Y2Mj1x1pzZGBTGOG5OM9KEDIm4frVa74wcYFWZCzPxxUVwC8OK2g7NC3M3PX/HpSKQOcZ9qa5AOPQ96ASx/pXVYXQCw55OM+lLkYzgmmOpJ4B/KlTO4+tUIdnJ465602QjyznJOMc04E7s5ye3pTSf3Tdc4xzQgPUmGegqQH93jHOaNg2A/nTo1yTzgD1r5m7bNBOhAPegHLZxwKeuMtmjbggdfpU3aAbznnp2pkikLweanOP1qOTqO2aE+4DY8Acj8ak4Jzmm8DApxHy+9E97gBAZgR0oBydtAztGO3agH5unNADx1weKTr36UhOOv607IC1L0AQEdTzSgDPNIMKo4pGU56UIB4J3jHApc7n5IHbim9VyfSlAyu7H41TdtwGMPmI7Cmk4XA65p7/6rg8k80xQADnrSSDUUfKMAcmhhwDSqdrHApF5bGeKTVxiYwo9aQuS47YpX+97UHluBTuraAKcKetDLn5qUgHB46UindkelLuIQfMx9KepwCB3pgyXOOw5p3GBx0o2VwYoHylSO1Q9CanZslcelRSLgZ70Xs7AID0JHapAdxyenaoAxPBH51Ko4Aarem4xd3ODTN3zjtUgAzz0puFBNSmwEyPMzS5zk4+lHBA9c0HrgDikIMbufSoXP7wY6d6mIKnp2qI/rRHe4DJHCDNEEqzn5TzTLiMyRMF7jiodMhZNxbiuhRTi2w62NAnDVMynYPQVEfl6c5qVMlVJbj0rBdbAVZ4VPzfnWXcSiNtijkmtt8FiODmsua1Bl3EVtSlrqILX5kzjkVciyUye1VEHlKc9KvRsvk5HWir712hj2wzAqOgoI3D1NKvyrn1pV+X5h+tYX11Aq3NssnJWpIohFCNtPkJB56mmR7gxDHirTdmrgSHK/lSgBRu9aazcY9KikZ96Lj5TUpX0AlxnNMAJxnsahNwySshHHrT1kAAJ71Xs2ldBuSkggj0qORTtHP60q4J5HFO4aPBHINJt9QEbggjpimt1yBxTn+6B3xTBwpFJ7jFk6KRTgyuMg9OtA5WqksbxRHYec1Ss9BMupjdtPNOHOV71VspmZNx6g1aHzMXXg45olGzsAmeCpozuQgnp0pB0Yn1oTrkDpUhuBI4BHamuCWzjjFO6t04zTWOMehpttgBOFC4ozhAPehz0FIThsdsYFJMYpIwpFDnJUikAyDjoBQgz1Pamrh0HKQSKGyGJ/GkXgGntzEDSAhDZQnvmng8YpnRRgd6kK9/Wm1pdADDB/Wn/AHvriozgNtJ+lK4O5cH8qVkICvI96jkTauBzUnXGeg60cE4Paq5gKrMRhe3eq8zgA+lWJ/lU8VUiBkjIYc1tHXUNyyhWNAfWlUbgTxg1WuG2xin2z749vcUcmjkPyJJ1xGVA5FVLCX52API61fIycmqscS27kr3p05JJp9RdS6cYzRGAVBoQ7lwO/Sk6DHpXO+wxT8w2/lTV4wMUu7aMd6MhgKpK4CL1LZp54IPr1pin5jnpTyN2KLgKw3EYPamu3K4HPepCArZHQ0zG0k0Jq+oB1bIob7/1pRw3SggbhxUrVgKTnj0pCwJH60oALkdqaT+8AHSnYBRyzGjGUYigLhsZzQSTwBQ1YCvksc1LEQsfvS7BgimopOFqub3bASK20Z65o6HjOTTSM/LnpSluR7VG9gAj06UFsjbjignDH0NIRk0WsG41jg7vTtTScAOKew6rUbJzgdBVx7gPD5fNH3lJqKJ8yY7U9SSSo6dqpxtqwAADGaCAVbjnrTh9wjHNIVIXNRzMBgDADI5NI53OcUrPhRn8KY0qqD/eq1dsCIuSfmHFSoueRxULMGIGeWq0oxGat3SsMRMgk9qTGMkmlUjFI+OmKz6iI3JPOOKruuELH8qts2F2gVXnyIsL+NXEBhJ8te1NdQcY5NOBDRgdKbG219h7VutxEIYLKNo+WnXDrJswOPSnAqHIpZFAHzY68VVxCRKoQtjiniPzJQxHSmhv4V6VaUDaueKJStdgKGLH5u3FQXQ3D3XpU24BSxNVJJhJkDqKxgm2NhMcW+4tz/Kq9qCAc9PWmSlpUGPxFWUXbEB3rotyx1FfURl3y59O9ROrFmx3q2RtAqs5KyED8amMgK82S6mtG3JMQPes0yN5uDWhanEeDgmtJp8gItphsD9aQn5iaBgc9zSOQFJ71xPUoUEbxzTGJ3nPSo42PPHNPwfvH8qtKwmRTYfaAaZBiMnvTiuWYmo1ZgeBke1arawuo4AO3tQoyuB+dKpznHHHNCHapGaLWAhOEPHbrUf2g54qYkbiRVE4EpI49a1ppS0FsXPMXGc5JqdSSgJPWs5XIYCrikDC1NSOlx3HhhgYFPjYBSf1qFH37gO1SrgDA6ismgJfNAGB0qlcQq7Eng+tWJB6daq3L+XGW6inTXVAyi0bdeo7VZ8pFjBBqqkhyT29KsxD+8a6ZXsJIDKA429PSrIfIWq4hCsH7UjSMR8o+tQ4qWiK1NPeqqDmgMME5qtCfMiHpVhcFNoHFYyikxkayBt3OKbIPlz0qXYEOcVBc7imBSjZsRWYnJJ6elNUBmGB+NI42qM8HvTEl2AnHNdcY32EWyq7gD17VMp2kqDz3qupyQxqdQRyTz3rCSsBUuBh2yeO9CRAAKvOepp88ah92frTTMNvHStIv3VYCzbxrEWI71MuMk1SjuVI2ZqcylWwORWMoyvqPYmABbOKZMpbocU6NjnDVI+Nvy81GqY9EVo4tpIAyRU6YIO7g1DIxADDgVLbuHXg805a6iHEZBU1A6bmyO1WDndjtUDnnGcUoN30AhlQOTgfWqyxeW5GMA1oKuV6cnrUbR72NbRnYLFWMEuewp6S4lKk55pjxsmec+lAQr8+OfStLJ63C5O/yvwOtNJAwMd6k5YZI571HI37vOOlZpagSqQGHep93y4xzVaGQHDEjNWN3fHtWc4tPUBB1wKUrnAowT2ppDBvapDQSZtvHengnYBSeXu560bCyU2kFxQcHGOacnBJ9KbkseKUEnK5xUvYAB4JxSA5z7U7A6daQk5/lQhlST5pd2OlTAZYjvTX4JJHNKhJUnHWtHsJDd3UHqO9V1ZUkIPJNTyLhwWpjKvJbseKtPqGo9SThcc9zT+oPGKjL7W4HXvU38B9ahgU2Qs/K8UqKwBTt2pzlwM44peqqR+NaJuwW6j1b5fmqVeVHFN2howKcv3B7VkwJCOBzSNyAaGPApMjHWosAzGMk1EWJTn8BUp45NNfkgirTAjLZwSOaMDaDinNtwefxpP4KoCZWBQetNI5Ge9EJB4NKy5yBUbMGCHB6VJk7feoQcEA9qlJ9KGDEOCOnNRueQT19KlIHaopBzREBCQXwBTZMbsAcU/qN3amHkc1ewCegHQVLkY96YMUvBUZpMNh6ikXv1pA3zZxx2p2cUmg2EJBOe1RPyOmMVNxjNQtznFEQG9RjGabtyMZz70qnAPFC+taMQsQ9yakPNM/hpwHH1qXqNEeOcgcZp2DjcBzSEnpindBg0wK8hIckipU52nrTZRzhaWA/KBV300Aex+fpmmjBY4H40vUE9xULkqgVTyaS13AeAS+BzSMCZAOgFNjBVgW606UEFcHGafUCYMOlMUZPzU7PH86Dg9D0qNmA4/NgUuMZA6CmlgGpVOWpagK4yKgY5UVPJ7ck1CcEAYpxAaxAiyaNxWHI/CmyEqcY4phkyoUdPWtEtAIvvHPpVlFJwRUSrgKB071OPulQeachEqnK89RTWO2hMgCkYjBHU1klqMQnKHFNiOFJ9acGBRh2pseCuPSqQWHSHg55pm3CcVIRuWmM+w7T/8AqoV9gEY5XA/Onbug6UhAAxjmnyJ8oPpSbQgOR+NAyWIoPBB7U1WAJotoMlYfMB2oYZbrSHmlPtUgLnAzSHgetJkgc9qdkEZosALwAM0rAAYzSLyc05uetLqAwk9u9RMP4c5qbPzHsKhdgHzVIRXYjmmJ93mlb754pVJYEYrboMfjHPen9AOaYOTj0p3RdxqWF0HReetCEBs1FIxyM9ai3MGO37tUo3QroknJZ1wOTUyHjFQrKoGDwanGOppS2sCJAeDxTX96VCNvrUcjZHFQlqMgkfLEDI4p9ueueaiYnOO5qaPhc5rR7CHfLtOKZkKWJpQ2WNRzAnOPypJajJEfdz2qQck5qtC2WwelWhzjFElYVyo/ysQTULMAT61OQA7fzqBkwPpWsSbsahyeRTkJ80gcA0qfe+lKcCUDuatjRYRgDzxSsxZuOgqMkE57VIQB0/E1g11H6Cjk4qNsAVID8xagDcf6Ug1Iwvz57Uu35s4p68sRmgjBPoKdxhnnilfnrSbhihjkDFSBXmJP4VCTuPSrD/dz+tVyMKBitoisSfeIPX1pXGGyRSJ8x/pTnyVz3FGzB+ZGnyjOPyqSMYJYnrTRgpg9aFIdSo7GhhsOxgMc9O9OB/cj1oxtjIPelblAD+FSBDKOQ470rcrkcNSvnIXtQxxKoFULzFJJUDpU3pUTHkD8qk4JB5qWMcpyRT8fLn3pvGKUNxg1mA8AYzj6Ujj5R3pc/LimlvlNIZE/TNO6A0nVTzyKQMCKsB6jPNKVyc5piHt+tSA8VLEKCNtJ0FKOVxSnoP5UhjGGVxjFN7AAVI+eKb17U0xCgDHFRyDJqUDkcUhUZOaE7MOhDzkjtilB5waU8c0oT8/WquAgHzj0qRcmgAA0oBHNSwAjjNIeoNNckHk0hlA4xTSAczbifaoH5fPaldtuTnntTQ2MGqURDXJB9abL8sRzUhXdhjUUvzRH09BVrcbZkOD83PGaUA5yDz9aU8k5z17dqRDyQVA9xXZ0Is7IXIz/AEpgPzfSn7VBPOPakUYJ/LBoBoUAenSmucIw9qftUdx7imP91+e3ShbjXZHrO0liO2KQDIIHXNKuQpI70qgAV8vdWLEGQMZ5pynH3vSkUEEtikYjfwR0pa3AUfN06CkcggDHNOTKk5puCzZNNANOCQO+OlSkcZyMimKNrZ7U4YZiSOKcmgBQR3701xg5HXPNOHPQ8ZobIqVoAuCRyOtKF4PpS5bbgjigEHAA4pys3oA0ZJ+lOxkEsaQ/fI9fSgYVcEd6SQCopKElhxTeSpUU0kgkY7dKUHamfWqQBjOAM8daQqGkAB4pwwikHqaIwAT7CktwGk4OAe9IQQRjvTwowSeDSf0pB0G9OCBQgIB460uMfMaD0BXIxTv1GJn5Dz3oA44zmlC5Smq+wkcGk1roIVOpApy4UlT1poGG3dAaCd2D3p6W8wHjJ69KRjuxkd6XuAaQrng9qV2gEKdB26inYO4A9KQZPHU5zTsZ5PahyuGwjdQB0pOgPuelOIDe2KjTlyKerAVgOAKQjI47UobD88gUm7BI9aQXEJPemEcEmnngbSOaYykKaa0AQHINPReMVGBhacoyOOtNjJA3ABHIp8attPtScYwR81OTpjPJpIQMMspHTvTWUOxH40qsPmB4pDkLx1NLrqBUmQAkVFDMrDHQ1ZnI3L+RqkQFmOO1b09dGBoqQy7geaVctn09KjyAqlalb7oNYyjZsYyTLocHkVBAHVjuPPpVgjb05pG68U1LTQQ5hubA4o27h15FDYwp70ikg+2aQ2QeR5ituOD2pyRAYB7VLn5umAKTI3k9u1PmfcQLjJz+FIOmO2aH++CDxTgByDT5QI3JRMnpTEb5enWlk5Qg9KamNoPanFICX+EDvTo1BUjvSAYWkj4c1D1GORFBI6d6XO0/zoAyT60YLNj0pCEYY6dKMcLjoaCcjj1pTkqV9Oc0K49gTjOaYV3fSnAg8DpR9zGKAGY3EdaacEj2qQjCcd6bxtyevrVXuwQmMKWB/ClH3NpHWkzlAKc3RfTFP0EO5WMjilYkRgY57UxmOAAOaeSSFHp04pMBv/LL3oUEr64pxORtAwRQBhT60mMQgMVJphIpzEbR7GonYb8AdaEriHh+T6GlzggGmBTtz370MpLBs9KqyuMhnO99oNN8vYAwNK6DzNwqRhvUkHtWiatZCKjxFwc8io7fPmY6VdUDgGoLoeW4x1zwauE9bBtqT5wNtVpDlwp/CpUJB55zUNwGChl6g81EV71gLa8KMGl9T1B7VFE2Y8dxUgJIApSjqMG+9z0FG3kjsaZIw3hTxUq8kjt60rWjcBuCF49akxiLPvUQbLYp+7cpHp3pbhe49+BjrkU052gDrQOW+gpj53YHQdaS3sBJkhQcUvBbNM4454pQec/pQAoPz8UhGCfQ0e9OcYIzzU7ANUZJweaQH5hxSgDk/pTcZwT+VUtQHqOCaBjYTSEndjrQTk7e1IBMjqKToxJ5FOHBKgcdKYw6gcE1XYBQetNB4PNKo+UjHNMHD4psCTGFz3zSMCinvmnbTsz60kh+UDuKQhoAXGByabkhzx9Ke5GQOhFIx+fPaqT6D6hnauaVs7MnpTTgpyeaGPAGeKVtQIpEJO7P4VDNbFyGDc1O2Q3HSnMemKtSa2DoVWgKMj+lWd+5iBwMU6Tnj16Uxh8y88U7uS1AM/KMetPIBUNnvTQAOF7U4j5R6VD1AZ3PrUQBO8HpU54fGO1QyHafTPWqiMhO0IAB0qF3UKTjpUjsF+Q9TUZUuxU9K3hbqTqQq5AyRzUxy0OD94UxFOfXHapeCWYHtWjkmLcjhAMh45FWWkKgA9RVa35feOtTy427u4qZavUFsPkG58KeGqg8LW4JJxmrUbEnI6VXuX3uA33KIX5tAGxZQAdm71YUMXyelQlwMqBnHSnxyb1HrVTvLUNiWRsR571CQAw461acApiq7jADd6iDswZSmX/SAPSrNq7Byp6Go7hgHB7k1Kpx8w610u/IBoBh16YoJ381Ci/Jmph/q89xXA12KKySZuj6VO65B9DVUN++BPerRq5XSQlqVZ5CgAAzmnQtztI5NJKQuM80yCYSPyuAKuLvHYWzJpVwSAMe9RJgKQOeasSA4znn6VAkZBqYtWB7jNoY5IqjNjJOMLV2T72DVKYoTszg10Ut7kuwoXgH061ZUgoDn5qqq5ZMA9Ks2/A+cA1dTzGrdCWDKq3r6VOvyjPekQdx0pzDjPp2rlbTdxq9hCB+VVpB+5IJ+lWM5YU25OBgd6IjMcfOMqMAHrVmIBo/emiRYVxt61KgBjyDgk9a629PISHRhypB6UmxQnIp/BXg/gKb1J7mso6sNie32lNq1MuOVB4FRwDCkAfjU+0AAA896xk1coUjnB5FV5ASenFT9V60xzhTnmpixGdMoLkZqqELHjoKtMMMT096rgtuNdkHZaElu3AKgDGRUrkKxyeKgjbZnFPJ4APNZSV5alPYNok+cmqrtgFSvTpVoDdtT86W6iUqCO3anGVnYXQzYgRID6VsoBsGRyaqQWwb5wMYq8FwQB0pVZqWg0NYbTSrxx2oI5+lRDdvJ7VluGwpCv8ALU0MYiU4pkILEcc1PjqB0FKbtoCI3PG6oCBwfWrLjcu2oDwhA60RYC/eOAaTIBIHWhcbsCkYDfuNV1ARkxnPWkxgjJoGXO7PWnEZK5HIqldagObAAxySKrTKTEc8H1qzt53ZNQSOGPtSj5AyCNWyMnFXCcEY6VGnAGTUinPaqm7isyT+Ie1GeST0FGDvNNzngnNZWGKSAPrQhOCB0o4202M4Jz0oAcSQpxSo2R05pHB6ZpoOGwBQBITxikdtpHvTlO7k8YpknI45NIHoR3ByvFNTIUA9KV9xXkcelImRwTWkdEA+TAb5ulQSkZIznNSswB65qCRhu+lOKAlUfLzUqnIORiqwuIwgBPIqwGBGRSmvIBkmWGKagwntTpd3G00gIKY701sA9Bhc4qRANtRhj5Z9acOABUMLC9vpUbZPQ4pzsBx3pA/OB0NCQCkjd0pGyD9KecAgnvTG+9xQgIyADS9Rn0puBnnr/KlGce1X0AerZIApzk5wDUScMABzUj8kY69KmSswIlPztnmrCsDzUBADZAwanAGOKJWAM/Nj0prdSDTjkd+TTSP/AK5qUAzI59KTGVJ7UuBjA700ZIOTxVgwJwcrTkGc5FQ78yAGpwRk/wAqcloAAnacDilLZHHUUnRvc9qV+DUgxD05pjfKSKaz7ifam55GetUogJyRQOQaaOMgf/qp6g4q+gh/YU45PHp6U0jrQh2rkd6zGIwy/FPYE8UnINKDg5NADHB6jmmRKcE1MOPrQo9apSsgInDKCR19KihRmJLdqsOQowajVscY/Gqi3YBSAWBxxTX5GadggfjSSqSmB1o0AcnTnpSLjefT0oT7opSABmp2YDhjOcUAHB7etCnaCSOtKueaQCDIXr9DUeP0qX2HaoyQeaaAjcFjtHGagSI78P2qz0fIHJqI484N09a0ixCMcOKnAGQf8moWK7+amj689KUnoPUeCA4GOBTJSVYkCnIcvyMUy5XtnFStxDRnnsKWLODikiB3cnpUi53e9NjFJ46VA5w/Az71ZIyfeon7+tKICZ3N04qYnIFV1bvzkVMpynvSkgCXnA70xcbsUpBIoA+YEU1sIdkZ5GOadnJOKZJyaccYyPzpWAbnkjuaeORyOlNA+bPenA0DFQ4bint2xTVPzfSlPr1qGA3PBNRMAByOtSDkGopASoIq0IgcckmmqePpUjZ2GoR1x29a2QEvQLjrQefloHT+lBIxnv61IIZMR8o70oUYORxUcpyelKrbl69e1aJe6LQSSME/L+dOVyMIefenKeenWo5jt5Bx74pLXRh5k8RyCelLjJx2HSoEmBH0qRXyOTzUuLTHsRmMtNntT+jc9DT14JJprDg0XAEYbulK4+Q7qjhyGJI5FSOfkIPFGzAqo21xV0EFj9KgEY4NSLjrTk0xoik+V9w6VC3arbjcu2opI1VCR0pxkIrxHDYzUrJl1b0FQIoV85q0i/KR1rVuyuKy2EA3En86eucULtUE5+tOXBQ5rFsoUAgmjoRjrSDGcDqac2c/TpUi2EGNx7mhhwcjNJxupckdaA0sRH5R/KlBJUDPFKy4NMGSx46dKvcAx8pFQuM/lUoOD0qKQ46ciqW4Ap+fAp5yvGKiX72e9OOcEd6bF5gPlYknrTgAHG3qaAm4YJyaGHPvQMUnL7SOKGxwmelIzcgd6CQZOlILBIpOAe1JIudpBzS5DAg8Ghj0x0HensHQcceYAe1OBO8+lM3hZfrUowN3FSwFX7+SacMF+elMUY46Gn7cLUMEOHJ9vWkJ4waXovApNu4dKQDCOMA1A2c8dKtFeRUbAAkflVRYESOd2D361Y3f/qqvHGA596sgY+nanKwIXilB4pg+YfKKXIJxUWAf/D/Wm4PbinduO1GOcUAJ6dqDyOBS5HOOcUA4HSkNjcccfpRgK1OJAxxUROXamtREm4HIxThxwelQqTye1SBw5HFDQCTKN2cVE+0pn1qWRtoPrVJ5cDiqgm0PcbKctknpSqwYYBzVCWYknFSW8mH5NdPs3Yi5odRjsKY7bV6UB8r1prcg461klZlWMtkG8gE5z9KjfIJx09qnmwWOPX86gLHGcYJ6H0rsjqK44ADrz3xQc5696QdR696djHuKCVuIO4wOOlI4/dnPGRTsqrAc0jY8puOcGhPUOp61ngrikOQMd6A2XOaaT82e3vXyy2NAbcBikPDAGhSxJJ7d6XYSxNUA8/Oc44FLngnFIpwp9elIM4PFJPoAgPy+9OBIU5HBpHGSOwpSD6+9GwCjhMZ5pTz+FIRkA5yacD2x+NJIBAfXrRuGOOtDgfU+1OTaFxjmq33ANvAPc9aGADKN2aFwRye9NcDcNv50Xa1AQkGQ84xSMwJwOaSQDbxS4Cxg96ndMBZCBgkcigMNxNRsxcj2FInKnIqktQJARznNCnORUchyqj3qVGCjpkkYpbAMYkj6dqUZI9u9Nx8xz3pwO0YIqbKwDmHQCoNuXz71M5G3IPWgjCg0LTRgxoO4gNSjAY4GcUYAXgc0A7WzTW4xQxYk0oJIoAIGexNOb5QD0zS1Qhv3QGH5UvHPOKU4PFRuQD9KPIBccE5po5BOOaX7wAzzUZfbnmmkA/hhkjmgfMeetMVgygg1IeoI4p2a3GJnkg9egphyMZNO3Enkc0npkc0t2LYYpyTTh94YFRA/vTk8VIgyx7elbNaXAm3bHyRzQx+fI4oQCRgDS5G8rxwetZdAI8YlyeRTgCvPbtTmUKPc0S4OMelN3sBBcDbgjvVRoyXyKt4Ofn9KJFCAMOaqMktgIwf3HvT1ZvKFVfMK5A71PAHz83INXKN9Rk+PkAPWkYHaOxpW+/3prtk8VjuxCkEhfbrQ/HIpyg+WPc0xgScGhxe4Di4/i9KaTwBSDknPao1YhiD07VfKmBK3y4oHzMBRKMLxyai3HbnuKSV9Bjm+YN7UR4yFPSmg4GfWlJxg09LASKMvgUKf3nTpQOACOvajq3HWpasA/PzkmhDgkiiQ9CKQ9RzxUq4huCGzj8KcMkM350A8ncOMUgyI2IpgJgcsOlDAFcjpUaMWU4p8ZOOaaigFbhdvY96ax4ApcHIzQ3DAHp60LUBF5IzwDTQSJMdqe/I+lGBsGetCdgAD58ntTsncCKbgYXH407lGApAOGDIc0i8E+9J1b3ofg7aV3cBGKiPPfNVZHGd1WpANgA4qi42MM1pCzYF1WO3kZBFI5wMUA/IAKWX5kGOtJppjZUkkCuBmpDjbxUEyHeOORSxMc88iteVOIEqnJ6UTRrIQPSlXr7GlPD4HTtUO6EQ8A8dRRKcRA04DEuT3pko3qQe1NDIrd97sM1bGc8dqp2kQic579Ku42dOpqqqV9AWxDLtdwQO9Shtq47VXdSjHFLGSBg0vs6MCQHGSaev3Mioidyg4qReAAO9JrTUBwPy596XcNx44pNu0Y96Vuowam/YAXHPpQpxn3oHyg8c0inJb1pWbAUcsRS5469OlIuQDzSgDHNDAROQWo6kA0JjaR0oBB/woYD1xz60i5C5NJnGQO9HPQ1L1ACMD3pH7ccilbAXFDEBc00AzJEgxSSDnI657UiNlQ3vUjNubOKp6AOUE4GeKacZ5pR97rxUbt60rAxMhpMnpQ/LYHamvhVHPalZxkfStOTqA4jLChsY296asg2k8Uhcldx4IqUgHMOw/GmY5OOlODlsk0o4B54p7AIzdO9Mc8YHWnKu1cn8KYrcZbqaaV3oFwjIP3uKeTxkjgVHKuOlO3gIAacrMBFYuN36U2Rdxx6VIfug0xvmII70luBXmiHnrURRlmbPSrbsu4HtVeblyfXtWsW7WAjQIG9fehMEuoP41XBMZJFWIPvbu/cV0ctldCIwxSURgZB61aK5TYTVZgBJuHXNW87iO5pSfuoSEVCsezvUbwho8ZqSYgsOcZprMIwAOcVnZp3Q+hUli8ssw6Gnxx4iU+vWpnjEkeT3o8v8A0cgVpKpdWYrEoXjIyeKjdSy59KkjYKAKjlyF9qxi7SGU54/3YLdaICfKyBkUlxJnauPr7URuAcDqa6nflJW5oRqBHx3qUACPrUUfMNPX5geea45IsrSKd6kDODVpu2O9VXJDEZqZDuAyatq6EQ3cYIGOtNSMKNwPWpZwSDgGktk/d4P3u9VFvluJrUmY7l+lRKTg5FThMD29ars2ckdqzT7DuRSAHtVSRdw3EfN296tyE7SAMnFVWDBcvXRS03ZLIo8t7VbWMrgkiq0R3PirUfPD9a0qXWwFpSABikbDH2oQDb0we9IR1IP0rl6lEco3rx0FK4LICRStwMdc08jdF6GnrENzHkb94wYd6nUho+eBUk4HQAbjUcS7UO/nFdKknESJ4vLDbR19aeqAMcdKhgGGz2qygDE4rKVkxkkS7c5qQ5AB600HBFOfhBmud6sYNwufWopiNuKerDb1/Cony3OBxVxi7i6FSXg4qu4G4DHWrDAk81GsRJJY5rpjZWEPUYAyP1pHKphzTlwTjJps0BbhuPek2uYTHRDe28dBU28M+3PSmBTHBsHekU+XnPUjipfvFdCwpwxUDk05uoANNTO3J60pwwweD2rG2oD8YJBGah6ZPpTuQfmqtLIyMUHeqgrgSpMq9+tWlJODjFZccT7s5rTjJMILdR1oqRSWgIVwAD71WcAHk1a2571VmADZ9KiG9gEVtq4PU0j4aLGKQHJJ9KegB6961egDIfmU9sUowO/1pyA4IxgVGzBTjtQruQa2JGJKkL1qDlkJ9KmWQAjA4NIFwSex7U1eL1Dcjj+bhu1TqOMA1XQ4Y1YjPPtSmCHFsEL3owB160hXnJp64K5rMBG5UetIowc09ThqQDL+1IHuKR8maixyMfnUmQVIPFNHTHemgHYCrzSA4BIpw+YY9KaeRwOlIdyI7scil2gDmnMegFI2AuSeBVXYiGXEZC1mXEw3nrjNaUwDgMOtUmtwTuK810U2luJ3K68NkfrV+3ZicZ4qq8YRhnjNWoCdwParqaoEXGAIA9qrqnluRU568Hiog2ZSc/SueNx2Jh93gUvGajh5Y5p7HHX1qXa4+ghOGJpgIDgE5NKxy+B0pjRbn396a8xE5IzTSMU49vSm9cnvUoGRsoByPxo3En2pHPOR2oQcZq0G4/7sgOaewBJNQnO4MenpU7Llc0mBCeX6damXGAP1qJwRg56U5W4xjk9qTAkx82aQ02Q7TgUHlenNKwDN2APrQSCpUUAfLwOaYoIyaoCMMfM9vWplOWxTei7u9JGwLE45q90FyZsebj0pHQsQc0oGcE96exJwKzuBCUIJ4pCAOlO3FmPtRtAXmqBkQADkY5pCSH9qccB800cOciqQWHsSaF4OO1KPU0DG4VIDt3AAo4Kim5yM+lKMHpRYNAbp9KTJCikZhk4oJBNNIWg2Y5Ge4qBJDuJbgVZYArgVWMDF+enpVws9AJPM3OMdKlkHyYFMhTGT+VPYZjOTxUu19B7jU+ZR6CnY+U+lJFjyiaXOVpPcBcnANIpzml7D1pGHAxSAUHn2qJyMEVIB8tMbBQ56imtwYZAQ849Kq5JlKjp71OqhlyegqvGu6YntW0VuLYdsAfd+FTITn1phQpwT9KfESWGKUthkyjGcUyYZXJ6inI2XPNJMN0ZAP41n1AqRSZcqOtWosA89aoRKUk9KtgncCK0muwJlnp7CoZAF59OtSk5HB4qC6OEJHaso72D1GJyR7mrKjA5qlb/OCT2q0GyKua1sJDiPn60mfmxnilABJ9aToTUooGBJ/nTd2D7UoPymo2O3FCFqTjJbIpcAdaYhyuacvXNJhuPTgnNKcsaj53Y709WGamwDf4yO1RyKAM9qCSMkUpyRiqAiIwhGaqKDvyOlXSm5cZqqVw4Fawe4iTPFNZgwAHQUmDzntTD0IppBcQHLFgKcSB0/OkixnrTSCzEenerFsLvOOlMdtwwRkUx92DjnPapYItxBPbpT+FXDccQQvHelhQlvpSuwU4qWEYzUN6APxgYFNZcEA1IOO1NwN/NZlELKUkBB4pkjHGKs4yduarXAO3Heri7sTFWQ/d7GpW+U5xVVAcqTmrLHgelOSsCHMeOlIeUxSEjAPelHdT0qQK8ygMuO9O3AZA/WkmGSc0gTDggVottQe4/O3AxUg4FV5G+fB6dhU0b7j7VMloBJxmjjPrikK5bPYUAnOTUDG8A59KXrzRnmhTgHjmmAMQFz1prD5c+tKxH4UpI200IhbGzPpUUhG0jvUx5BFVZeXGK0jqxPYcQF78Uqj585qNjubNTIB941T0QdSTbznNIcklhSjGPrSKx5rMZE4yd3OPanZy2R6U7AkUjGKjGFUjjNUhbjowCc8ZPahFySKEG07gOTSjh2Paga0QjACRRnNTk4qJTuJPcU5G3rnvUyAkI96fio1+Yg/oKeDnpUMY44pev0FN555pRwT/WkLoOGOaiccZxz6VKRyOaa3TrwO1JDIFbg8c9aeGyo5qOT5WORQGyOvNaNX1ESrkCne9IpyCKXHQVIEq8/WlA59aaBhsDvUjHDcfnWbBDCuAR700jgilz1NJkEcUw0GtngDpVdj8+3GRU74A5/KmABvmIx71otFcBFBKn2qTzAgBqJ5UjA5qlNckgkVSg5BsiSa4G4k/lVN5PMzjNJgkjPIbrT1X5mGBXRGKiSQKpySx/SnqoTOB7mp1jGTuGfenBFzkcfSm5hYcrYwBStkYGeKaMc9h60x2IA6f41na7HqVXGXYEVGRzgng1Kx5O38jTEUN14INdCECg5yTyOlOwDkDp60g+UHnJobcQRnigXmBz2ySPWmOMRHsSKeTkYPy464psn+rY4zximtxq56uU+XPrShAAMmnEcAdaRhwMnpXy/kWNIK49KCeSO1KDlfp3pdoCZo1egADtU+tCnjNAJx0zmkCsTjtik9QDOTyKUnoCMUhznFOOSR0zRfQBQvBOaUYYdecdKTJHFICN2KEluA5Uzyx6Ucg89DR8zEY4pApYk+lPW2iAOVH4Un3cZ7ilKkjcO3WhyrJ70nsBGx2jHvSPksPTHSlYZHTikUbsijpYByn5TkdqkQIISOMmoyflKAdKOeKIthYVVDcHpQq9e3NGcHFBB28VKd2A1ievvQz/MpoYfKBTJDkrjvVrXYB6kbuacchj6DrTWG1R69TUi8ofXvSauAw8sSOgqN87lI9eakUnBApSMYzQDFz2BHNKTlfcUi9enPbNKgHmEtS8gFXAGCOahPL5IqVz82QKiJOcEU7ajQp+8PSo3XcSKfkbqRsBwPWiOmoDIl27h2FPU5+gpXGG3DoRSD7nTinLcVwb5WDUM2efWlx8metJj5ASaVmMYRlzS9MHPSlB5OOtCjdkEVaYEnAAK96UDILH0pqc5Q+tOHyg9xU31EGQxHPNNcEEU1siXB6HvUsgw3PNDVldgV5JMkDpih+YSD+FJKvzE96jdiQMdfStLXtYOg1UIbJHarSgeV7jmqyPsT5+1SqSVLDpTabAlXkZY1WaQI2D36VY/hHuelN2qTtYc1CVnqA5STFjPeggZ5oUbDtoYc80m9QHBRhqYoxzjkVIpxz1zUZBOfepuArEbge1NZQzN6+1PXqFpmMMad2BC33gD2p7fnTWGWye1OTkEjpjmrSuBIgyME06MZY1GOoIpykBs9j1qemowOcYHelLY+UjmlwdwA9eKCcuOKnRC2AgfL60dyDwDSEE8+nejO5h6d6LgRlCGPb2p5A2gCnEfP1zikA3NkdKHqMQDC4PWh+MY5yKQ/Mfel68du1CYhGBAyKG+8D6jpSknO0018hue1NPRgDgqRjvUobByRmmEZIbOQKX7y59KV+wAThs8j0pH+Zs56UPnYM9jSA/L/tVVgHABmBPSoHiBYt3FTK3GPWmmMlsY5oW9wGRH5Sh//VTgccc0hTEnHNOU/NjvVSd3oBDKM81ApGKnkJ5U96rKvz471pBMNiVGyQKep5OajjUbyc1MpHORSkhkRBLc08qCuaUcmlA5x61mGgImQMDpUpBPbpTUyjbcVKQcgEUpNgVZkIyzDioDuf29KuXH+rwfWs9id/HStYRuIkU/KF71LGdwOe1QAhmyOwzTo32nHY1ck+oyyT8/tS9c46io2ODkU6QbRnpkVjYQ7lmz1oUZbOaEbCjnqKcg3KR0IoafQYwcZ9Kc3CDvTT12jtS5IBFK9wG5yvFOxlPemngDjilbAAAp7gKRgAjrTj8wJ70wfe96X5snFIAJyoBPSmEjbtbp60/jbTDgrg0IBkeFBXnHapSeQT2qEY346YqQSBhg/nVyu9wFY/NxwKimYAcVHuJZiD0pCTkk8iqUUAwuQu096Ad7j09aEB3HPNTIoBzirukA7YBkjpUbeijIqcHnkcU1lAYenes1uA0HKH1oXcqnNKAFcnHFIzZf2py1AcG3jZTVHJjPOOaMYfIP4U5FHmk55NGwDSp8w0yYHIYdKsY+bFRsmQxByKSbDYrmcDCkVLnLZqhKrD5s8DtT4rncu0/hW7pX23C5ZdBIpOe9MlG1VNISVTC802TLICOuahIRBLERg+tJCCJckYzU3DR7W+92piAlic8Ctud8tgtqRthJck1aQBgHU1WYfKzkZOasW53R8DgdRTndxuJDZBiQtk0nmB58Y5xzUqgNk+lRKP37cc1MGr6jJ+cdM0hyUpplIIGOKcRjoeDUWdtQG46N6Ukw3qMU5CdxU0PwCMcUk9QKZUMSBVeEDz+Tz3q1gKhPQnpUMSqkxXHWuqMrxsTuX42GcdqQLtJweKSEBQeelOVsoa52ncorycAv+lWIiCoPQe1QSZ2OTjFS2n7yMk8VTu43Ac43dKjjcq+3+GpG+VTj8aiLLGy9zUxEWjkx4HWq+0mL/aNSBy/T9KDuyG6CklYbKz/IpPeqcheX6CrV0W2lqohjtPzV00VfUlvoEQwx9e1WEJLDJ5FUmcgnpn1q3ZMd3NayTeoI0UwqfWmkY79eRTv4sdqCOhPSuFaFDed1PI+TA61G+4uCOnanqzGPJpsChNncRikjBZCCaLo4fJPSoQ5Y8Hg9q3ivdBXLpGIwE60K5D4PWqyl9uVJyO1WYxvAz96k4pbgWYySOaUjcOTUasRwTxTnZtmFH41g9xkaggmjkKTUqqFHvimOOcD8RTvqIrkZXJ4FR7jjnp2qeRSDtqpKPyFaxswsPhRjl88CmSXTNIUXmlilxGV7VWX75Y8GtVG71J2NGNDIgYmkkBBz6U+2I8kn0pJHXqe1YczUh62Hhj5eT1pBySCaEVmXOaeBUt6juOYZGKjaFNwbHNSZ49c0OCVBzipTaAhTAcrjip88EA8VXjcedip+OQKqQIePuAGoJVJz71YH3ajcE1EXqBXwQuaXODnPSmNJg46c0ON2CD+FbLzAkwNoA6GoZCMEDrUjcqMDimSxnGQeaI76hYI8HBHUU9WyzA9TTYkKfWhQASe9NtXDqIQF+lPRgR9Ka+CcUR8kqQcdzS3DYn3blI7Uo4HpTUwFPFKMNtFZBuLkUjHD0hOTgdKRvu/SmgJDgnikOSevNITjHHWhjgj+dKwDh0AoJwooA70gwc57ULUBp6E+lRAkK2elTkYUjvUeB0NUgGYyooZAOKeg5xnpTXxndVXdwIZoVdQKdHGI0C9aeMk5HOaXO007u1g0uJJgN1qqS6ONp+tWZGBJNVkYtLnHNVBdwLigKM+tEh4GDUKyZl2/pUhBZsgdKhrW4CN94djTlORmmSglgc0+M5AOMADpRbS4DuTn0pPurT+1R9qlANkHAI71EgcAgnipTlsZpr8HirWmgxcZOKkLY4qLd8owacRn60hA3KelNhB8w7qcCSOaaoJbINHQCXO+Slbk03IU073zUgM65A/Oo3JB4NSqMkntUZAyTVqyDYZk4460qnLH5elKoAzgcGlUhGqtgJU5XJo+YilHKmlA+XI61kBGoAOD1pSAc0gxnNOBGTVPQCuT8/SkfJansvzE9qZHnBOatdwF53Yp0fXGaaD3JyaVPvZpMB3U46CgZGfSgjj3FAO7A9KQDcYY04cdBScFjg0uTkAU9wFj68jrSnl+KMYqN5ApznvS3YCMwUkE01n3/KOlQTtmTIPFPt8sp9a05bK4FpQAgoxwM96bGeSvbtUjHIA7Vm1ZgR9ueaU8qKjZ8SY60jOQBxTswJW4AqrLKF7nFWM5UEGqF3k8L941dON3qJkytuGM8HtTIzsl202NW2c9acP9aCK02Yug6fLZHpSwnHI/E1Iybl+o61DGMORnge1JO6sO12WlGHz60E8GmocjB/CnFvkPFZMe5CYwwyeop5AKgCk5BHvT9u1qbYAj4TJpXO9OtAALGlPvU9QKqDZlR071KpBYA1H5RWQt2NSKg3+9aOzEtBf4mNAbk4p5XJNNZORU3QyPz93ygfWopHbIHaphGEckHOaSUcjFUmriHRyZIA6GpM7e9MUcn1pxB79Kh7jELEA+9Kr8AHtUbNg4pp3DJPSqsBMMGQgUEZBxUEUnzYFTsDtHrSaswIy2EOKrkncGPfrUsuQAB+dRbC4x0HrWkdri6Egxk5qHsTUzMRxUBHzEZ604q4bbjEOZMHp6VN1c+lQMMEHvilUszA549Ku3UTHv8rYqaHkcDFNjXcCevvT4+DhRxUN3VgEeMSEVKgCjFKqkPk0uOfYVm2VZhnBOKXHXFA9+tIAevrUgNx3o2gjJHNBJNC5I61QDdoAFIw4NSY/+tTTnPpQmIZj9KRSSeaDkGmg5+varSGK6Z60hXIx7U/HAqNnxnt7UK4iu+S4qWNQuPyoPTNPTjrxmrb0FYlY9hS9OvemrjOc0McVlYfQVm+fFMY4IFGctkHig/OSM/hiqsNsH6ZpyY2U3BwQRxTlGFx+tD0FYjZvl461UdCT14q02dvHFM28ZPari7BuVgu2XB6VZjGF5qFuSDUqtxgHmqlsIX1IJ+tNUZTNIVYKVJ60mSoAx1FJIB4bKEjr60gAwWNOJCADsaaQM89O4oDqAYbhzwOopUxuPfNMcKi4p0XB6Y9KH5AgDEbsCnxgRr05phPzbf8mpQM8EdKT2Cw8AKM+tOXkfWmqeBmnIeDWbKHYyaTdlqUE9qTkCkIf2pD93j8qTPcflSFu9IZDKuSciolbbIB29afO5HzYzUMbB33fpW8dULqXFwWHrUn3uKiiPOKlU881jIBeh9aXzMnFKBzUMgYMDikkmx3JG+tMJ9eTRv3NjOKVzgE07CE4I5qCedYUNMmuRHwAcmqDl5iMn61tCnfV7BewyR2d8k8HvSKu88Z21KEbAyanjjCrjkVu5JLQmwyOLbznmpAMZOacy/L6U0AkY7Vle40IWwMmlJwDk0nOcU2Rgg+tNILjXO19o79ajmHABGD656U+MEsfeo5ychR39a0S1EmRD5uPSnrGqrnv35psaEsMdB1qfyeP8KptJbjexAzZBA60hODwKeVKcL0zzTAPmyenWmrE2GggHk9KRiRE3PUHtSliGyDn0pJMeW4BGNvFUtx3R60cFvTFEgG4AHg9aNu3IbrSeWNpLd+lfL7blh/FgdKGBPfihclODk0pOAAepqnYBcjZjFIvTJppOcjt604AGL6GpSuwAH1/OkB+bPamB8E09SMD1NDWvugDfMM5FCKApbv8AWnHBbHTFNVdxbGapWEOJO0Y608cRnHU9qjAOMdxTgSp6VN0Ao6FR6U0AAGjeQ2cfWmk5fPepb7DQhBwCelAHU54oLnaaaudtKwwySxNNYkDK07o1IKaYhV65JpwPXmmZwSRSKeaVuox33uO1JtBbHpThgHB60i/f54PemriEzlsehp5yCcU1hjjvT+oGeuKd9BijCjmhPmakUnp2zTuAwI4peggLc9MUKQOTTWIHfilC9D2PalbqA7jjmmsAWJ5x2oKgdDSsBsB/OnZgR7eeP0pvl/MST0qUAAjBpjnGRTuwEBzkHpSdNwoK4GQaXGV3HrQwAHC4NCrnOe1Bw230ApeC557cUWuw2GH72e1IW4zSsNoI6Um3kA1UX7tgJAMR7u9OUrsOetNQZQ+1G3dgjtS2YDsfIWPUdKceFViOaavRgTQeFGfWh7ARStmcEfjUMqFjxnipyuSCBSSDLZHerUhIrkhm2nqBTHZk69DT5FO3I+8KQxl4wrVd1cYJOQAKlQs0hLdKrxwsrbW6VYKMpOD9KJW3AkP3hThzkd6hhLkfP17VKo5zWUkkxig/IR3pSOg9KQZK59+uKVhnvxiptroIaPl5PXtQecH3oxnAzQRiQU+oDGUD8aah2RnBqR/4vWq6H7wPergBIGJHvQAfxNIvAx05p7ABMd+9JoY/IEYPfuKUHKk96YvKetPQ8H1qOoAPuMDSK4AwRSr33elRrhiQaaAeDk5/OkDbTkVECyS9PlqZOhI7U3F2TEGAWpvJPFIrncacuc49aTjYAJ655OaRiCeetLIMNkHPpTEGd27mnYNR2SBxyacnIODzSIqkEHr2oHy5zSAQ5NRh8nbmpWXHOeahEeH3etNeYEqY3nnntSocsck/Wm42nP5U9ui4HJpDGYwcjp600LucntUxC+Vjv3FR/dU4zzVLQRDIOQajIHLZ5FSspGcjtxUJOMitI36gCDILd6lDZAI/Ko14GB0p4BPbipluUP6EntSbv4qXnJwKQAdO1ShD0yQSalRhtO48ioYyRkHvxUqjB2nvQ9NgGyjdGB2Peqfl4bpxV9/v4HaopAMFqIyshGeVKg4oTlsCpZxgKRn3qBHAlGO1dMNUBZZjtGO1SfeUZqHcPM57mpSAOaya0HqOQdWHTtSqxyx6H0pR8q4HSm4+bPrUaX1AUgfe9aMHbn1pB1AFJkg7e1AbiscijOGB7Up4XFI2RgAUkAoxv5pVzkkZ4pinLdaen38etDYCY5pr4OMDJ708HBwe9NGFfnpQtAGsAGGPxpGiDZFLn5+lKOSAOtUm0G5UVHTikCkLtPQ1ZcAt6etBQ5AJ4NXzAIqhU+tG5VWhuBtPSmAZHNTa4dSSI7kIIwc0S5ZR6Co4mw2M9Km++rDHNNoOgxyAoI60xuGGOlSFP3ZXv2NV2JC+4PAqlYL6kvfI61IACd2KhQhgD7VPnC4FTIYmSCTimB85qRiCdtNEYUGkthFW4QEhADg1CIBEygjPNX3RWYVC5/enPStYTltELDW5Jx0qIrtXJHepcl2K/wANOZflx7U/hEyBgS+R0HaqckpWQjtVsFtrAdRVZ0BXJraDtqwZIAZID0GafZtuhKHr61E52oMce9JBMFbgfjV2unYXUu5ARuxxTEJCFmHNMWTByw4NLlsEnpWPKx3IlPmDOcYPSrSjzMFe3aq/lhVOOveprfiPmiVrCFbg5XvSSMRgY+tSAgDOMim4y3NZrcZWJLzY9KqMwWVvWr7xhGJB5NZk8bCfHvXVT5WJ6FyybltzfhVwLx061nIh81QODWoOMAisqysxogeEspBPy063jCZANTEAxnjPNNUbVwRWfM7WYdRjjb71QmVjJkGr8nqfzqGbAOfWqptpiaJLchVGevepJG/d8dqgifccenWrJHH1FE781x9DNuZgy4H41nMGDEk9euKvXRKSnb0qlgjgjj1rro/Dch7kR5bg1cgJD9KqHKyAAfhV+JCzDHYVtL4QReiJOaAcZyM4pEOCKUDDmuBv3mXcRzwMHrSRsxB5pzAFT9aZEpBx69aOjAp3y5+c8VAMAjFW79Rt5qogVkIX71dFN3iSWgQiEKfm70isUxIDk+tVWJVvmPU0hkxgIOKah2HfsXN7yy/J071adisQKn61XtVAOW/SppmCt8w61hP4rIdiVJN6g96XaVPrUNvhsn8qWWRkk+7UOOthjpASM461VZecVcOGAzVeXGCacXZ2EyCKMCb61HcReU4wOtSxEk7qbe5YKQOO5zW6k7pEvREYuREuPWmLIZOTUYTc3rmrCQkAAitGkgRbgkPlewp4+YZyc1Gi7BgmnKcPyMKK5patsomRdoxTjgqeeaAwJyKTOM4rJ3ArM4jk4WpoxuHHeoXJMgwPxqwo5BBqm9AHjoFz9aH4OB070hz5nSk3fKfXtUICvNCHB2nFMRGSHn8qnPA9zSHgcdK0UnawWIg+5fenMpZcZqIHOQoxmpVyEGeTVNNBuAbAOaaBt429aWRgGpHDsR6CgBHwvOeaRZAO3WmzDCjJ61Cr/MAOauKTVwLjNhdvrTk469aiL889qkjYbsjpUSjZAG4gn1pRwppGXBNHaloA4EbOTQDkY/KlHykADmkGPM6cVIDgOozzSn5U6UDINKfu5NSwGkHbk0zBzkCnNkp6U0HBGaqwDHyFyBk9xSblcYB6VIQGJPaolTbuHrV6WDUeMquQaYVJH0p4BIwegpBy/HSktwI+GfPamlDvGPu1KFAJAGaQjOMfjV31DcQx4lDCpud34Uxn2uOKdjGaiV+oCH5mCntThw2AKjXlulSr6mkwBuGIHSmck4pzcD3NNYY696SQDHOFwOajzwM8ZqdjyD+VROMNWi2AauUwDU6jge1QZJYg8e9TgnbgdKUgGsepA4pqtjp0p0pIOO1MiAyaFtqBNgMvuaXHBHpSLzkk0A5Oe1QA08DGaQ42HmkbOTzyaXohqhiEBQAKYyktnNPU7lH1pD0z3z1qle4iZM7jnpijPU/pTVOcntTgDjPas2A0k5GBxQc9BS5Azmk4K5FMCKYjA5pg4TilI3HBHNOXBGavZBuRsMcetCHHBPWgg7s0H5nz2FV0Ak7YpjZBpXbBz60A7lPoKlALkcAfjSMCxAHHrUSuCfpUo4UEim1YFYglnZDtqsCZXz2q1dqXTJ71ViUrjmtqaViX2LJiyB9KZGfLfHaleQ8KKkjjGQxqXotRksZJyKkkBwD61ECFcD171KSABWTVmMg2sZc9hUgQZyaV/vZBpR0zRcAA+UgVXeDDgmrI4XjrTX4Ge9CbTFoQlQWxilAVeD1pGG2QEUEjeQTxV7oY5DxyOKhdh8xHWnEkZIP4VGuW+U8VSVhbIkTgZPU1LnOai2/MBU2Rg4/OpkNWIlO49OlOyc1Cg5IqcD5cnrQwHDGOnSlb09Kb3+tKTtGRUMBTjGKUgbQe9Mxzu608gHpwKABjhM+tRrgoCalbGz6UwKNvt6ULYXUarbjz26UmPnBobCsMUo5IOaYxcck+9KeoPagcg80MRgUgGSJn8KcqZTrS5GcU9fT0ptsLESQhXz3qQgngUpJzxS7gOKV2FipMSCPWkQfLmlmILY/WlU7YunNadBDHUbs5xUbJhuTUwG7kjkUsi7uad7aAynJnfjHJpEJDkHoKuMgC8DpVSQFXGRx61cZXDcnRvm2iplAB3Cq0bjeatA4xxx/KoldAO7UpGOnWmjGD70oP/wBas2MD1GKbuwev1p2MCmMMDPc00AxyMZzSoQVJHXFNdcj+dOXhRVdA6ik5OB1prZ6fnTj97pzTXP5d6ELcb1bBprD56cmd3U0jcniqTsGgcAk5qFfmfmpCT0phbD4qogKeF96OMD601z/Fjk0qAEjPFO2gEw6YzUcjktx2pWcgY601kLMCalLuD2HBsmnhcHA6VFtI5xT8nkdBQ/IfmKx4GDTlzgk/lUfIAqRR196lgRFcMaaTwcCpGPzkHrUQ6EVaXUWhC3cGhS350rr82RSLnJ54rToIkON/NJglval25OSfpTl6HNTew9BC2SuO1DDIHGaPuj1NL/H7UhEcoJ2qDkipAMLim4wxenjHUnrQ3oMawOOB0p6nB3evak/jIPQUtJgSEDGaUfLg9qaWDNjoacOmKgB5+XtSfQ9e9BPTPWjGRxSAO9IT7f8A16OQcimseBmmhsbKu6PHeqC7lbBHStIng1TA2ysTyK0pvQTJY5doyetWkYOKhjCOcd6sbAMAVE7XGh4HyioLglTmpzkH6VUuXBwamG4hoYM+5eoqO4uAq7Ry3pVaWcg7RxmoFjeQg54NdKp9WS30HAM8nzc5PFSrHluB+PrUiRgAZ5x2qT+HC/hTc+w1oRqMHkYNSr05oC5PNLjA/pWbdx7ARgc9ajbjtUpxnGfwqNzg4poZFuI5/Sq+TK1Pkbap55pYo2KZHBrVaK5D7EpwF9qpyku4BH45q46twRn6VVkXLADmiDBtpE0aqBxUmBgnFRmNjyDj1pwBCk57dKl6htoQP8pPPzVEpPIA47mpnGD/AEqM5OOOPc1rGwbEWM8np0+hobmN8nscZp2Dg5+6Ka4Xym47VfUL9z1jlnyfSlf51wOuaRR83oaeQUOfyr5m6epRHECqlc5INSDJXOOKApVS3c0uSI807rqHQjBLLgClx8p78UoHb1pMjfgdCKi19RkSoWzipUODzQq7TjOKAMNim5WtYBz/ADdOKFz64pp+9gUpxxg8+lRfW4WAEgk8daCTuGelGNw60mOlO9gFbggUEZxil4PJpoJIxxUsBhyR9KXjaR3NL04FIenFAxAM9aQHrxSk5GKQ8DFAheAcnpTcdTSjp1oIyuPTrTSAM85XpTjnkmkOAvpQT8nI5NDGLglATTiMjPpQv3KEH3vSjfQLCpgISetAGfr1oGCSO1ISAwIp20EKq5JzzjpTx9054pF6EigMAuOtJx6AKoXByaaeWx2owSOKN2CDjmhNDG88mmkbjmgc5oblRiiwgXn5fXpR325pq9QTTjgtkmm0hgQAQPTrQi5JOelAwz89KVTiRh/DimhDGbccmmg557U/HzEUzGWAIqk1bUNRyn5sdiKcp+Q1H1YcdKmbhcd6JACqPKyevrQrFgfak3HbjHBoQ4LDqMUJLmFqIhO8jqKaTzyKE46dTTHO4dOaQxMgtx0oYcCmvwQQevWnk5UH8KdwsIynzAaRH+bFLICcEUjZWPIGTVb6MB+QXPoKc52gY71QhEpkLkkD0q8DkEmiUUrO4Dl+4QD0o5WLJ70xMMCRTidwwTU6tgCjjg80oxuyaYOGP6UrcAYFJjAtnPp61WCHdgHmpWPYcE0u07hTjoriFUcrn8qcxGTxTCQGJzShtwzik3cdhF+VefWpOARg4qNfnUjoRTsjAz16UWuA9uVBFJjByKXg4oI+YAVIDMZz9KYpLqwU1I33+KjQeXKeaqLEKqYAyfanqMSDcenShugBqncmZeQOMetUlzMNi0CSzEH6CoomJJYiqcNw7SAHoetXV+XitXFxiCJtwJHvTnBJC9qYo+UDHNOBP3vSubqMRjk7fSoyS2No4FPLAKTjmmKdvanERIDvAUnmhcAnJ4FNIw2fWhxuPHGOtHUYMQqk01H3xc9QadtBXB60qjaQO1NNLQRGz7uOwqB15wetWGXD8de9MYAuPerUhsjQAHnpinocZFNK/PSKCPm9afNcCYfKhb9KYFyDUhGMZ70ZAU+tRsAgGOccVMB3z2qJSWBHenowDY6ik7AKmSSfzpGXg96Af3oI6UpPBBoa0EUbv5AKzXLLJvXoa17mPco3elVJIMlcdBXVSmosTHhN6A449anwARzxTUO0baVTuk56d6zlK7K6ko4bAox+89KRhtcAUsmQ2D3qLaCEYEPxSH5uT1pXOOlIACelSGgjngH0pCS6AjrSyr3qOOTIIxVR20AeTjHvStkYIpTggHtS/wCeaQxp5YGlPLnI4pMDH0NIz5IFHkApOc460DjkGmnIP6Yo74HSjQBJPvgipHwcY9KicHcMelOTLHNV0AJOBjvTFyxHFSck59KapA4pJgIAqk4HNPA2t9ahbhunNT53KHHQdarUABAwT1qs/L9MA1MefmzSBduSefShaMCNEwCMYA6VJnKbR370inLc0vqegFU3zMOg48EDvQwy+O1NJG3ec4pFYHLZ6VDTAceTmoSC8/T5cdaeG3KTmlQbI89aqLcdQK8XyI3GTQjh4izflVkgdfWqkysBlO/Wqi1J6i2IkYsGyOOlRsAF2kU35zkLxmlbcxGT0rpSFuIwJiCtUcSqJCpHSpScHnj2qJ22OWWrg90hOxZdzhQByDSu+6Ie3pUe/wDd5654qb7qrnv3qJKy1KKCXDGQg1Zhlwpzw1V7iLyptyjrUysN68VpJRcbolNltXyAKeDlckVCT0wKfvDYArkktdCwcZO7sKrSYLlsfWrBcbDiq0jjaX65rSEWJjoD5zA7cEVcRt7HmqMAIbKn5e9XY+ORSq2Ww1sPH3setITk4HakU/Nz2pykLKfQ1ihDX7nHIqlKSeKu9Gaqkykjito6uzExtsQDnvmrjNtGfWqUChCc8mrTPgYp1VdgtilMC0gyOe1MljIQ5xUjhi+GOBUJJaTaTwOlaw1QioQQ+cVdjbGMdPSqrABsY5zT4CdxJHvXQ9YiL45IwKmPcfrUUfPPenjJOK4pbli44wTSfdk4PanYBpCpUhvWldJgVr3AGDyaoINhBHetC7TIDHk9vaqULZUjHIPFdFLSNhPcJsSOMDkVGRiQKas+SMeY3BqA7HkJPWtY2toIsQOWlVAOB1q/Mi+WS3UVnWoKzD27VrFd3PY1y19JIpbFWBt2MDAqZ8McmmEGNsjpTtnmRVEnfUBSAB8vTtVaUduTVj+DnnFV7mTy1x39KIK7ExgYDCgUlwR5J4471CsuDlvypzO0wKr0rfke4mMs1V1BI6npVoxDzc9AKSKIRqMDgUrsTItKUm5aDWw7hucZpyueCeBQu1EbIqMvkZNZ7jJwxKkikYjPXFLG2V561E37wmptrqMI8lunAqYDa2QMColxGNoqVWBQDNOWokhSfmzTQ2eO1PK5AGKYM5PHFQrDI3JPSnBcKc9aAcdaVB144q76CIRw2e/pSgFgQetKF5yOaf05p82uoEZ2kA4pMEgDFOX5elJGxyd1PuBHIuY8mqavhwO49Kvv8ynA4rLYFZPxram76CZpIBtz606EDaSTUSsTHUkYzzms2tx6EvU8nimnk+lOOKY4yDzWVgJF6E5/CgcZpAMZx0oz8tIA3/LTwQBUa49PrT+McU2l0ATHNBXcc4p+COtMbI4zSQCHhSKjBxxT24OD0pGAAHrVANXIyDQAQ3NBYDA60rcgHv3p+YAcgE9KaoO0+tKwyOtKCF4zyaa8gGfMBg9aeTzmmAEgkdqVMyDPcUMALAHGetTA8+1VXBaQFulT56Y6UNaALIpbGfpSEYxQTuANGOcmpAQjgAUxhk57inp15ppOBmmtAGFdwPFOjJBxSDOM09FxzVN2QdRzL8pNQxj5uKlbPr1qNQA59qlaIBwyGx2pxzjAHHrQMhcmnD7oobDoRgAHpQww1IfXoKBgE9eaLAhQc9sUw5IIzS5IoKjBxVAPT/Vr7U8cDBqNQAoqUnnJqZbgNKgjk0oxtAAoBG04pB+lSHUjbGTxz3pnQGnOCH600D1/KrQDTx1702VtqgqKe5B/CkwGHPT0q0wIpD+796IX6gDrUnlggA/jQiBAad1awdRQgB6daXocelGcAZppLbSR1qdxbD2Xeo9qjW3wpPc1Ir4X3oLfJjuaE2h2IDEowPepT8q8UZ2getOUZ+lNtgIcjBPAp2DspJRkjPapMkjipb0BkT54pTxxmlIJc0mCTnHSi4D87cD3pJMk9OlByW9qG79sUluBDIu5/amsADwM+tOflxikb5V9hWiEIccj9KQELgnrUaSb+KkbAkUd6qzQEnIOaUcpnvUW7c23tUqn5dtS0MhUgOeKkHSotvzkdzUwwAB3okALwMYoyA1B56UEck9qkBR98g9DQWwPWmA8GmD5gfSmkHQsnmMc4xTcjFBOFwKZESck9KVgBjTScH+dOwQM00kbcUxXJei8d6bnPagHK0ihg3pSGOIANSADORUYOfenjkikwHHoTUbghfU1Ickmo5f7vekg6FaRhnJ6inD5lIFRyAdqWMHeeOPrW9tBX1FGB8uelSj7vuKaRg+xpw4HJqGw0DkDB61SkBGcnj1q9jcCKrSDOQOoqoPUbZXhQrJnPXpVxeoqso+fP6VYjcN9e9XNslJE2QTSk5NNQDnJyKF681jYoVskUh4AOc0vU03PehACrk4oXaCSTSA5BPrSj3FMBCO/em/eyT0pz9KjVsr1/SmgQgHcUoYkUwEKhoV8ZPQVdhBvxIBikAJbpTSSZAeamAHbr603oMrvgOcjmlzxxQRh8mjaSBT0EBAHIHWrAUgc96ixyPWpu+MVMmCQzbkY/KhQNhzS8bqRmz0qQQqr8pH5UqnI68UdF54+lICFBo3AYRliex7004Iz6U9jhelR4Iz6mrQMhcHPtSIVPB5p7AgkEcGmr3PTH6Va2EP6c9jT1IPPSmqQ2ewxSoc5AGMUmh9RzcDNMZ8J7mhwWYr0pONwQ/lSSBq4AgLgnP0oRsrweR7UBPnz0FOxjjoKbsFhV6bqeCepHFMjHymngHaP51LAXGSDTl5JpFPbGaFzvqBDySeKQHnp0ox270oGSPekPcMjOTTSMjJFOPvRjI9qAI2J49KgmBPTt2q0QMYIqu4DVpF6jEgIGBmriHrj86pRgB+KuoPSlUEhXOIyehqi6lgQeat3B+QDNU5G2g+lOmgsQMm3kjn0qRBhewHtTQwODUg4wPWtGxWFX7wBp2Pfn0pNvTt7U8KMn1qWCEAwM9qCDgYpQvSgYJ44zSuPcjducDrTdu/BoZeCR0p6ngc1QtypcqAyr79atQxgJnvVW5H7xTjNXEb5M9T6VUr8qFe7BgSOKoyApKQeM1oZHNUrvhgcUqe9hvVEyA9uh60qjnkVDFN8x7GpXkxyO9OzTGtSGRcd8YqBxx6Z60533c44z61Cc8kd+ma2itCdBCOcg/WkkJKMAOx/lS44O3j04ok4jOB25q1uC1Z62oDZNCsD978DTVbkgetO+VkYY5r5pJdChC2QR60pUgDnjNNI+XBpRkL83TFTowHIobJPQUkaAsfbvSruCfWgDauc9T0ppp2AF5bPbNIeHJAqRFXBJOM9Aah3EH2pSVrAOYA5PSolOfansdy4HSkK7QAO1KSQwSlwcmkHoKUZ6HNRoAjcexo4C4xStx70h5XpRsDEz8tNZgp9aceR9OKawHFNIYYyMikJ6Yp2MLSLxz2pXQCZG0+tKCQcUmPnwB0oGN2aYAPmPtQTkAntQvU05QSdo6UAK449sUqnaM560jEsMelPjUbcseaaW7EBXHIbmg5Dc9xSqoIJJ57UikYOT9KYCgkJgdO9KQI2HehMc7qAMhs0WugGuNo47mkP3SDilYgxjj5geKjyTL70kmApwI8c00fKoyKCSTTpMELjrQwIwcnkcU9+cnt2pikBgewp7/ePoaHsMU48oHHPrTSOAe5p5wowMYIpCuSD04p2EKRvbr2qHOc47d6kxhwM1EeHwOlUkmA8Hjd6U8MGRePrUfYCpFAHHT61IAOjBRxQo5JPSmjgk9DSqCQR+Jp6aANJAYk0jD5cjpQ4DR0IP3ZGTRoBCykoxoQfux7VLj92wpij93jHNXze6HUDkdOhp33geeKAAI8UwOAStK7egxMEEelE5ITg9aerEuRik2buvSi/cRDDuTCk8E1b/gOKbhfMBHalUkk805y5gBAPlz1zTh3zyaYxJxt9akIOBWbegEbKCMj1pC24EelPPtzTNvDsByRVJdAImJOB2zzUwUBQRUSAhDmpk+7k/lTbXQYBCoB7mnADuKcTlQaOqgd6zb10ARR1AphODT84bPf0o+8/pxSEJIDxzUboS2eo9akGTx1xSZyMe9NAMdjuwBkCnSKPLB65FOP+rBxSDIQA9KrcNSpJaDCsODR5mxtrDt1q0csg46GoJId446irUna0mBKGV1VgeKeDh6qoDEp5yKmVw6ggfWhpboBXG4k9qThh6Ypec7c01VwDUJ9wHjJB9qCuRvpEOYzmnYAA5yKVgE649aUDDAGkBycY5pT05qRjX+WQ4ppK9qVxjoahA2nIq0Ic5Cgcc0gGRg9OtMbd5g96UcHJq7DJxyBTSCTgClTl6AT5h9TUq19QBT8wqQFcnFRr1z6U5MbsgUrX0AdjjIpAP4jzT+Nwx0pg6lRSSYCSYZQeoBrOuJGjuQo6dc1fcjAqlNteXHetqT11QvQlJ4DdjUgxsHrVdlIIUHipU+9j2oaGTbdw3HtTWzjJ/OlzgYo3fLtNZq4hrg7qM5cc0p4OKEGW5FMYkvUioIhiQg1ZI+bFRAbZM4qotoOo8EMT6UOQBxTB1JpXGI/epsgFR85FAVCeahjfdkDgipSMYNNqzDcaQAc570ZIfNKwGMigE4PrT3QBgl8ihflzjvSEnAHenAYbHelsAgPJwPrSAgf40AEH6UgACfjTVgFKgyjNOz8rqOKilcIVYdqGcgBgM57VSTa0ARmJjG2ntlQDjmoiDxgY5qdssBxRLQCFCwwT1qUDcGB6GoImYuUYdKk3bGI7U3cBOPKwegpp4TI6Un3kbNSqAEX0pgRKeMKDmpeQMEVGg/fnI4qUZIIPXtUsBH+UD3qJ22MQTnNPcdDnpUbJ5hB9OtONluBVIKyA56+tJJsDgnrmnSgb+OtRuysM45roXQnYV1BlyBwKicfvhxwakgbO4scjtTHG0bs8D2q476jJoQGLZHy1MnIIYdOlMtQGRm9aeoOzFZzetgsRuiScdcVWK4m5+7VmCNo+Xqu+FlIB681Ud9xD58gAqTgdaWNxnOKcPng4PNCRfMCfxourNMB7YCEgc1Au5wVYcelWJAVXIqCTd8u0YpU9gCJ1HGMZ4q1b5/AVTmG5hg/WrUX+qGDz3p1LNaD6lgd+KYT1x1FKmSDnrUTt8xArnS1GPz61Tc4Zhnkd6u8EADtVSZQR1zWsNCXcjikO8kCp5OQDVJGMcpGeD2q2BlRk1rNbNAtis55IY9KpK+Hznr+lXLmIoDjk96qOpYDjGDWtJk9RJFOc9c96t2i7hyPrmq7YKkY4HepbdiDgelaSV4sFoy+ow2D1p/frTApPPr1p4PPFcStfUsBncdvXvQ5DYH505CQzUqr948VKauG5BLymazU++wB49a03G5TWaGVXO0cg1rS1uHUex2jaOSaBEAQxGSamjXJO7uKGZS2F5xVt9AJoUTf71cBHas8Ehdw4Pap4uGB7msJxv1C5YIBOO1MbKAjtTXYjOaQvviyaiw+pGsnpUN2pYAgVEZNpqz5itEGz2rdx5bNC3Rn7Gd+uF9auQoqEBOarO4YkDtU9sGA69a1lflJRMDnOKiY7pOKlCLjI61XfcDkHFYKzYyVyHjxRwsOSM461GhIBXHPrUiKfLIbqav4RiQyBycHBpTKVVuOe1V0Q+ZuA/CrPYFh160NLcNyrlyScn3qaEs7gnoKkUIwO0UsI2g+uaTlpawkiznBHrUbttzxSh++Ka7fLnuaxS1GMY55qRSfpUGDjc35U8NgDA61bQCkYOKUgkdaR2XoOvbNKjHIUiiwEQG3OD0pgc5OR8tLKdhbB5qszZKjdWkY3EXIyCp5zVSaIIxI71YhwoIFNePJyTTT5ZANjOQAalQfMR27URYC59KkHGT3qZO7HqMI5we1O2DAFKDk0773Ss22AfhSdjk9KdkDj8qTA6GkFgQAAnHNL0B5oBwmMUA5X6UasBfTNI1IDz0pu4bz3FCWoCPjdwKj5DjjipfvHPao1OWIb8quIdQOAee9PIwPSowCz8/hUij1FDAYSM4JppJyT0p5UF6bIMYGeaaYwjGVYd6dBwh55psZ+Q5606PGMZpydroRFMfn2jvTxlY8d6Jly2fSnAArg9qV1YABwnSjnIFO25B/SmkfpSQCg9QOlMYZXPWlBABz3oHQYoDYapyO9OVgQO5NNP3sY4pHPzKF4p7jJGOBTFznnpTmPIpq/M2KS0EyXOVoPPGaQcsfSlU4ye9SBGy84pFwXNOPc03ouRV3ATPzHI/GkJ54puSTnNKrcfSqSDqSo25cDtTz/ACqKEgEjFSdhgVElqAuMdutLntTGbBFO7Z71IDX5GajzjJp7ABck9ai6j2rRbB1Gn7pPcUKflzTiMRnimgYGe1MNiQfnSH72KUHg+pprZApdQIyMdc4px6DuDSMSCM0dMCqQh4ApCNv1oH3+vFMd+DnjFCuMcRup2cMB2qssrGTjpU465NNprcWg6QDqakRsDOKhdu57VIvMYqWtBiA/Nkmnr1zTGGXA/KpARtx6VLAOjcU1vvHPehT940HqaQEJ5J9aY5ypU8elOJ5zSMcngVql2AhKhQNvUdaUncQf1pdmG5NR5bYRmtFqKwscgWRs1YTnk+lVgMe59Ktx8xgmpnYaGMP32e2KcpBbOOKRyx/rS7gRxU9AHL60jcUq9SB+FDY2nNR1GRs3oKTdjihs01cnB/KtELUlRiwxjn0psYIYjNSREbsnrTWyj5qfIBpbnANMY8nIoJO8nvUTvkZ61cUBLE+c5qcHjFV4h+7yepqZeBmpklcExw4HFKp9RTMkt7U9BkVLQD6ZMMgcc04jnGetKcEUr9QsUXGxjnpSoMk4P406WPe5JqJTiQjNbLVCskWCOAKcOFpiLknJ+tS/wD2rNjEyAP6VA49Knblc1ERnJprQLaFcAg+3epUA4x+VRhSM+1ScDk/hVslEmeM05Ohpg7DNOz8pHeoKFycZ/WmtwvPWlPA9aGGRj1oQEKyDgf5NTdh3qMR4epe/HWm7AB6cniq4PJqc5xioHXnFERdBuRk+9NLDdtA/GmjcZsE/hUjAAg1psAw8MPX1qxu+QHAFVyA+DUyncu3vRLYEIE53flSSDamakX0prY289Km+oDcgsM9alUk/41FGN2cjmpQQMDFEguxjcHFIe3NKTlqQDv6UAO77aTp0BxTgeM1EmQTk/SgBzDOPSmAck09juGf8imErytNMBknGfSoh0PHXpUr4KkdqhX369jWkXoLcep2uMZ+lSIMHmm7uRn8qkU/OT1qWwE6Mc1CoJl3VKWyW7Z6VEjES4HXpTSGT4Bz6Un8OBz6UZy/J4HWmh+/UUrMBei4707kRjufemsQM470/B2c0gVhVbGTTl45qs24/KDxVhQAABz6UNaB1Hr6kUuc03OMYpeegrMBTjPek7+1GB260gxj3FAbjj3qvJyKmHAOOv0pjHA6c1S0GQICD1q7G2V+tVQOTjtU8ZyM96c9QI7nLEc4qpKC+Dnir8i5FUpxs61VNk7ESggj+dW1HPv2qsoyOlToTnmrlqBIRk+1AxvNBOMd/UUDqT61mMDnOAOlJgdc0xn2NkdKUMSuCKdgHYHpxTcDt0pQTyM0HPGBQg3K9xF8hI5pluxA2k9KnzkHmoNrbiRWsXdWZD0d0Sltpx2qvdMeD+dTfMeo4xVac5ypP6U4rUq+g6PGRx+NPG0n/ADxVQI5zgf8A16VFbzAeee1W49bivYlYbecZHeoxwN3p1qeTBT3zURPfvRFg1cj6kYJpGb903HalYEcgnFNdRsbB42nt7VpoB63JjCletLgBd/c01+ccY9zQPu4J5z2r5nTcoCW2jPSnEbk6/hTOwyeM0E88UlZMY8McBTSkcDPamtyB7UuCVz1NVLZoQ4gMoxniotpJ46U7OBkE0gxzziovroFgQ8cjpSk85NR8+tOyAR/Wle4xSfmBHSgDk5PNNJ5JJp3VciptcBMZHAoAPOaX+EDvSnhck0mAzpn3FJjmnYzxTQCGzT9QEYkkCgnAGR0pcDLH8abyw6UDBTgc9KFxyT3pANzEUqkA81S3AVcDINKozn1ocZOaByv86LX2EKBwQKVRuG2hFGT9Kchw5FNLYAA2kjrQRjA70HhuDTj94DvR0Aafv80jkgbQfypsr4kAxzSR5JyecmntqwHciMk9e1VFkfzyTVsdSDTVRSMntQpAKF/dliOc0xScU8c5BPFJkAGoeoCP2pWUqAfXpUbkqM+vrUu7egJ7Cq5PduAvVB60rNgD2oGAu7r9aQDcTilqtAHOvAI5qvIcHpUwYk4qGT7+CKqO6Aaj/vCPapWBIDCoEHJPc1LGcLkniqqWQXFDcgkUbgG4p2PlLYqBs7ge1KMbgOLcDHenjhAQKqoW3bPfg1YL4ylW0o6AK/OPcUhU7hjpSjkgZ6UqnqSazv0SAa/AKmqsn3wQKsO27vUbYGPp2qojJFIA9zSx9TmqsjSeYu0cHvVs4VAO5okrIW4g4B4ptuflJNPOAoBpcADFTsgBSFYGlZyvP86Q8U3cWAyOlC0Vg1Hg/uzxTQDt/GlLZDLjAIpQ/wB0Ypt9gEkUZG0cUKMsMfjS52nBHWlUbTyevpUt9gHNweORQ46EdqVxtYHqKTdl8djUtANY88Dk0053gUrEKxzSK3U55poY4/KxGetAABIPYUzPf0pUOcsafUQKTjB/ClJONp65pEPzZPGKU8yE/jTvZBuC8Ag+nFMAIPWn5xIDSMfm6UruwFeRGZ+DweoqWIBF29ac3ytzQvy/Q1XM+WwCMc9PxpEO7NOZQDTWyVytShiAkKVp+ACOaYFwpPehh8o9ab1EhynJ3AcinHDc1GhxnPenKCDUsYvGCCeagPGR3qaQfNkEVFJhXz69ataAJtBUEnpSOwxxzSKc9aSTmQY6VfLZgSg7Ywe/elUcEnqaYQS2B0pyjL8njNJ9wAE43dqkUhRu/SmgDdj+Gkz823tS2F0HjJII70icSHHeg/KwUfhS8BsHqetFgI35bBqnNtUbs8irZHJyeaoXSs0bAdRWlNXYdBLW5EjEbufergKhic/Suftt6Nu75rYt4mILuetb1qcUCd0W43DYIpEJMxBGB2NOjQIRnoac2CeK5U0ndDEYZYDjNPxz7ikJBINDHnjtUvzADwd1RspLZxT2OSCDSsQeRS2Agb/WYJxTnPOMVHI48wE9qeTuPFW1pcCOOPZI3pUhOW56U4jIzUZPQd6HqwBj82BTcEnjtTuB0FC8DPrR0ADy3vSt160DIGQaUYIyetIBGO003pgGlOGxx0pqkbue1NbAMn4QE1IwDRLjrSSKGGKFOTg9B0q4vQBCDkY6VKTiIU3H8J70Y5Kk/Sp3APl2H1PeopcgAiphgoeOe1Rvyue4prRgMjXGVzUrrnHNQK5389+lTYycntTd73YCsoUCjHfpTjymT07UjH5V680mtQI3PGe1NDADNOdwDtFMYhf601sBAxHmcCq+zaxyelWpeCCBxUDcE8fjW8HoTsNVdwwvUU1gqgKT97qKeAUPPeo5yGkUjgd6uL1AtW5EQVaHuAJdi9qqmTK57VVj8x3Lfzq/Y80m2DdrGy3zAbjWe5CzHPWrcDh0BI/Gq1yAJ92OKzjFRlyg9R8RABIp8MmSfSoI12knPWpYQVb2onFXHqSTHoM1VfduBzxVrhm2kc1UfLM2e3SiktROw7AbDDtVtW+VcDrVFdzKT6VcgclcsOcVVVaAtSdcsAc/WmsB1pY2xx601/ve1cyvcoCOKrzrhC3erfVeKgkU+WaqLsxNFAsWnx2q2FwOuKh24bOBn0qyil0ArpnJNIlFeZd3NVGbapB/OtCVAgxnPsaznbax4+gqqUr6A2NYFRx17UW52sQaEPyl2PIqNCTIDmt0mSbEWSOPxp4PJ4qrFJhsYqZpOORXJONpGl9CaJuTTsgE5HFMUjjFSFhjFYy0AjkXIOKzniUSZ/irTHqPSqJAaTHbvWlN22CyGPu2ZHbvTEkEZ9SetWSAVK44NU2jMbccitINPRiZdyWAIqxGOAarRnMSrirK8DA9KzqW1KFKbjk0uweXjGKcF7U4cde1YXAyJ7ZwcZ57ULE0a9e3pWhMcnA61XCl2yT+FdMaja1JtYqLDkjjI9asxrtTdQ425xUMbGQlc/Wru5ILE6t8tR78lhSsSuAKUkb/AKVCXUCNWKSHuKfuKt169qa8kYPv6U9kDgNnpVO63AU8fU0pIK4JyaeYwy5qIbFkzUaAO4XKjg+tKuVUg9aXrknoaIwQCTzSvoMkXGDmgcCkB+U+9RknpnA9ahIAZsqeaaSQuR+NBHoak424xVjImAxk8mnLnG4HtTWOBQvC554q3toIbL3x1NVwjF81ZOCN3Q1BuIB56043SsHUnTrxz7VIeQT0qCEkZ5qbI7nFTJO4Cx4HXpTwBmoxgGpFAFRIBpwpxTl46UjY4oDc4pWYDu/I+lKTk5pOf/rUp6YpPcBF5yDS4Apq4BwaXtTW4DQ2G9KT2xQeSKJDhR6U0AvVCKYwXb/tU7JKj1pHAVf60LQBoY7sGnx8k0xhkgr0oR8darcBTjzcelI4y2etIPXPNJKSo9/WhIOg0tzjtUqjac96rRyAtyeatICQc1T0AGxu5703r07U5hz14pDxwOtZgP6AAdKawA6U7BFMk3YGKSAjXlDmnbsp71G/3hzilU4bBrToA9TySe1NPBJxSkfOfSmt93Jo3DUUDIznjvSx/ezUat8xXPFSINvU02rAPBHqaUcDnpTF+8fSnkAVDAa9NPIwBTzwOetMBBBpoCuwbfx0qXq2BSYPJ7mkDYJq73DYlQAHj86ePvEE49aYnODmpD0qJaMBrjIBzinEig8+5pd2cD0qeg+pG3Kn9KYThORzUh5JH5U1jgEVSEMwAmT1NRlwi8808HI60wpuBPb0q0BLEeDmnNwPemJ8uPSnHIB5qHuAzG4c0wEc09SCCB0phIyOelWIUMS2e9Nlx37U/OGzjimyID1p9RlfcC6gDH0qwevWmrGN1Pbk5/WqckGorEAU9SdoAphBKU6HoPX1qAHscN15pFJZcUjDnOKcrADAqegB0AFKTgnFBHzc0VIFd+DSEEjavXuakZcnGKiQEOSelax2uDGTA5wB0quxbB96v7e9QbQXK1cZC3IQdxBFXUGFA/SoSgRelSqCzZ9KUtUCByQCeppFAKjJok4BNRgkKCKlLQZOucGgnNMz8opSwUe9KwDWHIpp5OegofgZqKJyxyRVpaXEWIcH61IVEg46g9ahjBD5J4qWP5ZCOoNQ9wEZOeOoqExZJxVk8ZppXjrzSUgZCP3eBgDFOznFJIuCSaFGD1qulwHMeKWInOKjfJHBxTkBU+1FtBk3VutAPJ4pp4P8qcT3qBERABqJwTUpyG64NIBlc4q07ahYRTg4zU2criq+cN71Onrn6UpAI3K4FRHHQ1Ng4JFQucDPWhDImbk96cCQBntUJfc3TmpVGeTWrWhKZICGNKBzmkGAPbvTic8VmVcXHvSZ+bOaUD1prHIPrSDqLj58nmkz3pGbkU78aYA2NvAqI9RxUp6Cozx0NNAQgZl+nelON2OtOzzu9aZICMe9XuxbDgMkjP0pUCg4HfpSKBinJwD3xQA8BjTWX/JoVju9vWlOKnZjGISHI7U8DPWmkbT9akXke1DFuQkfOTSBick/jUjEA8fjTPMC9cc9qpXsA9SCOOtRO2SQKabhQCB1p3UZppW3AOFQelMB+Y06Qfu+tRqMHNNIXQJjjimDaP6VJKOeBUUYGSR+OT1qlsDJQQVqRWwcgZFRdvWpF4GaljvYbKdo3Ec1CGCnJ6mppMnjHSo2ALKaqOwnoKxKnA6tSK+GxnmlkPyjt71C2Ac+tOKuCdiZXyTxzUqt8pJ6VAp5Pp61KpO0c0pLsNDlA7fnT1HPvSL09aUnaQT0rNgOzlvSnZ6nvTQec9qVTxnpipsAue5pM44o6tn0pjsRk/yosMfuCjmmHk1EH3ck9e1S8AVVrAkhvG6nK23gHNIF496QKcijQLaE+QSD+dV72McEVOo5H6VBfSccdqIL3hX6lIEgH/CpFkyo6mkjPRTyasLCpww4HWtm0tybArA4pd4J4p7J0AFV9jA8Dis1ZlBJ8zY71IgO3GcfhULkouc4xUiPuxg/nVNaBoSHGDmo3YqKex7+tN+8pqUDIz/eHX+dKBnnvSkEggHHvSAccVQlcaWPAI59aguVJXjr2qWc7efypm4NHuarjpqJlaLA3Bj34qwm09xVRhtc44PpVhEyMelaSXUVxzAjpyahyeQBVgoQTzUfIzjt3oixuxWPDZznmhzlH57H+VKAWGBxk9aZJgA4JHB/GtFuK2p63kkbccigj5x2pACDn1pQd2Se1fNJX26FvQe6qi4Peo+3ualA8xck9Kj53dOnSpaAG+6AOtOGccGkx83J96Oc5B4pdRgQcciocMj+oqY/N0PFKAGXHtQ7ARrzT8Z9vakQBQeOKXqOKlpboBrpu+lLjbxnpR070uO9T6AHOcYxSj3oGcj+dGQOo47UIY1vvcU122kDGacQc0jEAn1p9RDW6Zz1pAMLSt0prEZ4p+g9hASDntS/eb1FKQNopU647UXAVAWY56UKcuQO1A780sXAIxyaasIeFAYE/pSqvzHg0ijjJpVZgS1UlrqAoGXPtS7j5g9aGwEyDzUQJPJxkUtbWAJUPmc04gr1HApTkruNI2WIGabvsgGk857mlB+Tp1oYjODRwBzUdwG8YwOtI3TGOlOUYJNNzk5xRe6ASQDjI6UJyu3NEnzNmkTkfL2pq9mMlJO0DHFCnpilLYXB6dqRcJzihbiA4bIXr1qKZDww696lKnG9ehomXCAjqapJgVgQWHrTkOVIpGXOKdFy/IpttgPU/Jg9RSYDK3FKOOO5NPOMDipvZgVpUw6kccUpAySaeRkc9hVaV9oGTjNWnfYCRHwCe9JuY5I9KQAHBXv1py5DMD0xTtrYCFThDn72aeCrkKOCKfs+U8c0gQA7wOnWndNgSKoGB2NDDLEe1DEAZPFAORu9qzSe4xqHIGacMt9M0Agk/SkXg57elO4gYkPSNwMCnkBvm9KYPmwaLWvYB38J9T1pu7GG9KcBluOlNdSQyfkaF5gSJh8sTREMH5ulRxg7fmqVTxiiTAc/ByelHQbvWhj+7xR96IE9jUegDCNw5601Fw9SMOAR0oPy7fWgYxVyxBobC1IPlYsO9MIyoJxVS2EGMYBNI3p/Kl+8RxyBSevrUgB4IxQAWbk80o689RTQfm5ot0GGA2c9KXIKClJAJGOKbwFHFPoA7AK+lNUAqVpRzgHj3oBAI56URdmAmAWNBI3AmlIA6detNdT70bB1EwMkdqdkNHjvTU5i96kXGGB/OiSsxEfQZqKT5mBAqcfd96iIJ+tCeoETc8d6cRwMdTRtwc96QE5zmtVdjFUlDk0o3BsnpTipYkGkkOAoFIQ/nbx1pjL84B+tKO1IzfvNxpPRjJSuHBoYfNmhcuvqacPmcUndaAV5MtKO1V5QGibnkVdZd0pHSoNo3FSM5q4yYijDEjjpzVwAjCgVHCAszelTj5n4HHrVzkrDQ4k7lBHApzMFb2pCcyA0shBGO9ZW6gKW3YPakyMnjrSBsDHFOAG0HNKwDcHb7GlHyqB2pSNycdqTAKj1ofYDPvQyyqR3q0g3IBTpUDnkdKXaFQAd60c7w5RBzjFNPI4p2CQfWkYblA71mMjYZYCgnLY7048GmcFqu66DHsMYFKTtbpxQPmNB+Y471IDSCCSe9NUY60pyzjNI4G+qEgkIAApsTBiRSvyvNV7Zf9JY7jg1pCKa1AuMAZB6UNxn1pUHPPWm/eJz1rPYYdAvvUcqnI5wKehxww6d6R8upyKaEQgZcAVPjKkd6iVQsnPU1Lu+cgVUg3AnB29qQ5yBjijflulOPUAcnFTqkMjbkZxzTDh1HtUvIGOxqAkg4FNO60ENblcdu1QS5yKkeXA2kVG4Plhh61vBNCZFKWKAkcA0yUHyAR60/cURtw69KJDiEA9TW0U7iZGBiLDH8adCnysoPHrQqhwVPapIRgFc5rSUndq4EtmezU24UYIAyTUkY25HfFNX95wa5pO75h9CBFIYc1YiUE89qix5bZI4qYkBc9j6U5O7uITO12A5zVfJyeeKt8bjx1qs2QrdDSg11DqQxtgk9jU8LsHKgfLUKsqIy9eKngzsPvW07WdxIsYO/cO1SMBtJNRw5OSacx/dgCuSW5dxqk+WcCgnKZ6UqH5CO4pvIGTRYRVmO8jbx6ip4T8tQsgYsT1qSGQFPk6VvLWFkIbcHYMjrWbKjFgc9e9XZjycnioGQyDr0qqd4oTIXbChcU0JsO78qkAzJhuRSyxjcMHgdq6OazJHRvtbnvU7sSnHWogFAGcfSngE96znrqXcsQjKjnJqdlPrVWIlWwatZy3NcstxiA7jiqNxHtbcelaBAVvas28Zt3HSqpfEDE88rilWZJODUJ3MoWnJGwbn866LJbiLsOFUDsalMm08+lRxoGYe1SyqCRjoK5nbm1GSBsjd2oZ84wOKaGATaKeFG36VnJK4+pEVO6oypBOOlWQeDUTjByaak9hFKdjkY7UkaANxxUsyYOaiLe3St4t8uguo5346dKYGyMngDrT2HAIqJSVU54pxQEZwXx+tTAvwByKifaFBB5qxEQ0YYnirlewInJyoUHApGjG0HrTGcZGOaC5JHp6Vz2YxxAf5QKVV2qFHNAOTxSqQX96etg0GqcsRTGXeOOKcPvE4p/QdaARX4WnK3GT2o25Yk0xWO/iqQEwG4Akc0oB29KYCA3J4p5cgE4od7gRzDao9DVZ42KkjpVgnKZ6+1NMg8s5P4VcW0G4yFMEZPJoOAxIPSooXAbB71LjLn19Kpqz1F0Jxzz19qeCaiUYapAcsTWUhjzgJUSMC/I4NPDbj/WmEfMfSpQEgYE8dadnpxVUSNvJAqzuG0ZpSiA0nByAMHrS8lcCoZZSGCgVMvAGetNppAKMDHFMlAwCenpTgSD0prruYUluMaG3DIpTg8UbQpAHSiT5VJHWqvfQQDGMYpmPn4H40gcbQf1pytg5Jp2sMY6nePamXAzgZ5NWARgnHWmsmQeOKcXZkspqojk5/OryZHPY1WdA/fp0qePPlAHtTkrq4AWIY+nrTlGWyaGHTFJEQ27PWotdDJM7jwaRz1yaEI3HA4pDy2TUgQnl8YzSj/WH2p2Pn4oB+YgVoAknIyOlRSE7eBUnQGkxkfWhAMA2ndjmpMAnn0pGX5ccUEgKaY2ODZJxTiec9qYnIpx+77VL0YCNg85yKYOCeaQKduTSAlmzTEPHUkVDgqSSamxkZFMlX5c96admAQn5cZ59KmDYGDVZTgAipw2T0okgH9O9JkBqcRnGB1qNwd4IFQrMBcnJ9DSYB57UNuDcdP5UEZX6UwGdB7UiDJK9BSEkHjvTiQD6VYCggnHpQcDIpsbBzxUhXLE1LVnqBGAdp570wAYIJpwIzj86btG4gYqrAJwpwadtJPWkx8/vTwcmmwEPyk+9IQMZoJBBpOCOlIBc/IAoxSwZMf1600E7ScEelEQxyKb2AmJ2ggc0yP6U/gc9aah2nmoWwdSTqDQeFxTQTg+hp3GMUmBGxwajBz0H1qRmxmkB71a2BiMNy8fjUC/6/PoKHZgePWo8kMDnBPatIxsJksmdwz+VOUkOD2701lywNOBBJTBzSHZCt8wwO9RI43YI4p/Kv7U0belC0QEp54HalABHXNIi/LTxgCoYEMnFRDKjgcVNKMqajQkgZFWtgJVwU3dqdH2b86Yo2KRimI+GxSsBZzu57CkJJb6UqEbTTf4s1AEcg59jTEJJGT+FTtggjFQbfn71cdUIf7d6DuOMHFAOSOOKUk9e9IepI3H0pCefagnOB2pu/5+lSkArctRzjHrQeDQKYEO07sd6nT06VAxAcHFTL65py1QkKx+XioSOxqdh6ioiBmlEZDtyRjt1qTolJnkjNKQccDtzV3FsOjIwQadk9xxUanmnnjFS9xgPveopMAZzQuRz60vGTxQG5GxJ6U8EZ6U2QhQMdaUkYFO2gCkZqN/un0NODZNBXIoWgEaYI560kgLCnfKM5NHABx2quotyEk8YqcdMr+NRlRtOeaB8o56VT2BIeGywp2eeOlNQYwabIwDcGpsHQczjrmmtNnIU81C7cn3qAuRxWkYBcfJcYfGeRUTS7ifWoSuGJGcU4Zbp3rZRSRKfUkUkjOKuQn5BmqkakDn86sQZC8jPt6VnMq+oTk4HrTVbgYqdhkgdxTAuDUJ6WBoDzzioQNzEjirBzjJqru/eHAqoi6kyc9akUkr/Ko044p2cAED6VLB3FPB6flUaMFYIe9PZjwxpgyX3U0tAYrqSCvFR7CEw3Wp2YFsE8npUcwORQm9gshoznHSpF5Xr+lJtyf6+lSAAim5DBM5OOlOfBAAHNIgwx54pwOccVm9w3Dpkd+1LjJGOKOBQvc0gE5yT3qN1yp6U/J3HFKBwBRsHQrRk9MYqTJJxUhQY6daVUXFU5JjQqZ5pHB7U/ofejGRUXBhE2eTVC7YiQsOeelXgMDAqjc8vx6960pr3iXsQw/f8Al6VoR/MvPGapwgg4/Sra8ZxVVARIR3z+NNxnPrQGUL7UhIPQ4rICGchUI702HJTHFPkOTx096Yv3+OlarYLkg7DFIeh9qcSC2R+dMZRzzSQ0KOxNNbO4AdKXBxwf/r0Ekc0CEkVWBGOlRKu1SMZqV8Y579ai3DaQRzVxvYCk3LE7elTxMMeg9KgmH7zkHnvRHknPf61s1dA27FzIx9Kjdto69egpRHu+bPFMcjkdcd6iK1Aru2ScE0x1OxvYGpmAJPHA75qN8mIjqADWy3J1T1PWMcgA0MdhA6+uKHUKwOaR8AjJr5q3VlkwUeWMZyTxTDuU4xSgYUEGkcEsO/FDeoCBs0vIGCetIeMenpSlQcc0rdgEyRkDvTjjb700NjoMkUpPGTUMYDGDRxj3pucn6U7IHak9hi9eFpRnHTmmrxg04CpvqAgyOwxQTgc0o4HJpAcnBp2EMkY9vSolJL81LgM+DxSIgJJBq1puPQYck7e1IxG3jr3oPfFIvJ6UgFB7U/AAz1NIMNJzTjjcQKHsAo4OccGnL8jZPSkI+Qc9KBuYD0FPZ3QhxJVsgcGnFgflAIoUFxnoFoRd2STg+tMQj4KkD8ab1UAce9PPyx4PJNIxULgdaWzGN5yEpmfnGPypx5f3xTEJ8zIpsBT8r89TUjDLKPao3O58461I3zAYFSA1hl+P0pM5BwOfendFI700Cn5jGjoc96IgURvUmg4BBpof5z0xTiIlUbiBnpSsCx2+lIG5JApVzvBJ4xSVgAnK7R1ppPUHtSgfvcjpSyYOCOvemthEKg809PkcGm9QcUqDIGaNWxjtxD5IpM/OPrUm395g9OtMxhi3YGk1ZCBhkn0qpLbbkJPbpVzkIOPfNEn+r4696pNp6A7FRPkjA70nmFjSOG35GRTPLPmgg8VrG27Yy23KbR270ijC4PWncABaJMjjtWMrpgRyRArimMCpVQOKkbOeabn5xVK/QFuOXhuKcOuO1LwORxTen1qXuMc3AIxTUHyECnHG3Oee9MQlDjsapJgOBw30pCe+OT3pSoGOaQng1G24hSMtjsaXkDb3pEyVPqO9IM9c8mmkBKFO0elC8AjPakB4x6UA881KeoAeKb5uWC45FBPOc5oAH38dKaVtWAozn6mkJyCAODSjIyaRSQPeld7jEOc8delKuA3NBIzmhSC7emKa1EIh2sT27Ug9R69KUD5iMdqQMFfnp0prULCsMk4NIwIxSNkNkd6cTzj06VNgAn5M+lIy4ANLgvxSgYODzinYAAHA/ipXPz800MMgg9PWlPzEtQ77gIRtQDoc0Y2EH9KVycA4obsT3pDEABbnjPaoj8pNTN98YqOYbmyv400IjxhSe5pnAGAetOkJC5/KkBztOO9apXQyTORimOhZs+lKOZuPWpWGHIHSpd0wIkJ+bd+FGeOvJpSOcCl2cZ7CgB4GzGD9ad907vWmbvlzTS5ZcdaV7sAcFpDimADNSqo2k55FMHPJo2ApBsXG0VbiHPNRiILLnHJqULgk1cpJqwIROXI5GKR/mYUqkZ570m0ljjtU9QGsSTgCpEzjB6ikVT1NPQgZP503LSwDh/qyB1pq5C5xz2pegPFI/QYqNgEB5po+XP8AKlx+tBGSAeKYCAfMKH+9xStkLjHNI2NoHrQBDnDZ601zjkU7hetRls++K1VmBNEeBTlIzznNRRMWHFO/jqHuDHAfNzUQwSzH8qlPLdetRLg5BFUttQA8qBUaJiUEdalTl8U4KGbpyOlUpWCw8grTT8ykinZJyT2poIIOKgBMYB9TSj7uMc1E+7zMipVOSCaGMNueTwe1NUYBJ605wcjHT0ofp70IQ1Pv89KlXGSSKj5AAp64xt70MBpAOR37VBGu12J61YYDrURYb8U47AVcgsQVpjEmFlA6VK5UvgZHtTAOSvQVutFcXQgKs6AHkrT5UDoPUdqIj8xB69qVBgk9xWjdncQQrmUgDtzSiEplh+NLE371sd6kVsHBHFKTe6CxEj/vsVKvCnPWgRgSA4xTsAE896zlJMZHIAcLjIp2AeAOlN5Dk46U7GOe9HqAnQ4quB95e9TqBzmohjea0j1FqU2JR+elC3RBz/DUjxiR2JHAqB4xj0rdWejFd9DVt3Dx5HepBwrVUtQAnymrZOR+FclSNpOxRFbZLMx6dqex3Z44psAKKQaefuHND1DoZ80m0lQKbayEMRT5xhsjvTYOX5xx0xXTG3IT1FaTOQ1V3kxhVOc9KusgD4IyKqyqEJPpRBxvsDIcgHkYIo3BySKVl3ZbNRRKWYqPzrpST1ZJYblFPYHmphMCg4xioXzEqjOc1D5jZI7GsnFSRSdi9CRnpVouNo9az7eQ7severRbHOeBWNSm0xommb91xWXO5BGeTVwuWXHc1UnG87TwBRTVnqDI0+bBzgVakHAIPNQIoKYU5xUsci7QrcnNay1Asw5U5Peri4YGqyFWA549KnVwVz0+lc09WUNHB5p/qaQc45qOTnoeO9RZsQ9ehPekkI2jmmeZg47CoWfccZoSdwCVwpzVV2DNkdKmkIZcZqsM+Zg8ZrpgtLCZMWygxyfSmgAZZhn1pW4PB4FNMq4IahJ9AGlkIPHepNoMXyHj0pi42/KODT1x2FN3WwyuHZG2lutWcNtG01XkUBw2OKuwjKA9qJtWuJEiAleevWkACjPc0M4XjH4VBJuZhtPHpWMU2Msqo/E05RtBpqBsDJqXI2mok9QICo3Hmos8nI/WpXGwk+tRbx3FXHYACnnJ4p+MJtzn3pAQQcUbuNoOKYbEeDgqegqIAOeBxUoG0800FVYgVd3YREsZ3jHNWVwH560i5yD3oIw4PSk3fcFoPb1xRngA0oGRk0N7VmUGAc45PrTwAeKReAeKco4z3okxFSZCkgI6VIhwpzzU0qgjPeoEQ9Ce9WpJx1AUHPB59KmXrgmokj2sMdfWnlcHrUy3GOJ+bpxSA4JJoB2mkYen40la+oheM7u9NkIIpVwPY0yQcHFCSuMaoyOfwpX4A6fSk3naP50pGce9abPUVx0ePLHvTmBI+lInA560Y45qXuBXkjyflNTA7VAqJchzkcU9W5qne1hehI3AwOtCLsUClJ+amlgxC45rPUfUevB4pvcnNSZ+Wo3GVwO9C3AZGw6UrqBxmoYyQ5FSNkv7CrtZh0EYAdTzSg80cEkUuCWouADnk/ypp6EU9euO9MJ+aktwBKkPSo168U8k4FD3AjIxkdqAAD14pz4Pao8HHJprUALHdSn5vpTWBBBz+FAJIANVYBAoJqVRhsVFyG9KlHIyelJhoS8kZ7UcZPTFICCCO1J6AH86gNBxP86iLYHrUhNQnqc00Aj9QaRgcYFOLgHHalXnt9aq4ENsfmNWHxgnPSmhFQZFPZdw+tEnd3ApI+6Q/pU6r1JNI8YRhgU7IGKptdAGnCgtQhypPamzgkHFEOWj5oS0uw2GE7Mk9KSNi44FE/8AqmAqO3bAxWiStcRa2/KajjYhznJ9Kl/hAFMJwwyKzGWFXpz1pg5c0/OQOabtHJ71CAAM/WnADYSelMLBT6Uo+70o1BEZGc88Ug9OwpXb5SBQMBcE8mrQIZsDrn3600xKvPapB8mAOlGATk07hbQF+ZDzzRxn3pUOWI7UwkAZHUUAJyc+lQq/zEE81OT+7OBg1WC5bnp3qoiLSNyM04sC2B0qupAPBp5+4WBqWh7EhOR9Kb90fzpFYkc9aGAxnii1gFUjGcVF1kBJp5OMDtUcoBdeMU0tRPYsggHBOaeOvNQjG3PepVOMk9KhoYr8Cq5+9mrDEMTiqZVmY9qcAZOnJxjj0oJ601SR/jQSAp4zRbUCQcqCppNvzcU1DmPIGKdG3c0rAKwyaTFKeppB0PY0ARHlyCKlQgjimMucsacgAU4pvYCQ9Khdtufeph0PpUbKGH9aUQsRKwBBpfN46dahb72FpwA285rXlQhyupfHrUj8DIqBlwwap85jpSXVAhw7YFKMEmkX7tKBx1rMbI3AB6Uo5GPah22j3pBjZn1qugdRuSSSODTweOeKYRglqcDkCmw6jcAN15po6c9c9KUjnPWhuSOKYhGBAP60ifP8tSEZ4NRhcNkfjTTAdnav8qhA35PQGiZsYzTVkyOOlVFO1xOz0GumT9ajIxnsKnHP1pHTPTrVJ2Fr0KhTJyDyaFAVsevensMY2j8KXaTkH9a0voBIAFXGKliHBNRjA5J71JGwGazkmUSDhT60gUHHagOGOAaAeD2rMBHJA9jVV1wxbuamkcbqjYBjj16VpHQT3HRgtgHqO1TFCBio4jhqsBvl9PcVEm7j31IpFJGP0pmCMc5qSTO496Yc7j6U1sJ6sTr83YUud5B7CkwWXiiMAAimO47G5CVpwIVf500fX8KMds0hCkFuafn5f5UzPHHSndVyetJjFwQuM0AnoetCnCDinDkE0mA057ULxSjpj1pTweOKQC8MfagcseKO9IvekG48feHYilwDyeKQYAzRyPrSAG4Q1UmQMp9frVpuveoJV3YFXB2C5BHHg7s/WrA9KiAMeOfrQJBuq3qJaExGaRwCfSmebnn9KkU5HrU6obuNdfQ1EMg5qyOtRkAgUJgNOO/SgqrNnH4U8YA96B16U7huJnHXrTNoOT1Bp5XjFJjHWi4WuMZMjGfzqq6gmrzA46ZzUTL8uSKqMhGfIDnnt+tRE7dwHXqBU82MnjOPWoSM8kYFdMXoJ6FiDLKRQw2kZ6dqbbuw6dKHO5zu6Cp6huMk7Ddj14qJzhXHXg9anwV6YFQScI2OcjrWkRX1PWQuepprD9KU8jk4waT0+ua+YV3oaEysFTHem5x3/GiM4c5oP3vmHWnJN2aEJgdc0nBwc04AE+gxSfLxUtWGgGPwpOpo6cmg8c/pS1AOvTrQOue9HuDTlJweKS00AM8elPwDzntTTnGMfjS42gEdO9HmAnHU9MUh5ORnAFAPy4NKpIDDNMBMjOcUwHBNGflNNY5XFJXQEbDMoNKuA/Haoxlvm9O1SdOe9W1oMcDhxThy+e2aaOnPWnDoB61IDzgyHPSl5DEdjSZ5x3FKgyzHFCXYRIxKJweCKYThQoz1pVBZgpP405/lBz1NX5iYr4KDA9qiX7uWp53bADxmomBGAaV9bjEUH5j/ADpYsNIPrSt8q47UhGzGDQGw88yAADg0qgliBQpCnce9C5TJ9TQkuoAg3OQaaSFJHU08fIw496R+ZMgVIELZKkY5qvbI+4hqtjlxnijA3ccVafKrANB+enMcsMdO9KBhSRzQueazWugEYbMm0cYpTkOfelUfvCaRV3P83ar0sAwcAinKCFz70xvvknpS5JGO/wDOqAm+8QT+dIv8Q6ikHzAA9aRCSdvSpAepzlfTpSpjBzSIcNS/xVLVrMCCVOdw61WU/OB6d6utycH8Khe3/eZ6VcHbcCRSC9KFDMQD9KEHltg80oK7SwP4UuXW4PuRuQDg9j1qq86KSAeaS7lYZK1VijxJvfnNbwpKerA0x8yD3FIu7JJ6CnhgQGHSm7sZHY1jJWk0HoLkY9qTeN68U0nCEdqiaTbg9x2qoxu9ALTEcZpjNtfGeDSKQ4yT1pHTeMdxU8tnZgPzgZHFMJ+cc9KEBIIPbihsBwaErXGTdSNvSh8bQRSxEZJPSmuN33aT3FsM43bTUqHCEVBt2yhs5FTHLj5e3Wh2sMGJKbewNDLhl96G+YAelI2cZGaVgEfh89qb0B/pTmGQCaRsBwOg9KEIMYwT1pCBvz2xStycDtSHG0Yqthij5qAfnoUnFC8E+1QAvVuKF+8SRQCV5FOUfKfpQnYBgIDZxS8Y+tNLcEGlLbiBRYBxHyrnpSMMgDPFK3XbSHkYzRsAisPmNNOQpNOAw4B70MOdtOzAgcfIKahGQDSynLAHjFNTBkraKAkI/eCpS2BuX9ajPzHIpznK47VDvbULiROC3160rvsVgehpijawqRkEwIPajrYOgwt8tNhDBDupwTHy/jTgOAmeKW2iAVfuEgc5po5TJP4U4HaCtMb5TihW6gJg/e9KYzEtgd6ccmP0pEGBk1WgDYlYvhu3SpcYpkYbJNSSHABIpTeohGGBjNKw549KVyDjHejID+1LWwwP3R9aSTICj1ozklaM5OCfxo2AaMH60rEHI70i/K7E0vfNDYCZ680p+bA9KQ4LBu1BPJI6UaAQyjJyKiZNo471YYZyM1HwW+lXF6CGwKFyD1qXqpNIvegHOQKHeQwY/LkdaYXAyPWnnpgcAUwAeYcjIoQCKMvkU4na/PelA54/CgLubntTT1AVc4JPFKi9ew7UpHyHml7ik3oCIlPJBHSlQbVpSQM4poYFRzTAc2DyKUjJJFDgBQaXG3A9aQDSDg+tIu4vk9KCwB56Um4E9eBTW1wJeDwaql9r5q0ccYPaqN0pL8DjFOC1sAjujybqYeZM5psIble4qKeXy5MCt1FXshCggS5HQVOVVScGoA65yMUryB42K9aqUXsGwyF0R8jqetTkFuVrNRjvBP4VoxTDYa0qrRNCWxOGzgmkkPXimLlo8inycoK5mrDIlclTzzSo3ymopYimCDyTUyqFT3qtLXQajQSQSeKgCtvOelWgv7qo2T7hBqovohFVkIz+tQbtwZevpV5yd+McVEkas/A71vGdtxWH2Z+UDHPpV8L89QwRgPkVNkFvpXNVkm/dRYBQXIxxTHHJAqUEA5zUUhwd3as1cGUrkBXHrUcRLSnHbqakuORmmQfu37gEV2U37hPUtkfIWxVWRPNjBxV1W/d89aqzsTH8nBzWMW7gyk0eCN3ahNo3ZPSnsQ2Mdc1G8ILcHrXbTlbRkNBM+QpHamYyQam8sDBPSmhME0JroN7iq4iALU9rkCPd3z0qs6EvnPFPjiD5J61LStqxXJkuWJyV4pX2EFjR5ShT6CoScx4I4FQlF7Fj4YwO/XrVlLdAcmqsMgPHcVaabYvApT5tkCJAoUkY49KYkuJdoPBojYlC5qNCBKM4rPla3HoXtwXt1qMr3Jp7DA9hUYcP34rJJ9BgCGBqBlCk+tT7lBI/SoZR3HSqjvqIrFvm64FJkF+B1700jL570fdPFbpCJnXEWCefWq52qw3DmrEYyCzngdqaxDDGM0Xs7A7iIygkN0pROoLbVwMUxCW6UhRmbA4FOybC4/IOT1qeN/kxjpUSbfuin7CB15qJWHuOR9zMaE3F/anxqF4HU1NtCfWobS2AXjAx+VOOBimA8YBp3oayerH0GSjv61WlGcY4q031qBgrNVwYrajI2CjigtyMdacq8HjgUMvGMfStGwGNy+TTXwrYAz70pUlgCTilk7Y6VSYACQPm496cSMDH5U1lJUc9aYWGNoHSptdgTghlx0FCY5FIMbfajjOQPrUgSAf/AKqUEbeaZuwOO9BGVGamwWHscjFRkjsaU524HWgKe/NGwdQye/FBYBsnvSspPWkMakg9TRoAFhjIpcZFOICnihsbfehNXAaDk9KR8bTzRjBzSAg8mgCMgEccYpUG4Ybj2pCPnpwwpIrW+gDgefan/eHBqFOuKlHOQD1rNgRTZCHb1qOByeGU5qaQZOD2qPd8oIGcGtI/CGxISSeR0oUAj+VA+YDNNQ4fAOKgCVAcnPT0obk+1Kpwc0g9fSo3YEJT96T2peSxHbtSj5zjoaXgDAq79AIwdrc9KeT1IpCMnBpHI6elADhwuR1qMgKc1IOF/lTH5IzQgETO/tg1KfuiogMHipVHBpy1ADwKbgEH+VLz1zTWPGRxSAaw5wPwpAMAGhm+XOaTkgH07VS2AUnDc08nKgdajY/KDThyPY033Ak9hRnA/Dig8DA5NJk455qUFhSMjiozg96eOWqJxhiaIgL1IyKch+Y57UgIpozn2NUrNBqSk7uKUHI47UxeuOwp2cdB+VQwGsRv5qMEsTTm4BJ71FGcgjNaLYBzjOcUo+7j0prA4GKVR8uB1o3QakUobbjuabGhHXjFTM3zewqNnA59apPQViwB8tQOMvnvTlJyPSnONzdKlaDZIvOD7Upb5ulAO1TjpTc55POakCOQt5mOgqYD5eaikBZgQOlSA8dKb2EtxrYB6VHuAJzxUrYB4qBhg04oY5H3Ek03cCpqEzbd2OlRoW5461rGPULlgy7enekUneWHQ8YqHDl+elSxrlcHpTasIkb/AFYxUDbt/B61ZbO3HU1WHyvj+tTEY9xsUU/eCmPyqQRgqAetRTJsHH5VKaegCEg5AOTTh90VWGQ3b6VYiXI+bp1q5KyEOK7iKTlhg9KUNz9adnBOKi4MagOc9qeXwM9Kah+XND8AYpdRj0bK0hGc8U1Dt6d6lGDSegDEU4pCuTipDxQB1NK4WIiwVMfl70sRHTpUch+bGOabE4I+hq+XQC105puDzRnFBqAGOSelOXGBimsSPrTVbjGc1VroCfjBpoGKXnaDjijjGakCvIuGGKad2QccVO/ALUhO4da0TdhFcnf+FT9ExxUQVVzU/BUUSBCjAUU4HAzTQOMZpSOgrMY2Q+1Ih4+venEc9KBjbjFPoAxyNhoQge9KwyRj8aZHzk4qugAcBs+tJnDU8nj3pDjjNO4C9ADTE++c9/WpRnA5qJvlYEfiaSF1GXY/d5P5Gq8eCg56VLcyfux/WoIj8vTBraHw6ilYl6fWgnjg/Sg8mm7unc01sF+gnvigjAwO9CjB6fjTh6UwVhoBI4GKDkcE5pxHBpjDBPYUJsTFSQLIR1qwOV/pVRepBH4mpTL/AA4xUyWuhS0EeMg7l4pig59z3pwm+bB696RwxHA601fZi9By5Hep42461Xj4yCOKlRgDUSQ0SuuR6VAVwvJ+lS7uOnNRPhsc0o6CegoPGcU0Db04JpzZYbccUjYC81SAcgIA7elLgAE9D6U1eg6ClA570mNdhU5U5HJpW+78pNCtxwOtJ1A6/Sl1C9x2cpnNOx8oHemYwTT1OcikwYp4FJ2zj8KCevvQBjvQMUkheaFGOB0pBg55NOBIHFSIXGcdjTsdKao604DgY60gEzkUxwChOKeemKQY79KaApbtpO7pTSyk+9PuQoOSKhkdQmR6c10LVC2JVI49anjbI96qI+cHpT0O08damURls+9MJBI9KTd0NN71FguSdAeaYDzTSxz/AEpQSD0p2C/QkBOcg03jgZ5pgfB+tOVt2O3vSsA5gTx2pj4WnE54HTNQyE846U4oCnOAHAx3xUKAMQOvb6U+YjOR0FFsRvA7/wAq61dRIvcUIyqQvJprEqTkc1cPTPT2zUDkNwahSuNxGH7hyMj+dQPnY3PGO1TsV27R1qCXAVwccDitIb2A9ZwM5zz6UMQRkDp6UgQR7snmkQEqxJFfNdNCh6D5ge3vTpMlhzwO9R4woIPJ7U9sbMA0NtqwCNjOAaYx5AHWlxxgU0gr9aTkhjhz940pGee9AwQM/jTt3txUXYxuAO5pckKAo69aM5xkUAHPSmpai6Ducinhex6UwAk/hSgFuR0prYBrHDe1ICMnPShh3460h+YZHSp21AYwH50jj5cDrQee/So5GPamAdOKVTlcHkUwseBS4YHA709UMkTOSSKcPvbqTPy7adnsKmwDhyCe9SrjyyehqMNiPA4pwU+Xk4wapKzYhRnYTu5oPJA7ikBztAFB+8MjrSvcCVssFHeo5ADyeo4xQGKvSMTvNXJq2gWEbnB9qY2SfanydQO1NYYxz1qHe9gHHlMk96kYDyRUeMkelOY4zzT7gxxYZXio8kue1Sud0aniogPU80ndOwDGO18mnLg5Y9MUSgMP6UYwBST0DoKNypx3pSNq5INA6que9Eh3Zo6AMJwSad90Bu/WggbBikU8kU0Ax+Gz1zSIAHzSt70ik8mrT6gSEZbINI2RIuOvrQOOR0pCcncOlSgJFU+YQT1pB3yenShTucMOlJjBIFNgJy7inuxzyOaTO3vSPliCetJS0AaQMgmo9rCQ46Gpsbn56U5towQPahNoRl3ZIJHc0KuIQxqzcQ7jjFRBcrtbPFdNOa5LXAcswEIHegSEAE1AnLbT2q6sYK+1ZzaWoyMkncce9VZCQFY+vNaRAK49e9RvEuFUilGaT2ArF9irgcVKXzjs1OVU3YPY1IyIHB7ZqpTT6AIPvZPSmMucgVL0Y56VE+VyfSs43Yx6HC7elLnBGKZu3YpSuGPvSa6iELYkOR+FSxkDd6EVXIJbNTAgpiheQwU8N+lKMjAPNJg7efWlJ7+tT0EIxywHbFNk5/xoJ+bjp60rkEjHpTV0MYTkgHrQVOMU3O4huhp+ST1/CnuAbsqAcZFKSAvvQMbcnp2prdMHrQA4ngDtSk4GAarktkelPVucU5RAcwG7ignAAApHPf1oXkZPNSgJD97NI3BwPrQfuDmkY/Ju70bgwZgM5x7UZK4YmkHIPekIJ4Naq3L7wupWkBZy2aRcqcjvUzjt3qIqRJt7VUZaASqQG+tSDBB9qjXG8GpQMZPr1rNjImO9dw6inxk5B9etORcEj1FJEdrc4OKE0Ie3L8U1hlvlpw/1hOeKANj/AOFTfUYz3NDAE57U8jLHHSmsADweDS8wGN91iBxTV+5xUjDavqMVGn3DxV3AfnC570pUSIdx6DiohncATmpegGKWwiOPHT0oV8j3FOZfk+XrVNi8Tk4yKuK52F7FxcZpQOTUcThxuxUgJ5bHFS007MYwZxinA54pOg96Pu4NLcYgIxSA469KVhxk8ZpD0FCEIRtyc1BnH1NWH6AdahkXpmrXYQqH5frTlODzTEGTinkfPQ9xiHpz1peOKCNx96hbcFHtQrASscjApYyUyG71FnjrzUjYbB7U1oA8HjHrTR8w+lL3po4z61KsAMOM1BHIrMV7irAAIIqHyAH3DrVRa2Am+8AtKc5FRIwL/SpCcn3zQ7gIw+bGKjVcMQakBySaaDuOfShaICQngn0qOTBANP43kmoiN7EA8CiOrGRkYYkcE1RuYWb5u1W5CRkjqKjdso2elbwbVmS0VBEWj2g4PY1LHF5aMhPJpSo2hhxTs4+brWrk3oFiosWDtIqyqFSFA+U01lbzScdamikw3zU5zdhJWJ4lGCB0FLtDjbSrjdxTEPJOa57tjIn+aTaRwKCxJNOwCevNQyODIoB/CqjqBZ52cVFIcxfL1FBfKZB4FQGTfGdoyaqEW2BICejUIFBAXrUCuTgHrUiYDZH0q7MWhYV+pUc09WPFQsx3EJ2606NW3Vk43RROvr37VFdMQmV/KpsYwabJjn6VEdxPYzH3Mo5pIeHyxq1OqheOtV1j6HHWuqM7kvQuEh1IFVnJERx1qzGNqc9KidCayi0nqUyiB0559KMspxjjvQVEbnJ4pqylicDiutbXITvoSPlo8ZpyjeMd6gjOckj6VICEGSaGA5kwMevanoB0HSowwDk9zT1JAqGnbUZIzDaQaryJjPvVhcYHrQY8DJqFoMooCkhIqcz7uKinjYtnJpY1VRk8mtd1di1LO9vLx2qABt4I/OnI4PXirC/d3dqzV49CnqSsGZR+tRiIqQQTUgbv60pyO1Y8zWgbgU5xUb85FKzk/dHSm8t07UJMZXKBXI61Gx5AFWHjPYVVMbbuO9bQasSTpgjA9KQ4UketNi3qSDTTHlgzHPtVWTbC+gka4OPWpCAmcnmnfKACKjGXkP6Um2wWgyJv3oGc+tXJFYj5aqeS27n8DVvJEfBpTa6BrYEYKMnqO9PLiUYU1BsJx6GrMSbAfUiobGgjjINSHnHHNC5HWgH5sVi2wBhkVXPHbn1qweBzVcjEhz+FVAAGSvvQ2doPQ04ZB5NRs+XwelWtR9BB90t6U0MXTJ6U84CZBquy/wB3qKtK7sLYmDg8daR1AX61HG2G5HNTMuduTTfuvQBYxx1xTgMAYpqDA6U7PFZuz2AUY4FIw5HJpA2SMDmn7c/1qdgHDoQKRc4OaVVIFNUkEikFxxAI5pcZ4pG64pcZxSARuTxTQfmI/GnHFMGBTQAevPHvTdw34xSnnnNN4Y5q1tYBOj5pzAYDEc+lM5JNOzuTFUAK4xnH4U9cdqjUADbT0P5VLWohjn5iKarYbbQ7fMT3qJXBbB61SV0Mt4zimYO6nKRimk45qbCH7xuHpSknOPWmberetPWpsMbwox70AflSkDd9KaCTnNNIYh5PrTMDPvSrkZOetA6ZFMkU8KO9I3Az3pQd3HehhlcCn5D2IxkgnNSrkD+dRnHenKdw+lDQxSOuaQj5CDSk0E8ZPFFg6ER6DuKCDuz2o7kUgOR71RIMAcccU8YximbgQAKchBJ54otoMeODnvRjk8UZ3GjOSalXACTjjvTG4wvrTxkg1Geuc5p9QFAyuKB8vFNDcnPfpS09gJM457e9P+nWoieg/SpV+aoYWI5R61UQjnFXHXdk1SX5XP1rSGwmT5+UZ603dgZNID3PWg52k96dhjJOBzVcFt/X8KulAQCaidBz+lVGSQupInKj1p7LkEdDUa8YAFKDg5zxUPcZMgyoHelIAP8ASozIBnnpUbTE5xSUWw6kxkUAGojPwcCq+4sTk81LHFuPNVypbgPVieoqC6f5eOtWwoAFVp1yRnpRB6iuV0Q7ASPrUsfD4A4pQoP09Kdjaelayl0EkxcZfpwaFUq5/SnoCef0p275jkdKi7RQScHIFV9hZwf0qyDlQaZt5BqU7ATLwBnrQ6ZGKjVvnxTycios0wKEqlW4p6MSm0VI4DLz271E7BTx3rZO6ETJy2CeaXn8KhDESfWpuevrSaAI/vAVJgdMVHGMMPSpT97NQxkQBJ5NSxnnBpm7J+nelVuen50PUCTjbSE/Lim5BzzSMwBxSsAxu571XgB3Z7VaIBbAx+NRSKEb5ePWtIvSwrMlU4YnsacTg5HeoQ3Q+tPJzxU2GB6+tMU/Nz3PFOZiFNRqQTk9Ka2F1LAPGM9KXIIxTU6UA7e9SMVsc1XlbYuanJzVfKsxB61UUIjRmJGTmrg+7xUCqOw6VMn3ac3cLMcBjilOCeelA45pccA4rNjGt15pMgDA7UrGmg5OaEAh6E96iTuak3ZOMU0DB4q0ArcYoPBB7elJg4yaGyUyO1AvMeDhc0x/mTpQrZHPSnH2o2Yyg+STk05SdpxxT5owr5xTBwowOfWt07rQl6CD5qdjqKRsk4I5oBCnJyaNBLUcAB2/GmscMccUm45w3fpSnPXt60JAO4JwaY/GB3pwO4YBpG4wfbpT2HuQknv1qMNzwfrUxBzwfwqCVQuffpVrUNhGchsdu9W4vmFVtvBOM06J8EZ70pK60Fez1LKrgnPelUck9M0oAPHpRnPFZFCjk9elNPLhqlX7tMOMgdu9JMB2RnPTNRDLFs9DUxB2ECoMOg5oQbEiAAZzThwc0in8qTGTwaQXGqeeBxT+r+oHSkC85ApecU2IBkvk09cbuO9NXp9KeGx/SpYICCRSHgdKXPJpc9KTGA9aB97AFA4oJwB/jSBjh2py4xnNN42Uo5HvSGI/UdKQnAIx9KeVyQSaaRnoKBblK7UsKpYKjnpWjdsFHTis/IZgB+FdVO/KQ1rclj5P4/nU4jPXHQ1HGm3nrirI5T1qZMsYrAc5/CnBhu+tMkHcdBSr0B6ipsCJABjgUjLyWp2ckCg8D3qQIR15H51Io5pgIJOOlPAIOKpgDnByetRMM8561LtJzn86jcdcDjvQn2BaFCXAYg96SMFWyB+HrSSj96Sf/wBVEYbIJ7d66ehG5M7HG0du9VzjOCeg6VZLcE9qrsMZcg8d6Ih8xSCQRj/69QuDsbucVKB2OR6mmOAI2wMDB/lVR3KvY9XmBxuJ4J4pIBk4pzDcBnoKeq4O4A4r51WT0GJhdxpQBtPSgAsx29MUAAKRnms5Oz1QaEJbLDFKwzz2FB4/pS9jmjRjFU9cUwybTwKfjJGOBQMAnil6gOU5GQOKUdMk803IxxnrUbP1pO19AZZGNvuaTBUc9+lUPtuWIHarUU3nKGGOK09nK3MA5gVGD3qM5Un0p8hyN2aYD5hHrioYIawIxUcgz0/KpWO4/SmeppLcaIowQST0p6n56F+uKVQM02MenzPjNOBAf2qEfK2c8VIGAIAPJoaW4iRjufAFSYJXB4xTBywJNOf7w2nNJCHJ8oJHNIBkEk/hTh8mQfSmYxn3qthAq/KWJ6ULzlieabn93+NBIIwKV76jFA3n3FI2Nw4zSZwCTxScYHrUgSEjb1pwZeGI7cU0jIHrSsuAM1WquA4HkZ6GkYDJoY7sLjkd6b657c1N9LAKCGGSOlNGSOnenL8wwtIGABU9aYDUfBOevanDG7J7UhTDgNSsvzHHQdaFutAEyqyZPQ0m8FiR2pJQH+lNwBFuH0qtGvQBJGwPY0ikBfrRKpKD0pv8A9aFHQCYNiKjK+VjvUaN2p20Fc5qm0kDHKSIiaajfOMkUBmHy9hwacVBPHWk7JBqDN+8IHNLuBOewprfNgjimbSo6/hRbTcCZSQckdaVsDGfypAu4jn6UpBzz0FTIGDYZzkVTnYIpJNXA3JY4xiqN3CZUOOKa3V2BlvKVkyuea0rJ5GX5u9Rw2gAG5fpV6OMIBxzXTUnFKyBXsBJ8wAdKbcSBGHoamwCQeM1DLF5y4bseK51YHcjQjOaTzGzyM1IiKF2n8KftU7gattcwwVlc1FJlQTnilaLH3c0jcIATzUXV7oLkSPvBwORT1kzkd6rn5JwR0NTDCkselaSirXESKeRT+Mk9hUSg9e1SouH56Vk+wxWbc+B+VDEk49KVhtYY60OcNnpSvoA3B3fWlwCW/SkJG4YpzcsPehOzAjGADmlAGMg9qVgDkY4pq4II56VUWuoC8lRSY3ZGenSjndgdMUo6ZzSb1FuRnO3BoUc5p4+6cikAAB+lAxxO/r2oxxmjoMetL2xU3ADgKKQjC89aFz3pT8yDFPZ6ANU7VPrSgZGaReRilB6rRqAxzmQHHFRzfeBqdwOPamSKFTNUmIaRgZ9akBJUVGFJAqWP+770MYjcEAUfdOMdaV+WwO1KQCAe4qebqA37vJ709cYbNNcZI46U8FRwe4psBsf8XtTVyWIxSoMMTSg43MKSSuBE3QrTAcjANTFeCT0qFBzntVoByDpmnA9QaiYkuMVISOoNJq+oEmMDnpTCoGdwpT05ppBPGalaAwUY4HFO3fIR3pG4xTSPlyKpMBM5UZ6UjNkA0vQCjHJpoCBWJYg1KRgDHemAAU4ZBBqmwBmIbjpTXwzUODu56etRtlW65FCQEigdf1pxwMelMQ/KaFbik9Rilvx55pGU4pVIH40pOeKOoiDYQSaVZM/L371MQNoIqusbeYxxxWkdQJTJhaUjIBH402JMcGnnKnnoOtR6BcYuSCx607tmkX5mODSqcKSab7jGKuHyKFJLcmpAuFJoYAcii+uovQRvam7dpx2pzDd1700DeSO9O1gHYG7mothSYt1BqfPy8dqidgGx3oUmnoBFJ785qJsEAdqssvT19KiZcqwIxWkZWArSAhaRSRjJ4qDzixKkZpJWLOFFdCi0TdMsSMpcAd6cihSMc460x0wq5GKkBCAGpt7o+pYRtwOeo7U1HVXwe9IsysRz17YpkiYlGBnFZqK2DYZNOqSn0NVFDSTll596sXtufK3/lRaLsgy3cVtTajG6E9dCTBCkHv1qFkCDC9DVjJPBHU0uAcis1KwPcpbfmPPQUsEm/PHFHAZiT0pY2XGemfSuhtNC2LUA+bLflVgDBNVVH7wEdKtZ5HHFck9ShGPzdKZw3zGnuxbBXgVCScnH41MQEkIAJxVSViQMcVbKEjGc5qrJ8rEZraFriZbhbMIB64pxGUJPFVlZnT5eoqRGJUhutTKOrYynMuWyOtQKuOCOKtTYzkUIgK810QlaJO5DgqR6U1huP0qTOAVI496iYgJkfeq4LUQ5T3PQdqkUEtntVaIjd8w47g1a3gnCjilNNaIOhIMZBp+dxP0qNWVT71JnPI/KsnqyirNkyZHFRGTYaluQcqfzqFlBPzce9XHVICwhDAHb1qyoBjxVYAqo9Ksb98fyjms5jGA4YVYc5XHr2qFPl5NPILHINTJAG0jvSBQEwBT8A9etNC4HA61KDRjQ+Mgc1CQ2dx9alVcEk0hGetUnYGIFyCPWoyMcU5n28Z59KjYlh71aWoiPrJjPTtUy4ByeKjx5QPfNOR/M4I4q5agThRsGeaRI97HnikV1zt9KInILVm00MlCjgU7OTx0o7DFJ+FZ3uA8NkdKD1yBTV49qcM4/CkAZOCTURG589MVI5welRM2OvSnFMYvVx6UxupxS4A5pDkkiqQDA3ye9QMNpI6mpsEdT0pjNkbh9DWsXroJoSJeBmpnOQPSooehyacvJJP5US3AeWwRnrTwM/SoHG5RzyKlQkAE81NklcBQMPkVIOeemKbxk/zpWIXpWbB2HE5AxTMe9LGCOSeaDw/PQ0bAKSMUoPyk00jP1pwBI4pAMJOyhcEe9OYZH86YFKmmrWAQ9OlAwFp4PzcjimnkHHApgR9acAAOetNJ54HFIDnoasCRQB3pzfKOOtRghRjvTuScnpU9bgVnJQ4PelMW/BHWnzr8vvmkQFI85rS+lwsTIuxR3xTm5Sol5GQalf7vArNgNfoMdKepG3IppPy5x+FKv3elLoG4Hpx+lM27RxUnY84zTDnb1oTAYATnJpuSH46U4nHGOtJznNWgD7r5NOPPSkk55FIoz06dKGAdQSadGeOlNYYHHH9aEJJpPVAKw+b2pGGRSnJNJn1oQxhGATUZOOKkPf1qMjDgGqQhwAC80sbDPNMU56mnbcfWmw8yVeMnNJxQvSgEnqORU7AKDzio2BOMcY6U4cDmmMfk46GhLUHsKOvvRkscjpSD17UvfNUA7aTk1IBjHp3pij5SM8VICCRxxUMEI+Np9Kz14lPGc1fc7sgVTUDzSauGiYmPIIXmlCllwKXHNO2kCi4w7Uxhls54p27A703qRmhILgQNw9qidiQMdKWTIfPamAcEE1aEN2knJPSpUhLLnNNT72KtKQBg0Tl2AijgCmpQADgUZ5zTs56Vm22NIUdOlVLjr04FW+wqnccsPWnDcGJG27GakxximrHnDDqKQ5DHFaaMCUDaBTiMY44NM3Z4HSpTWctGAgXjgYFQyttHXmphntUUkXzcnrTT11DYIyXJbFPf7vFAXaABQAMn2pPuIYykioZgA+BVlQN7YPNBiDMTTUrAVChwCelSKxGOamdAUxio9mD/AI1fPdBYkjAJGetSSYaol+/uJ61ISAvvWbWoJIhRSoNAbOB0NS7cgVFIAoJHWmndgOBw34UFTnNNyRgGpByKHuG4nO7p+NRyZJ+nWpqikyO+M96FuBEhBcCrGzDVXiBDDNWDwQf1py3BIRh2qPG3qacWAYjvUZBLYJoQMnU5oC7m4psYwTUgwg471LGNxjPOCarCM7+atuKj245704yEMVSCRxUqZxx+ZqMkjJp6t8vFN6jWhIuQOeaUHI601TkYx+NP9qgQxhnk0g4GO9POTUZAFCGRPw4Henjg1CSQ/A5qVc55NaNaCHYNNbOMCn9un40wjPJqUMQjAx0NCE5IP50vDGlIx04p3EMlAOaqqMkjGO2KmnfC8HGagUH05rWGiE7XsKQcnntSdv5YoxtPPNL1UADNXuIMd80ZOD3oHueQfSmKMjmjcbHr1980hIVcU7kEZFNLZ4FJAM5/Oo5F3Mw7noakzz047Uwt09T6Va0YaPRjSRzxzTQ23DY4ps528BhnvTIzuO0HkVajoT1LaT4yp7EdKlXcTn1qvEhyRgDv9KuryoAGPQVjOy2KQoB/KkHXIp/fHeg8j5fxrMYhyFz+lMfIXANSAAYz1qJ+Tx0oQAMbeuPpT1x2qMD5gORmpl5H0oYLQauASDzSEAA8U4g5pMDB96LgJk5Bp55wCKYO5pw57cUmDFByeOlAPHHagcigj5RmkK44DaaUDcuaQcg5pcHBpDFB5pQMjpQcbaBwMCkA4HApuO3epOMYFMPse9JAQzR+YpBrMkhaFz/drXY5XpzVWRPMUg1vTlYUlfYjhYMOeanwAvP5VU+62QKkSQ4GTmqkru4yU5JpFPHpSlQwPTBpuOuDUgPB556048k+tMwfTrTwM9qlhfUZsO/PUmjzAsgU9ac7bcnjGKy5nLyEoefftWkI8wmzTZxjrx3qKSQY681TV3/HsKjYu2fm6VapjT1CRwX55NLG2MkDketRchgAB+dSRufrmtWtCSXrkHpTio24I4oXjoKlJO3HrWbY07lXnPce1Qyg7W9ADip3wCSetQSKQj4545rSItb2PWs5UD9alB+Qjtjk1EMkYA4p2Tt24r53mSuMQEhM570uAByetIQeBg0LgDmolYY3GSfakAzyaeTjt9KbnmpWwwAPXtSAYGacOgozx/SmwEbAXjrTGBK4x1qQ4ZuBQAT7UXewGY8flScjg1fs0BiI6fSnSRK2BingbAAOla+0SVmJCMoKtioYxjPpU7nGMflUbccCsrtIZHySfSmkkEinHgUh6Z70hjCR09aViQBTBwc09BuOT2qrAIxAZR61J0bikSMNJuboOlPUZkwOlHoIk5LAfzqTGHBFN6tgn6VI4K7aS7gxsmVc55pGJHHbtQwycHrmhhkgGknZXAaBgYPfpTB096lcYcAGmOCrZo2AaBuODR90+1KRg5z1pVwFzinpcAAOCakUFlBJpi/NkYp4JV+elN2Ac5DOCOw6VGTljTtw3ZH41GFJy3v0paboBVYowNGDu3kcGnHlQDjOaQ5I46CktgEcbmDjPSlfIUEd+tIpJXHoetCrvdhnimm3oAiDK80EDafSnZBYj8qY2R8po3QClf3Q+tRMQDjtUjH7uDUfc/Sr3YWEQZap2AYgjpUMZw3PpUo+/kDj0oklYBhBzt9+tOc4GB260pHAOO9NkO4+1RcBpOU465oY4QCmMSOO3Wnj5kI/I1bVwFLYIAPPc1KWO0rUIG1hmpASWyelS20gEHTGKQYJGfw96cf73bpUTMPMVc0kmBKQGbAHPtT8AqR6VGDiUbeaeud5z3FO4ERyG479KdyVKk0oOX46CnsFGee1JbAUW3CUc8VJHl2bPp1pZCVTpmnxAsg7GrvpogFU8dKq3BwS2elXWO2s64UsTjvRDcCIgvIj9s1ZlXcoYHjPNV8+WqocHNW9mIuDnitZ3sA0NtA7rUwJLBh0quo2pg+tSw/KcZ4rKSQx+fm9qGG4k9qToD60DgjJqE7BoCjB4605OSabj3+lOx8me9CWorjT1PtUQcAFqkLYJA9OahA+ciqSTC5KGzQSMjmoo1PmE9qXePNINFkMnznPHBqNjtFPBPA7daryvhgPehK7AlJGB60DLCmhacuOOaGl0GOHSlwOopufmpejcdKl7EijBYUh/TNKMZpN2ARijTqApG4EjtTT9z3p2eCvU0w8DBpoBUO3bT1ID57VEGxTd25jg8UW1GPZsN1o8wMeKgfdu6cetSRoBz69qrlSiBNkhvWlOATQucmojIM81CVwHZyp96B6U0cv7U7b85NNgB/u5NRIH6Y9qsHkZqN3VeRQmwECH0pFUGnSTAKCKaWwc4qmmgHsOcZpjHGRnigMQ3NM5DHNJRAVzjjNLkbBUGd759KGb5sA8CtOSyAkl6gKeaUtge9VyxEmc5qRD5gyePahxsgDIP0pGkXjBo2cMO3aqbwSM4IPHpVRinuBcdtygZ61FICFPPNAQpgZzimk7/cU7a6BYgjvOVU1cJymRVQQKHJxg1MjN07VU7S2AkHApxbJBpmQRinZBX6VnYBZAe3FJEcdetCkkn2pVOQQBzRsrASKNr+1JIPn45BpgJ3U8nK8VOwEKt+8I71Ix4A71FwHLE9afu3FfQVTjdgPydoB6UMBgAd6YW3Nx2pyDLcmlawA44FRqdj555qb1JxUWAc1UQ6gW2nNNK7ucUozzkU77oIxQAfxj2qOYNgkU9hldxPNRyv8oweD1poGYxBErE8VZWHcyMTU0sceMheTTM5C44xXX7S60I20JZAAcDoKjikDMQo4pHfn602Moj4qIrQosYAkVqlLZ5xzUKfM+O1TZUqcEcVDAGAdMGoyAEAx0qRG68cU1gGyRUJtAN5yBTOjMc05SVYE9KJRjJPetGBUYBgT+ZpsY25GPpSyvgYUUitlcE8Vur20JLULfLgnmraEAAd6pxLt5B47VYRgxBrnqLqUOduCOlRqOck1K/zA461BGMZFTG1hsdk5PpVRkJyas5zkYqBslDitIpoljIQQpUcNUq/dwetQw5EhJNS8N8/6VU/IS2IWB5GacVO0badIOjg8U4YK8U7uwEIIY4I4NVpBzgHI6VI25S2enrTJOgI6dzW8I21Qm9AxsHIqZWGMd6iAZ0wfypEfBYd6HG4XLCj1P51MrMSMCq6AllqwDggGsakSkVrncG60yNi3BFT3IU8n1qrGxLn3q6esQJmDHA9KUOyjavSgjByemKkjG4ZFTLQZLlTjPWnD5TxQgAOBUm0belYNoENBHNOK4HPWgYJGaY+7d14qVuAjjOKhfjk0kjsHwegpGG7vWsVYLjHAJz2pgP5U7gDmkZcgYxWiuLqPOTwPypsbZOMUFyiniq6SuXzTUbiL4jBbJFDbQeKQSgruHWhQzA56mstSiQN2PSms5DcZxUb5wAelPQke9Fla4D1JbnFSKcjFRAnGcYp6ZJPtWbAUmoieeRwakkHeo2YbM043QCNnPXimt2J4NPwWUHNNfBAq/IBrsAM9qhIBUgHk1K6/JTDsCVcdAY2JdoOTzT+2e5qOLkn0qTBUkHpTdwWpC7EPkVaU5T3qmwbfg1bUYX3pz+EXUfnPHcUE4I45pqkihQTznmsh9CXOTTSeeKTO4/SlK5NSkA8EYoPApFX0oPP4UuoC5AXNITkfypx6DHUUwgE0IBI+TzQehpcbAMGmtyc9M1T1BkbYUVGfv7uw61Owx1qJiFB3nirjqA9MN9aeOcLUaHABFPVSTk/nSkJojuCOB6VGp3Lz2qeRQetRDGDVRegx6D5TUrfcFRQ4IyDyO1SD5xUS31DoAycAU4cHFIBg8UoG7I/KpYCsAQSKaoAUc044PFJjk0gIyuTk00/MpqQ/MMVGoxkE1ovMBAcg8YoQY59qQ4xinA8DApgLjcMGmBwDjvT3OB0qIAFwRQgJM4HPOaTtnNPPPtUbc4pIBpU9RTGHNSnAbFJgd6dwI8Yx60AHfj0p+O4pMbaq4DxkmjODSKe/elzkc1DAaQSSKQjINKDgZ5+lJ/CeaoHYaAQc5oHHvTuAcnpTCRuIqlqGhN1p6dKiU9sU4ZK+1RYBwy2eKqlcS1a43CoWH7zNOAbjgcAgUqjI9qToM0q/dOKQCFd2cVXIYPk1ZYkIR61WdiOB1q4BYJSAOaYB68Ujxk9RzSj5c4NaJIT3DPUUqOS3FNwTnNEaHJ9qNLCLG4+uc09ScVGM55qYDismUKM4zTPLDNk1IDwOKCcD1qbiGBNpzUWOWOKmYnb/AEpuAAcd6qIyAsF9KeG4yQQKinPTAqRWBWra0uA9SSB6DrSyAn61GjYbPanMc59ai2oMXJ65pQc9qRRuGaOlAC9GyDTs9DTS3rS546UgFI4qIe/Spccfh0prjkUIQ0fMcUCombbKAOtSkgISOpqmg2HDimuc8Ecik7ik/jJ7GhIY3GV5p4wF9qbjK0dUpsXQXdgDHfpSScjPak4YAk084KDAoBDI12inHLD6dKcqgYJ/KlOADj8qV9QSsREdD2FNbk5FP6LimkgYNUgFFS9VAx0qJME8c09epFJoY9z8tRAjOKezYTIqNU+bdSSEJkMSooXKNyeKlCADg1G3Dj0p3voFrEuecZP+NPXv/WmAY5NPXr1qGMNvTv61HnGakbGCKgk+YYximtQIW/1mc/jU4POM1CqfMGI6VN0JNWxdRxIxgdPSmFgQRS8jnNJgdcYqdh9AQA0o96b/ABADj1px4wD2oF0ILhPwFQhg2atygPHj1qoijn0HWtYPTUT0FPJz6UnQcc0ODnilz6YB9K0DqNB4x36U7HY0gAyeevSnE8DPpSYCHp8vWonO0ADv+lOY5Umqsj9s4qoxDS5IrDv60h+8W7H0qMA7SR1qWJSevpxVvR3B2uiKVDznsMimxAgnAx6GrBjDnNRqgLYGeaOa6Jdnv1Ho373AH41ejGBjPNVoE+fJ6fzq0Tg7axqPWxSsKRk++KQcLz1o6DNBwG6VmMcOF5qFwe3PoKkB3dO3emk8cD601uG5EMZOe1Tx8jFVgBtzu/8Ar1LGcc96qSEnqSZH40hOQDTSx7AUo4GSOtTYLjgB37UDIBxSfxEUe9IBw4xnkUdBR1XOaDgrSGhdw25FO9hUa8CnqOcE0MXUdx07Uo6/0pBRnJ4NSV5juO1DZwDQe4pOwoFsJjDYxxio2X5enNSkjFQTPtGM/WqWouhXlAAye1RryvH/AOqo5bjBI9KdAQSQOldFmlqFydc5/rT9o6/0p4A49TSHGPasrhrYULjvx6UrEA+9A6Z/SkcA4J7UkOxHcEmMkcc1nAgnB4xWhOwMe3+VZjkBsA4Nb0loDQ5yQMbjmmAg4wMZpSGJ46mgRNnODittCUl0GqAMcY55pygnoOO5qURgJ1+tSRqNvp75qXIBUUg8dO1PJ3cY+tOQEc9qRlyCDnmsr3Yys+Cxz07mq8jZQkY6dqncFTjHH1qGXGw+4PFbwIW564mRk46U4ksuR+VNDdPTvSlgOAOa+clqjQGLAY6U1SoB459aSQkkZ4oDKQVAqXcLCvyRgcCmjoTjn1pXY7AMYxTRnHFIY7JGKABzSkDaOeaQsAwAqeoDs8DaMYoYMAOlKSc9BSMWOM003cAKEDqM0EfISTSbj3oJ+XHHNDcWIRsEA1Dnjk09jg4FRNypFTbqUB5AxTWPymiM5QnHtSn7mTwDVPQCtJLtOCKsJ9wetVJQdwNXo1B2/StHFcoXHD7mT1psZPmcUpGB/OkjypzWcbLVgTp8z81IGJbB6Cok+6xJ5FSZGCSOtVqkSDnc2RzikkJJUmnAMiZ7GmsrZXPQ1Owwdh5oIphy8n1pSuJDSDh/xob7gG3c5BpWGHxQcLJzSHg1LsgHJjfnpSgZfJJxTSCGGT1p7qBtwea0S0uAj4UqB1pGPzfL360Mdrcn2pDhWBzUNNgK+CB6jtTfMK8fnmjbtbdnqaGHz03eLAaHOSBSpk5NO4U88ihAVyTkZpq26ECNlySOlRs4Mg3VKBjJI4NRbAz0k9LMYMCo68VGzcgrUwXkg9PeoiApPFGiegDSdp981ZibnJ9KrgZPNSnjBFaN+6AFyVJPTNKfmApjDBx2NSEggAD61m1cEQMrF854xUoUAgdqHI3gD0pz4ABPXFVJu4hspy49qXduU8cEUDGDnv0pVClTzQmMGP7oDFVJoy2CDg5q2ATgEimlMgnPIpJ2dxEMLGMc81YUncCfumoDGTz2NWCDsXniqequxinAziiQgpx94U3OG5FPZSfmHSsr72AgJxjI5NSKdrH6UyTB2k9Aach3nI/StF3AcDuzu6etUpDtmP8Adq2zsqkYHWqUwyeKqO4ETDMq46GrQGOlQqCnzdqnjUHJz74rSbTVkBGygdetSDG30prEb6VDlsGsbXAlIG4YphIJ5p7H5/QCmsv3T2NTbUfQM8GlIwnP5032pzEkAHoKLWAb0Ge5qNQeTTmcBQB1z1pvXpTWghqPtcY/GgjMueKQ4yCKkXBwT3FUn1GP3bkx3qEqckkdKdnaS2KGIYZHSkr3uApOV6c0pAAHc0zOQOKXJJFGlgH7sDFAPrTD0zQGweRStcByn5sUpI3Uwkg5FJvJ6mjluBITluKYw3NmkDDOcUbsE0bAO25pEwM5pjvtXNQyTjGR3qop3Au8HPpTF4kqqsxwMH61IJec1Tg0BYLc8ZpHCkZ71Hv/ADNIzfnU8rWgEu4KuaHfCDmqxZsnPSlbsT+lPkAc023IBqNnIxmlb5lzSKrNgZqlGwDWYY2k1PG3QdsUxoxu5qTb27YptoAAJbJowdxzTzjaB+tNzjnNZi3Io4wu7NNWPDmrAIxmmHhjjpT5mOxFJGEOR36U6NCAM0SdqcDgAVXO7WAcVGRxSFRuxjinFgSAKR2IbFRd7MRE5AkIqBU25B6Zq1Jtbgjn1qMxc5zVqXQZABinpgHNSFAE96YU/OmncAABye1AUKCacq7uKjmyF45x2oWrsg2FHAJ/lSA7Du6GoUkLsBTuTLjqKu1g1JxnO49aN5BPFIQR8tSbCAOeahtBcjYcZpAoUD9ac64bOeKACCfSi9tgBEw+cU4ffoJApvIcE9KFqwHOeDUcRJGSO9SZJBPamgEOQB1o6BYftyMjoaj52VJnkCmngkHpS1AYQd209KhdMEKDwKsRqDnPWq85AB9auN07AyBxiXaelL3AIxTAys2T1pxcBxW3yEV5MlwPSmJ/rucVZZf3nPOah4M2CfpW0JaWJsydSE6fNUsSnkk8GoBADLknHoKsAbW2g1m7PYoUuNpHQ0kJIHXmm7eopsQIk5/Ks0tAJJUBUe1NkXcvTipTzkA8VGx3HZnpSQdShKTHgY60xgdv9KllQuSDxVMyMxIPUV2Q95EmnEf3QHHNWI9wPSsyGQlfmPSr8Uny7hzWFSDRaLH3c+9QJwW7Gp2PyA5qo7ESDB4NZQE9CZehyKrzgJkk4FWCTjk1Uu8lOtXH4hNEUZOSRU7jbARmqqjcOOPQ1cIDRDIrap3EtiqWYLtJ4PamJKyjGamuE5GOh4NQn5QAo/GtItOOwtRJz93J700ZKcGnyplVJ7VEWBxk49quL00QmKDlTtPTvQgBXjrmgkcYGD3ApYmC5Pp2qnzWAsw8EetWpNuOOapRt0YnirEbiQZ/SuWqne5StYhuBuAOeaqFmEmK0WX5Dj8qrtCS2TTpyXUfmNHK7Se1WUXanB5qqiZfGauYCr9Kmo1cYsLZJ9cVMDhTVaNTu3etWAwC9aykHQQn5eKXPIGKDwuR0pu4bc9/WpDQZIFY44ppGBSuuACOtIeh9q0W1kBTd9jAevSpgdsQ6c0jRbgCOlBHBGK1umhakJYFsfnSxYDnIpzIFG49aaGGfTNUndB1J1YFSAMVPGCFyagUqAc1ICSgI61lLyGhkn3ye1OjcEdKanMuDUqBeQBSlpoIVTxzTlbBOKjJJJA496kI4qLFA3I5ppXCYP4UZPehz8o5pCsMGd2DTmQYzjimgZxntTt4IxVh0GAZBFQMp6mpGfYcdaaPm4OatXAZFwcCpNucnmkQAMwFBdhkGm9XoBBIp356ip48lcelMGGJB71KiHbz1ob0swADNS7PlwODUacZYilDnoKh3uA5QM4zQ3HehR/31TZACM5pLcB6ml9aiVugzjFSHk0NajJB04phGTShs8DrSdwB1qRCEHHuOlGc0EdhxTWJ7de9VqwFds9O1QTLvIGeKlIO081GPpVx0BiKdp21IHP4U3HIJ603dh/ejcCRwdoqNgMe1PL8DPWo2Y7cUK4D4sDgdD1qwmMY7VShLKeenpVqM9x0omAuNpJ7ULkDHfrSt0wKjBJPtWa1DqSgYwRTc4BpV6U1hhTSAYGO45FNwRnOaZuIkwenpUw+ZcDtWmwEbYDY6cU9AvXvTSuWFPGM4Pei+gDGyVPrTFqU8sR6VFzmqWuwEpXI6UzZ71IBxnvSY4wajYLDW9x+FRuCTx2qQ8jn8aax+XA4xVJ6hcb3wKUqOmc00AbjxT88jmmFheq9eRQx4pF9PWggk4PSpATv7UmcDrThjNMLckHvVbgB5HSkA+YZ5xSMeOtJjofWqWwEg5NPUdzTE4NKrZzxUsCUHvioyuQSaeSM03GaleQABgHNNTIzSngYpVUA4p9AAgH6VXdSrEgfjVjoM5xTSMiqjKwtyA7mB5/SoyG5qyqfLxxTHUKvNWpJB5kaLnr075qzHGAPakjClQelPJ5+lRKTegWGPhff3pqknkinyEbRxTAcrz1oWwyVTzSFuM0iE9MUyQ7RStqA+RjjjrUbPgdaTfu4z0quwwSzHIq4oLjid7EHoae8Z28cVWGSpbHHapzMFTB5PetHF9BbsbEWBJJ/CrAAySeaZEN3Iqww3Lnpis5vUYwcN6UkvCcUm8rkn+dNd8jilqA12yBTg3vTMbh15pcjgimLqTg5P1psjEDjqKYD1IoLnG7tSS1GNwWJOOlKvTtTEfBOelKjc+1U0xCgEN3IqTGFBxzSZ4Bp/G2pbDqIy5FQt1IPSpck5BpGQt3oTsMiKnCipl6YNJtIHNOCjr6UN3FZXGFieMcmn4+XBoOB9aAMr+NK4xjLlc9qTYPxp/TrTZD0NNMBg4HvUqtxk9agTPU9KlU/LzTkA1n49KVGyT0wKZIPQfjRjaMev601sLYnQgj2xUBwZTx0pyPg4obOSQKVrMbJ15p64GeKgi5zmpf/ANVQ1qFgxnmoccn1NT4HSouMkHrQhDVXANNIxzUpA200AY9qq4aEbtxTh93P6VG64570KTtHHFVbQNxy8tSk849aTtkdKUdaLAxh+5iq4OGx+VW2AJziqzptc8cVUWHQRsbgTUYYZI9O1PkPzDnBqMH0Gea0Wwh5QuCQcGnqjEk4p8Ayv0qdcY4qHK2gWKRXAIqncKFOccfyrRucDHp3qtKgkOSOK1hLqxPsU4hu5xz2q6voOtNRdo6VPEmME05y6ghu0kgUm0gknmpWBHSo2OcelZp3HoPixtJyc04NgA9qiU4GCeKNxA6/hSauwuSg4HtRuyoJGKi80DC04twAe9HKO49Mksc4zTM9gefSkBKtg80woQ/FNIQ48dalUDHB6UhGV5OTSryM0myhGU5yKUEkDHOPSkPQ8fjQmRyaXQV9R5OSaTnaPX3oOBnvQuNvX8fWkJjj0z60dqDyBmgHjjtSGO6LnFKMgH0pp6D1pQev8qQCgknFAHzE4ozQPpxQA/IPU0nb0pAQCM0vU5pAJztI7VVnOVxjmrOep61BIvGRVx3B2sZkvyydPqalthlueCOgpZVAUmmwMHPB+tdLd4kmiD8vIoA5/pUatxg1J05zmudjWqF9TTJDkdeaUHOcGoZcjp0NOKuxleZyTz0HSqjDcwycjsTVgKxXr07VAcAnvkc11RSRDLEP3cUp4HQH0FJEoAwO3an/ACsvFQ3qOwik4564pQfTvTV+YHHP170Lk9en0pD0tqSq3JwfpTy2M1GpGTjqKc3GTnt3qWtQ9CGUgt6+1VpeEcY6irDv6/8A66ryL+7Y55x0raAbs9b3E4wOgoHIJNCk4wAeaQggbSfwr5zVMZBMztnb6Ult8o+frmpSmAeefSmY2nPr2ohKwyYjcc+lJgE/hTSTjvSKcDNQA/HGTTCcke1SqRjOOaa0e0g9zTjZMByhjhqVj8xHelfAXbnmmMOCO9DstBbjWBJP60Agoe560zc26ngYGO9SrDI87R15phxg+9ScbsEVGfv+1FgET7pHeg88UxSd+af0wDQMYFycHpUyrh+DTMdAetPjG4076gOXk4IHNAxv5HSmEkPgdqd1YGhASKwLsvY08k/dx+NVtxFx0/KpixPOe1U1oIe4Kgc8dhSuxbao/wD1UxcscenSgffHrSt26iB87hjrTZBginceZj0NJIAZRStoA4kEkkZOKQtnGR0pTgtj3okUYyKbTY7jMluT1HSn9WBFMIO3ipCvloDStoBGRvYmgAu9PGFHPQ0zplhUqSAXqDz0pT0DegpBwM9s0ZITb60730ACN/IpWOBs/Kk+6QKQcsT6VW2gD2bKbTximEgMFFOb5uQKheNgwIHFTuBLjgnNRHAHvUqKShye2aicYUU9mAitkHjmpTzHx1FQIc8kflU0ZGSD3q99AEOSBQMlOlPH3WB/ChP9UfWlrsIYfm/CgMXHTpTGJUU+FgR70o6jFAypGelOH3B9aZncxxxTkcBSO9SDQ5mycjsKQ87feoi2EI75p6HjJzmqkgQ8HBx2zSHO/A5pSQBn1psbASEN+FK7uA5zngdaUNn5c5FMPDe5pcbVz3oXcBk6jbtPSi2+VGHp3qG5LNH8vWpIiyxKG6kVcLqLYCliQxrPeR1l+taDAA7fUc1TeEtJmqpWi9RDsgrgn608NtAwapSbgce9TJKGQAVpKKtdDLDDA/lUiAbMn72ajU7toz3qZk2d81hayAB0pB8/ynsacSAopp7t3qAGt6DrUMjMGx2qcj8M1EPmJFVF2CxXEu6Tb37VMpIqIqqOTUsbBhzWj1WgAy8k04fdHrSgZB56U1QSO/AqNbAPyCCBTcDGKcANoPem45680kMM4OMUox6Un8dKy+lIAxyBignmg9BzTd3zc8UWAdwODTWKYFJ/HSMoBqgI2k2gYFRNcbc7qkZdxx2qFogzba1ikwImuGlUgA4p6KzcEVLHAqjGPrUm3DgCrbitEBCBsBpVyAKnaLcCO4qOG3bByalSVtQJOoBpWBJ4pSpIAxUkYIyO1Z3AiaJivqcUCL5cd6nUbiaVACxzU8zAgZeMDrQqlPmqYgAN60yQ4iIxyacWAzBZsnpT1x16U3GEHNGQAMUrgPIyB6CozgkD1qQL6ioyMSe1CAdjDDmkcgnAoUZPWkxmQkUaAM+7KCeafjLZpjr8me/aljJ2fN1qnqriHbucihvulscikPy5qPzMKQOaSAkY7gGI6ijOAKRiCgFJkYGadhj+ueO3SonJwPanAncc9DQxHIFPqDEVvmxQR1zQq5z607bux7UmwM2QMkrMBxUiShRz1qxKqk9OlU5AuCMYNbwfMrMRZWUM4J6VYznPPFZ8GRkv0q+MDGOhqakUthjMfMM9KceCTSMe1GcknH41HQBAQW4oJJJz0pUUde9IOe/SjcBF4SlVsEZobG3ikbBAoTuAoyZcnpSNkhh1NObqD3pG9utO4CgbVzVeZd3J71PJlgBwKaV5Ck1SVmG5n7NmCPWlYgSA0+VCsuAeKgmVhJuyMV0R13ESkn/CoAuHDHmplyYwB1qEZV8N1FOOjaFuWSm7BFSuNqhsVFHyvWp2GUC5rOV1ox2IAyjvQkhLUkseyQelN2sHwvNNW6hct9+KrnDSEZxUqnjjmopCVAcdT1qVo7ARyYDmqZRC2QKvGP5t/b0qqyLJkDgVtB2Je5GFKkYPWtGAqkYB61Q2ZcDPSrkSjjnmipdoaZO4JUbaiK/NhjU5bAAPeoSvzc1zobHA5H9arXWdvTpVkkDgdaguOY+RzVwVncUtingDBB5zWgjAwis8hgRgZq4jDywAK3qttaiViO4P7sfWq6jA45NXJMFDUESgsSadOaSB3GSLmLdVTouDWkF3IwI4FZzp85HY1dOV20yWCNllHc1OIvU8GoQm1lyeatbgxC9quTaWgCgFSABxU8CFTyKQL86gVZjxuINc05SKsV7lvJVmFUYrkljuHHetO6XKkEcVkoVHTqTTo8rTuge5aA+beOlSbyQQaIUbYT2NSOg21Mmk7FDEkAj5NRxTMW9vc1WYHeR2q6qgpwOlU0kvUVxwZyD/ACoztGKcnyqc03G881Ddx3YmWxlunrUciMRgVYkQlPYUz7q5NJPqgK8knljGOfWqr3DNwvWnXIbJOOKW3gUAHqe9dEEkuYi7eg8K7ICfrSFSSoq2VG0AUwqFGT3rNT1KsIFwvPSpFwVwD0quXJYr1FSpkHJ6VMkMdysgJHBqUH5uKiZ1LDPanI3GRUPYFoSYAppfaOBQpOMnn0pq4PANA0SLwvtSSH5QAKUfKPWhjlTSERpnoaR1IGRTec1ISDjmrsMruxLUqthgDTnjJPWlRRtP86pNWEMziTpTnG7tzUZc7+lPLHrTtYY3AC809GyuD0zUZPBz1qSPnHFFnYWwNkdOtA4oc4Of5U0N8vA5pJaAShhg0YGB3qJWwCCOakX7pyalqwBhcjFDHninKOBSsFAJpX1GIhI+lPBHWo1kU/LUgAIpO4hhBHvTDkCpCN34VHgnjPFNMB/Reec1E3Bz2qQ9MHpUbHjA6VUQELYGQKROWy3WmB93HfvUqYAqnoANgk8cdqrbhv8AarQGQeKryoB25og+gDk5Ge3rViMkg88VVUkEDtVlOBjFE1YCRj8uM9aYuRnPelPDdaTPb1rMB4zk5PNMJJYkdDT+q46VH9wUIBpUb805QQfagDgE04jI4NO4dBjHPanDJ574oJ2jgUinPIoAU45qMKeT3qToxNIeDTQC4xjNByRQO1KeOKQDCPfmmdcipWHT1qM8HjvTQDSQATTVbH0obrtPFOMYC1QxF5bn8KeDzg9KaBgYpy9R60mHUYDycVER8+fSpTnn0pmz5c/rVJkiHBpMDHWlYYGAKReEJNUMkXoTmlQ81GD2pVYCpaDoTLgGg8AnNIM4z60hOO3FSkA7BI96bkbqdztAzTSBxRuApIzgUjEKoGaHz1FRkZOM00hDy3yjFIcMOR+NAOV44pc8UDBTxgU4DIxmo48ljnrUo6cUmAjjg1GoLHmps8ZxQBgfzoTAjxznFJIvcU8DGaVhkE0X1Aq4IYAGgY2+tEq4bINRhwM+lbLVAMYnB9e1MaMsM09pRnAHSmlySMnBq1clli0bC7TU7HEftVFJNnfr2q0jFoxzWdSOtxrsNkYL160wE4OKGX5gSeRTAw8zaOlUlpoF9SePO0560hBA460qKTx2p6oTweRUN2GMUYj5NGAIscYpxjypGc0iIcc84ouLoRlMgD9acE2r71KwCgHFRseevAo5rjA5ODjilpwGfoKhmfYeDzQtdAJd3Oc0A4NQpMAc04uC3Io5bBuT53cU4Dn8KiDhjgcU8EYqWguIwy1A6gYFKVBBqNjgHHWhaghGI3EUzeMnNIqs5JNNSPdkGrSQhTLg47VIvzDpURjx1qYABR6ih2sAMR36imuRnmnMvGaRk3YPWkh9BIyH6VIScdsUyIAEgde9PfjAoe4IcCMA1LjsKhUgjHWpe2f0qGIRjxUK/ezjpUkmTUSgqORTQeo8cinUwN2HX0pQSeooYyGbmmDK4HrVhlzn1qNR8xBq09BCHrxTxgdajXAJzTicmgNh44U81C6bm6VIx6Y5xTCTmhDKsiEGmbW4B/MVZkTcTSCPIA71qpaE6IfBzGR/k0sWfMOaliQKDj1p4XDZrJy1YyOVMg1URWJIPbvWifmHFQCIq2KIysrBuNEQH4dabKcH0+lWtuF+lVJ1xzRF3YNaDCckHPTtTCcLk9ab5o34zilwWXHT3rW1haPYaD8g+Y0xmbcBjOaaQQ/HHrUqEMDjr2q9gdhoTeQc5qwEw3tSqBihSck5zmobbGJlSTxzQMlS39aZnk+pqTOABSasFhF6dKcPl5xSAd+9BOB1pBcbn5SB3pc4ApSOOgBppwOMc0wHnB/xpVPGKb/B+FC5LYPA7UrBdEmcdfxpAeMdKbgZ6frS5waVhdR2SM0EndwcCkP3sE5pVx60h6juc+lKOnqKQNjoeKTjANKwD+Bg0uTSCkJCtxxSAF6nGKYwxmjPze3rUUjlVJz+NUoiuQyjj6elQImCcHH9akaZSOTyKhL5cntXTG6Jb2J0kweetOkm2qCapsWJ9aGLbSSSc+o6UciDmLazgj+tMeYOOOcdagC9FNKq989aaggbHKMRD19agkPHFWOB3/GonKFsHOTVp6iVx8ROPmxT9pBAA4qMERrjGfX3phn6kfiam13oNNWLCrznigrjnmq6y7jtwQfSrIAGTgce9S01uVceEAAxihkypHpSJuU8dPSjqTzU2YJXKxXnJ/8A11FJ/qmx2FPkypznrxUTt+6kHoK2RKs9j15Mhsk8Ujnc+T07Ubj04pSCB75r5uTZYxl2t1zRsy3NOxtI700sQxqGrO4wIDZqPAyMVIGDE8d6DweKQAnUDtUjZ3DFRrwvv705fXNaWewmK33+aXAYFs/Skyo4xk+1MywyD096ba6gA2nPHNN4xk9e1ICN2R19KVwAaz66jI8ZJbmmEjHHWnt1I7VFnPSmhiA4YZp5+9imdX9qcSN2aYCsfmGPwp68DPrTAMtTo+OD6Ul5ALgnc1OQYGe1Io5Ipx5UA8UXtYQFCQHFSLGdhNC/dK+9OJwu3PPWra0EMAPUZ+tCDDHOaeG2qB60gHyse9JrawxvABPrSZJwT1FABORmkzhgewqdOgDs4wSORTiSydOaRm556UEk4obYCsuEH0obOBnp6Up+91GMUKS5BA6UaICN8nrTWztAqR8kkelRFht/lUIBd3Bx2NL95lpiE8gU5exB5zVLYBwXMgyaMAEjP0oORIKRhzkU1sA8DKEHg5ozwR2pOqD19qT8KcrXATOAMH61HI3Y/pTg3znPSoZjlqcbX1Acny09QfMpgJJ+lSJzk03fmBinhiPWlB29aYMlt3enbs8YqXo9QI5W+fPT2pYhjpTZl6mktpD93FWleN0HUlPGcdc0p+8Kb0cE9utNkchwR0NStWA0nbOBxtP6VN/CQPWo5AF+fjmmEs7fLnHrVPX0ETL8y49KI0O4t6U9EAGM9aTgZVe9RtsMdwZNx6DmmuCx4HFMdiq4zTw+Ex60JaagUWZtzeneratvUA/hTCuH56HrT0wrHAq21ZWQLYUsEYZpMDBPagkMQW6elIwGfYVF9dQM27Rg+B3pIF2DHUmpr6VVXJFYzXzeYAD9DXbRTlGwtjfhHPuKsL84yT3qjZy70O7rV1CNh9TXPUi4uxQ5tu7jpSqecYpmcDFPHy/eNYiI3BJIAJqJY23se1TBvnJoOd31q4tJAV2gG7JFSrGFU9KeflIFIp+cqelLmbDYbt2qTTgMgsD17UMDnAp2MAUr6AMUZYCkYAP0qQDnIqNgS9K+oAFyR6mkI5x3pSfmHY0pPzcdafQYzp16imspZgfSpSucGo3yORQmAi9TScDOTT8bVHrTGTDdadwIzjOeuKbHgsT3qYx5b2oCBTVXVgGqvzGkQbj9KM4zTk6U7gLu+c8U7dtO2mjgHNJEwb8DSQEy4BPGTQOv1pqsQxPan9Ez61ADyuwjb0NAwHz2pAMpmkUAxsSaS3FoBGG56VBIfmPp2qZQW4NRy8nAHFOLDoNPKAGjbnHtTTndxUy+/Wq2GJ1P9KZcDCAg05eT70ONwwB0oTswI4HBByOaeG2tk1BCcy4Iqx1baelVNANbG4GmMwLcdPSnuAH9qi2jecGkg3FlbcSR2qAkqucVIp4NNchkwOTWiELEcoT6UhORnNM8zYpAOaSJxIuf0puPWwyVWx1NODbhnpUZOFqTblRik7AJG21jnp0qVWyDTR64xT2YYwBUSsCGsmQWqnLHlySKu5O3B60x03A0RlZgUIUYscnjtVwDamPSmiILF+NPB+TGc1pOXMAvVAT3piA7SD1pV6YPSmlvmHtUJX0AXpjmkLYcY79aXI700jBB707AEhIbApyDOAetNBzhieakX17mh6aDGvwRk0Z+bpTZzjpSocx5prYWwrZxmmsOBz+NOP3KY5ATHpQlcHsQSJmTcD1qErvbHpUzKQwNRyttccVvF7IRHGCrjLU1yvmnJ5p4ZdxIqGbBfd19a1tdiJ4XG0461ZAywbP4VQt1LSbu1Xl3E4HBzzUVY2GtgcBs5GTTGVgvvU3B4PXvTHK4x61imFhkbZXGc0rYePHakjQAEdzTwByvaqejuBBI5AAHQ96hxjj1qxJHswewqFyNoxWid0Iruu2TIPNWoWO3LdTVJwWfPer9up8sbhyK1qNKIIs44FRkjzDnnipcjA20woAcmuSO5XQYVBbI6CmS/N7ipwAMgUyRQRg9au+wmUWwTnvUkSNtZh0qILgEnirULhY8VtK6RKK7ZJ5pu4AHHU1LKrbgPWq0g5wvaqhYRMkpMRX1qJ7fK5xwKUDIx0q4MeUFzVSajrEN0UAnQEd6lMe1hjp3xU/l4cDtnmnSlEBYYzU+1begWGp80intVmMjeeOKppIpjLd+tTKcYOTzWc1fcpDpWLMRVJrbEua0mGQPUVSunKZqaT6IPUe0qxgAmkSVZuFqgEZuSc1bslA6citZ0lFXY7hLEVbNWYxhKfIMnkVGp2Eqe9Yp8yswAqCam8sBcjrUa8HHarAwENRJ2BIZtxyehphAbjpxVSW8IYqOgqtJdPu6nFbRpSYi1cqCNoqOFNigY61CLkk5I5qaNy44q3FwVmGjJixUZx3qrI7MCKslgy/NxVUnd8u3AohoIdHHtwafLJtGBQrZwKGVSTnmk99SkRSMAgc9akjk3RZHX0qufmbB7dKPM2NgHIq+VNWYi2JGwPQ0+Hkk4x6VFEdx6fSp1+ViKyfYZJjoDSHnp0pd3ygd6aSQhrJANYcGmKCe9PB3LzQMDtWi0GIVJ5zzQAF70ryYxxxTCfekrtCGTADkUnVRjmmufmxjkUvQVpYBuOevHpUkfU/zpjKVGR3oiJOQetN7BsOf73tTgowCOlDDoTSgZHHQVPkA1E5JPSpGQBeKAe2enpThg4FS5X3CwyMMRzxTn4HFOC+nagjj2qb6gQxrl8mrDYz7VEB6daeAWGab1YDj+tNYYFKM59qRhuyCeKS3AhYHP1qNshD6mrBGeP0qIg7fU+laKQiBU2j3qUEEcGjGcZHFB4wapu4Eq9MYqKWnKec8UjfM5FTazBK4wAFRjrUwBCio1Uj6U5WyABQ2NCqCTnpTgvv+FCnbwetKOPepbuAoxzzTGPBGM048f1phORTSAdjK47U8cZpo+7gmhR8xNSAE7gRTIcDrT24HWoAcSY71UVuBOSMnsKb0zxilY5AprNg4oW4DgOnNL2yajjYk09uT1/GhoNxDlsgmmuvHWnDuKa2FHNAEeMN7U/PBPSmgZbNObGzFUwGZAI/lUnQ4xURAbBzTwO/SjQAI79KaRxTyc4qN32tzSVwGZ4IxTUyVIPJpVPznnilH3ulabaAISQvQYoBAUYNDEKh9aaucUbiLAPTPenY4wegpinIpRg1DGPJGPpUbk7RTzyh5qJiSBk0khMeTkAd6YQMEdu9OxUchHJP5VSGxeMHnnNOHvVUOxNSxsQMHk1Tiw0JQfmqYZIzVfIH1FTKcD61nIEKSBSA578U1uQc0seQKOgXAHBx1pTxx+dIBnOO9Lgde1AEbLnpVOcFc4FXehNRPHuYkmrg7MCl5e9sGiUY+YVbMfORSNCPTitudE2ZRAYNViBjjGc4qQ2+7mnRw7BxSlNNBbUUpx15pgjC/XNTEY70wYZSf0qE3YbJk4470vIHFIgA5I5p7cHArO+o9hrDgetGKcemTQQetIBCMjJqFk3ZINSOcd6YJAO9UrgxDnpVSdTv5PHpVoOOT+tRnDHnr2rSOjFpuVsEc54pwbC/41MiYz6Gl8oA8mrc4iHJ05qQNxTQu360o74rJ6j2HdRjvQATxTc/N6U4HAJpDGjCkjtTVH7zJPBppcEZqRfujinsLoRH7/WpFxnHYUxvvcj60NkE4NO1w2Hg7uvSm889qI+RjPWkPAyaOoCRnDDjFSSNhsURnIpXGcc/hSb1GOQDORUmePaol+9jt3qUkDt+NQ9xBjBpCuRTx39aa5KpnvSuMqf8ALU1KjE9KqecfMOatK2VArWUbCQ8YJNMbqTinA44prEkEd6lDaG4G6gnBx3pAADnv3pwwfrVMBvQ0jDsKeqnv1pONrE8ijqIZuxjBqWMDdioN3GccU+Jjuz2qmtAJ8DPWl6kYpgJz1p/Qdayegxc7R70dc0nBbI7UFucUWC1x38GT3qnOMA+lWc8YHSoZcODjtVQ0YdCgw+Y5FSJytSSR4YEDkCo2JAI6e9b3uiXbcic4yc9acjAR5700EsOuDTkAXtwap7BqTIACKVRweB70nfihiFHr7VmAOFLYz0qRVyD6VHuBUZPPrUyAhcmk9gSRGQB7Cmkgtj17VK3zDpzTCMDP60Jj2Gg8tTsFu/FMHTODShskEfhimxbEmMfSjPTH4UZpQMGpHqIR+VJx3NGaTqP50xeo4YwaCME470nOP5Uu7C88mgegozjB/KnAA9OlIGPQ0qng5qWHkP6D29KTbuBBFHJ5oHXIqQsNxgZIqGYAjGO1WMcVXmGauO4uhmypsfA6UKSuBjgZ/GrLDnP6UqoCuSK6ebQWi2K+Cx44A96lEec5HXtmnhFDZojkG8qRj+tLmfQdu4nknbnrSLGVTbnGKsF1UYPWo/vVCbFbXUqyIR0NRgMD83fvU0569wKjLbkwTz14reO2ovQUoQuQelM2qpx1HvxR8wOOMdzTsjrnOetBPkxFyW5xVjzMpz1qtjk4OKnSMumd1TK3Usk3Nsz3qEbjJjt2pz5hGMk0xZMkEk0ktLoPUHDMOvI61DLkIwJ5APap5MFTzjP61XkOY3weMdauIddT1wZyCB0FHJIbFKpA78UMMrntXzUk3qUIT83HSkYDJ4OaCMDdS4HBznPapTCxHEu07j605xubNSYBPoKa2ACO9NvXUYwdcA/WnBepzx0pAgDfWnDG4BelVuApwGUL1pShdgMDihSEbJpWYsxYVXzERbQCB0wcU2Qdwe9S5GQfzqLAyc9O1TLcEMcbcY71ER1xnpTz8w5PSoyeDg1PUYKQc+1CAUDpTv4ePWgYo6jNPUgZzTMYGe9C9cUJAPQZJ9qUffoUnNGCX9xSegCqTuOacD82fbimJksQKehAJ70+gh+B5ZY9aQMQuMUgO76UHAbOaE7gLwVJ7+lRMMYHc04k4yKQEEikgHNjANKW3LxjiiRM4GKa68AKec1Ss20BNAVxzSqOG29R1qFW2cdxUisQuAetW4pgIT8pJHPSqxYZx2qeVyEIxVRUbPzVnZAtyVW5609ME+lRqOop8fDYo3YxW5PvTnGFH60xT83POKcc7uad0hDsbV570jHAHrRywGe1KOWqdNguxh6YA5qCVCxGBzVgEqT6GkxliTTTtsG5FjC8UqE7gAfajBB5oHB7VdrgSMNrAYpWGDzxShgRu7ik3GQgnp0qH2AY+CvAppUqARUjD5sds02XjjFCbSAMADcelMkHyAinnds7YNNbIjxmqSS3DccVzECRTkHApM5iU96dnCAUp72AUEA8dKY4xg1IMBeetRuy7D61NgIwpbnPFTnGABUUZDDAqTBBz6VpJtKwAUG/nrjg1Hg9RUuSQfX1prcLxUX7AI/yjI/KomycjoKlKgqpzzmldRlaoDMvUZrZsDJrKisWZg2Oa3pGXcU7ZoRFBUAD3ropVZQWwmupDaQNGvP51bQENnPFOxtBFO4CfXmsJ1HLVlIGHoKaQWGTTk6Y/KjqoB61AhjDGAO/elIG09SfalOCfpQMAn1ppsBj42gClC4Xd3p2NoyRSE8daFqAHkbuhHajOWBNKRlh9KGIJGO1DdtAFBBFIuDmnPgMuPxpowH5FTa249xhUAA0rdc460Hqf60hJbBx0pq9hDtuEI70zH96l6ikZtwHrT1GNbtQecE9RSk4ODUTyBW5NCuwJlIIJPUUw89aRCGGRQ5wOOlHUCHPzGnqTtx3qMnc+B0zT14YjvWjWgDsZQj0oUqg4AwKQcDFLtAOccGk23oBKD8pPqKcgLL0pjEbcA09chM1IDhzkA8elJjj0NCqeeaRjuwKlALHkMf0prYIbjHtT3wgGDz61C8g2+mTTsIjyA2Kk9+gphwVB9KcWO3A6Vb0GKDt5HSlz8u7v6U0j5enSlw23ilcCHOxt2ODUynI3Uz+BgacqhgfaqbEIcMSfSod2CQKl6A+1QvndnFKLAEAU8nmk28496QH5/epGOMYFW29h9CuVAYk96SNQrHHQ1K4FMZtg6VSk2Ic2MYpynApg5XmhjyBUtDCWby8ZqZW8yPNVJ0LjaOuas26lYttVb3dA3HBjkH0pVP600A4JHNKp5xWe4A2dpB6mmAFeDU5XOKifk7T+FNANJAGRULDJ3g1LkE9cUwlQNvr0prRgAYZHvTycZqDoOetSA5wD3qmuoC8ZNOTlSRwKQDGaAQFNSAr4YUp+7gdqRe1KCckUPQAP3cVC33T3qYrgH1qHnYVNOIEUhHGKglUvznpVlkUDg1FJycgcVtF2Eyug9etI2Npz2FShQ5yO1MkkHStb3dwuNtXYHn8quJJtG6qcYO4kVbAxHz1qalgQ9vlYtxzVVlZZCR0qwSTjFRTE7sGoT6ALGxPPvUiNjg9ajiXCGpQBgHvRIAmG5SB1qtIMgADmrWRnp1qF1w5NKD6MCr5ZDj2qzE/zFe1VHcvNgdKlgBVznvW8o33CJYdyhAB4pyyAsM/lTH5wRQI+5rFbDLJwVBFRSoSOD1qVR+7wKbkJHkjnsKmLsxMoOCGKHp3polCqMU+dgRwKiiQEZrpirq8iSwThcnp61BkBjirBjLxk9qqhDnOOKULMGOIy3XAFTQqXH0pkcW/JParIj8tfl4zVTlFLzERuCG2iqUwO4hm5rRZcMGNUrqEvKGHftU0pgyOAgEr2q1FnADc1XjjMYy1XI1zj+dVUkNLQsg5+lVbpQ2c9at8Cqs2Cxya5Yb3LZThTeOeAKuRKsXtUMIO7GOKn4xitZtt2ZKHu+48UYDNUcfXBqdQFODWb00K3HMBtzSuf3fAqNnDNikMmRtqbMDJn+WU03bxk/hVy5i3HOOarFflPHPpXfTnFpXM3oQD7+ehqcS4+VRUaRkEkj6VOIwOTyKJNXHEcNxXcTx71G7HHy1IRhfY01mXZnbyKzvcGOgwG+Y1MxAzx0qvCDyTSzuyAdzScbyKWxWklJcjGKfborfMRg96gYYAPXmnoWCZHatnGyIuaKlCOO1SIQxz296rwA7eRUw+U4IrlktS0SdaZISBxUg5Gaay81C3Ag37SOOKkHHuDUUwLFecVIoyAKt2sAj/ADJ6HNRrkDHNSlc03aCue9CtYGMcZI9aQjkD0oYktTc7hk9ateYDmxnOfwoj4JIpqgmlUjoKbDQlXoSaUcH2pByMCgD5SM81mA/B6jpQvGecmjkCjOeaQC7uKVjlcCmds96cR8uaNA9RMc8fjSrxSr396XjrSuAnVaBzk0MQMCmnA5pIOggOCeabgHJpwORgjpTei8VYEcmSM9hQRhcntUhAIGKTacZP0q76agRq4ZsA07aSxHak2hT6ZpryhO/WjfYBxx0qQYVQQOaiyDj3p4wTik02A8etOXkdcCm4AzQvcD8qgBzD0pjdsfnUnUD0pp24HNCATnbz1p2Tg8U1s9qM/J9aYARkHPNV+A2SeRUoyFPFQyDHy+tVEGTlgRmmOCcEdaQDK+lPzkijYBY0xzmn7RkelJn0/GndKh3GNYY4AqNgAcGpTgDnvUbkAZPWmhDCQBgdqCcIc9TR05HU01xuIx1rTS4DQSBn1qQAkihkAAp4GD7UroNQI4461A4JIqc4FMY46UloGoxVGPSm9G+lKWwMjp60hPPPerVwEJyuO+Kh3MDj0qYZyTmoHJDVcRPuWoz8vNPGNvFRRHipO2azluMdyBimSEAcde1OJGMCo3OBzzSW4Dlb5ajf5kpQMjFDLgEZqla4iKMDrjinonOaVVwcHtTlzupydmFhzDFSryoprdKFPbFQ3dAKeR70g4HtTmOB1pu7dxUjDnPFKfSkJ+Yfyo4B7/WmBET+8Gadkc0HrnFR+ZmqtcB+cnHSmsfmxTULE5aiRW4INVbUVxscm5iMVJuJNJHHsU56+tKBkYFDt0AYT83FSKBkdqgYFT6VPCcqCTTktLodh+35vSnEcU0HqfWlBxwayAGOODTWPI/lSnG7k5FIWBbimBG7ZODVcttlz0qeRgr1FIhYZBxWkRDtwPAHFCoBJzTC3lDB605JATwOarVbDWpKDzwKG+Zc0oAINDD5MGsxCA4X60gYZzTcnAOac3H0FOwDuM8UvAU56VCXO7g049MfpRYaI3Xg5qQMSgxTWTt605BlR7VV9BdRMcZJphIwadKdq5FRZyuMULXUNiSIYxmnOece9C8DHU0SH5c0uoWGA4cgdKm759qhjfcOOtTHsKJbgOGAd1SA5brUPfrUkfP51DGSEc0yQkpTs/N14prYxUoDNKFWLYqwjED0p8y4x6VHGQzFe9b3uritYcpAY0uMPzTgvIpJB83sakYwsN+OlOXH403YN/v2pScEEUwHKTtPvSHpQen9aG6D2qRdCBzyQBREeevNPYZJqGP5Wx6Ctd0D0Lqr+Zp3T/CmgkqDinHmsQDgN7mjPze9KRjOabnB+tIYfdUnOM9KiVTyO1SgZU5NNqhbkcvXOfzqu+OasSAbsCoCPmNaRBkO0E0ZPUdB19qkKgio4ycY7dzWgtSRDlvfNMm9egp6YPIGM9akKBgcVN7MaZWjfcM9BVqNwRjt6ZqoYHUnHf8ASpocqmCOackmtCU7aMmBDY5pXpMdhQxAA5yTUFbDdnT1pxXHSgZJANOYUri0sM78dacOCT1o79s+lJ3AH50DBvTFNHI+lKfSkB44PtTDUdnjFJjnHpQOwNOHUEjmkLYTIHFKDkCgj+dJnH9TQDZN/WkI5HGDSAYHIFOwAc+tSAdcDvUUi/LgVL0NQSttJPrTjuMhkx5ZJHQdKrQykEhs1YmfK5FUc5fvk+grogrol76Fl3BHP51GDtZieAOlQliSc9aPMbpz/hWnIK9yVpNzYx1NP37VOCKh3bvlHOeppeQc9vWlZCGmQlyen4daEXOf605VG72oJGDgc9etUhXbDAxt4FJ0yDwKiy2S2R1pGbOCuce9OwyQ/wAQz06cVNbygghj0qsW655pMk4POaTjdDvY0coQapzYLEqeOmKkiRt2Qe/Bpk6bZQT0qYqzKuNWUtgH6UyQ5jbntnFSqiFcg81HKvytngYNWrX0E7nq5fA245qVQCmSfwqtn5yc81MFYrz0FfNvbVFC8bc5pAwbGO1Kpzx2pAQFIHrUO24x+G64oIAbp25pQ3QH0pVIDEYzRyoQibQ2TSfek+UdeaXG5unSnqSGzjkVSSbAiOdx4GaQHkj1qTo5Zh1pjlewNCutgGjhsdqjc/OakdtzDHYVA59qlqwxFOTzTc4JFNZiHGDT+jVTQxuOacOFNJgqee/She4qWA4DjPY0qYD80Z+QgdRRj5evNCdncGSZ79qTzNjH6Uw5JxximuRnGc0/MCaMkHPrxQPvn0qJd7Nx90dak3cHvRKy0ESKQG4pB9456Usf3CajbO1jSigHBsjb2NOOFZQKZGCG55PWpAAXO7im1oAM27BpduI8jrSg5BGOvejaC20mnFXd2JjGXKB/0pykBc85oZPm29qUnaxU/nTuA0tuwdtRHnNSNnqKbgEcdT2qHqhjdmMHPalA6Dv605sZFJ0cCkwA/KegzSsvJ9KCQT7ChuCOuKNWAq8MBQRhcj1pGPIYdKUk9uhpyARjkLSOMHFLj+IdKAfM+b09KGmwI3GACaaB8p5pzcn6UzkntWkWnECROVI9akBGAn5moEOODTkPPNQtNAHnG7p0psrZpx5cYpjDdJ+NNbAIz4wO3WmvzyKc6jO3tjg0mCpA6571VkA9QBEKVcgHIpgODsp4YliKJJN3EKSCnXnNRuyqMU7HOO5qKRMkA+tRa7sxj4wF6U4uSMdqa2FGBQRhR7VTaWgEzP8AKKBjyyTUWd2PamOzD6UnuBMfQ9DQSF+8M+9NL8KPypsuSKNtwM+4YJPkn6VahxvDA1Qv0JkDAc+lTW0jLGAa6eS8LpiRoqdxbpzRnOBjgVGHAwB3qQDA471yvQoVjyCvShjg8UAjkelCkbj7UvIBGXAJ6UxB3qQnIOahY+XjnrTihEmd/wAp6CkO0ggde1KMeXu70R4z0zTs9gF4GDntSABSffpSZ4weuaVidwzwKlW3AGHyjpmlGB8x5xSvxjHNMAIFN9QFBBJJHFIDjI7Um7g/zqPfzg9aSGSE4570wLkHNNDfN7U4tzmnawDSMnkZqOSIN161LgseKVQAcHr61S0B6iKm1MAUMvyZ/Sn9fmzSOegFRd3ArY2uPQ05RhjmklOXGOlOblgMVq72AX+L2oVeDmjkEc0ckHHFJNrYQoGBUq8rycVEgz1NPByMCpegDwSTjPBpOBkUmfmFKOWNShiMPlz3rNmkKy4J4rUPDYqhfwn7wHNa0mk9RN2EWQNFhamQ4i+brVKAmMkGrSSBlK45FaSir6ASI+5OetObKp1/Co1A24z1p+MqBmolYBoPOTUykAt71A4xj+VSj7nvS6AIc9uKhkblQDVhsACq7AClFpPUYMoJ3DjsaDzx6U5SCpNIB3Peqe4iIn5iP0pGCyrihxtbr1pNpjOe5qlsGg6MDG0GlAA59KExjigntSlqxjsc09RxnNRMcAVKmVXB6mpYCnkDHSmg5zTx0IFMUYzjihNAShh5ZB61ESCm40rMAoPbvUE8hB+Xp2ppXBjZXHAHWq7y7Tk1DOWC5rPV2eT5iea7IU00TcvGd5H4HFXEyqAmmwxL5YIp/Tj8qxm1sh7EoPFKVyMUwHpipB1JFZ2KGBsg460obC5zVSSby3OelTRHdGfetJR05hIsA/Luzz6VG5Hfv3pUAXJpjN8uD3rNbgRrgue/tTWAVjk8GnFNoznmhxnBNaiI412ZpkkeVxip2jIGR2NNEmcrjmmpPcLEUQAOD1pZHxwo4xTMfvenWiVmVsY4q3q0HQlydoNRMHM2T0pyMHHWnMvmIDnkUcvLuDHoPWnHAPTinAjYMikYjOBWO4DhjFV5xuGB1qXeMZpHcAb8U46MGil9nMRBJ4PepY2GQKZLO0w2qvHrUltGBgn71dVWSsTEm8sE4xTye2MCnHhc9xSSYZQc9K5U+5Q5WA4qKQ9QaVcYDZpWAYilomDM51yxPaktxuUrmn3ikN8vSktE+bnpXXGUVFsndl5cCHFQFcgYqwqDBB5pFjCLgnmudSS1GxioIzyeKsEAj2qtKwwM1HNdqoCryabUphZE1wy5H1qrNMEPB6dKXeZB92onQMeRWkEloxPcgW5Zid3T0q/bSK35Vk42tt7Vds2O7vzW9SmnDQIu7NI4AJxVVyHBxVh2+XjGKrs+OAOtccEVoPDoi8mq73agnGarXDO0mM0xYQwyprojTjvIPQmacA7kPPpUvnPIwI4NQxRkdfzqxChVsnpQ7dBCqr96mAxg5py98ilKgLWDd2Mbgk+9NdFIyRxTmbZ05zR1JoVwKsi7TgUhycAGp3UFuaiyN3Xj6Vom2LYZsJAGeKAoJwRzUmfnJoJGd2KpNgV2yG2rTjFvUHuBT+PNz7U2QndheBVJiIVjBBA5qdE+Q05VVR8opwz61Lk2NDwuE4pFY7gD1qQEFOKaAAQe9ZXuUPDYOKcxzTMHeSKc3zLUvcRXm+VlxRnHOcGkkVi30pdpJB7VqtgJMDaffpUaAqDjmlDEsR2pT14qVcCLJ3c00AnPFOcgMKjJbdnPFaJsB6mlbGc96RWxkYz70Zxgn8KXUCReRjNLyCAKQNwDT1bNSwHAcc00sAfrTiccDrTNuTz0qV5gLgnBp+MDFNyc/wAqdksOKTARetOJ55pM8e+aCoPQ0AIQGoxkDNOA4FJnGaQDP4uaGwAKc35VEcZ5q0MMkH+lPU5Hv6U3jGfzpwOFPbNDEULmchsZ4qHcWFLcgCUk8CkgySO/vXUklG4r6k8KMSKtKADmmomFye9P5zmsJO47DlYEdKUYxmo0+9jtUg+vFQAo65IpjdRinE9vWmjjrQgHEcU04GAKUdOaaCc49aAE3HNRt87cdcU88daZ0P1q4gKMjIpUUgEk5oRgeacxyuO9LULDwwxxSnge9MTAAqQ1L0AQ80xlBPXpTumTTT1JJoQDWAK4HWmKhHU/SpBwMUjZUfXpVp2AaT84p/U9KaQARTuTz+VD7gDgdqhBz1qQEnIJ570bVANGwajCmR7VA5bOKtH7tViAzVUWHUUE7c4wagcDuO9WBzx6VE5+bAH/ANerW4tyZPu1Jniol4xj86kB4IGc1DBA3OBSPywpCec5NB56/nSQbC/d60N60YJHPNHG3mgY1iCKUH5eKaQAPWl6jjrVW0AlBytKn3aReBTscZzUAJJ93dTU+brxmn4Dj3puAPwoTEGPmpSvGO9IBgmlJFJjEcYH1qvkL071JLIFBNQxt5nIrSK0uwJAOemKlUgjnqKiPyKeeKEcY96Grq4ErjOBUSuA2MUu4lfYVCGUnJJoSAllQuKdHhRtppbgAHmkjjbOWo6WYupI7EfWnIwKZJ6d6YE6mkMZxwamyGOyOT3oVgQTjvULIVY4NOXJ4Haq5VYV0OZA2SetB4Ge9O3ZXkYqCSTafrSSb0HogMW9yT+VOWEKeB+FEcwY4xUhwTiqu9g0GchuvWnnkUgOCfUdKcMYNSwGkYNMblCSaew3DNJsyvJoQiNFAIz1p5GD14prqRyD9Kduyg7VXmMGXI9/Wm5OMU5e/oe1Nbg+1JCGOCyEGmrxHn35NPc7I++O9Ip3Rc1pd2F1FGdwx0pSSSVpOMDilLDJGKkewzOH49anGGOQaqSPtbNPikJwegqnG6uGxZUZ5z0pwPPFM3Y4FSDGc1kxjzwMZpGGD70EgkU0ktUoOgyVSy81XiADehNWyM8GowgTpVp2VgsABC02RivQ805zx1qN+etC3FsIBnJJpxwR9KRjhKVfmXn8qbCwZJPXpSMcD2o4z/nmnYyeTTYDJM7eKiUYIyMn2qeTkcGol5YZPFOOwMsdhk08cHmmgjaPWhjWYDjzSMMsOaaGpcjr60WGIQTzQBigg9c9KRTuB7UC0Ijkufao365IqRuRnHWkKgjPNaA2VpSAQKjyW6c+uKnkUYqIDgY4rVPQLXJIzz6e1TgjGSKgBxwKkDfNtPbpUS1DbckwCM4puADnAockMB2FBBIB7Y5qQGM3ORSA5IpJRk4BpAMsBjjsKtbCZJkg565pScn1NBG1AT0FVzIQ3NJK4y1kYBqME59RSITtINAyBtY0WFcdu3Eg8UDIyBSHrjr7Uo+bBPekGw4Hg5FOHXr+lNzgdKcualhu7CHkHNAzmnAAk+9BWi49gbrTd3APrTivy47etRqCSaaAmByPeo5M9fzpyAjPNI5ytJbiWxTk5RhjiqbnGMD5u9XZAF78elUSwOTnOa6qYrXI3YkZJIoHIzk9e1KOTkD260oX5eR+NauyFcXByMDg5xSk54yOtHJYErjHvUiqGI7cUmxO9xynPPOO1IVXOMexFKpwcZz6AU4j3DH0qNtgXcqGMAn9aADuPPHYVdMYKknr2qoyAEYyapSTGtHdjMADA79qFLAEHp60pUgFvSlHA9qoV2XIWwo/lTZAp+lRRj5+SPyqZkDc5yfWsmrMtXK4yB3/ABpsi/u2/wB0mpipwSQORUMg+RvXBz+VWtxbNHq+3HIHepCfl5/Ko1J3c9O1SnJAPavnGpcrKGqOcAcY60BcPjtSgnB/wppOeBxio7APG0Nk0bvmLY7UgAODTiMgjH5U0rAN3EDgcetSAjYPWmtuwFx+FPRdrgdsc1cWmASMMAHpVcvhulWXcOeRwB2qoxycn1qXL3tBIM4yTx6VF1BbtUjdMk47VX84bghGD2oUb6LcoQjDZp/8Qoaojuzkd6S13GPbOSaFakOTS4GBg02rAkPX7pNO/i+lNHCjIpeq1L7ANkfvTTyPenMu7Hp6U4D5qNEA9RtBANPQYGO5pgO1yPWpIhucZ/GhJsQqfLx2NMPBxk+9SMecDtTRhkLd+1U1oG4DCvnsBQu53Y4pFPQNTs7WI7CmtvIBVOBjPSnuRtBX71RqFLDPrQfvGhvQVh+7CcdaY3zEE0bcjIPFITtYE0nrqwBwScZ4xxSBSCM05vnYYpCSWAPbrSluMMYk56UFQQWH/wCul4DEHmkweBnihpADEKvNJksv0of74B6U0n5uDx2qWgEIPTtS5IGM0Ofm9M0u0bsDp60krgA4+XPBFKjbRtHUmozywzwBTjw4I6VqvIPUSUbB+NMYDAJ/Gi4JbHHQ03JYD0oUVZtBYVAWYnNOU/MQemKRSuSPzpN3PAo5QJPf+VJGOSSetJGCWAJNKwG8DvmlZoBq5BJoRgxzSueOBzSYAQnPvT3AV/vj1p/3TnNRb8kH0p2coSPWhR0BjoxubOenekYfNk0oOFBH40jnnHbFDstgGSNuOBT3YGMAdajIG6nsuCDjtS0sMcgGw+tGAFx1zSLxknvSDdnJ69qGmxBjLgUrHPGKfvCHJFRsec5603dICGeLeQaYUWMgY69KsKMZJ61EyeYwOeBVJu6VwsKOAMjmphlSKizuk4qRid+Bx9aiVr2AGbDdKVWA/rTJG6+vrTYyAvPU0JaaASBscY5NMKg1Ky5ZaQ8vSvbYZHHn7pp6khsimDgnPPpTuSA3cUb63AVBmTNOk/1nP0xTSdmfWlb7u4nk0dBDjxwaZk4wBQOVDUrEDilqnoAhHp1FR7QXJp7ZCgetMU7OD36U0m9bjAgEgilbB+akAwppWHyg9qGAN14oJwM00ttHNNDbiMdaFENxwUhQc5p4GVJJ6Ug68UZKrikwIGG1gTSE/MSadIcuBUb5rWKugJFYE9aXPz8dKZ0xil3ce9JrsA88NgU/O0j9aYoywNGSzc/lU26ASY6vjmnpnaT7U0nkYNKT6cVICjnk9ajlw65NK+OB3qKQnIHbvVxuhFLa3mE9qmtU+Vieo/WmnLN7CplJBKgcVvey1AbHlnI7GpdvOAajwfNA7VMO/NZyew7aETn58+lAYl8kcUrLgGhecA0lYB5OTx0qMqM808n5qGwhFT5gRAY6UrHBwOlO6ZyKTqop3AryfeBHSgsCOegomwDios4wDWiV0BMnf0px4OfypqqM4pxBBIJ4pNahcQgE1KMlRzTE557U5sYOPwqQHj7uT1phPWnD7pz1qI5LA0JagOYjbjH0pgQMmKfgk57UoGAfQ1SaQFGWH5CR1rHETiU+hNdHjg+lMNqjc4raFZw3Ja1K0RaOPml3Z5FTPbgrUBiI6dqnmUtR3Hq3zelTRsFz3FVhyOnNSqWA96JKwyndoWlBHQ1dgAWMAVVmkC8NSxXKbQM1pZyjawkXgoUn0prRhlzULXQHGeTT/M+TJ71lyMfQUxqME1WmmEfY5p0l0uMZ5qhLulOM1pCDe4ttif7bvHyjINRStLEC6jOaVY9iAjkU6RiUVcVtFJfCha9SvHdlTlhzSyXDOTQbfqwqBlfOK05YvUObQnik24z+NWEuVXAqkM96bjI45ocb7grGt9oSReDTDnGc1mhmGMHirMc5CkEVLpK+gXuidiFXk0knzRj61WW4+cCpZSHA25zScOUWg6OQH5cdDVpXx2qiylOR160+FmKnNEknEfUvhwxIz1puByDTYVwpOcipGwR0rl0TKICxBwKk3biKjKnnH5U+IHZnPNaWVriIbhg2M9qjt2XdtxT7iI4yKrxHa3PH1rWEU0TfU1cHJFU7yfy8hvwqykowOOtZt8yMx9aypw96zG2VJbhnbg06I889aSO2ZhkCp4oiSQetdnNFIkuw7ducUsmNpbFCAhcGpGI2461xtrmuiuhlTQ45XrU0MZUZNTykIQMdaSTK89jWvPJqwkrEisHXHeknjAQEdqZ8yAEU9nLrway2ehTM6RSXwetO2+WoA5Jp8p8yUADpShx5mK6L2WokPXdj3p4YnvQSQpYVFCjtJ14PWs0O5eQ5UH1pSOKAMEgfnSsvHNYPcBpAJ9u9MBw3bFOJwKQqFU461SYMjcYqAZBGanPAz3qI4PWtYX3Aco3t1FP8kDvUasEG6mm7UnA7Ue8/hEPaM54HFEeD94Uz7TwDT1lDZx+dJp9Rhxv4psnLcdKeeme9NXgHJ/CjzAchJXnpTiMjioyWXinZOMYpNASKRyacvJOaaBkZx+FOHHSoYEUnDe1Ih5xnr0p833Ce9VSxHWriroCXIBIpArDvUYPze9SDJHFXqg0GTZ4FREnOBU2zGSxph5PFVFg9RVOD60c7vakXPU01nOcUWDoSbiO9ODbcDrTBhlpnRxzRZAXDnFCkA0IdwAJo2j1rHZ2YBnOcDinKdqmmEgE4qMuwPoKfLcNibPrTgfamKcjmgc1NgHZyvp7VHyT1p/rnrUfR8imgHvnbVaclORVkkkVHIm8DNOLsDK8W52wTVlBgdaiVCGOOgqQkg1ctXoBFJAJGye1RRxFGI9KuLzzTJF25xQpvYBf4etN5zTEl3cU4Zzg0rWBjwABS8YFA6deaQZUUgHBcDrnPemYGSacuTkZo6ZHajS4AG56fWl7E96ABkYHJpD97rgUuoCYz3qKQYHH51NkGmtyKpPUBIxhcnvTgBt65oxtFA5J/WluAg+tSZwtGOaXGRj9aTsMbvzgU1hl8ZpAo8zPpQSQafkIUn58Uw/ezntS/7XrSEc4NMGAzkg0/np2pMbR70qk55GaVw6AVw2B1NGOwpASTSjg8UMBCOMVSfIc45zV5hz71Rl5lORmrpg7DkOeCeTTJgd3AxSoT5hyMAU+VSSCOtabMQxcpjNShjge9NxkcikDZbpiluA4joKcvTJ6VGeWqUc496l7DvqKG45pCTjA/KjgNR3PvSQDeQetKuC/tSHg4pwHXtTAecZ4pSM4poJHWlXkntUAOztzimjjrThmmqcjNAAeo54psmduR1p2MkmkJAUk0IDOlZi2D96prcFR9agnYM/A/GiBjxzXU1eOhJdbBySfwpqAMaidzj0z7UxJCDnFQouwblsrlCBUDRHIAHNTRyFu1PxuYcVF2hkKQsDk1OoOAM05v5UicDkVLd0Ow7oMUHk/zoOSR70DjAx1qQALkkkZ9aUAAmjPpQBzS1AaUBOMAD+dRtCpOf0qXIyB3pGI7dKpNoTRX8sB/lp68c0rYAz0FNjYMxAqtWhjgOc0oB5FIeDwKRM80gF6nj8qO2KTGaUZ70BcQ85qF1O0gnpU3IHvTWGOeuapaCIg2MD0pZG5CihgAcdzTH/1mBV2QEhGUANRfc4xx6VK3zYFQsMSgHrinEGPH3uTSydAT8o/nSN8oBof5wKS3DyIWQFR7VLHwgpr8ADFKhziqeqEWBnjHQU8fNj2qPd0x1p4IHYZrFjHjgilYUinGT6UNn8KkYuOfaoj94jtUvbmo3+77U0BG3Ix+lIehpBk/hTugNWAxh8uO1KhwOaG4FInoae6DqO79KQ5yAKUH34pRwtAdBsp2rwetRoctgU6QkICaZHwwNNLQks4xjP4VEz4OByalIJFQujbsipiO4/cCM9zT8YGaaE5AHIFSbflwKTYxvXHUUjtjgU49M01hkgng0IOhHJlRnuaFOVpJAxPtSRggc8/Wr6CuMPWogpyec1Yxntmq7AqxAq46hqhUG48807kDOf0pF4xj8akcZXBND3ARSCD605SdmM1VaTZgDIp0Llh97vTcdLibVhz7vNPp/KpY8deaGT5c0q8Djoalu6Hsxzf6rAFU2TLfKePervVRnp61Ft2kgDOe1EXYHqCptHWhwDTud3SmMuecD86PMQ4AA8elLnGCKQZI9TTsDA45pFJicDmlUjPfNGCOtJuwee9FibEgJ/CnDr61HmpEI796hlCOPk68VBGc9everL/MvAquo+bJHWqjsSWF+Veo5qMkDOeakJxHzVBpv3hznaO9KMbg3YZcEk4qiVAGD+HvV2X52PpUXks5JPSuqDsgK4BLYUfhT1z0B5pXQxnPpUYfJ9x3ArRakkvIBK+tICB1GT70ihguQaRw3fqT2pEji5BxTt+TkkYFRkNycUjbwvU/lRZMN2Ww+7gVFJhjheDTYmIO48Z7U/kg4FTazLSbKeScjt6U8EEAZH0pH+8MDk9aaARxjJ7mtd0LoWY8glhyf51YA7E5z2x0qCIFW68d6mMnBbHespMFtcSRMA+p7VXm/wBUwPpU7SqvBOQD2qtcODkqOo6fhThe+pex60FwSKkAJbGe1KqlixzjjvTVGMknnHFfN67AIcqNvrTc7TjHNOPAz1qF2Z29s0oq4E6ZI2gd809VI3HOKYARjFSHIQqePWq5knqARxMx+Y4NKUwxUE5pPN2gAA59adhh+8PeqtFgM3bCVxk1ARuyMc9anJBdj3FQHJP1qXbZgiNjuA9BUXlg/OeoqVlP3aG9M8UlK2qGREbuvFNYYYAU89eelNXmkhiOMECkJ5AFPHOeaaRuHFUIkYYVfSmO23ApxGFAqNhlgaWjYE+35Ac01nCrmoyzH5V6UGEk/OeKaSb1AVHMj+1WEJQ5701VAAwMUpbkU7x2QCscEHPPenFQB7daaBl884pzfpSl6AIF4z3oK469TS46e1MkcEg+1CSaAmUDbg9ccUh+UBeDkUxSeuc+9OOCwNC10ACvzAD06U4oCcDoBTSSX3KKFyVJBqrboQHouOopABuGe5pD8pzQjfNg1MXcBQMsRmjq2CeRSDhiaUfMxIPapQDR9800D5ie9OXljkcCk9fTFJdBiA72ye1PRvmz61CMk5HapG4IA64pruAhAwfWnY2qCOaMZzkdqCNyfT0px6iGyAMnFQA/NjpVnkqQOwqDADnNNPsMB8pxUkiYVSKjHzN9KkB7DpRJgRlyHwOc04Kd27vTU6Nmp0YbMnriq0tcREetMY7WIz7YqTBC5quQfMLE8URtcB4yODUg4TBqHfuwcd+KerEqSeo6U23tcYrMY+OoNOPKD3qNSXbk81M42jGah2SAjbgjNSMSfpioZpPlH5A09X2x9eTT5dGA4ZLDrQ52r3yKajZJB603Jyc+lCi3YY12LY9O9K7ZwBUcY8wEelSKuOaqyQiTkgZqIZw1Sg/LjFGw5DDpips7jEjG1QxpxyX9c1CXLMAOgqRXwQT0puKYhJsg4HemW67SC3apiAW3dutMPBPpSTS0AlJ/ecU0EK5B5pRk4I64pMbmxio3GA5JzSYwQTSkDPFIeSKaEIeWyelIW3KaUncCO1AXjih2GAJ2jFOwMZzzTUHBFOXkmhsQMQwGOtRHcWFPYc8VE7YlAH40K7GSlCFA9aCeAO4pRkHJoHzuSTT0tcRDKMnmmKCvSp2Uk+wqCRueKafQZPHxzR1JPaoycRUsZIQ560ct9wGtz81RA8ZNSv8AKuKikxtqogKDuIOeKcuN+e1RJwpGOtPU4IFNx7ASkYbOeMUqLwTnmmHgAdakzhQeMVMrAOUgLyOe1LggcflTUBbntT1OSajRBqI/UcfNVeU8HPWpmPzZqCQ5c89OgqoaiKwynHrU3RSahDFmzjpUi5IJ7dq6HqgWxIDuGakGSM/nUMLlgQak3ZQjNZco7hJyfpSD60YJH9KEIJ560mktgHSqQAVqGV9qjmpiS3FRXCApwM+tCeohRICBj0ozjOTzUSbNgIPbpSAsG5PGK0lFX0HYZO38XcUyJtzkntTJ9yE45zTYnOeetXFe6JF4ctkdqGJ2mmK4GT604fMorNqwCpn8Kf0z3oAwMd6MhcVmMcMY560hGKX37ikOAM9qNAGsSDtFKeQMdaVBznvRjJppgNwduKdhhwfwpCc8U7O7A9Kd7hsJIO2elRFe+akcnd15prYA4pICJkAI4+tKVAPWnkAryaa33Kq9wKN9C0gyozWQrNGSO4ropJAsfSsd086beF5rroT0s9iWiJJjnLdqttMzW+AeaWKzU4YirIiCqeKqpUjsPoZBLq+WBNXLdi+dwqR4cH5ulD7UTAGPSn7RW0CxHMw6CoVkYLg81GxYvUqRkjnrWqUVGxPUa8rdBnNQtO2cEc0+bK/d4+tIYdy7ttEeULsjUs45NKA2OKesJCk09Ii3frTcogQfMeBUqRsRg96sra8CrEaBOO4rOVWP2QsU4oHDEY/OrccWFJYVLkbqbI23FYubb0HsRbD3pwTa317Up29c80bgCAetO7aAlUbRTw3y0xSApz1pQc9KyaKJFAOTTCxZiqinkfJ05NR52HnqaStcRDKpUEk5FVJHzINvT1q7NINpFZrfMOOtdFJX1ZMjRiOSAe9RTwAOfT0p1tvCgmp3GRn86iT5ZaD6EUYAUjp70gjZXz2p21QoBNS4AU1LlZ3Q7DdwY4FBwiEGgKUBOaqzXBzg0RV3ZBcjU+bKcnpU7rjAPNU4ZB5hUDr0q8FycmtZ+7YkYxbBGKjg3MW54qyy8nI4qu7CJuOnpUp3WiHqKse4k0xIyXxTvNAXg06OfJ570/eVwHGPbx1qQLt5HWoriTawA709HynNQ07XH1JucihjnOO9IrbxgdR1pxx0rNvUCJlyOtGDt68U5x2zUe7naTxVLYOojr8tVnJxgVM8u3j8arsxOa1gmIgmLEcUKpA61Kq5JB704R84ra62F1IghzntVmFccCkMWM56U6IPu46VE2nsCsSsuM1FIM4I4p8ofGarbzznkelZxjcq/QnyW28VITleKZE25OfyqRQc+xpS3DqA4PNPA4zn9KaBk4705ck81DAbICBUAAbg1Zc4OahKfPn1qosCu0OOnNOhk4Oe1SOdgz1qsf8AZ79a1jqrMNmTF9x2k8Gomyo96eI++abJ69qpb2AdwuDnmmsuTmo954BqTPy5HpRa2oxY8EkEc0rqBknr6URJwT3oYfKw/LNLqIfA4K8dfSpt4JxVCAt5pGatLjkt61M46iQ8D5u1IVz1oAwc5o5qNhioDnFSd6jUkVIcYzSeoCYJJpjH5s4qQY281FITihAODflSE89KVQSppCMDBPNNDGLzkUZ3HmkIxz2pM4P1qrCJEJznsO1LIQR7UxM5x+dP7EEc1PUCKJAhzSjvSdGJIpVOOapgLjGRmmK2SVJ5qQYPP50zYB81CdtwHIpBzUh496M/L70DjmpYDOAenWkJJJFSHHBqM8NTTCwAANgZ96G6YHWjdg0g65zTB6D+OhpM4yB1oORknrTAxDUkguPAI4PSnjoc9KYM4qVckY4pSGMxznFRvnrUvt1qNx17k00IQ84FKUximg5zn86cDnr2NO49BTzQMjJobhaQHPFHQQqjnJNHfd0FB6kigGkwEJOc1VkQlyx4FWmB2803HygelVF22FqQAfMGPFOkPIIpcZ4NLgBeOaoOgwtnA/Wo2U7sgU4kZFOJBHHX1poNBp4BIHJpYz8oBPJpcEjrSKOhzgUguOI+alB5xil6n0zTNxDZpDFIyc9RSDqfekDZ+lKpG7Bp7CJAOmetGMHNNU0oOT9KQx5Py0g4HvSUg5I9qkB/QVSuJDu21bY9Pao3jDNkde1XBpO7B+RUW3yck9af5BBATpVsIBjNLgYz+tV7RisQGP5OnNQsqoTzViV9inHWqmGkYknp0pxu9xE8DA544q0v0+tV0G1AAOPSrA479amW402NcnHApBkgU5umPSkDelR0GOxk4pRx3zikQjNKpzk9KQDhxQPU1HuHX07U4HjNABjPXrTW5agPSE5/CgBGXimpH5YNSAE80HGBVXYDcA59KB04oIIFKOMZoANvBpB3+lOPSm98dqLgIc4zmo5GyPpUhHy4zUDjAx/KqiAEgD61CvMjHqfWpG+705pF4X/61aLRE2FQkkAmgIA5z+FNOVY+1Ku52yegofceg4sCNp/Onqgxn9aYxAXgfhT05Xk1LBEUwABx0pEHGO/epHA59aZx15pp6BZ3JAM4x2pRuDYJ4pI+P4qepJJHSpYElOJ4pq525xQM1Ax54qvJk5GalbPrUZGRRHQRGvAxjpTnGO9IGyvtSknGat33H0GH754pOd3NOI5zikBxnNVewhwy3ApQMdqRRwTmlzgUh7jZBwajj5AJxT5eajjJDEZppaE2LKn8RSjoe9MB5wKAc554qGivUkyBQx469aj7YzSht31osLQUcZpGIA3ZpfvHNI4yCO1AxjZKgimoDtyD3pw6e1IeBxVi2YhXH0qIKdxOBUu4kDPWmEYbBpoBFB3EClYcH07UL9elPA3YyKGwuVnhPBqKMMHAz068VoYJ68mqrjMmV4xVxmS9FdEw9KUKCx547U1Dknn9KmjUDk1m9CgIB571C+T0qxj8aiPByKSAavH3ulNfJFOwefSmlCe9Ug6AoxgnrT8nOR+VNC5GQOPSl6DuAO1DDzDmkIpxFKASAfzpCG8/jTlyM0hBPGPrTgAO3FAyY8rkVWYYbjipweME1E4yuR1pK4vQZJIQlZ0rHce5q8TvGD26VW8vDHj8K2p6C1Yq84zxViJcJ6mo0XOMVaSPH07VM5DIJYgy5IzVZLUM/tmtJ8FCCKpA7XIJxThJ2JaV7k5hQrjFQLEowB371MXJUZpgx1A/+tSTY3qxoiGc4xUgtkcYxxSgkHnmp0B5pOTGVHsj1HWozbsOccmtPrz39abtBGaSqvqFrIx2tjli3A7GodhQj1rcZVPGKzrhAGyvNbQq30E1ZaEPQgDmmtKcYAx6e9P2hSTmopFBA2joOa0VmHMNTcxNNkVgjZ6YpcMDwPoM0khby2HUVa3BJbnsD5B2g8UnbHT1pep3E4ppPzZ7V8zfQoHAAHNLtCjHrSEYHr6U9QMDJqdbjEXP3j2qT5pCCRUZ657VNuwq5q1bYljc4k+YYpSWfC9hTxGJJM9gOKib5GYduxp6x0ewAwXaT3zUDNlQR0FTMQF56461BIMbcVMroYA4JOOcU1c57UEnIpHJBAWpv1GNboQKbtwuc05gPXk0xjnApWDoOGBk03GD1pT0ApD1pjBjx70L8wJx7UMM45z7UKNp+tHkIeMKOOtPJwcdaYBzipP4s9sYpbgAxsb1zSDG3J60dFz2pMZZaauwJQAEyep6UMwYikJy3Pb0pAQM5qgGuSzcfnTTEShY9qlQjkkUgVmzjkZqldsQiDaBx2pRzmnE4OKDjauOtKXvAOQBVOTyRTByPSnHkgHg0EEAil7wDQMnB60hOcYFLuwo45zjmmhTuzTt0AUEkY96Rc5OO1L2J5zSfdqeawwVsEkUwE5NDdMjrTGPBpIBSSKkzhN3eo0A4BNPJG/Hak1ZAO6qfU0oBVMD0zTeN49BS/x89KaYCxgkmoSD5hNTp94jOBUUgwT6U4vsBHnqelPQZUtjpUec5yOamX5YuOverdragMVQxNSkKAVHpTVXaRTHOSQPWp6bAOLbsAc1XuchaW3f97gnkVLKodqvkasxECr8gb0py8gmh0KNgfdprArHwetVvsMVCQ2egHenu4YZHX0qFZcjaetTFFCBu9Jq2rBkIYsMkdKlUEru7VG8ipwakSQGLAGab95bDAdSwPagZ2k9aMAcHinAKoPvUiKyy7G21Y3AjFVlQCTJFSFvmIHStLqysCJ1Pye+akZiIwB9KiJJAA7daec7fesmBC6lVIA60iHoucjvVhFLqc9qiRRuyRQnZWAfyXx2pHBDg9qRXO8r6U/q3NRcAYkDAo3YAx1oNBTpmpbswFGATnk0ik56UmeenNDN+dNDDI5GPxoLYxj8aQfqaXGccdDVCA47de9A9KaowxFOwFJzSYwwQTUbAB91PJIOajl4j45px1YiQHK0qZBxSI2EoXAGc896Tuhg4PQVAykMKn5P1NNkb+EU07gNT5229qfgE49KRMK3WlGAW4obVgGSjJGKgdix2+lSyEgY70zaQferhvdiI3LKach/i70o5bnpQqAufTFUmtUhkkYzyaXcT8uOM0hOF4peuCKz0eoDoyRkDpThwRTF+VTilU4Umk9QCQhcnPFV/NQjNTupday5AyMUHWtKUYy0YnoTlwGIHAp8YYpj0qHhQAKswj93npWs0loAirtOcVL0UUi4Oc9aUDceelYvUNgBxx60gXBPrQTgkdcUAE/MT2pMaHL0ODSBuGBHNKmTnigAFaNgKgG8kdMGjaO9PlABHakUjpWiYDWAftms66k8k4FahUjp0qjc2jTNnvVU5K+omMtZWY5atFTld3Sqtva7OvWrSj5cdqKkot6DH5yelRkFjxTx/wDroGM5x1rLYBygheaRjke1OByD/Kj7qk0uoDQcDIp2MgnNR55NPGWXrinZgMI+XimSP5ahql2nGO4qCXBBBFVHcOgPKB83c0GQAA0xlytMIJjwKtRT3DqSl8/SnMfk4quqnHBqUfdwacoroCYwgsoBP40wxBWytWfLFN2qDilzIBEX5OlMcHb05qUHB9qaCCTSTAjkAKioWTdwamcEjHYUxlNaXEVZE4HrSdgetPkHzAk9DUcY/eH0reL0uIV4sqOKkjCtDt7UYJU+/So42IXApfFe40JICqZzSWoJOaJB+6OR2otX2fK1UtrIXUuOuxd36VEr7vmI5NWicqQarMg2nFYx7MBqyjdzTyQxzVM8MQxwatRnC881bgrXAXAIz2oALNx096Vxj7tCZ3entRvcCRQS3PaiM5cjsKTzMHkZNToFIz3rNprUfoOcbcYqu45zU8oI4FMZeAKmLs7huV5UBANQEgH7tWXcRrjNVDJjk962hd9BFyKRCvPFOLhh1zVQEPgZxmnMCkeM8mk4a3ASWQdhU8Mm5MnpVF1JQAHBNSRlgmF6VTgnEV7Fp2JAxVOdckdfc1aAJUVHISeDSh7r0B3K8MAR8+9XVyMVDHuJ2449alIPFFRtvUa2FlYbR71nXI7hs1eZQ6kColt8sQ1EGlqJkESFouRU0aZOe/apNgUYFInHGOTVc172HYjMZkkO7oKlAKrj0pYzuJB4IoyN9K7tYY+BCp5PNSyHgYpjN83FNXLNgjisnq7h5Cs/OMc1GQc89alKbWBpH/MimrBcrPGd5JpH29hTmHJyfpUO8L3zmtVcRMq45AqUpxx1qGFyxq0Dipk3cfQjG7A3DinLgDimM5x1p0bZU+tKV2g1sLk7NpqqybXGKtseM5qqzZY4ohqAoYI+KmVt2SKqCMsdxP41YgQKu3tVSirASe4pwJUZ9qRhR261kBFuZnORxSMxD/WnsMD6VDIcGrWoyR4w6HPWoI4yT6YqUPmLpTEcE4XtVJNIXUkOAOlRXB4BxUpBP0qFuQaI7hqV9vQjrU2ML05FNIOCB+dOXOzmtbgLETsHpSH5iRjil3YWlA2gkcmp0vcXQRAB7mpFbPfFEafLlqeoA5NS3qAqg4yaUnp6Gk3ZGTTR8x4/KpsmMeOOAKeOB70nANCnOSahgJyOTSSZPSn47mo3ORgDiqTAUP8ALgdaY2AMURgnOfzoP3s0bBcCOMUhAAz+lNG49aeMNwRVPQBpIBDU4HOfT0prLk+1Jzn0oYCZAbngGnZBFIVycCgrgZ9KejELzx707ABoHzZNMyc9OKW49CX8KD0xmkHB96d0GKnYBD972pvTmnHmkK80IBpGabuweafnJAFNZRmquIcSO/X0qIjL5zzT8YalwAxo2GKeMVInTr+NRkAkU7hWx1FSwHE8/wCeaa+AfelPUY6UppIZA2e3SlUH/GnAfmaVVxyfzqr6CEcZTmmR4JqZgOKaqgEYpX0ATqSKeQB0HWmAjqetLu70WAVhxio8HBpznPIpgJKmmtgYwDHPekchV9zS44Ge1Nxubnt2rTqIjYEjI7VGSVOc/hVgAOp9+1RiEHI/WqUl1Cw+MjZ7YpRjBpqrgAUH7uKnRsOg89+1R8lafjB5pvQYoQxE6kUYOP8ACnquFzTTwMDNPqA8DinAAEVFkqQDTt4646VLTBEmMA+9MHP1pN5YYpWzt96LWAf1JJpBnk0zHGT1p2QAaQETSEEc0pct0HXvTXjzT4029etXokGoxgSOetNQc1My8GmKOTxRe4D8cYpwUChVxTl6HiobENbO32qEsRUzA9+lRkfNgetVEYoJ7H60/d1H600IMHB5psnB4/CjS4D8jGe9OHIqvhicZ61YRcLik1YBu08UAY607p3pMjrRe4DgMik6timoxHXvTgeD6+tJgIR3pFJwc/hSnp060o7c0+gC7cnNN4zxTmJ5ApMcUkA3pnBpjjj8acOPembixwwqkAxlyMmmDg1ZbG3B61A5XaMdatNvQGhAM+xoPyAAfjTQ3zZpzncuadncBNnU+tOHA96aMhR60jNjHr3osyVoS7gRzTMbuO1Nbk+tITzj1otYY/7g45qVT8uagxzxU4OVwKTWgD1604Dk0xegxzTs45zWbGIx4yOtRg9AetSYyc0xhgk4poNRg4GOtIcg0seSxLUj7izYq+otNxOoNH3j9aFG1fam+3U+tMeo9RwfT1pQO2aVRt4oz0qWA1+eM1EnLVK454FIgweetUmhWJV6cdqaByfWpB096bwvPrUXHuBOaZgdKUPwaMY59aYdRU4GDzSn9e1OUc80d+nFTcCI4UUhHygd6fIO4HX3phBP0qiRAMj3HSo2zvyfSpjgZqEjJHcVSYxwBx71Mo5piDjPU085yKTY9RrnAJquR0zjJ71YJ5+tMChucU46K4hyAYBAqQZIAzSdvpSgcYFQ2MQ8ZHWoz1HIzTwctg4zR7YprQCME5NRg7TzUh69OtQyZU88g1cQ6kg6EfyoUnOM1Hwe/wBaepIXFPoIeBntT1U8imB8DmlVyAfU9ahphYQ8Gl6jFML5ejtgU7C6jw2CecYpPb0qq0jB8HNNMrJyOQetWoML6k5Ta2QOKjc9s96jN1tbDHj1pwdX+YDpVWfUOmhYgXkfrVrHAqshAyMdKmWTpnFYzu2NWEk+Ug9qqyDLhs5FWZ24GKpMxy2etXBMUh+SB1pc55Azn9KaGyvPWhGGcDtVNASqD+FWF45qKNd31qUegrKTBDuMmk7g+lB7e9NOAakYjYH1NZU8g8xs8VpyH5Dmsec7nJNdFFXIlsBkJ6fnUYbLccZ6mlGCevA7U8Yx057+9dGwIjViHORgGhgTG/HQVKzjB2jPaoH3FCW44ORTWo0ewkADb39aazbQB2p3AYd+O9MI3Nk9O1fL8t1oUKDkZp4GW544pM/LtxSqWZhiqlFAB5b0p7AgjuKYW25XBNPUHZu/SnZcoD0bapNMJwpyMg0u7yyAe4oLbyMYwKt6oXUZj5SCO3FQtkDHepi5Y7umKiYndn9aiUk1ZDIyCQB3pCMnA/GnDOQT1NNGQxIrO9xiMMNntUeCCTTxk5Jpo6E0JDGnk0pI7UHGKTHU0XAXPzYpWX5uKRhjHrTkBOBmqAdnkd6d91s01PvAU7qxz0pLYQ3kg805OpPpSEjcfSjODgdKSAep5yT+FM3CSQBe3eoiWaUAdKmVQi4A5zWitbUCREzJzwBTnby3wvSkLDOBUM7nacU73W4CSS8gjNSxnjd2zVJAWIzmtGGPKdelEV0Bi9ZenamkNvPp3pyZ+9jODim+b8zYHWi7fUQnBXjsaQnLqQcCnKdob1ph4Ax1qXewAQDIQTTGyPlzSnA5FNPPzVLsMaxyelMzg5p3YmmZGDRYYBhkk0+NtzVERyPSpF4APfoadkwJhgyUH5mJoXggkUhB5NLZWEKpJU8USoVQc8e1KnK8ZpJCSgFUnZAyBh+dPiYhsdRTe3TmlTAJPeiTaYDmfDE1BnGW7mpeGLE0wx7Uz61UbdQIohh9xNOL/PweTUE29c7eRRA29ORW3K5K6AtOyjBPeq7Mx4I4p7YbhabgjaxPSojawEMihWHOG61dh5jwaozx+bIrDpV5AVRQf0q5O0BIjeEOfUimxHngYqznHP40wJhMnvWCk0PqMDF5OtOIJbFOCgEDvTshW96d7vQCJxxgVSeZlYgDpWgQDmoVgUyEkcelVTaT1D0HQEkKTVhuHx7VGgAbAHFK5KyH3pTam7oB0THaRnikxgHnmo0G0H60MzNj0qWk9gJAuMvSrjbuaoFduVzxTwxZMZxSaGSkjeCKWQhyMVCrAOfpUqjgmk9BCP1GBQRu5xQpySDTWyCOetHkAituYj0p4PB7ZqEQ+XKHzwe1THLNx0ptpIfQQcnPelYg/wBaRR83HTvRgc4qbCGMTjb+tJMwCg0p4H1qvctmIjvVRV2BYTBUHsaCewqG2cSQgZ6VOffrRONmMax2getMlJIBHWpHXI+lNkG1fU4oV1qAnIAB6+tSEbQMHk1BG/mPn0qUn5gcVUl3EIwwORTG4wafI2TkVGOue1JK6GDDCAjrQuFOe5FB5B/QVGuWIyK0juIlRSCc06MEgjGadxs/CkUkc1mwHA4BGKQjjNGepPek7ZzU7jEye9VpYt8m6rR56daZ0yaqLs9BFTySswJPFT7grY/Sg5d+elOaNQa0butQQo5YYqTAqNMbTjipFxtJ70pW2AiKEZYHmn5+WjPXtTNxKntzS1aGOB2gikU84pjv8oANIM9RSs7ARznnrVeFiJMA5HpUt2CEJFZ1rv8AP74HWuilG8RdTa7UqjNMU5I5p5UmsHZMYgXJOaMACnBfkyaUc/SpYDGGeaAOKdgHrSEnb0oACOc0rdMHvSFvlyPxpxORj0p2DYgcEHFPXjgdKaw+fPanrTDcQj5sjmoSpep9uTmhhgcU1K2wFXDFMelKF421ORkA8UzYN26q57gR+WVzUiD5PelkOI8gcioUl3CnrJXAmPQYqKQEjOalJ4HpSHgED8anVMCIDauTTkUbaV+lCnAxjrVdAGYDA5qGQ7eOlWCAoJqJwGX3pxEVChL81EY3HzCrAJVjkUAlgeK3i2tRPUiUnyiMdKISGxzTwAqmqyFhLwPlFXHW4PQteXjORxVdojuyOtXBIDx61EF/eZ7dxUJtO4NEqDMeO/rT40wMGhXXgCmmYLIA3FQ+Z6DIbi23S7u9SIm0AYqwcFsjpTG5bI6Uc7a5WHLYjOd2B+NRvhWz0qR/lOaoSTEzBetaRTb0Bl3erHmp42GQRVOKQbwD19KvdEJxUTTQx0jEjIqJpBtOelIXLJgdarmDnrzUwiurE2RsPMct29KRo/UjpVhY9vOKa5Gea1U9dBW7leJSeR0z1qxtDEDHTvSxbUp4dS3PSlKd2O10VpQQeBxRG4AwTViTaR1qusAJJBqoSTWpJIH3MRT1iVqETacn8acG2vkDms5PsMQQ7RgGlHXBp+7j3qF3w4FSrvcNiRYwDTjHkcUiOByTTt49aTuOxH5fGSaUIPSlaRQTjmq8t2E6U0pMLIsFAv1qHIduOBUH20Mw560rTgEHIrRQa3GWlRQRTtu1TmqjXPAwajeaQjrik6cnuLQneYK20mmtMB3zVLDGQHNS7SSRWqpxSC+oSyFuAOtMihySW6mp44QOo7dakwF7Uua2iFa7GqArY7VKCe54qo8m1iBR5oA9Go5W0MtBQTjtTlXaaiWT86cWJBFZtMQoO9se9NMIyadGpA3HqKfnuad7PQe5Gq8bTT44to9qcuCfanKcdDUtsLCN0x1qvLLtxnsatEdCapTxs2fY04Wb1AlWQOvFRtkDB6UyIBBz1prOe9aKOugXJkU7sZ4+lOaOMHNEb7h0pzDFQ27gMkPAwajYcZHWpwu8DpioxlTjtVRfQCE8DgUpJPHSpjg1G5AGKpSAhOTwakQHZiogCCQenrU0IIB9qb0WgbkoPy4IpAGIxwKReTzUikE1m9AGjpgjmmrlScHipHxn60iKMdc0c1kAmSOafg460DHTvS9ulQ2PcQnkD1peAMUm0YzmkwDQKw72zTHAAxQD6mmNnHB/CqQAp2jk05cEZqPkqcnmheAapoCQ/MSabt3NnNKWOKRQMAn8KWwD8en500rkEU8Y9OaZI2wdKAHAYXilVfWmx5I6VIM5qWG4z+LpQVJbNO4zxTenegAJ/wD10mc9+c0dsd6ZnOeaaAXPOP1pd1N+8RgUEbuRVAOXGD6etBGEzxQBTjjAx0pNgIrcDPWnY4yah3/nmpA2UzRYNBQSce1K3U+lNQ46ml3bgcUnuA3PzVI3K/TtTFxu68U45AJpMBG45qNnxnpRIx28VRlZunb0rSELsG7IuK+eO/pSGTGcVU8xsH17VIuTHg9f5VbhZ6gmTJKHXFSKOtU422NiriHcc1LVgQxlyORUSDOSasN93GKgJZTihO6AjZtp61Kjg5FQKhc5qdFC1UrbCF6kjuKTGOT0p6gUNjtUpjBuOn4U0jJwTinYzSlccEc0aCG7MAelDAKMk/Wn8Acmq00hfgChXbDQbNIMjb1pUUnk01IwxzU4O1cYq27aIBFXHB5p3THpTAfm5p3AOMVLAcM59cU0nJPH0pckdqaCc0hjvrS9/UU1MkHinjgUMLjGHJ54oH6U5hximH5eAeO9ACSOVP8AWpVJK8AU1VzweKd6f0odgFY8YFMwOD3zTyc03GT7UALwvQU3APJPNPIyOKbnj1pIBAoxnNOXgn6U0HjpTt3AoYdAJP4UjfWlIzzim/1oSENzzxS7uKbt+bI6CgDB7kVejGKT83U1IckcZpi8npxRuOTU2DqOHuKcG74pm7BpecUmAHp71GR9Kl4x0phFNANPTGagkT9KmPX3pnOT61cdBECg7iCOlS4+TjvRt7469KU9B79Kpu4/IhYEjNKCccnmkUMz9al7YHFVJ20JWw0qSvHSm7TUgOGyTTGcRtnGQaSY9Oo/jp3p4BWo0bJz0qTOGFSxj1HFKTt6d6UY6etNI+asw6DgccAUxidtOB/OjHFMViOPoc0vAX+lA5yRSgAjPemMj+oqNQN2e1SEckfpTMDNWiSVeuaPc9KRTxSk4HtUlDXIGTUIbBxmpJCSM461GFIGO3pVxWgiyj7iAaZOOOOc+9EZ5pzDcMflUbMZXi3BsEVbAzjmokTDc1N2AA6UTYrBgnjvRRna3tSA5k68VI3sOdcjnvTMfLwMmpDzTWGfpQmAwio9o3GpTyMDtTcYFUmAD5fWo3cKue/pUg71C8e4kelONriuOD7lyP0py80ijA4H5UqHjHSgYoyR6U/HHX8abnCjIGaQyADk0nqAE85ppbjrTWmXIqJpfmyBVKLYiUtyTnpUbngcn2NQtISSDTPM3AcH3rRQe4nJFkYIUg9P1pAxYAk5qINkccU8H5cZoaAkzkZ70pOSKRCSDkUoAFSMDyc96UcH60vfrmgdMmkD7kMsW8571VJdTtb860Dznr9KilRW524Iq4ytuLbUouoJJX8afCQowcUk0flscdKhDbWBOM+lbJcyJL/m7SR0FMe5G8KOMdapiUk/Pz9Ka8m4YH45pKnqVoWXuW4I6U0zbs1VUnJGMg0AnPAwavkQrtlou3QcEfrViEEnn16VQDEHAwRV+Jmbkjn69KiasgRcU4PB/wDrU7dnIquXPA/WlVty+9c3KMnHTmkKZOaZ5jDgUrPx70WYWGTOFjyayHPJJORWhdPx9KyzwSRnntXTRWgpIM4J44pVfnknjoKMB8s350bV3evethIATtzxgd6a4O084UDH1qXKFSccVHITsPfg4oT1Gm76HsRAznPNR4OMmnjaAQaax4wea+XSTV0UKpGR1+tPXJBOelRggmn8/wDAarZhcQZp8e7OaRkyN3b0qSMhUOeoHFaREJNGVIL96cAE5xxTfnm5x92lMmY8Y6HFCaTbENYbiVHT2qHopH61LxtyD9aiLKc4rN6FDDx+FMb9TTz93Hr3owMBc9qzUXYYw42Uw8cU/G1iKZj5jn8aAGjvSHpS/dzmmM2Oc80WGOJORnmnr93NRcnFSDgCm9AHr8rZ9u9KvJLGkUjdyO1Cn+VHSwgxnPajgKCOtOBG0r3pgGeKFZaAKvGM9fWn/ebj86Cg2jpThwoNVK+wXBfmzxzih4gV/GnRnbk449afy3Qcd6f2SRFijEZJHzY4pRlI844PWjK7QPelc5+XtVOXUBcgREetRkBADmnMQvyUyQfMMenNDvuhkMkpAO0cmktizZEn1p5G0n2pmGzuHfpRewIex5I65phOFI70/Gxveo34b+lYjGn7tR9M09uKjJ+Y1SGAORjpUoGAB60xcE8dqeASAaAJgQwC/wA6XggqKapwx9SKkiwAT396XKxAmFVlzTc7U570KM7jUbksuPSq1EVricRLmorW8V3Ksee1Q34IGelZwYxuCvSumnSi43Y09ToSMEmmu4EfJqi9/shH96onaa4Ktg7aXsJN76A2kLcXu+MovWnW+9UGOeKVLQYyRzVyOIIgzWknGmrRBX3YQYJb3prLgHJqyqKq571DMc8DvXPKTb0AjjKnJHarCsXQcdKpQxkOfrzV9BsQmiokg3QjjtTN3zbT+NPZc81XkzuBBrOPYCy+A2QeTUR3HketKT6U4HAwKb3GCngnHJpI2BJyelOkAKADrUJYKQp601G+gEgYGXFPwOvemqg25xQjfKc1MlbRAKeQeaEUbTzzQoBbFIpIcgmiIhBHnJGM0mwkZU1Io7Uu7blR06UuYBgGT05oLEZUUpG1c0ow3Pei4EauQcYqJ5txGanC8k1E0IOSOlXBoTHmUbBkfSpgQoBqsYyy47DvUjk7lAP1pNLoMehwSexoIwCRTSQDjt3pxIEeaVmwdhuO4qGQBh0qUEYNM+8SOMUJagMtkC5xU5PNRo3zYFPJGc05SuwDdls0mMnNHBJpVyBz0pbDItoV+OlPf5iPSmlTv5pT3p7oBsmAAPzpm0n6USNn+VSJxH61WyAQDnaeaUcNjtSjpzTaXQCTJzgCmsQeP0pVODmmsctmpsA7ohFHTAxyab1cCnk/NQAEgVFJIOlSHqWqGSPeQeacbXEDOq496V2G0U1YzgZ7Urf3e1VcY0feHNSk56UxVAORTgeTSbATpmmlTtxninsSOlIeR05pK4EbIMULyTScliKeqbeKu9twEbDJyKiEaxnIGM1ZYDb6VE4Cr60Rdw6iJg8ipyTt471DhQAB3qfrjNElYAHTGOKj2lX9qkGQaQkl/aoAAAPxpDjaaXrR3wBSDYYp7EU8Eg8imkYfpSueeM0wGkfNSBgARQ2TyKjTOSDzVWAm4x6U1TkdvpSMcLimgnjjFFgHdCR0pmCc089KaO5oEOHKYNU8AEgVbGMEVE0eW4q4N3GLu3cZ6UpbA460xE+enlTuxTfcZG0h6EU89BQSMAYpcDANICNhxQPugU915BFMINPdE3IJCQ3I4qNnKn2qaQZqJk/d+4rWLDUQD5Dg/hVfYVYnNPL4XilQDbz3rWLtcB0agVOoBBHemDaq4zTTOpOBWbTkAKm193bPFRspeTPGKimuewNQrcso9q2jGTVxaI1N/YCk39cGqa3AYc1OmG5zWbptbjCRg5x+dVGj8uUNVwsmaimZfKPP404NrRCsVxOouAe1aX2lJI8A81gsAzZoVpd3FaypKSBNGzFneQTmpjyc+lZSTFSDu71fgkWQZNYSg1qPQkYg1XeQAnPanTSquRms6W5BLAGnTpuQmxbiZtwwfypwnJQDPNVPmbnuatwQgNuJzXS4pRJW5PCJGA3Va4QYNR+aAMdMVA0zE9eK52nIrYttIAMZ70LIu3NUmkHrzVeW6KDC/nVKi3sF+5pvche9VZLrLcfhWYXkY5LUKxORmtI0EhcxdF2x9aUXbAZz9RVAggY6k0AFueoFaeziO5bN22Rg5zUMkpkXJOKjWMlhz0q1HaFlFDUY6ivcrrtUcE8CnbskCrUdmalNnjnvUupFAiqrMx2g/jVqOBiuSfxpEjAcE9avY+XgZrGcyiBIdp5p2PmB7VKU4/pTX+Vay5mwGMQMcVXlZieKncZIP51Fg7uBzVxBlZkdifWkSFy5LflVo4z0NPyOlac7SJsRcxtkVLCS3WoiRuHpmnrySOgqGroosBsjFNJJYD9aaqlc4pGJU+p7VFkCHDO4gnPqakA6EVEp9etTKPlwaUkAuRmkYAjGKacK3XpTqkCBo8DpzUDxMxB7VaZiWwRTgdw5q1JoLEcSAL05qRhTgOfb1pDjPFS3dhYjJwenFRySICMnHrTrmQopxWRJIzscnNbUocwXLc1x/dquC7OMn8KYvI75pwJ3YA6mt1Hl0E9WXQflAPalDbckmmRJ0JqRow4rJtJjew9TknjrU6qO1RoPlxUp+VQR+VZSYCbcnHaomODipsjHX8qic4pK9wANn3pwJJ9hUIc5xmpl54FNqwAaaxbPFPPcdvWoXYq2KSVwHE8Ypj9gKUcnjmnd85qtgGDJIA6d6NozilZgvJpEPrT13GDjkYFEa881IVyM02NMNzSvoIlHQChk3CnL1oc96zvqAwYApM5NMZxmonuVxweatRbAnBxkg0wtkkVWa5zxUqMDz1quRrcOg9ie350eopT0pQB1pBpcARtpFPXnilPJ4poXB9qEDHFsY9fSkLUjLggUbCDmjQCJgd1OHpzQwBFAORnvV30AmGOnWlHDYNIo9akxg5rJsBjHb3pTkrxTmXJGaMfLg0X0AqvJ1zVcgHmpJuW29BTcfNgc1tHREjFUMRirIUAcCkVAuDTgTjmlJj2IGGckinQyY6mlIwWBqEJgZXtV6NAWxIPb8ajdlLGqcjsg5NNVjnkmmqXUdy1EfmPpU+R0H5VWiGTUi53daiSuxXJ8Db703gfWnL6Uh5JqEMF7dvegk5px4wBTWIA560gIj8xOelB2gc0vt60NHleasXQaxA57UAhhxQ6HoKcigD3p6WATHrS9MYpSOetNOTn1pAKeKYevT9aVjnHc0LyefxpoNx2QB70A5yaByM5pRzn0qQDoOaa3Q4pTwMd6TJ6ZpjYqnrTlORyOaRQAOOKXOBz0oAUUDgelIOec0hP5UgHZ4z1pnTpS5+Umm84pgLnoByKcMdulMI+bpxQDyaGgHFsdKYxyRzQxI470uenaiwB069KYCec05iM4HSkZunNNAOBAFMJ60hPAOaQ9cE8U0gJFPBp/UGmK3egYHH8qVgJAOKRjSA8Un9aVgEyM9M00jninDk9KDwcfrTAjOOufypRgjNKQcccUJgqeaYDAApODiod+GPpU0mB+NRBcnBq492Ji53CkYFl9hUixjB7ZpT0/Ci9noHqQK7hxxVheW9artOqsR/F6VNESwxRLuFywB0ptJu460Z9KysMOhxQTTl57dKQjjrTAZkbsClHFMjHJ9acDycVTQkNP3j6UwDLnHSlYknk04Dgn1p3sDFHUdcCjOScdKCPk6dabjAwBzQMQj5hkflRjB4pW6g54+lN3YOKfQQ7GDinZOelRI25yc8VKOORn6UmG4pPP9aQP0peSajAO/tikkh6EpYmkB+UEUYJ4x0pyjnB5pbA7jg3Tihj/AA0o4HbNNI5yakQnQcGmnOQO1PJyaZnsaY3uJ3xSHn6+tGcH3qMyjBHersIkHHH503PPqKrvIxPFKrtjH6VXLoBKSQxFNlIZeOtIeQTnrTXyFppDKyKdxz0pVyWyKeoYYz19anVMdB171o5WISsViC2TipFTGM9qshB1o8sfSocxsiVOOB81NCHnvVjHFN2n8qnmCwxQx69Kdg45p5B60AZFLmGJgYpuQM4zT88+lG0HpRcH5DDJjj0603OeR+VBQk0u07eefanoGrKNyw3YJP4VVOG5JwfWr9zCX5A49PWqDQuCVB+h9a6abViW7jwMfN1NQtkhSQM+4pVyMA8VIQcbtx3H9K02ERLjPXn37U5RhiSP/rUqIM56mnKNwGOpobKs0TQoG65JBq5EhXgd+x7VFbrkirgUhee3euapLUSIZFw570xSM1MVz+FM8vIJ71KY7Meoz1P0oJzSBeQc0rkYx0qeobFO6I6Y6d6o5J59ferk5DFvWqTHIxgbfeuunsLfUO+OSP60/oOnXrSIpGRk01eAM8jr9aoaXkPTbjOeKax+Rgewp8eVzzlvpUcgwh+np1prcS0PXyc8gcE0jAsxPTipM4UD0NNPOS3pXy97bFjo1C4OOaePTHB7VGpOAtTR7Vzk9PWtIu+xIxj8xHbFSFlCYx1FRkBmLHj0prue/aqv7oEquyJgDrSSKVHJwDSK420OxfBJyaJJKIdRhIBwOgqHB25x3qyyL5RbPJ7VCXwoBHBNZOKWjGRtk4pqklsmnM3zkDpimZwTzQ1YYhfncaYCTkkcU8gbCeKQfcNFrAhpPFRs2DjvUmeM1EV6mpQyQHOCKB3NIv3ad1BGKAFDFjxTkG1uc0i4UignBAHemn0EPDfxUcZ4pgbDbacpG7J9KGugD+qjPank5XHYVGvztyeKec4IHrTVgF3EpgDip0kCx+9Q7sKAelOXCndnirTalcQowWJI60YDPhc4pkko3Z7CpIZVUhh0qkkwGkgMVI57Uudq7SOTTv8AWylgOMVG0p6nFQ1YW4YB6496R+AMYxUYkDZ579KA2eQOlTKL5RgxGBjk0xwcAmnH7uaac7RzUXGRHk5qNs5qQgc+tMJ700Meh5NPB+Q1GMCPpzSr02nvVAyWI981IpIBz0qIAKo9c1K2dmO1DvYQjSbEJzVeSdAmajnJdglVn+8FFbU4JxERXb+eQoPekFuNg44pDA3nc8e1W0DbcHpW8ZKMeWLDfUhWzBI4zV2OMABRwKdGoxmnkBWGK55Tdh6DTEAevFPI6GgH5PegtlQAKxd2BG74Q8VBGSzEmp5BkgVGeB8o71pHYNB4XC8VIeIQM00/dGBzTivyjPelLTcBrZVB6mqT5Rs54q/jgDPbrVVkLk04uzsA5JORzTvM2vz0qJF2jpUipk85xTaimMnA3tUckeZQ3cGnBtpHakcEt/hU3dxMexwQB0puArg+tNbj8Kr+eDMAapR5paAWZDtfINNDbskDNRXEo3YX8afAuAOeM0+XS4ImQe+OKAepPWlYbTgc0jYAA9axGOIBQZ60gyD9aOwHpSNnOaLLqIaXAYg04jauMjkU4qMZI5pPvDr0p2tYAXCwtnrUUZ3ZDdulSSHcmemKibIAx61Sd2Aq8kg96emApBGaajgnOOlKp+Ymh36ANUHk0xT8/tUnXI7VHkA/SkkwHKcSZ9afjrnpUZ4I96eDnAJ4pPUYqAY5P0oJJU+1BB3YFABBIIoXmAhOQfWo3OBTyQB71WlJDc8VUEmxCn3NODc+1MPVQT3qQjDYAq3sPW5GzNuyOlOMhI6UEe1MK4wfzo0sBaX7lJ0PIpqtuAHennpWctGAwEg8U8nOSRTM4BoVuMGklcBf5UvQ5IoU80D73PSgBcE81G4HJP41KSMio5B8tOIMFGBxQoyc9qUcKBSZ70a3ADzxTQQeKdnnFNGN2KLdAGAnfyKkHLc0ZGaa2Swx0pgPIycdqjKhjUwGB1poA3H1pK6YEW0lqlHYdqbtwSc5pyEY96bbAU5BpDgA807OOabng5HNJAKPu570uB1NRhuD6U/I2Y702Iafu0HITPanMCQKRz8pHakhkSnAI60mCrZpcfNjrQQT71d7AIwOM9M0gG4e9OZCUxmmqCtMBQGAPPSmqcNnNPYUzGM0AS/7X50xuhbFOXlcHpSYypFStwId43DFOD7uR1qEwYfcD+FSAADgc1q0rDAgk0gHI9aeuWwaaVwx9KkQ/IxwajPSnAjJoIzxS2BkLqMZJqJyAKLmQIvJ5rPluiQOa6KcG0K4TyhTgVGtwcYxVZgzSlsmnJGQvFdiglGzJuyV7k44PFQmVsnnr2p5hIB4xR5OQKqKV9A3ELDaPUnrSLLnjGMU9oCo3U0KGbPQihKLFrcVZAScnin+eVPBqMRnfyMCnmAu+MVDSTHfQTz2BPvTWnYrj0pxtZFXOKUQcZ9uar3Au2QHG3PU0q7guM9adtPmAY+taEMC9SMUpNRV2LcqJEzHPr2qwu6JeOM1c2BeMVHKMqABWDqXexWxSkV3IyeajW0Jznqe9aCRDIJ61OwGwDHNDrdETa5nx2oQDNWAgXHPSkXJcg/gae0ZY+3pQ5X3Y1axTnYiQMO1QlpGb5Qa0FhyxFDRKpyOlVGoloBmNFISevNQYO4jb+NbO0d+lV2iQHIHJ6VpCqJoohNxIxzSmPAAq1EmHJI4qWfYqkntVe0s7Bq1czhnI4qRQCcYpvLsdo49aswxBV3HmqewlqMRW346+lXhlUB71U37X4HNIZJep6GspKUtS1oacXC5PenORt5qtDKzJipScoc8mueULMZCGycircfOOKqRcZwKtqflGelE7AgIwaY6hgM9fWpfvYprAms7gVz1pCAOlSHOCaY3A4rRO+gEBJDEU8EkZxzSP7//AKqapzWnQER7vmxmpoDzkmohGFcA1YQYbpRJqwK5MW7VGyjfnrQSGfFNbIGAOazSsL1Jcgcmnq2ee1VAflwwqyCdmaUkPfUdwetGcDFIFJOcU4DOPWoAhZ8P0pRKoUkflUpQM2CKhMI357VSaDUesu5Tgc00PTwmKCnfFK6DUq3jfIQPSstY2J24ya2Wjz1pogAO7ArenVUVYlrUoxwE4zxVuO1AIJqfYAuRSqD1PSplUbHYTy1UdKQIAM0pOT1oXOzrUXYxUAAzTz05pgBz7U45P+FJ6gISD9BTCDjnrT87c+lNJHXH40ICLYcZNPjJUH1FLnn0xSZzwO1W5NqwEgIC5qCQE896lzjFBAI9KV7MZHEu0Z9aVjyaN3OAKVhkjPSi+t2IYQGyKaowduKXkHinjJ5qrtAOP3QO9Jt704LxyaOuAKgBVOBUE1wFz6VMzYU5rHuZyZNqirpw5mK6QSzs7fKTxURboSc+lSxR5Ap/kbjkcGutWig1IFyAOaswynkDpTQhKkYp8aLGamTQIk8xyeR0qQSZ4p4GRmoZcqCQKxtcCVZAXAFPYjAqpGT61aCkjpSkrMY5TmkLkHGKUDBApzDjOPpWYELKSp9TTYo2LZNT44pqNziqUnYCVcbvWlIyabkg+3anDmoAd0ajGc56UdT/AFo6CpYdSnKoV8k1BgliRwOOasXKZXmq5JQc9PSuiGwFiPO3BHSn57A1XWQgZNODgipcW9QsPf7hz3qBmCp1p0r/AC1VJLDGOtaRjfcGMfaxIzUkKZPP1pEgZjwKurDtXrVymkrEq7GhcDrTwgzx1pGJx17cUsbAn1rHUZKFHXuKXoD6UvTNIw4rPqNiHk1GT370pOBik3YNUkIafvbak7c9qRTlxmnA/lQxiCmn1pSeeOKQjA56mgBuctwKZk84pwyeaTv71SEITgZ70mO/rThy2KQY3HvTAXkUq8E470nH40ZwDzSGONN6D0+tKGzmmscEY/KhAx/TNMbJGO1J0HXmnDI49ae2oCqp5J60c80uemaXqMd6VwA9Bim4IwaXHakzxx3oQMa3UjsabjBJp+M80McHPamhDAQevNA5HNOHTNNByKY7IUHBOKQkE5oI6+9IeRQGwAZOaULg5NIDlRnpS7gBQAucDGKTr83p0pAfWnZ5x3psB4Py80YB57UmB/hQDhTUAJzkcYFB9RSM2DigA49qfmIaxwaXGAOcZpSOQO9IcZpj3Gnkg4poALfjTyPmyaAuWFO4hVGO3FNccHFS9qrSkjtnmiOrC9kRGAP0FTqCi56Gmxgg88VOOeO3aqlJiStsNXnrml5pcbeM8UduKgoFJ60yRmC5HSnAHHWlZQVxj8KOolcgVjnkc+lSA5phXDAnpTjyMiqeo12IZgc8daAxCj9aewLH+lLtz2qr6E9QD5X0qTPy0wL14FPxxxUuxQ0nkA00LzxTmHfGc9qVR8vvSAFTBpxyDjPAPalHGaD15FK+oDCev86Ve3c02Qjr3pqsDn17U7CJx7UA49s0zeelMJIfmlYZYJwKYSM8VG8oGTnkVB5vIwcZpqDEWd2Tj9KQthsHpUCvkmlzuPqRT5QHMwx9ahYgZNByfrTWBB69atKwmMVwcn9afvAHHSq7ME+UVG0rHHPI61ry3DZI0FIK47UxmCnmmQygL1prnLEZ49ai2o9CVeWyPx9qsKAcVUTGM56+tWUqZoESKMLS9s55FN3Ed6NwAxWdgHDGeeaYcZ5PWgNkYGfeg5wT3p2AO39KTgDOcClxxS4yoBpj0GFxjmmhvx9KVwD0NNAxxTS0FomOOSev40AAnvRwoJJ5pQcDrQG+4xjg0xl7EckU9ue/PrimGXjB6gVSDcqyW4aQMfwpsqALgflVljwGHeoJfmBA5raLfUkrANtPfHpUi4xyf8KkjUFff1pAo3Y7fSm5BsWbbODxwauHOMGoIAAOevtU+PmHv2rmm7sYwZBpB1xS55Ix+NIOaBilQF5FMYAAj17+tSNnbiq9x04496I6gU5iCevAqscgEYyfWpXOT83PvUYYAHt612RVkSmBycjPPrSKCVHrngGlc4zz07mnErjB6470wtqR45Jzz6U6TcEI6nFJjBOOc4FK5LRYIzgGn1A9gAHXg+tNZhkHHNJEBHnuTzTsZIz0FfMtW2GIDjPB59KXPB559KT7smVNOC8lj1NUlbQBzgBRg5PpTGUALjrmpk+9yM0jY8wg8GmlfYAUAcDoaXggqMZokZMKAORTGYEk9Dim97AI+AdvUVE2X64wOBUocBSCOSKhwxxnpWctQG5ByB2qMAEYqYKNxx6VDGfm57UrFD3AUBRUe7A5p0hONw/Co2GfrSd73Abzg8daQcLQTwRmg/dx+VAIUEHj1p4xjHeo1XA5pVJyTRqPceMkk+lKPmOe4pAeuKCwCjHahbCGM20lj+VRRz7iTTZm3KcUkcRCg4raOwIuxNlBjuanI+XIPWoYI9o/DinkkDmolo9AsRynHrSBmaM8/SnnHX3qM578VKi1oAMPk60wMR16UE547VHu8xtoHStVGTYGnER5Pynmo3TKEnvRbKSh7Ypx5GO9KbEVBFtYnPWpgMD2odeR2xSg5FLmbXKwGEfPQT1pXyHwetNAGOaxvbYYwjIzTcYFPGeTSY4NCGBAyKUgZGKM8D0pzcY9Kb3AVuozQzccHikzk9eKQZzzzTVkIiKgtnvTBCC+cc5qVvvggUEkcjvxVqTAYUBNLtAUjHepNnemZzmlzXAkjIBHpTmAZiRUa9ead0Gc0pdgGDJY+lKh2knHFOAGwnv2pIznNN20AcoDZNRrwSOOtPzgEDoaYcKeoogwELZc/wAqexyo/Sod2ZM9vWljcSMyfpVcrbsBOxwgwOtM24BY0rE8AfjUbPgYNT1DQcFAG4jjNGcjOagklwnWo1kIjOTntWkYXTGWQ+c96dG276+lQRHIIzzUi5VcihpLQQ9iBkGqLbVkzVuXkdarfZ9zYpQsncGh0gDjI61NbFsc1FEjqfm6DtViL5VIxTk1YEPBG/Pamvydw6Uu3cCPWhRtj2nk1jfQBUO7nuO9KBlh6UgGFAxS5K8dqYaDvvEj0oxyQDSNkMMfjQcA5BpO+7ACDt2kVBK2xcAVPIdwHHNRuoKgmmmBGrDbkdadn5eaAoUdKD1Hoabs3oA8HjNRHGead0OPShueKSAYRkgGpWHzcUxlO0Z608Kdoz1p9LDHAg/WjHPWmYJb6U5RubJqUrgAXk+lQOoY81KG5b0prL8hOeapIRTP+tAz0q2SCoGMYqsy4YHPNToMLmtZapMaEOM4oXBHvTWUknmkUbRxUICVPvcdqfnJqEMV5FSqc4btUtMBAOcUgHXjtTupNAHzZ9aVxCDAPPWgNzRj5qD81PpcYpALdaacmlGTTGYgmjcADc4J6U4N85AqIHLEdKUHnA61TQD8AE01cZ9/WnjkGkAytSgG9RnHWlX0xTsbaNu4jFMBaQ9fannpUZPSl1AXqxIoQDPtSvkYwaCeBxzT2QB/FxSMQTjBpTzimt1681KAaFwxxQ+WIx2pYyS3NOTnOatuwhecUnGDml5zTW+9ipSYxuDjigcgmhTzjvTDw3XrT12Af/D9O1JlTgd6UMCPUCoV3CQ+lUtwJW4AxTVHWlJpm7DYo3DqPzlPpTJH2j60McHANRhstjFNIGSKBjNGOeOlJjBx2p2BikA1ec00feIB5poOGIpVHzE1dlYZIqjrikb7uacegHHSoySOlK2oilNH5jc1Xa0XOQOlaBAJJIqDOGKmtoSa2Eyl5ADZxnFSpEvO0VLjGc9e1KrYU5HFaOTsJblVxg9KcwVUzSOw8wEdKXeDwRxWt76gL5YZRjpSG2DMMDB9qkiPO0D8anHBzWcpSiGhF9nywBFOSMK+OasowbtxTMck/rWbm29QsRMuW64qpIrISRV5zg/1qtLgjAqqbswKo5bk8mrkSnAx61Sbhsk4q9an8vWt535QRMAc4NLtycd6exHXPIqJc+Zz0rlvcY/H504rkGkbO7jpT+SOBU3AqsMZOBUkXzRc96ZcDC5pYCDHz+taWfLcQq5BprnjJ608gYwKicfKaUdWMj3dc9BUZ25z7cU5F5PWon+Vuvy1vHXQkcNrcioZ13d+KA6g4UUjYKnvWii09A6Do1VRnHbrT9yfdHrUcUZYVYWDkHFKcktxaEEcJ84Zqd04II+lSyOka5PXvVFrxpD8tT703oVoicPt9hUnmZHHeoF3OBn8qnVUA6c1Ml0GKuT1qzGBt5qurZbjgCrGfl4rKeoCnoTUJc8AVP2qMrUppB1GAE4pr5wOPwqYDGajkOF4pp6hsVSeefypyAKKYxAPTing5+lbdBdRAoL+tPb731oABpzAFc0r2YdCIIQ+4nvUvINAWlZSRwMVLlcYiqGOT1qcfcwKhUHPt71YB4NRJgA70i9SDThwD70mAO3FSAZ9aaw5GKeRzk0080JgAz36UpII4ppOVpqjB60AOGCTS+wPFMUEE0pwBmmAH0peAOtNDD1pTgc0xjWGQT3pIwe/al3c0of/APXRqAo5NHX60iknmkJOTRbUQFSQTTeD2pwOetIRn2o2GI3Y0ikKPenAgjmm7csfSmIVee/FRTFl5BqQcVG+SRVRvcdnYarMOg5qVjxTCQF6Uowee1N6iF4z0pfTFRlwGpykHmk0BKOV+tOA701CCeaR2wKkCGeTjFZ6pmX2qacEsTTYhu5HaumGiJ3HLhce1OySOD0qKVTv5PFEZ2Hnoe9X9kOpN5fepAgCAnrTBIrdDSq4Yc1lJSKJVORmmOM8EU6M5PSnMATUJ2Yirgq59anikz7U1kJNIxzkCqeoFrOfelJIBz17VFGMD5qQksTzWdtSiXotR425PrRuHakBzRZiZKRnHrS80meKUHPHSpuA7PHPekB5OetAOaQHnmhAhGG6q0sQZjzxVk+gpmMg1Sdtg8ik6MeAKVUK9fzq0AB0/Shh3rTnFYrbATipo4gMcU8IPTjvUgXNKUhoYFA4p31/KnMMHpQF4zWd7gQOpJPpTo02/WpSOM00tx0qrvoFh3Cmm5/Om5yePwpQCQTSS7gNPLZpHHHpUmBjmgJkc07gRqAMGntimDIwKcBkc03uHQQjmmS/KuacWwevSmyEOuP1oQMqiV87fWp88VAYipGBUqk7MVpKz2FqDPg5JpAepzxTDx9aXAFFgJAwI/rQcE9KYDzkdKeBlTStYegA0MxGeOOtJjtjFNwf0osLUMktz2605SOBTOgxSIQRg0wJt2CMDNO3ZNRE8nFAbnBpWGSluM03J/hqMNz7U8H060mrAO74zSHpQOSBSnAHvQGvQYfekxinHmmAkE56U9wHHtQByfekViOtKDk0WAOvQcUh5H0pd34Gk3cYIoQhBS9MY/OlwO3SgdfamwAN6/lQTu6Cl75JNC8cZqQIypLgnpUobjimkjPTJFHXinuAhJJI9utIDzgDmnlcU1QMk+tFxi9O/wCVCg+nHY0mSRz3pNxAIAoEPzkYqOUbiBTkyOSKVgQpNC0YMjAORUnPU8UlO/iptgBzyR2pOcZNPPIx1pmRyDUoewvP5U49BTStOoAYRn/GkDDGMdKcTjvUff8ArTQdRcbfrSZwxz3opD61QhS1LuAA561CeG5zTS/JJ5FPlGTFiCBmlzt9zVcyEnP60eZk/SjlEW94HWmF+CR1qAykAcdartcEOecetONJsL3LDucVGkp3DNVvOZm9qkzwMDr2rTktuK93oXBIM5zTC2CDUeTt5P6U9euPSotYYhbg56Gmbsk1Iee9MY/Ng8VSFcVSQDjk1IAMZqD+Ikfl60ru23IBFDVxN2Jt2CeOlRM+0dah3tn8OaR2BG3HH8qaj3G31Krvk55J9T2oUjJ29unFSpHk89DU6QBQCRWzmkrE6tEYyAAOT2BpwVtpJHFTqqg5P4U9Qo96ycx2uRxJ81Tq2RgfhUZZATimiU/lUO71HsTA8+9BbI/lUZk55phl568mlyjZOevuKN46Zqv5m7pTWfI5FPlFfQuFwRx0oUgdv1qp5mec/SlEh3D1P6UuQZa9jQQeajSQFiM04vk4pWaDYQ8Y5oOPbn3oJBWozlWAxTSBDuAcE1HKFPTH1pc4Yg8e1VpGySM8cdauMdReZMXUp8vBqsxAyfbNIHKLknJpjuWOcYrVRFceGJHHBNLArGT5s/SokXLDIycdM1owRAID1pTaih6Eu4R4B4JqTfnHrUUkW88duhp6jA9a53Zj6iEgnHc0uQOlM53nj9aUY6miwXuOZuDiqly+SB0xVlSec1SuRluv1q4LUPMquxUHqOwNRDLBjxxTzypycKajUMe+Djg11oWwoycZ696eM5y3Tvz1pirnGe3U9KkCAtjP4UNg3qIGA6Z3H9KRm/dMB3FISwJB70pB8o55G09KEHmewYJ6UAnOB34zQASePSnInzYP4mvl2m2MYVIGB1p2cKBg5pzDaM9aQIwUuTyfWtFdPQVxysR9aMYbceRSrjy8fxU0khCAOPU1pbuJjtpYMwGBTQny7m69qUSERe1PPCAv07VMopO4DCgIBPXFREbVGe/apM8MSfpTHxsBPWob6jK7MVJA9KjQgBs/nUrAtg0xowsRP50RaewwT5up4oYjNJGfkyKAv/6qmTGVycuacvXFKw2/MKOnI60g6Cg/NjtSgYHNMyOTTfNp6sZJI4jQkVArNLjsKGkG0g0+KQMvA5qkrICYRqEHFSBMgACmqcU+NuvFQ229RCs2wA9qhWXeM1JIMg1HGONuKtJu4XJYvmjOeaDDuUDPFPXCLjuOtM83Ax3zSQDVt9rnPNSwQIrZqWIAR7mPPahGG7dzjNW/MQFduQDQMKAachAJYjioy+WOBwDU6WEOfDydMY61HjMgAqUkMQR1qMkA570NX1GNcbn+lRdDipT2461GQAc9qze40M7YpQOcCjHP1pSMH2pdBie3pRnLkdqAcHJpOd3pmncQp6+tKRQOKUjIpdQI84b2pWwO9H8RFKeTjH1qnsAA4X2qHOGODxTyc8VGflfHanB6g0OVsmpM9jioMgOPc1KTRILBuIIoHD49aARmmO+BuzQr2AkY8gVFI+HwemKjSYEMT1qNX8zr1FXGD3GPDEvgHNOhOJjxzVeAESs2acZSsnArTld9BbIuOfnIJ5qJ3X7pqm0rtc/LzipD80oPen7BqzbC4koMmVpiqwG3rVokCTFOWMCQ0ubl9AsQW6OZsmr3AwDUYAVjihmAGe5qKlRzGO4zz0oRBvyaaj7gB3p4HzdeKjqIUjLYpRgEilTG40wH5j6ZpasB6HaeaazHdnHFOfBXjrQwG0Y/WhJgxhb5van5xRjPAHIoxuyPQUdBjsjZnvTR93j8aRhlcelJwqkCle6EPA3DjimdcqT0p+SuB61E3BzihbgKPmHvSFvmxxTkzg4quWHmYP4VaV3YCQ8UFsAcUuc8CmSZDUkgJG6DmlJ4WozwAO1SEcA0mrMYhbDcU7gJnvTQuXzmnNycdcUCaEf7vAqJGJJU1KxzxTMYOP1px3GVZsZA75qccAD2pqopck9qce2Olav4bMSFOOvtSHHHvQOVb0pigg89qiwxQSG2noalI+Woy3Kk04kZxmiTAORzSjpkflRnnFP4CH1qQGMKXGBSdRweacBheaQBtwRg4qFwMnnpUx6AioycqfWnEGR9hxSp1zSgfLnNGMJVBsOX170pPcUxWz/WpG4xxwahgCjg+tKrAdeppCSFwKaGycd6eu4EpIqJxk8U7k803APJoSYgYnK9xSk8CjGKaxAb0o0Yx+RtBP5VGxOeafwVHNNcZHFPdhYYpKufSng8VCGJbB/OpQOuOlNp2AkzwKackHHelPA4pjfKBmpTAjAOTmmyg7uOakxh+Kc3JORVp9REY3bKYr7TUuCQRigqMUJjIwx3dOKXGGPIzUgXANQk8/ShaiGvksCBTuMA0HJXGahZypAq1roMsnGfpQfmPSmpk4J7Cl6HP6Una+gbjMYPPWhUHWjndkCng5OCMUmMUdKRvu0p69abIfl60hETkAGq0mAwPrT5GyuKY4wgJ61tFWEyOVcENQw3LSMzMnI5pd+1eTWyug0IJVI4oRPlx71K5G7dSrjGauL01ELFhB05pdxyc8U75SRRIMjOKnmGPR+eacfu5zUSDjPepWBwRWUlqBC7hVGfxqFiGwBT3BL4xxUT4Q471pBWEQyplgc1PbsoAweaiHJJNNt8q5NdDV00LqaW0MpPcUwc8+lPjfPUdqcR8pxXG9Chok5GO9PH6mmxx4AzxQseDmpdrgxk4yME0yElTjHFSTnLinKynOKpN8obDGfDYqKQ5OKV1Bb3pzL8vNNe7ZgVgxTIB5pXjDKDio8/Mc9atGQeWOK2ldWsTuZ4G1/Y0m0s3H6VOB8wOPxp3CMcDk9K09oybD4E2cmppJQi88U2M5NJNCWXArCTvK8itjPnYznA4qVI1C5x9KsC0Ax3NOaLYvI571ftEtEBX3CPvQGLE4pGRmbBp6RgZJNU9hq7HxKepqdZATjvUSAHoamjHBxWMrdRkhbj3o60h+6TTYySDmsxi8YwaYw4pZATxR1jwapaiK5TA9BTUGDjkinuD171Gpw1ap6AStwQKVD0zzQy560oUjmo6APJHbrTgc803kY4pVGBUsFqOHy9Kf2Ipme9KGznFSMcB60gOBSjOM9qaenWkIUnIpuetN3dqd1707D3Gk7eaYrfMaewzmhRhelVdWEOBIH9aaTjml3Z4NRvkGhLUBwbJ5px6cdai3/MMdKkyCeDTdxic4Ge9RMMHg1NwD/KoJM5/lQtwZKjflSO/FRh/WmSbivFVy6iJQ2DUhPqaqLuByeatLyBRJJAI1LjrQ3BznikJxj1pdAGsM8UEelOLZHoaXtxSuwI2Xn3obAHWnHAHsaYcZ9/Sq3AYQC1IGYECnEDJoCkmqug2JhwBTHGQc1J0pnGcms1uHkQmMUgiCg+tTAc5/WggZ96vmdgKxjy2WPWo5IyT04q233qXAB4q1JoVivHCEFPK4GeR6VIPvH0pWAJ6n6VPMxkMb/MRUh5FGwDpxTgCBQ2AgXtmjCr0xQxqPOT/OlZgPJ5qPfg/WpCBg1Ey5HH51SQAJeeBTRkMST1p6hRzjmkOKegEofNPU5PXFQr0qZemah2Bj+poIyM0g680mepNRYCMsRn1o3VHM2w4pkZLtmtUtLhcnUgc9aOwzS+WCRmmuORg0t2BIpA6c1IvHWq8bEnJ/CpVNS1YESMRk+lN3cdaGBJxULKcjBpKwDieSKAQcYFIeOMdaFUnp0pgAPPWnD0pQAKFA6ilcAB5JPWnLzSYwDTlwAcmgCJ170hO2ns3JqCTkHFNagQ+YDIcnGKDPGO9Z87nzW9qgaQkcGupUrk8xrC4R8GpAFZeCKxFYg9cH0q/BI2zJPWidJLYL6k0nJ20gHUZqEuzN1+tKCVAPrS5bIZLnCnA5p6NhenNQo2QOPwqQ8cdqT7CvclJ4zmmhwffNRn72RTvukjH61Nh9QYnn+dMXI96Qgk+1Cg4JzzVW0EKxO4UjHv0o68im7dx+lNDHhhinBjjBpgAIz2xQScEeneiwEwfnrR5g4zVZXzjnr1pdwxU8obrQsbgpNISMetQd+eaFckc0+UVybOMY60ZC4xUJY7h7/rSl+gxRyjJiSScU0Y6Co9+ec596QOQOtFhk2fQ9KUvjNQl+T6UFwOfSiwrkwbJy1LuGOelQBsDnim+YMYzS5RFncCOuKM+9V/M4H60CTHIPFHKPYsFu3amA7h71D5gYkE0m/B4p8gixux160u7nkVXEhPHOO1LvOcdaOUe7LAPelBABP6VWDn8P5Uoc8g9aXKMmdsc5HvTgcDPSq7Pz9KPM+WjlEWGYg8/lTVPzZqEyZHHegTYAIxzRygy1nnilOAKrrKD7Gl8z5uKnlY0Skck1HkYppkBXOeveojJz/WqUWBKSBjntSFwRkd6geUYqt9rH41ooNivYtls9fx561HnGf1quLjzGwvBPFPRSwB71fJbcEx5J3Y9KkUcH0pChPANPHGfWpbsIcoIzzVK4P737vWrTuFOM+9VHkD5yOtVBO9wehEPlwQSPSp0kJO08n0qLYADgbsD6UqqflPP0rV2YuaxZzg5zxUq8rx+VUn3DIXp6VYg+4PX0rGUdLjJWOAM9RUbZPX8qkJyD7daaFBqUHkMb73oaCDjaBk0pGBycelIOc+nYVVxMYVA9/WmLhevFDlTnn6UwNlvUVaQXXUmVhjdjrSeeR1FMWM9h9KaYnxjt/KiyFclMxyB2J6Un2gMNuKabTAOOfemGFt3yjn60JR6DsL5mTkipB8wPP1FMWNmIBH1phDAEYPPenZBexZCgnNPCdgOKrLvU4IxmniQqOTyKhxYaXJNpUnj8BSHABJpN7EE9fcVHmQpyPzoSHew9jk5x1pp5APQimgM44H1pRG+TgYqrJC3HqSSD1NG5u3U96QcYI/KkYk8ev6UgH+YS4GMU7BDZHORUe3K4JqVWz0yPWkwWomecd6qzgg57VOJRuP6Cobg7zxz9aqKsw3IwSV64ApFG45x0HNPSLJz0q1HbDjPWqc0iuo2GPAOR0q/GNoqILtPSpVPWuecrgh/BXrio2wM4/KnnnOBxUZGQeKhCI8HOcdKfz170mc5447Um7JAzzV7gKT6dKoXbFec4Bq9t4rOunDenFaU9w3KrHg+lNOM5XrjmkIyMAdR1NGOOeorqESc7evBNPUnzPcdajbnOPTOadGSR15Pek0N26D2b58YGaYzfuiM9QaUqxbGDx1NRsAEY559B2oVhXR7GS4bpUgViMg4ojwCd3em7iMgYxXzj/MY4vhQMZxUbTnOCOvQVIMKrHHzVAv7yUOw5AqorWwicnLZXjtRKSgKjpjmm5Ltx1FKNzH1NTddBiKp2jJAp0yhdqk5JHrSRKzAjsKVFyrOxyRSlytABVRGF5LUkg2RDI60bsjcByKaQZVyx6HpTukhETAlAc4qGVSYmBPFTkFhgc4NV5iTgDpms9ShiNtUAGpFzyfaoZCFwv608yqPl9qppsYxyfwprH5c0Bg7FfSmuSAQBxU2Y+g0ZYjNNc7c0FyBxTdrSZq43QETucGrVsCFye9IIFUKTUwHp0q5SSVkA8cH2qVSAcCoieQKeh2kVgtdAsOJIbnrSqu3nBz2NLg7SxpoJPNPbQQMNzUeUpb6U9cK3zCgsFBI796OmgC7ucHkYpWk4A20yM96eASR70K+7ESbiUx2FLtAjyTzSIuTtzwKRvvbSehqk7L1BkTFsdKMM39adI+OAPxoDHZjpmnEBnfOelN6rj3qTgDB7imAYBrJ7jE7baTq1LyOT1oUAHJ5FC3AZjnFAyWp+NpPf0pEGc96QDQpDn0pxJ9OlOXqcDNDdRgUdAIVZg/zVIemfWgrmkIwAD609wQ0jrUDEAjPWpzwx9KryDLZqoLUNBrnvUgcMgNMcjbTcHywM1ppazAlamSfcNRs5XaOaVmJOD+VJR1HuRJ83IFSogVMipIUAXpxSg4JB6U3K4irbggvzyOakEe/wCbvTQwQsR3qRSQOD1q3K4kJ5QE24DtzTNu2bOfwqUtyPWhlBYHHPakpa7jGKpMxJp8e5pSCeAKcvOQevrSRsFbGfxovcALEycdqjLEy47d6MgzHBqZYwSc9cULRANjwp61YT5j7VVIwcetWF+UY71E9dbDH5wQB61G2Q+B3NPJBIpp4cE1EbgKDzinZwPrTB/rOelEjYbA6UMQ4HapPenI3BJ7imDOMnvQCQpHQU1a4Dh35pqjc+TQpxyelCNlumMGkkrAOf7w9qRhn+tByzfSndQSenpR1AjDYJB6VGFB5p4Xk5/CoIyRIyk8Grirhcl3E4pWIJqJWDOUPX1pW+VwvXNJxfUY84PGc8cU5XB+U+tRj73vRkB6bs1YCbo4xSjjNN6DPegfdz3qVsArcqPWmldozSg556UpHy81KdhEGNrexp68qRUL7g49KnTrg+la2uroBn3Rgd6Q/cpTy9NYnOD0FJXsMaPmOO4py857UYAINH3cEZNLcBd2SD0NOTnk9KaSBTgAFI9aHawArHfjtT2OTTV4OacDnnNKwCHJHpTAM8ZpC+TjP0pyLlu1Fna4CKuODTuCp9KXqD6UgICkUbiEwOgpzHC8ikVuc96GJJ4pDEHGaQx8gjrS+gNKh3HHpVLyAOi5PbtUW/B7YqU9wfwqrKccCqSETeYOKSRx0NV1Y96cxyBVOFh9CcNhRQxwpB71HGTkA05zk1FrMCIjDYPQ1OvAIzUKjDc5qUZyabYaEuMrULZOT1x2qQk7fao84WkrALGM8088n3qOMlQSaXdz160WuAAkZpWztz1NB5ekkyDjtQAhz61GynORUxxspCMjOeKaAiYEAU0oMg4zUnUe1JnJouIeeFzTQAR1pM5JFCDOTRYYoGKRe9K2M9KRepo6DDOfqKjb5+M07t/Wk4zTRJA64470jr8owKdK2TjvSA5JrXW1wK0hwyrio3+9jHNTTHGarrhnDVtF3VxDhgrRv42g8051VV4OKSJRvB7VSd9QJV5QDvUoGOM0DaGwKcVzWUpJjGID0obgY96cvAPFNc85xUptsBjAetV5RkgnpUmTgjt6USD5K2g7ahuVHUnlfpU1vEdpyeaiJKjikNyQu0cVu03HQnQuxk7j6Cp0clTVO2lO3nrip4TuOK5qkWtytycHjmjI2j1oA+Wmrhj9KysgGTDAPrRCwYEDjFJO2O9Rwt2FWl7oLckZTn296HYBOSamKhgTUDAgYqVroDKTFep6mpF+7z+FRSxl345FTIPlIIrpdktCRME9uKXaTyePSkeTC9aVX3L1wRS8wJbfkkkVMRluTUULAEipurZ6jvWUviGhdnPFMmX1PQVJnkjFQzKxJ5qFe4ysXAkx1FLIAwKgVWd9knPUVOj7sHvXRy2Vw3GoxRvXirEOc9e9RhDvzUinAJHFTJgWD0IzTcFcUqHcuaXsB3rGzBgRUbAcY/OpcVGw79hQmBCTg1Ew+birDL+dV2UluOtaxsIl3ZA9qeB6d6rgN0FWVBIpS0Ghw5HvSc4waUdPejFQIACAe9GML6ClOB9aTdwDSGL0GOtROQp68Gpc5GarS7RmqirsBxZcDJqQYxweKqRqH71aUYHXNVJJAOz8vYikHSgdc0magB4PHtSHApvVcd+9BI2Z/SnbqA3jIxSsM9D+NNz7UzeOmKuwEpcCkYfLzxmos/MKk3g0mrAMRRnpTu/Sg4z6CnEY70XAZg8mnAgUZ4xSj0P4UNgG7tnNMPJ9xTmGKMe9AAM45Gaf0zTScmg81IDGwaQL0J4pxAP1o+vSrT0CwYB5pEHXFL1PPagcLihXAVic9eaaWGcUhPHtTGUHFCS6gSg+h4oJAXrTF6UFuOeKLAhMZ5J6daA3Xig9PrTW+72NUlcQ4PycU4nnPaoI+GxUvJyKTQ+gEgj60o6HNJ2pcYHXmgQwnJwOtKhB570gXA60YJGRTGP6jApmP0pwGBikpLyAbilKetOYgDFNJ4FNXAQDpT0fnvxTMg09SAaGA4N8wH60rkA4pCwzjtTGYVNtQIJ+V6896kgXgUxdrNUxYAYrRvSyESsRgUw4POOlR7896UNkYJ69qi1hseqgc0ivzTS2BScEYxTsDJ9wJpD81Rj5RjvTg+B71Ng3EYZPtTx04pgbr60/IxxQwFzk0p9h+NIDS59OKQCMOaaxxTz1qNuvpQgEZsdKifleuKWTOOKjIyM/pWiQGbeIVkzjjvUAztJ7+1XZwGbnnHaqoQjJPP8AWuyL90h7jQDnC9vWp1yVABz7U0Ljrx70pPQ8ZoeuwOV3YkU5fJqbAwPWqwO5snoKtxComC7EijBFSKoIpgUr+NPz6H61gyiMjH096YSQMnp6VITxwKj3EjBqkCAnIpMbR1pVAH1pCw9qYCE4JHSgcUnc559KR8ZBH/6qaAB3oHAI9aco9e9MduB6U92K1hrDJ56etN7YHapCuOSc0gXnkVQW1E3E0pyp64pc4IHSlPPH6VIDenpTc5wPSlIJ4xTeg4P6U0HUdkgjHOKQnjikzwAKQkY/rTsN7XDJB5Oc9/WnbvfimYP5UEMBu4I70WTC49Tke9N6scUYJx+lM5C85A+lOwrkoPBxTcZXOaaQSQD0puW6ZosCHjOeT+dKe3NRjJHPal3HuMinYHsSA/lS5xxnk9KjP3duKepyue1S0AbthPPFIxwpIpzcrwaY4Jzj8TihA7kBmYEds0hmJ5zj2oYYJ3DBpBbhuT0/nWtl1J1JFcsRg4pWkwOG6VHuCDA4B71CTk8ceuKFG5V7Fn7SQ3PA7Uv2nC4qrgBTUm0MwA798UcqC4vnOWxnikLEoQeSKeYvQnBPX0p/lfKBnA9KLpbCd9iAkkY7U0IxySPpVoIA3HSmkNuHpQpBa+g2GAqTmrf3celV1O0cf/qpwYnHrUu7dxpjy/aoRMVHPXpintnBz1pgIxkmiNhiZMn+NOEQBx+tOUseR0qXp/PNJytsTa5GIcHOKcYuP507PHNIx+UZqbsegzZtAx0pm8xnmpVbseh6UrKpTNV6iQwXAwMn5jSiZM4HOaqSQ45B+opgAByOoq/ZpoXNqXvMUikLj16dqq7jkHoBQZSQSBjPelyDurXJhHu5I+lSKgUkdPQVAu4tjHUVKHO3nilK4nox4wvTpSb8ZJH09qRXyTk8fSkkwykg/SlYb0RJv3qcDj0pFAyTjp+tIoOTilUnHXoOtHoPsSgBSMUbck+9MBP8XfrShiF61I9FoI23nPaoQVd8461K4LdO3WmCFlA5zTVrC6gTgew6VEGJcKx+tSFWJ9+1IsGCCKpWQDwybCScA88VGJlwR3FOaE5znjvxTBbcZoXKLUQyBuAec0wnOcc/Wni1O3G7jNAtcsTn8KpOIN9xpJyBk8dKUZxnHA6VMYcckcHtTxHg47VLkugiBYjg571L5Izk9amGABxS/X8ahzbHbQaiKvOBz2p4OOeuaQ+xoQipAexAFCOSaRh8o74pkfUc80raFE5zjIpvQgdvWnZzQwOMVKDqRnqOeKYCM8YpzZANRuMjjtVokcS2Tisu4JYgdBntWiTiM84NZc7Dk9Se9bUlqD1IzgMQBnjpQoK+nNNwCSwJ6dKQA55roBEqnjg96njXLHIGF9arKQCBgHnqasQgEY/A1MtCuug8gjAHIqKY4jI9jVjG4e3pUEwARyCc49KmO4n8R68x5AU8+tAwDknmmOQicEZquJC0uCelfPK99QLTZYnnApR8pCrzkdaVWO3A9etM3eVnjNPmV7MCQMd+1efU0uQAQRz7U1GCYJHJpykNJuPX0p3BjRuyuchT1p4QM5GflA7Uksm8gdMUiybFIHU0rJSsAuMhggqHJAIPSnOzJHnpUCyZU1FrvQBytjIHNNMWUye5p4+VWIxk0hJ8vn16VOysyildKFI2noKqxMztkVemUFMY5qDy/KGQODW0ZpbCHIojBJ60AbsnNPA+X60fwn1rJ3KI9mTinqm08UA+tGcHNF9AI5ULMBnipFUjvQeWxTjxQ20gFK/N9Kdzx+lJnApw+7UiHYLEdcU5cbsjoKQEhcd6avy5JqkwFlbaC1RbnaLOKUqz89qsKB5ZFVG1ncCOJWCVYUYXJPFNQEqEHX1pSCCVPYU3MQBwoIqPDMd1IAfN5p+TnHSpi9AFx8o4+agqQQD6ULudueMU4/M2W7UbgMIHC96RhhCBSngknp2puMjg4FS2AzGPlNKRgr6U9RlifagDL5PSlqwEkUjGKRBtBYCnYYvjPHagKwcA1QCgYHPf0pAmc8j8ac4XI25zQR8uO/XNS3b5ARMCFOetVyGk5q464xzUbKFzimtBkPRDnrUEhO3INWMZJ9qq4IJz3qoXumApOVGKRTkH2pARjr1pYyNxFbSSY1cajB2IPahxh80Iw3kDrSt81Z7MQ+JwwwOlOP3qrwt8xAqc5zSkrPQZEQHYqRx2p6qCwwelOOA3tRkUXAGADdOaccMRSL8zAU48NSuAhG4jA5qNwEbvmncKc1EH8yTkcVUU2Ihc7JMjJ9auKRgGqk4YSALzmplJHBPFaS1iCJHXJB9KkA+YemKRgAooDgAVi77DHdXpjgs+PSng7cn1phyMtSQiTII+lNYbjmmrwM07kimxjWO3HpSs+4cU4rwAahLfN7CmImzlQooYgYPU9xUTMQwYDrUoG7mkwFx3oJyOKO3tSY4yeh61O7AQHj3qrcEhht6mp2BUnA6ioh86Et1rSN07gxIwA+4mnA7pOfwqEvtTdjjNRy3JjRT3JrVxlLYV0XGzvBHpUYBMnJ4pyuDGG705BhSc1nZq5Q5u3tT+gHHWoVOQfapQQQPaoYbi9DxStyuO1IeBmgHFTcQ04YD1FBPNIw2gmkHzKau2gwJB6daawzx60iEjOaUdM0LRhuIQMY70qnjmhgQu6gYNHQBwAPNAXIpBxxjmnZOQKT30EO6LjuaRx8mAaD1zSt0oGVApVzkcGp1O0H0pH5HvTMnO01d3JWAmPHOaQkAHFNkyvHU1C0vy4FSo3DYmjwRmnk5+tU4rnAwVwfrVlT8ue9VKDWoA570seAMjrTCMgZ6UqjCjHFJLQBWJJ60wqDnjNOcGmYAGe5oT7gMIAIwKVgFpF5Y88U5sd6tgKtI/tQmMdaVsmpaswEQ5O2pOnBqOIYPv71IFzn1oe4Ax4AppNPPQ0w/cPOKSuxbCheKZtUv6YpQTimsKFoxkg5YGnMMnJpiDHentyKLgxg5yKU8Ke9C0h4BzQAg5Xg0zIUc04A7TzUZwQR3qkr7ACsOTSpzu9KjjXGeeKfFnd1rTSwWHBsHI6UoIC+9IR70jHAz37VmMVuFNRBs5oV93y4oRdpJqlZbiGMhBzRkdqQvk4zzTec/yqkr7gUbl2EmOo70+MYC065ABB70kcZGCT+FdCa5SSScAj1plvw3PU+lPMgOQBUMEmZTmlG9mh9SYuRMeO1SqxznrTNu0kmljGGpOzErjycEmkJP45p0gGc9jSYAOe1Z6DsRFsA1G2SM549KlZATn0pO+OvrVqyAqnhTxUYCtk45qSRcrtqGL5W5Fbx1RJYgGD04qymA/B4pm0FOBilWIh8g8VhJ3K2LJJ7DtzTVODSA44pUyWrPRB1GSqCn9apRZQn0zWiy/Ic9MVB5a1cJWQnuOy3UGklY9Keq4HNKwHelfUZSJIOach5PFK/EhPpSomGyOma2vpqSMkjLDgUwxkcCr2wEdKYkYaQ/pSjVsgGohUc9asR8HFKIxuwacq4bArFzuMXaNxNQTvtBA9KsE4OO9QSrkdKiL1GZbLul9M04AhuuAOtSzDOCOtRQxlmJJrsTuriLAJJHpTmjJBFAOw4pxPB71kxksO0AAU8Dk1DFwB71P2rOa1GCng8UjLipFHFD+3XFR1EVmUkn+VRlRu5NStj14qNs1ogDGRyKcGx+NKoOOe9IQM5x1oDYUHtSAgZpOCOKaQQcYoQDmwc8/jScUuMgYpOowRTAUnjk/Sq8wyQan61FKqkGqjowGxMuCKlz6GqwTJzUyfKeKckg6EoOTz3oI7UBlGMim7hmpAewGKjb9KDjPFIeee9MAP3utJgY5pRx9KMfNxQBGAck08Djj9aD8vIFBPB5p3uAvHJxTCwPHNIWPSmjJHTr0ppAHmHJx2p285xUYU56c0pwDVNIGSM+D9aUn0qMc9evpS84FJoCUcCnZ+U4/KohnbzTx1qWAZxmm59KGBI/nTMkcUIOo7d+dNZ6YW4461F5gJNWohsWA249aU9BVbzQGGOnSpg4YUNAPBI6ml4xmmE4PrSls9TSYD8cUwjr/ACppYdjSM5xweaEgHhdpzS/hmoTJk00ygHGeaaTBljOeP50o6dOKriTnNL5nvRysCXOOtIXFVvMOelG/gk0+ULlkNz1xTS3aoPMDE85pBLzT5QuTl6bu5qMsSM0m/APrTsGiJsgZpDJ09DUO4AGkJyT/ADo5RPYnEmRmmsx4HU1FnqMU1iehzTUQuTJJjtSNIep/KoeT0NBwWyTzT5UDehKH464p3mHGaiHP4frS8jr0pNICQuRSq/BqMj0oPA4zmiwE+/NI0mag3HAzTC5I4PFJRC5PHKCfWpg5NUlGGAJqYFgKJRQXLYfmlVie9VdxVjk5pwfDe1RyjLJbJprMMAVGXAz60zeeSOlJRAkNROcLxUmSwxVaTKk/yq4K7E9iFjg/yqJsHOOop7kZ5PHeogCXPpXSloS7CluR6UD5kGOlBAHeml8gAcetMA5AwT7Vbt27ZqkFJOe9WIhs5JqZrQFcvNwMGjkdTVWS5weDwBTDdHArHkdirlljyaj255I5pkbFzk96nxtbFP4dA3GlcZ9ajbPQipTx1zULtgFe4oiFxS3zYA5pWAOPWoiCcevrS78qO2O9XYROcYIHXvTSAByMmlDALknjvQuH+9UbDGkcg9qXb1x3p4G+nhRtIApOQESoQQcfhSlMf4GpNpwBSEccmi4mQkHuPypBwelOJZmyelOcADLVQFWUhBk/n6U2N9zYIzSyKZTkURr5eQQa10sFyRVB7ZFSiJQMcZoR1wOKcCCCc8fSsm2NLQYU7YpApIqXdkD1ppPBPrSuwIyoPGM00ovpipBnBpnGapXFZDCoY4pAAenXtTgQG5pTgMWwKrUQ3bxxTQMr9KUNgHPJoDZXGM0agIrHbxxS5HLD8qBGzDjinLAS3J9zT0HuRfJuJx9aUxsRheKsLEFb1p42gnip5wsUvspbGeQP1o+zbVI6VfXFIy547Ue0YOJRFsA3OSG/SpdgBA71YK88800oFP0o529wSIDHnAPSnBCAOPxp5zuGOKaRzyfwouHXQQgdOh9KbsY9albgc9qMAHpx/Ki4EYi546ULCCRkcipeN3Wk9yaLseiGGIHJ596asajHGal5z9KE6g9KLsXUciKF/wA8UoRTkU4j5utGMAiob6jE8sYJpjKpJ71Nnpjt2pu0A9KSYaFBg2eB0qMzGMEcnHQ1deMc1WlTjrXRGSZOy1I0cvwc/Wm7WIwSQaVRgnHTvUi+4x9Kpuwk7jRFkAEcCpUgAx0P9KAOg/nUsdRKTK0YbACTj8aYdqnHpUjHjrVZuH64/rSWoEgC7jxSgYwO1JG2Rg+nWnLySO/rTfYW4FOcDpQB1wOKeR39aEABwc/T1qblWGZzk54pc5GSKe0Y+WneWBzSug1G4x3pCCeOx96k6DnpTcHAz19KVxeQzZxz1p3AxxSjrxSYPfqe9A9noDckkjpS57bfwpeCppF470C2G4Ofel/hwetKBgH3pD9KLjdhDjGKMkD+tHVevFA55NMSFBz9aXvwKaASRzUgTgHPWkxrYQA7Rnk05QB2p20f/WFGOTU3EhOSPeo8AHpg0884PXFJtBOcZx3poLkoI60Hr0oUfL9KTtUDYzIOTmoju38HmpiBu9feomUEnB/CrRPQZJjaRjNZU2N5J6E9K1peh4zWPMpMnBB/CuiiMYPlOM8d8UhG2TPDUbcA+vqaXgJggk10B5jlyxIOR6mrMcRbsAD2qvGpHOTuzwKtxYBwDzis5vsFhwX58ZIwfzqKcHyWJHarBByWHBqObPkyZHG0/wAqzi9RbM9VkT5WHfrVZYm8zdVskDJPSmptfJHrxXhR32GTAhYjnr2qosbPJycDqasMBu5OcU1zj2zSb1uAjttYBaeM/e6D1qNAOvanqSSQB26UWsA2XCgtz7UyGXzASB9KcVLsVJ4FEY8snaKUZILCOrSsATgCmIu3cB09alO7IJ4z1prDLY6Dv70k7DG5IGBzRtwmSeaNu0jA+XNK3GDnNS1oAwjKHjvTGTctSHjqcCk2/KcA4qbPoMhYcY9KaeBTmOeF6Uh4XpT1QCDG00nrS9qAcUxoaB789qcD+Ipp607g8CgBynJ9qeAC/HSm/wAWCaUDJ44qQHZLN7U3cGkwfWlA3Hg9KYRhgR1FXG3URKW5wBT1bCkYqMk9cU9eRg0WXcCaIEgmm5yeevrSBix2r09qdt+XrwOtCa6CGrkfNinE4HPU9KD93A9aaQSRgUX7APXhM560MAFPrim4P3R0xRnD5Y09L6ADKRGM+lJx5XvSO5bgdKOdwANS9GMcBlenFBwI8DrmkBOSAe1Ifu9eaHfcB8Yw6k9+tPYjeNvSoCRxg81KrKAMcmtYXtYTGgYdj6U9FGS/pSA7flx19acqYxzx3pKPkIY2GbOO1RhRuOTyKlkQoeDUbAe2c1OvUYwLySBVaRCckdKt85OOhpjY5FS3ZjKQh2g4NIqYYHPNTSNhsDFNJG8VrKUgRCUKy5oY/Lx1qww5NRBByKlTAbAuOfxqU5yTimRBtxBHFOJyMUTd2MY/TcOlMRiFyelLIDt47Ug5QjjNVFaCuTR569KXOck00/d20q8LzU2GRuxCk9qBhVDHinqoKEHpTDg8dqpPsIcy87sfhUZ6g+tTgZXbUUij5SB0ovqA9z8q/So8kyqKH9TSRks2SKdrgWSTgY7Uh547UsfJOKMAZrLyDcUn5SBS5OVxTQcrjvRnGM1T0AU5zzUbDrRIx3UY+bg1SVldgOTDJTlPBFN24UYI680uORiplZAO5C5NIxxEfWgk5AprjJPepQEO4k8mkjUiQ56GkLEnp0qVV3Dd0rRtrceg2SNXjKAVD5XyFXHParUQwWLVVuNxkBUcCqhJ3tckWAnkHtUhJIIAqCGQEMOntSwypK7LnmqlB8w76EsbYIBqwi847VXQHfg9ulSxsQ5waykkBNgA4P5UwjJpxbmmMcHIqFsMGGRTEXg54qTGaQ8JgdadxEQ5YjPAoAycHgGjaQM0EYx61QwYnGKBnPFLxj3pVHJoa0AcMZOetKMbc01TzzxSjGanYBQOM0gBKml6NwOKUcLxyKQiLaevf3o2A8nrUrcL25prfd4qrsYxsEe9QSD5eKlYfnTRnByOlUmDKxOVPHPep0bCDmoJM4JA5qaIYQZBNaP4QHlyG4pxJKgCmjG7+VOwScis2A/BximHpjFOJyBSEgdqlXYESoMe9KQMYOafgEE0gUEVSeoxikAVKKjK/lTl5HWhiAJ82cU4cDrTQTnrTuMUtbXAdxsJ9qi5YcCpCPl4ppOxT60RAOAme4pjkkUL09jTmwVxim9GAiNhj709V4yagPysOan5Cc0SAQj3pjkjtTxyKQ80k9Q2GqQFNQSKcjHFTtz0pj8DOKuLsxMiYEEAVNGoBP51CzgkGnbsMMHrVb7jJioyaYfanMD1zSDvUJANABPSowcOeamVcVEygEmmndgQyL+8zjFKX46UkoO04NNiO6tUtBDSgc/SkOAME49KexO7ioJDlyKuKuA1wFBI6U2HBJNTMgKAGq6/u34PWtF2F1LZJdeacAwApgJaPj8qmi4Xnk9qzewxSM4NKQCDgdKU800kA9ai/UNhjrk4FREFDkdKm6jpTG7800CKs+cZHeqyOSSCKuSITnPNV/LCcDrXTSasJ7li3k4weuKsg4qnErBcmpy+FyOtZT12H0JiOOO9KmUHPWmAnAPWn7i4BrJgSckGoRwxFTgcCoyRnNJMAyMYpGywIFAyvT8aeBjmi9gKzpk5p0IJwTUjDIzimxKwXBqm9BdSUDjFCoEORSgcYPWl6DFQNDwBnPbtSAZP40cbfegexqAArljUMuQMHrU5Bz70xhnOTVLQChICecVXiLZq9IBtwOBUMaha3jJWDqI54yfyoBbBI604puzzmkj44NHQOpIsgHHepwRtzVQEbuanVgMDtiolEZOn3fwpS26okmUZwKa8wC4xUuDuIWQYPtTcDFQtM3So3mYD+taKLDdlrIAx70x5AAOKqiZsdKd5jY6U+R9Q0JQ/zelODA96gG8jOOKASDnHSjlQE270PNJubOOKYW9BSFgPxNFgJASeT19ajkHFO3hh1ph+brQgGB9o471JuOODzUTYz7UgALCrshFgPkHjrTScHBpA2eM005Bx1qUhj94zjtRk9RTT1570DgemKYAxIpQCO/WmdutPAyBk0NA+4biMAmmuT3PFK4+XIFV9+WximlcB5bvninoTj0qAnPWgy7QABV8twuiyuSen0pMHNQK55GTTvMPpxU8o+hLgDnNJ1FRgjd707dx1oaaETZ44PNGcA81XD7SPSkaTpnilygTbuPemZJOaYWpuR1qkgFOOSarOx3EA1O5DCoAgPXpWkNBNgOcHpU0ZYDNIP1HSnMcDFDdxdR244/pSlsiouQeDS+p9Kmwx5bsKYSc8HmgnPApeARk49qQNjcAH0zTdvGRyaecdcUpwOT0qhdCIZVh6UpbLH+VKfujNGCPcUwDGR6U3BIx6U/j8qX5cZ9aVxkZA49KMZB9qUsvII5o7HBpiAZJOR+FAxn3pCeDTg4IGPyoHsJtx1IFAAzn06UnQ9acuGzzwaHcVhOc9OKaRkjNS9vr3ox24pXCxEqkdqeF6E08jBDU09SM0XuMac5xnHtSjIzg01jj3p2fwNMewfw5alGeKbjGQaTNAuo7gnNIQB2oHJpRgjBoAaRtOc0Hk5PegAd+tAU9T0oAMdT1p2TjkcGmEnOe1KGOaAJA+ByKcrL1B+lQ5IyaUHA6fjS5QLKnJznrVa5O1SQePWnliDwagnbcvXGKILUHsVs9yeKQSAIG45/SmOxD4H1qMnkcZArpSJuSu+5+OMU0EYGaaDzgdaeI88/hT0QtibOFOeTTWbJ/+tQ3yihQeOnvUpWHfUf5W4e9EUB4yf/rVLCQRg/rVpU5J6VnKbWgWGRjaOlOx3oKHNOIAHX8KybKQxuRzx9ag28kjr3qwVDdagYbTjPBqoiaVhOD0qJwQBUuCoNQSEt04q4gx6tkYzxinxrmoYWGQT/8Aqqw8owMdaJXWgJ3Q+Nwuc1LvUkY5rOkkK4z0Pf0pY5WTB60On1DmRqDHQ/nTHwopqSb0ywwabkdSe9ZWKHIp69qV0DDp0o3AGgZORnijXcQyOMAZxSSRA5xUm4JkE00yDJHQU7u9wtoVj+7fHUU7fhsetJOMtnPTtUJcDvjNaJXEWfMBz3pC+Vqqsjf41Iu4lePWm4WC455AAT29KaWyBxzT/I3Z3Hj6VKsYXtSukGpVWN8jIx9KlETNwT+FWWC96avX3pc7YW6EawgHPWnBFB6cetSjk8igCo5mOwgHtSEdu1OHt+dNPJOPzpAROTkihGO7mkOCeacMYz69asTvclAwcUhHUg0m7v6UufwqbMqwA5b0HajGc5ozwKP4uKQCYz34+lMZN3Sn5PBHSgg4+gqtgZHgsaApBx+VSZyKAOeCaLi8yLoSM5pcDOak2j6imEY+7TuAwkZAFCY3ZA4pWIBpI1PAp9A1RKDwBTs88jmm43Dmnlcg4/CoC40d+MmlHPb8acB2xT1HXNJsZBKAEIxVN8qDjk1dmPy5z9aqOBjJPFa09iWyEpwSOvrSqWxgj9aY+RxjApdwIIPatiCTcAMfmamTgevtVeHMjAH8qvCIDJFZzdtCyF/pgGq0wAO6rUgI+9VR+Tgj9acGS1dEkbqVxnPtUisCxwAagQd8ZPapE27Tg4+tDQ7kw6YFGSDnPPWheO9GM4qCxd+R6U7dxx+HvUY4AP6U7jrSsK4oPoeaBwcE/jSe9HU0AO7jPag5HIOKaSDj+eKQeueR3oCw89P60hAOaB2OeacO1ABt70hBwBUqr3pGXqTU3AiKjGCPpTenGakLgc55qmXLPxVxTYnZE5OCTmpc4AzUCISQSM1MYyUI7+tJ2DoPVstwcj0pcfe9ajjUr1p4IBxjrUtAkKMdccd/egcnqaU+vf0pEOQOaQ7WFI6im5J60pFMbg+lCE9xT0wOtRMMEAdKfk8jHIpMY61S0Dca+dhA61jy7tzBhk9Qa2iuVOeayLhSHJP6d63ovUbIWBI54yMmmk7gevHrTmY7ugxSZ3En+L09K6RDkblR3PtVpOWzknA7VBHww6delXFA5GBtFZzYuoCXBwT071HcSZidcnIU9KldRtyDVadSisAw5BzURSbHqz1ibLZUdKZCGRgv8IqQ/ez2pyRsxJ7A14MX2DS4uCG9eKbKu7BBz7VIAOhoUYYgHPoaSs2A2Mfwk0gYhjj8zSvkv070i/M2KN/UByKxUnHWkDE/KB+NPMhT5ccVHuI6daW2wA+7cA1K5AXaByaYA33m5NP3bm3EdBRfT1GK4woGOKgkJyMDipQGkzk8YquxJbHYUp2bAcRkjuaWdSsXHQ01OGzUsill+Y0l8LY7FOJwoYdcU4HI+lNlQhPl496RflFP4tQBieQKZuI4qTGSfamEZPSpQxAcjmpEA61HyT04p6HAzTuA5Qc5Jp4ICHnmmZ4zTSd2KnViJlO1SevHembuM96cPmGDnAp3BGPyqkAiDgFvyqQkYz3ppxwP0pWHQCi2lw3JEYBTgde9OCnbnPHemDO3inbjwM4NN22ENBOKNx65pSCFApGGPp1qG7IYvJ56UZVVweSaTnFIDg4xVWfUBwACAjr6UinJyaUnngcUBucDpip62C4053HAxxxS4AFOJwvFIqjv2otrYQjY28ChVwQTT/vMMDgU4j96F7DvWieoCNy4yOlGfmxk4oYln2j17UDAO0jk1TkriCUAEAGozt/Kntkde1Rbcvms27jHH7mR1qM9M1JjDEUADBNGtxlCQ4kJpqcuafMMN7mm4KnOK2jrEB5zmkAAFP6jNMOR1rDyGKvWm4G7igMCaTv0qgGHnNRxKQzEngVLj5vamqAFPrVJgOPPPalJwmDQg4wTxQwx+FLyEOXBQkUzgpwaRDimSZjGRVJIZIud2e1TbQQaqxSbmx+dWM4O3NKXYCOQfLxUaqUUHvUrc9ajfdxTjJiJ4uMsDine/akH3fwo/h96h76jAY3ewpfvEntUZ7U+MjNNbCGSDjg0kfBH9aeo+b2zTGB83FNOwEoBC7qX370pyVANBHSoYAx4yKRhkUfxHPSlA+bPal5jI9gz7mlxyBS5w2ccUrcsWpsBMYJB71VEbKWycirZwQGqLaS2atPl1J6lJ02gFOppIofKl3E9fWrTcsCB0NRTKWlGOn1rWM30CxZRACc96UDbiomdhtAFPaRVAyazkrlWJHOFBP0pAc8Gmkhsc5pz4FQAqtg8mgH5jmmOCV96agdVyapJNAPY81GSKfk7T/OoA2Tg0oxDQlzn5hSg5pm8dBTweMDrTatoFxTkMAO1LnjApOrc9KOjGk7gLuOM9KcuNpNR5OCKcDhMetFhC9aaATxml5xxSY2gEUkA08uM0h+UkdKfwxzQRnrVX6DK7RkkkU4HnGKlPHT8qbtx1HNNSATA7d6Vchce/Wk70ZP4UdAHhgT71ECNxGabKH3fLUSbhL0+tUoqwFkdDmjftHIpxGQMUx14yalWAQcnpxT1GO9Mj708HJFEgBgM8dKF5570dWIpR1NFwFGTTGTcMUZxQGwaS0AQLtUCm7SWWnnnnmm47jpTvdgM2/NUm4npTe/oaRHAfHrTDoPHBxS54xSYO70FKQAwFIOghA4qNsNkE0/qxwOlQyZzimtwIyB060pIJ4pu3aeQaVGw9agT7sKAetKCQcnqaTO5uvNDKSc1m+wAHy2KbJ1pxUHp1pMcc0IRUdgwwKb908U9gBJ9KYy4k4HArZWAkIAPvVeVgH69ambO7pVa4BLkjtVxs2A5SGJBqNsZ5FNjb35p4PPPSrasA6M4j4NWIjgcjmoo1G6rKKBgYqKjQIRO+DSk0McHA/GjbnmsvMNSPhcUjYJBApWAPA7UjEBKtMBhHHrUErZ4FWMcVXcfxVcHrqDJouUOegHSmhWY8dKcrARmlST5toAxQ7sESL0AqRcjjtSBCGBHSpCuTxWLdguGckioyNrYNSnA570hG7mpARRxS5wOacMKAKRhknNAxPSjvSjA47ZoXg0MQvB5NB9qQE0A0gHjoSKMjNMDkcUvAPQ0rAOycmo36mn5JpCN3FNBcryKWBxVZ42ZTirjHbkZqIdetaxbDchVCqnJowc1ISBnmm/ePB/Oqu2Awoevenw9cNzTs4GareYVuABVK7AugL1FDYxkU0EbaUnArPVsCNwR0qu+/HAGKtbhiml1IzVqTQblLc+fu8VIHbA4qclB2prbCMCr579BaksTbh0qQoueBUURAyM4xVjcODWUtNh2ZCYwD71GyYJ9atHnmkZe5pcwWKewM2AcUvlkGpDgNwPxqQDBycVTloIpyRnOKQKVANXNo3Gm+WPWjnGiqc56U3k5q0Y+KhMYFWpICHcccU5XJPP40xl29qRSO4q7XC5KOnPWpVAxnvVcMQevFSo+ealrQNh8h4xjtVE8nODV58spx+FMEWeWFEXYCoclvahuRmrDx8fSoCoxjtVxdxbjFbB29fWn7uB603aB+FBwO1VpcFcU9Ov1pOQME0DAXjrR1GcfQUDuKcke2KaT707GOaaOTTEOyMg0jDI+npQcg8mmZwcUrAL9fzpcDHA6UZwvGM00MRkkU7MCRWz3zS/wnFQbsmjzcnHajlYE5I/GkJGKgaVc0GTBwOD70KLB2RMD07CnHnvVXzCQaUznnA+tPkbAsFsdOlIGHc9areYzZxTN7dzVez6XEXd4JINMdwBVX5hyTzR1O16FAOZE4mB/OlL4XrzUBUAZz06UhHGBj6VXLG43oTkr3/E0hfHApnHWn7Sf6mpshX6gXAHPamiXBxtxjr7VIMbcY6UEjb0o0QCbgSMj8aCfakPtSgkHBHvQkMbvbA46dqXzXHWngjGfSoJGy2MU1qK7RJ5x4460hlJJ9+lReZ6c+1NDAnI5zVKKAsLIrfeGKcWXPvVXPOaQEqvJ3Ypcq6CbtuXN+7I/WjHriqscpBOTUol557VLjYpMeWI47fzoHJHqKTvgAY9KXryKQvUcTgk+tRtLg8flSsCeATxVVy2/AGRTirg2WFk3sAOlSDHP86rwZGAasBhnAoa10C44D5aQLz6CgcdaUMG6Go1GKFxUEy7jjr61ZXAbHQUhVe3NCdmG5R8kkYAxR5BY8DHrWisYxyKmWNQMjp6VTrWElYzPsjEknj0qN12vg1s4DEVRuEBlAFKNRt2YNWK2Aen0pNu3A75qQKVJBHTpQBnJ6H3rTmFcWCPLkg//AFqtqSBUEQ2kVYzWU3qUtB2DjmmkYOacDkYpM8VmPRkBYsxXP6U1VG7k0kn7s5ByfpSRncS1a200JYr4CnHJqpJ846YzV5lD9qq3BUNsH4+1XBgyFW4yTzTgSeD270xTl8k5H0qUKGAJP4VoxboeEVlwfxppjw3yntS/dz+tKjA8gdKm7HZjw4GAKlLE4x0qv5YY5/SpAxRNo71DSYvJkmd3fmnb1AIBqsSwzjrSiMk5bvRyp9R3B5N/Gf8A61IpboSaf5YGQKkCqBx1NDkraArkGw/xcg+tIsKuTkdOlTOV5PpSgqOn5UczsLlBYVB4FSYVRnvTA+AT+VMZ8jios3uVsPd6RJCR0qNlLHmnxghR6VVkkK+pKMn2pFGDmnZA6/lRnHSsxh0570uOD6U3IJxTj93jigYHI78U1m49aUnkA8UjDgcUwIcbe2CfSlIHTNOK8ADg/Wolfr6etWtSdESAjHtS4+YGoDJtHBxTo5AydeB+tDTBE5PcUZ+am7h0NRtKobipSuVYlHf0p2arGUEc8UnngcnpVcrYJljOc/ypc5FRrKG6U/dnPGKloBT796QdKbJ04OKiV22jjmmo3QvImI9etKFIHT9aZHyBg9P0pztgDFJ9hDwfl607dtGKreZyMflRJLyMU3EZZDZBNAODioI3wKUMSx44qeUZK+CvtVeRMoTVhVJUEjOO1NmUbCD+VOLsxPVFFwOeOTULK+X67SMdas4IXjk+lQkEnHY9q6YsgltUIP171eBwM4/CorZMR+ppZJREcdTWE/ekWtrjJzzxVOTjrjAqaSdc8daqM24cn6+1aQixN2HLkE449KnjwDzzk+lV0zuIYdKsoOcYxVSBbE4ORQRxQo/OlB4x27ViUhOgNN4HGeRS5+X+dNA+bIHWgVxVzgZPWlI4HpRgk9O9OPHbr2oDoIevsaReTg/jQxxSx/Mc4o6DY8daeq8UnAyaduz0zn0qGJijGP8AGkk4HFLk88Uxhnr1pLcEV9hJ4HTvT0hABOOan4A6UNVczENA+X0JpB0FKDwQeKQ9eBSHuO46CkI56UvQEnvSgjOTQAxsAjj60o5GaGAbIxSgHHP40CtdjWOKaTkZJpWORgVHyRTSAN4BJpSwU88H2poK8ev8qUEE9MelVYLjmzggGsu7ILn8xWmTweKzLo4JGMAda1pbg7lVsL78dDQuM7Qacyqq5PXHFMBUHdgkV1boFsTIQWBbIye4qwj8bTj2qsrgtnB/Kp1Qfe5ANZy8yWiQEljg/Q+tQTfcbH90nmnnccrmo5fuSEjPFKK1H0uevY4xxkUpYgfL0NMycUKQo+tfP6WGTooK/MR060xRg47etIDngCpJEO0EkAUNrdCG/KmVY9qajbeQKXAZlGOR1pd21yAOKpNtgNAYkyEU1uRuPX2pdxIYHgUxAPmJ59KjyH0HKDxk8U3Az1yKdGuQSeajLYJUcDNKWtgFyzcA8VGUCkgnj1qUHaCM/iahKktyeKl+Y0PT5TkDgUEkgZPFA4bAPFK/UelJ3tboMjfBUYqCYkFVHWrDdMKKibgknHSmuoEXO3HegnCAZ5pBnqetNcHg0DFwTilXrntSZwlOxkYoAaXO0nFPUEoPekYDbwOlPH3R6U07CHkgfLUnHFQjlsnpUqcjnoKV1qDFyC2aCTu9hSA/MaTO9uOlJO9gFQ75vapWwGBHNVlJViB0qeMEp9Kt7CH9SCfXijOc+3ejI4yMgU3dn7vT3qFtuPQUD/8AUaAN2cdqaQfWkz8u3170JoQ4n5QAQc0DjrTPunA/CnDPUmhRYx54AwMim4PJ7U4AnpQFJbBPFPluwBBg5JpyAFvmPehVAyxPSmgbhuAwM00SPYgPhRn3p20hTxk03YDkr9aEZiue1Unr7wDmOVB21C3+sUjpUkjMq4I4qIZyKUpWdkCHzAADafrTBwcEcGnOP3ee/pSHJwfyolJ3GVplHmAjpUEhIyKnn69KpzAkGqp2YEydAaHz2pkAJiGe1S1M0ouyK3K6AqxqVjgZpjMFFLkFcUavUFsNU8c0IeSKRuKEGVprQB6/ep3HOetR9FJ70KpYbvzpNCHkALgUxgSRkVLikI3Gi4yoEIm3DtVwevembOcGpACflJpuV0CGycjNN2/KDTip3YoccEUlZaAIhy3tTs8ZzUCkq2KmI2j602rAIfWgHn3prHkYo3AMKSAfwGz+dLgh8+9JkE0ozvANDAlPGOlIOBn3prk5GOaU5xipYhxAK0jdAOmKToBQfmbNIY1s9TTsgrQTnI9KTORTEGCVIpD0wBTu2DSfxdaLjIjy/ANIRnnHNSH71IyZII/GquIjMZLKc9Kjlh3GrTA5FMYYX3qrtMBAAqin9V96YTuIBpQeCBS6j3Ag9e1OHKmmA8DmnE5I9KLaCGjkEH8qgb2FT4+bikZRnpTTsDIfumplyq5ppADdKC3am32HsSZ3YpQOcnrQCOPXFHfpUNgKAAaMc0Nzj0pe39aXkAw5BJHSj71LjnBNAGDTYCKBnFOIHejGDkU4DjJpXAZjim4LHJNPI4xTX/Wi4ERyGGOlPxzxQ3anAALn0q7sBgUr070KvzdKcOtGctildsBG44AppBK+1SY+bNDDtSTERlCB/OgkDpzT8ZGOhoxim2MZngGj3pW64oYdCKPMBjDK8d6XjHA6UA4pVUZJ7GmAoHY80AdfU0c56Uhzk4NKzAQpx71HsAbNTD5h9Kay9aE2AZ680oPHvQwx0peM8U+lwuJjDE1GVy+albk8d6bjGTQmAxkUtzTFQYJ71LgHrQRkU7hoCoOtB5GKVemKaDk0gEAwvuKYzECnH7/ApCMA8c1QFcsGHI5pseGzkUrEDJxxTYjgHHStlsIc3WomAB59anyCvPX0qnKxViMZ9KIq4bCrGoqKY4bavfipFzn69qbNGF+YGtloxPYkT5AD7c1aUZTd3qpF+8XrjAq0jYXbUVENDVYBqkyCuRUT4IIpR93FRuCFwpximkbhgCmg7Cc05TnJx0otbUCKQ8YqEvx0qeQc5PSqxPzdOBW0WmIlSMvx7VZihCEcU2A5XNWFBHNZSmxgg5p5wT0puecfrTsceprFjEZcnmkGc4p5GVFNzjFNCY7aOtIeB60/IK4zUZODjoKnqADrmjdSZBbPpTWfaOnNV5BccT8pHSkwRz3qAzUizk8dM1XKwLOMnrSs4HIqsZMD71M870o5WxlkS0pfvVMzqPrSefkZzT5GBI75OaYCO5phfPSoXlC9+fWtIx6CJnYdzSIw9aptLuPHSlEoUZzWvswvqXGk25qi7h7lTnBFKbrAwfxqk74YuKqFNkyaVjdVxiguPyrEF04PB4zTvt0gU54pfV3crmRqPIajV89v1qj9qJ47U8XSg4qvZSitUK6TLgJJ60Z56VXW5Q4APNSiUcc1HK09hkq5XnJqdXGPYVCrjFPwO1ZS1Al3847CpCdy1SDkvUwbjFS4hoKXAbkVKCGXNQ5BHNPDbRgUaWAeFyMkYpoHalL0mcnikAuAKjbB5xUh6YpCQRihAV2X071A0eDk1bKjGDUbLkYrSLEyEKM9KeI8MMCpAvGKdgBT60cw2OBBH9ajZ+cA0wTLkqeDUayK7DacihRDcJ32qRnBqkZ1Xg461Pd8AmshmZpSB+VdNKndEydmXTcqT9aDcoFz3qkFzgdxSbCoG09fatuSIlLQti7z0Gc0iXDHI6n0qADpz+VWI0AOe9DSXQauxwldgRTQ7LkZ6VIDheBn1FRuctjsam4ncUSPnGenUUhO1uuKauN2BxS4JOcZxT0C45Rx1596U52nnikG7pinAnPI5qXcYgHWhh8xPftSMeTnvSZzTSYugoXvmjHy8d6AxHfP4UnOM00u4dBxz2FMwCcntS5AbGetKVw1GwhuOCf0pBw/PenYycZ+tIBkkg02CFwcc9aCmFA9KM+p5o5JznOKSfcLJiYG3GOtIQH4PSlVv7o/OjacH+VGw+hIi4HA/CnbjyKbnbz79aTJ59e9J6jQ/JpuSRwOM9KUc556UDByM/U0kkDG4yMd6eRjk/lSt7UwNk9OKWr1BDWfnb2poBOfepSvGab8wHv6U00BGU5yRS+U+04/OpAPmwR26VIoHTvRzWFZFQgrxjikDDPIq2qfjTWjBzQpLqBVXLZ3DFKF6AHn0NWGgBwO3tR9nb14p8yCyK4kYMR6VIs2ByOakEDAnP500wk8459KfNF7jJQQffFBjA5quA0Z2nipYnY5zUNW2HdMeq5zgUoB6flS+3ehjnI71ACYx260o+UUL8v5UZBHSmNBu6HNSIATkjmojg/jSCUp70WvsK5cCggc07cB0PHeqReU8DilAckbjx3GKjk8wcieSYYwo+Y0sEIY7jUkcQA6dfWpwMDOKmTtsG4xrdWOaja046VaB5PFGQKyUmFjLlhdGGeg9qaXI6itR1BznmoWjU8YrVVL7hZlLzgOaUyAg9jUrxJmoxCox6VSaYakLlSuetR7goIFTSAcAVEflHStFYQwyHPy9expqx/Md5z9aeynccD603azMc/hWielgQ4qoXp1poGMAjA9fWnqCcelKRxn07UnLoG7uREckk9MUkYKv8vTtTnOflFCYB6nijoD7kygBeBQTnr0FLHyME9KQA5OagBhYge9NEzHr1qYoD1oEQ3HinzIepFGWINS7PkOCcU5FwcDtTwB0xipcgjsQbSTntQV596mK9/SkYHgihSFuNCHbknn0pNuMe/6VKemBTCcg560XHYaGwcdxTlIP1A6UwgEk45oU9eME07BfoTc44ppOBjvS844zk0mMA4qBiKcnn8RUow3WowOfelXIOAOO9DEOHLYPNDMBjIoPr6Um7jO7n6UgI5DgZqq0mG+UfXNSTsRx69qruF28jk+9bwjoKTsriyOT8oP41CrlTtHamkkNjNKrBj+VaqOgr3JlkJHX60jEg8cUIOQccU8jg96nZg9UM55H5mm4IALcY6Y7VIF5JxUqpuGCPzobsCWhGMqCT0qxGxwPelWI45pyRAGs5STDqDjcuOtRDntVraANv60gT8qhSsU9wQcdKjKEtk1YUDbg/lSkDBNRzBZWKwiI4zTRGW68VaPI5HSmdT0qlJgNSEDin7AGp45B9fSkxnOam49gHHemyKWUUuCaBnFAtiExnb/ADNRPEFYHHHerm0+tNZMfQ1SmJoRGxHVObLnirjKBkjpVOQfN06GqhuDIjCAORUJUrgA5X0qyzEkY6ioGHXjgevet4tg3ZXHRnd179qtKMDpzUEf38ntVpfuiomwADHXqe9Lz1xTgCc8Ckwe1ZXHuIRkkHrSdT7U8LnqKAtCYxCcc0oII7+woGOlAHFIWpE2WODUi4C4HJoKnJ4pQMY9KbYW6CrwORilA+ejHenDHbpSYC4AGaTggjFITgeoNHJ5A4qQFJ4+lNLcfypSRjn9KTvTQADikx3FKwOcD8qF4HXigN9BetLjnH600MP/AK1BPHWgTQ7OMk0hGDxQBjjAzQcAc8+tIaGsecA1Eeeae2M55zUJ3DmrSDYEBB+bpTieeOtMOQxJpxJGMDjvVBfqOzhSf5Vn3Iyxz9a0Dzjn8KpXgJ5x7mrp6MHtoUW3bs8Z700MQMkn3FPOGbjkD0pdvzFQM/jXVexKfQenJwMgGp87RzkA1AhKk4OPepVJIGenpWchN21Hb8ZDce1QycRtjpjip2wT06VDIB5TfTNEdx36HrYXcNxPHpUgQD5iOKiRhjJ7dqk83zUKivDlDqNsdwHO3pSuu5uTkVH/ALJ61KcADB57ioWq0E9BrYjOU70zDO2TxzTnLuclcAdKAny5LfhU7uwBKoVVAOTTdwTAHJpygbgWprkGXKih737jELkKcd6iKkDJNPkDZ3dB6VH1YZ+tTe+4CsBgY6+9MGSMHnFOLKcr/KmGQIP60NDHq2FOegpAS3FRbgw69amGETOeTRKNtxi8dO9VrgMqYxyasJyc9+9VrqXHJpJARq2VB70D5sc1XEhKk0kU4J4/Or5G9g3LhXke1DUZpCenrWYwGQOtOY8DFJxnilYdPWhiFDc4qVAMGq4JBBNWAwxn3ptADLjHNIPkHXk0v3z7CmgbnwegpJNPQBq5wfWrMQCx8moT97AqQDjFVcB6cZJ6UhbIwB1peowelINuMd6SbWghv3QMnqKXhQe5IoOCP9r3oIUD3od90ABQDnPWl4LYpuOaeBzmjmvoDFwVzjvTygCHnmmqO+aX77n9ad+wCBcryacp2oc4xiglQCvftTSw2gY4pqXYLDwD5ZPtTUJVSBS9V9hQ3GAPTtTejuISQMVyTTMbulO3M3AHSmAkNx61F9RinoRn6UmSMGlAyp9aFUN1NJt30AguMkZxVcAMpq1OpKZB6VTBwCM9e9awS6jJUG1cUhB4xTo8bOaTHGe1RLe4yGSPfTR8px371Ow44phUFhQpAkRsDnNCjintjdikxgindgIQfzpyDaMUE5IpTkmluhATgikGck0uM8nj2pD8vGKBig7j9KX3piDGcmnK2402uwD8YUNTH9afzg47UnBPPel1AqMefSpQSVwTxSPHmQDvRjtVtrlENJy200vIHqKQjnHenIMAZNV0GHPGOtSg5xUZIVqbHIC5HapacgJi3c08HIzULkM4FSqDnFTyrYB2ecUA4BA60gPzZ6UHG72qAE6cUqrznt2obqKU8EYPWnZiAn5qM5JFNHehe/rRq9AEHXmlXjJJpuN55pwx0NPS2oxOS2cUHk07pyOtJ2zRcBrLz6UgAU/Wn5zyaianrcSGg8n0p0RJznpStHuGOtIowCvp3q1oMcOW4ol44ojI5okGTUrzAYcYBqM8H2p2crig4prRiG7iDmpBICfekCZBpqLim7SYyxntQfT8aaB37UdT1qLWAXHOaBz1pQec0hOTS1AXsaTdgYzQOhpAvB5p2uAH7oxSPzzRjt+VJn5fehAITnn8qeOFPP1pijtTh0IpvsAoGcnpTSOvP409ehoVRjFCdmAijLc0EZagdaceWqLgIQc5FJ1zSmk6dadwGEjv1oY5HFGwMxBoxgfSqvoFyJgQ2aeOeKjYHdkVKnIp9AH5wMZpAMjJpOjc9KXvikAinaKAaCRjNITQAdaDwelBxjNKDihoAIwOtJx+FKeRSDk0gEbg0hyBwaU4LUHBGKYDSCAeaYSRT2IA96TBIB7VSAFBBz1oI4JpwGB60houBWZRtNRg4XGKslABUMgHUmtFIRHuylRsAealyAMDvTAMHBGRitIiImbBqCYl+BUzqC+AfemyAKnvWsdAEtmPIJq2g2kk9BVOJcSH1NW+duDSqOzH0A8NntTxjBNV5JDs96fDJuXNZ8rsArJufIqQ/dA7ikRju5p4xnmlJtAQuO3XNR+WAMmpXUMc56UNytNNoBYVwo9Ksr0yTVeM4O081YAziond6gHByaaW2/SnY2mmNggg1KVxiq4JwDSM2CfpTQoUEioXlI5q4rXQRZDHA5pGlAAzWfJdlc8/SoUndzV+xvqBoGUKcioJrtBULq7LjNVHgc81pCnHqwuTTXvPFRi6dumKiW2OeVz7VIsPIz1Fa8sFsTzMaZ3PftRvcjrT2iJYDoKbt2gZPFO6C4qbmJyalXOBk8VFE4HPbtUySBsHHFTIa1H7tq57VUmIbJNSzNjjPHeqrsDwOlOEbu4PYiac5wo5oRJJTkjANSQWxkk3Y6dq1Et9oGBirnNQJV3uZ4tHIz09qj+zyR8Bcitryxkcc05Y1LYxzWPt2htGCEfB+TmpUspHxxwf1rdESA8rS4AHFL6y+iDld9zLTTBgcVItgq9eTWiAetJtz3rN1pvdhyozX05WOenvVWazkTlD+FbhGBimFVIpxrSQOKOcFzLG2DkEe1WUvHK+3rVq/tQ8ZKj5hWQMxPsauqLjUVxXknY0lugOtSC6Ukdc1mhsjPf0p23knPFDpxGpmmLgFh6VLvHWstSR1/Onm4wah0uw0aRkwuaUS47mqC3SMMelPE6Y681DpNLUroXN5bg0qsB07VU+0IQcNzSeenPNTyMC2zg4OacrLjk1RMqjPPNHng/xUezYrF3ehOaXAPGfwqiJOeoxQZTnOaPZsLCXihSSpwaqwyFc4FTyP5ikZqvv24GOK6KexNrDLiVyTnuKpgEcgZq2SDSBcsSASTWsXZCaV7lZXx9/ODVjerADPNPKAdQDThbIxolNMFdEIGD1HFKXIGOvtU/2QjADcU027pnjOetTzruBGJCM+1KHGzOc01oXwcjIPakKFRkrzVaBew8MVPB4qVCAuc81XLd9pFLng+3Wk43HdFtMAmmsN3JFQCUqMetSCU9PWp5XcLgseSSec+tOEAY+3pQJAq4PamibJIHTuKPeYWjYSRQozz9aMgjk8elLJ0Oep96aHXGPShC0EAO7GOnanFjj6elIRzk9u9AU9hxVXSHsxAcnHrTu3FO8o4AzT0Q4yBgUnJCISwGOOtIW3EcVI0argUjIB/hTukLUjAyW4qVUyBng0qITxnFShSKTkrjIcFT7ntSgbqU8ckcUq4wDgUr6DSEZOMH8qVVweetOJz+PeljKmou7FWQNFwDUZj2nNSyuUHyio1LOcnp6UJu1yW1e3UQfN1GBSYGOfzpc4bmlHNVfW4WBV5561IR7/hQF+WndhkVF2HUYeOtKMdcUp6UxRigFoSLz+NLsO72poOM+/epFxUu49wxg+lOAGAKB9Kd396hj2K8lsrke3Sojbkng1cwSc0xuKtTewrIqGKQZ60zEwB4zV4ng03PPTpVKbC3QqbXI608KxHPUVOpHcUzrnjmnzMXUYIC2eak8sKMhaVH4qQDk5qZSfUfmNwevehQRye3enn1oHvUXHsiZW6e1Sg5FVlPzY5qdeuO9Q0A/seeajMnOD0pWbjjmoDyaIruHQnDbhnPFIByQelIowvHNGcN70egDJFwagZ9uOeKnfnGagZdxNXERGRyP1prIBzUnbmomYk5JrRABA2jPWkZe1GQx5OaUjcQc8Uw0DGV96a/AyacR83B5pHIOe+OtHUNStIQrdenpT4yPw7UhAyMcCnrznPbpWj2Fr0JQcDIHWkVeScdacO1LwGx2rIdhi7snOakzwaazAcevao2lCrkHmna4bE6nJ96cMYGO1Ukuev8ASrETbge1KUWg9B7HPf8ACmPz0p+M5INVppAD9KIq7B2HySgDrmohPkc1GTuPHSkiQkgHn3rZRSQm7FvO4DIzTQQHwTTlABAPWm4G/NZhsShstTgueppowTx3oyc4qBjsY96YXC8UrnH1qu24/wCNVGNxExk5o7g44pioWOc5xUnlnpwKbSQ2QSjJz2HeoGXCk9MdhVxk5471BMAqntVRl0E0VHUNnPpTIlG33qR2OOAdvcY5pFDMfTPet+gnqToASQSaseXwf6VDGvT09atKOOfw9qwkx9SNUBHIqZEwBj86AAKevp2rOUmGwEAHjp3pyjJ5pSB0xmlHrUX0KE4PNN7mlJyMYpuemT+FNBsOPUU7IHH86jzSg8UrALnnNN6HOetAfmk700Go/I60DOCe+Pypp9M/hUgwF6UgADBoUdzSjg5PWlXkdKQDhjGaim4GelS9ODUE2SPc0o7gRCXJC5qCYYG4UrD5fx9aeo3R8jkV0WtqKxVJOMg5ppJ4z1NSEjOMcjtTH4GMcA1qifQZxuAxjtV2PlfU1V2EjA4Pp6VZtwc5zgVM9gXmT9DmjHp09KcRkH1pDx3+lYFDSCCaQ4JNOI/Okx05600OwzIHFO7nnvSHr246CnDvxTYthTnrQeAaQcg5pd3BGOvepAA1KMgemaYSM49aeCc9MmhoBSCce1Kee/NAHABoPTpSBDTxSY5J7UE85x0ozj+tMNQPqRikJOP5Up5XNRueOnSmkGg8ZBJWjPToaap/Wj1xRYBQR0PenFvlHNN2jIpTjGKAGPz+fNREcGpW5GM/pUbdKpB1uRrnrSgkDHrTNxQlscU84IwelXYlbEpwRuzg/Sqt1yoyMY6YqwDgD+dQynKnjp1ojox3RnMTxtH4ikU7ev4ZpWJJ4JxSYBX1J6D1rqEPR+R8vfv1qdCRklcYquBg9OTxirCA7eAeKmQa3sKzgkbRg1FMRtfAzgY+lS4/vDJqKYfuWGOgoja429T1NmO3A6UsLYzzSNznnpUKzB3ZQeBXhR5pAXlIDAZqcEZOBzVCNgG5P4VfAOzJ4p8rihDixMRLED0qMYABc9abyyZJpyqpIJPAqNW9AEQgyE0hG2UgChiBKdvTHakVinUc96hNLRjsI4O4bulRuQp9zTwWdj+dJjcdx7dqe+oFRyeT/KqwEjtitHarkkjoKREG/OPpRGo7pjKkMTLLz0FWz1x+VGDlm7DihCME96U5OSAOQv1qrPCZOKtE7vwppB2jFYqVndDM/wCzFAfpUMFqVlJ7VptUYGM1tGq+oJEZODj0pDkn6U8qN2BQF9qlyGKBhqd/FSZGaUfMc9qmwDcEnNSKw6UnGDQBgD170ISJAcKR0FKcKAKbnPSnEjIHoO9PcBiZMh4qVRl+PzphIDkAGnq2B0ovzbgObk7RQgIbkUgOZMgc96dnBJ/i9Ket7sQNguCtISGYYpR+7A45NMLMnJHJpMY7gngU8kbcelNxhQT3pyDkg/rT1EAGcE9KejAZpgGSfrUjBQvvVW6oQ0hSc06UKEXHXvinfKI896YFHUnijltpbcBrE7KagyeTSHnOOmac3A684qW7u/YewNwxxyKZyW4pT8uMdTSDOeKlvUB2cAgD60zoOOlOyMdMk0AjbiiO9gI5/lTHY1QLbeDVyU5BqjJ61pHshlkEFPagg49qZGfkHNSHJHNTPRgI3QGm4IGacfekxgVPoMjPLZP4UE4+lOI6GhhzVAICOtL/AHqaPvU4EcgCgBeSKYfu049PQ0E0uoDMEJiljXFOzhqUfe4FN3ADxijHPFGeePyozUiY0nkHvQeuaTvS5APNVfSwyAj5807r0pSvPFKlWmrAN27jg5pEi2nIqUDJz3oHXilzNaoLDQhZt3pUoYgkfhTc84p2cUuZ3FYTPfNL0APrScYNHbiptcAOD0pxPIFMzzTm5OaYAcg4NCjAFHVuaQH9P0oTGLjB96F5JpOpzmlBABNDAX+L60EDJXPFInSg4zmjYQdBTCM/U09j0pAM80dRgchfeozkE+9PzuPSkPLCqXmA1RhqWQHdSnj/AApzYIBpAQ4+lLjHJpcY5/SjopptoBw5xQQB0pwIKAjrSYqWmgDqOKMYWlwCcU0DmncAXGOaMUdz6UnXOO1ACqPlIzzTRuFKvp09KAw6UMBGP6UgyBRjcD29qOfzqkApNPXgZpn3uelCsRwKVmA4HJpwIzTM0hIxzS1b1AlJAJFNHpikHXPejOG56UrdAHDkcmmn07UZw3HSmtzzQAYweKQdDzSbsGgfpVdAG7fmpy4wQBSY5460ZwRQgHdPrR05FNZs4Helzxj0pgKBjNIBmkJoz1pILAOKXPakJ4NNz0oSEP3Y470g4ANIDuOaX+VFihcADIpB0zS9RRIfkoJ2I+M5pccVFvPc1IDk8VVmMdg+tI3AzTxyaimbI+lCVxEbtnvVeWQbsZxQWLZqEqMjPWtoxS3AlB5BoZgPfPekVgWxSMu1sk8VaS6gyFvvZPWgDeeR+FSyqDg9KjzhevBq07CIlHz/AEq8rbowf1qpGPmzVuNc4PalUCJWaNi3PSjaETA/SrUpHpiqjP8ANiiMm0BPD83U9KlIyKgSTA9zVlUJXJ6VE3Z3GRkcZBzQVO2kLEcjtSeeuPm6mhXaAenBz3qwjL/9aqUsygfL1qBLs8j0qvZtgrF+WdQartLkGqzs0p4p+z+/1qlTUUFyQTnbg96jlVtmegpyxhnx6VI4+XmldLYWrM8xEnOamRFjFLj5unSiRhtAHWtG2w2Juo7UgGFzimhsKBnNSdV/lWWwxuxT061EV2NxUpGDUMzbVyO1VFu+gh7AN8xqCSMbaRJGYH1pSpOMmqSaAgdQqetRibap4OR0qy20fd5FVWj+YnHetYtPcTVtiCW5fdyKfAplbO3qPzpWj8xgOh9a1La0EcIOMVUqighWbeotvGE4xVoYJpFRVIwKUnnNccpczLtbQUYzk9KVQMgikY7uKVT2qNQJSRTM5NLnp60hXmptYAzSHrnoaXsKRwTQG4hOeSajJAOMU4kE+lVpZAG5OPetYxBj5eR0yKzL2BcBgOe49anmut3yr1NC2hlIZ2reCcNWS7PTcoRptHy81ajtJZR68cVo29iqNk81oIiDgDFVKtZ6Aloc8dMlPNNbSZ2WujYgD6Ukb7jjFZLESFynM/2dOOPTvU0WlMcF24rpXj4xxVSQMicVary2YW8zDutPeNcoeAKzwj5yBnHSui89Spyc4rIlYlyRj2rWnOT0YW6lXYejH9aTDg5q2kDSHpnFKbchTx16VpzobXYpLIytgtxn0p6SMThqmeEHnv601Y0A47UcyYtRUfPfkdqCeh7U4pxjvTWjO08Urod7oeME/wAqlAxg/lUKAKw3VOGHXtUyY7kbgk544p0OGJPr1pQA3GetSIm0jFS3oCJF4al287j+NJn/APVS56g9ayCw7aCenNIY0PbmlHWnHqOKm4yPyVY4K/nTDaxnPy5qyDjrRjkjn2p8zCyZTNmOopv2NgnytV4jFKvAx60/aSFyozGt2zznmmiFiTxWptBGentSBQM56iqVVhYy5EYDAGDUYZkOGHNarFckmoxCG5OKuNTuCuZ7SnAOOP5U0zHAIHHr6Vp+QueBSNaR7fu/jQqsV0FqZ/2nA57U83YCkD160S2qk/Sq/l8/N2rVKL1DoTedtJJPJpGnDcqO9QbcDmnxJl/mPWnypahsTLcenX0qVZ8jpzSrEg5xk04xqATWTcR6kMkrZwBz6U3e/PHFTgDcSOtDDA6YJ607oNSuJTxkc+tAYg5z9cUED2wKVB8xJ709BX1LcZ3D61Jt4qNDgZGBTyeDWL3KuVZB8+PSnxjqce1RyNl85x6VLH92tHsT1JcA84x6UfxHFAx65o7YqLgNPP0pPp3704n1pD3oQaDDkdqkUjJzmmHk5701Tg8VbGWh7/pTh06cVEr0M/HXrWLiMm6CopOOtCy5HNMdsjOKajZhcQN701nwKapwffNKFyau3cV2KvByaM4J96T5gdtSBMkE8mm7C9Bu0EZqYfrTVXHelyazZVh24ED2pMkHFIB2ox39f0pAOzjHFSK+ME1AT3oEgDdeaGgLPWomHfP0oLdcUwv+JpJMB8ch3YJqXgDJqvkYyKkRiRyetDQDic96hc/KT2/nUhIJpjYGfbpQgIWPHNV3bauTnj0qwTuOc8VBIMjA/KtoibIlkAP1qbcGOBwO9VxEzA5/H3qXbgdat2EmSEgn+VQytnIz+FOw2MUyVSec49cUJK4eQwtgg4HPalWUEdMGm+VuGSelIpRSCPpV2QJu5P5+7p+dNkmOOmexGaapwd2eBTXwCABz6UklcBy7mwfWmGPnbn6VMgGMY4qULuCnrScrAl3KaRMJMY4FX4uBimJGASe/c1OoHH86mc7hG6A/dI7VTkXJzzxV3r7e1NKjdxURdh9CtsAHpmnKnPGan2DNKFGKfOJDVUAe4ppAJ6VLgA4pGAxj8qm+oWBV6+tIRtX3oQ/p3p55zSKIe2e9JgkipQO3b1pOvHrVIWlhFOBwKeDnk1Gy4HBxTFz0osA8881VmAz1qwxwvBwag2kvwMiqgKRAy5FMRDuU+nTNWmjyaEA71rzaCtpYfHGAOfyqYLzt701R78VKOvX61jJjI8gnipQT1qKTjH8qfGc8GlbS4yXgEe/WkOO4pQec9qQnB4rMY1iBxjik9zSsM+xppHHpVIBWG3BPU00t6dfSlckLnqfWod2Tz1PvVJXE9ESKcHFL1PvRGMnrxUvG3Pek2PWw0dOnNOHTJ5prORjNOU+lSxDgvYmjPpS8A0Z54qR9Bc4GSKhfJqUkHj86a4BFOO4Ge+VYgd6kh/1fSiVCrBu/cetC8Jn9K3buidCGXAbjj0xULg4BBqSUNuPb0qE4XnPH8q1iFug9BhgN3TtV2E5UHHH1qihLEj0HBq5AMIBU1FoBZ45OKYT83tQW560jNzWCAUnA9qQ4xQOR14xSA96ZQxiFIOOacPu9eM00sAcGlLbfxpk7McSemMfjSqOeTSHjpzmkX0Pal0Gr3GkfOTT8heSPxpCPmzjnpThgHANMOhIDwRnpSdAaAcGgg5PPFQIYwyaM8cmlJ3fTvSA5bNUNagzdwKTAPJ6UdRn+dKOuKBB/DjvSAcc9KQt83XmgdM96dtBjh1I5peVXikOARjvQeV4HB60gIXOMnHPamM3OPSlZhv8ASmMwBHStEhA3X5uaXGCFA4oXBHI/Omg87TVBezJARt5NRFgV4HtzTlOBg0rDP/66S0HdmfIew4AqHoMfwjoRVi4i2E9wepquuDzg4rpjsT0HnIOQAMd6tBtyAdCP1qsANwyeOuMVMnPPJzSkJ3Q53H3sD/CoHX904B7E+9Sy4HBP/wBeoC4KuM9Aefwpwj1HY9Ukwy49ajECRjI71Ps+Y5qG4kKjAr5+MmtUMtQRgvu/WrMgIYc8VTtZCVHv0qyVO7k9auV07CsSO6kBFAzTiEWHBPzVXAVZdxPFSEhsk9O1Q3u7AIny5YDJpo3O+TjH0oW4UfLnk0jPiMkDvU9E+gwkJDEKeT7Uw5RfXIpgl5J6mnISfmf8qm2oClsAAUhBxknHtTZHGVx+VOfNEotLmYxhO5cA8d6cdqjC0EEIAPzpqDjNZtsYuMLnmmsTgUuctz0FNblvak9gE6AGoz1NSNyaacY6UdQGigHn2pe2e9J/DTGJjmnn7oA600dOelOByBRdgJ9KVuwFIcbsClIwR0zTQDuny96eRtHNMGSfpTs71yaNxDhjBYj6Um7aMgde1OH3MDtUOTuGaa3sBJHuTmpFz1PeoWYs2T2p5k3KNp6dcVq46ASt8zDikdSGBPPakU4UknFIMkc5zUyESEZYAelLjKlR1FN5Qg55NOAAXd3p6Jh0BX2kDGaVhkFu+elLhfvd6azBgT2NTIQh5UYNHODnvRnCYAok5IHYik5LcZEshD4FDlsjn8KUoFfIPNJjLEn0obuA8cjrkUjE0qjHGfrR0BNLdcww7+9GMgse1IM4pWPy4/WktdRW0IWUc+mKoyDHWr7naCtUJyM9auD1GOiboKsHkZqrD1qznHBp1I2GN/iobOfajPPWkJwexqLAJ0B9aByMmkOdvPrS9QB2oAaAM80MORzTycnjrSGncQHGBmjHy0BQc80DIBpbsYDrzSqduaZkbjSjIGfWnbqA8HOTimg560ufl96QEChgNPWjOcUpAxnvSAYBzQwAnt1oU9x0pq+9OPTIoAUMMnFIGAb2poBXn+lKCDk4ptAP4BozSZzwaMVLAXtmkBwtGBQBk4PShAKvBzTcnJpwzg8UnGD3qkIQetA5oH3eOtHIxSYwJPTFLx6UBsDpQME5pgPOMY9aTjHvQvLGgnDetS7AGfl6fSgEEYoOB2o429O9PYBo4Y46ClI5zSjpikJ2j1piBiCBTcDAzQRjikbOBihjsLnrimkHHNKeBzS4G3P6UwBQAKXtSDJoBqWmAuenvQxGaGPIFDdaYDSQBTA/zHFK+ADUAyCTVwjcRYDDOKbnB61HuBHWkHUZ6U7DJtwP170mcGmBuRSk9f5UlFXAczce1KOnTmo+vFKCR1oaAcGIPI5ozikz+dIGBOKVrsCQEdqXIIAPWo84PNBbPSl11AkXAGKaxOKZuxSbx0p2uIceoFJmmE5PWgEjvTUbjH9OtKODzTAwalDAjGelKzuArDnIo6nryKTtzSAgCiwCgkmlLU0HjIoHBosGo4HJpDQMnj0pCeOlMBRwetOXpimDgcUoIBpPUBxbGR3prHjFIfvZNJjnPWiwiKQHHFSR8Jg00n5uKcD+VV0GSDio5D155qQHIPNRt6mktw6lDJWXB6USHJyOgp8yZIYUwjK8DNdKaF0ERfmBHNTTD5c1GgKjIFTD50wwxSdwtoVLiQCMAUqJuQN6U90QjFKhULtB4q00loGpEvEhFWISVJHUVA2A9OIZTuBoeqEWnwetVZIweacpY9T9abJl+AM1MYtMem4kYAHzVbMyrGcH2qiIZCeuBmnGBznmnOMXrcV7Ec1yTkCqw3OSSe1TPET26U5UwePpzXRDlitCXqR/Nt96YoKNk1YZDjpUbL83NCmrXAkjcZ9ventIG70xEB4/WpBCCaylyplDozxxQzn8KUR7SBnihjtzWe+wyAkgbu1V3fJyDz61LK4KnmqbSgduK3gmSTBiWBHAq6hyozWcsqgZGaeJmxkHjPNEoNjXcuu3Ydaj8vPzNTYTuByeanxhOlZtcug1qQNGF9cUiqwPH4VMcYI9KjB2jimm7CIyvGGFVWlG4gdKtyP8ntWeqs7Z9+la01fcTsmXbaMMynHFanGzGeKzocouO9XEIYD2rGrqxgTg8U4LkfhSYJ70x5PLBzWYyUuAOKUOB1qksvzEntStKz429qr2YFvzQKkzuxVCMtuwelX4xgYNTKKQIXHUE8U0kcenan4z1NVbhyinFStdAHy46+lUHRpnKLSiSSTIA/GrNvEFZSTz3rohHlerE9Sp9l8iTkc1YEiqoBrSmjiZTnrisiSEmQgZxVOXM7MlaGrbMrRe4qc+uOKz7XchwOB71oHlKwqx1uUV3cZI/KmQSHfyeaJVKZI/KoFZj0xTha2g2aDOCM56CqFzcYG0dTQfMLD0qrcgiT+VOMLu7FsQkEg4qDyS+QBzVuKIntzU8UGwn1rZ1Etg3ZXtYmjByMGnPg59RVll5x29agkUAday5uZ3YLRFORePY9ai2bfpU8hGckVCX+ccda6IthYQ8HrwfWnL6ZpxXOPeowCW46HrTEDDnINPjAAxmoJNzE47dhSxxOR1p201Y76kwZfXHpUivnBHT1qPaMD+dSJgCodhko6Z9Kaee1GCRkU5Rg8CsxDl7U7OB9aQDnrS5B+lSMcAc+lKR8wo6jn8aM85NSAoAzzQelL39+/FAGQKLgNxn+lBBNPAwelO6H3pXAo3C5kGDU0OMAGnSjJFORcDjirctBdRwUAHimlfl96kzxik/lUXHYoOoGRVNk25I/KtZ0Ug4qs0IPB710RnYmzM8pxRGuX5PIq5JCNmf8moliIfkc1rzpoHuSJxjnoKm8sE+tNTjtUgPpWMnqUxRGF4qKVBjips5oYDFSm0xW0KDocg4pir1yeOxq6yBhwKjEag+9aqYMFGf605s7e9JkDr2pA5JNSNFZlycjrU6Dap45HaljAJyaeOO2apyEAORzRvG33pT0xjmoAOvPH0pLURMTxmmFwaaSF781AzDJpqIyYn0poBxnH605MAZxUiAE/TpTvYTCMNtzijaTnnmnE547+tIM8jHFTdlaAgxinjk0oXHNDMAB6VN7i8hu3n270Y4GOPam78jqaSRyq/TtT5WArsFbGaQXCce9Vi+4kdAKYq5brx61ah3E5alzz1b7tSA5FQIgB9+lSgkcelQ49hjycHJNQmYAnPSnlgRxnmqzIGY5GKIpdQuOef5SB27VXSRs5OME1KQBwKVUyOlaqyE73H+eQMd6iFxnNSiMk81C6DJ44NJJAItyQTjrVtJuORmqakBhznnrV5NhApVEl0DUcsnPNNd+MCpPLHVqb5YUZrLQrYrgMTyadtwMkZNTYGQRTW6H361V7hZogJ6460LjAOKU8HqDSHGKsnd3EB5oOMYI/+tSNjP9KMHOD1pjGvgIe1VW+U8888VZkzsIAqvtxhSeKuInuKEyRzwTUwjG3n8/SiNQFxjIFWABspSl0AjHIwB17VMigKRUe0A5xxinocLWcgF2gMf5VKOlNxyT1oIPX+tQNaCkcGgEDk0p96CM0XAO3tRjkYpQOMUewpAhp+lK3qBQTg80zPzHPSmgAcnJp6g4zimkeoxTiw4FABwKbnknHNITg+ue1BPpTDoDetNPrTweKTtntQAx+nHemqvXmnkU3dximrgMbAGaRcdjkiiXJCgAcmljXDc9avoIlRQMZ6etSKBTVGM/pUgGKybKGN93kU1Adx6YqRs5PFCAAdefSjoAoOB701jgZzTunFRuCVOB16UIBQw4puQc+tRlgh5HNIJB+dXYWhIDkHHNMKfMSOlOjzzk08Yz6mlewAABkUpYkcdKRjx0qPf0pWuBMB2p6rt4qOMhslT/8AWqYZxUyGHqKQ8nIFBNJ2xnrSAARyT0puct14ob60jEAgU0hO3Ua4Dd6aVwAD+tPJ5x0pp5FUgKE/MmOg7iqrELuyM5HGTVy4U5Bzj+tVHTkA85711QegmSwYYjaeTWgvC9KzosxnjketaCcLyaioEdBw96RjgcHrS59qidwRWSVxkiNu6H86CN1NjYsMdu9PY4wBQ9wurDP4ulL/ABY7UZwM0p5AFAavYPcU6kA+XAFC8j0zSAU8DcaAMjJ6U057U/qeKNguPGCCTTWJ7Uq9gD9KRiBnPNT1B3GMQGJ70Z54xmhjk03PHrg1dgHKB9BRnGaFI2+tHGc+tICOU8Z9qI3BBx0omYYHvVeNzuxtwPrWiV0F1ctEZ4HJ9aASeOcVGG5yT+FSLycZ/CpYtCJo/mz2NMMJDHngdKstgZwOf5UzGff2p8zDQixxzmkxk5Bx7VKOuD1HamhRnPencLsQgc/zpO59M8U5+/FJn1GQOlAFWdR5fGcCqK8E8Vo3EnyseelZ4znkdfSumnsLS2o/ALA5y3pU4BC5PSogcfMAPTipYyCMDBHpRIe4jKq5BH1qEqBC2cdDT2fa/wBTSY/cMO+OaauguesyDb35NQeVnqM5qyygYzSLwCSOa+dSV9HYBIkESDA5PSpeSQx9KaW3EE9Kk3A7R0HrT5gIWU7ht61M7ARFAPm/lTcYbK8D1NNdWxkHn1o5kkJFaHm6K4yR3q1cFVjIzzUMcRhbfUmNx5+tVLltygVkjYvkg4qY4bHNOZ+yjim5AjBJrN67IYnlANup7nC5NIWBI9hVeRmZjik3dWGP835SAaVeEyfyqKBBuy34U8EEkk8elKUewxwPUkcCkHz8UuRt6cmgDame9TZJ6gISKYecntS44PejHy0r6gRsctjsaU8YFOxxxTDyx5oGLjIApx6AD9KaeoAp3TpTuIB14pM/N7GnEYT60nG3Hei4xQecA1IoAHNQ/dPFOXJPNHQCQHAx0qNs9TwKkxg4pjtngdKEIa5+XA6nqadEAuPp0prgbRjqKWL5jyeRXRfXTcCfaHYelSPt2hR1qKIFnPpUhGHA/I1Das/MQjcEA0uRjrSuM8k9KbkcHHNKVxdRWJHJ6GlV+qnoaVzkBaaV24PpUcvLLQYh4x6elKfm98UpILZPem/dU9s0mnZoY3GRzSA5Oc/pQclRS8Yx70JOwC/jzSkYXqMUmBzTsZxz+FFtAEx8hOabgkZFPGC3PSkL43Y6GnZCuQt8wqnMo39KtnINQTLn5hRHR3GRJgN9alc9qaI8YPen8mnKV2MTpSU4fepOrUgEJ4x3pV6ZoI5xTQdrEdKe4Ds8k03gmkL0ucEe9CjcAJy2P0p2cZGKYRg+9L39qaABg89aUdKaCe3FAY0rAO6k8U1W5zQKVSMYo0QB1pGP5UoPWkJyv86BCE80E8dKaO4p2NwJpjAc9zSHrxRuwMUinnn8qPMBwFLnI5pDnj9aOMcUPcBc/NS9TxTM8jP508HrSAM0Ht6Uh6jmg80XAUjp796AMDnpRn1oPqKAGofmIpw4poOKUZwTT8wuOViCaQnn8aMcE05QWGT+FCAa2fSgEsvvS45pBw2DTvcQoHegDP0oyFO2lU9qljG53dKTPb0pw45zTf601oA08nBpAKcBhjzSnkZpqwDVyGODTsfxCk9eOad2pt3AYOTk0pbgYo4pOaXqAx2AzmoM9fSppF+UetVXYr/WtICHA81IjZz3xVMNljzwBVmJu1ayjZDJMcinNyAe9MOQcY4pMkGsREg9fTrSjOaZk5A7UuRg4psBWJB5pOv1oJ4zTQcCpQxw5GSaB196TPHWkxjvTaEO5+tNxzmgnnrSA5PXFPVABIJ96QnHApDzSYOfmNACg4zSBgeR1oBweaOBQMcX46/hSZycgc03OKXp0/OmA7fzgCnFu/SoSQT70qtk89qGguS7uOBzS59ah3c4NDMetLlAm6mkJxTA+OtKGyeaVgHA5IJp+QeKjxgcUoIH0oYhGGP6UgGRTjyKTB+lA+g9Mf40jgsDQvHPrTiQaL9hWKbpwFpVXCVMR3pMHGRVcwWIxjpimhtrEDkCo2LiX/ZqZcZBNXayuMrSk1DHw+SasyPHg7uazrifa2E6VtTV9Cdty4Mbgc0/zcsBjis6Kb15q7BknPaqlCwyywHB7U8BVU+lMPzg0oUlQPzrnAnXDLkflQUGPemIdnXoalDZNQ07hYrPFzTNgU1O6s3TtVd1deSMmtIPTcTFdQVxUCxZPNEk20e9JHJuPNaxjJCsh5+TtUiHIzQDk+1SbMdKhsrYj55qOTd26VYxtGKQpkYNLmsBlmIufm4FRvbqEJJ6d6vyBY+Sc1m3DvK21OBXRBuWxDRTUFnyDlRVlSMbai2+X360sR6cV0sSdtC9C4U/1qxJKApqiHxzj9aYzMTx3rBwUncu9kS/aCXAHAqYOSuPyqvDEzZ46cVOIXzRLlWgk9NRWww69etNEX92niM5561ZRccY4rNy5dh76kAUKlPtjuYHtRcDjbUkEZSPpSbXLcNUyyAATimvCHNOHA+vendBjNc9+wysbVc9KdHCFHTNSsewpc85xVOUrBYaIx1xipB703p16UgPtUasew884qndKxHT6Vc60x1zmnF2YrFa3GB0qYuFAz1o4zVeQ5ya1XvSE9ETecS3HJp+z5siq6AbsrVpcnrVNaWQEkagg8VPjCZzVWM4apxlh/jUT1EJJjb0quEUEj1qyRx7ZqBhg596iLtsUNdWCfLzUCxFvv8AUVZcnZxVeMsxOTitEmkFiUAAY703pmgsVPPSonkGeKhJsGOJGSM1DKDg9zUm4EUxgc8dKtAUpFyM96jAAGPzqzIuTzUJBB54FbxegrAeg44qN8BcYpVO5sH9adKmRxVbMGQoMHPapkUgcUirgY7U9QRTbuLQQjHPamI5LYPSpjxTQuTyOfWpvoUPBzQGOT3oOFUZoXHpzU9AJE4BzSjG70pQuevFAwOlQGo9f1o74oDccCjA/H1pAPHXgUfhTM45phkHNCTYaEuSc0gOKhMnTFOd8gYP5U+UY4sD1HSnFgo4qIEAcmmGVaOW4tywCTRuAbk1V+0jkZqJpSOc96pUwZcaRRnFQlw1Qq24nPShRk4xzVKFtwJfvZ5oxgGhRhjT8DgH9KT0AjC+lSAZX6Udvc0oxSbuAgXHWkOT07UdsUjMPwoQhhPzH0HpUbSAtgdqkZhtyOtU3cljitIq4yxnJ600ADjPFMjRm65qfy+OnNVJWB3EUDHSpBx9Kaq4Ip2DtxnpUNgITg+9N200MQwzUiqcijYVrkTLkg1WdAXOehq8QMccVF5ZJPNXGQnqiAI3GKsQrtUen0p3l5OB3qVUORnrUyloFtSMk7sUo6DmnMMsMUgXav8AOpK6inHHHFBXd/hQM4PPFO6D3pbCIWXB9Khk54H41ZI5pnl9vzNXGQFIAdxk9qkjRjgEdal2AEH1qZQoXIq3PQVhoQDA60pGCKkXBzSEe1Y31KsMx8tQsTuPOB61ZIyvvUMydB2q4vUnVIi6/jT0HJNNXkZI6npT1XselaMTuKVOOlIVJx6VIvqfypcE/wBazuOxWEGDgcfSpootoBqQqOfSnryMZpOTaKsL1XrTWHy+1PJxyM/hTSdoOec9qzQmM+UcAUjABetRySgA5qIXJbOePStVFvUNNh5xk0xiCOKQuM5xzmkBOenSqSB22FGCMYwaeAKaq+vNSgcYzQ2JaDCMr05xVZkb0/8ArVdH3R7Uxl9BRGVhsiiTA5H0qcLgcGhE/SnkAAHNTKV2KxHt7YpyKTxS4BJNKp560rgkOxS4wKMkUAnuKkqwcdzSY4470oOTSdKBWAfeOaU45poxuwetKcnrQwEJAPJ+tIeP8aU8qRTRnaKaAfjC8c0nQDigE9KVhxikAwjP9aQ/Lk5oY5HBo71QbBx3PNL1BGKMcg80nQUADHI96j2jgenSlbkU0nLfSqQ9gzzk/h7VIFz0qLPTFTICV9/50MVyRRxj8qcBjApB8opRg9BWQC9Qab0/KnjgYxUbZoQC/wAPP5UmB+NGSegpN3pTGQSDDEnmq+eeB3qd+uMc96jVc4wOK3jsIsIOAM04mkVcYJpx5OayYDeM+tQPuDE9vSpzndjFMxng1UWNhEDtFWVP6VCmMjvipgRjioluLYM470jZFO7/AEphNJDGse469qQnI4pN3rTc8VVhdR/UVG3UjOSKcCBjjmmEk844NNIYhUEVWMQVuO/erfQc1FL1/pVxbJaIjGMcirKMoGKrk4Gex7U+FsjFOWqC5Mw9OnaoVHBGKmOCc4zTfu/SpTHuCqV6UucjBNIHB+g70zHJ9qLdxbIex9+KeOmMc00YAxtzS9O9IaF6LxyaB97imE5B9abG5z1osOxM3fFIrFTk0deO9IRjnoaQrWH5564pM8dKM96Q4xkDn0oAQ8riowu3A9alJGOv/wBamdqaYdQzt4xTg3OCKYclfrTQcMeOKdguPZdw6VXWI57nB5q3jjFKMA4/KhSsKzIsev0pw6ZAofG7OKZuxRuFyTovp6VGcdMfrTyR60x+mM9aEMQZHPelHDDPApFyrZNDHB+neqeor3FbAX60wdT9KVydpPU+9Rp8w9qEtB2IrgDYTzVAklu/rmtKT5kORms9yQTuwR6Zrop7C0sIpZT7fzqb+LIPHpUa4PQc/wAqkT5TjnbVsBhU9xknvmlk4iOcng1KcEBgMVCQTE3fg0luTc9d29Cxpxw4AHHvTSd3figsAcLXzztYoVtoUAmnKMgZx1ppUBdxOTjgUsZGck8dqV7MCR9uQF5PpTZAdygj8BSghTx1PNOB+cO3boKpxTEMKjblj36Uwn58AdaczbnY549KaXDN6YpNXVwGEbQc4qFlJA9KmclyPSmsvzYJ6UKQyPLbCe1NUEDJ+8alwMNk4ApoGQSelZtu4xqrhhnvSYw5x0pQMv8A1pvU7R69aGxkuAx6cClxleoyKaxwQo/GhiBxmlokA3oDz1oIxS9Mk9KQ9agBCcJjvUfTFPbO/BpHHYGq6XC4gGTmnDpSEe/NGQMUN9gHMOaCMD3puWyKcQSBipYCYxj1pU+9k0DkZJ6U8cJ71a2bGKTn64poAwT3pRx9TSqme/1pXs7iI++CKTlXFTcEqO4pjctitE0BMpwARx60oyGJzSYIj96Xb8nXk0S+LQQucqeaMjbnHSgphOvNIMgihyV9RDs7mH0pCc8HpSMcvkDig85P6VF0xgCMnPbpTdxPuKQ4B96AMniobbGBzjBoGQPelPt0oyduKWqeoBg569aeuMH1FMBIJpy9aqO4mOGNhqMd6kyFbpUZwZG960a1QiNxlie1McDAFDvgkDpSE4AqfMoQnt6Uh9aDx0FLgEYo8xjRwf5UAZJINOAHP8qQCnsg6jc4pj9M1KRlcd6iZARQhEeBu+tOH3s0pHAJ7UqjjIq210GIxO6jPGadgDNAxtxU30ATB25oBAAoIyooxkYo6AKDxijoKMYOM0p64qQEHSkBGcetL3xnimnue9NAB60o9O9N96DnOc1S2AX1oAHembhu60p5OKLMBwIH0pD0xSZHSg85o6iDr1NPXAzzTFFA4JpWGPHA5NJyOvSmjlqUnJ5/CgNRf4qUY5pM0hPGaQDugpf4cZ5phPIOKAccU0A8HjFKDkUwnHFKAexoAeOnvTQcjNAPzEHgUwHacDpTSEPzk9KT+Km57UqEHr2o16DFLYpRzwaTOSTSjOc96PUAPvSClZeR1zSH0zUvcBuCO9OB4ppHOacW+WqbvsIQdc0h4o6ikYnpSt1GNk5HB/Cqjx4+YnirMn3arTMxXA6VrT3EQMMNx3qxb9SPyqqqvvHFXo0KmtptJWAcQwJzQw4Ap5zgjvTCcnmsEMXbxR2xSEc8U4HsaTAaRgH+VIcjHpSkUhPqaYCE89KD0zikyM+1OwD3poRGTke9IBnvT8DOKdheKLgNAHXNNcHrT8YFNOSMZoQEeOCc04nNKRxgUAYFO4DScDHUUg+vFLk5NJjAzTAQjJpTx3oH86U4xnFFwGgnoKQE5ozk/Sg9jTAcGwcHpS9xTFPU96M8gUrDJg+PxpwPHFQDr1oyw6EUcuoJk2c8GlORzUKse/4U7cdwyaTQErNxjFRrnJpS2760iHHWi1kBJnPTpTScGlI4wKhlkCL1qUrgMkbAOBUUsm2IkdTSrKrLnvTZlZo+BXRGPRibVjJeUs3J6005YnB5HSnzRlT05pqfe+bkV2xemhF7k9vtUjctW3uFQfL2qpgbaQAsOamVNPVjWxdhuwTzUj3gQ896zSCMECkcknis3SiPmZekvsjA60sdy45yTWaAVOT3qzFcDgdqbppLQE2zRF2Sp7Unn7+OlUmfPQ4FRmcRnOcjNZuldhctTr9KhjQl/l6ZqeGZJkI9aVYwrZqk7K3UGtSzHHhcnpSlgTipEG6LAxisy6DwsxByKza5pWHsXTgNgUjHOcVUt7kSLxUVxLIuecZp+yd7MGyvdzmWbyhz61cEAEGQOar2lnulLtWo6qiEdsVVSajaMSUr6sx3iODkfWo1QjPNaGQ+RTWjCj61ftH1C1ioUzkAcVJFABxTguDyPrUqKMbifrVe0sg5SRBt4qX+tRh1GehpDOFIzXO7tlWFKfNUwBK9earC4BbirBbK0pJoBGj3nnmpVG3Ax24pqsAM9aA2anUB5OM5o3ZqN2wAaTfzz3pcoDyR2pymoS4HNCvyc0W0AnPQCkB6CgEEjmlxUjY7IOKTt0phIFO3dqLCIpDjgVWkzjPSrMiHiqz9a2gkJksAHfrVpPQVUiGCOasE7TkU5WbuCLCKAMHrUm0KvNUklO/iriyEr24qJK4IYc9RUT5Bz1qyMEZxj2qNx3xWetwK7k4qtEGWU571afhMDmqyk7ue3etI6ASTNwBnFV8DOc0+Y5OR09ahBJ571Udh6Eo5FKTnpTA23BAo/hyOlFhjGH51A+QcnrUrEngdaib5lPHPetUrEkCcuafkk4JqM8HjnmnLyDitGg0Js4XgZ/GmBzvp38OPSmEEHI60kDH7gepqQKG59OtVwf0qaINnjpSkh7kjjKcihPSldcr6UyHg4qEroOpMAT1FNOVapByQfWgj5hxUJjDHHTrTlH5UDg8U4GlcQwrnPvVdoyMjtVotz0wKTjbk04yaCxWKZOKcy7ASevapPw4pJOnvVcze4bFKR2bIHaq7KT94/jVl/kOKaFJOAK3TsSQbCTnNTBPl5H604pzxTyCB0p819xkakAYA6U5TyM9aFQ43dKEi5JPXvSbVhFhRjnOTTyOOlNBG3HOPWnA5rB7lDGaoHl5znge1SzA7T6VSGW5PWtIJMTkTibIPc0IhkJ3VJBEOufwqcgDp0pSklogs2V3G1cVURDI3SrrgE05IwvIojPlQmtRscQQcClIPWpTzzTThT04qbtu5VhmOaDkdTxS9/SlOAOetMCIjn6UhBJ9qfwfp60hbj/PNAhcHbn2oABOaQHoSOKdkCjUBwHNSHoMdaaOPpTwCRwKhjGFQD/WmEYPtUrdcmozycUJgMI5p/OBQfrmnZPancRE2SaC/ynA5pxH69aZjP9KoCo7Mx/zxSpvbjpipthDHinKME+1ac3YWo5AQKe2dvAxSYPbigg4NZ9Sgpp5PNKo6ZxijBP0oAj2jOO9O2+nXvTtvzY70rDHancVhRHjAx2pNvymlBz9aXO3jrU3YxmKevTrzQBkcClGO/buaTYCtnFVpXwMdfSp3bjsKrfecE1UEK5C0bbcmmxrxzVlj8oyOtMAwuRW12KxGxCnIFRqcnJ7+tPk+VeR1oSPrnv2p9B7kir8w9KlUdKYoPQfnTxwM1nIEtRhwAeKA2aac4x696FI2jNUkIkyAeaDk89qidweh5FPU/LSasO4rsAv1qGFiZCT0pZMnOKIVOckdfWqVkrg90TSPg06M575p3lbgeOtCptPoayurBqOH3uKYQMmnkZOfyFBzjAFIY3HpRg8j2oIxz+VBzj3piGkdhQAOnf1pxzjP6037tAwTP4UEe/ekwRinAgDOKAI8YOadkBiT1pH+7j9KbuFPcXUfnPFMZsDjik3EHjGKjJJ57GqSBXsKWycUc8Af/qoCknrx7VIFyKewXGIuTVhRkAioQQpxUyt2FTIFoPYEDHrSqMA0mcinDmsxingCmc496UnimdqEAjHB46U1TnoPwoPORnim9AB61dgRA4JY4qSLB4qOXoMc1JF93JHFav4RX1JRjPtRnnFNHAPr60ufl65rINUIeuaYc7icU/POOtKVzwRTQhqnjjrUiNxzUY6EninKvqaGrjuSg5PtSYAzkflTQccnmlP3eTU2GRt19QKaScEU/wB+1NY5J4zVIAzx04pCN3H6UuD1PHvSk85xQA0HPfiopCc471JnBOBURJzyKpLUT2I32kkevapolx1OTTHQkg+lSqCy5HFXJ6CWg8ZznFRz4KcH8KkOQvFQSSjODxmoitdB+Q1GwMmlEg2571CSFGRyKYjZLE/dB/OtOW4my35g6CpQw7jmqgXbz+dTRn5eehqZJBcewBJKnn1qLBzz19KUk46Y5p2Sfmx+FJaDtccOlObgcUdTn17UmecVIDgMd6QnB/pSjGOKacenXrQgYrDgHtTQc5H8qU5x703BAx1NP1DZijpt9qAvqfoaVRgDv60pPfrQKw5OB2GKAe44HtTATzTgcdD1qWhjJeW603b2JzTiM01+OPSqWwtN0BGOaFORUTs3OOPQ0uG4OeKu2gtLkhfHH8qgL5IKnFIzHGB2FRwtknPzEng01HS4/IflscmpI+DyQaaI+ehAp47460Nghso2oSDWbJkNyMDt71qSMNpXHXtWZO35noK0pCsMUHeS3rU4bJweR2FRJluAwyO/rU6D+InpxWkhtjiM4DDI9c1G/wDqH+h71NkBgMjjrUcoAjbjgA/yqVuDtoerZG2nIAPeq6MC3zflVgEGvAtpdBuKwAGT+VKq5AwOaMDqaVW24x0pat6oGOPyOO5oAKsc8k0gbqxGaX5nO7pxTVhEZXgt0weaEAc7jwKeEAA3ng0xiNxxwoHak0la4yNjgnHaow25yPSpGwMEdPehAACTjOKnpYCMDLYzkmlbKtihCFB/nQOWyfrU22GiJm296SNsnA605xuyTxk0BQqYFPpcEKD8570oXIyagExUnipEJYc9+1Jwdrj3QrZJFAPJPcVFPLsXApImOMmhR0uG5KDkknrmkHJyelIDk9aXv7UuoCEfNkdKXOTmk3AnI6UvcYo6ABJbqfwpwGBTeCfenqBnmpbuMEGcmgdcCk6dKP0NO4hcktj8qkU7RiowxznrTx82WPpTXcGKvJ3YoXLMT+dCMdpGKD8owO/WrT69BDlJY89O1OJAzxz2pijHX681Io4JNPS11uICcYyO1N4LUucjmmEH86zb6jH8de1N756CjOOPwpjEn5feovcY44LH6UKfzxTcc+4FKfTnJp9gEzx704A4FNHL89qcTjJFDCwuM/Wmu+BxwaXJHIppAJyTjmrWqAer5TJ64prN0bFCc5ApNowwPpVSs2BAAS+f1oYH8qeBhTTB1yajoCFJJYGkBycjtS9sdqBjaQKBh3zSDO405cKeaDkHOOtPazFoMzuOeg7UMBtzTsDpigDkjoKSdwISCV6YFAX5aefSkBAGO9abAIBkUvAUUmR0obg1AxB940qnmgfLnFCYyc9KrcBD1pd3tmkA+al4pAL1FJjj3pQMUjdaAG4BzSY3AHNDc5Ham5NXdANxninEnikzjtRu4NDDoKe1HekXrz0ox1PakgFB5oHFGMA0H7ufWgBD65o/hzjkUnIGO1I2VoAdnPXrQDmmbqXdRYB+aUmmAnIpwzmi3cB3UYNKRg8UgJz708joO9JsQmcj3pgyc5FSYwwxSnAaqVkMrnO7vk04AnipmXnNIfaldtACL8pPenZz0FKBnp2FKV+XND7AMPzfhTDnFSgcYphXvSbVxCAZHNGOfanAY70nekO431FNPTkc05yOabzimIY/IxUDJhamdc9MdaYRz14rSLtqMRRhRxUig4oGCKQnHFEncOopPNBABpC3pQ2WHtSsIaWCHFI0gzmnMmeRVeTCg5FVHXQB6tubBPWn45qqJRvBz1qypBPJ5qpRaAUjBI7GgLhqc454pDgGoGIRk0HkDnmjIzTWPPtQhCscAU3oKRiAozSE496aQDjyDzxTclsik3U0E5qkhjiccUhyBTC3ORQXBFOwiTPakJznJqLf2ppY9qfKDJC46UxnypFRM5x/WmNJsQ5qlECVJD0FWF/SqNqxLMaurnPJomrMFqScDvTW68cUbu1DMFHqKiwCHPXvTGJHWhZAxqOR+TxVpO9hkiyc47VL5mRzVbflvSl3E45ocbiuXAw6VXuIy/IHFQtMY2HcVZR969anlcdR2uVVhMfA5zVkfdxSk1CXG41V3IVrCTwqy5HWqYtgDg8VpAb1x3pkitjjrRGbWgmkUjbgNxzTsBIznAJqViQTmoplDjmtlJvqKxUZhyewo35HIpP4SKQrlfauhR00AThmIpRGcg4oUAcZ5q3DGG5x+NE7RjcW5TlLjAxUR3k4xW0IVfqvHal+yR56cVgq6Wg7MzbWOROVq9tcsCQasLGqAADHrUgA2Z6Vm6rbuNJojSYxjbWVey+ZJx0+tXLq4SNCD97tWScsxx36Ctaau+Zib6CK7QsdtWUWSUhnHHp1qNIWzzzWhEmEx+dXUmrBuSRDaAAKfOxINCDBocZGTXLfUZSjJMvNXSmQeKoGcLMFxV5WyvHSrqX0CL0K0qlTio2PAAqzLnBxVMc8Grhqgeg0Pg/0qNyxNSOMLkcgVGpBPP6Vqrbkt9CaJSB71dUkLVaNgpzU6uGweprGdykS54NKh3Z4prFiMClA2g+tZpDuEvSq+/A5qUksSB0FQMMN1/Cqj2AlDZHtSo459aYrcYP507HBwaGu4aFhCGGakz8tVogVJB6VY4I96yloxiHg5pVByKYQSakXmhiB8MMdqpSrtORzmrrHCiqkrZXjrVQvcTIo5SOvWreQy59OKohSCAPWrqLtUZPPpW07dAQqr0NW4sBCDVMPyAPxq3H0yfSsZXW4D+meajL5yM1IQc+lQsSW9R7VmkAx8496puzK47fWrjdTiqU8bE7gc1pB66h0JSm5evWmeWFzSjcF+nalUFzz1pjEKgj2pmOMdKlCkZpmMnihMCIj5jUbEhSKmbCkk84qCSRenGa1jdhcqupP0oyQR6UyVjkj14zSp9K3toJkuXJzk4peSw9KAOSfWpDxU37BYaYwOR3qWL1A6U1SG561IijqKiT0GSYyhFRQrhzx+FTr3pu3BrJMQ/1pm4sfTFPPQ+lNXAoQx4wQacACxpqHNOU857VLDqNckDio0ZiD6VPjJpTjPA/CnzBsQ4OeKRkPp1qfovNMPqO1HMBA0IJ570GIDgVIMk5xR7GndiGbR3FO2Dr60o579O9OU59qLj3I2QdulRqhxmrAOTjtRtGQB1oUg3IVjOcYqQJjp2qTofWnfh+NS5ARNHxiqzQjeRirnPfrUTYzVRk0JjNoX60j54/lSls4FOzkU9RjAuSM81J7fpTcmlIxj3o3AQ5zimscHjtSsOORTCcfjTQAOtNck4I54qN36U1W44q1EVx+cn1peCAMZpV7cfhQTyPT0oAXjIFOGAPeo2B4z+dPB7d6VgHhgTgDj1qVTxUSjFSDnFQxiNnqaj4AFPc8Y/Omheo6e9C03AamSfapApFHQ4FKWAHNG4DSAFyajDckYqU5xmo9uPxpoBmcEj9KcPb8qQjL5xTh1psEIBnr/wDrp2M54oPykCgZxk/hQFhBxzS4LHOKaGJbpSlwp9KAFxzx0prjkmnqxIprA4z601uDE7jinYFIp45pSQeO9LqAuRn2pCPy/lS9enemk8cUgBuRjtVcckdx61OQcnFRsB7ZqoiaYwnJx6UmSfalO0CkXHJA5NXcBpA3Z9O1PC8ggfWgKC3PenjheaGwGB8k4pSeCMcVHvA496cWBU+lFgsIQc8dKTac/SgOowO5pyndnHSndjIG/wBYBgYp7H5doFITljxT1TviqbJe5Bl93I61YjfLbaUpxx3ojQh92f0pNpoNS0pAApH6Zx0pVNIxOCBWK3H0BTk8d6QjgcUq8D1pxGQP60gIz05oA44NOI9OaTqDTAYc8jPTqKY2c9Dz0qRsd+lMzz/SrQAM9O1L2BpOhz3px9c/hSBDWP596iJHbnNPbrg9RUZ4OaqIxr5zjHWnBQBgUijNSdsdT9apsVmJg5p4OM5H6008MaAck7j0qRN2I5OGxUkcnPI+lNlXkMBSpjIGOneq3QWsyxuwgwM5pwyM0xAW61ITjisWVdjc8n9RTC2FzTXJycVEXJ461cYiuAb5z6U13K49aToBnmhwc4J+tWlqK+goG7qP1p54GPWkQENgj9aJCxOB25o3Y3rqOVsdPpT/AMOKr7iG64xVhamSsO6EHGTj9ajeQk/yNOb7uD/KmBcA00uorjg/y8jmpVbK1AGVcg/nUi4yKTQXFPXjtQTgcmndeaY5AGWpIYik4JPc0nXvj1oR8jJNI5IUmqS1JH55P8qMdOetIhzg96eACOe1SwI5BgEjtUEg3E84NWXwep4qu2AeOauIS13G5bac06KTK8dKQj5OB9c0wMUJHX6Vb1FcmZjv69ahIVzzzS5JOAMH60q+nP1otYbsiHPJLcf0pI+XJxnNSE9SFyPSmg/MXx161aegrEjtgjBNPQnknr601sbRTlGUyMVDegwI7dTipEXA56U1RuXqaerYx61DGgLcYpOcmnkDr2pjjnP5UkDuPAwOnApv3eDzSBs85pd2QKNh6MOSBgYo5UcUpIz0pGGRjFBPUT1wM07p24pAMZyaUZ259aB9BME5bPNOGCec/WkzgcGlyeffvQAwg0w4I+YcVIc44NRscHgVSFpYa2MkDpSgY4pT3OeaYW4prUaViNl4KjvUaKR2PFSgHIz0pueTk9PStEyVZkqYx16c4pxOOMcetNB28Dn39KGy3UHio6lbEcj7R71nyn5h696vShiCapE5Jz/KtqYnoEQIYE847CraDCjHSqgXa4B+gGatK3GBwAOKqeuo9mNC7pOc5PTNEi4ilIPG0j9KYZctyOh60MS0cmDn5Tn8qSTuK9z0/qRtqwvQmq0QIBI/WpgSQB3r5/oBIPugdTnmlDZIAqNuFx0p6YH0pxXYCUMfugce9LnOVX6UwsWYgdKdGAoPPvRdvQQ6UKu0bsnrUT4EeByxpyqZDk9hTMAHLdjxSlHcfqRkdAaQKQx9ae/LjFRncrnJ5rPRAITg89aQnGCOtQysyk85GadEWZeRk1XIrXuPUUZOc/lUcpzxnipnGw80zysrkjioT7jGRxZHNPjTkknpTzgKopf4fY1UpXYERQM/rSYw3oKcTgZpOpzRd2ARcA0AEnk0tKAdvI9qVwEIxxmlPWkA/ipeTz0pPQAzilyQuKQcAnvQOSfSkh6DzgAY9KQ8Y/rQD2oPqelN2EJg5xUuMAY6VECTyOlSLjqafQBc5IxxTx87DPb0qMcH/Cn7sOMU1pqwHvzgDtQTgDnigtlcdzSEHb/WnK1xIVe/86aB1JNGeOOuKQZ/Os2u4wJ4zTc/MfXvTXJHFAOV/nQkrXAfTGO386XJC57U1lyce1KKsMcCKeBmo1PFOGScZrRJiA8cDpRtyQOtOUANyfpS5wx6E1aXUQj4RuOp61GOpPUVITuBPOaiPAx61m3ZjEJJXgcU0fdINOYkIQKb0X60egAOVIFC4Apq9D60owcmlYY4A4LY60ZLEcdKUMdpFJnAGDQIAMsc0057dKcBxmkI+bjpSuMQrjmmYKnpTycj6UjHcOapCG4IpB9eaXtTGOPrVIY71FJnFA+7mkBIB/ziiwD8gDIpDgCgElaOD9KNWK4hbdig5pxGGOKaTjrSAa2e1Jg5pSfSkJ5xTsMTsaQHNKTnNCjJpgNIpxAI96Urx7UuPlovoAnQUhPy8U5cHgikA60rgxmeAM0h5GKVlyc008c96YhD1xQCAcUmTyaDyc1XkO48cGnKTv5NMBBpwbjmlqDHnOelPByo9aYOmaA3ynNSIVmwKcoLgEdKafmFSJhV/DimnfRgDcjHcUg4anDuMc0meR6/ypDFAw2aeuMY7U0jjjpTuMZ70aJgJnBIFNJ7dadgYz3ppAApCGg9aTPPvTf508ZPbpTtpcBpHtTDuGeO9T7ce9ManFgVZG281GzcZBqw8Zc47Gq/2dgSQTWsLaBew8H5etLgkdKgDEMfYVNHMC2DxVSi+g27j1UkU4DPFOxxSMNozWQughPaqczHkCrD5XLVWmnVYy2MGrhG4GbOzx/M1TW96f4/pVO4lEhAHSrMcQeMY6gV2SilHUUb3NMS5Gc5oaQYFZ6vJFgHpUouFOa53T6orcshxmkLHOKrLMpOaGkGOtLkAmL+pppk9e1Ql1PQ0bgDnt2quUROJM5xTCTg0zkD3pwBI4HNFhiBuBigH9KcI2zR5RxwaLoQ0k5+lIxzwOoqykeF6ZqMwk8gdOtCkrgReWS3NDW4YHPSrCxseTUnlA8Uc9tQKEKlGINWxgnJpkkRU5ApgL5wacnzaguzLGV6VHJjacnmmMWHbp2qF5GHbiko3BhC2HxnvS3TAfd5NVw53dMU4K784rZLW7J6CrISeTzS+Zz1/Co2yCeOntUKsdxB6etaOEeg7jriXOCevapba4YEA1TlySOMiljJDfhzQ4e7YaeptA5xzxTRxnFNgkQoOamUAnjmuZqw2OX5FNRNJknnrTpX2pg1S8z951opw5ndiLIIbryKp3LE5A4JqVZ85B4qEpvk29a2iuViexEke4Y709kKgirsNuq4PNSOo9qp1m5WQraGcsftzWhbrtX19KYybSMDinoxXAFZznzIErE60pwMetRbwp68U4ybugrHlZQjEluKqzzlR8venSSHOOlV5j8rZNawi7XYig0bSS7259quQQcZot0DdSMiryqAuK0qVGvdCKsV1jAOBU6jkU0j5s9KeozisG20MkHbmkc4BpxyMdqZL05qVqwMp1BnJHWr8C4j5NV5V5yBVmJwIwBW83eJMVZiS/Lk1TY9cGrbNvXGKqSDPA7UU1fQbIiGYEDgUwsqYz1PQVKDtJUCoyuWxitkyGSo2eCKnU7Rk1Cq5INSry3P5VEi0SpI3pT8Fj83SkjAJ6VMvB4rFuzGI42xnbVQ8nmrchyvFUZGCn3PSiAPQk49akAJqsG59asJ0Hoa0logJ1XnNSAdKi3YHtUinANYSC2o8dKMe/40A4x3pe/H4VACOMn1qlKDk81fGc1VnQ8npWkNGDIY8KeeeetTMwxUSICwJ4xTpDjjqK1avYQ5TluPSrUeOF71VgI49D0q7GuW3EcUpR0sIkfOP61DwCSKmc5GB0qu2QD71zlDAxqGUYwamUhmPv2ps6cZrTZgRjAPWnEgkc1WD4f6VIjhjxzVOPUehKTgVA3A4NTZ496iYf59KURXKz8nrUEjFTx+NWXHOKrzAMSBXRAHuVsEtuHOOgqYL83B/CiJCOD1zTjw4Fat9ETqSbSo46+1IydealU8fWo3YFhx9az1DoEdWAc/SqyhtxHQCp1YKPU1ElcpEq8in9etRo2Ting5HFZMAc8YpEUAGkJ59zUi9KNkDsLtAGegoXGc0mcA8UmevvSsBIp46073poPTIoyfy7VLBgeRzTcAjFOJA7/WmjBwe9MLDCQp9KRfmHFDckU5DjNXbQQ3BzTsf/WpfWg/KM0rgNz2pI+vJNLkYyabuGMU9xk3eg4xmmAgdDSFufWpsFyQMMc1EzfpQretKw7U7WAjzkjPNLjp70Y5pRj86oLCDOeKU4poPOO1OAwfr2pAIw6c1DKOtT57VE3LGqQWKxU559aCQOp61K1VnBOf5VqncG0TpIoHJ60oYE5B4qqFbdg//qqdANucn2ocUiehKVzg1Iq4PNNUjsKUc9qzbHohwp4btTcUqrmpdmh2FIzzSfj0p56U0gAZJqQEbucYNMILVJ0NJ3qkCGrypPT2ph4JGaf2PqaYw9OvrTQC98UnOeKDyDjrRnsaYCil7UzknjpTv4c5osFxi/eqCfIJPTFTDduHNEwyp7VS0ZO6GxyYqQyAj2qkoLOMdD+tWxGDVSiirgJADx+VSZ55FRqihs1J3qHYEO6k8UwjBp+cc9qYR3qUAZxk9qpTOd2B+Bq6wISqJ+8c9OwrSmKVrD1y7dcVJtPtRGOvrUwHBOKcpANA45HNRyEqp4/Cpc4OKhkXeKmOrAqEHjPU9hRhiOTUrBlU/LSqOMn/APVW9yddimS+7p371at+RycU4xbySBz60JGUP8hRKSaBEm3JBzgGpAABnPSoWfBwaTzOBgnNZ2bK0J9wHP601pgD15qsXcHCjjHNPRS3JH0p8q6ivcsxuX+lPJCj+VQq+HwPxpztvwAcelZtaj1JlJPJ/WnHtmmR4IwKeecenvUPcZGx4/pQD8uM0yQfPn0pyDGPWq0sIU5xmoh3BqYnn1phGW4oTGMxzz+FHQ5pxFN6cYqr3EIx596YUJPPSnkc9aTHcjpTTBjU6+tKV705V5p3bOOKL6h0EVQOTTXwOe9PwTg0yRSRzSW+oiMtlRzTFY5IpdrYyajQkH5uTWqSC5cjc1MDxjNQxgEDIwakzispLWwyKZsDNRKgY5B4qaRS1Ise1cH8KpNJB1GsmAPWm7WDZNTMuAcHmm8fhSTAag5yP1pHBx61JwQcdaQ4zz+lF9RkSIRg9zUo+X0oB7elPPQ0m7iRC4H8NIDx1p3G4ioyBkgdata6AyMsS2MUsZIfaehqTy/m9+9OSPvim5IB6nLEZ/So5z71IOp5pr4PBFQtw3K6t8x45qcfMMH8KaEAYY7VJnsBVSYIaoxUykn+lMAzTxkD0A71DBIbIm5eDVZwQcVabnmonTcPWnFjauQEgLz19aYDl+fwpznAFJHgr9a2W1yPIaxZO3NOQ5TJp2zcuAelIM7TxnHei66D3EIwMLwfaohwxp5bceAfeolyOgOPpTWwN2JmJOAtSpwvpUCsSDx0NSoRjtUyBEgB5BoyytyeaFIYZ60MOMg5/rUDJB6UgA70qjIpPXmkP1GkYH9KaOO/409jkZxxTRk4piFbIzj/APXS/nkdKXqSKMDjn8aQAD9KM8YH4Uh6nFNPqRQgHtwc9DSZwOmaTecYxz6UEbcHOR9KfkFxGP5U0c5waeeTgDFMUY5oQtUxCPzFJg4GRzmlOTn0pvIIBHFWhARhs84NIw4+lOY/OOtIc446CgCPf0GO/WnK2Rg+nWm7QGyetICd+Rxiq0HqiQ/c+lUpFHB61a3/AC5JyKqSsdxwBhauCaYxqrgcd+fXNTg7c5GOBxUanIx+lWlwQeMYpyZNrleRFdiRgHjkU0qwjk/3Dk/hVgAZPX8e1Mm4ikHX5Tn8qIyewI9KBwMVIuVHSoiST9amxhcdq8CzetyhA4bgmpAPQ1GkY3Fj+VSZ54wBTfcQ4uFGO9SIVx8x61CygENn6UAgHLenSlFvZgPLsEJXp61CrbuppksjucKPlqNiI8Bevfmmo3At5+b6VA7L5hwaRpCU2J+JpojKg45PejlSAaxDsM1MpChQPxpnlfLmnx4VST1qNNRgRufNBwTtFIASC1KAPvds1n0AQ4PHtSZ+XHWmO+0bvWmI/XI61WoyRhnFDdMCkBI+tL/CfU0tWgEx8maM/L1xQTheTTdwYU0mwFHPHalGKYrA/nTvu8CiwxWPTHWl6DBpnQ5NKRk0rX0AeANo9aQnJAFGMcZoX+dNbAGSGwP0p1IoyQc/hTgOdx9aGhDlIXrTgMuxximZJapFBZTjpVavcGOCYBz+dJvIGOtCsSOegpAcHOPpU+aEIDgGm+pPSgngH1pMjOe1S1djBucUDgYxnNI3zCjqaNgDHQH1oOOfalOM00nmlfoAq+45pyHOSeKRT6dqUc5q1JpagK/ykH3pYiAzMeaAAc7jyO1IAACc96vZXQhM5c0w8uTTgCVNMGVJrNN7DEJJNJj5vahgcZo+6nSmrAJkAkUgOMg96MZFNU9z2poCXkD60E4XrzSFvlHrScHAzUpDFJynFIDgUfWk3U9mIAT0xSHpijcKbkjmkMMZGKavU5pR6mkB56VonrqAp9qB1NKp4JNJgdqTYBj5falIzjFA756U0HjrQApO0009RSk9z3ooEIR0prZB9/WpGXIzTGHNNXGNBwcUoyKa3DUgY7qdm1cCYYxgmgHjHakByaXv1qQBl9KM4H1pCTuz1pW5WgXQYeeTURPXipH7elRtwapIdhoPHFNVjn1oIA4FNDbc4q7CJBz9afnPFRA7cHHXpTwTnJ6UmMlHTBp33TxUe7PSnDPWoaAeDzxTx1Ipq9c9qkXg5qfQAXhqFGWNKRg0Acmle4C+vNLnAHHWmggnk8U4kEYJoEIRgZpp54qQnAxTCVK+4ppPoAhAyOPancnnHakVs8HmgdcE8VrGKtqAu3IzSmMhQzYIoDAHPUd6V23DHaqskiRp27gccVG6fNjse5p0hGRikZ93XqKnmQ9Rj2qKR3qvdW6AqYzg+1WGmAGSKrPIc5HetFNrRC5RizMhxIOB3qVZklXg8ionYMvSoDEA3HHvTi11Q3cfLOQNuO9Z87Eodx61eCY61HJEpNVCSiw1ZiMMdO1SpdFD1q1LaoMnpxwKz5UCMe5FdcXGaCN0aKXKkDPelIVhnFZcecDFadoC2C1TKny7D3HCMEnFKbd27VorCvUdalwBwBXLKr2CxmpYsRkHGamW0A4NXars7btoHFR7SUgtbcjaNE4NPVFHQZFNlwRToj8oz0pvYLEgUflS7R7Ume9Bbjio1GhwUY68UdBTFbOR6daU/NRYBwGaCnNC8HBqQfSpvYACjHT61UKjzfaroqIxc+9ClYGRYGPemGNSCCM1NtyCKac1pF6i3K32ZM5wKmVAFIHSg+opAcnAptt7jKskPz7qb5Kk4K5q7sB78UxQN2a0U2hNJkBtQy4xVSWzKds1rilKhhg0lWkmLlMRUkA9u9XLeR1HzVbEQA6c0hQBTkc1XtbhqiC4cMh9azxuIwBnHY1qmMEE/pTBCoOcVVOoog9TPWMlgPzqyiFSDVgRqSfakOA2aHU5gsB4AI7UjN7U48gnv60xdxwcVNxiythOOtLHyvvTZFyvvTVYjHr3otdASPwKjDjGQaRice3pUYJ5qoqwDnIPTqO9QN8y4P508EDIzUROGPP1rWOr1J2EiJjOKvxjKfWqsQDkHFXEGF46VnUY1otSN+DxTom2/wCNJKvvSIeB3FZrVDJt1Eg+XmmE+9OzleamwFOUHooyRT4wQtK4AJ9adG4OBWrbsIaUIzVV1I4HSrjkAYJqqQMdeKcGFiFQcHB5p6jcSe1MU5brxTt2DwOK2dyR+7acY5Ap6csOagUlieOlSodxPrUtaDLSjA461Io4FRJyOvNWEOMelc8tCrIaRkZrPlxv5H0rRdsAis+b7x/zmrp7g9hka8kDt+tWlTIwO1QR8D+VWoxhetXNitoOA7fpUvbFRZwckUock/1rF3YyY9cHvTgelMUEgZ/CndagBx4GaryfMD3NWDUD46GnBag7lQEq3oD61ZxuQg9agOS1WRtHat5PQlECja/FXopDgCqsmAc44qWKTGMdc1L95aDLbn5feoXHy9e9PZiUzmoiSVJFZtWegyuqnzM5qR8snNRxvmQj0qZhx705X0uC20KhQZzimjCEEdasbe9RuuVOOKpSAj875sdfWlJ3DmmR/KSCKkYYGP0qnYLkDL/D27VEQFPP4VYfODjrUYXPX8KtPS4EYHPAoIBbJ7etTFQDxVeQls+lNO4noPJzx29KY5A470J8uc09kOM0/UNxU9e9AGT0pqAk5P4VKBzUvQe49DzgVLxio0XFSd81mw6COMEU4GkI454NOXGME0nawAcMBmkwc9aeO9IOuKm9hjuh60o6cd6bxigtxQIRs0A5HFK33aRCOuKOgbDMFmOacFAxzTuPTNNPQGqu2AE0HHpxTWIAGeoppY7eD1p2AVyNp5xioScCnuMqc0xeByKqItSVegJ604gdaahz1p38QqXuMQ5H1oPY/pQ5468VHuJzQtQJGPNBPHoKbzuHNP8AqO9AaibSBTWbbTiTk4/Gq8mS2KaV2BKpBB5pjE44/GhFx605sdjzT2ER5ITk/SomVR9alPp29KRl6DNUnYZGIwcZ6HrSquKkJxx2pD/ePai7FsxUxnFPzx7elMjGM+3SpguRz0PapYw5xn0py8Cg8Uo6danoAHkdKTsaAccfpS9/akFxpxgYpnGTmnnGcmmYyfQ1SAQ980x+RjPNSMKY3OKaDSwnbk0vQdM4pvfPanZBAFMQ3JFKGwOtNJ79qF5HIp2GO7njikcF4/enHAIpxHy0r2DYqQqd/I471bIDLn9KgUlDwKlDehq3qxa9RhDBuKd0GT3qQ4xmm5BFRcAHX2pcce9NyMYoDZxigYScJx1qlxkg8HP5VfJwMGqjqNxwKum7CaHwkkip8YHPU1FCMH27VMTzUy3H0GkDpjmkbC9s0vcGmucE5oQugwoD25p4jXGKRTmnEg809b2GIBjNMkJxxTyeoprAEihAUGJPTmliU5GfTinOmGzjGKdHnpW7ehGxYCgrzzSIp5xzmhTgGnoxrErQbs+YnvRs5xinbsA07qM9+1K7AkHyjpmlxkCm9qcDjtjNQxjGH0o6dKXrSdCKaFYQjn6UmCPyoJwSaCeeaY9Bp/yKRjmh+TTc4bH61VhMQkEe9OVeabwBingnGQM0MYuQpP8AKlpB8x4604nCntSF1EByKaemacB2pjDHr7ULcCNwTUapzj1p4fGeKRTnkcVohNJssIB2FSYyajQE09Saye47BjHf8aQgcDpT+n1ppHFK4W0GkUw5Bz3oY8YxSL9TVJA9RcHB9R0pvGeO9Oyd3+eKb0NMaBfUU/BamJ92nrkZoYCbcH5eKhAG8kipGYjnGTTSc81SExc4c5GPennrxTQRwe/rQScAnrSYIVjg5/yKj37mHtQxOM9qavJBqkrBcmAOfejac5ozj8KQPnGakdhQfm4HWndeKYMlsmpPwzSYugncelIADmlPoKAeuaBlZwASKi4UdaklHzHjmqwyOCevatoq6JLSdM5pSBjOaiViMHtSszDvxStqAYLPyeBTSgIx+opUZcYxwT+dKTwev1p6j3IWO0j5etOT5W5pzANz27VHgIQBye9UhK97E4+7wKUnABFMD/NkjjtTgM8nnNRYfkTdsjrQDx0+tAzilJ7GoB6DG6nPI9KROD1/HFGe4HWhcq3Qc9qfQSHA84zxRjt2oHH17+9IuQTx9KBi8gnjj1ppOTilyT049qTPHIoAXcQSPSgjIHcmlG7JpxY5x6+tIWrGHJHFNByM0uSM56UxgQc+tUgFOCufTvTGG09aeMkCmMPy9KaFfqMVsHr+FO/hOeDTNvOQf/r05SAPbvVAIw6t3HrUfITJPT0qQ43HBxTQpJzVJjbEJ2oTwPSqj9TgcCrEgADZ5XtVVy2ScDPrWkED1VieNAWxkg9x61K3AHNRRsCF4yacZOSR1HpSd2w06jWJPT8OaJWzC46jBAz9KRio4IweoprcwMfRTnNUkG+56g64wO9SAHGOcVAGOfWpA24Dsa+fV3FjsPDBepp28Y471WL5DA9O2KdE2CM0knayAsDO0k847mkEZkYsTgDmnCTOVA4zT144PSp1VhDSQONvGKiEIYZbirEhBwq1FISNqg/Wn6gNUBFJxSqCFLY470u0AD35oIJIXotTd7MY1/mwF9KCNvygU7gNnsKOfvmlp1AjY4G0/jUWSy4XpT2G/OTTHYRRgAdaEMjPTHXHWo1+Zzx0FKCQDz160oxGue5962UUhEigk5JpcgtgU2IkqSe5p+AB7msmMY+c4phypA7mpeAucc1Eeue9VbQLjoxg5PenoPmyaRR8uTTgeKi+twE3ZPSkJO6l7H1poBJOKAHe5604dQaaOeKU8cAUmMOhz3p+cEHvUY+XnrT+jBjVMRIMDnvS5O04pAcnI6UZ3ZGab1d+oiQDKhhTc8EDjvRzgdqR/k5HejZgMY9qQep4px69Kbx0I5rPQYHJJIpVODyKbnjij8aBjgQO9NJ/Ol7UcFqezEhy9OnNP+6M/wA6YPbigj5gCaelwHKBjJ6etBUgZH4U4EbCDTWYBeOtU11EDDC5BqM4ApztkKAefSmN8xA9KVtRiHlc005Kg05hgnHSkByvShMBgzSZABA/Onke1MJwDxTTVwBRkc0oGT70A9MUvt71LdmAfWmkYp3GRQcZovqGwz60hzjntTj60wk5qkMCeMUhbPFIBknNJ0FO1gHjBXFHoDTU61J2odgGv0+tRM23/CkkY5qHk+tXBaC0JwQTUqnsagjHfpUykGpkMcWx+NNxn6UoPJJoHFSnqIhcHpRgZA708gc0Beau4CinMOAKTPPFKW59qjVjEJ496VulI5549KVmyBSAYw496hYd6mJx1qEtyauIiJuSfam5GT6U9umR1qDdj+tarUCXOcVJuNVt+B1qWNhiiSH0LHanKcD61GSR3604HPFZ2AlXp7U9G4qMHr+lOB+U1LQEmSSMnmlHeoS+Bj16UjSbB1pqF0DHiT5smnB1B5P5VSeQcnNMefaOtbKndaiNBpwTgc5qEyhmyBiqcEu7PXNO80hwMdetVKCWgFxpskAUvmblwetVg+M9qjFwucE/jUKLeiQF8yfLwKQPxk1UE3PXin78d6lwklqO2pPnPPeoy3XAqEz4XAHJpmZGB460Km+omEzMRj1qFXwxDHIpCj78GmFCT0rVJBcm3c8/nSF+eaifOMVUkabeAOnrTjC4XSNBnwMnpVeW6UA46ioHkZgFPbrVdVO7mtFSXUbZI8jyH0FV5IC8m0jir0aj6/WpAoOSRzVKpyiaKsdpk5IrQgg2rz2psbhmxirK8YrOpUk9x26kiccetKW6+1JxTWwR+lc+4MTziWx6+lObGKhAO7BHNThdwqmkg3IGB5PXilTOfapmj6UoTtinz6CG45xTtvHSnKBnpTs4PNZ3AjWPnPalcbW4HFBky+AakYDr3od1uMiCk9qlBGKaWGKb0NLcCUHJ5FJgHr0pF5HWn9uaVhEeOaQrwTTnYLUTyZJA79apJjKksjb8DpSwyA8d6fNGNhIFVIwynIrojyyjYXUuyE7eKiyR3zTo5cggigD580mrASLkY/WpguTTAM9sVIBWMpDEI7ConHapc4NMfBXPoaIiGjrjqaYSQTTlIyc96bIeT+lXYCBmO481E77WFSAfvTmo5tuQa2ja9hdCVX/d/hTFZiOvFEPzDHapCuBim9GPUjw2DzmoA7bsEYNSs5Q7R0qEMS3NXBWF1JPMPJoyD070gGc0wDDGiwAy4OR19ahc5+lWMHuPpVd1IbbTixMfE+1M06O4LPt7VXkJxt6GltgVIyckdTV8qa1BMvyEkYp0a4Xmm9aeOABXM9ixfvdsU/GAaRCdp4pTyuahu4itMcZxUMEg3VYlHBqCFfn5HNax+El7kjfMTmq7jFWZAV6VCeh7mqiwexX3ZkpWJA759aVlAII6mm7jkAjk9a36XDyDcScAY+lPiGOTUZb0qaIE5zSkmkK5YVwDxVlDxVZEwOetTIQRXLNFJj25FZ1xkNkd60iP/r1m3ec1VLcb0EicY5NXEYEAVmpgD3NWocrk/lWs4KwJ6Fvn86ULtAJpvzE1J1xmudgOThacTgdKVFwDS7QDms29QEHSoZPUCpeeKiPOT0q42uFyD7rAY61PFgjPWoHXDZqW3bHGK1eqJFccnI4oDgZHT3qZ8FTxmoRt3Y6elT0BEq5bBNSEfKcUIuVPpTjgA1nN6lWK2wA5pxPU0MAG6000IBmeR60dQQaGcA0ZyOlUPYjkAB4pv3uaHXJ600nHAxVIQxmAJOaFOc0xuTz0pV6/SrtoArc/0pu3K80/qPem44weBTQETHawCj6inBiwIpkgB6HB9adGMc+tXfQWtyUDr2FPXjimA881Ivr1rNjFUgE08nkVHnBJpwA21DAcTQM/n2owKQHnHOKOgWHE4JpQQQP50hHIFL1NIBRjtQV3Z4pV6npT16cjNLYOg0AYwaABxgUvJzijI9OaEIQnkCmtwOe1OxwPamSNhSKFuPoQs+SfaiN/m55prEAHPBpkeQa2tpqMtSMCOKqOSOe1SeZ2qOUHB9qI6MXQdE4OBU696qQ9ff3q4Og44pT0YA3IHpUQU+nFTAcY4pMfLxUXAbgAe9A6Go2yPeljPHqadguOJ6mocnqaezckUjH3/CqDYA3OMUhbkDimg84pWOTmnYBNwzzR1JwKQAHvzSY4z3phsJn5ue1Pz0OKYgzzT+BnihgOUYPNSg+lRpz2qTt9KhgKeTkcEUc44pueaUPg4pWCwAAZ9qOD0pG547ULwCKBoXP8qbjFO9KYxwelMT1HHP4CoiMnipASRwelMb73AoQEYODgnApc9OO9MbjntRngdq0sApPalB96byTn9Kci85PTPSgRIMZye/pTs5BpvSn7sDmoYyIqOvemjpxUh5PIpjYBqkAO+BkmofOK8d+1PmJ2YUGoVAzVRQvQljz3NTAAA8U1BxxTicDPUVLeo+oZzxmmlQTmovNw59B1qUNvHT6U7NBdCqQeBS8dSaMYGBSA8+9SCFzhiTTZAcU453YxxTsYXH6UXtqLyIwOBmm43c9qeRjilQcU7jEVeD7U1j61MT8vAqBzjpTWr1DYhkTJH1pnO/2qdcscZppULwO3arUugh6ggUoAGaUHilHAxWYxm3nGakHA4pAeelPH0oYhOCoFL0FNLYPWgEZ5pdAuLn5s009aXo3tTWHOOlC3AcTweOveowf8mnN1xTW6cCqQDQfn4GTTW4OKcMHikIyTVANyOPSpFOBjHIppXH/1qXPIIoAePvUrAhQKavBxmpCSR9agaI1P50O3GCfyp2OBUT5x0prViIMc/wA/engYPt3zSKQxzjJFSKNy8/lWkmLQkR/kxUqnI561WHDYqUHBrOSGiTjFRyHHSn59B1qJj69PSktxjASQDmnDPUUqDuBTyQB9apsRGBkdaTgDdTs8Y70HpQFrDYwOnanZ4ODTQe9KTwTjigFoNY5wAfrUQPzZzxTweDUfrz1rRIdyYMGHNLwe+aYo2g5pQ2AePap9BCEY696EUY46mlyA2MUoIxRdoLsVkyOTTakI9vzpOppJ6A9xF4IHbFSjg1GBg9OtSD3ORmpYxD3xUQ6H3qYsKgYlTx0pxERSkk5z0qs7EH+tWXJxk1WlOCOtbwE0rkkTHAJ6VMw3oc85qCF+cHj8KsKefrSktQIzx9P5U3YDyB9KmYdTjPpTMkGkmOyZHyR04700nPuKmY5HA5pqru+o600xXIgh2YJzx09aljyGwOB6U7ABB/Kk65OeabdwskiYNz6GkznGBTBwTnpT1OSCT+NZ2He4Dg5NAIxz+lDMBTNxxn3osFx3IOfSnA8ctjNJ2PoaQMCue9OzsGt7jtvPFNPAHHWnBu/rTCckjtSQX7Dx79Palz8uc89qjyQ/+eaduBOMfQ0NBdJhjIznimAZPX/69OPynr+NIFweOlNAricEHIBHpSNhg2emadnrx9KacAZ9KZJHg49DTFypI71ITnIPHtUbMPbirQbDgwLA8cfpSswCBsnGOKYCpOKHb5Tina42yFuuR93PaqrNwTjj1qw5Y8mod4DHg4PatYgSIcjj05pfuqMn6Ui8A+h7VMQCCKGLbQjkBKZXqOhpm3bA47lTmpFGJD060SLmBmGMYNF7aBdrQ9KB4wOvrSc8AcmnAccc0dDnvXz7d9Cg8vC5IpyABaOV5PTtShNx3E4FTbogHQnLkdqsFs4A/E1Eqge1SoQBk8CnH3lYVrDWbkhTzULv5YC9WzUrEEsRwMdaiiiJfc54oitdQHqv8TH6Uq8/Mc0v32x2FI5BO30qWrsAABfJ6CkbDc/wik3LjA60yXPlYH1peQyN5V8zaKRwZAMDp6VAI2EgNXPuxYFU0orcCvjAIxzSEZGaecAZ75pp9KTb0uMcDkCgtk8Uh44PWgZIHale4CsRxio2OT7VIB8pJpu3k00+gCrzgU/vjHNMX2pecZ71LsAZwaB09KFGck0vXk0WsAfd5HWlHJyaQ44NISSfrSQDsDrTh83JpB9cYFC4JJ6D0p6gSJ0OaAMH+tIOWAp42nPPFPWwByFpMErnrxRnd8vag+lG4hpORTSOPenA9u1NNQ7lCZ5zjFJ2/pSn2pvXjvT6APAyPrT8AD3xTOgpSTjigGOB+alAG7PtTMfLyaVVyM0JdxD8bsnNRP14p54Xg8mm/wAPFUkragJ3z3FJnJJ/ClGW47UbMGm5degDSDzmlX7vf3p2M9aF4U4NRcBnUkYpvTI/Wn9elNx83NCWoDeRgYzSjg9acRTSOcU2rD6i8Zzimt14NLnnHekY0rMBh5OMUEGnNjHT8abjjnrVK9hDMUDO3GMUvvQxyMCqaYxyDA/xoJAHB6UDpTWNIBuARmm4GMCmSMegoQHqc1pYCRR9BUi8cetMPt0pCxzUWAkoGBUW4lvQ04nAJFFgDvyaUYAz2phbIpc8UxCk4o3Z601j2puQOlCQxxJ/AUc5603PB96TdtHFNIQpbBqEuMnnimzS+X3qk1xlvT1rSFNsC5u5qAnOT2qLzznHrSeYDndWsabQAzE8VYgPy+hqn5nz81JHJtzzVuGgF5pe1PWTj3rPRt5PtU/mYrFwAuK3HtSmUAdarLKAtN3MTwKlQ7oY/wA0lj7VIzhkHPFQrEetTRpgDNVdLYREIiz/AC1Klplue3WrEagE44p2ecjik6zvcLEAtkRse9SiKMMO9SFR161EW5PtUczmFkiGeIFuDxVGWBgc5rQkOQT2qCZlMZNawk4vQLEEcLFc7qkhUl9vPFRozBOvFPt5fmIHWtpXfULGlDBGpywoKqWOw9Ogpm4+WMn8aiV5FbAxya5+VsRMUG7kUx1Udqkd+Qe9Qs4LH2rPlYyKRBmq7qF5wcVM71A8obAzzWkbgV5tobNV93ORyamn7VXUAnIOMV1QWgFiI9PWphx3qKMfKT3qZD831rKQ2CJ8+ccVcU8VAoINTKcDC45rOWoloSA5XpUZJz0p65z7U4phs1CaGyJQc+lTAEfWgKCc08UnK4C7cHmkyCM96G54pQAoqdAGMQCT3qMlmGVpZweMdKcDtjHrVKyQtyJBtOT1qfqM1AzjPPWplJ2jBpy2HsBFNwetSE8UzGKhBYkjAxSnJGe3rSKcCgnC9aVwsV58nO3rTIxk1IxFNA5681tFuKsA2YNtz+lVEfD7cVdLAjrxVMMPO6VUNhPcs44GOtLxkcU7b0pVX1qLjHjpQc7eKAP5UHoKgBB90+tIw4pw7daaetO+ohozimOvfpUhGB/OmSE7apMHsVZAd429ahuF2YOatMAuSarXPJCk/jW8LuwrCwSYXmka5BOF696jH3MVGsfJNa8qvdhe2hN94gmmuecinqoJxmnsgHQdqlvW4dBqMQvPWmLhmp6jccHineWqnNVewCkZ6daqyfeyeKvcYJqnMNz9OBUQ1YyrcsD0x2psMuHwTTpkwpIqvE2ZCP4q6YJNWJvqbKt8oJp+4E1BA3y4NWQuTx0rlmuVljhwBTgMgigUcfnWIivLnjFRJkyccj1qeTG40xcFq0T0EPk2hRVQ4yTjp6VZk6ZqoX3sRj5e9XTWgDGfHbmowwLEHoBxTpSAcjtUMYy2fzroS0IbsyZdokHrU4Gen51AVw+ccCp04HXNTPyKROCQOe9OhYljSADqamj2jkfSueWiKHc1n3uCeBWgRnnODVOeMNnNKm9biexVij5+lXIlwTmo4+Cc96erfNx61rK7GWRweB+NOGMk9qiRjg5pwye/FYtDLCnindjUa4GRS5z9aza6gHUk+tRu3pTzyaYw3Dk81UREEg56cU6L8qSXA5NLCe5rWN2tBE7H5CKroDnn9anwKr5w3XnsKIaMLXLqlR0/SkOCOaI2G3pzikIzWUtxpDHUU0jmlbk4NGc4oSAjZOaQnjilZsZGKYxyKpAMOeTUT5OcmpsDA4qJwOBVxAj2flTgvUU4+h70dGqrgIBzxSPyvWnNwcg80zIH5UICIqAaemOxpp+YdcUL3FadBXuSA5Bp68io84p444qLDH4z71IMYqPPGacnJwKhoOo/gjHSm4HJxSjGaXjv0qXoAcge1IBjnNOOD+FKPz9aLhuA+lSg5XrUWM9OtSDnpUsBSAKbwAafgHFMIwTzxQApHBHaopAAD9etPz2ph4BpoCtIm7NNC7frU5Gaj4DcmtU+gAFDKTTXU447VLwQaYRkdOKFICsCVcH86uIxCnPWoBH8xGO1S9xTlqFtSUc/SnbeKYvsfwp69MVkDGbe3fNNxtGD3qUjvTD0z+lNMNiAgBsnNIeeOlOIwck0vVau4EaLzkdKGwRS0hHQjvTDoNU4HvQOhFA9cdaUAgnNMBBwxHandRn+dGOfp1pV5JyOKGwHIdoxmpeNvNQBwG5oZi5wDjnkUrXDoMlZi4VR+NTIpxycmjyxnPWngc+9DatZCGuB+NKpOBmlwMn2pMdakfUTPemMeTT/AGppHTNNAA456U1xk04cNjtSOMgcULcRExGKap7Gnlvbr2piHnvz3rRbBuPYDtSKeeuKfjOKYo65pICTg4NPblOBz65pq4IJ70o59hUBqIp5600gUp69aU4OcVQyKTbtx60xI8HmpCmWBp+BjIp3shAg7frSMpI6/WjIFOHPJ6VPUZRaMo2O5qZG29KmZFJyaURgEHtWjqJolJIazYXkVGj7jxUkuQM8UyJQMipVrD6k3GM4yaCe1ICQOaC3HNSF9RCcGlB/KoA2XKjqKlGQOadrBe4/8KYyjOaeOtIfUUrjY0CmNg54p+3k03uKpC6AMjntSxnKk5pGU44/Klz8uaNAFAyc0/OOB+NMGduKdnK/TvUjIpH4PemISzEmnP1xinIoA9qu9kIcTyAelJ296RsHmkPAGOtJIe4rCoz0/rUhPHqKaQWOc0IQ1SORnijqAMU7HIFJnnjqKoLi89+tJ2zinZyDUZIBwBmgLCBskYNTr0GelV1I3An8qmBytEkJD847VEyja3GRUm4elIO4qbsZAq9eKkVRjrmjjcetMyBnFWAZweOlKjbsZppOee9LGP8A9dNrQL2LAP8A9aomBzUoHFKRk1lewEfA7UM3FKVphx170wuHHr+dIfu4H50Hr9KOgIz+NUPqIMA9Pwpw74HXtTOvanL9elPoSIQAcZqFmIOBUr5PSmKFPzevemhvQAcjkU8EEU0kdOgPagEEcUMEPPByOKF64NNJB5PWnqMHIpbBYecljSAdqecYz2NMJHPrUAL0+hpR+VQl/elD9MU+VjJM8c9agbp71KTkZqMJ82SacREToR0+lQuDyc/nVuQZ6VVcgrlu/rWsWw6DIxzkHOKtKAT6Huaqoc5JIAJ/KraqCMfpTmJADuz9aTYA3NOAAGemaGAHJFZpj0GEYJGKF2kk9qRsHnHekXhxz0qxIUjJ460vGOBSA5Oe9ITnOT07UrBe4pbPOM0qkDNRgc4z+FSoAoI6UPYNWxJCQOAelJjp7U9hntQANgA/Olceop5HH5004De9O7c0mMjOelJCuAyCfSgqQMUY7Y4oxlQMnFMHsMfIz/I0pJ3UMuT1yKEUZ6U29AFAJY8Ug4yMfrTwcOR/k0uOeBzSuFuqGMcr/nimsMjHT2qY4J9aif5Bxyf50ILdyNumO/TNV/XPOe1TPyMjqKgAyfm5xWsRWuyVSpbA6daa2VHOMH9KkXnvk1FM2Mkcn0prVjtcgZhzz06VFgEjOc+nenHBHI6dqTnaP5k1ulZCuSqvAOORxUq4HB61GCSc9yafnYCRxnH4Vmx7bDHG046gUH5rZznsaa2Dltx5pmD5cgI52mqS0FbW56kpCqR3NNBABz1pOfWgcHOK+bum7F2Hg7uvTHSnhgRjtTAcrxThhfrTuIlAB4zx2pRg8DrTc9KkXAHvS3eiAQ5wB2zShCc5Py9aUkAU1jwBSd76iGj5c7aaqdSfrSu2CAtG4BNueatJpaAC43ZpXAILdhSZwtG35TmsvIZHs43fpQck47UgbzGx2p7t0C9uKNbAQt1pq4ByalkAUVFtJPPHemthh1k9qUD5j6UAc8U7pkDr600+oCY7ChhQOtGOck0rAAwOhoA+Xk80gPzZo6kk0MEKOBmgYpFGev4UZIPNLcA7UDPU0gyTjtTj6ULQYdcU/ge9M6mkJ7frT0sBKDlhgUpPzVGhyvH86kC5TOKtoQ7bximsSSB+dKT8vvTCQOlS2nogFIwOelR7ucU5m3VGB3qWkhjiacOhxUZYUqnqaLOwaEm45pB15oGOxoJPQUuoDsZal6cDpTO1KDjvwKrRsQvUdaaBxQx3fSkzzTbuA4DilLGmgnn1oB6/Sl5BYduBGDSA7RTQPSildAOzgHFNOPzoJ6elNJ7UeQC5496RjwDnpRn86aTzzV6MAByT7UD72ccUADP1oHBNJbjAjOfTrSYycUdDSGgNhcUgXvRnnmgNxiqVwEyR9aaSTlqU800/lSATA280KAAT+lI3UAUEGr6CFB45pmeelBz2pMEscmhIAU9z1p3Qd6aDxQx6DtVb6BYAc85pw6elNA5pc4yfwotdaDEdqaMEnjg0jOueaj87PAqlF2ESk84FRNKM80hkNV5mOKqMdbAyC7k3HHYVRMjAHjOOlTsC5yetMEJZsDpXZGKihXGpLk5weKd5pLcA89KuR6f3xzU62SjoKzdaC2HqUYYyxywwaJEdG5Faqwqq8VDPFkEVmq15BYrwfMvFWVjz1qG3Uo201a6EClN66BoCwingYPFHanA8DjisbsYoFPXGDTcHr2qVQNue9Q2AoGFx1zSgHp2NOUcZNO25/PrU37gG0dAM0piyOF5pwGCKlB6+ppX6iKjw5XHfFVZYMJj2rSxk1DMoxmqjNoDO8vaPrT7aACTPapjFlc/kafGpGMVt7S0W0BI0AKkdBVfKxEbhzV7B2c9KrtCrNk81lGfcGV2bfyPrUExYpwOT2q1NHgfLUIU7ea1UuwFfBCc9arYw5z3q/Im8VUaLBq4yuGhFIMjHamLEN1TnnjvQo6DFXzOwDEXr7VKgAP0p2zAoAGM1LdwJMZHFSoOKiXKgcVOgB/lWcthj06U4/d5o4FQyyHGB2rOKcnZA2OVsinB+OvOKpiRs4qYZ2+9bSghIXzCGp7TYHHXvVYqQ2asqgIyetKUUguI0m480jyAL+FOEYBpkiCp0GVXfLgVcRsJ9ahEAz05qbyzgDtVSlFqwgaUd6dvGAf0pjRmk2EDGanSwyQOMUeYMVGqGnhBwO1JpBsRM24nAqMyEHB7VaKgZzUUsIZsjiqjJCIHbiiGL5807ysHmpo0IbGeKpySWg2OA+bgdKd0OKBkHIoIwetZMQEndR17U0nnFLnGaEgFYcDPakGdvNBIJo78UhjfrSNkg0pBJpHPy1VgKjvuJyOBUEp3c4zU8n3Ccciqpc9K6qauIRSC2KkYg49agj+8R1+lWYoju56EVpPuJbEeDkYqbdjr1NPCjd8oPAqFwScHr61G7sAm4k5/SpHbcOlRO21cg1H5wx6j0qlFtDRP5hGQTSOuUz3pkbK5PWpT93jtUtW2AoSZywI4qBI/3ntVuRSwpoXAGRW8JaEMkQlSPQ9BVyPIUc1WjQ8HPFTIc8dqxqblllTnJOaco45pq4p4x061zMOpDIAckjiqrtsfHQ9qty4C47VRlwr8/hW1PUTLGCyc9Kr7cEmrEb7lPHFQSyDJApwT5rC03IJEYuSBxTQoIHrVxMFPeoHyrfjWqb2YWImJPtTgT9RRgAe9EftiqvoGpaXlalTpjPNRxJgZqygAFc83bYY8DiqN2xAOPwq8SMEVRu8+Wamn8Q3sVY3Ynk9O1TI+WPpVeLleehqcAqc966JJIV9Cfkd6niBxzUMZ3DnpVhTwQDWMmMf14pSwAAzSD2puzPXrWXWwMdznNNbmpB1xTWwB0paXGVZ+ePSkhfjnNPcjbyagj+VjXRFXTJLjfKMk1SJYy5Aqwx3jJ/KoGXDYHHpSjowZdt2bA3etTE81Xtx8vJqUd6zmrMpXZHLy1Jztp744zUe6pWoug0nbTCwz7U8jPJpMD8KtWAYWx24FMPPPSnnJPpTScmqQaCZ4JqMZLZp7HsOlNRBgk1SAHHzcVEBk9anIB9fpUfQnFOLCxHKW6UwHBJzTnyaQDGMelaWsIeDT1Py9KiLBV60quKVhosZGMU5RzkVAG3VKpPHas2raASk57UuOmaT+dDAk57VmAoOB/OnAYNNXsaf1GaTAVcYx3p45/KmqcU4GkIXI6d6Ycbs0pOKYTyKEMOA1MbnJx1pdwxSZ4pj6XI2+8MVEeD+NOk4bpTGOSCelbRsJ2FLnpinqeMY5pu0de9Knf1pMY4jJoPCin9ev503PT1qRAucVICcU1MZpw+tKQCuOAP5U3PanEflTcgc96SAjbOODTAflwaexHJ71EzZzx171otQuByeKVgNvIpAcdKVuRimwYwA7QD1pw6/Wl5zSZwelG4IRl4NPQACmk88VEznNCTYCscyccip1XA645qKEZGTUoPXIpy7BYdk8+lKCM/WmdyKQcYAqVsBKM80wnn0FKCepPSmnhvpSQbiqM96RvSlU4B96CQc+ooDzGjg57Ur8CkHDdMU5uB60wK7HH1pgIHPenPxxUG4hyOvPJrWKE9y2D8oOab/FmhSMYo6HrU+QxwOAQevtTx09zUYHORTgTnrSaAOc0u4c0nQ88CkIyDQ0g6jGYswwT9KlX7gNR7ecjjmpE4HNEthCDoQaVQCOBxSNxyKVRxikPpYdgYpO1LkDjsaU56/lSC5Cx654NNBGKc+DmkAAzgc1SAkHrSZyTjpmhcY9KMdfSkAgABJPFIxJNKR+RpACD04pgOTAOBSscDikGd9Drk570uoDC2DxTlyRnv3qJjyc8GpQQFxVNaCEbpTV+6T2pxFNwB3/KgBw5p38OBSL0zS4wBjpUjG4BbpSk4APejODkdKRvu9KAtoJjOMnvxSYwMAUE5JxS4wCaoQhHUU3bwM9Kkx8tJjpmhMBuNvf6Uz1PrUpBPFRng9KaYASR9c1GASTx0708+1AUECmtB7kXPOamjJI6UxhinD5R7027oXQkwBTSRimM2ATRnrzU2DqJuyeDTACXpR69qF2j61pcW4/bkY7inRjHNNzlxUyjqcc1DdkA8Y65pCecdqAcHNHQH3rMpCHntTPcCnEYHNNGKaENB+f6dqaemAOtL3IxTVHH1rQY4n16UKML0owCOBTWO3tQLcbLgjnoOtKvA9vWmDlcnOacThDkVdugvMSRQy4H50itgEH8DThyMd+1ROpCkjOKa10GTAqrEVIMFeOPaqBZlJP5Vat5MjBqZRtqC1JshFwarPId3HQVNK+Qe/rUHVScfh6UQXVgxhY8Z6U+Jt2efzpgB3bSM5p0YYMCV4NaNaCSTLfaoz8v+FBbimSFguAKxSG7Nj3I2cGqb8AYXg1MG9u3SmPIvOcf4VpFWHuQIDuAbkVbiODjGKpBsOD39BVxfugD+dXNEomHPX86hkYBTz9Kkzge9Vm6k8c1nFaj0HJICaeBzmoQSrEjpUgO7g/WqYiTGW570vl4BIJzSLyOeOOtKSCec4qNSrjCMZOORTkOT160hcZAwc0DIbrkd6fQSbJR16UjZC8DJ9Kcpxz3+lJxkelR1Bick9cUAnd04pc88DjtScjHrRcLsQdeeKUjpSEEqSDQpGMUw3Gv1HrQM7eB+tI6gg+1Kmcc8+1PoC3FV+hHFKDx069KaEwORTsZzx9BmjQOgE8gU1jtOMU/jNIRkYIzSQirLjkdagHL4B4HJxVwpnqMfjUJiIbcBz6VtFg3cVAcZycVBKDuIOcdqskkAe1VpTvLEjp0FVT3B7akP8ZzwD0pqrkDjvTjIuSO9EYBGBn6Vr0Fpoh6NjAK/U08jPA79qaxGMilOcDpkVG4J7WIsHOQM9sZprHYrDb1WrKqMZxz6UjhfszkdSpq+azsGraPSu3NIeeDVY6rpx/5iFr+Ey/40n9q6d3v7Xj/AKbL/jXzipTT2f3MssgkcU9TmqZ1XTs/8f8Aa/8Af5f8aF1XTh/y/wBr/wB/l/xp+yn/ACv7hF9TnipUGG45qgNV0wKMajaD/tsv+NK2sab0XULTPr5y/wCNNUqn8r+4C6WwT60DkZJ4xVBtW01cH+0LUn/rsv8AjUcut6eYyFvrbI/6ar/jTVGo73i/xALzURA2xeTmoYLl5Zdznj0rGlvrWS5LfaYfqZBVqDUbNCM3dv8A9/F/xrthQtC9hHSg7sAUsnygDPPes+PVtPCj/iYWuSP+ey/409tV00k/8TC0J/67L/jXBOlUv8L+4omBKE4GM0/dwKpnVNN24+32pP8A12X/ABqNtT09mA+32wA/6bL/AI1XspvdMC6xySTz6UvIU9Kqf2ppwYD7fa4/67L/AI0f2rpxbP2+1x/12X/GpdGe3K/uAuAYGaTdge5qq2q6cel/a8f9Nl/xqM6pp4GRf2uf+uy/41KozvrF/cBdIAPWkJyMVTGp6ecE6ha5/wCuy/40HVdO6i/tc/8AXZf8av2M19l/iMuGgH5TVNdU0/r9vtc/9dl/xpRqmncH7da/9/l/xpOjPt+DEi0CTzQeWx7VVbVdO4xfWuP+uy/403+1NP6/b7X/AL/L/jR7Gpvyv7gLq53cUbsEjvVVdU08db+1/wC/y/40DVNOzk39r9POX/Gl7Kpvyv7hploH9aUDqT1qoNU04DP9oWuf+uy/40o1TTuv9oWv085f8aapT/lf3CLiHg1ICdgHrWf/AGrp3/P/AGv/AH+X/GnjVtO24/tC0/7/AC/41fsqiVlH8BF0kYJNRjB5NVG1bTjx9vtcf9dl/wAaQ6rp2B/p9rx/02X/ABqPZVL/AAv7hosnPag5C4qr/aunn/l/tf8Av8v+NNbVbA8fb7X/AL/L/jSdGon8L+4Cck45NOQ4+lUm1Kwx/wAf1tn/AK7L/jSjVNPPW+th/wBtV/xqvZVP5X9wzRHWl75qiNV0/H/H9bf9/l/xp39rad/z/wBr/wB/l/xrP2NT+V/cBbNPGQM1T/tXTc/8f9p/3+X/ABp41bTc4N/aY/67L/jVqlUeri/uEWSpPajGO1Qf2rphIH9o2mP+u6/4006ppp66jaY/67r/AI0exqb2f3CJmUgjnrS8jIxVU6ppu7/kIWmP+u6/409tT0zHGo2hP/Xdf8aXsZ2vyv7guWNuB/OkA3dKrnVdNx/yELU/9t1/xpBqumg/8hC0/wC/6/40exm/sv7hlgdD/Kkwe4qD+1dN/wCgha/9/l/xpp1XTtpH2+1/7/L/AI0lTqdYv7gLH06UhHXP4VXGq6dj/j/teP8Apsv+NNOqadj/AI/7X/v8v+NP2dR7xf3AWB0NOANUxqmnZ/4/rYf9tl/xpTqun/8AP/a/9/V/xqvZTts/uAt8Y96QCqR1XT+MX1t/39X/ABpTqunkf8f1t/39X/Gj2FT+VjLe0flRjAqn/a9h/wA/tt/39X/GkbVrDHF9bf8Af1f8afsanVMC2SBTCcr71UfVLA9L23/7+r/jSHVLE/8AL7b/APf1f8afsJroBcxkZFBHI5qmNUsFGPttv/39X/Gmtq1izY+12+P+uop+xn2f3AXWUikPyn2qi2rWRPF3Bj/rqtNbV7IjH2uD/v4Kr2FTomIvYz0NAHJz0FURqtltz9rgz/10FKdVsT/y9wZ/66Cj2NRdGBbZuMAfhUDyHBxjNVJNStOcXcB/7aCq32+33f8AHzDj/roK3p4d7tCZbkkI5zxTFfnnpVaW9tiP+PqA/SQf41CLuAnAuYhnv5grojHo0w1RpLJlvpTZWBaooZ7Lq15b8/8ATVf8amFxp2cm8t/+/q/41nJ2ekX9waldvmOAMVbtbf8AiPJ+tQteWQbK3dv/AN/F/wAamXUbID/j8g/7+L/jUTlUaso/gxovdOKQnniqv9pWQ/5fLf8A7+r/AI0h1GyP/L5b/wDf1f8AGuX2U+z+4Zaz81RkBjVc39kR/wAflv8A9/V/xpft1jj/AI/bf/v6v+NV7KS6P7gLCpinhBt57VXGoWHe8tv+/q/407+0bDH/AB+2/wD39X/Gk6c+z+4CUAjipRjGMVUOpWPUXtv/AN/V/wAaUalYHre23/f1f8aTpz/lf3AW8cVIMgYxVE6nYg/8ftv/AN/V/wAacNUsNw/063/7+r/jU+yqdn9wMvqeMdKkUfNk1n/2nYZ/4/7b6+av+NTLqund761/7/L/AI0eznf4X9waFwAbj+lKCAT3NVG1XTecaha/9/l/xoTVNN5zqFr/AN/l/wAamNKp/K/uEWzycn0qCc5xUR1fTx/y/wBr/wB/l/xqF9U09j/x/W2P+uy/41UaU39l/cGhYyCo6+4oRcMPeqo1HT8j/TrX/v8AL/jUn9qad/z/AFqf+2y/41ThPpF/cMuk/LiomY5wOlV/7U0/H/H/AG3/AH+X/GmNqdhn/j+tv+/y/wCNZqlUv8L+4RZb1qtIRwKadTsMf8f1sf8Atqv+NVn1CxPS8t/+/q/41oqU77Md0WyuEOaruRiozqNlgf6Zb/8Af1f8age/s88XcH4SD/GrjSn2YMfs3H8acB82R2FVRqFqpyLmH/vsUC9tc83UPv8AvBW6hK2wi6DilOO/Sqf2+16/aoM/9dB/jR/aFrkf6VDx/wBNBU+zn2GXVap4zx1rO+32h/5eoMf9dF/xqRNQssHN3B/38H+NTKnLswL/AJg71FKeTiqrXtnnIvLf/v6v+NPN9ZFf+Py3z/11X/GpjSktbMTYKp3g96sqQc1TN7Z5GLy3/wC/i/41Il9Yg83lv/39X/Gm4TfT8ALXBOMVKBiqa6jYA/8AH7b/APf1f8acdSsd+fttv/39X/GsnTqPSz+4C13NMcZNQtqVgRj7bbf9/V/xph1CwAx9tt/+/q/40KlPs/uHcsDkipu3tWa2o2YOReW//f0f41OuqWRH/H5bD6yr/jVSpTavb8ALOQc00Dv2qudRsP8An9tv+/q/40g1GxH/AC+23/f1f8alU5/ysRYPymlB3HiqrahYt/y+23/f1f8AGlXUbEdL23HH/PVf8afs59n9wFpgT9O9NJPbpUJ1Kwzj7bb4P/TVf8ajbUrJeReW5/7ar/jSVKfZ/cBbwMcjrSYH0qqNSss/8ftuP+2q/wCNL/aVjjH223/7+r/jT9lP+V/cMsdKSq51Cx/5/Lf/AL+r/jSHUbHH/H5b/wDf1f8AGn7KfZiLPTtSZyarHUbLbxeW/wD39X/GoxqNnn/j7g/7+D/Gmqc/5WBcLYpvmqWxUB1CyK/8flv/AN/V/wAap/a7ZZf+PqDHf94v+NVGi3umGiNgDgmopM5zUKajZYwby3/7+r/jQ+oWJXH2y3/7+r/jU+ymnomDDHJyOtUJ28o8+tWH1C0Xn7VAR6CQGsm6vYJ5B++TGf7wrqo05dUJt9C5C26Q4HB61eUMMkVlW9zbJgm4i/FxV4X9pj/j6h/7+D/GnVi76JgiwWwwOKa7gnBqBr60yP8ASof+/g/xpGvbM4/0qHP++KzjCV7tMewyYKWxngVCsYzk8j3olubZiCLiHk8/vBTTc2pB/wBIiB/3xW8YOxJcjRcfLwTSt39KpR3Vuv8Ay8RD/gYqX7balh/pEP8A32KiVN32ZSFJ5x3FLtLDGaT7Vad7mH/v4Kj+02xbP2mHH++KEpdmTbQtxrgYqboMiqiXtqCM3MP4yCpDfWh/5eoP+/gqJU5vWxRbQ561InTOKqLe2QHF3AB/10H+NSDUbLj/AEqD/v4P8aycJ32Yx8vXms2df3mauSX1mV/4+4Cf+ug/xrPnurbPE8R+jitaUZ31REloW42yvHSmODkY71FDeWoXm5iB/wB8U6S7tcZW5h/BxVuMr7Buh+do559qrSTZbgd6jkubcgkTxZ9mFQi4hBJM0Z/4EK1jDyFclyccevFSL+7IPpUH2qDP+tT8WFKLmAc+fH/30OavkfYNTTSTI5qdXJGBWXHd2w58+Mf8DFWY7606m5h/FxXNOi76IabL+eME9qpXr5AAp639oAc3MP8A38FUbi8t2bieIj/fFTTpSvqhsEfgAHnFTodzc9DVBJrYP/rowPXcKsLd26t/r4v++xXRKm3shJGlHjHFSR9eaorfW23m5hH/AAMVJHfWu7JuYRj/AKaCueVKV72KNAAAe9Icjoah+32WP+PuDP8A10H+NIb+zP8Ay+W//fxf8ay9nPsxFkD5qGAC1WF/Z9PtcH/fxf8AGl+32RX/AI+7c/8AbVf8aXs532/AY2UAZ9QaiXO45FI97Z9rmE/9tB/jUa3druybqH/v4P8AGt4wlbVMTJXDLiossTnPIp73lmwH+lwcf9NB/jVb7Rag4+0w4/66CqSl2f3CL8YO3mrUfzCqcV7ZAc3cHP8A00H+NSm/shwLuD8JV/xrCUJt7MolcZBApm0DgVG1/ZjP+lwfhIv+NM+32mP+PqD/AL+D/Gkqc+zBEhJ7UD6/nVc3lpj/AI+of+/g/wAaQX1ox/4+ofxcVfs5dvwDcnx3xUbHPFM+3WnT7TD/AN/BUZvLToLmHn/bFCpy7MCQkE59KUcnPYVCbu17XEP/AH2KT7XbZ/4+Yef9sVpyStsBY9j1qE88Cj7ZbY/4+Yf+/gqJry2A/wCPiI/8DFEYyT2Ac74HPGahL5PFNM1u2C1xF/32KZ5tuD8txF/30K1UW+gr6EwXcB3NPCkdKiF1AG/18XH+2Kf9qt8c3EWf98VLjLsFh653HJ4qdO3oaqi5ti2TPF/32KmF5anrcQ/99iplGXYa2LQFOHvVcX1r0+1Q/wDfYo+32vT7TD/32KydOfYC0Ovt3oB9qrC+tMH/AEmH/v4KBfWuP+PqH/v4KPZz7AWxg8Uo+7xVUX1qBj7VBn/roKUX9pjH2qD/AL+Cl7OXZgWGOB7VEzAfWonvrTHF1D/38FQNe2o5+0Q59nFONOXZgWlJJp+OhBrPGoQA/wCvjx/vipBf25HE8X4yCrdKXYCw6g1GQB9O9NN9a/8APxD/AN9iozd22P8Aj4h5/wBsU4wkBLuzTgcEVWFzbf8APzD/AN9inLd23/PxD/32KHGXYNCxuzxikRjnuaqtd2+7ieLH++KkS8tsc3EQ/wCBij2crbAWlHPtUgz3qsLy0H/LzD/38H+NON5Z/wDP1D/38H+NQ6cuwblg560wjvioxfWv/P1B/wB/BTHvrUDi5hP/AG0FJU59guOYA8A9Kj25qM3tsf8Al4iz/vilF3ahcfaIs/74q+SS6DZIVPFBIAx+tRG8tv8An4i4/wBsUz7Vb54uIv8AvsU1CXYRPu6g9KTIPNVjdW/P+kR8f7YpBdwFv9fGB/vCn7OXYVyy74HXFN25HNQm4tjj/SIuO28U/wC1WwT/AI+Ij/wMU+SSWwbkqdMZqSPgYqsLu2/5+Iv++xQl5bj/AJeIv++xScJdhl0cdKaSOOai+22uMfaYc/74pjXdsf8Al4i9/nFLkl2EWQARnP4UxgCc4phvbXH/AB8Q59nFH2u1x/x8w/8AfwVPJLsx3JRR2wahN3ag4FzD/wB/BR9stj1uIv8Av4KOSXYCQjA44pzcrx2qBry16/aYv++xR9stQP8Aj4h/77FPkl2CwkyZ69Kj2g9PXikkurc/8vER/wCBimi4t/8AnvF/32K0UZJEvcmUck04KGNQfabcH/Xxn/gYpwu7fp9oi577xxScZdhlkLjFKR09KgF5bf8APxF/32KX7ZbHrcRcj/noKnkn2G7dCRjzR1qJrq24P2iIj03imtd2xA/0iLp/fFNQl2Alzh8U8H3qsbq3x/x8Rf8AfYpEubfdzcRY93FNwk+gFljj8aVRkVB9rtj1uIv++xSrd2w6XEX/AH2KXLLsFyz2o+tVxeWwPFzD/wB9ilN5bf8APzCf+Bip5JdhDnIXk00MCvHSmSXVsQf9IiPHZxVQ3USj/XR5/wB4Vcaba1Q7mhx0FKOO/OaqJdwZ+aeP/vsU/wC12x/5bxD/AIGKHTl2AnPBoHTFQNdWxJP2iL/vsUou7cDAuIv++xS5JdgLKDk0OKgW7th/y8Rf99ika9t8ECeLP++KXJO+whkmQc1LCdy1Xe5gOf38ft84pYrm3UY8+L/vsVpyO2wW1LTcqQKYvHWmG6tv+fiLn/bFM+12+T+/iH/AxUqDtsO5YHOcGncbeRVcXlt/z3j/AO+xTheW/JNxF9N4pOE+wdCVh3xQBjqKha7tv+fiL/vsUG9t8g+fH/32KOSXYNCQj0pVyOpqH7VbY/18X/fYoW7t8/8AHxEM/wC2KOSXYOpNxjikPB9qj+12/P8ApEX/AH2Kb9rt+9xF/wB9ihQl2ETNyRUb8dOvpTPtNvkfv4v++xTWuLcjieL/AL7FUoS7B0JVx+NPC4HIqoLqEf8ALWP/AL7FSrdW+QfPj4/2xQ4S7APkHy0gQnPqO9Ibm2PBnix/vihbm3Ax58WP98UcsuwIeFyfpShBzxTDc2xyPtEQ994p32q2H/LxF/32KTjLsApQAbQOlQFcPjp7077Xb7v9fFj/AHxTWntzkedF/wB9iqjGXYViUAB8mpcVWF1bleZ4s/74p4urfH/HxFn/AHxScZdhk9Jk5HP6VF9qtuf9Ii+m8Ufarcji4i/77FSoyXQZKTkZ5FMI5OKb9rtv+fiL/vsUxrm3P/LxF/32KahLsG4rcnIHIp2ML6elRG4tyeZ4v++xS/abfHE8X/fYqmpdhW00JDjjtioZX7DrQ11BggTRe3zioTNASf3yY/3hVKD7APj+4Oc+tP52/wBKiWeAf8tY/wDvsU4Twrn97H6/fFU79hJaWJVI60hPPSoxcwg586P/AL6FDXEBPE0f/fQo5H2C4hVWGMdqaCA+M80efEMYlj4/2hTRNCTzNGc/7QqreQrWRO2WXFNRcknHB7U4T25HM8f13ikWaDI/fxf99ipV10G11JGj7j8KZznn9af9qtxx58Z/4GKY09uT/r4z/wADFQubsNrsSYyBSY+am/arfn9/H/30KT7Tbr0mi/77FHLJ9AY6ZMDOKplQUxVt7m3ZSPOjx/viqPmx7yPMTGOu4VpTT6oTRIqgHJ9asoABjtVRZ0znzE/76FTC4hxxJGP+BCiUZAn2JieM44qu67mLelSG5hzzNGcf7QqNbiHJzKn/AH0KUU+w7CHIYHPTjApysM4/lTGmixxMmf8AeFCyQrx5qf8AfQq7abCtoTjgEYpc5Xnt15qIXEPOZU9huFKZbc5xMn03CocX2HeyH4BJIHWnIP19ajFzCDjzY8f7wpxuYc/66PHb5hSal2CxMB3600kAZPemrcwY/wBdGP8AgQpjzwY/10fH+0KlRl2DVEwIwMDig9z+VQC4h6ebHj/eFPFxBjHnR5/3hT5GLyHjj6d6CR1wcComuIeD50f/AH0KXz4f+e8eD/tCjlfYFuObJfinryearvPCOk0f13CnRzwBf9bGPbeKbi7DuyfHXNNHHJHXpTftNv08+PP+8KabiDgedH/30Knll1QEp46daa+RgZxTftEBHM0f/fQprXEHJE0f/fQpqL7BYGIJBxT8Z49e5qMS2/A82If8DFKbiEc+dGfQbxT5X2F6iOONvc9aoSKSvToeeavGeAj/AFsfP+0KryyxEcSJkd81pC66BYrhemeG9RToyQw3NkZNRbkxt3A0ilFOQ3J9+lb2BXsWyc4yetIRtPUZ/nUfmLwPMHB9aUSISAXQdccip5RW7j03Ek84zRKo8lz32nPFMSRQhHmLjHQsKc8sfkOokBypwKSjqGp//9k=",
+ "imageHeight": 2048,
+ "imageWidth": 2592
+}
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/pre_processing/config/sample_json/gemini_sample.json b/official/projects/waste_identification_ml/pre_processing/config/sample_json/gemini_sample.json
new file mode 100644
index 00000000000..53c7449d434
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/sample_json/gemini_sample.json
@@ -0,0 +1 @@
+{"info": {}, "licenses": [], "categories": [{"id": 1, "name": "Aluminium_Can"}, {"id": 2, "name": "Aluminium_Foil"}, {"id": 3, "name": "Battery"}, {"id": 4, "name": "Brush"}, {"id": 5, "name": "Bulb"}, {"id": 6, "name": "Fiber_Cardboard"}, {"id": 7, "name": "Fiber_Cup-&-glass"}, {"id": 8, "name": "Fiber_Paper"}, {"id": 9, "name": "Footwear"}, {"id": 10, "name": "Glass_Bottle"}, {"id": 11, "name": "Lighter"}, {"id": 12, "name": "Metals"}, {"id": 13, "name": "Metals_Bottle"}, {"id": 14, "name": "Metals_Container"}, {"id": 15, "name": "Metals_Lid"}, {"id": 16, "name": "Plastics-ABS_Electronics"}, {"id": 17, "name": "Plastics-HDPE_Bottle"}, {"id": 18, "name": "Plastics-HDPE_Container"}, {"id": 19, "name": "Plastics-HDPE_Lid"}, {"id": 20, "name": "Plastics-HDPE_Toys"}, {"id": 21, "name": "Plastics-HDPE_Tube"}, {"id": 22, "name": "Plastics-LDPE_Flexibles"}, {"id": 23, "name": "Plastics-MLP_Flexibles"}, {"id": 24, "name": "Plastics-PC_CD"}, {"id": 25, "name": "Plastics-PC_Goggles"}, {"id": 26, "name": "Plastics-PET_Blister-pack"}, {"id": 27, "name": "Plastics-PET_Bottle"}, {"id": 28, "name": "Plastics-PET_Container"}, {"id": 29, "name": "Plastics-PET_Cup-&-glass"}, {"id": 30, "name": "Plastics-PP_Comb"}, {"id": 31, "name": "Plastics-PP_Container"}, {"id": 32, "name": "Plastics-PP_Pen"}, {"id": 33, "name": "Plastics-PP_Spoon"}, {"id": 34, "name": "Plastics-PP_Straw"}, {"id": 35, "name": "Plastics-PP_Tray"}, {"id": 36, "name": "Plastics-PS"}, {"id": 37, "name": "Plastics-PS_Cup-&-glass"}, {"id": 38, "name": "Plastics-PS_Flexibles"}, {"id": 39, "name": "Plastics-PS_Hangers"}, {"id": 40, "name": "Plastics-PVC_Pipe"}, {"id": 41, "name": "Plastics-Tetrapak_Carton"}, {"id": 42, "name": "Textile_Clothes"}, {"id": 43, "name": "Textile_Flexibles"}, {"id": 44, "name": "Tire"}, {"id": 45, "name": "Wood"}], "images": [{"file_name": "client_recykal_20210826-105712_04d67df42a69988c6ee212d4f544ee577b8cc7be503a40a31dfea297ae0b9023.png", "id": 21490, "width": 1920, "height": 1080}], "annotations": [{"segmentation": [[166, 184, 130, 181, 82, 198, 86, 261, 103, 270, 102, 306, 97, 337, 112, 392, 147, 386, 172, 384, 190, 302, 183, 244]], "area": 17202.5, "bbox": [82, 181, 108, 211], "image_id": 21490, "category_id": 27, "id": 156386, "iscrowd": 0}, {"segmentation": [[253, 208, 255, 233, 266, 250, 261, 371, 242, 373, 224, 328, 201, 305, 187, 300, 187, 309, 197, 254, 217, 240, 227, 209, 240, 204]], "area": 7984.0, "bbox": [187, 204, 79, 169], "image_id": 21490, "category_id": 27, "id": 156387, "iscrowd": 0}, {"segmentation": [[270, 257, 318, 273, 332, 284, 324, 307, 310, 326, 294, 340, 270, 366, 259, 324]], "area": 4575.5, "bbox": [259, 257, 73, 109], "image_id": 21490, "category_id": 27, "id": 156388, "iscrowd": 0}, {"segmentation": [[386, 200, 373, 208, 351, 227, 339, 237, 323, 270, 338, 284, 332, 303, 343, 307, 402, 271, 394, 181]], "area": 5191.0, "bbox": [323, 181, 79, 126], "image_id": 21490, "category_id": 24, "id": 156389, "iscrowd": 0}, {"segmentation": [[461, 125, 468, 150, 483, 162, 474, 223, 465, 250, 425, 263, 406, 263, 399, 189, 405, 165, 420, 144, 425, 118, 450, 117]], "area": 9094.5, "bbox": [399, 117, 84, 146], "image_id": 21490, "category_id": 27, "id": 156390, "iscrowd": 0}, {"segmentation": [[459, 254, 549, 209, 566, 229, 592, 262, 607, 294, 540, 328, 424, 387, 358, 418, 297, 427, 258, 435, 236, 439, 233, 397, 265, 381, 292, 346, 324, 317, 356, 305, 414, 275]], "area": 35665.5, "bbox": [233, 209, 374, 230], "image_id": 21490, "category_id": 17, "id": 156391, "iscrowd": 0}, {"segmentation": [[510, 511, 485, 532, 483, 548, 480, 563, 469, 582, 484, 599, 515, 612, 534, 590, 548, 575, 618, 512]], "area": 7923.5, "bbox": [469, 511, 149, 101], "image_id": 21490, "category_id": 21, "id": 156392, "iscrowd": 0}, {"segmentation": [[656, 310, 636, 321, 630, 328, 621, 348, 647, 380, 667, 385, 669, 427, 665, 433, 672, 452, 696, 494, 728, 533, 748, 560, 790, 540, 831, 520, 744, 407, 715, 387, 698, 363, 671, 324]], "area": 19306.5, "bbox": [621, 310, 210, 250], "image_id": 21490, "category_id": 27, "id": 156393, "iscrowd": 0}, {"segmentation": [[724, 575, 671, 586, 626, 601, 656, 663, 664, 673, 779, 643, 822, 634, 806, 592, 821, 561, 823, 542, 808, 540, 770, 556]], "area": 14946.0, "bbox": [626, 540, 197, 133], "image_id": 21490, "category_id": 21, "id": 156394, "iscrowd": 0}, {"segmentation": [[347, 890, 316, 916, 308, 947, 321, 984, 355, 994, 393, 984, 415, 950, 409, 910, 380, 892]], "area": 8476.5, "bbox": [308, 890, 107, 104], "image_id": 21490, "category_id": 19, "id": 156395, "iscrowd": 0}, {"segmentation": [[134, 982, 121, 997, 121, 1017, 136, 1030, 145, 1056, 171, 1076, 206, 1076, 232, 1037, 220, 1013, 180, 995]], "area": 6790.0, "bbox": [121, 982, 111, 94], "image_id": 21490, "category_id": 27, "id": 156396, "iscrowd": 0}, {"segmentation": [[238, 1007, 228, 1050, 227, 1077, 290, 1074, 302, 1079, 335, 1079, 380, 1075, 382, 1009, 299, 999, 269, 999]], "area": 11087.0, "bbox": [227, 999, 155, 80], "image_id": 21490, "category_id": 31, "id": 156397, "iscrowd": 0}, {"segmentation": [[690, 683, 658, 703, 640, 724, 725, 849, 834, 971, 876, 1024, 894, 1011, 929, 973, 858, 886, 845, 870, 829, 853, 803, 822, 783, 836, 759, 824, 795, 798]], "area": 27929.5, "bbox": [640, 683, 289, 341], "image_id": 21490, "category_id": 30, "id": 156398, "iscrowd": 0}, {"segmentation": [[958, 575, 900, 575, 880, 575, 862, 656, 856, 717, 856, 770, 861, 846, 859, 882, 867, 894, 915, 899, 950, 897, 948, 849, 964, 790, 966, 754, 952, 749, 954, 731, 974, 721, 973, 665, 972, 635]], "area": 32406.0, "bbox": [856, 575, 118, 324], "image_id": 21490, "category_id": 17, "id": 156399, "iscrowd": 0}, {"segmentation": [[1238, 636, 1152, 632, 1142, 770, 1139, 824, 1150, 966, 1146, 966, 1159, 966, 1163, 975, 1187, 982, 1207, 971, 1231, 973, 1244, 803, 1244, 645]], "area": 32539.0, "bbox": [1139, 632, 105, 350], "image_id": 21490, "category_id": 17, "id": 156400, "iscrowd": 0}, {"segmentation": [[1262, 879, 1279, 1021, 1294, 1074, 1321, 1078, 1308, 1013, 1308, 966, 1284, 888, 1271, 873]], "area": 5744.5, "bbox": [1262, 873, 59, 205], "image_id": 21490, "category_id": 4, "id": 156401, "iscrowd": 0}, {"segmentation": [[1362, 533, 1340, 556, 1321, 568, 1288, 621, 1266, 643, 1277, 654, 1279, 684, 1290, 693, 1308, 665, 1341, 634, 1356, 604, 1369, 580, 1380, 547]], "area": 6980.5, "bbox": [1266, 533, 114, 160], "image_id": 21490, "category_id": 21, "id": 156402, "iscrowd": 0}, {"segmentation": [[1082, 297, 1069, 297, 1071, 304, 1050, 311, 1039, 326, 1047, 343, 1024, 341, 1012, 357, 1030, 402, 1050, 424, 1072, 462, 1083, 437, 1113, 479, 1161, 534, 1205, 579, 1229, 610, 1240, 608, 1262, 580, 1281, 593, 1271, 617, 1273, 626, 1286, 617, 1317, 566, 1332, 545, 1306, 483, 1220, 407, 1188, 409, 1133, 394, 1096, 361, 1078, 332, 1091, 302]], "area": 39289.0, "bbox": [1012, 297, 320, 329], "image_id": 21490, "category_id": 27, "id": 156403, "iscrowd": 0}, {"segmentation": [[940, 186, 857, 219, 779, 269, 770, 302, 763, 322, 790, 350, 835, 396, 875, 405, 942, 374, 962, 346, 980, 330, 1001, 335, 1024, 284, 1026, 256, 991, 203, 964, 186]], "area": 37557.5, "bbox": [763, 186, 263, 219], "image_id": 21490, "category_id": 18, "id": 156404, "iscrowd": 0}, {"segmentation": [[940, 398, 931, 411, 934, 448, 947, 472, 942, 494, 986, 555, 1008, 590, 1028, 621, 1060, 634, 1093, 610, 1120, 569, 1126, 553, 1120, 507, 1072, 470, 1047, 427, 1024, 402, 1006, 359, 1021, 337, 993, 341, 982, 334, 958, 357, 938, 385]], "area": 31748.0, "bbox": [931, 334, 195, 300], "image_id": 21490, "category_id": 2, "id": 156405, "iscrowd": 0}, {"segmentation": [[719, 369, 750, 391, 767, 407, 805, 439, 836, 466, 870, 499, 901, 523, 927, 556, 822, 487, 810, 460, 790, 453, 759, 413, 719, 389]], "area": 4095.0, "bbox": [719, 369, 208, 187], "image_id": 21490, "category_id": 4, "id": 156406, "iscrowd": 0}, {"segmentation": [[1677, 254, 1657, 271, 1653, 287, 1669, 295, 1680, 319, 1758, 361, 1835, 398, 1872, 339, 1813, 300, 1714, 258, 1697, 265]], "area": 14936.0, "bbox": [1653, 254, 219, 144], "image_id": 21490, "category_id": 27, "id": 156407, "iscrowd": 0}, {"segmentation": [[343, 474, 311, 477, 260, 503, 225, 527, 205, 558, 201, 584, 227, 614, 256, 647, 310, 713, 326, 722, 372, 704, 418, 673, 453, 634, 468, 599, 446, 584]], "area": 39782.5, "bbox": [201, 474, 267, 248], "image_id": 21490, "category_id": 13, "id": 156408, "iscrowd": 0}, {"segmentation": [[6, 74, 53, 70, 76, 66, 111, 50, 127, 35, 136, 29, 140, 18, 155, 18, 182, 39, 216, 48, 234, 74, 203, 87, 162, 87, 138, 94, 112, 109, 96, 147, 63, 138, 39, 144, 18, 127, 35, 118, 24, 111, 0, 90]], "area": 13294.0, "bbox": [0, 18, 234, 129], "image_id": 21490, "category_id": 2, "id": 156409, "iscrowd": 0}, {"segmentation": [[1309, 684, 1298, 849, 1313, 849, 1340, 649, 1320, 658]], "area": 4075.5, "bbox": [1298, 649, 42, 200], "image_id": 21490, "category_id": 32, "id": 156410, "iscrowd": 0}, {"segmentation": [[720, 978, 682, 996, 636, 1078, 731, 1076, 756, 1042, 767, 967, 722, 956]], "area": 9430.0, "bbox": [636, 956, 131, 122], "image_id": 21490, "category_id": 27, "id": 156411, "iscrowd": 0}]}
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/pre_processing/config/visualization.py b/official/projects/waste_identification_ml/pre_processing/config/visualization.py
new file mode 100644
index 00000000000..08dad779d08
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/config/visualization.py
@@ -0,0 +1,109 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""To visualize of the category distribution in an annotated JSON file."""
+
+#! /usr/bin/env python3
+
+import json
+import numpy as np
+import pandas as pd
+
+
+def data_creation(path: str) -> pd.DataFrame:
+ """Create a dataframe with the occurences of images and categories.
+
+ Args:
+ path: path to the annotated JSON file.
+
+ Returns:
+ dataset consisting of the counts of images and categories.
+ """
+ # get annotation file data into a variable
+ with open(path) as json_file:
+ data = json.load(json_file)
+
+ # count the occurance of each category and an image in the annotation file
+ category_names = [i['name'] for i in data['categories']]
+ category_ids = [i['category_id'] for i in data['annotations']]
+ image_ids = [i['image_id'] for i in data['annotations']]
+
+ # create a dataframe
+ df = pd.DataFrame(
+ list(zip(category_ids, image_ids)), columns=['category_ids', 'image_ids'])
+ df = df.groupby('category_ids').agg(
+ object_count=('category_ids', 'count'),
+ image_count=('image_ids', 'nunique'))
+ df = df.reindex(range(1, len(data['categories']) + 1), fill_value=0)
+ df.index = category_names
+ return df
+
+
+def visualize_detailed_counts_horizontally(path: str) -> None:
+ """Plot a vertical bar graph showing the counts of images & categories.
+
+ Args:
+ path: path to the annotated JSON file.
+ """
+ df = data_creation(path)
+ ax = df.plot(
+ kind='bar',
+ figsize=(40, 10),
+ xlabel='Categories',
+ ylabel='Counts',
+ width=0.8,
+ linewidth=1,
+ edgecolor='white') # rot = 0 for horizontal labeling
+ for p in ax.patches:
+ ax.annotate(
+ text=np.round(p.get_height()),
+ xy=(p.get_x() + p.get_width() / 2., p.get_height()),
+ ha='center',
+ va='top',
+ xytext=(4, 14),
+ textcoords='offset points')
+
+
+def visualize_detailed_counts_vertically(path: str) -> None:
+ """Plot a horizontal bar graph showing the counts of images & categories.
+
+ Args:
+ path: path to the annotated JSON file.
+ """
+ df = data_creation(path)
+ ax = df.plot(
+ kind='barh',
+ figsize=(15, 40),
+ xlabel='Categories',
+ ylabel='Counts',
+ width=0.6)
+ for p in ax.patches:
+ ax.annotate(
+ str(p.get_width()), (p.get_x() + p.get_width(), p.get_y()),
+ xytext=(4, 6),
+ textcoords='offset points')
+
+
+def visualize_annotation_file(path: str) -> None:
+ """Plot a bar graph showing the category distribution.
+
+ Args:
+ path: path to the annotated JSON file.
+ """
+ df = data_creation(path)
+ df['object_count'].plot.bar(
+ figsize=(20, 5),
+ width=0.5,
+ xlabel='Material types',
+ ylabel='count of material types')
diff --git a/official/projects/waste_identification_ml/pre_processing/deprecated_JSON_Generation_for_Training.ipynb b/official/projects/waste_identification_ml/pre_processing/deprecated_JSON_Generation_for_Training.ipynb
new file mode 100644
index 00000000000..9fb846a5a7b
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/deprecated_JSON_Generation_for_Training.ipynb
@@ -0,0 +1,750 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0JmF5ohLbPlF"
+ },
+ "source": [
+ "# Pre processing steps of a COCO JSON annotated file "
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "uXwwz3PlbUX2"
+ },
+ "source": [
+ "Given a single COCO annotated JSON file, your goal is to pre-process in order to remove noise and manipulate it into a form which is suitable for training a ML model. This script will also check if the annotated images are broken or missing."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "E1SxGZD2bv8E"
+ },
+ "source": [
+ "The COCO annotation file includes the following -\n",
+ "\n",
+ "1. Name of the images.\n",
+ "\n",
+ "2. Dimensions of the images.\n",
+ "\n",
+ "3. Classes in the image category.\n",
+ "\n",
+ "4. Name of the super categories of the classes.\n",
+ "\n",
+ "5. Area acquired by the segmented pixels in an image.\n",
+ "\n",
+ "6. Bounding box co-ordinates.\n",
+ "\n",
+ "7. Annotated segmentation coordinates."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "j0v31gxTbweO"
+ },
+ "source": [
+ "There is a lot of noise in the real world annotation file. The images name could be wrong. The images mentioned in an annotation file may not be present in the image folder, which will disrupt the model training procedure. The contents within an annotation file may not match with each other. Even the files present in an image folder may be broken or truncated, which will cause errors while reading image files. Our goal is to eradicate all these problems."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "PyFn96EKb7A-"
+ },
+ "source": [
+ "Our goal is to make sure that all information in the key values corresponds to each other correctly. This notebook will help you achieve this task."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "W6aXxxox0DDa"
+ },
+ "source": [
+ "## Import labels and sample JSON file \n",
+ "To import total classes for the material, material_form and plastic_type we will import the label files from the waste_identification_ml project from Tensorflow Model Garden.\n",
+ "We will also import a noisy sample JSON file to illustrate an example."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WluEHMZYm0zM",
+ "outputId": "b8c4738c-4636-4c56-c6ea-b4e4da92474c",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ }
+ },
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 3536 100 3536 0 0 29714 0 --:--:-- --:--:-- --:--:-- 29714\n",
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 2427 100 2427 0 0 18248 0 --:--:-- --:--:-- --:--:-- 18248\n",
+ " % Total % Received % Xferd Average Speed Time Time Time Current\n",
+ " Dload Upload Total Spent Left Speed\n",
+ "\r 0 0 0 0 0 0 0 0 --:--:-- --:--:-- --:--:-- 0\r100 3303k 100 3303k 0 0 14.1M 0 --:--:-- --:--:-- --:--:-- 14.2M\n"
+ ]
+ }
+ ],
+ "source": [
+ "%%bash\n",
+ "curl -O https://raw.githubusercontent.com/tensorflow/models/master/official/projects/waste_identification_ml/pre_processing/config/categories_list_of_dictionaries.py\n",
+ "\n",
+ "curl -O https://raw.githubusercontent.com/tensorflow/models/master/official/projects/waste_identification_ml/pre_processing/config/sample_json/dataset.json\n",
+ "\n",
+ "mkdir image_folder\n",
+ "\n",
+ "curl -o image_folder/image_2.png https://raw.githubusercontent.com/tensorflow/models/master/official/projects/waste_identification_ml/pre_processing/config/sample_images/image_2.png"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "MRhCAFlVcRm0"
+ },
+ "source": [
+ "## Import the required libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Mnxbo8GBcN2O"
+ },
+ "outputs": [],
+ "source": [
+ "import glob\n",
+ "import tqdm\n",
+ "import json\n",
+ "from PIL import Image\n",
+ "import subprocess\n",
+ "import copy\n",
+ "import os\n",
+ "from google.colab import files\n",
+ "from categories_list_of_dictionaries import *"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "f-05VwsL0mCi"
+ },
+ "outputs": [],
+ "source": [
+ "# reading labels \n",
+ "\n",
+ "images_folder_path = 'image_folder/' #@param {type:\"string\"}\n",
+ "list_of_material = build_material(MATERIAL_LIST,'material-types')\n",
+ "list_of_material_form = build_material(MATERIAL_FORM_LIST,'material-form-types')\n",
+ "list_of_plastic_type = build_material(PLASTICS_SUBCATEGORY_LIST,'plastic-types')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# common labeling typo errors\n",
+ "_KNOWN_TYPOS = {\n",
+ " 'and': '&',\n",
+ " 'Cassete': 'Cassette',\n",
+ " 'Toy':'Toys',\n",
+ " 'Mug-&-Tub':'tub',\n",
+ " 'Toyss':'toys'\n",
+ "}\n",
+ "_KNOWN_TYPOS"
+ ],
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "p7xwZDoc5rZU",
+ "outputId": "33d060a4-f4aa-475a-e140-30bc222553a3"
+ },
+ "execution_count": null,
+ "outputs": [
+ {
+ "output_type": "execute_result",
+ "data": {
+ "text/plain": [
+ "{'and': '&',\n",
+ " 'Cassete': 'Cassette',\n",
+ " 'Toy': 'Toys',\n",
+ " 'Mug-&-Tub': 'tub',\n",
+ " 'Toyss': 'toys'}"
+ ]
+ },
+ "metadata": {},
+ "execution_count": 4
+ }
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "958ZSjT_eG_b"
+ },
+ "source": [
+ "## Utility functions"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "tGOCdeiucUgq"
+ },
+ "outputs": [],
+ "source": [
+ "def read_json(file):\n",
+ " \"\"\"Read any JSON file.\n",
+ "\n",
+ " Args:\n",
+ " file: path to the file\n",
+ " \"\"\"\n",
+ " with open(file) as json_file:\n",
+ " data = json.load(json_file)\n",
+ " return data\n",
+ "\n",
+ "\n",
+ "def search_dict_value(dic, id):\n",
+ " \"\"\"Returns the key of the dictionary from its value'\n",
+ "\n",
+ " Args:\n",
+ " dic = Mapping to search by value.\n",
+ " id = Value to search.\n",
+ " \"\"\" \n",
+ " key_list = list(dic.keys())\n",
+ " val_list = list(dic.values())\n",
+ " position = val_list.index(id)\n",
+ " return key_list[position]\n",
+ "\n",
+ "\n",
+ "def delete_truncated_images(folder_path: str) -> None:\n",
+ " \"\"\"Find and delete truncated images.\n",
+ "\n",
+ " Args:\n",
+ " folder_path: path to the folder where images are saved.\n",
+ " \"\"\"\n",
+ " # path to the images folder to read its content\n",
+ " files = glob.glob(folder_path + '/*')\n",
+ " print('Total number of files in the folder:', len(files))\n",
+ "\n",
+ " num = 0\n",
+ "\n",
+ " # read all image files and remove them from the directory in case they are broken\n",
+ " for file in tqdm.tqdm(files):\n",
+ " if file.endswith(('.png','.jpg')):\n",
+ " try:\n",
+ " img = Image.open(file)\n",
+ " img.verify()\n",
+ " except:\n",
+ " num = num + 1\n",
+ " subprocess.run(['rm', file])\n",
+ " print('Broken file name: ' + file)\n",
+ " if num == 0:\n",
+ " print('\\nNo broken images found')\n",
+ " else:\n",
+ " print('Total number of broken images found:', num)\n",
+ "\n",
+ "\n",
+ "def spelling_correction(dic):\n",
+ " \"\"\"Correcting some common spelling mistakes.\"\"\"\n",
+ " for i in dic['categories']:\n",
+ " for old, new in _KNOWN_TYPOS.items():\n",
+ " i['name'].replace(old, new)\n",
+ "\n",
+ "\n",
+ "def labeling_correction(dic, num, labels_dict):\n",
+ " \"\"\"Matching annotated labels with the correct labels and correcting the mistakes.\n",
+ "\n",
+ " Mapping the modified labeling ID with the corresponding original ID for alignment\n",
+ " of categories.\n",
+ "\n",
+ " Args:\n",
+ " dic: JSON file read as a dictionary\n",
+ " num: keyword position inside the label\n",
+ " labels_dict: dictionary showing the labels ID of the original categories \n",
+ " \"\"\"\n",
+ " incorrect_labels = []\n",
+ " mapping_list = []\n",
+ " for i in dic['categories']:\n",
+ " if i['name'].split('_')[num].lower() in labels_dict.values():\n",
+ " id = i['id']\n",
+ " name = i['name'].split('_')[num]\n",
+ " id_match = search_dict_value(labels_dict, i['name'].split('_')[num].lower())\n",
+ " mapping_list.append((id, name, id_match))\n",
+ " else:\n",
+ " id = i['id']\n",
+ " incorrect_labels.append(id)\n",
+ " return mapping_list, incorrect_labels\n",
+ "\n",
+ "\n",
+ "def images_key(dic):\n",
+ " \"\"\"Align the data within the dictionary in the 'images' key.\n",
+ " \n",
+ " The 'image_id' parameter in the 'annotation' key is the same as 'id' in the 'images' key of the dictionary. This function \n",
+ " will also remove all image data from the 'images' key whose 'id' does not \n",
+ " match with 'image_id' in the 'annotation' key in the dictionary.\n",
+ "\n",
+ " Args:\n",
+ " dic: where the JSON file is read into\n",
+ " \"\"\"\n",
+ " image_ids = set(i['image_id'] for i in dic['annotations'])\n",
+ " new_images = [i for i in dic['images'] if i['id'] in image_ids]\n",
+ " return new_images\n",
+ "\n",
+ "\n",
+ "def annotations_key(dic, incorrect_labels, mapping_dict):\n",
+ " \"\"\"Align the data within the dictionary in the 'annotation' key.\n",
+ " \n",
+ " Notice that the 'category_id' in the 'annotation' key is same as 'id' \n",
+ " in the 'categories' key of the dictionary.\n",
+ "\n",
+ " Args:\n",
+ " dic: where the JSON file is read into\n",
+ " \"\"\"\n",
+ " new_annotation = []\n",
+ "\n",
+ " for i in dic['annotations']:\n",
+ " id = i['category_id']\n",
+ " if id not in incorrect_labels:\n",
+ " new_id = [i[2] for i in mapping_dict if i[0] == id][0]\n",
+ " i['category_id'] = new_id\n",
+ " new_annotation.append(i)\n",
+ " return new_annotation\n",
+ "\n",
+ "\n",
+ "def annotated_images(folder_path, dic):\n",
+ " \"\"\"Get images infromation that are mentioned in an annotation file but are not present in an image folder.\n",
+ "\n",
+ " Args:\n",
+ " folder_path: path of an image folder.\n",
+ " \"\"\"\n",
+ " # read the file names from the directory \n",
+ " files = glob.glob(folder_path + '/*')\n",
+ " files = set(map(os.path.basename, files))\n",
+ "\n",
+ " # list of images in an annotation file\n",
+ " dic['images'] = [i for i in dic['images'] if i['file_name'] in files]\n",
+ " return dic\n",
+ "\n",
+ "\n",
+ "def image_annotation_key(dic):\n",
+ " \"\"\"Check if same images are present in both \"images\" key and \"annotations\" key. \n",
+ "\n",
+ " List of the image IDs which are in the \"images\" key but NOT in \"annotation\" key.\n",
+ " Remove information if they are not present in both keys.\n",
+ "\n",
+ " Args:\n",
+ " dic: annotation file read as a dictionary\n",
+ " \"\"\"\n",
+ " images_id = [i['id'] for i in dic['images']]\n",
+ " annotation_id = [i['image_id'] for i in dic['annotations']]\n",
+ " common_list = set(images_id).intersection(annotation_id)\n",
+ " dic['images'] = [i for i in dic['images'] if i['id'] in common_list]\n",
+ " dic['annotations'] = [i for i in dic['annotations'] if i['image_id'] in common_list]\n",
+ " return dic"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0OoDmNC22ycz"
+ },
+ "source": [
+ "## Find and delete truncated images from the image folder."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "bUUu3F6I20w3",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "21ad1270-0486-4eb8-b4ed-07d7dd99b936"
+ },
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stdout",
+ "text": [
+ "Total number of files in the folder: 1\n"
+ ]
+ },
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "100%|██████████| 1/1 [00:00<00:00, 30.21it/s]"
+ ]
+ },
+ {
+ "output_type": "stream",
+ "name": "stdout",
+ "text": [
+ "\n",
+ "No broken images found\n"
+ ]
+ },
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "\n"
+ ]
+ }
+ ],
+ "source": [
+ "delete_truncated_images(images_folder_path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "65XuyPBSea7-"
+ },
+ "source": [
+ "## Perform operations on the file\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "l-uMtZK2edPY",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "428298ff-af02-44c2-c06b-7223646a70db"
+ },
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stdout",
+ "text": [
+ "dict_keys(['images', 'annotations', 'categories'])\n"
+ ]
+ }
+ ],
+ "source": [
+ "# read json file and it should contain at least the three keys as shown below\n",
+ "path_to_json = 'dataset.json' #@param {type:\"string\"}\n",
+ "data = read_json(path_to_json)\n",
+ "print(data.keys())\n",
+ "\n",
+ "# create a copy to compare the results in the end\n",
+ "data_preprocessing = copy.deepcopy(data)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "G8w7MfDtvDIq",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "4c33d866-e788-4b59-cdb2-b3b929c25f7b"
+ },
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "100%|██████████| 6/6 [00:00<00:00, 51463.85it/s]"
+ ]
+ },
+ {
+ "output_type": "stream",
+ "name": "stdout",
+ "text": [
+ "\n",
+ "Total number of wrong annotated labels are 5\n"
+ ]
+ },
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "\n"
+ ]
+ }
+ ],
+ "source": [
+ "# checking labeling mistakes as all annotated labels should have 6 keywords connected by '_' \n",
+ "num = 0\n",
+ "for i in tqdm.tqdm(data['categories']):\n",
+ " if len(i['name'].split('_')) != 6:\n",
+ " num += 1\n",
+ "print('\\nTotal number of wrong annotated labels are', num)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "q2jOWegZxPEp",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "0d44c754-aead-422a-c6bc-3fa304acd2b5"
+ },
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stdout",
+ "text": [
+ "\n",
+ "Total number of labels which has less than 6 keywords are 0\n"
+ ]
+ },
+ {
+ "output_type": "execute_result",
+ "data": {
+ "text/plain": [
+ "[{'id': 0,\n",
+ " 'name': 'plastics_HDPE_flexible_color_SAchets-&-pouch_pouch',\n",
+ " 'supercategory': 'plastics_HDPE_flexible_color_SAchets-&-pouch_pouch'},\n",
+ " {'id': 1,\n",
+ " 'name': 'Plastics_HDPE_Rigid_Blue_Lid_Bottle-Cap_Na_Na',\n",
+ " 'supercategory': 'Plastics_HDPE_Rigid_Blue_Lid_Bottle-Cap_Na_Na'},\n",
+ " {'id': 2,\n",
+ " 'name': 'Plastics_peTE_Na_Clear_Bottle_Shampoo-Bottle_250Ml_Vlcc',\n",
+ " 'supercategory': 'Plastics_peTE_Na_Clear_Bottle_Shampoo-Bottle_250Ml_Vlcc'},\n",
+ " {'id': 3,\n",
+ " 'name': 'Plastics_na_Rigid_Blue_Bottle_Hair-Oil-Bottle-500Ml_Parachute',\n",
+ " 'supercategory': 'Plastics_na_Rigid_Blue_Bottle_Hair-Oil-Bottle-500Ml_Parachute'},\n",
+ " {'id': 4,\n",
+ " 'name': 'Plastics_HDPE_Rigid_Na_Cosmetic_Comb_Na_Na',\n",
+ " 'supercategory': 'Plastics_HDPE_Rigid_Na_Cosmetic_Comb_Na_Na'},\n",
+ " {'id': 5,\n",
+ " 'name': 'Plastics_PETE_Na_Clear_Bottle_Energy-Drink-Bottle_250Ml_Sting-Energy',\n",
+ " 'supercategory': 'Plastics_PETE_Na_Clear_Bottle_Energy-Drink-Bottle_250Ml_Sting-Energy'}]"
+ ]
+ },
+ "metadata": {},
+ "execution_count": 9
+ }
+ ],
+ "source": [
+ "# remove category labels which has less than 6 keywords\n",
+ "categories = []\n",
+ "num = 0\n",
+ "for i in data['categories']:\n",
+ " if len(i['name'].split('_')) >= 6:\n",
+ " categories.append(i)\n",
+ " else:\n",
+ " num += 1\n",
+ "print('\\nTotal number of labels which has less than 6 keywords are', num)\n",
+ "data['categories'] = categories\n",
+ "\n",
+ "# display categories after removing the labels\n",
+ "data['categories']"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qup_-ReIz-iv",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "55963fc5-89f9-4767-b10e-b8ca8ea7a65d"
+ },
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "100%|██████████| 6/6 [00:00<00:00, 48026.38it/s]\n"
+ ]
+ },
+ {
+ "output_type": "execute_result",
+ "data": {
+ "text/plain": [
+ "[{'id': 0,\n",
+ " 'name': 'plastics_HDPE_flexible_color_SAchets-&-pouch_pouch',\n",
+ " 'supercategory': 'plastics_HDPE_flexible_color_SAchets-&-pouch_pouch'},\n",
+ " {'id': 1,\n",
+ " 'name': 'Plastics_HDPE_Rigid_Blue_Lid_Bottle-Cap-Na-Na',\n",
+ " 'supercategory': 'Plastics_HDPE_Rigid_Blue_Lid_Bottle-Cap_Na_Na'},\n",
+ " {'id': 2,\n",
+ " 'name': 'Plastics_peTE_Na_Clear_Bottle_Shampoo-Bottle-250Ml-Vlcc',\n",
+ " 'supercategory': 'Plastics_peTE_Na_Clear_Bottle_Shampoo-Bottle_250Ml_Vlcc'},\n",
+ " {'id': 3,\n",
+ " 'name': 'Plastics_na_Rigid_Blue_Bottle_Hair-Oil-Bottle-500Ml-Parachute',\n",
+ " 'supercategory': 'Plastics_na_Rigid_Blue_Bottle_Hair-Oil-Bottle-500Ml_Parachute'},\n",
+ " {'id': 4,\n",
+ " 'name': 'Plastics_HDPE_Rigid_Na_Cosmetic_Comb-Na-Na',\n",
+ " 'supercategory': 'Plastics_HDPE_Rigid_Na_Cosmetic_Comb_Na_Na'},\n",
+ " {'id': 5,\n",
+ " 'name': 'Plastics_PETE_Na_Clear_Bottle_Energy-Drink-Bottle-250Ml-Sting-Energy',\n",
+ " 'supercategory': 'Plastics_PETE_Na_Clear_Bottle_Energy-Drink-Bottle_250Ml_Sting-Energy'}]"
+ ]
+ },
+ "metadata": {},
+ "execution_count": 10
+ }
+ ],
+ "source": [
+ "# According to the collected data it was found that most issues occurs from the\n",
+ "# 6th keyword which are the sub category of the material form.\n",
+ "\n",
+ "for i in tqdm.tqdm(data['categories']):\n",
+ " l1 = i['name'].split('_')[:5]\n",
+ " l2 = i['name'].split('_')[5:]\n",
+ " l1.append('-'.join(l2))\n",
+ " i['name'] = '_'.join(l1)\n",
+ "\n",
+ "# display categories after making corrections\n",
+ "data['categories']"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Ro8KNGaGFv7k",
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "outputId": "2dcd1d5b-2387-49be-e1ab-45b796169e4b"
+ },
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stdout",
+ "text": [
+ "Dictionary characteristics before processing :\n",
+ "images: 2 categories: 6 annotations: 6\n",
+ "\n",
+ "Dictionary characteristics after processing of material_type_annotation :\n",
+ "images: 1 categories: 10 annotations: 5\n",
+ "\n",
+ "Dictionary characteristics after processing of material_form_type_annotation :\n",
+ "images: 1 categories: 34 annotations: 5\n",
+ "\n",
+ "Dictionary characteristics after processing of plastic_type_annotation :\n",
+ "images: 1 categories: 9 annotations: 4\n"
+ ]
+ }
+ ],
+ "source": [
+ "print('Dictionary characteristics before processing :')\n",
+ "print('images:',len(data_preprocessing['images']),'categories:', len(data_preprocessing['categories']),'annotations:',len(data_preprocessing['annotations']))\n",
+ "\n",
+ "list_of_categories = [(list_of_material,0,'material_type_annotation.json'),\\\n",
+ " (list_of_material_form,4,'material_form_type_annotation.json'),\\\n",
+ " (list_of_plastic_type,1,'plastic_type_annotation.json')]\n",
+ "\n",
+ "for m in list_of_categories:\n",
+ "\n",
+ " data_processing = copy.deepcopy(data)\n",
+ "\n",
+ " # create a dict showing TDs corresponding to the labels & convert all words\n",
+ " # to lower case in order to eliminate case sensitive issues\n",
+ " labels_dict = dict([(i['id'], i['name'].lower()) for i in m[0]])\n",
+ "\n",
+ " # correcting grammatical errors\n",
+ " spelling_correction(data_processing)\n",
+ "\n",
+ " # create a mapping table to map each label to the right label structure.\n",
+ " # find the incorrect labels.\n",
+ " mapping_dict, incorrect_labels = labeling_correction(data_processing, m[1], labels_dict) \n",
+ "\n",
+ " # change the 'categories' key\n",
+ " data_processing['categories'] = m[0]\n",
+ "\n",
+ " # change the 'annotation' key\n",
+ " data_processing['annotations'] = annotations_key(data_processing, incorrect_labels, mapping_dict)\n",
+ "\n",
+ " # change the 'images' key\n",
+ " data_processing['images'] = images_key(data_processing)\n",
+ "\n",
+ " # remove data from the 'images' key not present in the image folder\n",
+ " data_processing = annotated_images(images_folder_path, data_processing)\n",
+ "\n",
+ " # align 'images' and 'annotations' key\n",
+ " data_processing = image_annotation_key(data_processing)\n",
+ "\n",
+ " # write to a new JSON file\n",
+ " with open(m[2], 'w') as opened_file:\n",
+ " opened_file.write(json.dumps(data_processing, indent=4))\n",
+ "\n",
+ " print('\\nDictionary characteristics after processing of', m[2].replace('.json','') ,':')\n",
+ " print('images:',len(data_processing['images']),'categories:', len(data_processing['categories']),'annotations:',len(data_processing['annotations'])) "
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# View the final JSON file\n",
+ "try:\n",
+ " files.view(m[2]) # use files.download to download the file\n",
+ "except ImportError:\n",
+ " pass"
+ ],
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 17
+ },
+ "id": "nC6XzQYL15Ki",
+ "outputId": "ea8e0a7d-79bc-4321-8a63-4855afc190e1"
+ },
+ "execution_count": 13,
+ "outputs": [
+ {
+ "output_type": "display_data",
+ "data": {
+ "text/plain": [
+ ""
+ ],
+ "application/javascript": [
+ "\n",
+ " ((filepath) => {{\n",
+ " if (!google.colab.kernel.accessAllowed) {{\n",
+ " return;\n",
+ " }}\n",
+ " google.colab.files.view(filepath);\n",
+ " }})(\"/content/plastic_type_annotation.json\")"
+ ]
+ },
+ "metadata": {}
+ }
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "collapsed_sections": [],
+ "name": "json_preparation.ipynb",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
\ No newline at end of file
diff --git a/official/projects/waste_identification_ml/pre_processing/labelme_to_coco.ipynb b/official/projects/waste_identification_ml/pre_processing/labelme_to_coco.ipynb
new file mode 100644
index 00000000000..d0ca6a30542
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/labelme_to_coco.ipynb
@@ -0,0 +1,224 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "wINm_lPYZhlO"
+ },
+ "source": [
+ "# Convert label me annotations to COCO JSON format"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "jlmasaZNtJ6C"
+ },
+ "source": [
+ "Given the images and their corresponding annotated JSON files exported from the \"labelme\" tool. Our goal is to generate COCO format JSON files."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "M-LDdP-NtT3F"
+ },
+ "source": [
+ "We will use an open source library called labelme2coco to convert all the JSON files from \"labelme\" tool to COCO JSON format."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "RojnXi7lfLSA"
+ },
+ "source": [
+ "Put the annotated JSON files and their corresponding images in the same folder and create another folder for storing the output COCO JSON file and then use the labelme2coco tool to export the output."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Ud4OzlNMtHv7"
+ },
+ "outputs": [],
+ "source": [
+ "# install the library and RESTART runtime\n",
+ "!pip install labelme2coco"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nWgcQzpHt4hq"
+ },
+ "outputs": [],
+ "source": [
+ "# import the library\n",
+ "import labelme2coco\n",
+ "import json"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 1,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 56,
+ "status": "ok",
+ "timestamp": 1660342052001,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "3Wt3WELFXO5o"
+ },
+ "outputs": [],
+ "source": [
+ "# Import a sample annotation file exported from \"labelme\" tool and its corresponding image\n",
+ "!curl -O https://raw.githubusercontent.com/tensorflow/models/master/official/projects/\\\n",
+ "waste_identification_ml/pre_processing/config/sample_json/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.json\n",
+ "\n",
+ "# Import the corresponding image file mentioned in the annotation file\n",
+ "!curl -O https://raw.githubusercontent.com/tensorflow/models/master/official/projects/waste_identification_ml/\\\n",
+ "pre_processing/config/sample_images/ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "9A8aivRPXXqF",
+ "outputId": "047b6300-5e58-4d0d-c993-a3a29ce98966"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "dict_keys(['version', 'flags', 'shapes', 'imagePath', 'imageData', 'imageHeight', 'imageWidth'])"
+ ]
+ },
+ "execution_count": 65,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "# check the content of the file exported from the \"labelme\" tool\n",
+ "with open('ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.json') as json_file:\n",
+ " data = json.load(json_file)\n",
+ "data.keys()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "2wQKjve3XpG7"
+ },
+ "outputs": [],
+ "source": [
+ "%%bash\n",
+ "\n",
+ "# set input directory that contains labelme annotations and image files\n",
+ "mkdir labelme_folder\n",
+ "\n",
+ "# set output dir\n",
+ "mkdir export_dir\n",
+ "\n",
+ "# put all the annotation files exported from \"labelme\" tool and their corresponding images!ls in the labelme_folder\n",
+ "mv ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.json labelme_folder/\n",
+ "mv ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg labelme_folder/"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "i5G7dOOycExg",
+ "outputId": "f235fe9c-7ad8-439a-c89d-26a331ecc31c"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "There are 1 listed files in folder .\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ "Converting labelme annotations to COCO format: 100%|██████████| 1/1 [00:00\u003c00:00, 171.26it/s]\n"
+ ]
+ }
+ ],
+ "source": [
+ "input = '/content/labelme_folder/'\n",
+ "output = '/content/export_dir/'\n",
+ "\n",
+ "# it will combine all the JSON files and convert them to COCO JSON format\n",
+ "labelme2coco.convert(input, output)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "S6jqV6PtbOtA",
+ "outputId": "15e02124-6558-4427-ec88-5f0bccbbb33e"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "dict_keys(['images', 'annotations', 'categories'])\n",
+ "[{'id': 0, 'name': 'plastics_HDPE_flexible_color_SAchets-\u0026-pouch_pouch', 'supercategory': 'plastics_HDPE_flexible_color_SAchets-\u0026-pouch_pouch'}]\n",
+ "[{'height': 2048, 'width': 2592, 'id': 1, 'file_name': 'ffdeb4cd-43ba-4ca0-a1e6-aa5824005f44.jpg'}]\n",
+ "[832, 255, 729, 697]\n"
+ ]
+ }
+ ],
+ "source": [
+ "# check the content of the output COCO JSON format\n",
+ "with open('export_dir/dataset.json') as json_file:\n",
+ " data = json.load(json_file)\n",
+ "print(data.keys())\n",
+ "print(data['categories'])\n",
+ "print(data['images'])\n",
+ "print(data['annotations'][0]['bbox'])"
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "collapsed_sections": [],
+ "name": "labelme_to_coco.ipynb",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/merge_coco_files.ipynb b/official/projects/waste_identification_ml/pre_processing/merge_coco_files.ipynb
new file mode 100644
index 00000000000..3a940f1de3e
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/merge_coco_files.ipynb
@@ -0,0 +1,392 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0cTM_BOrBUSU"
+ },
+ "source": [
+ "# Merge multiple COCO annotation JSON files into one file. "
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cwJuts2DBaaU"
+ },
+ "source": [
+ "Given multiple COCO annotated JSON files, your goal is to merge them into one COCO annotated JSON file. \n",
+ "\n",
+ "A merged COCO annotated JSON file is required where all the data is in one place and it becomes easy to split it into a training and validation JSON file according to the percentage ratio. In case you already have a validated COCO annotated JSON file, then this notebook can be used to merge multiple files into one training COCO annotated JSON file. \n",
+ "\n",
+ "This notebook uses a third party library to accomplish this task. Recursion is used to combine multiple JSON files using a third party library. \n",
+ "\n",
+ "This notebook is an end to end example. When you run the notebook, it will take all the multiple JSON files and will output one JSON file. "
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "aDE5tjSUm1qu"
+ },
+ "source": [
+ "**Note** - In this example, we assume that all our data is saved on Google drive and we will also write our outputs to Google drive. We also assume that the script will be used as a Google Colab notebook. But this can be changed according to the needs of users. They can modify this in case they are working on their local workstation, remote server or any other database. This colab notebook can be changed to a regular jupyter notebook running on a local machine according to the need of the users."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "BM-tYHTlWhDQ"
+ },
+ "source": [
+ "## **MUST DO** - Install the package and restart runtime"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "zegGuOmQGBOr"
+ },
+ "outputs": [],
+ "source": [
+ "# install python object detection insights library to merge multiple COCO annotation files\n",
+ "!pip install pyodi\n",
+ "\n",
+ "# RESTART THE RUNTIME in order to use this library"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "l7eLOdQ5F33b"
+ },
+ "source": [
+ "## Run the below command to connect to your google drive"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "h5soS-6URktT"
+ },
+ "outputs": [],
+ "source": [
+ "# import other libraries\n",
+ "from google.colab import drive\n",
+ "import pyodi\n",
+ "import subprocess\n",
+ "import sys\n",
+ "import os\n",
+ "import json\n",
+ "import numpy as np\n",
+ "import pandas as pd"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "mTXqVbFdlqxi",
+ "outputId": "b12566b2-458f-4673-eb97-7cf30075e258"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Mounted at /content/gdrive\n",
+ "Successful\n"
+ ]
+ }
+ ],
+ "source": [
+ "# connect to google drive\n",
+ "drive.mount('/content/gdrive')\n",
+ "\n",
+ "# making an alias for the root path\n",
+ "try:\n",
+ " !ln -s /content/gdrive/My\\ Drive/ /mydrive\n",
+ " print('Successful')\n",
+ "except Exception as e:\n",
+ " print(e)\n",
+ " print('Not successful')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "U-80gvViT833"
+ },
+ "source": [
+ "## Visualization function"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "ipOuE69eT6Vk"
+ },
+ "outputs": [],
+ "source": [
+ "def data_creation(path: str) -\u003e pd.DataFrame:\n",
+ " \"\"\"Create a dataframe with the occurences of images and categories.\n",
+ " Args:\n",
+ " path: path to the annotated JSON file.\n",
+ " Returns:\n",
+ " dataset consisting of the counts of images and categories.\n",
+ " \"\"\"\n",
+ " # get annotation file data into a variable\n",
+ " with open(path) as json_file:\n",
+ " data = json.load(json_file)\n",
+ "\n",
+ " # count the occurance of each category and an image in the annotation file\n",
+ " category_names = [i['name'] for i in data['categories']]\n",
+ " category_ids = [i['category_id'] for i in data['annotations']]\n",
+ " image_ids = [i['image_id'] for i in data['annotations']]\n",
+ "\n",
+ " # create a dataframe\n",
+ " df = pd.DataFrame(\n",
+ " list(zip(category_ids, image_ids)), columns=['category_ids', 'image_ids'])\n",
+ " df = df.groupby('category_ids').agg(\n",
+ " object_count=('category_ids', 'count'),\n",
+ " image_count=('image_ids', 'nunique'))\n",
+ " df = df.reindex(range(1, len(data['categories']) + 1), fill_value=0)\n",
+ " df.index = category_names\n",
+ " return df\n",
+ "\n",
+ "def visualize_detailed_counts_horizontally(path: str) -\u003e None:\n",
+ " \"\"\"Plot a vertical bar graph showing the counts of images \u0026 categories.\n",
+ " Args:\n",
+ " path: path to the annotated JSON file.\n",
+ " \"\"\"\n",
+ " df = data_creation(path)\n",
+ " ax = df.plot(\n",
+ " kind='bar',\n",
+ " figsize=(40, 10),\n",
+ " xlabel='Categories',\n",
+ " ylabel='Counts',\n",
+ " width=0.8,\n",
+ " linewidth=1,\n",
+ " edgecolor='white') # rot = 0 for horizontal labeling\n",
+ " for p in ax.patches:\n",
+ " ax.annotate(\n",
+ " text=np.round(p.get_height()),\n",
+ " xy=(p.get_x() + p.get_width() / 2., p.get_height()),\n",
+ " ha='center',\n",
+ " va='top',\n",
+ " xytext=(4, 14),\n",
+ " textcoords='offset points')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "4hF5F_QE627R"
+ },
+ "source": [
+ "## Define the paths of inputs and outputs"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "qUE6cHse3zOT"
+ },
+ "outputs": [],
+ "source": [
+ "def list_full_paths(directory):\n",
+ " '''return the files names in a directory with absolute path.\n",
+ " Args:\n",
+ " directory: path where all the files that need to merge are saved.\n",
+ " '''\n",
+ " return [os.path.join(directory, file) for file in os.listdir(directory)]"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "a5UVAF--2nas"
+ },
+ "outputs": [],
+ "source": [
+ "folder_with_jsons = '/mydrive/TFHub/jsons/' #@param {type:\"string\"}\n",
+ "output_merged_file = '/mydrive/TFHub/jsons/merged.json' #@param {type:\"string\"}"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "80gZnQCU30gK",
+ "outputId": "d3fdb152-c63a-42b5-f5bf-f4d50eb16d02"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "8"
+ ]
+ },
+ "execution_count": 8,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "# get a list of all the JSON files that need to merge with their absolute paths\n",
+ "list_of_json_files = list_full_paths(folder_with_jsons)\n",
+ "len(list_of_json_files)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "6ot4VWOcSTWO"
+ },
+ "source": [
+ "# Merge the files"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "hngtw3K6S7Qx"
+ },
+ "outputs": [],
+ "source": [
+ "def merge_two_files(file1, file2, output_file):\n",
+ " \"\"\"Function to merge 2 files\n",
+ "\n",
+ " Args:\n",
+ " file1: path of the 1st COCO annotation json file\n",
+ " file2: path of the 2nd COCO annotation json file\n",
+ " output_file: path of the output COCO annotation json file after merge\n",
+ "\n",
+ " Returns:\n",
+ " Path of the merged COCO annotation json file\n",
+ " \"\"\"\n",
+ " subprocess.run(['pyodi', 'coco', 'merge', file1, file2, output_file])\n",
+ " return output_file\n",
+ "\n",
+ "def merge_multiple_files(list_of_files,output_file_path):\n",
+ " \"\"\"Recursive function to merge multiple files\n",
+ "\n",
+ " Args:\n",
+ " list_of_files: list of all the COCO annotation json files that need to be merged \n",
+ " output_file_path: path of the output COCO annotation json file after merge\n",
+ "\n",
+ " Returns:\n",
+ " Path of the merged COCO annotation json file\n",
+ " \"\"\"\n",
+ " if len(list_of_files) == 2:\n",
+ " return merge_two_files(list_of_files[0], list_of_files[1], output_file_path)\n",
+ "\n",
+ " else:\n",
+ " return merge_two_files(list_of_files[0], merge_multiple_files(list_of_files[1:], output_file_path), output_file_path)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "__Wj0rO2D-HT"
+ },
+ "source": [
+ "The output of the below code will be a merged COCO annotation file in the same directory."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "ZMKFOjchVMlH",
+ "outputId": "815a83a3-43ae-478d-abaa-58181f034b94"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Total number of files to merge : 8\n",
+ "Merge Done\n"
+ ]
+ }
+ ],
+ "source": [
+ "# call function to merge multiple files\n",
+ "print('Total number of files to merge :', len(list_of_json_files))\n",
+ "merge_multiple_files(list_of_json_files, output_merged_file)\n",
+ "\n",
+ "print('Merge Done')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "vf4ESHiDEQYF"
+ },
+ "source": [
+ "## Visualize the results"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/",
+ "height": 429
+ },
+ "id": "UBwMf-W0EqG1",
+ "outputId": "9bfbfb43-2ad2-4c18-8d75-05ea56a9573e"
+ },
+ "outputs": [
+ {
+ "data": {
+ "image/png": "iVBORw0KGgoAAAANSUhEUgAACPoAAAKeCAYAAAA21ecdAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjMuNCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8QVMy6AAAACXBIWXMAAAsTAAALEwEAmpwYAADFv0lEQVR4nOzdebxXVb0//tdickJSEwxFA3Jg9og4pZJDGA43xxyiUpQss3K4VvS795tUlmRqmVppjqmpaZnmPOdwRT0KTjimpKhXccABRQH37w8+ngsKisnhuPH5fDzOg73fe6293wu7t+G8HmuVqqoCAAAAAAAAAAB8tLVr6wYAAAAAAAAAAID3J+gDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA10KGtG1jUVlxxxapnz55t3QYAAAAAAAAAAMzTHXfc8VxVVV3fWf/YBX169uyZ5ubmtm4DAAAAAAAAAADmqZTyr3nVHd0FAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA10KGtGwAAAAAAAAAAYNGYMWNGJk+enOnTp7d1KyRZcskl06NHj3Ts2HGBxgv6AAAAAAAAAAB8TEyePDnLLrtsevbsmVJKW7fzsVZVVZ5//vlMnjw5vXr1WqA5ju4CAAAAAAAAAPiYmD59ej75yU8K+XwElFLyyU9+8gPtriToAwAAAAAAAADwMSLk89HxQf9ZCPoAAAAAAAAAAEANCPoAAAAAAAAAAHxMTZ8x6yPxvkmTJmXAgAHzfDZq1KhMnDjxA79zwoQJufTSS/+tfhaWSZMm5U9/+tNCe1+HhfYmAAAAAAAAAABqZcmO7dNz9CUL7X2Txm670N71tpNOOunfmjdhwoQ0Nzdnm222WcgdLbi3gz5f/vKXF8r77OgDAAAAAAAAbaRnz54ZOHBgmpqaMmTIkCTJXXfdlY022igDBw7Mf/zHf+Tll19Oklx11VVZd911M3DgwKy77rq59tprkySvvfZatt122/Tp0yf9+/fP6NGjW95/9NFHp1+/fhk0aFC23HLL/Otf/1r0iwSAeTj66KMzYMCADBgwIL/+9a+TJDNnzsyIESPSt2/f7LLLLnnttdeSJJtttlmam5uTJFdeeWU22mijDB48OF/60pfy6quvJkluv/32fPazn83aa6+d9ddfPy+99FJ+9KMf5dxzz01TU1POPffcefbx6quvZuTIkRk4cGAGDRqUv/zlL0mSs88+OwMHDsyAAQPygx/8oGV8586dW67PP//87LXXXkmSvfbaK9/97nfz2c9+Nr17987555+fJBk9enRuvPHGNDU15Ve/+tWH/nsT9AEAAAAAAIA2dN1117XsOJDMPp5k7Nixueeee7Ljjjvml7/8ZZJkxRVXzN///vfcc889Of300/PVr3615R2HHHJIHnjggYwfPz4333xzLrvssiTJOuusk+bm5tx9993ZZZdd8v3vf3/RLxAA3uGOO+7IqaeemltvvTXjxo3LH/7wh7z44ot58MEH861vfSv3339/unTpkt/+9rdzzXvuuedy2GGH5eqrr86dd96ZIUOG5Oijj86bb76Z3XbbLcccc0zuuuuuXH311VlmmWXyk5/8JLvttlsmTJiQ3XbbbZ69/PSnP80nPvGJ3HPPPbn77ruzxRZb5KmnnsoPfvCDXHvttZkwYUJuv/32/O1vf3vfdT399NO56aabcvHFF7cEb8eOHZtNN900EyZMyEEHHfSh/+4EfQAAAAAAAOAj5KGHHsrQoUOTJMOGDWvZWWCdddbJyiuvnCTp379/Xn/99bzxxhtZeumls/nmmydJOnXqlMGDB2fy5MlJks033zxLL710kmTDDTdsqQNAW7rpppuy4447Zplllknnzp2z00475cYbb8yqq66ajTfeOEnyla98JTfddNNc88aNG5eJEydm4403TlNTU04//fT861//yoMPPpju3btnvfXWS5J06dIlHTp0WKBerr766uy///4t98svv3xuv/32bLbZZunatWs6dOiQESNG5IYbbnjfd+2www5p165d+vXrl2eeeWZB/zo+kFYN+pRSliulnF9KeaCUcn8pZaNSygqllKtKKQ83/ly+MbaUUn5TSnmklHJ3KWXwHO/ZszH+4VLKnnPU1y2l3NOY85tSSmnN9QAAAAAAAMDCVErJVlttlXXXXTcnnnhiktkhngsvvDBJct555+WJJ55417y//OUvGTx4cJZYYom56lOnTs3f//73bLnllu+ac/LJJ2frrbduhVUAwMLxztjHO++rqsqwYcMyYcKETJgwIRMnTszJJ5+8KFucq6fp06fP9WzOf1+uqqpVvt/aO/ock+Tyqqr6JFk7yf1JRie5pqqqNZJc07hPkq2TrNH42TfJ75KklLJCkkOTbJBk/SSHvh0Oaoz5+hzzhrfyegAAAAAAAGChuemmm3LnnXfmsssuy/HHH58bbrghp5xySn77299m3XXXzSuvvJJOnTrNNee+++7LD37wg5xwwglz1WfOnJk99tgj3/3ud9O7d++5np155plpbm7O9773vVZfEwC8n0033TR/+9vf8tprr2XatGm54IILsummm+bxxx/PLbfckiT505/+lE022WSueRtuuGFuvvnmPPLII0mSadOm5aGHHspaa62Vp59+OrfffnuS5JVXXsnMmTOz7LLL5pVXXnnPXoYNG5bjjz++5f7FF1/M+uuvn3/84x957rnnMmvWrJx99tn53Oc+lyRZaaWVcv/99+ett97KBRdc8L5rXZAePogF26fo31BK+USSoUn2SpKqqt5M8mYpZfskmzWGnZ7k+iQ/SLJ9kj9WsyNN4xq7AXVvjL2qqqoXGu+9KsnwUsr1SbpUVTWuUf9jkh2SXNZaawIAAAAAAICFaZVVVkmSdOvWLTvuuGNuu+22HHLIIbnyyiuTzD7G65JLLmkZP3ny5Oy444754x//mM985jNzvWvffffNGmuskQMPPHCu+tVXX52f/exn+cc//vGuHYAAYPqMWZk0dtuF+r4lO7Z/zzGDBw/OXnvtlfXXXz9JMmrUqCy//PJZa621cvzxx2fvvfdOv379st9++7XMKaWka9euOe2007LHHnvkjTfeSJIcdthhWXPNNXPuuefmO9/5Tl5//fUstdRSufrqq7P55ptn7NixaWpqyg9/+MPstttu7+rlv//7v7P//vtnwIABad++fQ499NDstNNOGTt2bDbffPNUVZVtt90222+/fZJk7Nix2W677dK1a9cMGTIkr7766nuuddCgQWnfvn3WXnvt7LXXXjnooIM+0N/nO5XW2iqolNKU5MQkEzN7N587khyQ5MmqqpZrjClJXqyqarlSysVJxlZVdVPj2TWZHQDaLMmSVVUd1qj/vySvZ3ZAaGxVVZ9v1DdN8oOqqrZ7r76GDBlSNTc3L9S1AgAAAAAAwAc1bdq0vPXWW1l22WUzbdq0DBs2LD/60Y8yePDgdOvWLW+99Vb22muvbLbZZtl7770zderUfO5zn2v5BeSc/vu//zv3339/zjvvvLRr93+HeowfPz677LJLLr/88qyxxhqLeokAfATdf//96du3b1u38YEMHDgwF110UXr16tXWrbSKef0zKaXcUVXVkHeObc2juzokGZzkd1VVrZNkWv7vmK4kSWP3ntZJGs2hlLJvKaW5lNI8ZcqU1v4cAAAAAAAAvK9nnnkmm2yySdZee+2sv/762XbbbTN8+PCcffbZWXPNNdOnT5+svPLKGTlyZJLkuOOOyyOPPJKf/OQnaWpqSlNTU5599tlMnjw5P/vZzzJx4sQMHjw4TU1NOemkk5Ik3/ve9/Lqq6/mS1/6UpqamvLFL36xLZcMAB/YsGHDMnDgwMU25PNBteaOPp9KMq6qqp6N+00zO+izepLNqqp6unE01/VVVa1VSjmhcX12Y/yDmb2bz2aN8d9o1E/I7N18rk9yXVVVfRr1PeYcNz929AEAAAAAAAAAPq7quKPPwnDqqafmmGOOmau28cYb5/jjj2+jjv7PB9nRp0NrNVFV1f+WUp4opaxVVdWDSbbM7GO8JibZM8nYxp8XNqZclOTbpZRzkmyQ5KVGGOiKJD8vpSzfGLdVkh9WVfVCKeXlUsqGSW5N8rUkx7bWegAAAAAAAAAAqKeRI0e27JJXZ60W9Gn4TpKzSimdkjyaZGRmHxf251LKPkn+lWTXxthLk2yT5JEkrzXGphHo+WmS2xvjflJV1QuN628lOS3JUkkua/wAAAAAAAAAAMBip1WDPlVVTUjyrm2EMnt3n3eOrZLsP5/3nJLklHnUm5MM+HBdAgAAAAAAAADAR1+7tm4AAAAAAAAAAAB4f4I+AAAAAAAAAABQA4I+AAAAAAAAsAhNnzFrsfoOADU3Y/oif99nP/vZhfvNj7jTTjstTz311EJ5V4eF8hYAAAAAAABggSzZsX16jr6k1b8zaey2rf4NABYDHZdMxnxi4b1vzEvvO+R//ud/Ft73auC0007LgAEDsvLKK3/od9nRBwAAAAAAAACARaZz585Jkuuvvz6f+9znsv3226d3794ZPXp0zjrrrKy//voZOHBg/vnPfyZJ/v73v2eDDTbIOuusk89//vN55plnkiRTpkzJsGHD0r9//4waNSqf/vSn89xzzyVJzjzzzKy//vppamrKN77xjcyaNf+d7i6//PIMHjw4a6+9drbccsskyQsvvJAddtghgwYNyoYbbpi77747STJmzJgceeSRLXMHDBiQSZMmZdKkSenbt2++/vWvp3///tlqq63y+uuv5/zzz09zc3NGjBiRpqamvP766x/q707QBwAAAAAAAACANnHXXXfl97//fe6///6cccYZeeihh3Lbbbdl1KhROfbYY5Mkm2yyScaNG5fx48dn9913zxFHHJEk+fGPf5wtttgi9913X3bZZZc8/vjjSZL7778/5557bm6++eZMmDAh7du3z1lnnTXP70+ZMiVf//rX85e//CV33XVXzjvvvCTJoYcemnXWWSd33313fv7zn+drX/va+67l4Ycfzv7775/77rsvyy23XP7yl79kl112yZAhQ3LWWWdlwoQJWWqppT7U35ejuwAAAAAAAAAAaBPrrbdeunfvniT5zGc+k6222ipJMnDgwFx33XVJksmTJ2e33XbL008/nTfffDO9evVKktx000254IILkiTDhw/P8ssvnyS55pprcscdd2S99dZLkrz++uvp1q3bPL8/bty4DB06tOWdK6ywQsu7//KXvyRJtthiizz//PN5+eWX33MtvXr1SlNTU5Jk3XXXzaRJkz7w38f7EfQBAAAAAAAAAKBNLLHEEi3X7dq1a7lv165dZs6cmST5zne+k4MPPjhf/OIXc/3112fMmDHv+c6qqrLnnnvm8MMPX+j9dujQIW+99VbL/fTp01uu51xL+/btP/QxXfPi6C4AAAAAAAAAAD6yXnrppayyyipJktNPP72lvvHGG+fPf/5zkuTKK6/Miy++mCTZcsstc/755+fZZ59Nkrzwwgv517/+Nc93b7jhhrnhhhvy2GOPtYxNkk033bTluK/rr78+K664Yrp06ZKePXvmzjvvTJLceeedLfPey7LLLptXXnnlA697XuzoAwAAAAAAAADwcTVjejLmpYX7vo5LLrz3JRkzZky+9KUvZfnll88WW2zREq459NBDs8cee+SMM87IRhttlE996lNZdtlls+KKK+awww7LVlttlbfeeisdO3bM8ccfn09/+tPvenfXrl1z4oknZqeddspbb72Vbt265aqrrsqYMWOy9957Z9CgQVl66aVbAkY777xz/vjHP6Z///7ZYIMNsuaaa75v/3vttVe++c1vZqmllsott9ySpZZa6t/+uyhVVf3bk+toyJAhVXNzc1u3AQAAAAAAwMdYz9GXtPo3Jo3dttW/AUD93H///enbt29bt7FQvPHGG2nfvn06dOiQW265Jfvtt18mTJjQ1m19YPP6Z1JKuaOqqiHvHGtHHwAAAAAAAAAAaufxxx/PrrvumrfeeiudOnXKH/7wh7ZuqdUJ+gAAAAAAAAAAUDtrrLFGxo8fv8DjN9hgg7zxxhtz1c4444wMHDhwYbfWagR9AAAAAAAAAABY7N16661t3cKH1q6tGwAAAAAAAAAAYNGpqqqtW6Dhg/6zEPQBAAAAAAAAAPiYWHLJJfP8888L+3wEVFWV559/PksuueQCz3F0FwAAAAAAAADAx0SPHj0yefLkTJkypa1bIbODVz169Fjg8YI+AAAAAAAAAAAfEx07dkyvXr3aug3+TY7uAgAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGhD0AQAAAAAAAACAGmjVoE8pZVIp5Z5SyoRSSnOjtkIp5apSysONP5dv1Esp5TellEdKKXeXUgbP8Z49G+MfLqXsOUd93cb7H2nMLa25HgAAAAAAAAAAaCuLYkefzauqaqqqakjjfnSSa6qqWiPJNY37JNk6yRqNn32T/C6ZHQxKcmiSDZKsn+TQt8NBjTFfn2Pe8NZfDgAAAAAAAAAALHptcXTX9klOb1yfnmSHOep/rGYbl2S5Ukr3JF9IclVVVS9UVfVikquSDG8861JV1biqqqokf5zjXQAAAAAAAAAAsFhp7aBPleTKUsodpZR9G7WVqqp6unH9v0lWalyvkuSJOeZObtTeqz55HnUAAAAAAAAAAFjstHbQZ5OqqgZn9rFc+5dShs75sLETT9XKPaSUsm8ppbmU0jxlypSF+u5Zs2ZlnXXWyXbbbZckueaaazJ48OA0NTVlk002ySOPPJIk+f3vf5+BAwe21CdOnJgkmTRpUpZaaqk0NTWlqakp3/zmN5Mkr732Wrbddtv06dMn/fv3z+jRo+fdAAAAAAAAAAAAHwutGvSpqurJxp/PJrkgyfpJnmkcu5XGn882hj+ZZNU5pvdo1N6r3mMe9Xn1cWJVVUOqqhrStWvXD7usuRxzzDHp27dvy/1+++2Xs846KxMmTMiXv/zlHHbYYUmSL3/5y7nnnnsyYcKEfP/738/BBx/cMuczn/lMJkyYkAkTJuT3v/99S/2QQw7JAw88kPHjx+fmm2/OZZddtlB7BwAAAAAAAACgPlot6FNKWaaUsuzb10m2SnJvkouS7NkYtmeSCxvXFyX5WpltwyQvNY74uiLJVqWU5Uspyzfec0Xj2cullA1LKSXJ1+Z41yIxefLkXHLJJRk1alRLrZSSl19+OUny0ksvZeWVV06SdOnSpWXMtGnTMrvl+Vt66aWz+eabJ0k6deqUwYMHZ/Lkye85BwAAAAAAAACAxVeHVnz3SkkuaARaOiT5U1VVl5dSbk/y51LKPkn+lWTXxvhLk2yT5JEkryUZmSRVVb1QSvlpktsb435SVdULjetvJTktyVJJLmv8LDIHHnhgjjjiiLzyyisttZNOOinbbLNNllpqqXTp0iXjxo1reXb88cfn6KOPzptvvplrr722pf7YY49lnXXWSZcuXXLYYYdl0003nes7U6dOzd///vcccMABrb8oAAAAAAAAAAA+klptR5+qqh6tqmrtxk//qqp+1qg/X1XVllVVrVFV1effDu1Us+1fVdVnqqoaWFVV8xzvOqWqqtUbP6fOUW+uqmpAY863q6qqWms973TxxRenW7duWXfddeeq/+pXv8qll16ayZMnZ+TIkXMd0bX//vvnn//8Z37xi1+0HOnVvXv3PP744xk/fnyOPvrofPnLX27ZEShJZs6cmT322CPf/e5307t370WzOAAAAAAAAAAAPnJac0efxdrNN9+ciy66KJdeemmmT5+el19+Odtuu20eeOCBbLDBBkmS3XbbLcOHD3/X3N133z377bdfkmSJJZbIEksskSRZd91185nPfCYPPfRQhgwZkiTZd999s8Yaa+TAAw9cNAsDAAAAAAAAAOAjqdV29FncHX744Zk8eXImTZqUc845J1tssUUuvPDCvPTSS3nooYeSJFdddVX69u2bJHn44Ydb5l5yySVZY401kiRTpkzJrFmzkiSPPvpoHn744Zade/77v/87L730Un79618vwpUBAAAAAAAAAPBRZEefhahDhw75wx/+kJ133jnt2rXL8ssvn1NOOSVJctxxx+Xqq69Ox44ds/zyy+f0009Pktxwww350Y9+lI4dO6Zdu3b5/e9/nxVWWCGTJ0/Oz372s/Tp0yeDBw9Oknz729/OqFGj2mx9AAAAAAAAAAC0nVJVVVv3sEgNGTKkam5ubus2AAAAAAAA+BjrOfqSVv/GpLHbtvo3AIDWUUq5o6qqIe+sO7oLAAAAAAAAAABqQNAHAAAAAAAAAABqQNAHAAAAAAAAAABqQNAHAAAAAAAAAABqQNAHAAAAAAAAAABqQNDnA5o+Y9Zi9R0AAAAAAAAAAOqhQ1s3UDdLdmyfnqMvafXvTBq7bat/AwAAAAAAAACA+rCjDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ICgDwAAAAAAAAAA1ECrB31KKe1LKeNLKRc37nuVUm4tpTxSSjm3lNKpUV+icf9I43nPOd7xw0b9wVLKF+aoD2/UHimljG7ttQAAAAAAAAAAQFtZFDv6HJDk/jnuf5HkV1VVrZ7kxST7NOr7JHmxUf9VY1xKKf2S7J6kf5LhSX7bCA+1T3J8kq2T9EuyR2MsAAAAAAAAAAAsdlo16FNK6ZFk2yQnNe5Lki2SnN8YcnqSHRrX2zfu03i+ZWP89knOqarqjaqqHkvySJL1Gz+PVFX1aFVVbyY5pzEWAAAAAAAAAAAWO629o8+vk3w/yVuN+08mmVpV1czG/eQkqzSuV0nyRJI0nr/UGN9Sf8ec+dUBAAAAAAAAAGCx02pBn1LKdkmerarqjtb6xgfoZd9SSnMppXnKlClt3Q4AAAAAAAAAAHxgrbmjz8ZJvlhKmZTZx2ptkeSYJMuVUjo0xvRI8mTj+skkqyZJ4/knkjw/Z/0dc+ZXf5eqqk6sqmpIVVVDunbt+uFXBgAAAAAAAAAAi1irBX2qqvphVVU9qqrqmWT3JNdWVTUiyXVJdmkM2zPJhY3rixr3aTy/tqqqqlHfvZSyRCmlV5I1ktyW5PYka5RSepVSOjW+cVFrrQcAAAAAAAAAANpSh/cfstD9IMk5pZTDkoxPcnKjfnKSM0opjyR5IbODO6mq6r5Syp+TTEwyM8n+VVXNSpJSyreTXJGkfZJTqqq6b5GuBAAAAAAAAAAAFpFFEvSpqur6JNc3rh9Nsv48xkxP8qX5zP9Zkp/No35pkksXYqsAAAAAAAAAAPCR1GpHdwEAAAAAAAAAAAuPoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANTABw76lFKWL6UMao1mAAAAAAAAAACAeVugoE8p5fpSSpdSygpJ7kzyh1LK0a3bGgAAAAAAAAAA8LYF3dHnE1VVvZxkpyR/rKpqgySfb722AAAAAAAAAACAOS1o0KdDKaV7kl2TXNyK/QAAAAAAAAAAAPOwoEGfHye5IskjVVXdXkrpneTh1msLAAAAAAAAAACYU4cFHPd0VVWD3r6pqurRUsrRrdQTAAAAAAAAAADwDgu6o8+xC1gDAAAAAAAAAABawXvu6FNK2SjJZ5N0LaUcPMejLknat2ZjAAAAAAAAAADA/3m/o7s6JencGLfsHPWXk+zSWk0BAAAAAAAAAABze8+gT1VV/0jyj1LKaVVV/WsR9QQAAAAAAAAAALzD++3o87YlSiknJuk555yqqrZojaYAAAAAAAAAAIC5LWjQ57wkv09yUpJZrdcOAAAAAAAAAAAwLwsa9JlZVdXvWrUTAAAAAAAAAABgvtot4Li/l1K+VUrpXkpZ4e2fVu0MAAAAAAAAAABosaA7+uzZ+PN7c9SqJL0XbjsAAAAAAAAAAMC8LFDQp6qqXq3dCAAAAAAAAAAAMH8LFPQppXxtXvWqqv64cNsBAAAAAAAAAADmZUGP7lpvjuslk2yZ5M4kgj4AAAAAAAAAALAILOjRXd+Z876UslySc1qjIQAAAAAAAAAA4N3a/ZvzpiXptTAbAQAAAAAAAAAA5m+BdvQppfw9SdW4bZ+kb5I/t1ZTAAAAAAAAAADA3BYo6JPkyDmuZyb5V1VVk1uhHwAAAAAAAAAAYB4W6Oiuqqr+keSBJMsmWT7Jm+83p5SyZCnltlLKXaWU+0opP27Ue5VSbi2lPFJKObeU0qlRX6Jx/0jjec853vXDRv3BUsoX5qgPb9QeKaWM/kArBwAAAAAAAACAGlmgoE8pZdcktyX5UpJdk9xaStnlfaa9kWSLqqrWTtKUZHgpZcMkv0jyq6qqVk/yYpJ9GuP3SfJio/6rxriUUvol2T1J/yTDk/y2lNK+lNI+yfFJtk7SL8kejbEAAAAAAAAAALDYWaCgT5L/SrJeVVV7VlX1tSTrJ/l/7zWhmu3Vxm3Hxk+VZIsk5zfqpyfZoXG9feM+jedbllJKo35OVVVvVFX1WJJHGt9fP8kjVVU9WlXVm0nOaYwFAAAAAAAAAIDFzoIGfdpVVfXsHPfPL8jcxs47E5I8m+SqJP9MMrWqqpmNIZOTrNK4XiXJE0nSeP5Skk/OWX/HnPnVAQAAAAAAAABgsdNhAcddXkq5IsnZjfvdklz6fpOqqpqVpKmUslySC5L0+Xea/LBKKfsm2TdJVltttbZoAQAAAAAAAAAAPpT3DPqUUlZPslJVVd8rpeyUZJPGo1uSnLWgH6mqamop5bokGyVZrpTSobFrT48kTzaGPZlk1SSTSykdknwis3cOerv+tjnnzK/+zu+fmOTEJBkyZEi1oH0DAAAAAAAAAMBHxfsdv/XrJC8nSVVVf62q6uCqqg7O7N15fv1eE0spXRs7+aSUslSSYUnuT3Jdkl0aw/ZMcmHj+qLGfRrPr62qqmrUdy+lLFFK6ZVkjSS3Jbk9yRqllF6llE5Jdm+MBQAAAAAAAACAxc77Hd21UlVV97yzWFXVPaWUnu8zt3uS00sp7TM7UPTnqqouLqVMTHJOKeWwJOOTnNwYf3KSM0opjyR5IbODO6mq6r5Syp+TTEwyM8n+jSPBUkr5dpIrkrRPckpVVfe974oBAAAAAAAAAKCG3i/os9x7PFvqvSZWVXV3knXmUX80yfrzqE9P8qX5vOtnSX42j/qlSS59rz4AAAAAAAAAAGBx8H5HdzWXUr7+zmIpZVSSO1qnJQAAAAAAAAAA4J3eb0efA5NcUEoZkf8L9gxJ0inJjq3YFwAAAAAAAAAAMIf3DPpUVfVMks+WUjZPMqBRvqSqqmtbvTMAAAAAAAAAAKDF+x3dlSSpquq6qqqObfwI+ZAkeeKJJ7L55punX79+6d+/f4455pgkyW677ZampqY0NTWlZ8+eaWpqmmve448/ns6dO+fII49sqe29997p1q1bBgwYMNfYF154IcOGDcsaa6yRYcOG5cUXX2z1dQEAAAAAAAAAfBQtUNAH5qVDhw456qijMnHixIwbNy7HH398Jk6cmHPPPTcTJkzIhAkTsvPOO2ennXaaa97BBx+crbfeeq7aXnvtlcsvv/xd3xg7dmy23HLLPPzww9lyyy0zduzYVl0TAAAAAAAAAMBHlaAP/7bu3btn8ODBSZJll102ffv2zZNPPtnyvKqq/PnPf84ee+zRUvvb3/6WXr16pX///nO9a+jQoVlhhRXe9Y0LL7wwe+65Z5Jkzz33zN/+9rdWWAkAAAAAAAAAwEefoA8LxaRJkzJ+/PhssMEGLbUbb7wxK620UtZYY40kyauvvppf/OIXOfTQQxf4vc8880y6d++eJPnUpz6VZ555ZuE2DgAAAAAAAABQE4I+fGivvvpqdt555/z6179Oly5dWupnn332XLv5jBkzJgcddFA6d+78b32nlJJSyofuFwAAAAAAAACgjjq0dQPU24wZM7LzzjtnxIgR2WmnnVrqM2fOzF//+tfccccdLbVbb701559/fr7//e9n6tSpadeuXZZccsl8+9vfnu/7V1pppTz99NPp3r17nn766XTr1q1V1wMAAAAAAAAA8FEl6MO/raqq7LPPPunbt28OPvjguZ5dffXV6dOnT3r06NFSu/HGG1uux4wZk86dO79nyCdJvvjFL+b000/P6NGjc/rpp2f77bdfuIsAAAAAAAAAAKgJR3fxb7v55ptzxhln5Nprr01TU1Oamppy6aWXJknOOeecuY7tej977LFHNtpoozz44IPp0aNHTj755CTJ6NGjc9VVV2WNNdbI1VdfndGjR7fKWgAAAAAAAAAAPupKVVVt3cMiNWTIkKq5uflDvaPn6EsWUjfzN2nstq3+DQAAAAAAANqG3zcBAO+llHJHVVVD3lm3ow8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8AAAAAAAAAANSAoA8f2PQZsxar7wAAAAAAAAAA1EGHtm6A+lmyY/v0HH1Jq39n0thtW/0bAAAAAAAAAAB1YUcfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACoAUEfAAAAAAAAAACogVYL+pRSVi2lXFdKmVhKua+UckCjvkIp5apSysONP5dv1Esp5TellEdKKXeXUgbP8a49G+MfLqXsOUd93VLKPY05vymllNZaDwAAAAAAAAAAtKXW3NFnZpL/rKqqX5INk+xfSumXZHSSa6qqWiPJNY37JNk6yRqNn32T/C6ZHQxKcmiSDZKsn+TQt8NBjTFfn2Pe8FZcDwAAAAAAAAAAtJlWC/pUVfV0VVV3Nq5fSXJ/klWSbJ/k9Maw05Ps0LjePskfq9nGJVmulNI9yReSXFVV1QtVVb2Y5KokwxvPulRVNa6qqirJH+d4FwAAAAAAAAAALFZac0efFqWUnknWSXJrkpWqqnq68eh/k6zUuF4lyRNzTJvcqL1XffI86gAAAAAAAAAAsNhp9aBPKaVzkr8kObCqqpfnfNbYiadaBD3sW0ppLqU0T5kypbU/BwAAAAAAAAAAC12rBn1KKR0zO+RzVlVVf22Un2kcu5XGn8826k8mWXWO6T0atfeq95hH/V2qqjqxqqohVVUN6dq164dbFAAAAAAAAAAAtIFWC/qUUkqSk5PcX1XV0XM8uijJno3rPZNcOEf9a2W2DZO81Dji64okW5VSli+lLJ9kqyRXNJ69XErZsPGtr83xLgAAAAAAAAAAWKx0aMV3b5zkq0nuKaVMaNT+vyRjk/y5lLJPkn8l2bXx7NIk2yR5JMlrSUYmSVVVL5RSfprk9sa4n1RV9ULj+ltJTkuyVJLLGj8AAAAAAAAAALDYabWgT1VVNyUp83m85TzGV0n2n8+7TklyyjzqzUkGfIg2AQAAAAAAAACgFlrt6C4AAAAAAAAAAGDhEfQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAACAxdjee++dbt26ZcCAAS21733ve+nTp08GDRqUHXfcMVOnTk2SPP/889l8883TuXPnfPvb357rPWeffXYGDhyYQYMGZfjw4Xnuuedanh177LHp06dP+vfvn+9///uLZF0A8HEk6AMAAAAAAACLsb322iuXX375XLVhw4bl3nvvzd13350111wzhx9+eJJkySWXzE9/+tMceeSRc42fOXNmDjjggFx33XW5++67M2jQoBx33HFJkuuuuy4XXnhh7rrrrtx333055JBDFs3CAOBjSNAHAAAAAAAAFmNDhw7NCiusMFdtq622SocOHZIkG264YSZPnpwkWWaZZbLJJptkySWXnGt8VVWpqirTpk1LVVV5+eWXs/LKKydJfve732X06NFZYoklkiTdunVr7SUBwMeWoM9H3Ly2UjzvvPPSv3//tGvXLs3NzS31GTNmZM8998zAgQPTt2/fluT1E088kc033zz9+vVL//79c8wxx7TMmd+2jAAAAAAAAHw8nHLKKdl6663fc0zHjh3zu9/9LgMHDszKK6+ciRMnZp999kmSPPTQQ7nxxhuzwQYb5HOf+1xuv/32RdE2AHwsCfp8xM1rK8UBAwbkr3/9a4YOHTpX/bzzzssbb7yRe+65J3fccUdOOOGETJo0KR06dMhRRx2ViRMnZty4cTn++OMzceLEJPPflhEAAAAAAIDF389+9rN06NAhI0aMeM9xM2bMyO9+97uMHz8+Tz31VAYNGtTye6WZM2fmhRdeyLhx4/LLX/4yu+66a6qqWhTtA8DHjqDPR9y8tlLs27dv1lprrXeNLaVk2rRpmTlzZl5//fV06tQpXbp0Sffu3TN48OAkybLLLpu+ffvmySefTDL/bRkBAAAAAABYvJ122mm5+OKLc9ZZZ6WU8p5jJ0yYkCT5zGc+k1JKdt111/zP//xPkqRHjx7ZaaedUkrJ+uuvn3bt2uW5555r7fYB4GNJ0Gcxsssuu2SZZZZJ9+7ds9pqq+WQQw55V0ho0qRJGT9+fDbYYIN3zV+QbRkBAAAAAACov8svvzxHHHFELrrooiy99NLvO36VVVbJxIkTM2XKlCTJVVddlb59+yZJdthhh1x33XVJZh/j9eabb2bFFVdsveYB4GOsQ1s3wMJz2223pX379nnqqafy4osvZtNNN83nP//59O7dO0ny6quvZuedd86vf/3rdOnSZa65C7otIwAAAAAAAPWyxx575Prrr89zzz2XHj165Mc//nEOP/zwvPHGGxk2bFiS2Sc//P73v0+S9OzZMy+//HLefPPN/O1vf8uVV16Zfv365dBDD83QoUPTsWPHfPrTn85pp52WJNl7772z9957Z8CAAenUqVNOP/30990hCAD49wj6LEb+9Kc/Zfjw4enYsWO6deuWjTfeOM3Nzendu3dmzJiRnXfeOSNGjMhOO+0017y3t2W85ppr/IcuAAAAAACAxczZZ5/9rto+++wz3/GTJk2aZ/2b3/xmvvnNb76r3qlTp5x55pn/dn8AwIJzdNdiZLXVVsu1116bJJk2bVrGjRuXPn36pKqq7LPPPunbt28OPvjgueZ80G0ZAQAAAAAAAABoG4I+H3F77LFHNtpoozz44IPp0aNHTj755FxwwQXp0aNHbrnllmy77bb5whe+kCTZf//98+qrr6Z///5Zb731MnLkyAwaNCg333xzzjjjjFx77bVpampKU1NTLr300iTJt7/97bzyyisZNmxYmpqa5pnCBgAAAAAAAACg7Tm66yNuXlspJsmOO+74rlrnzp1z3nnnvau+ySabpKqqeb7nkUce+XANAgAAAAAAAACwSNjRBwAAAAAAAAAAakDQBwAAAAAAAAAAakDQBwAAAAAAABZHM6YvXt8BANKhrRsAAAAAAAAAWkHHJZMxn2j974x5qfW/AQAksaMPAAAAAAAAAADUgqDPR5WtFAEAAAAAAAAAmIOjuz6qbKUIAAAAAAAAAMAc7OgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA1IOgDAAAAAAAAAAA10GpBn1LKKaWUZ0sp985RW6GUclUp5eHGn8s36qWU8ptSyiOllLtLKYPnmLNnY/zDpZQ956ivW0q5pzHnN6WU0lprAQAAAAAAAACAttaaO/qclmT4O2qjk1xTVdUaSa5p3CfJ1knWaPzsm+R3yexgUJJDk2yQZP0kh74dDmqM+foc8975LQAAAAAAAAAAWGy0WtCnqqobkrzwjvL2SU5vXJ+eZIc56n+sZhuXZLlSSvckX0hyVVVVL1RV9WKSq5IMbzzrUlXVuKqqqiR/nONdAAAAAAAAAACw2GnNHX3mZaWqqp5uXP9vkpUa16skeWKOcZMbtfeqT55HHQAAAAAAAAAAFkuLOujTorETT7UovlVK2beU0lxKaZ4yZcqi+CQAAAAAAAAAACxUizro80zj2K00/ny2UX8yyapzjOvRqL1Xvcc86vNUVdWJVVUNqapqSNeuXT/0IgAAAAAAAAAAYFFb1EGfi5Ls2bjeM8mFc9S/VmbbMMlLjSO+rkiyVSll+VLK8km2SnJF49nLpZQNSyklydfmeBcAAAAAAAAAACx2OrTWi0spZyfZLMmKpZTJSQ5NMjbJn0sp+yT5V5JdG8MvTbJNkkeSvJZkZJJUVfVCKeWnSW5vjPtJVVUvNK6/leS0JEsluazxAwAAAAAAAAAAi6VWC/pUVbXHfB5tOY+xVZL95/OeU5KcMo96c5IBH6ZHAAAAAAAAAACoi0V9dBcAAAAAAAAAAPBvEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAAAAAAAAAIAaEPQBAACARWTq1KnZZZdd0qdPn/Tt2ze33HJLzjvvvPTv3z/t2rVLc3Pzu+Y8/vjj6dy5c4488siW2q9+9av0798/AwYMyB577JHp06cvymUAAAAAAG1E0AcAAAAWkQMOOCDDhw/PAw88kLvuuit9+/bNgAED8te//jVDhw6d55yDDz44W2+9dcv9k08+md/85jdpbm7Ovffem1mzZuWcc85ZVEsAAAAAANpQh7ZuAAAAAD4OXnrppdxwww057bTTkiSdOnVKp06dstxyy813zt/+9rf06tUryyyzzFz1mTNn5vXXX0/Hjh3z2muvZeWVV27FzgEAAACAjwo7+gAAAMAi8Nhjj6Vr164ZOXJk1llnnYwaNSrTpk2b7/hXX301v/jFL3LooYfOVV9llVVyyCGHZLXVVkv37t3ziU98IltttVVrtw8AAAAAfAQI+gAAAMAiMHPmzNx5553Zb7/9Mn78+CyzzDIZO3bsfMePGTMmBx10UDp37jxX/cUXX8yFF16Yxx57LE899VSmTZuWM888s7XbBwAAAAA+AhzdBQAAsAj07Nkzyy67bNq3b58OHTqkubk5Y8aMyR/+8Id07do1SfLzn/8822yzTc4666z88pe/bJl79913584770xTU1NL7Ytf/GIeffTR3HvvvYt6KfybevTokR49emSDDTZIkuyyyy7vGfS59dZbc/755+f73/9+pk6dmnbt2mXJJZfMSiutlF69erX862annXbK//zP/+QrX/nKIlkHAAAAANB2BH2onXn9guRtRx11VA455JBMmTIlK664Ys4666z84he/SFVVWXbZZfO73/0ua6+9dsv4WbNmZciQIVlllVVy8cUXt8VyAAD4GLnuuuuy4oorzlU76KCDcsghh8xVGzFiREaMGJEkueeee7LDDjvMFfL561//+q5dXvjo+9SnPpVVV101Dz74YNZaa61cc8016dev33zH33jjjS3XY8aMSefOnfPtb387t956a8aNG5fXXnstSy21VK655poMGTJkUSwBAAAAAGhjgj7U0rx+QfLEE0/kyiuvzGqrrdZS69WrV/7xj39k+eWXz2WXXZZ99903t956a8vzY445Jn379s3LL7+8yHoHAIAP4uyzz87uu+/ecv/qq6/m6KOPzoknnphdd921DTvj33HsscdmxIgRefPNN9O7d++ceuqpueCCC/Kd73wnU6ZMybbbbpumpqZcccUV833HBhtskF122SWDBw9Ohw4dss4662TfffddhKsAAAAAANqKoA+LjYMOOihHHHFEtt9++5baZz/72ZbrDTfcMJMnT265nzx5ci655JL813/9V44++uhF2isAAB8/pZRstdVWKaXkG9/4Rksw47jjjssf//jHDBkyJEcddVSWX375ueade+65ufDCC1vu/9//+3/5z//8zyy99NKLtH8Wjqamprl2JU2SHXfcMTvuuON7zhszZsxc9z/+8Y/z4x//eGG3BwAAAAB8xLVr6wbgg3r7FyTrrrtuTjzxxCTJhRdemFVWWWWuY7ne6eSTT87WW2/dcn/ggQfmiCOOSLt2/s8AAIDWd9NNN+XOO+/MZZddluOPPz433HBD9ttvv/zzn//MhAkT0r179/znf/7nXHNuvfXWLL300hkwYECSZMKECfnnP//5vqEQAAAAAAAWT3b0oXZuuummrLLKKnn22WczbNiw9OnTJz//+c9z5ZVXznfOddddl5NPPjk33XRTkuTiiy9Ot27dsu666+b6669fRJ0DAPBxtsoqqyRJunXrlh133DG33XZbhg4d2vL861//erbbbru55pxzzjnZY489Wu5vueWWNDc3p2fPnpk5c2aeffbZbLbZZv4zLQAAAADAx4StTKidd/6C5B//+Ecee+yxrL322unZs2cmT56cwYMH53//93+TJHfffXdGjRqVCy+8MJ/85CeTJDfffHMuuuii9OzZM7vvvnuuvfbafOUrX2mzNQEAsHibNm1aXnnllZbrK6+8MgMGDMjTTz/dMuaCCy5o2bknSd566638+c9/zu67795S22+//fLUU09l0qRJuemmm7LmmmsK+QAAAAAAfIzY0YdamTZtWt56660su+yyLb8g+dGPfpRnn322ZUzPnj3T3NycFVdcMY8//nh22mmnnHHGGVlzzTVbxhx++OE5/PDDkyTXX399jjzyyJx55pmLfD0AAHw8PPPMMy3Hbc2cOTNf/vKXM3z48Hz1q1/NhAkTUkpJz549c8IJJ7TMueGGG7Lqqqumd+/ebdU2AAAAAAAfMYI+1Mr8fkEyPz/5yU/y/PPP51vf+laSpEOHDmlubl4kvQIAwNt69+6du+666131M844Y75zNttss4wbN26+z3v27Jl77713ofQHAAAAAEA9CPpQK/P7BcmcJk2a1HJ90kkn5aSTTnrP8Ztttlk222yzhdAdAADAbNNnzMqSHdsvNt8BAAAAAD4aBH0AAABgIVuyY/v0HH1Jq39n0thtW/0b0FpmzZqVIUOGZJVVVsnFF1+cffbZJ83NzamqKmuuuWZOO+20dO7cOY8//nj23HPPTJ06NbNmzcrYsWOzzTbbZMaMGRk1alTuvPPOzJw5M1/72tfywx/+sK2XBQAAANCq2rV1AwAAAAB8/BxzzDHp27dvy/2vfvWr3HXXXbn77ruz2mqr5bjjjkuSHHbYYdl1110zfvz4nHPOOS3Hc5933nl54403cs899+SOO+7ICSecMNcuvwAAAACLI0EfAAAAABapyZMn55JLLsmoUaNaal26dEmSVFWV119/PaWUJEkpJS+//HKS5KWXXsrKK6/cUp82bVpmzpyZ119/PZ06dWp5BwAAAMDiStAHAAAAFkOzZs3KOuusk+222y5Jctxxx2X11VdPKSXPPffcu8bffvvt6dChQ84///yW2uOPP56tttoqffv2Tb9+/eyWwkJz4IEH5ogjjki7dnP/T1MjR47Mpz71qTzwwAP5zne+kyQZM2ZMzjzzzPTo0SPbbLNNjj322CTJLrvskmWWWSbdu3fPaqutlkMOOSQrrLDCIl8LAAAAwKIk6AMAALAQTZ8xa7H6DvX1zmORNt5441x99dX59Kc//a6xs2bNyg9+8INstdVWc9W/9rWv5Xvf+17uv//+3HbbbenWrVur983i7+KLL063bt2y7rrrvuvZqaeemqeeeip9+/bNueeemyQ5++yzs9dee2Xy5Mm59NJL89WvfjVvvfVWbrvttrRv3z5PPfVUHnvssRx11FF59NFHF/VyAAAAABapDm3dAMzXjOlJxyUXn+8AAPCxsGTH9uk5+pJW/86ksdu2+jeor7ePRfqv//qvHH300UmSddZZZ77jjz322Oy88865/fbbW2oTJ07MzJkzM2zYsCRJ586dW7dpPjZuvvnmXHTRRbn00kszffr0vPzyy/nKV76SM888M0nSvn377L777jniiCMycuTInHzyybn88suTJBtttFGmT5+e5557Ln/6058yfPjwdOzYMd26dcvGG2+c5ubm9O7duy2XBwAAANCqBH346Oq4ZDLmE63/nTEvtf43AAAAFqG3j0V65ZVX3nfsk08+mQsuuCDXXXfdXEGfhx56KMstt1x22mmnPPbYY/n85z+fsWPHpn379q3ZOh8Dhx9+eA4//PAkyfXXX58jjzwyZ5xxRh555JGsvvrqqaoqF110Ufr06ZMkWW211XLNNddkr732yv3335/p06ena9euWW211XLttdfmq1/9aqZNm5Zx48blwAMPbMOVAQAAALQ+R3cBAADAYuS9jkWalwMPPDC/+MUv0q7d3P8TwcyZM3PjjTfmyCOPzO23355HH300p512Wit0DElVVdlzzz0zcODADBw4ME8//XR+9KMfJUmOOuqo/OEPf8jaa6+dPfbYI6eddlpKKdl///3z6quvpn///llvvfUycuTIDBo0qI1XAgAAANC67OgDAAAAi5H3OxbpnZqbm7P77rsnSZ577rlceuml6dChQ3r06JGmpqaWY5B22GGHjBs3Lvvss88iWwuLv8022yybbbZZktn/2p2Xfv36zfNZ586dc95557VmewAAAAAfOYI+AAAAsBiZ17FI8wv5JMljjz3Wcr3XXntlu+22yw477JBZs2Zl6tSpmTJlSrp27Zprr702Q4YMafX+AQAAAID5c3QXAADAYmL69OlZf/31s/baa6d///459NBDkyQjRozIWmutlQEDBmTvvffOjBkzkiQvvvhidtxxxwwaNCjrr79+7r333iTJE088kc033zz9+vVL//79c8wxx7TZmlh4fvOb36RHjx6ZPHlyBg0alFGjRr3n+Pbt2+fII4/MlltumYEDB6aqqnz9619fRN0CAAAAAPNiRx8AAIDFxBJLLJFrr702nTt3zowZM7LJJptk6623zogRI1p2dPnyl7+ck046Kfvtt19+/vOfp6mpKRdccEEeeOCB7L///rnmmmvSoUOHHHXUURk8eHBeeeWVrLvuuhk2bFj69evXxivkg5rzWKTvfve7+e53v/ue40877bS57ocNG5a77767lboDAAAAAD4oO/oAAAAsJkop6dy5c5JkxowZmTFjRkop2WabbVJKSSkl66+/fiZPnpwkmThxYrbYYoskSZ8+fTJp0qQ888wz6d69ewYPHpwkWXbZZdO3b988+eSTbbMoAAAAAABaCPoAAAAsRmbNmpWmpqZ069Ytw4YNywYbbNDybMaMGTnjjDMyfPjwJMnaa6+dv/71r0mS2267Lf/6179aQkBvmzRpUsaPHz/XewDqYH7HEN51113ZaKONMnDgwPzHf/xHXn755SSz/3/knnvumYEDB6Zv3745/PDDW941derU7LLLLunTp0/69u2bW265pU3WBAAAACDoAwAAsBhp3759JkyYkMmTJ+e2227Lvffe2/LsW9/6VoYOHZpNN900STJ69OhMnTo1TU1NOfbYY7POOuukffv2LeNfffXV7Lzzzvn1r3+dLl26LPK1APU3fcasNvvO28cQTpw4MePGjcvxxx+fiRMnZtSoURk7dmzuueee7LjjjvnlL3+ZJDnvvPPyxhtv5J577skdd9yRE044IZMmTUqSHHDAARk+fHgeeOCB3HXXXenbt+8iWRcAAADAO3Vo6wYAAABY+JZbbrlsvvnmufzyyzNgwID8+Mc/zpQpU3LCCSe0jOnSpUtOPfXUJElVVenVq1d69+6dZPbOFjvvvHNGjBiRnXbaqU3WANTfkh3bp+foS1r9O5PGbvuuWvfu3dO9e/ckcx9D+NBDD2Xo0KFJkmHDhuULX/hCfvrTn6aUkmnTpmXmzJl5/fXX06lTp3Tp0iUvvfRSbrjhhpx22mlJkk6dOqVTp06tviYAAACAebGjDwAAwGJiypQpmTp1apLk9ddfz1VXXZU+ffrkpJNOyhVXXJGzzz477dr9338NnDp1at58880kyUknnZShQ4emS5cuqaoq++yzT/r27ZuDDz64LZYCsFDNeQxh//79c+GFFyaZvYvPE088kSTZZZddsswyy6R79+5ZbbXVcsghh2SFFVbIY489lq5du2bkyJFZZ511MmrUqEybNq0tlwMAAAB8jAn6AAAALCaefvrpbL755hk0aFDWW2+9DBs2LNttt12++c1v5plnnslGG22Upqam/OQnP0mS3H///RkwYEDWWmutXHbZZTnmmGOSJDfffHPOOOOMXHvttWlqakpTU1MuvfTStlwa8zNj+uL1HWgF7zyG8JRTTslvf/vbrLvuunnllVdadue57bbb0r59+zz11FN57LHHctRRR+XRRx/NzJkzc+edd2a//fbL+PHjs8wyy2Ts2LFtvCoAAADg48rRXQAAAIuJQYMGZfz48e+qz5w5c57jN9poozz00EPvqm+yySapqmqh90cr6LhkMuYTrf+dMS+1/jegFczrGMI+ffrkyiuvTJI89NBDueSS2UeL/elPf8rw4cPTsWPHdOvWLRtvvHGam5szdOjQ9OjRIxtssEGS2Tv/CPoAAAAAbcWOPgAAAAAsduZ3DOGzzz6bJHnrrbdy2GGH5Zvf/GaSZLXVVsu1116bJJk2bVrGjRuXPn365FOf+lRWXXXVPPjgg0mSa665Jv369VvEqwEAAACYzY4+AAAAACx23j6GcODAgWlqakqS/PznP8/DDz+c448/Pkmy0047ZeTIkUmS/fffPyNHjkz//v1TVVVGjhyZQYMGJUmOPfbYjBgxIm+++WZ69+6dU089tU3WBAAAACDoAwAAAMBi572OITzggAPeVevcuXPOO++8eY5vampKc3PzQu0PAAAA4N/h6C4AAAAAAAAAAKgBQR8AAAAAAAAAAKgBQR8AAIA6mjF98foOAAAAAADvq0NbNwAAAMC/oeOSyZhPtP53xrzU+t8AAAAAAGCB2NEHAAAAYBHae++9061btwwYMKClNmbMmKyyyippampKU1NTLr300pZnhx9+eFZfffWstdZaueKKK1rql19+edZaa62svvrqGTt27CJdAwAAAABtQ9AHAAAAYBHaa6+9cvnll7+rftBBB2XChAmZMGFCttlmmyTJxIkTc8455+S+++7L5Zdfnm9961uZNWtWZs2alf333z+XXXZZJk6cmLPPPjsTJ05c1Ev56HCcIQAAAPAx4eguAAAAgEVo6NChmTRp0gKNvfDCC7P77rtniSWWSK9evbL66qvntttuS5Ksvvrq6d27d5Jk9913z4UXXph+/folmb1r0MUXX5xu3brl3nvvTZJ873vfy9///vd06tQpn/nMZ3LqqadmueWWy2233ZZ99903SVJVVcaMGZMdd9wxSTJ16tSMGjUq9957b0opOeWUU7LRRhstzL+OhcNxhgAAAMDHhB19AAAAAD4CjjvuuAwaNCh77713XnzxxSTJk08+mVVXXbVlTI8ePfLkk0/Ot/62ee0aNGzYsNx77725++67s+aaa+bwww9PkgwYMCDNzc2ZMGFCLr/88nzjG9/IzJkzkyQHHHBAhg8fngceeCB33XVX+vbt22rrBwAAAOD9CfoAAAAAtLH99vv/27vvOCvK64/jny8gIiJiYgcRFQtSxYYlRpPYsAW7YomiiSXWaGLiL7HERGNi1Ni7WGIvGAsWFCyJIiBF7C0Bey8ICnh+f8zc5bLsUpd5lp3v+/XixZ2Zu8zZ4ZaZZ85zzuG8/vrrjB49mpVWWolf/epXC/TvbbHFFnzve9+bad0222xDixZZcec+ffowceJEAFq3bl2zfsqUKUgC4PPPP+fxxx9nwIABALRs2ZJ27dotUFxmZmZmZmZmZrZgnOhjZmZmZmZmZpbYCiusQPPmzWnWrBmHHnpoTXuu9u3bM2HChJrnTZw4kfbt29e7fm5dffXVbL/99jXLzzzzDF27dqV79+5ceumltGjRgjfffJPllluOgw46iPXWW49DDjmESZMmNcBva7Wdf/75dOvWja5du3LeeecBWau1ddZZhx49etCvXz8+++yzmuePHTuWTTbZpOb/bMqUKWkCNzMzMzMzM7PCOdHHzMzMFrqXX36ZXr161fxp27Yt5513HqNHj6ZPnz706tWLDTbYoOaGVkRw9NFH07lzZ3r06MGoUaMS/wZmZmZmC9e7775b8/iuu+6iW7duAOy8887cfPPNfPPNN7z55pu8+uqrbLTRRmy44Ya8+uqrvPnmm3z77bfcfPPN7LzzznO1rz/96U+0aNGC/v3716zbeOONGT9+PM8++yxnnnkmU6ZMYdq0aYwaNYrDDz+c5557jiWXXJKzzjqrYX9x4/nnn+eKK65g+PDhjBkzhnvvvZfXXnut3lZr06ZNY7/99uPSSy9l/PjxDB06lMUWW2yu9lXfeXnFOeecgyQ++uijmnVDhw6lV69edO3alR/+8IcN+rubmZmZmZmZ2bxrkToAMzMza/rWXnttRo8eDcD06dNp3749/fr149BDD+WUU05h++235/777+fXv/41Q4cO5YEHHuDVV1/l1Vdf5ZlnnuHwww/nmWeeSftLmJmZmTWQffbZh6FDh/LRRx/RoUMHTjvtNIYOHcro0aORRKdOnbjssssA6Nq1K3vuuSfrrrsuLVq04KKLLqJ58+YAXHjhhWy77bZMnz6dgw8+mK5du85x39deey333nsvQ4YMqWnRVa1Lly60adOG559/ng4dOtChQwc23nhjAHbffXcn+iwEL774IhtvvDGtW7cG4Ic//CF33nknv/71r2ue06dPH26//XYAHnroIXr06EHPnj0B+P73vz/X+6rvvBxgwoQJPPTQQ3Ts2LHm+Z999hlHHHEEgwcPpmPHjnzwwQcL9LuamZmZmZmZ2YJzoo/ZQtKpUyeWWmopmjdvTosWLRgxYgR77bUXL7/8MpANlrVr165mgG3s2LH84he/4IsvvqBZs2Y8++yztGrVKuFvYGa2cAwZMoQ11liDVVddFUl88cUXAHz++eesvPLKAAwaNIgDDjgASfTp04fPPvuMd999l5VWWmmO/35dn78nnngi//rXv2jZsiVrrLEG11xzDe3atWPq1KkccsghjBo1imnTpnHAAQfw29/+dqH+/mZmZmY33XTTLOsGDBhQ7/NPPvlkTj755FnW9+3bl759+871fgcPHszZZ5/NsGHDapJKAN58801WWWUVWrRowX//+19eeuklOnXqxLLLLssqq6zCyy+/zNprr82QIUNYd91153p/Nne6devGySefzMcff8wSSyzB/fffzwYbbDDTc66++mr22msvAF555RUkse222/Lhhx+y9957z5QUNLeqz8sBjjvuOM4++2x22WWXmuf885//ZNddd61J/ll++eXn99c0MzMzM5tvdY353nbbbZx66qm8+OKLDB8+fKZzaN9zM7Omzq27zBaixx57jNGjRzNixAgAbrnlFkaPHs3o0aPZbbfd2HXXXYEFK7sN2Sy89dZbjx133BHIWt6cfPLJrLXWWnTp0oV//OMfMz3/2WefpUWLFjWzAc3MinTzzTezzz77AHDeeedx4oknssoqq3DCCSfUtCN4++23WWWVVWp+pkOHDrz99ttzvY/an7/1tT247bbb+Oabbxg3bhwjR47ksssu46233mqg39TMzMwsnX322YdNNtmEl19+mQ4dOnDVVVfxy1/+ki+//JKtt96aXr16cdhhhwHw5JNP0rNnT3r16kW/fv24+OKLWXbZZQG44IIL6N+/Pz169GD06NH87ne/S/lrNUldunThN7/5Ddtssw3bbbcdvXr1qqnaBLO2Wps2bRpPPvkkN954I08++SR33XUXQ4YMmef9Vp+XDxo0iPbt29dUCap45ZVX+PTTT9lyyy1Zf/31ue666xbgNzWbO7XHufr378/aa69Nt27dOPjgg5k6depMz/c4l5lZ01f7u2HAgAH07NmTHj16sPvuu/PVV18B8L///Y+tttqK9dZbjx49enD//ffP1b8/ZcoUNtpoI3r27EnXrl055ZRTgCwxunfv3vTq1YvNN9+c1157DYBvvvmGvfbai86dO7Pxxht7PLEgtcd8u3Xrxp133skWW2wx0/Ma+p5bxdFHH02bNm1qluf39WZm1hCc6GOWQERw66231gyo1VV2u3pQb07OP/98unTpUrN87bXXMmHCBF566SVefPFF9t5775pt06dPrxlANFvY6rtAqm+Q7sYbb6RHjx50796dTTfdlDFjxqQM3xaCb7/9lnvuuYc99tgDgEsuuYRzzz2XCRMmcO655852JvuC2GabbWjRIitk2KdPHyZOnAiAJCZNmsS0adOYPHkyLVu2pG3btgslBjMzM7Mi3XTTTbz77rtMnTqViRMnMmDAAF577TUmTJhQMwHl0ksvBWD//fdn/PjxjB49mlGjRvHTn/605t/p1asXI0aMYOzYsdx9990ss8wyiX6jpm3AgAGMHDmSxx9/nGWWWYa11loLmNFq7cYbb6xptdahQwe22GILll12WVq3bk3fvn0ZNWrUPO2v+rz866+/5s9//jOnn376LM+bNm0aI0eO5L777uPBBx/kj3/8I6+88sqC/8Jms1F7nKt///689NJLjBs3jsmTJ3PllVfWbJvfca76xit+8IMf0KtXL3r16sXKK69c83no8Qozs7Rqfzece+65jBkzhrFjx9KxY0cuvPBCAM444wz23HNPnnvuOW6++WaOOOKIufr3F198cR599FHGjBnD6NGjGTx4ME8//TSHH344N954I6NHj2bffffljDPOAOCqq65imWWW4bXXXuO4447jN7/5TcP/0jZHXbp0Ye21155lfUPfcwMYMWIEn3766Uzr5vf1Zja/JkyYwFZbbcW6665L165dOf/88wEYPXo0ffr0oVevXmywwQYMHz4c8DlsU+dEH7OFRBLbbLMN66+/PpdffvlM25544glWWGEF1lxzTWDmstu9e/fm7LPPnuv9TJw4kfvuu49DDjmkZt0ll1zCH/7wB5o1y97i1aW1L7jgAnbbbTeX27ZC1HeBVN8g3WqrrcawYcMYN24cv//97/n5z38+1/s6+OCDWX755enWrVvNujFjxrDJJpvQvXt3dtppp5oWUVOnTuXAAw+ke/fudOnSpaa6iy18DzzwAL1792aFFVYAYODAgTXVzfbYY4+aE9D27dszYcKEmp+bOHEi7du3n6t9zO7zF7K2B9tvvz0Au+++O0suuSQrrbQSHTt25IQTTuB73/veAv2O9anvJPz3v/89PXr0oFevXmyzzTa88847NT8zdOhQevXqRdeuXfnhD3+4QPv55JNP2HrrrVlzzTXZeuutZ7ownZ/9mJmZmVnD+eCDD4BsVvCdd97JvvvuW9Nq7Z577pmp1dq2227LuHHj+Prrr5k2bRrDhg2b55Zq1eflr7/+Om+++SY9e/akU6dOTJw4kd69e/Pee+/RoUMHtt12W5ZcckmWXXZZtthiCw8O20JV1zhX3759kYQkNtpoo5qJGzD/41z1jVc88cQTNcmQm2yySc316oKMV5iZ2YKp67uhMlEvIpg8eXJNQrSkmjHgzz//nJVXXnmu9iGpplLL1KlTmTp1as13T13/3qBBgzjwwAOBbHxxyJAhREQD/LZWnzmN+VZr6Htu06dP58QTT5zl35nf19v8qG/M97bbbqNr1640a9asptJRxdixY9lkk03o2rUr3bt3Z8qUKXO1r7rutdSXSALFji3XFdupp55K+/bta5K1K5WVHn74YdZff326d+/O+uuvz6OPPrpQYytCixYtOOecc3jhhRd4+umnueiii3jhhRf49a9/zSmnnMLo0aM5/fTTa1o7L8g5bH2vuYpzzjkHSXz00UdA9nl89NFH07lzZ3r06DHPk1Fs3jnRx2whefLJJxk1ahQPPPAAF110EY8//njNtptuuqmmmg8sWNntY489lrPPPrsmqQfg9ddf55ZbbmGDDTZg++2359VXXwWyVjh33XUXhx9+eAP9lmazV98FUn2DdJtuumnNDOHqqitz42c/+xmDBw+ead0hhxzCWWedxbhx4+jXrx9//etfAbdrSqn259/KK6/MsGHDAHj00UdrEiB33nlnrrvuOiKCp59+mqWXXpqVVlpprvYxu8/f2m0Phg8fTvPmzXnnnXd48803Oeecc3jjjTca6tedSX0n4SeeeCJjx45l9OjR7LjjjjUzqT/77DOOOOII7rnnHsaPH89tt922QPs566yz+PGPf8yrr77Kj3/8Y84666wF2k9DOPfcc+natSvdunVjn332YcqUKXNsPzmvXn755ZqLvF69etG2bVvOO++8mu21L0gWtrouRhdGElZR+2kInTp1onv37jUDBQAnnngi66yzDj169KBfv3589tlnDbKvuS29bGZmVrTddtuNddddl5122omLLrqIdu3a1dtqbZllluH4449nww03pFevXvTu3ZsddthhnvZXfV7evXt3PvjgA9566y3eeustOnTowKhRo1hxxRXZZZddePLJJ5k2bRpff/01zzzzzCyzmxeWus4RZneDYX7VPj+or0WHFaOuca6KqVOncv3117PddtsBCzbOVd94RcUXX3zBo48+WlPRZ0HGK8zMbMHU991w0EEHseKKK/LSSy9x1FFHAdkN/xtuuIEOHTrQt29fLrjggrnez/Tp0+nVqxfLL788W2+9NRtvvDFXXnklffv2pUOHDlx//fWcdNJJQPYdtMoqqwDZWNzSSy/Nxx9/3EC/sdVldmO+tTX0PbcLL7yQnXfeeZbx6QV5vc2r+sZ8F0b7srrutdSXSFL02HJdsQEcd9xxNcnaffv2BWDZZZflX//6F+PGjWPgwIHsv//+c72fusZWZzdeeeaZZ9K5c2fWXnttHnzwwfn/BedgpZVWonfv3gAstdRSdOnShbfffrvepLMFOYet7zUHWRLQQw89RMeOHWue/8ADD/Dqq6/y6quvcvnll/tedAGc6GO2kFQqTyy//PL069evZvBp2rRp3Hnnney11141z53fstv33nsvyy+/POuvv/5M67/55htatWrFiBEjOPTQQzn44IOB7ATlL3/5S52DJUUbPHgwa6+9Np07d6652bwo72d+lOUY1HWBVFF7kK7aVVddVVN1ZW5sscUWs1RieeWVV2pOcLfeemvuuOMOIG27pnltZ9bY9zMvJk2axMMPP1wzIxLgiiuu4Fe/+hU9e/bkd7/7Xc1sjL59+7L66qvTuXNnDj30UC6++OK53k99n791tT345z//yXbbbcdiiy3G8ssvz2abbTbLzIeGUt9JePVrb9KkSTPFtuuuu9acLM/tDNX69lM90+jAAw/k7rvvXqD9LKi3336bf/zjH4wYMYLnn3+e6dOnc/PNN8+2/eT8WHvttWsu8kaOHEnr1q3p168fUPcFycJW18XowkjCKmo/DaV2j/Wtt96a559/nrFjx7LWWms1WOW1uS29bMWq7zuroqESsYraj82f2c0Uu+CCC1hnnXXo2rVrzWDiwtgPFJ8AalbxxBNP8MILLzBmzBh+/OMfA9Tbag1gv/32Y/z48Tz//PPzNEMZ6j4vr0+XLl3Ybrvt6NGjBxtttBGHHHLITIPdC1vtc4T6bjAsiNrnB/W16EitDGMI9Y1zVRxxxBFsscUW/OAHPwAWfJxrduMVd999Nz/+8Y/rHCuY1/GKBTE/FVsb834aQhneC3PiY+BjUCaz+2645ppreOedd+jSpQu33HILkCUz/+xnP2PixIncf//97L///nz33Xdzta/mzZszevRoJk6cyPDhw3n++ec599xzuf/++5k4cSIHHXQQxx9/fIP+fguqTO+F+sZ869KQ99zeeecdbrvttppksmoL8nqbV/WN+S6M9mV13WupL5Gk6LHlumKrz3rrrVcTZ9euXZk8eTLffPPNXP1sXWOr9Y1XvvDCC9x8882MHz+ewYMHc8QRRzB9+vR5+K3mz1tvvcVzzz3HxhtvzHnnnceJJ57IKquswgknnFDnWOq8nsPW95qDLLHq7LPPnilRftCgQRxwwAFIok+fPnz22We8++67C/hbzrsyfS6mv9tv1gRNmjSJL7/8subxQw89VDMQ9sgjj7DOOuvQoUOHmufPb9ntp556invuuYdOnTqx99578+ijj7LffvvRoUOHmgG7fv36MXbsWCC7kbX33nvTqVMnbr/9do444oiaG71Fmj59OkceeSQPPPAAL7zwAjfddFNNFuiiuJ/5UaZjUNcFUkXtQbqKxx57jKuuuoq//OUvC7Tvrl27MmjQICCr4lNpBVVku6ba5rWdWWPfz7xYcskl+fjjj1l66aVr1m2++eaMHDmSMWPG8Mwzz9RcREnioosu4vXXX2fcuHE1M3jnpL7P3/raHnTs2LGmZOekSZN4+umnWWeddRrqV65X9Uk4wMknn8wqq6zCjTfeWFPR55VXXuHTTz9lyy23ZP311+e6665boP28//77NbNOVlxxRd5///0G28/8qiTbVWaIr7zyyrNtP7mghgwZwhprrMGqq64K1H1BsrDVdTG6MJKwitrPwrLNNtvQokULoOFmTM9L6WUrVn3fWdCwiVhF7cfmT30zxR577DEGDRrEmDFjGD9+PCeccMJC2Q8snATQKVMX/uBekfuxpqGu8/Jqb731Fssuu2zN8oknnsgLL7zA888/z7HHHltQlHVr6PYIdZ0fFNmCYW6VZQyhvnEugNNOO40PP/yQv//97zXPX9BxrtmNV9SuRlvRUOMVc2teK7Y29v0sqLK8F2bHx8DHoGxm990A2Wf53nvvXTPB86qrrmLPPfcEYJNNNmHKlCnznMTfrl07ttpqKx544AHGjBlTM36311578e9//xvIkk4qY83Tpk3j888/5/vf//4C/77zokzvhdndc6tLQ95z69q1K6+99hqdO3emU6dOfP3113Tu3BlomNfb/Kg9tlyXBWlfVpf6EklSji1Xu/DCC+nRowcHH3xwneM7d9xxB71792bxxRefq3+vrrHV+sYrBw0axN57783iiy/OaqutRufOnRuk8ujsfPXVV+y2226cd955tG3blksuuYRzzz2XCRMmcO655zJgwICZnr+g57DVr7lBgwbRvn37miSyiupKZ5Al3FUSg4pSps9FcKKP2ULx/vvvs/nmm9OzZ0822mgjdthhh5qKJTfffPMsAwXzW3b7zDPPZOLEibz11lvcfPPN/OhHP+KGG27gpz/9KY899hgAw4YNY6211gLgzTffrCnFvfvuu3PxxRfXlCAu0vDhw+ncuTOrr746LVu2ZO+9965JxlgU9zM/yngMKhdIlSzougbpIOsbe8ghhzBo0KAFvji6+uqrufjii1l//fX58ssvadmyJVBsu6ba5rWdWWPfT2NT3+dvfW0PjjzySL766iu6du3KhhtuyEEHHUSPHj0Waoy1T8Ihayk2YcIE+vfvz4UXXghkgwQjR47kvvvu48EHH+SPf/wjr7zyygLtp6LyGmiI/cyv9u3bc8IJJ9CxY0dWWmklll56abbZZpt62082hOrv4PouSFIoKgmrMSZ7wZx7rF999dUNMmN6XkovW7Hq+85q6ESsovZj86e+mWKXXHIJJ510Us1g3IImI87rjLQF1Wqx5nQ66b6F/qfVYnM3M9NsUVLXOcLczFSdF3WdH9TXoiOlsowh1DfOdeWVV/Lggw9y0003zfR/1VDjXLXHKz766COGDx8+y9hcQ45XzK15rdja2PezoMryXpgdHwMfg7Kp67vh+uuvr2mtGRHcc889NRP3OnbsWNOi6cUXX2TKlCkst9xyc9zPhx9+WNOGZ/LkyTz88MN06dKFzz//vGacrLIOYOedd2bgwIEA3H777fzoRz8qdCIZlOu9UN+Y71133UWHDh34z3/+ww477MC2224LNOw9t08//ZT33nuv5pyjdevWNa+/+X29LYjZjflWW5D2ZXWpL5Ek1dhytcMPP5zXX3+d0aNHs9JKK/GrX/1qpu3jx4/nN7/5DZdddlmD7bN6vLLoBJepU6ey22670b9//5qiDwMHDqx5vMcee8yUaLSg57DVr7kWLVrw5z//uWaycmNTps9FgBaF79GsBFZffXXGjBlT57Zrr722zvX77bffTJnoC+Kkk06if//+nHvuubRp06bQih1zo64vvWeeeWaR3c/8KMsx+PDDD1lsscVo165dzQXSb37zm5pBuiFDhsw0SPe///2PXXfdleuvv74mQW1BrLPOOjz00ENAdiP7vvvuA+pv17T66qsv8D7nxvTp01l//fV57bXXOPLII+tsZ1a7jURj3k9jUt/nb+Xiq7Y2bdoU2qaorpPwav3796dv376cdtppdOjQge9///ssueSSLLnkkmyxxRaMGTNmrt4bde1nhRVW4N1332WllVbi3XffrblRuiD7WRCffvopgwYN4s0336Rdu3bsscce3HDDDTO1n7zzzjs5+OCDeeKJJxZ4f99++y333HMPZ555Jl9//TV//vOfaz4fGpO6krCGDBnC5MmT2WSTTejTp0+D/N8UtZ+58eSTT9K+fXs++OADtt56a9ZZZ52atot/+tOfaNGiBf3791+gfVSXXh46dCgwo/RyZdnSqus76/zzz2/wRKyi9mMLpnqm2IknnsgTTzzBySefTKtWrfjb3/7Ghhtu2OD7aUwJoPNl6hRYrFXT2Y8ZdZ8j3H777Zx77rnstttu3HrrrQwYMIBHHnlkvv79us4PgJoWHRtvvDF//etfOf7445OPq5RlDKE+hx12GKuuuiqbbLIJALvuuit/+MMfFujfrG+8ArIbtjvuuCOtWs34vGvo8Yr5MTcVWxel/cyPsr8XwMcAfAwsS+458MAD+eKLL4gIevbsySWXXAJkrXgPPfRQzj33XCRx7bXXzlUCzrvvvsuBBx7I9OnT+e6779hzzz3ZcccdueKKK9htt91o1qwZyyyzDFdffTUAAwYMYP/996dz585873vf4+abb16ov3NdyvReqG/Mt1+/fvTr16/On2nIe271md/X2/ya09hyter2ZUBN+7JKm+B5NXDgwJp7CXvssUdNRcxUY8vVVlhhhZrHhx56KDvuuGPN8sSJE+nXrx/XXXcda6yxRoPsr6HGK+dHRDBgwAC6dOkyUyvBlVdemWHDhrHlllvy6KOPsuaaawILfg5b+zU3btw43nzzzZqxk4kTJ9K7d2+GDx8+U6WzyrZKy72ilOlzEZzoY9ZkbLnllmy55ZZANhOpktBQn/oSjswaUn0XSC1atKhzkO7000/n448/5ogjjgCy0tEjRoyY7/1/8MEHLL/88nz33XecccYZNVVcKu2a9t9//5p2TUWWoK+UB//ss8/o168fzz//fE2p0framTXm/djcqe8k/NVXX6058R40aFDNDKRddtmFX/7yl0ybNo1vv/2WZ555huOOO26+91OZaXTSSScxcOBAdtlllwXaz4J65JFHWG211Wpmuey66678+9//nqX95EEHHdQg+3vggQfo3bs3K6ywwmwvSFZcccUG2d+8KCoJq7Ele1XU1WN9iy224Nprr+Xee+9lyJAhCzxIUim9fP/99zNlyhS++OILunbtyuKLL15TbrlSerm+xEBbuGp/Zz3++OMLJRGrqP3Y/Ks9O3HatGl88sknPP300zz77LPsueeevPHGGwv8uVDXjLTGmAA61xZrBafW3YapQZ36+cLfh1murnOE+m4wzI+6zg922GEHXnrppZladFQqNFuxqse5pk2bNsfnz+s4V33jFZBVAq1dyamhxyvm1dxWbF1U9mNmNj+qvxueeuqpOp+z7rrr1rttdnr06MFzzz03y/r6EklatWpV6ORBK171663aV199VfN4fl9v86O+Md/6bLvttpx99tl8/fXXtGzZkmHDhi3QmG99iSSpxparVcY7Ae66666aeyCfffYZO+ywA2eddRabbbZZg+yrrvHKIhNcnnrqKa6//nq6d+9Or169APjzn//MFVdcwTHHHMO0adNo1apVTUXUBTmHres11717dz744IOa53Tq1IkRI0aw7LLLsvPOO3PhhRey995788wzz7D00kt7Ut1C5kQfMytcUV96jSF7tD5lOQb1XSDVN0h35ZVXzvdMyX322YehQ4fy0Ucf0aFDB0477TS++uorLrroIiBLIqgkDBx55JEcdNBBdO3alYgopF1TXarLg3fr1q2mnVlDlpAscj82e/WdhF911VW8/PLLNGvWjFVXXZVLL70UgC5durDddtvRo0cPmjVrxiGHHDLb3tNz2s9JJ53EnnvuyVVXXcWqq67KrbfeukD7WVAdO3bk6aef5uuvv2aJJZZgyJAhbLDBBrRt25bHHnuM1VZbbab2kwvqpptuqmnbNbsLkhSKSsJqbMlekPVV/+6771hqqaVqeqz/4Q9/YPDgwZx99tkMGzaM1q1bL/B+zjzzzJrWHkOHDuVvf/sb995770zPadOmjZN8GoHKd9Zjjz3Ga6+9ttASsYraj82bumYnVhJAK21HmzVrxkcffbRA5dDnZUZaigRQM6v/HKG+Gwzzo67zg7vvvpsVV1yRV155hbXWWmumFh0plWUMoUj1jVcAdSYAL8h4xYKal4qti8J+FoTfCz4G4GNgVuH3QrnUN+b7zTffcNRRR/Hhhx+yww470KtXLx588MGZ2pdJom/fvnPVvgzqvtdSXyJJ0WPLdcU2dOhQRo8ejSQ6depUc9/jwgsv5LXXXuP000+vaTX10EMPzff5TH3jlTvvvDP77rsvxx9/PO+88w6vvvoqG2200YL/snXYfPPNiYg6t40cOXKWdQtyDlvfa65v3751Pr9v377cf//9dO7cmdatW3PNNdfM134XRNk+F53oY2aF23DDDXn11Vd58803ad++PTfffDP//Oc/F9n9zA8fg4Z300031bn+mGOOmWVd0e2aqs1rO7PGvh+be/WdhNd3Ygxw4okncuKJJzbIfoB6+zDPz34W1MYbb8zuu+9O7969adGiBeuttx4///nPmTx5coO3n5w0aRIPP/xwo0huq+tidGEkYRW1nwX1/vvv18yMmzZtGvvuuy/bbbcdnTt35ptvvmHrrbcGoE+fPjVJcNb01Ped9d5779U8pyESsYraj82f+mYn/vSnP+Wxxx5jq6224pVXXuHbb79doMTMeZ2RZtYQpkydTqvFmjeZ/RShvnOENm3a1HmDoaG0aNGi3hYdKXkMobzmtWJrY9/PgvJ7wccAfAzMKvxeKJfZjfk2dPuy+u611JVIAsWOLdcV24ABA+p87v/93//xf//3f/O1n7rGVs8888w6xyu7du3KnnvuybrrrkuLFi246KKLaN580b8um91rruKtt96qeSypZuJ9KmX7XHSij5kVrkWLFlx44YVsu+22TJ8+nYMPPpiuXbsusvuZHz4G5TWv7cwa+37ANy5s/p122mmcdtppM61bfPHF59h+cl4tueSSfPzxx/Vur74gWdjqu1Bu6CSsovazoOrrsb4wEy3mpvSyFWt2rTMWxf3Y/KlvptjBBx/MwQcfTLdu3WjZsiUDBw5coLYh8zojzawhtFqsOZ1Oatjzm7q89ccfAwWcL0+dkrWLW4jqO0fYfPPN673BsCCqzw/qa9GRkscQymteK7Y29v0sKL8XfAzAx8Cswu8Fs4VnXhKKAE4++WROPvnkhRmSzYWyfS5qTplYTc0GG2wQC9o/uZDBmbN2gFOXXuj74dTP5+vHfAzMzKxaYd8LZmaLCJ8vG/h1YBm/DnwMwMfAzMzM6uZzBDMzM5sdSSMjYoPa692zw2w+TJk6vUntx8xskTB1StPaj5mZmZmZmZmZmVnJ+Z6bmdm8W+Rbd0naDjifrDbxlRFxVuKQrARcdtvMLIHFWnn2kZmZmZmZ2SLKbZ/NzKy2oj6zY+oUVMQ9EN9rmS++52ZmNu8W6UQfSc2Bi4CtgYnAs5LuiYgX0kZm1kB8U9saOQ/SmZmZmZmZmdnc8E08MzOrrbDvBrcvM/A9N2v0fM/N5sUinegDbAS8FhFvAEi6GdgFcKKPmVkBPEhnlvHsIzMzMzMzs0bCN/HMzMxsEePxZSe5gO+52bxZ1BN92gMTqpYnAhsnisXMzBYWD9JZI+fZR8VeIPmC1McAGvcxMDMzs5n5/MDMzMzMrH4eX3aSS6F8z61JUESkjmG+Sdod2C4iDsmX9wc2johf1nrez4Gf54trAy8XGuj8WRb4KHUQifkY+BiAjwH4GICPAfgYgI8B+BhU+Dj4GICPAfgYgI8B+BiAjwH4GICPAfgYgI8B+BiAjwH4GFT4OPgYgI8B+BiAjwH4GICPAfgYgI8B+BjAonUMVo2I5WqvXNQr+rwNrFK13CFfN5OIuBy4vKigGoKkERGxQeo4UvIx8DEAHwPwMQAfA/AxAB8D8DGo8HHwMQAfA/AxAB8D8DEAHwPwMQAfA/AxAB8D8DEAHwPwMajwcfAxAB8D8DEAHwPwMQAfA/AxAB8D8DGApnEMmqUOYAE9C6wpaTVJLYG9gXsSx2RmZmZmZmZmZmZmZmZmZmZm1uAW6Yo+ETFN0i+BB8ma6V0dEeMTh2VmZmZmZmZmZmZmZmZmZmZm1uAW6UQfgIi4H7g/dRwLwSLVamwh8THwMQAfA/AxAB8D8DEAHwPwMajwcfAxAB8D8DEAHwPwMQAfA/AxAB8D8DEAHwPwMQAfA/AxqPBx8DEAHwPwMQAfA/AxAB8D8DEAHwPwMYAmcAwUEaljMDMzMzMzMzMzMzMzMzMzMzOzOWiWOgAzMzMzMzMzMzMzMzMzMzMzM5szJ/qYNRLKrJI6DjNrHCR1Tx2DmZmZmZmZNV6SmknaM3UcZmZmZmZmViwn+lijJKl16hiKFlkfvftTx2HWWEharY51G6aIJZGLJQ2XdISkpVMHk4Kk1pJ+L+mKfHlNSTumjsuKJ2lXSX+XdI6kfqnjscZFUsvUMZiZmVkakppLOi51HKlExHfAr1PHYWZmjVueGNo2dRxFcjJszTHYNHUcKUk6RlLbfKL9VZJGSdomdVxmZg1BWW6BNQaSlgQmR8R3ktYC1gEeiIipiUMrTH7ScSXQJiI6SuoJ/CIijkgcWiEkDQQujIhnU8eSkqTFgd2ATkCLyvqIOD1VTEWTtBxwKLMeg4NTxVQ0SaOAnSLi7Xz5h2Tvj9JUupG0JnAwsAcwHLgmIh5OG1VxJN0CjAQOiIhueRLovyOiV9rIilX2z0RJFwOdgZvyVXsBr0fEkemiKl6e4PRoRHyeL7cDtoyIu1PGVTRJQ4GfRcRb+fJGwBUR0TNlXEWRdAxwDfAl2TnzesBJEfFQ0sAKJun7wKnAZkAATwKnR8THKeMqiqSzgTOAycBgoAdwXETckDSwguXXzJcAK+TnCT2AnSPijMShFUbS8XWs/hwYGRGjCw7HCibpe7PbHhGfFBVLapKGR8RGqeNIRdJZwEfALcCkyvoyvQYAJN0JXEU2lvpd6nhSkbQ5sGZEXJOPLbWJiDdTx1UkSWsAEyPiG0lbkp0rXRcRn6WMq0iSjgJuiIhPU8dSNEm7zm57RNxZVCypSfoncBgwHXgWaAucHxF/TRpYgSSNiIgNUseRkqTnImK91HGkImlMRPSUtC3wC+D3wPUR0TtxaIWT1BxYgZnHl/+XLqLilfU8SdJsX+8RMaqoWFJrarkYTvRpRCSNBH4ALAM8RXby9W1E9E8aWIEkPQPsDtxTOfmQ9HxEdEsbWTEkvUR2M/O/ZIMzIiv20yNpYAWTNJh8gJrsQgSAiDgnWVAFk/Rv4AlmPQZ3JAuqYHn1nouBnYDewJnAjhExIWlgBctPwH8K/AP4guxz4XdlGJioXIxXX5BWLs5Sx1aksn8m5t+NXfLKd0hqBoyPiC5pIyuWpNG1k9zKOFiTD8ycT/aZ2B7YHjikLBekHqDKSHoYeByoJLb0J0t8+0m6qIpT+TzIEwB3BI4HHi/h9+Mw4ETgsjJeO0LNzZsNgH/lq3YExpIlB98WEWcnCm2hkzSOLNGvTmW4hpb0JtkxUB2bIyJWLzikZCSdCyzGrIkuZTk/qOvmRKleAwCSfgIcBPQBbiObKPNy2qiKJekUsu+FtSNiLUkrk30fbJY4tEJJGk12HDqRVU8fBHSNiL4JwyqUpDOAvYFRwNXAg5Vr6qZO0jWz2Rwlm0RZuW7oTza2ehJZQniTP0+qcDIsSPob8B/gzrJ8DlSTNDYiekg6HxgaEXeVdDztKOAU4H2gkhBdqnuPZT5PkvTYbDZHRPyosGASa2q5GC3m/BQrkCLia0kDgIsj4uz8wqRUImKCNNM41fT6ntsEbZs6gEaiQ0RslzqIxFpHxG9SB5FSRDwr6WjgIWAK8JOI+DBxWIXJZ6UfBOwAPExW3WhUfgL6H6DJJ/oA30pagvwmTj4r75u0ISVR9s/E14COZEmwAKvk68qmrpa7pTuXj4gHJR1G9rn4EbBeRLyXOKwiVU6S+5Il+IxXrRPnklgpIv5YtXyGpL2SRVO8ynt/B7JBqc/L+TKgdUQMr/W7T0sVTCIdgN4R8RXUDFzeB2xBliDcZBN9yJKaACoV/q7P/14kB+fmR0TM0uq4xHrlf1dXvAygFAPWfi1kIuIR4BFlra/3yR9PAK4gq2yySM7SnUf9yCo+jgKIiHckLZU2pCS+i4hpeVL0BRFxgaTnUgdVpIj4P0m/B7YhG1u6UNKtwFUR8Xra6BauiDgodQyNyGKSFiObQHhhREwt4XVD5Tqxuip0AGVKhv0F2eSQ6ZImM2NyeVlauY2U9BCwGvDb/HuxjJX/jiFLcClFJeR6lPY8KSK2Sh1DI9KkcjFKd3OgkZOkTcgGpgbk65onjCeFCXn7rshPQo8BXkwcU2Ei4r8AkpYHWiUOJ6V/S+oeEeNSB5LQvZL6RsT9qQMpmqR/MfPs3NZk1UyukkRE7JwmssJdQNaW5XcRMbmyMj8B/b90YRXqFLKWJKtIupGsRcvPkkaURik/E6s+C5YCXpQ0PF/emKyVXdmMkPR34KJ8+Uiym7ilkg9W70l2E7sHMFTSryLivrSRFcYDVJmHJO0N3Jov7w48mDCeot2bVzubDByel5uekjimFD7Kk4ArCcG7A++mDalwyzNzEvRUslZmkyU16eToqmvnrWvNxj1JWQvgk9JEVhyXX5+h7APXylocHw90jIifK2sBvXZE3Js4tMIpa++5H7A/8BxwI7A5cCCwZbrICvNtRISkynfjkqkDSmSqpH3I/t93ytctljCeJPLXwnvAe2TJ0MsAt0t6OCJ+nTa6hUfSfhFxg+pucUpE/L3omBK6DHgLGAM8LmlVsjHW0nAyLEREKRIZZmMAWVL4G/nN/e+TJUCWzQRK9v6vQ2nPkyT9KCIeVT3tLcvQPaJKk8rFcKJP43Is8Fvgrnx27urA7MppNUWHkbVjaA+8TVbJ48jZ/kQTImln4BxgZeADYFWyRKeuKeNKYHPgZ3n56W8oUQszSV8yo/z67/IB+qmUK9P+b6kDaAwi4oez2XZ9fduakoh4OL9R04fsPXBMRHyUOKwUyvqZ6M+CmR1F1qbplnz5YUp0jlTl+8BGeQLkf/LWdleSVbAog1IPUNU6TzqWGRU8mgNfASekiaxYEXGSpLOBzyNiuqRJwC6p40rgSOByYB1JbwNvkt3cLZMbgWckDcqXdwL+mQ9YvpAurEJJ0mYR8VS+sCl1V8FrimbXxrU01WwA8goup5AlAgMMA06PiLLczLiGLAF803z5bbLWVaVK9JF0F7A22fnBThFRSf68RdKIdJEV6lZJlwHtJB0KHExW0ahsDiIbY/5TRLwpaTVmnDeWgqRjgAPIqqBeCZyYV3NpBrwKNNlEH6By47bsyQ0A/4qIf1QWJP2P7HOhVCR1A9alamJ1RFyXLqLi5fedKudJQ0uWDHwbWQvD0QB5RZvSVLWpSnp8g2yy3H1UTRYpWfJjmc+Tfgg8yowE6GpBObpHVBxDE8rFUAlbMjZ6klpHxNep47DiSRpDNhj3SESsJ2krYL+IGDCHH21S8tkFs6jM2rRyyAdi3o2IKfnyEmQzlN9KGlhB8lmYZzLrhWiZSsuSZ5lvTnbC+WRE3JU4pML5M9FsZpJWADbMF4dHxAcp4ylS3qarP7B6RJwuqSOwYkSUscpVaUnaAxgcEV/mVf56A2eUqXpHtTyppVlEfJk6lhQkbciMm/tPRURZbmYDIGl9soH7pcmSAD8FDi7r+6GsJN0BPA8MzFftD/SMiDpnrDY1kkZExAaSnqtUuJI0JiJ6po6tKHnywu8i4ozUsaQmaWuydk0AD0XEwynjSSUfQ+oYES+njiUFSacC19Q1biCpS0SUpoJ+mUkaFRG9a60bGRHrp4qpaHlr2y3JxlfvB7YnG1/cPWVcRZJ0FtkYyo35qn2AERHx23RRFUfST8gSQPuQJf1cU6bvhvw9UK+IOK2oWBoDnyeVm6TmwF8ioslMFHSiTyOSl4q6CmgTER0l9QR+ERFHJA5toZN0ATO36plJRBxdYDjJVA3OjAHWi4jvyjY4U5G//n+QLz4REWNSxlM0ZX3EH63MQJTUDtgyIu5OGVeR8hl3m0bEt/lyS7KbFxvO/iebBklPks1KPZcs0/ogsptYf0gaWIEkXQx0Bm7KV+0FvB4RZaxiMktbx4j4X8JwClNVwQOgJVnJ9UklqXCGpPMi4tg62hoClKmdIVCT4PA3YCjZDd0fkM1OvT1lXEWRdAlZq64fRUQXScuQDUyU5bvRbWoASWMjooekzYEzgL8Cf4iIjROHVihJiwO7AZ2oqlYcEaeniimVsp4jVMsrulCiCi413LYJJI2OiF5zWtdUSfo38GOy6+XeeVvDmyJio8ShFao60anMJK0IbER27fBsRLyXOKTCSdqJ7JqhZUSsJqkXWZWvUlw75TexxkfEOqljSUnSQLLK0J/ly8sA50REk69oI2kdsg4BZwMnVm1qS3b9XJruAZLGAT2B5yKiZz5x6IaI2DpxaIWRNBboFRHf5cvNyY5HU68UPpP8emEf4GSyNlZXkL0WpiYNLIE8QbpNRHyROpailf08yZVQQdLTEdEndRwNxa27GpfzgG2BewAiYoykLWb7E01HqWYdzsZnktoATwA3SvoAmJQ4psLl5WUPZUa5uBskXR4RFyQMq2inVFcuiYjP8uzru9OFVLgWlSQfgIj4Nk/2KYslImKIJOUzsE6VNBIoTaIPWYWzLpFnJeeDNOPThlQ8lbytY1T1Es+rmexCNgunLCol5t3KLPN/wIaVKj6SlgMeAUqR6ANsnN/Aew4gIj4t2Xej29Rkpud/7wBcHhH3SSpjBYNBwOdk7Wq+mcNzm6Q6zhE6Ai9RgnMESftFxA1Vpegr64HSlaB32yaYLGnziHgSQNJmwOTEMRXpFGAwsIqkG4HNgJ8ljSiNIZJ2A+6sXEOWjaRDyMYMHiVLir9A0ukRcXXayAp3KtlNvKEAETE6b8tQCpG1dn1ZUscyJv9W6VFJ8oGaa6eyJAOuDewItGPmNi1fko25l8nkfDL1NEltyc6ZV0kdVALtgE/yx0snjCMJZW3P9yOr+vgcWXWjzYEDySo+NXmS/knW1nI68CzQVtL5EfHXtJEVx+dJQFYN93lgz3x5f7LryVJUQs09J+kesmvmmvvvEbFIti9zok8jExETKgNTuen1PbcpiYiBkM3SjojbqrflM7fLYmdgClmPwP3IsuxLVTovN4DsRtYkAEl/Af4DlCnRp1kd68r2mf2hpJ0j4h4ASbuQ9RYvi2/y7PpXJf2SbMC+TeKYivYa2Q2rSqnpVfJ1ZfNHssSWmdo6Jo5poZPUIiKmVa/LB+zvzhMfT0oTWbEiYmT+sFdEnF+9LU+MHVZ8VEk1q9Wq62Pq/s5sqqbms+8qCZDLkVX4KYWI2Cp1DI3E28r6ym8N/CWvbFOm90FFh4jYLnUQiZXyHCG3ZP73UnVsK9sN/jUiYi9J+wBExNeqNbBUAocB11UqO5G1cDswYTyFioiHJY0i+zwQWQWLMl07V/yCrLrVdEmTyY5FlKUSaO5EsgrhH0PNjc1/k93UKZOpEfF5rY/C0pwz55YBxksazsw3sUpR1SjXTNIyEfEpgKTvUZKx1YgYBAyStElE/Cd1PImNyCvlX0GWGP0V2X2GMjmT7Mb2Y2TfjVtQkjE1AEl3kSW/XQ/sFBHv5ptuyTsKlMW6EfGFpP7AA2SvgZFkFYLLwudJ2bXjblXLp0kanSqYRFqRjSdXTxYMZhSeWKSU4sRmETJB0qZASFqMLNmjbP1yf0uWRTendU1KrbYkNavzv/8g6XXg5IgYUmxkyYiZk9ymM+N4lMUISX8HLsqXjyQ78SqTw8gqW11I9v8/ATggbUiFOgZoDRxNdhNnK0oyWF3Vomgp4MV8YCqAjYHhKWNLZGpEfCypmaRmEfGYpPNSB1WA4UBvSdUzCpoBG5AlxZbNgcD5tdb9rI51Td1gSQ8yc0u/+xPGU7R/AHcBy0v6E7A7WZWjUpFU5/lARFxXdCyJ7AlsB/wtr/q4EjOX5C+Lf0vqHhHjUgeSUFnPEYiIy/K/Z5kYI+nYwgNK61tJSzAjCXQNSlTlKk+A3T9vx9EWoCxtCDRrS8vKjauOeSWPUrS0rKiuBFpiH5NV7Kj4Ml9XNuMl7Qs0V9bO8GiyG3ll8vvUATQC5wD/kXQb2bji7sCf0oZUuMMkvVjG9mUVEXFE/vBSSYOBthExNmVMRYuImyQNBSotv39TsnZF/4iIx+raEBEbFB1MQovl951/ClwYEVMllW2ChM+TXAmViDgodQwNSSWtZNooSVqW7GbNT8hOPh8im4XT5D9oJG0P9CUbtL6lalNbskzTUvUVr5YPWnUDboyIbqnjKUJefv1AshtZkJ18DIyIc5MFVTBJS5JdlP8kX/UwcEalylGZ5O3siIivUsdSlPx9/5eIOCF1LClI+uHstkdEqSqYSHqE7HPwTGBZsjLDG0bEprP7uUWdpFF5i6JrmJEMOw14C7giIj5MFlyB8tn5+5KVFH6iatNSwHcR8eMkgSWUJ39tni8+Ud3qsgwkrQP8mOx6YUhElG1iAJKqqzy2IjseoyJi90QhJSFpebLfH4CytWaQ9ALQGXiTLKmhUrmhR9LAClTWc4Q5kfS/iOiYOo6iSNqaLOlzXbJxpM2An0XE0JRxFUnS0xFRptauAOQz8+sTEVGWlpZATZvf/sBqEfFHSasAK0VEaSaLSLoO6E7W3jLI2h6Pzf+Upq2hpNbAycA2+aoHycbUSjVhRNKqwJoR8Uh+TJpHxJdz+rmmRFJXsolzAI9GxAsp4ymapOciYr05rWvKqr4bVo+I0yV1BFYs2XfDZsDoiJgkaT+gN3B+RPx3Dj/aJOTJLYeTVTKCrDL2pRExNV1UxZN0NPAbYAxZG/COwA0R8YOkgRXI50kgqSdwHTNa+H0KHFimBEhJawGXACtERDdJPYCdI+KMxKHNFyf6WKOQf7j0Ak4n65FY8SXwWKXEZplJ+kVl1mIZ5DPTqm/iPZcyHktD0g5AV2a+iXV6uogWvkq7orIOVtsMkjoDKwCjyTLrm5ENTqwK3FfV0qlJkjQR+DuzVnQLKMfFF9QMzq5GdhO3urTyl8DY2u3NyiRPkv84SnZBkyeDrkBVddayJXjUlpdiv7ksbZwk7Uw2Q3llssSOjsBLEdE1aWAFyz8fZ1GWAWuomRxQfY6wNNkEkSY/WWh2JE2IiFVSx1GkvOx8pW3T02Vr2yTpEqA9WTXo6hY1i2T5dZs/+evgO+BHEdElr1zxUERsOIcfbTKUtTiuV11V0KxpknQo8HPgexGxRl7Z6NKyTRQp+7WTpDHAlrXalw2LiO5pIyuOvxtA0ligJ9ADuAa4CtgzImY70bKpkHQlsBgwMF+1PzA9Ig5JF1XjULkPkTqOopT5PCmv9vm/quVSVUKtJmkYWVXsyyqJr5KeX1QLbbh1VyOQz0it9wZFRBxdYDhJRMQYYIykFSJiYPU2ScdQvrYUsyhZks/1EbE/MKqOdaWQz8yb5XOhTDPyJF1K1rpqK+BKshK7ZZhtMZxsZsVzku6hxIPVtdoatiS7KJsUEW3TRVWo84DfVlXy+g4YKKk78Gdgp1SBFaQ50IbytW6cSX7D+r/AJrVmZC4BLMHMJWebLEl9gLOAT8jaGV5PVr2imaQDImJwyviKIuko4BTgfWa0Ng2yAbsym0SWEFcWfyS7of9IRKwnaStgv8QxFUZS23wwqhSff7NTfY4g6T5KmPxYjzIeg1ZkszFbAOtKIiIeTxxTkVqRld2vvl4OoBTXTrVa3VZ8DoyLiA+KjiehjfOKoM8BRMSnklqmDqpITfkG1byQ9DCwR612RTdHxLZJAyvWkcBGwDMAEfFqXg2yNHztBLh9Gfi7AWBaRISkXYCLIuIqSQNSB7WwVSWxbBgRPas2PZonwZWKpD/Us6lJT6quVvLzpLvJ7jkh6Y6I2C1tOEm1jojhWcG3GotswpsTfRqHEakDaET2Bs6ute5nONGnbGaajZzPvlg/USypVLdsagXsxiL8ZTOfNo2IHpLGRsRpks4BHkgdVIGqB6uDGQMSpRisBoiIpSqP81K7u5Dd2CyLFSJiXO2VETFOUqcE8RTt3aZewWteVM/IBNYAOgCXkrUsKoMLgd+RVat4FNg+Ip7O21jdBJQi0Qc4Bljb1Tr0L2bcyG9G1q7m1nQRFW5qRHwsqZmkZhHxmKTzUgdVoH8COwIjmXGOVBHA6imCKpKTH2dJCJ9pE1kibGlI+guwFzCeLDEcsmNTikSffLzg47K2Pc4NADYBKq28tiT7jFxN0ukRcX2qwAo2NX89BICk5ZjxniiF/Hf+NbNWRi7NpLHcspUkH6i5sV+qJBfgm4j4tnITS1ILypcIW/prp4i4TtIIZiTC7lq29mX4uwHgS0m/JZscsoWkZmSTKZu6ymTa6ZLWiIjXASStTpb8VzaTqh63IrumLkUreEnnRcSxtcaSakTEzgnCKlr1uEmTHzOZg48krcGM74XdgXfThjT/nOjTCNRRwaZttro8PXMl7QPsSzYIcU/VprZkg5dWAvkJ5++AJSRVSsYJ+Ba4PFlgCdTRkucpSWWoZlNtcv7315JWJkt6WSlhPEVZXtLxwPPUffOqyaurbGg+O/3uvMTmSXX/ZJPTbjbbynADq9SVfOpQ9hmZLSLiIYD8htXTABHxUq0ZGE3dBLIZ+mX3t6rH04D/RsTEVMEk8JmkNmQ38W+U9AEzD9o1aRGxY/53mao41Vb65MfqhHDjp2Q3Mr9JHUgKETFd0map40isBdAlIt4HkLQCcB2wMdl3RVkSff4B3EV2Tf0nssoV/5c2pMLdCNxCdvPuMOBA4MOkEaXxXXWLirwyainGU6oMk1QZY90aOAL4V+KYilb6aydJHYGvgHuq15WpfRn+boAsIXxfYEBEvJe/Lv6aOKYiVAaLTgAek/RGvtwJOChJRAlFxDnVy5L+BjyYKJyiVc6F/zbbZzVtUc/jMjqS7H7zOpLeBt5kEa6S7USfRkTSBmQ9MpfKFvUZcHAdN/ybon+TZcwtS1ZSsuJLYGySiKxwEXFmPhvxyog4OHU8KeU9kyuakVU0WjpROKncK6kd2YXHKLITkCuSRlSM2bUrKstJ2HCgd60S9M2ADYApaUJKYoSkQyNipte9pEPIZug2dWWpVDO3yj4js3rG3eRa25r8ccgTQAHeAIbmLXpqbuhGxN+TBJZIRAyrPJa0LFkycJnsQvZ9eBzQn+wcsXQV0CQNAc6JiPur1l0eET9PGFZRnPxo1d4gm5VdykSf3OiStz1epZLkk/sgX/eJpKmpgipaRNwoaSTZdYSAn0ZEKWaqV/l+3pLlmPx8aZikZ1MHlcDJwJOShpG9Fn5AVh21TE4iq/Y1DvgFcH/tsYUS8LUT3MeM6+UlyNodv0ytavpNmb8bgOwe2/l5cvRaQGVyQFO3XNVYymVkY+6QVfNZjxmVEMuqNVm18DIYL+lYoDPZ9+JVtScZl0DPvLiCmLXQQkRE23ShFSsi3gB+ImlJoNmiXnTFiT6Ny9XAERHxBICkzckSf5p839iI+C/wX2CTfObRhvmmF0v4gVtqEfGdpA3n/Mwmr7oVwTSyrNIm3zu3WkT8MX94h6R7gVYRUYaZOG5XNMNOzBiQmAa8BZShlGbFscBdkvozI7FnA6Al0C9VUEWJCFf0m1nZZ2TO7oK0Vf0/1mRUKlf8L//TMv8DJUh0qnC7okxEVFfvGVjvE5u+1YDfSNowIk7L122QMqAClTr50TKSLiD7//6aLNFlCDPfyDw6VWwJVLc9rihT2+Oh+TXzbfnybvm6JYHPkkVVkFoTpT6g6ualpO+V7Lqiktj1rqQdgHfIWv+WSkQMltSbGa2/j42Ij1LGlMBREXE+VRPm8gSw8xPGVLS6rp1KJSK6Vy/n74sjEoWThKQ/klW3u7bWdVSZPA78QNIywEPAs2RVfvonjWrhq28ybQtmjLGUhqRxzLhWbA4sR3kmDA0kO0d6AtierP37MUkjKlhENJ/zs8qhKgGwsgxZBcCRETE6RUwLQlknDGsMJD0XEevVWjcqInqniqlokvYgK582lBkzLk6MiNtTxmXFkjQQuDAiyjjryHKSWpFdfG5OdhL6JHBJRDTpii51fReUjaSJwN+Z9UIsoHSzr5C0FdAtXxwfEY+mjMfSyHuoDwC2IXtvPEhWAc8n8yUiaY+IuG1O65oqSSOY0a7ocmq1K2rq35+SvqTuJI7SzcCC7FqZrKXhP4BVyEotP1aG62dJ08mqlohsdvbXlU1kyfGLpYrNiiPpwNltr90m3pouZaPTuwGVFmZPAXeU5TxR0pvM3Pa68ntXvh9XTxJYApJ2JLuJtQpwAdAWODUiSjFBQNI6eXW7Os8FImJU0TGlUtc9hbKON0lqHRFfz/mZ5SBpXO0EoKZM0kFk95g2Iats8wTweEQMShpYgSqfB5KOApaIiLMljYmInqljW5jKdm91TvI2lhXTgPfLUmSh+nMvr5A+3K+N8pL0T7JJYpXz4x3JOgt1Am6LiLMThTZfnOjTCFRdfBxANkh3E9lF6V7AlIg4vr6fbWokjQG2jogP8uXlgEea+kmHzUzSS2Rl9P7LjAHsiIgmX92qQtJiwOHAFvmqocBlEVGastuSbiW7ALshX7Uv0C4i9kgX1cJXwhmHs5D0LnAJdbcvo2rWvplZqdQzYF+awStJoyOiV/74xYjoUrWtlDcuyqz6/1zSz4BfActERFnKj5vNIp+pvUpElKIFuqRf5zerKtWNZlKyqkZmSNosIp6a07qmqtLCU1Jd7VgiIn5Ux/omRdI+ZONnm5MlNFQsBXwXEaVpkS1pE+AqoE1EdJTUE/hFRJSmok2tqgXNgN5kLf62TRRSMpJWBPYETiC7ZihNRRdJz5FNpj0XGBAR48uQ8OUxghkkNSebPLpO6lhSqD1uVqZxNJuVpMeBvhHxVb7chqzV5XZkVX3WTRnfvHLrrsbhnFrLp1Q9LlsmVrNKkk/uY7KTUCuX0l1s1OESYDHg4nx5/3zdIckiKl63Wl+qj0l6IVk0BSl7kk/O7cvMasln5/4RWJXsHL6UFTzKStL2QF+gvaR/VG1qSzYTqyzcrohZWpRUfFmmhPDcpZUHEXFtXor8yITxmCUhaShZe9sWZO1eP5D0VEkmjb2Y/z0iaRSJSdoV+AuwPNk5YinPEyUNqZ3EUNe6Ju4Cshv5c1rXJEXEz/O/t0odS0L/Bt4la29bfc/hS7LZ6mVyHtkY8z0AETFG0haz/YmmpzqZZRrZjcw7EsWShKQrydr0vE+W/LY7UJrqXrljgd8Cd+VJPqsDdSVENjVl+v6frYiYLullSR0j4n+p40mgp6Qv8scClsiXS3nObCxPVctrsrZuK0TEZEnf1PMzjZYTfRqBkl981DZY0oPM6Ke9F3B/wngsgYj4L4Ck5YFWicNJZcNalawezStelckoSX0i4mkASRtT8gHcEqmzko9ZyZ0H7AqMK0sbBpvJO2TfgTuT3cSt+BI4LklEafSsGoxZotZATZnOGUeRteT4lOx3bwe8J+l94NCIGDmbn20yIuIymOma4UPg1JQxmSWydER8IekQ4LqIOEVSKW7mVtoRuU0ZZwM7RcSLc3xmE5S3/V4SWDavalW5nmwLtE8WWIHyyiWbAsvVquDRFmieJqp08s/Am4BbI+L11PEUKR9T/a+kxyNiWPU2SX8BfpMmsjQiYkLW3bDG9FSxpOCK2AB8n+xz8DPgE+CjsrQrqsg/C4ZVLb8BNPmqh55MO4tlgPGShpN10QAgInZOF1IxIqJ050I2WzcCz0iqtHDcCfinpCWBRa7QgBN9GgFJ+0XEDbUuxGpExN+LjimViDgxn4m0eb7q8oi4K2VMVjxJO5PNOlkZ+ICsesGLQNeUcRVsuqQ1KgMSeaZ9qS5GgfWBf0uqZJl3BF7OZ2uXqpVbCXnGhdmsJgDPO8mnnCJiDDAm7yMtYK1808tlquLiwZkaDwO3R8SDAJK2AXYDriGrBrlxwtgKI2kn4O/MuGboSHbN0C1lXGYJtJC0Elk7ipNTB1MkSffMbnsZblzk3i9rkk/uF2TVClZm5ioNXwAXpggogZZAG7Kx/uoKHl+QVa8om53IJo/eKuk74BaypJ8yVTHYmlmTeravY11TNkHSpkBIWgw4hhmV4Jo0Sf9iNhVPS/T9SET0A5DUhazC02OSmpeh3a+k8yLi2PpeD2V6HRgAv08dgFljEBF/lDSYLEke4LCIqBQY6J8orPkm3ytIT9IvIuIySafUtb2smdeSlgU+9g2t8skr1/wIeCQi1pO0FbBfRAxIHFphJP2Y7GbNG2Q39FYFDoqIMpTVBEDSqrPbXqn8ZGZWBpI2JGvdNYyq8qJlSgg3kPRD4DrgLbLzg1WAAyPi8ZRxWbEkjYuI7rXWjY2IHpJGR0SvRKEVytcMZhlJe5AN3D8ZEUfkk0T+GhG7JQ5toZP0IVky9E3AM9SqDFq7mkVTJel8YEXgbmY+T7wzVUwpSDoqIi5IHUdKklatqpK9YkS8lzqm1CStSfYZ2b8MSeOSDgeOANYAXqvatBTwVETslySwBPJ7C+cDPwGaAQ8Cx0TEx0kDK0B+3QhZVeAVgRvy5X3IkkNLUxU2b4P+A2ALskqoTwNPRMTVKeMqgqT1I2Jk1ethJmU5T7JZ+d6rlZ2k5sAKVBXEWVQTwp3o08hJahkR36aOY2GT1Ac4i6x84h+B68l6CTcDDoiIwQnDs4JJGhERG+SD9+tFxHeSxtRqZdVkSVqOLLFnIlm/SMhm7C9y/SEbiqSfR8TlqeMwM0tF0kPAV8A44LvK+rImhJeVpJHAvhHxcr68FnBTRKyfNjIrUv55MAS4OV+1F9ms7e2AZyOid6rYilT2awYzqxmg3ZrsxmUP4D6y78XxSQMrmKRr6lgdEXFw4cEkJOmAutZHxHVFx9IYSBpVlnOCuuSTx/bK/0wHbomIc9JGtfBJWpqsPcuZwElVm750G5vyqZwvz2ldUybpQuAJsuSed1LHk0LejmZyRHyXLzcHFo+Ir9NGZkXwvVezmUk6CjgFeJ/sHFEswh1E3LqrEZE0FPhZRLyVL28IXAmUYaDyQuB3wNLAo8D2EfG0pHXIZmb5y6ZcPpPUBngcuFHSB1T1DW3KJB0C/Bl4HVgN+HlEzLYceUkcBjjRx8zKbOWIcDsaW6yS5AMQEa/kZeitXPYlG5S4O19+Ml/XnKx1T1mU9prBrJqk1YCjgE7MPCOxybdjiIjpZONFgyUtTpbwM1TSaRFRlpZNRMRBqWNoJDasetyKrCX0KLJqiGWkOT+laZL0DLAYcBuwR0S8kTikwkTE58DnZJ+HSFqe7P3QRlKbRXW2+vzIK9ydD/Qha1v0H+C4Mr0egCUlrV75nfNzhiUTx1SoiPhl5bGkHSPi3pTxJDKErLLVV/nyEsBDzGhbY02b772azewYYO2mUuHPFX0aEUnbkp18/gNoT9Y395CIGDXbH2wCqkvMS3oxIrpUbXsuItZLFpwVRlJnsnJpo4HJZFnF/cmq29wXESPTRVcMSc8DW0XEh/kF6Y0RsUnquFLz54CZlZ2ks8na0zyUOhZLR9LVZBWdKqXX+wPNyzZj32aQtFJEvJs6jhTymalTyG5k9icbuLyxqQzWmM2tvKrVVcxa9a8U7RjyBJ8dyG5qdwLuAa6OiLdTxlUkSa2AAUBXshv6AJT9/EBSO+DmiNgudSwpSDoiIi5OHUcKktauTo4vI0k7AX8HVgY+IBtbfTEiuiYNrECSngYuIruRDbA3cFREbJwuqmJJ2o5s4uQbZOfMq5JNKi3luEJZK53V1d65TC2fy873Xs1mJukxYOuImJY6lobgij6NSEQ8KOkw4GHgI7Ly42Xppfxd1ePJtbY5G608zgN+GxGVmbjfAQMldSercrNTqsAK9G1EfAgQEW/kg5alk5cQHR8R6+SryvB/b2Y2O4cDJ0j6BpjKjLKibdOGZQU7HDgSODpffgIo5Q0cq3EfULrBaoCqawaAgckCMUtvSkT8I3UQKUi6DugG3A+cFhHPJw4pleuBl4BtgdPJkh9fTBpR4zAJWD11EClIWreS5COpT0Q8nTqmgn0m6SqyqqjbS1oX2CQirkodWIHOIKtk80hErCdpK2C/xDEVrXVEXF+1fIOkE5NFk0BEDJa0JlAZX30pIr5JGVNiZa10NklS70pBAUkbMOs9OGu6fO/VbGZvkFWBvQ+o+U6MiL+nC2n+uaJPIyLp92Sl1n9O1lv8OOBXEXFf0sAKIGk62QW4yEoHVvqDCmgVEW5JUAKSno2IDevZNi4iuhcdU9HylgM3V63au3o5Io6e5YeaKEmDyGbalKassJmZmdm8KPMMPElfMuvA5OfACLLr6DK1ZbASk7QvsCZZC4bqgcoyVIf+jhkt+6o/D0qREC2pRURMq3wXSBobET3ytp5PRESf1DEWSdK/mPE6aAasC9waESeliyoNSfcCywCDyKrFr5U4pEJJegC4Bjg5InpKagE8V4ZxxQpJIyJig7zq23oR8Z2kMRHRM3VsRZH0F+BTsnHVAPYie1/8FSAiPkkXXTHy74PDgS3yVUOByyJiarKgCiZp8Upyk6SNImJ49boykLQh2fvgnXzVSsBeZeieYL73alabpFPqWh8RpxUdS0NwRZ/G5fvARhExGfiPpMHAlWSzNJu0iGieOgZrFNrNZtsSRQWRWO2ZJWU+4V4GGC9pODMGb4mIndOFZGaWhqTNgNERMUnSfmQVPM5zMmQ5SBrHbGZaRUSPAsOxxuWK1AEkdB4wEfgn2SDl3sAawCjgamDLVIGZFaw7sD/wI2bM2I18uUmLiGapY0hsONk5YeWG7WeSugHvAcsni6pgVW3g/1a1ehrZd0Mp2ltK6gR8EhFfAETEjpKOIjsm+6aMLZFlI+JWSb8FyBPipqcOqmCfSWoDPA7cmE8snDSHn2lq9sz//kWt9XuTfU+WoeLXJcBizKgCu3++7pBkERXvP+QVUCNieO11TVme4DMhIp6VtA7Ze2FXYDDwZtLgrDC+92o2s0U1oac+rujTyEhaAuhY9j7CVk6SbgIejYgraq0/hKxn4l5pIktL0oolauNXQ9IP61ofEcOKjsXMLDVJY4GeZFUfryVLBt8zIur8rLSmJS+3vgIwodamVYD3IuK14qOy1PK2HC/kj0vXlqOuWemSRkdEr7LNWLdyk/QasG5EfJs6FiuWpFER0TsfM7mDLOnrWqAN8PuIuCxlfEXJq9f8NiLG1VrfHfhzRDT5VuCSRgI/iojP8+WjyaqXHAJcFBFNPvGvmqShwG7Aw/l7pA/wlzJdO0lakqxFSzOydn5LAzdGxMdJA7NC1XO+XIrzZEkrAu2BG8gSHittu9oCl0bEOvX9bFMhaRTwk4j4RNIWZFV9jgJ6AV0iYveU8ZmZpSBpOeDXQFegVWX9onq+7Io+jYiknchmWrQEVpPUCzjd1SusRI4F7pLUnxmVbDYge0/0SxVUI3A/JZhlUFtEDJO0KrBmRDwiqTXgDHQzK6tpERGSdgEujIirJA1IHZQV5lyyG1j/rV4pqW2+rcnfwLI6nS2ppi0HUKq2HMDXkvYEbs+Xdwem5I89o8nK5Hmy6rgfJI7Dire8pOPzxwflf1+U/71kgnhSWaF2kg9ARIzLK92UQcuqJJ8/A+uRTZj7WtLSaUNL4njgHmANSU8By5GdJ5RGRFSq93wn6T7g4yjxjG9Jl0fEz1PHkcB0SWtExOsAklYHylLdalvgZ0AH4O9V678EfpcioASaV7Wo2wu4PCLuAO6QNDpdWGZmSd0I3ALsCBwGHAh8mDSiBeBEn8blVGAjsl6pRMTo/OTLrBQi4n1gU0lbAd3y1fdFxKMJw2oMNOenND2SDgV+DnyPrA1De+BS4Mcp4zIzS+TLvPT8/sAPJDUjK8Ft5eAbWOa2HLPqD5xP1ooggKeB/fIqub9MGZhZwdoBL0l6FvimstKTxkqhOVn1nrrGDMp0Q7/dbLaVpQ38a5KuIbuhvR6wdp7k0yVxXElExKi8SvTaZO+PlyNi6hx+rEnIqxedBXwC/BG4HlgWaCbpgIgYnDK+hDZIHUAiJwCPSXqD7L2wKjMSQ5u0iBgIDJS0W57cUkbNJbWIiGlk4+nVyW6+N2xmZfX9fALtMXn3kGH5tfQiyR/mjcvUiPhcmun6/Lv6nmzWVEXEY8BjqeNoRK6Y81OapCPJkh+fAYiIVyUtnzYkM7Nk9iK7kX9wRLwnqSPw18QxWXHazWZbWW5gWdaWpaaUcFVbjl5kFRxKNYAdEW9QfzWrJ4uMxSyxU1IHYMm8GxGnpw6iERgh6dB62sCPrOdnmpq9gT2Ab4E3gKGSPgTWIZulXAqSdq1n01qSiIg7Cw0ojQvJqpUsDTwKbB8RT0taB7gJKGuiT+mq3klqTtb+e02ypDfIkt6+qf+nmqQhkv4ObJEvDyProvF5wpiKchPZDeyPyFr5PQEgqTNQht/fzKwuleTvdyXtALxDVmxgkaQSV2xsdCRdBQwBTiLrI3w0sFhEHJY0MDNLRtK6EfFC/rhPRDydOqaiSHomIjaW9FxErCepBTAqInqkjs3MLAVJKwAb5ovDI6J0g5VlJekm4NF6bmBtHRF7pYnMiiRpXER0zx9X2nLsls/YHxkR66eNsFiSWgEDmLWv+sHJgjIzK1DlWjl1HKnl58h3kSW5zNIGPiLeSxVbKvl3ZHfg1Yj4LHE4hcmrGgEsD2xKlugCsBXw74jYMUlgBZI0OiJ65Y9fjIguVdtK/5khqVVETJnzM5sGScMjYqPUcaQk6Q6yNqcD81X7Az0jor7EwCYlr/K1EvBQpaWfpLWANhExKmlwZmYJSNqRLPFxFeACoC1wWkTckzSw+eREn0ZEUmvgZGAbslKKDwJ/LNPJp5nNTNK9wDLAIOCQiFgrcUiFkXQ28BlwAHAUcATwQkScnDIuM7MUJO1JVsFnKNl54g+AEyPi9pRxWTF8A8sAJN1Fdm5U3Zbj47wtxw0lTPS5DXiJrNrZ6WStvF6MiGOSBmZWsPwGzgVAF7LvhebApIhomzQwW+gkfS8iPkkdR2NRqw38+LK3gZd0akScmjqOFCQ9BBwYEe/myysB10bEtmkjW/gkjYqI3rUf17VcFnk7jn8CNwO3R8RmiUMqjKRzyVp+3wJMqqwvU4JHdfLb7NaZmZktipzoY2bWiEjqBHwSEV9UrTsK+Buwb5l6CktqRjZLuyb5sXYlAzOzspA0hqxyywf58nLAIxHRM21kViTfwCo3SYszc1uOa4CathwR8XDC8ApXVfVxbET0kLQY8ERE9Ekdm1mRJI0ga9tzG1kS6AHAWhHx26SBmVlSZU3qgDor2TQjO3fuMpsfaxIkTSdL6BBZi9+vK5uAVhGxWKrYUpG0LPBLspZmJ0TEPxKHVBhJj9WxOiLiR3Wsb5Ik/YdsktST+fJmwN8iYpO0kZmZWZEkXQDUmxQTEUcXGE6DaZE6AANJ/2L2L66dCwzHzNK6A6i52JJ0NLAX0Au4KN9eFkdFxPlATXKPpGPydWZmZdOsVquuj4FmqYKxNCLiMaCuwVorgYj4BrihsixpQ0rYlqNKpa/6Z5K6Ae+RteowK52IeE1S84iYDlwj6TnAiT5m5abUASQ0RNKDwE358l7AIwnjKUxENE8dQ2p5C7dTI+K/+aqlyZLlzwZ6JAssgYjYKnUMjcDhwEBJS5N9Ln4CHJg2JDMzS2BE1ePTgFNSBdKQnOjTOPwtdQBm1mi0jIjPAST9mawtw9YR8XV+QVImBwK1k3p+Vsc6M7MyGFzHYPX9CeMxs8QiYoqkHcralgO4XNIywO+Be4A2wB/ShmSWxNeSWgKj8/bH7+JkYDODUrX0rBYRv5TUD9giX3V5RNyVMiYrVO9Kko+k9cnadh0cEU9JGp42tGJJ+j7ZjczNySaaPwmcHhEfJw2sQBExGugpqW2+/MXsf8LMzJqiiBhYeSzp2OrlRZlbdzUCkjpGxP9Sx2Fm6Um6C/gM6ECW5LN2RHwsqQtwQ0Q0+YEaSfsA+5JdhD5RtaktMD0ifpwkMDOzBCR1BlbIByV3JftshOy74saIeD1ZcGaWXJnbcphZRtKqwAfAYsBxZJULLo6I15IGZmaFk7QWcAnZ9UM3ST2AnSPijMShJSNpx4i4N3UcVhxJo4GjgY7An4C+ETE+T4odU4YWbhWSHgYeZ0ZV0P7AlhHxk3RRFSufOHsKMxL/hpElO32eLiozM0upKY2lOdGnEah+QUm6IyJ2Sx2TmaUhaXGycrLfAm8A1wAfAusAB0bEwwnDK0Q+UL0acCZwUtWmL4GxETEtSWBmZglIuhf4bUSMq7W+O/DniNgpTWRm1hhIei4i1ksdRwqS2gEHAJ2oqla8qPZVNzMzW1CShgEnApdVzg8kPR8R3dJGlk5TupFjc0fSxmQJPt8CrwNLkCW77AU8HxG/SRheoep6/0saFxHdU8VUNEl3AM8DlcoN+wM9I2LXdFGZmVlKTen80K27GofqnsmrJ4vCzJKLiG+YMcsCSRsC3YFXI+KzVHEVKS+v+19JPwEmR8R3+ay0dYBxs/9pM7MmZ4XaST4AETFOUqcE8ZhZ49IkBibm0/3A02Tnh98ljsWscJLGkbXhqFNE9CgwHDNrHFpHxHCpeqiZsk+W0pyfYk1JRDwD1FSskbQzsC1wF3BVqrgSeUjS3sCt+fLuwIMJ40lhjVoT60/Lqz6ZmVmJSPqSGdfPrSVVWjkKiIhomyayBeNEn8Yh6nlsZiUXEVMk7RARp6aOJYHHgR9IWgZ4CHiWbPZN/6RRmZkVq91sti1RVBBm1nhIWh04H9gE+E7Sf4DjIuKNtJEVrlVEHJ86CLOEdgVWACbUWr8K8F7x4ZhZI/CRpDXIx5cl7Q68mzak5H6ROgBLKyLukTQxIkaljiWBQ4FjgevJbmQ2AyZJ+gWL8E3NeTRZ0uYR8SSApM2AyYljMjOzgkXEUqljWBiapQ7AAOgp6Ys8m6xH/vgLSV9WZZSZWXntnDqARBQRX5MNYF8cEXsAXRPHZGZWtBGSDq29UtIhwMgE8ZhZev8km5W7IrAycBtwU9KI0rhe0qGSVpL0vcqf1EGZFehc4POI+G/1H+DzfJuZlc+RwGXAOpLeJrvBf1jSiBKQdGTe4pO8wtEyko5IHJaldWXqAFKIiKUiollELBYRLfLHS+V/ypDkA9ln4EWS3pL0FnAhTgA0M7MmQhEuIGNm1phJeq7SW71MJD0HHEE2SD0gIsaXrY+0mZmkFchKjH/LjMSeDYCWQL+I8Ix9s5KRNLZ2Sx5JYyKiZ6qYUpB0JPAn4DNmVMaNiHA7bCsFSc9GxIb1bPN1k1kJSVotIt6UtCTQLCK+rKxLHVuRJI2OiF611pVybM0y/v8HSaeWqWK8pI4R8b+q5bYAEeGJ9WZm1mS4oo+ZWeO3fuoAEjkW+C1wV57kszrwWNqQzMyKFRHvR8SmwGnAW/mf0yJiEyf5mJXWA5JOktRJ0qqSfg3cX8KKNr8COkdEp4hYLf/jJB8rk3az2eb2nmbldAdAREyKiC/zdbcnjCeV5pJUWZDUnGyihJXXaakDaATKVjH+7soDSXdExBdO8jEzs6amReoAzMxsVpLOBs4g6xk8WFIP4LiIuCFtZMWJiGHAsKrlN4Cj00VkZpZORDyGkx3NLLNn/nftkvN7k1W2KUuyy2vA16mDMEtohKRDI+KK6pVu72lWPpLWIWt1vrSkXas2tQVapYkqqcHALZIuy5d/ka+zkpHUHlgV+ETSFgAR8XjaqIqRJ7t1iIgJlVUp40mg+vcty/WRmZmVjFt3mZk1QpUyw5L6ATsCxwOPl6Elg6TzIuJYSf9iRhuGGhFRthkoZmZmZlaLpLvIbmo+BnxTWR8RTgy3UnB7TzOrkLQL8FOyih33VG36Erg5Iv6dIq5UJDUjS+75cb7qYeDKiJieLiormqS/AHsBLwCV//so07hidStPSc0i4rvUMRVF0qiI6F37sZmZWVPiRB8zs0ZI0viI6CrpSuD2iBgsaUxJEn3Wj4iRkn5Y1/a80o+ZmZlZKUk6oK71EXFd0bGkJOnAutZHxMCiYzFLSdJWQLd8cXxEPJoyHjNLR9IWtauVSNosIp5KFZNZKpJeBnpExDdzfHITJWkgcGFEPJs6lqJJmg5MIqvsswQzKoGKLOGrbarYzMzMGooTfczMGiFJZ5HNxpoMbAS0A+6NiI0ThmVmZmZmiUm6oGqxFdls9VERsXuikJKT1DsiRqWOw8zMLKW6qlaUqZKFpFsjYk9J46i7QnSPBGFZIpIeAPaIiK9Sx5KKpJeAzsB/mZH0En4vmJmZNQ1O9DEza4QkLQ4sCXweEdMlLQm0iYj3E4e20NU3IFPhi1EzMzOzGSS1I2vLsV3qWFIp001MMzOz2iRtAmwKHAucW7WpLVkrvyZfHRpA0koR8a6kVevaHhH/LTomS0fSHUBPYAglbfPq94KZmVnT1iJ1AGZmVqf/VN+siIhJkp4AynADY8fUAZiZmZktQiYBq6cOIjGlDsDMzCyhlkAbsrH+parWfwGUpuJfRLyb/+0kBgO4J/9TWhHxX0mbA2tGxDWSliP7rDAzM7MmwIk+ZmaNiKQVgfbAEpLWY8ZNi7ZA62SBFaiuARlJywIfh8vQmZmZWclJ+hczqh82A9YFbk0XUaNwWuoAzMzMUomIYcAwSdfmN/ZbR8TXqeNKRdKuwF+A5cnG1SrtitomDcwKFREDJS0BdIyIl1PHk4KkU4ANgLWBa4DFgBuAzVLGZWZmZg3DiT5mZo3LtsDPgA7A36vWfwH8LkVARZPUBzgL+AT4I3A9sCzQTNIBETE4ZXxmZmZmKUjqDKwA/K1q9TSym1fvJgmqEZDUF3gwf7xrRNyZOCQzM7NUVpb0AFnFjo6SegK/iIgjEsdVtLOBnSLixdSBWDqSdiI7b24JrCapF3B6ROycNLBi9QPWA0YBRMQ7kpaa/Y+YmZnZosKJPmZmjUhEDAQGStotIu5IHU8iF5IlNS0NPApsHxFPS1oHuAlwoo+ZmZmV0XnAbyNiXPVKSd3zbTsliKkx6Av8QdIooA/gRB8zMyur88gmkN0DEBFjJG2RNKI03neSjwGnAhsBQwEiYrSksrW7/TYiQlIASFoydUBmZmbWcJqlDsDMzOr0lKSr8plYSFpX0oDUQRWkRUQ8FBG3Ae9FxNMAEfFS4rjMzMzMUlqhdpIPQL6uU/HhpCFpY0nLVZYj4pfA/cBeZDP4zczMSisiJtRaNT1JIGmNkHSLpH0k7Vr5kzooK9zUiPi81rrvkkSSzq2SLgPaSToUeAS4InFMZmZm1kCc6GNm1jhdQ9aCYOV8+RXg2GTRFKv6ontyrW1RZCBmZmZmjUi72WxboqggGoHLydraAiDp70AvYB3gl4liMjMzawwmSNoUCEmLSToBKGNlm7bA18A2ZBUPdwJ2TBqRpTBe0r5Ac0lrSroA+HfqoIoUEX8DbgfuANYC/hARF6SNyszMzBqKW3eZmTVOy0bErZJ+CxAR0ySVZRZWT0lfAAKWyB+TL7dKF5aZmZlZUiMkHRoRM83ClXQIMDJRTCm0iIhvJLUAriVLDN89Ir6T1DptaGZmZkkdBpwPtAfeBh4CjkwaUQIRcVDqGKxROAo4GfgG+CfZhMozkkaUxjiySQGRPzYzM7Mmwok+ZmaN0yRJ3yevYCOpD1C73GyTFBHNU8dgZmZm1ggdC9wlqT8zEns2AFoC/VIFlcCTkoYAKwJtgC3yJJ8fMms1SDMzs9KIiI+A/qnjSE1SB+ACYLN81RPAMRExMV1UVhRJrciS3jqTJbZsEhHT0kaVRj4h4A/Ao2QTKC+QdHpEXJ02MjMzM2sIinAXFDOzxkZSb7JBiW7A88ByZDOVxyYNzMzMzMySkrQV2TkiwPiIeDRlPClI2hz4FnifrB3Bsvmm3SJiVLLAzMzMEpJ0NlnFksnAYKAHcFxE3JA0sIJJepisgsv1+ar9gP4RsXW6qKwokm4BppIleG0PvBURxyYNKhFJLwObRsTH+fL3gX9HxNppIzMzM7OG4EQfM7NGKm9HsDbZjIuXI2Jq4pDMzMzMzBodSctFxIep4zAzM0tJ0uiI6CWpH7AjcDzweET0TBxaoSrHYU7rrGmSNC4iuuePWwDDI6J34rCSkPRvYMuI+DZfbgkMjYhN00ZmZmZmDcGtu8zMGq+NgE5kn9W9JRER16UNyczMzMys0fkT8PPUQZiZmSVWGevfAbgtIj6XlDKeVD6WtB9wU768D/BxwnisWDUTJSNiWknfAxWvAc9IGgQEsAswVtLxABHx95TBmZmZ2YJxoo+ZWSMk6XpgDWA0MD1fHYATfczMzMzMZrZB6gDMzMwagXslvUTWuutwScsBUxLHlMLBwAXAuWRjaf8GDkoakRWpp6Qv8scClsiXBUREtE0XWuFez/9UDMr/XipBLGZmZtbA3LrLzKwRkvQisG74Q9rMzMzMbLYkDY6I7VLHYWZmlpqk7wGfR8R0Sa2BthHxXuq4iiKpOXBdRPRPHYuZmZmZ2cLULHUAZmZWp+eBFVMHYWZmZmbWGElarfK4kuQjacN0EZmZmaUlaQ9gap7k83/ADcDKicMqVERMB1aV1DJ1LGapSVpO0l8l3S/p0cqf1HGZmZlZw3DrLjOzRkTSv8jKCi8FvCBpOPBNZXtE7JwqNjMzMzOzRuQOSTtFxNsAkn4IXAh0TxuWmZlZMr+PiNskbQ78BPgrcAmwcdqwCvcG8JSke4BJlZUR8fd0IZklcSNwC7AjcBhwIPBh0ojMzMyswTjRx8yscflb6gDMzMzMzBYBvwDulrQT0Bs4E+ibNiQzM7Okpud/7wBcHhH3STojZUCJvJ7/aUY2kQ6ySXVmZfP9iLhK0jERMQwYJunZ1EGZmZlZw3Cij5lZI5JfdFVaEbwbEVPy5SWAFVLGZmZmZmbWWETEs5KOBh4CpgA/iQjPUDYzszJ7W9JlwNbAXyQtTpbsUjYvRMRt1SvytmZmZTM1//tdSTsA7wDfSxiPmZmZNSBFOJndzKyxkTQC2DQivs2XWwJPRcSGaSMzMzMzM0unqtVtxbrAu8Cn4Fa3ZmZWXpJaA9sB4yLiVUkrAd0j4qHEoRVK0qiI6D2ndWZNnaQdgSeAVYALgLbAaRFxT9LAzMzMrEG4oo+ZWePUopLkAxAR3+bJPmZmZmZmZeZWt2ZmZnWIiK+BOyUtL6ljvvqllDEVSdL2ZG0820v6R9WmtsC0NFGZFU9SK+AwoDPQHrgqIrZKG5WZmZk1NCf6mJk1Th9K2rkyw0LSLsBHiWMyMzMzM0uq0uoWQNIKQKXi5fCI+CBNVGZmZulJ2hk4B1gZ+ADoSJbo0zVlXAV6BxgB7AyMrFr/JXBckojM0hhI1rbrCWB7sgqYxySNyMzMzBqcW3eZmTVCktYAbiQbnBEwATggIl5LGpiZmZmZWSMgaU/gr8BQsvPlHwAnRsTtKeMyMzNLRdIY4EfAIxGxnqStgP0iYkDi0AolqS0wKSKm58vNgcXzikdmTZ6kcRHRPX/cgiwh3q3rzMzMmhhX9DEza4Qi4nWgj6Q2+fJXiUMyMzMzM2tMTgY2rFTxkbQc8AjgRB8zMyurqRHxsaRmkppFxGOSzksdVAIPAT8BKmNpS+TrNk0WkVmxplYeRMQ0SSljMTMzs4XEiT5mZo2UpB3Iyiu3qlyQRcTpSYMyMzMzM2scmtVq1fUx0CxVMGZmZo3AZ/mEsceBGyV9AExKHFMKraonzEXEV5JapwzIrGA9JX2RPxawRL4sICKibbrQzMzMrKE40cfMrBGSdCnQGtgKuBLYHRieNCgzMzMzs8ZjsKQHgZvy5b2A+xPGY2ZmloSkzsAKwC7AZOA4oD+wKnBUwtBSmSSpd0SMApC0PtlxMSuFiGieOgYzMzNb+BQRqWMwM7NaJI2NiB5Vf7cBHoiIH6SOzczMzMysMZC0K7B5vvhERNyVMh4zM7MUJN0L/DYixtVa3x34c0TslCayNCRtCNwMvENWwWRFYK+IGJk0MDMzMzOzBuSKPmZmjVNlptHXklYGPgFWShiPmZmZmVlj8xQwFQhc/dLMzMprhdpJPgARMU5SpwTxJBURz0paB1g7X/VyRExNGZOZmZmZWUNz/3ozs8bpXkntgLOBkcCbzGhLYGZmZmZWapL2JEvu2R3YE3hG0u5pozIzM0ui3Wy2LVFUEI2FpNbAb4BjIuJ5oJOkHROHZWZmZmbWoFzRx8ysEcnLC0+IiD/my22AccBLwLkpYzMzMzMza0ROBjaMiA8AJC0HPALcnjQqMzOz4o2QdGhEXFG9UtIhZJPHyuYast97k3z5beA24N5kEZmZmZmZNTBFROoYzMwsJ2kU8JOI+ETSFmQ9xY8CegFdIsKzlM3MzMys9CSNi4juVcvNgDHV68zMzMpA0grAXcC3zEjs2QBoCfSLiPdSxZaCpBERsYGk5yJivXzdmIjomTo2MzMzM7OG4oo+ZmaNS/OI+CR/vBdweUTcAdwhaXS6sMzMzMzMGpXBkh5kRnvbvYD7E8ZjZmaWRES8D2wqaSugW776voh4NGFYKX0raQkgACStAXyTNiQzMzMzs4blij5mZo2IpOeBXhExTdJLwM8j4vHKtojoNvt/wczMzMysHCTtCmyeLz4REXeljMfMzMzSk7Q18H/AusBDwGbAzyJiaMq4zMzMzMwakhN9zMwaEUknA32Bj4COQO+ICEmdgYERsVnSAM3MzMzMGhlJywIfhwc4zMzMDJD0faAPIODpiPgocUhmZmZmZg3KiT5mZo2MpD7ASsBDETEpX7cW0CYiRiUNzszMzMwsofxc+SzgE+CPwPXAskAz4ICIGJwwPDMzM0tEUu/ZbfeYmpmZmZk1JU70MTMzMzMzM7NFgqQRwO+ApYHLge0j4mlJ6wA3RcR6SQM0MzOzJCQ9NpvNERE/KiwYMzMzM7OFzIk+ZmZmZmZmZrZIkDQ6Inrlj1+MiC5V255zoo+ZmZmZmZmZmTV1zVIHYGZmZmZmZmY2l76rejy51jbPZDIzMyspSb+uerxHrW1/Lj4iMzMzM7OFxxV9zMzMzMzMzGyRIGk6MAkQsATwdWUT0CoiFksVm5mZmaUjaVRE9K79uK5lMzMzM7NFXYvUAZiZmZmZmZmZzY2IaJ46BjMzM2uUVM/jupbNzMzMzBZpbt1lZmZmZmZmZmZmZmaLsqjncV3LZmZmZmaLNLfuMjMzMzMzMzMzMzOzRZbbe5qZmZlZmTjRx8zMzMzMzMzMzMzMzMzMzMxsEeDWXWZmZmZmZmZmZmZmZmZmZmZmiwAn+piZmZmZmZmZmZmZmZmZmZmZLQKc6GNmZmZmZmZm1kRJWlHSzZJelzRS0v2S1qrnue0kHVFQXIdJOqCIfZmZmZmZmZmZNSWKiNQxmJmZmZmZmZlZA5Mk4N/AwIi4NF/XE2gbEU/U8fxOwL0R0W0hx9UiIqYtzH2YmZmZmZmZmTVVruhjZmZmZmZmZtY0bQVMrST5AETEGOA5SUMkjZI0TtIu+eazgDUkjZb0VwBJJ0p6VtJYSadV/h1Jv5f0sqQnJd0k6YR8fS9JT+fPv0vSMvn6oZLOkzQCOEbSqVU/s4akwXnFoSckrZOv30PS85LGSHq8gONlZmZmZmZmZtbotUgdgJmZmZmZmZmZLRTdgJF1rJ8C9IuILyQtCzwt6R7gJKBbRPQCkLQNsCawESDgHklbAJOB3YCewGLAqKr9XAccFRHDJJ0OnAIcm29rGREb5P/2qVXxXA4cFhGvStoYuBj4EfAHYNuIeFtSuwU8FmZmZmZmZmZmTYITfczMzMzMzMzMykXAn/Okne+A9sAKdTxvm/zPc/lyG7LEn6WAQRExBZgi6V8AkpYG2kXEsPz5A4Hbqv69W2YJRGoDbArclnUaA2Dx/O+ngGsl3QrcOR+/p5mZmZmZmZlZk+NEHzMzMzMzMzOzpmk8sHsd6/sDywHrR8RUSW8Brep4noAzI+KymVZKx85nPJPqWNcM+KxSRahaRByWV/jZARgpaf2I+Hg+921mZmZmZmZm1iQ0Sx2AmZmZmZmZmZktFI8Ci0v6eWWFpB7AqsAHeZLPVvkywJdk1XoqHgQOzqvuIKm9pOXJKu3sJKlVvm1HgIj4HPhU0g/yn98fGMZsRMQXwJuS9sj3IUk988drRMQzEfEH4ENglfk+EmZmZmZmZmZmTYQr+piZmZmZmZmZNUEREZL6AedJ+g0wBXgLOBX4h6RxwAjgpfz5H0t6StLzwAMRcaKkLsB/8rZaXwH7RcSzku4BxgLvA+OAz/PdHghcKqk18AZw0FyE2h+4RNL/AYsBNwNjgL9KWpOsstCQfJ2ZmZmZmZmZWakpIlLHYGZmZmZmZmZmixBJbSLiqzyh53Hg5xExKnVcZmZmZmZmZmZNnSv6mJmZmZmZmZnZvLpc0rpAK2Cgk3zMzMzMzMzMzIrhij5mZmZmZmZmZmZmZmZmZmZmZouAZqkDMDMzMzMzMzMzMzMzMzMzMzOzOXOij5mZmZmZmZmZmZmZmZmZmZnZIsCJPmZmZmZmZmZmZmZmZmZmZmZmiwAn+piZmZmZmZmZmZmZmZmZmZmZLQKc6GNmZmZmZmZmZmZmZmZmZmZmtghwoo+ZmZmZmZmZmZmZmZmZmZmZ2SLg/wHwxycJm9mLtwAAAABJRU5ErkJggg==\n",
+ "text/plain": [
+ "\u003cFigure size 2880x720 with 1 Axes\u003e"
+ ]
+ },
+ "metadata": {
+ "needs_background": "light"
+ },
+ "output_type": "display_data"
+ }
+ ],
+ "source": [
+ "# visualize the merged COCO annotated JSON file\n",
+ "visualize_detailed_counts_horizontally(output_merged_file)"
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "collapsed_sections": [],
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/merge_coco_files_faster.ipynb b/official/projects/waste_identification_ml/pre_processing/merge_coco_files_faster.ipynb
new file mode 100644
index 00000000000..85b18d656a0
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/merge_coco_files_faster.ipynb
@@ -0,0 +1,202 @@
+{
+ "nbformat": 4,
+ "nbformat_minor": 0,
+ "metadata": {
+ "colab": {
+ "provenance": []
+ },
+ "kernelspec": {
+ "name": "python3",
+ "display_name": "Python 3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "source": [
+ "# Merge multiple COCO annotation JSON files into one file."
+ ],
+ "metadata": {
+ "id": "vAan5iCyEQp5"
+ }
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "Given multiple COCO annotated JSON files, your goal is to merge them into one COCO annotated JSON file.\n",
+ "\n",
+ "A merged COCO annotated JSON file is required where all the data is in one place and it becomes easy to split it into a training and validation JSON file according to the percentage ratio. In case you already have a validated COCO annotated JSON file, then this notebook can be used to merge multiple files into one training COCO annotated JSON file."
+ ],
+ "metadata": {
+ "id": "OXNMINymEVvW"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# Import necessary libraries\n",
+ "import tqdm\n",
+ "import json\n",
+ "import glob"
+ ],
+ "metadata": {
+ "id": "DWc_xka7ix0I"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "def merge_jsons(list_of_jsons):\n",
+ " \"\"\"\n",
+ " Merges a list of JSON files into a single JSON file.\n",
+ "\n",
+ " Args:\n",
+ " list_of_jsons: A list of JSON files to be merged.\n",
+ "\n",
+ " Returns:\n",
+ " A single JSON file containing the merged data.\n",
+ " \"\"\"\n",
+ "\n",
+ " num = 1\n",
+ " image_id = 0\n",
+ " images_list = []\n",
+ " categories_list = []\n",
+ " annotations_list = []\n",
+ " labels_dict = {}\n",
+ " mapping_images = {}\n",
+ " mapping_categories = {}\n",
+ "\n",
+ "\n",
+ " for i,json_file_path in tqdm.tqdm(enumerate(list_of_jsons)):\n",
+ " # read JSON file\n",
+ " with open(json_file_path) as json_file:\n",
+ " read_json = json.load(json_file)\n",
+ "\n",
+ " if len(read_json['images'][0]) != 0:\n",
+ " list_of_dic = []\n",
+ " list_of_dic_cat = []\n",
+ "\n",
+ "\n",
+ " # process images dictionary\n",
+ " for image in read_json['images']:\n",
+ " images_dict = {}\n",
+ " list_of_dic.append((image['id'], image_id))\n",
+ " images_dict['file_name'] = image['file_name']\n",
+ " images_dict['id'] = image_id\n",
+ " image_id += 1\n",
+ " images_dict['width'] = image['width']\n",
+ " images_dict['height'] = image['height']\n",
+ " images_list.append(images_dict)\n",
+ " mapping_images['file_{}'.format(i)] = dict(list_of_dic)\n",
+ "\n",
+ "\n",
+ " # process categories dictionary\n",
+ " for category in read_json['categories']:\n",
+ " list_of_dic_cat.append((category['id'], category['name']))\n",
+ " categories_dict = {}\n",
+ " if category['name'] not in labels_dict.keys():\n",
+ " if len(labels_dict.keys()) == 0:\n",
+ " labels_dict[read_json['categories'][0]['name']] = 1\n",
+ " else:\n",
+ " labels_dict[category['name']] = max(labels_dict.values()) + 1\n",
+ " categories_dict['supercategory'] = category['supercategory']\n",
+ " categories_dict['id'] = labels_dict[category['name']]\n",
+ " categories_dict['name'] = category['name']\n",
+ " categories_list.append(categories_dict)\n",
+ " else:\n",
+ " pass\n",
+ " mapping_categories['file_{}'.format(i)] = dict(list_of_dic_cat)\n",
+ "\n",
+ "\n",
+ " # process annotations dictionary\n",
+ " for annotation in read_json['annotations']:\n",
+ " annotations_dict = {}\n",
+ " annotations_dict['segmentation'] = annotation['segmentation']\n",
+ " annotations_dict['area'] = annotation['area']\n",
+ " annotations_dict['bbox'] = annotation['bbox']\n",
+ " annotations_dict['image_id'] = mapping_images['file_{}'.format(i)][annotation['image_id']]\n",
+ " annotations_dict['category_id'] = labels_dict[mapping_categories['file_{}'.format(i)][annotation['category_id']]]\n",
+ " annotations_dict['id'] = num\n",
+ " num +=1\n",
+ " annotations_dict['iscrowd'] = 0\n",
+ " annotations_list.append(annotations_dict)\n",
+ "\n",
+ "\n",
+ " final_json = {\n",
+ " 'images':images_list,\n",
+ " 'categories':categories_list,\n",
+ " 'annotations':annotations_list\n",
+ " }\n",
+ "\n",
+ " for i in final_json['annotations']:\n",
+ " i['segmentation'] = [max(i['segmentation'], key=len)]\n",
+ " i['bbox'] = list(map(lambda x: round(float(x)), i['bbox']))\n",
+ "\n",
+ " return final_json"
+ ],
+ "metadata": {
+ "id": "7kQVyWxYqVdj"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "files = glob.glob('/mydrive/sherman/**/*.json', recursive=True)\n",
+ "files"
+ ],
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "g-617H5ONlwk",
+ "outputId": "2bb90235-f1cb-49b8-d8cc-1e1d54f0a9f8"
+ },
+ "execution_count": null,
+ "outputs": [
+ {
+ "output_type": "execute_result",
+ "data": {
+ "text/plain": [
+ "['/mydrive/sherman/annotation_3.json',\n",
+ " '/mydrive/sherman/annotation_2.json',\n",
+ " '/mydrive/sherman/annotation_4.json',\n",
+ " '/mydrive/sherman/annotation_1.json']"
+ ]
+ },
+ "metadata": {},
+ "execution_count": 20
+ }
+ ]
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "data = merge_jsons(files)"
+ ],
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "XWTy85rMEka3",
+ "outputId": "85567bf2-27e0-4152-8ad1-6539b551a9fb"
+ },
+ "execution_count": null,
+ "outputs": [
+ {
+ "output_type": "stream",
+ "name": "stderr",
+ "text": [
+ "4it [00:00, 10.26it/s]\n"
+ ]
+ }
+ ]
+ }
+ ]
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/split_coco_files.ipynb b/official/projects/waste_identification_ml/pre_processing/split_coco_files.ipynb
new file mode 100644
index 00000000000..aef3b86674b
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/split_coco_files.ipynb
@@ -0,0 +1,319 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "m0qQu-luFmB5"
+ },
+ "source": [
+ "# Split one COCO annotation JSON file into training and validation JSON files."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "9NGkWKGrF3pc"
+ },
+ "source": [
+ "Given a single COCO annotated JSON file, your goal is to split them into training and validation COCO annotated JSON files.\n",
+ "\n",
+ " A single JSON file needs to be split into training and validation files. The output files will be further converted to TFRecord files using another notebook.\n",
+ "\n",
+ "This notebook uses a third party library to accomplish this task. The library can split the JSON files according to the ratio. We kept the validation file to contain 20% of the data.\n",
+ "\n",
+ "This notebook is an end to end example. When you run the notebook, it will take one JSON file and will split into a train and a val JSON file."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "GIjj-vE-n1e3"
+ },
+ "source": [
+ "**Note** - In this example, we assume that all our data is saved on Google drive and we will also write our outputs to Google drive. We also assume that the script will be used as a Google Colab notebook. But this can be changed according to the needs of users. They can modify this in case they are working on their local workstation, remote server or any other database. This colab notebook can be changed to a regular jupyter notebook running on a local machine according to the need of the users."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "QElyM7FtWv5E"
+ },
+ "source": [
+ "## **MUST DO** - Install and restart runtime"
+ ]
+ },
+ {
+ "metadata": {
+ "id": "H8pEG4JeS10V"
+ },
+ "cell_type": "code",
+ "source": [
+ "# Change the python version to 3.10\n",
+ "!wget -O mini.sh https://repo.anaconda.com/miniconda/Miniconda3-py310_25.1.1-2-Linux-x86_64.sh\n",
+ "!chmod +x mini.sh\n",
+ "!bash ./mini.sh -b -f -p /usr/local\n",
+ "!conda install -q -y jupyter\n",
+ "!conda install -q -y google-colab -c conda-forge\n",
+ "!python -m ipykernel install --name \"py310\" --user"
+ ],
+ "outputs": [],
+ "execution_count": null
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "WMy_xu64FJ1j"
+ },
+ "outputs": [],
+ "source": [
+ "# Install pyodi from the source code.\n",
+ "!git clone https://github.com/Gradiant/pyodi.git\n",
+ "%cd pyodi/\n",
+ "!pip install .\n",
+ "# RESTART THE RUNTIME in order to use this library"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "tySpWIuVFPj0"
+ },
+ "source": [
+ "## Run the below command to connect to your google drive"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "RfJAkMY9FSPz"
+ },
+ "outputs": [],
+ "source": [
+ "# import other libraries\n",
+ "from google.colab import drive\n",
+ "import pyodi\n",
+ "import subprocess\n",
+ "import sys\n",
+ "import os\n",
+ "import json\n",
+ "import numpy as np\n",
+ "import pandas as pd"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "AOLmsOOZFVdJ",
+ "outputId": "f7f6dba8-0872-4d21-d55d-2b95c42a06a4"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Mounted at /content/gdrive\n",
+ "Successful\n"
+ ]
+ }
+ ],
+ "source": [
+ "# connect to google drive\n",
+ "drive.mount('/content/gdrive')\n",
+ "\n",
+ "# making an alias for the root path\n",
+ "try:\n",
+ " !ln -s /content/gdrive/My\\ Drive/ /mydrive\n",
+ " print('Successful')\n",
+ "except Exception as e:\n",
+ " print(e)\n",
+ " print('Not successful')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "v0vJRt5qUOD_"
+ },
+ "source": [
+ "## Visualization function"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "HbNMcLBmUOZ2"
+ },
+ "outputs": [],
+ "source": [
+ "def data_creation(path: str) -\u003e pd.DataFrame:\n",
+ " \"\"\"Create a dataframe with the occurences of images and categories.\n",
+ " Args:\n",
+ " path: path to the annotated JSON file.\n",
+ " Returns:\n",
+ " dataset consisting of the counts of images and categories.\n",
+ " \"\"\"\n",
+ " # get annotation file data into a variable\n",
+ " with open(path) as json_file:\n",
+ " data = json.load(json_file)\n",
+ "\n",
+ " # count the occurance of each category and an image in the annotation file\n",
+ " category_names = [i['name'] for i in data['categories']]\n",
+ " category_ids = [i['category_id'] for i in data['annotations']]\n",
+ " image_ids = [i['image_id'] for i in data['annotations']]\n",
+ "\n",
+ " # create a dataframe\n",
+ " df = pd.DataFrame(\n",
+ " list(zip(category_ids, image_ids)), columns=['category_ids', 'image_ids'])\n",
+ " df = df.groupby('category_ids').agg(\n",
+ " object_count=('category_ids', 'count'),\n",
+ " image_count=('image_ids', 'nunique'))\n",
+ " df = df.reindex(range(1, len(data['categories']) + 1), fill_value=0)\n",
+ " df.index = category_names\n",
+ " return df\n",
+ "\n",
+ "def visualize_detailed_counts_horizontally(path: str) -\u003e None:\n",
+ " \"\"\"Plot a vertical bar graph showing the counts of images \u0026 categories.\n",
+ " Args:\n",
+ " path: path to the annotated JSON file.\n",
+ " \"\"\"\n",
+ " df = data_creation(path)\n",
+ " ax = df.plot(\n",
+ " kind='bar',\n",
+ " figsize=(40, 10),\n",
+ " xlabel='Categories',\n",
+ " ylabel='Counts',\n",
+ " width=0.8,\n",
+ " linewidth=1,\n",
+ " edgecolor='white') # rot = 0 for horizontal labeling\n",
+ " for p in ax.patches:\n",
+ " ax.annotate(\n",
+ " text=np.round(p.get_height()),\n",
+ " xy=(p.get_x() + p.get_width() / 2., p.get_height()),\n",
+ " ha='center',\n",
+ " va='top',\n",
+ " xytext=(4, 14),\n",
+ " textcoords='offset points')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "gmTfyRQo9pT3"
+ },
+ "source": [
+ "## Define the paths of inputs and outputs"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "nl5MrEPR9q9x"
+ },
+ "outputs": [],
+ "source": [
+ "input_file = '/mydrive/TFHub/jsons/merged.json' #@param {type:\"string\"}\n",
+ "output_folder = '/mydrive/TFHub/jsons/' #@param {type:\"string\"}"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "2E7P4_2eFaPB"
+ },
+ "source": [
+ "## Split coco annotation file into train and val COCO files"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "id": "9HLYrO4JGKFm",
+ "outputId": "a31b04fa-0d7c-4c22-cd18-58672d5a29e7"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "\u001b[32m2022-09-09 21:40:00.173\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36mpyodi.apps.coco.coco_split\u001b[0m:\u001b[36mrandom_split\u001b[0m:\u001b[36m183\u001b[0m - \u001b[1mGathering images...\u001b[0m\n",
+ "\u001b[32m2022-09-09 21:40:00.192\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36mpyodi.apps.coco.coco_split\u001b[0m:\u001b[36mrandom_split\u001b[0m:\u001b[36m194\u001b[0m - \u001b[1mGathering annotations...\u001b[0m\n",
+ "\u001b[32m2022-09-09 21:40:11.078\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36mpyodi.apps.coco.coco_split\u001b[0m:\u001b[36mrandom_split\u001b[0m:\u001b[36m218\u001b[0m - \u001b[1mSaving splits to file...\u001b[0m\n",
+ "/mydrive/TFHub/jsons/_train.json\n",
+ "/mydrive/TFHub/jsons/_val.json\n"
+ ]
+ }
+ ],
+ "source": [
+ "# split a COCO annotation file into train and val files\n",
+ "!pyodi coco random-split $input_file $output_folder --val-percentage 0.2\n",
+ "\n",
+ "# there will be two files with name '_train.json' and '_val.json' in the output_folder"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "wLnDJLIuMf8o"
+ },
+ "source": [
+ "## Visualization"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "2dNcl3XCMLDX"
+ },
+ "outputs": [],
+ "source": [
+ "# visualization of the input COCO annotated JSON file\n",
+ "visualize_detailed_counts_horizontally(input_file)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "GHZZ3aLbMO35"
+ },
+ "outputs": [],
+ "source": [
+ "# visualization of the training COCO annotated JSON file\n",
+ "print('Train JSON')\n",
+ "visualize_detailed_counts_horizontally(output_folder + '_train.json')\n",
+ "\n",
+ "print('Validation JSON')\n",
+ "# visualization of the validation COCO annotated JSON file\n",
+ "visualize_detailed_counts_horizontally(output_folder + '_val.json')"
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "provenance": []
+ },
+ "gpuClass": "standard",
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/projects/waste_identification_ml/pre_processing/visualize_instance_segmentation_annotations.ipynb b/official/projects/waste_identification_ml/pre_processing/visualize_instance_segmentation_annotations.ipynb
new file mode 100644
index 00000000000..d9d30eac24f
--- /dev/null
+++ b/official/projects/waste_identification_ml/pre_processing/visualize_instance_segmentation_annotations.ipynb
@@ -0,0 +1 @@
+{"nbformat":4,"nbformat_minor":0,"metadata":{"colab":{"provenance":[],"authorship_tag":"ABX9TyOUBLEtg5H/8iToDUvN0Zol"},"kernelspec":{"name":"python3","display_name":"Python 3"},"language_info":{"name":"python"}},"cells":[{"cell_type":"markdown","source":["# COCO Annotation Visualizer for Instance Segmentation"],"metadata":{"id":"kvDUF2GoWD74"}},{"cell_type":"markdown","source":["This Colab notebook is designed to visualize annotations from a single merged COCO JSON file containing instance segmentation annotations for multiple images. We will use Detectron2 to display images along with their corresponding segmentation masks, bounding boxes, and category labels."],"metadata":{"id":"YUTwkHOyWLk8"}},{"cell_type":"code","execution_count":null,"metadata":{"id":"8nBOXaLJVk_U"},"outputs":[],"source":["# Install and RESTART the runtime.\n","!git clone 'https://github.com/facebookresearch/detectron2'\n","!pip install 'git+https://github.com/facebookresearch/detectron2.git'"]},{"cell_type":"code","source":["#@title Imports\n","\n","from google.colab import drive\n","from detectron2.data import MetadataCatalog, DatasetCatalog\n","from detectron2.data.datasets import register_coco_instances\n","from detectron2.utils.visualizer import Visualizer\n","from google.colab.patches import cv2_imshow\n","import json\n","import random\n","import os\n","import cv2\n","from collections import Counter"],"metadata":{"id":"nEXQ2uV6Wjto"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":["# Connect to google drive in case you have data there.\n","drive.mount('/content/gdrive')\n","try:\n"," !ln -s /content/gdrive/My\\ Drive/ /mydrive\n"," print('Successful')\n","except Exception as e:\n"," print(e)\n"," print('Not successful')"],"metadata":{"id":"fu8n-_KCWvOn"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":["#@title Utils\n","\n","def unregister_dataset(dataset_name: str) -> None:\n"," \"\"\"\n"," Removes the specified dataset from the MetadataCatalog and DatasetCatalog.\n","\n"," Args:\n"," dataset_name: Name of the dataset to be removed.\n"," \"\"\"\n"," if dataset_name in MetadataCatalog.list():\n"," MetadataCatalog.pop(dataset_name)\n","\n"," if dataset_name in DatasetCatalog.list():\n"," DatasetCatalog.remove(dataset_name)\n","\n","\n","def read_coco_json(json_path: str) -> dict:\n"," \"\"\"\n"," Reads a COCO JSON file and returns its contents as a dictionary.\n","\n"," Args:\n"," json_path: Path to the COCO JSON file.\n","\n"," Returns:\n"," Parsed JSON data as a dictionary.\n"," \"\"\"\n"," with open(json_path) as file:\n"," return json.load(file)\n","\n","\n","def filter_dataset_by_category(dataset_dicts: list, category_id: int) -> list:\n"," \"\"\"\n"," Filters a dataset to only include images containing a specific category.\n","\n"," Args:\n"," dataset_dicts: List of dataset dictionaries containing image and annotation data.\n"," category_id: Category ID to filter annotations by.\n","\n"," Returns:\n"," A new dataset dictionary containing only images with the specified category.\n"," \"\"\"\n"," filtered_dataset_dicts = []\n","\n"," for data in dataset_dicts:\n"," # Filter annotations by the given category_id\n"," filtered_annotations = [ann for ann in data['annotations'] if ann['category_id'] == category_id]\n","\n"," # If there are annotations with the specified category_id, add the image data to filtered dataset\n"," if filtered_annotations:\n"," # Create a copy of the original dictionary to modify safely\n"," filtered_data = data.copy()\n"," # Update annotations to only include those with the specified category_id\n"," filtered_data['annotations'] = filtered_annotations\n"," filtered_dataset_dicts.append(filtered_data)\n","\n"," return filtered_dataset_dicts\n","\n","\n","def get_object_counts_per_category(data):\n"," \"\"\"\n"," Counts the number of objects for each category in a JSON file, including category ID,\n"," and sorts the output by ascending category ID.\n","\n"," Args:\n"," data: Parsed JSON data.\n","\n"," Returns:\n"," A list of tuples where each tuple contains (category_name, category_id, count),\n"," sorted by ascending category_id.\n"," \"\"\"\n"," # Map category IDs to names\n"," category_id_to_name = {category['id']: category['name'] for category in data.get('categories', [])}\n","\n"," # Count objects by category ID\n"," counts_by_id = Counter(anno['category_id'] for anno in data.get('annotations', []))\n","\n"," # Translate category IDs to names and include ID in the output\n"," counts_by_name_and_id = [\n"," (cat_id, category_id_to_name[cat_id], count)\n"," for cat_id, count in counts_by_id.items()\n"," ]\n","\n"," # Sort by category ID\n"," counts_by_name_and_id.sort(key=lambda x: x[0])\n","\n"," return counts_by_name_and_id"],"metadata":{"id":"tegfWXnHWwnQ","cellView":"form"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":["# Provide the path to images and corresponding merged coco json file.\n","images_dir = '/mydrive/circularnet/Client_ManualAnnotation_data/Benjamin/batch1/images/' #@param {type:\"string\"}\n","coco_json = '/mydrive/circularnet/Client_ManualAnnotation_data/Benjamin/batch1/annotations/bottle_vs_non-bottle.json' #@param {type:\"string\"}\n","data = read_coco_json(coco_json)"],"metadata":{"id":"cNs6DB_AXi1g"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":["## Visualize randoom images of all categories"],"metadata":{"id":"g6BKabiOguX5"}},{"cell_type":"code","source":["# Number of images to visualize randomly.\n","num_of_images = 25 #@param {type:\"integer\"}\n","\n","# Unregister the previous dataset\n","unregister_dataset(\"dataset\")\n","\n","register_coco_instances(\"dataset\", {}, coco_json, images_dir)\n","metadata = MetadataCatalog.get(\"dataset\")\n","dataset_dicts = DatasetCatalog.get(\"dataset\")\n","\n","# Visualize images randomly.\n","for image in random.sample(dataset_dicts, num_of_images):\n"," print(os.path.basename(image[\"file_name\"]))\n"," img = cv2.imread(image[\"file_name\"])\n"," visualizer = Visualizer(img[:, :, ::-1], metadata=metadata, scale=0.5)\n"," out = visualizer.draw_dataset_dict(image)\n"," cv2_imshow(out.get_image()[:, :, ::-1])\n"," print()"],"metadata":{"id":"E8fp6cDlYenX"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":["## Visualize by image name"],"metadata":{"id":"EYF5rl1ufMi7"}},{"cell_type":"code","source":["image_name = 'google_benjamin_669154948um_d5bcad7c584849d7557e74704701be2cd9bd6326097bb2e7902f919dc5b6284e.jpeg' #@param {type:\"string\"}\n","\n","# Find the image in the dataset\n","target_image = None\n","for d in dataset_dicts:\n"," if os.path.basename(d[\"file_name\"]) == image_name:\n"," target_image = d\n"," break\n","\n","# Visualize the target image if found\n","if target_image:\n"," print(\"Found Image:\", os.path.basename(target_image[\"file_name\"]))\n"," img = cv2.imread(target_image[\"file_name\"])\n"," visualizer = Visualizer(img[:, :, ::-1], metadata=metadata, scale=0.5)\n"," out = visualizer.draw_dataset_dict(target_image)\n"," cv2_imshow(out.get_image()[:, :, ::-1]) # Use cv2.imshow if running locally\n","else:\n"," print(f\"Image '{image_name}' not found in the dataset.\")"],"metadata":{"id":"MjjSq81lfPi4"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":["## Visualize random images of a single category"],"metadata":{"id":"KRlD3u_mg0K2"}},{"cell_type":"code","source":["# Check the category IDs.\n","data['categories']"],"metadata":{"id":"m-FnQZG_g3-N"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":["num_of_images = 10 #@param {type:\"integer\"}\n","category_id_to_filter = 1 #@param {type:\"integer\"}\n","\n","unregister_dataset(\"filtered_dataset\")\n","filtered_dataset_dicts = filter_dataset_by_category(dataset_dicts, category_id_to_filter - 1)\n","DatasetCatalog.register(\"filtered_dataset\", lambda: filtered_dataset_dicts)\n","MetadataCatalog.get(\"filtered_dataset\").set(thing_classes=MetadataCatalog.get(\"dataset\").thing_classes)\n","\n","for image in random.sample(filtered_dataset_dicts, num_of_images):\n"," print(os.path.basename(image[\"file_name\"]))\n"," img = cv2.imread(image[\"file_name\"])\n"," visualizer = Visualizer(img[:, :, ::-1], metadata=metadata, scale=0.5)\n"," out = visualizer.draw_dataset_dict(image)\n"," cv2_imshow(out.get_image()[:, :, ::-1])\n"," print()"],"metadata":{"id":"9yBlLda7jdmW"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":["## Visualize images of all categories"],"metadata":{"id":"5s8T2Vl9kKsq"}},{"cell_type":"code","source":["# Get the image counts per category in order to display the minimum images.\n","category_counts = get_object_counts_per_category(data)\n","category_counts"],"metadata":{"id":"4mxabYDRkQeC"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":["max_samples_per_category = 6 #@param {type:\"integer\"}\n","\n","# Define range of categories to visualize\n","CATEGORY_RANGE_START = 1 #@param {type:\"integer\"}\n","CATEGORY_RANGE_END = 1 #@param {type:\"integer\"}\n","CATEGORY_RANGE = (CATEGORY_RANGE_START, CATEGORY_RANGE_END) # User-defined range\n","\n","# Extract category names and image counts\n","category_names = [category[1] for category in category_counts] # List of category names\n","image_counts_per_category = [category[2] for category in category_counts] # Number of images per category\n","\n","# Define max number of images to visualize per category\n","max_samples_per_category = 6\n","\n","# Iterate over category indices and corresponding data\n","for category_index, (category_name, image_count) in enumerate(zip(category_names, image_counts_per_category), start=1):\n","\n"," # Ensure we only process categories within the defined range\n"," if category_index < CATEGORY_RANGE[0] or category_index > CATEGORY_RANGE[1]:\n"," continue\n","\n"," # Determine the number of images to sample\n"," num_images_to_sample = min(max_samples_per_category, image_count)\n","\n"," print(f\"Category ID: {category_index} | Name: {category_name} | Image Count: {image_count}\\n\")\n","\n"," # Unregister previous dataset if it exists\n"," unregister_dataset(\"filtered_dataset\")\n","\n"," # Filter dataset for the given category\n"," filtered_annotations = filter_dataset_by_category(dataset_dicts, category_index - 1)\n","\n"," # Register filtered dataset\n"," DatasetCatalog.register(\"filtered_dataset\", lambda d=filtered_annotations: d)\n"," MetadataCatalog.get(\"filtered_dataset\").set(thing_classes=MetadataCatalog.get(\"dataset\").thing_classes)\n","\n"," # Randomly sample images and visualize\n"," for annotation in random.sample(filtered_annotations, num_images_to_sample):\n"," print(os.path.basename(annotation[\"file_name\"]))\n"," image = cv2.imread(annotation[\"file_name\"])\n","\n"," # Create visualizers for original and annotated images\n"," visualizer_original = Visualizer(image[:, :, ::-1], metadata=metadata, scale=0.4)\n"," original_output = visualizer_original.get_output()\n","\n"," visualizer_annotated = Visualizer(image[:, :, ::-1], metadata=metadata, scale=0.4)\n"," annotated_output = visualizer_annotated.draw_dataset_dict(annotation)\n","\n"," # Combine original and annotated images\n"," combined_image = cv2.hconcat([original_output.get_image()[:, :, ::-1], annotated_output.get_image()[:, :, ::-1]])\n","\n"," # Display the combined image\n"," cv2_imshow(combined_image)\n"," print()\n"],"metadata":{"id":"88WpRUydj2Fb"},"execution_count":null,"outputs":[]}]}
\ No newline at end of file
diff --git a/official/projects/yolo/README.md b/official/projects/yolo/README.md
new file mode 100644
index 00000000000..7d011a3c884
--- /dev/null
+++ b/official/projects/yolo/README.md
@@ -0,0 +1,101 @@
+# YOLO Object Detectors, You Only Look Once
+
+[](https://arxiv.org/abs/1804.02767)
+[](https://arxiv.org/abs/2004.10934)
+[](https://arxiv.org/abs/2207.02696)
+
+This repository contains the implementation of the following papers.
+
+* YOLOv3: An Incremental Improvement:
+ [Paper](https://arxiv.org/abs/1804.02767)
+
+* YOLOv4: Optimal Speed and Accuracy of Object Detection:
+ [Paper](https://arxiv.org/abs/2004.10934)
+
+* YOLOv7: Trainable bag-of-freebies sets new state-of-the-art for real-time
+ object detectors: [Paper](https://arxiv.org/abs/2207.02696)
+
+## Description
+
+YOLO v1 the original implementation was released in 2015 providing a
+ground-breaking algorithm that would quickly process images and locate objects
+in a single pass through the detector. The original implementation used a
+backbone derived from state of the art object classifiers of the time, like
+[GoogLeNet](https://arxiv.org/abs/1409.4842) and
+[VGG](https://arxiv.org/abs/1409.1556). More attention was given to the novel
+YOLO Detection head that allowed for Object Detection with a single pass of an
+image. Though limited, the network could predict up to 90 bounding boxes per
+image, and was tested for about 80 classes per box. Also, the model can only
+make predictions at one scale. These attributes caused YOLO v1 to be more
+limited and less versatile, so as the year passed, the Developers continued to
+update and develop this model.
+
+In 2020, YOLO v3 and v4 serve as the upgrades of the YOLO network group. The
+model uses a custom backbone called Darknet53 that uses knowledge gained from
+the ResNet paper to improve its predictions. The new backbone also allows for
+objects to be detected at multiple scales. As for the new detection head, the
+model now predicts the bounding boxes using a set of anchor box priors (Anchor
+Boxes) as suggestions. Multiscale predictions in combination with Anchor boxes
+allow for the network to make up to 1000 object predictions on a single image.
+Finally, the new loss function forces the network to make better predictions by
+using Intersection over Union (IoU) to inform the model's confidence rather than
+relying on the mean squared error for the entire output.
+
+As of 2023, YOLOv7 further improves the previous versions of YOLOs by
+introducing ELAN and E-ELAN structures. These new architectures are designed
+to diversify the gradients ([Designing Network Design Strategies Through
+Gradient Path Analysis](https://arxiv.org/abs/2211.04800)) so that the learned
+models are more expressive. In addition, YOLOv7 introduces auxiliary losses to
+enhance training, as well as re-parameterization to improve inference speed.
+Apart from what is mentioning in the paper, YOLOv7 also uses OTA loss ([OTA:
+Optimal Transport Assignment for Object Detection](
+https://arxiv.org/abs/2103.14259)) which gives more gains on mAP.
+
+## Authors
+
+### YOLOv3 & v4
+
+* Vishnu Samardh Banna ([@GitHub vishnubanna](https://github.com/vishnubanna))
+* Anirudh Vegesana ([@GitHub anivegesana](https://github.com/anivegesana))
+* Akhil Chinnakotla ([@GitHub The-Indian-Chinna](https://github.com/The-Indian-Chinna))
+* Tristan Yan ([@GitHub Tyan3001](https://github.com/Tyan3001))
+* Naveen Vivek ([@GitHub naveen-vivek](https://github.com/naveen-vivek))
+* Jacob Zietek ([@GitHub jacob-zietek](https://github.com/jacob-zietek))
+
+### YOLOv7
+
+* Jiageng Zhang ([@Github Zarjagen](https://github.com/zarjagen))
+
+## Table of Contents
+
+* [Our Goal](#our-goal)
+* [Models in the library](#models-in-the-library)
+* [References](#references)
+
+
+## Our Goal
+
+Our goal with this model conversion is to provide implementation of the Backbone
+and YOLO Head. We have built the model in such a way that the YOLO head could be
+connected to a new, more powerful backbone if a person chose to.
+
+## Models in the library
+
+| Object Detectors | Classifiers |
+| :--------------: | :--------------: |
+| Yolo-v3 | Darknet53 |
+| Yolo-v3 tiny | CSPDarknet53 |
+| Yolo-v3 spp |
+| Yolo-v4 |
+| Yolo-v4 tiny |
+| Yolo-v4 csp |
+| Yolo-v4 large |
+| Yolo-v7 |
+| Yolo-v7-tiny |
+| Yolo-v7X |
+| Yolo-v7-nano |
+| Yolo-v7-pico |
+
+## Requirements
+[](https://github.com/tensorflow/tensorflow/releases/tag/v2.11.0)
+[](https://www.python.org/downloads/release/python-380/)
diff --git a/official/projects/yolo/__init__.py b/official/projects/yolo/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/yolo/common/__init__.py b/official/projects/yolo/common/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/common/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/yolo/common/registry_imports.py b/official/projects/yolo/common/registry_imports.py
new file mode 100644
index 00000000000..becaf7991c8
--- /dev/null
+++ b/official/projects/yolo/common/registry_imports.py
@@ -0,0 +1,40 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""All necessary imports for registration."""
+
+# pylint: disable=unused-import
+# pylint: disable=g-bad-import-order
+from official.vision import registry_imports
+
+# import configs
+from official.projects.yolo.configs import darknet_classification
+from official.projects.yolo.configs import yolo as yolo_config
+from official.projects.yolo.configs import yolov7 as yolov7_config
+
+# import modeling components
+from official.projects.yolo.modeling.backbones import darknet
+from official.projects.yolo.modeling.decoders import yolo_decoder
+from official.projects.yolo.modeling.backbones import yolov7 as yolov7_backbone
+from official.projects.yolo.modeling.decoders import yolov7 as yolov7_decoder
+
+# import tasks
+from official.projects.yolo.tasks import image_classification
+from official.projects.yolo.tasks import yolo as yolo_task
+from official.projects.yolo.tasks import yolov7 as yolov7_task
+
+# import optimization packages
+from official.projects.yolo.optimization import optimizer_factory
+from official.projects.yolo.optimization.configs import optimizer_config
+from official.projects.yolo.optimization.configs import optimization_config
diff --git a/official/projects/yolo/configs/__init__.py b/official/projects/yolo/configs/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/configs/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/configs/backbones.py b/official/projects/yolo/configs/backbones.py
similarity index 76%
rename from official/vision/beta/projects/yolo/configs/backbones.py
rename to official/projects/yolo/configs/backbones.py
index f397809a12d..75962ca01fb 100644
--- a/official/vision/beta/projects/yolo/configs/backbones.py
+++ b/official/projects/yolo/configs/backbones.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -31,6 +31,14 @@ class Darknet(hyperparams.Config):
use_reorg_input: bool = False
+@dataclasses.dataclass
+class YoloV7(hyperparams.Config):
+ model_id: str = 'yolov7'
+ min_level: int = 3
+ max_level: int = 5
+
+
@dataclasses.dataclass
class Backbone(backbones.Backbone):
- darknet: Darknet = Darknet()
+ darknet: Darknet = dataclasses.field(default_factory=Darknet)
+ yolov7: YoloV7 = dataclasses.field(default_factory=YoloV7)
diff --git a/official/vision/beta/projects/yolo/configs/darknet_classification.py b/official/projects/yolo/configs/darknet_classification.py
similarity index 52%
rename from official/vision/beta/projects/yolo/configs/darknet_classification.py
rename to official/projects/yolo/configs/darknet_classification.py
index f1bd29a0972..a1ac38f9218 100644
--- a/official/vision/beta/projects/yolo/configs/darknet_classification.py
+++ b/official/projects/yolo/configs/darknet_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,8 @@
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.modeling import hyperparams
-from official.vision.beta.projects.yolo.configs import backbones
+from official.modeling import optimization
+from official.projects.yolo.configs import backbones
from official.vision.configs import common
from official.vision.configs import image_classification as imc
@@ -30,10 +31,17 @@ class ImageClassificationModel(hyperparams.Config):
"""Image classification model config."""
num_classes: int = 0
input_size: List[int] = dataclasses.field(default_factory=lambda: [224, 224])
- backbone: backbones.Backbone = backbones.Backbone(
- type='darknet', darknet=backbones.Darknet())
+ backbone: backbones.Backbone = dataclasses.field(
+ # pylint: disable=g-long-lambda
+ default_factory=lambda: backbones.Backbone(
+ type='darknet', darknet=backbones.Darknet()
+ )
+ # pylint: enable=g-long-lambda
+ )
dropout_rate: float = 0.0
- norm_activation: common.NormActivation = common.NormActivation()
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
# Adds a Batch Normalization layer pre-GlobalAveragePooling in classification.
add_head_batch_norm: bool = False
kernel_initializer: str = 'VarianceScaling'
@@ -44,18 +52,28 @@ class Losses(hyperparams.Config):
one_hot: bool = True
label_smoothing: float = 0.0
l2_weight_decay: float = 0.0
+ loss_weight: float = 1.0
+ soft_labels: bool = False
+ use_binary_cross_entropy: bool = False
@dataclasses.dataclass
class ImageClassificationTask(cfg.TaskConfig):
"""The model config."""
- model: ImageClassificationModel = ImageClassificationModel()
- train_data: imc.DataConfig = imc.DataConfig(is_training=True)
- validation_data: imc.DataConfig = imc.DataConfig(is_training=False)
- evaluation: imc.Evaluation = imc.Evaluation()
- losses: Losses = Losses()
+ model: ImageClassificationModel = dataclasses.field(
+ default_factory=ImageClassificationModel
+ )
+ train_data: imc.DataConfig = dataclasses.field(
+ default_factory=lambda: imc.DataConfig(is_training=True)
+ )
+ validation_data: imc.DataConfig = dataclasses.field(
+ default_factory=lambda: imc.DataConfig(is_training=False)
+ )
+ evaluation: imc.Evaluation = dataclasses.field(default_factory=imc.Evaluation)
+ losses: Losses = dataclasses.field(default_factory=Losses)
gradient_clip_norm: float = 0.0
logging_dir: Optional[str] = None
+ freeze_backbone: bool = False
@exp_factory.register_config_factory('darknet_classification')
@@ -63,8 +81,23 @@ def darknet_classification() -> cfg.ExperimentConfig:
"""Image classification general."""
return cfg.ExperimentConfig(
task=ImageClassificationTask(),
- trainer=cfg.TrainerConfig(),
+ trainer=cfg.TrainerConfig(
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {'type': 'sgd', 'sgd': {'momentum': 0.9}},
+ 'learning_rate': {
+ 'type': 'polynomial',
+ 'initial_learning_rate': 0.1,
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_learning_rate': 0,
+ },
+ },
+ })
+ ),
restrictions=[
'task.train_data.is_training != None',
- 'task.validation_data.is_training != None'
- ])
+ 'task.validation_data.is_training != None',
+ ],
+ )
diff --git a/official/vision/beta/projects/yolo/configs/decoders.py b/official/projects/yolo/configs/decoders.py
old mode 100755
new mode 100644
similarity index 82%
rename from official/vision/beta/projects/yolo/configs/decoders.py
rename to official/projects/yolo/configs/decoders.py
index 2a796a1e29b..b38615f1bbb
--- a/official/vision/beta/projects/yolo/configs/decoders.py
+++ b/official/projects/yolo/configs/decoders.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -42,7 +42,14 @@ class YoloDecoder(hyperparams.Config):
activation: Optional[str] = 'same'
+@dataclasses.dataclass
+class YoloV7(hyperparams.Config):
+ model_id: str = 'yolov7'
+ use_separable_conv: bool = False
+
+
@dataclasses.dataclass
class Decoder(decoders.Decoder):
type: Optional[str] = 'yolo_decoder'
- yolo_decoder: YoloDecoder = YoloDecoder()
+ yolo_decoder: YoloDecoder = dataclasses.field(default_factory=YoloDecoder)
+ yolov7: YoloV7 = dataclasses.field(default_factory=YoloV7)
diff --git a/official/vision/beta/projects/yolo/configs/experiments/darknet/csp_darknet53.yaml b/official/projects/yolo/configs/experiments/darknet/csp_darknet53.yaml
similarity index 100%
rename from official/vision/beta/projects/yolo/configs/experiments/darknet/csp_darknet53.yaml
rename to official/projects/yolo/configs/experiments/darknet/csp_darknet53.yaml
diff --git a/official/vision/beta/projects/yolo/configs/experiments/darknet/csp_darknet53_tfds.yaml b/official/projects/yolo/configs/experiments/darknet/csp_darknet53_tfds.yaml
similarity index 100%
rename from official/vision/beta/projects/yolo/configs/experiments/darknet/csp_darknet53_tfds.yaml
rename to official/projects/yolo/configs/experiments/darknet/csp_darknet53_tfds.yaml
diff --git a/official/vision/beta/projects/yolo/configs/experiments/darknet/darknet53.yaml b/official/projects/yolo/configs/experiments/darknet/darknet53.yaml
similarity index 100%
rename from official/vision/beta/projects/yolo/configs/experiments/darknet/darknet53.yaml
rename to official/projects/yolo/configs/experiments/darknet/darknet53.yaml
diff --git a/official/vision/beta/projects/yolo/configs/experiments/darknet/darknet53_tfds.yaml b/official/projects/yolo/configs/experiments/darknet/darknet53_tfds.yaml
similarity index 100%
rename from official/vision/beta/projects/yolo/configs/experiments/darknet/darknet53_tfds.yaml
rename to official/projects/yolo/configs/experiments/darknet/darknet53_tfds.yaml
diff --git a/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_gpu.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_gpu.yaml
new file mode 100644
index 00000000000..725ab3408a7
--- /dev/null
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_gpu.yaml
@@ -0,0 +1,102 @@
+# --experiment_type=scaled_yolo
+# Use 8 V100 GPU for training and validation.
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float16'
+task:
+ gradient_clip_norm: 10
+ model:
+ input_size: [640, 640, 3]
+ backbone:
+ type: 'darknet'
+ darknet:
+ model_id: 'altered_cspdarknet53'
+ max_level: 5
+ min_level: 3
+ decoder:
+ type: yolo_decoder
+ yolo_decoder:
+ version: v4
+ type: csp
+ head:
+ smart_bias: true
+ detection_generator:
+ box_type:
+ 'all': scaled
+ scale_xy:
+ 'all': 2.0
+ max_boxes: 300
+ nms_version: iou
+ iou_thresh: 0.001
+ nms_thresh: 0.65
+ loss:
+ use_scaled_loss: true
+ update_on_repeat: true
+ box_loss_type:
+ 'all': ciou
+ ignore_thresh:
+ 'all': 0.0
+ iou_normalizer:
+ 'all': 0.05
+ cls_normalizer:
+ 'all': 0.3
+ object_normalizer:
+ '5': 0.28
+ '4': 0.70
+ '3': 2.80
+ objectness_smooth:
+ 'all': 1.0
+ norm_activation:
+ use_sync_bn: true
+ num_classes: 80
+ anchor_boxes:
+ anchors_per_scale: 3
+ boxes: [box: [12, 16], box: [19, 36], box: [40, 28],
+ box: [36, 75], box: [76, 55], box: [72, 146],
+ box: [142, 110], box: [192, 243], box: [459, 401]]
+ train_data:
+ input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/train*'
+ global_batch_size: 64
+ dtype: float16
+ shuffle_buffer_size: 10000
+ parser:
+ mosaic:
+ mosaic_frequency: 1.0
+ mixup_frequency: 0.0
+ mosaic_crop_mode: 'scale'
+ mosaic_center: 0.25
+ aug_scale_min: 0.1
+ aug_scale_max: 1.9
+ max_num_instances: 300
+ letter_box: true
+ random_flip: true
+ aug_rand_translate: 0.1
+ area_thresh: 0.1
+ validation_data:
+ input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/val*'
+ global_batch_size: 64
+ dtype: float16
+trainer:
+ best_checkpoint_eval_metric: 'AP'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ checkpoint_interval: 3696
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.2
+ decay_steps: 554400
+ initial_learning_rate: 0.04
+ optimizer:
+ sgd_torch:
+ warmup_steps: 5544
+ global_clipnorm: 10
+ warmup:
+ linear:
+ warmup_steps: 5544
+ steps_per_loop: 1848
+ summary_interval: 1848
+ train_steps: 554400
+ validation_interval: 3696
+ validation_steps: 79
diff --git a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_tpu.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_tpu.yaml
similarity index 91%
rename from official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_tpu.yaml
rename to official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_tpu.yaml
index e4a00b1d8c9..149b30b4e20 100644
--- a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_tpu.yaml
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_csp_640_tpu.yaml
@@ -1,4 +1,5 @@
# --experiment_type=scaled_yolo
+# Use 4x4 DF for training
# mAP 47.6
runtime:
distribution_strategy: 'tpu'
@@ -26,7 +27,7 @@ task:
scale_xy:
'all': 2.0
max_boxes: 300
- nms_type: iou
+ nms_version: iou
iou_thresh: 0.001
nms_thresh: 0.65
loss:
@@ -72,3 +73,7 @@ task:
area_thresh: 0.1
validation_data:
input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/val*'
+trainer:
+ best_checkpoint_eval_metric: 'AP'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
diff --git a/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_gpu.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_gpu.yaml
new file mode 100644
index 00000000000..20a876f2fb1
--- /dev/null
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_gpu.yaml
@@ -0,0 +1,102 @@
+# --experiment_type=scaled_yolo
+# Use 8 V100 GPU for training and validation.
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ input_size: [896, 896, 3]
+ backbone:
+ type: 'darknet'
+ darknet:
+ model_id: 'csp-large'
+ max_level: 5
+ min_level: 3
+ width_scale: 1.00
+ depth_scale: 1.00
+ decoder:
+ type: yolo_decoder
+ yolo_decoder:
+ version: v4
+ type: csp_large
+ head:
+ smart_bias: true
+ detection_generator:
+ box_type:
+ 'all': scaled
+ scale_xy:
+ 'all': 2.0
+ max_boxes: 300
+ nms_version: iou
+ iou_thresh: 0.001
+ nms_thresh: 0.65
+ loss:
+ use_scaled_loss: true
+ update_on_repeat: true
+ box_loss_type:
+ 'all': ciou
+ ignore_thresh:
+ 'all': 0.0
+ iou_normalizer:
+ 'all': 0.05
+ cls_normalizer:
+ 'all': 0.3
+ object_normalizer:
+ '5': 0.28
+ '4': 0.70
+ '3': 2.80
+ objectness_smooth:
+ 'all': 1.0
+ norm_activation:
+ use_sync_bn: true
+ num_classes: 80
+ anchor_boxes:
+ anchors_per_scale: 4
+ boxes: [box: [13, 17], box: [31, 25], box: [24, 51], box: [61, 45],
+ box: [48, 102], box: [119, 96], box: [97, 189], box: [217, 184],
+ box: [171, 384], box: [324, 451], box: [616, 618], box: [800, 800]]
+ train_data:
+ input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/train*'
+ global_batch_size: 16
+ dtype: float32
+ shuffle_buffer_size: 10000
+ parser:
+ mosaic:
+ mosaic_frequency: 1.0
+ mixup_frequency: 0.2
+ mosaic_crop_mode: 'scale'
+ mosaic_center: 0.25
+ aug_scale_min: 0.1
+ aug_scale_max: 1.9
+ max_num_instances: 300
+ letter_box: true
+ random_flip: true
+ aug_rand_translate: 0.1
+ area_thresh: 0.1
+ validation_data:
+ input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/val*'
+ global_batch_size: 64
+ dtype: float32
+trainer:
+ best_checkpoint_eval_metric: 'AP'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ checkpoint_interval: 7392
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.2
+ decay_steps: 1848000 # 250 epochs
+ initial_learning_rate: 0.000625
+ optimizer:
+ sgd_torch:
+ warmup_steps: 22176
+ warmup:
+ linear:
+ warmup_steps: 22176
+ steps_per_loop: 7392
+ summary_interval: 7392
+ train_steps: 1848000 # 250 epochs
+ validation_interval: 7392
+ validation_steps: 79
diff --git a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_tpu.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_tpu.yaml
similarity index 97%
rename from official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_tpu.yaml
rename to official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_tpu.yaml
index dcfb93cb9a7..36e67ecfaf4 100644
--- a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_tpu.yaml
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p5_896_tpu.yaml
@@ -1,5 +1,6 @@
# --experiment_type=scaled_yolo
-# mAP 51.1%
+# Use 4x4 DF for training
+# mAP 51.1
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'float32'
@@ -28,7 +29,7 @@ task:
scale_xy:
'all': 2.0
max_boxes: 300
- nms_type: iou
+ nms_version: iou
iou_thresh: 0.001
nms_thresh: 0.65
loss:
diff --git a/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_gpu.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_gpu.yaml
new file mode 100644
index 00000000000..02a911fce08
--- /dev/null
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_gpu.yaml
@@ -0,0 +1,103 @@
+# --experiment_type=scaled_yolo
+# Use 8 V100 GPU for training and validation.
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ input_size: [1280, 1280, 3]
+ backbone:
+ type: 'darknet'
+ darknet:
+ model_id: 'csp-large'
+ max_level: 6
+ min_level: 3
+ width_scale: 1.00
+ depth_scale: 1.00
+ decoder:
+ type: yolo_decoder
+ yolo_decoder:
+ version: v4
+ type: csp_large
+ head:
+ smart_bias: true
+ detection_generator:
+ box_type:
+ 'all': scaled
+ scale_xy:
+ 'all': 2.0
+ max_boxes: 300
+ nms_version: iou
+ iou_thresh: 0.001
+ nms_thresh: 0.65
+ loss:
+ use_scaled_loss: true
+ update_on_repeat: true
+ box_loss_type:
+ 'all': ciou
+ ignore_thresh:
+ 'all': 0.0
+ iou_normalizer:
+ 'all': 0.05
+ cls_normalizer:
+ 'all': 0.5
+ object_normalizer:
+ '6': 0.07
+ '5': 0.29
+ '4': 0.7
+ '3': 2.8
+ objectness_smooth:
+ 'all': 1.0
+ norm_activation:
+ use_sync_bn: true
+ num_classes: 80
+ anchor_boxes:
+ anchors_per_scale: 4
+ boxes: [box: [13, 17], box: [31, 25], box: [24, 51], box: [61, 45],
+ box: [61, 45], box: [48, 102], box: [119, 96], box: [97, 189],
+ box: [97, 189], box: [217, 184], box: [171, 384], box: [324, 451],
+ box: [324, 451], box: [545, 357], box: [616, 618], box: [1024, 1024]]
+ train_data:
+ input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/train*'
+ global_batch_size: 8
+ dtype: float32
+ parser:
+ mosaic:
+ mosaic_frequency: 1.0
+ mixup_frequency: 0.2
+ mosaic_crop_mode: 'scale'
+ mosaic_center: 0.25
+ aug_scale_min: 0.1
+ aug_scale_max: 1.9
+ max_num_instances: 300
+ letter_box: true
+ random_flip: true
+ aug_rand_translate: 0.1
+ area_thresh: 0.1
+ validation_data:
+ input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/val*'
+ global_batch_size: 32
+ dtype: float32
+trainer:
+ best_checkpoint_eval_metric: 'AP'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ checkpoint_interval: 7392
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.2
+ decay_steps: 3696000 # 250 epochs
+ initial_learning_rate: 0.0003
+ optimizer:
+ sgd_torch:
+ warmup_steps: 44352
+ warmup:
+ linear:
+ warmup_steps: 44352
+ steps_per_loop: 14784
+ summary_interval: 14784
+ train_steps: 3696000 # 250 epochs
+ validation_interval: 14784
+ validation_steps: 157
diff --git a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_tpu.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_tpu.yaml
similarity index 97%
rename from official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_tpu.yaml
rename to official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_tpu.yaml
index 56bbead3f47..c331aadeb65 100644
--- a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_tpu.yaml
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p6_1280_tpu.yaml
@@ -1,5 +1,6 @@
# --experiment_type=scaled_yolo
-# mAP 54%
+# Use 8x8 DF for training
+# mAP 53.9
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'float32'
@@ -28,7 +29,7 @@ task:
scale_xy:
'all': 2.0
max_boxes: 300
- nms_type: iou
+ nms_version: iou
iou_thresh: 0.001
nms_thresh: 0.65
loss:
diff --git a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p7_1536_tpu.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p7_1536_tpu.yaml
similarity index 96%
rename from official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p7_1536_tpu.yaml
rename to official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p7_1536_tpu.yaml
index 7273d9846d6..c70f08ae4a8 100644
--- a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p7_1536_tpu.yaml
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/detection/yolo_l_p7_1536_tpu.yaml
@@ -1,5 +1,6 @@
# --experiment_type=scaled_yolo
-# mAP 54.7%
+# Use 8x16 DF for training
+# mAP 54.8
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'float32'
@@ -28,7 +29,7 @@ task:
scale_xy:
'all': 2.0
max_boxes: 300
- nms_type: iou
+ nms_version: iou
iou_thresh: 0.001
nms_thresh: 0.65
loss:
@@ -62,6 +63,7 @@ task:
box: [812, 393], box: [477, 808], box: [1070, 908], box: [1408, 1408]]
train_data:
input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/train*'
+ dtype: float32
shuffle_buffer_size: 10000
parser:
mosaic:
@@ -78,6 +80,7 @@ task:
area_thresh: 0.1
validation_data:
input_path: '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/val*'
+ dtype: float32
trainer:
best_checkpoint_eval_metric: 'AP'
best_checkpoint_export_subdir: 'best_ckpt'
diff --git a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/tpu/640.yaml b/official/projects/yolo/configs/experiments/scaled-yolo/tpu/640.yaml
similarity index 98%
rename from official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/tpu/640.yaml
rename to official/projects/yolo/configs/experiments/scaled-yolo/tpu/640.yaml
index c2f106ad86f..8249bb467bd 100644
--- a/official/vision/beta/projects/yolo/configs/experiments/scaled-yolo/tpu/640.yaml
+++ b/official/projects/yolo/configs/experiments/scaled-yolo/tpu/640.yaml
@@ -26,7 +26,7 @@ task:
scale_xy:
'all': 2.0
max_boxes: 300
- nms_type: iou
+ nms_version: iou
iou_thresh: 0.001
nms_thresh: 0.60
loss:
diff --git a/official/projects/yolo/configs/experiments/yolov4/detection/scaled_yolov4_1280_gpu.yaml b/official/projects/yolo/configs/experiments/yolov4/detection/scaled_yolov4_1280_gpu.yaml
new file mode 100644
index 00000000000..d757147c161
--- /dev/null
+++ b/official/projects/yolo/configs/experiments/yolov4/detection/scaled_yolov4_1280_gpu.yaml
@@ -0,0 +1,116 @@
+# This config is for demonstration purpose only. Parameters such as batch_size, train_steps should be overridden when used.
+# --experiment_type=scaled_yolo
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: multi_worker_mirrored
+ mixed_precision_dtype: float16
+task:
+ max_num_eval_detections: 300
+ model:
+ input_size: [1280, 1280, 3]
+ backbone:
+ type: 'darknet'
+ darknet:
+ model_id: 'csp-large'
+ max_level: 6
+ min_level: 3
+ width_scale: 1.00
+ depth_scale: 1.00
+ decoder:
+ type: yolo_decoder
+ yolo_decoder:
+ version: v4
+ type: csp_large
+ head:
+ smart_bias: true
+ detection_generator:
+ box_type:
+ 'all': scaled
+ scale_xy:
+ 'all': 2.0
+ max_boxes: 300
+ nms_version: iou
+ iou_thresh: 0.001
+ nms_thresh: 0.65
+ loss:
+ use_scaled_loss: true
+ update_on_repeat: true
+ box_loss_type:
+ 'all': ciou
+ ignore_thresh:
+ 'all': 0.0
+ iou_normalizer:
+ 'all': 0.05
+ cls_normalizer:
+ 'all': 0.5
+ object_normalizer:
+ '6': 0.07
+ '5': 0.29
+ '4': 0.7
+ '3': 2.8
+ objectness_smooth:
+ 'all': 1.0
+ norm_activation:
+ use_sync_bn: true
+ num_classes: 2 # BG + product
+ anchor_boxes:
+ anchors_per_scale: 4
+ boxes: [box: [13, 17], box: [31, 25], box: [24, 51], box: [61, 45],
+ box: [61, 45], box: [48, 102], box: [119, 96], box: [97, 189],
+ box: [97, 189], box: [217, 184], box: [171, 384], box: [324, 451],
+ box: [324, 451], box: [545, 357], box: [616, 618], box: [1024, 1024]]
+ train_data:
+ file_type: tfrecord
+ global_batch_size: 2
+ shuffle_buffer_size: 5000
+ decoder:
+ type: simple_decoder
+ simple_decoder:
+ regenerate_source_id: true
+ coco91_to_80: false
+ parser:
+ mosaic:
+ mosaic_frequency: 1.0
+ mixup_frequency: 0.2
+ mosaic_crop_mode: 'scale'
+ mosaic_center: 0.25
+ aug_scale_min: 0.1
+ aug_scale_max: 1.9
+ max_num_instances: 300
+ letter_box: true
+ random_flip: true
+ aug_rand_translate: 0.1
+ area_thresh: 0.1
+ validation_data:
+ file_type: tfrecord
+ global_batch_size: 2
+ decoder:
+ type: simple_decoder
+ simple_decoder:
+ regenerate_source_id: true
+ coco91_to_80: false
+ parser:
+ max_num_instances: 300
+ annotation_file: null
+trainer:
+ best_checkpoint_eval_metric: 'AP50'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ train_steps: 500
+ validation_steps: -1
+ validation_interval: 50
+ steps_per_loop: 50
+ summary_interval: 50
+ checkpoint_interval: 50
+ optimizer_config:
+ ema: null
+ learning_rate:
+ cosine:
+ decay_steps: 50
+ initial_learning_rate: 0.001
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0.00067
+ warmup_steps: 20
+ type: linear
diff --git a/official/vision/beta/projects/yolo/configs/experiments/yolov4/detection/yolov4_512_tpu.yaml b/official/projects/yolo/configs/experiments/yolov4/detection/yolov4_512_tpu.yaml
old mode 100755
new mode 100644
similarity index 99%
rename from official/vision/beta/projects/yolo/configs/experiments/yolov4/detection/yolov4_512_tpu.yaml
rename to official/projects/yolo/configs/experiments/yolov4/detection/yolov4_512_tpu.yaml
index 121d9e6d630..c9586b397c3
--- a/official/vision/beta/projects/yolo/configs/experiments/yolov4/detection/yolov4_512_tpu.yaml
+++ b/official/projects/yolo/configs/experiments/yolov4/detection/yolov4_512_tpu.yaml
@@ -30,7 +30,7 @@ task:
'4': 1.1
'3': 1.2
max_boxes: 200
- nms_type: iou
+ nms_version: iou
iou_thresh: 0.001
nms_thresh: 0.60
loss:
diff --git a/official/vision/beta/projects/yolo/configs/experiments/yolov4/imagenet_pretraining/cspdarknet53_256_tpu.yaml b/official/projects/yolo/configs/experiments/yolov4/imagenet_pretraining/cspdarknet53_256_tpu.yaml
similarity index 100%
rename from official/vision/beta/projects/yolo/configs/experiments/yolov4/imagenet_pretraining/cspdarknet53_256_tpu.yaml
rename to official/projects/yolo/configs/experiments/yolov4/imagenet_pretraining/cspdarknet53_256_tpu.yaml
diff --git a/official/projects/yolo/configs/experiments/yolov4/yolov4_tiny_416_tpu.yaml b/official/projects/yolo/configs/experiments/yolov4/yolov4_tiny_416_tpu.yaml
new file mode 100644
index 00000000000..4533912bf3e
--- /dev/null
+++ b/official/projects/yolo/configs/experiments/yolov4/yolov4_tiny_416_tpu.yaml
@@ -0,0 +1,141 @@
+# --experiment_type=yolo_darknet
+# 21.21 AP
+# 41.68 AP50
+# 19.12 AP75
+# 29.59 APl
+# 23.94 APm
+# 9.67 APs
+
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ smart_bias_lr: 0.0
+ model:
+ darknet_based_model: true
+ input_size: [416, 416, 3]
+ backbone:
+ type: 'darknet'
+ darknet:
+ model_id: 'cspdarknettiny'
+ max_level: 5
+ min_level: 4
+ decoder:
+ type: yolo_decoder
+ yolo_decoder:
+ version: v4
+ type: tiny
+ head:
+ smart_bias: true
+ detection_generator:
+ box_type:
+ 'all': original
+ scale_xy:
+ 'all': 1.05
+ max_boxes: 300
+ nms_type: iou
+ iou_thresh: 0.001
+ nms_thresh: 0.60
+ loss:
+ use_scaled_loss: false
+ box_loss_type:
+ 'all': ciou
+ ignore_thresh:
+ 'all': 0.7
+ iou_normalizer:
+ 'all': 0.07
+ cls_normalizer:
+ 'all': 1.0
+ object_normalizer:
+ 'all': 1.0
+ objectness_smooth:
+ 'all': 0.0
+ max_delta:
+ 'all': .inf
+ norm_activation:
+ activation: leaky
+ norm_epsilon: 0.00001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ num_classes: 80
+ anchor_boxes:
+ anchors_per_scale: 3
+ boxes: [box: [10, 14], box: [23, 27], box: [37, 58], box: [81, 82], box: [135, 169], box: [344, 319]]
+ train_data:
+ prefetch_buffer_size: 32
+ global_batch_size: 512
+ dtype: float32
+ input_path: 'gs://cam2-datasets/coco/train*'
+ is_training: true
+ drop_remainder: true
+ seed: 1000
+ parser:
+ mosaic:
+ mosaic_frequency: 0.0
+ mixup_frequency: 0.0
+ max_num_instances: 300
+ letter_box: false
+ random_flip: true
+ aug_rand_saturation: 1.5
+ aug_rand_brightness: 1.5
+ aug_rand_hue: 0.1
+ aug_scale_min: 0.50
+ aug_scale_max: 1.5
+ aug_rand_translate: 0.0
+ jitter: 0.3
+ area_thresh: 0.0
+ random_pad: true
+ use_tie_breaker: false
+ best_match_only: false
+ anchor_thresh: 1.0
+ validation_data:
+ prefetch_buffer_size: 32
+ global_batch_size: 8
+ dtype: float32
+ input_path: 'gs://cam2-datasets/coco/val*'
+ is_training: false
+ drop_remainder: true
+ parser:
+ max_num_instances: 300
+ letter_box: false
+ use_tie_breaker: false
+ best_match_only: false
+ anchor_thresh: 1.0
+ weight_decay: 0.000
+ init_checkpoint: null
+ init_checkpoint_modules: null
+ annotation_file: null
+trainer:
+ best_checkpoint_eval_metric: 'AP'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ train_steps: 553126
+ validation_steps: 625
+ steps_per_loop: 920
+ summary_interval: 920
+ validation_interval: 9200
+ checkpoint_interval: 920
+ optimizer_config:
+ ema:
+ average_decay: 0.9998
+ trainable_weights_only: false
+ dynamic_decay: true
+ learning_rate:
+ type: stepwise
+ stepwise:
+ boundaries: [442500, 497814]
+ name: PiecewiseConstantDecay
+ values: [0.04176, 0.004176, 0.0004176]
+ optimizer:
+ type: sgd_torch
+ sgd_torch:
+ momentum: 0.9
+ momentum_start: 0.9
+ nesterov: true
+ warmup_steps: 2000
+ weight_decay: 0.0005
+ name: SGD
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_steps: 2000
diff --git a/official/projects/yolo/configs/experiments/yolov7/detection/yolov7.yaml b/official/projects/yolo/configs/experiments/yolov7/detection/yolov7.yaml
new file mode 100644
index 00000000000..887c25078dd
--- /dev/null
+++ b/official/projects/yolo/configs/experiments/yolov7/detection/yolov7.yaml
@@ -0,0 +1,147 @@
+# YOLOv7-P5
+# --experiment_type=coco_yolov7
+# mAP: 50.4
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: float32
+task:
+ smart_bias_lr: 0.1
+ weight_decay: 0.0
+ model:
+ input_size: [640, 640, 3]
+ min_level: 3
+ max_level: 5
+ num_classes: 80
+ anchor_boxes:
+ anchors_per_scale: 3
+ boxes: [box: [12, 16], box: [19, 36], box: [40, 28],
+ box: [36, 75], box: [76, 55], box: [72, 146],
+ box: [142, 110], box: [192, 243], box: [459, 401]]
+ backbone:
+ type: yolov7
+ yolov7:
+ max_level: 5
+ min_level: 3
+ model_id: yolov7
+ decoder:
+ type: yolov7
+ yolov7:
+ model_id: yolov7
+ head:
+ num_anchors: 3
+ detection_generator:
+ box_type:
+ '3': scaled
+ '4': scaled
+ '5': scaled
+ all: null
+ scale_xy:
+ '3': 2.0
+ '4': 2.0
+ '5': 2.0
+ all: null
+ path_scales:
+ '3': 8
+ '4': 16
+ '5': 32
+ all: null
+ nms_version: iou
+ iou_thresh: 0.001
+ nms_thresh: 0.7
+ max_boxes: 300
+ pre_nms_points: 5000
+ loss:
+ gamma: 0.0
+ box_weight: 0.05
+ obj_weight: 0.7
+ cls_weight: 0.3
+ label_smoothing: 0.0
+ anchor_threshold: 4.0
+ iou_mix_ratio: 1.0
+ auto_balance: false
+ use_ota: true
+ norm_activation:
+ activation: swish
+ norm_momentum: 0.03
+ norm_epsilon: 0.001
+ use_sync_bn: true
+ train_data:
+ input_path: /readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/train*
+ is_training: true
+ global_batch_size: 256
+ dtype: float32
+ parser:
+ max_num_instances: 300
+ letter_box: true
+ random_flip: true
+ random_pad: false
+ jitter: 0.0
+ aug_scale_min: 1.0
+ aug_scale_max: 1.0
+ aug_rand_translate: 0.2
+ aug_rand_saturation: 0.7
+ aug_rand_brightness: 0.4
+ aug_rand_hue: 0.015
+ aug_rand_angle: 0.0
+ aug_rand_perspective: 0.0
+ use_tie_breaker: true
+ best_match_only: true
+ anchor_thresh: 4.0
+ area_thresh: 0.0
+ mosaic:
+ mosaic_frequency: 1.0
+ mosaic9_frequency: 0.2
+ mixup_frequency: 0.15
+ mosaic_crop_mode: scale
+ mosaic_center: 0.25
+ aug_scale_min: 0.1
+ aug_scale_max: 1.9
+ validation_data:
+ input_path: /readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/val*
+ is_training: false
+ global_batch_size: 16
+ dtype: float32
+ drop_remainder: true
+ parser:
+ max_num_instances: 300
+ letter_box: true
+ use_tie_breaker: true
+ best_match_only: true
+ anchor_thresh: 4.0
+ area_thresh: 0.0
+trainer:
+ best_checkpoint_export_subdir: best_ckpt
+ best_checkpoint_eval_metric: AP
+ best_checkpoint_metric_comp: higher
+ train_steps: 138600
+ validation_steps: 312
+ validation_interval: 2310
+ steps_per_loop: 462
+ summary_interval: 462
+ checkpoint_interval: 462
+ optimizer_config:
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ dynamic_decay: true
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.01
+ alpha: 0.1
+ decay_steps: 138600
+ type: cosine
+ optimizer:
+ sgd_torch:
+ bias_keys: [bias, beta]
+ weight_keys: [kernel, weight]
+ momentum: 0.937
+ momentum_start: 0.8
+ nesterov: true
+ warmup_steps: 1386
+ weight_decay: 0.002
+ type: sgd_torch
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 1386
+ type: linear
diff --git a/official/projects/yolo/configs/experiments/yolov7/detection/yolov7_gpu.yaml b/official/projects/yolo/configs/experiments/yolov7/detection/yolov7_gpu.yaml
new file mode 100644
index 00000000000..24c72526b08
--- /dev/null
+++ b/official/projects/yolo/configs/experiments/yolov7/detection/yolov7_gpu.yaml
@@ -0,0 +1,158 @@
+# This config is for demonstration purpose only. Parameters such as batch_size, train_steps should
+# be overridden when used.
+# YOLOv7-P5
+# --experiment_type=yolov7_coco
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: multi_worker_mirrored
+ mixed_precision_dtype: float32
+task:
+ smart_bias_lr: 0.1
+ weight_decay: 0.0
+ model:
+ input_size: [640, 640, 3]
+ min_level: 3
+ max_level: 5
+ num_classes: 80
+ anchor_boxes:
+ anchors_per_scale: 3
+ boxes: [box: [12, 16], box: [19, 36], box: [40, 28],
+ box: [36, 75], box: [76, 55], box: [72, 146],
+ box: [142, 110], box: [192, 243], box: [459, 401]]
+ backbone:
+ type: yolov7
+ yolov7:
+ max_level: 5
+ min_level: 3
+ model_id: yolov7
+ decoder:
+ type: yolov7
+ yolov7:
+ model_id: yolov7
+ head:
+ num_anchors: 3
+ detection_generator:
+ box_type:
+ '3': scaled
+ '4': scaled
+ '5': scaled
+ all: null
+ scale_xy:
+ '3': 2.0
+ '4': 2.0
+ '5': 2.0
+ all: null
+ path_scales:
+ '3': 8
+ '4': 16
+ '5': 32
+ all: null
+ nms_version: iou
+ iou_thresh: 0.001
+ nms_thresh: 0.7
+ max_boxes: 300
+ pre_nms_points: 5000
+ loss:
+ gamma: 0.0
+ box_weight: 0.05
+ obj_weight: 0.7
+ cls_weight: 0.3
+ label_smoothing: 0.0
+ anchor_threshold: 4.0
+ iou_mix_ratio: 1.0
+ auto_balance: false
+ use_ota: true
+ norm_activation:
+ activation: swish
+ norm_momentum: 0.03
+ norm_epsilon: 0.001
+ use_sync_bn: true
+ train_data:
+ input_path: /readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/train*
+ shuffle_buffer_size: 5000
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ coco91_to_80: false
+ is_training: true
+ global_batch_size: 32
+ dtype: float32
+ parser:
+ max_num_instances: 100
+ letter_box: true
+ random_flip: true
+ random_pad: false
+ jitter: 0.0
+ aug_scale_min: 1.0
+ aug_scale_max: 1.0
+ aug_rand_translate: 0.2
+ aug_rand_saturation: 0.7
+ aug_rand_brightness: 0.4
+ aug_rand_hue: 0.015
+ aug_rand_angle: 0.0
+ aug_rand_perspective: 0.0
+ use_tie_breaker: true
+ best_match_only: true
+ anchor_thresh: 4.0
+ area_thresh: 0.0
+ mosaic:
+ mosaic_frequency: 1.0
+ mosaic9_frequency: 0.2
+ mixup_frequency: 0.15
+ mosaic_crop_mode: scale
+ mosaic_center: 0.25
+ aug_scale_min: 0.1
+ aug_scale_max: 1.9
+ validation_data:
+ input_path: /readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco/val*
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ coco91_to_80: false
+ is_training: false
+ global_batch_size: 8
+ dtype: float32
+ drop_remainder: true
+ parser:
+ max_num_instances: 100
+ letter_box: true
+ use_tie_breaker: true
+ best_match_only: true
+ anchor_thresh: 4.0
+ area_thresh: 0.0
+trainer:
+ best_checkpoint_export_subdir: best_ckpt
+ best_checkpoint_eval_metric: AP
+ best_checkpoint_metric_comp: higher
+ train_steps: 1108800
+ validation_steps: 625
+ validation_interval: 18480
+ steps_per_loop: 3696
+ summary_interval: 3696
+ checkpoint_interval: 3696
+ optimizer_config:
+ ema:
+ average_decay: 0.9999
+ trainable_weights_only: false
+ dynamic_decay: true
+ learning_rate:
+ cosine:
+ initial_learning_rate: 0.01
+ alpha: 0.1
+ decay_steps: 1108800
+ type: cosine
+ optimizer:
+ sgd_torch:
+ bias_keys: [bias, beta]
+ weight_keys: [kernel, weight]
+ momentum: 0.937
+ momentum_start: 0.8
+ nesterov: true
+ warmup_steps: 11088
+ weight_decay: 0.0005
+ type: sgd_torch
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0
+ warmup_steps: 11088
+ type: linear
diff --git a/official/vision/beta/projects/yolo/configs/yolo.py b/official/projects/yolo/configs/yolo.py
old mode 100755
new mode 100644
similarity index 86%
rename from official/vision/beta/projects/yolo/configs/yolo.py
rename to official/projects/yolo/configs/yolo.py
index 59dcbb05410..6229061d2da
--- a/official/vision/beta/projects/yolo/configs/yolo.py
+++ b/official/projects/yolo/configs/yolo.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,9 +22,9 @@
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.modeling import hyperparams
-from official.vision.beta.projects.yolo import optimization
-from official.vision.beta.projects.yolo.configs import backbones
-from official.vision.beta.projects.yolo.configs import decoders
+from official.projects.yolo import optimization
+from official.projects.yolo.configs import backbones
+from official.projects.yolo.configs import decoders
from official.vision.configs import common
@@ -76,15 +76,21 @@ class TfExampleDecoderLabelMap(hyperparams.Config):
@dataclasses.dataclass
class DataDecoder(hyperparams.OneOfConfig):
type: Optional[str] = 'simple_decoder'
- simple_decoder: TfExampleDecoder = TfExampleDecoder()
- label_map_decoder: TfExampleDecoderLabelMap = TfExampleDecoderLabelMap()
+ simple_decoder: TfExampleDecoder = dataclasses.field(
+ default_factory=TfExampleDecoder
+ )
+ label_map_decoder: TfExampleDecoderLabelMap = dataclasses.field(
+ default_factory=TfExampleDecoderLabelMap
+ )
@dataclasses.dataclass
class Mosaic(hyperparams.Config):
mosaic_frequency: float = 0.0
+ mosaic9_frequency: float = 0.0
mixup_frequency: float = 0.0
mosaic_center: float = 0.2
+ mosaic9_center: float = 0.33
mosaic_crop_mode: Optional[str] = None
aug_scale_min: float = 1.0
aug_scale_max: float = 1.0
@@ -110,7 +116,7 @@ class Parser(hyperparams.Config):
best_match_only: bool = False
anchor_thresh: float = -0.01
area_thresh: float = 0.1
- mosaic: Mosaic = Mosaic()
+ mosaic: Mosaic = dataclasses.field(default_factory=Mosaic)
@dataclasses.dataclass
@@ -120,11 +126,10 @@ class DataConfig(cfg.DataConfig):
input_path: str = ''
tfds_name: str = ''
tfds_split: str = ''
- global_batch_size: int = 1
is_training: bool = True
dtype: str = 'float16'
- decoder: DataDecoder = DataDecoder()
- parser: Parser = Parser()
+ decoder: DataDecoder = dataclasses.field(default_factory=DataDecoder)
+ parser: Parser = dataclasses.field(default_factory=Parser)
shuffle_buffer_size: int = 10000
tfds_download: bool = True
cache: bool = False
@@ -140,17 +145,21 @@ class YoloHead(hyperparams.Config):
@dataclasses.dataclass
class YoloDetectionGenerator(hyperparams.Config):
+ apply_nms: bool = True
box_type: FPNConfig = dataclasses.field(
default_factory=_build_dict(MIN_LEVEL, MAX_LEVEL, 'original'))
scale_xy: FPNConfig = dataclasses.field(
default_factory=_build_dict(MIN_LEVEL, MAX_LEVEL, 1.0))
path_scales: FPNConfig = dataclasses.field(
default_factory=_build_path_scales(MIN_LEVEL, MAX_LEVEL))
- nms_type: str = 'greedy'
+ # Choose from v1, v2, iou and greedy.
+ nms_version: str = 'greedy'
iou_thresh: float = 0.001
nms_thresh: float = 0.6
max_boxes: int = 200
pre_nms_points: int = 5000
+ # Only works when nms_version='v2'.
+ use_class_agnostic_nms: Optional[bool] = False
@dataclasses.dataclass
@@ -178,7 +187,7 @@ class YoloLoss(hyperparams.Config):
@dataclasses.dataclass
class Box(hyperparams.Config):
- box: List[int] = dataclasses.field(default=list)
+ box: List[int] = dataclasses.field(default_factory=list)
@dataclasses.dataclass
@@ -224,21 +233,32 @@ def set_boxes(self, boxes):
class Yolo(hyperparams.Config):
input_size: Optional[List[int]] = dataclasses.field(
default_factory=lambda: [512, 512, 3])
- backbone: backbones.Backbone = backbones.Backbone(
- type='darknet', darknet=backbones.Darknet(model_id='cspdarknet53'))
- decoder: decoders.Decoder = decoders.Decoder(
- type='yolo_decoder',
- yolo_decoder=decoders.YoloDecoder(version='v4', type='regular'))
- head: YoloHead = YoloHead()
- detection_generator: YoloDetectionGenerator = YoloDetectionGenerator()
- loss: YoloLoss = YoloLoss()
- norm_activation: common.NormActivation = common.NormActivation(
- activation='mish',
- use_sync_bn=True,
- norm_momentum=0.99,
- norm_epsilon=0.001)
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='darknet', darknet=backbones.Darknet(model_id='cspdarknet53')
+ )
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder( # pylint: disable=g-long-lambda
+ type='yolo_decoder',
+ yolo_decoder=decoders.YoloDecoder(version='v4', type='regular'),
+ )
+ )
+ head: YoloHead = dataclasses.field(default_factory=YoloHead)
+ detection_generator: YoloDetectionGenerator = dataclasses.field(
+ default_factory=YoloDetectionGenerator
+ )
+ loss: YoloLoss = dataclasses.field(default_factory=YoloLoss)
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ activation='mish',
+ use_sync_bn=True,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ )
+ )
num_classes: int = 80
- anchor_boxes: AnchorBoxes = AnchorBoxes()
+ anchor_boxes: AnchorBoxes = dataclasses.field(default_factory=AnchorBoxes)
darknet_based_model: bool = False
@@ -246,9 +266,13 @@ class Yolo(hyperparams.Config):
class YoloTask(cfg.TaskConfig):
per_category_metrics: bool = False
smart_bias_lr: float = 0.0
- model: Yolo = Yolo()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
+ model: Yolo = dataclasses.field(default_factory=Yolo)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
weight_decay: float = 0.0
annotation_file: Optional[str] = None
init_checkpoint: Optional[str] = None
@@ -256,6 +280,8 @@ class YoloTask(cfg.TaskConfig):
str, List[str]] = 'all' # all, backbone, and/or decoder
gradient_clip_norm: float = 0.0
seed = GLOBAL_SEED
+ # Sets maximum number of boxes to be evaluated by coco eval api.
+ max_num_eval_detections: int = 100
COCO_INPUT_PATH_BASE = 'coco'
@@ -399,7 +425,7 @@ def yolo_darknet() -> cfg.ExperimentConfig:
def scaled_yolo() -> cfg.ExperimentConfig:
"""COCO object detection with YOLOv4-csp and v4."""
train_batch_size = 256
- eval_batch_size = 8
+ eval_batch_size = 256
train_epochs = 300
warmup_epochs = 3
@@ -413,8 +439,8 @@ def scaled_yolo() -> cfg.ExperimentConfig:
task=YoloTask(
smart_bias_lr=0.1,
init_checkpoint_modules='',
- annotation_file=None,
weight_decay=0.0,
+ annotation_file=None,
model=Yolo(
darknet_based_model=False,
norm_activation=common.NormActivation(
@@ -462,7 +488,7 @@ def scaled_yolo() -> cfg.ExperimentConfig:
input_path=os.path.join(COCO_INPUT_PATH_BASE, 'val*'),
is_training=False,
global_batch_size=eval_batch_size,
- drop_remainder=True,
+ drop_remainder=False,
dtype='float32',
parser=Parser(
letter_box=True,
@@ -474,7 +500,7 @@ def scaled_yolo() -> cfg.ExperimentConfig:
))),
trainer=cfg.TrainerConfig(
train_steps=train_epochs * steps_per_epoch,
- validation_steps=COCO_VAL_EXAMPLES // eval_batch_size,
+ validation_steps=20,
validation_interval=validation_interval * steps_per_epoch,
steps_per_loop=steps_per_epoch,
summary_interval=steps_per_epoch,
diff --git a/official/projects/yolo/configs/yolov7.py b/official/projects/yolo/configs/yolov7.py
new file mode 100644
index 00000000000..786da5a7b35
--- /dev/null
+++ b/official/projects/yolo/configs/yolov7.py
@@ -0,0 +1,395 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""YOLOv7 configuration definition."""
+import dataclasses
+import os
+from typing import List, Optional, Union
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.yolo import optimization
+from official.projects.yolo.configs import backbones
+from official.projects.yolo.configs import decoders
+from official.projects.yolo.configs.yolo import AnchorBoxes
+from official.projects.yolo.configs.yolo import DataConfig
+from official.projects.yolo.configs.yolo import Mosaic
+from official.projects.yolo.configs.yolo import Parser
+from official.projects.yolo.configs.yolo import YoloDetectionGenerator
+from official.vision.configs import common
+
+
+# pytype: disable=annotation-type-mismatch
+
+MIN_LEVEL = 3
+MAX_LEVEL = 5
+GLOBAL_SEED = 1000
+
+
+def _build_dict(min_level, max_level, value):
+ vals = {str(key): value for key in range(min_level, max_level + 1)}
+ vals['all'] = None
+ return lambda: vals
+
+
+def _build_path_scales(min_level, max_level):
+ return lambda: {str(key): 2**key for key in range(min_level, max_level + 1)}
+
+
+# pylint: disable=missing-class-docstring
+@dataclasses.dataclass
+class TfExampleDecoder(hyperparams.Config):
+ regenerate_source_id: bool = False
+ coco91_to_80: bool = True
+
+
+@dataclasses.dataclass
+class TfExampleDecoderLabelMap(hyperparams.Config):
+ regenerate_source_id: bool = False
+ label_map: str = ''
+
+
+@dataclasses.dataclass
+class DataDecoder(hyperparams.OneOfConfig):
+ type: Optional[str] = 'simple_decoder'
+ simple_decoder: TfExampleDecoder = dataclasses.field(
+ default_factory=TfExampleDecoder
+ )
+ label_map_decoder: TfExampleDecoderLabelMap = dataclasses.field(
+ default_factory=TfExampleDecoderLabelMap
+ )
+
+
+@dataclasses.dataclass
+class YoloV7Head(hyperparams.Config):
+ """Parameterization for the YOLO Head."""
+ num_anchors: int = 3
+ use_separable_conv: bool = False
+
+
+@dataclasses.dataclass
+class YoloV7Loss(hyperparams.Config):
+ """Config or YOLOv7 loss."""
+ alpha: float = 0.0
+ gamma: float = 0.0
+ box_weight: float = 0.05
+ obj_weight: float = 0.7
+ cls_weight: float = 0.3
+ label_smoothing: float = 0.0
+ anchor_threshold: float = 4.0
+ iou_mix_ratio: float = 1.0
+ auto_balance: bool = False
+ use_ota: bool = True
+
+
+@dataclasses.dataclass
+class Box(hyperparams.Config):
+ box: List[int] = dataclasses.field(default_factory=list)
+
+
+@dataclasses.dataclass
+class YoloV7(hyperparams.Config):
+ input_size: Optional[List[int]] = dataclasses.field(
+ default_factory=lambda: [640, 640, 3]
+ )
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='yolov7', yolov7=backbones.YoloV7(model_id='yolov7')
+ )
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder( # pylint: disable=g-long-lambda
+ type='yolov7', yolo_decoder=decoders.YoloV7(model_id='yolov7')
+ )
+ )
+ head: YoloV7Head = dataclasses.field(default_factory=YoloV7Head)
+ detection_generator: YoloDetectionGenerator = dataclasses.field(
+ default_factory=lambda: YoloDetectionGenerator( # pylint: disable=g-long-lambda
+ box_type=_build_dict(MIN_LEVEL, MAX_LEVEL, 'scaled')(),
+ scale_xy=_build_dict(MIN_LEVEL, MAX_LEVEL, 2.0)(),
+ path_scales=_build_path_scales(MIN_LEVEL, MAX_LEVEL)(),
+ nms_version='iou',
+ iou_thresh=0.001,
+ nms_thresh=0.7,
+ max_boxes=300,
+ pre_nms_points=5000,
+ )
+ )
+ loss: YoloV7Loss = dataclasses.field(default_factory=YoloV7Loss)
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ activation='swish',
+ use_sync_bn=True,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ )
+ )
+ num_classes: int = 80
+ min_level: int = 3
+ max_level: int = 5
+ anchor_boxes: AnchorBoxes = dataclasses.field(default_factory=AnchorBoxes)
+
+
+@dataclasses.dataclass
+class YoloV7Task(cfg.TaskConfig):
+ per_category_metrics: bool = False
+ smart_bias_lr: float = 0.0
+ model: YoloV7 = dataclasses.field(default_factory=YoloV7)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ weight_decay: float = 0.0
+ annotation_file: Optional[str] = None
+ init_checkpoint: Optional[str] = None
+ init_checkpoint_modules: Union[str, List[str]] = (
+ 'all' # all, backbone, and/or decoder
+ )
+ gradient_clip_norm: float = 0.0
+ seed = GLOBAL_SEED
+ # Sets maximum number of boxes to be evaluated by coco eval api.
+ max_num_eval_detections: int = 100
+
+
+COCO_INPUT_PATH_BASE = (
+ '/readahead/200M/placer/prod/home/tensorflow-performance-data/datasets/coco'
+)
+COCO_TRAIN_EXAMPLES = 118287
+COCO_VAL_EXAMPLES = 5000
+
+
+@exp_factory.register_config_factory('yolov7')
+def yolov7() -> cfg.ExperimentConfig:
+ """YOLOv7 general config."""
+ return cfg.ExperimentConfig(
+ task=YoloV7Task(),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ],
+ )
+
+
+@exp_factory.register_config_factory('coco_yolov7')
+def coco_yolov7() -> cfg.ExperimentConfig:
+ """COCO object detection with YOLOv7."""
+ train_batch_size = 256
+ eval_batch_size = 256
+ train_epochs = 300
+ steps_per_epoch = COCO_TRAIN_EXAMPLES // train_batch_size
+ validation_interval = 5
+ warmup_steps = 3 * steps_per_epoch
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(mixed_precision_dtype='float32'),
+ task=YoloV7Task(
+ init_checkpoint='',
+ init_checkpoint_modules='backbone',
+ annotation_file=None,
+ weight_decay=0.0,
+ model=YoloV7(
+ norm_activation=common.NormActivation(
+ activation='swish',
+ norm_momentum=0.03,
+ norm_epsilon=0.001,
+ use_sync_bn=True),
+ head=YoloV7Head(),
+ loss=YoloV7Loss(),
+ anchor_boxes=AnchorBoxes(
+ anchors_per_scale=3,
+ boxes=[
+ Box(box=[12, 16]),
+ Box(box=[19, 36]),
+ Box(box=[40, 28]),
+ Box(box=[36, 75]),
+ Box(box=[76, 55]),
+ Box(box=[72, 146]),
+ Box(box=[142, 110]),
+ Box(box=[192, 243]),
+ Box(box=[459, 401]),
+ ],
+ ),
+ ),
+ train_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ dtype='float32',
+ parser=Parser(
+ max_num_instances=300,
+ letter_box=True,
+ random_flip=True,
+ random_pad=False,
+ jitter=0.0,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ aug_rand_translate=0.2,
+ aug_rand_saturation=0.7,
+ aug_rand_brightness=0.4,
+ aug_rand_hue=0.015,
+ aug_rand_angle=0.0,
+ aug_rand_perspective=0.0,
+ use_tie_breaker=True,
+ best_match_only=True,
+ anchor_thresh=4.0,
+ area_thresh=0.0,
+ mosaic=Mosaic(
+ mosaic_frequency=1.0,
+ mosaic9_frequency=0.2,
+ mixup_frequency=0.15,
+ mosaic_crop_mode='scale',
+ mosaic_center=0.25,
+ mosaic9_center=0.33,
+ aug_scale_min=0.1,
+ aug_scale_max=1.9,
+ ),
+ ),
+ ),
+ validation_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=True,
+ dtype='float32',
+ parser=Parser(
+ max_num_instances=300,
+ letter_box=True,
+ use_tie_breaker=True,
+ best_match_only=True,
+ anchor_thresh=4.0,
+ area_thresh=0.0,
+ ),
+ ),
+ smart_bias_lr=0.1,
+ ),
+ trainer=cfg.TrainerConfig(
+ best_checkpoint_export_subdir='best_ckpt',
+ best_checkpoint_eval_metric='AP',
+ best_checkpoint_metric_comp='higher',
+ train_steps=train_epochs * steps_per_epoch,
+ validation_steps=COCO_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=validation_interval * steps_per_epoch,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'ema': {
+ 'average_decay': 0.9999,
+ 'trainable_weights_only': False,
+ 'dynamic_decay': True,
+ },
+ 'optimizer': {
+ 'type': 'sgd_torch',
+ 'sgd_torch': {
+ 'momentum': 0.937,
+ 'momentum_start': 0.8,
+ 'nesterov': True,
+ 'warmup_steps': warmup_steps,
+ # Scale up the weight decay by batch size.
+ 'weight_decay': 0.0005 * train_batch_size / 64,
+ },
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.01,
+ 'alpha': 0.1,
+ 'decay_steps': train_epochs * steps_per_epoch,
+ },
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': warmup_steps,
+ 'warmup_learning_rate': 0.0,
+ },
+ },
+ }),
+ ),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ],
+ )
+
+ return config
+
+
+@exp_factory.register_config_factory('coco_yolov7tiny')
+def coco_yolov7_tiny() -> cfg.ExperimentConfig:
+ """COCO object detection with YOLOv7-tiny."""
+ config = coco_yolov7()
+ config.task.model.input_size = [416, 416, 3]
+ config.task.model.backbone.yolov7.model_id = 'yolov7-tiny'
+ config.task.model.decoder.yolov7.model_id = 'yolov7-tiny'
+ config.task.model.norm_activation.activation = 'leaky'
+ config.task.model.anchor_boxes.boxes = [
+ Box(box=[10, 13]),
+ Box(box=[16, 30]),
+ Box(box=[33, 23]),
+ Box(box=[30, 61]),
+ Box(box=[62, 45]),
+ Box(box=[59, 119]),
+ Box(box=[116, 90]),
+ Box(box=[156, 198]),
+ Box(box=[373, 326]),
+ ]
+
+ config.task.model.loss.cls_weight = 0.5
+ config.task.model.loss.obj_weight = 1.0
+ config.task.train_data.parser.aug_rand_translate = 0.1
+ config.task.train_data.parser.mosaic.mixup_frequency = 0.05
+ config.task.train_data.parser.mosaic.aug_scale_min = 0.5
+ config.task.train_data.parser.mosaic.aug_scale_max = 1.5
+ config.trainer.optimizer_config.learning_rate.cosine.alpha = 0.01
+ return config
+
+
+@exp_factory.register_config_factory('coco91_yolov7tiny')
+def coco91_yolov7_tiny() -> cfg.ExperimentConfig:
+ """COCO object detection with YOLOv7-tiny using 91 classes."""
+ config = coco_yolov7_tiny()
+ config.task.model.num_classes = 91
+ config.task.model.decoder.yolov7.use_separable_conv = True
+ config.task.model.head.use_separable_conv = True
+ config.task.train_data.coco91_to_80 = False
+ config.task.validation_data.coco91_to_80 = False
+ return config
+
+
+@exp_factory.register_config_factory('coco_yolov7x')
+def coco_yolov7x() -> cfg.ExperimentConfig:
+ config = coco_yolov7()
+ config.task.model.backbone.yolov7.model_id = 'yolov7x'
+ config.task.model.decoder.yolov7.model_id = 'yolov7x'
+ return config
+
+
+@exp_factory.register_config_factory('coco_yolov7_nano')
+def coco_yolov7_nano() -> cfg.ExperimentConfig:
+ config = coco_yolov7()
+ config.task.model.backbone.yolov7.model_id = 'yolov7-nano'
+ config.task.model.decoder.yolov7.model_id = 'yolov7-nano'
+ return config
+
+
+@exp_factory.register_config_factory('coco_yolov7_pico')
+def coco_yolov7_pico() -> cfg.ExperimentConfig:
+ config = coco_yolov7()
+ config.task.model.backbone.yolov7.model_id = 'yolov7-pico'
+ config.task.model.decoder.yolov7.model_id = 'yolov7-pico'
+ return config
diff --git a/official/projects/yolo/darknet_image_calssification.ipynb b/official/projects/yolo/darknet_image_calssification.ipynb
new file mode 100644
index 00000000000..4572b9e2da5
--- /dev/null
+++ b/official/projects/yolo/darknet_image_calssification.ipynb
@@ -0,0 +1,1376 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "rG2rwLmwh4Rc"
+ },
+ "source": [
+ "#### Copyright 2023 The TensorFlow Authors."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "yOzhgsd6hz6X"
+ },
+ "outputs": [],
+ "source": [
+ "#@title Licensed under the Apache License, Version 2.0 (the \"License\");\n",
+ "# you may not use this file except in compliance with the License.\n",
+ "# You may obtain a copy of the License at\n",
+ "#\n",
+ "# https://www.apache.org/licenses/LICENSE-2.0\n",
+ "#\n",
+ "# Unless required by applicable law or agreed to in writing, software\n",
+ "# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
+ "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
+ "# See the License for the specific language governing permissions and\n",
+ "# limitations under the License."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "VfjRf_jVg43D"
+ },
+ "source": [
+ "# YOLO Image classification\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "keCteDOniuMA"
+ },
+ "source": [
+ "This tutorial trains [Darkent](https://arxiv.org/abs/1911.11929) from Tensorflow Model Garden package ([tf-models-official](https://pypi.org/project/tf-models-official/)) to classify images in the [cats_vs_dogs](https://www.tensorflow.org/datasets/catalog/cats_vs_dogs) dataset.\n",
+ "\n",
+ "\n",
+ "[Model Garden](https://github.com/tensorflow/models/tree/master/official) contains a collection of state-of-the-art vision models, implemented with TensorFlow's high-level APIs. The implementations demonstrate the best practices for modeling, letting users to take full advantage of TensorFlow for their research and product development.\n",
+ "\n",
+ "**Dataset:** cats_vs_dogs\n",
+ "* A large set of images of cats and dogs.\n",
+ "\n",
+ "This tutorial demonstrates how to:\n",
+ "* Use models from the [TensorFlow Models package](https://pypi.org/project/tf-models-official/)\n",
+ "* Train/Fine-tune a pre-built Darkent variations for Image Classification\n",
+ "* Export the trained/tuned darknet model"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "41irlZ-FlMhi"
+ },
+ "source": [
+ "## Clone the model-garden repository"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "snEeWhVIdWgi"
+ },
+ "outputs": [],
+ "source": [
+ "! git clone -q https://github.com/tensorflow/models.git"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "DVSglGuXmrjN"
+ },
+ "outputs": [],
+ "source": [
+ "! pip install -q -U tensorflow_datasets\n",
+ "! pip install -q --user -r models/official/requirements.txt"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "y_THKv-WwjmF"
+ },
+ "source": [
+ "**Note:** Please restart runtime and continue with running the notebook"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XVQ1By43lc_k"
+ },
+ "outputs": [],
+ "source": [
+ "import os\n",
+ "import sys\n",
+ "\n",
+ "import os\n",
+ "os.environ['PYTHONPATH'] += \":/content/models\"\n",
+ "\n",
+ "sys.path.append(\"/content/models\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "zEwBh_eplVZg"
+ },
+ "source": [
+ "## Import necessary libraries"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "FTOX9q8tZE9R"
+ },
+ "outputs": [],
+ "source": [
+ "import pprint\n",
+ "import logging\n",
+ "import matplotlib.pyplot as plt\n",
+ "import tensorflow as tf\n",
+ "import tensorflow_datasets as tfds\n",
+ "\n",
+ "from official import core\n",
+ "from official.vision.data import tfrecord_lib\n",
+ "from official.vision import configs\n",
+ "from official.vision.configs import common\n",
+ "from official.projects.yolo.common import registry_imports\n",
+ "from official.projects.yolo.serving import export_saved_model\n",
+ "from official.projects.yolo.serving import export_module_factory\n",
+ "from official.vision.serving import export_saved_model_lib\n",
+ "\n",
+ "logging.disable(logging.WARNING)\n",
+ "pp = pprint.PrettyPrinter(indent=4)\n",
+ "%matplotlib inline"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Obap4N9HlZzf"
+ },
+ "source": [
+ "## Load dataset from Tensorflow Datasets(tfds)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 1551,
+ "status": "ok",
+ "timestamp": 1699997729595,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "_k5x1Tt8G9VV",
+ "outputId": "e78c2377-65ec-4524-8c8f-25490a705256"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/plain": [
+ "tfds.core.DatasetInfo(\n",
+ " name='cats_vs_dogs',\n",
+ " full_name='cats_vs_dogs/4.0.0',\n",
+ " description=\"\"\"\n",
+ " A large set of images of cats and dogs. There are 1738 corrupted images that are dropped.\n",
+ " \"\"\",\n",
+ " homepage='https://www.microsoft.com/en-us/download/details.aspx?id=54765',\n",
+ " data_dir='/root/tensorflow_datasets/cats_vs_dogs/4.0.0',\n",
+ " file_format=tfrecord,\n",
+ " download_size=786.67 MiB,\n",
+ " dataset_size=689.64 MiB,\n",
+ " features=FeaturesDict({\n",
+ " 'image': Image(shape=(None, None, 3), dtype=uint8),\n",
+ " 'image/filename': Text(shape=(), dtype=string),\n",
+ " 'label': ClassLabel(shape=(), dtype=int64, num_classes=2),\n",
+ " }),\n",
+ " supervised_keys=('image', 'label'),\n",
+ " disable_shuffling=False,\n",
+ " splits={\n",
+ " 'train': \u003cSplitInfo num_examples=23262, num_shards=8\u003e,\n",
+ " },\n",
+ " citation=\"\"\"@Inproceedings (Conference){asirra-a-captcha-that-exploits-interest-aligned-manual-image-categorization,\n",
+ " author = {Elson, Jeremy and Douceur, John (JD) and Howell, Jon and Saul, Jared},\n",
+ " title = {Asirra: A CAPTCHA that Exploits Interest-Aligned Manual Image Categorization},\n",
+ " booktitle = {Proceedings of 14th ACM Conference on Computer and Communications Security (CCS)},\n",
+ " year = {2007},\n",
+ " month = {October},\n",
+ " publisher = {Association for Computing Machinery, Inc.},\n",
+ " url = {https://www.microsoft.com/en-us/research/publication/asirra-a-captcha-that-exploits-interest-aligned-manual-image-categorization/},\n",
+ " edition = {Proceedings of 14th ACM Conference on Computer and Communications Security (CCS)},\n",
+ " }\"\"\",\n",
+ ")"
+ ]
+ },
+ "execution_count": 3,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "(train_ds, validation_ds, test_ds), ds_info = tfds.load(\n",
+ " name='cats_vs_dogs',\n",
+ " split=['train[:70%]', 'train[70%:90%]', 'train[90%:100%]'],\n",
+ " with_info=True)\n",
+ "label_info = ds_info.features['label']\n",
+ "ds_info"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "el_oHjBAlreD"
+ },
+ "source": [
+ "## Write data to TFrecords"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "qMp811pgl1_0"
+ },
+ "source": [
+ "### Helper functions to preproces the data"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "JDIRTV1Bh4FA"
+ },
+ "outputs": [],
+ "source": [
+ "def process_record(record):\n",
+ " \"\"\"\n",
+ " Process a single record for TFRecords.\n",
+ "\n",
+ " This function takes a record, typically containing image and label data,\n",
+ " and converts it into a TFRecord example. Detailed explaination is available here\n",
+ " https://www.tensorflow.org/api_docs/python/tf/train/Example\n",
+ "\n",
+ " Args:\n",
+ " record (dict): A dictionary containing the record data with the following keys:\n",
+ " - 'image': A tensor representing the image data.\n",
+ " - 'label': A tensor representing the label associated with the image.\n",
+ "\n",
+ " Returns:\n",
+ " tf.train.Example: A TFRecord example containing the processed data with\n",
+ " the following features:\n",
+ " - 'image/encoded': The encoded image data as a feature.\n",
+ " - 'image/class/label': The label data as a feature.\n",
+ " \"\"\"\n",
+ " keys_to_features = {\n",
+ " 'image/encoded': tfrecord_lib.convert_to_feature(\n",
+ " tf.io.encode_jpeg(record['image']).numpy()),\n",
+ " 'image/class/label': tfrecord_lib.convert_to_feature(\n",
+ " record['label'].numpy())\n",
+ " }\n",
+ " example = tf.train.Example(features=tf.train.Features(feature=keys_to_features))\n",
+ " return example\n"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "pUQmR7seiCBF"
+ },
+ "outputs": [],
+ "source": [
+ "def write_tfrecords(dataset, output_path, num_shards=1):\n",
+ " \"\"\"\n",
+ " Write a dataset to TFRecords files.\n",
+ "\n",
+ " This function takes a dataset and writes it to one or more TFRecords files,\n",
+ " splitting the data into shards if specified.\n",
+ "\n",
+ " Args:\n",
+ " dataset (iterable): An iterable containing the data records to be written\n",
+ " to TFRecords. Each record should be in a format suitable for processing\n",
+ " with the 'process_record' function.\n",
+ " output_path (str): The base path where the TFRecords files will be saved.\n",
+ " If 'num_shards' is greater than 1, a unique suffix for each shard will\n",
+ " be added to the base path.\n",
+ " num_shards (int, optional): The number of TFRecords files to split the data\n",
+ " into. Defaults to 1, indicating no sharding.\n",
+ "\n",
+ " Reuturns:\n",
+ " None\n",
+ " \"\"\"\n",
+ " writers = [\n",
+ " tf.io.TFRecordWriter(\n",
+ " output_path + '-%05d-of-%05d.tfrecord' % (i, num_shards))\n",
+ " for i in range(num_shards)\n",
+ " ]\n",
+ " for idx, record in enumerate(dataset):\n",
+ " if idx % LOG_EVERY == 0:\n",
+ " print('On image %d' % idx)\n",
+ " tf_example = process_record(record)\n",
+ " writers[idx % num_shards].write(tf_example.SerializeToString())\n",
+ "\n",
+ "LOG_EVERY = 1000\n",
+ "output_dir = './cat_vs_dogs_tfrecords/'\n",
+ "if not os.path.exists(output_dir):\n",
+ " os.mkdir(output_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "K35c-4M1ojdC"
+ },
+ "source": [
+ "### Writing training data to TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 23362,
+ "status": "ok",
+ "timestamp": 1699997756512,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "xiya20KViEBx",
+ "outputId": "ab75364e-be3f-46f8-d0b1-e3186b0906d1"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "On image 0\n",
+ "On image 1000\n",
+ "On image 2000\n",
+ "On image 3000\n",
+ "On image 4000\n",
+ "On image 5000\n",
+ "On image 6000\n",
+ "On image 7000\n",
+ "On image 8000\n",
+ "On image 9000\n",
+ "On image 10000\n",
+ "On image 11000\n",
+ "On image 12000\n",
+ "On image 13000\n",
+ "On image 14000\n",
+ "On image 15000\n",
+ "On image 16000\n"
+ ]
+ }
+ ],
+ "source": [
+ "output_train_tfrecs = output_dir + 'train'\n",
+ "write_tfrecords(train_ds, output_train_tfrecs,\n",
+ " num_shards=int(train_ds.cardinality().numpy() * 0.1))"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "oAOuGouiolhq"
+ },
+ "source": [
+ "### Writing validation data to TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 14426,
+ "status": "ok",
+ "timestamp": 1699997770935,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "BdU2-OlLiHNh",
+ "outputId": "fcffd93d-f24f-407b-d1c6-26880b1ac690"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "On image 0\n",
+ "On image 1000\n",
+ "On image 2000\n",
+ "On image 3000\n",
+ "On image 4000\n"
+ ]
+ }
+ ],
+ "source": [
+ "output_validation_tfrecs = output_dir + 'validation'\n",
+ "write_tfrecords(validation_ds, output_validation_tfrecs,\n",
+ " num_shards=int(validation_ds.cardinality().numpy() *0.1))"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "iAU3OLKjosAd"
+ },
+ "source": [
+ "### Writing testing data to TFRecords"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 32,
+ "status": "ok",
+ "timestamp": 1699997770937,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "8bcmO71biJR4",
+ "outputId": "8c8bf0d9-8661-42ed-9e70-8ba30eb1a88b"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "On image 0\n",
+ "On image 1000\n",
+ "On image 2000\n"
+ ]
+ }
+ ],
+ "source": [
+ "output_test_tfrecs = output_dir + 'test'\n",
+ "write_tfrecords(test_ds, output_test_tfrecs,\n",
+ " num_shards=int(test_ds.cardinality().numpy() *0.1))"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Lc6Wf_oMozOJ"
+ },
+ "source": [
+ "## Experiment Configuration"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "FL5aunZjpIl_"
+ },
+ "source": [
+ "### Load the existing configuration"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "tbqB6L5wldIE"
+ },
+ "outputs": [],
+ "source": [
+ "exp_config = core.exp_factory.get_exp_config('darknet_classification')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "09iJzu27pNaf"
+ },
+ "source": [
+ "### Change the configuration parameters for custom dataset"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "38duTtEKiSdr"
+ },
+ "outputs": [],
+ "source": [
+ "BATCH_SIZE = 16\n",
+ "IMG_SIZE = 224\n",
+ "\n",
+ "epochs = 10\n",
+ "steps_per_epoch = int(train_ds.cardinality().numpy() / BATCH_SIZE)\n",
+ "validation_steps = int(validation_ds.cardinality().numpy() / BATCH_SIZE)\n",
+ "num_steps = epochs * steps_per_epoch\n",
+ "\n",
+ "lr = 0.012\n",
+ "warmpup_lr = 0.1 * lr\n",
+ "\n",
+ "exp_config.task.model.input_size = [IMG_SIZE, IMG_SIZE, 3]\n",
+ "exp_config.task.model.num_classes = ds_info.features['label'].num_classes\n",
+ "\n",
+ "exp_config.task.train_data.input_path = f'{output_train_tfrecs}*'\n",
+ "exp_config.task.train_data.global_batch_size = BATCH_SIZE\n",
+ "\n",
+ "exp_config.task.validation_data.input_path = f'{output_validation_tfrecs}*'\n",
+ "exp_config.task.validation_data.global_batch_size = BATCH_SIZE\n",
+ "\n",
+ "exp_config.trainer.checkpoint_interval = steps_per_epoch\n",
+ "exp_config.trainer.best_checkpoint_export_subdir = 'best_ckpt'\n",
+ "exp_config.trainer.optimizer_config.optimizer.type = 'sgd'\n",
+ "exp_config.trainer.optimizer_config.optimizer.sgd.momentum = 0.9\n",
+ "exp_config.trainer.optimizer_config.learning_rate.type = 'cosine'\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.decay_steps = num_steps\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.initial_learning_rate = lr\n",
+ "exp_config.trainer.optimizer_config.warmup.type = 'linear'\n",
+ "exp_config.trainer.optimizer_config.warmup.linear.warmup_learning_rate = warmpup_lr\n",
+ "exp_config.trainer.optimizer_config.warmup.linear.warmup_steps = int(0.1 * steps_per_epoch)\n",
+ "\n",
+ "exp_config.trainer.train_steps = num_steps\n",
+ "exp_config.trainer.steps_per_loop = steps_per_epoch\n",
+ "exp_config.trainer.validation_steps = validation_steps\n",
+ "exp_config.trainer.validation_interval = steps_per_epoch\n",
+ "exp_config.trainer.summary_interval = steps_per_epoch"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "FINslYHhqaPt"
+ },
+ "source": [
+ "### Set up the distribution strategy"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 30,
+ "status": "ok",
+ "timestamp": 1699997770939,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "3-e4HjyKqZfw",
+ "outputId": "5977b054-ee4d-4e8e-ae18-55d903a5547f"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Running on single GPU /device:GPU:0\n",
+ "Number of accelerators: 1\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Detect hardware\n",
+ "try:\n",
+ " tpu_resolver = tf.distribute.cluster_resolver.TPUClusterResolver() # TPU detection\n",
+ "except ValueError:\n",
+ " tpu_resolver = None\n",
+ " gpus = tf.config.experimental.list_logical_devices(\"GPU\")\n",
+ "\n",
+ "# Select appropriate distribution strategy\n",
+ "if tpu_resolver:\n",
+ " tf.config.experimental_connect_to_cluster(tpu_resolver)\n",
+ " tf.tpu.experimental.initialize_tpu_system(tpu_resolver)\n",
+ " distribution_strategy = tf.distribute.experimental.TPUStrategy(tpu_resolver)\n",
+ " print('Running on TPU ', tpu_resolver.cluster_spec().as_dict()['worker'])\n",
+ "elif len(gpus) \u003e 1:\n",
+ " distribution_strategy = tf.distribute.MirroredStrategy([gpu.name for gpu in gpus])\n",
+ " print('Running on multiple GPUs ', [gpu.name for gpu in gpus])\n",
+ "elif len(gpus) == 1:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on single GPU ', gpus[0].name)\n",
+ "else:\n",
+ " distribution_strategy = tf.distribute.get_strategy() # default strategy that works on CPU and single GPU\n",
+ " print('Running on CPU')\n",
+ "\n",
+ "print(\"Number of accelerators: \", distribution_strategy.num_replicas_in_sync)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "ugiuA0-0qq4Z"
+ },
+ "source": [
+ "### Check the new configuration"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 29,
+ "status": "ok",
+ "timestamp": 1699997770940,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "a5CO8qTqHqGH",
+ "outputId": "d9d97856-051f-457b-a9d7-bd0cb5135042"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "{'runtime': {'all_reduce_alg': None,\n",
+ " 'batchnorm_spatial_persistent': False,\n",
+ " 'dataset_num_private_threads': None,\n",
+ " 'default_shard_dim': -1,\n",
+ " 'distribution_strategy': 'mirrored',\n",
+ " 'enable_xla': False,\n",
+ " 'gpu_thread_mode': None,\n",
+ " 'loss_scale': None,\n",
+ " 'mixed_precision_dtype': None,\n",
+ " 'num_cores_per_replica': 1,\n",
+ " 'num_gpus': 0,\n",
+ " 'num_packs': 1,\n",
+ " 'per_gpu_thread_count': 0,\n",
+ " 'run_eagerly': False,\n",
+ " 'task_index': -1,\n",
+ " 'tpu': None,\n",
+ " 'tpu_enable_xla_dynamic_padder': None,\n",
+ " 'use_tpu_mp_strategy': False,\n",
+ " 'worker_hosts': None},\n",
+ " 'task': {'allow_image_summary': False,\n",
+ " 'differential_privacy_config': None,\n",
+ " 'evaluation': {'precision_and_recall_thresholds': None,\n",
+ " 'report_per_class_precision_and_recall': False,\n",
+ " 'top_k': 5},\n",
+ " 'freeze_backbone': False,\n",
+ " 'gradient_clip_norm': 0.0,\n",
+ " 'init_checkpoint': '',\n",
+ " 'logging_dir': None,\n",
+ " 'losses': {'l2_weight_decay': 0.0,\n",
+ " 'label_smoothing': 0.0,\n",
+ " 'loss_weight': 1.0,\n",
+ " 'one_hot': True,\n",
+ " 'soft_labels': False,\n",
+ " 'use_binary_cross_entropy': False},\n",
+ " 'model': {'add_head_batch_norm': False,\n",
+ " 'backbone': {'darknet': {'depth_scale': 1.0,\n",
+ " 'dilate': False,\n",
+ " 'max_level': 5,\n",
+ " 'min_level': 3,\n",
+ " 'model_id': 'cspdarknet53',\n",
+ " 'use_reorg_input': False,\n",
+ " 'use_separable_conv': False,\n",
+ " 'width_scale': 1.0},\n",
+ " 'type': 'darknet'},\n",
+ " 'dropout_rate': 0.0,\n",
+ " 'input_size': [224, 224, 3],\n",
+ " 'kernel_initializer': 'VarianceScaling',\n",
+ " 'norm_activation': {'activation': 'relu',\n",
+ " 'norm_epsilon': 0.001,\n",
+ " 'norm_momentum': 0.99,\n",
+ " 'use_sync_bn': True},\n",
+ " 'num_classes': 2},\n",
+ " 'name': None,\n",
+ " 'train_data': {'apply_tf_data_service_before_batching': False,\n",
+ " 'aug_crop': True,\n",
+ " 'aug_policy': None,\n",
+ " 'aug_rand_hflip': True,\n",
+ " 'aug_type': None,\n",
+ " 'autotune_algorithm': None,\n",
+ " 'block_length': 1,\n",
+ " 'cache': False,\n",
+ " 'center_crop_fraction': 0.875,\n",
+ " 'color_jitter': 0.0,\n",
+ " 'crop_area_range': (0.08, 1.0),\n",
+ " 'cycle_length': 10,\n",
+ " 'decode_jpeg_only': True,\n",
+ " 'decoder': {'simple_decoder': {'attribute_names': [],\n",
+ " 'mask_binarize_threshold': None,\n",
+ " 'regenerate_source_id': False},\n",
+ " 'type': 'simple_decoder'},\n",
+ " 'deterministic': None,\n",
+ " 'drop_remainder': True,\n",
+ " 'dtype': 'float32',\n",
+ " 'enable_shared_tf_data_service_between_parallel_trainers': False,\n",
+ " 'enable_tf_data_service': False,\n",
+ " 'file_type': 'tfrecord',\n",
+ " 'global_batch_size': 16,\n",
+ " 'image_field_key': 'image/encoded',\n",
+ " 'input_path': './cat_vs_dogs_tfrecords/train*',\n",
+ " 'is_multilabel': False,\n",
+ " 'is_training': True,\n",
+ " 'label_field_key': 'image/class/label',\n",
+ " 'mixup_and_cutmix': None,\n",
+ " 'prefetch_buffer_size': None,\n",
+ " 'randaug_magnitude': 10,\n",
+ " 'random_erasing': None,\n",
+ " 'repeated_augment': None,\n",
+ " 'seed': None,\n",
+ " 'sharding': True,\n",
+ " 'shuffle_buffer_size': 10000,\n",
+ " 'tf_data_service_address': None,\n",
+ " 'tf_data_service_job_name': None,\n",
+ " 'tf_resize_method': 'bilinear',\n",
+ " 'tfds_as_supervised': False,\n",
+ " 'tfds_data_dir': '',\n",
+ " 'tfds_name': '',\n",
+ " 'tfds_skip_decoding_feature': '',\n",
+ " 'tfds_split': '',\n",
+ " 'three_augment': False,\n",
+ " 'trainer_id': None,\n",
+ " 'weights': None},\n",
+ " 'validation_data': {'apply_tf_data_service_before_batching': False,\n",
+ " 'aug_crop': True,\n",
+ " 'aug_policy': None,\n",
+ " 'aug_rand_hflip': True,\n",
+ " 'aug_type': None,\n",
+ " 'autotune_algorithm': None,\n",
+ " 'block_length': 1,\n",
+ " 'cache': False,\n",
+ " 'center_crop_fraction': 0.875,\n",
+ " 'color_jitter': 0.0,\n",
+ " 'crop_area_range': (0.08, 1.0),\n",
+ " 'cycle_length': 10,\n",
+ " 'decode_jpeg_only': True,\n",
+ " 'decoder': {'simple_decoder': {'attribute_names': [],\n",
+ " 'mask_binarize_threshold': None,\n",
+ " 'regenerate_source_id': False},\n",
+ " 'type': 'simple_decoder'},\n",
+ " 'deterministic': None,\n",
+ " 'drop_remainder': True,\n",
+ " 'dtype': 'float32',\n",
+ " 'enable_shared_tf_data_service_between_parallel_trainers': False,\n",
+ " 'enable_tf_data_service': False,\n",
+ " 'file_type': 'tfrecord',\n",
+ " 'global_batch_size': 16,\n",
+ " 'image_field_key': 'image/encoded',\n",
+ " 'input_path': './cat_vs_dogs_tfrecords/validation*',\n",
+ " 'is_multilabel': False,\n",
+ " 'is_training': False,\n",
+ " 'label_field_key': 'image/class/label',\n",
+ " 'mixup_and_cutmix': None,\n",
+ " 'prefetch_buffer_size': None,\n",
+ " 'randaug_magnitude': 10,\n",
+ " 'random_erasing': None,\n",
+ " 'repeated_augment': None,\n",
+ " 'seed': None,\n",
+ " 'sharding': True,\n",
+ " 'shuffle_buffer_size': 10000,\n",
+ " 'tf_data_service_address': None,\n",
+ " 'tf_data_service_job_name': None,\n",
+ " 'tf_resize_method': 'bilinear',\n",
+ " 'tfds_as_supervised': False,\n",
+ " 'tfds_data_dir': '',\n",
+ " 'tfds_name': '',\n",
+ " 'tfds_skip_decoding_feature': '',\n",
+ " 'tfds_split': '',\n",
+ " 'three_augment': False,\n",
+ " 'trainer_id': None,\n",
+ " 'weights': None}},\n",
+ " 'trainer': {'allow_tpu_summary': False,\n",
+ " 'best_checkpoint_eval_metric': '',\n",
+ " 'best_checkpoint_export_subdir': 'best_ckpt',\n",
+ " 'best_checkpoint_metric_comp': 'higher',\n",
+ " 'checkpoint_interval': 1017,\n",
+ " 'continuous_eval_timeout': 3600,\n",
+ " 'eval_tf_function': True,\n",
+ " 'eval_tf_while_loop': False,\n",
+ " 'loss_upper_bound': 1000000.0,\n",
+ " 'max_to_keep': 5,\n",
+ " 'optimizer_config': {'ema': None,\n",
+ " 'learning_rate': {'cosine': {'alpha': 0.0,\n",
+ " 'decay_steps': 10170,\n",
+ " 'initial_learning_rate': 0.012,\n",
+ " 'name': 'CosineDecay',\n",
+ " 'offset': 0},\n",
+ " 'type': 'cosine'},\n",
+ " 'optimizer': {'sgd': {'clipnorm': None,\n",
+ " 'clipvalue': None,\n",
+ " 'decay': 0.0,\n",
+ " 'global_clipnorm': None,\n",
+ " 'momentum': 0.9,\n",
+ " 'name': 'SGD',\n",
+ " 'nesterov': False},\n",
+ " 'type': 'sgd'},\n",
+ " 'warmup': {'linear': {'name': 'linear',\n",
+ " 'warmup_learning_rate': 0.0012000000000000001,\n",
+ " 'warmup_steps': 101},\n",
+ " 'type': 'linear'}},\n",
+ " 'preemption_on_demand_checkpoint': True,\n",
+ " 'recovery_begin_steps': 0,\n",
+ " 'recovery_max_trials': 0,\n",
+ " 'steps_per_loop': 1017,\n",
+ " 'summary_interval': 1017,\n",
+ " 'train_steps': 10170,\n",
+ " 'train_tf_function': True,\n",
+ " 'train_tf_while_loop': True,\n",
+ " 'validation_interval': 1017,\n",
+ " 'validation_steps': 290,\n",
+ " 'validation_summary_subdir': 'validation'}}\n"
+ ]
+ }
+ ],
+ "source": [
+ "pprint.pprint(exp_config.as_dict())"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "s21NhXxwqzcD"
+ },
+ "source": [
+ "## Create the Task object `(tfm.core.base_task.Task)` from the `config_definitions.TaskConfig`.\n",
+ "\n",
+ "The Task object has all the methods necessary for **building the dataset, building the model, and running training \u0026 evaluation**. These methods are driven by `tfm.core.train_lib.run_experiment`."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "n6TYLP2mIh7H"
+ },
+ "outputs": [],
+ "source": [
+ "model_dir = './trained_model/'\n",
+ "with distribution_strategy.scope():\n",
+ " task = core.task_factory.get_task(exp_config.task, logging_dir=model_dir)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 29,
+ "status": "ok",
+ "timestamp": 1699997770943,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "cdC3Plx7rMTi",
+ "outputId": "bf02a8e2-3446-4c09-e15d-373fa928bc6b"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "images.shape: (16, 224, 224, 3) images.dtype: tf.float32\n",
+ "labels.shape: (16,) labels.dtype: tf.int32\n"
+ ]
+ }
+ ],
+ "source": [
+ "for images, labels in task.build_inputs(exp_config.task.train_data).take(1):\n",
+ " print(f'images.shape: {str(images.shape):16} images.dtype: {images.dtype!r}')\n",
+ " print(f'labels.shape: {str(labels.shape):16} labels.dtype: {labels.dtype!r}')"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "0-KFVW4KA_vU"
+ },
+ "source": [
+ "## Save the configuration"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "onc4YJbBBE3e"
+ },
+ "outputs": [],
+ "source": [
+ "core.train_utils.serialize_config(exp_config, model_dir)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "KasKyG_hrYac"
+ },
+ "source": [
+ "## Visualize the training data"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "emn7QfE1rfQc"
+ },
+ "source": [
+ "### Use `ds_info` (which is an instance of tfds.core.DatasetInfo) to lookup the text descriptions of each class ID."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "MmZg2Z3drToR"
+ },
+ "outputs": [],
+ "source": [
+ "label_info = ds_info.features['label']"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "S3_xSuXPrnX_"
+ },
+ "source": [
+ "### Visualize a batch of the data."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "CQBTmQpzrrGb"
+ },
+ "outputs": [],
+ "source": [
+ "def show_batch(images, labels, predictions=None):\n",
+ " plt.figure(figsize=(10, 10))\n",
+ " min = images.numpy().min()\n",
+ " max = images.numpy().max()\n",
+ " delta = max - min\n",
+ "\n",
+ " for i in range(BATCH_SIZE):\n",
+ " plt.subplot(4, 4, i + 1)\n",
+ " plt.imshow((images[i]-min) / delta)\n",
+ " if predictions is None:\n",
+ " plt.title(label_info.int2str(labels[i]))\n",
+ " else:\n",
+ " if labels[i] == predictions[i]:\n",
+ " color = 'g'\n",
+ " else:\n",
+ " color = 'r'\n",
+ " plt.title(label_info.int2str(predictions[i]), color=color)\n",
+ " plt.axis(\"off\")\n",
+ " plt.show()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "Du8GwFayrydM"
+ },
+ "outputs": [],
+ "source": [
+ "for images, labels in task.build_inputs(exp_config.task.validation_data).take(1):\n",
+ " show_batch(images, labels)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Oz9M3EDZro3W"
+ },
+ "source": [
+ "## Train and Evaluate"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 1214820,
+ "status": "ok",
+ "timestamp": 1699997521229,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "492B5k0cIo0e",
+ "outputId": "cb486eb5-5504-4e6c-e46d-df1e807d83e1"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "restoring or initializing model...\n",
+ "train | step: 0 | training until step 1017...\n",
+ "train | step: 1017 | steps/sec: 6.4 | output: \n",
+ " {'accuracy': 0.54013026,\n",
+ " 'learning_rate': 0.011706339,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.7053006}\n",
+ "saved checkpoint to ./trained_model/ckpt-1017.\n",
+ " eval | step: 1017 | running 290 steps of evaluation...\n",
+ " eval | step: 1017 | steps/sec: 17.6 | eval time: 16.4 sec | output: \n",
+ " {'accuracy': 0.57693964,\n",
+ " 'steps_per_second': 17.634644012635317,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.67286325}\n",
+ "train | step: 1017 | training until step 2034...\n",
+ "train | step: 2034 | steps/sec: 8.6 | output: \n",
+ " {'accuracy': 0.58898723,\n",
+ " 'learning_rate': 0.010854102,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.6713662}\n",
+ "saved checkpoint to ./trained_model/ckpt-2034.\n",
+ " eval | step: 2034 | running 290 steps of evaluation...\n",
+ " eval | step: 2034 | steps/sec: 37.3 | eval time: 7.8 sec | output: \n",
+ " {'accuracy': 0.6153017,\n",
+ " 'steps_per_second': 37.2743481981908,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.652928}\n",
+ "train | step: 2034 | training until step 3051...\n",
+ "train | step: 3051 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.6157817,\n",
+ " 'learning_rate': 0.009526712,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.65529084}\n",
+ "saved checkpoint to ./trained_model/ckpt-3051.\n",
+ " eval | step: 3051 | running 290 steps of evaluation...\n",
+ " eval | step: 3051 | steps/sec: 37.2 | eval time: 7.8 sec | output: \n",
+ " {'accuracy': 0.51681036,\n",
+ " 'steps_per_second': 37.18629141886086,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 1.0405935}\n",
+ "train | step: 3051 | training until step 4068...\n",
+ "train | step: 4068 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.6411013,\n",
+ " 'learning_rate': 0.007854102,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.63446975}\n",
+ "saved checkpoint to ./trained_model/ckpt-4068.\n",
+ " eval | step: 4068 | running 290 steps of evaluation...\n",
+ " eval | step: 4068 | steps/sec: 37.3 | eval time: 7.8 sec | output: \n",
+ " {'accuracy': 0.68297416,\n",
+ " 'steps_per_second': 37.306989215139936,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.59876406}\n",
+ "train | step: 4068 | training until step 5085...\n",
+ "train | step: 5085 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.6549287,\n",
+ " 'learning_rate': 0.0059999996,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.6212429}\n",
+ "saved checkpoint to ./trained_model/ckpt-5085.\n",
+ " eval | step: 5085 | running 290 steps of evaluation...\n",
+ " eval | step: 5085 | steps/sec: 37.7 | eval time: 7.7 sec | output: \n",
+ " {'accuracy': 0.67349136,\n",
+ " 'steps_per_second': 37.65721170518001,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.59417003}\n",
+ "train | step: 5085 | training until step 6102...\n",
+ "train | step: 6102 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.6778515,\n",
+ " 'learning_rate': 0.004145897,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.5970689}\n",
+ "saved checkpoint to ./trained_model/ckpt-6102.\n",
+ " eval | step: 6102 | running 290 steps of evaluation...\n",
+ " eval | step: 6102 | steps/sec: 37.3 | eval time: 7.8 sec | output: \n",
+ " {'accuracy': 0.68168104,\n",
+ " 'steps_per_second': 37.324507959229,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.6371893}\n",
+ "train | step: 6102 | training until step 7119...\n",
+ "train | step: 7119 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.7034169,\n",
+ " 'learning_rate': 0.0024732884,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.56852204}\n",
+ "saved checkpoint to ./trained_model/ckpt-7119.\n",
+ " eval | step: 7119 | running 290 steps of evaluation...\n",
+ " eval | step: 7119 | steps/sec: 37.4 | eval time: 7.7 sec | output: \n",
+ " {'accuracy': 0.7521552,\n",
+ " 'steps_per_second': 37.433325893948016,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.52360195}\n",
+ "train | step: 7119 | training until step 8136...\n",
+ "train | step: 8136 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.7234513,\n",
+ " 'learning_rate': 0.0011458977,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.54219157}\n",
+ "saved checkpoint to ./trained_model/ckpt-8136.\n",
+ " eval | step: 8136 | running 290 steps of evaluation...\n",
+ " eval | step: 8136 | steps/sec: 37.5 | eval time: 7.7 sec | output: \n",
+ " {'accuracy': 0.787931,\n",
+ " 'steps_per_second': 37.471471146938974,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.45562845}\n",
+ "train | step: 8136 | training until step 9153...\n",
+ "train | step: 9153 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.7457596,\n",
+ " 'learning_rate': 0.0002936611,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.50991994}\n",
+ "saved checkpoint to ./trained_model/ckpt-9153.\n",
+ " eval | step: 9153 | running 290 steps of evaluation...\n",
+ " eval | step: 9153 | steps/sec: 37.7 | eval time: 7.7 sec | output: \n",
+ " {'accuracy': 0.80625,\n",
+ " 'steps_per_second': 37.66015917575601,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.41287535}\n",
+ "train | step: 9153 | training until step 10170...\n",
+ "train | step: 10170 | steps/sec: 9.3 | output: \n",
+ " {'accuracy': 0.7517822,\n",
+ " 'learning_rate': 0.0,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'training_loss': 0.5036062}\n",
+ "saved checkpoint to ./trained_model/ckpt-10170.\n",
+ " eval | step: 10170 | running 290 steps of evaluation...\n",
+ " eval | step: 10170 | steps/sec: 37.6 | eval time: 7.7 sec | output: \n",
+ " {'accuracy': 0.81314653,\n",
+ " 'steps_per_second': 37.58710303226575,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.40270108}\n",
+ " eval | step: 10170 | running 290 steps of evaluation...\n",
+ " eval | step: 10170 | steps/sec: 37.3 | eval time: 7.8 sec | output: \n",
+ " {'accuracy': 0.81314653,\n",
+ " 'steps_per_second': 37.3229583933457,\n",
+ " 'top_5_accuracy': 1.0,\n",
+ " 'validation_loss': 0.40270108}\n"
+ ]
+ }
+ ],
+ "source": [
+ "model, eval_logs = core.train_lib.run_experiment(\n",
+ " distribution_strategy=distribution_strategy,\n",
+ " task=task,\n",
+ " mode='train_and_eval',\n",
+ " params=exp_config,\n",
+ " model_dir=model_dir,\n",
+ " run_post_eval=True)"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Yz8-qUIcFDKt"
+ },
+ "source": [
+ "## Export the trained model"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "base_uri": "https://localhost:8080/"
+ },
+ "executionInfo": {
+ "elapsed": 99764,
+ "status": "ok",
+ "timestamp": 1699997870693,
+ "user": {
+ "displayName": "Siva Sravana Kumar Neeli",
+ "userId": "06669604936988620923"
+ },
+ "user_tz": 480
+ },
+ "id": "SDX0RWMBtZB2",
+ "outputId": "c0a78a30-b278-4266-b8c2-41a68dd84cd0"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "2023-11-14 21:36:10.722529: W tensorflow/compiler/tf2tensorrt/utils/py_utils.cc:38] TF-TRT Warning: Could not find TensorRT\n",
+ "2023-11-14 21:36:18.463386: W tensorflow/core/common_runtime/gpu/gpu_bfc_allocator.cc:47] Overriding orig_value setting because the TF_FORCE_GPU_ALLOW_GROWTH environment variable is set. Original config value was 0.\n",
+ "I1114 21:36:49.038166 140374729912320 signature_serialization.py:148] Function `serve` contains input name(s) resource with unsupported characters which will be renamed to dense_biasadd_readvariableop_resource in the SavedModel.\n",
+ "I1114 21:37:34.619264 140374729912320 save.py:274] Found untraced functions such as serve_eval, conv2d_layer_call_fn, conv2d_layer_call_and_return_conditional_losses, _jit_compiled_convolution_op, conv_bn_layer_call_fn while saving (showing 5 of 415). These functions will not be directly callable after loading.\n",
+ "INFO:tensorflow:Assets written to: ./exported_model//saved_model/assets\n",
+ "I1114 21:37:46.884536 140374729912320 builder_impl.py:804] Assets written to: ./exported_model//saved_model/assets\n",
+ "I1114 21:37:47.379255 140374729912320 fingerprinting_utils.py:48] Writing fingerprint to ./exported_model//saved_model/fingerprint.pb\n",
+ "I1114 21:37:48.188559 140374729912320 train_utils.py:400] Saving experiment configuration to ./exported_model//params.yaml\n"
+ ]
+ }
+ ],
+ "source": [
+ "EXPORT_DIR_PATH = \"./exported_model/\"\n",
+ "!python -m official.projects.yolo.serving.export_saved_model \\\n",
+ " --experiment=\"darknet_classification\" \\\n",
+ " --export_dir=$EXPORT_DIR_PATH/ \\\n",
+ " --checkpoint_path=$model_dir \\\n",
+ " --config_file=$model_dir/params.yaml \\\n",
+ " --batch_size=$BATCH_SIZE \\\n",
+ " --input_type=\"image_tensor\" \\\n",
+ " --input_image_size=$IMG_SIZE,$IMG_SIZE"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "h64DRrJEr7ls"
+ },
+ "source": [
+ "### Test the exported model."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "JKcuFx17FBeu"
+ },
+ "source": [
+ "### Importing SavedModel"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "2-GADU74zkjK"
+ },
+ "outputs": [],
+ "source": [
+ "imported = tf.saved_model.load('/content/exported_model/saved_model')\n",
+ "model_fn = imported.signatures['serving_default']"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "tNKz-BD-sFHJ"
+ },
+ "source": [
+ "### Visualize the test predictions."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "9lE9sOYGxIkf"
+ },
+ "outputs": [],
+ "source": [
+ "def resize_image(record):\n",
+ " image = tf.image.resize(record['image'], size=(IMG_SIZE, IMG_SIZE))\n",
+ " image = tf.cast(image, tf.uint8)\n",
+ " return image, record['label']"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "81gCeHCava1w"
+ },
+ "outputs": [],
+ "source": [
+ "test_ds_resized = test_ds.map(resize_image).shuffle(100)\n",
+ "test_ds_batched = test_ds_resized.batch(BATCH_SIZE)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "P7b5ApE2bRFW"
+ },
+ "outputs": [],
+ "source": [
+ "for images, labels in test_ds_batched.take(1):\n",
+ " predictions = model_fn(inputs=images)['logits']\n",
+ " predictions = tf.argmax(predictions, axis=-1)\n",
+ "\n",
+ "show_batch(images, labels, predictions)"
+ ]
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "T4",
+ "provenance": [
+ {
+ "file_id": "1rt6hpJWOSKa2c746Ss58VERmnaA3gWys",
+ "timestamp": 1700078413508
+ }
+ ],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/vision/beta/projects/yolo/dataloaders/__init__.py b/official/projects/yolo/dataloaders/__init__.py
similarity index 89%
rename from official/vision/beta/projects/yolo/dataloaders/__init__.py
rename to official/projects/yolo/dataloaders/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/vision/beta/projects/yolo/dataloaders/__init__.py
+++ b/official/projects/yolo/dataloaders/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/beta/projects/yolo/dataloaders/classification_input.py b/official/projects/yolo/dataloaders/classification_input.py
old mode 100755
new mode 100644
similarity index 84%
rename from official/vision/beta/projects/yolo/dataloaders/classification_input.py
rename to official/projects/yolo/dataloaders/classification_input.py
index e1737dba354..c7bcffb6fe3
--- a/official/vision/beta/projects/yolo/dataloaders/classification_input.py
+++ b/official/projects/yolo/dataloaders/classification_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,9 @@
# limitations under the License.
"""Classification decoder and parser."""
-import tensorflow as tf
+from typing import List
+
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import classification_input
from official.vision.ops import preprocess_ops
@@ -90,3 +92,19 @@ def _parse_eval_image(self, decoded_tensors):
image = tf.image.convert_image_dtype(image, self._dtype)
image = image / 255.0
return image
+
+ @classmethod
+ def inference_fn(
+ cls, image: tf.Tensor, input_image_size: List[int], num_channels: int = 3
+ ) -> tf.Tensor:
+ """Builds image model inputs for serving."""
+
+ image = tf.cast(image, dtype=tf.float32)
+ image = preprocess_ops.center_crop_image(image)
+ image = tf.image.resize(
+ image, input_image_size, method=tf.image.ResizeMethod.BILINEAR
+ )
+
+ image.set_shape(input_image_size + [num_channels])
+ image = image / 255.0
+ return image
diff --git a/official/vision/beta/projects/yolo/dataloaders/tf_example_decoder.py b/official/projects/yolo/dataloaders/tf_example_decoder.py
similarity index 98%
rename from official/vision/beta/projects/yolo/dataloaders/tf_example_decoder.py
rename to official/projects/yolo/dataloaders/tf_example_decoder.py
index 8578b663f5c..c172634ef85 100644
--- a/official/vision/beta/projects/yolo/dataloaders/tf_example_decoder.py
+++ b/official/projects/yolo/dataloaders/tf_example_decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
A decoder to decode string tensors containing serialized tensorflow.Example
protos for object detection.
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import tf_example_decoder
diff --git a/official/vision/beta/projects/yolo/dataloaders/yolo_input.py b/official/projects/yolo/dataloaders/yolo_input.py
old mode 100755
new mode 100644
similarity index 93%
rename from official/vision/beta/projects/yolo/dataloaders/yolo_input.py
rename to official/projects/yolo/dataloaders/yolo_input.py
index a6591aa0097..323b7bdf9d7
--- a/official/vision/beta/projects/yolo/dataloaders/yolo_input.py
+++ b/official/projects/yolo/dataloaders/yolo_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,10 +13,10 @@
# limitations under the License.
"""Detection Data parser and processing for YOLO."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import anchor
-from official.vision.beta.projects.yolo.ops import preprocessing_ops
+from official.projects.yolo.ops import anchor
+from official.projects.yolo.ops import preprocessing_ops
from official.vision.dataloaders import parser
from official.vision.dataloaders import utils
from official.vision.ops import box_ops as bbox_ops
@@ -339,32 +339,26 @@ def _build_label(self,
'source_id': utils.process_source_id(data['source_id']),
'bbox': tf.cast(boxes, dtype=self._dtype),
'classes': tf.cast(classes, dtype=self._dtype),
+ # For OTA loss.
+ 'image_info': info,
})
# Update the labels dictionary.
if not is_training:
# Sets up groundtruth data for evaluation.
groundtruths = {
- 'source_id':
- labels['source_id'],
- 'height':
- data['height'],
- 'width':
- data['width'],
- 'num_detections':
- tf.shape(data['groundtruth_boxes'])[0],
- 'image_info':
- info,
- 'boxes':
- bbox_ops.denormalize_boxes(
- data['groundtruth_boxes'],
- tf.cast([data['height'], data['width']], gt_boxes.dtype)),
- 'classes':
- data['groundtruth_classes'],
- 'areas':
- data['groundtruth_area'],
- 'is_crowds':
- tf.cast(tf.gather(data['groundtruth_is_crowd'], inds), tf.int32),
+ 'source_id': labels['source_id'],
+ 'height': data['height'],
+ 'width': data['width'],
+ 'num_detections': tf.shape(data['groundtruth_boxes'])[0],
+ 'image_info': info,
+ 'boxes': bbox_ops.denormalize_boxes(
+ data['groundtruth_boxes'],
+ tf.cast([data['height'], data['width']], gt_boxes.dtype)),
+ 'classes': data['groundtruth_classes'],
+ 'areas': data['groundtruth_area'],
+ 'is_crowds': tf.cast(
+ tf.gather(data['groundtruth_is_crowd'], inds), tf.int32),
}
groundtruths['source_id'] = utils.process_source_id(
groundtruths['source_id'])
diff --git a/official/projects/yolo/losses/__init__.py b/official/projects/yolo/losses/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/losses/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/losses/yolo_loss.py b/official/projects/yolo/losses/yolo_loss.py
old mode 100755
new mode 100644
similarity index 95%
rename from official/vision/beta/projects/yolo/losses/yolo_loss.py
rename to official/projects/yolo/losses/yolo_loss.py
index a123a04e2ab..b2cd349e596
--- a/official/vision/beta/projects/yolo/losses/yolo_loss.py
+++ b/official/projects/yolo/losses/yolo_loss.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,11 +17,11 @@
import collections
import functools
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import box_ops
-from official.vision.beta.projects.yolo.ops import loss_utils
-from official.vision.beta.projects.yolo.ops import math_ops
+from official.projects.yolo.ops import box_ops
+from official.projects.yolo.ops import loss_utils
+from official.projects.yolo.ops import math_ops
class YoloLossBase(object, metaclass=abc.ABCMeta):
@@ -323,7 +323,7 @@ def _compute_loss(self, true_counts, inds, y_true, boxes, classes, y_pred):
grid_points = tf.stop_gradient(grid_points)
anchor_grid = tf.stop_gradient(anchor_grid)
- # Split all the ground truths to use as seperate items in loss computation.
+ # Split all the ground truths to use as separate items in loss computation.
(true_box, ind_mask, true_class) = tf.split(y_true, [4, 1, 1], axis=-1)
true_conf = tf.squeeze(true_conf, axis=-1)
true_class = tf.squeeze(true_class, axis=-1)
@@ -411,7 +411,7 @@ def _compute_loss(self, true_counts, inds, y_true, boxes, classes, y_pred):
# make training more unstable but may also return higher APs.
pred_class = loss_utils.apply_mask(
ind_mask, tf.gather_nd(pred_class, inds, batch_dims=1))
- class_loss = tf.keras.losses.binary_crossentropy(
+ class_loss = tf_keras.losses.binary_crossentropy(
tf.expand_dims(true_class, axis=-1),
tf.expand_dims(pred_class, axis=-1),
label_smoothing=self._label_smoothing,
@@ -431,7 +431,7 @@ def _compute_loss(self, true_counts, inds, y_true, boxes, classes, y_pred):
# Mask the confidence loss and take the sum across all the grid cells.
if self._ignore_thresh != 0.0:
- bce = loss_utils.apply_mask(obj_mask, bce)
+ bce = loss_utils.apply_mask(obj_mask, bce) # pyrefly: ignore[unbound-name]
conf_loss = tf.cast(tf.reduce_sum(bce, axis=(1, 2, 3)), dtype=y_pred.dtype)
# Apply the weights to each loss.
@@ -556,16 +556,16 @@ def _compute_loss(self, true_counts, inds, y_true, boxes, classes, y_pred):
true_conf = tf.squeeze(true_conf, axis=-1)
# Compute the cross entropy loss for the confidence map.
- bce = tf.keras.losses.binary_crossentropy(
+ bce = tf_keras.losses.binary_crossentropy(
tf.expand_dims(true_conf, axis=-1), pred_conf, from_logits=True)
if self._ignore_thresh != 0.0:
- bce = loss_utils.apply_mask(obj_mask, bce)
+ bce = loss_utils.apply_mask(obj_mask, bce) # pyrefly: ignore[unbound-name]
conf_loss = tf.reduce_sum(bce) / tf.reduce_sum(obj_mask)
else:
conf_loss = tf.reduce_mean(bce)
# Compute the cross entropy loss for the class maps.
- class_loss = tf.keras.losses.binary_crossentropy(
+ class_loss = tf_keras.losses.binary_crossentropy(
true_class,
pred_class,
label_smoothing=self._label_smoothing,
@@ -692,17 +692,17 @@ def __init__(self,
self._loss_dict[key] = losses[loss_type](
classes=classes,
anchors=anchors[key],
- truth_thresh=truth_thresholds[key],
- ignore_thresh=ignore_thresholds[key],
- loss_type=loss_types[key],
- iou_normalizer=iou_normalizers[key],
- cls_normalizer=cls_normalizers[key],
- object_normalizer=object_normalizers[key],
- box_type=box_types[key],
- objectness_smooth=objectness_smooths[key],
- max_delta=max_deltas[key],
- path_stride=path_strides[key],
- scale_x_y=scale_xys[key],
+ truth_thresh=truth_thresholds[key], # pyrefly: ignore[unsupported-operation]
+ ignore_thresh=ignore_thresholds[key], # pyrefly: ignore[unsupported-operation]
+ loss_type=loss_types[key], # pyrefly: ignore[unsupported-operation]
+ iou_normalizer=iou_normalizers[key], # pyrefly: ignore[unsupported-operation]
+ cls_normalizer=cls_normalizers[key], # pyrefly: ignore[unsupported-operation]
+ object_normalizer=object_normalizers[key], # pyrefly: ignore[unsupported-operation]
+ box_type=box_types[key], # pyrefly: ignore[unsupported-operation]
+ objectness_smooth=objectness_smooths[key], # pyrefly: ignore[unsupported-operation]
+ max_delta=max_deltas[key], # pyrefly: ignore[unsupported-operation]
+ path_stride=path_strides[key], # pyrefly: ignore[unsupported-operation]
+ scale_x_y=scale_xys[key], # pyrefly: ignore[unsupported-operation]
update_on_repeat=update_on_repeat,
label_smoothing=label_smoothing)
diff --git a/official/vision/beta/projects/yolo/losses/yolo_loss_test.py b/official/projects/yolo/losses/yolo_loss_test.py
old mode 100755
new mode 100644
similarity index 93%
rename from official/vision/beta/projects/yolo/losses/yolo_loss_test.py
rename to official/projects/yolo/losses/yolo_loss_test.py
index 9a8f8a7816d..99c1766ff5d
--- a/official/vision/beta/projects/yolo/losses/yolo_loss_test.py
+++ b/official/projects/yolo/losses/yolo_loss_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
"""Tests for yolo heads."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.losses import yolo_loss
+from official.projects.yolo.losses import yolo_loss
class YoloDecoderTest(parameterized.TestCase, tf.test.TestCase):
@@ -35,7 +35,7 @@ def inpdict(input_shape, dtype=tf.float32):
inputs[key] = tf.ones(input_shape[key], dtype=dtype)
return inputs
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
input_shape = {
'3': [1, 52, 52, 255],
'4': [1, 26, 26, 255],
diff --git a/official/projects/yolo/losses/yolov7_loss.py b/official/projects/yolo/losses/yolov7_loss.py
new file mode 100644
index 00000000000..cf3bcc68d0f
--- /dev/null
+++ b/official/projects/yolo/losses/yolov7_loss.py
@@ -0,0 +1,954 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""YOLOv7 loss function."""
+
+import tensorflow as tf, tf_keras
+
+from official.projects.yolo.ops import box_ops
+from official.vision.losses import focal_loss
+
+_LAYER_BALANCE = {
+ '3': [4.0, 1.0, 0.4],
+ '5': [4.0, 1.0, 0.25, 0.06, 0.02],
+}
+
+
+def smooth_bce_targets(eps=0.1):
+ """Computes positive, negative label smoothing BCE targets.
+
+ https://arxiv.org/pdf/1902.04103.pdf equation 3.
+
+ Args:
+ eps: a float number from [0, 1] representing label smoothing factor.
+
+ Returns:
+ Positive and negative targets after label smoothing.
+ """
+ return 1.0 - 0.5 * eps, 0.5 * eps
+
+
+def merge_labels(labels):
+ """Converts the ground-truth labels into loss targets."""
+ boxes = box_ops.yxyx_to_xcycwh(labels['bbox'])
+ classes = tf.cast(labels['classes'], boxes.dtype)
+ return tf.concat([classes[..., None], boxes], axis=-1)
+
+
+class YoloV7Loss(tf_keras.losses.Loss):
+ """YOLOv7 loss function."""
+
+ def __init__(
+ self,
+ anchors,
+ strides,
+ input_size,
+ alpha=0.25,
+ gamma=1.5,
+ box_weight=0.05,
+ obj_weight=0.7,
+ cls_weight=0.3,
+ label_smoothing=0.0,
+ anchor_threshold=4.0,
+ iou_mix_ratio=1.0,
+ num_classes=80,
+ auto_balance=False,
+ reduction=tf_keras.losses.Reduction.NONE,
+ name=None,
+ ):
+ """Constructor for YOLOv7 loss.
+
+ Follows the implementation here:
+ https://github.com/WongKinYiu/yolov7/blob/main/utils/loss.py#L422
+
+ Args:
+ anchors: a 2D array represents different anchors used at each level.
+ strides: a 1D array represents the strides. Note that all numbers should
+ be a power of 2, and they usually start with level 3 and end at level
+ 5 or 7. Therefore, the list should usually be [8, 16, 32] or
+ [8, 16, 32, 64, 128].
+ input_size: a list containing the height and width of the input image.
+ alpha: alpha for focal loss.
+ gamma: gamma for focal loss. If set to 0, focal loss will be disabled.
+ box_weight: float weight scalar applied to bounding box loss.
+ obj_weight: float weight scalar applied to objectness loss.
+ cls_weight: float weight scalar applied to class loss.
+ label_smoothing: small float number used to compute positive and negative
+ targets. If set to 0, the positive targets will be 1 and negative
+ targets will be 0.
+ anchor_threshold: threshold for the anchor matching. Larger number allows
+ more displacements between anchors and targets.
+ iou_mix_ratio: float ratio to mix the IoU score with the positive target,
+ which is 1.
+ num_classes: number of classes.
+ auto_balance: a boolean flag that indicates whether auto balance should be
+ used. If used, the default balance factors will automatically update
+ for each batch.
+ reduction: Reduction method. Should be set to None at all time as this
+ loss module always output a loss scalar.
+ name: Optional name for the loss.
+ """
+ # Loss required fields.
+ self._num_classes = num_classes
+ self._num_layers = len(strides)
+ self._num_anchors = len(anchors[0])
+ self._anchors = anchors
+ self._strides = strides
+ self._input_size = input_size
+ self._iou_mix_ratio = iou_mix_ratio
+
+ # Scale down anchors by the strides to match the feature map.
+ for i, stride in enumerate(strides):
+ self._anchors[i] = tf.constant(self._anchors[i], tf.float32) / stride
+
+ self._anchor_threshold = anchor_threshold
+
+ self._pos_targets, self._neg_targets = smooth_bce_targets(label_smoothing)
+ if gamma > 0:
+ self._cls_loss_fn = focal_loss.FocalLoss(
+ alpha=alpha, gamma=gamma, reduction=reduction, name='cls_loss')
+ self._obj_loss_fn = focal_loss.FocalLoss(
+ alpha=alpha, gamma=gamma, reduction=reduction, name='obj_loss')
+ else:
+ self._cls_loss_fn = tf.nn.sigmoid_cross_entropy_with_logits
+ self._obj_loss_fn = tf.nn.sigmoid_cross_entropy_with_logits
+
+ # Weight to combine losses
+ self._box_weight = box_weight
+ self._obj_weight = obj_weight * input_size[0] / 640 * input_size[1] / 640
+ self._cls_weight = cls_weight * num_classes / 80
+
+ # Layer balance scalar
+ self._balance = _LAYER_BALANCE[str(self._num_layers)][:]
+ for i, bal in enumerate(self._balance):
+ self._balance[i] = tf.constant(bal, tf.float32)
+ self._auto_balance = auto_balance
+ assert 16 in strides, (
+ 'Expect level 4 (stride of 16) always exist in the strides, received %s'
+ % strides
+ )
+ self._ssi = list(strides).index(16) if auto_balance else 0 # stride 16 idx
+
+ super().__init__(reduction=reduction, name=name)
+
+ def call(self, labels, predictions):
+ labels = merge_labels(labels)
+ p = {}
+ for key in predictions:
+ # [batch_size, num_anchors, height, width, num_classes + boxes + obj]
+ p[key] = tf.transpose(predictions[key], [0, 3, 1, 2, 4])
+ cls_loss, box_loss, obj_loss, iou_metric = [tf.zeros(1) for _ in range(4)]
+ total_num_matchings = tf.zeros(1)
+ total_num_gts = tf.reduce_sum(tf.cast(labels[..., 0] != -1, tf.float32))
+
+ masks, indices, anchors, cls_targets, box_targets = self._build_targets(
+ labels, p)
+
+ batch_size = tf.shape(indices)[0]
+ layer_shape = [batch_size, self._num_layers, -1]
+ # [anchor_indices, grid_js, grid_is]
+ masks = tf.reshape(masks, layer_shape)
+ indices = tf.reshape(indices, [*layer_shape, 3])
+ anchors = tf.reshape(anchors, [*layer_shape, 2])
+ cls_targets = tf.reshape(cls_targets, layer_shape)
+ box_targets = tf.reshape(box_targets, [*layer_shape, 4])
+
+ # Losses
+ for layer_key, layer_pred in p.items():
+ i = int(layer_key) - 3
+
+ obj_targets = tf.zeros_like(layer_pred[..., 0])
+
+ layer_masks = masks[:, i]
+ num_matchings = tf.reduce_sum(tf.cast(layer_masks, tf.int32))
+ total_num_matchings += tf.cast(num_matchings, tf.float32)
+
+ if num_matchings > 0:
+ layer_indices = indices[:, i]
+ batch_indices = tf.tile(
+ tf.range(batch_size)[:, None], [1, tf.shape(layer_indices)[1]]
+ )[..., None]
+ layer_indices = tf.concat([batch_indices, layer_indices], axis=-1)
+ layer_indices = tf.boolean_mask(layer_indices, layer_masks)
+ layer_anchors = tf.boolean_mask(anchors[:, i], layer_masks)
+
+ layer_cls_targets = tf.boolean_mask(cls_targets[:, i], layer_masks)
+ layer_box_targets = tf.boolean_mask(box_targets[:, i], layer_masks)
+
+ # In the same shape of layer_target.
+ matched_pred = tf.gather_nd(layer_pred, layer_indices)
+
+ pred_xcyc = tf.sigmoid(matched_pred[..., :2]) * 2 - 0.5
+ pred_wh = (
+ tf.square(tf.sigmoid(matched_pred[..., 2:4]) * 2) * layer_anchors)
+ pred_xcycwh = tf.concat([pred_xcyc, pred_wh], axis=-1)
+ _, ciou = box_ops.compute_ciou(pred_xcycwh, layer_box_targets)
+
+ box_loss += tf.reduce_mean(1.0 - ciou)
+ iou_metric += tf.reduce_mean(ciou)
+
+ # Compute classification loss.
+ if self._num_classes > 1: # cls loss (only if multiple classes)
+ t = tf.one_hot(
+ layer_cls_targets,
+ self._num_classes,
+ on_value=self._pos_targets,
+ off_value=self._neg_targets,
+ )
+ cls_loss += tf.reduce_mean(
+ self._cls_loss_fn(t, matched_pred[..., 5:]))
+
+ # Compute objectness loss.
+ iou_ratio = tf.cast(
+ (1.0 - self._iou_mix_ratio)
+ + (self._iou_mix_ratio * tf.maximum(tf.stop_gradient(ciou), 0)),
+ obj_targets.dtype,
+ )
+ obj_targets = tf.tensor_scatter_nd_max(
+ obj_targets, layer_indices, iou_ratio
+ )
+ layer_obj_loss = tf.reduce_mean(
+ self._obj_loss_fn(obj_targets, layer_pred[..., 4])
+ )
+ obj_loss += layer_obj_loss * self._balance[i]
+ # Updates the balance factor, which is a moving average of previous
+ # factor at the same level.
+ if self._auto_balance:
+ self._balance[i] = self._balance[
+ i
+ ] * 0.9999 + 0.0001 / tf.stop_gradient(layer_obj_loss)
+
+ # Re-balance the factors so that stride at self._ssi always receives 1.
+ if self._auto_balance:
+ self._balance = [x / self._balance[self._ssi] for x in self._balance]
+
+ box_loss *= self._box_weight
+ obj_loss *= self._obj_weight
+ cls_loss *= self._cls_weight
+
+ self._box_loss = tf.stop_gradient(box_loss)
+ self._obj_loss = tf.stop_gradient(obj_loss)
+ self._cls_loss = tf.stop_gradient(cls_loss)
+ self._iou = tf.stop_gradient(iou_metric) / self._num_layers
+ self._num_matchings = tf.stop_gradient(
+ total_num_matchings) / tf.cast(batch_size, tf.float32)
+ self._num_gts = tf.stop_gradient(
+ total_num_gts) / tf.cast(batch_size, tf.float32)
+
+ loss = box_loss + obj_loss + cls_loss
+ return loss * tf.cast(batch_size, loss.dtype)
+
+ def _build_targets(self, labels, predictions):
+ """Finds three matching anchors for each ground-truth."""
+ label_shape = tf.shape(labels)
+ batch_size, max_boxes = label_shape[0], label_shape[1]
+ masks, indices, anch = [], [], []
+ cls_targets, box_targets = [], []
+ anchor_indices = tf.tile(
+ tf.range(self._num_anchors, dtype=tf.float32)[None, None],
+ [batch_size, max_boxes, 1],
+ )
+ # Append anchor indices to labels.
+ labels = tf.tile(labels[:, :, None], [1, 1, self._num_anchors, 1])
+ labels = tf.concat([labels, anchor_indices[..., None]], axis=-1)
+
+ # Bias is used to determine the matching. 0.5 means matching anchors that
+ # fall in the 0.5 differences in the feature map. For instance, a box
+ # coordinates of (15.6, 35.4) will match the anchors at [15, 35], [16, 35],
+ # and [15, 34].
+ bias = 0.5 # bias
+ off = (
+ tf.constant(
+ [
+ [0, 0],
+ [1, 0], [0, 1], [-1, 0], [0, -1], # j, k, l, m
+ # [1, 1], [1, -1], [-1, 1], [-1, -1], # jk,jm,lk,lm
+ ],
+ tf.float32,
+ )
+ * bias
+ ) # offsets
+
+ for i in range(self._num_layers):
+ anchors = self._anchors[i]
+ _, _, h, w, _ = predictions[str(i + 3)].get_shape().as_list()
+ gain = tf.constant([1, w, h, w, h, 1], dtype=tf.float32)
+
+ t = labels * gain
+
+ # Filter out targets that do not match the current anchors.
+ wh_ratio = t[..., 3:5] / tf.cast(anchors[None, None], tf.float32)
+ labels_mask = tf.less(
+ tf.reduce_max(tf.maximum(wh_ratio, 1.0 / wh_ratio), axis=-1),
+ self._anchor_threshold,
+ )[..., None]
+ # Compute valid mask for ground-truths.
+ labels_mask = tf.logical_and(t[..., :1] != -1, labels_mask)
+
+ labels_mask = tf.reshape(labels_mask, [batch_size, -1])
+ t = tf.reshape(t, [batch_size, -1, 6])
+
+ # Find the matching offsets for valid labels.
+ gxy = t[..., 1:3] # grid xy
+ gxi = gain[1:3] - gxy # inverse
+ j, k = tf.split((gxy % 1.0 < bias) & (gxy >= 1.0), 2, axis=-1)
+ l, m = tf.split((gxi % 1.0 < bias) & (gxi >= 1.0), 2, axis=-1)
+
+ j, k, l, m = j[..., 0], k[..., 0], l[..., 0], m[..., 0]
+
+ # Note that j and l, k and m are conjugate to each other, so at most one
+ # of them will be True during running. Therefore, we can reduce memory
+ # usage by gathering the selected index.
+ x_map = tf.cast(tf.stack([j, l], axis=-1), tf.int8)
+ y_map = tf.cast(tf.stack([k, m], axis=-1), tf.int8)
+
+ # Add the indices offsets.
+ x_indices = tf.argmax(x_map, axis=-1) * 2 + 1
+ y_indices = tf.argmax(y_map, axis=-1) * 2 + 2
+ three_targets_indices = tf.stack(
+ [tf.zeros_like(x_indices), x_indices, y_indices], axis=-1
+ )[..., None]
+
+ # Gather the selected 3 targets from the 5-target map.
+ j = tf.stack([tf.ones_like(j), j, k, l, m], axis=-1)
+ three_targets_mask = tf.gather_nd(j, three_targets_indices, batch_dims=2)
+
+ labels_mask = tf.tile(labels_mask[:, :, None], [1, 1, 5])
+ t = tf.tile(t[:, :, None], [1, 1, 5, 1])
+
+ labels_mask = tf.gather_nd(
+ labels_mask, three_targets_indices, batch_dims=2
+ )
+ t = tf.gather_nd(t, three_targets_indices, batch_dims=2)
+
+ offsets = tf.zeros_like(gxy)[:, :, None] + off[None, None]
+ offsets = tf.gather_nd(offsets, three_targets_indices, batch_dims=2)
+
+ cls_target = tf.cast(t[..., 0], tf.int32)
+ gxy, gwh = t[..., 1:3], t[..., 3:5]
+ # Find the actual grid locations.
+ gij = tf.cast(gxy - offsets * 2, tf.int32)
+ gi, gj = tf.split(gij, 2, axis=-1)
+ gi, gj = gi[..., 0], gj[..., 0]
+
+ # Append the result.
+ anchor_idx = tf.cast(t[..., 5], tf.int32)
+ gain = tf.cast(gain, tf.int32)
+ gi = tf.clip_by_value(gi, 0, gain[2] - 1)
+ gj = tf.clip_by_value(gj, 0, gain[3] - 1)
+ gij = tf.stack([gi, gj], axis=-1)
+
+ labels_mask = tf.logical_and(labels_mask, three_targets_mask)
+ masks.append(labels_mask)
+ indices.append(tf.stack([anchor_idx, gj, gi], axis=-1))
+ anch.append(tf.gather(anchors, anchor_idx))
+ cls_targets.append(cls_target)
+ box_targets.append(
+ tf.concat([gxy - tf.cast(gij, tf.float32), gwh], axis=-1)) # box
+
+ # [batch_size, num_layers, num_anchors * max_boxes, num_targets]
+ masks = tf.stack(masks, axis=1)
+ indices = tf.stack(indices, axis=1)
+ anch = tf.stack(anch, axis=1)
+ cls_targets = tf.stack(cls_targets, axis=1)
+ box_targets = tf.stack(box_targets, axis=1)
+ return masks, indices, anch, cls_targets, box_targets
+
+ def report_separate_losses(self):
+ return {
+ 'box_loss': self._box_loss,
+ 'obj_loss': self._obj_loss,
+ 'cls_loss': self._cls_loss,
+ 'iou': self._iou,
+ }
+
+ def report_stats(self):
+ return {
+ 'num_gts': self._num_gts,
+ 'num_matchings': self._num_matchings,
+ # No duplicates.
+ 'num_duplicates': tf.constant(0),
+ }
+
+ def get_config(self):
+ config = {
+ 'alpha': self._alpha,
+ 'gamma': self._gamma,
+ 'box_weight': self._box_weight,
+ 'obj_weight': self._obj_weight,
+ 'cls_weight': self._cls_weight,
+ 'pos_targets': self._pos_targets,
+ 'neg_targets': self._neg_targets,
+ 'num_classes': self._num_classes,
+ 'num_layers': self._num_layers,
+ 'num_anchors': self._num_anchors,
+ 'auto_balance': self._auto_balance,
+ 'balance': self._balance,
+ 'strides': self._strides,
+ 'anchors': self._anchors,
+ 'input_size': self._input_size,
+ 'anchor_threshold': self._anchor_threshold,
+ }
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(config.items()))
+
+
+class YoloV7LossOTA(tf_keras.losses.Loss):
+ """YOLOv7 loss function with OTA.
+
+ OTA (Optimal Transport Assignment) uses Sinkhorn-Knopp algorithm to copmute
+ a matching between anchors and ground-truth labels.
+
+ Paper: https://arxiv.org/pdf/2103.14259.pdf
+ """
+
+ def __init__(
+ self,
+ anchors,
+ strides,
+ input_size,
+ alpha=0.25,
+ gamma=1.5,
+ box_weight=0.05,
+ obj_weight=0.7,
+ cls_weight=0.3,
+ iou_weight=3.0,
+ label_smoothing=0.0,
+ anchor_threshold=4.0,
+ iou_mix_ratio=1.0,
+ num_classes=80,
+ auto_balance=False,
+ reduction=tf_keras.losses.Reduction.NONE,
+ name=None,
+ ):
+ """Constructor for YOLOv7 loss OTA.
+
+ Follows the implementation here:
+ https://github.com/WongKinYiu/yolov7/blob/main/utils/loss.py#L556
+
+ Args:
+ anchors: a 2D array represents different anchors used at each level.
+ strides: a 1D array represents the strides. Note that all numbers should
+ be a power of 2, and they usually start with level 3 and end at level 5
+ or 7. Therefore, the list should usually be [8, 16, 32] or [8, 16, 32,
+ 64, 128].
+ input_size: a list containing the height and width of the input image.
+ alpha: alpha for focal loss.
+ gamma: gamma for focal loss. If set to 0, focal loss will be disabled.
+ box_weight: float weight scalar applied to bounding box loss.
+ obj_weight: float weight scalar applied to objectness loss.
+ cls_weight: float weight scalar applied to class loss.
+ iou_weight: float weight scalar to mix class loss and IoU class to
+ construct the cost matrix.
+ label_smoothing: small float number used to compute positive and negative
+ targets. If set to 0, the positive targets will be 1 and negative
+ targets will be 0.
+ anchor_threshold: threshold for the anchor matching. Larger number allows
+ more displacements between anchors and targets.
+ iou_mix_ratio: float ratio to mix the IoU score with the positive target,
+ which is 1.
+ num_classes: number of classes.
+ auto_balance: a boolean flag that indicates whether auto balance should be
+ used. If used, the default balance factors will automatically update for
+ each batch.
+ reduction: Reduction method. Should be set to None at all time as this
+ loss module always output a loss scalar.
+ name: Optional name for the loss.
+ """
+ # Loss required fields.
+ self._num_classes = num_classes
+ self._num_layers = len(strides)
+ self._num_anchors = len(anchors[0])
+ self._anchors = []
+ self._strides = strides
+ self._input_size = input_size
+ self._iou_mix_ratio = iou_mix_ratio
+
+ # Scale down anchors by the strides to match the feature map.
+ for i, stride in enumerate(strides):
+ self._anchors.append(tf.constant(anchors[i], tf.float32) / stride)
+
+ self._anchor_threshold = anchor_threshold
+
+ self._pos_targets, self._neg_targets = smooth_bce_targets(label_smoothing)
+ if gamma > 0:
+ self._cls_loss_fn = focal_loss.FocalLoss(
+ alpha=alpha, gamma=gamma, reduction=reduction, name='cls_loss')
+ self._obj_loss_fn = focal_loss.FocalLoss(
+ alpha=alpha, gamma=gamma, reduction=reduction, name='obj_loss')
+ else:
+ self._cls_loss_fn = tf.nn.sigmoid_cross_entropy_with_logits
+ self._obj_loss_fn = tf.nn.sigmoid_cross_entropy_with_logits
+
+ # Weight to combine losses
+ self._box_weight = box_weight
+ self._obj_weight = obj_weight * input_size[0] / 640 * input_size[1] / 640
+ self._cls_weight = cls_weight * num_classes / 80
+
+ # Weight to construct cost matrix
+ self._iou_weight = iou_weight
+
+ # Layer balance scalar
+ self._balance = _LAYER_BALANCE[str(self._num_layers)][:]
+ for i, bal in enumerate(self._balance):
+ self._balance[i] = tf.constant(bal, tf.float32)
+ self._auto_balance = auto_balance
+ assert 16 in strides, (
+ 'Expect level 4 (stride of 16) always exist in the strides, received %s'
+ % strides
+ )
+ self._ssi = list(strides).index(16) if auto_balance else 0 # stride 16 idx
+
+ super().__init__(reduction=reduction, name=name)
+
+ def call(self, labels, predictions):
+ """Comptues the OTA loss.
+
+ Args:
+ labels: a dictionary contains the following required keys:
+ - classes: class indices in shape [batch_size, max_num_instances].
+ - bbox: bounding boxes in shape [batch_size, max_num_instances, 4].
+ - image_info: image info in shape [batch_size, 4, 2].
+ predictions: a dictionary contains model outputs at different layers.
+ They are in shape of [batch_size, h_at_level, w_at_level, num_anchors,
+ num_classes + 4 (box coordinates) + 1 (objectness)].
+
+ Returns:
+ The scaled loss (up by batch size) from OTA.
+ """
+ image_info = labels['image_info']
+ # Convert labels dictionary into tensors.
+ labels = merge_labels(labels)
+ p = {}
+ for key in predictions:
+ # [batch_size, num_anchors, height, width, num_classes + boxes + obj]
+ p[key] = tf.transpose(predictions[key], [0, 3, 1, 2, 4])
+
+ cls_loss, box_loss, obj_loss, iou_metric = [tf.zeros(1) for _ in range(4)]
+ total_num_matchings = tf.zeros(1)
+ total_num_gts = tf.reduce_sum(tf.cast(labels[..., 0] != -1, tf.float32))
+ (matched_indices, matched_anchors, matched_mask, matched_targets,
+ num_duplicates) = self._build_targets(labels, p, image_info)
+ # Get height and width for each layers.
+ pre_gen_gains = [
+ tf.gather(tf.shape(p[str(i + 3)]), [3, 2, 3, 2])
+ for i in range(self._num_layers)
+ ]
+
+ batch_size = tf.shape(matched_indices)[0]
+ layer_shape = [batch_size, self._num_layers, -1]
+ # [anchor_indices, grid_js, grid_is]
+ masks = tf.reshape(matched_mask, layer_shape)
+ indices = tf.reshape(matched_indices, [*layer_shape, 3])
+ anchors = tf.reshape(matched_anchors, [*layer_shape, 2])
+ targets = tf.reshape(matched_targets, [*layer_shape, 5])
+
+ # Losses
+ for layer_idx, layer_pred in p.items():
+ # Always assume the output level starts with 3.
+ i = int(layer_idx) - 3
+
+ obj_targets = tf.zeros_like(layer_pred[..., 0])
+
+ # Get layer inputs
+ layer_masks = masks[:, i]
+ num_matchings = tf.reduce_sum(tf.cast(layer_masks, tf.int32))
+ total_num_matchings += tf.cast(num_matchings, tf.float32)
+
+ if num_matchings > 0:
+ layer_indices = indices[:, i]
+ batch_indices = tf.tile(
+ tf.range(batch_size)[:, None], [1, tf.shape(layer_indices)[1]]
+ )[..., None]
+ layer_indices = tf.concat([batch_indices, layer_indices], axis=-1)
+ layer_indices = tf.boolean_mask(layer_indices, layer_masks)
+ layer_anchors = tf.boolean_mask(anchors[:, i], layer_masks)
+
+ layer_targets = tf.boolean_mask(targets[:, i], layer_masks)
+ layer_cls_targets = tf.cast(layer_targets[:, 0], tf.int32)
+ layer_box_targets = layer_targets[:, 1:]
+
+ # In the same shape of layer_target.
+ matched_pred = tf.gather_nd(layer_pred, layer_indices)
+
+ pred_xcyc = tf.sigmoid(matched_pred[..., :2]) * 2 - 0.5
+ pred_wh = (
+ tf.square(tf.sigmoid(matched_pred[..., 2:4]) * 2) * layer_anchors)
+ pred_xcycwh = tf.concat([pred_xcyc, pred_wh], axis=-1)
+
+ grid = tf.cast(
+ tf.stack(
+ [
+ layer_indices[:, 3], # gi
+ layer_indices[:, 2], # gj
+ tf.zeros_like(layer_indices[:, 0]),
+ tf.zeros_like(layer_indices[:, 0]),
+ ],
+ axis=-1,
+ ),
+ tf.float32,
+ )
+ target_xcycwh = layer_box_targets * tf.cast(
+ pre_gen_gains[i], layer_targets.dtype
+ )
+ target_xcycwh -= grid
+ _, ciou = box_ops.compute_ciou(target_xcycwh, pred_xcycwh)
+
+ box_loss += tf.reduce_mean(1.0 - ciou)
+ iou_metric += tf.reduce_mean(ciou)
+
+ # Compute classification loss.
+ if self._num_classes > 1: # cls loss (only if multiple classes)
+ t = tf.one_hot(
+ layer_cls_targets,
+ self._num_classes,
+ on_value=self._pos_targets,
+ off_value=self._neg_targets,
+ )
+ cls_loss += tf.reduce_mean(
+ self._cls_loss_fn(t, matched_pred[..., 5:]))
+
+ # Compute objectness loss.
+ iou_ratio = tf.cast(
+ (1.0 - self._iou_mix_ratio)
+ + (self._iou_mix_ratio * tf.maximum(tf.stop_gradient(ciou), 0)),
+ obj_targets.dtype,
+ )
+ obj_targets = tf.tensor_scatter_nd_max(
+ obj_targets, layer_indices, iou_ratio
+ )
+ layer_obj_loss = tf.reduce_mean(
+ self._obj_loss_fn(obj_targets, layer_pred[..., 4])
+ )
+ obj_loss += layer_obj_loss * self._balance[i]
+ # Updates the balance factor, which is a moving average of previous
+ # factor at the same level.
+ if self._auto_balance:
+ self._balance[i] = self._balance[
+ i
+ ] * 0.9999 + 0.0001 / tf.stop_gradient(layer_obj_loss)
+
+ # Re-balance the factors so that stride at self._ssi always receives 1.
+ if self._auto_balance:
+ self._balance = [x / self._balance[self._ssi] for x in self._balance]
+
+ # Keep separate losses for summary purpose.
+ box_loss *= self._box_weight
+ obj_loss *= self._obj_weight
+ cls_loss *= self._cls_weight
+
+ self._iou = tf.stop_gradient(iou_metric) / self._num_layers
+ self._num_matchings = tf.stop_gradient(
+ total_num_matchings) / tf.cast(batch_size, tf.float32)
+ self._num_gts = total_num_gts / tf.cast(batch_size, tf.float32)
+ self._num_duplicates = tf.stop_gradient(
+ num_duplicates) / tf.cast(batch_size, tf.float32)
+ self._box_loss = tf.stop_gradient(box_loss)
+ self._obj_loss = tf.stop_gradient(obj_loss)
+ self._cls_loss = tf.stop_gradient(cls_loss)
+
+ loss = box_loss + obj_loss + cls_loss
+
+ # Scale up the loss by batch size.
+ return loss * tf.cast(batch_size, loss.dtype)
+
+ def _build_targets(self, labels, predictions, image_info):
+ """Finds the matching targets using Sinkhorn-Knopp."""
+ # Find the three positives matching first for predictions.
+ masks, indices, anchors = self._find_three_positives(labels, predictions)
+
+ batch_size = tf.shape(masks)[0]
+
+ # Collect the predictions.
+ p_box, p_cls, p_obj = [], [], []
+ for layer_key, layer_p in predictions.items():
+ # Always assume level starts from 3.
+ i = int(layer_key) - 3
+ layer_indices = tf.reshape(indices[:, i], [batch_size, -1, 3])
+ anchor = tf.reshape(anchors[:, i], [batch_size, -1, 2])
+
+ fg_pred = tf.gather_nd(layer_p, layer_indices, batch_dims=1)
+
+ grid = tf.stack([layer_indices[..., 2], layer_indices[..., 1]], axis=-1)
+ grid = tf.cast(grid, fg_pred.dtype)
+
+ pxy = (tf.sigmoid(fg_pred[..., :2]) * 2 - 0.5 + grid) * self._strides[i]
+ pwh = (
+ tf.square(tf.sigmoid(fg_pred[..., 2:4]) * 2)
+ * anchor
+ * self._strides[i]
+ )
+ pxywh = tf.concat([pxy, pwh], axis=-1)
+
+ p_box.append(pxywh)
+ p_obj.append(fg_pred[..., 4:5])
+ p_cls.append(fg_pred[..., 5:])
+
+ p_box = tf.concat(p_box, axis=1)
+ p_cls = tf.concat(p_cls, axis=1)
+ p_obj = tf.concat(p_obj, axis=1)
+
+ # Compute valid masks for both targets and predictions.
+ t_mask = labels[..., 0] != -1
+ p_mask = tf.reshape(masks, [batch_size, -1])
+ # [anchor_idx, gj, gi]
+ indices = tf.reshape(indices, [batch_size, -1, 3])
+ anchors = tf.reshape(anchors, [batch_size, -1, 2])
+
+ num_preds = tf.shape(p_box)[1]
+ num_gts = tf.shape(labels)[1]
+
+ # Computes pair-wise IoU.
+ t_box = labels[..., 1:5] * tf.tile(image_info[0, 1], [2])
+
+ pair_wise_iou = box_ops.compute_iou(t_box[:, :, None], p_box[:, None])
+ pair_wise_iou_loss = -tf.math.log(pair_wise_iou + 1e-8)
+
+ # Computes pair-wise class loss.
+ y = tf.sqrt(tf.sigmoid(p_cls) * tf.sigmoid(p_obj))
+ # Add 1e-9 to avoid nan.
+ logits = tf.math.log(y / (1 - y + 1e-9) + 1e-9)
+ logits = tf.tile(logits[:, None], [1, num_gts, 1, 1])
+
+ t_cls = tf.cast(labels[..., 0], tf.int32)
+ class_labels = tf.one_hot(t_cls, self._num_classes, dtype=tf.float32)
+ class_labels = tf.tile(class_labels[:, :, None], [1, 1, num_preds, 1])
+
+ pair_wise_cls_loss = tf.reduce_sum(
+ tf.nn.sigmoid_cross_entropy_with_logits(class_labels, logits), axis=-1
+ )
+
+ # Compute the cost matrix and its corresponding valid mask.
+ cost_mask = tf.logical_and(t_mask[..., None], p_mask[:, None])
+ cost = tf.stop_gradient(pair_wise_cls_loss + 3 * pair_wise_iou_loss)
+ largest_cost = tf.reduce_max(cost)
+
+ # Set invalid IoU to 0.0 for top_k.
+ valid_iou = tf.where(cost_mask, pair_wise_iou, tf.zeros_like(pair_wise_iou))
+
+ # Compute top-10 IoUs from valid IoUs for each target.
+ # When matched predictions is smaller than 10, we only want the top-k where
+ # k is the total size of the matched predictions (k < 10).
+ top_k_mask = tf.less(
+ tf.range(10)[None],
+ tf.minimum(10, tf.reduce_sum(tf.cast(p_mask, tf.int32), axis=-1))[
+ :, None
+ ],
+ )
+ top_k_mask = tf.logical_and(top_k_mask[:, None], t_mask[..., None])
+ top_k, _ = tf.nn.top_k(valid_iou, k=10)
+ top_k = tf.where(top_k_mask, top_k, tf.zeros_like(top_k))
+
+ # Use top_k to compute the dynamic ks for target matching. Each target_i can
+ # match to k_i predictions, and k_i is computed based on the pair-wise
+ # valid IoU.
+ dynamic_ks = tf.maximum(tf.cast(tf.reduce_sum(top_k, axis=-1), tf.int32), 1)
+ dynamic_ks = tf.where(t_mask, dynamic_ks, tf.zeros_like(dynamic_ks))
+ dynamic_ks = tf.stop_gradient(dynamic_ks)
+ dynamic_mask = tf.range(10)[None, None] < dynamic_ks[..., None]
+
+ # Set the invalid field to maximum cost so that they won't be selected
+ # during matching.
+ cost = tf.where(cost_mask, cost, tf.ones_like(cost) * (largest_cost + 1))
+
+ matching_matrix = tf.zeros_like(cost, dtype=tf.int32)
+ _, pred_idx = tf.nn.top_k(-cost, k=10)
+
+ # Update matching matrix.
+ # [batch_size, num_gts, 10]
+ batch_idx = tf.tile(tf.range(batch_size)[:, None, None], [1, num_gts, 10])
+ gt_idx = tf.tile(tf.range(num_gts)[None, :, None], [batch_size, 1, 10])
+ matched_indices = tf.stack([batch_idx, gt_idx, pred_idx], axis=-1)
+ matching_matrix = tf.tensor_scatter_nd_add(
+ matching_matrix,
+ matched_indices,
+ tf.cast(dynamic_mask, matching_matrix.dtype),
+ )
+
+ # Detect if there is a detection matches to multiple targets, if so, we
+ # assign it to the target with minimum cost.
+ duplicate_mask = tf.reduce_sum(matching_matrix, axis=1) > 1
+ num_duplicates = tf.reduce_sum(tf.cast(duplicate_mask, tf.float32))
+ cost_argmin = tf.argmin(cost, axis=1, output_type=tf.int32)
+
+ remove_mask = tf.tile(duplicate_mask[:, None], [1, num_gts, 1])
+ matching_matrix = tf.where(
+ remove_mask, tf.zeros_like(matching_matrix), matching_matrix)
+
+ min_mask = tf.equal(
+ tf.tile(tf.range(num_gts)[None, :, None], [batch_size, 1, num_preds]),
+ cost_argmin[:, None],
+ )
+ update_mask = tf.logical_and(min_mask, duplicate_mask[:, None])
+ matching_matrix = tf.where(
+ update_mask, tf.ones_like(matching_matrix), matching_matrix)
+
+ # Find the final matching and collect the matched targets.
+ matched_gt_indices = tf.argmax(
+ matching_matrix, axis=1, output_type=tf.int32
+ )
+ matched_mask = tf.reduce_sum(matching_matrix, axis=1) > 0
+ matched_targets = tf.gather_nd(
+ labels, matched_gt_indices[..., None], batch_dims=1
+ )
+ return indices, anchors, matched_mask, matched_targets, num_duplicates
+
+ def _find_three_positives(self, labels, predictions):
+ """Finds three matching anchors for each ground-truth."""
+ label_shape = tf.shape(labels)
+ batch_size, max_boxes = label_shape[0], label_shape[1]
+ masks, indices, anch = [], [], []
+ anchor_indices = tf.tile(
+ tf.range(self._num_anchors, dtype=tf.float32)[None, None],
+ [batch_size, max_boxes, 1],
+ )
+ # Append anchor indices to labels.
+ labels = tf.tile(labels[:, :, None], [1, 1, self._num_anchors, 1])
+ labels = tf.concat([labels, anchor_indices[..., None]], axis=-1)
+
+ # Bias is used to determine the matching. 0.5 means matching anchors that
+ # fall in the 0.5 differences in the feature map. For instance, a box
+ # coordinates of (15.6, 35.4) will match the anchors at [15, 35], [16, 35],
+ # and [15, 34].
+ bias = 0.5 # bias
+ off = (
+ tf.constant(
+ [
+ [0, 0],
+ [1, 0], [0, 1], [-1, 0], [0, -1], # j, k, l, m
+ # [1, 1], [1, -1], [-1, 1], [-1, -1], # jk,jm,lk,lm
+ ],
+ tf.float32,
+ )
+ * bias
+ ) # offsets
+
+ for i in range(self._num_layers):
+ anchors = self._anchors[i]
+ _, _, h, w, _ = predictions[str(i + 3)].get_shape().as_list()
+ gain = tf.constant([1, w, h, w, h, 1], dtype=tf.float32)
+
+ t = labels * gain
+
+ # Filter out targets that do not match the current anchors.
+ wh_ratio = t[..., 3:5] / tf.cast(anchors[None, None], tf.float32)
+ labels_mask = tf.less(
+ tf.reduce_max(tf.maximum(wh_ratio, 1.0 / wh_ratio), axis=-1),
+ self._anchor_threshold,
+ )[..., None]
+ # Compute valid mask for ground-truths.
+ labels_mask = tf.logical_and(t[..., :1] != -1, labels_mask)
+
+ labels_mask = tf.reshape(labels_mask, [batch_size, -1])
+ t = tf.reshape(t, [batch_size, -1, 6])
+
+ # Find the matching offsets for valid labels.
+ gxy = t[..., 1:3] # grid xy
+ gxi = gain[1:3] - gxy # inverse
+ j, k = tf.split((gxy % 1.0 < bias) & (gxy >= 1.0), 2, axis=-1)
+ l, m = tf.split((gxi % 1.0 < bias) & (gxi >= 1.0), 2, axis=-1)
+
+ j, k, l, m = j[..., 0], k[..., 0], l[..., 0], m[..., 0]
+
+ # Note that j and l, k and m are conjugate to each other, so at most one
+ # of them will be True during running. Therefore, we can reduce memory
+ # usage by gathering the selected index.
+ x_map = tf.cast(tf.stack([j, l], axis=-1), tf.int8)
+ y_map = tf.cast(tf.stack([k, m], axis=-1), tf.int8)
+
+ # Add the indices offsets.
+ x_indices = tf.argmax(x_map, axis=-1) * 2 + 1
+ y_indices = tf.argmax(y_map, axis=-1) * 2 + 2
+ three_targets_indices = tf.stack(
+ [tf.zeros_like(x_indices), x_indices, y_indices], axis=-1
+ )[..., None]
+
+ # Gather the selected 3 targets from the 5-target map.
+ j = tf.stack([tf.ones_like(j), j, k, l, m], axis=-1)
+ three_targets_mask = tf.gather_nd(j, three_targets_indices, batch_dims=2)
+
+ labels_mask = tf.tile(labels_mask[:, :, None], [1, 1, 5])
+ t = tf.tile(t[:, :, None], [1, 1, 5, 1])
+
+ labels_mask = tf.gather_nd(
+ labels_mask, three_targets_indices, batch_dims=2
+ )
+ t = tf.gather_nd(t, three_targets_indices, batch_dims=2)
+
+ offsets = tf.zeros_like(gxy)[:, :, None] + off[None, None]
+ offsets = tf.gather_nd(offsets, three_targets_indices, batch_dims=2)
+
+ gxy = t[..., 1:3]
+ # Find the actual grid locations.
+ gij = tf.cast(gxy - offsets * 2, tf.int32)
+ gi, gj = tf.split(gij, 2, axis=-1)
+ gi, gj = gi[..., 0], gj[..., 0]
+
+ # Append the result.
+ anchor_idx = tf.cast(t[..., 5], tf.int32)
+ gain = tf.cast(gain, tf.int32)
+ gi = tf.clip_by_value(gi, 0, gain[2] - 1)
+ gj = tf.clip_by_value(gj, 0, gain[3] - 1)
+
+ labels_mask = tf.logical_and(labels_mask, three_targets_mask)
+ masks.append(labels_mask)
+ indices.append(tf.stack([anchor_idx, gj, gi], axis=-1))
+ anch.append(tf.gather(anchors, anchor_idx))
+
+ # [batch_size, num_layers, num_anchors * max_boxes, num_targets]
+ masks = tf.stack(masks, axis=1)
+ indices = tf.stack(indices, axis=1)
+ anch = tf.stack(anch, axis=1)
+ return masks, indices, anch
+
+ def report_stats(self):
+ return {
+ 'num_gts': self._num_gts,
+ 'num_matchings': self._num_matchings,
+ 'num_duplicates': self._num_duplicates,
+ }
+
+ def report_separate_losses(self):
+ """Returns separate losses that construct the reported loss."""
+ return {
+ 'iou': self._iou,
+ 'box_loss': self._box_loss,
+ 'obj_loss': self._obj_loss,
+ 'cls_loss': self._cls_loss,
+ }
+
+ def get_config(self):
+ """Configs for the loss constructor."""
+ config = {
+ 'alpha': self._alpha,
+ 'gamma': self._gamma,
+ 'box_weight': self._box_weight,
+ 'obj_weight': self._obj_weight,
+ 'cls_weight': self._cls_weight,
+ 'iou_weight': self._iou_weight,
+ 'iou_mix_ratio': self._iou_mix_ratio,
+ 'pos_targets': self._pos_targets,
+ 'neg_targets': self._neg_targets,
+ 'num_classes': self._num_classes,
+ 'num_layers': self._num_layers,
+ 'num_anchors': self._num_anchors,
+ 'auto_balance': self._auto_balance,
+ 'balance': self._balance,
+ 'strides': self._strides,
+ 'anchors': self._anchors,
+ 'input_size': self._input_size,
+ 'anchor_threshold': self._anchor_threshold,
+ }
+ base_config = super().get_config()
+ return dict(list(base_config.items()) + list(config.items()))
diff --git a/official/projects/yolo/losses/yolov7_loss_test.py b/official/projects/yolo/losses/yolov7_loss_test.py
new file mode 100644
index 00000000000..0dfe2083db5
--- /dev/null
+++ b/official/projects/yolo/losses/yolov7_loss_test.py
@@ -0,0 +1,144 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for yolo heads."""
+
+from absl import logging
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.yolo.losses import yolov7_loss
+from official.projects.yolo.ops import box_ops
+
+_HEIGHT, _WIDTH = 640, 640
+_BATCH_SIZE = 8
+_NUM_GTS = 100
+_NUM_LAYERS, _NUM_ANCHORS = 3, 3
+_NUM_CLASSES = 80
+
+
+def build_labels():
+ image_info = tf.constant(
+ [
+ [[640, 640], [640, 640], [1.0, 1.0], [0.0, 0.0]]
+ for _ in range(_BATCH_SIZE)
+ ], dtype=tf.float32
+ )
+ box_y1x1 = np.random.rand(_BATCH_SIZE, _NUM_GTS, 2).astype(np.float32)
+ box_y2x2 = (
+ np.random.rand(_BATCH_SIZE, _NUM_GTS, 2).astype(np.float32)
+ * (1 - box_y1x1)
+ + box_y1x1
+ )
+ boxes_yxyx = tf.concat([box_y1x1, box_y2x2], axis=-1)
+ num_detections = np.random.randint(_NUM_GTS, size=[_BATCH_SIZE])
+ classes = np.arange(_NUM_GTS * _BATCH_SIZE).reshape([_BATCH_SIZE, -1])
+ for i in range(_BATCH_SIZE):
+ classes[i, num_detections[i]:] = -1
+ classes = tf.constant(classes, dtype=tf.int32)
+ return {'image_info': image_info, 'classes': classes, 'bbox': boxes_yxyx}
+
+
+def build_predictions():
+ # Scale down by 2^3 because prediction outputs start at level 3.
+ h, w = _HEIGHT // 8, _WIDTH // 8
+
+ predictions = {}
+ for i in range(_NUM_LAYERS):
+ shape = [_BATCH_SIZE, h // (2**i), w // (2**i), _NUM_ANCHORS]
+ p_y1x1 = tf.constant(np.random.rand(*shape, 2), dtype=tf.float32)
+ p_y2x2 = tf.constant(np.random.rand(*shape, 2), dtype=tf.float32)
+ # Transform the box from yxyx to xywh.
+ p_box = box_ops.yxyx_to_xcycwh(tf.concat([p_y1x1, p_y2x2], axis=-1))
+ p_obj = tf.constant(np.random.rand(*shape, 1), dtype=tf.float32)
+ p_cls = tf.constant(np.random.rand(*shape, _NUM_CLASSES), dtype=tf.float32)
+ predictions[str(i + 3)] = tf.concat([p_box, p_obj, p_cls], axis=-1)
+
+ return predictions
+
+
+class YoloV7LossTest(parameterized.TestCase, tf.test.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ np.random.seed(42)
+ self._anchors = [
+ [[12, 16], [19, 36], [40, 28]], # Level 3
+ [[36, 75], [76, 55], [72, 146]], # Level 4
+ [[142, 110], [192, 243], [459, 401]], # Level 5
+ ]
+ self._strides = [8, 16, 32]
+
+ @parameterized.product(
+ gamma=(0.0, 1.5), label_smoothing=(0.0, 0.2), auto_balance=(True, False)
+ )
+ def test_loss(self, gamma, label_smoothing, auto_balance):
+ """Test YOLOv7 normal loss."""
+ labels = build_labels()
+ predictions = build_predictions()
+ loss = yolov7_loss.YoloV7Loss(
+ anchors=self._anchors,
+ strides=self._strides,
+ input_size=[_HEIGHT, _WIDTH],
+ gamma=gamma,
+ label_smoothing=label_smoothing,
+ num_classes=_NUM_CLASSES,
+ auto_balance=auto_balance,
+ )
+
+ loss_val = loss(labels, predictions)
+ losses = loss.report_separate_losses()
+ logging.info('loss_val: %.6f', loss_val)
+ logging.info('box_loss: %.6f', losses['box_loss'])
+ logging.info('obj_loss: %.6f', losses['obj_loss'])
+ logging.info('cls_loss: %.6f', losses['cls_loss'])
+
+ expected_loss_val = (
+ losses['box_loss'] + losses['obj_loss'] + losses['cls_loss']
+ ) * _BATCH_SIZE
+ self.assertNear(loss_val, expected_loss_val, err=1e-6)
+
+ @parameterized.product(
+ gamma=(0.0, 1.5), label_smoothing=(0.0, 0.2), auto_balance=(True, False)
+ )
+ def test_loss_ota(self, gamma, label_smoothing, auto_balance):
+ """Test YOLOv7 OTA loss."""
+ labels = build_labels()
+ predictions = build_predictions()
+ loss = yolov7_loss.YoloV7LossOTA(
+ anchors=self._anchors,
+ strides=self._strides,
+ input_size=[_HEIGHT, _WIDTH],
+ gamma=gamma,
+ label_smoothing=label_smoothing,
+ num_classes=_NUM_CLASSES,
+ auto_balance=auto_balance,
+ )
+
+ loss_val = loss(labels, predictions)
+ losses = loss.report_separate_losses()
+ logging.info('loss_val: %.6f', loss_val)
+ logging.info('box_loss: %.6f', losses['box_loss'])
+ logging.info('obj_loss: %.6f', losses['obj_loss'])
+ logging.info('cls_loss: %.6f', losses['cls_loss'])
+
+ expected_loss_val = (
+ losses['box_loss'] + losses['obj_loss'] + losses['cls_loss']
+ ) * _BATCH_SIZE
+ self.assertNear(loss_val, expected_loss_val, err=1e-6)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/yolo/modeling/__init__.py b/official/projects/yolo/modeling/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/modeling/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/yolo/modeling/backbones/__init__.py b/official/projects/yolo/modeling/backbones/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/modeling/backbones/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/modeling/backbones/darknet.py b/official/projects/yolo/modeling/backbones/darknet.py
similarity index 88%
rename from official/vision/beta/projects/yolo/modeling/backbones/darknet.py
rename to official/projects/yolo/modeling/backbones/darknet.py
index 4f19379875e..b62a9bbb0d2 100644
--- a/official/vision/beta/projects/yolo/modeling/backbones/darknet.py
+++ b/official/projects/yolo/modeling/backbones/darknet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -37,10 +37,10 @@
import collections
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
-from official.vision.beta.projects.yolo.modeling.layers import nn_blocks
+from official.projects.yolo.modeling.layers import nn_blocks
from official.vision.modeling.backbones import factory
@@ -104,7 +104,7 @@ class LayerBuilder:
def __init__(self):
self._layer_dict = {
'ConvBN': (nn_blocks.ConvBN, self.conv_bn_config_todict),
- 'MaxPool': (tf.keras.layers.MaxPool2D, self.maxpool_config_todict)
+ 'MaxPool': (tf_keras.layers.MaxPool2D, self.maxpool_config_todict)
}
def conv_bn_config_todict(self, config, kwargs):
@@ -372,13 +372,13 @@ def __call__(self, config, kwargs):
}
-class Darknet(tf.keras.Model):
+class Darknet(tf_keras.Model):
"""The Darknet backbone architecture."""
def __init__(
self,
model_id='darknet53',
- input_specs=tf.keras.layers.InputSpec(shape=[None, None, None, 3]),
+ input_specs=tf_keras.layers.InputSpec(shape=[None, None, None, 3]),
min_level=None,
max_level=5,
width_scale=1.0,
@@ -400,7 +400,7 @@ def __init__(
self._model_name = model_id
self._splits = splits
- self._input_shape = input_specs
+ self._input_specs = input_specs
self._registry = LayerBuilder()
# default layer look up
@@ -435,13 +435,11 @@ def __init__(
'name': None
}
- inputs = tf.keras.layers.Input(shape=self._input_shape.shape[1:])
+ inputs = tf_keras.Input(shape=input_specs.shape[1:])
output = self._build_struct(layer_specs, inputs)
- super().__init__(inputs=inputs, outputs=output, name=self._model_name)
-
- @property
- def input_specs(self):
- return self._input_shape
+ super().__init__(
+ inputs=inputs, outputs=output, name=self._model_name, **kwargs
+ )
@property
def output_specs(self):
@@ -490,11 +488,11 @@ def _build_struct(self, net, inputs):
stack_outputs[config.route], config, name=f'{config.layer}_{i}')
stack_outputs.append(x)
if (config.is_output and self._min_size is None):
- endpoints[str(config.output_name)] = x
+ endpoints[str(config.output_name)] = x # pyrefly: ignore[unbound-name]
elif (self._min_size is not None and
config.output_name >= self._min_size and
config.output_name <= self._max_size):
- endpoints[str(config.output_name)] = x
+ endpoints[str(config.output_name)] = x # pyrefly: ignore[unbound-name]
self._output_specs = {l: endpoints[l].get_shape() for l in endpoints.keys()}
return endpoints
@@ -514,7 +512,7 @@ def _csp_stack(self, inputs, config, name):
residual_filter_scale = 1
scale_filters = 2
self._default_dict['activation'] = self._get_activation(config.activation)
- self._default_dict['name'] = f'{name}_csp_down'
+ self._default_dict['name'] = f'{name}_csp_down' # pyrefly: ignore[bad-assignment]
if self._dilate:
self._default_dict['dilation_rate'] = config.dilation_rate
degrid = int(tf.math.log(float(config.dilation_rate)) / tf.math.log(2.))
@@ -532,36 +530,39 @@ def _csp_stack(self, inputs, config, name):
dilated_reps = config.repetitions - degrid
for i in range(dilated_reps):
- self._default_dict['name'] = f'{name}_{i}'
+ self._default_dict['name'] = f'{name}_{i}' # pyrefly: ignore[bad-assignment]
x = nn_blocks.DarkResidual(
filters=config.filters // scale_filters,
filter_scale=residual_filter_scale,
- **self._default_dict)(
- x)
+ **self._default_dict,
+ )(x)
for i in range(dilated_reps, config.repetitions):
self._default_dict['dilation_rate'] = max(
- 1, self._default_dict['dilation_rate'] // 2)
- self._default_dict[
- 'name'] = f"{name}_{i}_degridded_{self._default_dict['dilation_rate']}"
+ 1, self._default_dict['dilation_rate'] // 2
+ )
+ self._default_dict['name'] = (
+ f"{name}_{i}_degridded_{self._default_dict['dilation_rate']}" # pyrefly: ignore[bad-assignment]
+ )
x = nn_blocks.DarkResidual(
filters=config.filters // scale_filters,
filter_scale=residual_filter_scale,
- **self._default_dict)(
- x)
+ **self._default_dict,
+ )(x)
- self._default_dict['name'] = f'{name}_csp_connect'
+ self._default_dict['name'] = f'{name}_csp_connect' # pyrefly: ignore[bad-assignment]
output = nn_blocks.CSPConnect(
filters=config.filters,
filter_scale=csp_filter_scale,
- **self._default_dict)([x, x_route])
+ **self._default_dict,
+ )([x, x_route])
self._default_dict['activation'] = self._activation
self._default_dict['name'] = None
return output
def _csp_tiny_stack(self, inputs, config, name):
self._default_dict['activation'] = self._get_activation(config.activation)
- self._default_dict['name'] = f'{name}_csp_tiny'
+ self._default_dict['name'] = f'{name}_csp_tiny' # pyrefly: ignore[bad-assignment]
x, x_route = nn_blocks.CSPTiny(
filters=config.filters, **self._default_dict)(
inputs)
@@ -570,7 +571,7 @@ def _csp_tiny_stack(self, inputs, config, name):
return x, x_route
def _tiny_stack(self, inputs, config, name):
- x = tf.keras.layers.MaxPool2D(
+ x = tf_keras.layers.MaxPool2D(
pool_size=2,
strides=config.strides,
padding='same',
@@ -578,7 +579,7 @@ def _tiny_stack(self, inputs, config, name):
name=f'{name}_tiny/pool')(
inputs)
self._default_dict['activation'] = self._get_activation(config.activation)
- self._default_dict['name'] = f'{name}_tiny/conv'
+ self._default_dict['name'] = f'{name}_tiny/conv' # pyrefly: ignore[bad-assignment]
x = nn_blocks.ConvBN(
filters=config.filters,
kernel_size=(3, 3),
@@ -592,7 +593,7 @@ def _tiny_stack(self, inputs, config, name):
def _residual_stack(self, inputs, config, name):
self._default_dict['activation'] = self._get_activation(config.activation)
- self._default_dict['name'] = f'{name}_residual_down'
+ self._default_dict['name'] = f'{name}_residual_down' # pyrefly: ignore[bad-assignment]
if self._dilate:
self._default_dict['dilation_rate'] = config.dilation_rate
if config.repetitions < 8:
@@ -601,25 +602,28 @@ def _residual_stack(self, inputs, config, name):
self._default_dict['dilation_rate'] = 1
x = nn_blocks.DarkResidual(
- filters=config.filters, downsample=True, **self._default_dict)(
- inputs)
+ filters=config.filters, downsample=True, **self._default_dict
+ )(inputs)
- dilated_reps = config.repetitions - self._default_dict[
- 'dilation_rate'] // 2 - 1
+ dilated_reps = (
+ config.repetitions - self._default_dict['dilation_rate'] // 2 - 1
+ )
for i in range(dilated_reps):
- self._default_dict['name'] = f'{name}_{i}'
- x = nn_blocks.DarkResidual(
- filters=config.filters, **self._default_dict)(
- x)
+ self._default_dict['name'] = f'{name}_{i}' # pyrefly: ignore[bad-assignment]
+ x = nn_blocks.DarkResidual(filters=config.filters, **self._default_dict)(
+ x
+ )
for i in range(dilated_reps, config.repetitions - 1):
- self._default_dict[
- 'dilation_rate'] = self._default_dict['dilation_rate'] // 2
- self._default_dict[
- 'name'] = f"{name}_{i}_degridded_{self._default_dict['dilation_rate']}"
- x = nn_blocks.DarkResidual(
- filters=config.filters, **self._default_dict)(
- x)
+ self._default_dict['dilation_rate'] = (
+ self._default_dict['dilation_rate'] // 2
+ )
+ self._default_dict['name'] = (
+ f"{name}_{i}_degridded_{self._default_dict['dilation_rate']}" # pyrefly: ignore[bad-assignment]
+ )
+ x = nn_blocks.DarkResidual(filters=config.filters, **self._default_dict)(
+ x
+ )
self._default_dict['activation'] = self._activation
self._default_dict['name'] = None
@@ -631,7 +635,7 @@ def _build_block(self, inputs, config, name):
i = 0
self._default_dict['activation'] = self._get_activation(config.activation)
while i < config.repetitions:
- self._default_dict['name'] = f'{name}_{i}'
+ self._default_dict['name'] = f'{name}_{i}' # pyrefly: ignore[bad-assignment]
layer = self._registry(config, self._default_dict)
x = layer(x)
i += 1
@@ -672,11 +676,11 @@ def get_config(self):
@factory.register_backbone_builder('darknet')
def build_darknet(
- input_specs: tf.keras.layers.InputSpec,
+ input_specs: tf_keras.layers.InputSpec,
backbone_config: hyperparams.Config,
norm_activation_config: hyperparams.Config,
- l2_regularizer: tf.keras.regularizers.Regularizer = None
-) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ l2_regularizer: tf_keras.regularizers.Regularizer = None # pyrefly: ignore[bad-function-definition]
+) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
"""Builds darknet."""
backbone_config = backbone_config.get()
diff --git a/official/vision/beta/projects/yolo/modeling/backbones/darknet_test.py b/official/projects/yolo/modeling/backbones/darknet_test.py
similarity index 85%
rename from official/vision/beta/projects/yolo/modeling/backbones/darknet_test.py
rename to official/projects/yolo/modeling/backbones/darknet_test.py
index 00e7352f37c..81418b915ff 100644
--- a/official/vision/beta/projects/yolo/modeling/backbones/darknet_test.py
+++ b/official/projects/yolo/modeling/backbones/darknet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,11 +16,11 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
-from official.vision.beta.projects.yolo.modeling.backbones import darknet
+from official.projects.yolo.modeling.backbones import darknet
class DarknetTest(parameterized.TestCase, tf.test.TestCase):
@@ -34,13 +34,13 @@ class DarknetTest(parameterized.TestCase, tf.test.TestCase):
def test_network_creation(self, input_size, model_id, endpoint_filter_scale,
scale_final, dilate):
"""Test creation of ResNet family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
network = darknet.Darknet(
model_id=model_id, min_level=3, max_level=5, dilate=dilate)
self.assertEqual(network.model_id, model_id)
- inputs = tf.keras.Input(shape=(input_size, input_size, 3), batch_size=1)
+ inputs = tf_keras.Input(shape=(input_size, input_size, 3), batch_size=1)
endpoints = network(inputs)
if dilate:
@@ -78,22 +78,27 @@ def test_sync_bn_multiple_devices(self, strategy, use_sync_bn):
"""Test for sync bn on TPU and GPU devices."""
inputs = np.random.rand(1, 224, 224, 3)
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
with strategy.scope():
- network = darknet.Darknet(model_id='darknet53', min_size=3, max_size=5)
+ network = darknet.Darknet(
+ model_id='darknet53',
+ min_level=3,
+ max_level=5,
+ use_sync_bn=use_sync_bn,
+ )
_ = network(inputs)
@parameterized.parameters(1, 3, 4)
def test_input_specs(self, input_dim):
"""Test different input feature dimensions."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
- input_specs = tf.keras.layers.InputSpec(shape=[None, None, None, input_dim])
+ input_specs = tf_keras.layers.InputSpec(shape=[None, None, None, input_dim])
network = darknet.Darknet(
model_id='darknet53', min_level=3, max_level=5, input_specs=input_specs)
- inputs = tf.keras.Input(shape=(224, 224, input_dim), batch_size=1)
+ inputs = tf_keras.Input(shape=(224, 224, input_dim), batch_size=1)
_ = network(inputs)
def test_serialize_deserialize(self):
diff --git a/official/projects/yolo/modeling/backbones/yolov7.py b/official/projects/yolo/modeling/backbones/yolov7.py
new file mode 100644
index 00000000000..df88ec18c9e
--- /dev/null
+++ b/official/projects/yolo/modeling/backbones/yolov7.py
@@ -0,0 +1,455 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains backbone architectures for YOLOv7 families.
+
+The models are built with ELAN and E-ELAN.
+
+ELAN was proposed in:
+[1] Wang, Chien-Yao and Liao, Hong-Yuan Mark and Yeh, I-Hau
+ Designing Network Design Strategies Through Gradient Path Analysis
+ arXiv:2211.04800
+
+E-ELAN is proposed in YOLOv7 paper:
+[1] Wang, Chien-Yao and Bochkovskiy, Alexey and Liao, Hong-Yuan Mark
+ YOLOv7: Trainable bag-of-freebies sets new state-of-the-art for real-time
+ object detectors
+ arXiv:2207.02696
+"""
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import hyperparams
+from official.projects.yolo.modeling.layers import nn_blocks
+from official.projects.yolo.ops import initializer_ops
+from official.vision.modeling.backbones import factory
+
+# Required block functions for YOLOv7 backbone familes.
+_BLOCK_FNS = {
+ 'convbn': nn_blocks.ConvBN,
+ 'maxpool2d': tf_keras.layers.MaxPooling2D,
+ 'concat': tf_keras.layers.Concatenate,
+}
+
+# Names for key arguments needed by each block function.
+_BLOCK_SPEC_SCHEMAS = {
+ 'convbn': [
+ 'block_fn',
+ 'from',
+ 'kernel_size',
+ 'strides',
+ 'filters',
+ 'is_output',
+ ],
+ 'maxpool2d': [
+ 'block_fn',
+ 'from',
+ 'pool_size',
+ 'strides',
+ 'padding',
+ 'is_output',
+ ],
+ 'concat': [
+ 'block_fn',
+ 'from',
+ 'axis',
+ 'is_output',
+ ]
+}
+
+# Define YOLOv7-pico variant.
+_YoloV7Pico = [
+ ['convbn', -1, 3, 2, 8, False], # 0-P1/2
+
+ ['convbn', -1, 3, 2, 16, False], # 1-P2/4
+
+ ['convbn', -1, 1, 1, 8, False],
+ ['convbn', -2, 1, 1, 8, False],
+ ['convbn', -1, 3, 1, 8, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 16, False], # 7
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 8-P3/8
+ ['convbn', -1, 1, 1, 16, False],
+ ['convbn', -2, 1, 1, 16, False],
+ ['convbn', -1, 3, 1, 16, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 32, True], # 14
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 15-P4/16
+ ['convbn', -1, 1, 1, 32, False],
+ ['convbn', -2, 1, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 64, True], # 21
+]
+
+# Define YOLOv7-nano variant.
+
+_YoloV7Nano = [
+ ['convbn', -1, 3, 2, 8, False], # 0-P1/2
+
+ ['convbn', -1, 3, 2, 16, False], # 1-P2/4
+
+ ['convbn', -1, 1, 1, 8, False],
+ ['convbn', -2, 1, 1, 8, False],
+ ['convbn', -1, 3, 1, 8, False],
+ ['convbn', -1, 3, 1, 8, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 16, False], # 7
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 8-P3/8
+ ['convbn', -1, 1, 1, 16, False],
+ ['convbn', -2, 1, 1, 16, False],
+ ['convbn', -1, 3, 1, 16, False],
+ ['convbn', -1, 3, 1, 16, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 32, True], # 14
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 15-P4/16
+ ['convbn', -1, 1, 1, 32, False],
+ ['convbn', -2, 1, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 64, True], # 21
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 22-P5/32
+ ['convbn', -1, 1, 1, 64, False],
+ ['convbn', -2, 1, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 128, True], # 28
+]
+
+# Define YOLOv7-tiny variant.
+_YoloV7Tiny = [
+ ['convbn', -1, 3, 2, 32, False], # 0-P1/2
+
+ ['convbn', -1, 3, 2, 64, False], # 1-P2/4
+
+ ['convbn', -1, 1, 1, 32, False],
+ ['convbn', -2, 1, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 64, False], # 7
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 8-P3/8
+ ['convbn', -1, 1, 1, 64, False],
+ ['convbn', -2, 1, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 128, True], # 14
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 15-P4/16
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -2, 1, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 256, True], # 21
+
+ ['maxpool2d', -1, 2, 2, 'same', False], # 22-P5/32
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 512, True], # 28
+]
+
+# Define YOLOv7 variant.
+_YoloV7 = [
+ ['convbn', -1, 3, 1, 32, False], # 0
+
+ ['convbn', -1, 3, 2, 64, False], # 1-P1/2
+ ['convbn', -1, 3, 1, 64, False],
+
+ ['convbn', -1, 3, 2, 128, False], # 3-P2/4
+ ['convbn', -1, 1, 1, 64, False],
+ ['convbn', -2, 1, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['concat', [-1, -3, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 256, False], # 11
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -3, 1, 1, 128, False],
+ ['convbn', -1, 3, 2, 128, False],
+ ['concat', [-1, -3], -1, False], # 16-P3/8
+
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -2, 1, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['concat', [-1, -3, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 512, True], # 24
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -3, 1, 1, 256, False],
+ ['convbn', -1, 3, 2, 256, False],
+ ['concat', [-1, -3], -1, False], # 29-P4/16
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -3, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 1024, True], # 37
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 512, False],
+ ['convbn', -3, 1, 1, 512, False],
+ ['convbn', -1, 3, 2, 512, False],
+ ['concat', [-1, -3], -1, False], # 42-P5/32
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -3, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 1024, True], # 50
+]
+
+_YoloV7X = [
+ ['convbn', -1, 3, 1, 40, False], # 0
+
+ ['convbn', -1, 3, 2, 80, False], # 1-P1/2
+ ['convbn', -1, 3, 1, 80, False],
+
+ ['convbn', -1, 3, 2, 160, False], # 3-P2/4
+ ['convbn', -1, 1, 1, 64, False],
+ ['convbn', -2, 1, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 320, False], # 13
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 160, False],
+ ['convbn', -3, 1, 1, 160, False],
+ ['convbn', -1, 3, 2, 160, False],
+ ['concat', [-1, -3], -1, False], # 18-P3/8
+
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -2, 1, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 640, True], # 28
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 320, False],
+ ['convbn', -3, 1, 1, 320, False],
+ ['convbn', -1, 3, 2, 320, False],
+ ['concat', [-1, -3], -1, False], # 33-P4/16
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 1280, True], # 43
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 640, False],
+ ['convbn', -3, 1, 1, 640, False],
+ ['convbn', -1, 3, 2, 640, False],
+ ['concat', [-1, -3], -1, False], # 48-P5/32
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 1280, True], # 58
+]
+
+# Aggregates all variants for YOLOv7 backbones.
+BACKBONES = {
+ 'yolov7-nano': _YoloV7Nano,
+ 'yolov7-pico': _YoloV7Pico,
+ 'yolov7-tiny': _YoloV7Tiny,
+ 'yolov7': _YoloV7,
+ 'yolov7x': _YoloV7X,
+}
+
+
+class YoloV7(tf_keras.Model):
+ """YOLOv7 backbone architecture."""
+
+ def __init__(
+ self,
+ model_id='yolov7',
+ input_specs=tf_keras.layers.InputSpec(shape=[None, None, None, 3]),
+ use_sync_bn=False,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ activation='swish',
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_initializer='zeros',
+ bias_regularizer=None,
+ **kwargs):
+ """Initializes the YOLOv7 backbone.
+
+ Args:
+ model_id: a `str` represents the model variants.
+ input_specs: a `tf_keras.layers.InputSpec` of the input tensor.
+ use_sync_bn: if set to `True`, use synchronized batch normalization.
+ norm_momentum: a `float` of normalization momentum for the moving average.
+ norm_epsilon: a small `float` added to variance to avoid dividing by zero.
+ activation: a `str` name of the activation function.
+ kernel_initializer: a `str` for kernel initializer of convolutional
+ layers.
+ kernel_regularizer: a `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default to None.
+ bias_initializer: a `str` for bias initializer of convolutional layers.
+ bias_regularizer: a `tf_keras.regularizers.Regularizer` object for Conv2D.
+ Default to None.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+
+ self._model_id = model_id
+ self._input_specs = input_specs
+ self._use_sync_bn = use_sync_bn
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._activation = activation
+
+ self._kernel_initializer = initializer_ops.pytorch_kernel_initializer(
+ kernel_initializer
+ )
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_initializer = bias_initializer
+ self._bias_regularizer = bias_regularizer
+
+ inputs = tf_keras.layers.Input(shape=input_specs.shape[1:])
+
+ block_specs = BACKBONES[model_id.lower()]
+ outputs = []
+ endpoints = {}
+ level = 3
+ for spec in block_specs:
+ block_kwargs = dict(zip(_BLOCK_SPEC_SCHEMAS[spec[0]], spec)) # pyrefly: ignore[bad-index]
+
+ block_fn_str = block_kwargs.pop('block_fn')
+ from_index = block_kwargs.pop('from')
+ is_output = block_kwargs.pop('is_output')
+
+ if not outputs:
+ x = inputs
+ elif isinstance(from_index, int):
+ x = outputs[from_index]
+ else:
+ x = [outputs[idx] for idx in from_index] # pyrefly: ignore[bad-index]
+
+ if block_fn_str in ['convbn']:
+ block_kwargs.update({ # pyrefly: ignore[no-matching-overload]
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'activation': self._activation,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_initializer': self._bias_initializer,
+ 'bias_regularizer': self._bias_regularizer,
+ })
+ block_fn = _BLOCK_FNS[block_fn_str](**block_kwargs) # pyrefly: ignore[bad-index]
+
+ x = block_fn(x)
+ outputs.append(x)
+ if is_output:
+ endpoints[str(level)] = x
+ level += 1
+ self._output_specs = {k: v.get_shape() for k, v in endpoints.items()}
+ super().__init__(inputs=inputs, outputs=endpoints, **kwargs)
+
+ def get_config(self):
+ config_dict = {
+ 'model_id': self._model_id,
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'activation': self._activation,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_initializer': self._bias_initializer,
+ 'bias_regularizer': self._bias_regularizer,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+ @property
+ def output_specs(self):
+ """A dict of {level: TensorShape} pairs for the model output."""
+ return self._output_specs
+
+
+@factory.register_backbone_builder('yolov7')
+def build_yolov7(
+ input_specs: tf_keras.layers.InputSpec,
+ backbone_config: hyperparams.Config,
+ norm_activation_config: hyperparams.Config,
+ l2_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
+) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ """Builds YOLOv7."""
+
+ assert backbone_config.type == 'yolov7', (
+ f'Inconsistent backbone type {backbone_config.type}.')
+ backbone_config = backbone_config.get()
+ assert backbone_config.model_id in BACKBONES, (
+ f'Unsupported backbone {backbone_config.model_id}.')
+ model = YoloV7(
+ model_id=backbone_config.model_id,
+ input_specs=input_specs,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ activation=norm_activation_config.activation,
+ kernel_regularizer=l2_regularizer,
+ )
+ return model
diff --git a/official/projects/yolo/modeling/backbones/yolov7_test.py b/official/projects/yolo/modeling/backbones/yolov7_test.py
new file mode 100644
index 00000000000..098fe0d9197
--- /dev/null
+++ b/official/projects/yolo/modeling/backbones/yolov7_test.py
@@ -0,0 +1,92 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for yolov7 backbone."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from tensorflow.python.distribute import combinations
+from tensorflow.python.distribute import strategy_combinations
+from official.projects.yolo.modeling.backbones import yolov7
+
+_INPUT_SIZE = (224, 224)
+
+
+class YoloV7BackboneTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ('yolov7',),
+ )
+ def test_network_creation(self, model_id):
+ """Tests declaration of YOLOv7 backbone variants."""
+ tf_keras.backend.set_image_data_format('channels_last')
+
+ network = yolov7.YoloV7(model_id)
+ self.assertEqual(network.get_config()['model_id'], model_id)
+
+ inputs = tf_keras.Input(shape=(*_INPUT_SIZE, 3), batch_size=1)
+ outputs = network(inputs)
+
+ for level, level_output in outputs.items():
+ scale = 2**int(level)
+ input_size = (_INPUT_SIZE[0] // scale, _INPUT_SIZE[1] // scale)
+ self.assertAllEqual((1, *input_size), level_output.shape.as_list()[:-1])
+
+ @combinations.generate(
+ combinations.combine(
+ strategy=[
+ strategy_combinations.cloud_tpu_strategy,
+ strategy_combinations.one_device_strategy_gpu,
+ ],
+ )
+ )
+ def test_sync_bn_multiple_devices(self, strategy):
+ """Test for sync bn on TPU and GPU devices."""
+ inputs = np.random.rand(1, *_INPUT_SIZE, 3)
+
+ tf_keras.backend.set_image_data_format('channels_last')
+
+ with strategy.scope():
+ network = yolov7.YoloV7(model_id='yolov7')
+ _ = network(inputs)
+
+ def test_serialize_deserialize(self):
+ # Create a network object that sets all of its config options.
+ kwargs = dict(
+ model_id='yolov7',
+ use_sync_bn=False,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ activation='swish',
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_initializer='zeros',
+ bias_regularizer=None,
+ )
+ network = yolov7.YoloV7(**kwargs)
+
+ # Create another network object from the first object's config.
+ new_network = yolov7.YoloV7.from_config(network.get_config())
+
+ # Validate that the config can be forced to JSON.
+ _ = new_network.to_json()
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(network.get_config(), new_network.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/yolo/modeling/decoders/__init__.py b/official/projects/yolo/modeling/decoders/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/modeling/decoders/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/modeling/decoders/yolo_decoder.py b/official/projects/yolo/modeling/decoders/yolo_decoder.py
similarity index 95%
rename from official/vision/beta/projects/yolo/modeling/decoders/yolo_decoder.py
rename to official/projects/yolo/modeling/decoders/yolo_decoder.py
index 3aa0dfa44b2..0c9f1364967 100644
--- a/official/vision/beta/projects/yolo/modeling/decoders/yolo_decoder.py
+++ b/official/projects/yolo/modeling/decoders/yolo_decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,10 @@
"""Feature Pyramid Network and Path Aggregation variants used in YOLO."""
from typing import Mapping, Optional, Union
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import hyperparams
-from official.vision.beta.projects.yolo.modeling.layers import nn_blocks
+from official.projects.yolo.modeling.layers import nn_blocks
from official.vision.modeling.decoders import factory
# model configurations
@@ -84,13 +84,13 @@
}
-class _IdentityRoute(tf.keras.layers.Layer):
+class _IdentityRoute(tf_keras.layers.Layer):
def call(self, inputs): # pylint: disable=arguments-differ
return None, inputs
-class YoloFPN(tf.keras.layers.Layer):
+class YoloFPN(tf_keras.layers.Layer):
"""YOLO Feature pyramid network."""
def __init__(self,
@@ -128,8 +128,8 @@ def __init__(self,
norm_epsilon: `float`, small float added to variance to avoid dividing by
zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
**kwargs: keyword arguments to be passed.
"""
@@ -246,7 +246,7 @@ def call(self, inputs):
return outputs
-class YoloPAN(tf.keras.layers.Layer):
+class YoloPAN(tf_keras.layers.Layer):
"""YOLO Path Aggregation Network."""
def __init__(self,
@@ -282,8 +282,8 @@ def __init__(self,
norm_epsilon: `float`, small float added to variance to avoid dividing
by zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
fpn_input: `bool`, for whether the input into this fucntion is an FPN or
a backbone.
fpn_filter_scale: `int`, scaling factor for the FPN filters.
@@ -438,7 +438,7 @@ def call(self, inputs):
return outputs
-class YoloDecoder(tf.keras.Model):
+class YoloDecoder(tf_keras.Model):
"""Darknet Backbone Decoder."""
def __init__(self,
@@ -489,8 +489,8 @@ def __init__(self,
norm_epsilon: `float`, small float added to variance to avoid dividing by
zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
**kwargs: keyword arguments to be passed.
"""
@@ -533,7 +533,7 @@ def __init__(self,
**self._base_config)
inputs = {
- key: tf.keras.layers.Input(shape=value[1:])
+ key: tf_keras.layers.Input(shape=value[1:])
for key, value in input_specs.items()
}
if self._use_fpn:
@@ -575,20 +575,20 @@ def from_config(cls, config, custom_objects=None):
def build_yolo_decoder(
input_specs: Mapping[str, tf.TensorShape],
model_config: hyperparams.Config,
- l2_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- **kwargs) -> Union[None, tf.keras.Model, tf.keras.layers.Layer]:
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs) -> Union[None, tf_keras.Model, tf_keras.layers.Layer]:
"""Builds Yolo FPN/PAN decoder from a config.
Args:
input_specs: A `dict` of input specifications. A dictionary consists of
{level: TensorShape} from a backbone.
model_config: A OneOfConfig. Model config.
- l2_regularizer: A `tf.keras.regularizers.Regularizer` instance. Default to
+ l2_regularizer: A `tf_keras.regularizers.Regularizer` instance. Default to
None.
**kwargs: Additional kwargs arguments.
Returns:
- A `tf.keras.Model` instance of the Yolo FPN/PAN decoder.
+ A `tf_keras.Model` instance of the Yolo FPN/PAN decoder.
"""
decoder_cfg = model_config.decoder.get()
norm_activation_config = model_config.norm_activation
@@ -613,7 +613,7 @@ def build_yolo_decoder(
'{yolo_model.YOLO_MODELS[decoder_cfg.version].keys()}'
'or specify a custom decoder config using YoloDecoder.')
- base_model = YOLO_MODELS[decoder_cfg.version][decoder_cfg.type]
+ base_model = YOLO_MODELS[decoder_cfg.version][decoder_cfg.type].copy()
cfg_dict = decoder_cfg.as_dict()
for key in base_model:
@@ -629,6 +629,6 @@ def build_yolo_decoder(
norm_epsilon=norm_activation_config.norm_epsilon,
kernel_regularizer=l2_regularizer)
- base_model.update(base_dict)
+ base_model.update(base_dict) # pyrefly: ignore[no-matching-overload]
model = YoloDecoder(input_specs, **base_model, **kwargs)
return model
diff --git a/official/vision/beta/projects/yolo/modeling/decoders/yolo_decoder_test.py b/official/projects/yolo/modeling/decoders/yolo_decoder_test.py
similarity index 90%
rename from official/vision/beta/projects/yolo/modeling/decoders/yolo_decoder_test.py
rename to official/projects/yolo/modeling/decoders/yolo_decoder_test.py
index 85887312454..7fbc37a6b8f 100644
--- a/official/vision/beta/projects/yolo/modeling/decoders/yolo_decoder_test.py
+++ b/official/projects/yolo/modeling/decoders/yolo_decoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,12 @@
"""Tests for YOLO."""
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.distribute import combinations
from tensorflow.python.distribute import strategy_combinations
-from official.vision.beta.projects.yolo.modeling.decoders import yolo_decoder as decoders
+from official.projects.yolo.modeling.decoders import yolo_decoder as decoders
class YoloDecoderTest(parameterized.TestCase, tf.test.TestCase):
@@ -66,7 +65,7 @@ def _build_yolo_decoder(self, input_specs, name='1'):
@parameterized.parameters('1', '6spp', '6sppfpn', '6')
def test_network_creation(self, version):
"""Test creation of ResNet family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
input_shape = {
'3': [1, 52, 52, 256],
'4': [1, 26, 26, 512],
@@ -94,7 +93,7 @@ def test_network_creation(self, version):
def test_sync_bn_multiple_devices(self, strategy, use_sync_bn):
"""Test for sync bn on TPU and GPU devices."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
with strategy.scope():
input_shape = {
@@ -113,7 +112,7 @@ def test_sync_bn_multiple_devices(self, strategy, use_sync_bn):
@parameterized.parameters(1, 3, 4)
def test_input_specs(self, input_dim):
"""Test different input feature dimensions."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
input_shape = {
'3': [1, 52, 52, 256],
@@ -129,7 +128,7 @@ def test_input_specs(self, input_dim):
def test_serialize_deserialize(self):
"""Create a network object that sets all of its config options."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
input_shape = {
'3': [1, 52, 52, 256],
diff --git a/official/projects/yolo/modeling/decoders/yolov7.py b/official/projects/yolo/modeling/decoders/yolov7.py
new file mode 100644
index 00000000000..03c38314faa
--- /dev/null
+++ b/official/projects/yolo/modeling/decoders/yolov7.py
@@ -0,0 +1,553 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains decoder architectures for YOLOv7 families.
+
+The models are built with ELAN and E-ELAN.
+
+ELAN was proposed in:
+[1] Wang, Chien-Yao and Liao, Hong-Yuan Mark and Yeh, I-Hau
+ Designing Network Design Strategies Through Gradient Path Analysis
+ arXiv:2211.04800
+
+E-ELAN is proposed in YOLOv7 paper:
+[1] Wang, Chien-Yao and Bochkovskiy, Alexey and Liao, Hong-Yuan Mark
+ YOLOv7: Trainable bag-of-freebies sets new state-of-the-art for real-time
+ object detectors
+ arXiv:2207.02696
+"""
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import hyperparams
+from official.projects.yolo.modeling.layers import nn_blocks
+from official.projects.yolo.ops import initializer_ops
+from official.vision.modeling.decoders import factory
+
+# Required block functions for YOLOv7 decoder familes.
+_BLOCK_FNS = {
+ 'convbn': nn_blocks.ConvBN,
+ 'upsample2d': tf_keras.layers.UpSampling2D,
+ 'maxpool2d': tf_keras.layers.MaxPooling2D,
+ 'concat': tf_keras.layers.Concatenate,
+ 'sppcspc': nn_blocks.SPPCSPC,
+ 'repconv': nn_blocks.RepConv,
+}
+
+# Names for key arguments needed by each block function.
+# Note that for field `from`, it can be either an integer or a str. Use of int
+# means that the previous layer comes from a decoder intermediate output, while
+# str means that the previous layer comes from the backbone output at a specific
+# level.
+_BLOCK_SPEC_SCHEMAS = {
+ 'convbn': [
+ 'block_fn',
+ 'from',
+ 'kernel_size',
+ 'strides',
+ 'filters',
+ 'is_output',
+ ],
+ 'upsample2d': [
+ 'block_fn',
+ 'from',
+ 'size',
+ 'interpolation',
+ 'is_output',
+ ],
+ 'maxpool2d': [
+ 'block_fn',
+ 'from',
+ 'pool_size',
+ 'strides',
+ 'padding',
+ 'is_output',
+ ],
+ 'concat': [
+ 'block_fn',
+ 'from',
+ 'axis',
+ 'is_output',
+ ],
+ 'sppcspc': ['block_fn', 'from', 'filters', 'is_output'],
+ 'repconv': [
+ 'block_fn',
+ 'from',
+ 'kernel_size',
+ 'strides',
+ 'filters',
+ 'is_output',
+ ],
+}
+
+# Define specs for YOLOv7-pico variant. It is recommended to use together with
+# YOLOv7-pico backbone.
+_YoloV7Pico = [
+ ['convbn', '4', 1, 1, 16, False],
+ ['convbn', -1, 1, 1, 16, False],
+
+ ['convbn', -1, 1, 1, 8, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '3', 1, 1, 8, False],
+ ['concat', [-1, -2], -1, False],
+ ['convbn', -1, 1, 1, 8, False],
+ ['convbn', -1, 1, 1, 8, False],
+
+ ['convbn', -1, 3, 2, 16, False],
+ ['concat', [-1, 1], -1, False],
+ ['convbn', -1, 1, 1, 16, False],
+
+ ['convbn', -1, 3, 2, 32, False],
+ ['convbn', -1, 1, 1, 32, False],
+
+ ['convbn', 7, 1, 1, 8, True],
+ ['convbn', 10, 1, 1, 16, True],
+ ['convbn', 12, 1, 1, 32, True],
+]
+
+# Define specs for YOLOv7-nano variant. It is recommended to use together with
+# YOLOv7-nano backbone.
+_YoloV7Nano = [
+ ['convbn', -1, 1, 1, 64, False],
+ ['convbn', -2, 1, 1, 64, False],
+ ['maxpool2d', -1, 5, 1, 'same', False],
+ ['maxpool2d', -2, 9, 1, 'same', False],
+ ['maxpool2d', -3, 13, 1, 'same', False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 64, False],
+ ['concat', [-1, -7], -1, False],
+ ['convbn', -1, 1, 1, 64, False], # 8
+
+ ['convbn', -1, 1, 1, 32, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '4', 1, 1, 32, False], # route from backbone P4
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 16, False],
+ ['convbn', -2, 1, 1, 16, False],
+ ['convbn', -1, 3, 1, 16, False],
+ ['convbn', -1, 3, 1, 16, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 32, False], # 18
+
+ ['convbn', -1, 1, 1, 16, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '3', 1, 1, 16, False], # route from backbone P3
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 8, False],
+ ['convbn', -2, 1, 1, 8, False],
+ ['convbn', -1, 3, 1, 8, False],
+ ['convbn', -1, 3, 1, 8, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 16, False], # 28
+
+ ['convbn', -1, 3, 2, 32, False],
+ ['concat', [-1, 18], -1, False],
+
+ ['convbn', -1, 1, 1, 16, False],
+ ['convbn', -2, 1, 1, 16, False],
+ ['convbn', -1, 3, 1, 16, False],
+ ['convbn', -1, 3, 1, 16, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 32, False], # 36
+
+ ['convbn', -1, 3, 2, 64, False],
+ ['concat', [-1, 8], -1, False],
+
+ ['convbn', -1, 1, 1, 32, False],
+ ['convbn', -2, 1, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 64, False], # 44
+
+ ['convbn', 28, 1, 1, 32, True],
+ ['convbn', 36, 1, 1, 64, True],
+ ['convbn', 44, 1, 1, 128, True],
+]
+
+# Define specs for YOLOv7-tiny variant. It is recommended to use together with
+# YOLOv7-tiny backbone.
+_YoloV7Tiny = [
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['maxpool2d', -1, 5, 1, 'same', False],
+ ['maxpool2d', -2, 9, 1, 'same', False],
+ ['maxpool2d', -3, 13, 1, 'same', False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 256, False],
+ ['concat', [-1, -7], -1, False],
+ ['convbn', -1, 1, 1, 256, False], # 8
+
+ ['convbn', -1, 1, 1, 128, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '4', 1, 1, 128, False], # route from backbone P4
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 64, False],
+ ['convbn', -2, 1, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 128, False], # 18
+
+ ['convbn', -1, 1, 1, 64, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '3', 1, 1, 64, False], # route from backbone P3
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 32, False],
+ ['convbn', -2, 1, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['convbn', -1, 3, 1, 32, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 64, False], # 28
+
+ ['convbn', -1, 3, 2, 128, False],
+ ['concat', [-1, 18], -1, False],
+
+ ['convbn', -1, 1, 1, 64, False],
+ ['convbn', -2, 1, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 128, False], # 36
+
+ ['convbn', -1, 3, 2, 256, False],
+ ['concat', [-1, 8], -1, False],
+
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -2, 1, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['concat', [-1, -2, -3, -4], -1, False],
+ ['convbn', -1, 1, 1, 256, False], # 44
+
+ ['convbn', 28, 1, 1, 128, True],
+ ['convbn', 36, 1, 1, 256, True],
+ ['convbn', 44, 1, 1, 512, True],
+]
+
+
+# Define specs YOLOv7 variant. The spec schema is defined above.
+# It is recommended to use together with YOLOv7 backbone.
+_YoloV7 = [
+ ['sppcspc', -1, 512, False], # 0
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '4', 1, 1, 256, False], # route from backbone P4
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['concat', [-1, -2, -3, -4, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 256, False], # 12
+
+ ['convbn', -1, 1, 1, 128, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '3', 1, 1, 128, False], # route from backbone P3
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -2, 1, 1, 128, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['convbn', -1, 3, 1, 64, False],
+ ['concat', [-1, -2, -3, -4, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 128, False], # 24
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -3, 1, 1, 128, False],
+ ['convbn', -1, 3, 2, 128, False],
+ ['concat', [-1, -3, 12], -1, False],
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['concat', [-1, -2, -3, -4, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 256, False], # 37
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -3, 1, 1, 256, False],
+ ['convbn', -1, 3, 2, 256, False],
+ ['concat', [-1, -3, 0], -1, False],
+
+ ['convbn', -1, 1, 1, 512, False],
+ ['convbn', -2, 1, 1, 512, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -2, -3, -4, -5, -6], -1, False],
+ ['convbn', -1, 1, 1, 512, False], # 50
+
+ ['repconv', 24, 3, 1, 256, True],
+ ['repconv', 37, 3, 1, 512, True],
+ ['repconv', 50, 3, 1, 1024, True],
+]
+
+_YoloV7X = [
+ ['sppcspc', -1, 640, False], # 0
+
+ ['convbn', -1, 1, 1, 320, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '4', 1, 1, 320, False], # route from backbone P4
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 320, False], # 14
+
+ ['convbn', -1, 1, 1, 160, False],
+ ['upsample2d', -1, 2, 'nearest', False],
+ ['convbn', '3', 1, 1, 160, False], # route from backbone P3
+ ['concat', [-1, -2], -1, False],
+
+ ['convbn', -1, 1, 1, 128, False],
+ ['convbn', -2, 1, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['convbn', -1, 3, 1, 128, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 160, False], # 28
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 160, False],
+ ['convbn', -3, 1, 1, 160, False],
+ ['convbn', -1, 3, 2, 160, False],
+ ['concat', [-1, -3, 14], -1, False],
+
+ ['convbn', -1, 1, 1, 256, False],
+ ['convbn', -2, 1, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['convbn', -1, 3, 1, 256, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 320, False], # 43
+
+ ['maxpool2d', -1, 2, 2, 'same', False],
+ ['convbn', -1, 1, 1, 320, False],
+ ['convbn', -3, 1, 1, 320, False],
+ ['convbn', -1, 3, 2, 320, False],
+ ['concat', [-1, -3, 0], -1, False],
+
+ ['convbn', -1, 1, 1, 512, False],
+ ['convbn', -2, 1, 1, 512, False],
+ ['convbn', -1, 3, 1, 512, False],
+ ['convbn', -1, 3, 1, 512, False],
+ ['convbn', -1, 3, 1, 512, False],
+ ['convbn', -1, 3, 1, 512, False],
+ ['convbn', -1, 3, 1, 512, False],
+ ['convbn', -1, 3, 1, 512, False],
+ ['concat', [-1, -3, -5, -7, -8], -1, False],
+ ['convbn', -1, 1, 1, 640, False], # 58
+
+ ['repconv', 28, 3, 1, 320, True],
+ ['repconv', 43, 3, 1, 640, True],
+ ['repconv', 58, 3, 1, 1280, True],
+]
+
+# Aggregates all variants for YOLOv7 decoders.
+DECODERS = {
+ 'yolov7-nano': _YoloV7Nano,
+ 'yolov7-pico': _YoloV7Pico,
+ 'yolov7-tiny': _YoloV7Tiny,
+ 'yolov7': _YoloV7,
+ 'yolov7x': _YoloV7X,
+}
+
+
+class YoloV7(tf_keras.Model):
+ """YOLOv7 decoder architecture."""
+
+ def __init__(
+ self,
+ input_specs,
+ model_id='yolov7',
+ use_sync_bn=False,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ activation='swish',
+ use_separable_conv=False,
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_initializer='zeros',
+ bias_regularizer=None,
+ **kwargs,
+ ):
+ """Initializes the YOLOv7 decoder.
+
+ Args:
+ input_specs: a dictionary of `tf.TensorShape` from backbone outputs.
+ model_id: a `str` represents the model variants.
+ use_sync_bn: if set to `True`, use synchronized batch normalization.
+ norm_momentum: a `float` of normalization momentum for the moving average.
+ norm_epsilon: a small `float` added to variance to avoid dividing by zero.
+ activation: a `str` name of the activation function.
+ use_separable_conv: `bool` wether to use separable convs.
+ kernel_initializer: a `str` for kernel initializer of convolutional
+ layers.
+ kernel_regularizer: a `tf_keras.regularizers.Regularizer` object for
+ Conv2D. Default to None.
+ bias_initializer: a `str` for bias initializer of convolutional layers.
+ bias_regularizer: a `tf_keras.regularizers.Regularizer` object for Conv2D.
+ Default to None.
+ **kwargs: Additional keyword arguments to be passed.
+ """
+
+ self._input_specs = input_specs
+ self._model_id = model_id
+ self._use_sync_bn = use_sync_bn
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._activation = activation
+ self._use_separable_conv = use_separable_conv
+
+ self._kernel_initializer = initializer_ops.pytorch_kernel_initializer(
+ kernel_initializer
+ )
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_initializer = bias_initializer
+ self._bias_regularizer = bias_regularizer
+
+ inputs = self._generate_inputs(input_specs)
+ outputs = []
+ endpoints = {}
+ level = int(min(inputs.keys()))
+ block_specs = DECODERS[model_id.lower()]
+
+ for spec in block_specs:
+ block_kwargs = dict(zip(_BLOCK_SPEC_SCHEMAS[spec[0]], spec)) # pyrefly: ignore[bad-index]
+ block_fn_str = block_kwargs.pop('block_fn')
+ from_index = block_kwargs.pop('from')
+ is_output = block_kwargs.pop('is_output')
+
+ x = self._group_layer_inputs(from_index, inputs, outputs)
+
+ if block_fn_str in ['convbn', 'sppcspc', 'repconv']:
+ block_kwargs.update({ # pyrefly: ignore[no-matching-overload]
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'activation': self._activation,
+ 'use_separable_conv': self._use_separable_conv,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_initializer': self._bias_initializer,
+ 'bias_regularizer': self._bias_regularizer,
+ })
+ block_fn = _BLOCK_FNS[block_fn_str](**block_kwargs) # pyrefly: ignore[bad-index]
+
+ x = block_fn(x)
+ outputs.append(x)
+ if is_output:
+ endpoints[str(level)] = x
+ level += 1
+ self._output_specs = {k: v.get_shape() for k, v in endpoints.items()}
+ super().__init__(inputs=inputs, outputs=endpoints, **kwargs)
+
+ def _generate_inputs(self, input_specs):
+ inputs = {}
+ for level, input_shape in input_specs.items():
+ inputs[level] = tf_keras.layers.Input(shape=input_shape[1:])
+ return inputs
+
+ def _group_layer_inputs(self, from_index, inputs, outputs):
+ if isinstance(from_index, list):
+ return [self._group_layer_inputs(i, inputs, outputs) for i in from_index]
+
+ if isinstance(from_index, int):
+ # Need last layer output from backbone.
+ if len(outputs) + from_index == -1:
+ return inputs[max(inputs.keys())]
+ return outputs[from_index]
+ return inputs[from_index] # from_index is a string.
+
+ def get_config(self):
+ config_dict = {
+ 'input_specs': self._input_specs,
+ 'model_id': self._model_id,
+ 'use_sync_bn': self._use_sync_bn,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'activation': self._activation,
+ 'kernel_initializer': self._kernel_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_initializer': self._bias_initializer,
+ 'bias_regularizer': self._bias_regularizer,
+ }
+ return config_dict
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
+
+ @property
+ def output_specs(self):
+ """A dict of {level: TensorShape} pairs for the model output."""
+ return self._output_specs
+
+
+@factory.register_decoder_builder('yolov7')
+def build_yolov7(
+ input_specs: tf_keras.layers.InputSpec,
+ model_config: hyperparams.Config,
+ l2_regularizer: tf_keras.regularizers.Regularizer = None, # pyrefly: ignore[bad-function-definition]
+) -> tf_keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
+ """Builds YOLOv7 decoder."""
+ decoder_config = model_config.decoder
+ norm_activation_config = model_config.norm_activation
+ assert (
+ decoder_config.type == 'yolov7'
+ ), f'Inconsistent decoder type {decoder_config.type}.'
+ decoder_config = decoder_config.get()
+ assert (
+ decoder_config.model_id in DECODERS
+ ), f'Unsupported decoder {decoder_config.model_id}.'
+ model = YoloV7(
+ model_id=decoder_config.model_id,
+ input_specs=input_specs,
+ use_sync_bn=norm_activation_config.use_sync_bn,
+ norm_momentum=norm_activation_config.norm_momentum,
+ norm_epsilon=norm_activation_config.norm_epsilon,
+ activation=norm_activation_config.activation,
+ kernel_regularizer=l2_regularizer,
+ use_separable_conv=decoder_config.use_separable_conv,
+ )
+ return model
diff --git a/official/projects/yolo/modeling/decoders/yolov7_test.py b/official/projects/yolo/modeling/decoders/yolov7_test.py
new file mode 100644
index 00000000000..dfb77a09fd8
--- /dev/null
+++ b/official/projects/yolo/modeling/decoders/yolov7_test.py
@@ -0,0 +1,98 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for yolov7 decoder."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from tensorflow.python.distribute import combinations
+from tensorflow.python.distribute import strategy_combinations
+from official.projects.yolo.modeling.backbones import yolov7 as backbone
+from official.projects.yolo.modeling.decoders import yolov7 as decoder
+
+_INPUT_SIZE = (224, 224)
+
+
+class YoloV7DecoderTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ('yolov7',),
+ )
+ def test_network_creation(self, model_id):
+ """Tests declaration of YOLOv7 decoder variants."""
+ tf_keras.backend.set_image_data_format('channels_last')
+
+ backbone_network = backbone.YoloV7(model_id)
+ decoder_network = decoder.YoloV7(backbone_network.output_specs, model_id)
+ self.assertEqual(decoder_network.get_config()['model_id'], model_id)
+
+ inputs = tf_keras.Input(shape=(*_INPUT_SIZE, 3), batch_size=1)
+ outputs = decoder_network(backbone_network(inputs))
+
+ for level, level_output in outputs.items():
+ scale = 2 ** int(level)
+ input_size = (_INPUT_SIZE[0] // scale, _INPUT_SIZE[1] // scale)
+ self.assertAllEqual((1, *input_size), level_output.shape.as_list()[:-1])
+
+ @combinations.generate(
+ combinations.combine(
+ strategy=[
+ strategy_combinations.cloud_tpu_strategy,
+ strategy_combinations.one_device_strategy_gpu,
+ ],
+ )
+ )
+ def test_sync_bn_multiple_devices(self, strategy):
+ """Test for sync bn on TPU and GPU devices."""
+ inputs = np.random.rand(1, *_INPUT_SIZE, 3)
+
+ tf_keras.backend.set_image_data_format('channels_last')
+
+ with strategy.scope():
+ backbone_network = backbone.YoloV7(model_id='yolov7', use_sync_bn=True)
+ decoder_network = decoder.YoloV7(
+ backbone_network.output_specs, model_id='yolov7', use_sync_bn=True)
+ _ = decoder_network(backbone_network(inputs))
+
+ def test_serialize_deserialize(self):
+ # Create a network object that sets all of its config options.
+ kwargs = dict(
+ model_id='yolov7',
+ use_sync_bn=False,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ activation='swish',
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_initializer='zeros',
+ bias_regularizer=None,
+ )
+ backbone_network = backbone.YoloV7(**kwargs)
+ kwargs['input_specs'] = backbone_network.output_specs
+ decoder_network = decoder.YoloV7(**kwargs)
+
+ # Create another network object from the first object's config.
+ new_network = decoder.YoloV7.from_config(decoder_network.get_config())
+
+ # Validate that the config can be forced to JSON.
+ _ = new_network.to_json()
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(decoder_network.get_config(), new_network.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/vision/beta/projects/yolo/modeling/factory.py b/official/projects/yolo/modeling/factory.py
similarity index 54%
rename from official/vision/beta/projects/yolo/modeling/factory.py
rename to official/projects/yolo/modeling/factory.py
index 1fbbecb4986..79dae40d5bc 100644
--- a/official/vision/beta/projects/yolo/modeling/factory.py
+++ b/official/projects/yolo/modeling/factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,10 +16,12 @@
from absl import logging
-from official.vision.beta.projects.yolo.configs import yolo
-from official.vision.beta.projects.yolo.modeling import yolo_model
-from official.vision.beta.projects.yolo.modeling.heads import yolo_head
-from official.vision.beta.projects.yolo.modeling.layers import detection_generator
+from official.projects.yolo.configs import yolo
+from official.projects.yolo.modeling import yolo_model
+from official.projects.yolo.modeling import yolov7_model
+from official.projects.yolo.modeling.heads import yolo_head
+from official.projects.yolo.modeling.heads import yolov7_head
+from official.projects.yolo.modeling.layers import detection_generator
from official.vision.modeling.backbones import factory as backbone_factory
from official.vision.modeling.decoders import factory as decoder_factory
@@ -33,7 +35,7 @@ def build_yolo_detection_generator(model_config: yolo.Yolo, anchor_boxes):
nms_thresh=model_config.detection_generator.nms_thresh,
max_boxes=model_config.detection_generator.max_boxes,
pre_nms_points=model_config.detection_generator.pre_nms_points,
- nms_type=model_config.detection_generator.nms_type,
+ nms_version=model_config.detection_generator.nms_version,
box_type=model_config.detection_generator.box_type.get(),
path_scale=model_config.detection_generator.path_scales.get(),
scale_xy=model_config.detection_generator.scale_xy.get(),
@@ -47,7 +49,9 @@ def build_yolo_detection_generator(model_config: yolo.Yolo, anchor_boxes):
cls_normalizer=model_config.loss.cls_normalizer.get(),
object_normalizer=model_config.loss.object_normalizer.get(),
ignore_thresh=model_config.loss.ignore_thresh.get(),
- objectness_smooth=model_config.loss.objectness_smooth.get())
+ objectness_smooth=model_config.loss.objectness_smooth.get(),
+ use_class_agnostic_nms=model_config.detection_generator.use_class_agnostic_nms,
+ )
return model
@@ -87,9 +91,81 @@ def build_yolo(input_specs, model_config, l2_regularization):
decoder=decoder,
head=head,
detection_generator=detection_generator_obj)
- model.build(input_specs.shape)
-
- model.summary(print_fn=logging.info)
losses = detection_generator_obj.get_losses()
return model, losses
+
+
+def build_yolov7_detection_generator(model_config: yolo.Yolo, anchor_boxes):
+ """Builds yolo detection generator."""
+ model = detection_generator.YoloLayer(
+ classes=model_config.num_classes,
+ anchors=anchor_boxes,
+ iou_thresh=model_config.detection_generator.iou_thresh,
+ nms_thresh=model_config.detection_generator.nms_thresh,
+ max_boxes=model_config.detection_generator.max_boxes,
+ pre_nms_points=model_config.detection_generator.pre_nms_points,
+ nms_version=model_config.detection_generator.nms_version,
+ box_type=model_config.detection_generator.box_type.get(),
+ path_scale=model_config.detection_generator.path_scales.get(),
+ scale_xy=model_config.detection_generator.scale_xy.get(),
+ use_class_agnostic_nms=model_config.detection_generator.use_class_agnostic_nms,
+ )
+ return model
+
+
+def build_yolov7(input_specs, model_config, l2_regularization):
+ """Builds yolov7 model."""
+ norm_activation_config = model_config.norm_activation
+ backbone = backbone_factory.build_backbone(
+ input_specs,
+ model_config.backbone,
+ norm_activation_config,
+ l2_regularization,
+ )
+ decoder = decoder_factory.build_decoder(
+ backbone.output_specs,
+ model_config,
+ l2_regularization,
+ )
+
+ decoder_output_specs = decoder.output_specs
+ min_level = min(map(int, decoder_output_specs.keys()))
+ max_level = max(map(int, decoder_output_specs.keys()))
+ if min_level != model_config.min_level:
+ logging.warning(
+ (
+ 'The `min_level` does not match! Expects min_level=%d but got '
+ 'min_level=%d. Expected value will be used.'
+ ),
+ min_level,
+ model_config.min_level,
+ )
+ if max_level != model_config.max_level:
+ logging.warning(
+ (
+ 'The `max_level` does not match! Expects max_level=%d but got'
+ 'max_level=%d. Expected value will be used.'
+ ),
+ max_level,
+ model_config.max_level,
+ )
+ anchor_dict, _ = model_config.anchor_boxes.get(min_level, max_level)
+ num_anchors = len(anchor_dict[str(min_level)])
+ head = yolov7_head.YoloV7DetectionHead(
+ model_config.num_classes,
+ min_level,
+ max_level,
+ num_anchors,
+ kernel_regularizer=l2_regularization,
+ use_separable_conv=model_config.head.use_separable_conv,
+ )
+ # YOLOv7 and YOLOv4 share the same detection generator.
+ detection_generator_obj = build_yolov7_detection_generator(
+ model_config, anchor_dict
+ )
+ model = yolov7_model.YoloV7(
+ backbone, decoder, head, detection_generator=detection_generator_obj
+ )
+
+ return model
diff --git a/official/projects/yolo/modeling/factory_test.py b/official/projects/yolo/modeling/factory_test.py
new file mode 100644
index 00000000000..647a279f9b2
--- /dev/null
+++ b/official/projects/yolo/modeling/factory_test.py
@@ -0,0 +1,107 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for factory.py."""
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+# pylint: disable=unused-import
+from official.projects.yolo.configs import backbones
+from official.projects.yolo.configs import yolo
+from official.projects.yolo.configs import yolov7
+from official.projects.yolo.modeling import factory
+from official.projects.yolo.modeling.backbones import darknet
+from official.projects.yolo.modeling.backbones import yolov7 as yolov7_backbone
+from official.projects.yolo.modeling.decoders import yolo_decoder
+from official.projects.yolo.modeling.decoders import yolov7 as yolov7_decoder
+from official.projects.yolo.modeling.heads import yolo_head as heads
+from official.projects.yolo.modeling.heads import yolov7_head
+from official.projects.yolo.modeling.layers import detection_generator
+# pylint: enable=unused-import
+
+
+class FactoryTest(tf.test.TestCase):
+
+ def test_yolo_builder(self):
+ num_classes = 3
+ input_size = 640
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None, input_size, input_size, 3])
+ model_config = yolo.Yolo(
+ num_classes=num_classes,
+ head=yolo.YoloHead(smart_bias=True),
+ anchor_boxes=yolo.AnchorBoxes(
+ anchors_per_scale=3,
+ boxes=[
+ yolo.Box(box=[12, 16]),
+ yolo.Box(box=[19, 36]),
+ yolo.Box(box=[40, 28]),
+ yolo.Box(box=[36, 75]),
+ yolo.Box(box=[76, 55]),
+ yolo.Box(box=[72, 146]),
+ yolo.Box(box=[142, 110]),
+ yolo.Box(box=[192, 243]),
+ yolo.Box(box=[459, 401])
+ ]))
+ l2_regularizer = tf_keras.regularizers.l2(5e-5)
+
+ yolo_model, _ = factory.build_yolo(
+ input_specs=input_specs,
+ model_config=model_config,
+ l2_regularization=l2_regularizer)
+
+ # Do forward pass.
+ inputs = np.random.rand(2, input_size, input_size, 3)
+ _ = yolo_model(inputs)
+
+ def test_yolov7_builder(self):
+ num_classes = 3
+ input_size = 640
+ input_specs = tf_keras.layers.InputSpec(
+ shape=[None, input_size, input_size, 3]
+ )
+ model_config = yolov7.YoloV7(
+ num_classes=num_classes,
+ head=yolov7.YoloV7Head(),
+ anchor_boxes=yolo.AnchorBoxes(
+ anchors_per_scale=3,
+ boxes=[
+ yolo.Box(box=[12, 16]),
+ yolo.Box(box=[19, 36]),
+ yolo.Box(box=[40, 28]),
+ yolo.Box(box=[36, 75]),
+ yolo.Box(box=[76, 55]),
+ yolo.Box(box=[72, 146]),
+ yolo.Box(box=[142, 110]),
+ yolo.Box(box=[192, 243]),
+ yolo.Box(box=[459, 401]),
+ ],
+ ),
+ )
+ l2_regularizer = tf_keras.regularizers.l2(5e-5)
+
+ yolo_model = factory.build_yolov7(
+ input_specs=input_specs,
+ model_config=model_config,
+ l2_regularization=l2_regularizer,
+ )
+
+ # Do forward pass.
+ inputs = np.random.rand(2, input_size, input_size, 3)
+ _ = yolo_model(inputs)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/yolo/modeling/heads/__init__.py b/official/projects/yolo/modeling/heads/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/modeling/heads/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/modeling/heads/yolo_head.py b/official/projects/yolo/modeling/heads/yolo_head.py
similarity index 91%
rename from official/vision/beta/projects/yolo/modeling/heads/yolo_head.py
rename to official/projects/yolo/modeling/heads/yolo_head.py
index 70a6a76cfd0..6b83ca74fdf 100644
--- a/official/vision/beta/projects/yolo/modeling/heads/yolo_head.py
+++ b/official/projects/yolo/modeling/heads/yolo_head.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,11 @@
"""Yolo heads."""
-import tensorflow as tf
-from official.vision.beta.projects.yolo.modeling.layers import nn_blocks
+import tensorflow as tf, tf_keras
+from official.projects.yolo.modeling.layers import nn_blocks
-class YoloHead(tf.keras.layers.Layer):
+class YoloHead(tf_keras.layers.Layer):
"""YOLO Prediction Head."""
def __init__(self,
@@ -50,8 +50,8 @@ def __init__(self,
norm_epsilon: `float`, small float added to variance to avoid dividing by
zero.
kernel_initializer: kernel_initializer for convolutional layers.
- kernel_regularizer: tf.keras.regularizers.Regularizer object for Conv2D.
- bias_regularizer: tf.keras.regularizers.Regularizer object for Conv2d.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
activation: `str`, the activation function to use typically leaky or mish.
smart_bias: `bool`, whether to use smart bias.
use_separable_conv: `bool` wether to use separable convs.
@@ -82,7 +82,7 @@ def __init__(self,
kernel_regularizer=kernel_regularizer,
bias_regularizer=bias_regularizer)
- self._conv_config = dict(
+ self._conv_config = dict( # pyrefly: ignore[bad-unpacking]
filters=self._output_conv,
kernel_size=(1, 1),
strides=(1, 1),
@@ -94,7 +94,7 @@ def __init__(self,
def bias_init(self, scale, inshape, isize=640, no_per_conf=8):
def bias(shape, dtype):
- init = tf.keras.initializers.Zeros()
+ init = tf_keras.initializers.Zeros()
base = init(shape, dtype=dtype)
if self._smart_bias:
base = tf.reshape(base, [self._boxes_per_level, -1])
diff --git a/official/vision/beta/projects/yolo/modeling/heads/yolo_head_test.py b/official/projects/yolo/modeling/heads/yolo_head_test.py
similarity index 86%
rename from official/vision/beta/projects/yolo/modeling/heads/yolo_head_test.py
rename to official/projects/yolo/modeling/heads/yolo_head_test.py
index 37818743a70..fd5bf388819 100644
--- a/official/vision/beta/projects/yolo/modeling/heads/yolo_head_test.py
+++ b/official/projects/yolo/modeling/heads/yolo_head_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,18 +14,17 @@
"""Tests for yolo heads."""
-# Import libraries
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.modeling.heads import yolo_head as heads
+from official.projects.yolo.modeling.heads import yolo_head as heads
class YoloDecoderTest(parameterized.TestCase, tf.test.TestCase):
def test_network_creation(self):
"""Test creation of YOLO family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
input_shape = {
'3': [1, 52, 52, 256],
'4': [1, 26, 26, 512],
@@ -49,7 +48,7 @@ def test_network_creation(self):
def test_serialize_deserialize(self):
# Create a network object that sets all of its config options.
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
input_shape = {
'3': [1, 52, 52, 256],
'4': [1, 26, 26, 512],
diff --git a/official/projects/yolo/modeling/heads/yolov7_head.py b/official/projects/yolo/modeling/heads/yolov7_head.py
new file mode 100644
index 00000000000..6a427072027
--- /dev/null
+++ b/official/projects/yolo/modeling/heads/yolov7_head.py
@@ -0,0 +1,157 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""YOLOv7 heads."""
+
+import tensorflow as tf, tf_keras
+from official.projects.yolo.ops import initializer_ops
+
+
+class YoloV7DetectionHead(tf_keras.layers.Layer):
+ """YOLOv7 Detection Head."""
+
+ def __init__(
+ self,
+ num_classes=80,
+ min_level=3,
+ max_level=5,
+ num_anchors=3,
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_initializer='zeros',
+ bias_regularizer=None,
+ use_separable_conv=False,
+ **kwargs,
+ ):
+ """Initializes YOLOv7 head.
+
+ Args:
+ num_classes: integer.
+ min_level: minimum feature level.
+ max_level: maximum feature level.
+ num_anchors: integer for number of anchors at each location.
+ kernel_initializer: kernel_initializer for convolutional layers.
+ kernel_regularizer: tf_keras.regularizers.Regularizer object for Conv2D.
+ bias_initializer: bias initializer for convolutional layers.
+ bias_regularizer: tf_keras.regularizers.Regularizer object for Conv2d.
+ use_separable_conv: `bool` wether to use separable convs.
+ **kwargs: other keyword arguments.
+ """
+ super().__init__(**kwargs)
+ self._num_classes = num_classes
+ self._min_level = min_level
+ self._max_level = max_level
+ self._num_anchors = num_anchors
+
+ self._kernel_initializer = initializer_ops.pytorch_kernel_initializer(
+ kernel_initializer
+ )
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_initializer = bias_initializer
+ self._bias_regularizer = bias_regularizer
+ self._use_separable_conv = use_separable_conv
+
+ def _bias_init(self, scale, in_channels, isize=640, no_per_conf=8):
+
+ def bias(shape, dtype):
+ init = tf_keras.initializers.VarianceScaling(
+ scale=1 / 3, mode='fan_in', distribution='uniform')
+ base = init([in_channels, *shape], dtype=dtype)[0]
+
+ base = tf.reshape(base, [self._num_anchors, -1])
+ box, conf, classes = tf.split(base, [4, 1, -1], axis=-1)
+ conf += tf.math.log(no_per_conf / ((isize / scale)**2))
+ classes += tf.math.log(0.6 / (self._num_classes - 0.99))
+ base = tf.concat([box, conf, classes], axis=-1)
+ base = tf.reshape(base, [-1])
+ return base
+
+ return bias
+
+ def build(self, input_shape):
+ self._convs = []
+ self._implicit_adds = []
+ self._implicit_muls = []
+ conv_op = (
+ tf_keras.layers.SeparableConv2D
+ if self._use_separable_conv
+ else tf_keras.layers.Conv2D
+ )
+ for level in range(self._min_level, self._max_level + 1):
+ # Note that we assume height == width.
+ h = input_shape[str(level)][2]
+ scale = 2 ** int(level)
+ in_channels = input_shape[str(level)][-1]
+ # Outputs are num_classes + 5 (box coordinates + objectness score)
+ self._convs.append(
+ conv_op(
+ (self._num_classes + 5) * self._num_anchors,
+ kernel_size=1,
+ padding='same',
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_initializer=self._bias_init(scale, in_channels, h * scale),
+ )
+ )
+ self._implicit_adds.append(
+ self.add_weight(
+ name=f'implicit_adds_l{level}',
+ shape=[1, 1, 1, in_channels],
+ initializer=tf_keras.initializers.random_normal(
+ mean=0.0, stddev=0.02
+ ),
+ trainable=True,
+ )
+ )
+ self._implicit_muls.append(
+ self.add_weight(
+ name=f'implicit_muls_l{level}',
+ shape=[1, 1, 1, (self._num_classes + 5) * self._num_anchors],
+ initializer=tf_keras.initializers.random_normal(
+ mean=1.0, stddev=0.02
+ ),
+ trainable=True,
+ )
+ )
+ super().build(input_shape)
+
+ def call(self, inputs, training=False):
+ outputs = {}
+ for i, level in enumerate(range(self._min_level, self._max_level + 1)):
+ x = inputs[str(level)]
+ x = self._implicit_adds[i] + x
+ x = self._convs[i](x)
+ x = self._implicit_muls[i] * x
+ _, h, w, _ = x.get_shape().as_list()
+ x = tf.reshape(x, [-1, h, w, self._num_anchors, self._num_classes + 5])
+ outputs[str(level)] = x
+ return outputs
+
+ def get_config(self):
+ config = dict(
+ num_classes=self._num_classes,
+ min_level=self._min_level,
+ max_level=self._max_level,
+ num_anchors=self._num_anchors,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_initializer=self._bias_initializer,
+ bias_regularizer=self._bias_regularizer,
+ use_separable_conv=self._use_separable_conv,
+ )
+ return config
+
+ @classmethod
+ def from_config(cls, config, custom_objects=None):
+ return cls(**config)
diff --git a/official/projects/yolo/modeling/heads/yolov7_head_test.py b/official/projects/yolo/modeling/heads/yolov7_head_test.py
new file mode 100644
index 00000000000..24999b0e433
--- /dev/null
+++ b/official/projects/yolo/modeling/heads/yolov7_head_test.py
@@ -0,0 +1,76 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for yolov7 heads."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.yolo.modeling.backbones import yolov7 as backbone
+from official.projects.yolo.modeling.decoders import yolov7 as decoder
+from official.projects.yolo.modeling.heads import yolov7_head as head
+
+_INPUT_SIZE = (224, 224)
+
+
+class YoloV7DetectionHeadTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.parameters(
+ ('yolov7',),
+ )
+ def test_network_creation(self, model_id):
+ """Tests declaration of YOLOv7 detection head."""
+ tf_keras.backend.set_image_data_format('channels_last')
+
+ backbone_network = backbone.YoloV7(model_id)
+ decoder_network = decoder.YoloV7(backbone_network.output_specs, model_id)
+ head_network = head.YoloV7DetectionHead()
+
+ inputs = tf_keras.Input(shape=(*_INPUT_SIZE, 3), batch_size=1)
+ outputs = head_network(decoder_network(backbone_network(inputs)))
+
+ for level, level_output in outputs.items():
+ scale = 2 ** int(level)
+ input_size = (_INPUT_SIZE[0] // scale, _INPUT_SIZE[1] // scale)
+ head_config = head_network.get_config()
+ num_classes = head_config['num_classes']
+ num_anchors = head_config['num_anchors']
+ self.assertAllEqual(
+ (1, *input_size, num_anchors, num_classes + 5),
+ level_output.shape.as_list(),
+ )
+
+ def test_serialize_deserialize(self):
+ # Create a network object that sets all of its config options.
+ kwargs = dict(
+ num_classes=3,
+ min_level=3,
+ max_level=5,
+ num_anchors=3,
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_initializer='zeros',
+ bias_regularizer=None,
+ )
+ network = head.YoloV7DetectionHead(**kwargs)
+
+ # Create another network object from the first object's config.
+ new_network = head.YoloV7DetectionHead.from_config(network.get_config())
+
+ # If the serialization was successful, the new config should match the old.
+ self.assertAllEqual(network.get_config(), new_network.get_config())
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/yolo/modeling/layers/__init__.py b/official/projects/yolo/modeling/layers/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/modeling/layers/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/modeling/layers/detection_generator.py b/official/projects/yolo/modeling/layers/detection_generator.py
similarity index 66%
rename from official/vision/beta/projects/yolo/modeling/layers/detection_generator.py
rename to official/projects/yolo/modeling/layers/detection_generator.py
index e3df866aa06..258b1ef1d50 100644
--- a/official/vision/beta/projects/yolo/modeling/layers/detection_generator.py
+++ b/official/projects/yolo/modeling/layers/detection_generator.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,46 +13,52 @@
# limitations under the License.
"""Contains common building blocks for yolo layer (detection layer)."""
-import tensorflow as tf
+from typing import Optional
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.losses import yolo_loss
-from official.vision.beta.projects.yolo.ops import box_ops
-from official.vision.beta.projects.yolo.ops import loss_utils
+from official.projects.yolo.losses import yolo_loss
+from official.projects.yolo.ops import box_ops
+from official.projects.yolo.ops import loss_utils
from official.vision.modeling.layers import detection_generator
-class YoloLayer(tf.keras.Model):
+class YoloLayer(tf_keras.layers.Layer):
"""Yolo layer (detection generator)."""
- def __init__(self,
- anchors,
- classes,
- iou_thresh=0.0,
- ignore_thresh=0.7,
- truth_thresh=1.0,
- nms_thresh=0.6,
- max_delta=10.0,
- loss_type='ciou',
- iou_normalizer=1.0,
- cls_normalizer=1.0,
- object_normalizer=1.0,
- use_scaled_loss=False,
- update_on_repeat=False,
- pre_nms_points=5000,
- label_smoothing=0.0,
- max_boxes=200,
- box_type='original',
- path_scale=None,
- scale_xy=None,
- nms_type='greedy',
- objectness_smooth=False,
- **kwargs):
+ def __init__(
+ self,
+ anchors,
+ classes,
+ apply_nms=True,
+ iou_thresh=0.0,
+ ignore_thresh=0.7,
+ truth_thresh=1.0,
+ nms_thresh=0.6,
+ max_delta=10.0,
+ loss_type='ciou',
+ iou_normalizer=1.0,
+ cls_normalizer=1.0,
+ object_normalizer=1.0,
+ use_scaled_loss=False,
+ update_on_repeat=False,
+ pre_nms_points=5000,
+ label_smoothing=0.0,
+ max_boxes=200,
+ box_type='original',
+ path_scale=None,
+ scale_xy=None,
+ nms_version='greedy',
+ objectness_smooth=False,
+ use_class_agnostic_nms: Optional[bool] = False,
+ **kwargs
+ ):
"""Parameters for the loss functions used at each detection head output.
Args:
anchors: `List[List[int]]` for the anchor boxes that are used in the
model.
classes: `int` for the number of classes.
+ apply_nms: A boolean indicating whether to apply NMS.
iou_thresh: `float` to use many anchors per object if IoU(Obj, Anchor) >
iou_thresh.
ignore_thresh: `float` for the IOU value over which the loss is not
@@ -61,15 +67,15 @@ def __init__(self,
despite a detection being made'.
nms_thresh: `float` for the minimum IOU value for an overlap.
max_delta: gradient clipping to apply to the box loss.
- loss_type: `str` for the typeof iou loss to use with in {ciou, diou,
- giou, iou}.
+ loss_type: `str` for the typeof iou loss to use with in {ciou, diou, giou,
+ iou}.
iou_normalizer: `float` for how much to scale the loss on the IOU or the
boxes.
cls_normalizer: `float` for how much to scale the loss on the classes.
object_normalizer: `float` for how much to scale loss on the detection
map.
- use_scaled_loss: `bool` for whether to use the scaled loss
- or the traditional loss.
+ use_scaled_loss: `bool` for whether to use the scaled loss or the
+ traditional loss.
update_on_repeat: `bool` indicating how you would like to handle repeated
indexes in a given [j, i] index. Setting this to True will give more
consistent MAP, setting it to falls will improve recall by 1-2% but will
@@ -94,17 +100,21 @@ def __init__(self,
scale_xy: dictionary `float` values inidcating how far each pixel can see
outside of its containment of 1.0. a value of 1.2 indicates there is a
20% extended radius around each pixel that this specific pixel can
- predict values for a center at. the center can range from 0 - value/2
- to 1 + value/2, this value is set in the yolo filter, and resused here.
+ predict values for a center at. the center can range from 0 - value/2 to
+ 1 + value/2, this value is set in the yolo filter, and resused here.
there should be one value for scale_xy for each level from min_level to
max_level.
- nms_type: `str` for which non max suppression to use.
+ nms_version: `str` for which non max suppression to use.
objectness_smooth: `float` for how much to smooth the loss on the
detection map.
+ use_class_agnostic_nms: A `bool` of whether non max suppression is
+ operated on all the boxes using max scores across all classes. Only
+ valid when nms_version is v2.
**kwargs: Addtional keyword arguments.
"""
super().__init__(**kwargs)
self._anchors = anchors
+ self._apply_nms = apply_nms
self._thresh = iou_thresh
self._ignore_thresh = ignore_thresh
self._truth_thresh = truth_thresh
@@ -117,6 +127,7 @@ def __init__(self,
self._max_delta = max_delta
self._classes = classes
self._loss_type = loss_type
+ self._use_class_agnostic_nms = use_class_agnostic_nms
self._use_scaled_loss = use_scaled_loss
self._update_on_repeat = update_on_repeat
@@ -129,7 +140,7 @@ def __init__(self,
self._box_type = box_type
self._path_scale = path_scale or {key: 2**int(key) for key in self._keys}
- self._nms_type = nms_type
+ self._nms_version = nms_version
self._scale_xy = scale_xy or {key: 1.0 for key, _ in anchors.items()}
self._generator = {}
@@ -153,25 +164,25 @@ def parse_prediction_path(self, key, inputs):
len_mask = self._len_mask[key]
scale_xy = self._scale_xy[key]
- # reshape the yolo output to (batchsize,
+ # Reshape the yolo output to (batchsize,
# width,
# height,
# number_anchors,
# remaining_points)
data = tf.reshape(inputs, [-1, height, width, len_mask, self._classes + 5])
- # use the grid generator to get the formatted anchor boxes and grid points
- # in shape [1, height, width, 2]
+ # Use the grid generator to get the formatted anchor boxes and grid points
+ # in shape [1, height, width, 2].
centers, anchors = generator(height, width, batchsize, dtype=data.dtype)
- # split the yolo detections into boxes, object score map, classes
+ # Split the yolo detections into boxes, object score map, classes.
boxes, obns_scores, class_scores = tf.split(
data, [4, 1, self._classes], axis=-1)
- # determine the number of classes
+ # Determine the number of classes.
classes = class_scores.get_shape().as_list()[-1]
- # configurable to use the new coordinates in scaled Yolo v4 or not
+ # Configurable to use the new coordinates in scaled Yolo v4 or not.
_, _, boxes = loss_utils.get_predicted_box(
tf.cast(height, data.dtype),
tf.cast(width, data.dtype),
@@ -183,23 +194,23 @@ def parse_prediction_path(self, key, inputs):
darknet=False,
box_type=self._box_type[key])
- # convert boxes from yolo(x, y, w. h) to tensorflow(ymin, xmin, ymax, xmax)
+ # Convert boxes from yolo(x, y, w. h) to tensorflow(ymin, xmin, ymax, xmax).
boxes = box_ops.xcycwh_to_yxyx(boxes)
- # activate and detection map
+ # Activate and detection map
obns_scores = tf.math.sigmoid(obns_scores)
- # convert detection map to class detection probabailities
+ # Convert detection map to class detection probabilities.
class_scores = tf.math.sigmoid(class_scores) * obns_scores
- # platten predictions to [batchsize, N, -1] for non max supression
+ # Flatten predictions to [batchsize, N, -1] for non max supression.
fill = height * width * len_mask
boxes = tf.reshape(boxes, [-1, fill, 4])
class_scores = tf.reshape(class_scores, [-1, fill, classes])
obns_scores = tf.reshape(obns_scores, [-1, fill])
return obns_scores, boxes, class_scores
- def call(self, inputs):
+ def __call__(self, inputs):
boxes = []
class_scores = []
object_scores = []
@@ -207,7 +218,7 @@ def call(self, inputs):
min_level = int(min(levels))
max_level = int(max(levels))
- # aggregare boxes over each scale
+ # Aggregate boxes over each scale.
for i in range(min_level, max_level + 1):
key = str(i)
object_scores_, boxes_, class_scores_ = self.parse_prediction_path(
@@ -216,58 +227,79 @@ def call(self, inputs):
class_scores.append(class_scores_)
object_scores.append(object_scores_)
- # colate all predicitons
+ # Collate all predicitons.
boxes = tf.concat(boxes, axis=1)
object_scores = tf.concat(object_scores, axis=1)
class_scores = tf.concat(class_scores, axis=1)
- # get masks to threshold all the predicitons
+ # Get masks to threshold all the predicitons.
object_mask = tf.cast(object_scores > self._thresh, object_scores.dtype)
class_mask = tf.cast(class_scores > self._thresh, class_scores.dtype)
- # apply thresholds mask to all the predicitons
+ # Apply thresholds mask to all the predictions.
object_scores *= object_mask
class_scores *= (tf.expand_dims(object_mask, axis=-1) * class_mask)
- # apply nms
- if self._nms_type == 'greedy':
- # greedy NMS
- boxes = tf.cast(boxes, dtype=tf.float32)
- class_scores = tf.cast(class_scores, dtype=tf.float32)
- boxes, object_scores_, class_scores, num_detections = (
+ # Make a copy of the original dtype.
+ dtype = object_scores.dtype
+
+ if not self._apply_nms:
+ return {
+ 'bbox': tf.expand_dims(tf.cast(boxes, dtype=tf.float32), axis=-2),
+ 'classes': tf.cast(class_scores, dtype=tf.float32),
+ 'confidence': object_scores,
+ 'num_detections': self._max_boxes,
+ }
+
+ # Apply nms.
+ if self._nms_version == 'greedy':
+ # Greedy NMS.
+ boxes, object_scores, class_scores, num_detections = (
tf.image.combined_non_max_suppression(
- tf.expand_dims(boxes, axis=-2),
- class_scores,
+ tf.expand_dims(tf.cast(boxes, dtype=tf.float32), axis=-2),
+ tf.cast(class_scores, dtype=tf.float32),
self._pre_nms_points,
self._max_boxes,
iou_threshold=self._nms_thresh,
- score_threshold=self._thresh))
- # cast the boxes and predicitons abck to original datatype
- boxes = tf.cast(boxes, object_scores.dtype)
- class_scores = tf.cast(class_scores, object_scores.dtype)
- object_scores = tf.cast(object_scores_, object_scores.dtype)
- else:
- # TPU NMS
- boxes = tf.cast(boxes, dtype=tf.float32)
- class_scores = tf.cast(class_scores, dtype=tf.float32)
- (boxes, confidence, classes,
- num_detections) = detection_generator._generate_detections_v2( # pylint:disable=protected-access
- tf.expand_dims(boxes, axis=-2),
- class_scores,
- pre_nms_top_k=self._pre_nms_points,
- max_num_detections=self._max_boxes,
- nms_iou_threshold=self._nms_thresh,
- pre_nms_score_threshold=self._thresh)
- boxes = tf.cast(boxes, object_scores.dtype)
- class_scores = tf.cast(classes, object_scores.dtype)
- object_scores = tf.cast(confidence, object_scores.dtype)
-
- # format and return
+ score_threshold=self._thresh,
+ )
+ )
+ elif self._nms_version == 'v1':
+ (boxes, object_scores, class_scores, num_detections, _) = (
+ detection_generator._generate_detections_v1( # pylint:disable=protected-access
+ tf.expand_dims(tf.cast(boxes, dtype=tf.float32), axis=-2),
+ tf.cast(class_scores, dtype=tf.float32),
+ pre_nms_top_k=self._pre_nms_points,
+ max_num_detections=self._max_boxes,
+ nms_iou_threshold=self._nms_thresh,
+ pre_nms_score_threshold=self._thresh,
+ )
+ )
+
+ elif self._nms_version == 'v2' or self._nms_version == 'iou':
+ (boxes, object_scores, class_scores, num_detections) = (
+ detection_generator._generate_detections_v2( # pylint:disable=protected-access
+ tf.expand_dims(tf.cast(boxes, dtype=tf.float32), axis=-2),
+ tf.cast(class_scores, dtype=tf.float32),
+ pre_nms_top_k=self._pre_nms_points,
+ max_num_detections=self._max_boxes,
+ nms_iou_threshold=self._nms_thresh,
+ pre_nms_score_threshold=self._thresh,
+ use_class_agnostic_nms=self._use_class_agnostic_nms,
+ )
+ )
+
+ # Cast the boxes and predicitons back to original datatype.
+ boxes = tf.cast(boxes, dtype)
+ class_scores = tf.cast(class_scores, dtype)
+ object_scores = tf.cast(object_scores, dtype)
+
+ # Format and return
return {
'bbox': boxes,
'classes': class_scores,
'confidence': object_scores,
- 'num_detections': num_detections,
+ 'num_detections': num_detections, # pyrefly: ignore[unbound-name]
}
def get_losses(self):
diff --git a/official/vision/beta/projects/yolo/modeling/layers/detection_generator_test.py b/official/projects/yolo/modeling/layers/detection_generator_test.py
similarity index 71%
rename from official/vision/beta/projects/yolo/modeling/layers/detection_generator_test.py
rename to official/projects/yolo/modeling/layers/detection_generator_test.py
index fc8e6100335..fca7fbb130f 100644
--- a/official/vision/beta/projects/yolo/modeling/layers/detection_generator_test.py
+++ b/official/projects/yolo/modeling/layers/detection_generator_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,20 +14,22 @@
"""Tests for yolo detection generator."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.modeling.layers import detection_generator as dg
+from official.projects.yolo.modeling.layers import detection_generator
class YoloDecoderTest(parameterized.TestCase, tf.test.TestCase):
@parameterized.parameters(
- (True),
- (False),
+ ('v1', None),
+ ('v2', False),
+ ('v2', True),
+ ('greedy', None),
)
- def test_network_creation(self, nms):
+ def test_network_creation(self, nms_version, use_class_agnostic_nms):
"""Test creation of ResNet family models."""
- tf.keras.backend.set_image_data_format('channels_last')
+ tf_keras.backend.set_image_data_format('channels_last')
input_shape = {
'3': [1, 52, 52, 255],
'4': [1, 26, 26, 255],
@@ -42,7 +44,14 @@ def test_network_creation(self, nms):
box_type = {key: 'scaled' for key in anchors.keys()}
- layer = dg.YoloLayer(anchors, classes, box_type=box_type, max_boxes=10)
+ layer = detection_generator.YoloLayer(
+ anchors,
+ classes,
+ box_type=box_type,
+ max_boxes=10,
+ use_class_agnostic_nms=use_class_agnostic_nms,
+ nms_version=nms_version,
+ )
inputs = {}
for key in input_shape:
diff --git a/official/vision/beta/projects/yolo/modeling/layers/nn_blocks.py b/official/projects/yolo/modeling/layers/nn_blocks.py
similarity index 84%
rename from official/vision/beta/projects/yolo/modeling/layers/nn_blocks.py
rename to official/projects/yolo/modeling/layers/nn_blocks.py
index 5598d1580a1..c2eb209272a 100644
--- a/official/vision/beta/projects/yolo/modeling/layers/nn_blocks.py
+++ b/official/projects/yolo/modeling/layers/nn_blocks.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,21 +13,22 @@
# limitations under the License.
"""Contains common building blocks for yolo neural networks."""
+import functools
from typing import Callable, List, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.modeling import tf_utils
from official.vision.ops import spatial_transform_ops
-class Identity(tf.keras.layers.Layer):
+class Identity(tf_keras.layers.Layer):
def call(self, inputs):
return inputs
-class ConvBN(tf.keras.layers.Layer):
+class ConvBN(tf_keras.layers.Layer):
"""ConvBN block.
Modified Convolution layer to match that of the Darknet Library.
@@ -101,7 +102,7 @@ def __init__(self,
if kernel_initializer == 'VarianceScaling':
# to match pytorch initialization method
- self._kernel_initializer = tf.keras.initializers.VarianceScaling(
+ self._kernel_initializer = tf_keras.initializers.VarianceScaling(
scale=1 / 3, mode='fan_in', distribution='uniform')
else:
self._kernel_initializer = kernel_initializer
@@ -122,16 +123,13 @@ def __init__(self,
if not isinstance(ksize, List) and not isinstance(ksize, Tuple):
ksize = [ksize]
if use_separable_conv and not all([a == 1 for a in ksize]):
- self._conv_base = tf.keras.layers.SeparableConv2D
+ self._conv_base = tf_keras.layers.SeparableConv2D
else:
- self._conv_base = tf.keras.layers.Conv2D
+ self._conv_base = tf_keras.layers.Conv2D
- if use_sync_bn:
- self._bn_base = tf.keras.layers.experimental.SyncBatchNormalization
- else:
- self._bn_base = tf.keras.layers.BatchNormalization
+ self._bn_base = tf_keras.layers.BatchNormalization
- if tf.keras.backend.image_data_format() == 'channels_last':
+ if tf_keras.backend.image_data_format() == 'channels_last':
# format: (batch_size, height, width, channels)
self._bn_axis = -1
else:
@@ -164,12 +162,13 @@ def build(self, input_shape):
self.bn = self._bn_base(
momentum=self._norm_momentum,
epsilon=self._norm_epsilon,
- axis=self._bn_axis)
+ axis=self._bn_axis,
+ synchronized=self._use_sync_bn)
else:
self.bn = None
if self._activation == 'leaky':
- self._activation_fn = tf.keras.layers.LeakyReLU(alpha=self._leaky_alpha)
+ self._activation_fn = tf_keras.layers.LeakyReLU(alpha=self._leaky_alpha)
elif self._activation == 'mish':
self._activation_fn = lambda x: x * tf.math.tanh(tf.math.softplus(x))
else:
@@ -178,7 +177,7 @@ def build(self, input_shape):
def call(self, x):
x = self.conv(x)
if self._use_bn and not self._fuse:
- x = self.bn(x)
+ x = self.bn(x) # pyrefly: ignore[not-callable]
x = self._activation_fn(x)
return x
@@ -239,7 +238,7 @@ def get_config(self):
return layer_config
-class DarkResidual(tf.keras.layers.Layer):
+class DarkResidual(tf_keras.layers.Layer):
"""Darknet block with Residual connection for Yolo v3 Backbone."""
def __init__(self,
@@ -365,9 +364,9 @@ def build(self, input_shape):
padding='same',
**dark_conv_args)
- self._shortcut = tf.keras.layers.Add()
+ self._shortcut = tf_keras.layers.Add()
if self._sc_activation == 'leaky':
- self._activation_fn = tf.keras.layers.LeakyReLU(alpha=self._leaky_alpha)
+ self._activation_fn = tf_keras.layers.LeakyReLU(alpha=self._leaky_alpha)
elif self._sc_activation == 'mish':
self._activation_fn = lambda x: x * tf.math.tanh(tf.math.softplus(x))
else:
@@ -403,13 +402,13 @@ def get_config(self):
return layer_config
-class CSPTiny(tf.keras.layers.Layer):
+class CSPTiny(tf_keras.layers.Layer):
"""CSP Tiny layer.
A Small size convolution block proposed in the CSPNet. The layer uses
shortcuts, routing(concatnation), and feature grouping in order to improve
gradient variablity and allow for high efficency, low power residual learning
- for small networtf.keras.
+ for small networtf_keras.
Cross Stage Partial networks (CSPNets) were proposed in:
[1] Chien-Yao Wang, Hong-Yuan Mark Liao, I-Hau Yeh, Yueh-Hua Wu,
Ping-Yang Chen, Jun-Wei Hsieh
@@ -534,7 +533,7 @@ def build(self, input_shape):
**dark_conv_args)
if self._downsample:
- self._maxpool = tf.keras.layers.MaxPool2D(
+ self._maxpool = tf_keras.layers.MaxPool2D(
pool_size=2, strides=2, padding='same', data_format=None)
super().build(input_shape)
@@ -552,7 +551,7 @@ def call(self, inputs, training=None):
return x, x5
-class CSPRoute(tf.keras.layers.Layer):
+class CSPRoute(tf_keras.layers.Layer):
"""CSPRoute block.
Down sampling layer to take the place of down sampleing done in Residual
@@ -691,7 +690,7 @@ def call(self, inputs, training=None):
return (x, y)
-class CSPConnect(tf.keras.layers.Layer):
+class CSPConnect(tf_keras.layers.Layer):
"""CSPConnect block.
Sister Layer to the CSPRoute layer. Merges the partial feature stacks
@@ -794,7 +793,7 @@ def build(self, input_shape):
kernel_size=self._kernel_size,
strides=(1, 1),
**dark_conv_args)
- self._concat = tf.keras.layers.Concatenate(axis=-1)
+ self._concat = tf_keras.layers.Concatenate(axis=-1)
if not self._drop_final:
self._conv2 = ConvBN(
@@ -815,7 +814,7 @@ def call(self, inputs, training=None):
return x
-class CSPStack(tf.keras.layers.Layer):
+class CSPStack(tf_keras.layers.Layer):
"""CSP Stack layer.
CSP full stack, combines the route and the connect in case you dont want to
@@ -823,7 +822,7 @@ class CSPStack(tf.keras.layers.Layer):
make it a cross stage partial. Added for ease of use. you should be able
to wrap any layer stack with a CSP independent of wether it belongs
to the Darknet family. if filter_scale = 2, then the blocks in the stack
- passed into the the CSP stack should also have filters = filters/filter_scale
+ passed into the CSP stack should also have filters = filters/filter_scale
Cross Stage Partial networks (CSPNets) were proposed in:
[1] Chien-Yao Wang, Hong-Yuan Mark Liao, I-Hau Yeh, Yueh-Hua Wu,
@@ -935,7 +934,7 @@ def call(self, inputs, training=None):
return x
-class PathAggregationBlock(tf.keras.layers.Layer):
+class PathAggregationBlock(tf_keras.layers.Layer):
"""Path Aggregation block."""
def __init__(self,
@@ -1089,7 +1088,7 @@ def build(self, input_shape):
else:
self._build_regular(input_shape, dark_conv_args)
- self._concat = tf.keras.layers.Concatenate()
+ self._concat = tf_keras.layers.Concatenate()
super().build(input_shape)
def _call_regular(self, inputs, training=None):
@@ -1125,7 +1124,7 @@ def call(self, inputs, training=None):
return self._call_regular(inputs, training=training)
-class SPP(tf.keras.layers.Layer):
+class SPP(tf_keras.layers.Layer):
"""Spatial Pyramid Pooling.
A non-agregated SPP layer that uses Pooling.
@@ -1134,14 +1133,14 @@ class SPP(tf.keras.layers.Layer):
def __init__(self, sizes, **kwargs):
self._sizes = list(reversed(sizes))
if not sizes:
- raise ValueError('More than one maxpool should be specified in SSP block')
+ raise ValueError('More than one maxpool should be specified in SPP block')
super().__init__(**kwargs)
def build(self, input_shape):
maxpools = []
for size in self._sizes:
maxpools.append(
- tf.keras.layers.MaxPool2D(
+ tf_keras.layers.MaxPool2D(
pool_size=(size, size),
strides=(1, 1),
padding='same',
@@ -1154,7 +1153,7 @@ def call(self, inputs, training=None):
for maxpool in self._maxpools:
outputs.append(maxpool(inputs))
outputs.append(inputs)
- concat_output = tf.keras.layers.concatenate(outputs)
+ concat_output = tf_keras.layers.concatenate(outputs)
return concat_output
def get_config(self):
@@ -1163,7 +1162,7 @@ def get_config(self):
return layer_config
-class SAM(tf.keras.layers.Layer):
+class SAM(tf_keras.layers.Layer):
"""Spatial Attention Model.
[1] Sanghyun Woo, Jongchan Park, Joon-Young Lee, In So Kweon
@@ -1225,7 +1224,7 @@ def build(self, input_shape):
self._filters = input_shape[-1]
self._conv = ConvBN(filters=self._filters, **self.dark_conv_args)
if self._output_activation == 'leaky':
- self._activation_fn = tf.keras.layers.LeakyReLU(alpha=self._leaky_alpha)
+ self._activation_fn = tf_keras.layers.LeakyReLU(alpha=self._leaky_alpha)
elif self._output_activation == 'mish':
self._activation_fn = lambda x: x * tf.math.tanh(tf.math.softplus(x))
else:
@@ -1243,7 +1242,7 @@ def call(self, inputs, training=None):
return self._activation_fn(inputs * attention_mask)
-class CAM(tf.keras.layers.Layer):
+class CAM(tf_keras.layers.Layer):
"""Channel Attention Model.
[1] Sanghyun Woo, Jongchan Park, Joon-Young Lee, In So Kweon
@@ -1270,16 +1269,12 @@ def __init__(self,
self._reduction_ratio = reduction_ratio
- # use_pooling
- if use_sync_bn:
- self._bn = tf.keras.layers.experimental.SyncBatchNormalization
- else:
- self._bn = tf.keras.layers.BatchNormalization
-
if not use_bn:
self._bn = Identity
self._bn_args = {}
else:
+ self._bn = functools.partial(
+ tf_keras.layers.BatchNormalization, synchronized=use_sync_bn)
self._bn_args = {
'momentum': norm_momentum,
'epsilon': norm_epsilon,
@@ -1302,18 +1297,18 @@ def __init__(self,
def build(self, input_shape):
self._filters = input_shape[-1]
- self._mlp = tf.keras.Sequential([
- tf.keras.layers.Dense(self._filters, **self._mlp_args),
+ self._mlp = tf_keras.Sequential([
+ tf_keras.layers.Dense(self._filters, **self._mlp_args),
self._bn(**self._bn_args),
- tf.keras.layers.Dense(
+ tf_keras.layers.Dense(
int(self._filters * self._reduction_ratio), **self._mlp_args),
self._bn(**self._bn_args),
- tf.keras.layers.Dense(self._filters, **self._mlp_args),
+ tf_keras.layers.Dense(self._filters, **self._mlp_args),
self._bn(**self._bn_args),
])
if self._activation == 'leaky':
- self._activation_fn = tf.keras.layers.LeakyReLU(alpha=self._leaky_alpha)
+ self._activation_fn = tf_keras.layers.LeakyReLU(alpha=self._leaky_alpha)
elif self._activation == 'mish':
self._activation_fn = lambda x: x * tf.math.tanh(tf.math.softplus(x))
else:
@@ -1330,7 +1325,7 @@ def call(self, inputs, training=None):
return inputs * attention_mask
-class CBAM(tf.keras.layers.Layer):
+class CBAM(tf_keras.layers.Layer):
"""Convolutional Block Attention Module.
[1] Sanghyun Woo, Jongchan Park, Joon-Young Lee, In So Kweon
@@ -1403,7 +1398,7 @@ def call(self, inputs, training=None):
return self._sam(self._cam(inputs))
-class DarkRouteProcess(tf.keras.layers.Layer):
+class DarkRouteProcess(tf_keras.layers.Layer):
"""Dark Route Process block.
Process darknet outputs and connect back bone to head more generalizably
@@ -1701,7 +1696,7 @@ def call(self, inputs, training=None):
return self._call_regular(inputs)
-class Reorg(tf.keras.layers.Layer):
+class Reorg(tf_keras.layers.Layer):
"""Splits a high resolution image into 4 lower resolution images.
Used in YOLOR to process very high resolution inputs efficiently.
@@ -1716,3 +1711,245 @@ def call(self, x, training=None):
x[..., 1::2, 1::2, :]
],
axis=-1)
+
+
+class SPPCSPC(tf_keras.layers.Layer):
+ """Cross-stage partial network with spatial pyramid pooling.
+
+ This module is used in YOLOv7 to process backbone feature at the highest
+ level. SPPCSPC uses fusion-first CSP block and it uses SPP within
+ the dense block.
+ """
+
+ def __init__(
+ self,
+ filters,
+ pool_sizes=(5, 9, 13),
+ scale=0.5,
+ kernel_initializer='VarianceScaling',
+ bias_initializer='zeros',
+ kernel_regularizer=None,
+ bias_regularizer=None,
+ use_separable_conv=False,
+ use_bn=True,
+ use_sync_bn=False,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ activation='swish',
+ **kwargs):
+ """Initializes SPPCSPC block.
+
+ Args:
+ filters: an `int` for filters used in Conv2D.
+ pool_sizes: a tuple of `int` for maxpool layer used in the dense block.
+ scale: a `float` scale that applies on the filters to determine the
+ internal Conv2D filters within CSP block.
+ kernel_initializer: string to indicate which function to use to initialize
+ weights in Conv2D.
+ bias_initializer: string to indicate which function to use to initialize
+ bias.
+ kernel_regularizer: string to indicate which function to use to
+ regularizer weights in Conv2D.
+ bias_regularizer: string to indicate which function to use to regularizer
+ bias.
+ use_separable_conv: `bool` wether to use separable convs.
+ use_bn: boolean for whether to use batch normalization.
+ use_sync_bn: boolean for whether sync batch normalization statistics
+ of all batch norm layers to the models global statistics
+ (across all input batches).
+ norm_momentum: float for moment to use for batch normalization.
+ norm_epsilon: float for batch normalization epsilon.
+ activation: string to indicate the activation function used after each
+ Conv2D.
+ **kwargs: other keyword arguments.
+ """
+ super().__init__(**kwargs)
+ self._filters = filters
+ self._pool_sizes = pool_sizes
+ self._scale = scale
+ self._kernel_initializer = kernel_initializer
+ self._bias_initializer = bias_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_regularizer = bias_regularizer
+ self._use_separable_conv = use_separable_conv
+ self._use_bn = use_bn
+ self._use_sync_bn = use_sync_bn
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._activation = activation
+
+ def build(self, input_shape):
+ filters = self._filters * 2 * self._scale
+ conv_op = functools.partial(
+ ConvBN,
+ activation=self._activation,
+ use_separable_conv=self._use_separable_conv,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_initializer=self._bias_initializer,
+ bias_regularizer=self._bias_regularizer,
+ use_bn=self._use_bn,
+ use_sync_bn=self._use_sync_bn,
+ norm_momentum=self._norm_momentum,
+ norm_epsilon=self._norm_epsilon,
+ )
+ self._conv1_1 = conv_op(filters, kernel_size=1, strides=1)
+ self._conv1_2 = conv_op(filters, kernel_size=3, strides=1)
+ self._conv1_3 = conv_op(filters, kernel_size=1, strides=1)
+ self._poolings = [
+ tf_keras.layers.MaxPooling2D(pool_size, strides=1, padding='same')
+ for pool_size in self._pool_sizes
+ ]
+ self._conv1_4 = conv_op(filters, kernel_size=1, strides=1)
+ self._conv1_5 = conv_op(filters, kernel_size=3, strides=1)
+
+ self._conv2_1 = conv_op(filters, kernel_size=1, strides=1)
+
+ self._merge_conv = conv_op(self._filters, kernel_size=1, strides=1)
+ super().build(input_shape)
+
+ def call(self, inputs, training=None):
+ x = self._conv1_3(self._conv1_2(self._conv1_1(inputs)))
+ x = self._conv1_5(
+ self._conv1_4(
+ tf.concat([x] + [pooling(x) for pooling in self._poolings], -1)
+ )
+ )
+ y = self._conv2_1(inputs)
+ return self._merge_conv(tf.concat([x, y], axis=-1))
+
+ def get_config(self):
+ # used to store/share parameters to reconstruct the model
+ layer_config = {
+ 'filters': self._filters,
+ 'pool_sizes': self._pool_sizes,
+ 'scale': self._scale,
+ 'kernel_initializer': self._kernel_initializer,
+ 'bias_initializer': self._bias_initializer,
+ 'kernel_regularizer': self._kernel_regularizer,
+ 'bias_regularizer': self._bias_regularizer,
+ 'use_bn': self._use_bn,
+ 'use_sync_bn': self._use_sync_bn,
+ 'use_separable_conv': self._use_separable_conv,
+ 'norm_momentum': self._norm_momentum,
+ 'norm_epsilon': self._norm_epsilon,
+ 'activation': self._activation,
+ }
+ layer_config.update(super().get_config())
+ return layer_config
+
+
+class RepConv(tf_keras.layers.Layer):
+ """Represented convolution.
+
+ https://arxiv.org/abs/2101.03697
+ """
+
+ def __init__(
+ self,
+ filters,
+ kernel_size=3,
+ strides=1,
+ padding='same',
+ activation='swish',
+ use_separable_conv=False,
+ use_sync_bn=False,
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ kernel_initializer='VarianceScaling',
+ kernel_regularizer=None,
+ bias_initializer='zeros',
+ bias_regularizer=None,
+ **kwargs
+ ):
+ """Initializes RepConv layer.
+
+ Args:
+ filters: integer for output depth, or the number of features to learn.
+ kernel_size: integer or tuple for the shape of the weight matrix or kernel
+ to learn.
+ strides: integer of tuple how much to move the kernel after each kernel
+ use.
+ padding: string 'valid' or 'same', if same, then pad the image, else do
+ not.
+ activation: string or None for activation function to use in layer, if
+ None activation is replaced by linear.
+ use_separable_conv: `bool` wether to use separable convs.
+ use_sync_bn: boolean for whether sync batch normalization statistics of
+ all batch norm layers to the models global statistics (across all input
+ batches).
+ norm_momentum: float for moment to use for batch normalization.
+ norm_epsilon: float for batch normalization epsilon.
+ kernel_initializer: string to indicate which function to use to initialize
+ weights.
+ kernel_regularizer: string to indicate which function to use to
+ regularizer weights.
+ bias_initializer: string to indicate which function to use to initialize
+ bias.
+ bias_regularizer: string to indicate which function to use to regularizer
+ bias.
+ **kwargs: other keyword arguments.
+ """
+ super().__init__(**kwargs)
+ self._filters = filters
+ self._kernel_size = kernel_size
+ self._strides = strides
+ self._padding = padding
+ self._activation = activation
+ self._use_separable_conv = use_separable_conv
+ self._use_sync_bn = use_sync_bn
+ self._norm_momentum = norm_momentum
+ self._norm_epsilon = norm_epsilon
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_initializer = bias_initializer
+ self._bias_regularizer = bias_regularizer
+ # For deploy.
+ self._fuse = False
+
+ def build(self, input_shape):
+ conv_op = functools.partial(
+ tf_keras.layers.SeparableConv2D
+ if self._use_separable_conv
+ else tf_keras.layers.Conv2D,
+ filters=self._filters,
+ strides=self._strides,
+ padding=self._padding,
+ kernel_initializer=self._kernel_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ bias_initializer=self._bias_initializer,
+ bias_regularizer=self._bias_regularizer,
+ )
+ bn_op = functools.partial(
+ tf_keras.layers.BatchNormalization,
+ synchronized=self._use_sync_bn,
+ momentum=self._norm_momentum,
+ epsilon=self._norm_epsilon,
+ )
+
+ self._activation_fn = tf_utils.get_activation(self._activation)
+ self._rbr_reparam = conv_op(kernel_size=self._kernel_size, use_bias=True)
+ if input_shape[-1] == self._filters and self._strides == 1:
+ self._rbr_identity = bn_op()
+ self._rbr_dense = conv_op(kernel_size=self._kernel_size, use_bias=False)
+ self._rbr_dense_bn = bn_op()
+ self._rbr_1x1 = conv_op(kernel_size=1, use_bias=False)
+ self._rbr_1x1_bn = bn_op()
+
+ def call(self, inputs, training=None):
+ if self._fuse:
+ return self._activation_fn(self._rbr_reparam(inputs))
+
+ id_out = 0
+ if hasattr(self, '_rbr_identity'):
+ id_out = self._rbr_identity(inputs)
+
+ x = self._rbr_dense_bn(self._rbr_dense(inputs))
+ y = self._rbr_1x1_bn(self._rbr_1x1(inputs))
+ return self._activation_fn(x + y + id_out)
+
+ def fuse(self):
+ if self._fuse:
+ return
+ # TODO(b/264495198): Implement fuse for RepConv.
+ raise NotImplementedError()
diff --git a/official/vision/beta/projects/yolo/modeling/layers/nn_blocks_test.py b/official/projects/yolo/modeling/layers/nn_blocks_test.py
similarity index 71%
rename from official/vision/beta/projects/yolo/modeling/layers/nn_blocks_test.py
rename to official/projects/yolo/modeling/layers/nn_blocks_test.py
index 5e88c09f4a2..04729f23b68 100644
--- a/official/vision/beta/projects/yolo/modeling/layers/nn_blocks_test.py
+++ b/official/projects/yolo/modeling/layers/nn_blocks_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,9 +14,9 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.modeling.layers import nn_blocks
+from official.projects.yolo.modeling.layers import nn_blocks
class CSPConnectTest(tf.test.TestCase, parameterized.TestCase):
@@ -24,7 +24,7 @@ class CSPConnectTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.named_parameters(('same', 224, 224, 64, 1),
('downsample', 224, 224, 64, 2))
def test_pass_through(self, width, height, filters, mod):
- x = tf.keras.Input(shape=(width, height, filters))
+ x = tf_keras.Input(shape=(width, height, filters))
test_layer = nn_blocks.CSPRoute(filters=filters, filter_scale=mod)
test_layer2 = nn_blocks.CSPConnect(filters=filters, filter_scale=mod)
outx, px = test_layer(x)
@@ -39,8 +39,8 @@ def test_pass_through(self, width, height, filters, mod):
@parameterized.named_parameters(('same', 224, 224, 64, 1),
('downsample', 224, 224, 128, 2))
def test_gradient_pass_though(self, filters, width, height, mod):
- loss = tf.keras.losses.MeanSquaredError()
- optimizer = tf.keras.optimizers.SGD()
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
test_layer = nn_blocks.CSPRoute(filters, filter_scale=mod)
path_layer = nn_blocks.CSPConnect(filters, filter_scale=mod)
@@ -68,7 +68,7 @@ class CSPRouteTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.named_parameters(('same', 224, 224, 64, 1),
('downsample', 224, 224, 64, 2))
def test_pass_through(self, width, height, filters, mod):
- x = tf.keras.Input(shape=(width, height, filters))
+ x = tf_keras.Input(shape=(width, height, filters))
test_layer = nn_blocks.CSPRoute(filters=filters, filter_scale=mod)
outx, _ = test_layer(x)
print(outx)
@@ -81,8 +81,8 @@ def test_pass_through(self, width, height, filters, mod):
@parameterized.named_parameters(('same', 224, 224, 64, 1),
('downsample', 224, 224, 128, 2))
def test_gradient_pass_though(self, filters, width, height, mod):
- loss = tf.keras.losses.MeanSquaredError()
- optimizer = tf.keras.optimizers.SGD()
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
test_layer = nn_blocks.CSPRoute(filters, filter_scale=mod)
path_layer = nn_blocks.CSPConnect(filters, filter_scale=mod)
@@ -115,7 +115,7 @@ def test_pass_through(self, kernel_size, padding, strides):
pad_const = 1
else:
pad_const = 0
- x = tf.keras.Input(shape=(224, 224, 3))
+ x = tf_keras.Input(shape=(224, 224, 3))
test_layer = nn_blocks.ConvBN(
filters=64,
kernel_size=kernel_size,
@@ -134,8 +134,8 @@ def test_pass_through(self, kernel_size, padding, strides):
@parameterized.named_parameters(('filters', 3))
def test_gradient_pass_though(self, filters):
- loss = tf.keras.losses.MeanSquaredError()
- optimizer = tf.keras.optimizers.SGD()
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
with tf.device('/CPU:0'):
test_layer = nn_blocks.ConvBN(filters, kernel_size=(3, 3), padding='same')
@@ -162,7 +162,7 @@ def test_pass_through(self, width, height, filters, downsample):
mod = 1
if downsample:
mod = 2
- x = tf.keras.Input(shape=(width, height, filters))
+ x = tf_keras.Input(shape=(width, height, filters))
test_layer = nn_blocks.DarkResidual(filters=filters, downsample=downsample)
outx = test_layer(x)
print(outx)
@@ -176,8 +176,8 @@ def test_pass_through(self, width, height, filters, downsample):
('downsample', 32, 223, 223, True),
('oddball', 32, 223, 223, False))
def test_gradient_pass_though(self, filters, width, height, downsample):
- loss = tf.keras.losses.MeanSquaredError()
- optimizer = tf.keras.optimizers.SGD()
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
test_layer = nn_blocks.DarkResidual(filters, downsample=downsample)
if downsample:
@@ -209,7 +209,7 @@ class DarkSppTest(tf.test.TestCase, parameterized.TestCase):
('test1', 300, 300, 10, [2, 3, 4, 5]),
('test2', 256, 256, 5, [10]))
def test_pass_through(self, width, height, channels, sizes):
- x = tf.keras.Input(shape=(width, height, channels))
+ x = tf_keras.Input(shape=(width, height, channels))
test_layer = nn_blocks.SPP(sizes=sizes)
outx = test_layer(x)
self.assertAllEqual(outx.shape.as_list(),
@@ -220,8 +220,8 @@ def test_pass_through(self, width, height, channels, sizes):
('test1', 300, 300, 10, [2, 3, 4, 5]),
('test2', 256, 256, 5, [10]))
def test_gradient_pass_though(self, width, height, channels, sizes):
- loss = tf.keras.losses.MeanSquaredError()
- optimizer = tf.keras.optimizers.SGD()
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
test_layer = nn_blocks.SPP(sizes=sizes)
init = tf.random_normal_initializer()
@@ -249,7 +249,7 @@ class DarkRouteProcessTest(tf.test.TestCase, parameterized.TestCase):
('test1', 224, 224, 64, 7, False), ('test2', 223, 223, 32, 3, False),
('tiny', 223, 223, 16, 1, False), ('spp', 224, 224, 64, 7, False))
def test_pass_through(self, width, height, filters, repetitions, spp):
- x = tf.keras.Input(shape=(width, height, filters))
+ x = tf_keras.Input(shape=(width, height, filters))
test_layer = nn_blocks.DarkRouteProcess(
filters=filters, repetitions=repetitions, insert_spp=spp)
outx = test_layer(x)
@@ -270,8 +270,8 @@ def test_pass_through(self, width, height, filters, repetitions, spp):
('test1', 224, 224, 64, 7, False), ('test2', 223, 223, 32, 3, False),
('tiny', 223, 223, 16, 1, False), ('spp', 224, 224, 64, 7, False))
def test_gradient_pass_though(self, width, height, filters, repetitions, spp):
- loss = tf.keras.losses.MeanSquaredError()
- optimizer = tf.keras.optimizers.SGD()
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
test_layer = nn_blocks.DarkRouteProcess(
filters=filters, repetitions=repetitions, insert_spp=spp)
@@ -301,5 +301,80 @@ def test_gradient_pass_though(self, width, height, filters, repetitions, spp):
return
+class SPPCSPCTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(('SPPCSPC', 224, 224, 8, [5, 9, 13], 0.5),
+ ('test1', 300, 300, 32, [2, 3, 4, 5], 1.0),
+ ('test2', 256, 256, 16, [10], 2.0))
+ def test_pass_through(self, width, height, filters, pool_sizes, scale):
+ x = tf_keras.Input(shape=(width, height, filters))
+ test_layer = nn_blocks.SPPCSPC(filters, pool_sizes, scale)
+ out = test_layer(x)
+ self.assertAllEqual(out.shape.as_list(), [None, width, height, filters])
+
+ @parameterized.named_parameters(('SPPCSPC', 224, 224, 8, [5, 9, 13], 0.5),
+ ('test1', 300, 300, 32, [2, 3, 4, 5], 1.0),
+ ('test2', 256, 256, 16, [10], 2.0))
+ def test_gradient_pass_though(
+ self, width, height, filters, pool_sizes, scale):
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
+ test_layer = nn_blocks.SPPCSPC(filters, pool_sizes, scale)
+
+ init = tf.random_normal_initializer()
+ x = tf.Variable(
+ initial_value=init(shape=(1, width, height, filters), dtype=tf.float32))
+ y = tf.Variable(
+ initial_value=init(shape=(1, width, height, filters), dtype=tf.float32))
+
+ with tf.GradientTape() as tape:
+ x_hat = test_layer(x)
+ grad_loss = loss(x_hat, y)
+ grad = tape.gradient(grad_loss, test_layer.trainable_variables)
+ optimizer.apply_gradients(zip(grad, test_layer.trainable_variables))
+
+ self.assertNotIn(None, grad)
+ return
+
+
+class RepConvTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(('RepConv', 224, 224, 8, 1),
+ ('test1', 300, 300, 32, 2),
+ ('test2', 256, 256, 16, 4))
+ def test_pass_through(self, width, height, filters, strides):
+ x = tf_keras.Input(shape=(width, height, filters))
+ test_layer = nn_blocks.RepConv(filters, strides=strides)
+ out = test_layer(x)
+ self.assertAllEqual(out.shape.as_list(),
+ [None, width // strides, height // strides, filters])
+
+ @parameterized.named_parameters(('RepConv', 224, 224, 8, 1),
+ ('test1', 300, 300, 32, 2),
+ ('test2', 256, 256, 16, 4))
+ def test_gradient_pass_though(self, width, height, filters, strides):
+ loss = tf_keras.losses.MeanSquaredError()
+ optimizer = tf_keras.optimizers.SGD()
+ test_layer = nn_blocks.RepConv(filters, strides=strides)
+
+ init = tf.random_normal_initializer()
+ x = tf.Variable(
+ initial_value=init(shape=(1, width, height, filters), dtype=tf.float32))
+ y = tf.Variable(
+ initial_value=init(
+ shape=(1, width // strides, height // strides, filters),
+ dtype=tf.float32,
+ )
+ )
+
+ with tf.GradientTape() as tape:
+ x_hat = test_layer(x)
+ grad_loss = loss(x_hat, y)
+ grad = tape.gradient(grad_loss, test_layer.trainable_variables)
+ optimizer.apply_gradients(zip(grad, test_layer.trainable_variables))
+
+ self.assertNotIn(None, grad)
+ return
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/vision/beta/projects/yolo/modeling/yolo_model.py b/official/projects/yolo/modeling/yolo_model.py
similarity index 79%
rename from official/vision/beta/projects/yolo/modeling/yolo_model.py
rename to official/projects/yolo/modeling/yolo_model.py
index 9748bba660e..3cdf49173de 100644
--- a/official/vision/beta/projects/yolo/modeling/yolo_model.py
+++ b/official/projects/yolo/modeling/yolo_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,12 +14,12 @@
"""Yolo models."""
-from typing import Mapping, Union
-import tensorflow as tf
-from official.vision.beta.projects.yolo.modeling.layers import nn_blocks
+from typing import Mapping, Union, Any, Dict
+import tensorflow as tf, tf_keras
+from official.projects.yolo.modeling.layers import nn_blocks
-class Yolo(tf.keras.Model):
+class Yolo(tf_keras.Model):
"""The YOLO model class."""
def __init__(self,
@@ -31,8 +31,8 @@ def __init__(self,
"""Detection initialization function.
Args:
- backbone: `tf.keras.Model` a backbone network.
- decoder: `tf.keras.Model` a decoder network.
+ backbone: `tf_keras.Model` a backbone network.
+ decoder: `tf_keras.Model` a decoder network.
head: `RetinaNetHead`, the RetinaNet head.
detection_generator: the detection generator.
**kwargs: keyword arguments to be passed.
@@ -52,9 +52,11 @@ def __init__(self,
self._head = head
self._detection_generator = detection_generator
self._fused = False
- return
- def call(self, inputs, training=False):
+ def call(self, # pytype: disable=annotation-type-mismatch
+ inputs: tf.Tensor,
+ training: bool = None, # pyrefly: ignore[bad-function-definition]
+ mask: Any = None) -> Dict[str, tf.Tensor]:
maps = self.backbone(inputs)
decoded_maps = self.decoder(maps)
raw_predictions = self.head(decoded_maps)
@@ -86,12 +88,12 @@ def get_config(self):
return self._config_dict
@classmethod
- def from_config(cls, config):
+ def from_config(cls, config): # pyrefly: ignore[bad-override]
return cls(**config)
@property
def checkpoint_items(
- self) -> Mapping[str, Union[tf.keras.Model, tf.keras.layers.Layer]]:
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
"""Returns a dictionary of items to be additionally checkpointed."""
items = dict(backbone=self.backbone, head=self.head)
if self.decoder is not None:
diff --git a/official/projects/yolo/modeling/yolov7_model.py b/official/projects/yolo/modeling/yolov7_model.py
new file mode 100644
index 00000000000..61ad93aafd6
--- /dev/null
+++ b/official/projects/yolo/modeling/yolov7_model.py
@@ -0,0 +1,109 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""YOLOv7 models."""
+
+from typing import Mapping, Union, Any, Dict
+from absl import logging
+import tensorflow as tf, tf_keras
+from official.projects.yolo.modeling.layers import nn_blocks
+
+
+class YoloV7(tf_keras.Model):
+ """The YOLOv7 model class."""
+
+ def __init__(self, backbone, decoder, head, detection_generator, **kwargs):
+ """Detection initialization function.
+
+ Args:
+ backbone: `tf_keras.Model` a backbone network.
+ decoder: `tf_keras.Model` a decoder network.
+ head: `RetinaNetHead`, the RetinaNet head.
+ detection_generator: the detection generator.
+ **kwargs: keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+
+ self._config_dict = {
+ 'backbone': backbone,
+ 'decoder': decoder,
+ 'head': head,
+ 'detection_generator': detection_generator
+ }
+
+ # model components
+ self._backbone = backbone
+ self._decoder = decoder
+ self._head = head
+ self._detection_generator = detection_generator
+ self._fused = False
+ return
+
+ def call(self, # pytype: disable=annotation-type-mismatch
+ inputs: tf.Tensor,
+ training: bool = None, # pyrefly: ignore[bad-function-definition]
+ mask: Any = None) -> Dict[str, tf.Tensor]:
+ backbone_outputs = self.backbone(inputs)
+ decoder_outputs = self.decoder(backbone_outputs)
+ raw_outputs = self.head(decoder_outputs)
+ if training:
+ return {'raw_output': raw_outputs}
+ else:
+ # Post-processing.
+ predictions = self.detection_generator(raw_outputs)
+ predictions.update({'raw_output': raw_outputs})
+ return predictions
+
+ @property
+ def backbone(self):
+ return self._backbone
+
+ @property
+ def decoder(self):
+ return self._decoder
+
+ @property
+ def head(self):
+ return self._head
+
+ @property
+ def detection_generator(self):
+ return self._detection_generator
+
+ def get_config(self):
+ return self._config_dict
+
+ @classmethod
+ def from_config(cls, config): # pyrefly: ignore[bad-override]
+ return cls(**config)
+
+ @property
+ def checkpoint_items(
+ self) -> Mapping[str, Union[tf_keras.Model, tf_keras.layers.Layer]]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ items = dict(backbone=self.backbone, head=self.head)
+ if self.decoder is not None:
+ items.update(decoder=self.decoder)
+ return items
+
+ def fuse(self):
+ """Performs re-parameterization on ConvBN and RepConv layers."""
+ logging.info('Fusing ConvBN and RepConv layers.')
+ if not self._fused:
+ self._fused = True
+ for layer in self.submodules:
+ if isinstance(layer, (nn_blocks.ConvBN, nn_blocks.RepConv)):
+ layer.fuse()
+ self.summary()
+ return
diff --git a/official/vision/beta/projects/yolo/ops/__init__.py b/official/projects/yolo/ops/__init__.py
similarity index 89%
rename from official/vision/beta/projects/yolo/ops/__init__.py
rename to official/projects/yolo/ops/__init__.py
index ba97902e7ec..41caa388f95 100644
--- a/official/vision/beta/projects/yolo/ops/__init__.py
+++ b/official/projects/yolo/ops/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/beta/projects/yolo/ops/anchor.py b/official/projects/yolo/ops/anchor.py
similarity index 98%
rename from official/vision/beta/projects/yolo/ops/anchor.py
rename to official/projects/yolo/ops/anchor.py
index b73c44f3585..b32cf55b9b9 100644
--- a/official/vision/beta/projects/yolo/ops/anchor.py
+++ b/official/projects/yolo/ops/anchor.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,11 +14,11 @@
"""Yolo Anchor labler."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import box_ops
-from official.vision.beta.projects.yolo.ops import loss_utils
-from official.vision.beta.projects.yolo.ops import preprocessing_ops
+from official.projects.yolo.ops import box_ops
+from official.projects.yolo.ops import loss_utils
+from official.projects.yolo.ops import preprocessing_ops
INF = 10000000
diff --git a/official/vision/beta/projects/yolo/ops/box_ops.py b/official/projects/yolo/ops/box_ops.py
similarity index 98%
rename from official/vision/beta/projects/yolo/ops/box_ops.py
rename to official/projects/yolo/ops/box_ops.py
index a674c927363..3a46710ac68 100644
--- a/official/vision/beta/projects/yolo/ops/box_ops.py
+++ b/official/projects/yolo/ops/box_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,8 +14,8 @@
"""Yolo box ops."""
import math
-import tensorflow as tf
-from official.vision.beta.projects.yolo.ops import math_ops
+import tensorflow as tf, tf_keras
+from official.projects.yolo.ops import math_ops
def yxyx_to_xcycwh(box: tf.Tensor):
diff --git a/official/vision/beta/projects/yolo/ops/box_ops_test.py b/official/projects/yolo/ops/box_ops_test.py
similarity index 92%
rename from official/vision/beta/projects/yolo/ops/box_ops_test.py
rename to official/projects/yolo/ops/box_ops_test.py
index 17a83c4b008..a683595fbd1 100644
--- a/official/vision/beta/projects/yolo/ops/box_ops_test.py
+++ b/official/projects/yolo/ops/box_ops_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
"""box_ops tests."""
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import box_ops
+from official.projects.yolo.ops import box_ops
class InputUtilsTest(parameterized.TestCase, tf.test.TestCase):
diff --git a/official/projects/yolo/ops/initializer_ops.py b/official/projects/yolo/ops/initializer_ops.py
new file mode 100644
index 00000000000..f3dbe6c8e79
--- /dev/null
+++ b/official/projects/yolo/ops/initializer_ops.py
@@ -0,0 +1,26 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Yolo initializer ops."""
+
+import tensorflow as tf, tf_keras
+
+
+def pytorch_kernel_initializer(kernel_initializer):
+ """Prepare kernel weights initializer to match PyTorch implementation."""
+ if kernel_initializer == 'VarianceScaling':
+ return tf_keras.initializers.VarianceScaling(
+ scale=1 / 3, mode='fan_in', distribution='uniform'
+ )
+ return kernel_initializer
diff --git a/official/vision/beta/projects/yolo/ops/kmeans_anchors.py b/official/projects/yolo/ops/kmeans_anchors.py
similarity index 98%
rename from official/vision/beta/projects/yolo/ops/kmeans_anchors.py
rename to official/projects/yolo/ops/kmeans_anchors.py
index 0f4a6117108..252303cb77f 100644
--- a/official/vision/beta/projects/yolo/ops/kmeans_anchors.py
+++ b/official/projects/yolo/ops/kmeans_anchors.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,10 +16,10 @@
import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import input_reader
-from official.vision.beta.projects.yolo.ops import box_ops
+from official.projects.yolo.ops import box_ops
def _iou(x, centroids_x, iou_type="iou"):
@@ -290,7 +290,7 @@ def read(self,
k-means for.
num_samples: `Optional[int]` for the number of samples to use for kmeans,
typically about 5000 samples are all that are needed, but for the best
- results use None to run the entire dataset.
+ results use -1 to run the entire dataset.
Returns:
boxes: `List[List[int]]` of shape [k, 2] for the anchor boxes to use for
diff --git a/official/vision/beta/projects/yolo/ops/kmeans_anchors_test.py b/official/projects/yolo/ops/kmeans_anchors_test.py
similarity index 89%
rename from official/vision/beta/projects/yolo/ops/kmeans_anchors_test.py
rename to official/projects/yolo/ops/kmeans_anchors_test.py
index 894da7542f1..5fda4cacd00 100644
--- a/official/vision/beta/projects/yolo/ops/kmeans_anchors_test.py
+++ b/official/projects/yolo/ops/kmeans_anchors_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
"""kmeans_test tests."""
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import kmeans_anchors
+from official.projects.yolo.ops import kmeans_anchors
class KMeansTest(parameterized.TestCase, tf.test.TestCase):
diff --git a/official/vision/beta/projects/yolo/ops/loss_utils.py b/official/projects/yolo/ops/loss_utils.py
old mode 100755
new mode 100644
similarity index 98%
rename from official/vision/beta/projects/yolo/ops/loss_utils.py
rename to official/projects/yolo/ops/loss_utils.py
index 53909ae2b7f..3d4d4018beb
--- a/official/vision/beta/projects/yolo/ops/loss_utils.py
+++ b/official/projects/yolo/ops/loss_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,10 @@
"""Yolo loss utility functions."""
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import box_ops
-from official.vision.beta.projects.yolo.ops import math_ops
+from official.projects.yolo.ops import box_ops
+from official.projects.yolo.ops import math_ops
@tf.custom_gradient
@@ -43,7 +43,7 @@ def sigmoid_bce(y, x_prime, label_smoothing):
This simplification is used by the darknet library in order to improve
training stability. The gradient is almost the same
- as tf.keras.losses.binary_crossentropy but varies slightly and
+ as tf_keras.losses.binary_crossentropy but varies slightly and
yields different performance.
Args:
@@ -171,7 +171,7 @@ def __init__(self, anchors, scale_anchors=None):
scale_anchors: An `int` for how much to scale this level to get the
original input shape.
"""
- self.dtype = tf.keras.backend.floatx()
+ self.dtype = tf_keras.backend.floatx()
self._scale_anchors = scale_anchors
self._anchors = tf.convert_to_tensor(anchors)
return
@@ -206,7 +206,7 @@ def _extend_batch(self, grid, batch_size):
def __call__(self, height, width, batch_size, dtype=None):
if dtype is None:
- self.dtype = tf.keras.backend.floatx()
+ self.dtype = tf_keras.backend.floatx()
else:
self.dtype = dtype
grid_points = self._build_grid_points(height, width, self._anchors,
diff --git a/official/vision/beta/projects/yolo/ops/math_ops.py b/official/projects/yolo/ops/math_ops.py
similarity index 94%
rename from official/vision/beta/projects/yolo/ops/math_ops.py
rename to official/projects/yolo/ops/math_ops.py
index 7a42288c15c..9af75322e5b 100644
--- a/official/vision/beta/projects/yolo/ops/math_ops.py
+++ b/official/projects/yolo/ops/math_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""A set of private math operations used to safely implement the YOLO loss."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
def rm_nan_inf(x, val=0.0):
diff --git a/official/vision/beta/projects/yolo/ops/mosaic.py b/official/projects/yolo/ops/mosaic.py
old mode 100755
new mode 100644
similarity index 57%
rename from official/vision/beta/projects/yolo/ops/mosaic.py
rename to official/projects/yolo/ops/mosaic.py
index e3982de2df8..e4a95a54ac3
--- a/official/vision/beta/projects/yolo/ops/mosaic.py
+++ b/official/projects/yolo/ops/mosaic.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,40 +15,46 @@
"""Mosaic op."""
import random
-import tensorflow as tf
-import tensorflow_addons as tfa
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import preprocessing_ops
+from official.projects.yolo.ops import preprocessing_ops
+from official.vision.ops import augment
from official.vision.ops import box_ops
from official.vision.ops import preprocess_ops
class Mosaic:
- """Stitch together sets of 4 images to generate samples with more boxes."""
-
- def __init__(self,
- output_size,
- mosaic_frequency=1.0,
- mixup_frequency=0.0,
- letter_box=True,
- jitter=0.0,
- mosaic_crop_mode='scale',
- mosaic_center=0.25,
- aug_scale_min=1.0,
- aug_scale_max=1.0,
- aug_rand_angle=0.0,
- aug_rand_perspective=0.0,
- aug_rand_translate=0.0,
- random_pad=False,
- random_flip=False,
- area_thresh=0.1,
- pad_value=preprocessing_ops.PAD_VALUE,
- seed=None):
+ """Stitch together sets of 4 (2x2) or 9 (3x3) images to generate samples with more boxes."""
+
+ def __init__(
+ self,
+ output_size,
+ mosaic_frequency=1.0,
+ mosaic9_frequency=0.0,
+ mixup_frequency=0.0,
+ letter_box=True,
+ jitter=0.0,
+ mosaic_crop_mode='scale',
+ mosaic_center=0.25,
+ mosaic9_center=0.33,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ aug_rand_angle=0.0,
+ aug_rand_perspective=0.0,
+ aug_rand_translate=0.0,
+ random_pad=False,
+ random_flip=False,
+ area_thresh=0.1,
+ pad_value=preprocessing_ops.PAD_VALUE,
+ seed=None,
+ ):
"""Initializes parameters for mosaic.
Args:
output_size: `Tensor` or `List` for [height, width] of output image.
mosaic_frequency: `float` indicating how often to apply mosaic.
+ mosaic9_frequency: `float` indicating how often to apply a 3x3 mosaic
+ instead of 2x2.
mixup_frequency: `float` indicating how often to apply mixup.
letter_box: `boolean` indicating whether upon start of the datapipeline
regardless of the preprocessing ops that are used, the aspect ratio of
@@ -61,7 +67,9 @@ def __init__(self,
image, and None will default to scale and apply no post processing to
the created mosaic.
mosaic_center: `float` indicating how much to randomly deviate from the
- from the center of the image when creating a mosaic.
+ center of the image when creating a mosaic.
+ mosaic9_center: `float` indicating how much to randomly deviate from the
+ center of the image when creating a mosaic9.
aug_scale_min: `float` indicating the minimum scaling value for image
scale jitter.
aug_scale_max: `float` indicating the maximum scaling value for image
@@ -85,6 +93,7 @@ def __init__(self,
self._area_thresh = area_thresh
self._mosaic_frequency = mosaic_frequency
+ self._mosaic9_frequency = mosaic9_frequency
self._mixup_frequency = mixup_frequency
self._letter_box = letter_box
@@ -92,6 +101,7 @@ def __init__(self,
self._mosaic_crop_mode = mosaic_crop_mode
self._mosaic_center = mosaic_center
+ self._mosaic9_center = mosaic9_center
self._aug_scale_min = aug_scale_min
self._aug_scale_max = aug_scale_max
@@ -105,10 +115,10 @@ def __init__(self,
self._deterministic = seed is not None
self._seed = seed if seed is not None else random.randint(0, 2**30)
- def _generate_cut(self):
+ def _generate_cut(self, num_tiles, mosaic_center):
"""Generate a random center to use for slicing and patching the images."""
if self._mosaic_crop_mode == 'crop':
- min_offset = self._mosaic_center
+ min_offset = mosaic_center
cut_x = preprocessing_ops.random_uniform_strong(
self._output_size[1] * min_offset,
self._output_size[1] * (1 - min_offset),
@@ -122,25 +132,28 @@ def _generate_cut(self):
[self._output_size[0], self._output_size[1], 3])
else:
cut = None
- ishape = tf.convert_to_tensor(
- [self._output_size[0] * 2, self._output_size[1] * 2, 3])
+ ishape = tf.convert_to_tensor([
+ self._output_size[0] * num_tiles,
+ self._output_size[1] * num_tiles,
+ 3,
+ ])
return cut, ishape
- def scale_boxes(self, patch, ishape, boxes, classes, xs, ys):
+ def scale_boxes(self, patch, ishape, boxes, x_offset, y_offset):
"""Scale and translate the boxes for each image prior to patching."""
- xs = tf.cast(xs, boxes.dtype)
- ys = tf.cast(ys, boxes.dtype)
+ x_offset = tf.cast(x_offset, boxes.dtype)
+ y_offset = tf.cast(y_offset, boxes.dtype)
pshape = tf.cast(tf.shape(patch), boxes.dtype)
ishape = tf.cast(ishape, boxes.dtype)
- translate = tf.cast((ishape - pshape), boxes.dtype)
+ y_offset = ishape[0] * y_offset
+ x_offset = ishape[1] * x_offset
boxes = box_ops.denormalize_boxes(boxes, pshape[:2])
- boxes = boxes + tf.cast([
- translate[0] * ys, translate[1] * xs, translate[0] * ys,
- translate[1] * xs
- ], boxes.dtype)
+ boxes = boxes + tf.cast(
+ [y_offset, x_offset, y_offset, x_offset], boxes.dtype
+ )
boxes = box_ops.normalize_boxes(boxes, ishape[:2])
- return boxes, classes
+ return boxes
def _select_ind(self, inds, *args):
items = []
@@ -148,15 +161,18 @@ def _select_ind(self, inds, *args):
items.append(tf.gather(item, inds))
return items
- def _augment_image(self,
- image,
- boxes,
- classes,
- is_crowd,
- area,
- xs=0.0,
- ys=0.0,
- cut=None):
+ def _augment_image(
+ self,
+ image,
+ boxes,
+ classes,
+ is_crowd,
+ area,
+ xs=0.0,
+ ys=0.0,
+ cut=None,
+ letter_box=False,
+ ):
"""Process a single image prior to the application of patching."""
if self._random_flip:
# Randomly flip the image horizontally.
@@ -165,14 +181,16 @@ def _augment_image(self,
# Augment the image without resizing
image, infos, crop_points = preprocessing_ops.resize_and_jitter_image(
- image, [self._output_size[0], self._output_size[1]],
+ image,
+ [self._output_size[0], self._output_size[1]],
random_pad=False,
- letter_box=self._letter_box,
+ letter_box=letter_box,
jitter=self._random_crop,
shiftx=xs,
shifty=ys,
cut=cut,
- seed=self._seed)
+ seed=self._seed,
+ )
# Clip and clean boxes.
boxes, inds = preprocessing_ops.transform_and_clip_boxes(
@@ -185,11 +203,12 @@ def _augment_image(self,
classes, is_crowd, area = self._select_ind(inds, classes, is_crowd, area) # pylint:disable=unbalanced-tuple-unpacking
return image, boxes, classes, is_crowd, area, crop_points
- def _mosaic_crop_image(self, image, boxes, classes, is_crowd, area):
+ def _mosaic_crop_image(
+ self, image, boxes, classes, is_crowd, area, mosaic_center):
"""Process a patched image in preperation for final output."""
if self._mosaic_crop_mode != 'crop':
shape = tf.cast(preprocessing_ops.get_image_shape(image), tf.float32)
- center = shape * self._mosaic_center
+ center = shape * mosaic_center
# shift the center of the image by applying a translation to the whole
# image
@@ -200,8 +219,10 @@ def _mosaic_crop_image(self, image, boxes, classes, is_crowd, area):
preprocessing_ops.random_uniform_strong(
-center[1], center[1], seed=self._seed))
- # clip the boxes to those with in the image
- image = tfa.image.translate(image, [cw, ch], fill_value=self._pad_value)
+ # clip the boxes to fit within the image
+ image = augment.translate(
+ image, [cw, ch], fill_value=self._pad_value, fill_mode='constant'
+ )
boxes = box_ops.denormalize_boxes(boxes, shape[:2])
boxes = boxes + tf.cast([ch, cw, ch, cw], boxes.dtype)
boxes = box_ops.clip_boxes(boxes, shape[:2])
@@ -235,16 +256,22 @@ def _mosaic_crop_image(self, image, boxes, classes, is_crowd, area):
return image, boxes, classes, is_crowd, area, area
# mosaic full frequency doubles model speed
- def _process_image(self, sample, shiftx, shifty, cut, ishape):
- """Process and augment each image."""
+ def _process_image(self, sample, shiftx, shifty, cut, letter_box):
+ """Process and augment an image."""
(image, boxes, classes, is_crowd, area, crop_points) = self._augment_image(
- sample['image'], sample['groundtruth_boxes'],
- sample['groundtruth_classes'], sample['groundtruth_is_crowd'],
- sample['groundtruth_area'], shiftx, shifty, cut)
-
- (boxes, classes) = self.scale_boxes(image, ishape, boxes, classes,
- 1 - shiftx, 1 - shifty)
-
+ sample['image'],
+ sample['groundtruth_boxes'],
+ sample['groundtruth_classes'],
+ sample['groundtruth_is_crowd'],
+ sample['groundtruth_area'],
+ shiftx,
+ shifty,
+ cut,
+ letter_box,
+ )
+
+ # Make a copy so this method is functional.
+ sample = sample.copy()
sample['image'] = image
sample['groundtruth_boxes'] = boxes
sample['groundtruth_classes'] = classes
@@ -255,44 +282,31 @@ def _process_image(self, sample, shiftx, shifty, cut, ishape):
sample['crop_points'] = crop_points
return sample
- def _patch2(self, one, two):
- """Stitch together 2 images in totality."""
- sample = one
- sample['image'] = tf.concat([one['image'], two['image']], axis=-2)
-
- sample['groundtruth_boxes'] = tf.concat(
- [one['groundtruth_boxes'], two['groundtruth_boxes']], axis=0)
- sample['groundtruth_classes'] = tf.concat(
- [one['groundtruth_classes'], two['groundtruth_classes']], axis=0)
- sample['groundtruth_is_crowd'] = tf.concat(
- [one['groundtruth_is_crowd'], two['groundtruth_is_crowd']], axis=0)
- sample['groundtruth_area'] = tf.concat(
- [one['groundtruth_area'], two['groundtruth_area']], axis=0)
- return sample
-
- def _patch(self, one, two):
- """Build the full 4 patch of images from sets of 2 images."""
- image = tf.concat([one['image'], two['image']], axis=-3)
- boxes = tf.concat([one['groundtruth_boxes'], two['groundtruth_boxes']],
- axis=0)
- classes = tf.concat(
- [one['groundtruth_classes'], two['groundtruth_classes']], axis=0)
- is_crowd = tf.concat(
- [one['groundtruth_is_crowd'], two['groundtruth_is_crowd']], axis=0)
- area = tf.concat([one['groundtruth_area'], two['groundtruth_area']], axis=0)
+ def _update_patched_sample(
+ self, sample, image, boxes, classes, is_crowds, areas, mosaic_center
+ ):
+ """Returns a shallow copy of sample with updated values."""
+ boxes = tf.concat(boxes, axis=0)
+ classes = tf.concat(classes, axis=0)
+ is_crowds = tf.concat(is_crowds, axis=0)
+ areas = tf.concat(areas, axis=0)
if self._mosaic_crop_mode is not None:
- image, boxes, classes, is_crowd, area, _ = self._mosaic_crop_image(
- image, boxes, classes, is_crowd, area)
+ image, boxes, classes, is_crowds, areas, _ = self._mosaic_crop_image(
+ image, boxes, classes, is_crowds, areas, mosaic_center
+ )
- sample = one
height, width = preprocessing_ops.get_image_shape(image)
+ # Shallow copy of dict is needed to keep this method functional and
+ # AutoGraph happy.
+ sample = sample.copy()
sample['image'] = tf.cast(image, tf.uint8)
sample['groundtruth_boxes'] = boxes
- sample['groundtruth_area'] = area
- sample['groundtruth_classes'] = tf.cast(classes,
- sample['groundtruth_classes'].dtype)
- sample['groundtruth_is_crowd'] = tf.cast(is_crowd, tf.bool)
+ sample['groundtruth_area'] = areas
+ sample['groundtruth_classes'] = tf.cast(
+ classes, sample['groundtruth_classes'].dtype
+ )
+ sample['groundtruth_is_crowd'] = tf.cast(is_crowds, tf.bool)
sample['width'] = tf.cast(width, sample['width'].dtype)
sample['height'] = tf.cast(height, sample['height'].dtype)
sample['num_detections'] = tf.shape(sample['groundtruth_boxes'])[1]
@@ -301,29 +315,112 @@ def _patch(self, one, two):
del sample['shiftx']
del sample['shifty']
del sample['crop_points']
+
return sample
- def _mosaic(self, one, two, three, four):
- """Stitch together 4 images to build a mosaic."""
+ def _patch(self, patches, ishape, num_rows, num_cols, mosaic_center):
+ """Combines patches into a num_patches x num_patches mosaic and translates the bounding boxes."""
+ rows = []
+ for row_idx in range(num_rows):
+ row_patches = [
+ patches[row_idx * num_cols + col_idx]['image']
+ for col_idx in range(num_cols)
+ ]
+ rows.append(tf.concat(row_patches, axis=-2))
+ image = tf.concat(rows, axis=-3)
+
+ boxes = []
+ classes = []
+ is_crowds = []
+ areas = []
+ # Shift boxes to their new coordinates in the mosaic.
+ for row_idx in range(num_rows):
+ for col_idx in range(num_cols):
+ patch = patches[row_idx * num_cols + col_idx]
+ transformed_boxes = self.scale_boxes(
+ patch['image'],
+ ishape,
+ patch['groundtruth_boxes'],
+ col_idx / num_cols,
+ row_idx / num_rows,
+ )
+ boxes.append(transformed_boxes)
+ classes.append(patch['groundtruth_classes'])
+ is_crowds.append(patch['groundtruth_is_crowd'])
+ areas.append(patch['groundtruth_area'])
+
+ return self._update_patched_sample(
+ patches[0], image, boxes, classes, is_crowds, areas, mosaic_center
+ )
+
+ def _mosaic(self, *patch_samples):
+ """Builds a 2x2 or 3x3 mosaic."""
if self._mosaic_frequency >= 1.0:
- domo = 1.0
+ mosaic_prob = 1.0
else:
- domo = preprocessing_ops.random_uniform_strong(
- 0.0, 1.0, dtype=tf.float32, seed=self._seed)
- noop = one.copy()
-
- if domo >= (1 - self._mosaic_frequency):
- cut, ishape = self._generate_cut()
- one = self._process_image(one, 1.0, 1.0, cut, ishape)
- two = self._process_image(two, 0.0, 1.0, cut, ishape)
- three = self._process_image(three, 1.0, 0.0, cut, ishape)
- four = self._process_image(four, 0.0, 0.0, cut, ishape)
- patch1 = self._patch2(one, two)
- patch2 = self._patch2(three, four)
- stitched = self._patch(patch1, patch2)
- return stitched
+ mosaic_prob = preprocessing_ops.random_uniform_strong(
+ 0.0, 1.0, dtype=tf.float32, seed=self._seed
+ )
+ sample = patch_samples[0].copy()
+
+ if mosaic_prob >= (1 - self._mosaic_frequency):
+ mosaic9_prob = preprocessing_ops.random_uniform_strong(
+ 0.0, 1.0, dtype=tf.float32, seed=self._seed + 1
+ )
+ if self._mosaic9_frequency > 0 and mosaic9_prob >= (
+ 1 - self._mosaic9_frequency
+ ):
+ return self._mosaic9(*patch_samples)
+ else:
+ return self._mosaic4(*patch_samples)
else:
- return self._add_param(noop)
+ return self._add_param(sample)
+
+ def _mosaic4(self, *samples):
+ """Stitches together 4 images to build a 2x2 mosaic."""
+ cut, ishape = self._generate_cut(2, self._mosaic_center)
+ samples = [
+ self._process_image(
+ samples[0], 1.0, 1.0, cut, letter_box=self._letter_box
+ ),
+ self._process_image(
+ samples[1], 0.0, 1.0, cut, letter_box=self._letter_box
+ ),
+ self._process_image(
+ samples[2], 1.0, 0.0, cut, letter_box=self._letter_box
+ ),
+ self._process_image(
+ samples[3], 0.0, 0.0, cut, letter_box=self._letter_box
+ ),
+ ]
+ stitched = self._patch(samples, ishape, 2, 2, self._mosaic_center)
+ return stitched
+
+ def _mosaic9(self, *samples):
+ """Stitches together 9 images to build a 3x3 mosaic."""
+ cut, ishape = self._generate_cut(3, self._mosaic9_center)
+ # Only corner images can be letterboxed to prevent gaps in the image.
+ samples = [
+ self._process_image(
+ samples[0], 1.0, 1.0, cut, letter_box=self._letter_box
+ ),
+ self._process_image(samples[1], 0.0, 0.0, cut, letter_box=False),
+ self._process_image(
+ samples[2], 0.0, 1.0, cut, letter_box=self._letter_box
+ ),
+ self._process_image(samples[3], 0.0, 0.0, cut, letter_box=False),
+ self._process_image(samples[4], 0.0, 0.0, cut, letter_box=False),
+ self._process_image(samples[5], 0.0, 0.0, cut, letter_box=False),
+ self._process_image(
+ samples[6], 1.0, 0.0, cut, letter_box=self._letter_box
+ ),
+ self._process_image(samples[7], 0.0, 0.0, cut, letter_box=False),
+ self._process_image(
+ samples[8], 0.0, 0.0, cut, letter_box=self._letter_box
+ ),
+ ]
+ stitched = self._patch(samples, ishape, 3, 3, self._mosaic9_center)
+ return stitched
def _beta(self, alpha, beta):
"""Generates a random number using the beta distribution."""
@@ -364,7 +461,8 @@ def _mixup(self, one, two):
def _add_param(self, sample):
"""Add parameters to handle skipped images."""
- sample['is_mosaic'] = tf.cast(0.0, tf.bool)
+ if 'is_mosaic' not in sample:
+ sample['is_mosaic'] = tf.cast(0.0, tf.bool)
sample['num_detections'] = tf.shape(sample['groundtruth_boxes'])[0]
return sample
@@ -372,23 +470,29 @@ def _apply(self, dataset):
"""Apply mosaic to an input dataset."""
determ = self._deterministic
dataset = dataset.prefetch(tf.data.AUTOTUNE)
- one = dataset.shuffle(100, seed=self._seed, reshuffle_each_iteration=True)
- two = dataset.shuffle(
- 100, seed=self._seed + 1, reshuffle_each_iteration=True)
- three = dataset.shuffle(
- 100, seed=self._seed + 2, reshuffle_each_iteration=True)
- four = dataset.shuffle(
- 100, seed=self._seed + 3, reshuffle_each_iteration=True)
-
- dataset = tf.data.Dataset.zip((one, two, three, four))
+
+ patch_datasets = []
+ num_patches = 9 if self._mosaic9_frequency > 0.0 else 4
+ for i in range(num_patches):
+ patch_datasets.append(
+ dataset.shuffle(
+ 100, seed=self._seed + i, reshuffle_each_iteration=True
+ )
+ )
+
+ dataset = tf.data.Dataset.zip(tuple(patch_datasets))
dataset = dataset.map(
self._mosaic, num_parallel_calls=tf.data.AUTOTUNE, deterministic=determ)
if self._mixup_frequency > 0:
one = dataset.shuffle(
- 100, seed=self._seed + 4, reshuffle_each_iteration=True)
+ 100, seed=self._seed + num_patches, reshuffle_each_iteration=True
+ )
two = dataset.shuffle(
- 100, seed=self._seed + 5, reshuffle_each_iteration=True)
+ 100,
+ seed=self._seed + num_patches + 1,
+ reshuffle_each_iteration=True,
+ )
dataset = tf.data.Dataset.zip((one, two))
dataset = dataset.map(
self._mixup,
diff --git a/official/vision/beta/projects/yolo/ops/preprocessing_ops.py b/official/projects/yolo/ops/preprocessing_ops.py
old mode 100755
new mode 100644
similarity index 99%
rename from official/vision/beta/projects/yolo/ops/preprocessing_ops.py
rename to official/projects/yolo/ops/preprocessing_ops.py
index 93c8b156922..b35778106bb
--- a/official/vision/beta/projects/yolo/ops/preprocessing_ops.py
+++ b/official/projects/yolo/ops/preprocessing_ops.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,9 @@
import random
import numpy as np
-import tensorflow as tf
-import tensorflow_addons as tfa
+import tensorflow as tf, tf_keras
+from official.vision.ops import augment
from official.vision.ops import box_ops as bbox_ops
PAD_VALUE = 114
@@ -669,12 +669,14 @@ def affine_warp_image(image,
affine = tf.cast(affine[:-1], tf.float32)
# Apply the transformation to image.
- image = tfa.image.transform(
+ image = augment.transform(
image,
affine,
fill_value=PAD_VALUE,
output_shape=desired_size,
- interpolation='bilinear')
+ interpolation='bilinear',
+ fill_mode='constant',
+ )
desired_size = tf.cast(desired_size, tf.float32)
affine_info = [image_size, desired_size, affine_boxes]
diff --git a/official/vision/beta/projects/yolo/ops/preprocessing_ops_test.py b/official/projects/yolo/ops/preprocessing_ops_test.py
old mode 100755
new mode 100644
similarity index 97%
rename from official/vision/beta/projects/yolo/ops/preprocessing_ops_test.py
rename to official/projects/yolo/ops/preprocessing_ops_test.py
index c2f477684ba..91aace3d1bb
--- a/official/vision/beta/projects/yolo/ops/preprocessing_ops_test.py
+++ b/official/projects/yolo/ops/preprocessing_ops_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
"""Tests for preprocessing_ops.py."""
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import preprocessing_ops
+from official.projects.yolo.ops import preprocessing_ops
from official.vision.ops import box_ops as bbox_ops
diff --git a/official/vision/beta/projects/yolo/optimization/__init__.py b/official/projects/yolo/optimization/__init__.py
old mode 100755
new mode 100644
similarity index 68%
rename from official/vision/beta/projects/yolo/optimization/__init__.py
rename to official/projects/yolo/optimization/__init__.py
index 06a9588cae8..8a84e4ad689
--- a/official/vision/beta/projects/yolo/optimization/__init__.py
+++ b/official/projects/yolo/optimization/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,6 +17,6 @@
# pylint: disable=wildcard-import
from official.modeling.optimization.configs.learning_rate_config import *
from official.modeling.optimization.ema_optimizer import ExponentialMovingAverage
-from official.vision.beta.projects.yolo.optimization.configs.optimization_config import *
-from official.vision.beta.projects.yolo.optimization.configs.optimizer_config import *
-from official.vision.beta.projects.yolo.optimization.optimizer_factory import OptimizerFactory as YoloOptimizerFactory
+from official.projects.yolo.optimization.configs.optimization_config import *
+from official.projects.yolo.optimization.configs.optimizer_config import *
+from official.projects.yolo.optimization.optimizer_factory import OptimizerFactory as YoloOptimizerFactory
diff --git a/official/projects/yolo/optimization/configs/__init__.py b/official/projects/yolo/optimization/configs/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/optimization/configs/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/optimization/configs/optimization_config.py b/official/projects/yolo/optimization/configs/optimization_config.py
old mode 100755
new mode 100644
similarity index 83%
rename from official/vision/beta/projects/yolo/optimization/configs/optimization_config.py
rename to official/projects/yolo/optimization/configs/optimization_config.py
index 7314a9c2db9..5902e769f78
--- a/official/vision/beta/projects/yolo/optimization/configs/optimization_config.py
+++ b/official/projects/yolo/optimization/configs/optimization_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,7 +22,7 @@
from typing import Optional
from official.modeling.optimization.configs import optimization_config as optimization_cfg
-from official.vision.beta.projects.yolo.optimization.configs import optimizer_config as opt_cfg
+from official.projects.yolo.optimization.configs import optimizer_config as opt_cfg
@dataclasses.dataclass
@@ -38,7 +38,9 @@ class OptimizerConfig(optimization_cfg.OptimizerConfig):
rmsprop: rmsprop optimizer.
"""
type: Optional[str] = None
- sgd_torch: opt_cfg.SGDTorchConfig = opt_cfg.SGDTorchConfig()
+ sgd_torch: opt_cfg.SGDTorchConfig = dataclasses.field(
+ default_factory=opt_cfg.SGDTorchConfig
+ )
@dataclasses.dataclass
@@ -53,4 +55,6 @@ class OptimizationConfig(optimization_cfg.OptimizationConfig):
warmup: warmup oneof config.
"""
type: Optional[str] = None
- optimizer: OptimizerConfig = OptimizerConfig()
+ optimizer: OptimizerConfig = dataclasses.field(
+ default_factory=OptimizerConfig
+ )
diff --git a/official/vision/beta/projects/yolo/optimization/configs/optimizer_config.py b/official/projects/yolo/optimization/configs/optimizer_config.py
old mode 100755
new mode 100644
similarity index 94%
rename from official/vision/beta/projects/yolo/optimization/configs/optimizer_config.py
rename to official/projects/yolo/optimization/configs/optimizer_config.py
index 46c9609649c..021e35b952b
--- a/official/vision/beta/projects/yolo/optimization/configs/optimizer_config.py
+++ b/official/projects/yolo/optimization/configs/optimizer_config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -41,7 +41,7 @@ class BaseOptimizerConfig(base_config.Config):
class SGDTorchConfig(optimizer_config.BaseOptimizerConfig):
"""Configuration for SGD optimizer.
- The attributes for this class matches the arguments of tf.keras.optimizer.SGD.
+ The attributes for this class matches the arguments of tf_keras.optimizer.SGD.
Attributes:
name: name of the optimizer.
diff --git a/official/vision/beta/projects/yolo/optimization/optimizer_factory.py b/official/projects/yolo/optimization/optimizer_factory.py
old mode 100755
new mode 100644
similarity index 91%
rename from official/vision/beta/projects/yolo/optimization/optimizer_factory.py
rename to official/projects/yolo/optimization/optimizer_factory.py
index e66082d62ef..55b075c98c0
--- a/official/vision/beta/projects/yolo/optimization/optimizer_factory.py
+++ b/official/projects/yolo/optimization/optimizer_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,13 +18,13 @@
from official.modeling.optimization import ema_optimizer
from official.modeling.optimization import optimizer_factory
-from official.vision.beta.projects.yolo.optimization import sgd_torch
+from official.projects.yolo.optimization import sgd_torch
-optimizer_factory.OPTIMIZERS_CLS.update({
+optimizer_factory.LEGACY_OPTIMIZERS_CLS.update({
'sgd_torch': sgd_torch.SGDTorch,
})
-OPTIMIZERS_CLS = optimizer_factory.OPTIMIZERS_CLS
+OPTIMIZERS_CLS = optimizer_factory.LEGACY_OPTIMIZERS_CLS
LR_CLS = optimizer_factory.LR_CLS
WARMUP_CLS = optimizer_factory.WARMUP_CLS
@@ -73,7 +73,7 @@ def get_bias_lr_schedule(self, bias_lr):
bias_lr: learning rate config.
Returns:
- tf.keras.optimizers.schedules.LearningRateSchedule instance. If
+ tf_keras.optimizers.schedules.LearningRateSchedule instance. If
learning rate type is consant, lr_config.learning_rate is returned.
"""
if self._lr_type == 'constant':
diff --git a/official/vision/beta/projects/yolo/optimization/sgd_torch.py b/official/projects/yolo/optimization/sgd_torch.py
similarity index 98%
rename from official/vision/beta/projects/yolo/optimization/sgd_torch.py
rename to official/projects/yolo/optimization/sgd_torch.py
index 5f372a2c5b6..975f3c55e8e 100644
--- a/official/vision/beta/projects/yolo/optimization/sgd_torch.py
+++ b/official/projects/yolo/optimization/sgd_torch.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,9 @@
import re
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-LearningRateSchedule = tf.keras.optimizers.schedules.LearningRateSchedule
+LearningRateSchedule = tf_keras.optimizers.schedules.LearningRateSchedule
def _var_key(var):
@@ -43,7 +43,7 @@ def _var_key(var):
return var._unique_id
-class SGDTorch(tf.keras.optimizers.Optimizer):
+class SGDTorch(tf_keras.optimizers.legacy.Optimizer):
"""Optimizer that simulates the SGD module used in pytorch.
diff --git a/official/projects/yolo/serving/__init__.py b/official/projects/yolo/serving/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/serving/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/projects/yolo/serving/export_module_factory.py b/official/projects/yolo/serving/export_module_factory.py
new file mode 100644
index 00000000000..ac9fdfd30b8
--- /dev/null
+++ b/official/projects/yolo/serving/export_module_factory.py
@@ -0,0 +1,264 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Factory for YOLO export modules."""
+
+from typing import Any, Callable, Dict, List, Optional, Text, Union
+
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import export_base
+from official.projects.yolo.configs import darknet_classification
+from official.projects.yolo.configs import yolo
+from official.projects.yolo.configs import yolov7
+from official.projects.yolo.dataloaders import classification_input
+from official.projects.yolo.modeling import factory as yolo_factory
+from official.projects.yolo.modeling.backbones import darknet # pylint: disable=unused-import
+from official.projects.yolo.modeling.decoders import yolo_decoder # pylint: disable=unused-import
+from official.projects.yolo.serving import model_fn as yolo_model_fn
+from official.vision.modeling import factory
+from official.vision.serving import export_utils
+
+
+class ExportModule(export_base.ExportModule):
+ """Base Export Module."""
+
+ def __init__(self,
+ params: cfg.ExperimentConfig,
+ model: tf_keras.Model,
+ input_signature: Union[tf.TensorSpec, Dict[str, tf.TensorSpec]],
+ preprocessor: Optional[Callable[..., Any]] = None,
+ inference_step: Optional[Callable[..., Any]] = None,
+ postprocessor: Optional[Callable[..., Any]] = None,
+ eval_postprocessor: Optional[Callable[..., Any]] = None):
+ """Initializes a module for export.
+
+ Args:
+ params: A dataclass for parameters to the module.
+ model: A tf_keras.Model instance to be exported.
+ input_signature: tf.TensorSpec, e.g. tf.TensorSpec(shape=[None, 224, 224,
+ 3], dtype=tf.uint8)
+ preprocessor: An optional callable to preprocess the inputs.
+ inference_step: An optional callable to forward-pass the model.
+ postprocessor: An optional callable to postprocess the model outputs.
+ eval_postprocessor: An optional callable to postprocess model outputs used
+ for model evaluation.
+ """
+ super().__init__(
+ params,
+ model=model,
+ preprocessor=preprocessor,
+ inference_step=inference_step,
+ postprocessor=postprocessor)
+ self.eval_postprocessor = eval_postprocessor
+ self.input_signature = input_signature
+
+ @tf.function
+ def serve(self, inputs: Any) -> Any:
+ x = self.preprocessor(inputs=inputs) if self.preprocessor else inputs
+ x = self.inference_step(x)
+ x = self.postprocessor(x) if self.postprocessor else x
+ return x
+
+ @tf.function
+ def serve_eval(self, inputs: Any) -> Any:
+ x = self.preprocessor(inputs=inputs) if self.preprocessor else inputs
+ x = self.inference_step(x)
+ x = self.eval_postprocessor(x) if self.eval_postprocessor else x
+ return x
+
+ def get_inference_signatures(
+ self, function_keys: Dict[Text, Text]):
+ """Gets defined function signatures.
+
+ Args:
+ function_keys: A dictionary with keys as the function to create signature
+ for and values as the signature keys when returns.
+
+ Returns:
+ A dictionary with key as signature key and value as concrete functions
+ that can be used for tf.saved_model.save.
+ """
+ signatures = {}
+ for _, def_name in function_keys.items():
+ if 'eval' in def_name and self.eval_postprocessor:
+ signatures[def_name] = self.serve_eval.get_concrete_function(
+ self.input_signature)
+ else:
+ signatures[def_name] = self.serve.get_concrete_function(
+ self.input_signature)
+ return signatures
+
+
+def create_classification_export_module(
+ params: cfg.ExperimentConfig,
+ input_type: str,
+ batch_size: int,
+ input_image_size: List[int],
+ num_channels: int = 3,
+ input_name: Optional[str] = None) -> ExportModule:
+ """Creates classification export module."""
+ input_signature = export_utils.get_image_input_signatures(
+ input_type, batch_size, input_image_size, num_channels, input_name)
+ input_specs = tf_keras.layers.InputSpec(shape=[batch_size] +
+ input_image_size + [num_channels])
+
+ model = factory.build_classification_model(
+ input_specs=input_specs,
+ model_config=params.task.model,
+ l2_regularizer=None)
+
+ def preprocess_fn(inputs):
+ image_tensor = export_utils.parse_image(inputs, input_type,
+ input_image_size, num_channels)
+ # If input_type is `tflite`, do not apply image preprocessing.
+ if input_type == 'tflite':
+ return image_tensor
+
+ def preprocess_image_fn(inputs):
+ return classification_input.Parser.inference_fn(inputs, input_image_size,
+ num_channels)
+
+ images = tf.map_fn(
+ preprocess_image_fn,
+ elems=image_tensor,
+ fn_output_signature=tf.TensorSpec(
+ shape=input_image_size + [num_channels], dtype=tf.float32))
+
+ return images
+
+ def postprocess_fn(logits):
+ probs = tf.nn.softmax(logits)
+ return {'logits': logits, 'probs': probs}
+
+ export_module = ExportModule(
+ params,
+ model=model,
+ input_signature=input_signature,
+ preprocessor=preprocess_fn,
+ postprocessor=postprocess_fn)
+ return export_module
+
+
+def create_yolo_export_module(
+ params: cfg.ExperimentConfig,
+ input_type: str,
+ batch_size: int,
+ input_image_size: List[int],
+ num_channels: int = 3,
+ input_name: Optional[str] = None) -> ExportModule:
+ """Creates YOLO export module."""
+ input_signature = export_utils.get_image_input_signatures(
+ input_type, batch_size, input_image_size, num_channels, input_name)
+ input_specs = tf_keras.layers.InputSpec(shape=[batch_size] +
+ input_image_size + [num_channels])
+ if isinstance(params.task, yolo.YoloTask):
+ model, _ = yolo_factory.build_yolo(
+ input_specs=input_specs,
+ model_config=params.task.model,
+ l2_regularization=None)
+ elif isinstance(params.task, yolov7.YoloV7Task):
+ model = yolo_factory.build_yolov7(
+ input_specs=input_specs,
+ model_config=params.task.model,
+ l2_regularization=None)
+
+ def preprocess_fn(inputs):
+ image_tensor = export_utils.parse_image(inputs, input_type,
+ input_image_size, num_channels)
+
+ def normalize_image_fn(inputs):
+ image = tf.cast(inputs, dtype=tf.float32)
+ return image / 255.0
+
+ # If input_type is `tflite`, do not apply image preprocessing. Only apply
+ # normalization.
+ if input_type == 'tflite':
+ return normalize_image_fn(image_tensor), None
+
+ def preprocess_image_fn(inputs):
+ image = normalize_image_fn(inputs)
+ (image, image_info) = yolo_model_fn.letterbox(
+ image,
+ input_image_size,
+ letter_box=params.task.validation_data.parser.letter_box)
+ return image, image_info
+
+ images_spec = tf.TensorSpec(shape=input_image_size + [3], dtype=tf.float32)
+
+ image_info_spec = tf.TensorSpec(shape=[4, 2], dtype=tf.float32)
+
+ images, image_info = tf.nest.map_structure(
+ tf.identity,
+ tf.map_fn(
+ preprocess_image_fn,
+ elems=image_tensor,
+ fn_output_signature=(images_spec, image_info_spec),
+ parallel_iterations=32))
+
+ return images, image_info
+
+ def inference_steps(inputs, model):
+ images, image_info = inputs
+ detection = model.call(images, training=False)
+ if input_type != 'tflite':
+ detection['bbox'] = yolo_model_fn.undo_info(
+ detection['bbox'],
+ detection['num_detections'],
+ image_info,
+ expand=False,
+ )
+
+ final_outputs = {
+ 'detection_boxes': detection['bbox'],
+ 'detection_scores': detection['confidence'],
+ 'detection_classes': detection['classes'],
+ 'num_detections': detection['num_detections']
+ }
+
+ return final_outputs
+
+ export_module = ExportModule(
+ params,
+ model=model, # pyrefly: ignore[unbound-name]
+ input_signature=input_signature,
+ preprocessor=preprocess_fn,
+ inference_step=inference_steps)
+
+ return export_module
+
+
+def get_export_module(params: cfg.ExperimentConfig,
+ input_type: str,
+ batch_size: Optional[int],
+ input_image_size: List[int],
+ num_channels: int = 3,
+ input_name: Optional[str] = None) -> ExportModule:
+ """Factory for export modules."""
+ if isinstance(params.task,
+ darknet_classification.ImageClassificationTask):
+ export_module = create_classification_export_module(params, input_type,
+ batch_size, # pyrefly: ignore[bad-argument-type]
+ input_image_size,
+ num_channels,
+ input_name)
+ elif isinstance(params.task, (yolo.YoloTask, yolov7.YoloV7Task)):
+ export_module = create_yolo_export_module(params, input_type, batch_size, # pyrefly: ignore[bad-argument-type]
+ input_image_size, num_channels,
+ input_name)
+ else:
+ raise ValueError('Export module not implemented for {} task.'.format(
+ type(params.task)))
+ return export_module
diff --git a/official/vision/beta/projects/yolo/serving/export_saved_model.py b/official/projects/yolo/serving/export_saved_model.py
similarity index 74%
rename from official/vision/beta/projects/yolo/serving/export_saved_model.py
rename to official/projects/yolo/serving/export_saved_model.py
index 42bd79cf490..7bf9abc4acd 100644
--- a/official/vision/beta/projects/yolo/serving/export_saved_model.py
+++ b/official/projects/yolo/serving/export_saved_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -38,9 +38,8 @@
from official.core import exp_factory
from official.modeling import hyperparams
-from official.vision.beta.projects.yolo.configs import yolo as cfg # pylint: disable=unused-import
-from official.vision.beta.projects.yolo.serving import export_module_factory
-from official.vision.beta.projects.yolo.tasks import yolo as task # pylint: disable=unused-import
+from official.projects.yolo.common import registry_imports # pylint: disable=unused-import
+from official.projects.yolo.serving import export_module_factory
from official.vision.serving import export_saved_model_lib
FLAGS = flags.FLAGS
@@ -69,18 +68,36 @@
'input_image_size', '224,224',
'The comma-separated string of two integers representing the height,width '
'of the input to the model.')
+_EXPORT_SAVED_MODEL_SUBDIR = flags.DEFINE_string(
+ 'export_saved_model_subdir', 'saved_model',
+ 'The subdirectory for saved model.')
+_INPUT_NAME = flags.DEFINE_string(
+ 'input_name', None,
+ 'Input tensor name in signature def. Default at None which'
+ 'produces input tensor name `inputs`.')
def main(_):
params = exp_factory.get_exp_config(FLAGS.experiment)
for config_file in FLAGS.config_file or []:
- params = hyperparams.override_params_dict(
- params, config_file, is_strict=True)
+ try:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=True
+ )
+ except KeyError:
+ params = hyperparams.override_params_dict(
+ params, config_file, is_strict=False
+ )
if FLAGS.params_override:
- params = hyperparams.override_params_dict(
- params, FLAGS.params_override, is_strict=True)
-
+ try:
+ params = hyperparams.override_params_dict(
+ params, FLAGS.params_override, is_strict=True
+ )
+ except KeyError:
+ params = hyperparams.override_params_dict(
+ params, FLAGS.params_override, is_strict=False
+ )
params.validate()
params.lock()
@@ -91,7 +108,8 @@ def main(_):
input_type=FLAGS.input_type,
batch_size=FLAGS.batch_size,
input_image_size=[int(x) for x in FLAGS.input_image_size.split(',')],
- num_channels=3)
+ num_channels=3,
+ input_name=_INPUT_NAME.value)
export_saved_model_lib.export_inference_graph(
input_type=FLAGS.input_type,
@@ -100,7 +118,8 @@ def main(_):
params=params,
checkpoint_path=FLAGS.checkpoint_path,
export_dir=FLAGS.export_dir,
- export_module=export_module)
+ export_module=export_module,
+ export_saved_model_subdir=_EXPORT_SAVED_MODEL_SUBDIR.value)
if __name__ == '__main__':
diff --git a/official/projects/yolo/serving/export_tflite.py b/official/projects/yolo/serving/export_tflite.py
new file mode 100644
index 00000000000..46c5bbc805f
--- /dev/null
+++ b/official/projects/yolo/serving/export_tflite.py
@@ -0,0 +1,23 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Binary to convert a saved model to TFLite model for Yolo model."""
+
+from absl import app
+
+from official.projects.yolo.common import registry_imports # pylint: disable=unused-import
+from official.vision.serving import export_tflite
+
+if __name__ == '__main__':
+ app.run(export_tflite.main)
diff --git a/official/vision/beta/projects/yolo/serving/model_fn.py b/official/projects/yolo/serving/model_fn.py
similarity index 94%
rename from official/vision/beta/projects/yolo/serving/model_fn.py
rename to official/projects/yolo/serving/model_fn.py
index c1b78ca907e..e03d4251d6a 100644
--- a/official/vision/beta/projects/yolo/serving/model_fn.py
+++ b/official/projects/yolo/serving/model_fn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,9 +16,9 @@
from typing import List, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
-from official.vision.beta.projects.yolo.ops import preprocessing_ops
+from official.projects.yolo.ops import preprocessing_ops
from official.vision.ops import box_ops
diff --git a/official/projects/yolo/tasks/__init__.py b/official/projects/yolo/tasks/__init__.py
new file mode 100644
index 00000000000..e7e7c21950e
--- /dev/null
+++ b/official/projects/yolo/tasks/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
diff --git a/official/vision/beta/projects/yolo/tasks/image_classification.py b/official/projects/yolo/tasks/image_classification.py
similarity index 91%
rename from official/vision/beta/projects/yolo/tasks/image_classification.py
rename to official/projects/yolo/tasks/image_classification.py
index b69bee618e7..80554ccd27e 100644
--- a/official/vision/beta/projects/yolo/tasks/image_classification.py
+++ b/official/projects/yolo/tasks/image_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,8 +15,8 @@
"""Image classification task definition."""
from official.common import dataset_fn
from official.core import task_factory
-from official.vision.beta.projects.yolo.configs import darknet_classification as exp_cfg
-from official.vision.beta.projects.yolo.dataloaders import classification_input
+from official.projects.yolo.configs import darknet_classification as exp_cfg
+from official.projects.yolo.dataloaders import classification_input
from official.vision.dataloaders import classification_input as classification_input_base
from official.vision.dataloaders import input_reader_factory
from official.vision.dataloaders import tfds_factory
diff --git a/official/vision/beta/projects/yolo/tasks/task_utils.py b/official/projects/yolo/tasks/task_utils.py
similarity index 89%
rename from official/vision/beta/projects/yolo/tasks/task_utils.py
rename to official/projects/yolo/tasks/task_utils.py
index 9a14f49104b..b12675bcfac 100644
--- a/official/vision/beta/projects/yolo/tasks/task_utils.py
+++ b/official/projects/yolo/tasks/task_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,7 @@
# limitations under the License.
"""Utils for yolo task."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
class ListMetrics:
@@ -29,7 +29,7 @@ def build_metric(self):
metric_names = self._metric_names
metrics = []
for name in metric_names:
- metrics.append(tf.keras.metrics.Mean(name, dtype=tf.float32))
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
return metrics
def update_state(self, loss_metrics):
diff --git a/official/vision/beta/projects/yolo/tasks/yolo.py b/official/projects/yolo/tasks/yolo.py
old mode 100755
new mode 100644
similarity index 90%
rename from official/vision/beta/projects/yolo/tasks/yolo.py
rename to official/projects/yolo/tasks/yolo.py
index 9d6c5943ec9..f27e3def262
--- a/official/vision/beta/projects/yolo/tasks/yolo.py
+++ b/official/projects/yolo/tasks/yolo.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
from typing import Optional
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import dataset_fn
from official.core import base_task
@@ -26,15 +26,15 @@
from official.core import input_reader
from official.core import task_factory
from official.modeling import performance
-from official.vision.beta.projects.yolo import optimization
-from official.vision.beta.projects.yolo.configs import yolo as exp_cfg
-from official.vision.beta.projects.yolo.dataloaders import tf_example_decoder
-from official.vision.beta.projects.yolo.dataloaders import yolo_input
-from official.vision.beta.projects.yolo.modeling import factory
-from official.vision.beta.projects.yolo.ops import kmeans_anchors
-from official.vision.beta.projects.yolo.ops import mosaic
-from official.vision.beta.projects.yolo.ops import preprocessing_ops
-from official.vision.beta.projects.yolo.tasks import task_utils
+from official.projects.yolo import optimization
+from official.projects.yolo.configs import yolo as exp_cfg
+from official.projects.yolo.dataloaders import tf_example_decoder
+from official.projects.yolo.dataloaders import yolo_input
+from official.projects.yolo.modeling import factory
+from official.projects.yolo.ops import kmeans_anchors
+from official.projects.yolo.ops import mosaic
+from official.projects.yolo.ops import preprocessing_ops
+from official.projects.yolo.tasks import task_utils
from official.vision.dataloaders import tfds_factory
from official.vision.dataloaders import tf_example_label_map_decoder
from official.vision.evaluation import coco_evaluator
@@ -84,7 +84,8 @@ def generate_anchors(self):
dataset.global_batch_size = 1
box_reader = kmeans_anchors.BoxGenInputReader(
dataset,
- dataset_fn=tf.data.TFRecordDataset,
+ dataset_fn=dataset_fn.pick_dataset_fn(
+ self.task_config.train_data.file_type),
decoder_fn=decoder.decode)
boxes = box_reader.read(
@@ -112,11 +113,13 @@ def build_model(self):
l2_weight_decay = self.task_config.weight_decay / 2.0
input_size = model_base_cfg.input_size.copy()
- input_specs = tf.keras.layers.InputSpec(shape=[None] + input_size)
+ input_specs = tf_keras.layers.InputSpec(shape=[None] + input_size)
l2_regularizer = (
- tf.keras.regularizers.l2(l2_weight_decay) if l2_weight_decay else None)
+ tf_keras.regularizers.l2(l2_weight_decay) if l2_weight_decay else None)
model, losses = factory.build_yolo(
input_specs, model_base_cfg, l2_regularizer)
+ model.build(input_specs.shape)
+ model.summary(print_fn=logging.info)
# save for later usage within the task.
self._loss_fn = losses
@@ -239,13 +242,14 @@ def build_metrics(self, training=True):
annotation_file=annotation_file,
include_mask=False,
need_rescale_bboxes=False,
- per_category_metrics=self._task_config.per_category_metrics)
+ per_category_metrics=self._task_config.per_category_metrics,
+ max_num_eval_detections=self.task_config.max_num_eval_detections)
return metrics
def build_losses(self, outputs, labels, aux_losses=None):
"""Build YOLO losses."""
- return self._loss_fn(labels, outputs)
+ return self._loss_fn(labels, outputs) # pyrefly: ignore[not-callable]
def train_step(self, inputs, model, optimizer, metrics=None):
"""Train Step.
@@ -275,7 +279,7 @@ def train_step(self, inputs, model, optimizer, metrics=None):
loss_metrics) = self.build_losses(y_pred['raw_output'], label)
# Scale the loss for numerical stability
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
# Compute the gradient
@@ -283,7 +287,7 @@ def train_step(self, inputs, model, optimizer, metrics=None):
gradients = tape.gradient(scaled_loss, train_vars)
# Get unscaled loss if we are using the loss scale optimizer on fp16
- if isinstance(optimizer, tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
gradients = optimizer.get_unscaled_gradients(gradients)
# Apply gradients to the model
@@ -356,7 +360,7 @@ def validation_step(self, inputs, model, metrics=None):
# Compute all metrics
if metrics:
logs.update(
- {self.coco_metric.name: (label['groundtruths'], coco_model_outputs)})
+ {self.coco_metric.name: (label['groundtruths'], coco_model_outputs)}) # pyrefly: ignore[missing-attribute]
for m in metrics:
m.update_state(loss_metrics[m.name])
logs.update({m.name: m.result()})
@@ -365,18 +369,18 @@ def validation_step(self, inputs, model, metrics=None):
def aggregate_logs(self, state=None, step_outputs=None):
"""Get Metric Results."""
if not state:
- self.coco_metric.reset_states()
+ self.coco_metric.reset_states() # pyrefly: ignore[missing-attribute]
state = self.coco_metric
- self.coco_metric.update_state(step_outputs[self.coco_metric.name][0],
- step_outputs[self.coco_metric.name][1])
+ self.coco_metric.update_state(step_outputs[self.coco_metric.name][0], # pyrefly: ignore[missing-attribute, unsupported-operation]
+ step_outputs[self.coco_metric.name][1]) # pyrefly: ignore[unsupported-operation]
return state
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
"""Reduce logs and remove unneeded items. Update with COCO results."""
- res = self.coco_metric.result()
+ res = self.coco_metric.result() # pyrefly: ignore[missing-attribute]
return res
- def initialize(self, model: tf.keras.Model):
+ def initialize(self, model: tf_keras.Model):
"""Loading pretrained checkpoint."""
if not self.task_config.init_checkpoint:
@@ -428,7 +432,7 @@ def create_optimizer(self,
optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
optimizer.set_bias_lr(
opt_factory.get_bias_lr_schedule(self._task_config.smart_bias_lr))
- optimizer.search_and_set_variable_groups(self._model.trainable_variables)
+ optimizer.search_and_set_variable_groups(self._model.trainable_variables) # pyrefly: ignore[missing-attribute]
else:
optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
opt_factory._use_ema = ema
diff --git a/official/projects/yolo/tasks/yolov7.py b/official/projects/yolo/tasks/yolov7.py
new file mode 100644
index 00000000000..64471a9bfb3
--- /dev/null
+++ b/official/projects/yolo/tasks/yolov7.py
@@ -0,0 +1,479 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains classes used to train Yolo."""
+
+from typing import Optional
+
+from absl import logging
+import tensorflow as tf, tf_keras
+
+from official.common import dataset_fn
+from official.core import base_task
+from official.core import config_definitions
+from official.core import input_reader
+from official.core import task_factory
+from official.modeling import performance
+from official.projects.yolo import optimization
+from official.projects.yolo.configs import yolov7 as exp_cfg
+from official.projects.yolo.dataloaders import tf_example_decoder
+from official.projects.yolo.dataloaders import yolo_input
+from official.projects.yolo.losses import yolov7_loss
+from official.projects.yolo.modeling import factory
+from official.projects.yolo.ops import kmeans_anchors
+from official.projects.yolo.ops import mosaic
+from official.projects.yolo.ops import preprocessing_ops
+from official.projects.yolo.tasks import task_utils
+from official.vision.dataloaders import tfds_factory
+from official.vision.dataloaders import tf_example_label_map_decoder
+from official.vision.evaluation import coco_evaluator
+from official.vision.ops import box_ops
+
+OptimizationConfig = optimization.OptimizationConfig
+RuntimeConfig = config_definitions.RuntimeConfig
+
+
+@task_factory.register_task_cls(exp_cfg.YoloV7Task)
+class YoloV7Task(base_task.Task):
+ """A single-replica view of training procedure.
+
+ YOLO task provides artifacts for training/evalution procedures, including
+ loading/iterating over Datasets, initializing the model, calculating the loss,
+ post-processing, and customized metrics with reduction.
+ """
+
+ def __init__(self, params, logging_dir: Optional[str] = None):
+ super().__init__(params, logging_dir)
+
+ min_level = self.task_config.model.min_level
+ max_level = self.task_config.model.max_level
+ anchors_dict = self.task_config.model.anchor_boxes.get(
+ min_level, max_level)[0]
+ anchors, strides = [], []
+ for level in range(min_level, max_level + 1):
+ anchors.append(anchors_dict[str(level)])
+ strides.append(2 ** level)
+
+ loss_config = self.task_config.model.loss
+ if loss_config.use_ota:
+ loss_fn = yolov7_loss.YoloV7LossOTA
+ else:
+ loss_fn = yolov7_loss.YoloV7Loss
+ self._loss_fn = loss_fn(
+ anchors=anchors,
+ strides=strides,
+ input_size=self.task_config.model.input_size[:2],
+ alpha=loss_config.alpha,
+ gamma=loss_config.gamma,
+ box_weight=loss_config.box_weight,
+ obj_weight=loss_config.obj_weight,
+ cls_weight=loss_config.cls_weight,
+ label_smoothing=loss_config.label_smoothing,
+ anchor_threshold=loss_config.anchor_threshold,
+ iou_mix_ratio=loss_config.iou_mix_ratio,
+ num_classes=self.task_config.model.num_classes,
+ auto_balance=loss_config.auto_balance,
+ )
+ self._coco_91_to_80 = False
+ self._metrics = []
+
+ # globally set the random seed
+ preprocessing_ops.set_random_seeds(seed=params.seed)
+
+ if self.task_config.model.anchor_boxes.generate_anchors:
+ self.generate_anchors()
+ return
+
+ def generate_anchors(self):
+ """Generate Anchor boxes for an arbitrary object detection dataset."""
+ input_size = self.task_config.model.input_size
+ anchor_cfg = self.task_config.model.anchor_boxes
+ backbone = self.task_config.model.backbone.get()
+
+ dataset = self.task_config.train_data
+ decoder = self._get_data_decoder(dataset)
+
+ num_anchors = backbone.max_level - backbone.min_level + 1
+ num_anchors *= anchor_cfg.anchors_per_scale
+
+ gbs = dataset.global_batch_size
+ dataset.global_batch_size = 1
+ box_reader = kmeans_anchors.BoxGenInputReader(
+ dataset,
+ dataset_fn=dataset_fn.pick_dataset_fn(
+ self.task_config.train_data.file_type),
+ decoder_fn=decoder.decode)
+
+ boxes = box_reader.read(
+ k=num_anchors,
+ anchors_per_scale=anchor_cfg.anchors_per_scale,
+ image_resolution=input_size,
+ scaling_mode=anchor_cfg.scaling_mode,
+ box_generation_mode=anchor_cfg.box_generation_mode,
+ num_samples=anchor_cfg.num_samples)
+
+ dataset.global_batch_size = gbs
+
+ with open('anchors.txt', 'w') as f:
+ f.write(f'input resolution: {input_size} \n boxes: \n {boxes}')
+ logging.info('INFO: boxes will be saved to anchors.txt, mack sure to save'
+ 'them and update the boxes feild in you yaml config file.')
+
+ anchor_cfg.set_boxes(boxes)
+ return boxes
+
+ def build_model(self):
+ """Build an instance of Yolo."""
+
+ model_base_cfg = self.task_config.model
+ l2_weight_decay = self.task_config.weight_decay / 2.0
+
+ input_size = model_base_cfg.input_size.copy()
+ input_specs = tf_keras.layers.InputSpec(shape=[None] + input_size)
+ l2_regularizer = (
+ tf_keras.regularizers.l2(l2_weight_decay) if l2_weight_decay else None)
+ model = factory.build_yolov7(input_specs, model_base_cfg, l2_regularizer)
+ model.build(input_specs.shape)
+ model.summary(print_fn=logging.info)
+
+ # save for later usage within the task.
+ self._model = model
+ return model
+
+ def _get_data_decoder(self, params):
+ """Get a decoder object to decode the dataset."""
+ if params.tfds_name:
+ decoder = tfds_factory.get_detection_decoder(params.tfds_name)
+ else:
+ decoder_cfg = params.decoder.get()
+ if params.decoder.type == 'simple_decoder':
+ self._coco_91_to_80 = decoder_cfg.coco91_to_80
+ decoder = tf_example_decoder.TfExampleDecoder(
+ coco91_to_80=decoder_cfg.coco91_to_80,
+ regenerate_source_id=decoder_cfg.regenerate_source_id)
+ elif params.decoder.type == 'label_map_decoder':
+ decoder = tf_example_label_map_decoder.TfExampleDecoderLabelMap(
+ label_map=decoder_cfg.label_map,
+ regenerate_source_id=decoder_cfg.regenerate_source_id)
+ else:
+ raise ValueError('Unknown decoder type: {}!'.format(
+ params.decoder.type))
+ return decoder
+
+ def build_inputs(self, params, input_context=None):
+ """Build input dataset."""
+ model = self.task_config.model
+
+ # get anchor boxes dict based on models min and max level
+ backbone = model.backbone.get()
+ anchor_dict, level_limits = model.anchor_boxes.get(backbone.min_level,
+ backbone.max_level)
+
+ params.seed = self.task_config.seed
+ # set shared patamters between mosaic and yolo_input
+ base_config = dict(
+ letter_box=params.parser.letter_box,
+ aug_rand_translate=params.parser.aug_rand_translate,
+ aug_rand_angle=params.parser.aug_rand_angle,
+ aug_rand_perspective=params.parser.aug_rand_perspective,
+ area_thresh=params.parser.area_thresh,
+ random_flip=params.parser.random_flip,
+ seed=params.seed,
+ )
+
+ # get the decoder
+ decoder = self._get_data_decoder(params)
+
+ # init Mosaic
+ sample_fn = mosaic.Mosaic(
+ output_size=model.input_size,
+ mosaic_frequency=params.parser.mosaic.mosaic_frequency,
+ mosaic9_frequency=params.parser.mosaic.mosaic9_frequency,
+ mixup_frequency=params.parser.mosaic.mixup_frequency,
+ jitter=params.parser.mosaic.jitter,
+ mosaic_center=params.parser.mosaic.mosaic_center,
+ mosaic9_center=params.parser.mosaic.mosaic9_center,
+ mosaic_crop_mode=params.parser.mosaic.mosaic_crop_mode,
+ aug_scale_min=params.parser.mosaic.aug_scale_min,
+ aug_scale_max=params.parser.mosaic.aug_scale_max,
+ **base_config)
+
+ # init Parser
+ parser = yolo_input.Parser(
+ output_size=model.input_size,
+ anchors=anchor_dict,
+ use_tie_breaker=params.parser.use_tie_breaker,
+ jitter=params.parser.jitter,
+ aug_scale_min=params.parser.aug_scale_min,
+ aug_scale_max=params.parser.aug_scale_max,
+ aug_rand_hue=params.parser.aug_rand_hue,
+ aug_rand_saturation=params.parser.aug_rand_saturation,
+ aug_rand_brightness=params.parser.aug_rand_brightness,
+ max_num_instances=params.parser.max_num_instances,
+ scale_xy=model.detection_generator.scale_xy.get(),
+ expanded_strides=model.detection_generator.path_scales.get(),
+ darknet=False,
+ best_match_only=params.parser.best_match_only,
+ anchor_t=params.parser.anchor_thresh,
+ random_pad=params.parser.random_pad,
+ level_limits=level_limits,
+ dtype=params.dtype,
+ **base_config,
+ )
+
+ # init the dataset reader
+ reader = input_reader.InputReader(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ sample_fn=sample_fn.mosaic_fn(is_training=params.is_training),
+ parser_fn=parser.parse_fn(params.is_training))
+ dataset = reader.read(input_context=input_context)
+ return dataset
+
+ def build_metrics(self, training=True):
+ """Build detection metrics."""
+ metrics = []
+
+ metrics = [
+ task_utils.ListMetrics(
+ ['box_loss', 'obj_loss', 'cls_loss', 'iou'], 'separate_losses'
+ ),
+ task_utils.ListMetrics(
+ ['num_matchings', 'num_gts', 'num_duplicates'], 'stats'
+ ),
+ ]
+
+ self._metrics = metrics
+ if not training:
+ annotation_file = self.task_config.annotation_file
+ if self._coco_91_to_80:
+ annotation_file = None
+ self.coco_metric = coco_evaluator.COCOEvaluator(
+ annotation_file=annotation_file,
+ include_mask=False,
+ need_rescale_bboxes=False,
+ per_category_metrics=self._task_config.per_category_metrics,
+ max_num_eval_detections=self.task_config.max_num_eval_detections)
+
+ return metrics
+
+ def build_losses(self, outputs, labels, aux_losses=None):
+ """Build YOLOv7 losses."""
+ return self._loss_fn(labels, outputs)
+
+ def train_step(self, inputs, model, optimizer, metrics=None):
+ """Train Step.
+
+ Forward step and backwards propagate the model.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the model, forward pass definition.
+ optimizer: the optimizer for this training step.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ image, label = inputs
+
+ with tf.GradientTape(persistent=False) as tape:
+ # Compute a prediction
+ y_pred = model(image, training=True)
+
+ # Cast to float32 for gradietn computation
+ y_pred = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), y_pred)
+
+ # Get the total loss
+ loss = self.build_losses(y_pred['raw_output'], label)
+ scaled_loss = loss
+
+ # Scale the loss for numerical stability
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ scaled_loss = optimizer.get_scaled_loss(scaled_loss)
+
+ # Compute the gradient
+ train_vars = model.trainable_variables
+ gradients = tape.gradient(scaled_loss, train_vars)
+
+ # Get unscaled loss if we are using the loss scale optimizer on fp16
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
+ gradients = optimizer.get_unscaled_gradients(gradients)
+
+ # Apply gradients to the model
+ optimizer.apply_gradients(zip(gradients, train_vars))
+ logs = {self.loss: loss}
+
+ # Compute all metrics
+ if metrics:
+ metrics[0].update_state(self._loss_fn.report_separate_losses())
+ logs.update({metrics[0].name: metrics[0].result()})
+
+ metrics[1].update_state(self._loss_fn.report_stats())
+ logs.update({metrics[1].name: metrics[1].result()})
+ return logs
+
+ def _reorg_boxes(self, boxes, info, num_detections):
+ """Scale and Clean boxes prior to Evaluation."""
+ mask = tf.sequence_mask(num_detections, maxlen=tf.shape(boxes)[1])
+ mask = tf.cast(tf.expand_dims(mask, axis=-1), boxes.dtype)
+
+ # Denormalize the boxes by the shape of the image
+ inshape = tf.expand_dims(info[:, 1, :], axis=1)
+ ogshape = tf.expand_dims(info[:, 0, :], axis=1)
+ scale = tf.expand_dims(info[:, 2, :], axis=1)
+ offset = tf.expand_dims(info[:, 3, :], axis=1)
+
+ boxes = box_ops.denormalize_boxes(boxes, inshape)
+ boxes = box_ops.clip_boxes(boxes, inshape)
+ boxes += tf.tile(offset, [1, 1, 2])
+ boxes /= tf.tile(scale, [1, 1, 2])
+ boxes = box_ops.clip_boxes(boxes, ogshape)
+
+ # Mask the boxes for usage
+ boxes *= mask
+ boxes += (mask - 1)
+ return boxes
+
+ def validation_step(self, inputs, model, metrics=None):
+ """Validatation step.
+
+ Args:
+ inputs: a dictionary of input tensors.
+ model: the keras.Model.
+ metrics: a nested structure of metrics objects.
+
+ Returns:
+ A dictionary of logs.
+ """
+ image, label = inputs
+
+ # Step the model once
+ y_pred = model(image, training=False)
+ y_pred = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), y_pred)
+ loss_val = self.build_losses(y_pred['raw_output'], label)
+ logs = {self.loss: loss_val}
+
+ # Reorganize and rescale the boxes
+ info = label['groundtruths']['image_info']
+ boxes = self._reorg_boxes(y_pred['bbox'], info, y_pred['num_detections'])
+
+ # Build the input for the coc evaluation metric
+ coco_model_outputs = {
+ 'detection_boxes': boxes,
+ 'detection_scores': y_pred['confidence'],
+ 'detection_classes': y_pred['classes'],
+ 'num_detections': y_pred['num_detections'],
+ 'source_id': label['groundtruths']['source_id'],
+ 'image_info': label['groundtruths']['image_info']
+ }
+
+ # Compute all metrics
+ if metrics:
+ logs.update(
+ {self.coco_metric.name: (label['groundtruths'], coco_model_outputs)})
+ if metrics:
+ metrics[0].update_state(self._loss_fn.report_separate_losses())
+ logs.update({metrics[0].name: metrics[0].result()})
+
+ metrics[1].update_state(self._loss_fn.report_stats())
+ logs.update({metrics[1].name: metrics[1].result()})
+ return logs
+
+ def aggregate_logs(self, state=None, step_outputs=None):
+ """Get Metric Results."""
+ if not state:
+ self.coco_metric.reset_states()
+ state = self.coco_metric
+ self.coco_metric.update_state(step_outputs[self.coco_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ step_outputs[self.coco_metric.name][1]) # pyrefly: ignore[unsupported-operation]
+ return state
+
+ def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
+ """Reduce logs and remove unneeded items. Update with COCO results."""
+ res = self.coco_metric.result()
+ return res
+
+ def initialize(self, model: tf_keras.Model):
+ """Loading pretrained checkpoint."""
+
+ if not self.task_config.init_checkpoint:
+ logging.info('Training from Scratch.')
+ return
+
+ ckpt_dir_or_file = self.task_config.init_checkpoint
+ if tf.io.gfile.isdir(ckpt_dir_or_file):
+ ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
+
+ # Restoring checkpoint.
+ if self.task_config.init_checkpoint_modules == 'all':
+ ckpt = tf.train.Checkpoint(**model.checkpoint_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+ else:
+ ckpt_items = {}
+ if 'backbone' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(backbone=model.backbone)
+ if 'decoder' in self.task_config.init_checkpoint_modules:
+ ckpt_items.update(decoder=model.decoder)
+
+ ckpt = tf.train.Checkpoint(**ckpt_items)
+ status = ckpt.read(ckpt_dir_or_file)
+ status.expect_partial().assert_existing_objects_matched()
+
+ logging.info('Finished loading pretrained checkpoint from %s',
+ ckpt_dir_or_file)
+
+ def create_optimizer(self,
+ optimizer_config: OptimizationConfig,
+ runtime_config: Optional[RuntimeConfig] = None):
+ """Creates an TF optimizer from configurations.
+
+ Args:
+ optimizer_config: the parameters of the Optimization settings.
+ runtime_config: the parameters of the runtime.
+
+ Returns:
+ A tf.optimizers.Optimizer object.
+ """
+ opt_factory = optimization.YoloOptimizerFactory(optimizer_config)
+ # pylint: disable=protected-access
+ ema = opt_factory._use_ema
+ opt_factory._use_ema = False
+
+ opt_type = opt_factory._optimizer_type
+ if opt_type == 'sgd_torch':
+ optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
+ optimizer.set_bias_lr(
+ opt_factory.get_bias_lr_schedule(self._task_config.smart_bias_lr))
+ optimizer.search_and_set_variable_groups(self._model.trainable_variables)
+ else:
+ optimizer = opt_factory.build_optimizer(opt_factory.build_learning_rate())
+ opt_factory._use_ema = ema
+
+ if ema:
+ logging.info('EMA is enabled.')
+ optimizer = opt_factory.add_ema(optimizer)
+
+ # pylint: enable=protected-access
+
+ if runtime_config and runtime_config.loss_scale:
+ use_float16 = runtime_config.mixed_precision_dtype == 'float16'
+ optimizer = performance.configure_optimizer(
+ optimizer,
+ use_float16=use_float16,
+ loss_scale=runtime_config.loss_scale)
+
+ return optimizer
diff --git a/official/vision/beta/projects/yolo/train.py b/official/projects/yolo/train.py
similarity index 79%
rename from official/vision/beta/projects/yolo/train.py
rename to official/projects/yolo/train.py
index 85c9f215d52..e9eb992fd99 100644
--- a/official/vision/beta/projects/yolo/train.py
+++ b/official/projects/yolo/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,8 +18,8 @@
from absl import flags
from official.common import flags as tfm_flags
-from official.vision.beta import train
-from official.vision.beta.projects.yolo.common import registry_imports # pylint: disable=unused-import
+from official.projects.yolo.common import registry_imports # pylint: disable=unused-import
+from official.vision import train
FLAGS = flags.FLAGS
diff --git a/official/projects/yolo/yolo_object_detection.ipynb b/official/projects/yolo/yolo_object_detection.ipynb
new file mode 100644
index 00000000000..28c1d0208c8
--- /dev/null
+++ b/official/projects/yolo/yolo_object_detection.ipynb
@@ -0,0 +1,821 @@
+{
+ "nbformat": 4,
+ "nbformat_minor": 0,
+ "metadata": {
+ "colab": {
+ "provenance": [],
+ "gpuType": "T4"
+ },
+ "kernelspec": {
+ "name": "python3",
+ "display_name": "Python 3"
+ },
+ "language_info": {
+ "name": "python"
+ },
+ "accelerator": "GPU"
+ },
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "source": [
+ "# Object Detection using YOLOv7"
+ ],
+ "metadata": {
+ "id": "HbJG-0vkfOWt"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "!pip install -q tf-models-official"
+ ],
+ "metadata": {
+ "id": "uo0ZQ1wBBSHf"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Import Necessary Libraries"
+ ],
+ "metadata": {
+ "id": "SU3iw-I5fTll"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "import os\n",
+ "import logging\n",
+ "import yaml\n",
+ "import tempfile\n",
+ "import pprint\n",
+ "import numpy as np\n",
+ "import matplotlib.pyplot as plt\n",
+ "from PIL import Image\n",
+ "from six import BytesIO\n",
+ "from IPython import display\n",
+ "from urllib.request import urlopen\n",
+ "\n",
+ "logging.disable(logging.WARNING)"
+ ],
+ "metadata": {
+ "id": "z_yMmIayFgKJ"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "import orbit\n",
+ "import tensorflow as tf\n",
+ "import tensorflow_models as tfm\n",
+ "\n",
+ "from official.core import exp_factory\n",
+ "from official.core import config_definitions as cfg\n",
+ "from official.projects.yolo.common import registry_imports\n",
+ "from official.projects.yolo.configs import yolov7\n",
+ "from official.vision.serving import export_saved_model_lib\n",
+ "from official.vision.ops.preprocess_ops import normalize_image\n",
+ "from official.vision.ops.preprocess_ops import resize_and_crop_image\n",
+ "from official.vision.utils.object_detection import visualization_utils\n",
+ "from official.vision.dataloaders.tf_example_decoder import TfExampleDecoder\n",
+ "\n",
+ "\n",
+ "pp = pprint.PrettyPrinter(indent=4) # Set Pretty Print Indentation\n",
+ "print(tf.__version__) # Check the version of tensorflow used\n",
+ "\n",
+ "%matplotlib inline"
+ ],
+ "metadata": {
+ "id": "G34F8TBOF1Se"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Download BCCD(Blood Cells) dataset."
+ ],
+ "metadata": {
+ "id": "MfhQ5C-bfV_v"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "!curl -L 'https://public.roboflow.com/ds/ZpYLqHeT0W?key=ZXfZLRnhsc' > './BCCD.v1-bccd.coco.zip'\n",
+ "!unzip -q -o './BCCD.v1-bccd.coco.zip' -d './BCC.v1-bccd.coco/'\n",
+ "!rm './BCCD.v1-bccd.coco.zip'"
+ ],
+ "metadata": {
+ "id": "DzXHG3CnGPJv"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Convert COCO format dataset to tfrecords"
+ ],
+ "metadata": {
+ "id": "Lgw_SGaVfcX3"
+ }
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Training TFRecords"
+ ],
+ "metadata": {
+ "id": "9IzGpt2JfdEj"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "TRAIN_DATA_DIR='./BCC.v1-bccd.coco/train'\n",
+ "TRAIN_ANNOTATION_FILE_DIR='./BCC.v1-bccd.coco/train/_annotations.coco.json'\n",
+ "OUTPUT_TFRECORD_TRAIN='./bccd_coco_tfrecords/train'\n",
+ "\n",
+ "# Need to provide\n",
+ " # 1. image_dir: where images are present\n",
+ " # 2. object_annotations_file: where annotations are listed in json format\n",
+ " # 3. output_file_prefix: where to write output convered TFRecords files\n",
+ "!python -m official.vision.data.create_coco_tf_record --logtostderr \\\n",
+ " --image_dir={TRAIN_DATA_DIR} \\\n",
+ " --object_annotations_file={TRAIN_ANNOTATION_FILE_DIR} \\\n",
+ " --output_file_prefix={OUTPUT_TFRECORD_TRAIN} \\\n",
+ " --num_shards=1"
+ ],
+ "metadata": {
+ "id": "ka8AuE9KGu9c"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Validation TFRecords"
+ ],
+ "metadata": {
+ "id": "x8WGWESzff5c"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "VALID_DATA_DIR='./BCC.v1-bccd.coco/valid'\n",
+ "VALID_ANNOTATION_FILE_DIR='./BCC.v1-bccd.coco/valid/_annotations.coco.json'\n",
+ "OUTPUT_TFRECORD_VALID='./bccd_coco_tfrecords/valid'\n",
+ "\n",
+ "!python -m official.vision.data.create_coco_tf_record --logtostderr \\\n",
+ " --image_dir={VALID_DATA_DIR} \\\n",
+ " --object_annotations_file={VALID_ANNOTATION_FILE_DIR} \\\n",
+ " --output_file_prefix={OUTPUT_TFRECORD_VALID} \\\n",
+ " --num_shards=1"
+ ],
+ "metadata": {
+ "id": "gqLhou6_HKo0"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Test TFRecords"
+ ],
+ "metadata": {
+ "id": "f2cnxT0XfjpY"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "TEST_DATA_DIR='./BCC.v1-bccd.coco/test'\n",
+ "TEST_ANNOTATION_FILE_DIR='./BCC.v1-bccd.coco/test/_annotations.coco.json'\n",
+ "OUTPUT_TFRECORD_TEST='./bccd_coco_tfrecords/test'\n",
+ "\n",
+ "!python -m official.vision.data.create_coco_tf_record --logtostderr \\\n",
+ " --image_dir=$TEST_DATA_DIR \\\n",
+ " --object_annotations_file=$TEST_ANNOTATION_FILE_DIR \\\n",
+ " --output_file_prefix=$OUTPUT_TFRECORD_TEST \\\n",
+ " --num_shards=1"
+ ],
+ "metadata": {
+ "id": "0UHWrMtnHY-Q"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Configure the YOLOv7"
+ ],
+ "metadata": {
+ "id": "fxNVx9kkfmkZ"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "train_data_input_path = './bccd_coco_tfrecords/train-00000-of-00001.tfrecord'\n",
+ "valid_data_input_path = './bccd_coco_tfrecords/valid-00000-of-00001.tfrecord'\n",
+ "test_data_input_path = './bccd_coco_tfrecords/test-00000-of-00001.tfrecord'\n",
+ "model_dir = './trained_model/'\n",
+ "export_dir ='./exported_model/'"
+ ],
+ "metadata": {
+ "id": "9bRtYTviIO9s"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Download YOLOv7 config file"
+ ],
+ "metadata": {
+ "id": "3C99W7mJfr-Q"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "!wget https://raw.githubusercontent.com/tensorflow/models/master/official/projects/yolo/configs/experiments/yolov7/detection/yolov7_gpu.yaml"
+ ],
+ "metadata": {
+ "id": "RAoNvLSLPlwq"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "In Model Garden, the collections of parameters that define a model are called configs. Model Garden can create a config based on a known set of parameters via a factory.\n",
+ "\n",
+ "Use the yolo_darknet experiment configuration, all the configurations are present [here]().\n",
+ "\n",
+ "The configuration defines an experiment to train a RetinanNet with Resnet-50 as backbone, FPN as decoder. Default Configuration is trained on COCO train2017 and evaluated on COCO val2017.\n",
+ "\n",
+ "There are also other alternative experiments available such as retinanet_resnetfpn_coco, retinanet_spinenet_coco, fasterrcnn_resnetfpn_coco and more. One can switch to them by changing the experiment name argument to the get_exp_config function.\n",
+ "\n",
+ "We are going to fine tune the Resnet-50 backbone checkpoint which is already present in the default configuration.\n",
+ "\n",
+ "\n"
+ ],
+ "metadata": {
+ "id": "M-DgdoJWfv46"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "exp_config = exp_factory.get_exp_config('yolo_darknet')"
+ ],
+ "metadata": {
+ "id": "nHz6-t-aIXz4"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "with open('./yolov7_gpu.yaml') as f:\n",
+ " params = yaml.full_load(f)\n",
+ "exp_config.override(params, is_strict=False)"
+ ],
+ "metadata": {
+ "id": "xAzD4cZSPpOI"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Adjust task configuration which includes model, train_data and validation_data."
+ ],
+ "metadata": {
+ "id": "ef_ZA-pbgFAa"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "batch_size = 8\n",
+ "num_classes = 3\n",
+ "\n",
+ "HEIGHT, WIDTH = 416, 416\n",
+ "IMG_SIZE = [HEIGHT, WIDTH, 3]\n",
+ "\n",
+ "\n",
+ "# Backbone config.\n",
+ "exp_config.task.init_checkpoint = ''\n",
+ "exp_config.task.freeze_backbone = False\n",
+ "exp_config.task.annotation_file = ''\n",
+ "\n",
+ "# Model config.\n",
+ "exp_config.task.model.input_size = IMG_SIZE\n",
+ "exp_config.task.model.num_classes = num_classes\n",
+ "\n",
+ "# Training data config.\n",
+ "exp_config.task.train_data.input_path = train_data_input_path\n",
+ "exp_config.task.train_data.global_batch_size = batch_size\n",
+ "exp_config.task.train_data.parser.aug_scale_max = 1.0\n",
+ "exp_config.task.train_data.parser.aug_scale_min = 1.0\n",
+ "\n",
+ "# Validation data config.\n",
+ "exp_config.task.validation_data.input_path = valid_data_input_path\n",
+ "exp_config.task.validation_data.global_batch_size = batch_size"
+ ],
+ "metadata": {
+ "id": "KgDas9WmI_u2"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Adjust trainer configuration"
+ ],
+ "metadata": {
+ "id": "OwIapj-ngGb4"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "train_steps = 2000\n",
+ "\n",
+ "exp_config.trainer.steps_per_loop = 500 # steps_per_loop = num_of_training_examples // train_batch_size\n",
+ "exp_config.trainer.summary_interval = 500\n",
+ "exp_config.trainer.checkpoint_interval = 500\n",
+ "exp_config.trainer.validation_interval = 500\n",
+ "exp_config.trainer.validation_steps = 500 # validation_steps = num_of_validation_examples // eval_batch_size\n",
+ "exp_config.trainer.train_steps = train_steps\n",
+ "exp_config.trainer.optimizer_config.warmup.linear.warmup_steps = 500\n",
+ "exp_config.trainer.optimizer_config.learning_rate.type = 'cosine'\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.decay_steps = train_steps\n",
+ "exp_config.trainer.optimizer_config.learning_rate.cosine.initial_learning_rate = 3.2e-5"
+ ],
+ "metadata": {
+ "id": "sOOtAun_JJCh"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Checkout the changed the configuration parameters and default parameters for further customization of model."
+ ],
+ "metadata": {
+ "id": "0ZK_AfAngJpk"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "pp.pprint(exp_config.as_dict())"
+ ],
+ "metadata": {
+ "id": "VfrJhSrfIvqu"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Setup the distribution strategy"
+ ],
+ "metadata": {
+ "id": "kGEcZTbHgNUm"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "logical_device_names = [logical_device.name for logical_device in tf.config.list_logical_devices()]\n",
+ "\n",
+ "if exp_config.runtime.mixed_precision_dtype == tf.float16:\n",
+ " tf.keras.mixed_precision.set_global_policy('mixed_float16')\n",
+ "\n",
+ "if 'GPU' in ''.join(logical_device_names):\n",
+ " distribution_strategy = tf.distribute.MirroredStrategy()\n",
+ "elif 'TPU' in ''.join(logical_device_names):\n",
+ " tf.tpu.experimental.initialize_tpu_system()\n",
+ " tpu = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='/device:TPU_SYSTEM:0')\n",
+ " distribution_strategy = tf.distribute.experimental.TPUStrategy(tpu)\n",
+ "else:\n",
+ " print('Warning: this will be really slow.')\n",
+ " distribution_strategy = tf.distribute.OneDeviceStrategy(logical_device_names[0])"
+ ],
+ "metadata": {
+ "id": "u5zfMsq_JsT7"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Create the `Task` object `(tfm.core.base_task.Task)` from the `config_definitions.TaskConfig`.\n",
+ "\n",
+ "The Task object has all the methods necessary for building the dataset, building the model, and running training & evaluation. These methods are driven by `tfm.core.train_lib.run_experimen`t."
+ ],
+ "metadata": {
+ "id": "RR4uiiAzg3kh"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "with distribution_strategy.scope():\n",
+ " task = tfm.core.task_factory.get_task(exp_config.task, logging_dir=model_dir)"
+ ],
+ "metadata": {
+ "id": "6H_pcsn3JoXt"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Build inputs from given training tfrecords"
+ ],
+ "metadata": {
+ "id": "MnFUsZsJg6QY"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "for images, labels in task.build_inputs(exp_config.task.train_data).take(1):\n",
+ " print()\n",
+ " print(f'images.shape: {str(images.shape):16} images.dtype: {images.dtype!r}')\n",
+ " print(f'labels.keys: {labels.keys()}')"
+ ],
+ "metadata": {
+ "id": "xXE_sjHQJvwW"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Create category index for each label"
+ ],
+ "metadata": {
+ "id": "Ce_0Prspg7Av"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "tf_ex_decoder = TfExampleDecoder() # define tf example decoder\n",
+ "\n",
+ "category_index={\n",
+ " 1: {\n",
+ " 'id': 1,\n",
+ " 'name': 'Platelets'\n",
+ " },\n",
+ " 2: {\n",
+ " 'id': 2,\n",
+ " 'name': 'RBC'\n",
+ " },\n",
+ " 3: {\n",
+ " 'id': 3,\n",
+ " 'name': 'WBC'\n",
+ " }\n",
+ "}"
+ ],
+ "metadata": {
+ "id": "O4pypN0jg9aM"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Helper function for visualizing the results from TFRecords."
+ ],
+ "metadata": {
+ "id": "-IPcyGWgg_Tg"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "def show_batch(raw_records, num_of_examples):\n",
+ " plt.figure(figsize=(20, 20))\n",
+ " use_normalized_coordinates=True\n",
+ " min_score_thresh = 0.30\n",
+ " for i, serialized_example in enumerate(raw_records):\n",
+ " plt.subplot(1, num_of_examples, i + 1)\n",
+ " decoded_tensors = tf_ex_decoder.decode(serialized_example)\n",
+ " image = decoded_tensors['image'].numpy().astype('uint8')\n",
+ " scores = np.ones(shape=(len(decoded_tensors['groundtruth_boxes'])))\n",
+ " visualization_utils.visualize_boxes_and_labels_on_image_array(\n",
+ " image,\n",
+ " decoded_tensors['groundtruth_boxes'].numpy(),\n",
+ " decoded_tensors['groundtruth_classes'].numpy().astype('int'),\n",
+ " scores,\n",
+ " category_index=category_index,\n",
+ " use_normalized_coordinates=use_normalized_coordinates,\n",
+ " max_boxes_to_draw=200,\n",
+ " min_score_thresh=min_score_thresh,\n",
+ " agnostic_mode=False,\n",
+ " instance_masks=None,\n",
+ " line_thickness=4)\n",
+ "\n",
+ " plt.imshow(image)\n",
+ " plt.axis('off')\n",
+ " plt.title(f'Image-{i+1}')\n",
+ " plt.show()"
+ ],
+ "metadata": {
+ "id": "sn3dOH61hDI4"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Visualize the training data samples"
+ ],
+ "metadata": {
+ "id": "an1PwnGShEdt"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "buffer_size = 20\n",
+ "num_of_examples = 3\n",
+ "\n",
+ "raw_records = tf.data.TFRecordDataset(\n",
+ " exp_config.task.train_data.input_path).shuffle(\n",
+ " buffer_size=buffer_size).take(num_of_examples)\n",
+ "show_batch(raw_records, num_of_examples)"
+ ],
+ "metadata": {
+ "id": "2Moh4jfqhGQU"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Train and Evaluate the model using `tfm.core.train_lib.run_experiment`."
+ ],
+ "metadata": {
+ "id": "T_DzGO6zhIB6"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "model, eval_logs = tfm.core.train_lib.run_experiment(\n",
+ " distribution_strategy=distribution_strategy,\n",
+ " task=task,\n",
+ " mode='train_and_eval',\n",
+ " params=exp_config,\n",
+ " model_dir=model_dir,\n",
+ " run_post_eval=True)"
+ ],
+ "metadata": {
+ "id": "Gk2kn8qyJ5Yk"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Save the trained experiment configuration"
+ ],
+ "metadata": {
+ "id": "5Um3JlvXhKTW"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "# save config file\n",
+ "tfm.core.train_utils.serialize_config(exp_config, model_dir)"
+ ],
+ "metadata": {
+ "id": "pt1w5i2qhMd2"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Export the trained model"
+ ],
+ "metadata": {
+ "id": "-iOdrkPwhOs1"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "EXPORT_DIR_PATH = \"./exported_model/\"\n",
+ "CHECKPOINT_PATH = \"/content/trained_model/ckpt-1000\"\n",
+ "CONFIG_FILE_PATH = \"/content/trained_model/params.yaml\""
+ ],
+ "metadata": {
+ "id": "QNK411CBhRF-"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "!python -m official.projects.yolo.serving.export_saved_model --export_dir={EXPORT_DIR_PATH}/ \\\n",
+ " --checkpoint_path={CHECKPOINT_PATH} \\\n",
+ " --config_file={CONFIG_FILE_PATH} \\\n",
+ " --batch_size=1 \\\n",
+ " --input_image_size={HEIGHT},{WIDTH}"
+ ],
+ "metadata": {
+ "id": "l1PsfVKUhSnK"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "## Load the exported model for inference"
+ ],
+ "metadata": {
+ "id": "aRAJM37zhWQU"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "imported = tf.saved_model.load(\"/content/exported_model/saved_model\")\n",
+ "model_fn = imported.signatures['serving_default']"
+ ],
+ "metadata": {
+ "id": "tUnJI6b1hW63"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Helper functions for inference"
+ ],
+ "metadata": {
+ "id": "Y2EpOLg6hYb0"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "def load_image_into_numpy_array(path):\n",
+ " \"\"\"Load an image from file into a numpy array.\n",
+ "\n",
+ " Puts image into numpy array to feed into tensorflow graph.\n",
+ " Note that by convention we put it into a numpy array with shape\n",
+ " (height, width, channels), where channels=3 for RGB.\n",
+ "\n",
+ " Args:\n",
+ " path: the file path to the image\n",
+ "\n",
+ " Returns:\n",
+ " uint8 numpy array with shape (img_height, img_width, 3)\n",
+ " \"\"\"\n",
+ " image = None\n",
+ " if(path.startswith('http')):\n",
+ " response = urlopen(path)\n",
+ " image_data = response.read()\n",
+ " image_data = BytesIO(image_data)\n",
+ " image = Image.open(image_data)\n",
+ " else:\n",
+ " image_data = tf.io.gfile.GFile(path, 'rb').read()\n",
+ " image = Image.open(BytesIO(image_data))\n",
+ "\n",
+ " (im_width, im_height) = image.size\n",
+ " return np.array(image.getdata()).reshape(\n",
+ " (1, im_height, im_width, 3)).astype(np.uint8)\n",
+ "\n",
+ "\n",
+ "\n",
+ "def build_inputs_for_object_detection(image, input_image_size):\n",
+ " \"\"\"Builds Object Detection model inputs for serving.\"\"\"\n",
+ " image, _ = resize_and_crop_image(\n",
+ " image,\n",
+ " input_image_size,\n",
+ " padded_size=input_image_size,\n",
+ " aug_scale_min=1.0,\n",
+ " aug_scale_max=1.0)\n",
+ " return image"
+ ],
+ "metadata": {
+ "id": "Ql177vJ3hUXj"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Visualize original test data"
+ ],
+ "metadata": {
+ "id": "pkyLaTNkhbz0"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "num_of_examples = 3\n",
+ "\n",
+ "test_ds = tf.data.TFRecordDataset(\n",
+ " '/content/bccd_coco_tfrecords/test-00000-of-00001.tfrecord').take(\n",
+ " num_of_examples)\n",
+ "show_batch(test_ds, num_of_examples)"
+ ],
+ "metadata": {
+ "id": "nWqc5lKiheQa"
+ },
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "markdown",
+ "source": [
+ "### Inference on test data"
+ ],
+ "metadata": {
+ "id": "iWynYjEOhhv_"
+ }
+ },
+ {
+ "cell_type": "code",
+ "source": [
+ "input_image_size = (HEIGHT, WIDTH)\n",
+ "plt.figure(figsize=(20, 20))\n",
+ "min_score_thresh = 0.3 # Change minimum score for threshold to see all bounding boxes confidences.\n",
+ "\n",
+ "for i, serialized_example in enumerate(test_ds):\n",
+ " plt.subplot(1, 3, i+1)\n",
+ " decoded_tensors = tf_ex_decoder.decode(serialized_example)\n",
+ " image = build_inputs_for_object_detection(decoded_tensors['image'], input_image_size)\n",
+ " image = tf.expand_dims(image, axis=0)\n",
+ " image = tf.cast(image, dtype = tf.uint8)\n",
+ " image_np = image[0].numpy()\n",
+ " result = model_fn(image)\n",
+ " visualization_utils.visualize_boxes_and_labels_on_image_array(\n",
+ " image_np,\n",
+ " result['detection_boxes'][0].numpy(),\n",
+ " result['detection_classes'][0].numpy().astype(int),\n",
+ " result['detection_scores'][0].numpy(),\n",
+ " category_index=category_index,\n",
+ " use_normalized_coordinates=True,\n",
+ " max_boxes_to_draw=200,\n",
+ " min_score_thresh=min_score_thresh,\n",
+ " agnostic_mode=False,\n",
+ " instance_masks=None,\n",
+ " line_thickness=4)\n",
+ " plt.imshow(image_np)\n",
+ " plt.axis('off')\n",
+ "\n",
+ "plt.show()"
+ ],
+ "metadata": {
+ "id": "9Gvvk1G1hgCD"
+ },
+ "execution_count": null,
+ "outputs": []
+ }
+ ]
+}
\ No newline at end of file
diff --git a/official/projects/yt8m/__init__.py b/official/projects/yt8m/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/yt8m/__init__.py
+++ b/official/projects/yt8m/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/yt8m/configs/__init__.py b/official/projects/yt8m/configs/__init__.py
index d34bc0957cd..d41cd4ceef2 100644
--- a/official/projects/yt8m/configs/__init__.py
+++ b/official/projects/yt8m/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/yt8m/configs/yt8m.py b/official/projects/yt8m/configs/yt8m.py
index cfb4e033b0e..7734f673ea6 100644
--- a/official/projects/yt8m/configs/yt8m.py
+++ b/official/projects/yt8m/configs/yt8m.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,6 @@
"""Video classification configuration definition."""
import dataclasses
from typing import Optional, Tuple
-from absl import flags
from official.core import config_definitions as cfg
from official.core import exp_factory
@@ -23,7 +22,6 @@
from official.modeling import optimization
from official.vision.configs import common
-FLAGS = flags.FLAGS
YT8M_TRAIN_EXAMPLES = 3888919
YT8M_VAL_EXAMPLES = 1112356
@@ -35,44 +33,77 @@
@dataclasses.dataclass
class DataConfig(cfg.DataConfig):
- """The base configuration for building datasets."""
+ """The base configuration for building datasets.
+
+ Attributes:
+ name: Dataset name.
+ split: dataset split, 'train' or 'valid'.
+ feature_sizes: shape(length) of each feature specified in the feature_names.
+ feature_names: names of the features in the tf.SequenceExample.
+ feature_sources: if the feature from 'context' or 'features'.
+ feature_dtypes: dtype of decoded feature.
+ feature_from_bytes: decode feature from bytes or as dtype list.
+ label_fields: name of field to read from tf.SequenceExample.
+ segment_size: Number of frames in each segment.
+ segment_labels: Use segment level label. Default: False, video level label.
+ include_video_id: `True` means include video id (string) in the input to the
+ model.
+ temporal_stride: Not used. Need to deprecated.
+ max_frames: Maxim Number of frames in a input example. It is used to crop
+ the input in the temporal dimension.
+ sample_random_frames: If sample random frames or random sequence.
+ num_sample_frames: Number of frames to sample for each input example. No
+ frame sampling if None.
+ num_classes: Number of classes to classify. Assuming it is a classification
+ task.
+ num_devices: Not used. To be deprecated.
+ input_path: The path to the input.
+ is_training: Whether this data is used for training or not.
+ num_examples: Number of examples in the dataset. It is used to compute the
+ steps for train or eval. set the value to `-1` to make the experiment run
+ until the end of dataset.
+ file_type: type of input files.
+ """
name: Optional[str] = 'yt8m'
split: Optional[str] = None
feature_sizes: Tuple[int, ...] = (1024, 128)
feature_names: Tuple[str, ...] = ('rgb', 'audio')
+ feature_sources: Tuple[str, ...] = ('feature', 'feature')
+ feature_dtypes: Tuple[str, ...] = ('uint8', 'uint8')
+ feature_from_bytes: Tuple[bool, ...] = (True, True)
+ label_field: str = 'labels'
segment_size: int = 1
segment_labels: bool = False
+ include_video_id: bool = False
temporal_stride: int = 1
- max_frames: int = 300
- num_frames: int = 300 # set smaller to allow random sample (Parser)
+ max_frames: int = 300 # Cap input frames.
+ sample_random_frames: bool = True
+ # Sample random frames if not None. No sampling in inference.
+ num_sample_frames: Optional[int] = 300
+ input_per_feature_l2_norm: bool = False
+ prefetch_buffer_size: int = 100
+ shuffle_buffer_size: int = 100
num_classes: int = 3862
num_devices: int = 1
input_path: str = ''
is_training: bool = True
- random_seed: int = 123
num_examples: int = -1
+ file_type: str = 'tfrecord'
def yt8m(is_training):
"""YT8M dataset configs."""
+ # pylint: disable=unexpected-keyword-arg
return DataConfig(
- num_frames=30,
temporal_stride=1,
segment_labels=False,
segment_size=5,
is_training=is_training,
split='train' if is_training else 'valid',
+ drop_remainder=is_training, # pytype: disable=wrong-keyword-args
num_examples=YT8M_TRAIN_EXAMPLES if is_training else YT8M_VAL_EXAMPLES,
input_path=YT8M_TRAIN_PATH if is_training else YT8M_VAL_PATH)
-
-
-@dataclasses.dataclass
-class MoeModel(hyperparams.Config):
- """The model config."""
- num_mixtures: int = 5
- l2_penalty: float = 1e-5
- use_input_context_gate: bool = False
- use_output_context_gate: bool = False
+ # pylint: enable=unexpected-keyword-arg
@dataclasses.dataclass
@@ -81,14 +112,65 @@ class DbofModel(hyperparams.Config):
cluster_size: int = 3000
hidden_size: int = 2000
add_batch_norm: bool = True
- sample_random_frames: bool = True
+ pooling_method: str = 'average'
use_context_gate_cluster_layer: bool = False
context_gate_cluster_bottleneck_size: int = 0
- pooling_method: str = 'average'
- yt8m_agg_classifier_model: str = 'MoeModel'
- agg_model: hyperparams.Config = MoeModel()
- norm_activation: common.NormActivation = common.NormActivation(
- activation='relu', use_sync_bn=False)
+
+
+@dataclasses.dataclass
+class Backbone(hyperparams.OneOfConfig):
+ """Configuration for backbones.
+
+ Attributes:
+ type: 'str', type of backbone be used, one of the fields below.
+ dbof: dbof backbone config.
+ """
+ type: Optional[str] = None
+ dbof: DbofModel = dataclasses.field(default_factory=DbofModel)
+
+
+@dataclasses.dataclass
+class MoeModel(hyperparams.Config):
+ """The MoE model config."""
+
+ num_mixtures: int = 5
+ vocab_as_last_dim: bool = False
+ use_input_context_gate: bool = False
+ use_output_context_gate: bool = False
+
+
+@dataclasses.dataclass
+class LogisticModel(hyperparams.Config):
+ """The logistic model config."""
+ return_logits: bool = False
+
+
+@dataclasses.dataclass
+class Head(hyperparams.OneOfConfig):
+ """Configuration for aggreagation heads.
+
+ Attributes:
+ type: 'str', type of head be used, one of the fields below.
+ moe: MoE head config.
+ logistic: Logistic head config.
+ """
+ type: Optional[str] = None
+ moe: MoeModel = dataclasses.field(default_factory=MoeModel)
+ logistic: LogisticModel = dataclasses.field(default_factory=LogisticModel)
+
+
+@dataclasses.dataclass
+class VideoClassificationModel(hyperparams.Config):
+ """The classifier model config."""
+ backbone: Backbone = dataclasses.field(
+ default_factory=lambda: Backbone(type='dbof')
+ )
+ head: Head = dataclasses.field(default_factory=lambda: Head(type='moe'))
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ activation='relu', use_sync_bn=False
+ )
+ )
@dataclasses.dataclass
@@ -99,17 +181,37 @@ class Losses(hyperparams.Config):
l2_weight_decay: float = 1e-5
+@dataclasses.dataclass
+class AveragePrecisionConfig(hyperparams.Config):
+ top_k: int = 20
+ top_n: Optional[int] = None
+ return_per_class_ap: bool = False
+
+
+@dataclasses.dataclass
+class Evaluation(hyperparams.Config):
+ average_precision: Optional[AveragePrecisionConfig] = None
+
+
@dataclasses.dataclass
class YT8MTask(cfg.TaskConfig):
"""The task config."""
- model: DbofModel = DbofModel()
- train_data: DataConfig = yt8m(is_training=True)
- validation_data: DataConfig = yt8m(is_training=False)
- losses: Losses = Losses()
+ model: VideoClassificationModel = dataclasses.field(
+ default_factory=VideoClassificationModel
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: yt8m(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: yt8m(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ evaluation: Evaluation = dataclasses.field(
+ default_factory=lambda: Evaluation( # pylint: disable=g-long-lambda
+ average_precision=AveragePrecisionConfig()
+ )
+ )
gradient_clip_norm: float = 1.0
- num_readers: int = 8
- top_k: int = 20
- top_n: Optional[int] = None
def add_trainer(
@@ -118,24 +220,26 @@ def add_trainer(
eval_batch_size: int,
learning_rate: float = 0.0001,
train_epochs: int = 50,
-):
- """Add and config a trainer to the experiment config."""
- if YT8M_TRAIN_EXAMPLES <= 0:
+ num_train_examples: int = YT8M_TRAIN_EXAMPLES,
+ num_val_examples: int = YT8M_VAL_EXAMPLES,
+) -> cfg.ExperimentConfig:
+ """Adds and config a trainer to the experiment config."""
+ if num_train_examples <= 0:
raise ValueError('Wrong train dataset size {!r}'.format(
experiment.task.train_data))
- if YT8M_VAL_EXAMPLES <= 0:
+ if num_val_examples <= 0:
raise ValueError('Wrong validation dataset size {!r}'.format(
experiment.task.validation_data))
experiment.task.train_data.global_batch_size = train_batch_size
experiment.task.validation_data.global_batch_size = eval_batch_size
- steps_per_epoch = YT8M_TRAIN_EXAMPLES // train_batch_size
- steps_per_loop = 30
+ steps_per_epoch = num_train_examples // train_batch_size
+ steps_per_loop = 500
experiment.trainer = cfg.TrainerConfig(
steps_per_loop=steps_per_loop,
summary_interval=steps_per_loop,
checkpoint_interval=steps_per_loop,
train_steps=train_epochs * steps_per_epoch,
- validation_steps=YT8M_VAL_EXAMPLES // eval_batch_size,
+ validation_steps=num_val_examples // eval_batch_size,
validation_interval=steps_per_loop,
optimizer_config=optimization.OptimizationConfig({
'optimizer': {
@@ -176,14 +280,16 @@ def yt8m_experiment() -> cfg.ExperimentConfig:
'task.train_data.num_classes == task.validation_data.num_classes',
'task.train_data.feature_sizes != None',
'task.train_data.feature_names != None',
+ 'task.train_data.feature_sources != None',
+ 'task.train_data.feature_dtypes != None',
])
# Per TPUv3 Core batch size 16GB HBM. `factor` in range(1, 26)
factor = 1
- num_cores = 32 # for TPU 4x4
+ num_cores = 32 # for TPUv3 4x4
train_per_core_bs = 32 * factor
train_bs = train_per_core_bs * num_cores
- eval_per_core_bs = 32 * 50 # multiplier<=100
+ eval_per_core_bs = 4 * 50 # multiplier<=100
eval_bs = eval_per_core_bs * num_cores
# based lr=0.0001 for bs=512
return add_trainer(
diff --git a/official/projects/yt8m/configs/yt8m_test.py b/official/projects/yt8m/configs/yt8m_test.py
new file mode 100644
index 00000000000..00f6333db10
--- /dev/null
+++ b/official/projects/yt8m/configs/yt8m_test.py
@@ -0,0 +1,40 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.core import config_definitions as cfg
+from official.core import exp_factory
+from official.modeling import hyperparams
+from official.projects.yt8m.configs import yt8m # pylint: disable=unused-import
+from official.projects.yt8m.configs.yt8m import yt8m as exp_cfg
+
+
+class YT8MTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.parameters(
+ ('yt8m_experiment',),)
+ def test_yt8m_configs(self, config_name):
+ config = exp_factory.get_exp_config(config_name)
+ self.assertIsInstance(config, cfg.ExperimentConfig)
+ self.assertIsInstance(config.task, cfg.TaskConfig)
+ self.assertIsInstance(config.task.model, hyperparams.Config)
+ self.assertIsInstance(config.task.train_data, cfg.DataConfig)
+ config.task.train_data.is_training = None
+ with self.assertRaises(KeyError):
+ config.validate()
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/yt8m/dataloaders/utils.py b/official/projects/yt8m/dataloaders/utils.py
index 9a5534e36f0..3bf971eb8c5 100644
--- a/official/projects/yt8m/dataloaders/utils.py
+++ b/official/projects/yt8m/dataloaders/utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,11 +15,12 @@
"""Contains a collection of util functions for training and evaluating."""
from absl import logging
-import numpy
-import tensorflow as tf
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.vision.dataloaders import tfexample_utils
-def Dequantize(feat_vector, max_quantized_value=2, min_quantized_value=-2):
+def dequantize(feat_vector, max_quantized_value=2, min_quantized_value=-2):
"""Dequantize the feature from the byte format to the float format.
Args:
@@ -37,7 +38,7 @@ def Dequantize(feat_vector, max_quantized_value=2, min_quantized_value=-2):
return feat_vector * scalar + bias
-def MakeSummary(name, value):
+def make_summary(name, value):
"""Creates a tf.Summary proto with the given name and value."""
summary = tf.Summary()
val = summary.value.add()
@@ -46,10 +47,10 @@ def MakeSummary(name, value):
return summary
-def AddGlobalStepSummary(summary_writer,
- global_step_val,
- global_step_info_dict,
- summary_scope="Eval"):
+def add_global_step_summary(summary_writer,
+ global_step_val,
+ global_step_info_dict,
+ summary_scope="Eval"):
"""Add the global_step summary to the Tensorboard.
Args:
@@ -68,19 +69,19 @@ def AddGlobalStepSummary(summary_writer,
examples_per_second = global_step_info_dict.get("examples_per_second", -1)
summary_writer.add_summary(
- MakeSummary("GlobalStep/" + summary_scope + "_Hit@1", this_hit_at_one),
+ make_summary("GlobalStep/" + summary_scope + "_Hit@1", this_hit_at_one),
global_step_val)
summary_writer.add_summary(
- MakeSummary("GlobalStep/" + summary_scope + "_Perr", this_perr),
+ make_summary("GlobalStep/" + summary_scope + "_Perr", this_perr),
global_step_val)
summary_writer.add_summary(
- MakeSummary("GlobalStep/" + summary_scope + "_Loss", this_loss),
+ make_summary("GlobalStep/" + summary_scope + "_Loss", this_loss),
global_step_val)
if examples_per_second != -1:
summary_writer.add_summary(
- MakeSummary("GlobalStep/" + summary_scope + "_Example_Second",
- examples_per_second), global_step_val)
+ make_summary("GlobalStep/" + summary_scope + "_Example_Second",
+ examples_per_second), global_step_val)
summary_writer.flush()
info = (
@@ -91,10 +92,10 @@ def AddGlobalStepSummary(summary_writer,
return info
-def AddEpochSummary(summary_writer,
- global_step_val,
- epoch_info_dict,
- summary_scope="Eval"):
+def add_epoch_summary(summary_writer,
+ global_step_val,
+ epoch_info_dict,
+ summary_scope="Eval"):
"""Add the epoch summary to the Tensorboard.
Args:
@@ -113,21 +114,21 @@ def AddEpochSummary(summary_writer,
avg_loss = epoch_info_dict["avg_loss"]
aps = epoch_info_dict["aps"]
gap = epoch_info_dict["gap"]
- mean_ap = numpy.mean(aps)
+ mean_ap = np.mean(aps)
summary_writer.add_summary(
- MakeSummary("Epoch/" + summary_scope + "_Avg_Hit@1", avg_hit_at_one),
+ make_summary("Epoch/" + summary_scope + "_Avg_Hit@1", avg_hit_at_one),
global_step_val)
summary_writer.add_summary(
- MakeSummary("Epoch/" + summary_scope + "_Avg_Perr", avg_perr),
+ make_summary("Epoch/" + summary_scope + "_Avg_Perr", avg_perr),
global_step_val)
summary_writer.add_summary(
- MakeSummary("Epoch/" + summary_scope + "_Avg_Loss", avg_loss),
+ make_summary("Epoch/" + summary_scope + "_Avg_Loss", avg_loss),
global_step_val)
summary_writer.add_summary(
- MakeSummary("Epoch/" + summary_scope + "_MAP", mean_ap), global_step_val)
+ make_summary("Epoch/" + summary_scope + "_MAP", mean_ap), global_step_val)
summary_writer.add_summary(
- MakeSummary("Epoch/" + summary_scope + "_GAP", gap), global_step_val)
+ make_summary("Epoch/" + summary_scope + "_GAP", gap), global_step_val)
summary_writer.flush()
info = ("epoch/eval number {0} | Avg_Hit@1: {1:.3f} | Avg_PERR: {2:.3f} "
@@ -137,7 +138,7 @@ def AddEpochSummary(summary_writer,
return info
-def GetListOfFeatureNamesAndSizes(feature_names, feature_sizes):
+def get_list_of_feature_names_and_sizes(feature_names, feature_sizes):
"""Extract the list of feature names and the dimensionality.
Args:
@@ -163,53 +164,126 @@ def GetListOfFeatureNamesAndSizes(feature_names, feature_sizes):
return list_of_feature_names, list_of_feature_sizes
-def ClipGradientNorms(gradients_to_variables, max_norm):
- """Clips the gradients by the given value.
+def make_yt8m_example(
+ num_segment: int = 5, num_frames: int = 120
+) -> tf.train.SequenceExample:
+ """Generate fake data for unit tests."""
+ rgb = np.random.randint(low=256, size=1024, dtype=np.uint8)
+ audio = np.random.randint(low=256, size=128, dtype=np.uint8)
+
+ seq_example = tf.train.SequenceExample()
+ seq_example.context.feature["id"].bytes_list.value[:] = [b"id001"]
+ seq_example.context.feature["labels"].int64_list.value[:] = [1, 2, 3, 4]
+ seq_example.context.feature["segment_labels"].int64_list.value[:] = (
+ [4] * num_segment)
+ seq_example.context.feature["segment_start_times"].int64_list.value[:] = [
+ i * 5 for i in range(num_segment)
+ ]
+ seq_example.context.feature["segment_scores"].float_list.value[:] = (
+ [0.5] * num_segment)
+ tfexample_utils.put_bytes_list_to_feature(
+ seq_example, rgb.tobytes(), key="rgb", repeat_num=num_frames)
+ tfexample_utils.put_bytes_list_to_feature(
+ seq_example, audio.tobytes(), key="audio", repeat_num=num_frames)
+
+ return seq_example
+
+
+# TODO(yeqing): Move the test related functions to test_utils.
+def make_example_with_float_features(
+ num_segment: int = 5) -> tf.train.SequenceExample:
+ """Generate fake data for unit tests."""
+ rgb = np.random.rand(1, 2048).astype(np.float32)
+ audio = np.random.rand(256).astype(np.float32)
+
+ seq_example = tf.train.SequenceExample()
+ seq_example.context.feature["id"].bytes_list.value[:] = [b"id001"]
+ seq_example.context.feature["clip/label/index"].int64_list.value[:] = [
+ 1, 2, 3, 4
+ ]
+ seq_example.context.feature["segment_labels"].int64_list.value[:] = (
+ [4] * num_segment)
+ seq_example.context.feature["segment_start_times"].int64_list.value[:] = [
+ i * 5 for i in range(num_segment)
+ ]
+ seq_example.context.feature["segment_scores"].float_list.value[:] = (
+ [0.] * num_segment)
+ seq_example.context.feature[
+ "VIDEO_EMBEDDING/context_feature/floats"].float_list.value[:] = (
+ audio.tolist())
+
+ tfexample_utils.put_float_list_to_feature(
+ seq_example, rgb.tolist(), key="FEATURE/feature/floats")
+
+ return seq_example
+
+
+def sample_random_sequence(batch_video_matrix, num_frames, num_samples):
+ """Samples a random sequence of frames of size num_samples.
Args:
- gradients_to_variables: A list of gradient to variable pairs (tuples).
- max_norm: the maximum norm value.
+ batch_video_matrix: tensor of shape [batch_size x max_frames x feature_size]
+ num_frames: tensor of shape [batch_size x 1]
+ num_samples: a scalar indicating the number of samples
Returns:
- A list of clipped gradient to variable pairs.
+ reshaped batch_video_matrix in [batch_size x 'num_samples' x feature_size]
"""
- clipped_grads_and_vars = []
- for grad, var in gradients_to_variables:
- if grad is not None:
- if isinstance(grad, tf.IndexedSlices):
- tmp = tf.clip_by_norm(grad.values, max_norm)
- grad = tf.IndexedSlices(tmp, grad.indices, grad.dense_shape)
- else:
- grad = tf.clip_by_norm(grad, max_norm)
- clipped_grads_and_vars.append((grad, var))
- return clipped_grads_and_vars
-
-def CombineGradients(tower_grads):
- """Calculate the combined gradient for each shared variable across all towers.
-
- Note that this function provides a synchronization point across all towers.
+ batch_size = tf.shape(batch_video_matrix)[0]
+ frame_index_offset = tf.tile(
+ tf.expand_dims(tf.range(num_samples), 0), [batch_size, 1])
+ max_start_frame_index = tf.maximum(num_frames - num_samples, 0)
+ start_frame_index = tf.cast(
+ tf.multiply(
+ tf.random.uniform([batch_size, 1]),
+ tf.cast(max_start_frame_index + 1, tf.float32)), tf.int32)
+ frame_index = tf.minimum(start_frame_index + frame_index_offset,
+ tf.cast(num_frames - 1, tf.int32))
+ batch_index = tf.tile(
+ tf.expand_dims(tf.range(batch_size), 1), [1, num_samples])
+ index = tf.stack([batch_index, frame_index], 2)
+ return tf.gather_nd(batch_video_matrix, index)
+
+
+def sample_random_frames(batch_video_matrix, num_frames, num_samples):
+ """Samples a random set of frames of size num_samples.
Args:
- tower_grads: List of lists of (gradient, variable) tuples. The outer list is
- over individual gradients. The inner list is over the gradient calculation
- for each tower.
+ batch_video_matrix: tensor of shape [batch_size x max_frames x feature_size]
+ num_frames: tensor of shape [batch_size x 1]
+ num_samples (int): a scalar indicating the number of samples
Returns:
- List of pairs of (gradient, variable) where the gradient has been summed
- across all towers.
+ reshaped batch_video_matrix in [batch_size x 'num_samples' x feature_size]
"""
- filtered_grads = [
- [x for x in grad_list if x[0] is not None] for grad_list in tower_grads
- ]
- final_grads = []
- for i in range(len(filtered_grads[0])):
- grads = [filtered_grads[t][i] for t in range(len(filtered_grads))]
- grad = tf.stack([x[0] for x in grads], 0)
- grad = tf.reduce_sum(grad, 0)
- final_grads.append((
- grad,
- filtered_grads[0][i][1],
- ))
-
- return final_grads
+ batch_size = tf.shape(batch_video_matrix)[0]
+ frame_index = tf.cast(
+ tf.multiply(
+ tf.random.uniform([batch_size, num_samples]),
+ tf.tile(num_frames, [1, num_samples])), tf.int32)
+ batch_index = tf.tile(
+ tf.expand_dims(tf.range(batch_size), 1), [1, num_samples])
+ index = tf.stack([batch_index, frame_index], 2)
+ return tf.gather_nd(batch_video_matrix, index)
+
+
+def sample_video_frames(
+ batch_video_matrix: tf.Tensor,
+ num_frames: tf.Tensor,
+ random_frames: bool = True,
+ num_sample_frames: int = 25,
+):
+ """Preprocesses input to sample frames."""
+
+ # Sample random frames / random sequence.
+ num_frames = tf.cast(num_frames, tf.float32)
+ if random_frames:
+ batch_video_matrix = sample_random_frames(
+ batch_video_matrix, num_frames, num_sample_frames
+ )
+ else:
+ batch_video_matrix = sample_random_sequence(
+ batch_video_matrix, num_frames, num_sample_frames
+ )
+ return batch_video_matrix
diff --git a/official/projects/yt8m/dataloaders/yt8m_input.py b/official/projects/yt8m/dataloaders/yt8m_input.py
index c2d57481b28..0a4a8f320c0 100644
--- a/official/projects/yt8m/dataloaders/yt8m_input.py
+++ b/official/projects/yt8m/dataloaders/yt8m_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -22,10 +22,9 @@
back into a range between min_quantized_value and max_quantized_value.
link for details: https://research.google.com/youtube8m/download.html
"""
+from typing import Any, Dict
-from typing import Dict
-
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.yt8m.dataloaders import utils
from official.vision.configs import video_classification as exp_cfg
from official.vision.dataloaders import decoder
@@ -33,7 +32,7 @@
def resize_axis(tensor, axis, new_size, fill_value=0):
- """Truncates or pads a tensor to new_size on on a given axis.
+ """Truncates or pads a tensor to new_size on a given axis.
Truncate or extend tensor such that tensor.shape[axis] == new_size. If the
size increases, the padding will be performed at the end, using fill_value.
@@ -81,13 +80,14 @@ def _process_segment_and_label(video_matrix, num_frames, contexts,
num_frames: Number of frames per subclip.
contexts: context information extracted from decoder
segment_labels: if we read segment labels instead.
- segment_size: the segment_size used for reading segments.
+ segment_size: the segment_size used for reading segments. Segment length.
num_classes: a positive integer for the number of classes.
Returns:
output: dictionary containing batch information
"""
# Partition frame-level feature matrix to segment-level feature matrix.
+ batch_video_ids = None
if segment_labels:
start_times = contexts["segment_start_times"].values
# Here we assume all the segments that started at the same start time has
@@ -101,11 +101,12 @@ def _process_segment_and_label(video_matrix, num_frames, contexts,
batch_video_matrix = tf.gather_nd(video_matrix,
tf.expand_dims(range_mtx, axis=-1))
num_segment = tf.shape(batch_video_matrix)[0]
- batch_video_ids = tf.reshape(
- tf.tile([contexts["id"]], [num_segment]), (num_segment,))
+ if "id" in contexts:
+ batch_video_ids = tf.reshape(
+ tf.tile([contexts["id"]], [num_segment]), (num_segment,))
batch_frames = tf.reshape(
- tf.tile([segment_size], [num_segment]), (num_segment,))
- batch_frames = tf.cast(tf.expand_dims(batch_frames, 1), tf.float32)
+ tf.tile([segment_size], [num_segment]), (num_segment, 1))
+ batch_frames = tf.cast(batch_frames, tf.int32)
# For segment labels, all labels are not exhaustively rated. So we only
# evaluate the rated labels.
@@ -123,43 +124,49 @@ def _process_segment_and_label(video_matrix, num_frames, contexts,
(num_segment, num_classes))
batch_label_weights = tf.sparse.to_dense(
sparse_label_weights, validate_indices=False)
- # output_dict = utils.get_segments(batch_video_matrix, batch_frames, 5)
+
else:
# Process video-level labels.
label_indices = contexts["labels"].values
sparse_labels = tf.sparse.SparseTensor(
tf.expand_dims(label_indices, axis=-1),
- tf.ones_like(contexts["labels"].values, dtype=tf.bool), (num_classes,))
+ tf.ones_like(contexts["labels"].values, dtype=tf.float32),
+ (num_classes,),
+ )
labels = tf.sparse.to_dense(
- sparse_labels, default_value=False, validate_indices=False)
+ sparse_labels, validate_indices=False)
# convert to batch format.
- batch_video_ids = tf.expand_dims(contexts["id"], 0)
+ if "id" in contexts:
+ batch_video_ids = tf.expand_dims(contexts["id"], 0)
batch_video_matrix = tf.expand_dims(video_matrix, 0)
batch_labels = tf.expand_dims(labels, 0)
batch_frames = tf.expand_dims(num_frames, 0)
batch_label_weights = None
output_dict = {
- "video_ids": batch_video_ids,
"video_matrix": batch_video_matrix,
"labels": batch_labels,
"num_frames": batch_frames,
}
+ if batch_video_ids is not None:
+ output_dict["video_ids"] = batch_video_ids
if batch_label_weights is not None:
output_dict["label_weights"] = batch_label_weights
return output_dict
-def _get_video_matrix(features, feature_size, max_frames, max_quantized_value,
- min_quantized_value):
+# TODO(allenyan, zhengxu): Adds a unit test for this function.
+def _get_video_matrix(features, feature_size, dtype, max_frames,
+ max_quantized_value, min_quantized_value):
"""Decodes features from an input string and quantizes it.
Args:
- features: raw feature values
- feature_size: length of each frame feature vector
- max_frames: number of frames (rows) in the output feature_matrix
+ features: raw feature values.
+ feature_size: length of each frame feature vector.
+ dtype: raw type of the feature.
+ max_frames: number of frames (rows) in the output feature_matrix.
max_quantized_value: the maximum of the quantized value.
min_quantized_value: the minimum of the quantized value.
@@ -167,28 +174,41 @@ def _get_video_matrix(features, feature_size, max_frames, max_quantized_value,
feature_matrix: matrix of all frame-features
num_frames: number of frames in the sequence
"""
- decoded_features = tf.reshape(
- tf.cast(tf.io.decode_raw(features, tf.uint8), tf.float32),
- [-1, feature_size])
+ decoded_features = tf.reshape(features, [-1, feature_size])
+
+ if dtype.is_integer:
+ feature_matrix = utils.dequantize(decoded_features, max_quantized_value,
+ min_quantized_value)
+ else:
+ feature_matrix = decoded_features
num_frames = tf.math.minimum(tf.shape(decoded_features)[0], max_frames)
- feature_matrix = utils.Dequantize(decoded_features, max_quantized_value,
- min_quantized_value)
- feature_matrix = resize_axis(feature_matrix, 0, max_frames)
+ feature_matrix = feature_matrix[:num_frames]
+
return feature_matrix, num_frames
-def _concat_features(features, feature_names, feature_sizes, max_frames,
- max_quantized_value, min_quantized_value):
+def _concat_features(
+ features,
+ feature_names,
+ feature_sizes,
+ feature_dtypes,
+ max_frames,
+ max_quantized_value,
+ min_quantized_value,
+ per_feature_l2_norm=False,
+):
"""Loads (potentially) different types of features and concatenates them.
Args:
features: raw feature values
feature_names: list of feature names
feature_sizes: list of features sizes
+ feature_dtypes: dtype of the feature.
max_frames: number of frames in the sequence
max_quantized_value: the maximum of the quantized value.
min_quantized_value: the minimum of the quantized value.
+ per_feature_l2_norm: whether to l2 normalize each feature.
Returns:
video_matrix: different features concatenated into one matrix
@@ -201,29 +221,35 @@ def _concat_features(features, feature_names, feature_sizes, max_frames,
assert len(feature_names) == len(feature_sizes), (
"length of feature_names (={}) != length of feature_sizes (={})".format(
len(feature_names), len(feature_sizes)))
+ assert len(feature_names) == len(feature_dtypes), (
+ "length of feature_names (={}) != length of feature_sizes (={})".format(
+ len(feature_names), len(feature_dtypes)))
- num_frames = -1 # the number of frames in the video
+ # the number of common frames of all features in the video
+ num_common_frames = 1080000 # set max to a 10-hour video at 30fps
feature_matrices = [None] * num_features # an array of different features
- for feature_index in range(num_features):
+ for i in range(num_features):
feature_matrix, num_frames_in_this_feature = _get_video_matrix(
- features[feature_names[feature_index]], feature_sizes[feature_index],
- max_frames, max_quantized_value, min_quantized_value)
- if num_frames == -1:
- num_frames = num_frames_in_this_feature
-
- feature_matrices[feature_index] = feature_matrix
-
- # cap the number of frames at self.max_frames
- num_frames = tf.minimum(num_frames, max_frames)
-
- # concatenate different features
+ features[feature_names[i]], feature_sizes[i],
+ tf.dtypes.as_dtype(feature_dtypes[i]), max_frames, max_quantized_value,
+ min_quantized_value)
+ num_common_frames = tf.math.minimum(num_frames_in_this_feature,
+ num_common_frames)
+ if per_feature_l2_norm:
+ feature_matrix = tf.math.l2_normalize(feature_matrix, axis=-1)
+ feature_matrices[i] = feature_matrix
+
+ for i in range(num_features):
+ feature_matrices[i] = feature_matrices[i][:num_common_frames] # pyrefly: ignore[unsupported-operation]
+
+ # Concatenate different features.
video_matrix = tf.concat(feature_matrices, 1)
- return video_matrix, num_frames
+ return video_matrix, num_common_frames
class Decoder(decoder.Decoder):
- """A tf.Example decoder for classification task."""
+ """A tf.train.SequeneExample decoder for classification task."""
def __init__(
self,
@@ -232,9 +258,22 @@ def __init__(
self._segment_labels = input_params.segment_labels
self._feature_names = input_params.feature_names
- self._context_features = {
- "id": tf.io.FixedLenFeature([], tf.string),
- }
+ self._feature_sources = input_params.feature_sources
+ self._feature_sizes = input_params.feature_sizes
+ self._feature_dtypes = input_params.feature_dtypes
+ self._feature_from_bytes = input_params.feature_from_bytes
+ self._include_video_id = input_params.include_video_id
+ self._label_field = input_params.label_field
+
+ assert len(self._feature_names) == len(self._feature_sources), (
+ "length of feature_names (={}) != length of feature_sizes (={})".format(
+ len(self._feature_names), len(self._feature_sources)))
+
+ self._context_features = {}
+ self._sequence_features = {}
+ if self._include_video_id:
+ self._context_features["id"] = tf.io.FixedLenFeature([], tf.string)
+
if self._segment_labels:
self._context_features.update({
# There is no need to read end-time given we always assume the segment
@@ -244,22 +283,51 @@ def __init__(
"segment_scores": tf.io.VarLenFeature(tf.float32)
})
else:
- self._context_features.update({"labels": tf.io.VarLenFeature(tf.int64)})
-
- self._sequence_features = {
- feature_name: tf.io.FixedLenSequenceFeature([], dtype=tf.string)
- for feature_name in self._feature_names
- }
-
- def decode(self, serialized_example):
- """Parses a single tf.Example into image and label tensors."""
+ self._add_labels_specification()
+ for i, name in enumerate(self._feature_names):
+ if self._feature_from_bytes[i]:
+ feature_type = tf.io.FixedLenSequenceFeature([], dtype=tf.string)
+ else:
+ dtype = tf.dtypes.as_dtype(self._feature_dtypes[i])
+ feature_shape = [self._feature_sizes[i]]
+ if self._feature_sources[i] == "feature":
+ feature_type = tf.io.FixedLenSequenceFeature(feature_shape, dtype)
+ else:
+ feature_type = tf.io.FixedLenFeature(feature_shape, dtype)
+ if self._feature_sources[i] == "feature":
+ self._sequence_features[name] = feature_type
+ elif self._feature_sources[i] == "context":
+ self._context_features[name] = feature_type
+ else:
+ raise ValueError(
+ f"Unknown feature source {self._feature_sources[i]} for {name}")
+
+ def _add_labels_specification(self):
+ if not self._label_field:
+ raise ValueError(f"Invalid label field: {self._label_field}!")
+ self._context_features.update(
+ {self._label_field: tf.io.VarLenFeature(tf.int64)})
+
+ def decode(self,
+ serialized_example: tf.train.SequenceExample) -> Dict[str, Any]:
+ """Parses a single tf.train.SequenceExample into video and label tensors."""
contexts, features = tf.io.parse_single_sequence_example(
serialized_example,
context_features=self._context_features,
sequence_features=self._sequence_features)
-
- return {"contexts": contexts, "features": features}
+ decoded_tensor = {**contexts, **features}
+ for i, name in enumerate(self._feature_names):
+ # Convert the VarLen feature to dense tensor.
+ if self._feature_from_bytes[i]:
+ dtype = tf.dtypes.as_dtype(self._feature_dtypes[i])
+ decoded_tensor[name] = tf.cast(
+ tf.io.decode_raw(decoded_tensor[name], out_type=dtype), tf.float32
+ )
+ else:
+ if isinstance(decoded_tensor[name], tf.SparseTensor):
+ decoded_tensor[name] = tf.sparse.to_dense(decoded_tensor[name])
+ return decoded_tensor
class Parser(parser.Parser):
@@ -278,42 +346,96 @@ def __init__(
min_quantized_value=-2,
):
self._num_classes = input_params.num_classes
+ self._label_field = input_params.label_field
self._segment_size = input_params.segment_size
self._segment_labels = input_params.segment_labels
+ self._include_video_id = input_params.include_video_id
self._feature_names = input_params.feature_names
+ self._feature_sources = input_params.feature_sources
self._feature_sizes = input_params.feature_sizes
- self.stride = input_params.temporal_stride
+ self._feature_dtypes = input_params.feature_dtypes
self._max_frames = input_params.max_frames
- self._num_frames = input_params.num_frames
- self._seed = input_params.random_seed
+ self._sample_random_frames = input_params.sample_random_frames
+ self._num_sample_frames = input_params.num_sample_frames
self._max_quantized_value = max_quantized_value
self._min_quantized_value = min_quantized_value
+ self._input_per_feature_l2_norm = input_params.input_per_feature_l2_norm
def _parse_train_data(self, decoded_tensors):
"""Parses data for training."""
# loads (potentially) different types of features and concatenates them
- self.video_matrix, self.num_frames = _concat_features(
- decoded_tensors["features"], self._feature_names, self._feature_sizes,
- self._max_frames, self._max_quantized_value, self._min_quantized_value)
- output_dict = _process_segment_and_label(self.video_matrix, self.num_frames,
- decoded_tensors["contexts"],
- self._segment_labels,
- self._segment_size,
- self._num_classes)
- return output_dict
+ video_matrix, num_frames = _concat_features(
+ decoded_tensors, self._feature_names, self._feature_sizes,
+ self._feature_dtypes, self._max_frames, self._max_quantized_value,
+ self._min_quantized_value, self._input_per_feature_l2_norm)
+ if not self._include_video_id and "id" in decoded_tensors:
+ del decoded_tensors["id"]
+
+ # Valid `num_frames` comes from _concat_features().
+ outputs = self._process_label(video_matrix, num_frames, decoded_tensors)
+ if self._num_sample_frames is None:
+ # Padding to max_frames.
+ outputs["video_matrix"] = resize_axis(
+ outputs["video_matrix"], 1, self._max_frames
+ )
+ else:
+ outputs["video_matrix"] = utils.sample_video_frames(
+ outputs["video_matrix"],
+ tf.reshape(outputs["num_frames"], [-1, 1]),
+ random_frames=self._sample_random_frames,
+ num_sample_frames=self._num_sample_frames,
+ )
+ outputs["num_frames"] = (
+ tf.ones_like(outputs["num_frames"]) * self._num_sample_frames
+ )
+ return outputs
def _parse_eval_data(self, decoded_tensors):
"""Parses data for evaluation."""
# loads (potentially) different types of features and concatenates them
- self.video_matrix, self.num_frames = _concat_features(
- decoded_tensors["features"], self._feature_names, self._feature_sizes,
- self._max_frames, self._max_quantized_value, self._min_quantized_value)
- output_dict = _process_segment_and_label(self.video_matrix, self.num_frames,
- decoded_tensors["contexts"],
+ video_matrix, num_frames = _concat_features(
+ decoded_tensors, self._feature_names, self._feature_sizes,
+ self._feature_dtypes, self._max_frames, self._max_quantized_value,
+ self._min_quantized_value, self._input_per_feature_l2_norm)
+ if not self._include_video_id and "id" in decoded_tensors:
+ del decoded_tensors["id"]
+
+ outputs = self._process_label(video_matrix, num_frames, decoded_tensors)
+ if self._num_sample_frames is None:
+ # Padding to max_frames.
+ outputs["video_matrix"] = resize_axis(
+ outputs["video_matrix"], 1, self._max_frames
+ )
+ else:
+ outputs["video_matrix"] = utils.sample_video_frames(
+ outputs["video_matrix"],
+ tf.reshape(outputs["num_frames"], [-1, 1]),
+ random_frames=self._sample_random_frames,
+ num_sample_frames=self._num_sample_frames,
+ )
+ outputs["num_frames"] = (
+ tf.ones_like(outputs["num_frames"]) * self._num_sample_frames
+ )
+ return outputs
+
+ def _process_label(self, video_matrix, num_frames, contexts):
+ """Processes a batched Tensor of frames.
+
+ Args:
+ video_matrix: video feature matric.
+ num_frames: number of frames in this video.
+ contexts: context information extracted from decoder.
+
+ Returns:
+ output: dictionary containing batch information
+ """
+ if self._label_field and not self._segment_labels:
+ contexts["labels"] = contexts[self._label_field]
+ output_dict = _process_segment_and_label(video_matrix, num_frames, contexts,
self._segment_labels,
self._segment_size,
self._num_classes)
- return output_dict # batched
+ return output_dict
def parse_fn(self, is_training):
"""Returns a parse fn that reads and parses raw tensors from the decoder.
@@ -329,6 +451,27 @@ def parse_fn(self, is_training):
def parse(decoded_tensors):
"""Parses the serialized example data."""
+
+ # Concatenate video features to all frames if there are both video-level
+ # (context) and frame-level (feature) features.
+ if "feature" in self._feature_sources:
+ # Take first frame feature matrix, any feature matrix should be fine
+ # since assume all frame features have same number of frames.
+ feature_idx = self._feature_sources.index("feature")
+ num_frames = tf.shape(
+ decoded_tensors[self._feature_names[feature_idx]]
+ )[0]
+ for feature_idx, feature_source in enumerate(self._feature_sources):
+ if feature_source == "context":
+ feature_name = self._feature_names[feature_idx]
+ context_tensor = tf.reshape(
+ decoded_tensors[feature_name],
+ shape=(1, self._feature_sizes[feature_idx]),
+ )
+ decoded_tensors[feature_name] = tf.tile(
+ context_tensor, [num_frames, 1]
+ )
+
if is_training:
return self._parse_train_data(decoded_tensors)
else:
@@ -337,83 +480,84 @@ def parse(decoded_tensors):
return parse
+class TransformBatcher():
+ """Performs manual batching on input dataset."""
+
+ def __init__(self, input_params: exp_cfg.DataConfig):
+ self._segment_labels = input_params.segment_labels
+ self._global_batch_size = input_params.global_batch_size
+ self._is_training = input_params.is_training
+ self._include_video_id = input_params.include_video_id
+ self._drop_remainder = input_params.drop_remainder
+
+ def batch_fn(self, dataset, input_context):
+ """Add padding when segment_labels is true."""
+ per_replica_batch_size = input_context.get_per_replica_batch_size(
+ self._global_batch_size) if input_context else self._global_batch_size
+ # Add padding specifications.
+ pad_values = {
+ "video_matrix": 0.0,
+ "labels": -1.0,
+ "num_frames": 0,
+ }
+ if self._include_video_id:
+ pad_values["video_ids"] = None # pyrefly: ignore[bad-assignment]
+ if self._segment_labels:
+ pad_values["label_weights"] = 0.0
+ dataset = dataset.padded_batch(
+ per_replica_batch_size,
+ padding_values=pad_values,
+ drop_remainder=self._drop_remainder,
+ )
+ return dataset
+
+
class PostBatchProcessor():
"""Processes a video and label dataset which is batched."""
def __init__(self, input_params: exp_cfg.DataConfig):
self.segment_labels = input_params.segment_labels
self.num_classes = input_params.num_classes
- self.segment_size = input_params.segment_size
+ self.num_batched_frames = (
+ input_params.num_sample_frames or input_params.max_frames
+ )
+ self.num_features = sum(input_params.feature_sizes)
- def post_fn(self, batched_tensors):
+ def post_fn(self, batched_tensors: Dict[str,
+ tf.Tensor]) -> Dict[str, tf.Tensor]:
"""Processes batched Tensors."""
- video_ids = batched_tensors["video_ids"]
+ video_ids = batched_tensors.get("video_ids", None)
video_matrix = batched_tensors["video_matrix"]
labels = batched_tensors["labels"]
num_frames = batched_tensors["num_frames"]
- label_weights = None
if self.segment_labels:
- # [batch x num_segment x segment_size x num_features]
- # -> [batch * num_segment x segment_size x num_features]
- video_ids = tf.reshape(video_ids, [-1])
- video_matrix = tf.reshape(video_matrix, [-1, self.segment_size, 1152])
+ # [batch x num_segment x num_batched_frames x num_features]
+ # -> [batch * num_segment x num_batched_frames x num_features]
+ if video_ids is not None:
+ video_ids = tf.reshape(video_ids, [-1])
+ video_matrix = tf.reshape(
+ video_matrix, [-1, self.num_batched_frames, self.num_features]
+ )
labels = tf.reshape(labels, [-1, self.num_classes])
num_frames = tf.reshape(num_frames, [-1, 1])
-
- label_weights = tf.reshape(batched_tensors["label_weights"],
- [-1, self.num_classes])
-
+ batched_tensors["label_weights"] = tf.reshape(
+ batched_tensors["label_weights"], [-1, self.num_classes])
else:
- video_matrix = tf.squeeze(video_matrix)
- labels = tf.squeeze(labels)
+ # NOTE(b/237445211): Must provide axis argument to tf.squeeze.
+ video_matrix = tf.squeeze(video_matrix, axis=1)
+ labels = tf.squeeze(labels, axis=1)
+ num_frames = tf.reshape(num_frames, [-1, 1])
+ if "label_weights" in batched_tensors:
+ batched_tensors["label_weights"] = tf.squeeze(
+ batched_tensors["label_weights"], axis=1)
- batched_tensors = {
- "video_ids": video_ids,
+ batched_tensors.update({
"video_matrix": video_matrix,
"labels": labels,
"num_frames": num_frames,
- }
-
- if label_weights is not None:
- batched_tensors["label_weights"] = label_weights
+ })
+ if video_ids is not None:
+ batched_tensors["video_ids"] = video_ids
return batched_tensors
-
-
-class TransformBatcher():
- """Performs manual batching on input dataset."""
-
- def __init__(self, input_params: exp_cfg.DataConfig):
- self._segment_labels = input_params.segment_labels
- self._global_batch_size = input_params.global_batch_size
- self._is_training = input_params.is_training
-
- def batch_fn(self, dataset, input_context):
- """Add padding when segment_labels is true."""
- per_replica_batch_size = input_context.get_per_replica_batch_size(
- self._global_batch_size) if input_context else self._global_batch_size
- if not self._segment_labels:
- dataset = dataset.batch(per_replica_batch_size, drop_remainder=True)
- else:
- # add padding
- pad_shapes = {
- "video_ids": [None],
- "video_matrix": [None, None, None],
- "labels": [None, None],
- "num_frames": [None, None],
- "label_weights": [None, None]
- }
- pad_values = {
- "video_ids": None,
- "video_matrix": 0.0,
- "labels": -1.0,
- "num_frames": 0.0,
- "label_weights": 0.0
- }
- dataset = dataset.padded_batch(
- per_replica_batch_size,
- padded_shapes=pad_shapes,
- drop_remainder=True,
- padding_values=pad_values)
- return dataset
diff --git a/official/projects/yt8m/dataloaders/yt8m_input_test.py b/official/projects/yt8m/dataloaders/yt8m_input_test.py
new file mode 100644
index 00000000000..c61bef91229
--- /dev/null
+++ b/official/projects/yt8m/dataloaders/yt8m_input_test.py
@@ -0,0 +1,248 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import os
+
+from absl import logging
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.core import input_reader
+from official.projects.yt8m.configs import yt8m as yt8m_configs
+from official.projects.yt8m.dataloaders import utils
+from official.projects.yt8m.dataloaders import yt8m_input
+from official.vision.dataloaders import tfexample_utils
+
+
+class Yt8mInputTest(parameterized.TestCase, tf.test.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ self._model_dir = os.path.join(self.get_temp_dir(), 'model_dir')
+ tf.io.gfile.makedirs(self._model_dir)
+
+ data_dir = os.path.join(self.get_temp_dir(), 'data')
+ tf.io.gfile.makedirs(data_dir)
+ self.data_path = os.path.join(data_dir, 'data.tfrecord')
+ self.num_segment = 6
+ examples = [
+ utils.make_yt8m_example(self.num_segment, 120 + i) for i in range(8)
+ ]
+ tfexample_utils.dump_to_tfrecord(self.data_path, tf_examples=examples)
+
+ def create_input_reader(self, params):
+ decoder = yt8m_input.Decoder(input_params=params)
+ decoder_fn = decoder.decode
+ parser = yt8m_input.Parser(input_params=params)
+ parser_fn = parser.parse_fn(params.is_training)
+ postprocess = yt8m_input.PostBatchProcessor(input_params=params)
+ postprocess_fn = postprocess.post_fn
+ transform_batch = yt8m_input.TransformBatcher(input_params=params)
+ batch_fn = transform_batch.batch_fn
+
+ return input_reader.InputReader(
+ params,
+ dataset_fn=tf.data.TFRecordDataset,
+ decoder_fn=decoder_fn,
+ parser_fn=parser_fn,
+ postprocess_fn=postprocess_fn,
+ transform_and_batch_fn=batch_fn)
+
+ @parameterized.parameters((True, 20), (False, 20), (False, None))
+ def test_read_video_level_input(self, include_video_id, num_sample_frames):
+ params = yt8m_configs.yt8m(is_training=False)
+ params.global_batch_size = 4
+ params.segment_labels = False
+ params.input_path = self.data_path
+ params.include_video_id = include_video_id
+ params.max_frames = 122
+ params.num_sample_frames = num_sample_frames
+ reader = self.create_input_reader(params)
+
+ dataset = reader.read()
+ iterator = iter(dataset)
+ example = next(iterator)
+
+ for k, v in example.items():
+ logging.info('DEBUG read example %r %r %r', k, v.shape, type(v))
+ if include_video_id:
+ self.assertCountEqual(
+ ['video_matrix', 'labels', 'num_frames', 'video_ids'], example.keys())
+ else:
+ self.assertCountEqual(['video_matrix', 'labels', 'num_frames'],
+ example.keys())
+ batch_size = params.global_batch_size
+ expected_num_frames = num_sample_frames or params.max_frames
+ self.assertEqual(
+ example['video_matrix'].shape.as_list(),
+ [batch_size, expected_num_frames, sum(params.feature_sizes)],
+ )
+ self.assertEqual(
+ example['labels'].shape.as_list(), [batch_size, params.num_classes]
+ )
+ # Check non empty labels.
+ self.assertGreater(np.nonzero(example['labels'][0].numpy())[0].shape[0], 0)
+
+ if num_sample_frames:
+ self.assertAllEqual(
+ example['num_frames'].numpy(),
+ [[num_sample_frames]] * batch_size,
+ )
+ else:
+ self.assertAllEqual(
+ example['num_frames'].numpy(),
+ [[120], [121], [122], [122]],
+ )
+
+ if include_video_id:
+ self.assertEqual(example['video_ids'].shape.as_list(), [batch_size, 1])
+
+ @parameterized.parameters((True, 20), (False, 20), (False, None))
+ def test_read_segment_level_input(self, include_video_id, num_sample_frames):
+ params = yt8m_configs.yt8m(is_training=False)
+ params.global_batch_size = 2
+ params.segment_labels = True
+ params.segment_size = 24
+ params.input_path = self.data_path
+ params.include_video_id = include_video_id
+ params.num_sample_frames = num_sample_frames
+ reader = self.create_input_reader(params)
+
+ dataset = reader.read()
+ iterator = iter(dataset)
+ example = next(iterator)
+
+ for k, v in example.items():
+ logging.info('DEBUG read example %r %r %r', k, v.shape, type(v))
+ if include_video_id:
+ self.assertCountEqual([
+ 'video_matrix', 'labels', 'num_frames', 'label_weights', 'video_ids'
+ ], example.keys())
+ else:
+ self.assertCountEqual(
+ ['video_matrix', 'labels', 'num_frames', 'label_weights'],
+ example.keys())
+ batch_size = params.global_batch_size * self.num_segment
+ expected_num_frames = num_sample_frames or params.max_frames
+ self.assertEqual(
+ example['video_matrix'].shape.as_list(),
+ [batch_size, expected_num_frames, sum(params.feature_sizes)],
+ )
+ self.assertEqual(example['labels'].shape.as_list(),
+ [batch_size, params.num_classes])
+ self.assertGreater(np.nonzero(example['labels'][0].numpy())[0].shape[0], 0)
+ self.assertEqual(example['label_weights'].shape.as_list(),
+ [batch_size, params.num_classes])
+
+ if num_sample_frames:
+ self.assertAllEqual(
+ example['num_frames'].numpy(),
+ [[num_sample_frames]] * batch_size,
+ )
+ else:
+ self.assertAllEqual(
+ example['num_frames'].numpy(),
+ [[params.segment_size]] * batch_size,
+ )
+
+ if include_video_id:
+ self.assertEqual(example['video_ids'].shape.as_list(), [batch_size])
+
+ @parameterized.parameters(
+ (True, 4, False),
+ (False, 4, False),
+ (False, None, False),
+ (True, 4, True),
+ (False, 4, True),
+ (False, None, True),
+ )
+ def test_read_video_level_float_input(
+ self, include_video_id, num_sample_frames, per_feature_l2_norm
+ ):
+ data_dir = os.path.join(self.get_temp_dir(), 'data2')
+ tf.io.gfile.makedirs(data_dir)
+ data_path = os.path.join(data_dir, 'data2.tfrecord')
+ examples = [
+ utils.make_example_with_float_features(self.num_segment)
+ for _ in range(8)
+ ]
+ tfexample_utils.dump_to_tfrecord(data_path, tf_examples=examples)
+
+ params = yt8m_configs.yt8m(is_training=False)
+ params.global_batch_size = 4
+ params.segment_labels = False
+ params.input_path = data_path
+ params.num_frames = 2
+ params.max_frames = 2
+ params.num_sample_frames = num_sample_frames
+ params.feature_names = ('VIDEO_EMBEDDING/context_feature/floats',
+ 'FEATURE/feature/floats')
+ params.feature_sources = ('context', 'feature')
+ params.feature_dtypes = ('float32', 'float32')
+ params.feature_sizes = (256, 2048)
+ params.feature_from_bytes = (False, False)
+ params.label_field = 'clip/label/index'
+ params.include_video_id = include_video_id
+ params.input_per_feature_l2_norm = per_feature_l2_norm
+ reader = self.create_input_reader(params)
+
+ dataset = reader.read()
+ iterator = iter(dataset)
+ example = next(iterator)
+
+ for k, v in example.items():
+ logging.info('DEBUG read example %r %r %r', k, v.shape, type(v))
+ logging.info('DEBUG read example %r', example['video_matrix'][0, 0, :])
+ if include_video_id:
+ self.assertCountEqual(
+ ['video_matrix', 'labels', 'num_frames', 'video_ids'], example.keys())
+ else:
+ self.assertCountEqual(['video_matrix', 'labels', 'num_frames'],
+ example.keys())
+
+ # Check tensor values.
+ expected_context = examples[0].context.feature[
+ 'VIDEO_EMBEDDING/context_feature/floats'].float_list.value
+ expected_feature = examples[0].feature_lists.feature_list[
+ 'FEATURE/feature/floats'].feature[0].float_list.value
+ expected_labels = examples[0].context.feature[
+ params.label_field].int64_list.value
+ if per_feature_l2_norm:
+ expected_feature = tf.math.l2_normalize(expected_feature, axis=-1)
+ expected_context = tf.math.l2_normalize(expected_context, axis=-1)
+ self.assertAllEqual(expected_feature,
+ example['video_matrix'][0, 0, params.feature_sizes[0]:])
+ self.assertAllEqual(expected_context,
+ example['video_matrix'][0, 0, :params.feature_sizes[0]])
+ self.assertAllEqual(
+ np.nonzero(example['labels'][0, :].numpy())[0], expected_labels)
+ self.assertGreater(np.nonzero(example['labels'][0].numpy())[0].shape[0], 0)
+
+ # Check tensor shape.
+ batch_size = params.global_batch_size
+ expected_num_frames = params.num_sample_frames or params.max_frames
+ self.assertEqual(
+ example['video_matrix'].shape.as_list(),
+ [batch_size, expected_num_frames, sum(params.feature_sizes)],
+ )
+ self.assertEqual(example['labels'].shape.as_list(),
+ [batch_size, params.num_classes])
+ self.assertEqual(example['num_frames'].shape.as_list(), [batch_size, 1])
+ if include_video_id:
+ self.assertEqual(example['video_ids'].shape.as_list(), [batch_size, 1])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/yt8m/eval_utils/average_precision_calculator.py b/official/projects/yt8m/eval_utils/average_precision_calculator.py
index 16cb71a8109..ff96b3cdbba 100644
--- a/official/projects/yt8m/eval_utils/average_precision_calculator.py
+++ b/official/projects/yt8m/eval_utils/average_precision_calculator.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -114,7 +114,7 @@ def accumulate(self, predictions, actuals, num_positives=None):
raise ValueError("the shape of predictions and actuals does not match.")
if num_positives is not None:
- if not isinstance(num_positives, numbers.Number) or num_positives < 0:
+ if not isinstance(num_positives, numbers.Number) or num_positives < 0: # pyrefly: ignore[unsupported-operation]
raise ValueError(
"'num_positives' was provided but it was a negative number.")
@@ -268,6 +268,5 @@ def _zero_one_normalize(predictions, epsilon=1e-7):
The normalized prediction.
"""
denominator = numpy.max(predictions) - numpy.min(predictions)
- ret = (predictions - numpy.min(predictions)) / numpy.max(
- denominator, epsilon)
+ ret = (predictions - numpy.min(predictions)) / max(denominator, epsilon)
return ret
diff --git a/official/projects/yt8m/eval_utils/eval_util.py b/official/projects/yt8m/eval_utils/eval_util.py
index d61d9cdfcdb..b884a4de923 100644
--- a/official/projects/yt8m/eval_utils/eval_util.py
+++ b/official/projects/yt8m/eval_utils/eval_util.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,8 +13,9 @@
# limitations under the License.
"""Provides functions to help with evaluating models."""
+import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.yt8m.eval_utils import average_precision_calculator as ap_calculator
from official.projects.yt8m.eval_utils import mean_average_precision_calculator as map_calculator
@@ -57,6 +58,9 @@ def calculate_precision_at_equal_recall_rate(predictions, actuals):
"""
aggregated_precision = 0.0
num_videos = actuals.shape[0]
+ if num_videos == 0:
+ logging.warning("Num_videos is 0, returning 0.0 aggregated_precision.")
+ return aggregated_precision
for row in np.arange(num_videos):
num_labels = int(np.sum(actuals[row]))
top_indices = np.argpartition(predictions[row], -num_labels)[-num_labels:]
@@ -99,8 +103,8 @@ def top_k_by_class(predictions, labels, k=20):
Args:
predictions: A numpy matrix containing the outputs of the model. Dimensions
are 'batch' x 'num_classes'.
- labels: A numpy matrix containing the ground truth labels.
- Dimensions are 'batch' x 'num_classes'.
+ labels: A numpy matrix containing the ground truth labels. Dimensions are
+ 'batch' x 'num_classes'.
k: the top k non-zero entries to preserve in each prediction.
Returns:
@@ -139,9 +143,10 @@ def top_k_triplets(predictions, labels, k=20):
Args:
predictions: A numpy matrix containing the outputs of the model. Dimensions
are 'batch' x 'num_classes'.
- labels: A numpy matrix containing the ground truth labels.
- Dimensions are 'batch' x 'num_classes'.
+ labels: A numpy matrix containing the ground truth labels. Dimensions are
+ 'batch' x 'num_classes'.
k: The number top predictions to pick.
+
Returns:
a sparse list of tuples in (prediction, class) format.
"""
@@ -171,7 +176,7 @@ def __init__(self, num_class, top_k, top_n):
self.sum_hit_at_one = 0.0
self.sum_perr = 0.0
self.map_calculator = map_calculator.MeanAveragePrecisionCalculator(
- num_class, top_n=top_n)
+ num_class, filter_empty_classes=False, top_n=top_n)
self.global_ap_calculator = ap_calculator.AveragePrecisionCalculator()
self.top_k = top_k
self.num_examples = 0
@@ -213,9 +218,13 @@ def accumulate(self, predictions, labels):
return {"hit_at_one": mean_hit_at_one, "perr": mean_perr}
- def get(self):
+ def get(self, return_per_class_ap=False):
"""Calculate the evaluation metrics for the whole epoch.
+ Args:
+ return_per_class_ap: a bool variable to determine whether return the
+ detailed class-wise ap for more detailed analysis. Default is `False`.
+
Raises:
ValueError: If no examples were accumulated.
@@ -232,13 +241,19 @@ def get(self):
aps = self.map_calculator.peek_map_at_n()
mean_ap = sum(aps) / self.num_class
gap = self.global_ap_calculator.peek_ap_at_n()
+ lw_map = self.map_calculator.peek_log_weighted_map_at_n()
epoch_info_dict = {
"avg_hit_at_one": avg_hit_at_one,
"avg_perr": avg_perr,
"map": mean_ap,
- "gap": gap
+ "gap": gap,
+ "lw_map": lw_map
}
+
+ if return_per_class_ap:
+ epoch_info_dict["per_class_ap"] = aps
+
return epoch_info_dict
def clear(self):
@@ -265,5 +280,5 @@ def _convert_to_numpy(self, groundtruths, predictions):
else:
outputs = predictions
- labels = labels * 1
+ labels = labels * 1 # pyrefly: ignore[unsupported-operation]
return outputs, labels
diff --git a/official/projects/yt8m/eval_utils/eval_util_test.py b/official/projects/yt8m/eval_utils/eval_util_test.py
new file mode 100644
index 00000000000..791db4a60bb
--- /dev/null
+++ b/official/projects/yt8m/eval_utils/eval_util_test.py
@@ -0,0 +1,71 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from absl import logging
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.projects.yt8m.eval_utils.average_precision_calculator import AveragePrecisionCalculator
+
+
+class YT8MAveragePrecisionCalculatorTest(parameterized.TestCase,
+ tf.test.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ self.prediction = np.array([
+ [0.98, 0.88, 0.77, 0.65, 0.64, 0.59, 0.45, 0.43, 0.20, 0.05],
+ [0.878, 0.832, 0.759, 0.621, 0.458, 0.285, 0.134],
+ [0.98],
+ [0.56],
+ ], dtype=object)
+ self.raw_prediction = np.random.rand(5, 10) + np.random.randint(
+ low=0, high=10, size=(5, 10))
+ self.ground_truth = np.array([[1, 1, 0, 0, 0, 1, 1, 0, 0, 1],
+ [1, 0, 1, 0, 0, 1, 0], [1], [0]],
+ dtype=object)
+
+ self.expected_ap = np.array([
+ 0.714,
+ 0.722,
+ 1.000,
+ 0.000,
+ ])
+
+ def test_ap_calculator_ap(self):
+
+ # Compare Expected Average Precision with function expected
+ for i, _ in enumerate(self.ground_truth):
+ calculator = AveragePrecisionCalculator()
+ ap = calculator.ap(self.prediction[i], self.ground_truth[i])
+ logging.info('DEBUG %dth AP: %r', i + 1, ap)
+
+ def test_ap_calculator_zero_one_normalize(self):
+ for i, _ in enumerate(self.raw_prediction):
+ calculator = AveragePrecisionCalculator()
+ logging.error('%r', self.raw_prediction[i])
+ normalized_score = calculator._zero_one_normalize(self.raw_prediction[i])
+ self.assertAllInRange(normalized_score, lower_bound=0.0, upper_bound=1.0)
+
+ @parameterized.parameters((None,), (3,), (5,), (10,), (20,))
+ def test_ap_calculator_ap_at_n(self, n):
+ for i, _ in enumerate(self.ground_truth):
+ calculator = AveragePrecisionCalculator(n)
+ ap = calculator.ap_at_n(self.prediction[i], self.ground_truth[i], n)
+ logging.info('DEBUG %dth AP: %r', i + 1, ap)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/projects/yt8m/eval_utils/mean_average_precision_calculator.py b/official/projects/yt8m/eval_utils/mean_average_precision_calculator.py
index a5ed00000c9..f79bfa3d392 100644
--- a/official/projects/yt8m/eval_utils/mean_average_precision_calculator.py
+++ b/official/projects/yt8m/eval_utils/mean_average_precision_calculator.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -37,6 +37,7 @@
```
"""
+import numpy as np
from official.projects.yt8m.eval_utils import average_precision_calculator
@@ -56,7 +57,7 @@ def __init__(self, num_class, filter_empty_classes=True, top_n=None):
ValueError: An error occurred when num_class is not a positive integer;
or the top_n_array is not a list of positive integers.
"""
- if not isinstance(num_class, int) or num_class <= 1:
+ if not isinstance(num_class, int) or num_class < 1:
raise ValueError("num_class must be a positive integer.")
self._ap_calculators = [] # member of AveragePrecisionCalculator
@@ -112,3 +113,24 @@ def peek_map_at_n(self):
ap = self._ap_calculators[i].peek_ap_at_n()
aps.append(ap)
return aps
+
+ def peek_log_weighted_map_at_n(self):
+ """Peek the non-interpolated log weighted mean average precision at n.
+
+ Returns:
+ Log weighted mean average precision.
+ """
+ sum_log_weighted_ap = 0
+ sum_log_weights = 0
+ for i in range(self._num_class):
+ pos = self._ap_calculators[i].num_accumulated_positives
+ if not self._filter_empty_classes or pos > 0:
+ ap = self._ap_calculators[i].peek_ap_at_n()
+ # TODO(b/286928055)
+ log_pos = np.log(1 + pos)
+ sum_log_weights += log_pos
+ sum_log_weighted_ap += ap * log_pos
+
+ if not sum_log_weights:
+ return 0
+ return sum_log_weighted_ap / sum_log_weights
diff --git a/official/projects/yt8m/experiments/yt8m.yaml b/official/projects/yt8m/experiments/yt8m.yaml
index c099f23f90b..f1354c78b57 100644
--- a/official/projects/yt8m/experiments/yt8m.yaml
+++ b/official/projects/yt8m/experiments/yt8m.yaml
@@ -22,12 +22,11 @@ task:
segment_labels: false
temporal_stride: 1
max_frames: 300
- num_frames: 300
+ num_sample_frames: 300
num_classes: 3862
num_devices: 1
input_path: 'gs://youtube8m-ml/2/frame/train/train*.tfrecord'
is_training: true
- random_seed: 123
validation_data:
name: 'yt8m'
split: 'train'
@@ -41,12 +40,11 @@ task:
segment_labels: true
temporal_stride: 1
max_frames: 300
- num_frames: 300
+ num_sample_frames: 300
num_classes: 3862
num_devices: 1
input_path: 'gs://youtube8m-ml/3/frame/validate/validate*.tfrecord'
is_training: false
- random_seed: 123
losses:
name: 'binary_crossentropy'
from_logits: false
diff --git a/official/projects/yt8m/modeling/__init__.py b/official/projects/yt8m/modeling/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/projects/yt8m/modeling/__init__.py
+++ b/official/projects/yt8m/modeling/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/yt8m/modeling/backbones/__init__.py b/official/projects/yt8m/modeling/backbones/__init__.py
new file mode 100644
index 00000000000..e39385ba424
--- /dev/null
+++ b/official/projects/yt8m/modeling/backbones/__init__.py
@@ -0,0 +1,17 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Backbones package definition."""
+
+from official.projects.yt8m.modeling.backbones import dbof
diff --git a/official/projects/yt8m/modeling/backbones/dbof.py b/official/projects/yt8m/modeling/backbones/dbof.py
new file mode 100644
index 00000000000..8abdbdbd5c5
--- /dev/null
+++ b/official/projects/yt8m/modeling/backbones/dbof.py
@@ -0,0 +1,186 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Dbof model definitions."""
+
+import functools
+from typing import Any, Optional
+
+import tensorflow as tf, tf_keras
+
+from official.modeling import hyperparams
+from official.modeling import tf_utils
+from official.projects.yt8m.configs import yt8m as yt8m_cfg
+from official.projects.yt8m.modeling import nn_layers
+from official.projects.yt8m.modeling import yt8m_model_utils
+from official.vision.configs import common
+from official.vision.modeling.backbones import factory
+
+
+layers = tf_keras.layers
+
+
+class Dbof(layers.Layer):
+ """A YT8M model class builder.
+
+ Creates a Deep Bag of Frames model.
+ The model projects the features for each frame into a higher dimensional
+ 'clustering' space, pools across frames in that space, and then
+ uses a configurable video-level model to classify the now aggregated features.
+ The model will randomly sample either frames or sequences of frames during
+ training to speed up convergence.
+ """
+
+ def __init__(
+ self,
+ input_specs: layers.InputSpec = layers.InputSpec(
+ shape=[None, None, 1152]
+ ),
+ params: yt8m_cfg.DbofModel = yt8m_cfg.DbofModel(),
+ norm_activation: common.NormActivation = common.NormActivation(),
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs,
+ ):
+ """YT8M initialization function.
+
+ Args:
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
+ [batch_size x num_frames x num_features].
+ params: model configuration parameters.
+ norm_activation: Model normalization and activation configs.
+ l2_regularizer: An optional kernel weight regularizer.
+ **kwargs: keyword arguments to be passed.
+ """
+ super().__init__(**kwargs)
+ self._input_specs = input_specs
+ self._params = params
+ self._norm_activation = norm_activation
+ self._l2_regularizer = l2_regularizer
+ self._act_fn = tf_utils.get_activation(self._norm_activation.activation)
+ self._norm = functools.partial(
+ layers.BatchNormalization,
+ momentum=self._norm_activation.norm_momentum,
+ epsilon=self._norm_activation.norm_epsilon,
+ synchronized=self._norm_activation.use_sync_bn,
+ )
+ feature_size = input_specs.shape[-1]
+
+ # Configure model batch norm layer.
+ if self._params.add_batch_norm:
+ self._input_bn = self._norm(name="input_bn")
+ self._cluster_bn = self._norm(name="cluster_bn")
+ self._hidden_bn = self._norm(name="hidden_bn")
+ else:
+ self._hidden_biases = self.add_weight(
+ name="hidden_biases",
+ shape=[self._params.hidden_size],
+ initializer=tf.random_normal_initializer(stddev=0.01),
+ )
+ self._cluster_biases = self.add_weight(
+ name="cluster_biases",
+ shape=[self._params.cluster_size],
+ initializer=tf.random_normal_initializer(
+ stddev=1.0 / tf.math.sqrt(feature_size)
+ ),
+ )
+
+ if self._params.use_context_gate_cluster_layer:
+ self._context_gate = nn_layers.ContextGate(
+ normalizer_fn=self._norm,
+ pooling_method=None,
+ hidden_layer_size=self._params.context_gate_cluster_bottleneck_size,
+ kernel_regularizer=self._l2_regularizer,
+ name="context_gate_cluster",
+ )
+
+ self._hidden_dense = layers.Dense(
+ self._params.hidden_size,
+ kernel_regularizer=self._l2_regularizer,
+ kernel_initializer=tf.random_normal_initializer(
+ stddev=1.0 / tf.sqrt(tf.cast(self._params.cluster_size, tf.float32))
+ ),
+ name="hidden_dense",
+ )
+
+ if self._params.cluster_size > 0:
+ self._cluster_dense = layers.Dense(
+ self._params.cluster_size,
+ kernel_regularizer=self._l2_regularizer,
+ kernel_initializer=tf.random_normal_initializer(
+ stddev=1.0 / tf.sqrt(tf.cast(feature_size, tf.float32))
+ ),
+ name="cluster_dense",
+ )
+
+ def call(
+ self, inputs: tf.Tensor, num_frames: Any = None,
+ ) -> tf.Tensor:
+ # L2 normalize input features
+ activation = tf.nn.l2_normalize(inputs, -1)
+
+ if self._params.add_batch_norm:
+ activation = self._input_bn(activation)
+
+ if self._params.cluster_size > 0:
+ activation = self._cluster_dense(activation)
+ if self._params.add_batch_norm:
+ activation = self._cluster_bn(activation)
+ if not self._params.add_batch_norm:
+ activation += self._cluster_biases
+
+ activation = self._act_fn(activation)
+
+ if self._params.use_context_gate_cluster_layer:
+ activation = self._context_gate(activation)
+
+ activation = yt8m_model_utils.frame_pooling(
+ activation,
+ method=self._params.pooling_method,
+ num_frames=num_frames,
+ )
+
+ activation = self._hidden_dense(activation)
+ if self._params.add_batch_norm:
+ activation = self._hidden_bn(activation)
+ else:
+ activation += self._hidden_biases
+
+ activation = self._act_fn(activation)
+ return activation
+
+
+@factory.register_backbone_builder("dbof")
+def build_dbof(
+ input_specs: tf_keras.layers.InputSpec,
+ backbone_config: hyperparams.Config,
+ norm_activation_config: hyperparams.Config,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs,
+) -> tf_keras.Model:
+ """Builds a dbof backbone from a config."""
+ backbone_type = backbone_config.type
+ backbone_cfg = backbone_config.get()
+ assert backbone_type == "dbof", f"Inconsistent backbone type {backbone_type}"
+
+ dbof = Dbof(
+ input_specs=input_specs,
+ params=backbone_cfg,
+ norm_activation=norm_activation_config,
+ l2_regularizer=l2_regularizer,
+ **kwargs,
+ )
+
+ # Warmup calls to build model variables.
+ dbof(tf_keras.Input(input_specs.shape[1:]))
+ return dbof # pyrefly: ignore[bad-return]
diff --git a/official/projects/yt8m/modeling/backbones/dbof_test.py b/official/projects/yt8m/modeling/backbones/dbof_test.py
new file mode 100644
index 00000000000..9d43f3464cf
--- /dev/null
+++ b/official/projects/yt8m/modeling/backbones/dbof_test.py
@@ -0,0 +1,58 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for dbof."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.yt8m.configs import yt8m as yt8m_cfg
+from official.projects.yt8m.modeling.backbones import dbof
+
+
+class DbofTest(parameterized.TestCase, tf.test.TestCase):
+ """Class for testing nn_layers."""
+
+ @parameterized.product(
+ pooling_method=["average", "max", "swap"],
+ use_context_gate_cluster_layer=[True, False],
+ context_gate_cluster_bottleneck_size=[0, 8],
+ )
+ def test_dbof_backbone(
+ self,
+ pooling_method,
+ use_context_gate_cluster_layer,
+ context_gate_cluster_bottleneck_size,
+ ):
+ """Test for creation of a context gate layer."""
+
+ model_cfg = yt8m_cfg.DbofModel(
+ cluster_size=30,
+ hidden_size=20,
+ pooling_method=pooling_method,
+ use_context_gate_cluster_layer=use_context_gate_cluster_layer,
+ context_gate_cluster_bottleneck_size=context_gate_cluster_bottleneck_size,
+ )
+ backbone = dbof.Dbof(
+ input_specs=tf_keras.layers.InputSpec(shape=[None, None, 32]),
+ params=model_cfg,
+ )
+
+ inputs = tf.ones([2, 24, 32], dtype=tf.float32)
+ outputs = backbone(inputs, num_frames=tf.constant([24, 16]))
+ self.assertAllEqual(outputs.shape.as_list(), [2, 20])
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/yt8m/modeling/heads/__init__.py b/official/projects/yt8m/modeling/heads/__init__.py
new file mode 100644
index 00000000000..96e03e4801e
--- /dev/null
+++ b/official/projects/yt8m/modeling/heads/__init__.py
@@ -0,0 +1,18 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Aggregation heads package definition."""
+
+from official.projects.yt8m.modeling.heads.logistic import LogisticModel
+from official.projects.yt8m.modeling.heads.moe import MoeModel
diff --git a/official/projects/yt8m/modeling/heads/logistic.py b/official/projects/yt8m/modeling/heads/logistic.py
new file mode 100644
index 00000000000..caace8c0d8a
--- /dev/null
+++ b/official/projects/yt8m/modeling/heads/logistic.py
@@ -0,0 +1,66 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Logistic model definitions."""
+
+from typing import Optional
+
+import tensorflow as tf, tf_keras
+
+
+layers = tf_keras.layers
+
+
+class LogisticModel(layers.Layer):
+ """Logistic prediction head model with L2 regularization."""
+
+ def __init__(
+ self,
+ vocab_size: int = 3862,
+ return_logits: bool = False,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs,
+ ):
+ """Creates a logistic model.
+
+ Args:
+ vocab_size: The number of classes in the dataset.
+ return_logits: if True also return logits.
+ l2_regularizer: An optional L2 weight regularizer.
+ **kwargs: extra key word args.
+ """
+ super().__init__(**kwargs)
+ self._return_logits = return_logits
+ self._dense = layers.Dense(vocab_size, kernel_regularizer=l2_regularizer)
+
+ def call(
+ self,
+ inputs: tf.Tensor,
+ ):
+ """Logistic model forward call.
+
+ Args:
+ inputs: 'batch' x 'num_features' matrix of input features.
+
+ Returns:
+ A dictionary with a tensor containing the probability predictions of the
+ model in the 'predictions' key. The dimensions of the tensor are
+ batch_size x num_classes.
+ """
+
+ logits = self._dense(inputs)
+ outputs = {"predictions": tf.nn.sigmoid(logits)}
+ if self._return_logits:
+ outputs.update({"logits": logits})
+ return outputs
diff --git a/official/projects/yt8m/modeling/heads/moe.py b/official/projects/yt8m/modeling/heads/moe.py
new file mode 100644
index 00000000000..57841e66342
--- /dev/null
+++ b/official/projects/yt8m/modeling/heads/moe.py
@@ -0,0 +1,152 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""MoE model definitions."""
+
+
+from typing import Any, Optional
+
+import tensorflow as tf, tf_keras
+
+from official.projects.yt8m.modeling import nn_layers
+
+
+layers = tf_keras.layers
+
+
+class MoeModel(layers.Layer):
+ """A softmax over a mixture of logistic models (with L2 regularization)."""
+
+ def __init__(
+ self,
+ vocab_size: int = 3862,
+ num_mixtures: int = 2,
+ use_input_context_gate: bool = False,
+ use_output_context_gate: bool = False,
+ normalizer_params: Optional[dict[str, Any]] = None,
+ vocab_as_last_dim: bool = False,
+ l2_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ **kwargs,
+ ):
+ """Creates a Mixture of (Logistic) Experts model.
+
+ The model consists of a per-class softmax distribution over a
+ configurable number of logistic classifiers. One of the classifiers
+ in the mixture is not trained, and always predicts 0.
+
+ Args:
+ vocab_size: The number of classes in the dataset.
+ num_mixtures: The number of mixtures (excluding a dummy 'expert' that
+ always predicts the non-existence of an entity).
+ use_input_context_gate: if True apply context gate layer to the input.
+ use_output_context_gate: if True apply context gate layer to the output.
+ normalizer_params: parameters of the batch normalization.
+ vocab_as_last_dim: if True reshape `activations` and make `vocab_size` as
+ the last dimension to avoid small `num_mixtures` as the last dimension.
+ XLA pads up the dimensions of tensors: typically the last dimension will
+ be padded to 128, and the second to last will be padded to 8.
+ l2_regularizer: An optional L2 weight regularizer.
+ **kwargs: extra key word args.
+ """
+ super().__init__(**kwargs)
+ self._vocab_size = vocab_size
+ self._num_mixtures = num_mixtures
+ self._use_input_context_gate = use_input_context_gate
+ self._use_output_context_gate = use_output_context_gate
+ self._vocab_as_last_dim = vocab_as_last_dim
+ self._normalizer_params = normalizer_params
+ self._l2_regularizer = l2_regularizer
+
+ if use_input_context_gate:
+ self._input_context_gate = nn_layers.ContextGate(
+ normalizer_fn=layers.BatchNormalization,
+ normalizer_params=normalizer_params,
+ name="input_context_gate",
+ )
+ if use_output_context_gate:
+ self._output_context_gate = nn_layers.ContextGate(
+ normalizer_fn=layers.BatchNormalization,
+ normalizer_params=normalizer_params,
+ name="output_context_gate",
+ )
+
+ self._gate_dense = layers.Dense(
+ vocab_size * (num_mixtures + 1),
+ activation=None,
+ bias_initializer=None,
+ kernel_regularizer=l2_regularizer,
+ name="gate",
+ )
+
+ self._expert_dense = layers.Dense(
+ vocab_size * num_mixtures,
+ activation=None,
+ kernel_regularizer=l2_regularizer,
+ name="expert",
+ )
+
+ def call(self, inputs: tf.Tensor) -> dict[str, tf.Tensor]:
+ """MoE forward call.
+
+ Args:
+ inputs: 'batch_size' x 'num_features' matrix of input features.
+
+ Returns:
+ A dictionary with a tensor containing the probability predictions
+ of the model in the 'predictions' key. The dimensions of the tensor
+ are batch_size x num_classes.
+ """
+
+ if self._use_input_context_gate:
+ inputs = self._input_context_gate(inputs)
+
+ gate_activations = self._gate_dense(inputs)
+ expert_activations = self._expert_dense(inputs)
+
+ if self._vocab_as_last_dim:
+ # Batch x (num_mixtures + 1) x #Labels
+ gate_activations = tf.reshape(
+ gate_activations, [-1, self._num_mixtures + 1, self._vocab_size]
+ )
+ # Batch x num_mixtures x #Labels
+ expert_activations = tf.reshape(
+ expert_activations,
+ [-1, self._num_mixtures, self._vocab_size],
+ )
+ else:
+ # (Batch * #Labels) x (num_mixtures + 1)
+ gate_activations = tf.reshape(
+ gate_activations,
+ [-1, self._num_mixtures + 1],
+ )
+ # (Batch * #Labels) x num_mixtures
+ expert_activations = tf.reshape(
+ expert_activations,
+ [-1, self._num_mixtures],
+ )
+
+ gating_distribution = tf.nn.softmax(gate_activations, axis=1)
+ expert_distribution = tf.nn.sigmoid(expert_activations)
+ final_probabilities = tf.reduce_sum(
+ gating_distribution[:, : self._num_mixtures] * expert_distribution,
+ axis=1,
+ )
+
+ if not self._vocab_as_last_dim:
+ final_probabilities = tf.reshape(
+ final_probabilities,
+ [-1, self._vocab_size],
+ )
+
+ return {"predictions": final_probabilities}
diff --git a/official/projects/yt8m/modeling/nn_layers.py b/official/projects/yt8m/modeling/nn_layers.py
new file mode 100644
index 00000000000..7436e57cec4
--- /dev/null
+++ b/official/projects/yt8m/modeling/nn_layers.py
@@ -0,0 +1,162 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Contains a collection of util functions for model construction."""
+
+from typing import Any, Dict, Optional, Union
+
+import tensorflow as tf, tf_keras
+
+from official.projects.yt8m.modeling import yt8m_model_utils
+
+
+class ContextGate(tf_keras.layers.Layer):
+ """Context Gating. More details: https://arxiv.org/pdf/1706.06905.pdf."""
+
+ def __init__(
+ self,
+ normalizer_fn=None,
+ normalizer_params: Optional[Dict[str, Any]] = None,
+ kernel_initializer: Union[
+ str, tf_keras.regularizers.Regularizer
+ ] = "glorot_uniform",
+ kernel_regularizer: Optional[tf_keras.regularizers.Regularizer] = None,
+ bias_initializer: Union[str, tf_keras.regularizers.Regularizer] = "zeros",
+ hidden_layer_size: int = 0,
+ pooling_method: Optional[str] = None,
+ additive_residual: bool = False,
+ name: Optional[str] = None,
+ ):
+ """Initialization of context gate.
+
+ Args:
+ normalizer_fn: Normalization function to use instead of `biases` (e.g.
+ tf.contrib.layers.batch_norm). If None, bias is added.
+ normalizer_params: Normalization function parameters.
+ kernel_initializer: Weight initializer to use instead of Xavier (e.g.
+ tf.contrib.layers.variance_scaling_initializer).
+ kernel_regularizer: Weight regularizer to use instead of None (e.g.,
+ tf.contrib.layers.l2_regularizer(l2_penalty)).
+ bias_initializer: Biases initializer to use (default tf.zeros_initializer)
+ hidden_layer_size: Dimensionality of the context gating hidden layer size,
+ if any. If None, will apply a fully-connected context gating layer with
+ shape [input_size x input_size]. If set to an int N, will factorize the
+ context gating layer into [input_size x N] x [N x input_size] as in the
+ squeeze-and-excitation block from https://arxiv.org/pdf/1709.01507.pdf.
+ pooling_method: Whether to perform global pooling of the local features
+ before applying the context gating layer. This is relevant only if the
+ input_features tensor has rank > 2, e.g., it's a sequence of frame
+ features, [batch_size, num_frames, feature_dim], or spatial convolution
+ features, [batch_size*num_frames, h, w, feature_dim]. If the inputs are
+ a set of local features and pooling_method is not None, will pool
+ features across all but the batch_size dimension using the specified
+ pooling method, and pass the aggregated features as context to the
+ gating layer. For a list of pooling methods, see the frame_pooling()
+ function.
+ additive_residual: If true, will use ReLu6-activated (additive) residual
+ connections instead of Sigmoid-activated (multiplicative) connections
+ when combining the input_features with the context gating branch.
+ name: Optional `str` name of the module.
+
+ Returns:
+ A tensor with the same shape as input_features.
+ """
+ super().__init__(name=name)
+ self._normalizer_fn = normalizer_fn
+ self._normalizer_params = normalizer_params or {}
+ self._kernel_initializer = kernel_initializer
+ self._kernel_regularizer = kernel_regularizer
+ self._bias_initializer = bias_initializer
+ self._hidden_layer_size = hidden_layer_size
+ self._pooling_method = pooling_method
+ self._additive_residual = additive_residual
+
+ if hidden_layer_size >= 2:
+ self._gates_bottleneck = tf_keras.layers.Dense(
+ hidden_layer_size,
+ activation="relu6",
+ kernel_initializer=kernel_initializer,
+ bias_initializer=bias_initializer,
+ kernel_regularizer=kernel_regularizer,
+ name="bottleneck",
+ )
+ if self._normalizer_fn:
+ self._gates_bottleneck_norm = self._normalizer_fn(
+ **self._normalizer_params,
+ name="bottleneck_norm",
+ )
+
+ def build(self, input_shape):
+ super().build(input_shape)
+ feature_size = input_shape[-1]
+ activation_fn = tf.nn.relu6 if self._additive_residual else tf.nn.sigmoid
+ self._gates = tf_keras.layers.Dense(
+ feature_size,
+ activation=activation_fn,
+ kernel_initializer=self._kernel_initializer,
+ bias_initializer=self._bias_initializer,
+ kernel_regularizer=self._kernel_regularizer,
+ name="gates_dense",
+ )
+ if self._normalizer_fn:
+ self._gates_norm = self._normalizer_fn(
+ **self._normalizer_params,
+ name="gates_norm",
+ )
+
+ def call(self, inputs: tf.Tensor):
+ num_dimensions = len(inputs.shape.as_list())
+ feature_size = inputs.shape.as_list()[-1]
+
+ if self._pooling_method:
+ assert num_dimensions > 2
+ # Collapse the inner axes of the original features shape into a 3D tensor
+ original_shape = tf.shape(inputs)
+ # The last dimension will change after concatenating the context
+ new_shape = tf.concat(
+ [original_shape[:-1], tf.constant([2 * feature_size])], 0
+ )
+ batch_size = original_shape[0]
+ reshaped_features = tf.reshape(inputs, [batch_size, -1, feature_size])
+ num_features = tf.shape(reshaped_features)[1]
+ # Pool the feature channels across the inner axes to get global context
+ context_features = yt8m_model_utils.frame_pooling(
+ reshaped_features, self._pooling_method
+ )
+ context_features = tf.expand_dims(context_features, 1)
+ # Replicate the global context features and concat to the local features.
+ context_features = tf.tile(context_features, [1, num_features, 1])
+ context_features = tf.concat([reshaped_features, context_features], 2)
+ context_features = tf.reshape(context_features, shape=new_shape)
+ else:
+ # num_dimensions should be 2
+ context_features = tf.identity(inputs)
+
+ if self._hidden_layer_size >= 2:
+ gates_bottleneck = self._gates_bottleneck(context_features)
+ if self._normalizer_fn:
+ gates_bottleneck = self._gates_bottleneck_norm(gates_bottleneck)
+ else:
+ gates_bottleneck = tf.identity(context_features)
+
+ gates = self._gates(gates_bottleneck)
+ if self._normalizer_fn:
+ gates = self._gates_norm(gates)
+
+ if self._additive_residual:
+ inputs += tf.cast(gates, inputs.dtype)
+ else:
+ inputs *= tf.cast(gates, inputs.dtype)
+
+ return inputs
diff --git a/official/projects/yt8m/modeling/nn_layers_test.py b/official/projects/yt8m/modeling/nn_layers_test.py
new file mode 100644
index 00000000000..812bb7207c6
--- /dev/null
+++ b/official/projects/yt8m/modeling/nn_layers_test.py
@@ -0,0 +1,60 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for nn_layers."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.yt8m.modeling import nn_layers
+
+
+class NNLayersTest(parameterized.TestCase, tf.test.TestCase):
+ """Class for testing nn_layers."""
+
+ @parameterized.product(
+ hidden_layer_size=(0, 8, 16),
+ additive_residual=(True, False),
+ pooling_method=["average", "max", "swap", "none", None],
+ )
+ def test_context_gate(
+ self, hidden_layer_size, additive_residual, pooling_method
+ ):
+ """Test for creation of a context gate layer."""
+
+ context_gate = nn_layers.ContextGate(
+ normalizer_fn=tf_keras.layers.BatchNormalization,
+ hidden_layer_size=hidden_layer_size,
+ additive_residual=additive_residual,
+ pooling_method=pooling_method,
+ )
+
+ if pooling_method is None:
+ inputs = tf.ones([2, 32], dtype=tf.float32)
+ elif pooling_method == "none":
+ inputs = tf.ones([2, 1, 32], dtype=tf.float32)
+ else:
+ inputs = tf.ones([2, 24, 32], dtype=tf.float32)
+
+ outputs = context_gate(inputs)
+ self.assertShapeEqual(inputs, outputs)
+
+ context_vars_len = 12 if hidden_layer_size else 6
+ context_trainable_vars_len = 8 if hidden_layer_size else 4
+ self.assertLen(context_gate.variables, context_vars_len)
+ self.assertLen(context_gate.trainable_variables, context_trainable_vars_len)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/yt8m/modeling/yt8m_agg_models.py b/official/projects/yt8m/modeling/yt8m_agg_models.py
deleted file mode 100644
index 67638db3a8c..00000000000
--- a/official/projects/yt8m/modeling/yt8m_agg_models.py
+++ /dev/null
@@ -1,119 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Contains model definitions."""
-from typing import Any, Dict, Optional
-
-import tensorflow as tf
-from official.projects.yt8m.modeling import yt8m_model_utils as utils
-
-layers = tf.keras.layers
-
-
-class LogisticModel():
- """Logistic model with L2 regularization."""
-
- def create_model(self, model_input, vocab_size, l2_penalty=1e-8):
- """Creates a logistic model.
-
- Args:
- model_input: 'batch' x 'num_features' matrix of input features.
- vocab_size: The number of classes in the dataset.
- l2_penalty: L2 weight regularization ratio.
-
- Returns:
- A dictionary with a tensor containing the probability predictions of the
- model in the 'predictions' key. The dimensions of the tensor are
- batch_size x num_classes.
- """
- output = layers.Dense(
- vocab_size,
- activation=tf.nn.sigmoid,
- kernel_regularizer=tf.keras.regularizers.l2(l2_penalty))(
- model_input)
- return {"predictions": output}
-
-
-class MoeModel():
- """A softmax over a mixture of logistic models (with L2 regularization)."""
-
- def create_model(self,
- model_input,
- vocab_size,
- num_mixtures: int = 2,
- use_input_context_gate: bool = False,
- use_output_context_gate: bool = False,
- normalizer_fn=None,
- normalizer_params: Optional[Dict[str, Any]] = None,
- l2_penalty: float = 1e-5):
- """Creates a Mixture of (Logistic) Experts model.
-
- The model consists of a per-class softmax distribution over a
- configurable number of logistic classifiers. One of the classifiers
- in the mixture is not trained, and always predicts 0.
- Args:
- model_input: 'batch_size' x 'num_features' matrix of input features.
- vocab_size: The number of classes in the dataset.
- num_mixtures: The number of mixtures (excluding a dummy 'expert' that
- always predicts the non-existence of an entity).
- use_input_context_gate: if True apply context gate layer to the input.
- use_output_context_gate: if True apply context gate layer to the output.
- normalizer_fn: normalization op constructor (e.g. batch norm).
- normalizer_params: parameters to the `normalizer_fn`.
- l2_penalty: How much to penalize the squared magnitudes of parameter
- values.
-
- Returns:
- A dictionary with a tensor containing the probability predictions
- of the model in the 'predictions' key. The dimensions of the tensor
- are batch_size x num_classes.
- """
- if use_input_context_gate:
- model_input = utils.context_gate(
- model_input,
- normalizer_fn=normalizer_fn,
- normalizer_params=normalizer_params,
- )
-
- gate_activations = layers.Dense(
- vocab_size * (num_mixtures + 1),
- activation=None,
- bias_initializer=None,
- kernel_regularizer=tf.keras.regularizers.l2(l2_penalty))(
- model_input)
- expert_activations = layers.Dense(
- vocab_size * num_mixtures,
- activation=None,
- kernel_regularizer=tf.keras.regularizers.l2(l2_penalty))(
- model_input)
-
- gating_distribution = tf.nn.softmax(
- tf.reshape(
- gate_activations,
- [-1, num_mixtures + 1])) # (Batch * #Labels) x (num_mixtures + 1)
- expert_distribution = tf.nn.sigmoid(
- tf.reshape(expert_activations,
- [-1, num_mixtures])) # (Batch * #Labels) x num_mixtures
-
- final_probabilities_by_class_and_batch = tf.reduce_sum(
- gating_distribution[:, :num_mixtures] * expert_distribution, 1)
- final_probabilities = tf.reshape(final_probabilities_by_class_and_batch,
- [-1, vocab_size])
- if use_output_context_gate:
- final_probabilities = utils.context_gate(
- final_probabilities,
- normalizer_fn=normalizer_fn,
- normalizer_params=normalizer_params,
- )
- return {"predictions": final_probabilities}
diff --git a/official/projects/yt8m/modeling/yt8m_model.py b/official/projects/yt8m/modeling/yt8m_model.py
index 80a7bf10919..a02d38e005e 100644
--- a/official/projects/yt8m/modeling/yt8m_model.py
+++ b/official/projects/yt8m/modeling/yt8m_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -12,23 +12,28 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-"""YT8M model definition."""
-from typing import Optional
+"""YT8M prediction model definition."""
+
+import functools
+from typing import Any, Optional
+
+from absl import logging
+import tensorflow as tf, tf_keras
-import tensorflow as tf
-from official.modeling import tf_utils
from official.projects.yt8m.configs import yt8m as yt8m_cfg
-from official.projects.yt8m.modeling import yt8m_agg_models
-from official.projects.yt8m.modeling import yt8m_model_utils as utils
+from official.projects.yt8m.modeling import backbones # pylint: disable=unused-import
+from official.projects.yt8m.modeling import heads
+from official.vision.modeling.backbones import factory
+
-layers = tf.keras.layers
+layers = tf_keras.layers
-class DbofModel(tf.keras.Model):
- """A YT8M model class builder.
+class VideoClassificationModel(tf_keras.Model):
+ """A video classification model class builder.
- Creates a Deep Bag of Frames model.
- The model projects the features for each frame into a higher dimensional
+ The model consists of a backbone (dbof) and a classification head.
+ The dbof backbone projects features for each frame into a higher dimensional
'clustering' space, pools across frames in that space, and then
uses a configurable video-level model to classify the now aggregated features.
The model will randomly sample either frames or sequences of frames during
@@ -37,161 +42,109 @@ class DbofModel(tf.keras.Model):
def __init__(
self,
- params: yt8m_cfg.DbofModel,
- num_frames=30,
- num_classes=3862,
- input_specs=layers.InputSpec(shape=[None, None, 1152]),
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- activation: str = "relu",
- use_sync_bn: bool = False,
- norm_momentum: float = 0.99,
- norm_epsilon: float = 0.001,
- **kwargs):
- """YT8M initialization function.
+ params: yt8m_cfg.VideoClassificationModel,
+ backbone: Optional[tf_keras.Model] = None,
+ num_classes: int = 3862,
+ input_specs: layers.InputSpec = layers.InputSpec(
+ shape=[None, None, 1152]
+ ),
+ l2_weight_decay: Optional[float] = None,
+ **kwargs,
+ ):
+ """YT8M video classification model initialization function.
Args:
- params: model configuration parameters
- num_frames: `int` number of frames in a single input.
+ params: Model configuration parameters.
+ backbone: Optional backbone model. Will build a backbone if None.
num_classes: `int` number of classes in dataset.
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
+ input_specs: `tf_keras.layers.InputSpec` specs of the input tensor.
[batch_size x num_frames x num_features]
- kernel_regularizer: tf.keras.regularizers.Regularizer object. Default to
- None.
- activation: A `str` of name of the activation function.
- use_sync_bn: If True, use synchronized batch normalization.
- norm_momentum: A `float` of normalization momentum for the moving average.
- norm_epsilon: A `float` added to variance to avoid dividing by zero.
+ l2_weight_decay: An optional `float` of kernel regularizer weight decay.
**kwargs: keyword arguments to be passed.
"""
-
- self._self_setattr_tracking = False
+ super().__init__()
+ self._params = params
+ self._num_classes = num_classes
+ self._input_specs = input_specs
+ self._l2_weight_decay = l2_weight_decay
self._config_dict = {
+ "params": params,
"input_specs": input_specs,
"num_classes": num_classes,
- "num_frames": num_frames,
- "params": params
+ "l2_weight_decay": l2_weight_decay,
}
- self._num_classes = num_classes
- self._input_specs = input_specs
- self._act_fn = tf_utils.get_activation(activation)
- if use_sync_bn:
- self._norm = layers.experimental.SyncBatchNormalization
- else:
- self._norm = layers.BatchNormalization
- if tf.keras.backend.image_data_format() == "channels_last":
- bn_axis = -1
- else:
- bn_axis = 1
-
- # [batch_size x num_frames x num_features]
- feature_size = input_specs.shape[-1]
- # shape 'excluding' batch_size
- model_input = tf.keras.Input(shape=self._input_specs.shape[1:])
- reshaped_input = tf.reshape(model_input, [-1, feature_size])
- tf.summary.histogram("input_hist", model_input)
-
- # configure model
- if params.add_batch_norm:
- reshaped_input = self._norm(
- axis=bn_axis,
- momentum=norm_momentum,
- epsilon=norm_epsilon,
- name="input_bn")(
- reshaped_input)
-
- # activation = reshaped input * cluster weights
- if params.cluster_size > 0:
- activation = layers.Dense(
- params.cluster_size,
- kernel_regularizer=kernel_regularizer,
- kernel_initializer=tf.random_normal_initializer(
- stddev=1 / tf.sqrt(tf.cast(feature_size, tf.float32))))(
- reshaped_input)
-
- if params.add_batch_norm:
- activation = self._norm(
- axis=bn_axis,
- momentum=norm_momentum,
- epsilon=norm_epsilon,
- name="cluster_bn")(
- activation)
- else:
- cluster_biases = tf.Variable(
- tf.random_normal_initializer(stddev=1 / tf.math.sqrt(feature_size))(
- shape=[params.cluster_size]),
- name="cluster_biases")
- tf.summary.histogram("cluster_biases", cluster_biases)
- activation += cluster_biases
-
- activation = self._act_fn(activation)
- tf.summary.histogram("cluster_output", activation)
-
- if params.use_context_gate_cluster_layer:
- pooling_method = None
- norm_args = dict(
- axis=bn_axis,
- momentum=norm_momentum,
- epsilon=norm_epsilon,
- name="context_gate_bn")
- activation = utils.context_gate(
- activation,
- normalizer_fn=self._norm,
- normalizer_params=norm_args,
- pooling_method=pooling_method,
- hidden_layer_size=params.context_gate_cluster_bottleneck_size,
- kernel_regularizer=kernel_regularizer)
- activation = tf.reshape(activation, [-1, num_frames, params.cluster_size])
- activation = utils.frame_pooling(activation, params.pooling_method)
-
- # activation = activation * hidden1_weights
- activation = layers.Dense(
- params.hidden_size,
- kernel_regularizer=kernel_regularizer,
- kernel_initializer=tf.random_normal_initializer(
- stddev=1 / tf.sqrt(tf.cast(params.cluster_size, tf.float32))))(
- activation)
-
- if params.add_batch_norm:
- activation = self._norm(
- axis=bn_axis,
- momentum=norm_momentum,
- epsilon=norm_epsilon,
- name="hidden1_bn")(
- activation)
+ if backbone is None:
+ # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
+ # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
+ # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
+ l2_regularizer = (
+ tf_keras.regularizers.l2(l2_weight_decay / 2.0)
+ if l2_weight_decay
+ else None
+ )
+ backbone = factory.build_backbone(
+ input_specs=input_specs,
+ backbone_config=params.backbone,
+ norm_activation_config=params.norm_activation,
+ l2_regularizer=l2_regularizer,
+ **kwargs,
+ )
+
+ self.backbone = backbone
+ self.build_head()
+
+ def build_head(self):
+ logging.info("Build DbofModel with %s.", self._params.head.type)
+ head_cfg = self._params.head.get()
+ if self._params.head.type == "moe":
+ normalizer_params = dict(
+ synchronized=self._params.norm_activation.use_sync_bn,
+ momentum=self._params.norm_activation.norm_momentum,
+ epsilon=self._params.norm_activation.norm_epsilon,
+ )
+ aggregation_head = functools.partial(
+ heads.MoeModel, normalizer_params=normalizer_params
+ )
+ elif self._params.head.type == "logistic":
+ aggregation_head = heads.LogisticModel
else:
- hidden1_biases = tf.Variable(
- tf.random_normal_initializer(stddev=0.01)(shape=[params.hidden_size]),
- name="hidden1_biases")
-
- tf.summary.histogram("hidden1_biases", hidden1_biases)
- activation += hidden1_biases
-
- activation = self._act_fn(activation)
- tf.summary.histogram("hidden1_output", activation)
-
- aggregated_model = getattr(yt8m_agg_models,
- params.yt8m_agg_classifier_model)
- norm_args = dict(axis=bn_axis, momentum=norm_momentum, epsilon=norm_epsilon)
- output = aggregated_model().create_model(
- model_input=activation,
+ logging.warn("Skip build head type: %s", self._params.head.type)
+ return
+
+ l2_regularizer = (
+ tf_keras.regularizers.l2(self._l2_weight_decay / 2.0)
+ if self._l2_weight_decay
+ else None
+ )
+ self.head = aggregation_head(
vocab_size=self._num_classes,
- num_mixtures=params.agg_model.num_mixtures,
- normalizer_fn=self._norm,
- normalizer_params=norm_args,
- l2_penalty=params.agg_model.l2_penalty)
-
- super().__init__(
- inputs=model_input, outputs=output.get("predictions"), **kwargs)
-
- @property
- def checkpoint_items(self):
- """Returns a dictionary of items to be additionally checkpointed."""
- return dict()
+ l2_regularizer=l2_regularizer,
+ **head_cfg.as_dict(),
+ )
def get_config(self):
return self._config_dict
@classmethod
- def from_config(cls, config):
+ def from_config(cls, config): # pyrefly: ignore[bad-override]
return cls(**config)
+
+ def call(
+ self,
+ inputs: tf.Tensor,
+ num_frames: Any = None,
+ training: Any = None,
+ ) -> dict[str, tf.Tensor]:
+ features = self.backbone(
+ inputs,
+ num_frames=num_frames,
+ training=training,
+ )
+ outputs = self.head(features, training=training)
+ return outputs
+
+ @property
+ def checkpoint_items(self) -> dict[str, Any]:
+ """Returns a dictionary of items to be additionally checkpointed."""
+ return dict(backbone=self.backbone, head=self.head)
diff --git a/official/projects/yt8m/modeling/yt8m_model_test.py b/official/projects/yt8m/modeling/yt8m_model_test.py
index 50883fa5ab1..0302ce1fcda 100644
--- a/official/projects/yt8m/modeling/yt8m_model_test.py
+++ b/official/projects/yt8m/modeling/yt8m_model_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.projects.yt8m.configs import yt8m as yt8m_cfg
from official.projects.yt8m.modeling import yt8m_model
@@ -26,33 +26,51 @@ class YT8MNetworkTest(parameterized.TestCase, tf.test.TestCase):
"""Class for testing yt8m network."""
# test_yt8m_network_creation arbitrary params
- @parameterized.parameters((32, 1152)) # 1152 = 1024 + 128
- def test_yt8m_network_creation(self, num_frames, feature_dims):
+ @parameterized.product(
+ num_sample_frames=(None, 16, 32),
+ pooling_method=('average', 'max', 'swap'),
+ )
+ def test_yt8m_network_creation(
+ self, num_sample_frames, pooling_method
+ ):
"""Test for creation of a YT8M Model.
Args:
- num_frames: number of frames.
- feature_dims: indicates total dimension size of the features.
+ num_sample_frames: indicates number of frames to sample.
+ pooling_method: str of frame pooling method.
"""
- input_specs = tf.keras.layers.InputSpec(shape=[num_frames, feature_dims])
+ num_frames = 24
+ feature_dims = 52
+ num_classes = 45
+ input_specs = tf_keras.layers.InputSpec(shape=[None, None, feature_dims])
- num_classes = 3862
- model = yt8m_model.DbofModel(
- params=yt8m_cfg.YT8MTask.model,
- num_frames=num_frames,
+ params = yt8m_cfg.YT8MTask().model
+ params.backbone.dbof.pooling_method = pooling_method
+ model = yt8m_model.VideoClassificationModel(
+ params=params,
num_classes=num_classes,
- input_specs=input_specs)
+ input_specs=input_specs,
+ )
- # batch = 2 -> arbitrary value for test
- inputs = np.random.rand(2 * num_frames, feature_dims)
- logits = model(inputs)
- self.assertAllEqual([2, num_classes], logits.numpy().shape)
+ # batch = 2 -> arbitrary value for test.
+ if num_sample_frames:
+ inputs = np.random.rand(2, num_sample_frames, feature_dims)
+ num_frames = tf.constant([num_sample_frames, num_sample_frames])
+ else:
+ # Add padding frames.
+ inputs = np.random.rand(2, num_frames + 4, feature_dims)
+ num_frames = tf.constant([num_frames, num_frames + 1])
+
+ predictions = model(inputs, num_frames=num_frames)['predictions']
+ self.assertAllEqual([2, num_classes], predictions.numpy().shape)
def test_serialize_deserialize(self):
- model = yt8m_model.DbofModel(params=yt8m_cfg.YT8MTask.model)
+ model = yt8m_model.VideoClassificationModel(
+ params=yt8m_cfg.YT8MTask().model
+ )
config = model.get_config()
- new_model = yt8m_model.DbofModel.from_config(config)
+ new_model = yt8m_model.VideoClassificationModel.from_config(config)
# If the serialization was successful,
# the new config should match the old.
diff --git a/official/projects/yt8m/modeling/yt8m_model_utils.py b/official/projects/yt8m/modeling/yt8m_model_utils.py
index d56fe44ba9e..4470fbf73ae 100644
--- a/official/projects/yt8m/modeling/yt8m_model_utils.py
+++ b/official/projects/yt8m/modeling/yt8m_model_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,68 +13,82 @@
# limitations under the License.
"""Contains a collection of util functions for model construction."""
-from typing import Any, Dict, Optional, Union
-import tensorflow as tf
+from typing import Optional
+import tensorflow as tf, tf_keras
-def sample_random_sequence(model_input, num_frames, num_samples):
- """Samples a random sequence of frames of size num_samples.
+
+def _large_compatible_negative(tensor_type):
+ """Large negative number as Tensor.
+
+ This function is necessary because the standard value for epsilon
+ in this module (-1e9) cannot be represented using `tf.float16`.
+
+ Args:
+ tensor_type: A dtype to determine the type.
+
+ Returns:
+ A large negative number.
+ """
+ if tensor_type == tf.float16:
+ return tf.float16.min
+ return -1e9
+
+
+def weighted_average_pooling(features, weights, axis):
+ """Weighted average pooling.
Args:
- model_input: tensor of shape [batch_size x max_frames x feature_size]
- num_frames: tensor of shape [batch_size x 1]
- num_samples: a scalar indicating the number of samples
+ features: a tensor of at least rank 1.
+ weights: a weight tensor whose shape is broadcast compatible with features.
+ It doesn't have to be normalized.
+ axis: the dimensions to reduce.
Returns:
- reshaped model_input in [batch_size x 'num_samples' x feature_size]
+ The reduced tensor.
"""
+ return tf.math.divide_no_nan(
+ tf.reduce_sum(weights * features, axis), # numerator.
+ tf.reduce_sum(weights, axis), # denominator.
+ )
+
+
+def frame_swap(
+ frames: tf.Tensor, frame_mask: Optional[tf.Tensor] = None
+) -> tf.Tensor:
+ """Self-weighted average pooling over all frames of a video.
- batch_size = tf.shape(model_input)[0]
- frame_index_offset = tf.tile(
- tf.expand_dims(tf.range(num_samples), 0), [batch_size, 1])
- max_start_frame_index = tf.maximum(num_frames - num_samples, 0)
- start_frame_index = tf.cast(
- tf.multiply(
- tf.random_uniform([batch_size, 1]),
- tf.cast(max_start_frame_index + 1, tf.float32)), tf.int32)
- frame_index = tf.minimum(start_frame_index + frame_index_offset,
- tf.cast(num_frames - 1, tf.int32))
- batch_index = tf.tile(
- tf.expand_dims(tf.range(batch_size), 1), [1, num_samples])
- index = tf.stack([batch_index, frame_index], 2)
- return tf.gather_nd(model_input, index)
-
-
-def sample_random_frames(model_input, num_frames, num_samples):
- """Samples a random set of frames of size num_samples.
+ It does the following operation independently for each feature:
+ x_pooled = (sum_i x_i * |x_i|) / (sum_i |x_i|).
+ Basically the weight for the feature in each frame is determined by the
+ magnitude of the feature itself.
+
+ Paper: https://research.google/pubs/pub48351/
Args:
- model_input: tensor of shape [batch_size x max_frames x feature_size]
- num_frames: tensor of shape [batch_size x 1]
- num_samples (int): a scalar indicating the number of samples
+ frames: A tensor with shape [batch_size, max_frames, feature_size].
+ frame_mask: A tensor with shape [batch_size, max_frames, 1].
Returns:
- reshaped model_input in [batch_size x 'num_samples' x feature_size]
+ A tensor with shape [batch_size, feature_size].
"""
- batch_size = tf.shape(model_input)[0]
- frame_index = tf.cast(
- tf.multiply(
- tf.random.uniform([batch_size, num_samples]),
- tf.tile(tf.cast(num_frames, tf.float32), [1, num_samples])), tf.int32)
- batch_index = tf.tile(
- tf.expand_dims(tf.range(batch_size), 1), [1, num_samples])
- index = tf.stack([batch_index, frame_index], 2)
- return tf.gather_nd(model_input, index)
-
-
-def frame_pooling(frames, method):
+ weights = tf.abs(frames)
+ if frame_mask is not None:
+ weights *= tf.cast(frame_mask, weights.dtype)
+ # We set axis to 1 to reduce the dimension corresponding to max_frames.
+ return weighted_average_pooling(frames, weights, axis=1)
+
+
+def frame_pooling(frames, method="average", num_frames=None):
"""Pools over the frames of a video.
Args:
frames: tensor of shape [batch_size, num_frames, feature_size].
method: string indicating pooling method, one of: "average", "max",
"attention", or "none".
+ num_frames: optional tensor of shape [batch_size] indicating valid number of
+ frames for each video.
Returns:
tensor of shape [batch_size, feature_size] for average, max, or
@@ -84,121 +98,40 @@ def frame_pooling(frames, method):
ValueError: if method is other than "average", "max", "attention", or
"none".
"""
+ frame_mask = None
+ if num_frames is not None:
+ max_frames = frames.shape.as_list()[1]
+ # Generate binary mask from number of frames.
+ frame_mask = tf.sequence_mask(num_frames, max_frames, frames.dtype)
+ frame_mask = tf.expand_dims(frame_mask, axis=2)
+
if method == "average":
- reduced = tf.reduce_mean(frames, 1)
+ if num_frames is None:
+ reduced = tf.reduce_mean(frames, 1)
+ else:
+ num_frames = tf.reshape(tf.cast(num_frames, frames.dtype), [-1, 1])
+ reduced = tf.reduce_sum(frames * frame_mask, 1) / num_frames
elif method == "max":
- reduced = tf.reduce_max(frames, 1)
+ if num_frames is not None:
+ frame_mask = tf.cast(frame_mask, tf.bool)
+ frames = tf.where(
+ frame_mask,
+ frames,
+ tf.ones_like(frames, dtype=frames.dtype)
+ * _large_compatible_negative(frames.dtype),
+ )
+ # Magic to avoid loss NaN when bfloat16 is enabled.
+ # See yaqs/5377152819545505792 and b/214396297 for more discussion.
+ reduced = tf.reduce_max(frames, 1) + tf.reduce_mean(frames, 1) * 0
+ elif method == "swap":
+ # Note we assume the frames are in the shape of
+ # [batch_size, num_frames, feature_size]. Otherwise this function might
+ # fail.
+ reduced = frame_swap(frames, frame_mask)
elif method == "none":
- feature_size = frames.shape_as_list()[2]
+ feature_size = frames.shape.as_list()[2]
reduced = tf.reshape(frames, [-1, feature_size])
else:
raise ValueError("Unrecognized pooling method: %s" % method)
return reduced
-
-
-def context_gate(
- input_features,
- normalizer_fn=None,
- normalizer_params: Optional[Dict[str, Any]] = None,
- kernel_initializer: Union[
- str, tf.keras.regularizers.Regularizer] = "glorot_uniform",
- kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
- bias_initializer: Union[str, tf.keras.regularizers.Regularizer] = "zeros",
- hidden_layer_size: int = 0,
- pooling_method: Optional[str] = None,
- additive_residual: bool = False):
- """Context Gating.
-
- More details: https://arxiv.org/pdf/1706.06905.pdf.
-
- Args:
- input_features: a tensor of at least rank 2.
- normalizer_fn: Normalization function to use instead of `biases` (e.g.
- tf.contrib.layers.batch_norm). If None, bias is added.
- normalizer_params: Normalization function parameters.
- kernel_initializer: Weight initializer to use instead of Xavier (e.g.
- tf.contrib.layers.variance_scaling_initializer).
- kernel_regularizer: Weight regularizer to use instead of None (e.g.,
- tf.contrib.layers.l2_regularizer(l2_penalty)).
- bias_initializer: Biases initializer to use (default tf.zeros_initializer)
- hidden_layer_size: Dimensionality of the context gating hidden layer size,
- if any. If None, will apply a fully-connected context gating layer with
- shape [input_size x input_size]. If set to an int N, will factorize the
- context gating layer into [input_size x N] x [N x input_size] as in the
- squeeze-and-excitation block from https://arxiv.org/pdf/1709.01507.pdf.
- pooling_method: Whether to perform global pooling of the local features
- before applying the context gating layer. This is relevant only if the
- input_features tensor has rank > 2, e.g., it's a sequence of frame
- features, [batch_size, num_frames, feature_dim], or spatial convolution
- features, [batch_size*num_frames, h, w, feature_dim]. If the inputs are a
- set of local features and pooling_method is not None, will pool features
- across all but the batch_size dimension using the specified pooling
- method, and pass the aggregated features as context to the gating layer.
- For a list of pooling methods, see the frame_pooling() function.
- additive_residual: If true, will use ReLu6-activated (additive) residual
- connections instead of Sigmoid-activated (multiplicative) connections when
- combining the input_features with the context gating branch.
-
- Returns:
- A tensor with the same shape as input_features.
- """
- if normalizer_params is None:
- normalizer_params = {}
- with tf.name_scope("ContextGating"):
- num_dimensions = len(input_features.shape.as_list())
- feature_size = input_features.shape.as_list()[-1]
- if pooling_method:
- assert num_dimensions > 2
- # Collapse the inner axes of the original features shape into a 3D tensor
- original_shape = tf.shape(input_features)
- # The last dimension will change after concatenating the context
- new_shape = tf.concat(
- [original_shape[:-1],
- tf.constant([2 * feature_size])], 0)
- batch_size = original_shape[0]
- reshaped_features = tf.reshape(input_features,
- [batch_size, -1, feature_size])
- num_features = tf.shape(reshaped_features)[1]
- # Pool the feature channels across the inner axes to get global context
- context_features = frame_pooling(reshaped_features, pooling_method)
- context_features = tf.expand_dims(context_features, 1)
- # Replicate the global context features and concat to the local features.
- context_features = tf.tile(context_features, [1, num_features, 1])
- context_features = tf.concat([reshaped_features, context_features], 2)
- context_features = tf.reshape(context_features, shape=new_shape)
- else:
- context_features = input_features
-
- if hidden_layer_size >= 2:
- gates_bottleneck = tf.keras.layers.Dense(
- hidden_layer_size,
- activation="relu6",
- kernel_initializer=kernel_initializer,
- bias_initializer=bias_initializer,
- kernel_regularizer=kernel_regularizer,
- )(
- context_features)
- if normalizer_fn:
- gates_bottleneck = normalizer_fn(**normalizer_params)(gates_bottleneck)
- else:
- gates_bottleneck = context_features
-
- activation_fn = (tf.nn.relu6 if additive_residual else tf.nn.sigmoid)
- gates = tf.keras.layers.Dense(
- feature_size,
- activation=activation_fn,
- kernel_initializer=kernel_initializer,
- bias_initializer=bias_initializer,
- kernel_regularizer=kernel_regularizer,
- )(
- gates_bottleneck)
- if normalizer_fn:
- gates = normalizer_fn(**normalizer_params)(gates)
-
- if additive_residual:
- input_features += gates
- else:
- input_features *= gates
-
- return input_features
diff --git a/official/projects/yt8m/modeling/yt8m_model_utils_test.py b/official/projects/yt8m/modeling/yt8m_model_utils_test.py
new file mode 100644
index 00000000000..fd24f65ee1f
--- /dev/null
+++ b/official/projects/yt8m/modeling/yt8m_model_utils_test.py
@@ -0,0 +1,56 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for YT8M modeling utilities."""
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.projects.yt8m.modeling import yt8m_model_utils
+
+
+class Yt8MModelUtilsTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.product(
+ frame_pooling=("average", "max", "swap", "none"),
+ use_frame_mask=(True, False),
+ )
+ def test_frame_pooling(self, frame_pooling, use_frame_mask):
+ frame = tf.constant([
+ [[0.0, 0.0, 0.0], [0.0, 1.0, -1.0]],
+ [[0.0, 0.0, 0.0], [0.0, 2.0, -2.0]],
+ ])
+ num_frames = tf.constant([2, 2]) if use_frame_mask else None
+ pooled_frame = yt8m_model_utils.frame_pooling(
+ frame, method=frame_pooling, num_frames=num_frames
+ )
+ if frame_pooling == "swap":
+ self.assertAllClose([[0.0, 1.0, -1.0], [0.0, 2.0, -2.0]], pooled_frame)
+ elif frame_pooling == "average":
+ self.assertAllClose([[0.0, 0.5, -0.5], [0.0, 1.0, -1.0]], pooled_frame)
+ elif frame_pooling == "max":
+ self.assertAllClose([[0.0, 1.0, 0.0], [0.0, 2.0, 0.0]], pooled_frame)
+ elif frame_pooling == "none":
+ self.assertAllClose(
+ [
+ [0.0, 0.0, 0.0],
+ [0.0, 1.0, -1.0],
+ [0.0, 0.0, 0.0],
+ [0.0, 2.0, -2.0],
+ ],
+ pooled_frame,
+ )
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/projects/yt8m/tasks/__init__.py b/official/projects/yt8m/tasks/__init__.py
index 85df31a45b8..c2093eea0da 100644
--- a/official/projects/yt8m/tasks/__init__.py
+++ b/official/projects/yt8m/tasks/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/yt8m/tasks/yt8m_task.py b/official/projects/yt8m/tasks/yt8m_task.py
index ee26017d607..53eda1d5cc1 100644
--- a/official/projects/yt8m/tasks/yt8m_task.py
+++ b/official/projects/yt8m/tasks/yt8m_task.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,18 +13,19 @@
# limitations under the License.
"""Video classification task definition."""
+from typing import Dict, List, Optional, Tuple
+
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
-from official.core import input_reader
from official.core import task_factory
from official.modeling import tf_utils
from official.projects.yt8m.configs import yt8m as yt8m_cfg
from official.projects.yt8m.dataloaders import yt8m_input
from official.projects.yt8m.eval_utils import eval_util
-from official.projects.yt8m.modeling import yt8m_model_utils as utils
-from official.projects.yt8m.modeling.yt8m_model import DbofModel
+from official.projects.yt8m.modeling import yt8m_model
+from official.core import input_reader
@task_factory.register_task_cls(yt8m_cfg.YT8MTask)
@@ -37,29 +38,60 @@ def build_model(self):
common_input_shape = [None, sum(train_cfg.feature_sizes)]
# [batch_size x num_frames x num_features]
- input_specs = tf.keras.layers.InputSpec(shape=[None] + common_input_shape)
+ input_specs = tf_keras.layers.InputSpec(shape=[None] + common_input_shape)
logging.info('Build model input %r', common_input_shape)
l2_weight_decay = self.task_config.losses.l2_weight_decay
- # Divide weight decay by 2.0 to match the implementation of tf.nn.l2_loss.
- # (https://www.tensorflow.org/api_docs/python/tf/keras/regularizers/l2)
- # (https://www.tensorflow.org/api_docs/python/tf/nn/l2_loss)
- l2_regularizer = (
- tf.keras.regularizers.l2(l2_weight_decay /
- 2.0) if l2_weight_decay else None)
# Model configuration.
model_config = self.task_config.model
- norm_activation_config = model_config.norm_activation
- model = DbofModel(
+ model = yt8m_model.VideoClassificationModel(
params=model_config,
input_specs=input_specs,
- num_frames=train_cfg.num_frames,
num_classes=train_cfg.num_classes,
- activation=norm_activation_config.activation,
- use_sync_bn=norm_activation_config.use_sync_bn,
- norm_momentum=norm_activation_config.norm_momentum,
- norm_epsilon=norm_activation_config.norm_epsilon,
- kernel_regularizer=l2_regularizer)
+ l2_weight_decay=l2_weight_decay,
+ )
+
+ # Warmup calls to build model variables.
+ _ = model(
+ inputs=tf_keras.Input(common_input_shape, dtype=tf.float32),
+ num_frames=tf_keras.Input([], dtype=tf.float32),
+ )
+
+ non_trainable_batch_norm_variables = []
+ non_trainable_extra_variables = []
+ for var in model.non_trainable_variables:
+ if 'moving_mean' in var.name or 'moving_variance' in var.name:
+ non_trainable_batch_norm_variables.append(var)
+ else:
+ non_trainable_extra_variables.append(var)
+
+ logging.info(
+ 'Trainable model variables:\n%s',
+ '\n'.join(
+ [f'{var.name}\t{var.shape}' for var in model.trainable_variables]
+ ),
+ )
+ logging.info(
+ (
+ 'Non-trainable batch norm variables (get updated in training'
+ ' mode):\n%s'
+ ),
+ '\n'.join(
+ [
+ f'{var.name}\t{var.shape}'
+ for var in non_trainable_batch_norm_variables
+ ]
+ ),
+ )
+ logging.info(
+ 'Non-trainable frozen model variables:\n%s',
+ '\n'.join(
+ [
+ f'{var.name}\t{var.shape}'
+ for var in non_trainable_extra_variables
+ ]
+ ),
+ )
return model
def build_inputs(self, params: yt8m_cfg.DataConfig, input_context=None):
@@ -89,37 +121,52 @@ def build_inputs(self, params: yt8m_cfg.DataConfig, input_context=None):
decoder_fn=decoder_fn,
parser_fn=parser_fn,
postprocess_fn=postprocess_fn,
- transform_and_batch_fn=batch_fn)
+ transform_and_batch_fn=batch_fn,
+ )
dataset = reader.read(input_context=input_context)
return dataset
- def build_losses(self, labels, model_outputs, aux_losses=None):
+ def build_losses(
+ self, labels, model_outputs, label_weights=None, aux_losses=None
+ ):
"""Sigmoid Cross Entropy.
Args:
labels: tensor containing truth labels.
- model_outputs: output logits of the classifier.
+ model_outputs: output probabilities of the classifier.
+ label_weights: optional tensor of label weights.
aux_losses: tensor containing auxiliarly loss tensors, i.e. `losses` in
keras.Model.
Returns:
- Tensors: The total loss, model loss tensors.
+ A dict of tensors contains total loss, model loss tensors.
"""
losses_config = self.task_config.losses
- model_loss = tf.keras.losses.binary_crossentropy(
- labels,
- model_outputs,
+ model_loss = tf_keras.losses.binary_crossentropy(
+ tf.expand_dims(labels, axis=-1),
+ tf.expand_dims(model_outputs, axis=-1),
from_logits=losses_config.from_logits,
- label_smoothing=losses_config.label_smoothing)
+ label_smoothing=losses_config.label_smoothing,
+ axis=-1,
+ )
+ if label_weights is None:
+ model_loss = tf_utils.safe_mean(model_loss)
+ else:
+ model_loss = model_loss * label_weights
+ # Manutally compute weighted mean loss.
+ total_loss = tf.reduce_sum(model_loss)
+ total_weight = tf.cast(
+ tf.reduce_sum(label_weights), dtype=total_loss.dtype
+ )
+ model_loss = tf.math.divide_no_nan(total_loss, total_weight)
- model_loss = tf_utils.safe_mean(model_loss)
total_loss = model_loss
if aux_losses:
total_loss += tf.add_n(aux_losses)
- return total_loss, model_loss
+ return {'total_loss': total_loss, 'model_loss': model_loss}
def build_metrics(self, training=True):
"""Gets streaming metrics for training/validation.
@@ -130,34 +177,117 @@ def build_metrics(self, training=True):
top_n: A positive Integer specifying the average precision at n, or None
to use all provided data points.
Args:
- training: bool value, true for training mode, false for eval/validation.
+ training: Bool value, true for training mode, false for eval/validation.
Returns:
- list of strings that indicate metrics to be used
+ A list of metrics to be used.
"""
metrics = []
metric_names = ['total_loss', 'model_loss']
for name in metric_names:
- metrics.append(tf.keras.metrics.Mean(name, dtype=tf.float32))
+ metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
- if not training: # Cannot run in train step.
+ if (
+ self.task_config.evaluation.average_precision is not None
+ and not training
+ ):
+ # Cannot run in train step.
num_classes = self.task_config.validation_data.num_classes
- top_k = self.task_config.top_k
- top_n = self.task_config.top_n
+ top_k = self.task_config.evaluation.average_precision.top_k
+ top_n = self.task_config.evaluation.average_precision.top_n
self.avg_prec_metric = eval_util.EvaluationMetrics(
- num_classes, top_k=top_k, top_n=top_n)
+ num_classes, top_k=top_k, top_n=top_n
+ )
return metrics
+ def process_metrics(
+ self,
+ metrics: List[tf_keras.metrics.Metric],
+ labels: tf.Tensor,
+ outputs: tf.Tensor,
+ model_losses: Optional[Dict[str, tf.Tensor]] = None,
+ label_weights: Optional[tf.Tensor] = None,
+ training: bool = True,
+ **kwargs,
+ ) -> Dict[str, Tuple[tf.Tensor, ...]]:
+ """Updates metrics.
+
+ Args:
+ metrics: Evaluation metrics to be updated.
+ labels: A tensor containing truth labels.
+ outputs: Model output logits of the classifier.
+ model_losses: An optional dict of model losses.
+ label_weights: Optional label weights, can be broadcast into shape of
+ outputs/labels.
+ training: Bool indicates if in training mode.
+ **kwargs: Additional input arguments.
+
+ Returns:
+ Updated dict of metrics log.
+ """
+ if model_losses is None:
+ model_losses = {}
+
+ logs = {}
+ if (
+ self.task_config.evaluation.average_precision is not None
+ and not training
+ ):
+ logs.update({self.avg_prec_metric.name: (labels, outputs)})
+
+ for m in metrics:
+ if m.name in model_losses:
+ m.update_state(model_losses[m.name])
+ logs[m.name] = m.result()
+ return logs
+
+ def _preprocess_model_inputs(
+ self,
+ inputs: dict[str, tf.Tensor],
+ require_num_frames: bool = True,
+ training: bool = True,
+ ):
+ """Preprocesses input tensors before model on device."""
+ extra_inputs = {
+ 'num_frames': (
+ tf.reshape(inputs['num_frames'], [-1])
+ if require_num_frames
+ else None
+ ),
+ 'training': training,
+ }
+ return inputs['video_matrix'], extra_inputs
+
+ def _preprocess_labels(
+ self, inputs: dict[str, tf.Tensor], training: bool = True
+ ):
+ """Preprocesses labels."""
+ del training # training is unused in _preprocess_labels in YT8M.
+ labels = inputs['labels']
+ label_weights = inputs.get('label_weights', None)
+
+ return labels, label_weights
+
+ def _postprocess_outputs(
+ self, outputs, labels, label_weights, training: bool = True
+ ):
+ """Postprocess model outputs (inputs / labels / label_weights)."""
+ if not training and self.task_config.validation_data.segment_labels:
+ # workaround to ignore the unrated labels.
+ outputs *= label_weights
+ # remove padding
+ outputs = outputs[~tf.reduce_all(labels == -1, axis=1)]
+ labels = labels[~tf.reduce_all(labels == -1, axis=1)]
+ return outputs, labels, label_weights
+
def train_step(self, inputs, model, optimizer, metrics=None):
"""Does forward and backward.
Args:
- inputs: a dictionary of input tensors. output_dict = {
- "video_ids": batch_video_ids,
- "video_matrix": batch_video_matrix,
- "labels": batch_labels,
- "num_frames": batch_frames, }
+ inputs: a dictionary of input tensors. output_dict = { "video_ids":
+ batch_video_ids, "video_matrix": batch_video_matrix, "labels":
+ batch_labels, "num_frames": batch_frames, }
model: the model, forward pass definition.
optimizer: the optimizer for this training step.
metrics: a nested structure of metrics objects.
@@ -165,133 +295,142 @@ def train_step(self, inputs, model, optimizer, metrics=None):
Returns:
a dictionary of logs.
"""
- features, labels = inputs['video_matrix'], inputs['labels']
- num_frames = inputs['num_frames']
-
- # Normalize input features.
- feature_dim = len(features.shape) - 1
- features = tf.nn.l2_normalize(features, feature_dim)
-
- # sample random frames / random sequence
- num_frames = tf.cast(num_frames, tf.float32)
- sample_frames = self.task_config.train_data.num_frames
- if self.task_config.model.sample_random_frames:
- features = utils.sample_random_frames(features, num_frames, sample_frames)
- else:
- features = utils.sample_random_sequence(features, num_frames,
- sample_frames)
+ # Will require `num_frames` if `num_sample_frames` is None since
+ # video_matrix is padded to max_frames in this case.
+ require_num_frames = self.task_config.train_data.num_sample_frames is None
+ inputs_tensor, extra_inputs = self._preprocess_model_inputs(
+ inputs,
+ require_num_frames=require_num_frames,
+ training=True,
+ )
+ labels, label_weights = self._preprocess_labels(inputs, training=True)
num_replicas = tf.distribute.get_strategy().num_replicas_in_sync
with tf.GradientTape() as tape:
- outputs = model(features, training=True)
+ outputs = model(inputs_tensor, **extra_inputs)['predictions']
# Casting output layer as float32 is necessary when mixed_precision is
# mixed_float16 or mixed_bfloat16 to ensure output is casted as float32.
outputs = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), outputs)
+ # Post-process model / label outputs.
+ outputs, labels, label_weights = self._postprocess_outputs(
+ outputs, labels, label_weights, training=True
+ )
# Computes per-replica loss
- loss, model_loss = self.build_losses(
- model_outputs=outputs, labels=labels, aux_losses=model.losses)
+ all_losses = self.build_losses(
+ model_outputs=outputs,
+ labels=labels,
+ label_weights=label_weights,
+ aux_losses=model.losses,
+ )
+
+ loss = all_losses['total_loss']
# Scales loss as the default gradients allreduce performs sum inside the
# optimizer.
scaled_loss = loss / num_replicas
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
# scaled for numerical stability.
- if isinstance(optimizer,
- tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
tvars = model.trainable_variables
grads = tape.gradient(scaled_loss, tvars)
# Scales back gradient before apply_gradients when LossScaleOptimizer is
# used.
- if isinstance(optimizer,
- tf.keras.mixed_precision.LossScaleOptimizer):
+ if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
grads = optimizer.get_unscaled_gradients(grads)
# Apply gradient clipping.
if self.task_config.gradient_clip_norm > 0:
- grads, _ = tf.clip_by_global_norm(grads,
- self.task_config.gradient_clip_norm)
+ grads, _ = tf.clip_by_global_norm(
+ grads, self.task_config.gradient_clip_norm
+ )
optimizer.apply_gradients(list(zip(grads, tvars)))
logs = {self.loss: loss}
-
- all_losses = {'total_loss': loss, 'model_loss': model_loss}
-
- if metrics:
- for m in metrics:
- m.update_state(all_losses[m.name])
- logs.update({m.name: m.result()})
-
+ logs.update(
+ self.process_metrics(
+ metrics, # pyrefly: ignore[bad-argument-type]
+ labels=labels,
+ outputs=outputs,
+ model_losses=all_losses,
+ label_weights=label_weights,
+ training=True,
+ )
+ )
return logs
def validation_step(self, inputs, model, metrics=None):
"""Validatation step.
Args:
- inputs: a dictionary of input tensors. output_dict = {
- "video_ids": batch_video_ids,
- "video_matrix": batch_video_matrix,
- "labels": batch_labels,
- "num_frames": batch_frames, }
- model: the model, forward definition
+ inputs: a dictionary of input tensors. output_dict = { "video_ids":
+ batch_video_ids, "video_matrix": batch_video_matrix, "labels":
+ batch_labels, "num_frames": batch_frames}.
+ model: the model, forward definition.
metrics: a nested structure of metrics objects.
Returns:
a dictionary of logs.
"""
- features, labels = inputs['video_matrix'], inputs['labels']
- num_frames = inputs['num_frames']
-
- # Normalize input features.
- feature_dim = len(features.shape) - 1
- features = tf.nn.l2_normalize(features, feature_dim)
-
- # sample random frames (None, 5, 1152) -> (None, 30, 1152)
- sample_frames = self.task_config.validation_data.num_frames
- if self.task_config.model.sample_random_frames:
- features = utils.sample_random_frames(features, num_frames, sample_frames)
- else:
- features = utils.sample_random_sequence(features, num_frames,
- sample_frames)
-
- outputs = self.inference_step(features, model)
+ # Will require `num_frames` if `num_sample_frames` is None since
+ # video_matrix is padded to max_frames in this case.
+ require_num_frames = (
+ self.task_config.validation_data.num_sample_frames is None
+ )
+ outputs = self.inference_step(
+ model, inputs, require_num_frames=require_num_frames
+ )['predictions']
outputs = tf.nest.map_structure(lambda x: tf.cast(x, tf.float32), outputs)
- if self.task_config.validation_data.segment_labels:
- # workaround to ignore the unrated labels.
- outputs *= inputs['label_weights']
- # remove padding
- outputs = outputs[~tf.reduce_all(labels == -1, axis=1)]
- labels = labels[~tf.reduce_all(labels == -1, axis=1)]
- loss, model_loss = self.build_losses(
- model_outputs=outputs, labels=labels, aux_losses=model.losses)
-
- logs = {self.loss: loss}
-
- all_losses = {'total_loss': loss, 'model_loss': model_loss}
-
- logs.update({self.avg_prec_metric.name: (labels, outputs)})
+ labels, label_weights = self._preprocess_labels(inputs, training=False)
+ outputs, labels, label_weights = self._postprocess_outputs(
+ outputs, labels, label_weights, training=False
+ )
+
+ all_losses = self.build_losses(
+ labels=labels,
+ model_outputs=outputs,
+ label_weights=label_weights,
+ aux_losses=model.losses,
+ )
+
+ logs = {self.loss: all_losses['total_loss']}
+ logs.update(
+ self.process_metrics(
+ metrics, # pyrefly: ignore[bad-argument-type]
+ labels=labels,
+ outputs=outputs,
+ model_losses=all_losses,
+ label_weights=inputs.get('label_weights', None),
+ training=False,
+ )
+ )
- if metrics:
- for m in metrics:
- m.update_state(all_losses[m.name])
- logs.update({m.name: m.result()})
return logs
- def inference_step(self, inputs, model):
+ def inference_step(self, model, inputs, require_num_frames=True):
"""Performs the forward step."""
- return model(inputs, training=False)
+ model_inputs, extra_inputs = self._preprocess_model_inputs(
+ inputs, require_num_frames=require_num_frames, training=False
+ )
+ return model(model_inputs, **extra_inputs)
def aggregate_logs(self, state=None, step_logs=None):
- if state is None:
- state = self.avg_prec_metric
- self.avg_prec_metric.accumulate(
- labels=step_logs[self.avg_prec_metric.name][0],
- predictions=step_logs[self.avg_prec_metric.name][1])
+ if self.task_config.evaluation.average_precision is not None:
+ if state is None:
+ state = self.avg_prec_metric
+ self.avg_prec_metric.accumulate(
+ labels=step_logs[self.avg_prec_metric.name][0], # pyrefly: ignore[unsupported-operation]
+ predictions=step_logs[self.avg_prec_metric.name][1], # pyrefly: ignore[unsupported-operation]
+ )
return state
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
- avg_prec_metrics = self.avg_prec_metric.get()
- self.avg_prec_metric.clear()
- return avg_prec_metrics
+ if self.task_config.evaluation.average_precision is not None:
+ avg_prec_metrics = self.avg_prec_metric.get(
+ self.task_config.evaluation.average_precision.return_per_class_ap
+ )
+ self.avg_prec_metric.clear()
+ return avg_prec_metrics
+ return None
diff --git a/official/projects/yt8m/train.py b/official/projects/yt8m/train.py
index 2145a182651..a0a9af03e42 100644
--- a/official/projects/yt8m/train.py
+++ b/official/projects/yt8m/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/projects/yt8m/train_test.py b/official/projects/yt8m/train_test.py
index 258ccd4a18e..a392397d29d 100644
--- a/official/projects/yt8m/train_test.py
+++ b/official/projects/yt8m/train_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -12,49 +12,78 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-
import json
import os
from absl import flags
from absl.testing import flagsaver
-import numpy as np
-import tensorflow as tf
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
from official.projects.yt8m import train as train_lib
+from official.projects.yt8m.dataloaders import utils
from official.vision.dataloaders import tfexample_utils
FLAGS = flags.FLAGS
-def make_yt8m_example():
- rgb = np.random.randint(low=256, size=1024, dtype=np.uint8)
- audio = np.random.randint(low=256, size=128, dtype=np.uint8)
-
- seq_example = tf.train.SequenceExample()
- seq_example.context.feature['id'].bytes_list.value[:] = [b'id001']
- seq_example.context.feature['labels'].int64_list.value[:] = [1, 2, 3, 4]
- tfexample_utils.put_bytes_list_to_feature(
- seq_example, rgb.tobytes(), key='rgb', repeat_num=120)
- tfexample_utils.put_bytes_list_to_feature(
- seq_example, audio.tobytes(), key='audio', repeat_num=120)
-
- return seq_example
-
-
-class TrainTest(tf.test.TestCase):
+class TrainTest(parameterized.TestCase, tf.test.TestCase):
def setUp(self):
- super(TrainTest, self).setUp()
+ super().setUp()
self._model_dir = os.path.join(self.get_temp_dir(), 'model_dir')
tf.io.gfile.makedirs(self._model_dir)
data_dir = os.path.join(self.get_temp_dir(), 'data')
tf.io.gfile.makedirs(data_dir)
self._data_path = os.path.join(data_dir, 'data.tfrecord')
- examples = [make_yt8m_example() for _ in range(8)]
+ examples = [utils.make_yt8m_example() for _ in range(8)]
tfexample_utils.dump_to_tfrecord(self._data_path, tf_examples=examples)
- def test_run(self):
+ @parameterized.named_parameters(
+ dict(
+ testcase_name='segment_with_avg_precison',
+ use_segment_level_labels=True,
+ use_average_precision_metric=True,
+ num_sample_frames=24,
+ ),
+ dict(
+ testcase_name='video_with_avg_precison',
+ use_segment_level_labels=False,
+ use_average_precision_metric=True,
+ num_sample_frames=24,
+ ),
+ dict(
+ testcase_name='segment',
+ use_segment_level_labels=True,
+ use_average_precision_metric=False,
+ num_sample_frames=24,
+ ),
+ dict(
+ testcase_name='video',
+ use_segment_level_labels=False,
+ use_average_precision_metric=False,
+ num_sample_frames=24,
+ ),
+ dict(
+ testcase_name='segment_without_sampling_frames',
+ use_segment_level_labels=True,
+ use_average_precision_metric=False,
+ num_sample_frames=None,
+ ),
+ dict(
+ testcase_name='video_without_sampling_frames',
+ use_segment_level_labels=False,
+ use_average_precision_metric=False,
+ num_sample_frames=None,
+ ),
+ )
+ def test_train_and_eval(
+ self,
+ use_segment_level_labels,
+ use_average_precision_metric,
+ num_sample_frames,
+ ):
saved_flag_values = flagsaver.save_flag_values()
train_lib.tfm_flags.define_flags()
FLAGS.mode = 'train'
@@ -62,41 +91,56 @@ def test_run(self):
FLAGS.experiment = 'yt8m_experiment'
FLAGS.tpu = ''
+ average_precision = {'top_k': 20} if use_average_precision_metric else None
params_override = json.dumps({
'runtime': {
'distribution_strategy': 'mirrored',
'mixed_precision_dtype': 'float32',
},
'trainer': {
- 'train_steps': 1,
- 'validation_steps': 1,
+ 'train_steps': 2,
+ 'validation_steps': 2,
},
'task': {
'model': {
- 'cluster_size': 16,
- 'hidden_size': 16,
- 'use_context_gate_cluster_layer': True,
- 'agg_model': {
- 'use_input_context_gate': True,
- 'use_output_context_gate': True,
+ 'backbone': {
+ 'type': 'dbof',
+ 'dbof': {
+ 'cluster_size': 16,
+ 'hidden_size': 16,
+ 'use_context_gate_cluster_layer': True,
+ },
+ },
+ 'head': {
+ 'type': 'moe',
+ 'moe': {
+ 'use_input_context_gate': True,
+ 'use_output_context_gate': True,
+ },
},
},
'train_data': {
'input_path': self._data_path,
'global_batch_size': 4,
+ 'num_sample_frames': num_sample_frames,
},
'validation_data': {
'input_path': self._data_path,
+ 'segment_labels': use_segment_level_labels,
'global_batch_size': 4,
- }
- }
+ 'num_sample_frames': num_sample_frames,
+ },
+ 'evaluation': {
+ 'average_precision': average_precision,
+ },
+ },
})
FLAGS.params_override = params_override
- train_lib.train.main('unused_args')
+ with train_lib.train.gin.unlock_config():
+ train_lib.train.main('unused_args')
FLAGS.mode = 'eval'
-
with train_lib.train.gin.unlock_config():
train_lib.train.main('unused_args')
diff --git a/official/recommendation/README.md b/official/recommendation/README.md
index 59c73b5a8a4..2a1e51fd10b 100644
--- a/official/recommendation/README.md
+++ b/official/recommendation/README.md
@@ -1,10 +1,10 @@
# Recommendation Model
## Overview
-This is an implementation of the Neural Collaborative Filtering (NCF) framework with Neural Matrix Factorization (NeuMF) model as described in the [Neural Collaborative Filtering](https://arxiv.org/abs/1708.05031) paper. Current implementation is based on the code from the authors' [NCF code](https://github.com/hexiangnan/neural_collaborative_filtering) and the Stanford implementation in the [MLPerf Repo](https://github.com/mlperf/reference/tree/master/recommendation/pytorch).
+This is an implementation of the Neural Collaborative Filtering (NCF) framework with the Neural Matrix Factorization (NeuMF) model as described in the [Neural Collaborative Filtering](https://arxiv.org/abs/1708.05031) paper. The current implementation is based on the code from the authors' [NCF code](https://github.com/hexiangnan/neural_collaborative_filtering) and the Stanford implementation in the [MLPerf Repo](https://github.com/mlperf).
-NCF is a general framework for collaborative filtering of recommendations in which a neural network architecture is used to model user-item interactions. Unlike traditional models, NCF does not resort to Matrix Factorization (MF) with an inner product on latent features of users and items. It replaces the inner product with a multi-layer perceptron that can learn an arbitrary function from data.
+NCF is a general framework for the collaborative filtering of recommendations in which a neural network architecture is used to model user-item interactions. Unlike traditional models, NCF does not resort to Matrix Factorization (MF) with an inner product on latent features of users and items. It replaces the inner product with a multi-layer perceptron that can learn an arbitrary function from data.
-Two instantiations of NCF are Generalized Matrix Factorization (GMF) and Multi-Layer Perceptron (MLP). GMF applies a linear kernel to model the latent feature interactions, and MLP uses a nonlinear kernel to learn the interaction function from data. NeuMF is a fused model of GMF and MLP to better model the complex user-item interactions, and unifies the strengths of linearity of MF and non-linearity of MLP for modeling the user-item latent structures. NeuMF allows GMF and MLP to learn separate embeddings, and combines the two models by concatenating their last hidden layer. [neumf_model.py](neumf_model.py) defines the architecture details.
+Two instantiations of NCF are Generalized Matrix Factorization (GMF) and Multi-Layer Perceptron (MLP). GMF applies a linear kernel to model the latent feature interactions, and MLP uses a nonlinear kernel to learn the interaction function from data. NeuMF is a fused model of GMF and MLP to better model the complex user-item interactions and unifies the strengths of linearity of MF and non-linearity of MLP for modeling the user-item latent structures. NeuMF allows GMF and MLP to learn separate embeddings and combines the two models by concatenating their last hidden layer. [neumf_model.py](neumf_model.py) defines the architecture details.
Some abbreviations used the code base include:
- NCF: Neural Collaborative Filtering
@@ -20,7 +20,7 @@ Some abbreviations used the code base include:
The [MovieLens datasets](https://files.grouplens.org/datasets/movielens/) are used for model training and evaluation. Specifically, we use two datasets: **ml-1m** (short for MovieLens 1 million) and **ml-20m** (short for MovieLens 20 million).
### ml-1m
-ml-1m dataset contains 1,000,209 anonymous ratings of approximately 3,706 movies made by 6,040 users who joined MovieLens in 2000. All ratings are contained in the file "ratings.dat" without header row, and are in the following format:
+ml-1m dataset contains 1,000,209 anonymous ratings of approximately 3,706 movies made by 6,040 users who joined MovieLens in 2000. All ratings are contained in the file "ratings.dat" without a header row, and are in the following format:
```
UserID::MovieID::Rating::Timestamp
```
@@ -69,4 +69,4 @@ Arguments:
* `--dataset`: The dataset name to be downloaded and preprocessed. By default, it is `ml-1m`.
* `--num_gpus`: The number of GPUs used for training/evaluation of the model. Use CPU if this flag is 0. By default, it is 1.
-There are other arguments about models and training process. Refer to the [Flags package](https://abseil.io/docs/python/guides/flags) documentation or use the `--helpfull` flag to get a full list of possible arguments with detailed descriptions.
+There are other arguments about models and the training processes. Refer to the [Flags package](https://abseil.io/docs/python/guides/flags) documentation or use the `--helpful` flag to get a full list of possible arguments with detailed descriptions.
diff --git a/official/recommendation/__init__.py b/official/recommendation/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/recommendation/__init__.py
+++ b/official/recommendation/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/constants.py b/official/recommendation/constants.py
index bfbcf52ccce..a8130c2b471 100644
--- a/official/recommendation/constants.py
+++ b/official/recommendation/constants.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/create_ncf_data.py b/official/recommendation/create_ncf_data.py
index 013d0499740..f224e194e79 100644
--- a/official/recommendation/create_ncf_data.py
+++ b/official/recommendation/create_ncf_data.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,10 +17,9 @@
import json
# pylint: disable=g-bad-import-order
-# Import libraries
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: enable=g-bad-import-order
from official.recommendation import movielens
diff --git a/official/recommendation/data_pipeline.py b/official/recommendation/data_pipeline.py
index 78f2a892f27..a3094d71bd7 100644
--- a/official/recommendation/data_pipeline.py
+++ b/official/recommendation/data_pipeline.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -32,7 +32,7 @@
from absl import logging
import numpy as np
from six.moves import queue
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.tpu.datasets import StreamingFilesDataset
from official.recommendation import constants as rconst
@@ -317,14 +317,14 @@ def get_dataset(self, batch_size, epochs_between_evals):
}
if self._is_training:
- types[rconst.VALID_POINT_MASK] = np.bool
+ types[rconst.VALID_POINT_MASK] = bool
shapes[rconst.VALID_POINT_MASK] = tf.TensorShape([batch_size, 1])
- types = (types, np.bool)
+ types = (types, bool)
shapes = (shapes, tf.TensorShape([batch_size, 1]))
else:
- types[rconst.DUPLICATE_MASK] = np.bool
+ types[rconst.DUPLICATE_MASK] = bool
shapes[rconst.DUPLICATE_MASK] = tf.TensorShape([batch_size, 1])
data_generator = functools.partial(
@@ -637,7 +637,7 @@ def _assemble_eval_batch(users, positive_items, negative_items,
users = np.concatenate([users, padding.astype(users.dtype)], axis=0)
items = np.concatenate([items, padding.astype(items.dtype)], axis=0)
- duplicate_mask = stat_utils.mask_duplicates(items, axis=1).astype(np.bool)
+ duplicate_mask = stat_utils.mask_duplicates(items, axis=1).astype(bool)
items[:, (0, -1)] = items[:, (-1, 0)]
duplicate_mask[:, (0, -1)] = duplicate_mask[:, (-1, 0)]
diff --git a/official/recommendation/data_preprocessing.py b/official/recommendation/data_preprocessing.py
index 394935acf26..0bac4a3e9d3 100644
--- a/official/recommendation/data_preprocessing.py
+++ b/official/recommendation/data_preprocessing.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -28,7 +28,7 @@
from absl import logging
import numpy as np
import pandas as pd
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.recommendation import constants as rconst
from official.recommendation import data_pipeline
diff --git a/official/recommendation/data_test.py b/official/recommendation/data_test.py
index 841be8e5818..348575fb1a6 100644
--- a/official/recommendation/data_test.py
+++ b/official/recommendation/data_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,7 +26,7 @@
import numpy as np
import scipy.stats
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.recommendation import constants as rconst
from official.recommendation import data_preprocessing
@@ -190,7 +190,7 @@ def _test_end_to_end(self, constructor_type):
train_examples[l].add((u_raw, i_raw))
counts[(u_raw, i_raw)] += 1
- self.assertRegexpMatches(md5.hexdigest(), END_TO_END_TRAIN_MD5)
+ self.assertRegex(md5.hexdigest(), END_TO_END_TRAIN_MD5)
num_positives_seen = len(train_examples[True])
self.assertEqual(producer._train_pos_users.shape[0], num_positives_seen)
@@ -254,7 +254,7 @@ def _test_end_to_end(self, constructor_type):
# from the negatives.
assert (u_raw, i_raw) not in self.seen_pairs
- self.assertRegexpMatches(md5.hexdigest(), END_TO_END_EVAL_MD5)
+ self.assertRegex(md5.hexdigest(), END_TO_END_EVAL_MD5)
def _test_fresh_randomness(self, constructor_type):
train_epochs = 5
@@ -300,7 +300,7 @@ def _test_fresh_randomness(self, constructor_type):
else:
negative_counts[(u, i)] += 1
- self.assertRegexpMatches(md5.hexdigest(), FRESH_RANDOMNESS_MD5)
+ self.assertRegex(md5.hexdigest(), FRESH_RANDOMNESS_MD5)
# The positive examples should appear exactly once each epoch
self.assertAllEqual(
diff --git a/official/recommendation/movielens.py b/official/recommendation/movielens.py
index fb9e595176c..b1ec533af74 100644
--- a/official/recommendation/movielens.py
+++ b/official/recommendation/movielens.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -27,7 +27,6 @@
import zipfile
# pylint: disable=g-bad-import-order
-# Import libraries
import numpy as np
import pandas as pd
import six
@@ -35,7 +34,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: enable=g-bad-import-order
from official.utils.flags import core as flags_core
diff --git a/official/recommendation/ncf_common.py b/official/recommendation/ncf_common.py
index f1677bf15ad..f15b8ae96ab 100644
--- a/official/recommendation/ncf_common.py
+++ b/official/recommendation/ncf_common.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -24,7 +24,7 @@
from absl import flags
from absl import logging
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.recommendation import constants as rconst
diff --git a/official/recommendation/ncf_input_pipeline.py b/official/recommendation/ncf_input_pipeline.py
index 194a83866f4..b95c72846d8 100644
--- a/official/recommendation/ncf_input_pipeline.py
+++ b/official/recommendation/ncf_input_pipeline.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import functools
# pylint: disable=g-bad-import-order
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: enable=g-bad-import-order
from official.recommendation import constants as rconst
@@ -152,23 +152,23 @@ def create_ncf_input_data(params,
train_dataset = create_dataset_from_tf_record_files(
params["train_dataset_path"],
- input_meta_data["train_prebatch_size"],
+ input_meta_data["train_prebatch_size"], # pyrefly: ignore[unsupported-operation]
params["batch_size"],
is_training=True,
rebatch=False)
# Re-batch evaluation dataset for TPU Pods.
# TODO (b/162341937) remove once it's fixed.
- eval_rebatch = (params["use_tpu"] and strategy.num_replicas_in_sync > 8)
+ eval_rebatch = (params["use_tpu"] and strategy.num_replicas_in_sync > 8) # pyrefly: ignore[missing-attribute]
eval_dataset = create_dataset_from_tf_record_files(
params["eval_dataset_path"],
- input_meta_data["eval_prebatch_size"],
+ input_meta_data["eval_prebatch_size"], # pyrefly: ignore[unsupported-operation]
params["eval_batch_size"],
is_training=False,
rebatch=eval_rebatch)
- num_train_steps = int(input_meta_data["num_train_steps"])
- num_eval_steps = int(input_meta_data["num_eval_steps"])
+ num_train_steps = int(input_meta_data["num_train_steps"]) # pyrefly: ignore[unsupported-operation]
+ num_eval_steps = int(input_meta_data["num_eval_steps"]) # pyrefly: ignore[unsupported-operation]
else:
if params["use_tpu"]:
raise ValueError("TPU training does not support data producer yet. "
diff --git a/official/recommendation/ncf_keras_main.py b/official/recommendation/ncf_keras_main.py
index 268ce09343c..fcdf34e8f53 100644
--- a/official/recommendation/ncf_keras_main.py
+++ b/official/recommendation/ncf_keras_main.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,7 +26,7 @@
from absl import app
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
# pylint: enable=g-bad-import-order
from official.common import distribute_utils
@@ -51,7 +51,7 @@ def metric_fn(logits, dup_mask, match_mlperf):
return in_top_k, metric_weights
-class MetricLayer(tf.keras.layers.Layer):
+class MetricLayer(tf_keras.layers.Layer):
"""Custom layer of metrics for NCF model."""
def __init__(self, match_mlperf):
@@ -81,14 +81,14 @@ def call(self, inputs, training=False):
return logits
-class LossLayer(tf.keras.layers.Layer):
+class LossLayer(tf_keras.layers.Layer):
"""Pass-through loss layer for NCF model."""
def __init__(self, loss_normalization_factor):
# The loss may overflow in float16, so we use float32 instead.
super(LossLayer, self).__init__(dtype="float32")
self.loss_normalization_factor = loss_normalization_factor
- self.loss = tf.keras.losses.SparseCategoricalCrossentropy(
+ self.loss = tf_keras.losses.SparseCategoricalCrossentropy(
from_logits=True, reduction="sum")
def get_config(self):
@@ -107,7 +107,7 @@ def call(self, inputs):
return logits
-class IncrementEpochCallback(tf.keras.callbacks.Callback):
+class IncrementEpochCallback(tf_keras.callbacks.Callback):
"""A callback to increase the requested epoch for the data producer.
The reason why we need this is because we can only buffer a limited amount of
@@ -122,7 +122,7 @@ def on_epoch_begin(self, epoch, logs=None):
self._producer.increment_request_epoch()
-class CustomEarlyStopping(tf.keras.callbacks.Callback):
+class CustomEarlyStopping(tf_keras.callbacks.Callback):
"""Stop training has reached a desired hit rate."""
def __init__(self, monitor, desired_value):
@@ -157,28 +157,28 @@ def _get_keras_model(params):
"""Constructs and returns the model."""
batch_size = params["batch_size"]
- user_input = tf.keras.layers.Input(
+ user_input = tf_keras.layers.Input(
shape=(1,), name=movielens.USER_COLUMN, dtype=tf.int32)
- item_input = tf.keras.layers.Input(
+ item_input = tf_keras.layers.Input(
shape=(1,), name=movielens.ITEM_COLUMN, dtype=tf.int32)
- valid_pt_mask_input = tf.keras.layers.Input(
+ valid_pt_mask_input = tf_keras.layers.Input(
shape=(1,), name=rconst.VALID_POINT_MASK, dtype=tf.bool)
- dup_mask_input = tf.keras.layers.Input(
+ dup_mask_input = tf_keras.layers.Input(
shape=(1,), name=rconst.DUPLICATE_MASK, dtype=tf.int32)
- label_input = tf.keras.layers.Input(
+ label_input = tf_keras.layers.Input(
shape=(1,), name=rconst.TRAIN_LABEL_KEY, dtype=tf.bool)
base_model = neumf_model.construct_model(user_input, item_input, params)
logits = base_model.output
- zeros = tf.keras.layers.Lambda(lambda x: x * 0)(logits)
+ zeros = tf_keras.layers.Lambda(lambda x: x * 0)(logits)
- softmax_logits = tf.keras.layers.concatenate([zeros, logits], axis=-1)
+ softmax_logits = tf_keras.layers.concatenate([zeros, logits], axis=-1)
# Custom training loop calculates loss and metric as a part of
# training/evaluation step function.
@@ -190,7 +190,7 @@ def _get_keras_model(params):
softmax_logits = LossLayer(batch_size)(
[softmax_logits, label_input, valid_pt_mask_input])
- keras_model = tf.keras.Model(
+ keras_model = tf_keras.Model(
inputs={
movielens.USER_COLUMN: user_input,
movielens.ITEM_COLUMN: item_input,
@@ -216,7 +216,7 @@ def run_ncf(_):
model_helpers.apply_clean(FLAGS)
if FLAGS.dtype == "fp16" and FLAGS.fp16_implementation == "keras":
- tf.keras.mixed_precision.set_global_policy("mixed_float16")
+ tf_keras.mixed_precision.set_global_policy("mixed_float16")
strategy = distribute_utils.get_distribution_strategy(
distribution_strategy=FLAGS.distribution_strategy,
@@ -265,7 +265,7 @@ def run_ncf(_):
with distribute_utils.get_strategy_scope(strategy):
keras_model = _get_keras_model(params)
- optimizer = tf.keras.optimizers.Adam(
+ optimizer = tf_keras.optimizers.Adam(
learning_rate=params["learning_rate"],
beta_1=params["beta1"],
beta_2=params["beta2"],
@@ -283,9 +283,9 @@ def run_ncf(_):
# here for the case where a custom training loop or fixed loss scale is
# used.
if loss_scale == "dynamic":
- optimizer = tf.keras.mixed_precision.LossScaleOptimizer(optimizer)
+ optimizer = tf_keras.mixed_precision.LossScaleOptimizer(optimizer)
else:
- optimizer = tf.keras.mixed_precision.LossScaleOptimizer(
+ optimizer = tf_keras.mixed_precision.LossScaleOptimizer(
optimizer, dynamic=False, initial_scale=loss_scale)
if params["keras_use_ctl"]:
@@ -306,10 +306,10 @@ def run_ncf(_):
if not FLAGS.ml_perf:
# Create Tensorboard summary and checkpoint callbacks.
summary_dir = os.path.join(FLAGS.model_dir, "summaries")
- summary_callback = tf.keras.callbacks.TensorBoard(
+ summary_callback = tf_keras.callbacks.TensorBoard(
summary_dir, profile_batch=0)
checkpoint_path = os.path.join(FLAGS.model_dir, "checkpoint")
- checkpoint_callback = tf.keras.callbacks.ModelCheckpoint(
+ checkpoint_callback = tf_keras.callbacks.ModelCheckpoint(
checkpoint_path, save_weights_only=True)
callbacks += [summary_callback, checkpoint_callback]
@@ -342,7 +342,7 @@ def run_ncf(_):
train_history = history.history
train_loss = train_history["loss"][-1]
- stats = build_stats(train_loss, eval_results, time_callback)
+ stats = build_stats(train_loss, eval_results, time_callback) # pyrefly: ignore[unbound-name]
return stats
@@ -375,7 +375,7 @@ def run_ncf_custom_training(params,
Returns:
A tuple of train loss and a list of training and evaluation results.
"""
- loss_object = tf.keras.losses.SparseCategoricalCrossentropy(
+ loss_object = tf_keras.losses.SparseCategoricalCrossentropy(
reduction="sum", from_logits=True)
train_input_iterator = iter(
strategy.experimental_distribute_dataset(train_input_dataset))
@@ -499,7 +499,7 @@ def step_fn(features):
hr_sum / hr_count)
if eval_summary_writer:
with eval_summary_writer.as_default():
- tf.summary.scalar("hit_rate", hr_sum / hr_count, step=current_step)
+ tf.summary.scalar("hit_rate", hr_sum / hr_count, step=current_step) # pyrefly: ignore[unbound-name]
if (FLAGS.early_stopping and
float(hr_sum / hr_count) > params["hr_threshold"]):
diff --git a/official/recommendation/ncf_test.py b/official/recommendation/ncf_test.py
index 2f6f0865cb4..2ef9969ee8b 100644
--- a/official/recommendation/ncf_test.py
+++ b/official/recommendation/ncf_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import unittest
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.eager import context # pylint: disable=ungrouped-imports
from official.recommendation import constants as rconst
from official.recommendation import ncf_common
diff --git a/official/recommendation/neumf_model.py b/official/recommendation/neumf_model.py
index b739546ed13..039ba336bc7 100644
--- a/official/recommendation/neumf_model.py
+++ b/official/recommendation/neumf_model.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -35,8 +35,8 @@
import sys
-from six.moves import xrange # pylint: disable=redefined-builtin
-import tensorflow as tf
+from six.moves import xrange # pylint: disable=redefined-builtin # pyrefly: ignore[missing-source-for-stubs]
+import tensorflow as tf, tf_keras
from tensorflow import estimator as tf_estimator
from typing import Any, Dict, Text
@@ -79,8 +79,8 @@ def neumf_model_fn(features, labels, mode, params):
users = features[movielens.USER_COLUMN]
items = features[movielens.ITEM_COLUMN]
- user_input = tf.keras.layers.Input(tensor=users)
- item_input = tf.keras.layers.Input(tensor=items)
+ user_input = tf_keras.layers.Input(tensor=users)
+ item_input = tf_keras.layers.Input(tensor=items)
logits = construct_model(user_input, item_input, params).output
# Softmax with the first column of zeros is equivalent to sigmoid.
@@ -136,7 +136,7 @@ def _strip_first_and_last_dimension(x, batch_size):
def construct_model(user_input: tf.Tensor, item_input: tf.Tensor,
- params: Dict[Text, Any]) -> tf.keras.Model:
+ params: Dict[Text, Any]) -> tf_keras.Model:
"""Initialize NeuMF model.
Args:
@@ -175,59 +175,59 @@ def mlp_slice_fn(x):
# It turns out to be significantly more effecient to store the MF and MLP
# embedding portions in the same table, and then slice as needed.
- embedding_user = tf.keras.layers.Embedding(
+ embedding_user = tf_keras.layers.Embedding(
num_users,
mf_dim + model_layers[0] // 2,
embeddings_initializer=embedding_initializer,
- embeddings_regularizer=tf.keras.regularizers.l2(mf_regularization),
+ embeddings_regularizer=tf_keras.regularizers.l2(mf_regularization),
input_length=1,
name="embedding_user")(
user_input)
- embedding_item = tf.keras.layers.Embedding(
+ embedding_item = tf_keras.layers.Embedding(
num_items,
mf_dim + model_layers[0] // 2,
embeddings_initializer=embedding_initializer,
- embeddings_regularizer=tf.keras.regularizers.l2(mf_regularization),
+ embeddings_regularizer=tf_keras.regularizers.l2(mf_regularization),
input_length=1,
name="embedding_item")(
item_input)
# GMF part
- mf_user_latent = tf.keras.layers.Lambda(
+ mf_user_latent = tf_keras.layers.Lambda(
mf_slice_fn, name="embedding_user_mf")(
embedding_user)
- mf_item_latent = tf.keras.layers.Lambda(
+ mf_item_latent = tf_keras.layers.Lambda(
mf_slice_fn, name="embedding_item_mf")(
embedding_item)
# MLP part
- mlp_user_latent = tf.keras.layers.Lambda(
+ mlp_user_latent = tf_keras.layers.Lambda(
mlp_slice_fn, name="embedding_user_mlp")(
embedding_user)
- mlp_item_latent = tf.keras.layers.Lambda(
+ mlp_item_latent = tf_keras.layers.Lambda(
mlp_slice_fn, name="embedding_item_mlp")(
embedding_item)
# Element-wise multiply
- mf_vector = tf.keras.layers.multiply([mf_user_latent, mf_item_latent])
+ mf_vector = tf_keras.layers.multiply([mf_user_latent, mf_item_latent])
# Concatenation of two latent features
- mlp_vector = tf.keras.layers.concatenate([mlp_user_latent, mlp_item_latent])
+ mlp_vector = tf_keras.layers.concatenate([mlp_user_latent, mlp_item_latent])
num_layer = len(model_layers) # Number of layers in the MLP
for layer in xrange(1, num_layer):
- model_layer = tf.keras.layers.Dense(
+ model_layer = tf_keras.layers.Dense(
model_layers[layer],
- kernel_regularizer=tf.keras.regularizers.l2(mlp_reg_layers[layer]),
+ kernel_regularizer=tf_keras.regularizers.l2(mlp_reg_layers[layer]),
activation="relu")
mlp_vector = model_layer(mlp_vector)
# Concatenate GMF and MLP parts
- predict_vector = tf.keras.layers.concatenate([mf_vector, mlp_vector])
+ predict_vector = tf_keras.layers.concatenate([mf_vector, mlp_vector])
# Final prediction layer
- logits = tf.keras.layers.Dense(
+ logits = tf_keras.layers.Dense(
1,
activation=None,
kernel_initializer="lecun_uniform",
@@ -235,11 +235,11 @@ def mlp_slice_fn(x):
predict_vector)
# Print model topology.
- model = tf.keras.models.Model([user_input, item_input], logits)
+ model = tf_keras.models.Model([user_input, item_input], logits)
model.summary()
sys.stdout.flush()
- return model
+ return model # pyrefly: ignore[bad-return]
def _get_estimator_spec_with_metrics(logits: tf.Tensor,
diff --git a/official/recommendation/popen_helper.py b/official/recommendation/popen_helper.py
index 4004c207fab..66107a864a4 100644
--- a/official/recommendation/popen_helper.py
+++ b/official/recommendation/popen_helper.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/ranking/__init__.py b/official/recommendation/ranking/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/recommendation/ranking/__init__.py
+++ b/official/recommendation/ranking/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/ranking/common.py b/official/recommendation/ranking/common.py
index f7bdf49ea5a..2f1c907d89b 100644
--- a/official/recommendation/ranking/common.py
+++ b/official/recommendation/ranking/common.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Flags and common definitions for Ranking Models."""
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import flags as tfm_flags
@@ -44,8 +44,8 @@ def define_flags() -> None:
'overhead, and the output file can be gigantic if profiling many steps.')
-@tf.keras.utils.register_keras_serializable(package='RANKING')
-class WarmUpAndPolyDecay(tf.keras.optimizers.schedules.LearningRateSchedule):
+@tf_keras.utils.register_keras_serializable(package='RANKING')
+class WarmUpAndPolyDecay(tf_keras.optimizers.schedules.LearningRateSchedule):
"""Learning rate callable for the embeddings.
Linear warmup on [0, warmup_steps] then
diff --git a/official/recommendation/ranking/configs/__init__.py b/official/recommendation/ranking/configs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/recommendation/ranking/configs/__init__.py
+++ b/official/recommendation/ranking/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/ranking/configs/config.py b/official/recommendation/ranking/configs/config.py
index d7fa5807dc1..26af84531c9 100644
--- a/official/recommendation/ranking/configs/config.py
+++ b/official/recommendation/ranking/configs/config.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -52,9 +52,15 @@ class LearningRateConfig(hyperparams.Config):
@dataclasses.dataclass
class OptimizationConfig(hyperparams.Config):
- """Embedding Optimizer config."""
- lr_config: LearningRateConfig = LearningRateConfig()
+ """Embedding and dense optimizer configs."""
+ lr_config: LearningRateConfig = dataclasses.field(
+ default_factory=LearningRateConfig
+ )
+ dense_sgd_config: LearningRateConfig = dataclasses.field(
+ default_factory=lambda: LearningRateConfig(warmup_steps=0)
+ )
embedding_optimizer: str = 'SGD'
+ dense_optimizer: str = 'Adam'
@dataclasses.dataclass
@@ -68,6 +74,7 @@ class DataConfig(hyperparams.Config):
cycle_length: int = 10
sharding: bool = True
num_shards_per_host: int = 8
+ use_cached_data: bool = False
@dataclasses.dataclass
@@ -78,6 +85,10 @@ class ModelConfig(hyperparams.Config):
num_dense_features: Number of dense features.
vocab_sizes: Vocab sizes for each of the sparse features. The order agrees
with the order of the input data.
+ use_multi_hot: Flag to determine if enabling multi-hot data loading and
+ training used for DLRM V2
+ multi_hot_sizes: List to pass in the multi hot size of each sparse embedding
+ feature
embedding_dim: An integer or a list of embedding table dimensions.
If it's an integer then all tables will have the same embedding dimension.
If it's a list then the length should match with `vocab_sizes`.
@@ -89,15 +100,46 @@ class ModelConfig(hyperparams.Config):
features.
top_mlp: The sizes of hidden layers for top MLP.
interaction: Interaction can be on of the following:
- 'dot', 'cross'.
+ 'dot', 'cross', 'multi_layer_dcn'.
+ concat_dense: Weather to concatenate output from interaction module with
+ dense output again
+ dcn_num_layers: Number of Stacked DCN layers used in the dcn interaction
+ module
+ dcn_low_rank_dim: Project dimension in stacked DCN layers for dcn
+ interaction module
+ dcn_kernel_initializer: the kernel initializer used for the dcn interaction
+ module
+ dcn_bias_initializer: the bias initializer used for the dcn interaction
+ module
+ dcn_use_bias: Flag to determine whether to use bias for the dcn interaction
+ module
+ use_partial_tpu_embedding: Flag to determine whether to use partial tpu
+ embedding layer or not.
+ max_ids_per_chip_per_sample: Maximum number of ids per chip per sample.
+ max_ids_per_table: Maximum number of ids per table.
+ max_unique_ids_per_table: Maximum number of unique ids per table.
"""
num_dense_features: int = 13
vocab_sizes: List[int] = dataclasses.field(default_factory=list)
+ use_multi_hot: bool = False
+ multi_hot_sizes: List[int] = dataclasses.field(default_factory=list)
embedding_dim: Union[int, List[int]] = 8
size_threshold: int = 50_000
bottom_mlp: List[int] = dataclasses.field(default_factory=list)
top_mlp: List[int] = dataclasses.field(default_factory=list)
interaction: str = 'dot'
+ concat_dense: bool = True
+ dcn_num_layers: int = 3
+ dcn_low_rank_dim: int = 512
+ dcn_kernel_initializer: str = 'truncated_normal'
+ dcn_bias_initializer: str = 'zeros'
+ dcn_use_bias: bool = True
+ use_partial_tpu_embedding: bool = True
+ max_ids_per_chip_per_sample: int | None = None
+ max_ids_per_table: Union[int, List[int]] | None = None
+ max_unique_ids_per_table: Union[int, List[int]] | None = None
+ allow_id_dropping: bool = False
+ initialize_tables_on_host: bool = False
@dataclasses.dataclass
@@ -115,11 +157,16 @@ class Loss(hyperparams.Config):
class Task(hyperparams.Config):
"""The model config."""
init_checkpoint: str = ''
- model: ModelConfig = ModelConfig()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- loss: Loss = Loss()
+ model: ModelConfig = dataclasses.field(default_factory=ModelConfig)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ loss: Loss = dataclasses.field(default_factory=Loss)
use_synthetic_data: bool = False
+ use_tf_record_reader: bool = False
@dataclasses.dataclass
@@ -148,16 +195,25 @@ class TrainerConfig(cfg.TrainerConfig):
time_history: Config of TimeHistory callback.
optimizer_config: An `OptimizerConfig` instance for embedding optimizer.
Defaults to None.
+ pipeline_sparse_and_dense_exeuction: Whether to pipeline embedding and
+ dense execution. This is a performance optimization.
"""
train_steps: int = 0
# Sets validation steps to be -1 to evaluate the entire dataset.
validation_steps: int = -1
validation_interval: int = 70000
- callbacks: CallbacksConfig = CallbacksConfig()
+ callbacks: CallbacksConfig = dataclasses.field(
+ default_factory=CallbacksConfig
+ )
use_orbit: bool = False
enable_metrics_in_training: bool = True
- time_history: TimeHistoryConfig = TimeHistoryConfig(log_steps=5000)
- optimizer_config: OptimizationConfig = OptimizationConfig()
+ time_history: TimeHistoryConfig = dataclasses.field(
+ default_factory=lambda: TimeHistoryConfig(log_steps=5000)
+ )
+ optimizer_config: OptimizationConfig = dataclasses.field(
+ default_factory=OptimizationConfig
+ )
+ pipeline_sparse_and_dense_execution: bool = False
NUM_TRAIN_EXAMPLES = 4195197692
@@ -167,10 +223,61 @@ class TrainerConfig(cfg.TrainerConfig):
eval_batch_size = 16384
steps_per_epoch = NUM_TRAIN_EXAMPLES // train_batch_size
vocab_sizes = [
- 39884406, 39043, 17289, 7420, 20263, 3, 7120, 1543, 63, 38532951,
- 2953546, 403346, 10, 2208, 11938, 155, 4, 976, 14, 39979771, 25641295,
- 39664984, 585935, 12972, 108, 36
- ]
+ 39884406,
+ 39043,
+ 17289,
+ 7420,
+ 20263,
+ 3,
+ 7120,
+ 1543,
+ 63,
+ 38532951,
+ 2953546,
+ 403346,
+ 10,
+ 2208,
+ 11938,
+ 155,
+ 4,
+ 976,
+ 14,
+ 39979771,
+ 25641295,
+ 39664984,
+ 585935,
+ 12972,
+ 108,
+ 36,
+]
+multi_hot_sizes = [
+ 3,
+ 2,
+ 1,
+ 2,
+ 6,
+ 1,
+ 1,
+ 1,
+ 1,
+ 7,
+ 3,
+ 8,
+ 1,
+ 6,
+ 9,
+ 5,
+ 1,
+ 1,
+ 1,
+ 12,
+ 100,
+ 27,
+ 10,
+ 3,
+ 1,
+ 1,
+]
@dataclasses.dataclass
@@ -184,26 +291,35 @@ class Config(hyperparams.Config):
task: `Task` instance.
trainer: A `TrainerConfig` instance.
"""
- runtime: cfg.RuntimeConfig = cfg.RuntimeConfig()
- task: Task = Task(
- model=ModelConfig(
- embedding_dim=8,
- vocab_sizes=vocab_sizes,
- bottom_mlp=[64, 32, 8],
- top_mlp=[64, 32, 1]),
- loss=Loss(label_smoothing=0.0),
- train_data=DataConfig(
- is_training=True,
- global_batch_size=train_batch_size),
- validation_data=DataConfig(
- is_training=False,
- global_batch_size=eval_batch_size))
- trainer: TrainerConfig = TrainerConfig(
- train_steps=2 * steps_per_epoch,
- validation_interval=steps_per_epoch,
- validation_steps=NUM_EVAL_EXAMPLES // eval_batch_size,
- enable_metrics_in_training=True,
- optimizer_config=OptimizationConfig())
+ runtime: cfg.RuntimeConfig = dataclasses.field(
+ default_factory=cfg.RuntimeConfig
+ )
+ task: Task = dataclasses.field(
+ default_factory=lambda: Task( # pylint: disable=g-long-lambda
+ model=ModelConfig(
+ embedding_dim=8,
+ vocab_sizes=vocab_sizes,
+ bottom_mlp=[64, 32, 8],
+ top_mlp=[64, 32, 1],
+ ),
+ loss=Loss(label_smoothing=0.0),
+ train_data=DataConfig(
+ is_training=True, global_batch_size=train_batch_size
+ ),
+ validation_data=DataConfig(
+ is_training=False, global_batch_size=eval_batch_size
+ ),
+ )
+ )
+ trainer: TrainerConfig = dataclasses.field(
+ default_factory=lambda: TrainerConfig( # pylint: disable=g-long-lambda
+ train_steps=2 * steps_per_epoch,
+ validation_interval=steps_per_epoch,
+ validation_steps=NUM_EVAL_EXAMPLES // eval_batch_size,
+ enable_metrics_in_training=True,
+ optimizer_config=OptimizationConfig(),
+ )
+ )
restrictions: dataclasses.InitVar[Optional[List[str]]] = None
@@ -301,3 +417,44 @@ def dcn_criteo_tb_config() -> Config:
'task.train_data.is_training != None',
'task.validation_data.is_training != None',
])
+
+
+@exp_factory.register_config_factory('dlrm_dcn_v2_criteo')
+def dlrm_dcn_v2_criteo_tb_config() -> Config:
+ return Config(
+ runtime=cfg.RuntimeConfig(),
+ task=Task(
+ model=ModelConfig(
+ num_dense_features=13,
+ vocab_sizes=vocab_sizes,
+ bottom_mlp=[512, 256, 64],
+ embedding_dim=64,
+ top_mlp=[1024, 1024, 512, 256, 1],
+ interaction='multi_layer_dcn',
+ dcn_num_layers=3,
+ dcn_low_rank_dim=512,
+ dcn_use_bias=True,
+ concat_dense=False,
+ use_multi_hot=True,
+ use_partial_tpu_embedding=False,
+ multi_hot_sizes=multi_hot_sizes,
+ ),
+ loss=Loss(label_smoothing=0.0),
+ train_data=DataConfig(
+ global_batch_size=train_batch_size,
+ is_training=True,
+ sharding=True),
+ validation_data=DataConfig(
+ global_batch_size=eval_batch_size,
+ is_training=False,
+ sharding=False)),
+ trainer=TrainerConfig(
+ train_steps=steps_per_epoch,
+ validation_interval=steps_per_epoch // 2,
+ validation_steps=NUM_EVAL_EXAMPLES // eval_batch_size,
+ enable_metrics_in_training=True,
+ optimizer_config=OptimizationConfig()),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ])
diff --git a/official/recommendation/ranking/configs/config_test.py b/official/recommendation/ranking/configs/config_test.py
index 890a1943b8e..0aa2592274e 100644
--- a/official/recommendation/ranking/configs/config_test.py
+++ b/official/recommendation/ranking/configs/config_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Unit tests for DLRM config."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.recommendation.ranking.configs import config
diff --git a/official/recommendation/ranking/data/__init__.py b/official/recommendation/ranking/data/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/recommendation/ranking/data/__init__.py
+++ b/official/recommendation/ranking/data/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/ranking/data/data_pipeline.py b/official/recommendation/ranking/data/data_pipeline.py
index 8a8a4a1b6e8..d8bd2063160 100644
--- a/official/recommendation/ranking/data/data_pipeline.py
+++ b/official/recommendation/ranking/data/data_pipeline.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
"""
from typing import List
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.recommendation.ranking.configs import config
@@ -38,12 +38,14 @@ def __init__(self,
params: config.DataConfig,
num_dense_features: int,
vocab_sizes: List[int],
- use_synthetic_data: bool = False):
+ use_synthetic_data: bool = False,
+ use_cached_data: bool = False):
self._file_pattern = file_pattern
self._params = params
self._num_dense_features = num_dense_features
self._vocab_sizes = vocab_sizes
self._use_synthetic_data = use_synthetic_data
+ self._use_cached_data = use_cached_data
def __call__(self, ctx: tf.distribute.InputContext) -> tf.data.Dataset:
params = self._params
@@ -117,6 +119,8 @@ def make_dataset(shard_index):
num_parallel_calls=tf.data.experimental.AUTOTUNE)
dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)
+ if self._use_cached_data:
+ dataset = dataset.take(1).cache().repeat()
return dataset
@@ -173,6 +177,9 @@ def _generate_synthetic_data(self, ctx: tf.distribute.InputContext,
if params.is_training:
dataset = dataset.repeat()
+ if self._use_cached_data:
+ dataset = dataset.take(1).cache().repeat()
+
return dataset.batch(batch_size, drop_remainder=True)
diff --git a/official/recommendation/ranking/data/data_pipeline_multi_hot.py b/official/recommendation/ranking/data/data_pipeline_multi_hot.py
new file mode 100644
index 00000000000..b14efab6b29
--- /dev/null
+++ b/official/recommendation/ranking/data/data_pipeline_multi_hot.py
@@ -0,0 +1,354 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data pipeline for the Ranking model.
+
+This module defines various input datasets for the Ranking model.
+"""
+
+from typing import List
+import tensorflow as tf, tf_keras
+
+from official.recommendation.ranking.configs import config
+
+
+class CriteoTsvReaderMultiHot:
+ """Input reader callable for pre-processed Multi Hot Criteo data.
+
+ Raw Criteo data is assumed to be preprocessed in the following way:
+ 1. Missing values are replaced with zeros.
+ 2. Negative values are replaced with zeros.
+ 3. Integer features are transformed by log(x+1) and are hence tf.float32.
+ 4. Categorical data is bucketized and are hence tf.int32.
+
+ Implements a TsvReaderMultiHot for reading data from a criteo dataset and
+ generate multi hot synthetic data using the provided vocab_sizes and
+ multi_hot_sizes, also includes a complete synthetic data generator as well as
+ a TFRecordReader to read data from pre materialized multi hot synthetic
+ dataset that converted to TFRecords
+ """
+
+ def __init__(self,
+ file_pattern: str,
+ params: config.DataConfig,
+ num_dense_features: int,
+ vocab_sizes: List[int],
+ multi_hot_sizes: List[int],
+ use_synthetic_data: bool = False):
+ self._file_pattern = file_pattern
+ self._params = params
+ self._num_dense_features = num_dense_features
+ self._vocab_sizes = vocab_sizes
+ self._use_synthetic_data = use_synthetic_data
+ self._multi_hot_sizes = multi_hot_sizes
+
+ def __call__(self, ctx: tf.distribute.InputContext) -> tf.data.Dataset:
+ params = self._params
+ # Per replica batch size.
+ batch_size = ctx.get_per_replica_batch_size(
+ params.global_batch_size) if ctx else params.global_batch_size
+ if self._use_synthetic_data:
+ return self._generate_synthetic_data(ctx, batch_size)
+
+ @tf.function
+ def _parse_fn(example: tf.Tensor):
+ """Parser function for pre-processed Criteo TSV records."""
+ label_defaults = [[0.0]]
+ dense_defaults = [
+ [0.0] for _ in range(self._num_dense_features)
+ ]
+ num_sparse_features = len(self._vocab_sizes)
+ categorical_defaults = [
+ [0] for _ in range(num_sparse_features)
+ ]
+ record_defaults = label_defaults + dense_defaults + categorical_defaults
+ fields = tf.io.decode_csv(
+ example, record_defaults, field_delim='\t', na_value='-1')
+
+ num_labels = 1
+ label = tf.reshape(fields[0], [batch_size, 1])
+
+ features = {}
+ num_dense = len(dense_defaults)
+
+ dense_features = []
+ offset = num_labels
+ for idx in range(num_dense):
+ dense_features.append(fields[idx + offset])
+ features['dense_features'] = tf.stack(dense_features, axis=1)
+
+ offset += num_dense
+ features['sparse_features'] = {}
+
+ sparse_tensors = []
+ for idx, (vocab_size, multi_hot_size) in enumerate(
+ zip(self._vocab_sizes, self._multi_hot_sizes)
+ ):
+ sparse_tensor = tf.reshape(fields[idx + offset], [batch_size, 1])
+ sparse_tensor_synthetic = tf.random.uniform(
+ shape=(batch_size, multi_hot_size - 1),
+ maxval=int(vocab_size),
+ dtype=tf.int32,
+ )
+ sparse_tensors.append(
+ tf.sparse.from_dense(
+ tf.concat([sparse_tensor, sparse_tensor_synthetic], axis=1)
+ )
+ )
+
+ sparse_tensor_elements = {
+ str(i): sparse_tensors[i] for i in range(len(sparse_tensors))
+ }
+
+ features['sparse_features'] = sparse_tensor_elements
+
+ return features, label
+
+ filenames = tf.data.Dataset.list_files(self._file_pattern, shuffle=False)
+
+ # Shard the full dataset according to host number.
+ # Each host will get 1 / num_of_hosts portion of the data.
+ if params.sharding and ctx and ctx.num_input_pipelines > 1:
+ filenames = filenames.shard(ctx.num_input_pipelines,
+ ctx.input_pipeline_id)
+
+ num_shards_per_host = 1
+ if params.sharding:
+ num_shards_per_host = params.num_shards_per_host
+
+ def make_dataset(shard_index):
+ filenames_for_shard = filenames.shard(num_shards_per_host, shard_index)
+ dataset = tf.data.TextLineDataset(filenames_for_shard)
+ if params.is_training:
+ dataset = dataset.repeat()
+ dataset = dataset.batch(batch_size, drop_remainder=True)
+ dataset = dataset.map(_parse_fn,
+ num_parallel_calls=tf.data.experimental.AUTOTUNE)
+ return dataset
+
+ indices = tf.data.Dataset.range(num_shards_per_host)
+ dataset = indices.interleave(
+ map_func=make_dataset,
+ cycle_length=params.cycle_length,
+ num_parallel_calls=tf.data.experimental.AUTOTUNE)
+
+ dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)
+ if self._params.use_cached_data:
+ dataset = dataset.take(1).cache().repeat()
+
+ return dataset
+
+ def _generate_synthetic_data(self, ctx: tf.distribute.InputContext,
+ batch_size: int) -> tf.data.Dataset:
+ """Creates synthetic data based on the parameter batch size.
+
+ Args:
+ ctx: Input Context
+ batch_size: per replica batch size.
+
+ Returns:
+ The synthetic dataset.
+ """
+ params = self._params
+ num_dense = self._num_dense_features
+ num_replicas = ctx.num_replicas_in_sync if ctx else 1
+
+ if params.is_training:
+ dataset_size = 50 * batch_size * num_replicas
+ else:
+ dataset_size = 50 * batch_size * num_replicas
+ dense_tensor = tf.random.uniform(
+ shape=(dataset_size, num_dense), maxval=1.0, dtype=tf.float32
+ )
+
+ sparse_tensors = []
+ for vocab_size, multi_hot_size in zip(
+ self._vocab_sizes, self._multi_hot_sizes
+ ):
+ sparse_tensors.append(
+ tf.sparse.from_dense(
+ tf.random.uniform(
+ shape=(dataset_size, multi_hot_size),
+ maxval=int(vocab_size),
+ dtype=tf.int32,
+ )
+ )
+ )
+
+ sparse_tensor_elements = {
+ str(i): sparse_tensors[i] for i in range(len(sparse_tensors))
+ }
+
+ # the mean is in [0, 1] interval.
+ dense_tensor_mean = tf.math.reduce_mean(dense_tensor, axis=1)
+
+ # the label is in [0, 1] interval.
+ label_tensor = (dense_tensor_mean)
+ # Using the threshold 0.5 to convert to 0/1 labels.
+ label_tensor = tf.cast(label_tensor + 0.5, tf.int32)
+
+ input_elem = {'dense_features': dense_tensor,
+ 'sparse_features': sparse_tensor_elements}, label_tensor
+
+ dataset = tf.data.Dataset.from_tensor_slices(input_elem)
+ dataset = dataset.cache()
+ if params.is_training:
+ dataset = dataset.repeat()
+
+ return dataset.batch(batch_size, drop_remainder=True)
+
+
+class CriteoTFRecordReader(object):
+ """Input reader fn for TFRecords that have been serialized in batched form."""
+
+ def __init__(self,
+ file_pattern: str,
+ params: config.DataConfig,
+ num_dense_features: int,
+ vocab_sizes: List[int],
+ multi_hot_sizes: List[int],
+ use_cached_data: bool = False):
+ self._file_pattern = file_pattern
+ self._params = params
+ self._num_dense_features = num_dense_features
+ self._vocab_sizes = vocab_sizes
+ self._multi_hot_sizes = multi_hot_sizes
+ self._use_cached_data = use_cached_data
+
+ self.label_features = 'label'
+ self.dense_features = ['dense-feature-%d' % x for x in range(1, 14)]
+ self.sparse_features = ['sparse-feature-%d' % x for x in range(14, 40)]
+
+ def __call__(self, ctx: tf.distribute.InputContext):
+ params = self._params
+ # Per replica batch size.
+ batch_size = (
+ ctx.get_per_replica_batch_size(params.global_batch_size)
+ if ctx
+ else params.global_batch_size
+ )
+
+ def _get_feature_spec():
+ feature_spec = {}
+ feature_spec[self.label_features] = tf.io.FixedLenFeature(
+ [], dtype=tf.int64
+ )
+ for dense_feat in self.dense_features:
+ feature_spec[dense_feat] = tf.io.FixedLenFeature(
+ [],
+ dtype=tf.float32,
+ )
+ for i, sparse_feat in enumerate(self.sparse_features):
+ feature_spec[sparse_feat] = tf.io.FixedLenFeature(
+ [self._multi_hot_sizes[i]], dtype=tf.int64
+ )
+ return feature_spec
+
+ def _parse_fn(serialized_example):
+ feature_spec = _get_feature_spec()
+ parsed_features = tf.io.parse_single_example(
+ serialized_example, feature_spec
+ )
+ label = parsed_features[self.label_features]
+ features = {}
+ int_features = []
+ for dense_ft in self.dense_features:
+ int_features.append(parsed_features[dense_ft])
+ features['dense_features'] = tf.stack(int_features)
+
+ features['sparse_features'] = {}
+ for i, sparse_ft in enumerate(self.sparse_features):
+ features['sparse_features'][str(i)] = tf.sparse.from_dense(
+ parsed_features[sparse_ft]
+ )
+
+ return features, label
+
+ filenames = tf.data.Dataset.list_files(self._file_pattern, shuffle=False)
+ # Shard the full dataset according to host number.
+ # Each host will get 1 / num_of_hosts portion of the data.
+ if params.sharding and ctx and ctx.num_input_pipelines > 1:
+ filenames = filenames.shard(ctx.num_input_pipelines,
+ ctx.input_pipeline_id)
+
+ num_shards_per_host = 1
+ if params.sharding:
+ num_shards_per_host = params.num_shards_per_host
+
+ def make_dataset(shard_index):
+ filenames_for_shard = filenames.shard(num_shards_per_host, shard_index)
+ dataset = tf.data.TFRecordDataset(
+ filenames_for_shard
+ )
+ if params.is_training:
+ dataset = dataset.repeat()
+ dataset = dataset.map(
+ _parse_fn, num_parallel_calls=tf.data.experimental.AUTOTUNE
+ )
+ return dataset
+
+ indices = tf.data.Dataset.range(num_shards_per_host)
+ dataset = indices.interleave(
+ map_func=make_dataset,
+ cycle_length=params.cycle_length,
+ num_parallel_calls=tf.data.experimental.AUTOTUNE,
+ )
+
+ dataset = dataset.batch(
+ batch_size,
+ drop_remainder=True,
+ num_parallel_calls=tf.data.experimental.AUTOTUNE,
+ )
+ dataset = dataset.prefetch(buffer_size=tf.data.experimental.AUTOTUNE)
+ if self._use_cached_data:
+ dataset = dataset.take(1).cache().repeat()
+
+ return dataset
+
+
+def train_input_fn(params: config.Task) -> CriteoTsvReaderMultiHot:
+ """Returns callable object of batched training examples.
+
+ Args:
+ params: hyperparams to create input pipelines.
+
+ Returns:
+ CriteoTsvReader callable for training dataset.
+ """
+ return CriteoTsvReaderMultiHot(
+ file_pattern=params.train_data.input_path,
+ params=params.train_data,
+ vocab_sizes=params.model.vocab_sizes,
+ num_dense_features=params.model.num_dense_features,
+ multi_hot_sizes=params.model.multi_hot_sizes,
+ use_synthetic_data=params.use_synthetic_data)
+
+
+def eval_input_fn(params: config.Task) -> CriteoTsvReaderMultiHot:
+ """Returns callable object of batched eval examples.
+
+ Args:
+ params: hyperparams to create input pipelines.
+
+ Returns:
+ CriteoTsvReader callable for eval dataset.
+ """
+
+ return CriteoTsvReaderMultiHot(
+ file_pattern=params.validation_data.input_path,
+ params=params.validation_data,
+ vocab_sizes=params.model.vocab_sizes,
+ num_dense_features=params.model.num_dense_features,
+ multi_hot_sizes=params.model.multi_hot_sizes,
+ use_synthetic_data=params.use_synthetic_data)
diff --git a/official/recommendation/ranking/data/data_pipeline_multi_hot_test.py b/official/recommendation/ranking/data/data_pipeline_multi_hot_test.py
new file mode 100644
index 00000000000..55c24354fd7
--- /dev/null
+++ b/official/recommendation/ranking/data/data_pipeline_multi_hot_test.py
@@ -0,0 +1,78 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Unit tests for data_pipeline."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.recommendation.ranking.configs import config
+from official.recommendation.ranking.data import data_pipeline_multi_hot
+
+
+class DataPipelineTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.named_parameters(
+ ('TrainCached', True, True),
+ ('EvalNotCached', False, False),
+ ('TrainNotCached', True, False),
+ ('EvalCached', False, True),
+ )
+ def testSyntheticDataPipeline(self, is_training, use_cached_data):
+ task = config.Task(
+ model=config.ModelConfig(
+ embedding_dim=4,
+ num_dense_features=8,
+ vocab_sizes=[40, 12, 11, 13, 2, 5],
+ multi_hot_sizes=[2, 3, 1, 1, 3, 2],
+ use_multi_hot=True,
+ concat_dense=False,
+ interaction='multi_layer_dcn',
+ dcn_num_layers=3,
+ dcn_low_rank_dim=64,
+ bottom_mlp=[64, 32, 4],
+ top_mlp=[64, 32, 1]),
+ train_data=config.DataConfig(global_batch_size=16,
+ use_cached_data=use_cached_data),
+ validation_data=config.DataConfig(global_batch_size=16,
+ use_cached_data=use_cached_data),
+ use_synthetic_data=True)
+
+ num_dense_features = task.model.num_dense_features
+ num_sparse_features = len(task.model.vocab_sizes)
+ batch_size = task.train_data.global_batch_size
+
+ if is_training:
+ dataset = data_pipeline_multi_hot.train_input_fn(task)
+ else:
+ dataset = data_pipeline_multi_hot.eval_input_fn(task)
+
+ dataset_iter = iter(dataset(ctx=None))
+ print('task model', task.model)
+ # Consume full batches and validate shapes.
+ for _ in range(10):
+ features, label = next(dataset_iter)
+ dense_features = features['dense_features']
+ sparse_features = features['sparse_features']
+ self.assertEqual(dense_features.shape, [batch_size, num_dense_features])
+ self.assertLen(sparse_features, num_sparse_features)
+ for idx, (_, val) in enumerate(sparse_features.items()):
+ self.assertEqual(
+ val.shape, [batch_size, task.model.multi_hot_sizes[idx]]
+ )
+ self.assertEqual(label.shape, [batch_size])
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/recommendation/ranking/data/data_pipeline_test.py b/official/recommendation/ranking/data/data_pipeline_test.py
index d33f1564da3..472ae8c2fbc 100644
--- a/official/recommendation/ranking/data/data_pipeline_test.py
+++ b/official/recommendation/ranking/data/data_pipeline_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Unit tests for data_pipeline."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.recommendation.ranking.configs import config
from official.recommendation.ranking.data import data_pipeline
@@ -23,19 +23,29 @@
class DataPipelineTest(parameterized.TestCase, tf.test.TestCase):
- @parameterized.named_parameters(('Train', True),
- ('Eval', False))
- def testSyntheticDataPipeline(self, is_training):
+ @parameterized.named_parameters(
+ ('TrainCached', True, True),
+ ('EvalNotCached', False, False),
+ ('TrainNotCached', True, False),
+ ('EvalCached', False, True),
+ )
+ def testSyntheticDataPipeline(self, is_training, use_cached_data):
task = config.Task(
model=config.ModelConfig(
embedding_dim=4,
num_dense_features=8,
vocab_sizes=[40, 12, 11, 13, 2, 5],
bottom_mlp=[64, 32, 4],
- top_mlp=[64, 32, 1]),
- train_data=config.DataConfig(global_batch_size=16),
- validation_data=config.DataConfig(global_batch_size=16),
- use_synthetic_data=True)
+ top_mlp=[64, 32, 1],
+ ),
+ train_data=config.DataConfig(
+ global_batch_size=16, use_cached_data=use_cached_data
+ ),
+ validation_data=config.DataConfig(
+ global_batch_size=16, use_cached_data=use_cached_data
+ ),
+ use_synthetic_data=True,
+ )
num_dense_features = task.model.num_dense_features
num_sparse_features = len(task.model.vocab_sizes)
diff --git a/official/recommendation/ranking/preprocessing/criteo_preprocess.py b/official/recommendation/ranking/preprocessing/criteo_preprocess.py
index 7f0f5ae5e47..5f7074419ab 100644
--- a/official/recommendation/ranking/preprocessing/criteo_preprocess.py
+++ b/official/recommendation/ranking/preprocessing/criteo_preprocess.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -36,7 +36,7 @@
import apache_beam as beam
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_transform as tft
import tensorflow_transform.beam as tft_beam
from tensorflow_transform.tf_metadata import dataset_metadata
diff --git a/official/recommendation/ranking/preprocessing/setup.py b/official/recommendation/ranking/preprocessing/setup.py
index 37184cdddc7..91aab859d5b 100644
--- a/official/recommendation/ranking/preprocessing/setup.py
+++ b/official/recommendation/ranking/preprocessing/setup.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/ranking/preprocessing/shard_rebalancer.py b/official/recommendation/ranking/preprocessing/shard_rebalancer.py
index 51465025952..8d3eddef9d3 100644
--- a/official/recommendation/ranking/preprocessing/shard_rebalancer.py
+++ b/official/recommendation/ranking/preprocessing/shard_rebalancer.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
import os
import apache_beam as beam
-import tensorflow as tf
+import tensorflow as tf, tf_keras
parser = argparse.ArgumentParser()
diff --git a/official/recommendation/ranking/task.py b/official/recommendation/ranking/task.py
index 105aa0c3d5b..929acd4d789 100644
--- a/official/recommendation/ranking/task.py
+++ b/official/recommendation/ranking/task.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,9 @@
"""Task for the Ranking model."""
import math
-from typing import Dict, List, Optional, Union
+from typing import Dict, List, Optional, Union, Tuple
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_recommenders as tfrs
from official.core import base_task
@@ -25,6 +25,8 @@
from official.recommendation.ranking import common
from official.recommendation.ranking.configs import config
from official.recommendation.ranking.data import data_pipeline
+from official.recommendation.ranking.data import data_pipeline_multi_hot
+
RuntimeConfig = config_definitions.RuntimeConfig
@@ -32,8 +34,17 @@
def _get_tpu_embedding_feature_config(
vocab_sizes: List[int],
embedding_dim: Union[int, List[int]],
- table_name_prefix: str = 'embedding_table'
-) -> Dict[str, tf.tpu.experimental.embedding.FeatureConfig]:
+ table_name_prefix: str = 'embedding_table',
+ batch_size: Optional[int] = None,
+ max_ids_per_chip_per_sample: Optional[int] = None,
+ max_ids_per_table: Optional[Union[int, List[int]]] = None,
+ max_unique_ids_per_table: Optional[Union[int, List[int]]] = None,
+ allow_id_dropping: bool = False,
+ initialize_tables_on_host: bool = False,
+) -> Tuple[
+ Dict[str, tf.tpu.experimental.embedding.FeatureConfig],
+ Optional[tf.tpu.experimental.embedding.SparseCoreEmbeddingConfig],
+]:
"""Returns TPU embedding feature config.
i'th table config will have vocab size of vocab_sizes[i] and embedding
@@ -43,6 +54,16 @@ def _get_tpu_embedding_feature_config(
vocab_sizes: List of sizes of categories/id's in the table.
embedding_dim: An integer or a list of embedding table dimensions.
table_name_prefix: a prefix for embedding tables.
+ batch_size: Per-replica batch size.
+ max_ids_per_chip_per_sample: Maximum number of embedding ids per chip per
+ sample.
+ max_ids_per_table: Maximum number of embedding ids per table.
+ max_unique_ids_per_table: Maximum number of unique embedding ids per table.
+ allow_id_dropping: bool to allow id dropping.
+ initialize_tables_on_host: bool : if the embedding table size is more than
+ what HBM can handle, this flag will help initialize the full embedding
+ tables on host and then copy shards to HBM.
+
Returns:
A dictionary of feature_name, FeatureConfig pairs.
"""
@@ -50,14 +71,49 @@ def _get_tpu_embedding_feature_config(
if len(vocab_sizes) != len(embedding_dim):
raise ValueError(
f'length of vocab_sizes: {len(vocab_sizes)} is not equal to the '
- f'length of embedding_dim: {len(embedding_dim)}')
+ f'length of embedding_dim: {len(embedding_dim)}'
+ )
elif isinstance(embedding_dim, int):
embedding_dim = [embedding_dim] * len(vocab_sizes)
else:
- raise ValueError('embedding_dim is not either a list or an int, got '
- f'{type(embedding_dim)}')
+ raise ValueError(
+ 'embedding_dim is not either a list or an int, got '
+ f'{type(embedding_dim)}'
+ )
+
+ if isinstance(max_ids_per_table, List):
+ if len(vocab_sizes) != len(max_ids_per_table):
+ raise ValueError(
+ f'length of vocab_sizes: {len(vocab_sizes)} is not equal to the '
+ f'length of max_ids_per_table: {len(max_ids_per_table)}'
+ )
+ elif isinstance(max_ids_per_table, int):
+ max_ids_per_table = [max_ids_per_table] * len(vocab_sizes)
+ elif max_ids_per_table is not None:
+ raise ValueError(
+ 'max_ids_per_table is not either a list or an int or None, got '
+ f'{type(max_ids_per_table)}'
+ )
+
+ if isinstance(max_unique_ids_per_table, List):
+ if len(vocab_sizes) != len(max_unique_ids_per_table):
+ raise ValueError(
+ f'length of vocab_sizes: {len(vocab_sizes)} is not equal to the '
+ 'length of max_unique_ids_per_table: '
+ f'{len(max_unique_ids_per_table)}'
+ )
+ elif isinstance(max_unique_ids_per_table, int):
+ max_unique_ids_per_table = [max_unique_ids_per_table] * len(vocab_sizes)
+ elif max_unique_ids_per_table is not None:
+ raise ValueError(
+ 'max_unique_ids_per_table is not either a list or an int or None, '
+ f'got {type(max_unique_ids_per_table)}'
+ )
feature_config = {}
+ sparsecore_config = None
+ max_ids_per_table_dict = {}
+ max_unique_ids_per_table_dict = {}
for i, vocab_size in enumerate(vocab_sizes):
table_config = tf.tpu.experimental.embedding.TableConfig(
@@ -65,12 +121,36 @@ def _get_tpu_embedding_feature_config(
dim=embedding_dim[i],
combiner='mean',
initializer=tf.initializers.TruncatedNormal(
- mean=0.0, stddev=1 / math.sqrt(embedding_dim[i])),
- name=table_name_prefix + '_%s' % i)
+ mean=0.0, stddev=1 / math.sqrt(embedding_dim[i])
+ ),
+ name=table_name_prefix + '_%02d' % i,
+ )
feature_config[str(i)] = tf.tpu.experimental.embedding.FeatureConfig(
- table=table_config)
+ name=str(i),
+ table=table_config,
+ output_shape=[batch_size] if batch_size else None,
+ )
+ if max_ids_per_table:
+ max_ids_per_table_dict[str(table_name_prefix + '_%02d' % i)] = (
+ max_ids_per_table[i]
+ )
+ if max_unique_ids_per_table:
+ max_unique_ids_per_table_dict[str(table_name_prefix + '_%02d' % i)] = (
+ max_unique_ids_per_table[i]
+ )
+
+ if all((max_ids_per_chip_per_sample, max_ids_per_table,
+ max_unique_ids_per_table)):
+ sparsecore_config = tf.tpu.experimental.embedding.SparseCoreEmbeddingConfig(
+ disable_table_stacking=False,
+ max_ids_per_chip_per_sample=max_ids_per_chip_per_sample,
+ max_ids_per_table=max_ids_per_table_dict,
+ max_unique_ids_per_table=max_unique_ids_per_table_dict,
+ allow_id_dropping=allow_id_dropping,
+ initialize_tables_on_host=initialize_tables_on_host,
+ )
- return feature_config
+ return feature_config, sparsecore_config
class RankingTask(base_task.Task):
@@ -78,7 +158,7 @@ class RankingTask(base_task.Task):
def __init__(self,
params: config.Task,
- optimizer_config: config.OptimizationConfig,
+ trainer_config: config.TrainerConfig,
logging_dir: Optional[str] = None,
steps_per_execution: int = 1,
name: Optional[str] = None):
@@ -86,7 +166,7 @@ def __init__(self,
Args:
params: the RankingModel task configuration instance.
- optimizer_config: Optimizer configuration instance.
+ trainer_config: Trainer configuration instance.
logging_dir: a string pointing to where the model, summaries etc. will be
saved.
steps_per_execution: Int. Defaults to 1. The number of batches to run
@@ -94,18 +174,35 @@ def __init__(self,
name: the task name.
"""
super().__init__(params, logging_dir, name=name)
- self._optimizer_config = optimizer_config
+ self._trainer_config = trainer_config
+ self._optimizer_config = trainer_config.optimizer_config
self._steps_per_execution = steps_per_execution
def build_inputs(self, params, input_context=None):
"""Builds classification input."""
-
- dataset = data_pipeline.CriteoTsvReader(
- file_pattern=params.input_path,
- params=params,
- vocab_sizes=self.task_config.model.vocab_sizes,
- num_dense_features=self.task_config.model.num_dense_features,
- use_synthetic_data=self.task_config.use_synthetic_data)
+ if self.task_config.model.use_multi_hot:
+ if self.task_config.use_tf_record_reader:
+ dataset = data_pipeline_multi_hot.CriteoTFRecordReader(
+ file_pattern=params.input_path,
+ params=params,
+ vocab_sizes=self.task_config.model.vocab_sizes,
+ multi_hot_sizes=self.task_config.model.multi_hot_sizes,
+ num_dense_features=self.task_config.model.num_dense_features)
+ else:
+ dataset = data_pipeline_multi_hot.CriteoTsvReaderMultiHot(
+ file_pattern=params.input_path,
+ params=params,
+ vocab_sizes=self.task_config.model.vocab_sizes,
+ multi_hot_sizes=self.task_config.model.multi_hot_sizes,
+ num_dense_features=self.task_config.model.num_dense_features,
+ use_synthetic_data=self.task_config.use_synthetic_data)
+ else:
+ dataset = data_pipeline.CriteoTsvReader(
+ file_pattern=params.input_path,
+ params=params,
+ vocab_sizes=self.task_config.model.vocab_sizes,
+ num_dense_features=self.task_config.model.num_dense_features,
+ use_synthetic_data=self.task_config.use_synthetic_data)
return dataset(input_context)
@@ -115,7 +212,7 @@ def create_optimizer(cls, optimizer_config: config.OptimizationConfig,
"""See base class. Return None, optimizer is set in `build_model`."""
return None
- def build_model(self) -> tf.keras.Model:
+ def build_model(self) -> tf_keras.Model:
"""Creates Ranking model architecture and Optimizers.
The RankingModel uses different optimizers/learning rates for embedding
@@ -132,41 +229,89 @@ def build_model(self) -> tf.keras.Model:
warmup_steps=lr_config.warmup_steps,
decay_steps=lr_config.decay_steps,
decay_start_steps=lr_config.decay_start_steps)
-
- dense_optimizer = tf.keras.optimizers.Adam()
- embedding_optimizer = tf.keras.optimizers.get(
- self.optimizer_config.embedding_optimizer)
+ embedding_optimizer = tf_keras.optimizers.get(
+ self.optimizer_config.embedding_optimizer, use_legacy_optimizer=True)
embedding_optimizer.learning_rate = lr_callable
- feature_config = _get_tpu_embedding_feature_config(
- embedding_dim=self.task_config.model.embedding_dim,
- vocab_sizes=self.task_config.model.vocab_sizes)
+ dense_optimizer = tf_keras.optimizers.get(
+ self.optimizer_config.dense_optimizer, use_legacy_optimizer=True)
+ if self.optimizer_config.dense_optimizer == 'SGD':
+ dense_lr_config = self.optimizer_config.dense_sgd_config
+ dense_lr_callable = common.WarmUpAndPolyDecay(
+ batch_size=self.task_config.train_data.global_batch_size,
+ decay_exp=dense_lr_config.decay_exp,
+ learning_rate=dense_lr_config.learning_rate,
+ warmup_steps=dense_lr_config.warmup_steps,
+ decay_steps=dense_lr_config.decay_steps,
+ decay_start_steps=dense_lr_config.decay_start_steps)
+ dense_optimizer.learning_rate = dense_lr_callable
+
+ feature_config, sparse_core_embedding_config = (
+ _get_tpu_embedding_feature_config(
+ embedding_dim=self.task_config.model.embedding_dim,
+ vocab_sizes=self.task_config.model.vocab_sizes,
+ batch_size=self.task_config.train_data.global_batch_size
+ // tf.distribute.get_strategy().num_replicas_in_sync,
+ max_ids_per_chip_per_sample=self.task_config.model.max_ids_per_chip_per_sample,
+ max_ids_per_table=self.task_config.model.max_ids_per_table,
+ max_unique_ids_per_table=self.task_config.model.max_unique_ids_per_table,
+ allow_id_dropping=self.task_config.model.allow_id_dropping,
+ initialize_tables_on_host=self.task_config.model.initialize_tables_on_host,
+ )
+ )
- embedding_layer = tfrs.experimental.layers.embedding.PartialTPUEmbedding(
- feature_config=feature_config,
- optimizer=embedding_optimizer,
- size_threshold=self.task_config.model.size_threshold)
+ # to work around PartialTPUEmbedding issue in v5p and to enable multi hot
+ # features
+ if self.task_config.model.use_partial_tpu_embedding:
+ embedding_layer = tfrs.experimental.layers.embedding.PartialTPUEmbedding(
+ feature_config=feature_config,
+ optimizer=embedding_optimizer,
+ pipeline_execution_with_tensor_core=self.trainer_config.pipeline_sparse_and_dense_execution,
+ size_threshold=self.task_config.model.size_threshold,
+ )
+ else:
+ embedding_layer = tfrs.layers.embedding.tpu_embedding_layer.TPUEmbedding(
+ feature_config=feature_config,
+ optimizer=embedding_optimizer,
+ pipeline_execution_with_tensor_core=self.trainer_config.pipeline_sparse_and_dense_execution,
+ sparse_core_embedding_config=sparse_core_embedding_config,
+ )
if self.task_config.model.interaction == 'dot':
feature_interaction = tfrs.layers.feature_interaction.DotInteraction(
skip_gather=True)
elif self.task_config.model.interaction == 'cross':
- feature_interaction = tf.keras.Sequential([
- tf.keras.layers.Concatenate(),
+ feature_interaction = tf_keras.Sequential([
+ tf_keras.layers.Concatenate(),
tfrs.layers.feature_interaction.Cross()
])
+ elif self.task_config.model.interaction == 'multi_layer_dcn':
+ feature_interaction = tf_keras.Sequential([
+ tf_keras.layers.Concatenate(),
+ tfrs.layers.feature_interaction.MultiLayerDCN(
+ projection_dim=self.task_config.model.dcn_low_rank_dim,
+ num_layers=self.task_config.model.dcn_num_layers,
+ use_bias=self.task_config.model.dcn_use_bias,
+ kernel_initializer=self.task_config.model.dcn_kernel_initializer,
+ bias_initializer=self.task_config.model.dcn_bias_initializer,
+ ),
+ ])
else:
raise ValueError(
- f'params.task.model.interaction {self.task_config.model.interaction} '
- f'is not supported it must be either \'dot\' or \'cross\'.')
+ f' {self.task_config.model.interaction} is not supported it must be'
+ " either 'dot' or 'cross' or 'multi_layer_dcn'."
+ )
model = tfrs.experimental.models.Ranking(
embedding_layer=embedding_layer,
bottom_stack=tfrs.layers.blocks.MLP(
- units=self.task_config.model.bottom_mlp, final_activation='relu'),
+ units=self.task_config.model.bottom_mlp, final_activation='relu'
+ ),
feature_interaction=feature_interaction,
top_stack=tfrs.layers.blocks.MLP(
- units=self.task_config.model.top_mlp, final_activation='sigmoid'),
+ units=self.task_config.model.top_mlp, final_activation='sigmoid'
+ ),
+ concat_dense=self.task_config.model.concat_dense,
)
optimizer = tfrs.experimental.optimizers.CompositeOptimizer([
(embedding_optimizer, lambda: model.embedding_trainable_variables),
@@ -179,9 +324,9 @@ def build_model(self) -> tf.keras.Model:
def train_step(
self,
inputs: Dict[str, tf.Tensor],
- model: tf.keras.Model,
- optimizer: tf.keras.optimizers.Optimizer,
- metrics: Optional[List[tf.keras.metrics.Metric]] = None) -> tf.Tensor:
+ model: tf_keras.Model,
+ optimizer: tf_keras.optimizers.Optimizer,
+ metrics: Optional[List[tf_keras.metrics.Metric]] = None) -> tf.Tensor:
"""See base class."""
# All metrics need to be passed through the RankingModel.
assert metrics == model.metrics
@@ -190,15 +335,17 @@ def train_step(
def validation_step(
self,
inputs: Dict[str, tf.Tensor],
- model: tf.keras.Model,
- metrics: Optional[List[tf.keras.metrics.Metric]] = None) -> tf.Tensor:
+ model: tf_keras.Model,
+ metrics: Optional[List[tf_keras.metrics.Metric]] = None) -> tf.Tensor:
"""See base class."""
# All metrics need to be passed through the RankingModel.
assert metrics == model.metrics
return model.test_step(inputs)
+ @property
+ def trainer_config(self) -> config.TrainerConfig:
+ return self._trainer_config
+
@property
def optimizer_config(self) -> config.OptimizationConfig:
return self._optimizer_config
-
-
diff --git a/official/recommendation/ranking/task_test.py b/official/recommendation/ranking/task_test.py
index 426f468d217..9b66f86b8d1 100644
--- a/official/recommendation/ranking/task_test.py
+++ b/official/recommendation/ranking/task_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,20 +15,28 @@
"""Unit tests for task."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import exp_factory
from official.recommendation.ranking import task
from official.recommendation.ranking.data import data_pipeline
+from official.recommendation.ranking.data import data_pipeline_multi_hot
class TaskTest(parameterized.TestCase, tf.test.TestCase):
- @parameterized.parameters(('dlrm_criteo', True),
- ('dlrm_criteo', False),
- ('dcn_criteo', True),
- ('dcn_criteo', False))
- def test_task(self, config_name, is_training):
+ @parameterized.parameters(('dlrm_criteo', True, False),
+ ('dlrm_criteo', False, False),
+ ('dcn_criteo', True, False),
+ ('dcn_criteo', False, False),
+ ('dlrm_criteo', True, True),
+ ('dlrm_criteo', False, True),
+ ('dcn_criteo', True, True),
+ ('dcn_criteo', False, True),
+ ('dlrm_dcn_v2_criteo', True, True),
+ ('dlrm_dcn_v2_criteo', False, True),
+ )
+ def test_task(self, config_name, is_training, use_multi_hot):
params = exp_factory.get_exp_config(config_name)
params.task.train_data.global_batch_size = 16
@@ -40,12 +48,18 @@ def test_task(self, config_name, is_training):
params.task.model.num_dense_features = 5
ranking_task = task.RankingTask(params.task,
- params.trainer.optimizer_config)
+ params.trainer)
- if is_training:
- dataset = data_pipeline.train_input_fn(params.task)
+ if use_multi_hot:
+ if is_training:
+ dataset = data_pipeline_multi_hot.train_input_fn(params.task)
+ else:
+ dataset = data_pipeline_multi_hot.eval_input_fn(params.task)
else:
- dataset = data_pipeline.eval_input_fn(params.task)
+ if is_training:
+ dataset = data_pipeline.train_input_fn(params.task)
+ else:
+ dataset = data_pipeline.eval_input_fn(params.task)
iterator = iter(dataset(ctx=None))
model = ranking_task.build_model()
diff --git a/official/recommendation/ranking/train.py b/official/recommendation/ranking/train.py
index 5ae322a71e6..bfb2251e564 100644
--- a/official/recommendation/ranking/train.py
+++ b/official/recommendation/ranking/train.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -20,7 +20,7 @@
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.common import distribute_utils
from official.core import base_trainer
@@ -73,7 +73,7 @@ def main(_) -> None:
task = RankingTask(
params=params.task,
- optimizer_config=params.trainer.optimizer_config,
+ trainer_config=params.trainer,
logging_dir=model_dir,
steps_per_execution=params.trainer.steps_per_loop,
name='RankingTask')
@@ -150,7 +150,7 @@ def get_dataset_fn(params):
callbacks = [checkpoint_callback, time_callback]
if enable_tensorboard:
- tensorboard_callback = tf.keras.callbacks.TensorBoard(
+ tensorboard_callback = tf_keras.callbacks.TensorBoard(
log_dir=model_dir,
update_freq=min(1000, params.trainer.validation_interval),
profile_batch=FLAGS.profile_steps)
diff --git a/official/recommendation/ranking/train_test.py b/official/recommendation/ranking/train_test.py
index 81d9f718d97..ac7d9e97f08 100644
--- a/official/recommendation/ranking/train_test.py
+++ b/official/recommendation/ranking/train_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,7 +18,7 @@
import os
from absl import flags
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.recommendation.ranking import common
from official.recommendation.ranking import train
@@ -27,9 +27,16 @@
def _get_params_override(vocab_sizes,
+ multi_hot_sizes=None,
+ use_multi_hot=False,
interaction='dot',
use_orbit=True,
- strategy='mirrored'):
+ strategy='mirrored',
+ concat_dense=True,
+ dcn_num_layers=3,
+ dcn_low_rank_dim=64,
+ use_partial_tpu_embedding=True,
+ ):
# Update `data_dir` if `synthetic_data=False`.
data_dir = ''
@@ -43,6 +50,12 @@ def _get_params_override(vocab_sizes,
'embedding_dim': [8] * len(vocab_sizes),
'bottom_mlp': [64, 32, 8],
'interaction': interaction,
+ 'concat_dense': concat_dense,
+ 'dcn_num_layers': dcn_num_layers,
+ 'dcn_low_rank_dim': dcn_low_rank_dim,
+ 'use_multi_hot': use_multi_hot,
+ 'use_partial_tpu_embedding': use_partial_tpu_embedding,
+ 'multi_hot_sizes': multi_hot_sizes,
},
'train_data': {
'input_path': os.path.join(data_dir, 'train/*'),
@@ -79,59 +92,220 @@ def tearDown(self):
super().tearDown()
@parameterized.named_parameters(
- ('DlrmOneDeviceCTL', 'one_device', 'dot', True),
- ('DlrmOneDevice', 'one_device', 'dot', False),
- ('DcnOneDeviceCTL', 'one_device', 'cross', True),
- ('DcnOneDevice', 'one_device', 'cross', False),
- ('DlrmTPUCTL', 'tpu', 'dot', True),
- ('DlrmTPU', 'tpu', 'dot', False),
- ('DcnTPUCTL', 'tpu', 'cross', True),
- ('DcnTPU', 'tpu', 'cross', False),
- ('DlrmMirroredCTL', 'Mirrored', 'dot', True),
- ('DlrmMirrored', 'Mirrored', 'dot', False),
- ('DcnMirroredCTL', 'Mirrored', 'cross', True),
- ('DcnMirrored', 'Mirrored', 'cross', False),
+ ('DlrmOneDeviceCTL', 'one_device', 'dot', True, True, 3, 64, False, True),
+ ('DlrmOneDevice', 'one_device', 'dot', False, True, 3, 64, False, True),
+ (
+ 'DcnOneDeviceCTL',
+ 'one_device',
+ 'cross',
+ True,
+ True,
+ 3,
+ 64,
+ False,
+ True,
+ ),
+ ('DcnOneDevice', 'one_device', 'cross', False, True, 3, 64, False, True),
+ (
+ 'DlrmDcnV2OneDeviceCTL',
+ 'one_device',
+ 'multi_layer_dcn',
+ True,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ (
+ 'DlrmDcnV2OneDevice',
+ 'one_device',
+ 'multi_layer_dcn',
+ False,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ ('DlrmTPUCTL', 'tpu', 'dot', True, True, 3, 64, False, True),
+ ('DlrmTPU', 'tpu', 'dot', False, True, 3, 64, False, True),
+ ('DcnTPUCTL', 'tpu', 'cross', True, True, 3, 64, False, True),
+ ('DcnTPU', 'tpu', 'cross', False, True, 3, 64, False, True),
+ (
+ 'DlrmDcnV2TPUCTL',
+ 'tpu',
+ 'multi_layer_dcn',
+ True,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ (
+ 'DlrmDcnV2TPU',
+ 'tpu',
+ 'multi_layer_dcn',
+ False,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ ('DlrmMirroredCTL', 'Mirrored', 'dot', True, True, 3, 64, False, True),
+ ('DlrmMirrored', 'Mirrored', 'dot', False, True, 3, 64, False, True),
+ ('DcnMirroredCTL', 'Mirrored', 'cross', True, True, 3, 64, False, True),
+ ('DcnMirrored', 'Mirrored', 'cross', False, True, 3, 64, False, True),
+ (
+ 'DlrmDcnV2MirroredCTL',
+ 'Mirrored',
+ 'multi_layer_dcn',
+ True,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ (
+ 'DlrmDcnV2Mirrored',
+ 'Mirrored',
+ 'multi_layer_dcn',
+ False,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
)
- def testTrainEval(self, strategy, interaction, use_orbit=True):
+ def testTrainEval(
+ self,
+ strategy,
+ interaction,
+ use_orbit=True,
+ concat_dense=True,
+ dcn_num_layers=3,
+ dcn_low_rank_dim=64,
+ use_multi_hot=False,
+ use_partial_tpu_embedding=True,
+ ):
# Set up simple trainer with synthetic data.
# By default the mode must be `train_and_eval`.
self.assertEqual(FLAGS.mode, 'train_and_eval')
vocab_sizes = [40, 12, 11, 13]
+ multi_hot_sizes = [1, 2, 3, 1]
+
+ FLAGS.params_override = _get_params_override(
+ vocab_sizes=vocab_sizes,
+ multi_hot_sizes=multi_hot_sizes,
+ use_multi_hot=use_multi_hot,
+ interaction=interaction,
+ use_orbit=use_orbit,
+ strategy=strategy,
+ concat_dense=concat_dense,
+ dcn_num_layers=dcn_num_layers,
+ dcn_low_rank_dim=dcn_low_rank_dim,
+ use_partial_tpu_embedding=use_partial_tpu_embedding,
+ )
- FLAGS.params_override = _get_params_override(vocab_sizes=vocab_sizes,
- interaction=interaction,
- use_orbit=use_orbit,
- strategy=strategy)
train.main('unused_args')
self.assertNotEmpty(
- tf.io.gfile.glob(os.path.join(self._model_dir, 'params.yaml')))
+ tf.io.gfile.glob(os.path.join(self._model_dir, 'params.yaml'))
+ )
@parameterized.named_parameters(
- ('DlrmTPUCTL', 'tpu', 'dot', True),
- ('DlrmTPU', 'tpu', 'dot', False),
- ('DcnTPUCTL', 'tpu', 'cross', True),
- ('DcnTPU', 'tpu', 'cross', False),
- ('DlrmMirroredCTL', 'Mirrored', 'dot', True),
- ('DlrmMirrored', 'Mirrored', 'dot', False),
- ('DcnMirroredCTL', 'Mirrored', 'cross', True),
- ('DcnMirrored', 'Mirrored', 'cross', False),
+ ('DlrmTPUCTL', 'tpu', 'dot', True, True, 3, 64, False, True),
+ ('DlrmTPU', 'tpu', 'dot', False, True, 3, 64, False, True),
+ ('DcnTPUCTL', 'tpu', 'cross', True, True, 3, 64, False, True),
+ ('DcnTPU', 'tpu', 'cross', False, True, 3, 64, False, True),
+ (
+ 'DlrmDcnV2TPUCTL',
+ 'tpu',
+ 'multi_layer_dcn',
+ True,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ (
+ 'DlrmDcnV2TPU',
+ 'tpu',
+ 'multi_layer_dcn',
+ False,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ ('DlrmMirroredCTL', 'Mirrored', 'dot', True, True, 3, 64, False, True),
+ ('DlrmMirrored', 'Mirrored', 'dot', False, True, 3, 64, False, True),
+ ('DcnMirroredCTL', 'Mirrored', 'cross', True, True, 3, 64, False, True),
+ ('DcnMirrored', 'Mirrored', 'cross', False, True, 3, 64, False, True),
+ (
+ 'DlrmDcnV2MirroredCTL',
+ 'Mirrored',
+ 'multi_layer_dcn',
+ True,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
+ (
+ 'DlrmDcnV2Mirrored',
+ 'Mirrored',
+ 'multi_layer_dcn',
+ False,
+ False,
+ 3,
+ 64,
+ True,
+ False,
+ ),
)
- def testTrainThenEval(self, strategy, interaction, use_orbit=True):
+ def testTrainThenEval(
+ self,
+ strategy,
+ interaction,
+ use_orbit=True,
+ concat_dense=True,
+ dcn_num_layers=3,
+ dcn_low_rank_dim=64,
+ use_multi_hot=False,
+ use_partial_tpu_embedding=True,
+ ):
# Set up simple trainer with synthetic data.
vocab_sizes = [40, 12, 11, 13]
+ multi_hot_sizes = [1, 2, 3, 1]
- FLAGS.params_override = _get_params_override(vocab_sizes=vocab_sizes,
- interaction=interaction,
- use_orbit=use_orbit,
- strategy=strategy)
+ FLAGS.params_override = _get_params_override(
+ vocab_sizes=vocab_sizes,
+ multi_hot_sizes=multi_hot_sizes,
+ interaction=interaction,
+ use_orbit=use_orbit,
+ strategy=strategy,
+ concat_dense=concat_dense,
+ dcn_num_layers=dcn_num_layers,
+ dcn_low_rank_dim=dcn_low_rank_dim,
+ use_multi_hot=use_multi_hot,
+ use_partial_tpu_embedding=use_partial_tpu_embedding,
+ )
default_mode = FLAGS.mode
# Training.
FLAGS.mode = 'train'
train.main('unused_args')
self.assertNotEmpty(
- tf.io.gfile.glob(os.path.join(self._model_dir, 'params.yaml')))
+ tf.io.gfile.glob(os.path.join(self._model_dir, 'params.yaml'))
+ )
# Evaluation.
FLAGS.mode = 'eval'
diff --git a/official/recommendation/stat_utils.py b/official/recommendation/stat_utils.py
index a565ce9df26..c534d5579d1 100644
--- a/official/recommendation/stat_utils.py
+++ b/official/recommendation/stat_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/recommendation/uplift/README.md b/official/recommendation/uplift/README.md
new file mode 100644
index 00000000000..55a7838bfac
--- /dev/null
+++ b/official/recommendation/uplift/README.md
@@ -0,0 +1,106 @@
+# Uplift Modeling Library
+
+Uplift modeling is a predictive modeling technique that models the incremental
+impact of a treatment on an individual. This repository contains a suite of
+tools to build, train, evaluate and export an uplift model in TensorFlow 2. All
+components are built in a modular fashion and integrate well with other TF-based
+systems such as TFX, Tensorflow Transform (TFT) and Tensorflow Model Analysis
+(TFMA). We use Keras as our modeling framework.
+
+The library is divided into the following key directories:
+
+- **layers**: Keras layers related to uplift modeling.
+- **models**: Keras models that contain all the necessary logic to train,
+ evaluate and do inference on an uplift model.
+- **metrics**: Keras metrics used in uplift modeling.
+- **losses**: Keras losses used in uplift modeling.
+
+See the [uplift_modeling_intro.ipynb](uplift_modeling_intro.ipynb) for an
+introduction to uplift modeling and how to use the library to build, train and
+evaluate an uplift model on large scale data.
+
+## Two Tower Uplift Model
+The initial release focuses on the family of models that follow a two tower
+uplift network architecture. The architecture draws inspiration from several
+related works in the field of treatment effect estimation [^1] [^2] [^3] [^4]
+[^5].
+
+
+
+
+
+- **Inputs**: a mapping from feature names to feature tensors. The tensors can
+be of different types, eg `tf.Tensor`, `tf.SparseTensor` and `tf.RaggedTensor`.
+- **Backbone**: a trainable network that computes an embedding from the inputs
+shared between the control and treatment arms.
+- **Control & Treatment Feature Encoders**: trainable networks that compute
+embeddings from control and treatment speficic features.
+- **Control & Treatment Feature Combiners**: methods to combine the backbone's
+shared embedding with the control/treatment specific embeddings.
+- **Control Tower**: trainable network with zero or more hidden layers that
+learns from control examples only.
+- **Treatment Tower**: trainable network with zero or more hidden layers that
+learns from treatment examples only.
+- **Logits Head**: computes control and treatment logits. At training time, the
+gradient flows from the control logits for control examples and from the
+treatment logits for the treatment examples.
+
+Example usage:
+
+```python
+# Create a two tower uplift model.
+uplift_network = two_tower_uplift_network.TwoTowerUpliftNetwork(
+ backbone=encoders.concat_features.ConcatFeatures(
+ feature_names=["shared_feature_1", "shared_feature_2"]
+ ),
+ treatment_feature_encoder=encoders.concat_features.ConcatFeatures(
+ feature_names=["treatment_feature_1", "treatment_feature_2"]
+ ),
+ treatment_input_combiner=tf.keras.layers.Concatenate(),
+ treatment_tower=tf.keras.Sequential([
+ tf.keras.layers.Dense(64, activation="relu"),
+ tf.keras.layers.Dropout(0.1),
+ ]),
+ control_tower=tf.keras.Sequential([
+ tf.keras.layers.Dense(128, activation="relu"),
+ tf.keras.layers.Dropout(0.1),
+ tf.keras.layers.Dense(32, activation="relu")
+ ]),
+ logits_head=two_tower_logits_head.TwoTowerLogitsHead(
+ control_head=tf.keras.layers.Dense(1),
+ treatment_head=tf.keras.layers.Dense(1),
+ ),
+)
+model = two_tower_uplift_model.TwoTowerUpliftModel(
+ treatment_indicator_feature_name="is_treatment",
+ uplift_network=uplift_network,
+)
+
+# Compile and train the model.
+model.compile(
+ optimizer=tf.keras.optimizers.Adagrad(learning_rate=0.05),
+ loss=true_logits_loss.TrueLogitsLoss(tf.keras.losses.mean_squared_error),
+ metrics=[
+ treatment_fraction.TreatmentFraction(),
+ uplift_mean.UpliftMean(),
+ label_mean.LabelMean(),
+ label_variance.LabelVariance(),
+ ]
+)
+model.fit(dataset, epochs=10)
+```
+
+[^1]: Johansson, Fredrik, Uri Shalit, and David Sontag. "Learning
+ representations for counterfactual inference." *International conference
+ on machine learning*. PMLR, 2016.
+[^2]: Shalit, Uri, Fredrik D. Johansson, and David Sontag. "Estimating
+ individual treatment effect: generalization bounds and algorithms."
+ *International conference on machine learning*. PMLR, 2017.
+[^3]: Johansson, Fredrik D., et al. "Learning weighted representations for
+ generalization across designs." *arXiv preprint arXiv:1802.08598* (2018).
+[^4]: Hassanpour, Negar and Russell Greiner. “CounterFactual Regression with
+ Importance Sampling Weights.” *International Joint Conference on
+ Artificial Intelligence* (2019).
+[^5]: Hassanpour, Negar and Russell Greiner. “Learning Disentangled
+ Representations for CounterFactual Regression.” *International Conference
+ on Learning Representations* (2020).
diff --git a/official/recommendation/uplift/__init__.py b/official/recommendation/uplift/__init__.py
new file mode 100644
index 00000000000..0e787ed00c5
--- /dev/null
+++ b/official/recommendation/uplift/__init__.py
@@ -0,0 +1,23 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Uplift package definition."""
+
+from official.recommendation.uplift import keys
+from official.recommendation.uplift import layers
+from official.recommendation.uplift import losses
+from official.recommendation.uplift import metrics
+from official.recommendation.uplift import models
+from official.recommendation.uplift import types
+from official.recommendation.uplift import utils
diff --git a/official/recommendation/uplift/keras_test_case.py b/official/recommendation/uplift/keras_test_case.py
new file mode 100644
index 00000000000..fab9cabd78a
--- /dev/null
+++ b/official/recommendation/uplift/keras_test_case.py
@@ -0,0 +1,182 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Testing framework for Keras extending tf.TestCase.
+
+Keras layers have several nuances. Usually, we expect:
+
+1. A Keras layer to be stable. If no learning happens, they should return
+ the same result. See assertLayerStable.
+
+2. A Keras layer to be savable. We should be able to save, then load
+ a Keras layer. It should have the same variables, and return the same result.
+ See assertLayerSavable.
+
+Various nuances, such as defining a correct get_config function that doesn't
+forget any variables, and making sure that any sub-layers are referenced as
+fields (or in fields that are dictionaries), are necessary for these tests to
+succeed.
+"""
+
+import json
+from typing import Any, Mapping, Sequence
+
+import tensorflow as tf, tf_keras
+
+
+# pylint: disable=invalid-name
+class KerasTestCase(tf.test.TestCase):
+ """Adds methods for testing Keras layers."""
+
+ def assertNestedEqual(self, data1, data2):
+ """Asserts that both inputs have identical structure and contents."""
+ tf.nest.assert_same_structure(data1, data2)
+ for a, b in zip(tf.nest.flatten(data1), tf.nest.flatten(data2)):
+ if isinstance(a, tf.SparseTensor):
+ self.assertAllEqual(a.indices, b.indices)
+ self.assertAllEqual(a.values, b.values)
+ self.assertAllEqual(a.dense_shape, b.dense_shape)
+ else:
+ self.assertAllEqual(a, b)
+
+ def assertLayerStable(self, inputs, layer, **kwargs):
+ """Layer called twice with the same inputs returns the same result.
+
+ Layer must return a result appropriate for assertAllEqual.
+
+ Args:
+ inputs: inputs to the layer.
+ layer: layer to test.
+ **kwargs: auxiliary inputs passed to the layer.
+ """
+ output1 = layer(inputs, **kwargs)
+ output2 = layer(inputs, **kwargs)
+ self.assertNestedEqual(output1, output2)
+
+ def toKerasInputs(self, inputs): # pylint:disable=[invalid-name]
+ """Generate tf_keras.Input for inputs.
+
+ Args:
+ inputs: a StructuredTensor, Tensor, RecordTensor, RaggedTensor, or nested
+ structures of these.
+
+ Returns:
+ tf_keras.Input representing the input.
+ """
+
+ def map_to_keras_input(x):
+ if isinstance(x, tf.Tensor):
+ return tf_keras.Input(x.shape[1:], dtype=x.dtype)
+ ts = tf.type_spec_from_value(x)
+ return tf_keras.Input(type_spec=ts)
+
+ return tf.nest.map_structure(map_to_keras_input, inputs)
+
+ def assertLayerSavable(
+ self,
+ inputs,
+ layer,
+ keras_inputs=None,
+ custom_objects=None,
+ save_format="tf",
+ **kwargs
+ ):
+ """Layer can be saved and loaded in a model.
+
+ Args:
+ inputs: an input to the layer.
+ layer: a layer to save.
+ keras_inputs: if inputs._type_spec won't create a tf_keras.Input.
+ custom_objects: passed to load_model.
+ save_format: save_format ("tf" or "h5")
+ **kwargs: auxiliary inputs passed to the layer.
+ """
+
+ def _make_model(inputs, layer, keras_inputs):
+ if keras_inputs is None:
+ # TODO(martinz): This is not a generic solution.
+ keras_inputs = self.toKerasInputs(inputs)
+ keras_outputs = layer(keras_inputs, **kwargs)
+ return tf_keras.Model(keras_inputs, keras_outputs)
+
+ model = _make_model(inputs, layer, keras_inputs)
+ self.assertModelSavable(
+ inputs, model, custom_objects=custom_objects, save_format=save_format
+ )
+
+ def assertModelSavable(
+ self, inputs, model, custom_objects=None, save_format="tf"
+ ):
+ """Model can be saved and loaded.
+
+ Args:
+ inputs: an input to the layer.
+ model: a model to save.
+ custom_objects: passed to load_model.
+ save_format: save_format ("tf" or "h5")
+ """
+ if custom_objects is None:
+ custom_objects = {}
+
+ src_output = model(inputs)
+ model_path = self.get_temp_dir() + "/tmp_model"
+ tf_keras.models.save_model(model, model_path, save_format=save_format)
+ reloaded_model = tf_keras.models.load_model(
+ model_path, custom_objects=custom_objects
+ )
+ self.assertEqual(
+ len(model.trainable_variables), len(reloaded_model.trainable_variables)
+ )
+ for src_v, loaded_v in zip(
+ model.trainable_variables, reloaded_model.trainable_variables
+ ):
+ self.assertAllEqual(src_v, loaded_v)
+
+ loaded_output = reloaded_model(inputs)
+
+ self.assertNestedEqual(src_output, loaded_output)
+
+ def assertLayerConfigurable(
+ self, layer: tf_keras.layers.Layer, serializable: bool = True, **kwargs
+ ):
+ """Layer can be reconstructed using get_config and from_config.
+
+ Args:
+ layer: layer with get_config and from_config methods.
+ serializable: boolean to indicate if the config should be tested for json
+ serializability.
+ **kwargs: keyword arguments to pass to layer's `__call__` method. These
+ are used for testing the correctness of the call output. If no keywords
+ are passed then this part of the test will not be executed.
+ """
+ config = layer.get_config()
+ from_config_layer = layer.__class__.from_config(
+ config if not serializable else json.loads(json.dumps(config))
+ )
+ self.assertDictEqual(config, from_config_layer.get_config())
+
+ if kwargs:
+ self.assertNestedEqual(layer(**kwargs), from_config_layer(**kwargs))
+
+
+def layer_dict_from_classes(classes: Sequence[Any]) -> Mapping[str, Any]:
+ """Construct a dictionary for custom_objects while saving keras layers.
+
+ Args:
+ classes: a sequence of layer classes.
+
+ Returns:
+ A dictionary of layer classes, keyed by name.
+ """
+ return {x.__name__: x for x in classes}
diff --git a/official/recommendation/uplift/keys.py b/official/recommendation/uplift/keys.py
new file mode 100644
index 00000000000..4d3379b36e0
--- /dev/null
+++ b/official/recommendation/uplift/keys.py
@@ -0,0 +1,28 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Defines enum keys used by the keras uplift modeling library."""
+
+import enum
+
+
+class TwoTowerOutputKeys(str, enum.Enum):
+ """Keys for training and inference output tensors."""
+
+ CONTROL_PREDICTIONS = "control_predictions"
+ TREATMENT_PREDICTIONS = "treatment_predictions"
+ UPLIFT_PREDICTIONS = "uplift_predictions"
+ IS_TREATMENT = "is_treatment"
+ TRUE_LOGITS = "true_logits"
+ TRUE_PREDICTIONS = "true_predictions"
diff --git a/official/recommendation/uplift/layers/__init__.py b/official/recommendation/uplift/layers/__init__.py
new file mode 100644
index 00000000000..e1e9de3c83c
--- /dev/null
+++ b/official/recommendation/uplift/layers/__init__.py
@@ -0,0 +1,19 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Layers package definitions."""
+
+from official.recommendation.uplift.layers import encoders
+from official.recommendation.uplift.layers import heads
+from official.recommendation.uplift.layers import uplift_networks
diff --git a/official/recommendation/uplift/layers/encoders/__init__.py b/official/recommendation/uplift/layers/encoders/__init__.py
new file mode 100644
index 00000000000..10873faebfc
--- /dev/null
+++ b/official/recommendation/uplift/layers/encoders/__init__.py
@@ -0,0 +1,17 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Uplift heads package definitions."""
+
+from official.recommendation.uplift.layers.encoders import concat_features
diff --git a/official/recommendation/uplift/layers/encoders/concat_features.py b/official/recommendation/uplift/layers/encoders/concat_features.py
new file mode 100644
index 00000000000..9519a57f6fd
--- /dev/null
+++ b/official/recommendation/uplift/layers/encoders/concat_features.py
@@ -0,0 +1,113 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Defines an encoder for concatenating input features into a single tensor."""
+
+from typing import Mapping, Sequence
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class ConcatFeatures(tf_keras.layers.Layer):
+ """Concatenates features into a single dense tensor.
+
+ Takes a dictionary of feature tensors as input and concatenates the specified
+ features into a single tensor. The tensors are concatenated along their last
+ axis. Sparse and ragged tensors are converted to dense tensors before being
+ concatenated.
+ """
+
+ def __init__(self, feature_names: Sequence[str], **kwargs):
+ """Initializes a feature concatenation encoder.
+
+ Args:
+ feature_names: names of the input features to concatenate together.
+ **kwargs: base layer keyword arguments.
+ """
+ super().__init__(**kwargs)
+ self._feature_names = feature_names
+
+ # Validate feature names.
+ if not feature_names:
+ raise ValueError(
+ "feature_names must be a non-empty list of strings but got"
+ f" {feature_names} instead."
+ )
+ if not all(isinstance(name, str) for name in feature_names):
+ raise TypeError(
+ "feature_names must be a list of strings, but got types"
+ f" {list(map(type, feature_names))}"
+ )
+
+ def build(self, input_shapes: Mapping[str, tf.TensorShape]) -> None:
+ missing_features = set(self._feature_names) - input_shapes.keys()
+ if missing_features:
+ raise ValueError(f"Layer inputs is missing features: {missing_features}")
+
+ feature_shapes = {
+ feature_name: tensor_shape
+ for feature_name, tensor_shape in input_shapes.items()
+ if feature_name in self._feature_names
+ }
+
+ most_specific_shape = tf.TensorShape(None)
+ for feature_name, shape in feature_shapes.items():
+ if not isinstance(shape, tf.TensorShape):
+ raise TypeError(
+ f"Got unsupported tensor shape type for feature {feature_name}. The"
+ " feature tensor must be one of `tf.Tensor`, `tf.SparseTensor` or"
+ " `tf.RaggedTensor`, with a well defined tensor shape but got shape"
+ f" {shape} instead."
+ )
+
+ shape = shape[:-1]
+ if shape.is_subtype_of(most_specific_shape):
+ most_specific_shape = shape
+
+ elif not most_specific_shape.is_subtype_of(shape):
+ raise ValueError(
+ "All features from the feature_names set must be tensors with the"
+ " same shape except for the last dimension, but got features with"
+ f" incompatible shapes {feature_shapes}"
+ )
+
+ super().build(input_shapes)
+
+ def call(self, inputs: types.DictOfTensors) -> tf.Tensor:
+ features = []
+
+ for feature_name, feature in inputs.items():
+ if feature_name in self._feature_names:
+ if isinstance(feature, tf.Tensor):
+ features.append(feature)
+ elif isinstance(feature, tf.SparseTensor):
+ features.append(tf.sparse.to_dense(feature))
+ elif isinstance(feature, tf.RaggedTensor):
+ features.append(feature.to_tensor())
+ else:
+ raise TypeError(
+ f"Got unsupported tensor type for feature {feature_name}. The"
+ " feature tensor must be one of `tf.Tensor`, `tf.SparseTensor` or"
+ f" `tf.RaggedTensor`, but got {feature} instead."
+ )
+
+ return tf.concat(features, axis=-1)
+
+ def get_config(self):
+ config = super().get_config()
+ config.update({"feature_names": self._feature_names})
+ return config
diff --git a/official/recommendation/uplift/layers/encoders/concat_features_test.py b/official/recommendation/uplift/layers/encoders/concat_features_test.py
new file mode 100644
index 00000000000..d2291452dd0
--- /dev/null
+++ b/official/recommendation/uplift/layers/encoders/concat_features_test.py
@@ -0,0 +1,275 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for concat_feature_encoder."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift.layers.encoders import concat_features
+
+
+class ConcatFeaturesTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "single_dense",
+ "feature_names": ["feature"],
+ "inputs": {"feature": tf.ones((3, 1))},
+ "expected_output": tf.ones((3, 1)),
+ },
+ {
+ "testcase_name": "single_sparse",
+ "feature_names": ["feature"],
+ "inputs": {
+ "feature": tf.sparse.SparseTensor(
+ indices=[[0, 0], [1, 2]], values=[1, 2], dense_shape=[3, 4]
+ )
+ },
+ "expected_output": tf.constant(
+ [[1, 0, 0, 0], [0, 0, 2, 0], [0, 0, 0, 0]]
+ ),
+ },
+ {
+ "testcase_name": "single_ragged",
+ "feature_names": ["feature"],
+ "inputs": {"feature": tf.ragged.constant([[5, 7], [0, 3, 1], [6]])},
+ "expected_output": tf.constant([[5, 7, 0], [0, 3, 1], [6, 0, 0]]),
+ },
+ {
+ "testcase_name": "excess_features",
+ "feature_names": ["feature3"],
+ "inputs": {
+ "feature1": tf.ones((3, 1)),
+ "feature2": 2 * tf.ones((3, 1)),
+ "feature3": 3 * tf.ones((3, 1)),
+ },
+ "expected_output": 3 * tf.ones((3, 1)),
+ },
+ {
+ "testcase_name": "one_dimensional_features",
+ "feature_names": ["feature1", "feature2"],
+ "inputs": {
+ "feature1": tf.ones((1, 3)),
+ "feature2": 2 * tf.ones((1, 2)),
+ },
+ "expected_output": tf.constant([1, 1, 1, 2, 2], shape=(1, 5)),
+ },
+ {
+ "testcase_name": "mixed_features",
+ "feature_names": ["dense", "sparse", "ragged"],
+ "inputs": {
+ "dense": tf.constant([-1.4, 2.0], shape=(2, 1)),
+ "sparse": tf.sparse.SparseTensor(
+ indices=[[0, 1], [1, 0]],
+ values=[2.718, 3.14],
+ dense_shape=[2, 2],
+ ),
+ "ragged": tf.ragged.constant([[5, 7.77], [8]]),
+ "other_feature": tf.ones((2, 5)),
+ },
+ "expected_output": tf.constant(
+ [[-1.4, 0, 2.718, 5, 7.77], [2.0, 3.14, 0, 8, 0]]
+ ),
+ },
+ )
+ def test_layer_correctness(self, feature_names, inputs, expected_output):
+ layer = concat_features.ConcatFeatures(feature_names=feature_names)
+ self.assertAllClose(expected_output, layer(inputs))
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "none_dimensions",
+ "inputs": {
+ "x1": tf_keras.Input(shape=(2, None, 1, 3)),
+ "x2": tf_keras.Input(shape=(2, None, 1, None)),
+ },
+ "expected_shape": [None, 2, None, 1, None],
+ },
+ {
+ "testcase_name": "dense_sparse_ragged",
+ "inputs": {
+ "dense": tf_keras.Input(shape=(2, None, 1, 3), batch_size=10),
+ "sparse": tf_keras.Input(shape=(2, None, 1, 3), sparse=True),
+ "ragged": tf_keras.Input(shape=(2, None, None, 1), ragged=True),
+ },
+ "expected_shape": [10, 2, None, 1, None],
+ },
+ )
+ def test_layer_correctness_keras_inputs(self, inputs, expected_shape):
+ layer = concat_features.ConcatFeatures(feature_names=list(inputs.keys()))
+ output = layer(inputs)
+
+ KerasTensor = tf_keras.Input(shape=(1,)).__class__ # pylint: disable=invalid-name
+ self.assertIsInstance(output, KerasTensor)
+ self.assertEqual(tf.TensorShape(expected_shape), output.shape)
+
+ def test_layer_stability(self):
+ layer = concat_features.ConcatFeatures(
+ feature_names=["dense", "sparse", "ragged"]
+ )
+ inputs = {
+ "dense": tf.constant([-1.4, 2.0], shape=(2, 1)),
+ "sparse": tf.sparse.SparseTensor(
+ indices=[[0, 1], [1, 0]],
+ values=[2.718, 3.14],
+ dense_shape=[2, 2],
+ ),
+ "ragged": tf.ragged.constant([[5, 7.77], [8]]),
+ "other_feature": tf.ones((2, 5)),
+ }
+ self.assertLayerStable(inputs=inputs, layer=layer)
+
+ def test_layer_savable(self):
+ layer = concat_features.ConcatFeatures(
+ feature_names=["dense", "sparse", "ragged"]
+ )
+ inputs = {
+ "dense": tf.constant([-1.4, 2.0], shape=(2, 1)),
+ "sparse": tf.sparse.SparseTensor(
+ indices=[[0, 1], [1, 0]],
+ values=[2.718, 3.14],
+ dense_shape=[2, 2],
+ ),
+ "ragged": tf.ragged.constant([[5, 7.77], [8]]),
+ "other_feature": tf.ones((2, 5)),
+ }
+ self.assertLayerSavable(inputs=inputs, layer=layer)
+
+ def test_missing_input_features(self):
+ layer = concat_features.ConcatFeatures(feature_names=["feature"])
+
+ with self.assertRaisesRegex(
+ ValueError, "Layer inputs is missing features*"
+ ):
+ layer({"other_feature": tf.ones((3, 1))})
+
+ def test_unsupported_tensor_type(self):
+ class TestType(tf.experimental.ExtensionType):
+ tensor: tf.Tensor
+
+ layer = concat_features.ConcatFeatures(feature_names=["feature"])
+ with self.assertRaisesRegex(TypeError, "Got unsupported tensor shape type"):
+ layer({
+ "feature": TestType(tensor=tf.ones((3, 1))),
+ "other_feature": tf.ones((3, 1)),
+ })
+
+ def test_empty_feature_names_list(self):
+ with self.assertRaisesRegex(
+ ValueError, "feature_names must be a non-empty list"
+ ):
+ concat_features.ConcatFeatures(feature_names=[])
+
+ def test_non_string_feature_name(self):
+ with self.assertRaisesRegex(
+ TypeError, "feature_names must be a list of strings"
+ ):
+ concat_features.ConcatFeatures(feature_names=["x", 1])
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "different_shapes_dense",
+ "inputs": {
+ "x1": tf.ones((2, 4)),
+ "x2": tf.ones((1, 4)),
+ },
+ },
+ {
+ "testcase_name": "different_shapes_sparse",
+ "inputs": {
+ "x1": tf.ones((10, 4)),
+ "x2": tf.sparse.SparseTensor(
+ indices=[[0, 0], [1, 2]], values=[1, 2], dense_shape=[3, 4]
+ ),
+ },
+ },
+ {
+ "testcase_name": "different_shapes_ragged",
+ "inputs": {
+ "x1": tf.ones((2, 2, 2)),
+ "x2": tf.ragged.constant([[5, 7], [0, 3, 1], [6]]),
+ },
+ },
+ {
+ "testcase_name": "keras_input_batch_size",
+ "inputs": {
+ "x1": tf_keras.Input(shape=(2, 3), batch_size=10),
+ "x2": tf_keras.Input(shape=(2, 3), batch_size=4),
+ },
+ },
+ )
+ def test_shape_mismatch(self, inputs):
+ layer = concat_features.ConcatFeatures(feature_names=list(inputs.keys()))
+ with self.assertRaisesRegex(
+ ValueError,
+ (
+ "All features from the feature_names set must be tensors with the"
+ " same shape except for the last dimension"
+ ),
+ ):
+ layer(inputs)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "different_ranks_dense",
+ "inputs": {
+ "x1": tf.ones((2, 4)),
+ "x2": tf.ones((2, 4, 6)),
+ },
+ },
+ {
+ "testcase_name": "different_ranks_sparse",
+ "inputs": {
+ "x1": tf.ones((3, 4, 1)),
+ "x2": tf.sparse.SparseTensor(
+ indices=[[0, 0], [1, 2]], values=[1, 2], dense_shape=[3, 4]
+ ),
+ },
+ },
+ {
+ "testcase_name": "different_ranks_ragged",
+ "inputs": {
+ "x1": tf.ones((2, 2, 2, 2)),
+ "x2": tf.ragged.constant([[5, 7], [0, 3, 1], [6]]),
+ },
+ },
+ {
+ "testcase_name": "keras_input",
+ "inputs": {
+ "x1": tf_keras.Input(shape=(2, 3, 4), batch_size=4),
+ "x2": tf_keras.Input(shape=(2, 3), batch_size=4),
+ },
+ },
+ )
+ def test_rank_mismatch(self, inputs):
+ layer = concat_features.ConcatFeatures(feature_names=list(inputs.keys()))
+ with self.assertRaisesRegex(
+ ValueError,
+ (
+ "All features from the feature_names set must be tensors with the"
+ " same shape except for the last dimension"
+ ),
+ ):
+ layer(inputs)
+
+ def test_layer_config(self):
+ layer = concat_features.ConcatFeatures(
+ feature_names=["feature1", "feature2"], name="encoder", dtype=tf.float64
+ )
+ self.assertLayerConfigurable(layer=layer, serializable=True)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/layers/heads/__init__.py b/official/recommendation/uplift/layers/heads/__init__.py
new file mode 100644
index 00000000000..67fe1dde0d4
--- /dev/null
+++ b/official/recommendation/uplift/layers/heads/__init__.py
@@ -0,0 +1,17 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Uplift heads package definitions."""
+
+from official.recommendation.uplift.layers.heads import two_tower_logits_head
diff --git a/official/recommendation/uplift/layers/heads/two_tower_logits_head.py b/official/recommendation/uplift/layers/heads/two_tower_logits_head.py
new file mode 100644
index 00000000000..d8b34b5c4fe
--- /dev/null
+++ b/official/recommendation/uplift/layers/heads/two_tower_logits_head.py
@@ -0,0 +1,188 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builds a TwoTowerLogitsHead layer."""
+
+from __future__ import annotations
+
+import dataclasses
+import enum
+from typing import Any
+
+import tensorflow as tf, tf_keras
+
+
+@enum.unique
+class LayeringMethod(str, enum.Enum):
+ """Layering method between the control and treatment towers."""
+
+ # No layering.
+ NONE = "none"
+
+ # The treatment logits are adjusted by the following function:
+ # treatment_logits += tf.stop_gradient(control_logits)
+ LOGIT_SUM = "logit_sum"
+
+ # The treatment embedding is adjusted by the following function:
+ # treatment_embedding += stop_gradient(control_embedding) * W
+ # Where "W" is a learnable weight matrix of shape CxT where C is the control
+ # embedding dimension and T is the treatment embedding dimension.
+ LINEAR_LAYERING = "linear_layering"
+
+
+@dataclasses.dataclass(frozen=True)
+class LinearLayeringConfig:
+ """Configuration for the linear layering method.
+
+ Attributes:
+ kernel_initializer: kernel initializer for the learnable weight matrix.
+ kernel_regularizer: kernel regularizer for the learnable weight matrix.
+ """
+
+ kernel_initializer: str = "glorot_uniform"
+ kernel_regularizer: str | None = None
+
+
+@dataclasses.dataclass(frozen=True)
+class LayeringConfig:
+ """Configuration for all layering methods.
+
+ Attributes:
+ layering_method: specifies what layering method to apply. Defaults to
+ `LayeringMethod.NONE`.
+ linear_layering_config: configuration for the linear layering method. Will
+ only be used if the `layering_method` is set to
+ `LayeringMethod.LINEAR_LAYERING`.
+ """
+
+ layering_method: LayeringMethod = LayeringMethod.NONE
+ linear_layering_config: LinearLayeringConfig | None = None
+
+ @classmethod
+ def from_dict(cls, config: dict[str, Any]) -> LayeringConfig:
+ linear_layering_config = (
+ LinearLayeringConfig(**config["linear_layering_config"])
+ if config["linear_layering_config"] is not None
+ else None
+ )
+ return cls(
+ layering_method=LayeringMethod(config["layering_method"]),
+ linear_layering_config=linear_layering_config,
+ )
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class TwoTowerLogitsHead(tf_keras.layers.Layer):
+ """Computes control and treatment logits from their respective embeddings.
+
+ Takes as input a tuple of control and treatment embeddings and computes
+ control and treatment logits.
+ """
+
+ def __init__(
+ self,
+ control_head: tf_keras.layers.Layer,
+ treatment_head: tf_keras.layers.Layer,
+ layering_config: LayeringConfig = LayeringConfig(),
+ **kwargs,
+ ):
+ """Initializes the instance.
+
+ The control and treatment heads must compute logits of the same shape.
+
+ Args:
+ control_head: computes control logits from the control embedding. Its
+ input and output is expected to be a dense tensor.
+ treatment_head: computes treatment logits from the treatment embedding.
+ Its input and output is expected to be a dense tensor.
+ layering_config: configuration for the layering method. Defaults to no
+ layering.
+ **kwargs: base layer keyword arguments.
+ """
+ super().__init__(**kwargs)
+
+ self._control_head = control_head
+ self._treatment_head = treatment_head
+ self._layering_config = layering_config
+
+ def build(self, input_shapes: tuple[tf.TensorShape, tf.TensorShape]):
+ if self._layering_config.layering_method == LayeringMethod.LINEAR_LAYERING:
+ if self._layering_config.linear_layering_config is None:
+ raise ValueError(
+ "The linear layering config cannot be `None` when using the linear"
+ " layering method."
+ )
+ # Build a learnable weight matrix that projects from the control embedding
+ # space to the treatment embedding space.
+ _, treatment_embedding_shape = input_shapes
+ self._linear_layering = tf_keras.layers.Dense(
+ units=treatment_embedding_shape[-1],
+ activation=None,
+ use_bias=True,
+ kernel_initializer=(
+ self._layering_config.linear_layering_config.kernel_initializer
+ ),
+ kernel_regularizer=(
+ self._layering_config.linear_layering_config.kernel_regularizer
+ ),
+ )
+ super().build(input_shapes)
+
+ def call(
+ self, inputs: tuple[tf.Tensor, tf.Tensor]
+ ) -> tuple[tf.Tensor, tf.Tensor]:
+ control_embedding, treatment_embedding = inputs
+
+ if self._layering_config.layering_method == LayeringMethod.LINEAR_LAYERING:
+ treatment_embedding += self._linear_layering(
+ tf.stop_gradient(control_embedding)
+ )
+
+ control_logits = self._control_head(control_embedding)
+ treatment_logits = self._treatment_head(treatment_embedding)
+
+ if control_logits.shape != treatment_logits.shape:
+ raise ValueError(
+ "The control logits and treatment logits computed by the control and"
+ " treatment heads must be tensors of the same shape, but got shape"
+ f" {control_logits.shape} for the control logits and shape"
+ f" {treatment_logits.shape} for the treatment logits."
+ )
+
+ if self._layering_config.layering_method == LayeringMethod.LOGIT_SUM:
+ treatment_logits += tf.stop_gradient(control_logits)
+
+ return control_logits, treatment_logits
+
+ def get_config(self) -> dict[str, Any]:
+ config = super().get_config()
+ config["layering_config"] = dataclasses.asdict(self._layering_config)
+
+ for layer_name, layer in (
+ ("control_head", self._control_head),
+ ("treatment_head", self._treatment_head),
+ ):
+ config[layer_name] = tf_keras.utils.serialize_keras_object(layer)
+
+ return config
+
+ @classmethod
+ def from_config(cls, config: dict[str, Any]) -> TwoTowerLogitsHead:
+ config["layering_config"] = LayeringConfig.from_dict(
+ config["layering_config"]
+ )
+ for layer_name in ("control_head", "treatment_head"):
+ config[layer_name] = tf_keras.layers.deserialize(config[layer_name])
+
+ return cls(**config)
diff --git a/official/recommendation/uplift/layers/heads/two_tower_logits_head_test.py b/official/recommendation/uplift/layers/heads/two_tower_logits_head_test.py
new file mode 100644
index 00000000000..64878b75275
--- /dev/null
+++ b/official/recommendation/uplift/layers/heads/two_tower_logits_head_test.py
@@ -0,0 +1,188 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for two_tower_logits_head."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift.layers.heads import two_tower_logits_head
+
+
+class TwoTowerLogitsHeadTest(
+ keras_test_case.KerasTestCase, parameterized.TestCase
+):
+
+ def _get_layer(
+ self,
+ control_head=tf_keras.layers.Dense(1),
+ treatment_head=tf_keras.layers.Dense(1),
+ layering_method=two_tower_logits_head.LayeringMethod.NONE,
+ linear_layering_config=two_tower_logits_head.LinearLayeringConfig(),
+ **kwargs
+ ):
+ logits_head = two_tower_logits_head.TwoTowerLogitsHead(
+ control_head=control_head,
+ treatment_head=treatment_head,
+ layering_config=two_tower_logits_head.LayeringConfig(
+ layering_method=layering_method,
+ linear_layering_config=linear_layering_config,
+ ),
+ **kwargs
+ )
+ return tf_keras.models.clone_model(logits_head)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "layering_method_none",
+ "layering_method": two_tower_logits_head.LayeringMethod.NONE,
+ "expected_logits": (10 * tf.ones((5, 1)), 6 * tf.ones((5, 1))),
+ },
+ {
+ "testcase_name": "layering_method_logit_sum",
+ "layering_method": two_tower_logits_head.LayeringMethod.LOGIT_SUM,
+ "expected_logits": (10 * tf.ones((5, 1)), (6 + 10) * tf.ones((5, 1))),
+ },
+ {
+ "testcase_name": "layering_method_linear_layering",
+ "layering_method": (
+ two_tower_logits_head.LayeringMethod.LINEAR_LAYERING
+ ),
+ "expected_logits": (
+ 10 * tf.ones((5, 1)),
+ ((10 + 1) * 6) * tf.ones((5, 1)),
+ ),
+ "linear_layering_config": two_tower_logits_head.LinearLayeringConfig(
+ kernel_initializer="ones"
+ ),
+ },
+ )
+ def test_layer_correctness(
+ self, layering_method, expected_logits, linear_layering_config=None
+ ):
+ control_embedding = tf.ones((5, 10))
+ treatment_embedding = tf.ones((5, 6))
+ inputs = (control_embedding, treatment_embedding)
+ layer = self._get_layer(
+ control_head=tf_keras.layers.Dense(1, kernel_initializer="ones"),
+ treatment_head=tf_keras.layers.Dense(1, kernel_initializer="ones"),
+ layering_method=layering_method,
+ linear_layering_config=linear_layering_config,
+ )
+ logits = layer(inputs)
+ self.assertAllEqual(expected_logits, logits)
+
+ @parameterized.parameters(
+ two_tower_logits_head.LayeringMethod.NONE,
+ two_tower_logits_head.LayeringMethod.LOGIT_SUM,
+ two_tower_logits_head.LayeringMethod.LINEAR_LAYERING,
+ )
+ def test_layer_stable(self, layering_method):
+ layer = self._get_layer(layering_method=layering_method)
+ inputs = (tf.random.normal((10, 12)), tf.random.normal((10, 6)))
+ self.assertLayerStable(inputs=inputs, layer=layer)
+
+ @parameterized.parameters(
+ two_tower_logits_head.LayeringMethod.NONE,
+ two_tower_logits_head.LayeringMethod.LOGIT_SUM,
+ two_tower_logits_head.LayeringMethod.LINEAR_LAYERING,
+ )
+ def test_layer_savable(self, layering_method):
+ layer = self._get_layer(layering_method=layering_method, name="logits_head")
+ inputs = (tf.random.normal((5, 10)), tf.random.normal((5, 6)))
+ self.assertLayerSavable(inputs=inputs, layer=layer)
+
+ @parameterized.parameters(
+ two_tower_logits_head.LayeringMethod.NONE,
+ two_tower_logits_head.LayeringMethod.LOGIT_SUM,
+ two_tower_logits_head.LayeringMethod.LINEAR_LAYERING,
+ )
+ def test_stop_gradient_does_not_update_control_tower(self, layering_method):
+ linear_layering_config = None
+ if layering_method == two_tower_logits_head.LayeringMethod.LINEAR_LAYERING:
+ linear_layering_config = two_tower_logits_head.LinearLayeringConfig()
+
+ # Build dummy model.
+ inputs = tf_keras.Input(shape=(1,))
+ control_embedding = tf_keras.layers.Dense(
+ 10, use_bias=False, kernel_initializer="ones", name="control_tower"
+ )(inputs)
+ treatment_embedding = tf_keras.layers.Dense(
+ 6, use_bias=False, kernel_initializer="ones", name="treatment_tower"
+ )(inputs)
+ logits = two_tower_logits_head.TwoTowerLogitsHead(
+ control_head=tf_keras.layers.Dense(
+ 1, use_bias=False, kernel_initializer="ones", name="control_head"
+ ),
+ treatment_head=tf_keras.layers.Dense(
+ 1, use_bias=False, kernel_initializer="ones", name="treatment_head"
+ ),
+ layering_config=two_tower_logits_head.LayeringConfig(
+ layering_method=layering_method,
+ linear_layering_config=linear_layering_config,
+ ),
+ name="logits_head",
+ )((control_embedding, treatment_embedding))
+ model = tf_keras.Model(inputs=inputs, outputs=logits)
+
+ def _get_kernel(name: str):
+ if name == "control_head":
+ return model.get_layer("logits_head")._control_head.weights[0]
+ elif name == "treatment_head":
+ return model.get_layer("logits_head")._treatment_head.weights[0]
+ return model.get_layer(name=name).weights[0]
+
+ self.assertAllClose(_get_kernel("control_tower"), tf.ones((1, 10)))
+ self.assertAllClose(_get_kernel("treatment_tower"), tf.ones((1, 6)))
+ self.assertAllClose(_get_kernel("control_head"), tf.ones((10, 1)))
+ self.assertAllClose(_get_kernel("treatment_head"), tf.ones((6, 1)))
+
+ # Train for one step on treatment examples only. The control weights should
+ # remain unchanged due to the stop_gradient op.
+ with tf.GradientTape() as tape:
+ _, treatment_logits = model(tf.ones((3, 1)), training=True)
+ loss = tf.reduce_mean(
+ tf_keras.losses.mse(tf.ones((3, 1)), treatment_logits)
+ )
+ optimizer = tf_keras.optimizers.SGD(learning_rate=0.1)
+ grads = tape.gradient(loss, model.trainable_weights)
+ optimizer.apply_gradients(zip(grads, model.trainable_weights))
+
+ self.assertAllClose(_get_kernel("control_tower"), tf.ones((1, 10)))
+ self.assertNotAllClose(_get_kernel("treatment_tower"), tf.ones((1, 6)))
+ self.assertAllClose(_get_kernel("control_head"), tf.ones((10, 1)))
+ self.assertNotAllClose(_get_kernel("treatment_head"), tf.ones((6, 1)))
+
+ @parameterized.parameters(
+ two_tower_logits_head.LayeringMethod.NONE,
+ two_tower_logits_head.LayeringMethod.LOGIT_SUM,
+ two_tower_logits_head.LayeringMethod.LINEAR_LAYERING,
+ )
+ def test_layer_configurable(self, layering_method):
+ layer = self._get_layer(layering_method=layering_method)
+ self.assertLayerConfigurable(layer=layer)
+
+ def test_different_logit_shapes_raises_error(self):
+ layer = two_tower_logits_head.TwoTowerLogitsHead(
+ control_head=tf_keras.layers.Dense(1),
+ treatment_head=tf_keras.layers.Dense(2),
+ )
+ inputs = (tf.zeros(3, 1), tf.ones(3, 1))
+ with self.assertRaises(ValueError):
+ layer(inputs=inputs)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/layers/uplift_networks/__init__.py b/official/recommendation/uplift/layers/uplift_networks/__init__.py
new file mode 100644
index 00000000000..9d613b1fd9b
--- /dev/null
+++ b/official/recommendation/uplift/layers/uplift_networks/__init__.py
@@ -0,0 +1,19 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Uplift networks package defitions."""
+
+from official.recommendation.uplift.layers.uplift_networks import base_uplift_networks
+from official.recommendation.uplift.layers.uplift_networks import two_tower_output_head
+from official.recommendation.uplift.layers.uplift_networks import two_tower_uplift_network
diff --git a/official/recommendation/uplift/layers/uplift_networks/base_uplift_networks.py b/official/recommendation/uplift/layers/uplift_networks/base_uplift_networks.py
new file mode 100644
index 00000000000..199dff0c61e
--- /dev/null
+++ b/official/recommendation/uplift/layers/uplift_networks/base_uplift_networks.py
@@ -0,0 +1,39 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Defines base abstract uplift network layers."""
+
+import abc
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+
+
+class BaseTwoTowerUpliftNetwork(tf_keras.layers.Layer, metaclass=abc.ABCMeta):
+ """Abstract class for uplift layers that compute control and treatment logits.
+
+ A TwoTowerUpliftNetwork layer is expected to take in a dictionary of feature
+ tensors and compute two sets of logits: one for the control group and one for
+ treatment group.
+ """
+
+ @abc.abstractmethod
+ def call(
+ self,
+ inputs: types.DictOfTensors,
+ training: bool | None = None,
+ mask: tf.Tensor | None = None,
+ ) -> types.TwoTowerTrainingOutputs:
+ raise NotImplementedError()
diff --git a/official/recommendation/uplift/layers/uplift_networks/two_tower_output_head.py b/official/recommendation/uplift/layers/uplift_networks/two_tower_output_head.py
new file mode 100644
index 00000000000..6c274a4eda8
--- /dev/null
+++ b/official/recommendation/uplift/layers/uplift_networks/two_tower_output_head.py
@@ -0,0 +1,180 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builds a TwoTowerOutputHead layer."""
+
+from __future__ import annotations
+
+from typing import Any, Callable, Mapping, MutableMapping
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift import utils
+from official.recommendation.uplift.layers.uplift_networks import base_uplift_networks
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class TwoTowerOutputHead(tf_keras.layers.Layer):
+ """Two tower training and inference output computation.
+
+ This layer is intended to be used in conjunction with a two tower uplift
+ network layer to compute training and inference outputs. It passes the input
+ dictionary of feature tensors to the uplift network and computes training and
+ inference tensors from the uplift network's outputs.
+
+ The control, treatment, and uplift predictions are computed from the uplift
+ network's control and treatment logits. Additionally, if the treatment
+ indicator tensor is part of the inputs, the true logits are also computed to
+ be used for loss computation. The treatment indicator tensor should therefore
+ be present in the inputs dictionary during training and evaluation.
+ """
+
+ def __init__(
+ self,
+ treatment_indicator_feature_name: str,
+ uplift_network: base_uplift_networks.BaseTwoTowerUpliftNetwork,
+ inverse_link_fn: Callable[[tf.Tensor], tf.Tensor] | None = None,
+ **kwargs,
+ ):
+ """Initializes the instance.
+
+ Args:
+ treatment_indicator_feature_name: the name of the feature representing the
+ treatment indicator tensor, which should be castable to a boolean tensor
+ (False for control and True for treatment). When this tensor is part of
+ the layer inputs, additional true logits are returned for the loss
+ computation. In this case the layer returns a `TwoTowerTrainingOutputs`
+ instance, otherwise a `TwoTowerPredictionOutputs` instance is returned.
+ uplift_network: a layer for computing control and treatment logits. Its
+ input is expected to be a dictionary of feature tensors and its output
+ is exptected to be a `TwoTowerNetworkOutputs` instance.
+ inverse_link_fn: a function for computing the control and treatment
+ predictions from their respective logits. If left as `None` it is
+ functionally equivalent to the idendity function.
+ **kwargs: base layer keyword arguments.
+ """
+ super().__init__(**kwargs)
+
+ self._treatment_indicator_feature_name = treatment_indicator_feature_name
+ self._uplift_network = uplift_network
+ self._inverse_link_fn = inverse_link_fn
+
+ def call(
+ self,
+ inputs: types.DictOfTensors,
+ training: bool | None = None,
+ mask: tf.Tensor | None = None,
+ ) -> types.TwoTowerPredictionOutputs | types.TwoTowerTrainingOutputs:
+ """Computes two tower inference and training outputs.
+
+ Args:
+ inputs: feature tensors to be passed to the uplift network. For training
+ and evaluation, the treatment indicator tensor must be given as part of
+ the inputs.
+ training: optional boolean training flag to pass to the uplift network.
+ mask: optional tensor mask to pass to the uplift network.
+
+ Returns:
+ `TwoTowerTrainingOutputs` when the treatment indicator tensor is part of
+ the input features, otherwise `TwoTowerPredictionOutputs`.
+ """
+ outputs: types.TwoTowerNetworkOutputs = self._uplift_network(
+ inputs=inputs, training=training, mask=mask
+ )
+
+ if self._inverse_link_fn is None:
+ control_predictions = outputs.control_logits
+ treatment_predictions = outputs.treatment_logits
+ else:
+ control_predictions = self._inverse_link_fn(outputs.control_logits)
+ treatment_predictions = self._inverse_link_fn(outputs.treatment_logits)
+
+ uplift = treatment_predictions - control_predictions
+
+ # If the treatment indicator tensor is not in the layer inputs return
+ # inference output.
+ if self._treatment_indicator_feature_name not in inputs:
+ return types.TwoTowerPredictionOutputs(
+ shared_embedding=outputs.shared_embedding,
+ control_logits=outputs.control_logits,
+ treatment_logits=outputs.treatment_logits,
+ control_predictions=control_predictions,
+ treatment_predictions=treatment_predictions,
+ uplift=uplift,
+ )
+
+ # Compute the true logits only when the treatment_indicator tensor is given.
+ is_treatment = tf.cast(
+ inputs[self._treatment_indicator_feature_name], tf.bool
+ )
+
+ # Expand treatment_indicator tensor to match the logits and predictions.
+ # This is done to prevent tf.where from broadcasting the condition. For
+ # example, if is_treatment is of shape (3,) and the predictions are of shape
+ # (3, 1) then tf.where will return a tensor of shape (3, 3). However,
+ # tf.where will return a tensor of shape (3, 1) if the is_treatment tensor
+ # is also of that shape.
+ # Also, note that for generalizability the logits are allowed to have a
+ # different shape than the predictions, however the control/treatment logits
+ # shapes must be equal to each other and likewise for the control/treatment
+ # prediction shapes.
+ true_logits = tf.where(
+ utils.expand_to_match_rank(is_treatment, outputs.control_logits),
+ outputs.treatment_logits,
+ outputs.control_logits,
+ )
+ true_predictions = tf.where(
+ utils.expand_to_match_rank(is_treatment, control_predictions),
+ treatment_predictions,
+ control_predictions,
+ )
+
+ # Create a new tensor since ExtensionTypes are immutable.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=outputs.shared_embedding,
+ control_logits=outputs.control_logits,
+ treatment_logits=outputs.treatment_logits,
+ control_predictions=control_predictions,
+ treatment_predictions=treatment_predictions,
+ uplift=uplift,
+ true_logits=true_logits,
+ true_predictions=true_predictions,
+ is_treatment=is_treatment,
+ )
+
+ def get_config(self) -> Mapping[str, Any]:
+ config = super().get_config()
+ config.update({
+ "treatment_indicator_feature_name": (
+ self._treatment_indicator_feature_name
+ ),
+ "uplift_network": tf_keras.utils.serialize_keras_object(
+ self._uplift_network
+ ),
+ "inverse_link_fn": tf_keras.utils.serialize_keras_object(
+ self._inverse_link_fn
+ ),
+ })
+ return config
+
+ @classmethod
+ def from_config(cls, config: MutableMapping[str, Any]) -> TwoTowerOutputHead:
+ config["uplift_network"] = tf_keras.layers.deserialize(
+ config["uplift_network"]
+ )
+ config["inverse_link_fn"] = tf_keras.utils.deserialize_keras_object(
+ config["inverse_link_fn"]
+ )
+ return cls(**config)
diff --git a/official/recommendation/uplift/layers/uplift_networks/two_tower_output_head_test.py b/official/recommendation/uplift/layers/uplift_networks/two_tower_output_head_test.py
new file mode 100644
index 00000000000..273abdefe89
--- /dev/null
+++ b/official/recommendation/uplift/layers/uplift_networks/two_tower_output_head_test.py
@@ -0,0 +1,253 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for two_tower_output_head."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import types
+from official.recommendation.uplift.layers.uplift_networks import two_tower_output_head
+
+
+_DEFAULT_SHARED_EMBEDDING = tf.ones((3, 4, 5), dtype=tf.float32)
+_DEFAULT_CONTROL_LOGITS = tf.constant([-1, 0, 3], dtype=tf.float32)
+_DEFAULT_TREATMENT_LOGITS = tf.constant([2, -1, 1], dtype=tf.float32)
+_DEFAULT_IS_TREATMENT = tf.constant([1, 1, 0], dtype=tf.float32)
+
+
+@tf_keras.utils.register_keras_serializable()
+class TwoTowerUpliftNetworkMock(tf_keras.layers.Layer):
+
+ def call(self, inputs, training=None, mask=None):
+ return types.TwoTowerNetworkOutputs(
+ shared_embedding=inputs["shared_embedding"],
+ control_logits=inputs["control_logits"],
+ treatment_logits=inputs["treatment_logits"],
+ )
+
+
+class TwoTowerOutputHeadTest(
+ keras_test_case.KerasTestCase, parameterized.TestCase
+):
+
+ def _get_layer(self, inverse_link_fn=None):
+ layer = two_tower_output_head.TwoTowerOutputHead(
+ treatment_indicator_feature_name="is_treatment",
+ uplift_network=TwoTowerUpliftNetworkMock(),
+ inverse_link_fn=inverse_link_fn,
+ )
+ return layer
+
+ def _get_inputs(
+ self,
+ shared_embedding=_DEFAULT_SHARED_EMBEDDING,
+ control_logits=_DEFAULT_CONTROL_LOGITS,
+ treatment_logits=_DEFAULT_TREATMENT_LOGITS,
+ is_treatment=None,
+ ):
+ inputs = {
+ "shared_embedding": shared_embedding,
+ "control_logits": control_logits,
+ "treatment_logits": treatment_logits,
+ }
+ if is_treatment is not None:
+ inputs["is_treatment"] = is_treatment
+ return inputs
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "without_inverse_link",
+ "inverse_link_fn": None,
+ },
+ {
+ "testcase_name": "with_inverse_link",
+ "inverse_link_fn": tf.math.sigmoid,
+ },
+ )
+ def test_training_outputs_are_returned_when_treatment_indicator_in_inputs(
+ self, inverse_link_fn
+ ):
+ layer = self._get_layer(inverse_link_fn=inverse_link_fn)
+ inputs = self._get_inputs(is_treatment=_DEFAULT_IS_TREATMENT)
+ outputs = layer(inputs)
+ self.assertIsInstance(outputs, types.TwoTowerTrainingOutputs)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "without_inverse_link",
+ "inverse_link_fn": None,
+ },
+ {
+ "testcase_name": "with_inverse_link",
+ "inverse_link_fn": tf.math.sigmoid,
+ },
+ )
+ def test_prediction_outputs_are_returned_when_treatment_indicator_not_in_inputs(
+ self, inverse_link_fn
+ ):
+ layer = self._get_layer(inverse_link_fn=inverse_link_fn)
+ inputs = self._get_inputs(is_treatment=None)
+ outputs = layer(inputs)
+ self.assertIsInstance(outputs, types.TwoTowerPredictionOutputs)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "treatment_indicator_without_inverse_link",
+ "is_treatment": _DEFAULT_IS_TREATMENT,
+ "inverse_link_fn": None,
+ },
+ {
+ "testcase_name": "treatment_indicator_with_inverse_link",
+ "is_treatment": _DEFAULT_IS_TREATMENT,
+ "inverse_link_fn": tf.math.exp,
+ },
+ {
+ "testcase_name": "without_treatment_indicator_and_inverse_link",
+ "is_treatment": None,
+ "inverse_link_fn": None,
+ },
+ {
+ "testcase_name": "without_treatment_indicator_but_with_inverse_link",
+ "is_treatment": None,
+ "inverse_link_fn": tf.math.exp,
+ },
+ )
+ def test_uplift_network_outputs_remain_unchanged(
+ self, is_treatment, inverse_link_fn
+ ):
+ layer = self._get_layer(inverse_link_fn=inverse_link_fn)
+ inputs = self._get_inputs(is_treatment=is_treatment)
+ outputs = layer(inputs)
+
+ self.assertAllEqual(_DEFAULT_SHARED_EMBEDDING, outputs.shared_embedding)
+ self.assertAllEqual(_DEFAULT_CONTROL_LOGITS, outputs.control_logits)
+ self.assertAllEqual(_DEFAULT_TREATMENT_LOGITS, outputs.treatment_logits)
+
+ @parameterized.product(
+ (dict(is_treatment=None), dict(is_treatment=_DEFAULT_IS_TREATMENT)),
+ (
+ dict(
+ inverse_link_fn=None,
+ expected_control_predictions=_DEFAULT_CONTROL_LOGITS,
+ expected_treatment_predictions=_DEFAULT_TREATMENT_LOGITS,
+ expected_uplift=_DEFAULT_TREATMENT_LOGITS
+ - _DEFAULT_CONTROL_LOGITS,
+ ),
+ dict(
+ inverse_link_fn=tf.math.exp,
+ expected_control_predictions=tf.math.exp(_DEFAULT_CONTROL_LOGITS),
+ expected_treatment_predictions=tf.math.exp(
+ _DEFAULT_TREATMENT_LOGITS
+ ),
+ expected_uplift=tf.math.exp(_DEFAULT_TREATMENT_LOGITS)
+ - tf.math.exp(_DEFAULT_CONTROL_LOGITS),
+ ),
+ ),
+ )
+ def test_uplift_predictions_are_computed_from_logits(
+ self,
+ is_treatment,
+ inverse_link_fn,
+ expected_control_predictions,
+ expected_treatment_predictions,
+ expected_uplift,
+ ):
+ layer = self._get_layer(inverse_link_fn=inverse_link_fn)
+ inputs = self._get_inputs(
+ control_logits=_DEFAULT_CONTROL_LOGITS,
+ treatment_logits=_DEFAULT_TREATMENT_LOGITS,
+ is_treatment=is_treatment,
+ )
+ outputs = layer(inputs)
+
+ self.assertAllClose(
+ expected_control_predictions, outputs.control_predictions
+ )
+ self.assertAllClose(
+ expected_treatment_predictions, outputs.treatment_predictions
+ )
+ self.assertAllClose(expected_uplift, outputs.uplift)
+
+ def test_true_logits_correspond_to_control_and_treatment_logits(self):
+ layer = self._get_layer()
+ inputs = self._get_inputs(
+ control_logits=tf.constant([-1, 0, 3]),
+ treatment_logits=tf.constant([2, -1, 1]),
+ is_treatment=tf.constant([1, 1, 0]),
+ )
+ outputs = layer(inputs)
+ self.assertAllEqual(tf.constant([2, -1, 3]), outputs.true_logits)
+
+ def test_true_preds_correspond_to_control_and_treatment_preds(self):
+ layer = self._get_layer(inverse_link_fn=tf.nn.relu)
+ inputs = self._get_inputs(
+ control_logits=tf.constant([2, 0, 3]),
+ treatment_logits=tf.constant([-1, 2, 1]),
+ is_treatment=tf.constant([1, 1, 0]),
+ )
+ outputs = layer(inputs)
+ self.assertAllEqual(tf.constant([0, 2, 3]), outputs.true_predictions)
+
+ def test_is_treatment_tensor_gets_converted_to_boolean_tensor(self):
+ layer = self._get_layer()
+ inputs = self._get_inputs(is_treatment=tf.constant([1, 1, 0]))
+ outputs = layer(inputs)
+ self.assertAllEqual(tf.constant([True, True, False]), outputs.is_treatment)
+
+ def test_true_logits_correctness_with_logits_rank_mismatch(self):
+ layer = self._get_layer()
+ inputs = self._get_inputs(
+ control_logits=tf.constant([[2], [0], [3]]),
+ treatment_logits=tf.constant([[-1], [2], [1]]),
+ is_treatment=tf.constant([1, 1, 0]),
+ )
+ outputs = layer(inputs)
+ self.assertAllEqual(tf.constant([[-1], [2], [3]]), outputs.true_predictions)
+
+ @parameterized.product(
+ (dict(is_treatment=None), dict(is_treatment=_DEFAULT_IS_TREATMENT)),
+ (dict(inverse_link_fn=None), dict(inverse_link_fn=tf.math.sigmoid)),
+ )
+ def test_multiple_layer_calls_with_same_input_returns_same_output(
+ self, is_treatment, inverse_link_fn
+ ):
+ layer = self._get_layer(inverse_link_fn=inverse_link_fn)
+ inputs = self._get_inputs(is_treatment=is_treatment)
+ self.assertLayerStable(layer=layer, inputs=inputs)
+
+ @parameterized.product(
+ (dict(is_treatment=None), dict(is_treatment=_DEFAULT_IS_TREATMENT)),
+ (dict(inverse_link_fn=None), dict(inverse_link_fn=tf.math.sigmoid)),
+ )
+ def test_layer_saving_succeeds(self, is_treatment, inverse_link_fn):
+ layer = self._get_layer(inverse_link_fn=inverse_link_fn)
+ inputs = self._get_inputs(is_treatment=is_treatment)
+ self.assertLayerSavable(layer=layer, inputs=inputs)
+
+ @parameterized.product(
+ (dict(is_treatment=None), dict(is_treatment=_DEFAULT_IS_TREATMENT)),
+ (dict(inverse_link_fn=None), dict(inverse_link_fn=tf.math.sigmoid)),
+ )
+ def test_from_config_layer_returns_same_output_as_original_layer(
+ self, is_treatment, inverse_link_fn
+ ):
+ layer = self._get_layer(inverse_link_fn=inverse_link_fn)
+ inputs = self._get_inputs(is_treatment=is_treatment)
+ self.assertLayerConfigurable(layer=layer, inputs=inputs)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network.py b/official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network.py
new file mode 100644
index 00000000000..291d4c21b03
--- /dev/null
+++ b/official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network.py
@@ -0,0 +1,212 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builds a TwoTowerUpliftNetwork layer."""
+
+from __future__ import annotations
+
+from typing import Any, Dict
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift.layers.uplift_networks import base_uplift_networks
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class TwoTowerUpliftNetwork(base_uplift_networks.BaseTwoTowerUpliftNetwork):
+ """Computes control and treatment logits in two separate towers.
+
+ Takes as input a dictionary of feature tensors and computes logits for the
+ control and treatment groups. The layer returns a `TwoTowerNetworkOutputs`
+ which contains the main tensors computed during the feed forward pass of the
+ network.
+
+ See ../../README.md for more details on the network's architecture.
+ """
+
+ def __init__(
+ self,
+ backbone: tf_keras.layers.Layer,
+ control_tower: tf_keras.layers.Layer,
+ treatment_tower: tf_keras.layers.Layer,
+ logits_head: tf_keras.layers.Layer,
+ control_feature_encoder: tf_keras.layers.Layer | None = None,
+ control_input_combiner: tf_keras.layers.Layer | None = None,
+ treatment_feature_encoder: tf_keras.layers.Layer | None = None,
+ treatment_input_combiner: tf_keras.layers.Layer | None = None,
+ **kwargs,
+ ):
+ """Initializes a TwoTowerUpliftNetwork layer.
+
+ Args:
+ backbone: encodes input features into a shared embedding for the control
+ and treatment towers. Its input is a dictionary with the input features
+ and its output is expected to be a single dense embedding.
+ control_tower: computes an embedding for the control group. Its input and
+ output is a single dense tensor.
+ treatment_tower: computes an embedding for the treatment group. Its input
+ and output is a single dense tensor.
+ logits_head: computes control and treatment logits. Its inputs is a tuple
+ of (control_embedding, treatment_embedding) and its output is expected
+ to be a tuple of (control_logits, treatment_logits).
+ control_feature_encoder: encodes control specific input features into a
+ dense tensor. Its input is the entire inputs dictionary. If this layer
+ is provided, a control_input_combiner layer must also be given.
+ control_input_combiner: combines the shared embedding from the backbone
+ with the control_feature_encoder output embedding. Its input is a list
+ with the shared and control specific embedding, and its output is a
+ single dense embedding to be fed into the control_tower layer. This
+ layer should only be provided if a control_feature_encoder is also
+ provided.
+ treatment_feature_encoder: encodes treatment specific input features into
+ a dense tensor. Its input is the entire inputs dictionary. If this layer
+ is provided, a treatment_input_combiner layer must also be given.
+ treatment_input_combiner: combines the shared embedding from the backbone
+ with the treatment_feature_encoder output embedding. Its input is a list
+ with the shared and treatment specific embedding, and its output is a
+ single dense embedding to be fed into the treatment_tower layer. This
+ layer should only be provided if a treatment_feature_encoder is also
+ provided.
+ **kwargs: base layer keyword arguments.
+
+ Raises:
+ ValueError if only one of the control encoder or combiner layers is set.
+ Both must be a Keras layer or None.
+ ValueError if only one of the treatment encoder or combiner layers is set.
+ Both must be a Keras layer or None.
+ """
+ super().__init__(**kwargs)
+
+ self._backbone = backbone
+ self._control_tower = control_tower
+ self._treatment_tower = treatment_tower
+ self._logits_head = logits_head
+ self._control_feature_encoder = control_feature_encoder
+ self._control_input_combiner = control_input_combiner
+ self._treatment_feature_encoder = treatment_feature_encoder
+ self._treatment_input_combiner = treatment_input_combiner
+
+ self._validate_encoder_combiner_layers(
+ control_feature_encoder, control_input_combiner, "control"
+ )
+ self._validate_encoder_combiner_layers(
+ treatment_feature_encoder, treatment_input_combiner, "treatment"
+ )
+
+ def _validate_encoder_combiner_layers(
+ self,
+ encoder: tf_keras.layer.Layer,
+ combiner: tf_keras.layer.Layer,
+ name: str,
+ ) -> None:
+ if encoder is not None and combiner is None:
+ raise ValueError(
+ f"The {name}_input_combiner layer must be specified if the"
+ f" {name}_feature_encoder is not None. Consider using"
+ " tf_keras.layers.Concatenate() as a combiner layer."
+ )
+ if encoder is None and combiner is not None:
+ raise ValueError(
+ f"The {name}_feature_encoder layer must be specified if the"
+ f" {name}_input_combiner is not None."
+ )
+
+ def call(
+ self,
+ inputs: types.DictOfTensors,
+ training: bool | None = None,
+ mask: tf.Tensor | None = None,
+ ) -> types.TwoTowerNetworkOutputs:
+ """Computes control and treatment logits.
+
+ Args:
+ inputs: dictionary of feature tensors to be passed to the backbone and
+ control/treatment feature encoder layers when set.
+ training: unused optional training flag.
+ mask: unused optional mask.
+
+ Returns:
+ A `TwoTowerNetworkOutputs` ExtensionType which should be used as a
+ dataclass to access the main tensors computed during the feed forward pass
+ of the network.
+ """
+ # Compute shared embedding for control and treatment towers.
+ shared_embedding = self._backbone(inputs)
+
+ # Compute control embedding.
+ if self._control_feature_encoder is not None:
+ control_feature_encoding = self._control_feature_encoder(inputs)
+ control_tower_input = self._control_input_combiner( # pyrefly: ignore[not-callable]
+ [shared_embedding, control_feature_encoding]
+ )
+ else:
+ control_tower_input = shared_embedding
+ control_embedding = self._control_tower(control_tower_input)
+
+ # Compute treatment embedding.
+ if self._treatment_feature_encoder is not None:
+ treatment_feature_encoding = self._treatment_feature_encoder(inputs)
+ treatment_tower_input = self._treatment_input_combiner( # pyrefly: ignore[not-callable]
+ [shared_embedding, treatment_feature_encoding]
+ )
+ else:
+ treatment_tower_input = shared_embedding
+ treatment_embedding = self._treatment_tower(treatment_tower_input)
+
+ # Compute control and treatment logits.
+ control_logits, treatment_logits = self._logits_head(
+ (control_embedding, treatment_embedding)
+ )
+
+ return types.TwoTowerNetworkOutputs(
+ shared_embedding=shared_embedding,
+ control_logits=control_logits,
+ treatment_logits=treatment_logits,
+ )
+
+ def get_config(self) -> Dict[str, Any]:
+ config = super().get_config()
+
+ for layer_name, layer in (
+ ("backbone", self._backbone),
+ ("control_tower", self._control_tower),
+ ("treatment_tower", self._treatment_tower),
+ ("logits_head", self._logits_head),
+ ("control_feature_encoder", self._control_feature_encoder),
+ ("control_input_combiner", self._control_input_combiner),
+ ("treatment_feature_encoder", self._treatment_feature_encoder),
+ ("treatment_input_combiner", self._treatment_input_combiner),
+ ):
+ config[layer_name] = tf_keras.utils.serialize_keras_object(layer)
+
+ return config
+
+ @classmethod
+ def from_config(cls, config: Dict[str, Any]) -> TwoTowerUpliftNetwork:
+ for layer_name in (
+ "backbone",
+ "control_tower",
+ "treatment_tower",
+ "logits_head",
+ "control_feature_encoder",
+ "control_input_combiner",
+ "treatment_feature_encoder",
+ "treatment_input_combiner",
+ ):
+ # layers.deserialize does not accept empty config
+ if config.get(layer_name):
+ config[layer_name] = tf_keras.layers.deserialize(config[layer_name])
+
+ return cls(**config)
diff --git a/official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network_test.py b/official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network_test.py
new file mode 100644
index 00000000000..44edcf2531c
--- /dev/null
+++ b/official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network_test.py
@@ -0,0 +1,227 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for two_tower_uplift_network."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift.layers.uplift_networks import two_tower_uplift_network
+
+
+class TwoTowerUpliftNetworkTest(
+ keras_test_case.KerasTestCase, parameterized.TestCase
+):
+
+ def _get_full_layer(self, **kwargs):
+ layer = two_tower_uplift_network.TwoTowerUpliftNetwork(
+ backbone=tf_keras.layers.Lambda(lambda inputs: inputs["shared_inputs"]),
+ control_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(-1.0)
+ ),
+ treatment_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(2.0)
+ ),
+ logits_head=tf_keras.layers.Identity(),
+ control_feature_encoder=kwargs.get(
+ "control_feature_encoder",
+ tf_keras.layers.Lambda(lambda inputs: inputs["control_inputs"]),
+ ),
+ control_input_combiner=kwargs.get(
+ "control_input_combiner", tf_keras.layers.Concatenate()
+ ),
+ treatment_feature_encoder=kwargs.get(
+ "treatment_feature_encoder",
+ tf_keras.layers.Lambda(lambda inputs: inputs["treatment_inputs"]),
+ ),
+ treatment_input_combiner=kwargs.get(
+ "treatment_input_combiner", tf_keras.layers.Concatenate()
+ ),
+ )
+ return tf_keras.models.clone_model(layer)
+
+ def _get_full_layer_inputs(self):
+ return {
+ "shared_inputs": tf.ones((3, 3)),
+ "control_inputs": tf.ones((3, 10)),
+ "treatment_inputs": tf.ones((3, 10)),
+ }
+
+ def test_forward_pass_no_control_or_treatment_encoders(self):
+ layer = two_tower_uplift_network.TwoTowerUpliftNetwork(
+ backbone=tf_keras.layers.Lambda(lambda inputs: inputs["shared_inputs"]),
+ control_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(-1.0)
+ ),
+ treatment_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(2.0)
+ ),
+ logits_head=tf_keras.layers.Identity(),
+ )
+ inputs = {"shared_inputs": tf.ones((3, 3))}
+ outputs = layer(inputs)
+
+ self.assertAllClose(tf.ones((3, 3)), outputs.shared_embedding)
+ self.assertAllClose(3 * -1 * tf.ones((3, 1)), outputs.control_logits)
+ self.assertAllClose(3 * 2 * tf.ones((3, 1)), outputs.treatment_logits)
+
+ def test_forward_pass_only_control_encoder(self):
+ layer = two_tower_uplift_network.TwoTowerUpliftNetwork(
+ backbone=tf_keras.layers.Lambda(lambda inputs: inputs["shared_inputs"]),
+ control_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(-1.0)
+ ),
+ treatment_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(2.0)
+ ),
+ logits_head=tf_keras.layers.Identity(),
+ control_feature_encoder=tf_keras.layers.Lambda(
+ lambda inputs: inputs["control_inputs"]
+ ),
+ control_input_combiner=tf_keras.layers.Concatenate(),
+ )
+ inputs = {
+ "shared_inputs": tf.ones((3, 3)),
+ "control_inputs": tf.ones((3, 10)),
+ }
+ outputs = layer(inputs)
+
+ self.assertAllClose(tf.ones((3, 3)), outputs.shared_embedding)
+ self.assertAllClose((3 + 10) * -1 * tf.ones((3, 1)), outputs.control_logits)
+ self.assertAllClose(3 * 2 * tf.ones((3, 1)), outputs.treatment_logits)
+
+ def test_forward_pass_only_treatment_encoder(self):
+ layer = two_tower_uplift_network.TwoTowerUpliftNetwork(
+ backbone=tf_keras.layers.Lambda(lambda inputs: inputs["shared_inputs"]),
+ control_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(-1.0)
+ ),
+ treatment_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(2.0)
+ ),
+ logits_head=tf_keras.layers.Identity(),
+ treatment_feature_encoder=tf_keras.layers.Lambda(
+ lambda inputs: inputs["treatment_inputs"]
+ ),
+ treatment_input_combiner=tf_keras.layers.Concatenate(),
+ )
+ inputs = {
+ "shared_inputs": tf.ones((3, 3)),
+ "treatment_inputs": tf.ones((3, 10)),
+ }
+ outputs = layer(inputs)
+
+ self.assertAllClose(tf.ones((3, 3)), outputs.shared_embedding)
+ self.assertAllClose(3 * -1 * tf.ones((3, 1)), outputs.control_logits)
+ self.assertAllClose(
+ (3 + 10) * 2 * tf.ones((3, 1)), outputs.treatment_logits
+ )
+
+ def test_forward_pass_both_control_and_treatment_encoders(self):
+ layer = two_tower_uplift_network.TwoTowerUpliftNetwork(
+ backbone=tf_keras.layers.Lambda(lambda inputs: inputs["shared_inputs"]),
+ control_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(-1.0)
+ ),
+ treatment_tower=tf_keras.layers.Dense(
+ 1, kernel_initializer=tf_keras.initializers.Constant(2.0)
+ ),
+ logits_head=tf_keras.layers.Identity(),
+ control_feature_encoder=tf_keras.layers.Lambda(
+ lambda inputs: inputs["control_inputs"]
+ ),
+ control_input_combiner=tf_keras.layers.Concatenate(),
+ treatment_feature_encoder=tf_keras.layers.Lambda(
+ lambda inputs: inputs["treatment_inputs"]
+ ),
+ treatment_input_combiner=tf_keras.layers.Concatenate(),
+ )
+ inputs = {
+ "shared_inputs": tf.ones((3, 3)),
+ "control_inputs": tf.ones((3, 10)),
+ "treatment_inputs": tf.ones((3, 10)),
+ }
+ outputs = layer(inputs)
+
+ self.assertAllClose(tf.ones((3, 3)), outputs.shared_embedding)
+ self.assertAllClose((3 + 10) * -1 * tf.ones((3, 1)), outputs.control_logits)
+ self.assertAllClose(
+ (3 + 10) * 2 * tf.ones((3, 1)), outputs.treatment_logits
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "encoder_without_combiner",
+ "control_feature_encoder": tf_keras.layers.Lambda(
+ lambda inputs: inputs["control_inputs"]
+ ),
+ "control_input_combiner": None,
+ },
+ {
+ "testcase_name": "combiner_without_encoder",
+ "control_feature_encoder": None,
+ "control_input_combiner": tf_keras.layers.Concatenate(),
+ },
+ )
+ def test_invalid_control_encoder_combiner_combination_raises_error(
+ self, control_feature_encoder, control_input_combiner
+ ):
+ with self.assertRaises(ValueError):
+ self._get_full_layer(
+ control_feature_encoder=control_feature_encoder,
+ control_input_combiner=control_input_combiner,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "encoder_without_combiner",
+ "treatment_feature_encoder": tf_keras.layers.Lambda(
+ lambda inputs: inputs["treatment_inputs"]
+ ),
+ "treatment_input_combiner": None,
+ },
+ {
+ "testcase_name": "combiner_without_encoder",
+ "treatment_feature_encoder": None,
+ "treatment_input_combiner": tf_keras.layers.Concatenate(),
+ },
+ )
+ def test_invalid_treatment_encoder_combiner_combination_raises_error(
+ self, treatment_feature_encoder, treatment_input_combiner
+ ):
+ with self.assertRaises(ValueError):
+ self._get_full_layer(
+ treatment_feature_encoder=treatment_feature_encoder,
+ treatment_input_combiner=treatment_input_combiner,
+ )
+
+ def test_multiple_layer_calls_with_same_input_returns_same_output(self):
+ layer = self._get_full_layer()
+ inputs = self._get_full_layer_inputs()
+ self.assertLayerStable(layer=layer, inputs=inputs)
+
+ def test_layer_saving_succeeds(self):
+ layer = self._get_full_layer()
+ inputs = self._get_full_layer_inputs()
+ self.assertLayerSavable(layer=layer, inputs=inputs)
+
+ def test_from_config_layer_returns_same_output_as_original_layer(self):
+ layer = self._get_full_layer()
+ inputs = self._get_full_layer_inputs()
+ self.assertLayerConfigurable(layer=layer, inputs=inputs)
+
+
+if __name__ == "__main__":
+ tf_keras.__internal__.enable_unsafe_deserialization()
+ tf.test.main()
diff --git a/official/recommendation/uplift/losses/__init__.py b/official/recommendation/uplift/losses/__init__.py
new file mode 100644
index 00000000000..cc4745ffcf3
--- /dev/null
+++ b/official/recommendation/uplift/losses/__init__.py
@@ -0,0 +1,17 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Losses package definitions."""
+
+from official.recommendation.uplift.losses import true_logits_loss
diff --git a/official/recommendation/uplift/losses/true_logits_loss.py b/official/recommendation/uplift/losses/true_logits_loss.py
new file mode 100644
index 00000000000..2984ebf1dec
--- /dev/null
+++ b/official/recommendation/uplift/losses/true_logits_loss.py
@@ -0,0 +1,102 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Wrapper to apply any loss function on the true logits tensor."""
+
+from __future__ import annotations
+
+from typing import Any, Callable, Mapping, MutableMapping
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class TrueLogitsLoss(tf_keras.__internal__.losses.LossFunctionWrapper):
+ """Computes any arbitrary loss between the labels and the true logits tensor.
+
+ Note that the prediction tensor is expected to be of a tensor of type
+ `TwoTowerTrainingOutputs`.
+
+ Example standalone usage:
+
+ >>> y_true = tf.ones((3, 1))
+ >>> y_pred = types.TwoTowerTrainingOutputs(
+ ... control_logits=tf.constant([[0], [1], [0]]),
+ ... treatment_logits=tf.constant([[1], [0], [1]]),
+ ... true_logits=tf.ones((3, 1)),
+ ... is_treatment=tf.constant([[True], [False], [True]])
+ ... )
+ >>> loss = TrueLogitsLoss(
+ ... loss_fn=tf_keras.losses.mean_squared_error,
+ ... name="mean_squared_error",
+ ... reduction=tf_keras.losses.Reduction.SUM,
+ ... )
+ >>> loss(y_true, y_pred)
+ 0.0
+
+ Example usage with the `compile()` API:
+
+ ```python
+ model.compile(
+ optimizer='sgd'.
+ loss=TrueLogitsLoss(
+ loss_fn=tf_keras.losses.categorical_crossentropy,
+ name="categorical_crossentropy",
+ from_logits=True
+ )
+ )
+ ```
+ """
+
+ def __init__(
+ self,
+ loss_fn: Callable[[Any, tf.Tensor], tf.Tensor],
+ name: str = "true_logits_loss",
+ reduction=tf_keras.losses.Reduction.AUTO,
+ **loss_fn_kwargs,
+ ):
+ """Initialize `TrueLogitsLoss` instance.
+
+ Args:
+ loss_fn: The loss function to apply between the labels and true logits
+ tensor, with signature `loss_fn(y_true, y_pred, **loss_fn_kwargs)`.
+ name: Optional name for the instance.
+ reduction: Type of `tf_keras.losses.Reduction` to apply to loss. Default
+ value is `AUTO`. `AUTO` indicates that the reduction option will be
+ determined by the usage context. For almost all cases this defaults to
+ `SUM_OVER_BATCH_SIZE`. When used under a `tf.distribute.Strategy`,
+ except via `Model.compile()` and `Model.fit()`, using `AUTO` or
+ `SUM_OVER_BATCH_SIZE` will raise an error.
+ **loss_fn_kwargs: The keyword arguments that are passed on to `loss_fn`.
+ """
+ super().__init__(
+ fn=loss_fn, name=name, reduction=reduction, **loss_fn_kwargs
+ )
+
+ def call(
+ self, y_true: Any, y_pred: types.TwoTowerTrainingOutputs
+ ) -> tf.Tensor:
+ return super().call(y_true, y_pred.true_logits)
+
+ def get_config(self) -> Mapping[str, Any]:
+ config = super().get_config()
+ config["loss_fn"] = config.pop("fn")
+ return config
+
+ @classmethod
+ def from_config(cls, config: MutableMapping[str, Any]) -> TrueLogitsLoss:
+ config["loss_fn"] = tf_keras.losses.get(config["loss_fn"])
+ return cls(**config)
diff --git a/official/recommendation/uplift/losses/true_logits_loss_test.py b/official/recommendation/uplift/losses/true_logits_loss_test.py
new file mode 100644
index 00000000000..84a34f078c4
--- /dev/null
+++ b/official/recommendation/uplift/losses/true_logits_loss_test.py
@@ -0,0 +1,93 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for true_logits_loss."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift.losses import true_logits_loss
+
+
+class TrueLogitsLossTest(tf.test.TestCase, parameterized.TestCase):
+
+ def _get_y_pred(self, **kwargs):
+ # The shared embedding and control/treatment/uplift/true predictions are
+ # distracting from the test logic.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=tf.zeros((3, 1)),
+ control_predictions=tf.zeros((3, 1)),
+ treatment_predictions=tf.zeros((3, 1)),
+ true_predictions=tf.zeros((3, 1)),
+ uplift=tf.zeros((3, 1)),
+ **kwargs,
+ )
+
+ @parameterized.product(
+ (
+ dict(
+ reduction_strategy=tf_keras.losses.Reduction.NONE,
+ reduction_op=tf.identity,
+ ),
+ dict(
+ reduction_strategy=tf_keras.losses.Reduction.SUM,
+ reduction_op=tf.reduce_sum,
+ ),
+ dict(
+ reduction_strategy=tf_keras.losses.Reduction.SUM_OVER_BATCH_SIZE,
+ reduction_op=tf.reduce_mean,
+ ),
+ ),
+ (
+ dict(
+ loss_fn=tf_keras.losses.mean_squared_error, loss_fn_kwargs=dict()
+ ),
+ dict(
+ loss_fn=tf_keras.losses.mean_absolute_percentage_error,
+ loss_fn_kwargs=dict(),
+ ),
+ dict(
+ loss_fn=tf_keras.losses.huber,
+ loss_fn_kwargs=dict(delta=0.2),
+ ),
+ dict(
+ loss_fn=tf_keras.losses.categorical_crossentropy,
+ loss_fn_kwargs=dict(from_logits=True),
+ ),
+ ),
+ )
+ def test_correctness(
+ self, reduction_strategy, reduction_op, loss_fn, loss_fn_kwargs
+ ):
+ loss = true_logits_loss.TrueLogitsLoss(
+ loss_fn=loss_fn,
+ reduction=reduction_strategy,
+ **loss_fn_kwargs,
+ )
+ y_true = tf.constant([[0.4], [1.0], [0.0]])
+ y_pred = self._get_y_pred(
+ control_logits=tf.constant([[0.6], [4.3], [-0.3]]),
+ treatment_logits=tf.constant([[-2.0], [-0.1], [0.5]]),
+ true_logits=tf.constant([[-2.0], [4.3], [0.5]]),
+ is_treatment=tf.constant([[True], [False], [True]]),
+ )
+ expected_loss = reduction_op(
+ loss_fn(y_true, y_pred.true_logits, **loss_fn_kwargs)
+ )
+ self.assertAllEqual(expected_loss, loss(y_true, y_pred))
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/__init__.py b/official/recommendation/uplift/metrics/__init__.py
new file mode 100644
index 00000000000..0e132a50be8
--- /dev/null
+++ b/official/recommendation/uplift/metrics/__init__.py
@@ -0,0 +1,26 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Metrics package definition."""
+
+from official.recommendation.uplift.metrics import label_mean
+from official.recommendation.uplift.metrics import label_variance
+from official.recommendation.uplift.metrics import loss_metric
+from official.recommendation.uplift.metrics import metric_configs
+from official.recommendation.uplift.metrics import poisson_metrics
+from official.recommendation.uplift.metrics import sliced_metric
+from official.recommendation.uplift.metrics import treatment_fraction
+from official.recommendation.uplift.metrics import treatment_sliced_metric
+from official.recommendation.uplift.metrics import uplift_mean
+from official.recommendation.uplift.metrics import variance
diff --git a/official/recommendation/uplift/metrics/label_mean.py b/official/recommendation/uplift/metrics/label_mean.py
new file mode 100644
index 00000000000..dba457c5ed8
--- /dev/null
+++ b/official/recommendation/uplift/metrics/label_mean.py
@@ -0,0 +1,102 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for computing the label mean sliced by treatment group."""
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import treatment_sliced_metric
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class LabelMean(tf_keras.metrics.Metric):
+ """Computes the overall and treatment sliced label mean.
+
+ Note that the prediction tensor is expected to be of type
+ `TwoTowerTrainingOutputs`.
+
+ Example standalone usage:
+
+ >>> label_mean = LabelMean()
+ >>> y_true = tf.constant([[1], [2], [3], [4]])
+ >>> y_pred = types.TwoTowerTrainingOutputs(
+ ... control_logits=tf.zeros((4, 1)),
+ ... treatment_logits=tf.zeros((4, 1)),
+ ... true_logits=tf.zeros((4, 1)),
+ ... is_treatment=tf.constant([[True], [False], [True], [False]]),
+ ... )
+ >>> label_mean(y_true=y_true, y_pred=y_pred)
+ {
+ "label/mean": 2.5,
+ "label/mean/control": 3.0
+ "label/mean/treatment": 2.0
+ }
+
+ Example usage with the `compile()` API:
+
+ >>> model.compile(
+ ... optimizer="sgd",
+ ... loss=TrueLogitsLoss(tf_keras.losses.mean_squared_error),
+ ... metrics=[LabelMean()]
+ ... )
+ """
+
+ def __init__(self, name: str = "label/mean", **kwargs):
+ """Initializes a `LabelMean` instance.
+
+ Args:
+ name: name for the overall label mean metric result. The control and
+ treatment label means will have "/control" and "/treatment" appended to
+ the result name.
+ **kwargs: other base metric keyword arguments.
+ """
+ super().__init__(name=name, **kwargs)
+ self._sliced_mean = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.Mean(name=name, **kwargs)
+ )
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs,
+ sample_weight: tf.Tensor | None = None,
+ ):
+ """Updates the overall, control and treatment label means.
+
+ Args:
+ y_true: tensor labels.
+ y_pred: prediction logits. The treatment indicator tensor is used to slice
+ the labels into control and treatment groups.
+ sample_weight: optional sample weight to compute weighted label means. If
+ given, the sample weight will also be sliced by the treatment indicator
+ tensor to compute the weighted control and treatment label means.
+
+ Raises:
+ TypeError: if y_pred is not of type `TwoTowerTrainingOutputs`.
+ """
+ if not isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise TypeError(
+ "y_pred must be of type `TwoTowerTrainingOutputs` but got type"
+ f" {type(y_pred)} instead."
+ )
+
+ self._sliced_mean.update_state(
+ values=y_true,
+ is_treatment=y_pred.is_treatment,
+ sample_weight=sample_weight,
+ )
+
+ def result(self) -> dict[str, tf.Tensor]:
+ return self._sliced_mean.result()
diff --git a/official/recommendation/uplift/metrics/label_mean_test.py b/official/recommendation/uplift/metrics/label_mean_test.py
new file mode 100644
index 00000000000..9c4180020a7
--- /dev/null
+++ b/official/recommendation/uplift/metrics/label_mean_test.py
@@ -0,0 +1,206 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for label_mean."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import label_mean
+
+
+class LabelMeanTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ def _get_y_pred(
+ self, is_treatment: tf.Tensor
+ ) -> types.TwoTowerTrainingOutputs:
+ # All tensors except the is_treatment tensor is distracting from the
+ # testing logic.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=tf.ones_like(is_treatment),
+ control_predictions=tf.ones_like(is_treatment),
+ treatment_predictions=tf.ones_like(is_treatment),
+ uplift=tf.ones_like(is_treatment),
+ control_logits=tf.ones_like(is_treatment),
+ treatment_logits=tf.ones_like(is_treatment),
+ true_logits=tf.ones_like(is_treatment),
+ true_predictions=tf.ones_like(is_treatment),
+ is_treatment=is_treatment,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "y_true": tf.constant([0, 1, 5, 6]),
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": None,
+ "expected_result": {
+ "label/mean": 3.0,
+ "label/mean/control": 3.5,
+ "label/mean/treatment": 2.5,
+ },
+ },
+ {
+ "testcase_name": "weighted",
+ "y_true": tf.constant([0, 1, 5, 6, -7]),
+ "is_treatment": tf.constant(
+ [[True], [False], [True], [True], [False]]
+ ),
+ "sample_weight": tf.constant([0.5, 0.5, 0, 0.7, 1.8]),
+ "expected_result": {
+ "label/mean": np.average(
+ np.array([0, 1, 5, 6, -7]),
+ weights=np.array([0.5, 0.5, 0, 0.7, 1.8]),
+ ),
+ "label/mean/control": np.average(
+ np.array([1, -7]), weights=np.array([0.5, 1.8])
+ ),
+ "label/mean/treatment": np.average(
+ np.array([0, 5, 6]), weights=np.array([0.5, 0, 0.7])
+ ),
+ },
+ },
+ {
+ "testcase_name": "only_control",
+ "y_true": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[False], [False], [False]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_result": {
+ "label/mean": 2.5,
+ "label/mean/control": 2.5,
+ "label/mean/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "only_treatment",
+ "y_true": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[True], [True], [True]]),
+ "sample_weight": tf.constant([0, 1, 1]),
+ "expected_result": {
+ "label/mean": 3.0,
+ "label/mean/control": 0.0,
+ "label/mean/treatment": 3.0,
+ },
+ },
+ {
+ "testcase_name": "one_entry",
+ "y_true": tf.constant([2.5]),
+ "is_treatment": tf.constant([True]),
+ "sample_weight": tf.constant([1]),
+ "expected_result": {
+ "label/mean": 2.5,
+ "label/mean/control": 0.0,
+ "label/mean/treatment": 2.5,
+ },
+ },
+ {
+ "testcase_name": "no_entry",
+ "y_true": tf.constant([]),
+ "is_treatment": tf.constant([], dtype=tf.bool),
+ "sample_weight": tf.constant([]),
+ "expected_result": {
+ "label/mean": 0.0,
+ "label/mean/control": 0.0,
+ "label/mean/treatment": 0.0,
+ },
+ },
+ )
+ def test_treatment_sliced_metric(
+ self, y_true, is_treatment, sample_weight, expected_result
+ ):
+ metric = label_mean.LabelMean()
+ y_pred = self._get_y_pred(is_treatment)
+ metric(y_true, y_pred, sample_weight=sample_weight)
+ self.assertEqual(expected_result, metric.result())
+
+ def test_multiple_batches(self):
+ metric = label_mean.LabelMean(name="label")
+
+ metric.update_state(
+ y_true=tf.constant([[1], [2], [4]]),
+ y_pred=self._get_y_pred(tf.constant([[True], [True], [True]])),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.constant([[-3], [0], [5]]),
+ y_pred=self._get_y_pred(tf.constant([[False], [False], [False]])),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.constant([[0], [1], [-5]]),
+ y_pred=self._get_y_pred(tf.constant([[True], [False], [True]])),
+ sample_weight=tf.constant([0.3, 0.25, 0.7]),
+ )
+
+ expected_results = {
+ "label": np.average(
+ np.array([1, 2, 4, -3, 0, 5, 0, 1, -5]),
+ weights=np.array([1, 1, 1, 1, 1, 1, 0.3, 0.25, 0.7]),
+ ),
+ "label/control": np.average(
+ np.array([-3, 0, 5, 1]), weights=np.array([1, 1, 1, 0.25])
+ ),
+ "label/treatment": np.average(
+ np.array([1, 2, 4, 0, -5]), weights=np.array([1, 1, 1, 0.3, 0.7])
+ ),
+ }
+ self.assertEqual(expected_results, metric.result())
+
+ def test_metric_states(self):
+ metric = label_mean.LabelMean()
+
+ expected_initial_result = {
+ "label/mean": 0.0,
+ "label/mean/control": 0.0,
+ "label/mean/treatment": 0.0,
+ }
+ self.assertEqual(expected_initial_result, metric.result())
+
+ metric(
+ y_true=tf.constant([1, 2, 6]),
+ y_pred=self._get_y_pred(tf.constant([[True], [False], [True]])),
+ )
+ self.assertEqual(
+ {
+ "label/mean": 3.0,
+ "label/mean/control": 2.0,
+ "label/mean/treatment": 3.5,
+ },
+ metric.result(),
+ )
+
+ metric.reset_states()
+ self.assertEqual(expected_initial_result, metric.result())
+
+ def test_metric_config(self):
+ metric = label_mean.LabelMean(name="test_name", dtype=tf.float16)
+ y_true = tf.constant([[1], [2], [3], [4]])
+ y_pred = self._get_y_pred(
+ is_treatment=tf.constant([[True], [False], [True], [False]]),
+ )
+ self.assertLayerConfigurable(layer=metric, y_true=y_true, y_pred=y_pred)
+
+ def test_invalid_prediction_tensor_type(self):
+ metric = label_mean.LabelMean()
+
+ with self.assertRaisesRegex(
+ TypeError, "y_pred must be of type `TwoTowerTrainingOutputs`"
+ ):
+ metric.update_state(y_true=tf.ones((3, 1)), y_pred=tf.ones((3, 1)))
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/label_variance.py b/official/recommendation/uplift/metrics/label_variance.py
new file mode 100644
index 00000000000..d48f7c86e05
--- /dev/null
+++ b/official/recommendation/uplift/metrics/label_variance.py
@@ -0,0 +1,104 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for computing the label variance sliced by treatment group."""
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import treatment_sliced_metric
+from official.recommendation.uplift.metrics import variance
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class LabelVariance(tf_keras.metrics.Metric):
+ """Computes the overall and treatment sliced label variance.
+
+ Note that the prediction tensor is expected to be of type
+ `TwoTowerTrainingOutputs`.
+
+ Example standalone usage:
+
+ >>> label_variance = LabelVariance()
+ >>> y_true = tf.constant([1, 2, 3, 10])
+ >>> y_pred = types.TwoTowerTrainingOutputs(
+ ... control_logits=tf.zeros(4),
+ ... treatment_logits=tf.zeros(4),
+ ... true_logits=tf.zeros(4),
+ ... is_treatment=tf.constant([True, False, True, False]),
+ ... )
+ >>> label_variance(y_true=y_true, y_pred=y_pred)
+ {
+ "label/variance": 12.5
+ "label/variance/control": 16.0
+ "label/variance/treatment": 1.0
+ }
+
+ Example usage with the `model.compile()` API:
+
+ >>> model.compile(
+ ... optimizer="sgd",
+ ... loss=TrueLogitsLoss(tf_keras.losses.mean_squared_error),
+ ... metrics=[LabelVariance()]
+ ... )
+ """
+
+ def __init__(self, name: str = "label/variance", **kwargs):
+ """Initializes the instance.
+
+ Args:
+ name: name for the overall label variance metric result. The control and
+ treatment label variances will have "/control" and "/treatment" appended
+ to the result name.
+ **kwargs: other base metric keyword arguments.
+ """
+ super().__init__(name=name, **kwargs)
+ self._sliced_variance = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=variance.Variance(name=name, **kwargs)
+ )
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs,
+ sample_weight: tf.Tensor | None = None,
+ ):
+ """Updates the overall, control and treatment label variances.
+
+ Args:
+ y_true: tensor labels.
+ y_pred: two tower training outputs. The treatment indicator tensor is used
+ to slice the labels into control and treatment groups.
+ sample_weight: optional sample weight to compute weighted label variances.
+ If given, the sample weight will also be sliced by the treatment
+ indicator tensor to compute the weighted control and treatment label
+ variances.
+
+ Raises:
+ TypeError: if y_pred is not of type `TwoTowerTrainingOutputs`.
+ """
+ if not isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise TypeError(
+ "y_pred must be of type `TwoTowerTrainingOutputs` but got type"
+ f" {type(y_pred)} instead."
+ )
+
+ self._sliced_variance.update_state(
+ values=y_true,
+ is_treatment=y_pred.is_treatment,
+ sample_weight=sample_weight,
+ )
+
+ def result(self) -> dict[str, tf.Tensor]:
+ return self._sliced_variance.result()
diff --git a/official/recommendation/uplift/metrics/label_variance_test.py b/official/recommendation/uplift/metrics/label_variance_test.py
new file mode 100644
index 00000000000..c3252a4d8b8
--- /dev/null
+++ b/official/recommendation/uplift/metrics/label_variance_test.py
@@ -0,0 +1,214 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for label_variance."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import label_variance
+
+
+def _weighted_var(values, weights):
+ values = np.array(values)
+ weights = np.array(weights)
+ weighted_mean = np.average(values, weights=weights)
+ return np.average((values - weighted_mean) ** 2, weights=weights)
+
+
+class LabelVarianceTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ def _get_y_pred(
+ self, is_treatment: tf.Tensor
+ ) -> types.TwoTowerTrainingOutputs:
+ # Only the is_treatment tensor is required for testing.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=tf.ones_like(is_treatment),
+ control_predictions=tf.ones_like(is_treatment),
+ treatment_predictions=tf.ones_like(is_treatment),
+ uplift=tf.ones_like(is_treatment),
+ control_logits=tf.ones_like(is_treatment),
+ treatment_logits=tf.ones_like(is_treatment),
+ true_logits=tf.ones_like(is_treatment),
+ true_predictions=tf.ones_like(is_treatment),
+ is_treatment=is_treatment,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "y_true": tf.constant([0, 1, 5, 6]),
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": None,
+ "expected_result": {
+ "label/variance": np.array([0, 1, 5, 6]).var(),
+ "label/variance/control": np.array([1, 6]).var(),
+ "label/variance/treatment": np.array([0, 5]).var(),
+ },
+ },
+ {
+ "testcase_name": "weighted",
+ "y_true": tf.constant([0, 1, 5, 6, -7]),
+ "is_treatment": tf.constant(
+ [[True], [False], [True], [True], [False]]
+ ),
+ "sample_weight": tf.constant([0.5, 0.5, 0, 0.7, 1.8]),
+ "expected_result": {
+ "label/variance": _weighted_var(
+ values=[0, 1, 5, 6, -7], weights=[0.5, 0.5, 0, 0.7, 1.8]
+ ),
+ "label/variance/control": _weighted_var(
+ values=[1, -7], weights=[0.5, 1.8]
+ ),
+ "label/variance/treatment": _weighted_var(
+ values=[0, 5, 6], weights=[0.5, 0, 0.7]
+ ),
+ },
+ },
+ {
+ "testcase_name": "only_control",
+ "y_true": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[False], [False], [False]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_result": {
+ "label/variance": np.array([0, 5]).var(),
+ "label/variance/control": np.array([0, 5]).var(),
+ "label/variance/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "only_treatment",
+ "y_true": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[True], [True], [True]]),
+ "sample_weight": tf.constant([0, 1, 1]),
+ "expected_result": {
+ "label/variance": np.array([1, 5]).var(),
+ "label/variance/control": 0.0,
+ "label/variance/treatment": np.array([1, 5]).var(),
+ },
+ },
+ {
+ "testcase_name": "one_entry",
+ "y_true": tf.constant([2.5]),
+ "is_treatment": tf.constant([True]),
+ "sample_weight": tf.constant([1]),
+ "expected_result": {
+ "label/variance": 0.0,
+ "label/variance/control": 0.0,
+ "label/variance/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "no_entry",
+ "y_true": tf.constant([]),
+ "is_treatment": tf.constant([], dtype=tf.bool),
+ "sample_weight": tf.constant([]),
+ "expected_result": {
+ "label/variance": 0.0,
+ "label/variance/control": 0.0,
+ "label/variance/treatment": 0.0,
+ },
+ },
+ )
+ def test_label_variance_computes_sliced_variances(
+ self, y_true, is_treatment, sample_weight, expected_result
+ ):
+ metric = label_variance.LabelVariance()
+ y_pred = self._get_y_pred(is_treatment)
+ metric(y_true, y_pred, sample_weight=sample_weight)
+ self.assertEqual(expected_result, metric.result())
+
+ def test_multiple_update_batches_returns_aggregated_label_variances(self):
+ metric = label_variance.LabelVariance(name="label")
+
+ metric.update_state(
+ y_true=tf.constant([[1], [2], [4]]),
+ y_pred=self._get_y_pred(tf.constant([[True], [True], [True]])),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.constant([[-3], [0], [5]]),
+ y_pred=self._get_y_pred(tf.constant([[False], [False], [False]])),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.constant([[0], [1], [-5]]),
+ y_pred=self._get_y_pred(tf.constant([[True], [False], [True]])),
+ sample_weight=tf.constant([0.3, 0.25, 0.7]),
+ )
+
+ expected_results = {
+ "label": _weighted_var(
+ values=[1, 2, 4, -3, 0, 5, 0, 1, -5],
+ weights=[1, 1, 1, 1, 1, 1, 0.3, 0.25, 0.7],
+ ),
+ "label/control": _weighted_var(
+ values=[-3, 0, 5, 1], weights=[1, 1, 1, 0.25]
+ ),
+ "label/treatment": _weighted_var(
+ values=[1, 2, 4, 0, -5], weights=[1, 1, 1, 0.3, 0.7]
+ ),
+ }
+ self.assertEqual(expected_results, metric.result())
+
+ def test_initial_and_reset_state_return_zero_label_variances(self):
+ metric = label_variance.LabelVariance()
+
+ expected_initial_result = {
+ "label/variance": 0.0,
+ "label/variance/control": 0.0,
+ "label/variance/treatment": 0.0,
+ }
+ self.assertEqual(expected_initial_result, metric.result())
+
+ metric(
+ y_true=tf.constant([1, 2, 6]),
+ y_pred=self._get_y_pred(tf.constant([[True], [False], [True]])),
+ )
+ self.assertEqual(
+ {
+ "label/variance": np.array([1, 2, 6]).var(),
+ "label/variance/control": 0.0,
+ "label/variance/treatment": np.array([1, 6]).var(),
+ },
+ metric.result(),
+ )
+
+ metric.reset_states()
+ self.assertEqual(expected_initial_result, metric.result())
+
+ def test_metric_config_is_serializable(self):
+ metric = label_variance.LabelVariance(name="test_name", dtype=tf.float16)
+ y_true = tf.constant([[1], [2], [3], [4]])
+ y_pred = self._get_y_pred(
+ is_treatment=tf.constant([[True], [False], [True], [False]]),
+ )
+ self.assertLayerConfigurable(
+ layer=metric, y_true=y_true, y_pred=y_pred, serializable=True
+ )
+
+ def test_invalid_prediction_tensor_type_raises_type_error(self):
+ metric = label_variance.LabelVariance()
+
+ with self.assertRaisesRegex(
+ TypeError, "y_pred must be of type `TwoTowerTrainingOutputs`"
+ ):
+ metric.update_state(y_true=tf.ones((3, 1)), y_pred=tf.ones((3, 1)))
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/loss_metric.py b/official/recommendation/uplift/metrics/loss_metric.py
new file mode 100644
index 00000000000..8665b303873
--- /dev/null
+++ b/official/recommendation/uplift/metrics/loss_metric.py
@@ -0,0 +1,202 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for computing a loss sliced by treatment group."""
+
+from __future__ import annotations
+
+import inspect
+from typing import Any, Callable
+
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import treatment_sliced_metric
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class LossMetric(tf_keras.metrics.Metric):
+ """Computes a loss sliced by treatment group.
+
+ Example standalone usage:
+
+ >>> sliced_loss = LossMetric(tf_keras.losses.mean_squared_error)
+ >>> y_true = tf.constant([0, 0, 2, 2])
+ >>> y_pred = types.TwoTowerTrainingOutputs(
+ ... true_logits=tf.constant([1, 2, 3, 4])
+ ... is_treatment=tf.constant([True, False, True, False]),
+ ... )
+ >>> sliced_loss(y_true=y_true, y_pred=y_pred)
+ {
+ "loss": 2.5
+ "loss/control": 4.0
+ "loss/treatment": 1.0
+ }
+
+ Example usage with the `model.compile()` API:
+
+ >>> model.compile(
+ ... optimizer="sgd",
+ ... loss=TrueLogitsLoss(tf_keras.losses.mean_squared_error),
+ ... metrics=[LossMetric(tf_keras.losses.mean_squared_error)]
+ ... )
+ """
+
+ def __init__(
+ self,
+ loss_fn: (
+ Callable[[tf.Tensor, tf.Tensor], tf.Tensor] | tf_keras.metrics.Metric
+ ),
+ from_logits: bool = True,
+ slice_by_treatment: bool = True,
+ name: str = "loss",
+ dtype: tf.DType = tf.float32,
+ **loss_fn_kwargs,
+ ):
+ """Initializes the instance.
+
+ Args:
+ loss_fn: The loss function or Keras metric to apply with call signature
+ `__call__(y_true: tf,Tensor, y_pred: tf.Tensor, **loss_fn_kwargs)`. Note
+ that the `loss_fn_kwargs` will not be passed to the `__call__` method if
+ `loss_fn` is a Keras metric.
+ from_logits: When `y_pred` is of type `TwoTowerTrainingOutputs`, specifies
+ whether the true logits or true predictions should be used to compute
+ the loss (defaults to using the true logits). Othwerwise, this argument
+ will be ignored if `y_pred` is of type `tf.Tensor`.
+ slice_by_treatment: Specifies whether the loss should be sliced by the
+ treatment indicator tensor. If `True`, `loss_fn` will be wrapped in a
+ `TreatmentSlicedMetric` to report the loss values sliced by the
+ treatment group.
+ name: Optional name for the instance. If `loss_fn` is a Keras metric then
+ its name will be used instead.
+ dtype: Optional data type for the instance. If `loss_fn` is a Keras metric
+ then its `dtype` will be used instead.
+ **loss_fn_kwargs: The keyword arguments that are passed on to `loss_fn`.
+ These arguments will be ignored if `loss_fn` is a Keras metric.
+ """
+ # Do not accept Loss objects as they reduce tensors before weighting.
+ if isinstance(loss_fn, tf_keras.losses.Loss):
+ raise TypeError(
+ "`loss_fn` cannot be a Keras `Loss` object, pass a non-reducing loss"
+ " function or a metric instance instead."
+ )
+
+ if isinstance(loss_fn, tf_keras.metrics.Metric):
+ name = loss_fn.name
+ dtype = loss_fn.dtype
+
+ super().__init__(name=name, dtype=dtype)
+
+ self._loss_fn = loss_fn
+ self._from_logits = from_logits
+ self._loss_fn_kwargs = loss_fn_kwargs
+ self._slice_by_treatment = slice_by_treatment
+
+ if isinstance(loss_fn, tf_keras.metrics.Metric):
+ metric_from_logits = loss_fn.get_config().get("from_logits", from_logits)
+ if from_logits != metric_from_logits:
+ raise ValueError(
+ f"Value passed to `from_logits` ({from_logits}) is conflicting with"
+ " the `from_logits` value passed to the `loss_fn` metric"
+ f" ({metric_from_logits}). Ensure that they have the same value."
+ )
+ loss_metric = loss_fn
+
+ else:
+ if "from_logits" in inspect.signature(loss_fn).parameters:
+ self._loss_fn_kwargs.update({"from_logits": from_logits})
+ loss_metric = tf_keras.metrics.Mean(name=name, dtype=dtype)
+
+ if slice_by_treatment:
+ self._loss = treatment_sliced_metric.TreatmentSlicedMetric(loss_metric)
+ else:
+ self._loss = loss_metric
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor | np.ndarray,
+ sample_weight: tf.Tensor | None = None,
+ ):
+ """Updates the overall, control and treatment losses.
+
+ Args:
+ y_true: A `tf.Tensor` with the targets.
+ y_pred: Model outputs. If of type `TwoTowerTrainingOutputs`, the treatment
+ indicator tensor is used to slice the true logits or true predictions
+ into control and treatment losses.
+ sample_weight: Optional sample weight to compute weighted losses. If
+ given, the sample weight will also be sliced by the treatment indicator
+ tensor to compute the weighted control and treatment losses.
+
+ Raises:
+ TypeError: if `y_pred` is not of type `TwoTowerTrainingOutputs`.
+ """
+ if isinstance(y_pred, (tf.Tensor, np.ndarray)):
+ if self._slice_by_treatment:
+ raise ValueError(
+ "`slice_by_treatment` must be False when y_pred is a `tf.Tensor` or"
+ " `np.ndarray`."
+ )
+ pred = y_pred
+ elif isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ pred = (
+ y_pred.true_logits if self._from_logits else y_pred.true_predictions
+ )
+ else:
+ raise TypeError(
+ "y_pred must be of type `TwoTowerTrainingOutputs`, `tf.Tensor` or"
+ f" `np.ndarray` but got type {type(y_pred)} instead."
+ )
+
+ is_treatment = {}
+ if self._slice_by_treatment:
+ is_treatment["is_treatment"] = y_pred.is_treatment
+
+ if isinstance(self._loss_fn, tf_keras.metrics.Metric):
+ self._loss.update_state(
+ y_true,
+ y_pred=pred, # pyrefly: ignore[unexpected-keyword]
+ sample_weight=sample_weight,
+ **is_treatment,
+ )
+ else:
+ self._loss.update_state(
+ values=self._loss_fn(y_true, pred, **self._loss_fn_kwargs), # pyrefly: ignore[bad-argument-type]
+ sample_weight=sample_weight,
+ **is_treatment,
+ )
+
+ def result(self) -> tf.Tensor | dict[str, tf.Tensor]:
+ return self._loss.result()
+
+ def reset_state(self):
+ self._loss.reset_state()
+
+ def get_config(self) -> dict[str, Any]:
+ config = super().get_config()
+ config["loss_fn"] = tf_keras.utils.serialize_keras_object(self._loss_fn)
+ config["from_logits"] = self._from_logits
+ config["slice_by_treatment"] = self._slice_by_treatment
+ config.update(self._loss_fn_kwargs)
+ return config
+
+ @classmethod
+ def from_config(cls, config: dict[str, Any]) -> LossMetric:
+ config["loss_fn"] = tf_keras.utils.deserialize_keras_object(
+ config["loss_fn"]
+ )
+ return cls(**config)
diff --git a/official/recommendation/uplift/metrics/loss_metric_test.py b/official/recommendation/uplift/metrics/loss_metric_test.py
new file mode 100644
index 00000000000..868ea3c68c2
--- /dev/null
+++ b/official/recommendation/uplift/metrics/loss_metric_test.py
@@ -0,0 +1,451 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for loss_metric."""
+
+from typing import Callable
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import loss_metric
+
+
+class LossMetricTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ def _get_outputs(
+ self,
+ true_logits: tf.Tensor,
+ true_predictions: tf.Tensor,
+ is_treatment: tf.Tensor,
+ ) -> types.TwoTowerTrainingOutputs:
+ # Only the true_logits, true_predictions and is_treatment tensors are used
+ # for testing.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=tf.ones_like(is_treatment),
+ control_predictions=tf.ones_like(is_treatment),
+ treatment_predictions=tf.ones_like(is_treatment),
+ uplift=tf.ones_like(is_treatment),
+ control_logits=tf.ones_like(is_treatment),
+ treatment_logits=tf.ones_like(is_treatment),
+ true_logits=true_logits,
+ true_predictions=true_predictions,
+ is_treatment=is_treatment,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "loss_fn": tf_keras.losses.mean_squared_error,
+ "from_logits": False,
+ "y_true": tf.constant([[0], [0], [2], [2]]),
+ "y_pred": tf.constant([[1], [2], [3], [4]]),
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": None,
+ "expected_losses": {
+ "loss": 2.5,
+ "loss/control": 4.0,
+ "loss/treatment": 1.0,
+ },
+ },
+ {
+ "testcase_name": "unweighted_metric",
+ "loss_fn": tf_keras.metrics.MeanSquaredError(name="loss"),
+ "from_logits": False,
+ "y_true": tf.constant([[0], [0], [2], [2]]),
+ "y_pred": tf.constant([[1], [2], [3], [4]]),
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": None,
+ "expected_losses": {
+ "loss": 2.5,
+ "loss/control": 4.0,
+ "loss/treatment": 1.0,
+ },
+ },
+ {
+ "testcase_name": "weighted",
+ "loss_fn": tf_keras.losses.mean_absolute_error,
+ "from_logits": False,
+ "y_true": tf.constant([[0], [0], [2], [7]]),
+ "y_pred": tf.constant([[1], [2], [3], [4]]),
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": tf.constant([0.5, 0.5, 0.7, 1.8]),
+ "expected_losses": {
+ "loss": np.average([1, 2, 1, 3], weights=[0.5, 0.5, 0.7, 1.8]),
+ "loss/control": np.average([2, 3], weights=[0.5, 1.8]),
+ "loss/treatment": 1.0,
+ },
+ },
+ {
+ "testcase_name": "weighted_keras_metric",
+ "loss_fn": tf_keras.metrics.MeanAbsoluteError(name="loss"),
+ "from_logits": False,
+ "y_true": tf.constant([[0], [0], [2], [7]]),
+ "y_pred": tf.constant([[1], [2], [3], [4]]),
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": tf.constant([[0.5], [0.5], [0.7], [1.8]]),
+ "expected_losses": {
+ "loss": np.average([1, 2, 1, 3], weights=[0.5, 0.5, 0.7, 1.8]),
+ "loss/control": np.average([2, 3], weights=[0.5, 1.8]),
+ "loss/treatment": 1.0,
+ },
+ },
+ {
+ "testcase_name": "only_control",
+ "loss_fn": tf_keras.metrics.mean_squared_error,
+ "from_logits": False,
+ "y_true": tf.constant([[0], [1], [5]]),
+ "y_pred": tf.constant([[1], [2], [5]]),
+ "is_treatment": tf.constant([[False], [False], [False]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_losses": {
+ "loss": 0.5,
+ "loss/control": 0.5,
+ "loss/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "only_control_metric",
+ "loss_fn": tf_keras.metrics.MeanSquaredError(name="loss"),
+ "from_logits": False,
+ "y_true": tf.constant([[0], [1], [5]]),
+ "y_pred": tf.constant([[1], [2], [5]]),
+ "is_treatment": tf.constant([[False], [False], [False]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_losses": {
+ "loss": 0.5,
+ "loss/control": 0.5,
+ "loss/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "only_treatment",
+ "loss_fn": tf_keras.metrics.mean_absolute_error,
+ "from_logits": False,
+ "y_true": tf.constant([[0], [1], [5]]),
+ "y_pred": tf.constant([[1], [2], [5]]),
+ "is_treatment": tf.constant([[True], [True], [True]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_losses": {
+ "loss": 0.5,
+ "loss/control": 0.0,
+ "loss/treatment": 0.5,
+ },
+ },
+ {
+ "testcase_name": "only_treatment_metric",
+ "loss_fn": tf_keras.metrics.MeanAbsoluteError(name="loss"),
+ "from_logits": False,
+ "y_true": tf.constant([[0], [1], [5]]),
+ "y_pred": tf.constant([[1], [2], [5]]),
+ "is_treatment": tf.constant([[True], [True], [True]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_losses": {
+ "loss": 0.5,
+ "loss/control": 0.0,
+ "loss/treatment": 0.5,
+ },
+ },
+ {
+ "testcase_name": "one_entry",
+ "loss_fn": tf.nn.log_poisson_loss,
+ "from_logits": True,
+ "y_true": tf.constant([0.2]),
+ "y_pred": tf.constant([2.5]),
+ "is_treatment": tf.constant([True]),
+ "sample_weight": tf.constant([1]),
+ "expected_losses": {
+ "loss": tf.nn.log_poisson_loss(
+ tf.constant([0.2]), tf.constant([2.5])
+ ),
+ "loss/control": 0.0,
+ "loss/treatment": tf.nn.log_poisson_loss(
+ tf.constant([0.2]), tf.constant([2.5])
+ ),
+ },
+ },
+ {
+ "testcase_name": "no_entry",
+ "loss_fn": tf_keras.losses.binary_crossentropy,
+ "from_logits": True,
+ "y_true": tf.constant([[]]),
+ "y_pred": tf.constant([[]]),
+ "is_treatment": tf.constant([[]]),
+ "sample_weight": tf.constant([[]]),
+ "expected_losses": {
+ "loss": 0.0,
+ "loss/control": 0.0,
+ "loss/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "no_entry_metric",
+ "loss_fn": tf_keras.metrics.BinaryCrossentropy(name="loss"),
+ "from_logits": False,
+ "y_true": tf.constant([[]]),
+ "y_pred": tf.constant([[]]),
+ "is_treatment": tf.constant([[]]),
+ "sample_weight": tf.constant([[]]),
+ "expected_losses": {
+ "loss": 0.0,
+ "loss/control": 0.0,
+ "loss/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "auc_metric",
+ "loss_fn": tf_keras.metrics.AUC(from_logits=True, name="loss"),
+ "from_logits": True,
+ "y_true": tf.constant([[0], [0], [1], [1]]),
+ "y_pred": tf.constant([[0], [0.5], [0.3], [0.9]]),
+ "is_treatment": tf.constant([[1], [1], [1], [1]]),
+ "sample_weight": None,
+ "expected_losses": {
+ "loss": 0.75,
+ "loss/control": 0.0,
+ "loss/treatment": 0.75,
+ },
+ },
+ {
+ "testcase_name": "loss_fn_with_from_logits",
+ "loss_fn": tf_keras.losses.binary_crossentropy,
+ "from_logits": True,
+ "y_true": tf.constant([[0.0, 1.0]]),
+ "y_pred": tf.constant([[0.0, 1.0]]),
+ "is_treatment": tf.constant([[0], [0]]),
+ "sample_weight": None,
+ "expected_losses": {
+ "loss": 0.50320446,
+ "loss/control": 0.50320446,
+ "loss/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "no_treatment_slice",
+ "loss_fn": tf_keras.losses.binary_crossentropy,
+ "from_logits": True,
+ "y_true": tf.constant([[0.0, 1.0]]),
+ "y_pred": tf.constant([[0.0, 1.0]]),
+ "is_treatment": tf.constant([[0], [0]]),
+ "sample_weight": None,
+ "expected_losses": 0.50320446,
+ "slice_by_treatment": False,
+ },
+ {
+ "testcase_name": "no_treatment_slice_metric",
+ "loss_fn": tf_keras.metrics.BinaryCrossentropy(from_logits=False),
+ "from_logits": False,
+ "y_true": tf.constant([[0.0, 1.0]]),
+ "y_pred": tf.constant([[0.0, 1.0]]),
+ "is_treatment": tf.constant([[0], [0]]),
+ "sample_weight": None,
+ "expected_losses": 0,
+ "slice_by_treatment": False,
+ },
+ )
+ def test_metric_computes_sliced_losses(
+ self,
+ loss_fn: Callable[[tf.Tensor, tf.Tensor], tf.Tensor],
+ from_logits: bool,
+ y_true: tf.Tensor,
+ y_pred: tf.Tensor,
+ is_treatment: tf.Tensor,
+ sample_weight: tf.Tensor | None,
+ expected_losses: float | dict[str, float],
+ slice_by_treatment: bool = True,
+ ):
+ if from_logits:
+ true_logits = y_pred
+ true_predictions = tf.zeros_like(y_pred) # Irrelevant for testing.
+ else:
+ true_logits = tf.zeros_like(y_pred) # Irrelevant for testing.
+ true_predictions = y_pred
+
+ metric = loss_metric.LossMetric(
+ loss_fn=loss_fn,
+ from_logits=from_logits,
+ slice_by_treatment=slice_by_treatment,
+ )
+ outputs = self._get_outputs(
+ true_logits=true_logits,
+ true_predictions=true_predictions,
+ is_treatment=is_treatment,
+ )
+ metric(y_true, outputs, sample_weight=sample_weight)
+ self.assertEqual(expected_losses, metric.result())
+
+ def test_metric_with_y_pred_tensor(self):
+ y_true = tf.constant([[0], [0], [2], [7]])
+ y_pred = tf.constant([[1], [2], [3], [4]])
+ sample_weight = tf.constant([[0.5], [0.5], [0.7], [1.8]])
+
+ metric = loss_metric.LossMetric(
+ loss_fn=tf_keras.metrics.mae, slice_by_treatment=False
+ )
+ metric(y_true, y_pred, sample_weight)
+
+ expected_loss = np.average([1, 2, 1, 3], weights=[0.5, 0.5, 0.7, 1.8])
+ self.assertAllClose(expected_loss, metric.result())
+
+ def test_multiple_update_batches_returns_aggregated_sliced_losses(self):
+ metric = loss_metric.LossMetric(
+ loss_fn=tf_keras.losses.mean_absolute_error,
+ from_logits=False,
+ name="mean_absolute_error",
+ )
+
+ metric.update_state(
+ y_true=tf.constant([[0], [0], [2]]),
+ y_pred=self._get_outputs(
+ true_logits=tf.constant([[1], [2], [4]]),
+ true_predictions=tf.constant([[1], [2], [4]]),
+ is_treatment=tf.constant([[True], [True], [True]]),
+ ),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.constant([[0], [1], [5]]),
+ y_pred=self._get_outputs(
+ true_logits=tf.constant([[-3], [0], [5]]),
+ true_predictions=tf.constant([[-3], [0], [5]]),
+ is_treatment=tf.constant([[False], [False], [False]]),
+ ),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.constant([[2], [3], [-4]]),
+ y_pred=self._get_outputs(
+ true_logits=tf.constant([[0], [1], [-5]]),
+ true_predictions=tf.constant([[0], [1], [-5]]),
+ is_treatment=tf.constant([[True], [False], [True]]),
+ ),
+ sample_weight=tf.constant([0.3, 0.25, 0.7]),
+ )
+
+ expected_results = {
+ "mean_absolute_error": np.average(
+ [1, 2, 2, 3, 1, 0, 2, 2, 1],
+ weights=[1, 1, 1, 1, 1, 1, 0.3, 0.25, 0.7],
+ ),
+ "mean_absolute_error/control": np.average(
+ [3, 1, 0, 2], weights=[1, 1, 1, 0.25]
+ ),
+ "mean_absolute_error/treatment": np.average(
+ [1, 2, 2, 2, 1], weights=[1, 1, 1, 0.3, 0.7]
+ ),
+ }
+ self.assertEqual(expected_results, metric.result())
+
+ def test_initial_and_reset_state_return_zero_losses(self):
+ metric = loss_metric.LossMetric(
+ tf_keras.losses.binary_crossentropy, from_logits=True
+ )
+
+ expected_initial_result = {
+ "loss": 0.0,
+ "loss/control": 0.0,
+ "loss/treatment": 0.0,
+ }
+ self.assertEqual(expected_initial_result, metric.result())
+
+ metric.update_state(
+ y_true=tf.constant([[1], [1], [0]]),
+ y_pred=self._get_outputs(
+ true_logits=tf.constant([[2.3], [0.5], [-3.3]]),
+ true_predictions=tf.random.normal((3, 1)), # Will not be used.
+ is_treatment=tf.constant([[True], [False], [True]]),
+ ),
+ )
+ metric.reset_states()
+ self.assertEqual(expected_initial_result, metric.result())
+
+ @parameterized.product(
+ loss_fn=(
+ tf_keras.losses.binary_crossentropy,
+ tf_keras.metrics.BinaryCrossentropy(
+ from_logits=True, name="bce_loss"
+ ),
+ ),
+ slice_by_treatment=(True, False),
+ )
+ def test_metric_is_configurable(
+ self,
+ loss_fn: (
+ Callable[[tf.Tensor, tf.Tensor], tf.Tensor] | tf_keras.metrics.Metric
+ ),
+ slice_by_treatment: bool,
+ ):
+ metric = loss_metric.LossMetric(
+ loss_fn,
+ from_logits=True,
+ slice_by_treatment=slice_by_treatment,
+ name="bce_loss",
+ )
+ self.assertLayerConfigurable(
+ layer=metric,
+ y_true=tf.constant([[1], [1], [0]]),
+ y_pred=self._get_outputs(
+ true_logits=tf.constant([[2.3], [0.5], [-3.3]]),
+ true_predictions=tf.constant([[2.3], [0.5], [-3.3]]),
+ is_treatment=tf.constant([[True], [False], [True]]),
+ ),
+ serializable=True,
+ )
+
+ def test_invalid_prediction_tensor_type_raises_type_error(self):
+ metric = loss_metric.LossMetric(
+ tf_keras.metrics.mean_absolute_percentage_error
+ )
+ y_true = tf.ones((3, 1))
+ y_pred = types.TwoTowerNetworkOutputs(
+ shared_embedding=tf.ones((3, 5)),
+ control_logits=tf.ones((3, 1)),
+ treatment_logits=tf.ones((3, 1)),
+ )
+ with self.assertRaisesRegex(
+ TypeError,
+ "y_pred must be of type `TwoTowerTrainingOutputs`, `tf.Tensor` or"
+ " `np.ndarray`",
+ ):
+ metric.update_state(y_true=y_true, y_pred=y_pred)
+
+ def test_slice_by_treatment_with_y_pred_tensor_raises_error(self):
+ metric = loss_metric.LossMetric(
+ tf_keras.metrics.mae, slice_by_treatment=True
+ )
+ with self.assertRaisesRegex(
+ ValueError,
+ "`slice_by_treatment` must be False when y_pred is a `tf.Tensor`.",
+ ):
+ metric.update_state(y_true=tf.ones((3, 1)), y_pred=tf.ones((3, 1)))
+
+ def test_passing_loss_object_raises_error(self):
+ with self.assertRaisesRegex(
+ TypeError, "`loss_fn` cannot be a Keras `Loss` object"
+ ):
+ loss_metric.LossMetric(loss_fn=tf_keras.losses.MeanAbsoluteError())
+
+ def test_conflicting_from_logits_values_raises_error(self):
+ with self.assertRaises(ValueError):
+ loss_metric.LossMetric(
+ loss_fn=tf_keras.metrics.BinaryCrossentropy(from_logits=True),
+ from_logits=False,
+ )
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/metric_configs.py b/official/recommendation/uplift/metrics/metric_configs.py
new file mode 100644
index 00000000000..eb7f06dbe6c
--- /dev/null
+++ b/official/recommendation/uplift/metrics/metric_configs.py
@@ -0,0 +1,47 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Metric configurations for TF Model Garden."""
+
+from collections.abc import Mapping
+import dataclasses
+from typing import Any
+
+from official.core.config_definitions import base_config
+
+
+@dataclasses.dataclass(kw_only=True)
+class SlicedMetricConfig(base_config.Config):
+ """Sliced metric configuration.
+
+ Attributes:
+ slicing_feature: The feature whose values to slice the metric on. Required.
+ slicing_spec: A mapping from the name of the slice to the value to slice on.
+ The name will be displayed on TB. Required.
+ slicing_feature_dtype: Optional dtype to cast the slicing feature and the
+ values to slice on.
+ """
+
+ slicing_feature: str | None = None
+ slicing_spec: Mapping[str, int] | None = None
+ slicing_feature_dtype: str | None = None
+
+ def __post_init__( # pyrefly: ignore[bad-function-definition]
+ self, default_params: dict[str, Any], restrictions: list[str]
+ ):
+ if not restrictions:
+ restrictions = ['slicing_feature != None', 'slicing_spec != None']
+ super().__post_init__(
+ default_params=default_params, restrictions=restrictions
+ )
diff --git a/official/recommendation/uplift/metrics/poisson_metrics.py b/official/recommendation/uplift/metrics/poisson_metrics.py
new file mode 100644
index 00000000000..f73b58cb762
--- /dev/null
+++ b/official/recommendation/uplift/metrics/poisson_metrics.py
@@ -0,0 +1,407 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Poisson regression metrics."""
+
+from __future__ import annotations
+
+from typing import Any
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import loss_metric
+from official.recommendation.uplift.metrics import treatment_sliced_metric
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class LogLoss(loss_metric.LossMetric):
+ """Computes the (weighted) poisson log loss sliced by treatment group.
+
+ Given labels `y` and the model's predictions `x`, the loss is computed as:
+ `loss = x - y * log(x) + [y * log(y) - y + 0.5 * log(2 * pi * y)]`
+
+ Note that from a numerical perspective it is preferred to compute the loss
+ from the model's logits as opposed to directly using its predictions, where
+ `logits = log(x) = log(prediction)`. In this case the loss is computed as:
+ `loss = exp(logits) - y * logits + [y * log(y) - y + 0.5 * log(2 * pi * y)]`
+
+ Example standalone usage:
+
+ >>> poisson_loss = poisson_metrics.LogLoss()
+ >>> y_true = tf.constant([[1.0], [0.0]])
+ >>> y_pred = types.TwoTowerTrainingOutputs(
+ true_logits=tf.constant([[1.0], [0.0]]),
+ is_treatment=tf.constant([[1], [0]]),
+ )
+ >>> poisson_loss(y_true=y_true, y_pred=y_pred)
+ {
+ "poisson_log_loss/treatment": 1.7182817 # exp(1) - 1 * 1
+ "poisson_log_loss/control": 1.0 # exp(0) - 0 * 0
+ "poisson_log_loss": 1.3591409 # (1.7182817 + 1.0) / 2
+ }
+
+ Example usage with the `model.compile()` API:
+
+ >>> model.compile(
+ ... optimizer="sgd",
+ ... loss=TrueLogitsLoss(tf.nn.log_poisson_loss),
+ ... metrics=[poisson_metrics.LogLoss()]
+ ... )
+ """
+
+ def __init__(
+ self,
+ from_logits: bool = True,
+ compute_full_loss: bool = False,
+ slice_by_treatment: bool = True,
+ name: str = "poisson_log_loss",
+ dtype: tf.DType = tf.float32,
+ ):
+ """Initializes the instance.
+
+ Args:
+ from_logits: When `y_pred` is of type `tf.Tensor`, specifies whether
+ `y_pred` represents the model's logits or predictions. Otherwise, when
+ `y_pred` is of type `TwoTowerTrainingOutputs`, set this to `True` in
+ order to compute the loss using the true logits.
+ compute_full_loss: Specifies whether the full log loss will be computed.
+ If `True`, the expression `[y_true * log(y_true) - y_true + 0.5 * log(2
+ * pi * y_true)]` will be added to the loss, otherwise the loss will be
+ computed solely by the expression `[y_pred - y_true * log(y_pred)]`.
+ slice_by_treatment: Specifies whether the loss should be sliced by the
+ treatment indicator tensor. If `True`, the metric's result will return
+ the loss values sliced by the treatment group. Note that this can only
+ be set to `True` when `y_pred` is of type `TwoTowerTrainingOutputs`.
+ name: Optional name for the instance.
+ dtype: Optional data type for the instance.
+ """
+ super().__init__(
+ loss_fn=tf.nn.log_poisson_loss,
+ from_logits=from_logits,
+ compute_full_loss=compute_full_loss,
+ slice_by_treatment=slice_by_treatment,
+ name=name,
+ dtype=dtype,
+ )
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor,
+ sample_weight: tf.Tensor | None = None,
+ ):
+ if not self._from_logits:
+ if isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise ValueError(
+ "`from_logits` must be set to `True` when `y_pred` is of type"
+ " TwoTowerTrainingOutputs. Note that the true logits and true"
+ " predictions are assumed to be linked to each other through the"
+ " log link function: `true_logits = tf.math.log(true_predictions)."
+ )
+ y_pred = tf.math.log(y_pred)
+
+ super().update_state(y_true, y_pred, sample_weight)
+
+ def get_config(self) -> dict[str, Any]:
+ config = super().get_config()
+ del config["loss_fn"]
+ return config
+
+ @classmethod
+ def from_config(cls, config: dict[str, Any]) -> LogLoss:
+ return cls(**config)
+
+
+def _safe_x_minus_xlogx(x: tf.Tensor) -> tf.Tensor:
+ """Computes x - x * log(x) with 0 as its continuity point when x equals 0."""
+ values = x * (1.0 - tf.math.log(x))
+ return tf.where(tf.equal(x, 0.0), tf.zeros_like(x), values)
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class LogLossMeanBaseline(tf_keras.metrics.Metric):
+ """Computes the (weighted) poisson log loss for a mean predictor."""
+
+ def __init__(
+ self,
+ compute_full_loss: bool = False,
+ slice_by_treatment: bool = True,
+ name: str = "poisson_log_loss_mean_baseline",
+ dtype: tf.DType = tf.float32,
+ ):
+ """Initializes the instance.
+
+ Args:
+ compute_full_loss: Specifies whether to compute the full poisson log loss
+ for the mean predictor or not. Defaults to `False`.
+ slice_by_treatment: Specifies whether the loss should be sliced by the
+ treatment indicator tensor. If `True`, the metric's result will return
+ the loss values sliced by the treatment group. Note that this can only
+ be set to `True` when `y_pred` is of type `TwoTowerTrainingOutputs`.
+ name: Optional name for the instance.
+ dtype: Optional data type for the instance.
+ """
+ super().__init__(name=name, dtype=dtype)
+
+ if compute_full_loss:
+ raise NotImplementedError("Full loss computation is not yet supported.")
+
+ self._compute_full_loss = compute_full_loss
+ self._slice_by_treatment = slice_by_treatment
+
+ if slice_by_treatment:
+ self._mean_label = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.Mean(name=name, dtype=dtype)
+ )
+ else:
+ self._mean_label = tf_keras.metrics.Mean(name=name, dtype=dtype)
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor | None = None,
+ sample_weight: tf.Tensor | None = None,
+ ):
+ is_treatment = {}
+ if self._slice_by_treatment:
+ if not isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise ValueError(
+ "`slice_by_treatment` must be set to `False` when `y_pred` is not"
+ " of type `TwoTowerTrainingOutputs`."
+ )
+ is_treatment["is_treatment"] = y_pred.is_treatment
+
+ self._mean_label.update_state(
+ y_true, sample_weight=sample_weight, **is_treatment
+ )
+
+ def result(self) -> tf.Tensor | dict[str, tf.Tensor]:
+ return tf.nest.map_structure(_safe_x_minus_xlogx, self._mean_label.result())
+
+ def get_config(self) -> dict[str, Any]:
+ config = super().get_config()
+ config["compute_full_loss"] = self._compute_full_loss
+ config["slice_by_treatment"] = self._slice_by_treatment
+ return config
+
+ @classmethod
+ def from_config(cls, config: dict[str, Any]) -> LogLossMeanBaseline:
+ return cls(**config)
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class LogLossMinimum(tf_keras.metrics.Metric):
+ """Computes the minimum achievable (weighted) poisson log loss.
+
+ Given labels `y` and the model's predictions `x`, the minimum loss is obtained
+ when `x` equals `y`. In this case the loss is computed as:
+ `loss = y - y * log(y) + [y * log(y) - y + 0.5 * log(2 * pi * y)]`
+
+ Note that `[y * log(y) - y + 0.5 * log(2 * pi * y)]` is only computed if
+ `compute_full_loss` is set to `True`.
+ """
+
+ def __init__(
+ self,
+ compute_full_loss: bool = False,
+ slice_by_treatment: bool = True,
+ name: str = "poisson_log_loss_minimum",
+ dtype: tf.DType = tf.float32,
+ ):
+ """Initializes the instance.
+
+ Args:
+ compute_full_loss: Specifies whether to compute the full minimum log loss
+ or not. Defaults to `False`.
+ slice_by_treatment: Specifies whether the loss should be sliced by the
+ treatment indicator tensor. If `True`, the metric's result will return
+ the loss values sliced by the treatment group. Note that this can only
+ be set to `True` when `y_pred` is of type `TwoTowerTrainingOutputs`.
+ name: Optional name for the instance.
+ dtype: Optional data type for the instance.
+ """
+ super().__init__(name=name, dtype=dtype)
+
+ if compute_full_loss:
+ raise NotImplementedError("Full loss computation is not yet supported.")
+
+ self._compute_full_loss = compute_full_loss
+ self._slice_by_treatment = slice_by_treatment
+
+ if slice_by_treatment:
+ self._loss = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.Mean(name=name, dtype=dtype)
+ )
+ else:
+ self._loss = tf_keras.metrics.Mean(name=name, dtype=dtype)
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor | None = None,
+ sample_weight: tf.Tensor | None = None,
+ ):
+ is_treatment = {}
+ if self._slice_by_treatment:
+ if not isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise ValueError(
+ "`slice_by_treatment` must be set to `False` when `y_pred` is not"
+ " of type `TwoTowerTrainingOutputs`."
+ )
+ is_treatment["is_treatment"] = y_pred.is_treatment
+
+ self._loss.update_state(
+ _safe_x_minus_xlogx(y_true), sample_weight=sample_weight, **is_treatment
+ )
+
+ def result(self) -> tf.Tensor | dict[str, tf.Tensor]:
+ return self._loss.result()
+
+ def get_config(self) -> dict[str, Any]:
+ config = super().get_config()
+ config["compute_full_loss"] = self._compute_full_loss
+ config["slice_by_treatment"] = self._slice_by_treatment
+ return config
+
+ @classmethod
+ def from_config(cls, config: dict[str, Any]) -> LogLossMinimum:
+ return cls(**config)
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class PseudoRSquared(tf_keras.metrics.Metric):
+ """Computes the pseudo R-squared metric for poisson regression.
+
+ The pseudo R-squared is computed from log likelihoods of three models:
+ 1) LLbaseline: log likelihood of a mean baseline predictor.
+ 2) LLfit: log likelihood of the fitted model.
+ 3) LLmax: maximum achievable log likelihood, which occurs when the predictions
+ equal to the labels.
+
+ The equation that computes the pseudo R-squared is:
+ >>> R_squared = (LLfit - LLbaseline) / (LLmax - LLbaseline)
+ """
+
+ def __init__(
+ self,
+ from_logits: bool = True,
+ slice_by_treatment: bool = True,
+ name: str = "pseudo_r_squared",
+ dtype: tf.DType = tf.float32,
+ ):
+ """Initializes the instance.
+
+ Args:
+ from_logits: When `y_pred` is of type `tf.Tensor`, specifies whether
+ `y_pred` represents the model's logits or predictions. Otherwise, when
+ `y_pred` is of type `TwoTowerTrainingOutputs`, set this to `True` in
+ order to compute the loss using the true logits.
+ slice_by_treatment: Specifies whether the loss should be sliced by the
+ treatment indicator tensor. If `True`, the metric's result will return
+ the loss values sliced by the treatment group. Note that this can only
+ be set to `True` when `y_pred` is of type `TwoTowerTrainingOutputs`.
+ name: Optional name for the instance.
+ dtype: Optional data type for the instance.
+ """
+ super().__init__(name=name, dtype=dtype)
+
+ self._from_logits = from_logits
+ self._slice_by_treatment = slice_by_treatment
+
+ # Since log_loss = -1 * log_likelihood we can just accumulate the losses.
+ loss = LogLoss(
+ from_logits=from_logits,
+ compute_full_loss=False,
+ slice_by_treatment=False,
+ name=name,
+ dtype=dtype,
+ )
+ minimum_loss = LogLossMinimum(
+ compute_full_loss=False,
+ slice_by_treatment=False,
+ name=name,
+ dtype=dtype,
+ )
+ mean_baseline_loss = LogLossMeanBaseline(
+ compute_full_loss=False,
+ slice_by_treatment=False,
+ name=name,
+ dtype=dtype,
+ )
+
+ if slice_by_treatment:
+ self._model_loss = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=loss
+ )
+ self._minimum_loss = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=minimum_loss
+ )
+ self._mean_baseline_loss = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=mean_baseline_loss
+ )
+ else:
+ self._model_loss = loss
+ self._minimum_loss = minimum_loss
+ self._mean_baseline_loss = mean_baseline_loss
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor,
+ sample_weight: tf.Tensor | None = None,
+ ):
+ is_treatment = {}
+ if self._slice_by_treatment:
+ if not isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise ValueError(
+ "`slice_by_treatment` must be set to `False` when `y_pred` is not"
+ " of type `TwoTowerTrainingOutputs`."
+ )
+ is_treatment["is_treatment"] = y_pred.is_treatment
+
+ self._model_loss.update_state(
+ y_true, y_pred=y_pred, sample_weight=sample_weight, **is_treatment
+ )
+ self._minimum_loss.update_state(
+ y_true, y_pred=y_pred, sample_weight=sample_weight, **is_treatment
+ )
+ self._mean_baseline_loss.update_state(
+ y_true, y_pred=y_pred, sample_weight=sample_weight, **is_treatment
+ )
+
+ def result(self) -> tf.Tensor | dict[str, tf.Tensor]:
+ def _pseudo_r_squared(
+ loss_model: tf.Tensor, loss_baseline: tf.Tensor, loss_min: tf.Tensor
+ ) -> tf.Tensor:
+ return tf.math.divide_no_nan(
+ loss_model - loss_baseline, loss_min - loss_baseline
+ )
+
+ return tf.nest.map_structure(
+ _pseudo_r_squared,
+ self._model_loss.result(),
+ self._mean_baseline_loss.result(),
+ self._minimum_loss.result(),
+ )
+
+ def get_config(self) -> dict[str, Any]:
+ config = super().get_config()
+ config["from_logits"] = self._from_logits
+ config["slice_by_treatment"] = self._slice_by_treatment
+ return config
+
+ @classmethod
+ def from_config(cls, config: dict[str, Any]) -> PseudoRSquared:
+ return cls(**config)
diff --git a/official/recommendation/uplift/metrics/poisson_metrics_test.py b/official/recommendation/uplift/metrics/poisson_metrics_test.py
new file mode 100644
index 00000000000..29a07ab9b56
--- /dev/null
+++ b/official/recommendation/uplift/metrics/poisson_metrics_test.py
@@ -0,0 +1,449 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for poisson regression metrics."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import poisson_metrics
+
+
+def _get_two_tower_outputs(
+ is_treatment: tf.Tensor,
+ true_logits: tf.Tensor | None = None,
+) -> types.TwoTowerTrainingOutputs:
+ # Only the true_logits and is_treatment tensors are needed for testing.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=tf.ones_like(is_treatment),
+ control_predictions=tf.ones_like(is_treatment),
+ treatment_predictions=tf.ones_like(is_treatment),
+ uplift=tf.ones_like(is_treatment),
+ control_logits=tf.ones_like(is_treatment),
+ treatment_logits=tf.ones_like(is_treatment),
+ true_logits=(
+ true_logits if true_logits is not None else tf.ones_like(is_treatment)
+ ),
+ true_predictions=tf.ones_like(is_treatment),
+ is_treatment=is_treatment,
+ )
+
+
+class LogLossTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "two_tower_outputs_not_sliced",
+ "from_logits": True,
+ "compute_full_loss": False,
+ "slice_by_treatment": False,
+ "y_true": tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ true_logits=tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ is_treatment=tf.constant([[1], [0], [1], [0]]),
+ ),
+ "expected_loss": tf.reduce_mean(
+ tf.nn.log_poisson_loss(
+ tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ )
+ ),
+ },
+ {
+ "testcase_name": "two_tower_outputs_sliced",
+ "from_logits": True,
+ "compute_full_loss": False,
+ "slice_by_treatment": True,
+ "y_true": tf.constant([[1], [0]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ true_logits=tf.constant([[1], [0]], dtype=tf.float32),
+ is_treatment=tf.constant([[1], [0]]),
+ ),
+ "expected_loss": {
+ "poisson_log_loss/treatment": (
+ tf.math.exp(1.0) - 1 # exp(1) - 1 * 1
+ ),
+ "poisson_log_loss/control": 1.0, # exp(0) - 0 * 0
+ "poisson_log_loss": ((tf.math.exp(1.0) - 1) + 1) / 2,
+ },
+ },
+ {
+ "testcase_name": "tensor_outputs_from_logits",
+ "from_logits": True,
+ "compute_full_loss": False,
+ "slice_by_treatment": False,
+ "y_true": tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ "y_pred": tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ "expected_loss": tf.reduce_mean(
+ tf.nn.log_poisson_loss(
+ tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ )
+ ),
+ },
+ {
+ "testcase_name": "tensor_outputs_from_logits_full_loss",
+ "from_logits": True,
+ "compute_full_loss": True,
+ "slice_by_treatment": False,
+ "y_true": tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ "y_pred": tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ "expected_loss": tf.reduce_mean(
+ tf.nn.log_poisson_loss(
+ tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ compute_full_loss=True,
+ )
+ ),
+ },
+ {
+ "testcase_name": "tensor_outputs_from_predictions",
+ "from_logits": False,
+ "compute_full_loss": False,
+ "slice_by_treatment": False,
+ "y_true": tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ "y_pred": tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ "expected_loss": tf.reduce_mean(
+ tf.nn.log_poisson_loss(
+ tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ tf.math.log(
+ tf.constant([[1], [2], [3], [4]], dtype=tf.float32)
+ ),
+ )
+ ),
+ },
+ {
+ "testcase_name": "tensor_outputs_from_predictions_full_loss",
+ "from_logits": False,
+ "compute_full_loss": True,
+ "slice_by_treatment": False,
+ "y_true": tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ "y_pred": tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ "expected_loss": tf.reduce_mean(
+ tf.nn.log_poisson_loss(
+ tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ tf.math.log(
+ tf.constant([[1], [2], [3], [4]], dtype=tf.float32)
+ ),
+ compute_full_loss=True,
+ )
+ ),
+ },
+ )
+ def test_metric_computes_correct_loss(
+ self,
+ from_logits: bool,
+ compute_full_loss: bool,
+ slice_by_treatment: bool,
+ y_true: tf.Tensor,
+ y_pred: tf.Tensor,
+ expected_loss: tf.Tensor,
+ ):
+ metric = poisson_metrics.LogLoss(
+ from_logits=from_logits,
+ compute_full_loss=compute_full_loss,
+ slice_by_treatment=slice_by_treatment,
+ )
+ metric.update_state(y_true, y_pred)
+ self.assertAllClose(expected_loss, metric.result())
+
+ @parameterized.product(
+ from_logits=(True, False),
+ compute_full_loss=(True, False),
+ )
+ def test_metric_is_configurable(
+ self, from_logits: bool, compute_full_loss: bool
+ ):
+ metric = poisson_metrics.LogLoss(
+ from_logits=from_logits,
+ compute_full_loss=compute_full_loss,
+ slice_by_treatment=False,
+ )
+ self.assertLayerConfigurable(
+ layer=metric,
+ y_true=tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ y_pred=tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ serializable=True,
+ )
+
+
+class LogLossMeanBaselineTest(
+ keras_test_case.KerasTestCase, parameterized.TestCase
+):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "label_zero",
+ "expected_loss": 0.0,
+ "y_true": tf.constant([0], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "small_positive_label",
+ "expected_loss": 0.0,
+ "y_true": tf.constant([1e-10], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "label_one",
+ "expected_loss": 1.0,
+ "y_true": tf.constant([1], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "weighted_loss",
+ "expected_loss": 1.0,
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "sample_weight": tf.constant([[0], [1]], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "two_tower_outputs",
+ "expected_loss": 0.5 - 0.5 * tf.math.log(0.5),
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ is_treatment=tf.constant([[0], [1]], dtype=tf.float32),
+ ),
+ },
+ {
+ "testcase_name": "two_tower_outputs_sliced_loss",
+ "expected_loss": {
+ "loss": 0.5 - 0.5 * tf.math.log(0.5),
+ "loss/control": 0.0,
+ "loss/treatment": 1.0,
+ },
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ is_treatment=tf.constant([[0], [1]], dtype=tf.float32),
+ ),
+ "slice_by_treatment": True,
+ },
+ )
+ def test_metric_computes_correct_loss(
+ self,
+ expected_loss: tf.Tensor,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor | None = None,
+ sample_weight: tf.Tensor | None = None,
+ slice_by_treatment: bool = False,
+ ):
+ metric = poisson_metrics.LogLossMeanBaseline(
+ slice_by_treatment=slice_by_treatment, name="loss"
+ )
+ metric.update_state(y_true, y_pred, sample_weight=sample_weight)
+ self.assertAllClose(expected_loss, metric.result())
+
+ def test_negative_label_returns_nan_loss(self):
+ metric = poisson_metrics.LogLossMeanBaseline(slice_by_treatment=False)
+ metric.update_state(tf.constant([-1.0]))
+ self.assertTrue(tf.math.is_nan(metric.result()).numpy().item())
+
+ def test_metric_is_configurable(self):
+ metric = poisson_metrics.LogLossMeanBaseline(slice_by_treatment=False)
+ self.assertLayerConfigurable(
+ layer=metric,
+ y_true=tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ serializable=True,
+ )
+
+
+class LogLossMinimumTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "label_zero",
+ "expected_loss": 0.0,
+ "y_true": tf.constant([0], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "small_positive_label",
+ "expected_loss": 0.0,
+ "y_true": tf.constant([1e-10], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "label_one",
+ "expected_loss": 1.0,
+ "y_true": tf.constant([1], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "weighted_loss",
+ "expected_loss": 1.0,
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "sample_weight": tf.constant([[0], [1]], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "two_tower_outputs",
+ "expected_loss": 0.5,
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ is_treatment=tf.constant([[0], [1]], dtype=tf.float32),
+ ),
+ },
+ {
+ "testcase_name": "two_tower_outputs_sliced_loss",
+ "expected_loss": {
+ "loss": 0.5,
+ "loss/control": 0.0,
+ "loss/treatment": 1.0,
+ },
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ is_treatment=tf.constant([[0], [1]], dtype=tf.float32),
+ ),
+ "slice_by_treatment": True,
+ },
+ )
+ def test_metric_computes_correct_loss(
+ self,
+ expected_loss: tf.Tensor,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor | None = None,
+ sample_weight: tf.Tensor | None = None,
+ slice_by_treatment: bool = False,
+ ):
+ metric = poisson_metrics.LogLossMinimum(
+ slice_by_treatment=slice_by_treatment, name="loss"
+ )
+ metric.update_state(y_true, y_pred, sample_weight=sample_weight)
+ self.assertAllClose(expected_loss, metric.result())
+
+ def test_negative_label_returns_nan_loss(self):
+ metric = poisson_metrics.LogLossMinimum(slice_by_treatment=False)
+ metric.update_state(tf.constant([-1.0]))
+ self.assertTrue(tf.math.is_nan(metric.result()).numpy().item())
+
+ def test_metric_is_configurable(self):
+ metric = poisson_metrics.LogLossMinimum(slice_by_treatment=False)
+ self.assertLayerConfigurable(
+ layer=metric,
+ y_true=tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ serializable=True,
+ )
+
+
+class PseudoRSquaredTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "no_data",
+ "expected_loss": 0.0,
+ "y_true": tf.constant([], dtype=tf.float32),
+ "y_pred": tf.constant([], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "one_correct_prediction",
+ "expected_loss": 0.0,
+ "y_true": tf.constant([1], dtype=tf.float32),
+ "y_pred": tf.constant([1], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "one_wrong_prediction",
+ "expected_loss": 0.0, # LLmax and LLbaseline are equal.
+ "y_true": tf.constant([0], dtype=tf.float32),
+ "y_pred": tf.constant([1], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "all_correct_predictions",
+ "expected_loss": 1.0,
+ "y_true": tf.constant([[1], [2], [3]], dtype=tf.float32),
+ "y_pred": tf.constant([[1], [2], [3]], dtype=tf.float32),
+ "from_logits": False,
+ },
+ {
+ "testcase_name": "almost_correct_predictions",
+ "expected_loss": 1.0,
+ "y_true": tf.constant([[1], [2], [3]], dtype=tf.float32),
+ "y_pred": tf.constant([[1], [1.9999], [3.0001]], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "from_logits",
+ "expected_loss": (
+ (tf.math.exp(1.0) / 2) - (0.5 - 0.5 * tf.math.log(0.5))
+ ) / (0.5 - (0.5 - 0.5 * tf.math.log(0.5))),
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "y_pred": tf.constant([[0], [1]], dtype=tf.float32),
+ "from_logits": True,
+ },
+ {
+ "testcase_name": "two_tower_outputs",
+ "expected_loss": (
+ ((tf.math.exp(1.0) - 1) + 1) / 2 - (0.5 - 0.5 * tf.math.log(0.5))
+ ) / (0.5 - (0.5 - 0.5 * tf.math.log(0.5))),
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ true_logits=tf.constant([[0], [1]], dtype=tf.float32),
+ is_treatment=tf.constant([[0], [1]], dtype=tf.float32),
+ ),
+ "from_logits": True,
+ },
+ {
+ "testcase_name": "two_tower_outputs_sliced_loss",
+ "expected_loss": {
+ "r2": (
+ ((tf.math.exp(1.0) - 1) + 1) / 2 # LLfit
+ - (0.5 - 0.5 * tf.math.log(0.5)) # LLbaseline
+ ) / (0.5 - (0.5 - 0.5 * tf.math.log(0.5))),
+ "r2/control": 0.0,
+ "r2/treatment": 0.0,
+ },
+ "y_true": tf.constant([[0], [1]], dtype=tf.float32),
+ "y_pred": _get_two_tower_outputs(
+ true_logits=tf.constant([[0], [1]], dtype=tf.float32),
+ is_treatment=tf.constant([[0], [1]], dtype=tf.float32),
+ ),
+ "from_logits": True,
+ "slice_by_treatment": True,
+ },
+ )
+ def test_metric_computation_is_correct(
+ self,
+ expected_loss: tf.Tensor,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs | tf.Tensor,
+ sample_weight: tf.Tensor | None = None,
+ from_logits: bool = False,
+ slice_by_treatment: bool = False,
+ ):
+ metric = poisson_metrics.PseudoRSquared(
+ from_logits=from_logits,
+ slice_by_treatment=slice_by_treatment,
+ name="r2",
+ )
+ metric.update_state(y_true, y_pred, sample_weight=sample_weight)
+ self.assertAllClose(expected_loss, metric.result())
+
+ def test_slicing_raises_error_when_input_is_tensor(self):
+ metric = poisson_metrics.PseudoRSquared()
+ y_true = tf.constant([[0], [0], [2], [7]], dtype=tf.float32)
+ y_pred = tf.constant([[1], [2], [3], [4]], dtype=tf.float32)
+ with self.assertRaisesRegex(
+ ValueError,
+ "`slice_by_treatment` must be set to `False` when `y_pred` is not of"
+ " type `TwoTowerTrainingOutputs`.",
+ ):
+ metric(y_true, y_pred)
+
+ @parameterized.parameters(True, False)
+ def test_metric_is_configurable(self, from_logits: bool):
+ metric = poisson_metrics.PseudoRSquared(
+ from_logits=from_logits, slice_by_treatment=False
+ )
+ self.assertLayerConfigurable(
+ layer=metric,
+ y_true=tf.constant([[0], [0], [2], [7]], dtype=tf.float32),
+ y_pred=tf.constant([[1], [2], [3], [4]], dtype=tf.float32),
+ serializable=True,
+ )
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/sliced_metric.py b/official/recommendation/uplift/metrics/sliced_metric.py
new file mode 100644
index 00000000000..dcbbfb59ddd
--- /dev/null
+++ b/official/recommendation/uplift/metrics/sliced_metric.py
@@ -0,0 +1,215 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for reporting metrics sliced by a feature."""
+
+import copy
+
+import tensorflow as tf, tf_keras
+
+
+class SlicedMetric(tf_keras.metrics.Metric):
+ """A metric sliced by integer, boolean, or string features.
+
+ A metric wrapper that computes a metric for different slices of an arbitrary
+ feature. The slicing is specified via a slicing spec, which is a dictionary
+ from a slice name to the unique value to be sliced on. For each pair of
+ `slice_name`, `slicing_value` passed, the suffix `/slice_name` will be
+ appended to the name of the result of the corresponding slice.
+ An overall result is also computed without any slicing applied.
+
+ In order for this to work correctly, the given metric must support passing
+ `sample_weight` to its `update_state()` method. Additionally, the slicing
+ feature must also be passed to `update_state()` method of this class and
+ must be of a broadcastable shape to the metric inputs.
+ This wrapper creates a deep copy of the metric passed to it for each slice.
+ At every call to `update_state()`, the wrapper will call the `update_state()`
+ method of the metric for each slice with the `sample_weights` set to zero
+ where the slicing feature is not equal to the corresponding slicing value.
+
+ If the given metric returns a tensor, the result of this metric will be a
+ dictionary mapping from the sliced metric's name to the result for that slice.
+ If the given metric returns a dictionary of tensors, the result of this metric
+ will be a flattened dictionary consisting of each of the sliced metrics'
+ results for every slice.
+
+ Example usage:
+
+ >>> sliced_metric = SlicedMetric(
+ ... tf_keras.metrics.Accuracy('accuracy'),
+ ... slicing_spec={"control": False, "treatment": True},
+ ... )
+ >>> sliced_metric.update_state(
+ ... y_true=tf.constant([[0], [1], [0], [1]]),
+ ... y_pred=tf.constant([[1], [0], [1], [1]]),
+ ... slicing_feature=tf.constant([[True], [False], [True], [False]]),
+ ... )
+ >>> sliced_metric.result()
+ {
+ "accuracy": 0.25,
+ "accuracy/control": 0.5,
+ "accuracy/treatment": 0
+ }
+ """
+
+ def __init__(
+ self,
+ metric: tf_keras.metrics.Metric,
+ slicing_spec: dict[str, str] | dict[str, int],
+ slicing_feature_dtype: tf.DType | None = None,
+ name: str | None = None,
+ ):
+ """Initializes the instance.
+
+ Args:
+ metric: A `tf_keras.metrics.Metric` instance.
+ slicing_spec: A dictionary that maps from string slice names, to one of
+ integer, boolean, or string slicing values.
+ slicing_feature_dtype: The expected dtype of the slicing feature. The
+ values in the slicing spec are casted to this type if passed. If None,
+ the dtype of the slicing feature is inferred based on the values in the
+ slicing spec.
+ name: The name of the wrapper metric. Defaults to `sliced_{metric.name}`.
+
+ Raises:
+ A ValueError if `slicing_spec` is empty, contains duplicate slicing
+ values, or has slicing values of different types.
+ """
+ super().__init__(name=name or f"sliced_{metric.name}", dtype=metric.dtype)
+
+ if not slicing_spec:
+ raise ValueError("The slicing spec must be a non-empty dictionary.")
+
+ slice_names, slicing_values = zip(*slicing_spec.items())
+ if not isinstance(slicing_values[0], (int, bool, str)) or not all(
+ isinstance(k, type(slicing_values[0])) for k in slicing_values
+ ):
+ raise ValueError(
+ "All slicing values in the slicing spec must be one of `int`, "
+ "`bool`, or `str`, and all values must have the same type. "
+ f"Got types: {list(map(type, slicing_values))}."
+ )
+
+ if len(slicing_values) > len(set(slicing_values)):
+ raise ValueError(
+ "The slicing values passed to the slicing spec must be unique. Got "
+ f"{slicing_values}."
+ )
+
+ # TODO(b/276811843): Look into validating whether `metric` accepts
+ # `sample_weights` in its `update_state` method.
+
+ # Instance fully owns a deep copy of the metric.
+ self._metric = copy.deepcopy(metric)
+ self._slice_names = list(slice_names)
+ self._slicing_values = list(slicing_values)
+ self._slicing_values_tensors = [
+ tf.constant(v, slicing_feature_dtype) for v in slicing_values
+ ]
+ self._slicing_feature_dtype = self._slicing_values_tensors[0].dtype
+ self._sliced_metrics = [copy.deepcopy(metric) for _ in self._slicing_values]
+
+ def update_state(
+ self,
+ *args: tf.Tensor,
+ sample_weight: tf.Tensor | None = None,
+ slicing_feature: tf.Tensor,
+ **kwargs,
+ ):
+ """Updates the state of the metrics for each slice.
+
+ Args:
+ *args: A variable amount of `tf.Tensor` instances that will be passed to
+ the `update_state` method of each metric.
+ sample_weight: An optional `tf.Tensor` used to weight the sample. Its
+ dimensions must be broadcastable to the shape(s) of *args.
+ slicing_feature: A `tf.Tensor` consisting of the feature to be sliced on.
+ Its dimensions must be broadcastable to the shape(s) of *args.
+ **kwargs: Keyword arguments that will be passed to the `update_state`
+ method of each metric.
+ """
+
+ if slicing_feature.dtype != self._slicing_feature_dtype:
+ raise ValueError(
+ "The `slicing_feature` and slicing values in `slicing_spec` must "
+ "have the same type. Got types: "
+ f"{(slicing_feature.dtype, self._slicing_feature_dtype)}."
+ )
+
+ if sample_weight is not None:
+ for _ in range(len(slicing_feature.shape) - len(sample_weight.shape)):
+ sample_weight = tf.expand_dims(sample_weight, axis=-1)
+
+ for _ in range(len(sample_weight.shape) - len(slicing_feature.shape)):
+ slicing_feature = tf.expand_dims(slicing_feature, axis=-1)
+
+ self._metric.update_state(*args, sample_weight=sample_weight, **kwargs)
+ for slicing_val, metric in zip(
+ self._slicing_values_tensors, self._sliced_metrics
+ ):
+ slice_mask = tf.cast(slicing_feature == slicing_val, dtype=tf.float32)
+ if sample_weight is not None:
+ weight = slice_mask * tf.cast(sample_weight, dtype=tf.float32)
+ else:
+ weight = slice_mask
+ metric.update_state(*args, sample_weight=weight, **kwargs)
+
+ def result(self) -> dict[str, tf.Tensor]:
+ """Aggregates all the metrics' results into a flattened dictionary."""
+ metric_name = self._metric.name
+ metric_result = self._metric.result()
+ slice_results = [metric.result() for metric in self._sliced_metrics]
+
+ if isinstance(metric_result, tf.Tensor):
+ results = {metric_name: metric_result}
+ slice_names = (f"{metric_name}/{name}" for name in self._slice_names)
+ results.update(zip(slice_names, slice_results))
+ return results
+
+ if isinstance(metric_result, dict) and all(
+ isinstance(result, tf.Tensor) for result in metric_result.values()
+ ):
+ results = {**metric_result}
+ for slice_name, slice_result in zip(self._slice_names, slice_results):
+ result_names, result_values = zip(*slice_result.items())
+ slice_names = [f"{name}/{slice_name}" for name in result_names]
+ results.update(zip(slice_names, result_values))
+ return results
+
+ raise ValueError(
+ "The output of the given metric must either be a `tf.Tensor` or "
+ "a `dict[str, tf.Tensor]`, but got unsupported output: "
+ f"{metric_result}."
+ )
+
+ def reset_state(self):
+ self._metric.reset_state()
+ for metric in self._sliced_metrics:
+ metric.reset_state()
+
+ def get_config(self):
+ return {
+ "name": self.name,
+ "metric": tf_keras.metrics.serialize(self._metric),
+ "slicing_spec": dict(zip(self._slice_names, self._slicing_values)),
+ "slicing_feature_dtype": self._slicing_feature_dtype.name,
+ }
+
+ @classmethod
+ def from_config(cls, config):
+ config["metric"] = tf_keras.metrics.deserialize(config["metric"])
+ config["slicing_feature_dtype"] = tf.as_dtype(
+ config["slicing_feature_dtype"]
+ )
+ return cls(**config)
diff --git a/official/recommendation/uplift/metrics/sliced_metric_test.py b/official/recommendation/uplift/metrics/sliced_metric_test.py
new file mode 100644
index 00000000000..f1b0e179f37
--- /dev/null
+++ b/official/recommendation/uplift/metrics/sliced_metric_test.py
@@ -0,0 +1,361 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for sliced metrics."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift.metrics import sliced_metric
+
+
+class MeanSquared(tf_keras.metrics.Mean):
+
+ def result(self):
+ mean = super().result()
+ return {
+ self.name + "/mean": mean,
+ self.name + "/squared": tf.math.square(mean),
+ }
+
+
+class SlicedMetricTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "bool_slicing",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "predictions": tf.constant([1, 0, 1, 1], dtype=tf.int32),
+ "slicing_spec": {"treatment": True, "control": False},
+ "slicing_feature": tf.constant([True, False, True, False]),
+ "expected_result": {
+ "accuracy": 0.25,
+ "accuracy/treatment": 0,
+ "accuracy/control": 0.5,
+ },
+ },
+ {
+ "testcase_name": "int_slicing",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "predictions": tf.constant([1, 0, 1, 1], dtype=tf.int32),
+ "slicing_spec": {"app_usage": 3, "install": 4, "purchase": 5},
+ "slicing_feature": tf.constant([4, 5, 4, 3], dtype=tf.int32),
+ "expected_result": {
+ "accuracy": 0.25,
+ "accuracy/install": 0,
+ "accuracy/purchase": 0,
+ "accuracy/app_usage": 1,
+ },
+ },
+ {
+ "testcase_name": "str_slicing",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "predictions": tf.constant([1, 0, 1, 1], dtype=tf.int32),
+ "slicing_spec": {"install": "install", "purchase": "purchase"},
+ "slicing_feature": tf.constant(
+ ["install", "purchase", "install", "app_usage"]
+ ),
+ "expected_result": {
+ "accuracy": 0.25,
+ "accuracy/install": 0,
+ "accuracy/purchase": 0,
+ },
+ },
+ {
+ "testcase_name": "int_slicing_weighted",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "predictions": tf.constant([1, 0, 1, 1], dtype=tf.int32),
+ "slicing_spec": {"app_usage": 3, "install": 4, "purchase": 5},
+ "slicing_feature": tf.constant([4, 5, 4, 3], dtype=tf.int32),
+ "weights": tf.constant([0, 0, 1, 1], dtype=tf.float32),
+ "expected_result": {
+ "accuracy": 0.5,
+ "accuracy/install": 0,
+ "accuracy/purchase": 0,
+ "accuracy/app_usage": 1,
+ },
+ },
+ )
+ def test_binary_sliced_metric(
+ self,
+ labels: tf.Tensor,
+ predictions: tf.Tensor,
+ slicing_spec: dict[str, int | str | bool],
+ slicing_feature: tf.Tensor,
+ expected_result: dict[str, float],
+ weights: tf.Tensor | None = None,
+ ):
+ metric = sliced_metric.SlicedMetric(
+ tf_keras.metrics.Accuracy("accuracy"),
+ slicing_spec=slicing_spec,
+ )
+ metric.update_state(
+ labels,
+ predictions,
+ sample_weight=weights,
+ slicing_feature=slicing_feature,
+ )
+ self.assertDictEqual(expected_result, metric.result())
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "bool_slicing",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "slicing_spec": {"treatment": True, "control": False},
+ "slicing_feature": tf.constant([True, False, True, False]),
+ "expected_result": {
+ "mean": 0.5,
+ "mean/treatment": 0,
+ "mean/control": 1,
+ },
+ },
+ {
+ "testcase_name": "int_slicing",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "slicing_spec": {"app_usage": 3, "install": 4, "purchase": 5},
+ "slicing_feature": tf.constant([4, 5, 4, 3], dtype=tf.int32),
+ "expected_result": {
+ "mean": 0.5,
+ "mean/install": 0,
+ "mean/purchase": 1,
+ "mean/app_usage": 1,
+ },
+ },
+ )
+ def test_unary_sliced_metrics(
+ self,
+ labels: tf.Tensor,
+ slicing_spec: dict[str, int | str | bool],
+ slicing_feature: tf.Tensor,
+ expected_result: dict[str, float],
+ ):
+ metric = sliced_metric.SlicedMetric(
+ tf_keras.metrics.Mean("mean"),
+ slicing_spec=slicing_spec,
+ )
+ metric.update_state(labels, slicing_feature=slicing_feature)
+ self.assertDictEqual(expected_result, metric.result())
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "int_slicing",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "slicing_spec": {"app_usage": 3, "install": 4, "purchase": 5},
+ "slicing_feature": tf.constant([4, 5, 4, 3], dtype=tf.int32),
+ "expected_result": {
+ "msq/mean": 0.5,
+ "msq/mean/install": 0,
+ "msq/mean/purchase": 1,
+ "msq/mean/app_usage": 1,
+ "msq/squared": 0.25,
+ "msq/squared/install": 0,
+ "msq/squared/purchase": 1,
+ "msq/squared/app_usage": 1,
+ },
+ },
+ {
+ "testcase_name": "str_slicing",
+ "labels": tf.constant([0, 1, 0, 1], dtype=tf.int32),
+ "slicing_spec": {"install": "install", "purchase": "purchase"},
+ "slicing_feature": tf.constant(
+ ["install", "purchase", "install", "app_usage"]
+ ),
+ "expected_result": {
+ "msq/mean": 0.5,
+ "msq/mean/install": 0,
+ "msq/mean/purchase": 1,
+ "msq/squared": 0.25,
+ "msq/squared/install": 0,
+ "msq/squared/purchase": 1,
+ },
+ },
+ )
+ def test_metric_with_dict_result(
+ self,
+ labels: tf.Tensor,
+ slicing_spec: dict[str, int | str | bool],
+ slicing_feature: tf.Tensor,
+ expected_result: dict[str, float],
+ ):
+ metric = sliced_metric.SlicedMetric(
+ MeanSquared("msq"), slicing_spec=slicing_spec
+ )
+ metric.update_state(labels, slicing_feature=slicing_feature)
+ self.assertDictEqual(expected_result, metric.result())
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "empty_slicing_spec",
+ "slicing_spec": {},
+ },
+ {
+ "testcase_name": "incompatible_types",
+ "slicing_spec": {"a": True, "b": 1, "c": 0},
+ },
+ {
+ "testcase_name": "duplicate_slicing_values",
+ "slicing_spec": {"a": True, "b": False, "c": True},
+ },
+ )
+ def test_invalid_inputs(self, slicing_spec):
+ self.assertRaises(
+ ValueError,
+ sliced_metric.SlicedMetric,
+ tf_keras.metrics.Mean(),
+ slicing_spec=slicing_spec,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "float_to_int",
+ "slicing_spec": {"a": 1, "b": 2},
+ "slicing_feature": tf.constant([1, 2, 1, 0], dtype=tf.float32),
+ },
+ {
+ "testcase_name": "int_to_bool",
+ "slicing_spec": {"a": False, "b": True},
+ "slicing_feature": tf.constant([1, 0, 0, 1], dtype=tf.int32),
+ },
+ )
+ def test_invalid_update(self, slicing_spec, slicing_feature):
+ # Invalid cast
+ metric = sliced_metric.SlicedMetric(
+ tf_keras.metrics.Mean(), slicing_spec=slicing_spec
+ )
+ self.assertRaises(
+ ValueError,
+ metric.update_state,
+ values=tf.constant([1, 0, 1, 0], dtype=tf.int32),
+ slicing_feature=slicing_feature,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "inputs_2x2_slicing_feature_2",
+ "slicing_feature": tf.constant([False, True]),
+ "expected_result": {
+ "accuracy": 0.5,
+ "accuracy/control": 0.5,
+ "accuracy/treatment": 0.5,
+ },
+ },
+ {
+ "testcase_name": "inputs_2x2_slicing_feature_1x2",
+ "slicing_feature": tf.constant([[False, True]]),
+ "expected_result": {
+ "accuracy": 0.5,
+ "accuracy/control": 0,
+ "accuracy/treatment": 1.0,
+ },
+ },
+ {
+ "testcase_name": "inputs_2x2_slicing_feature_1x2_weight_1",
+ "slicing_feature": tf.constant([[False, True]]),
+ "sample_weight": tf.constant([1.0]),
+ "expected_result": {
+ "accuracy": 0.5,
+ "accuracy/control": 0,
+ "accuracy/treatment": 1.0,
+ },
+ },
+ {
+ "testcase_name": "inputs_2x2_slicing_feature_2_weight_2",
+ "slicing_feature": tf.constant([False, True]),
+ "sample_weight": tf.constant([0.5, 0.5]),
+ "expected_result": {
+ "accuracy": 0.5,
+ "accuracy/control": 0.5,
+ "accuracy/treatment": 0.5,
+ },
+ },
+ )
+ def test_broadcastable_weights_and_slicing_feature(
+ self,
+ slicing_feature: tf.Tensor,
+ expected_result: dict[str, float],
+ sample_weight: tf.Tensor | None = None,
+ ):
+ metric = sliced_metric.SlicedMetric(
+ tf_keras.metrics.Accuracy("accuracy"),
+ slicing_spec={"control": False, "treatment": True},
+ )
+ metric.update_state(
+ tf.constant([[0, 1], [0, 1]], dtype=tf.int32),
+ tf.constant([[1, 1], [1, 1]], dtype=tf.int32),
+ sample_weight=sample_weight,
+ slicing_feature=slicing_feature,
+ )
+ self.assertDictEqual(expected_result, metric.result())
+
+ def test_batched_inputs(self):
+ metric = sliced_metric.SlicedMetric(
+ tf_keras.metrics.Accuracy("accuracy"),
+ slicing_spec={"install": 4, "purchase": 5},
+ )
+ metric.update_state(
+ tf.constant([[0, 1], [1, 0]], dtype=tf.int32),
+ tf.constant([[1, 1], [1, 1]], dtype=tf.int32),
+ slicing_feature=tf.constant([[4, 5], [4, 3]], dtype=tf.int32),
+ )
+ expected_result = {
+ "accuracy": 0.5,
+ "accuracy/install": 0.5,
+ "accuracy/purchase": 1.0,
+ }
+ self.assertDictEqual(expected_result, metric.result())
+
+ def test_reset_state(self):
+ metric = sliced_metric.SlicedMetric(
+ metric=tf_keras.metrics.AUC(curve="PR", from_logits=False, name="auc"),
+ slicing_spec={"control": False, "treatment": True},
+ )
+
+ expected_initial_result = {
+ "auc": 0.0,
+ "auc/control": 0.0,
+ "auc/treatment": 0.0,
+ }
+ self.assertAllClose(expected_initial_result, metric.result())
+
+ metric.update_state(
+ tf.constant([[0], [0], [1], [1]]), # y_true
+ tf.constant([[0.2], [0.6], [0.3], [0.7]]), # y_pred
+ slicing_feature=tf.constant([[True], [False], [True], [False]]),
+ )
+
+ result = metric.result()
+ self.assertGreater(result["auc"], 0.0)
+ self.assertGreater(result["auc/control"], 0.0)
+ self.assertGreater(result["auc/treatment"], 0.0)
+
+ metric.reset_state()
+ self.assertAllClose(expected_initial_result, metric.result())
+
+ def test_metric_config(self):
+ metric = sliced_metric.SlicedMetric(
+ tf_keras.metrics.SparseTopKCategoricalAccuracy(k=2, name="accuracy@2"),
+ slicing_spec={"a": False, "b": True},
+ slicing_feature_dtype=tf.bool,
+ name="sliced_accuracy",
+ )
+ y_true = tf.constant([1, 0, 1, 0])
+ y_pred = tf.constant([[0.1, 0.9], [0.8, 0.2], [0.7, 0.3], [0.6, 0.4]])
+ slicing_feature = tf.constant([True, False, False, True])
+ self.assertLayerConfigurable(
+ metric, y_true=y_true, y_pred=y_pred, slicing_feature=slicing_feature
+ )
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/treatment_fraction.py b/official/recommendation/uplift/metrics/treatment_fraction.py
new file mode 100644
index 00000000000..69e4845cfd5
--- /dev/null
+++ b/official/recommendation/uplift/metrics/treatment_fraction.py
@@ -0,0 +1,86 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for computing fraction of treated examples."""
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class TreatmentFraction(tf_keras.metrics.Metric):
+ """Computes the fraction of treated examples.
+
+ Note that the prediction tensor is expected to be of type
+ `TwoTowerTrainingOutputs`.
+
+ Example standalone usage:
+
+ >>> treatment_fraction = TreatmentFraction()
+ >>> y_pred = types.TwoTowerTrainingOutputs(
+ ... is_treatment=tf.constant([True, False, True, True]),
+ ... )
+ >>> treatment_fraction(y_true=tf.zeros(4), y_pred=y_pred)
+ 0.75
+
+ Example usage with the `model.compile()` API:
+
+ >>> model.compile(
+ ... optimizer="sgd",
+ ... loss=TrueLogitsLoss(tf_keras.losses.mean_squared_error),
+ ... metrics=[TreatmentFraction()]
+ ... )
+ """
+
+ def __init__(self, **kwargs):
+ """Initializes the instance.
+
+ Args:
+ **kwargs: base metric keyword arguments.
+ """
+ super().__init__(**kwargs)
+ self._treatment_fraction = tf_keras.metrics.Mean(**kwargs)
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs,
+ sample_weight: tf.Tensor | None = None,
+ ) -> None:
+ """Updates the treatment fraction.
+
+ Args:
+ y_true: tensor labels.
+ y_pred: two tower training outputs. The treatment indicator tensor is used
+ update the treatment fraction.
+ sample_weight: optional sample weight tensor for computing the weighted
+ treatment fraction. The unweighted treatment fraction is computed
+ instead if it is left as `None`.
+
+ Raises:
+ TypeError: if y_pred is not of type `TwoTowerTrainingOutputs`.
+ """
+ if not isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise TypeError(
+ "y_pred must be of type `TwoTowerTrainingOutputs` but got type"
+ f" {type(y_pred)} instead."
+ )
+
+ self._treatment_fraction.update_state(
+ values=y_pred.is_treatment, sample_weight=sample_weight
+ )
+
+ def result(self) -> tf.Tensor:
+ return self._treatment_fraction.result()
diff --git a/official/recommendation/uplift/metrics/treatment_fraction_test.py b/official/recommendation/uplift/metrics/treatment_fraction_test.py
new file mode 100644
index 00000000000..5bedf99cca3
--- /dev/null
+++ b/official/recommendation/uplift/metrics/treatment_fraction_test.py
@@ -0,0 +1,156 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for treatment_fraction."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import treatment_fraction
+
+
+class TreatmentFractionTest(
+ keras_test_case.KerasTestCase, parameterized.TestCase
+):
+
+ def _get_y_pred(
+ self, is_treatment: tf.Tensor
+ ) -> types.TwoTowerTrainingOutputs:
+ # Only the is_treatment tensor is required for testing.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=tf.ones_like(is_treatment),
+ control_predictions=tf.ones_like(is_treatment),
+ treatment_predictions=tf.ones_like(is_treatment),
+ uplift=tf.ones_like(is_treatment),
+ control_logits=tf.ones_like(is_treatment),
+ treatment_logits=tf.ones_like(is_treatment),
+ true_logits=tf.ones_like(is_treatment),
+ true_predictions=tf.ones_like(is_treatment),
+ is_treatment=is_treatment,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": None,
+ "expected_result": 0.5,
+ },
+ {
+ "testcase_name": "weighted",
+ "is_treatment": tf.constant(
+ [[True], [False], [True], [True], [False]]
+ ),
+ "sample_weight": tf.constant([0.5, 0.5, 0, 0.7, 1.8]),
+ "expected_result": np.average(
+ [1, 0, 1, 1, 0], weights=[0.5, 0.5, 0, 0.7, 1.8]
+ ),
+ },
+ {
+ "testcase_name": "only_control",
+ "is_treatment": tf.constant([[False], [False], [False]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_result": 0.0,
+ },
+ {
+ "testcase_name": "only_treatment",
+ "is_treatment": tf.constant([[True], [True], [True]]),
+ "sample_weight": tf.constant([0, 1, 1]),
+ "expected_result": 1.0,
+ },
+ {
+ "testcase_name": "one_entry",
+ "is_treatment": tf.constant([True]),
+ "sample_weight": None,
+ "expected_result": 1.0,
+ },
+ {
+ "testcase_name": "no_entry",
+ "is_treatment": tf.constant([], dtype=tf.bool),
+ "sample_weight": tf.constant([]),
+ "expected_result": 0.0,
+ },
+ )
+ def test_treatment_fraction_computes_weighted_mean_of_is_treatment_tensor(
+ self, is_treatment, sample_weight, expected_result
+ ):
+ metric = treatment_fraction.TreatmentFraction()
+ y_true = tf.zeros_like(is_treatment)
+ y_pred = self._get_y_pred(is_treatment)
+ metric.update_state(
+ y_true=y_true, y_pred=y_pred, sample_weight=sample_weight
+ )
+ self.assertEqual(expected_result, metric.result())
+
+ def test_multiple_update_batches_returns_aggregated_treatment_fractions(self):
+ metric = treatment_fraction.TreatmentFraction()
+
+ metric.update_state(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(tf.constant([[True], [True], [True]])),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(tf.constant([[False], [False], [False]])),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(tf.constant([[True], [False], [True]])),
+ sample_weight=tf.constant([0.3, 0.25, 0.7]),
+ )
+
+ expected_treatment_fraction = np.average(
+ [1, 1, 1, 0, 0, 0, 1, 0, 1], weights=[1, 1, 1, 1, 1, 1, 0.3, 0.25, 0.7]
+ )
+ self.assertEqual(expected_treatment_fraction, metric.result())
+
+ def test_initial_and_reset_state_return_zero_treatment_fraction(self):
+ metric = treatment_fraction.TreatmentFraction()
+ self.assertEqual(0.0, metric.result())
+
+ metric(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(tf.constant([[True], [False], [True]])),
+ )
+ self.assertEqual(2 / 3, metric.result())
+
+ metric.reset_states()
+ self.assertEqual(0.0, metric.result())
+
+ def test_metric_config_is_serializable(self):
+ metric = treatment_fraction.TreatmentFraction(
+ name="test_name", dtype=tf.float16
+ )
+ y_pred = self._get_y_pred(
+ is_treatment=tf.constant([[True], [False], [True], [False]]),
+ )
+ self.assertLayerConfigurable(
+ layer=metric, y_true=tf.zeros(4), y_pred=y_pred, serializable=True
+ )
+
+ def test_invalid_prediction_tensor_type_raises_type_error(self):
+ metric = treatment_fraction.TreatmentFraction()
+
+ with self.assertRaisesRegex(
+ TypeError, "y_pred must be of type `TwoTowerTrainingOutputs`"
+ ):
+ metric.update_state(y_true=tf.ones((3, 1)), y_pred=tf.ones((3, 1)))
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/treatment_sliced_metric.py b/official/recommendation/uplift/metrics/treatment_sliced_metric.py
new file mode 100644
index 00000000000..954edbb809d
--- /dev/null
+++ b/official/recommendation/uplift/metrics/treatment_sliced_metric.py
@@ -0,0 +1,103 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for computing sliced metric values based on treatment group."""
+
+from typing import Optional
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift.metrics import sliced_metric
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class TreatmentSlicedMetric(sliced_metric.SlicedMetric):
+ """Computes (weighted) sliced metric values based on treatment group.
+
+ Two copies are made of the given metric to compute the same metric for the
+ control and treatment groups. The metric name is used for the output results,
+ with a "/control" and "/treatment" suffix for the control and treatment
+ groups. If the given metric outputs a dictionary of tensors the result will be
+ a flattened dictionary with all the sliced metric results.
+
+ Example usage:
+
+ >>> sliced_metric = TreatmentSlicedMetric(
+ ... metric=tf_keras.metrics.Mean(name="example/mean")
+ ... )
+ >>> sliced_metric.update_state(
+ ... values=tf.constant([[0], [1], [5]]),
+ ... is_treatment=tf.constant([[1], [0], [1]]),
+ ... )
+ >>> sliced_metric.result()
+ {
+ "example/mean": 2.0,
+ "example/mean/control": 1.0,
+ "example/mean/treatment": 2.5
+ }
+ """
+
+ def __init__(self, metric: tf_keras.metrics.Metric):
+ """Initializes a `TreatmentSlicedMetric` instance.
+
+ Args:
+ metric: A `keras.metrics.Metric` instance with implemented get_config and
+ from_config methods. Its update_state method is expected have a
+ `update_state(self, values, sample_weight=None)` signature. The control
+ and treatment metrics will have the same name as the given metric, with
+ "/control" and "/treatment" suffixes.
+ """
+ super().__init__(
+ metric=metric,
+ slicing_spec={"control": False, "treatment": True},
+ name=f"treatment_sliced_{metric.name}",
+ )
+
+ def update_state(
+ self,
+ values: tf.Tensor,
+ is_treatment: tf.Tensor,
+ sample_weight: Optional[tf.Tensor] = None,
+ **kwargs,
+ ):
+ """Computes aggregate, control and treatment metric updates.
+
+ Args:
+ values: A `tf.Tensor` instance of shape (D0, ..., DN) passed to the
+ `update_state` method of the treatment, control, and overall metrics.
+ is_treatment: a `tf.Tensor` of shape (D0,) or (D0, 1) castable to boolean
+ indicating if the example belongs to the treatment group (True) or
+ control group (False).
+ sample_weight: optional `tf.Tensor` for the sample weight. If given it
+ will also be sliced by the is_treatment tensor.
+ **kwargs: Keyword arguments that will be passed to the `update_state`
+ method of each metric.
+ """
+ slicing_feature = tf.reshape(tf.cast(is_treatment, tf.bool), [-1])
+ if sample_weight is not None:
+ sample_weight = tf.reshape(sample_weight, [-1])
+ super().update_state(
+ values,
+ sample_weight=sample_weight,
+ slicing_feature=slicing_feature,
+ **kwargs,
+ )
+
+ def get_config(self):
+ return {"metric": tf_keras.metrics.serialize(self._metric)}
+
+ @classmethod
+ def from_config(cls, config):
+ config["metric"] = tf_keras.metrics.deserialize(config["metric"])
+ return cls(**config)
diff --git a/official/recommendation/uplift/metrics/treatment_sliced_metric_test.py b/official/recommendation/uplift/metrics/treatment_sliced_metric_test.py
new file mode 100644
index 00000000000..142c16b73b8
--- /dev/null
+++ b/official/recommendation/uplift/metrics/treatment_sliced_metric_test.py
@@ -0,0 +1,194 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for treatment_sliced_metric."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift.metrics import treatment_sliced_metric
+
+
+class MeanSquared(tf_keras.metrics.Mean):
+
+ def result(self):
+ mean = super().result()
+ return {
+ self.name + "/mean": mean,
+ self.name + "/squared": tf.math.square(mean),
+ }
+
+
+# TODO(b/271487910): Add test case to ensure the right inputs are passed to the
+# sliced metrics.
+class TreatmentSlicedMetricTest(
+ keras_test_case.KerasTestCase, parameterized.TestCase
+):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "values": tf.constant([0, 1, 5, 6]),
+ "is_treatment": tf.constant([1, 0, 1, 0]),
+ "sample_weight": None,
+ "expected_result": {
+ "test/mean": 3.0,
+ "test/mean/control": 3.5,
+ "test/mean/treatment": 2.5,
+ },
+ },
+ {
+ "testcase_name": "weighted",
+ "values": tf.constant([0, 1, 5, 6, -7]),
+ "is_treatment": tf.constant([1, 0, 1, 1, 0]),
+ "sample_weight": tf.constant([0.5, 0.5, 0, 0.7, 1.8]),
+ "expected_result": {
+ "test/mean": np.average(
+ np.array([0, 1, 5, 6, -7]),
+ weights=np.array([0.5, 0.5, 0, 0.7, 1.8]),
+ ),
+ "test/mean/control": np.average(
+ np.array([1, -7]), weights=np.array([0.5, 1.8])
+ ),
+ "test/mean/treatment": np.average(
+ np.array([0, 5, 6]), weights=np.array([0.5, 0, 0.7])
+ ),
+ },
+ },
+ {
+ "testcase_name": "only_control",
+ "values": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[0], [0], [0]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_result": {
+ "test/mean": 2.5,
+ "test/mean/control": 2.5,
+ "test/mean/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "only_treatment",
+ "values": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[1], [1], [1]]),
+ "sample_weight": tf.constant([0, 1, 1]),
+ "expected_result": {
+ "test/mean": 3.0,
+ "test/mean/control": 0.0,
+ "test/mean/treatment": 3.0,
+ },
+ },
+ )
+ def test_treatment_sliced_metric(
+ self, values, is_treatment, sample_weight, expected_result
+ ):
+ sliced_metric = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.Mean(name="test/mean")
+ )
+ sliced_metric(values, is_treatment, sample_weight=sample_weight)
+ self.assertDictEqual(expected_result, sliced_metric.result())
+
+ def test_multiple_batches(self):
+ sliced_metric = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.Mean(name="test/mean")
+ )
+
+ sliced_metric(
+ values=tf.constant([[1], [2], [4]]),
+ is_treatment=tf.ones((3, 1)),
+ sample_weight=None,
+ )
+ sliced_metric(
+ values=tf.constant([[-3], [0], [5]]),
+ is_treatment=tf.zeros((3, 1)),
+ sample_weight=None,
+ )
+ sliced_metric(
+ values=tf.constant([[0], [1], [-5]]),
+ is_treatment=tf.constant([1, 0, 1]),
+ sample_weight=tf.constant([0.3, 0.25, 0.7]),
+ )
+
+ expected_results = {
+ "test/mean": np.average(
+ np.array([1, 2, 4, -3, 0, 5, 0, 1, -5]),
+ weights=np.array([1, 1, 1, 1, 1, 1, 0.3, 0.25, 0.7]),
+ ),
+ "test/mean/control": np.average(
+ np.array([-3, 0, 5, 1]), weights=np.array([1, 1, 1, 0.25])
+ ),
+ "test/mean/treatment": np.average(
+ np.array([1, 2, 4, 0, -5]), weights=np.array([1, 1, 1, 0.3, 0.7])
+ ),
+ }
+ self.assertDictEqual(expected_results, sliced_metric.result())
+
+ def test_metric_states(self):
+ sliced_metric = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.Mean(name="test/mean")
+ )
+
+ expected_initial_result = {
+ "test/mean": 0.0,
+ "test/mean/control": 0.0,
+ "test/mean/treatment": 0.0,
+ }
+ self.assertDictEqual(expected_initial_result, sliced_metric.result())
+
+ sliced_metric(tf.constant([1, 2, 6]), tf.constant([1, 0, 1]))
+ self.assertDictEqual(
+ {
+ "test/mean": 3.0,
+ "test/mean/control": 2.0,
+ "test/mean/treatment": 3.5,
+ },
+ sliced_metric.result(),
+ )
+
+ sliced_metric.reset_state()
+ self.assertDictEqual(expected_initial_result, sliced_metric.result())
+
+ def test_metric_config(self):
+ sliced_metric = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.BinaryCrossentropy(
+ name="loss/bc", from_logits=True
+ )
+ )
+ self.assertLayerConfigurable(layer=sliced_metric)
+
+ def test_multi_output_result(self):
+ sliced_metric = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=MeanSquared(name="test_metric")
+ )
+
+ x1 = np.array([1, 2, 3])
+ x2 = np.array([-1, 4, -6])
+ x = np.concatenate([x1, x2], axis=0)
+
+ sliced_metric(tf.convert_to_tensor(x1), tf.zeros((3, 1)))
+ sliced_metric(tf.convert_to_tensor(x2), tf.ones((3, 1)))
+
+ expected_result = {
+ "test_metric/mean": x.mean(),
+ "test_metric/squared": x.mean() ** 2,
+ "test_metric/mean/control": x1.mean(),
+ "test_metric/squared/control": x1.mean() ** 2,
+ "test_metric/mean/treatment": x2.mean(),
+ "test_metric/squared/treatment": x2.mean() ** 2,
+ }
+ self.assertDictEqual(expected_result, sliced_metric.result())
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/uplift_mean.py b/official/recommendation/uplift/metrics/uplift_mean.py
new file mode 100644
index 00000000000..677548947d4
--- /dev/null
+++ b/official/recommendation/uplift/metrics/uplift_mean.py
@@ -0,0 +1,101 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for computing the mean uplift sliced by treatment group."""
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import treatment_sliced_metric
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class UpliftMean(tf_keras.metrics.Metric):
+ """Computes the overall and treatment sliced uplift mean.
+
+ Note that the prediction tensor is expected to be of type
+ `TwoTowerTrainingOutputs`.
+
+ Example standalone usage:
+
+ >>> uplift_mean = UpliftMean()
+ >>> y_pred = types.TwoTowerTrainingOutputs(
+ ... uplift=tf.constant([1, 2, 3, 4])
+ ... is_treatment=tf.constant([True, False, True, False]),
+ ... )
+ >>> uplift_mean(y_true=tf.zeros(4), y_pred=y_pred)
+ {
+ "uplift/mean": 2.5
+ "uplift/mean/control": 3.0
+ "uplift/mean/treatment": 2.0
+ }
+
+ Example usage with the `model.compile()` API:
+
+ >>> model.compile(
+ ... optimizer="sgd",
+ ... loss=TrueLogitsLoss(tf_keras.losses.mean_squared_error),
+ ... metrics=[UpliftMean()]
+ ... )
+ """
+
+ def __init__(self, name: str = "uplift/mean", **kwargs):
+ """Initializes the instance.
+
+ Args:
+ name: name for the overall uplift mean metric result. The control and
+ treatment uplift means will have "/control" and "/treatment" appended to
+ the result name.
+ **kwargs: other base metric keyword arguments.
+ """
+ super().__init__(name=name, **kwargs)
+ self._sliced_uplift = treatment_sliced_metric.TreatmentSlicedMetric(
+ metric=tf_keras.metrics.Mean(name=name, **kwargs)
+ )
+
+ def update_state(
+ self,
+ y_true: tf.Tensor,
+ y_pred: types.TwoTowerTrainingOutputs,
+ sample_weight: tf.Tensor | None = None,
+ ) -> None:
+ """Updates the overall, control and treatment uplift means.
+
+ Args:
+ y_true: tensor labels.
+ y_pred: two tower training outputs. The treatment indicator tensor is used
+ to slice the uplift prediction into control and treatment groups.
+ sample_weight: optional sample weight to compute weighted uplift means. If
+ given, the sample weight will also be sliced by the treatment indicator
+ tensor to compute the weighted control and treatment uplift means.
+
+ Raises:
+ TypeError: if y_pred is not of type `TwoTowerTrainingOutputs`.
+ """
+ del y_true
+
+ if not isinstance(y_pred, types.TwoTowerTrainingOutputs):
+ raise TypeError(
+ "y_pred must be of type `TwoTowerTrainingOutputs` but got type"
+ f" {type(y_pred)} instead."
+ )
+
+ self._sliced_uplift.update_state(
+ values=y_pred.uplift,
+ is_treatment=y_pred.is_treatment,
+ sample_weight=sample_weight,
+ )
+
+ def result(self) -> dict[str, tf.Tensor]:
+ return self._sliced_uplift.result()
diff --git a/official/recommendation/uplift/metrics/uplift_mean_test.py b/official/recommendation/uplift/metrics/uplift_mean_test.py
new file mode 100644
index 00000000000..bc4a379d5e4
--- /dev/null
+++ b/official/recommendation/uplift/metrics/uplift_mean_test.py
@@ -0,0 +1,216 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for uplift_mean."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import types
+from official.recommendation.uplift.metrics import uplift_mean
+
+
+class UpliftMeanTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ def _get_y_pred(
+ self, uplift: tf.Tensor, is_treatment: tf.Tensor
+ ) -> types.TwoTowerTrainingOutputs:
+ # Only the uplift and is_treatment tensors are required for testing.
+ return types.TwoTowerTrainingOutputs(
+ shared_embedding=tf.ones_like(is_treatment),
+ control_predictions=tf.ones_like(is_treatment),
+ treatment_predictions=tf.ones_like(is_treatment),
+ uplift=uplift,
+ control_logits=tf.ones_like(is_treatment),
+ treatment_logits=tf.ones_like(is_treatment),
+ true_logits=tf.ones_like(is_treatment),
+ true_predictions=tf.ones_like(is_treatment),
+ is_treatment=is_treatment,
+ )
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "uplift": tf.constant([0, 1, 5, 6]),
+ "is_treatment": tf.constant([[True], [False], [True], [False]]),
+ "sample_weight": None,
+ "expected_result": {
+ "uplift/mean": 3.0,
+ "uplift/mean/control": 3.5,
+ "uplift/mean/treatment": 2.5,
+ },
+ },
+ {
+ "testcase_name": "weighted",
+ "uplift": tf.constant([0, 1, 5, 6, -7]),
+ "is_treatment": tf.constant(
+ [[True], [False], [True], [True], [False]]
+ ),
+ "sample_weight": tf.constant([0.5, 0.5, 0, 0.7, 1.8]),
+ "expected_result": {
+ "uplift/mean": np.average(
+ [0, 1, 5, 6, -7], weights=[0.5, 0.5, 0, 0.7, 1.8]
+ ),
+ "uplift/mean/control": np.average([1, -7], weights=[0.5, 1.8]),
+ "uplift/mean/treatment": np.average(
+ [0, 5, 6], weights=[0.5, 0, 0.7]
+ ),
+ },
+ },
+ {
+ "testcase_name": "only_control",
+ "uplift": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[False], [False], [False]]),
+ "sample_weight": tf.constant([1, 0, 1]),
+ "expected_result": {
+ "uplift/mean": 2.5,
+ "uplift/mean/control": 2.5,
+ "uplift/mean/treatment": 0.0,
+ },
+ },
+ {
+ "testcase_name": "only_treatment",
+ "uplift": tf.constant([[0], [1], [5]]),
+ "is_treatment": tf.constant([[True], [True], [True]]),
+ "sample_weight": tf.constant([0, 1, 1]),
+ "expected_result": {
+ "uplift/mean": 3.0,
+ "uplift/mean/control": 0.0,
+ "uplift/mean/treatment": 3.0,
+ },
+ },
+ {
+ "testcase_name": "one_entry",
+ "uplift": tf.constant([2.5]),
+ "is_treatment": tf.constant([True]),
+ "sample_weight": tf.constant([1]),
+ "expected_result": {
+ "uplift/mean": 2.5,
+ "uplift/mean/control": 0.0,
+ "uplift/mean/treatment": 2.5,
+ },
+ },
+ {
+ "testcase_name": "no_entry",
+ "uplift": tf.constant([]),
+ "is_treatment": tf.constant([], dtype=tf.bool),
+ "sample_weight": tf.constant([]),
+ "expected_result": {
+ "uplift/mean": 0.0,
+ "uplift/mean/control": 0.0,
+ "uplift/mean/treatment": 0.0,
+ },
+ },
+ )
+ def test_metric_computes_sliced_uplift_means(
+ self, uplift, is_treatment, sample_weight, expected_result
+ ):
+ metric = uplift_mean.UpliftMean()
+ y_pred = self._get_y_pred(uplift=uplift, is_treatment=is_treatment)
+ metric(
+ y_true=tf.zeros_like(uplift), y_pred=y_pred, sample_weight=sample_weight
+ )
+ self.assertEqual(expected_result, metric.result())
+
+ def test_multiple_update_batches_returns_aggregated_uplift_means(self):
+ metric = uplift_mean.UpliftMean(name="uplift")
+
+ metric.update_state(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(
+ uplift=tf.constant([[1], [2], [4]]),
+ is_treatment=tf.constant([[True], [True], [True]]),
+ ),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(
+ uplift=tf.constant([[-3], [0], [5]]),
+ is_treatment=tf.constant([[False], [False], [False]]),
+ ),
+ sample_weight=None,
+ )
+ metric.update_state(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(
+ uplift=tf.constant([[0], [1], [-5]]),
+ is_treatment=tf.constant([[True], [False], [True]]),
+ ),
+ sample_weight=tf.constant([0.3, 0.25, 0.7]),
+ )
+
+ expected_results = {
+ "uplift": np.average(
+ [1, 2, 4, -3, 0, 5, 0, 1, -5],
+ weights=[1, 1, 1, 1, 1, 1, 0.3, 0.25, 0.7],
+ ),
+ "uplift/control": np.average([-3, 0, 5, 1], weights=[1, 1, 1, 0.25]),
+ "uplift/treatment": np.average(
+ [1, 2, 4, 0, -5], weights=[1, 1, 1, 0.3, 0.7]
+ ),
+ }
+ self.assertEqual(expected_results, metric.result())
+
+ def test_initial_and_reset_state_return_zero_uplift_means(self):
+ metric = uplift_mean.UpliftMean()
+
+ expected_initial_result = {
+ "uplift/mean": 0.0,
+ "uplift/mean/control": 0.0,
+ "uplift/mean/treatment": 0.0,
+ }
+ self.assertEqual(expected_initial_result, metric.result())
+
+ metric(
+ y_true=tf.zeros(3),
+ y_pred=self._get_y_pred(
+ uplift=tf.constant([1, 2, 6]),
+ is_treatment=tf.constant([[True], [False], [True]]),
+ ),
+ )
+ self.assertEqual(
+ {
+ "uplift/mean": 3.0,
+ "uplift/mean/control": 2.0,
+ "uplift/mean/treatment": 3.5,
+ },
+ metric.result(),
+ )
+
+ metric.reset_states()
+ self.assertEqual(expected_initial_result, metric.result())
+
+ def test_metric_config_is_serializable(self):
+ metric = uplift_mean.UpliftMean(name="test_name", dtype=tf.float16)
+ y_pred = self._get_y_pred(
+ uplift=tf.constant([[1], [2], [3], [4]]),
+ is_treatment=tf.constant([[True], [False], [True], [False]]),
+ )
+ self.assertLayerConfigurable(
+ layer=metric, y_true=tf.zeros(4), y_pred=y_pred, serializable=True
+ )
+
+ def test_invalid_prediction_tensor_type_raises_type_error(self):
+ metric = uplift_mean.UpliftMean()
+
+ with self.assertRaisesRegex(
+ TypeError, "y_pred must be of type `TwoTowerTrainingOutputs`"
+ ):
+ metric.update_state(y_true=tf.ones((3, 1)), y_pred=tf.ones((3, 1)))
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/metrics/variance.py b/official/recommendation/uplift/metrics/variance.py
new file mode 100644
index 00000000000..449448f1b73
--- /dev/null
+++ b/official/recommendation/uplift/metrics/variance.py
@@ -0,0 +1,72 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Keras metric for computing the (weighted) variance of a tensor."""
+
+from typing import Optional
+import tensorflow as tf, tf_keras
+
+
+class Variance(tf_keras.metrics.Metric):
+ """Computes the (weighted) variance of the given values.
+
+ For example, if values is [1, 2, 1, 4] then the variance is 1.5.
+ If the weights were specified as [1, 0, 1, 0] then the variance would be 0.
+
+ If `sample_weight` is `None`, weights default to 1.
+ Use `sample_weight` of 0 to mask values.
+
+ Standalone usage:
+
+ >>> m = Variance()
+ >>> m.update_state([1, 2, 1, 4])
+ >>> m.result().numpy()
+ 1.5
+ >>> m.reset_state()
+ >>> m.update_state([1, 2, 1, 4], sample_weight=[1, 0, 1, 0])
+ >>> m.result().numpy()
+ 0.0
+
+ Usage within a Keras layer:
+
+ ```python
+ layer.add_metric(Variance(name="variance")(values))
+ ```
+ """
+
+ def __init__(self, name: str = "variance", dtype: Optional[tf.DType] = None):
+ """Initializes a Variance metric instance.
+
+ Args:
+ name: (Optional) string name of the metric instance.
+ dtype: (Optional) data type of the metric result.
+ """
+ super().__init__(name=name, dtype=dtype)
+ self._first_moment = tf_keras.metrics.Mean(name="first_moment", dtype=dtype)
+ self._second_moment = tf_keras.metrics.Mean(
+ name="second_moment", dtype=dtype
+ )
+
+ def update_state(
+ self, values: tf.Tensor, sample_weight: Optional[tf.Tensor] = None
+ ):
+ self._first_moment.update_state(values=values, sample_weight=sample_weight)
+ self._second_moment.update_state(
+ values=tf.math.square(values), sample_weight=sample_weight
+ )
+
+ def result(self) -> tf.Tensor:
+ return self._second_moment.result() - tf.math.square(
+ self._first_moment.result()
+ )
diff --git a/official/recommendation/uplift/metrics/variance_test.py b/official/recommendation/uplift/metrics/variance_test.py
new file mode 100644
index 00000000000..90c4c0fe413
--- /dev/null
+++ b/official/recommendation/uplift/metrics/variance_test.py
@@ -0,0 +1,239 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for variance metric."""
+
+from typing import Optional
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift.metrics import variance
+
+
+class VarianceTest(keras_test_case.KerasTestCase, parameterized.TestCase):
+
+ def _compute_variance(
+ self, values: tf.Tensor, weights: Optional[tf.Tensor] = None
+ ) -> float:
+ values = values.numpy()
+
+ if weights is None:
+ return values.var()
+
+ weights = weights.numpy()
+ weights = np.broadcast_to(weights, shape=values.shape)
+ weighted_mean = np.average(values, weights=weights)
+ return np.average((values - weighted_mean) ** 2, weights=weights)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "values": tf.constant([-2, 0, 3, 5]),
+ "sample_weight": None,
+ },
+ {
+ "testcase_name": "weighted",
+ "values": tf.constant([-2, 0, 3, 5]),
+ "sample_weight": tf.constant([1, 0.3, 0.0, 1.5]),
+ },
+ {
+ "testcase_name": "negative_weights",
+ "values": tf.constant([-2, 0, 3, 5]),
+ "sample_weight": tf.constant([1, 0.3, 0.0, -1.5]),
+ },
+ )
+ def test_single_batch_correctness(self, values, sample_weight):
+ metric = variance.Variance()
+ metric(values=values, sample_weight=sample_weight)
+
+ expected_variance = self._compute_variance(values, sample_weight)
+ self.assertAllClose(expected_variance, metric.result())
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unweighted",
+ "values_batches": [tf.constant([-2, 0, 3, 5]), tf.constant([10])],
+ "sample_weight_batches": [None, None],
+ "all_values": tf.constant([-2, 0, 3, 5, 10]),
+ "all_weights": tf.ones((5,)),
+ },
+ {
+ "testcase_name": "weighted",
+ "values_batches": [tf.constant([-2, 0, 3, 5]), tf.constant([10, -4])],
+ "sample_weight_batches": [
+ tf.constant([1, 0.3, 0.0, -1.5]),
+ tf.constant([-4.0]),
+ ],
+ "all_values": tf.constant([-2, 0, 3, 5, 10, -4]),
+ "all_weights": tf.constant([1, 0.3, 0.0, -1.5, -4.0, -4.0]),
+ },
+ {
+ "testcase_name": "mix_weighted_and_unweighted",
+ "values_batches": [
+ tf.constant([-2.2, 0, 3, 5]),
+ tf.constant([10.5, -4]),
+ tf.ones((3,), dtype=tf.float32),
+ ],
+ "sample_weight_batches": [
+ tf.constant([1, 0.3, 0.0, -1.5]),
+ None,
+ None,
+ ],
+ "all_values": tf.constant([-2.2, 0, 3, 5, 10.5, -4, 1, 1, 1]),
+ "all_weights": tf.constant([1, 0.3, 0.0, -1.5, 1, 1, 1, 1, 1]),
+ },
+ )
+ def test_multi_batch_correctness(
+ self, values_batches, sample_weight_batches, all_values, all_weights
+ ):
+ metric = variance.Variance()
+
+ for values, sample_weight in zip(values_batches, sample_weight_batches):
+ metric(values=values, sample_weight=sample_weight)
+
+ expected_variance = self._compute_variance(all_values, all_weights)
+ self.assertAllClose(expected_variance, metric.result())
+ self.assertAllGreaterEqual(metric.result(), 0.0)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "unit_weight",
+ "values": tf.constant([0, 1, 2, 3]),
+ "sample_weight": tf.constant([1.0]),
+ "expected_variance": 1.25,
+ },
+ {
+ "testcase_name": "zero_weight",
+ "values": tf.constant([0, 1, 2, 3]),
+ "sample_weight": tf.constant([0.0]),
+ "expected_variance": 0.0,
+ },
+ {
+ "testcase_name": "decimal_weight",
+ "values": tf.constant([0, 1, 2, 3]),
+ "sample_weight": tf.constant([0.2]),
+ "expected_variance": 1.25,
+ },
+ {
+ "testcase_name": "negative_weight",
+ "values": tf.constant([0, 1, 2, 3]),
+ "sample_weight": tf.constant([-0.2]),
+ "expected_variance": 1.25,
+ },
+ )
+ def test_float_sample_weight(self, values, sample_weight, expected_variance):
+ metric = variance.Variance()
+ metric(values, sample_weight=sample_weight)
+ self.assertAllClose(expected_variance, metric.result())
+
+ def test_empty_input(self):
+ metric = variance.Variance()
+ values = tf.constant([0, 1, 2, 3])
+ metric(values)
+ self.assertAllClose(1.25, metric.result())
+ metric(tf.ones(shape=(0,)), sample_weight=None)
+ self.assertAllClose(1.25, metric.result())
+
+ def test_initial_state(self):
+ metric = variance.Variance()
+ self.assertAllClose(0.0, metric.result())
+
+ def test_dtype_correctness(self):
+ # 1 << 128 overflows for float32 but fits in float64.
+ value = tf.constant([1 << 128], dtype=tf.float64)
+
+ metric = variance.Variance(dtype=tf.float32)
+ metric(value)
+ self.assertAllEqual(np.nan, metric.result().numpy())
+
+ metric = variance.Variance(dtype=tf.float64)
+ metric(value)
+ self.assertAllEqual(0.0, metric.result().numpy())
+
+ def test_invalid_dtype(self):
+ with self.assertRaises(ValueError):
+ metric = variance.Variance(dtype=tf.string)
+ metric(tf.constant(["hello, world!"], tf.string))
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "squeeze_dimension_invalid",
+ "values": tf.ones((10, 10)),
+ "weights": tf.ones((10, 10, 10)),
+ },
+ {
+ "testcase_name": "dimension_mismatch",
+ "values": tf.ones((10, 10)),
+ "weights": tf.ones((10, 7)),
+ },
+ )
+ def test_invalid_weight_shape(self, values, weights):
+ metric = variance.Variance()
+ with self.assertRaises(tf.errors.InvalidArgumentError):
+ metric(values, weights)
+
+ def test_name(self):
+ metric = variance.Variance(name="test_name")
+ self.assertEqual("test_name", metric.name)
+
+ def test_multiple_result_calls(self):
+ metric = variance.Variance()
+
+ values = tf.constant([1, 2, 1, 4])
+ metric.update_state(values)
+
+ self.assertAllClose(values.numpy().var(), metric.result())
+ self.assertAllClose(values.numpy().var(), metric.result())
+
+ metric.update_state(tf.constant([-1, -2, 0]))
+
+ self.assertAllClose(
+ np.array([1, 2, 1, 4, -1, -2, 0]).var(), metric.result()
+ )
+
+ def test_reset_state(self):
+ metric = variance.Variance()
+ values = tf.constant([1, 2, 1, 4])
+
+ metric.update_state(values)
+ self.assertAllClose(1.5, metric.result())
+
+ metric.reset_state()
+
+ metric.update_state(values, sample_weight=tf.constant([1, 0, 1, 0]))
+ self.assertAllClose(0.0, metric.result())
+
+ def test_numpy_correctness(self):
+ metric = variance.Variance()
+
+ values = np.array([-1.3, 2.4, 1, 4])
+ weights = np.array([0.7, 0, 1.3, 1.0])
+
+ metric.update_state(values, weights)
+
+ expected_variance = self._compute_variance(
+ tf.convert_to_tensor(values), tf.convert_to_tensor(weights)
+ )
+ self.assertAllClose(expected_variance, metric.result())
+
+ def test_metric_config(self):
+ metric = variance.Variance()
+ self.assertLayerConfigurable(layer=metric)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/models/__init__.py b/official/recommendation/uplift/models/__init__.py
new file mode 100644
index 00000000000..bfb8d394002
--- /dev/null
+++ b/official/recommendation/uplift/models/__init__.py
@@ -0,0 +1,17 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Models package definition."""
+
+from official.recommendation.uplift.models import two_tower_uplift_model
diff --git a/official/recommendation/uplift/models/two_tower_uplift_model.py b/official/recommendation/uplift/models/two_tower_uplift_model.py
new file mode 100644
index 00000000000..af257976bde
--- /dev/null
+++ b/official/recommendation/uplift/models/two_tower_uplift_model.py
@@ -0,0 +1,131 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Defines a Keras model for the `TwoTowerUpliftNetwork` layer."""
+
+from __future__ import annotations
+
+from typing import Any, Callable, Mapping, MutableMapping
+
+import tensorflow as tf, tf_keras
+
+from official.recommendation.uplift import keys
+from official.recommendation.uplift import types
+from official.recommendation.uplift.layers.uplift_networks import base_uplift_networks
+from official.recommendation.uplift.layers.uplift_networks import two_tower_output_head
+
+
+@tf_keras.utils.register_keras_serializable(package="Uplift")
+class TwoTowerUpliftModel(tf_keras.Model):
+ """Training and inference model for a `BaseTwoTowerUpliftNetwork` layer."""
+
+ def __init__(
+ self,
+ treatment_indicator_feature_name: str,
+ uplift_network: base_uplift_networks.BaseTwoTowerUpliftNetwork,
+ inverse_link_fn: Callable[[tf.Tensor], tf.Tensor] | None = None,
+ **kwargs,
+ ):
+ """Initializes the instance.
+
+ Args:
+ treatment_indicator_feature_name: the name of the feature representing the
+ treatment_indicator tensor, which should be castable to a boolean tensor
+ (False for control and True for treatment). This tensor is required
+ during training and evaluation to compute the true logits needed for
+ loss computation.
+ uplift_network: a layer for computing control and treatment logits. Its
+ input is expected to be a dictionary of feature tensors and its output
+ is exptected to be a `TwoTowerNetworkOutputs` instance.
+ inverse_link_fn: a function for computing the control and treatment
+ predictions from their respective logits. If left as `None` it is
+ functionally equivalent to the identity function.
+ **kwargs: base model keyword arguments.
+ """
+ super().__init__(**kwargs)
+
+ self._treatment_indicator_feature_name = treatment_indicator_feature_name
+ self._uplift_network = uplift_network
+ self._inverse_link_fn = inverse_link_fn
+
+ self._output_head = two_tower_output_head.TwoTowerOutputHead(
+ treatment_indicator_feature_name=treatment_indicator_feature_name,
+ uplift_network=uplift_network,
+ inverse_link_fn=inverse_link_fn,
+ )
+
+ def call(
+ self,
+ inputs: types.DictOfTensors,
+ training: bool | None = None,
+ mask: tf.Tensor | None = None,
+ ) -> types.TwoTowerPredictionOutputs | types.TwoTowerTrainingOutputs:
+ return self._output_head(inputs=inputs, training=training, mask=mask)
+
+ def _assert_treatment_indicator_in_data(self, data):
+ inputs, _, _ = tf_keras.utils.unpack_x_y_sample_weight(data)
+
+ if self._treatment_indicator_feature_name not in inputs:
+ raise ValueError(
+ "The treatment_indicator feature (specified as"
+ f" '{self._treatment_indicator_feature_name}') must be part of the"
+ " inputs during training and evaluation, but got input features"
+ f" {set(inputs.keys())} instead."
+ )
+
+ def train_step(self, data) -> types.TwoTowerTrainingOutputs:
+ self._assert_treatment_indicator_in_data(data)
+ return super().train_step(data)
+
+ def test_step(self, data) -> types.TwoTowerTrainingOutputs:
+ self._assert_treatment_indicator_in_data(data)
+ return super().test_step(data)
+
+ def predict_step(self, data) -> dict[str, tf.Tensor]:
+ outputs = super().predict_step(data)
+
+ return {
+ keys.TwoTowerOutputKeys.CONTROL_PREDICTIONS: (
+ outputs.control_predictions
+ ),
+ keys.TwoTowerOutputKeys.TREATMENT_PREDICTIONS: (
+ outputs.treatment_predictions
+ ),
+ keys.TwoTowerOutputKeys.UPLIFT_PREDICTIONS: outputs.uplift,
+ }
+
+ def get_config(self) -> Mapping[str, Any]:
+ config = super().get_config()
+ config.update({
+ "treatment_indicator_feature_name": (
+ self._treatment_indicator_feature_name
+ ),
+ "uplift_network": tf_keras.utils.serialize_keras_object(
+ self._uplift_network
+ ),
+ "inverse_link_fn": tf_keras.utils.serialize_keras_object(
+ self._inverse_link_fn
+ ),
+ })
+ return config
+
+ @classmethod
+ def from_config(cls, config: MutableMapping[str, Any]) -> TwoTowerUpliftModel: # pyrefly: ignore[bad-override]
+ config["uplift_network"] = tf_keras.layers.deserialize(
+ config["uplift_network"]
+ )
+ config["inverse_link_fn"] = tf_keras.utils.deserialize_keras_object(
+ config["inverse_link_fn"]
+ )
+ return cls(**config)
diff --git a/official/recommendation/uplift/models/two_tower_uplift_model_test.py b/official/recommendation/uplift/models/two_tower_uplift_model_test.py
new file mode 100644
index 00000000000..a9081917f9c
--- /dev/null
+++ b/official/recommendation/uplift/models/two_tower_uplift_model_test.py
@@ -0,0 +1,278 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for two_tower_uplift_model."""
+
+from absl.testing import parameterized
+import numpy as np
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import keras_test_case
+from official.recommendation.uplift import keys
+from official.recommendation.uplift.layers.uplift_networks import two_tower_uplift_network
+from official.recommendation.uplift.losses import true_logits_loss
+from official.recommendation.uplift.metrics import loss_metric
+from official.recommendation.uplift.models import two_tower_uplift_model
+
+
+class TwoTowerUpliftModelTest(
+ keras_test_case.KerasTestCase, parameterized.TestCase
+):
+
+ def _get_uplift_network(self, **kwargs):
+ network = two_tower_uplift_network.TwoTowerUpliftNetwork(
+ backbone=kwargs.get(
+ "backbone",
+ tf_keras.layers.Lambda(lambda inputs: inputs["shared_feature"]),
+ ),
+ control_tower=kwargs.get("control_tower", tf_keras.layers.Dense(1)),
+ treatment_tower=kwargs.get("treatment_tower", tf_keras.layers.Dense(1)),
+ logits_head=tf_keras.layers.Identity(),
+ control_feature_encoder=kwargs.get(
+ "control_feature_encoder",
+ tf_keras.layers.Lambda(lambda inputs: inputs["control_feature"]),
+ ),
+ control_input_combiner=kwargs.get(
+ "control_input_combiner", tf_keras.layers.Concatenate()
+ ),
+ treatment_feature_encoder=kwargs.get(
+ "treatment_feature_encoder",
+ tf_keras.layers.Lambda(lambda inputs: inputs["treatment_feature"]),
+ ),
+ treatment_input_combiner=kwargs.get(
+ "treatment_input_combiner", tf_keras.layers.Concatenate()
+ ),
+ )
+ return network
+
+ def _get_compiled_model(self, **kwargs):
+ model = two_tower_uplift_model.TwoTowerUpliftModel(
+ treatment_indicator_feature_name="is_treatment",
+ uplift_network=self._get_uplift_network(**kwargs),
+ )
+ model.compile(
+ optimizer=tf_keras.optimizers.SGD(0.1),
+ loss=true_logits_loss.TrueLogitsLoss(
+ tf_keras.losses.mean_squared_error
+ ),
+ )
+ return model
+
+ def _get_inputs(self):
+ return {
+ "shared_feature": tf.ones((3, 1)),
+ "control_feature": tf.ones((3, 1)) * -2.0,
+ "treatment_feature": tf.ones((3, 1)) * 3.0,
+ }
+
+ def test_model_training_and_inference(self):
+ tf_keras.utils.set_random_seed(1)
+
+ # Create MSE uplift model.
+ uplift_network = self._get_uplift_network(
+ control_feature_encoder=None, control_input_combiner=None
+ )
+ model = two_tower_uplift_model.TwoTowerUpliftModel(
+ treatment_indicator_feature_name="is_treatment",
+ uplift_network=uplift_network,
+ )
+ model.compile(
+ optimizer=tf_keras.optimizers.SGD(0.1),
+ loss=true_logits_loss.TrueLogitsLoss(
+ tf_keras.losses.mean_squared_error
+ ),
+ )
+
+ # Create toy regression dataset.
+ shared_feature, treatment_feature = np.ones((10, 1)), 2 * np.ones((10, 1))
+ treatment = tf.constant([[1], [1], [0], [1], [1], [1], [0], [1], [0], [1]])
+ y = (shared_feature + treatment_feature) * treatment
+ dataset = tf.data.Dataset.from_tensor_slices((
+ {
+ "shared_feature": shared_feature,
+ "treatment_feature": treatment_feature,
+ "is_treatment": treatment,
+ },
+ y,
+ )).batch(5)
+
+ # Test model training.
+ history = model.fit(dataset, epochs=100)
+ self.assertIn("loss", history.history)
+ self.assertLen(history.history["loss"], 100)
+ self.assertTrue(
+ history.history["loss"][0] >= history.history["loss"][-1] >= 0.0
+ )
+
+ # Test model evaluation.
+ loss = model.evaluate(dataset)
+ self.assertLessEqual(loss, 1e-10)
+ self.assertGreaterEqual(loss, 0.0)
+ self.assertAllClose(history.history["loss"][-1], loss)
+
+ # Test model inference predictions.
+ expected_predictions = {
+ keys.TwoTowerOutputKeys.CONTROL_PREDICTIONS: tf.zeros((10, 1)),
+ keys.TwoTowerOutputKeys.TREATMENT_PREDICTIONS: 3 * tf.ones((10, 1)),
+ keys.TwoTowerOutputKeys.UPLIFT_PREDICTIONS: 3 * tf.ones((10, 1)),
+ }
+ self.assertAllClose(expected_predictions, model.predict(dataset))
+
+ def test_classification_model_trains(self):
+ tf_keras.utils.set_random_seed(1)
+
+ # Create binary classifier uplift model.
+ uplift_network = self._get_uplift_network(
+ control_feature_encoder=None, control_input_combiner=None
+ )
+ model = two_tower_uplift_model.TwoTowerUpliftModel(
+ treatment_indicator_feature_name="is_treatment",
+ uplift_network=uplift_network,
+ inverse_link_fn=tf.math.sigmoid,
+ )
+ model.compile(
+ optimizer=tf_keras.optimizers.SGD(0.1),
+ loss=true_logits_loss.TrueLogitsLoss(
+ loss_fn=tf_keras.losses.binary_crossentropy, from_logits=True
+ ),
+ metrics=[
+ loss_metric.LossMetric(
+ tf_keras.metrics.AUC(curve="PR", from_logits=True, name="aucpr")
+ ),
+ ],
+ )
+
+ # Create toy classification dataset.
+ treatment = tf.constant([[1], [1], [0], [1], [1], [1], [0], [1], [0], [1]])
+ y = treatment
+ dataset = tf.data.Dataset.from_tensor_slices((
+ {
+ "shared_feature": np.random.normal(size=(10, 1)),
+ "treatment_feature": np.random.normal(size=(10, 1)),
+ "is_treatment": treatment,
+ },
+ y,
+ )).batch(5)
+
+ # Test model training.
+ history = model.fit(dataset, epochs=100)
+ self.assertIn("loss", history.history)
+ self.assertLen(history.history["loss"], 100)
+ self.assertBetween(
+ history.history["loss"][-1], 0.0, history.history["loss"][0]
+ )
+ self.assertIn("aucpr", history.history)
+ self.assertLess(history.history["aucpr"][0], 1.0)
+ self.assertEqual(history.history["aucpr"][-1], 1.0)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "identity",
+ "inverse_link_fn": tf.identity,
+ "expected_predictions": {
+ keys.TwoTowerOutputKeys.CONTROL_PREDICTIONS: (
+ tf.ones((3, 1)) * -1.0
+ ), # 1 - 2 = -1
+ keys.TwoTowerOutputKeys.TREATMENT_PREDICTIONS: (
+ tf.ones((3, 1)) * 4.0
+ ), # 1 + 3 = 4
+ keys.TwoTowerOutputKeys.UPLIFT_PREDICTIONS: tf.ones((3, 1)) * 5.0,
+ },
+ },
+ {
+ "testcase_name": "abs",
+ "inverse_link_fn": tf.math.abs,
+ "expected_predictions": {
+ keys.TwoTowerOutputKeys.CONTROL_PREDICTIONS: (
+ tf.ones((3, 1)) * 1.0
+ ),
+ keys.TwoTowerOutputKeys.TREATMENT_PREDICTIONS: (
+ tf.ones((3, 1)) * 4.0
+ ),
+ keys.TwoTowerOutputKeys.UPLIFT_PREDICTIONS: tf.ones((3, 1)) * 3.0,
+ },
+ },
+ {
+ "testcase_name": "relu",
+ "inverse_link_fn": tf_keras.activations.relu,
+ "expected_predictions": {
+ keys.TwoTowerOutputKeys.CONTROL_PREDICTIONS: (
+ tf.ones((3, 1)) * 0.0
+ ),
+ keys.TwoTowerOutputKeys.TREATMENT_PREDICTIONS: (
+ tf.ones((3, 1)) * 4.0
+ ),
+ keys.TwoTowerOutputKeys.UPLIFT_PREDICTIONS: tf.ones((3, 1)) * 4.0,
+ },
+ },
+ )
+ def test_predict_step(self, inverse_link_fn, expected_predictions):
+ uplift_network = self._get_uplift_network(
+ control_tower=tf_keras.layers.Dense(1, kernel_initializer="ones"),
+ treatment_tower=tf_keras.layers.Dense(1, kernel_initializer="ones"),
+ )
+ model = two_tower_uplift_model.TwoTowerUpliftModel(
+ treatment_indicator_feature_name="is_treatment",
+ uplift_network=uplift_network,
+ inverse_link_fn=inverse_link_fn,
+ )
+ inputs = {
+ "shared_feature": tf.ones((3, 1)),
+ "control_feature": tf.ones((3, 1)) * -2.0,
+ "treatment_feature": tf.ones((3, 1)) * 3.0,
+ }
+ self.assertAllClose(expected_predictions, model.predict_step(inputs))
+
+ def test_missing_treatment_indicator_from_inputs_during_training_raises_value_error(
+ self,
+ ):
+ model = self._get_compiled_model()
+ inputs = {"x": tf.ones((3, 1))}
+ dataset = tf.data.Dataset.from_tensor_slices((inputs, tf.ones((3, 1))))
+
+ with self.assertRaises(ValueError):
+ model.fit(dataset)
+
+ def test_missing_treatment_indicator_from_inputs_during_evaluation_raises_value_error(
+ self,
+ ):
+ model = self._get_compiled_model()
+ inputs = {"x": tf.ones((3, 1))}
+ dataset = tf.data.Dataset.from_tensor_slices((inputs, tf.ones((3, 1))))
+
+ with self.assertRaises(ValueError):
+ model.evaluate(dataset)
+
+ def test_model_is_stable(self):
+ model = self._get_compiled_model()
+ inputs = self._get_inputs()
+ self.assertLayerStable(layer=model, inputs=inputs)
+
+ def test_model_is_savable(self):
+ model = self._get_compiled_model()
+ inputs = self._get_inputs()
+ self.assertModelSavable(model=model, inputs=inputs)
+
+ def test_layer_configurable(self):
+ # Cannot use lambda layers since they are not serializable.
+ model = self._get_compiled_model(
+ backbone=tf_keras.layers.Identity(),
+ control_feature_encoder=tf_keras.layers.Identity(),
+ treatment_feature_encoder=tf_keras.layers.Identity(),
+ inverse_link_fn=tf.math.sigmoid,
+ )
+ self.assertLayerConfigurable(layer=model)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/recommendation/uplift/two_tower_uplift_network.svg b/official/recommendation/uplift/two_tower_uplift_network.svg
new file mode 100644
index 00000000000..b0941f05803
--- /dev/null
+++ b/official/recommendation/uplift/two_tower_uplift_network.svg
@@ -0,0 +1 @@
+
\ No newline at end of file
diff --git a/official/recommendation/uplift/types.py b/official/recommendation/uplift/types.py
new file mode 100644
index 00000000000..6f48c1d7166
--- /dev/null
+++ b/official/recommendation/uplift/types.py
@@ -0,0 +1,121 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Defines types used by the keras uplift modeling library."""
+
+import tensorflow as tf, tf_keras
+
+TensorType = tf.Tensor | tf.SparseTensor | tf.RaggedTensor
+
+ListOfTensors = list[TensorType]
+TupleOfTensors = tuple[TensorType, ...]
+DictOfTensors = dict[str, TensorType]
+
+CollectionOfTensors = ListOfTensors | TupleOfTensors | DictOfTensors
+
+
+class TwoTowerNetworkOutputs(tf.experimental.ExtensionType):
+ """Tensors computed by a `BaseTwoTowerUpliftNetwork` layer.
+
+ Attributes:
+ shared_embedding: embedding computed by the backbone layer and used as a
+ shared representation for the control and treatment towers.
+ control_logits: logits for the control group. Its shape and dtype must be
+ the same as the treatment logits.
+ treatment_logits: logits for the treatment group. Its shape and dtype must
+ be the same as the control logits.
+ """
+
+ __name__ = "TwoTowerNetworkOutputs"
+
+ shared_embedding: tf.Tensor
+ control_logits: tf.Tensor
+ treatment_logits: tf.Tensor
+
+ # TODO(b/281776818): Override __validate__ to assert control and treatment
+ # logits are of the same dtype and shape. Also add validation tests.
+
+ # The `model.compile()` API casts and expands labels plus sample weights to
+ # match the shape and dtype of the model outputs. By setting the dtype and
+ # shape to that of the control and treatment logits the labels and sample
+ # weights get casted to the same dtype and shape as that of the logits.
+ dtype = property(lambda self: self.control_logits.dtype)
+ shape = property(lambda self: self.control_logits.shape)
+
+ class Spec:
+ """Tensor spec.
+
+ Note that this ExtensionType does not have a well defined shape since its
+ intended use is the same as that of a dataclass. A noteworthy case of when
+ the spec's shape attribute is needed is when a `KerasTensor` is initialized
+ during the construction of a functional Keras model, which expects all
+ tensors to have a spec with a shape attribute.
+ """
+
+ shape = property(lambda self: tf.TensorShape(None))
+
+
+class TwoTowerPredictionOutputs(TwoTowerNetworkOutputs):
+ """Inference tensors computed by a `TwoTowerUpliftNetwork` layer.
+
+ Attributes:
+ control_predictions: predictions for the control group. Its shape and dtype
+ must be the same as the treatment predictions.
+ treatment_predictions: predictions for the treatment group. Its shape and
+ dtype must be the same as the control predictions.
+ uplift: difference between the treatment and control predictions.
+ """
+
+ __name__ = "TwoTowerPredictionOutputs"
+
+ control_predictions: tf.Tensor
+ treatment_predictions: tf.Tensor
+ uplift: tf.Tensor
+
+ # TODO(b/281776818): Override __validate__ to assert control and treatment
+ # predictions are of the same dtype and shape as the control and treatment
+ # logits. Also add validation tests.
+
+ class Spec: # pyrefly: ignore[bad-override]
+ shape = property(lambda self: tf.TensorShape(None))
+
+
+class TwoTowerTrainingOutputs(TwoTowerPredictionOutputs):
+ """Training tensors computed by a `TwoTowerUpliftNetwork` layer.
+
+ Attributes:
+ true_logits: logits for either the control or treatment group, depending on
+ the corresponding value in the `is_treatment` tensor. It will contain
+ treatment group logits for the `is_treatment == 1` entries and control
+ group logits otherwise.
+ true_predictions: predictions for either the control or treatment group,
+ depending on the corresponding value in the `is_treatment` tensor. It will
+ contain treatment group predictions for the `is_treatment == 1` entries
+ and control group predictions otherwise.
+ is_treatment: a boolean `tf.Tensor` indicating if the example belongs to the
+ treatment group (True) or control group (False).
+ """
+
+ __name__ = "TwoTowerTrainingOutputs"
+
+ true_logits: tf.Tensor
+ true_predictions: tf.Tensor
+ is_treatment: tf.Tensor
+
+ # TODO(b/281776818): Override __validate__ to assert that the true logits is
+ # of the same rank as the control and treatment logits, and that the
+ # is_treatment tensor is a boolean tensor. Also add validation tests.
+
+ class Spec: # pyrefly: ignore[bad-override]
+ shape = property(lambda self: tf.TensorShape(None))
diff --git a/official/recommendation/uplift/uplift_modeling_intro.ipynb b/official/recommendation/uplift/uplift_modeling_intro.ipynb
new file mode 100644
index 00000000000..7f3e130b0c0
--- /dev/null
+++ b/official/recommendation/uplift/uplift_modeling_intro.ipynb
@@ -0,0 +1,1668 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Cho6oVk5Mvfd"
+ },
+ "source": [
+ "# Imports"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XCcXCYDWKAyU"
+ },
+ "outputs": [],
+ "source": [
+ "!pip3 install -q tf-models-nightly\n",
+ "# Fix Colab default opencv problem\n",
+ "!pip3 install -q opencv-python-headless==4.1.2.30"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "3669AfhUp2Z1"
+ },
+ "outputs": [],
+ "source": [
+ "import tensorflow_datasets as tfds\n",
+ "import tensorflow as tf\n",
+ "\n",
+ "import matplotlib.pyplot as plt\n",
+ "import numpy as np\n",
+ "\n",
+ "tf.keras.utils.set_random_seed(0)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "G6oWRL4bO_Hc"
+ },
+ "outputs": [],
+ "source": [
+ "concat_features = tfm.uplift.layers.encoders.concat_features\n",
+ "two_tower_logits_head = tfm.uplift.layers.heads.two_tower_logits_head\n",
+ "two_tower_uplift_network = tfm.uplift.layers.uplift_networks.two_tower_uplift_network\n",
+ "two_tower_uplift_model = tfm.uplift.models.two_tower_uplift_model\n",
+ "true_logits_loss = tfm.uplift.losses.true_logits_loss\n",
+ "treatment_fraction = tfm.uplift.metrics.treatment_fraction\n",
+ "uplift_mean = tfm.uplift.metrics.uplift_mean\n",
+ "label_mean = tfm.uplift.metrics.label_mean"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "cRD_we8jD5Bn"
+ },
+ "source": [
+ "# Introduction\n",
+ "\n",
+ "Uplift modeling is a crucial research area focused on estimating the causal influence of a treatment on an individual's behavior. It predicts how a customer's actions differ with and without the treatment. Within digital advertising for example, this treatment might involve exposure to various ads. Uplift modeling can then optimize marketing efforts by targeting users who demonstrate the most significant potential return on investment. This approach is essential because customers exhibit varied responses to treatments:\n",
+ "\n",
+ "- *persuadables*: these individuals consistently react positively to marketing, mostly purchasing only when exposed to the treatment\n",
+ "- *do not disturbs*: this group has a strongly negative reaction to marketing; they are likely to purchase if left untreated\n",
+ "- *lost causes*: these customers won't purchase regardless of marketing efforts, making spending ineffective\n",
+ "- *sure things*: these consumers will purchase irrespective of marketing, eliminating the need for targeted spending\n",
+ "\n",
+ "Therefore, uplift modeling aims to pinpoint the *persuadables*, conserve resources by avoiding *sure things* and *lost causes*, and prevent negative experiences for the *do not disturbs*.\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "bmqbMyHubN02"
+ },
+ "source": [
+ "# Dataset\n",
+ "\n",
+ "This colab walks through a very simple overview of how to use the uplift modeling library to build, train and evaluate an uplift model on large scale data. The dataset used in this colab is Criteo's uplift prediction [dataset](https://ailab.criteo.com/criteo-uplift-prediction-dataset/). This dataset is constructed by assembling data resulting from several incrementality tests, a particular randomized trial procedure where a random part of the population is prevented from being targeted by advertising.\n",
+ "\n",
+ "The data is collected such that at a pre-defined moment, users are randomly assigned to either a treated or control group. Before this assignment, user features (mainly related to prior activity) are collected. After assignment, the treated group receives personalized advertising while the control group does not. Ad visits and online conversions are logged for two weeks following assignment. Finally, the initial user features, treatment status, ad exposure, visits, and conversions are combined for analysis. Note that it is possible for a user in the treatment group to never get exposed to the treatment, which is indicated by the \"exposure\" feature. Nevertheless, we will ignore this feature in this analysis in order to keep the validity of the randomized control trial setting and to ensure that the uplift $u(x)$ is calculated by the difference of the following conditional expectations:\n",
+ "\n",
+ "\n",
+ "$$u(x) = \\mathbb{E}[Y | T=1, X=x] - \\mathbb{E}[Y | T=0, X=x]$$\n",
+ "\n",
+ "\n",
+ "The dataset consists of 14M rows, each one representing a user with eleven features, a treatment indicator and two possible labels (visits and conversions). The data fields are as follows:\n",
+ "\n",
+ "- f0, f1, f2, f3, f4, f5, f6, f7, f8, f9, f10, f11: anonymized feature values\n",
+ "- treatment: treatment group (0 = control, 1 = treatment)\n",
+ "- visit: whether a visit occured for this user (0 = did not visit, 1 = visited)\n",
+ "- conversion: whether a conversion occured for this user (0 = did not purchase, 1 = purchased)\n",
+ "- exposure: whether the user was exposed to a treatment during the experiment (0 = not exposed, 1 = exposed)\n",
+ "\n",
+ "A more comprehensive introduction to the Criteo dataset can be found in their accompanying papers [[1](https://hal.science/hal-02515860v1/document)] and [[2](https://arxiv.org/pdf/2111.10106.pdf)]."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 3218,
+ "status": "ok",
+ "timestamp": 1712877008190,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "oCAzKYhu9N6X",
+ "outputId": "d7d77311-dab5-4b16-9dba-9dc79b01e9db"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Number of datapoints: 13979592\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Load dataset (all of the data is stored under the \"train\" split).\n",
+ "full_dataset = tfds.load(\"criteo\")[\"train\"]\n",
+ "full_dataset = full_dataset.shuffle(\n",
+ " buffer_size=10_000,\n",
+ " seed=0,\n",
+ " reshuffle_each_iteration=False,\n",
+ ")\n",
+ "print(f\"Number of datapoints: {full_dataset.cardinality().numpy()}\")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 2,
+ "status": "ok",
+ "timestamp": 1712877008294,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "J-lzb-b39oOm",
+ "outputId": "3f4cc1ef-5e50-48f8-8b4e-8596c6912e14"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Number of train datapoints: 12979592\n",
+ "Number of test datapoints: 1000000\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Use 1M examples for testing and keep the rest for training.\n",
+ "N_TEST = 1_000_000\n",
+ "\n",
+ "test_dataset = full_dataset.take(N_TEST)\n",
+ "train_dataset = full_dataset.skip(N_TEST)\n",
+ "\n",
+ "print(f\"Number of train datapoints: {train_dataset.cardinality().numpy()}\")\n",
+ "print(f\"Number of test datapoints: {test_dataset.cardinality().numpy()}\")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "height": 363
+ },
+ "executionInfo": {
+ "elapsed": 25120,
+ "status": "ok",
+ "timestamp": 1712877033536,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "erj7SbIwsOFh",
+ "outputId": "d6470785-9161-4c96-db64-e89319ff6088"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/html": [
+ "\n",
+ " \u003cdiv id=\"df-cf51283c-b4cf-4a6e-9a57-e8d8043cb5c6\" class=\"colab-df-container\"\u003e\n",
+ " \u003cdiv\u003e\n",
+ "\u003cstyle scoped\u003e\n",
+ " .dataframe tbody tr th:only-of-type {\n",
+ " vertical-align: middle;\n",
+ " }\n",
+ "\n",
+ " .dataframe tbody tr th {\n",
+ " vertical-align: top;\n",
+ " }\n",
+ "\n",
+ " .dataframe thead th {\n",
+ " text-align: right;\n",
+ " }\n",
+ "\u003c/style\u003e\n",
+ "\u003ctable border=\"1\" class=\"dataframe\"\u003e\n",
+ " \u003cthead\u003e\n",
+ " \u003ctr style=\"text-align: right;\"\u003e\n",
+ " \u003cth\u003e\u003c/th\u003e\n",
+ " \u003cth\u003econversion\u003c/th\u003e\n",
+ " \u003cth\u003eexposure\u003c/th\u003e\n",
+ " \u003cth\u003ef0\u003c/th\u003e\n",
+ " \u003cth\u003ef1\u003c/th\u003e\n",
+ " \u003cth\u003ef10\u003c/th\u003e\n",
+ " \u003cth\u003ef11\u003c/th\u003e\n",
+ " \u003cth\u003ef2\u003c/th\u003e\n",
+ " \u003cth\u003ef3\u003c/th\u003e\n",
+ " \u003cth\u003ef4\u003c/th\u003e\n",
+ " \u003cth\u003ef5\u003c/th\u003e\n",
+ " \u003cth\u003ef6\u003c/th\u003e\n",
+ " \u003cth\u003ef7\u003c/th\u003e\n",
+ " \u003cth\u003ef8\u003c/th\u003e\n",
+ " \u003cth\u003ef9\u003c/th\u003e\n",
+ " \u003cth\u003etreatment\u003c/th\u003e\n",
+ " \u003cth\u003evisit\u003c/th\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003c/thead\u003e\n",
+ " \u003ctbody\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e0\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e26.003653\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-5.987667\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e1\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e19.585320\u003c/td\u003e\n",
+ " \u003ctd\u003e10.679513\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.217696\u003c/td\u003e\n",
+ " \u003ctd\u003e1.114982\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e3.013064\u003c/td\u003e\n",
+ " \u003ctd\u003e-10.527786\u003c/td\u003e\n",
+ " \u003ctd\u003e8.367581\u003c/td\u003e\n",
+ " \u003ctd\u003e3.809904\u003c/td\u003e\n",
+ " \u003ctd\u003e42.775658\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eTrue\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e2\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e12.616364\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.261647\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e0.294443\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.906514\u003c/td\u003e\n",
+ " \u003ctd\u003e25.240993\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e3\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e24.760006\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-4.595459\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e4\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e23.326521\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-7.301017\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e5\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e16.174679\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.425407\u003c/td\u003e\n",
+ " \u003ctd\u003e2.934780\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-3.282109\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.869313\u003c/td\u003e\n",
+ " \u003ctd\u003e29.214161\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eTrue\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e6\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e20.825842\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e3.013064\u003c/td\u003e\n",
+ " \u003ctd\u003e-1.288207\u003c/td\u003e\n",
+ " \u003ctd\u003e8.996727\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e7\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e21.404404\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e6.425741\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.800342\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e11.029585\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-7.301017\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.899112\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e8\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e24.595268\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-2.411115\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e9\u003c/th\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003ctd\u003e22.833614\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-2.411115\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1\u003c/td\u003e\n",
+ " \u003ctd\u003eFalse\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003c/tbody\u003e\n",
+ "\u003c/table\u003e\n",
+ "\u003c/div\u003e\n",
+ " \u003cdiv class=\"colab-df-buttons\"\u003e\n",
+ "\n",
+ " \u003cdiv class=\"colab-df-container\"\u003e\n",
+ " \u003cbutton class=\"colab-df-convert\" onclick=\"convertToInteractive('df-cf51283c-b4cf-4a6e-9a57-e8d8043cb5c6')\"\n",
+ " title=\"Convert this dataframe to an interactive table.\"\n",
+ " style=\"display:none;\"\u003e\n",
+ "\n",
+ " \u003csvg xmlns=\"http://www.w3.org/2000/svg\" height=\"24px\" viewBox=\"0 -960 960 960\"\u003e\n",
+ " \u003cpath d=\"M120-120v-720h720v720H120Zm60-500h600v-160H180v160Zm220 220h160v-160H400v160Zm0 220h160v-160H400v160ZM180-400h160v-160H180v160Zm440 0h160v-160H620v160ZM180-180h160v-160H180v160Zm440 0h160v-160H620v160Z\"/\u003e\n",
+ " \u003c/svg\u003e\n",
+ " \u003c/button\u003e\n",
+ "\n",
+ " \u003cstyle\u003e\n",
+ " .colab-df-container {\n",
+ " display:flex;\n",
+ " gap: 12px;\n",
+ " }\n",
+ "\n",
+ " .colab-df-convert {\n",
+ " background-color: #E8F0FE;\n",
+ " border: none;\n",
+ " border-radius: 50%;\n",
+ " cursor: pointer;\n",
+ " display: none;\n",
+ " fill: #1967D2;\n",
+ " height: 32px;\n",
+ " padding: 0 0 0 0;\n",
+ " width: 32px;\n",
+ " }\n",
+ "\n",
+ " .colab-df-convert:hover {\n",
+ " background-color: #E2EBFA;\n",
+ " box-shadow: 0px 1px 2px rgba(60, 64, 67, 0.3), 0px 1px 3px 1px rgba(60, 64, 67, 0.15);\n",
+ " fill: #174EA6;\n",
+ " }\n",
+ "\n",
+ " .colab-df-buttons div {\n",
+ " margin-bottom: 4px;\n",
+ " }\n",
+ "\n",
+ " [theme=dark] .colab-df-convert {\n",
+ " background-color: #3B4455;\n",
+ " fill: #D2E3FC;\n",
+ " }\n",
+ "\n",
+ " [theme=dark] .colab-df-convert:hover {\n",
+ " background-color: #434B5C;\n",
+ " box-shadow: 0px 1px 3px 1px rgba(0, 0, 0, 0.15);\n",
+ " filter: drop-shadow(0px 1px 2px rgba(0, 0, 0, 0.3));\n",
+ " fill: #FFFFFF;\n",
+ " }\n",
+ " \u003c/style\u003e\n",
+ "\n",
+ " \u003cscript\u003e\n",
+ " const buttonEl =\n",
+ " document.querySelector('#df-cf51283c-b4cf-4a6e-9a57-e8d8043cb5c6 button.colab-df-convert');\n",
+ " buttonEl.style.display =\n",
+ " google.colab.kernel.accessAllowed ? 'block' : 'none';\n",
+ "\n",
+ " async function convertToInteractive(key) {\n",
+ " const element = document.querySelector('#df-cf51283c-b4cf-4a6e-9a57-e8d8043cb5c6');\n",
+ " const dataTable =\n",
+ " await google.colab.kernel.invokeFunction('convertToInteractive',\n",
+ " [key], {});\n",
+ " if (!dataTable) return;\n",
+ "\n",
+ " const docLinkHtml = 'Like what you see? Visit the ' +\n",
+ " '\u003ca target=\"_blank\" href=https://colab.research.google.com/notebooks/data_table.ipynb\u003edata table notebook\u003c/a\u003e'\n",
+ " + ' to learn more about interactive tables.';\n",
+ " element.innerHTML = '';\n",
+ " dataTable['output_type'] = 'display_data';\n",
+ " await google.colab.output.renderOutput(dataTable, element);\n",
+ " const docLink = document.createElement('div');\n",
+ " docLink.innerHTML = docLinkHtml;\n",
+ " element.appendChild(docLink);\n",
+ " }\n",
+ " \u003c/script\u003e\n",
+ " \u003c/div\u003e\n",
+ "\n",
+ "\n",
+ "\u003cdiv id=\"df-05e44069-9608-4cf5-a0c5-2da213d958b4\"\u003e\n",
+ " \u003cbutton class=\"colab-df-quickchart\" onclick=\"quickchart('df-05e44069-9608-4cf5-a0c5-2da213d958b4')\"\n",
+ " title=\"Suggest charts\"\n",
+ " style=\"display:none;\"\u003e\n",
+ "\n",
+ "\u003csvg xmlns=\"http://www.w3.org/2000/svg\" height=\"24px\"viewBox=\"0 0 24 24\"\n",
+ " width=\"24px\"\u003e\n",
+ " \u003cg\u003e\n",
+ " \u003cpath d=\"M19 3H5c-1.1 0-2 .9-2 2v14c0 1.1.9 2 2 2h14c1.1 0 2-.9 2-2V5c0-1.1-.9-2-2-2zM9 17H7v-7h2v7zm4 0h-2V7h2v10zm4 0h-2v-4h2v4z\"/\u003e\n",
+ " \u003c/g\u003e\n",
+ "\u003c/svg\u003e\n",
+ " \u003c/button\u003e\n",
+ "\n",
+ "\u003cstyle\u003e\n",
+ " .colab-df-quickchart {\n",
+ " --bg-color: #E8F0FE;\n",
+ " --fill-color: #1967D2;\n",
+ " --hover-bg-color: #E2EBFA;\n",
+ " --hover-fill-color: #174EA6;\n",
+ " --disabled-fill-color: #AAA;\n",
+ " --disabled-bg-color: #DDD;\n",
+ " }\n",
+ "\n",
+ " [theme=dark] .colab-df-quickchart {\n",
+ " --bg-color: #3B4455;\n",
+ " --fill-color: #D2E3FC;\n",
+ " --hover-bg-color: #434B5C;\n",
+ " --hover-fill-color: #FFFFFF;\n",
+ " --disabled-bg-color: #3B4455;\n",
+ " --disabled-fill-color: #666;\n",
+ " }\n",
+ "\n",
+ " .colab-df-quickchart {\n",
+ " background-color: var(--bg-color);\n",
+ " border: none;\n",
+ " border-radius: 50%;\n",
+ " cursor: pointer;\n",
+ " display: none;\n",
+ " fill: var(--fill-color);\n",
+ " height: 32px;\n",
+ " padding: 0;\n",
+ " width: 32px;\n",
+ " }\n",
+ "\n",
+ " .colab-df-quickchart:hover {\n",
+ " background-color: var(--hover-bg-color);\n",
+ " box-shadow: 0 1px 2px rgba(60, 64, 67, 0.3), 0 1px 3px 1px rgba(60, 64, 67, 0.15);\n",
+ " fill: var(--button-hover-fill-color);\n",
+ " }\n",
+ "\n",
+ " .colab-df-quickchart-complete:disabled,\n",
+ " .colab-df-quickchart-complete:disabled:hover {\n",
+ " background-color: var(--disabled-bg-color);\n",
+ " fill: var(--disabled-fill-color);\n",
+ " box-shadow: none;\n",
+ " }\n",
+ "\n",
+ " .colab-df-spinner {\n",
+ " border: 2px solid var(--fill-color);\n",
+ " border-color: transparent;\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " animation:\n",
+ " spin 1s steps(1) infinite;\n",
+ " }\n",
+ "\n",
+ " @keyframes spin {\n",
+ " 0% {\n",
+ " border-color: transparent;\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " border-left-color: var(--fill-color);\n",
+ " }\n",
+ " 20% {\n",
+ " border-color: transparent;\n",
+ " border-left-color: var(--fill-color);\n",
+ " border-top-color: var(--fill-color);\n",
+ " }\n",
+ " 30% {\n",
+ " border-color: transparent;\n",
+ " border-left-color: var(--fill-color);\n",
+ " border-top-color: var(--fill-color);\n",
+ " border-right-color: var(--fill-color);\n",
+ " }\n",
+ " 40% {\n",
+ " border-color: transparent;\n",
+ " border-right-color: var(--fill-color);\n",
+ " border-top-color: var(--fill-color);\n",
+ " }\n",
+ " 60% {\n",
+ " border-color: transparent;\n",
+ " border-right-color: var(--fill-color);\n",
+ " }\n",
+ " 80% {\n",
+ " border-color: transparent;\n",
+ " border-right-color: var(--fill-color);\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " }\n",
+ " 90% {\n",
+ " border-color: transparent;\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " }\n",
+ " }\n",
+ "\u003c/style\u003e\n",
+ "\n",
+ " \u003cscript\u003e\n",
+ " async function quickchart(key) {\n",
+ " const quickchartButtonEl =\n",
+ " document.querySelector('#' + key + ' button');\n",
+ " quickchartButtonEl.disabled = true; // To prevent multiple clicks.\n",
+ " quickchartButtonEl.classList.add('colab-df-spinner');\n",
+ " try {\n",
+ " const charts = await google.colab.kernel.invokeFunction(\n",
+ " 'suggestCharts', [key], {});\n",
+ " } catch (error) {\n",
+ " console.error('Error during call to suggestCharts:', error);\n",
+ " }\n",
+ " quickchartButtonEl.classList.remove('colab-df-spinner');\n",
+ " quickchartButtonEl.classList.add('colab-df-quickchart-complete');\n",
+ " }\n",
+ " (() =\u003e {\n",
+ " let quickchartButtonEl =\n",
+ " document.querySelector('#df-05e44069-9608-4cf5-a0c5-2da213d958b4 button');\n",
+ " quickchartButtonEl.style.display =\n",
+ " google.colab.kernel.accessAllowed ? 'block' : 'none';\n",
+ " })();\n",
+ " \u003c/script\u003e\n",
+ "\u003c/div\u003e\n",
+ " \u003c/div\u003e\n",
+ " \u003c/div\u003e\n"
+ ],
+ "text/plain": [
+ " conversion exposure f0 ... f9 treatment visit\n",
+ "0 False False 26.003653 ... 13.190056 1 False\n",
+ "1 False False 19.585320 ... 42.775658 1 True\n",
+ "2 False False 12.616364 ... 25.240993 1 False\n",
+ "3 False False 24.760006 ... 13.190056 1 False\n",
+ "4 False False 23.326521 ... 13.190056 1 False\n",
+ "5 False False 16.174679 ... 29.214161 1 True\n",
+ "6 False False 20.825842 ... 13.190056 1 False\n",
+ "7 False False 21.404404 ... 13.190056 1 False\n",
+ "8 False False 24.595268 ... 13.190056 1 False\n",
+ "9 False False 22.833614 ... 13.190056 1 False\n",
+ "\n",
+ "[10 rows x 16 columns]"
+ ]
+ },
+ "execution_count": 5,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "# Take a small sample for exploratory data analysis.\n",
+ "sample_df = tfds.as_dataframe(train_dataset.take(10_000))\n",
+ "sample_df.head(10)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "height": 300
+ },
+ "executionInfo": {
+ "elapsed": 54,
+ "status": "ok",
+ "timestamp": 1712877033715,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "EhuI-cgq5cpI",
+ "outputId": "ecc6620a-8c3c-43c7-e4ad-bacbf3e15597"
+ },
+ "outputs": [
+ {
+ "data": {
+ "text/html": [
+ "\n",
+ " \u003cdiv id=\"df-2c8694ae-813f-48ee-9f9f-8c7746030f28\" class=\"colab-df-container\"\u003e\n",
+ " \u003cdiv\u003e\n",
+ "\u003cstyle scoped\u003e\n",
+ " .dataframe tbody tr th:only-of-type {\n",
+ " vertical-align: middle;\n",
+ " }\n",
+ "\n",
+ " .dataframe tbody tr th {\n",
+ " vertical-align: top;\n",
+ " }\n",
+ "\n",
+ " .dataframe thead th {\n",
+ " text-align: right;\n",
+ " }\n",
+ "\u003c/style\u003e\n",
+ "\u003ctable border=\"1\" class=\"dataframe\"\u003e\n",
+ " \u003cthead\u003e\n",
+ " \u003ctr style=\"text-align: right;\"\u003e\n",
+ " \u003cth\u003e\u003c/th\u003e\n",
+ " \u003cth\u003econversion\u003c/th\u003e\n",
+ " \u003cth\u003ef0\u003c/th\u003e\n",
+ " \u003cth\u003ef1\u003c/th\u003e\n",
+ " \u003cth\u003ef10\u003c/th\u003e\n",
+ " \u003cth\u003ef11\u003c/th\u003e\n",
+ " \u003cth\u003ef2\u003c/th\u003e\n",
+ " \u003cth\u003ef3\u003c/th\u003e\n",
+ " \u003cth\u003ef4\u003c/th\u003e\n",
+ " \u003cth\u003ef5\u003c/th\u003e\n",
+ " \u003cth\u003ef6\u003c/th\u003e\n",
+ " \u003cth\u003ef7\u003c/th\u003e\n",
+ " \u003cth\u003ef8\u003c/th\u003e\n",
+ " \u003cth\u003ef9\u003c/th\u003e\n",
+ " \u003cth\u003etreatment\u003c/th\u003e\n",
+ " \u003cth\u003evisit\u003c/th\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003c/thead\u003e\n",
+ " \u003ctbody\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003ecount\u003c/th\u003e\n",
+ " \u003ctd\u003e10000.00000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e10000.000000\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003emean\u003c/th\u003e\n",
+ " \u003ctd\u003e0.00250\u003c/td\u003e\n",
+ " \u003ctd\u003e19.599846\u003c/td\u003e\n",
+ " \u003ctd\u003e10.069551\u003c/td\u003e\n",
+ " \u003ctd\u003e5.332988\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.171190\u003c/td\u003e\n",
+ " \u003ctd\u003e8.448245\u003c/td\u003e\n",
+ " \u003ctd\u003e4.177443\u003c/td\u003e\n",
+ " \u003ctd\u003e10.338730\u003c/td\u003e\n",
+ " \u003ctd\u003e4.031238\u003c/td\u003e\n",
+ " \u003ctd\u003e-4.131609\u003c/td\u003e\n",
+ " \u003ctd\u003e5.097363\u003c/td\u003e\n",
+ " \u003ctd\u003e3.932956\u003c/td\u003e\n",
+ " \u003ctd\u003e16.088783\u003c/td\u003e\n",
+ " \u003ctd\u003e0.849700\u003c/td\u003e\n",
+ " \u003ctd\u003e0.046800\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003estd\u003c/th\u003e\n",
+ " \u003ctd\u003e0.04994\u003c/td\u003e\n",
+ " \u003ctd\u003e5.362421\u003c/td\u003e\n",
+ " \u003ctd\u003e0.094223\u003c/td\u003e\n",
+ " \u003ctd\u003e0.167144\u003c/td\u003e\n",
+ " \u003ctd\u003e0.024777\u003c/td\u003e\n",
+ " \u003ctd\u003e0.299626\u003c/td\u003e\n",
+ " \u003ctd\u003e1.336848\u003c/td\u003e\n",
+ " \u003ctd\u003e0.340552\u003c/td\u003e\n",
+ " \u003ctd\u003e0.412840\u003c/td\u003e\n",
+ " \u003ctd\u003e4.535934\u003c/td\u003e\n",
+ " \u003ctd\u003e1.191975\u003c/td\u003e\n",
+ " \u003ctd\u003e0.057004\u003c/td\u003e\n",
+ " \u003ctd\u003e7.101997\u003c/td\u003e\n",
+ " \u003ctd\u003e0.357383\u003c/td\u003e\n",
+ " \u003ctd\u003e0.211221\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003emin\u003c/th\u003e\n",
+ " \u003ctd\u003e0.00000\u003c/td\u003e\n",
+ " \u003ctd\u003e12.616364\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.827690\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e-5.106985\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e-2.404006\u003c/td\u003e\n",
+ " \u003ctd\u003e-23.423904\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.682439\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e0.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e0.000000\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e25%\u003c/th\u003e\n",
+ " \u003ctd\u003e0.00000\u003c/td\u003e\n",
+ " \u003ctd\u003e12.616364\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-6.699321\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.910792\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e0.000000\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e50%\u003c/th\u003e\n",
+ " \u003ctd\u003e0.00000\u003c/td\u003e\n",
+ " \u003ctd\u003e21.919387\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.214383\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e-2.411115\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e0.000000\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003e75%\u003c/th\u003e\n",
+ " \u003ctd\u003e0.00000\u003c/td\u003e\n",
+ " \u003ctd\u003e24.411213\u003c/td\u003e\n",
+ " \u003ctd\u003e10.059654\u003c/td\u003e\n",
+ " \u003ctd\u003e5.300375\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e8.732892\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e10.280525\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e0.294443\u003c/td\u003e\n",
+ " \u003ctd\u003e4.833815\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e13.190056\u003c/td\u003e\n",
+ " \u003ctd\u003e1.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e0.000000\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003ctr\u003e\n",
+ " \u003cth\u003emax\u003c/th\u003e\n",
+ " \u003ctd\u003e1.00000\u003c/td\u003e\n",
+ " \u003ctd\u003e26.745085\u003c/td\u003e\n",
+ " \u003ctd\u003e12.520096\u003c/td\u003e\n",
+ " \u003ctd\u003e6.473310\u003c/td\u003e\n",
+ " \u003ctd\u003e-0.168679\u003c/td\u003e\n",
+ " \u003ctd\u003e9.051891\u003c/td\u003e\n",
+ " \u003ctd\u003e4.679882\u003c/td\u003e\n",
+ " \u003ctd\u003e16.369278\u003c/td\u003e\n",
+ " \u003ctd\u003e4.115453\u003c/td\u003e\n",
+ " \u003ctd\u003e0.294443\u003c/td\u003e\n",
+ " \u003ctd\u003e11.991045\u003c/td\u003e\n",
+ " \u003ctd\u003e3.971858\u003c/td\u003e\n",
+ " \u003ctd\u003e57.158524\u003c/td\u003e\n",
+ " \u003ctd\u003e1.000000\u003c/td\u003e\n",
+ " \u003ctd\u003e1.000000\u003c/td\u003e\n",
+ " \u003c/tr\u003e\n",
+ " \u003c/tbody\u003e\n",
+ "\u003c/table\u003e\n",
+ "\u003c/div\u003e\n",
+ " \u003cdiv class=\"colab-df-buttons\"\u003e\n",
+ "\n",
+ " \u003cdiv class=\"colab-df-container\"\u003e\n",
+ " \u003cbutton class=\"colab-df-convert\" onclick=\"convertToInteractive('df-2c8694ae-813f-48ee-9f9f-8c7746030f28')\"\n",
+ " title=\"Convert this dataframe to an interactive table.\"\n",
+ " style=\"display:none;\"\u003e\n",
+ "\n",
+ " \u003csvg xmlns=\"http://www.w3.org/2000/svg\" height=\"24px\" viewBox=\"0 -960 960 960\"\u003e\n",
+ " \u003cpath d=\"M120-120v-720h720v720H120Zm60-500h600v-160H180v160Zm220 220h160v-160H400v160Zm0 220h160v-160H400v160ZM180-400h160v-160H180v160Zm440 0h160v-160H620v160ZM180-180h160v-160H180v160Zm440 0h160v-160H620v160Z\"/\u003e\n",
+ " \u003c/svg\u003e\n",
+ " \u003c/button\u003e\n",
+ "\n",
+ " \u003cstyle\u003e\n",
+ " .colab-df-container {\n",
+ " display:flex;\n",
+ " gap: 12px;\n",
+ " }\n",
+ "\n",
+ " .colab-df-convert {\n",
+ " background-color: #E8F0FE;\n",
+ " border: none;\n",
+ " border-radius: 50%;\n",
+ " cursor: pointer;\n",
+ " display: none;\n",
+ " fill: #1967D2;\n",
+ " height: 32px;\n",
+ " padding: 0 0 0 0;\n",
+ " width: 32px;\n",
+ " }\n",
+ "\n",
+ " .colab-df-convert:hover {\n",
+ " background-color: #E2EBFA;\n",
+ " box-shadow: 0px 1px 2px rgba(60, 64, 67, 0.3), 0px 1px 3px 1px rgba(60, 64, 67, 0.15);\n",
+ " fill: #174EA6;\n",
+ " }\n",
+ "\n",
+ " .colab-df-buttons div {\n",
+ " margin-bottom: 4px;\n",
+ " }\n",
+ "\n",
+ " [theme=dark] .colab-df-convert {\n",
+ " background-color: #3B4455;\n",
+ " fill: #D2E3FC;\n",
+ " }\n",
+ "\n",
+ " [theme=dark] .colab-df-convert:hover {\n",
+ " background-color: #434B5C;\n",
+ " box-shadow: 0px 1px 3px 1px rgba(0, 0, 0, 0.15);\n",
+ " filter: drop-shadow(0px 1px 2px rgba(0, 0, 0, 0.3));\n",
+ " fill: #FFFFFF;\n",
+ " }\n",
+ " \u003c/style\u003e\n",
+ "\n",
+ " \u003cscript\u003e\n",
+ " const buttonEl =\n",
+ " document.querySelector('#df-2c8694ae-813f-48ee-9f9f-8c7746030f28 button.colab-df-convert');\n",
+ " buttonEl.style.display =\n",
+ " google.colab.kernel.accessAllowed ? 'block' : 'none';\n",
+ "\n",
+ " async function convertToInteractive(key) {\n",
+ " const element = document.querySelector('#df-2c8694ae-813f-48ee-9f9f-8c7746030f28');\n",
+ " const dataTable =\n",
+ " await google.colab.kernel.invokeFunction('convertToInteractive',\n",
+ " [key], {});\n",
+ " if (!dataTable) return;\n",
+ "\n",
+ " const docLinkHtml = 'Like what you see? Visit the ' +\n",
+ " '\u003ca target=\"_blank\" href=https://colab.research.google.com/notebooks/data_table.ipynb\u003edata table notebook\u003c/a\u003e'\n",
+ " + ' to learn more about interactive tables.';\n",
+ " element.innerHTML = '';\n",
+ " dataTable['output_type'] = 'display_data';\n",
+ " await google.colab.output.renderOutput(dataTable, element);\n",
+ " const docLink = document.createElement('div');\n",
+ " docLink.innerHTML = docLinkHtml;\n",
+ " element.appendChild(docLink);\n",
+ " }\n",
+ " \u003c/script\u003e\n",
+ " \u003c/div\u003e\n",
+ "\n",
+ "\n",
+ "\u003cdiv id=\"df-656d0ac2-96c9-427f-a622-f7e5d831f9ce\"\u003e\n",
+ " \u003cbutton class=\"colab-df-quickchart\" onclick=\"quickchart('df-656d0ac2-96c9-427f-a622-f7e5d831f9ce')\"\n",
+ " title=\"Suggest charts\"\n",
+ " style=\"display:none;\"\u003e\n",
+ "\n",
+ "\u003csvg xmlns=\"http://www.w3.org/2000/svg\" height=\"24px\"viewBox=\"0 0 24 24\"\n",
+ " width=\"24px\"\u003e\n",
+ " \u003cg\u003e\n",
+ " \u003cpath d=\"M19 3H5c-1.1 0-2 .9-2 2v14c0 1.1.9 2 2 2h14c1.1 0 2-.9 2-2V5c0-1.1-.9-2-2-2zM9 17H7v-7h2v7zm4 0h-2V7h2v10zm4 0h-2v-4h2v4z\"/\u003e\n",
+ " \u003c/g\u003e\n",
+ "\u003c/svg\u003e\n",
+ " \u003c/button\u003e\n",
+ "\n",
+ "\u003cstyle\u003e\n",
+ " .colab-df-quickchart {\n",
+ " --bg-color: #E8F0FE;\n",
+ " --fill-color: #1967D2;\n",
+ " --hover-bg-color: #E2EBFA;\n",
+ " --hover-fill-color: #174EA6;\n",
+ " --disabled-fill-color: #AAA;\n",
+ " --disabled-bg-color: #DDD;\n",
+ " }\n",
+ "\n",
+ " [theme=dark] .colab-df-quickchart {\n",
+ " --bg-color: #3B4455;\n",
+ " --fill-color: #D2E3FC;\n",
+ " --hover-bg-color: #434B5C;\n",
+ " --hover-fill-color: #FFFFFF;\n",
+ " --disabled-bg-color: #3B4455;\n",
+ " --disabled-fill-color: #666;\n",
+ " }\n",
+ "\n",
+ " .colab-df-quickchart {\n",
+ " background-color: var(--bg-color);\n",
+ " border: none;\n",
+ " border-radius: 50%;\n",
+ " cursor: pointer;\n",
+ " display: none;\n",
+ " fill: var(--fill-color);\n",
+ " height: 32px;\n",
+ " padding: 0;\n",
+ " width: 32px;\n",
+ " }\n",
+ "\n",
+ " .colab-df-quickchart:hover {\n",
+ " background-color: var(--hover-bg-color);\n",
+ " box-shadow: 0 1px 2px rgba(60, 64, 67, 0.3), 0 1px 3px 1px rgba(60, 64, 67, 0.15);\n",
+ " fill: var(--button-hover-fill-color);\n",
+ " }\n",
+ "\n",
+ " .colab-df-quickchart-complete:disabled,\n",
+ " .colab-df-quickchart-complete:disabled:hover {\n",
+ " background-color: var(--disabled-bg-color);\n",
+ " fill: var(--disabled-fill-color);\n",
+ " box-shadow: none;\n",
+ " }\n",
+ "\n",
+ " .colab-df-spinner {\n",
+ " border: 2px solid var(--fill-color);\n",
+ " border-color: transparent;\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " animation:\n",
+ " spin 1s steps(1) infinite;\n",
+ " }\n",
+ "\n",
+ " @keyframes spin {\n",
+ " 0% {\n",
+ " border-color: transparent;\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " border-left-color: var(--fill-color);\n",
+ " }\n",
+ " 20% {\n",
+ " border-color: transparent;\n",
+ " border-left-color: var(--fill-color);\n",
+ " border-top-color: var(--fill-color);\n",
+ " }\n",
+ " 30% {\n",
+ " border-color: transparent;\n",
+ " border-left-color: var(--fill-color);\n",
+ " border-top-color: var(--fill-color);\n",
+ " border-right-color: var(--fill-color);\n",
+ " }\n",
+ " 40% {\n",
+ " border-color: transparent;\n",
+ " border-right-color: var(--fill-color);\n",
+ " border-top-color: var(--fill-color);\n",
+ " }\n",
+ " 60% {\n",
+ " border-color: transparent;\n",
+ " border-right-color: var(--fill-color);\n",
+ " }\n",
+ " 80% {\n",
+ " border-color: transparent;\n",
+ " border-right-color: var(--fill-color);\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " }\n",
+ " 90% {\n",
+ " border-color: transparent;\n",
+ " border-bottom-color: var(--fill-color);\n",
+ " }\n",
+ " }\n",
+ "\u003c/style\u003e\n",
+ "\n",
+ " \u003cscript\u003e\n",
+ " async function quickchart(key) {\n",
+ " const quickchartButtonEl =\n",
+ " document.querySelector('#' + key + ' button');\n",
+ " quickchartButtonEl.disabled = true; // To prevent multiple clicks.\n",
+ " quickchartButtonEl.classList.add('colab-df-spinner');\n",
+ " try {\n",
+ " const charts = await google.colab.kernel.invokeFunction(\n",
+ " 'suggestCharts', [key], {});\n",
+ " } catch (error) {\n",
+ " console.error('Error during call to suggestCharts:', error);\n",
+ " }\n",
+ " quickchartButtonEl.classList.remove('colab-df-spinner');\n",
+ " quickchartButtonEl.classList.add('colab-df-quickchart-complete');\n",
+ " }\n",
+ " (() =\u003e {\n",
+ " let quickchartButtonEl =\n",
+ " document.querySelector('#df-656d0ac2-96c9-427f-a622-f7e5d831f9ce button');\n",
+ " quickchartButtonEl.style.display =\n",
+ " google.colab.kernel.accessAllowed ? 'block' : 'none';\n",
+ " })();\n",
+ " \u003c/script\u003e\n",
+ "\u003c/div\u003e\n",
+ " \u003c/div\u003e\n",
+ " \u003c/div\u003e\n"
+ ],
+ "text/plain": [
+ " conversion f0 ... treatment visit\n",
+ "count 10000.00000 10000.000000 ... 10000.000000 10000.000000\n",
+ "mean 0.00250 19.599846 ... 0.849700 0.046800\n",
+ "std 0.04994 5.362421 ... 0.357383 0.211221\n",
+ "min 0.00000 12.616364 ... 0.000000 0.000000\n",
+ "25% 0.00000 12.616364 ... 1.000000 0.000000\n",
+ "50% 0.00000 21.919387 ... 1.000000 0.000000\n",
+ "75% 0.00000 24.411213 ... 1.000000 0.000000\n",
+ "max 1.00000 26.745085 ... 1.000000 1.000000\n",
+ "\n",
+ "[8 rows x 15 columns]"
+ ]
+ },
+ "execution_count": 6,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "sample_df[[\"visit\", \"conversion\"]] = sample_df[[\"visit\", \"conversion\"]].astype(int)\n",
+ "sample_df.describe()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 1,
+ "status": "ok",
+ "timestamp": 1712877033835,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "p0Asbitw6cJe",
+ "outputId": "1ba45783-fb60-4383-822c-6efabe609172"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Treatment fraction: 0.8497\n"
+ ]
+ }
+ ],
+ "source": [
+ "print(f\"Treatment fraction: {sample_df.treatment.mean()}\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "U4VNVQPMAi9O"
+ },
+ "source": [
+ "The data is highly imbalanced since 85% of the examples belong to the treatment group. This does not necessarily pose a problem for uplift modeling though, as we will discuss in greater detail in the next section."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 52,
+ "status": "ok",
+ "timestamp": 1712877034019,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "hV16GZeg758K",
+ "outputId": "035d0805-b866-4172-a870-eafd15a936e1"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Average visit rate: 0.0468\n",
+ "Average conversion rate: 0.0025\n"
+ ]
+ }
+ ],
+ "source": [
+ "print(f\"Average visit rate: {sample_df.visit.mean()}\")\n",
+ "print(f\"Average conversion rate: {sample_df.conversion.mean()}\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "SuJgzrRV8GXo"
+ },
+ "source": [
+ "The conversion rate is significantly lower than the visit rate, providing a sparser and therefore more challenging dataset for model training. Since a conversion follows from a visit, in this analysis we will focus on predicting the likelihood of visit occuring and leave the conversion estimation to a later stage."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 55,
+ "status": "ok",
+ "timestamp": 1712877034172,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "YXckAmwo65Sj",
+ "outputId": "ab11cad0-3ba9-4c32-872b-1973147b2429"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Visit rate (overall): 4.68%\n",
+ "Visit rate (control): 3.53%\n",
+ "Visit rate (treatment): 4.88%\n",
+ "Average treatment effect: 1.36%\n"
+ ]
+ }
+ ],
+ "source": [
+ "ctrl_visit_rate = sample_df[sample_df.treatment == 0].visit.mean() * 100\n",
+ "trmt_visit_rate = sample_df[sample_df.treatment == 1].visit.mean() * 100\n",
+ "visit_rate = sample_df.visit.mean() * 100\n",
+ "\n",
+ "print(f\"Visit rate (overall): {visit_rate:.2f}%\")\n",
+ "print(f\"Visit rate (control): {ctrl_visit_rate:.2f}%\")\n",
+ "print(f\"Visit rate (treatment): {trmt_visit_rate:.2f}%\")\n",
+ "print(f\"Average treatment effect: {trmt_visit_rate - ctrl_visit_rate:.2f}%\")"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "m215SfwI_ICf"
+ },
+ "source": [
+ "As expected, the treatment group has a higher visit rate (4.88%) than the control group (3.53%). The average treatment effect (1.36%) suggests that the treatment is effective at increasing the visit rate, and is a good starting point for uplift modeling. An uplift model is typically used to identify the set of users whose likelihood of visiting increases the most when exposed to the treatment."
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "YuM6Zo0CJL8F"
+ },
+ "source": [
+ "# Randomized Control Trial Test\n",
+ "\n",
+ "Randomized controlled trials are considered the gold standard for estimating causal effects because they help mitigate two major threats to drawing causal conclusions from data: confounding and selection bias. The randomization balances out the distribution of confounding variables across the control and treatment groups, which helps isolate the true treatment effect.\n",
+ "\n",
+ "Since uplift models aim to predict the difference in outcome between receiving a treatment and not receiving it, if the data is biased the model might learn patterns that do not reflect the true impact of the treatment. We can test if the data is truly random by training a classifier to predict whether an example belongs to the treatment group from its feature set. In a randomized control trial setting it should not be possible for a model to predict if an example belongs to the treatment group or not."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 499191,
+ "status": "ok",
+ "timestamp": 1712877533455,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "Y4wR_bTtJOFJ",
+ "outputId": "3e125912-b4ea-42ae-8fb1-662d4fd12d53"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Epoch 1/3\n",
+ "12676/12676 [==============================] - 203s 15ms/step - loss: 0.4236 - binary_accuracy: 0.8499 - auc: 0.5059 - val_loss: 0.4224 - val_binary_accuracy: 0.8502 - val_auc: 0.5096\n",
+ "Epoch 2/3\n",
+ "12676/12676 [==============================] - 155s 11ms/step - loss: 0.4228 - binary_accuracy: 0.8500 - auc: 0.5070 - val_loss: 0.4223 - val_binary_accuracy: 0.8502 - val_auc: 0.5097\n",
+ "Epoch 3/3\n",
+ "12676/12676 [==============================] - 141s 10ms/step - loss: 0.4227 - binary_accuracy: 0.8500 - auc: 0.5073 - val_loss: 0.4222 - val_binary_accuracy: 0.8502 - val_auc: 0.5099\n"
+ ]
+ }
+ ],
+ "source": [
+ "FEATURE_NAMES = [f\"f{i}\" for i in range(12)]\n",
+ "TREATMENT_NAME = \"treatment\"\n",
+ "\n",
+ "def preprocess(inputs: dict[str, tf.Tensor]) -\u003e tuple[tf.Tensor, tf.Tensor]:\n",
+ " features = tf.stack([inputs[name] for name in FEATURE_NAMES], axis=-1)\n",
+ " label = inputs[TREATMENT_NAME]\n",
+ " return features, label\n",
+ "\n",
+ "class Classifier(tf.keras.Model):\n",
+ " def __init__(self):\n",
+ " super().__init__()\n",
+ " self._mlp = tf.keras.Sequential([\n",
+ " tf.keras.layers.Dense(64, activation=\"relu\"),\n",
+ " tf.keras.layers.Dense(32, activation=\"relu\"),\n",
+ " tf.keras.layers.Dense(1)\n",
+ " ])\n",
+ "\n",
+ " def call(self, inputs: tf.Tensor) -\u003e tf.Tensor:\n",
+ " return self._mlp(inputs)\n",
+ "\n",
+ "classifier = Classifier()\n",
+ "classifier.compile(\n",
+ " optimizer=tf.keras.optimizers.SGD(),\n",
+ " loss=tf.keras.losses.BinaryCrossentropy(from_logits=True),\n",
+ " metrics=[\n",
+ " tf.keras.metrics.BinaryAccuracy(),\n",
+ " tf.keras.metrics.AUC(curve=\"ROC\", from_logits=True),\n",
+ " ]\n",
+ ")\n",
+ "\n",
+ "classifier.fit(\n",
+ " train_dataset.map(preprocess).batch(1024),\n",
+ " validation_data=test_dataset.map(preprocess).batch(1024),\n",
+ " epochs=3,\n",
+ ");"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Y0M9etg4ur5P"
+ },
+ "source": [
+ "The AUC value on the test split is indeed 0.5, indicating that the model cannot distinguish between examples in the control and treatment group and therefore does no better than random guessing. This validates the randomized control trial setting and ensures we have unbiased data that is perfect for uplift modeling!"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Fj9pqFESPgy3"
+ },
+ "source": [
+ "# Uplift Modeling\n",
+ "\n",
+ "The initial release of the uplift modeling library focuses on the family of models that follow the two tower uplift network architecture illustrated below. The library is written in Keras and is designed in a modular, layered manner such that it can easily be extended with custom layers and components. We offer a suite of tools (losses, metrics etc.) that can be hooked up with any Keras-compatible training framework (like Keras' training [API](https://keras.io/api/models/model_training_apis/), [Orbit](https://www.tensorflow.org/tfmodels/orbit) and [TFRS](https://www.tensorflow.org/recommenders)) in order to train an uplift model on potentially billions of datapoints.\n",
+ "\n",
+ "The two tower uplift model is composed of the following components:\n",
+ "- Inputs: a mapping from feature names to feature tensors. The tensors can be of different types, eg `tf.Tensor`, `tf.SparseTensor` and `tf.RaggedTensor`.\n",
+ "- Backbone: a trainable network that computes an embedding from the inputs shared between the control and treatment arms.\n",
+ "- Control and treatment feature encoders: trainable networks that compute embeddings from control and treatment speficic features.\n",
+ "- Control and treatment feature combiners: methods to combine the backbone's shared embedding with the control/treatment specific embeddings.\n",
+ "- Control tower: trainable network with zero or more hidden layers that learns from control examples only.\n",
+ "- Treatment tower: trainable network with zero or more hidden layers that learns from treatment examples only.\n",
+ "- Logits head: computes control and treatment logits. At training time, the gradient flows from the control logits for control examples and from the treatment logits for the treatment examples.\n",
+ "- Model outputs: contains the predicted control and treatment outcomes. The uplift is computed as the difference between the predicted treatment and control outcomes.\n",
+ "\n"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "PvyTGPkUTEUU"
+ },
+ "source": [
+ ""
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "Vx-zXSbDacY3"
+ },
+ "source": [
+ "Since we are interested in measuring the increase in probability of a visit occuring from the treatment, we will train an uplift model using the binary crossentropy loss. The treatment and control heads are two seperate classification heads that estimate the probability of a visit occuring with and without the treatment respectively. The uplift is then computed as the difference of the estimated treatment and control visit probabilities."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 509462,
+ "status": "ok",
+ "timestamp": 1712878043141,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "wC_ss1fCPhyP",
+ "outputId": "62e22a79-7d81-4cee-8066-717d0266f87f"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "Epoch 1/3\n",
+ "12676/12676 [==============================] - 169s 12ms/step - loss: 0.1194 - treatment_fraction: 0.8500 - uplift/mean: 0.0045 - uplift/mean/control: 0.0016 - uplift/mean/treatment: 0.0022 - label/mean: 0.0470 - label/mean/control: 0.0384 - label/mean/treatment: 0.0485 - val_loss: 0.1134 - val_treatment_fraction: 0.8502 - val_uplift/mean: 0.0055 - val_uplift/mean/control: 0.0051 - val_uplift/mean/treatment: 0.0056 - val_label/mean: 0.0467 - val_label/mean/control: 0.0379 - val_label/mean/treatment: 0.0483\n",
+ "Epoch 2/3\n",
+ "12676/12676 [==============================] - 170s 13ms/step - loss: 0.1143 - treatment_fraction: 0.8500 - uplift/mean: 0.0066 - uplift/mean/control: 0.0058 - uplift/mean/treatment: 0.0064 - label/mean: 0.0470 - label/mean/control: 0.0384 - label/mean/treatment: 0.0485 - val_loss: 0.1105 - val_treatment_fraction: 0.8502 - val_uplift/mean: 0.0069 - val_uplift/mean/control: 0.0064 - val_uplift/mean/treatment: 0.0070 - val_label/mean: 0.0467 - val_label/mean/control: 0.0379 - val_label/mean/treatment: 0.0483\n",
+ "Epoch 3/3\n",
+ "12676/12676 [==============================] - 170s 12ms/step - loss: 0.1125 - treatment_fraction: 0.8500 - uplift/mean: 0.0070 - uplift/mean/control: 0.0064 - uplift/mean/treatment: 0.0071 - label/mean: 0.0470 - label/mean/control: 0.0384 - label/mean/treatment: 0.0485 - val_loss: 0.1093 - val_treatment_fraction: 0.8502 - val_uplift/mean: 0.0062 - val_uplift/mean/control: 0.0057 - val_uplift/mean/treatment: 0.0063 - val_label/mean: 0.0467 - val_label/mean/control: 0.0379 - val_label/mean/treatment: 0.0483\n"
+ ]
+ }
+ ],
+ "source": [
+ "FEATURE_NAMES = [f\"f{i}\" for i in range(12)]\n",
+ "TREATMENT_NAME = \"treatment\"\n",
+ "LABEL_NAME = \"visit\"\n",
+ "\n",
+ "def preprocess(inputs: dict[str, tf.Tensor]) -\u003e dict[str, tf.Tensor]:\n",
+ " inputs = tf.nest.map_structure(lambda x: tf.expand_dims(x, axis=-1), inputs)\n",
+ " features = {name: inputs[name] for name in FEATURE_NAMES}\n",
+ " features[TREATMENT_NAME] = inputs[TREATMENT_NAME]\n",
+ " label = tf.cast(inputs[LABEL_NAME], tf.float32)\n",
+ " return features, label\n",
+ "\n",
+ "uplift_network = two_tower_uplift_network.TwoTowerUpliftNetwork(\n",
+ " backbone=tf.keras.Sequential([\n",
+ " concat_features.ConcatFeatures(feature_names=FEATURE_NAMES),\n",
+ " tf.keras.layers.Dense(64, activation=\"relu\"),\n",
+ " tf.keras.layers.Dropout(0.1),\n",
+ " tf.keras.layers.Dense(32, activation=\"relu\"),\n",
+ " tf.keras.layers.Dropout(0.1),\n",
+ " ]),\n",
+ " control_tower=tf.keras.Sequential([\n",
+ " tf.keras.layers.Dense(16, activation=\"relu\"),\n",
+ " tf.keras.layers.Dropout(0.1),\n",
+ " ]),\n",
+ " treatment_tower=tf.keras.Sequential([\n",
+ " tf.keras.layers.Dense(16, activation=\"relu\"),\n",
+ " tf.keras.layers.Dropout(0.1),\n",
+ " ]),\n",
+ " logits_head=two_tower_logits_head.TwoTowerLogitsHead(\n",
+ " control_head=tf.keras.layers.Dense(1),\n",
+ " treatment_head=tf.keras.layers.Dense(1),\n",
+ " ),\n",
+ ")\n",
+ "\n",
+ "uplift_model = two_tower_uplift_model.TwoTowerUpliftModel(\n",
+ " treatment_indicator_feature_name=TREATMENT_NAME,\n",
+ " uplift_network=uplift_network,\n",
+ " inverse_link_fn=tf.math.sigmoid,\n",
+ ")\n",
+ "\n",
+ "uplift_model.compile(\n",
+ " optimizer=tf.keras.optimizers.Adagrad(learning_rate=0.05),\n",
+ " loss=true_logits_loss.TrueLogitsLoss(\n",
+ " loss_fn=tf.keras.losses.binary_crossentropy,\n",
+ " from_logits=True,\n",
+ " ),\n",
+ " metrics=[\n",
+ " treatment_fraction.TreatmentFraction(),\n",
+ " uplift_mean.UpliftMean(),\n",
+ " label_mean.LabelMean(),\n",
+ " ],\n",
+ ")\n",
+ "\n",
+ "uplift_model.fit(\n",
+ " train_dataset.map(preprocess).batch(1024),\n",
+ " validation_data=test_dataset.map(preprocess).batch(1024),\n",
+ " epochs=3,\n",
+ ");"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "gAqAN6aTVEfJ"
+ },
+ "source": [
+ "# Uplift Evaluation\n",
+ "\n",
+ "There are various ways to evaluate uplift models. The *uplift curve* is a powerful visual tool that demonstrates the effectiveness of an uplift model in identifying the right users to target, compared to simply using random selection. While there are different versions of uplift curve, we will focus on the absolute cumulative uplift curve (introduced in Equation 17 of [[3](https://proceedings.mlr.press/v67/gutierrez17a/gutierrez17a.pdf)]) because of its relative simplicity.\n",
+ "\n",
+ "As stated before, the goal of an uplift model is to identify the set of users whose likelihood of visiting increases the most when exposed to the treatment. For any number of users *n*, we can estimate the expected number of incremental visits ($\\Delta$V) if *n* users were to be treated by:\n",
+ "\n",
+ "$$\\Delta V(n) = (\\frac{\\sum_{i=1}^n{V^T}}{N^T(n)} - \\frac{\\sum_{i=1}^n{V^C}}{N^C(n)}) \\cdot n $$\n",
+ "\n",
+ "where $N^T(n)$ is the number of treatment examples amongst the *n* individuals and $N^C(n)$ is the number of control examples amongst the *n* individuals.\n",
+ "\n",
+ "To compute the uplift curve we sort the test dataset in descending order of the predicted uplift, and compute the expected number of incremental visits if the top *n* users were to be treated, where *n* ranges from one to the entire test dataset. Since a good uplift model will sort the users with the highest treatment effect first, we would expect the uplift curve to show that the majority of incremental visits can be captured by treating just a small portion of targeted individuals. In particular, we would expect the number of incremental visits to always be greater than the incremental visits gained by random selection (which we can visualize by *not* sorting the test dataset)."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "hyfv6Tn3Fq7Z"
+ },
+ "outputs": [],
+ "source": [
+ "def uplift_curve(visit: np.ndarray, treatment: np.ndarray, uplift: np.ndarray | None = None) -\u003e np.ndarray:\n",
+ " \"\"\"Computes the incremental visits by different number of treated individuals.\"\"\"\n",
+ " # Sort by the uplift predictions if given. Otherwise the computed uplift curve\n",
+ " # will be that of a random selection strategy.\n",
+ " if uplift is not None:\n",
+ " sorted_indices = np.argsort(uplift)[::-1]\n",
+ " visit, treatment = visit[sorted_indices], treatment[sorted_indices]\n",
+ "\n",
+ " # Compute the cumulative number of control and treatment visits.\n",
+ " ctrl_visit, trmt_visit = visit.copy(), visit.copy()\n",
+ " ctrl_visit[treatment == 1] = 0\n",
+ " trmt_visit[treatment == 0] = 0\n",
+ " ctrl_visits = np.cumsum(ctrl_visit)\n",
+ " trmt_visits = np.cumsum(trmt_visit)\n",
+ "\n",
+ " # Compute the cumulative number of control and treatment individuals.\n",
+ " num_ctrl = np.cumsum(treatment == 0)\n",
+ " num_trmt = np.cumsum(treatment == 1)\n",
+ "\n",
+ " # Compute the visit rate for top n individuals, with n ranging from 1 to all.\n",
+ " avg_ctrl_visits = np.divide(ctrl_visits, num_ctrl, out=np.zeros_like(ctrl_visits, dtype=np.float32), where=num_ctrl \u003e 0)\n",
+ " avg_trmt_visits = np.divide(trmt_visits, num_trmt, out=np.zeros_like(trmt_visits, dtype=np.float32), where=num_trmt \u003e 0)\n",
+ "\n",
+ " # Estimate the expected number of incremental visits.\n",
+ " avg_treatment_effect = avg_trmt_visits - avg_ctrl_visits\n",
+ " expected_incremental_visits = avg_treatment_effect * (num_trmt + num_ctrl)\n",
+ " return expected_incremental_visits"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "XPIqu1OcCaEf"
+ },
+ "outputs": [],
+ "source": [
+ "# Extract label and treatment indicator from test data.\n",
+ "selector = lambda x: {TREATMENT_NAME: x[TREATMENT_NAME], LABEL_NAME: x[LABEL_NAME]}\n",
+ "test_df = tfds.as_dataframe(test_dataset.map(selector))\n",
+ "visit = test_df[LABEL_NAME].to_numpy().astype(np.int64)\n",
+ "treatment = test_df[TREATMENT_NAME].to_numpy().astype(np.int64)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "executionInfo": {
+ "elapsed": 17771,
+ "status": "ok",
+ "timestamp": 1712878216327,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "StBY-H2PgN0O",
+ "outputId": "872fd6e0-5edb-4d99-eaba-2e09c436542d"
+ },
+ "outputs": [
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "977/977 [==============================] - 17s 13ms/step\n"
+ ]
+ }
+ ],
+ "source": [
+ "# Compute uplift predictions on the test dataset.\n",
+ "predictions = uplift_model.predict(test_dataset.map(preprocess).batch(1024))\n",
+ "uplift_predictions = predictions[\"uplift_predictions\"].squeeze()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "id": "IrpPIGmMNUQ2"
+ },
+ "outputs": [],
+ "source": [
+ "# Compute incremental visists for random selection and for uplift targeting.\n",
+ "random_incremental_visits = uplift_curve(visit, treatment)\n",
+ "uplift_incremental_visits = uplift_curve(visit, treatment, uplift_predictions)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {
+ "colab": {
+ "height": 511
+ },
+ "executionInfo": {
+ "elapsed": 755,
+ "status": "ok",
+ "timestamp": 1712878217541,
+ "user": {
+ "displayName": "",
+ "userId": ""
+ },
+ "user_tz": 420
+ },
+ "id": "VZDIeXjlMJ5v",
+ "outputId": "57e0149c-bdd6-4c4d-f730-1fd52dc6487e"
+ },
+ "outputs": [
+ {
+ "data": {
+ "image/png": "iVBORw0KGgoAAAANSUhEUgAABcAAAAPdCAYAAACtK8udAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90\nbGliIHZlcnNpb24zLjYuMSwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/av/WaAAAACXBIWXMAABYl\nAAAWJQFJUiTwAAEAAElEQVR4nOzddZhc1f3H8ffZeIgRCAQCIQmQ4O7ukqBF69CWUkrxHy4JECBQ\nSimU0tIWKxQp2uLuHiC4hCQ4xN1Wzu+PO5uszOzO7o7szr5fzzPPzD33nHu/q4TPnjknxBiRJEmS\nJEmSJKnUlBW7AEmSJEmSJEmS8sEAXJIkSZIkSZJUkgzAJUmSJEmSJEklyQBckiRJkiRJklSSDMAl\nSZIkSZIkSSXJAFySJEmSJEmSVJIMwCVJkiRJkiRJJckAXJIkSZIkSZJUkgzAJUmSJEmSJEklyQBc\nkiRJkiRJklSSDMAlSZIkSZIkSSXJAFySJEmSJEmSVJIMwCVJkiRJkiRJJckAXJIkqYBCCKNCCDGE\ncGOxa2lMCGFSqtadil2L8i/1tY4hhEEFvu+g6nsX8r5tWaavVQjhiFT7MxnG9QwhXBFC+CyEsDjV\nd1IBSpYkSSoaA3BJkqQWCCEcUCOMeqzY9bR2IYQ+qT8CjCp2LQAhhM1CCFeGEN4OIUwNIZSHEKaF\nEF4JIVwWQliv2DW2Zamfj1Gl9EeUGj/vR+Syb4HcA5wEDAEWAN8DU2BJeD4qhLBRLm4UQugYQvhp\nCOHOEMLEEMLcEMKCEMKXIYT/hRCOCSH0ycW9JEmSGtKx2AVIkiS1cT+v8XrXEMIqMcavilZN69cH\nGJl6PapYRYQQegD/AA6r0VwJzAJ6A1umHqeGEG6PMf6w8FWWhANY+jPyTIY+5cDHhSimHZhF8rn8\nou6JEMK6wG4kn+8dYoyv1OlyBLAjMAl4uyVFhBC2Am4BVq/RvABYBKySeuwDXBxCOCbGeHtL7idJ\nktQQZ4BLkiQ1UwhhOWAEMB/4N8m/rX5S1KLUqBBCT+AFkvC7ArgW2BzoHGNcDugMbABcQBIoHl6k\nUtuFGOPXMca1YoxrFbuWti7GeG/qc/mzNKfXTT2/kyb8zpkQwh7A0yTh99fAb4EBMcbuMcY+wDIk\nvzfvJfmD2F75qkWSJAkMwCVJklriR0An4H7gb6m2n2furlbiOmBDYCEwIsb42xjjGzHGKoAYY1WM\n8d0Y40iSEO/hItYq5Uq31PPcfN0ghLAyyR8DuwJvABvGGK+NMX5T3SfGOD/G+FCM8QfALsB3+apH\nkiQJDMAlSZJaojrsvhV4nmTZgbVCCFtkMziEUBZCOCmEMC6EMC+19vR/GxofQhgcQrg2hPBJaj3d\n+SGEz0MIz4QQzgwhLJ9h3M4hhHtCCN+lNr/7LoRwbwhhl6Z+0Nls5BlCuDHVZ1SNtmeAiTWOY53H\nqDTXGRRCuDqE8HHqY50TQhgbQjg9hLBMM2rfhKUzus+JMTa4bnuMcRrJUg01r/FMY+s6Z9pAtObn\nLiSODSG8lfr6fxtCuCmEsEqN/mum2r4KISwMIbwXQjgqwz0b3ACx7v0b+rjrjOmQ+v75U+pz/33q\ne+ibTN9DIYSdQrKpZfXPyMi6X+8afdNughlC+DTV/rtG6ns01e+KNOc6hxB+F0J4PoQwPYSwKPXz\ncn0IYe1sPwe5VvdrFUL4eUjWnZ8dQpgVQngyhNDkmdHpvgeqv+bAjammHet8LY5Ind8xdf6GOucn\nNaGEM4HlSEL2g1M/PxnFGJ9OjamutdENUau/t9LVVfPnLoQwIITwlxDChNTX/e0Qwnap84tDCH0b\nuMeAEEJlqu8Gac7n9PeSJEnKLwNwSZKkZgjJerqbAtOAx2KMEbgtdTqbWeAB+A9wBbAOybq8fYF9\ngZdCCIfVG5CEt+OA3wBrAh1I1tQdSBJeXQxslmbcaOAp4EBgBWBe6vkA4MkQwiXZfMw5MB2YWuP4\n+zqPWjNTQwg/AD4EfgcMTTV3BjYBxgAvhxBWbGINv6lRyzXZDKieGZ4HtwF/Jvn6R6A/8DPguRBC\nv5Cso/xaqq0Hyce+LnBdCOHUPNWUztok3z/Hk3zuewOLgZVY+j10Vp0xi0m+pgtTx/Oo//VuTPXP\n048ydQghrADsmjr8d51zK5F8/q4GtkvVXf3zciTwZup7rKhCCH8kCac3J1mHvifJzOiHQwj/l4Nb\nzCX5fM9OHZdT++uwWuq5PHV+dp3zU7L8ODqTfF4B/hVj/Dybcanfnbk2lGQd82OAFVn6sb0IfE7y\nzp2DGxh/GMn/K78fY3yn5ok8/V6SJEl5ZAAuSZLUPNUh950xxupw5dbU8+GpMKgh+6ceJwO9Umvj\nrgE8ThJs3xBCWL3OmMtJwrFXgU1ijJ1jjMuSrKm7OXAlyZrVS4QQDgfOTh3+GVghNaYfSTAIcEYI\nIe9rl6eWPNi8xnH/Oo/La9S9OXA7SVB1KUlItwzQHdiK5HOwPnBzE8vYOfX8eIxxYYM98+sAknWQ\nf0LyNe0J7ECyHMRg4EKSj/8FYPXU90cf4K+p8ReEZA36QlhM8seafUlC+m4xxh4kweK5JKHt6BDC\nltUDYowvxRj7A3ekmi6v+/XO4r7VP09bhxAGZehzCMnPy6cxxjeqG0MI1UsTbQg8R/K57RZj7JX6\nGP5AskzHv9L8nBXSRsCJJN/jfVM/mwNY+rFfFkLYriU3iDFenvp8n5BqeqnO1+L81PmXUudPqHN+\n8/RXrmdzkp9RgP+2pOYc+APwLbBtjHGZ1PfrwamwvXrDzYY2tq0+V/ePKvn6vSRJkvLIAFySJKmJ\nQggdWLrZ5ZKAJMb4LvAuS2dyN6Q3MDLG+McY44LU+M+A/YCPSdbrPbPOmK1SzyfEGN+qcd/5qTWs\nT4oxvlyjzkASpALcHmM8LsY4NTVmWozxeJbOsh0dQmhN/zb8I0nIdGqM8YwY4xcxURljfBXYG/gG\n2COEUG/WezqpUHSN1OG4vFSdvd7A72KMt8YYF6c+tueB01LnjyYJng+MMU4AiDHOBo4FxpOEtyMK\nUWiM8ZMY46ExxgdijN9Xz9iNMU6OMY4Gzid5R8NvGrxQ0+/7MVD9fZ4prEwbVJL8gWpz4HVgjxjj\n8zHGxanrfh9j/D+SzU+7Ayflsu4m6g38I/U9PitV37fAT0k2kgzAqOKV1yQ1l5Qp9s9XBbB7jLE6\n1CfGOD71svqPCzuEEAbUHRhCWIOl76S5rc7pnP9ekiRJ+dea/idHkiSprdiDZPmHz0neUl9TdbjS\n2DIo80lmbNeSmpX8h9ThQakQu1r1EgYrZVnnRiwNfEdn6HN+6nk1IKu1y/MtNSN3W2ABS2c81xJj\nnMHSzSl3z/LSNdf8nd7sAnPjK+BfadqfqPH69zHGiponU8uxPJ06XC9PtTXV/1LP2+bh2tXBdr0A\nPIQwENimTr9q1T9/18QYFzVy7Wy/f/Ll4roNqT8yVC9NtEtD61W3IjXfkVDsn6+bY4xpl9lJ/aHy\nPZL/F6631BRLv9dejjHW3LMgX7+XJElSnhmAS5IkNV11uHZbmvVrbyNZz3nvEEK/Bq7xRoxxXoZz\nz6ae+5Ash1HtodTzzSGEMSGErVKzmjPZJPU8Jcb4froOqVm2X9fpX2zVoWZnYGJINuys92DpZpar\nFqfMFvkgw9rik2u8fi/D2Opgb9nclpRZCKFbSDZsfSaEMDmEUF5js8LqWdor5+HWtwFVwPohhLqB\n/w9JZkiPjTF+UqPWjiz9Y84VDXz/3JvqU8zvny9qhqx1vECyvEwg+WOWsvdyI+cbWl8+07sK2sPv\nJUmSSpIBuCRJUhOEEHqTrN0N9QMSYoxfAM8DHWlg8z6Whs6NnasZop9Ksk5vT+B0kpBndgjhqRDC\nMSGEbnWuUz22oXtBMhu57r2KqXqGeweSdaYzParXG+6e5XVrzkot9ozab9M1xhgrG+tDEopCshRD\n3qU2k3ybZMPWHUm+TxaRbI74PUs3Nl0m3fiWiDF+TfLzBPVngWcKKvuShJTVrzN9/yyf6lP356aQ\nMv5sppZGmpE6bC0/mw2ZVuN1sX++Gtu4s/p7ZtMQwprVjSGEjUiWcqkE7qwzJl+/lyRJUp4ZgEuS\nJDXNYSTrLwO8Uz0LtuaDZMM9aHwZlExCusYY4zRgO5K31l9FMvO2M8nGjn8B3gshrJJmaJdm1lEs\n1f9GfSvGGLJ4HJHNRWOyWWn1OsAb5qPwEnUlMBSYABxEslljjxjjCqnNE7dqaHAO1FsGJYSwNsnX\nsIqlmxpWq/n/OBtm8z3UxHqql1RpMDgPIdQMQBc08R5LLtPMccXwYY3Xxf75qmzoZIxxEktnidf8\nQ2X199gTMcbJ1JaX30uSJCn/DMAlSZKapimh9sYhhPUznGtouYiaa3zXmsmY2nDtiRjjCTHGTUhm\nsR5NMrt5CMkmbXXHDmykzurQvLFZk9Wq16Xu2kCf3lleK53qJT7WTC1nkUvV62fvHkJoqP6G5Pvj\nb66c1xVC6MzSdzz8OMZ4T2qd45pWbMo1m+E/JBuCDg4hVIft1UHlszHGb+r0n8bSAHSdPNRTPdO5\nsbX4a56flqFPxt8Dqe/PPqnDbH82i+l1oHpZp/2aeY0la9438POZq5+tWn9YSe23cHidczXl8/eS\nJEnKIwNwSZKkLIUQ1mDpOrAbkazBnOlRvTFgpsB88zozRGvaMfU8E8i0PjCQbLoWY7wOOKvOWIA3\nU8/LhBDSbnAZQhgKDKjTvzEzU8/pZptXB0mbZhhbVadfOtUzM3uQbDiaS9elnvsCv81mQJo6Z6ae\nM338a7A0uCykmanntHWlbN7Eay7P0ncQvJWhz24NjK/+ejd7JnMqcH80dVg9WzdjUJma6f9G6vAH\nzb1vA6o/D9s02Kv2+bcz9FkthDAow7ntSJbbiA2Mz6UWfa1ijIuBG1OHP01tUtqoOj9fM2u8zvR9\n3NTv4UzuIAnch4UQNiH5eg0kma1/b5r++fy9JEmS8sgAXJIkKXvVYfa4GOO4GOPMTA+SWasAPw4h\ndEhzre7ACXUbQwhdgJNTh3dVb7IZQihrZNZh9RILNZc7eZulS36cRXqjUs+TgNcauH5N76aeN0+t\nD13Xj8m8AdzsGq/7pOsQY/wIeCV1eGkIIePa0qnNGbNe4iXG+AZL1/a9KISwe0P9QwjLAQ/Waa7+\n+DPNcj0j23pyrLquASGEen+ACCFsD2zbxGvOJglgAeq9myH19T+ukfHQ8j8IVAfdh6Zmga9JMiv8\n7gz9b0w9HxRC2LmhC4cQmrqZaPU9d04Fp+mu2ZGlP9/PxRinpuuXcmaa8YGl30dPxhin1+2TB7n4\nWl1C8m6UHsBdIYQG1wJPfW0uqT6OMc4l+V0ES995ULP/csCvWlDfEjHGKcCTqcMfsvSPKw/EGOek\n6Z+330uSJCm/DMAlSZKykAqkfpo6vCeLIf8DyoH+wJ5pzs8CLgwhnFC9eWUIYQhwP8kmbAuBMTX6\n9wLGhxDODiGsXx2qp4LxXYGLUv2qZ8qSCs/PSR3uH0K4OhUgEUJYLoRwFUuXkjgnxrhkdnYjXgS+\nIVl//LYQwuDUNbuHEI4G/s7SzftqSf1xoHrJiiMbuMdxJGstrwc8H0LYrfoPAKmPed0QwjnAZzS+\nFEVdR5GExV2Bh0II14QQNgkhlKWuH0II64UQzk9df+864+8iCYXXDyH8KYTQJzVuhdTn9KfA/CbW\n1GIxxs9Z+keMG6uX3wkhdAohHALcR4avSwPXnMvS0O/61CaBNb/vnqXhGcPvp573yvDHkmz9F5hL\nstzKNam2h9Msx1Ltn6m6y4AHUj9nS8LY1NfqhyGEZ0jzh6hG/Bt4L3Xth0IIh1e/myP1vbMx8ADJ\nuyAqWfozmM5s4NchhItTG+wSQugP3ATsSvJ9dn4T62uu6q/VD6praarUpqU/JvnZ3RwYF5INepd8\n7VO/J/YOIdwNPEXyO7Km6j9QnRNC2K/Gz/1WwBMs3eA0F6r/sHI4cEidtnTy+XtJkiTlS4zRhw8f\nPnz48OHDRyMPko0mY+qxbpZjHkn1v6NG26hU200kQXokmck6o8b1K4DD61yrT43z1WOmpfpWt30G\nrJKmjtE1+lSSzNCsrNF2SYb6J6XO75Tm3IF1rjGLJPCPJOHjjanXo9KMPb/GuOoZn5OAE+v025tk\nSYTqvouAqamPvebnYrVmfD17kczSr3mditTntLxGWxVwU5rxV9QZOyPVtwI4ItPnrsbX/8YGaqu+\n5qAM5zNeA9iSJHyvvsac1OctkvxxZHQDY9PeN80159Y4nkYyUzeS+ptLnbHLp/pUf+99W/31rtFn\nUKbxda51S53P+aGN9F8BeKHO13J66nNS8zojm/H9M4Rk08ea3ztTSd6JUfP79YgM449I9XmGZN3+\n6mtMT9VZfY3/a8r3SM3rNnTPDNdcq8b3Sjnwdepr9UIzPj/bkmyaWvPzPI/aP8+RZG3zg+qMXZbk\nd1l1n4Wp77kIfA78JPV6Upr7TiLD76wMdfas8709A+jSyJi8/V7y4cOHDx8+fOTn4QxwSZKk7Pw8\n9fxJjPH9BnsuVb1Uwv7Vs4RriCQzDk8mCdI6k4QvDwDbxBhvr9N/NrAPcCXJLN8pJOHNPJLN584G\nNooxflW3iBjjOSSzSe8nCWp6kISS/wV2izHWW4KhMTHGe0nWwX2aJFDsQLLkyq9ijL9sZPgFwOnA\nOySzh1dLPfrUucfDwFCS0PZNkiCsD8nn4iXgPGDtmMx8bmr9s2OMhwBbAFeTzAifTRKMzyaZPTwm\ndf2fp7nEKSRriI9L1VUdMO8SY7yxqfXkSozxVZK1o/9HEtJ1BD4BTgWGU2OTwSZec2uWziDvBEwG\n/kayFv64BsZOJfnj0T0k37P9WPr1bqpba7yey9J19jPdezLJmvg/Bh5K1dyD5HvuI5I/1AwHLm5q\nITHGCcAmwO9IltGYTrI5YznJ99KfgHWy+V6IMZ5EEk6PJfl6zSX5udo7xnh5U2trrpgs8bE7yR/u\nZpHMzF6NhteUz3StF4FhJL837yYJriF518VXJL/nfg0MjjHeXWfsDJL1uK8jebdIGcnvq6tJPuf1\nfsc1V0yWOqn5fXR3jHFRI2Py9ntJkiTlR4gxFrsGSZIkSWo3QghHADcAz8YYdypuNZIkSaXNGeCS\nJEmSJEmSpJJkAC5JkiRJkiRJKkkG4JIkSZIkSZKkkmQALkmSJEmSJEkqSW6CKUmSJEmSJEkqSc4A\nlyRJkiRJkiSVJANwSZIkSZIkSVJJMgCXJEmSJEmSJJUkA3BJkiRJkiRJUkkyAJckSZIkSZIklaSO\nxS5A+RFCmAj0AiYVuRRJkiRJkiRJaolBwOwY4+CmDjQAL129unXr1nfttdfuW+xCJEmSJEmSJKm5\nPvzwQxYsWNCssQbgpWvS2muv3Xfs2LHFrkOSJEmSJEmSmm3TTTflzTffnNScsa4BLkmSJEmSJEkq\nSQbgkiRJkiRJkqSSZAAuSZIkSZIkSSpJBuCSJEmSJEmSpJJkAC5JkiRJkiRJKkkG4JIkSZIkSZKk\nkmQALkmSJEmSJEkqSQbgkiRJkiRJkqSSZAAuSZIkSZIkSSpJBuCSJEmSJEmSpJJkAC5JkiRJkiRJ\nKkkG4JIkSZIkSZKkkmQALkmSJEmSJEkqSQbgkiRJkiRJkqSSZAAuSZIkSZIkSSpJBuCSJEmSJEmS\npJJkAC5JkiRJkiRJKkkG4JIkSZIkSZKkkmQALkmSJEmSJEkqSQbgkiRJkiRJkqSSZAAuSZIkSZIk\nSSpJBuCSJEmSJEmSpJJkAC5JkiRJkiRJKkkG4JIkSZIkSZKkkmQALkmSJEmSJEkqSQbgkiRJkiRJ\nkqSSZAAuSZIkSZIkSSpJBuCSJEmSJEmSpJJkAC5JkiRJkiRJKkkG4JIkSZIkSZKkkmQALkmSJEmS\nJEkqSQbgkiRJkiRJkqSSZAAuSZIkSZIkSSpJBuCSJEmSJEmSpJJkAC5JkiRJkiRJKkkG4JIkSZIk\nSZKUJ5VVkfmLK5gxb3GxS2mXOha7AEmSJEmSJEkqRYPOeLBe28RLhhNCKEI17ZMzwCVJkiRJkiQp\nxx5859u07UPOeqjAlbRvBuCSJEmSJEmS1ExVVZEYY732Y//9Ztr+MSYzw+csLKeisirf5bV7LoEi\nSZIkSZIkSRksWFzJRQ99wC2vfAHAzsP60b1LRzp3KOPet75u9nXXH/UYAPcfuy0brtonF6UqDQNw\nSZIkSZIkSW1OjJEYoaxs6XraM+cv5r/jvuGwzVelS8cOtfovrqji1lc/52/PTuDhE7Zn2WU6p73u\ntLmLOO2ud3jyo8lpzz/98ZQm1/rqWbuy5cVPpj23/zUvMmnMiCZfU9kxAJckSZIkSZJUEPMXV/DD\nv7/KuC9nMuHi4bXC67oWllcye0E5K/TqytMfTebIG1/P+j7n3f9+g+c3vvDxJa9XW647n0+bn/W1\nm+rAjQewYq+uLN+jC1PnLkrb55KHPuTM4WvnrYb2zABckiRJkiRJUrNVz8R++uPJ/PKmNwCYcPFw\nHv/we47+19iM42puBrl6v2X4bMq8vNeaTj7Db4A/HrYRAG+csxuQrP9d19+em8Dpe63V4B8E1DwG\n4JIkSZIkSZIyqqqKvDB+Kj+7/jU2XKU3476a1eiYmuF2NooVfjfVoZutwp1vfLXkeIWeXXjpjF34\nYvp8dvnDswD8ZKuB3PLKFzx/2s6s2rd7vWs8f9rObH/Z0/Xah5z1EK+cuSv9e3fN3wfQDoV0O5Sq\n7QshjN1kk002GTs281/ZJEmSJEmS1LqN+3Im+1/zYq22fTZYidEHrMfsBRUMXK5+wNpSFZVVrHH2\nwzm/bq5tPLAPb30xs177IZuuwn/GflV/ALDcMp1ZuU833v06c4j/v99tx7D+PencsSxXpdbz03++\nyvOfTk17zvXA69t00015880334wxbtrUsQbgJcoAXJIkSZIkqfhijCwsr6JDWWDuogp6du1Ipw5J\nsPrkh9/TvXNH1hvQi5H3v889b32d93r2Xq8/D7/33ZLjrYcsx6eT5zB17uK83xvgqO0Hs+0ay3PE\nDdmv513t/fP3ZJkuuVvQIsZICIH/+8847hr7Fe+O2oOeXTvl7PoNqaqKGWfJD1+/P3/5cZNz3pJm\nAK56DMAlSZIkSVIhvTh+Kj/+x6vc8eut2HLIcs2+zsLySrp26lCrbdrcRSzbvTNlZYEYI1PmLuLb\nmQvZcNU+Lao53VrMap5d1lqBQzdbhZ2GrcBh173CuC9nAnDuPuvwy+0GNzr+va9nsaC8ks1WW5YQ\n2s862FPnLmKz0U/Ua083C7xmaP7RhXvV+zkpZS0JwF0DXJIkSZIktRpfTp/Pv1/7ghN2XZOunTpQ\nXlnFW1/MZIvBfYtdmtKIMTL4zNqzWA+77pUiVQMXH7g+A/t257z/vsfROwzh9LvfLVotpej+Y7dl\njRV60LVTBzo0sFnj/cdu2+Rrrzegd0tKa7OW79GFu4/ZmoOufblW+9czF7BCzy4cccNrvDh+Wr1x\na537iEulZMkAXJIkSZIkFcW3sxaw9SVPpT137TOfNTr+V9sNpkfXjlz5xKf8fOvVOH//9aiorGJ+\neSW9CrSMQV3XvzCRCx74oME+9x27LRu1cOZyMa1z3iPMX1xZ7DLSOuvepYF3qYTfWw9Zjlt/tSWV\nMdKxLDBh6jx2/cOz7LVufx55/7vGL9AC74zao2g/S+3JpqvV/wPftmPS/26s6euZCxjQp1s+Siop\nLoFSolwCRZIkSZJKx7xFFUyaNo91V247MyTTzQwutIdP2J7VlutOt04dCCFZOqMq0uDM1cZUVFYB\nMGdhBb26deLBd7/l+NveanGtEy4eTlmGumYvLOeQa1/m1D2H8aub32jwOhfsvy4/3nI1Xhw/lZ9d\n/1qL61LuDOm3DEdsM4jz7n+/Vvta/Xsyceo8Hjphe1ZZthtdOuZuWYvyyqol640DzJi3mF/e9DqX\nHrQBa67Yk+nzFnPC7W/xmx1XZ9s1ls/ZfdV0zV2Op73MAncNcNVjAC5JkiRJbd97X89in6tfSHtu\n+Pr9+fMPNyEE6q2XO3dRBeuNfDTtuH02WIkH3vm2ybU8fML2rL1SryXH0+YuYtMa69Yesc0genXt\nyFVPjW/ytYvl30dtyTarNxz6tYYgvy1589zd2fH3TzNnYUWxS8la3e9tqRgmTZ3HTpc/0/RxBuCN\nMgAvUQbgkiRJklqjGCMLyiuZPHsRyy7Tmd7dfGt9NTfjW+oHmwzgnje/LnYZaoIPL9iLbp1zuyFf\njDGrzRBnLSinS8cyOncoY/KcRXw3eyEbrtKb/477hqEr9jTcVpvR2H8Hdi0by4AwlWM73s+KYSaX\nlR/KaRf9vUDVFZebYEqSJEmSWr0rHvu4ybOD685s+2bmAsY8/BH/HfdNrfZ7frsNmwxctsU1QvrQ\n7btZC9nqkicBeOOc3Vi+RxcAKqsiq5+1dHZwv55deO2sXZcstzFvcSXPfTKFvz8/gbe+mLmkX5/u\nnZg5vzwn9bZ1p++1FkfvMKTe8h9XHLpRg+NijLw6cTpVVZHf/vvNVvX5fOLkHVhjhZ712quqIkPO\navuzybddYzn+8uNN8/4HrGzCb6BWHf17d6V/764A7L/RgLzUJeXLRQeux9n3vpf23KEdnuayTrXD\n7tM63cncWZfSo7ebBDfEGeAlyhngkiRJkoqtNc1ovuygDTh081V55uPJHHHD6wDsv9HKjD5gPc7/\n3wfcNfaremP69ezC62fv1qo+jraoGG/Pn7+4gu9mLeSkO95m3Fez8n6/l8/chf69umYd2NaU7bIH\nP9h4AOuv0puBfbuz69orpu2zsLySo25+g+c/nQrAj7YcyMUHrt/kmqD+2tGSmqmqEsqye3dERWUV\na5z9cK22jVbtw51Hb03n0en/yPvZDx5h9Q22bnGZrZ1LoKgeA3BJkiRJhfDvV7/grHvfXXJ85LaD\nuOHFScUrqAQ9f9rOrNq3e732va58jo++m9Pg2ImXDCeEwKKKSjqWlbHWuQ9TXhm58cjN2WnYCkv6\nlVdWUV5ZRffOHXn6o8kcf/tbvHXu7nTsUNbgGtjn77cuP99mEHMWlvP97EVMnDqP3dZeoVlBcKF9\nO2sBW1/yVNb9qz+X+baoopLOqeC5pRt2SiqS6RPgqo1rt43K7o9xsxeWs8Gox5YcT7xkOGHc7XDf\nb9IPyPK6bZ0BuOoxAJckSZKUa1VVkb8+9xmXPfJxsUspWRcfuD4/3GJVZi+o4K0vZ9QKqRvz5fT5\n9O/dlfmLK5mzsJxVlq0fmkuS8qh8IXxwP9z76/rndh0J25+89HjKJ3DLQTDri+T47O+gU7clpydO\nncdqfbtTVjEfLl45/f1+8RgM3DKHH0Dr5RrgkiRJkqS0Pp82jx1//wwADxy3HSv17sref3qeyXMW\n1ep38y+24GfXv5b3em7+xRbsMLRfrbbqiVnVs2vTvQW8pisP24hNV1t2yazo8+5/j5tf/jxPFWfn\nk9F707EsZFzfed8NV2bnYf0Yvv5KdO5QxtgvZnDIX1/mqh9uzH4b1g82enfv1KTwG1jy+ejdrczN\nRSWp0Eb1bvj8k+cnj81/Ba//o/75i/ons7mrquDLVxl8w15J+/LD0l9vk5+3m/C7pZwBXqKcAS5J\nkiS1TeWVVfzqpjd49pMp9c6dM2JtfrX9kIxjD/3ry7w2aXo+y2uyR0/cgWH9628GmI0YIzHC/901\njnve/JoHj9+OdVduJGBIY+b8xWx0weNLjs/cey2O3nH1JWt7H7/LGhy94+p079xhSQhfd9mPF07f\n2RnVkqT0Ggu/c+kH/4BeK8Nq20AbWG4qV1wCRfUYgEuSJEmtwysTpnH4da8AcNAmqzB8/f788qY3\nlpz/ZPTefDF9PgP6dGPt8x7J6po/2Wogt7zyRV7qba5JY0YwZ2E5j7z3Hduv2Y8Ve3VpE+tAS5LU\nIi9cCU+MLNz92sma33UZgKseA3BJkiSpuKbOXcRmo58odhl58epZu/LdrIUs16Mz8xZVNnuGtyRJ\nbV5Ds7+3OR5euip39xo5s13N+q7JNcAlSZIkqRV5afxUfvSPV4tdxhL/PmpLTr5jHN/NXrikbcQG\nK/HnH27MM59M4cgbXl/S/odDNuSgTVdp9Jor9uqal1olSWozPnksfftOZ8E6+8MKa8G422He5PT9\nTnwP/rEbzP0uu/u10/C7pZwBXqKcAS5JkiQV1v/GfcNxt72Vs+uNv2hv9rn6BT76bg5/+fEm/PbW\nN7Ma99GFe9G1U4ec1SFJkjKoO/v75I+g10r1+73yV/juHVhmeXjxT9B/AzjkRlhudYgRzu/T+L3a\n8exvcAa4JEmSJBVE3WVNPh69Fxc/+CE3vfx5g+NO2HVNdhi6PAdd+zIAmwzsw5tfzEzbd9KYEUte\nP3LiDrXa73vra0684+20fSVJUgG9/e/6benCb4CtfrP09e4X1D4XQrtd17tQDMAlSZIkqQGvTpjG\nYalNLOsadk7jm1Z+etHedOpQBrQ8sD5g4wEcsPGAFl1DkiS10Ov/gAdPqd328/8VpxY1ygBckiRJ\nUrvwpyc+5Y9PfFKvvW4ofeUTn3DlE5/m5J7O0JYkqQTVDb8BBu9Qv02tggG4JEmSpJL3xAffpw2/\nAQad8WDO7jOgTze+nrkAgLfP2z1n15UkSa3EpBfrt53e8FJoKi4DcEmSJEltWlVVZMhZDy05Pn7X\nNTlulzV46qPJHP2v3G8Kf/MvtmDeogqOqbMp5fiL9qZjaqkTSZJUom4cXvt453OgW5+ilKLsGIBL\nkiRJapMyzdy+6slPuerJ3CxhUtPES4YTQlhyfP+x27L/NS9y3U83ZY91++f8fpIkqQ3Y8dRiV6BG\nGIBLkiRJKroYI0CtgDnGyKKKKspCYOg5D+fsXhcfuD6r91sm48aWAJuutix3H7NNg9fZcNU+rvEt\nSVJ7ct3OtY9Per84dahJDMAlSZIkFU2MkcFnPlSrbWDf7nwxfX5e7vfJ6L3p3DFZpqQ6vK6sisxZ\nWE6HskDPrp3ycl9JklQCvqm9/Bm9VylOHWoSA3BJkiRJBTNz/mJ+8s9X+dPhGzOgTzfWOveRen1a\nEn6vsUIPxk+eW6utsVnaHcoCfbp3bvY9JUlSOzB/eu3j9Q8pTh1qMgNwSZIkSQVRc83uXf/wbE6u\nednBG7DHOivSpWMHunXukJNrSpIk1XPZ4NrHB/2jOHWoyQzAJUmSJOXc/MUVdO3YgSFnPdR45yw9\nf9rOrNq3e86uJ0mSlJV5U4tdgVrAAFySJElSk3w5fT7bX/Y0AFccuiH7bbgy8xZV0qtbR76euYDt\nLn26WdfdcnBfXp249O3F4y/am44dynJSsyRJUrOM6l2/bevfFb4ONVubC8BDCAcDOwIbARsCPYFb\nY4w/aWDMNsA5wFZAV2A8cD1wdYyxMsOYnwPHAusAlcBbwOUxxgcy9O8GnAEcDqwGzAaeAUbGGD/M\nMGYV4AJgL2A54FvgPuD8GOOMTB+PJEmSVAgLyyvTrtFd08l3juPkO8e16D5PnbIjQ/r1aNE1JEmS\ncu6P66dv3/OiwtahFmlzAThJkL0hMBf4Cliroc4hhP2Bu4GFwB3AdGBf4I/AtkC9FetDCJcDp6Su\n/3egM0mw/b8QwnExxj/X6d8FeDx1vTeAPwGrpq49IoSwS4zx1TpjVgdeAlYA7gc+ArYATgD2CiFs\nG2Oclt2nRJIkSWrc2M9n8MqEafz+0Y+LWsezp+7EasstU9QaJEmSGjR3Csz6on77sIY311br0xYD\n8JNIgunxJDPBM76/MoTQiyTArgR2ijG+kWo/F3gKODiEcHiM8fYaY7YhCb8/AzavnokdQvg9MBa4\nPITwQIxxUo1bnUwSft8FHBZjrEqNuYNkRvf1IYT1q9tT/kISfh8fY7y6xv2vSH2MFwG/adqnRpIk\nSVoqxshjH3zPGiv0yNmmk81x21FbsfXqyxXt/pIkSU2SbtmTaj/8d+HqUE60uQA8xrgk8A4hNNb9\nYKAfcHN1+J26xsIQwjnAk8AxwO01xlSHzhfVXIYkxjgphHANcC5wJDAyVUOoMea0miF3jPH+EMLz\nwPbUCOtDCEOAPYBJwDV1ah4J/Br4aQjhlBjjvMY+SEmSJLVflVWR8soqunbqwOKKKs65713ufOOr\notY0pN8yPHHSjpSVNfrvdUmSpNYlU/h93gwoc2+StqjNBeBNtEvqOd3Chc8B84FtQghdYoyLshjz\nMEkAvgupABxYHRgIfBJjnJhhzPapMdXhffU9HqszK5wY45wQwoskAflWJCG9JEmSRGVVZO7CCiZN\nm8ev//UG389e1PigHFq2eye2HLwcf/3ppgAsrqiiY1kw6JYkSW1fZTlcuHz6c7uONPxuw0o9AB+W\nev6k7okYY0UIYSKwLjAE+DCEsAwwAJgbY/w2zfU+TT0PzeYeLRyzR2pMgwF4CGFshlMNro0uSZKk\n1q2isoo1zn44b9fffs3lGbnvuvTr0YUpcxfSq2snTrzjbV76LNmG5qID12P63MX8bpc1Mr7zsnNH\n/0dQkiS1IbO+gucuTzaxXDgbvn8fbj2o4TGrbgnbn1yY+pQXpR6AV79nYVaG89XtfZrZv5BjJEmS\nVOJijKxz3qMsKK/M2z0mjam/cVPv7p0A+PdRW+XtvpIkSUW1aA78cd3k9dgbshuz7Qmw+wX5q0kF\nUeoBeGOqp7LEJo5rSv/m3CPrMTHGTdNeIJkZvkkT7ilJkqQiWbC4krXPS7cCX8t8cMGedO/c3v/J\nL0mS2r0Y4ZJVmjbmzK+gS8/81KOCKvV/DVfPpM60dWuvOv0a659u5nZT79HcMZIkSSpBMcZmh9/P\nnbozA5frnuOKJEmSSsyrf2ta/5//z/C7hJR6AP4xsBnJWtq11soOIXQEBgMVwASAGOO8EMLXwIAQ\nwkpp1gFfM/Vcc+3uj1PPQ0kvV2MkSZJUQgad8WDWfa/50SaM2GClPFYjSZJUgl64Ep4Y2bQxp3wC\nPVfMSzkqjlLfteap1PNeac7tAHQHXooxLspyzN51+gB8BnwBDA0hDM5yzNOp5z1CCLW+BiGEnsC2\nwALglTTXkyRJUhv26oRpjYbfP9lqIADnjFibSWNGGH5LkiQ1VYzZh99r7QPnTIFRswy/S1CpzwC/\nC7gUODyEcHWM8Q2AEEJXYHSqz7V1xvwV+ClwdgjhvhjjjNSYQcCxwCJgyUr5McYYQvgrcDFwWQjh\nsBhjVWrM/sD2wAfAszXGfBZCeAzYI3XNq2vc/3xgGeBvMcZ5Lf8USJIkqRBijIQQarVVVUWGnPUQ\nADf9Ygt+fv1rDV5j6Io9eOykHQEYfcD6+SlUkiSpVD18Orz614b7nPIJdO4O5QvhwZNgmX6w9++h\nQ6nHpO1Xm/vKhhAOAA5IHfZPPW8dQrgx9XpqjPH/AGKMs0MIR5EE4c+EEG4HpgP7AcNS7XfUvH6M\n8aUQwhXAycA7IYS7gM7AYUBf4LgY46Q6ZV0B7AMcDLwaQngSGAgcAswHflEditfwW+Al4KoQwq7A\nh8CWwM4kS5+c3aRPjCRJkoqivLKKNc9+uF77e+fvyXojH11y3Fj4fcruQzlu1zUb7CNJkqQMqiob\nD7+79Fo6w7tLTzjslvzXpaILMcZi19AkIYRRQEPvX/g8xjiozphtSQLlrYGuwHjgeuCqGGNlhvv8\nHPgdsA5QBbwJ/D7G+ECG/t2AM4AfkYTfs4FngJExxg8yjFkVuIBkuZXlgG+B+4DzY4zTG/gYGxVC\nGLvJJptsMnbs2MY7S5IkqVku+N8HXP/ixGaP32JQX+78zdY5rEiSJKkdqiyHC5dvuM8RD8Kg7QpT\nj3Ju00035c0333wzxrhpU8e2uQBc2TEAlyRJarnyyioWllcyfvJcNlylD2VlS5c42Wz0E0ydu6iB\n0ZltMagvN/9yC7p26pCrUiVJktqnxfPg4pUb7jNqVmFqUd60JABvc0ugSJIkSYWQaaPKSWNGNLqJ\nZSZbDu7LHUc741uSJKnFqqpg8RwYM7Dhfmd/V5h61GoZgEuSJEl1/OWZ8RnPZQq/HzhuOx59/zuu\nfqr22NP2GsZvd1ojp/VJkiS1ay9dDY+d03i/nc6ETt3yX49aNQNwSZIkCYgxMvjMh5o1dsLFwykr\nC6w3oDen7DGMyqpIZVWkc8eyHFcpSZLUTs35HpZZHi7o23C/kz+Cnv2TdcE7di5MbWrVDMAlSZLU\nbs2Yt5iNL3y8RdeYNGZEvbYOZYEONdYLlyRJUguM6p1dv5EzIaT+DWb4rRQDcEmSJLU7u13xLOMn\nz82q78RLhvPAO99y3G1v1TuXLvyWJElSjiyeDxevlF3fH/1nafgt1WAALkmSpHbj4oc+5LrnJmTd\n//ojNiOEwL4brsy+G66cx8okSZLaucXzoaxj7ZnbjYXfJ74HVeXQY0XovEx+61ObZQAuSZKkkjf2\n8+kcdO3LWfd3ZrckSVIBpJvhfcxLsOK6MPvbhsce/zb0WTVvpal0uCuPJEmSStZVT37KoDMebDT8\n/t/vtmP7NZcH4NOL9i5EaZIkSUo3w/vabWDuZLhirYbH9h2cn5pUcpwBLkmSpJJ0yyufc8Xjn2Q8\nv8Eqvbn/2G0JqbUi//XLLQtVmiRJkhpy+Zr1286dBuMfh5f+DD+6o/A1qc0yAJckSVJJiTGy5cVP\nMnnOorTnd1lrBa4/YvMCVyVJkqRaRvXOvu/P/wcdOsKwvZOH1AQG4JIkSWpzYoyMfvBD/vnCRB44\nbjv2ufqFrMZNvGT4khnfkiRJKoLF8+Dz7PdmAWDwDvmpRe2CAbgkSZLalHe+msl+f35xyXE24fe4\n8/agd/dO+SxLkiRJjVk8Hy5eOf25kz6AP65Tv33kzLyWpNLnJpiSJElqM2KMtcLvbJy/37qG35Ik\nScW0aC489/v0m14CnP099B4Ah/6rdvtOZ4Hv3lMLOQNckiRJbUKMkcFnPtSkMa+etSsr9uqap4ok\nSZKUlUsGNHy+U+rfa+vsB+dMgSkfQf/1Db+VEwbgkiRJavXuGvsV//efcQ32+cW2gzl1z2F069yh\nQFVJkiSpUVM+afj8qFm1jzt2hpU2yF89ancMwCVJktTq3D32K05pJPDeYJXe3H/strw6cTqrLded\nlXp3K1B1kiRJalRVJVzQt+E+Z39fmFrUrhmAS5IkqagmTJnLLn94FoCTdhvKH59oZJZQyn9/tx0A\nWw1ZLm+1SZIkqYkqFsPofg33Ofp5Z3mrYAzAJUmSVDSLKiqXhN9A1uH3pDEj8lWSJEmSmmtU7+z6\nGX6rgAzAJUmSVHAxRn510xs8+dHkJo99/KQd8lCRJEmSWmT+9Mb7DNgUfvD3/Nci1WAALkmSpIIb\nfOZDWfU7fa+12G3tFaiKMHTFHoQQ8lyZJEmSmmzmF3Dl+g332fRI2PfKgpQj1WQALkmSpIKIMWYd\nfG8xuC93Hr11niuSJElSi1Uszhx+j5wJTmBQkRmAS5IkKe9qbnSZiet6S5IktUE37FW/beezYcfT\nCl+LlIYBuCRJkvJq0BkPNtpn4iXDC1CJJEmScu7rsfXbDL/VihiAS5IkKW8aC79P3n0ox++6ZoGq\nkSRJUk4tmFG/7bSJha9DaoABuCRJkpqsorKKNc5+uFZb3SVMGgu/V+3bzfBbkiSprZr1Ffxx3dpt\nP7wDuvctTj1SBgbgkiRJJS7GSMjB5kNfTJvPDr9/OuP5QWc8yMRLhvPtrIVsM+aptH12GtaPG47Y\nnBihrMwNkSRJktqkUb3Ttw/ds7B1SFkwAJckSSphox/4gH+8sPRtqA8ctx3rDcjwPywZfDl9Pttf\nljn4rmnwmQ9lPFdzhngO8nhJkiQVQ6bwG/xHnlolA3BJkqQSFWOsFX4D7HP1C3x4wV5069yh0fHZ\nbF6ZrfEX7Z2za0mSJKlIMoXfPVeCUz4qbC1SlgzAJUmSSlSm2dhrn/cIEy4eXm8Jkq9mzGe7S7Ob\n6d0UddcGlyRJUhvz3O/hqdGZzxt+qxUzAJckSSpB212afg3uakPOemhJMP30x5M58obXs7523Rnk\nmWaK77/Ryvzp8I2zvq4kSZJamYdOhfFPwPQJmfucN71w9UjNYAAuSZJUQj74ZjbDr3o+q77/euVz\nzr3vvSZdf9zIPeotn/Lx6L3Y56oXuOjA9dlicN8mXU+SJEmt1K2HwKePZT4/albhapFawABckiSp\njZu1oJwNz2/gf05IliGpO1M72/D70RN3YFj/nhnPd+nYgcdP3jGra0mSJKmVmToeHj0Lhu4BoQN8\ncD9MaGRZPMNvtSEG4JIkSW3YP1+YyIUPfNBgn88uHg7AxEuGZ1wXvKbRB6zH8PVXou8ynXNSoyRJ\nklqhP20IMyYtPf700ezGnT6p0S5Sa2IALkmS1EZlWnu7pg8v2IsOqc0uQwi8de7ubHzh42n7Hrvz\n6pyw61A6dyzLaZ2SJElqZZ6/onb4nS1nfqsNMgCXJElqY977ehbn3d/w8iVjz9mN5Xp0qde+7DKd\nOWyzVbnjjS9rtVdviClJkqQSN+UTePL8po3Z6Cew18X5qUfKMwNwSZKkVqayKvL2lzPZYJXedOqw\ndDb2wvJK1jr3kQbHTrxkOCGEBvtcevAGXHrwBjmpVZIkSW3E/b+Dt/7VtDFr7gk/vjM/9UgFYgAu\nSZLUyqx+1tJ1uqsD7alzF7HZ6Ccyjnn5zF1YqXe3QpQnSZKktmb8k9mF3z1WhB/eBittBGUd8l6W\nVAgG4JIkSa3IP56fUOs4m00rP7t4+JJ1viVJkqRaPnoIbv9h5vOu660SZwAuSZJUIN/PXsj//Wcc\nz386FYAnTt6RNVbowfA/Pc8H385u9nUNvyVJkpRRQ+H3CusWrg6pSAzAJUmS8uzBd77l2H+/Wa99\ntyuebdb1DtpkFf5w6IYtLUuSJEmlbsonmc+VdYLfvlS4WqQiMQCXJEnKk+nzFrPJhY/n/LqG35Ik\nSWrUqN4Z2l3yRO1LWbELkCRJKjULyysZdMaDLQq/LzpwvSWv916vPwC9unZk0pgRLa5PkiRJJa6y\nIn274bfaIWeAS5IktVCMkaoIn06ew7czF3Lkja+36Hpvnrs7fZfpzI+3XC1HFUqSJKndWDgLxgys\n337kw4WvRWoFDMAlSZKa4flPp/DTf77W5HGfXrQ3HcsCp/xnHPe8+fWS9o9H70WXjh1yWaIkSZLa\nm0zLnux8Dqy2TWFrkVoJA3BJkqQm2vnyZ5g4dV6Txky4eDhlZWHJ8RWHbsQVh25EjJEQQgMjJUmS\npCxM+Th9+9r7wo6nFrYWqRUxAJckSWqClz+b1qTw+6ML96Jrp8wzuw2/JUmSlBPXbJG+/bBbCluH\n1MoYgEuSJDXBD//+StZ9P7yg4fBbkiRJarFpn8Enj9Rv3/EM2PnMwtcjtTIG4JIkSTn0399tywar\n9Cl2GZIkSWoP/rQhzJhUv33wDobfUooBuCRJUhZijAw+86Fabcftsgan7DGsSBVJkiSp3amqgonP\nwOCdoKwsffgNLnsi1WAALkmS1IiF5ZWsdW79t5UafkuSJKlgvn0H/rZ9dn279s5vLVIbYgAuSZLU\ngIOvfYk3Pp9R7DIkSZLUXlUsgg6dsw+/f/VUfuuR2hgDcEmSpBqqqiILyitZd+SjDfabNGZEgSqS\nJElSu3XzATDh6ez7H3w9rLJp3sqR2iIDcEmS1K7NnL+YyqpIz66dGHrOw1mNMfyWJElSXlQshjt/\nBp88DD1XhjnfZDfu3KnQoVN+a5PaKANwSZLULv133Dccf9tbTR5n+C1JkqS8Gd1v6etsw+9Rs/JT\ni1QiDMAlSVK7M2dhebPC77/+xLeTSpIkKccWzIRLV8u+/9nfwbOXwucvw0/uyltZUqkwAJckSe3O\nBf/7oEn9LzpwPfbZYGV6d/NtpZIkScqxbMLvYSNg+TVh15FQVga7jcp7WVKpMACXJEntyox5i/nP\n2K8a7LNy765c97PNWG9A7wJVJUmSpHYnRji/T3Z9f/jvvJYilbKyYhcgSZKUS1/PXMDC8spabbPm\nl/PFtPlMnr2QjS98vN6YrYb0XfL6iG0G8dKZuxp+S5IkKX+qKrMPv0fOzGclUslzBrgkSSoZg854\nsMlj/vqTTdlrvf55qEaSJEnK4IK+mc8dcC1s9KPC1SKVOANwSZJUEna/4tlmjTP8liRJUkFULILR\nKzTcZ809DL+lHDMAlyRJbV6MkU8nz23yuH/8bLM8VCNJkiSl0VD4/fP/weAdCleL1I4YgEuSpDbv\niQ8nN3nMZxcPp0NZyEM1kiRJUhOMmlXsCqSSZgAuSZLavKNufqNJ/SeNGZGnSiRJkqQ0PsqwV83Z\n3xe2DqkdMgCXJElt2ox5i+u1TRozgqqqyHG3v8UGA3pz9I6rs7C8kpnzy+nfu2sRqpQkSVJJiRFC\nE95NeHuadb1/+Th08t+mUr4ZgEuSpDZt4wsfr3X8n99sDUBZWeCaH22ypL1rpw70792hoLVJkiSp\nhLz+T3jw5KXHP7oThu7Z+Lg539Vvc9kTqWAMwCVJUpsVY6zXtvmgvkWoRJIkSSWtsqJ2+A3w70OX\nvt7lXFhtW+i2LHz2FGx2JMz6CvoOgT8Mqz3uxPfyX6+kJQzAJUlSmzX4zIdqHe+29opFqkSSJEkl\n7cLlGj7/1IW1jx89M3PfPqu2vB5JWTMAlyRJbc4j730H1J/9/Y+fb1b4YiRJklTaFs7O3bW2/7/c\nXUtSVgzAJUlSm7L5RU8wZc6iYpchSZKkUjR9Aly1cf6uv+u5+bu2pLQMwCVJUpuSKfx+7/wsNiCS\nJEmS6ooRvh0HE5+Dx/MYUJ8+KX/XlpSRAbgkSWozrn9hYsZzPbr4zxpJkiQ10bjb4d6js+9/0D9h\n/YPhvbvhwVPggL/Ciusm565cL/O4n96bbJApqeD8P0VJktTqzV9cwfA/Pc+kafPTnr/j11sVuCJJ\nkiSVhKaE35CE3wDrHZQ8aho5ExbPhe/ehc7LQP8NYOKz0Gcg9B2Sk3IlNZ0BuCRJatVmzS9nwwse\ny3h+0HLd2XLIcgWsSJIkSSVh1ldN6z9qVsPnQ4AuPWG1bZa2DdmpyWVJyi0DcEmS1CpVVUVG/vd9\n/vXK52nPP3/azqzat3uBq5IkSVKb8cpf4ZHToUNnOHdK7XN/2ghmZF5ejwP+CsP2gksHJce/fSVf\nVUrKMwNwSZLU6nw7awFbX/JUg30MvyVJkpTRvcfAuH8nrysXwxPnw7Tx8Onj0L0vzP46/bjzpkNZ\nh6XHjc36ltTqGYBLkqRWp7Hwe9KYEQWqRJIkSW1Sdfhd7YUrlr7OFH4bdkslyQBckiS1Kp9+Pyfj\nufEX7U3HDmUFrEaSJEntguG3VLL8P0hJkpRXMUb2//MLDDrjQb6aMb/Bvv965XN2/+Nzac9NGjPC\n8FuSJEmNG9U7+769VzX8lkqcM8AlSVJeDT7zoSWvt7v0aQD2XHdFfrfzmoz63/ucv9+6jJ88l3++\nMJF3v67/Px+fXTycDmWhYPVKkiSpjVo4C8YMzL5/555w0nv5q0dSq2AALkmSCu7R97/n0fe/B2Cf\nq1/I2G+nYf0MvyVJktS48gWZw++D/gnLLA+DdoAy31EotTcG4JIkKadijMQIq5/9EDG27Fo3HrlF\nboqSJElSabuof/p2lzeR2j0DcEmSlDOLK6oYes7DObnWxEuG5+Q6kiRJKnHP/b7YFUhqxXzfhyRJ\nyol5iypyFn5/dOFehODSJ5IkScrCU6PTtzv7WxLOAJckSTmy7shHGzz/7Kk7MWHKPI688XUADtpk\nFYb178HFD30EwLjz9oAAvbp2NPyWJElSZrO/hSvWynx+5Ezw35OSUgzAJUlSk82Yt5iNL3w86/6P\nnLg9qy23DKsttwyTxoyode6o7YcYeEuSJKlhC2ZAt2WT1w2F3+dMMfyWVIsBuCRJysqs+eVseMFj\nTRrzxMk7ssYKPRrsY/gtSZKkjCor4MLlsu/fsXP+apHUJhmAS5KkBsUYGXzmQ00a8/AJ2zNsxZ6U\nlRluS5IkqQWaEn6fNyN/dUhqs9wEU5IkZdSc8Btg7ZV6GX5LkiQps6nj4T9HwqyvoapyaXuMsHh+\n8rp8YfbXGzULyoy5JNXnDHBJklRPjJGz7n2X2177ssljXzxjlzxUJEmSpJIQI5zfZ+nx+/c0/1pu\ndikpCwbgkiSpnqbO+n721J1Ybbll8lSNJEmSSsZN++bmOmWdDL8lZcUAXJKkdm7+4gq6duxAWVmg\nsirSoYGlS+749VZsOaQJ6zBKkiRJNU16vnnjlukHq24JI66AnivmtiZJJc0AXJKkduzc+97jX698\nnlXfd0btQa+unfJckSRJkkpSjHDxgOaPP3V87mqR1K4YgEuS1E4NOuPBrPpd99NN2WPd/nmuRpIk\nSSXtopWgYkHttt1GwROjGh97yif5qEhSO2EALklSCZm7qIKPvp3N+qv0pkvHDgBUVFaxxtkPA7Du\nyr345883Z9q8RVlf0/BbkiRJLTLhmfrhN8B2JyUPScojA3BJkkrEe1/PYp+rX2iwz/vfzGarS57M\n+pofXbhXS8uSJElSKYkx2XxyVO/keMjO8LP7Gu5/8/712zf6SV7Kk6S6DMAlSSoRjYXfTXX8rmvS\ntVOHnF5TkiRJbcCiuXD/sfDBfbDlMbDXJfD6P+Ch/6vfd8LTUL4AOnVLf625k+u3HXYrrL1PTkuW\npEwMwCVJasMWLK5klz88wzor9Wr2NfZatz+PvP8dNxyxOTsN60cIIYcVSpIkqc344lW4fo/aba9e\nmzwaclF/GDUr/bk/DK3fZvgtqYAMwCVJasPWPu8RAL6dtbDZ1/jrTzfNVTmSJElqqyZ/VD/8bqrF\n82H+VCDAvw+Dnc6o3ydTUC5JeWIALklSG1ReWcW65z3a5HGvnbUrW1y8dA3wX243OJdlSZIkqS2K\nEf6yZcuucX5fiJW12+78ae3jvS9r2T0kqRkMwCVJaoPWPPvhBs9fedhGnHjH2wBcdvAGrN2/F+uv\nkmxUNGnMCBZXVAHQuWNZXuuUJElSK1KxCD5/ETr3hFU2g/P75O7adcPvdLY8Onf3k6QsGYBLktTG\nHP2vNxo8/9IZu7Byn24csPGAjH0MviVJktqZa7aEKR9l1/eAa6HvEFh+KFxW4x2DG/0E3r4leX3y\nh3DF2tnfv//62feVpBwyAJckqY1YsLhyyZrfmTx0/Pas3KdbgSqSJElSmzCqd9P6b/Sjpa9HzoS7\njoTNfgGDd4ADrmleDb95oXnjJKmFDMAlSWoDyiurMobfR20/mLNHrFPgiiRJklRUVZVQWQ6dukL5\nAvjsKRi6N5Sl3unX1NC72siZtY9DgENuTN+37xCYPqHxa7rxpaQiMgCXJKkNyLTm9xMn78AaK/Qs\ncDWSJEkqqobC7dM/h0tXa/o1z50KHTo1bcyB18E/d6tTm2G3pNbFBUAlSWrlFlWk31CoLGD4LUmS\n1N5881bD5xsLv1fbtvbxqFnJo6nhN8Cqm8NeY5LXh91q+C2pVXIGuCRJrdywc+ovffLRhXvRtVOH\nIlQjSZKkgqssT4LvlTaC63Zq/nXOm5E8X7M5TBsPZ3/X8tq2OiZ5SFIrZQAuSVIrtvpZD9VrmzRm\nRBEqkSRJUlHM+Bz+tEHLr3PCuKXrgx83tuXXk6Q2wgBckqRWKMbI4DPrh9+SJElqJyorYMxAKJ+X\nm+stOyg315GkNsYAXJKkIiuvrFqyyeVLZ+zCyn26ZQy/nzt150KWJkmSpGL4+y7wdQtnaZ8zBWIV\nTP0Y+udgBrkktVEG4JIkFVl1+A2wzZinMvbbYnBfBi7XvRAlSZIkqVheuTa78LvmhpMxwgXLQayE\nwTvCT+5euqnlShvmp05JaiMMwCVJKqIjbngtq36/P3gDDtls1TxXI0mSpKJ75IzM50b8AdY/FLr2\nqt0eAoycngThIeS3PklqYwzAJUkqki+nz+eZj6c02m/iJcMJ/o+MJElSaamqgoqF0Lk7vHQ1PHZO\nw/1rzvjOxH8zSlI9BuCSJBXB3WO/4pT/jMuqr+G3JElSiVk0By5ZJbu+254Iu5+f13IkqZQZgEuS\nVGAxxqzD788uHp7naiRJklRw2Ybf2cz6liQ1yABckqQCG3zmQ/XaLtx/XQiB9VbuxYF/eYlbf7Ul\n266xfBGqkyRJUpONfzLZfLJDFjHLqN7ZXdPwW5JywgBckqQC+XrmAh5659t67VsM7stPtx605HjS\nmBEFrEqSJEktUjPQPndawyH4gpmNX+/gG2DY3i0uS5KUMACXJCnPFldU8dwnU/jVzW+kPX/n0VsX\nuCJJkiTlxOv/qH184XKZZ25P/gj+smXmax37GvQblrvaJEmAAbgkSXk19JyHWVxRlfH8uJF7FLAa\nSZIk5USMcH6f9OceOwf2GF277dLBsGB6/b7LDoYZE2GfPxp+S1KeGIBLkpRHDYXfAL27dSpQJZIk\nScqZTOE3wEtXw6pbwXKrwwprw0OnpQ+/D70Z1tk/byVKkhIG4JIk5cndY79q8PzES4YXqBJJkiQV\n1B0/bryP4bckFURZsQuQJKlUnfKfcRnPbTywDyGEAlYjSZKknLhk1ZaNH7p35nXCJUk55wxwSZJa\nKMbIP1+YyG2vfcFjJ+1Ih7LA5Y9+nLbv6XutxTE7rV7gCiVJktQiMULFInj1r7Bodu1zZ38PnbrC\nqN6NX+ewW2DtffNToyQpLQNwSZJa6Lz73+dfr3wOwOpnPcSR2w7ihhcn1epz6p7DOHLbQXTv7H96\nJUmS2oQY4Y6fwEcPNNyvU9fk+RePwvV7NtzX8FuSCs7/C5ckqYWqw+9qdcNvgN/utLpLnkiSJLVm\nMUII8Ng5yUaW2ai5lMnAreDE9+DK9dL3PW1iy2uUJDWZAbgkSS3w6PvfNdpnxPorGX5LkiS1Ztks\nX1LXkQ/Xb+uzahKK37QfTHx2afs6+0P3vs2vT5LUbG6CKUlSCxz9r7GN9rnmx5sUoBJJkiQ1y6WD\nmj7ml4/DattkPv+jO2sfH3pz0+8hScoJZ4BLktRMVzyWfqPLalsO7ssdR29doGokSZLUqIpFEMqg\nQ6fk+Ju3YcGM7MefNjG7mdydusLJH8Lr/4Rtj29WqZKk3DAAlySpGb6cPp+rnhpfq21An26M3Hcd\ndllrBTp28E1WkiRJrcK8afD7IbXbfv1MEnz/68DsrnHUU7DyJska4dnqtTLsem72/SVJeWEALklS\nEz32/nf8Os3SJy+esUsRqpEkSVJGE56Bm/ev337dTo2P3eY4ePs2OG4sdOuT48IkSYViAC5JUhP8\n+B+v8OL4afXaV+zVpQjVSJIkCUiWNhm9QvL6lE+g8zJw/Z7w/XvZX+Nn/4UhO8KiOdC5RzLbe4/R\n+alXklQwBuCSJDXgmFvG8vB733HYZqtyxxtfZuz36lm7FbAqSZIk1VIdfgP8YWjzrjFkx+S5S8+W\n1yNJajVcoFSSpAyOv+0tHn7vO4BGwu9dC1WSJEmS8uGsb4tdgSQpTwzAJUlKY+6iCv477psG+2y4\nSm9u/dWWrNira4GqkiRJUj2zvmq8z8BtYNQsOPPr+udGzYLO3XNflySpVXAJFEmS6rj55Umcd//7\nDfZ58pQdWb1fjwJVJEmSpLS+Ggv/aGQj8h1Ph53PSl536QEjZ8J/j4PxT8IJb+e7QklSkRmAS5JU\nwy9ufJ2nPprcYJ99N1zZ8FuSJKnYyhc0HH5vc1z6TSxDgP3/nL+6JEmtSrtZAiWEMCKE8FgI4asQ\nwoIQwoQQwn9CCFtn6L9NCOGhEML0EML8EMI7IYQTQwgdGrjHz0MIr4UQ5oYQZoUQngkh7NNA/24h\nhPNDCB+HEBaGECaHEO4MIaydi49ZktQ0j3/wfaPhN8BlB21QgGokSZJUy80HwB/WhpmpvVku6t9w\n/3ThtySp3WkXAXgI4VLgAWAT4BHgT8CbwP7AiyGEn9Tpvz/wHLADcC9wDdAZ+CNwe4Z7XA7cCKwE\n/B24BVgf+F8I4Xdp+ncBHgfOA2ananoCOBB4I4SwZUs+ZklS0x118xtp2987f89ax906Z/xbqCRJ\nkvLh7dtgwtMw5xu4cj148sL0/c6ZAqdNTNb1liSJdrAESgihP/B/wPfABjHGyTXO7Qw8BVxAElgT\nQuhFEmBXAjvFGN9ItZ+b6ntwCOHwGOPtNa6zDXAK8BmweYxxRqr998BY4PIQwgMxxkk1SjsZ2Ba4\nCzgsxliVGnMHcB9wfQhh/ep2SVJ+jZ88J2372iv1okeXjkwaM6LAFUmSJAmAb96G+35Tu+35y+v3\nO28GlJVBx74FKUuS1Da0hxngq5F8nK/WDL8BYoxPA3OAfjWaD04d314dfqf6LgTOSR0eU+ce1f8l\nvqg6/E6NmUQye7wLcGR1ewgh1BhzWs2QO8Z4P/A8sA6wY1M+UElS9sorq7jkoQ95/IPv+XL6fHa7\n4rl6fV4/ezcePmH7IlQnSZKkJa7L4n+NR81Kwm9Jkuoo+RngwKfAYmCLEMLyMcap1SdCCDsAPUlm\nXFer3kHjkTTXeg6YD2wTQugSY1yUxZiHgXNTfUam2lYHBgKfxBgnZhizfWrM0w1+dJKkJjv5jre5\n562vAfjbcxPS9vnT4RvRr2eXQpYlSZKkuhbObrzPAdfmvw5JUptV8gF4jHF6COF04ArggxDCfcA0\nkhB6P5J1uI+uMWRY6vmTNNeqCCFMBNYFhgAfhhCWAQYAc2OM36Yp4dPU89Bs7tHAmLRCCGMznFqr\nsbGS1B6N/Xz6kvA7k6dO2ZEh/XoUqCJJkiSlNekFuDGLZeg2+lH+a5EktVklH4ADxBivDCFMAq4H\njqpxajxwY52lUXqnnjPtmFHd3qeZ/Zs7RpLUQnMXVXDQtS832s/wW5IkqRVIF35vfwq8exfM/Bz6\nDoHj3yp8XZKkNqVdBOAhhNOAi4GrgD8D35HMkL4EuDWEsFGM8bRsL5d6jk0soyn9s75HjHHTtBdI\nZoZv0oR7SlLJW2/ko432mXjJ8AJUIkmSpAbNn16/bftTYNfzkockSVkq+QA8hLATcClwb4zx5Bqn\n3gwhHEiyDMkpIYS/xhgnsHT2dW/S65V6nlXnOVP/dLO9m3oPSVILvfnFjEb7fHrR3iT7FEuSJKng\nYoSJz8I9R8Pc7+qfN/iWJDVDyQfgwD6p53qbScYY54cQXgMOBDYGJgAfA5uRrL9da33tEEJHYDBQ\nkepLjHFeCOFrYEAIYaU064CvmXquud73x6nnTGt8pxsjSWqipz+ezJE3vJ7xfL+eXZgyJ9nPeNKY\nLNaXlCRJUn5ULIbR/TKfH+X8MElS87SHALxL6jnTf0mr2xennp8CfgzsBdxWp+8OQHfguRjjohrt\nTwE/TY25oc6YvWv0qfYZ8AUwNIQwOMY4MYsxkqQmiDE2GH7/6fCN2H+jAQWsSJIkSRk1FH4fXPd/\nsyVJyl5ZsQsogOdTz78OIdRKOkIIewPbAguBl1LNdwFTgcNDCJvV6NsVGJ06vLbOPf6aej47hLBs\njTGDgGOBRdQIxmOMscaYy0IIZTXG7A9sD3wAPNuUD1SStNTQcx5u8LzhtyRJUitRvrDh8+v9oDB1\nSJJKUnuYAX4X8ASwG/BhCOFekk0w1yZZHiUAZ8QYpwHEGGeHEI5KjXsmhHA7MB3YDxiWar+j5g1i\njC+FEK4ATgbeCSHcBXQGDgP6AsfFGCfVqeuK1P0PBl4NITwJDAQOAeYDv4gxVuXyEyFJ7cU3MxdQ\nXpl5H2GXO5EkSWolpn0GV2+S+fw5kwtXiySpJJV8AB5jrAohDCeZiX04yXrf3UlC7YeAq2KMj9UZ\nc18IYUfgbOAgoCswniTgvio1g7vufU4JIbwD/A74NVAFvAn8Psb4QJr+i0IIuwFnAD8CTgJmA/cB\nI2OMH+Tgw5ekdmmbMa4gJUmS1OpVVaUPv396LwzaHjp0KnxNkqSSU/IBOECMsRy4MvXIdsyLwPAm\n3ucm4KYm9F8AjEw9JEkF8PZ5uxe7BEmSJAFcsGz69tV3KWwdkqSS1i4CcElS+/Hl9Pn12iaNGcGC\nxZV069yhCBVJkiSpnlG907efN6OwdUiSSp4BuCSpzfnk+zns8cfnlhw/euIODOvfE4DtL3u6Vt9L\nfrA+gOG3JElSazG6f/r2jX4MZWWFrUWSVPL8L4skqU0Z+/mMWuE3wJ5XJseDzniwXv8fbjGwIHVJ\nkiQpSxUL6reNmgUH/KXwtUiSSp4BuCSpTTno2pfStv/j+Qn12gb27Z7vciRJktSQGGH2t8mSJ+/c\nCeUL6/c55ePC1yVJajdcAkWS1GYsLK/MeG70gx/Wa3vutJ3zWY4kSZIaMnU8/HnTpcf3HAUcVbvP\nsa9DzwxLokiSlAPOAJcktRlrnftI1n1fPGOXPFYiSZKkRtUMvzPpNzT/dUiS2jUDcElSydlhaD8G\n9OlW7DIkSZLar1f+2nifrX+X/zokSe2eAbgkqU24962v6rWNO2+PtH3/8uNN8l2OJEmS0pk7OVnv\n+5HTG++750X5r0eS1O4ZgEuS2oST7hhX6/je325D7+6dmDRmBA+fsP2S9hV6dqFHF7e4kCRJKorL\n18yu31FP57cOSZJSTAgkSa3ejHmL67VtPHDZJa/XXqkXk8aMKGRJkiRJqqui/r/ZADjgWtjgMLig\n79K2Ab5jT5JUGAbgkqRWb+MLH691POYH6xepEkmSJGV02ZD6bb8bC8uvkbweNQtihBAKW5ckqV1z\nCRRJUqv26fdz6rUdvsXAIlQiSZIkAJ7/Q7LO9+cvLW2b/S0srvPvtlGzlobf1Qy/JUkF5gxwSVKr\nFWNk9z8+V+wyJEmSVG3SC/DkBcnrG/ZOQu5Rvev3221UQcuSJCkTA3BJUqs1c355vbYXTt+5CJVI\nkiQpbdD99Zvp+25zfH5rkSQpSwbgkqRW5+0vZ3Lo315mcUVVvXOrLNu9CBVJkiS1czGmb/97hskJ\nZR3yV4skSU1gAC5JanUOuObFtO2PnbRDgSuRJEkSAOf3yb7vyJn5qkKSpCZzE0xJUpswoE83hq7Y\ns9hlSJIktS8xwmt/z67vjmck4bcbXUqSWhFngEuSWpU7X/8ybfuLZ+xS4EokSZLasY8fhg/uh3G3\nZT9m5zPzV48kSc1kAC5JalVOu/udYpcgSZLUPk3+CP6yZeP9Rs2CcbfDvUcnx+dMho5d8lubJEnN\nZAAuSWo1bn3187Ttlx28QYErkSRJamcqFmUXflfb8PDkIUlSK2cALklqFd7/ZhZn3/terbaHjt+e\ndVbuVaSKJEmSStwbN8ADJzZtzHkz8lKKJEn5YgAuSWoVRlz1Qr02w29JkqQcq6yAC5dr3tjVd4Gy\nstzWI0lSnhmAS5IkSZLUHpQvhItWbNqYA66FjX6UBOcdjBAkSW2P//WSJBXdlDmL6rWNPWe3IlQi\nSZJUQhbOhjGrNn3cyR9BjxVrz/Y2/JYktVH+F0ySVHSbX/REreMnTt6R5Xp0KVI1kiRJJaK54Xev\nlXJfiyRJRWIALklqddZYoUexS5AkSWrbKsuz7/ubF2Hyh7DBIfmrR5KkIjEAlyQV1VVPflrsEiRJ\nkkrPP3Zt+PxZ38KUD2GljaCsA/RfryBlSZJUaAbgkqSiWVRRyRWPf1KrbeIlw4tUjSRJUht3729g\n3G3pzx30T3jnDtj7Uug7JGkbsGnhapMkqUgMwCVJeRFj5F+vfM6lD3/E3b/dhjX69aBjh2QjpVH/\nfZ8bX5qUdlwIoYBVSpIktUHlC2HhTOjZPzme/Q1csXbm/mvtA+sfnDwkSWpnDMAlSXnx6Pvfc979\n7wOw15XPZzfmxB3yWZIkSVLbFiOc36fp4w6/NeelSJLUVpQVuwBJUmn6zS1jmzxmWP+eeahEkiSp\nRDQn/B60fc7LkCSpLXEGuCQp58orq4pdgiRJko59DfoNK3YVkiQVlTPAJUk5MXXuIqqqIk988D1r\nnv1wk8f/5zdb56EqSZKkEvGnjbLr17Fb8vzD2w2/JUnCGeCSpBZ696tZ7PvnF5o87rWzduWc+97j\nsQ++B2DzQX1zXZokSVLb9vlLcMPe2fX9xWMwcMv81iNJUhtkAC5JapFswu/d11mRx1NBN8ADx23H\nCr26ct3PNstnaZIkSW3Xy3+BR8/MfH7ULFg4Cx45C7Y9AfoNLVxtkiS1IQbgkqS8OmKbQYzab12+\nmbmAX//rDU7dcy3WG9C72GVJkiS1bg2F3799NXnu2hsOuKYw9UiS1EYZgEuSsvKP5ycw+sEPlxxf\ndvAGrLFCjwbH3PSLLdhxaD8AVu7TjQeO2z6vNUqSJLV5VZUw59uG+6ywVmFqkSSpBBiAS5KyUjP8\nBjjtrnca7H/3Mduw6WrL5rMkSZKk0vHxw3Db4Y33Gzkz76VIklRKDMAlSY2aPm9xk/o/f9rOrNq3\ne56qkSRJKkHZhN87nQkh5L8WSZJKiAG4JKlRe//puaz6TRozIs+VSJIklZjyhXDRipnP//x/MHiH\nwtUjSVKJKSt2AZKk1u/72Ysa7fPiGbsUoBJJkqQS01D4DYbfkiS1kAG4JKlB389emFW/AX265bkS\nSZKkEvPuXQ2fHzWrMHVIklTCXAJFktSgLS9+stbxz7ZejaO2H8I3Mxcwc0E55ZVV7LPBykWqTpIk\nqQ27+5f1245/G/oOLngpkiSVKgNwSVJGiyoq67Wdt886dOxQ5iaXkiRJLVFVVb/t9EnQbdmClyJJ\nUilzCRRJUkbDznmkXlvHDv6nQ5IkqUUqK+CCOkH3KR8bfkuSlAemGJKktNKt/X3W8LWKUIkkSVIJ\nmTcVLlyufnvP/oWvRZKkdiAnS6CEEDoAXWKM8+u07wLsD8wHrosxTszF/SRJ+bfHH5+r1/brHVYv\nQiWSJEltUPkC+PxFWGO35PjbcfC3HYpbkyRJ7VCu1gC/HDgmhLBijHEWQAjhcOBWIKT6/CqEsEmM\n8csc3VOSlEezFpTXOp40ZkSRKpEkSWpDXrkWHjmjdtsB18J9x2Qe83+f5rcmSZLasVwtgbID8HR1\n+J0yEpgJ/Aw4DegDnJyj+0mSJEmS1PrUDb+h4fD7zK+gxwr5q0eSpHYuVwH4qsD46oMQwhBgGHB1\njPGWGOPlwMPAXjm6nyQpj8orq2od/2ZHlz6RJElq1C0HNa3/qFnQpWd+apEkSUDulkDpBcyucbwt\nEIFHarS9D+yco/tJkvJozbMfrnV84m5rFqkSSZKkNmT8E9n1O/5t6Ds4r6VIkqRErgLwb4Ga//Xe\nDVgAjK3R1gOoyNH9JEkF1LVTh2KXIEmS1LpVVWbXb9SsxvtIkqScydUSKK8A+4UQ9gkh7AYcDDwV\nY6y5g9oQ4Osc3U+SJEmSpNahshwu6Fu77djX6/db/9DC1CNJkpbI1Qzwi4H9gftTx1XARdUnQwi9\ngJ2A23N0P0lSnrwxaXqt4xuO2LxIlUiSJLURFy5fv63f0GS299Tx0LEzLLMCdOpa+NokSWrnchKA\nxxjfDSFsCfw81XRHjLHmn7s3AB4DbsvF/SRJuVFRWcXdb37F6Xe/y74brswR2wzi4L++XKvPzmut\nUKTqJEmS2oDX/1m/beezl75efo3C1SJJkurJ1QxwYozvAv+X4dwLwAu5upckKTdufGkSox/8EID/\njfuG/437psgVSZIktXKfPg6fvwjbHA/d+8KDJ9c+v9WxsONpxalNkiTVk5M1wEMIT4UQftZIn5+E\nEJ7Kxf0kSblRHX5ncuqewwpUiSRJUhvw6nVw68Hwwh/hssEQY/0+e11c+LokSVJGuZoBvhPwTCN9\nVgN2zNH9JEktNHnOwkb7HLuzb9mVJEkCYFTv+m3n96l9vOclBSlFkiRlLyczwLPUDago4P0kSQ3Y\n4qInGzw//qK9C1SJJElSKzd/euN9ALb+bX7rkCRJTZazNcCBNO/9ghBCAAYCw4Evc3g/SVKePHj8\ndnTsUMi/kUqSJLViz/2+2BVIkqRmanYAHkKoonboPSqEMKqhIYCLoUlSKzBp6ryM5x44bjvWXTnN\nW3wlSZLao3RLn6Rz9nf5rUOSJDVLS2aAP8fSAHwH4AtgUpp+lcA04EngHy24nyQpR3a6/Jlaxy+c\nvjOrLNu9OMVIkiS1Vnf8NH37qFnJBpiL58J378GATaBjl8LWJkmSstLsADzGuFP169Rs8BtijBfk\noihJUmEZfkuSJNUxfzp8+N/67b9+JnkOAbr0hNW2LmhZkiSpaXK1BvhgYGaOriVJyqOKyqpilyBJ\nktT6XTY4ffvKGxe2DkmS1CI52eEsxvh5jHFWLq4lScqfisoq1jj74Vpt1/54kyJVI0mS1ErNn16/\nbaezYOTMgpciSZJaplkzwEMI55Gs/31NjHF66jgbMcZ4YXPuKUlquT2vfK5e217r9S9CJZIkSa1Y\n3dnf258CO51enFokSVKLNHcJlFEkAfgdwPTUcTYiYAAuSUXy2ZR59dpCCEWoRJIkqZWKsX7brtnO\n+ZIkSa1NcwPwnVPPX9Q5liS1If/93bbFLkGSJKllKhZDx87128sXwEWpd7r98A4YtlfyOkZYPBfe\nvw8+fhh+cB106bF03Pl9al9n9wvyUbUkSSqQZgXgMcZnGzqWJLUNG6zSp9glSJIkNU/FYhjdb+nx\nGrvDT+6CRXOSxxVrLz1322HJ88kfwRVr1b7OJQNg1CyoqoLy+u+WY9sTcl+7JEkqmObOAJcktRFf\nz1zAtmOeqtc++oD1ilCNJElSjtQMvwHGPw6jejc8pm74Xe1vO8C343JTlyRJalXKcnGREMKgEMLw\nEMIyNdo6hhDODyGMCyG8FEI4MBf3kiRlb3FFVdrwG+DHWw4scDWSJEk5MOUTeOXa3F4zU/h97rTc\n3keSJBVcrmaAjwT2A1as0XYOcG6N4ztDCNvHGF/J0T0lSY0Yes7DGc+5+aUkSWpVJjwL9xwFh90K\nj58LW/wa1vtB7Vndyw+DqR8XrqYOvmlakqS2Llf/Nd8aeDLGWAEQQigDfgt8BOwB9AeeAE4CDsvR\nPSVJzbTOSr2KXYIkSdJS86fDzfslr/+5W/L8xcsw+5va/Zoafq+4PvQdDB/+t+k1/e6Npo+RJEmt\nTk6WQCGZ+f15jeONgOWBa2KMX8UY3wDuBzbP0f0kSY2IMWY8d/vRWxWwEkmSpEZcNjh9+2NnNz52\n9wvTt+92PhzzAhx8Q/rz250E256Y/tyoWbD8mo3fW5IktXq5mgHeCaiZtGybOq658OxXwEo5up8k\nqRFT5iyqddyxLDBo+WU4e8Ta9OraqUhVSZIk5djmv4Jtj4epn0LXPtCjzuaYHTrCJj+HN29Kjvus\nBie+s/T8bqPgkTMgRtjzYpc9kSSpxOTqv+xfARvUOB4OTI0xflijbQVgdo7uJ0lqxBYXP1nrePzF\nw4tUiSRJUgPeubNl4zt3T54bmrG931XJI50QYO9LW1aDJElqtXIVgD8AnBRCuBxYCOwO1H2f2VrU\nXiZFkpQn738zq9glSJIkZeeeo5o/9syvcleHJEkqSbkKwC8DDgBOTh1/DYysPhlCWA3YBvhjju4n\nSWrAiKteKHYJkiRJjXv/3vpto1J/yB/Ve2nbMS/DCmvDwlnQNdUeQv7rkyRJbV5OAvAY4+QQwvrA\nrqmmZ2OMc2p06UESjj+ai/tJkjIbdMaD9dr+/KONi1CJJElSA+ZOhv8cUbvtZ/cvfT1qFsz5Hrr3\nhQ6p/Uu69SlUdZIkqUTkbHePGOMCkqVQ0p17H3g/V/eSJDXNPhusXOwSJEmSlpryCVyzef32ITvV\nPu65YkHKkSRJpaus2AVIknKnqirWa5t4iZtfSpKkVqSqMn34veMZha9FkiSVvGbNAA8hXA9E4KwY\n4/ep42zEGOMvm3NPSVLjhpz1UK3jUfuuQ3B9TEmS1Frcdyy8fUv6czufWdhaJElSu9DcJVCOIAnA\nLwW+Tx1nIwIG4JJUID/fZlCxS5AkSUq8fVvm8Pvs7wtbiyRJajeaG4APTj1/XedYklQkU+Ysqtfm\n7G9JktQqzJ0C9/0m/bnTP4dOXQtbjyRJajeaFYDHGD8PIWwQY/y8+ji3ZUmSmurPT31a6/is4WsV\nqRJJkqQaxj8Jt/wg/bkzvoCuvQtbjyRJaleaOwMc4O0QwmvAdcDtMcb5OapJktQMN71c+2+Rv95h\n9SJVIkmS2r3Z38D0CXDjiMx9tjzG8FuSJOVdSwLwT4EtgM2BP4YQ/g38Pcb4Zk4qkyRJkiS1PROf\nh5v2yXz+oH/C+gcXrh5JktSulTV3YIxxGLAT8G+gE3A08HoI4fUQwlEhhB65KVGS1Jg5C8uLXYIk\nSRI8f0XD4TcYfkuSpIJqdgAOEGN8Lsb4U2Al4HjgXWBT4K/ANyGE60IIm7e8TElSQ9Yf9Vit4ysO\n3bBIlUiSpHZr9rfw5PkN9xk1qzC1SJIkpbQoAK8WY5wVY/xzjHEjkmVR/glE4FfAKyGEt0IIx4QQ\neuXifpKkhh2w0YBilyBJktqTic/BFY1swH3qhMLUIkmSVENOAvCaYoxvxBh/TTIr/CjgNWBD4M/A\n17m+nySpvrKyUOwSJElSezH5I7hp3/rtfVeH1XdJXm/5G1hmucLWJUmSRMs2wWxQjHE+8M8QwjvA\n5cD2QPd83U+S2qvbXvui1vGGq/YpTiGSJKl9mfUVTPsMbt6v/rn9/wIb/7jwNUmSJNWRlwA8hNAb\n+AnJDPD1gQDMBW7Px/0kqT078553ax2fO2LtIlUiSZJK3lOj4bnfN97P8FuSJLUSOQ3AQwg7kKz7\nfRDQlST4Hgv8Hfh3jHFuLu8nSe3d+Mn1f61uMnDZIlQiSZJK3v2/g7f+1Xi/Q7PoI0mSVCAtDsBD\nCP2AI4BfAmuShN5zgOuA62KMb7X0HpKk9Ha74tl6ba7/LUmScu61v2cXfo+alf9aJEmSmqDZAXgI\nYU+SJU72TV0nkGx4eR1we2oNcElSnsxZWF6vbeIlw4tQiSRJKnkP/V/jffpvkP86JEmSmqglM8Af\nTj3PAm4hme39bgP9JUktNG9RBcfd9hZPfTSZoSv2qHc+BGd/S5KkHLs9y/W8f/N8fuuQJElqhpYE\n4K8AfwPuiDEuzFE9kqQGrDvy0SWvP/m+9vrfk8aMKHQ5kiSp1E37DD56oHbbGV/C/Gmw7KDkeNaX\n0HvVgpcmSZKUjWYH4DHGbXJZiCSpYcff5pYKkiSpwK7epH5b117Jo1qfgYWrR5IkqYnKil2AJCk7\n/x33TcZzHdz4UpIk5dr379dvW2WLwtchSZLUAi1ZAkWSVCB/furTBs8/efKOBapEkiSVtFlfwx/X\nSX9u6F7wozsKW48kSVILOQNcklq55z6ZwuWPfdJgn0HLL1OgaiRJUknLFH6D4bckSWqTnAEuSa3Y\nwde+xBufz6jXvsqy3fhqxgJO3XMYx+68RhEqkyRJJWX8E3DLQZnPH/t64WqRJEnKIQNwSWqlZi0o\nTxt+A7xw+i4FrkaSJJWsF/8Ej5/XcJ9+QwtTiyRJUo4ZgEtSK7Xh+Y8VuwRJklTqqiozh9/Hvg7L\nrQFlrpwpSZLaLgNwSWpjHjhuu2KXIEmSSsHcyXD5munPnTfD4FuSJJWEZgXgIYTrm3m/GGP8ZTPH\nSlK79+Dx27Huyr2LXYYkSWrrxt0B9/46/bkeKxp+S5KkktHcGeBHNHNcBAzAJakRsxaUp203/JYk\nSS1WWZ45/B45E0IoaDmSJEn51NwAfHBOq5Ak1bLbFc/WOn7xjF0Y0KdbkaqRJElt3vSJcNVGDfc5\n/XPDb0mSVHKaFYDHGD/PdSGS1F5NnbuI397yJr/cfjB7rtufL6fPZ8qcRbX6GH5LkqSslC+Ai/on\nr/utDb94GC4d1Pg4Z35LkqQS5SaYklRkm41+AoDXJk2nU4dAeWUsckWSJKlNinFp+A0w5cPswu8z\nvzL8liRJJSvnAXgIoQOwPNAl3fkY4xe5vqcktTVzFpaz/qjH6rWnC783XMV1vyVJUiNihPP7NH3c\nqFk5L0WSJKk1yVkAHkJYHxgD7EyG8JtkE0xnnUtq99KF35nc/7vt8liJJEkqCc0Jv8+bkfMyJEmS\nWpuchNEhhLWAl1KHjwP7AuOA74FNSGaEPw04+1tSu7eoojLrvq+etWseK5EkSSXh3buy7+uMb0mS\n1M6U5eg65wKdgG1ijPun2u6NMe4FDAZuANYBzsvR/SSpzbru2QlZ9Ttpt6Gs2KtrnquRJElt3t2/\nrN9Wc3Z3p2XgzK8NvyVJUruUq+VIdgIeiDG+W6MtAMQY54UQjgbeAS4EjsjRPSWpzVlYXskfHv8k\nq77H77pGnquRJEltXkyzefYZX0BZmYG3JEkSuQvAlwc+rXFcAXSvPogxVoQQngYOzNH9JKnNmTR1\nHjtd/kzG8xMvGU5lVWTG/HKW79GZEELhipMkSW3T3b+qfXzaROjqBtqSJEnVcrUEynSgR43jqcDA\nOn0WA/5LTFLJW1RRyU0vTWL2wnIAYmpmVkPhN0AIgY4dyujXs4vhtyRJatz0ifBenfW/u/ctTi2S\nJEmtVK5mgH8GDKpxPBbYPYSwQoxxcghhGWB/YGKO7idJrdawcx4BYOR/32+07/VHbMb0eeXss8FK\n+S5LkiSVmqs2KnYFkiRJrV6uAvDHgNNCCMvEGOcBfwVGAG+FEF4CNgVWA07J0f0kqVUadMaDTeq/\ny1or5qkSSZLUalVWwIXLLT1e7yD4wd+hrEP9vlWV8N27sOJ60KFjsub39+/DX7et3/eU7PYZkSRJ\nak9yFYD/HfgY6AbMizE+GEI4ERgFHATMBy4FrsrR/SSpVZk2dxGbjn4i6/6HbLoKYw7aII8VSZKk\nVqtm+A3w3t3Jo+6mlfOmwe+HZHfNbY6Hnv5hXZIkqa6cBOAxxm+BO+q0XRVCuIZkg8zJMabbnlyS\nSkNTwm+AS36wPh3KXOdbkqR2Z853mc+98x/Y4JDk9UUrQ/m87K+7x4Utq0uSJKlE5WQTzBDCwBBC\nr7rtMcbKGOP3McYYQugZQqi7MaYktXlVVU3/+17HDrnag1iSJLUZ370HfxiW+fw9v0qWPHnu900L\nv0+d0PLaJEmSSlSuEpiJwAmN9DkeN8GUVIJ2/sMzTeo/8ZLh+SlEkiS1XjO/TL9ud10X9IWnRmd/\n3bO+hWWWa7yfJElSO5WrNcBD6iFJ7c7n0+bXOl65d1e+mbWwVtu4kXvQu1unQpYlSZJai/nT4cr1\n6ref/CE8ciZ8cF/TrrfTmbDTGTkpTZIkqdTlKgDPxopAE97HJ0lt00tn7srcRRWcdc+7HLzpKuww\ntF+xS5IkSfm0cBaMSa32OGRn6DUA9v9zcvzJI3Db4fXHrLAu9FoZDr0JRvXOfO29L4Mtfg3B+UaS\nJEnN0ewAPITwszpNG6VpA+gADAR+Crzb3PtJUmv04Dvfpm3v0aUjV/1w4wJXI0mSCq584dLwG2DC\n08nz27c0PO6YF5e+3vMSePTM+n3KOsGWR7e8RkmSpHasJTPAbwSqd36LwP6pR13VUxXmA+e34H6S\n1Ooc++83ax2fM2LtIlUiSZKK4s50c4AaMXJm7RndW/8WVtkM/rl77X7nTmlRaZIkSWpZAH5k6jkA\n1wP3Afen6VcJTANejjHObMH9JKlVmbeool7bL7cbXIRKJElS0Xz6aNP6j5qVvn3VLeC8GXDBssnx\nSR+47IkkSVIONDsAjzHeVP06hPBz4L4Y4805qUqS2oBJ0+pvaxD8H1VJktqPr8c2rf8xLzV8vqws\nc0AuSZKkZinLxUVijDsbfktqb0Y/8GGt45t+sUWRKpEkSQX33t3w911qt620Yf1+u18AoQxOeAdW\nXLcwtUmSJGmJnATgktQevTxhWq3jHYf2K1IlkiSpoL4eC3f9on770c/BlsfUbtv2BBg5A5ZdrTC1\nSZIkqZaWrAFeSwhhR+BUYAtgWdKH6zHGmLN7SlKxLK6oKnYJkiSpEL4dB8sPhU7dkuNRvdP3G355\n8rz3GKiqgHG3wanjC1OjJEmSMspJGB1CGEGyCWYH4AvgY6D+7nCS1IY88t53/OaWsZyy+1CO23VN\nAGKMhBAYes7Dtfou36NzMUqUJEnNsXB2ElJ375v+fFUVvHotPHrW0rZfPgH/3C3zNTc9cunrEZcn\nD0mSJBVdrmZjjwLKgRExxsdydM2cCyFsD5wIbAP0BaYD7wJXxhgfqtN3G+AcYCugKzAeuB64OsZY\nmeH6PweOBdYBKoG3gMtjjA9k6N8NOAM4HFgNmA08A4yMMX6YboykwvnNLcnGVn94/BMefPdbPvpu\nTsa+L5y+S8ZzkiSpFbnlYBj/ePJ6wGZw1JPJ6xmT4Lnfw1u3pB/XUPgN0ME3ukqSJLVGufpX2nrA\n7a08/D4HuBCYCjwAfAssD2wM7AQ8VKPv/sDdwELgDpKgfF/gj8C2wCFprn85cArwFfB3oDNJsP2/\nEMJxMcY/1+nfBXg8db03gD8Bq6auPSKEsEuM8dXcfPSSmuqrGfNrHTcUfgN07dQhn+VIkqSW+vwl\nuGHv2m1fvwFfvgb/3L351z3on7D2fi2rTZIkSXmTqwB8LklI3CqFEA4hCb+fAH4QY5xT53ynGq97\nkQTYlcBOMcY3Uu3nAk8BB4cQDo8x3l5jzDYk4fdnwOYxxhmp9t8DY4HLQwgPxBgn1bjtySTh913A\nYTHGqtSYO0iWk7k+hLB+dbukwrrisU+y7tuhLOSxEkmS1CIVi2HMQKhYkP58c8PvUbOaX5MkSZIK\nJt1Glc3xJLB1jq6VUyGEMuBSYD7wo7rhN0CMsbzG4cFAP5IZ7W/U6LOQZEkUgDpbu/Ob1PNF1eF3\naswk4BqgC7BkUcAQQqgx5rSaIXeM8X7geZJlVHbM+gOVlDMTpszlnre+zrr/hxfslcdqJElSk82d\nDJetDi9fA6P7ZQ6/m+v4t3N7PUmSJOVNrmaAnw68llpm5KIYY8zRdXNhG2AwyUzrGakNO9cjWd7k\ntRjjy3X6Vy/k+0iaaz1HEqRvE0LoEmNclMWYh4FzU31GptpWBwYCn8QYJ2YYs31qzNMNf3iScmH3\nK57l08lzmzW2c8dc/S1RkiQ125RP4JrNoe/qMP2zpK3mJpbNsetI2OZ4+P49uG5H6NgNTngbevZv\ncbmSJEkqjFwF4COB94HzgV+EEN4GZqbpF2OMv8zRPbO1eer5e+BNYP2aJ0MIzwEHxxinpJqGpZ7r\nrX8QY6wIIUwE1gWGAB+GEJYBBgBzY4zfprn/p6nnoTXaMt6jgTFphRDGZji1VmNjpfZucUUVv7jx\ndWYuWNzs8HvDVXrnuCpJktQs16T+2V8dfmeywjqwyc/gkTMy9zn8Nlhr+NLjlTdyyRNJkqQ2KlcB\n+BE1Xg9KPdKJQKED8BVSz78BJgK7Aa8CqwF/APYE/kOyESZAdZqV6V+41e19mtm/uWMk5ciVT3zC\ny59N49WJTd+64IML9mTG/HKOvfVNDthoZQ7ZbNU8VChJkppkzvfZ9z3mJQgBBm4F1+20tH3Q9vDj\nu6BT15yXJ0mSpOLJVQA+OEfXyYcOqedAMtN7XOr4/RDCgSSzsHcMIWydZjmUdKp3u2vqMi9N6Z/1\nPWKMm6a9QDIzfJMm3FNqF87/3/vc8OKkZo/v3rkj3Tt35L5jt81dUZIkqWX+0OgbJ5NZ3/tdvfR4\n5Y1h6F7wySOw+wWw7Qn5q0+SJElFk5MAPMb4eS6ukyfVm1JOqBF+AxBjXBBCeJRkVvoWwMssnX2d\naV2DXqnnWXWeM/VPN9u7qfeQlCNNCb9/se1gzhy+Fqf+ZxxbDVmOwzZ3trckSa3OO//Jrt++V9Vv\n+9Edua1FkiRJrU6uZoC3Zh+nnmdmOF8dkHer0X8zkvW3a62vHULoSDLbvQKYABBjnBdC+BoYEEJY\nKc064Gumnmuu911dU6apKunGSGqhQWc8mHXf187alRV6JW+BvvLwjfNVkiRJaomZX8A9v6rdNnAb\n+OKl2m09V06WPZEkSVK7k9MAPISwL/BjYG1gmRjjGqn2tYF9gVtjjF/n8p5ZeI4ksF4zhNA5xri4\nzvn1Us+TUs9PkXwMewG31em7A9AdeC7GuKhG+1PAT1NjbqgzZu8afap9BnwBDA0hDI4xTsxijKQW\n+HL6/Eb7HLHNIEbtt24BqpEkSS126yHw6WP123/xcOFrkSRJUqtVlouLhMRNwH3AIcDq1F4XfAZw\nMfCTXNyvKWKMU4E7SJYbOa/muRDC7iSbYM4CHkk13wVMBQ4PIWxWo29XYHTq8No6t/lr6vnsEMKy\nNcYMAo4FFlEjGI8xxhpjLgshlNUYsz+wPfAB8GzTPlpJmez75xcaPH/1Dzdm5L7rFKgaSZLUIl+N\nTR9+/9/4wtciSZKkVi1XM8B/SzID+nrgFOAk4NzqkzHG70IILwIjgEtzdM+mOBnYkiSg3gF4DVgN\nOBCoBI6KMc5M1To7hHAUSRD+TAjhdmA6sB8wLNVea7HAGONLIYQrUvd5J4RwF9AZOAzoCxwXY5xU\np6YrgH2Ag4FXQwhPAgNJ/oAwH/hFjLEql58EqT2bOb8847n3z9+TZbq0hxWhJElq4z56CG7/Yfpz\nP7kHevQrbD2SJElq9XKV+PwSGEcSJMcQQkzT51OS2dYFF2OcHELYEjiHJPTeCpgDPAhcEmN8pU7/\n+0IIOwJnAwcBXYHxJAH3VakZ3HXvcUoI4R3gd8CvgSrgTeD3McYH0vRfFELYDTgD+BHJHw1mk8yi\nHxlj/CAXH7uk9HYY2o+bf7FFscuQJElNkSn8Blhj18LVIUmSpDYjVwH4MOBv6YLhGiYDRZuSEWOc\nThJgn5xl/xeB4U28x03ATU3ovwAYmXpIKiDDb0mS2phRvTOfO2964eqQJElSm5KTNcBJNpns2kif\nAcDcHN1PkrKyYHEl97z5Va22W365ZZGqkSRJzTI1w9reW/waRs2Csg6FrUeSJEltRq5mgH8A7BRC\nCOlmgac2kNwFeCtH95OkRsUYWfu8R+q1bzG4bxGqkSRJzfbnTeu3nf0ddOpW+FokSZLUpuRqBvi/\ngLWAP4YQal0zhNCBZMPHlYEbc3Q/SWrUvn9+IW175465+tUnSZLybtGc+m3nTTf8liRJUlZyNQP8\nb8B+wPHAISQbTBJCuItkw8mVgftjjLfm6H6S1KApcxbx3tezi12GJElqqUtWqX282S9d8kSSJElZ\ny8k0yBhjJbAPcAHQGRgKBOAHQHfgQpJgXJIKYvOLnih2CZIkqaXSzf7e54rC1yFJkqQ2K1czwIkx\nVgCjQgjnkwTgywGzgI9SAbkkFcT5/3s/47kfbjGwgJVIkqQW+dsOtY83PbI4dUiSJKnNylkAXi21\nCebHub6uJGXrhhcnZTx3xt5rFa4QSZK0VIww5WPo0An6rAYdGvlfkRmTYPqE2m37Xpmv6iRJklSi\nch6AS1Jrc+qew7jj9S+5+MD16d2tU7HLkSSp/SlfCBetWLtt1Kz6/WKEuZPhD0Prnxu6d35qkyRJ\nUknLWQAeQlgFOAnYCFgFSJcyxRjj6rm6pyRl49id1+DYndcodhmSJLVPFYvrh98AH9wP6+yfvH7n\nP3DPrxq+zuH/zn1tkiRJKnk5CcBDCDsBDwFdgQrg+9Rzva65uJ8kZevaH29S7BIkSWrfRvdL337n\nz5p2nbKyltciSZKkdidXM8AvAzoAPwP+HWOsytF1JalJ/jvum1rHe6zbv0iVSJKknPnl48WuQJIk\nSW1UrgLw9YHbYoy35Oh6ktQsx9/2Vq3jDmW+8USSpKJ57vKWX+PEd6HPwJZfR5IkSe1SrgLwGcD0\nHF1Lkprl65kLil2CJEkCmD4Rrtqofvspn6Tf4DKd82a47IkkSZJaLFcB+APAjjm6liQ1y7Zjnip2\nCZIktV+3HgqfPgr91oYpH6bv03NFGDULFs6GMavWPz9qFsQIwXdwSZIkKTdyFYCfBbwSQrgG+H/2\n7jtOqur+//jrLFVAQEDFgiJ2k6iIvWIv2MtXfzGWxFiiJhJNjJ3FnmiMPYkxsaVoEhUiYhcMdgVL\nVCxUUbGBrEgT2PP7Y5YyzGy/M3d35vV8PPYxc8q99z2LI8tnz5x7boxxTkLnlaQGmbMg9767vbq0\nTyGJJEllaMTZmeI31F78Puy2Zc87ds0UuwGqPoKPXoXNDsm0LX5LkiQpQYkUwGOMX4YQ9gNeAo4P\nIbwPVOWfGvdM4pqStLwDbhyT0/fqRXunkESSpDJTXQ2v/rn+eZsdnL+/29qZL0mSJKkAEimAhxC+\nA4wCVqnp6l/L1JjE9SRpRVNnzM1qv3HJPiklkSSpzFy6Sv1zANqtVNgckiRJUh5J3VXmOqAncAmw\nLtAuxliR56tNQteTpDp169Qu7QiSJJW+WM/6lj7bQa+NYcisosSRJEmSVpTUHuA7AA/EGC9P6HyS\n1GBfzfk27QiSJJWnod1z+zY7FHY5B9bYvNhpJEmSpBxJFcC/BaYkdC5JapT+lz2R1Z545QEpJZEk\nqYx88npuX2W+2wBJkiRJ6UmqAD4a2Dahc0lSs7SpCGlHkCSpNI35LTx1adopJEmSpAZLqgB+LvBS\nCOE84Ncx1rcZoCQlo7ra/91IklQUld3qHr9kZnFySJIkSY2QVAH8IuAt4Arg5BDC60C+zz/GGONJ\nCV1Tkrj7hSlpR5AkqfTNmlb3+NnvQoX3u5ckSVLLk1QB/MTlnq9X85VPBCyAS2qWGd8s4Mnxn7HT\nBr2ofOidrLEd1++ZUipJkkrIvFlQ0RY6dMm0r/9u3fO7rlHwSJIkSVJTJFUAr63gLUmJG3D5k7WO\n/e3H2xUxiSRJJeSbz+HaDbP7dvwZPH9j3cf9amrhMkmSJEnNlEgBPMboT72SiqK+WwyE4A0wJUlq\nkhWL35C/+N2uM1z4SeHzSJIkSQmoSDuAJDVG1byFtY7ddtyAIiaRJKmENOYe9qc/X7gckiRJUsIS\nLYCHEA4KIdwbQngjhDBhuf5NQwjnhhDWSvJ6ksrPNY+9V+vYPt/pXcQkkiSVkP/8tOFzV+lbsBiS\nJElS0hLZAiVk9hy4E/hBTdc8YKXlpnwFXAkE4NdJXFNSeamau5Cz7nuN0e99kXYUSZJKz2v3NGze\nRf49LEmSpNYlqRXgpwPHAXcAPYBrlx+MMX4KPAcMSuh6ksrMFpc+bvFbkqRCuH2v3L7TnsvtGzIL\n2rYveBxJkiQpSUkVwE8C3gBOjjFWAfk2EfwAWC+h60kSAGfsvj47rt+TD67YP+0okiS1PosWwEev\nZPdd9Dn0/i78clKmvdbWUFkF3mhakiRJrVAiW6AAGwN/jLHOu+d8Dqya0PUklZHq6vz/a/lf5T6s\n3LFdkdNIklQiFi+Cy1fL7W/bIfPYuWem8C1JkiS1YkmtAF8EdKxnzlrANwldT1IZ6XfByLz9Fr8l\nSWqiqo/gsp65/UNmFT2KJEmSVEhJFcDfAQbW3AwzRwihI7AH8FpC15NUJn7z6Lt5+9+4ZJ8iJ5Ek\nqYT87jv5+93mRJIkSSUmqQL4PcAmwO9CCFnnDCG0Aa4D1gTuTOh6ksrA9Kp53Dp6Yt6xbp1c/S1J\nUqNNGg2V3fKPXfJVUaNIkiRJxZDUHuB/BA4GfgYcBcwGCCH8G9ieTPF7eIzxbwldT1IZ2OGqp/P2\n3/+THYqcRJKkElBdDXcfkn/szLFQkdTaGEmSJKnlSOSn3BjjYuBA4FKgPbAREIDDgU7AZWQK45LU\nILeOnlDr2IB1exQxiSRJJeLSVfL3V1ZBrw2Km0WSJEkqkqRWgBNjXARUhhCGkimA9wSqgHdrCuSS\n1GC/efS9tCNIktT61bbdyRLBVd+SJEkqbYkUwEMIi4H7YozfjzFGwMqVpCb78psFefsP3mJNrjlq\n8yKnkSSpFbpyLfj2m/rnnfBQ4bNIkiRJKUpqBfhsYGpC55JU5ra+/MmcvvGX7sdK7dukkEaSpBbs\nus3g64+h3+5w2B+gU0+4rFfDj++7c+GySZIkSS1AUgXw14DNEjqXJGUJAYvfkiSt6LN3MsVvgEmj\n4LcbN+y4vrvAiSMKl0uSJElqQZLa9O/XwAEhhL0TOp+kMrVocXVO3+SrBqWQRJKkFu73OzT+mN6b\nW/yWJElSWUlqBfhqwKPAIyGEYcArwKdAXHFijPHuhK4pqQT9a+xHaUeQJKnl+3Zuw+aFNrD8/ehP\nG1OYPJIkSVILlVQB/E4yxe4AHF7zBdkF8FDTtgAuCYAPZ8zlrhemsEO/nuy12eqMnfoV5z/wv6w5\nVx/+vZTSSZKUks/Hw63bw56XwM5nw5Qx8PUn8NYDcMBvYJW+cOUa9Z/n4hnQpi3Mr4I27aHdSgWP\nLkmSJLU0SRXAf5jQeSSVkV2vGQXAn5+dzI3/rz8/+8drOXOO3qZPsWNJkpSexYsyxW+Apy7NfC3v\nhsfyH7fLOTDmt5nn+10N2/9k2VjHbsnnlCRJklqJRArgMca7kjiPpPJRNXdhVjtf8RsghFCMOJIk\ntQyX9Wz8MRvslVktvuclyeeRJEmSWrmkboIpSY2yxaWP1zvnrD03LEISSZJagOlvQmUTV2of++9k\ns0iSJEklJJECeAhhQAjhkhDC6rWM964Z3zKJ60lq3aqrc+6Pm+N7a3Xj53tvVIQ0kiSl7IlL4I+7\nNO3YIbPAT0tJkiRJtUpqBfg5wI+Bz2sZ/ww4CTg7oetJasU+njWv3jnDz9ipCEkkSUrZhy/Cczc0\n7dgOXS1+S5IkSfVI6iaYOwCjYox5l3XGGGMI4Wlg14SuJ6kVe+jNT+oc3+87vamo8B/0kqQSFSMM\n7V73nCUruz95Hdq0g1U3gYo28MlrcNvAzJxfTS1sTkmSJKkEJFUA7w18VM+cT4A1ErqepFasfZu6\nP3zyi33d+kSSVGK+nQPPXg/9fwA3bF733KP/umxl95pbZo+t2R8qqwqRUJIkSSpJSRXA5wKr1jNn\nVWBBQteT1Ipd/vD4OsfXXqVTkZJIklQkV66Zefzvb+qeZ3FbkiRJSlRSe4C/DhwSQuiSbzCE0BU4\npGaepDL2/MQv6xy/8IBN6diuTZHSSJJUYDFCZbeGzd3n8sJmkSRJkspQUivAbwP+ATwRQjg1xvjm\nkoEQwhbAH4FeNfMklbHv/+mlnL6JVx7Aq1Nmsv5qXejVpUMKqSRJKpD69voG+P4/Ya2toXPPgseR\nJEmSyk0iBfAY430hhP2B44HXQgifAR8DawGrAwG4K8b4jySuJ6l1+nr+wpy+3l070qYisF0//9Ev\nSSoxixfVPf7DR2DdHYuTRZIkSSpTSa0AJ8Z4YgjheeCnwHfI3BgT4C3gxhjj7UldS1LrtHnl4zl9\nz5+3RwpJJEkqoM/fhYd+BtNyP/UEwC7nwJ6XFDeTJEmSVKYSK4ADxBhvA24LIXQCugOzYoxzk7yG\npNbpmwW5q+Du/OE2VFSEFNJIklQgk0bD3YfUPj5kFgT/7pMkSZKKJdEC+BI1RW8L35KW+s/rn+T0\nDdx4tRSSSJJUIJPH1F38rqwqXhZJkiRJQIEK4JK0otvHTEo7giRJhVPZre7xH+VuAyZJkiSp8CqS\nOlEIYbcQwogQwuchhIUhhMV5vuq5E5CkUjXpyzlZ7Z036JVSEkmSEvb+Y/XPWWe7wueQJEmSlCOR\nFeAhhEHAMKAN8CHwHmCxW1Kt7jlp27QjSJKUjL//X+1jh/4etvx+8bJIkiRJypLUFiiVwEJgUIzR\nz3dKyvLhjNxbAgRvACZJKgUx5u8/dzJ06lHcLJIkSZJyJFUA/y5wr8VvSfnses2orPaZu2+QUhJJ\nkhI2tHt2+9j7Ya2tLH5LkiRJLURSBfBvgJkJnUtSiTtnn43SjiBJUvPlu/HlhnsVP4ckSZKkWiV1\nE8yngB0SOpekEjK9al5On9ufSJJavTf/mdu35lbFzyFJkiSpTkkVwH8FrB9CuChY2ZJUo7o6ssNV\nT6cdQ5KkZH00Fh44Obf/ZP/OkyRJklqapLZAGQK8DQwFfhRCeB2YlWdejDGelNA1JbVwP7zzlZy+\nD67YP4UkkiQl6PY9cvvOehNcByJJkiS1OEkVwE9c7nnfmq98ImABXCoDp97zKs+8/0VOf7s2SX3w\nRJKkApo7E8b8FqoXwUt/WNa/yYG5cy+YDu07FS+bJEmSpAZLqgC+XkLnkVQCZs75lsfe/iynv9tK\n7VJII0lSI301FW7YPP/YuyOy22tsafFbkiRJasESKYDHGKcmcR5JpWGry57I2z/mV7sXOYkkSY20\nYHbtxe98Tn2mcFkkSZIkNVtSK8AlCYDzH/hf3v6XLtiTrh1dAS5JasEu7w2L5jV8/rmTC5dFkiRJ\nUiIsgEtKzPyFi/nHyx/m9G+8+sqs3rVjCokkSWqgym6NP6ZTj+RzSJIkSUpUk+9GF0JY3ISvRUmG\nl9SybHLxo3n7H/v5rkVOIklSI8ybVff4tqdAZRUc/ddlfbtfWNBIkiRJkpLRnBXgoUjHSGoFFi6u\nztv/8702KnISSZIa6dfr5u+/4BOoaAdt22famx6UKYRLkiRJajWaXACPMTZ59bik0rPhhY/k7T95\n1/WKnESSpAT8dBy075x2CkmSJEnN5B7gkpqtttXf7dtW0Km9/5uRJLVgM1e4keV2P4H9r04niyRJ\nkqTEuYpbUrPEGPOu/t5941V599L9UkgkSVIj3LhldnvfK1OJIUmSJKkwLIBLapb/fZx/L9Q/n7AN\nFRVu+y9JasFizO2r8MdjSZIkqZT4E76kZhk39au8/Ra/JUkt3pRns9v7ufWJJEmSVGosgEtqlvmL\ncvf/vvzQ76aQRJKkRqheDHcdmN23/U/SySJJkiSpYLw7naRGW7BoMef88w122bAXL06akTP+g+3X\nTSGVJEmNcGmPtBNIkiRJKgIL4JIabeOLHgVgxJvTc8ZevmDPYseRJKlxFi3I7Tt+ePFzSJIkSSo4\nt0CRlKjVunZMO4IkSXW7fLXcvn4Dix5DkiRJUuElugI8hNAO2BPYFOgSY7yspr8j0BX4MsaYu2Gw\npFbjnU++TjuCJEkN9+u+MK/mhs3tOsHCublzhswqZiJJkiRJRZTYCvAQwn7AFOBh4LdA5XLDWwLT\ngaOTup6kdBxz2wu1jo36xcDiBZEkqT7Vi5cVvyF/8RsghOLkkSRJklR0iRTAQwhbA8OACPwc+Pvy\n4zHGF4HJwGFJXE9SejZfu3utY+v16ly8IJIkrShGmDkZFi+EebMadqPLi3Nv5ixJkiSpdCS1BcrF\nwFxg6xjjpyGEIXnmvAJsldD1JKWk/zrdeXbClzn9Fw3aNIU0kiTViBGGdm/8cW28J7wkSZJUypLa\nAmUnYFiM8dM65kwD1kjoepJSMuz1j/P2/3iXfkVOIknScppS/Hbvb0mSJKnkJbXkpQuQuyQ0WycS\n3HNcUjqmzZyX1b73lO3Zpm8DPmIuSVLSpr8Jf9ylYXPP/whe+yusvyesulFhc0mSJElqMZIqgH8M\nfKeeOVsCkxK6nqQWYvt+PdOOIEkqN6OvhtFXNXx+ZVXmcfufFCaPJEmSpBYrqRXZjwD7hhB2zjcY\nQtgf2BEYkdD1JKXgjWmzstrt2/ihDklSChpT/L7ws8LlkCRJktTiJbUC/CrgGODxEMJNQF+AEMIg\nYFfgDGA6cF1C15NURDFGrhw5nj+NmZzV/8TZu6aUSJJUlqqr4dJV6p6z489gs0NgjS2gTbvi5JIk\nSZLUYiVSAI8xfhxC2Af4J/DL5Yb+AwRgInB4jLG+fcIltUDrnT8yb/86PToVOYkkqazVV/wG2Oey\nwueQJEmS1GoktQKcGOO4EMLGwCBgB6AnUAW8CAyPMS5K6lqSiqfveQ/XOhZCKGISSVJZe3Jo7WPH\n/B02GVS8LJIkSZJajcQK4AAxxsVkVn3/J8nzSiquqrkL2eLSx+ucM/oXA4sTRpKkV/8Cz9ayk95m\nh1r8liRJklSrRArgIYSfAPfGGL9K4nyS0lVf8Rug20ruqypJKqBX74ARg2sfP+d9WHn1osWRJEmS\n1DoltQL8FuC6EMJDwF3AIzHG6oTOLakFsgAuSSqIuTPhN+vVPefcydCpR3HySJIkSWrVKhI6zwXA\nZOBIMtuffBJCuDaEsHlC55fUgjz8s52pqHD/b0lSwl64tf7iN1j8liRJktRgiRTAY4xXxxg3A7YF\nfg+0Ac4GXgshjAsh/CyEsGoS15JUWF/MXlDr2GWHfpeJVx7Ad9bsVsREkqSSV70Yrv8ePHZ+3fO6\nrg2VVcXJJEmSJKkkJLUCHIAY46sxxjOBNcmsBh8BfAe4HvgohDAsyetJSt4Jf3m51rHjtl+XNq78\nliQl5fPxUNkNLu0Bsz6se+7x/4Gz3y5OLkmSJEklI6k9wLPEGBcCDwAPhBB6Aj8BLgYOKsT1JCXn\nnelf5+3vv0734gaRJJW+W7eve/z0F2G1TYuTRZIkSVJJKkgBHCCEEIC9gROAQ4B2wOJCXU9SYd1/\n2o5pR5AklZJhZ9Q9/vO3odvaxckiSZIkqWQlXgAPIWxKpuj9A2ANIAAfAHfXfElqob5ZsChv/xM/\n39WbXkqSklNZx70k+h8Hh9xcvCySJEmSSloiBfAQQg/g/5EpfA8gU/T+GvgzcGeM8fkkriOpsL47\n5LGcvilXD0ohiSSpZFR9BK//Hbb/CXRYue7iN8CBvytOLkmSJEllIakV4NNrzhWBJ4E7gQdjjPMT\nOr+kFPz5hK3TjiBJao2eugzGXJvdN+oKGHRd/vk/fCRzzFF3QJt2hc8nSZIkqWwkVQCfTKbofXeM\n8ZOEzimpiL5dVJ3Tt8cmq6WQRJLUqv15H5j2Uv6xh8/O7Tv0D7DujvCjRwqbS5IkSVJZSqQAHmPc\nJInzSErPRhdlFx66d2pH5l62kiQ10D/+X+3F73wu+gLati9cHkmSJEllryLtAJJapscH75p2BElS\na/PeyIbPPXeyxW9JkiRJBdekFeAhhEvI7Pd9S4xxZk27IWKM8bKmXFNS4Tw1/rOcvtW6dkwhiSSp\n1ar6qHHzO/UoTA5JkiRJWk5Tt0CpJFMAvw+YWdNuiAhYAJdamJPuejXtCJKk1urLCXDzgNz+iz6H\nth2Wtae+AHfsBwfdAANOLFo8SZIkSeWtqQXw3WseP1yhLakEvDV037QjSJJai3zFb8gufgOsuwNU\nVhU+jyRJkiQtp0kF8BjjM3W1JbVs3yxYxHeHPAbA9UdvmTPepUMi98eVJJWahfPhDzvBjAlw8tPw\npz3yzzv73eLmkiRJkqRaJFLlCiEcD7weY3yzjjnfBbaKMd6dxDUlNd0OVz619Png+17PGtuk98pF\nTiNJapE+GguTn4FND4IPX4D//DR7vLbi92qbQdc1Cp9PkiRJkhogqWWed5LZB7zWAjhwCHApYAFc\nStnsBYtqHXt08K5FTCJJapEquy17/tTQxh17+gvJZpEkSZKkZijmPgdtyNwEU5IkSS3Vlx807biL\nv4Q27ZLNIkmSJEnNVMwC+EbAV0W8nqQ8fvmvN9KOIElqyW7euvHHXDwD2nj/CEmSJEktT5P/pRJC\n+MsKXYeGEPrmmdoGWAfYBXi4qdeT1HQxRu59ZRpDH3qb+Qur044jSWqpqhc3bN42J8Mrf8o8H3Ci\nxW9JkiRJLVZz/rVy4nLPI7BlzVc+EXgJ+Hkzriepie59ZRrnP/C/euftuH7PIqSRJLVIy+/7vUS/\ngTBp9HJzqpY93/tS+HYOdFm10MkkSZIkqcmaUwBfr+YxAJOA64Eb8sxbDHwVY5zTjGtJaoaGFL8B\nfrXfJgVOIklqkR4+J3//8cMzj/NmwUrds8fad8p8SZIkSVIL1uQCeIxx6pLnIYShwKjl+yS1fFus\n3Y03Plq2mm/ztfOs/pMklbaJo+CV2+ues2LxW5IkSZJaiUQ2bIwxDk3iPJKS9/ns+bWODTtjJ554\n5zNOuWcsd/xwG0IIRUwmSUrVzdvCl+/VPn7u5OJlkSRJkqQCSfyORSGENkAvoEO+8Rjjh0lfU1J+\nMUa2veKpWsdDCOzznd5MuXpQEVNJklLz+EXw/E11z+m5IZwyCjqsXJxMkiRJklRAiRXAQwjfA64G\ndqeW4jeZm2EmXnSXlCvGyHrnj0w7hiSppch3k8sVff+fsNG+hc8iSZIkSUWSSDE6hLAJ8HxN8wng\nIOAN4DNgKzIrwkcBrv6WiqS+4vcrF+5VpCSSpNTN+6ph8yx+S5IkSSoxSa3GvhhoB2wTY/xfCKEa\neDDGeGkIoTNwI3AAcGJC15PUBF07tuXQ/mux0wa9WHXl2j6oIUkqOX89sv45Q2YVPIYkSZIkFVtS\nBfCBwIgY4/+W6wsAMcY5IYRTgTeBy7AILhXcosXVeftfvnAvOrZrU+Q0kqTUffxq/v7KKpgzAzr1\nAG+ELEmSJKkEJVUA7wV8sFx7EdBpSSPGuCiEMAo4LKHrSVrBcxO+5NjbX6p1/KUL9rT4LUnlaOyd\n2e1VN4UzXlzW7tyzqHEkSZIkqZgqEjrPTKDLcu0vgXVWmPMt0IC7L0lqirqK36fu2o/Vu3YsYhpJ\nUovx0FnZ7ZMeTyeHJEmSJKUgqQL4RKDvcu2xwN4hhNUAavYBPwSYnND1JDXC+QdsmnYESVIabs9z\nw+OOXYufQ5IkSZJSklQB/HFg95pCN8AfgB7AayGEfwH/A9YFbk/oepIkSarPR69kt096Ip0ckiRJ\nkpSSpArgfwJOAlYCiDE+DAyuaR8BrAb8GrgxoetJZS/GyPufzWZxdUw7iiSpWD55HSq7Zb6W7O39\n+MWZ9jUbZtqLF8G1G2X6VtRn22IllSRJkqQWIZGbYMYYpwP3rdB3YwjhFjI3yPw8xmiVTkpIjJH1\nzh+5tD3m3N1rnfvuZfsVI5IkqdAWzofbdlvWfugsePtBmDQ6057zef6i9xK/nFjQeJIkSZLUEiVS\nAK9NjHEx8FkhryGVm77nPZzTt8tvRuWde/ePtqVjuzaFjiRJKoYrVs/tW1L8bojOvRKLIkmSJEmt\nRUEL4JKS9dbHVQ2e+9Q5u7H+ql0KmEaS1Gpc/GXaCSRJkiQpFU0qgIcQnm7i9WKMcc8mHpuYEMJx\nwN01zZNjjDk35wwh7AhcBGwPdAQmAH8BbqpZ2Z7vvCcAZwCbAYuB14BrY4wjapm/EnAecAyZm4R+\nDYwGhsQYxzf19am0xBi5ePhbjHhzOu3aNGzb/ilXDypwKklSUU1s4o9eh/4etvx+slkkSZIkqRVp\n6grwgU08LvV9wEMIfYCbgG+AvMtjQwiHAPcD88nsbT4TOAj4HbATcFSeY64FzgE+InNT0PZkCtsP\nhRB+GmO8eYX5HYAnas73KnAD0Kfm3INCCHvEGF9q7utV63ffK9P464sfph1DklRoX34AnXpC+y7w\nz+Pgu0fAnC/gsQsad54BJ0LPDWH7n0CF22BJkiRJKm9NKoDHGBu2DLWFCSEE4A5gBvAA8Is8c7qS\nKWAvBgbGGF+t6b8YeBo4MoRwTIzx3uWO2ZFM8XsisE2M8aua/muAscC1IYQRMcYpy13qbDLF738D\nR8cYq2uOuQ8YBvwlhPC9Jf0qX+c98L9Gzd97szx7xEqSWrZ8N698/9E65jd8SyxJkiRJKmetspDd\nDD8D9gB+CMypZc6RwKrAvUuK3wAxxvlktkQB+MkKx5xW83jFkuJ3zTFTgFuADjXXBJYW4pccc+7y\nRe4Y43BgDJltVHZrxGuTAPj1EZunHUGStLyHf5EpcF++Oow4GxZ8kz3+zeeNO9/Z7yaXTZIkSZJK\nXNkUwEMImwJXAzfEGP9bx9Q9ah7zLbv6LzAX2LFmC5OGHPPICnMA1gfWAd6PMU5u4DFSg/To3D7t\nCJKkJf55PLzyp8zzRfPh1T/DVWtBrNkVLka4dsPGnbPrGslmlCRJkqQS1tQ9wLOEEHZt6Nx6is8F\nEUJoC9wDfAjUt5HmxjWP7684EGNcFEKYDHwH6AeMDyF0BtYCvokxTs9zvg9qHjdqyDXqOCavEMLY\nWoY2qe9YtX5jL9qLAZc/ubQ96coDUkwjScoy7h54Z3j+saHdm3bOCz9tchxJkiRJKkeJFMCB0TT8\nBpdp3I3pEqA/sHOMcV49c5dswlnb5ppL+rs3cX5Tj5GyTL7qAEIITLl6ELPnL2Tlju3SjiRJAlg4\nD65eFxYvSOZ8A8+HTQ+G1TdL5nySJEmSVEaSKoBfSv4CeHdgG2BH4CFgXELXa7AQwrZkVn3/Nsb4\nQhKnrHlsaMF/icbMb/A1YowD8p4gszJ8q0ZcUy3Q+5/Nztv/yFm7kNlKPsPitySlqHoxPHoerNkf\ntvw+XNG7+efc/zeZc7XvAsv9/16SJEmS1DiJFMBjjJV1jYcQTgRuAi5M4noNtdzWJ+8DFzfwsCWr\nr7vVMt51hXn1zc+32rux11AZ+t9HVRx087M5/Q+duTObrtE1zxGSpFRc2mPZ82Er3id7OWttDR+/\nWvv4kFnw+t+h30DotlZS6SRJkiSprBXlJpgxxjuBF4Ari3G95XQhs4/2psD8EEJc8gUMqZnzp5q+\n62va79U85uy/XVNQXw9YBEwCiDHOAT4GuoQQ8t2VasmdrZbf77vWa9RxjMpMvuL3G0P24Xtr1/Z7\nE0lS0f1204bNO+c9OPkpaJPnRsWrfQcqqzIrvfsfa/FbkiRJkhKU1BYoDfEGcHIRrwewAPhzLWNb\nkdkX/FkyBekl26M8DRwL7Af8Y4VjdgU6Af+NMS6/sefTwHE1x9yxwjH7LzdniYlkbsi5UQhhvRjj\n5AYcI9FtJbc6kaQWIcaG3chyw33h6HugbYdM++IvMo+Vy/0y8yfPJR5PkiRJkpRRzAJ4nyJfj5ob\nXv4431gIoZJMAfyuGOPtyw39G/g1cEwI4aYY46s18zsCl9fM+f0Kp/sDmQL4hSGEYTHGr2qO6Quc\nQaYQv7QwHmOMIYQ/kFkR/5sQwtExxuqaYw4BdgHeAZ5p4kuXJEmFcFUfWPB1w+b+chJ07pl/rNJd\nziRJkiSpGApekA4htAF+CBxJZrV1ixZj/DqEcDKZQvjoEMK9wEzgYGDjmv77Vjjm+RDCdcDZwJsh\nhH8D7YGjgR7AT2OMU1a41HXAgWS+Ly+FEJ4C1gGOAuYCP1pSFFf5WbTYP3pJalGqq+HSVRp3TG3F\nb0mSJElS0SRSAA8hTKrj/KvXPH4LXJDE9QotxjgshLAbmZt2HgF0BCaQKXDfGGOMeY45J4TwJnAm\ncApQDYwDrokxjsgzf0EIYS/gPOD7wM+Br4FhwJAY4zuFeG1qHT74/JucvklXHpBCEkkSMdZf/K5o\nC9WLlrVd4S1JkiRJLUJSK8ArgJyiMLAQ+B/wMnBTjHF8QtdrthhjJVBZx/hzQKMqjjHGu4C7GjF/\nHpmbcQ6pb67KS4e22fenfbNyHyoqQkppJKmMzf4MflvbPatrXPT5sj2+JUmSJEktSiIF8Bhj3yTO\nIynj/nEfZbW7dvTml5JUNAvnwRW9GzbXld6SJEmS1KIV9aaUkhrm/rEfpx1BksrTwvn1F7/X3wP+\n7x7o0KU4mSRJkiRJTWYBXGqB+q3amU+/np92DEkqP389vO7xIbMguCWVJEmSJLUWiRbAQwgHAVsC\nawP59myIMcaTkrymVIqenzgj7QiSVH4+/R9Mfa728V9NsfgtSZIkSa1MIgXwEMK6wAhgM6CufxlG\nwAK4JElqOb6aCp+Ph38cXfucY++HlVYpXiZJkiRJUiKSWgF+I/Ad4C/A3cDHwKKEzi2Vlfc/m512\nBEkqH8/8BkZdUf+8DfcqfBZJkiRJUuKSKoDvATwWY/xxQueTytY+v/tvVrtLB7fql6SCqa34vfel\nsO2p8OV70Hvz4maSJEmSJCUmqcraQuB/CZ1L0nLuOWnbtCNIUvnZ6azM4xpbpJtDkiRJktQsFQmd\n5znguwmdSypL3y6q5vi/vJzT338d95yVpIL40575+y+ZWdwckiRJkqSCSaoAfgmwawjhmITOJ5Wd\nm0dN4L/vf5F2DEkqD19NhY9fze2vrIKKNsXPI0mSJEkqiES2QIkxvhZC2BN4OIRwKjAOqMo/NV6W\nxDWlUnPjUx/k9J2ww7opJJGkEjN3Jtw2EGZNha1Pglf/nDsnVMCQr4oeTZIkSZJUWIkUwEMI3YCr\ngB7AbjVf+UTAArjUQEMPcWchSWq236y37Hm+4jfAxTOKk0WSJEmSVFRJ3QTzd8BA4EngHuATYFFC\n55ZK3k15Vn9LkhJQ2a3+OT9/GyqS2hVOkiRJktSSJFUAPxB4Psa4T0Lnk8rKb594P+0IklR63h7W\nsHnd1i5oDEmSJElSepIqgK8EPJ/QuSQBL1+wZ9oRJKn1enIoPHtd/fNOHVP4LJIkSZKk1CT1ed/X\ngH4JnUsqKwsWLc7p2/+7vVmta8cU0khSCahe3LDi96YHwRqbFz6PJEmSJCk1Sa0AvwwYEULYOcb4\nbELnlMrCobfkfnji1mO3SiGJJJWIS3vk9p05FnptUPwskiRJkqRUJVUAXwMYATwdQvg7MBaoyjcx\nxnh3QteUSsL46V9ntV+5cC9CCCmlkaRWbubk/P0WvyVJkiSpLCVVAL8TiEAAjq/5iivMCTV9FsCl\nOqy6coe0I0hS63XjltntbU6GQdemEkWSJEmSlL6kCuA/TOg8kiRJTXP1url9Fr8lSZIkqawlUgCP\nMd6VxHmkcnL5iHe4/dlaPqovSWqcRd/C/FnZff93TypRJEmSJEktR0XaAaRyVDVvocVvSUrS5avm\n9m12cPFzSJIkSZJaFAvgUgr+9eq0vP0jfrpzkZNIUitVXb3s+cL5ueOXzCxeFkmSJElSi5XIFigh\nhEkNnBpjjOsncU2pNftzLau/v7tWtyInkaRW6NqN4JvPMs8v+QoeuyB3TkWb4maSJEmSJLVISd0E\nswKIefq7Ad1rnn8CLEzoelKrNr0qz2pFSVL95sxYVvwGuHSV3DmVVcXLI0mSJElq0ZK6CWbf2sZC\nCBsANwKdgX2TuJ4kSSpT1/RLO4EkSZIkqRUp+B7gMcYJwOHAWsCQQl9Paq1u+f5WaUeQpJbtiQb8\nGPGrKQWPIUmSJElqPYpyE8wY43zgCeD/FeN6UksWY+5uQat0asf+3+2dQhpJaiVihOeur3vOMX+H\nlfJsiSJJkiRJKltJ7QHeEIsAK3wqe+udPzKrfdLO6/GzPTekoiKklEiSWoFbt69/ziaDCp9DkiRJ\nktSqFGUFeAihF3AYMK0Y15Nak4sP3IxuK7VLO4YktVxTX4Av3s3uu/Az2PWXmedtO3rjS0mSJElS\nXomsAA8hXFLH+fsAhwDdgPOTuJ4kSSoxixbA5atlnm9+NBx+W+b5zMlwx36589t1hD0uynxJkiRJ\nklSLpLZAqaxn/Gvg8hjjbxK6ntQq5dv/W5LK3ievwW0Dl7XfvA+2OgFW2xRu3DJ3/nl+oEySJEmS\n1DBJFcB3r6W/GvgKeDfGuCiha0mt1reLq9OOIEnpixHeuh+evwmmv55/ziPnwmdv5R/r2LVg0SRJ\nkiRJpSWRAniM8ZkkziOVunnfLs5q3/L9rVJKIkkpeek2eOSX9c+rrfj9q6nJ5pEkSZIklbSkVoBL\naoC5yxXAV2rXhkGbr5FiGkkqohhhaPfmneOSmVDRJpE4kiRJkqTyUNHUA0MIHUIIL4cQngohtKtj\nXvuaOS/WNU8qB8sXwOctXFzHTEkqMc0tfldWWfyWJEmSJDVakwvgwLHAAOC3McaFtU2KMX4LXANs\nW3OMVLZGvft52hEkqfgquzVs3iUz8/ef815yWSRJkiRJZaU5BfDDgUkxxpH1TYwxPgp8ABzVjOtJ\nrd4VI8enHUGSiuuJS+oe33NIZnX3khXeg65bNrbZoXDxl7By74JGlCRJkiSVrubsAd4fqLf4vZz/\nAgc043qSJKm1WLQA/nsNPHdD7lhde3lvcxJsdgi07wLtOhY2oyRJkiSp5DWnAN4L+KwR8z8Dejbj\nepIkqbW4fLX8/Rd+Vv9e3p17JZ9HkiRJklSWmrMFyjygSyPmdwHmN+N6kiSpNfj3Sfn7fzXFVd2S\nJEmSpKJqTgF8GrBNI+ZvDXzYjOtJJeX6o7dMO4IkJW/YGfDWv/OPrbRKcbNIkiRJkspecwrgo4Ht\nQwhb1zcxhDAA2BEY1YzrSSVl2/V6pB1BkpIVI7z+1/xjew0tbhZJkiRJkmheAfxmIAL/CiFsWtuk\nEMImwL+AxcCtzbieVFJWalfPHriS1Nrcc2jtYzsPLlYKSZIkSZKWavJNMGOM74UQLgUqgddCCP8G\nngY+IlMYXxvYEzgC6ABcEmN8r9mJpVbq86+zt8Dv3KE596CVpBamuhomjc7tr6wqehRJkiRJkpZo\nVgUuxnhpCGERMAT4PvD/VpgSgIXAhTHGq5pzLam1u/zh8Vnt9m2b8wEMSWpBFi+Cy3rm9lv8liRJ\nkiSlrNlLUGOMV4YQ/gb8CNgJWINM4fsT4Fngjhjj1OZeR2rtZs9fmHYESUpWjDC0e/6x8z8uahRJ\nkiRJkvJJZA+GmgL3kCTOJZWqFybNSDuCJDVfXUXvJXb8KXToUpQ4kiRJkiTVxT0YpCL46Ku5zF9Y\nnXYMSWqeD56ov/gNsM/lBY8iSZIkSVJDWACXiuD4v7ycdgRJap7qavjbkfXPW2vrwmeRJEmSJKmB\nEtkCRVLdJn0xJ6v926O2SCmJJDXRG/+oe7z35jDgRNjmpKLEkSRJkiSpISyASynwhpiSWp3hp9c+\ndslXUOGHyiRJkiRJLY8FcCkFm67RNe0IktRw7z6c2/ez1+Czt2GDvSx+S5IkSZJaLAvgUgq2Xa9H\n2hEkqX5jroOnhub2XzITKtpAj37FzyRJkiRJUiNYAJcKLMaY1b7+6C0JIaSURpIaqLJb7WMVbYqX\nQ5IkSZKkZvAzy1KBXf7w+Kz2vt/pnVISScojRnj7wUzB+6o+8Pawuovf255StGiSJEmSJDWXK8Cl\nAvvri1Oz2iu1d+WkpBZi8hi468Bl7QVfw79OqH3+zmfDXkMKn0uSJEmSpIRYAJcKbMGi6rQjSFLG\nv0+Ct/7d9OMtfkuSJEmSWhm3QJEkqRw8Wdm84nf/HyQWRZIkSZKkYnEFuCRJ5eDZ3zX+mJ+/Ay/e\nCmttBd89IvlMkiRJkiQVmAVwqUDmL1zME+98lnYMSeXsjfvgwQbctHLAiTD2zuy+yqrM475XJJ1K\nkiRJkqSisQAuFcgmFz+adgRJ5ezNfzWs+A0w6Do48HoYcy302wPWHlDQaJIkSZIkFYsFcKkAbhk1\nIW//X0/arshJJJWlym71jFfBF+/Dv06AI/4MFW0y/bv+svDZJEmSJEkqIgvgUgGM/N/0vP07b9ir\nyEkklZX6Ct8AQ2ZlHlfdCE5/oaBxJEmSJElKW0XaAaRSNHv+orQjSCo3DSl+nz0eQih8FkmSJEmS\nWghXgEsJiTFy/ZMf8OHMuXxv7W58OHNu1vifT9g6pWSSytr+18B2DdwLXJIkSZKkEmMBXErInr99\nhklfzql1fI9NVitiGkmt3jefw+Jvodva9c+9so45Fr8lSZIkSWXMAriUkLqK31OuHlTEJJJapdmf\nQudVIVTA0O7L+tffA457sPbjFsyGb2fnH7vkq0QjSpIkSZLU2lgAlyQpbfceC++OyD828Wn43Xfh\n52/ljr3/OPz9qPzHrbcbVHirD0mSJElSebMALiVg2gr7fS9vn81WL2ISSa1SbcXvJaqmwS3bwxkv\nwuzP4Lcb1T63sirZbJIkSZIktWIWwKUEjJ1a+zYDfzxuQBGTSGp1HruwYfO+GA+37Q6fjCtsHkmS\nJEmSSogFcCkBHdu1qXUshFDEJJJahRFnw6t/bvxx9RW/9/9N0/JIkiRJklSi3BxUSsAjb03P2//T\nPTYochJJLd43nzes+N1zw8ad93tHwXanNi2TJEmSJEklyhXgUgKGv/5JVvuJn+/K9Kr57LJhr5QS\nSWqRFi+Ea+sobB/6exj2E1izP5wyOtP30h/hkXPrPu/5H0OHLonFlCRJkiSpVFgAlwpgw9VXZsPV\nV047hqSW5rJ6fim25fczX8vb7lSY/zWMujy7f/szYMCJ0GtDcKslSZIkSZLysgAuSVIxVHare/yi\nz2sf2+2Xma+F86FNO6io/b4DkiRJkiRpGQvgkiQVWoz5+y+YDrOnQ49+DVvF3a5jsrkkSZIkSSpx\n3gRTkqRCGnkuDO2e23/JV9C+E/Rc3y1MJEmSJEkqEFeASwkbvFcdN7iTVF6++Rxe/mNuf2VV8bNI\nkiRJklSGXAEuJWzbvj3SjiCppbg2zy/ENj24+DkkSZIkSSpTFsClZvrymwVZ7X6rdkkpiaRW4eh7\n0k4gSZIkSVLZcAsUqZlmzf02q7161w4pJZGUqudugCcuyTz/wQPw4Qu5cy78tLiZJEmSJEkqcxbA\npWa67on3s9rBm9lJ5aV6MVy6wtZHfz08d96QWd7sUpIkSZKkInMLFKmZRv7PFZ1SWVux+F0bi9+S\nJEmSJBWdBXBJkppq4byGzfv524XNIUmSJEmS8rIALklSU7z9IFzRu/55/Y+DbmsXPo8kSZIkScrh\nHuBSM6x4A8zt+zVwKwRJrd+/Tszt22kw7HQWdOwGFW2KnUiSJEmSJK3AArjUDA+9OT2r/Y+Tt08p\niaSiquyWv3/vocXNIUmSJEmS6uQWKFIz9OrcPqsdvMmdVPpqK36f8Upxc0iSJEmSpHq5Alxqhrc+\nqUo7gqRiGn5m/v5K/18gSZIkSVJL5ApwqRluGTUx7QiSiuXZ6+G1e3L7LX5LkiRJktRiWQCXJKk+\n1dXw5JDc/otnFD+LJEmSJElqMAvgkiTV59JVcvs2PQjauJOYJEmSJEktmf9ylyRpRYsWwOWr1T7+\ny0nQuWfx8kiSJEmSpCZxBbiUkC4d/H2SVBLmzKi7+B3aWPyWJEmSJKmVsGInJaRPj05pR5DUHP+9\nBp6+vP55v3i/8FkkSZIkSVIiXAEuNdG0mXOz2tv365FSEklNtmhB5jHGhhW/f/AAdO5V2EySJEmS\nJCkxrgCXmujeVz7Mal94wKYpJZHUKDHCm/+EB09p3HGXzISKNoXJJEmSJEmSCsICuNRE02fNz2q3\nbeMHKqQW7+Nx8KfdGza3TQe48FOo8L0tSZIkSVJrZQFcaqIHXvs47QiSGiPGhhe/L/gE2ncubB5J\nkiRJklRwFsAlSaVv+Bnw2l8bNven4yx+S5IkSZJUIiyAS01QXR3TjiCpIb6aCjds3vD5bdpDz/UL\nl0eSJEmSJBWVBXCpCb6evzCrfc2RjSiwSSqsD1+Ev+xb/7xzJ0OnHpnnsz+FUAFdVitsNkmSJEmS\nVFQWwKUmWLzCCvAjB6ydUhJJORpS/K6sym6v3LswWSRJkiRJUqosgEtNcO8r07LaIYSUkkhlrLoa\nHv457Pgz+PrjzL7dc2bUf9xFnxc+myRJkiRJahEsgEtN8N6ns9OOIJW3GOHSVTLPx97Z8OPOmwZt\nOxQkkiRJkiRJankq0g4gtUYr7gEuqciGdm/8MXsNhY5dE48iSZIkSZJaLleAS03wyax5aUeQytfi\nRQ2bt91PYPufwCrrFjaPJEmSJElqsSyAS03w/mffpB1BKl839m/YvP2vLmwOSZIkSZLU4rkFiiSp\ndan6MLu96qa5c/ofV5wskiRJkiSpRXMFuCSp9Xjswuz29qfDfldlnn/wBPztSDh5FKy1VfGzSZIk\nSZKkFscCuNRI8xcuTjuCVJ7euA9euDm7b98rlz3fcG+orCpuJkmSJEmS1KK5BYrUSDPnfJt2BKk8\nPXhKbl8Ixc8hSZIkSZJaDQvgUiOtWAB/+YI9U0oilaiZk+CZa+DLD5b1VXbLnXfRF8XLJEmSJEmS\nWiW3QJEa6bon3s9qr9a1Y0pJpBK0eBHc2D/zfNTlmcdT/5s77/j/QNv2xcslSZIkSZJaJVeAS43U\nqX2btCNIpSlGuKxnbv8fd83t67db4fNIkiRJkqRWzwK41EjTq+anHUEqPfO+gqHdGzb3/I8KGkWS\nJEmSJJUOC+BSI02fNS/tCFLp+XXfhs078HrosHIhk0iSJEmSpBJiAVxqpE9cAS4l68/7Nnzu1j8s\nXA5JkiRJklRyLIBLktLzzecw7cXc/soq+N5R2X39di9OJkmSJEmSVDIsgEuS0nPthrl953+ceTz0\n98v6VuoBxz1YnEySJEmSJKlktE07gNSaVFfHtCNIpeOGLXP7KquWPW/TLrstSZIkSZLUSBbApUbo\nd8HItCNIrU91NYQAs6fD/SfD1Gfzz9v+jOLmkiRJkiRJJc8CuNQMm/ReOe0IUss150u4Zv2Gz9/v\nysJlkSRJkiRJZck9wKVm+MU+G6cdQWq5GlP8HjKrYDEkSZIkSVL5sgAuNUPnDn6IQsox/2uo7Nbw\n+Z16ZrZIkSRJkiRJSpjVO6mBxn34VU6fNTspj6v7NGze6t+Fk0dB2/aFzSNJkiRJksqWBXCpgQ6/\n9fmcvn69OqeQRGrBFs6ve/zw22H1zaB9Z1ilb1EiSZIkSZKk8mUBXGqiiwZtympdO6YdQ2o5Fn0L\nV6ye2z9klh+XkCRJkiRJqbAALjXRj3fpl3YEqWW5fNXcvsqq4ueQJEmSJEmq4U0wpSbYY5PV0o4g\ntXy9v5d2AkmSJEmSVOZcAS41QIwxq927m1ufqMzFCP/7N6y2Kfxhp/xzTnqyuJkkSZIkSZJWYAFc\naoBvF1dntb+7ZreUkkgtxNDudY+fOgba+YsiSZIkSZKULrdAkRrgjuemZLXXX7VzOkGktM2vgsoG\n/AJojc0Ln0WSJEmSJKkeFsClBhj++idZ7W3X65FSEilF7wyHq9epf97/u7fwWSRJkiRJkhqg5Avg\nIYSeIYQfhxAeDCFMCCHMCyFUhRCeDSGcFELI+z0IIewYQhgZQpgZQpgbQngzhDA4hNCmjmudEEJ4\nOYTwTc01RocQDqxj/kohhKEhhPdCCPNDCJ+HEP4ZQtg0ideu5Iyf/nVWO4SQUhIpBTFmVn3/8/i6\n563UAy7+Ejbevzi5JEmSJEmS6lEOe4AfBfwemA6MAj4EVgcOB24H9g8hHBWXu8thCOEQ4H5gPnAf\nMBM4CPgdsFPNObOEEK4FzgE+Av4EtAeOAR4KIfw0xnjzCvM7AE/UnO9V4AagT825B4UQ9ogxvpTQ\n90DN8NyEL9OOIBXfO8PrL3gv74ePwro7FC6PJEmSJElSE5RDAfx94GDg4Rjj0jsZhhAuAF4GjiBT\nDL+/pr8rmQL2YmBgjPHVmv6LgaeBI0MIx8QY713uXDuSKX5PBLaJMX5V038NMBa4NoQwIsY4Zblc\nZ5Mpfv8bOHpJthDCfcAw4C8hhO8tn1npOPZ2fw+hMtSQ4vdaA2CtreF7R0KfbQufSZIkSZIkqZFK\nfguUGOPTMcaHViwkxxg/Bf5Q0xy43NCRwKrAvUuK3zXz5wMX1TR/ssJlTqt5vGJJ8bvmmCnALUAH\n4IdL+kNm/4wlx5y7fLYY43BgDLAZsFuDX6gkJaGyW8NucllZBSc/DQf8xuK3JEmSJElqsUq+AF6P\nhTWPi5br26Pm8dE88/8LzAV2rNnCpCHHPLLCHID1gXWA92OMkxt4jFqIET/dOe0IUvK+ndOwwjdk\nit+SJEmSJEmtQDlsgZJXCKEtsOQz/ssXrjeueXx/xWNijItCCJOB7wD9gPEhhM7AWsA3McbpeS71\nQc3jRg25Rh3H5BVCGFvL0Cb1Hav6ffTV3Jy+767VwCKh1Bq89Ed45NyGzT13MnTqUdg8kiRJkiRJ\nCSrbAjhwNfBdYGSM8bHl+pdUN2tb4rikv3sT5zf1GKXgvPv/l9Ue+bNdUkoiFcD4EfUXvwe/BV1W\ng7Yd6p4nSZIkSZLUApVlATyE8DMyN618FziusYfXPMZGHteY+Q2+RoxxQN4TZFaGb9WIayqPFVeA\nx0b/sUst1Ohfw+gr657zw0ege5/i5JEkSZIkSSqAsiuAhxDOAG4A3gH2jDHOXGHKktXXte1z0XWF\nefXNz7fau7HXUEqmzMgugG+2RtdaZkqtyAOnwJv31T7+46dh7by/W5MkSZIkSWpVyqoAHkIYDPwO\neItM8fvzPNPeA7Yms/921v7aNfuGr0fmppmTAGKMc0IIHwNrhRDWyLMP+IY1j8vv9/1ezWNte3zn\nO0YtQAih/klSSxQjDO3esLkWvyVJkiRJUomoSDtAsYQQfkWm+P06sHstxW+Ap2se98sztivQCXg+\nxriggcfsv8IcgInAh8BGIYT1GniMJDXNQ4MbXvyu9IMnkiRJkiSpdJRFATyEcDGZm16OJbPy+8s6\npv8b+BI4JoSw9XLn6AhcXtP8/QrH/KHm8cIQwirLHdMXOANYANyxpD/GGJc75jchhIrljjkE2IXM\nFi3PNPAlqghOH7h+2hGkphl7R+1je1wEF3+ZKXxb/JYkSZIkSSWm5LdACSGcAFwKLAbGAD/Ls43F\nlBjjnQAxxq9DCCeTKYSPDiHcC8wEDgY2runP2jw3xvh8COE64GzgzRDCv4H2wNFAD+CnMcYpK1zz\nOuBA4EjgpRDCU8A6wFHAXOBHMcbqZn8DlJgB665S/ySppfjifZg5Caqm1T7n/90LG+9f+7gkSZIk\nSVIrV/IFcDJ7dgO0AQbXMucZ4M4ljRjjsBDCbsCFwBFAR2ACmQL3jTUruLPEGM8JIbwJnAmcAlQD\n44BrYowj8sxfEELYCzgP+D7wc+BrYBgwJMb4TmNfqJK14h/z/IX+PkKtwFdT4dHz4L2Rtc9ZpS+c\n9UbRIkmSJEmSJKWl5AvgMcZKoLIJxz0HHNDIY+4C7mrE/HnAkJovtTAffTUvq7161w4pJZEa6KHB\ndW93AjBkFngzV0mSJEmSVCbKYg9wqSl2+c2orPbWfXuklERqoPqK38f83eK3JEmSJEkqKxbAJakU\nTB5T/5xNBhU+hyRJkiRJUgtiAVySWqPX/w6V3TJfc2fCXQfWPf+0Z4uTS5IkSZIkqQWxAC5Jrc1D\nZ8Gwnyxr/2a93DnH/hs6dIUfPgqVVdD7e8XLJ0mSJEmS1EKU/E0wpSSct/8maUeQMqqrYeyddc85\n7dlMwfv8aUWJJEmSJEmS1FK5AlzK4/3PZme1T921X0pJpBW8N7Lu8e7ruNpbkiRJkiSphgVwKY//\nvv9FVjuEkFISaQX3HVv3+OD/FSeHJEmSJElSK+AWKFIelz88Pu0I0jKLvoXffQfmfF77nD0uhl1/\nUbxMkiRJkiRJrYAFcElqyebOzH+TS4A9L4FdziluHkmSJEmSpFbELVAkqaVavKj24jdY/JYkSZIk\nSaqHBXApj3V6dFr6/Oy9N0oxicrWJ6/DZT1rH7+oju1QJEmSJEmSBLgFipTXhzPnLn3eu1vHFJOo\n7CxaAJevVvecs9+Fth2Kk0eSJEmSJKkVswAurWDR4uqsdrs2IaUkKiuzP4MPn4cHTq19zqljYI3N\ni5dJkiRJkiSplbMALq1gxpxvs9obr941pSQqG9PfgD/uWvec4x60+C1JkiRJktRIFsClFbz24VdZ\n7X6rdk4piUrasNPh9b81bO7PXoMe/QqbR5IkSZIkqQRZAJdW8N8Pvsxqd2zXJqUkKgkxZr7efSjT\n3uwQqOzWsGMvngFt/N+0JEmSJElSU1lZkVbw95c+TDuCSsWf9oSPX2368Ra/JUmSJEmSmqUi7QCS\nVJJmTWt68btDVxgyK9E4kiRJkiRJ5cjlhZJUCNd/t3Hzzx4PXdcsTBZJkiRJkqQyZQFcqsMem6yW\ndgS1Rgu+adi8n70OPdYraBRJkiRJkqRyZgFcqsPVR3wv7Qhqja5aK7u99UnQc/1MYXz0lZm+U0Zb\n/JYkSZIkSSowC+DSchYtrs5qd2rvW0SN9MApuX0HXrfs+cBfFS+LJEmSJElSmfMmmNJyhvzn7ax2\n5/ZtUkqiFuXtYVDZDd4ZXv/cN+/Lbm92aCESSZIkSZIkqQEsgEvL+dtLH2a1QwgpJVGL8dVU+NcJ\nmef/PB6eu6H2ufO/zu37v7sKk0uSJEmSJEn1cn8Hqcbi6ph2BLU0C2bDDZtn9z1xCex0Vnbfn/eB\naS/lHn/Gy4XLJkmSJEmSpHpZAJdq/Oaxd9OOoJbmDzvXPf7Fe3DLtrWPr7pxsnkkSZIkSZLUKG6B\nItX44zOT0o6gluarKfn7J42GGOsufg/4YSESSZIkSZIkqRFcAS7VYsIV+6cdQWlZOB+uWL328bsP\nqf8cB12fWBxJkiRJkiQ1jQVwqRZt2/gBibJVV/G7LhvsBQMvgLUHJJtHkiRJkiRJTWIBXJKSUFmV\ndgJJkiRJkiStwCWuEjDjmwVZ7Tcr90kpiVK3aEFu32nPwgkP1X7M2eMLl0eSJEmSJElN5gpwCRjz\nwZdZ7a4d26WURKla9C1cvlpuf+/v1X7Mj5+CrmsWLpMkSZIkSZKazBXgEjD4vtfTjqC0ffE+XL5q\nbv/yW5ucOXbZ8w33yYytvXXhs0mSJEmSJKlJXAEuqTx89jb8fsfM819OgrYdoEOXZeO3bJN7zE6D\ns9u9Nshsd7LgG1h1o4JFlSRJkiRJUjIsgEsqfdNehj/vvax9Tb/s8TW3yn/cXpW5fW53IkmSJEmS\n1GpYAFfZm141L6vdq0uHlJIoccPPgNf+Wv+8T8bl9g2ZBSEkHkmSJEmSJEnF4x7gKntzFizKao/+\n5cB0gihZk8c0rPidz8UzLH5LkiRJkiSVAFeAq+ydcs/YrHaXDr4tWr3Kbk0/9qIvoI3/DUiSJEmS\nJJUCV4Cr7E36Yk7aEZSkxYvqn1OXtu2TySFJkiRJkqTUucxRUstUXQ2XrrKsfcC1sO3JdR8z50u4\nZv3c/lX6wllvZPfFCH/ZD6a9uKzvgk+aHFeSJEmSJEktjwVwSS3T8sVvgJG/gKnPwdsPQtuOcOGn\ny/bpnjERbtqq9nOtWPyGzLE/ehQu7QlxMZzxCrTvnFx+SZIkSZIkpc4CuKSW55oN8ve//WDmcdF8\nGNodznkffrtR3eeqrKp9LAQYMrNJESVJkiRJktTyuQe4ytpzE77Mav98r3qKqSq8F26FOV80bG59\nxe+LPm9+HkmSJEmSJLVargBXWTv29pey2jus3zOlJAJgwTfw2PnJnKuuld+SJEmSJEkqC64Al5az\n0epd0o5QvmKEq9bK7f/puMad5yfPW/yWJEmSJEkS4ApwKUu3ldqlHaH8TH8D/rhr/rFz3oeVV4ef\nvwO/2yzTt9t58MzVuXOHzFp2U0xJkiRJkiQJC+BSlmABtXievwkev6juOSuvnnnstlb2qu7dz4f3\nH4O//1+mvcHeFr8lSZIkSZKUwwK4ylaMMe0I5Wv2Z/UXv+vbxmSjfd3qRJIkSZIkSXVyD3CVrap5\nC7PaD5y+Y0pJStSCb2Dm5Pxjv92o7mMv/jL5PJIkSZIkSSo7rgBX2Zo559us9lbrrJJSkhJU2S27\n3aU3nP4CdOoBV/Wp+9iz34U27sUuSZIkSZKk5rMArrL128ffTztCaVowO7fvm0/hN+vBmv1hwde5\n40f/FTY9qPDZJEmSJEmSVFYsgKtsPfy/6WlHKC3vDId/Hl/3nE9ey+1zH29JkiRJkiQViHuAS2q+\nLyfUX/zO5+x3k88iSZIkSZIk1bAALqn5bh7Q+GMOuRW6rpF8FkmSJEmSJKmGBXCVpUWLq7Pavz7i\neyklKQH59vyuz0b7Qf9jk88iSZIkSZIkLcc9wFWWJn05J6t91IA+KSUpAU8Ozd//iw+gy2oQI/z1\nCJj4VKb/gunQvlPx8kmSJEmSJKlsWQBXWbrtv5Oy2hUVIaUkrdicGXBNv/xj25+eKX4DhADHPVC8\nXJIkSZIkSVINC+AqS4/8b3raEVq3ym51jFUVL4ckSZIkSZJUB/cAV1ma8+3itCO0Xg8Nrn3sV1OL\nFkOSJEmSJEmqjwVwlZ1vF1XXP0n5vfMfGHtH/rFND4aVuhc1jiRJkiRJklQXt0BR2fng89lpR2i9\n/nlc/v4TR0LfnYqbRZIkSZIkSaqHBXCVnUE3Ppt2hNbpq1q2N/nZa9CjlpthSpIkSZIkSSmyAK6y\n9/ole6cdoeWLEW7YPLuv9+Zw2ph08kiSJEmSJEkN4B7gKnvdO7VPO0LLVr0YhnbP7bf4LUmSJEmS\npBbOArjKyoTPv8lqH/C93iklaUUu7ZHbt+svi59DkiRJkiRJaiQL4Cort46akN0+dkBKSVqJRd/m\n79/jouLmkCRJkiRJkprAPcBVVh547eO0I7R8//s33H9S7eOVVcXLIkmSJEmSJDWDK8AlLbNgdt3F\n7zNeLl4WSZIkSZIkqZksgEvK+OYLuGrt2sfX2xVW3bh4eSRJkiRJkqRmsgAuKePaDeoeP/4/xckh\nSZIkSZIkJcQCuMpGjDHtCK3XDx+BENJOIUmSJEmSJDWKBXCVjcfe/jTtCC1XZbfaxzqvCuvuWLws\nkiRJkiRJUkLaph1AKpY3PqpKO0LLMWcGXNOv9vEhs1zxLUmSJEmSpFbPFeAqG9Nmzs1qP/7zXVNK\n0gLUVfwGi9+SJEmSJEkqCRbAVTZGvDk9q73R6iunlCRlf/u/uscHv1WcHJIkSZIkSVKBWQCXStXr\n/4DbBsKMicv6YoQPHqv9mNNfhO59Ch5NkiRJkiRJKgb3AFdZ6tiuxH/38/hF8PxNmec3bVX//B8+\nCuvuUNhMkiRJkiRJUpFZAFdZ2na9nmlHKIwYoXrxsuJ3fY74M2ywF6zUvaCxJEmSJEmSpDRYAFdZ\niDFmtX+883opJSmAB0+DN/7RtGO/d2SyWSRJkiRJkqQWpMT3gZAynp84I6s98YtvUkqSsA9fanrx\n++Ivk80iSZIkSZIktTCuAFdZOPb2l7LafVbplFKShP1ln8Yf84sJ0GXV5LNIkiRJkiRJLYwFcJWl\n3TdZLe0IzfflB/XPOew22OLowmeRJEmSJEmSWiC3QFFZalMR0o7QfDdvXf8ci9+SJEmSJEkqYxbA\npdaoujq3r7IKLvlqWfvovxYvjyRJkiRJktQCuQWKys4Zu6+fdoTmmf4m/HGX7L59r8w8VlRkCuGS\nJEmSJEmSXAGu0rdwcfZq6WO3WzelJAkYPyK3+A2w9Y+Kn0WSJEmSJElq4SyAq+RN/nJOVnvN7iul\nlKSZ5lfBfcfmH2vXSl+TJEmSJEmSVEAWwFXyPp41L+0IzbdwHly9Tv4xtzyRJEmSJEmS8rIArpL3\nn9c/STtC813RO3//kFlFjSFJkiRJkiS1JhbAVfIefO3jtCM0zzUb5PZ16plZ+R1C8fNIkiRJkiRJ\nrYQFcKkle2IIzPkit//cScXPIkmSJEmSJLUyFsClluy563P7Lvq86DEkSZIkSZKk1sgCuNRSTXk2\nf3/bDsXNIUmSJEmSJLVSFsCllujTt+DOQbn9W/+o+FkkSZIkSZKkVqpt2gEk5fGHnXL7LpkJFW2K\nn0WSJEmSJElqpVwBrrJy4o59045Qv3uPze07e7zFb0mSJEmSJKmRLICrrPxy343TjlC3SaPh3RHZ\nfbucA13XTCWOJEmSJEmS1JpZAFdJe33arKx2p/YteBX12Lvg7kNy+/e8pPhZJEmSJEmSpBLgHuAq\naYfe8lxWO4SQUpJ63HccjP9Pbv+QWUWPIkmSJEmSJJUKV4BLLUG+4jdASy3YS5IkSZIkSa2ABXAp\nbTMm5u+vrCpuDkmSJEmSJKnEWABXyfrv+1+kHaFhbtoqt8+tTyRJkiRJkqRmcw9wlazj//Jy2hFq\nt3ghxGpYMDt3zJXfkiRJkiRJUiIsgKtstG/TQj7wMOJsePXP+cfOaMFFe0mSJEmSJKmVsQCusvHo\n4F3SjgCLF9Ve/AZYdePiZZEkSZIkqYWqrq5m5syZzJ49mwULFhBjTDuSpISEEOjQoQMrr7wyPXr0\noKKisItWLYCrbPRbtUvaEeD3O9Q+tssvipdDkiRJkqQWqrq6mmnTpjF37ty0o0gqgBgj8+fPZ/78\n+cyZM4c+ffoUtAhuAVxl4aAt1kw7Anw8Fr58v/bxPS8uXhZJkiRJklqomTNnMnfuXNq2bUvv3r3p\n3LlzwVeISiqe6upq5syZw6effsrcuXOZOXMmvXr1Ktj1/L+HSlLV3IVZ7aEHfyelJMv50x61jx11\nV/FySJIkSZLUgs2ePRuA3r17s/LKK1v8lkpMRUUFK6+8Mr179waWvecLdr2Cnl1KyRaXPp7VXrlj\nyh92qK6ufWyrE+A7hxYtiiRJkiRJLdmCBQsA6Ny5c8pJJBXSkvf4kvd8obgFispCuzYp/67nXydk\nt/tsByc9nn+uJEmSJEllbMkNL135LZW2EAJAwW9yawFcJWf+wsVpR8hW2S23z+K3JEmSJEmSytiS\nAnihWQBXyflwZgu5S/Sib+H1v6WdQpIkSZIkSSpbFsBVcq57/P2s9k8Grl/8EDHC5avmHxsyq6hR\nJEmSJEmSpHLlZkoqOY++/WlW+5f7bFz8EE9flr+/S28o0sc7JEmSJEmS6nPiiScSQmDKlClL+6ZM\nmUIIgRNPPDFn/gcffMBhhx1G7969CSHQvXv3omUtVfn+DJQcV4Cr5FVUFLHgXF0Nl65S+/gv3ite\nFkmSJEmSpAQtXryYQw89lAkTJnDcccex9tpr07FjRwD69u0L0OwiblLnaUkqKysZOnQoo0aNYuDA\ngWnHKTsWwFVSCn3X2Fq9/neo+ghGXVH7HLc+kSRJkiRJrcBaa63F+PHj6datW1b/5MmTeeeddzj5\n5JO57bbbUkpXeq666irOO+881lprrbSjlCQL4CopVz3yblb7ysO+V/iLjv41jL6y7jlrbuXWJ5Ik\nSZIkqVVo164dm2yySU7/J598AsCaa65Z7EglbY011mCNNdZIO0bJcg9wlZTb/jspq/397dYp/EXr\nK36f/DScMqrwOSRJkiRJUkkZPXo0IQQqKyvzjvft23fpliEAd955JyEE7rzzTh5++GF23HFHOnfu\nzCqrrMKRRx7JBx980KDr5tsDPITAbrvtBsDQoUMJIRBCYODAgYQQmDp1KlOnTl3aX9se4vW91vrO\nM2zYMH7wgx+w0UYb0blzZ7p06cKAAQO48cYbqa6uzjnvkv21J02axE033cTmm2/OSiutlLUVyfvv\nv88RRxzBKqusQufOndlxxx15+OGHs76fK/roo48488wz6devHx06dKBnz54cfPDBvPLKK1nz+vbt\ny9ChQwHYfffds17Xihlr24d9ypQpHHPMMfTq1YuOHTuy9dZbM2LEiLzfx6qqKgYPHrx0e5pNNtmE\n6667jkmTJjX6z6RUuAJcao7KbrWPnTsZOvUoXhZJkiRJkiTggQce4JFHHuGwww5j4MCBvP7669x/\n//2MGjWK559/no033rjR5xwyZAhTpkzhrrvuYrfddltaQO7bty8DBw7k+uuvB2Dw4MFLj9lyyy0b\nfP6+ffsyZMiQes9z3nnnUVFRwXbbbcdaa61FVVUVTz/9NGeddRavvPIK99xzT97zn3XWWYwZM4ZB\ngwZxwAEH0KZNGwDeffdddtppJ2bOnMmgQYPYfPPNmTRpEocddhgHHHBA3nONGzeOffbZh5kzZ7Lv\nvvty+OGH8+WXXzJs2DB23nlnHnzwwaXHDh48mGHDhvHMM89wwgknZP3CoiGmTp3KtttuS79+/Tju\nuOOYOXMm9913H4cccghPPvkku++++9K58+fPZ4899mDcuHH079+fY489lqqqKq644grGjBnTqOuW\nEgvgUlP979+1j+38c4vfkiRJkiQpFQ899BAPPfQQBx544NK+G264gcGDB3P66afz1FNPNfqclZWV\njB49mrvuuouBAwfmrEpfskq6ttXq9enbty+VlZX1nufhhx9m/fXXz+qrrq7mhz/8IXfffTdnnnkm\n2223Xc5x48aN47XXXmO99dbL6j/jjDOYOXMmt956Kz/5yU+W9j/yyCN5C+CLFi3i//7v//jmm28Y\nNWrU0lXxkNkiZptttuGkk05iypQpdOjQgcGDBzNr1iyeeeYZTjzxxEbfBHP06NFUVlYyZMiQpX3f\n//732W+//bjmmmuyCuDXXHMN48aN45hjjuHvf//70lXmF154IVtttVWjrltKLIBLTXX/SbWP7VVZ\ntBiSJEmSJJWTvuc9nHaEBpty9aBUrrvHHntkFb8BzjzzTG666Saefvpppk6dyrrrrptKtuZasfgN\nUFFRwVlnncXdd9/NY489lrcAfu655+YUv6dNm8bTTz/NBhtswKmnnpo1tv/++7PXXnvx5JNPZvU/\n/PDDTJw4kV/84hdZxW/I7I1+7rnnMnjwYJ566qlaV5A3xrrrrstFF12U1bfvvvuyzjrr8PLLL2f1\n33XXXVRUVHDVVVdlbbHSp08fBg8enHOecmEBXGqKP+2Z23fJV7B4AbRbqfh5JEmSJEmSaqxYmAVo\n06YNO++8MxMnTuS1115rtQXwGTNmcM011zBy5EgmTZrEnDlzssY//vjjvMdtu+22OX2vv/46ADvs\nsAMVFbm3Stx5551zCuAvvPACkNmaJN8q9SX7rI8fPz6RAviWW265dLuW5fXp02dpFoCvv/6aiRMn\n0qdPn7zbrOy8887NztJaWQCXmuLjV3P7KiqgwuK3JEmSJElK1+qrr563v3fv3kDmRomt0axZs9hm\nm22YPHky2267Lccffzw9evSgbdu2zJo1ixtuuIEFCxbkPXbJa1/eku9Dbd+vfP0zZswA4F//+led\nWb/55ps6xxuqe/fuefvbtm2bddPPr7/+GmjcaykXFsBVsn6+10aFOXG+G1+emacgLkmSJEmSEpfW\ntiJpWLIqedGiRXnHq6qq6NYtt07x2Wef5Z3/6aefAuQ9pjW4/fbbmTx5MkOGDMlZff3CCy9www03\n1Hrs8luCLNG1a1eg9u9Xvv4l37vhw4dz8MEHNzR6wTXltZSL3LX9Uok4eMs1kz3hzEn5i99DZkGv\nDZO9liRJkiRJKnurrLIKkNmrekUTJkxg1qxZeY975plncvoWL17Ms88+C0D//v2TC1mjTZs2LF68\nuKDnmTBhAgBHHHFEzli+11yfJd+HF154IWs19RJLvl/L23777QEYM2ZMg6+zZAuTJL4/tenatSv9\n+vXj448/ZsqUKTnj+V5LubAArpI1Z0H+3442SWU3uLGWvxzy/AZRkiRJkiSpuTbZZBO6du3K8OHD\n+fzzz5f2z5s3j5/97Ge1Hvf0008zYsSIrL6bb76ZiRMnsvvuuxdk/++ePXvyxRdfMG/evIKdZ8ne\n1qNHj87qf+2117jqqqsafa0+ffowcOBAJkyYwB//+MessUcffTRn/2+AQw45hPXXX59bbrmFkSNH\n5j3vCy+8wNy5c5e2e/bsCcCHH37Y6IyNcfzxx1NdXc35559PjHFp/7Rp07j++usLeu2WzC1QVDIW\nLMr+Ldr6q3Zp3gnnV8Ezv4EXbq59zoXl+/ERSZIkSZJUWO3ateOss87isssuo3///hx22GEsWrSI\nJ554gjXXXJM118z/6feDDjqIww47jMMOO4wNNtiAN954g5EjR9KjRw9uvfXWgmTdc889eeWVV9hv\nv/3Ydddd6dChA1tssQUHHXRQYuc5/vjjueaaaxg8eDCjRo1iww035IMPPmDEiBEcfvjh3HfffY3O\nfcstt7DTTjtx+umnM3LkSDbffHMmTZrE/fffzyGHHMLw4cOzbpDZrl07HnjgAfbdd18GDRrEjjvu\nyJZbbkmnTp2YNm0ar7zyCpMmTWL69Ol06tQJgN13352KigrOP/983nrrraUr+y+66KJG563Lueee\ny7Bhw7j33nt577332GeffaiqquKf//wnu+66K8OGDct7s89SZwFcJWP6rPlZ7ZXa594ht8EWzIar\n16l/XruOTb+GJEmSJElSPYYOHUqnTp3405/+xG233Ubv3r055phjqKysZLPNNst7zOGHH84pp5zC\nFVdcwcMPP0y7du04/PDDueqqq9hoo8LcM+2iiy5i1qxZPPTQQzz33HMsXryYE044odEF8LrOs+aa\nazJmzBjOO+88nn32WR577DE22WQTbr31Vvbaa68mFcA322wzXnjhBS644AKefvppnn76aTbffHMe\nfPBBxo8fz/Dhw5fur73E5ptvzhtvvMF1113HiBEjuOOOO6ioqGCNNdagf//+DB06lF69ei2dv+mm\nm3LXXXdx7bXXcuuttzJ//vylrzVJK620EqNGjeKSSy7h3//+N7/73e9Yb731uOCCC9hll10YNmxY\nzmspB2H55fAqHSGEsVtttdVWY8eOTTtK0Xz01Vx2/vWope1m3RQj317fKzphBKy3S9OvIUmSJEmS\ncowfPx7IFA3VOHfeeSc//OEPueOOOzjxxBPTjtPqHXvssfz973/n3XffZeONN047TrP86U9/4pRT\nTuEPf/gDp556atpxlmro+33AgAGMGzduXIxxQGOvUX5r3lWyRr37ef2T6jJrGty+d93F75OehPX3\nhJ+/Y/FbkiRJkiSplauurubTTz/N6X/qqae477772GyzzVpV8fuTTz7J6Zs2bRqXXXYZbdu25cAD\nD0whVbrcAkUl48tvvm36wb/fCT57q+45R/8N+mwDxz3Q9OtIkiRJkiSpxfj222/p06cPu+++O5ts\nsglt27bl7bff5oknnqB9+/bccsstaUdslCOOOIKFCxcyYMAAunfvzpQpUxgxYgRz587lqquuYq21\n1ko7YtFZAFfJaN+2iR9oGHtn/cXvHc6ETcvvN2SSJEmSJEnNUVlZ2aB5hx56KFtuuWVBs+TTrl07\nTjvtNJ5++mleeukl5s6dS69evTjqqKM477zz6N+/f9EzNcdxxx3HPffcw/33309VVRVdunRhu+22\n48wzz+Twww9PO14q3AO8RJXjHuCbXfIoc79dvLTdoD3A69vr+9zJ0KlHM5NJkiRJkqSGcg/w0hJC\naNA89y0vT8XYA9wV4CoZ/dfpznMTZgCw4WpdsgdnToYx18Lel0EIsNIqsGB23Sfc7FCL35IkSZIk\nSc3g4lulzQJ4ykIIawOXAvsBPYHpwDBgaIzxqxSjtTpLit8APTq3Xzbw7Vy4ccvM89f+2sCzBfi/\nuxLLJkmSJEmSJKn4LICnKISwPvA8sBowHHgX2BY4C9gvhLBTjHFGHadQLV6aPHNZ48o1GnbQT8dB\nh64w7ytYdaPCBJMkSZIkSZJUNBbA03UrmeL3z2KMNy3pDCFcB/wcuAI4LaVsrcqnVfMBaMNi1ggz\n+aLNapmB6uqGneC7R0DP9TPPu6xagISSJEmSJEmSis0CeEpCCP2AfYApwC0rDA8BTgGOCyGcE2Oc\nU+R4rc72Vz1FGxYzseNxyzo/fxFu3b5hJzjyL4UJJkmSJEmSJCk1FWkHKGN71Dw+HmPMWqYcY5wN\nPAd0AhpYwS1vnZifXfyGhhe/j/5b8oEkSZIkSZIkpc4V4OnZuObx/VrGPyCzQnwj4KnaThJCGFvL\n0CZNj9b6nNr2oYZNrKwqbBBJkiRJkiRJLYYrwNPTreaxtorskv7uhY/S+u1Q8U79k4bMKngOSZIk\nSZIkSS2HK8BbrlDzGOuaFGMckPfgzMrwrZIO1VJNP/jvMCLvtwJ2Ggx7DoEQ8o9LkiRJkiRJKkkW\nwNOzZIV3t1rGu64wT3U4ZOsN+Gj96cxfuJgNOlTB776TGTjlGVhzy1SzSZIkSZIkSUqHW6Ck572a\nx41qGd+w5rG2PcK1grVX6cQGq60M3dbO7PVdWWXxW5IkSZIkqUDuvPNOQgjceeedaUdRA5144omE\nEJgyZUraUYrGAnh6RtU87hNCyPpzCCGsDOwEzANeLHYwSZIkSZIkqdT17duXvn37ph0jUZWVlYQQ\nGD16dNpRWgwL4CmJMU4EHgf6AmesMDwU6AzcHWOcU+RokiRJkiRJkkrQVVddxfjx41lrrbXSjlI0\n7gGertOB54EbQwh7AuOB7YDdyWx9cmGK2SRJkiRJkiSVkDXWWIM11lgj7RhF5QrwFNWsAt8auJNM\n4fscYH3gRmCHGOOM9NJJkiRJkiQpTVOmTCGEwIknnsj777/P0UcfzWqrrUZFRQWjR49m7NixnHXW\nWWyxxRb06NGDjh07suGGG3LOOefw1Vdf5Zxv+T27R40axcCBA1l55ZXp2rUrgwYNYvz48XlzTJgw\ngaOOOopVVlmFzp07s+OOO/Lwww/XmX3s2LEcccQRrLbaanTo0IF1112X008/nenTp+fMXbIv9eTJ\nk7n55pvZbLPN6NixI3379uXKK68kxgjAv/71L7bddls6d+7Maqutxplnnsn8+fMb/X0dPXo0IQSm\nTp3K1KlTCSEs/TrxxBOXzhs2bBg/+MEP2GijjejcuTNdunRhwIAB3HjjjVRXV9f6OiZNmsRNN93E\n5ptvzkorrcTAgQOXznn//fc54ogjcr6Xde2n/tFHH3HmmWfSr18/OnToQM+ePTn44IN55ZVXsub1\n7duXoUOHArD77rtnva4VMy6/B/jy/51NmTKFY445hl69etGxY0e23nprRowYkff7WFVVxeDBg1l7\n7bXp2LEjm2yyCddddx2TJk3K+V6myRXgKYsxTgN+mHYOSZIkSZIktUwTJ05ku+22Y6ONNuLYY49l\n3rx5dO3aldtuu40HH3yQ3Xbbjb322ovFixczbtw4rrvuOh555BFeeuklVl555ZzzjRgxguHDh7P/\n/vtz2mmn8c477zBy5EheeeUV3nnnHXr16rV07gcffMAOO+zAjBkz2H///dlyyy2ZMGEChx56KPvv\nv3/evCNGjOCII44gxsiRRx7Juuuuy9ixY/n973/P8OHDee655/Luvf2LX/yC0aNHc9BBB7HPPvvw\nn//8hwsvvJBvv/2WHj16cN5553HooYeyyy678MQTT3DLLbewePFifv/73zfq+9m3b1+GDBnC9ddf\nD8DgwYOXjm255ZZLn5933nlUVFSw3XbbsdZaa1FVVcXTTz/NWWedxSuvvMI999yT9/xnnXUWY8aM\nYdCgQRxwwAG0adMGgHfffZeddtqJmTNnMmjQIDbffHMmTZrEYYcdxgEHHJD3XOPGjWOfffZh5syZ\n7Lvvvhx++OF8+eWXDBs2jJ133pkHH3xw6bGDBw9m2LBhPPPMM5xwwgmN3t986tSpbLvttvTr14/j\njjuOmTNnct9993HIIYfw5JNPsvvuuy+dO3/+fPbYYw/GjRtH//79OfbYY6mqquKKK65gzJgxjbpu\nwcUY/SrBL2DsVlttFSVJkiRJklqTd955J77zzjtpx2gRJk+eHIEIxPPPPz9nfMqUKXHRokU5/bff\nfnsE4tVXX53Vf8cdd0QgtmnTJj755JNZY+edd14E4q9//eus/r333jsC8frrr8/qHzZs2NJsd9xx\nx9L+2bNnx549e8aKior43//+N+uYq6++OgJx7733zuo/4YQTIhDXXXfd+NFHHy3t/+qrr2LPnj1j\np06dYq9evbL+u5g/f37cdNNNY/v27eNnn32W8z1oiHXXXTeuu+66tY5PmDAhp2/x4sXx+OOPj0B8\n8cUX876ONddcM06aNCnn2D322CMC8dZbb83qHzlyZN7v5cKFC+P6668fO3ToEEePHp11zMcffxzX\nXHPN2Lt37zh//vyl/UOGDIlAHDVqVN7XtCTj5MmTl/Yt/99ZZWVl1vxHH300AnH//ffP6r/00ksj\nEI855phYXV29tP/DDz+MvXr1ikA84YQT8mZYXkPf71tttVUExsYm1EldAS5JkiRJkqTWo7Jb2gka\nrrIqkdOsvvrqDBkyJKd/3XXXzTv/Rz/6EWeffTaPPfYYv/rVr3LGjznmGPbcc8+svlNOOYWrr76a\nl19+eWnfRx99xBNPPMF6663HmWeemTX/kEMOYbfdduOZZ57J6h8+fDgzZszg//2//8cuu+ySNXbO\nOefwhz/8gSeeeIIPP/yQddZZJ2v84osvzro5Y/fu3Tn44IO54447OOecc9h0002XjnXo0IGjjz6a\nyspKxo8fz2qrrZb3e9Ec66+/fk5fRUUFZ511FnfffTePPfYY2223Xc6cc889l/XWWy+rb9q0aTz9\n9NNssMEGnHrqqVlj+++/P3vttRdPPvlkVv/DDz/MxIkT+cUvfsFuu+2WNbbmmmty7rnnMnjwYJ56\n6qlaV5A3xrrrrstFF12U1bfvvvuyzjrrZP13AXDXXXdRUVHBVVddlbXFSp8+fRg8eHDOedJkAVyS\nJEmSJElqwbbYYgs6dOiQ079w4UL++Mc/cu+99/LOO+9QVVWVtTf1xx9/nPd8W2+9dU5fnz59ALL2\nDn/ttdcA2HnnnZdu47G8gQMH5hTAx40bB8Aee+yRM79t27bsuuuuTJkyhddeey2nAJ4v15prrgnA\ngAEDcsaWFMs/+uijnLEkzJgxg2uuuYaRI0cyadIk5syZkzVe2/d32223zel7/fXXAdhhhx2oqMi9\nLePOO++cUwB/4YUXgMzWJJWVlTnHfPDBBwCMHz8+kQL4lltumffPuU+fPkuzAHz99ddMnDiRPn36\n5N1mZeedd252liRZAJckSZIkSZJasN69e+ftP/roo3nwwQfp168fhxxyCL17915aKL/++utZsGBB\n3uO6d++e09e2baZMuHjx4qV9VVWZFeyrr756g3MtOWaNNdbIe8yS/lmzZuWMdeuWu7p/Sa66xhYu\nXJj3Ws0xa9YsttlmGyZPnsy2227L8ccfT48ePWjbti2zZs3ihhtuqPX7W9f3pbbvZb7+GTNmAJmb\nf9blm2++qXO8ofL9dwGZ7/Pyv1j5+uuvgca9ljRZAJckSZIkSVLrkdC2Iq3J8ltMLPHqq6/y4IMP\nstdeezFy5EjatWu3dKy6uprf/OY3zb7ukqLzZ599lnf8008/rfWYfGMA06dPz5rXUt1+++1MnjyZ\nIUOG5Ky+fuGFF7jhhhtqPTbfn1fXrl2B2r+X+fqXfI+GDx/OwQcf3NDoBdeU15Km3PX2kiRJkiRJ\nklq0CRMmAHDwwQdnFb8BXn75ZebNm9fsa/Tv3x+AZ599Nmtl+BKjR4+u9Zh8Y4sWLeLZZ58FYKut\ntmp2vuZq06ZN3tcFy76/RxxxRM7Yitu+NMSS78sLL7yQtZp6iSXfl+Vtv/32AIwZM6bB11myhUlt\nrysJXbt2pV+/fnz88cdMmTIlZzzfa0mTBXBJkiRJkiSplVmy9/KKhebPP/+cM844I5FrrL322uy9\n995MnjyZm2++OWts+PDheQvBhx56KD169OAf//gHL774YtbY9ddfz6RJk9hrr71y9v9OQ8+ePfni\niy/y/rKgtu/va6+9xlVXXdXoa/Xp04eBAwcyYcIE/vjHP2aNPfroozn7f0PmRqPrr78+t9xyCyNH\njsx73hdeeIG5c+cubffs2ROADz/8sNEZG+P444+nurqa888/nxjj0v5p06Zx/fXXF/TajeUWKJIk\nSZIkSVIrs80227DTTjvxwAMPsOOOO7Lzzjvz2Wef8cgjj7DxxhsvvXlkc91yyy3ssMMODB48mMcf\nf5wtttiCCRMm8OCDD3LQQQfx0EMPZc3v0qULf/nLXzjqqKPYbbfdOOqoo1hnnXUYO3Ysjz/+OL17\n984pAKdlzz335JVXXmG//fZj1113pUOHDmyxxRYcdNBBHH/88VxzzTUMHjyYUaNGseGGG/LBBx8w\nYsQIDj/8cO67775GX++WW25hp5124vTTT2fkyJFsvvnmTJo0ifvvv59DDjmE4cOHZ90gs127djzw\nwAPsu+++DBo0iB133JEtt9ySTp06MW3aNF555RUmTZrE9OnT6dSpEwC77747FRUVnH/++bz11lus\nssoqAFx00UXJfNNqnHvuuQwbNox7772X9957j3322Yeqqir++c9/suuuuzJs2LC8N/tMQ8tIIUmS\nJEmSJKnB2rRpw3/+8x9+8pOf8Mknn3DjjTfy7LPP8uMf/5jHHnssZ1uUptpwww158cUXOeKII3ju\nuee44YYbmDZtGsOGDePwww/Pe8whhxzCc889xwEHHMBjjz3Gtddey/jx4znttNMYO3Ys/fr1SyRb\nc1100UWcdtppTJw4kauuuoqLL76Y+++/H4A111yTMWPGMGjQIJ599lluvvlmpk6dyq233srVV1/d\npOttttlmvPDCCxx22GGMGTOG66+/nilTpvDggw+y8847A8v2115i880354033uBXv/oVVVVV3HHH\nHfz+979n7Nix9O/fn3vuuYdevXotnb/pppty11130bt3b2699VYuvvhiLr744iZ+h2q30korMWrU\nKH7605/y6aef8rvf/Y5Ro0ZxwQUXcP755+d9LWkJyy9RV+kIIYzdaquttho7dmzaUSRJkiRJkhps\n/PjxQKaQJ5WLY489lr///e+8++67bLzxxmnHaZY//elPnHLKKfzhD3/g1FNPrXNuQ9/vAwYMYNy4\nceNijAMam8cV4JIkSZIkSZJUYNXV1Xz66ac5/U899RT33Xcfm222Wasqfn/yySc5fdOmTeOyyy6j\nbdu2HHjggSmkyuUe4JIkSZIkSZJUYN9++y19+vRh9913Z5NNNqFt27a8/fbbPPHEE7Rv355bbrkl\n7YiNcsQRR7Bw4UIGDBhA9+7dmTJlCiNGjGDu3LlcddVVrLXWWmlHBCyAS5IkSZIkSSoBlZWVDZp3\n6KGHsuWWWxY0Sz7t2rXjtNNO4+mnn+all15i7ty59OrVi6OOOorzzjuP/v37Fz1Tcxx33HHcc889\n3H///VRVVdGlSxe22247zjzzzFr3h0+DBXBJkiRJkiRJrd7QoUMbNK9v376pFMDbtGnDTTfdVPTr\nFsrpp5/O6aefnnaMelkAlyRJkiRJktTqxRjTjqAWyJtgSpIkSZIkSZJKkgVwSZIkSZIkSVJJsgAu\nSZIkSZIkSSqqYm1ZYwFckiRJkiRJLUYIAYDq6uqUk0gqpCUF8CXv+UKxAC5JkiRJkqQWo0OHDgDM\nmTMn5SSSCmnJe3zJe75QLIBLkiRJkiSpxVh55ZUB+PTTT5k9ezbV1dVF2ypBUmHFGKmurmb27Nl8\n+umnwLL3fKG0LejZJUmSJEmSpEbo0aMHc+bMYe7cuXz00Udpx5FUQJ06daJHjx4FvYYFcEmSJEmS\nJLUYFRUV9OnTh5kzZzJ79mwWLFjgCnCphIQQ6NChAyuvvDI9evSgoqKwm5RYAJckSZIkSVKLUlFR\nQa9evejVq1faUSS1cu4BLkmSJEmSJEkqSRbAJUmSJEmSJEklyQK4JEmSJEmSJKkkWQCXJEmSJEmS\nJJUkC+CSJEmSJEmSpJJkAVySJEmSJEmSVJIsgEuSJEmSJEmSSlKIMaadQQUQQpix0kor9dh0003T\njiJJkiRJkiRJTTZ+/HjmzZs3M8bYs7HHWgAvUSGEyUBXYErKUYppk5rHd1NNISkJvp+l0uB7WSoN\nvpel0uH7WSoN5fhe7gt8HWNcr7EHWgBXyQghjAWIMQ5IO4uk5vH9LJUG38tSafC9LJUO389SafC9\n3DjuAS5JkiRJkiRJKkkWwCVJkiRJkiRJJckCuCRJkiRJkiSpJFkAlyRJkiRJkiSVJAvgkiRJkiRJ\nkqSSFGKMaWeQJEmSJEmSJClxrgCXJEmSJEmSJJUkC+CSJEmSJEmSpJJkAVySJEmSJEmSVJIsgEuS\nJEmSJEmSSpIFcEmSJEmSJElSSbIALkmSJEmSJEkqSRbAJUmSJEmSJEklyQK4WrQQwtohhL+EED4J\nISwIIUwJIVwfQlgljfNIaprmvgdDCD1DCD8OITwYQpgQQpgXQqgKITwbQjgphODfZ1KRFOLv1BDC\ncSGEWPP14yTzSsovyfdyCGGXEML9IYTpNeeaHkJ4PIRwQCGyS1omwX8zD6p5335U87P2pBDCv0II\nOxQqu6RlQghHhhBuCiGMCSF8XfNz8V+beC5rYCsIMca0M0h5hRDWB54HVgOGA+8C2wK7A+8BO8UY\nZxTrPJKaJon3YAjhNOD3wHRgFPAhsDpwONANuB84KvqXmlRQhfg7NYTQB/gf0AboApwcY7w9ydyS\nsiX5Xg4hXARcBnwJjCDzd3UvoD8wKsZ4buIvQBKQ6L+Zfw2cC8wAhpF5P28AHAy0BY6PMTapECep\nYUIIrwNbAN8AHwGbAH+LMf6gkeexBpaHBXC1WCGEx4B9gJ/FGG9arv864OfAH2OMpxXrPJKaJon3\nYAhhD6Az8HCMsXq5/t7Ay0Af4MgY4/0FeAmSaiT9d2oIIQBPAOsBDwC/wAK4VHAJ/px9FPBP4Eng\n8Bjj7BXG28UYFyYaXtJSCf2c3Rv4GPgC2DzG+PlyY7sDTwOTY4z9CvASJNWoeb99BEwAdiOz8Ksp\nBXBrYHlYAFeLFELoB0wEpgDrr1DwWpnMypIArBZjnFPo80hqmmK8B0MIFwBXADfHGH/a7NCS8irE\n+zmEcBbwO2AgsAcwBAvgUkEl+HN2BZl/pK8O9I0xflHI3JKyJfhe3g54EfhPjPGQPONfk6kdrZzs\nK5BUmxDCQJpQALcGVjv3TFVLtUfN4+PLv2EBalaWPAd0ArYv0nkkNU0x3oNLVpYtasY5JNUv0fdz\nCGFT4Grghhjjf5MMKqlOSb2XdyTz6Y2RwFc1+wf/KoRwlnsGS0WR1Hv5A+BbYNsQQq/lB0IIuwIr\nk/mUh6SWzxpYLSyAq6XauObx/VrGP6h53KhI55HUNAV9D4YQ2gLH1zQfbco5JDVYYu/nmvfuPWT2\n87+g+dEkNUJS7+Vtah4/A8aR2f/7auB64PkQwjMhhFWbkVNS3RJ5L8cYZwK/IvNpjndCCLeFEK4K\nIfwTeJzMVmWnJpBXUuFZA6tF27QDSLXoVvNYVcv4kv7uRTqPpKYp9HvwauC7wMgY42NNPIekhkny\n/XwJmRvk7RxjnNfMXJIaJ6n38mo1j6cBk4G9gJeAdYHfAvsC/yKzxZGk5CX293KM8foQwhTgL8DJ\nyw1NAO5cfl9wSS2aNbBauAJcrVWoeWzuJvZJnUdS0zT5PRhC+BlwDpm7Wh+XZChJTdKg93MIYVsy\nq75/G2N8oeCpJDVWQ/9ubrPc/CNjjE/FGL+JMb4NHEbmRl67uR2KlJoG/5wdQjgX+DdwJ7A+mZvP\nDwAmAX8LIfymQBklFVfZ1sAsgKulWvJbqW61jHddYV6hzyOpaQryHgwhnAHcALwD7F7z0U1JhdXs\n9/NyW5+8D1ycXDRJjZDU381f1TxOijG+sfxAzSc7lnwya9tGJ5TUEIm8l2tutvdrMjfBPDvGOCnG\nODfGOI7ML7M+Bs6pubmepJbNGlgtLICrpXqv5rG2fYk2rHmsbV+jpM8jqWkSfw+GEAYDNwNvkSl+\nf9rkdJIaI4n3c5ea4zcF5ocQ4pIvYEjNnD/V9F3f3MCS8kr65+xZtYwvKZCv1LBYkhopqffygTWP\no1YciDHOBV4mUzvq39iAkorOGlgt3ANcLdWSv3z3CSFULH/32hDCysBOwDzgxSKdR1LTJPoeDCH8\nisy+368De8cYv0w2rqQ6JPF+XgD8uZaxrcj84/pZMj+8uz2KVBhJ/d38X2ARsGEIoX2M8dsVxr9b\n8zil+ZEl5ZHUe7lDzWNtN61d0r/ie1xSy2MNrBauAFeLFGOcSOaO032BM1YYHkpmT7K7Y4xzAEII\n7UIIm4QQ1m/OeSQlK6n3cs3YxWSK32OBPS1+S8WVxPs5xjgvxvjjfF/Af2qm3VXTd1/BX5RUhhL8\nOftL4D4yH7O+ZPmxEMLeZG6CWQU8WoCXIZW9BH/OHlPzeEoIYa3lB0II+5MpmM0Hnk/2FUhqKmtg\njRdiLLt9z9VK1LyRnydzh/nhwHhgO2B3Mh/X2DHGOKNmbl8yd5+fGmPs29TzSEpeEu/lEMIJZG7K\nsxi4ifx7lk2JMd5ZoJchieT+bq7l3JVktkE5OcZ4ewHiS6qR4M/ZqwHPARuQKaK9DKxLZt/gCHw/\nxvivwr8iqTwl9HN2BZk9+/cCZgMPAp+S2a7sQDI3zRscY7yhKC9KKlMhhEOBQ2uavcn8InkSy35J\n9WWM8Rc1c/tiDaxR3AJFLVaMcWIIYWvgUmA/4ABgOnAjMLShN71L6jySmiah9+B6NY9tgMG1zHmG\nTJFcUoH4d6pUGhL8OfvzEMJ2wEVkit7bkymgPQxcFWMsu49YS8WUxHs5xlgdQjiAzGrRY8i8lzsB\nM4GRwI0xxscL9BIkLbMlcMIKff1qvgCmAr+o7yT+vJ6fK8AlSZIkSZIkSSXJPcAlSZIkSZIkSSXJ\nArgkSZIkSZIkqSRZAJckSZIkSZIklSQL4JIkSZIkSZKkkmQBXJIkSZIkSZJUkiyAS5IkSZIkSZJK\nkgVwSZIkSZIkSVJJsgAuSZIkSZIkSSpJFsAlSZIkSZIkSSXJArgkSZIkSZIkqSRZAJckSZIkSZIk\n5QghHBlCuCmEMCaE8HUIIYYQ/lqA63wvhHB3CGFaCGFBCOHzEMIzIYTjm3tuC+CSJElqNUIIo0MI\nMe0cSQohbBhCeDCE8GnNPyhmpZ2pNQoh3Fnz/eubdpbl5ftvNoQwsCZrZTPPPSWE/9/enUfbUVV5\nHP/+ZEZoptAooIwiEZAh0gtE7ESIDCKRUZAp0AgoMsuQBkwwLGSQudulOBBggSJhJtA0gkFkshlC\nGmyVKciggoxCmNn9xz6VFDd137sv970AL7/PWqziVZ176ty6VVl19921j6b1of2KZb8Tutnv7JA0\nrux7+Jzet5mZmXXlGOBbwDrAkwOxA0mjgXuBrwC3AKcCEwEBW3bb/7zddmBmZmZmHyy1YNyfgU9G\nxGsNbaYBKwDzRcRbc3B4cxVJ8wBXAKsCFwBPALN8HrX2fQ3+7xkRE2Z3fN0ogehHgfMiYvR7MYae\nlCDwHryHx8jMzMzsA+AQ8h71IeBfgV/3Z+eSNgB+AtwPbB4Rf23ZPl+3+3AA3MzMzGzu9XHgYODE\n93gcc7OVgE8BP46IfTpof1zDuoOBxYAzgRdatk3pYmw2cH4HDAX+3mU/m/TDWMzMzMzaiogZAW9J\nHb1G0s7APmTW+EJkUsSFwCkR8XpL85OBeYBdW4PfZf9vztbAaxwANzMzM5s7PQ8EMEbSTyKi20Cc\nzZ5ly/KpThpHxLjWdeWR0cWAMyJiWn8NzAZOREwH/tAP/TzcD8MxMzMz6zeSfgrsRWaNX0YmaGwA\njAc2kTSyesJU0vLAxsBdwAOSRgDDyO8pU4BfR8Q73Y7JNcDNzMzM5k7TyZvQfwLGdvKC3uoWN9Uj\nljS6vGa0pJFl8pyXJT0j6VxJi5d260q6RtLzZftVPdVylrSApOMlPVomyXlY0lhJ87dpv3qpEV1N\nqvM3SRdJ+mRD26qW9MqSDpA0VdKrkiZ3eJyGSbq0TNzzuqTHJP1A0kdb2gVwc/lzbNln13Wha/1P\nLv3NL+k7kv5YxjOh1mZ5Sf8h6ZGy7dly7Ndv6G/Z0s+tpV75G5KeKsdxaEvbcWSmD8AetfcWJWBf\nb7uZpGsl/b32WZ5SnRsN49i0nEevSHpO0hWSVu/uaL2r/+q4zSvp3yU9WMb1uKSTejjHdpJ0dzlX\nnpZ0gaRl27Sd5VqS9IdyTIe0ec1R5TX719Y11gCXtKik0yQ9Iem10vehtPn+px5q69ev4Zb1IySd\nI+n3ygmxXpV0f7kOF2zqq03/G0u6uoz19XJu3SGpo3+XzMzM7P2j3C/sBVwOrBYR/xYRh0XERuST\njMOB/Wsvqe45HwRuKv+dAnwf+BUwRdKq3Y7LGeBmZmZmc6//JCe02VfS2RHxpwHc19bAVsA1wA+B\nzwKjgZUkHQXcSE5481NgLeDLwCqS1mqT9fFL8oZ5IvAmMAoYB3xG0tYRMSOYJ2lzMvtkPuBqsn7h\n8sC2wJckjYiIexr2cSaZkTIJuBZ4u7c3KWkr4FJywp6JwGNkFss3gFGSNqplaR8HrEjWob4ZmFzW\nT6Z/XUoeq+vIeuNPl7GuB/w3sCRwPXmMhpCTD/1W0jYRcW2tn88DR5F1Hy8FXgY+AWwPbF3e2321\n97A4cBBwX9lvZUr1P5K+Qx6H58hz42ng08C3gS0lbRgRL9Xabw9cDLxRln8BPgfcDkydjWPTk4vI\nz/864CVyAqYjgH8G9qw3lHQIcBqZ4XR+WW4G3Aa82OH+zgNOAHYGzm7Yvjv5vn/RUyeSFiCvp/XJ\nY38h+VkcS9bt7C9HAquT73ESsCCwEXkdDpe0aUT0eM2Ua3MSeXyvIifWWpIsD/NNmkv+mJmZ2fvX\nQcBbwF4R8WrLtvHkd49dyPtsyPsqgB3J0nDbkvcxS5NJOrsBk8p3gjdmd1AOgJuZmZnNpSLizRJ8\nvoSsA77tAO5ua2CTiLgZQNKHyKDrpmRweZ+IuLBqrJmPTn4ZuLKhv6HAGhHxfGl/NBmY3QrYlZxQ\nEklLAD8nM94/HxG/r+1jDeBOctKd9Rr2sR6wbkQ82rBtFpIWASaQ99jDI+KW2rYjyWN8DvBFyHIm\nkoaTAfDJTeVN+skKwJr1MjeS5iV/RFgEGFF9LmXbssD/AD+VtGKtTuNNwDIR8Y9655LWBm4l398W\nABExuWQmHwRMaVO6ZQQZ4Lwd2DIiXqhtGw2cW7YfUtYtAvwIeAfYOCLuqrU/nayF3p9WIc+x58o+\njiYDyrtLGlPVqFQ+qXAiWVZoveoHDkljyGur0+vqfOB48nx4VwC8ZOQPBS6LiGd76ecwMvh9GbBD\n9QOSpBOBuzscSye+CTxa/7Gp7Gc8cAz5w8jFvfTxdTIrfXjtx5Oqn8ZMeDMzM3t/krQwsDYZyD5Y\nzfXCXyfvaSrz1JZ7R8Q15e+XJO1R2n4G2I68p58tLoFiZmZmNheLiIlkAHIbSZ8bwF39vB5kLUG5\nC8qf99eD38X5ZblOm/7GV8Hv0t9rwJjy5161druT2a9j68Hv8poHgB8D60r6VMM+Tu40+F2MApYC\nLq4Hv4tTgWnASEkf70Of/eHYhhrvXyIDvGfXPxeAiHiKnIzoI9QmWYyIp1uD32X9fWRwfISk+fow\nrgPL8uv14HfpcwKZKb5LbfUoMjv4onrwuxhH55nWnTqyCn6XMb1CZlN/iPwiVtkFmJ88ltNq7d8B\nDicD9r2KiCfJjKdh5ceZuj3K8rwOutqz7POI+tMT5Vw+q5OxdCIiHmkNfhdnlOVmfeiuNUMMz0tg\nZmb2gbME+RRklb3d9N+yZAJGpbqff51Mipmh3GdUiTD/0s3AnAFuZmZmZoeRZQxOlbRBm6BWt1oD\nljBz4semrNQny3L5Nv3d3LDuFvKRy3Vr6zYsy7XVXFt7tbIcCvy+Zdvv2uy7nSqL/KbWDRHxlqTf\nkCVP1gX+3Me+u9H0PqrjskKb4/KJshxK7cuIpC8B+5EB4CHM+n1iCFmWpBMbkuVrdpC0Q8P2+YGl\nJS1Vsp6r4zvLZx8RL0qaQv+W+Gg6Zx8vyyVq63oa1yOSHiez8DsxARhJBryPACg1x3cCnqHli2Er\nSYsCqwKPt5kgczId1vzvjaQPkxn+25DX0aLkl97Kch10cyGZIX+npIvJpzhujYgn+mOMZmZmNkdV\nyQj3RkTT05VN/liW/2hT9rAKkC/UzcAcADczMzOby0XE7ZImkiULdqT3sgWzoyk7960OtrXLKP5b\n64qIeFvSs8ysJQiZkQ1ZaqEnizSs+2svr2m1WFm2CwBX6xfvY7/danof1XFpCjzXzTgukg4k6zU+\nD9xABvGnA0HWDV8bWKAP41qK/D7SW0B2EeBZZh7fWT77oq+fV49as9KL6rycp7auk3F1GgC/nKyH\nvWsps/I2WdZnKeCMiHirx1fPoWNUMv1vIrOx7if/zXiG/EED8jPt9VyIiMtK3fzDyCc39i393w2M\niYgb+mO8ZmZmNvAi4mVJDwBrSFqy/iRdD6aSJVOGSFomIlrvYdYsy2ndjM0lUMzMzMwMcnLDN4Hv\nlYzTJlVWRrskisXarB8Iy7SukDQPGSh8qba6Cq6vHRHq4b+m0hJ9zYSv9vWRNts/2tJujmiT0V+N\nYVQvx+U4mFEz/DgygLpGRHw1Ig6PiLGlvne7gGtPXgSe72X/iojHWsY8y2dftDvuA63fxlUmi/ol\nea6MLKv7Uv5kdsdS1QlvurYXb1g3igx+nxcRa0XEPhFxdDkXftTBOGeIiEkR8QUyq34T4HRgDeCa\nNqWJzMzM7P3rNPIpvp9JWrx1o6QlykTsQD4lycx7h5PLPEFV27WA0WQCwsRuBuUAuJmZmZlRyiX8\nAFgJOKBNs+oRxI+1bpC0KnM2s7mp1MXGZHD+3tq6O2rbBlq13+GtG0pgsaqxfs8cGEtv+npchpCf\n720R8a4M9zI5ZdNjrm+X5TwN26oxLNFQ77qd6rjN8tlLWoz29eIHWk/jWpmG66UXE8pyjzIR5BbA\n1IiY0tsLS432h4DlJK3S0GR4m5e2vbZ5d73zyqpleWnDttkqQxMRr0TETRFxKHAC+eV5i9npy8zM\nzPqPpK9ImiBpApk0A7BhtU7S96u2EfEz8jvFKOBhSRdJOlHSOZJuIJMp9mnZxQnkfeHuwF2STpN0\nATlZ/YLkvCwPdfMeHAA3MzMzs8p3gReAo2kuCfIHMrt6lKQZZUYkLUQ/Tq7XoWMlzajDLGlB4Hvl\nz3Nr7c4l39NYSbNMniPpQ5KG99OYrgCeA3aWtEHLtoOBlYFfRcScrP/dzpXAw8D+krZsaiBpQ0kL\nlz+fJsudDCsB76rNfGRZlCENXTxPZtG3m/Tz9LL8saRlG/b/4ZbjeGXp82uSWoOy45izTyDUXUg+\nPXGApBWrlSWD6RT6+J0rIm4FHiS/OH6DLAM0oQ9dnFv2eVJLFtVKzJx4tFVVJ/5dpYIkbQLs3NB+\nWlkOb2m/MnBSpwOVtEn596NVlcE+vdO+zMzMbMCsQz6RtgczJ7leubZu+3rjiNgf+DJwO7ApcCiw\nNXmvdgozJ8yu2k8nnwI7DlgY2L+0vw3YMiJO6/YNuAa4mZmZmQEQEc9JOgE4uc32NyWdCRwL3Cvp\ncvJ+ciQ5oeVTTa8bIP8HPFBql79JBgtXASYBF9TG/Kyk7cnayndIuhF4gCz58HFyIsalyOySrpS6\nh3sBlwA3S7qErJM9DPgimfGyb7f76Q/ls9wWuB6YJOk2YAoZcPwYsD75xeajwPSIeEfSWWTWz/9K\nupLM0B0BLElOXjiiZR8vS7oT2FjShcCfyKzwqyJiakTcKOko8oeLByVdCzxK/viyAplJ/Ftg81p/\n+5D1pm8pkyb+hcysXxP4DfD5/j9aPYuIaeV9nEpeFxeTpUg2I7PmpwKf7mO35wPjyWvtLeCiPrz2\nVLIm+3bAPZKuJ79wfpU8Rls3vOZc4HBgjKS1yQlhVyMzsC8vfdVdTWaaH1oeT76XvJ62Iq/Bdj96\nNI11RUmTyaD6G+T18gXgMeAXHfZjZmZmA6SUOBvXx9dcA1zTh/bTyz76tJ9OOQPczMzMzOrOoudJ\nZsYCY4DXyMcXtyTLIGzGzAnw5oQdgZ+R2SXfIu9rxwHbtda8jogbyQDkD4AVgf2Avcmg6U3ATv01\nqIi4EtgIuJY8Jt8GhgI/BIZFxCP9ta9uRcRUcuLKk8gA6Z5kxvEwMqC5GzkpUeVYcrLCV8lA/rbA\nXWQt6HZZ7buRAdHNyXNnPLVyKRFxEhm0nkQet4PJiTmXA84BjmkZ88TS193kObAfmXW/IRk8f0+U\nzKSvlTGMJid0vB/4LDPLi/TF+eSPNPMB/9UwIVRPY3mdzLY6HVgaOIjM1D4eOKTNa54mf3C4jvw8\nvkGeEyNp+PIaEa+QQeqLyHrdB5LX2Hhg107HSj7yfF3pY2/y81ymrF8/Imbn2JmZmZm9i5rnxDEz\nMzMzMzMzMzMz+2BzBriZmZmZmZmZmZmZDUoOgJuZmZmZmZmZmZnZoOQAuJmZmZmZmZmZmZkNSg6A\nm5mZmZmZmZmZmdmg5AC4mZmZmZmZmZmZmQ1KDoCbmZmZmZmZmZmZ2aDkALiZmZmZmZmZmZmZDUoO\ngJuZnPJ2wQAAAJNJREFUmZmZmZmZmZnZoOQAuJmZmZmZmZmZmZkNSg6Am5mZmZmZmZmZmdmg5AC4\nmZmZmZmZmZmZmQ1KDoCbmZmZmZmZmZmZ2aDkALiZmZmZmZmZmZmZDUoOgJuZmZmZmZmZmZnZoOQA\nuJmZmZmZmZmZmZkNSg6Am5mZmZmZmZmZmdmg5AC4mZmZmZmZmZmZmQ1K/w8CgQx9JIgMwAAAAABJ\nRU5ErkJggg==\n",
+ "text/plain": [
+ "\u003cFigure size 1200x800 with 1 Axes\u003e"
+ ]
+ },
+ "metadata": {
+ "image/png": {
+ "height": 494,
+ "width": 736
+ }
+ },
+ "output_type": "display_data"
+ }
+ ],
+ "source": [
+ "plt.figure(figsize=(12, 8))\n",
+ "num_individuals = range(len(uplift_predictions))\n",
+ "plt.plot(num_individuals, uplift_incremental_visits, label=\"uplift_targeting\")\n",
+ "plt.plot(num_individuals, random_incremental_visits, label=\"random_targeting\")\n",
+ "plt.title(\"Absolute Cumulative Uplift Curve\")\n",
+ "plt.ylabel(\"Cumulative Incremental Visits\")\n",
+ "plt.xlabel(\"Number of Treated Individuals\")\n",
+ "plt.legend(loc=\"lower right\")\n",
+ "plt.rc('font', size=14)\n",
+ "plt.show()"
+ ]
+ },
+ {
+ "cell_type": "markdown",
+ "metadata": {
+ "id": "JYXC8LK8RKLQ"
+ },
+ "source": [
+ "The uplift curve demonstrates that targeting customers with the uplift model can lead to significantly more visits than random selection. For example, uplift targeting could bring in 5000 extra visits when treating 100k individuals, while random selection would likely yield less than 2000 extra visits.\n",
+ "\n",
+ "As mentioned earlier, the uplift curve is just one evaluation method which can also be turned into a more rigurous, quantifiable metric by measuring the area under the uplift curve (AUUC). See this paper [[4](https://arxiv.org/pdf/2002.05897.pdf)] for a more comprehensive overview of the various metrics used in uplift modeling."
+ ]
+ }
+ ],
+ "metadata": {
+ "colab": {
+ "last_runtime": {
+ "build_target": "//learning/grp/tools/ml_python:ml_notebook",
+ "kind": "private"
+ },
+ "provenance": [
+ {
+ "file_id": "1QmPiH1YvWKUk-MkUsJvTlfk3gmhlcXlQ",
+ "timestamp": 1711525016304
+ }
+ ],
+ "toc_visible": true
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
diff --git a/official/recommendation/uplift/utils.py b/official/recommendation/uplift/utils.py
new file mode 100644
index 00000000000..d1b515089d8
--- /dev/null
+++ b/official/recommendation/uplift/utils.py
@@ -0,0 +1,98 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Common utilities for the Keras uplift library."""
+
+from typing import Tuple
+import tensorflow as tf, tf_keras
+
+
+def expand_to_match_rank(a: tf.Tensor, b: tf.Tensor) -> tf.Tensor:
+ """Expands tensor a to match the rank of tensor b.
+
+ Args:
+ a: a `tf.Tensor` of shape (D0, D1, ..., Dn).
+ b: a `tf.Tensor` of shape (D0, D1, ..., Dn, Dn+1, ... Dn+m).
+
+ Returns:
+ A `tf.Tensor` of shape (D0, D1, ..., DN, 1, ..., 1) if b has a higher rank
+ than a, otherwise a `tf.Tensor` of shape (D0, D1, ..., Dn)
+ """
+ rank_deficit = b.shape.rank - a.shape.rank
+ for _ in range(rank_deficit):
+ a = tf.expand_dims(a, axis=-1)
+ return a
+
+
+def split_by_treatment(
+ values: tf.Tensor, is_treatment: tf.Tensor
+) -> Tuple[tf.Tensor, tf.Tensor]:
+ """Splits a tensor into control and treatment tensors.
+
+ Args:
+ values: a `tf.Tensor` of shape (D0, D1, ..., DN).
+ is_treatment: a `tf.Tensor` of shape (D0,) or (D0, 1) castable to boolean
+ indicating if the example belongs to the treatment group (True) or control
+ group (False).
+
+ Returns:
+ A tuple with control and treatment values sliced by the is_treatment tensor.
+ """
+ if is_treatment.shape.rank > 2 or (
+ is_treatment.shape == 2 and is_treatment.shape[1] != 1 # pyrefly: ignore[unsupported-operation]
+ ):
+ raise ValueError(
+ "is_treatment tensor must be a tensor of shape (D0,) (D0, 1) but got a"
+ f" tensor of shape {is_treatment.shape} instead."
+ )
+
+ if values.shape[0] != is_treatment.shape[0]:
+ raise ValueError(
+ "values and is_treatment must be tensors of shapes (D0, D1, ..., DN)"
+ f" and (D0, 1) (or (D0,)), but got tensors of shapes {values.shape} and"
+ f" {is_treatment.shape} respectively."
+ )
+
+ if is_treatment.dtype == tf.string:
+ raise ValueError(
+ "is_treatment must be a tensor castable to boolean but got tensor"
+ f" {is_treatment} of dtype {is_treatment.dtype} instead."
+ )
+
+ # Assert is_treatment tensor containss only 0 or 1 values.
+ if is_treatment.dtype != tf.bool:
+ is_treatment_float = tf.cast(is_treatment, tf.float32)
+ tf.debugging.assert_equal(
+ tf.reduce_all(
+ tf.logical_or(is_treatment_float == 1.0, is_treatment_float == 0.0)
+ ),
+ tf.convert_to_tensor(True),
+ message=(
+ "When is_treatment is not a boolean tensor all of its values must"
+ f" either be 0 or 1, but got tensor {is_treatment} instead."
+ ),
+ )
+
+ if is_treatment.shape.rank == 1:
+ is_treatment = tf.expand_dims(is_treatment, axis=1)
+
+ is_treatment = tf.cast(is_treatment, tf.bool)
+
+ control_indices = tf.cast(tf.where(~is_treatment)[:, 0], dtype=tf.int32)
+ treatment_indices = tf.cast(tf.where(is_treatment)[:, 0], dtype=tf.int32)
+
+ control_values = tf.gather(values, control_indices)
+ treatment_values = tf.gather(values, treatment_indices)
+
+ return control_values, treatment_values
diff --git a/official/recommendation/uplift/utils_test.py b/official/recommendation/uplift/utils_test.py
new file mode 100644
index 00000000000..a2186f4c98f
--- /dev/null
+++ b/official/recommendation/uplift/utils_test.py
@@ -0,0 +1,136 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for utils."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+from official.recommendation.uplift import utils
+
+
+class UtilsTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "same_rank",
+ "a": tf.zeros((3, 1)),
+ "b": tf.zeros((3, 2)),
+ "expected_output": tf.zeros((3, 1)),
+ },
+ {
+ "testcase_name": "higher_rank",
+ "a": tf.zeros((3, 1)),
+ "b": tf.zeros((3,)),
+ "expected_output": tf.zeros((3, 1)),
+ },
+ {
+ "testcase_name": "one_less_rank",
+ "a": tf.zeros((3,)),
+ "b": tf.zeros((3, 2)),
+ "expected_output": tf.zeros((3, 1)),
+ },
+ {
+ "testcase_name": "multiple_rank_difference",
+ "a": tf.zeros((3, 4)),
+ "b": tf.zeros((3, 1, 2, 4)),
+ "expected_output": tf.zeros((3, 4, 1, 1)),
+ },
+ )
+ def test_expand_to_match_rank(self, a, b, expected_output):
+ self.assertAllEqual(expected_output, utils.expand_to_match_rank(a, b))
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "treatment_only",
+ "values": tf.constant([1.5, 2, 3]),
+ "is_treatment": tf.ones((3, 1)),
+ "expected_control_values": tf.zeros((0,)),
+ "expected_treatment_values": tf.constant([1.5, 2, 3]),
+ },
+ {
+ "testcase_name": "control_only",
+ "values": tf.constant([1.5, 2, 3]),
+ "is_treatment": tf.zeros((3, 1)),
+ "expected_control_values": tf.constant([1.5, 2, 3]),
+ "expected_treatment_values": tf.zeros((0,)),
+ },
+ {
+ "testcase_name": "control_and_treatment",
+ "values": tf.concat(
+ values=[tf.ones((2, 2, 3)), 2.0 * tf.ones((1, 2, 3))], axis=0
+ ),
+ "is_treatment": tf.constant([[0], [0], [1]]),
+ "expected_control_values": tf.ones((2, 2, 3)),
+ "expected_treatment_values": 2.0 * tf.ones((1, 2, 3)),
+ },
+ {
+ "testcase_name": "one_dimensional_is_treatment",
+ "values": tf.concat(
+ values=[tf.ones((2, 2, 3)), 2.0 * tf.ones((1, 2, 3))], axis=0
+ ),
+ "is_treatment": tf.constant([0, 0, 1]),
+ "expected_control_values": tf.ones((2, 2, 3)),
+ "expected_treatment_values": 2.0 * tf.ones((1, 2, 3)),
+ },
+ {
+ "testcase_name": "empty_values",
+ "values": tf.raw_ops.Empty(shape=(0,), dtype=tf.float32, init=True),
+ "is_treatment": tf.raw_ops.Empty(
+ shape=(0,), dtype=tf.float32, init=True
+ ),
+ "expected_control_values": tf.zeros((0,)),
+ "expected_treatment_values": tf.zeros((0,)),
+ },
+ )
+ def test_split_by_treatment(
+ self,
+ values,
+ is_treatment,
+ expected_control_values,
+ expected_treatment_values,
+ ):
+ control_values, treatment_values = utils.split_by_treatment(
+ values=values, is_treatment=is_treatment
+ )
+ self.assertAllEqual(expected_control_values, control_values)
+ self.assertAllEqual(expected_treatment_values, treatment_values)
+
+ @parameterized.named_parameters(
+ {
+ "testcase_name": "decimal_values",
+ "is_treatment": tf.constant([1.0, 0.3, 0.0]),
+ "expected_error": tf.errors.InvalidArgumentError,
+ },
+ {
+ "testcase_name": "string_values",
+ "is_treatment": tf.constant(["a", "b", "c"]),
+ "expected_error": ValueError,
+ },
+ )
+ def test_invalid_treatment_indicator_tensor(
+ self, is_treatment, expected_error
+ ):
+ values = tf.ones((3, 1))
+ with self.assertRaises(expected_error):
+ utils.split_by_treatment(values, is_treatment)
+
+ def test_shape_mismatch(self):
+ values = tf.ones((4, 1))
+ is_treatment = tf.constant([0, 0, 1])
+ with self.assertRaises(ValueError):
+ utils.split_by_treatment(values, is_treatment)
+
+
+if __name__ == "__main__":
+ tf.test.main()
diff --git a/official/requirements.txt b/official/requirements.txt
index bff65abbe00..fecd055d445 100644
--- a/official/requirements.txt
+++ b/official/requirements.txt
@@ -10,14 +10,13 @@ scipy>=0.19.1
tensorflow-hub>=0.6.0
tensorflow-model-optimization>=0.4.1
tensorflow-datasets
-tensorflow-addons
-dataclasses;python_version<"3.7"
+tf-keras>=2.16.0
gin-config
tf_slim>=1.1.0
Cython
matplotlib
# Loader becomes a required positional argument in 6.0 in yaml.load
-pyyaml>=5.1,<6.0
+pyyaml>=6.0.0
# CV related dependencies
opencv-python-headless
Pillow
@@ -26,3 +25,6 @@ pycocotools
seqeval
sentencepiece
sacrebleu
+# Projects/vit dependencies
+immutabledict
+ai-edge-litert>=1.0.1
\ No newline at end of file
diff --git a/official/utils/__init__.py b/official/utils/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/utils/__init__.py
+++ b/official/utils/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/docs/__init__.py b/official/utils/docs/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/utils/docs/__init__.py
+++ b/official/utils/docs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/docs/build_orbit_api_docs.py b/official/utils/docs/build_orbit_api_docs.py
index c7f25715e8c..cd52f4662a0 100644
--- a/official/utils/docs/build_orbit_api_docs.py
+++ b/official/utils/docs/build_orbit_api_docs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -25,7 +25,7 @@
import orbit
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow_docs.api_generator import doc_controls
from tensorflow_docs.api_generator import generate_lib
from tensorflow_docs.api_generator import public_api
@@ -57,8 +57,8 @@ def hide_module_model_and_layer_methods():
complex layers.
"""
module_contents = list(tf.Module.__dict__.items())
- model_contents = list(tf.keras.Model.__dict__.items())
- layer_contents = list(tf.keras.layers.Layer.__dict__.items())
+ model_contents = list(tf_keras.Model.__dict__.items())
+ layer_contents = list(tf_keras.layers.Layer.__dict__.items())
for name, obj in module_contents + layer_contents + model_contents:
if name == '__init__':
diff --git a/official/utils/docs/build_tfm_api_docs.py b/official/utils/docs/build_tfm_api_docs.py
index 75a401b9935..edba914779a 100644
--- a/official/utils/docs/build_tfm_api_docs.py
+++ b/official/utils/docs/build_tfm_api_docs.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -26,10 +26,13 @@
from absl import flags
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow_docs.api_generator import doc_controls
from tensorflow_docs.api_generator import generate_lib
+from tensorflow_docs.api_generator import parser
from tensorflow_docs.api_generator import public_api
+from tensorflow_docs.api_generator.pretty_docs import base_page
+from tensorflow_docs.api_generator.pretty_docs import function_page
import tensorflow_models as tfm
@@ -52,6 +55,44 @@
PROJECT_FULL_NAME = 'TensorFlow Modeling Library'
+class ExpFactoryInfo(function_page.FunctionPageInfo):
+ """Customize the page for the experiment factory."""
+
+ def collect_docs(self):
+ super().collect_docs()
+ self.doc.docstring_parts.append(self.make_factory_options_table())
+
+ def make_factory_options_table(self):
+ lines = [
+ '',
+ 'Allowed values for `exp_name`:',
+ '',
+ # The indent is important here, it keeps the site's markdown parser
+ # from switching to HTML mode.
+ '
\n',
+ '
exp_name
Description
',
+ ]
+ reference_resolver = self.parser_config.reference_resolver
+ api_tree = self.parser_config.api_tree
+ for name, fn in sorted(tfm.core.exp_factory._REGISTERED_CONFIGS.items()): # pylint: disable=protected-access
+ fn_api_node = api_tree.node_for_object(fn)
+ if fn_api_node is None:
+ location = parser.get_defined_in(self.py_object, self.parser_config)
+ link = base_page.small_source_link(location, name)
+ else:
+ link = reference_resolver.python_link(name, fn_api_node.full_name)
+ doc = fn.__doc__
+ if doc:
+ doc = doc.splitlines()[0]
+ else:
+ doc = ''
+
+ lines.append(f'
{link}
{doc}
')
+
+ lines.append('
')
+ return '\n'.join(lines)
+
+
def hide_module_model_and_layer_methods():
"""Hide methods and properties defined in the base classes of Keras layers.
@@ -61,8 +102,8 @@ def hide_module_model_and_layer_methods():
complex layers.
"""
module_contents = list(tf.Module.__dict__.items())
- model_contents = list(tf.keras.Model.__dict__.items())
- layer_contents = list(tf.keras.layers.Layer.__dict__.items())
+ model_contents = list(tf_keras.Model.__dict__.items())
+ layer_contents = list(tf_keras.layers.Layer.__dict__.items())
for name, obj in module_contents + layer_contents + model_contents:
if name == '__init__':
@@ -103,6 +144,9 @@ def gen_api_docs(code_url_prefix, site_path, output_dir, project_short_name,
del tfm.nlp.layers.MultiHeadAttention
del tfm.nlp.layers.EinsumDense
+ doc_controls.set_custom_page_builder_cls(tfm.core.exp_factory.get_exp_config,
+ ExpFactoryInfo)
+
url_parts = code_url_prefix.strip('/').split('/')
url_parts = url_parts[:url_parts.index('tensorflow_models')]
url_parts.append('official')
diff --git a/official/utils/flags/README.md b/official/utils/flags/README.md
index beb3b2a1e1d..f015a9c2672 100644
--- a/official/utils/flags/README.md
+++ b/official/utils/flags/README.md
@@ -1,13 +1,13 @@
# Adding Abseil (absl) flags quickstart
-**WARNING** This module is deprecated. We no long use it in new models and
+**WARNING** This module is deprecated. We no longer use it in new models and
your projects should not depend on it. We will remove this module when
all models using it are deprecated which may take time.
## Defining a flag
absl flag definitions are similar to argparse, although they are defined on a global namespace.
-For instance defining a string flag looks like:
+For instance, defining a string flag looks like this:
```$xslt
from absl import flags
flags.DEFINE_string(
@@ -17,13 +17,13 @@ flags.DEFINE_string(
)
```
-All three arguments are required, but default may be `None`. A common optional argument is
-short_name for defining abreviations. Certain `DEFINE_*` methods will have other required arguments.
-For instance `DEFINE_enum` requires the `enum_values` argument to be specified.
+All three arguments are required, but the default may be `None`. A common optional argument is
+short_name for defining abbreviations. Certain `DEFINE_*` methods will have other required arguments.
+For instance, `DEFINE_enum` requires the `enum_values` argument to be specified.
## Key Flags
absl has the concept of a key flag. Any flag defined in `__main__` is considered a key flag by
-default. Key flags are displayed in `--help`, others only appear in `--helpfull`. In order to
+default. Key flags are displayed in `--help`, others only appear in `--helpful`. To
handle key flags that are defined outside the module in question, absl provides the
`flags.adopt_module_key_flags()` method. This adds the key flags of a different module to one's own
key flags. For example:
@@ -54,14 +54,14 @@ absl_app.run(main, [__file__, "-h"]
when `my_module.py` is run it will show the help text for `my_flag`. Because not all flags defined
in a file are equally important, `official/utils/flags/core.py` (generally imported as flags_core)
-provides an abstraction for handling key flag declaration in an easy way through the
+provides an abstraction for easily handling key flag declaration through the
`register_key_flags_in_core()` function, which allows a module to make a single
`adopt_key_flags(flags_core)` call when using the util flag declaration functions.
## Validators
Often the constraints on a flag are complicated. absl provides the validator decorator to allow
one to mark a function as a flag validation function. Suppose we want users to provide a flag
-which is a palindrome.
+that is a palindrome.
```$xslt
from absl import flags
diff --git a/official/utils/flags/__init__.py b/official/utils/flags/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/utils/flags/__init__.py
+++ b/official/utils/flags/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/flags/_base.py b/official/utils/flags/_base.py
index 8e8f5b4d3cf..a2dff77cdd8 100644
--- a/official/utils/flags/_base.py
+++ b/official/utils/flags/_base.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Flags which will be nearly universal across models."""
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.utils.flags._conventions import help_wrap
diff --git a/official/utils/flags/_benchmark.py b/official/utils/flags/_benchmark.py
index 97adf563268..942ac954d6f 100644
--- a/official/utils/flags/_benchmark.py
+++ b/official/utils/flags/_benchmark.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/flags/_conventions.py b/official/utils/flags/_conventions.py
index fa3d186d540..5b10162eab0 100644
--- a/official/utils/flags/_conventions.py
+++ b/official/utils/flags/_conventions.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/flags/_device.py b/official/utils/flags/_device.py
index 09e004b720d..8cf630d2ec9 100644
--- a/official/utils/flags/_device.py
+++ b/official/utils/flags/_device.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/flags/_distribution.py b/official/utils/flags/_distribution.py
index 76ec5bb283e..b32fb7b401b 100644
--- a/official/utils/flags/_distribution.py
+++ b/official/utils/flags/_distribution.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Flags related to distributed execution."""
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.utils.flags._conventions import help_wrap
diff --git a/official/utils/flags/_misc.py b/official/utils/flags/_misc.py
index fc25d7bbf93..68d0bb6bf25 100644
--- a/official/utils/flags/_misc.py
+++ b/official/utils/flags/_misc.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/flags/_performance.py b/official/utils/flags/_performance.py
index 6ccfd4a9f89..4205a9c82f5 100644
--- a/official/utils/flags/_performance.py
+++ b/official/utils/flags/_performance.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import multiprocessing
from absl import flags # pylint: disable=g-bad-import-order
-import tensorflow as tf # pylint: disable=g-bad-import-order
+import tensorflow as tf, tf_keras # pylint: disable=g-bad-import-order
from official.utils.flags._conventions import help_wrap
@@ -198,11 +198,11 @@ def _check_loss_scale(loss_scale):
flags.DEFINE_enum(
name="fp16_implementation",
default="keras",
- enum_values=("keras', 'graph_rewrite"),
+ enum_values=("keras", "graph_rewrite"),
help=help_wrap(
"When --dtype=fp16, how fp16 should be implemented. This has no "
"impact on correctness. 'keras' uses the "
- "tf.keras.mixed_precision API. 'graph_rewrite' uses the "
+ "tf_keras.mixed_precision API. 'graph_rewrite' uses the "
"tf.compat.v1.mixed_precision."
"enable_mixed_precision_graph_rewrite API."))
diff --git a/official/utils/flags/core.py b/official/utils/flags/core.py
index 36a244da239..eac05d08c2a 100644
--- a/official/utils/flags/core.py
+++ b/official/utils/flags/core.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/flags/flags_test.py b/official/utils/flags/flags_test.py
index f8c639c3963..d456d510ed8 100644
--- a/official/utils/flags/flags_test.py
+++ b/official/utils/flags/flags_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
import unittest
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.utils.flags import core as flags_core # pylint: disable=g-bad-import-order
diff --git a/official/utils/hyperparams_flags.py b/official/utils/hyperparams_flags.py
index d3428e0f9b8..ebb0b75b757 100644
--- a/official/utils/hyperparams_flags.py
+++ b/official/utils/hyperparams_flags.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/misc/__init__.py b/official/utils/misc/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/utils/misc/__init__.py
+++ b/official/utils/misc/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/misc/keras_utils.py b/official/utils/misc/keras_utils.py
index c3e8d12b038..46c4e6001e7 100644
--- a/official/utils/misc/keras_utils.py
+++ b/official/utils/misc/keras_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -19,7 +19,7 @@
import time
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.eager import monitoring
@@ -47,7 +47,7 @@ def __repr__(self):
self.batch_index, self.timestamp)
-class TimeHistory(tf.keras.callbacks.Callback):
+class TimeHistory(tf_keras.callbacks.Callback):
"""Callback for Keras models."""
def __init__(self, batch_size, log_steps, initial_step=0, logdir=None):
@@ -165,7 +165,7 @@ def on_epoch_end(self, epoch, logs=None):
self.steps_in_epoch = 0
-class SimpleCheckpoint(tf.keras.callbacks.Callback):
+class SimpleCheckpoint(tf_keras.callbacks.Callback):
"""Keras callback to save tf.train.Checkpoints."""
def __init__(self, checkpoint_manager):
diff --git a/official/utils/misc/model_helpers.py b/official/utils/misc/model_helpers.py
index f5065ceaef1..d6c5a85e12a 100644
--- a/official/utils/misc/model_helpers.py
+++ b/official/utils/misc/model_helpers.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import numbers
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from tensorflow.python.util import nest
# pylint:disable=logging-format-interpolation
diff --git a/official/utils/misc/model_helpers_test.py b/official/utils/misc/model_helpers_test.py
index 6d5a3e84e22..a21c484c17b 100644
--- a/official/utils/misc/model_helpers_test.py
+++ b/official/utils/misc/model_helpers_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Tests for Model Helper functions."""
-import tensorflow as tf # pylint: disable=g-bad-import-order
+import tensorflow as tf, tf_keras # pylint: disable=g-bad-import-order
from official.utils.misc import model_helpers
@@ -82,7 +82,7 @@ def test_generate_synethetic_data(self):
for n in range(5):
inp, lab = sess.run((input_element, label_element))
self.assertAllClose(inp, [123., 123., 123., 123., 123.])
- self.assertEquals(lab, 456)
+ self.assertEqual(lab, 456)
def test_generate_only_input_data(self):
d = model_helpers.generate_synthetic_data(
@@ -111,7 +111,7 @@ def test_generate_nested_data(self):
element = tf.compat.v1.data.make_one_shot_iterator(d).get_next()
self.assertIn('a', element)
self.assertIn('b', element)
- self.assertEquals(len(element['b']), 2)
+ self.assertEqual(len(element['b']), 2)
self.assertIn('c', element['b'])
self.assertIn('d', element['b'])
self.assertNotIn('c', element)
diff --git a/official/utils/testing/__init__.py b/official/utils/testing/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/utils/testing/__init__.py
+++ b/official/utils/testing/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/testing/integration.py b/official/utils/testing/integration.py
index 84af32a0151..8f2dff19f38 100644
--- a/official/utils/testing/integration.py
+++ b/official/utils/testing/integration.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/utils/testing/mock_task.py b/official/utils/testing/mock_task.py
index dd7e493e197..8fd1a61887b 100644
--- a/official/utils/testing/mock_task.py
+++ b/official/utils/testing/mock_task.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
import dataclasses
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.core import base_task
from official.core import config_definitions as cfg
@@ -25,13 +25,13 @@
from official.modeling.hyperparams import base_config
-class MockModel(tf.keras.Model):
+class MockModel(tf_keras.Model):
def __init__(self, network):
super().__init__()
self.network = network
- def call(self, inputs):
+ def call(self, inputs): # pytype: disable=signature-mismatch # overriding-parameter-count-checks
outputs = self.network(inputs)
self.add_loss(tf.reduce_mean(outputs))
return outputs
@@ -50,18 +50,18 @@ def __init__(self, params=None, logging_dir=None, name=None):
super().__init__(params=params, logging_dir=logging_dir, name=name)
def build_model(self, *arg, **kwargs):
- inputs = tf.keras.layers.Input(shape=(2,), name="random", dtype=tf.float32)
- outputs = tf.keras.layers.Dense(
- 1, bias_initializer=tf.keras.initializers.Ones(), name="dense_0")(
+ inputs = tf_keras.layers.Input(shape=(2,), name="random", dtype=tf.float32)
+ outputs = tf_keras.layers.Dense(
+ 1, bias_initializer=tf_keras.initializers.Ones(), name="dense_0")(
inputs)
- network = tf.keras.Model(inputs=inputs, outputs=outputs)
+ network = tf_keras.Model(inputs=inputs, outputs=outputs)
return MockModel(network)
def build_metrics(self, training: bool = True):
del training
- return [tf.keras.metrics.Accuracy(name="acc")]
+ return [tf_keras.metrics.Accuracy(name="acc")]
- def validation_step(self, inputs, model: tf.keras.Model, metrics=None):
+ def validation_step(self, inputs, model: tf_keras.Model, metrics=None):
logs = super().validation_step(inputs, model, metrics)
logs["counter"] = tf.constant(1, dtype=tf.float32)
return logs
diff --git a/official/utils/testing/scripts/presubmit.sh b/official/utils/testing/scripts/presubmit.sh
index b94683f48aa..855a26fd299 100755
--- a/official/utils/testing/scripts/presubmit.sh
+++ b/official/utils/testing/scripts/presubmit.sh
@@ -32,7 +32,7 @@ py_test() {
echo "===========Running Python test============"
# Skipping Ranking tests, TODO(b/189265753) remove it once the issue is fixed.
- for test_file in `find official/ -name '*test.py' -print | grep -v 'official/recommendation/ranking'`
+ for test_file in `find official/ -name '*test.py' -print | grep -v -E 'official/(recommendation/ranking|legacy)'`
do
echo "####=======Testing ${test_file}=======####"
${PY_BINARY} "${test_file}"
@@ -47,7 +47,7 @@ py_test() {
}
py3_test() {
- local PY_BINARY=python3.9
+ local PY_BINARY=python3.11
py_test "$PY_BINARY"
return $?
}
diff --git a/official/vision/Images/Runtime_Config.png b/official/vision/Images/Runtime_Config.png
new file mode 100644
index 00000000000..71f30d52611
Binary files /dev/null and b/official/vision/Images/Runtime_Config.png differ
diff --git a/official/vision/Images/faster_rcnn.png b/official/vision/Images/faster_rcnn.png
new file mode 100644
index 00000000000..8aac4c118ad
Binary files /dev/null and b/official/vision/Images/faster_rcnn.png differ
diff --git a/official/vision/Images/mask_rcnn.png b/official/vision/Images/mask_rcnn.png
new file mode 100644
index 00000000000..f3c3940196c
Binary files /dev/null and b/official/vision/Images/mask_rcnn.png differ
diff --git a/official/vision/Images/retinanet.png b/official/vision/Images/retinanet.png
new file mode 100644
index 00000000000..722810c1a46
Binary files /dev/null and b/official/vision/Images/retinanet.png differ
diff --git a/official/vision/Images/yolo.png b/official/vision/Images/yolo.png
new file mode 100644
index 00000000000..01e78ec0931
Binary files /dev/null and b/official/vision/Images/yolo.png differ
diff --git a/official/vision/MODEL_GARDEN.md b/official/vision/MODEL_GARDEN.md
index 0ce10df86c3..da3fb4d268d 100644
--- a/official/vision/MODEL_GARDEN.md
+++ b/official/vision/MODEL_GARDEN.md
@@ -1,9 +1,13 @@
# TF-Vision Model Garden
-⚠️ Disclaimer: All datasets hyperlinked from this page are not owned or
-distributed by Google. The dataset is made available by third parties.
-Please review the terms and conditions made available by the third parties
-before using the data.
+⚠️ Disclaimer: Checkpoints are based on training with publicly available datasets. Some datasets contain limitations, including non-commercial use limitations.
+Please review the terms and conditions made available by third parties before using
+the datasets provided. Checkpoints are licensed under
+[Apache 2.0](https://github.com/tensorflow/models/blob/master/LICENSE).
+
+⚠️ Disclaimer: Datasets hyperlinked from this page are not owned or distributed
+by Google. Such datasets are made available by third parties. Please review the
+terms and conditions made available by the third parties before using the data.
## Introduction
@@ -25,9 +29,9 @@ segmentation.
| Model | Resolution | Epochs | Top-1 | Top-5 | Download |
| ------------ |:-------------:|--------:|--------:|--------:|---------:|
| ResNet-50 | 224x224 | 90 | 76.1 | 92.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml) |
-| ResNet-50 | 224x224 | 200 | 77.1 | 93.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml) |
-| ResNet-101 | 224x224 | 200 | 78.3 | 94.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml) |
-| ResNet-152 | 224x224 | 200 | 78.7 | 94.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml) |
+| ResNet-50 | 224x224 | 200 | 77.1 | 93.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet/resnet-50-i224.tar.gz) |
+| ResNet-101 | 224x224 | 200 | 78.3 | 94.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet/resnet-101-i224.tar.gz) |
+| ResNet-152 | 224x224 | 200 | 78.7 | 94.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet/resnet-152-i224.tar.gz) |
#### ResNet-RS models trained with various settings
@@ -36,37 +40,37 @@ classification models with features:
* ResNet-RS architectural changes and Swish activation. (Note that ResNet-RS
adopts ReLU activation in the paper.)
-* Regularization methods including Random Augment, 4e-5 weight decay, stochastic
-depth, label smoothing and dropout.
-* New training methods including a 350-epoch schedule, cosine learning rate and
+* Regularization methods include Random Augment, 4e-5 weight decay, stochastic
+depth, label smoothing , and dropout.
+* New training methods including a 350-epoch schedule, cosine learning rate , and
EMA.
* Configs are in this [directory](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification).
| Model | Resolution | Params (M) | Top-1 | Top-5 | Download |
| --------- | :--------: | ---------: | ----: | ----: | --------:|
-| ResNet-RS-50 | 160x160 | 35.7 | 79.1 | 94.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-50-i160.tar.gz) |
-| ResNet-RS-101 | 160x160 | 63.7 | 80.2 | 94.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i160.tar.gz) |
-| ResNet-RS-101 | 192x192 | 63.7 | 81.3 | 95.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i192.tar.gz) |
-| ResNet-RS-152 | 192x192 | 86.8 | 81.9 | 95.8 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i192.tar.gz) |
-| ResNet-RS-152 | 224x224 | 86.8 | 82.5 | 96.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i224.tar.gz) |
-| ResNet-RS-152 | 256x256 | 86.8 | 83.1 | 96.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i256.tar.gz) |
-| ResNet-RS-200 | 256x256 | 93.4 | 83.5 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-200-i256.tar.gz) |
-| ResNet-RS-270 | 256x256 | 130.1 | 83.6 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-270-i256.tar.gz) |
-| ResNet-RS-350 | 256x256 | 164.3 | 83.7 | 96.7 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i256.tar.gz) |
-| ResNet-RS-350 | 320x320 | 164.3 | 84.2 | 96.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i320.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i320.tar.gz) |
+| ResNet-RS-50 | 160x160 | 35.7 | 79.1 | 94.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-50-i160.tar.gz) |
+| ResNet-RS-101 | 160x160 | 63.7 | 80.2 | 94.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i160.tar.gz) |
+| ResNet-RS-101 | 192x192 | 63.7 | 81.3 | 95.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i192.tar.gz) |
+| ResNet-RS-152 | 192x192 | 86.8 | 81.9 | 95.8 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i192.tar.gz) |
+| ResNet-RS-152 | 224x224 | 86.8 | 82.5 | 96.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i224.tar.gz) |
+| ResNet-RS-152 | 256x256 | 86.8 | 83.1 | 96.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i256.tar.gz) |
+| ResNet-RS-200 | 256x256 | 93.4 | 83.5 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-200-i256.tar.gz) |
+| ResNet-RS-270 | 256x256 | 130.1 | 83.6 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-270-i256.tar.gz) |
+| ResNet-RS-350 | 256x256 | 164.3 | 83.7 | 96.7 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i256.tar.gz) |
+| ResNet-RS-350 | 320x320 | 164.3 | 84.2 | 96.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i320.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i320.tar.gz) |
#### Vision Transformer (ViT)
-We support [ViT](https://arxiv.org/abs/2010.11929) and [DEIT](https://arxiv.org/abs/2012.12877) implementations in a TF
-Vision
-[project](https://github.com/tensorflow/models/tree/master/official/projects/vit). ViT models trained under the DEIT settings:
+We support [ViT](https://arxiv.org/abs/2010.11929) and [DEIT](https://arxiv.org/abs/2012.12877) implementations.
+ViT models trained under the DEIT settings:
-model | resolution | Top-1 | Top-5 |
---------- | :--------: | ----: | ----: |
-ViT-s16 | 224x224 | 79.4 | 94.7 |
-ViT-b16 | 224x224 | 81.8 | 95.8 |
-ViT-l16 | 224x224 | 82.2 | 95.8 |
+model | resolution | Top-1 | Top-5 | Download |
+--------- | :--------: | ----: | ----: | :-------: |
+ViT-ti16 | 224x224 | 73.4 | 91.9 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-ti16.tar.gz) |
+ViT-s16 | 224x224 | 79.4 | 94.7 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-s16.tar.gz) |
+ViT-b16 | 224x224 | 81.8 | 95.8 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-b16.tar.gz) |
+ViT-l16 | 224x224 | 82.2 | 95.8 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-l16.tar.gz) |
## Object Detection and Instance Segmentation
@@ -81,17 +85,17 @@ ViT-l16 | 224x224 | 82.2 | 95.8 |
* Models are all trained on [COCO](https://cocodataset.org/) train2017 and
evaluated on [COCO](https://cocodataset.org/) val2017.
* Training details:
- * Models finetuned from [ImageNet](https://www.image-net.org/) pretrained
+ * Models finetuned from [ImageNet](https://www.image-net.org/) pre-trained
checkpoints adopt the 12 or 36 epochs schedule. Models trained from scratch
adopt the 350 epochs schedule.
* The default training data augmentation implements horizontal flipping and
scale jittering with a random scale between [0.5, 2.0].
* Unless noted, all models are trained with l2 weight regularization and ReLU
activation.
- * We use batch size 256 and stepwise learning rate that decays at the last 30
- and 10 epoch.
- * We use square image as input by resizing the long side of an image to the
- target size then padding the short side with zeros.
+ * We use batch size 256 and a stepwise learning rate that decays at the last 30
+ and 10 epochs.
+ * We use a square image as input by resizing the long side of an image to the
+ target size and then padding the short side with zeros.
### COCO Object Detection Baselines
@@ -100,7 +104,7 @@ evaluated on [COCO](https://cocodataset.org/) val2017.
| Backbone | Resolution | Epochs | FLOPs (B) | Params (M) | Box AP | Download |
| ------------ |:-------------:| -------:|--------------:|-----------:|-------:|---------:|
| R50-FPN | 640x640 | 12 | 97.0 | 34.0 | 34.3 | config|
-| R50-FPN | 640x640 | 72 | 97.0 | 34.0 | 36.8 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/retinanet.py#L187-L258) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/retinanet/retinanet-resnet50fpn.tar.gz) |
+| R50-FPN | 640x640 | 72 | 97.0 | 34.0 | 36.8 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/retinanet.py#L187-L258) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/retinanet/retinanet-resnet50fpn.tar.gz) |
#### RetinaNet (Trained from scratch) with training features including:
@@ -109,16 +113,16 @@ evaluated on [COCO](https://cocodataset.org/) val2017.
| Backbone | Resolution | Epochs | FLOPs (B) | Params (M) | Box AP | Download |
| ------------ |:-------------:| -------:|--------------:|-----------:|--------:|---------:|
-| SpineNet-49 | 640x640 | 500 | 85.4| 28.5 | 44.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_tpu.yaml) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
-| SpineNet-96 | 1024x1024 | 500 | 265.4 | 43.0 | 48.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet96_tpu.yaml) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
-| SpineNet-143 | 1280x1280 | 500 | 524.0 | 67.0 | 50.0 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet143_tpu.yaml) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
+| SpineNet-49 | 640x640 | 500 | 85.4| 28.5 | 44.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/spinenet/spinenet-49-i640.tar.gz) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
+| SpineNet-96 | 1024x1024 | 500 | 265.4 | 43.0 | 48.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet96_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/spinenet/spinenet-96-i1024.tar.gz) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
+| SpineNet-143 | 1280x1280 | 500 | 524.0 | 67.0 | 50.0 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet143_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/spinenet/spinenet-143-i1280.tar.gz) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
#### Mobile-size RetinaNet (Trained from scratch):
| Backbone | Resolution | Epochs | FLOPs (B) | Params (M) | Box AP | Download |
| ----------- | :--------: | -----: | --------: | ---------: | -----: | --------:|
| MobileNetv2 | 256x256 | 600 | - | 2.27 | 23.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_mobilenetv2_tpu.yaml) |
-| Mobile SpineNet-49 | 384x384 | 600 | 1.0 | 2.32 | 28.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/retinanet/spinenet49mobile.tar.gz) |
+| Mobile SpineNet-49 | 384x384 | 600 | 1.0 | 2.32 | 28.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/retinanet/spinenet49mobile.tar.gz) |
### Instance Segmentation Baselines
@@ -152,7 +156,7 @@ evaluated on [COCO](https://cocodataset.org/) val2017.
| Model | Backbone | Resolution | Steps | mIoU | Download |
| ---------- | :----------------: | :--------: | ----: | ---: | --------:|
| DeepLabV3 | Dilated Resnet-101 | 512x512 | 30k | 78.7 | |
-| DeepLabV3+ | Dilated Resnet-101 | 512x512 | 30k | 79.2 | |
+| DeepLabV3+ | Dilated Resnet-101 | 512x512 | 30k | 79.2 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/deeplabv3plus/dilated-resnet-101-deeplabv3plus.tar.gz) |
### CITYSCAPES
@@ -177,9 +181,9 @@ evaluated on [COCO](https://cocodataset.org/) val2017.
* Training and evaluation details (SlowFast and ResNet):
* All models are trained from scratch with vision modality (RGB) for 200
epochs.
- * We use batch size of 1024 and cosine learning rate decay with linear warmup
- in first 5 epochs.
- * We follow [SlowFast](https://arxiv.org/abs/1812.03982) to perform 30-view
+ * We use a batch size of 1024 and cosine learning rate decay with a linear warmup
+ in the first 5 epochs.
+ * We follow [SlowFast](https://arxiv.org/abs/1812.03982) to perform a 30-view
evaluation.
### Kinetics-400 Action Recognition Baselines
diff --git a/official/vision/README.md b/official/vision/README.md
index 57365b3c16f..676e3a17a20 100644
--- a/official/vision/README.md
+++ b/official/vision/README.md
@@ -1,13 +1,21 @@
# TF-Vision Model Garden
-⚠️ Disclaimer: All datasets hyperlinked from this page are not owned or
-distributed by Google. The dataset is made available by third parties.
-Please review the terms and conditions made available by the third parties
-before using the data.
+⚠️ Disclaimer: Checkpoints are based on training with publicly available
+datasets. Some datasets contain limitations, including non-commercial use
+limitations. Please review the terms and conditions made available by third
+parties before using the datasets provided. Checkpoints are licensed under
+[Apache 2.0](https://github.com/tensorflow/models/blob/master/LICENSE).
+
+⚠️ Disclaimer: Datasets hyperlinked from this page are not owned or distributed
+by Google. Such datasets are made available by third parties. Please review the
+terms and conditions made available by the third parties before using the data.
## Table of Contents
- [Introduction](#introduction)
+- [Backbones](#backbones)
+- [Decoders](#decoders)
+- [Heads](#heads)
- [Image Classification](#image-classification)
* [ResNet models trained with vanilla settings](#resnet-models-trained-with-vanilla-settings)
* [ResNet-RS models trained with various settings](#resnet-rs-models-trained-with-various-settings)
@@ -18,6 +26,7 @@ before using the data.
* [RetinaNet (ImageNet pretrained)](#RetinaNet-ImageNet-pretrained)
* [RetinaNet (Trained from scratch)](#RetinaNet-Trained-from-scratch)
* [Mobile-size RetinaNet (Trained from scratch)](#Mobile-size-RetinaNet-Trained-from-scratch))
+ * [YOLOv7 (Trained from scratch)](#yolov7-trained-from-scratch)
- [Instance Segmentation Baselines](#Instance-Segmentation-Baselines)
* [Mask R-CNN (Trained from scratch)](#Mask-R-CNN-Trained-from-scratch)
* [Cascade RCNN-RS (Trained from scratch)](#Cascade-RCNN-RS-Trained-from-scratch)
@@ -35,6 +44,40 @@ TF-Vision modeling library for computer vision provides a collection of
baselines and checkpoints for image classification, object detection, and
segmentation.
+## Backbones
+
+| Backbones |
+| ---------------- |
+| [DilatedResNet](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/DilatedResNet) |
+| [EfficientNet](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/EfficientNet) |
+| [MobileDet](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/MobileDet) |
+| [MobileNet](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/MobileNet) |
+| [ResNet](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/ResNet) |
+| [ResNet3D](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/ResNet3D) |
+| [RevNet](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/RevNet) |
+| [SpineNet](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/SpineNet) |
+| [SpineNetMobile](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/SpineNetMobile) |
+| [VisionTransformer](https://www.tensorflow.org/api_docs/python/tfm/vision/backbones/VisionTransformer) |
+
+## Decoders
+
+| Decoders |
+| --------------- |
+| [ASPP](https://www.tensorflow.org/api_docs/python/tfm/vision/decoders/ASPP) |
+| [FPN](https://www.tensorflow.org/api_docs/python/tfm/vision/decoders/FPN) |
+| [NASFPN](https://www.tensorflow.org/api_docs/python/tfm/vision/decoders/NASFPN) |
+
+## Heads
+
+| Heads |
+| --------------- |
+| [DetectionHead](https://www.tensorflow.org/api_docs/python/tfm/vision/heads/DetectionHead) |
+| [MaskHead](https://www.tensorflow.org/api_docs/python/tfm/vision/heads/MaskHead) |
+| [MaskScoring](https://www.tensorflow.org/api_docs/python/tfm/vision/heads/MaskScoring) |
+| [RPNHead](https://www.tensorflow.org/api_docs/python/tfm/vision/heads/RPNHead) |
+| [RetinaNetHead](https://www.tensorflow.org/api_docs/python/tfm/vision/heads/RetinaNetHead) |
+| [SegmentationHead](https://www.tensorflow.org/api_docs/python/tfm/vision/heads/SegmentationHead) |
+
## Image Classification
### ResNet models trained with vanilla settings
@@ -49,9 +92,9 @@ segmentation.
| Model | Resolution | Epochs | Top-1 | Top-5 | Download |
| ------------ |:-------------:|--------:|--------:|--------:|---------:|
| ResNet-50 | 224x224 | 90 | 76.1 | 92.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml) |
-| ResNet-50 | 224x224 | 200 | 77.1 | 93.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml) |
-| ResNet-101 | 224x224 | 200 | 78.3 | 94.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml) |
-| ResNet-152 | 224x224 | 200 | 78.7 | 94.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml) |
+| ResNet-50 | 224x224 | 200 | 77.1 | 93.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet/resnet-50-i224.tar.gz) |
+| ResNet-101 | 224x224 | 200 | 78.3 | 94.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet/resnet-101-i224.tar.gz) |
+| ResNet-152 | 224x224 | 200 | 78.7 | 94.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet/resnet-152-i224.tar.gz) |
@@ -72,16 +115,16 @@ depth, label smoothing and dropout.
| Model | Resolution | Params (M) | Top-1 | Top-5 | Download |
| --------- | :--------: | ---------: | ----: | ----: | --------:|
-| ResNet-RS-50 | 160x160 | 35.7 | 79.1 | 94.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-50-i160.tar.gz) |
-| ResNet-RS-101 | 160x160 | 63.7 | 80.2 | 94.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i160.tar.gz) |
-| ResNet-RS-101 | 192x192 | 63.7 | 81.3 | 95.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i192.tar.gz) |
-| ResNet-RS-152 | 192x192 | 86.8 | 81.9 | 95.8 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i192.tar.gz) |
-| ResNet-RS-152 | 224x224 | 86.8 | 82.5 | 96.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i224.tar.gz) |
-| ResNet-RS-152 | 256x256 | 86.8 | 83.1 | 96.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i256.tar.gz) |
-| ResNet-RS-200 | 256x256 | 93.4 | 83.5 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-200-i256.tar.gz) |
-| ResNet-RS-270 | 256x256 | 130.1 | 83.6 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-270-i256.tar.gz) |
-| ResNet-RS-350 | 256x256 | 164.3 | 83.7 | 96.7 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i256.tar.gz) |
-| ResNet-RS-350 | 320x320 | 164.3 | 84.2 | 96.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs420_i256.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i320.tar.gz) |
+| ResNet-RS-50 | 160x160 | 35.7 | 79.1 | 94.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-50-i160.tar.gz) |
+| ResNet-RS-101 | 160x160 | 63.7 | 80.2 | 94.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i160.tar.gz) |
+| ResNet-RS-101 | 192x192 | 63.7 | 81.3 | 95.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-101-i192.tar.gz) |
+| ResNet-RS-152 | 192x192 | 86.8 | 81.9 | 95.8 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i192.tar.gz) |
+| ResNet-RS-152 | 224x224 | 86.8 | 82.5 | 96.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i224.tar.gz) |
+| ResNet-RS-152 | 256x256 | 86.8 | 83.1 | 96.3 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-152-i256.tar.gz) |
+| ResNet-RS-200 | 256x256 | 93.4 | 83.5 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-200-i256.tar.gz) |
+| ResNet-RS-270 | 256x256 | 130.1 | 83.6 | 96.6 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-270-i256.tar.gz) |
+| ResNet-RS-350 | 256x256 | 164.3 | 83.7 | 96.7 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i256.tar.gz) |
+| ResNet-RS-350 | 320x320 | 164.3 | 84.2 | 96.9 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/image_classification/imagenet_resnetrs420_i256.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/resnet-rs/resnet-rs-350-i320.tar.gz) |
@@ -89,15 +132,16 @@ depth, label smoothing and dropout.
-We support [ViT](https://arxiv.org/abs/2010.11929) and [DEIT](https://arxiv.org/abs/2012.12877) implementations in a TF
-Vision
-[project](https://github.com/tensorflow/models/tree/master/official/projects/vit). ViT models trained under the DEIT settings:
+We support [ViT](https://arxiv.org/abs/2010.11929) and
+[DEIT](https://arxiv.org/abs/2012.12877) implementations. ViT models trained
+under the DEIT settings:
-model | resolution | Top-1 | Top-5 |
---------- | :--------: | ----: | ----: |
-ViT-s16 | 224x224 | 79.4 | 94.7 |
-ViT-b16 | 224x224 | 81.8 | 95.8 |
-ViT-l16 | 224x224 | 82.2 | 95.8 |
+model | resolution | Top-1 | Top-5 | Download |
+--------- | :--------: | ----: | ----: | :-------: |
+ViT-ti16 | 224x224 | 73.4 | 91.9 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-ti16.tar.gz) |
+ViT-s16 | 224x224 | 79.4 | 94.7 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-s16.tar.gz) |
+ViT-b16 | 224x224 | 81.8 | 95.8 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-b16.tar.gz) |
+ViT-l16 | 224x224 | 82.2 | 95.8 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/vit/vit-deit-imagenet-l16.tar.gz) |
@@ -117,6 +161,15 @@ ViT-l16 | 224x224 | 82.2 | 95.8 |
[Cascade RCNN-RS](https://arxiv.org/abs/2107.00057)
* Models are all trained on [COCO](https://cocodataset.org/) train2017 and
evaluated on [COCO](https://cocodataset.org/) val2017.
+ * The checkpoints were trained on annotations
+ [owned and licensed by the COCO Consortium](https://cocodataset.org/#termsofuse)
+ under a
+ [Creative Commons Attribution 4.0 License](https://creativecommons.org/licenses/by/4.0/legalcode).
+ * The COCO Consortium does not own the copyright of the images
+ corresponding to the annotations. The images are
+ [made available by Flickr](https://www.flickr.com/creativecommons/) under
+ various Creative Commons licenses, and users of the images accept full
+ responsibility for the use of the dataset.
* Training details:
* Models finetuned from [ImageNet](https://www.image-net.org/) pretrained
checkpoints adopt the 12 or 36 epochs schedule. Models trained from
@@ -141,7 +194,7 @@ ViT-l16 | 224x224 | 82.2 | 95.8 |
| Backbone | Resolution | Epochs | FLOPs (B) | Params (M) | Box AP | Download |
| ------------ |:-------------:| -------:|--------------:|-----------:|-------:|---------:|
| R50-FPN | 640x640 | 12 | 97.0 | 34.0 | 34.3 | config|
-| R50-FPN | 640x640 | 72 | 97.0 | 34.0 | 36.8 | config \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/retinanet/retinanet-resnet50fpn.tar.gz) |
+| R50-FPN | 640x640 | 72 | 97.0 | 34.0 | 36.8 | config \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/retinanet/retinanet-resnet50fpn.tar.gz) |
@@ -155,9 +208,9 @@ training features including:
| Backbone | Resolution | Epochs | FLOPs (B) | Params (M) | Box AP | Download |
| ------------ |:-------------:| -------:|--------------:|-----------:|--------:|---------:|
-| SpineNet-49 | 640x640 | 500 | 85.4| 28.5 | 44.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_tpu.yaml) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
-| SpineNet-96 | 1024x1024 | 500 | 265.4 | 43.0 | 48.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet96_tpu.yaml) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
-| SpineNet-143 | 1280x1280 | 500 | 524.0 | 67.0 | 50.0 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet143_tpu.yaml) \| [TB.dev](https://tensorboard.dev/experiment/n2UN83TkTdyKZn3slCWulg/#scalars&_smoothingWeight=0)|
+| SpineNet-49 | 640x640 | 500 | 85.4| 28.5 | 44.2 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/spinenet/spinenet-49-i640.tar.gz) |
+| SpineNet-96 | 1024x1024 | 500 | 265.4 | 43.0 | 48.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet96_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/spinenet/spinenet-96-i1024.tar.gz) |
+| SpineNet-143 | 1280x1280 | 500 | 524.0 | 67.0 | 50.0 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet143_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/spinenet/spinenet-143-i1280.tar.gz) |
@@ -168,7 +221,17 @@ training features including:
| Backbone | Resolution | Epochs | FLOPs (B) | Params (M) | Box AP | Download |
| ----------- | :--------: | -----: | --------: | ---------: | -----: | --------:|
| MobileNetv2 | 256x256 | 600 | - | 2.27 | 23.5 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_mobilenetv2_tpu.yaml) |
-| Mobile SpineNet-49 | 384x384 | 600 | 1.0 | 2.32 | 28.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml) \| [ckpt](https://storage.cloud.google.com/tf_model_garden/vision/retinanet/spinenet49mobile.tar.gz) |
+| Mobile SpineNet-49 | 384x384 | 600 | 1.0 | 2.32 | 28.1 | [config](https://github.com/tensorflow/models/blob/master/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/retinanet/spinenet49mobile.tar.gz) |
+
+
+
+### YOLOv7 (Trained from scratch)
+
+
+
+| Variant | Resolution | Epochs | FLOPs (B) | Params (M) | Box AP | Download |
+| ----------- | :--------: | -----: | --------: | ---------: | -----: | --------:|
+| YOLOv7 | 640x640 | 300 | 53.16 | 44.57 | 50.5 | [config](https://github.com/tensorflow/models/blob/master/official/projects/yolo/configs/experiments/yolov7/detection/yolov7.yaml) \| [ckpt](https://storage.googleapis.com/tf_model_garden/vision/yolo/yolov7/yolov7.tar.gz) |
@@ -213,7 +276,7 @@ training features including:
| Model | Backbone | Resolution | Steps | mIoU | Download |
| ---------- | :----------------: | :--------: | ----: | ---: | --------:|
| DeepLabV3 | Dilated Resnet-101 | 512x512 | 30k | 78.7 | |
-| DeepLabV3+ | Dilated Resnet-101 | 512x512 | 30k | 79.2 | |
+| DeepLabV3+ | Dilated Resnet-101 | 512x512 | 30k | 79.2 | [ckpt](https://storage.googleapis.com/tf_model_garden/vision/deeplabv3plus/dilated-resnet-101-deeplabv3plus.tar.gz) |
@@ -293,3 +356,8 @@ training features including:
| MoViNet-A4-Base | 80 x 3 | 83.48 | 96.16 | [config](https://github.com/tensorflow/models/blob/master/official/projects/movinet/configs/yaml/movinet_a4_k600_8x8.yaml) |
| MoViNet-A5-Base | 120 x 2 | 84.27 | 96.39 | [config](https://github.com/tensorflow/models/blob/master/official/projects/movinet/configs/yaml/movinet_a5_k600_8x8.yaml) |
+
+## More Documentations
+
+Please read through the references in the
+[examples/starter](examples/starter).
diff --git a/official/vision/__init__.py b/official/vision/__init__.py
index bf92ea45e0a..7d3fc60c4e2 100644
--- a/official/vision/__init__.py
+++ b/official/vision/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,7 +13,6 @@
# limitations under the License.
"""Vision package definition."""
-# Lint as: python3
# pylint: disable=unused-import
from official.vision import configs
from official.vision import tasks
diff --git a/official/vision/beta/README.md b/official/vision/beta/README.md
deleted file mode 100644
index 2dbb9631c50..00000000000
--- a/official/vision/beta/README.md
+++ /dev/null
@@ -1,2 +0,0 @@
-Contents of this `beta` folder is going to be deprecated soon and most of the
-content has been moved to[official/vision](https://github.com/tensorflow/models/tree/master/official/vision).
diff --git a/official/vision/beta/projects/README.md b/official/vision/beta/projects/README.md
deleted file mode 100644
index 9c20f07fc60..00000000000
--- a/official/vision/beta/projects/README.md
+++ /dev/null
@@ -1,3 +0,0 @@
-Here are a few projects that are built on tf.vision. They are build and maintain
-by different parties. They can be used as examples of how to build your own
-projects based on tf.vision.
diff --git a/official/vision/beta/projects/movinet/README.md b/official/vision/beta/projects/movinet/README.md
deleted file mode 100644
index d3d6e308b22..00000000000
--- a/official/vision/beta/projects/movinet/README.md
+++ /dev/null
@@ -1 +0,0 @@
-The MoViNet project has moved to [official/projects/movinet](https://github.com/tensorflow/models/tree/master/official/projects/movinet).
\ No newline at end of file
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/__init__.py b/official/vision/beta/projects/panoptic_maskrcnn/__init__.py
deleted file mode 100644
index 00e8f8abe41..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/__init__.py
+++ /dev/null
@@ -1,27 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-# Copyright 2021 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/configs/__init__.py b/official/vision/beta/projects/panoptic_maskrcnn/configs/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/configs/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/modeling/factory.py b/official/vision/beta/projects/panoptic_maskrcnn/modeling/factory.py
deleted file mode 100644
index ca14cfd2ce4..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/modeling/factory.py
+++ /dev/null
@@ -1,144 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Factory method to build panoptic segmentation model."""
-
-import tensorflow as tf
-
-from official.projects.deepmac_maskrcnn.tasks import deep_mask_head_rcnn
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as panoptic_maskrcnn_cfg
-from official.vision.beta.projects.panoptic_maskrcnn.modeling import panoptic_maskrcnn_model
-from official.vision.beta.projects.panoptic_maskrcnn.modeling.layers import panoptic_segmentation_generator
-from official.vision.modeling import backbones
-from official.vision.modeling.decoders import factory as decoder_factory
-from official.vision.modeling.heads import segmentation_heads
-
-
-def build_panoptic_maskrcnn(
- input_specs: tf.keras.layers.InputSpec,
- model_config: panoptic_maskrcnn_cfg.PanopticMaskRCNN,
- l2_regularizer: tf.keras.regularizers.Regularizer = None) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
- """Builds Panoptic Mask R-CNN model.
-
- This factory function builds the mask rcnn first, builds the non-shared
- semantic segmentation layers, and finally combines the two models to form
- the panoptic segmentation model.
-
- Args:
- input_specs: `tf.keras.layers.InputSpec` specs of the input tensor.
- model_config: Config instance for the panoptic maskrcnn model.
- l2_regularizer: Optional `tf.keras.regularizers.Regularizer`, if specified,
- the model is built with the provided regularization layer.
- Returns:
- tf.keras.Model for the panoptic segmentation model.
- """
- norm_activation_config = model_config.norm_activation
- segmentation_config = model_config.segmentation_model
-
- # Builds the maskrcnn model.
- maskrcnn_model = deep_mask_head_rcnn.build_maskrcnn(
- input_specs=input_specs,
- model_config=model_config,
- l2_regularizer=l2_regularizer)
-
- # Builds the semantic segmentation branch.
- if not model_config.shared_backbone:
- segmentation_backbone = backbones.factory.build_backbone(
- input_specs=input_specs,
- backbone_config=segmentation_config.backbone,
- norm_activation_config=norm_activation_config,
- l2_regularizer=l2_regularizer)
- segmentation_decoder_input_specs = segmentation_backbone.output_specs
- else:
- segmentation_backbone = None
- segmentation_decoder_input_specs = maskrcnn_model.backbone.output_specs
-
- if not model_config.shared_decoder:
- segmentation_decoder = decoder_factory.build_decoder(
- input_specs=segmentation_decoder_input_specs,
- model_config=segmentation_config,
- l2_regularizer=l2_regularizer)
- decoder_config = segmentation_decoder.get_config()
- else:
- segmentation_decoder = None
- decoder_config = maskrcnn_model.decoder.get_config()
-
- segmentation_head_config = segmentation_config.head
- detection_head_config = model_config.detection_head
- postprocessing_config = model_config.panoptic_segmentation_generator
-
- segmentation_head = segmentation_heads.SegmentationHead(
- num_classes=segmentation_config.num_classes,
- level=segmentation_head_config.level,
- num_convs=segmentation_head_config.num_convs,
- prediction_kernel_size=segmentation_head_config.prediction_kernel_size,
- num_filters=segmentation_head_config.num_filters,
- upsample_factor=segmentation_head_config.upsample_factor,
- feature_fusion=segmentation_head_config.feature_fusion,
- decoder_min_level=segmentation_head_config.decoder_min_level,
- decoder_max_level=segmentation_head_config.decoder_max_level,
- low_level=segmentation_head_config.low_level,
- low_level_num_filters=segmentation_head_config.low_level_num_filters,
- activation=norm_activation_config.activation,
- use_sync_bn=norm_activation_config.use_sync_bn,
- norm_momentum=norm_activation_config.norm_momentum,
- norm_epsilon=norm_activation_config.norm_epsilon,
- num_decoder_filters=decoder_config['num_filters'],
- kernel_regularizer=l2_regularizer)
-
- if model_config.generate_panoptic_masks:
- max_num_detections = model_config.detection_generator.max_num_detections
- mask_binarize_threshold = postprocessing_config.mask_binarize_threshold
- panoptic_segmentation_generator_obj = panoptic_segmentation_generator.PanopticSegmentationGenerator(
- output_size=postprocessing_config.output_size,
- max_num_detections=max_num_detections,
- stuff_classes_offset=model_config.stuff_classes_offset,
- mask_binarize_threshold=mask_binarize_threshold,
- score_threshold=postprocessing_config.score_threshold,
- things_overlap_threshold=postprocessing_config.things_overlap_threshold,
- things_class_label=postprocessing_config.things_class_label,
- stuff_area_threshold=postprocessing_config.stuff_area_threshold,
- void_class_label=postprocessing_config.void_class_label,
- void_instance_id=postprocessing_config.void_instance_id,
- rescale_predictions=postprocessing_config.rescale_predictions)
- else:
- panoptic_segmentation_generator_obj = None
-
- # Combines maskrcnn, and segmentation models to build panoptic segmentation
- # model.
-
- model = panoptic_maskrcnn_model.PanopticMaskRCNNModel(
- backbone=maskrcnn_model.backbone,
- decoder=maskrcnn_model.decoder,
- rpn_head=maskrcnn_model.rpn_head,
- detection_head=maskrcnn_model.detection_head,
- roi_generator=maskrcnn_model.roi_generator,
- roi_sampler=maskrcnn_model.roi_sampler,
- roi_aligner=maskrcnn_model.roi_aligner,
- detection_generator=maskrcnn_model.detection_generator,
- panoptic_segmentation_generator=panoptic_segmentation_generator_obj,
- mask_head=maskrcnn_model.mask_head,
- mask_sampler=maskrcnn_model.mask_sampler,
- mask_roi_aligner=maskrcnn_model.mask_roi_aligner,
- segmentation_backbone=segmentation_backbone,
- segmentation_decoder=segmentation_decoder,
- segmentation_head=segmentation_head,
- class_agnostic_bbox_pred=detection_head_config.class_agnostic_bbox_pred,
- cascade_class_ensemble=detection_head_config.cascade_class_ensemble,
- min_level=model_config.min_level,
- max_level=model_config.max_level,
- num_scales=model_config.anchor.num_scales,
- aspect_ratios=model_config.anchor.aspect_ratios,
- anchor_size=model_config.anchor.anchor_size)
- return model
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/modeling/factory_test.py b/official/vision/beta/projects/panoptic_maskrcnn/modeling/factory_test.py
deleted file mode 100644
index 9a70d1f5721..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/modeling/factory_test.py
+++ /dev/null
@@ -1,66 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Tests for factory.py."""
-
-from absl.testing import parameterized
-import numpy as np
-import tensorflow as tf
-
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as panoptic_maskrcnn_cfg
-from official.vision.beta.projects.panoptic_maskrcnn.modeling import factory
-from official.vision.configs import backbones
-from official.vision.configs import decoders
-from official.vision.configs import semantic_segmentation
-
-
-class PanopticMaskRCNNBuilderTest(parameterized.TestCase, tf.test.TestCase):
-
- @parameterized.parameters(
- ('resnet', (640, 640), 'dilated_resnet', 'fpn'),
- ('resnet', (640, 640), 'dilated_resnet', 'aspp'),
- ('resnet', (640, 640), None, 'fpn'),
- ('resnet', (640, 640), None, 'aspp'),
- ('resnet', (640, 640), None, None),
- ('resnet', (None, None), 'dilated_resnet', 'fpn'),
- ('resnet', (None, None), 'dilated_resnet', 'aspp'),
- ('resnet', (None, None), None, 'fpn'),
- ('resnet', (None, None), None, 'aspp'),
- ('resnet', (None, None), None, None))
- def test_builder(self, backbone_type, input_size, segmentation_backbone_type,
- segmentation_decoder_type):
- num_classes = 2
- input_specs = tf.keras.layers.InputSpec(
- shape=[None, input_size[0], input_size[1], 3])
- segmentation_output_stride = 16
- level = int(np.math.log2(segmentation_output_stride))
- segmentation_model = semantic_segmentation.SemanticSegmentationModel(
- num_classes=2,
- backbone=backbones.Backbone(type=segmentation_backbone_type),
- decoder=decoders.Decoder(type=segmentation_decoder_type),
- head=semantic_segmentation.SegmentationHead(level=level))
- model_config = panoptic_maskrcnn_cfg.PanopticMaskRCNN(
- num_classes=num_classes,
- segmentation_model=segmentation_model,
- backbone=backbones.Backbone(type=backbone_type),
- shared_backbone=segmentation_backbone_type is None,
- shared_decoder=segmentation_decoder_type is None)
- l2_regularizer = tf.keras.regularizers.l2(5e-5)
- _ = factory.build_panoptic_maskrcnn(
- input_specs=input_specs,
- model_config=model_config,
- l2_regularizer=l2_regularizer)
-
-if __name__ == '__main__':
- tf.test.main()
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/panoptic_segmentation_generator.py b/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/panoptic_segmentation_generator.py
deleted file mode 100644
index 69900866b35..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/panoptic_segmentation_generator.py
+++ /dev/null
@@ -1,321 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Contains definition for postprocessing layer to genrate panoptic segmentations."""
-
-from typing import List, Optional
-
-import tensorflow as tf
-
-from official.vision.beta.projects.panoptic_maskrcnn.modeling.layers import paste_masks
-
-
-class PanopticSegmentationGenerator(tf.keras.layers.Layer):
- """Panoptic segmentation generator layer."""
-
- def __init__(
- self,
- output_size: List[int],
- max_num_detections: int,
- stuff_classes_offset: int,
- mask_binarize_threshold: float = 0.5,
- score_threshold: float = 0.5,
- things_overlap_threshold: float = 0.5,
- stuff_area_threshold: float = 4096,
- things_class_label: int = 1,
- void_class_label: int = 0,
- void_instance_id: int = -1,
- rescale_predictions: bool = False,
- **kwargs):
- """Generates panoptic segmentation masks.
-
- Args:
- output_size: A `List` of integers that represent the height and width of
- the output mask.
- max_num_detections: `int` for maximum number of detections.
- stuff_classes_offset: An `int` that is added to the output of the
- semantic segmentation mask to make sure that the stuff class ids do not
- ovelap with the thing class ids of the MaskRCNN outputs.
- mask_binarize_threshold: A `float`
- score_threshold: A `float` representing the threshold for deciding
- when to remove objects based on score.
- things_overlap_threshold: A `float` representing a threshold for deciding
- to ignore a thing if overlap is above the threshold.
- stuff_area_threshold: A `float` representing a threshold for deciding to
- to ignore a stuff class if area is below certain threshold.
- things_class_label: An `int` that represents a single merged category of
- all thing classes in the semantic segmentation output.
- void_class_label: An `int` that is used to represent empty or unlabelled
- regions of the mask
- void_instance_id: An `int` that is used to denote regions that are not
- assigned to any thing class. That is, void_instance_id are assigned to
- both stuff regions and empty regions.
- rescale_predictions: `bool`, whether to scale back prediction to original
- image sizes. If True, image_info is used to rescale predictions.
- **kwargs: additional kewargs arguments.
- """
- self._output_size = output_size
- self._max_num_detections = max_num_detections
- self._stuff_classes_offset = stuff_classes_offset
- self._mask_binarize_threshold = mask_binarize_threshold
- self._score_threshold = score_threshold
- self._things_overlap_threshold = things_overlap_threshold
- self._stuff_area_threshold = stuff_area_threshold
- self._things_class_label = things_class_label
- self._void_class_label = void_class_label
- self._void_instance_id = void_instance_id
- self._rescale_predictions = rescale_predictions
-
- self._config_dict = {
- 'output_size': output_size,
- 'max_num_detections': max_num_detections,
- 'stuff_classes_offset': stuff_classes_offset,
- 'mask_binarize_threshold': mask_binarize_threshold,
- 'score_threshold': score_threshold,
- 'things_class_label': things_class_label,
- 'void_class_label': void_class_label,
- 'void_instance_id': void_instance_id,
- 'rescale_predictions': rescale_predictions
- }
- super(PanopticSegmentationGenerator, self).__init__(**kwargs)
-
- def build(self, input_shape):
- grid_sampler = paste_masks.BilinearGridSampler(align_corners=False)
- self._paste_masks_fn = paste_masks.PasteMasks(
- output_size=self._output_size, grid_sampler=grid_sampler)
-
- def _generate_panoptic_masks(self, boxes, scores, classes, detections_masks,
- segmentation_mask):
- """Generates panoptic masks for a single image.
-
- This function implements the following steps to merge instance and semantic
- segmentation masks described in https://arxiv.org/pdf/1901.02446.pdf
- Steps:
- 1. resolving overlaps between different instances based on their
- confidence scores
- 2. resolving overlaps between instance and semantic segmentation
- outputs in favor of instances
- 3. removing any stuff regions labeled other or under a given area
- threshold.
- Args:
- boxes: A `tf.Tensor` of shape [num_rois, 4], representing the bounding
- boxes for detected objects.
- scores: A `tf.Tensor` of shape [num_rois], representing the
- confidence scores for each object.
- classes: A `tf.Tensor` of shape [num_rois], representing the class
- for each object.
- detections_masks: A `tf.Tensor` of shape
- [num_rois, mask_height, mask_width, 1], representing the cropped mask
- for each object.
- segmentation_mask: A `tf.Tensor` of shape [height, width], representing
- the semantic segmentation output.
- Returns:
- Dict with the following keys:
- - category_mask: A `tf.Tensor` for category masks.
- - instance_mask: A `tf.Tensor for instance masks.
- """
-
- # Offset stuff class predictions
- segmentation_mask = tf.where(
- tf.logical_or(
- tf.equal(segmentation_mask, self._things_class_label),
- tf.equal(segmentation_mask, self._void_class_label)),
- segmentation_mask,
- segmentation_mask + self._stuff_classes_offset
- )
- # sort instances by their scores
- sorted_indices = tf.argsort(scores, direction='DESCENDING')
-
- mask_shape = self._output_size + [1]
- category_mask = tf.ones(mask_shape,
- dtype=tf.float32) * self._void_class_label
- instance_mask = tf.ones(
- mask_shape, dtype=tf.float32) * self._void_instance_id
-
- # filter instances with low confidence
- sorted_scores = tf.sort(scores, direction='DESCENDING')
-
- valid_indices = tf.where(sorted_scores > self._score_threshold)
-
- # if no instance has sufficient confidence score, skip merging
- # instance segmentation masks
- if tf.shape(valid_indices)[0] > 0:
- loop_end_idx = valid_indices[-1, 0] + 1
- loop_end_idx = tf.minimum(
- tf.cast(loop_end_idx, dtype=tf.int32),
- self._max_num_detections)
- pasted_masks = self._paste_masks_fn((
- detections_masks[:loop_end_idx],
- boxes[:loop_end_idx]))
-
- # add things segmentation to panoptic masks
- for i in range(loop_end_idx):
- # we process instances in decending order, which will make sure
- # the overlaps are resolved based on confidence score
- instance_idx = sorted_indices[i]
-
- pasted_mask = pasted_masks[instance_idx]
-
- class_id = tf.cast(classes[instance_idx], dtype=tf.float32)
-
- # convert sigmoid scores to binary values
- binary_mask = tf.greater(
- pasted_mask, self._mask_binarize_threshold)
-
- # filter empty instance masks
- if not tf.reduce_sum(tf.cast(binary_mask, tf.float32)) > 0:
- continue
-
- overlap = tf.logical_and(
- binary_mask,
- tf.not_equal(category_mask, self._void_class_label))
- binary_mask_area = tf.reduce_sum(
- tf.cast(binary_mask, dtype=tf.float32))
- overlap_area = tf.reduce_sum(
- tf.cast(overlap, dtype=tf.float32))
-
- # skip instance that have a big enough overlap with instances with
- # higer scores
- if overlap_area / binary_mask_area > self._things_overlap_threshold:
- continue
-
- # fill empty regions in category_mask represented by
- # void_class_label with class_id of the instance.
- category_mask = tf.where(
- tf.logical_and(
- binary_mask, tf.equal(category_mask, self._void_class_label)),
- tf.ones_like(category_mask) * class_id, category_mask)
-
- # fill empty regions in the instance_mask represented by
- # void_instance_id with the id of the instance, starting from 1
- instance_mask = tf.where(
- tf.logical_and(
- binary_mask,
- tf.equal(instance_mask, self._void_instance_id)),
- tf.ones_like(instance_mask) *
- tf.cast(instance_idx + 1, tf.float32), instance_mask)
-
- stuff_class_ids = tf.unique(tf.reshape(segmentation_mask, [-1])).y
- for stuff_class_id in stuff_class_ids:
- if stuff_class_id == self._things_class_label:
- continue
-
- stuff_mask = tf.logical_and(
- tf.equal(segmentation_mask, stuff_class_id),
- tf.equal(category_mask, self._void_class_label))
-
- stuff_mask_area = tf.reduce_sum(
- tf.cast(stuff_mask, dtype=tf.float32))
-
- if stuff_mask_area < self._stuff_area_threshold:
- continue
-
- category_mask = tf.where(
- stuff_mask,
- tf.ones_like(category_mask) * stuff_class_id,
- category_mask)
-
- results = {
- 'category_mask': category_mask[:, :, 0],
- 'instance_mask': instance_mask[:, :, 0]
- }
- return results
-
- def _resize_and_pad_masks(self, mask, image_info):
- """Resizes masks to match the original image shape and pads to`output_size`.
-
- Args:
- mask: a padded mask tensor.
- image_info: a tensor that holds information about original and
- preprocessed images.
- Returns:
- resized and padded masks: tf.Tensor.
- """
- rescale_size = tf.cast(
- tf.math.ceil(image_info[1, :] / image_info[2, :]), tf.int32)
- image_shape = tf.cast(image_info[0, :], tf.int32)
- offsets = tf.cast(image_info[3, :], tf.int32)
-
- mask = tf.image.resize(
- mask,
- rescale_size,
- method='bilinear')
- mask = tf.image.crop_to_bounding_box(
- mask,
- offsets[0], offsets[1],
- image_shape[0],
- image_shape[1])
- mask = tf.image.pad_to_bounding_box(
- mask, 0, 0, self._output_size[0], self._output_size[1])
- return mask
-
- def call(self, inputs: tf.Tensor, image_info: Optional[tf.Tensor] = None):
- detections = inputs
-
- batched_scores = detections['detection_scores']
- batched_classes = detections['detection_classes']
- batched_detections_masks = tf.expand_dims(
- detections['detection_masks'], axis=-1)
- batched_boxes = detections['detection_boxes']
- batched_segmentation_masks = tf.cast(
- detections['segmentation_outputs'], dtype=tf.float32)
-
- if self._rescale_predictions:
- scale = tf.tile(
- tf.cast(image_info[:, 2:3, :], dtype=batched_boxes.dtype),
- multiples=[1, 1, 2])
- batched_boxes /= scale
-
- batched_segmentation_masks = tf.map_fn(
- fn=lambda x: self._resize_and_pad_masks(x[0], x[1]),
- elems=(
- batched_segmentation_masks,
- image_info),
- fn_output_signature=tf.float32,
- parallel_iterations=32)
- else:
- batched_segmentation_masks = tf.image.resize(
- batched_segmentation_masks,
- size=self._output_size,
- method='bilinear')
-
- batched_segmentation_masks = tf.expand_dims(tf.cast(
- tf.argmax(batched_segmentation_masks, axis=-1),
- dtype=tf.float32), axis=-1)
-
- panoptic_masks = tf.map_fn(
- fn=lambda x: self._generate_panoptic_masks( # pylint:disable=g-long-lambda
- x[0], x[1], x[2], x[3], x[4]),
- elems=(
- batched_boxes,
- batched_scores,
- batched_classes,
- batched_detections_masks,
- batched_segmentation_masks),
- fn_output_signature={
- 'category_mask': tf.float32,
- 'instance_mask': tf.float32
- }, parallel_iterations=32)
-
- for k, v in panoptic_masks.items():
- panoptic_masks[k] = tf.cast(v, dtype=tf.int32)
-
- return panoptic_masks
-
- def get_config(self):
- return self._config_dict
-
- @classmethod
- def from_config(cls, config):
- return cls(**config)
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/panoptic_segmentation_generator_test.py b/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/panoptic_segmentation_generator_test.py
deleted file mode 100644
index d746f3d8dae..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/modeling/layers/panoptic_segmentation_generator_test.py
+++ /dev/null
@@ -1,145 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Tests for panoptic_segmentation_generator.py."""
-
-from absl.testing import parameterized
-import numpy as np
-import tensorflow as tf
-from tensorflow.python.distribute import combinations
-from tensorflow.python.distribute import strategy_combinations
-
-from official.vision.beta.projects.panoptic_maskrcnn.modeling.layers import panoptic_segmentation_generator
-
-PANOPTIC_SEGMENTATION_GENERATOR = panoptic_segmentation_generator.PanopticSegmentationGenerator
-
-
-class PanopticSegmentationGeneratorTest(
- parameterized.TestCase, tf.test.TestCase):
-
- def test_serialize_deserialize(self):
- config = {
- 'output_size': [640, 640],
- 'max_num_detections': 100,
- 'stuff_classes_offset': 90,
- 'mask_binarize_threshold': 0.5,
- 'score_threshold': 0.005,
- 'things_class_label': 1,
- 'void_class_label': 0,
- 'void_instance_id': -1,
- 'rescale_predictions': False,
- }
- generator = PANOPTIC_SEGMENTATION_GENERATOR(**config)
-
- expected_config = dict(config)
- self.assertEqual(generator.get_config(), expected_config)
-
- new_generator = PANOPTIC_SEGMENTATION_GENERATOR.from_config(
- generator.get_config())
-
- self.assertAllEqual(generator.get_config(), new_generator.get_config())
-
- @combinations.generate(
- combinations.combine(
- strategy=[
- strategy_combinations.default_strategy,
- strategy_combinations.one_device_strategy_gpu,
- ]))
- def test_outputs(self, strategy):
-
- # 0 represents the void class label
- thing_class_ids = [0, 1, 2, 3, 4]
- stuff_class_ids = [0, 5, 6, 7, 8, 9, 10]
- all_class_ids = set(thing_class_ids + stuff_class_ids)
-
- num_thing_classes = len(thing_class_ids)
- num_stuff_classes = len(stuff_class_ids)
- num_classes_for_segmentation = num_stuff_classes + 1
-
- # all thing classes are mapped to class_id=1, stuff class ids are offset
- # such that the stuff class_ids start from 2, this means the semantic
- # segmentation head will have ground truths with class_ids belonging to
- # [0, 1, 2, 3, 4, 5, 6, 7]
-
- config = {
- 'output_size': [640, 640],
- 'max_num_detections': 100,
- 'stuff_classes_offset': 3,
- 'mask_binarize_threshold': 0.5,
- 'score_threshold': 0.005,
- 'things_class_label': 1,
- 'void_class_label': 0,
- 'void_instance_id': -1,
- 'rescale_predictions': False,
- }
- generator = PANOPTIC_SEGMENTATION_GENERATOR(**config)
-
- crop_height = 112
- crop_width = 112
-
- boxes = tf.constant([[
- [167, 398, 342, 619],
- [192, 171, 363, 449],
- [211, 1, 382, 74]
- ]])
-
- num_detections = boxes.get_shape().as_list()[1]
- scores = tf.random.uniform([1, num_detections], 0, 1)
- classes = tf.random.uniform(
- [1, num_detections],
- 1, num_thing_classes, dtype=tf.int32)
- masks = tf.random.normal(
- [1, num_detections, crop_height, crop_width])
-
- segmentation_mask = tf.random.uniform(
- [1, *config['output_size']],
- 0, num_classes_for_segmentation, dtype=tf.int32)
- segmentation_mask_one_hot = tf.one_hot(
- segmentation_mask, depth=num_stuff_classes + 1)
-
- inputs = {
- 'detection_boxes': boxes,
- 'detection_scores': scores,
- 'detection_classes': classes,
- 'detection_masks': masks,
- 'num_detections': tf.constant([num_detections]),
- 'segmentation_outputs': segmentation_mask_one_hot
- }
-
- def _run(inputs):
- return generator(inputs=inputs)
-
- @tf.function
- def _distributed_run(inputs):
- outputs = strategy.run(_run, args=((inputs,)))
- return strategy.gather(outputs, axis=0)
-
- outputs = _distributed_run(inputs)
-
- self.assertIn('category_mask', outputs)
- self.assertIn('instance_mask', outputs)
-
- self.assertAllEqual(
- outputs['category_mask'][0].get_shape().as_list(),
- config['output_size'])
-
- self.assertAllEqual(
- outputs['instance_mask'][0].get_shape().as_list(),
- config['output_size'])
-
- for category_id in np.unique(outputs['category_mask']):
- self.assertIn(category_id, all_class_ids)
-
-if __name__ == '__main__':
- tf.test.main()
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/modeling/panoptic_maskrcnn_model_test.py b/official/vision/beta/projects/panoptic_maskrcnn/modeling/panoptic_maskrcnn_model_test.py
deleted file mode 100644
index c8ba8366064..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/modeling/panoptic_maskrcnn_model_test.py
+++ /dev/null
@@ -1,554 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Tests for panoptic_maskrcnn_model.py."""
-
-import os
-
-from absl.testing import parameterized
-import tensorflow as tf
-
-from tensorflow.python.distribute import combinations
-from tensorflow.python.distribute import strategy_combinations
-from official.vision.beta.projects.panoptic_maskrcnn.modeling import panoptic_maskrcnn_model
-from official.vision.beta.projects.panoptic_maskrcnn.modeling.layers import panoptic_segmentation_generator
-from official.vision.modeling.backbones import resnet
-from official.vision.modeling.decoders import aspp
-from official.vision.modeling.decoders import fpn
-from official.vision.modeling.heads import dense_prediction_heads
-from official.vision.modeling.heads import instance_heads
-from official.vision.modeling.heads import segmentation_heads
-from official.vision.modeling.layers import detection_generator
-from official.vision.modeling.layers import mask_sampler
-from official.vision.modeling.layers import roi_aligner
-from official.vision.modeling.layers import roi_generator
-from official.vision.modeling.layers import roi_sampler
-from official.vision.ops import anchor
-
-
-class PanopticMaskRCNNModelTest(parameterized.TestCase, tf.test.TestCase):
-
- @combinations.generate(
- combinations.combine(
- use_separable_conv=[True, False],
- build_anchor_boxes=[True, False],
- shared_backbone=[True, False],
- shared_decoder=[True, False],
- is_training=[True,]))
- def test_build_model(self,
- use_separable_conv,
- build_anchor_boxes,
- shared_backbone,
- shared_decoder,
- is_training=True):
- num_classes = 3
- min_level = 2
- max_level = 6
- num_scales = 3
- aspect_ratios = [1.0]
- anchor_size = 3
- resnet_model_id = 50
- segmentation_resnet_model_id = 50
- aspp_dilation_rates = [6, 12, 18]
- aspp_decoder_level = 2
- fpn_decoder_level = 2
- num_anchors_per_location = num_scales * len(aspect_ratios)
- image_size = 128
- images = tf.random.normal([2, image_size, image_size, 3])
- image_info = tf.convert_to_tensor(
- [[[image_size, image_size], [image_size, image_size], [1, 1], [0, 0]],
- [[image_size, image_size], [image_size, image_size], [1, 1], [0, 0]]])
- shared_decoder = shared_decoder and shared_backbone
- if build_anchor_boxes or not is_training:
- anchor_boxes = anchor.Anchor(
- min_level=min_level,
- max_level=max_level,
- num_scales=num_scales,
- aspect_ratios=aspect_ratios,
- anchor_size=3,
- image_size=(image_size, image_size)).multilevel_boxes
- for l in anchor_boxes:
- anchor_boxes[l] = tf.tile(
- tf.expand_dims(anchor_boxes[l], axis=0), [2, 1, 1, 1])
- else:
- anchor_boxes = None
-
- backbone = resnet.ResNet(model_id=resnet_model_id)
- decoder = fpn.FPN(
- input_specs=backbone.output_specs,
- min_level=min_level,
- max_level=max_level,
- use_separable_conv=use_separable_conv)
- rpn_head = dense_prediction_heads.RPNHead(
- min_level=min_level,
- max_level=max_level,
- num_anchors_per_location=num_anchors_per_location,
- num_convs=1)
- detection_head = instance_heads.DetectionHead(num_classes=num_classes)
- roi_generator_obj = roi_generator.MultilevelROIGenerator()
- roi_sampler_obj = roi_sampler.ROISampler()
- roi_aligner_obj = roi_aligner.MultilevelROIAligner()
- detection_generator_obj = detection_generator.DetectionGenerator()
- panoptic_segmentation_generator_obj = panoptic_segmentation_generator.PanopticSegmentationGenerator(
- output_size=[image_size, image_size],
- max_num_detections=100,
- stuff_classes_offset=90)
- mask_head = instance_heads.MaskHead(
- num_classes=num_classes, upsample_factor=2)
- mask_sampler_obj = mask_sampler.MaskSampler(
- mask_target_size=28, num_sampled_masks=1)
- mask_roi_aligner_obj = roi_aligner.MultilevelROIAligner(crop_size=14)
-
- if shared_backbone:
- segmentation_backbone = None
- else:
- segmentation_backbone = resnet.ResNet(
- model_id=segmentation_resnet_model_id)
- if not shared_decoder:
- feature_fusion = 'deeplabv3plus'
- level = aspp_decoder_level
- segmentation_decoder = aspp.ASPP(
- level=level, dilation_rates=aspp_dilation_rates)
- else:
- feature_fusion = 'panoptic_fpn_fusion'
- level = fpn_decoder_level
- segmentation_decoder = None
- segmentation_head = segmentation_heads.SegmentationHead(
- num_classes=2, # stuff and common class for things,
- level=level,
- feature_fusion=feature_fusion,
- decoder_min_level=min_level,
- decoder_max_level=max_level,
- num_convs=2)
-
- model = panoptic_maskrcnn_model.PanopticMaskRCNNModel(
- backbone,
- decoder,
- rpn_head,
- detection_head,
- roi_generator_obj,
- roi_sampler_obj,
- roi_aligner_obj,
- detection_generator_obj,
- panoptic_segmentation_generator_obj,
- mask_head,
- mask_sampler_obj,
- mask_roi_aligner_obj,
- segmentation_backbone=segmentation_backbone,
- segmentation_decoder=segmentation_decoder,
- segmentation_head=segmentation_head,
- min_level=min_level,
- max_level=max_level,
- num_scales=num_scales,
- aspect_ratios=aspect_ratios,
- anchor_size=anchor_size)
-
- gt_boxes = tf.convert_to_tensor(
- [[[10, 10, 15, 15], [2.5, 2.5, 7.5, 7.5], [-1, -1, -1, -1]],
- [[100, 100, 150, 150], [-1, -1, -1, -1], [-1, -1, -1, -1]]],
- dtype=tf.float32)
- gt_classes = tf.convert_to_tensor([[2, 1, -1], [1, -1, -1]], dtype=tf.int32)
- gt_masks = tf.ones((2, 3, 100, 100))
-
- # Results will be checked in test_forward.
- _ = model(
- images,
- image_info,
- anchor_boxes,
- gt_boxes,
- gt_classes,
- gt_masks,
- training=is_training)
-
- @combinations.generate(
- combinations.combine(
- strategy=[
- strategy_combinations.one_device_strategy,
- strategy_combinations.one_device_strategy_gpu,
- ],
- shared_backbone=[True, False],
- shared_decoder=[True, False],
- training=[True, False],
- generate_panoptic_masks=[True, False]))
- def test_forward(self, strategy, training,
- shared_backbone, shared_decoder,
- generate_panoptic_masks):
- num_classes = 3
- min_level = 2
- max_level = 6
- num_scales = 3
- aspect_ratios = [1.0]
- anchor_size = 3
- segmentation_resnet_model_id = 101
- aspp_dilation_rates = [6, 12, 18]
- aspp_decoder_level = 2
- fpn_decoder_level = 2
-
- class_agnostic_bbox_pred = False
- cascade_class_ensemble = False
-
- image_size = (256, 256)
- images = tf.random.normal([2, image_size[0], image_size[1], 3])
- image_info = tf.convert_to_tensor(
- [[[224, 100], [224, 100], [1, 1], [0, 0]],
- [[224, 100], [224, 100], [1, 1], [0, 0]]])
- shared_decoder = shared_decoder and shared_backbone
- with strategy.scope():
-
- anchor_boxes = anchor.Anchor(
- min_level=min_level,
- max_level=max_level,
- num_scales=num_scales,
- aspect_ratios=aspect_ratios,
- anchor_size=anchor_size,
- image_size=image_size).multilevel_boxes
-
- num_anchors_per_location = len(aspect_ratios) * num_scales
-
- input_specs = tf.keras.layers.InputSpec(shape=[None, None, None, 3])
- backbone = resnet.ResNet(model_id=50, input_specs=input_specs)
- decoder = fpn.FPN(
- min_level=min_level,
- max_level=max_level,
- input_specs=backbone.output_specs)
- rpn_head = dense_prediction_heads.RPNHead(
- min_level=min_level,
- max_level=max_level,
- num_anchors_per_location=num_anchors_per_location)
- detection_head = instance_heads.DetectionHead(
- num_classes=num_classes,
- class_agnostic_bbox_pred=class_agnostic_bbox_pred)
- roi_generator_obj = roi_generator.MultilevelROIGenerator()
-
- roi_sampler_cascade = []
- roi_sampler_obj = roi_sampler.ROISampler()
- roi_sampler_cascade.append(roi_sampler_obj)
- roi_aligner_obj = roi_aligner.MultilevelROIAligner()
- detection_generator_obj = detection_generator.DetectionGenerator()
-
- if generate_panoptic_masks:
- panoptic_segmentation_generator_obj = panoptic_segmentation_generator.PanopticSegmentationGenerator(
- output_size=list(image_size),
- max_num_detections=100,
- stuff_classes_offset=90)
- else:
- panoptic_segmentation_generator_obj = None
-
- mask_head = instance_heads.MaskHead(
- num_classes=num_classes, upsample_factor=2)
- mask_sampler_obj = mask_sampler.MaskSampler(
- mask_target_size=28, num_sampled_masks=1)
- mask_roi_aligner_obj = roi_aligner.MultilevelROIAligner(crop_size=14)
-
- if shared_backbone:
- segmentation_backbone = None
- else:
- segmentation_backbone = resnet.ResNet(
- model_id=segmentation_resnet_model_id)
- if not shared_decoder:
- feature_fusion = 'deeplabv3plus'
- level = aspp_decoder_level
- segmentation_decoder = aspp.ASPP(
- level=level, dilation_rates=aspp_dilation_rates)
- else:
- feature_fusion = 'panoptic_fpn_fusion'
- level = fpn_decoder_level
- segmentation_decoder = None
- segmentation_head = segmentation_heads.SegmentationHead(
- num_classes=2, # stuff and common class for things,
- level=level,
- feature_fusion=feature_fusion,
- decoder_min_level=min_level,
- decoder_max_level=max_level,
- num_convs=2)
-
- model = panoptic_maskrcnn_model.PanopticMaskRCNNModel(
- backbone,
- decoder,
- rpn_head,
- detection_head,
- roi_generator_obj,
- roi_sampler_obj,
- roi_aligner_obj,
- detection_generator_obj,
- panoptic_segmentation_generator_obj,
- mask_head,
- mask_sampler_obj,
- mask_roi_aligner_obj,
- segmentation_backbone=segmentation_backbone,
- segmentation_decoder=segmentation_decoder,
- segmentation_head=segmentation_head,
- class_agnostic_bbox_pred=class_agnostic_bbox_pred,
- cascade_class_ensemble=cascade_class_ensemble,
- min_level=min_level,
- max_level=max_level,
- num_scales=num_scales,
- aspect_ratios=aspect_ratios,
- anchor_size=anchor_size)
-
- gt_boxes = tf.convert_to_tensor(
- [[[10, 10, 15, 15], [2.5, 2.5, 7.5, 7.5], [-1, -1, -1, -1]],
- [[100, 100, 150, 150], [-1, -1, -1, -1], [-1, -1, -1, -1]]],
- dtype=tf.float32)
- gt_classes = tf.convert_to_tensor(
- [[2, 1, -1], [1, -1, -1]], dtype=tf.int32)
- gt_masks = tf.ones((2, 3, 100, 100))
-
- results = model(
- images,
- image_info,
- anchor_boxes,
- gt_boxes,
- gt_classes,
- gt_masks,
- training=training)
-
- self.assertIn('rpn_boxes', results)
- self.assertIn('rpn_scores', results)
- if training:
- self.assertIn('class_targets', results)
- self.assertIn('box_targets', results)
- self.assertIn('class_outputs', results)
- self.assertIn('box_outputs', results)
- self.assertIn('mask_outputs', results)
- else:
- self.assertIn('detection_boxes', results)
- self.assertIn('detection_scores', results)
- self.assertIn('detection_classes', results)
- self.assertIn('num_detections', results)
- self.assertIn('detection_masks', results)
- self.assertIn('segmentation_outputs', results)
-
- self.assertAllEqual(
- [2, image_size[0] // (2**level), image_size[1] // (2**level), 2],
- results['segmentation_outputs'].numpy().shape)
-
- if generate_panoptic_masks:
- self.assertIn('panoptic_outputs', results)
- self.assertIn('category_mask', results['panoptic_outputs'])
- self.assertIn('instance_mask', results['panoptic_outputs'])
- self.assertAllEqual(
- [2, image_size[0], image_size[1]],
- results['panoptic_outputs']['category_mask'].numpy().shape)
- self.assertAllEqual(
- [2, image_size[0], image_size[1]],
- results['panoptic_outputs']['instance_mask'].numpy().shape)
- else:
- self.assertNotIn('panoptic_outputs', results)
-
- @combinations.generate(
- combinations.combine(
- shared_backbone=[True, False], shared_decoder=[True, False]))
- def test_serialize_deserialize(self, shared_backbone, shared_decoder):
- input_specs = tf.keras.layers.InputSpec(shape=[None, None, None, 3])
- backbone = resnet.ResNet(model_id=50, input_specs=input_specs)
- decoder = fpn.FPN(
- min_level=3, max_level=7, input_specs=backbone.output_specs)
- rpn_head = dense_prediction_heads.RPNHead(
- min_level=3, max_level=7, num_anchors_per_location=3)
- detection_head = instance_heads.DetectionHead(num_classes=2)
- roi_generator_obj = roi_generator.MultilevelROIGenerator()
- roi_sampler_obj = roi_sampler.ROISampler()
- roi_aligner_obj = roi_aligner.MultilevelROIAligner()
- detection_generator_obj = detection_generator.DetectionGenerator()
- panoptic_segmentation_generator_obj = panoptic_segmentation_generator.PanopticSegmentationGenerator(
- output_size=[None, None],
- max_num_detections=100,
- stuff_classes_offset=90)
- segmentation_resnet_model_id = 101
- aspp_dilation_rates = [6, 12, 18]
- min_level = 2
- max_level = 6
- aspp_decoder_level = 2
- fpn_decoder_level = 2
- shared_decoder = shared_decoder and shared_backbone
- mask_head = instance_heads.MaskHead(num_classes=2, upsample_factor=2)
- mask_sampler_obj = mask_sampler.MaskSampler(
- mask_target_size=28, num_sampled_masks=1)
- mask_roi_aligner_obj = roi_aligner.MultilevelROIAligner(crop_size=14)
-
- if shared_backbone:
- segmentation_backbone = None
- else:
- segmentation_backbone = resnet.ResNet(
- model_id=segmentation_resnet_model_id)
- if not shared_decoder:
- feature_fusion = 'deeplabv3plus'
- level = aspp_decoder_level
- segmentation_decoder = aspp.ASPP(
- level=level, dilation_rates=aspp_dilation_rates)
- else:
- feature_fusion = 'panoptic_fpn_fusion'
- level = fpn_decoder_level
- segmentation_decoder = None
- segmentation_head = segmentation_heads.SegmentationHead(
- num_classes=2, # stuff and common class for things,
- level=level,
- feature_fusion=feature_fusion,
- decoder_min_level=min_level,
- decoder_max_level=max_level,
- num_convs=2)
-
- model = panoptic_maskrcnn_model.PanopticMaskRCNNModel(
- backbone,
- decoder,
- rpn_head,
- detection_head,
- roi_generator_obj,
- roi_sampler_obj,
- roi_aligner_obj,
- detection_generator_obj,
- panoptic_segmentation_generator_obj,
- mask_head,
- mask_sampler_obj,
- mask_roi_aligner_obj,
- segmentation_backbone=segmentation_backbone,
- segmentation_decoder=segmentation_decoder,
- segmentation_head=segmentation_head,
- min_level=min_level,
- max_level=max_level,
- num_scales=3,
- aspect_ratios=[1.0],
- anchor_size=3)
-
- config = model.get_config()
- new_model = panoptic_maskrcnn_model.PanopticMaskRCNNModel.from_config(
- config)
-
- # Validate that the config can be forced to JSON.
- _ = new_model.to_json()
-
- # If the serialization was successful, the new config should match the old.
- self.assertAllEqual(model.get_config(), new_model.get_config())
-
- @combinations.generate(
- combinations.combine(
- shared_backbone=[True, False], shared_decoder=[True, False]))
- def test_checkpoint(self, shared_backbone, shared_decoder):
- input_specs = tf.keras.layers.InputSpec(shape=[None, None, None, 3])
- backbone = resnet.ResNet(model_id=50, input_specs=input_specs)
- decoder = fpn.FPN(
- min_level=3, max_level=7, input_specs=backbone.output_specs)
- rpn_head = dense_prediction_heads.RPNHead(
- min_level=3, max_level=7, num_anchors_per_location=3)
- detection_head = instance_heads.DetectionHead(num_classes=2)
- roi_generator_obj = roi_generator.MultilevelROIGenerator()
- roi_sampler_obj = roi_sampler.ROISampler()
- roi_aligner_obj = roi_aligner.MultilevelROIAligner()
- detection_generator_obj = detection_generator.DetectionGenerator()
- panoptic_segmentation_generator_obj = panoptic_segmentation_generator.PanopticSegmentationGenerator(
- output_size=[None, None],
- max_num_detections=100,
- stuff_classes_offset=90)
- segmentation_resnet_model_id = 101
- aspp_dilation_rates = [6, 12, 18]
- min_level = 2
- max_level = 6
- aspp_decoder_level = 2
- fpn_decoder_level = 2
- shared_decoder = shared_decoder and shared_backbone
- mask_head = instance_heads.MaskHead(num_classes=2, upsample_factor=2)
- mask_sampler_obj = mask_sampler.MaskSampler(
- mask_target_size=28, num_sampled_masks=1)
- mask_roi_aligner_obj = roi_aligner.MultilevelROIAligner(crop_size=14)
-
- if shared_backbone:
- segmentation_backbone = None
- else:
- segmentation_backbone = resnet.ResNet(
- model_id=segmentation_resnet_model_id)
- if not shared_decoder:
- feature_fusion = 'deeplabv3plus'
- level = aspp_decoder_level
- segmentation_decoder = aspp.ASPP(
- level=level, dilation_rates=aspp_dilation_rates)
- else:
- feature_fusion = 'panoptic_fpn_fusion'
- level = fpn_decoder_level
- segmentation_decoder = None
- segmentation_head = segmentation_heads.SegmentationHead(
- num_classes=2, # stuff and common class for things,
- level=level,
- feature_fusion=feature_fusion,
- decoder_min_level=min_level,
- decoder_max_level=max_level,
- num_convs=2)
-
- model = panoptic_maskrcnn_model.PanopticMaskRCNNModel(
- backbone,
- decoder,
- rpn_head,
- detection_head,
- roi_generator_obj,
- roi_sampler_obj,
- roi_aligner_obj,
- detection_generator_obj,
- panoptic_segmentation_generator_obj,
- mask_head,
- mask_sampler_obj,
- mask_roi_aligner_obj,
- segmentation_backbone=segmentation_backbone,
- segmentation_decoder=segmentation_decoder,
- segmentation_head=segmentation_head,
- min_level=max_level,
- max_level=max_level,
- num_scales=3,
- aspect_ratios=[1.0],
- anchor_size=3)
- expect_checkpoint_items = dict(
- backbone=backbone,
- decoder=decoder,
- rpn_head=rpn_head,
- detection_head=[detection_head])
- expect_checkpoint_items['mask_head'] = mask_head
- if not shared_backbone:
- expect_checkpoint_items['segmentation_backbone'] = segmentation_backbone
- if not shared_decoder:
- expect_checkpoint_items['segmentation_decoder'] = segmentation_decoder
- expect_checkpoint_items['segmentation_head'] = segmentation_head
- self.assertAllEqual(expect_checkpoint_items, model.checkpoint_items)
-
- # Test save and load checkpoints.
- ckpt = tf.train.Checkpoint(model=model, **model.checkpoint_items)
- save_dir = self.create_tempdir().full_path
- ckpt.save(os.path.join(save_dir, 'ckpt'))
-
- partial_ckpt = tf.train.Checkpoint(backbone=backbone)
- partial_ckpt.read(tf.train.latest_checkpoint(
- save_dir)).expect_partial().assert_existing_objects_matched()
-
- partial_ckpt_mask = tf.train.Checkpoint(
- backbone=backbone, mask_head=mask_head)
- partial_ckpt_mask.restore(tf.train.latest_checkpoint(
- save_dir)).expect_partial().assert_existing_objects_matched()
-
- if not shared_backbone:
- partial_ckpt_segmentation = tf.train.Checkpoint(
- segmentation_backbone=segmentation_backbone,
- segmentation_decoder=segmentation_decoder,
- segmentation_head=segmentation_head)
- elif not shared_decoder:
- partial_ckpt_segmentation = tf.train.Checkpoint(
- segmentation_decoder=segmentation_decoder,
- segmentation_head=segmentation_head)
- else:
- partial_ckpt_segmentation = tf.train.Checkpoint(
- segmentation_head=segmentation_head)
-
- partial_ckpt_segmentation.restore(tf.train.latest_checkpoint(
- save_dir)).expect_partial().assert_existing_objects_matched()
-
-
-if __name__ == '__main__':
- tf.test.main()
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/serving/panoptic_segmentation_test.py b/official/vision/beta/projects/panoptic_maskrcnn/serving/panoptic_segmentation_test.py
deleted file mode 100644
index 51d14cfc2f1..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/serving/panoptic_segmentation_test.py
+++ /dev/null
@@ -1,126 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Test for panoptic image segmentation export lib."""
-import io
-import os
-
-from absl.testing import parameterized
-import numpy as np
-from PIL import Image
-import tensorflow as tf
-
-from official.core import exp_factory
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as cfg # pylint: disable=unused-import
-from official.vision.beta.projects.panoptic_maskrcnn.modeling import factory
-from official.vision.beta.projects.panoptic_maskrcnn.serving import panoptic_segmentation
-from official.vision.beta.projects.panoptic_maskrcnn.tasks import panoptic_maskrcnn as task # pylint: disable=unused-import
-
-
-class PanopticSegmentationExportTest(tf.test.TestCase, parameterized.TestCase):
-
- def _get_panoptic_segmentation_module(self, experiment_name):
- params = exp_factory.get_exp_config(experiment_name)
- params.task.model.backbone.resnet.model_id = 18
- params.task.model.detection_generator.nms_version = 'batched'
- input_specs = tf.keras.layers.InputSpec(shape=[1, 128, 128, 3])
- model = factory.build_panoptic_maskrcnn(
- input_specs=input_specs, model_config=params.task.model)
- panoptic_segmentation_module = panoptic_segmentation.PanopticSegmentationModule(
- params, model=model, batch_size=1, input_image_size=[128, 128])
- return panoptic_segmentation_module
-
- def _export_from_module(self, module, input_type, save_directory):
- signatures = module.get_inference_signatures(
- {input_type: 'serving_default'})
- tf.saved_model.save(module, save_directory, signatures=signatures)
-
- def _get_dummy_input(self, input_type, batch_size, image_size):
- """Get dummy input for the given input type."""
- h, w = image_size
-
- if input_type == 'image_tensor':
- return tf.zeros((batch_size, h, w, 3), dtype=np.uint8)
- elif input_type == 'image_bytes':
- image = Image.fromarray(np.zeros((h, w, 3), dtype=np.uint8))
- byte_io = io.BytesIO()
- image.save(byte_io, 'PNG')
- return [byte_io.getvalue() for b in range(batch_size)]
- elif input_type == 'tf_example':
- image_tensor = tf.zeros((h, w, 3), dtype=tf.uint8)
- encoded_jpeg = tf.image.encode_jpeg(tf.constant(image_tensor)).numpy()
- example = tf.train.Example(
- features=tf.train.Features(
- feature={
- 'image/encoded':
- tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[encoded_jpeg])),
- })).SerializeToString()
- return [example for b in range(batch_size)]
-
- @parameterized.parameters(
- ('image_tensor', 'panoptic_fpn_coco'),
- ('image_bytes', 'panoptic_fpn_coco'),
- ('tf_example', 'panoptic_fpn_coco'),
- )
- def test_export(self, input_type, experiment_name):
- tmp_dir = self.get_temp_dir()
- module = self._get_panoptic_segmentation_module(experiment_name)
-
- self._export_from_module(module, input_type, tmp_dir)
-
- self.assertTrue(os.path.exists(os.path.join(tmp_dir, 'saved_model.pb')))
- self.assertTrue(
- os.path.exists(os.path.join(tmp_dir, 'variables', 'variables.index')))
- self.assertTrue(
- os.path.exists(
- os.path.join(tmp_dir, 'variables',
- 'variables.data-00000-of-00001')))
-
- imported = tf.saved_model.load(tmp_dir)
- detection_fn = imported.signatures['serving_default']
-
- images = self._get_dummy_input(
- input_type, batch_size=1, image_size=[128, 128])
-
- processed_images, anchor_boxes, image_info = module._build_inputs(
- tf.zeros((128, 128, 3), dtype=tf.uint8))
- image_info = tf.expand_dims(image_info, 0)
- processed_images = tf.expand_dims(processed_images, 0)
- for l, l_boxes in anchor_boxes.items():
- anchor_boxes[l] = tf.expand_dims(l_boxes, 0)
-
- expected_outputs = module.model(
- images=processed_images,
- image_info=image_info,
- anchor_boxes=anchor_boxes,
- training=False)
- outputs = detection_fn(tf.constant(images))
-
- self.assertAllClose(outputs['num_detections'].numpy(),
- expected_outputs['num_detections'].numpy())
-
- def test_build_model_fail_with_none_batch_size(self):
- params = exp_factory.get_exp_config('panoptic_fpn_coco')
- input_specs = tf.keras.layers.InputSpec(shape=[1, 128, 128, 3])
- model = factory.build_panoptic_maskrcnn(
- input_specs=input_specs, model_config=params.task.model)
- with self.assertRaisesRegex(
- ValueError,
- 'batch_size cannot be None for panoptic segmentation model.'):
- _ = panoptic_segmentation.PanopticSegmentationModule(
- params, model=model, batch_size=None, input_image_size=[128, 128])
-
-if __name__ == '__main__':
- tf.test.main()
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/tasks/__init__.py b/official/vision/beta/projects/panoptic_maskrcnn/tasks/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/tasks/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/panoptic_maskrcnn/tasks/panoptic_maskrcnn_test.py b/official/vision/beta/projects/panoptic_maskrcnn/tasks/panoptic_maskrcnn_test.py
deleted file mode 100644
index daa86552ac6..00000000000
--- a/official/vision/beta/projects/panoptic_maskrcnn/tasks/panoptic_maskrcnn_test.py
+++ /dev/null
@@ -1,69 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Tests for panoptic_maskrcnn.py."""
-import os
-
-from absl.testing import parameterized
-import tensorflow as tf
-
-from official.vision.beta.projects.panoptic_maskrcnn.configs import panoptic_maskrcnn as cfg
-from official.vision.beta.projects.panoptic_maskrcnn.tasks import panoptic_maskrcnn
-from official.vision.configs import decoders as decoder_cfg
-from official.vision.configs import semantic_segmentation as segmentation_cfg
-
-
-class PanopticMaskRCNNTaskTest(tf.test.TestCase, parameterized.TestCase):
-
- @parameterized.parameters(
- (['all'],),
- (['backbone'],),
- (['segmentation_backbone'],),
- (['segmentation_decoder'],),
- (['backbone', 'segmentation_backbone'],),
- (['segmentation_backbone', 'segmentation_decoder'],))
- def test_model_initializing(self, init_checkpoint_modules):
-
- shared_backbone = ('segmentation_backbone' not in init_checkpoint_modules)
- shared_decoder = ('segmentation_decoder' not in init_checkpoint_modules and
- shared_backbone)
-
- task_config = cfg.PanopticMaskRCNNTask(
- model=cfg.PanopticMaskRCNN(
- num_classes=2,
- input_size=[640, 640, 3],
- segmentation_model=segmentation_cfg.SemanticSegmentationModel(
- decoder=decoder_cfg.Decoder(type='fpn')),
- shared_backbone=shared_backbone,
- shared_decoder=shared_decoder))
-
- task = panoptic_maskrcnn.PanopticMaskRCNNTask(task_config)
- model = task.build_model()
-
- ckpt = tf.train.Checkpoint(**model.checkpoint_items)
- ckpt_save_dir = self.create_tempdir().full_path
- ckpt.save(os.path.join(ckpt_save_dir, 'ckpt'))
-
- if (init_checkpoint_modules == ['all'] or
- 'backbone' in init_checkpoint_modules):
- task._task_config.init_checkpoint = ckpt_save_dir
- if ('segmentation_backbone' in init_checkpoint_modules or
- 'segmentation_decoder' in init_checkpoint_modules):
- task._task_config.segmentation_init_checkpoint = ckpt_save_dir
-
- task._task_config.init_checkpoint_modules = init_checkpoint_modules
- task.initialize(model)
-
-if __name__ == '__main__':
- tf.test.main()
diff --git a/official/vision/beta/projects/yolo/README.md b/official/vision/beta/projects/yolo/README.md
deleted file mode 100644
index a39fdf0d288..00000000000
--- a/official/vision/beta/projects/yolo/README.md
+++ /dev/null
@@ -1,88 +0,0 @@
-DISCLAIMER: this YOLO implementation is still under development. No support will
-be provided during the development phase.
-
-# YOLO Object Detectors, You Only Look Once
-
-[](https://arxiv.org/abs/1804.02767)
-[](https://arxiv.org/abs/2004.10934)
-
-This repository is the unofficial implementation of the following papers.
-However, we spent painstaking hours ensuring that every aspect that we
-constructed was the exact same as the original paper and the original
-repository.
-
-* YOLOv3: An Incremental Improvement: [YOLOv3: An Incremental Improvement](https://arxiv.org/abs/1804.02767)
-
-* YOLOv4: Optimal Speed and Accuracy of Object Detection: [YOLOv4: Optimal Speed and Accuracy of Object Detection](https://arxiv.org/abs/2004.10934)
-
-## Description
-
-YOLO v1 the original implementation was released in 2015 providing a
-ground breaking algorithm that would quickly process images and locate objects
-in a single pass through the detector. The original implementation used a
-backbone derived from state of the art object classifiers of the time, like
-[GoogLeNet](https://arxiv.org/abs/1409.4842) and
-[VGG](https://arxiv.org/abs/1409.1556). More attention was given to the novel
-YOLO Detection head that allowed for Object Detection with a single pass of an
-image. Though limited, the network could predict up to 90 bounding boxes per
-image, and was tested for about 80 classes per box. Also, the model can only
-make predictions at one scale. These attributes caused YOLO v1 to be more
-limited and less versatile, so as the year passed, the Developers continued to
-update and develop this model.
-
-YOLO v3 and v4 serve as the most up to date and capable versions of the YOLO
-network group. This model uses a custom backbone called Darknet53 that uses
-knowledge gained from the ResNet paper to improve its predictions. The new
-backbone also allows for objects to be detected at multiple scales. As for the
-new detection head, the model now predicts the bounding boxes using a set of
-anchor box priors (Anchor Boxes) as suggestions. Multiscale predictions in
-combination with Anchor boxes allow for the network to make up to 1000 object
-predictions on a single image. Finally, the new loss function forces the network
-to make better predictions by using Intersection Over Union (IOU) to inform the
-model's confidence rather than relying on the mean squared error for the entire
-output.
-
-
-## Authors
-
-* Vishnu Samardh Banna ([@GitHub vishnubanna](https://github.com/vishnubanna))
-* Anirudh Vegesana ([@GitHub anivegesana](https://github.com/anivegesana))
-* Akhil Chinnakotla ([@GitHub The-Indian-Chinna](https://github.com/The-Indian-Chinna))
-* Tristan Yan ([@GitHub Tyan3001](https://github.com/Tyan3001))
-* Naveen Vivek ([@GitHub naveen-vivek](https://github.com/naveen-vivek))
-
-## Table of Contents
-
-* [Our Goal](#our-goal)
-* [Models in the library](#models-in-the-library)
-* [References](#references)
-
-
-## Our Goal
-
-Our goal with this model conversion is to provide implementation of the Backbone
-and YOLO Head. We have built the model in such a way that the YOLO head could be
-connected to a new, more powerful backbone if a person chose to.
-
-## Models in the library
-
-| Object Detectors | Classifiers |
-| :--------------: | :--------------: |
-| Yolo-v3 | Darknet53 |
-| Yolo-v3 tiny | CSPDarknet53 |
-| Yolo-v3 spp |
-| Yolo-v4 |
-| Yolo-v4 tiny |
-| Yolo-v4 csp |
-| Yolo-v4 large |
-
-## Models Zoo
-
-
-## Requirements
-[](https://github.com/tensorflow/tensorflow/releases/tag/v2.6.0)
-[](https://www.python.org/downloads/release/python-380/)
-
-
-DISCLAIMER: this YOLO implementation is still under development. No support
-will be provided during the development phase.
diff --git a/official/vision/beta/projects/yolo/__init__.py b/official/vision/beta/projects/yolo/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/yolo/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/yolo/common/__init__.py b/official/vision/beta/projects/yolo/common/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/yolo/common/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/yolo/common/registry_imports.py b/official/vision/beta/projects/yolo/common/registry_imports.py
deleted file mode 100644
index 28218cdfa81..00000000000
--- a/official/vision/beta/projects/yolo/common/registry_imports.py
+++ /dev/null
@@ -1,36 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""All necessary imports for registration."""
-
-# pylint: disable=unused-import
-# pylint: disable=g-bad-import-order
-from official.vision import registry_imports
-
-# import configs
-from official.vision.beta.projects.yolo.configs import darknet_classification
-from official.vision.beta.projects.yolo.configs import yolo as yolo_config
-
-# import modeling components
-from official.vision.beta.projects.yolo.modeling.backbones import darknet
-from official.vision.beta.projects.yolo.modeling.decoders import yolo_decoder
-
-# import tasks
-from official.vision.beta.projects.yolo.tasks import image_classification
-from official.vision.beta.projects.yolo.tasks import yolo as yolo_task
-
-# import optimization packages
-from official.vision.beta.projects.yolo.optimization import optimizer_factory
-from official.vision.beta.projects.yolo.optimization.configs import optimizer_config
-from official.vision.beta.projects.yolo.optimization.configs import optimization_config
diff --git a/official/vision/beta/projects/yolo/configs/__init__.py b/official/vision/beta/projects/yolo/configs/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/yolo/configs/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/yolo/losses/__init__.py b/official/vision/beta/projects/yolo/losses/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/yolo/losses/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/yolo/modeling/decoders/__init__.py b/official/vision/beta/projects/yolo/modeling/decoders/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/yolo/modeling/decoders/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/yolo/modeling/heads/__init__.py b/official/vision/beta/projects/yolo/modeling/heads/__init__.py
deleted file mode 100644
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/yolo/modeling/heads/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/yolo/optimization/configs/__init__.py b/official/vision/beta/projects/yolo/optimization/configs/__init__.py
deleted file mode 100755
index 310bfb28f0c..00000000000
--- a/official/vision/beta/projects/yolo/optimization/configs/__init__.py
+++ /dev/null
@@ -1,14 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
diff --git a/official/vision/beta/projects/yolo/serving/export_module_factory.py b/official/vision/beta/projects/yolo/serving/export_module_factory.py
deleted file mode 100644
index 5c8fc5ae3f4..00000000000
--- a/official/vision/beta/projects/yolo/serving/export_module_factory.py
+++ /dev/null
@@ -1,175 +0,0 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
-#
-# Licensed under the Apache License, Version 2.0 (the "License");
-# you may not use this file except in compliance with the License.
-# You may obtain a copy of the License at
-#
-# http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing, software
-# distributed under the License is distributed on an "AS IS" BASIS,
-# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-# See the License for the specific language governing permissions and
-# limitations under the License.
-
-"""Factory for YOLO export modules."""
-
-from typing import List, Optional
-
-import tensorflow as tf
-
-from official.core import config_definitions as cfg
-from official.vision import configs
-from official.vision.beta.projects.yolo.configs.yolo import YoloTask
-from official.vision.beta.projects.yolo.modeling import factory as yolo_factory
-from official.vision.beta.projects.yolo.modeling.backbones import darknet # pylint: disable=unused-import
-from official.vision.beta.projects.yolo.modeling.decoders import yolo_decoder # pylint: disable=unused-import
-from official.vision.beta.projects.yolo.serving import model_fn as yolo_model_fn
-from official.vision.dataloaders import classification_input
-from official.vision.modeling import factory
-from official.vision.serving import export_base_v2 as export_base
-from official.vision.serving import export_utils
-
-
-def create_classification_export_module(
- params: cfg.ExperimentConfig,
- input_type: str,
- batch_size: int,
- input_image_size: List[int],
- num_channels: int = 3) -> export_base.ExportModule:
- """Creates classification export module."""
- input_signature = export_utils.get_image_input_signatures(
- input_type, batch_size, input_image_size, num_channels)
- input_specs = tf.keras.layers.InputSpec(shape=[batch_size] +
- input_image_size + [num_channels])
-
- model = factory.build_classification_model(
- input_specs=input_specs,
- model_config=params.task.model,
- l2_regularizer=None)
-
- def preprocess_fn(inputs):
- image_tensor = export_utils.parse_image(inputs, input_type,
- input_image_size, num_channels)
- # If input_type is `tflite`, do not apply image preprocessing.
- if input_type == 'tflite':
- return image_tensor
-
- def preprocess_image_fn(inputs):
- return classification_input.Parser.inference_fn(inputs, input_image_size,
- num_channels)
-
- images = tf.map_fn(
- preprocess_image_fn,
- elems=image_tensor,
- fn_output_signature=tf.TensorSpec(
- shape=input_image_size + [num_channels], dtype=tf.float32))
-
- return images
-
- def postprocess_fn(logits):
- probs = tf.nn.softmax(logits)
- return {'logits': logits, 'probs': probs}
-
- export_module = export_base.ExportModule(
- params,
- model=model,
- input_signature=input_signature,
- preprocessor=preprocess_fn,
- postprocessor=postprocess_fn)
- return export_module
-
-
-def create_yolo_export_module(
- params: cfg.ExperimentConfig,
- input_type: str,
- batch_size: int,
- input_image_size: List[int],
- num_channels: int = 3) -> export_base.ExportModule:
- """Creates YOLO export module."""
- input_signature = export_utils.get_image_input_signatures(
- input_type, batch_size, input_image_size, num_channels)
- input_specs = tf.keras.layers.InputSpec(shape=[batch_size] +
- input_image_size + [num_channels])
- model, _ = yolo_factory.build_yolo(
- input_specs=input_specs,
- model_config=params.task.model,
- l2_regularization=None)
-
- def preprocess_fn(inputs):
- image_tensor = export_utils.parse_image(inputs, input_type,
- input_image_size, num_channels)
- # If input_type is `tflite`, do not apply image preprocessing.
- if input_type == 'tflite':
- return image_tensor
-
- def preprocess_image_fn(inputs):
- image = tf.cast(inputs, dtype=tf.float32)
- image = image / 255.
- (image, image_info) = yolo_model_fn.letterbox(
- image,
- input_image_size,
- letter_box=params.task.validation_data.parser.letter_box)
- return image, image_info
-
- images_spec = tf.TensorSpec(shape=input_image_size + [3], dtype=tf.float32)
-
- image_info_spec = tf.TensorSpec(shape=[4, 2], dtype=tf.float32)
-
- images, image_info = tf.nest.map_structure(
- tf.identity,
- tf.map_fn(
- preprocess_image_fn,
- elems=image_tensor,
- fn_output_signature=(images_spec, image_info_spec),
- parallel_iterations=32))
-
- return images, image_info
-
- def inference_steps(inputs, model):
- images, image_info = inputs
- detection = model(images, training=False)
- detection['bbox'] = yolo_model_fn.undo_info(
- detection['bbox'],
- detection['num_detections'],
- image_info,
- expand=False)
-
- final_outputs = {
- 'detection_boxes': detection['bbox'],
- 'detection_scores': detection['confidence'],
- 'detection_classes': detection['classes'],
- 'num_detections': detection['num_detections']
- }
-
- return final_outputs
-
- export_module = export_base.ExportModule(
- params,
- model=model,
- input_signature=input_signature,
- preprocessor=preprocess_fn,
- inference_step=inference_steps)
-
- return export_module
-
-
-def get_export_module(params: cfg.ExperimentConfig,
- input_type: str,
- batch_size: Optional[int],
- input_image_size: List[int],
- num_channels: int = 3) -> export_base.ExportModule:
- """Factory for export modules."""
- if isinstance(params.task,
- configs.image_classification.ImageClassificationTask):
- export_module = create_classification_export_module(params, input_type,
- batch_size,
- input_image_size,
- num_channels)
- elif isinstance(params.task, YoloTask):
- export_module = create_yolo_export_module(params, input_type, batch_size,
- input_image_size, num_channels)
- else:
- raise ValueError('Export module not implemented for {} task.'.format(
- type(params.task)))
- return export_module
diff --git a/official/vision/configs/__init__.py b/official/vision/configs/__init__.py
index 7ba5793215c..b152d58e877 100644
--- a/official/vision/configs/__init__.py
+++ b/official/vision/configs/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/configs/backbones.py b/official/vision/configs/backbones.py
index 337844c4b1d..46caf5ff547 100644
--- a/official/vision/configs/backbones.py
+++ b/official/vision/configs/backbones.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,48 @@
"""Backbones configurations."""
import dataclasses
-from typing import Optional, List
-
-# Import libraries
+from typing import List, Optional, Tuple
from official.modeling import hyperparams
+@dataclasses.dataclass
+class Transformer(hyperparams.Config):
+ """Transformer config."""
+ mlp_dim: int = 1
+ num_heads: int = 1
+ num_layers: int = 1
+ attention_dropout_rate: float = 0.0
+ dropout_rate: float = 0.0
+
+
+@dataclasses.dataclass
+class VisionTransformer(hyperparams.Config):
+ """VisionTransformer config."""
+ model_name: str = 'vit-b16'
+ # pylint: disable=line-too-long
+ pooler: str = 'token' # 'token', 'gap' or 'none'. If set to 'token', an extra classification token is added to sequence.
+ # pylint: enable=line-too-long
+ representation_size: int = 0
+ hidden_size: int = 1
+ patch_size: int = 16
+ transformer: Transformer = dataclasses.field(default_factory=Transformer)
+ init_stochastic_depth_rate: float = 0.0
+ original_init: bool = True
+ pos_embed_shape: Optional[Tuple[int, int]] = None
+ # If output encoded tokens sequence when pooler is `none`.
+ output_encoded_tokens: bool = True
+ # If output encoded tokens 2D feature map.
+ output_2d_feature_maps: bool = False
+
+ # Adding Layerscale to each Encoder block https://arxiv.org/abs/2204.07118
+ layer_scale_init_value: float = 0.0
+ # Transformer encoder spatial partition dimensions.
+ transformer_partition_dims: Optional[Tuple[int, int, int, int]] = None
+ # If True, output attention scores.
+ output_attention_scores: bool = False
+
+
@dataclasses.dataclass
class ResNet(hyperparams.Config):
"""ResNet config."""
@@ -45,6 +80,8 @@ class DilatedResNet(hyperparams.Config):
last_stage_repeats: int = 1
se_ratio: float = 0.0
stochastic_depth_drop_rate: float = 0.0
+ resnetd_shortcut: bool = False
+ replace_stem_max_pool: bool = False
@dataclasses.dataclass
@@ -61,6 +98,10 @@ class MobileNet(hyperparams.Config):
model_id: str = 'MobileNetV2'
filter_size_scale: float = 1.0
stochastic_depth_drop_rate: float = 0.0
+ # Whether to apply a fixed and common stochastic depth drop rate to all
+ # blocks, instead to linearly scale it from zero to maximum value (standard
+ # behaviour for stochastic depth). Set to True for backward compatibility.
+ flat_stochastic_depth_drop_rate: bool = True
output_stride: Optional[int] = None
output_intermediate_endpoints: bool = False
@@ -118,14 +159,19 @@ class Backbone(hyperparams.OneOfConfig):
spinenet_mobile: mobile spinenet backbone config.
mobilenet: mobilenet backbone config.
mobiledet: mobiledet backbone config.
+ vit: vision transformer backbone config.
"""
type: Optional[str] = None
- resnet: ResNet = ResNet()
- dilated_resnet: DilatedResNet = DilatedResNet()
- revnet: RevNet = RevNet()
- efficientnet: EfficientNet = EfficientNet()
- spinenet: SpineNet = SpineNet()
- spinenet_mobile: SpineNetMobile = SpineNetMobile()
- mobilenet: MobileNet = MobileNet()
- mobiledet: MobileDet = MobileDet()
-
+ resnet: ResNet = dataclasses.field(default_factory=ResNet)
+ dilated_resnet: DilatedResNet = dataclasses.field(
+ default_factory=DilatedResNet
+ )
+ revnet: RevNet = dataclasses.field(default_factory=RevNet)
+ efficientnet: EfficientNet = dataclasses.field(default_factory=EfficientNet)
+ spinenet: SpineNet = dataclasses.field(default_factory=SpineNet)
+ spinenet_mobile: SpineNetMobile = dataclasses.field(
+ default_factory=SpineNetMobile
+ )
+ mobilenet: MobileNet = dataclasses.field(default_factory=MobileNet)
+ mobiledet: MobileDet = dataclasses.field(default_factory=MobileDet)
+ vit: VisionTransformer = dataclasses.field(default_factory=VisionTransformer)
diff --git a/official/vision/configs/backbones_3d.py b/official/vision/configs/backbones_3d.py
index 436a3b1be66..38dfa791e4b 100644
--- a/official/vision/configs/backbones_3d.py
+++ b/official/vision/configs/backbones_3d.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,6 @@
"""3D Backbones configurations."""
from typing import Optional, Tuple
-# Import libraries
import dataclasses
from official.modeling import hyperparams
@@ -88,10 +87,6 @@ class ResNet3DRS(ResNet3D):
use_self_gating=True))
-_RESNET3D50_DEFAULT_CFG = ResNet3D50()
-_RESNET3DRS_DEFAULT_CFG = ResNet3DRS()
-
-
@dataclasses.dataclass
class Backbone3D(hyperparams.OneOfConfig):
"""Configuration for backbones.
@@ -102,5 +97,5 @@ class Backbone3D(hyperparams.OneOfConfig):
resnet_3d_rs: resnet3d-rs backbone config.
"""
type: Optional[str] = None
- resnet_3d: ResNet3D = _RESNET3D50_DEFAULT_CFG
- resnet_3d_rs: ResNet3D = _RESNET3DRS_DEFAULT_CFG
+ resnet_3d: ResNet3D = dataclasses.field(default_factory=ResNet3D50)
+ resnet_3d_rs: ResNet3D = dataclasses.field(default_factory=ResNet3DRS)
diff --git a/official/vision/configs/common.py b/official/vision/configs/common.py
index 6731715bc50..a1091b8cad3 100644
--- a/official/vision/configs/common.py
+++ b/official/vision/configs/common.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,9 +15,7 @@
"""Common configurations."""
import dataclasses
-from typing import List, Optional
-
-# Import libraries
+from typing import List, Optional, Sequence
from official.core import config_definitions as cfg
from official.modeling import hyperparams
@@ -28,6 +26,7 @@ class TfExampleDecoder(hyperparams.Config):
"""A simple TF Example decoder config."""
regenerate_source_id: bool = False
mask_binarize_threshold: Optional[float] = None
+ attribute_names: List[str] = dataclasses.field(default_factory=list)
@dataclasses.dataclass
@@ -48,8 +47,12 @@ class DataDecoder(hyperparams.OneOfConfig):
label_map_decoder: TF Example decoder with label map config.
"""
type: Optional[str] = 'simple_decoder'
- simple_decoder: TfExampleDecoder = TfExampleDecoder()
- label_map_decoder: TfExampleDecoderLabelMap = TfExampleDecoderLabelMap()
+ simple_decoder: TfExampleDecoder = dataclasses.field(
+ default_factory=TfExampleDecoder
+ )
+ label_map_decoder: TfExampleDecoderLabelMap = dataclasses.field(
+ default_factory=TfExampleDecoderLabelMap
+ )
@dataclasses.dataclass
@@ -95,6 +98,35 @@ class MixupAndCutmix(hyperparams.Config):
label_smoothing: float = 0.1
+@dataclasses.dataclass
+class SSDRandomCropParam(hyperparams.Config):
+ min_object_covered: float = 0.0
+ min_box_overlap: float = 0.5
+ prob_to_apply: float = 0.85
+
+
+@dataclasses.dataclass
+class SSDRandomCrop(hyperparams.Config):
+ """Configuration for SSDRandomCrop.
+
+ Liu et al., SSD: Single shot multibox detector
+ https://arxiv.org/abs/1512.02325.
+ """
+ ssd_random_crop_params: Sequence[SSDRandomCropParam] = dataclasses.field(
+ default_factory=lambda: (
+ SSDRandomCropParam(min_object_covered=0.0),
+ SSDRandomCropParam(min_object_covered=0.1),
+ SSDRandomCropParam(min_object_covered=0.3),
+ SSDRandomCropParam(min_object_covered=0.5),
+ SSDRandomCropParam(min_object_covered=0.7),
+ SSDRandomCropParam(min_object_covered=0.9),
+ SSDRandomCropParam(min_object_covered=1.0),
+ )
+ )
+ aspect_ratio_range: tuple[float, float] = (0.5, 2.0)
+ area_range: tuple[float, float] = (0.1, 1.0)
+
+
@dataclasses.dataclass
class Augmentation(hyperparams.OneOfConfig):
"""Configuration for input data augmentation.
@@ -105,8 +137,18 @@ class Augmentation(hyperparams.OneOfConfig):
autoaug: AutoAugment config.
"""
type: Optional[str] = None
- randaug: RandAugment = RandAugment()
- autoaug: AutoAugment = AutoAugment()
+ randaug: RandAugment = dataclasses.field(default_factory=RandAugment)
+ autoaug: AutoAugment = dataclasses.field(default_factory=AutoAugment)
+ ssd_random_crop: SSDRandomCrop = dataclasses.field(
+ default_factory=SSDRandomCrop
+ )
+
+
+@dataclasses.dataclass
+class RandJpegQuality(hyperparams.Config):
+ min_quality: int = 20
+ max_quality: int = 100
+ prob_to_apply: float = 0.6
@dataclasses.dataclass
@@ -138,6 +180,7 @@ class PseudoLabelDataConfig(cfg.DataConfig):
@dataclasses.dataclass
class TFLitePostProcessingConfig(hyperparams.Config):
+ """TFLite Post Processing config for inference."""
max_detections: int = 200
max_classes_per_detection: int = 5
# Regular NMS run in a multi-class fashion and is slow. Setting it to False
@@ -145,3 +188,17 @@ class TFLitePostProcessingConfig(hyperparams.Config):
use_regular_nms: bool = False
nms_score_threshold: float = 0.1
nms_iou_threshold: float = 0.5
+ # Whether to normalize coordinates of anchors to [0, 1]. If setting to True,
+ # coordinates of output boxes is also normalized but latency increases.
+ normalize_anchor_coordinates: Optional[bool] = False
+ # Whether to omit the final nms placeholder op. If set to True, the output
+ # will be a tuple of boxes, scores result right before the NMS operation.
+ omit_nms: Optional[bool] = False
+ # The number of detections per class when using regular NMS.
+ detections_per_class: Optional[int] = 5
+ # Box scaling factors. It should agree with `box_coder_weights` defined in
+ # `DetectionGenerator`, which is in the format of [y, x, w, h].
+ y_scale: float = 1.0
+ x_scale: float = 1.0
+ w_scale: float = 1.0
+ h_scale: float = 1.0
diff --git a/official/vision/configs/decoders.py b/official/vision/configs/decoders.py
index 4081d429ecd..c156f2b24aa 100644
--- a/official/vision/configs/decoders.py
+++ b/official/vision/configs/decoders.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,8 +16,6 @@
import dataclasses
from typing import List, Optional
-# Import libraries
-
from official.modeling import hyperparams
@@ -33,6 +31,7 @@ class FPN(hyperparams.Config):
num_filters: int = 256
fusion_type: str = 'sum'
use_separable_conv: bool = False
+ use_keras_layer: bool = False
@dataclasses.dataclass
@@ -65,7 +64,7 @@ class Decoder(hyperparams.OneOfConfig):
fpn: fpn config.
"""
type: Optional[str] = None
- fpn: FPN = FPN()
- nasfpn: NASFPN = NASFPN()
- identity: Identity = Identity()
- aspp: ASPP = ASPP()
+ fpn: FPN = dataclasses.field(default_factory=FPN)
+ nasfpn: NASFPN = dataclasses.field(default_factory=NASFPN)
+ identity: Identity = dataclasses.field(default_factory=Identity)
+ aspp: ASPP = dataclasses.field(default_factory=ASPP)
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv1_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv1_tpu.yaml
new file mode 100644
index 00000000000..258c46ae71f
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv1_tpu.yaml
@@ -0,0 +1,48 @@
+# MobileNetV1_1.0 ImageNet classification. 72.33% top-1 and 90.65% top-5 accuracy.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [224, 224, 3]
+ backbone:
+ type: 'mobilenet'
+ mobilenet:
+ model_id: 'MobileNetV1'
+ filter_size_scale: 1.0
+ dropout_rate: 0.2
+ losses:
+ l2_weight_decay: 0.00001
+ one_hot: true
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 156000 # 500 epochs
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312 # NUM_EXAMPLES (1281167) // global_batch_size
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ type: 'exponential'
+ exponential:
+ initial_learning_rate: 0.256 # 0.008 * batch_size / 128
+ decay_steps: 780 # 2.5 * steps_per_epoch
+ decay_rate: 0.94
+ staircase: true
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_steps: 1560
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_gpu.yaml
index ff1a0719e6f..35ea33df68a 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_gpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_gpu.yaml
@@ -18,12 +18,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 1024 # 128 * 8
dtype: 'float16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 1024 # 128 * 8
dtype: 'float16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_tpu.yaml
index b5df9d6e74a..a5ec4420b07 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv2_tpu.yaml
@@ -17,12 +17,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_avg_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_avg_tpu.yaml
index 666bc5b7e90..d35ea55f157 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_avg_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_avg_tpu.yaml
@@ -13,22 +13,22 @@ task:
stochastic_depth_drop_rate: 0.0
type: mobilenet
dropout_rate: 0.2
- losses:
- l2_weight_decay: 0.00001
- one_hot: true
- label_smoothing: 0.1
norm_activation:
activation: relu
norm_epsilon: 0.001
norm_momentum: 0.997
use_sync_bn: false
+ losses:
+ l2_weight_decay: 0.00001
+ one_hot: true
+ label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_avgseg_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_avgseg_tpu.yaml
new file mode 100644
index 00000000000..15fd099a177
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3.5_avgseg_tpu.yaml
@@ -0,0 +1,60 @@
+# MobileNetV3.5 AVG Seg ImageNet classification. 73.35% top-1 and 91.49% top-5 accuracy.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [224, 224, 3]
+ backbone:
+ mobilenet:
+ filter_size_scale: 1.0
+ model_id: MobileNetMultiAVGSeg
+ stochastic_depth_drop_rate: 0.0
+ type: mobilenet
+ dropout_rate: 0.2
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
+ use_sync_bn: false
+ losses:
+ l2_weight_decay: 0.00001
+ one_hot: true
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 156000 # 500 epochs
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312 # NUM_EXAMPLES (1281167) // global_batch_size
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ type: 'exponential'
+ exponential:
+ initial_learning_rate: 0.256 # 0.008 * batch_size / 128
+ decay_steps: 780 # 2.5 * steps_per_epoch
+ decay_rate: 0.96
+ staircase: true
+ optimizer:
+ type: rmsprop
+ rmsprop:
+ epsilon: 0.002
+ momentum: 0.9
+ rho: 0.9
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_steps: 1560
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3large_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3large_tpu.yaml
index a8fb0c95d55..5c088720fe5 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3large_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3large_tpu.yaml
@@ -1,4 +1,4 @@
-# MobileNetV3-large_1.0 ImageNet classification: 74.96% top-1.
+# MobileNetV3-large_1.0 ImageNet classification: ~75.7% top-1.
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'bfloat16'
@@ -11,28 +11,34 @@ task:
mobilenet:
model_id: 'MobileNetV3Large'
filter_size_scale: 1.0
+ kernel_initializer: 'random_uniform'
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
dropout_rate: 0.2
losses:
- l2_weight_decay: 0.00001
+ l2_weight_decay: 0.0
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
- # Enables Inception-style pre-processing.
- decode_jpeg_only: false
+ aug_type:
+ autoaug:
+ augmentation_name: 'v0'
+ cutout_const: 100
+ translate_const: 250
+ type: 'autoaug'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
drop_remainder: false
- # Enables Inception-style pre-processing.
- decode_jpeg_only: false
trainer:
- train_steps: 156000 # 500 epochs
+ train_steps: 218000 # 700 epochs
validation_steps: 13
validation_interval: 312
steps_per_loop: 312 # NUM_EXAMPLES (1281167) // global_batch_size
@@ -40,14 +46,25 @@ trainer:
checkpoint_interval: 312
optimizer_config:
learning_rate:
- type: 'cosine'
cosine:
alpha: 0.0
- decay_steps: 156000
- initial_learning_rate: 0.5
+ decay_steps: 218000
+ initial_learning_rate: 0.004
name: CosineDecay
offset: 0
- warmup:
- type: 'linear'
- linear:
- warmup_steps: 5000
+ type: 'cosine'
+ optimizer:
+ adamw:
+ amsgrad: false
+ beta_1: 0.9
+ beta_2: 0.999
+ clipnorm: null
+ clipvalue: null
+ epsilon: 1.0e-07
+ exclude_from_weight_decay: ['batch_normalization']
+ global_clipnorm: null
+ gradient_clip_norm: 0.0
+ include_in_weight_decay: null
+ name: 'AdamWeightDecay'
+ weight_decay_rate: 0.1
+ type: 'adamw'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3small_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3small_tpu.yaml
index 57466795214..a9959c33689 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3small_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv3small_tpu.yaml
@@ -1,4 +1,4 @@
-# MobileNetV3Small ImageNet classification. 67.5% top-1 and 87.6% top-5 accuracy.
+# MobileNetV3Small ImageNet classification. 67.5% top-1 and 87.7% top-5 accuracy.
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'bfloat16'
@@ -22,19 +22,19 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
drop_remainder: false
trainer:
train_steps: 312000 # 1000 epochs
- validation_steps: 12
+ validation_steps: 13
validation_interval: 312
steps_per_loop: 312 # NUM_EXAMPLES (1281167) // global_batch_size
summary_interval: 312
@@ -49,7 +49,7 @@ trainer:
learning_rate:
type: 'exponential'
exponential:
- initial_learning_rate: 0.01
+ initial_learning_rate: 0.426 # 0.02 * (batch_size / 192)
decay_steps: 936 # 3 * steps_per_epoch
decay_rate: 0.99
staircase: true
@@ -60,4 +60,4 @@ trainer:
type: 'linear'
linear:
warmup_steps: 1560
- warmup_learning_rate: 0.001
+ warmup_learning_rate: 0.0
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_large_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_large_tpu.yaml
new file mode 100644
index 00000000000..e3944d1cf4c
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_large_tpu.yaml
@@ -0,0 +1,73 @@
+# Top-1 Acc: 82.92% @ 500 epoches
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [384, 384, 3]
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetV4ConvLarge'
+ flat_stochastic_depth_drop_rate: true
+ stochastic_depth_drop_rate: 0.2
+ type: 'mobilenet'
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
+ dropout_rate: 0.2
+ losses:
+ l2_weight_decay: 0.0
+ one_hot: false
+ soft_labels: true
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 16384
+ dtype: 'bfloat16'
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: ['Cutout']
+ magnitude: 15
+ type: 'randaug'
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 0.2
+ switch_prob: 0.5
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 39000 # 500 epochs
+ validation_steps: 13
+ validation_interval: 78
+ steps_per_loop: 78 # NUM_EXAMPLES (1281167) // global_batch_size
+ summary_interval: 78
+ checkpoint_interval: 78
+ max_to_keep: 1
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 39000
+ initial_learning_rate: 0.004
+ name: 'CosineDecay'
+ type: 'cosine'
+ optimizer:
+ adamw:
+ exclude_from_weight_decay: ['batch_normalization']
+ gradient_clip_norm: 0.0
+ weight_decay_rate: 0.2
+ name: 'AdamWeightDecay'
+ type: 'adamw'
+ warmup:
+ linear:
+ warmup_steps: 1560
+ name: 'linear'
+ type: 'linear'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_medium_seg_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_medium_seg_tpu.yaml
new file mode 100644
index 00000000000..d5ca5d491d2
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_medium_seg_tpu.yaml
@@ -0,0 +1,65 @@
+# Top-1 Acc: 77.15%/ 66.57% Val/Train @ 500 epochs
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [256, 256, 3]
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetV4ConvMediumSeg'
+ flat_stochastic_depth_drop_rate: false
+ stochastic_depth_drop_rate: 0.075
+ type: 'mobilenet'
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
+ dropout_rate: 0.2
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ aug_type:
+ randaug:
+ cutout_const: 20
+ exclude_ops: ['Cutout']
+ magnitude: 15
+ prob_to_apply: 0.7
+ type: 'randaug'
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 156000
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 156000
+ initial_learning_rate: 0.004
+ name: 'CosineDecay'
+ type: 'cosine'
+ optimizer:
+ adamw:
+ exclude_from_weight_decay: ['batch_normalization']
+ gradient_clip_norm: 0.0
+ weight_decay_rate: 0.1
+ name: 'AdamWeightDecay'
+ type: 'adamw'
+ warmup:
+ linear:
+ warmup_steps: 1560
+ name: 'linear'
+ type: 'linear'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_medium_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_medium_tpu.yaml
new file mode 100644
index 00000000000..f2b27473a26
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_medium_tpu.yaml
@@ -0,0 +1,65 @@
+# Top-1 Acc: 79.9% @ 500 epochs
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [256, 256, 3]
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetV4ConvMedium'
+ flat_stochastic_depth_drop_rate: false
+ stochastic_depth_drop_rate: 0.075
+ type: 'mobilenet'
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
+ dropout_rate: 0.2
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ aug_type:
+ randaug:
+ cutout_const: 20
+ exclude_ops: ['Cutout']
+ magnitude: 15
+ prob_to_apply: 0.7
+ type: 'randaug'
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 156000
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 156000
+ initial_learning_rate: 0.004
+ name: 'CosineDecay'
+ type: 'cosine'
+ optimizer:
+ adamw:
+ exclude_from_weight_decay: ['batch_normalization']
+ gradient_clip_norm: 0.0
+ weight_decay_rate: 0.1
+ name: 'AdamWeightDecay'
+ type: 'adamw'
+ warmup:
+ linear:
+ warmup_steps: 1560
+ name: 'linear'
+ type: 'linear'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_small_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_small_tpu.yaml
new file mode 100644
index 00000000000..deecaf80900
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_conv_small_tpu.yaml
@@ -0,0 +1,74 @@
+# Top-1 Acc: 73.8% @ ~9623 epochs
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'float32'
+task:
+ model:
+ num_classes: 1001
+ input_size: [224, 224, 3]
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetV4ConvSmall'
+ type: 'mobilenet'
+ norm_activation:
+ norm_epsilon: 1.0e-06
+ norm_momentum: 0.95
+ dropout_rate: 0.3
+ losses:
+ l2_weight_decay: 1.0e-05
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'float32'
+ aug_type:
+ randaug:
+ cutout_const: 20
+ exclude_ops: ['Cutout']
+ magnitude: 9
+ prob_to_apply: 0.5
+ translate_const: 20
+ type: 'randaug'
+ prefetch_buffer_size: 50
+ enable_tf_data_service: true
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'float32'
+ prefetch_buffer_size: 50
+ drop_remainder: false
+trainer:
+ train_steps: 3002368
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ ema:
+ name: 'ExponentialMovingAverage'
+ average_decay: 0.9999
+ start_step: 1560
+ dynamic_decay: false
+ learning_rate:
+ cosine:
+ decay_steps: 3002368
+ initial_learning_rate: 0.002
+ name: 'CosineDecay'
+ type: 'cosine'
+ optimizer:
+ adamw:
+ beta_1: 0.6
+ epsilon: 1.0e-06
+ exclude_from_weight_decay: ['batch_normalization']
+ gradient_clip_norm: 0.0
+ weight_decay_rate: 0.01
+ name: 'AdamWeightDecay'
+ type: 'adamw'
+ warmup:
+ linear:
+ warmup_steps: 1560
+ name: 'linear'
+ type: 'linear'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_hybrid_large_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_hybrid_large_tpu.yaml
new file mode 100644
index 00000000000..20ef374c875
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_hybrid_large_tpu.yaml
@@ -0,0 +1,71 @@
+# Top-1 Acc: 83.7% @ 500 epochs
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [384, 384, 3]
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetV4HybridLarge'
+ flat_stochastic_depth_drop_rate: false
+ stochastic_depth_drop_rate: 0.35
+ type: 'mobilenet'
+ norm_activation:
+ norm_epsilon: 0.001
+ norm_momentum: 0.997
+ dropout_rate: 0.2
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 16384
+ dtype: 'bfloat16'
+ aug_type:
+ randaug:
+ exclude_ops: ['Cutout']
+ magnitude: 15
+ type: 'randaug'
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ mixup_alpha: 0.8
+ prob: 0.3
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 39000
+ validation_steps: 13
+ validation_interval: 78
+ steps_per_loop: 78
+ summary_interval: 78
+ checkpoint_interval: 78
+ max_to_keep: 1
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 39000
+ initial_learning_rate: 0.01
+ alpha: 0.001
+ name: 'CosineDecay'
+ type: 'cosine'
+ optimizer:
+ adamw:
+ exclude_from_weight_decay: ['batch_normalization']
+ gradient_clip_norm: 0.0
+ weight_decay_rate: 0.2
+ name: 'AdamWeightDecay'
+ type: 'adamw'
+ warmup:
+ linear:
+ warmup_steps: 1560
+ name: 'linear'
+ type: 'linear'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_hybrid_medium_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_hybrid_medium_tpu.yaml
new file mode 100644
index 00000000000..37558aff28f
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_mobilenetv4_hybrid_medium_tpu.yaml
@@ -0,0 +1,67 @@
+# Top-1 Acc: 80.7% @ 500 epochs
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [256, 256, 3]
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetV4HybridMedium'
+ flat_stochastic_depth_drop_rate: false
+ stochastic_depth_drop_rate: 0.075
+ type: 'mobilenet'
+ norm_activation:
+ norm_momentum: 0.997
+ dropout_rate: 0.2
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 16384
+ dtype: 'bfloat16'
+ aug_type:
+ randaug:
+ cutout_const: 20
+ exclude_ops: ['Cutout']
+ magnitude: 15
+ magnitude_std: 0.0
+ num_layers: 2
+ prob_to_apply: 0.7
+ translate_const: 10
+ type: 'randaug'
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 16384
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 39000
+ validation_steps: 4
+ validation_interval: 156
+ steps_per_loop: 156
+ summary_interval: 156
+ checkpoint_interval: 156
+ optimizer_config:
+ learning_rate:
+ cosine:
+ decay_steps: 39000
+ initial_learning_rate: 0.016
+ name: 'CosineDecay'
+ type: 'cosine'
+ optimizer:
+ adamw:
+ exclude_from_weight_decay: ['batch_normalization']
+ gradient_clip_norm: 0.0
+ weight_decay_rate: 0.1
+ name: 'AdamWeightDecay'
+ type: 'adamw'
+ warmup:
+ linear:
+ name: 'linear'
+ warmup_steps: 1560
+ type: 'linear'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnet101_deeplab_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnet101_deeplab_tpu.yaml
index 5d7d2959637..cb86a8e7f11 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnet101_deeplab_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnet101_deeplab_tpu.yaml
@@ -23,13 +23,13 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
aug_policy: 'randaug'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml
index 2600f58faa5..c9fea870939 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnet101_tpu.yaml
@@ -17,12 +17,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml
index 1c81953e2f6..5093ce1100b 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnet152_tpu.yaml
@@ -17,12 +17,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnet26_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnet26_tpu.yaml
new file mode 100644
index 00000000000..df8189ccafd
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnet26_tpu.yaml
@@ -0,0 +1,47 @@
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ model:
+ num_classes: 1001
+ input_size: [224, 224, 3]
+ backbone:
+ type: 'resnet'
+ resnet:
+ model_id: 26
+ losses:
+ l2_weight_decay: 0.0001
+ one_hot: true
+ label_smoothing: 0.1
+ train_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
+ is_training: true
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ validation_data:
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
+ is_training: false
+ global_batch_size: 4096
+ dtype: 'bfloat16'
+ drop_remainder: false
+trainer:
+ train_steps: 28080
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ optimizer:
+ type: 'sgd'
+ sgd:
+ momentum: 0.9
+ learning_rate:
+ type: 'stepwise'
+ stepwise:
+ boundaries: [9360, 18720, 24960]
+ values: [1.6, 0.16, 0.016, 0.0016]
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_steps: 1560
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnet50_deeplab_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnet50_deeplab_tpu.yaml
index 11bdafbc35d..5054d1a7eaf 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnet50_deeplab_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnet50_deeplab_tpu.yaml
@@ -17,12 +17,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnet50_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnet50_gpu.yaml
index f1202242a80..f31fde1ddd2 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnet50_gpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnet50_gpu.yaml
@@ -15,7 +15,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 2048
dtype: 'float16'
@@ -23,7 +23,7 @@ task:
# static value: 32. See b/218880025.
prefetch_buffer_size: 32
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 2048
dtype: 'float16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml
index 358cafb6df9..8421b3411d5 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnet50_tpu.yaml
@@ -14,12 +14,12 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml
index 7c9e7b80a02..915f94fc3af 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i160.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml
index 576c4862505..10af95d6f3f 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs101_i192.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml
index b1c8edc463f..7427de95261 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i192.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml
index 2ec14bae5ab..a31c09bd872 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i224.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml
index 91b53d6217f..12990520e50 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs152_i256.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml
index 9d76c010170..d8f7e841e14 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs200_i256.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml
index b7c6a644e2c..25735b23cfa 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs270_i256.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml
index 3b2d3fe261c..fd8f9f57812 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i256.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i320.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i320.yaml
index 36cdba7bb43..a83420d663a 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i320.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs350_i320.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs420_i320.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs420_i320.yaml
index 9b02b7e006a..fea2190c864 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs420_i320.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs420_i320.yaml
@@ -24,7 +24,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -33,7 +33,7 @@ task:
randaug:
magnitude: 15
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml
index a57f41f3908..f7f7b2560b7 100644
--- a/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160.yaml
@@ -25,7 +25,7 @@ task:
one_hot: true
label_smoothing: 0.1
train_data:
- input_path: 'imagenet-2012-tfrecord/train*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
is_training: true
global_batch_size: 4096
dtype: 'bfloat16'
@@ -34,7 +34,7 @@ task:
randaug:
magnitude: 10
validation_data:
- input_path: 'imagenet-2012-tfrecord/valid*'
+ input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
is_training: false
global_batch_size: 4096
dtype: 'bfloat16'
diff --git a/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160_gpu.yaml
new file mode 100644
index 00000000000..99cfcfa26e5
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_resnetrs50_i160_gpu.yaml
@@ -0,0 +1,66 @@
+# ResNet-RS-50 ImageNet classification.
+# This config is for demonstration purpose only. Parameters such as batch_size, train_steps should be overridden when used.
+# --experiment_type=resnet_rs_imagenet
+runtime:
+ distribution_strategy: 'mirrored'
+ mixed_precision_dtype: 'float16'
+ loss_scale: 'dynamic'
+task:
+ model:
+ num_classes: 1001
+ input_size: [160, 160, 3]
+ backbone:
+ type: 'resnet'
+ resnet:
+ model_id: 50
+ replace_stem_max_pool: true
+ resnetd_shortcut: true
+ se_ratio: 0.25
+ stem_type: 'v1'
+ stochastic_depth_drop_rate: 0.0
+ norm_activation:
+ activation: 'swish'
+ norm_momentum: 0.0
+ use_sync_bn: false
+ dropout_rate: 0.25
+ losses:
+ l2_weight_decay: 0.00004
+ one_hot: true
+ label_smoothing: 0.1
+ train_data:
+ is_training: true
+ global_batch_size: 64
+ dtype: 'float16'
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 10
+ validation_data:
+ is_training: false
+ global_batch_size: 64
+ dtype: 'float16'
+ drop_remainder: false
+trainer:
+ best_checkpoint_eval_metric: accuracy
+ best_checkpoint_export_subdir: best_ckpt
+ train_steps: 109200
+ validation_steps: 13
+ validation_interval: 312
+ steps_per_loop: 312
+ summary_interval: 312
+ checkpoint_interval: 312
+ optimizer_config:
+ ema: null
+ optimizer:
+ type: 'sgd'
+ sgd:
+ momentum: 0.9
+ learning_rate:
+ type: 'cosine'
+ cosine:
+ initial_learning_rate: 1.6
+ decay_steps: 109200
+ warmup:
+ type: 'linear'
+ linear:
+ warmup_steps: 1560
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitb16_i224_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i224_gpu.yaml
new file mode 100644
index 00000000000..39dc5a9da84
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i224_gpu.yaml
@@ -0,0 +1,81 @@
+# ViT-B16 ImageNet classification using 224x224 input.
+# Use 8 V100 GPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-b16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: float16
+ global_batch_size: 512
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: float16
+ global_batch_size: 2048
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 2496
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 374400
+ initial_learning_rate: 0.0005
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 12480
+ type: linear
+ steps_per_loop: 2496
+ summary_interval: 2496
+ train_steps: 374400
+ validation_interval: 2496
+ validation_steps: 25
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitb16_i224_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i224_tpu.yaml
new file mode 100644
index 00000000000..f6e1a43f73c
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i224_tpu.yaml
@@ -0,0 +1,81 @@
+# ViT-B16 ImageNet classification using 224x224 input. Top-1 accuracy 81.3.
+# Use 4x4 DF TPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: bfloat16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-b16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 93600
+ initial_learning_rate: 0.004
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 1560
+ type: linear
+ steps_per_loop: 312
+ summary_interval: 312
+ train_steps: 93600
+ validation_interval: 312
+ validation_steps: 12
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitb16_i384_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i384_gpu.yaml
new file mode 100644
index 00000000000..167d5e62657
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i384_gpu.yaml
@@ -0,0 +1,84 @@
+# ViT-B16 ImageNet classification using 384x384 input.
+# Use 8 V100 GPU for training and validation.
+# --experiment_type=vit_imagenet_finetune
+runtime:
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-b16
+ original_init: false
+ patch_size: 16
+ representation_size: 0
+ dropout_rate: 0.0
+ input_size: [384, 384, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ is_multilabel: false
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: float16
+ global_batch_size: 128
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ is_multilabel: false
+ dtype: float16
+ global_batch_size: 1024
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 4992
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 1497600
+ initial_learning_rate: 0.000175
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 49800
+ type: linear
+ steps_per_loop: 4992
+ summary_interval: 4992
+ train_steps: 1497600 # 300 epochs
+ validation_interval: 4992
+ validation_steps: 49
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitb16_i384_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i384_tpu.yaml
new file mode 100644
index 00000000000..940c7164f3d
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitb16_i384_tpu.yaml
@@ -0,0 +1,84 @@
+# ViT-B16 ImageNet classification using 384x384 input.
+# Use 4x4 DF TPU for training and validation.
+# --experiment_type=vit_imagenet_finetune
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: bfloat16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-b16
+ original_init: false
+ patch_size: 16
+ representation_size: 0
+ dropout_rate: 0.0
+ input_size: [384, 384, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ is_multilabel: false
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ is_multilabel: false
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 93600
+ initial_learning_rate: 0.004
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 1560
+ type: linear
+ steps_per_loop: 312
+ summary_interval: 312
+ train_steps: 93600
+ validation_interval: 312
+ validation_steps: 12
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vith14_i384_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vith14_i384_gpu.yaml
new file mode 100644
index 00000000000..f9143a958f3
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vith14_i384_gpu.yaml
@@ -0,0 +1,84 @@
+# ViT-H14 ImageNet classification using 384x384 input.
+# Use 8 V100 GPU for training and validation.
+# --experiment_type=vit_imagenet_finetune
+runtime:
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-h14
+ original_init: false
+ patch_size: 16
+ representation_size: 0
+ dropout_rate: 0.0
+ input_size: [384, 384, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ is_multilabel: false
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: float16
+ global_batch_size: 8
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ is_multilabel: false
+ dtype: float16
+ global_batch_size: 128
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 9984
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 1198080
+ initial_learning_rate: 0.00001
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 99600
+ type: linear
+ steps_per_loop: 9984
+ summary_interval: 9984
+ train_steps: 1198080 # 15 epochs
+ validation_interval: 9984
+ validation_steps: 316
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vith14_i384_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vith14_i384_tpu.yaml
new file mode 100644
index 00000000000..dfd19f814a1
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vith14_i384_tpu.yaml
@@ -0,0 +1,84 @@
+# ViT-H14 ImageNet classification using 384x384 input.
+# Use 8x16 DF TPU for training and validation.
+# --experiment_type=vit_imagenet_finetune
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: bfloat16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-h14
+ original_init: false
+ patch_size: 16
+ representation_size: 0
+ dropout_rate: 0.0
+ input_size: [384, 384, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ is_multilabel: false
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: bfloat16
+ global_batch_size: 1024
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ is_multilabel: false
+ dtype: bfloat16
+ global_batch_size: 1024
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 1248
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 187200
+ initial_learning_rate: 0.001
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 6240
+ type: linear
+ steps_per_loop: 1248
+ summary_interval: 1248
+ train_steps: 187200
+ validation_interval: 1248
+ validation_steps: 49
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitl16_i224_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i224_gpu.yaml
new file mode 100644
index 00000000000..7f2d89dc98d
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i224_gpu.yaml
@@ -0,0 +1,80 @@
+# ViT-L16 ImageNet classification using 224x224 input.
+# Use 8 V100 GPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-l16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: float16
+ global_batch_size: 128
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: float16
+ global_batch_size: 1024
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+trainer:
+ checkpoint_interval: 4992
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 1497600
+ initial_learning_rate: 0.000175
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 49800
+ type: linear
+ steps_per_loop: 4992
+ summary_interval: 4992
+ train_steps: 1497600 # 300 epochs
+ validation_interval: 4992
+ validation_steps: 49
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitl16_i224_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i224_tpu.yaml
new file mode 100644
index 00000000000..2f69fd689ac
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i224_tpu.yaml
@@ -0,0 +1,80 @@
+# ViT-L16 ImageNet classification using 224x224 input. Top-1 accuracy 82.2.
+# Use 8x8 DF TPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: bfloat16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-l16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+trainer:
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 93600
+ initial_learning_rate: 0.004
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 1560
+ type: linear
+ steps_per_loop: 312
+ summary_interval: 312
+ train_steps: 93600
+ validation_interval: 312
+ validation_steps: 12
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitl16_i384_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i384_gpu.yaml
new file mode 100644
index 00000000000..b00c3e34550
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i384_gpu.yaml
@@ -0,0 +1,84 @@
+# ViT-L16 ImageNet classification using 384x384 input.
+# Use 8 V100 GPU for training and validation.
+# --experiment_type=vit_imagenet_finetune
+runtime:
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-l16
+ original_init: false
+ patch_size: 16
+ representation_size: 0
+ dropout_rate: 0.0
+ input_size: [384, 384, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ is_multilabel: false
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: float16
+ global_batch_size: 32
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ is_multilabel: false
+ dtype: float16
+ global_batch_size: 256
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 9984
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 599040
+ initial_learning_rate: 0.000044
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 49800
+ type: linear
+ steps_per_loop: 9984
+ summary_interval: 9984
+ train_steps: 599040 # 30 epochs
+ validation_interval: 9984
+ validation_steps: 158
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitl16_i384_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i384_tpu.yaml
new file mode 100644
index 00000000000..386cff10334
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitl16_i384_tpu.yaml
@@ -0,0 +1,84 @@
+# ViT-L16 ImageNet classification using 384x384 input.
+# Use 4x8 DF TPU for training and validation.
+# --experiment_type=vit_imagenet_finetune
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: bfloat16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-l16
+ original_init: false
+ patch_size: 16
+ representation_size: 0
+ dropout_rate: 0.0
+ input_size: [384, 384, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ is_multilabel: false
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ is_multilabel: false
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 93600
+ initial_learning_rate: 0.004
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 1560
+ type: linear
+ steps_per_loop: 312
+ summary_interval: 312
+ train_steps: 93600
+ validation_interval: 312
+ validation_steps: 12
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vits16_i224_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vits16_i224_gpu.yaml
new file mode 100644
index 00000000000..4f2f210eb6d
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vits16_i224_gpu.yaml
@@ -0,0 +1,81 @@
+# ViT-S16 ImageNet classification using 224x224 input.
+# Use 8 V100 GPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-s16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: float16
+ global_batch_size: 1024
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: float16
+ global_batch_size: 2048
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 624
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 187200
+ initial_learning_rate: 0.002
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 3120
+ type: linear
+ steps_per_loop: 624
+ summary_interval: 624
+ train_steps: 187200
+ validation_interval: 624
+ validation_steps: 25
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vits16_i224_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vits16_i224_tpu.yaml
new file mode 100644
index 00000000000..b36378ac383
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vits16_i224_tpu.yaml
@@ -0,0 +1,81 @@
+# ViT-S16 ImageNet classification using 224x224 input. Top-1 accuracy 79.4.
+# Use 4x8 DF TPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: bfloat16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-s16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 93600
+ initial_learning_rate: 0.004
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 1560
+ type: linear
+ steps_per_loop: 312
+ summary_interval: 312
+ train_steps: 93600
+ validation_interval: 312
+ validation_steps: 12
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitti16_i224_gpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitti16_i224_gpu.yaml
new file mode 100644
index 00000000000..526c5d4bb16
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitti16_i224_gpu.yaml
@@ -0,0 +1,81 @@
+# ViT-Ti16 ImageNet classification using 224x224 input.
+# Use 8 V100 GPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-ti16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: float16
+ global_batch_size: 2048
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: float16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 624
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 187200
+ initial_learning_rate: 0.002
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 3120
+ type: linear
+ steps_per_loop: 624
+ summary_interval: 624
+ train_steps: 187200
+ validation_interval: 624
+ validation_steps: 12
diff --git a/official/vision/configs/experiments/image_classification/imagenet_vitti16_i224_tpu.yaml b/official/vision/configs/experiments/image_classification/imagenet_vitti16_i224_tpu.yaml
new file mode 100644
index 00000000000..ff3401e33da
--- /dev/null
+++ b/official/vision/configs/experiments/image_classification/imagenet_vitti16_i224_tpu.yaml
@@ -0,0 +1,81 @@
+# ViT-Ti16 ImageNet classification using 224x224 input. Top-1 accuracy 73.4.
+# Use 4x4 DF TPU for training and validation.
+# --experiment_type=deit_imagenet_pretrain
+runtime:
+ distribution_strategy: tpu
+ mixed_precision_dtype: bfloat16
+task:
+ losses:
+ l2_weight_decay: 0.0
+ label_smoothing: 0.1
+ one_hot: false
+ soft_labels: true
+ model:
+ backbone:
+ type: vit
+ vit:
+ init_stochastic_depth_rate: 0.1
+ model_name: vit-ti16
+ original_init: false
+ patch_size: 16
+ dropout_rate: 0.0
+ input_size: [224, 224, 3]
+ kernel_initializer: zeros
+ norm_activation:
+ activation: relu
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: false
+ num_classes: 1001
+ train_data:
+ aug_rand_hflip: true
+ aug_type:
+ randaug:
+ cutout_const: 40
+ exclude_ops: [Cutout]
+ magnitude: 9
+ magnitude_std: 0.0
+ num_layers: 2
+ translate_const: 10
+ type: randaug
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*
+ is_training: true
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+ randaug_magnitude: 10
+ validation_data:
+ dtype: bfloat16
+ global_batch_size: 4096
+ input_path: gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*
+ is_training: false
+ drop_remainder: false
+trainer:
+ checkpoint_interval: 312
+ optimizer_config:
+ learning_rate:
+ cosine:
+ alpha: 0.0
+ decay_steps: 312000
+ initial_learning_rate: 0.004
+ type: cosine
+ optimizer:
+ adamw:
+ include_in_weight_decay: .*(kernel|weight):0$
+ weight_decay_rate: 0.05
+ type: adamw
+ warmup:
+ linear:
+ warmup_learning_rate: 0
+ warmup_steps: 1560
+ type: linear
+ steps_per_loop: 312
+ summary_interval: 312
+ train_steps: 312000
+ validation_interval: 312
+ validation_steps: 12
diff --git a/official/vision/configs/experiments/maskrcnn/coco_mobilenetv2_mrcnn_tpu.yaml b/official/vision/configs/experiments/maskrcnn/coco_mobilenetv2_mrcnn_tpu.yaml
new file mode 100644
index 00000000000..6380eafb8fc
--- /dev/null
+++ b/official/vision/configs/experiments/maskrcnn/coco_mobilenetv2_mrcnn_tpu.yaml
@@ -0,0 +1,20 @@
+# Expect to reach: box mAP: 33.3%, mask mAP: 29.4% on COCO
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ init_checkpoint: gs://**/mobilenetv2_gpu/22984194/ckpt-625500
+ init_checkpoint_modules: 'backbone'
+ train_data:
+ parser:
+ aug_rand_hflip: true
+ aug_scale_min: 0.1
+ aug_scale_max: 2.0
+ losses:
+ l2_weight_decay: 0.00004
+ model:
+ anchor:
+ anchor_size: 3.0
+ num_scales: 3
+ detection_generator:
+ pre_nms_top_k: 1000
diff --git a/official/vision/configs/experiments/retinanet/coco_mobilenetv2_tpu.yaml b/official/vision/configs/experiments/retinanet/coco_mobilenetv2_tpu.yaml
index 9e27bfe8c72..60a97bc0588 100644
--- a/official/vision/configs/experiments/retinanet/coco_mobilenetv2_tpu.yaml
+++ b/official/vision/configs/experiments/retinanet/coco_mobilenetv2_tpu.yaml
@@ -21,6 +21,7 @@ task:
fpn:
num_filters: 128
use_separable_conv: true
+ use_keras_layer: true
head:
num_convs: 4
num_filters: 128
@@ -43,8 +44,9 @@ task:
aug_scale_min: 0.5
validation_data:
dtype: 'bfloat16'
- global_batch_size: 8
+ global_batch_size: 256
is_training: false
+ drop_remainder: false
trainer:
optimizer_config:
learning_rate:
@@ -59,4 +61,4 @@ trainer:
steps_per_loop: 462
train_steps: 277200
validation_interval: 462
- validation_steps: 625
+ validation_steps: 20
diff --git a/official/vision/configs/experiments/retinanet/coco_mobilenetv3.5_avg_tpu.yaml b/official/vision/configs/experiments/retinanet/coco_mobilenetv3.5_avg_tpu.yaml
new file mode 100644
index 00000000000..3d8531335fb
--- /dev/null
+++ b/official/vision/configs/experiments/retinanet/coco_mobilenetv3.5_avg_tpu.yaml
@@ -0,0 +1,65 @@
+# --experiment_type=retinanet_mobile_coco
+# COCO AP 24.92%
+# Use 4x4 DF for training.
+runtime:
+ distribution_strategy: 'tpu'
+ mixed_precision_dtype: 'bfloat16'
+task:
+ losses:
+ l2_weight_decay: 3.0e-05
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ mobilenet:
+ model_id: 'MobileNetMultiAVG'
+ filter_size_scale: 1.0
+ type: 'mobilenet'
+ decoder:
+ type: 'fpn'
+ fpn:
+ num_filters: 128
+ use_separable_conv: true
+ use_keras_layer: true
+ head:
+ num_convs: 4
+ num_filters: 128
+ use_separable_conv: true
+ input_size: [256, 256, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'relu6'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'bfloat16'
+ global_batch_size: 256
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ dtype: 'bfloat16'
+ global_batch_size: 256
+ is_training: false
+ drop_remainder: false
+trainer:
+ optimizer_config:
+ learning_rate:
+ stepwise:
+ boundaries: [263340, 272580]
+ values: [0.32, 0.032, 0.0032]
+ type: 'stepwise'
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0067
+ warmup_steps: 2000
+ steps_per_loop: 462
+ train_steps: 277200
+ validation_interval: 462
+ validation_steps: 20
diff --git a/official/vision/configs/experiments/retinanet/coco_spinenet143_gpu_multiworker_mirrored.yaml b/official/vision/configs/experiments/retinanet/coco_spinenet143_gpu_multiworker_mirrored.yaml
new file mode 100644
index 00000000000..8481145d292
--- /dev/null
+++ b/official/vision/configs/experiments/retinanet/coco_spinenet143_gpu_multiworker_mirrored.yaml
@@ -0,0 +1,77 @@
+# SpineNet-143 COCO detection with protocol C config. Expect best AP50 at 5.73%.
+# This is a demostration only, not a full training to reproduce numbers reported in the paper.
+# --experiment_type=retinanet_spinenet_coco
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: multi_worker_mirrored
+ mixed_precision_dtype: float16
+task:
+ annotation_file: null
+ losses:
+ l2_weight_decay: 4.0e-05
+ model:
+ anchor:
+ anchor_size: 4
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ spinenet:
+ stochastic_depth_drop_rate: 0.2
+ model_id: '143'
+ type: 'spinenet'
+ decoder:
+ type: 'identity'
+ head:
+ num_convs: 4
+ num_filters: 256
+ input_size: [1280, 1280, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'float16'
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ type: simple_decoder
+ global_batch_size: 8
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.1
+ validation_data:
+ dtype: 'float16'
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ type: simple_decoder
+ drop_remainder: false
+ global_batch_size: 8
+ is_training: false
+trainer:
+ best_checkpoint_export_subdir: "best_ckpt"
+ best_checkpoint_eval_metric: "AP50"
+ checkpoint_interval: 462
+ optimizer_config:
+ optimizer:
+ sgd:
+ global_clipnorm: 10.0
+ momentum: 0.9
+ learning_rate:
+ cosine:
+ decay_steps: 23100
+ initial_learning_rate: 0.032
+ type: 'cosine'
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0067
+ warmup_steps: 2000
+ steps_per_loop: 462
+ train_steps: 23100
+ validation_interval: 462
+ validation_steps: 625
diff --git a/official/vision/configs/experiments/retinanet/coco_spinenet49_gpu_multiworker_mirrored.yaml b/official/vision/configs/experiments/retinanet/coco_spinenet49_gpu_multiworker_mirrored.yaml
new file mode 100644
index 00000000000..4711a20ed42
--- /dev/null
+++ b/official/vision/configs/experiments/retinanet/coco_spinenet49_gpu_multiworker_mirrored.yaml
@@ -0,0 +1,77 @@
+# Example of SpineNet-49 COCO detection with protocol C config. Expect best AP50 at 38.06%.
+# This is a demostration only, not a full training to reproduce numbers reported in the paper.
+# --experiment_type=retinanet_spinenet_coco
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: multi_worker_mirrored
+ mixed_precision_dtype: float16
+task:
+ annotation_file: null
+ losses:
+ l2_weight_decay: 4.0e-05
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ spinenet:
+ stochastic_depth_drop_rate: 0.2
+ model_id: '49'
+ type: 'spinenet'
+ decoder:
+ type: 'identity'
+ head:
+ num_convs: 4
+ num_filters: 256
+ input_size: [640, 640, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'float16'
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ type: simple_decoder
+ global_batch_size: 32
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.1
+ validation_data:
+ dtype: 'float16'
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ type: simple_decoder
+ drop_remainder: false
+ global_batch_size: 8
+ is_training: false
+trainer:
+ best_checkpoint_export_subdir: "best_ckpt"
+ best_checkpoint_eval_metric: "AP50"
+ checkpoint_interval: 462
+ optimizer_config:
+ optimizer:
+ sgd:
+ global_clipnorm: 10.0
+ momentum: 0.9
+ learning_rate:
+ cosine:
+ decay_steps: 23100
+ initial_learning_rate: 0.32
+ type: 'cosine'
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0067
+ warmup_steps: 2000
+ steps_per_loop: 462
+ train_steps: 23100
+ validation_interval: 462
+ validation_steps: 625
diff --git a/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml b/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml
index e1e14b321f0..27f613f9518 100644
--- a/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml
+++ b/official/vision/configs/experiments/retinanet/coco_spinenet49_mobile_tpu.yaml
@@ -1,4 +1,5 @@
# --experiment_type=retinanet_mobile_coco
+# COCO mAP: 27.6
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'bfloat16'
@@ -26,7 +27,7 @@ task:
max_level: 7
min_level: 3
norm_activation:
- activation: 'swish'
+ activation: 'hard_swish'
norm_epsilon: 0.001
norm_momentum: 0.99
use_sync_bn: true
@@ -40,8 +41,9 @@ task:
aug_scale_min: 0.5
validation_data:
dtype: 'bfloat16'
- global_batch_size: 8
+ global_batch_size: 256
is_training: false
+ drop_remainder: false
trainer:
checkpoint_interval: 462
optimizer_config:
@@ -57,4 +59,4 @@ trainer:
steps_per_loop: 462
train_steps: 277200
validation_interval: 462
- validation_steps: 625
+ validation_steps: 20
diff --git a/official/vision/configs/experiments/retinanet/coco_spinenet49s_mobile_tpu.yaml b/official/vision/configs/experiments/retinanet/coco_spinenet49s_mobile_tpu.yaml
index 9f854ccf44c..4e82339475f 100644
--- a/official/vision/configs/experiments/retinanet/coco_spinenet49s_mobile_tpu.yaml
+++ b/official/vision/configs/experiments/retinanet/coco_spinenet49s_mobile_tpu.yaml
@@ -1,4 +1,5 @@
# --experiment_type=retinanet_mobile_coco
+# COCO mAP: 23.5
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'bfloat16'
@@ -26,7 +27,7 @@ task:
max_level: 7
min_level: 3
norm_activation:
- activation: 'swish'
+ activation: 'hard_swish'
norm_epsilon: 0.001
norm_momentum: 0.99
use_sync_bn: true
@@ -40,8 +41,9 @@ task:
aug_scale_min: 0.5
validation_data:
dtype: 'bfloat16'
- global_batch_size: 8
+ global_batch_size: 256
is_training: false
+ drop_remainder: false
trainer:
checkpoint_interval: 462
optimizer_config:
@@ -57,4 +59,4 @@ trainer:
steps_per_loop: 462
train_steps: 277200
validation_interval: 462
- validation_steps: 625
+ validation_steps: 20
diff --git a/official/vision/configs/experiments/retinanet/coco_spinenet49xs_mobile_tpu.yaml b/official/vision/configs/experiments/retinanet/coco_spinenet49xs_mobile_tpu.yaml
index 926bd62097a..be895c3feef 100644
--- a/official/vision/configs/experiments/retinanet/coco_spinenet49xs_mobile_tpu.yaml
+++ b/official/vision/configs/experiments/retinanet/coco_spinenet49xs_mobile_tpu.yaml
@@ -1,4 +1,5 @@
# --experiment_type=retinanet_mobile_coco
+# COCO mAP: 16.5
runtime:
distribution_strategy: 'tpu'
mixed_precision_dtype: 'bfloat16'
@@ -26,7 +27,7 @@ task:
max_level: 7
min_level: 3
norm_activation:
- activation: 'swish'
+ activation: 'hard_swish'
norm_epsilon: 0.001
norm_momentum: 0.99
use_sync_bn: true
@@ -40,8 +41,9 @@ task:
aug_scale_min: 0.5
validation_data:
dtype: 'bfloat16'
- global_batch_size: 8
+ global_batch_size: 256
is_training: false
+ drop_remainder: false
trainer:
checkpoint_interval: 462
optimizer_config:
@@ -57,4 +59,4 @@ trainer:
steps_per_loop: 462
train_steps: 277200
validation_interval: 462
- validation_steps: 625
+ validation_steps: 20
diff --git a/official/vision/configs/experiments/retinanet/coco_spinenet96_gpu_multiworker_mirrored.yaml b/official/vision/configs/experiments/retinanet/coco_spinenet96_gpu_multiworker_mirrored.yaml
new file mode 100644
index 00000000000..2776a3779cb
--- /dev/null
+++ b/official/vision/configs/experiments/retinanet/coco_spinenet96_gpu_multiworker_mirrored.yaml
@@ -0,0 +1,77 @@
+# Example of SpineNet-96 COCO detection with protocol C config. Expect best AP50 at 22.97%.
+# This is a demostration only, not a full training to reproduce numbers reported in the paper.
+# --experiment_type=retinanet_spinenet_coco
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: multi_worker_mirrored
+ mixed_precision_dtype: float16
+task:
+ annotation_file: null
+ losses:
+ l2_weight_decay: 4.0e-05
+ model:
+ anchor:
+ anchor_size: 3
+ aspect_ratios: [0.5, 1.0, 2.0]
+ num_scales: 3
+ backbone:
+ spinenet:
+ stochastic_depth_drop_rate: 0.2
+ model_id: '96'
+ type: 'spinenet'
+ decoder:
+ type: 'identity'
+ head:
+ num_convs: 4
+ num_filters: 256
+ input_size: [1024, 1024, 3]
+ max_level: 7
+ min_level: 3
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ train_data:
+ dtype: 'float16'
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ type: simple_decoder
+ global_batch_size: 16
+ is_training: true
+ parser:
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.1
+ validation_data:
+ dtype: 'float16'
+ decoder:
+ simple_decoder:
+ regenerate_source_id: true
+ type: simple_decoder
+ drop_remainder: false
+ global_batch_size: 8
+ is_training: false
+trainer:
+ best_checkpoint_export_subdir: "best_ckpt"
+ best_checkpoint_eval_metric: "AP50"
+ checkpoint_interval: 462
+ optimizer_config:
+ optimizer:
+ sgd:
+ global_clipnorm: 10.0
+ momentum: 0.9
+ learning_rate:
+ cosine:
+ decay_steps: 23100
+ initial_learning_rate: 0.032
+ type: 'cosine'
+ warmup:
+ linear:
+ warmup_learning_rate: 0.0067
+ warmup_steps: 2000
+ steps_per_loop: 462
+ train_steps: 23100
+ validation_interval: 462
+ validation_steps: 625
diff --git a/official/vision/configs/experiments/semantic_segmentation/deeplabv3plus_resnet101_cityscapes_gpu.yaml b/official/vision/configs/experiments/semantic_segmentation/deeplabv3plus_resnet101_cityscapes_gpu.yaml
new file mode 100644
index 00000000000..861f6b228e2
--- /dev/null
+++ b/official/vision/configs/experiments/semantic_segmentation/deeplabv3plus_resnet101_cityscapes_gpu.yaml
@@ -0,0 +1,80 @@
+# Use your own cityscapes preprocessed dataset.
+# --experiment_type=seg_deeplabv3plus_cityscapes.
+# Use 8 V100 GPU for training and validation.
+# Achieves 69.4% meanIoU when training from scratch.
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: mirrored
+ mixed_precision_dtype: float16
+task:
+ init_checkpoint: null
+ model:
+ num_classes: 19
+ input_size: [null, null, 3]
+ backbone:
+ type: 'dilated_resnet'
+ dilated_resnet:
+ model_id: 101
+ output_stride: 16
+ stem_type: 'v1'
+ se_ratio: 0.25
+ stochastic_depth_drop_rate: 0.2
+ multigrid: [1, 2, 4]
+ last_stage_repeats: 1
+ decoder:
+ aspp:
+ pool_kernel_size: [512, 1024]
+ head:
+ feature_fusion: 'deeplabv3plus'
+ low_level: 2
+ low_level_num_filters: 48
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ losses:
+ top_k_percent_pixels: 1.0 # only backpropagate loss for the topk 100% pixels.
+ train_data:
+ output_size: [1024, 2048]
+ crop_size: [512, 1024]
+ is_training: true
+ global_batch_size: 16
+ dtype: 'float16'
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ output_size: [1024, 2048]
+ is_training: false
+ global_batch_size: 16
+ dtype: 'float16'
+ drop_remainder: false
+ resize_eval_groundtruth: true
+trainer:
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 90000
+ initial_learning_rate: 0.01
+ power: 0.9
+ type: polynomial
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 925
+ type: linear
+ steps_per_loop: 185
+ summary_interval: 185
+ train_steps: 90000
+ validation_interval: 185
+ validation_steps: 31
+ checkpoint_interval: 185
diff --git a/official/vision/configs/experiments/semantic_segmentation/deeplabv3plus_resnet101_cityscapes_gpu_multiworker_mirrored.yaml b/official/vision/configs/experiments/semantic_segmentation/deeplabv3plus_resnet101_cityscapes_gpu_multiworker_mirrored.yaml
new file mode 100644
index 00000000000..60ceb192adf
--- /dev/null
+++ b/official/vision/configs/experiments/semantic_segmentation/deeplabv3plus_resnet101_cityscapes_gpu_multiworker_mirrored.yaml
@@ -0,0 +1,77 @@
+# Use your own cityscapes preprocessed dataset. 79% meanIoU.
+# --experiment_type=seg_deeplabv3plus_cityscapes
+runtime:
+ all_reduce_alg: nccl
+ distribution_strategy: multi_worker_mirrored
+ mixed_precision_dtype: float16
+task:
+ init_checkpoint: null
+ model:
+ num_classes: 19
+ input_size: [null, null, 3]
+ backbone:
+ type: 'dilated_resnet'
+ dilated_resnet:
+ model_id: 101
+ output_stride: 16
+ stem_type: 'v1'
+ se_ratio: 0.25
+ stochastic_depth_drop_rate: 0.2
+ multigrid: [1, 2, 4]
+ last_stage_repeats: 1
+ decoder:
+ aspp:
+ pool_kernel_size: [512, 1024]
+ head:
+ feature_fusion: 'deeplabv3plus'
+ low_level: 2
+ low_level_num_filters: 48
+ norm_activation:
+ activation: 'swish'
+ norm_epsilon: 0.001
+ norm_momentum: 0.99
+ use_sync_bn: true
+ losses:
+ top_k_percent_pixels: 1.0 # only backpropagate loss for the topk 100% pixels.
+ train_data:
+ output_size: [1024, 2048]
+ is_training: true
+ global_batch_size: 16
+ dtype: 'float32'
+ aug_rand_hflip: true
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+ validation_data:
+ output_size: [1024, 2048]
+ is_training: false
+ global_batch_size: 16
+ dtype: 'float32'
+ drop_remainder: false
+ resize_eval_groundtruth: true
+trainer:
+ best_checkpoint_eval_metric: 'mean_iou'
+ best_checkpoint_export_subdir: 'best_ckpt'
+ best_checkpoint_metric_comp: 'higher'
+ optimizer_config:
+ learning_rate:
+ polynomial:
+ decay_steps: 90000
+ initial_learning_rate: 0.01
+ power: 0.9
+ type: polynomial
+ optimizer:
+ sgd:
+ momentum: 0.9
+ type: sgd
+ warmup:
+ linear:
+ name: linear
+ warmup_learning_rate: 0
+ warmup_steps: 925
+ type: linear
+ steps_per_loop: 185
+ summary_interval: 185
+ train_steps: 90000
+ validation_interval: 185
+ validation_steps: 31
+ checkpoint_interval: 185
diff --git a/official/vision/configs/experiments/video_classification/k400_resnet3drs_50_tpu.yaml b/official/vision/configs/experiments/video_classification/k400_resnet3drs_50_tpu.yaml
index 83875d1273a..baff531b6d4 100644
--- a/official/vision/configs/experiments/video_classification/k400_resnet3drs_50_tpu.yaml
+++ b/official/vision/configs/experiments/video_classification/k400_resnet3drs_50_tpu.yaml
@@ -42,7 +42,6 @@ task:
is_training: true
min_image_size: 256
name: kinetics400
- num_channels: 3
num_classes: 400
num_examples: 215570
num_test_clips: 1
@@ -67,7 +66,6 @@ task:
is_training: false
min_image_size: 256
name: kinetics400
- num_channels: 3
num_classes: 400
num_examples: 17706
num_test_clips: 10
diff --git a/official/vision/configs/image_classification.py b/official/vision/configs/image_classification.py
index ab4bdb56a29..b7301cacbd2 100644
--- a/official/vision/configs/image_classification.py
+++ b/official/vision/configs/image_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Image classification configuration definition."""
import dataclasses
import os
-from typing import List, Optional
+from typing import List, Optional, Tuple, Union, Sequence
from official.core import config_definitions as cfg
from official.core import exp_factory
@@ -28,7 +28,8 @@
@dataclasses.dataclass
class DataConfig(cfg.DataConfig):
"""Input config for training."""
- input_path: str = ''
+ input_path: Union[Sequence[str], str, hyperparams.Config] = ''
+ weights: Optional[hyperparams.base_config.Config] = None
global_batch_size: int = 0
is_training: bool = True
dtype: str = 'float32'
@@ -36,8 +37,11 @@ class DataConfig(cfg.DataConfig):
cycle_length: int = 10
is_multilabel: bool = False
aug_rand_hflip: bool = True
+ aug_crop: Optional[bool] = True
+ crop_area_range: Optional[Tuple[float, float]] = (0.08, 1.0)
aug_type: Optional[
common.Augmentation] = None # Choose from AutoAugment and RandAugment.
+ three_augment: bool = False
color_jitter: float = 0.
random_erasing: Optional[common.RandomErasing] = None
file_type: str = 'tfrecord'
@@ -45,11 +49,21 @@ class DataConfig(cfg.DataConfig):
label_field_key: str = 'image/class/label'
decode_jpeg_only: bool = True
mixup_and_cutmix: Optional[common.MixupAndCutmix] = None
- decoder: Optional[common.DataDecoder] = common.DataDecoder()
+ decoder: Optional[common.DataDecoder] = dataclasses.field(
+ default_factory=common.DataDecoder
+ )
# Keep for backward compatibility.
aug_policy: Optional[str] = None # None, 'autoaug', or 'randaug'.
randaug_magnitude: Optional[int] = 10
+ # Determines ratio between the side of the cropped image and the short side of
+ # the original image.
+ center_crop_fraction: Optional[float] = 0.875
+ # Interpolation method for resizing image in Parser for both training and eval
+ tf_resize_method: str = 'bilinear'
+ # Repeat augmentation puts multiple augmentations of the same image in a batch
+ # https://arxiv.org/abs/1902.05509
+ repeated_augment: Optional[int] = None
@dataclasses.dataclass
@@ -57,14 +71,20 @@ class ImageClassificationModel(hyperparams.Config):
"""The model config."""
num_classes: int = 0
input_size: List[int] = dataclasses.field(default_factory=list)
- backbone: backbones.Backbone = backbones.Backbone(
- type='resnet', resnet=backbones.ResNet())
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='resnet', resnet=backbones.ResNet()
+ )
+ )
dropout_rate: float = 0.0
- norm_activation: common.NormActivation = common.NormActivation(
- use_sync_bn=False)
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation(use_sync_bn=False)
+ )
# Adds a BatchNormalization layer pre-GlobalAveragePooling in classification
add_head_batch_norm: bool = False
kernel_initializer: str = 'random_uniform'
+ # Whether to output softmax results instead of logits.
+ output_softmax: bool = False
@dataclasses.dataclass
@@ -74,25 +94,41 @@ class Losses(hyperparams.Config):
label_smoothing: float = 0.0
l2_weight_decay: float = 0.0
soft_labels: bool = False
+ # Converts multi-class classification to multi-label classification. Weights
+ # each object class equally in the loss function, ignoring their size.
+ use_binary_cross_entropy: bool = False
@dataclasses.dataclass
class Evaluation(hyperparams.Config):
top_k: int = 5
+ precision_and_recall_thresholds: Optional[List[float]] = None
+ report_per_class_precision_and_recall: bool = False
@dataclasses.dataclass
class ImageClassificationTask(cfg.TaskConfig):
"""The task config."""
- model: ImageClassificationModel = ImageClassificationModel()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- losses: Losses = Losses()
- evaluation: Evaluation = Evaluation()
+ model: ImageClassificationModel = dataclasses.field(
+ default_factory=ImageClassificationModel
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ evaluation: Evaluation = dataclasses.field(default_factory=Evaluation)
+ train_input_partition_dims: Optional[List[int]] = dataclasses.field(
+ default_factory=list)
+ eval_input_partition_dims: Optional[List[int]] = dataclasses.field(
+ default_factory=list)
init_checkpoint: Optional[str] = None
init_checkpoint_modules: str = 'all' # all or backbone
model_output_keys: Optional[List[int]] = dataclasses.field(
default_factory=list)
+ freeze_backbone: bool = False
@exp_factory.register_config_factory('image_classification')
@@ -395,3 +431,201 @@ def image_classification_imagenet_mobilenet() -> cfg.ExperimentConfig:
])
return config
+
+
+@exp_factory.register_config_factory('deit_imagenet_pretrain')
+def image_classification_imagenet_deit_pretrain() -> cfg.ExperimentConfig:
+ """Image classification on imagenet with vision transformer."""
+ train_batch_size = 4096 # originally was 1024 but 4096 better for tpu v3-32
+ eval_batch_size = 4096 # originally was 1024 but 4096 better for tpu v3-32
+ label_smoothing = 0.1
+ steps_per_epoch = IMAGENET_TRAIN_EXAMPLES // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=ImageClassificationTask(
+ model=ImageClassificationModel(
+ num_classes=1001,
+ input_size=[224, 224, 3],
+ kernel_initializer='zeros',
+ backbone=backbones.Backbone(
+ type='vit',
+ vit=backbones.VisionTransformer(
+ model_name='vit-b16',
+ representation_size=768,
+ init_stochastic_depth_rate=0.1,
+ original_init=False,
+ transformer=backbones.Transformer(
+ dropout_rate=0.0, attention_dropout_rate=0.0)))),
+ losses=Losses(
+ l2_weight_decay=0.0,
+ label_smoothing=label_smoothing,
+ one_hot=False,
+ soft_labels=True),
+ train_data=DataConfig(
+ input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ aug_type=common.Augmentation(
+ type='randaug',
+ randaug=common.RandAugment(
+ magnitude=9, exclude_ops=['Cutout'])),
+ mixup_and_cutmix=common.MixupAndCutmix(
+ label_smoothing=label_smoothing)),
+ validation_data=DataConfig(
+ input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'valid*'),
+ is_training=False,
+ global_batch_size=eval_batch_size)),
+ trainer=cfg.TrainerConfig(
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ train_steps=300 * steps_per_epoch,
+ validation_steps=IMAGENET_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.05,
+ 'include_in_weight_decay': r'.*(kernel|weight):0$',
+ 'gradient_clip_norm': 0.0
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.0005 * train_batch_size / 512,
+ 'decay_steps': 300 * steps_per_epoch,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 5 * steps_per_epoch,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+
+ return config
+
+
+@exp_factory.register_config_factory('vit_imagenet_pretrain')
+def image_classification_imagenet_vit_pretrain() -> cfg.ExperimentConfig:
+ """Image classification on imagenet with vision transformer."""
+ train_batch_size = 4096
+ eval_batch_size = 4096
+ steps_per_epoch = IMAGENET_TRAIN_EXAMPLES // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=ImageClassificationTask(
+ model=ImageClassificationModel(
+ num_classes=1001,
+ input_size=[224, 224, 3],
+ kernel_initializer='zeros',
+ backbone=backbones.Backbone(
+ type='vit',
+ vit=backbones.VisionTransformer(
+ model_name='vit-b16', representation_size=768))),
+ losses=Losses(l2_weight_decay=0.0),
+ train_data=DataConfig(
+ input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size),
+ validation_data=DataConfig(
+ input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'valid*'),
+ is_training=False,
+ global_batch_size=eval_batch_size)),
+ trainer=cfg.TrainerConfig(
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ train_steps=300 * steps_per_epoch,
+ validation_steps=IMAGENET_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'adamw',
+ 'adamw': {
+ 'weight_decay_rate': 0.3,
+ 'include_in_weight_decay': r'.*(kernel|weight):0$',
+ 'gradient_clip_norm': 0.0
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.003 * train_batch_size / 4096,
+ 'decay_steps': 300 * steps_per_epoch,
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 10000,
+ 'warmup_learning_rate': 0
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+
+ return config
+
+
+@exp_factory.register_config_factory('vit_imagenet_finetune')
+def image_classification_imagenet_vit_finetune() -> cfg.ExperimentConfig:
+ """Image classification on imagenet with vision transformer."""
+ train_batch_size = 512
+ eval_batch_size = 512
+ steps_per_epoch = IMAGENET_TRAIN_EXAMPLES // train_batch_size
+ config = cfg.ExperimentConfig(
+ task=ImageClassificationTask(
+ model=ImageClassificationModel(
+ num_classes=1001,
+ input_size=[384, 384, 3],
+ backbone=backbones.Backbone(
+ type='vit',
+ vit=backbones.VisionTransformer(model_name='vit-b16'))),
+ losses=Losses(l2_weight_decay=0.0),
+ train_data=DataConfig(
+ input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size),
+ validation_data=DataConfig(
+ input_path=os.path.join(IMAGENET_INPUT_PATH_BASE, 'valid*'),
+ is_training=False,
+ global_batch_size=eval_batch_size)),
+ trainer=cfg.TrainerConfig(
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ train_steps=20000,
+ validation_steps=IMAGENET_VAL_EXAMPLES // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9,
+ 'global_clipnorm': 1.0,
+ }
+ },
+ 'learning_rate': {
+ 'type': 'cosine',
+ 'cosine': {
+ 'initial_learning_rate': 0.003,
+ 'decay_steps': 20000,
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None'
+ ])
+
+ return config
diff --git a/official/vision/configs/image_classification_test.py b/official/vision/configs/image_classification_test.py
index 3c58c055398..1198dbca5cc 100644
--- a/official/vision/configs/image_classification_test.py
+++ b/official/vision/configs/image_classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for image_classification."""
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -29,7 +29,10 @@ class ImageClassificationConfigTest(tf.test.TestCase, parameterized.TestCase):
('resnet_imagenet',),
('resnet_rs_imagenet',),
('revnet_imagenet',),
- ('mobilenet_imagenet'),
+ ('mobilenet_imagenet',),
+ ('deit_imagenet_pretrain',),
+ ('vit_imagenet_pretrain',),
+ ('vit_imagenet_finetune',),
)
def test_image_classification_configs(self, config_name):
config = exp_factory.get_exp_config(config_name)
diff --git a/official/vision/configs/maskrcnn.py b/official/vision/configs/maskrcnn.py
index 768d871210b..1bbb5ec3ef4 100644
--- a/official/vision/configs/maskrcnn.py
+++ b/official/vision/configs/maskrcnn.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
import dataclasses
import os
-from typing import List, Optional, Union
+from typing import List, Optional, Sequence, Union
from official.core import config_definitions as cfg
from official.core import exp_factory
@@ -34,8 +34,11 @@ class Parser(hyperparams.Config):
match_threshold: float = 0.5
unmatched_threshold: float = 0.5
aug_rand_hflip: bool = False
+ aug_rand_vflip: bool = False
aug_scale_min: float = 1.0
aug_scale_max: float = 1.0
+ aug_type: Optional[
+ common.Augmentation] = None # Choose from AutoAugment and RandAugment.
skip_crowd_during_training: bool = True
max_num_instances: int = 100
rpn_match_threshold: float = 0.7
@@ -43,17 +46,32 @@ class Parser(hyperparams.Config):
rpn_batch_size_per_im: int = 256
rpn_fg_fraction: float = 0.5
mask_crop_size: int = 112
+ pad: bool = True # Only support `pad = True`.
+ keep_aspect_ratio: bool = True # Only support `keep_aspect_ratio = True`.
+
+ def __post_init__(self, *args, **kwargs):
+ """Validates the configuration."""
+ if not self.pad:
+ raise ValueError('`maskrcnn.Parser` only supports `pad = True`.')
+ if not self.keep_aspect_ratio:
+ raise ValueError(
+ '`maskrcnn.Parser` only supports `keep_aspect_ratio = True`.'
+ )
+ super().__post_init__(*args, **kwargs)
@dataclasses.dataclass
class DataConfig(cfg.DataConfig):
"""Input config for training."""
- input_path: str = ''
+ input_path: Union[Sequence[str], str, hyperparams.Config] = ''
+ weights: Optional[hyperparams.Config] = None
global_batch_size: int = 0
is_training: bool = False
dtype: str = 'bfloat16'
- decoder: common.DataDecoder = common.DataDecoder()
- parser: Parser = Parser()
+ decoder: common.DataDecoder = dataclasses.field(
+ default_factory=common.DataDecoder
+ )
+ parser: Parser = dataclasses.field(default_factory=Parser)
shuffle_buffer_size: int = 10000
file_type: str = 'tfrecord'
drop_remainder: bool = True
@@ -133,6 +151,7 @@ class DetectionGenerator(hyperparams.Config):
nms_version: str = 'v2' # `v2`, `v1`, `batched`
use_cpu_nms: bool = False
soft_nms_sigma: Optional[float] = None # Only works when nms_version='v1'.
+ use_sigmoid_probability: bool = False
@dataclasses.dataclass
@@ -161,25 +180,39 @@ class MaskRCNN(hyperparams.Config):
input_size: List[int] = dataclasses.field(default_factory=list)
min_level: int = 2
max_level: int = 6
- anchor: Anchor = Anchor()
+ anchor: Anchor = dataclasses.field(default_factory=Anchor)
include_mask: bool = True
- backbone: backbones.Backbone = backbones.Backbone(
- type='resnet', resnet=backbones.ResNet())
- decoder: decoders.Decoder = decoders.Decoder(
- type='fpn', fpn=decoders.FPN())
- rpn_head: RPNHead = RPNHead()
- detection_head: DetectionHead = DetectionHead()
- roi_generator: ROIGenerator = ROIGenerator()
- roi_sampler: ROISampler = ROISampler()
- roi_aligner: ROIAligner = ROIAligner()
- detection_generator: DetectionGenerator = DetectionGenerator()
- mask_head: Optional[MaskHead] = MaskHead()
- mask_sampler: Optional[MaskSampler] = MaskSampler()
- mask_roi_aligner: Optional[MaskROIAligner] = MaskROIAligner()
- norm_activation: common.NormActivation = common.NormActivation(
- norm_momentum=0.997,
- norm_epsilon=0.0001,
- use_sync_bn=True)
+ outer_boxes_scale: float = 1.0
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone(
+ type='resnet', resnet=backbones.ResNet()
+ )
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder(type='fpn', fpn=decoders.FPN())
+ )
+ rpn_head: RPNHead = dataclasses.field(default_factory=RPNHead)
+ detection_head: DetectionHead = dataclasses.field(
+ default_factory=DetectionHead
+ )
+ roi_generator: ROIGenerator = dataclasses.field(default_factory=ROIGenerator)
+ roi_sampler: ROISampler = dataclasses.field(default_factory=ROISampler)
+ roi_aligner: ROIAligner = dataclasses.field(default_factory=ROIAligner)
+ detection_generator: DetectionGenerator = dataclasses.field(
+ default_factory=DetectionGenerator
+ )
+ mask_head: Optional[MaskHead] = dataclasses.field(default_factory=MaskHead)
+ mask_sampler: Optional[MaskSampler] = dataclasses.field(
+ default_factory=MaskSampler
+ )
+ mask_roi_aligner: Optional[MaskROIAligner] = dataclasses.field(
+ default_factory=MaskROIAligner
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation( # pylint: disable=g-long-lambda
+ norm_momentum=0.997, norm_epsilon=0.0001, use_sync_bn=True
+ )
+ )
@dataclasses.dataclass
@@ -187,21 +220,29 @@ class Losses(hyperparams.Config):
loss_weight: float = 1.0
rpn_huber_loss_delta: float = 1. / 9.
frcnn_huber_loss_delta: float = 1.
+ frcnn_class_use_binary_cross_entropy: bool = False
+ frcnn_class_loss_top_k_percent: float = 1.
l2_weight_decay: float = 0.0
rpn_score_weight: float = 1.0
rpn_box_weight: float = 1.0
frcnn_class_weight: float = 1.0
frcnn_box_weight: float = 1.0
mask_weight: float = 1.0
+ class_weights: Optional[List[float]] = None
@dataclasses.dataclass
class MaskRCNNTask(cfg.TaskConfig):
- model: MaskRCNN = MaskRCNN()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False,
- drop_remainder=False)
- losses: Losses = Losses()
+ model: MaskRCNN = dataclasses.field(default_factory=MaskRCNN)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ is_training=False, drop_remainder=False
+ )
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
init_checkpoint: Optional[str] = None
init_checkpoint_modules: Union[
str, List[str]] = 'all' # all, backbone, and/or decoder
@@ -213,6 +254,9 @@ class MaskRCNNTask(cfg.TaskConfig):
use_coco_metrics: bool = True
# If set, the Waymo Open Dataset evaluator would be used.
use_wod_metrics: bool = False
+ # If set, use instance metrics (AP, mask AP, etc.) computed by an efficient
+ # approximation algorithm with TPU compatible operations.
+ use_approx_instance_metrics: bool = False
# If set, freezes the backbone during training.
# TODO(crisnv) Add paper link when available.
@@ -524,3 +568,91 @@ def cascadercnn_spinenet_coco() -> cfg.ExperimentConfig:
'task.model.max_level == task.model.backbone.spinenet.max_level',
])
return config
+
+
+@exp_factory.register_config_factory('maskrcnn_mobilenet_coco')
+def maskrcnn_mobilenet_coco() -> cfg.ExperimentConfig:
+ """COCO object detection with Mask R-CNN with MobileNet backbone."""
+ steps_per_epoch = 232
+ coco_val_samples = 5000
+ train_batch_size = 512
+ eval_batch_size = 512
+
+ config = cfg.ExperimentConfig(
+ runtime=cfg.RuntimeConfig(mixed_precision_dtype='bfloat16'),
+ task=MaskRCNNTask(
+ annotation_file=os.path.join(COCO_INPUT_PATH_BASE,
+ 'instances_val2017.json'),
+ model=MaskRCNN(
+ backbone=backbones.Backbone(
+ type='mobilenet',
+ mobilenet=backbones.MobileNet(model_id='MobileNetV2')),
+ decoder=decoders.Decoder(
+ type='fpn',
+ fpn=decoders.FPN(num_filters=128, use_separable_conv=True)),
+ rpn_head=RPNHead(use_separable_conv=True,
+ num_filters=128), # 1/2 of original channels.
+ detection_head=DetectionHead(
+ use_separable_conv=True, num_filters=128,
+ fc_dims=512), # 1/2 of original channels.
+ mask_head=MaskHead(use_separable_conv=True,
+ num_filters=128), # 1/2 of original channels.
+ anchor=Anchor(anchor_size=3),
+ norm_activation=common.NormActivation(
+ activation='relu6',
+ norm_momentum=0.99,
+ norm_epsilon=0.001,
+ use_sync_bn=True),
+ num_classes=91,
+ input_size=[512, 512, 3],
+ min_level=3,
+ max_level=6,
+ include_mask=True),
+ losses=Losses(l2_weight_decay=0.00004),
+ train_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'train*'),
+ is_training=True,
+ global_batch_size=train_batch_size,
+ parser=Parser(
+ aug_rand_hflip=True, aug_scale_min=0.5, aug_scale_max=2.0)),
+ validation_data=DataConfig(
+ input_path=os.path.join(COCO_INPUT_PATH_BASE, 'val*'),
+ is_training=False,
+ global_batch_size=eval_batch_size,
+ drop_remainder=False)),
+ trainer=cfg.TrainerConfig(
+ train_steps=steps_per_epoch * 350,
+ validation_steps=coco_val_samples // eval_batch_size,
+ validation_interval=steps_per_epoch,
+ steps_per_loop=steps_per_epoch,
+ summary_interval=steps_per_epoch,
+ checkpoint_interval=steps_per_epoch,
+ optimizer_config=optimization.OptimizationConfig({
+ 'optimizer': {
+ 'type': 'sgd',
+ 'sgd': {
+ 'momentum': 0.9
+ }
+ },
+ 'learning_rate': {
+ 'type': 'stepwise',
+ 'stepwise': {
+ 'boundaries': [
+ steps_per_epoch * 320, steps_per_epoch * 340
+ ],
+ 'values': [0.32, 0.032, 0.0032],
+ }
+ },
+ 'warmup': {
+ 'type': 'linear',
+ 'linear': {
+ 'warmup_steps': 2000,
+ 'warmup_learning_rate': 0.0067
+ }
+ }
+ })),
+ restrictions=[
+ 'task.train_data.is_training != None',
+ 'task.validation_data.is_training != None',
+ ])
+ return config
diff --git a/official/vision/configs/maskrcnn_test.py b/official/vision/configs/maskrcnn_test.py
index 6f7af1f3473..e0a767d55c3 100644
--- a/official/vision/configs/maskrcnn_test.py
+++ b/official/vision/configs/maskrcnn_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for maskrcnn."""
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -39,7 +39,7 @@ def test_maskrcnn_configs(self, config_name):
self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
config.validate()
config.task.train_data.is_training = None
- with self.assertRaisesRegex(KeyError, 'Found inconsistncy between key'):
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
config.validate()
diff --git a/official/vision/configs/retinanet.py b/official/vision/configs/retinanet.py
index 426117d921e..0ceb9856c9c 100644
--- a/official/vision/configs/retinanet.py
+++ b/official/vision/configs/retinanet.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,12 +16,13 @@
import dataclasses
import os
-from typing import List, Optional, Union
+from typing import Optional, List, Sequence, Union
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.modeling import hyperparams
from official.modeling import optimization
+from official.modeling.hyperparams import base_config
from official.vision.configs import common
from official.vision.configs import decoders
from official.vision.configs import backbones
@@ -52,12 +53,15 @@ class Parser(hyperparams.Config):
match_threshold: float = 0.5
unmatched_threshold: float = 0.5
aug_rand_hflip: bool = False
+ aug_rand_jpeg: common.RandJpegQuality | None = None
aug_scale_min: float = 1.0
aug_scale_max: float = 1.0
skip_crowd_during_training: bool = True
max_num_instances: int = 100
# Can choose AutoAugment and RandAugment.
aug_type: Optional[common.Augmentation] = None
+ pad: bool = True
+ keep_aspect_ratio: bool = True
# Keep for backward compatibility. Not used.
aug_policy: Optional[str] = None
@@ -65,13 +69,21 @@ class Parser(hyperparams.Config):
@dataclasses.dataclass
class DataConfig(cfg.DataConfig):
- """Input config for training."""
- input_path: str = ''
+ """Input config for training.
+
+ Attributes:
+ weights: Sampling weights for each corresponding input_path. If used, then
+ input_path must be a config with matching keys.
+ """
+ input_path: Union[Sequence[str], str, base_config.Config] = ''
+ weights: Optional[base_config.Config] = None
global_batch_size: int = 0
is_training: bool = False
dtype: str = 'bfloat16'
- decoder: common.DataDecoder = common.DataDecoder()
- parser: Parser = Parser()
+ decoder: common.DataDecoder = dataclasses.field(
+ default_factory=common.DataDecoder
+ )
+ parser: Parser = dataclasses.field(default_factory=Parser)
shuffle_buffer_size: int = 10000
file_type: str = 'tfrecord'
@@ -99,6 +111,17 @@ class AttributeHead(hyperparams.Config):
name: str = ''
type: str = 'regression'
size: int = 1
+ # Attribute heads of the same "prediction_tower_name" will share the same
+ # prediction tower. If unspecified, they will use their individual prediction
+ # tower.
+ prediction_tower_name: str = ''
+ # If `num_convs` or `num_filters` are not provided, it will use the parameters
+ # from RetinaNetHead. When several attributes share the head through setting
+ # the same `prediction_tower_name`, we only respect `num_convs` and
+ # `num_filters` from the first attribute that use the shared prediction tower
+ # name.
+ num_convs: Optional[int] = None
+ num_filters: Optional[int] = None
@dataclasses.dataclass
@@ -107,11 +130,14 @@ class RetinaNetHead(hyperparams.Config):
num_filters: int = 256
use_separable_conv: bool = False
attribute_heads: List[AttributeHead] = dataclasses.field(default_factory=list)
+ share_classification_heads: bool = False
+ share_level_convs: Optional[bool] = True
@dataclasses.dataclass
class DetectionGenerator(hyperparams.Config):
apply_nms: bool = True
+ decode_boxes: bool = True
pre_nms_top_k: int = 5000
pre_nms_score_threshold: float = 0.05
nms_iou_threshold: float = 0.5
@@ -123,8 +149,16 @@ class DetectionGenerator(hyperparams.Config):
# When nms_version = `tflite`, values from tflite_post_processing need to be
# specified. They are compatible with the input arguments used by TFLite
# custom NMS op and override above parameters.
- tflite_post_processing: common.TFLitePostProcessingConfig = common.TFLitePostProcessingConfig(
+ tflite_post_processing: common.TFLitePostProcessingConfig = dataclasses.field(
+ default_factory=common.TFLitePostProcessingConfig
)
+ # Return decoded boxes/scores even if apply_nms is set `True`.
+ return_decoded: Optional[bool] = None
+ # Only works when nms_version='v2'.
+ use_class_agnostic_nms: Optional[bool] = False
+ # Weights or scales when encode and decode boxes coordinates. For Faster RCNN,
+ # the open-source implementation recommends using [10.0, 10.0, 5.0, 5.0].
+ box_coder_weights: Optional[List[float]] = None
@dataclasses.dataclass
@@ -133,14 +167,22 @@ class RetinaNet(hyperparams.Config):
input_size: List[int] = dataclasses.field(default_factory=list)
min_level: int = 3
max_level: int = 7
- anchor: Anchor = Anchor()
- backbone: backbones.Backbone = backbones.Backbone(
- type='resnet', resnet=backbones.ResNet())
- decoder: decoders.Decoder = decoders.Decoder(
- type='fpn', fpn=decoders.FPN())
- head: RetinaNetHead = RetinaNetHead()
- detection_generator: DetectionGenerator = DetectionGenerator()
- norm_activation: common.NormActivation = common.NormActivation()
+ anchor: Anchor = dataclasses.field(default_factory=Anchor)
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='resnet', resnet=backbones.ResNet()
+ )
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder(type='fpn', fpn=decoders.FPN())
+ )
+ head: RetinaNetHead = dataclasses.field(default_factory=RetinaNetHead)
+ detection_generator: DetectionGenerator = dataclasses.field(
+ default_factory=DetectionGenerator
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
@dataclasses.dataclass
@@ -148,25 +190,37 @@ class ExportConfig(hyperparams.Config):
output_normalized_coordinates: bool = False
cast_num_detections_to_float: bool = False
cast_detection_classes_to_float: bool = False
+ output_intermediate_features: bool = False
@dataclasses.dataclass
class RetinaNetTask(cfg.TaskConfig):
- model: RetinaNet = RetinaNet()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- losses: Losses = Losses()
+ model: RetinaNet = dataclasses.field(default_factory=RetinaNet)
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
init_checkpoint: Optional[str] = None
init_checkpoint_modules: Union[
str, List[str]] = 'all' # all, backbone, and/or decoder
annotation_file: Optional[str] = None
per_category_metrics: bool = False
- export_config: ExportConfig = ExportConfig()
+ export_config: ExportConfig = dataclasses.field(default_factory=ExportConfig)
# If set, the COCO metrics will be computed.
use_coco_metrics: bool = True
# If set, the Waymo Open Dataset evaluator would be used.
use_wod_metrics: bool = False
+ # If set, freezes the backbone during training.
+ # TODO(crisnv) Add paper link when available.
+ freeze_backbone: bool = False
+
+ # Sets maximum number of boxes to be evaluated by coco eval api.
+ max_num_eval_detections: int = 100
+
@exp_factory.register_config_factory('retinanet')
def retinanet() -> cfg.ExperimentConfig:
diff --git a/official/vision/configs/retinanet_test.py b/official/vision/configs/retinanet_test.py
index f9605e6dbd9..fee26999ccf 100644
--- a/official/vision/configs/retinanet_test.py
+++ b/official/vision/configs/retinanet_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for retinanet."""
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -38,7 +38,7 @@ def test_retinanet_configs(self, config_name):
self.assertIsInstance(config.task.train_data, exp_cfg.DataConfig)
config.validate()
config.task.train_data.is_training = None
- with self.assertRaisesRegex(KeyError, 'Found inconsistncy between key'):
+ with self.assertRaisesRegex(KeyError, 'Found inconsistency between key'):
config.validate()
diff --git a/official/vision/configs/semantic_segmentation.py b/official/vision/configs/semantic_segmentation.py
index 1449c133b35..ce171ea5556 100644
--- a/official/vision/configs/semantic_segmentation.py
+++ b/official/vision/configs/semantic_segmentation.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,10 @@
"""Semantic segmentation configuration definition."""
import dataclasses
+import math
import os
-from typing import List, Optional, Union
+from typing import List, Optional, Sequence, Union
-import numpy as np
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.modeling import hyperparams
@@ -25,16 +25,49 @@
from official.vision.configs import common
from official.vision.configs import decoders
from official.vision.configs import backbones
+from official.vision.ops import preprocess_ops
+
+
+@dataclasses.dataclass
+class DenseFeatureConfig(hyperparams.Config):
+ """Config for dense features, such as RGB pixels, masks, heatmaps.
+
+ The dense features are encoded images in TF examples. Thus they are
+ 1-, 3- or 4-channel. For features with another channel number (e.g.
+ optical flow), they could be encoded in multiple 1-channel features.
+ The default config is for RGB input, with mean and stddev from ImageNet
+ datasets. Only supports 8-bit encoded features with the maximum value = 255.
+
+ Attributes:
+ feature_name: The key of the feature in TF examples.
+ num_channels: An `int` specifying the number of channels of the feature.
+ mean: A list of floats in the range of [0, 255] representing the mean value
+ of each channel. The length of the list should match num_channels.
+ stddev: A list of floats in the range of [0, 255] representing the standard
+ deviation of each channel. The length should match num_channels.
+ """
+ feature_name: str = 'image/encoded'
+ num_channels: int = 3
+ mean: List[float] = dataclasses.field(
+ default_factory=lambda: preprocess_ops.MEAN_RGB
+ )
+ stddev: List[float] = dataclasses.field(
+ default_factory=lambda: preprocess_ops.STDDEV_RGB
+ )
@dataclasses.dataclass
class DataConfig(cfg.DataConfig):
"""Input config for training."""
+ image_feature: DenseFeatureConfig = dataclasses.field(
+ default_factory=DenseFeatureConfig
+ )
output_size: List[int] = dataclasses.field(default_factory=list)
# If crop_size is specified, image will be resized first to
# output_size, then crop of size crop_size will be cropped.
crop_size: List[int] = dataclasses.field(default_factory=list)
- input_path: str = ''
+ input_path: Union[Sequence[str], str, hyperparams.Config] = None
+ weights: Optional[hyperparams.Config] = None
global_batch_size: int = 0
is_training: bool = True
dtype: str = 'float32'
@@ -52,7 +85,15 @@ class DataConfig(cfg.DataConfig):
aug_policy: Optional[str] = None
drop_remainder: bool = True
file_type: str = 'tfrecord'
- decoder: Optional[common.DataDecoder] = common.DataDecoder()
+ decoder: Optional[common.DataDecoder] = dataclasses.field(
+ default_factory=common.DataDecoder
+ )
+ additional_dense_features: List[DenseFeatureConfig] = dataclasses.field(
+ default_factory=list)
+ # If `centered_crop` is set to True, then resized crop
+ # (if smaller than padded size) is place in the center of the image.
+ # Default behaviour is to place it at left top corner.
+ centered_crop: bool = False
@dataclasses.dataclass
@@ -64,6 +105,7 @@ class SegmentationHead(hyperparams.Config):
use_depthwise_convolution: bool = False
prediction_kernel_size: int = 1
upsample_factor: int = 1
+ logit_activation: Optional[str] = None # None, 'sigmoid', or 'softmax'.
feature_fusion: Optional[
str] = None # None, deeplabv3plus, panoptic_fpn_fusion or pyramid_fusion
# deeplabv3plus feature fusion params
@@ -82,6 +124,7 @@ class MaskScoringHead(hyperparams.Config):
fc_input_size: List[int] = dataclasses.field(default_factory=list)
num_fcs: int = 2
fc_dims: int = 1024
+ use_depthwise_convolution: bool = False
@dataclasses.dataclass
@@ -91,46 +134,74 @@ class SemanticSegmentationModel(hyperparams.Config):
input_size: List[int] = dataclasses.field(default_factory=list)
min_level: int = 3
max_level: int = 6
- head: SegmentationHead = SegmentationHead()
- backbone: backbones.Backbone = backbones.Backbone(
- type='resnet', resnet=backbones.ResNet())
- decoder: decoders.Decoder = decoders.Decoder(type='identity')
+ head: SegmentationHead = dataclasses.field(default_factory=SegmentationHead)
+ backbone: backbones.Backbone = dataclasses.field(
+ default_factory=lambda: backbones.Backbone( # pylint: disable=g-long-lambda
+ type='resnet', resnet=backbones.ResNet()
+ )
+ )
+ decoder: decoders.Decoder = dataclasses.field(
+ default_factory=lambda: decoders.Decoder(type='identity')
+ )
mask_scoring_head: Optional[MaskScoringHead] = None
- norm_activation: common.NormActivation = common.NormActivation()
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=common.NormActivation
+ )
@dataclasses.dataclass
class Losses(hyperparams.Config):
+ """Loss function config."""
loss_weight: float = 1.0
label_smoothing: float = 0.0
ignore_label: int = 255
+ gt_is_matting_map: bool = False
class_weights: List[float] = dataclasses.field(default_factory=list)
l2_weight_decay: float = 0.0
use_groundtruth_dimension: bool = True
+ # If true, use binary cross entropy (sigmoid) in loss, otherwise, use
+ # categorical cross entropy (softmax).
+ use_binary_cross_entropy: bool = False
top_k_percent_pixels: float = 1.0
+ mask_scoring_weight: float = 1.0
@dataclasses.dataclass
class Evaluation(hyperparams.Config):
+ """Evaluation config."""
report_per_class_iou: bool = True
report_train_mean_iou: bool = True # Turning this off can speed up training.
+@dataclasses.dataclass
+class ExportConfig(hyperparams.Config):
+ """Model export config."""
+ # Whether to rescale the predicted mask to the original image size.
+ rescale_output: bool = False
+
+
@dataclasses.dataclass
class SemanticSegmentationTask(cfg.TaskConfig):
"""The model config."""
- model: SemanticSegmentationModel = SemanticSegmentationModel()
- train_data: DataConfig = DataConfig(is_training=True)
- validation_data: DataConfig = DataConfig(is_training=False)
- losses: Losses = Losses()
- evaluation: Evaluation = Evaluation()
+ model: SemanticSegmentationModel = dataclasses.field(
+ default_factory=SemanticSegmentationModel
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=False)
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ evaluation: Evaluation = dataclasses.field(default_factory=Evaluation)
train_input_partition_dims: List[int] = dataclasses.field(
default_factory=list)
- eval_input_partition_dims: List[int] = dataclasses.field(
- default_factory=list)
+ eval_input_partition_dims: List[int] = dataclasses.field(default_factory=list)
init_checkpoint: Optional[str] = None
init_checkpoint_modules: Union[
str, List[str]] = 'all' # all, backbone, and/or decoder
+ export_config: ExportConfig = dataclasses.field(default_factory=ExportConfig)
+ allow_image_summary: bool = True
@exp_factory.register_config_factory('semantic_segmentation')
@@ -144,6 +215,7 @@ def semantic_segmentation() -> cfg.ExperimentConfig:
'task.validation_data.is_training != None'
])
+
# PASCAL VOC 2012 Dataset
PASCAL_TRAIN_EXAMPLES = 10582
PASCAL_VAL_EXAMPLES = 1449
@@ -160,18 +232,22 @@ def seg_deeplabv3_pascal() -> cfg.ExperimentConfig:
aspp_dilation_rates = [12, 24, 36] # [6, 12, 18] if output_stride = 16
multigrid = [1, 2, 4]
stem_type = 'v1'
- level = int(np.math.log2(output_stride))
+ level = int(math.log2(output_stride))
config = cfg.ExperimentConfig(
task=SemanticSegmentationTask(
model=SemanticSegmentationModel(
num_classes=21,
- input_size=[None, None, 3],
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
backbone=backbones.Backbone(
- type='dilated_resnet', dilated_resnet=backbones.DilatedResNet(
- model_id=101, output_stride=output_stride,
- multigrid=multigrid, stem_type=stem_type)),
+ type='dilated_resnet',
+ dilated_resnet=backbones.DilatedResNet(
+ model_id=101,
+ output_stride=output_stride,
+ multigrid=multigrid,
+ stem_type=stem_type)),
decoder=decoders.Decoder(
- type='aspp', aspp=decoders.ASPP(
+ type='aspp',
+ aspp=decoders.ASPP(
level=level, dilation_rates=aspp_dilation_rates)),
head=SegmentationHead(level=level, num_convs=0),
norm_activation=common.NormActivation(
@@ -248,16 +324,19 @@ def seg_deeplabv3plus_pascal() -> cfg.ExperimentConfig:
aspp_dilation_rates = [6, 12, 18]
multigrid = [1, 2, 4]
stem_type = 'v1'
- level = int(np.math.log2(output_stride))
+ level = int(math.log2(output_stride))
config = cfg.ExperimentConfig(
task=SemanticSegmentationTask(
model=SemanticSegmentationModel(
num_classes=21,
- input_size=[None, None, 3],
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
backbone=backbones.Backbone(
- type='dilated_resnet', dilated_resnet=backbones.DilatedResNet(
- model_id=101, output_stride=output_stride,
- stem_type=stem_type, multigrid=multigrid)),
+ type='dilated_resnet',
+ dilated_resnet=backbones.DilatedResNet(
+ model_id=101,
+ output_stride=output_stride,
+ stem_type=stem_type,
+ multigrid=multigrid)),
decoder=decoders.Decoder(
type='aspp',
aspp=decoders.ASPP(
@@ -349,8 +428,7 @@ def seg_resnetfpn_pascal() -> cfg.ExperimentConfig:
decoder=decoders.Decoder(type='fpn', fpn=decoders.FPN()),
head=SegmentationHead(level=3, num_convs=3),
norm_activation=common.NormActivation(
- activation='swish',
- use_sync_bn=True)),
+ activation='swish', use_sync_bn=True)),
losses=Losses(l2_weight_decay=1e-4),
train_data=DataConfig(
input_path=os.path.join(PASCAL_INPUT_PATH_BASE, 'train_aug*'),
@@ -413,14 +491,14 @@ def mnv2_deeplabv3_pascal() -> cfg.ExperimentConfig:
steps_per_epoch = PASCAL_TRAIN_EXAMPLES // train_batch_size
output_stride = 16
aspp_dilation_rates = []
- level = int(np.math.log2(output_stride))
+ level = int(math.log2(output_stride))
pool_kernel_size = []
config = cfg.ExperimentConfig(
task=SemanticSegmentationTask(
model=SemanticSegmentationModel(
num_classes=21,
- input_size=[None, None, 3],
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
backbone=backbones.Backbone(
type='mobilenet',
mobilenet=backbones.MobileNet(
@@ -514,22 +592,26 @@ def seg_deeplabv3plus_cityscapes() -> cfg.ExperimentConfig:
aspp_dilation_rates = [6, 12, 18]
multigrid = [1, 2, 4]
stem_type = 'v1'
- level = int(np.math.log2(output_stride))
+ level = int(math.log2(output_stride))
config = cfg.ExperimentConfig(
task=SemanticSegmentationTask(
model=SemanticSegmentationModel(
# Cityscapes uses only 19 semantic classes for train/evaluation.
# The void (background) class is ignored in train and evaluation.
num_classes=19,
- input_size=[None, None, 3],
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
backbone=backbones.Backbone(
- type='dilated_resnet', dilated_resnet=backbones.DilatedResNet(
- model_id=101, output_stride=output_stride,
- stem_type=stem_type, multigrid=multigrid)),
+ type='dilated_resnet',
+ dilated_resnet=backbones.DilatedResNet(
+ model_id=101,
+ output_stride=output_stride,
+ stem_type=stem_type,
+ multigrid=multigrid)),
decoder=decoders.Decoder(
type='aspp',
aspp=decoders.ASPP(
- level=level, dilation_rates=aspp_dilation_rates,
+ level=level,
+ dilation_rates=aspp_dilation_rates,
pool_kernel_size=[512, 1024])),
head=SegmentationHead(
level=level,
@@ -611,14 +693,14 @@ def mnv2_deeplabv3_cityscapes() -> cfg.ExperimentConfig:
aspp_dilation_rates = []
pool_kernel_size = [512, 1024]
- level = int(np.math.log2(output_stride))
+ level = int(math.log2(output_stride))
config = cfg.ExperimentConfig(
task=SemanticSegmentationTask(
model=SemanticSegmentationModel(
# Cityscapes uses only 19 semantic classes for train/evaluation.
# The void (background) class is ignored in train and evaluation.
num_classes=19,
- input_size=[None, None, 3],
+ input_size=[None, None, 3], # pyrefly: ignore[bad-argument-type]
backbone=backbones.Backbone(
type='mobilenet',
mobilenet=backbones.MobileNet(
diff --git a/official/vision/configs/semantic_segmentation_test.py b/official/vision/configs/semantic_segmentation_test.py
index fb25dbdc899..1884e597762 100644
--- a/official/vision/configs/semantic_segmentation_test.py
+++ b/official/vision/configs/semantic_segmentation_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -26,8 +26,14 @@
class ImageSegmentationConfigTest(tf.test.TestCase, parameterized.TestCase):
- @parameterized.parameters(('seg_deeplabv3_pascal',),
- ('seg_deeplabv3plus_pascal',))
+ @parameterized.parameters(
+ ('seg_deeplabv3_pascal',),
+ ('seg_deeplabv3plus_pascal',),
+ ('mnv2_deeplabv3plus_cityscapes',),
+ ('mnv2_deeplabv3_cityscapes',),
+ ('mnv2_deeplabv3_pascal',),
+ ('seg_resnetfpn_pascal',),
+ )
def test_semantic_segmentation_configs(self, config_name):
config = exp_factory.get_exp_config(config_name)
self.assertIsInstance(config, cfg.ExperimentConfig)
diff --git a/official/vision/configs/video_classification.py b/official/vision/configs/video_classification.py
index d896f0f5622..73e910cb176 100644
--- a/official/vision/configs/video_classification.py
+++ b/official/vision/configs/video_classification.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""Video classification configuration definition."""
import dataclasses
-from typing import Optional, Tuple
+from typing import Optional, Tuple, Union
from official.core import config_definitions as cfg
from official.core import exp_factory
from official.modeling import hyperparams
@@ -41,14 +41,16 @@ class DataConfig(cfg.DataConfig):
global_batch_size: int = 128
data_format: str = 'channels_last'
dtype: str = 'float32'
+ label_dtype: str = 'int32'
one_hot: bool = True
shuffle_buffer_size: int = 64
cache: bool = False
- input_path: str = ''
+ input_path: Union[str, cfg.base_config.Config] = ''
is_training: bool = True
cycle_length: int = 10
drop_remainder: bool = True
min_image_size: int = 256
+ zero_centering_image: bool = False
is_multilabel: bool = False
output_audio: bool = False
audio_feature: str = ''
@@ -57,13 +59,17 @@ class DataConfig(cfg.DataConfig):
aug_max_aspect_ratio: float = 2.0
aug_min_area_ratio: float = 0.49
aug_max_area_ratio: float = 1.0
- aug_type: Optional[str] = None # 'autoaug', 'randaug', or None
+ aug_random_rotation: bool = False
+ aug_type: Optional[
+ common.Augmentation] = None # AutoAugment and RandAugment.
+ mixup_and_cutmix: Optional[common.MixupAndCutmix] = None
image_field_key: str = 'image/encoded'
label_field_key: str = 'clip/label/index'
+ input_image_format: str = 'jpeg'
def kinetics400(is_training):
- """Generated Kinectics 400 dataset configs."""
+ """Generated Kinetics 400 dataset configs."""
return DataConfig(
name='kinetics400',
num_classes=400,
@@ -75,7 +81,7 @@ def kinetics400(is_training):
def kinetics600(is_training):
- """Generated Kinectics 600 dataset configs."""
+ """Generated Kinetics 600 dataset configs."""
return DataConfig(
name='kinetics600',
num_classes=600,
@@ -87,7 +93,7 @@ def kinetics600(is_training):
def kinetics700(is_training):
- """Generated Kinectics 600 dataset configs."""
+ """Generated Kinetics 600 dataset configs."""
return DataConfig(
name='kinetics700',
num_classes=700,
@@ -99,7 +105,7 @@ def kinetics700(is_training):
def kinetics700_2020(is_training):
- """Generated Kinectics 600 dataset configs."""
+ """Generated Kinetics 600 dataset configs."""
return DataConfig(
name='kinetics700',
num_classes=700,
@@ -114,10 +120,14 @@ def kinetics700_2020(is_training):
class VideoClassificationModel(hyperparams.Config):
"""The model config."""
model_type: str = 'video_classification'
- backbone: backbones_3d.Backbone3D = backbones_3d.Backbone3D(
- type='resnet_3d', resnet_3d=backbones_3d.ResNet3D50())
- norm_activation: common.NormActivation = common.NormActivation(
- use_sync_bn=False)
+ backbone: backbones_3d.Backbone3D = dataclasses.field(
+ default_factory=lambda: backbones_3d.Backbone3D( # pylint: disable=g-long-lambda
+ type='resnet_3d', resnet_3d=backbones_3d.ResNet3D50()
+ )
+ )
+ norm_activation: common.NormActivation = dataclasses.field(
+ default_factory=lambda: common.NormActivation(use_sync_bn=False)
+ )
dropout_rate: float = 0.2
aggregate_endpoints: bool = False
require_endpoints: Optional[Tuple[str, ...]] = None
@@ -138,14 +148,22 @@ class Metrics(hyperparams.Config):
@dataclasses.dataclass
class VideoClassificationTask(cfg.TaskConfig):
"""The task config."""
- model: VideoClassificationModel = VideoClassificationModel()
- train_data: DataConfig = DataConfig(is_training=True, drop_remainder=True)
- validation_data: DataConfig = DataConfig(
- is_training=False, drop_remainder=False)
- losses: Losses = Losses()
- metrics: Metrics = Metrics()
+ model: VideoClassificationModel = dataclasses.field(
+ default_factory=VideoClassificationModel
+ )
+ train_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig(is_training=True, drop_remainder=True)
+ )
+ validation_data: DataConfig = dataclasses.field(
+ default_factory=lambda: DataConfig( # pylint: disable=g-long-lambda
+ is_training=False, drop_remainder=False
+ )
+ )
+ losses: Losses = dataclasses.field(default_factory=Losses)
+ metrics: Metrics = dataclasses.field(default_factory=Metrics)
init_checkpoint: Optional[str] = None
init_checkpoint_modules: str = 'all' # all or backbone
+ freeze_backbone: bool = False
# Spatial Partitioning fields.
train_input_partition_dims: Optional[Tuple[int, ...]] = None
eval_input_partition_dims: Optional[Tuple[int, ...]] = None
@@ -268,7 +286,7 @@ def video_classification_ucf101() -> cfg.ExperimentConfig:
@exp_factory.register_config_factory('video_classification_kinetics400')
def video_classification_kinetics400() -> cfg.ExperimentConfig:
- """Video classification on Kinectics 400 with resnet."""
+ """Video classification on Kinetics 400 with resnet."""
train_dataset = kinetics400(is_training=True)
validation_dataset = kinetics400(is_training=False)
task = VideoClassificationTask(
@@ -294,7 +312,7 @@ def video_classification_kinetics400() -> cfg.ExperimentConfig:
@exp_factory.register_config_factory('video_classification_kinetics600')
def video_classification_kinetics600() -> cfg.ExperimentConfig:
- """Video classification on Kinectics 600 with resnet."""
+ """Video classification on Kinetics 600 with resnet."""
train_dataset = kinetics600(is_training=True)
validation_dataset = kinetics600(is_training=False)
task = VideoClassificationTask(
@@ -320,7 +338,7 @@ def video_classification_kinetics600() -> cfg.ExperimentConfig:
@exp_factory.register_config_factory('video_classification_kinetics700')
def video_classification_kinetics700() -> cfg.ExperimentConfig:
- """Video classification on Kinectics 700 with resnet."""
+ """Video classification on Kinetics 700 with resnet."""
train_dataset = kinetics700(is_training=True)
validation_dataset = kinetics700(is_training=False)
task = VideoClassificationTask(
@@ -346,7 +364,7 @@ def video_classification_kinetics700() -> cfg.ExperimentConfig:
@exp_factory.register_config_factory('video_classification_kinetics700_2020')
def video_classification_kinetics700_2020() -> cfg.ExperimentConfig:
- """Video classification on Kinectics 700 2020 with resnet."""
+ """Video classification on Kinetics 700 2020 with resnet."""
train_dataset = kinetics700_2020(is_training=True)
validation_dataset = kinetics700_2020(is_training=False)
task = VideoClassificationTask(
diff --git a/official/vision/configs/video_classification_test.py b/official/vision/configs/video_classification_test.py
index 6c8053e14e3..514bb4cfe44 100644
--- a/official/vision/configs/video_classification_test.py
+++ b/official/vision/configs/video_classification_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -16,7 +16,7 @@
# pylint: disable=unused-import
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official import vision
from official.core import config_definitions as cfg
@@ -26,8 +26,14 @@
class VideoClassificationConfigTest(tf.test.TestCase, parameterized.TestCase):
- @parameterized.parameters(('video_classification',),
- ('video_classification_kinetics600',))
+ @parameterized.parameters(
+ ('video_classification',),
+ ('video_classification_ucf101',),
+ ('video_classification_kinetics400',),
+ ('video_classification_kinetics600',),
+ ('video_classification_kinetics700',),
+ ('video_classification_kinetics700_2020',),
+ )
def test_video_classification_configs(self, config_name):
config = exp_factory.get_exp_config(config_name)
self.assertIsInstance(config, cfg.ExperimentConfig)
diff --git a/official/vision/data/__init__.py b/official/vision/data/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/vision/data/__init__.py
+++ b/official/vision/data/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/data/create_coco_tf_record.py b/official/vision/data/create_coco_tf_record.py
index 27848d4e899..a552bc1e2f6 100644
--- a/official/vision/data/create_coco_tf_record.py
+++ b/official/vision/data/create_coco_tf_record.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -37,7 +37,7 @@
import numpy as np
from pycocotools import mask
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import multiprocessing as mp
from official.vision.data import tfrecord_lib
@@ -66,6 +66,9 @@
'include_panoptic_masks', False, 'Whether to include category and '
'instance masks in the result. These are required to run the PQ evaluator '
'default: False.')
+flags.DEFINE_boolean(
+ 'panoptic_skip_crowd', False, 'Whether to skip crowd or not for panoptic '
+ 'annotations. default: False.')
flags.DEFINE_string('output_file_prefix', '/tmp/train', 'Path to output file')
flags.DEFINE_integer('num_shards', 32, 'Number of shards for output file.')
_NUM_PROCESSES = flags.DEFINE_integer(
@@ -112,7 +115,7 @@ def generate_coco_panoptics_masks(segments_info, mask_path,
represent "stuff" and "things" classes respectively.
Returns:
- A dict with with keys: [u'semantic_segmentation_mask', u'category_mask',
+ A dict with keys: [u'semantic_segmentation_mask', u'category_mask',
u'instance_mask']. The dict contains 'category_mask' and 'instance_mask'
only if `include_panoptic_eval_masks` is set to True.
"""
@@ -134,7 +137,9 @@ def generate_coco_panoptics_masks(segments_info, mask_path,
for idx, segment in enumerate(segments_info):
segment_id = segment['id']
category_id = segment['category_id']
-
+ is_crowd = segment['iscrowd']
+ if FLAGS.panoptic_skip_crowd and is_crowd:
+ continue
if is_category_thing[category_id]:
encoded_category_id = _THING_CLASS_ID
instance_id = idx + 1
@@ -146,8 +151,8 @@ def generate_coco_panoptics_masks(segments_info, mask_path,
semantic_segmentation_mask[segment_mask] = encoded_category_id
if include_panoptic_masks:
- category_mask[segment_mask] = category_id
- instance_mask[segment_mask] = instance_id
+ category_mask[segment_mask] = category_id # pyrefly: ignore[unbound-name]
+ instance_mask[segment_mask] = instance_id # pyrefly: ignore[unbound-name]
outputs = {
'semantic_segmentation_mask': tfrecord_lib.encode_mask_as_png(
@@ -156,8 +161,8 @@ def generate_coco_panoptics_masks(segments_info, mask_path,
if include_panoptic_masks:
outputs.update({
- 'category_mask': tfrecord_lib.encode_mask_as_png(category_mask),
- 'instance_mask': tfrecord_lib.encode_mask_as_png(instance_mask)
+ 'category_mask': tfrecord_lib.encode_mask_as_png(category_mask), # pyrefly: ignore[unbound-name]
+ 'instance_mask': tfrecord_lib.encode_mask_as_png(instance_mask) # pyrefly: ignore[unbound-name]
})
return outputs
@@ -210,27 +215,30 @@ def bbox_annotations_to_feature_dict(
data, num_skipped = coco_annotations_to_lists(
bbox_annotations, id_to_name_map, image_height, image_width,
include_masks)
- feature_dict = {
- 'image/object/bbox/xmin':
- tfrecord_lib.convert_to_feature(data['xmin']),
- 'image/object/bbox/xmax':
- tfrecord_lib.convert_to_feature(data['xmax']),
- 'image/object/bbox/ymin':
- tfrecord_lib.convert_to_feature(data['ymin']),
- 'image/object/bbox/ymax':
- tfrecord_lib.convert_to_feature(data['ymax']),
- 'image/object/class/text':
- tfrecord_lib.convert_to_feature(data['category_names']),
- 'image/object/class/label':
- tfrecord_lib.convert_to_feature(data['category_id']),
- 'image/object/is_crowd':
- tfrecord_lib.convert_to_feature(data['is_crowd']),
- 'image/object/area':
- tfrecord_lib.convert_to_feature(data['area']),
- }
- if include_masks:
- feature_dict['image/object/mask'] = (
- tfrecord_lib.convert_to_feature(data['encoded_mask_png']))
+ feature_dict = {}
+ if len(bbox_annotations) != num_skipped:
+ feature_dict = {
+ 'image/object/bbox/xmin': tfrecord_lib.convert_to_feature(data['xmin']),
+ 'image/object/bbox/xmax': tfrecord_lib.convert_to_feature(data['xmax']),
+ 'image/object/bbox/ymin': tfrecord_lib.convert_to_feature(data['ymin']),
+ 'image/object/bbox/ymax': tfrecord_lib.convert_to_feature(data['ymax']),
+ 'image/object/class/text': tfrecord_lib.convert_to_feature(
+ data['category_names']
+ ),
+ 'image/object/class/label': tfrecord_lib.convert_to_feature(
+ data['category_id']
+ ),
+ 'image/object/is_crowd': tfrecord_lib.convert_to_feature(
+ data['is_crowd']
+ ),
+ 'image/object/area': tfrecord_lib.convert_to_feature(
+ data['area'], 'float_list'
+ ),
+ }
+ if include_masks:
+ feature_dict['image/object/mask'] = tfrecord_lib.convert_to_feature(
+ data['encoded_mask_png']
+ )
return feature_dict, num_skipped
@@ -331,7 +339,7 @@ def create_tf_example(image,
if panoptic_annotation:
segments_info = panoptic_annotation['segments_info']
- panoptic_mask_filename = os.path.join(
+ panoptic_mask_filename = os.path.join( # pyrefly: ignore[no-matching-overload]
panoptic_masks_dir,
panoptic_annotation['file_name'])
encoded_panoptic_masks = generate_coco_panoptics_masks(
diff --git a/official/vision/data/fake_feature_generator.py b/official/vision/data/fake_feature_generator.py
new file mode 100644
index 00000000000..71faafa0ada
--- /dev/null
+++ b/official/vision/data/fake_feature_generator.py
@@ -0,0 +1,128 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Generates fake feature for testing and validation."""
+
+import collections
+from typing import Optional, Tuple, Union
+
+import numpy as np
+
+_RGB_CHANNELS = 3
+
+
+def generate_image_np(height: int,
+ width: int,
+ num_channels: int = _RGB_CHANNELS) -> np.ndarray:
+ """Returns a fake numpy image matrix array."""
+ return np.reshape(
+ np.mod(np.arange(height * width * num_channels), 255).astype(np.uint8),
+ (height, width, num_channels),
+ )
+
+
+def generate_normalized_boxes_np(num_boxes: int) -> np.ndarray:
+ """Returns a fake numpy normalized boxes array."""
+ xmins = np.reshape(np.arange(num_boxes) / (2 * num_boxes), (num_boxes, 1))
+ ymins = np.reshape(np.arange(num_boxes) / (2 * num_boxes), (num_boxes, 1))
+ xmaxs = xmins + .5
+ ymaxs = ymins + .5
+ return np.concatenate((ymins, xmins, ymaxs, xmaxs), axis=-1)
+
+
+def generate_boxes_np(height: int, width: int, num_boxes: int) -> np.ndarray:
+ """Returns a fake numpy absolute boxes array."""
+ normalized_boxes = generate_normalized_boxes_np(num_boxes)
+ normalized_boxes[:, 1::2] *= height
+ normalized_boxes[:, 0::2] *= width
+ return normalized_boxes
+
+
+def generate_classes_np(num_classes: int,
+ size: Optional[int] = None) -> Union[int, np.ndarray]:
+ """Returns a fake class or a fake numpy classes array."""
+ if size is None:
+ return num_classes - 1
+
+ return np.arange(size) % num_classes
+
+
+def generate_confidences_np(
+ size: Optional[int] = None) -> Union[float, np.ndarray]:
+ """Returns a fake confidence score or a fake numpy confidence score array."""
+ if size is None:
+ return 0.5
+
+ return np.arange(size) / size
+
+
+def generate_instance_masks_np(height: int,
+ width: int,
+ boxes_np: np.ndarray,
+ normalized: bool = True) -> np.ndarray:
+ """Returns a fake numpy instance mask matrices array."""
+ num_boxes = len(boxes_np)
+ instance_masks_np = np.zeros((num_boxes, height, width, 1))
+ if normalized:
+ boxes_np[:, 1::2] *= height
+ boxes_np[:, ::2] *= width
+ xmins = boxes_np[:, 0].astype(int)
+ ymins = boxes_np[:, 1].astype(int)
+ box_widths = boxes_np[:, 2].astype(int) - xmins
+ box_heights = boxes_np[:, 3].astype(int) - ymins
+
+ for i, (x, y, w, h) in enumerate(zip(xmins, ymins, box_widths, box_heights)):
+ instance_masks_np[i, y:y + h, x:x + w, :] = np.reshape(np.mod(np.arange(h * w), 2).astype(np.uint8), (h, w, 1))
+ return instance_masks_np
+
+
+def generate_semantic_mask_np(height: int, width: int,
+ num_classes: int) -> np.ndarray:
+ """Returns a fake numpy semantic mask array."""
+ out = generate_image_np(height, width, num_channels=1)
+ if np.iinfo(out.dtype).max > num_classes:
+ out = out % num_classes
+ return out
+
+
+def generate_panoptic_masks_np(
+ semantic_mask: np.ndarray, instance_masks: np.ndarray,
+ instance_classes: np.ndarray,
+ stuff_classes_offset: int) -> Tuple[np.ndarray, np.ndarray]:
+ """Returns fake numpy panoptic category and instance mask arrays."""
+ panoptic_category_mask = np.zeros_like(semantic_mask)
+ panoptic_instance_mask = np.zeros_like(semantic_mask)
+ instance_ids = collections.defaultdict(int)
+ for instance_mask, instance_class in zip(instance_masks, instance_classes):
+ if instance_class == 0:
+ continue
+ instance_ids[instance_class] += 1
+ # If a foreground pixel is labelled previously, replace the old category
+ # class and instance ID with the new one.
+ foreground_indices = np.where(np.equal(instance_mask, 1))
+ # Note that instance class start from index 1.
+ panoptic_category_mask[foreground_indices] = instance_class + 1
+ panoptic_instance_mask[foreground_indices] = instance_ids[instance_class]
+
+ # If there are pixels remains unlablled (labelled as background), then the
+ # semantic labels will be used (if it has one).
+ # Note that in panoptic FPN, the panoptic labels are expected in this order,
+ # 0 (background), 1 ..., N (stuffs), N + 1, ..., N + M - 2 (things)
+ # N classes for stuff classes, without background class, and M classes for
+ # thing classes, with 0 representing the background class and 1 representing
+ # all stuff classes.
+ background_indices = np.where(np.equal(panoptic_category_mask, 0))
+ panoptic_category_mask[background_indices] = (
+ semantic_mask[background_indices] + stuff_classes_offset)
+ return panoptic_category_mask, panoptic_instance_mask
diff --git a/official/vision/data/image_utils.py b/official/vision/data/image_utils.py
new file mode 100644
index 00000000000..ab8dcd358b0
--- /dev/null
+++ b/official/vision/data/image_utils.py
@@ -0,0 +1,113 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Image-related utilities that are useful to prepare dataset."""
+
+import dataclasses
+import imghdr
+import io
+from typing import Optional, Tuple
+
+import numpy as np
+from PIL import Image
+
+
+@dataclasses.dataclass
+class ImageFormat:
+ """Supported image formats.
+
+ For model development, this library should support the same image formats as
+ `tf.io.decode_image`[1].
+
+ [1]: https://www.tensorflow.org/api_docs/python/tf/io/decode_image
+ """
+ bmp: str = 'BMP'
+ png: str = 'PNG'
+ jpeg: str = 'JPEG'
+ raw: str = 'RAW'
+
+
+def validate_image_format(format_str: str) -> str:
+ """Validates `format_str` and returns canonical format.
+
+ This function accepts image format in lower case and will returns the upper
+ case string as canonical format.
+
+ Args:
+ format_str: Image format string.
+
+ Returns:
+ Canonical image format string.
+
+ Raises:
+ ValueError: If the canonical format is not listed in `ImageFormat`.
+ """
+ canonical_format = format_str.upper()
+ if canonical_format in dataclasses.asdict(ImageFormat()).values():
+ return canonical_format
+ raise ValueError(f'Image format is invalid: {format_str}')
+
+
+def encode_image(image_np: np.ndarray, image_format: str) -> bytes:
+ """Encodes `image_np` specified by `image_format`.
+
+ Args:
+ image_np: Numpy image array.
+ image_format: An enum specifying the format of the generated image.
+
+ Returns:
+ Encoded image string.
+ """
+ if image_format == 'RAW':
+ return image_np.tobytes()
+
+ if len(image_np.shape) > 2 and image_np.shape[2] == 1:
+ image_pil = Image.fromarray(np.squeeze(image_np), 'L')
+ else:
+ image_pil = Image.fromarray(image_np)
+ with io.BytesIO() as output:
+ image_pil.save(output, format=validate_image_format(image_format))
+ return output.getvalue()
+
+
+def decode_image(image_bytes: bytes,
+ image_format: Optional[str] = None,
+ image_dtype: str = 'uint8') -> np.ndarray:
+ """Decodes image_bytes into numpy array."""
+ if image_format == 'RAW':
+ return np.frombuffer(image_bytes, dtype=image_dtype)
+ image_pil = Image.open(io.BytesIO(image_bytes))
+ image_np = np.array(image_pil)
+ if len(image_np.shape) < 3:
+ image_np = image_np[..., np.newaxis]
+ return image_np
+
+
+def decode_image_metadata(image_bytes: bytes) -> Tuple[int, int, int, str]:
+ """Decodes image metadata from encoded image string.
+
+ Note that if the image is encoded in RAW format, the metadata cannot be
+ inferred from the image bytes.
+
+ Args:
+ image_bytes: Encoded image string.
+
+ Returns:
+ A tuple of height, width, number of channels, and encoding format.
+ """
+ image_np = decode_image(image_bytes)
+ # https://pillow.readthedocs.io/en/stable/reference/Image.html#image-attributes
+ height, width, num_channels = image_np.shape
+ image_format = imghdr.what(file=None, h=image_bytes)
+ return height, width, num_channels, validate_image_format(image_format)
diff --git a/official/vision/data/image_utils_test.py b/official/vision/data/image_utils_test.py
new file mode 100644
index 00000000000..a15f8cf4fd6
--- /dev/null
+++ b/official/vision/data/image_utils_test.py
@@ -0,0 +1,87 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for image_utils."""
+import imghdr
+from unittest import mock
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+
+from official.vision.data import fake_feature_generator
+from official.vision.data import image_utils
+
+
+class ImageUtilsTest(parameterized.TestCase, tf.test.TestCase):
+
+ @parameterized.named_parameters(
+ ('RGB_PNG', 128, 64, 3, 'PNG'), ('RGB_JPEG', 2, 1, 3, 'JPEG'),
+ ('GREY_BMP', 32, 32, 1, 'BMP'), ('GREY_PNG', 128, 128, 1, 'png'))
+ def test_encode_image_then_decode_image(self, height, width, num_channels,
+ image_format):
+ image_np = fake_feature_generator.generate_image_np(height, width,
+ num_channels)
+ image_str = image_utils.encode_image(image_np, image_format)
+ actual_image_np = image_utils.decode_image(image_str)
+
+ # JPEG encoding does not keep the pixel value.
+ if image_format != 'JPEG':
+ self.assertAllClose(actual_image_np, image_np)
+ self.assertEqual(actual_image_np.shape, image_np.shape)
+
+ @parameterized.named_parameters(
+ ('RGB_RAW', 128, 64, 3, tf.bfloat16.as_numpy_dtype),
+ ('GREY_RAW', 32, 32, 1, tf.uint8.as_numpy_dtype))
+ def test_encode_raw_image_then_decode_raw_image(self, height, width,
+ num_channels, image_dtype):
+ image_np = fake_feature_generator.generate_image_np(height, width,
+ num_channels)
+ image_np = image_np.astype(image_dtype)
+ image_str = image_utils.encode_image(image_np, 'RAW')
+ actual_image_np = image_utils.decode_image(image_str, 'RAW', image_dtype)
+ actual_image_np = actual_image_np.reshape([height, width, num_channels])
+
+ self.assertAllClose(actual_image_np, image_np)
+ self.assertEqual(actual_image_np.shape, image_np.shape)
+
+ @parameterized.named_parameters(
+ ('RGB_PNG', 128, 64, 3, 'PNG'), ('RGB_JPEG', 64, 128, 3, 'JPEG'),
+ ('GREY_BMP', 32, 32, 1, 'BMP'), ('GREY_PNG', 128, 128, 1, 'png'))
+ def test_encode_image_then_decode_image_metadata(self, height, width,
+ num_channels, image_format):
+ image_np = fake_feature_generator.generate_image_np(height, width,
+ num_channels)
+ image_str = image_utils.encode_image(image_np, image_format)
+ (actual_height, actual_width, actual_num_channels, actual_format) = (
+ image_utils.decode_image_metadata(image_str))
+
+ self.assertEqual(actual_height, height)
+ self.assertEqual(actual_width, width)
+ self.assertEqual(actual_num_channels, num_channels)
+ self.assertEqual(actual_format, image_format.upper())
+
+ def test_encode_image_raise_error_with_invalid_image_format(self):
+ with self.assertRaisesRegex(ValueError, 'Image format is invalid: foo'):
+ image_np = fake_feature_generator.generate_image_np(2, 2, 1)
+ image_utils.encode_image(image_np, 'foo')
+
+ @mock.patch.object(imghdr, 'what', return_value='foo', autospec=True)
+ def test_decode_image_raise_error_with_invalid_image_format(self, _):
+ image_np = fake_feature_generator.generate_image_np(1, 1, 3)
+ image_str = image_utils.encode_image(image_np, 'PNG')
+ with self.assertRaisesRegex(ValueError, 'Image format is invalid: foo'):
+ image_utils.decode_image_metadata(image_str)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/vision/data/process_coco_few_shot_json_files.py b/official/vision/data/process_coco_few_shot_json_files.py
index 7a918c5117d..e0c8783a2f4 100644
--- a/official/vision/data/process_coco_few_shot_json_files.py
+++ b/official/vision/data/process_coco_few_shot_json_files.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -29,7 +29,7 @@
from absl import app
from absl import flags
-import tensorflow as tf
+import tensorflow as tf, tf_keras
logger = tf.get_logger()
logger.setLevel(logging.INFO)
diff --git a/official/vision/data/process_coco_panoptic.sh b/official/vision/data/process_coco_panoptic.sh
index fd70039741e..8bc9210e8c8 100644
--- a/official/vision/data/process_coco_panoptic.sh
+++ b/official/vision/data/process_coco_panoptic.sh
@@ -16,7 +16,7 @@ mkdir $DATA_DIR/zips && mv $DATA_DIR/*.zip $DATA_DIR/zips;
unzip $DATA_DIR/annotations/panoptic_train2017.zip -d $DATA_DIR
unzip $DATA_DIR/annotations/panoptic_val2017.zip -d $DATA_DIR
-python3 official/vision/beta/data/create_coco_tf_record.py \
+python3 official/vision/data/create_coco_tf_record.py \
--logtostderr \
--image_dir="$DATA_DIR/val2017" \
--object_annotations_file="$DATA_DIR/annotations/instances_val2017.json" \
@@ -28,7 +28,7 @@ python3 official/vision/beta/data/create_coco_tf_record.py \
--include_panoptic_masks
-python3 official/vision/beta/data/create_coco_tf_record.py \
+python3 official/vision/data/create_coco_tf_record.py \
--logtostderr \
--image_dir="$DATA_DIR/train2017" \
--object_annotations_file="$DATA_DIR/annotations/instances_train2017.json" \
diff --git a/official/vision/data/tf_example_builder.py b/official/vision/data/tf_example_builder.py
new file mode 100644
index 00000000000..0745fa8baa1
--- /dev/null
+++ b/official/vision/data/tf_example_builder.py
@@ -0,0 +1,492 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Builder class for preparing tf.train.Example in vision tasks."""
+
+# https://www.python.org/dev/peps/pep-0563/#enabling-the-future-behavior-in-python-3-7
+from __future__ import annotations
+
+import hashlib
+from typing import Optional, Sequence, Union
+import numpy as np
+
+from official.core import tf_example_builder
+from official.vision.data import image_utils
+from official.vision.data import tf_example_feature_key
+
+BytesValueType = Union[bytes, Sequence[bytes], str, Sequence[str]]
+
+_to_array = lambda v: [v] if not isinstance(v, (list, np.ndarray)) else v
+_to_bytes = lambda v: v.encode() if isinstance(v, str) else v
+_to_bytes_array = lambda v: list(map(_to_bytes, _to_array(v)))
+
+
+class TfExampleBuilder(tf_example_builder.TfExampleBuilder):
+ """Builder class for preparing tf.train.Example in vision task.
+
+ Read API doc at https://www.tensorflow.org/api_docs/python/tf/train/Example.
+ """
+
+ def add_image_matrix_feature(
+ self,
+ image_matrix: np.ndarray,
+ image_format: str = 'PNG',
+ image_source_id: Optional[bytes] = None,
+ feature_prefix: Optional[str] = None,
+ label: Optional[Union[int, Sequence[int]]] = None) -> 'TfExampleBuilder':
+ """Encodes and adds image features to the example.
+
+ See `tf_example_feature_key.EncodedImageFeatureKey` for list of feature keys
+ that will be added to the example.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ * For adding RGB image feature with PNG encoding:
+ >>> example_builder.add_image_matrix_feature(image_matrix)
+ * For adding RGB image feature with a pre-generated source ID.
+ >>> example_builder.add_image_matrix_feature(
+ image_matrix, image_source_id=image_source_id)
+ * For adding single-channel depth image feature with JPEG encoding:
+ >>> example_builder.add_image_matrix_feature(
+ image_matrix, image_format=ImageFormat.JPEG,
+ feature_prefix='depth')
+
+ Args:
+ image_matrix: Numpy image matrix with shape (height, width, channels)
+ image_format: Image format string, defaults to 'PNG'.
+ image_source_id: Unique string ID to identify the image. Hashed image will
+ be used if the field is not provided.
+ feature_prefix: Feature prefix for image features.
+ label: the label or a list of labels for the image.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ encoded_image = image_utils.encode_image(image_matrix, image_format)
+ height, width, num_channels = image_matrix.shape
+
+ return self.add_encoded_image_feature(encoded_image, image_format, height,
+ width, num_channels, image_source_id,
+ feature_prefix, label)
+
+ def add_encoded_image_feature(
+ self,
+ encoded_image: bytes,
+ image_format: Optional[str] = None,
+ height: Optional[int] = None,
+ width: Optional[int] = None,
+ num_channels: Optional[int] = None,
+ image_source_id: Optional[bytes] = None,
+ feature_prefix: Optional[str] = None,
+ label: Optional[Union[int, Sequence[int]]] = None) -> 'TfExampleBuilder':
+ """Adds encoded image features to the example.
+
+ See `tf_example_feature_key.EncodedImageFeatureKey` for list of feature keys
+ that will be added to the example.
+
+ Image format, height, width, and channels are inferred from the encoded
+ image bytes if any of them is not provided. Hashed image will be used if
+ pre-generated source ID is not provided.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ * For adding RGB image feature:
+ >>> example_builder.add_encoded_image_feature(image_bytes)
+ * For adding RGB image feature with pre-generated source ID:
+ >>> example_builder.add_encoded_image_feature(
+ image_bytes, image_source_id=image_source_id)
+ * For adding single-channel depth image feature:
+ >>> example_builder.add_encoded_image_feature(
+ image_bytes, feature_prefix='depth')
+
+ Args:
+ encoded_image: Encoded image string.
+ image_format: Image format string.
+ height: Number of rows.
+ width: Number of columns.
+ num_channels: Number of channels.
+ image_source_id: Unique string ID to identify the image.
+ feature_prefix: Feature prefix for image features.
+ label: the label or a list of labels for the image.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ if image_format == 'RAW':
+ if not (height and width and num_channels):
+ raise ValueError('For raw image feature, height, width and '
+ 'num_channels fields are required.')
+ if not all((height, width, num_channels, image_format)):
+ (height, width, num_channels, image_format) = (
+ image_utils.decode_image_metadata(encoded_image))
+ else:
+ image_format = image_utils.validate_image_format(image_format)
+
+ feature_key = tf_example_feature_key.EncodedImageFeatureKey(feature_prefix)
+
+ # If source ID is not provided, we use hashed encoded image as the source
+ # ID. Note that we only keep 24 bits to be consistent with the Model Garden
+ # requirement, which will transform the source ID into float32.
+ if not image_source_id:
+ hashed_image = int(hashlib.blake2s(encoded_image).hexdigest(), 16)
+ image_source_id = _to_bytes(str(hashed_image % ((1 << 24) + 1)))
+
+ if label is not None:
+ self.add_ints_feature(feature_key.label, label)
+
+ return (
+ self.add_bytes_feature(feature_key.encoded, encoded_image)
+ .add_bytes_feature(feature_key.format, image_format)
+ .add_ints_feature(feature_key.height, [height])
+ .add_ints_feature(feature_key.width, [width])
+ .add_ints_feature(feature_key.num_channels, num_channels)
+ .add_bytes_feature(feature_key.source_id, image_source_id))
+
+ def add_boxes_feature(
+ self,
+ xmins: Sequence[float],
+ xmaxs: Sequence[float],
+ ymins: Sequence[float],
+ ymaxs: Sequence[float],
+ labels: Sequence[int],
+ confidences: Optional[Sequence[float]] = None,
+ normalized: bool = True,
+ feature_prefix: Optional[str] = None) -> 'TfExampleBuilder':
+ """Adds box and label features to the example.
+
+ Four features will be generated for xmin, ymin, xmax, and ymax. One feature
+ will be generated for label. Different feature keys will be used for
+ normalized boxes and pixel-value boxes, depending on the value of
+ `normalized`.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ >>> example_builder.add_boxes_feature(xmins, xmaxs, ymins, ymaxs, labels)
+
+ Args:
+ xmins: A list of minimum X coordinates.
+ xmaxs: A list of maximum X coordinates.
+ ymins: A list of minimum Y coordinates.
+ ymaxs: A list of maximum Y coordinates.
+ labels: The labels of added boxes.
+ confidences: The confidences of added boxes.
+ normalized: Indicate if the coordinates of boxes are normalized.
+ feature_prefix: Feature prefix for added box features.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ if normalized:
+ feature_key = tf_example_feature_key.BoxFeatureKey(feature_prefix)
+ else:
+ feature_key = tf_example_feature_key.BoxPixelFeatureKey(feature_prefix)
+
+ self.add_floats_feature(feature_key.xmin, xmins)
+ self.add_floats_feature(feature_key.xmax, xmaxs)
+ self.add_floats_feature(feature_key.ymin, ymins)
+ self.add_floats_feature(feature_key.ymax, ymaxs)
+ self.add_ints_feature(feature_key.label, labels)
+ if confidences is not None:
+ self.add_floats_feature(feature_key.confidence, confidences)
+ return self
+
+ def _compute_mask_areas(
+ self, instance_mask_matrices: np.ndarray) -> Sequence[float]:
+ return np.sum(
+ instance_mask_matrices, axis=(1, 2, 3),
+ dtype=float).flatten().tolist()
+
+ def add_instance_mask_matrices_feature(
+ self,
+ instance_mask_matrices: np.ndarray,
+ feature_prefix: Optional[str] = None) -> 'TfExampleBuilder':
+ """Encodes and adds instance mask features to the example.
+
+ See `tf_example_feature_key.EncodedInstanceMaskFeatureKey` for list of
+ feature keys that will be added to the example. Please note that all masks
+ will be encoded as PNG images.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ >>> example_builder.add_instance_mask_matrices_feature(
+ instance_mask_matrices)
+
+ TODO(b/223653024): Provide a way to generate visualization mask from
+ feature mask.
+
+ Args:
+ instance_mask_matrices: Numpy instance mask matrices with shape
+ (num_instance, height, width, 1) or (num_instance, height, width).
+ feature_prefix: Feature prefix for instance mask features.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ if len(instance_mask_matrices.shape) == 3:
+ instance_mask_matrices = instance_mask_matrices[..., np.newaxis]
+
+ mask_areas = self._compute_mask_areas(instance_mask_matrices)
+ encoded_instance_masks = list(
+ map(lambda x: image_utils.encode_image(x, 'PNG'),
+ instance_mask_matrices))
+
+ return self.add_encoded_instance_masks_feature(encoded_instance_masks,
+ mask_areas, feature_prefix)
+
+ def add_encoded_instance_masks_feature(
+ self,
+ encoded_instance_masks: Sequence[bytes],
+ mask_areas: Optional[Sequence[float]] = None,
+ feature_prefix: Optional[str] = None) -> 'TfExampleBuilder':
+ """Adds encoded instance mask features to the example.
+
+ See `tf_example_feature_key.EncodedInstanceMaskFeatureKey` for list of
+ feature keys that will be added to the example.
+
+ Image area is inferred from the encoded instance mask bytes if not provided.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ >>> example_builder.add_encoded_instance_masks_feature(
+ instance_mask_bytes)
+
+ TODO(b/223653024): Provide a way to generate visualization mask from
+ feature mask.
+
+ Args:
+ encoded_instance_masks: A list of encoded instance mask string. Note that
+ the encoding is not changed in this function and it always assumes the
+ image is in "PNG" format.
+ mask_areas: Areas for each instance masks.
+ feature_prefix: Feature prefix for instance mask features.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ encoded_instance_masks = _to_bytes_array(encoded_instance_masks)
+
+ if mask_areas is None:
+ instance_mask_matrices = np.array(
+ list(map(image_utils.decode_image, encoded_instance_masks)))
+ mask_areas = self._compute_mask_areas(instance_mask_matrices)
+
+ feature_key = tf_example_feature_key.EncodedInstanceMaskFeatureKey(
+ feature_prefix)
+ return (
+ self.add_bytes_feature(feature_key.mask, encoded_instance_masks)
+ .add_floats_feature(feature_key.area, mask_areas))
+
+ def add_semantic_mask_matrix_feature(
+ self,
+ mask_matrix: np.ndarray,
+ mask_format: str = 'PNG',
+ visualization_mask_matrix: Optional[np.ndarray] = None,
+ visualization_mask_format: str = 'PNG',
+ feature_prefix: Optional[str] = None) -> 'TfExampleBuilder':
+ """Encodes and adds semantic mask features to the example.
+
+ See `tf_example_feature_key.EncodedSemanticMaskFeatureKey` for list of
+ feature keys that will be added to the example.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ * For adding semantic mask feature:
+ >>> example_builder.add_semantic_mask_matrix_feature(
+ semantic_mask_matrix)
+ * For adding semantic mask feature and visualization mask feature:
+ >>> example_builder.add_semantic_mask_matrix_feature(
+ semantic_mask_matrix,
+ visualization_mask_matrix=visualization_mask_matrix)
+ * For adding predicted semantic mask feature with visualization mask:
+ >>> example_builder.add_encoded_semantic_mask_feature(
+ predicted_mask_matrix,
+ visualization_mask_matrix=predicted_visualization_mask_matrix,
+ feature_prefix='predicted')
+
+ TODO(b/223653024): Provide a way to generate visualization mask from
+ feature mask.
+
+ Args:
+ mask_matrix: Numpy semantic mask matrix with shape (height, width, 1) or
+ (height, width).
+ mask_format: Mask format string, defaults to 'PNG'.
+ visualization_mask_matrix: Numpy visualization mask matrix for semantic
+ mask with shape (height, width, 3).
+ visualization_mask_format: Visualization mask format string, defaults to
+ 'PNG'.
+ feature_prefix: Feature prefix for semantic mask features.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ if len(mask_matrix.shape) == 2:
+ mask_matrix = mask_matrix[..., np.newaxis]
+ encoded_mask = image_utils.encode_image(mask_matrix, mask_format)
+
+ encoded_visualization_mask = None
+ if visualization_mask_matrix is not None:
+ encoded_visualization_mask = image_utils.encode_image(
+ visualization_mask_matrix, visualization_mask_format)
+
+ return self.add_encoded_semantic_mask_feature(encoded_mask, mask_format,
+ encoded_visualization_mask,
+ visualization_mask_format,
+ feature_prefix)
+
+ def add_encoded_semantic_mask_feature(
+ self, encoded_mask: bytes,
+ mask_format: str = 'PNG',
+ encoded_visualization_mask: Optional[bytes] = None,
+ visualization_mask_format: str = 'PNG',
+ feature_prefix: Optional[str] = None) -> 'TfExampleBuilder':
+ """Adds encoded semantic mask features to the example.
+
+ See `tf_example_feature_key.EncodedSemanticMaskFeatureKey` for list of
+ feature keys that will be added to the example.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ * For adding semantic mask feature:
+ >>> example_builder.add_encoded_semantic_mask_feature(semantic_mask_bytes)
+ * For adding semantic mask feature and visualization mask feature:
+ >>> example_builder.add_encoded_semantic_mask_feature(
+ semantic_mask_bytes,
+ encoded_visualization_mask=visualization_mask_bytes)
+ * For adding predicted semantic mask feature with visualization mask:
+ >>> example_builder.add_encoded_semantic_mask_feature(
+ predicted_mask_bytes,
+ encoded_visualization_mask=predicted_visualization_mask_bytes,
+ feature_prefix='predicted')
+
+ TODO(b/223653024): Provide a way to generate visualization mask from
+ feature mask.
+
+ Args:
+ encoded_mask: Encoded semantic mask string.
+ mask_format: Semantic mask format string, defaults to 'PNG'.
+ encoded_visualization_mask: Encoded visualization mask string.
+ visualization_mask_format: Visualization mask format string, defaults to
+ 'PNG'.
+ feature_prefix: Feature prefix for semantic mask features.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ feature_key = tf_example_feature_key.EncodedSemanticMaskFeatureKey(
+ feature_prefix)
+ example_builder = (
+ self.add_bytes_feature(feature_key.mask, encoded_mask)
+ .add_bytes_feature(feature_key.mask_format, mask_format))
+ if encoded_visualization_mask is not None:
+ example_builder = (
+ example_builder.add_bytes_feature(
+ feature_key.visualization_mask, encoded_visualization_mask)
+ .add_bytes_feature(
+ feature_key.visualization_mask_format, visualization_mask_format))
+ return example_builder
+
+ def add_panoptic_mask_matrix_feature(
+ self,
+ panoptic_category_mask_matrix: np.ndarray,
+ panoptic_instance_mask_matrix: np.ndarray,
+ panoptic_category_mask_format: str = 'PNG',
+ panoptic_instance_mask_format: str = 'PNG',
+ feature_prefix: Optional[str] = None) -> 'TfExampleBuilder':
+ """Encodes and adds panoptic mask features to the example.
+
+ See `tf_example_feature_key.EncodedPanopticMaskFeatureKey` for list of
+ feature keys that will be added to the example.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ >>> example_builder.add_panoptic_mask_matrix_feature(
+ panoptic_category_mask_matrix, panoptic_instance_mask_matrix)
+
+ TODO(b/223653024): Provide a way to generate visualization mask from
+ feature mask.
+
+ Args:
+ panoptic_category_mask_matrix: Numpy panoptic category mask matrix with
+ shape (height, width, 1) or (height, width).
+ panoptic_instance_mask_matrix: Numpy panoptic instance mask matrix with
+ shape (height, width, 1) or (height, width).
+ panoptic_category_mask_format: Panoptic category mask format string,
+ defaults to 'PNG'.
+ panoptic_instance_mask_format: Panoptic instance mask format string,
+ defaults to 'PNG'.
+ feature_prefix: Feature prefix for panoptic mask features.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ if len(panoptic_category_mask_matrix.shape) == 2:
+ panoptic_category_mask_matrix = (
+ panoptic_category_mask_matrix[..., np.newaxis])
+ if len(panoptic_instance_mask_matrix.shape) == 2:
+ panoptic_instance_mask_matrix = (
+ panoptic_instance_mask_matrix[..., np.newaxis])
+ encoded_panoptic_category_mask = image_utils.encode_image(
+ panoptic_category_mask_matrix, panoptic_category_mask_format)
+ encoded_panoptic_instance_mask = image_utils.encode_image(
+ panoptic_instance_mask_matrix, panoptic_instance_mask_format)
+
+ return self.add_encoded_panoptic_mask_feature(
+ encoded_panoptic_category_mask, encoded_panoptic_instance_mask,
+ panoptic_category_mask_format, panoptic_instance_mask_format,
+ feature_prefix)
+
+ def add_encoded_panoptic_mask_feature(
+ self,
+ encoded_panoptic_category_mask: bytes,
+ encoded_panoptic_instance_mask: bytes,
+ panoptic_category_mask_format: str = 'PNG',
+ panoptic_instance_mask_format: str = 'PNG',
+ feature_prefix: Optional[str] = None) -> 'TfExampleBuilder':
+ """Adds encoded panoptic mask features to the example.
+
+ See `tf_example_feature_key.EncodedPanopticMaskFeatureKey` for list of
+ feature keys that will be added to the example.
+
+ Example usages:
+ >>> example_builder = TfExampleBuilder()
+ >>> example_builder.add_encoded_panoptic_mask_feature(
+ encoded_panoptic_category_mask, encoded_panoptic_instance_mask)
+
+ TODO(b/223653024): Provide a way to generate visualization mask from
+ feature mask.
+
+ Args:
+ encoded_panoptic_category_mask: Encoded panoptic category mask string.
+ encoded_panoptic_instance_mask: Encoded panoptic instance mask string.
+ panoptic_category_mask_format: Panoptic category mask format string,
+ defaults to 'PNG'.
+ panoptic_instance_mask_format: Panoptic instance mask format string,
+ defaults to 'PNG'.
+ feature_prefix: Feature prefix for panoptic mask features.
+
+ Returns:
+ The builder object for subsequent method calls.
+ """
+ feature_key = tf_example_feature_key.EncodedPanopticMaskFeatureKey(
+ feature_prefix)
+ return (
+ self.add_bytes_feature(
+ feature_key.category_mask, encoded_panoptic_category_mask)
+ .add_bytes_feature(
+ feature_key.category_mask_format, panoptic_category_mask_format)
+ .add_bytes_feature(
+ feature_key.instance_mask, encoded_panoptic_instance_mask)
+ .add_bytes_feature(
+ feature_key.instance_mask_format, panoptic_instance_mask_format))
+
diff --git a/official/vision/data/tf_example_builder_test.py b/official/vision/data/tf_example_builder_test.py
new file mode 100644
index 00000000000..a9da0be74d6
--- /dev/null
+++ b/official/vision/data/tf_example_builder_test.py
@@ -0,0 +1,641 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for tf_example_builder."""
+
+from absl.testing import parameterized
+import tensorflow as tf, tf_keras
+from official.vision.data import fake_feature_generator
+from official.vision.data import image_utils
+from official.vision.data import tf_example_builder
+
+
+class TfExampleBuilderTest(tf.test.TestCase, parameterized.TestCase):
+
+ @parameterized.named_parameters(('RGB_PNG', 128, 64, 3, 'PNG', 3),
+ ('RGB_RAW', 128, 128, 3, 'RAW', 0),
+ ('RGB_JPEG', 64, 128, 3, 'JPEG', [2, 5]))
+ def test_add_image_matrix_feature_success(self, height, width, num_channels,
+ image_format, label):
+ # Prepare test data.
+ image_np = fake_feature_generator.generate_image_np(height, width,
+ num_channels)
+ expected_image_bytes = image_utils.encode_image(image_np, image_format)
+ hashed_image = bytes('10242048', 'ascii')
+
+ # Run code logic.
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_image_matrix_feature(
+ image_np, image_format, hashed_image, label=label)
+ example = example_builder.example
+
+ # Verify outputs.
+ # Prefer to use string literal for feature keys to directly display the
+ # structure of the expected tf.train.Example.
+ if isinstance(label, int):
+ expected_labels = [label]
+ else:
+ expected_labels = label
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[expected_image_bytes])),
+ 'image/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ 'image/height':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[height])),
+ 'image/width':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[width])),
+ 'image/channels':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(
+ value=[num_channels])),
+ 'image/source_id':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[hashed_image])),
+ 'image/class/label':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(
+ value=expected_labels)),
+ })), example)
+
+ def test_add_image_matrix_feature_with_feature_prefix_success(self):
+ height = 64
+ width = 64
+ num_channels = 1
+ image_format = 'PNG'
+ feature_prefix = 'depth'
+ label = 8
+ image_np = fake_feature_generator.generate_image_np(height, width,
+ num_channels)
+ expected_image_bytes = image_utils.encode_image(image_np, image_format)
+ hashed_image = bytes('10242048', 'ascii')
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_image_matrix_feature(
+ image_np,
+ image_format,
+ hashed_image,
+ feature_prefix=feature_prefix,
+ label=label)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'depth/image/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[expected_image_bytes])),
+ 'depth/image/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ 'depth/image/height':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[height])),
+ 'depth/image/width':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[width])),
+ 'depth/image/channels':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(
+ value=[num_channels])),
+ 'depth/image/source_id':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[hashed_image])),
+ 'depth/image/class/label':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[label]))
+ })), example)
+
+ def test_add_encoded_raw_image_feature_success(self):
+ height = 128
+ width = 128
+ num_channels = 3
+ image_format = 'RAW'
+ expected_image_bytes = bytes('image', 'ascii')
+ hashed_image = bytes('16188651', 'ascii')
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_encoded_image_feature(expected_image_bytes, 'RAW',
+ height, width, num_channels)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[expected_image_bytes])),
+ 'image/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ 'image/height':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[height])),
+ 'image/width':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[width])),
+ 'image/channels':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(
+ value=[num_channels])),
+ 'image/source_id':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[hashed_image]))
+ })), example)
+
+ def test_add_encoded_raw_image_feature_valueerror(self):
+ image_format = 'RAW'
+ image_bytes = tf.bfloat16.as_numpy_dtype
+ image_np = fake_feature_generator.generate_image_np(1, 1, 1)
+ image_np = image_np.astype(image_bytes)
+ expected_image_bytes = image_utils.encode_image(image_np, image_format)
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ with self.assertRaises(ValueError):
+ example_builder.add_encoded_image_feature(expected_image_bytes,
+ image_format)
+
+ @parameterized.product(
+ miss_image_format=(True, False),
+ miss_height=(True, False),
+ miss_width=(True, False),
+ miss_num_channels=(True, False),
+ miss_label=(True, False))
+ def test_add_encoded_image_feature_success(self, miss_image_format,
+ miss_height, miss_width,
+ miss_num_channels,
+ miss_label):
+ height = 64
+ width = 64
+ num_channels = 3
+ image_format = 'PNG'
+ image_np = fake_feature_generator.generate_image_np(height, width,
+ num_channels)
+ image_bytes = image_utils.encode_image(image_np, image_format)
+ # We don't test on image_source_id because encoding process becomes
+ # non-deterministic.
+ hashed_image = bytes('10242048', 'ascii')
+ label = 5
+
+ image_format = None if miss_image_format else image_format
+ height = None if miss_height else height
+ width = None if miss_width else width
+ num_channels = None if miss_num_channels else num_channels
+ label = None if miss_label else label
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_encoded_image_feature(
+ image_bytes,
+ image_format=image_format,
+ height=height,
+ width=width,
+ num_channels=num_channels,
+ image_source_id=hashed_image,
+ label=label)
+ example = example_builder.example
+
+ expected_features = {
+ 'image/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[image_bytes])),
+ 'image/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes('PNG', 'ascii')])),
+ 'image/height':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[64])),
+ 'image/width':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[64])),
+ 'image/channels':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[3])),
+ 'image/source_id':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[hashed_image]))}
+ if not miss_label:
+ expected_features.update({
+ 'image/class/label':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[label]))})
+ self.assertProtoEquals(
+ tf.train.Example(features=tf.train.Features(feature=expected_features)),
+ example)
+
+ @parameterized.named_parameters(('no_box', 0), ('10_boxes', 10))
+ def test_add_normalized_boxes_feature(self, num_boxes):
+ normalized_boxes_np = fake_feature_generator.generate_normalized_boxes_np(
+ num_boxes)
+ ymins, xmins, ymaxs, xmaxs = normalized_boxes_np.T.tolist()
+ labels = fake_feature_generator.generate_classes_np(
+ 2, size=num_boxes).tolist()
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_boxes_feature(
+ xmins, xmaxs, ymins, ymaxs, labels=labels, normalized=True)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/object/bbox/xmin':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmins)),
+ 'image/object/bbox/ymin':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymins)),
+ 'image/object/bbox/xmax':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmaxs)),
+ 'image/object/bbox/ymax':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymaxs)),
+ 'image/object/class/label':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=labels)),
+ })), example)
+
+ @parameterized.named_parameters(('no_box', 0), ('10_boxes', 10))
+ def test_add_box_pixels_feature(self, num_boxes):
+ height, width = 10, 10
+ boxes_np = fake_feature_generator.generate_boxes_np(height, width,
+ num_boxes)
+ ymins, xmins, ymaxs, xmaxs = boxes_np.T.tolist()
+ labels = fake_feature_generator.generate_classes_np(
+ 2, size=num_boxes).tolist()
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_boxes_feature(
+ xmins, xmaxs, ymins, ymaxs, labels=labels, normalized=False)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/object/bbox/xmin_pixels':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmins)),
+ 'image/object/bbox/ymin_pixels':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymins)),
+ 'image/object/bbox/xmax_pixels':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmaxs)),
+ 'image/object/bbox/ymax_pixels':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymaxs)),
+ 'image/object/class/label':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=labels)),
+ })), example)
+
+ @parameterized.named_parameters(('no_box', 0), ('10_boxes', 10))
+ def test_add_normalized_boxes_feature_with_confidence_and_prefix(
+ self, num_boxes):
+ normalized_boxes_np = fake_feature_generator.generate_normalized_boxes_np(
+ num_boxes)
+ ymins, xmins, ymaxs, xmaxs = normalized_boxes_np.T.tolist()
+ labels = fake_feature_generator.generate_classes_np(
+ 2, size=num_boxes).tolist()
+ confidences = fake_feature_generator.generate_confidences_np(
+ size=num_boxes).tolist()
+ feature_prefix = 'predicted'
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_boxes_feature(
+ xmins,
+ xmaxs,
+ ymins,
+ ymaxs,
+ labels=labels,
+ confidences=confidences,
+ normalized=True,
+ feature_prefix=feature_prefix)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'predicted/image/object/bbox/xmin':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmins)),
+ 'predicted/image/object/bbox/ymin':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymins)),
+ 'predicted/image/object/bbox/xmax':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmaxs)),
+ 'predicted/image/object/bbox/ymax':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymaxs)),
+ 'predicted/image/object/class/label':
+ tf.train.Feature(
+ int64_list=tf.train.Int64List(value=labels)),
+ 'predicted/image/object/bbox/confidence':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(value=confidences)),
+ })), example)
+
+ @parameterized.named_parameters(('no_mask', 128, 64, 0),
+ ('10_masks', 64, 128, 10))
+ def test_add_instance_mask_matrices_feature_success(self, height, width,
+ num_masks):
+ # Prepare test data.
+ instance_masks_np = fake_feature_generator.generate_instance_masks_np(
+ height,
+ width,
+ boxes_np=fake_feature_generator.generate_boxes_np(
+ height, width, num_masks),
+ normalized=False)
+ expected_instance_masks_bytes = list(
+ map(lambda x: image_utils.encode_image(x, 'PNG'), instance_masks_np))
+
+ # Run code logic.
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_instance_mask_matrices_feature(instance_masks_np)
+ example = example_builder.example
+
+ # Verify outputs.
+ # Prefer to use string literal for feature keys to directly display the
+ # structure of the expected tf.train.Example.
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/object/mask':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=expected_instance_masks_bytes)),
+ 'image/object/area':
+ # The box area is 4x smaller than the image, and the
+ # mask area is 2x smaller than the box.
+ tf.train.Feature(
+ float_list=tf.train.FloatList(
+ value=[height * width / 8] * num_masks))
+ })), example)
+
+ @parameterized.named_parameters(('with_mask_areas', True),
+ ('without_mask_areas', False))
+ def test_add_encoded_instance_masks_feature_success(self, has_mask_areas):
+ height = 64
+ width = 64
+ image_format = 'PNG'
+ mask_np = fake_feature_generator.generate_semantic_mask_np(height, width, 2)
+ mask_bytes = image_utils.encode_image(mask_np, image_format)
+
+ test_masks = [mask_bytes for _ in range(2)]
+ mask_areas = [2040., 2040.] if has_mask_areas else None
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_encoded_instance_masks_feature(
+ test_masks, mask_areas=mask_areas)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/object/mask':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=test_masks)),
+ 'image/object/area':
+ tf.train.Feature(
+ float_list=tf.train.FloatList(
+ value=[2040., 2040.])),
+ })), example)
+
+ @parameterized.named_parameters(
+ ('with_visualization_mask', 128, 64, True),
+ ('without_visualization_mask', 64, 128, False))
+ def test_add_semantic_mask_matrices_feature_success(self, height, width,
+ has_visualization_mask):
+ # Prepare test data.
+ semantic_mask_np = fake_feature_generator.generate_semantic_mask_np(
+ height, width, 2)
+ image_format = 'PNG'
+ expected_feature_dict = {
+ 'image/segmentation/class/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[
+ image_utils.encode_image(semantic_mask_np, image_format)
+ ])),
+ 'image/segmentation/class/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ }
+ visualization_mask_np = None
+ if has_visualization_mask:
+ visualization_mask_np = fake_feature_generator.generate_image_np(
+ height, width)
+ expected_feature_dict.update({
+ 'image/segmentation/class/visualization/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[
+ image_utils.encode_image(visualization_mask_np,
+ image_format)
+ ])),
+ 'image/segmentation/class/visualization/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ })
+
+ # Run code logic.
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_semantic_mask_matrix_feature(semantic_mask_np,
+ image_format,
+ visualization_mask_np,
+ image_format)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(feature=expected_feature_dict)), example)
+
+ @parameterized.named_parameters(('with_visualization_mask', True),
+ ('without_visualization_mask', False))
+ def test_add_encoded_semantic_mask_feature_success(self,
+ has_visualization_mask):
+ height, width = 64, 64
+ semantic_mask_np = fake_feature_generator.generate_semantic_mask_np(
+ height, width, 2)
+ image_format = 'PNG'
+ encoded_semantic_mask = image_utils.encode_image(semantic_mask_np,
+ image_format)
+ expected_feature_dict = {
+ 'image/segmentation/class/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[encoded_semantic_mask])),
+ 'image/segmentation/class/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ }
+ encoded_visualization_mask = None
+ if has_visualization_mask:
+ visualization_mask_np = fake_feature_generator.generate_image_np(
+ height, width)
+ encoded_visualization_mask = image_utils.encode_image(
+ visualization_mask_np, image_format)
+ expected_feature_dict.update({
+ 'image/segmentation/class/visualization/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[encoded_visualization_mask])),
+ 'image/segmentation/class/visualization/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ })
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_encoded_semantic_mask_feature(
+ encoded_semantic_mask, image_format, encoded_visualization_mask,
+ image_format)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(feature=expected_feature_dict)), example)
+
+ def test_add_panoptic_mask_matrices_feature_success(self):
+ # Prepare test data.
+ height, width, num_instances = 64, 64, 10
+ num_thing_classes, num_semantic_segmentation_classes = 3, 6
+ image_format = 'PNG'
+
+ normalized_boxes_np = fake_feature_generator.generate_normalized_boxes_np(
+ num_instances)
+ instance_masks_np = fake_feature_generator.generate_instance_masks_np(
+ height, width, normalized_boxes_np)
+ instance_classes_np = fake_feature_generator.generate_classes_np(
+ num_thing_classes, num_instances)
+ semantic_mask_np = fake_feature_generator.generate_semantic_mask_np(
+ height, width, num_semantic_segmentation_classes)
+ panoptic_category_mask_np, panoptic_instance_mask_np = (
+ fake_feature_generator.generate_panoptic_masks_np(
+ semantic_mask_np, instance_masks_np, instance_classes_np,
+ num_thing_classes - 1))
+
+ # Run code logic.
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_panoptic_mask_matrix_feature(panoptic_category_mask_np,
+ panoptic_instance_mask_np,
+ image_format,
+ image_format)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/panoptic/category/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[
+ image_utils.encode_image(
+ panoptic_category_mask_np, image_format)
+ ])),
+ 'image/panoptic/category/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ 'image/panoptic/instance/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[
+ image_utils.encode_image(
+ panoptic_instance_mask_np, image_format)
+ ])),
+ 'image/panoptic/instance/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ })), example)
+
+ def test_add_encoded_panoptic_mask_feature_success(self):
+ # Prepare test data.
+ height, width, num_instances = 64, 64, 10
+ num_thing_classes, num_semantic_segmentation_classes = 3, 6
+ image_format = 'PNG'
+
+ normalized_boxes_np = fake_feature_generator.generate_normalized_boxes_np(
+ num_instances)
+ instance_masks_np = fake_feature_generator.generate_instance_masks_np(
+ height, width, normalized_boxes_np)
+ instance_classes_np = fake_feature_generator.generate_classes_np(
+ num_thing_classes, num_instances)
+ semantic_mask_np = fake_feature_generator.generate_semantic_mask_np(
+ height, width, num_semantic_segmentation_classes)
+ panoptic_category_mask_np, panoptic_instance_mask_np = (
+ fake_feature_generator.generate_panoptic_masks_np(
+ semantic_mask_np, instance_masks_np, instance_classes_np,
+ num_thing_classes - 1))
+
+ encoded_panoptic_category_mask = image_utils.encode_image(
+ panoptic_category_mask_np, image_format)
+ encoded_panoptic_instance_mask = image_utils.encode_image(
+ panoptic_instance_mask_np, image_format)
+
+ example_builder = tf_example_builder.TfExampleBuilder()
+ example_builder.add_encoded_panoptic_mask_feature(
+ encoded_panoptic_category_mask, encoded_panoptic_instance_mask,
+ image_format, image_format)
+ example = example_builder.example
+
+ self.assertProtoEquals(
+ tf.train.Example(
+ features=tf.train.Features(
+ feature={
+ 'image/panoptic/category/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[encoded_panoptic_category_mask])),
+ 'image/panoptic/category/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ 'image/panoptic/instance/encoded':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[encoded_panoptic_instance_mask])),
+ 'image/panoptic/instance/format':
+ tf.train.Feature(
+ bytes_list=tf.train.BytesList(
+ value=[bytes(image_format, 'ascii')])),
+ })), example)
+
+
+if __name__ == '__main__':
+ tf.test.main()
diff --git a/official/vision/data/tf_example_feature_key.py b/official/vision/data/tf_example_feature_key.py
new file mode 100644
index 00000000000..062309e56a3
--- /dev/null
+++ b/official/vision/data/tf_example_feature_key.py
@@ -0,0 +1,174 @@
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Data classes for tf.Example proto feature keys in vision tasks.
+
+Feature keys are grouped by feature types. Key names follow conventions in
+go/tf-example.
+"""
+import dataclasses
+import functools
+
+from official.core import tf_example_feature_key
+
+# Disable init function to use the one defined in base class.
+dataclass = functools.partial(dataclasses.dataclass(init=False)) # pyrefly: ignore[bad-argument-type]
+
+
+@dataclass
+class EncodedImageFeatureKey(tf_example_feature_key.TfExampleFeatureKeyBase):
+ """Feature keys for a single encoded image.
+
+ The image matrix is expected to be in the shape of (height, width,
+ num_channels).
+
+ Attributes:
+ encoded: encoded image bytes.
+ format: format string, e.g. 'PNG'.
+ height: number of rows.
+ width: number of columns.
+ num_channels: number of channels.
+ source_id: Unique string ID to identify the image.
+ label: the label or a list of labels for the image.
+ """
+ encoded: str = 'image/encoded'
+ format: str = 'image/format'
+ height: str = 'image/height'
+ width: str = 'image/width'
+ num_channels: str = 'image/channels'
+ source_id: str = 'image/source_id'
+ label: str = 'image/class/label'
+
+
+@dataclass
+class BoxFeatureKey(tf_example_feature_key.TfExampleFeatureKeyBase):
+ """Feature keys for normalized boxes representing objects in an image.
+
+ Each box is defined by ((ymin, xmin), (ymax, xmax)).
+
+ The origin point of an image matrix is top left.
+
+ Note: The coordinate values are normalized to [0, 1], this is commonly adopted
+ by most model implementations.
+
+ Attributes:
+ xmin: The x coordinate (column) of top-left corner.
+ xmax: The x coordinate (column) of bottom-right corner.
+ ymin: The y coordinate (row) of top-left corner.
+ ymax: The y coordinate (row) of bottom-right corner.
+ label: The class id.
+ confidence: The confidence score of the box, could be prior score (for
+ training) or predicted score (for prediction).
+ """
+ xmin: str = 'image/object/bbox/xmin'
+ xmax: str = 'image/object/bbox/xmax'
+ ymin: str = 'image/object/bbox/ymin'
+ ymax: str = 'image/object/bbox/ymax'
+ label: str = 'image/object/class/label'
+ confidence: str = 'image/object/bbox/confidence'
+
+
+@dataclass
+class BoxPixelFeatureKey(tf_example_feature_key.TfExampleFeatureKeyBase):
+ """Feature keys for boxes in pixel values representing objects in an image.
+
+ Each box is defined by ((ymin, xmin), (ymax, xmax)).
+
+ Note: The coordinate values are in the scale of the context image. The image
+ size is usually stored in `EncodedImageFeatureKey`.
+
+ Attributes:
+ xmin: The x coordinate (column) of top-left corner.
+ xmax: The x coordinate (column) of bottom-right corner.
+ ymin: The y coordinate (row) of top-left corner.
+ ymax: The y coordinate (row) of bottom-right corner.
+ label: The class id.
+ confidence: The confidence score of the box, could be prior score (for
+ training) or predicted score (for prediction).
+ """
+ xmin: str = 'image/object/bbox/xmin_pixels'
+ xmax: str = 'image/object/bbox/xmax_pixels'
+ ymin: str = 'image/object/bbox/ymin_pixels'
+ ymax: str = 'image/object/bbox/ymax_pixels'
+ label: str = 'image/object/class/label'
+ confidence: str = 'image/object/bbox/confidence'
+
+
+@dataclass
+class EncodedInstanceMaskFeatureKey(
+ tf_example_feature_key.TfExampleFeatureKeyBase):
+ """Feature keys for a single encoded instance mask.
+
+ The instance mask matrices are expected to be in the shape of (num_instances,
+ height, width, 1) or (num_instance, height, width). The height and width
+ correspond to the image height and width. For each instance mask, the pixel
+ value is either 0, representing a background, or 1, representing the object.
+
+ TODO(b/223653024): Add keys for visualization mask as well.
+
+ Attributes:
+ mask: Encoded instance mask bytes.
+ area: Total number of pixels that are marked as objects.
+ """
+ mask: str = 'image/object/mask'
+ area: str = 'image/object/area'
+
+
+@dataclass
+class EncodedSemanticMaskFeatureKey(
+ tf_example_feature_key.TfExampleFeatureKeyBase):
+ """Feature keys for a encoded semantic mask and its associated images.
+
+ The semantic mask matrix is expected to be in the shape of (height, width, 1)
+ or (height, width). The visualization mask matrix is expected to be in the
+ shape of (height, width, 3). The height and width correspond to the image
+ height and width. Each pixel in the semantic mask respresents a class.
+
+ Attributes:
+ mask: Encoded semantic mask bytes.
+ mask_format: Format string for semantic mask, e.g. 'PNG'.
+ visualization_mask: Encoded visualization mask bytes.
+ visualization_mask_format: Format string for visualization mask, e.g.
+ 'PNG'.
+ """
+ mask: str = 'image/segmentation/class/encoded'
+ mask_format: str = 'image/segmentation/class/format'
+ visualization_mask: str = 'image/segmentation/class/visualization/encoded'
+ visualization_mask_format: str = 'image/segmentation/class/visualization/format'
+
+
+@dataclass
+class EncodedPanopticMaskFeatureKey(
+ tf_example_feature_key.TfExampleFeatureKeyBase):
+ """Feature keys for encoded panoptic category and instance masks.
+
+ Both panoptic mask matrices are expected to be in the shape of (height, width,
+ 1) or (height, width). The height and width correspond to the image height and
+ width. For category mask, each pixel represents a class ID, and for instance
+ mask, each pixel represents an instance ID.
+
+ TODO(b/223653024): Add keys for visualization mask as well.
+
+ Attributes:
+ category_mask: Encoded panoptic category mask bytes.
+ category_mask_format: Format string for panoptic category mask, e.g.
+ 'PNG'.
+ instance_mask: Encoded panoptic instance mask bytes.
+ instance_mask_format: Format string for panoptic instance mask, e.g.
+ 'PNG'.
+ """
+ category_mask: str = 'image/panoptic/category/encoded'
+ category_mask_format: str = 'image/panoptic/category/format'
+ instance_mask: str = 'image/panoptic/instance/encoded'
+ instance_mask_format: str = 'image/panoptic/instance/format'
diff --git a/official/vision/data/tfrecord_lib.py b/official/vision/data/tfrecord_lib.py
index 96c7c5dc33b..ddbb159521e 100644
--- a/official/vision/data/tfrecord_lib.py
+++ b/official/vision/data/tfrecord_lib.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -21,7 +21,7 @@
from absl import logging
import numpy as np
from PIL import Image
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import multiprocessing as mp
@@ -171,7 +171,7 @@ def write_tf_record_dataset(output_path, annotation_iterator,
writers[idx % num_shards].write(tf_example.SerializeToString())
if multiple_processes is None or multiple_processes > 0:
- pool.close()
+ pool.close() # pyrefly: ignore[unbound-name]
pool.join()
for writer in writers:
diff --git a/official/vision/data/tfrecord_lib_test.py b/official/vision/data/tfrecord_lib_test.py
index 6825ff8cd4a..d918efb9493 100644
--- a/official/vision/data/tfrecord_lib_test.py
+++ b/official/vision/data/tfrecord_lib_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,8 +18,9 @@
from absl import flags
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.vision.data import create_coco_tf_record as create_coco_tf_record_lib
from official.vision.data import tfrecord_lib
@@ -88,6 +89,89 @@ def test_convert_to_feature_bytes_list(self):
proto = tfrecord_lib.convert_to_feature([b'123', b'456'])
self.assertSequenceAlmostEqual(proto.bytes_list.value, [b'123', b'456'])
+ def test_obj_annotation_tf_example(self):
+ images = [
+ {
+ 'id': 0,
+ 'file_name': 'example1.jpg',
+ 'height': 512,
+ 'width': 512,
+ },
+ {
+ 'id': 1,
+ 'file_name': 'example2.jpg',
+ 'height': 512,
+ 'width': 512,
+ },
+ ]
+ img_to_obj_annotation = {
+ 0: [{
+ 'id': 0,
+ 'image_id': 0,
+ 'category_id': 1,
+ 'bbox': [3, 1, 511, 510],
+ 'area': 260610.00,
+ 'segmentation': [],
+ 'iscrowd': 0,
+ }],
+ 1: [{
+ 'id': 1,
+ 'image_id': 1,
+ 'category_id': 1,
+ 'bbox': [1, 1, 100, 150],
+ 'area': 15000.00,
+ 'segmentation': [],
+ 'iscrowd': 0,
+ }],
+ }
+ id_to_name_map = {
+ 0: 'Super-Class',
+ 1: 'Class-1',
+ }
+
+ temp_dir = FLAGS.test_tmpdir
+ image_dir = os.path.join(temp_dir, 'data')
+ if not os.path.exists(image_dir):
+ os.mkdir(image_dir)
+ for image in images:
+ image_path = os.path.join(image_dir, image['file_name'])
+ tf_keras.utils.save_img(
+ image_path,
+ tf.ones(shape=(image['height'], image['width'], 3)).numpy(),
+ )
+
+ output_path = os.path.join(image_dir, 'train')
+ coco_annotations_iter = create_coco_tf_record_lib.generate_annotations(
+ images=images,
+ image_dirs=[image_dir],
+ panoptic_masks_dir=None,
+ img_to_obj_annotation=img_to_obj_annotation,
+ img_to_caption_annotation=None,
+ img_to_panoptic_annotation=None,
+ is_category_thing=None,
+ id_to_name_map=id_to_name_map,
+ include_panoptic_masks=False,
+ include_masks=False,
+ )
+
+ tfrecord_lib.write_tf_record_dataset(
+ output_path,
+ coco_annotations_iter,
+ create_coco_tf_record_lib.create_tf_example,
+ 1,
+ multiple_processes=0,
+ )
+ tfrecord_files = tf.io.gfile.glob(output_path + '*')
+
+ self.assertLen(tfrecord_files, 1)
+
+ ds = tf.data.TFRecordDataset(tfrecord_files)
+ assertion_count = 0
+ for _ in ds:
+ assertion_count += 1
+
+ self.assertEqual(assertion_count, 2)
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/vision/dataloaders/__init__.py b/official/vision/dataloaders/__init__.py
index 310bfb28f0c..e7e7c21950e 100644
--- a/official/vision/dataloaders/__init__.py
+++ b/official/vision/dataloaders/__init__.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/dataloaders/classification_input.py b/official/vision/dataloaders/classification_input.py
index aab8e29e43b..847c991c232 100644
--- a/official/vision/dataloaders/classification_input.py
+++ b/official/vision/dataloaders/classification_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -13,9 +13,8 @@
# limitations under the License.
"""Classification decoder and parser."""
-from typing import Any, Dict, List, Optional
-# Import libraries
-import tensorflow as tf
+from typing import Any, Dict, List, Optional, Tuple
+import tensorflow as tf, tf_keras
from official.vision.configs import common
from official.vision.dataloaders import decoder
@@ -23,9 +22,6 @@
from official.vision.ops import augment
from official.vision.ops import preprocess_ops
-MEAN_RGB = (0.485 * 255, 0.456 * 255, 0.406 * 255)
-STDDEV_RGB = (0.229 * 255, 0.224 * 255, 0.225 * 255)
-
DEFAULT_IMAGE_FIELD_KEY = 'image/encoded'
DEFAULT_LABEL_FIELD_KEY = 'image/class/label'
@@ -54,8 +50,8 @@ def __init__(self,
self._keys_to_features = keys_to_features
def decode(self, serialized_example):
- return tf.io.parse_single_example(
- serialized_example, self._keys_to_features)
+ return tf.io.parse_single_example(serialized_example,
+ self._keys_to_features)
class Parser(parser.Parser):
@@ -68,23 +64,32 @@ def __init__(self,
label_field_key: str = DEFAULT_LABEL_FIELD_KEY,
decode_jpeg_only: bool = True,
aug_rand_hflip: bool = True,
+ aug_crop: Optional[bool] = True,
aug_type: Optional[common.Augmentation] = None,
color_jitter: float = 0.,
random_erasing: Optional[common.RandomErasing] = None,
is_multilabel: bool = False,
- dtype: str = 'float32'):
+ dtype: str = 'float32',
+ crop_area_range: Optional[Tuple[float, float]] = (0.08, 1.0),
+ center_crop_fraction: Optional[
+ float] = preprocess_ops.CENTER_CROP_FRACTION,
+ tf_resize_method: str = 'bilinear',
+ three_augment: bool = False):
"""Initializes parameters for parsing annotations in the dataset.
Args:
output_size: `Tensor` or `list` for [height, width] of output image. The
output_size should be divided by the largest feature stride 2^max_level.
num_classes: `float`, number of classes.
- image_field_key: `str`, the key name to encoded image in tf.Example.
+ image_field_key: `str`, the key name to encoded image or decoded image
+ matrix in tf.Example.
label_field_key: `str`, the key name to label in tf.Example.
decode_jpeg_only: `bool`, if True, only JPEG format is decoded, this is
faster than decoding other types. Default is True.
- aug_rand_hflip: `bool`, if True, augment training with random
- horizontal flip.
+ aug_rand_hflip: `bool`, if True, augment training with random horizontal
+ flip.
+ aug_crop: `bool`, if True, perform random cropping during training and
+ center crop during validation.
aug_type: An optional Augmentation object to choose from AutoAugment and
RandAugment.
color_jitter: Magnitude of color jitter. If > 0, the value is used to
@@ -95,9 +100,18 @@ def __init__(self,
is_multilabel: A `bool`, whether or not each example has multiple labels.
dtype: `str`, cast output image in dtype. It can be 'float32', 'float16',
or 'bfloat16'.
+ crop_area_range: An optional `tuple` of (min_area, max_area) for image
+ random crop function to constraint crop operation. The cropped areas
+ of the image must contain a fraction of the input image within this
+ range. The default area range is (0.08, 1.0).
+ https://arxiv.org/abs/2204.07118.
+ center_crop_fraction: center_crop_fraction.
+ tf_resize_method: A `str`, interpolation method for resizing image.
+ three_augment: A bool, whether to apply three augmentations.
"""
self._output_size = output_size
self._aug_rand_hflip = aug_rand_hflip
+ self._aug_crop = aug_crop
self._num_classes = num_classes
self._image_field_key = image_field_key
if dtype == 'float32':
@@ -143,6 +157,10 @@ def __init__(self,
self._random_erasing = None
self._is_multilabel = is_multilabel
self._decode_jpeg_only = decode_jpeg_only
+ self._crop_area_range = crop_area_range
+ self._center_crop_fraction = center_crop_fraction
+ self._tf_resize_method = tf_resize_method
+ self._three_augment = three_augment
def _parse_train_data(self, decoded_tensors):
"""Parses data for training."""
@@ -167,29 +185,42 @@ def _parse_eval_data(self, decoded_tensors):
def _parse_train_image(self, decoded_tensors):
"""Parses image data for training."""
image_bytes = decoded_tensors[self._image_field_key]
-
- if self._decode_jpeg_only:
+ require_decoding = (
+ not tf.is_tensor(image_bytes) or image_bytes.dtype == tf.dtypes.string
+ )
+
+ if (
+ require_decoding
+ and self._decode_jpeg_only
+ and self._aug_crop
+ ):
image_shape = tf.image.extract_jpeg_shape(image_bytes)
# Crops image.
cropped_image = preprocess_ops.random_crop_image_v2(
- image_bytes, image_shape)
+ image_bytes, image_shape, area_range=self._crop_area_range)
image = tf.cond(
tf.reduce_all(tf.equal(tf.shape(cropped_image), image_shape)),
lambda: preprocess_ops.center_crop_image_v2(image_bytes, image_shape),
lambda: cropped_image)
else:
- # Decodes image.
- image = tf.io.decode_image(image_bytes, channels=3)
- image.set_shape([None, None, 3])
+ if require_decoding:
+ # Decodes image.
+ image = tf.io.decode_image(image_bytes, channels=3)
+ image.set_shape([None, None, 3])
+ else:
+ # Already decoded image matrix
+ image = image_bytes
# Crops image.
- cropped_image = preprocess_ops.random_crop_image(image)
+ if self._aug_crop:
+ cropped_image = preprocess_ops.random_crop_image(
+ image, area_range=self._crop_area_range)
- image = tf.cond(
- tf.reduce_all(tf.equal(tf.shape(cropped_image), tf.shape(image))),
- lambda: preprocess_ops.center_crop_image(image),
- lambda: cropped_image)
+ image = tf.cond(
+ tf.reduce_all(tf.equal(tf.shape(cropped_image), tf.shape(image))),
+ lambda: preprocess_ops.center_crop_image(image),
+ lambda: cropped_image)
if self._aug_rand_hflip:
image = tf.image.random_flip_left_right(image)
@@ -202,17 +233,23 @@ def _parse_train_image(self, decoded_tensors):
# Resizes image.
image = tf.image.resize(
- image, self._output_size, method=tf.image.ResizeMethod.BILINEAR)
+ image, self._output_size, method=self._tf_resize_method)
image.set_shape([self._output_size[0], self._output_size[1], 3])
# Apply autoaug or randaug.
if self._augmenter is not None:
image = self._augmenter.distort(image)
+ # Three augmentation
+ if self._three_augment:
+ image = augment.AutoAugment(
+ augmentation_name='deit3_three_augment',
+ translate_const=20,
+ ).distort(image)
+
# Normalizes image with mean and std pixel values.
- image = preprocess_ops.normalize_image(image,
- offset=MEAN_RGB,
- scale=STDDEV_RGB)
+ image = preprocess_ops.normalize_image(
+ image, offset=preprocess_ops.MEAN_RGB, scale=preprocess_ops.STDDEV_RGB)
# Random erasing after the image has been normalized
if self._random_erasing is not None:
@@ -226,34 +263,52 @@ def _parse_train_image(self, decoded_tensors):
def _parse_eval_image(self, decoded_tensors):
"""Parses image data for evaluation."""
image_bytes = decoded_tensors[self._image_field_key]
-
- if self._decode_jpeg_only:
+ require_decoding = (
+ not tf.is_tensor(image_bytes) or image_bytes.dtype == tf.dtypes.string
+ )
+
+ if (
+ require_decoding
+ and self._decode_jpeg_only
+ and self._aug_crop
+ ):
image_shape = tf.image.extract_jpeg_shape(image_bytes)
# Center crops.
- image = preprocess_ops.center_crop_image_v2(image_bytes, image_shape)
+ image = preprocess_ops.center_crop_image_v2(
+ image_bytes, image_shape, self._center_crop_fraction)
else:
- # Decodes image.
- image = tf.io.decode_image(image_bytes, channels=3)
- image.set_shape([None, None, 3])
+ if require_decoding:
+ # Decodes image.
+ image = tf.io.decode_image(image_bytes, channels=3)
+ image.set_shape([None, None, 3])
+ else:
+ # Already decoded image matrix
+ image = image_bytes
# Center crops.
- image = preprocess_ops.center_crop_image(image)
+ if self._aug_crop:
+ image = preprocess_ops.center_crop_image(
+ image, self._center_crop_fraction)
image = tf.image.resize(
- image, self._output_size, method=tf.image.ResizeMethod.BILINEAR)
+ image, self._output_size, method=self._tf_resize_method)
image.set_shape([self._output_size[0], self._output_size[1], 3])
# Normalizes image with mean and std pixel values.
- image = preprocess_ops.normalize_image(image,
- offset=MEAN_RGB,
- scale=STDDEV_RGB)
+ image = preprocess_ops.normalize_image(
+ image, offset=preprocess_ops.MEAN_RGB, scale=preprocess_ops.STDDEV_RGB)
# Convert image to self._dtype.
image = tf.image.convert_image_dtype(image, self._dtype)
return image
+ def parse_train_image(self, decoded_tensors: Dict[str,
+ tf.Tensor]) -> tf.Tensor:
+ """Public interface for parsing image data for training."""
+ return self._parse_train_image(decoded_tensors)
+
@classmethod
def inference_fn(cls,
image: tf.Tensor,
@@ -268,6 +323,6 @@ def inference_fn(cls,
# Normalizes image with mean and std pixel values.
image = preprocess_ops.normalize_image(
- image, offset=MEAN_RGB, scale=STDDEV_RGB)
+ image, offset=preprocess_ops.MEAN_RGB, scale=preprocess_ops.STDDEV_RGB)
image.set_shape(input_image_size + [num_channels])
return image
diff --git a/official/vision/dataloaders/decoder.py b/official/vision/dataloaders/decoder.py
index 821083f0f09..0afcda8e2a0 100644
--- a/official/vision/dataloaders/decoder.py
+++ b/official/vision/dataloaders/decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,11 +17,9 @@
import abc
-class Decoder(object):
+class Decoder(metaclass=abc.ABCMeta):
"""Decodes the raw data into tensors."""
- __metaclass__ = abc.ABCMeta
-
@abc.abstractmethod
def decode(self, serialized_example):
"""Decodes the serialized example into tensors.
diff --git a/official/vision/dataloaders/input_reader.py b/official/vision/dataloaders/input_reader.py
index fba7dc2772f..c2cbf0bc01b 100644
--- a/official/vision/dataloaders/input_reader.py
+++ b/official/vision/dataloaders/input_reader.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,16 +14,74 @@
"""Dataset reader for vision model garden."""
-from typing import Any, Callable, Optional, Tuple
+from typing import Any, Callable, Mapping, Optional, Tuple, Union
-import tensorflow as tf
+from absl import logging
+import tensorflow as tf, tf_keras
from official.core import config_definitions as cfg
from official.core import input_reader
+InputReader = input_reader.InputReader
+
+
+def build_weighted_sampling_combine_fn(
+ weights: Mapping[Any, Any], stop_on_empty_dataset=True
+) -> Callable[[tf.data.Dataset], tf.data.Dataset]:
+ """Builds a combine_fn using weighted sampling."""
+
+ def combine_fn(datasets: Mapping[Any, tf.data.Dataset]) -> tf.data.Dataset:
+ """Combines multiple datasets using weighted sampling."""
+ ds = []
+ ws = []
+ for k, dataset in datasets.items():
+ ds.append(dataset)
+ ws.append(weights[k])
+ return tf.data.Dataset.sample_from_datasets(
+ ds, ws, stop_on_empty_dataset=stop_on_empty_dataset)
+
+ return combine_fn # pyrefly: ignore[bad-return]
+
+
+def create_combine_fn(
+ params: cfg.DataConfig
+) -> Union[None, Callable[[tf.data.Dataset], tf.data.Dataset]]:
+ """Creates and returns a combine_fn for dataset mixing."""
+ if (
+ hasattr(params, 'stop_on_empty_dataset')
+ and params.stop_on_empty_dataset is not None
+ ):
+ stop_on_empty_dataset = params.stop_on_empty_dataset
+ else:
+ stop_on_empty_dataset = True
+
+ if params.weights:
+ # Combine multiple datasets using weighted sampling.
+ if (not isinstance(params.input_path, cfg.base_config.Config) or
+ not isinstance(params.weights, cfg.base_config.Config)):
+ raise ValueError(
+ 'input_path and weights must both be a Config to use weighted '
+ 'sampling.')
+ input_paths = params.input_path.as_dict()
+ weights = params.weights.as_dict()
+ if len(input_paths) != len(weights):
+ raise ValueError(
+ 'The number of input_path and weights must be the same, but got %d '
+ 'input_paths and %d weights.' % (len(input_paths), len(weights)))
+
+ for k in input_paths.keys():
+ if k not in weights:
+ raise ValueError(
+ 'input_path key \'%s\' does not have a corresponding weight.' % k)
+
+ return build_weighted_sampling_combine_fn(weights, stop_on_empty_dataset)
+ return None
+
+
def calculate_batch_sizes(total_batch_size: int,
- pseudo_label_ratio: float) -> Tuple[int, int]:
+ pseudo_label_ratio: float,
+ pseudo_label_batch_size: int = 0) -> Tuple[int, int]:
"""Calculates labeled and pseudo-labeled dataset batch sizes.
Returns (labeled_batch_size, pseudo_labeled_batch_size) given a
@@ -31,26 +89,36 @@ def calculate_batch_sizes(total_batch_size: int,
Args:
total_batch_size: The total batch size for all data.
- pseudo_label_ratio: A non-negative float ratio of pseudo-labeled
- to labeled data in a batch.
+ pseudo_label_ratio: A float ratio of pseudo-labeled to labeled data in a
+ batch. If it is negative, use `pseudo_label_batch_size` instead.
+ pseudo_label_batch_size: The batch size of pseudo-labeled data. It is ignored
+ if `pseudo_label_ratio` is valid. If not, it will be used and it cannot be
+ larger than total global batch size or less than 0 if pseudo_label_ratio is
+ also less than 0.
Returns:
(labeled_batch_size, pseudo_labeled_batch_size) as ints.
Raises:
- ValueError: If total_batch_size is negative.
- ValueError: If pseudo_label_ratio is negative.
+ ValueError: If total_batch_size is negative, or both If pseudo_label_ratio
+ is negative and pseudo-label global_batch_size is negative or larger than
+ total batch size.
"""
if total_batch_size < 0:
raise ValueError('Invalid total_batch_size: {}'.format(total_batch_size))
- if pseudo_label_ratio < 0.0:
- raise ValueError(
- 'Invalid pseudo_label_ratio: {}'.format(pseudo_label_ratio))
-
- ratio_factor = pseudo_label_ratio / (1.0 + pseudo_label_ratio)
- pseudo_labeled_batch_size = int(round(total_batch_size * ratio_factor))
- labeled_batch_size = total_batch_size - pseudo_labeled_batch_size
- return labeled_batch_size, pseudo_labeled_batch_size
+ if pseudo_label_ratio >= 0.0:
+ ratio_factor = pseudo_label_ratio / (1.0 + pseudo_label_ratio)
+ pseudo_label_batch_size = int(total_batch_size * ratio_factor)
+ label_batch_size = total_batch_size - pseudo_label_batch_size
+ else:
+ if pseudo_label_batch_size > total_batch_size or pseudo_label_batch_size < 0:
+ raise ValueError(
+ 'The batch size of pseudo-label dataset should not be larger than '
+ 'total global batch size.')
+ logging.info('data_ratio for pseudo-label dataset is less than 0. '
+ 'Use global_batch_size from pseudo_label data config instead.')
+ label_batch_size = total_batch_size - pseudo_label_batch_size
+ return label_batch_size, pseudo_label_batch_size
class CombinationDatasetInputReader(input_reader.InputReader):
@@ -61,6 +129,7 @@ def __init__(self,
dataset_fn=tf.data.TFRecordDataset,
pseudo_label_dataset_fn=tf.data.TFRecordDataset,
decoder_fn: Optional[Callable[..., Any]] = None,
+ combine_fn: Optional[Callable[..., Any]] = None,
sample_fn: Optional[Callable[..., Any]] = None,
parser_fn: Optional[Callable[..., Any]] = None,
transform_and_batch_fn: Optional[Callable[
@@ -83,6 +152,9 @@ def __init__(self,
files. For example, it can be `tf.data.TFRecordDataset`.
decoder_fn: An optional `callable` that takes the serialized data string
and decodes them into the raw tensor dictionary.
+ combine_fn: An optional `callable` that takes a dictionarty of
+ `tf.data.Dataset` objects as input and outputs a combined dataset. It
+ will be executed after the decoder_fn and before the sample_fn.
sample_fn: An optional `callable` that takes a `tf.data.Dataset` object as
input and outputs the transformed dataset. It performs sampling on the
decoded raw tensors dict before the parser_fn.
@@ -101,17 +173,20 @@ def __init__(self,
Raises:
ValueError: If drop_remainder is False.
"""
- super().__init__(params=params,
- dataset_fn=dataset_fn,
- decoder_fn=decoder_fn,
- sample_fn=sample_fn,
- parser_fn=parser_fn,
- transform_and_batch_fn=transform_and_batch_fn,
- postprocess_fn=postprocess_fn)
+ super().__init__(
+ params=params,
+ dataset_fn=dataset_fn,
+ decoder_fn=decoder_fn,
+ combine_fn=combine_fn,
+ sample_fn=sample_fn,
+ parser_fn=parser_fn,
+ transform_and_batch_fn=transform_and_batch_fn,
+ postprocess_fn=postprocess_fn)
self._pseudo_label_file_pattern = params.pseudo_label_data.input_path
self._pseudo_label_dataset_fn = pseudo_label_dataset_fn
self._pseudo_label_data_ratio = params.pseudo_label_data.data_ratio
+ self._pseudo_label_batch_size = params.pseudo_label_data.global_batch_size
self._pseudo_label_matched_files = input_reader.match_files(
self._pseudo_label_file_pattern)
if not self._drop_remainder:
@@ -125,7 +200,8 @@ def read(
"""Generates a tf.data.Dataset object."""
labeled_batch_size, pl_batch_size = calculate_batch_sizes(
- self._global_batch_size, self._pseudo_label_data_ratio)
+ self._global_batch_size, self._pseudo_label_data_ratio,
+ self._pseudo_label_batch_size)
if not labeled_batch_size and pl_batch_size:
raise ValueError(
@@ -134,24 +210,21 @@ def read(
self._global_batch_size, self._pseudo_label_data_ratio))
def _read_decode_and_parse_dataset(matched_files, dataset_fn, batch_size,
- input_context, tfds_builder):
- dataset = self._read_data_source(matched_files, dataset_fn, input_context,
- tfds_builder)
+ input_context):
+ dataset = self._read_data_source(matched_files, dataset_fn, input_context)
return self._decode_and_parse_dataset(dataset, batch_size, input_context)
labeled_dataset = _read_decode_and_parse_dataset(
matched_files=self._matched_files,
dataset_fn=self._dataset_fn,
batch_size=labeled_batch_size,
- input_context=input_context,
- tfds_builder=self._tfds_builder)
+ input_context=input_context)
pseudo_labeled_dataset = _read_decode_and_parse_dataset(
matched_files=self._pseudo_label_matched_files,
dataset_fn=self._pseudo_label_dataset_fn,
batch_size=pl_batch_size,
- input_context=input_context,
- tfds_builder=False)
+ input_context=input_context)
def concat_fn(d1, d2):
return tf.nest.map_structure(
diff --git a/official/vision/dataloaders/input_reader_factory.py b/official/vision/dataloaders/input_reader_factory.py
index 6af7c1b1491..71b75fd048f 100644
--- a/official/vision/dataloaders/input_reader_factory.py
+++ b/official/vision/dataloaders/input_reader_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/dataloaders/maskrcnn_input.py b/official/vision/dataloaders/maskrcnn_input.py
index ddad3a95846..257f10a7044 100644
--- a/official/vision/dataloaders/maskrcnn_input.py
+++ b/official/vision/dataloaders/maskrcnn_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,13 +14,15 @@
"""Data parser and processing for Mask R-CNN."""
-# Import libraries
+from typing import Optional
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.vision.configs import common
from official.vision.dataloaders import parser
from official.vision.dataloaders import utils
from official.vision.ops import anchor
+from official.vision.ops import augment
from official.vision.ops import box_ops
from official.vision.ops import preprocess_ops
@@ -40,11 +42,14 @@ def __init__(self,
rpn_batch_size_per_im=256,
rpn_fg_fraction=0.5,
aug_rand_hflip=False,
+ aug_rand_vflip=False,
aug_scale_min=1.0,
aug_scale_max=1.0,
+ aug_type: Optional[common.Augmentation] = None,
skip_crowd_during_training=True,
max_num_instances=100,
include_mask=False,
+ outer_boxes_scale=1.0,
mask_crop_size=112,
dtype='float32'):
"""Initializes parameters for parsing annotations in the dataset.
@@ -57,7 +62,7 @@ def __init__(self,
num_scales: `int` number representing intermediate scales added
on each level. For instances, num_scales=2 adds one additional
intermediate anchor scales [2^0, 2^0.5] on each level.
- aspect_ratios: `list` of float numbers representing the aspect raito
+ aspect_ratios: `list` of float numbers representing the aspect ratio
anchors added on each level. The number indicates the ratio of width to
height. For instances, aspect_ratios=[1.0, 2.0, 0.5] adds three anchors
on each scale level.
@@ -67,18 +72,25 @@ def __init__(self,
rpn_unmatched_threshold:
rpn_batch_size_per_im:
rpn_fg_fraction:
- aug_rand_hflip: `bool`, if True, augment training with random
- horizontal flip.
+ aug_rand_hflip: `bool`, if True, augment training with random horizontal
+ flip.
+ aug_rand_vflip: `bool`, if True, augment training with random vertical
+ flip.
aug_scale_min: `float`, the minimum scale applied to `output_size` for
data augmentation during training.
aug_scale_max: `float`, the maximum scale applied to `output_size` for
data augmentation during training.
+ aug_type: An optional Augmentation object with params for AutoAugment.
+ The AutoAug policy should not use rotation/translation/shear.
+ Only in-place augmentations can be used.
skip_crowd_during_training: `bool`, if True, skip annotations labeled with
`is_crowd` equals to 1.
max_num_instances: `int` number of maximum number of instances in an
- image. The groundtruth data will be padded to `max_num_instances`.
- include_mask: a bool to indicate whether parse mask groundtruth.
- mask_crop_size: the size which groundtruth mask is cropped to.
+ image. The ground-truth data will be padded to `max_num_instances`.
+ include_mask: a bool to indicate whether parse mask ground-truth.
+ outer_boxes_scale: a float to scale up the bounding boxes to generate
+ more inclusive masks. The scale is expected to be >=1.0.
+ mask_crop_size: the size which ground-truth mask is cropped to.
dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
"""
@@ -101,11 +113,33 @@ def __init__(self,
# Data augmentation.
self._aug_rand_hflip = aug_rand_hflip
+ self._aug_rand_vflip = aug_rand_vflip
self._aug_scale_min = aug_scale_min
self._aug_scale_max = aug_scale_max
+ if aug_type and aug_type.type:
+ if aug_type.type == 'autoaug':
+ self._augmenter = augment.AutoAugment(
+ augmentation_name=aug_type.autoaug.augmentation_name,
+ cutout_const=aug_type.autoaug.cutout_const,
+ translate_const=aug_type.autoaug.translate_const)
+ elif aug_type.type == 'randaug':
+ self._augmenter = augment.RandAugment(
+ num_layers=aug_type.randaug.num_layers,
+ magnitude=aug_type.randaug.magnitude,
+ cutout_const=aug_type.randaug.cutout_const,
+ translate_const=aug_type.randaug.translate_const,
+ prob_to_apply=aug_type.randaug.prob_to_apply,
+ exclude_ops=aug_type.randaug.exclude_ops)
+ else:
+ raise ValueError('Augmentation policy {} not supported.'.format(
+ aug_type.type))
+ else:
+ self._augmenter = None
+
# Mask.
self._include_mask = include_mask
+ self._outer_boxes_scale = outer_boxes_scale
self._mask_crop_size = mask_crop_size
# Image output dtype.
@@ -137,11 +171,11 @@ def _parse_train_data(self, data):
shape [height_l, width_l, anchors_per_location * 4]. The height_l and
width_l represent the dimension of bounding box regression output at
l-th level.
- gt_boxes: Groundtruth bounding box annotations. The box is represented
+ gt_boxes: Ground-truth bounding box annotations. The box is represented
in [y1, x1, y2, x2] format. The coordinates are w.r.t the scaled
image that is fed to the network. The tennsor is padded with -1 to
the fixed dimension [self._max_num_instances, 4].
- gt_classes: Groundtruth classes annotations. The tennsor is padded
+ gt_classes: Ground-truth classes annotations. The tennsor is padded
with -1 to the fixed dimension [self._max_num_instances].
gt_masks: groundtrugh masks cropped by the bounding box and
resized to a fixed size determined by mask_crop_size.
@@ -163,23 +197,31 @@ def _parse_train_data(self, data):
classes = tf.gather(classes, indices)
boxes = tf.gather(boxes, indices)
if self._include_mask:
- masks = tf.gather(masks, indices)
+ masks = tf.gather(masks, indices) # pyrefly: ignore[unbound-name]
# Gets original image and its size.
image = data['image']
+ if self._augmenter is not None:
+ image = self._augmenter.distort(image)
+
image_shape = tf.shape(image)[0:2]
# Normalizes image with mean and std pixel values.
image = preprocess_ops.normalize_image(image)
# Flips image randomly during training.
- if self._aug_rand_hflip:
- if self._include_mask:
- image, boxes, masks = preprocess_ops.random_horizontal_flip(
- image, boxes, masks)
- else:
- image, boxes, _ = preprocess_ops.random_horizontal_flip(
- image, boxes)
+ image, boxes, masks = preprocess_ops.random_horizontal_flip(
+ image,
+ boxes,
+ masks=None if not self._include_mask else masks, # pyrefly: ignore[unbound-name]
+ prob=tf.where(self._aug_rand_hflip, 0.5, 0.0),
+ )
+ image, boxes, masks = preprocess_ops.random_vertical_flip(
+ image,
+ boxes,
+ masks=None if not self._include_mask else masks,
+ prob=tf.where(self._aug_rand_vflip, 0.5, 0.0),
+ )
# Converts boxes from normalized coordinates to pixel coordinates.
# Now the coordinates of boxes are w.r.t. the original image.
@@ -202,14 +244,17 @@ def _parse_train_data(self, data):
boxes = preprocess_ops.resize_and_crop_boxes(
boxes, image_scale, image_info[1, :], offset)
- # Filters out ground truth boxes that are all zeros.
+ # Filters out ground-truth boxes that are all zeros.
indices = box_ops.get_non_empty_box_indices(boxes)
boxes = tf.gather(boxes, indices)
classes = tf.gather(classes, indices)
if self._include_mask:
+ outer_boxes = box_ops.compute_outer_boxes(boxes, image_info[1, :],
+ self._outer_boxes_scale)
masks = tf.gather(masks, indices)
# Transfer boxes to the original image space and do normalization.
- cropped_boxes = boxes + tf.tile(tf.expand_dims(offset, axis=0), [1, 2])
+ cropped_boxes = outer_boxes + tf.tile(
+ tf.expand_dims(offset, axis=0), [1, 2])
cropped_boxes /= tf.tile(tf.expand_dims(image_scale, axis=0), [1, 2])
cropped_boxes = box_ops.normalize_boxes(cropped_boxes, image_shape)
num_masks = tf.shape(masks)[0]
@@ -242,29 +287,29 @@ def _parse_train_data(self, data):
# Casts input image to self._dtype
image = tf.cast(image, dtype=self._dtype)
+ boxes = preprocess_ops.clip_or_pad_to_fixed_size(
+ boxes, self._max_num_instances, -1)
+ classes = preprocess_ops.clip_or_pad_to_fixed_size(
+ classes, self._max_num_instances, -1)
# Packs labels for model_fn outputs.
labels = {
- 'anchor_boxes':
- anchor_boxes,
- 'image_info':
- image_info,
- 'rpn_score_targets':
- rpn_score_targets,
- 'rpn_box_targets':
- rpn_box_targets,
- 'gt_boxes':
- preprocess_ops.clip_or_pad_to_fixed_size(boxes,
- self._max_num_instances,
- -1),
- 'gt_classes':
- preprocess_ops.clip_or_pad_to_fixed_size(classes,
- self._max_num_instances,
- -1),
+ 'anchor_boxes': anchor_boxes,
+ 'image_info': image_info,
+ 'rpn_score_targets': rpn_score_targets,
+ 'rpn_box_targets': rpn_box_targets,
+ 'gt_boxes': boxes,
+ 'gt_classes': classes,
}
if self._include_mask:
- labels['gt_masks'] = preprocess_ops.clip_or_pad_to_fixed_size(
+ outer_boxes = preprocess_ops.clip_or_pad_to_fixed_size(
+ outer_boxes, self._max_num_instances, -1) # pyrefly: ignore[unbound-name]
+ masks = preprocess_ops.clip_or_pad_to_fixed_size(
masks, self._max_num_instances, -1)
+ labels.update({
+ 'gt_outer_boxes': outer_boxes,
+ 'gt_masks': masks,
+ })
return image, labels
@@ -275,13 +320,13 @@ def _parse_eval_data(self, data):
data: the decoded tensor dictionary from TfExampleDecoder.
Returns:
- A dictionary of {'images': image, 'labels': labels} where
+ A tuple of (image, labels) where
image: image tensor that is preproessed to have normalized value and
dimension [output_size[0], output_size[1], 3]
labels: a dictionary of tensors used for training. The following
describes {key: value} pairs in the dictionary.
source_ids: Source image id. Default value -1 if the source id is
- empty in the groundtruth annotation.
+ empty in the ground-truth annotation.
image_info: a 2D `Tensor` that encodes the information of the image
and the applied preprocessing. It is in the format of
[[original_height, original_width], [scaled_height, scaled_width],
@@ -341,5 +386,19 @@ def _parse_eval_data(self, data):
groundtruths['source_id'])
groundtruths = utils.pad_groundtruths_to_fixed_size(
groundtruths, self._max_num_instances)
+ if self._include_mask:
+ masks = data['groundtruth_instance_masks']
+ masks = tf.image.crop_and_resize(
+ tf.expand_dims(masks, axis=-1),
+ boxes=data['groundtruth_boxes'],
+ box_indices=tf.range(tf.shape(masks)[0], dtype=tf.int32),
+ crop_size=[self._mask_crop_size, self._mask_crop_size],
+ method='bilinear',
+ )
+ masks = tf.squeeze(masks, axis=-1)
+ groundtruths['masks'] = preprocess_ops.clip_or_pad_to_fixed_size(
+ masks, self._max_num_instances, -1
+ )
+
labels['groundtruths'] = groundtruths
return image, labels
diff --git a/official/vision/dataloaders/parser.py b/official/vision/dataloaders/parser.py
index 2a415cb018a..4d42f0af7f5 100644
--- a/official/vision/dataloaders/parser.py
+++ b/official/vision/dataloaders/parser.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/dataloaders/retinanet_input.py b/official/vision/dataloaders/retinanet_input.py
index c3c4f75284e..60f6bc1c894 100644
--- a/official/vision/dataloaders/retinanet_input.py
+++ b/official/vision/dataloaders/retinanet_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,14 +14,16 @@
"""Data parser and processing for RetinaNet.
-Parse image and ground truths in a dataset to training targets and package them
+Parse image and ground-truths in a dataset to training targets and package them
into (image, labels) tuple for RetinaNet.
"""
-# Import libraries
+from typing import Optional
+
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.vision.configs import common as cfg
from official.vision.dataloaders import parser
from official.vision.dataloaders import utils
from official.vision.ops import anchor
@@ -35,15 +37,17 @@ class Parser(parser.Parser):
def __init__(self,
output_size,
- min_level,
+ min_level: int | None,
max_level,
- num_scales,
- aspect_ratios,
- anchor_size,
+ num_scales: int | None,
+ aspect_ratios: list[float] | None,
+ anchor_size: float | None,
match_threshold=0.5,
unmatched_threshold=0.5,
+ box_coder_weights=None,
aug_type=None,
aug_rand_hflip=False,
+ aug_rand_jpeg: cfg.RandJpegQuality | None = None,
aug_scale_min=1.0,
aug_scale_max=1.0,
use_autoaugment=False,
@@ -51,33 +55,49 @@ def __init__(self,
skip_crowd_during_training=True,
max_num_instances=100,
dtype='bfloat16',
- mode=None):
+ resize_first: Optional[bool] = None,
+ mode=None,
+ pad=True,
+ keep_aspect_ratio=True):
"""Initializes parameters for parsing annotations in the dataset.
+ If one provides `input_anchor` when calling `_parse_eval_data()` and
+ `_parse_train_data()`, the `min_level`, `num_scales`, `aspect_ratios`, and
+ `anchor_size` can be `None`.
+
Args:
output_size: `Tensor` or `list` for [height, width] of output image. The
output_size should be divided by the largest feature stride 2^max_level.
min_level: `int` number of minimum level of the output feature pyramid.
+ Can be `None` if `input_anchor` is provided in `_parse_*_data()`.
max_level: `int` number of maximum level of the output feature pyramid.
num_scales: `int` number representing intermediate scales added on each
level. For instances, num_scales=2 adds one additional intermediate
- anchor scales [2^0, 2^0.5] on each level.
- aspect_ratios: `list` of float numbers representing the aspect raito
+ anchor scales [2^0, 2^0.5] on each level. Can be `None` if
+ `input_anchor` is provided in `_parse_*_data()`.
+ aspect_ratios: `list` of float numbers representing the aspect ratio
anchors added on each level. The number indicates the ratio of width to
height. For instances, aspect_ratios=[1.0, 2.0, 0.5] adds three anchors
- on each scale level.
+ on each scale level. Can be `None` if `input_anchor` is provided in
+ `_parse_*_data()`.
anchor_size: `float` number representing the scale of size of the base
- anchor to the feature stride 2^level.
+ anchor to the feature stride 2^level. Can be `None` if `input_anchor` is
+ provided in `_parse_*_data()`.
match_threshold: `float` number between 0 and 1 representing the
lower-bound threshold to assign positive labels for anchors. An anchor
with a score over the threshold is labeled positive.
unmatched_threshold: `float` number between 0 and 1 representing the
upper-bound threshold to assign negative labels for anchors. An anchor
with a score below the threshold is labeled negative.
+ box_coder_weights: Optional `list` of 4 positive floats to scale y, x, h,
+ and w when encoding box coordinates. If set to None, does not perform
+ scaling. For Faster RCNN, the open-source implementation recommends
+ using [10.0, 10.0, 5.0, 5.0].
aug_type: An optional Augmentation object to choose from AutoAugment and
RandAugment.
aug_rand_hflip: `bool`, if True, augment training with random horizontal
flip.
+ aug_rand_jpeg: if not None, apply random JPEG quality change augmentation.
aug_scale_min: `float`, the minimum scale applied to `output_size` for
data augmentation during training.
aug_scale_max: `float`, the maximum scale applied to `output_size` for
@@ -91,8 +111,20 @@ def __init__(self,
max_num_instances: `int` number of maximum number of instances in an
image. The groundtruth data will be padded to `max_num_instances`.
dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
+ resize_first: Optional `bool`, if True, resize the image before the
+ augmentations; computationally more efficient.
mode: a ModeKeys. Specifies if this is training, evaluation, prediction or
- prediction with groundtruths in the outputs.
+ prediction with ground-truths in the outputs.
+ pad: A bool indicating whether to pad the input image to make it
+ size a factor of 2**max_level. The padded size will be the smallest
+ rectangle, such that each dimension is the smallest multiple of
+ 2**max_level which is larger than the desired output size. For example,
+ if desired output size = (320, 320) and max_level = 7, the output padded
+ size = (384, 384). This is necessary when using FPN as it assumes each
+ lower feature map is 2x size of its higher neighbor. Without padding,
+ such relationship may be invalidated. The backbone may produce 5x5 and
+ 2x2 consecutive feature maps, which does not work with FPN.
+ keep_aspect_ratio: `bool`, if True, keep the aspect ratio when resizing.
"""
self._mode = mode
self._max_num_instances = max_num_instances
@@ -107,9 +139,11 @@ def __init__(self,
self._anchor_size = anchor_size
self._match_threshold = match_threshold
self._unmatched_threshold = unmatched_threshold
+ self._box_coder_weights = box_coder_weights
# Data augmentation.
self._aug_rand_hflip = aug_rand_hflip
+ self._aug_rand_jpeg = aug_rand_jpeg
self._aug_scale_min = aug_scale_min
self._aug_scale_max = aug_scale_max
@@ -131,6 +165,13 @@ def __init__(self,
translate_const=aug_type.randaug.translate_const,
prob_to_apply=aug_type.randaug.prob_to_apply,
exclude_ops=aug_type.randaug.exclude_ops)
+ elif aug_type.type == 'ssd_random_crop':
+ logging.info('Using SSD Random Crop.')
+ self._augmenter = augment.SSDRandomCrop(
+ params=aug_type.ssd_random_crop.ssd_random_crop_params,
+ aspect_ratio_range=aug_type.ssd_random_crop.aspect_ratio_range,
+ area_range=aug_type.ssd_random_crop.area_range,
+ )
else:
raise ValueError(f'Augmentation policy {aug_type.type} not supported.')
@@ -141,24 +182,55 @@ def __init__(self,
# Data type.
self._dtype = dtype
- def _parse_train_data(self, data):
+ # Input pipeline optimization.
+ self._resize_first = resize_first
+
+ # Whether to pad image to make its size the smallest factor of 2*max_level.
+ # This is needed when using FPN decoder.
+ self._pad = pad
+
+ self._keep_aspect_ratio = keep_aspect_ratio
+
+ def _resize_and_crop_image_and_boxes(self, image, boxes, pad=True):
+ """Resizes and crops image and boxes, optionally with padding."""
+ # Resizes and crops image.
+ padded_size = None
+ if pad:
+ padded_size = preprocess_ops.compute_padded_size(self._output_size,
+ 2**self._max_level)
+ image, image_info = preprocess_ops.resize_and_crop_image(
+ image,
+ self._output_size,
+ padded_size=padded_size,
+ aug_scale_min=self._aug_scale_min,
+ aug_scale_max=self._aug_scale_max,
+ keep_aspect_ratio=self._keep_aspect_ratio,
+ )
+
+ # Resizes and crops boxes.
+ image_scale = image_info[2, :]
+ offset = image_info[3, :]
+ boxes = preprocess_ops.resize_and_crop_boxes(boxes, image_scale,
+ image_info[1, :], offset)
+ return image, boxes, image_info
+
+ def _parse_train_data(self, data, anchor_labeler=None, input_anchor=None):
"""Parses data for training and evaluation."""
classes = data['groundtruth_classes']
boxes = data['groundtruth_boxes']
# If not empty, `attributes` is a dict of (name, ground_truth) pairs.
- # `ground_gruth` of attributes is assumed in shape [N, attribute_size].
- # TODO(xianzhi): support parsing attributes weights.
+ # `ground_truth` of attributes is assumed in shape [N, attribute_size].
attributes = data.get('groundtruth_attributes', {})
is_crowds = data['groundtruth_is_crowd']
# Skips annotations with `is_crowd` = True.
if self._skip_crowd_during_training:
- num_groundtrtuhs = tf.shape(input=classes)[0]
- with tf.control_dependencies([num_groundtrtuhs, is_crowds]):
+ num_groundtruths = tf.shape(input=classes)[0]
+ with tf.control_dependencies([num_groundtruths, is_crowds]):
indices = tf.cond(
pred=tf.greater(tf.size(input=is_crowds), 0),
true_fn=lambda: tf.where(tf.logical_not(is_crowds))[:, 0],
- false_fn=lambda: tf.cast(tf.range(num_groundtrtuhs), tf.int64))
+ false_fn=lambda: tf.cast(tf.range(num_groundtruths), tf.int64))
classes = tf.gather(classes, indices)
boxes = tf.gather(boxes, indices)
for k, v in attributes.items():
@@ -166,10 +238,37 @@ def _parse_train_data(self, data):
# Gets original image.
image = data['image']
+ image_size = tf.cast(tf.shape(image)[0:2], tf.float32)
+
+ less_output_pixels = (
+ self._output_size[0] * self._output_size[1]
+ ) < image_size[0] * image_size[1]
+
+ # Resizing first can reduce augmentation computation if the original image
+ # has more pixels than the desired output image.
+ # There might be a smarter threshold to compute less_output_pixels as
+ # we keep the padding to the very end, i.e., a resized image likely has less
+ # pixels than self._output_size[0] * self._output_size[1].
+ resize_first = self._resize_first and less_output_pixels
+ if resize_first:
+ image, boxes, image_info = self._resize_and_crop_image_and_boxes(
+ image, boxes, pad=False
+ )
+ image = tf.cast(image, dtype=tf.uint8)
# Apply autoaug or randaug.
if self._augmenter is not None:
image, boxes = self._augmenter.distort_with_boxes(image, boxes)
+
+ # Apply random jpeg quality change.
+ if self._aug_rand_jpeg is not None:
+ image = preprocess_ops.random_jpeg_quality(
+ image,
+ min_quality=self._aug_rand_jpeg.min_quality,
+ max_quality=self._aug_rand_jpeg.max_quality,
+ prob_to_apply=self._aug_rand_jpeg.prob_to_apply,
+ )
+
image_shape = tf.shape(input=image)[0:2]
# Normalizes image with mean and std pixel values.
@@ -182,22 +281,26 @@ def _parse_train_data(self, data):
# Converts boxes from normalized coordinates to pixel coordinates.
boxes = box_ops.denormalize_boxes(boxes, image_shape)
- # Resizes and crops image.
- image, image_info = preprocess_ops.resize_and_crop_image(
- image,
- self._output_size,
- padded_size=preprocess_ops.compute_padded_size(self._output_size,
- 2**self._max_level),
- aug_scale_min=self._aug_scale_min,
- aug_scale_max=self._aug_scale_max)
+ if self._pad:
+ padded_size = preprocess_ops.compute_padded_size(
+ self._output_size, 2**self._max_level
+ )
+ else:
+ padded_size = self._output_size
+
+ if not resize_first:
+ image, boxes, image_info = (
+ self._resize_and_crop_image_and_boxes(image, boxes, pad=self._pad)
+ )
+
+ image = tf.image.pad_to_bounding_box(
+ image, 0, 0, padded_size[0], padded_size[1]
+ )
+ image = tf.ensure_shape(image, padded_size + [3])
+
image_height, image_width, _ = image.get_shape().as_list()
- # Resizes and crops boxes.
- image_scale = image_info[2, :]
- offset = image_info[3, :]
- boxes = preprocess_ops.resize_and_crop_boxes(boxes, image_scale,
- image_info[1, :], offset)
- # Filters out ground truth boxes that are all zeros.
+ # Filters out ground-truth boxes that are all zeros.
indices = box_ops.get_non_empty_box_indices(boxes)
boxes = tf.gather(boxes, indices)
classes = tf.gather(classes, indices)
@@ -205,15 +308,22 @@ def _parse_train_data(self, data):
attributes[k] = tf.gather(v, indices)
# Assigns anchors.
- input_anchor = anchor.build_anchor_generator(
- min_level=self._min_level,
- max_level=self._max_level,
- num_scales=self._num_scales,
- aspect_ratios=self._aspect_ratios,
- anchor_size=self._anchor_size)
+ if input_anchor is None:
+ input_anchor = anchor.build_anchor_generator(
+ min_level=self._min_level,
+ max_level=self._max_level,
+ num_scales=self._num_scales,
+ aspect_ratios=self._aspect_ratios,
+ anchor_size=self._anchor_size,
+ )
+
anchor_boxes = input_anchor(image_size=(image_height, image_width))
- anchor_labeler = anchor.AnchorLabeler(self._match_threshold,
- self._unmatched_threshold)
+ if anchor_labeler is None:
+ anchor_labeler = anchor.AnchorLabeler(
+ match_threshold=self._match_threshold,
+ unmatched_threshold=self._unmatched_threshold,
+ box_coder_weights=self._box_coder_weights,
+ )
(cls_targets, box_targets, att_targets, cls_weights,
box_weights) = anchor_labeler.label_anchors(
anchor_boxes, boxes, tf.expand_dims(classes, axis=1), attributes)
@@ -228,20 +338,19 @@ def _parse_train_data(self, data):
'anchor_boxes': anchor_boxes,
'cls_weights': cls_weights,
'box_weights': box_weights,
- 'image_info': image_info,
+ 'image_info': image_info, # pyrefly: ignore[unbound-name]
}
if att_targets:
labels['attribute_targets'] = att_targets
return image, labels
- def _parse_eval_data(self, data):
+ def _parse_eval_data(self, data, anchor_labeler=None, input_anchor=None):
"""Parses data for training and evaluation."""
- groundtruths = {}
+
classes = data['groundtruth_classes']
boxes = data['groundtruth_boxes']
# If not empty, `attributes` is a dict of (name, ground_truth) pairs.
- # `ground_gruth` of attributes is assumed in shape [N, attribute_size].
- # TODO(xianzhi): support parsing attributes weights.
+ # `ground_truth` of attributes is assumed in shape [N, attribute_size].
attributes = data.get('groundtruth_attributes', {})
# Gets original image and its size.
@@ -255,13 +364,22 @@ def _parse_eval_data(self, data):
boxes = box_ops.denormalize_boxes(boxes, image_shape)
# Resizes and crops image.
+ if self._pad:
+ padded_size = preprocess_ops.compute_padded_size(
+ self._output_size, 2**self._max_level
+ )
+ else:
+ padded_size = self._output_size
+
image, image_info = preprocess_ops.resize_and_crop_image(
image,
self._output_size,
- padded_size=preprocess_ops.compute_padded_size(self._output_size,
- 2**self._max_level),
+ padded_size=padded_size,
aug_scale_min=1.0,
- aug_scale_max=1.0)
+ aug_scale_max=1.0,
+ keep_aspect_ratio=self._keep_aspect_ratio,
+ )
+ image = tf.ensure_shape(image, padded_size + [3])
image_height, image_width, _ = image.get_shape().as_list()
# Resizes and crops boxes.
@@ -269,7 +387,7 @@ def _parse_eval_data(self, data):
offset = image_info[3, :]
boxes = preprocess_ops.resize_and_crop_boxes(boxes, image_scale,
image_info[1, :], offset)
- # Filters out ground truth boxes that are all zeros.
+ # Filters out ground-truth boxes that are all zeros.
indices = box_ops.get_non_empty_box_indices(boxes)
boxes = tf.gather(boxes, indices)
classes = tf.gather(classes, indices)
@@ -277,15 +395,22 @@ def _parse_eval_data(self, data):
attributes[k] = tf.gather(v, indices)
# Assigns anchors.
- input_anchor = anchor.build_anchor_generator(
- min_level=self._min_level,
- max_level=self._max_level,
- num_scales=self._num_scales,
- aspect_ratios=self._aspect_ratios,
- anchor_size=self._anchor_size)
+ if input_anchor is None:
+ input_anchor = anchor.build_anchor_generator(
+ min_level=self._min_level,
+ max_level=self._max_level,
+ num_scales=self._num_scales,
+ aspect_ratios=self._aspect_ratios,
+ anchor_size=self._anchor_size,
+ )
+
anchor_boxes = input_anchor(image_size=(image_height, image_width))
- anchor_labeler = anchor.AnchorLabeler(self._match_threshold,
- self._unmatched_threshold)
+ if anchor_labeler is None:
+ anchor_labeler = anchor.AnchorLabeler(
+ match_threshold=self._match_threshold,
+ unmatched_threshold=self._unmatched_threshold,
+ box_coder_weights=self._box_coder_weights,
+ )
(cls_targets, box_targets, att_targets, cls_weights,
box_weights) = anchor_labeler.label_anchors(
anchor_boxes, boxes, tf.expand_dims(classes, axis=1), attributes)
@@ -293,7 +418,7 @@ def _parse_eval_data(self, data):
# Casts input image to desired data type.
image = tf.cast(image, dtype=self._dtype)
- # Sets up groundtruth data for evaluation.
+ # Sets up ground-truth data for evaluation.
groundtruths = {
'source_id': data['source_id'],
'height': data['height'],
diff --git a/official/vision/dataloaders/segmentation_input.py b/official/vision/dataloaders/segmentation_input.py
index ae0d0199bed..27aad100b8f 100644
--- a/official/vision/dataloaders/segmentation_input.py
+++ b/official/vision/dataloaders/segmentation_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,44 +14,65 @@
"""Data parser and processing for segmentation datasets."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+from official.vision.configs import semantic_segmentation as config_lib
from official.vision.dataloaders import decoder
from official.vision.dataloaders import parser
+from official.vision.dataloaders import utils
from official.vision.ops import preprocess_ops
class Decoder(decoder.Decoder):
"""A tf.Example decoder for segmentation task."""
- def __init__(self):
+ def __init__(
+ self,
+ image_feature=config_lib.DenseFeatureConfig(),
+ additional_dense_features=None,
+ ):
self._keys_to_features = {
'image/encoded': tf.io.FixedLenFeature((), tf.string, default_value=''),
'image/height': tf.io.FixedLenFeature((), tf.int64, default_value=0),
'image/width': tf.io.FixedLenFeature((), tf.int64, default_value=0),
- 'image/segmentation/class/encoded':
- tf.io.FixedLenFeature((), tf.string, default_value='')
+ 'image/segmentation/class/encoded': tf.io.FixedLenFeature(
+ (), tf.string, default_value=''
+ ),
+ image_feature.feature_name: tf.io.FixedLenFeature(
+ (), tf.string, default_value=''
+ ),
}
+ if additional_dense_features:
+ for feature in additional_dense_features:
+ self._keys_to_features[feature.feature_name] = tf.io.FixedLenFeature(
+ (), tf.string, default_value=''
+ )
def decode(self, serialized_example):
return tf.io.parse_single_example(
- serialized_example, self._keys_to_features)
+ serialized_example, self._keys_to_features
+ )
class Parser(parser.Parser):
- """Parser to parse an image and its annotations into a dictionary of tensors.
- """
-
- def __init__(self,
- output_size,
- crop_size=None,
- resize_eval_groundtruth=True,
- groundtruth_padded_size=None,
- ignore_label=255,
- aug_rand_hflip=False,
- preserve_aspect_ratio=True,
- aug_scale_min=1.0,
- aug_scale_max=1.0,
- dtype='float32'):
+ """Parser to parse an image and its annotations into a dictionary of tensors."""
+
+ def __init__(
+ self,
+ output_size,
+ crop_size=None,
+ resize_eval_groundtruth=True,
+ gt_is_matting_map=False,
+ groundtruth_padded_size=None,
+ ignore_label=255,
+ aug_rand_hflip=False,
+ preserve_aspect_ratio=True,
+ aug_scale_min=1.0,
+ aug_scale_max=1.0,
+ dtype='float32',
+ image_feature=config_lib.DenseFeatureConfig(),
+ additional_dense_features=None,
+ centered_crop=False,
+ ):
"""Initializes parameters for parsing annotations in the dataset.
Args:
@@ -61,15 +82,18 @@ def __init__(self,
specified a training crop of size crop_size is returned. This is useful
for cropping original images during training while evaluating on
original image sizes.
- resize_eval_groundtruth: `bool`, if True, eval groundtruth masks are
+ resize_eval_groundtruth: `bool`, if True, eval ground-truth masks are
resized to output_size.
+ gt_is_matting_map: `bool`, if True, the expected mask is in the range
+ between 0 and 255. The parser will normalize the value of the mask into
+ the range between 0 and 1.
groundtruth_padded_size: `Tensor` or `list` for [height, width]. When
- resize_eval_groundtruth is set to False, the groundtruth masks are
+ resize_eval_groundtruth is set to False, the ground-truth masks are
padded to this size.
ignore_label: `int` the pixel with ignore label will not used for training
and evaluation.
- aug_rand_hflip: `bool`, if True, augment training with random
- horizontal flip.
+ aug_rand_hflip: `bool`, if True, augment training with random horizontal
+ flip.
preserve_aspect_ratio: `bool`, if True, the aspect ratio is preserved,
otherwise, the image is resized to output_size.
aug_scale_min: `float`, the minimum scale applied to `output_size` for
@@ -77,13 +101,24 @@ def __init__(self,
aug_scale_max: `float`, the maximum scale applied to `output_size` for
data augmentation during training.
dtype: `str`, data type. One of {`bfloat16`, `float32`, `float16`}.
+ image_feature: the config for the image input (usually RGB). Defaults to
+ the config for a 3-channel image with key = `image/encoded` and ImageNet
+ dataset mean/stddev.
+ additional_dense_features: `list` of DenseFeatureConfig for additional
+ dense features.
+ centered_crop: If `centered_crop` is set to True, then resized crop (if
+ smaller than padded size) is place in the center of the image. Default
+ behaviour is to place it at left top corner.
"""
self._output_size = output_size
self._crop_size = crop_size
self._resize_eval_groundtruth = resize_eval_groundtruth
if (not resize_eval_groundtruth) and (groundtruth_padded_size is None):
- raise ValueError('groundtruth_padded_size ([height, width]) needs to be'
- 'specified when resize_eval_groundtruth is False.')
+ raise ValueError(
+ 'groundtruth_padded_size ([height, width]) needs to be'
+ 'specified when resize_eval_groundtruth is False.'
+ )
+ self._gt_is_matting_map = gt_is_matting_map
self._groundtruth_padded_size = groundtruth_padded_size
self._ignore_label = ignore_label
self._preserve_aspect_ratio = preserve_aspect_ratio
@@ -96,32 +131,89 @@ def __init__(self,
# dtype.
self._dtype = dtype
+ self._image_feature = image_feature
+ self._additional_dense_features = additional_dense_features
+ self._centered_crop = centered_crop
+ if self._centered_crop and not self._resize_eval_groundtruth:
+ raise ValueError(
+ 'centered_crop is only supported when resize_eval_groundtruth is'
+ ' True.'
+ )
+
def _prepare_image_and_label(self, data):
"""Prepare normalized image and label."""
- image = tf.io.decode_image(data['image/encoded'], channels=3)
- label = tf.io.decode_image(data['image/segmentation/class/encoded'],
- channels=1)
height = data['image/height']
width = data['image/width']
- image = tf.reshape(image, (height, width, 3))
+ label = tf.io.decode_image(
+ data['image/segmentation/class/encoded'], channels=1
+ )
label = tf.reshape(label, (1, height, width))
label = tf.cast(label, tf.float32)
- # Normalizes image with mean and std pixel values.
- image = preprocess_ops.normalize_image(image)
+
+ image = tf.io.decode_image(
+ data[self._image_feature.feature_name],
+ channels=self._image_feature.num_channels,
+ dtype=tf.uint8,
+ )
+ image = tf.reshape(image, (height, width, self._image_feature.num_channels))
+ # Normalizes the image feature.
+ # The mean and stddev values are divided by 255 to ensure correct
+ # normalization, as the input `uint8` image is automatically converted to
+ # `float32` and rescaled to values in the range [0, 1] before the
+ # normalization happens (as a pre-processing step). So, we re-scale the
+ # mean and stddev values to the range [0, 1] beforehand.
+ # See `preprocess_ops.normalize_image` for details on the expected ranges
+ # for the image mean (`offset`) and stddev (`scale`).
+ image = preprocess_ops.normalize_image(
+ image,
+ [mean / 255.0 for mean in self._image_feature.mean],
+ [stddev / 255.0 for stddev in self._image_feature.stddev],
+ )
+
+ if self._additional_dense_features:
+ input_list = [image]
+ for feature_cfg in self._additional_dense_features:
+ feature = tf.io.decode_image(
+ data[feature_cfg.feature_name],
+ channels=feature_cfg.num_channels,
+ dtype=tf.uint8,
+ )
+ feature = tf.reshape(feature, (height, width, feature_cfg.num_channels))
+ feature = preprocess_ops.normalize_image(
+ feature,
+ [mean / 255.0 for mean in feature_cfg.mean],
+ [stddev / 255.0 for stddev in feature_cfg.stddev],
+ )
+ input_list.append(feature)
+ concat_input = tf.concat(input_list, axis=2)
+ else:
+ concat_input = image
if not self._preserve_aspect_ratio:
label = tf.reshape(label, [data['image/height'], data['image/width'], 1])
- image = tf.image.resize(image, self._output_size, method='bilinear')
+ concat_input = tf.image.resize(
+ concat_input, self._output_size, method='bilinear'
+ )
label = tf.image.resize(label, self._output_size, method='nearest')
label = tf.reshape(label[:, :, -1], [1] + self._output_size)
- return image, label
+ return concat_input, label
def _parse_train_data(self, data):
"""Parses data for training and evaluation."""
image, label = self._prepare_image_and_label(data)
+ # Normalize the label into the range of 0 and 1 for matting ground-truth.
+ # Note that the input ground-truth labels must be 0 to 255, and do not
+ # contain ignore_label. For gt_is_matting_map case, ignore_label is only
+ # used for padding the labels.
+ if self._gt_is_matting_map:
+ scale = tf.constant(255.0, dtype=tf.float32)
+ scale = tf.expand_dims(scale, axis=0)
+ scale = tf.expand_dims(scale, axis=0)
+ label = tf.cast(label, tf.float32) / scale
+
if self._crop_size:
label = tf.reshape(label, [data['image/height'], data['image/width'], 1])
@@ -132,15 +224,17 @@ def _parse_train_data(self, data):
label = tf.image.resize(label, self._output_size, method='nearest')
image_mask = tf.concat([image, label], axis=2)
- image_mask_crop = tf.image.random_crop(image_mask,
- self._crop_size + [4])
+ image_mask_crop = tf.image.random_crop(
+ image_mask, self._crop_size + [tf.shape(image_mask)[-1]]
+ )
image = image_mask_crop[:, :, :-1]
label = tf.reshape(image_mask_crop[:, :, -1], [1] + self._crop_size)
# Flips image randomly during training.
if self._aug_rand_hflip:
image, _, label = preprocess_ops.random_horizontal_flip(
- image, masks=label)
+ image, masks=label
+ )
train_image_size = self._crop_size if self._crop_size else self._output_size
# Resizes and crops image.
@@ -149,7 +243,9 @@ def _parse_train_data(self, data):
train_image_size,
train_image_size,
aug_scale_min=self._aug_scale_min,
- aug_scale_max=self._aug_scale_max)
+ aug_scale_max=self._aug_scale_max,
+ centered_crop=self._centered_crop,
+ )
# Resizes and crops boxes.
image_scale = image_info[2, :]
@@ -160,12 +256,19 @@ def _parse_train_data(self, data):
label += 1
label = tf.expand_dims(label, axis=3)
label = preprocess_ops.resize_and_crop_masks(
- label, image_scale, train_image_size, offset)
+ label,
+ image_scale,
+ train_image_size,
+ offset,
+ centered_crop=self._centered_crop,
+ )
label -= 1
- label = tf.where(tf.equal(label, -1),
- self._ignore_label * tf.ones_like(label), label)
+ label = tf.where(
+ tf.equal(label, -1), self._ignore_label * tf.ones_like(label), label
+ )
label = tf.squeeze(label, axis=0)
valid_mask = tf.not_equal(label, self._ignore_label)
+
labels = {
'masks': label,
'valid_masks': valid_mask,
@@ -180,36 +283,70 @@ def _parse_train_data(self, data):
def _parse_eval_data(self, data):
"""Parses data for training and evaluation."""
image, label = self._prepare_image_and_label(data)
+
+ # Binarize mask if ground-truth is a matting map
+ if self._gt_is_matting_map:
+ label = tf.divide(tf.cast(label, dtype=tf.float32), 255.0)
+ label = utils.binarize_matting_map(label)
+
# The label is first offset by +1 and then padded with 0.
label += 1
label = tf.expand_dims(label, axis=3)
# Resizes and crops image.
image, image_info = preprocess_ops.resize_and_crop_image(
- image, self._output_size, self._output_size)
+ image,
+ self._output_size,
+ self._output_size,
+ centered_crop=self._centered_crop,
+ )
if self._resize_eval_groundtruth:
# Resizes eval masks to match input image sizes. In that case, mean IoU
# is computed on output_size not the original size of the images.
image_scale = image_info[2, :]
offset = image_info[3, :]
- label = preprocess_ops.resize_and_crop_masks(label, image_scale,
- self._output_size, offset)
+ label = preprocess_ops.resize_and_crop_masks(
+ label,
+ image_scale,
+ self._output_size,
+ offset,
+ centered_crop=self._centered_crop,
+ )
else:
- label = tf.image.pad_to_bounding_box(
- label, 0, 0, self._groundtruth_padded_size[0],
- self._groundtruth_padded_size[1])
+ if self._centered_crop:
+ label_size = tf.cast(tf.shape(label)[0:2], tf.int32)
+ label = tf.image.pad_to_bounding_box(
+ label,
+ tf.maximum(
+ (self._groundtruth_padded_size[0] - label_size[0]) // 2, 0 # pyrefly: ignore[unsupported-operation]
+ ),
+ tf.maximum(
+ (self._groundtruth_padded_size[1] - label_size[1]) // 2, 0 # pyrefly: ignore[unsupported-operation]
+ ),
+ self._groundtruth_padded_size[0], # pyrefly: ignore[unsupported-operation]
+ self._groundtruth_padded_size[1], # pyrefly: ignore[unsupported-operation]
+ )
+ else:
+ label = tf.image.pad_to_bounding_box(
+ label,
+ 0,
+ 0,
+ self._groundtruth_padded_size[0], # pyrefly: ignore[unsupported-operation]
+ self._groundtruth_padded_size[1], # pyrefly: ignore[unsupported-operation]
+ )
label -= 1
- label = tf.where(tf.equal(label, -1),
- self._ignore_label * tf.ones_like(label), label)
+ label = tf.where(
+ tf.equal(label, -1), self._ignore_label * tf.ones_like(label), label
+ )
label = tf.squeeze(label, axis=0)
valid_mask = tf.not_equal(label, self._ignore_label)
labels = {
'masks': label,
'valid_masks': valid_mask,
- 'image_info': image_info
+ 'image_info': image_info,
}
# Cast image as self._dtype
diff --git a/official/vision/dataloaders/tf_example_decoder.py b/official/vision/dataloaders/tf_example_decoder.py
index 9ec0c91688b..6d98a9c417a 100644
--- a/official/vision/dataloaders/tf_example_decoder.py
+++ b/official/vision/dataloaders/tf_example_decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
A decoder to decode string tensors containing serialized tensorflow.Example
protos for object detection.
"""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import decoder
@@ -31,16 +31,19 @@ def _generate_source_id(image_bytes):
class TfExampleDecoder(decoder.Decoder):
"""Tensorflow Example proto decoder."""
- def __init__(self,
- include_mask=False,
- regenerate_source_id=False,
- mask_binarize_threshold=None):
+ def __init__(
+ self,
+ include_mask=False,
+ regenerate_source_id=False,
+ mask_binarize_threshold=None,
+ attribute_names=None,
+ ):
self._include_mask = include_mask
self._regenerate_source_id = regenerate_source_id
self._keys_to_features = {
'image/encoded': tf.io.FixedLenFeature((), tf.string),
- 'image/height': tf.io.FixedLenFeature((), tf.int64),
- 'image/width': tf.io.FixedLenFeature((), tf.int64),
+ 'image/height': tf.io.FixedLenFeature((), tf.int64, -1),
+ 'image/width': tf.io.FixedLenFeature((), tf.int64, -1),
'image/object/bbox/xmin': tf.io.VarLenFeature(tf.float32),
'image/object/bbox/xmax': tf.io.VarLenFeature(tf.float32),
'image/object/bbox/ymin': tf.io.VarLenFeature(tf.float32),
@@ -49,6 +52,13 @@ def __init__(self,
'image/object/area': tf.io.VarLenFeature(tf.float32),
'image/object/is_crowd': tf.io.VarLenFeature(tf.int64),
}
+ attribute_names = attribute_names or []
+ for attr_name in attribute_names:
+ self._keys_to_features[f'image/object/attribute/{attr_name}'] = (
+ tf.io.VarLenFeature(tf.int64)
+ )
+ self._attribute_names = attribute_names
+
self._mask_binarize_threshold = mask_binarize_threshold
if include_mask:
self._keys_to_features.update({
@@ -76,6 +86,14 @@ def _decode_boxes(self, parsed_tensors):
def _decode_classes(self, parsed_tensors):
return parsed_tensors['image/object/class/label']
+ def _decode_attributes(self, parsed_tensors):
+ attribute_dict = dict()
+ for attr_name in self._attribute_names:
+ attr_array = parsed_tensors[f'image/object/attribute/{attr_name}']
+ # TODO(b/269654135): Support decoding of fully 2D attributes.
+ attribute_dict[attr_name] = tf.expand_dims(attr_array, -1)
+ return attribute_dict
+
def _decode_areas(self, parsed_tensors):
xmin = parsed_tensors['image/object/bbox/xmin']
xmax = parsed_tensors['image/object/bbox/xmax']
@@ -148,6 +166,20 @@ def decode(self, serialized_example):
boxes = self._decode_boxes(parsed_tensors)
classes = self._decode_classes(parsed_tensors)
areas = self._decode_areas(parsed_tensors)
+
+ attributes = self._decode_attributes(parsed_tensors)
+
+ decode_image_shape = tf.logical_or(
+ tf.equal(parsed_tensors['image/height'], -1),
+ tf.equal(parsed_tensors['image/width'], -1))
+ image_shape = tf.cast(tf.shape(image), dtype=tf.int64)
+
+ parsed_tensors['image/height'] = tf.where(decode_image_shape,
+ image_shape[0],
+ parsed_tensors['image/height'])
+ parsed_tensors['image/width'] = tf.where(decode_image_shape, image_shape[1],
+ parsed_tensors['image/width'])
+
is_crowds = tf.cond(
tf.greater(tf.shape(parsed_tensors['image/object/is_crowd'])[0], 0),
lambda: tf.cast(parsed_tensors['image/object/is_crowd'], dtype=tf.bool),
@@ -168,9 +200,11 @@ def decode(self, serialized_example):
'groundtruth_area': areas,
'groundtruth_boxes': boxes,
}
+ if self._attribute_names:
+ decoded_tensors.update({'groundtruth_attributes': attributes})
if self._include_mask:
decoded_tensors.update({
- 'groundtruth_instance_masks': masks,
+ 'groundtruth_instance_masks': masks, # pyrefly: ignore[unbound-name]
'groundtruth_instance_masks_png': parsed_tensors['image/object/mask'],
})
return decoded_tensors
diff --git a/official/vision/dataloaders/tf_example_decoder_test.py b/official/vision/dataloaders/tf_example_decoder_test.py
index 935baeb2dc2..299f1068a61 100644
--- a/official/vision/dataloaders/tf_example_decoder_test.py
+++ b/official/vision/dataloaders/tf_example_decoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,9 @@
"""Tests for tf_example_decoder.py."""
-# Import libraries
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import tf_example_decoder
from official.vision.dataloaders import tfexample_utils
@@ -26,18 +25,21 @@
class TfExampleDecoderTest(tf.test.TestCase, parameterized.TestCase):
@parameterized.parameters(
- (100, 100, 0, True),
- (100, 100, 1, True),
- (100, 100, 2, True),
- (100, 100, 0, False),
- (100, 100, 1, False),
- (100, 100, 2, False),
+ (100, 100, 0, True, True),
+ (100, 100, 1, True, True),
+ (100, 100, 2, True, True),
+ (100, 100, 0, False, True),
+ (100, 100, 1, False, True),
+ (100, 100, 2, False, True),
+ (100, 100, 0, True, False),
+ (100, 100, 1, True, False),
+ (100, 100, 2, True, False),
+ (100, 100, 0, False, False),
+ (100, 100, 1, False, False),
+ (100, 100, 2, False, False),
)
- def test_result_shape(self,
- image_height,
- image_width,
- num_instances,
- regenerate_source_id):
+ def test_result_shape(self, image_height, image_width, num_instances,
+ regenerate_source_id, fill_image_size):
decoder = tf_example_decoder.TfExampleDecoder(
include_mask=True, regenerate_source_id=regenerate_source_id)
@@ -45,7 +47,9 @@ def test_result_shape(self,
image_height=image_height,
image_width=image_width,
image_channel=3,
- num_instances=num_instances).SerializeToString()
+ num_instances=num_instances,
+ fill_image_size=fill_image_size,
+ ).SerializeToString()
decoded_tensors = decoder.decode(
tf.convert_to_tensor(value=serialized_example))
@@ -72,7 +76,9 @@ def test_result_shape(self,
(num_instances,), results['groundtruth_instance_masks_png'].shape)
def test_result_content(self):
- decoder = tf_example_decoder.TfExampleDecoder(include_mask=True)
+ decoder = tf_example_decoder.TfExampleDecoder(
+ include_mask=True, attribute_names=['attr1', 'attr2']
+ )
image_content = [[[0, 0, 0], [0, 0, 0], [0, 0, 0], [0, 0, 0]],
[[0, 0, 0], [255, 255, 255], [255, 255, 255], [0, 0, 0]],
@@ -87,6 +93,8 @@ def test_result_content(self):
ymins = [0, 0]
ymaxs = [0.5, 1.0]
labels = [3, 1]
+ attr1 = np.array([[0], [2]])
+ attr2 = np.array([[1], [3]])
areas = [
0.25 * image_height * image_width, 0.75 * image_height * image_width
]
@@ -106,32 +114,53 @@ def test_result_content(self):
serialized_example = tf.train.Example(
features=tf.train.Features(
feature={
- 'image/encoded': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[image]))),
- 'image/source_id': (tf.train.Feature(
+ 'image/encoded': tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=[image])
+ ),
+ 'image/source_id': tf.train.Feature(
bytes_list=tf.train.BytesList(
- value=[tfexample_utils.DUMP_SOURCE_ID]))),
- 'image/height': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=[image_height]))),
- 'image/width': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=[image_width]))),
- 'image/object/bbox/xmin': (tf.train.Feature(
- float_list=tf.train.FloatList(value=xmins))),
- 'image/object/bbox/xmax': (tf.train.Feature(
- float_list=tf.train.FloatList(value=xmaxs))),
- 'image/object/bbox/ymin': (tf.train.Feature(
- float_list=tf.train.FloatList(value=ymins))),
- 'image/object/bbox/ymax': (tf.train.Feature(
- float_list=tf.train.FloatList(value=ymaxs))),
- 'image/object/class/label': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=labels))),
- 'image/object/is_crowd': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=is_crowds))),
- 'image/object/area': (tf.train.Feature(
- float_list=tf.train.FloatList(value=areas))),
- 'image/object/mask': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=masks))),
- })).SerializeToString()
+ value=[tfexample_utils.DUMP_SOURCE_ID]
+ )
+ ),
+ 'image/height': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[image_height])
+ ),
+ 'image/width': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=[image_width])
+ ),
+ 'image/object/bbox/xmin': tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmins)
+ ),
+ 'image/object/bbox/xmax': tf.train.Feature(
+ float_list=tf.train.FloatList(value=xmaxs)
+ ),
+ 'image/object/bbox/ymin': tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymins)
+ ),
+ 'image/object/bbox/ymax': tf.train.Feature(
+ float_list=tf.train.FloatList(value=ymaxs)
+ ),
+ 'image/object/class/label': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=labels)
+ ),
+ 'image/object/is_crowd': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=is_crowds)
+ ),
+ 'image/object/area': tf.train.Feature(
+ float_list=tf.train.FloatList(value=areas)
+ ),
+ 'image/object/mask': tf.train.Feature(
+ bytes_list=tf.train.BytesList(value=masks)
+ ),
+ 'image/object/attribute/attr1': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=attr1.flatten())
+ ),
+ 'image/object/attribute/attr2': tf.train.Feature(
+ int64_list=tf.train.Int64List(value=attr2.flatten())
+ ),
+ }
+ )
+ ).SerializeToString()
decoded_tensors = decoder.decode(
tf.convert_to_tensor(value=serialized_example))
@@ -158,8 +187,13 @@ def test_result_content(self):
(num_instances,), results['groundtruth_instance_masks_png'].shape)
self.assertAllEqual(
[3, 1], results['groundtruth_classes'])
- self.assertAllEqual(
- [True, False], results['groundtruth_is_crowd'])
+ np.testing.assert_array_equal(
+ attr1, results['groundtruth_attributes']['attr1']
+ )
+ np.testing.assert_array_equal(
+ attr2, results['groundtruth_attributes']['attr2']
+ )
+ self.assertAllEqual([True, False], results['groundtruth_is_crowd'])
self.assertNDArrayNear(
[0.25 * image_height * image_width, 0.75 * image_height * image_width],
results['groundtruth_area'], 1e-4)
diff --git a/official/vision/dataloaders/tf_example_label_map_decoder.py b/official/vision/dataloaders/tf_example_label_map_decoder.py
index a2f04477b19..146aebc2a43 100644
--- a/official/vision/dataloaders/tf_example_label_map_decoder.py
+++ b/official/vision/dataloaders/tf_example_label_map_decoder.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -18,8 +18,7 @@
protos for object detection.
"""
import csv
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import tf_example_decoder
diff --git a/official/vision/dataloaders/tf_example_label_map_decoder_test.py b/official/vision/dataloaders/tf_example_label_map_decoder_test.py
index 3ff9a8b3c14..00c73d99141 100644
--- a/official/vision/dataloaders/tf_example_label_map_decoder_test.py
+++ b/official/vision/dataloaders/tf_example_label_map_decoder_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,10 +15,9 @@
"""Tests for tf_example_label_map_decoder.py."""
import os
-# Import libraries
from absl.testing import parameterized
import numpy as np
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import tf_example_label_map_decoder
from official.vision.dataloaders import tfexample_utils
diff --git a/official/vision/dataloaders/tfds_classification_decoders.py b/official/vision/dataloaders/tfds_classification_decoders.py
index a27cdf5f3ce..a87c0031370 100644
--- a/official/vision/dataloaders/tfds_classification_decoders.py
+++ b/official/vision/dataloaders/tfds_classification_decoders.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""TFDS Classification decoders."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import decoder
@@ -35,4 +35,5 @@ def decode(self, serialized_example):
'cifar10': ClassificationDecorder,
'cifar100': ClassificationDecorder,
'imagenet2012': ClassificationDecorder,
+ 'imagenet2012_fewshot/10shot': ClassificationDecorder,
}
diff --git a/official/vision/dataloaders/tfds_detection_decoders.py b/official/vision/dataloaders/tfds_detection_decoders.py
index 4c270128a12..77b9ad966db 100644
--- a/official/vision/dataloaders/tfds_detection_decoders.py
+++ b/official/vision/dataloaders/tfds_detection_decoders.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""TFDS detection decoders."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import decoder
@@ -25,7 +25,7 @@ def decode(self, serialized_example):
"""Decode the serialized example.
Args:
- serialized_example: a dictonary example produced by tfds.
+ serialized_example: a dictionary example produced by tfds.
Returns:
decoded_tensors: a dictionary of tensors with the following fields:
@@ -56,5 +56,6 @@ def decode(self, serialized_example):
TFDS_ID_TO_DECODER_MAP = {
'coco/2017': MSCOCODecoder,
'coco/2014': MSCOCODecoder,
- 'coco': MSCOCODecoder
+ 'coco': MSCOCODecoder,
+ 'scenic:objects365': MSCOCODecoder,
}
diff --git a/official/vision/dataloaders/tfds_factory.py b/official/vision/dataloaders/tfds_factory.py
index 8f4c877043c..52da9154ac3 100644
--- a/official/vision/dataloaders/tfds_factory.py
+++ b/official/vision/dataloaders/tfds_factory.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
diff --git a/official/vision/dataloaders/tfds_factory_test.py b/official/vision/dataloaders/tfds_factory_test.py
index c9397a9724d..1f6295f42ec 100644
--- a/official/vision/dataloaders/tfds_factory_test.py
+++ b/official/vision/dataloaders/tfds_factory_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,7 +15,7 @@
"""Tests for tfds factory functions."""
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import decoder as base_decoder
from official.vision.dataloaders import tfds_factory
diff --git a/official/vision/dataloaders/tfds_segmentation_decoders.py b/official/vision/dataloaders/tfds_segmentation_decoders.py
index 9613d0f00ff..3b22f00f3b9 100644
--- a/official/vision/dataloaders/tfds_segmentation_decoders.py
+++ b/official/vision/dataloaders/tfds_segmentation_decoders.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,7 +14,7 @@
"""TFDS Semantic Segmentation decoders."""
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import decoder
diff --git a/official/vision/dataloaders/tfexample_utils.py b/official/vision/dataloaders/tfexample_utils.py
index 00647851521..6aa6c7c7d1b 100644
--- a/official/vision/dataloaders/tfexample_utils.py
+++ b/official/vision/dataloaders/tfexample_utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -42,32 +42,32 @@ def test_foo(self):
```
"""
-import io
-from typing import Sequence, Union
+from typing import Mapping, Optional, Sequence, Union
import numpy as np
-from PIL import Image
-import tensorflow as tf
+import tensorflow as tf, tf_keras
+
+from official.core import file_writers
+from official.vision.data import fake_feature_generator
+from official.vision.data import image_utils
+from official.vision.data import tf_example_builder
IMAGE_KEY = 'image/encoded'
CLASSIFICATION_LABEL_KEY = 'image/class/label'
-DISTILATION_LABEL_KEY = 'image/class/soft_labels'
+DISTILLATION_LABEL_KEY = 'image/class/soft_labels'
LABEL_KEY = 'clip/label/index'
AUDIO_KEY = 'features/audio'
-DUMP_SOURCE_ID = b'123'
+DUMP_SOURCE_ID = b'7435790'
-def encode_image(image_array: np.array, fmt: str) -> bytes:
- image = Image.fromarray(image_array)
- with io.BytesIO() as output:
- image.save(output, format=fmt)
- return output.getvalue()
+def encode_image(image_array: np.ndarray, fmt: str) -> bytes:
+ return image_utils.encode_image(image_array, fmt)
def make_image_bytes(shape: Sequence[int], fmt: str = 'JPEG') -> bytes:
"""Generates image and return bytes in specified format."""
- random_image = np.random.randint(0, 256, size=shape, dtype=np.uint8)
- return encode_image(random_image, fmt=fmt)
+ image = fake_feature_generator.generate_image_np(*shape)
+ return encode_image(image, fmt=fmt)
def put_int64_to_context(seq_example: tf.train.SequenceExample,
@@ -111,89 +111,145 @@ def make_video_test_example(image_shape: Sequence[int] = (263, 320, 3),
return seq_example
-def dump_to_tfrecord(record_file: str,
- tf_examples: Sequence[Union[tf.train.Example,
- tf.train.SequenceExample]]):
- """Writes serialized Example to TFRecord file with path."""
- with tf.io.TFRecordWriter(record_file) as writer:
- for tf_example in tf_examples:
- writer.write(tf_example.SerializeToString())
+def dump_to_tfrecord(
+ record_file: str,
+ tf_examples: Sequence[Union[tf.train.Example, tf.train.SequenceExample]],
+ file_type: str = 'tfrecord'):
+ """Writes serialized Example to TFRecord file with path.
+ Note that examples are expected to be not seriazlied.
-def _encode_image(image_array: np.ndarray, fmt: str) -> bytes:
- """Util function to encode an image."""
- image = Image.fromarray(image_array)
- with io.BytesIO() as output:
- image.save(output, format=fmt)
- return output.getvalue()
+ Args:
+ record_file: The name of the output file.
+ tf_examples: A list of examples to be stored.
+ file_type: A string indicating the file format, could be: 'tfrecord',
+ 'tfrecords', 'tfrecord_compressed', 'tfrecords_gzip', 'riegeli'. The
+ string is case insensitive.
+ """
+ file_writers.write_small_dataset(tf_examples, record_file, file_type)
def create_classification_example(
image_height: int,
image_width: int,
image_format: str = 'JPEG',
- is_multilabel: bool = False) -> tf.train.Example:
- """Creates image and labels for image classification input pipeline."""
- image = _encode_image(
- np.uint8(np.random.rand(image_height, image_width, 3) * 255),
- fmt=image_format)
- labels = [0, 1] if is_multilabel else [0]
- serialized_example = tf.train.Example(
- features=tf.train.Features(
- feature={
- IMAGE_KEY: (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[image]))),
- CLASSIFICATION_LABEL_KEY: (tf.train.Feature(
- int64_list=tf.train.Int64List(value=labels))),
- })).SerializeToString()
- return serialized_example
+ is_multilabel: bool = False,
+ output_serialized_example: bool = True) -> tf.train.Example:
+ """Creates image and labels for image classification input pipeline.
+
+ Args:
+ image_height: The height of test image.
+ image_width: The width of test image.
+ image_format: The format of test image.
+ is_multilabel: A boolean flag represents whether the test image can have
+ multiple labels.
+ output_serialized_example: A boolean flag represents whether to return a
+ serialized example.
+
+ Returns:
+ A tf.train.Example for testing.
+ """
+ image = fake_feature_generator.generate_image_np(image_height, image_width)
+ labels = fake_feature_generator.generate_classes_np(2,
+ int(is_multilabel) +
+ 1).tolist()
+ builder = tf_example_builder.TfExampleBuilder()
+ example = builder.add_image_matrix_feature(image, image_format,
+ DUMP_SOURCE_ID).add_ints_feature(
+ CLASSIFICATION_LABEL_KEY,
+ labels).example
+ if output_serialized_example:
+ return example.SerializeToString()
+ return example
def create_distillation_example(
image_height: int,
image_width: int,
num_labels: int,
- image_format: str = 'JPEG') -> tf.train.Example:
- """Creates image and labels for image classification with distillation."""
- image = _encode_image(
- np.uint8(np.random.rand(image_height, image_width, 3) * 255),
- fmt=image_format)
- soft_labels = [0.6] * num_labels
- serialized_example = tf.train.Example(
- features=tf.train.Features(
- feature={
- IMAGE_KEY: (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[image]))),
- DISTILATION_LABEL_KEY: (tf.train.Feature(
- float_list=tf.train.FloatList(value=soft_labels))),
- })).SerializeToString()
- return serialized_example
-
-
-def create_3d_image_test_example(image_height: int, image_width: int,
- image_volume: int,
- image_channel: int) -> tf.train.Example:
- """Creates 3D image and label."""
- images = np.random.rand(image_height, image_width, image_volume,
- image_channel)
- images = images.astype(np.float32)
-
- labels = np.random.randint(
- low=2, size=(image_height, image_width, image_volume, image_channel))
- labels = labels.astype(np.float32)
-
- feature = {
- IMAGE_KEY: (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[images.tobytes()]))),
- CLASSIFICATION_LABEL_KEY: (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[labels.tobytes()])))
- }
- return tf.train.Example(features=tf.train.Features(feature=feature))
-
-
-def create_detection_test_example(image_height: int, image_width: int,
- image_channel: int,
- num_instances: int) -> tf.train.Example:
+ image_format: str = 'JPEG',
+ output_serialized_example: bool = True) -> tf.train.Example:
+ """Creates image and labels for image classification with distillation.
+
+ Args:
+ image_height: The height of test image.
+ image_width: The width of test image.
+ num_labels: The number of labels used in test image.
+ image_format: The format of test image.
+ output_serialized_example: A boolean flag represents whether to return a
+ serialized example.
+
+ Returns:
+ A tf.train.Example for testing.
+ """
+ image = fake_feature_generator.generate_image_np(image_height, image_width)
+ labels = fake_feature_generator.generate_classes_np(2, 1).tolist()
+ soft_labels = (fake_feature_generator.generate_classes_np(1, num_labels) +
+ 0.6).tolist()
+ builder = tf_example_builder.TfExampleBuilder()
+ example = builder.add_image_matrix_feature(image, image_format,
+ DUMP_SOURCE_ID).add_ints_feature(
+ CLASSIFICATION_LABEL_KEY,
+ labels).add_floats_feature(
+ DISTILLATION_LABEL_KEY,
+ soft_labels).example
+ if output_serialized_example:
+ return example.SerializeToString()
+ return example
+
+
+def create_3d_image_test_example(
+ image_height: int,
+ image_width: int,
+ image_volume: int,
+ image_channel: int,
+ num_classes: int = 2,
+ output_serialized_example: bool = False) -> tf.train.Example:
+ """Creates 3D image and label.
+
+ Args:
+ image_height: The height of test 3D image.
+ image_width: The width of test 3D image.
+ image_volume: The volume of test 3D image.
+ image_channel: The channel of test 3D image.
+ num_classes: The number of classes of the test 3D label.
+ output_serialized_example: A boolean flag represents whether to return a
+ serialized example.
+
+ Returns:
+ A tf.train.Example for testing.
+ """
+ image = fake_feature_generator.generate_image_np(image_height, image_width,
+ image_channel)
+ images = image[:, :, np.newaxis, :]
+ images = np.tile(images, [1, 1, image_volume, 1]).astype(np.float32)
+
+ label_shape = [image_height, image_width, image_volume, num_classes]
+ labels = (
+ fake_feature_generator.generate_classes_np(
+ num_classes, np.prod(label_shape)
+ )
+ .reshape(label_shape)
+ .astype(np.float32)
+ )
+
+ builder = tf_example_builder.TfExampleBuilder()
+ example = builder.add_bytes_feature(IMAGE_KEY,
+ images.tobytes()).add_bytes_feature(
+ CLASSIFICATION_LABEL_KEY,
+ labels.tobytes()).example
+ if output_serialized_example:
+ return example.SerializeToString()
+ return example
+
+
+def create_detection_test_example(
+ image_height: int,
+ image_width: int,
+ image_channel: int,
+ num_instances: int,
+ fill_image_size: bool = True,
+ output_serialized_example: bool = False) -> tf.train.Example:
"""Creates and returns a test example containing box and mask annotations.
Args:
@@ -201,90 +257,80 @@ def create_detection_test_example(image_height: int, image_width: int,
image_width: The width of test image.
image_channel: The channel of test image.
num_instances: The number of object instances per image.
+ fill_image_size: If image height and width will be added to the example.
+ output_serialized_example: A boolean flag represents whether to return a
+ serialized example.
Returns:
A tf.train.Example for testing.
"""
- image = make_image_bytes([image_height, image_width, image_channel])
- if num_instances == 0:
- xmins = []
- xmaxs = []
- ymins = []
- ymaxs = []
- labels = []
- areas = []
- is_crowds = []
- masks = []
- labels_text = []
- else:
- xmins = list(np.random.rand(num_instances))
- xmaxs = list(np.random.rand(num_instances))
- ymins = list(np.random.rand(num_instances))
- ymaxs = list(np.random.rand(num_instances))
- labels_text = [b'class_1'] * num_instances
- labels = list(np.random.randint(100, size=num_instances))
- areas = [(xmax - xmin) * (ymax - ymin) * image_height * image_width
- for xmin, xmax, ymin, ymax in zip(xmins, xmaxs, ymins, ymaxs)]
- is_crowds = [0] * num_instances
- masks = []
- for _ in range(num_instances):
- mask = make_image_bytes([image_height, image_width], fmt='PNG')
- masks.append(mask)
- return tf.train.Example(
- features=tf.train.Features(
- feature={
- 'image/encoded': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[image]))),
- 'image/source_id': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[DUMP_SOURCE_ID]))),
- 'image/height': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=[image_height]))),
- 'image/width': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=[image_width]))),
- 'image/object/bbox/xmin': (tf.train.Feature(
- float_list=tf.train.FloatList(value=xmins))),
- 'image/object/bbox/xmax': (tf.train.Feature(
- float_list=tf.train.FloatList(value=xmaxs))),
- 'image/object/bbox/ymin': (tf.train.Feature(
- float_list=tf.train.FloatList(value=ymins))),
- 'image/object/bbox/ymax': (tf.train.Feature(
- float_list=tf.train.FloatList(value=ymaxs))),
- 'image/object/class/label': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=labels))),
- 'image/object/class/text': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=labels_text))),
- 'image/object/is_crowd': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=is_crowds))),
- 'image/object/area': (tf.train.Feature(
- float_list=tf.train.FloatList(value=areas))),
- 'image/object/mask': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=masks))),
- }))
-
-
-def create_segmentation_test_example(image_height: int, image_width: int,
- image_channel: int) -> tf.train.Example:
+ image = fake_feature_generator.generate_image_np(image_height, image_width,
+ image_channel)
+ boxes = fake_feature_generator.generate_normalized_boxes_np(num_instances)
+ ymins, xmins, ymaxs, xmaxs = boxes.T.tolist()
+ is_crowds = [0] * num_instances
+ labels = fake_feature_generator.generate_classes_np(
+ 2, size=num_instances).tolist()
+ labels_text = [b'class_1'] * num_instances
+ masks = fake_feature_generator.generate_instance_masks_np(
+ image_height, image_width, boxes)
+
+ builder = tf_example_builder.TfExampleBuilder()
+
+ example = builder.add_image_matrix_feature(
+ image, image_source_id=DUMP_SOURCE_ID).add_boxes_feature(
+ xmins, xmaxs, ymins, ymaxs,
+ labels).add_instance_mask_matrices_feature(masks).add_ints_feature(
+ 'image/object/is_crowd',
+ is_crowds).add_bytes_feature('image/object/class/text',
+ labels_text).example
+ if not fill_image_size:
+ del example.features.feature['image/height']
+ del example.features.feature['image/width']
+
+ if output_serialized_example:
+ return example.SerializeToString()
+ return example
+
+
+def create_segmentation_test_example(
+ image_height: int,
+ image_width: int,
+ image_channel: int,
+ output_serialized_example: bool = False,
+ dense_features: Optional[Mapping[str, int]] = None) -> tf.train.Example:
"""Creates and returns a test example containing mask annotations.
Args:
image_height: The height of test image.
image_width: The width of test image.
image_channel: The channel of test image.
-
+ output_serialized_example: A boolean flag represents whether to return a
+ serialized example.
+ dense_features: An optional dictionary of additional dense features, where
+ the key is the prefix of the feature key in tf.Example and the value is
+ the number of the channels of this feature.
Returns:
A tf.train.Example for testing.
"""
- image = make_image_bytes([image_height, image_width, image_channel])
- mask = make_image_bytes([image_height, image_width], fmt='PNG')
- return tf.train.Example(
- features=tf.train.Features(
- feature={
- 'image/encoded': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[image]))),
- 'image/segmentation/class/encoded': (tf.train.Feature(
- bytes_list=tf.train.BytesList(value=[mask]))),
- 'image/height': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=[image_height]))),
- 'image/width': (tf.train.Feature(
- int64_list=tf.train.Int64List(value=[image_width])))
- }))
+ image = fake_feature_generator.generate_image_np(image_height, image_width,
+ image_channel)
+ mask = fake_feature_generator.generate_semantic_mask_np(
+ image_height, image_width, 3)
+ builder = tf_example_builder.TfExampleBuilder()
+ builder.add_image_matrix_feature(
+ image,
+ image_source_id=DUMP_SOURCE_ID).add_semantic_mask_matrix_feature(mask)
+
+ if dense_features:
+ for prefix, channel in dense_features.items():
+ dense_feature = fake_feature_generator.generate_semantic_mask_np(
+ image_height, image_width, channel)
+ builder.add_semantic_mask_matrix_feature(
+ dense_feature, feature_prefix=prefix)
+
+ example = builder.example
+
+ if output_serialized_example:
+ return example.SerializeToString()
+ return example
diff --git a/official/vision/dataloaders/utils.py b/official/vision/dataloaders/utils.py
index 408522abe29..1625f03a91d 100644
--- a/official/vision/dataloaders/utils.py
+++ b/official/vision/dataloaders/utils.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,8 +15,7 @@
"""Data loader utils."""
from typing import Dict
-# Import libraries
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.ops import preprocess_ops
@@ -67,3 +66,20 @@ def pad_groundtruths_to_fixed_size(groundtruths: Dict[str, tf.Tensor],
groundtruths['attributes'][k] = preprocess_ops.clip_or_pad_to_fixed_size(
v, size, -1)
return groundtruths
+
+
+def binarize_matting_map(matting_map: tf.Tensor,
+ threshold: float = 0.5) -> tf.Tensor:
+ """Binarizes a matting map.
+
+ If the matting_map value is above a threshold, set it as 1 otherwise 0. The
+ binarization is done for every element in the matting_map.
+
+ Args:
+ matting_map: The groundtruth in the matting map format.
+ threshold: The threshold used to binarize the matting map.
+
+ Returns:
+ The binarized labels (0 for BG, 1 for FG) as tf.float32.
+ """
+ return tf.cast(tf.greater(matting_map, threshold), tf.float32)
diff --git a/official/vision/dataloaders/utils_test.py b/official/vision/dataloaders/utils_test.py
index 8622b9b414c..df1eec33a4c 100644
--- a/official/vision/dataloaders/utils_test.py
+++ b/official/vision/dataloaders/utils_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -14,10 +14,8 @@
"""Tests for dataloader utils functions."""
-# Import libraries
-
from absl.testing import parameterized
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.dataloaders import utils
diff --git a/official/vision/dataloaders/video_input.py b/official/vision/dataloaders/video_input.py
index ba3160ee4a1..c5b01fa5368 100644
--- a/official/vision/dataloaders/video_input.py
+++ b/official/vision/dataloaders/video_input.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -17,7 +17,7 @@
from typing import Dict, Optional, Tuple, Union
from absl import logging
-import tensorflow as tf
+import tensorflow as tf, tf_keras
from official.vision.configs import video_classification as exp_cfg
from official.vision.dataloaders import decoder
@@ -36,23 +36,26 @@ def process_image(image: tf.Tensor,
random_stride_range: int = 0,
num_test_clips: int = 1,
min_resize: int = 256,
- crop_size: int = 224,
+ crop_size: Union[int, Tuple[int, int]] = 224,
+ num_channels: int = 3,
num_crops: int = 1,
zero_centering_image: bool = False,
min_aspect_ratio: float = 0.5,
max_aspect_ratio: float = 2,
min_area_ratio: float = 0.49,
max_area_ratio: float = 1.0,
+ random_rotation: bool = False,
augmenter: Optional[augment.ImageAugment] = None,
- seed: Optional[int] = None) -> tf.Tensor:
+ seed: Optional[int] = None,
+ input_image_format: Optional[str] = 'jpeg') -> tf.Tensor:
"""Processes a serialized image tensor.
Args:
- image: Input Tensor of shape [timesteps] and type tf.string of serialized
+ image: Input Tensor of shape [time-steps] and type tf.string of serialized
frames.
is_training: Whether or not in training mode. If True, random sample, crop
and left right flip is used.
- num_frames: Number of frames per subclip.
+ num_frames: Number of frames per sub clip.
stride: Temporal stride to sample frames.
random_stride_range: An int indicating the min and max bounds to uniformly
sample different strides from the video. E.g., a value of 1 with stride=2
@@ -62,10 +65,12 @@ def process_image(image: tf.Tensor,
num_test_clips: Number of test clips (1 by default). If more than 1, this
will sample multiple linearly spaced clips within each video at test time.
If 1, then a single clip in the middle of the video is sampled. The clips
- are aggreagated in the batch dimension.
+ are aggregated in the batch dimension.
min_resize: Frames are resized so that min(height, width) is min_resize.
- crop_size: Final size of the frame after cropping the resized frames. Both
- height and width are the same.
+ crop_size: Final size of the frame after cropping the resized frames.
+ Optionally, specify a tuple of (crop_height, crop_width) if
+ crop_height != crop_width.
+ num_channels: Number of channels of the clip.
num_crops: Number of crops to perform on the resized frames.
zero_centering_image: If True, frames are normalized to values in [-1, 1].
If False, values in [0, 1].
@@ -73,12 +78,15 @@ def process_image(image: tf.Tensor,
max_aspect_ratio: The maximum aspect range for cropping.
min_area_ratio: The minimum area range for cropping.
max_area_ratio: The maximum area range for cropping.
+ random_rotation: Use uniform random rotation augmentation or not.
augmenter: Image augmenter to distort each image.
seed: A deterministic seed to use when sampling.
+ input_image_format: The format of input image which could be jpeg, png or
+ none for unknown or mixed datasets.
Returns:
Processed frames. Tensor of shape
- [num_frames * num_test_clips, crop_size, crop_size, 3].
+ [num_frames * num_test_clips, crop_height, crop_width, num_channels].
"""
# Validate parameters.
if is_training and num_test_clips != 1:
@@ -90,6 +98,14 @@ def process_image(image: tf.Tensor,
raise ValueError('Random stride range should be >= 0, got {}'.format(
random_stride_range))
+ if input_image_format not in ('jpeg', 'png', 'none'):
+ raise ValueError('Unknown input image format: {}'.format(
+ input_image_format))
+
+ if isinstance(crop_size, int):
+ crop_size = (crop_size, crop_size)
+ crop_height, crop_width = crop_size
+
# Temporal sampler.
if is_training:
if random_stride_range > 0:
@@ -113,15 +129,17 @@ def process_image(image: tf.Tensor,
# Decode JPEG string to tf.uint8.
if image.dtype == tf.string:
- image = preprocess_ops_3d.decode_jpeg(image, 3)
+ image = preprocess_ops_3d.decode_image(image, num_channels)
if is_training:
# Standard image data augmentation: random resized crop and random flip.
image = preprocess_ops_3d.random_crop_resize(
- image, crop_size, crop_size, num_frames, 3,
+ image, crop_height, crop_width, num_frames, num_channels,
(min_aspect_ratio, max_aspect_ratio),
(min_area_ratio, max_area_ratio))
image = preprocess_ops_3d.random_flip_left_right(image, seed)
+ if random_rotation:
+ image = preprocess_ops_3d.random_rotation(image, seed)
if augmenter is not None:
image = augmenter.distort(image)
@@ -129,7 +147,7 @@ def process_image(image: tf.Tensor,
# Resize images (resize happens only if necessary to save compute).
image = preprocess_ops_3d.resize_smallest(image, min_resize)
# Crop of the frames.
- image = preprocess_ops_3d.crop_image(image, crop_size, crop_size, False,
+ image = preprocess_ops_3d.crop_image(image, crop_height, crop_width, False,
num_crops)
# Cast the frames in float32, normalizing according to zero_centering_image.
@@ -146,17 +164,17 @@ def postprocess_image(image: tf.Tensor,
The same parameters used in process should be used here.
Args:
- image: Input Tensor of shape [batch, timesteps, height, width, 3].
+ image: Input Tensor of shape [batch, time-steps, height, width, 3].
is_training: Whether or not in training mode. If True, random sample, crop
and left right flip is used.
- num_frames: Number of frames per subclip.
+ num_frames: Number of frames per sub clip.
num_test_clips: Number of test clips (1 by default). If more than 1, this
will sample multiple linearly spaced clips within each video at test time.
If 1, then a single clip in the middle of the video is sampled. The clips
- are aggreagated in the batch dimension.
+ are aggregated in the batch dimension.
num_test_crops: Number of test crops (1 by default). If more than 1, there
are multiple crops for each clip at test time. If 1, there is a single
- central crop. The crops are aggreagated in the batch dimension.
+ central crop. The crops are aggregated in the batch dimension.
Returns:
Processed frames. Tensor of shape
@@ -164,7 +182,7 @@ def postprocess_image(image: tf.Tensor,
"""
num_views = num_test_clips * num_test_crops
if num_views > 1 and not is_training:
- # In this case, multiple views are merged together in batch dimenstion which
+ # In this case, multiple views are merged together in batch dimension which
# will be batch * num_views.
image = tf.reshape(image, [-1, num_frames] + image.shape[2:].as_list())
@@ -173,15 +191,16 @@ def postprocess_image(image: tf.Tensor,
def process_label(label: tf.Tensor,
one_hot_label: bool = True,
- num_classes: Optional[int] = None) -> tf.Tensor:
+ num_classes: Optional[int] = None,
+ label_dtype: tf.DType = tf.int32) -> tf.Tensor:
"""Processes label Tensor."""
# Validate parameters.
if one_hot_label and not num_classes:
raise ValueError(
'`num_classes` should be given when requesting one hot label.')
- # Cast to tf.int32.
- label = tf.cast(label, dtype=tf.int32)
+ # Cast to label_dtype (default = tf.int32).
+ label = tf.cast(label, dtype=label_dtype)
if one_hot_label:
# Replace label index by one hot representation.
@@ -269,34 +288,56 @@ def __init__(self,
self._random_stride_range = input_params.random_stride_range
self._num_test_clips = input_params.num_test_clips
self._min_resize = input_params.min_image_size
- self._crop_size = input_params.feature_shape[1]
+ crop_height = input_params.feature_shape[1]
+ crop_width = input_params.feature_shape[2]
+ self._crop_size = crop_height if crop_height == crop_width else (
+ crop_height, crop_width)
+ self._num_channels = input_params.feature_shape[3]
self._num_crops = input_params.num_test_crops
+ self._zero_centering_image = input_params.zero_centering_image
self._one_hot_label = input_params.one_hot
self._num_classes = input_params.num_classes
self._image_key = image_key
self._label_key = label_key
self._dtype = tf.dtypes.as_dtype(input_params.dtype)
+ self._label_dtype = tf.dtypes.as_dtype(input_params.label_dtype)
self._output_audio = input_params.output_audio
self._min_aspect_ratio = input_params.aug_min_aspect_ratio
self._max_aspect_ratio = input_params.aug_max_aspect_ratio
self._min_area_ratio = input_params.aug_min_area_ratio
self._max_area_ratio = input_params.aug_max_area_ratio
+ self._input_image_format = input_params.input_image_format
+ self._random_rotation = input_params.aug_random_rotation
if self._output_audio:
self._audio_feature = input_params.audio_feature
self._audio_shape = input_params.audio_feature_shape
- self._augmenter = None
- if input_params.aug_type is not None:
- aug_type = input_params.aug_type
- if aug_type == 'autoaug':
+ aug_type = input_params.aug_type
+ if aug_type is not None:
+ if aug_type.type == 'autoaug':
logging.info('Using AutoAugment.')
- self._augmenter = augment.AutoAugment()
- elif aug_type == 'randaug':
+ self._augmenter = augment.AutoAugment(
+ augmentation_name=aug_type.autoaug.augmentation_name,
+ cutout_const=aug_type.autoaug.cutout_const,
+ translate_const=aug_type.autoaug.translate_const)
+ elif aug_type.type == 'randaug':
logging.info('Using RandAugment.')
- self._augmenter = augment.RandAugment()
+ self._augmenter = augment.RandAugment(
+ num_layers=aug_type.randaug.num_layers,
+ magnitude=aug_type.randaug.magnitude,
+ cutout_const=aug_type.randaug.cutout_const,
+ translate_const=aug_type.randaug.translate_const,
+ prob_to_apply=aug_type.randaug.prob_to_apply,
+ exclude_ops=aug_type.randaug.exclude_ops)
+ else:
+ raise ValueError(
+ 'Augmentation policy {} not supported.'.format(aug_type.type))
+ else:
+ self._augmenter = None
+ if self._random_rotation:
+ logging.info('Using standard augmentation with rotation.')
else:
- raise ValueError('Augmentation policy {} is not supported.'.format(
- aug_type))
+ logging.info('Using standard augmentation without rotation.')
def _parse_train_data(
self, decoded_tensors: Dict[str, tf.Tensor]
@@ -313,17 +354,22 @@ def _parse_train_data(
num_test_clips=self._num_test_clips,
min_resize=self._min_resize,
crop_size=self._crop_size,
+ num_channels=self._num_channels,
min_aspect_ratio=self._min_aspect_ratio,
max_aspect_ratio=self._max_aspect_ratio,
min_area_ratio=self._min_area_ratio,
max_area_ratio=self._max_area_ratio,
- augmenter=self._augmenter)
+ random_rotation=self._random_rotation,
+ augmenter=self._augmenter,
+ zero_centering_image=self._zero_centering_image,
+ input_image_format=self._input_image_format)
image = tf.cast(image, dtype=self._dtype)
features = {'image': image}
label = decoded_tensors[self._label_key]
- label = process_label(label, self._one_hot_label, self._num_classes)
+ label = process_label(label, self._one_hot_label, self._num_classes,
+ self._label_dtype)
if self._output_audio:
audio = decoded_tensors[self._audio_feature]
@@ -349,12 +395,16 @@ def _parse_eval_data(
num_test_clips=self._num_test_clips,
min_resize=self._min_resize,
crop_size=self._crop_size,
- num_crops=self._num_crops)
+ num_channels=self._num_channels,
+ num_crops=self._num_crops,
+ zero_centering_image=self._zero_centering_image,
+ input_image_format=self._input_image_format)
image = tf.cast(image, dtype=self._dtype)
features = {'image': image}
label = decoded_tensors[self._label_key]
- label = process_label(label, self._one_hot_label, self._num_classes)
+ label = process_label(label, self._one_hot_label, self._num_classes,
+ self._label_dtype)
if self._output_audio:
audio = decoded_tensors[self._audio_feature]
diff --git a/official/vision/dataloaders/video_input_test.py b/official/vision/dataloaders/video_input_test.py
index 2f58e642117..db8b70b1df6 100644
--- a/official/vision/dataloaders/video_input_test.py
+++ b/official/vision/dataloaders/video_input_test.py
@@ -1,4 +1,4 @@
-# Copyright 2022 The TensorFlow Authors. All Rights Reserved.
+# Copyright 2026 The TensorFlow Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -15,12 +15,12 @@
import io
-# Import libraries
import numpy as np
from PIL import Image
-import tensorflow as tf
+import tensorflow as tf, tf_keras
import tensorflow_datasets as tfds
+from official.vision.configs import common
from official.vision.configs import video_classification as exp_cfg
from official.vision.dataloaders import video_input
@@ -173,7 +173,8 @@ def test_video_input_augmentation_returns_shape(self):
params.min_image_size = 224
params.temporal_stride = 2
- params.aug_type = 'autoaug'
+ params.aug_type = common.Augmentation(
+ type='autoaug', autoaug=common.AutoAugment())
decoder = video_input.Decoder()
parser = video_input.Parser(params).parse_fn(params.is_training)
@@ -189,6 +190,28 @@ def test_video_input_augmentation_returns_shape(self):
self.assertAllEqual(image.shape, (2, 224, 224, 3))
self.assertAllEqual(label.shape, (600,))
+ def test_video_input_image_shape_label_type(self):
+ params = exp_cfg.kinetics600(is_training=True)
+ params.feature_shape = (2, 168, 224, 1)
+ params.min_image_size = 168
+ params.label_dtype = 'float32'
+ params.one_hot = False
+
+ decoder = video_input.Decoder()
+ parser = video_input.Parser(params).parse_fn(params.is_training)
+
+ seq_example, label = fake_seq_example()
+
+ input_tensor = tf.constant(seq_example.SerializeToString())
+ decoded_tensors = decoder.decode(input_tensor)
+ output_tensor = parser(decoded_tensors)
+ image_features, label = output_tensor
+ image = image_features['image']
+
+ self.assertAllEqual(image.shape, (2, 168, 224, 1))
+ self.assertAllEqual(label.shape, (1,))
+ self.assertDTypeEqual(label, tf.float32)
+
if __name__ == '__main__':
tf.test.main()
diff --git a/official/vision/docs/aug_reg.md b/official/vision/docs/aug_reg.md
new file mode 100644
index 00000000000..3ed571edc15
--- /dev/null
+++ b/official/vision/docs/aug_reg.md
@@ -0,0 +1,251 @@
+# TF-Vision Data Augmentation and Model Regularization
+
+## Data augmentation methods
+
+TF-Vision provides a rich collection of advanced SoTA data augmentation methods
+for various vision tasks. Default augmentation methods such as random flipping
+or cropping will not be discussed here.
+
+### RandAugment and AutoAugment
+
+[RandAugment](https://arxiv.org/abs/1909.13719) or
+[AutoAugment](https://arxiv.org/abs/1805.09501) combines multiple common image
+augmentations/transformations that adjust the contrast, orientation, color,
+brightness and sharpness. It randomly selects the augmentations as well as the
+augmentation strength. The TF-Vision implementations are designed to work out of
+the box for a wide range of datasets and come with low overhead. RandAugment
+and AutoAugment nowadays are commonly adopted for SoTA image classification
+model training.
+
+Supported vision tasks:
+
+* Image classification
+* Video action classification
+* Object detection
+
+[Definitions](https://github.com/tensorflow/models/blob/master/official/vision/configs/common.py)
+in TF-Vision:
+
+```python
+@dataclasses.dataclass
+class RandAugment(hyperparams.Config):
+ """Configuration for RandAugment."""
+ num_layers: int = 2
+ magnitude: float = 10
+ cutout_const: float = 40
+ translate_const: float = 10
+ magnitude_std: float = 0.0
+ prob_to_apply: Optional[float] = None
+ exclude_ops: List[str] = dataclasses.field(default_factory=list)
+
+
+@dataclasses.dataclass
+class AutoAugment(hyperparams.Config):
+ """Configuration for AutoAugment."""
+ augmentation_name: str = 'v0'
+ cutout_const: float = 100
+ translate_const: float = 250
+```
+
+Yaml template of using the techniques in TF-Vision:
+
+```python
+task:
+ train_data:
+ aug_type:
+ type: 'randaug'
+ randaug:
+ magnitude: 10
+ magnitude_std: 0.0
+ num_layers: 2
+```
+
+### Image scale jittering
+
+Image scale jittering randomly selects a scale from a user-specified range. It
+either upsamples the image if the scale is greater than 1.0 or downsamples the image if
+the scale is smaller than 1.0. Randomly cropping or zero-padding is further
+performed to tailor the scaled image to a desired size. Scale jittering is one
+of the most effective augmentation methods for training SoTA object detection
+model. It has been widely used in recent publications on object detection such
+as [SpineNet](https://arxiv.org/abs/1912.05027),
+[EfficientDet](https://arxiv.org/abs/1911.09070),
+[Detection-RS](https://arxiv.org/pdf/2107.00057.pdf) and
+[Copy-Paste](https://arxiv.org/abs/2012.07177).
+
+Supported vision tasks:
+
+* Object detection
+* Semantic segmentation
+
+[Definition](https://github.com/tensorflow/models/blob/ae79c473faf78bfa068ecbb57504483eb81eef21/official/vision/configs/retinanet.py#L51C1-L67C35)
+in TF-Vision:
+
+```python
+@dataclasses.dataclass
+class Parser(hyperparams.Config):
+ aug_scale_min: float = 1.0
+ aug_scale_max: float = 1.0
+```
+
+Yaml template of using the techniques in TF-Vision:
+
+```python
+task:
+ train_data:
+ parser:
+ aug_scale_max: 2.0
+ aug_scale_min: 0.5
+```
+
+### Mixup and Cutmix
+
+[Mixup](https://arxiv.org/pdf/1710.09412.pdf) combines two input images through
+a random convex combination:
+
+```python
+img = img_1 * a + img_2 * (1.0 - a),
+```
+
+where a is drawn from a beta distribution). Analogously, the labels are
+calculated with:
+
+```python
+label = label_1 * a + label_2 * (1.0 - a).
+```
+
+[Cutmix](https://arxiv.org/pdf/1905.04899.pdf) also combines two training
+instances. Instead of linearly interpolating, here we paste a random rectangular
+area of img_2 and into img_1. The labels are derived similarly to mixup:
+
+```python
+label = label_1 * a + label_2 * (1.0 - a),
+```
+
+where a compensates the area of the inserted rectangle. The two methods are key
+factors in training transformer based image models
+([ViT](https://arxiv.org/abs/2010.11929),
+[DEIT](https://arxiv.org/abs/2012.12877)) with limited labeled data.
+
+Supported vision tasks:
+
+* Image classification
+
+[Definitions](https://github.com/tensorflow/models/blob/ae79c473faf78bfa068ecbb57504483eb81eef21/official/vision/configs/common.py#L93C1-L100C40)
+in TF-Vision:
+
+```python
+@dataclasses.dataclass
+class MixupAndCutmix(hyperparams.Config):
+ """Configuration for MixupAndCutmix."""
+ mixup_alpha: float = .8
+ cutmix_alpha: float = 1.
+ prob: float = 1.0
+ switch_prob: float = 0.5
+ label_smoothing: float = 0.1
+```
+
+Yaml template of using the methods in TF-Vision:
+
+```python
+task:
+ train_data:
+ mixup_and_cutmix:
+ cutmix_alpha: 1.0
+ label_smoothing: 0.1
+ mixup_alpha: 0.8
+ prob: 1.0
+ switch_prob: 0.5
+```
+
+### Random erasing
+
+[Random erasing](https://arxiv.org/pdf/1708.04896.pdf) samples random rectangles
+of different aspect ratios (default is one per image). Then, the pixels in these
+rectangles are replaced by random Gaussian noise. Note that random erasing is
+applied after normalizing the input images to zero mean and a variance of one.
+To apply random erasing , we exclude the Cutout augmentation in RandAugment.
+Cutout fills random square regions of the image.
+
+Supported vision tasks:
+
+* Image classification
+
+[Definition](https://github.com/tensorflow/models/blob/ae79c473faf78bfa068ecbb57504483eb81eef21/official/vision/configs/common.py#L80C1-L90C14)
+in TF-Vision:
+
+```python
+@dataclasses.dataclass
+class RandomErasing(hyperparams.Config):
+ """Configuration for RandomErasing."""
+ probability: float = 0.25
+ min_area: float = 0.02
+ max_area: float = 1 / 3
+ min_aspect: float = 0.3
+ max_aspect = None
+ min_count = 1
+ max_count = 1
+ trials = 10
+```
+
+Yaml template of using the methods in TF-Vision:
+
+```python
+task:
+ train_data:
+ random_erasing:
+ min_area: float: 0.02
+ max_area: float: 0.33
+```
+
+## Model Regularization methods
+
+TF-Vision provides a rich collection of advanced SoTA model regularization
+methods for different models and tasks. Default regularization methods such as
+weight decay and dropout will not be described here.
+
+### Stochastic depth
+
+[Stochastic depth](https://arxiv.org/abs/1603.09382) complements the idea of
+residual networks and is conceptually similar to dropout. In each mini-batch a
+random set of layers is bypassed/replaced with the identity function (aka
+residual connection). The probability of bypassing a layer increases linearly
+with the depth of the network. This procedure reduces the training time and
+improves generalization.
+
+Supported models:
+
+* ResNet, ResNet-RS
+* SpineNet, SpineNet-mobile, SpineNet-seg
+* Vision Transformer (ViT)
+* ResNet-RS-3D
+
+Yaml template of using the methods in TF-Vision:
+
+```python
+task:
+ model:
+ backbone:
+ resnet:
+ init_stochastic_depth_rate: 0.2
+```
+
+### Label smoothing
+
+[Label smoothing](https://arxiv.org/pdf/1512.00567.pdf) aims to reduce the issue
+of overconfidence and consequently overfitting. Instead of one-hot labeling, we
+assign a high value x (close to 1) to the correct class and evenly distribute
+(1 - x) to the remaining classes.
+
+Supported models:
+
+* Image classification
+* Video action classification
+
+Yaml template of using the methods in TF-Vision:
+
+```python
+task
+ losses:
+ label_smoothing: 0.1
+```
diff --git a/official/vision/docs/customize_input_pipeline.md b/official/vision/docs/customize_input_pipeline.md
new file mode 100644
index 00000000000..4dc27547a6c
--- /dev/null
+++ b/official/vision/docs/customize_input_pipeline.md
@@ -0,0 +1,370 @@
+# Customize Input Pipeline
+
+
+
+
+
+
+## Overview
+
+
+A task is a class that encapsulates the logic of loading data, building models,
+performing one-step training and validation, etc. It connects all components
+together and is called by the base
+[Trainer](https://github.com/tensorflow/models/blob/master/official/core/base_trainer.py).
+You can create your own task by inheriting from base
+[Task](https://github.com/tensorflow/models/blob/master/official/core/base_task.py),
+or from one of the
+[tasks](https://github.com/tensorflow/models/tree/master/official/vision/tasks)
+we already defined, if most of the operations can be reused. An `ExampleTask`
+inheriting from
+[ImageClassificationTask](https://github.com/tensorflow/models/blob/master/official/vision/tasks/image_classification.py#L31)
+can be found
+[here](https://github.com/tensorflow/models/blob/master/official/vision/examples/starter/example_task.py).
+
+
+In a task class, the `build_inputs` method is responsible for building the input
+pipeline for training and evaluation. Specifically, it will instantiate a
+Decoder object and a Parser object, which are used to create an `InputReader`
+that will generate a `tf.data.Dataset` object.
+
+
+Here's an example code snippet that demonstrates how to create a custom
+`build_inputs` method:
+
+
+```python
+def build_inputs(
+ self,
+ params: exp_cfg.DataConfig,
+ input_context: Optional[tf.distribute.InputContext] = None
+) -> tf.data.Dataset:
+ ....
+
+
+ decoder = sample_input.Decoder()
+ parser = sample_input.Parser(
+ output_size=..., num_classes=...)
+ reader = input_reader_factory.input_reader_generator(
+ params,
+ dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
+ decoder_fn=decoder.decode,
+ parser_fn=parser.parse_fn(params.is_training))
+ ....
+
+
+ dataset = reader.read(input_context=input_context)
+ return dataset
+```
+
+
+The class being responsible for building the input pipeline is
+[InputReader](https://github.com/tensorflow/models/blob/b1a7752c5137822a32bd0dd70a0cb96e807ea411/official/core/input_reader.py#L214)
+with interface
+
+```python
+class InputReader:
+ """Input reader that returns a tf.data.Dataset instance."""
+
+ def __init__(
+ self,
+ params: cfg.DataConfig,
+ dataset_fn=tf.data.TFRecordDataset,
+ decoder_fn: Optional[Callable[..., Any]] = None,
+ combine_fn: Optional[Callable[..., Any]] = None,
+ sample_fn: Optional[Callable[..., Any]] = None,
+ parser_fn: Optional[Callable[..., Any]] = None,
+ filter_fn: Optional[Callable[..., tf.Tensor]] = None,
+ transform_and_batch_fn: Optional[
+ Callable[
+ [tf.data.Dataset, Optional[tf.distribute.InputContext]],
+ tf.data.Dataset,
+ ]
+ ] = None,
+ postprocess_fn: Optional[Callable[..., Any]] = None,
+ ):
+ ....
+
+ def read(self,
+ input_context: Optional[tf.distribute.InputContext] = None,
+ dataset: Optional[tf.data.Dataset] = None) -> tf.data.Dataset:
+ """Generates a tf.data.Dataset object."""
+ if dataset is None:
+ dataset = self._read_data_source(self._matched_files, self._dataset_fn,
+ input_context)
+ dataset = self._decode_and_parse_dataset(dataset, self._global_batch_size,
+ input_context)
+ dataset = _maybe_map_fn(dataset, self._postprocess_fn)
+ if not (self._enable_shared_tf_data_service_between_parallel_trainers and
+ self._apply_tf_data_service_before_batching):
+ dataset = self._maybe_apply_data_service(dataset, input_context)
+
+ if self._deterministic is not None:
+ options = tf.data.Options()
+ options.deterministic = self._deterministic
+ dataset = dataset.with_options(options)
+ if self._autotune_algorithm:
+ options = tf.data.Options()
+ options.autotune.autotune_algorithm = (
+ tf.data.experimental.AutotuneAlgorithm[self._autotune_algorithm])
+ dataset = dataset.with_options(options)
+ return dataset.prefetch(self._prefetch_buffer_size)
+```
+
+Therefore, customizing the input pipeline is equivalent to having customized
+versions of `dataset_fn`, `decoder_fn`, etc. The execution order is generally
+as:
+
+```
+dataset_fn -> decoder_fn -> combine_fn -> parser_fn -> filter_fn ->
+transform_and_batch_fn -> postprocess_fn
+```
+
+The `transform_and_batch_fn` is an optional function that merges multiple
+examples into a batch and its default behavior to `dataset.batch` if not
+specified. In this workflow, the functions before `transform_and_batch_fn`, e.g.
+`dataset_fn`, `decoder_fn`, consume tensors without the batch dimension, while
+`postprocess_fn` will consume tensors with the batch dimension.
+
+We have essentially covered
+[decoder_fn](https://github.com/tensorflow/models/blob/master/official/vision/docs/read_custom_datasets.md#decoder),
+and `parser_fn` is another very important one that takes the decoded raw tensors
+dict and parses them into a dictionary of tensors that can be consumed by the
+model. It will be executed after decoder_fn.
+
+It is also worth noting that optimizing of the input pipeline through
+batching, shuffling and prefetching is also implemented in this class.
+
+## Parser
+
+A custom data loader can also be useful if you want to take advantage of
+features such as data augmentation.
+
+Customizing preprocessing is useful because it allows the user to tailor the
+preprocessing steps to suit the specific requirements of the task. While there
+are standard preprocessing techniques that are commonly used, different
+applications may require different preprocessing steps. Additionally, custom
+preprocessing can also improve the efficiency and accuracy of the model by
+removing unnecessary steps, reducing computational resources or adding steps
+that are important to the specific task being addressed.
+
+For example, tasks such as object detection or segmentation may require
+additional preprocessing steps such as resizing, cropping, or data augmentation
+to improve the robustness of the model. Below are some essential steps to
+customize a parser.
+
+### Instructions
+
+* **Create a Subclass**
+
+
+
+
+ Like Decoder, create `class Parser(parser.Parser)` in the same file.The
+`Parser` class should be a childclass of the
+[generic parser interface](https://github.com/tensorflow/models/blob/master/official/vision/dataloaders/parser.py)
+and must implement all the abstract methods. It should have the implementation
+of abstract methods `_parse_train_data` and `_parse_eval_data`, to generate
+images and labels for model training and evaluation respectively. The below example
+takes only two arguments but one can freely add as many arguments as needed.
+
+```python
+class Parser(parser.Parser):
+
+ def __init__(self, output_size: List[int], num_classes: float):
+
+ self._output_size = output_size
+ self._num_classes = num_classes
+ self._dtype = tf.float32
+
+ ....
+```
+
+
+
+Refer to the data parser and processing [class](https://github.com/tensorflow/models/blob/master/official/vision/dataloaders/maskrcnn_input.py) for Mask R-CNN for more complex cases. The class has multiple parameters related to data augmentation, masking, anchor boxes, data type of output image and more.
+
+
+
+
+
+* **Complete Abstract Methods**
+
+
+
+To define your own Parser, the user should override abstract functions
+[_parse_train_data](https://github.com/tensorflow/models/blob/b1a7752c5137822a32bd0dd70a0cb96e807ea411/official/vision/dataloaders/parser.py#L26)
+and
+[_parse_eval_data](https://github.com/tensorflow/models/blob/b1a7752c5137822a32bd0dd70a0cb96e807ea411/official/vision/dataloaders/parser.py#L39)
+of the
+[parser](https://github.com/tensorflow/models/blob/master/official/vision/dataloaders/parser.py)
+interface in the subclass, where decoded tensors are parsed with pre-processing
+steps for training and evaluation respectively. The output from the two
+functions can be any structure like a tuple, list or dictionary.
+
+```python
+ @abc.abstractmethod
+ def _parse_train_data(self, decoded_tensors):
+ """Generates images and labels that are usable for model training.
+
+ Args:
+ decoded_tensors: a dict of Tensors produced by the decoder.
+
+ Returns:
+ images: the image tensor.
+ labels: a dict of Tensors that contains labels.
+ """
+ pass
+
+ @abc.abstractmethod
+ def _parse_eval_data(self, decoded_tensors):
+ """Generates images and labels that are usable for model evaluation.
+
+ Args:
+ decoded_tensors: a dict of Tensors produced by the decoder.
+
+ Returns:
+ images: the image tensor.
+ labels: a dict of Tensors that contains labels.
+ """
+ pass
+
+```
+
+The input of `_parse_train_data` and `_parse_eval_data` is a dict of Tensors
+produced by the decoder; the output of these two functions is typically a tuple
+of (processe_image, processed_label). The user may perform any processing steps
+in these two functions as long as the interface is aligned. Note that the
+processing steps in `_parse_train_data` and `_parse_eval_data` are typically
+different since data augmentation is usually only applied to training. For
+Example, refer to the
+[Data parser](https://github.com/tensorflow/models/blob/b1a7752c5137822a32bd0dd70a0cb96e807ea411/official/vision/dataloaders/classification_input.py#L166)
+and processing steps for classification. We can observe that
+
+
+
+-For `_parse_train_data`, the following steps are performed
+
+ - Image decoding
+ - Random cropping
+ - Random flipping
+ - Color jittering
+ - Image resizing
+ - Auto-augmentation with autoaug, randaug etc.
+ - Image normalization
+
+