Skip to content

Feat (quant_tensor)!: convert QuantTensor to Tensor subclass - #1579

Open
Giuseppe5 wants to merge 26 commits into
masterfrom
giuseppe/quant_tensor_subclass
Open

Giuseppe5 wants to merge 26 commits into
masterfrom
giuseppe/quant_tensor_subclass

Conversation

@Giuseppe5

@Giuseppe5 Giuseppe5 commented Aug 14, 2026 •

Copy link
Copy Markdown
Collaborator

Migrate QuantTensor to Tensor Subclasses

Summary

Migrate Brevitas quant tensors from NamedTuple containers to native torch.Tensor subclasses.
This preserves tensor behavior and autograd while storing quantization metadata as private attributes. The change also updates groupwise storage, runtime proxies, caching, pooling, export paths, and tests.

Changes

  • Convert QuantTensor and concrete quant tensor types to torch.Tensor subclasses.
  • Store metadata through _quant_tensor_metadata and private attributes such as _scale, _zero_point, and _bit_width.
  • Add shared reconstruction and tensor-operation support for set, detach, to, cpu, cuda, device checks, and PyTorch fallback dispatch.
  • Centralize generic arithmetic in QuantTensor; retain metadata-aware arithmetic for IntQuantTensor pairs.
  • Consolidate integer and floating-point metadata properties in IntMixin and FloatMixin.
  • Refactor groupwise quant tensors around shared compressed-storage handling.
  • Make groupwise shape operations operate on expanded values and return plain tensors.
  • Update metadata-only caching to use empty tensor storage instead of None.
  • Add metadata-based truncation for average-pooling paths, avoiding temporary quant-tensor reconstruction during tracing/export.
  • Update proxy and export code for the new metadata representation.
  • Correct float runtime metadata ordering during passthrough reconstruction.
  • Remove legacy fields and APIs.

TruncAvgPool Fix

  • Avoid reconstructing a temporary IntQuantTensor from cached metadata during tracing and export.
  • Pass the pooled tensor value and quantization metadata directly to the truncation proxy.
  • Preserve truncation behavior while avoiding Tensor-subclass reconstruction issues with FakeTensors.
  • Apply the same fix to TruncAvgPool2d and TruncAdaptiveAvgPool2d.

Breaking Changes

  • Quant tensors are no longer tuples or NamedTuple instances.
  • Tuple unpacking, _fields, and tuple-based serialization are no longer supported.
  • The .tensor property has been removed; use .value.
  • Legacy metadata fields have been removed:
    • value_
    • scale_
    • zero_point_
    • signed_t
    • training_t
    • saturating_t
  • Generic arithmetic with non-matching operands may now return a plain Tensor.
  • Unsupported PyTorch operations return plain tensors and may discard quantization metadata.
  • Metadata-only cached values are now empty tensors instead of None.

Tests

Added coverage for Tensor-subclass behavior, metadata reconstruction, autograd, device/dtype movement, arithmetic, groupwise storage, PyTorch dispatch, caching, and export-related paths.

@Giuseppe5
Giuseppe5 force-pushed the giuseppe/quant_tensor_subclass branch 3 times, most recently from 76ad16c to 52e0e98 Compare August 20, 2026 09:30
@Giuseppe5
Giuseppe5 force-pushed the giuseppe/quant_tensor_subclass branch 3 times, most recently from 0b10189 to 5448b5a Compare August 26, 2026 08:45
@Giuseppe5 Giuseppe5 self-assigned this Aug 26, 2026
@Giuseppe5
Giuseppe5 force-pushed the giuseppe/quant_tensor_subclass branch from 48eed14 to 7ee8613 Compare September 1, 2026 12:46

@nickfraser nickfraser left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

My feeling is that we want to change the design pattern here. Current QuantTensor's methods are relatively complicated so that they may be used unmodified in the all the derived classes. My personal preference is to keep QuantTensor simple / readable and to override these methods in the derived classes through a Mixin or similar - unless there is a good reason to avoid this that I'm not seeing.

Comment thread notebooks/minifloat_mx_tutorial.ipynb Outdated
Comment thread src/brevitas/quant_tensor/base_quant_tensor.py
Comment thread src/brevitas/quant_tensor/base_quant_tensor.py
Comment thread src/brevitas/utils/quant_utils.py Outdated
if metadata_only:
self.value = None
self.quant_tensor = quant_tensor.set(value=None)
self.quant_tensor = quant_tensor.set(value=torch.empty(0))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Any side-effects of this? Worth auditing anywhere we check for qt.value is None in the code or similar?

@Giuseppe5
Giuseppe5 force-pushed the giuseppe/quant_tensor_subclass branch from 7ee8613 to b929a77 Compare September 2, 2026 13:56
@Giuseppe5
Giuseppe5 requested a review from nickfraser September 8, 2026 09:32
@Giuseppe5
Giuseppe5 force-pushed the giuseppe/quant_tensor_subclass branch from 5ed2434 to 1eb67b3 Compare September 9, 2026 13:51
Comment thread src/brevitas/quant_tensor/base_quant_tensor.py
@Giuseppe5
Giuseppe5 force-pushed the giuseppe/quant_tensor_subclass branch from 45dd1bb to 6e813ea Compare September 23, 2026 14:34

@nickfraser nickfraser left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Quite a few comments - we need to be a bit cautious here. This PR is growing legs from what originally was a sideways movement from NamedTuple -> QuantTensor, with some cleanup included. We seem to be also changing:

  • some behaviour
  • some public APIs

