import unittest import torch from h3_blackwell_runtime.block import gate_segments, modulate_segments class SegmentIndexContracts(unittest.TestCase): def test_segment_index_covers_rows_and_reuses_layout(self): try: from h3_blackwell_runtime.h3_fusion import segment_index except ImportError: self.skipTest("Triton is not installed") segments = [(0, 2, 1), (2, 5, 4), (5, 8, 0)] first = segment_index(8, segments, torch.device("cpu")) second = segment_index(8, segments, torch.device("cpu")) self.assertIs(first, second) torch.testing.assert_close(first, torch.tensor([1, 1, 4, 4, 4, 0, 0, 0], dtype=torch.int32)) def test_segment_index_rejects_gaps(self): try: from h3_blackwell_runtime.h3_fusion import segment_index except ImportError: self.skipTest("Triton is not installed") with self.assertRaisesRegex(ValueError, "ordered, contiguous"): segment_index(4, [(0, 2, 0), (3, 4, 1)], torch.device("cpu")) @unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") class FusedElementwiseParity(unittest.TestCase): @classmethod def setUpClass(cls): try: from h3_blackwell_runtime.h3_fusion import ( fused_gate_add_, fused_modulate_, segment_index, ) except ImportError as error: raise unittest.SkipTest("Triton is not installed") from error cls.fused_gate_add = staticmethod(fused_gate_add_) cls.fused_modulate = staticmethod(fused_modulate_) cls.segment_index = staticmethod(segment_index) def setUp(self): torch.manual_seed(440420) self.segments = [(0, 19, 1), (19, 100, 4), (100, 201, 8), (201, 259, 2)] self.row_index = self.segment_index(259, self.segments, torch.device("cuda")) def test_modulation_is_bit_exact(self): x = torch.randn(259, 1024, device="cuda", dtype=torch.bfloat16) shift = torch.randn(9, 1024, device="cuda", dtype=torch.float32) scale = torch.randn(9, 1024, device="cuda", dtype=torch.float32) expected = modulate_segments(x, shift, scale, self.segments) actual = self.fused_modulate(x.clone(), shift, scale, self.row_index) self.assertTrue(torch.equal(actual, expected)) def test_gate_add_is_bit_exact(self): residual = torch.randn(259, 1024, device="cuda", dtype=torch.bfloat16) update = torch.randn_like(residual) gate = torch.randn(9, 1024, device="cuda", dtype=torch.float32) expected = gate_segments(residual, update, gate, self.segments) actual = self.fused_gate_add(residual.clone(), update, gate, self.row_index) self.assertTrue(torch.equal(actual, expected)) if __name__ == "__main__": unittest.main()