File size: 8,281 Bytes
b78342b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
"""Focused checks for balanced scheduling and arbitrary-count node splitting."""
import pytest
import torch
from h3_sparse_attention.reblock_hierarchy import build_reblock_hierarchy
from h3_sparse_attention.landmark_v2_terminal import node_split_reference, route_scores, split_topology


def test_seventeen_blocks_use_two_balanced_children():
    h = build_reblock_hierarchy(17*64, (16,), fanout_mode='arbitrary_fanout')
    assert h.budgets(0,17) == (9,8)
    assert h.budgets(1,9) == (1,)*9
    assert h.budgets(1,8) == (1,)*8
    assert len(h.levels) == 3


def test_default_is_strict_fanout16_with_child_reconstruction():
    from h3_sparse_attention import H3SparseAttentionConfig
    for config in (H3SparseAttentionConfig.sol(20), H3SparseAttentionConfig.spark(20)):
        assert config.landmark_tree_v2_children == 16
        assert config.landmark_tree_v2_fanout_mode == 'power_of_two_fanout'
    strict = build_reblock_hierarchy(3 * 64)
    assert strict.fanout == strict.root_fanout == strict.final_fanout == 16
    assert strict.budgets(0, 3) == (2, 1)
    assert strict.budgets(1, 2) == (1, 1)
    hybrid = build_reblock_hierarchy(3 * 64, fanout_mode='power_of_two_arbitrary_final')
    assert hybrid.budgets(0, 3) == (1, 1, 1)
    with pytest.raises(ValueError, match='separate final_fanout'):
        build_reblock_hierarchy(3 * 64, final_fanout=8)


@pytest.mark.parametrize('leaves,fanout,root', [(33,16,(11,11,11)), (17,16,(9,8)), (5,8,(1,)*5)])
def test_arbitrary_fanout_with_matching_final_fanout(leaves, fanout, root):
    h = build_reblock_hierarchy(leaves*64,(fanout,),final_fanout=fanout,fanout_mode='arbitrary_fanout')
    assert h.budgets(0,leaves) == root
    for round_ in h.split_budgets:
        for total, capacities in round_:
            assert sum(capacities) == total
            assert max(capacities)-min(capacities) <= 1
            assert 2 <= len(capacities) <= fanout


def test_ten_second_hierarchy_is_balanced_and_published():
    for fanout, expected_root in [(8,(567,567)), (16,(227,227,227,227,226))]:
        h = build_reblock_hierarchy(72576,(fanout,),final_fanout=16,fanout_mode='arbitrary_fanout')
        assert h.budgets(0,1134) == expected_root
        assert h.levels[-1] == tuple((i,i+1) for i in range(1134))
        for round_ in h.split_budgets:
            for total, capacities in round_:
                assert sum(capacities) == total
                assert max(capacities)-min(capacities) <= 1


def test_named_fanout_modes_control_the_complete_hierarchy():
    power = build_reblock_hierarchy(
        1134*64, (16,), fanout_mode='power_of_two_fanout')
    arbitrary = build_reblock_hierarchy(
        1134*64, (16,), fanout_mode='arbitrary_fanout')
    assert [len(level) for level in power.levels] == [1,16,256,1024,1134]
    assert [len(level) for level in arbitrary.levels] == [1,5,75,1134]
    assert power.budgets(0,1134) == (71,)*7+(70,)+(71,)*7+(70,)
    assert arbitrary.budgets(0,1134) == (227,)*4+(226,)
    assert power.final_fanout == 16
    assert power.metadata()['fanout_mode'] == 'power_of_two_fanout'
    assert arbitrary.metadata()['fanout_mode'] == 'arbitrary_fanout'
    assert arbitrary.final_fanout == 16


@pytest.mark.parametrize('final', [3, 8, 16])
def test_power_of_two_nonfinal_rounds_and_arbitrary_terminal_round(final):
    h = build_reblock_hierarchy(1134*64, fanout=8, root_fanout=16,
                                final_fanout=final, fanout_mode='power_of_two_arbitrary_final')
    for depth, round_ in enumerate(h.split_budgets):
        for leaves, capacities in round_:
            if leaves <= final:
                assert capacities == (1,) * leaves
            else:
                count = len(capacities)
                assert count & (count-1) == 0
                assert count <= (16 if depth == 0 else 8)
    assert h.levels[-1] == tuple((i,i+1) for i in range(1134))


@pytest.mark.parametrize('bad', [None, '', 'power2', 'balanced', True])
def test_named_fanout_modes_are_strict(bad):
    with pytest.raises(ValueError, match='fanout_mode'):
        build_reblock_hierarchy(17*64,(16,),fanout_mode=bad)