The more we change, the more it feels like we should "do everything" in this PR, which would make it unwieldy and difficult to review. We might want to talk about scope here.

Essentially, we've deduped some code and left others duped, which leaves this PR is some "halfway" point between a proper refactor or not.

Comment thread src/brevitas/export/onnx/manager.py Outdated
except Exception:
if enabled:
model.apply(lambda module: _override_cache_class(module, enabled=False))
raise

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Does a bare raise like this still produce a human-readable error?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

No, this must be changed.

Comment thread src/brevitas/export/onnx/manager.py Outdated
raise
finally:
if not enabled:
model.apply(lambda module: _override_cache_class(module, enabled=False))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It's not clear to be if there is a change in behaviour here or if we're just guarding against future failures. Can you elaborate?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is a new issues introduced with Tensor subclass instead of NamedTuple; it is a bit lengthy to explain here but the main issue is what happens when:

  1. caching is enabled
  2. and we try to perform operations like quant_tensor.set(value=new_value)

With NamedTuple this was not an issue.

x.exponent_bit_width,
x.mantissa_bit_width,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Was this a bug before?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes

@@ -279,7 +279,7 @@ def forward(
self._cached_act = cached_inp

if self.is_quant_enabled:
if quant_input is None or isinstance(quant_input, Tensor):
if quant_input is None or not isinstance(quant_input, IntQuantTensor):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This logic changed here - it looks intentional, but thought I'd flag it.

Plain Tensors / QuantTensors will resolve to True and will be replaced with the cached value... Do we want this to occur for a QuantTensor that is not an IntQuantTensor?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is only compatible with IntQuantTensor inputs (and it only produces IntQuantTensor weights, since this inherits from DecoupledWeightQuantWithInputProxyFromInjector which inherits from WeightQuantProxyFromInjector, which is the Int-specific implementation of weight proxy).

I'll add a specific assert to reject other QT

signed_t: Tensor
training_t: Tensor
dequant_shape: Optional[Tuple] = None
class GroupwiseQuantTensorMixin:

@nickfraser nickfraser Sep 24, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why are there so many methods defined here that are unrelated to Groupwise-specific logic. My hope was to avoid lots of duplicated logic my introducing the Mixins. Is there a reason why these should live in the Groupwise Mixin?

Methods that don't seem to be related to groupsize-specific logic:

  • signed
  • training
  • __torch_function__
  • device
  • view
  • reshape
  • flatten
  • transpose
  • permute
  • __add__
  • __mul__
  • __truediv__

A follow-up question is: should view, reshape, flatten, transpose, permute have groupwise-specific implementations, but we can leave that for a follow-up PR.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I've cleaned up a lot of them, but all the views have to be Groupwise specific.

For non-groupwise quanttensor we made the decision of allowing unsafe operations, i.e. you can flatten a tensor but the quant metadata will not receive the same view, even though a new QuantTensor is generated.

# Regular IntQuantTensor: unsafe reshape is allowed.
qt = IntQuantTensor(
    value=torch.ones(2, 4),
    scale=torch.ones(1, 4),  # Per-channel metadata.
    zero_point=torch.zeros(1, 4),
    bit_width=torch.tensor(8.),
    signed=True,
    training=False,
)

flat = qt.view(-1)

type(flat)       # IntQuantTensor
flat.value.shape # torch.Size([8])
flat.scale.shape # torch.Size([1, 4])  <- unchanged, no longer aligned with value

For groupwise quant tensor, generating a new QT with an unsafe view will cause issues everytime .value is called (or .scale, .zero_point), which is why we override and explicitly demote to normal Tensor.
If we preserved QT structure:

# GroupwiseIntQuantTensor: raw storage is compressed.
qt = GroupwiseIntQuantTensor(
    value=compressed_value,     # e.g. shape [2, 4, 8]
    scale=compressed_scale,
    zero_point=compressed_zp,
    group_size=8,
    group_dim=1,
    bit_width=torch.tensor(8.),
    signed=True,
    training=False,
    dequant_shape=(2, 32),      # Logical expanded shape.
)

flat = qt.view(-1)
bad_qt.value       # Expansion uses invalid storage/group geometry.
bad_qt.scale       # Same issue.
bad_qt.zero_point  # Same issue.


@property
def signed(self):
return self.signed_t.item()
return self._signed.item()

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Other renaming of attributes have just been related to private values. Does the old version no longer exist? Do we want a breaking change like this with this PR?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think so; The point is that we have two distinct ways to indicate a "private" attribute currently in master: value_ (NamedTuple doesn't like leading underscore), and training_t (since we store a tensor but we always retrieve a bool).

The idea is to have leading underscore variables as storage variables in all cases; I can also revert training_t/signed_t and leave them as an exception, but since we're changing stuff here (and the fact that this will be a breaking change is almost unavoidable anyway), I'd rather do it here.


@property
def signed(self):
return self.signed_t.item()
def exponent_bit_width(self):

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

These should go into floatmixin

args = _unpack_quant_tensor(args)
kwargs = _unpack_quant_tensor(kwargs)
return func(*args, **kwargs)
def bit_width(self):

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Maybe move to IntMixin

@Giuseppe5 Giuseppe5 changed the title Giuseppe/quant tensor subclass Feat (quant_tensor)!: convert QuantTensor to Tensor subclass Sep 30, 2026

This branch has not been deployed

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants