MCPcopy Create free account
hub / github.com/BorealisAI/scaleformer / transition

Function transition

models/FiLMMS.py:35–90  ·  view source on GitHub ↗

A, B transition matrices for different measures. measure: the type of measure legt - Legendre (translated) legs - Legendre (scaled) glagt - generalized Laguerre (translated) lagt, tlagt - previous versions of (tilted) Laguerre with slightly different normalization

(measure, N, **measure_args)

Source from the content-addressed store, hash-verified

33 return x
34
35def transition(measure, N, **measure_args):
36 """ A, B transition matrices for different measures.
37
38 measure: the type of measure
39 legt - Legendre (translated)
40 legs - Legendre (scaled)
41 glagt - generalized Laguerre (translated)
42 lagt, tlagt - previous versions of (tilted) Laguerre with slightly different normalization
43 """
44 # Laguerre (translated)
45 if measure == 'lagt':
46 b = measure_args.get('beta', 1.0)
47 A = np.eye(N) / 2 - np.tril(np.ones((N, N)))
48 B = b * np.ones((N, 1))
49 if measure == 'tlagt':
50 # beta = 1 corresponds to no tilt
51 b = measure_args.get('beta', 1.0)
52 A = (1.-b)/2 * np.eye(N) - np.tril(np.ones((N, N)))
53 B = b * np.ones((N, 1))
54 # Generalized Laguerre
55 # alpha 0, beta small is most stable (limits to the 'lagt' measure)
56 # alpha 0, beta 1 has transition matrix A = [lower triangular 1]
57 if measure == 'glagt':
58 alpha = measure_args.get('alpha', 0.0)
59 beta = measure_args.get('beta', 0.01)
60 A = -np.eye(N) * (1 + beta) / 2 - np.tril(np.ones((N, N)), -1)
61 B = ss.binom(alpha + np.arange(N), np.arange(N))[:, None]
62
63 L = np.exp(.5 * (ss.gammaln(np.arange(N)+alpha+1) - ss.gammaln(np.arange(N)+1)))
64 A = (1./L[:, None]) * A * L[None, :]
65 B = (1./L[:, None]) * B * np.exp(-.5 * ss.gammaln(1-alpha)) * beta**((1-alpha)/2)
66 # Legendre (translated)
67 elif measure == 'legt':
68 Q = np.arange(N, dtype=np.float64)
69 R = (2*Q + 1) ** .5
70 j, i = np.meshgrid(Q, Q)
71 A = R[:, None] * np.where(i < j, (-1.)**(i-j), 1) * R[None, :]
72 B = R[:, None]
73 A = -A
74 # LMU: equivalent to LegT up to normalization
75 elif measure == 'lmu':
76 Q = np.arange(N, dtype=np.float64)
77 R = (2*Q + 1)[:, None] # / theta
78 j, i = np.meshgrid(Q, Q)
79 A = np.where(i < j, -1, (-1.)**(i-j+1)) * R
80 B = (-1.)**Q[:, None] * R
81 # Legendre (scaled)
82 elif measure == 'legs':
83 q = np.arange(N, dtype=np.float64)
84 col, row = np.meshgrid(q, q)
85 r = 2 * q + 1
86 M = -(np.where(row >= col, r, 0) - np.diag(q))
87 T = np.sqrt(np.diag(2 * q + 1))
88 A = T @ M @ np.linalg.inv(T)
89 B = np.diag(T)[:, None]
90 return A, B
91
92class HiPPO_LegT(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected