nltools 0.6.0.dev0__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.
Files changed (95) hide show
  1. nltools/__init__.py +55 -0
  2. nltools/algorithms/__init__.py +90 -0
  3. nltools/algorithms/alignment/__init__.py +21 -0
  4. nltools/algorithms/alignment/procrustes.py +565 -0
  5. nltools/algorithms/alignment/srm.py +758 -0
  6. nltools/algorithms/backends.py +1059 -0
  7. nltools/algorithms/corrections.py +177 -0
  8. nltools/algorithms/decoding.py +327 -0
  9. nltools/algorithms/inference/__init__.py +50 -0
  10. nltools/algorithms/inference/bootstrap.py +1386 -0
  11. nltools/algorithms/inference/correlation.py +373 -0
  12. nltools/algorithms/inference/intersubject.py +422 -0
  13. nltools/algorithms/inference/isc.py +1554 -0
  14. nltools/algorithms/inference/matrix.py +602 -0
  15. nltools/algorithms/inference/one_sample.py +288 -0
  16. nltools/algorithms/inference/random.py +122 -0
  17. nltools/algorithms/inference/timeseries.py +347 -0
  18. nltools/algorithms/inference/two_sample.py +212 -0
  19. nltools/algorithms/inference/utils.py +58 -0
  20. nltools/algorithms/inference/validation.py +282 -0
  21. nltools/algorithms/neighborhoods.py +207 -0
  22. nltools/algorithms/outliers.py +308 -0
  23. nltools/algorithms/regression.py +83 -0
  24. nltools/algorithms/signal.py +303 -0
  25. nltools/algorithms/similarity.py +234 -0
  26. nltools/algorithms/validation.py +151 -0
  27. nltools/cross_validation.py +72 -0
  28. nltools/data/__init__.py +30 -0
  29. nltools/data/adjacency/__init__.py +875 -0
  30. nltools/data/adjacency/io.py +111 -0
  31. nltools/data/adjacency/modeling.py +569 -0
  32. nltools/data/adjacency/plotting.py +174 -0
  33. nltools/data/adjacency/state.py +349 -0
  34. nltools/data/adjacency/stats.py +596 -0
  35. nltools/data/adjacency/utils.py +79 -0
  36. nltools/data/atlases/__init__.py +23 -0
  37. nltools/data/atlases/labeling.py +158 -0
  38. nltools/data/atlases/loading.py +76 -0
  39. nltools/data/atlases/registry.py +96 -0
  40. nltools/data/atlases/reporting.py +456 -0
  41. nltools/data/braindata/__init__.py +2170 -0
  42. nltools/data/braindata/analysis.py +1381 -0
  43. nltools/data/braindata/bootstrap.py +398 -0
  44. nltools/data/braindata/io.py +896 -0
  45. nltools/data/braindata/modeling.py +594 -0
  46. nltools/data/braindata/plotting.py +501 -0
  47. nltools/data/braindata/prediction.py +1250 -0
  48. nltools/data/braindata/utils.py +348 -0
  49. nltools/data/braindata/validation.py +197 -0
  50. nltools/data/braindata/viewer.js +266 -0
  51. nltools/data/braindata/viewer.py +770 -0
  52. nltools/data/combine.py +27 -0
  53. nltools/data/designmatrix/__init__.py +1032 -0
  54. nltools/data/designmatrix/append.py +518 -0
  55. nltools/data/designmatrix/diagnostics.py +248 -0
  56. nltools/data/designmatrix/io.py +356 -0
  57. nltools/data/designmatrix/plotting.py +291 -0
  58. nltools/data/designmatrix/regressors.py +463 -0
  59. nltools/data/designmatrix/transforms.py +200 -0
  60. nltools/data/designmatrix/utils.py +350 -0
  61. nltools/data/ownership.py +129 -0
  62. nltools/data/results.py +291 -0
  63. nltools/data/roc/__init__.py +398 -0
  64. nltools/data/simulator/__init__.py +927 -0
  65. nltools/data/simulator/haxby.py +124 -0
  66. nltools/data/validation.py +83 -0
  67. nltools/datasets.py +218 -0
  68. nltools/io/__init__.py +10 -0
  69. nltools/io/events.py +67 -0
  70. nltools/io/h5.py +246 -0
  71. nltools/mask.py +403 -0
  72. nltools/models/__init__.py +11 -0
  73. nltools/models/glm.py +543 -0
  74. nltools/models/results.py +49 -0
  75. nltools/models/ridge.py +1303 -0
  76. nltools/models/validation.py +26 -0
  77. nltools/plotting/__init__.py +32 -0
  78. nltools/plotting/adjacency.py +421 -0
  79. nltools/plotting/brain.py +669 -0
  80. nltools/plotting/decomposition.py +111 -0
  81. nltools/plotting/prediction.py +110 -0
  82. nltools/resources/covariates_example.csv +161 -0
  83. nltools/resources/onsets_example.csv +40 -0
  84. nltools/templates/__init__.py +51 -0
  85. nltools/templates/config.py +144 -0
  86. nltools/templates/fetch.py +260 -0
  87. nltools/templates/matching.py +183 -0
  88. nltools/templates/paths.py +106 -0
  89. nltools/templates/registry.py +25 -0
  90. nltools/utils.py +230 -0
  91. nltools/version.py +13 -0
  92. nltools-0.6.0.dev0.dist-info/METADATA +95 -0
  93. nltools-0.6.0.dev0.dist-info/RECORD +95 -0
  94. nltools-0.6.0.dev0.dist-info/WHEEL +4 -0
  95. nltools-0.6.0.dev0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,111 @@
