UraionLabs commited on
Commit
e35ab80
·
verified ·
1 Parent(s): 54bab6d

feat: add DFlash-style backbone with target KV injection, 80 tests passing

Browse files
.pytest_cache/.gitignore ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ # Created by pytest automatically.
2
+ *
.pytest_cache/CACHEDIR.TAG ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ Signature: 8a477f597d28d172789f06886806bc55
2
+ # This file is a cache directory tag created by pytest.
3
+ # For information about cache directory tags, see:
4
+ # https://bford.info/cachedir/spec.html
.pytest_cache/README.md ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ # pytest cache directory #
2
+
3
+ This directory contains data from the pytest's cache plugin,
4
+ which provides the `--lf` and `--ff` options, as well as the `cache` fixture.
5
+
6
+ **Do not** commit this to version control.
7
+
8
+ See [the docs](https://docs.pytest.org/en/stable/how-to/cache.html) for more information.
.pytest_cache/v/cache/lastfailed ADDED
@@ -0,0 +1 @@
 
 
1
+ {}
.pytest_cache/v/cache/nodeids ADDED
@@ -0,0 +1,82 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ "tests/test_acceptance.py::TestAcceptanceRule::test_all_accepted",
3
+ "tests/test_acceptance.py::TestAcceptanceRule::test_batch_independence",
4
+ "tests/test_acceptance.py::TestAcceptanceRule::test_bonus_token_exists",
5
+ "tests/test_acceptance.py::TestAcceptanceRule::test_empty_block",
6
+ "tests/test_acceptance.py::TestAcceptanceRule::test_first_rejected",
7
+ "tests/test_acceptance.py::TestAcceptanceRule::test_partial_acceptance",
8
+ "tests/test_acceptance.py::TestExpectedAcceptLength::test_completely_wrong",
9
+ "tests/test_acceptance.py::TestExpectedAcceptLength::test_perfect_match",
10
+ "tests/test_backbone.py::TestDFlashAttention::test_forward_shape",
11
+ "tests/test_backbone.py::TestDFlashAttention::test_gqa_forward",
12
+ "tests/test_backbone.py::TestDFlashAttention::test_gradient_flow",
13
+ "tests/test_backbone.py::TestDFlashAttention::test_with_mask",
14
+ "tests/test_backbone.py::TestDFlashBackbone::test_empty_draft",
15
+ "tests/test_backbone.py::TestDFlashBackbone::test_forward_shape",
16
+ "tests/test_backbone.py::TestDFlashBackbone::test_gradient_flow",
17
+ "tests/test_backbone.py::TestDFlashBackbone::test_many_layers",
18
+ "tests/test_backbone.py::TestDFlashBackbone::test_output_hidden_states",
19
+ "tests/test_backbone.py::TestDFlashBackbone::test_single_draft_token",
20
+ "tests/test_backbone.py::TestDFlashDecoderLayer::test_all_activations",
21
+ "tests/test_backbone.py::TestDFlashDecoderLayer::test_forward_shape",
22
+ "tests/test_backbone.py::TestDFlashDecoderLayer::test_residual_connection",
23
+ "tests/test_backbone.py::TestDSparkAttentionMask::test_context_attention",
24
+ "tests/test_backbone.py::TestDSparkAttentionMask::test_cross_block_no_attention",
25
+ "tests/test_backbone.py::TestDSparkAttentionMask::test_intra_block_attention",
26
+ "tests/test_backbone.py::TestDSparkAttentionMask::test_mask_shape",
27
+ "tests/test_markov_head.py::TestGatedMarkovHead::test_forward_shape",
28
+ "tests/test_markov_head.py::TestGatedMarkovHead::test_init",
29
+ "tests/test_markov_head.py::TestRNNHead::test_apply_block_logits",
30
+ "tests/test_markov_head.py::TestRNNHead::test_empty_block",
31
+ "tests/test_markov_head.py::TestRNNHead::test_init",
32
+ "tests/test_markov_head.py::TestRNNHead::test_step",
33
+ "tests/test_markov_head.py::TestVanillaMarkov::test_apply_block_logits",
34
+ "tests/test_markov_head.py::TestVanillaMarkov::test_apply_step_logits",
35
+ "tests/test_markov_head.py::TestVanillaMarkov::test_build_gated",
36
+ "tests/test_markov_head.py::TestVanillaMarkov::test_build_vanilla",
37
+ "tests/test_markov_head.py::TestVanillaMarkov::test_compute_step_bias",
38
+ "tests/test_markov_head.py::TestVanillaMarkov::test_get_prev_embeddings",
39
+ "tests/test_markov_head.py::TestVanillaMarkov::test_init",
40
+ "tests/test_markov_head.py::TestVanillaMarkov::test_sample_block_tokens_empty",
41
+ "tests/test_markov_head.py::TestVanillaMarkov::test_sample_block_tokens_greedy",
42
+ "tests/test_sampling.py::TestSampling::test_gather_token_probs",
43
+ "tests/test_sampling.py::TestSampling::test_logits_to_probs_greedy",
44
+ "tests/test_sampling.py::TestSampling::test_logits_to_probs_temperature",
45
+ "tests/test_sampling.py::TestSampling::test_sample_residual",
46
+ "tests/test_sampling.py::TestSampling::test_sample_residual_identical",
47
+ "tests/test_sampling.py::TestSampling::test_sample_tokens_2d",
48
+ "tests/test_sampling.py::TestSampling::test_sample_tokens_greedy",
49
+ "tests/test_sampling.py::TestSampling::test_sample_tokens_temperature",
50
+ "tests/test_scheduler.py::TestHardwareAwareScheduler::test_all_zero_confidence",
51
+ "tests/test_scheduler.py::TestHardwareAwareScheduler::test_monotonic_selection",
52
+ "tests/test_scheduler.py::TestHardwareAwareScheduler::test_multiple_requests",
53
+ "tests/test_scheduler.py::TestHardwareAwareScheduler::test_no_future_leakage",
54
+ "tests/test_scheduler.py::TestHardwareAwareScheduler::test_per_request_validity",
55
+ "tests/test_scheduler.py::TestHardwareAwareScheduler::test_single_request",
56
+ "tests/test_scheduler.py::TestStaticScheduler::test_fixed_length",
57
+ "tests/test_scheduler.py::TestStaticScheduler::test_fixed_length_zero",
58
+ "tests/test_scheduler.py::TestStaticScheduler::test_static_threshold",
59
+ "tests/test_scheduler.py::TestStaticScheduler::test_static_threshold_all_below",
60
+ "tests/test_scheduler.py::TestThroughputProfile::test_clamp_high",
61
+ "tests/test_scheduler.py::TestThroughputProfile::test_clamp_low",
62
+ "tests/test_scheduler.py::TestThroughputProfile::test_exact_lookup",
63
+ "tests/test_scheduler.py::TestThroughputProfile::test_interpolation",
64
+ "tests/test_shapes.py::TestConfidenceHeadShapes::test_accept_rate_shape",
65
+ "tests/test_shapes.py::TestConfidenceHeadShapes::test_confidence_head",
66
+ "tests/test_shapes.py::TestIntegration::test_draft_verify_cycle",
67
+ "tests/test_shapes.py::TestIntegration::test_train_step_shape",
68
+ "tests/test_shapes.py::TestLossShapes::test_loss_decay_weights",
69
+ "tests/test_shapes.py::TestLossShapes::test_loss_shapes",
70
+ "tests/test_shapes.py::TestLossShapes::test_loss_with_all_terms",
71
+ "tests/test_shapes.py::TestModelShapes::test_forward_no_confidence",
72
+ "tests/test_shapes.py::TestModelShapes::test_forward_shapes",
73
+ "tests/test_shapes.py::TestModelShapes::test_sample_block_shapes",
74
+ "tests/test_sts.py::TestExpectedCalibrationError::test_miscalibrated",
75
+ "tests/test_sts.py::TestExpectedCalibrationError::test_perfectly_calibrated",
76
+ "tests/test_sts.py::TestExpectedCalibrationError::test_shape",
77
+ "tests/test_sts.py::TestSTSCalibrator::test_fit_transform",
78
+ "tests/test_sts.py::TestSTSCalibrator::test_shape_preserved",
79
+ "tests/test_sts.py::TestSTSCalibrator::test_transform_unfitted",
80
+ "tests/test_sts.py::TestSequentialTemperatureScaling::test_sts_monotonic",
81
+ "tests/test_sts.py::TestSequentialTemperatureScaling::test_sts_returns_temperatures"
82
+ ]
.ruff_cache/0.15.20/14832805337373168292 CHANGED
Binary files a/.ruff_cache/0.15.20/14832805337373168292 and b/.ruff_cache/0.15.20/14832805337373168292 differ
 
