`tf.split` or `tf.transpose` cause errors for quantize-aware training with `quantize_apply`
@cdh4696 已经在做这个了。
开始于 2023年5月2日。
评估
这个 Issue 还没有评估数据。
描述
Describe the bug
We are trying to implement some network like ShuffleNetV2 but encounter some error when quantize_apply the model.

I believe ShuffleNet or related ideas are popular in edge devices, please kindly help us to resolve this proble.
Any advice is welcome.
System information
TensorFlow version (installed from source or binary): 2.7.0
TensorFlow Model Optimization version (installed from source or binary): 0.7.0
Python version: 3.8.13
Describe the expected behavior
Just add quantization-aware operator in to the model.
Describe the current behavior
When running the provided code, either the tf.transpose or tf.split will cause error to Tensorflow Model Optimization.
The error message due to tf.split before convolution layers:
ValueError: Exception encountered when calling layer "bn3" (type BatchNormalization).
Shape must be rank 4 but is rank 5 for '{{node bn3/FusedBatchNormV3}} = FusedBatchNormV3[T=DT_FLOAT, U=DT_FLOAT, data_format="NHWC", epsilon=0.001, exponential_avg_factor=1, is_training=false](Placeholder, bn3/ReadVariableOp, bn3/ReadVariableOp_1, bn3/FusedBatchNormV3/ReadVariableOp, bn3/FusedBatchNormV3/ReadVariableOp_1)' with input shapes: [1,?,128,128,32], [32], [32], [32], [32].
The error message due to tf.transpose:
ValueError: Exception encountered when calling layer "tf.compat.v1.transpose" (type TFOpLambda).
Dimension must be 6 but is 5 for '{{node tf.compat.v1.transpose/transpose}} = Transpose[T=DT_FLOAT, Tperm=DT_INT32](tf.compat.v1.transpose/transpose/a, tf.compat.v1.transpose/transpose/perm)' with input shapes: [1,?,128,128,2,32], [5].
Code to reproduce the issue
Just run the following code you will get the error message due to tf.split.
from __future__ import annotations
from typing import Callable, Optional
import tensorflow as tf
import tensorflow_model_optimization as tfmot
from tensorflow.keras import layers
SKIP_LAYER = [
"resize",
"Resize",
"reshape",
"Reshape",
"concat",
"Concat" "ExpandDims",
"Repeats",
"Shape",
"strided_slice",
"Tile",
]
def quantize_model(
model: tf.keras.Model,
annotate: Optional[Callable] = None,
quantize_scope: Optional[dict[str, tf.keras.layers.Layer]] = None,
) -> tf.keras.Model:
quantize_scope = {} if quantize_scope is None else quantize_scope
def annotate(layer):
if any([name in layer.name for name in SKIP_LAYER]):
return layer
else:
return tfmot.quantization.keras.quantize_annotate_layer(layer)
anno_model = tf.keras.models.clone_model(model, clone_function=annotate)
with tfmot.quantization.keras.quantize_scope(quantize_scope):
model = tfmot.quantization.keras.quantize_apply(anno_model)
return model
def channel_shuffle(tensor: tf.Tensor, groups: int = 2) -> tf.Tensor:
"""Channel shuffle operation."""
_, height, width, num_channels = tensor.shape.as_list()
assert num_channels % groups == 0
tensor = tf.reshape(tensor, [-1, height, width, groups, num_channels // groups])
tensor = tf.transpose(tensor, [0, 1, 2, 4, 3])
tensor = tf.identity(tensor, name="channel_shuffle")
tensor = tf.reshape(tensor, [-1, height, width, num_channels])
return tensor
def simple_nn(img_input: tf.Tensor) -> tf.Tensor:
latent = layers.Conv2D(32, 1, padding="same", use_bias=False, name="conv1")(img_input)
latent = layers.BatchNormalization(name="bn1")(latent)
latent = layers.ReLU(name="relu1")(latent)
latent = layers.DepthwiseConv2D(3, 1, padding="same", name="conv2")(img_input)
latent = layers.BatchNormalization(name="bn2")(latent)
latent = layers.Conv2D(32, 1, padding="same", use_bias=False, name="conv3")(img_input)
latent = layers.BatchNormalization(name="bn3")(latent)
latent = layers.ReLU(name="relu3")(latent)
return latent
def split_like_nn(img_input: tf.Tensor) -> tf.Tensor:
latent = layers.Conv2D(64, 1, padding="same", use_bias=False, name="conv0")(img_input)
latent = layers.BatchNormalization(name="bn0")(latent)
latent = layers.ReLU(name="relu0")(latent)
latent_0, latent_1 = tf.split(latent, 2, axis=-1)
latent_0 = simple_nn(latent_0)
latent = tf.concat([latent_0, latent_1], axis=-1)
latent = channel_shuffle(latent)
return latent
if __name__ == "__main__":
img_input = tf.keras.Input((128, 128, 1), dtype=tf.float32, name="img")
outputs = split_like_nn(img_input)
model = tf.keras.Model(inputs=img_input, outputs=outputs, name="PoseNetV2")
model.summary()
model_qat = quantize_model(model)
model_qat.summary()
You can just comment the following three lines of code will get the error message from tf.transpose.
latent_0, latent_1 = tf.split(latent, 2, axis=-1)
latent_0 = simple_nn(latent_0)
latent = tf.concat([latent_0, latent_1], axis=-1)
- 主要语言
- Python
- 星标
- 1.6k
- 派生
- 349
- PR 合并指标
- 30 天内没有已合并 PR
环境准备
- 没有 Dockerfile 或 Docker Compose 文件
- 没有 Pull Request 模板
- 阅读贡献指南
从这里开始
- 先读完整个 Issue,再读项目的贡献指南。
- 在 Issue 下留言说明你要接手 —— 这能避免两个人做同样的事。
- Fork 仓库,在一个分支上完成修改。
- 提交 Pull Request,并在描述里引用这个 Issue 编号。
tensorflow/model-optimization 的其他 Issue
-
难度 2/5 1-3 小时 新手友好度 74/100
tensorflow/model-optimization#1301 · 1 条评论 ·
-
bug
难度 4/5 3-5 天 新手友好度 48/100
tensorflow/model-optimization#1272 ·
-
bug
难度 3/5 1-2 天 新手友好度 50/100
tensorflow/model-optimization#1270 · 2 条评论 ·
-
bug
难度 2/5 1-3 小时 新手友好度 48/100
tensorflow/model-optimization#1241 ·
-
bug
难度 4/5 3-5 天 新手友好度 35/100
tensorflow/model-optimization#1182 · 1 条评论 ·
查看 tensorflow/model-optimization 的全部 Issue
相似的 Issue
-
bug status/needs-triage
难度 2/5 1-3 小时 新手友好度 86/100
prowler-cloud/prowler#12887 · 1 条评论 ·
维护者通常 1 天内回复
-
area: desktop platform: macos priority: p3 status: ready type: enhancement
难度 1/5 1 小时以内 新手友好度 92/100
use-agent-os/agent-os#3484 ·
维护者通常 2 天内回复
-
bug
难度 2/5 1-3 小时 新手友好度 86/100
open-telemetry/opentelemetry-python-contrib#5113 · 2 条评论 · 2 个 reaction ·
维护者通常 1 天内回复
-
external
难度 2/5 1-3 小时 新手友好度 68/100
langchain-ai/docs#6255 ·
维护者通常 1 天内回复
-
难度 2/5 1-3 小时 新手友好度 72/100
维护者通常 1 天内回复