plot3 0.4.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.
plot3/io.py ADDED
@@ -0,0 +1,76 @@
1
+ """Point-cloud IO and optional CRAFT SSH helpers."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import shlex
6
+ import subprocess
7
+
8
+ import numpy as np
9
+ import pandas as pd
10
+
11
+ try:
12
+ from IPython import get_ipython
13
+ except Exception: # pragma: no cover
14
+ get_ipython = None
15
+
16
+ def ssh_cfg():
17
+ """CRAFT's SSH config from the IPython user namespace (optional)."""
18
+ nss = [globals()]
19
+ if get_ipython is not None:
20
+ try:
21
+ nss.append(get_ipython().user_ns)
22
+ except Exception:
23
+ pass
24
+ for ns in nss:
25
+ if isinstance(ns, dict) and ns.get("SSH_HOST"):
26
+ return ns["SSH_HOST"], ns.get("SSH_OPTS", "")
27
+ raise RuntimeError("SSH_HOST not found — load CRAFT and run %gpu first.")
28
+
29
+
30
+ def ssh_bytes(remote_cmd: str) -> bytes:
31
+ host, opts = ssh_cfg()
32
+ cmd = ["ssh", *shlex.split(opts or ""), host, remote_cmd]
33
+ proc = subprocess.run(cmd, capture_output=True, check=False)
34
+ if proc.returncode != 0:
35
+ err = (proc.stderr or b"").decode("utf-8", "replace").strip()
36
+ raise RuntimeError(f"ssh failed (rc={proc.returncode}): {err[:300]}")
37
+ return proc.stdout or b""
38
+
39
+
40
+ def read_bin(path, *, stride=5, columns=None, remote=False, sub=1,
41
+ max_points=500_000) -> pd.DataFrame:
42
+ """Load an (N, stride) float32 point-cloud file as a DataFrame.
43
+
44
+ columns: names for the stride columns (default x,y,z,intensity,c4,...).
45
+ remote=True streams (and thins on) the CRAFT GPU host over SSH.
46
+ """
47
+ stride = max(3, int(stride))
48
+ sub = max(1, int(sub or 1))
49
+ if remote:
50
+ thin = (
51
+ "python3 -c " + shlex.quote(
52
+ "import sys,numpy as np;"
53
+ f"a=np.fromfile({str(path)!r},dtype=np.float32);"
54
+ f"s={stride};n=a.size//s;a=a[:n*s].reshape(n,s);"
55
+ f"m={int(max_points)};sb={sub};"
56
+ "sb=max(sb,(n+m-1)//m) if m>0 else sb;"
57
+ "sys.stdout.buffer.write("
58
+ "np.ascontiguousarray(a[::sb],dtype=np.float32).tobytes())"
59
+ )
60
+ )
61
+ raw = ssh_bytes(thin)
62
+ arr = np.frombuffer(raw, dtype=np.float32)
63
+ else:
64
+ arr = np.fromfile(str(path), dtype=np.float32)
65
+ n = arr.size // stride
66
+ arr = arr[: n * stride].reshape(n, stride)
67
+ if not remote:
68
+ if max_points and n > max_points * sub:
69
+ sub = max(sub, (n + max_points - 1) // max_points)
70
+ if sub > 1:
71
+ arr = np.ascontiguousarray(arr[::sub])
72
+ if columns is None:
73
+ base = ["x", "y", "z", "intensity"]
74
+ columns = (base + [f"c{i}" for i in range(4, stride)])[:stride]
75
+ return pd.DataFrame(arr, columns=list(columns)[:stride])
76
+
plot3/jupyter.py ADDED
@@ -0,0 +1,514 @@
1
+ """IPython/SolveIt integration: hide-from-AI, %plot3, R-style aes, extension load."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import shlex
7
+ from typing import Any
8
+
9
+ import numpy as np
10
+ import pandas as pd
11
+
12
+ try:
13
+ from IPython import get_ipython
14
+ except Exception: # pragma: no cover
15
+ get_ipython = None
16
+
17
+ from plot3.io import ssh_bytes
18
+ from plot3.masking import (
19
+ BT_NAME,
20
+ Plot3MaskTransformer,
21
+ is_backtick_transformer,
22
+ is_mask_transformer,
23
+ is_tidy3_backtick_transformer,
24
+ plot3_backtick_transform,
25
+ )
26
+
27
+ # R-style bare-name / backtick masking for aes() / facet_wrap().
28
+ _R_STYLE_ON = True
29
+
30
+
31
+ def _bt_fallback(name: str) -> str:
32
+ """If a backtick sentinel escapes AST rewrite, treat it as a column name."""
33
+ return str(name)
34
+
35
+
36
+ def enable_r_style(ipython: Any | None = None) -> bool:
37
+ """Enable bare-name + backtick masking for ``aes`` / ``facet_wrap``.
38
+
39
+ Jupyter / SolveIt only. After this::
40
+
41
+ aes(x=wt, y=mpg, colour=cyl)
42
+ aes(x=`First Name`, y=`Age (%)`)
43
+ facet_wrap(cyl)
44
+
45
+ become string column names, matching ggplot2 / tidy3 conventions.
46
+ """
47
+ global _R_STYLE_ON
48
+ _R_STYLE_ON = True
49
+ if ipython is None:
50
+ try:
51
+ ipython = get_ipython() if get_ipython is not None else None
52
+ except Exception:
53
+ ipython = None
54
+ if ipython is None:
55
+ return False
56
+
57
+ ns = getattr(ipython, "user_ns", None)
58
+ if ns is not None:
59
+ # Escape hatch if AST missed a sentinel (still a valid column string).
60
+ ns[BT_NAME] = _bt_fallback
61
+ # When tidy3 is not loaded, leave tidy3 sentinel unset; plot3 AST
62
+ # still rewrites __tidy3_bt__ if tidy3's preparser ran first.
63
+
64
+ # Source preparser: only install plot3's if tidy3 did not already provide one.
65
+ for attr in ("input_transformers_cleanup", "input_transformers_post"):
66
+ transformers = getattr(ipython, attr, None)
67
+ if not isinstance(transformers, list):
68
+ continue
69
+ has_tidy3_bt = any(is_tidy3_backtick_transformer(t) for t in transformers)
70
+ transformers[:] = [
71
+ t for t in transformers if not is_backtick_transformer(t)
72
+ ]
73
+ if not has_tidy3_bt:
74
+ transformers.insert(0, plot3_backtick_transform)
75
+
76
+ # AST pass for bare names inside aes / facet_wrap.
77
+ ast_transformers = getattr(ipython, "ast_transformers", None)
78
+ if isinstance(ast_transformers, list):
79
+ ast_transformers[:] = [
80
+ t for t in ast_transformers if not is_mask_transformer(t)
81
+ ]
82
+ ast_transformers.append(Plot3MaskTransformer())
83
+ return True
84
+
85
+
86
+ def disable_r_style(ipython: Any | None = None) -> None:
87
+ """Disable bare-name / backtick masking for plot3 aesthetics."""
88
+ global _R_STYLE_ON
89
+ _R_STYLE_ON = False
90
+ if ipython is None:
91
+ try:
92
+ ipython = get_ipython() if get_ipython is not None else None
93
+ except Exception:
94
+ ipython = None
95
+ if ipython is None:
96
+ return
97
+ for attr in ("input_transformers_cleanup", "input_transformers_post"):
98
+ transformers = getattr(ipython, attr, None)
99
+ if isinstance(transformers, list):
100
+ transformers[:] = [
101
+ t for t in transformers if not is_backtick_transformer(t)
102
+ ]
103
+ ast_transformers = getattr(ipython, "ast_transformers", None)
104
+ if isinstance(ast_transformers, list):
105
+ ast_transformers[:] = [
106
+ t for t in ast_transformers if not is_mask_transformer(t)
107
+ ]
108
+
109
+
110
+ def r_style_enabled() -> bool:
111
+ return _R_STYLE_ON
112
+
113
+ def find_caller_msg_id():
114
+ import inspect
115
+
116
+ frame = inspect.currentframe()
117
+ try:
118
+ f = frame.f_back if frame is not None else None
119
+ while f is not None:
120
+ for ns in (f.f_locals, f.f_globals):
121
+ mid = ns.get("__msg_id") if isinstance(ns, dict) else None
122
+ if mid:
123
+ return str(mid)
124
+ f = f.f_back
125
+ finally:
126
+ del frame
127
+ try:
128
+ ip = get_ipython()
129
+ for ns_name in ("user_ns", "user_global_ns"):
130
+ ns = getattr(ip, ns_name, None) or {}
131
+ mid = ns.get("__msg_id") if isinstance(ns, dict) else None
132
+ if mid:
133
+ return str(mid)
134
+ except Exception:
135
+ pass
136
+ try:
137
+ from safepyrun import find_var # type: ignore
138
+
139
+ mid = find_var("__msg_id")
140
+ if mid:
141
+ return str(mid)
142
+ except Exception:
143
+ pass
144
+ return None
145
+
146
+
147
+ def hide_caller_from_ai(mid=None):
148
+ """Best-effort ``skipped=1`` on the calling cell; no-op outside SolveIt."""
149
+ try:
150
+ from dialoghelper.core import update_msg
151
+ except Exception:
152
+ return
153
+
154
+ async def _run():
155
+ import inspect
156
+
157
+ m = mid or find_caller_msg_id()
158
+ if not m:
159
+ try:
160
+ from dialoghelper.core import read_msg
161
+
162
+ msg = read_msg(n=0, relative=True)
163
+ if inspect.iscoroutine(msg):
164
+ msg = await msg
165
+ m = msg.get("id") if isinstance(msg, dict) else getattr(msg, "id", None)
166
+ except Exception:
167
+ m = None
168
+ if not m:
169
+ try:
170
+ from dialoghelper.core import find_msgs
171
+
172
+ msgs = find_msgs(msg_type="code", re_pattern=r"%plot3",
173
+ include_output=False, include_meta=True,
174
+ include_skipped=True, use_regex=True)
175
+ if inspect.iscoroutine(msgs):
176
+ msgs = await msgs
177
+ for msg in msgs or []:
178
+ m = msg.get("id") if isinstance(msg, dict) else getattr(msg, "id", None)
179
+ except Exception:
180
+ m = None
181
+ if not m:
182
+ print("plot3: hide-from-ai failed — could not resolve msg id "
183
+ "(pass hide=0 to silence)")
184
+ return
185
+ m = str(m)
186
+ err = None
187
+ for cand in (m, m[1:] if m.startswith("_") else "_" + m):
188
+ try:
189
+ res = update_msg(id=cand, skipped=1)
190
+ if inspect.iscoroutine(res):
191
+ await res
192
+ return
193
+ except Exception as e:
194
+ err = e
195
+ print(f"plot3: hide-from-ai failed — update_msg({m}): {err}")
196
+
197
+ try:
198
+ import asyncio
199
+
200
+ try:
201
+ loop = asyncio.get_running_loop()
202
+ except RuntimeError:
203
+ loop = None
204
+ if loop is None:
205
+ asyncio.run(_run())
206
+ return
207
+ try:
208
+ import nest_asyncio
209
+
210
+ nest_asyncio.apply()
211
+ loop.run_until_complete(_run())
212
+ except Exception:
213
+ import concurrent.futures
214
+
215
+ with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
216
+ pool.submit(lambda: asyncio.run(_run())).result()
217
+ except Exception as e:
218
+ print(f"plot3: hide-from-ai failed — {e}")
219
+
220
+
221
+ # ═════════════════════════════════════════════════════════════════════════════
222
+ # %plot3 magic — remote (CRAFT) or local DataFrame snapshot -> figure
223
+ # ═════════════════════════════════════════════════════════════════════════════
224
+
225
+
226
+ def remote_df(expr: str, cols: dict, max_points: int) -> pd.DataFrame:
227
+ """Snapshot mapped columns of a remote DataFrame over CRAFT's SSH pipe.
228
+
229
+ Legacy host path — prefer :func:`remote_plot3_payload` so stats run on the
230
+ GPU and only a PlotPayload crosses the wire.
231
+ """
232
+ import io
233
+ import uuid
234
+
235
+ ip = get_ipython()
236
+ rr = (ip.user_ns or {}).get("remote_run_") if ip is not None else None
237
+ if not callable(rr):
238
+ raise RuntimeError("remote_run_ missing — load CRAFT and run %gpu")
239
+ tmp = f"/tmp/plot3_{uuid.uuid4().hex}.npz"
240
+ code = f"""
241
+ import numpy as _np, pandas as _pd, json as _json
242
+ _df = eval({expr!r})
243
+ _m = {int(max_points)}
244
+ if _m > 0 and len(_df) > _m:
245
+ _df = _df.iloc[::(len(_df) + _m - 1) // _m]
246
+ _out, _meta = {{}}, {{}}
247
+ for _a, _c in {cols!r}.items():
248
+ _s = _df[_c]
249
+ if str(_s.dtype).startswith("datetime"):
250
+ _out[_a] = _s.astype("datetime64[ns]").astype("int64").to_numpy()
251
+ _meta[_a] = ["dt", _c]
252
+ elif _s.dtype.kind in "ifub":
253
+ _out[_a] = _s.to_numpy(_np.float32); _meta[_a] = ["num", _c]
254
+ else:
255
+ _codes, _cats = _pd.factorize(_s.astype(str))
256
+ _out[_a] = _codes.astype(_np.int32)
257
+ _meta[_a] = ["cat", _c, [str(x) for x in _cats]]
258
+ _np.savez({tmp!r}, **_out)
259
+ print(_json.dumps(_meta))
260
+ """
261
+ out = rr(code, max_chars=8000).strip()
262
+ meta = json.loads(out.splitlines()[-1])
263
+ try:
264
+ raw = ssh_bytes("cat -- " + shlex.quote(tmp))
265
+ finally:
266
+ try:
267
+ ssh_bytes("rm -f -- " + shlex.quote(tmp))
268
+ except Exception:
269
+ pass
270
+ z = np.load(io.BytesIO(raw))
271
+ data = {}
272
+ for a, m in meta.items():
273
+ v = z[a]
274
+ col = m[1]
275
+ if m[0] == "dt":
276
+ data[col] = pd.to_datetime(v.astype("int64"), unit="ns")
277
+ elif m[0] == "cat":
278
+ data[col] = pd.Categorical.from_codes(
279
+ np.clip(v, 0, len(m[2]) - 1), categories=m[2])
280
+ else:
281
+ data[col] = v
282
+ return pd.DataFrame(data)
283
+
284
+
285
+ def remote_plot3_payload(
286
+ expr: str,
287
+ mapping: dict,
288
+ *,
289
+ kind: str = "point",
290
+ size: float | None = None,
291
+ max_points: int = 200_000,
292
+ theme: str = "dark",
293
+ height: str = "480px",
294
+ ) -> dict:
295
+ """Build a ggplot on the remote kernel and return a PlotPayload.
296
+
297
+ Stats / encoding run where the data lives; only the compact payload is
298
+ pulled to the host (via :func:`plot3.remote.fetch_remote_payload`).
299
+ """
300
+ from plot3.remote import fetch_remote_payload
301
+
302
+ # Remote source: load data, optional stride, assemble figure, to_payload.
303
+ size_lit = "None" if size is None else repr(float(size))
304
+ source = f"""
305
+ from plot3 import (
306
+ ggplot, aes, geom_point, geom_line, geom_path, theme_light, theme_dark,
307
+ )
308
+ import numpy as _np
309
+
310
+ _df = eval({expr!r})
311
+ _m = {int(max_points)}
312
+ # Stride large tables before encode (coord_3d max_points is separate).
313
+ try:
314
+ _n = len(_df)
315
+ except Exception:
316
+ _n = getattr(_df, "height", None) or getattr(getattr(_df, "shape", None), "__getitem__", lambda i: None)(0)
317
+ if _m and _n and int(_n) > _m:
318
+ try:
319
+ _step = max(1, (int(_n) + _m - 1) // _m)
320
+ _df = _df.iloc[::_step]
321
+ except Exception:
322
+ try:
323
+ import polars as _pl
324
+ if isinstance(_df, _pl.DataFrame):
325
+ _df = _df.with_row_index("_i").filter((_pl.col("_i") % _step) == 0).drop("_i")
326
+ except Exception:
327
+ pass
328
+
329
+ _aes = aes(**{mapping!r})
330
+ _fig = ggplot(_df, _aes, height={height!r}, hide=False)
331
+ _kind = {kind!r}
332
+ _size = {size_lit}
333
+ for _part in _kind.split("+"):
334
+ _part = _part.strip()
335
+ if _part == "point":
336
+ _fig = _fig + (geom_point(size=_size) if _size is not None else geom_point())
337
+ elif _part == "line":
338
+ _fig = _fig + geom_line()
339
+ elif _part == "path":
340
+ _fig = _fig + geom_path()
341
+ else:
342
+ raise ValueError(f"unknown kind {{_part!r}}")
343
+ if {theme!r} != "dark":
344
+ _fig = _fig + theme_light()
345
+ _plot3_payload = _fig.to_payload()
346
+ """
347
+ return fetch_remote_payload(source)
348
+
349
+
350
+ def run_plot3_from_magic(line: str = ""):
351
+ parts = shlex.split(line or "")
352
+ if not parts:
353
+ raise ValueError(
354
+ "%plot3 is an optional host helper; prefer ggplot under %gpu:\n"
355
+ " ggplot(df, aes(x=wt, y=mpg)) + geom_point()\n"
356
+ "usage: %plot3 <df_expr> x=col y=col [z=col] [color=col] "
357
+ "[group=col] [kind=point|line|path|point+line] [size=F] "
358
+ "[max_points=N] [theme=dark|light] [height=Npx] [hide=0|1]"
359
+ )
360
+ expr = parts[0]
361
+ m: dict = {}
362
+ kind, size, hide, theme = "point", None, True, "dark"
363
+ max_points, height = 200_000, "480px"
364
+ for tok in parts[1:]:
365
+ k, _, v = tok.partition("=")
366
+ if k in ("x", "y", "z", "color", "colour", "group"):
367
+ m["color" if k == "colour" else k] = v
368
+ elif k == "kind":
369
+ kind = v
370
+ elif k == "size":
371
+ size = float(v)
372
+ elif k == "max_points":
373
+ max_points = int(v)
374
+ elif k == "theme":
375
+ theme = v
376
+ elif k == "height":
377
+ height = v if v.endswith("px") else f"{int(v)}px"
378
+ elif k == "hide":
379
+ hide = v.lower() in ("1", "true", "yes")
380
+ else:
381
+ raise ValueError(f"unknown option {tok!r}")
382
+ if "x" not in m or "y" not in m:
383
+ raise ValueError("%plot3 needs x= and y=")
384
+
385
+ mid = find_caller_msg_id() if hide else None
386
+
387
+ ip = get_ipython() if get_ipython is not None else None
388
+ ns = (ip.user_ns or {}) if ip is not None else {}
389
+
390
+ from plot3.geoms import aes, geom_line, geom_path, geom_point, theme_light
391
+ from plot3.ggplot import ggplot
392
+
393
+ if callable(ns.get("remote_run_")):
394
+ # Phase D: stats/encode on remote; host only renders the payload.
395
+ try:
396
+ payload = remote_plot3_payload(
397
+ expr,
398
+ m,
399
+ kind=kind,
400
+ size=size,
401
+ max_points=max_points,
402
+ theme=theme,
403
+ height=height,
404
+ )
405
+ fig = ggplot.from_payload(payload, height=height, hide=False)
406
+ except Exception as e:
407
+ # Fall back to legacy column pull if remote payload path fails.
408
+ print(f"plot3: remote payload path failed ({e}); falling back to remote_df", flush=True)
409
+ df = remote_df(expr, m, max_points)
410
+ fig = ggplot(df, aes(**m), height=height, hide=False)
411
+ for part in kind.split("+"):
412
+ part = part.strip()
413
+ if part == "point":
414
+ fig = fig + (geom_point(size=size) if size else geom_point())
415
+ elif part == "line":
416
+ fig = fig + geom_line()
417
+ elif part == "path":
418
+ fig = fig + geom_path()
419
+ else:
420
+ raise ValueError(f"unknown kind {part!r}")
421
+ if theme != "dark":
422
+ fig = fig + theme_light()
423
+ else:
424
+ df = eval(expr, ns) # local fallback (plain Jupyter)
425
+ if not isinstance(df, pd.DataFrame):
426
+ try:
427
+ from plot3.table import as_table
428
+
429
+ df = as_table(df)
430
+ except Exception:
431
+ df = pd.DataFrame(df)
432
+ if max_points and hasattr(df, "__len__") and len(df) > max_points:
433
+ try:
434
+ df = df.iloc[:: (len(df) + max_points - 1) // max_points]
435
+ except Exception:
436
+ pass
437
+ fig = ggplot(df, aes(**m), height=height, hide=False)
438
+ for part in kind.split("+"):
439
+ part = part.strip()
440
+ if part == "point":
441
+ fig = fig + (geom_point(size=size) if size else geom_point())
442
+ elif part == "line":
443
+ fig = fig + geom_line()
444
+ elif part == "path":
445
+ fig = fig + geom_path()
446
+ else:
447
+ raise ValueError(f"unknown kind {part!r}")
448
+ if theme != "dark":
449
+ fig = fig + theme_light()
450
+
451
+ # Prefer fig display path (opens system browser under VS Code; iframe in SolveIt).
452
+ if hide:
453
+ fig.hide = False
454
+ fig.show()
455
+ hide_caller_from_ai(mid)
456
+ else:
457
+ fig.show()
458
+ return None
459
+
460
+
461
+ # ═════════════════════════════════════════════════════════════════════════════
462
+ # Registration (addon contract: register everything via get_ipython)
463
+ # ═════════════════════════════════════════════════════════════════════════════
464
+
465
+
466
+ def register_plot3(*, quiet=True, r_style: bool = True) -> bool:
467
+ if get_ipython is None:
468
+ return False
469
+ ip = get_ipython()
470
+ if ip is None:
471
+ return False
472
+ ok = False
473
+ try:
474
+ ip.magics_manager.register_function(
475
+ run_plot3_from_magic, magic_kind="line", magic_name="plot3")
476
+ ok = True
477
+ except Exception as e:
478
+ if not quiet:
479
+ print(f"plot3: magic registration failed: {e}")
480
+ # host-local under %gpu (CRAFT hook, when present)
481
+ try:
482
+ reg = (ip.user_ns or {}).get("register_local_magic")
483
+ if callable(reg):
484
+ reg("%plot3")
485
+ except Exception:
486
+ pass
487
+ # public API into user_ns (never rely on %run leaking module globals)
488
+ try:
489
+ import plot3 as plot3_pkg
490
+
491
+ ns = ip.user_ns
492
+ for name in plot3_pkg.__all__:
493
+ if name == "load_ipython_extension":
494
+ continue
495
+ ns[name] = getattr(plot3_pkg, name)
496
+ except Exception:
497
+ pass
498
+ # R-style bare names / backticks in aes() (ggplot2 parity with tidy3)
499
+ if r_style:
500
+ try:
501
+ enable_r_style(ip)
502
+ except Exception as e:
503
+ if not quiet:
504
+ print(f"plot3: R-style masking not enabled: {e}")
505
+ if ok and not quiet:
506
+ print("plot3 ready")
507
+ print(" ggplot(df, aes(x=wt, y=mpg[, colour=cyl])) + geom_point()")
508
+ print(" aes(x=`First Name`, y=mpg) # backticks for spaced names")
509
+ print(" %plot3 df x=a y=b [z=c] [color=d] read_bin(path) ggsave()")
510
+ return ok
511
+
512
+
513
+ def load_ipython_extension(ip=None):
514
+ register_plot3(quiet=True)