itrails 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
itrails/__init__.py ADDED
@@ -0,0 +1,4 @@
1
+ try:
2
+ from ._version import version as __version__ # type: ignore
3
+ except ImportError:
4
+ __version__ = "unknown"
itrails/_version.py ADDED
@@ -0,0 +1,21 @@
1
+ # file generated by setuptools-scm
2
+ # don't change, don't track in version control
3
+
4
+ __all__ = ["__version__", "__version_tuple__", "version", "version_tuple"]
5
+
6
+ TYPE_CHECKING = False
7
+ if TYPE_CHECKING:
8
+ from typing import Tuple
9
+ from typing import Union
10
+
11
+ VERSION_TUPLE = Tuple[Union[int, str], ...]
12
+ else:
13
+ VERSION_TUPLE = object
14
+
15
+ version: str
16
+ __version__: str
17
+ __version_tuple__: VERSION_TUPLE
18
+ version_tuple: VERSION_TUPLE
19
+
20
+ __version__ = version = '0.1.0'
21
+ __version_tuple__ = version_tuple = (0, 1, 0)
@@ -0,0 +1,140 @@
1
+ import numba as nb
2
+ import numpy as np
3
+
4
+
5
+ def combine_states(
6
+ state_dict_1, state_dict_2, state_dict_sum, final_probs_1, final_probs_2
7
+ ):
8
+ """
9
+ Function that combines two dictionaries of individual states and their final probabilities into a single dictionary of combined states and starting probabilities.
10
+
11
+ :param state_dict_1: Dictionary of states and indices for the first CTMC, can be 1 sequence or 2 sequence CTMC.
12
+ :type state_dict_1: Numba Dictionary of Key: Tuple of int64 and Value: int64.
13
+ :param state_dict_2: Dictionary of states and indices for the second CTMC, can only be 1 sequence CTMC.
14
+ :type state_dict_2: Numba Dictionary of Key: Tuple(int64, int64) and Value: int64
15
+ :param state_dict_sum: Dictionary of states and indices for the combined CTMC.
16
+ :type state_dict_sum: Numba Dictionary of Key: Tuple of int64 and Value: int64.
17
+ :param final_probs_1: Array of final probabilities for a single key of the first CTMC.
18
+ :type final_probs_1: Array of type: float64[:, :]
19
+ :param final_probs_2: Array of final probabilities for the only key of the second CTMC.
20
+ :type final_probs_2: Array of type: float64[:, :]
21
+ :return: Array of starting probabilities for each combined state in the combined CTMC.
22
+ :rtype: Array of type: float64[:, :]
23
+ """
24
+ len_array = len(list(state_dict_1.keys())[0]) + len(list(state_dict_2.keys())[0])
25
+ init_comb_dict = {}
26
+ init_comb_all = np.zeros((len(state_dict_sum.keys())), dtype=np.float64)
27
+ for key_1, index_1 in state_dict_1.items():
28
+ left_1 = key_1[: len(key_1) // 2]
29
+ right_1 = key_1[len(key_1) // 2 :]
30
+ for key_2, index_2 in state_dict_2.items():
31
+ left_2 = key_2[: len(key_2) // 2]
32
+ right_2 = key_2[len(key_2) // 2 :]
33
+ comb_state = np.zeros((len_array), dtype=np.int64)
34
+ used_values_1 = {}
35
+ used_values_2 = {}
36
+ current_value = 1
37
+ index = 0
38
+ for value in left_1:
39
+ if value in used_values_1:
40
+ comb_state[index] = used_values_1[value]
41
+ else:
42
+ used_values_1[value] = current_value
43
+ comb_state[index] = current_value
44
+ current_value += 1
45
+ index += 1
46
+
47
+ for value in left_2:
48
+ if value in used_values_2:
49
+ comb_state[index] = used_values_2[value]
50
+ else:
51
+ used_values_2[value] = current_value
52
+ comb_state[index] = current_value
53
+ current_value += 1
54
+ index += 1
55
+
56
+ for value in right_1:
57
+ if value in used_values_1:
58
+ comb_state[index] = used_values_1[value]
59
+ else:
60
+ used_values_1[value] = current_value
61
+ comb_state[index] = current_value
62
+ current_value += 1
63
+ index += 1
64
+
65
+ for value in right_2:
66
+ if value in used_values_2:
67
+ comb_state[index] = used_values_2[value]
68
+ else:
69
+ used_values_2[value] = current_value
70
+ comb_state[index] = current_value
71
+ current_value += 1
72
+ index += 1
73
+
74
+ init_comb_dict[tuple(comb_state)] = (
75
+ final_probs_1[0, index_1] * final_probs_2[0, index_2]
76
+ )
77
+ for state in init_comb_dict.keys():
78
+ index_AB = state_dict_sum[state]
79
+ init_comb_all[index_AB] = init_comb_dict[state]
80
+ return init_comb_all
81
+
82
+
83
+ def combine_states_wrapper(
84
+ state_dict_1,
85
+ state_dict_2,
86
+ state_dict_sum,
87
+ final_probs_1,
88
+ final_probs_2,
89
+ ):
90
+ """
91
+ Wrapper function that combines dictionaries of states and their final probabilities into a single dictionary of combined states and starting probabilities.
92
+
93
+ :param state_dict_1: Dictionary of states and indices for the first CTMC, can be 1 sequence or 2 sequence CTMC.
94
+ :type state_dict_1: Numba Dictionary of Key: Tuple of int64 and Value: int64.
95
+ :param state_dict_2: Dictionary of states and indices for the second CTMC, can only be 1 sequence CTMC.
96
+ :type state_dict_2: Numba Dictionary of Key: Tuple(int64, int64) and Value: int64.
97
+ :param state_dict_sum: Dictionary of states and indices for the combined CTMC.
98
+ :type state_dict_sum: Numba Dictionary of Key: Tuple of int64 and Value: int64.
99
+ :param final_probs_1: Final probability dictionary for the first CTMC.
100
+ :type final_probs_1: Numba Dictionary of Key: UniTuple(nb.types.UniTuple(int64, 3), 2) and Value: float64[:, :].
101
+ :param final_probs_2: Final probability dictionary for the second CTMC.
102
+ :type final_probs_2: Numba Dictionary of Key: UniTuple(nb.types.UniTuple(int64, 3), 2) and Value: float64[:, :].
103
+ :raises NotImplementedError: Not implemented for more than 2 species in state_dict_1 or and more than 1 species in state_dict_2.
104
+ :raises Exception: Fallback if invalid format in final_probs_1 or final_probs_2.
105
+ :return: Dictionary of combined states and starting probabilities for each state.
106
+ :rtype: Numba Dictionary of Key: UniTuple(nb.types.UniTuple(int64, 3), 2) and Value: float64[:, :].
107
+ """
108
+ pi_dict = nb.typed.Dict.empty(
109
+ key_type=nb.types.UniTuple(nb.types.UniTuple(nb.types.int64, 3), 2),
110
+ value_type=nb.types.float64[:, :],
111
+ )
112
+
113
+ start_placeholder = ((-1, -1, -1), (-1, -1, -1))
114
+
115
+ if len(final_probs_1.keys()) > 1 and len(final_probs_2.keys()) > 1:
116
+ raise NotImplementedError
117
+
118
+ elif len(final_probs_1.keys()) > 1 and len(final_probs_2.keys()) == 1:
119
+ prob2 = final_probs_2[start_placeholder]
120
+ for path, prob1 in final_probs_1.items():
121
+ pi_vector_combined = combine_states(
122
+ state_dict_1, state_dict_2, state_dict_sum, prob1, prob2
123
+ )
124
+
125
+ pi_dict[path] = pi_vector_combined.reshape(1, -1)
126
+
127
+ return pi_dict
128
+
129
+ elif len(final_probs_1.keys()) == 1 and len(final_probs_2.keys()) == 1:
130
+
131
+ prob1 = final_probs_1[start_placeholder]
132
+ prob2 = final_probs_2[start_placeholder]
133
+ pi_vector_combined = combine_states(
134
+ state_dict_1, state_dict_2, state_dict_sum, prob1, prob2
135
+ )
136
+ pi_dict[start_placeholder] = pi_vector_combined.reshape(1, -1)
137
+ return pi_dict
138
+
139
+ else:
140
+ raise Exception("Invalid final_probs_1 or final_probs_2")
itrails/cutpoints.py ADDED
@@ -0,0 +1,65 @@
1
+ import numpy as np
2
+ from scipy.stats import expon, truncexpon
3
+
4
+
5
+ def cutpoints_AB(n_int_AB, t_AB, coal_AB):
6
+ """
7
+ This function returns the cutpoints for the
8
+ intervals for the two-sequence CTMC. The cutpoints
9
+ will be defined by the quantiles of a truncated
10
+ exponential distribution.
11
+
12
+ :param n_int_AB: Number of intervals in the two-sequence CTMC.
13
+ :type n_int_AB: int
14
+ :param t_AB: Total time interval of the two-sequence CTMC
15
+ :type t_AB: float
16
+ :param coal_AB: coalescent rate of the two-sequence CTMC.
17
+ :type coal_AB: float
18
+ :return: cut_AB
19
+ :rtype: np.array
20
+ """
21
+ quantiles_AB = np.array(list(range(n_int_AB + 1))) / n_int_AB
22
+ lower, upper, scale = 0, t_AB, 1 / coal_AB
23
+ cut_AB = truncexpon.ppf(
24
+ quantiles_AB, b=(upper - lower) / scale, loc=lower, scale=scale
25
+ )
26
+ return cut_AB
27
+
28
+
29
+ def cutpoints_ABC(n_int_ABC, coal_ABC):
30
+ """
31
+ This function returns the cutpoints for the
32
+ intervals for the three-sequence CTMC. The cutpoints
33
+ will be defined by the quantiles of an exponential
34
+ distribution.
35
+
36
+ :param n_int_ABC: Number of intervals in the three-sequence CTMC.
37
+ :type n_int_ABC: int
38
+ :param coal_ABC: coalescent rate of the three-sequence CTMC.
39
+ :type coal_ABC: float
40
+ :return: cut_ABC
41
+ :rtype: np.array
42
+ """
43
+ quantiles_AB = np.array(list(range(n_int_ABC + 1))) / n_int_ABC
44
+ cut_ABC = expon.ppf(quantiles_AB, scale=1 / coal_ABC)
45
+ return cut_ABC
46
+
47
+
48
+ def get_times(cut, intervals):
49
+ """
50
+ Returns a list of time differences for each specified interval.
51
+
52
+ This function computes the duration for each interval defined by the given indices.
53
+ It does so by subtracting the earlier cutpoint from the subsequent one based on the indices
54
+ provided in the intervals list.
55
+
56
+ :param cut: List of ordered cutpoints.
57
+ :type cut: list[float]
58
+ :param intervals: Ordered indices of cutpoints that define the intervals.
59
+ :type intervals: list[int]
60
+ :return: List of time differences between consecutive cutpoints as defined by intervals.
61
+ :rtype: list[float]
62
+ """
63
+ return [
64
+ cut[intervals[i + 1]] - cut[intervals[i]] for i in range(len(intervals) - 1)
65
+ ]
itrails/deepest_ti.py ADDED
@@ -0,0 +1,256 @@
1
+ import numpy as np
2
+
3
+
4
+ def deep_identify(
5
+ current,
6
+ absorbing_state,
7
+ omega_nonrev_counts,
8
+ inverted_omega_nonrev_counts,
9
+ path,
10
+ all_paths_dict,
11
+ by_l=-1,
12
+ by_r=-1,
13
+ ):
14
+ """
15
+ Recursive function that identifies, in each iteration, the possible paths that can be taken from the current state to the final state. The function is called recursively until the final state is reached. The function stores the paths in the paths_array and all_paths_array arrays.
16
+
17
+ :param current: Current omega state in the recursion
18
+ :type current: Tuple of int
19
+ :param absorbing_state: Absorbing state
20
+ :type absorbing_state: Tuple of int
21
+ :param omega_nonrev_counts: Dictionary containing the number of non-reversible coalescents (value) for each omega state (key)
22
+ :type omega_nonrev_counts: Numba typed Dict
23
+ :param inverted_omega_nonrev_counts: Dictionary containing the omega states (value) for each number of non-reversible coalescents (key)
24
+ :type inverted_omega_nonrev_counts: Numba typed Dict
25
+ :param path: Current path in the recursion
26
+ :type path: Numpy array
27
+ :param all_paths_dict: Dictionary that recursively gets filled up with all the paths
28
+ :type all_paths_dict: Numpy array
29
+ :param by_l: Current omega left subpath, defaults to -1 (initial placeholder value)
30
+ :type by_l: int, optional
31
+ :param by_r: Current omega right subpath, defaults to -1 (initial placeholder value)
32
+ :type by_r: int, optional
33
+ """
34
+ # Calculate differences
35
+ diff_l = omega_nonrev_counts[absorbing_state[0]] - omega_nonrev_counts[current[0]]
36
+ diff_r = omega_nonrev_counts[absorbing_state[1]] - omega_nonrev_counts[current[1]]
37
+
38
+ # Termination condition
39
+ if diff_l <= 1 and diff_r <= 1:
40
+ key = (by_l, by_r)
41
+ if key not in all_paths_dict:
42
+ all_paths_dict[key] = []
43
+ # Create a copy of the current path
44
+ path_copy = [tuple(p) for p in path]
45
+ all_paths_dict[key].append(path_copy)
46
+ return
47
+
48
+ start_l = omega_nonrev_counts[current[0]]
49
+ start_r = omega_nonrev_counts[current[1]]
50
+ end_l = omega_nonrev_counts[absorbing_state[0]]
51
+ end_r = omega_nonrev_counts[absorbing_state[1]]
52
+
53
+ # Explore next states for left
54
+ if start_l < end_l:
55
+ next_states_l = inverted_omega_nonrev_counts[start_l + 1]
56
+ for left in next_states_l:
57
+ new_state = (left, current[1])
58
+ new_by_l = (
59
+ by_l
60
+ if by_l != -1
61
+ else (
62
+ left
63
+ if omega_nonrev_counts[left] == 1 and start_l + 1 != end_l
64
+ else -1
65
+ )
66
+ )
67
+ path.append(new_state)
68
+ deep_identify(
69
+ new_state,
70
+ absorbing_state,
71
+ omega_nonrev_counts,
72
+ inverted_omega_nonrev_counts,
73
+ path,
74
+ all_paths_dict,
75
+ new_by_l,
76
+ by_r,
77
+ )
78
+ path.pop()
79
+
80
+ # Explore next states for right
81
+ if start_r < end_r:
82
+ next_states_r = inverted_omega_nonrev_counts[start_r + 1]
83
+ for right in next_states_r:
84
+ new_state = (current[0], right)
85
+ new_by_r = (
86
+ by_r
87
+ if by_r != -1
88
+ else (
89
+ right
90
+ if omega_nonrev_counts[right] == 1 and start_r + 1 != end_r
91
+ else -1
92
+ )
93
+ )
94
+ path.append(new_state)
95
+ deep_identify(
96
+ new_state,
97
+ absorbing_state,
98
+ omega_nonrev_counts,
99
+ inverted_omega_nonrev_counts,
100
+ path,
101
+ all_paths_dict,
102
+ by_l,
103
+ new_by_r,
104
+ )
105
+ path.pop()
106
+
107
+ # Explore next states for both left and right
108
+ if start_l < end_l and start_r < end_r:
109
+ next_states_l = inverted_omega_nonrev_counts[start_l + 1]
110
+ next_states_r = inverted_omega_nonrev_counts[start_r + 1]
111
+ for left in next_states_l:
112
+ for right in next_states_r:
113
+ if omega_nonrev_counts[right] > start_r:
114
+ new_state = (left, right)
115
+ new_by_l = (
116
+ by_l
117
+ if by_l != -1
118
+ else (
119
+ left
120
+ if omega_nonrev_counts[left] == 1 and start_l + 1 != end_l
121
+ else -1
122
+ )
123
+ )
124
+ new_by_r = (
125
+ by_r
126
+ if by_r != -1
127
+ else (
128
+ right
129
+ if omega_nonrev_counts[right] == 1 and start_r + 1 != end_r
130
+ else -1
131
+ )
132
+ )
133
+ path.append(new_state)
134
+ deep_identify(
135
+ new_state,
136
+ absorbing_state,
137
+ omega_nonrev_counts,
138
+ inverted_omega_nonrev_counts,
139
+ path,
140
+ all_paths_dict,
141
+ new_by_l,
142
+ new_by_r,
143
+ )
144
+ path.pop()
145
+
146
+
147
+ def deep_identify_wrapper(
148
+ omega_init,
149
+ absorbing_state,
150
+ omega_nonrev_counts,
151
+ inverted_omega_nonrev_counts,
152
+ path_to_convert,
153
+ ):
154
+ """
155
+ Wrapper function for the deep_identify function. This function initializes the arrays and dictionaries needed for the recursion and calls the deep_identify function. In the end it returns the transformed keys and the paths.
156
+
157
+ :param omega_init: Initial omega state
158
+ :type omega_init: Tuple of int
159
+ :param absorbing_state: Absorbing state
160
+ :type absorbing_state: Tuple of int
161
+ :param omega_nonrev_counts: Dictionary containing the number of non-reversible coalescents (value) for each omega state (key)
162
+ :type omega_nonrev_counts: Numba typed Dict
163
+ :param inverted_omega_nonrev_counts: Dictionary containing the omega states (value) for each number of non-reversible coalescents (key)
164
+ :type inverted_omega_nonrev_counts: Numba typed Dict
165
+ :return: Resulting keys and paths
166
+ :rtype: Tuple(Array, Array, Array, int)
167
+ """
168
+
169
+ all_paths_dict = {}
170
+ path = [omega_init]
171
+ deep_identify(
172
+ omega_init,
173
+ absorbing_state,
174
+ omega_nonrev_counts,
175
+ inverted_omega_nonrev_counts,
176
+ path,
177
+ all_paths_dict,
178
+ )
179
+
180
+ # Convert dictionary to keys_array and paths_array
181
+ keys_array = np.array(list(all_paths_dict.keys()))
182
+ keys_array_final = np.zeros((len(keys_array), 6))
183
+ path_to_convert_array = np.array(path_to_convert)
184
+ flattened_array = np.hstack(path_to_convert_array.ravel())
185
+ coded_dict = {3: 1, 5: 2, 6: 3}
186
+ for i, key in enumerate(keys_array):
187
+ keys_array_final[i] = flattened_array
188
+ if key[0] != -1 and keys_array_final[i][0] == -1:
189
+ keys_array_final[i][0] = coded_dict[key[0]]
190
+ if key[1] != -1 and keys_array_final[i][3] == -1:
191
+ keys_array_final[i][3] = coded_dict[key[1]]
192
+
193
+ # Find the maximum number of paths and subpaths for padding
194
+ max_paths = max(len(paths) for paths in all_paths_dict.values())
195
+ max_subpaths = max(
196
+ max(len(subpath) for subpath in paths) for paths in all_paths_dict.values()
197
+ )
198
+
199
+ # Initialize paths_array with zeros
200
+ paths_array = np.zeros((len(all_paths_dict), max_paths, max_subpaths, 2))
201
+
202
+ # Initialize path_lengths_array to store the length of each path
203
+ path_lengths_array = np.zeros((len(all_paths_dict), max_paths))
204
+
205
+ # Fill paths_array and path_lengths_array
206
+ for i, (key, paths) in enumerate(all_paths_dict.items()):
207
+ for j, subpath in enumerate(paths):
208
+ path_lengths_array[i, j] = len(subpath) # Store the length of the path
209
+ for k, point in enumerate(subpath):
210
+ paths_array[i, j, k] = point
211
+
212
+ return keys_array_final, paths_array, path_lengths_array, max_subpaths
213
+
214
+
215
+ def deepest_ti(trans_mat_noabs, omega_dict_noabs, path):
216
+ """
217
+ This function calculated the integral of matrix exponentials with an infinite time limit.
218
+
219
+ :param trans_mat_noabs: Transition matrix without absorbing states
220
+ :type trans_mat_noabs: Numpy array
221
+ :param omega_dict_noabs: Omega dictionary without absorbing states
222
+ :type omega_dict_noabs: Numpy array
223
+ :param path: Path of omega states
224
+ :type path: Numpy array
225
+ :return: Result of the integral of the series of multiplying matrix exponentials.
226
+ :rtype: Numpy array
227
+ """
228
+ steps = len(path) - 1
229
+ n = trans_mat_noabs.shape[0]
230
+ if steps == 1:
231
+ C_mat = trans_mat_noabs
232
+
233
+ elif steps > 1:
234
+ C_mat = np.zeros((n * steps, n * steps))
235
+ C_mat[0:n, 0:n] = trans_mat_noabs
236
+ for idx in range(1, steps):
237
+
238
+ sub_om_init = (path[idx - 1, 0], path[idx - 1, 1])
239
+ sub_om_fin = (path[idx, 0], path[idx, 1])
240
+ A_mat = (
241
+ np.diag(omega_dict_noabs[sub_om_init].astype(np.float64))
242
+ @ trans_mat_noabs
243
+ @ np.diag(omega_dict_noabs[sub_om_fin].astype(np.float64))
244
+ )
245
+ C_mat[n * idx : n * (idx + 1), n * idx : n * (idx + 1)] = trans_mat_noabs
246
+ C_mat[n * (idx - 1) : n * idx, n * idx : n * (idx + 1)] = A_mat
247
+
248
+ sub_om_init = (path[-2, 0], path[-2, 1])
249
+ sub_om_fin = (path[-1, 0], path[-1, 1])
250
+ A_mat = np.ascontiguousarray(
251
+ np.diag(omega_dict_noabs[sub_om_init].astype(np.float64))
252
+ @ trans_mat_noabs
253
+ @ np.diag(omega_dict_noabs[sub_om_fin].astype(np.float64))
254
+ )
255
+ result = (-np.linalg.inv(C_mat))[:n, -n:] @ A_mat
256
+ return result
@@ -0,0 +1,21 @@
1
+ # Example configuration for itrails-optimize
2
+ fixed_parameters:
3
+ mu: 2e-8
4
+
5
+ optimized_parameters: # [starting, min, max]
6
+ N_AB: [50000, 5000, 500000]
7
+ N_ABC: [50000, 5000, 500000]
8
+ t_1: [240000, 24000, 2400000]
9
+ t_2: [40000, 4000, 400000]
10
+ t_3: [800000, 80000, 8000000]
11
+ t_upper: [745069.3855665945, 74506.93855665945, 7450693.855665945]
12
+ r: [1e-8, 1e-9, 1e-7]
13
+
14
+ settings:
15
+ input_maf: # Path to the MAF alignment file (overwritten by console argument)
16
+ output_name: # Path to the output directory (overwritten by console argument)
17
+ n_cpu: 64
18
+ method: "Nelder-Mead"
19
+ species_list: ["hg38", "panTro5", "gorGor5", "ponAbe2"]
20
+ n_int_AB: 3
21
+ n_int_ABC: 3
itrails/expm.py ADDED
@@ -0,0 +1,166 @@
1
+ from __future__ import division, print_function
2
+
3
+ import math
4
+
5
+ import numba as nb
6
+ import numpy as np
7
+
8
+
9
+ @nb.jit(nopython=True, parallel=False, fastmath=True)
10
+ def expm(A):
11
+ """
12
+ Calculates matrix exponential of a square matrix A.
13
+ Adapted from https://github.com/michael-hartmann/expm/blob/master/python/expm.py
14
+ Algorithm 10.20 from unctions of Matrices: Theory and Computation, Nicholas J. Higham, 2008
15
+ :param A: square matrix
16
+ :type A: np.array
17
+ :return: matrix exponential of A
18
+ :rtype: np.array
19
+ """
20
+ theta3 = 1.5e-2
21
+ theta5 = 2.5e-1
22
+ theta7 = 9.5e-1
23
+ theta9 = 2.1e0
24
+ theta13 = 5.4e0
25
+ # calculate the norm of A
26
+ norm = np.linalg.norm(A, ord=1)
27
+
28
+ if norm < theta3:
29
+ dtype = A.dtype
30
+ dim, dim = A.shape
31
+ b = [120, 60, 12, 1]
32
+
33
+ U = b[1] * np.eye(dim, dtype=dtype)
34
+ V = b[0] * np.eye(dim, dtype=dtype)
35
+
36
+ A2 = A @ A
37
+ A2n = np.eye(dim, dtype=dtype)
38
+
39
+ # evaluate (10.33)
40
+ for i in range(1, 3 // 2 + 1):
41
+ A2n = A2n @ A2
42
+ U += b[2 * i + 1] * A2n
43
+ V += b[2 * i] * A2n
44
+
45
+ U = A @ U
46
+ return np.linalg.solve(V - U, V + U)
47
+
48
+ elif norm < theta5:
49
+ dtype = A.dtype
50
+ dim, dim = A.shape
51
+ b = [30240, 15120, 3360, 420, 30, 1]
52
+
53
+ U = b[1] * np.eye(dim, dtype=dtype)
54
+ V = b[0] * np.eye(dim, dtype=dtype)
55
+
56
+ A2 = A @ A
57
+ A2n = np.eye(dim, dtype=dtype)
58
+
59
+ # evaluate (10.33)
60
+ for i in range(1, 5 // 2 + 1):
61
+ A2n = A2n @ A2
62
+ U += b[2 * i + 1] * A2n
63
+ V += b[2 * i] * A2n
64
+
65
+ U = A @ U
66
+ return np.linalg.solve(V - U, V + U)
67
+
68
+ elif norm < theta7:
69
+ dtype = A.dtype
70
+ dim, dim = A.shape
71
+ b = [17297280, 8648640, 1995840, 277200, 25200, 1512, 56, 1]
72
+
73
+ U = b[1] * np.eye(dim, dtype=dtype)
74
+ V = b[0] * np.eye(dim, dtype=dtype)
75
+
76
+ A2 = A @ A
77
+ A2n = np.eye(dim, dtype=dtype)
78
+
79
+ # evaluate (10.33)
80
+ for i in range(1, 7 // 2 + 1):
81
+ A2n = A2n @ A2
82
+ U += b[2 * i + 1] * A2n
83
+ V += b[2 * i] * A2n
84
+
85
+ U = A @ U
86
+ return np.linalg.solve(V - U, V + U)
87
+
88
+ elif norm < theta9:
89
+ dtype = A.dtype
90
+ dim, dim = A.shape
91
+ b = [
92
+ 17643225600,
93
+ 8821612800,
94
+ 2075673600,
95
+ 302702400,
96
+ 30270240,
97
+ 2162160,
98
+ 110880,
99
+ 3960,
100
+ 90,
101
+ 1,
102
+ ]
103
+
104
+ U = b[1] * np.eye(dim, dtype=dtype)
105
+ V = b[0] * np.eye(dim, dtype=dtype)
106
+
107
+ A2 = A @ A
108
+ A2n = np.eye(dim, dtype=dtype)
109
+
110
+ # evaluate (10.33)
111
+ for i in range(1, 9 // 2 + 1):
112
+ A2n = A2n @ A2
113
+ U += b[2 * i + 1] * A2n
114
+ V += b[2 * i] * A2n
115
+
116
+ U = A @ U
117
+ return np.linalg.solve(V - U, V + U)
118
+
119
+ else:
120
+
121
+ # algorithm 10.20, from line 7
122
+ dim, dim = A.shape
123
+ b = [
124
+ 64764752532480000,
125
+ 32382376266240000,
126
+ 7771770303897600,
127
+ 1187353796428800,
128
+ 129060195264000,
129
+ 10559470521600,
130
+ 670442572800,
131
+ 33522128640,
132
+ 1323241920,
133
+ 40840800,
134
+ 960960,
135
+ 16380,
136
+ 182,
137
+ 1,
138
+ ]
139
+
140
+ s = max(0, int(math.ceil(math.log(norm / theta13) / math.log(2))))
141
+ if s > 0:
142
+ A /= 2**s
143
+
144
+ Id = np.eye(dim)
145
+ A2 = A @ A
146
+ A4 = A2 @ A2
147
+ A6 = A2 @ A4
148
+
149
+ U = A @ (
150
+ (A6 @ (b[13] * A6 + b[11] * A4 + b[9] * A2))
151
+ + b[7] * A6
152
+ + b[5] * A4
153
+ + b[3] * A2
154
+ + b[1] * Id
155
+ )
156
+
157
+ V = (
158
+ A6 @ (b[12] * A6 + b[10] * A4 + b[8] * A2)
159
+ + b[6] * A6
160
+ + b[4] * A4
161
+ + b[2] * A2
162
+ + b[0] * Id
163
+ )
164
+ # added this
165
+ r13 = np.ascontiguousarray(np.linalg.solve(V - U, V + U))
166
+ return np.linalg.matrix_power(r13, 2**s)