lessPython 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.
Files changed (82) hide show
  1. lessPy/ANOVA.py +680 -0
  2. lessPy/Chart.py +1055 -0
  3. lessPy/Correlation.py +236 -0
  4. lessPy/Flows.py +116 -0
  5. lessPy/Logit.py +615 -0
  6. lessPy/Prop_test.py +267 -0
  7. lessPy/Regression.py +1491 -0
  8. lessPy/VariableLabels.py +119 -0
  9. lessPy/X.py +426 -0
  10. lessPy/XY.py +2007 -0
  11. lessPy/__init__.py +60 -0
  12. lessPy/anova_rmd.py +227 -0
  13. lessPy/bc_plotly.py +575 -0
  14. lessPy/bubble_plotly.py +470 -0
  15. lessPy/corCFA.py +316 -0
  16. lessPy/corEFA.py +220 -0
  17. lessPy/corPrint.py +45 -0
  18. lessPy/corProp.py +73 -0
  19. lessPy/corRead.py +48 -0
  20. lessPy/corReflect.py +72 -0
  21. lessPy/corReorder.py +161 -0
  22. lessPy/corScree.py +87 -0
  23. lessPy/data/Anova_1way.csv +25 -0
  24. lessPy/data/Anova_2way.csv +49 -0
  25. lessPy/data/Anova_rb.csv +8 -0
  26. lessPy/data/Anova_rbf.csv +49 -0
  27. lessPy/data/Anova_sp.csv +57 -0
  28. lessPy/data/BodyMeas.csv +341 -0
  29. lessPy/data/Cars93.csv +94 -0
  30. lessPy/data/Employee.csv +38 -0
  31. lessPy/data/Employee_lbl.csv +9 -0
  32. lessPy/data/FreqTable99.csv +5 -0
  33. lessPy/data/Jackets.csv +1026 -0
  34. lessPy/data/Learn.csv +35 -0
  35. lessPy/data/Mach4.csv +352 -0
  36. lessPy/data/Mach4_lbl.csv +21 -0
  37. lessPy/data/Reading.csv +101 -0
  38. lessPy/data/StockPrice.csv +1489 -0
  39. lessPy/data/WeightLoss.csv +11 -0
  40. lessPy/datasets.py +46 -0
  41. lessPy/date_infer.py +112 -0
  42. lessPy/details.py +314 -0
  43. lessPy/dn_plotly.py +495 -0
  44. lessPy/dot_plotly.py +385 -0
  45. lessPy/freq_poly_plotly.py +324 -0
  46. lessPy/getColors.py +399 -0
  47. lessPy/hier_plotly.py +352 -0
  48. lessPy/hs_plotly.py +395 -0
  49. lessPy/logit_rmd.py +410 -0
  50. lessPy/order_by.py +94 -0
  51. lessPy/pie_plotly.py +292 -0
  52. lessPy/pivot.py +158 -0
  53. lessPy/plotly_utils.py +787 -0
  54. lessPy/plt_add.py +129 -0
  55. lessPy/plt_contour.py +192 -0
  56. lessPy/plt_contour_facet.py +194 -0
  57. lessPy/plt_forecast.py +677 -0
  58. lessPy/plt_mat_plotly.py +201 -0
  59. lessPy/plt_plotly.py +216 -0
  60. lessPy/plt_smooth.py +170 -0
  61. lessPy/plt_time.py +143 -0
  62. lessPy/prob_norm.py +111 -0
  63. lessPy/prob_tcut.py +131 -0
  64. lessPy/prob_znorm.py +110 -0
  65. lessPy/radar_plotly.py +201 -0
  66. lessPy/reg_rmd.py +754 -0
  67. lessPy/rename.py +33 -0
  68. lessPy/reshape.py +95 -0
  69. lessPy/showColors.py +130 -0
  70. lessPy/simCImean.py +165 -0
  71. lessPy/simCLT.py +265 -0
  72. lessPy/simFlips.py +104 -0
  73. lessPy/simMeans.py +146 -0
  74. lessPy/stats_out.py +189 -0
  75. lessPy/ttest.py +641 -0
  76. lessPy/utils.py +235 -0
  77. lessPy/vbs_plotly.py +545 -0
  78. lesspython-0.1.0.dist-info/METADATA +93 -0
  79. lesspython-0.1.0.dist-info/RECORD +82 -0
  80. lesspython-0.1.0.dist-info/WHEEL +5 -0
  81. lesspython-0.1.0.dist-info/licenses/LICENSE +338 -0
  82. lesspython-0.1.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,201 @@
