Repository navigation
feat(casting): add dynamic activation range calibration in FP16 casting (#7) - #117
Conversation
b17607b to
4f627ae
Compare
|
|
||
|
|
||
| # ============================================================================= | ||
| # Dynamic Activation Range Calibration tests (Issue #7) |
There was a problem hiding this comment.
Done, removed the issue references from the comments and docstrings.
|
|
||
| def _check(self, n: torch.fx.Node, val: Any) -> None: | ||
| if isinstance(val, torch.Tensor) and val.is_floating_point() and val.numel() > 0: | ||
| if not torch.isfinite(val).all() or (val.detach().abs().max() > self.threshold): |
There was a problem hiding this comment.
Instead of defining our own checking logic here, can we reuse check_tensor_overflow_fp16? This also buys us 2 things:
-
We would be checking that the dtype is fp32 before adding to
self.direct_overflow_nodes. The check should only flag nodes which would have been downcast by this casting pass and not nodes which were downcast by the original model definition. -
We can treat
infthe same way we do for placeholder tensors where we don't actually mark it as overflowing. We would only prevent casting to fp16 if the range was > FP16_MAX but < fp32's max. If true inf is seen even in fp32, we can treat that as an inf in fp16 as well and allow the casting to be done.
Would be good to add a unit test to check the inf case too.
There was a problem hiding this comment.
Updated to reuse check_tensor_overflow_fp16 directly. It filters for FP32 dtype and handles finite vs inf values as suggested. Also added a unit test for the inf case.
| curr = queue.pop() | ||
| for user in curr.users: | ||
| if user.op == "call_function" and user not in all_overflow_nodes: | ||
| all_overflow_nodes.add(user) |
There was a problem hiding this comment.
Similar to the above comment, we can check that user outputs fp32 before adding it to all_overflow_nodes. That way we only act on nodes which would have been affected as part of our pass.
There was a problem hiding this comment.
Done, added a check so consumer nodes are only added if their output dtype is FP32.
| out = _run_ep(ep, x.half()) | ||
|
|
||
| assert not torch.isinf(out).any(), "Expected finite output with calibration_data" | ||
| assert torch.allclose(out.float(), ref, atol=1e-3) |
There was a problem hiding this comment.
Can we also check that there are cast to fp32 and cast back to fp16 nodes in the expected places in the graph, and that the overflowing ops are still outputting fp32 dtype?
Would be good to enhance the checks for the other unit tests as well that are testing that overflowing ops are remaining in fp32.
There was a problem hiding this comment.
Added assertions to verify the overflowing ops keep FP32 dtype, along with the _to_copy boundary casts (cast to FP32 before, and cast back to FP16 after).
| class _OverflowDetector(Interpreter): | ||
| """Interpreter that tracks call_function nodes producing out-of-range FP16 values.""" | ||
|
|
||
| def __init__(self, ep: torch.export.ExportedProgram, threshold: float = _FP16_MAX) -> None: |
There was a problem hiding this comment.
Currently there is no way to change threshold from user callable code and I don't see a need to use any other value other than FP16_MAX anyways. Can we remove the argument completely from _OverflowDetector and find_overflowing_nodes and just hardcode the usage of _FP16_MAX within _OverflowDetector?
There was a problem hiding this comment.
Done, removed threshold from _OverflowDetector and find_overflowing_nodes and hardcoded _FP16_MAX internally.
| yield _to_inputs(calibration_data) | ||
| return | ||
|
|
||
| if ( |
There was a problem hiding this comment.
This bit of code is a little hard to follow and seems to be trying to account for ambiguity in how calibration_data is presented where we are not sure if it's meant to be a list of samples, or instead just a single sample that happens to be a list.
Can we simplify the interface by mandating that when calibration data is iterated on, it always presents an element that represents a sample? So if a user wanted to pass a single sample for calibration, they should first wrap it as a single element list/tuple and present that as calibration_data.
To be precise, I'd like to define the interface as such:
calibration_data is always an iterable of samples (or None if no dynamic calibration is to be done). Iterating it yields one sample at a time where each sample is either a Tensor (one positional input), a tuple/list (positional inputs), or a dict (kwargs).
And as a result of that, calibration_data should never itself be a torch.Tensor or a dict, even though those are Iterables, so we can add some guards against that to prevent users accidentally passing those in.
There was a problem hiding this comment.
Simplified the interface so calibration_data must always be an iterable of samples (and raises a TypeError if passed a bare Tensor or dict).
| elif isinstance(val, (tuple, list)): | ||
| for item in val: | ||
| self._check(n, item) | ||
| elif isinstance(val, dict): |
There was a problem hiding this comment.
To my knowledge, torch graph nodes can't output dicts. Let me know if this is incorrect, but I think we can remove this part as it would be dead code.
There was a problem hiding this comment.
Done, removed the dict branch.
| f"Calibration sample provided {len(sample)} input(s), " | ||
| f"but model expects {num_inputs} input(s): {user_names}." | ||
| ) | ||
| return list(sample[:num_inputs]) |
There was a problem hiding this comment.
Let's not try to handle samples with labels using this slicing logic since it also allows any malformed input with more elements than expected to pass through. Related to the comment for line 920, let's just restrict calibration samples to be well-structured so we don't try to handle different corner cases ourselves.
There was a problem hiding this comment.
Removed the slicing logic; samples now require an exact match with the expected number of inputs, otherwise a ValueError is raised.
| # ============================================================================= | ||
| # Dynamic Activation Range Calibration tests (Issue #7) | ||
| # ============================================================================= | ||
| class TestDynamicActivationCalibration: |
There was a problem hiding this comment.
Can we have one more test in which there are multiple distinct samples in calibration_samples, each of which would result in a different part of the graph marked as overflow? So the expectation is that the union of all nodes marked as overflow should be run in fp32 after the cast pass is all finished.
There was a problem hiding this comment.
Added test_dynamic_calibration_union_of_overflowing_nodes_across_samples, which passes two samples overflowing separate branches and checks that both branches are kept in FP32 with appropriate boundary casts.
4f627ae to
ab04f6c
Compare
crowbat
left a comment
There was a problem hiding this comment.
Thanks for your changes @rohith500 , looks good.
Closes #7.
Adds dynamic activation range calibration via an optional
calibration_dataparameter incast_fp32_to_fp16andcast_to_16_bit_precision. Intermediate activations exceeding the IEEE 754 float16 representable range (> 65,504.0 or NaN/Inf) during calibration simulation are automatically preserved in FP32 precision along with their downstream consumers, with native support for standard PyTorch DataLoaders yielding(inputs, labels).