Skip to main content
Back to timeline
arXivSource publication:

Tsinghua team turns the SIS proposal into a learning problem with GFlowNets: one network zero-shot matches or beats the post-hoc best of 31 analytic proposals on 1,190 unseen margins

Synopsis

The work shows that the zero-variance sequential importance sampling (SIS) proposal for binary matrices with fixed margins is exactly the policy of a unit-reward GFlowNet, and proposes MarginFlow, a set transformer that reads the remaining margins and, trained on 1,904 margins, runs zero-shot on 1,190 held-out margins, matching or beating the post-hoc best of 31 analytically designed configurations on 1,187 of them with a median effective sample fraction of 99.8%.

AI-generated editorial illustration: One Proposal for Every Margin: Zero-Shot Amortized Sequential Importance Sampling for Binary Matrices

Interpretation

The paper establishes an equivalence between the zero-variance SIS proposal and the forward policy of a GFlowNet with unit reward on every binary matrix with the given margins: the flow at any state equals its number of completions, the flow at the initial state equals the total count, and selecting children in proportion to flow yields the ideal proposal. Previously SIS proposals were analytically designed in closed form and their accuracy varied substantially with the margins; this work turns proposal construction from analytic design into a learning problem and shows training needs no knowledge of the count. The equivalence is given as Lemma 3.1 and Theorem 3.2 with proofs; Theorem 3.2 further shows the trajectory-balance residual of any proposal equals its log weight up to a constant, and the loss vanishes exactly when weights are constant.

MarginFlow amortizes this learning across margins by exploiting self-similarity: every partial matrix is itself an instance with reduced margins, so one set transformer that reads the remaining margins serves every margin, with no per-instance training, tuning, or selection. Prior learned proposals were typically trained one network per model or conditioned per instance; here row-by-row construction and the symmetry of Proposition 3.6 let one network receive gradients at every state of every margin in the pool. The network is a four-layer set transformer of width 256 (about 3.3M parameters) with no positional encoding, so the symmetry holds by construction; the training pool holds 1,904 margins (1,688 synthetic, 216 real), and three seeds each train for about 25 hours on eight A100s.

On 1,190 held-out margins (synthetic and real, from 3x3 to 870x6), MarginFlow matches or beats the post-hoc best of 31 analytically designed configurations on 1,187, with a median effective sample fraction of 99.8%; on the 56 margins where that best loses more than one nat, it wins every one and lifts the median from 10.3% to 94.1%. Analytic proposals give out on hard margins, and the post-hoc best requires knowing in advance which of the 31 configurations each margin needs, a choice no user can make; the learned proposal holds on these margins. Baselines are 31 configurations (five proposals at six exponents plus uniform), with the post-hoc best selected on 8 repetitions and reported on 16; the effective sample fraction is computed exactly by dynamic programming on 681 margins and estimated from draws with a median over 16 repetitions on the other 509.

Error compounds over rows: a matrix weight is a product over rows, so a fixed per-row error compounds linearly on tall tables, which is why analytic proposals lose fit as rows grow while MarginFlow errs less on every row. This explains why analytic design fails on tall tables while the learned proposal holds, giving a verifiable mechanism rather than only aggregate metrics. Theorem C.1 and Corollary C.2 bound the per-row error; Figure 4 shows the Harrison-Miller proposal falling from 98.7% at 6 rows to 0.6% at 840 rows while MarginFlow stays below 0.1 nat at 840 rows; on the family of Bezakova et al. it extrapolates to five times the rows it trained on.

Perspective

The result is aimed at researchers who need to count or uniformly sample binary matrices with prescribed row and column sums, in settings such as ecological co-occurrence analysis, Rasch-model conditional inference, and social-network motif testing, where the core operation is a weighted average over draws. At deployment the network makes one forward pass per state with no training on new margins, and one checkpoint serves every experiment; a user with a single hard margin can also continue training on it, watching the variance of the log weights fall without ever knowing the count. The paper further notes that the same construction extends to integer-entry contingency tables, graphs with prescribed degrees, and tables with structural zeros, where the network must read more of a column than its reduced sum.

The network reads feasible row types, and the number of types grows with the number of distinct reduced column sums; training is limited to states under one type cap and evaluation under another, and where types are far more numerous the paper's group-by-group exact policy costs about nine times longer per training step, making it the natural next version. On the 509 margins above the exact-computation cap the effective sample fraction is itself estimated from draws, and the paper uses the unbiased count estimate as a cross-check, reporting agreement within Monte Carlo error between MarginFlow and the post-hoc best on 94% of those margins. In addition, this evidence bundle is full text, but figures and tables appear as prose descriptions, so specific numerical details still require the original appendices.

Sources