dmx-learn 1.0.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.
Files changed (106) hide show
  1. dmx/__init__.py +3 -0
  2. dmx/arithmetic.py +28 -0
  3. dmx/bexamples/dpm_auto_example1.py +19 -0
  4. dmx/bexamples/dpm_auto_example2.py +39 -0
  5. dmx/bstats/__init__.py +491 -0
  6. dmx/bstats/bernoulli.py +253 -0
  7. dmx/bstats/bestimation.py +267 -0
  8. dmx/bstats/beta.py +69 -0
  9. dmx/bstats/catdirichlet.py +84 -0
  10. dmx/bstats/categorical.py +293 -0
  11. dmx/bstats/composite.py +294 -0
  12. dmx/bstats/conditional.py +298 -0
  13. dmx/bstats/dirac.py +134 -0
  14. dmx/bstats/dirichlet.py +335 -0
  15. dmx/bstats/dmvn.py +278 -0
  16. dmx/bstats/dpm.py +428 -0
  17. dmx/bstats/exponential.py +209 -0
  18. dmx/bstats/gamma.py +172 -0
  19. dmx/bstats/gaussian.py +292 -0
  20. dmx/bstats/geometric.py +259 -0
  21. dmx/bstats/ignored.py +122 -0
  22. dmx/bstats/intrange.py +329 -0
  23. dmx/bstats/mixture.py +382 -0
  24. dmx/bstats/mvngamma.py +108 -0
  25. dmx/bstats/normgamma.py +96 -0
  26. dmx/bstats/nulldist.py +154 -0
  27. dmx/bstats/optional.py +287 -0
  28. dmx/bstats/pdist.py +194 -0
  29. dmx/bstats/poisson.py +259 -0
  30. dmx/bstats/sequence.py +340 -0
  31. dmx/bstats/setdist.py +230 -0
  32. dmx/bstats/symdirichlet.py +49 -0
  33. dmx/mpi4py/bstats/__init__.py +285 -0
  34. dmx/mpi4py/stats/__init__.py +309 -0
  35. dmx/mpi4py/utils/automatic.py +69 -0
  36. dmx/mpi4py/utils/bestimation.py +177 -0
  37. dmx/mpi4py/utils/estimation.py +296 -0
  38. dmx/mpi4py/utils/humap.py +87 -0
  39. dmx/mpi4py/utils/optsutil.py +16 -0
  40. dmx/stats/__init__.py +692 -0
  41. dmx/stats/binomial.py +608 -0
  42. dmx/stats/categorical.py +515 -0
  43. dmx/stats/catmultinomial.py +731 -0
  44. dmx/stats/composite.py +569 -0
  45. dmx/stats/conditional.py +839 -0
  46. dmx/stats/dirac_length.py +803 -0
  47. dmx/stats/dirichlet.py +711 -0
  48. dmx/stats/dmvn.py +606 -0
  49. dmx/stats/dmvn_mixture.py +887 -0
  50. dmx/stats/exponential.py +454 -0
  51. dmx/stats/gamma.py +511 -0
  52. dmx/stats/gaussian.py +504 -0
  53. dmx/stats/geometric.py +455 -0
  54. dmx/stats/gmm.py +780 -0
  55. dmx/stats/heterogeneous_mixture.py +817 -0
  56. dmx/stats/hidden_association.py +424 -0
  57. dmx/stats/hidden_markov.py +1859 -0
  58. dmx/stats/hmixture.py +785 -0
  59. dmx/stats/icltree.py +407 -0
  60. dmx/stats/ignored.py +300 -0
  61. dmx/stats/int_edit_setdist.py +560 -0
  62. dmx/stats/int_edit_stepsetdist.py +470 -0
  63. dmx/stats/int_hidden_association.py +886 -0
  64. dmx/stats/int_markovchain.py +707 -0
  65. dmx/stats/int_plsi.py +870 -0
  66. dmx/stats/int_spike.py +483 -0
  67. dmx/stats/intmultinomial.py +639 -0
  68. dmx/stats/intrange.py +495 -0
  69. dmx/stats/intsetdist.py +404 -0
  70. dmx/stats/jmixture.py +649 -0
  71. dmx/stats/lda.py +822 -0
  72. dmx/stats/log_gaussian.py +357 -0
  73. dmx/stats/look_back_hmm.py +874 -0
  74. dmx/stats/markovchain.py +826 -0
  75. dmx/stats/mixture.py +696 -0
  76. dmx/stats/mvn.py +404 -0
  77. dmx/stats/null_dist.py +297 -0
  78. dmx/stats/optional.py +518 -0
  79. dmx/stats/pdist.py +435 -0
  80. dmx/stats/poisson.py +336 -0
  81. dmx/stats/rdd_sampler.py +99 -0
  82. dmx/stats/select.py +224 -0
  83. dmx/stats/sequence.py +624 -0
  84. dmx/stats/setdist.py +396 -0
  85. dmx/stats/sparse_markov_transform.py +606 -0
  86. dmx/stats/spearman_rho.py +312 -0
  87. dmx/stats/ss_mixture.py +544 -0
  88. dmx/stats/tree_hmm.py +1602 -0
  89. dmx/stats/vmf.py +498 -0
  90. dmx/stats/weighted.py +247 -0
  91. dmx/utils/__init__.py +1 -0
  92. dmx/utils/automatic.py +436 -0
  93. dmx/utils/builder.py +83 -0
  94. dmx/utils/estimation.py +459 -0
  95. dmx/utils/htsne.py +496 -0
  96. dmx/utils/humap.py +69 -0
  97. dmx/utils/metrics.py +173 -0
  98. dmx/utils/optsutil.py +242 -0
  99. dmx/utils/pvalues.py +102 -0
  100. dmx/utils/special.py +148 -0
  101. dmx/utils/vector.py +688 -0
  102. dmx_learn-1.0.0.dist-info/METADATA +81 -0
  103. dmx_learn-1.0.0.dist-info/RECORD +106 -0
  104. dmx_learn-1.0.0.dist-info/WHEEL +5 -0
  105. dmx_learn-1.0.0.dist-info/licenses/LICENSE +27 -0
  106. dmx_learn-1.0.0.dist-info/top_level.txt +1 -0
