[Fix][Relax][Frontend][Torch] Apply the scale and input_scale arguments of aten.elu - #20515
Arthur031221 wants to merge 2 commits into
Conversation
tlopex
left a comment
There was a problem hiding this comment.
This ReLU formulation only works when input_scale >= 0. For example, with x=1.5, alpha=0.5, scale=2, and input_scale=-1, PyTorch returns 3.0, but this implementation returns approximately 2.22313. Could we select the branch based on the original input (where(x < 0, alpha * (exp(input_scale * x) - 1), x)), then apply scale?
And could you move the test to the end of the file and make it a structural check instead of numerical check?
|
Thanks, you are right about negative For The PR description refers to the numerical test from the earlier revision. This follow-up replaces that test with the structural check you requested. |
Anyone importing a model that uses
torch.nn.SELUorF.seluthroughfrom_exported_programgets outputs that are about 5% too small, because the importer drops the SELU scale.from_exported_programrunsrun_decompositions()by default, which rewritesaten.selutoaten.elu(x, alpha, scale). Theelu.defaultconverter only readsalpha, soscale(1.0507...) andinput_scaleare ignored. The expected IR intest_extended_unary_opsfor SELU had the same omission, so the test passed on the wrong result.Reproduction on
F.selu(torch.tensor([[-2.0, -0.5, 0.0, 1.5]])), compiled with llvm and run on the VM:Changes:
ExportedProgramImporter._elunow appliesscaleandinput_scalefromaten.elu. With the defaults (1 and 1) it still calls the existing converter, so plainF.eluandnn.ELUproduce the same IR as before. The fx importer is untouched.test_extended_unary_opsgets the final multiply by the SELU scale.test_selu_and_elu_scale_argumentscompares against PyTorch numerically forF.seluand foraten.elu(x, 0.5, 2.0, 3.0). It fails without the converter change (3 of 4 elements mismatch, max absolute difference 0.076) and passes with it.Tests:
tests/python/relax/test_frontend_from_exported_program.pyandtest_frontend_from_fx.pygive 437 passed, 1 failed, 3 skipped with the change; the one failure istest_norm, which fails the same way before it.