Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion grail/protocol/grail_verifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -253,14 +253,23 @@ def create_commitment(self, hidden_state: torch.Tensor, r_vec: torch.Tensor) ->
"sketch": sketch_val,
}

def create_commitments_batch(self, h_layer: torch.Tensor, r_vec: torch.Tensor) -> list[dict]:
# ===================== UPDATED FUNCTION =====================
def create_commitments_batch(
self,
h_layer: torch.Tensor,
r_vec: torch.Tensor,
projected_s_vals: torch.Tensor | None = None
) -> list[dict]:
"""Create commitments for all positions at once (vectorized).

Produces bit-identical results to calling create_commitment() in a loop.

Args:
h_layer: Hidden states [seq_len, hidden_dim] on any device.
r_vec: Coefficient vector [topk] (int8, typically on CPU).
projected_s_vals: Optional pre-computed s_vals (for deterministic mode).
Currently not used; exists for compatibility with the miner's
fixed-point projection path.

Returns:
List of commitment dicts, one per position.
Expand Down Expand Up @@ -296,6 +305,7 @@ def create_commitments_batch(self, h_layer: torch.Tensor, r_vec: torch.Tensor) -

# Step 4: Package as list of dicts
return [{"sketch": sketch_vals[pos]} for pos in range(seq_len)]
# ============================================================

def verify_commitment(
self,
Expand Down