dmx/__init__.py ADDED
@@ -0,0 +1,3 @@
1
+ __all__ = ['stats', 'utils','src']
2
+
3
+
dmx/arithmetic.py ADDED
@@ -0,0 +1,28 @@
1
+ """
2
+ This module defines mathematical constants and imports commonly used functions from NumPy.
3
+
4
+ The constants and functions provided here can be used for various mathematical operations
5
+ such as logarithms, exponentiation, and calculations involving pi, square roots, or infinity.
6
+
7
+ """
8
+
9
+ from numpy import (
10
+ log, # Natural logarithm
11
+ exp, # Exponential function
12
+ pi, # Mathematical constant π
13
+ sqrt, # Square root function
14
+ abs, # Absolute value function
15
+ dot, # Dot product of two arrays
16
+ isnan, # Check for NaN values
17
+ isinf # Check for infinite values
18
+ )
19
+
20
+ # Constants
21
+ maxint = 2**31 - 1 # Maximum value for a signed 32-bit integer
22
+ maxrandint = 2**31 - 1 # Maximum random integer value for signed 32-bit range
23
+ one = 1.0 # Floating-point representation of 1
24
+ zero = 0.0 # Floating-point representation of 0
25
+ two = 2.0 # Floating-point representation of 2
26
+ half = 0.5 # Floating-point representation of 0.5
27
+ inf = float('inf') # Floating-point representation of infinity
28
+ eps = 1.0e-8 # Small value for numerical precision
@@ -0,0 +1,19 @@
1
+ """Fitting a DPM with an automatic estimator determined from the data."""
2
+ from dmx.utils.automatic import get_dpm_mixture, get_estimator
3
+ from dmx.bstats import *
4
+ import numpy as np
5
+
6
+
7
+ if __name__ == '__main__':
8
+
9
+ d1 = DiagonalGaussianDistribution([-1, -1, -1], [5, 5, 5])
10
+ d2 = DiagonalGaussianDistribution([0, 0, 0], [0.1, 0.1, 0.1])
11
+ d3 = DiagonalGaussianDistribution([2, 2, 2], [1, 1, 1])
12
+ d4 = DiagonalGaussianDistribution([4, 4, 4], [1, 1, 1])
13
+ dist1 = MixtureDistribution([d1, d2, d3, d4], [0.3, 0.3, 0.2, 0.2])
14
+
15
+ data = dist1.sampler(seed=1).sample(1000)
16
+ est = get_estimator(data, use_bstats=True)
17
+ model = get_dpm_mixture(data, rng=np.random.RandomState(1))
18
+
19
+ print(str(model))
@@ -0,0 +1,39 @@
1
+ """Fitting a DPM with an automatic estimator determined from the data."""
2
+ from dmx.utils.automatic import get_dpm_mixture, get_estimator
3
+ from dmx.bstats import *
4
+ import numpy as np
5
+
6
+
7
+ if __name__ == '__main__':
8
+
9
+ rng = np.random.RandomState(2)
10
+ m = 10
11
+ n = 7
12
+ cc = 2.0
13
+ ss = 1.5
14
+ components = []
15
+ w = np.zeros(m)
16
+ pvec = np.ones(m) * 0.5
17
+ for i in range(m):
18
+ len_dist = IntegerCategoricalDistribution([0.2, 0.3, 0.3, 0.2], min_index=3)
19
+ dist1 = GaussianDistribution((i + 1) * ss, 1)
20
+ dist2 = IntegerCategoricalDistribution((np.eye(m)[i, :] + cc) / (m * cc + 1))
21
+ dist3 = CategoricalDistribution({str(j): ((1.0 + cc) if i == j else cc) / (m * cc + 1) for j in range(m)})
22
+ dist4 = OptionalDistribution(PoissonDistribution((i + 1) * ss), p=0.1)
23
+ dist = SequenceDistribution(CompositeDistribution((dist1, dist2, dist3, dist4)), len_dist)
24
+ components.append(dist)
25
+ w[i] = np.prod(1 - pvec[:i]) * pvec[i]
26
+
27
+ w[n:] = 1.0e-16
28
+ w /= w.sum()
29
+
30
+ dist = MixtureDistribution(components, w)
31
+ data = dist.sampler(seed=1).sample(300)
32
+
33
+ est = get_estimator(data, use_bstats=True)
34
+ model = get_dpm_mixture(data, rng=np.random.RandomState(1))
35
+
36
+ print(str(model))
37
+ print(model.num_components)
38
+ for u in model.components:
39
+ print(str(u))
dmx/bstats/__init__.py ADDED
@@ -0,0 +1,491 @@
1
+ __all__ = ['BernoulliDistribution', 'BernoulliEstimator', 'BernoulliSampler',
2
+ 'BernoulliSetDistribution', 'BernoulliSetEstimator', 'BernoulliSetSampler',
3
+ 'BetaDistribution', 'BetaSampler',
4
+ 'CategoricalDistribution', 'CategoricalEstimator', 'CategoricalSampler',
5
+ 'CompositeDistribution', 'CompositeEstimator', 'CompositeSampler',
6
+ 'DiagonalGaussianDistribution', 'DiagonalGaussianEstimator', 'DiagonalGaussianSampler',
7
+ 'DictDirichletDistribution',
8
+ 'DirichletDistribution', 'DirichletEstimator', 'DirichletSampler',
9
+ 'DirichletProcessMixtureDistribution', 'DirichletProcessMixtureEstimator', 'DirichletProcessMixtureSampler',
10
+ 'ExponentialDistribution', 'ExponentialEstimator', 'ExponentialSampler',
11
+ 'GaussianDistribution', 'GaussianEstimator', 'GaussianSampler',
12
+ 'GammaDistribution', 'GammaEstimator', 'GammaSampler',
13
+ 'GeometricDistribution', 'GeometricEstimator', 'GeometricSampler',
14
+ 'IgnoredDistribution', 'IgnoredEstimator', 'IgnoredSampler',
15
+ 'IntegerCategoricalDistribution', 'IntegerCategoricalEstimator', 'IntegerCategoricalSampler',
16
+ 'MixtureDistribution', 'MixtureEstimator', 'MixtureSampler',
17
+ 'MultivariateNormalGammaDistribution', 'MultivariateNormalGammaSampler',
18
+ 'NullDistribution', 'NullEstimator', 'NullSampler',
19
+ 'OptionalDistribution', 'OptionalEstimator', 'OptionalSampler',
20
+ 'PoissonDistribution', 'PoissonEstimator', 'PoissonSampler',
21
+ 'SequenceDistribution', 'SequenceEstimator', 'SequenceSampler',
22
+ 'estimate', 'seq_estimate', 'initialize', 'seq_log_density_sum', 'seq_encode', 'seq_log_density']
23
+
24
+ from typing import Any, Tuple, Sequence, Dict, Optional, Union, List
25
+
26
+ from dmx.arithmetic import *
27
+
28
+ from dmx.bstats.pdist import ParameterEstimator, DataSequenceEncoder, EncodedDataSequence, ProbabilityDistribution
29
+
30
+ from dmx.bstats.beta import BetaDistribution, BetaSampler
31
+ from dmx.bstats.bernoulli import BernoulliDistribution, BernoulliEstimator, BernoulliSampler
32
+ from dmx.bstats.categorical import CategoricalDistribution, CategoricalEstimator, CategoricalSampler
33
+ from dmx.bstats.composite import CompositeDistribution, CompositeEstimator, CompositeSampler
34
+ from dmx.bstats.catdirichlet import DictDirichletDistribution
35
+ from dmx.bstats.dirichlet import DirichletDistribution, DirichletEstimator, DirichletSampler
36
+ from dmx.bstats.dmvn import DiagonalGaussianDistribution, DiagonalGaussianEstimator, DiagonalGaussianSampler
37
+ from dmx.bstats.exponential import ExponentialDistribution, ExponentialEstimator, ExponentialSampler
38
+ from dmx.bstats.gaussian import GaussianDistribution, GaussianEstimator, GaussianSampler
39
+ from dmx.bstats.gamma import GammaDistribution, GammaEstimator, GammaSampler
40
+ from dmx.bstats.geometric import GeometricDistribution, GeometricEstimator, GeometricSampler
41
+ from dmx.bstats.ignored import IgnoredDistribution, IgnoredEstimator, IgnoredSampler
42
+ from dmx.bstats.intrange import IntegerCategoricalDistribution, IntegerCategoricalEstimator, IntegerCategoricalSampler
43
+ from dmx.bstats.mixture import MixtureDistribution, MixtureEstimator, MixtureSampler
44
+ from dmx.bstats.mvngamma import MultivariateNormalGammaDistribution, MultivariateNormalGammaSampler
45
+ from dmx.bstats.nulldist import NullDistribution, NullEstimator, NullSampler
46
+ from dmx.bstats.optional import OptionalDistribution, OptionalEstimator, OptionalSampler
47
+ from dmx.bstats.poisson import PoissonDistribution, PoissonEstimator, PoissonSampler
48
+ from dmx.bstats.sequence import SequenceDistribution, SequenceEstimator, SequenceSampler
49
+ from dmx.bstats.setdist import BernoulliSetDistribution, BernoulliSetEstimator, BernoulliSetSampler
50
+
51
+ from dmx.bstats.dpm import DirichletProcessMixtureDistribution, DirichletProcessMixtureEstimator, DirichletProcessMixtureSampler
52
+
53
+
54
+ import numpy as np
55
+ from numpy.random import RandomState
56
+ from pyspark import RDD
57
+ import pickle
58
+ import pandas as pd
59
+
60
+ def load_models(x: str) -> ProbabilityDistribution:
61
+ """Read in ProbabilityDistribution from string representaiton.
62
+
63
+ Args:
64
+ x (str): String representation of ProbabilityDistribution.
65
+
66
+ Returns:
67
+ ProbabilityDistribution
68
+ """
69
+ return eval(x)
70
+
71
+ def dump_models(x: ProbabilityDistribution) -> str:
72
+ """Serialize Probability Distribution.
73
+
74
+ Args:
75
+ x (ProbabilityDistribution): distribution to serialize.
76
+
77
+ Returns:
78
+ str
79
+ """
80
+ return str(x)
81
+
82
+
83
+
84
+ def _local_estimate(
85
+ data: Sequence[Any],
86
+ estimator: ParameterEstimator,
87
+ prev_estimate: Optional[ProbabilityDistribution] = None
88
+ ) -> ProbabilityDistribution:
89
+ """
90
+ Perform a local estimation of parameters using the provided data and estimator.
91
+
92
+ Args:
93
+ data (Sequence[Any]): The input sequence of data points to be used for estimation.
94
+ estimator (ParameterEstimator): The estimator object that provides methods for creating accumulators
95
+ and estimating parameters.
96
+ prev_estimate (Optional[ProbabilityDistribution]): An optional previous probability distribution
97
+ estimate to guide the current estimation process. Defaults to None.
98
+
99
+ Returns:
100
+ Any: The result of the estimation process, as determined by the estimator's `estimate` method.
101
+ """
102
+ idata = iter(data)
103
+ accumulator = estimator.accumulator_factory().make()
104
+ nobs = 0.0
105
+
106
+ for x in idata:
107
+ nobs += 1.0
108
+ accumulator.update(x, 1.0, estimate=prev_estimate)
109
+
110
+ stats_dict = dict()
111
+ accumulator.key_merge(stats_dict)
112
+ accumulator.key_replace(stats_dict)
113
+
114
+ return estimator.estimate(accumulator.value())
115
+
116
+
117
+ def estimate(
118
+ data: Union[RDD, pd.DataFrame, Sequence[Any]],
119
+ estimator: ParameterEstimator,
120
+ prev_estimate: Optional[ProbabilityDistribution] = None
121
+ ) -> ProbabilityDistribution:
122
+ """
123
+ Estimate parameters based on the input data, using the given estimator and optional previous estimate.
124
+
125
+ This function supports data in multiple formats, including PySpark RDDs, Pandas DataFrames, and iterables.
126
+
127
+ Args:
128
+ data (Union[RDD, pd.DataFrame, Sequence[Any]]): The input data for estimation. It can be:
129
+ - A PySpark RDD (`pyspark.rdd.RDD`).
130
+ - A Pandas DataFrame (`pandas.core.frame.DataFrame`).
131
+ - Any iterable object.
132
+ estimator (ParameterEstimator): The estimator object that provides methods for creating accumulators
133
+ and estimating parameters.
134
+ prev_estimate (Optional[ProbabilityDistribution]): An optional previous probability distribution
135
+ estimate to guide the current estimation process. Defaults to None.
136
+
137
+ Returns:
138
+ ProbabilityDistribution: Variational inference update
139
+ """
140
+ if 'pyspark.rdd' in str(type(data)):
141
+ sc = data.context
142
+ factory = estimator.accumulatorFactory()
143
+ estimatorBroadcast = sc.broadcast(estimator)
144
+
145
+ temp_estimate = pickle.dumps(prev_estimate, protocol=0)
146
+ temp_estimateB = sc.broadcast(temp_estimate)
147
+
148
+ def acc(splitIndex, itr):
149
+ accumulatorForSplit = estimatorBroadcast.value.accumulatorFactory().make()
150
+ countsForSplit = 0.0
151
+ loc_prev_estimate = pickle.loads(temp_estimateB.value)
152
+
153
+ for x in itr:
154
+ countsForSplit += 1.0
155
+ accumulatorForSplit.update(x, 1.0, estimate=loc_prev_estimate)
156
+
157
+ return iter([(countsForSplit, accumulatorForSplit.value())])
158
+
159
+ temp = data.mapPartitionsWithIndex(acc, True)
160
+ nobs = 0.0
161
+ accumulator = factory.make()
162
+
163
+ for nobsForSplit, statsForSplit in temp.collect():
164
+ nobs += nobsForSplit
165
+ accumulator.combine(statsForSplit)
166
+
167
+ return estimator.estimate(nobs, accumulator.value())
168
+
169
+ elif 'pandas.core.frame.DataFrame' in str(type(data)):
170
+ accumulator = estimator.accumulatorFactory().make()
171
+ accumulator.df_update(data, np.ones(len(data)), estimate=prev_estimate)
172
+ return estimator.estimate(None, accumulator.value())
173
+
174
+ elif hasattr(data, '__iter__'):
175
+ return _local_estimate(data, estimator, prev_estimate)
176
+
177
+
178
+ def initialize(
179
+ data: Union[Sequence[Any], RDD, pd.DataFrame],
180
+ estimator: ParameterEstimator,
181
+ rng: RandomState,
182
+ p: float
183
+ ) -> Any:
184
+ """
185
+ Initialize parameters based on the input data, using the given estimator, random state, and sampling probability.
186
+
187
+ This function supports data in multiple formats, including PySpark RDDs, Pandas DataFrames, and iterables.
188
+
189
+ Args:
190
+ data (Union[Sequence[Any], RDD, pd.DataFrame]): The input data for initialization. It can be:
191
+ - A PySpark RDD (`pyspark.rdd.RDD`).
192
+ - A Pandas DataFrame (`pandas.core.frame.DataFrame`).
193
+ - Any iterable object.
194
+ estimator (ParameterEstimator): The estimator object that provides methods for creating accumulators
195
+ and estimating parameters.
196
+ rng (RandomState): A NumPy random state object used for random sampling.
197
+ p (float): The probability for random sampling.
198
+
199
+ Returns:
200
+ Any: The result of the initialization process, as determined by the estimator's `estimate` method.
201
+ """
202
+ if 'pyspark.rdd' in str(type(data)):
203
+ factory = estimator.accumulator_factory()
204
+ sc = data.context
205
+
206
+ num_partitions = data.getNumPartitions()
207
+ seeds = rng.randint(np.iinfo(np.int32).max, size=num_partitions)
208
+
209
+ estimatorBroadcast = sc.broadcast(estimator)
210
+ seedsBroadcast = sc.broadcast(seeds)
211
+
212
+ def acc(splitIndex, itr):
213
+ accumulatorForSplit = estimatorBroadcast.value.accumulator_factory().make()
214
+ countsForSplit = 0.0
215
+ rng_loc = np.random.RandomState(seedsBroadcast.value[splitIndex])
216
+
217
+ for x in itr:
218
+ w = 1.0 if rng_loc.rand() <= p else 0.0
219
+ countsForSplit += w
220
+ accumulatorForSplit.initialize(x, w, rng_loc)
221
+
222
+ return iter([(countsForSplit, accumulatorForSplit.value())])
223
+
224
+ temp = data.mapPartitionsWithIndex(acc, True)
225
+ nobs = 0.0
226
+ accumulator = factory.make()
227
+
228
+ for nobsForSplit, statsForSplit in temp.collect():
229
+ nobs += nobsForSplit
230
+ accumulator.combine(statsForSplit)
231
+
232
+ stats_dict = dict()
233
+ accumulator.key_merge(stats_dict)
234
+ accumulator.key_replace(stats_dict)
235
+
236
+ return estimator.estimate(nobs, accumulator.value())
237
+
238
+ elif 'pandas.core.frame.DataFrame' in str(type(data)):
239
+ accumulator = estimator.accumulator_factory().make()
240
+ accumulator.df_initialize(data, rng.rand(len(data)) * p, rng)
241
+ return estimator.estimate(None, accumulator.value())
242
+
243
+ elif hasattr(data, '__iter__'):
244
+ idata = iter(data)
245
+ accumulator = estimator.accumulator_factory().make()
246
+ nobs = 0.0
247
+
248
+ for x in idata:
249
+ w = 1.0 if rng.rand() <= p else 0.0
250
+ nobs += w
251
+ accumulator.initialize(x, w, rng)
252
+
253
+ stats_dict = dict()
254
+ accumulator.key_merge(stats_dict)
255
+ accumulator.key_replace(stats_dict)
256
+
257
+ return estimator.estimate(accumulator.value())
258
+
259
+
260
+ def seq_encode(
261
+ data: Union[Sequence[Any], RDD],
262
+ model: ProbabilityDistribution,
263
+ num_chunks: int = 1,
264
+ chunk_size: Optional[int] = None
265
+ ) -> Union[RDD, List[Tuple[int, Any]]]:
266
+ """
267
+ Encodes a sequence of data using a given probability distribution model.
268
+
269
+ Args:
270
+ data (Union[Sequence[Any], RDD]): The input data to encode. This can be a standard Python sequence
271
+ or a PySpark RDD.
272
+ model (ProbabilityDistribution): The model used for encoding the data. It must implement a `seq_encode` method.
273
+ num_chunks (int, optional): The number of chunks to divide the data into for encoding. Defaults to 1.
274
+ chunk_size (Optional[int], optional): The size of each chunk. If provided, it overrides `num_chunks`. Defaults to None.
275
+
276
+ Returns:
277
+ Union[RDD, List[Tuple[int, Any]]]:
278
+ - If `data` is a PySpark RDD, returns an RDD where each element is a tuple containing the length of the chunk
279
+ and the encoded data.
280
+ - If `data` is a sequence, returns a list of tuples, where each tuple contains the length of the chunk
281
+ and the encoded data.
282
+ """
283
+ if 'pyspark.rdd' in str(type(data)):
284
+ sc = data.context
285
+
286
+ temp_model = pickle.dumps(model, protocol=0)
287
+ modelBroadcast = sc.broadcast(temp_model)
288
+
289
+ enc_data = data.glom().map(lambda x: list(x)).map(
290
+ lambda x: (len(x), pickle.loads(modelBroadcast.value).seq_encode(x))
291
+ )
292
+
293
+ return enc_data
294
+
295
+ else:
296
+ sz = len(data)
297
+ if chunk_size is not None:
298
+ num_chunks_loc = int(np.ceil(float(sz) / float(chunk_size)))
299
+ else:
300
+ num_chunks_loc = num_chunks
301
+
302
+ rv = []
303
+ for i in range(num_chunks_loc):
304
+ data_loc = [data[j] for j in range(i, sz, num_chunks_loc)]
305
+ enc_data = model.seq_encode(data_loc)
306
+ rv.append((len(data_loc), enc_data))
307
+
308
+ return rv
309
+
310
+ def seq_estimate(
311
+ enc_data: Union[RDD, Sequence[tuple[int, Any]]],
312
+ estimator: ParameterEstimator,
313
+ prev_estimate: Any
314
+ ) -> Any:
315
+ """
316
+ Sequentially estimate parameters based on encoded data, using the given estimator and previous estimate.
317
+
318
+ This function supports data in multiple formats, including PySpark RDDs and sequences of tuples.
319
+
320
+ Args:
321
+ enc_data (Union[RDD, Sequence[tuple[int, Any]]]): The encoded data for estimation. It can be:
322
+ - A PySpark RDD (`pyspark.rdd.RDD`) containing tuples of size and data.
323
+ - A sequence of tuples, where each tuple contains an integer size and associated data.
324
+ estimator (ParameterEstimator): The estimator object that provides methods for creating accumulators
325
+ and estimating parameters.
326
+ prev_estimate (Any): The previous estimate to guide the current estimation process.
327
+
328
+ Returns:
329
+ Any: The result of the sequential estimation process, as determined by the estimator's `estimate` method.
330
+ """
331
+
332
+ if 'pyspark.rdd' in str(type(enc_data)):
333
+ sc = enc_data.context
334
+
335
+ estimatorBroadcast = sc.broadcast(estimator)
336
+ estimateBroadcast = sc.broadcast(pickle.dumps(prev_estimate, protocol=0))
337
+
338
+ def acc(splitIndex, itr):
339
+ accumulatorForSplit = estimatorBroadcast.value.accumulatorFactory().make()
340
+ countsForSplit = zero
341
+ local_estimate = pickle.loads(estimateBroadcast.value)
342
+
343
+ for sz, x in itr:
344
+ countsForSplit = countsForSplit + sz
345
+ accumulatorForSplit.seq_update(x, np.ones(sz), local_estimate)
346
+
347
+ rv = pickle.dumps((countsForSplit, accumulatorForSplit.value()), protocol=0)
348
+ #return [(countsForSplit, accumulatorForSplit.value())]
349
+ return [rv]
350
+
351
+ def red(x, y):
352
+ """Reduce function to combine accros partitions."""
353
+ xx = pickle.loads(x)
354
+ yy = pickle.loads(y)
355
+ accumulator = estimatorBroadcast.value.accumulatorFactory().make()
356
+ nobs = xx[0] + yy[0]
357
+ vals = accumulator.from_value(xx[1]).combine(yy[1]).value()
358
+ rv = pickle.dumps((nobs, vals))
359
+ #return (nobs, vals)
360
+ return rv
361
+
362
+
363
+ temp = enc_data.mapPartitionsWithIndex(acc, True).cache()
364
+
365
+ nobs = zero
366
+ accumulator = estimator.accumulatorFactory().make()
367
+
368
+ for stuff in temp.collect():
369
+ nobsForSplit, statsForSplit = pickle.loads(stuff)
370
+ nobs = nobs + nobsForSplit
371
+ accumulator.combine(statsForSplit)
372
+
373
+
374
+ stats_dict = dict()
375
+ accumulator.key_merge(stats_dict)
376
+ accumulator.key_replace(stats_dict)
377
+
378
+
379
+ estimateBroadcast.destroy()
380
+ estimatorBroadcast.destroy()
381
+ temp.unpersist()
382
+ enc_data.localCheckpoint()
383
+
384
+ return(estimator.estimate(nobs, accumulator.value()))
385
+
386
+ else:
387
+
388
+ accumulator = estimator.accumulator_factory().make()
389
+ nobs = 0.0
390
+
391
+ data_update = []
392
+
393
+ for sz, x in enc_data:
394
+ nobs += sz
395
+ accumulator.seq_update(x, np.ones(sz), prev_estimate)
396
+ #x_update = accumulator.seq_update(x, np.ones(sz), prev_estimate)
397
+ #data_update.append((sz, x_update))
398
+
399
+ stats_dict = dict()
400
+ accumulator.key_merge(stats_dict)
401
+ accumulator.key_replace(stats_dict)
402
+
403
+ return estimator.estimate(accumulator.value())
404
+
405
+
406
+ def seq_log_density(
407
+ enc_data: Union[RDD, Sequence[Tuple[int, EncodedDataSequence]]],
408
+ estimate: Union[ProbabilityDistribution, List[ProbabilityDistribution]],
409
+ is_list: bool = False
410
+ ) -> List[np.ndarray]:
411
+ """
412
+ Compute the sequential log density for encoded data using the given estimate.
413
+
414
+ This function supports data in multiple formats, including PySpark RDDs and sequences of tuples.
415
+
416
+ Args:
417
+ enc_data (Union[RDD, Sequence[Tuple[int, EncodedDataSequence]]]): The encoded data for density computation. It can be:
418
+ - A PySpark RDD (`pyspark.rdd.RDD`) containing tuples of size and data.
419
+ - A sequence of tuples, where each tuple contains an integer size and associated EncodedDataSequence.
420
+ estimate (Union[ProbabilityDistribution, List[ProbabilityDistribution]]): The estimate object or a list of estimates used for log density computation.
421
+ is_list (bool): Whether the `estimate` is a list of estimates. Defaults to `False`.
422
+
423
+ Returns:
424
+ List[np.ndarray]: A list of log density values computed for the encoded data.
425
+ """
426
+
427
+ if 'pyspark.rdd' in str(type(enc_data)):
428
+ sc = enc_data.context
429
+ temp_estimate = pickle.dumps(estimate, protocol=0)
430
+ estimateBroadcast = sc.broadcast(temp_estimate)
431
+
432
+ def acc(itr):
433
+ loc_estimate = pickle.loads(estimateBroadcast.value)
434
+ if is_list:
435
+ return [np.asarray([ee.seq_log_density(x) for ee in loc_estimate]) for sz, x in itr]
436
+ else:
437
+ return [loc_estimate.seq_log_density(x) for sz,x in itr]
438
+
439
+ return enc_data.mapPartitions(acc).collect()
440
+
441
+ else:
442
+
443
+ if is_list:
444
+ return [np.asarray([ee.seq_log_density(u[1]) for ee in estimate]) for u in enc_data]
445
+ else:
446
+ return [estimate.seq_log_density(u[1]) for u in enc_data]
447
+
448
+
449
+
450
+ def seq_log_density_sum(
451
+ enc_data: Union[RDD, Sequence[Tuple[int, EncodedDataSequence]]],
452
+ estimate: ProbabilityDistribution
453
+ ) -> Tuple[float, float]:
454
+ """
455
+ Compute the sum of sequential log densities for encoded data using the given estimate.
456
+
457
+ This function supports data in multiple formats, including PySpark RDDs and sequences of tuples.
458
+
459
+ Args:
460
+ enc_data (Union[RDD, Sequence[Tuple[int, EncodedDataSequence]]]): The encoded data for density computation.
461
+ It can be:
462
+ - A PySpark RDD (`pyspark.rdd.RDD`) containing tuples of size and data.
463
+ - A sequence of tuples, where each tuple contains an integer size and associated data.
464
+ estimate (ProbabilityDistribution): The ProbabilityDistribution used for log density computation.
465
+
466
+ Returns:
467
+ Tuple[float, float]: A tuple containing:
468
+ - The total count of observations (`cnt`).
469
+ - The sum of log densities (`rv`).
470
+ """
471
+
472
+ if 'pyspark.rdd' in str(type(enc_data)):
473
+ sc = enc_data.context
474
+ estimate_broadcast = sc.broadcast(pickle.dumps(estimate, protocol=0))
475
+
476
+ def acc(itr):
477
+
478
+ rv = 0.0
479
+ cnt = 0.0
480
+ estimate_loc = pickle.loads(estimate_broadcast.value)
481
+ for sz, x in itr:
482
+ rv += estimate_loc.seq_log_density(x).sum()
483
+ cnt += sz
484
+
485
+ return [(cnt, rv)]
486
+
487
+ return enc_data.mapPartitions(acc).reduce(lambda a,b : (a[0]+b[0], a[1]+b[1]))
488
+
489
+ else:
490
+
491
+ return sum([u[0] for u in enc_data]), sum([estimate.seq_log_density(u[1]).sum() for u in enc_data])