Conversation
76ad16c to
52e0e98
Compare
0b10189 to
5448b5a
Compare
48eed14 to
7ee8613
Compare
nickfraser
left a comment
There was a problem hiding this comment.
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.
| if metadata_only: | ||
| self.value = None | ||
| self.quant_tensor = quant_tensor.set(value=None) | ||
| self.quant_tensor = quant_tensor.set(value=torch.empty(0)) |
There was a problem hiding this comment.
Any side-effects of this? Worth auditing anywhere we check for qt.value is None in the code or similar?
7ee8613 to
b929a77
Compare
5ed2434 to
1eb67b3
Compare
45dd1bb to
6e813ea
Compare
nickfraser
left a comment
There was a problem hiding this comment.
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.
| except Exception: | ||
| if enabled: | ||
| model.apply(lambda module: _override_cache_class(module, enabled=False)) | ||
| raise |
There was a problem hiding this comment.
Does a bare raise like this still produce a human-readable error?
There was a problem hiding this comment.
No, this must be changed.
| raise | ||
| finally: | ||
| if not enabled: | ||
| model.apply(lambda module: _override_cache_class(module, enabled=False)) |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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:
- caching is enabled
- 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, |
There was a problem hiding this comment.
Was this a bug before?
| @@ -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): | |||
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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: |
There was a problem hiding this comment.
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:
signedtraining__torch_function__deviceviewreshapeflattentransposepermute__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.
There was a problem hiding this comment.
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 valueFor 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() |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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): |
There was a problem hiding this comment.
These should go into floatmixin
| args = _unpack_quant_tensor(args) | ||
| kwargs = _unpack_quant_tensor(kwargs) | ||
| return func(*args, **kwargs) | ||
| def bit_width(self): |
There was a problem hiding this comment.
Maybe move to IntMixin
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
TruncAvgPool Fix
Breaking Changes
Tests
Added coverage for Tensor-subclass behavior, metadata reconstruction, autograd, device/dtype movement, arithmetic, groupwise storage, PyTorch dispatch, caching, and export-related paths.