Organizations: Department of Statistics, University of Washington · Center for Data Science, Zhejiang University · Inria, ´Ecole Normale Sup´erieure, PSL Research University
We consider estimating the conditional distribution of a multivariate outcome given covariates when its coordinates may be continuous, binary, categorical, ordinal or rankings, and are conditionally dependent on one another. Different statistical methods have been developed for each outcome type, and most of them target a summary of the conditional distribution, such as the mean of each coordinate, rather than the joint distribution of the outcome vector. We develop generalized engression models, a unified nonparametric distributional regression framework for outcomes of any type. The proposed method builds upon engression, a scoring-rule-based deep generative model, and introduces a data-type-specific link function and a stochastic perturbation that smooths the loss, enabling gradient-based training even with discontinuous links. We establish universal representation results for continuous, discrete and mixed outcomes. In simulations and in two applications, 242 species in a community ecology benchmark and a 17-dimensional mixed-type health outcome, the method matches type-specific models on marginal scores, improves on them on the joint distribution, and matches or exceeds purpose-built state-of-the-art joint species distribution models. Software is available in Python.
Figures & tables
Figure 1: GEMs on various outcome types, standard type-specific models. (a) Distribution of the eight label patterns and correlation between the labels: logistic regressions have the correct marginals but, drawing the labels independently, spread mass over all patterns and miss the correlation; GEM recovers both. (b) Conditional distribution and mean of the continuous coordinate and the four class probabilities as functions of x , recovered from one model of the joint outcome. (c) Probability of each of the six orderings as a function of x : the Plackett–Luce model cannot match them with linear or with spline scores, whereas GEM does, with a total-variation (TV) error of 0.015 . GEMs are fitted with the same algorithm in all three panels; only the link differs.
data type
Y
h(z)
representation
continuous, k variables
Rk
z
Theorem 1
multiple binary, k labels
{−1,+1}k
sign(z) , coordinate-wise
Theorem 3
categorical, k classes
{e1,…,ek}
ei∗ with i∗=argmaxizi
Theorem 3
ordinal, L levels
{1,…,L}
1+∑l=1L−11{z>l+21}
Theorem 3
ranking of k items
permutations of (1,…,k)
h(z)i=1+#{l:zl>zi}
Theorem 3
mixed
product of the above
concatenation of the above
Theorem 4
Table 1: Links by data type. ei is the i th unit vector, an ordinal level is coded by its value, and a ranking by its rank vector, whose i th entry is the position of item i . The ordinal link rounds z to the nearest level. Ties, which have probability zero, are broken by a fixed rule.
Figure 2: Latent-lift construction for discrete outcomes. Independent Gaussian noise ε is transformed by the continuous generator into Z=g(X,ε) , which the link h maps to Y=h(Z) . For a fixed x , the middle and right panels show the slice a(x,⋅) of the latent density and the joint class subdensities py(x) . Each cell Ay=h−1({y}) has nonempty interior and supports a bump density by . Together with a positive Gaussian background (not shown), these bumps form a latent density satisfying ∫Aya(x,z)dz=py(x) , as indicated by the dashed arrow. The conditional density of Z given X=x and the conditional probabilities of Y given X=x are obtained by dividing a(x,z) and py(x) by pX(x)=∑ypy(x)>0 .
Figure 3: Multi-label prediction: joint energy score, mean per-label CRPS, per-label cross-entropy and co-occurrence error (left to right) against ρ . Lower is better.
Figure 4: Conditional distribution of the number of positive labels at a fixed value x0 of X , for five values of ρ , from 4,000 draws of each method and of the true mechanism.
Figure 5: Mixed-type simulation: joint energy score, continuous CRPS, binary and categorical cross-entropy, and the two cross-type correlation errors, against the sample size n . Lower is better.
Figure 6: Metrics for the ranking outcome regression against sample size, in the independent setting (left panel of each metric) and the coupled setting (right panel). Lower is better except for Kendall’s τ .
Figure 7: Joint species distribution modeling on the vegetation data: the 17 models, each averaged over the five splits; the labels abbreviate the model names of Table 4 in Appendix C . For all metrics, lower is better. (a) The richness interval error, the calibration error of the central 50% interval, separates the ten models that draw the species independently (squares) from the seven with a dependence structure. (b) The richness CRPS against the marginal log score: the models that draw species independently are good on the marginals and poor on richness, the joint baselines better on richness and worse on the marginals, and GEM best on both. Figure 10 in Appendix C gives the split-by-split differences.
richness
composition
occurrence
model
energy ↓
CRPS ↓
int. ↓
CRPS ↓
int. ↓
NLL ↓
AUC ↑
overall ↑
fit time ↓
GEM, 5-seed ensemble
2.782
6.17
0.031
0.091
0.036
56.1
0.738
1.09
5 × 19 s
GEM, single fit
2.795
6.23
0.039
0.091
0.039
57.3
0.727
0.91
19 s
Random forest + coord.
2.792
7.33
0.301
0.105
0.290
56.6
0.734
0.24
18 min
GAM + spatial smooth
2.817
7.42
0.286
0.112
0.309
56.5
0.738
0.05
2 min
GAM
2.837
7.46
0.293
0.113
0.314
57.1
0.726
− 0.32
15 s
Table 2: Joint species distribution modeling on the vegetation data: mean over the five splits; Table 4 in Appendix C gives the standard deviations. Columns: energy score; CRPS and interval error of richness and of composition (the β -Sørensen dissimilarity between two plots); marginal log score (NLL, the negative log-likelihood of the observed occurrences) and AUC of species occurrence; the benchmark’s overall score, an average of nine of its measures, each standardized across the models, corrected as in Appendix C ; time of one fit (GEM on one GPU, baselines on CPU).
Figure 8: Health examination survey data from NHANES: paired per-split differences, baseline minus GEM, for the joint energy score and the per-block marginal metrics over 20 random splits. A positive value favors GEM. Each dot is one split; the line is the mean across splits and the band extends two standard errors on either side.
Appendix figures & tables10 assets
Supplementary material from the paper’s appendix.
Appendix
multi-label
mixed-type
ranking
species
NHANES
links
sign
identity, sign, argmax
ranking
sign
identity, sign, ordinal
d , k
5 , 10
5 , 31
5 , 4
4 , 242
29 , 13
#layers of g
4
4
4
4
4
width of g
512
512
512
256
512
scale σ
common
common
common
common
coordinate-wise
refinement (Section 3.3 )
none
pathwise, g
none
control variate
pathwise, g and σ
Appendix
Table 3: Settings of GEM in the five experiments; the common settings are in the text. d and k are the dimensions of the covariates and of the coded outcome.
Figure 9: Metrics for multiple binary labels at n=100 (top) and n=1,000 (bottom).
model
coord.
energy ↓
rich. CRPS ↓
rich. int. ↓
β CRPS ↓
β int. ↓
NLL ↓
AUC ↑
dep. gain ↑
co- occ. ↑
overall ↑
GEM, 5-seed ensemble
no
2.782
6.17
0.031
0.091
0.036
56.1
0.738
0.0078
0.81
1.09
(0.038)
(0.16)
(0.022)
(0.003)
(0.028)
(1.2)
(0.004)
(0.06)
Random forest + coord.
yes
2.792
7.33
0.301
0.105
0.290
56.6
0.734
− 0.0006
− 0.01
0.24
(0.028)
(0.22)
(0.010)
(0.002)
(0.033)
(1.1)
(0.007)
(0.14)
GEM, single fit
no
2.795
6.23
0.039
0.091
0.039
57.3
0.727
0.0056
0.77
0.91
(0.039)
(0.17)
(0.019)
(0.003)
(0.020)
(1.4)
(0.009)
(0.13)
Appendix
Table 4: Vegetation data: every model of the comparison, mean (standard deviation) over the five splits (the benchmark’s split only for HMSC with spatial latent factors), sorted by the energy score. Columns as in Table 2 except the fit time, with in addition whether the model uses the plot coordinates, the dependence gain and the correlation between predicted and observed residual co-occurrence.
variant
energy ↓
rich. CRPS ↓
rich. int. ↓
NLL ↓
AUC ↑
AUC rare ↑
AUC common ↑
dep. gain ↑
co- occ. ↑
seed SD ↓
fit time ↓
GEM without
2.916
6.89
0.052
64.4
0.652
0.639
0.670
0.0043
0.50
0.0110
3 min
control variate
(0.048)
(0.11)
(0.030)
(2.1)
GEM, single fit
2.795
6.23
0.039
57.3
0.727
0.697
0.752
0.0056
0.77
0.0071
19 s
(0.039)
(0.17)
(0.019)
(1.4)
GEM, 5-seed ensemble
2.782
6.17
0.031
56.1
0.738
0.717
0.755
0.0078
0.81
—
—
(0.038)
(0.16)
(0.022)
(1.2)
Appendix
Table 5: GEM on the vegetation data: the estimator of Proposition 1 without and with the control variate. Mean (standard deviation) over the five splits; AUC rare and AUC common are over the species with training prevalence below 5% and at least 10% . The seed SD is the standard deviation of the energy score across the five seeds within a split, averaged over splits. NLL is the marginal log score of Table 2 .
model
occurrence
richness
composition
acc. ↓
AUC ↑
calib. ↓
prec.
acc. ↓
disc. ↑
calib. ↓
prec.
acc. ↓
disc. ↑
calib. ↓
prec.
GEM, 5-seed ensemble
0.135
0.739
0.032
0.192
12.72
0.461
0.060
11.75
0.191
0.209
0.053
0.183
GEM, single fit
0.135
0.728
0.032
0.188
12.69
0.463
0.059
11.62
0.191
0.206
0.054
0.181
Random forest + coord.
0.136
0.734
0.032
0.193
9.56
0.601
0.263
3.96
0.147
0.315
0.269
0.074
GAM + spatial smooth
0.137
0.739
0.033
0.193
9.66
0.588
0.256
3.98
0.152
0.270
0.277
0.075
GAM
0.140
0.728
0.033
0.200
9.74
0.571
0.254
4.03
0.153
0.258
0.283
0.076
Appendix
Table 6: The benchmark’s twelve measures, exactly as its code computes them, on the vegetation data, averaged over the five splits (the benchmark’s split only for HMSC with spatial latent factors). Accuracy, discrimination, calibration and precision at the three levels of Norberg et al., (2019) ; the composition columns average the three β -diversity indices. Accuracy at the richness and composition levels is the mean absolute error of individual draws and discrimination there is a per-draw rank correlation, so models that draw species independently score well on them; see the text. Precision, a predictive standard deviation, has no arrow because neither direction is better on its own.
Figure 10: Joint species distribution modeling on the vegetation data: paired per-split differences between nine of the fifteen baselines and the GEM ensemble on the same split, signed so that a positive value favors GEM. Each dot is one split; the vertical bar is the mean over splits. HMSC with spatial latent factors contributes one split.
model
AUC ↑
NLL ↓
rare
uncommon
common
rare
uncommon
common
GEM, 5-seed ensemble
0.717
0.767
0.755
0.0948
0.2100
0.4571
GEM, single fit
0.697
0.763
0.752
0.1002
0.2150
0.4608
GAM + spatial smooth
0.719
0.768
0.751
0.0948
0.2101
0.4619
Random forest + coord.
0.705
0.764
0.763
0.0997
0.2126
0.4539
Boosted trees + coord.
0.645
0.742
0.753
0.0980
0.2134
0.4595
Appendix
Table 7: Vegetation data: species-level accuracy by training prevalence, below 5% , 5 to 10% and at least 10% , averaged over the five splits (the benchmark’s split only for HMSC with spatial latent factors). The classes are formed within each split and hold on average 120.6, 43.8 and 77.6 of the 242 species. NLL is the marginal log score of Table 2 .
joint
continuous
binary
ordinal
method
energy score ↓
CRPS ↓
log-loss ↓
RPS ↓
fit time (s) ↓
GEM
2.0880
0.4963
0.4534
0.4980
62.095
(0.0021)
(0.0007)
(0.0008)
(0.0010)
(0.312)
DRF
2.1034
0.5001
0.4549
0.5057
243.8
(0.0023)
(0.0008)
(0.0009)
(0.0010)
(16.2)
Gaussian copula
2.0951
0.4982
0.4567
0.4954
0.4278
Appendix
Table 8: Absolute values behind Figure 8 , as mean (standard error) over the 20 splits: joint energy score, CRPS averaged over the seven continuous coordinates on the standardized scale, Bernoulli log-loss averaged over the five binary labels, ranked probability score of self-rated health, and fit time in seconds (GEM and the stacked marginals on one GPU, DRF on CPU). The standard errors are larger than those of the paired differences because most of the split-to-split variation is shared by all methods (see the text of this appendix). The Gaussian copula is fitted on top of the stacked marginals, so its total cost is that of the stacked marginals plus its own fit, which takes under a second and is the time shown. Bold: the best mean in each of the four metric columns, ties at the printed precision included.
outcome
GEM score
DRF
Copula
Stacked
BMI (kg/m 2 )
3.428
+0.078 ± 0.003 ∗
+0.014 ± 0.002 ∗
+0.008 ± 0.002 ∗
waist (cm)
8.271
+0.189 ± 0.008 ∗
+0.035 ± 0.005 ∗
+0.021 ± 0.004 ∗
systolic BP (mmHg)
8.545
+0.004 ± 0.004
+0.033 ± 0.005 ∗
+0.018 ± 0.005 ∗
diastolic BP (mmHg)
5.968
+0.034 ± 0.004 ∗
+0.022 ± 0.004 ∗
+0.010 ± 0.004 ∗
HbA1c (%)
0.4172
− 0.0105 ± 0.0004 ∗
+0.0016 ± 0.0003 ∗
+0.0009 ± 0.0003 ∗
total cholesterol (mg/dL)
21.61
+0.07 ± 0.02 ∗
+0.08 ± 0.01 ∗
+0.03 ± 0.01 ∗
Appendix
Table 9: Marginal CRPS of the seven continuous outcomes, in original units. The “GEM score” column is the CRPS of GEM averaged over the 20 splits; each baseline column is the paired difference baseline − GEM (mean ± standard error over the splits), in the same units. Positive means GEM has the lower CRPS. ∗ marks ∣Δ∣>2SE .
outcome
GEM score
DRF
Copula
Stacked
diabetes diagnosis
0.3316
+0.0011 ± 0.0007
+0.0027 ± 0.0007 ∗
+0.0026 ± 0.0006 ∗
hypertension diagnosis
0.5403
+0.0033 ± 0.0006 ∗
+0.0019 ± 0.0004 ∗
+0.0018 ± 0.0005 ∗
BP medication
0.4476
+0.0037 ± 0.0006 ∗
+0.0039 ± 0.0005 ∗
+0.0040 ± 0.0005 ∗
high-cholesterol diagnosis
0.5587
+0.0005 ± 0.0005
+0.0041 ± 0.0006 ∗
+0.0042 ± 0.0007 ∗
cholesterol medication
0.3887
− 0.0011 ± 0.0006
+0.0040 ± 0.0008 ∗
+0.0043 ± 0.0007 ∗
Appendix
Table 10: Bernoulli log-loss of the five binary outcomes, from the smoothed marginal probabilities p^j(x)=(n++1)/(M+2) of the M=500 draws. Columns as in Table 9 : a positive entry means GEM has the lower log-loss; ∗ marks ∣Δ∣>2SE .
Engression is a recently proposed and effective framework for conditional distribution learning. Its multi-step Reverse Markov extension further improves generative flexibility by decomposing complex conditional sampling into sequential reverse transitions. Despite their strong empirical performance, rigorous finite-sample statistical guarantees for these methods remain unavailable. In this paper, under deep neural network parameterizations, we establish nonasymptotic convergence bounds for Engression by directly controlling the Energy Distance between the learned and target conditional distributions. For the Reverse Markov framework, we further develop an Energy-Distance-based chain rule that enables a rigorous analysis of error propagation across reverse steps. Our analysis yields corresponding excess-risk bounds that are near-optimal up to logarithmic factors relative to the classical minimax rate over a general Hölder class.
Jiaqi Huang, Gongjun Xu, Ji Zhu
Department of Statistics, University of Michigan · Ann Arbor, MI 48109, U.S.A.
Modern conditional generative models face significant challenges when learning complex covariate dependencies. While sufficient dimension reduction (SDR) provides a principled approach to compress these dependencies, traditional SDR frameworks were not formulated for conditional generation. To bridge this gap, we propose Belted Engression, a unified and architecturally parameter-efficient framework for generative distributional regression. Our approach establishes an end-to-end compress-then-generate paradigm driven by sufficient representation learning, embedding a structural bottleneck into the generative architecture. Theoretically, we prove that the standard SDR condition is equivalent to a law-preserving generative factorization, which is achieved at the global optimum of the population Belted Engression objective. Furthermore, by uncovering a localized Bernstein-type control for the energy-score loss, we establish finite-sample convergence rates that are sharper than those of existing results. We also prove that this belted architecture is strictly smaller, operating with an asymptotically vanishing parameter count relative to the unstructured baseline. Extensive simulations and real-world applications demonstrate that Belted Engression achieves superior distributional prediction and SDR recovery with fewer trainable parameters.
Wenxi Tan, Bing Li, Lingzhou Xue
Department of Statistics, The Pennsylvania State University
Engression (Shen and Meinshausen, 2024) learns a conditional distribution by fitting a generative model Y=f(X,ε) under the energy score, a strictly proper scoring rule. We provide a theoretical error analysis of engression implemented with deep neural networks. We decompose the excess risk into three components: the approximation error, the stochastic error, and the Monte Carlo error. Based on this decomposition, we establish convergence rates under the assumption that the target conditional generator admits a compositional smoothness structure.
Juntong Chen, Zijian Guo, Xinwei Shen
School of Mathematical Sciences, Xiamen University · Center for Data Science, Zhejiang University · Department of Statistics, University of Washington