Convert float tensor to spike sequence via integer quantizer + LIF encoder. Args: x (torch.Tensor): Float input tensor. quantizer (Callable): Float-to-int quantization function. lif_quantizer (SpikeCountBaseLIFNode): Spike encoder. x_zero (torch.Tensor, opti
(
x: torch.Tensor,
quantizer,
lif_quantizer: SpikeCountBaseLIFNode,
x_zero: torch.Tensor = None
)
| 378 | return x_reconstructed |
| 379 | |
| 380 | def quant( |
| 381 | x: torch.Tensor, |
| 382 | quantizer, |
| 383 | lif_quantizer: SpikeCountBaseLIFNode, |
| 384 | x_zero: torch.Tensor = None |
| 385 | ) -> torch.Tensor: |
| 386 | """ |
| 387 | Convert float tensor to spike sequence via integer quantizer + LIF encoder. |
| 388 | |
| 389 | Args: |
| 390 | x (torch.Tensor): Float input tensor. |
| 391 | quantizer (Callable): Float-to-int quantization function. |
| 392 | lif_quantizer (SpikeCountBaseLIFNode): Spike encoder. |
| 393 | x_zero (torch.Tensor, optional): Optional zero-point tensor (modified in-place if needed). |
| 394 | |
| 395 | Returns: |
| 396 | torch.Tensor: Spike sequence with shape [T, *x.shape] |
| 397 | """ |
| 398 | if quantizer is None: |
| 399 | raise ValueError("A quantizer must be provided to map float → int.") |
| 400 | |
| 401 | spike_count = quantizer(x) |
| 402 | return spike_quant(spike_count, lif_quantizer, x_zero) |
| 403 | |
| 404 | def dequant( |
| 405 | spike: torch.Tensor, |
no test coverage detected