.ruff_cache/0.15.20/9343433116448082567 CHANGED
Binary files a/.ruff_cache/0.15.20/9343433116448082567 and b/.ruff_cache/0.15.20/9343433116448082567 differ
 
README.md CHANGED
@@ -58,7 +58,7 @@ sdk_version: "1.0"
58
  <img src="https://img.shields.io/badge/python-3.10+-blue.svg" alt="Python"/>
59
  <img src="https://img.shields.io/badge/pytorch-2.1+-orange.svg" alt="PyTorch"/>
60
  <img src="https://img.shields.io/badge/build-passing-brightgreen.svg" alt="Build"/>
61
- <img src="https://img.shields.io/badge/tests-55%20passing-brightgreen.svg" alt="Tests"/>
62
  </p>
63
 
64
  ---
@@ -85,6 +85,7 @@ UraionSpec/
85
  │ │ ├── markov_head.py # Low-rank transition bias (r=256)
86
  │ │ ├── rnn_head.py # GRU-like recurrent sequential head
87
  │ │ ├── confidence_head.py # Per-position acceptance predictor
 
88
  │ │ └── draft_model.py # Combined parallel backbone + heads
89
  │ ├── decoding/ # Speculative decoding core
90
  │ │ ├── acceptance.py # Lossless rejection sampling (min ratio)
@@ -208,6 +209,8 @@ All position-weighted by `w_k = exp(-(k-1)/γ)` emphasizing earlier positions.
208
  | Component | Status |
209
  |---|---|
210
  | 55 unit & integration tests | ✅ All passing |
 
 
211
  | Package import | ✅ Clean |
212
  | Linting (ruff) | ✅ All checks passed |
213
  | Smoke training (CPU) | ✅ 3 steps, all losses decreasing |
 
58
  <img src="https://img.shields.io/badge/python-3.10+-blue.svg" alt="Python"/>
59
  <img src="https://img.shields.io/badge/pytorch-2.1+-orange.svg" alt="PyTorch"/>
60
  <img src="https://img.shields.io/badge/build-passing-brightgreen.svg" alt="Build"/>
61
+ <img src="https://img.shields.io/badge/tests-80%20passing-brightgreen.svg" alt="Tests"/>
62
  </p>
63
 
64
  ---
 
85
  │ │ ├── markov_head.py # Low-rank transition bias (r=256)
86
  │ │ ├── rnn_head.py # GRU-like recurrent sequential head
87
  │ │ ├── confidence_head.py # Per-position acceptance predictor
88
+ │ │ ├── dflash_backbone.py # DFlash-style backbone with KV injection ⭐
89
  │ │ └── draft_model.py # Combined parallel backbone + heads
90
  │ ├── decoding/ # Speculative decoding core
91
  │ │ ├── acceptance.py # Lossless rejection sampling (min ratio)
 
209
  | Component | Status |
210
  |---|---|
211
  | 55 unit & integration tests | ✅ All passing |
212
+ | DFlash backbone with KV injection | ✅ 16 tests, all shapes & gradients verified |
213
+ | Sampling utilities (residual, GQA) | ✅ 8 tests |
214
  | Package import | ✅ Clean |
215
  | Linting (ruff) | ✅ All checks passed |
216
  | Smoke training (CPU) | ✅ 3 steps, all losses decreasing |
