Skip to content

feat(casting): add dynamic activation range calibration in FP16 casting (#7) - #117

Merged
crowbat merged 1 commit into
apple:mainfrom
rohith500:feat/casting-activation-calibration
Oct 2, 2026
Merged

crowbat merged 1 commit into
apple:mainfrom
rohith500:feat/casting-activation-calibration

Conversation

@rohith500

@rohith500 rohith500 commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Closes #7.

Adds dynamic activation range calibration via an optional calibration_data parameter in cast_fp32_to_fp16 and cast_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).

Comment thread tests/casting/test_casting.py Outdated


# =============================================================================
# Dynamic Activation Range Calibration tests (Issue #7)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Let's remove "Issue #7" from the comment, keep the code comments self contained.
I also missed last time that "#7" was added to TestSelectiveOpSkipping's class docstring above, please remove that as well.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Done, removed the issue references from the comments and docstrings.

Comment thread src/coreai_opt/_utils/casting_utils.py Outdated

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):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Instead of defining our own checking logic here, can we reuse check_tensor_overflow_fp16? This also buys us 2 things:

  1. 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.

  2. We can treat inf the 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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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.

Comment thread src/coreai_opt/_utils/casting_utils.py Outdated
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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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).

Comment thread src/coreai_opt/_utils/casting_utils.py Outdated
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:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Done, removed threshold from _OverflowDetector and find_overflowing_nodes and hardcoded _FP16_MAX internally.

Comment thread src/coreai_opt/_utils/casting_utils.py Outdated
yield _to_inputs(calibration_data)
return

if (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Simplified the interface so calibration_data must always be an iterable of samples (and raises a TypeError if passed a bare Tensor or dict).

Comment thread src/coreai_opt/_utils/casting_utils.py Outdated
elif isinstance(val, (tuple, list)):
for item in val:
self._check(n, item)
elif isinstance(val, dict):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Done, removed the dict branch.

Comment thread src/coreai_opt/_utils/casting_utils.py Outdated
f"Calibration sample provided {len(sample)} input(s), "
f"but model expects {num_inputs} input(s): {user_names}."
)
return list(sample[:num_inputs])

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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.

@rohith500
rohith500 force-pushed the feat/casting-activation-calibration branch from 4f627ae to ab04f6c Compare October 1, 2026 22:53

@crowbat crowbat left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Thanks for your changes @rohith500 , looks good.

@crowbat
crowbat merged commit 4982c4b into apple:main Oct 2, 2026
14 checks passed
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.

FP16 casting pass does not guard against activation-level overflow (softplus, exp, logsumexp)

2 participants