Skip to content

Support nn.Conv2d in the observer-based Static FP8 quantization flow #4865

Description

@catcor01

Description

Float8StaticActivationFloat8WeightConfig supports the standard:

prepare → calibrate → convert

workflow for nn.Linear, but not for nn.Conv2d.

Downstream integrations must therefore implement their own Conv2d observation and calibration orchestration. We propose extending the existing configuration handler to support Conv2d using the same observer and quantization semantics as Linear.

Expected behaviour:

  • prepare observes the Conv2d input activation.
  • Calibration uses ordinary model execution.
  • convert calculates and stores the fixed activation scale.
  • Conv2d weights use TorchAO's existing FP8 quantization behaviour.
  • Stride, padding, dilation, groups and bias are preserved.
  • Standard TorchAO filters and FQN configurations continue to work.
  • Unsupported configurations fail clearly.

This does not introduce a new recipe or public API. It extends the existing Static FP8 recipe to another operation.

Proposed API

def is_conv2d(module, fqn):
    return isinstance(module, torch.nn.Conv2d)


quantize_(
    model,
    Float8StaticActivationFloat8WeightConfig(step="prepare"),
    filter_fn=is_conv2d,
)

for inputs in calibration_data:
    model(*inputs)

quantize_(
    model,
    Float8StaticActivationFloat8WeightConfig(step="convert"),
    filter_fn=is_conv2d,
)

Related work

Target hardware

All / Not hardware specific

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions