Allow infinite log odds in categorical_logit and multinomial_logit - #3412
godking123 wants to merge 11 commits into
Conversation
…tests and case comments
|
Hi, @godking123, and thanks for the PR. Given that this is in our C++ code, we're going to need someone like @SteveBronder to review it. I'll leave some comments, but don't feel confident enough about our C++ style and library to do the final review. |
bob-carpenter
left a comment
There was a problem hiding this comment.
I was confused in reading the code about why the checks were laid out how they were. The documentation also needs to change for the top-level function to indicate what does and doesn't throw.
Otherwise, I'm going to leave this to @SteveBronder.
| // Autodiff args must be finite, data can be +/-inf | ||
| if constexpr (is_constant<T_prob>::value) { | ||
| check_not_nan(function, "log odds parameter", beta_ref); | ||
| check_greater(function, "log odds parameter", beta_ref.maxCoeff(), |
There was a problem hiding this comment.
Why is this allowing +infinity, but not -infinity? I thought the issue was to have -infinity turn into zero.
There was a problem hiding this comment.
This check is for the all -inf case, it throws when the maximum in the beta vector is -inf. I realize this test isn't very straightforward I can reword it or add a comment
| check_greater(function, "log odds parameter", beta_ref.maxCoeff(), | ||
| NEGATIVE_INFTY); | ||
| } else { | ||
| check_finite(function, "log odds parameter", beta_ref); |
There was a problem hiding this comment.
Why are we checking the log odds parameter is finite? Isn't this allowed to be infinite?
There was a problem hiding this comment.
This is in the autodiff case, when it's a data variable it's allowed to be infinite
| if constexpr (is_constant<T_prob>::value) { | ||
| int num_infty = (beta_ref.array() == INFTY).count(); | ||
| if (num_infty > 0) { | ||
| return beta_ref.coeff(n - 1) == INFTY ? -std::log(num_infty) |
There was a problem hiding this comment.
Why are we taking the negative log of the number of infinite coefficients?
There was a problem hiding this comment.
We need log(1/num_infty) for log probability since the probability is 1/num_infty. -log(num_infty) is an equivalent expression to this
| ref_type_t<T_prob> beta_ref = beta; | ||
| check_finite(function, "log odds parameter", beta_ref); | ||
|
|
||
| // Autodiff args must be finite, data can be +/-inf |
There was a problem hiding this comment.
Thanks---this comment is useful---maybe something like this above to explain the range checks?
| // Autodiff args must be finite, data can be +/-inf | ||
| if constexpr (is_constant<T_prob>::value) { | ||
| check_not_nan(function, "log odds parameter", beta_ref); | ||
| if (beta_ref.size() > 0) { |
There was a problem hiding this comment.
Is beta_ref.size() == 0 legal at this point or was it short-circuited earlier?
There was a problem hiding this comment.
Yes should be legal it's reachable by an empty beta vector
There was a problem hiding this comment.
See my other comment. I think we should check this much earlier so we do not have to litter the code with these if statements.
|
Hi @bob-carpenter, thanks for the review. I've added comments clarifying the bounds checks and and updated the docstrings to spell out what throws, could you restart the CI checks for the latest commit whenever you get the chance? |
SteveBronder
left a comment
There was a problem hiding this comment.
Thanks! Most of the comments are just to clean some pieces up. I'm not super familiar on how to handle infinites for these distributions. Would you mind adding some docs or citing some sources on the reasoning for these? I just recently did a review of our std_normal_lcdf code that had a custom approximation. I was digging around trying to find where he took his impl from and after finding the original PR it turns out the guy just freestyle'd it and only had a few sentences of explanation in the PR lol. So if you would not mind citing your reasoning in the doxygen that would be great.
Overall looks good!
|
|
||
| // softmax is NaN with +inf, so split probability evenly over the | ||
| // +inf entries and give every other entry zero probability | ||
| if constexpr (is_constant<T_prob>::value) { |
There was a problem hiding this comment.
This applies for all is_constant<...>::value
| if constexpr (is_constant<T_prob>::value) { | |
| if constexpr (is_constant_v<T_prob>) { |
| if constexpr (is_constant<T_prob>::value) { | ||
| // Data Case: Throws in nan and all -inf case | ||
| check_not_nan(function, "log odds parameter", beta_ref); | ||
| check_greater(function, "log odds parameter", beta_ref.maxCoeff(), |
There was a problem hiding this comment.
I think we should check that beta_ref.size() > 0 much earlier in the code. Then we can get rid of all of the if(beta_ref.size() > 0) littered everywhere else
| // Autodiff args must be finite, data can be +/-inf | ||
| if constexpr (is_constant<T_prob>::value) { | ||
| check_not_nan(function, "log odds parameter", beta_ref); | ||
| if (beta_ref.size() > 0) { |
There was a problem hiding this comment.
See my other comment. I think we should check this much earlier so we do not have to litter the code with these if statements.
| lp += beta_ref.coeff(n - 1) == INFTY ? -std::log(num_infty) | ||
| : NEGATIVE_INFTY; |
There was a problem hiding this comment.
If you are going to add NEGATIVE_INFTY to lp then I think you can just return early with NEGATIVE_INFTY
| check_greater(function, "Log odds parameter", beta.maxCoeff(), | ||
| NEGATIVE_INFTY); |
There was a problem hiding this comment.
We should add the same early > 0 size check here as well
| // Data Case: Throws in nan and all -inf case | ||
| check_not_nan(function, "log-probabilities parameter", beta_ref); | ||
| // maxCoeff() is undefined for an empty beta | ||
| if (beta_ref.size() > 0) { |
There was a problem hiding this comment.
same size check scheme here as well
|
@SteveBronder: The reasoning for all of these is the same---it's the natural limiting value. That is, for The same thing works when there are two values approaching infinity at the same rate, we have And so on. Is this enough or are you looking for more justification? We don't have this level of doc elsewhere for our limiting conditions, so we should probably add it. |
|
Lol oh sorry yeah that makes total sense. Yeah just a one liner for anyone passing by would still be nice to have |
|
Thanks @SteveBronder and @bob-carpenter for the review! I made the requested changes. I replaced each is_constant::value with is_constant_v, moved the empty beta checks up to an early return, return -inf immediately when an outcome has zero probability, and added a comment to the docs explaining the softmax limits for infinite values. |
Jenkins Console Log Machine informationDistributor ID: Ubuntu Description: Ubuntu 20.04.3 LTS Release: 20.04 Codename: focal CPU: Architecture: x86_64 CPU op-mode(s): 32-bit, 64-bit Byte Order: Little Endian Address sizes: 52 bits physical, 57 bits virtual CPU(s): 192 On-line CPU(s) list: 0-191 Thread(s) per core: 2 Core(s) per socket: 48 Socket(s): 2 NUMA node(s): 2 Vendor ID: AuthenticAMD CPU family: 25 Model: 17 Model name: AMD EPYC 9474F 48-Core Processor Stepping: 1 Frequency boost: enabled CPU MHz: 1497.368 CPU max MHz: 4114.4229 CPU min MHz: 1500.0000 BogoMIPS: 7189.39 Virtualization: AMD-V L1d cache: 3 MiB L1i cache: 3 MiB L2 cache: 96 MiB L3 cache: 512 MiB NUMA node0 CPU(s): 0-47,96-143 NUMA node1 CPU(s): 48-95,144-191 Vulnerability Gather data sampling: Not affected Vulnerability Indirect target selection: Not affected Vulnerability Itlb multihit: Not affected Vulnerability L1tf: Not affected Vulnerability Mds: Not affected Vulnerability Meltdown: Not affected Vulnerability Mmio stale data: Not affected Vulnerability Old microcode: Not affected Vulnerability Reg file data sampling: Not affected Vulnerability Retbleed: Not affected Vulnerability Spec rstack overflow: Mitigation; Safe RET Vulnerability Spec store bypass: Mitigation; Speculative Store Bypass disabled via prctl Vulnerability Spectre v1: Mitigation; usercopy/swapgs barriers and __user pointer sanitization Vulnerability Spectre v2: Mitigation; Enhanced / Automatic IBRS; IBPB conditional; STIBP always-on; PBRSB-eIBRS Not affected; BHI Not affected Vulnerability Srbds: Not affected Vulnerability Tsa: Mitigation; Clear CPU buffers Vulnerability Tsx async abort: Not affected Vulnerability Vmscape: Mitigation; IBPB before exit to userspace Flags: fpu vme de pse tsc msr pae mce cx8 apic sep mtrr pge mca cmov pat pse36 clflush mmx fxsr sse sse2 ht syscall nx mmxext fxsr_opt pdpe1gb rdtscp lm constant_tsc rep_good amd_lbr_v2 nopl xtopology nonstop_tsc cpuid extd_apicid aperfmperf rapl pni pclmulqdq monitor ssse3 fma cx16 pcid sse4_1 sse4_2 x2apic movbe popcnt aes xsave avx f16c rdrand lahf_lm cmp_legacy svm extapic cr8_legacy abm sse4a misalignsse 3dnowprefetch osvw ibs skinit wdt tce topoext perfctr_core perfctr_nb bpext perfctr_llc mwaitx cpb cat_l3 cdp_l3 hw_pstate ssbd mba perfmon_v2 ibrs ibpb stibp ibrs_enhanced vmmcall fsgsbase bmi1 avx2 smep bmi2 erms invpcid cqm rdt_a avx512f avx512dq rdseed adx smap avx512ifma clflushopt clwb avx512cd sha_ni avx512bw avx512vl xsaveopt xsavec xgetbv1 xsaves cqm_llc cqm_occup_llc cqm_mbm_total cqm_mbm_local user_shstk avx512_bf16 clzero irperf xsaveerptr rdpru wbnoinvd amd_ppin cppc arat npt lbrv svm_lock nrip_save tsc_scale vmcb_clean flushbyasid decodeassists pausefilter pfthreshold avic v_vmsave_vmload vgif x2avic v_spec_ctrl vnmi avx512vbmi umip pku ospke avx512_vbmi2 gfni vaes vpclmulqdq avx512_vnni avx512_bitalg avx512_vpopcntdq la57 rdpid overflow_recov succor smca fsrm flush_l1d debug_swap G++: g++ (Ubuntu 9.4.0-1ubuntu1~20.04) 9.4.0 Copyright (C) 2019 Free Software Foundation, Inc. This is free software; see the source for copying conditions. There is NO warranty; not even for MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. Clang: clang version 10.0.0-4ubuntu1 Target: x86_64-pc-linux-gnu Thread model: posix InstalledDir: /usr/bin |
|
The doc looks good to me. Is this OK with you to merge, @SteveBronder? |
|
Hi @SteveBronder just wanted to check in and make sure this PR looks good and can be merged |
|
@SteveBronder's out of the office and will hopefully be able to get to this by the end of next week or early the following week. Sorry for the delay. I don't see any obstacle to merging this, so I wouldn't worry about it. I just wanted to get Steve's blessing rather than doing it myself. |
|
No problem makes sense, thanks again with your help with this PR! |
Jenkins Console Log Machine informationDistributor ID: Ubuntu Description: Ubuntu 20.04.3 LTS Release: 20.04 Codename: focal CPU: Architecture: x86_64 CPU op-mode(s): 32-bit, 64-bit Byte Order: Little Endian Address sizes: 43 bits physical, 48 bits virtual CPU(s): 256 On-line CPU(s) list: 0-255 Thread(s) per core: 2 Core(s) per socket: 64 Socket(s): 2 NUMA node(s): 2 Vendor ID: AuthenticAMD CPU family: 23 Model: 49 Model name: AMD EPYC 7742 64-Core Processor Stepping: 0 Frequency boost: enabled CPU MHz: 1497.881 CPU max MHz: 3416.0681 CPU min MHz: 1500.0000 BogoMIPS: 4491.56 Virtualization: AMD-V L1d cache: 4 MiB L1i cache: 4 MiB L2 cache: 64 MiB L3 cache: 512 MiB NUMA node0 CPU(s): 0-63,128-191 NUMA node1 CPU(s): 64-127,192-255 Vulnerability Gather data sampling: Not affected Vulnerability Indirect target selection: Not affected Vulnerability Itlb multihit: Not affected Vulnerability L1tf: Not affected Vulnerability Mds: Not affected Vulnerability Meltdown: Not affected Vulnerability Mmio stale data: Not affected Vulnerability Old microcode: Not affected Vulnerability Reg file data sampling: Not affected Vulnerability Retbleed: Mitigation; untrained return thunk; SMT enabled with STIBP protection Vulnerability Spec rstack overflow: Mitigation; Safe RET Vulnerability Spec store bypass: Mitigation; Speculative Store Bypass disabled via prctl Vulnerability Spectre v1: Mitigation; usercopy/swapgs barriers and __user pointer sanitization Vulnerability Spectre v2: Mitigation; Retpolines; IBPB conditional; STIBP always-on; RSB filling; PBRSB-eIBRS Not affected; BHI Not affected Vulnerability Srbds: Not affected Vulnerability Tsa: Not affected Vulnerability Tsx async abort: Not affected Vulnerability Vmscape: Mitigation; IBPB before exit to userspace Flags: fpu vme de pse tsc msr pae mce cx8 apic sep mtrr pge mca cmov pat pse36 clflush mmx fxsr sse sse2 ht syscall nx mmxext fxsr_opt pdpe1gb rdtscp lm constant_tsc rep_good nopl xtopology nonstop_tsc cpuid extd_apicid aperfmperf rapl pni pclmulqdq monitor ssse3 fma cx16 sse4_1 sse4_2 x2apic movbe popcnt aes xsave avx f16c rdrand lahf_lm cmp_legacy svm extapic cr8_legacy abm sse4a misalignsse 3dnowprefetch osvw ibs skinit wdt tce topoext perfctr_core perfctr_nb bpext perfctr_llc mwaitx cpb cat_l3 cdp_l3 hw_pstate ssbd mba ibrs ibpb stibp vmmcall fsgsbase bmi1 avx2 smep bmi2 cqm rdt_a rdseed adx smap clflushopt clwb sha_ni xsaveopt xsavec xgetbv1 xsaves cqm_llc cqm_occup_llc cqm_mbm_total cqm_mbm_local clzero irperf xsaveerptr rdpru wbnoinvd amd_ppin arat npt lbrv svm_lock nrip_save tsc_scale vmcb_clean flushbyasid decodeassists pausefilter pfthreshold avic v_vmsave_vmload vgif v_spec_ctrl umip rdpid overflow_recov succor smca sev sev_es G++: g++ (Ubuntu 9.4.0-1ubuntu1~20.04) 9.4.0 Copyright (C) 2019 Free Software Foundation, Inc. This is free software; see the source for copying conditions. There is NO warranty; not even for MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. Clang: clang version 10.0.0-4ubuntu1 Target: x86_64-pc-linux-gnu Thread model: posix InstalledDir: /usr/bin |
Closes #3201.
Summary
categorical_logitandmultinomial_logitrejected any infinite value inbeta, althoughsoftmaxhandles-inf, socategorical_logit(beta)was stricter thancategorical(softmax(beta)). For databeta, the_rngand_lpmffunctions (and so_lupmf) now accept infinite values, as discussed in the issue:-infentries have zero probability.+infentries split the probability evenly between them. A single+infentry is deterministic.-infinput and NaN throw astd::domain_error.Autodiff
betain the_lpmffunctions must still be finite, since infinite values would break the gradients.When there are
+infentries, the probabilities are computed directly, becausesoftmax/log_sum_expwould give NaN. Otherwise the existing code path is used unchanged.categorical_logit_rng's sampling loop now uses>=instead of>, so a uniform draw of exactly0.0can't land on a zero-probability first entry.Tests
+/-infand to throw on all--inf.+infentries, for both_rngand_lpmf.betawith-infstill throws.Side Effects
None
Release notes
categorical_logitandmultinomial_logitnow accept infinite log odds whenbetais data:
-infentries get zero probability, and +inf entries split the probability evenly among themselves.Checklist
Copyright holder: Rajit Sareen
The copyright holder is typically you or your assignee, such as a university or company. By submitting this pull request, the copyright holder is agreeing to the license the submitted work under the following licenses:
- Code: BSD 3-clause (https://opensource.org/licenses/BSD-3-Clause)
- Documentation: CC-BY 4.0 (https://creativecommons.org/licenses/by/4.0/)
the basic tests are passing
./runTests.py test/unit)make test-headers)make test-math-dependencies)make doxygen)make cpplint)the code is written in idiomatic C++ and changes are documented in the doxygen
the new changes are tested