MCPcopy Create free account
hub / github.com/FlashSampling/FlashSampling / make_synthetic_inputs

Function make_synthetic_inputs

src/fused_mm_sampling/testing.py:27–72  ·  view source on GitHub ↗

Build weights and hidden_states that produce known logits. Creates up to two hidden states: one with ascending logits (favors high token indices) and one with descending logits (favors low token indices). All logits are shifted negative via :func:`shift_logits_negative`.

(
    vocab_size: int = 256,
    hidden_size: int = 10,
    n_hidden_states: int = 2,
    device: torch.device = torch.device("cuda"),
    tp: TPInfo = TP1,
)

Source from the content-addressed store, hash-verified

25 vocab_size: int
26 hidden_size: int
27
28
29def make_synthetic_inputs(
30 vocab_size: int = 256,
31 hidden_size: int = 10,
32 n_hidden_states: int = 2,
33 device: torch.device = torch.device("cuda"),
34 tp: TPInfo = TP1,
35) -> SyntheticInputs:
36 """Build weights and hidden_states that produce known logits.
37
38 Creates up to two hidden states: one with ascending logits (favors high
39 token indices) and one with descending logits (favors low token indices).
40 All logits are shifted negative via :func:`shift_logits_negative`.
41 """
42 logits1 = torch.arange(-vocab_size / 2, vocab_size / 2, dtype=torch.float32)[None, :]
43 logits2 = torch.arange(vocab_size / 2, -vocab_size / 2, step=-1, dtype=torch.float32)[None, :]
44 all_logits = [logits1, logits2]
45 logits = torch.cat(all_logits[:n_hidden_states], dim=0).to(device)
46 n_hidden_states = logits.shape[0]
47
48 U, _, _ = torch.linalg.svd(logits, full_matrices=False) # noqa: N806
49
50 torch.manual_seed(0)
51 hidden_states = torch.cat(
52 [U, torch.rand((n_hidden_states, hidden_size - n_hidden_states), device=device)],
53 dim=1,
54 ).to(device)
55 weights = torch.linalg.pinv(hidden_states) @ logits # [D, V]
56
57 weights_bf16 = weights.bfloat16().T.contiguous() # [V, D]
58 hidden_states_bf16 = hidden_states.bfloat16()
59 weights_bf16, hidden_states_bf16 = shift_logits_negative(
60 weights_bf16,
61 hidden_states_bf16,
62 offset=float(vocab_size),
63 )
64
65 weights_bf16, hidden_states_bf16 = pad_to_tma_alignment(weights_bf16, hidden_states_bf16)
66 weights_bf16 = shard_weights(weights_bf16, tp)
67
68 return SyntheticInputs(
69 weights=weights_bf16,
70 hidden_states=hidden_states_bf16,
71 logits=logits,
72 vocab_size=vocab_size,
73 hidden_size=weights_bf16.shape[1],
74 )
75

Callers 5

test_top_k_top_pFunction · 0.90
test_greedy_samplingFunction · 0.90
verify_greedy_tp2Function · 0.85

Calls 4

shift_logits_negativeFunction · 0.85
pad_to_tma_alignmentFunction · 0.85
shard_weightsFunction · 0.85
SyntheticInputsClass · 0.85

Tested by

no test coverage detected