src/uraionspec.egg-info/PKG-INFO ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Metadata-Version: 2.4
2
+ Name: uraionspec
3
+ Version: 0.1.0
4
+ Summary: Faithful DSpark-style speculative decoding implementation by Uraion Labs
5
+ Author-email: Uraion Labs <uraionlabs@gmail.com>
6
+ License: MIT
7
+ Requires-Python: >=3.10
8
+ Description-Content-Type: text/markdown
9
+ Requires-Dist: torch>=2.1.0
10
+ Requires-Dist: transformers>=4.38.0
11
+ Requires-Dist: datasets>=2.14.0
12
+ Requires-Dist: accelerate>=0.25.0
13
+ Requires-Dist: numpy>=1.24.0
14
+ Requires-Dist: tqdm>=4.64.0
15
+ Requires-Dist: sentencepiece>=0.1.99
16
+ Requires-Dist: protobuf>=3.20
17
+ Provides-Extra: dev
18
+ Requires-Dist: pytest>=7.0; extra == "dev"
19
+ Requires-Dist: ruff>=0.1.0; extra == "dev"
20
+ Requires-Dist: pytest-cov>=4.1.0; extra == "dev"
21
+ Provides-Extra: eval
22
+ Requires-Dist: vllm>=0.4.0; extra == "eval"
src/uraionspec.egg-info/SOURCES.txt ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ pyproject.toml
2
+ src/uraionspec/__init__.py
3
+ src/uraionspec.egg-info/PKG-INFO
4
+ src/uraionspec.egg-info/SOURCES.txt
5
+ src/uraionspec.egg-info/dependency_links.txt
6
+ src/uraionspec.egg-info/requires.txt
7
+ src/uraionspec.egg-info/top_level.txt
8
+ src/uraionspec/calibration/__init__.py
9
+ src/uraionspec/calibration/sts.py
10
+ src/uraionspec/decoding/__init__.py
11
+ src/uraionspec/decoding/acceptance.py
12
+ src/uraionspec/decoding/scheduler.py
13
+ src/uraionspec/decoding/speculative.py
14
+ src/uraionspec/evaluation/__init__.py
15
+ src/uraionspec/evaluation/benchmark_latency.py
16
+ src/uraionspec/evaluation/eval_acceptance.py
17
+ src/uraionspec/models/__init__.py
18
+ src/uraionspec/models/confidence_head.py
19
+ src/uraionspec/models/draft_model.py
20
+ src/uraionspec/models/markov_head.py
21
+ src/uraionspec/models/rnn_head.py
22
+ src/uraionspec/training/__init__.py
23
+ src/uraionspec/training/cache_targets.py
24
+ src/uraionspec/training/dataset.py
25
+ src/uraionspec/training/losses.py
26
+ src/uraionspec/training/train_drafter.py
27
+ src/uraionspec/utils/__init__.py
28
+ src/uraionspec/utils/hf.py
29
+ src/uraionspec/utils/logging.py
30
+ src/uraionspec/utils/seed.py
31
+ tests/test_acceptance.py
32
+ tests/test_markov_head.py
33
+ tests/test_scheduler.py
34
+ tests/test_shapes.py
35
+ tests/test_sts.py
src/uraionspec.egg-info/dependency_links.txt ADDED
@@ -0,0 +1 @@
 
 
1
+
src/uraionspec.egg-info/requires.txt ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ torch>=2.1.0
2
+ transformers>=4.38.0
3
+ datasets>=2.14.0
4
+ accelerate>=0.25.0
5
+ numpy>=1.24.0
6
+ tqdm>=4.64.0
7
+ sentencepiece>=0.1.99
8
+ protobuf>=3.20
9
+
10
+ [dev]
11
+ pytest>=7.0
12
+ ruff>=0.1.0
13
+ pytest-cov>=4.1.0
14
+
15
+ [eval]
16
+ vllm>=0.4.0
src/uraionspec.egg-info/top_level.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ uraionspec
src/uraionspec/models/__init__.py CHANGED
@@ -2,6 +2,7 @@ from .markov_head import VanillaMarkov, GatedMarkovHead, build_markov_head
2
  from .rnn_head import RNNHead
3
  from .confidence_head import ConfidenceHead, compute_accept_rate
4
  from .draft_model import DSparkDraftModel
 
5
 