1
+ """I/O functions for Adjacency objects."""
2
+
3
+ from pathlib import Path
4
+
5
+ import numpy as np
6
+ import polars as pl
7
+
8
+ from nltools.io.h5 import _read_polars_frame, _reject_legacy_h5, _require_h5
9
+
10
+
11
+ def _write(adj, file_name, method="long"):
12
+ """Write an Adjacency to a `.csv` or `.h5` file.
13
+
14
+ HDF5 is the round-trip format: it stores the matrix values, the matrix kind,
15
+ the node labels, and `Y`. CSV stores values only. Reading a CSV back
16
+ therefore loses the node labels, `Y`, and the matrix kind, and needs an
17
+ explicit `matrix_type` wherever the flat layout is ambiguous. Square CSV
18
+ output is single-matrix only.
19
+
20
+ Args:
21
+ adj (Adjacency): Adjacency object to write.
22
+ file_name (str | Path): Output path; an `.h5`/`.hdf5` suffix writes HDF5.
23
+ method (str): Layout for CSV output, `'long'` (vectorized rows) or `'square'`
24
+ (single matrix only).
25
+ """
26
+ from nltools.io.h5 import _is_h5_path, _to_h5
27
+
28
+ if method not in ["long", "square"]:
29
+ raise ValueError('Make sure method is ["long","square"].')
30
+
31
+ if isinstance(file_name, Path):
32
+ file_name = str(file_name)
33
+
34
+ if _is_h5_path(file_name):
35
+ if method == "square":
36
+ raise NotImplementedError('Saving as hdf5 does not support method="square"')
37
+ _to_h5(adj, file_name, obj_type="adjacency")
38
+ else:
39
+ if method == "long":
40
+ _write_2d_csv(adj.data, file_name)
41
+ elif adj.is_single_matrix and method == "square":
42
+ _write_2d_csv(np.asarray(adj.squareform()), file_name)
43
+ elif not adj.is_single_matrix and method == "square":
44
+ raise NotImplementedError(
45
+ "Need to decide how we should write out multiple matrices. As separate files?"
46
+ )
47
+
48
+
49
+ def _write_2d_csv(arr: np.ndarray, file_name: str) -> None:
50
+ """Write an array to CSV via polars with numeric-index column names.
51
+
52
+ Matches the pandas ``pd.DataFrame(arr).to_csv(index=None)`` layout:
53
+ 1-D arrays become a single-column frame with ``len(arr)`` rows.
54
+ """
55
+ arr = np.asarray(arr)
56
+ if arr.ndim == 1:
57
+ arr = arr.reshape(-1, 1)
58
+ schema = [str(i) for i in range(arr.shape[1])]
59
+ pl.DataFrame(arr, schema=schema).write_csv(file_name)
60
+
61
+
62
+ def _to_graph(adj):
63
+ """Convert Adjacency into networkx graph.
64
+
65
+ Only works on single matrices for now.
66
+
67
+ Args:
68
+ adj (Adjacency): Adjacency instance (must be a single matrix).
69
+
70
+ Returns:
71
+ networkx.Graph or networkx.DiGraph: Graph representation of the
72
+ adjacency matrix. Uses DiGraph for directed matrices.
73
+ """
74
+
75
+ import networkx as nx
76
+
77
+ if adj.is_single_matrix:
78
+ if adj.matrix_type == "directed":
79
+ G = nx.DiGraph(adj.squareform())
80
+ else:
81
+ G = nx.Graph(adj.squareform())
82
+ if adj.labels:
83
+ labels = dict(zip(G.nodes, adj.labels))
84
+ nx.relabel_nodes(G, labels, copy=False)
85
+ return G
86
+ raise NotImplementedError("This function currently only works on single matrices.")
87
+
88
+
89
+ def _read_h5(file_name):
90
+ """Read the current vector layout into a normalized Adjacency."""
91
+ from . import Adjacency
92
+
93
+ _require_h5()
94
+ import h5py
95
+
96
+ with h5py.File(file_name, "r") as source:
97
+ _reject_legacy_h5(source, "Y_columns")
98
+ kind = source["matrix_type"][()].decode()
99
+ values = np.array(source["data"])
100
+ labels_ds = source["labels"]
101
+ labels = (
102
+ labels_ds.asstr()[()].tolist()
103
+ if h5py.check_string_dtype(labels_ds.dtype) is not None
104
+ else labels_ds[()].tolist()
105
+ )
106
+ return Adjacency(
107
+ None if kind == "empty" else values,
108
+ matrix_type=None if kind == "empty" else kind + "_flat",
109
+ labels=labels,
110
+ Y=_read_polars_frame(source, "Y"),
111
+ )
@@ -0,0 +1,569 @@
1
+ """Provide standalone modeling and inference functions for Adjacency matrices.
2
+
3
+ Each function takes an Adjacency instance as its first argument (`adj`).
4
+ """
5
+
6
+ import numpy as np
7
+
8
+
9
+ def _bootstrap(
10
+ adj,
11
+ statistic,
12
+ *,
13
+ n_samples=5000,
14
+ confidence_level=0.95,
15
+ return_samples=False,
16
+ n_jobs=-1,
17
+ random_state=None,
18
+ progress_bar=False,
19
+ ):
20
+ """Bootstrap an aggregate statistic across a stack of matrices.
21
+
22
+ Resamples matrices with replacement and aggregates the replicates as they
23
+ complete, so what the run holds is the retained tail — about
24
+ `(1 - confidence_level)` of the replicates per edge — plus one dispatch
25
+ window, rather than all `n_samples` matrices.
26
+
27
+ Args:
28
+ adj (Adjacency): Adjacency instance containing multiple matrices.
29
+ statistic (str): Statistic to bootstrap: `'mean'`, `'median'`, `'std'`,
30
+ `'sum'`, `'min'`, or `'max'` — each the corresponding NumPy
31
+ reduction over matrices, with `'std'` at `ddof=0`.
32
+ n_samples (int): Number of bootstrap replicates, at least two. Default
33
+ 5000.
34
+ confidence_level (float): Confidence level of the reported interval,
35
+ strictly between zero and one. Default 0.95.
36
+ return_samples (bool): Retain and return every replicate. Default
37
+ False.
38
+ n_jobs (int): CPU worker ceiling. -1 (default) means all cores.
39
+ random_state (int | None): Random seed for reproducibility.
40
+ progress_bar (bool): If True, show a progress bar. Default False.
41
+
42
+ Returns:
43
+ BootstrapResult: `estimate`, `standard_error`, `ci_lower` and
44
+ `ci_upper` as single-matrix `Adjacency` objects, plus `samples` as
45
+ a NumPy array with the bootstrap axis first when
46
+ `return_samples=True`.
47
+
48
+ Raises:
49
+ ValueError: If `statistic` is unknown, an argument is out of range, or
50
+ the retained output cannot fit the measured memory budget.
51
+
52
+ Examples:
53
+ ```python
54
+ boot = bootstrap(adj, "mean", n_samples=1000)
55
+ boot.estimate # → Adjacency
56
+ ```
57
+ """
58
+ from nltools.algorithms.inference.bootstrap import (
59
+ _bootstrap_simple_cpu_parallel,
60
+ )
61
+
62
+ SIMPLE_STATS = ["mean", "median", "std", "sum", "min", "max"]
63
+ if statistic not in SIMPLE_STATS:
64
+ raise ValueError(
65
+ f"Unsupported statistic '{statistic}'. "
66
+ f"Supported basic statistics: {SIMPLE_STATS}."
67
+ )
68
+
69
+ # Adjacency.data shape: (n_matrices, n_edges)
70
+ result = _bootstrap_simple_cpu_parallel(
71
+ adj.data,
72
+ method=statistic,
73
+ n_samples=n_samples,
74
+ confidence_level=confidence_level,
75
+ return_samples=return_samples,
76
+ n_jobs=n_jobs,
77
+ random_state=random_state,
78
+ progress_bar=progress_bar,
79
+ )
80
+
81
+ return _convert_bootstrap_results_to_adjacency(adj, result)
82
+
83
+
84
+ def _convert_bootstrap_results_to_adjacency(adj, result):
85
+ """Wrap an engine's arrays as a `BootstrapResult` of single-matrix `Adjacency`.
86
+
87
+ Args:
88
+ adj (Adjacency): Instance supplying matrix kind and node labels.
89
+ result (dict): Engine output with `'estimate'`, `'standard_error'`,
90
+ `'ci_lower'`, `'ci_upper'`, and optionally `'samples'`.
91
+
92
+ Returns:
93
+ BootstrapResult: The four summaries as `Adjacency`, and the retained
94
+ replicates as a NumPy array when present.
95
+ """
96
+ import polars as pl
97
+
98
+ from nltools.data.results import BootstrapResult
99
+
100
+ from .state import _common_labels, _result as adjacency_result
101
+
102
+ labels = _common_labels(adj)
103
+
104
+ def _map(values):
105
+ return adjacency_result(
106
+ adj,
107
+ np.asarray(values).reshape(-1),
108
+ labels=labels,
109
+ Y=pl.DataFrame(),
110
+ )
111
+
112
+ return BootstrapResult(
113
+ estimate=_map(result["estimate"]),
114
+ standard_error=_map(result["standard_error"]),
115
+ ci_lower=_map(result["ci_lower"]),
116
+ ci_upper=_map(result["ci_upper"]),
117
+ samples=result.get("samples"),
118
+ )
119
+
120
+
121
+ def _regress(adj, X, *, tail=2):
122
+ """Run a regression on an adjacency instance.
123
+
124
+ Pass an `Adjacency` as `X` to decompose `adj` with other matrices, or a
125
+ `DesignMatrix` to regress each cell across a stack of matrices.
126
+
127
+ Args:
128
+ adj (Adjacency): Adjacency instance.
129
+ X (Adjacency | DesignMatrix): Design matrix.
130
+ tail (int | str): `2`/`'two'` for two-tailed (default); `1`/`'one'` for
131
+ one-tailed (beta > 0; negate a regressor for the other direction).
132
+
133
+ Returns:
134
+ dict: Coefficient fields `beta`, `sigma` (coefficient standard error),
135
+ `t`, and `p` are predictor Adjacency maps for DesignMatrix input and
136
+ native predictor arrays/scalars for Adjacency input. `df` is scalar;
137
+ `residual` is an Adjacency retaining the response shape and metadata.
138
+ """
139
+ import polars as pl
140
+ from nltools.algorithms.regression import regress as ols_regress
141
+ from nltools.data.adjacency import Adjacency
142
+ from nltools.data.designmatrix import DesignMatrix
143
+ from nltools.algorithms.validation import _validate_tail_parameter
144
+ from .state import _common_labels, _result, _validate_compatible
145
+
146
+ _validate_tail_parameter(tail)
147
+ if isinstance(X, Adjacency):
148
+ if not adj.is_single_matrix:
149
+ raise ValueError("Adjacency predictors require a single response matrix.")
150
+ _validate_compatible(adj, X)
151
+ response_labels = _common_labels(adj)
152
+ predictor_labels = _common_labels(X)
153
+ if response_labels != predictor_labels or X.labels and not predictor_labels:
154
+ raise ValueError("Predictor and response node ordering must match.")
155
+ design = np.atleast_2d(X.data).T
156
+ response = adj.data[:, None]
157
+ elif isinstance(X, DesignMatrix):
158
+ if X.shape[0] != len(adj):
159
+ raise ValueError(
160
+ "Design matrix must have same number of observations as Adjacency"
161
+ )
162
+ design = X.to_numpy()
163
+ response = np.atleast_2d(adj.data)
164
+ else:
165
+ raise ValueError("X must be a DesignMatrix or Adjacency Instance.")
166
+
167
+ # The shared OLS is the single implementation; it squeezes every output, so
168
+ # restore the (n_regressors, n_targets) and (n_samples, n_targets) shapes the
169
+ # result assembly below indexes by axis.
170
+ beta, stderr, t, p, _, residual = ols_regress(design, response, tail=tail)
171
+ coefficient_shape = (design.shape[1], response.shape[1])
172
+ beta, stderr, t, p = (
173
+ np.reshape(value, coefficient_shape) for value in (beta, stderr, t, p)
174
+ )
175
+ residual = np.reshape(residual, (design.shape[0], response.shape[1]))
176
+ df = int(design.shape[0] - design.shape[1])
177
+ stats = {"df": df}
178
+ for key, values in [("beta", beta), ("sigma", stderr), ("t", t), ("p", p)]:
179
+ if isinstance(X, Adjacency):
180
+ stats[key] = values[:, 0].copy() if len(values) > 1 else values[0, 0].item()
181
+ else:
182
+ stats[key] = _result(
183
+ adj,
184
+ values[0] if len(values) == 1 else values,
185
+ labels=_common_labels(adj),
186
+ Y=pl.DataFrame(),
187
+ )
188
+ residual_values = (
189
+ residual[:, 0]
190
+ if isinstance(X, Adjacency)
191
+ else residual[0]
192
+ if adj.is_single_matrix
193
+ else residual
194
+ )
195
+ stats["residual"] = _result(adj, residual_values, labels=adj.labels, Y=adj.Y)
196
+ return stats
197
+
198
+
199
+ def _social_relations_model(adj, summarize_results=True, nan_replace=True):
200
+ """Estimate the social relations model from a matrix for a round-robin design.
201
+
202
+ $$X_{ij} = m + \\alpha_i + \\beta_j + g_{ij} + \\epsilon_{ijl}$$
203
+
204
+ where $X_{ij}$ is the score for person i rating person j, $m$ is the group mean,
205
+ $\\alpha_i$ is person i's actor effect, $\\beta_j$ is person j's partner effect, $g_{ij}$
206
+ is the relationship effect and $\\epsilon_{ijl}$ is the error in measure l for actor i and partner j.
207
+
208
+ This model is primarily concerned with partitioning the variance of the various
209
+ effects. The implementation follows Chapter 8 of Kenny, Kashy, & Cook (2006) and
210
+ the tests replicate the book's examples. Actor scores are rows (lower triangle)
211
+ and partner scores are columns (upper triangle). The minimal sample size to
212
+ estimate these effects is 4.
213
+
214
+ **Model assumptions:** social interactions are exclusively dyadic; people are
215
+ randomly sampled from the population; there are no order effects; the effects
216
+ combine additively and relationships are linear.
217
+
218
+ Args:
219
+ adj (Adjacency): A single matrix, or one matrix per group.
220
+ summarize_results (bool): If True, print a formatted summary of model results.
221
+ nan_replace (bool): If True, replace NaN values with row and column means.
222
+
223
+ Returns:
224
+ pd.Series | pd.DataFrame: All of the effects estimated using SRM (a Series
225
+ for a single matrix, a DataFrame with one row per matrix otherwise).
226
+
227
+ References:
228
+ Kenny, D. A., Kashy, D. A., & Cook, W. L. (2006). *Dyadic data analysis*.
229
+ Guilford Press.
230
+ """
231
+ import pandas as pd
232
+
233
+ from nltools.data.adjacency import Adjacency
234
+ from scipy.spatial.distance import squareform
235
+ import scipy.stats as scipy_stats
236
+
237
+ def mean_square_between(x1, x2=None, df="standard"):
238
+ """Calculate between-dyad variance."""
239
+
240
+ if df == "standard":
241
+ n = len(x1)
242
+ df = n - 1
243
+ elif df == "relationship":
244
+ n = len(squareform(x1))
245
+ df = ((n - 1) * (n - 2) / 2) - 1
246
+ else:
247
+ raise ValueError("df can only be ['standard', 'relationship']")
248
+ if x2 is not None:
249
+ return (
250
+ 2 * np.nansum((((x1 + x2) / 2) - np.nanmean((x1 + x2) / 2)) ** 2) / df
251
+ )
252
+ return np.nansum((x1 - np.nanmean(x1)) ** 2) / df
253
+
254
+ def mean_square_within(x1, x2, df="standard"):
255
+ """Calculate within-dyad variance."""
256
+
257
+ if df == "standard":
258
+ n = len(x1)
259
+ df = n
260
+ elif df == "relationship":
261
+ n = len(squareform(x1))
262
+ df = (n - 1) * (n - 2) / 2
263
+ else:
264
+ raise ValueError("df can only be ['standard', 'relationship']")
265
+ return np.nansum((x1 - x2) ** 2) / (2 * df)
266
+
267
+ def estimate_person_effect(n, x1_mean, x2_mean, grand_mean):
268
+ """Calculate actor, partner, and relationship effects."""
269
+ return (
270
+ ((n - 1) ** 2 / (n * (n - 2))) * x1_mean
271
+ + ((n - 1) / (n * (n - 2))) * x2_mean
272
+ - ((n - 1) / (n - 2)) * grand_mean
273
+ )
274
+
275
+ def estimate_person_variance(x, ms_b, ms_w):
276
+ """Calculate variance for a specific dyad member, such as actor or partner."""
277
+ n = len(x)
278
+ return mean_square_between(x) - (ms_b / (2 * (n - 2))) - (ms_w / (2 * n))
279
+
280
+ def estimate_srm(data):
281
+ """Estimate a Social Relations Model from a single matrix."""
282
+
283
+ if not data.is_single_matrix:
284
+ raise ValueError(
285
+ "This function only operates on single matrix Adjacency instances."
286
+ )
287
+
288
+ n = data.n_nodes
289
+ if n < 4:
290
+ raise ValueError(
291
+ "The Social Relations Model cannot be estimated when sample size is less than 4."
292
+ )
293
+ grand_mean = data.mean()
294
+ dat = data.squareform().copy()
295
+ np.fill_diagonal(dat, np.nan)
296
+ actor_mean = np.nanmean(dat, axis=1)
297
+ partner_mean = np.nanmean(dat, axis=0)
298
+
299
+ a = estimate_person_effect(
300
+ n, actor_mean, partner_mean, grand_mean
301
+ ) # Actor effects
302
+ b = estimate_person_effect(
303
+ n, partner_mean, actor_mean, grand_mean
304
+ ) # Partner effects
305
+
306
+ # Relationship effects
307
+ g = np.ones(dat.shape) * np.nan
308
+ for i in range(n):
309
+ for j in range(n):
310
+ if i != j:
311
+ g[i, j] = dat[i, j] - a[i] - b[j] - grand_mean
312
+
313
+ # Estimate Variance
314
+ x1 = g[np.tril_indices(n, k=-1)]
315
+ x2 = g[np.triu_indices(n, k=1)]
316
+ ms_b = mean_square_between(x1, x2, df="relationship")
317
+ ms_w = mean_square_within(x1, x2, df="relationship")
318
+ actor_variance = estimate_person_variance(a, ms_b, ms_w)
319
+ partner_variance = estimate_person_variance(b, ms_b, ms_w)
320
+ relationship_variance = (ms_b + ms_w) / 2
321
+ dyadic_reciprocity_covariance = (ms_b - ms_w) / 2
322
+ dyadic_reciprocity_correlation = (ms_b - ms_w) / (ms_b + ms_w)
323
+ actor_partner_covariance = (
324
+ (np.sum(a * b) / (n - 1)) - (ms_b / (2 * (n - 2))) + (ms_w / (2 * n))
325
+ )
326
+ actor_partner_correlation = actor_partner_covariance / (
327
+ np.sqrt(actor_variance * partner_variance)
328
+ )
329
+ actor_reliability = actor_variance / (
330
+ actor_variance
331
+ + (relationship_variance / (n - 1))
332
+ - (dyadic_reciprocity_covariance / ((n - 1) ** 2))
333
+ )
334
+ partner_reliability = partner_variance / (
335
+ partner_variance
336
+ + (relationship_variance / (n - 1))
337
+ - (dyadic_reciprocity_covariance / ((n - 1) ** 2))
338
+ )
339
+ adjusted_dyadic_reciprocity_correlation = actor_partner_correlation * np.sqrt(
340
+ actor_reliability * partner_reliability
341
+ )
342
+ total_variance = actor_variance + partner_variance + relationship_variance
343
+
344
+ return pd.Series(
345
+ {
346
+ "grand_mean": grand_mean,
347
+ "actor_effect": a,
348
+ "partner_effect": b,
349
+ "relationship_effect": g,
350
+ "actor_variance": actor_variance,
351
+ "partner_variance": partner_variance,
352
+ "relationship_variance": relationship_variance,
353
+ "actor_partner_covariance": actor_partner_covariance,
354
+ "actor_partner_correlation": actor_partner_correlation,
355
+ "dyadic_reciprocity_covariance": dyadic_reciprocity_covariance,
356
+ "dyadic_reciprocity_correlation": dyadic_reciprocity_correlation,
357
+ "adjusted_dyadic_reciprocity_correlation": adjusted_dyadic_reciprocity_correlation,
358
+ "actor_reliability": actor_reliability,
359
+ "partner_reliability": partner_reliability,
360
+ "total_variance": total_variance,
361
+ }
362
+ )
363
+
364
+ def summarize_srm_results(results):
365
+ """Summarize Social Relations Model results."""
366
+
367
+ def estimate_srm_stats(results, var_name, tailed=1):
368
+ """Compute mean estimate, standard error, t-statistic, and p-value for an SRM variance component.
369
+
370
+ Args:
371
+ results: DataFrame of SRM results across groups, or Series for a single group.
372
+ var_name: Name of the variance component column to summarize.
373
+ tailed: Number of tails for the t-test (1 or 2).
374
+
375
+ Returns:
376
+ Tuple of (estimate, standardized, se, t, p).
377
+ """
378
+ estimate = results[var_name].mean()
379
+ standardized = (results[var_name] / results["total_variance"]).mean()
380
+ se = results[var_name].std() / np.sqrt(len(results[var_name]))
381
+ with np.errstate(invalid="ignore", divide="ignore"):
382
+ t = estimate / se
383
+ if tailed == 1:
384
+ p = 1 - scipy_stats.t.cdf(t, len(results[var_name]) - 1)
385
+ elif tailed == 2:
386
+ p = 2 * (1 - scipy_stats.t.cdf(t, len(results[var_name]) - 1))
387
+ else:
388
+ raise ValueError("tailed can only be [1,2]")
389
+ return (estimate, standardized, se, t, p)
390
+
391
+ def print_srm_stats(results, var_name, tailed=1):
392
+ """Print a formatted summary row for an SRM variance component across multiple groups.
393
+
394
+ Args:
395
+ results: DataFrame of SRM results across groups.
396
+ var_name: Name of the variance component column to print.
397
+ tailed: Number of tails for the t-test (1 or 2).
398
+ """
399
+ estimate, standardized, se, t, p = estimate_srm_stats(
400
+ results, var_name, tailed
401
+ )
402
+ print(
403
+ f"{var_name:<40} {estimate:^10.2f}{standardized:^10.2f} {se:^10.2f} {t:^10.2f} {p:^10.4f}"
404
+ )
405
+
406
+ def print_single_group_srm_stats(results, var_name):
407
+ """Print a formatted summary row for an SRM variance component for a single group.
408
+
409
+ Inference statistics (se, t, p) are printed as NaN since they require multiple groups.
410
+
411
+ Args:
412
+ results: Series of SRM results for a single group.
413
+ var_name: Name of the variance component to print.
414
+ """
415
+ estimate = results[var_name].mean()
416
+ standardized = (results[var_name] / results["total_variance"]).mean()
417
+ print(
418
+ f"{var_name:<40} {estimate:^10.2f}{standardized:^10.2f} {np.nan:^10.2f} {np.nan:^10.2f} {np.nan:^10.4f}"
419
+ )
420
+
421
+ def print_srm_covariances(results, var_name):
422
+ """Print a formatted summary row for an SRM covariance component across multiple groups.
423
+
424
+ Uses the covariance estimate for inference and correlation as the standardized effect size.
425
+
426
+ Args:
427
+ results: DataFrame of SRM results across groups.
428
+ var_name: Name of the covariance component (without '_covariance' or '_correlation' suffix).
429
+ """
430
+ estimate, _, se, t, p = estimate_srm_stats(
431
+ results, f"{var_name}_covariance", tailed=2
432
+ )
433
+ standardized = results[f"{var_name}_correlation"].mean()
434
+ print(
435
+ f"{var_name:<40} {estimate:^10.2f}{standardized:^10.2f} {se:^10.2f} {t:^10.2f} {p:^10.4f}"
436
+ )
437
+
438
+ def print_single_srm_covariances(results, var_name):
439
+ """Print a formatted summary row for an SRM covariance component for a single group.
440
+
441
+ Inference statistics (se, t, p) are printed as NaN since they require multiple groups.
442
+
443
+ Args:
444
+ results: Series of SRM results for a single group.
445
+ var_name: Name of the covariance component (without '_covariance' or '_correlation' suffix).
446
+ """
447
+ estimate = results[f"{var_name}_covariance"].mean()
448
+ standardized = results[f"{var_name}_correlation"].mean()
449
+ print(
450
+ f"{var_name:<40} {estimate:^10.2f}{standardized:^10.2f} {np.nan:^10.2f} {np.nan:^10.2f} {np.nan:^10.4f}"
451
+ )
452
+
453
+ if isinstance(results, pd.Series):
454
+ n_groups = 1
455
+ group_size = results["actor_effect"].shape[0]
456
+ elif isinstance(results, pd.DataFrame):
457
+ n_groups = len(results)
458
+ group_size = np.mean([x.shape for x in results["actor_effect"]])
459
+
460
+ print("Social Relations Model: Results")
461
+ print("\n")
462
+ print(f"Number of Groups: {n_groups:<20}")
463
+ print(f"Average Group Size: {group_size:<20}")
464
+ print("\n")
465
+ print(
466
+ f"{'':<40} {'Estimate':<10} {'Standardized':<10} {'se':<10} {'t':<10} {'p':<10}"
467
+ )
468
+ if isinstance(results, pd.Series):
469
+ print_single_group_srm_stats(results, "actor_variance")
470
+ print_single_group_srm_stats(results, "partner_variance")
471
+ print_single_group_srm_stats(results, "relationship_variance")
472
+ print_single_srm_covariances(results, "actor_partner")
473
+ print_single_srm_covariances(results, "dyadic_reciprocity")
474
+ elif isinstance(results, pd.DataFrame):
475
+ print_srm_stats(results, "actor_variance")
476
+ print_srm_stats(results, "partner_variance")
477
+ print_srm_stats(results, "relationship_variance")
478
+ print_srm_covariances(results, "actor_partner")
479
+ print_srm_covariances(results, "dyadic_reciprocity")
480
+ print("\n")
481
+ print(f"{'Actor Reliability':<20} {results['actor_reliability'].mean():^20.2f}")
482
+ print(
483
+ f"{'Partner Reliability':<20} {results['partner_reliability'].mean():^20.2f}"
484
+ )
485
+ print("\n")
486
+
487
+ def replace_missing(data):
488
+ """Replace missing data with row and column means and return missing coordinates."""
489
+
490
+ def fix_missing(data):
491
+ """Replace NaN off-diagonal entries with the mean of their row and column.
492
+
493
+ Args:
494
+ data: Adjacency matrix with possible NaN values.
495
+
496
+ Returns:
497
+ Tuple of (Adjacency with NaNs replaced, (row, col) coordinates of replaced values).
498
+ """
499
+ X = data.squareform().copy()
500
+ x, y = np.where(np.isnan(X))
501
+ for i, j in zip(x, y):
502
+ if i != j:
503
+ X[i, j] = (np.nanmean(X[i, :]) + np.nanmean(X[:, j])) / 2
504
+ X = Adjacency(X, matrix_type=data.matrix_type)
505
+ return (X, (x, y))
506
+
507
+ if data.is_single_matrix:
508
+ X, coord = fix_missing(data)
509
+ else:
510
+ X = []
511
+ coord = []
512
+ for d in data:
513
+ m, c = fix_missing(d)
514
+ X.append(m)
515
+ coord.append(c)
516
+ X = Adjacency(X)
517
+ return (X, coord)
518
+
519
+ if nan_replace:
520
+ data, _ = replace_missing(adj)
521
+ else:
522
+ data = adj.copy()
523
+
524
+ if adj.is_single_matrix:
525
+ results = estimate_srm(data)
526
+ else:
527
+ results = pd.DataFrame([estimate_srm(x) for x in data])
528
+
529
+ if summarize_results:
530
+ summarize_srm_results(results)
531
+
532
+ return results
533
+
534
+
535
+ def _generate_permutations(adj, n_permute, random_state=None):
536
+ """Generate permuted versions of an Adjacency instance lazily.
537
+
538
+ This is useful for iterative comparisons.
539
+
540
+ Args:
541
+ adj (Adjacency): Adjacency instance.
542
+ n_permute (int): Number of permutations.
543
+ random_state (int | np.random.RandomState, optional): Random seed for
544
+ reproducibility. Defaults to None.
545
+
546
+ Yields:
547
+ Adjacency: Permuted version of `adj`.
548
+
549
+ Examples:
550
+ ```python
551
+ for perm in generate_permutations(adj, 1000):
552
+ out = neural_distance_mat.similarity(perm)
553
+ ```
554
+ """
555
+ from nltools.data.adjacency import Adjacency
556
+ from sklearn.utils import check_random_state
557
+
558
+ random_state = check_random_state(random_state)
559
+
560
+ for _ in range(n_permute):
561
+ # Get squareform as numpy array (no pandas conversion needed)
562
+ dat = adj.squareform()
563
+ # Generate random permutation indices
564
+ permuted_idx = random_state.choice(
565
+ dat.shape[0], size=dat.shape[0], replace=False
566
+ )
567
+ # Permute rows and columns using numpy advanced indexing (faster than pandas)
568
+ dat = dat[np.ix_(permuted_idx, permuted_idx)]
569
+ yield Adjacency(dat)