tensorflow / tensorflow/tflite-support

Inference time for tflite quantized model is high

Open
#980 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
C++
Stars
441
Forks
146
PR merge metrics
No merged PRs in 30d

Description

I followed a tutorial on MediaPipe and their model, https://storage.googleapis.com/mediapipe-models/image_classifier/efficientnet_lite0/float32/1/efficientnet_lite0.tflite, has inference time of milliseconds.

I used ai-edge-torch to convert a PyTorch efficientnet to tflite, but the inference time is 2-3 seconds.

Here's my code:

efficientnet = torchvision.models.efficientnet_b3(torchvision.models.EfficientNet_B3_Weights.IMAGENET1K_V1).eval()

class PermuteInput(nn.Module):
    def __init__(self):
        super(PermuteInput, self).__init__()

    def forward(self, x):
        # Permute from (batch, height, width, channels) to (batch, channels, height, width)
        return x.permute(0, 3, 1, 2)


import torch.nn.functional as F

class PermuteOutput(nn.Module):
    def __init__(self):
        super(PermuteOutput, self).__init__()

    def forward(self, x):
        return F.normalize(x)

efficientnet_with_reshape = nn.Sequential(
    PermuteInput(),
    efficientnet,
    PermuteOutput()
)


edge_model = efficientnet_with_reshape.eval()

sample_input = (torch.rand((1, 224, 224, 3), dtype=torch.float32),)

edge_model = ai_edge_torch.convert(edge_model.eval(), sample_input)

edge_model.export("/home/user/efficientnet.tflite")

# QUANTIZE TFLITE MODEL

pt2e_quantizer = PT2EQuantizer().set_global(
    get_symmetric_quantization_config(is_per_channel=True, is_dynamic=True)
)

pt2e_torch_model = capture_pre_autograd_graph(efficientnet_with_reshape.eval(),sample_input)
pt2e_torch_model = prepare_pt2e(pt2e_torch_model, pt2e_quantizer)

# Run the prepared model with sample input data to ensure that internal observers are populated with correct values
pt2e_torch_model(*sample_input)

# Convert the prepared model to a quantized model
pt2e_torch_model = convert_pt2e(pt2e_torch_model, fold_quantize=False)

# Convert to an ai_edge_torch model
pt2e_drq_model = ai_edge_torch.convert(pt2e_torch_model, sample_input, quant_config=QuantConfig(pt2e_quantizer=pt2e_quantizer))


pt2e_drq_model.export("/home/user/efficientnet_quantized.tflite")

I properly added metadata to tflite, labels and also added a CORS policy to the bucket.

Is this a quantization issue or a bucket bandwidth issue? Because with the supported model, the inference is really fast.

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with the ai_edge_torch.convert calls and the PT2E quantization setup shown in the issue, then compare inference timing for the two exported TFLite models. Check whether the slowdown occurs during model inference or only during deployment, and finish by identifying whether conversion, quantization, or bandwidth accounts for the reported difference.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning, performance
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
30/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.