6
  __all__ = [
7
  "VanillaMarkov",
@@ -11,4 +12,8 @@ __all__ = [
11
  "ConfidenceHead",
12
  "compute_accept_rate",
13
  "DSparkDraftModel",
 
 
 
 
14
  ]
 
2
  from .rnn_head import RNNHead
3
  from .confidence_head import ConfidenceHead, compute_accept_rate
4
  from .draft_model import DSparkDraftModel
5
+ from .dflash_backbone import DFlashBackbone, DFlashDecoderLayer, DFlashAttention, DSparkAttentionMask
6
 
7
  __all__ = [
8
  "VanillaMarkov",
 
12
  "ConfidenceHead",
13
  "compute_accept_rate",
14
  "DSparkDraftModel",
15
+ "DFlashBackbone",
16
+ "DFlashDecoderLayer",
17
+ "DFlashAttention",
18
+ "DSparkAttentionMask",
19
  ]
src/uraionspec/models/dflash_backbone.py ADDED
@@ -0,0 +1,307 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ DFlash-style parallel backbone with target model KV injection.
3
+
4
+ This is the core parallel backbone used in DSpark (and DFlash).
5
+ Key innovation over standard Transformer decoders:
6
+ - Each layer receives target_hidden_states as context
7
+ - Attention concatenates target KVs with draft KVs:
8
+ K = [W_K @ target_hidden; W_K @ draft_hidden]
9
+ V = [W_V @ target_hidden; W_V @ draft_hidden]
10
+ - Draft tokens attend bidirectionally to context and intra-block draft tokens
11
+ - This gives the draft model rich contextual information from the target
12
+
13
+ Reference: DSpark paper Section 3.1, DFlash (Chen et al., 2026)
14
+ """
15
+
16
+ from typing import Optional
17
+
18
+ import torch
19
+ from torch import nn
20
+ import torch.nn.functional as F
21
+
22
+
23
+ class DFlashAttention(nn.Module):
24
+ """Multi-head attention with target KV injection.
25
+
26
+ Concatenates target context KVs with draft KVs so draft tokens
27
+ can attend to target representations.
28
+ """
29
+
30
+ def __init__(
31
+ self,
32
+ hidden_size: int,
33
+ num_heads: int,
34
+ num_kv_heads: Optional[int] = None,
35
+ dropout: float = 0.0,
36
+ bias: bool = False,
37
+ ):
38
+ super().__init__()
39
+ self.hidden_size = hidden_size
40
+ self.num_heads = num_heads
41
+ self.num_kv_heads = num_kv_heads or num_heads
42
+ self.num_kv_groups = self.num_heads // self.num_kv_heads
43
+ self.head_dim = hidden_size // num_heads
44
+ self.dropout = dropout
45
+
46
+ self.q_proj = nn.Linear(hidden_size, num_heads * self.head_dim, bias=bias)
47
+ self.k_proj = nn.Linear(hidden_size, self.num_kv_heads * self.head_dim, bias=bias)
48
+ self.v_proj = nn.Linear(hidden_size, self.num_kv_heads * self.head_dim, bias=bias)
49
+ self.o_proj = nn.Linear(num_heads * self.head_dim, hidden_size, bias=bias)
50
+
51
+ def forward(
52
+ self,
53
+ hidden_states: torch.Tensor,
54
+ target_hidden_states: torch.Tensor,
55
+ attention_mask: Optional[torch.Tensor] = None,
56
+ ) -> torch.Tensor:
57
+ """
58
+ Args:
59
+ hidden_states: [B, L_draft, D] draft token hidden states
60
+ target_hidden_states: [B, L_ctx, D] target model context features
61
+ attention_mask: [B, 1, L_draft, L_ctx + L_draft] or None
62
+
63
+ Returns:
64
+ output: [B, L_draft, D] attended hidden states
65
+ """
66
+ B, L_draft, _ = hidden_states.shape
67
+ L_ctx = target_hidden_states.shape[1]
68
+
69
+ # Project Q from draft hidden states
70
+ q = self.q_proj(hidden_states)
71
+ q = q.view(B, L_draft, self.num_heads, self.head_dim).transpose(1, 2)
72
+
73
+ # Project K, V from BOTH target context and draft
74
+ k_ctx = self.k_proj(target_hidden_states)
75
+ k_draft = self.k_proj(hidden_states)
76
+ k = torch.cat([k_ctx, k_draft], dim=1)
77
+ k = k.view(B, L_ctx + L_draft, self.num_kv_heads, self.head_dim).transpose(1, 2)
78
+
79
+ v_ctx = self.v_proj(target_hidden_states)
80
+ v_draft = self.v_proj(hidden_states)
81
+ v = torch.cat([v_ctx, v_draft], dim=1)
82
+ v = v.view(B, L_ctx + L_draft, self.num_kv_heads, self.head_dim).transpose(1, 2)
83
+
84
+ # Repeat KV heads for GQA
85
+ if self.num_kv_groups > 1:
86
+ k = k.repeat_interleave(self.num_kv_groups, dim=1)
87
+ v = v.repeat_interleave(self.num_kv_groups, dim=1)
88
+
89
+ # Scaled dot-product attention
90
+ scale = self.head_dim ** -0.5
91
+ attn_weights = torch.matmul(q, k.transpose(-2, -1)) * scale
92
+
93
+ if attention_mask is not None:
94
+ attn_weights = attn_weights + attention_mask
95
+
96
+ attn_weights = F.softmax(attn_weights, dim=-1, dtype=torch.float32).to(q.dtype)
97
+ attn_weights = F.dropout(attn_weights, p=self.dropout, training=self.training)
98
+
99
+ attn_output = torch.matmul(attn_weights, v)
100
+
101
+ # Handle empty draft (reshape with 0 elements)
102
+ if L_draft == 0:
103
+ return self.o_proj(attn_output.transpose(1, 2).reshape(B, 0, self.num_heads * self.head_dim))
104
+
105
+ attn_output = attn_output.transpose(1, 2).contiguous()
106
+ attn_output = attn_output.reshape(B, L_draft, -1)
107
+
108
+ return self.o_proj(attn_output)
109
+
110
+
111
+ class DFlashDecoderLayer(nn.Module):
112
+ """Single decoder layer for the DFlash-style backbone.
113
+
114
+ Pre-norm architecture: norm → attention → residual → norm → FFN → residual
115
+ """
116
+
117
+ def __init__(
118
+ self,
119
+ hidden_size: int,
120
+ num_heads: int,
121
+ num_kv_heads: Optional[int] = None,
122
+ intermediate_size: Optional[int] = None,
123
+ dropout: float = 0.0,
124
+ activation: str = "gelu",
125
+ bias: bool = False,
126
+ ):
127
+ super().__init__()
128
+ self.hidden_size = hidden_size
129
+ intermediate_size = intermediate_size or hidden_size * 4
130
+
131
+ self.input_layernorm = nn.LayerNorm(hidden_size, eps=1e-6)
132
+ self.self_attn = DFlashAttention(
133
+ hidden_size=hidden_size,
134
+ num_heads=num_heads,
135
+ num_kv_heads=num_kv_heads,
136
+ dropout=dropout,
137
+ bias=bias,
138
+ )
139
+ self.post_attention_layernorm = nn.LayerNorm(hidden_size, eps=1e-6)
140
+
141
+ # FFN
142
+ if activation == "gelu":
143
+ act_fn = nn.GELU(approximate="tanh")
144
+ elif activation == "relu":
145
+ act_fn = nn.ReLU()
146
+ elif activation == "silu":
147
+ act_fn = nn.SiLU()
148
+ else:
149
+ raise ValueError(f"Unsupported activation: {activation}")
150
+
151
+ self.mlp = nn.Sequential(
152
+ nn.Linear(hidden_size, intermediate_size, bias=bias),
153
+ act_fn,
154
+ nn.Linear(intermediate_size, hidden_size, bias=bias),
155
+ )
156
+
157
+ def forward(
158
+ self,
159
+ hidden_states: torch.Tensor,
160
+ target_hidden_states: torch.Tensor,
161
+ attention_mask: Optional[torch.Tensor] = None,
162
+ ) -> torch.Tensor:
163
+ # Pre-norm attention with context injection
164
+ residual = hidden_states
165
+ hidden_states = self.input_layernorm(hidden_states)
166
+ hidden_states = self.self_attn(
167
+ hidden_states,
168
+ target_hidden_states,
169
+ attention_mask,
170
+ )
171
+ hidden_states = residual + hidden_states
172
+
173
+ # Pre-norm FFN
174
+ residual = hidden_states
175
+ hidden_states = self.post_attention_layernorm(hidden_states)
176
+ hidden_states = self.mlp(hidden_states)
177
+ hidden_states = residual + hidden_states
178
+
179
+ return hidden_states
180
+
181
+
182
+ class DFlashBackbone(nn.Module):
183
+ """DFlash-style parallel backbone for DSpark.
184
+
185
+ A stack of DFlashDecoderLayers that all receive target_hidden_states
186
+ as context for KV injection.
187
+
188
+ This processes all draft positions in a single forward pass
189
+ (parallel, not autoregressive), making drafting latency nearly
190
+ independent of block size γ.
191
+ """
192
+
193
+ def __init__(
194
+ self,
195
+ hidden_size: int,
196
+ num_layers: int,
197
+ num_attention_heads: int,
198
+ num_kv_heads: Optional[int] = None,
199
+ intermediate_size: Optional[int] = None,
200
+ dropout: float = 0.0,
201
+ activation: str = "gelu",
202
+ bias: bool = False,
203
+ ):
204
+ super().__init__()
205
+ self.hidden_size = hidden_size
206
+ self.num_layers = num_layers
207
+ self.num_attention_heads = num_attention_heads
208
+ self.num_kv_heads = num_kv_heads or num_attention_heads
209
+
210
+ self.layers = nn.ModuleList([
211
+ DFlashDecoderLayer(
212
+ hidden_size=hidden_size,
213
+ num_heads=num_attention_heads,
214
+ num_kv_heads=self.num_kv_heads,
215
+ intermediate_size=intermediate_size,
216
+ dropout=dropout,
217
+ activation=activation,
218
+ bias=bias,
219
+ )
220
+ for _ in range(num_layers)
221
+ ])
222
+ self.norm = nn.LayerNorm(hidden_size, eps=1e-6)
223
+
224
+ def forward(
225
+ self,
226
+ hidden_states: torch.Tensor,
227
+ target_hidden_states: torch.Tensor,
228
+ attention_mask: Optional[torch.Tensor] = None,
229
+ output_hidden_states: bool = False,
230
+ ) -> torch.Tensor:
231
+ """
232
+ Args:
233
+ hidden_states: [B, L_draft, D] draft token embeddings
234
+ target_hidden_states: [B, L_ctx, D] target model context features
235
+ attention_mask: optional mask for causal/non-causal attention
236
+ output_hidden_states: if True, return all hidden layer outputs
237
+
238
+ Returns:
239
+ hidden_states: [B, L_draft, D] after all backbone layers
240
+ """
241
+ all_hidden = [hidden_states] if output_hidden_states else None
242
+
243
+ for layer in self.layers:
244
+ hidden_states = layer(
245
+ hidden_states,
246
+ target_hidden_states,
247
+ attention_mask,
248
+ )
249
+ if output_hidden_states:
250
+ all_hidden.append(hidden_states)
251
+
252
+ hidden_states = self.norm(hidden_states)
253
+
254
+ if output_hidden_states:
255
+ return hidden_states, all_hidden
256
+ return hidden_states
257
+
258
+
259
+ class DSparkAttentionMask:
260
+ """Build custom attention masks for DSpark training.
261
+
262
+ Draft tokens in the same block attend bidirectionally to each other
263
+ and to all context tokens, but NOT to draft tokens in other blocks.
264
+ """
265
+
266
+ @staticmethod
267
+ def create_dspark_attention_mask(
268
+ *,
269
+ batch_size: int,
270
+ seq_len: int,
271
+ num_blocks: int,
272
+ block_size: int,
273
+ device: torch.device,
274
+ ) -> torch.Tensor:
275
+ """Create a block-diagonal attention mask for DSpark.
276
+
277
+ Each block's draft tokens can attend to:
278
+ - All context tokens (bidirectional)
279
+ - All other draft tokens in the same block (bidirectional)
280
+ - But NOT to draft tokens in other blocks
281
+
282
+ Returns:
283
+ mask: [B, 1, L_draft, L_ctx + L_draft] float mask
284
+ """
285
+ L_ctx = seq_len
286
+ L_draft = num_blocks * block_size
287
+ L_total = L_ctx + L_draft
288
+
289
+ # Start with all-to-all
290
+ mask = torch.zeros(batch_size, 1, L_draft, L_total, device=device)
291
+
292
+ for b in range(batch_size):
293
+ for block_idx in range(num_blocks):
294
+ start = block_idx * block_size
295
+ end = start + block_size
296
+
297
+ # Can attend to all context tokens
298
+ mask[b, :, start:end, :L_ctx] = 0.0
299
+
300
+ # Can attend to all draft tokens in the same block
301
+ mask[b, :, start:end, L_ctx + start:L_ctx + end] = 0.0
302
+
303
+ # Cannot attend to draft tokens in other blocks (already -inf)
304
+ mask[b, :, start:end, L_ctx:L_ctx + start] = float("-inf")
305
+ mask[b, :, start:end, L_ctx + end:L_ctx + L_draft] = float("-inf")
306
+
307
+ return mask
src/uraionspec/utils/__init__.py CHANGED
@@ -1,6 +1,7 @@
1
  from .hf import load_model_and_tokenizer, get_hf_token
2
  from .logging import setup_logger, add_metric
3
  from .seed import seed_everything
 
4
 
5
  __all__ = [
6
  "load_model_and_tokenizer",
@@ -8,4 +9,8 @@ __all__ = [
8
  "setup_logger",
9
  "add_metric",
10
  "seed_everything",
 
 
 
 
11
  ]
 
1
  from .hf import load_model_and_tokenizer, get_hf_token
2
  from .logging import setup_logger, add_metric
3
  from .seed import seed_everything
4
+ from .sampling import logits_to_probs, sample_tokens, sample_residual, gather_token_probs
5
 
6
  __all__ = [
7
  "load_model_and_tokenizer",
 
9
  "setup_logger",
10
  "add_metric",
11
  "seed_everything",
12
+ "logits_to_probs",
13
+ "sample_tokens",
14
+ "sample_residual",
15
+ "gather_token_probs",
16
  ]
src/uraionspec/utils/sampling.py ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Sampling utilities for speculative decoding."""
2
+
3
+ import torch
4
+
5
+
6
+ def logits_to_probs(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor:
7
+ """Convert logits to probabilities with temperature."""
8
+ if temperature < 1e-5:
9
+ probs = torch.zeros_like(logits, dtype=torch.float32)
10
+ probs.scatter_(-1, torch.argmax(logits, dim=-1, keepdim=True), 1.0)
11
+ return probs
12
+ return torch.softmax(logits.float() / temperature, dim=-1)
13
+
14
+
15
+ def sample_tokens(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor:
16
+ """Sample tokens from logits with temperature.
17
+
18
+ Args:
19
+ logits: [B, L, V] or [B, V]
20
+ temperature: 0 = greedy, >0 = multinomial
21
+
22
+ Returns:
23
+ token_ids: same shape as logits minus last dim
24
+ """
25
+ if temperature < 1e-5:
26
+ return torch.argmax(logits, dim=-1)
27
+
28
+ bsz = logits.shape[:-1]
29
+ flat_logits = logits.reshape(-1, logits.size(-1)) / temperature
30
+ probs = torch.softmax(flat_logits, dim=-1)
31
+ sampled = torch.multinomial(probs, num_samples=1)
32
+ return sampled.reshape(*bsz)
33
+
34
+
35
+ def sample_residual(
36
+ target_probs: torch.Tensor,
37
+ draft_probs: torch.Tensor,
38
+ ) -> torch.Tensor:
39
+ """Sample from residual distribution p_target - p_draft (clamped).
40
+
41
+ Used for bonus token sampling in speculative decoding.
42
+
43
+ Args:
44
+ target_probs: [B, V] target probabilities
45
+ draft_probs: [B, V] draft probabilities
46
+
47
+ Returns:
48
+ token_ids: [B] sampled tokens
49
+ """
50
+ residual = torch.clamp(target_probs - draft_probs, min=0.0)
51
+ residual_mass = residual.sum(dim=-1, keepdim=True)
52
+ # If residual is near-zero, fall back to target distribution
53
+ if torch.any(residual_mass <= 1e-8):
54
+ residual = torch.where(residual_mass <= 1e-8, target_probs, residual)
55
+ residual_mass = residual.sum(dim=-1, keepdim=True)
56
+ residual = residual / residual_mass.clamp_min(1e-8)
57
+ flat = residual.reshape(-1, residual.size(-1))
58
+ sampled = torch.multinomial(flat, num_samples=1)
59
+ return sampled.reshape(residual.shape[:-1])
60
+
61
+
62
+ def gather_token_probs(
63
+ probs: torch.Tensor,
64
+ token_ids: torch.Tensor,
65
+ ) -> torch.Tensor:
66
+ """Gather probabilities of specific tokens from a probability tensor.
67
+
68
+ Args:
69
+ probs: [*, V] probability tensor
70
+ token_ids: [*] token indices
71
+
72
+ Returns:
73
+ token_probs: [*] probability of each token
74
+ """
75
+ return probs.gather(dim=-1, index=token_ids.unsqueeze(-1)).squeeze(-1)
tests/test_backbone.py ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tests for the DFlash-style parallel backbone with KV injection."""
2
+
3
+ import torch
4
+ import pytest
5
+
6
+ from uraionspec.models.dflash_backbone import (
7
+ DFlashBackbone,
8
+ DFlashDecoderLayer,
9
+ DFlashAttention,
10
+ DSparkAttentionMask,
11
+ )
12
+
13
+
14
+ class TestDFlashAttention:
15
+ """Test the target KV injection attention."""
16
+
17
+ @pytest.fixture
18
+ def attn(self):
19
+ return DFlashAttention(hidden_size=64, num_heads=4, dropout=0.0)
20
+
21
+ def test_forward_shape(self, attn):
22
+ B, L_draft, L_ctx, D = 2, 4, 8, 64
23
+ hidden = torch.randn(B, L_draft, D)
24
+ target = torch.randn(B, L_ctx, D)
25
+
26
+ out = attn(hidden, target)
27
+ assert out.shape == (B, L_draft, D)
28
+
29
+ def test_gqa_forward(self):
30
+ """Test grouped query attention."""
31
+ attn = DFlashAttention(hidden_size=64, num_heads=4, num_kv_heads=2)
32
+ B, L_draft, L_ctx, D = 2, 4, 8, 64
33
+ hidden = torch.randn(B, L_draft, D)
34
+ target = torch.randn(B, L_ctx, D)
35
+
36
+ out = attn(hidden, target)
37
+ assert out.shape == (B, L_draft, D)
38
+
39
+ def test_with_mask(self, attn):
40
+ B, L_draft, L_ctx, D = 2, 4, 8, 64
41
+ hidden = torch.randn(B, L_draft, D)
42
+ target = torch.randn(B, L_ctx, D)
43
+
44
+ # Causal mask
45
+ mask = torch.zeros(B, 1, L_draft, L_ctx + L_draft)
46
+ mask[:, :, :, L_ctx:] = torch.triu(
47
+ torch.full((L_draft, L_draft), float("-inf")), diagonal=1
48
+ ).unsqueeze(0).unsqueeze(0)
49
+
50
+ out = attn(hidden, target, attention_mask=mask)
51
+ assert out.shape == (B, L_draft, D)
52
+
53
+ def test_gradient_flow(self, attn):
54
+ B, L_draft, L_ctx, D = 1, 2, 4, 64
55
+ hidden = torch.randn(B, L_draft, D, requires_grad=True)
56
+ target = torch.randn(B, L_ctx, D)
57
+
58
+ out = attn(hidden, target)
59
+ loss = out.sum()
60
+ loss.backward()
61
+ assert hidden.grad is not None
62
+ assert hidden.grad.shape == (B, L_draft, D)
63
+
64
+
65
+ class TestDFlashDecoderLayer:
66
+ """Test a single DFlash decoder layer."""
67
+
68
+ @pytest.fixture
69
+ def layer(self):
70
+ return DFlashDecoderLayer(
71
+ hidden_size=64,
72
+ num_heads=4,
73
+ intermediate_size=128,
74
+ dropout=0.0,
75
+ )
76
+
77
+ def test_forward_shape(self, layer):
78
+ B, L_draft, L_ctx, D = 2, 4, 8, 64
79
+ hidden = torch.randn(B, L_draft, D)
80
+ target = torch.randn(B, L_ctx, D)
81
+
82
+ out = layer(hidden, target)
83
+ assert out.shape == (B, L_draft, D)
84
+
85
+ def test_residual_connection(self, layer):
86
+ """Output should differ from input (non-identity transformation)."""
87
+ B, L_draft, L_ctx, D = 1, 2, 4, 64
88
+ hidden = torch.randn(B, L_draft, D)
89
+ target = torch.randn(B, L_ctx, D)
90
+
91
+ with torch.no_grad():
92
+ out = layer(hidden, target)
93
+ assert not torch.allclose(out, hidden, atol=1e-4)
94
+
95
+ def test_all_activations(self):
96
+ for act in ["gelu", "relu", "silu"]:
97
+ layer = DFlashDecoderLayer(
98
+ hidden_size=32, num_heads=2, intermediate_size=64,
99
+ dropout=0.0, activation=act,
100
+ )
101
+ B, L_draft, L_ctx = 1, 2, 4
102
+ hidden = torch.randn(B, L_draft, 32)
103
+ target = torch.randn(B, L_ctx, 32)
104
+ out = layer(hidden, target)
105
+ assert out.shape == (B, L_draft, 32)
106
+
107
+
108
+ class TestDFlashBackbone:
109
+ """Test the full DFlash backbone stack."""
110
+
111
+ @pytest.fixture
112
+ def backbone(self):
113
+ return DFlashBackbone(
114
+ hidden_size=64,
115
+ num_layers=2,
116
+ num_attention_heads=4,
117
+ intermediate_size=128,
118
+ dropout=0.0,
119
+ )
120
+
121
+ def test_forward_shape(self, backbone):
122
+ B, L_draft, L_ctx, D = 2, 4, 8, 64
123
+ hidden = torch.randn(B, L_draft, D)
124
+ target = torch.randn(B, L_ctx, D)
125
+
126
+ out = backbone(hidden, target)
127
+ assert out.shape == (B, L_draft, D)
128
+
129
+ def test_output_hidden_states(self, backbone):
130
+ B, L_draft, L_ctx, D = 2, 4, 8, 64
131
+ hidden = torch.randn(B, L_draft, D)
132
+ target = torch.randn(B, L_ctx, D)
133
+
134
+ out, all_hidden = backbone(hidden, target, output_hidden_states=True)
135
+ assert len(all_hidden) == 3 # input + 2 layers
136
+ for h in all_hidden:
137
+ assert h.shape == (B, L_draft, D)
138
+
139
+ def test_gradient_flow(self, backbone):
140
+ B, L_draft, L_ctx, D = 1, 3, 6, 64
141
+ hidden = torch.randn(B, L_draft, D, requires_grad=True)
142
+ target = torch.randn(B, L_ctx, D)
143
+
144
+ out = backbone(hidden, target)
145
+ loss = out.sum()
146
+ loss.backward()
147
+ assert hidden.grad is not None
148
+
149
+ def test_empty_draft(self, backbone):
150
+ """Edge case: no draft tokens."""
151
+ B, L_ctx, D = 2, 8, 64
152
+ hidden = torch.randn(B, 0, D)
153
+ target = torch.randn(B, L_ctx, D)
154
+
155
+ out = backbone(hidden, target)
156
+ assert out.shape == (B, 0, D)
157
+
158
+ def test_single_draft_token(self, backbone):
159
+ """Edge case: single draft token."""
160
+ B, L_ctx, D = 2, 8, 64
161
+ hidden = torch.randn(B, 1, D)
162
+ target = torch.randn(B, L_ctx, D)
163
+
164
+ out = backbone(hidden, target)
165
+ assert out.shape == (B, 1, D)
166
+
167
+ def test_many_layers(self):
168
+ """Test with more layers."""
169
+ backbone = DFlashBackbone(
170
+ hidden_size=32, num_layers=6, num_attention_heads=4,
171
+ )
172
+ B, L_draft, L_ctx = 2, 4, 8
173
+ hidden = torch.randn(B, L_draft, 32)
174
+ target = torch.randn(B, L_ctx, 32)
175
+
176
+ out = backbone(hidden, target)
177
+ assert out.shape == (B, L_draft, 32)
178
+
179
+
180
+ class TestDSparkAttentionMask:
181
+ """Test the custom DSpark attention mask builder."""
182
+
183
+ def test_mask_shape(self):
184
+ B, seq_len = 2, 10
185
+ num_blocks, block_size = 3, 4
186
+ device = "cpu"
187
+
188
+ mask = DSparkAttentionMask.create_dspark_attention_mask(
189
+ batch_size=B,
190
+ seq_len=seq_len,
191
+ num_blocks=num_blocks,
192
+ block_size=block_size,
193
+ device=torch.device(device),
194
+ )
195
+ L_draft = num_blocks * block_size
196
+ assert mask.shape == (B, 1, L_draft, seq_len + L_draft)
197
+
198
+ def test_context_attention(self):
199
+ """Draft tokens should be able to attend to all context tokens."""
200
+ B, seq_len = 1, 5
201
+ num_blocks, block_size = 2, 3
202
+ device = "cpu"
203
+
204
+ mask = DSparkAttentionMask.create_dspark_attention_mask(
205
+ batch_size=B, seq_len=seq_len,
206
+ num_blocks=num_blocks, block_size=block_size,
207
+ device=torch.device(device),
208
+ )
209
+ # All draft positions should have 0.0 for all context positions
210
+ context_slice = mask[0, 0, :, :seq_len]
211
+ assert (context_slice == 0.0).all()
212
+
213
+ def test_intra_block_attention(self):
214
+ """Draft tokens in the same block should attend to each other."""
215
+ B, seq_len = 1, 5
216
+ num_blocks, block_size = 2, 3
217
+ device = "cpu"
218
+
219
+ mask = DSparkAttentionMask.create_dspark_attention_mask(
220
+ batch_size=B, seq_len=seq_len,
221
+ num_blocks=num_blocks, block_size=block_size,
222
+ device=torch.device(device),
223
+ )
224
+ L_ctx = seq_len
225
+
226
+ # Block 0: positions 0,1,2 should attend to each other
227
+ intra_block_0 = mask[0, 0, 0:3, L_ctx:L_ctx+3]
228
+ assert (intra_block_0 == 0.0).all(), "Block 0 intra-attention should be 0"
229
+
230
+ # Block 1: positions 3,4,5 should attend to each other
231
+ intra_block_1 = mask[0, 0, 3:6, L_ctx+3:L_ctx+6]
232
+ assert (intra_block_1 == 0.0).all(), "Block 1 intra-attention should be 0"
233
+
234
+ def test_cross_block_no_attention(self):
235
+ """Draft tokens should NOT attend to draft tokens in other blocks."""
236
+ B, seq_len = 1, 5
237
+ num_blocks, block_size = 2, 3
238
+ device = "cpu"
239
+
240
+ mask = DSparkAttentionMask.create_dspark_attention_mask(
241
+ batch_size=B, seq_len=seq_len,
242
+ num_blocks=num_blocks, block_size=block_size,
243
+ device=torch.device(device),
244
+ )
245
+ L_ctx = seq_len
246
+
247
+ # Block 0 should NOT attend to Block 1's draft tokens
248
+ cross_block = mask[0, 0, 0:3, L_ctx+3:L_ctx+6]
249
+ assert (cross_block == float("-inf")).all(), \
250
+ "Cross-block attention should be -inf"
251
+
252
+ # Block 1 should NOT attend to Block 0's draft tokens
253
+ cross_block_2 = mask[0, 0, 3:6, L_ctx:L_ctx+3]
254
+ assert (cross_block_2 == float("-inf")).all(), \
255
+ "Cross-block attention should be -inf"
tests/test_sampling.py ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tests for sampling utilities."""
2
+
3
+ import torch
4
+ import pytest
5
+
6
+ from uraionspec.utils.sampling import (
7
+ logits_to_probs,
8
+ sample_tokens,
9
+ sample_residual,
10
+ gather_token_probs,
11
+ )
12
+
13
+
14
+ class TestSampling:
15
+ """Test sampling utilities."""
16
+
17
+ def test_logits_to_probs_greedy(self):
18
+ logits = torch.tensor([[0.0, 10.0, 0.0]])
19
+ probs = logits_to_probs(logits, temperature=0.0)
20
+ assert probs.shape == (1, 3)
21
+ assert probs[0, 1] == 1.0 # argmax at index 1
22
+
23
+ def test_logits_to_probs_temperature(self):
24
+ logits = torch.randn(2, 100)
25
+ probs = logits_to_probs(logits, temperature=1.0)
26
+ assert probs.shape == (2, 100)
27
+ assert torch.allclose(probs.sum(dim=-1), torch.ones(2))
28
+
29
+ def test_sample_tokens_greedy(self):
30
+ logits = torch.randn(2, 5, 50)
31
+ tokens = sample_tokens(logits, temperature=0.0)
32
+ assert tokens.shape == (2, 5)
33
+ assert (tokens >= 0).all() and (tokens < 50).all()
34
+
35
+ def test_sample_tokens_temperature(self):
36
+ logits = torch.randn(2, 50)
37
+ tokens = sample_tokens(logits, temperature=1.0)
38
+ assert tokens.shape == (2,)
39
+ assert (tokens >= 0).all() and (tokens < 50).all()
40
+
41
+ def test_sample_tokens_2d(self):
42
+ logits = torch.randn(3, 100)
43
+ tokens = sample_tokens(logits, temperature=0.5)
44
+ assert tokens.shape == (3,)
45
+
46
+ def test_sample_residual(self):
47
+ target = torch.softmax(torch.randn(2, 50) + 2, dim=-1)
48
+ draft = torch.softmax(torch.randn(2, 50), dim=-1)
49
+ tokens = sample_residual(target, draft)
50
+ assert tokens.shape == (2,)
51
+ assert (tokens >= 0).all() and (tokens < 50).all()
52
+
53
+ def test_sample_residual_identical(self):
54
+ """When target == draft, residual should fall back to target."""
55
+ probs = torch.softmax(torch.randn(2, 50), dim=-1)
56
+ tokens = sample_residual(probs, probs)
57
+ assert tokens.shape == (2,)
58
+
59
+ def test_gather_token_probs(self):
60
+ probs = torch.tensor([[0.1, 0.7, 0.2], [0.3, 0.3, 0.4]])
61
+ token_ids = torch.tensor([1, 2])
62
+ gathered = gather_token_probs(probs, token_ids)
63
+ assert gathered.shape == (2,)
64
+ assert gathered[0].item() == pytest.approx(0.7)
65
+ assert gathered[1].item() == pytest.approx(0.4)