jaxsnn.event.solver.lif_analytical_double_time
Classes
|
Functions
-
jaxsnn.event.solver.lif_analytical_double_time.safe_log(x: jax.Array, eps: float | None = None) → jax.Array
-
jaxsnn.event.solver.lif_analytical_double_time.safe_sqrt(x: jax.Array, eps: float | None = None) → jax.Array
-
jaxsnn.event.solver.lif_analytical_double_time.smallest_normal_like(x: jax.Array, mul: float = 1.0) → jax.Array
-
jaxsnn.event.solver.lif_analytical_double_time.ttfs_double_time(tau_mem: float, tau_syn: float, v_th: float, v_leak: float, state: jaxsnn.event.states.LIFState, t_max: float) → jax.Array
-
jaxsnn.event.solver.lif_analytical_double_time.ttfs_double_time_nonzero_current(state: jaxsnn.event.states.LIFState, tau_mem: float, t_max: float, v_th: float, v_leak: float) → jax.Array
-
jaxsnn.event.solver.lif_analytical_double_time.ttfs_double_time_zero_current(state: jaxsnn.event.states.LIFState, tau_mem: float, t_max: float, v_th: float, v_leak: float) → jax.Array
-
jaxsnn.event.solver.lif_analytical_double_time.ttfs_over_threshold_bwd(res, g) Backward pass for the custom VJP function.
-
jaxsnn.event.solver.lif_analytical_double_time.ttfs_over_threshold_fwd(tau_mem: float, v_th: float, v_leak: float, state: jaxsnn.event.states.LIFState, t_max: float) → Tuple[jax.Array, Tuple] Forward pass for the custom VJP function.
-
jaxsnn.event.solver.lif_analytical_double_time.ttfs_subthreshold(tau_mem: float, v_th: float, v_leak: float, state: jaxsnn.event.states.LIFState, t_max: float, eps: float | None = None) → jax.Array