🤖 AI Summary
This work addresses the severe memory bottleneck of traditional Semi-CRFs when applied to long sequences and large label sets, where explicit instantiation of the edge potential tensor renders scaling to ultra-long sequences—such as those in speech or genomics—infeasible. To overcome this limitation, the authors propose a memory-efficient exact inference framework that computes edge potentials on-the-fly via prefix sums, enabling a streaming forward–backward algorithm with sublinear memory complexity. The approach incorporates a zero-centered, numerically stable cumulative scoring mechanism and an adaptive duration prior. Implemented within a custom Triton fused kernel, the method supports exact inference on sequences exceeding 100,000 time steps, drastically reducing memory consumption while preserving gradient accuracy and thereby breaking the scalability barrier of conventional Semi-CRFs.
📝 Abstract
Semi-Markov Conditional Random Fields (semi-CRFs) assign labels to segments of a sequence rather than to individual positions, enabling exact inference over segment-level features and principled uncertainty estimates at their boundaries. However, existing implementations must materialize a large edge potential tensor whose size grows with sequence length, maximum segment length, and label count, becoming prohibitive for speech-scale state spaces and intractable at genomic scales where sequences can exceed 100,000 positions. This memory bottleneck has limited the adoption of exact segment-level inference for long sequences and large label sets. We identify that the core inefficiency is materializing edge potentials that can instead be evaluated on-the-fly from a compact prefix-sum array, and make several improvements. First, replacing the stored edge tensor with prefix-sum lookup reduces the memory footprint by a factor proportional to the product of segment length and label count. Second, a streaming forward-backward pass with checkpoint-boundary normalization keeps working memory sublinear in sequence length while preserving exact gradients. Third, zero-centered cumulative scores control numerical drift and induce an adaptive duration prior under label imbalance. We integrate these ideas into Flash-SemiCRF, a fused Triton kernel that enables exact semi-CRF inference on previously intractable problem sizes. Available at https://github.com/biobenkj/flash-semicrf.