tf-models-nightly 2.17.0.dev20240528__py2.py3-none-any.whl → 2.20.0.dev20251205__py2.py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- official/__init__.py +1 -1
- official/common/__init__.py +1 -1
- official/common/dataset_fn.py +1 -1
- official/common/distribute_utils.py +27 -3
- official/common/distribute_utils_test.py +13 -12
- official/common/flags.py +28 -1
- official/common/registry_imports.py +1 -1
- official/common/streamz_counters.py +1 -1
- official/core/__init__.py +1 -1
- official/core/actions.py +1 -1
- official/core/actions_test.py +1 -1
- official/core/base_task.py +1 -1
- official/core/base_trainer.py +1 -1
- official/core/base_trainer_test.py +1 -1
- official/core/config_definitions.py +1 -1
- official/core/exp_factory.py +1 -1
- official/core/export_base.py +1 -1
- official/core/export_base_test.py +1 -1
- official/core/file_writers.py +1 -1
- official/core/file_writers_test.py +1 -1
- official/core/input_reader.py +1 -1
- official/core/registry.py +1 -1
- official/core/registry_test.py +1 -1
- official/core/savedmodel_checkpoint_manager.py +1 -1
- official/core/savedmodel_checkpoint_manager_test.py +1 -1
- official/core/task_factory.py +1 -1
- official/core/test_utils.py +1 -1
- official/core/tf_example_builder.py +1 -1
- official/core/tf_example_builder_test.py +1 -1
- official/core/tf_example_feature_key.py +1 -1
- official/core/tf_example_feature_key_test.py +1 -1
- official/core/train_lib.py +1 -3
- official/core/train_lib_test.py +1 -1
- official/core/train_utils.py +1 -1
- official/core/train_utils_test.py +1 -1
- official/legacy/__init__.py +1 -1
- official/legacy/albert/__init__.py +1 -1
- official/legacy/albert/configs.py +1 -1
- official/legacy/bert/__init__.py +1 -1
- official/legacy/bert/bert_models.py +1 -1
- official/legacy/bert/bert_models_test.py +1 -1
- official/legacy/bert/common_flags.py +1 -1
- official/legacy/bert/configs.py +1 -1
- official/legacy/bert/export_tfhub.py +1 -2
- official/legacy/bert/export_tfhub_test.py +1 -1
- official/legacy/bert/input_pipeline.py +1 -1
- official/legacy/bert/model_saving_utils.py +1 -1
- official/legacy/bert/model_training_utils.py +1 -1
- official/legacy/bert/model_training_utils_test.py +1 -1
- official/legacy/bert/run_classifier.py +1 -2
- official/legacy/bert/run_pretraining.py +1 -2
- official/legacy/bert/run_squad.py +1 -2
- official/legacy/bert/run_squad_helper.py +1 -1
- official/legacy/bert/serving.py +1 -1
- official/legacy/detection/__init__.py +1 -1
- official/legacy/detection/configs/__init__.py +1 -1
- official/legacy/detection/configs/base_config.py +1 -1
- official/legacy/detection/configs/factory.py +1 -1
- official/legacy/detection/configs/maskrcnn_config.py +1 -1
- official/legacy/detection/configs/olnmask_config.py +1 -1
- official/legacy/detection/configs/retinanet_config.py +1 -1
- official/legacy/detection/configs/shapemask_config.py +1 -1
- official/legacy/detection/dataloader/__init__.py +1 -1
- official/legacy/detection/dataloader/anchor.py +1 -1
- official/legacy/detection/dataloader/factory.py +1 -1
- official/legacy/detection/dataloader/input_reader.py +1 -1
- official/legacy/detection/dataloader/maskrcnn_parser.py +1 -1
- official/legacy/detection/dataloader/mode_keys.py +1 -1
- official/legacy/detection/dataloader/olnmask_parser.py +1 -1
- official/legacy/detection/dataloader/retinanet_parser.py +1 -1
- official/legacy/detection/dataloader/shapemask_parser.py +1 -1
- official/legacy/detection/dataloader/tf_example_decoder.py +1 -1
- official/legacy/detection/evaluation/__init__.py +1 -1
- official/legacy/detection/evaluation/coco_evaluator.py +1 -1
- official/legacy/detection/evaluation/coco_utils.py +1 -1
- official/legacy/detection/evaluation/factory.py +1 -1
- official/legacy/detection/executor/__init__.py +1 -1
- official/legacy/detection/executor/detection_executor.py +1 -1
- official/legacy/detection/executor/distributed_executor.py +1 -1
- official/legacy/detection/main.py +1 -1
- official/legacy/detection/modeling/__init__.py +1 -1
- official/legacy/detection/modeling/architecture/__init__.py +1 -1
- official/legacy/detection/modeling/architecture/factory.py +1 -1
- official/legacy/detection/modeling/architecture/fpn.py +1 -1
- official/legacy/detection/modeling/architecture/heads.py +1 -1
- official/legacy/detection/modeling/architecture/identity.py +1 -1
- official/legacy/detection/modeling/architecture/nn_blocks.py +1 -1
- official/legacy/detection/modeling/architecture/nn_ops.py +1 -1
- official/legacy/detection/modeling/architecture/resnet.py +1 -1
- official/legacy/detection/modeling/architecture/spinenet.py +1 -1
- official/legacy/detection/modeling/base_model.py +1 -1
- official/legacy/detection/modeling/checkpoint_utils.py +1 -1
- official/legacy/detection/modeling/factory.py +1 -1
- official/legacy/detection/modeling/learning_rates.py +1 -1
- official/legacy/detection/modeling/losses.py +1 -1
- official/legacy/detection/modeling/maskrcnn_model.py +1 -1
- official/legacy/detection/modeling/olnmask_model.py +1 -1
- official/legacy/detection/modeling/optimizers.py +1 -1
- official/legacy/detection/modeling/retinanet_model.py +1 -1
- official/legacy/detection/modeling/shapemask_model.py +1 -1
- official/legacy/detection/ops/__init__.py +1 -1
- official/legacy/detection/ops/nms.py +1 -1
- official/legacy/detection/ops/postprocess_ops.py +1 -1
- official/legacy/detection/ops/roi_ops.py +1 -1
- official/legacy/detection/ops/spatial_transform_ops.py +1 -1
- official/legacy/detection/ops/target_ops.py +1 -1
- official/legacy/detection/utils/__init__.py +1 -1
- official/legacy/detection/utils/box_utils.py +1 -1
- official/legacy/detection/utils/class_utils.py +1 -1
- official/legacy/detection/utils/dataloader_utils.py +1 -1
- official/legacy/detection/utils/input_utils.py +1 -1
- official/legacy/detection/utils/mask_utils.py +1 -1
- official/legacy/image_classification/__init__.py +1 -1
- official/legacy/image_classification/augment.py +1 -1
- official/legacy/image_classification/augment_test.py +1 -1
- official/legacy/image_classification/callbacks.py +1 -1
- official/legacy/image_classification/classifier_trainer.py +1 -1
- official/legacy/image_classification/classifier_trainer_test.py +1 -1
- official/legacy/image_classification/classifier_trainer_util_test.py +1 -1
- official/legacy/image_classification/configs/__init__.py +1 -1
- official/legacy/image_classification/configs/base_configs.py +1 -1
- official/legacy/image_classification/configs/configs.py +1 -1
- official/legacy/image_classification/dataset_factory.py +1 -1
- official/legacy/image_classification/efficientnet/__init__.py +1 -1
- official/legacy/image_classification/efficientnet/common_modules.py +1 -1
- official/legacy/image_classification/efficientnet/efficientnet_config.py +1 -1
- official/legacy/image_classification/efficientnet/efficientnet_model.py +1 -1
- official/legacy/image_classification/efficientnet/tfhub_export.py +1 -1
- official/legacy/image_classification/learning_rate.py +1 -1
- official/legacy/image_classification/learning_rate_test.py +1 -1
- official/legacy/image_classification/mnist_main.py +1 -2
- official/legacy/image_classification/mnist_test.py +1 -1
- official/legacy/image_classification/optimizer_factory.py +1 -1
- official/legacy/image_classification/optimizer_factory_test.py +1 -1
- official/legacy/image_classification/preprocessing.py +1 -1
- official/legacy/image_classification/resnet/__init__.py +1 -1
- official/legacy/image_classification/resnet/common.py +1 -1
- official/legacy/image_classification/resnet/imagenet_preprocessing.py +1 -1
- official/legacy/image_classification/resnet/resnet_config.py +1 -1
- official/legacy/image_classification/resnet/resnet_ctl_imagenet_main.py +1 -2
- official/legacy/image_classification/resnet/resnet_model.py +1 -1
- official/legacy/image_classification/resnet/resnet_runnable.py +1 -1
- official/legacy/image_classification/resnet/tfhub_export.py +1 -2
- official/legacy/image_classification/test_utils.py +1 -1
- official/legacy/image_classification/vgg/__init__.py +1 -1
- official/legacy/image_classification/vgg/vgg_config.py +1 -1
- official/legacy/image_classification/vgg/vgg_model.py +1 -1
- official/legacy/transformer/__init__.py +1 -1
- official/legacy/transformer/attention_layer.py +1 -1
- official/legacy/transformer/beam_search_v1.py +1 -1
- official/legacy/transformer/compute_bleu.py +1 -1
- official/legacy/transformer/compute_bleu_test.py +1 -1
- official/legacy/transformer/data_download.py +1 -1
- official/legacy/transformer/data_pipeline.py +1 -1
- official/legacy/transformer/embedding_layer.py +1 -1
- official/legacy/transformer/ffn_layer.py +1 -1
- official/legacy/transformer/metrics.py +1 -1
- official/legacy/transformer/misc.py +1 -1
- official/legacy/transformer/model_params.py +1 -1
- official/legacy/transformer/model_utils.py +1 -1
- official/legacy/transformer/model_utils_test.py +1 -1
- official/legacy/transformer/optimizer.py +1 -1
- official/legacy/transformer/transformer.py +1 -1
- official/legacy/transformer/transformer_forward_test.py +1 -1
- official/legacy/transformer/transformer_layers_test.py +1 -1
- official/legacy/transformer/transformer_main.py +1 -5
- official/legacy/transformer/transformer_main_test.py +1 -1
- official/legacy/transformer/transformer_test.py +1 -1
- official/legacy/transformer/translate.py +1 -2
- official/legacy/transformer/utils/__init__.py +1 -1
- official/legacy/transformer/utils/metrics.py +1 -1
- official/legacy/transformer/utils/tokenizer.py +1 -1
- official/legacy/transformer/utils/tokenizer_test.py +1 -1
- official/legacy/xlnet/__init__.py +1 -1
- official/legacy/xlnet/classifier_utils.py +1 -1
- official/legacy/xlnet/common_flags.py +1 -1
- official/legacy/xlnet/data_utils.py +1 -1
- official/legacy/xlnet/optimization.py +1 -1
- official/legacy/xlnet/preprocess_classification_data.py +1 -2
- official/legacy/xlnet/preprocess_pretrain_data.py +1 -2
- official/legacy/xlnet/preprocess_squad_data.py +1 -2
- official/legacy/xlnet/preprocess_utils.py +1 -1
- official/legacy/xlnet/run_classifier.py +1 -2
- official/legacy/xlnet/run_pretrain.py +1 -2
- official/legacy/xlnet/run_squad.py +1 -2
- official/legacy/xlnet/squad_utils.py +1 -1
- official/legacy/xlnet/training_utils.py +1 -1
- official/legacy/xlnet/xlnet_config.py +1 -1
- official/legacy/xlnet/xlnet_modeling.py +1 -1
- official/modeling/__init__.py +1 -1
- official/modeling/activations/__init__.py +1 -1
- official/modeling/activations/gelu.py +1 -1
- official/modeling/activations/gelu_test.py +1 -1
- official/modeling/activations/mish.py +1 -1
- official/modeling/activations/mish_test.py +1 -1
- official/modeling/activations/relu.py +1 -1
- official/modeling/activations/relu_test.py +1 -1
- official/modeling/activations/sigmoid.py +1 -1
- official/modeling/activations/sigmoid_test.py +1 -1
- official/modeling/activations/swish.py +1 -1
- official/modeling/activations/swish_test.py +1 -1
- official/modeling/grad_utils.py +1 -1
- official/modeling/grad_utils_test.py +1 -1
- official/modeling/hyperparams/__init__.py +1 -1
- official/modeling/hyperparams/base_config.py +27 -19
- official/modeling/hyperparams/base_config_test.py +32 -1
- official/modeling/hyperparams/oneof.py +1 -1
- official/modeling/hyperparams/oneof_test.py +1 -1
- official/modeling/hyperparams/params_dict.py +1 -1
- official/modeling/hyperparams/params_dict_test.py +1 -1
- official/modeling/multitask/__init__.py +1 -1
- official/modeling/multitask/base_model.py +1 -1
- official/modeling/multitask/base_trainer.py +1 -1
- official/modeling/multitask/base_trainer_test.py +1 -1
- official/modeling/multitask/configs.py +3 -3
- official/modeling/multitask/evaluator.py +1 -1
- official/modeling/multitask/evaluator_test.py +1 -1
- official/modeling/multitask/interleaving_trainer.py +1 -1
- official/modeling/multitask/interleaving_trainer_test.py +1 -1
- official/modeling/multitask/multitask.py +1 -1
- official/modeling/multitask/task_sampler.py +1 -1
- official/modeling/multitask/task_sampler_test.py +1 -1
- official/modeling/multitask/test_utils.py +1 -1
- official/modeling/multitask/train_lib.py +81 -14
- official/modeling/multitask/train_lib_test.py +1 -1
- official/modeling/optimization/__init__.py +1 -1
- official/modeling/optimization/adafactor_optimizer.py +1 -1
- official/modeling/optimization/configs/__init__.py +1 -1
- official/modeling/optimization/configs/learning_rate_config.py +1 -1
- official/modeling/optimization/configs/optimization_config.py +1 -1
- official/modeling/optimization/configs/optimization_config_test.py +1 -1
- official/modeling/optimization/configs/optimizer_config.py +1 -1
- official/modeling/optimization/ema_optimizer.py +1 -1
- official/modeling/optimization/lamb.py +1 -1
- official/modeling/optimization/lamb_test.py +1 -1
- official/modeling/optimization/lars.py +1 -1
- official/modeling/optimization/legacy_adamw.py +1 -1
- official/modeling/optimization/lr_schedule.py +1 -1
- official/modeling/optimization/lr_schedule_test.py +1 -1
- official/modeling/optimization/optimizer_factory.py +1 -1
- official/modeling/optimization/optimizer_factory_test.py +1 -1
- official/modeling/optimization/slide_optimizer.py +1 -1
- official/modeling/performance.py +1 -1
- official/modeling/privacy/__init__.py +1 -1
- official/modeling/privacy/configs.py +1 -1
- official/modeling/privacy/configs_test.py +1 -1
- official/modeling/privacy/ops.py +1 -1
- official/modeling/privacy/ops_test.py +1 -1
- official/modeling/tf_utils.py +1 -1
- official/modeling/tf_utils_test.py +1 -1
- official/nlp/__init__.py +1 -1
- official/nlp/configs/__init__.py +1 -1
- official/nlp/configs/bert.py +1 -1
- official/nlp/configs/electra.py +1 -1
- official/nlp/configs/encoders.py +1 -1
- official/nlp/configs/encoders_test.py +1 -1
- official/nlp/configs/experiment_configs.py +1 -1
- official/nlp/configs/finetuning_experiments.py +1 -1
- official/nlp/configs/pretraining_experiments.py +1 -1
- official/nlp/configs/wmt_transformer_experiments.py +1 -1
- official/nlp/continuous_finetune_lib.py +1 -1
- official/nlp/continuous_finetune_lib_test.py +1 -1
- official/nlp/data/__init__.py +1 -1
- official/nlp/data/classifier_data_lib.py +1 -1
- official/nlp/data/classifier_data_lib_test.py +1 -1
- official/nlp/data/create_finetuning_data.py +1 -2
- official/nlp/data/create_pretraining_data.py +1 -3
- official/nlp/data/create_pretraining_data_test.py +1 -1
- official/nlp/data/create_xlnet_pretraining_data.py +1 -3
- official/nlp/data/create_xlnet_pretraining_data_test.py +1 -1
- official/nlp/data/data_loader.py +1 -1
- official/nlp/data/data_loader_factory.py +1 -1
- official/nlp/data/data_loader_factory_test.py +1 -1
- official/nlp/data/dual_encoder_dataloader.py +1 -1
- official/nlp/data/dual_encoder_dataloader_test.py +1 -1
- official/nlp/data/pretrain_dataloader.py +1 -1
- official/nlp/data/pretrain_dataloader_test.py +1 -1
- official/nlp/data/pretrain_dynamic_dataloader.py +1 -1
- official/nlp/data/pretrain_dynamic_dataloader_test.py +1 -1
- official/nlp/data/pretrain_text_dataloader.py +1 -1
- official/nlp/data/question_answering_dataloader.py +1 -1
- official/nlp/data/question_answering_dataloader_test.py +1 -1
- official/nlp/data/sentence_prediction_dataloader.py +1 -1
- official/nlp/data/sentence_prediction_dataloader_test.py +1 -1
- official/nlp/data/sentence_retrieval_lib.py +1 -1
- official/nlp/data/squad_lib.py +1 -1
- official/nlp/data/squad_lib_sp.py +1 -1
- official/nlp/data/tagging_data_lib.py +1 -1
- official/nlp/data/tagging_data_lib_test.py +1 -1
- official/nlp/data/tagging_dataloader.py +1 -1
- official/nlp/data/tagging_dataloader_test.py +1 -1
- official/nlp/data/train_sentencepiece.py +1 -1
- official/nlp/data/wmt_dataloader.py +1 -1
- official/nlp/data/wmt_dataloader_test.py +1 -1
- official/nlp/metrics/__init__.py +1 -1
- official/nlp/metrics/bleu.py +1 -1
- official/nlp/metrics/bleu_test.py +1 -1
- official/nlp/modeling/__init__.py +1 -1
- official/nlp/modeling/layers/__init__.py +3 -1
- official/nlp/modeling/layers/attention.py +1 -1
- official/nlp/modeling/layers/attention_test.py +1 -1
- official/nlp/modeling/layers/bigbird_attention.py +1 -1
- official/nlp/modeling/layers/bigbird_attention_test.py +1 -1
- official/nlp/modeling/layers/block_diag_feedforward.py +1 -1
- official/nlp/modeling/layers/block_diag_feedforward_test.py +1 -1
- official/nlp/modeling/layers/block_sparse_attention.py +359 -0
- official/nlp/modeling/layers/block_sparse_attention_test.py +433 -0
- official/nlp/modeling/layers/cls_head.py +1 -1
- official/nlp/modeling/layers/cls_head_test.py +1 -1
- official/nlp/modeling/layers/factorized_embedding.py +1 -1
- official/nlp/modeling/layers/factorized_embedding_test.py +1 -1
- official/nlp/modeling/layers/gated_feedforward.py +2 -2
- official/nlp/modeling/layers/gated_feedforward_test.py +1 -1
- official/nlp/modeling/layers/gaussian_process.py +1 -1
- official/nlp/modeling/layers/gaussian_process_test.py +1 -1
- official/nlp/modeling/layers/kernel_attention.py +1 -1
- official/nlp/modeling/layers/kernel_attention_test.py +1 -1
- official/nlp/modeling/layers/masked_lm.py +1 -1
- official/nlp/modeling/layers/masked_lm_test.py +1 -1
- official/nlp/modeling/layers/masked_softmax.py +1 -1
- official/nlp/modeling/layers/masked_softmax_test.py +1 -1
- official/nlp/modeling/layers/mat_mul_with_margin.py +1 -2
- official/nlp/modeling/layers/mat_mul_with_margin_test.py +1 -1
- official/nlp/modeling/layers/mixing.py +1 -1
- official/nlp/modeling/layers/mixing_test.py +1 -1
- official/nlp/modeling/layers/mobile_bert_layers.py +1 -1
- official/nlp/modeling/layers/mobile_bert_layers_test.py +1 -1
- official/nlp/modeling/layers/moe.py +1 -1
- official/nlp/modeling/layers/moe_test.py +1 -1
- official/nlp/modeling/layers/multi_channel_attention.py +1 -1
- official/nlp/modeling/layers/multi_channel_attention_test.py +1 -1
- official/nlp/modeling/layers/multi_query_attention.py +426 -0
- official/nlp/modeling/layers/multi_query_attention_test.py +415 -0
- official/nlp/modeling/layers/on_device_embedding.py +1 -1
- official/nlp/modeling/layers/on_device_embedding_test.py +1 -1
- official/nlp/modeling/layers/pack_optimization.py +9 -1
- official/nlp/modeling/layers/pack_optimization_test.py +1 -1
- official/nlp/modeling/layers/per_dim_scale_attention.py +1 -1
- official/nlp/modeling/layers/per_dim_scale_attention_test.py +1 -1
- official/nlp/modeling/layers/position_embedding.py +1 -1
- official/nlp/modeling/layers/position_embedding_test.py +1 -1
- official/nlp/modeling/layers/relative_attention.py +1 -1
- official/nlp/modeling/layers/relative_attention_test.py +1 -1
- official/nlp/modeling/layers/reuse_attention.py +1 -1
- official/nlp/modeling/layers/reuse_attention_test.py +1 -1
- official/nlp/modeling/layers/reuse_transformer.py +1 -1
- official/nlp/modeling/layers/reuse_transformer_test.py +1 -1
- official/nlp/modeling/layers/rezero_transformer.py +89 -21
- official/nlp/modeling/layers/rezero_transformer_test.py +64 -1
- official/nlp/modeling/layers/routing.py +1 -1
- official/nlp/modeling/layers/routing_test.py +1 -1
- official/nlp/modeling/layers/self_attention_mask.py +1 -1
- official/nlp/modeling/layers/spectral_normalization.py +1 -1
- official/nlp/modeling/layers/spectral_normalization_test.py +1 -1
- official/nlp/modeling/layers/talking_heads_attention.py +1 -1
- official/nlp/modeling/layers/talking_heads_attention_test.py +1 -1
- official/nlp/modeling/layers/text_layers.py +1 -1
- official/nlp/modeling/layers/text_layers_test.py +1 -1
- official/nlp/modeling/layers/tn_expand_condense.py +1 -1
- official/nlp/modeling/layers/tn_expand_condense_test.py +1 -1
- official/nlp/modeling/layers/tn_transformer_expand_condense.py +1 -3
- official/nlp/modeling/layers/tn_transformer_test.py +1 -1
- official/nlp/modeling/layers/transformer.py +1 -1
- official/nlp/modeling/layers/transformer_encoder_block.py +313 -52
- official/nlp/modeling/layers/transformer_encoder_block_test.py +291 -9
- official/nlp/modeling/layers/transformer_scaffold.py +1 -1
- official/nlp/modeling/layers/transformer_scaffold_test.py +1 -1
- official/nlp/modeling/layers/transformer_test.py +1 -1
- official/nlp/modeling/layers/transformer_xl.py +1 -1
- official/nlp/modeling/layers/transformer_xl_test.py +1 -1
- official/nlp/modeling/layers/util.py +1 -1
- official/nlp/modeling/losses/__init__.py +1 -1
- official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy.py +1 -1
- official/nlp/modeling/losses/weighted_sparse_categorical_crossentropy_test.py +1 -1
- official/nlp/modeling/models/__init__.py +1 -1
- official/nlp/modeling/models/bert_classifier.py +1 -1
- official/nlp/modeling/models/bert_classifier_test.py +1 -1
- official/nlp/modeling/models/bert_pretrainer.py +1 -1
- official/nlp/modeling/models/bert_pretrainer_test.py +1 -1
- official/nlp/modeling/models/bert_span_labeler.py +1 -1
- official/nlp/modeling/models/bert_span_labeler_test.py +1 -1
- official/nlp/modeling/models/bert_token_classifier.py +1 -1
- official/nlp/modeling/models/bert_token_classifier_test.py +1 -1
- official/nlp/modeling/models/dual_encoder.py +1 -1
- official/nlp/modeling/models/dual_encoder_test.py +1 -1
- official/nlp/modeling/models/electra_pretrainer.py +1 -1
- official/nlp/modeling/models/electra_pretrainer_test.py +1 -1
- official/nlp/modeling/models/seq2seq_transformer.py +1 -1
- official/nlp/modeling/models/seq2seq_transformer_test.py +1 -1
- official/nlp/modeling/models/t5.py +1 -1
- official/nlp/modeling/models/t5_test.py +1 -1
- official/nlp/modeling/models/xlnet.py +1 -1
- official/nlp/modeling/models/xlnet_test.py +1 -1
- official/nlp/modeling/networks/__init__.py +1 -1
- official/nlp/modeling/networks/albert_encoder.py +1 -1
- official/nlp/modeling/networks/albert_encoder_test.py +1 -1
- official/nlp/modeling/networks/bert_dense_encoder_test.py +1 -2
- official/nlp/modeling/networks/bert_encoder.py +1 -1
- official/nlp/modeling/networks/bert_encoder_test.py +1 -2
- official/nlp/modeling/networks/classification.py +1 -1
- official/nlp/modeling/networks/classification_test.py +1 -1
- official/nlp/modeling/networks/encoder_scaffold.py +1 -1
- official/nlp/modeling/networks/encoder_scaffold_test.py +1 -1
- official/nlp/modeling/networks/fnet.py +1 -1
- official/nlp/modeling/networks/fnet_test.py +1 -1
- official/nlp/modeling/networks/funnel_transformer.py +1 -1
- official/nlp/modeling/networks/funnel_transformer_test.py +1 -1
- official/nlp/modeling/networks/mobile_bert_encoder.py +6 -4
- official/nlp/modeling/networks/mobile_bert_encoder_test.py +1 -1
- official/nlp/modeling/networks/packed_sequence_embedding.py +1 -1
- official/nlp/modeling/networks/packed_sequence_embedding_test.py +1 -3
- official/nlp/modeling/networks/span_labeling.py +1 -1
- official/nlp/modeling/networks/span_labeling_test.py +1 -1
- official/nlp/modeling/networks/sparse_mixer.py +1 -1
- official/nlp/modeling/networks/sparse_mixer_test.py +1 -1
- official/nlp/modeling/networks/xlnet_base.py +1 -1
- official/nlp/modeling/networks/xlnet_base_test.py +1 -1
- official/nlp/modeling/ops/__init__.py +1 -1
- official/nlp/modeling/ops/beam_search.py +1 -1
- official/nlp/modeling/ops/beam_search_test.py +1 -1
- official/nlp/modeling/ops/decoding_module.py +1 -1
- official/nlp/modeling/ops/decoding_module_test.py +1 -1
- official/nlp/modeling/ops/sampling_module.py +3 -3
- official/nlp/modeling/ops/segment_extractor.py +1 -1
- official/nlp/modeling/ops/segment_extractor_test.py +1 -1
- official/nlp/optimization.py +1 -1
- official/nlp/serving/__init__.py +1 -1
- official/nlp/serving/export_savedmodel.py +1 -1
- official/nlp/serving/export_savedmodel_test.py +1 -1
- official/nlp/serving/export_savedmodel_util.py +1 -1
- official/nlp/serving/serving_modules.py +1 -1
- official/nlp/serving/serving_modules_test.py +1 -1
- official/nlp/tasks/__init__.py +1 -1
- official/nlp/tasks/dual_encoder.py +1 -2
- official/nlp/tasks/dual_encoder_test.py +1 -1
- official/nlp/tasks/electra_task.py +1 -1
- official/nlp/tasks/electra_task_test.py +1 -1
- official/nlp/tasks/masked_lm.py +1 -1
- official/nlp/tasks/masked_lm_determinism_test.py +1 -1
- official/nlp/tasks/masked_lm_test.py +1 -1
- official/nlp/tasks/question_answering.py +1 -1
- official/nlp/tasks/question_answering_test.py +1 -1
- official/nlp/tasks/sentence_prediction.py +1 -1
- official/nlp/tasks/sentence_prediction_test.py +1 -1
- official/nlp/tasks/tagging.py +1 -1
- official/nlp/tasks/tagging_test.py +1 -1
- official/nlp/tasks/translation.py +1 -1
- official/nlp/tasks/translation_test.py +1 -1
- official/nlp/tasks/utils.py +1 -1
- official/nlp/tools/__init__.py +1 -1
- official/nlp/tools/export_tfhub.py +1 -1
- official/nlp/tools/export_tfhub_lib.py +1 -2
- official/nlp/tools/export_tfhub_lib_test.py +1 -1
- official/nlp/tools/squad_evaluate_v1_1.py +1 -1
- official/nlp/tools/squad_evaluate_v2_0.py +1 -1
- official/nlp/tools/tf1_bert_checkpoint_converter_lib.py +1 -1
- official/nlp/tools/tf2_albert_encoder_checkpoint_converter.py +1 -1
- official/nlp/tools/tf2_bert_encoder_checkpoint_converter.py +1 -1
- official/nlp/tools/tokenization.py +1 -1
- official/nlp/tools/tokenization_test.py +1 -1
- official/nlp/train.py +1 -1
- official/projects/__init__.py +1 -1
- official/projects/bigbird/__init__.py +1 -1
- official/projects/bigbird/encoder.py +1 -1
- official/projects/bigbird/encoder_test.py +1 -1
- official/projects/bigbird/experiment_configs.py +1 -1
- official/projects/bigbird/recompute_grad.py +1 -1
- official/projects/bigbird/recomputing_dropout.py +1 -1
- official/projects/bigbird/stateless_dropout.py +1 -1
- official/projects/centernet/__init__.py +1 -1
- official/projects/centernet/common/__init__.py +1 -1
- official/projects/centernet/common/registry_imports.py +1 -1
- official/projects/centernet/configs/__init__.py +1 -1
- official/projects/centernet/configs/backbones.py +1 -1
- official/projects/centernet/configs/centernet.py +1 -1
- official/projects/centernet/configs/centernet_test.py +1 -1
- official/projects/centernet/dataloaders/__init__.py +1 -1
- official/projects/centernet/dataloaders/centernet_input.py +1 -1
- official/projects/centernet/losses/__init__.py +1 -1
- official/projects/centernet/losses/centernet_losses.py +1 -1
- official/projects/centernet/losses/centernet_losses_test.py +1 -1
- official/projects/centernet/modeling/__init__.py +1 -1
- official/projects/centernet/modeling/backbones/__init__.py +1 -1
- official/projects/centernet/modeling/backbones/hourglass.py +1 -1
- official/projects/centernet/modeling/backbones/hourglass_test.py +1 -1
- official/projects/centernet/modeling/centernet_model.py +2 -2
- official/projects/centernet/modeling/centernet_model_test.py +1 -1
- official/projects/centernet/modeling/heads/__init__.py +1 -1
- official/projects/centernet/modeling/heads/centernet_head.py +2 -2
- official/projects/centernet/modeling/heads/centernet_head_test.py +1 -1
- official/projects/centernet/modeling/layers/__init__.py +1 -1
- official/projects/centernet/modeling/layers/cn_nn_blocks.py +1 -1
- official/projects/centernet/modeling/layers/cn_nn_blocks_test.py +1 -1
- official/projects/centernet/modeling/layers/detection_generator.py +1 -1
- official/projects/centernet/modeling/layers/detection_generator_test.py +1 -1
- official/projects/centernet/ops/__init__.py +1 -1
- official/projects/centernet/ops/box_list.py +1 -1
- official/projects/centernet/ops/box_list_ops.py +1 -1
- official/projects/centernet/ops/loss_ops.py +1 -1
- official/projects/centernet/ops/nms_ops.py +1 -1
- official/projects/centernet/ops/preprocess_ops.py +1 -1
- official/projects/centernet/ops/target_assigner.py +1 -1
- official/projects/centernet/ops/target_assigner_test.py +1 -1
- official/projects/centernet/tasks/__init__.py +1 -1
- official/projects/centernet/tasks/centernet.py +1 -1
- official/projects/centernet/train.py +1 -1
- official/projects/centernet/utils/__init__.py +1 -1
- official/projects/centernet/utils/checkpoints/__init__.py +1 -1
- official/projects/centernet/utils/checkpoints/config_classes.py +1 -1
- official/projects/centernet/utils/checkpoints/config_data.py +1 -1
- official/projects/centernet/utils/checkpoints/load_weights.py +1 -1
- official/projects/centernet/utils/checkpoints/read_checkpoints.py +1 -1
- official/projects/centernet/utils/tf2_centernet_checkpoint_converter.py +1 -1
- official/projects/deepmac_maskrcnn/__init__.py +1 -1
- official/projects/deepmac_maskrcnn/common/__init__.py +1 -1
- official/projects/deepmac_maskrcnn/common/registry_imports.py +1 -1
- official/projects/deepmac_maskrcnn/configs/__init__.py +1 -1
- official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn.py +1 -1
- official/projects/deepmac_maskrcnn/configs/deep_mask_head_rcnn_config_test.py +1 -1
- official/projects/deepmac_maskrcnn/modeling/__init__.py +1 -1
- official/projects/deepmac_maskrcnn/modeling/heads/__init__.py +1 -1
- official/projects/deepmac_maskrcnn/modeling/heads/hourglass_network.py +1 -1
- official/projects/deepmac_maskrcnn/modeling/heads/instance_heads.py +1 -3
- official/projects/deepmac_maskrcnn/modeling/heads/instance_heads_test.py +1 -2
- official/projects/deepmac_maskrcnn/modeling/maskrcnn_model.py +1 -3
- official/projects/deepmac_maskrcnn/modeling/maskrcnn_model_test.py +1 -3
- official/projects/deepmac_maskrcnn/serving/__init__.py +1 -1
- official/projects/deepmac_maskrcnn/serving/detection.py +1 -1
- official/projects/deepmac_maskrcnn/serving/detection_test.py +1 -1
- official/projects/deepmac_maskrcnn/serving/export_saved_model.py +1 -1
- official/projects/deepmac_maskrcnn/tasks/__init__.py +1 -1
- official/projects/deepmac_maskrcnn/tasks/deep_mask_head_rcnn.py +1 -1
- official/projects/deepmac_maskrcnn/train.py +1 -1
- official/projects/detr/__init__.py +14 -0
- official/projects/detr/configs/__init__.py +14 -0
- official/projects/detr/configs/detr.py +277 -0
- official/projects/detr/configs/detr_test.py +51 -0
- official/projects/detr/dataloaders/__init__.py +14 -0
- official/projects/detr/dataloaders/coco.py +157 -0
- official/projects/detr/dataloaders/coco_test.py +111 -0
- official/projects/detr/dataloaders/detr_input.py +175 -0
- official/projects/detr/experiments/__init__.py +14 -0
- official/projects/detr/modeling/__init__.py +14 -0
- official/projects/detr/modeling/detr.py +345 -0
- official/projects/detr/modeling/detr_test.py +70 -0
- official/projects/detr/modeling/transformer.py +849 -0
- official/projects/detr/modeling/transformer_test.py +263 -0
- official/projects/detr/ops/__init__.py +14 -0
- official/projects/detr/ops/matchers.py +489 -0
- official/projects/detr/ops/matchers_test.py +95 -0
- official/projects/detr/optimization.py +151 -0
- official/projects/detr/serving/__init__.py +14 -0
- official/projects/detr/serving/export_module.py +103 -0
- official/projects/detr/serving/export_module_test.py +98 -0
- official/projects/detr/serving/export_saved_model.py +109 -0
- official/projects/detr/tasks/__init__.py +14 -0
- official/projects/detr/tasks/detection.py +433 -0
- official/projects/detr/tasks/detection_test.py +203 -0
- official/projects/detr/train.py +70 -0
- official/projects/maskconver/__init__.py +14 -0
- official/projects/maskconver/configs/__init__.py +14 -0
- official/projects/maskconver/configs/backbones.py +43 -0
- official/projects/maskconver/configs/decoders.py +36 -0
- official/projects/maskconver/configs/maskconver.py +523 -0
- official/projects/maskconver/configs/multiscale_maskconver.py +215 -0
- official/projects/maskconver/tasks/__init__.py +14 -0
- official/projects/maskconver/tasks/maskconver.py +641 -0
- official/projects/maskconver/tasks/multiscale_maskconver.py +278 -0
- official/projects/maskconver/train.py +30 -0
- official/projects/maxvit/__init__.py +1 -1
- official/projects/maxvit/configs/__init__.py +1 -1
- official/projects/maxvit/configs/backbones.py +1 -1
- official/projects/maxvit/configs/image_classification.py +1 -1
- official/projects/maxvit/configs/image_classification_test.py +1 -1
- official/projects/maxvit/configs/rcnn.py +1 -1
- official/projects/maxvit/configs/rcnn_test.py +1 -1
- official/projects/maxvit/configs/retinanet.py +1 -1
- official/projects/maxvit/configs/retinanet_test.py +1 -1
- official/projects/maxvit/configs/semantic_segmentation.py +1 -1
- official/projects/maxvit/configs/semantic_segmentation_test.py +1 -1
- official/projects/maxvit/modeling/__init__.py +1 -1
- official/projects/maxvit/modeling/common_ops.py +14 -1
- official/projects/maxvit/modeling/layers.py +1 -1
- official/projects/maxvit/modeling/maxvit.py +2 -2
- official/projects/maxvit/modeling/maxvit_test.py +1 -1
- official/projects/maxvit/registry_imports.py +1 -1
- official/projects/maxvit/train.py +1 -1
- official/projects/maxvit/train_test.py +1 -1
- official/projects/mobilebert/__init__.py +1 -1
- official/projects/mobilebert/distillation.py +1 -1
- official/projects/mobilebert/distillation_test.py +1 -1
- official/projects/mobilebert/export_tfhub.py +1 -1
- official/projects/mobilebert/model_utils.py +1 -1
- official/projects/mobilebert/run_distillation.py +1 -1
- official/projects/mobilebert/tf2_model_checkpoint_converter.py +1 -1
- official/projects/mobilebert/utils.py +1 -1
- official/projects/movinet/__init__.py +1 -1
- official/projects/movinet/configs/__init__.py +1 -1
- official/projects/movinet/configs/movinet.py +1 -1
- official/projects/movinet/configs/movinet_test.py +1 -1
- official/projects/movinet/modeling/__init__.py +1 -1
- official/projects/movinet/modeling/movinet.py +1 -1
- official/projects/movinet/modeling/movinet_layers.py +1 -1
- official/projects/movinet/modeling/movinet_layers_test.py +1 -1
- official/projects/movinet/modeling/movinet_model.py +1 -1
- official/projects/movinet/modeling/movinet_model_test.py +1 -1
- official/projects/movinet/modeling/movinet_test.py +1 -1
- official/projects/movinet/tools/__init__.py +1 -1
- official/projects/movinet/tools/convert_3d_2plus1d.py +1 -1
- official/projects/movinet/tools/convert_3d_2plus1d_test.py +1 -1
- official/projects/movinet/tools/export_saved_model.py +1 -1
- official/projects/movinet/tools/export_saved_model_test.py +6 -3
- official/projects/movinet/tools/quantize_movinet.py +1 -1
- official/projects/movinet/train.py +1 -1
- official/projects/movinet/train_test.py +1 -1
- official/projects/nhnet/__init__.py +1 -1
- official/projects/nhnet/configs.py +1 -1
- official/projects/nhnet/configs_test.py +1 -1
- official/projects/nhnet/decoder.py +1 -1
- official/projects/nhnet/decoder_test.py +1 -1
- official/projects/nhnet/evaluation.py +1 -3
- official/projects/nhnet/input_pipeline.py +1 -1
- official/projects/nhnet/models.py +1 -1
- official/projects/nhnet/models_test.py +1 -1
- official/projects/nhnet/optimizer.py +1 -1
- official/projects/nhnet/raw_data_process.py +1 -1
- official/projects/nhnet/raw_data_processor.py +1 -1
- official/projects/nhnet/trainer.py +1 -6
- official/projects/nhnet/trainer_test.py +1 -1
- official/projects/nhnet/utils.py +1 -1
- official/projects/panoptic/__init__.py +1 -1
- official/projects/panoptic/configs/__init__.py +1 -1
- official/projects/panoptic/configs/panoptic_deeplab.py +5 -6
- official/projects/panoptic/configs/panoptic_maskrcnn.py +1 -1
- official/projects/panoptic/tasks/__init__.py +1 -1
- official/projects/panoptic/tasks/panoptic_deeplab.py +1 -1
- official/projects/panoptic/tasks/panoptic_maskrcnn.py +3 -1
- official/projects/panoptic/train.py +1 -1
- official/projects/qat/__init__.py +1 -1
- official/projects/qat/nlp/__init__.py +1 -1
- official/projects/qat/nlp/configs/__init__.py +1 -1
- official/projects/qat/nlp/configs/finetuning_experiments.py +1 -1
- official/projects/qat/nlp/modeling/__init__.py +1 -1
- official/projects/qat/nlp/modeling/layers/__init__.py +1 -1
- official/projects/qat/nlp/modeling/layers/mobile_bert_layers.py +1 -1
- official/projects/qat/nlp/modeling/layers/multi_head_attention.py +1 -1
- official/projects/qat/nlp/modeling/layers/transformer_encoder_block.py +1 -1
- official/projects/qat/nlp/modeling/layers/transformer_encoder_block_test.py +1 -1
- official/projects/qat/nlp/modeling/models/__init__.py +1 -1
- official/projects/qat/nlp/modeling/models/bert_span_labeler.py +1 -1
- official/projects/qat/nlp/modeling/networks/__init__.py +1 -1
- official/projects/qat/nlp/modeling/networks/span_labeling.py +1 -1
- official/projects/qat/nlp/pretrained_checkpoint_converter.py +1 -3
- official/projects/qat/nlp/quantization/__init__.py +1 -1
- official/projects/qat/nlp/quantization/configs.py +1 -1
- official/projects/qat/nlp/quantization/configs_test.py +1 -2
- official/projects/qat/nlp/quantization/helper.py +1 -1
- official/projects/qat/nlp/quantization/schemes.py +1 -3
- official/projects/qat/nlp/quantization/wrappers.py +1 -1
- official/projects/qat/nlp/registry_imports.py +1 -1
- official/projects/qat/nlp/tasks/__init__.py +1 -1
- official/projects/qat/nlp/tasks/question_answering.py +1 -1
- official/projects/qat/nlp/tasks/question_answering_test.py +1 -1
- official/projects/qat/nlp/train.py +1 -1
- official/projects/qat/vision/__init__.py +1 -1
- official/projects/qat/vision/configs/__init__.py +1 -1
- official/projects/qat/vision/configs/common.py +1 -1
- official/projects/qat/vision/configs/image_classification.py +1 -1
- official/projects/qat/vision/configs/image_classification_test.py +1 -1
- official/projects/qat/vision/configs/retinanet.py +1 -1
- official/projects/qat/vision/configs/retinanet_test.py +1 -1
- official/projects/qat/vision/configs/semantic_segmentation.py +1 -1
- official/projects/qat/vision/configs/semantic_segmentation_test.py +1 -1
- official/projects/qat/vision/modeling/__init__.py +1 -1
- official/projects/qat/vision/modeling/factory.py +1 -3
- official/projects/qat/vision/modeling/factory_test.py +1 -3
- official/projects/qat/vision/modeling/heads/__init__.py +1 -1
- official/projects/qat/vision/modeling/heads/dense_prediction_heads.py +1 -3
- official/projects/qat/vision/modeling/heads/dense_prediction_heads_test.py +1 -2
- official/projects/qat/vision/modeling/layers/__init__.py +1 -1
- official/projects/qat/vision/modeling/layers/nn_blocks.py +1 -3
- official/projects/qat/vision/modeling/layers/nn_blocks_test.py +1 -2
- official/projects/qat/vision/modeling/layers/nn_layers.py +2 -2
- official/projects/qat/vision/modeling/layers/nn_layers_test.py +1 -2
- official/projects/qat/vision/modeling/segmentation_model.py +1 -2
- official/projects/qat/vision/n_bit/__init__.py +1 -1
- official/projects/qat/vision/n_bit/configs.py +1 -1
- official/projects/qat/vision/n_bit/configs_test.py +1 -3
- official/projects/qat/vision/n_bit/nn_blocks.py +1 -3
- official/projects/qat/vision/n_bit/nn_blocks_test.py +1 -2
- official/projects/qat/vision/n_bit/nn_layers.py +1 -1
- official/projects/qat/vision/n_bit/schemes.py +1 -3
- official/projects/qat/vision/quantization/__init__.py +1 -1
- official/projects/qat/vision/quantization/configs.py +1 -1
- official/projects/qat/vision/quantization/configs_test.py +1 -3
- official/projects/qat/vision/quantization/helper.py +1 -1
- official/projects/qat/vision/quantization/helper_test.py +1 -1
- official/projects/qat/vision/quantization/layer_transforms.py +1 -1
- official/projects/qat/vision/quantization/schemes.py +1 -3
- official/projects/qat/vision/registry_imports.py +1 -1
- official/projects/qat/vision/serving/__init__.py +1 -1
- official/projects/qat/vision/serving/export_module.py +1 -1
- official/projects/qat/vision/serving/export_saved_model.py +1 -1
- official/projects/qat/vision/serving/export_tflite.py +1 -1
- official/projects/qat/vision/tasks/__init__.py +1 -1
- official/projects/qat/vision/tasks/image_classification.py +1 -1
- official/projects/qat/vision/tasks/image_classification_test.py +1 -1
- official/projects/qat/vision/tasks/retinanet.py +1 -1
- official/projects/qat/vision/tasks/retinanet_test.py +1 -1
- official/projects/qat/vision/tasks/semantic_segmentation.py +1 -1
- official/projects/qat/vision/train.py +1 -1
- official/projects/roformer/__init__.py +1 -1
- official/projects/roformer/roformer.py +1 -1
- official/projects/roformer/roformer_attention.py +1 -1
- official/projects/roformer/roformer_attention_test.py +1 -1
- official/projects/roformer/roformer_encoder.py +1 -1
- official/projects/roformer/roformer_encoder_block.py +1 -1
- official/projects/roformer/roformer_encoder_block_test.py +1 -1
- official/projects/roformer/roformer_encoder_test.py +1 -1
- official/projects/roformer/roformer_experiments.py +1 -1
- official/projects/roformer/train.py +1 -1
- official/projects/teams/__init__.py +1 -1
- official/projects/teams/teams.py +1 -1
- official/projects/teams/teams_experiments.py +1 -1
- official/projects/teams/teams_pretrainer.py +1 -1
- official/projects/teams/teams_pretrainer_test.py +1 -1
- official/projects/teams/teams_task.py +1 -1
- official/projects/teams/teams_task_test.py +1 -1
- official/projects/teams/train.py +1 -1
- official/projects/triviaqa/__init__.py +1 -1
- official/projects/triviaqa/dataset.py +1 -1
- official/projects/triviaqa/download_and_prepare.py +1 -1
- official/projects/triviaqa/evaluate.py +1 -1
- official/projects/triviaqa/evaluation.py +1 -1
- official/projects/triviaqa/inputs.py +1 -1
- official/projects/triviaqa/modeling.py +1 -1
- official/projects/triviaqa/predict.py +1 -1
- official/projects/triviaqa/prediction.py +1 -1
- official/projects/triviaqa/preprocess.py +1 -1
- official/projects/triviaqa/sentencepiece_pb2.py +1 -1
- official/projects/triviaqa/train.py +1 -1
- official/projects/video_ssl/__init__.py +1 -1
- official/projects/video_ssl/configs/__init__.py +1 -1
- official/projects/video_ssl/configs/video_ssl.py +1 -1
- official/projects/video_ssl/configs/video_ssl_test.py +1 -1
- official/projects/video_ssl/dataloaders/__init__.py +1 -1
- official/projects/video_ssl/dataloaders/video_ssl_input.py +1 -1
- official/projects/video_ssl/dataloaders/video_ssl_input_test.py +1 -2
- official/projects/video_ssl/losses/__init__.py +1 -1
- official/projects/video_ssl/losses/losses.py +1 -2
- official/projects/video_ssl/modeling/__init__.py +1 -1
- official/projects/video_ssl/modeling/video_ssl_model.py +1 -3
- official/projects/video_ssl/ops/__init__.py +1 -1
- official/projects/video_ssl/ops/video_ssl_preprocess_ops.py +1 -1
- official/projects/video_ssl/ops/video_ssl_preprocess_ops_test.py +1 -1
- official/projects/video_ssl/tasks/__init__.py +1 -1
- official/projects/video_ssl/tasks/linear_eval.py +1 -1
- official/projects/video_ssl/tasks/pretrain.py +1 -1
- official/projects/video_ssl/tasks/pretrain_test.py +1 -1
- official/projects/video_ssl/train.py +1 -1
- official/projects/volumetric_models/__init__.py +1 -1
- official/projects/volumetric_models/configs/__init__.py +1 -1
- official/projects/volumetric_models/configs/backbones.py +1 -1
- official/projects/volumetric_models/configs/decoders.py +1 -1
- official/projects/volumetric_models/configs/semantic_segmentation_3d.py +1 -1
- official/projects/volumetric_models/configs/semantic_segmentation_3d_test.py +1 -1
- official/projects/volumetric_models/dataloaders/__init__.py +1 -1
- official/projects/volumetric_models/dataloaders/segmentation_input_3d.py +1 -1
- official/projects/volumetric_models/dataloaders/segmentation_input_3d_test.py +1 -1
- official/projects/volumetric_models/evaluation/__init__.py +1 -1
- official/projects/volumetric_models/evaluation/segmentation_metrics.py +1 -1
- official/projects/volumetric_models/evaluation/segmentation_metrics_test.py +1 -1
- official/projects/volumetric_models/losses/__init__.py +1 -1
- official/projects/volumetric_models/losses/segmentation_losses.py +1 -1
- official/projects/volumetric_models/losses/segmentation_losses_test.py +1 -1
- official/projects/volumetric_models/modeling/__init__.py +1 -1
- official/projects/volumetric_models/modeling/backbones/__init__.py +1 -1
- official/projects/volumetric_models/modeling/backbones/unet_3d.py +1 -2
- official/projects/volumetric_models/modeling/backbones/unet_3d_test.py +1 -2
- official/projects/volumetric_models/modeling/decoders/__init__.py +1 -1
- official/projects/volumetric_models/modeling/decoders/factory.py +1 -3
- official/projects/volumetric_models/modeling/decoders/factory_test.py +1 -1
- official/projects/volumetric_models/modeling/decoders/unet_3d_decoder.py +1 -1
- official/projects/volumetric_models/modeling/decoders/unet_3d_decoder_test.py +1 -2
- official/projects/volumetric_models/modeling/factory.py +1 -3
- official/projects/volumetric_models/modeling/factory_test.py +1 -1
- official/projects/volumetric_models/modeling/heads/__init__.py +1 -1
- official/projects/volumetric_models/modeling/heads/segmentation_heads_3d.py +1 -1
- official/projects/volumetric_models/modeling/heads/segmentation_heads_3d_test.py +1 -1
- official/projects/volumetric_models/modeling/nn_blocks_3d.py +2 -3
- official/projects/volumetric_models/modeling/nn_blocks_3d_test.py +1 -2
- official/projects/volumetric_models/modeling/segmentation_model_test.py +1 -1
- official/projects/volumetric_models/registry_imports.py +1 -1
- official/projects/volumetric_models/serving/__init__.py +1 -1
- official/projects/volumetric_models/serving/export_saved_model.py +1 -1
- official/projects/volumetric_models/serving/semantic_segmentation_3d.py +1 -1
- official/projects/volumetric_models/serving/semantic_segmentation_3d_test.py +3 -3
- official/projects/volumetric_models/tasks/__init__.py +1 -1
- official/projects/volumetric_models/tasks/semantic_segmentation_3d.py +1 -1
- official/projects/volumetric_models/tasks/semantic_segmentation_3d_test.py +1 -1
- official/projects/volumetric_models/train.py +1 -1
- official/projects/volumetric_models/train_test.py +1 -1
- official/projects/waste_identification_ml/__init__.py +1 -1
- official/projects/waste_identification_ml/data_generation/__init__.py +1 -1
- official/projects/waste_identification_ml/data_generation/utils.py +1 -1
- official/projects/waste_identification_ml/data_generation/utils_test.py +1 -1
- official/projects/yolo/__init__.py +1 -1
- official/projects/yolo/common/__init__.py +1 -1
- official/projects/yolo/common/registry_imports.py +1 -1
- official/projects/yolo/configs/__init__.py +1 -1
- official/projects/yolo/configs/backbones.py +1 -1
- official/projects/yolo/configs/darknet_classification.py +1 -1
- official/projects/yolo/configs/decoders.py +1 -1
- official/projects/yolo/configs/yolo.py +1 -1
- official/projects/yolo/configs/yolov7.py +17 -1
- official/projects/yolo/dataloaders/__init__.py +1 -1
- official/projects/yolo/dataloaders/classification_input.py +1 -1
- official/projects/yolo/dataloaders/tf_example_decoder.py +1 -1
- official/projects/yolo/dataloaders/yolo_input.py +1 -1
- official/projects/yolo/losses/__init__.py +1 -1
- official/projects/yolo/losses/yolo_loss.py +1 -1
- official/projects/yolo/losses/yolo_loss_test.py +1 -1
- official/projects/yolo/losses/yolov7_loss.py +1 -1
- official/projects/yolo/losses/yolov7_loss_test.py +1 -1
- official/projects/yolo/modeling/__init__.py +1 -1
- official/projects/yolo/modeling/backbones/__init__.py +1 -1
- official/projects/yolo/modeling/backbones/darknet.py +1 -1
- official/projects/yolo/modeling/backbones/darknet_test.py +1 -1
- official/projects/yolo/modeling/backbones/yolov7.py +69 -1
- official/projects/yolo/modeling/backbones/yolov7_test.py +1 -1
- official/projects/yolo/modeling/decoders/__init__.py +1 -1
- official/projects/yolo/modeling/decoders/yolo_decoder.py +1 -1
- official/projects/yolo/modeling/decoders/yolo_decoder_test.py +1 -2
- official/projects/yolo/modeling/decoders/yolov7.py +90 -1
- official/projects/yolo/modeling/decoders/yolov7_test.py +1 -1
- official/projects/yolo/modeling/factory.py +1 -1
- official/projects/yolo/modeling/factory_test.py +1 -1
- official/projects/yolo/modeling/heads/__init__.py +1 -1
- official/projects/yolo/modeling/heads/yolo_head.py +1 -1
- official/projects/yolo/modeling/heads/yolo_head_test.py +1 -2
- official/projects/yolo/modeling/heads/yolov7_head.py +1 -1
- official/projects/yolo/modeling/heads/yolov7_head_test.py +1 -1
- official/projects/yolo/modeling/layers/__init__.py +1 -1
- official/projects/yolo/modeling/layers/detection_generator.py +1 -1
- official/projects/yolo/modeling/layers/detection_generator_test.py +1 -1
- official/projects/yolo/modeling/layers/nn_blocks.py +1 -1
- official/projects/yolo/modeling/layers/nn_blocks_test.py +1 -1
- official/projects/yolo/modeling/yolo_model.py +2 -2
- official/projects/yolo/modeling/yolov7_model.py +2 -2
- official/projects/yolo/ops/__init__.py +1 -1
- official/projects/yolo/ops/anchor.py +1 -1
- official/projects/yolo/ops/box_ops.py +1 -1
- official/projects/yolo/ops/box_ops_test.py +1 -1
- official/projects/yolo/ops/initializer_ops.py +1 -1
- official/projects/yolo/ops/kmeans_anchors.py +1 -1
- official/projects/yolo/ops/kmeans_anchors_test.py +1 -1
- official/projects/yolo/ops/loss_utils.py +1 -1
- official/projects/yolo/ops/math_ops.py +1 -1
- official/projects/yolo/ops/mosaic.py +1 -1
- official/projects/yolo/ops/preprocessing_ops.py +1 -1
- official/projects/yolo/ops/preprocessing_ops_test.py +1 -1
- official/projects/yolo/optimization/__init__.py +1 -1
- official/projects/yolo/optimization/configs/__init__.py +1 -1
- official/projects/yolo/optimization/configs/optimization_config.py +1 -1
- official/projects/yolo/optimization/configs/optimizer_config.py +1 -1
- official/projects/yolo/optimization/optimizer_factory.py +1 -1
- official/projects/yolo/optimization/sgd_torch.py +1 -1
- official/projects/yolo/serving/__init__.py +1 -1
- official/projects/yolo/serving/export_module_factory.py +1 -1
- official/projects/yolo/serving/export_saved_model.py +1 -1
- official/projects/yolo/serving/export_tflite.py +1 -1
- official/projects/yolo/serving/model_fn.py +1 -1
- official/projects/yolo/tasks/__init__.py +1 -1
- official/projects/yolo/tasks/image_classification.py +1 -1
- official/projects/yolo/tasks/task_utils.py +1 -1
- official/projects/yolo/tasks/yolo.py +1 -1
- official/projects/yolo/tasks/yolov7.py +1 -1
- official/projects/yolo/train.py +1 -1
- official/projects/yt8m/__init__.py +1 -1
- official/projects/yt8m/configs/__init__.py +1 -1
- official/projects/yt8m/configs/yt8m.py +1 -1
- official/projects/yt8m/configs/yt8m_test.py +1 -1
- official/projects/yt8m/modeling/__init__.py +1 -1
- official/projects/yt8m/modeling/backbones/__init__.py +1 -1
- official/projects/yt8m/modeling/backbones/dbof.py +1 -1
- official/projects/yt8m/modeling/backbones/dbof_test.py +1 -1
- official/projects/yt8m/modeling/heads/__init__.py +1 -1
- official/projects/yt8m/modeling/heads/logistic.py +1 -1
- official/projects/yt8m/modeling/heads/moe.py +1 -1
- official/projects/yt8m/modeling/nn_layers.py +1 -1
- official/projects/yt8m/modeling/nn_layers_test.py +1 -1
- official/projects/yt8m/modeling/yt8m_model.py +1 -1
- official/projects/yt8m/modeling/yt8m_model_test.py +1 -1
- official/projects/yt8m/modeling/yt8m_model_utils.py +1 -1
- official/projects/yt8m/modeling/yt8m_model_utils_test.py +1 -1
- official/projects/yt8m/tasks/__init__.py +1 -1
- official/projects/yt8m/tasks/yt8m_task.py +1 -1
- official/projects/yt8m/train.py +1 -1
- official/projects/yt8m/train_test.py +1 -1
- official/recommendation/__init__.py +1 -1
- official/recommendation/constants.py +1 -1
- official/recommendation/create_ncf_data.py +1 -2
- official/recommendation/data_pipeline.py +1 -1
- official/recommendation/data_preprocessing.py +1 -1
- official/recommendation/data_test.py +4 -4
- official/recommendation/movielens.py +1 -2
- official/recommendation/ncf_common.py +1 -1
- official/recommendation/ncf_input_pipeline.py +1 -1
- official/recommendation/ncf_keras_main.py +1 -1
- official/recommendation/ncf_test.py +1 -1
- official/recommendation/neumf_model.py +1 -1
- official/recommendation/popen_helper.py +1 -1
- official/recommendation/ranking/__init__.py +1 -1
- official/recommendation/ranking/common.py +1 -1
- official/recommendation/ranking/configs/__init__.py +1 -1
- official/recommendation/ranking/configs/config.py +14 -1
- official/recommendation/ranking/configs/config_test.py +1 -1
- official/recommendation/ranking/data/__init__.py +1 -1
- official/recommendation/ranking/data/data_pipeline.py +9 -2
- official/recommendation/ranking/data/data_pipeline_multi_hot.py +8 -2
- official/recommendation/ranking/data/data_pipeline_multi_hot_test.py +12 -6
- official/recommendation/ranking/data/data_pipeline_test.py +18 -8
- official/recommendation/ranking/task.py +102 -19
- official/recommendation/ranking/task_test.py +1 -1
- official/recommendation/ranking/train.py +1 -1
- official/recommendation/ranking/train_test.py +76 -31
- official/recommendation/stat_utils.py +1 -1
- official/recommendation/uplift/__init__.py +1 -1
- official/recommendation/uplift/keras_test_case.py +1 -1
- official/recommendation/uplift/keys.py +1 -1
- official/recommendation/uplift/layers/__init__.py +1 -1
- official/recommendation/uplift/layers/encoders/__init__.py +1 -1
- official/recommendation/uplift/layers/encoders/concat_features.py +1 -1
- official/recommendation/uplift/layers/encoders/concat_features_test.py +1 -1
- official/recommendation/uplift/layers/heads/__init__.py +1 -1
- official/recommendation/uplift/layers/heads/two_tower_logits_head.py +1 -1
- official/recommendation/uplift/layers/heads/two_tower_logits_head_test.py +1 -1
- official/recommendation/uplift/layers/uplift_networks/__init__.py +1 -1
- official/recommendation/uplift/layers/uplift_networks/base_uplift_networks.py +1 -1
- official/recommendation/uplift/layers/uplift_networks/two_tower_output_head.py +1 -1
- official/recommendation/uplift/layers/uplift_networks/two_tower_output_head_test.py +1 -1
- official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network.py +1 -1
- official/recommendation/uplift/layers/uplift_networks/two_tower_uplift_network_test.py +1 -1
- official/recommendation/uplift/losses/__init__.py +1 -1
- official/recommendation/uplift/losses/true_logits_loss.py +1 -1
- official/recommendation/uplift/losses/true_logits_loss_test.py +1 -1
- official/recommendation/uplift/metrics/__init__.py +1 -1
- official/recommendation/uplift/metrics/label_mean.py +1 -1
- official/recommendation/uplift/metrics/label_mean_test.py +1 -1
- official/recommendation/uplift/metrics/label_variance.py +1 -1
- official/recommendation/uplift/metrics/label_variance_test.py +1 -1
- official/recommendation/uplift/metrics/loss_metric.py +1 -1
- official/recommendation/uplift/metrics/loss_metric_test.py +1 -1
- official/recommendation/uplift/metrics/metric_configs.py +1 -1
- official/recommendation/uplift/metrics/poisson_metrics.py +1 -1
- official/recommendation/uplift/metrics/poisson_metrics_test.py +1 -1
- official/recommendation/uplift/metrics/sliced_metric.py +1 -1
- official/recommendation/uplift/metrics/sliced_metric_test.py +1 -1
- official/recommendation/uplift/metrics/treatment_fraction.py +1 -1
- official/recommendation/uplift/metrics/treatment_fraction_test.py +1 -1
- official/recommendation/uplift/metrics/treatment_sliced_metric.py +1 -1
- official/recommendation/uplift/metrics/treatment_sliced_metric_test.py +1 -1
- official/recommendation/uplift/metrics/uplift_mean.py +1 -1
- official/recommendation/uplift/metrics/uplift_mean_test.py +1 -1
- official/recommendation/uplift/metrics/variance.py +1 -1
- official/recommendation/uplift/metrics/variance_test.py +12 -10
- official/recommendation/uplift/models/__init__.py +1 -1
- official/recommendation/uplift/models/two_tower_uplift_model.py +1 -1
- official/recommendation/uplift/models/two_tower_uplift_model_test.py +1 -1
- official/recommendation/uplift/types.py +1 -1
- official/recommendation/uplift/utils.py +3 -3
- official/recommendation/uplift/utils_test.py +1 -1
- official/utils/__init__.py +1 -1
- official/utils/docs/__init__.py +1 -1
- official/utils/docs/build_orbit_api_docs.py +1 -1
- official/utils/docs/build_tfm_api_docs.py +1 -1
- official/utils/flags/__init__.py +1 -1
- official/utils/flags/_base.py +1 -1
- official/utils/flags/_benchmark.py +1 -1
- official/utils/flags/_conventions.py +1 -1
- official/utils/flags/_device.py +1 -1
- official/utils/flags/_distribution.py +1 -1
- official/utils/flags/_misc.py +1 -1
- official/utils/flags/_performance.py +1 -1
- official/utils/flags/core.py +1 -1
- official/utils/flags/flags_test.py +1 -1
- official/utils/hyperparams_flags.py +1 -1
- official/utils/misc/__init__.py +1 -1
- official/utils/misc/keras_utils.py +1 -1
- official/utils/misc/model_helpers.py +1 -1
- official/utils/misc/model_helpers_test.py +3 -3
- official/utils/testing/__init__.py +1 -1
- official/utils/testing/integration.py +1 -1
- official/utils/testing/mock_task.py +1 -1
- official/vision/__init__.py +1 -1
- official/vision/configs/__init__.py +1 -1
- official/vision/configs/backbones.py +3 -1
- official/vision/configs/backbones_3d.py +1 -2
- official/vision/configs/common.py +1 -3
- official/vision/configs/decoders.py +1 -3
- official/vision/configs/image_classification.py +1 -1
- official/vision/configs/image_classification_test.py +1 -1
- official/vision/configs/maskrcnn.py +1 -1
- official/vision/configs/maskrcnn_test.py +1 -1
- official/vision/configs/retinanet.py +2 -1
- official/vision/configs/retinanet_test.py +1 -1
- official/vision/configs/semantic_segmentation.py +7 -8
- official/vision/configs/semantic_segmentation_test.py +1 -1
- official/vision/configs/video_classification.py +1 -1
- official/vision/configs/video_classification_test.py +1 -1
- official/vision/data/__init__.py +1 -1
- official/vision/data/create_coco_tf_record.py +1 -1
- official/vision/data/fake_feature_generator.py +5 -2
- official/vision/data/image_utils.py +1 -1
- official/vision/data/image_utils_test.py +1 -1
- official/vision/data/process_coco_few_shot_json_files.py +1 -1
- official/vision/data/tf_example_builder.py +1 -1
- official/vision/data/tf_example_builder_test.py +1 -1
- official/vision/data/tf_example_feature_key.py +1 -1
- official/vision/data/tfrecord_lib.py +1 -1
- official/vision/data/tfrecord_lib_test.py +1 -1
- official/vision/dataloaders/__init__.py +1 -1
- official/vision/dataloaders/classification_input.py +1 -2
- official/vision/dataloaders/decoder.py +1 -1
- official/vision/dataloaders/input_reader.py +1 -1
- official/vision/dataloaders/input_reader_factory.py +1 -1
- official/vision/dataloaders/maskrcnn_input.py +1 -2
- official/vision/dataloaders/parser.py +1 -1
- official/vision/dataloaders/retinanet_input.py +1 -3
- official/vision/dataloaders/segmentation_input.py +9 -4
- official/vision/dataloaders/tf_example_decoder.py +1 -1
- official/vision/dataloaders/tf_example_decoder_test.py +1 -2
- official/vision/dataloaders/tf_example_label_map_decoder.py +1 -2
- official/vision/dataloaders/tf_example_label_map_decoder_test.py +1 -2
- official/vision/dataloaders/tfds_classification_decoders.py +1 -1
- official/vision/dataloaders/tfds_detection_decoders.py +1 -1
- official/vision/dataloaders/tfds_factory.py +1 -1
- official/vision/dataloaders/tfds_factory_test.py +1 -1
- official/vision/dataloaders/tfds_segmentation_decoders.py +1 -1
- official/vision/dataloaders/tfexample_utils.py +1 -1
- official/vision/dataloaders/utils.py +1 -2
- official/vision/dataloaders/utils_test.py +1 -3
- official/vision/dataloaders/video_input.py +1 -1
- official/vision/dataloaders/video_input_test.py +1 -2
- official/vision/evaluation/__init__.py +1 -1
- official/vision/evaluation/coco_evaluator.py +1 -2
- official/vision/evaluation/coco_utils.py +1 -3
- official/vision/evaluation/coco_utils_test.py +1 -1
- official/vision/evaluation/instance_metrics.py +1 -1
- official/vision/evaluation/instance_metrics_test.py +1 -1
- official/vision/evaluation/iou.py +1 -1
- official/vision/evaluation/iou_test.py +1 -1
- official/vision/evaluation/panoptic_quality.py +1 -1
- official/vision/evaluation/panoptic_quality_evaluator.py +1 -1
- official/vision/evaluation/panoptic_quality_evaluator_test.py +1 -1
- official/vision/evaluation/panoptic_quality_test.py +1 -1
- official/vision/evaluation/segmentation_metrics.py +1 -1
- official/vision/evaluation/segmentation_metrics_test.py +1 -1
- official/vision/evaluation/wod_detection_evaluator.py +1 -1
- official/vision/losses/__init__.py +1 -1
- official/vision/losses/focal_loss.py +1 -1
- official/vision/losses/loss_utils.py +1 -1
- official/vision/losses/maskrcnn_losses.py +1 -2
- official/vision/losses/maskrcnn_losses_test.py +1 -1
- official/vision/losses/retinanet_losses.py +1 -2
- official/vision/losses/segmentation_losses.py +1 -1
- official/vision/losses/segmentation_losses_test.py +1 -1
- official/vision/modeling/__init__.py +1 -1
- official/vision/modeling/backbones/__init__.py +1 -1
- official/vision/modeling/backbones/efficientnet.py +1 -3
- official/vision/modeling/backbones/efficientnet_test.py +1 -2
- official/vision/modeling/backbones/factory.py +1 -3
- official/vision/modeling/backbones/factory_test.py +1 -2
- official/vision/modeling/backbones/mobiledet.py +1 -1
- official/vision/modeling/backbones/mobiledet_test.py +1 -1
- official/vision/modeling/backbones/mobilenet.py +73 -3
- official/vision/modeling/backbones/mobilenet_test.py +12 -3
- official/vision/modeling/backbones/resnet.py +1 -2
- official/vision/modeling/backbones/resnet_3d.py +1 -2
- official/vision/modeling/backbones/resnet_3d_test.py +1 -2
- official/vision/modeling/backbones/resnet_deeplab.py +5 -4
- official/vision/modeling/backbones/resnet_deeplab_test.py +21 -10
- official/vision/modeling/backbones/resnet_test.py +1 -2
- official/vision/modeling/backbones/resnet_unet.py +1 -2
- official/vision/modeling/backbones/resnet_unet_test.py +1 -3
- official/vision/modeling/backbones/revnet.py +1 -2
- official/vision/modeling/backbones/revnet_test.py +1 -2
- official/vision/modeling/backbones/spinenet.py +1 -3
- official/vision/modeling/backbones/spinenet_mobile.py +1 -3
- official/vision/modeling/backbones/spinenet_mobile_test.py +1 -2
- official/vision/modeling/backbones/spinenet_test.py +1 -2
- official/vision/modeling/backbones/vit.py +53 -27
- official/vision/modeling/backbones/vit_specs.py +1 -1
- official/vision/modeling/backbones/vit_test.py +12 -1
- official/vision/modeling/classification_model.py +1 -2
- official/vision/modeling/classification_model_test.py +1 -2
- official/vision/modeling/decoders/__init__.py +1 -1
- official/vision/modeling/decoders/aspp.py +1 -3
- official/vision/modeling/decoders/aspp_test.py +1 -2
- official/vision/modeling/decoders/factory.py +1 -3
- official/vision/modeling/decoders/factory_test.py +1 -1
- official/vision/modeling/decoders/fpn.py +1 -2
- official/vision/modeling/decoders/fpn_test.py +1 -2
- official/vision/modeling/decoders/nasfpn.py +1 -3
- official/vision/modeling/decoders/nasfpn_test.py +1 -2
- official/vision/modeling/factory.py +1 -1
- official/vision/modeling/factory_3d.py +1 -2
- official/vision/modeling/factory_test.py +1 -2
- official/vision/modeling/heads/__init__.py +1 -1
- official/vision/modeling/heads/dense_prediction_heads.py +1 -3
- official/vision/modeling/heads/dense_prediction_heads_test.py +1 -3
- official/vision/modeling/heads/instance_heads.py +3 -4
- official/vision/modeling/heads/instance_heads_test.py +1 -2
- official/vision/modeling/heads/segmentation_heads.py +2 -2
- official/vision/modeling/heads/segmentation_heads_test.py +1 -2
- official/vision/modeling/layers/__init__.py +1 -1
- official/vision/modeling/layers/box_sampler.py +1 -2
- official/vision/modeling/layers/deeplab.py +1 -1
- official/vision/modeling/layers/deeplab_test.py +1 -1
- official/vision/modeling/layers/detection_generator.py +1 -3
- official/vision/modeling/layers/detection_generator_test.py +1 -3
- official/vision/modeling/layers/edgetpu.py +1 -1
- official/vision/modeling/layers/edgetpu_test.py +1 -1
- official/vision/modeling/layers/mask_sampler.py +1 -2
- official/vision/modeling/layers/nn_blocks.py +1 -2
- official/vision/modeling/layers/nn_blocks_3d.py +1 -2
- official/vision/modeling/layers/nn_blocks_3d_test.py +1 -2
- official/vision/modeling/layers/nn_blocks_test.py +1 -3
- official/vision/modeling/layers/nn_layers.py +1 -1
- official/vision/modeling/layers/nn_layers_test.py +1 -2
- official/vision/modeling/layers/roi_aligner.py +7 -5
- official/vision/modeling/layers/roi_aligner_test.py +1 -2
- official/vision/modeling/layers/roi_generator.py +1 -2
- official/vision/modeling/layers/roi_sampler.py +1 -2
- official/vision/modeling/maskrcnn_model.py +1 -1
- official/vision/modeling/maskrcnn_model_test.py +1 -2
- official/vision/modeling/models/__init__.py +1 -1
- official/vision/modeling/retinanet_model.py +9 -8
- official/vision/modeling/retinanet_model_test.py +1 -2
- official/vision/modeling/segmentation_model.py +4 -4
- official/vision/modeling/segmentation_model_test.py +1 -1
- official/vision/modeling/video_classification_model.py +1 -1
- official/vision/modeling/video_classification_model_test.py +1 -2
- official/vision/ops/__init__.py +1 -1
- official/vision/ops/anchor.py +1 -3
- official/vision/ops/anchor_generator.py +1 -1
- official/vision/ops/anchor_generator_test.py +1 -1
- official/vision/ops/anchor_test.py +1 -2
- official/vision/ops/augment.py +4 -16
- official/vision/ops/augment_test.py +1 -1
- official/vision/ops/box_matcher.py +1 -1
- official/vision/ops/box_matcher_test.py +1 -1
- official/vision/ops/box_ops.py +1 -2
- official/vision/ops/iou_similarity.py +1 -1
- official/vision/ops/iou_similarity_test.py +1 -1
- official/vision/ops/mask_ops.py +1 -3
- official/vision/ops/mask_ops_test.py +1 -2
- official/vision/ops/nms.py +1 -2
- official/vision/ops/preprocess_ops.py +40 -11
- official/vision/ops/preprocess_ops_3d.py +6 -3
- official/vision/ops/preprocess_ops_3d_test.py +1 -1
- official/vision/ops/preprocess_ops_test.py +13 -7
- official/vision/ops/sampling_ops.py +1 -2
- official/vision/ops/spatial_transform_ops.py +1 -1
- official/vision/ops/target_gather.py +1 -1
- official/vision/ops/target_gather_test.py +1 -1
- official/vision/registry_imports.py +1 -1
- official/vision/serving/__init__.py +1 -1
- official/vision/serving/detection.py +84 -33
- official/vision/serving/detection_test.py +39 -1
- official/vision/serving/export_base.py +1 -1
- official/vision/serving/export_base_v2.py +1 -1
- official/vision/serving/export_base_v2_test.py +1 -1
- official/vision/serving/export_module_factory.py +1 -1
- official/vision/serving/export_module_factory_test.py +1 -1
- official/vision/serving/export_saved_model.py +1 -1
- official/vision/serving/export_saved_model_lib.py +47 -30
- official/vision/serving/export_saved_model_lib_test.py +1 -1
- official/vision/serving/export_saved_model_lib_v2.py +1 -1
- official/vision/serving/export_tfhub.py +1 -2
- official/vision/serving/export_tfhub_lib.py +1 -3
- official/vision/serving/export_tflite.py +1 -1
- official/vision/serving/export_tflite_lib.py +1 -1
- official/vision/serving/export_utils.py +1 -1
- official/vision/serving/image_classification.py +1 -1
- official/vision/serving/image_classification_test.py +1 -1
- official/vision/serving/semantic_segmentation.py +6 -3
- official/vision/serving/semantic_segmentation_test.py +71 -7
- official/vision/serving/video_classification.py +1 -1
- official/vision/serving/video_classification_test.py +1 -1
- official/vision/tasks/__init__.py +1 -1
- official/vision/tasks/image_classification.py +1 -1
- official/vision/tasks/maskrcnn.py +1 -1
- official/vision/tasks/retinanet.py +1 -1
- official/vision/tasks/semantic_segmentation.py +1 -1
- official/vision/tasks/video_classification.py +1 -1
- official/vision/train.py +1 -1
- official/vision/train_spatial_partitioning.py +1 -1
- official/vision/utils/__init__.py +1 -1
- official/vision/utils/object_detection/__init__.py +1 -1
- official/vision/utils/object_detection/argmax_matcher.py +1 -1
- official/vision/utils/object_detection/balanced_positive_negative_sampler.py +1 -1
- official/vision/utils/object_detection/box_coder.py +1 -1
- official/vision/utils/object_detection/box_list.py +1 -1
- official/vision/utils/object_detection/box_list_ops.py +1 -1
- official/vision/utils/object_detection/faster_rcnn_box_coder.py +1 -1
- official/vision/utils/object_detection/matcher.py +1 -1
- official/vision/utils/object_detection/minibatch_sampler.py +1 -1
- official/vision/utils/object_detection/ops.py +1 -1
- official/vision/utils/object_detection/preprocessor.py +1 -1
- official/vision/utils/object_detection/region_similarity_calculator.py +1 -1
- official/vision/utils/object_detection/shape_utils.py +1 -1
- official/vision/utils/object_detection/target_assigner.py +1 -1
- official/vision/utils/object_detection/visualization_utils.py +6 -1
- official/vision/utils/ops_test.py +1 -1
- official/vision/utils/summary_manager.py +1 -1
- orbit/__init__.py +1 -1
- orbit/actions/__init__.py +1 -1
- orbit/actions/conditional_action.py +3 -2
- orbit/actions/conditional_action_test.py +1 -1
- orbit/actions/export_saved_model.py +1 -1
- orbit/actions/export_saved_model_test.py +1 -1
- orbit/actions/new_best_metric.py +2 -2
- orbit/actions/new_best_metric_test.py +2 -2
- orbit/actions/save_checkpoint_if_preempted.py +1 -1
- orbit/controller.py +1 -1
- orbit/controller_test.py +1 -1
- orbit/examples/__init__.py +1 -1
- orbit/examples/single_task/__init__.py +1 -1
- orbit/examples/single_task/single_task_evaluator.py +1 -1
- orbit/examples/single_task/single_task_evaluator_test.py +1 -1
- orbit/examples/single_task/single_task_trainer.py +1 -1
- orbit/examples/single_task/single_task_trainer_test.py +1 -1
- orbit/runner.py +1 -1
- orbit/standard_runner.py +1 -1
- orbit/standard_runner_test.py +1 -1
- orbit/utils/__init__.py +1 -1
- orbit/utils/common.py +1 -1
- orbit/utils/common_test.py +1 -1
- orbit/utils/epoch_helper.py +1 -1
- orbit/utils/loop_fns.py +7 -2
- orbit/utils/summary_manager.py +1 -1
- orbit/utils/summary_manager_interface.py +1 -1
- orbit/utils/tpu_summaries.py +1 -1
- orbit/utils/tpu_summaries_test.py +1 -1
- tensorflow_models/__init__.py +1 -1
- tensorflow_models/nlp/__init__.py +1 -1
- tensorflow_models/tensorflow_models_test.py +1 -1
- tensorflow_models/uplift/__init__.py +1 -1
- tensorflow_models/vision/__init__.py +1 -1
- {tf_models_nightly-2.17.0.dev20240528.dist-info → tf_models_nightly-2.20.0.dev20251205.dist-info}/METADATA +1 -1
- tf_models_nightly-2.20.0.dev20251205.dist-info/RECORD +1256 -0
- tf_models_nightly-2.17.0.dev20240528.dist-info/RECORD +0 -1216
- {tf_models_nightly-2.17.0.dev20240528.dist-info → tf_models_nightly-2.20.0.dev20251205.dist-info}/AUTHORS +0 -0
- {tf_models_nightly-2.17.0.dev20240528.dist-info → tf_models_nightly-2.20.0.dev20251205.dist-info}/LICENSE +0 -0
- {tf_models_nightly-2.17.0.dev20240528.dist-info → tf_models_nightly-2.20.0.dev20251205.dist-info}/WHEEL +0 -0
- {tf_models_nightly-2.17.0.dev20240528.dist-info → tf_models_nightly-2.20.0.dev20251205.dist-info}/top_level.txt +0 -0
|
@@ -0,0 +1,433 @@
|
|
|
1
|
+
# Copyright 2025 The TensorFlow Authors. All Rights Reserved.
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
"""DETR detection task definition."""
|
|
16
|
+
from typing import Optional
|
|
17
|
+
|
|
18
|
+
from absl import logging
|
|
19
|
+
import tensorflow as tf, tf_keras
|
|
20
|
+
|
|
21
|
+
from official.common import dataset_fn
|
|
22
|
+
from official.core import base_task
|
|
23
|
+
from official.core import task_factory
|
|
24
|
+
from official.projects.detr.configs import detr as detr_cfg
|
|
25
|
+
from official.projects.detr.dataloaders import coco
|
|
26
|
+
from official.projects.detr.dataloaders import detr_input
|
|
27
|
+
from official.projects.detr.modeling import detr
|
|
28
|
+
from official.projects.detr.ops import matchers
|
|
29
|
+
from official.vision.dataloaders import input_reader_factory
|
|
30
|
+
from official.vision.dataloaders import tf_example_decoder
|
|
31
|
+
from official.vision.dataloaders import tfds_factory
|
|
32
|
+
from official.vision.dataloaders import tf_example_label_map_decoder
|
|
33
|
+
from official.vision.evaluation import coco_evaluator
|
|
34
|
+
from official.vision.modeling import backbones
|
|
35
|
+
from official.vision.ops import box_ops
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@task_factory.register_task_cls(detr_cfg.DetrTask)
|
|
39
|
+
class DetectionTask(base_task.Task):
|
|
40
|
+
"""A single-replica view of training procedure.
|
|
41
|
+
|
|
42
|
+
DETR task provides artifacts for training/evalution procedures, including
|
|
43
|
+
loading/iterating over Datasets, initializing the model, calculating the loss,
|
|
44
|
+
post-processing, and customized metrics with reduction.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
def build_model(self):
|
|
48
|
+
"""Build DETR model."""
|
|
49
|
+
|
|
50
|
+
input_specs = tf_keras.layers.InputSpec(shape=[None] +
|
|
51
|
+
self._task_config.model.input_size)
|
|
52
|
+
|
|
53
|
+
backbone = backbones.factory.build_backbone(
|
|
54
|
+
input_specs=input_specs,
|
|
55
|
+
backbone_config=self._task_config.model.backbone,
|
|
56
|
+
norm_activation_config=self._task_config.model.norm_activation)
|
|
57
|
+
|
|
58
|
+
model = detr.DETR(backbone,
|
|
59
|
+
self._task_config.model.backbone_endpoint_name,
|
|
60
|
+
self._task_config.model.num_queries,
|
|
61
|
+
self._task_config.model.hidden_size,
|
|
62
|
+
self._task_config.model.num_classes,
|
|
63
|
+
self._task_config.model.num_encoder_layers,
|
|
64
|
+
self._task_config.model.num_decoder_layers)
|
|
65
|
+
return model
|
|
66
|
+
|
|
67
|
+
def initialize(self, model: tf_keras.Model):
|
|
68
|
+
"""Loading pretrained checkpoint."""
|
|
69
|
+
if not self._task_config.init_checkpoint:
|
|
70
|
+
return
|
|
71
|
+
|
|
72
|
+
ckpt_dir_or_file = self._task_config.init_checkpoint
|
|
73
|
+
|
|
74
|
+
# Restoring checkpoint.
|
|
75
|
+
if tf.io.gfile.isdir(ckpt_dir_or_file):
|
|
76
|
+
ckpt_dir_or_file = tf.train.latest_checkpoint(ckpt_dir_or_file)
|
|
77
|
+
|
|
78
|
+
if self._task_config.init_checkpoint_modules == 'all':
|
|
79
|
+
ckpt = tf.train.Checkpoint(**model.checkpoint_items)
|
|
80
|
+
status = ckpt.restore(ckpt_dir_or_file)
|
|
81
|
+
status.assert_consumed()
|
|
82
|
+
elif self._task_config.init_checkpoint_modules == 'backbone':
|
|
83
|
+
ckpt = tf.train.Checkpoint(backbone=model.backbone)
|
|
84
|
+
status = ckpt.restore(ckpt_dir_or_file)
|
|
85
|
+
status.expect_partial().assert_existing_objects_matched()
|
|
86
|
+
|
|
87
|
+
logging.info('Finished loading pretrained checkpoint from %s',
|
|
88
|
+
ckpt_dir_or_file)
|
|
89
|
+
|
|
90
|
+
def build_inputs(self,
|
|
91
|
+
params,
|
|
92
|
+
input_context: Optional[tf.distribute.InputContext] = None):
|
|
93
|
+
"""Build input dataset."""
|
|
94
|
+
if isinstance(params, coco.COCODataConfig):
|
|
95
|
+
dataset = coco.COCODataLoader(params).load(input_context)
|
|
96
|
+
else:
|
|
97
|
+
if params.tfds_name:
|
|
98
|
+
decoder = tfds_factory.get_detection_decoder(params.tfds_name)
|
|
99
|
+
else:
|
|
100
|
+
decoder_cfg = params.decoder.get()
|
|
101
|
+
if params.decoder.type == 'simple_decoder':
|
|
102
|
+
decoder = tf_example_decoder.TfExampleDecoder(
|
|
103
|
+
regenerate_source_id=decoder_cfg.regenerate_source_id)
|
|
104
|
+
elif params.decoder.type == 'label_map_decoder':
|
|
105
|
+
decoder = tf_example_label_map_decoder.TfExampleDecoderLabelMap(
|
|
106
|
+
label_map=decoder_cfg.label_map,
|
|
107
|
+
regenerate_source_id=decoder_cfg.regenerate_source_id)
|
|
108
|
+
else:
|
|
109
|
+
raise ValueError('Unknown decoder type: {}!'.format(
|
|
110
|
+
params.decoder.type))
|
|
111
|
+
|
|
112
|
+
parser = detr_input.Parser(
|
|
113
|
+
class_offset=self._task_config.losses.class_offset,
|
|
114
|
+
output_size=self._task_config.model.input_size[:2],
|
|
115
|
+
)
|
|
116
|
+
|
|
117
|
+
reader = input_reader_factory.input_reader_generator(
|
|
118
|
+
params,
|
|
119
|
+
dataset_fn=dataset_fn.pick_dataset_fn(params.file_type),
|
|
120
|
+
decoder_fn=decoder.decode,
|
|
121
|
+
parser_fn=parser.parse_fn(params.is_training))
|
|
122
|
+
dataset = reader.read(input_context=input_context)
|
|
123
|
+
|
|
124
|
+
return dataset
|
|
125
|
+
|
|
126
|
+
def _compute_cost(self, cls_outputs, box_outputs, cls_targets, box_targets):
|
|
127
|
+
# Approximate classification cost with 1 - prob[target class].
|
|
128
|
+
# The 1 is a constant that doesn't change the matching, it can be ommitted.
|
|
129
|
+
# background: 0
|
|
130
|
+
cls_cost = self._task_config.losses.lambda_cls * tf.gather(
|
|
131
|
+
-tf.nn.softmax(cls_outputs), cls_targets, batch_dims=1, axis=-1
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
# Compute the L1 cost between boxes,
|
|
135
|
+
paired_differences = self._task_config.losses.lambda_box * tf.abs(
|
|
136
|
+
tf.expand_dims(box_outputs, 2) - tf.expand_dims(box_targets, 1)
|
|
137
|
+
)
|
|
138
|
+
box_cost = tf.reduce_sum(paired_differences, axis=-1)
|
|
139
|
+
|
|
140
|
+
# Compute the giou cost betwen boxes
|
|
141
|
+
giou_cost = (
|
|
142
|
+
self._task_config.losses.lambda_giou
|
|
143
|
+
* -box_ops.bbox_generalized_overlap(
|
|
144
|
+
box_ops.cycxhw_to_yxyx(box_outputs),
|
|
145
|
+
box_ops.cycxhw_to_yxyx(box_targets),
|
|
146
|
+
)
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
total_cost = cls_cost + box_cost + giou_cost
|
|
150
|
+
|
|
151
|
+
max_cost = (
|
|
152
|
+
self._task_config.losses.lambda_cls * 1.0
|
|
153
|
+
+ self._task_config.losses.lambda_box * 4.0
|
|
154
|
+
+ self._task_config.losses.lambda_giou * 1.0
|
|
155
|
+
)
|
|
156
|
+
|
|
157
|
+
# Set pads to large constant
|
|
158
|
+
valid = tf.expand_dims(
|
|
159
|
+
tf.cast(tf.not_equal(cls_targets, 0), dtype=total_cost.dtype), axis=1
|
|
160
|
+
)
|
|
161
|
+
total_cost = (1 - valid) * max_cost + valid * total_cost
|
|
162
|
+
|
|
163
|
+
# Set inf of nan to large constant
|
|
164
|
+
total_cost = tf.where(
|
|
165
|
+
tf.logical_or(tf.math.is_nan(total_cost), tf.math.is_inf(total_cost)),
|
|
166
|
+
max_cost * tf.ones_like(total_cost, dtype=total_cost.dtype),
|
|
167
|
+
total_cost,
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
return total_cost
|
|
171
|
+
|
|
172
|
+
def build_losses(self, outputs, labels, aux_losses=None):
|
|
173
|
+
"""Builds DETR losses."""
|
|
174
|
+
cls_outputs = outputs['cls_outputs']
|
|
175
|
+
box_outputs = outputs['box_outputs']
|
|
176
|
+
cls_targets = labels['classes']
|
|
177
|
+
box_targets = labels['boxes']
|
|
178
|
+
|
|
179
|
+
cost = self._compute_cost(
|
|
180
|
+
cls_outputs, box_outputs, cls_targets, box_targets)
|
|
181
|
+
|
|
182
|
+
_, indices = matchers.hungarian_matching(cost)
|
|
183
|
+
indices = tf.stop_gradient(indices)
|
|
184
|
+
|
|
185
|
+
target_index = tf.math.argmax(indices, axis=1)
|
|
186
|
+
cls_assigned = tf.gather(cls_outputs, target_index, batch_dims=1, axis=1)
|
|
187
|
+
box_assigned = tf.gather(box_outputs, target_index, batch_dims=1, axis=1)
|
|
188
|
+
|
|
189
|
+
background = tf.equal(cls_targets, 0)
|
|
190
|
+
num_boxes = tf.reduce_sum(
|
|
191
|
+
tf.cast(tf.logical_not(background), tf.float32), axis=-1)
|
|
192
|
+
|
|
193
|
+
# Down-weight background to account for class imbalance.
|
|
194
|
+
xentropy = tf.nn.sparse_softmax_cross_entropy_with_logits(
|
|
195
|
+
labels=cls_targets, logits=cls_assigned)
|
|
196
|
+
cls_loss = self._task_config.losses.lambda_cls * tf.where(
|
|
197
|
+
background, self._task_config.losses.background_cls_weight * xentropy,
|
|
198
|
+
xentropy)
|
|
199
|
+
cls_weights = tf.where(
|
|
200
|
+
background,
|
|
201
|
+
self._task_config.losses.background_cls_weight * tf.ones_like(cls_loss),
|
|
202
|
+
tf.ones_like(cls_loss))
|
|
203
|
+
|
|
204
|
+
# Box loss is only calculated on non-background class.
|
|
205
|
+
l_1 = tf.reduce_sum(tf.abs(box_assigned - box_targets), axis=-1)
|
|
206
|
+
box_loss = self._task_config.losses.lambda_box * tf.where(
|
|
207
|
+
background, tf.zeros_like(l_1), l_1)
|
|
208
|
+
|
|
209
|
+
# Giou loss is only calculated on non-background class.
|
|
210
|
+
giou = tf.linalg.diag_part(1.0 - box_ops.bbox_generalized_overlap(
|
|
211
|
+
box_ops.cycxhw_to_yxyx(box_assigned),
|
|
212
|
+
box_ops.cycxhw_to_yxyx(box_targets)
|
|
213
|
+
))
|
|
214
|
+
giou_loss = self._task_config.losses.lambda_giou * tf.where(
|
|
215
|
+
background, tf.zeros_like(giou), giou)
|
|
216
|
+
|
|
217
|
+
# Consider doing all reduce once in train_step to speed up.
|
|
218
|
+
num_boxes_per_replica = tf.reduce_sum(num_boxes)
|
|
219
|
+
cls_weights_per_replica = tf.reduce_sum(cls_weights)
|
|
220
|
+
replica_context = tf.distribute.get_replica_context()
|
|
221
|
+
num_boxes_sum, cls_weights_sum = replica_context.all_reduce(
|
|
222
|
+
tf.distribute.ReduceOp.SUM,
|
|
223
|
+
[num_boxes_per_replica, cls_weights_per_replica])
|
|
224
|
+
cls_loss = tf.math.divide_no_nan(
|
|
225
|
+
tf.reduce_sum(cls_loss), cls_weights_sum)
|
|
226
|
+
box_loss = tf.math.divide_no_nan(
|
|
227
|
+
tf.reduce_sum(box_loss), num_boxes_sum)
|
|
228
|
+
giou_loss = tf.math.divide_no_nan(
|
|
229
|
+
tf.reduce_sum(giou_loss), num_boxes_sum)
|
|
230
|
+
|
|
231
|
+
aux_losses = tf.add_n(aux_losses) if aux_losses else 0.0
|
|
232
|
+
|
|
233
|
+
total_loss = cls_loss + box_loss + giou_loss + aux_losses
|
|
234
|
+
return total_loss, cls_loss, box_loss, giou_loss
|
|
235
|
+
|
|
236
|
+
def build_metrics(self, training=True):
|
|
237
|
+
"""Builds detection metrics."""
|
|
238
|
+
metrics = []
|
|
239
|
+
metric_names = ['cls_loss', 'box_loss', 'giou_loss']
|
|
240
|
+
for name in metric_names:
|
|
241
|
+
metrics.append(tf_keras.metrics.Mean(name, dtype=tf.float32))
|
|
242
|
+
|
|
243
|
+
if not training:
|
|
244
|
+
self.coco_metric = coco_evaluator.COCOEvaluator(
|
|
245
|
+
annotation_file=self._task_config.annotation_file,
|
|
246
|
+
include_mask=False,
|
|
247
|
+
need_rescale_bboxes=True,
|
|
248
|
+
per_category_metrics=self._task_config.per_category_metrics)
|
|
249
|
+
return metrics
|
|
250
|
+
|
|
251
|
+
def train_step(self, inputs, model, optimizer, metrics=None):
|
|
252
|
+
"""Does forward and backward.
|
|
253
|
+
|
|
254
|
+
Args:
|
|
255
|
+
inputs: a dictionary of input tensors.
|
|
256
|
+
model: the model, forward pass definition.
|
|
257
|
+
optimizer: the optimizer for this training step.
|
|
258
|
+
metrics: a nested structure of metrics objects.
|
|
259
|
+
|
|
260
|
+
Returns:
|
|
261
|
+
A dictionary of logs.
|
|
262
|
+
"""
|
|
263
|
+
features, labels = inputs
|
|
264
|
+
with tf.GradientTape() as tape:
|
|
265
|
+
outputs = model(features, training=True)
|
|
266
|
+
|
|
267
|
+
loss = 0.0
|
|
268
|
+
cls_loss = 0.0
|
|
269
|
+
box_loss = 0.0
|
|
270
|
+
giou_loss = 0.0
|
|
271
|
+
|
|
272
|
+
for output in outputs:
|
|
273
|
+
# Computes per-replica loss.
|
|
274
|
+
layer_loss, layer_cls_loss, layer_box_loss, layer_giou_loss = (
|
|
275
|
+
self.build_losses(
|
|
276
|
+
outputs=output, labels=labels, aux_losses=model.losses
|
|
277
|
+
)
|
|
278
|
+
)
|
|
279
|
+
loss += layer_loss
|
|
280
|
+
cls_loss += layer_cls_loss
|
|
281
|
+
box_loss += layer_box_loss
|
|
282
|
+
giou_loss += layer_giou_loss
|
|
283
|
+
|
|
284
|
+
# Consider moving scaling logic from build_losses to here.
|
|
285
|
+
scaled_loss = loss
|
|
286
|
+
# For mixed_precision policy, when LossScaleOptimizer is used, loss is
|
|
287
|
+
# scaled for numerical stability.
|
|
288
|
+
if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
|
|
289
|
+
scaled_loss = optimizer.get_scaled_loss(scaled_loss)
|
|
290
|
+
|
|
291
|
+
tvars = model.trainable_variables
|
|
292
|
+
grads = tape.gradient(scaled_loss, tvars)
|
|
293
|
+
# Scales back gradient when LossScaleOptimizer is used.
|
|
294
|
+
if isinstance(optimizer, tf_keras.mixed_precision.LossScaleOptimizer):
|
|
295
|
+
grads = optimizer.get_unscaled_gradients(grads)
|
|
296
|
+
optimizer.apply_gradients(list(zip(grads, tvars)))
|
|
297
|
+
|
|
298
|
+
# Multiply for logging.
|
|
299
|
+
# Since we expect the gradient replica sum to happen in the optimizer,
|
|
300
|
+
# the loss is scaled with global num_boxes and weights.
|
|
301
|
+
# To have it more interpretable/comparable we scale it back when logging.
|
|
302
|
+
num_replicas_in_sync = tf.distribute.get_strategy().num_replicas_in_sync
|
|
303
|
+
loss *= num_replicas_in_sync
|
|
304
|
+
cls_loss *= num_replicas_in_sync
|
|
305
|
+
box_loss *= num_replicas_in_sync
|
|
306
|
+
giou_loss *= num_replicas_in_sync
|
|
307
|
+
|
|
308
|
+
# Trainer class handles loss metric for you.
|
|
309
|
+
logs = {self.loss: loss}
|
|
310
|
+
|
|
311
|
+
all_losses = {
|
|
312
|
+
'cls_loss': cls_loss,
|
|
313
|
+
'box_loss': box_loss,
|
|
314
|
+
'giou_loss': giou_loss,
|
|
315
|
+
}
|
|
316
|
+
|
|
317
|
+
# Metric results will be added to logs for you.
|
|
318
|
+
if metrics:
|
|
319
|
+
for m in metrics:
|
|
320
|
+
m.update_state(all_losses[m.name])
|
|
321
|
+
return logs
|
|
322
|
+
|
|
323
|
+
def validation_step(self, inputs, model, metrics=None):
|
|
324
|
+
"""Validatation step.
|
|
325
|
+
|
|
326
|
+
Args:
|
|
327
|
+
inputs: a dictionary of input tensors.
|
|
328
|
+
model: the keras.Model.
|
|
329
|
+
metrics: a nested structure of metrics objects.
|
|
330
|
+
|
|
331
|
+
Returns:
|
|
332
|
+
A dictionary of logs.
|
|
333
|
+
"""
|
|
334
|
+
features, labels = inputs
|
|
335
|
+
|
|
336
|
+
outputs = model(features, training=False)[-1]
|
|
337
|
+
loss, cls_loss, box_loss, giou_loss = self.build_losses(
|
|
338
|
+
outputs=outputs, labels=labels, aux_losses=model.losses)
|
|
339
|
+
|
|
340
|
+
# Multiply for logging.
|
|
341
|
+
# Since we expect the gradient replica sum to happen in the optimizer,
|
|
342
|
+
# the loss is scaled with global num_boxes and weights.
|
|
343
|
+
# To have it more interpretable/comparable we scale it back when logging.
|
|
344
|
+
num_replicas_in_sync = tf.distribute.get_strategy().num_replicas_in_sync
|
|
345
|
+
loss *= num_replicas_in_sync
|
|
346
|
+
cls_loss *= num_replicas_in_sync
|
|
347
|
+
box_loss *= num_replicas_in_sync
|
|
348
|
+
giou_loss *= num_replicas_in_sync
|
|
349
|
+
|
|
350
|
+
# Evaluator class handles loss metric for you.
|
|
351
|
+
logs = {self.loss: loss}
|
|
352
|
+
|
|
353
|
+
# This is for backward compatibility.
|
|
354
|
+
if 'detection_boxes' not in outputs:
|
|
355
|
+
detection_boxes = box_ops.cycxhw_to_yxyx(
|
|
356
|
+
outputs['box_outputs']) * tf.expand_dims(
|
|
357
|
+
tf.concat([
|
|
358
|
+
labels['image_info'][:, 1:2, 0], labels['image_info'][:, 1:2,
|
|
359
|
+
1],
|
|
360
|
+
labels['image_info'][:, 1:2, 0], labels['image_info'][:, 1:2,
|
|
361
|
+
1]
|
|
362
|
+
],
|
|
363
|
+
axis=1),
|
|
364
|
+
axis=1)
|
|
365
|
+
else:
|
|
366
|
+
detection_boxes = outputs['detection_boxes']
|
|
367
|
+
|
|
368
|
+
detection_scores = tf.math.reduce_max(
|
|
369
|
+
tf.nn.softmax(outputs['cls_outputs'])[:, :, 1:], axis=-1
|
|
370
|
+
) if 'detection_scores' not in outputs else outputs['detection_scores']
|
|
371
|
+
|
|
372
|
+
if 'detection_classes' not in outputs:
|
|
373
|
+
detection_classes = tf.math.argmax(
|
|
374
|
+
outputs['cls_outputs'][:, :, 1:], axis=-1) + 1
|
|
375
|
+
else:
|
|
376
|
+
detection_classes = outputs['detection_classes']
|
|
377
|
+
|
|
378
|
+
if 'num_detections' not in outputs:
|
|
379
|
+
num_detections = tf.reduce_sum(
|
|
380
|
+
tf.cast(
|
|
381
|
+
tf.math.greater(
|
|
382
|
+
tf.math.reduce_max(outputs['cls_outputs'], axis=-1), 0),
|
|
383
|
+
tf.int32),
|
|
384
|
+
axis=-1)
|
|
385
|
+
else:
|
|
386
|
+
num_detections = outputs['num_detections']
|
|
387
|
+
|
|
388
|
+
predictions = {
|
|
389
|
+
'detection_boxes': detection_boxes,
|
|
390
|
+
'detection_scores': detection_scores,
|
|
391
|
+
'detection_classes': detection_classes,
|
|
392
|
+
'num_detections': num_detections,
|
|
393
|
+
'source_id': labels['id'],
|
|
394
|
+
'image_info': labels['image_info']
|
|
395
|
+
}
|
|
396
|
+
|
|
397
|
+
ground_truths = {
|
|
398
|
+
'source_id': labels['id'],
|
|
399
|
+
'height': labels['image_info'][:, 0:1, 0],
|
|
400
|
+
'width': labels['image_info'][:, 0:1, 1],
|
|
401
|
+
'num_detections': tf.reduce_sum(
|
|
402
|
+
tf.cast(tf.math.greater(labels['classes'], 0), tf.int32), axis=-1),
|
|
403
|
+
'boxes': labels['gt_boxes'],
|
|
404
|
+
'classes': labels['classes'],
|
|
405
|
+
'is_crowds': labels['is_crowd']
|
|
406
|
+
}
|
|
407
|
+
logs.update({'predictions': predictions,
|
|
408
|
+
'ground_truths': ground_truths})
|
|
409
|
+
|
|
410
|
+
all_losses = {
|
|
411
|
+
'cls_loss': cls_loss,
|
|
412
|
+
'box_loss': box_loss,
|
|
413
|
+
'giou_loss': giou_loss,
|
|
414
|
+
}
|
|
415
|
+
|
|
416
|
+
# Metric results will be added to logs for you.
|
|
417
|
+
if metrics:
|
|
418
|
+
for m in metrics:
|
|
419
|
+
m.update_state(all_losses[m.name])
|
|
420
|
+
return logs
|
|
421
|
+
|
|
422
|
+
def aggregate_logs(self, state=None, step_outputs=None):
|
|
423
|
+
if state is None:
|
|
424
|
+
self.coco_metric.reset_states()
|
|
425
|
+
state = self.coco_metric
|
|
426
|
+
|
|
427
|
+
state.update_state(
|
|
428
|
+
step_outputs['ground_truths'],
|
|
429
|
+
step_outputs['predictions'])
|
|
430
|
+
return state
|
|
431
|
+
|
|
432
|
+
def reduce_aggregated_logs(self, aggregated_logs, global_step=None):
|
|
433
|
+
return aggregated_logs.result()
|
|
@@ -0,0 +1,203 @@
|
|
|
1
|
+
# Copyright 2025 The TensorFlow Authors. All Rights Reserved.
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
"""Tests for detection."""
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
import tensorflow as tf, tf_keras
|
|
19
|
+
import tensorflow_datasets as tfds
|
|
20
|
+
|
|
21
|
+
from official.projects.detr import optimization
|
|
22
|
+
from official.projects.detr.configs import detr as detr_cfg
|
|
23
|
+
from official.projects.detr.dataloaders import coco
|
|
24
|
+
from official.projects.detr.tasks import detection
|
|
25
|
+
from official.vision.configs import backbones
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
_NUM_EXAMPLES = 10
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _gen_fn():
|
|
32
|
+
h = np.random.randint(0, 300)
|
|
33
|
+
w = np.random.randint(0, 300)
|
|
34
|
+
num_boxes = np.random.randint(0, 50)
|
|
35
|
+
return {
|
|
36
|
+
'image': np.ones(shape=(h, w, 3), dtype=np.uint8),
|
|
37
|
+
'image/id': np.random.randint(0, 100),
|
|
38
|
+
'image/filename': 'test',
|
|
39
|
+
'objects': {
|
|
40
|
+
'is_crowd': np.ones(shape=(num_boxes), dtype=bool),
|
|
41
|
+
'bbox': np.ones(shape=(num_boxes, 4), dtype=np.float32),
|
|
42
|
+
'label': np.ones(shape=(num_boxes), dtype=np.int64),
|
|
43
|
+
'id': np.ones(shape=(num_boxes), dtype=np.int64),
|
|
44
|
+
'area': np.ones(shape=(num_boxes), dtype=np.int64),
|
|
45
|
+
}
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _as_dataset(self, *args, **kwargs):
|
|
50
|
+
del args
|
|
51
|
+
del kwargs
|
|
52
|
+
return tf.data.Dataset.from_generator(
|
|
53
|
+
lambda: (_gen_fn() for i in range(_NUM_EXAMPLES)),
|
|
54
|
+
output_types=self.info.features.dtype,
|
|
55
|
+
output_shapes=self.info.features.shape,
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
class DetectionTest(tf.test.TestCase):
|
|
60
|
+
|
|
61
|
+
def test_train_step(self):
|
|
62
|
+
config = detr_cfg.DetrTask(
|
|
63
|
+
model=detr_cfg.Detr(
|
|
64
|
+
input_size=[1333, 1333, 3],
|
|
65
|
+
num_encoder_layers=1,
|
|
66
|
+
num_decoder_layers=1,
|
|
67
|
+
num_classes=81,
|
|
68
|
+
backbone=backbones.Backbone(
|
|
69
|
+
type='resnet',
|
|
70
|
+
resnet=backbones.ResNet(model_id=10, bn_trainable=False))
|
|
71
|
+
),
|
|
72
|
+
train_data=coco.COCODataConfig(
|
|
73
|
+
tfds_name='coco/2017',
|
|
74
|
+
tfds_split='validation',
|
|
75
|
+
is_training=True,
|
|
76
|
+
global_batch_size=2,
|
|
77
|
+
))
|
|
78
|
+
with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
|
|
79
|
+
task = detection.DetectionTask(config)
|
|
80
|
+
model = task.build_model()
|
|
81
|
+
dataset = task.build_inputs(config.train_data)
|
|
82
|
+
iterator = iter(dataset)
|
|
83
|
+
opt_cfg = optimization.OptimizationConfig({
|
|
84
|
+
'optimizer': {
|
|
85
|
+
'type': 'detr_adamw',
|
|
86
|
+
'detr_adamw': {
|
|
87
|
+
'weight_decay_rate': 1e-4,
|
|
88
|
+
'global_clipnorm': 0.1,
|
|
89
|
+
}
|
|
90
|
+
},
|
|
91
|
+
'learning_rate': {
|
|
92
|
+
'type': 'stepwise',
|
|
93
|
+
'stepwise': {
|
|
94
|
+
'boundaries': [120000],
|
|
95
|
+
'values': [0.0001, 1.0e-05]
|
|
96
|
+
}
|
|
97
|
+
},
|
|
98
|
+
})
|
|
99
|
+
optimizer = detection.DetectionTask.create_optimizer(opt_cfg)
|
|
100
|
+
task.train_step(next(iterator), model, optimizer)
|
|
101
|
+
|
|
102
|
+
def test_validation_step(self):
|
|
103
|
+
config = detr_cfg.DetrTask(
|
|
104
|
+
model=detr_cfg.Detr(
|
|
105
|
+
input_size=[1333, 1333, 3],
|
|
106
|
+
num_encoder_layers=1,
|
|
107
|
+
num_decoder_layers=1,
|
|
108
|
+
num_classes=81,
|
|
109
|
+
backbone=backbones.Backbone(
|
|
110
|
+
type='resnet',
|
|
111
|
+
resnet=backbones.ResNet(model_id=10, bn_trainable=False))
|
|
112
|
+
),
|
|
113
|
+
validation_data=coco.COCODataConfig(
|
|
114
|
+
tfds_name='coco/2017',
|
|
115
|
+
tfds_split='validation',
|
|
116
|
+
is_training=False,
|
|
117
|
+
global_batch_size=2,
|
|
118
|
+
))
|
|
119
|
+
|
|
120
|
+
with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
|
|
121
|
+
task = detection.DetectionTask(config)
|
|
122
|
+
model = task.build_model()
|
|
123
|
+
metrics = task.build_metrics(training=False)
|
|
124
|
+
dataset = task.build_inputs(config.validation_data)
|
|
125
|
+
iterator = iter(dataset)
|
|
126
|
+
logs = task.validation_step(next(iterator), model, metrics)
|
|
127
|
+
state = task.aggregate_logs(step_outputs=logs)
|
|
128
|
+
task.reduce_aggregated_logs(state)
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
class DetectionTFDSTest(tf.test.TestCase):
|
|
132
|
+
|
|
133
|
+
def test_train_step(self):
|
|
134
|
+
config = detr_cfg.DetrTask(
|
|
135
|
+
model=detr_cfg.Detr(
|
|
136
|
+
input_size=[1333, 1333, 3],
|
|
137
|
+
num_encoder_layers=1,
|
|
138
|
+
num_decoder_layers=1,
|
|
139
|
+
backbone=backbones.Backbone(
|
|
140
|
+
type='resnet',
|
|
141
|
+
resnet=backbones.ResNet(model_id=10, bn_trainable=False))
|
|
142
|
+
),
|
|
143
|
+
losses=detr_cfg.Losses(class_offset=1),
|
|
144
|
+
train_data=detr_cfg.DataConfig(
|
|
145
|
+
tfds_name='coco/2017',
|
|
146
|
+
tfds_split='validation',
|
|
147
|
+
is_training=True,
|
|
148
|
+
global_batch_size=2,
|
|
149
|
+
))
|
|
150
|
+
with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
|
|
151
|
+
task = detection.DetectionTask(config)
|
|
152
|
+
model = task.build_model()
|
|
153
|
+
dataset = task.build_inputs(config.train_data)
|
|
154
|
+
iterator = iter(dataset)
|
|
155
|
+
opt_cfg = optimization.OptimizationConfig({
|
|
156
|
+
'optimizer': {
|
|
157
|
+
'type': 'detr_adamw',
|
|
158
|
+
'detr_adamw': {
|
|
159
|
+
'weight_decay_rate': 1e-4,
|
|
160
|
+
'global_clipnorm': 0.1,
|
|
161
|
+
}
|
|
162
|
+
},
|
|
163
|
+
'learning_rate': {
|
|
164
|
+
'type': 'stepwise',
|
|
165
|
+
'stepwise': {
|
|
166
|
+
'boundaries': [120000],
|
|
167
|
+
'values': [0.0001, 1.0e-05]
|
|
168
|
+
}
|
|
169
|
+
},
|
|
170
|
+
})
|
|
171
|
+
optimizer = detection.DetectionTask.create_optimizer(opt_cfg)
|
|
172
|
+
task.train_step(next(iterator), model, optimizer)
|
|
173
|
+
|
|
174
|
+
def test_validation_step(self):
|
|
175
|
+
config = detr_cfg.DetrTask(
|
|
176
|
+
model=detr_cfg.Detr(
|
|
177
|
+
input_size=[1333, 1333, 3],
|
|
178
|
+
num_encoder_layers=1,
|
|
179
|
+
num_decoder_layers=1,
|
|
180
|
+
backbone=backbones.Backbone(
|
|
181
|
+
type='resnet',
|
|
182
|
+
resnet=backbones.ResNet(model_id=10, bn_trainable=False))
|
|
183
|
+
),
|
|
184
|
+
losses=detr_cfg.Losses(class_offset=1),
|
|
185
|
+
validation_data=detr_cfg.DataConfig(
|
|
186
|
+
tfds_name='coco/2017',
|
|
187
|
+
tfds_split='validation',
|
|
188
|
+
is_training=False,
|
|
189
|
+
global_batch_size=2,
|
|
190
|
+
))
|
|
191
|
+
|
|
192
|
+
with tfds.testing.mock_data(as_dataset_fn=_as_dataset):
|
|
193
|
+
task = detection.DetectionTask(config)
|
|
194
|
+
model = task.build_model()
|
|
195
|
+
metrics = task.build_metrics(training=False)
|
|
196
|
+
dataset = task.build_inputs(config.validation_data)
|
|
197
|
+
iterator = iter(dataset)
|
|
198
|
+
logs = task.validation_step(next(iterator), model, metrics)
|
|
199
|
+
state = task.aggregate_logs(step_outputs=logs)
|
|
200
|
+
task.reduce_aggregated_logs(state)
|
|
201
|
+
|
|
202
|
+
if __name__ == '__main__':
|
|
203
|
+
tf.test.main()
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
# Copyright 2025 The TensorFlow Authors. All Rights Reserved.
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
"""TensorFlow Model Garden Vision training driver."""
|
|
16
|
+
|
|
17
|
+
from absl import app
|
|
18
|
+
from absl import flags
|
|
19
|
+
import gin
|
|
20
|
+
|
|
21
|
+
from official.common import distribute_utils
|
|
22
|
+
from official.common import flags as tfm_flags
|
|
23
|
+
from official.core import task_factory
|
|
24
|
+
from official.core import train_lib
|
|
25
|
+
from official.core import train_utils
|
|
26
|
+
from official.modeling import performance
|
|
27
|
+
# pylint: disable=unused-import
|
|
28
|
+
from official.projects.detr.configs import detr
|
|
29
|
+
from official.projects.detr.tasks import detection
|
|
30
|
+
# pylint: enable=unused-import
|
|
31
|
+
|
|
32
|
+
FLAGS = flags.FLAGS
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def main(_):
|
|
36
|
+
gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
|
|
37
|
+
params = train_utils.parse_configuration(FLAGS)
|
|
38
|
+
model_dir = FLAGS.model_dir
|
|
39
|
+
if 'train' in FLAGS.mode:
|
|
40
|
+
# Pure eval modes do not output yaml files. Otherwise continuous eval job
|
|
41
|
+
# may race against the train job for writing the same file.
|
|
42
|
+
train_utils.serialize_config(params, model_dir)
|
|
43
|
+
|
|
44
|
+
# Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
|
|
45
|
+
# can have significant impact on model speeds by utilizing float16 in case of
|
|
46
|
+
# GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
|
|
47
|
+
# dtype is float16
|
|
48
|
+
if params.runtime.mixed_precision_dtype:
|
|
49
|
+
performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
|
|
50
|
+
distribution_strategy = distribute_utils.get_distribution_strategy(
|
|
51
|
+
distribution_strategy=params.runtime.distribution_strategy,
|
|
52
|
+
all_reduce_alg=params.runtime.all_reduce_alg,
|
|
53
|
+
num_gpus=params.runtime.num_gpus,
|
|
54
|
+
tpu_address=params.runtime.tpu)
|
|
55
|
+
with distribution_strategy.scope():
|
|
56
|
+
task = task_factory.get_task(params.task, logging_dir=model_dir)
|
|
57
|
+
|
|
58
|
+
train_lib.run_experiment(
|
|
59
|
+
distribution_strategy=distribution_strategy,
|
|
60
|
+
task=task,
|
|
61
|
+
mode=FLAGS.mode,
|
|
62
|
+
params=params,
|
|
63
|
+
model_dir=model_dir)
|
|
64
|
+
|
|
65
|
+
train_utils.save_gin_config(FLAGS.mode, model_dir)
|
|
66
|
+
|
|
67
|
+
if __name__ == '__main__':
|
|
68
|
+
tfm_flags.define_flags()
|
|
69
|
+
flags.mark_flags_as_required(['experiment', 'mode', 'model_dir'])
|
|
70
|
+
app.run(main)
|