@pytest.mark.parametrize('children', [3,5,17,32])
def test_general_route_exact_capacities_and_ties(children):
    caps = tuple(2+i%3 for i in range(children))
    n=sum(caps)
    ids=torch.randperm(n)[None]
    labels=route_scores(torch.zeros(1,n,children-1),ids,caps)
    offset=0
    for child,cap in enumerate(caps):
        assert torch.equal(labels == child, (ids >= offset)&(ids < offset+cap))
        offset+=cap
    assert len(split_topology(caps)) == children-1


def test_cpu_recursive_seventeen_block_split():
    from h3_sparse_attention.landmark_tree_v2 import recursive_landmark_tree_v2_reference
    n=17*64
    out=recursive_landmark_tree_v2_reference(torch.zeros(1,n,8),grid_shape=(1,1,n),max_children=16,fanout_mode='arbitrary_fanout')
    assert out.hierarchy.budgets(0,17) == (9,8)
    assert torch.equal(out.permutation,torch.arange(n)[None])


def test_recursive_modes_select_matching_scheduler_and_splitter():
    from h3_sparse_attention.landmark_tree_v2 import recursive_landmark_tree_v2_reference
    n=17*64; samples=torch.zeros(1,n,8)
    power=recursive_landmark_tree_v2_reference(
        samples,grid_shape=(1,1,n),max_children=16,
        fanout_mode='power_of_two_fanout')
    arbitrary=recursive_landmark_tree_v2_reference(
        samples,grid_shape=(1,1,n),max_children=16,
        fanout_mode='arbitrary_fanout')
    assert [item.children for item in power.split_stats] == [16,2]
    assert [item.children for item in arbitrary.split_stats] == [2,9,8]
    assert power.hierarchy.fanout_mode == 'power_of_two_fanout'
    assert arbitrary.hierarchy.fanout_mode == 'arbitrary_fanout'
    assert torch.equal(power.permutation,torch.arange(n)[None])
    assert torch.equal(arbitrary.permutation,torch.arange(n)[None])


@pytest.mark.skipif(not torch.cuda.is_available(),reason='CUDA required')
def test_cuda_general_scoring_route_and_graph(monkeypatch):
    from h3_sparse_attention.landmark_v2_cosine_fast import build_cosine_directions, fused_cosine_scores
    from h3_sparse_attention.landmark_tree_clustering import _stable_counting_partition
    from h3_sparse_attention.landmark_tree_v2 import PreparedLandmarkTreeV2Permutation
    import h3_sparse_attention.landmark_v2_fused_node as fused
    torch.manual_seed(27)
    caps=(96,80,64)
    samples=torch.randn(1,sum(caps),128,device='cuda',dtype=torch.bfloat16)
    centers=samples[:,:32].contiguous()
    weights=torch.tensor([[8]*16+[7]*16],device='cuda')
    ids=torch.randperm(sum(caps),device='cuda')[None]
    directions=build_cosine_directions(centers,weights,caps)
    scores=fused_cosine_scores(samples,directions,'tf32x3')
    labels=route_scores(scores,ids,caps)
    actual=_stable_counting_partition(labels,caps,validate=True,source_indices=ids)
    expected=node_split_reference(samples,ids,centers,weights,caps)
    assert torch.equal(actual,expected)
    # Exercise large-node dispatch and graph replay without compiling many
    # full-node shapes: full-node fusion is checked separately below.
    monkeypatch.setattr(fused,'FUSED_NODE_ENABLED',False)
    n=17*64
    source=torch.zeros(1,n,128,device='cuda',dtype=torch.bfloat16)
    plan=PreparedLandmarkTreeV2Permutation(batch=1,tokens=n,dim=128,grid_shape=(1,1,n),device='cuda',max_children=16,fanout_mode='arbitrary_fanout')
    plan.run(source)
    perm,inv=plan.run(source)
    assert plan.graph_active
    assert plan.hierarchy.budgets(0,17)==(9,8)
    assert torch.equal(perm,torch.arange(n,device='cuda')[None])
    assert torch.equal(perm.gather(1,inv),perm)


@pytest.mark.skipif(not torch.cuda.is_available(),reason='CUDA required')
def test_existing_fused_split_accepts_unequal_three_way_capacities():
    from h3_sparse_attention.landmark_v2_fused_node import fused_node_split
    caps=(96,80,64);n=sum(caps)
    ids=torch.randperm(n,device='cuda')[None]
    source=torch.zeros(n,128,device='cuda',dtype=torch.bfloat16)
    out=fused_node_split(source,ids,torch.zeros(1,device='cuda',dtype=torch.long),n,caps,midpoint=True,mode='fp16')
    offset=0
    for cap in caps:
        assert torch.equal(out[0,offset:offset+cap],ids[0][(ids[0]>=offset)&(ids[0]<offset+cap)])
        offset+=cap