popgym.baselines.models.s4d =========================== .. py:module:: popgym.baselines.models.s4d .. autoapi-nested-parse:: Standalone version of Structured (Sequence) State Space (S4) model. Attributes ---------- .. autoapisummary:: popgym.baselines.models.s4d.contract popgym.baselines.models.s4d.contract_expression popgym.baselines.models.s4d.log popgym.baselines.models.s4d.has_cauchy_extension popgym.baselines.models.s4d.has_pykeops popgym.baselines.models.s4d.combinations Classes ------- .. autoapisummary:: popgym.baselines.models.s4d.DropoutNd popgym.baselines.models.s4d.OptimModule popgym.baselines.models.s4d.SSKernelNPLR popgym.baselines.models.s4d.SSKernelDiag popgym.baselines.models.s4d.SSKernel popgym.baselines.models.s4d.S4 Functions --------- .. autoapisummary:: popgym.baselines.models.s4d.get_logger popgym.baselines.models.s4d.Activation popgym.baselines.models.s4d.LinearActivation popgym.baselines.models.s4d.power popgym.baselines.models.s4d.transition popgym.baselines.models.s4d.rank_correction popgym.baselines.models.s4d.nplr popgym.baselines.models.s4d.dplr popgym.baselines.models.s4d.ssm popgym.baselines.models.s4d.combination Module Contents --------------- .. py:data:: contract .. py:data:: contract_expression .. py:function:: get_logger(name=__name__, level=logging.INFO) -> logging.Logger Initializes multi-GPU-friendly python logger. .. py:data:: log Cauchy and Vandermonde kernels .. py:data:: has_cauchy_extension :value: True .. py:data:: has_pykeops :value: True .. py:function:: Activation(activation=None, dim=-1) .. py:function:: LinearActivation(d_input, d_output, bias=True, transposed=False, activation=None, activate=False, **kwargs) Returns a linear nn.Module with control over axes order, initialization, and activation .. py:class:: DropoutNd(p: float = 0.5, tie=True, transposed=True) Bases: :py:obj:`torch.nn.Module` .. py:attribute:: p :value: 0.5 .. py:attribute:: tie :value: True .. py:attribute:: transposed :value: True .. py:attribute:: binomial .. py:method:: forward(X) X: (batch, dim, lengths...) .. py:function:: power(L, A, v=None) Compute A^L and the scan sum_i A^i v_i A: (..., N, N) v: (..., N, L) .. py:function:: transition(measure, N) A, B transition matrices for different measures .. py:function:: rank_correction(measure, N, rank=1, dtype=torch.float) Return low-rank matrix L such that A + L is normal .. py:function:: nplr(measure, N, rank=1, dtype=torch.float, diagonalize_precision=True) Return w, p, q, V, B such that (w - p q^*, B) is unitarily equivalent to the original HiPPO A, B by the matrix V i.e. A = V[w - p q^*]V^*, B = V B .. py:function:: dplr(scaling, N, rank=1, H=1, dtype=torch.float, real_scale=1.0, imag_scale=1.0, random_real=False, random_imag=False, normalize=False, diagonal=True, random_B=False) .. py:function:: ssm(measure, N, R, H, **ssm_args) Dispatcher to create single SSM initialization N: state size R: rank (for DPLR parameterization) H: number of independent SSM copies .. py:data:: combinations .. py:function:: combination(measures, N, R, S, **ssm_args) .. py:class:: OptimModule Bases: :py:obj:`torch.nn.Module` Interface for Module that allows registering buffers/parameters with configurable optimizer hyperparameters .. py:method:: register(name, tensor, lr=None) Register a tensor with a configurable learning rate and 0 weight decay .. py:class:: SSKernelNPLR(w, P, B, C, log_dt, L=None, lr=None, verbose=False, keops=False, real_type='exp', real_tolerance=0.001, bandlimit=None) Bases: :py:obj:`OptimModule` Stores a representation of and computes the SSKernel function K_L(A^dt, B^dt, C) corresponding to a discretized state space, where A is Normal + Low Rank (NPLR) .. py:attribute:: verbose :value: False .. py:attribute:: keops :value: False .. py:attribute:: bandlimit :value: None .. py:attribute:: real_type :value: 'exp' .. py:attribute:: real_tolerance :value: 0.001 .. py:attribute:: rank .. py:attribute:: H .. py:attribute:: N .. py:attribute:: n_ssm .. py:attribute:: broadcast .. py:attribute:: C .. py:attribute:: l_max :value: None .. py:method:: forward(state=None, rate=1.0, L=None) state: (B, H, N) initial state rate: sampling rate factor L: target length returns: (C, H, L) convolution kernel (generally C=1) (B, H, L) output from initial state .. py:method:: default_state(*batch_shape) .. py:method:: step(u, state) Must have called self._setup_step() and created state with self.default_state() before calling this .. py:class:: SSKernelDiag(A, B, C, log_dt, L=None, disc='bilinear', real_type='exp', lr=None, bandlimit=None) Bases: :py:obj:`OptimModule` Version using (complex) diagonal state matrix (S4D) .. py:attribute:: L :value: None .. py:attribute:: disc :value: 'bilinear' .. py:attribute:: bandlimit :value: None .. py:attribute:: real_type :value: 'exp' .. py:attribute:: H .. py:attribute:: N .. py:attribute:: n_ssm .. py:attribute:: repeat .. py:attribute:: channels .. py:attribute:: C .. py:method:: forward(L, state=None, rate=1.0, u=None) state: (B, H, N) initial state rate: sampling rate factor L: target length returns: (C, H, L) convolution kernel (generally C=1) (B, H, L) output from initial state .. py:method:: default_state(*batch_shape) .. py:method:: step(u, state) .. py:method:: forward_state(u, state) .. py:class:: SSKernel(H, N=64, L=None, measure='legs', rank=1, channels=1, dt_min=0.001, dt_max=0.1, deterministic=False, lr=None, mode='nplr', n_ssm=None, verbose=False, measure_args={}, **kernel_args) Bases: :py:obj:`torch.nn.Module` Wrapper around SSKernel parameterizations. The SSKernel is expected to support the interface forward() default_state() _setup_step() step() .. py:attribute:: N :value: 64 .. py:attribute:: H .. py:attribute:: channels :value: 1 .. py:attribute:: n_ssm .. py:attribute:: mode :value: 'nplr' .. py:attribute:: verbose :value: False .. py:attribute:: kernel_args .. py:method:: forward(state=None, L=None, rate=None) .. py:method:: forward_state(u, state) Forward the state through a sequence, i.e. computes the state after passing chunk through SSM state: (B, H, N) u: (B, H, L) Returns: (B, H, N) .. py:method:: step(u, state, **kwargs) .. py:method:: default_state(*args, **kwargs) .. py:class:: S4(d_model, d_state=64, l_max=None, channels=1, bidirectional=False, activation='gelu', postact='glu', hyper_act=None, dropout=0.0, tie_dropout=False, bottleneck=None, gate=None, transposed=True, verbose=False, **kernel_args) Bases: :py:obj:`torch.nn.Module` .. py:attribute:: d_model .. py:attribute:: H .. py:attribute:: N :value: 64 .. py:attribute:: L :value: None .. py:attribute:: bidirectional :value: False .. py:attribute:: channels :value: 1 .. py:attribute:: transposed :value: True .. py:attribute:: gate :value: None .. py:attribute:: bottleneck :value: None .. py:attribute:: hyper .. py:attribute:: D .. py:attribute:: kernel .. py:attribute:: activation .. py:attribute:: dropout .. py:attribute:: output_linear .. py:method:: forward(u, state=None, rate=1.0, lengths=None, **kwargs) u: (B H L) if self.transposed else (B L H) state: (H N) never needed unless you know what you're doing Returns: same shape as u .. py:method:: setup_step(**kwargs) .. py:method:: step(u, state) Step one time step as a recurrent model. Intended to be used during validation. u: (B H) state: (B H N) Returns: output (B H), state (B H N) .. py:method:: default_state(*batch_shape, device=None) .. py:property:: d_output