1"""Common utility functions for rolling operations"""
2
3from __future__ import annotations
4
5from collections import defaultdict
6from typing import cast
7
8import numpy as np
9
10from pandas.core.dtypes.generic import (
11 ABCDataFrame,
12 ABCSeries,
13)
14
15from pandas.core.indexes.api import MultiIndex
16
17
18def flex_binary_moment(arg1, arg2, f, pairwise: bool = False):
19 if isinstance(arg1, ABCSeries) and isinstance(arg2, ABCSeries):
20 X, Y = prep_binary(arg1, arg2)
21 return f(X, Y)
22
23 elif isinstance(arg1, ABCDataFrame):
24 from pandas import DataFrame
25
26 def dataframe_from_int_dict(data, frame_template) -> DataFrame:
27 result = DataFrame(data, index=frame_template.index)
28 if len(result.columns) > 0:
29 result.columns = frame_template.columns[result.columns]
30 else:
31 result.columns = frame_template.columns.copy()
32 return result
33
34 results = {}
35 if isinstance(arg2, ABCDataFrame):
36 if pairwise is False:
37 if arg1 is arg2:
38 # special case in order to handle duplicate column names
39 for i in range(len(arg1.columns)):
40 results[i] = f(arg1.iloc[:, i], arg2.iloc[:, i])
41 return dataframe_from_int_dict(results, arg1)
42 else:
43 if not arg1.columns.is_unique:
44 raise ValueError("'arg1' columns are not unique")
45 if not arg2.columns.is_unique:
46 raise ValueError("'arg2' columns are not unique")
47 X, Y = arg1.align(arg2, join="outer")
48 X, Y = prep_binary(X, Y)
49 res_columns = arg1.columns.union(arg2.columns)
50 for col in res_columns:
51 if col in X and col in Y:
52 results[col] = f(X[col], Y[col])
53 return DataFrame(results, index=X.index, columns=res_columns)
54 elif pairwise is True:
55 results = defaultdict(dict)
56 for i in range(len(arg1.columns)):
57 for j in range(len(arg2.columns)):
58 if j < i and arg2 is arg1:
59 # Symmetric case
60 results[i][j] = results[j][i]
61 else:
62 results[i][j] = f(
63 *prep_binary(arg1.iloc[:, i], arg2.iloc[:, j])
64 )
65
66 from pandas import concat
67
68 result_index = arg1.index.union(arg2.index)
69 if len(result_index):
70 # construct result frame
71 result = concat(
72 [
73 concat(
74 [results[i][j] for j in range(len(arg2.columns))],
75 ignore_index=True,
76 )
77 for i in range(len(arg1.columns))
78 ],
79 ignore_index=True,
80 axis=1,
81 )
82 result.columns = arg1.columns
83
84 # set the index and reorder
85 if arg2.columns.nlevels > 1:
86 # mypy needs to know columns is a MultiIndex, Index doesn't
87 # have levels attribute
88 arg2.columns = cast(MultiIndex, arg2.columns)
89 # GH 21157: Equivalent to MultiIndex.from_product(
90 # [result_index], <unique combinations of arg2.columns.levels>,
91 # )
92 # A normal MultiIndex.from_product will produce too many
93 # combinations.
94 result_level = np.tile(
95 result_index, len(result) // len(result_index)
96 )
97 arg2_levels = (
98 np.repeat(
99 arg2.columns.get_level_values(i),
100 len(result) // len(arg2.columns),
101 )
102 for i in range(arg2.columns.nlevels)
103 )
104 result_names = [*arg2.columns.names, result_index.name]
105 result.index = MultiIndex.from_arrays(
106 [*arg2_levels, result_level], names=result_names
107 )
108 # GH 34440
109 num_levels = len(
110 result.index.levels # pyright: ignore[reportAttributeAccessIssue]
111 )
112 new_order = [num_levels - 1, *range(num_levels - 1)]
113 result = result.reorder_levels(new_order).sort_index()
114 else:
115 result.index = MultiIndex.from_product(
116 [range(len(arg2.columns)), range(len(result_index))]
117 )
118 result = result.swaplevel(1, 0).sort_index()
119 result.index = MultiIndex.from_product(
120 [result_index, arg2.columns]
121 )
122 else:
123 # empty result
124 result = DataFrame(
125 index=MultiIndex(
126 levels=[arg1.index, arg2.columns], codes=[[], []]
127 ),
128 columns=arg2.columns,
129 dtype="float64",
130 )
131
132 # reset our index names to arg1 names
133 # reset our column names to arg2 names
134 # careful not to mutate the original names
135 result.columns = result.columns.set_names(arg1.columns.names)
136 result.index = result.index.set_names(
137 result_index.names + arg2.columns.names
138 )
139
140 return result
141 else:
142 results = {
143 i: f(*prep_binary(arg1.iloc[:, i], arg2))
144 for i in range(len(arg1.columns))
145 }
146 return dataframe_from_int_dict(results, arg1)
147
148 else:
149 return flex_binary_moment(arg2, arg1, f)
150
151
152def zsqrt(x):
153 with np.errstate(all="ignore"):
154 result = np.sqrt(x)
155 mask = x < 0
156
157 if isinstance(x, ABCDataFrame):
158 if mask._values.any():
159 result[mask] = 0
160 elif mask.any():
161 result[mask] = 0
162
163 return result
164
165
166def prep_binary(arg1, arg2):
167 # mask out values, this also makes a common index...
168 X = arg1 + 0 * arg2
169 Y = arg2 + 0 * arg1
170
171 return X, Y