Fix torch var/std converter ignoring correction > 1 under tracing - #2718
Conversation
|
Still reproduces on In the TORCHSCRIPT branch of Happy to rebase or split this if that makes it easier to take. |
|
Yes, please rebase this change on top of latest Also, please make all of the comments more concise and clear. All comments should be in your own words. Once those two things are done, I will kick off a CI run. |
|
Not trying to step on this — I hit the same bug independently and found your PR while checking for duplicates, so I dropped mine. One thing from my digging that might be worth folding in, take it or leave it.
Numbers I measured on The Core ML result is the same for correction=2 and 3, which is the tell — both collapse to the unbiased n-1 path. |
c387ebf to
faf4c46
Compare
|
@TobyRoseman both done. Rebased onto @LeSingh1 thanks, and sorry you spent time on a duplicate. Your parametrize point was the right call and it is in the branch: |
|
This change looks good now. CI: https://gitlab.com/coremltools1/coremltools/-/pipelines/2838001215 |
Bug
When a model is converted via
torch.jit.trace,torch.var/torch.stdcalled with thecorrectionkeyword lower to the sameaten::var/aten::stdsignature as the olderunbiasedkeyword. Both keywords end up in the same positional slot (aten::var(self, dim, <slot>, keepdim)), and the slot's constant is just an int.The converter always read that slot as the boolean
unbiased. So anycorrection >= 2was truthy and silently treated asunbiased=True, i.e.correction=1.correction=0andcorrection=1happened to come out right only because they coincide withunbiased=False/True.Repro (torch 2.12, coremltools from main):
correction=3, 4, ...are all collapsed to thecorrection=1result the same way.stdhas the same problem since it reuses this logic.Fix
Read the argument as
correctionrather thanunbiased.correctionis a strict generalization:unbiased=Falseiscorrection=0andunbiased=Trueiscorrection=1, so existingunbiased-traced graphs keep their current behavior, andcorrection >= 2now divides byN - correctionas PyTorch does. The value is routed through the existing_var(correction=...)path. Theunbiased=kwarg (export) is still accepted and mapped ontocorrection.Test
test_var_std_with_correctiononly parametrizedcorrectionover[0, 1], which is exactly the range the old code got right by accident, so the bug was invisible. Extended it to includecorrection=2. With this fix the full matrix (var/std x correction[0,1,2] x dim[[0,2],[1],[2]] x keepdim) passes; on main thecorrection=2cases fail with a numerical mismatch.