1
+ # plt_mat_plotly.py — the scatterplot matrix (SPLOM).
2
+ #
3
+ # Shared by Regression() and XY(), mirroring the single R helper
4
+ # .plt.mat (plt.mat.R). The layout matches R's pairs() panels:
5
+ # lower triangle scatter of (col var, row var) with a fit line
6
+ # and its 95% confidence band
7
+ # upper triangle the correlation coefficient, as text
8
+ # diagonal the variable name, on the se_fill background
9
+ # Regression() calls it with fit="lm"; XY() passes its own fit=
10
+ # (default "off", so no line). Points are dropped listwise across
11
+ # all variables first (R's na.omit).
12
+
13
+ import numpy as np
14
+ import plotly.graph_objects as go
15
+ from plotly.subplots import make_subplots
16
+ from scipy import stats as sps
17
+
18
+ from .plotly_utils import (
19
+ as_plotly_color, axis_format, make_trans, plotly_style,
20
+ to_hex)
21
+ from .utils import fmt, get_option, pretty
22
+
23
+
24
+ def _axis_ref(k, kind):
25
+ """Domain reference string for the k-th subplot axis, as used
26
+ by shapes and annotations: 'x domain', 'x2 domain', ..."""
27
+ n = "" if k == 1 else str(k)
28
+ return f"{kind}{n} domain"
29
+
30
+
31
+ def scatter_matrix(df, fit="lm", digits_d=2, cor_coef=True,
32
+ main=None, band=True):
33
+ """Scatterplot matrix of the columns of df, in column order.
34
+ fit: "off" (no line, XY default), "lm" (Regression) or
35
+ "loess". cor_coef=True puts the correlation in the upper
36
+ triangle and the scatter below (~ .plt.mat, Regression);
37
+ cor_coef=False draws the scatter in every off-diagonal cell,
38
+ the symmetric matrix R's Logit uses (logit.4Pred). band adds
39
+ the fit's 95% confidence band. Returns a plotly Figure."""
40
+ cols = list(df.columns)
41
+ n = len(cols)
42
+ if n < 2:
43
+ raise ValueError(
44
+ "a scatterplot matrix needs at least 2 variables")
45
+
46
+ data = df.dropna() # R na.omit: listwise
47
+ vals = {c: data[c].to_numpy(dtype=float) for c in cols}
48
+ rng = {c: pretty(float(v.min()), float(v.max()))
49
+ for c, v in vals.items()}
50
+
51
+ style = plotly_style()
52
+ panel_fill = get_option("panel_fill", "white")
53
+ window_fill = get_option("window_fill", "white")
54
+ bg = window_fill if panel_fill == "transparent" else panel_fill
55
+ diag_fill = get_option("se_fill", "#1A1A1A19")
56
+ se_fill = get_option("se_fill", "#1A1A1A19")
57
+ border = to_hex(style["panel_border"])
58
+ pt_fill = get_option("pt_fill", get_option("pt_color",
59
+ "#324E5C"))
60
+ pt_color = get_option("pt_color", "#324E5C")
61
+ lab_color = to_hex(style["lab_color"])
62
+ fit_color = to_hex(get_option("fit_color", "#5C4032"))
63
+ fit_lwd = get_option("fit_lwd", 2)
64
+
65
+ # sizes shrink with the variable count, as R's cex.adj
66
+ px = 0.5 * max(2.5, 7.25 * (0.80 - 0.048 * n) / 0.75)
67
+ txt_size = max(9, round(10 * (1.6 - 0.065 * n)))
68
+
69
+ fig = make_subplots(rows=n, cols=n, horizontal_spacing=0.008,
70
+ vertical_spacing=0.008)
71
+
72
+ do_fit = fit in ("lm", "loess")
73
+
74
+ for r in range(1, n + 1): # r: row (top = 1)
75
+ for c in range(1, n + 1): # c: column (left = 1)
76
+ k = (r - 1) * n + c # subplot / axis index
77
+ xr, yr = _axis_ref(k, "x"), _axis_ref(k, "y")
78
+
79
+ # scatter in the lower triangle, and also in the
80
+ # upper triangle when no correlations are requested
81
+ if r != c and (r > c or not cor_coef):
82
+ _cell_bg(fig, xr, yr, bg, border)
83
+ _scatter_cell(fig, r, c, vals[cols[c - 1]],
84
+ vals[cols[r - 1]], fit, do_fit, px,
85
+ pt_fill, pt_color, fit_color,
86
+ fit_lwd, se_fill, band)
87
+ continue
88
+
89
+ # text cells (diagonal, upper correlation): an
90
+ # invisible point anchors the axes so the domain
91
+ # shape/text render
92
+ _anchor(fig, r, c)
93
+ if r == c: # diagonal: variable name
94
+ _cell_bg(fig, xr, yr, diag_fill, border)
95
+ fig.add_annotation(
96
+ xref=xr, yref=yr, x=0.5, y=0.5, text=cols[r - 1],
97
+ showarrow=False,
98
+ font=dict(size=txt_size, color=lab_color))
99
+ else: # upper: correlation
100
+ _cell_bg(fig, xr, yr, bg, border)
101
+ a, b = vals[cols[r - 1]], vals[cols[c - 1]]
102
+ rr = float(np.corrcoef(a, b)[0, 1])
103
+ fig.add_annotation(
104
+ xref=xr, yref=yr, x=0.5, y=0.5,
105
+ text=fmt(rr, 2), showarrow=False,
106
+ font=dict(size=txt_size, color="black"))
107
+
108
+ # per-column x range (col var) and per-row y range (row var)
109
+ for c in range(1, n + 1):
110
+ rc = rng[cols[c - 1]]
111
+ fig.update_xaxes(range=[rc[0], rc[-1]], showgrid=False,
112
+ zeroline=False, showticklabels=False,
113
+ ticks="", col=c)
114
+ for r in range(1, n + 1):
115
+ rr = rng[cols[r - 1]]
116
+ fig.update_yaxes(range=[rr[0], rr[-1]], showgrid=False,
117
+ zeroline=False, showticklabels=False,
118
+ ticks="", row=r)
119
+
120
+ # outer scales only: x on the bottom row, y on the left
121
+ # column, interior ticks (drop the endpoints, which would
122
+ # collide at the panel seams)
123
+ for c in range(1, n + 1):
124
+ t = rng[cols[c - 1]][1:-1]
125
+ fig.update_xaxes(
126
+ showticklabels=True, tickvals=t,
127
+ ticktext=axis_format(t, digits_d),
128
+ tickfont=dict(size=8), row=n, col=c)
129
+ for r in range(1, n + 1):
130
+ t = rng[cols[r - 1]][1:-1]
131
+ fig.update_yaxes(
132
+ showticklabels=True, tickvals=t,
133
+ ticktext=axis_format(t, digits_d),
134
+ tickfont=dict(size=8), row=r, col=1)
135
+
136
+ fig.update_layout(
137
+ template=None, showlegend=False,
138
+ plot_bgcolor=to_hex(bg),
139
+ paper_bgcolor=to_hex(window_fill))
140
+ if main:
141
+ fig.update_layout(title=dict(
142
+ text=main, x=0.5, xanchor="center",
143
+ font=dict(size=round(16 * get_option("main_size",
144
+ 1)))))
145
+ return fig
146
+
147
+
148
+ def _anchor(fig, r, c):
149
+ """A single invisible point, so a text-only cell's axes exist
150
+ and its domain-referenced shape and label are drawn."""
151
+ fig.add_trace(go.Scatter(
152
+ x=[0.5], y=[0.5], mode="markers",
153
+ marker=dict(opacity=0), hoverinfo="skip",
154
+ showlegend=False), row=r, col=c)
155
+
156
+
157
+ def _cell_bg(fig, xr, yr, fill, border):
158
+ """Fill the cell and outline it, spanning the full domain."""
159
+ fig.add_shape(type="rect", xref=xr, yref=yr,
160
+ x0=0, x1=1, y0=0, y1=1, layer="below",
161
+ fillcolor=as_plotly_color(fill),
162
+ line=dict(color=border, width=1))
163
+
164
+
165
+ def _scatter_cell(fig, r, c, xv, yv, fit, do_fit, px,
166
+ pt_fill, pt_color, fit_color, fit_lwd, se_fill,
167
+ band=True):
168
+ """Points, and (when fit is on) the fit line and, if band,
169
+ its 95% confidence band, in subplot (row r, col c)."""
170
+ from .XY import _loess, _plt_fit, _se_band
171
+
172
+ if do_fit and len(xv) >= 2:
173
+ if fit == "loess":
174
+ xs, _, f, se_f = _loess(xv, yv, 2 / 3)
175
+ else: # lm
176
+ xs, ys_s, f = _plt_fit(xv, yv, "lm", 1)
177
+ if band:
178
+ if fit == "loess":
179
+ tq = sps.t.ppf((1 + 0.95) / 2, len(xs) - 1)
180
+ poly = (np.concatenate([xs, xs[::-1]]),
181
+ np.concatenate([f + tq * se_f,
182
+ (f - tq * se_f)[::-1]]))
183
+ else:
184
+ poly = _se_band(xs, ys_s, f, 0.95)
185
+ fig.add_trace(go.Scatter(
186
+ x=poly[0], y=poly[1], mode="none", fill="toself",
187
+ fillcolor=as_plotly_color(se_fill),
188
+ hoverinfo="skip", showlegend=False), row=r, col=c)
189
+
190
+ fig.add_trace(go.Scatter(
191
+ x=xv, y=yv, mode="markers",
192
+ marker=dict(symbol="circle", size=px, sizemode="diameter",
193
+ color=make_trans(pt_fill, 0.9), opacity=1,
194
+ line=dict(color=to_hex(pt_color), width=0.5)),
195
+ hoverinfo="x+y", showlegend=False), row=r, col=c)
196
+
197
+ if do_fit and len(xv) >= 2:
198
+ fig.add_trace(go.Scatter(
199
+ x=xs, y=f, mode="lines",
200
+ line=dict(color=fit_color, width=fit_lwd),
201
+ hoverinfo="skip", showlegend=False), row=r, col=c)
lessPy/plt_plotly.py ADDED
@@ -0,0 +1,216 @@
1
+ # plt_plotly.py — analog of plt.plotly.R
2
+ #
3
+ # Renders the XY() scatterplot / time-series view. Deviations from
4
+ # the R renderer, whose interface reflects the base-R path it
5
+ # shares state with:
6
+ # - fit lines, SE bands, and ellipses arrive as prepared
7
+ # line/polygon coordinate lists (XY computes them), not as
8
+ # per-point vectors threaded through a shared data frame
9
+ # - bubble mode (size=) and the outlier/MD path are not ported:
10
+ # bubble displays are internal to Chart(), and outlier
11
+ # flagging awaits the accompanying-statistics port
12
+ # - a date x-axis uses plotly's native date ticks rather than
13
+ # pretty() tick vectors
14
+
15
+ import numpy as np
16
+ import plotly.graph_objects as go
17
+
18
+ from .plotly_utils import (
19
+ as_plotly_color, axis_base, axis_num, get_tick_fmt, make_trans,
20
+ plot_border, plotly_style, sym_at, to_hex, x_grid, y_grid,
21
+ )
22
+ from .utils import get_option
23
+
24
+
25
+ def _hover_fmt(xv, yv, x_lab, y_lab, digits_d, is_date):
26
+ """Hover template for one trace. R analog: .hover.fmt()"""
27
+ fmty = get_tick_fmt(yv, digits_d)
28
+ ypart = f"%{{y:{fmty}}}" if fmty else "%{y}"
29
+ if is_date:
30
+ x_lab, xpart = "Date", "%{x|%Y-%m-%d}"
31
+ else:
32
+ fmtx = get_tick_fmt(xv, digits_d)
33
+ xpart = f"%{{x:{fmtx}}}" if fmtx else "%{x}"
34
+ return (f"{x_lab}: {xpart}<br>{y_lab.strip()}: {ypart}"
35
+ "<extra></extra>")
36
+
37
+
38
+ def plt_plotly(groups, by_name=None,
39
+ fill=("#324E5C",), border=None, shape="circle",
40
+ pt_size=1,
41
+ x_lab="", y_lab="", ax=None, gridT1=None,
42
+ gridT2=None, main=None, digits_d=2,
43
+ connect=False, is_date=False, pt_opacity=0.90,
44
+ area_polys=None,
45
+ fit_lines=None, fit_color=None, fit_lwd=None,
46
+ se_polys=None, se_fill=None,
47
+ ellipses=None, ellipse_fill=None,
48
+ ellipse_color=None, ellipse_lwd=1,
49
+ style_opts=None):
50
+ """Scatterplot / time-series renderer for XY().
51
+
52
+ groups: list of (name, x, y) — one entry, name None, when
53
+ there is no by variable. R analog: plt.plotly()
54
+ """
55
+ if style_opts is None:
56
+ style_opts = plotly_style()
57
+ if fit_color is None:
58
+ fit_color = get_option("fit_color", "#5C4032")
59
+ if fit_lwd is None:
60
+ fit_lwd = get_option("fit_lwd", 2)
61
+
62
+ n_grp = len(groups)
63
+ has_groups = n_grp > 1
64
+
65
+ def recycle(cols):
66
+ cols = (list(cols) if isinstance(cols, (list, tuple))
67
+ else [cols])
68
+ return [cols[i % len(cols)] for i in range(n_grp)]
69
+
70
+ fills = recycle(fill)
71
+ borders = fills if border is None else recycle(border)
72
+
73
+ fig = go.Figure()
74
+
75
+ # areas under/between time series: prepared polygons from
76
+ # XY() (ts_area_fill / ts_stack), drawn beneath the lines
77
+ if area_polys:
78
+ for pax, pay, pcol in area_polys:
79
+ fig.add_trace(go.Scatter(
80
+ x=pax, y=pay, mode="none", fill="toself",
81
+ fillcolor=pcol, hoverinfo="skip",
82
+ showlegend=False))
83
+
84
+ # points: pixel diameter, R's single/grouped scale factors
85
+ px = float(pt_size) * (6.5 if has_groups else 7.25)
86
+ if not np.isfinite(px) or px < 0:
87
+ px = 5
88
+ mode_pts = "lines+markers" if connect else "markers"
89
+ if connect and px == 0:
90
+ mode_pts = "lines"
91
+
92
+ for i, (nm, xv, yv) in enumerate(groups):
93
+ hover = _hover_fmt(xv, yv, x_lab, y_lab, digits_d, is_date)
94
+ if has_groups:
95
+ hover = hover.replace(
96
+ "<extra></extra>", f"<br>{by_name}: {nm}"
97
+ "<extra></extra>")
98
+ marker = dict(
99
+ symbol=sym_at(shape, i), size=px, sizemode="diameter",
100
+ color=make_trans(fills[i], pt_opacity), opacity=1,
101
+ line=dict(color=to_hex(borders[i]), width=1),
102
+ ) if "markers" in mode_pts else None
103
+ fig.add_trace(go.Scatter(
104
+ x=xv, y=yv, mode=mode_pts,
105
+ name=None if nm is None else str(nm),
106
+ legendgroup=None if nm is None else str(nm),
107
+ marker=marker,
108
+ line=(dict(color=to_hex(borders[i]), width=1.5)
109
+ if connect else None),
110
+ hovertemplate=hover,
111
+ # connected series show a legend-only line sample below
112
+ showlegend=has_groups and not connect,
113
+ ))
114
+ if has_groups and connect and len(xv) >= 2:
115
+ fig.add_trace(go.Scatter(
116
+ x=xv[:2], y=yv[:2], mode="lines",
117
+ name=str(nm), legendgroup=str(nm),
118
+ line=dict(color=to_hex(borders[i]), width=1.5),
119
+ hoverinfo="skip", showlegend=True,
120
+ ))
121
+
122
+ # ellipses: border per group when grouped, else ellipse_color
123
+ if ellipses:
124
+ for k, (ex, ey) in enumerate(ellipses):
125
+ edge = (to_hex(borders[k]) if has_groups
126
+ else to_hex(ellipse_color))
127
+ fig.add_trace(go.Scatter(
128
+ x=ex, y=ey, mode="lines",
129
+ line=dict(color=edge, width=ellipse_lwd),
130
+ fill="toself",
131
+ fillcolor=as_plotly_color(ellipse_fill),
132
+ hoverinfo="skip", showlegend=False,
133
+ ))
134
+
135
+ # SE bands under their fit lines
136
+ if se_polys:
137
+ for bx, bnd_y in se_polys:
138
+ fig.add_trace(go.Scatter(
139
+ x=bx, y=bnd_y, mode="none", fill="toself",
140
+ fillcolor=as_plotly_color(se_fill),
141
+ hoverinfo="skip", showlegend=False,
142
+ ))
143
+
144
+ # fit lines: single fit in fit_color; group fits in group hues
145
+ if fit_lines:
146
+ fmtx_all = np.concatenate([f["x"] for f in fit_lines])
147
+ fmtx = get_tick_fmt(fmtx_all, digits_d)
148
+ xpart = f"%{{x:{fmtx}}}" if fmtx else "%{x}"
149
+ for i, fl in enumerate(fit_lines):
150
+ if len(fl["x"]) < 2:
151
+ continue
152
+ fmty = get_tick_fmt(fl["y"], digits_d)
153
+ ypart = f"%{{y:{fmty}}}" if fmty else "%{y}"
154
+ nm = fl["name"]
155
+ single = nm is None
156
+ lbl = "Fit" if single else f"Fit ({nm})"
157
+ fig.add_trace(go.Scatter(
158
+ x=fl["x"], y=fl["y"], mode="lines",
159
+ name="Fit" if single else f"Fit: {nm}",
160
+ legendgroup="fit" if single else str(nm),
161
+ line=dict(
162
+ color=(to_hex(fit_color) if single
163
+ else to_hex(fills[i])),
164
+ width=fit_lwd),
165
+ hovertemplate=(
166
+ f"{x_lab}: {xpart}<br>{lbl}: {ypart}"
167
+ "<extra></extra>"),
168
+ showlegend=True,
169
+ ))
170
+
171
+ # axes, grids, legend, background
172
+ if is_date:
173
+ ax_x = axis_base()
174
+ ax_x["title"]["text"] = x_lab
175
+ ax_x.update(showgrid=True,
176
+ gridcolor=to_hex(style_opts["grid_col"]),
177
+ gridwidth=1)
178
+ shapes = y_grid(gridT2) + plot_border()
179
+ else:
180
+ ax_x = axis_num(x_lab, ax["axT1"], ax["axL1"])
181
+ shapes = x_grid(gridT1) + y_grid(gridT2) + plot_border()
182
+ ax_y = axis_num(y_lab, ax["axT2"], ax["axL2"])
183
+
184
+ leg_color = to_hex(style_opts["lab_color"])
185
+ fig.update_layout(
186
+ xaxis=ax_x, yaxis=ax_y,
187
+ shapes=shapes,
188
+ template=None,
189
+ legend=dict(
190
+ title=dict(
191
+ text=by_name if has_groups else None,
192
+ font=dict(
193
+ size=round(15 * get_option("lab_size", 1)),
194
+ family="Arial", color=leg_color)),
195
+ font=dict(
196
+ size=round(16 * get_option("axis_size", 0.9)),
197
+ family="Arial", color=leg_color),
198
+ orientation="v", x=1.05, y=0.5,
199
+ xanchor="left", yanchor="middle",
200
+ bgcolor=to_hex(style_opts["window_fill"]),
201
+ bordercolor="#CCCCCC", borderwidth=1,
202
+ itemsizing="constant",
203
+ ),
204
+ plot_bgcolor=to_hex(style_opts["panel_fill"]),
205
+ paper_bgcolor=to_hex(style_opts["window_fill"]),
206
+ )
207
+
208
+ if main:
209
+ title_size = round(16 * get_option("main_size", 1))
210
+ fig.update_layout(
211
+ title=dict(text=main, x=0.5, xanchor="center",
212
+ y=0.99, yanchor="top",
213
+ font=dict(size=title_size)),
214
+ margin=dict(t=round(title_size * 2.2)),
215
+ )
216
+ return fig
lessPy/plt_smooth.py ADDED
@@ -0,0 +1,170 @@
1
+ # plt_smooth.py — analog of the form="smooth" branch of
2
+ # plt.main.R (~795-808), which delegates to smoothScatter()
3
+ #
4
+ # form="smooth" for XY(): the 2-D binned kernel density of
5
+ # <x, y> as a smoothed color raster, densities transformed by
6
+ # z^smooth_power, with the smooth_points lowest-density points
7
+ # overplotted. _bkde2d() ports KernSmooth::bkde2D() directly
8
+ # (linear binning, separable normal kernel) with smoothScatter's
9
+ # bandwidth, (q95 - q05)/25 per axis, so the density grid
10
+ # matches R. Renderer conventions follow plt_plotly.py: fit
11
+ # lines, SE bands, and data ellipses arrive as prepared
12
+ # coordinates computed by XY(). Deviation: the fill ramp is
13
+ # fixed at the default-theme clr.den, colorRampPalette(
14
+ # c(window_fill, hcl(240, 80, 16))) — lessPy does not port
15
+ # themes.
16
+
17
+ import numpy as np
18
+ import plotly.graph_objects as go
19
+ from scipy import stats as sps
20
+ from scipy.signal import fftconvolve
21
+
22
+ from .plotly_utils import (
23
+ as_plotly_color, axis_num, get_tick_fmt, plot_border,
24
+ plotly_style, to_hex, x_grid, y_grid)
25
+ from .utils import get_option
26
+
27
+ # colorRampPalette(c("white", hcl(240, 80, 16))): the density
28
+ # ramp of the default lessR theme (plt.main.R ~799)
29
+ _RAMP = [[0, "#FFFFFF"], [1, "#0041A5"]]
30
+
31
+
32
+ def _linbin2d(x, y, gx, gy):
33
+ """Linear binning: each point's weight split among its four
34
+ surrounding grid nodes. R analog: KernSmooth's linbin2D()"""
35
+ m1, m2 = len(gx), len(gy)
36
+ lx = (x - gx[0]) / (gx[1] - gx[0])
37
+ ly = (y - gy[0]) / (gy[1] - gy[0])
38
+ i, j = np.floor(lx).astype(int), np.floor(ly).astype(int)
39
+ rx, ry = lx - i, ly - j
40
+ g = np.zeros((m1, m2))
41
+ for di, wx in ((0, 1 - rx), (1, rx)):
42
+ for dj, wy in ((0, 1 - ry), (1, ry)):
43
+ ii, jj = i + di, j + dj
44
+ ok = (ii >= 0) & (ii < m1) & (jj >= 0) & (jj < m2)
45
+ np.add.at(g, (ii[ok], jj[ok]), (wx * wy)[ok])
46
+ return g
47
+
48
+
49
+ def _bkde2d(x, y, h, n_grid, lims):
50
+ """Binned 2-D kernel density estimate on an n_grid x n_grid
51
+ grid over lims=(x_lo, x_hi, y_lo, y_hi), normal kernel with
52
+ sd h=(hx, hy) truncated at tau=3.4 sd.
53
+ R analog: KernSmooth::bkde2D()"""
54
+ tau = 3.4
55
+ gx = np.linspace(lims[0], lims[1], n_grid)
56
+ gy = np.linspace(lims[2], lims[3], n_grid)
57
+ g = _linbin2d(x, y, gx, gy)
58
+ kern = []
59
+ for (a, b), hh in (((lims[0], lims[1]), h[0]),
60
+ ((lims[2], lims[3]), h[1])):
61
+ L = min(int(tau * hh * (n_grid - 1) / (b - a)),
62
+ n_grid - 1)
63
+ fac = (b - a) / (hh * (n_grid - 1))
64
+ half = sps.norm.pdf(np.arange(L + 1) * fac) / hh
65
+ full = np.concatenate([half[:0:-1], half])
66
+ kern.append(full / (full.sum() * fac * hh))
67
+ z = fftconvolve(g, np.outer(kern[0], kern[1]),
68
+ mode="same") / len(x)
69
+ return gx, gy, np.clip(z, 0, None)
70
+
71
+
72
+ def plt_smooth(xv, yv, smooth_points, smooth_size, smooth_power,
73
+ smooth_bins, x_lab, y_lab, main, digits_d,
74
+ ax, gridT1, gridT2, x_lim, y_lim,
75
+ fit_lines, fit_color, fit_lwd, se_polys, se_fill,
76
+ ellipses, ellipse_fill, ellipse_color,
77
+ ellipse_lwd):
78
+ """Smoothed-density display of the scatter of x and y with
79
+ the lowest-density points overplotted, plus the standard
80
+ scatter overlays. R analog: smoothScatter() via plt.main.R"""
81
+
82
+ # smoothScatter bandwidth and grid: data range extended by
83
+ # 1.5 * bandwidth per axis (the bkde2D default range.x)
84
+ h = []
85
+ for v in (xv, yv):
86
+ q05, q95 = np.percentile(v, [5, 95])
87
+ hh = (q95 - q05) / 25
88
+ h.append(hh if hh > 0 else 1.0)
89
+ gx, gy, z = _bkde2d(xv, yv, h, smooth_bins,
90
+ (xv.min() - 1.5 * h[0],
91
+ xv.max() + 1.5 * h[0],
92
+ yv.min() - 1.5 * h[1],
93
+ yv.max() + 1.5 * h[1]))
94
+ dens = z ** smooth_power # compress the peaks
95
+
96
+ style_opts = plotly_style()
97
+ fig = go.Figure()
98
+
99
+ fmtx = get_tick_fmt(xv, digits_d)
100
+ fmty = get_tick_fmt(yv, digits_d)
101
+ xpart = f"%{{x:{fmtx}}}" if fmtx else "%{x}"
102
+ ypart = f"%{{y:{fmty}}}" if fmty else "%{y}"
103
+ fig.add_trace(go.Heatmap(
104
+ x=gx, y=gy, z=dens.T,
105
+ colorscale=_RAMP, zsmooth="best", showscale=False,
106
+ hovertemplate=(f"{x_lab}: {xpart}<br>{y_lab}: {ypart}"
107
+ "<extra></extra>"),
108
+ ))
109
+
110
+ # the smooth_points points in the lowest-density grid cells
111
+ # overplot the raster (smoothScatter nrpoints selection)
112
+ n_pts = min(len(xv), int(np.ceil(smooth_points)))
113
+ if n_pts > 0:
114
+ ix = ((len(gx) - 1) * (xv - gx[0])
115
+ / (gx[-1] - gx[0])).astype(int)
116
+ iy = ((len(gy) - 1) * (yv - gy[0])
117
+ / (gy[-1] - gy[0])).astype(int)
118
+ sel = np.argsort(dens[ix, iy], kind="stable")[:n_pts]
119
+ px = 2 * float(smooth_size) # pch="." speck
120
+ if not np.isfinite(px) or px <= 0:
121
+ px = 2
122
+ fig.add_trace(go.Scatter(
123
+ x=xv[sel], y=yv[sel], mode="markers",
124
+ marker=dict(symbol="circle", size=px,
125
+ sizemode="diameter", color="#000000"),
126
+ hoverinfo="skip", showlegend=False))
127
+
128
+ # standard scatter overlays, as in plt_plotly
129
+ for ex, ey in ellipses or []:
130
+ fig.add_trace(go.Scatter(
131
+ x=ex, y=ey, mode="lines",
132
+ line=dict(color=to_hex(ellipse_color),
133
+ width=ellipse_lwd),
134
+ fill="toself",
135
+ fillcolor=as_plotly_color(ellipse_fill),
136
+ hoverinfo="skip", showlegend=False))
137
+ for bx, bnd in se_polys or []:
138
+ fig.add_trace(go.Scatter(
139
+ x=bx, y=bnd, mode="none", fill="toself",
140
+ fillcolor=as_plotly_color(se_fill),
141
+ hoverinfo="skip", showlegend=False))
142
+ for fl in fit_lines or []:
143
+ if len(fl["x"]) < 2:
144
+ continue
145
+ fig.add_trace(go.Scatter(
146
+ x=fl["x"], y=fl["y"], mode="lines", name="Fit",
147
+ legendgroup="fit",
148
+ line=dict(color=to_hex(fit_color), width=fit_lwd),
149
+ hoverinfo="skip", showlegend=True))
150
+
151
+ ax_x = axis_num(x_lab, ax["axT1"], ax["axL1"])
152
+ ax_y = axis_num(y_lab, ax["axT2"], ax["axL2"])
153
+ ax_x["range"] = x_lim
154
+ ax_y["range"] = y_lim
155
+ fig.update_layout(
156
+ xaxis=ax_x, yaxis=ax_y,
157
+ shapes=x_grid(gridT1) + y_grid(gridT2) + plot_border(),
158
+ template=None,
159
+ plot_bgcolor=to_hex(style_opts["panel_fill"]),
160
+ paper_bgcolor=to_hex(style_opts["window_fill"]),
161
+ )
162
+
163
+ if main:
164
+ title_size = round(16 * get_option("main_size", 1))
165
+ fig.update_layout(
166
+ title=dict(text=main, x=0.5, xanchor="center",
167
+ y=0.99, yanchor="top",
168
+ font=dict(size=title_size)),
169
+ margin=dict(t=round(title_size * 2.2)))
170
+ return fig